mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-31 16:21:50 +03:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3be12cf134 | ||
|
|
1f50570600 | ||
|
|
4f6cf1fac2 | ||
|
|
28500fffd9 | ||
|
|
33e6aa329b | ||
|
|
c11ddbe491 | ||
|
|
2020645f18 | ||
|
|
e9d811e609 | ||
|
|
20d7defa03 |
@@ -188,7 +188,7 @@ These variables are process-level switches. Set them in the same terminal, servi
|
|||||||
|
|
||||||
| Variable | Default | Description |
|
| Variable | Default | Description |
|
||||||
|----------|---------|-------------|
|
|----------|---------|-------------|
|
||||||
| `NANOBOT_MAX_CONCURRENT_REQUESTS` | Unlimited | Maximum concurrently running inbound agent requests. Set a positive integer to apply a cap; unset, `0`, or a negative value means unlimited. |
|
| `NANOBOT_MAX_CONCURRENT_REQUESTS` | `3` | Maximum concurrently running inbound agent requests. Must be an integer; set `0` or a negative value for unlimited. |
|
||||||
| `NANOBOT_LLM_TIMEOUT_S` | `300` | Wall-clock timeout, in seconds. Ordinary requests use this value; streaming requests use the greater of 300 seconds or twice this value. Set `0` to disable. Sustained-goal turns bypass this wall-clock cap. |
|
| `NANOBOT_LLM_TIMEOUT_S` | `300` | Wall-clock timeout, in seconds. Ordinary requests use this value; streaming requests use the greater of 300 seconds or twice this value. Set `0` to disable. Sustained-goal turns bypass this wall-clock cap. |
|
||||||
| `NANOBOT_STREAM_IDLE_TIMEOUT_S` | `90` | Streaming idle timeout, in seconds, used by streaming providers. Invalid or non-positive values are ignored; values above `3600` are clamped. |
|
| `NANOBOT_STREAM_IDLE_TIMEOUT_S` | `90` | Streaming idle timeout, in seconds, used by streaming providers. Invalid or non-positive values are ignored; values above `3600` are clamped. |
|
||||||
| `NANOBOT_OPENAI_COMPAT_TIMEOUT_S` | `120` | HTTP request timeout, in seconds, for OpenAI-compatible providers. Invalid or non-positive values are ignored. |
|
| `NANOBOT_OPENAI_COMPAT_TIMEOUT_S` | `120` | HTTP request timeout, in seconds, for OpenAI-compatible providers. Invalid or non-positive values are ignored. |
|
||||||
@@ -2213,7 +2213,7 @@ The notification gate runs on a built-in system prompt. Advanced users can overr
|
|||||||
|
|
||||||
## Subagent Concurrency
|
## Subagent Concurrency
|
||||||
|
|
||||||
By default, nanobot allows four subagents to run at the same time. Additional subagents wait for capacity instead of being rejected. Lower the limit if a local model server cannot hold multiple KV caches, or raise it when the provider can handle more parallel work:
|
By default, nanobot only allows one spawned subagent at a time. When the limit is reached, the `spawn` tool returns an error so the agent can decide to wait or rearrange its work. This protects local LLM servers from loading multiple KV caches at once. If your provider can handle more parallel work, raise the limit:
|
||||||
|
|
||||||
```json
|
```json
|
||||||
{
|
{
|
||||||
@@ -2229,7 +2229,7 @@ The deprecated `agents.defaults.failOnToolError` field is silently ignored when
|
|||||||
|
|
||||||
| Option | Default | Description |
|
| Option | Default | Description |
|
||||||
|--------|---------|-------------|
|
|--------|---------|-------------|
|
||||||
| `agents.defaults.maxConcurrentSubagents` | `4` | Maximum number of subagents that may run at the same time. Additional tasks wait for capacity. |
|
| `agents.defaults.maxConcurrentSubagents` | `1` | Maximum number of spawned subagents that may run at the same time. Attempts to spawn beyond this limit return an error. |
|
||||||
|
|
||||||
|
|
||||||
## Auto Compact
|
## Auto Compact
|
||||||
|
|||||||
+3
-1
@@ -29,7 +29,9 @@ Memory moves through nanobot in two stages.
|
|||||||
|
|
||||||
### Stage 1: Consolidator
|
### Stage 1: Consolidator
|
||||||
|
|
||||||
When a conversation grows large, the `Consolidator` summarizes older turns and appends the result to `memory/history.jsonl`, while keeping recent conversation available. Each summary preserves useful long-term facts and a short handoff for active work.
|
When a conversation grows large enough to pressure the context window, nanobot does not try to carry every old message forever.
|
||||||
|
|
||||||
|
Instead, the `Consolidator` summarizes the oldest safe slice of the conversation and appends that summary to `memory/history.jsonl`.
|
||||||
|
|
||||||
This file is:
|
This file is:
|
||||||
|
|
||||||
|
|||||||
@@ -48,7 +48,7 @@ class AutoCompact:
|
|||||||
|
|
||||||
def _has_unarchived_messages(self, key: str) -> bool:
|
def _has_unarchived_messages(self, key: str) -> bool:
|
||||||
session = self.sessions.get_or_create(key)
|
session = self.sessions.get_or_create(key)
|
||||||
return session.last_archived < len(session.messages)
|
return session.last_consolidated < len(session.messages)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _is_internal_session(cls, key: str) -> bool:
|
def _is_internal_session(cls, key: str) -> bool:
|
||||||
|
|||||||
+30
-11
@@ -38,9 +38,10 @@ from nanobot.agent.subagent import SubagentManager
|
|||||||
from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context
|
from nanobot.agent.tools.context import RequestContext, bind_request_context, reset_request_context
|
||||||
from nanobot.agent.tools.exec_session import ExecSessionManager
|
from nanobot.agent.tools.exec_session import ExecSessionManager
|
||||||
from nanobot.agent.tools.file_state import FileStateStore, bind_file_states, reset_file_states
|
from nanobot.agent.tools.file_state import FileStateStore, bind_file_states, reset_file_states
|
||||||
from nanobot.agent.tools.message import capture_message_deliveries
|
from nanobot.agent.tools.message import MessageTool
|
||||||
from nanobot.agent.tools.registry import ToolRegistry
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.agent.tools.runtime_control import AgentRuntimeControl
|
from nanobot.agent.tools.runtime_control import AgentRuntimeControl
|
||||||
|
from nanobot.agent.tools.self import MyTool
|
||||||
from nanobot.agent.turn_delivery import (
|
from nanobot.agent.turn_delivery import (
|
||||||
TurnDelivery,
|
TurnDelivery,
|
||||||
TurnDeliveryFactory,
|
TurnDeliveryFactory,
|
||||||
@@ -144,6 +145,7 @@ class TurnContext:
|
|||||||
final_content: str | None = None
|
final_content: str | None = None
|
||||||
all_messages: list[dict[str, Any]] = field(default_factory=list)
|
all_messages: list[dict[str, Any]] = field(default_factory=list)
|
||||||
stop_reason: str = ""
|
stop_reason: str = ""
|
||||||
|
had_injections: bool = False
|
||||||
streamed_content: bool = False
|
streamed_content: bool = False
|
||||||
|
|
||||||
input_persisted_early: bool = False
|
input_persisted_early: bool = False
|
||||||
@@ -273,6 +275,7 @@ class AgentLoop:
|
|||||||
channels_config: ChannelsConfig | None = None,
|
channels_config: ChannelsConfig | None = None,
|
||||||
timezone: str | None = None,
|
timezone: str | None = None,
|
||||||
session_ttl_minutes: int = 0,
|
session_ttl_minutes: int = 0,
|
||||||
|
consolidation_ratio: float = 0.5,
|
||||||
hooks: list[AgentHook] | None = None,
|
hooks: list[AgentHook] | None = None,
|
||||||
hook_factories: list[AgentTurnHookFactory] | None = None,
|
hook_factories: list[AgentTurnHookFactory] | None = None,
|
||||||
unified_session: bool = False,
|
unified_session: bool = False,
|
||||||
@@ -429,8 +432,8 @@ class AgentLoop:
|
|||||||
("cron", self._cron_turns),
|
("cron", self._cron_turns),
|
||||||
("local trigger", self._local_trigger_turns),
|
("local trigger", self._local_trigger_turns),
|
||||||
)
|
)
|
||||||
# NANOBOT_MAX_CONCURRENT_REQUESTS: unset or <=0 means unlimited.
|
# NANOBOT_MAX_CONCURRENT_REQUESTS: <=0 means unlimited; default 3.
|
||||||
_max = int(os.environ.get("NANOBOT_MAX_CONCURRENT_REQUESTS", "0"))
|
_max = int(os.environ.get("NANOBOT_MAX_CONCURRENT_REQUESTS", "3"))
|
||||||
self._concurrency_gate: asyncio.Semaphore | None = (
|
self._concurrency_gate: asyncio.Semaphore | None = (
|
||||||
asyncio.Semaphore(_max) if _max > 0 else None
|
asyncio.Semaphore(_max) if _max > 0 else None
|
||||||
)
|
)
|
||||||
@@ -443,6 +446,7 @@ class AgentLoop:
|
|||||||
workspace_scopes=self.workspace_scopes,
|
workspace_scopes=self.workspace_scopes,
|
||||||
unified_session=unified_session,
|
unified_session=unified_session,
|
||||||
),
|
),
|
||||||
|
consolidation_ratio=consolidation_ratio,
|
||||||
unified_session=unified_session,
|
unified_session=unified_session,
|
||||||
)
|
)
|
||||||
self.auto_compact = AutoCompact(
|
self.auto_compact = AutoCompact(
|
||||||
@@ -515,6 +519,7 @@ class AgentLoop:
|
|||||||
disabled_skills=defaults.disabled_skills,
|
disabled_skills=defaults.disabled_skills,
|
||||||
session_ttl_minutes=defaults.session_ttl_minutes,
|
session_ttl_minutes=defaults.session_ttl_minutes,
|
||||||
idle_compact_check_interval_seconds=defaults.idle_compact_check_interval_seconds,
|
idle_compact_check_interval_seconds=defaults.idle_compact_check_interval_seconds,
|
||||||
|
consolidation_ratio=defaults.consolidation_ratio,
|
||||||
tools_config=config.tools,
|
tools_config=config.tools,
|
||||||
model_presets=preset_helpers.configured_model_presets(config),
|
model_presets=preset_helpers.configured_model_presets(config),
|
||||||
model_preset=defaults.model_preset,
|
model_preset=defaults.model_preset,
|
||||||
@@ -639,11 +644,20 @@ class AgentLoop:
|
|||||||
timezone=self.context.timezone or "UTC",
|
timezone=self.context.timezone or "UTC",
|
||||||
workspace_sandbox=self.workspace_scopes.sandbox_status,
|
workspace_sandbox=self.workspace_scopes.sandbox_status,
|
||||||
runtime_events=self.runtime_events,
|
runtime_events=self.runtime_events,
|
||||||
runtime_control=AgentRuntimeControl(self),
|
|
||||||
)
|
)
|
||||||
loader = ToolLoader()
|
loader = ToolLoader()
|
||||||
registered = loader.load(ctx, self.tools)
|
registered = loader.load(ctx, self.tools)
|
||||||
|
|
||||||
|
# MyTool receives only the explicit runtime-control capability.
|
||||||
|
if self.tools_config.my.enable:
|
||||||
|
self.tools.register(
|
||||||
|
MyTool(
|
||||||
|
runtime_control=AgentRuntimeControl(self),
|
||||||
|
modify_allowed=self.tools_config.my.allow_set,
|
||||||
|
)
|
||||||
|
)
|
||||||
|
registered.append("my")
|
||||||
|
|
||||||
logger.info("Registered {} tools: {}", len(registered), registered)
|
logger.info("Registered {} tools: {}", len(registered), registered)
|
||||||
|
|
||||||
def register_runtime_context_provider(
|
def register_runtime_context_provider(
|
||||||
@@ -1719,12 +1733,18 @@ class AgentLoop:
|
|||||||
msg: InboundMessage,
|
msg: InboundMessage,
|
||||||
final_content: str,
|
final_content: str,
|
||||||
stop_reason: str,
|
stop_reason: str,
|
||||||
|
had_injections: bool,
|
||||||
streamed_content: bool,
|
streamed_content: bool,
|
||||||
*,
|
*,
|
||||||
log_content: bool = True,
|
log_content: bool = True,
|
||||||
turn_latency_ms: int | None = None,
|
turn_latency_ms: int | None = None,
|
||||||
) -> OutboundMessage | None:
|
) -> OutboundMessage | None:
|
||||||
"""Assemble the final outbound message from turn results."""
|
"""Assemble the final outbound message from turn results."""
|
||||||
|
# MessageTool suppression
|
||||||
|
if (mt := self.tools.get("message")) and isinstance(mt, MessageTool) and mt._sent_in_turn:
|
||||||
|
if not had_injections or stop_reason == "empty_final_response":
|
||||||
|
return None
|
||||||
|
|
||||||
if log_content:
|
if log_content:
|
||||||
preview = final_content[:120] + "..." if len(final_content) > 120 else final_content
|
preview = final_content[:120] + "..." if len(final_content) > 120 else final_content
|
||||||
logger.info("Response to {}:{}: {}", msg.channel, msg.sender_id, preview)
|
logger.info("Response to {}:{}: {}", msg.channel, msg.sender_id, preview)
|
||||||
@@ -1880,6 +1900,10 @@ class AgentLoop:
|
|||||||
)
|
)
|
||||||
is_subagent = ctx.kind is TurnKind.SYSTEM and ctx.msg.sender_id == "subagent"
|
is_subagent = ctx.kind is TurnKind.SYSTEM and ctx.msg.sender_id == "subagent"
|
||||||
|
|
||||||
|
if ctx.kind is TurnKind.USER and (message_tool := self.tools.get("message")):
|
||||||
|
if isinstance(message_tool, MessageTool):
|
||||||
|
message_tool.start_turn()
|
||||||
|
|
||||||
_hist_kwargs: dict[str, Any] = {
|
_hist_kwargs: dict[str, Any] = {
|
||||||
"max_tokens": self._replay_token_budget(runtime),
|
"max_tokens": self._replay_token_budget(runtime),
|
||||||
"extend_to_user": is_subagent,
|
"extend_to_user": is_subagent,
|
||||||
@@ -1980,7 +2004,6 @@ class AgentLoop:
|
|||||||
if ctx.visible_run_started_at is None:
|
if ctx.visible_run_started_at is None:
|
||||||
ctx.visible_run_started_at = time.time()
|
ctx.visible_run_started_at = time.time()
|
||||||
await ctx.delivery.running(started_at=ctx.visible_run_started_at)
|
await ctx.delivery.running(started_at=ctx.visible_run_started_at)
|
||||||
with capture_message_deliveries() as message_sends:
|
|
||||||
result = await self._run_agent_loop(
|
result = await self._run_agent_loop(
|
||||||
ctx.initial_messages,
|
ctx.initial_messages,
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
@@ -2002,12 +2025,7 @@ class AgentLoop:
|
|||||||
ctx.final_content = result.final_content
|
ctx.final_content = result.final_content
|
||||||
ctx.all_messages = result.messages
|
ctx.all_messages = result.messages
|
||||||
ctx.stop_reason = result.stop_reason
|
ctx.stop_reason = result.stop_reason
|
||||||
if (
|
ctx.had_injections = result.had_injections
|
||||||
ctx.kind is TurnKind.USER
|
|
||||||
and (ctx.delivery.route.channel, ctx.delivery.route.chat_id) in message_sends
|
|
||||||
and (not result.had_injections or result.stop_reason == "empty_final_response")
|
|
||||||
):
|
|
||||||
ctx.suppress_response = True
|
|
||||||
ctx.usage = result.usage
|
ctx.usage = result.usage
|
||||||
ctx.delivery.record_usage(ctx.usage)
|
ctx.delivery.record_usage(ctx.usage)
|
||||||
if ctx.kind is TurnKind.USER:
|
if ctx.kind is TurnKind.USER:
|
||||||
@@ -2076,6 +2094,7 @@ class AgentLoop:
|
|||||||
ctx.delivery.delivery_message,
|
ctx.delivery.delivery_message,
|
||||||
cast(str, ctx.final_content),
|
cast(str, ctx.final_content),
|
||||||
ctx.stop_reason,
|
ctx.stop_reason,
|
||||||
|
ctx.had_injections,
|
||||||
ctx.streamed_content,
|
ctx.streamed_content,
|
||||||
log_content=ctx.require_session().policy.log_content,
|
log_content=ctx.require_session().policy.log_content,
|
||||||
turn_latency_ms=ctx.turn_latency_ms,
|
turn_latency_ms=ctx.turn_latency_ms,
|
||||||
|
|||||||
+183
-209
@@ -1,4 +1,4 @@
|
|||||||
"""Memory storage, transcript archiving, and legacy consolidation coordination."""
|
"""Memory system: pure file I/O store and lightweight Consolidator."""
|
||||||
|
|
||||||
# Tool schemas are installed by the ``@tool_parameters`` class decorator at
|
# Tool schemas are installed by the ``@tool_parameters`` class decorator at
|
||||||
# runtime; static analyzers cannot observe that it clears ``parameters`` from
|
# runtime; static analyzers cannot observe that it clears ``parameters`` from
|
||||||
@@ -32,6 +32,7 @@ from nanobot.utils.gitstore import GitStore
|
|||||||
from nanobot.utils.helpers import (
|
from nanobot.utils.helpers import (
|
||||||
content_with_media_breadcrumbs,
|
content_with_media_breadcrumbs,
|
||||||
ensure_dir,
|
ensure_dir,
|
||||||
|
estimate_message_tokens,
|
||||||
estimate_prompt_tokens_chain,
|
estimate_prompt_tokens_chain,
|
||||||
strip_think,
|
strip_think,
|
||||||
truncate_text,
|
truncate_text,
|
||||||
@@ -784,7 +785,7 @@ class MemoryStore:
|
|||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
# Memory ingestion and legacy context-pressure coordination
|
# Consolidator — lightweight token-budget triggered consolidation
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
# Individual history.jsonl writers cap their own payloads tightly; the
|
# Individual history.jsonl writers cap their own payloads tightly; the
|
||||||
@@ -795,165 +796,10 @@ _ARCHIVE_SUMMARY_MAX_CHARS = 8_000 # LLM-produced consolidation summary
|
|||||||
_HISTORY_ENTRY_HARD_CAP = 64_000 # emergency cap in append_history
|
_HISTORY_ENTRY_HARD_CAP = 64_000 # emergency cap in append_history
|
||||||
|
|
||||||
|
|
||||||
class MemoryArchiver:
|
|
||||||
"""Write durable transcript batches to the Memory ingestion journal.
|
|
||||||
|
|
||||||
The archiver deliberately has no SessionManager dependency: it may read a
|
|
||||||
captured transcript batch and append to history.jsonl, but it cannot mutate
|
|
||||||
provider continuation state or advance a session watermark.
|
|
||||||
"""
|
|
||||||
|
|
||||||
def __init__(
|
|
||||||
self,
|
|
||||||
store: MemoryStore,
|
|
||||||
build_messages: Callable[..., list[dict[str, Any]]],
|
|
||||||
get_tool_definitions: Callable[[], list[dict[str, Any]]],
|
|
||||||
resolve_prompt_context: Callable[[Session], tuple[str | None, Path | None]] | None = None,
|
|
||||||
unified_session: bool = False,
|
|
||||||
) -> None:
|
|
||||||
self.store = store
|
|
||||||
self._build_messages = build_messages
|
|
||||||
self._get_tool_definitions = get_tool_definitions
|
|
||||||
self._resolve_prompt_context = resolve_prompt_context
|
|
||||||
self.unified_session = unified_session
|
|
||||||
|
|
||||||
async def archive(
|
|
||||||
self,
|
|
||||||
messages: list[dict[str, Any]],
|
|
||||||
*,
|
|
||||||
runtime: LLMRuntime,
|
|
||||||
session_key: str,
|
|
||||||
request_messages: list[dict[str, Any]],
|
|
||||||
request_tools: list[dict[str, Any]],
|
|
||||||
) -> str | None:
|
|
||||||
"""Execute a prepared archive request and persist its result."""
|
|
||||||
if not messages:
|
|
||||||
return None
|
|
||||||
try:
|
|
||||||
with llm_usage_source("dream"):
|
|
||||||
response = await runtime.provider.chat_with_retry(
|
|
||||||
model=runtime.model,
|
|
||||||
messages=request_messages,
|
|
||||||
tools=request_tools,
|
|
||||||
tool_choice="none",
|
|
||||||
temperature=runtime.generation.temperature,
|
|
||||||
max_tokens=runtime.generation.max_tokens,
|
|
||||||
reasoning_effort=runtime.generation.reasoning_effort,
|
|
||||||
)
|
|
||||||
except Exception:
|
|
||||||
logger.warning("Memory archive provider call failed, raw-dumping to history")
|
|
||||||
self.store.raw_archive(messages, session_key=session_key)
|
|
||||||
return None
|
|
||||||
if response.finish_reason in {"error", "length"}:
|
|
||||||
logger.warning(
|
|
||||||
"Memory archive provider did not complete ({}), raw-dumping to history",
|
|
||||||
response.finish_reason,
|
|
||||||
)
|
|
||||||
self.store.raw_archive(messages, session_key=session_key)
|
|
||||||
return None
|
|
||||||
if response.has_tool_calls is True:
|
|
||||||
logger.warning("Memory archive provider returned tool calls, raw-dumping to history")
|
|
||||||
self.store.raw_archive(messages, session_key=session_key)
|
|
||||||
return None
|
|
||||||
summary = response.content
|
|
||||||
if not summary or not summary.strip():
|
|
||||||
logger.warning("Memory archive provider returned no summary, raw-dumping to history")
|
|
||||||
self.store.raw_archive(messages, session_key=session_key)
|
|
||||||
return None
|
|
||||||
if summary.strip() == "(nothing)":
|
|
||||||
return "(nothing)"
|
|
||||||
self.store.append_history(
|
|
||||||
summary,
|
|
||||||
max_chars=_ARCHIVE_SUMMARY_MAX_CHARS,
|
|
||||||
session_key=session_key,
|
|
||||||
)
|
|
||||||
return summary
|
|
||||||
|
|
||||||
async def archive_session(
|
|
||||||
self,
|
|
||||||
session: Session,
|
|
||||||
*,
|
|
||||||
archive_end: int,
|
|
||||||
runtime: LLMRuntime,
|
|
||||||
input_token_budget: int,
|
|
||||||
) -> str | None:
|
|
||||||
"""Archive a captured session prefix without mutating the session."""
|
|
||||||
messages = list(session.messages[session.last_archived:archive_end])
|
|
||||||
if not messages:
|
|
||||||
return None
|
|
||||||
if input_token_budget <= 0:
|
|
||||||
logger.debug(
|
|
||||||
"Memory archive has no safe input budget for {}; raw-dumping",
|
|
||||||
session.key,
|
|
||||||
)
|
|
||||||
self.store.raw_archive(messages, session_key=session.key)
|
|
||||||
return None
|
|
||||||
prefix = Session(
|
|
||||||
key=session.key,
|
|
||||||
messages=list(session.messages[:archive_end]),
|
|
||||||
last_consolidated=session.last_archived,
|
|
||||||
)
|
|
||||||
history = prefix.get_history(max_tokens=input_token_budget)
|
|
||||||
archive_history = Session(
|
|
||||||
key=session.key,
|
|
||||||
messages=messages,
|
|
||||||
).get_history()
|
|
||||||
if not archive_history or history[-len(archive_history):] != archive_history:
|
|
||||||
logger.debug(
|
|
||||||
"Memory archive cannot replay the full chunk for {}; raw-dumping",
|
|
||||||
session.key,
|
|
||||||
)
|
|
||||||
self.store.raw_archive(messages, session_key=session.key)
|
|
||||||
return None
|
|
||||||
prompt = render_template(
|
|
||||||
"agent/consolidator_archive.md",
|
|
||||||
strip=True,
|
|
||||||
archive_count=len(archive_history),
|
|
||||||
)
|
|
||||||
channel = session.key.split(":", 1)[0] if ":" in session.key else None
|
|
||||||
workspace: Path | None = None
|
|
||||||
if self._resolve_prompt_context is not None:
|
|
||||||
channel, workspace = self._resolve_prompt_context(session)
|
|
||||||
request_messages = self._build_messages(
|
|
||||||
history=history,
|
|
||||||
current_message=prompt,
|
|
||||||
channel=channel,
|
|
||||||
session_summary=session_summary_from_metadata(
|
|
||||||
session.metadata,
|
|
||||||
fallback_last_active=session.updated_at,
|
|
||||||
),
|
|
||||||
workspace=workspace,
|
|
||||||
session_key=session.key,
|
|
||||||
unified_session=self.unified_session,
|
|
||||||
)
|
|
||||||
tools = self._get_tool_definitions()
|
|
||||||
estimated, source = estimate_prompt_tokens_chain(
|
|
||||||
runtime.provider,
|
|
||||||
runtime.model,
|
|
||||||
request_messages,
|
|
||||||
tools,
|
|
||||||
)
|
|
||||||
if estimated > input_token_budget:
|
|
||||||
logger.debug(
|
|
||||||
"Memory archive prefix exceeds budget for {}; raw-dumping: {}/{} via {}",
|
|
||||||
session.key,
|
|
||||||
estimated,
|
|
||||||
input_token_budget,
|
|
||||||
source,
|
|
||||||
)
|
|
||||||
self.store.raw_archive(messages, session_key=session.key)
|
|
||||||
return None
|
|
||||||
return await self.archive(
|
|
||||||
messages,
|
|
||||||
runtime=runtime,
|
|
||||||
session_key=session.key,
|
|
||||||
request_messages=request_messages,
|
|
||||||
request_tools=tools,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
class Consolidator:
|
class Consolidator:
|
||||||
"""Legacy context-pressure coordinator backed by a MemoryArchiver."""
|
"""Summarize compacted messages into history.jsonl."""
|
||||||
|
|
||||||
|
_MAX_CONSOLIDATION_ROUNDS = 5
|
||||||
|
|
||||||
_SAFETY_BUFFER = 1024 # extra headroom for tokenizer estimation drift
|
_SAFETY_BUFFER = 1024 # extra headroom for tokenizer estimation drift
|
||||||
|
|
||||||
@@ -964,21 +810,16 @@ class Consolidator:
|
|||||||
build_messages: Callable[..., list[dict[str, Any]]],
|
build_messages: Callable[..., list[dict[str, Any]]],
|
||||||
get_tool_definitions: Callable[[], list[dict[str, Any]]],
|
get_tool_definitions: Callable[[], list[dict[str, Any]]],
|
||||||
resolve_prompt_context: Callable[[Session], tuple[str | None, Path | None]] | None = None,
|
resolve_prompt_context: Callable[[Session], tuple[str | None, Path | None]] | None = None,
|
||||||
|
consolidation_ratio: float = 0.5,
|
||||||
unified_session: bool = False,
|
unified_session: bool = False,
|
||||||
):
|
):
|
||||||
self.store = store
|
self.store = store
|
||||||
self.sessions = sessions
|
self.sessions = sessions
|
||||||
|
self.consolidation_ratio = consolidation_ratio
|
||||||
self.unified_session = unified_session
|
self.unified_session = unified_session
|
||||||
self._build_messages = build_messages
|
self._build_messages = build_messages
|
||||||
self._get_tool_definitions = get_tool_definitions
|
self._get_tool_definitions = get_tool_definitions
|
||||||
self._resolve_prompt_context = resolve_prompt_context
|
self._resolve_prompt_context = resolve_prompt_context
|
||||||
self.archiver = MemoryArchiver(
|
|
||||||
store=store,
|
|
||||||
build_messages=build_messages,
|
|
||||||
get_tool_definitions=get_tool_definitions,
|
|
||||||
resolve_prompt_context=resolve_prompt_context,
|
|
||||||
unified_session=unified_session,
|
|
||||||
)
|
|
||||||
self._locks: weakref.WeakValueDictionary[str, asyncio.Lock] = (
|
self._locks: weakref.WeakValueDictionary[str, asyncio.Lock] = (
|
||||||
weakref.WeakValueDictionary()
|
weakref.WeakValueDictionary()
|
||||||
)
|
)
|
||||||
@@ -990,19 +831,24 @@ class Consolidator:
|
|||||||
def pick_consolidation_boundary(
|
def pick_consolidation_boundary(
|
||||||
self,
|
self,
|
||||||
session: Session,
|
session: Session,
|
||||||
) -> int | None:
|
tokens_to_remove: int,
|
||||||
"""Return the fixed user-led boundary before the recent replay tail."""
|
) -> tuple[int, int] | None:
|
||||||
if not session.messages:
|
"""Pick a user-turn boundary that removes enough old prompt tokens."""
|
||||||
|
start = session.last_consolidated
|
||||||
|
if start >= len(session.messages) or tokens_to_remove <= 0:
|
||||||
return None
|
return None
|
||||||
boundary = max(0, len(session.messages) - MIN_COMPACTED_REPLAY_MESSAGES)
|
|
||||||
while boundary > 0 and session.messages[boundary].get("role") != "user":
|
removed_tokens = 0
|
||||||
boundary -= 1
|
last_boundary: tuple[int, int] | None = None
|
||||||
if (
|
for idx in range(start, len(session.messages)):
|
||||||
boundary <= session.last_archived
|
message = session.messages[idx]
|
||||||
or session.messages[boundary].get("role") != "user"
|
if idx > start and message.get("role") == "user":
|
||||||
):
|
last_boundary = (idx, removed_tokens)
|
||||||
return None
|
if removed_tokens >= tokens_to_remove:
|
||||||
return boundary
|
return last_boundary
|
||||||
|
removed_tokens += estimate_message_tokens(message)
|
||||||
|
|
||||||
|
return last_boundary
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
def _full_replay_history(
|
def _full_replay_history(
|
||||||
@@ -1066,14 +912,48 @@ class Consolidator:
|
|||||||
request_messages: list[dict[str, Any]],
|
request_messages: list[dict[str, Any]],
|
||||||
request_tools: list[dict[str, Any]],
|
request_tools: list[dict[str, Any]],
|
||||||
) -> str | None:
|
) -> str | None:
|
||||||
"""Compatibility wrapper for the extracted MemoryArchiver."""
|
"""Execute a prepared consolidation request and persist its result."""
|
||||||
return await self.archiver.archive(
|
if not messages:
|
||||||
messages,
|
return None
|
||||||
runtime=runtime,
|
try:
|
||||||
session_key=session_key,
|
with llm_usage_source("dream"):
|
||||||
request_messages=request_messages,
|
response = await runtime.provider.chat_with_retry(
|
||||||
request_tools=request_tools,
|
model=runtime.model,
|
||||||
|
messages=request_messages,
|
||||||
|
tools=request_tools,
|
||||||
|
tool_choice="none",
|
||||||
|
temperature=runtime.generation.temperature,
|
||||||
|
max_tokens=runtime.generation.max_tokens,
|
||||||
|
reasoning_effort=runtime.generation.reasoning_effort,
|
||||||
)
|
)
|
||||||
|
except Exception:
|
||||||
|
logger.warning("Consolidation provider call failed, raw-dumping to history")
|
||||||
|
self.store.raw_archive(messages, session_key=session_key)
|
||||||
|
return None
|
||||||
|
if response.finish_reason in {"error", "length"}:
|
||||||
|
logger.warning(
|
||||||
|
"Consolidation provider did not complete ({}), raw-dumping to history",
|
||||||
|
response.finish_reason,
|
||||||
|
)
|
||||||
|
self.store.raw_archive(messages, session_key=session_key)
|
||||||
|
return None
|
||||||
|
if response.has_tool_calls is True:
|
||||||
|
logger.warning("Consolidation provider returned tool calls, raw-dumping to history")
|
||||||
|
self.store.raw_archive(messages, session_key=session_key)
|
||||||
|
return None
|
||||||
|
summary = response.content
|
||||||
|
if not summary or not summary.strip():
|
||||||
|
logger.warning("Consolidation provider returned no summary, raw-dumping to history")
|
||||||
|
self.store.raw_archive(messages, session_key=session_key)
|
||||||
|
return None
|
||||||
|
if summary.strip() == "(nothing)":
|
||||||
|
return "(nothing)"
|
||||||
|
self.store.append_history(
|
||||||
|
summary,
|
||||||
|
max_chars=_ARCHIVE_SUMMARY_MAX_CHARS,
|
||||||
|
session_key=session_key,
|
||||||
|
)
|
||||||
|
return summary
|
||||||
|
|
||||||
async def archive_session(
|
async def archive_session(
|
||||||
self,
|
self,
|
||||||
@@ -1082,12 +962,82 @@ class Consolidator:
|
|||||||
archive_end: int,
|
archive_end: int,
|
||||||
runtime: LLMRuntime,
|
runtime: LLMRuntime,
|
||||||
) -> str | None:
|
) -> str | None:
|
||||||
"""Compatibility wrapper for the extracted MemoryArchiver."""
|
"""Archive a session prefix by appending a consolidation instruction."""
|
||||||
return await self.archiver.archive_session(
|
messages = list(session.messages[session.last_consolidated:archive_end])
|
||||||
session,
|
if not messages:
|
||||||
archive_end=archive_end,
|
return None
|
||||||
|
budget = self._input_token_budget(runtime)
|
||||||
|
if budget <= 0:
|
||||||
|
logger.debug(
|
||||||
|
"Consolidation has no safe input budget for {}; raw-dumping",
|
||||||
|
session.key,
|
||||||
|
)
|
||||||
|
self.store.raw_archive(messages, session_key=session.key)
|
||||||
|
return None
|
||||||
|
prefix = Session(
|
||||||
|
key=session.key,
|
||||||
|
messages=list(session.messages[:archive_end]),
|
||||||
|
last_consolidated=session.last_consolidated,
|
||||||
|
)
|
||||||
|
history = prefix.get_history(max_tokens=budget)
|
||||||
|
archive_history = Session(
|
||||||
|
key=session.key,
|
||||||
|
messages=messages,
|
||||||
|
).get_history()
|
||||||
|
if (
|
||||||
|
not archive_history
|
||||||
|
or history[-len(archive_history):] != archive_history
|
||||||
|
):
|
||||||
|
logger.debug(
|
||||||
|
"Consolidation cannot replay the full chunk for {}; raw-dumping",
|
||||||
|
session.key,
|
||||||
|
)
|
||||||
|
self.store.raw_archive(messages, session_key=session.key)
|
||||||
|
return None
|
||||||
|
prompt = render_template(
|
||||||
|
"agent/consolidator_archive.md",
|
||||||
|
strip=True,
|
||||||
|
archive_count=len(archive_history),
|
||||||
|
)
|
||||||
|
channel = session.key.split(":", 1)[0] if ":" in session.key else None
|
||||||
|
workspace: Path | None = None
|
||||||
|
if self._resolve_prompt_context is not None:
|
||||||
|
channel, workspace = self._resolve_prompt_context(session)
|
||||||
|
request_messages = self._build_messages(
|
||||||
|
history=history,
|
||||||
|
current_message=prompt,
|
||||||
|
channel=channel,
|
||||||
|
session_summary=session_summary_from_metadata(
|
||||||
|
session.metadata,
|
||||||
|
fallback_last_active=session.updated_at,
|
||||||
|
),
|
||||||
|
workspace=workspace,
|
||||||
|
session_key=session.key,
|
||||||
|
unified_session=self.unified_session,
|
||||||
|
)
|
||||||
|
tools = self._get_tool_definitions()
|
||||||
|
estimated, source = estimate_prompt_tokens_chain(
|
||||||
|
runtime.provider,
|
||||||
|
runtime.model,
|
||||||
|
request_messages,
|
||||||
|
tools,
|
||||||
|
)
|
||||||
|
if estimated > budget:
|
||||||
|
logger.debug(
|
||||||
|
"Consolidation prefix exceeds budget for {}; raw-dumping: {}/{} via {}",
|
||||||
|
session.key,
|
||||||
|
estimated,
|
||||||
|
budget,
|
||||||
|
source,
|
||||||
|
)
|
||||||
|
self.store.raw_archive(messages, session_key=session.key)
|
||||||
|
return None
|
||||||
|
return await self.archive(
|
||||||
|
messages,
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
input_token_budget=self._input_token_budget(runtime),
|
session_key=session.key,
|
||||||
|
request_messages=request_messages,
|
||||||
|
request_tools=tools,
|
||||||
)
|
)
|
||||||
|
|
||||||
async def maybe_consolidate_by_tokens(
|
async def maybe_consolidate_by_tokens(
|
||||||
@@ -1096,7 +1046,7 @@ class Consolidator:
|
|||||||
*,
|
*,
|
||||||
runtime: LLMRuntime,
|
runtime: LLMRuntime,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Archive one fixed old prefix when the prompt exceeds the safe budget.
|
"""Loop: archive old messages until prompt fits within safe budget.
|
||||||
|
|
||||||
The budget reserves space for completion tokens and a safety buffer
|
The budget reserves space for completion tokens and a safety buffer
|
||||||
so the LLM request never exceeds the context window.
|
so the LLM request never exceeds the context window.
|
||||||
@@ -1114,6 +1064,7 @@ class Consolidator:
|
|||||||
return
|
return
|
||||||
|
|
||||||
budget = self._input_token_budget(runtime)
|
budget = self._input_token_budget(runtime)
|
||||||
|
target = int(budget * self.consolidation_ratio)
|
||||||
last_summary: str | None = None
|
last_summary: str | None = None
|
||||||
estimated, source = self.estimate_session_prompt_tokens(
|
estimated, source = self.estimate_session_prompt_tokens(
|
||||||
session,
|
session,
|
||||||
@@ -1123,32 +1074,40 @@ class Consolidator:
|
|||||||
self._persist_last_summary(session, last_summary)
|
self._persist_last_summary(session, last_summary)
|
||||||
return
|
return
|
||||||
if estimated < budget:
|
if estimated < budget:
|
||||||
unarchived_count = len(session.messages) - session.last_archived
|
unconsolidated_count = len(session.messages) - session.last_consolidated
|
||||||
logger.debug(
|
logger.debug(
|
||||||
"Token consolidation idle {}: {}/{} via {}, msgs={}",
|
"Token consolidation idle {}: {}/{} via {}, msgs={}",
|
||||||
session.key,
|
session.key,
|
||||||
estimated,
|
estimated,
|
||||||
runtime.context_window_tokens,
|
runtime.context_window_tokens,
|
||||||
source,
|
source,
|
||||||
unarchived_count,
|
unconsolidated_count,
|
||||||
)
|
)
|
||||||
self._persist_last_summary(session, last_summary)
|
self._persist_last_summary(session, last_summary)
|
||||||
return
|
return
|
||||||
|
|
||||||
end_idx = self.pick_consolidation_boundary(session)
|
for round_num in range(self._MAX_CONSOLIDATION_ROUNDS):
|
||||||
if end_idx is None:
|
if estimated <= target:
|
||||||
logger.debug(
|
break
|
||||||
"Token consolidation: no safe fixed boundary for {}",
|
|
||||||
session.key,
|
|
||||||
)
|
|
||||||
return
|
|
||||||
|
|
||||||
chunk = session.messages[session.last_archived:end_idx]
|
boundary = self.pick_consolidation_boundary(session, max(1, estimated - target))
|
||||||
|
if boundary is None:
|
||||||
|
logger.debug(
|
||||||
|
"Token consolidation: no safe boundary for {} (round {})",
|
||||||
|
session.key,
|
||||||
|
round_num,
|
||||||
|
)
|
||||||
|
break
|
||||||
|
|
||||||
|
end_idx = boundary[0]
|
||||||
|
|
||||||
|
chunk = session.messages[session.last_consolidated:end_idx]
|
||||||
if not chunk:
|
if not chunk:
|
||||||
return
|
break
|
||||||
|
|
||||||
logger.info(
|
logger.info(
|
||||||
"Token consolidation for {}: {}/{} via {}, chunk={} msgs",
|
"Token consolidation round {} for {}: {}/{} via {}, chunk={} msgs",
|
||||||
|
round_num,
|
||||||
session.key,
|
session.key,
|
||||||
estimated,
|
estimated,
|
||||||
runtime.context_window_tokens,
|
runtime.context_window_tokens,
|
||||||
@@ -1160,12 +1119,26 @@ class Consolidator:
|
|||||||
archive_end=end_idx,
|
archive_end=end_idx,
|
||||||
runtime=runtime,
|
runtime=runtime,
|
||||||
)
|
)
|
||||||
# Advance either way: archive_session raw-archives on degradation,
|
# Advance the cursor either way: on success the chunk was
|
||||||
# and replaying the same chunk would duplicate Memory material.
|
# summarized; on failure archive_session() raw-archived it as
|
||||||
|
# a breadcrumb. Re-archiving the same chunk on the next call
|
||||||
|
# would just emit duplicate [RAW] entries.
|
||||||
if summary:
|
if summary:
|
||||||
last_summary = summary
|
last_summary = summary
|
||||||
session.last_archived = end_idx
|
session.last_consolidated = end_idx
|
||||||
|
session.provider_state = None
|
||||||
self.sessions.save(session)
|
self.sessions.save(session)
|
||||||
|
if not summary:
|
||||||
|
# LLM is degraded — stop hammering it this call;
|
||||||
|
# the next invocation can retry a fresh chunk.
|
||||||
|
break
|
||||||
|
|
||||||
|
estimated, source = self.estimate_session_prompt_tokens(
|
||||||
|
session,
|
||||||
|
runtime=runtime,
|
||||||
|
)
|
||||||
|
if estimated <= 0:
|
||||||
|
break
|
||||||
|
|
||||||
# Persist the last summary to session metadata so it can be injected
|
# Persist the last summary to session metadata so it can be injected
|
||||||
# into the runtime context on the next prepare_session() call, aligning
|
# into the runtime context on the next prepare_session() call, aligning
|
||||||
@@ -1197,7 +1170,7 @@ class Consolidator:
|
|||||||
self.sessions.invalidate(session_key)
|
self.sessions.invalidate(session_key)
|
||||||
session = self.sessions.get_or_create(session_key)
|
session = self.sessions.get_or_create(session_key)
|
||||||
|
|
||||||
archive_start = session.last_archived
|
archive_start = session.last_consolidated
|
||||||
messages_to_archive = list(session.messages[archive_start:])
|
messages_to_archive = list(session.messages[archive_start:])
|
||||||
if not messages_to_archive:
|
if not messages_to_archive:
|
||||||
return ""
|
return ""
|
||||||
@@ -1218,7 +1191,8 @@ class Consolidator:
|
|||||||
|
|
||||||
# A turn can append while the provider call is in flight. Advance only
|
# A turn can append while the provider call is in flight. Advance only
|
||||||
# through the captured batch so new messages remain eligible next time.
|
# through the captured batch so new messages remain eligible next time.
|
||||||
session.last_archived = archive_end
|
session.last_consolidated = archive_end
|
||||||
|
session.provider_state = None
|
||||||
self.sessions.save(session)
|
self.sessions.save(session)
|
||||||
|
|
||||||
visible = session.get_history(
|
visible = session.get_history(
|
||||||
|
|||||||
+332
-24
@@ -19,8 +19,7 @@ from nanobot.agent.context_governance import (
|
|||||||
ContextGovernor,
|
ContextGovernor,
|
||||||
)
|
)
|
||||||
from nanobot.agent.hook import AgentHook, AgentHookContext, AgentRunHookContext
|
from nanobot.agent.hook import AgentHook, AgentHookContext, AgentRunHookContext
|
||||||
from nanobot.agent.tools.execution import execute_tool_calls
|
from nanobot.agent.tools.registry import ToolRegistry, is_tool_error_result
|
||||||
from nanobot.agent.tools.registry import ToolRegistry
|
|
||||||
from nanobot.llm_usage.context import (
|
from nanobot.llm_usage.context import (
|
||||||
LLMUsageSource,
|
LLMUsageSource,
|
||||||
bind_llm_usage_source,
|
bind_llm_usage_source,
|
||||||
@@ -33,6 +32,7 @@ from nanobot.providers.base import (
|
|||||||
LLMUsage,
|
LLMUsage,
|
||||||
ProviderCallContext,
|
ProviderCallContext,
|
||||||
ProviderConversationState,
|
ProviderConversationState,
|
||||||
|
ToolCallRequest,
|
||||||
)
|
)
|
||||||
from nanobot.providers.conversation_state import (
|
from nanobot.providers.conversation_state import (
|
||||||
ProviderConversationStateController,
|
ProviderConversationStateController,
|
||||||
@@ -46,11 +46,13 @@ from nanobot.runtime_context import (
|
|||||||
from nanobot.session.history_visibility import is_hidden_history_message
|
from nanobot.session.history_visibility import is_hidden_history_message
|
||||||
from nanobot.session.recovery import PENDING_FOLLOWUP_ID_KEY
|
from nanobot.session.recovery import PENDING_FOLLOWUP_ID_KEY
|
||||||
from nanobot.utils.helpers import (
|
from nanobot.utils.helpers import (
|
||||||
|
IncrementalThinkExtractor,
|
||||||
build_assistant_message,
|
build_assistant_message,
|
||||||
estimate_message_tokens,
|
estimate_message_tokens,
|
||||||
estimate_prompt_tokens_chain,
|
estimate_prompt_tokens_chain,
|
||||||
extract_reasoning,
|
extract_reasoning,
|
||||||
strip_reasoning_tags,
|
strip_reasoning_tags,
|
||||||
|
strip_think,
|
||||||
)
|
)
|
||||||
from nanobot.utils.llm_runtime import LLMRuntime
|
from nanobot.utils.llm_runtime import LLMRuntime
|
||||||
from nanobot.utils.prompt_templates import render_template
|
from nanobot.utils.prompt_templates import render_template
|
||||||
@@ -60,9 +62,12 @@ from nanobot.utils.runtime import (
|
|||||||
build_finalization_retry_message,
|
build_finalization_retry_message,
|
||||||
build_length_recovery_message,
|
build_length_recovery_message,
|
||||||
is_blank_text,
|
is_blank_text,
|
||||||
|
repeated_external_lookup_error,
|
||||||
|
repeated_workspace_violation_error,
|
||||||
)
|
)
|
||||||
|
|
||||||
ContinuationCallback = Callable[[], str | None]
|
ContinuationCallback = Callable[[], str | None]
|
||||||
|
ProgressCallback = Callable[[str], Awaitable[None]]
|
||||||
RetryWaitCallback = Callable[[str], Awaitable[None]]
|
RetryWaitCallback = Callable[[str], Awaitable[None]]
|
||||||
CheckpointCallback = Callable[[dict[str, Any]], Awaitable[None]]
|
CheckpointCallback = Callable[[dict[str, Any]], Awaitable[None]]
|
||||||
InjectionCallback = Callable[..., Awaitable[Iterable[Any] | None]]
|
InjectionCallback = Callable[..., Awaitable[Iterable[Any] | None]]
|
||||||
@@ -107,6 +112,7 @@ class AgentRunSpec:
|
|||||||
session_key: str | None = None
|
session_key: str | None = None
|
||||||
context_block_limit: int | None = None
|
context_block_limit: int | None = None
|
||||||
provider_retry_mode: str = "standard"
|
provider_retry_mode: str = "standard"
|
||||||
|
progress_callback: ProgressCallback | None = None
|
||||||
retry_wait_callback: RetryWaitCallback | None = None
|
retry_wait_callback: RetryWaitCallback | None = None
|
||||||
checkpoint_callback: CheckpointCallback | None = None
|
checkpoint_callback: CheckpointCallback | None = None
|
||||||
injection_callback: InjectionCallback | None = None
|
injection_callback: InjectionCallback | None = None
|
||||||
@@ -584,14 +590,13 @@ class AgentRunner:
|
|||||||
|
|
||||||
await hook.before_execute_tools(context)
|
await hook.before_execute_tools(context)
|
||||||
|
|
||||||
results, new_events = await execute_tool_calls(
|
results, new_events = await self._execute_tools(
|
||||||
spec.tools,
|
spec,
|
||||||
response.tool_calls,
|
response.tool_calls,
|
||||||
concurrent=spec.concurrent_tools,
|
external_lookup_counts,
|
||||||
external_lookup_counts=external_lookup_counts,
|
workspace_violation_counts,
|
||||||
workspace_violation_counts=workspace_violation_counts,
|
hook,
|
||||||
hook=hook,
|
context,
|
||||||
context=context,
|
|
||||||
)
|
)
|
||||||
tool_events.extend(new_events)
|
tool_events.extend(new_events)
|
||||||
tools_used.extend(
|
tools_used.extend(
|
||||||
@@ -946,9 +951,15 @@ class AgentRunner:
|
|||||||
tools=spec.tools.get_definitions(),
|
tools=spec.tools.get_definitions(),
|
||||||
)
|
)
|
||||||
wants_streaming = hook.wants_streaming()
|
wants_streaming = hook.wants_streaming()
|
||||||
|
progress_callback = spec.progress_callback
|
||||||
|
wants_progress_streaming = (
|
||||||
|
not wants_streaming
|
||||||
|
and progress_callback is not None
|
||||||
|
and getattr(spec.runtime.provider, "supports_progress_deltas", False) is True
|
||||||
|
)
|
||||||
|
|
||||||
|
progress_state: dict[str, bool] | None = None
|
||||||
active_hosted_tools: dict[str, dict[str, Any]] = {}
|
active_hosted_tools: dict[str, dict[str, Any]] = {}
|
||||||
native_reasoning_open = False
|
|
||||||
request_started_at = 0.0
|
request_started_at = 0.0
|
||||||
first_output_at: float | None = None
|
first_output_at: float | None = None
|
||||||
generation_started_at: float | None = None
|
generation_started_at: float | None = None
|
||||||
@@ -971,17 +982,9 @@ class AgentRunner:
|
|||||||
generation_elapsed_s += max(0.0, time.perf_counter() - generation_started_at)
|
generation_elapsed_s += max(0.0, time.perf_counter() - generation_started_at)
|
||||||
generation_started_at = None
|
generation_started_at = None
|
||||||
|
|
||||||
async def _close_native_reasoning() -> None:
|
|
||||||
nonlocal native_reasoning_open
|
|
||||||
if not native_reasoning_open:
|
|
||||||
return
|
|
||||||
native_reasoning_open = False
|
|
||||||
await hook.emit_reasoning_end()
|
|
||||||
|
|
||||||
async def _provider_tool_event(event: dict[str, Any]) -> None:
|
async def _provider_tool_event(event: dict[str, Any]) -> None:
|
||||||
if event.get("kind") != "hosted_tool":
|
if event.get("kind") != "hosted_tool":
|
||||||
return
|
return
|
||||||
await _close_native_reasoning()
|
|
||||||
await hook.on_provider_tool_event(context, event)
|
await hook.on_provider_tool_event(context, event)
|
||||||
call_id = event.get("call_id")
|
call_id = event.get("call_id")
|
||||||
if not call_id:
|
if not call_id:
|
||||||
@@ -999,11 +1002,10 @@ class AgentRunner:
|
|||||||
_generation_delta(delta)
|
_generation_delta(delta)
|
||||||
if delta:
|
if delta:
|
||||||
context.streamed_content = True
|
context.streamed_content = True
|
||||||
await _close_native_reasoning()
|
|
||||||
await hook.on_stream(context, delta)
|
await hook.on_stream(context, delta)
|
||||||
|
|
||||||
async def _thinking(delta: str) -> None:
|
async def _thinking(delta: str) -> None:
|
||||||
nonlocal native_reasoning_open, thinking_buf
|
nonlocal thinking_buf
|
||||||
if not delta:
|
if not delta:
|
||||||
return
|
return
|
||||||
_generation_delta(delta)
|
_generation_delta(delta)
|
||||||
@@ -1013,12 +1015,10 @@ class AgentRunner:
|
|||||||
incremental = new_clean[len(prev_clean):]
|
incremental = new_clean[len(prev_clean):]
|
||||||
if incremental:
|
if incremental:
|
||||||
context.streamed_reasoning = True
|
context.streamed_reasoning = True
|
||||||
native_reasoning_open = True
|
|
||||||
await hook.emit_reasoning(incremental)
|
await hook.emit_reasoning(incremental)
|
||||||
|
|
||||||
async def _stream_recover() -> None:
|
async def _stream_recover() -> None:
|
||||||
_pause_generation()
|
_pause_generation()
|
||||||
await _close_native_reasoning()
|
|
||||||
await hook.on_stream_end(context, resuming=True)
|
await hook.on_stream_end(context, resuming=True)
|
||||||
|
|
||||||
coro = spec.runtime.provider.chat_stream_with_retry(
|
coro = spec.runtime.provider.chat_stream_with_retry(
|
||||||
@@ -1029,6 +1029,40 @@ class AgentRunner:
|
|||||||
on_tool_call_delta=_provider_tool_event,
|
on_tool_call_delta=_provider_tool_event,
|
||||||
on_stream_recover=_stream_recover,
|
on_stream_recover=_stream_recover,
|
||||||
)
|
)
|
||||||
|
elif wants_progress_streaming:
|
||||||
|
stream_buf = ""
|
||||||
|
think_extractor = IncrementalThinkExtractor()
|
||||||
|
progress_state = {"reasoning_open": False}
|
||||||
|
|
||||||
|
async def _stream_progress(delta: str) -> None:
|
||||||
|
nonlocal stream_buf
|
||||||
|
if not delta:
|
||||||
|
return
|
||||||
|
_generation_delta(delta)
|
||||||
|
prev_clean = strip_think(stream_buf)
|
||||||
|
stream_buf += delta
|
||||||
|
new_clean = strip_think(stream_buf)
|
||||||
|
incremental = new_clean[len(prev_clean):]
|
||||||
|
|
||||||
|
if await think_extractor.feed(stream_buf, hook.emit_reasoning):
|
||||||
|
context.streamed_reasoning = True
|
||||||
|
progress_state["reasoning_open"] = True
|
||||||
|
|
||||||
|
if incremental:
|
||||||
|
if progress_state["reasoning_open"]:
|
||||||
|
await hook.emit_reasoning_end()
|
||||||
|
progress_state["reasoning_open"] = False
|
||||||
|
context.streamed_content = True
|
||||||
|
callback = progress_callback
|
||||||
|
if callback is not None:
|
||||||
|
await callback(incremental)
|
||||||
|
|
||||||
|
coro = spec.runtime.provider.chat_stream_with_retry(
|
||||||
|
**kwargs,
|
||||||
|
provider_context=provider_context,
|
||||||
|
on_content_delta=_stream_progress,
|
||||||
|
on_tool_call_delta=_provider_tool_event,
|
||||||
|
)
|
||||||
else:
|
else:
|
||||||
coro = spec.runtime.provider.chat_with_retry(
|
coro = spec.runtime.provider.chat_with_retry(
|
||||||
**kwargs,
|
**kwargs,
|
||||||
@@ -1040,9 +1074,10 @@ class AgentRunner:
|
|||||||
# very slow deltas can still run forever. Use a more generous wall-clock
|
# very slow deltas can still run forever. Use a more generous wall-clock
|
||||||
# timeout for streaming while preserving NANOBOT_LLM_TIMEOUT_S=0 as an
|
# timeout for streaming while preserving NANOBOT_LLM_TIMEOUT_S=0 as an
|
||||||
# opt-out for all LLM wall-clock timeouts.
|
# opt-out for all LLM wall-clock timeouts.
|
||||||
|
is_streaming_request = wants_streaming or wants_progress_streaming
|
||||||
outer_timeout_s = (
|
outer_timeout_s = (
|
||||||
max(300.0, timeout_s * 2)
|
max(300.0, timeout_s * 2)
|
||||||
if wants_streaming and timeout_s is not None
|
if is_streaming_request and timeout_s is not None
|
||||||
else timeout_s
|
else timeout_s
|
||||||
)
|
)
|
||||||
request_started_at = time.perf_counter()
|
request_started_at = time.perf_counter()
|
||||||
@@ -1065,7 +1100,6 @@ class AgentRunner:
|
|||||||
error_kind="timeout",
|
error_kind="timeout",
|
||||||
)
|
)
|
||||||
_pause_generation()
|
_pause_generation()
|
||||||
await _close_native_reasoning()
|
|
||||||
if first_output_at is not None:
|
if first_output_at is not None:
|
||||||
response.ttft_ms = max(0, round((first_output_at - request_started_at) * 1000))
|
response.ttft_ms = max(0, round((first_output_at - request_started_at) * 1000))
|
||||||
if generation_elapsed_s > 0:
|
if generation_elapsed_s > 0:
|
||||||
@@ -1081,6 +1115,8 @@ class AgentRunner:
|
|||||||
"error": response.content
|
"error": response.content
|
||||||
or "Model request failed before the provider-hosted tool completed.",
|
or "Model request failed before the provider-hosted tool completed.",
|
||||||
})
|
})
|
||||||
|
if progress_state and progress_state.get("reasoning_open"):
|
||||||
|
await hook.emit_reasoning_end()
|
||||||
dropped, all_dropped, original_finish_reason = (
|
dropped, all_dropped, original_finish_reason = (
|
||||||
self._drop_malformed_tool_calls(response)
|
self._drop_malformed_tool_calls(response)
|
||||||
)
|
)
|
||||||
@@ -1384,6 +1420,253 @@ class AgentRunner:
|
|||||||
return left
|
return left
|
||||||
return left + right
|
return left + right
|
||||||
|
|
||||||
|
async def _execute_tools(
|
||||||
|
self,
|
||||||
|
spec: AgentRunSpec,
|
||||||
|
tool_calls: list[ToolCallRequest],
|
||||||
|
external_lookup_counts: dict[str, int],
|
||||||
|
workspace_violation_counts: dict[str, int],
|
||||||
|
hook: AgentHook | None = None,
|
||||||
|
context: AgentHookContext | None = None,
|
||||||
|
) -> tuple[list[Any], list[dict[str, str]]]:
|
||||||
|
hook = hook or AgentHook()
|
||||||
|
context = context or AgentHookContext(iteration=0, messages=[])
|
||||||
|
batches = self._partition_tool_batches(spec, tool_calls)
|
||||||
|
tool_results: list[tuple[Any, dict[str, str]]] = []
|
||||||
|
for batch in batches:
|
||||||
|
if spec.concurrent_tools and len(batch) > 1:
|
||||||
|
batch_results = await asyncio.gather(*(
|
||||||
|
self._run_tool(
|
||||||
|
spec,
|
||||||
|
tool_call,
|
||||||
|
external_lookup_counts,
|
||||||
|
workspace_violation_counts,
|
||||||
|
hook,
|
||||||
|
context,
|
||||||
|
)
|
||||||
|
for tool_call in batch
|
||||||
|
))
|
||||||
|
tool_results.extend(batch_results)
|
||||||
|
else:
|
||||||
|
batch_results: list[tuple[Any, dict[str, str]]] = []
|
||||||
|
for tool_call in batch:
|
||||||
|
result = await self._run_tool(
|
||||||
|
spec,
|
||||||
|
tool_call,
|
||||||
|
external_lookup_counts,
|
||||||
|
workspace_violation_counts,
|
||||||
|
hook,
|
||||||
|
context,
|
||||||
|
)
|
||||||
|
tool_results.append(result)
|
||||||
|
batch_results.append(result)
|
||||||
|
|
||||||
|
results: list[Any] = []
|
||||||
|
events: list[dict[str, str]] = []
|
||||||
|
for result, event in tool_results:
|
||||||
|
results.append(result)
|
||||||
|
events.append(event)
|
||||||
|
return results, events
|
||||||
|
|
||||||
|
async def _run_tool(
|
||||||
|
self,
|
||||||
|
spec: AgentRunSpec,
|
||||||
|
tool_call: ToolCallRequest,
|
||||||
|
external_lookup_counts: dict[str, int],
|
||||||
|
workspace_violation_counts: dict[str, int],
|
||||||
|
hook: AgentHook | None = None,
|
||||||
|
context: AgentHookContext | None = None,
|
||||||
|
) -> tuple[Any, dict[str, str]]:
|
||||||
|
hook = hook or AgentHook()
|
||||||
|
context = context or AgentHookContext(iteration=0, messages=[])
|
||||||
|
hint = "\n\n[Analyze the error above and try a different approach.]"
|
||||||
|
lookup_error = repeated_external_lookup_error(
|
||||||
|
tool_call.name,
|
||||||
|
tool_call.arguments,
|
||||||
|
external_lookup_counts,
|
||||||
|
)
|
||||||
|
if lookup_error:
|
||||||
|
event = {
|
||||||
|
"name": tool_call.name,
|
||||||
|
"status": "error",
|
||||||
|
"detail": "repeated external lookup blocked",
|
||||||
|
}
|
||||||
|
return lookup_error + hint, event
|
||||||
|
prepare_call = cast(
|
||||||
|
Callable[[str, Any], object] | None,
|
||||||
|
getattr(spec.tools, "prepare_call", None),
|
||||||
|
)
|
||||||
|
tool, params, prep_error = None, tool_call.arguments, None
|
||||||
|
if callable(prepare_call):
|
||||||
|
prepared = prepare_call(tool_call.name, tool_call.arguments)
|
||||||
|
if isinstance(prepared, tuple):
|
||||||
|
prepared_tuple = cast(tuple[object, ...], prepared)
|
||||||
|
if len(prepared_tuple) == 3:
|
||||||
|
tool, params, prep_error = cast(tuple[Any, Any, str | None], prepared_tuple)
|
||||||
|
if prep_error:
|
||||||
|
event = {
|
||||||
|
"name": tool_call.name,
|
||||||
|
"status": "error",
|
||||||
|
"detail": prep_error.split(": ", 1)[-1][:120],
|
||||||
|
}
|
||||||
|
handled = self._classify_violation(
|
||||||
|
raw_text=prep_error,
|
||||||
|
soft_payload=prep_error + hint,
|
||||||
|
event=event,
|
||||||
|
tool_call=tool_call,
|
||||||
|
workspace_violation_counts=workspace_violation_counts,
|
||||||
|
)
|
||||||
|
if handled is not None:
|
||||||
|
return handled
|
||||||
|
return prep_error + hint, event
|
||||||
|
await hook.before_execute_tool(context, tool_call, tool, params)
|
||||||
|
try:
|
||||||
|
if tool is not None:
|
||||||
|
result = await tool.execute(**params)
|
||||||
|
else:
|
||||||
|
result = await spec.tools.execute(tool_call.name, params)
|
||||||
|
except asyncio.CancelledError:
|
||||||
|
raise
|
||||||
|
except Exception as exc:
|
||||||
|
await hook.on_execute_tool_error(context, tool_call, tool, params, exc)
|
||||||
|
event = {
|
||||||
|
"name": tool_call.name,
|
||||||
|
"status": "error",
|
||||||
|
"detail": str(exc),
|
||||||
|
}
|
||||||
|
payload = f"Error: {type(exc).__name__}: {exc}"
|
||||||
|
handled = self._classify_violation(
|
||||||
|
raw_text=str(exc),
|
||||||
|
# Preserve legacy exception payloads without the retry hint.
|
||||||
|
soft_payload=payload,
|
||||||
|
event=event,
|
||||||
|
tool_call=tool_call,
|
||||||
|
workspace_violation_counts=workspace_violation_counts,
|
||||||
|
)
|
||||||
|
if handled is not None:
|
||||||
|
return handled
|
||||||
|
return payload, event
|
||||||
|
|
||||||
|
if is_tool_error_result(result):
|
||||||
|
await hook.on_execute_tool_error(context, tool_call, tool, params, result)
|
||||||
|
event = {
|
||||||
|
"name": tool_call.name,
|
||||||
|
"status": "error",
|
||||||
|
"detail": result.replace("\n", " ").strip()[:120],
|
||||||
|
}
|
||||||
|
handled = self._classify_violation(
|
||||||
|
raw_text=result,
|
||||||
|
soft_payload=result + hint,
|
||||||
|
event=event,
|
||||||
|
tool_call=tool_call,
|
||||||
|
workspace_violation_counts=workspace_violation_counts,
|
||||||
|
)
|
||||||
|
if handled is not None:
|
||||||
|
return handled
|
||||||
|
return result + hint, event
|
||||||
|
|
||||||
|
await hook.after_execute_tool(context, tool_call, tool, params, result)
|
||||||
|
|
||||||
|
detail = "" if result is None else str(result)
|
||||||
|
detail = detail.replace("\n", " ").strip()
|
||||||
|
if not detail:
|
||||||
|
detail = "(empty)"
|
||||||
|
elif len(detail) > 120:
|
||||||
|
detail = detail[:120] + "..."
|
||||||
|
return result, {"name": tool_call.name, "status": "ok", "detail": detail}
|
||||||
|
|
||||||
|
# SSRF is a hard security block at the tool boundary, but the agent turn
|
||||||
|
# should recover conversationally instead of aborting the runtime.
|
||||||
|
_SSRF_MARKERS: tuple[str, ...] = (
|
||||||
|
"internal/private url detected",
|
||||||
|
"private/internal address",
|
||||||
|
"private address",
|
||||||
|
)
|
||||||
|
_SSRF_BOUNDARY_NOTE: str = (
|
||||||
|
"This is a non-bypassable security boundary. Stop trying to access "
|
||||||
|
"private/internal URLs. Do not retry with curl, wget, encoded IPs, "
|
||||||
|
"alternate DNS, redirects, proxies, or another tool. Ask the user for "
|
||||||
|
"local files, logs, screenshots, or an explicit safe public URL instead. "
|
||||||
|
"If the user explicitly trusts this private URL, ask them to whitelist "
|
||||||
|
"the exact IP/CIDR via tools.ssrfWhitelist."
|
||||||
|
)
|
||||||
|
|
||||||
|
# Non-SSRF boundary markers returned to the LLM as recoverable tool errors.
|
||||||
|
_WORKSPACE_VIOLATION_MARKERS: tuple[str, ...] = (
|
||||||
|
"outside the configured workspace",
|
||||||
|
"outside allowed directory",
|
||||||
|
"working_dir is outside",
|
||||||
|
"working_dir could not be resolved",
|
||||||
|
"path outside working dir",
|
||||||
|
"path traversal detected",
|
||||||
|
)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _is_ssrf_violation(cls, text: str) -> bool:
|
||||||
|
if not text:
|
||||||
|
return False
|
||||||
|
lowered = text.lower()
|
||||||
|
return any(marker in lowered for marker in cls._SSRF_MARKERS)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _is_workspace_violation(cls, text: str) -> bool:
|
||||||
|
"""True when *text* looks like any policy boundary rejection."""
|
||||||
|
if not text:
|
||||||
|
return False
|
||||||
|
lowered = text.lower()
|
||||||
|
if cls._is_ssrf_violation(lowered):
|
||||||
|
return True
|
||||||
|
return any(marker in lowered for marker in cls._WORKSPACE_VIOLATION_MARKERS)
|
||||||
|
|
||||||
|
def _classify_violation(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
raw_text: str,
|
||||||
|
soft_payload: str,
|
||||||
|
event: dict[str, str],
|
||||||
|
tool_call: ToolCallRequest,
|
||||||
|
workspace_violation_counts: dict[str, int],
|
||||||
|
) -> tuple[Any, dict[str, str]] | None:
|
||||||
|
"""Classify safety-boundary failures, or return ``None`` to pass through."""
|
||||||
|
if self._is_ssrf_violation(raw_text):
|
||||||
|
logger.warning(
|
||||||
|
"Tool {} blocked by SSRF guard; returning non-retryable tool error: {}",
|
||||||
|
tool_call.name,
|
||||||
|
raw_text.replace("\n", " ").strip()[:200],
|
||||||
|
)
|
||||||
|
event["detail"] = self._event_detail("ssrf_violation: ", raw_text)
|
||||||
|
return self._ssrf_soft_payload(raw_text), event
|
||||||
|
|
||||||
|
if self._is_workspace_violation(raw_text):
|
||||||
|
escalation = repeated_workspace_violation_error(
|
||||||
|
tool_call.name,
|
||||||
|
tool_call.arguments,
|
||||||
|
workspace_violation_counts,
|
||||||
|
)
|
||||||
|
event["detail"] = self._event_detail("workspace_violation: ", raw_text)
|
||||||
|
if escalation is not None:
|
||||||
|
logger.warning(
|
||||||
|
"Tool {} hit workspace boundary repeatedly; escalating hint",
|
||||||
|
tool_call.name,
|
||||||
|
)
|
||||||
|
event["detail"] = self._event_detail(
|
||||||
|
"workspace_violation_escalated: ",
|
||||||
|
raw_text,
|
||||||
|
)
|
||||||
|
return escalation, event
|
||||||
|
return soft_payload, event
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _ssrf_soft_payload(cls, raw_text: str) -> str:
|
||||||
|
text = raw_text.strip() or "Error: request blocked by SSRF guard"
|
||||||
|
return f"{text}\n\n{cls._SSRF_BOUNDARY_NOTE}"
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _event_detail(prefix: str, text: str, limit: int = 160) -> str:
|
||||||
|
return (prefix + text.replace("\n", " ").strip())[:limit]
|
||||||
|
|
||||||
async def _emit_checkpoint(
|
async def _emit_checkpoint(
|
||||||
self,
|
self,
|
||||||
spec: AgentRunSpec,
|
spec: AgentRunSpec,
|
||||||
@@ -1413,3 +1696,28 @@ class AgentRunner:
|
|||||||
if messages and messages[-1].get("role") == "assistant" and not messages[-1].get("tool_calls"):
|
if messages and messages[-1].get("role") == "assistant" and not messages[-1].get("tool_calls"):
|
||||||
return
|
return
|
||||||
messages.append(build_assistant_message(_PERSISTED_MODEL_ERROR_PLACEHOLDER))
|
messages.append(build_assistant_message(_PERSISTED_MODEL_ERROR_PLACEHOLDER))
|
||||||
|
|
||||||
|
def _partition_tool_batches(
|
||||||
|
self,
|
||||||
|
spec: AgentRunSpec,
|
||||||
|
tool_calls: list[ToolCallRequest],
|
||||||
|
) -> list[list[ToolCallRequest]]:
|
||||||
|
if not spec.concurrent_tools:
|
||||||
|
return [[tool_call] for tool_call in tool_calls]
|
||||||
|
|
||||||
|
batches: list[list[ToolCallRequest]] = []
|
||||||
|
current: list[ToolCallRequest] = []
|
||||||
|
for tool_call in tool_calls:
|
||||||
|
get_tool = cast(Callable[[str], Any] | None, getattr(spec.tools, "get", None))
|
||||||
|
tool = get_tool(tool_call.name) if callable(get_tool) else None
|
||||||
|
can_batch = bool(tool and tool.concurrency_safe)
|
||||||
|
if can_batch:
|
||||||
|
current.append(tool_call)
|
||||||
|
continue
|
||||||
|
if current:
|
||||||
|
batches.append(current)
|
||||||
|
current = []
|
||||||
|
batches.append([tool_call])
|
||||||
|
if current:
|
||||||
|
batches.append(current)
|
||||||
|
return batches
|
||||||
|
|||||||
@@ -55,8 +55,7 @@ class SubagentStatus:
|
|||||||
label: str
|
label: str
|
||||||
task_description: str
|
task_description: str
|
||||||
started_at: float # time.monotonic()
|
started_at: float # time.monotonic()
|
||||||
# queued | initializing | awaiting_tools | tools_completed | final_response | done | error
|
phase: str = "initializing" # initializing | awaiting_tools | tools_completed | final_response | done | error
|
||||||
phase: str = "initializing"
|
|
||||||
iteration: int = 0
|
iteration: int = 0
|
||||||
tool_events: list[dict[str, str]] = field(default_factory=list)
|
tool_events: list[dict[str, str]] = field(default_factory=list)
|
||||||
usage: LLMUsage | None = None
|
usage: LLMUsage | None = None
|
||||||
@@ -148,7 +147,6 @@ class SubagentManager:
|
|||||||
if max_concurrent_subagents is not None
|
if max_concurrent_subagents is not None
|
||||||
else defaults.max_concurrent_subagents
|
else defaults.max_concurrent_subagents
|
||||||
)
|
)
|
||||||
self._run_slots = asyncio.Semaphore(self.max_concurrent_subagents)
|
|
||||||
self.runner = AgentRunner()
|
self.runner = AgentRunner()
|
||||||
self._exec_session_manager = ExecSessionManager()
|
self._exec_session_manager = ExecSessionManager()
|
||||||
self._llm_wall_timeout_for_session = llm_wall_timeout_for_session
|
self._llm_wall_timeout_for_session = llm_wall_timeout_for_session
|
||||||
@@ -365,35 +363,6 @@ class SubagentManager:
|
|||||||
workspace_scope: WorkspaceScope | None = None,
|
workspace_scope: WorkspaceScope | None = None,
|
||||||
*,
|
*,
|
||||||
announce: bool = True,
|
announce: bool = True,
|
||||||
) -> str:
|
|
||||||
"""Wait for capacity, then execute one subagent task."""
|
|
||||||
status.phase = "queued"
|
|
||||||
async with self._run_slots:
|
|
||||||
status.phase = "initializing"
|
|
||||||
return await self._run_admitted_subagent(
|
|
||||||
task_id,
|
|
||||||
task,
|
|
||||||
label,
|
|
||||||
origin,
|
|
||||||
status,
|
|
||||||
runtime,
|
|
||||||
origin_message_id,
|
|
||||||
workspace_scope,
|
|
||||||
announce=announce,
|
|
||||||
)
|
|
||||||
|
|
||||||
async def _run_admitted_subagent(
|
|
||||||
self,
|
|
||||||
task_id: str,
|
|
||||||
task: str,
|
|
||||||
label: str,
|
|
||||||
origin: _SubagentOrigin,
|
|
||||||
status: SubagentStatus,
|
|
||||||
runtime: LLMRuntime,
|
|
||||||
origin_message_id: str | None = None,
|
|
||||||
workspace_scope: WorkspaceScope | None = None,
|
|
||||||
*,
|
|
||||||
announce: bool = True,
|
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Execute the subagent task and announce the result."""
|
"""Execute the subagent task and announce the result."""
|
||||||
logger.info("Subagent [{}] starting task: {}", task_id, label)
|
logger.info("Subagent [{}] starting task: {}", task_id, label)
|
||||||
|
|||||||
@@ -11,7 +11,6 @@ if TYPE_CHECKING:
|
|||||||
from nanobot.agent.subagent import SubagentManager
|
from nanobot.agent.subagent import SubagentManager
|
||||||
from nanobot.agent.tools.exec_session import ExecSessionManager
|
from nanobot.agent.tools.exec_session import ExecSessionManager
|
||||||
from nanobot.agent.tools.file_state import FileStates
|
from nanobot.agent.tools.file_state import FileStates
|
||||||
from nanobot.agent.tools.runtime_control import RuntimeControl
|
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.bus.runtime_events import RuntimeEventBus
|
from nanobot.bus.runtime_events import RuntimeEventBus
|
||||||
from nanobot.config.schema import ProviderConfig, ToolsConfig
|
from nanobot.config.schema import ProviderConfig, ToolsConfig
|
||||||
@@ -91,4 +90,3 @@ class ToolContext:
|
|||||||
timezone: str = "UTC"
|
timezone: str = "UTC"
|
||||||
workspace_sandbox: WorkspaceSandboxStatus | None = None
|
workspace_sandbox: WorkspaceSandboxStatus | None = None
|
||||||
runtime_events: RuntimeEventBus | None = None
|
runtime_events: RuntimeEventBus | None = None
|
||||||
runtime_control: RuntimeControl | None = None
|
|
||||||
|
|||||||
@@ -1,285 +0,0 @@
|
|||||||
"""Execute tool calls and turn their outcomes into model observations."""
|
|
||||||
|
|
||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
from collections.abc import Callable
|
|
||||||
from typing import Any, cast
|
|
||||||
|
|
||||||
from loguru import logger
|
|
||||||
|
|
||||||
from nanobot.agent.hook import AgentHook, AgentHookContext
|
|
||||||
from nanobot.agent.tools.registry import ToolRegistry, is_tool_error_result
|
|
||||||
from nanobot.providers.base import ToolCallRequest
|
|
||||||
from nanobot.utils.runtime import (
|
|
||||||
repeated_external_lookup_error,
|
|
||||||
repeated_workspace_violation_error,
|
|
||||||
)
|
|
||||||
|
|
||||||
_RETRY_HINT = "\n\n[Analyze the error above and try a different approach.]"
|
|
||||||
# SSRF is a hard security block at the tool boundary, but the agent turn
|
|
||||||
# should recover conversationally instead of aborting the runtime.
|
|
||||||
_SSRF_MARKERS: tuple[str, ...] = (
|
|
||||||
"internal/private url detected",
|
|
||||||
"private/internal address",
|
|
||||||
"private address",
|
|
||||||
)
|
|
||||||
_SSRF_BOUNDARY_NOTE = (
|
|
||||||
"This is a non-bypassable security boundary. Stop trying to access "
|
|
||||||
"private/internal URLs. Do not retry with curl, wget, encoded IPs, "
|
|
||||||
"alternate DNS, redirects, proxies, or another tool. Ask the user for "
|
|
||||||
"local files, logs, screenshots, or an explicit safe public URL instead. "
|
|
||||||
"If the user explicitly trusts this private URL, ask them to whitelist "
|
|
||||||
"the exact IP/CIDR via tools.ssrfWhitelist."
|
|
||||||
)
|
|
||||||
# Non-SSRF boundary markers returned to the model as recoverable tool errors.
|
|
||||||
_WORKSPACE_VIOLATION_MARKERS: tuple[str, ...] = (
|
|
||||||
"outside the configured workspace",
|
|
||||||
"outside allowed directory",
|
|
||||||
"working_dir is outside",
|
|
||||||
"working_dir could not be resolved",
|
|
||||||
"path outside working dir",
|
|
||||||
"path traversal detected",
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
async def execute_tool_calls(
|
|
||||||
tools: ToolRegistry,
|
|
||||||
tool_calls: list[ToolCallRequest],
|
|
||||||
*,
|
|
||||||
concurrent: bool,
|
|
||||||
external_lookup_counts: dict[str, int],
|
|
||||||
workspace_violation_counts: dict[str, int],
|
|
||||||
hook: AgentHook,
|
|
||||||
context: AgentHookContext,
|
|
||||||
) -> tuple[list[Any], list[dict[str, str]]]:
|
|
||||||
"""Execute one model response's tool calls in stable result order."""
|
|
||||||
tool_results: list[tuple[Any, dict[str, str]]] = []
|
|
||||||
for batch in _partition_tool_batches(tools, tool_calls, concurrent=concurrent):
|
|
||||||
if concurrent and len(batch) > 1:
|
|
||||||
batch_results = await asyncio.gather(*(
|
|
||||||
_execute_tool_call(
|
|
||||||
tools,
|
|
||||||
tool_call,
|
|
||||||
external_lookup_counts,
|
|
||||||
workspace_violation_counts,
|
|
||||||
hook,
|
|
||||||
context,
|
|
||||||
)
|
|
||||||
for tool_call in batch
|
|
||||||
))
|
|
||||||
tool_results.extend(batch_results)
|
|
||||||
else:
|
|
||||||
for tool_call in batch:
|
|
||||||
result = await _execute_tool_call(
|
|
||||||
tools,
|
|
||||||
tool_call,
|
|
||||||
external_lookup_counts,
|
|
||||||
workspace_violation_counts,
|
|
||||||
hook,
|
|
||||||
context,
|
|
||||||
)
|
|
||||||
tool_results.append(result)
|
|
||||||
|
|
||||||
results = [result for result, _event in tool_results]
|
|
||||||
events = [event for _result, event in tool_results]
|
|
||||||
return results, events
|
|
||||||
|
|
||||||
|
|
||||||
async def _execute_tool_call(
|
|
||||||
tools: ToolRegistry,
|
|
||||||
tool_call: ToolCallRequest,
|
|
||||||
external_lookup_counts: dict[str, int],
|
|
||||||
workspace_violation_counts: dict[str, int],
|
|
||||||
hook: AgentHook,
|
|
||||||
context: AgentHookContext,
|
|
||||||
) -> tuple[Any, dict[str, str]]:
|
|
||||||
lookup_error = repeated_external_lookup_error(
|
|
||||||
tool_call.name,
|
|
||||||
tool_call.arguments,
|
|
||||||
external_lookup_counts,
|
|
||||||
)
|
|
||||||
if lookup_error:
|
|
||||||
event = {
|
|
||||||
"name": tool_call.name,
|
|
||||||
"status": "error",
|
|
||||||
"detail": "repeated external lookup blocked",
|
|
||||||
}
|
|
||||||
return lookup_error + _RETRY_HINT, event
|
|
||||||
|
|
||||||
prepare_call = cast(
|
|
||||||
Callable[[str, Any], object] | None,
|
|
||||||
getattr(tools, "prepare_call", None),
|
|
||||||
)
|
|
||||||
tool, params, prep_error = None, tool_call.arguments, None
|
|
||||||
if callable(prepare_call):
|
|
||||||
prepared = prepare_call(tool_call.name, tool_call.arguments)
|
|
||||||
if isinstance(prepared, tuple):
|
|
||||||
prepared_tuple = cast(tuple[object, ...], prepared)
|
|
||||||
if len(prepared_tuple) == 3:
|
|
||||||
tool, params, prep_error = cast(tuple[Any, Any, str | None], prepared_tuple)
|
|
||||||
if prep_error:
|
|
||||||
event = {
|
|
||||||
"name": tool_call.name,
|
|
||||||
"status": "error",
|
|
||||||
"detail": prep_error.split(": ", 1)[-1][:120],
|
|
||||||
}
|
|
||||||
handled = _classify_violation(
|
|
||||||
raw_text=prep_error,
|
|
||||||
soft_payload=prep_error + _RETRY_HINT,
|
|
||||||
event=event,
|
|
||||||
tool_call=tool_call,
|
|
||||||
workspace_violation_counts=workspace_violation_counts,
|
|
||||||
)
|
|
||||||
if handled is not None:
|
|
||||||
return handled
|
|
||||||
return prep_error + _RETRY_HINT, event
|
|
||||||
|
|
||||||
await hook.before_execute_tool(context, tool_call, tool, params)
|
|
||||||
try:
|
|
||||||
if tool is not None:
|
|
||||||
result = await tool.execute(**params)
|
|
||||||
else:
|
|
||||||
result = await tools.execute(tool_call.name, params)
|
|
||||||
except asyncio.CancelledError:
|
|
||||||
raise
|
|
||||||
except Exception as exc:
|
|
||||||
await hook.on_execute_tool_error(context, tool_call, tool, params, exc)
|
|
||||||
event = {
|
|
||||||
"name": tool_call.name,
|
|
||||||
"status": "error",
|
|
||||||
"detail": str(exc),
|
|
||||||
}
|
|
||||||
payload = f"Error: {type(exc).__name__}: {exc}"
|
|
||||||
handled = _classify_violation(
|
|
||||||
raw_text=str(exc),
|
|
||||||
# Preserve legacy exception payloads without the retry hint.
|
|
||||||
soft_payload=payload,
|
|
||||||
event=event,
|
|
||||||
tool_call=tool_call,
|
|
||||||
workspace_violation_counts=workspace_violation_counts,
|
|
||||||
)
|
|
||||||
if handled is not None:
|
|
||||||
return handled
|
|
||||||
return payload, event
|
|
||||||
|
|
||||||
if is_tool_error_result(result):
|
|
||||||
await hook.on_execute_tool_error(context, tool_call, tool, params, result)
|
|
||||||
event = {
|
|
||||||
"name": tool_call.name,
|
|
||||||
"status": "error",
|
|
||||||
"detail": result.replace("\n", " ").strip()[:120],
|
|
||||||
}
|
|
||||||
handled = _classify_violation(
|
|
||||||
raw_text=result,
|
|
||||||
soft_payload=result + _RETRY_HINT,
|
|
||||||
event=event,
|
|
||||||
tool_call=tool_call,
|
|
||||||
workspace_violation_counts=workspace_violation_counts,
|
|
||||||
)
|
|
||||||
if handled is not None:
|
|
||||||
return handled
|
|
||||||
return result + _RETRY_HINT, event
|
|
||||||
|
|
||||||
await hook.after_execute_tool(context, tool_call, tool, params, result)
|
|
||||||
|
|
||||||
detail = "" if result is None else str(result)
|
|
||||||
detail = detail.replace("\n", " ").strip()
|
|
||||||
if not detail:
|
|
||||||
detail = "(empty)"
|
|
||||||
elif len(detail) > 120:
|
|
||||||
detail = detail[:120] + "..."
|
|
||||||
return result, {"name": tool_call.name, "status": "ok", "detail": detail}
|
|
||||||
|
|
||||||
|
|
||||||
def is_ssrf_violation(text: str) -> bool:
|
|
||||||
"""Return whether a tool error describes a blocked private-network request."""
|
|
||||||
if not text:
|
|
||||||
return False
|
|
||||||
lowered = text.lower()
|
|
||||||
return any(marker in lowered for marker in _SSRF_MARKERS)
|
|
||||||
|
|
||||||
|
|
||||||
def _is_workspace_violation(text: str) -> bool:
|
|
||||||
"""Return whether text describes any workspace or network boundary rejection."""
|
|
||||||
if not text:
|
|
||||||
return False
|
|
||||||
lowered = text.lower()
|
|
||||||
if is_ssrf_violation(lowered):
|
|
||||||
return True
|
|
||||||
return any(marker in lowered for marker in _WORKSPACE_VIOLATION_MARKERS)
|
|
||||||
|
|
||||||
|
|
||||||
def _classify_violation(
|
|
||||||
*,
|
|
||||||
raw_text: str,
|
|
||||||
soft_payload: str,
|
|
||||||
event: dict[str, str],
|
|
||||||
tool_call: ToolCallRequest,
|
|
||||||
workspace_violation_counts: dict[str, int],
|
|
||||||
) -> tuple[Any, dict[str, str]] | None:
|
|
||||||
if is_ssrf_violation(raw_text):
|
|
||||||
logger.warning(
|
|
||||||
"Tool {} blocked by SSRF guard; returning non-retryable tool error: {}",
|
|
||||||
tool_call.name,
|
|
||||||
raw_text.replace("\n", " ").strip()[:200],
|
|
||||||
)
|
|
||||||
event["detail"] = _event_detail("ssrf_violation: ", raw_text)
|
|
||||||
return _ssrf_soft_payload(raw_text), event
|
|
||||||
|
|
||||||
if _is_workspace_violation(raw_text):
|
|
||||||
escalation = repeated_workspace_violation_error(
|
|
||||||
tool_call.name,
|
|
||||||
tool_call.arguments,
|
|
||||||
workspace_violation_counts,
|
|
||||||
)
|
|
||||||
event["detail"] = _event_detail("workspace_violation: ", raw_text)
|
|
||||||
if escalation is not None:
|
|
||||||
logger.warning(
|
|
||||||
"Tool {} hit workspace boundary repeatedly; escalating hint",
|
|
||||||
tool_call.name,
|
|
||||||
)
|
|
||||||
event["detail"] = _event_detail(
|
|
||||||
"workspace_violation_escalated: ",
|
|
||||||
raw_text,
|
|
||||||
)
|
|
||||||
return escalation, event
|
|
||||||
return soft_payload, event
|
|
||||||
|
|
||||||
return None
|
|
||||||
|
|
||||||
|
|
||||||
def _ssrf_soft_payload(raw_text: str) -> str:
|
|
||||||
text = raw_text.strip() or "Error: request blocked by SSRF guard"
|
|
||||||
return f"{text}\n\n{_SSRF_BOUNDARY_NOTE}"
|
|
||||||
|
|
||||||
|
|
||||||
def _event_detail(prefix: str, text: str, limit: int = 160) -> str:
|
|
||||||
return (prefix + text.replace("\n", " ").strip())[:limit]
|
|
||||||
|
|
||||||
|
|
||||||
def _partition_tool_batches(
|
|
||||||
tools: ToolRegistry,
|
|
||||||
tool_calls: list[ToolCallRequest],
|
|
||||||
*,
|
|
||||||
concurrent: bool,
|
|
||||||
) -> list[list[ToolCallRequest]]:
|
|
||||||
if not concurrent:
|
|
||||||
return [[tool_call] for tool_call in tool_calls]
|
|
||||||
|
|
||||||
batches: list[list[ToolCallRequest]] = []
|
|
||||||
current: list[ToolCallRequest] = []
|
|
||||||
for tool_call in tool_calls:
|
|
||||||
get_tool = cast(Callable[[str], Any] | None, getattr(tools, "get", None))
|
|
||||||
tool = get_tool(tool_call.name) if callable(get_tool) else None
|
|
||||||
can_batch = bool(tool and tool.concurrency_safe)
|
|
||||||
if can_batch:
|
|
||||||
current.append(tool_call)
|
|
||||||
continue
|
|
||||||
if current:
|
|
||||||
batches.append(current)
|
|
||||||
current = []
|
|
||||||
batches.append([tool_call])
|
|
||||||
if current:
|
|
||||||
batches.append(current)
|
|
||||||
return batches
|
|
||||||
@@ -2,11 +2,9 @@
|
|||||||
|
|
||||||
# pyright: reportIncompatibleMethodOverride=false
|
# pyright: reportIncompatibleMethodOverride=false
|
||||||
|
|
||||||
from collections.abc import Awaitable, Callable, Generator
|
|
||||||
from contextlib import contextmanager
|
|
||||||
from contextvars import ContextVar, Token
|
from contextvars import ContextVar, Token
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
from typing import Any, cast
|
from typing import Any, Awaitable, Callable, cast
|
||||||
|
|
||||||
from loguru import logger
|
from loguru import logger
|
||||||
|
|
||||||
@@ -18,22 +16,6 @@ from nanobot.bus.events import OutboundMessage
|
|||||||
from nanobot.config.paths import get_workspace_path
|
from nanobot.config.paths import get_workspace_path
|
||||||
from nanobot.security.workspace_access import current_tool_workspace
|
from nanobot.security.workspace_access import current_tool_workspace
|
||||||
|
|
||||||
_CURRENT_MESSAGE_SENDS: ContextVar[set[tuple[str, str]] | None] = ContextVar(
|
|
||||||
"message_sends",
|
|
||||||
default=None,
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
@contextmanager
|
|
||||||
def capture_message_deliveries() -> Generator[set[tuple[str, str]], None, None]:
|
|
||||||
"""Record successful MessageTool targets within one agent run."""
|
|
||||||
sends: set[tuple[str, str]] = set()
|
|
||||||
token = _CURRENT_MESSAGE_SENDS.set(sends)
|
|
||||||
try:
|
|
||||||
yield sends
|
|
||||||
finally:
|
|
||||||
_CURRENT_MESSAGE_SENDS.reset(token)
|
|
||||||
|
|
||||||
|
|
||||||
@tool_parameters(
|
@tool_parameters(
|
||||||
tool_parameters_schema(
|
tool_parameters_schema(
|
||||||
@@ -86,6 +68,7 @@ class MessageTool(Tool):
|
|||||||
self._fallback_chat_id = default_chat_id
|
self._fallback_chat_id = default_chat_id
|
||||||
self._fallback_message_id = default_message_id
|
self._fallback_message_id = default_message_id
|
||||||
self._fallback_metadata: dict[str, Any] = {}
|
self._fallback_metadata: dict[str, Any] = {}
|
||||||
|
self._sent_in_turn_var: ContextVar[bool] = ContextVar("message_sent_in_turn", default=False)
|
||||||
self._suppress_delivery_var: ContextVar[bool] = ContextVar(
|
self._suppress_delivery_var: ContextVar[bool] = ContextVar(
|
||||||
"message_suppress_delivery",
|
"message_suppress_delivery",
|
||||||
default=False,
|
default=False,
|
||||||
@@ -104,6 +87,10 @@ class MessageTool(Tool):
|
|||||||
"""Set the callback for sending messages."""
|
"""Set the callback for sending messages."""
|
||||||
self._send_callback = callback
|
self._send_callback = callback
|
||||||
|
|
||||||
|
def start_turn(self) -> None:
|
||||||
|
"""Reset per-turn send tracking."""
|
||||||
|
self._sent_in_turn = False
|
||||||
|
|
||||||
def set_suppress_delivery(self, active: bool) -> Token[bool]:
|
def set_suppress_delivery(self, active: bool) -> Token[bool]:
|
||||||
"""Acknowledge but don't deliver tool sends (heartbeat internal check)."""
|
"""Acknowledge but don't deliver tool sends (heartbeat internal check)."""
|
||||||
return self._suppress_delivery_var.set(active)
|
return self._suppress_delivery_var.set(active)
|
||||||
@@ -112,6 +99,14 @@ class MessageTool(Tool):
|
|||||||
"""Restore previous delivery-suppression state."""
|
"""Restore previous delivery-suppression state."""
|
||||||
self._suppress_delivery_var.reset(token)
|
self._suppress_delivery_var.reset(token)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def _sent_in_turn(self) -> bool:
|
||||||
|
return self._sent_in_turn_var.get()
|
||||||
|
|
||||||
|
@_sent_in_turn.setter
|
||||||
|
def _sent_in_turn(self, value: bool) -> None:
|
||||||
|
self._sent_in_turn_var.set(value)
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def name(self) -> str:
|
def name(self) -> str:
|
||||||
return "message"
|
return "message"
|
||||||
@@ -249,9 +244,8 @@ class MessageTool(Tool):
|
|||||||
|
|
||||||
try:
|
try:
|
||||||
await self._send_callback(msg)
|
await self._send_callback(msg)
|
||||||
sends = _CURRENT_MESSAGE_SENDS.get()
|
if channel == default_channel and chat_id == default_chat_id:
|
||||||
if sends is not None:
|
self._sent_in_turn = True
|
||||||
sends.add((channel, chat_id))
|
|
||||||
media_info = f" with {len(media)} attachments" if media else ""
|
media_info = f" with {len(media)} attachments" if media else ""
|
||||||
button_info = (
|
button_info = (
|
||||||
f" with {sum(len(row) for row in button_rows)} button(s)"
|
f" with {sum(len(row) for row in button_rows)} button(s)"
|
||||||
|
|||||||
@@ -58,6 +58,7 @@ def _is_string_mapping(value: object) -> TypeGuard[Mapping[str, object]]:
|
|||||||
class MyTool(Tool):
|
class MyTool(Tool):
|
||||||
"""Check and set the agent loop's runtime configuration."""
|
"""Check and set the agent loop's runtime configuration."""
|
||||||
|
|
||||||
|
_plugin_discoverable = False # Requires AgentLoop reference; registered manually
|
||||||
config_key = "my"
|
config_key = "my"
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@@ -66,16 +67,7 @@ class MyTool(Tool):
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def enabled(cls, ctx: ToolContext) -> bool:
|
def enabled(cls, ctx: ToolContext) -> bool:
|
||||||
return ctx.runtime_control is not None and ctx.config.my.enable
|
return ctx.config.my.enable
|
||||||
|
|
||||||
@classmethod
|
|
||||||
def create(cls, ctx: ToolContext) -> Tool:
|
|
||||||
if ctx.runtime_control is None:
|
|
||||||
raise RuntimeError("MyTool requires a runtime control capability")
|
|
||||||
return cls(
|
|
||||||
runtime_control=ctx.runtime_control,
|
|
||||||
modify_allowed=ctx.config.my.allow_set,
|
|
||||||
)
|
|
||||||
|
|
||||||
BLOCKED = frozenset({
|
BLOCKED = frozenset({
|
||||||
# Core infrastructure
|
# Core infrastructure
|
||||||
|
|||||||
@@ -73,11 +73,6 @@ class SpawnTool(Tool):
|
|||||||
"and use a dedicated subdirectory when helpful."
|
"and use a dedicated subdirectory when helpful."
|
||||||
)
|
)
|
||||||
|
|
||||||
@property
|
|
||||||
def concurrency_safe(self) -> bool:
|
|
||||||
"""Each call owns its task state; the manager serializes capacity admission."""
|
|
||||||
return True
|
|
||||||
|
|
||||||
async def execute(
|
async def execute(
|
||||||
self,
|
self,
|
||||||
task: str,
|
task: str,
|
||||||
@@ -87,6 +82,14 @@ class SpawnTool(Tool):
|
|||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> str:
|
) -> str:
|
||||||
"""Spawn a subagent to execute the given task."""
|
"""Spawn a subagent to execute the given task."""
|
||||||
|
running = self._manager.get_running_count()
|
||||||
|
limit = self._manager.max_concurrent_subagents
|
||||||
|
if running >= limit:
|
||||||
|
return (
|
||||||
|
f"Cannot spawn subagent: concurrency limit reached "
|
||||||
|
f"({running}/{limit} running). Wait for a running subagent "
|
||||||
|
f"to complete before spawning a new one."
|
||||||
|
)
|
||||||
request_ctx = current_request_context()
|
request_ctx = current_request_context()
|
||||||
if request_ctx is None or request_ctx.runtime is None:
|
if request_ctx is None or request_ctx.runtime is None:
|
||||||
return ToolResult.error("Error: spawn requires an active model runtime")
|
return ToolResult.error("Error: spawn requires an active model runtime")
|
||||||
|
|||||||
@@ -21,7 +21,6 @@ from nanobot.cli.process_identity import named_executable
|
|||||||
from nanobot.cli.runtime_config import _model_display
|
from nanobot.cli.runtime_config import _model_display
|
||||||
from nanobot.cli.webui_support import (
|
from nanobot.cli.webui_support import (
|
||||||
_gateway_health_ready,
|
_gateway_health_ready,
|
||||||
_gateway_health_url,
|
|
||||||
_gateway_instance_command,
|
_gateway_instance_command,
|
||||||
_host_for_local_browser,
|
_host_for_local_browser,
|
||||||
_webui_endpoint_reachable,
|
_webui_endpoint_reachable,
|
||||||
@@ -97,10 +96,6 @@ def launch_tui(
|
|||||||
env.update(
|
env.update(
|
||||||
{
|
{
|
||||||
"NANOBOT_TUI_BOOTSTRAP_URL": f"{base_url}/webui/bootstrap",
|
"NANOBOT_TUI_BOOTSTRAP_URL": f"{base_url}/webui/bootstrap",
|
||||||
"NANOBOT_TUI_HEALTH_URL": _gateway_health_url(
|
|
||||||
config.gateway.host,
|
|
||||||
config.gateway.port,
|
|
||||||
),
|
|
||||||
"NANOBOT_TUI_API_URL": base_url,
|
"NANOBOT_TUI_API_URL": base_url,
|
||||||
"NANOBOT_TUI_MODEL": _model_display(config)[0],
|
"NANOBOT_TUI_MODEL": _model_display(config)[0],
|
||||||
"NANOBOT_TUI_MODEL_PRESET": config.agents.defaults.model_preset or "default",
|
"NANOBOT_TUI_MODEL_PRESET": config.agents.defaults.model_preset or "default",
|
||||||
|
|||||||
@@ -311,7 +311,7 @@ async def cmd_new(ctx: CommandContext) -> OutboundMessage:
|
|||||||
snapshot = list(session.messages)
|
snapshot = list(session.messages)
|
||||||
archive_snapshot = None
|
archive_snapshot = None
|
||||||
runtime = None
|
runtime = None
|
||||||
if session.last_archived < len(snapshot):
|
if session.last_consolidated < len(snapshot):
|
||||||
runtime = ctx.runtime or loop.runtime_for_session(session)
|
runtime = ctx.runtime or loop.runtime_for_session(session)
|
||||||
archive_snapshot = replace(
|
archive_snapshot = replace(
|
||||||
session,
|
session,
|
||||||
|
|||||||
@@ -128,7 +128,7 @@ class AgentDefaults(Base):
|
|||||||
temperature: float = 0.1
|
temperature: float = 0.1
|
||||||
fallback_models: list[FallbackCandidate] = Field(default_factory=list)
|
fallback_models: list[FallbackCandidate] = Field(default_factory=list)
|
||||||
max_tool_iterations: int = 200
|
max_tool_iterations: int = 200
|
||||||
max_concurrent_subagents: int = Field(default=4, ge=1)
|
max_concurrent_subagents: int = Field(default=1, ge=1)
|
||||||
max_tool_result_chars: int = 16_000
|
max_tool_result_chars: int = 16_000
|
||||||
provider_retry_mode: Literal["standard", "persistent"] = "standard"
|
provider_retry_mode: Literal["standard", "persistent"] = "standard"
|
||||||
tool_hint_max_length: int = Field(
|
tool_hint_max_length: int = Field(
|
||||||
@@ -155,6 +155,13 @@ class AgentDefaults(Base):
|
|||||||
default=60,
|
default=60,
|
||||||
ge=0,
|
ge=0,
|
||||||
) # Minimum interval in seconds between scans for idle sessions
|
) # Minimum interval in seconds between scans for idle sessions
|
||||||
|
consolidation_ratio: float = Field(
|
||||||
|
default=0.5,
|
||||||
|
ge=0.1,
|
||||||
|
le=0.95,
|
||||||
|
validation_alias=AliasChoices("consolidationRatio"),
|
||||||
|
serialization_alias="consolidationRatio",
|
||||||
|
) # Consolidation target ratio (0.5 = 50% of budget retained after compression)
|
||||||
dream: DreamConfig = Field(default_factory=DreamConfig)
|
dream: DreamConfig = Field(default_factory=DreamConfig)
|
||||||
|
|
||||||
@model_validator(mode="before")
|
@model_validator(mode="before")
|
||||||
|
|||||||
@@ -603,6 +603,8 @@ _SYNTHETIC_USER_CONTENT = "(conversation continued)"
|
|||||||
class LLMProvider(ABC):
|
class LLMProvider(ABC):
|
||||||
"""Base class for LLM providers."""
|
"""Base class for LLM providers."""
|
||||||
|
|
||||||
|
supports_progress_deltas = False
|
||||||
|
|
||||||
_CHAT_RETRY_DELAYS = (1, 2, 4)
|
_CHAT_RETRY_DELAYS = (1, 2, 4)
|
||||||
_PERSISTENT_MAX_DELAY = 60
|
_PERSISTENT_MAX_DELAY = 60
|
||||||
_PERSISTENT_IDENTICAL_ERROR_LIMIT = 10
|
_PERSISTENT_IDENTICAL_ERROR_LIMIT = 10
|
||||||
|
|||||||
@@ -157,6 +157,10 @@ class FallbackProvider(LLMProvider):
|
|||||||
super().set_llm_call_observer(observer)
|
super().set_llm_call_observer(observer)
|
||||||
self._primary.set_llm_call_observer(observer)
|
self._primary.set_llm_call_observer(observer)
|
||||||
|
|
||||||
|
@property
|
||||||
|
def supports_progress_deltas(self) -> bool:
|
||||||
|
return bool(getattr(self._primary, "supports_progress_deltas", False))
|
||||||
|
|
||||||
def can_resume_conversation_state(
|
def can_resume_conversation_state(
|
||||||
self,
|
self,
|
||||||
state: ProviderConversationState,
|
state: ProviderConversationState,
|
||||||
|
|||||||
@@ -44,6 +44,8 @@ _COMPACTION_RETAINED_CHAR_BUDGET = 256_000
|
|||||||
class OpenAICodexProvider(LLMProvider):
|
class OpenAICodexProvider(LLMProvider):
|
||||||
"""Use Codex OAuth to call the Responses API."""
|
"""Use Codex OAuth to call the Responses API."""
|
||||||
|
|
||||||
|
supports_progress_deltas = True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
default_model: str = "openai-codex/gpt-5.6-sol",
|
default_model: str = "openai-codex/gpt-5.6-sol",
|
||||||
|
|||||||
@@ -63,6 +63,8 @@ def _is_named_x_search_tool(value: object) -> bool:
|
|||||||
class XAIGrokProvider(LLMProvider):
|
class XAIGrokProvider(LLMProvider):
|
||||||
"""Call xAI's subscription proxy and expose supported hosted tools."""
|
"""Call xAI's subscription proxy and expose supported hosted tools."""
|
||||||
|
|
||||||
|
supports_progress_deltas = True
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
default_model: str = DEFAULT_XAI_GROK_MODEL,
|
default_model: str = DEFAULT_XAI_GROK_MODEL,
|
||||||
|
|||||||
+25
-39
@@ -82,15 +82,6 @@ def _json_object(value: object) -> dict[str, Any]:
|
|||||||
return cast(dict[str, Any], value)
|
return cast(dict[str, Any], value)
|
||||||
|
|
||||||
|
|
||||||
def _archive_offset(data: dict[str, Any]) -> int:
|
|
||||||
"""Read the Memory archive watermark across the field-name migration."""
|
|
||||||
for key in ("last_archived", "last_consolidated"):
|
|
||||||
offset = cast(object, data.get(key))
|
|
||||||
if isinstance(offset, int) and not isinstance(offset, bool):
|
|
||||||
return offset
|
|
||||||
return 0
|
|
||||||
|
|
||||||
|
|
||||||
# TODO(0.3.2): Remove the write_stdin replay migration after 0.3.1.
|
# TODO(0.3.2): Remove the write_stdin replay migration after 0.3.1.
|
||||||
def _migrate_legacy_exec_arguments(container: dict[str, Any]) -> bool:
|
def _migrate_legacy_exec_arguments(container: dict[str, Any]) -> bool:
|
||||||
raw_arguments = cast(object, container.get("arguments"))
|
raw_arguments = cast(object, container.get("arguments"))
|
||||||
@@ -286,10 +277,7 @@ class Session:
|
|||||||
created_at: datetime = field(default_factory=datetime.now)
|
created_at: datetime = field(default_factory=datetime.now)
|
||||||
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)
|
||||||
# Legacy storage name for the Memory ingestion watermark. New code should
|
last_consolidated: int = 0 # Number of messages already consolidated to files
|
||||||
# use ``last_archived`` so this progress is not confused with model-context
|
|
||||||
# compaction. Keep the field while persisted sessions and SDK callers migrate.
|
|
||||||
last_consolidated: int = 0
|
|
||||||
provider_state: ProviderConversationState | None = field(default=None, repr=False)
|
provider_state: ProviderConversationState | None = field(default=None, repr=False)
|
||||||
policy: SessionPolicy = field(default_factory=SessionPolicy, repr=False, compare=False)
|
policy: SessionPolicy = field(default_factory=SessionPolicy, repr=False, compare=False)
|
||||||
|
|
||||||
@@ -307,15 +295,6 @@ class Session:
|
|||||||
):
|
):
|
||||||
self.last_consolidated = 0
|
self.last_consolidated = 0
|
||||||
|
|
||||||
@property
|
|
||||||
def last_archived(self) -> int:
|
|
||||||
"""Number of transcript messages already written to the Memory journal."""
|
|
||||||
return self.last_consolidated
|
|
||||||
|
|
||||||
@last_archived.setter
|
|
||||||
def last_archived(self, value: int) -> None:
|
|
||||||
self.last_consolidated = value
|
|
||||||
|
|
||||||
def add_message(self, role: str, content: str, **kwargs: Any) -> None:
|
def add_message(self, role: str, content: str, **kwargs: Any) -> None:
|
||||||
"""Add a message to the session."""
|
"""Add a message to the session."""
|
||||||
msg = {
|
msg = {
|
||||||
@@ -340,9 +319,9 @@ class Session:
|
|||||||
A positive ``max_messages`` applies an explicit caller-owned count
|
A positive ``max_messages`` applies an explicit caller-owned count
|
||||||
limit. The normal model path relies on ``max_tokens`` instead.
|
limit. The normal model path relies on ``max_tokens`` instead.
|
||||||
"""
|
"""
|
||||||
replay_start = self.last_archived
|
replay_start = self.last_consolidated
|
||||||
if replay_start:
|
if replay_start:
|
||||||
# ``last_archived`` is archive progress, not a replay boundary.
|
# ``last_consolidated`` is archive progress, not a replay boundary.
|
||||||
# Keep a small raw suffix for continuity, extending back to the user
|
# Keep a small raw suffix for continuity, extending back to the user
|
||||||
# that started an assistant/tool sequence when necessary.
|
# that started an assistant/tool sequence when necessary.
|
||||||
recent_start = recent_message_start_index(
|
recent_start = recent_message_start_index(
|
||||||
@@ -356,8 +335,8 @@ class Session:
|
|||||||
if max_messages <= 0:
|
if max_messages <= 0:
|
||||||
start_idx = 0
|
start_idx = 0
|
||||||
else:
|
else:
|
||||||
unarchived_count = len(self.messages) - self.last_archived
|
unarchived_count = len(self.messages) - self.last_consolidated
|
||||||
if replay_start < self.last_archived and unarchived_count < max_messages:
|
if replay_start < self.last_consolidated and unarchived_count < max_messages:
|
||||||
# The archived replay suffix can exceed the nominal count when one
|
# The archived replay suffix can exceed the nominal count when one
|
||||||
# tool-heavy turn spans the boundary. Preserve that complete turn.
|
# tool-heavy turn spans the boundary. Preserve that complete turn.
|
||||||
start_idx = 0
|
start_idx = 0
|
||||||
@@ -480,7 +459,7 @@ class Session:
|
|||||||
def clear(self) -> None:
|
def clear(self) -> None:
|
||||||
"""Clear all messages and reset session to initial state."""
|
"""Clear all messages and reset session to initial state."""
|
||||||
self.messages = []
|
self.messages = []
|
||||||
self.last_archived = 0
|
self.last_consolidated = 0
|
||||||
self.provider_state = None
|
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)
|
||||||
@@ -495,11 +474,11 @@ class Session:
|
|||||||
|
|
||||||
Returns a RetentionResult with dropped messages and how many of those
|
Returns a RetentionResult with dropped messages and how many of those
|
||||||
were in the already-consolidated prefix. This method mutates
|
were in the already-consolidated prefix. This method mutates
|
||||||
self.messages and self.last_archived in place.
|
self.messages and self.last_consolidated in place.
|
||||||
"""
|
"""
|
||||||
if max_messages <= 0:
|
if max_messages <= 0:
|
||||||
dropped = list(self.messages)
|
dropped = list(self.messages)
|
||||||
lc = self.last_archived
|
lc = self.last_consolidated
|
||||||
self.clear()
|
self.clear()
|
||||||
return RetentionResult(
|
return RetentionResult(
|
||||||
dropped=dropped,
|
dropped=dropped,
|
||||||
@@ -512,7 +491,7 @@ class Session:
|
|||||||
)
|
)
|
||||||
|
|
||||||
original = list(self.messages)
|
original = list(self.messages)
|
||||||
before_lc = self.last_archived
|
before_lc = self.last_consolidated
|
||||||
|
|
||||||
start_idx = max(0, len(self.messages) - max_messages)
|
start_idx = max(0, len(self.messages) - max_messages)
|
||||||
if extend_to_user:
|
if extend_to_user:
|
||||||
@@ -572,7 +551,7 @@ class Session:
|
|||||||
if i < before_lc and id(m) not in retained_ids
|
if i < before_lc and id(m) not in retained_ids
|
||||||
)
|
)
|
||||||
|
|
||||||
# New last_archived = count of retained messages that were inside
|
# New last_consolidated = count of retained messages that were inside
|
||||||
# the old consolidated prefix.
|
# the old consolidated prefix.
|
||||||
new_lc = sum(
|
new_lc = sum(
|
||||||
1 for i, m in enumerate(original)
|
1 for i, m in enumerate(original)
|
||||||
@@ -580,7 +559,7 @@ class Session:
|
|||||||
)
|
)
|
||||||
|
|
||||||
self.messages = retained
|
self.messages = retained
|
||||||
self.last_archived = new_lc
|
self.last_consolidated = new_lc
|
||||||
if dropped:
|
if dropped:
|
||||||
self.provider_state = None
|
self.provider_state = None
|
||||||
self.updated_at = datetime.now()
|
self.updated_at = datetime.now()
|
||||||
@@ -1188,7 +1167,12 @@ class JsonlSessionStore:
|
|||||||
if isinstance(updated_at_value, str) and updated_at_value
|
if isinstance(updated_at_value, str) and updated_at_value
|
||||||
else None
|
else None
|
||||||
)
|
)
|
||||||
last_consolidated = _archive_offset(data)
|
offset = cast(object, data.get("last_consolidated", 0))
|
||||||
|
last_consolidated = (
|
||||||
|
offset
|
||||||
|
if isinstance(offset, int) and not isinstance(offset, bool)
|
||||||
|
else 0
|
||||||
|
)
|
||||||
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
|
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
|
||||||
provider_state = ProviderConversationState.from_private_record(
|
provider_state = ProviderConversationState.from_private_record(
|
||||||
data.get("state")
|
data.get("state")
|
||||||
@@ -1270,7 +1254,12 @@ class JsonlSessionStore:
|
|||||||
if isinstance(updated_at_value, str) and updated_at_value:
|
if isinstance(updated_at_value, str) and updated_at_value:
|
||||||
with suppress(ValueError):
|
with suppress(ValueError):
|
||||||
updated_at = datetime.fromisoformat(updated_at_value)
|
updated_at = datetime.fromisoformat(updated_at_value)
|
||||||
last_consolidated = _archive_offset(data)
|
offset = cast(object, data.get("last_consolidated", 0))
|
||||||
|
last_consolidated = (
|
||||||
|
offset
|
||||||
|
if isinstance(offset, int) and not isinstance(offset, bool)
|
||||||
|
else 0
|
||||||
|
)
|
||||||
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
|
elif record_type == _PROVIDER_STATE_RECORD_TYPE:
|
||||||
candidate = ProviderConversationState.from_private_record(
|
candidate = ProviderConversationState.from_private_record(
|
||||||
data.get("state")
|
data.get("state")
|
||||||
@@ -1430,9 +1419,6 @@ class JsonlSessionStore:
|
|||||||
"created_at": session.created_at.isoformat(),
|
"created_at": session.created_at.isoformat(),
|
||||||
"updated_at": session.updated_at.isoformat(),
|
"updated_at": session.updated_at.isoformat(),
|
||||||
"metadata": session.metadata,
|
"metadata": session.metadata,
|
||||||
"last_archived": session.last_archived,
|
|
||||||
# Keep old nanobot releases able to read sessions written
|
|
||||||
# during the field-name migration.
|
|
||||||
"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")
|
||||||
@@ -2025,8 +2011,8 @@ class SessionManager:
|
|||||||
for key in _FORK_VOLATILE_METADATA_KEYS:
|
for key in _FORK_VOLATILE_METADATA_KEYS:
|
||||||
metadata.pop(key, None)
|
metadata.pop(key, None)
|
||||||
|
|
||||||
last_consolidated = min(source.last_archived, len(copied))
|
last_consolidated = min(source.last_consolidated, len(copied))
|
||||||
if source.last_archived > len(copied):
|
if source.last_consolidated > len(copied):
|
||||||
metadata.pop("_last_summary", None)
|
metadata.pop("_last_summary", None)
|
||||||
last_consolidated = 0
|
last_consolidated = 0
|
||||||
|
|
||||||
|
|||||||
@@ -6,8 +6,6 @@ Use [skip] unless a fact meets all SNIP criteria:
|
|||||||
- Important: prevents rework or captures preferences / rules
|
- Important: prevents rework or captures preferences / rules
|
||||||
- Persistent: still relevant after 2 weeks
|
- Persistent: still relevant after 2 weeks
|
||||||
|
|
||||||
Also preserve a compact working-state handoff even when it is not Persistent: the active objective, current status, completed steps, unresolved blockers, next action, and exact identifiers needed to continue without rework. Mark these facts [ephemeral].
|
|
||||||
|
|
||||||
Format each fact as:
|
Format each fact as:
|
||||||
- [mark] fact content
|
- [mark] fact content
|
||||||
|
|
||||||
|
|||||||
@@ -44,7 +44,7 @@ def session_context_payload(session: Session) -> dict[str, Any]:
|
|||||||
"schema_version": 1,
|
"schema_version": 1,
|
||||||
"session_key": session.key,
|
"session_key": session.key,
|
||||||
"total_messages": len(session.messages),
|
"total_messages": len(session.messages),
|
||||||
"archived_messages": min(session.last_archived, len(session.messages)),
|
"archived_messages": min(session.last_consolidated, len(session.messages)),
|
||||||
"replay_messages": len(replay),
|
"replay_messages": len(replay),
|
||||||
"estimated_replay_tokens": replay_tokens,
|
"estimated_replay_tokens": replay_tokens,
|
||||||
"estimated_summary_tokens": summary_tokens,
|
"estimated_summary_tokens": summary_tokens,
|
||||||
|
|||||||
@@ -88,11 +88,11 @@ def _make_fake_compact(
|
|||||||
state["count"] += 1
|
state["count"] += 1
|
||||||
session = loop.sessions.get_or_create(key)
|
session = loop.sessions.get_or_create(key)
|
||||||
|
|
||||||
tail = list(session.messages[session.last_archived:])
|
tail = list(session.messages[session.last_consolidated:])
|
||||||
if not tail:
|
if not tail:
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
return ""
|
return ""
|
||||||
archive_end = session.last_archived + len(tail)
|
archive_end = session.last_consolidated + len(tail)
|
||||||
archive_msgs = tail
|
archive_msgs = tail
|
||||||
|
|
||||||
last_active = session.updated_at
|
last_active = session.updated_at
|
||||||
@@ -109,7 +109,7 @@ def _make_fake_compact(
|
|||||||
"last_active": last_active.isoformat(),
|
"last_active": last_active.isoformat(),
|
||||||
}
|
}
|
||||||
|
|
||||||
session.last_archived = archive_end
|
session.last_consolidated = archive_end
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
return s
|
return s
|
||||||
|
|
||||||
@@ -399,12 +399,12 @@ class TestAutoCompact:
|
|||||||
await loop.aclose()
|
await loop.aclose()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_auto_compact_respects_last_archived(self, tmp_path):
|
async def test_auto_compact_respects_last_consolidated(self, tmp_path):
|
||||||
"""_archive should process only unarchived messages."""
|
"""_archive should only archive un-consolidated messages."""
|
||||||
loop = _make_loop(tmp_path, session_ttl_minutes=15)
|
loop = _make_loop(tmp_path, session_ttl_minutes=15)
|
||||||
session = loop.sessions.get_or_create("cli:test")
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
_add_turns(session, 14)
|
_add_turns(session, 14)
|
||||||
session.last_archived = 18
|
session.last_consolidated = 18
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
archived_messages = []
|
archived_messages = []
|
||||||
|
|||||||
@@ -16,7 +16,7 @@ def _runtime(_session: Session | None = None):
|
|||||||
def _make_session(
|
def _make_session(
|
||||||
key: str = "cli:test",
|
key: str = "cli:test",
|
||||||
messages: list | None = None,
|
messages: list | None = None,
|
||||||
last_archived: int = 0,
|
last_consolidated: int = 0,
|
||||||
updated_at: datetime | None = None,
|
updated_at: datetime | None = None,
|
||||||
metadata: dict | None = None,
|
metadata: dict | None = None,
|
||||||
) -> Session:
|
) -> Session:
|
||||||
@@ -25,8 +25,8 @@ def _make_session(
|
|||||||
key=key,
|
key=key,
|
||||||
messages=messages or [],
|
messages=messages or [],
|
||||||
metadata=metadata or {},
|
metadata=metadata or {},
|
||||||
|
last_consolidated=last_consolidated,
|
||||||
)
|
)
|
||||||
session.last_archived = last_archived
|
|
||||||
if updated_at is not None:
|
if updated_at is not None:
|
||||||
session.updated_at = updated_at
|
session.updated_at = updated_at
|
||||||
return session
|
return session
|
||||||
@@ -408,7 +408,7 @@ class TestCheckExpired:
|
|||||||
last_active = datetime(2026, 1, 1, 10, 0, 0)
|
last_active = datetime(2026, 1, 1, 10, 0, 0)
|
||||||
session = _make_session("cli:done", updated_at=last_active)
|
session = _make_session("cli:done", updated_at=last_active)
|
||||||
_add_turns(session, 2)
|
_add_turns(session, 2)
|
||||||
session.last_archived = len(session.messages)
|
session.last_consolidated = len(session.messages)
|
||||||
mock_sm.list_sessions.return_value = [
|
mock_sm.list_sessions.return_value = [
|
||||||
{"key": "cli:done", "updated_at": last_active.isoformat()},
|
{"key": "cli:done", "updated_at": last_active.isoformat()},
|
||||||
]
|
]
|
||||||
|
|||||||
@@ -0,0 +1,650 @@
|
|||||||
|
"""Test session management with cache-friendly message handling."""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from collections.abc import Coroutine
|
||||||
|
from pathlib import Path
|
||||||
|
from typing import Any
|
||||||
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
|
||||||
|
from nanobot.session.manager import Session, SessionManager
|
||||||
|
|
||||||
|
# Test constants
|
||||||
|
MEMORY_WINDOW = 50
|
||||||
|
KEEP_COUNT = MEMORY_WINDOW // 2 # 25
|
||||||
|
|
||||||
|
|
||||||
|
def create_session_with_messages(key: str, count: int, role: str = "user") -> Session:
|
||||||
|
"""Create a session and add the specified number of messages.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
key: Session identifier
|
||||||
|
count: Number of messages to add
|
||||||
|
role: Message role (default: "user")
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Session with the specified messages
|
||||||
|
"""
|
||||||
|
session = Session(key=key)
|
||||||
|
for i in range(count):
|
||||||
|
session.add_message(role, f"msg{i}")
|
||||||
|
return session
|
||||||
|
|
||||||
|
|
||||||
|
def assert_messages_content(messages: list, start_index: int, end_index: int) -> None:
|
||||||
|
"""Assert that messages contain expected content from start to end index.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
messages: List of message dictionaries
|
||||||
|
start_index: Expected first message index
|
||||||
|
end_index: Expected last message index
|
||||||
|
"""
|
||||||
|
assert len(messages) > 0
|
||||||
|
assert messages[0]["content"] == f"msg{start_index}"
|
||||||
|
assert messages[-1]["content"] == f"msg{end_index}"
|
||||||
|
|
||||||
|
|
||||||
|
def get_old_messages(session: Session, last_consolidated: int, keep_count: int) -> list:
|
||||||
|
"""Extract messages that would be consolidated using the standard slice logic.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
session: The session containing messages
|
||||||
|
last_consolidated: Index of last consolidated message
|
||||||
|
keep_count: Number of recent messages to keep
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
List of messages that would be consolidated
|
||||||
|
"""
|
||||||
|
return session.messages[last_consolidated:-keep_count]
|
||||||
|
|
||||||
|
|
||||||
|
class TestSessionLastConsolidated:
|
||||||
|
"""Test last_consolidated tracking to avoid duplicate processing."""
|
||||||
|
|
||||||
|
def test_initial_last_consolidated_zero(self) -> None:
|
||||||
|
"""Test that new session starts with last_consolidated=0."""
|
||||||
|
session = Session(key="test:initial")
|
||||||
|
assert session.last_consolidated == 0
|
||||||
|
|
||||||
|
def test_last_consolidated_persistence(self, tmp_path) -> None:
|
||||||
|
"""Test that last_consolidated persists across save/load."""
|
||||||
|
manager = SessionManager(Path(tmp_path))
|
||||||
|
session1 = create_session_with_messages("test:persist", 20)
|
||||||
|
session1.last_consolidated = 15
|
||||||
|
manager.save(session1)
|
||||||
|
|
||||||
|
session2 = manager.get_or_create("test:persist")
|
||||||
|
assert session2.last_consolidated == 15
|
||||||
|
assert len(session2.messages) == 20
|
||||||
|
|
||||||
|
def test_clear_resets_last_consolidated(self) -> None:
|
||||||
|
"""Test that clear() resets last_consolidated to 0."""
|
||||||
|
session = create_session_with_messages("test:clear", 10)
|
||||||
|
session.last_consolidated = 5
|
||||||
|
|
||||||
|
session.clear()
|
||||||
|
assert len(session.messages) == 0
|
||||||
|
assert session.last_consolidated == 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestSessionImmutableHistory:
|
||||||
|
"""Test Session message immutability for cache efficiency."""
|
||||||
|
|
||||||
|
def test_initial_state(self) -> None:
|
||||||
|
"""Test that new session has empty messages list."""
|
||||||
|
session = Session(key="test:initial")
|
||||||
|
assert len(session.messages) == 0
|
||||||
|
|
||||||
|
def test_add_messages_appends_only(self) -> None:
|
||||||
|
"""Test that adding messages only appends, never modifies."""
|
||||||
|
session = Session(key="test:preserve")
|
||||||
|
session.add_message("user", "msg1")
|
||||||
|
session.add_message("assistant", "resp1")
|
||||||
|
session.add_message("user", "msg2")
|
||||||
|
assert len(session.messages) == 3
|
||||||
|
assert session.messages[0]["content"] == "msg1"
|
||||||
|
|
||||||
|
def test_get_history_returns_most_recent(self) -> None:
|
||||||
|
"""Test get_history returns the most recent messages."""
|
||||||
|
session = Session(key="test:history")
|
||||||
|
for i in range(10):
|
||||||
|
session.add_message("user", f"msg{i}")
|
||||||
|
session.add_message("assistant", f"resp{i}")
|
||||||
|
|
||||||
|
history = session.get_history(max_messages=6)
|
||||||
|
assert len(history) == 6
|
||||||
|
assert history[0]["content"] == "msg7"
|
||||||
|
assert history[-1]["content"] == "resp9"
|
||||||
|
|
||||||
|
def test_get_history_with_all_messages(self) -> None:
|
||||||
|
"""Test get_history with max_messages larger than actual."""
|
||||||
|
session = create_session_with_messages("test:all", 5)
|
||||||
|
history = session.get_history(max_messages=100)
|
||||||
|
assert len(history) == 5
|
||||||
|
assert history[0]["content"] == "msg0"
|
||||||
|
|
||||||
|
def test_get_history_stable_for_same_session(self) -> None:
|
||||||
|
"""Test that get_history returns same content for same max_messages."""
|
||||||
|
session = create_session_with_messages("test:stable", 20)
|
||||||
|
history1 = session.get_history(max_messages=10)
|
||||||
|
history2 = session.get_history(max_messages=10)
|
||||||
|
assert history1 == history2
|
||||||
|
|
||||||
|
def test_messages_list_never_modified(self) -> None:
|
||||||
|
"""Test that messages list is never modified after creation."""
|
||||||
|
session = create_session_with_messages("test:immutable", 5)
|
||||||
|
original_len = len(session.messages)
|
||||||
|
|
||||||
|
session.get_history(max_messages=2)
|
||||||
|
assert len(session.messages) == original_len
|
||||||
|
|
||||||
|
for _ in range(10):
|
||||||
|
session.get_history(max_messages=3)
|
||||||
|
assert len(session.messages) == original_len
|
||||||
|
|
||||||
|
|
||||||
|
class TestSessionPersistence:
|
||||||
|
"""Test Session persistence and reload."""
|
||||||
|
|
||||||
|
@pytest.fixture
|
||||||
|
def temp_manager(self, tmp_path):
|
||||||
|
return SessionManager(Path(tmp_path))
|
||||||
|
|
||||||
|
def test_persistence_roundtrip(self, temp_manager):
|
||||||
|
"""Test that messages persist across save/load."""
|
||||||
|
session1 = create_session_with_messages("test:persistence", 20)
|
||||||
|
temp_manager.save(session1)
|
||||||
|
|
||||||
|
session2 = temp_manager.get_or_create("test:persistence")
|
||||||
|
assert len(session2.messages) == 20
|
||||||
|
assert session2.messages[0]["content"] == "msg0"
|
||||||
|
assert session2.messages[-1]["content"] == "msg19"
|
||||||
|
|
||||||
|
def test_get_history_after_reload(self, temp_manager):
|
||||||
|
"""Test that get_history works correctly after reload."""
|
||||||
|
session1 = create_session_with_messages("test:reload", 30)
|
||||||
|
temp_manager.save(session1)
|
||||||
|
|
||||||
|
session2 = temp_manager.get_or_create("test:reload")
|
||||||
|
history = session2.get_history(max_messages=10)
|
||||||
|
assert len(history) == 10
|
||||||
|
assert history[0]["content"] == "msg20"
|
||||||
|
assert history[-1]["content"] == "msg29"
|
||||||
|
|
||||||
|
def test_clear_resets_session(self, temp_manager):
|
||||||
|
"""Test that clear() properly resets session."""
|
||||||
|
session = create_session_with_messages("test:clear", 10)
|
||||||
|
assert len(session.messages) == 10
|
||||||
|
|
||||||
|
session.clear()
|
||||||
|
assert len(session.messages) == 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestConsolidationTriggerConditions:
|
||||||
|
"""Test consolidation trigger conditions and logic."""
|
||||||
|
|
||||||
|
def test_consolidation_needed_when_messages_exceed_window(self):
|
||||||
|
"""Test consolidation logic: should trigger when messages exceed the window."""
|
||||||
|
session = create_session_with_messages("test:trigger", 60)
|
||||||
|
|
||||||
|
total_messages = len(session.messages)
|
||||||
|
messages_to_process = total_messages - session.last_consolidated
|
||||||
|
|
||||||
|
assert total_messages > MEMORY_WINDOW
|
||||||
|
assert messages_to_process > 0
|
||||||
|
|
||||||
|
expected_consolidate_count = total_messages - KEEP_COUNT
|
||||||
|
assert expected_consolidate_count == 35
|
||||||
|
|
||||||
|
def test_consolidation_skipped_when_within_keep_count(self):
|
||||||
|
"""Test consolidation skipped when total messages <= keep_count."""
|
||||||
|
session = create_session_with_messages("test:skip", 20)
|
||||||
|
|
||||||
|
total_messages = len(session.messages)
|
||||||
|
assert total_messages <= KEEP_COUNT
|
||||||
|
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||||
|
assert len(old_messages) == 0
|
||||||
|
|
||||||
|
def test_consolidation_skipped_when_no_new_messages(self):
|
||||||
|
"""Test consolidation skipped when messages_to_process <= 0."""
|
||||||
|
session = create_session_with_messages("test:already_consolidated", 40)
|
||||||
|
session.last_consolidated = len(session.messages) - KEEP_COUNT # 15
|
||||||
|
|
||||||
|
# Add a few more messages
|
||||||
|
for i in range(40, 42):
|
||||||
|
session.add_message("user", f"msg{i}")
|
||||||
|
|
||||||
|
total_messages = len(session.messages)
|
||||||
|
messages_to_process = total_messages - session.last_consolidated
|
||||||
|
assert messages_to_process > 0
|
||||||
|
|
||||||
|
# Simulate last_consolidated catching up
|
||||||
|
session.last_consolidated = total_messages - KEEP_COUNT
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||||
|
assert len(old_messages) == 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestLastConsolidatedEdgeCases:
|
||||||
|
"""Test last_consolidated edge cases and data corruption scenarios."""
|
||||||
|
|
||||||
|
def test_last_consolidated_exceeds_message_count(self):
|
||||||
|
"""Test behavior when last_consolidated > len(messages) (data corruption)."""
|
||||||
|
session = create_session_with_messages("test:corruption", 10)
|
||||||
|
session.last_consolidated = 20
|
||||||
|
|
||||||
|
total_messages = len(session.messages)
|
||||||
|
messages_to_process = total_messages - session.last_consolidated
|
||||||
|
assert messages_to_process <= 0
|
||||||
|
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, 5)
|
||||||
|
assert len(old_messages) == 0
|
||||||
|
|
||||||
|
def test_last_consolidated_negative_value(self):
|
||||||
|
"""Test behavior with negative last_consolidated (invalid state)."""
|
||||||
|
session = create_session_with_messages("test:negative", 10)
|
||||||
|
session.last_consolidated = -5
|
||||||
|
|
||||||
|
keep_count = 3
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, keep_count)
|
||||||
|
|
||||||
|
# messages[-5:-3] with 10 messages gives indices 5,6
|
||||||
|
assert len(old_messages) == 2
|
||||||
|
assert old_messages[0]["content"] == "msg5"
|
||||||
|
assert old_messages[-1]["content"] == "msg6"
|
||||||
|
|
||||||
|
def test_messages_added_after_consolidation(self):
|
||||||
|
"""Test correct behavior when new messages arrive after consolidation."""
|
||||||
|
session = create_session_with_messages("test:new_messages", 40)
|
||||||
|
session.last_consolidated = len(session.messages) - KEEP_COUNT # 15
|
||||||
|
|
||||||
|
# Add new messages after consolidation
|
||||||
|
for i in range(40, 50):
|
||||||
|
session.add_message("user", f"msg{i}")
|
||||||
|
|
||||||
|
total_messages = len(session.messages)
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||||
|
expected_consolidate_count = total_messages - KEEP_COUNT - session.last_consolidated
|
||||||
|
|
||||||
|
assert len(old_messages) == expected_consolidate_count
|
||||||
|
assert_messages_content(old_messages, 15, 24)
|
||||||
|
|
||||||
|
def test_slice_behavior_when_indices_overlap(self):
|
||||||
|
"""Test slice behavior when last_consolidated >= total - keep_count."""
|
||||||
|
session = create_session_with_messages("test:overlap", 30)
|
||||||
|
session.last_consolidated = 12
|
||||||
|
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, 20)
|
||||||
|
assert len(old_messages) == 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestArchiveAllMode:
|
||||||
|
"""Test archive_all mode (used by /new command)."""
|
||||||
|
|
||||||
|
def test_archive_all_consolidates_everything(self):
|
||||||
|
"""Test archive_all=True consolidates all messages."""
|
||||||
|
session = create_session_with_messages("test:archive_all", 50)
|
||||||
|
|
||||||
|
archive_all = True
|
||||||
|
if archive_all:
|
||||||
|
old_messages = session.messages
|
||||||
|
assert len(old_messages) == 50
|
||||||
|
|
||||||
|
assert session.last_consolidated == 0
|
||||||
|
|
||||||
|
def test_archive_all_resets_last_consolidated(self):
|
||||||
|
"""Test that archive_all mode resets last_consolidated to 0."""
|
||||||
|
session = create_session_with_messages("test:reset", 40)
|
||||||
|
session.last_consolidated = 15
|
||||||
|
|
||||||
|
archive_all = True
|
||||||
|
if archive_all:
|
||||||
|
session.last_consolidated = 0
|
||||||
|
|
||||||
|
assert session.last_consolidated == 0
|
||||||
|
assert len(session.messages) == 40
|
||||||
|
|
||||||
|
def test_archive_all_vs_normal_consolidation(self):
|
||||||
|
"""Test difference between archive_all and normal consolidation."""
|
||||||
|
# Normal consolidation
|
||||||
|
session1 = create_session_with_messages("test:normal", 60)
|
||||||
|
session1.last_consolidated = len(session1.messages) - KEEP_COUNT
|
||||||
|
|
||||||
|
# archive_all mode
|
||||||
|
session2 = create_session_with_messages("test:all", 60)
|
||||||
|
session2.last_consolidated = 0
|
||||||
|
|
||||||
|
assert session1.last_consolidated == 35
|
||||||
|
assert len(session1.messages) == 60
|
||||||
|
assert session2.last_consolidated == 0
|
||||||
|
assert len(session2.messages) == 60
|
||||||
|
|
||||||
|
|
||||||
|
class TestCacheImmutability:
|
||||||
|
"""Test that consolidation doesn't modify session.messages (cache safety)."""
|
||||||
|
|
||||||
|
def test_consolidation_does_not_modify_messages_list(self):
|
||||||
|
"""Test that consolidation leaves messages list unchanged."""
|
||||||
|
session = create_session_with_messages("test:immutable", 50)
|
||||||
|
|
||||||
|
original_messages = session.messages.copy()
|
||||||
|
original_len = len(session.messages)
|
||||||
|
session.last_consolidated = original_len - KEEP_COUNT
|
||||||
|
|
||||||
|
assert len(session.messages) == original_len
|
||||||
|
assert session.messages == original_messages
|
||||||
|
|
||||||
|
def test_get_history_does_not_modify_messages(self):
|
||||||
|
"""Test that get_history doesn't modify messages list."""
|
||||||
|
session = create_session_with_messages("test:history_immutable", 40)
|
||||||
|
original_messages = [m.copy() for m in session.messages]
|
||||||
|
|
||||||
|
for _ in range(5):
|
||||||
|
history = session.get_history(max_messages=10)
|
||||||
|
assert len(history) == 10
|
||||||
|
|
||||||
|
assert len(session.messages) == 40
|
||||||
|
for i, msg in enumerate(session.messages):
|
||||||
|
assert msg["content"] == original_messages[i]["content"]
|
||||||
|
|
||||||
|
def test_consolidation_only_updates_last_consolidated(self):
|
||||||
|
"""Test that consolidation only updates last_consolidated field."""
|
||||||
|
session = create_session_with_messages("test:field_only", 60)
|
||||||
|
|
||||||
|
original_messages = session.messages.copy()
|
||||||
|
original_key = session.key
|
||||||
|
original_metadata = session.metadata.copy()
|
||||||
|
|
||||||
|
session.last_consolidated = len(session.messages) - KEEP_COUNT
|
||||||
|
|
||||||
|
assert session.messages == original_messages
|
||||||
|
assert session.key == original_key
|
||||||
|
assert session.metadata == original_metadata
|
||||||
|
assert session.last_consolidated == 35
|
||||||
|
|
||||||
|
|
||||||
|
class TestSliceLogic:
|
||||||
|
"""Test the slice logic: messages[last_consolidated:-keep_count]."""
|
||||||
|
|
||||||
|
def test_slice_extracts_correct_range(self):
|
||||||
|
"""Test that slice extracts the correct message range."""
|
||||||
|
session = create_session_with_messages("test:slice", 60)
|
||||||
|
|
||||||
|
old_messages = get_old_messages(session, 0, KEEP_COUNT)
|
||||||
|
|
||||||
|
assert len(old_messages) == 35
|
||||||
|
assert_messages_content(old_messages, 0, 34)
|
||||||
|
|
||||||
|
remaining = session.messages[-KEEP_COUNT:]
|
||||||
|
assert len(remaining) == 25
|
||||||
|
assert_messages_content(remaining, 35, 59)
|
||||||
|
|
||||||
|
def test_slice_with_partial_consolidation(self):
|
||||||
|
"""Test slice when some messages already consolidated."""
|
||||||
|
session = create_session_with_messages("test:partial", 70)
|
||||||
|
|
||||||
|
last_consolidated = 30
|
||||||
|
old_messages = get_old_messages(session, last_consolidated, KEEP_COUNT)
|
||||||
|
|
||||||
|
assert len(old_messages) == 15
|
||||||
|
assert_messages_content(old_messages, 30, 44)
|
||||||
|
|
||||||
|
def test_slice_with_various_keep_counts(self):
|
||||||
|
"""Test slice behavior with different keep_count values."""
|
||||||
|
session = create_session_with_messages("test:keep_counts", 50)
|
||||||
|
|
||||||
|
test_cases = [(10, 40), (20, 30), (30, 20), (40, 10)]
|
||||||
|
|
||||||
|
for keep_count, expected_count in test_cases:
|
||||||
|
old_messages = session.messages[0:-keep_count]
|
||||||
|
assert len(old_messages) == expected_count
|
||||||
|
|
||||||
|
def test_slice_when_keep_count_exceeds_messages(self):
|
||||||
|
"""Test slice when keep_count > len(messages)."""
|
||||||
|
session = create_session_with_messages("test:exceed", 10)
|
||||||
|
|
||||||
|
old_messages = session.messages[0:-20]
|
||||||
|
assert len(old_messages) == 0
|
||||||
|
|
||||||
|
|
||||||
|
class TestEmptyAndBoundarySessions:
|
||||||
|
"""Test empty sessions and boundary conditions."""
|
||||||
|
|
||||||
|
def test_empty_session_consolidation(self):
|
||||||
|
"""Test consolidation behavior with empty session."""
|
||||||
|
session = Session(key="test:empty")
|
||||||
|
|
||||||
|
assert len(session.messages) == 0
|
||||||
|
assert session.last_consolidated == 0
|
||||||
|
|
||||||
|
messages_to_process = len(session.messages) - session.last_consolidated
|
||||||
|
assert messages_to_process == 0
|
||||||
|
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||||
|
assert len(old_messages) == 0
|
||||||
|
|
||||||
|
def test_single_message_session(self):
|
||||||
|
"""Test consolidation with single message."""
|
||||||
|
session = Session(key="test:single")
|
||||||
|
session.add_message("user", "only message")
|
||||||
|
|
||||||
|
assert len(session.messages) == 1
|
||||||
|
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||||
|
assert len(old_messages) == 0
|
||||||
|
|
||||||
|
def test_exactly_keep_count_messages(self):
|
||||||
|
"""Test session with exactly keep_count messages."""
|
||||||
|
session = create_session_with_messages("test:exact", KEEP_COUNT)
|
||||||
|
|
||||||
|
assert len(session.messages) == KEEP_COUNT
|
||||||
|
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||||
|
assert len(old_messages) == 0
|
||||||
|
|
||||||
|
def test_just_over_keep_count(self):
|
||||||
|
"""Test session with one message over keep_count."""
|
||||||
|
session = create_session_with_messages("test:over", KEEP_COUNT + 1)
|
||||||
|
|
||||||
|
assert len(session.messages) == 26
|
||||||
|
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||||
|
assert len(old_messages) == 1
|
||||||
|
assert old_messages[0]["content"] == "msg0"
|
||||||
|
|
||||||
|
def test_very_large_session(self):
|
||||||
|
"""Test consolidation with very large message count."""
|
||||||
|
session = create_session_with_messages("test:large", 1000)
|
||||||
|
|
||||||
|
assert len(session.messages) == 1000
|
||||||
|
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||||
|
assert len(old_messages) == 975
|
||||||
|
assert_messages_content(old_messages, 0, 974)
|
||||||
|
|
||||||
|
remaining = session.messages[-KEEP_COUNT:]
|
||||||
|
assert len(remaining) == 25
|
||||||
|
assert_messages_content(remaining, 975, 999)
|
||||||
|
|
||||||
|
def test_session_with_gaps_in_consolidation(self):
|
||||||
|
"""Test session with potential gaps in consolidation history."""
|
||||||
|
session = create_session_with_messages("test:gaps", 50)
|
||||||
|
session.last_consolidated = 10
|
||||||
|
|
||||||
|
# Add more messages
|
||||||
|
for i in range(50, 60):
|
||||||
|
session.add_message("user", f"msg{i}")
|
||||||
|
|
||||||
|
old_messages = get_old_messages(session, session.last_consolidated, KEEP_COUNT)
|
||||||
|
|
||||||
|
expected_count = 60 - KEEP_COUNT - 10
|
||||||
|
assert len(old_messages) == expected_count
|
||||||
|
assert_messages_content(old_messages, 10, 34)
|
||||||
|
|
||||||
|
|
||||||
|
class TestNewCommandArchival:
|
||||||
|
"""Test /new archival behavior with the simplified consolidation flow."""
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _make_loop(tmp_path: Path):
|
||||||
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.providers.base import GenerationSettings, LLMResponse
|
||||||
|
|
||||||
|
bus = MessageBus()
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
provider.estimate_prompt_tokens.return_value = (10_000, "test")
|
||||||
|
provider.generation = GenerationSettings(max_tokens=100)
|
||||||
|
loop = AgentLoop(
|
||||||
|
bus=bus,
|
||||||
|
provider=provider,
|
||||||
|
workspace=tmp_path,
|
||||||
|
model="test-model",
|
||||||
|
context_window_tokens=1,
|
||||||
|
)
|
||||||
|
loop.provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="ok", tool_calls=[]))
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
return loop
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_new_clears_session_immediately_even_if_archive_fails(self, tmp_path: Path) -> None:
|
||||||
|
"""/new clears session immediately; archive is fire-and-forget."""
|
||||||
|
from nanobot.bus.events import InboundMessage
|
||||||
|
|
||||||
|
loop = self._make_loop(tmp_path)
|
||||||
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
|
for i in range(5):
|
||||||
|
session.add_message("user", f"msg{i}")
|
||||||
|
session.add_message("assistant", f"resp{i}")
|
||||||
|
loop.sessions.save(session)
|
||||||
|
|
||||||
|
call_count = 0
|
||||||
|
expected_runtime = loop.llm_runtime()
|
||||||
|
|
||||||
|
async def _failing_summarize(session, *, archive_end, runtime) -> None:
|
||||||
|
nonlocal call_count
|
||||||
|
assert runtime is expected_runtime
|
||||||
|
assert session.key == "cli:test"
|
||||||
|
assert archive_end == len(session.messages)
|
||||||
|
call_count += 1
|
||||||
|
|
||||||
|
loop.consolidator.archive_session = _failing_summarize # type: ignore[method-assign]
|
||||||
|
|
||||||
|
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||||
|
response = await loop._process_message(new_msg, runtime=expected_runtime)
|
||||||
|
|
||||||
|
assert response is not None
|
||||||
|
assert "new session started" in response.content.lower()
|
||||||
|
|
||||||
|
session_after = loop.sessions.get_or_create("cli:test")
|
||||||
|
assert len(session_after.messages) == 0
|
||||||
|
|
||||||
|
await loop.aclose()
|
||||||
|
assert call_count == 1
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_new_reuses_replay_prefix_and_archives_only_unconsolidated_messages(
|
||||||
|
self,
|
||||||
|
tmp_path: Path,
|
||||||
|
) -> None:
|
||||||
|
from nanobot.bus.events import InboundMessage
|
||||||
|
|
||||||
|
loop = self._make_loop(tmp_path)
|
||||||
|
loop.set_runtime_context_window(128_000)
|
||||||
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
|
for i in range(5):
|
||||||
|
session.add_message("user", f"msg{i}")
|
||||||
|
session.add_message("assistant", f"resp{i}")
|
||||||
|
session.last_consolidated = len(session.messages) - 2
|
||||||
|
ordinary_history = session.get_history()
|
||||||
|
assert [message["content"] for message in ordinary_history] == [
|
||||||
|
"msg1",
|
||||||
|
"resp1",
|
||||||
|
"msg2",
|
||||||
|
"resp2",
|
||||||
|
"msg3",
|
||||||
|
"resp3",
|
||||||
|
"msg4",
|
||||||
|
"resp4",
|
||||||
|
]
|
||||||
|
loop.sessions.save(session)
|
||||||
|
|
||||||
|
expected_runtime = loop.llm_runtime()
|
||||||
|
scheduled: list[Coroutine[Any, Any, object]] = []
|
||||||
|
loop.schedule_background = scheduled.append # type: ignore[method-assign]
|
||||||
|
|
||||||
|
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||||
|
response = await loop._process_message(new_msg, runtime=expected_runtime)
|
||||||
|
|
||||||
|
assert response is not None
|
||||||
|
assert "new session started" in response.content.lower()
|
||||||
|
|
||||||
|
assert len(scheduled) == 1
|
||||||
|
await scheduled[0]
|
||||||
|
await loop.aclose()
|
||||||
|
sent = loop.provider.chat_with_retry.call_args.kwargs["messages"]
|
||||||
|
assert sent[1:-1] == ordinary_history
|
||||||
|
assert "final 2 conversation messages" in sent[-1]["content"]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_new_clears_session_and_responds(self, tmp_path: Path) -> None:
|
||||||
|
from nanobot.bus.events import InboundMessage
|
||||||
|
|
||||||
|
loop = self._make_loop(tmp_path)
|
||||||
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
|
for i in range(3):
|
||||||
|
session.add_message("user", f"msg{i}")
|
||||||
|
session.add_message("assistant", f"resp{i}")
|
||||||
|
loop.sessions.save(session)
|
||||||
|
expected_runtime = loop.llm_runtime()
|
||||||
|
|
||||||
|
async def _ok_summarize(session, *, archive_end, runtime) -> str:
|
||||||
|
assert runtime is expected_runtime
|
||||||
|
assert session.key == "cli:test"
|
||||||
|
assert archive_end == len(session.messages)
|
||||||
|
return "Summary."
|
||||||
|
|
||||||
|
loop.consolidator.archive_session = _ok_summarize # type: ignore[method-assign]
|
||||||
|
|
||||||
|
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||||
|
response = await loop._process_message(new_msg, runtime=expected_runtime)
|
||||||
|
|
||||||
|
assert response is not None
|
||||||
|
assert "new session started" in response.content.lower()
|
||||||
|
assert loop.sessions.get_or_create("cli:test").messages == []
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_aclose_drains_background_tasks(self, tmp_path: Path) -> None:
|
||||||
|
"""aclose waits for background tasks to complete."""
|
||||||
|
from nanobot.bus.events import InboundMessage
|
||||||
|
|
||||||
|
loop = self._make_loop(tmp_path)
|
||||||
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
|
for i in range(3):
|
||||||
|
session.add_message("user", f"msg{i}")
|
||||||
|
session.add_message("assistant", f"resp{i}")
|
||||||
|
loop.sessions.save(session)
|
||||||
|
|
||||||
|
archived = asyncio.Event()
|
||||||
|
release_archive = asyncio.Event()
|
||||||
|
expected_runtime = loop.llm_runtime()
|
||||||
|
|
||||||
|
async def _slow_summarize(session, *, archive_end, runtime) -> str:
|
||||||
|
assert runtime is expected_runtime
|
||||||
|
assert session.key == "cli:test"
|
||||||
|
assert archive_end == len(session.messages)
|
||||||
|
await release_archive.wait()
|
||||||
|
archived.set()
|
||||||
|
return "Summary."
|
||||||
|
|
||||||
|
loop.consolidator.archive_session = _slow_summarize # type: ignore[method-assign]
|
||||||
|
|
||||||
|
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
||||||
|
await loop._process_message(new_msg, runtime=expected_runtime)
|
||||||
|
|
||||||
|
assert not archived.is_set()
|
||||||
|
release_archive.set()
|
||||||
|
await loop.aclose()
|
||||||
|
assert archived.is_set()
|
||||||
@@ -0,0 +1,112 @@
|
|||||||
|
"""Tests for configurable consolidation_ratio."""
|
||||||
|
|
||||||
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from pydantic import ValidationError
|
||||||
|
|
||||||
|
import nanobot.agent.memory as memory_module
|
||||||
|
from nanobot.agent.loop import AgentLoop
|
||||||
|
from nanobot.bus.queue import MessageBus
|
||||||
|
from nanobot.config.schema import AgentDefaults
|
||||||
|
from nanobot.providers.base import GenerationSettings, LLMResponse
|
||||||
|
|
||||||
|
|
||||||
|
def _make_loop(
|
||||||
|
tmp_path,
|
||||||
|
*,
|
||||||
|
estimated_tokens: int = 0,
|
||||||
|
context_window_tokens: int = 200,
|
||||||
|
consolidation_ratio: float = 0.5,
|
||||||
|
) -> AgentLoop:
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
provider.generation = GenerationSettings(max_tokens=0)
|
||||||
|
provider.estimate_prompt_tokens.return_value = (estimated_tokens, "test-counter")
|
||||||
|
_response = LLMResponse(content="ok", tool_calls=[])
|
||||||
|
provider.chat_with_retry = AsyncMock(return_value=_response)
|
||||||
|
provider.chat_stream_with_retry = AsyncMock(return_value=_response)
|
||||||
|
|
||||||
|
loop = AgentLoop(
|
||||||
|
bus=MessageBus(),
|
||||||
|
provider=provider,
|
||||||
|
workspace=tmp_path,
|
||||||
|
model="test-model",
|
||||||
|
context_window_tokens=context_window_tokens,
|
||||||
|
consolidation_ratio=consolidation_ratio,
|
||||||
|
)
|
||||||
|
loop.tools.get_definitions = MagicMock(return_value=[])
|
||||||
|
loop.consolidator._SAFETY_BUFFER = 0
|
||||||
|
return loop
|
||||||
|
|
||||||
|
|
||||||
|
def _session_with_turns(loop: AgentLoop, *, turns: int):
|
||||||
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
|
session.messages = []
|
||||||
|
for i in range(turns):
|
||||||
|
session.messages.append({"role": "user", "content": f"u{i}", "timestamp": f"2026-01-01T00:00:{i:02d}"})
|
||||||
|
session.messages.append({"role": "assistant", "content": f"a{i}", "timestamp": f"2026-01-01T00:01:{i:02d}"})
|
||||||
|
loop.sessions.save(session)
|
||||||
|
return session
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
@pytest.mark.parametrize(
|
||||||
|
("ratio", "context_window_tokens", "estimates", "expected_archives"),
|
||||||
|
[
|
||||||
|
(0.5, 200, [250, 90], 1),
|
||||||
|
(0.1, 1000, [1200, 800, 400, 50], 2),
|
||||||
|
(0.9, 200, [300, 175], 1),
|
||||||
|
],
|
||||||
|
)
|
||||||
|
async def test_consolidation_ratio_controls_target(
|
||||||
|
tmp_path,
|
||||||
|
monkeypatch,
|
||||||
|
ratio: float,
|
||||||
|
context_window_tokens: int,
|
||||||
|
estimates: list[int],
|
||||||
|
expected_archives: int,
|
||||||
|
) -> None:
|
||||||
|
loop = _make_loop(
|
||||||
|
tmp_path,
|
||||||
|
context_window_tokens=context_window_tokens,
|
||||||
|
consolidation_ratio=ratio,
|
||||||
|
)
|
||||||
|
loop.consolidator.archive_session = AsyncMock(return_value=True) # type: ignore[method-assign]
|
||||||
|
session = _session_with_turns(loop, turns=10)
|
||||||
|
|
||||||
|
remaining_estimates = list(estimates)
|
||||||
|
|
||||||
|
runtime = loop.llm_runtime()
|
||||||
|
|
||||||
|
def mock_estimate(_session, *, runtime):
|
||||||
|
return (remaining_estimates.pop(0), "test")
|
||||||
|
|
||||||
|
loop.consolidator.estimate_session_prompt_tokens = mock_estimate # type: ignore[method-assign]
|
||||||
|
monkeypatch.setattr(memory_module, "estimate_message_tokens", lambda _m: 100)
|
||||||
|
|
||||||
|
await loop.consolidator.maybe_consolidate_by_tokens(
|
||||||
|
session,
|
||||||
|
runtime=runtime,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert loop.consolidator.archive_session.await_count == expected_archives
|
||||||
|
|
||||||
|
|
||||||
|
def test_ratio_propagated_from_config_schema() -> None:
|
||||||
|
defaults = AgentDefaults()
|
||||||
|
assert defaults.consolidation_ratio == 0.5
|
||||||
|
|
||||||
|
defaults = AgentDefaults.model_validate({"consolidationRatio": 0.3})
|
||||||
|
assert defaults.consolidation_ratio == 0.3
|
||||||
|
|
||||||
|
dumped = defaults.model_dump(by_alias=True)
|
||||||
|
assert dumped["consolidationRatio"] == 0.3
|
||||||
|
|
||||||
|
|
||||||
|
def test_ratio_validation_rejects_out_of_range() -> None:
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
AgentDefaults(consolidation_ratio=0.05)
|
||||||
|
|
||||||
|
with pytest.raises(ValidationError):
|
||||||
|
AgentDefaults(consolidation_ratio=1.0)
|
||||||
@@ -232,19 +232,17 @@ class TestConsolidatorSummarize:
|
|||||||
|
|
||||||
|
|
||||||
class TestConsolidatorPromptContract:
|
class TestConsolidatorPromptContract:
|
||||||
def test_archive_prompt_preserves_working_state_with_memory_facts(self):
|
def test_archive_prompt_outputs_attribute_tags_without_missing_context_claims(self):
|
||||||
prompt = render_template("agent/consolidator_archive.md", strip=True, archive_count=4)
|
prompt = render_template("agent/consolidator_archive.md", strip=True, archive_count=4)
|
||||||
|
|
||||||
assert "SNIP" in prompt
|
assert "SNIP" in prompt
|
||||||
assert "final 4 conversation messages" in prompt
|
assert "final 4 conversation messages" in prompt
|
||||||
for mark in ("[permanent]", "[durable]", "[ephemeral]", "[correction]", "[skip]"):
|
for mark in ("[permanent]", "[durable]", "[ephemeral]", "[correction]", "[skip]"):
|
||||||
assert mark in prompt
|
assert mark in prompt
|
||||||
assert "working-state handoff" in prompt
|
assert "check context below" not in prompt.lower()
|
||||||
assert "exact identifiers needed to continue without rework" in prompt
|
|
||||||
assert "Do not output facts already present in the system prompt's Recent History" in prompt
|
assert "Do not output facts already present in the system prompt's Recent History" in prompt
|
||||||
assert "Do not mark something [skip] merely because it might already exist" in prompt
|
assert "Do not mark something [skip] merely because it might already exist" in prompt
|
||||||
|
|
||||||
|
|
||||||
class TestConsolidatorArchiveErrorHandling:
|
class TestConsolidatorArchiveErrorHandling:
|
||||||
"""archive() must fall back when the LLM does not complete its overview.
|
"""archive() must fall back when the LLM does not complete its overview.
|
||||||
|
|
||||||
@@ -344,7 +342,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
):
|
):
|
||||||
"""No consolidation when tokens are within budget."""
|
"""No consolidation when tokens are within budget."""
|
||||||
session = MagicMock()
|
session = MagicMock()
|
||||||
session.last_archived = 0
|
session.last_consolidated = 0
|
||||||
session.messages = [{"role": "user", "content": "hi"}]
|
session.messages = [{"role": "user", "content": "hi"}]
|
||||||
session.key = "test:key"
|
session.key = "test:key"
|
||||||
consolidator.sessions._session_cache[session.key] = session
|
consolidator.sessions._session_cache[session.key] = session
|
||||||
@@ -364,7 +362,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
with pytest.raises(RuntimeError, match="counter failed"):
|
with pytest.raises(RuntimeError, match="counter failed"):
|
||||||
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
|
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
|
||||||
|
|
||||||
async def test_estimate_uses_full_unarchived_tail(self, consolidator, runtime):
|
async def test_estimate_uses_full_unconsolidated_tail(self, consolidator, runtime):
|
||||||
"""Consolidation pressure must account for the full unarchived tail."""
|
"""Consolidation pressure must account for the full unarchived tail."""
|
||||||
session = Session(key="test:full-tail")
|
session = Session(key="test:full-tail")
|
||||||
for i in range(160):
|
for i in range(160):
|
||||||
@@ -387,7 +385,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
session = Session(key="test:archived-replay")
|
session = Session(key="test:archived-replay")
|
||||||
for i in range(10):
|
for i in range(10):
|
||||||
session.add_message("user", f"msg-{i}")
|
session.add_message("user", f"msg-{i}")
|
||||||
session.last_archived = len(session.messages)
|
session.last_consolidated = len(session.messages)
|
||||||
|
|
||||||
captured: dict[str, list[dict]] = {}
|
captured: dict[str, list[dict]] = {}
|
||||||
|
|
||||||
@@ -422,8 +420,8 @@ class TestConsolidatorTokenBudget:
|
|||||||
consolidator.estimate_session_prompt_tokens = MagicMock(
|
consolidator.estimate_session_prompt_tokens = MagicMock(
|
||||||
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
|
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
|
||||||
)
|
)
|
||||||
consolidator.pick_consolidation_boundary = MagicMock(return_value=50)
|
consolidator.pick_consolidation_boundary = MagicMock(return_value=(50, 800))
|
||||||
consolidator.archiver._build_messages = MagicMock(side_effect=_build_test_messages)
|
consolidator._build_messages = MagicMock(side_effect=_build_test_messages)
|
||||||
mock_provider.estimate_prompt_tokens.return_value = (100, "test-counter")
|
mock_provider.estimate_prompt_tokens.return_value = (100, "test-counter")
|
||||||
mock_provider.chat_with_retry.return_value = LLMResponse(
|
mock_provider.chat_with_retry.return_value = LLMResponse(
|
||||||
content="Token overflow summary.",
|
content="Token overflow summary.",
|
||||||
@@ -439,10 +437,10 @@ class TestConsolidatorTokenBudget:
|
|||||||
assert "final 50 conversation messages" in request["messages"][-1]["content"]
|
assert "final 50 conversation messages" in request["messages"][-1]["content"]
|
||||||
assert request["tools"] == []
|
assert request["tools"] == []
|
||||||
assert request["tool_choice"] == "none"
|
assert request["tool_choice"] == "none"
|
||||||
assert session.last_archived == 50
|
assert session.last_consolidated == 50
|
||||||
assert session.provider_state == _provider_state()
|
assert session.provider_state is None
|
||||||
|
|
||||||
async def test_raw_archive_fallback_advances_archive_watermark(
|
async def test_raw_archive_fallback_advances_last_consolidated(
|
||||||
self, consolidator, runtime
|
self, consolidator, runtime
|
||||||
):
|
):
|
||||||
"""When archive() falls back to raw-archive (LLM failed), the cursor
|
"""When archive() falls back to raw-archive (LLM failed), the cursor
|
||||||
@@ -450,12 +448,14 @@ class TestConsolidatorTokenBudget:
|
|||||||
on every subsequent maybe_consolidate_by_tokens() call, spamming
|
on every subsequent maybe_consolidate_by_tokens() call, spamming
|
||||||
duplicate [RAW] entries into history.jsonl."""
|
duplicate [RAW] entries into history.jsonl."""
|
||||||
consolidator._SAFETY_BUFFER = 0
|
consolidator._SAFETY_BUFFER = 0
|
||||||
session = Session(key="test:key")
|
session = MagicMock()
|
||||||
session.provider_state = _provider_state()
|
session.last_consolidated = 0
|
||||||
|
session.key = "test:key"
|
||||||
session.messages = [
|
session.messages = [
|
||||||
{"role": "user" if i in {0, 50} else "assistant", "content": f"m{i}"}
|
{"role": "user" if i in {0, 50} else "assistant", "content": f"m{i}"}
|
||||||
for i in range(70)
|
for i in range(70)
|
||||||
]
|
]
|
||||||
|
session.metadata = {}
|
||||||
consolidator.sessions._session_cache[session.key] = session
|
consolidator.sessions._session_cache[session.key] = session
|
||||||
consolidator.estimate_session_prompt_tokens = MagicMock(
|
consolidator.estimate_session_prompt_tokens = MagicMock(
|
||||||
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
|
side_effect=[(1200, "tiktoken"), (400, "tiktoken")]
|
||||||
@@ -467,10 +467,8 @@ class TestConsolidatorTokenBudget:
|
|||||||
|
|
||||||
consolidator.archive_session.assert_awaited_once()
|
consolidator.archive_session.assert_awaited_once()
|
||||||
# The chunk is considered "materialized" (as a raw-archive breadcrumb),
|
# The chunk is considered "materialized" (as a raw-archive breadcrumb),
|
||||||
# so the archive watermark must have moved past it without touching
|
# so last_consolidated must have moved past it.
|
||||||
# the provider-owned continuation state.
|
assert session.last_consolidated == 50
|
||||||
assert session.last_archived == 50
|
|
||||||
assert session.provider_state == _provider_state()
|
|
||||||
|
|
||||||
async def test_raw_archive_fallback_breaks_round_loop(
|
async def test_raw_archive_fallback_breaks_round_loop(
|
||||||
self, consolidator, runtime
|
self, consolidator, runtime
|
||||||
@@ -479,7 +477,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
same maybe_consolidate_by_tokens invocation — bail after one fallback."""
|
same maybe_consolidate_by_tokens invocation — bail after one fallback."""
|
||||||
consolidator._SAFETY_BUFFER = 0
|
consolidator._SAFETY_BUFFER = 0
|
||||||
session = MagicMock()
|
session = MagicMock()
|
||||||
session.last_archived = 0
|
session.last_consolidated = 0
|
||||||
session.key = "test:key"
|
session.key = "test:key"
|
||||||
session.messages = [
|
session.messages = [
|
||||||
{"role": "user" if i in {0, 20, 40, 60} else "assistant", "content": f"m{i}"}
|
{"role": "user" if i in {0, 20, 40, 60} else "assistant", "content": f"m{i}"}
|
||||||
@@ -495,7 +493,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
|
|
||||||
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
|
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
|
||||||
|
|
||||||
# The fixed policy archives at most one prefix per call.
|
# Exactly one fallback per call — not _MAX_CONSOLIDATION_ROUNDS.
|
||||||
assert consolidator.archive_session.await_count == 1
|
assert consolidator.archive_session.await_count == 1
|
||||||
|
|
||||||
async def test_boundary_respected_when_no_intermediate_user_turn(
|
async def test_boundary_respected_when_no_intermediate_user_turn(
|
||||||
@@ -504,7 +502,7 @@ class TestConsolidatorTokenBudget:
|
|||||||
"""When boundary points past a long tool chain, the full chunk is archived."""
|
"""When boundary points past a long tool chain, the full chunk is archived."""
|
||||||
consolidator._SAFETY_BUFFER = 0
|
consolidator._SAFETY_BUFFER = 0
|
||||||
session = MagicMock()
|
session = MagicMock()
|
||||||
session.last_archived = 0
|
session.last_consolidated = 0
|
||||||
session.key = "test:key"
|
session.key = "test:key"
|
||||||
session.messages = [
|
session.messages = [
|
||||||
{
|
{
|
||||||
@@ -522,8 +520,8 @@ class TestConsolidatorTokenBudget:
|
|||||||
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
|
await consolidator.maybe_consolidate_by_tokens(session, runtime=runtime)
|
||||||
|
|
||||||
consolidator.archive_session.assert_awaited_once()
|
consolidator.archive_session.assert_awaited_once()
|
||||||
# The fixed recent tail expands backward to the user at idx=61.
|
# pick_consolidation_boundary finds the only boundary at idx=61
|
||||||
assert session.last_archived == 61
|
assert session.last_consolidated == 61
|
||||||
|
|
||||||
|
|
||||||
class TestCompactIdleSession:
|
class TestCompactIdleSession:
|
||||||
@@ -577,8 +575,8 @@ class TestCompactIdleSession:
|
|||||||
reloaded = sessions.get_or_create("cli:test")
|
reloaded = sessions.get_or_create("cli:test")
|
||||||
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_archived == 40
|
assert reloaded.last_consolidated == 40
|
||||||
assert reloaded.provider_state == _provider_state()
|
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"
|
||||||
@@ -610,7 +608,7 @@ class TestCompactIdleSession:
|
|||||||
mock_provider.chat_with_retry.assert_awaited_once()
|
mock_provider.chat_with_retry.assert_awaited_once()
|
||||||
assert len(store.read_unprocessed_history(since_cursor=0)) == 1
|
assert len(store.read_unprocessed_history(since_cursor=0)) == 1
|
||||||
reloaded = sessions.get_or_create("cli:short")
|
reloaded = sessions.get_or_create("cli:short")
|
||||||
assert reloaded.last_archived == 2
|
assert reloaded.last_consolidated == 2
|
||||||
assert [message["content"] for message in reloaded.get_history()] == ["hello", "hi"]
|
assert [message["content"] for message in reloaded.get_history()] == ["hello", "hi"]
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -642,7 +640,7 @@ class TestCompactIdleSession:
|
|||||||
"second assistant",
|
"second assistant",
|
||||||
]
|
]
|
||||||
assert "final 2 conversation messages" in latest_messages[-1]["content"]
|
assert "final 2 conversation messages" in latest_messages[-1]["content"]
|
||||||
assert sessions.get_or_create("cli:incremental").last_archived == 4
|
assert sessions.get_or_create("cli:incremental").last_consolidated == 4
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_concurrent_append_remains_unarchived(
|
async def test_concurrent_append_remains_unarchived(
|
||||||
@@ -666,13 +664,13 @@ class TestCompactIdleSession:
|
|||||||
|
|
||||||
reloaded = sessions.get_or_create("cli:concurrent")
|
reloaded = sessions.get_or_create("cli:concurrent")
|
||||||
assert len(reloaded.messages) == 4
|
assert len(reloaded.messages) == 4
|
||||||
assert reloaded.last_archived == 2
|
assert reloaded.last_consolidated == 2
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_summarizes_retained_suffix_not_just_dropped_prefix(
|
async def test_summarizes_retained_suffix_not_just_dropped_prefix(
|
||||||
self, real_consolidator, mock_provider, runtime
|
self, real_consolidator, mock_provider, runtime
|
||||||
):
|
):
|
||||||
"""idleCompact must summarize over the full unarchived tail, including
|
"""idleCompact must summarize over the full unconsolidated tail, including
|
||||||
the recent suffix it retains. Otherwise a late user correction / final
|
the recent suffix it retains. Otherwise a late user correction / final
|
||||||
result that lands in the kept suffix is excluded from the persisted
|
result that lands in the kept suffix is excluded from the persisted
|
||||||
summary, leaving a stale wrong conclusion in history. Regression for #4264."""
|
summary, leaving a stale wrong conclusion in history. Regression for #4264."""
|
||||||
@@ -707,7 +705,6 @@ class TestCompactIdleSession:
|
|||||||
mock_provider.chat_with_retry.side_effect = RuntimeError("LLM unavailable")
|
mock_provider.chat_with_retry.side_effect = RuntimeError("LLM unavailable")
|
||||||
sessions = real_consolidator.sessions
|
sessions = real_consolidator.sessions
|
||||||
session = sessions.get_or_create("cli:rawdrop")
|
session = sessions.get_or_create("cli:rawdrop")
|
||||||
session.provider_state = _provider_state()
|
|
||||||
for i in range(18):
|
for i in range(18):
|
||||||
session.add_message("user", f"user msg {i}")
|
session.add_message("user", f"user msg {i}")
|
||||||
session.add_message("assistant", f"assistant msg {i}")
|
session.add_message("assistant", f"assistant msg {i}")
|
||||||
@@ -726,7 +723,6 @@ class TestCompactIdleSession:
|
|||||||
reloaded = sessions.get_or_create("cli:rawdrop")
|
reloaded = sessions.get_or_create("cli:rawdrop")
|
||||||
assert len(reloaded.messages) == 38
|
assert len(reloaded.messages) == 38
|
||||||
assert reloaded.messages[-1]["content"] == "RETAINED_SUFFIX_marker"
|
assert reloaded.messages[-1]["content"] == "RETAINED_SUFFIX_marker"
|
||||||
assert reloaded.provider_state == _provider_state()
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_idle_compact_writes_session_key_to_history(
|
async def test_idle_compact_writes_session_key_to_history(
|
||||||
@@ -822,7 +818,7 @@ class TestCompactIdleSession:
|
|||||||
reloaded = sessions.get_or_create("cli:fail")
|
reloaded = sessions.get_or_create("cli:fail")
|
||||||
assert len(reloaded.messages) == 20
|
assert len(reloaded.messages) == 20
|
||||||
assert reloaded.messages[0]["content"] == "u0"
|
assert reloaded.messages[0]["content"] == "u0"
|
||||||
assert reloaded.last_archived == 20
|
assert reloaded.last_consolidated == 20
|
||||||
assert [m["content"] for m in reloaded.get_history(max_messages=20)] == [
|
assert [m["content"] for m in reloaded.get_history(max_messages=20)] == [
|
||||||
"u6",
|
"u6",
|
||||||
"a6",
|
"a6",
|
||||||
@@ -835,10 +831,10 @@ class TestCompactIdleSession:
|
|||||||
]
|
]
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_respects_last_archived(
|
async def test_respects_last_consolidated(
|
||||||
self, real_consolidator, mock_provider, runtime
|
self, real_consolidator, mock_provider, runtime
|
||||||
):
|
):
|
||||||
"""30 turns with last_archived=50 → only the unarchived tail is considered."""
|
"""30 turns with last_consolidated=50 → only unconsolidated tail considered."""
|
||||||
mock_provider.chat_with_retry.return_value = MagicMock(
|
mock_provider.chat_with_retry.return_value = MagicMock(
|
||||||
content="Tail summary.", finish_reason="stop"
|
content="Tail summary.", finish_reason="stop"
|
||||||
)
|
)
|
||||||
@@ -847,7 +843,7 @@ class TestCompactIdleSession:
|
|||||||
for i in range(30):
|
for i in range(30):
|
||||||
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}")
|
||||||
session.last_archived = 50 # Only 10 messages remain unarchived
|
session.last_consolidated = 50 # Only 10 messages unconsolidated
|
||||||
sessions.save(session)
|
sessions.save(session)
|
||||||
|
|
||||||
result = await real_consolidator.compact_idle_session(
|
result = await real_consolidator.compact_idle_session(
|
||||||
@@ -856,10 +852,10 @@ class TestCompactIdleSession:
|
|||||||
assert result == "Tail summary."
|
assert result == "Tail summary."
|
||||||
reloaded = sessions.get_or_create("cli:offset")
|
reloaded = sessions.get_or_create("cli:offset")
|
||||||
assert len(reloaded.messages) == 60
|
assert len(reloaded.messages) == 60
|
||||||
assert reloaded.last_archived == 60
|
assert reloaded.last_consolidated == 60
|
||||||
|
|
||||||
# Verify only the unarchived tail was processed:
|
# Verify only the unconsolidated tail was processed:
|
||||||
# All 10 unarchived messages (50-59) are archived exactly once.
|
# All 10 unconsolidated messages (50-59) are archived exactly once.
|
||||||
archived_call = mock_provider.chat_with_retry.call_args
|
archived_call = mock_provider.chat_with_retry.call_args
|
||||||
sent_messages = archived_call.kwargs["messages"]
|
sent_messages = archived_call.kwargs["messages"]
|
||||||
sent_content = [message.get("content") for message in sent_messages]
|
sent_content = [message.get("content") for message in sent_messages]
|
||||||
@@ -894,7 +890,7 @@ class TestCompactIdleSession:
|
|||||||
|
|
||||||
reloaded = sessions.get_or_create("cli:noncontiguous")
|
reloaded = sessions.get_or_create("cli:noncontiguous")
|
||||||
assert len(reloaded.messages) == 25
|
assert len(reloaded.messages) == 25
|
||||||
assert reloaded.last_archived == 25
|
assert reloaded.last_consolidated == 25
|
||||||
assert [m["content"] for m in reloaded.get_history(max_messages=25)] == [
|
assert [m["content"] for m in reloaded.get_history(max_messages=25)] == [
|
||||||
"user-14",
|
"user-14",
|
||||||
"assistant-00",
|
"assistant-00",
|
||||||
@@ -909,7 +905,7 @@ class TestCompactIdleSession:
|
|||||||
"assistant-09",
|
"assistant-09",
|
||||||
]
|
]
|
||||||
|
|
||||||
# #4264: idle compaction now summarizes the full unarchived tail, so
|
# #4264: idle compaction now summarizes the full unconsolidated tail, so
|
||||||
# the dropped head (user-00) and retained suffix (user-14 through
|
# the dropped head (user-00) and retained suffix (user-14 through
|
||||||
# assistant-09) are all summarized.
|
# assistant-09) are all summarized.
|
||||||
archived_call = mock_provider.chat_with_retry.call_args
|
archived_call = mock_provider.chat_with_retry.call_args
|
||||||
@@ -927,7 +923,7 @@ class TestCompactIdleSession:
|
|||||||
runtime,
|
runtime,
|
||||||
):
|
):
|
||||||
tools = [{"type": "function", "function": {"name": "lookup"}}]
|
tools = [{"type": "function", "function": {"name": "lookup"}}]
|
||||||
real_consolidator.archiver._get_tool_definitions.return_value = tools
|
real_consolidator._get_tool_definitions.return_value = tools
|
||||||
mock_provider.chat_with_retry.return_value = LLMResponse(
|
mock_provider.chat_with_retry.return_value = LLMResponse(
|
||||||
content="Overview from the temporary turn.",
|
content="Overview from the temporary turn.",
|
||||||
finish_reason="stop",
|
finish_reason="stop",
|
||||||
@@ -1001,7 +997,7 @@ class TestCompactIdleSession:
|
|||||||
assert len(entries) == 1
|
assert len(entries) == 1
|
||||||
assert entries[0]["content"].startswith("[RAW] ")
|
assert entries[0]["content"].startswith("[RAW] ")
|
||||||
assert "important answer" in entries[0]["content"]
|
assert "important answer" in entries[0]["content"]
|
||||||
assert sessions.get_or_create("cli:unexpected-tool").last_archived == 2
|
assert sessions.get_or_create("cli:unexpected-tool").last_consolidated == 2
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_empty_response_uses_raw_fallback(
|
async def test_empty_response_uses_raw_fallback(
|
||||||
@@ -1031,7 +1027,7 @@ class TestCompactIdleSession:
|
|||||||
assert len(entries) == 1
|
assert len(entries) == 1
|
||||||
assert entries[0]["content"].startswith("[RAW] ")
|
assert entries[0]["content"].startswith("[RAW] ")
|
||||||
assert "important answer" in entries[0]["content"]
|
assert "important answer" in entries[0]["content"]
|
||||||
assert sessions.get_or_create("cli:empty-summary").last_archived == 2
|
assert sessions.get_or_create("cli:empty-summary").last_consolidated == 2
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_oversized_prefix_raw_archives_without_flattened_llm_retry(
|
async def test_oversized_prefix_raw_archives_without_flattened_llm_retry(
|
||||||
@@ -1057,7 +1053,7 @@ class TestCompactIdleSession:
|
|||||||
entries = store.read_unprocessed_history(since_cursor=0)
|
entries = store.read_unprocessed_history(since_cursor=0)
|
||||||
assert len(entries) == 1
|
assert len(entries) == 1
|
||||||
assert entries[0]["content"].startswith("[RAW] ")
|
assert entries[0]["content"].startswith("[RAW] ")
|
||||||
assert sessions.get_or_create("sdk:oversized").last_archived == 1
|
assert sessions.get_or_create("sdk:oversized").last_consolidated == 1
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_incremental_scope_counts_only_model_visible_messages(
|
async def test_incremental_scope_counts_only_model_visible_messages(
|
||||||
@@ -1074,7 +1070,7 @@ class TestCompactIdleSession:
|
|||||||
session = sessions.get_or_create("cli:commands")
|
session = sessions.get_or_create("cli:commands")
|
||||||
session.add_message("user", "already archived user")
|
session.add_message("user", "already archived user")
|
||||||
session.add_message("assistant", "already archived answer")
|
session.add_message("assistant", "already archived answer")
|
||||||
session.last_archived = 2
|
session.last_consolidated = 2
|
||||||
session.add_message("user", "/status", _command=True)
|
session.add_message("user", "/status", _command=True)
|
||||||
session.add_message("assistant", "status output", _command=True)
|
session.add_message("assistant", "status output", _command=True)
|
||||||
session.add_message("user", "new user")
|
session.add_message("user", "new user")
|
||||||
@@ -1282,7 +1278,7 @@ class TestConsolidatorSessionRefresh:
|
|||||||
|
|
||||||
session_after = sessions.get_or_create("cli:test")
|
session_after = sessions.get_or_create("cli:test")
|
||||||
assert len(session_after.messages) == 40
|
assert len(session_after.messages) == 40
|
||||||
assert session_after.last_archived == 40
|
assert session_after.last_consolidated == 40
|
||||||
assert len(session_after.get_history(max_messages=40)) == 8
|
assert len(session_after.get_history(max_messages=40)) == 8
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,47 +0,0 @@
|
|||||||
from __future__ import annotations
|
|
||||||
|
|
||||||
from types import SimpleNamespace
|
|
||||||
from unittest.mock import MagicMock
|
|
||||||
|
|
||||||
import pytest
|
|
||||||
|
|
||||||
|
|
||||||
def _provider() -> MagicMock:
|
|
||||||
provider = MagicMock()
|
|
||||||
provider.get_default_model.return_value = "test-model"
|
|
||||||
provider.generation = SimpleNamespace(
|
|
||||||
max_tokens=4096,
|
|
||||||
temperature=0.1,
|
|
||||||
reasoning_effort=None,
|
|
||||||
)
|
|
||||||
return provider
|
|
||||||
|
|
||||||
|
|
||||||
def test_request_concurrency_is_unlimited_by_default(
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
loop_factory,
|
|
||||||
) -> None:
|
|
||||||
monkeypatch.delenv("NANOBOT_MAX_CONCURRENT_REQUESTS", raising=False)
|
|
||||||
|
|
||||||
loop = loop_factory(provider=_provider(), patch_deps=True)
|
|
||||||
|
|
||||||
assert loop._concurrency_gate is None
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_positive_request_concurrency_keeps_explicit_cap(
|
|
||||||
monkeypatch: pytest.MonkeyPatch,
|
|
||||||
loop_factory,
|
|
||||||
) -> None:
|
|
||||||
monkeypatch.setenv("NANOBOT_MAX_CONCURRENT_REQUESTS", "2")
|
|
||||||
loop = loop_factory(provider=_provider(), patch_deps=True)
|
|
||||||
gate = loop._concurrency_gate
|
|
||||||
|
|
||||||
assert gate is not None
|
|
||||||
for _ in range(2):
|
|
||||||
await gate.acquire()
|
|
||||||
try:
|
|
||||||
assert gate.locked()
|
|
||||||
finally:
|
|
||||||
for _ in range(2):
|
|
||||||
gate.release()
|
|
||||||
@@ -2,6 +2,7 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
|
import nanobot.agent.memory as memory_module
|
||||||
from nanobot.agent.loop import AgentLoop
|
from nanobot.agent.loop import AgentLoop
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
from nanobot.providers.base import LLMResponse
|
from nanobot.providers.base import LLMResponse
|
||||||
@@ -40,16 +41,17 @@ async def test_prompt_below_threshold_does_not_consolidate(tmp_path) -> None:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_prompt_above_threshold_triggers_consolidation(tmp_path) -> None:
|
async def test_prompt_above_threshold_triggers_consolidation(tmp_path, monkeypatch) -> None:
|
||||||
loop = _make_loop(tmp_path, estimated_tokens=1000, context_window_tokens=200)
|
loop = _make_loop(tmp_path, estimated_tokens=1000, context_window_tokens=200)
|
||||||
loop.consolidator.archive_session = AsyncMock(return_value=True) # type: ignore[method-assign]
|
loop.consolidator.archive_session = AsyncMock(return_value=True) # type: ignore[method-assign]
|
||||||
session = loop.sessions.get_or_create("cli:test")
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
session.messages = [
|
session.messages = [
|
||||||
{"role": role, "content": f"{role[0]}{turn}"}
|
{"role": "user", "content": "u1", "timestamp": "2026-01-01T00:00:00"},
|
||||||
for turn in range(10)
|
{"role": "assistant", "content": "a1", "timestamp": "2026-01-01T00:00:01"},
|
||||||
for role in ("user", "assistant")
|
{"role": "user", "content": "u2", "timestamp": "2026-01-01T00:00:02"},
|
||||||
]
|
]
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
monkeypatch.setattr(memory_module, "estimate_message_tokens", lambda _message: 500)
|
||||||
|
|
||||||
await loop.process_direct("hello", session_key="cli:test")
|
await loop.process_direct("hello", session_key="cli:test")
|
||||||
|
|
||||||
@@ -57,18 +59,23 @@ async def test_prompt_above_threshold_triggers_consolidation(tmp_path) -> None:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_prompt_above_threshold_uses_fixed_recent_tail(tmp_path) -> None:
|
async def test_prompt_above_threshold_archives_until_next_user_boundary(tmp_path, monkeypatch) -> None:
|
||||||
loop = _make_loop(tmp_path, estimated_tokens=1000, context_window_tokens=200)
|
loop = _make_loop(tmp_path, estimated_tokens=1000, context_window_tokens=200)
|
||||||
loop.consolidator.archive_session = AsyncMock(return_value=True) # type: ignore[method-assign]
|
loop.consolidator.archive_session = AsyncMock(return_value=True) # type: ignore[method-assign]
|
||||||
|
|
||||||
session = loop.sessions.get_or_create("cli:test")
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
session.messages = [
|
session.messages = [
|
||||||
{"role": role, "content": f"{role[0]}{turn}"}
|
{"role": "user", "content": "u1", "timestamp": "2026-01-01T00:00:00"},
|
||||||
for turn in range(10)
|
{"role": "assistant", "content": "a1", "timestamp": "2026-01-01T00:00:01"},
|
||||||
for role in ("user", "assistant")
|
{"role": "user", "content": "u2", "timestamp": "2026-01-01T00:00:02"},
|
||||||
|
{"role": "assistant", "content": "a2", "timestamp": "2026-01-01T00:00:03"},
|
||||||
|
{"role": "user", "content": "u3", "timestamp": "2026-01-01T00:00:04"},
|
||||||
]
|
]
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
|
token_map = {"u1": 120, "a1": 120, "u2": 120, "a2": 120, "u3": 120}
|
||||||
|
monkeypatch.setattr(memory_module, "estimate_message_tokens", lambda message: token_map[message["content"]])
|
||||||
|
|
||||||
await loop.consolidator.maybe_consolidate_by_tokens(
|
await loop.consolidator.maybe_consolidate_by_tokens(
|
||||||
session,
|
session,
|
||||||
runtime=loop.llm_runtime(),
|
runtime=loop.llm_runtime(),
|
||||||
@@ -76,29 +83,112 @@ async def test_prompt_above_threshold_uses_fixed_recent_tail(tmp_path) -> None:
|
|||||||
|
|
||||||
archive_end = loop.consolidator.archive_session.await_args.kwargs["archive_end"]
|
archive_end = loop.consolidator.archive_session.await_args.kwargs["archive_end"]
|
||||||
archived_chunk = session.messages[:archive_end]
|
archived_chunk = session.messages[:archive_end]
|
||||||
assert [message["content"] for message in archived_chunk] == [
|
assert [message["content"] for message in archived_chunk] == ["u1", "a1", "u2", "a2"]
|
||||||
"u0", "a0", "u1", "a1", "u2", "a2", "u3", "a3", "u4", "a4", "u5", "a5",
|
assert session.last_consolidated == 4
|
||||||
]
|
|
||||||
assert session.last_archived == 12
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_consolidation_persists_summary_for_next_prepare_session(tmp_path) -> None:
|
async def test_consolidation_loops_until_target_met(tmp_path, monkeypatch) -> None:
|
||||||
|
"""Verify maybe_consolidate_by_tokens keeps looping until under threshold."""
|
||||||
|
loop = _make_loop(tmp_path, estimated_tokens=0, context_window_tokens=200)
|
||||||
|
loop.consolidator.archive_session = AsyncMock(return_value=True) # type: ignore[method-assign]
|
||||||
|
|
||||||
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
|
session.messages = [
|
||||||
|
{"role": "user", "content": "u1", "timestamp": "2026-01-01T00:00:00"},
|
||||||
|
{"role": "assistant", "content": "a1", "timestamp": "2026-01-01T00:00:01"},
|
||||||
|
{"role": "user", "content": "u2", "timestamp": "2026-01-01T00:00:02"},
|
||||||
|
{"role": "assistant", "content": "a2", "timestamp": "2026-01-01T00:00:03"},
|
||||||
|
{"role": "user", "content": "u3", "timestamp": "2026-01-01T00:00:04"},
|
||||||
|
{"role": "assistant", "content": "a3", "timestamp": "2026-01-01T00:00:05"},
|
||||||
|
{"role": "user", "content": "u4", "timestamp": "2026-01-01T00:00:06"},
|
||||||
|
]
|
||||||
|
loop.sessions.save(session)
|
||||||
|
|
||||||
|
call_count = [0]
|
||||||
|
def mock_estimate(_session, *, runtime):
|
||||||
|
call_count[0] += 1
|
||||||
|
if call_count[0] == 1:
|
||||||
|
return (500, "test")
|
||||||
|
if call_count[0] == 2:
|
||||||
|
return (300, "test")
|
||||||
|
return (80, "test")
|
||||||
|
|
||||||
|
loop.consolidator.estimate_session_prompt_tokens = mock_estimate # type: ignore[method-assign]
|
||||||
|
monkeypatch.setattr(memory_module, "estimate_message_tokens", lambda _m: 100)
|
||||||
|
|
||||||
|
await loop.consolidator.maybe_consolidate_by_tokens(
|
||||||
|
session,
|
||||||
|
runtime=loop.llm_runtime(),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert loop.consolidator.archive_session.await_count == 2
|
||||||
|
assert session.last_consolidated == 6
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_consolidation_continues_below_trigger_until_half_target(tmp_path, monkeypatch) -> None:
|
||||||
|
"""Once triggered, consolidation should continue until it drops below half threshold."""
|
||||||
|
loop = _make_loop(tmp_path, estimated_tokens=0, context_window_tokens=200)
|
||||||
|
loop.consolidator.archive_session = AsyncMock(return_value=True) # type: ignore[method-assign]
|
||||||
|
|
||||||
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
|
session.messages = [
|
||||||
|
{"role": "user", "content": "u1", "timestamp": "2026-01-01T00:00:00"},
|
||||||
|
{"role": "assistant", "content": "a1", "timestamp": "2026-01-01T00:00:01"},
|
||||||
|
{"role": "user", "content": "u2", "timestamp": "2026-01-01T00:00:02"},
|
||||||
|
{"role": "assistant", "content": "a2", "timestamp": "2026-01-01T00:00:03"},
|
||||||
|
{"role": "user", "content": "u3", "timestamp": "2026-01-01T00:00:04"},
|
||||||
|
{"role": "assistant", "content": "a3", "timestamp": "2026-01-01T00:00:05"},
|
||||||
|
{"role": "user", "content": "u4", "timestamp": "2026-01-01T00:00:06"},
|
||||||
|
]
|
||||||
|
loop.sessions.save(session)
|
||||||
|
|
||||||
|
call_count = [0]
|
||||||
|
|
||||||
|
def mock_estimate(_session, *, runtime):
|
||||||
|
call_count[0] += 1
|
||||||
|
if call_count[0] == 1:
|
||||||
|
return (500, "test")
|
||||||
|
if call_count[0] == 2:
|
||||||
|
return (150, "test")
|
||||||
|
return (80, "test")
|
||||||
|
|
||||||
|
loop.consolidator.estimate_session_prompt_tokens = mock_estimate # type: ignore[method-assign]
|
||||||
|
monkeypatch.setattr(memory_module, "estimate_message_tokens", lambda _m: 100)
|
||||||
|
|
||||||
|
await loop.consolidator.maybe_consolidate_by_tokens(
|
||||||
|
session,
|
||||||
|
runtime=loop.llm_runtime(),
|
||||||
|
)
|
||||||
|
|
||||||
|
assert loop.consolidator.archive_session.await_count == 2
|
||||||
|
assert session.last_consolidated == 6
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_consolidation_persists_summary_for_next_prepare_session(tmp_path, monkeypatch) -> None:
|
||||||
loop = _make_loop(tmp_path, estimated_tokens=0, context_window_tokens=200)
|
loop = _make_loop(tmp_path, estimated_tokens=0, context_window_tokens=200)
|
||||||
loop.consolidator.archive_session = AsyncMock(return_value="User discussed project status.") # type: ignore[method-assign]
|
loop.consolidator.archive_session = AsyncMock(return_value="User discussed project status.") # type: ignore[method-assign]
|
||||||
|
|
||||||
session = loop.sessions.get_or_create("cli:test")
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
session.messages = [
|
session.messages = [
|
||||||
{"role": role, "content": f"{role[0]}{turn}"}
|
{"role": "user", "content": "u1", "timestamp": "2026-01-01T00:00:00"},
|
||||||
for turn in range(5)
|
{"role": "assistant", "content": "a1", "timestamp": "2026-01-01T00:00:01"},
|
||||||
for role in ("user", "assistant")
|
{"role": "user", "content": "u2", "timestamp": "2026-01-01T00:00:02"},
|
||||||
]
|
]
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
|
||||||
|
call_count = [0]
|
||||||
|
|
||||||
def mock_estimate(_session, *, runtime):
|
def mock_estimate(_session, *, runtime):
|
||||||
|
call_count[0] += 1
|
||||||
|
if call_count[0] == 1:
|
||||||
return (500, "test")
|
return (500, "test")
|
||||||
|
return (80, "test")
|
||||||
|
|
||||||
loop.consolidator.estimate_session_prompt_tokens = mock_estimate # type: ignore[method-assign]
|
loop.consolidator.estimate_session_prompt_tokens = mock_estimate # type: ignore[method-assign]
|
||||||
|
monkeypatch.setattr(memory_module, "estimate_message_tokens", lambda _m: 150)
|
||||||
|
|
||||||
await loop.consolidator.maybe_consolidate_by_tokens(
|
await loop.consolidator.maybe_consolidate_by_tokens(
|
||||||
session,
|
session,
|
||||||
@@ -145,7 +235,7 @@ async def test_preflight_consolidation_receives_pending_summary(tmp_path) -> Non
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_preflight_consolidation_before_llm_call(tmp_path) -> None:
|
async def test_preflight_consolidation_before_llm_call(tmp_path, monkeypatch) -> None:
|
||||||
"""Verify preflight consolidation runs before the LLM call in process_direct."""
|
"""Verify preflight consolidation runs before the LLM call in process_direct."""
|
||||||
order: list[str] = []
|
order: list[str] = []
|
||||||
|
|
||||||
@@ -168,11 +258,13 @@ async def test_preflight_consolidation_before_llm_call(tmp_path) -> None:
|
|||||||
|
|
||||||
session = loop.sessions.get_or_create("cli:test")
|
session = loop.sessions.get_or_create("cli:test")
|
||||||
session.messages = [
|
session.messages = [
|
||||||
{"role": role, "content": f"{role[0]}{turn}"}
|
{"role": "user", "content": "u1", "timestamp": "2026-01-01T00:00:00"},
|
||||||
for turn in range(10)
|
{"role": "assistant", "content": "a1", "timestamp": "2026-01-01T00:00:01"},
|
||||||
for role in ("user", "assistant")
|
{"role": "user", "content": "u2", "timestamp": "2026-01-01T00:00:02"},
|
||||||
]
|
]
|
||||||
loop.sessions.save(session)
|
loop.sessions.save(session)
|
||||||
|
monkeypatch.setattr(memory_module, "estimate_message_tokens", lambda _m: 500)
|
||||||
|
|
||||||
call_count = [0]
|
call_count = [0]
|
||||||
def mock_estimate(_session, *, runtime):
|
def mock_estimate(_session, *, runtime):
|
||||||
call_count[0] += 1
|
call_count[0] += 1
|
||||||
|
|||||||
@@ -373,6 +373,7 @@ class TestToolEventProgress:
|
|||||||
"""The /goal command rewrites the prompt but must not bypass WebUI file-edit progress."""
|
"""The /goal command rewrites the prompt but must not bypass WebUI file-edit progress."""
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
provider.get_default_model.return_value = "test-model"
|
provider.get_default_model.return_value = "test-model"
|
||||||
call_count = 0
|
call_count = 0
|
||||||
|
|
||||||
@@ -459,6 +460,7 @@ class TestToolEventProgress:
|
|||||||
"""Non-streaming channels should get one final reply, not token progress spam."""
|
"""Non-streaming channels should get one final reply, not token progress spam."""
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
provider.get_default_model.return_value = "openai-codex/gpt-5.5"
|
provider.get_default_model.return_value = "openai-codex/gpt-5.5"
|
||||||
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="Hello", tool_calls=[]))
|
provider.chat_with_retry = AsyncMock(return_value=LLMResponse(content="Hello", tool_calls=[]))
|
||||||
provider.chat_stream_with_retry = AsyncMock()
|
provider.chat_stream_with_retry = AsyncMock()
|
||||||
@@ -491,6 +493,7 @@ class TestToolEventProgress:
|
|||||||
"""Streaming channels still receive provider deltas through stream events."""
|
"""Streaming channels still receive provider deltas through stream events."""
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
provider.get_default_model.return_value = "openai-codex/gpt-5.5"
|
provider.get_default_model.return_value = "openai-codex/gpt-5.5"
|
||||||
|
|
||||||
async def chat_stream_with_retry(*, on_content_delta, **kwargs):
|
async def chat_stream_with_retry(*, on_content_delta, **kwargs):
|
||||||
@@ -541,6 +544,7 @@ class TestToolEventProgress:
|
|||||||
) -> None:
|
) -> None:
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
provider.get_default_model.return_value = "test-model"
|
provider.get_default_model.return_value = "test-model"
|
||||||
responses = iter([
|
responses = iter([
|
||||||
LLMResponse(content="first-", finish_reason="length"),
|
LLMResponse(content="first-", finish_reason="length"),
|
||||||
@@ -586,6 +590,7 @@ class TestToolEventProgress:
|
|||||||
) -> None:
|
) -> None:
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
provider.get_default_model.return_value = "test-model"
|
provider.get_default_model.return_value = "test-model"
|
||||||
call_count = 0
|
call_count = 0
|
||||||
|
|
||||||
@@ -632,6 +637,7 @@ class TestToolEventProgress:
|
|||||||
) -> None:
|
) -> None:
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
provider.get_default_model.return_value = "test-model"
|
provider.get_default_model.return_value = "test-model"
|
||||||
|
|
||||||
async def chat_stream_with_retry(*, on_content_delta, **kwargs):
|
async def chat_stream_with_retry(*, on_content_delta, **kwargs):
|
||||||
@@ -722,6 +728,7 @@ class TestToolEventProgress:
|
|||||||
"""A no-tools finalization must not be dropped after empty stream retries."""
|
"""A no-tools finalization must not be dropped after empty stream retries."""
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
provider.get_default_model.return_value = "openai-codex/gpt-5.5"
|
provider.get_default_model.return_value = "openai-codex/gpt-5.5"
|
||||||
provider.chat_stream_with_retry = AsyncMock(side_effect=[
|
provider.chat_stream_with_retry = AsyncMock(side_effect=[
|
||||||
LLMResponse(content=None, tool_calls=[]),
|
LLMResponse(content=None, tool_calls=[]),
|
||||||
@@ -769,6 +776,7 @@ class TestToolEventProgress:
|
|||||||
) -> None:
|
) -> None:
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
provider.get_default_model.return_value = "openai-codex/gpt-5.5"
|
provider.get_default_model.return_value = "openai-codex/gpt-5.5"
|
||||||
first_request_started = asyncio.Event()
|
first_request_started = asyncio.Event()
|
||||||
release_first_request = asyncio.Event()
|
release_first_request = asyncio.Event()
|
||||||
@@ -927,6 +935,7 @@ class TestToolEventProgress:
|
|||||||
"""Recovered streaming output should use a new stream segment."""
|
"""Recovered streaming output should use a new stream segment."""
|
||||||
bus = MessageBus()
|
bus = MessageBus()
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
provider.get_default_model.return_value = "openai-codex/gpt-5.5"
|
provider.get_default_model.return_value = "openai-codex/gpt-5.5"
|
||||||
|
|
||||||
async def chat_stream_with_retry(*, on_content_delta, on_stream_recover, **kwargs):
|
async def chat_stream_with_retry(*, on_content_delta, on_stream_recover, **kwargs):
|
||||||
@@ -979,12 +988,13 @@ class TestToolEventProgress:
|
|||||||
provider.chat_with_retry.assert_not_awaited()
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_streamed_content_is_not_repeated_before_tool_execution(
|
async def test_streamed_progress_is_not_repeated_before_tool_execution(
|
||||||
self,
|
self,
|
||||||
tmp_path: Path,
|
tmp_path: Path,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""If content was already streamed, tool setup should not repeat it."""
|
"""If content was already streamed as progress, tool setup should not repeat it."""
|
||||||
loop = _make_loop(tmp_path)
|
loop = _make_loop(tmp_path)
|
||||||
|
loop.provider.supports_progress_deltas = True
|
||||||
tool_call = ToolCallRequest(id="call1", name="custom_tool", arguments={"path": "foo.txt"})
|
tool_call = ToolCallRequest(id="call1", name="custom_tool", arguments={"path": "foo.txt"})
|
||||||
calls = iter([
|
calls = iter([
|
||||||
LLMResponse(content="I will inspect it.", tool_calls=[tool_call]),
|
LLMResponse(content="I will inspect it.", tool_calls=[tool_call]),
|
||||||
|
|||||||
@@ -1,181 +0,0 @@
|
|||||||
"""Test /new archival behavior."""
|
|
||||||
|
|
||||||
import asyncio
|
|
||||||
from collections.abc import Coroutine
|
|
||||||
from pathlib import Path
|
|
||||||
from typing import Any
|
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
|
||||||
|
|
||||||
import pytest
|
|
||||||
|
|
||||||
|
|
||||||
class TestNewCommandArchival:
|
|
||||||
"""Test /new archival behavior with the structured archive flow."""
|
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _make_loop(tmp_path: Path):
|
|
||||||
from nanobot.agent.loop import AgentLoop
|
|
||||||
from nanobot.bus.queue import MessageBus
|
|
||||||
from nanobot.providers.base import GenerationSettings, LLMResponse
|
|
||||||
|
|
||||||
bus = MessageBus()
|
|
||||||
provider = MagicMock()
|
|
||||||
provider.get_default_model.return_value = "test-model"
|
|
||||||
provider.estimate_prompt_tokens.return_value = (10_000, "test")
|
|
||||||
provider.generation = GenerationSettings(max_tokens=100)
|
|
||||||
loop = AgentLoop(
|
|
||||||
bus=bus,
|
|
||||||
provider=provider,
|
|
||||||
workspace=tmp_path,
|
|
||||||
model="test-model",
|
|
||||||
context_window_tokens=1,
|
|
||||||
)
|
|
||||||
loop.provider.chat_with_retry = AsyncMock(
|
|
||||||
return_value=LLMResponse(content="ok", tool_calls=[])
|
|
||||||
)
|
|
||||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
|
||||||
return loop
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_new_clears_session_immediately_even_if_archive_fails(
|
|
||||||
self,
|
|
||||||
tmp_path: Path,
|
|
||||||
) -> None:
|
|
||||||
"""/new clears session immediately; archive is fire-and-forget."""
|
|
||||||
from nanobot.bus.events import InboundMessage
|
|
||||||
|
|
||||||
loop = self._make_loop(tmp_path)
|
|
||||||
session = loop.sessions.get_or_create("cli:test")
|
|
||||||
for i in range(5):
|
|
||||||
session.add_message("user", f"msg{i}")
|
|
||||||
session.add_message("assistant", f"resp{i}")
|
|
||||||
loop.sessions.save(session)
|
|
||||||
|
|
||||||
call_count = 0
|
|
||||||
expected_runtime = loop.llm_runtime()
|
|
||||||
|
|
||||||
async def _failing_summarize(session, *, archive_end, runtime) -> None:
|
|
||||||
nonlocal call_count
|
|
||||||
assert runtime is expected_runtime
|
|
||||||
assert session.key == "cli:test"
|
|
||||||
assert archive_end == len(session.messages)
|
|
||||||
call_count += 1
|
|
||||||
|
|
||||||
loop.consolidator.archive_session = _failing_summarize # type: ignore[method-assign]
|
|
||||||
|
|
||||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
|
||||||
response = await loop._process_message(new_msg, runtime=expected_runtime)
|
|
||||||
|
|
||||||
assert response is not None
|
|
||||||
assert "new session started" in response.content.lower()
|
|
||||||
|
|
||||||
session_after = loop.sessions.get_or_create("cli:test")
|
|
||||||
assert len(session_after.messages) == 0
|
|
||||||
|
|
||||||
await loop.aclose()
|
|
||||||
assert call_count == 1
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_new_reuses_replay_prefix_and_archives_only_unarchived_messages(
|
|
||||||
self,
|
|
||||||
tmp_path: Path,
|
|
||||||
) -> None:
|
|
||||||
from nanobot.bus.events import InboundMessage
|
|
||||||
|
|
||||||
loop = self._make_loop(tmp_path)
|
|
||||||
loop.set_runtime_context_window(128_000)
|
|
||||||
session = loop.sessions.get_or_create("cli:test")
|
|
||||||
for i in range(5):
|
|
||||||
session.add_message("user", f"msg{i}")
|
|
||||||
session.add_message("assistant", f"resp{i}")
|
|
||||||
session.last_archived = len(session.messages) - 2
|
|
||||||
ordinary_history = session.get_history()
|
|
||||||
assert [message["content"] for message in ordinary_history] == [
|
|
||||||
"msg1",
|
|
||||||
"resp1",
|
|
||||||
"msg2",
|
|
||||||
"resp2",
|
|
||||||
"msg3",
|
|
||||||
"resp3",
|
|
||||||
"msg4",
|
|
||||||
"resp4",
|
|
||||||
]
|
|
||||||
loop.sessions.save(session)
|
|
||||||
|
|
||||||
expected_runtime = loop.llm_runtime()
|
|
||||||
scheduled: list[Coroutine[Any, Any, object]] = []
|
|
||||||
loop.schedule_background = scheduled.append # type: ignore[method-assign]
|
|
||||||
|
|
||||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
|
||||||
response = await loop._process_message(new_msg, runtime=expected_runtime)
|
|
||||||
|
|
||||||
assert response is not None
|
|
||||||
assert "new session started" in response.content.lower()
|
|
||||||
|
|
||||||
assert len(scheduled) == 1
|
|
||||||
await scheduled[0]
|
|
||||||
await loop.aclose()
|
|
||||||
sent = loop.provider.chat_with_retry.call_args.kwargs["messages"]
|
|
||||||
assert sent[1:-1] == ordinary_history
|
|
||||||
assert "final 2 conversation messages" in sent[-1]["content"]
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_new_clears_session_and_responds(self, tmp_path: Path) -> None:
|
|
||||||
from nanobot.bus.events import InboundMessage
|
|
||||||
|
|
||||||
loop = self._make_loop(tmp_path)
|
|
||||||
session = loop.sessions.get_or_create("cli:test")
|
|
||||||
for i in range(3):
|
|
||||||
session.add_message("user", f"msg{i}")
|
|
||||||
session.add_message("assistant", f"resp{i}")
|
|
||||||
loop.sessions.save(session)
|
|
||||||
expected_runtime = loop.llm_runtime()
|
|
||||||
|
|
||||||
async def _ok_summarize(session, *, archive_end, runtime) -> str:
|
|
||||||
assert runtime is expected_runtime
|
|
||||||
assert session.key == "cli:test"
|
|
||||||
assert archive_end == len(session.messages)
|
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive_session = _ok_summarize # type: ignore[method-assign]
|
|
||||||
|
|
||||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
|
||||||
response = await loop._process_message(new_msg, runtime=expected_runtime)
|
|
||||||
|
|
||||||
assert response is not None
|
|
||||||
assert "new session started" in response.content.lower()
|
|
||||||
assert loop.sessions.get_or_create("cli:test").messages == []
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_aclose_drains_background_tasks(self, tmp_path: Path) -> None:
|
|
||||||
"""aclose waits for background tasks to complete."""
|
|
||||||
from nanobot.bus.events import InboundMessage
|
|
||||||
|
|
||||||
loop = self._make_loop(tmp_path)
|
|
||||||
session = loop.sessions.get_or_create("cli:test")
|
|
||||||
for i in range(3):
|
|
||||||
session.add_message("user", f"msg{i}")
|
|
||||||
session.add_message("assistant", f"resp{i}")
|
|
||||||
loop.sessions.save(session)
|
|
||||||
|
|
||||||
archived = asyncio.Event()
|
|
||||||
release_archive = asyncio.Event()
|
|
||||||
expected_runtime = loop.llm_runtime()
|
|
||||||
|
|
||||||
async def _slow_summarize(session, *, archive_end, runtime) -> str:
|
|
||||||
assert runtime is expected_runtime
|
|
||||||
assert session.key == "cli:test"
|
|
||||||
assert archive_end == len(session.messages)
|
|
||||||
await release_archive.wait()
|
|
||||||
archived.set()
|
|
||||||
return "Summary."
|
|
||||||
|
|
||||||
loop.consolidator.archive_session = _slow_summarize # type: ignore[method-assign]
|
|
||||||
|
|
||||||
new_msg = InboundMessage(channel="cli", sender_id="user", chat_id="test", content="/new")
|
|
||||||
await loop._process_message(new_msg, runtime=expected_runtime)
|
|
||||||
|
|
||||||
assert not archived.is_set()
|
|
||||||
release_archive.set()
|
|
||||||
await loop.aclose()
|
|
||||||
assert archived.is_set()
|
|
||||||
@@ -798,6 +798,64 @@ async def test_runner_times_out_never_ending_streaming_request():
|
|||||||
provider.chat_with_retry.assert_not_awaited()
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runner_closes_progress_reasoning_on_streaming_wall_timeout():
|
||||||
|
from nanobot.agent.hook import AgentHook
|
||||||
|
from nanobot.agent.runner import AgentRunner
|
||||||
|
|
||||||
|
provider = MagicMock(spec=LLMProvider)
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
|
events: list[tuple[str, str | None]] = []
|
||||||
|
|
||||||
|
async def chat_stream_with_retry(*, on_content_delta, **kwargs):
|
||||||
|
try:
|
||||||
|
await on_content_delta("<think>working...</think>")
|
||||||
|
await asyncio.sleep(3600)
|
||||||
|
finally:
|
||||||
|
events.append(("provider_cancelled", None))
|
||||||
|
|
||||||
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
|
tools = MagicMock()
|
||||||
|
tools.get_definitions.return_value = []
|
||||||
|
|
||||||
|
class ProgressReasoningHook(AgentHook):
|
||||||
|
async def emit_reasoning(self, reasoning_content: str | None) -> None:
|
||||||
|
if reasoning_content:
|
||||||
|
events.append(("reasoning", reasoning_content))
|
||||||
|
|
||||||
|
async def emit_reasoning_end(self) -> None:
|
||||||
|
events.append(("reasoning_end", None))
|
||||||
|
|
||||||
|
real_wait_for = asyncio.wait_for
|
||||||
|
|
||||||
|
async def fake_wait_for(coro, *, timeout):
|
||||||
|
assert timeout == 300.0
|
||||||
|
return await real_wait_for(coro, timeout=0.01)
|
||||||
|
|
||||||
|
runner = AgentRunner()
|
||||||
|
with patch("nanobot.agent.runner.asyncio.wait_for", fake_wait_for):
|
||||||
|
result = await runner.run(make_run_spec(provider,
|
||||||
|
initial_messages=[{"role": "user", "content": "think forever"}],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=1,
|
||||||
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
hook=ProgressReasoningHook(),
|
||||||
|
progress_callback=AsyncMock(),
|
||||||
|
llm_timeout_s=1,
|
||||||
|
))
|
||||||
|
|
||||||
|
assert result.stop_reason == "error"
|
||||||
|
assert result.final_content == "Error calling LLM: timed out after 300s"
|
||||||
|
assert events == [
|
||||||
|
("reasoning", "working..."),
|
||||||
|
("provider_cancelled", None),
|
||||||
|
("reasoning_end", None),
|
||||||
|
]
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_replaces_empty_tool_result_with_marker():
|
async def test_runner_replaces_empty_tool_result_with_marker():
|
||||||
from nanobot.agent.runner import AgentRunner
|
from nanobot.agent.runner import AgentRunner
|
||||||
@@ -1227,8 +1285,13 @@ async def test_runner_accumulates_usage_and_preserves_cache_reads():
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_binds_on_retry_wait_callback():
|
async def test_runner_binds_on_retry_wait_to_retry_callback_not_progress():
|
||||||
"""Provider retry heartbeats use the explicitly supplied callback."""
|
"""Regression: provider retry heartbeats must route through
|
||||||
|
``retry_wait_callback``, not ``progress_callback``. Binding them to
|
||||||
|
the progress callback (as an earlier runtime refactor did) caused
|
||||||
|
internal retry diagnostics like "Model request failed, retry in 1s"
|
||||||
|
to leak to end-user channels as normal progress updates.
|
||||||
|
"""
|
||||||
from nanobot.agent.runner import AgentRunner
|
from nanobot.agent.runner import AgentRunner
|
||||||
|
|
||||||
captured: dict = {}
|
captured: dict = {}
|
||||||
@@ -1242,6 +1305,7 @@ async def test_runner_binds_on_retry_wait_callback():
|
|||||||
tools = MagicMock()
|
tools = MagicMock()
|
||||||
tools.get_definitions.return_value = []
|
tools.get_definitions.return_value = []
|
||||||
|
|
||||||
|
progress_cb = AsyncMock()
|
||||||
retry_wait_cb = AsyncMock()
|
retry_wait_cb = AsyncMock()
|
||||||
|
|
||||||
runner = AgentRunner()
|
runner = AgentRunner()
|
||||||
@@ -1254,10 +1318,12 @@ async def test_runner_binds_on_retry_wait_callback():
|
|||||||
model="test-model",
|
model="test-model",
|
||||||
max_iterations=1,
|
max_iterations=1,
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
retry_wait_callback=retry_wait_cb,
|
retry_wait_callback=retry_wait_cb,
|
||||||
))
|
))
|
||||||
|
|
||||||
assert captured["on_retry_wait"] is retry_wait_cb
|
assert captured["on_retry_wait"] is retry_wait_cb
|
||||||
|
assert captured["on_retry_wait"] is not progress_cb
|
||||||
|
|
||||||
|
|
||||||
# ---------------------------------------------------------------------------
|
# ---------------------------------------------------------------------------
|
||||||
|
|||||||
@@ -9,9 +9,7 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from agent.runner_helpers import make_run_spec
|
from agent.runner_helpers import make_run_spec
|
||||||
from nanobot.agent.hook import AgentHook, AgentHookContext
|
|
||||||
from nanobot.agent.tools import ToolResult
|
from nanobot.agent.tools import ToolResult
|
||||||
from nanobot.agent.tools.execution import execute_tool_calls
|
|
||||||
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, ToolCallRequest
|
||||||
|
|
||||||
@@ -57,7 +55,11 @@ async def test_runner_returns_tool_exception_to_model_for_recovery():
|
|||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.parametrize("control_error", [KeyboardInterrupt, SystemExit])
|
@pytest.mark.parametrize("control_error", [KeyboardInterrupt, SystemExit])
|
||||||
async def test_tool_execution_propagates_control_flow_exceptions(control_error: type[BaseException]):
|
async def test_runner_propagates_tool_control_flow_exceptions(control_error: type[BaseException]):
|
||||||
|
from nanobot.agent.runner import AgentRunner
|
||||||
|
|
||||||
|
provider = MagicMock(spec=LLMProvider)
|
||||||
|
|
||||||
async def execute(_name, _args):
|
async def execute(_name, _args):
|
||||||
raise control_error("stop")
|
raise control_error("stop")
|
||||||
|
|
||||||
@@ -65,15 +67,22 @@ async def test_tool_execution_propagates_control_flow_exceptions(control_error:
|
|||||||
get_definitions=lambda: [],
|
get_definitions=lambda: [],
|
||||||
execute=execute,
|
execute=execute,
|
||||||
)
|
)
|
||||||
|
runner = AgentRunner()
|
||||||
|
spec = make_run_spec(
|
||||||
|
provider,
|
||||||
|
initial_messages=[],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=1,
|
||||||
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
)
|
||||||
|
|
||||||
with pytest.raises(control_error):
|
with pytest.raises(control_error):
|
||||||
await execute_tool_calls(
|
await runner._run_tool(
|
||||||
tools,
|
spec,
|
||||||
[ToolCallRequest(id="call_1", name="list_dir", arguments={})],
|
ToolCallRequest(id="call_1", name="list_dir", arguments={}),
|
||||||
concurrent=False,
|
|
||||||
external_lookup_counts={},
|
external_lookup_counts={},
|
||||||
workspace_violation_counts={},
|
workspace_violation_counts={},
|
||||||
hook=AgentHook(),
|
|
||||||
context=AgentHookContext(iteration=0, messages=[]),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
"""Tests for runner progress hooks and provider event routing."""
|
"""Tests for provider progress delta routing in the shared runner."""
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
@@ -6,6 +6,7 @@ from unittest.mock import AsyncMock, MagicMock
|
|||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from agent.runner_helpers import make_run_spec
|
from agent.runner_helpers import make_run_spec
|
||||||
|
from nanobot.agent.hook import CompositeHook
|
||||||
from nanobot.agent.hooks import FileEditActivityHook
|
from nanobot.agent.hooks import FileEditActivityHook
|
||||||
from nanobot.agent.progress_hook import AgentProgressHook
|
from nanobot.agent.progress_hook import AgentProgressHook
|
||||||
from nanobot.agent.runner import AgentRunner
|
from nanobot.agent.runner import AgentRunner
|
||||||
@@ -16,9 +17,45 @@ from nanobot.providers.base import LLMResponse, ToolCallRequest
|
|||||||
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
|
_MAX_TOOL_RESULT_CHARS = AgentDefaults().max_tool_result_chars
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_runner_streams_provider_progress_deltas_by_default():
|
||||||
|
"""Direct runner users keep the existing opt-in provider progress behavior."""
|
||||||
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
|
|
||||||
|
async def chat_stream_with_retry(*, on_content_delta, **kwargs):
|
||||||
|
await on_content_delta("he")
|
||||||
|
await on_content_delta("llo")
|
||||||
|
return LLMResponse(content="hello", tool_calls=[], usage=None)
|
||||||
|
|
||||||
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
|
tools = MagicMock()
|
||||||
|
tools.get_definitions.return_value = []
|
||||||
|
progress_cb = AsyncMock()
|
||||||
|
|
||||||
|
runner = AgentRunner()
|
||||||
|
result = await runner.run(make_run_spec(provider,
|
||||||
|
initial_messages=[
|
||||||
|
{"role": "system", "content": "system"},
|
||||||
|
{"role": "user", "content": "hi"},
|
||||||
|
],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=1,
|
||||||
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
|
))
|
||||||
|
|
||||||
|
assert result.final_content == "hello"
|
||||||
|
assert [call.args[0] for call in progress_cb.await_args_list] == ["he", "llo"]
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_routes_hosted_tool_events_to_structured_progress():
|
async def test_runner_routes_hosted_tool_events_to_structured_progress():
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
|
|
||||||
async def chat_stream_with_retry(*, on_content_delta, on_tool_call_delta, **kwargs):
|
async def chat_stream_with_retry(*, on_content_delta, on_tool_call_delta, **kwargs):
|
||||||
await on_tool_call_delta({
|
await on_tool_call_delta({
|
||||||
@@ -51,17 +88,13 @@ async def test_runner_routes_hosted_tool_events_to_structured_progress():
|
|||||||
tools.get_definitions.return_value = []
|
tools.get_definitions.return_value = []
|
||||||
progress_events: list[dict] = []
|
progress_events: list[dict] = []
|
||||||
progress_text: list[str] = []
|
progress_text: list[str] = []
|
||||||
streamed_text: list[str] = []
|
|
||||||
|
|
||||||
async def progress_cb(content, *, tool_events=None, **kwargs):
|
async def progress_cb(content, *, tool_events=None, **kwargs):
|
||||||
progress_text.append(content)
|
progress_text.append(content)
|
||||||
if tool_events:
|
if tool_events:
|
||||||
progress_events.extend(tool_events)
|
progress_events.extend(tool_events)
|
||||||
|
|
||||||
async def stream_cb(content: str) -> None:
|
hook = CompositeHook([AgentProgressHook(on_progress=progress_cb)])
|
||||||
streamed_text.append(content)
|
|
||||||
|
|
||||||
hook = AgentProgressHook(on_progress=progress_cb, on_stream=stream_cb)
|
|
||||||
result = await AgentRunner().run(make_run_spec(
|
result = await AgentRunner().run(make_run_spec(
|
||||||
provider,
|
provider,
|
||||||
initial_messages=[{"role": "user", "content": "search X"}],
|
initial_messages=[{"role": "user", "content": "search X"}],
|
||||||
@@ -69,6 +102,7 @@ async def test_runner_routes_hosted_tool_events_to_structured_progress():
|
|||||||
model="test-model",
|
model="test-model",
|
||||||
max_iterations=1,
|
max_iterations=1,
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
hook=hook,
|
hook=hook,
|
||||||
))
|
))
|
||||||
|
|
||||||
@@ -99,14 +133,14 @@ async def test_runner_routes_hosted_tool_events_to_structured_progress():
|
|||||||
"embeds": [],
|
"embeds": [],
|
||||||
},
|
},
|
||||||
]
|
]
|
||||||
assert progress_text == ['search X "nanobot oauth"', ""]
|
assert progress_text == ['search X "nanobot oauth"', "", "done"]
|
||||||
assert streamed_text == ["done"]
|
|
||||||
provider.chat_with_retry.assert_not_awaited()
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_fails_pending_hosted_tool_when_model_request_fails():
|
async def test_runner_fails_pending_hosted_tool_when_model_request_fails():
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
|
|
||||||
async def chat_stream_with_retry(*, on_tool_call_delta, **kwargs):
|
async def chat_stream_with_retry(*, on_tool_call_delta, **kwargs):
|
||||||
await on_tool_call_delta({
|
await on_tool_call_delta({
|
||||||
@@ -132,10 +166,7 @@ async def test_runner_fails_pending_hosted_tool_when_model_request_fails():
|
|||||||
if tool_events:
|
if tool_events:
|
||||||
progress_events.extend(tool_events)
|
progress_events.extend(tool_events)
|
||||||
|
|
||||||
async def stream_cb(_content: str) -> None:
|
hook = CompositeHook([AgentProgressHook(on_progress=progress_cb)])
|
||||||
pass
|
|
||||||
|
|
||||||
hook = AgentProgressHook(on_progress=progress_cb, on_stream=stream_cb)
|
|
||||||
result = await AgentRunner().run(make_run_spec(
|
result = await AgentRunner().run(make_run_spec(
|
||||||
provider,
|
provider,
|
||||||
initial_messages=[{"role": "user", "content": "search X"}],
|
initial_messages=[{"role": "user", "content": "search X"}],
|
||||||
@@ -143,6 +174,7 @@ async def test_runner_fails_pending_hosted_tool_when_model_request_fails():
|
|||||||
model="test-model",
|
model="test-model",
|
||||||
max_iterations=1,
|
max_iterations=1,
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
hook=hook,
|
hook=hook,
|
||||||
))
|
))
|
||||||
|
|
||||||
@@ -168,6 +200,7 @@ async def test_runner_fails_pending_hosted_tool_when_model_request_fails():
|
|||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_emits_write_file_diff_from_tool_execution_snapshots(tmp_path):
|
async def test_runner_emits_write_file_diff_from_tool_execution_snapshots(tmp_path):
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
call_count = 0
|
call_count = 0
|
||||||
progress_events: list[dict] = []
|
progress_events: list[dict] = []
|
||||||
(tmp_path / "big.txt").write_text("old\n", encoding="utf-8")
|
(tmp_path / "big.txt").write_text("old\n", encoding="utf-8")
|
||||||
@@ -185,7 +218,7 @@ async def test_runner_emits_write_file_diff_from_tool_execution_snapshots(tmp_pa
|
|||||||
def prepare_call(self, name, params):
|
def prepare_call(self, name, params):
|
||||||
return tool, params, None
|
return tool, params, None
|
||||||
|
|
||||||
async def chat_with_retry(**kwargs):
|
async def chat_stream_with_retry(**kwargs):
|
||||||
nonlocal call_count
|
nonlocal call_count
|
||||||
call_count += 1
|
call_count += 1
|
||||||
if call_count == 1:
|
if call_count == 1:
|
||||||
@@ -202,7 +235,8 @@ async def test_runner_emits_write_file_diff_from_tool_execution_snapshots(tmp_pa
|
|||||||
)
|
)
|
||||||
return LLMResponse(content="done", tool_calls=[], usage=None)
|
return LLMResponse(content="done", tool_calls=[], usage=None)
|
||||||
|
|
||||||
provider.chat_with_retry = chat_with_retry
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
tools = Tools()
|
tools = Tools()
|
||||||
|
|
||||||
runner = AgentRunner()
|
runner = AgentRunner()
|
||||||
@@ -212,6 +246,7 @@ async def test_runner_emits_write_file_diff_from_tool_execution_snapshots(tmp_pa
|
|||||||
model="test-model",
|
model="test-model",
|
||||||
max_iterations=2,
|
max_iterations=2,
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
hook=FileEditActivityHook(on_progress=progress_cb, workspace=tmp_path),
|
hook=FileEditActivityHook(on_progress=progress_cb, workspace=tmp_path),
|
||||||
))
|
))
|
||||||
@@ -228,11 +263,13 @@ async def test_runner_emits_write_file_diff_from_tool_execution_snapshots(tmp_pa
|
|||||||
and event["diff"]["format"] == "unified"
|
and event["diff"]["format"] == "unified"
|
||||||
for event in progress_events
|
for event in progress_events
|
||||||
)
|
)
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_emits_edit_file_diff_from_tool_execution_snapshots(tmp_path):
|
async def test_runner_emits_edit_file_diff_from_tool_execution_snapshots(tmp_path):
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
call_count = 0
|
call_count = 0
|
||||||
progress_events: list[dict] = []
|
progress_events: list[dict] = []
|
||||||
target = tmp_path / "notes.txt"
|
target = tmp_path / "notes.txt"
|
||||||
@@ -251,7 +288,7 @@ async def test_runner_emits_edit_file_diff_from_tool_execution_snapshots(tmp_pat
|
|||||||
def prepare_call(self, name, params):
|
def prepare_call(self, name, params):
|
||||||
return tool, params, None
|
return tool, params, None
|
||||||
|
|
||||||
async def chat_with_retry(**kwargs):
|
async def chat_stream_with_retry(**kwargs):
|
||||||
nonlocal call_count
|
nonlocal call_count
|
||||||
call_count += 1
|
call_count += 1
|
||||||
if call_count == 1:
|
if call_count == 1:
|
||||||
@@ -272,7 +309,8 @@ async def test_runner_emits_edit_file_diff_from_tool_execution_snapshots(tmp_pat
|
|||||||
)
|
)
|
||||||
return LLMResponse(content="done", tool_calls=[], usage=None)
|
return LLMResponse(content="done", tool_calls=[], usage=None)
|
||||||
|
|
||||||
provider.chat_with_retry = chat_with_retry
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
tools = Tools()
|
tools = Tools()
|
||||||
|
|
||||||
runner = AgentRunner()
|
runner = AgentRunner()
|
||||||
@@ -282,6 +320,7 @@ async def test_runner_emits_edit_file_diff_from_tool_execution_snapshots(tmp_pat
|
|||||||
model="test-model",
|
model="test-model",
|
||||||
max_iterations=2,
|
max_iterations=2,
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
hook=FileEditActivityHook(on_progress=progress_cb, workspace=tmp_path),
|
hook=FileEditActivityHook(on_progress=progress_cb, workspace=tmp_path),
|
||||||
))
|
))
|
||||||
@@ -296,11 +335,13 @@ async def test_runner_emits_edit_file_diff_from_tool_execution_snapshots(tmp_pat
|
|||||||
and event["diff"]["format"] == "unified"
|
and event["diff"]["format"] == "unified"
|
||||||
for event in progress_events
|
for event in progress_events
|
||||||
)
|
)
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_marks_file_edit_activity_failed_when_tool_errors(tmp_path):
|
async def test_runner_marks_file_edit_activity_failed_when_tool_errors(tmp_path):
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
call_count = 0
|
call_count = 0
|
||||||
progress_events: list[dict] = []
|
progress_events: list[dict] = []
|
||||||
|
|
||||||
@@ -317,7 +358,7 @@ async def test_runner_marks_file_edit_activity_failed_when_tool_errors(tmp_path)
|
|||||||
def prepare_call(self, name, params):
|
def prepare_call(self, name, params):
|
||||||
return tool, params, None
|
return tool, params, None
|
||||||
|
|
||||||
async def chat_with_retry(**kwargs):
|
async def chat_stream_with_retry(**kwargs):
|
||||||
nonlocal call_count
|
nonlocal call_count
|
||||||
call_count += 1
|
call_count += 1
|
||||||
if call_count == 1:
|
if call_count == 1:
|
||||||
@@ -334,7 +375,8 @@ async def test_runner_marks_file_edit_activity_failed_when_tool_errors(tmp_path)
|
|||||||
)
|
)
|
||||||
return LLMResponse(content="done", tool_calls=[], usage=None)
|
return LLMResponse(content="done", tool_calls=[], usage=None)
|
||||||
|
|
||||||
provider.chat_with_retry = chat_with_retry
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
tools = Tools()
|
tools = Tools()
|
||||||
|
|
||||||
runner = AgentRunner()
|
runner = AgentRunner()
|
||||||
@@ -344,6 +386,7 @@ async def test_runner_marks_file_edit_activity_failed_when_tool_errors(tmp_path)
|
|||||||
model="test-model",
|
model="test-model",
|
||||||
max_iterations=2,
|
max_iterations=2,
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
hook=FileEditActivityHook(on_progress=progress_cb, workspace=tmp_path),
|
hook=FileEditActivityHook(on_progress=progress_cb, workspace=tmp_path),
|
||||||
))
|
))
|
||||||
@@ -352,11 +395,13 @@ async def test_runner_marks_file_edit_activity_failed_when_tool_errors(tmp_path)
|
|||||||
assert progress_events[-1]["path"] == "aborted.txt"
|
assert progress_events[-1]["path"] == "aborted.txt"
|
||||||
assert progress_events[-1]["phase"] == "error"
|
assert progress_events[-1]["phase"] == "error"
|
||||||
assert progress_events[-1]["status"] == "error"
|
assert progress_events[-1]["status"] == "error"
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_marks_file_edit_activity_failed_when_cancelled(tmp_path):
|
async def test_runner_marks_file_edit_activity_failed_when_cancelled(tmp_path):
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
progress_events: list[dict] = []
|
progress_events: list[dict] = []
|
||||||
executing = asyncio.Event()
|
executing = asyncio.Event()
|
||||||
target = tmp_path / "cancelled.txt"
|
target = tmp_path / "cancelled.txt"
|
||||||
@@ -381,7 +426,7 @@ async def test_runner_marks_file_edit_activity_failed_when_cancelled(tmp_path):
|
|||||||
def prepare_call(self, name, params):
|
def prepare_call(self, name, params):
|
||||||
return tool, params, None
|
return tool, params, None
|
||||||
|
|
||||||
async def chat_with_retry(**kwargs):
|
async def chat_stream_with_retry(**kwargs):
|
||||||
return LLMResponse(
|
return LLMResponse(
|
||||||
content=None,
|
content=None,
|
||||||
tool_calls=[
|
tool_calls=[
|
||||||
@@ -394,7 +439,8 @@ async def test_runner_marks_file_edit_activity_failed_when_cancelled(tmp_path):
|
|||||||
usage=None,
|
usage=None,
|
||||||
)
|
)
|
||||||
|
|
||||||
provider.chat_with_retry = chat_with_retry
|
provider.chat_stream_with_retry = chat_stream_with_retry
|
||||||
|
provider.chat_with_retry = AsyncMock()
|
||||||
tools = Tools()
|
tools = Tools()
|
||||||
|
|
||||||
runner = AgentRunner()
|
runner = AgentRunner()
|
||||||
@@ -404,6 +450,7 @@ async def test_runner_marks_file_edit_activity_failed_when_cancelled(tmp_path):
|
|||||||
model="test-model",
|
model="test-model",
|
||||||
max_iterations=2,
|
max_iterations=2,
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
progress_callback=progress_cb,
|
||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
hook=FileEditActivityHook(on_progress=progress_cb, workspace=tmp_path),
|
hook=FileEditActivityHook(on_progress=progress_cb, workspace=tmp_path),
|
||||||
)))
|
)))
|
||||||
@@ -417,3 +464,4 @@ async def test_runner_marks_file_edit_activity_failed_when_cancelled(tmp_path):
|
|||||||
assert progress_events[-1]["path"] == "cancelled.txt"
|
assert progress_events[-1]["path"] == "cancelled.txt"
|
||||||
assert progress_events[-1]["status"] == "error"
|
assert progress_events[-1]["status"] == "error"
|
||||||
assert progress_events[-1]["error"] == "Task interrupted before this tool finished."
|
assert progress_events[-1]["error"] == "Task interrupted before this tool finished."
|
||||||
|
provider.chat_with_retry.assert_not_awaited()
|
||||||
|
|||||||
@@ -9,14 +9,12 @@ channels, gated by ``context.streamed_reasoning`` rather than
|
|||||||
|
|
||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
from typing import Any
|
|
||||||
from unittest.mock import AsyncMock, MagicMock
|
from unittest.mock import AsyncMock, MagicMock
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from agent.runner_helpers import make_run_spec
|
from agent.runner_helpers import make_run_spec
|
||||||
from nanobot.agent.hook import AgentHook, AgentHookContext
|
from nanobot.agent.hook import AgentHook, AgentHookContext
|
||||||
from nanobot.agent.progress_hook import AgentProgressHook
|
|
||||||
from nanobot.config.schema import AgentDefaults
|
from nanobot.config.schema import AgentDefaults
|
||||||
from nanobot.providers.base import LLMResponse, LLMUsage, ToolCallRequest
|
from nanobot.providers.base import LLMResponse, LLMUsage, ToolCallRequest
|
||||||
|
|
||||||
@@ -37,51 +35,6 @@ class _RecordingHook(AgentHook):
|
|||||||
self.end_calls += 1
|
self.end_calls += 1
|
||||||
|
|
||||||
|
|
||||||
class _StreamRecordingHook(_RecordingHook):
|
|
||||||
def __init__(self) -> None:
|
|
||||||
super().__init__()
|
|
||||||
self.streamed: list[str] = []
|
|
||||||
|
|
||||||
def wants_streaming(self) -> bool:
|
|
||||||
return True
|
|
||||||
|
|
||||||
async def on_stream(self, _ctx: AgentHookContext, delta: str) -> None:
|
|
||||||
self.streamed.append(delta)
|
|
||||||
|
|
||||||
|
|
||||||
class _LifecycleRecordingHook(AgentHook):
|
|
||||||
def __init__(self) -> None:
|
|
||||||
super().__init__()
|
|
||||||
self.events: list[str] = []
|
|
||||||
|
|
||||||
def wants_streaming(self) -> bool:
|
|
||||||
return True
|
|
||||||
|
|
||||||
async def emit_reasoning(self, reasoning_content: str | None) -> None:
|
|
||||||
if reasoning_content:
|
|
||||||
self.events.append(f"reasoning:{reasoning_content}")
|
|
||||||
|
|
||||||
async def emit_reasoning_end(self) -> None:
|
|
||||||
self.events.append("reasoning_end")
|
|
||||||
|
|
||||||
async def on_stream(self, _ctx: AgentHookContext, delta: str) -> None:
|
|
||||||
self.events.append(f"content:{delta}")
|
|
||||||
|
|
||||||
async def on_stream_end(self, _ctx: AgentHookContext, *, resuming: bool) -> None:
|
|
||||||
self.events.append(f"stream_end:{resuming}")
|
|
||||||
|
|
||||||
async def before_execute_tools(self, context: AgentHookContext) -> None:
|
|
||||||
names = ",".join(call.name for call in context.tool_calls)
|
|
||||||
self.events.append(f"local_tools:{names}")
|
|
||||||
|
|
||||||
async def on_provider_tool_event(
|
|
||||||
self,
|
|
||||||
_context: AgentHookContext,
|
|
||||||
event: dict[str, Any],
|
|
||||||
) -> None:
|
|
||||||
self.events.append(f"hosted_tool:{event.get('phase')}")
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_preserves_reasoning_fields_in_assistant_history():
|
async def test_runner_preserves_reasoning_fields_in_assistant_history():
|
||||||
"""Reasoning fields ride along on the persisted assistant message so
|
"""Reasoning fields ride along on the persisted assistant message so
|
||||||
@@ -248,6 +201,7 @@ async def test_runner_emits_reasoning_content_even_when_answer_was_streamed():
|
|||||||
from nanobot.agent.runner import AgentRunner
|
from nanobot.agent.runner import AgentRunner
|
||||||
|
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
|
|
||||||
async def chat_stream_with_retry(*, on_content_delta=None, **kwargs):
|
async def chat_stream_with_retry(*, on_content_delta=None, **kwargs):
|
||||||
if on_content_delta:
|
if on_content_delta:
|
||||||
@@ -264,7 +218,12 @@ async def test_runner_emits_reasoning_content_even_when_answer_was_streamed():
|
|||||||
tools = MagicMock()
|
tools = MagicMock()
|
||||||
tools.get_definitions.return_value = []
|
tools.get_definitions.return_value = []
|
||||||
|
|
||||||
hook = _StreamRecordingHook()
|
progress_calls: list[str] = []
|
||||||
|
|
||||||
|
async def _progress(content: str, **_kwargs):
|
||||||
|
progress_calls.append(content)
|
||||||
|
|
||||||
|
hook = _RecordingHook()
|
||||||
runner = AgentRunner()
|
runner = AgentRunner()
|
||||||
result = await runner.run(make_run_spec(provider,
|
result = await runner.run(make_run_spec(provider,
|
||||||
initial_messages=[{"role": "user", "content": "question"}],
|
initial_messages=[{"role": "user", "content": "question"}],
|
||||||
@@ -273,10 +232,11 @@ async def test_runner_emits_reasoning_content_even_when_answer_was_streamed():
|
|||||||
max_iterations=3,
|
max_iterations=3,
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
hook=hook,
|
hook=hook,
|
||||||
|
progress_callback=_progress,
|
||||||
))
|
))
|
||||||
|
|
||||||
assert result.final_content == "The answer."
|
assert result.final_content == "The answer."
|
||||||
assert hook.streamed == ["The ", "answer."]
|
assert progress_calls, "answer should have streamed via progress callback"
|
||||||
assert hook.emitted == ["step-by-step deduction"]
|
assert hook.emitted == ["step-by-step deduction"]
|
||||||
|
|
||||||
|
|
||||||
@@ -287,6 +247,7 @@ async def test_runner_does_not_double_emit_when_inline_think_already_streamed():
|
|||||||
from nanobot.agent.runner import AgentRunner
|
from nanobot.agent.runner import AgentRunner
|
||||||
|
|
||||||
provider = MagicMock()
|
provider = MagicMock()
|
||||||
|
provider.supports_progress_deltas = True
|
||||||
|
|
||||||
async def chat_stream_with_retry(*, on_content_delta=None, **kwargs):
|
async def chat_stream_with_retry(*, on_content_delta=None, **kwargs):
|
||||||
if on_content_delta:
|
if on_content_delta:
|
||||||
@@ -302,16 +263,10 @@ async def test_runner_does_not_double_emit_when_inline_think_already_streamed():
|
|||||||
tools = MagicMock()
|
tools = MagicMock()
|
||||||
tools.get_definitions.return_value = []
|
tools.get_definitions.return_value = []
|
||||||
|
|
||||||
reasoning_events: list[str] = []
|
async def _progress(content: str, **_kwargs):
|
||||||
|
|
||||||
async def _progress(content: str, *, reasoning: bool = False, **_kwargs):
|
|
||||||
if reasoning:
|
|
||||||
reasoning_events.append(content)
|
|
||||||
|
|
||||||
async def _stream(_content: str) -> None:
|
|
||||||
pass
|
pass
|
||||||
|
|
||||||
hook = AgentProgressHook(on_progress=_progress, on_stream=_stream)
|
hook = _RecordingHook()
|
||||||
runner = AgentRunner()
|
runner = AgentRunner()
|
||||||
result = await runner.run(make_run_spec(provider,
|
result = await runner.run(make_run_spec(provider,
|
||||||
initial_messages=[{"role": "user", "content": "question"}],
|
initial_messages=[{"role": "user", "content": "question"}],
|
||||||
@@ -320,10 +275,12 @@ async def test_runner_does_not_double_emit_when_inline_think_already_streamed():
|
|||||||
max_iterations=3,
|
max_iterations=3,
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
hook=hook,
|
hook=hook,
|
||||||
|
progress_callback=_progress,
|
||||||
))
|
))
|
||||||
|
|
||||||
assert result.final_content == "The answer."
|
assert result.final_content == "The answer."
|
||||||
assert reasoning_events == ["working..."]
|
assert hook.emitted == ["working..."]
|
||||||
|
assert hook.end_calls >= 1, "reasoning stream must be closed once the answer starts"
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -363,6 +320,14 @@ async def test_runner_closes_reasoning_stream_after_one_shot_response():
|
|||||||
assert hook.end_calls == 1
|
assert hook.end_calls == 1
|
||||||
|
|
||||||
|
|
||||||
|
class _StreamRecordingHook(_RecordingHook):
|
||||||
|
def wants_streaming(self) -> bool:
|
||||||
|
return True
|
||||||
|
|
||||||
|
async def on_stream(self, _ctx: AgentHookContext, delta: str) -> None:
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_streams_native_thinking_deltas_without_post_hoc_dup():
|
async def test_runner_streams_native_thinking_deltas_without_post_hoc_dup():
|
||||||
"""Anthropic-style ``on_thinking_delta`` should fan out to ``emit_reasoning``;
|
"""Anthropic-style ``on_thinking_delta`` should fan out to ``emit_reasoning``;
|
||||||
@@ -405,155 +370,6 @@ async def test_runner_streams_native_thinking_deltas_without_post_hoc_dup():
|
|||||||
assert hook.emitted == ["part1", "part2"]
|
assert hook.emitted == ["part1", "part2"]
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_runner_closes_native_reasoning_before_streaming_answer():
|
|
||||||
from nanobot.agent.runner import AgentRunner
|
|
||||||
|
|
||||||
provider = MagicMock()
|
|
||||||
|
|
||||||
async def chat_stream_with_retry(
|
|
||||||
*, on_content_delta=None, on_thinking_delta=None, **kwargs
|
|
||||||
):
|
|
||||||
if on_thinking_delta:
|
|
||||||
await on_thinking_delta("inspect")
|
|
||||||
if on_content_delta:
|
|
||||||
await on_content_delta("done")
|
|
||||||
return LLMResponse(content="done", tool_calls=[], usage=None)
|
|
||||||
|
|
||||||
provider.chat_stream_with_retry = chat_stream_with_retry
|
|
||||||
tools = MagicMock()
|
|
||||||
tools.get_definitions.return_value = []
|
|
||||||
hook = _LifecycleRecordingHook()
|
|
||||||
|
|
||||||
result = await AgentRunner().run(make_run_spec(
|
|
||||||
provider,
|
|
||||||
initial_messages=[{"role": "user", "content": "q"}],
|
|
||||||
tools=tools,
|
|
||||||
model="test-model",
|
|
||||||
max_iterations=1,
|
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
|
||||||
hook=hook,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert result.final_content == "done"
|
|
||||||
assert hook.events == [
|
|
||||||
"reasoning:inspect",
|
|
||||||
"reasoning_end",
|
|
||||||
"content:done",
|
|
||||||
"stream_end:False",
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_runner_closes_native_reasoning_before_local_tool_execution():
|
|
||||||
from nanobot.agent.runner import AgentRunner
|
|
||||||
|
|
||||||
provider = MagicMock()
|
|
||||||
responses = iter([
|
|
||||||
LLMResponse(
|
|
||||||
content="",
|
|
||||||
finish_reason="tool_calls",
|
|
||||||
tool_calls=[ToolCallRequest(id="call-1", name="list_dir", arguments={"path": "."})],
|
|
||||||
usage=None,
|
|
||||||
),
|
|
||||||
LLMResponse(content="done", tool_calls=[], usage=None),
|
|
||||||
])
|
|
||||||
|
|
||||||
async def chat_stream_with_retry(
|
|
||||||
*, on_content_delta=None, on_thinking_delta=None, **kwargs
|
|
||||||
):
|
|
||||||
response = next(responses)
|
|
||||||
if response.tool_calls:
|
|
||||||
if on_thinking_delta:
|
|
||||||
await on_thinking_delta("inspect")
|
|
||||||
elif on_content_delta:
|
|
||||||
await on_content_delta("done")
|
|
||||||
return response
|
|
||||||
|
|
||||||
provider.chat_stream_with_retry = chat_stream_with_retry
|
|
||||||
tools = MagicMock()
|
|
||||||
tools.get_definitions.return_value = []
|
|
||||||
tools.execute = AsyncMock(return_value="tool result")
|
|
||||||
hook = _LifecycleRecordingHook()
|
|
||||||
|
|
||||||
result = await AgentRunner().run(make_run_spec(
|
|
||||||
provider,
|
|
||||||
initial_messages=[{"role": "user", "content": "inspect"}],
|
|
||||||
tools=tools,
|
|
||||||
model="test-model",
|
|
||||||
max_iterations=2,
|
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
|
||||||
hook=hook,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert result.final_content == "done"
|
|
||||||
assert hook.events == [
|
|
||||||
"reasoning:inspect",
|
|
||||||
"reasoning_end",
|
|
||||||
"stream_end:True",
|
|
||||||
"local_tools:list_dir",
|
|
||||||
"content:done",
|
|
||||||
"stream_end:False",
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_runner_closes_native_reasoning_before_hosted_tool_event():
|
|
||||||
from nanobot.agent.runner import AgentRunner
|
|
||||||
|
|
||||||
provider = MagicMock()
|
|
||||||
|
|
||||||
async def chat_stream_with_retry(
|
|
||||||
*, on_content_delta=None, on_thinking_delta=None, on_tool_call_delta=None, **kwargs
|
|
||||||
):
|
|
||||||
if on_thinking_delta:
|
|
||||||
await on_thinking_delta("search")
|
|
||||||
if on_tool_call_delta:
|
|
||||||
await on_tool_call_delta({
|
|
||||||
"kind": "hosted_tool",
|
|
||||||
"phase": "start",
|
|
||||||
"call_id": "search-1",
|
|
||||||
"name": "web_search",
|
|
||||||
"arguments": {"query": "nanobot"},
|
|
||||||
})
|
|
||||||
await on_tool_call_delta({
|
|
||||||
"kind": "hosted_tool",
|
|
||||||
"phase": "end",
|
|
||||||
"call_id": "search-1",
|
|
||||||
"name": "web_search",
|
|
||||||
"arguments": {"query": "nanobot"},
|
|
||||||
"result": {"count": 1},
|
|
||||||
})
|
|
||||||
if on_content_delta:
|
|
||||||
await on_content_delta("done")
|
|
||||||
return LLMResponse(content="done", tool_calls=[], usage=None)
|
|
||||||
|
|
||||||
provider.chat_stream_with_retry = chat_stream_with_retry
|
|
||||||
tools = MagicMock()
|
|
||||||
tools.get_definitions.return_value = []
|
|
||||||
hook = _LifecycleRecordingHook()
|
|
||||||
|
|
||||||
result = await AgentRunner().run(make_run_spec(
|
|
||||||
provider,
|
|
||||||
initial_messages=[{"role": "user", "content": "search"}],
|
|
||||||
tools=tools,
|
|
||||||
model="test-model",
|
|
||||||
max_iterations=1,
|
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
|
||||||
hook=hook,
|
|
||||||
))
|
|
||||||
|
|
||||||
assert result.final_content == "done"
|
|
||||||
assert hook.events == [
|
|
||||||
"reasoning:search",
|
|
||||||
"reasoning_end",
|
|
||||||
"hosted_tool:start",
|
|
||||||
"hosted_tool:end",
|
|
||||||
"content:done",
|
|
||||||
"stream_end:False",
|
|
||||||
]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_runner_strips_thinking_tags_from_native_thinking_deltas():
|
async def test_runner_strips_thinking_tags_from_native_thinking_deltas():
|
||||||
from nanobot.agent.runner import AgentRunner
|
from nanobot.agent.runner import AgentRunner
|
||||||
|
|||||||
@@ -9,7 +9,6 @@ import pytest
|
|||||||
from agent.runner_helpers import make_run_spec
|
from agent.runner_helpers import make_run_spec
|
||||||
from nanobot.agent.runner import AgentRunner
|
from nanobot.agent.runner import AgentRunner
|
||||||
from nanobot.agent.tools import ToolResult
|
from nanobot.agent.tools import ToolResult
|
||||||
from nanobot.agent.tools.execution import is_ssrf_violation
|
|
||||||
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, ToolCallRequest
|
||||||
|
|
||||||
@@ -67,20 +66,20 @@ async def test_runner_does_not_abort_on_workspace_violation_anymore():
|
|||||||
def test_is_ssrf_violation_recognizes_private_url_blocks():
|
def test_is_ssrf_violation_recognizes_private_url_blocks():
|
||||||
"""SSRF rejections are classified separately from workspace boundaries."""
|
"""SSRF rejections are classified separately from workspace boundaries."""
|
||||||
ssrf_msg = "Error: Command blocked by safety guard (internal/private URL detected)"
|
ssrf_msg = "Error: Command blocked by safety guard (internal/private URL detected)"
|
||||||
assert is_ssrf_violation(ssrf_msg) is True
|
assert AgentRunner._is_ssrf_violation(ssrf_msg) is True
|
||||||
assert is_ssrf_violation(
|
assert AgentRunner._is_ssrf_violation(
|
||||||
"URL validation failed: Blocked: host resolves to private/internal address 192.168.1.2"
|
"URL validation failed: Blocked: host resolves to private/internal address 192.168.1.2"
|
||||||
) is True
|
) is True
|
||||||
|
|
||||||
# Workspace-bound markers are NOT classified as SSRF.
|
# Workspace-bound markers are NOT classified as SSRF.
|
||||||
assert is_ssrf_violation(
|
assert AgentRunner._is_ssrf_violation(
|
||||||
"Error: Command blocked by safety guard (path outside working dir)"
|
"Error: Command blocked by safety guard (path outside working dir)"
|
||||||
) is False
|
) is False
|
||||||
assert is_ssrf_violation(
|
assert AgentRunner._is_ssrf_violation(
|
||||||
"Path /tmp/x is outside allowed directory /ws"
|
"Path /tmp/x is outside allowed directory /ws"
|
||||||
) is False
|
) is False
|
||||||
# Deny / allowlist filter messages stay non-fatal too.
|
# Deny / allowlist filter messages stay non-fatal too.
|
||||||
assert is_ssrf_violation(
|
assert AgentRunner._is_ssrf_violation(
|
||||||
"Error: Command blocked by deny pattern filter"
|
"Error: Command blocked by deny pattern filter"
|
||||||
) is False
|
) is False
|
||||||
|
|
||||||
|
|||||||
@@ -3,17 +3,14 @@
|
|||||||
from __future__ import annotations
|
from __future__ import annotations
|
||||||
|
|
||||||
import asyncio
|
import asyncio
|
||||||
from typing import Any
|
|
||||||
from unittest.mock import AsyncMock, MagicMock, patch
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
|
||||||
import pytest
|
import pytest
|
||||||
|
|
||||||
from agent.runner_helpers import make_run_spec
|
from agent.runner_helpers import make_run_spec
|
||||||
from nanobot.agent.hook import AgentHook, AgentHookContext
|
|
||||||
from nanobot.agent.runner import AgentRunner
|
from nanobot.agent.runner import AgentRunner
|
||||||
from nanobot.agent.tools.base import Tool, ToolResult
|
from nanobot.agent.tools.base import Tool, ToolResult
|
||||||
from nanobot.agent.tools.context import ToolContext
|
from nanobot.agent.tools.context import ToolContext
|
||||||
from nanobot.agent.tools.execution import execute_tool_calls
|
|
||||||
from nanobot.agent.tools.loader import ToolLoader
|
from nanobot.agent.tools.loader import ToolLoader
|
||||||
from nanobot.agent.tools.registry import ToolRegistry
|
from nanobot.agent.tools.registry import ToolRegistry
|
||||||
from nanobot.config.schema import AgentDefaults
|
from nanobot.config.schema import AgentDefaults
|
||||||
@@ -153,69 +150,31 @@ def _tool_message(result, tool_call_id: str) -> dict:
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_tool_execution_propagates_preparation_failure():
|
async def test_runner_propagates_tool_preparation_failure():
|
||||||
tools = MagicMock()
|
tools = MagicMock()
|
||||||
tools.prepare_call.side_effect = RuntimeError("tool preparation failed")
|
tools.prepare_call.side_effect = RuntimeError("tool preparation failed")
|
||||||
tools.execute = AsyncMock()
|
tools.execute = AsyncMock()
|
||||||
|
|
||||||
with pytest.raises(RuntimeError, match="tool preparation failed"):
|
with pytest.raises(RuntimeError, match="tool preparation failed"):
|
||||||
await execute_tool_calls(
|
await AgentRunner()._run_tool(
|
||||||
tools,
|
make_run_spec(
|
||||||
[ToolCallRequest(id="call-1", name="demo", arguments={})],
|
MagicMock(),
|
||||||
concurrent=False,
|
initial_messages=[],
|
||||||
external_lookup_counts={},
|
tools=tools,
|
||||||
workspace_violation_counts={},
|
model="test-model",
|
||||||
hook=AgentHook(),
|
max_iterations=1,
|
||||||
context=AgentHookContext(iteration=0, messages=[]),
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
),
|
||||||
|
ToolCallRequest(id="call-1", name="demo", arguments={}),
|
||||||
|
{},
|
||||||
|
{},
|
||||||
)
|
)
|
||||||
|
|
||||||
tools.execute.assert_not_awaited()
|
tools.execute.assert_not_awaited()
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_tool_execution_propagates_cancellation_without_error_hook():
|
async def test_runner_batches_read_only_tools_before_exclusive_work():
|
||||||
tools = MagicMock()
|
|
||||||
tools.prepare_call.return_value = (None, {}, None)
|
|
||||||
tools.execute = AsyncMock(side_effect=asyncio.CancelledError)
|
|
||||||
|
|
||||||
events: list[str] = []
|
|
||||||
|
|
||||||
class RecordingHook(AgentHook):
|
|
||||||
async def before_execute_tool(
|
|
||||||
self,
|
|
||||||
context: AgentHookContext,
|
|
||||||
tool_call: ToolCallRequest,
|
|
||||||
tool: Any,
|
|
||||||
params: Any,
|
|
||||||
) -> None:
|
|
||||||
events.append("before")
|
|
||||||
|
|
||||||
async def on_execute_tool_error(
|
|
||||||
self,
|
|
||||||
context: AgentHookContext,
|
|
||||||
tool_call: ToolCallRequest,
|
|
||||||
tool: Any,
|
|
||||||
params: Any,
|
|
||||||
error: Any,
|
|
||||||
) -> None:
|
|
||||||
events.append("error")
|
|
||||||
|
|
||||||
with pytest.raises(asyncio.CancelledError):
|
|
||||||
await execute_tool_calls(
|
|
||||||
tools,
|
|
||||||
[ToolCallRequest(id="call-1", name="demo", arguments={})],
|
|
||||||
concurrent=False,
|
|
||||||
external_lookup_counts={},
|
|
||||||
workspace_violation_counts={},
|
|
||||||
hook=RecordingHook(),
|
|
||||||
context=AgentHookContext(iteration=0, messages=[]),
|
|
||||||
)
|
|
||||||
|
|
||||||
assert events == ["before"]
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_tool_execution_batches_read_only_tools_before_exclusive_work():
|
|
||||||
tools = ToolRegistry()
|
tools = ToolRegistry()
|
||||||
shared_events: list[str] = []
|
shared_events: list[str] = []
|
||||||
read_a = _DelayTool("read_a", delay=0.05, read_only=True, shared_events=shared_events)
|
read_a = _DelayTool("read_a", delay=0.05, read_only=True, shared_events=shared_events)
|
||||||
@@ -225,18 +184,24 @@ async def test_tool_execution_batches_read_only_tools_before_exclusive_work():
|
|||||||
tools.register(read_b)
|
tools.register(read_b)
|
||||||
tools.register(write_a)
|
tools.register(write_a)
|
||||||
|
|
||||||
await execute_tool_calls(
|
provider = MagicMock()
|
||||||
tools,
|
runner = AgentRunner()
|
||||||
|
await runner._execute_tools(
|
||||||
|
make_run_spec(provider,
|
||||||
|
initial_messages=[],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=1,
|
||||||
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
concurrent_tools=True,
|
||||||
|
),
|
||||||
[
|
[
|
||||||
ToolCallRequest(id="ro1", name="read_a", arguments={}),
|
ToolCallRequest(id="ro1", name="read_a", arguments={}),
|
||||||
ToolCallRequest(id="ro2", name="read_b", arguments={}),
|
ToolCallRequest(id="ro2", name="read_b", arguments={}),
|
||||||
ToolCallRequest(id="rw1", name="write_a", arguments={}),
|
ToolCallRequest(id="rw1", name="write_a", arguments={}),
|
||||||
],
|
],
|
||||||
concurrent=True,
|
{},
|
||||||
external_lookup_counts={},
|
{},
|
||||||
workspace_violation_counts={},
|
|
||||||
hook=AgentHook(),
|
|
||||||
context=AgentHookContext(iteration=0, messages=[]),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert shared_events[0:2] == ["start:read_a", "start:read_b"]
|
assert shared_events[0:2] == ["start:read_a", "start:read_b"]
|
||||||
@@ -247,7 +212,7 @@ async def test_tool_execution_batches_read_only_tools_before_exclusive_work():
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_tool_execution_does_not_batch_exclusive_read_only_tools():
|
async def test_runner_does_not_batch_exclusive_read_only_tools():
|
||||||
tools = ToolRegistry()
|
tools = ToolRegistry()
|
||||||
shared_events: list[str] = []
|
shared_events: list[str] = []
|
||||||
read_a = _DelayTool("read_a", delay=0.03, read_only=True, shared_events=shared_events)
|
read_a = _DelayTool("read_a", delay=0.03, read_only=True, shared_events=shared_events)
|
||||||
@@ -263,18 +228,24 @@ async def test_tool_execution_does_not_batch_exclusive_read_only_tools():
|
|||||||
tools.register(ddg_like)
|
tools.register(ddg_like)
|
||||||
tools.register(read_b)
|
tools.register(read_b)
|
||||||
|
|
||||||
await execute_tool_calls(
|
provider = MagicMock()
|
||||||
tools,
|
runner = AgentRunner()
|
||||||
|
await runner._execute_tools(
|
||||||
|
make_run_spec(provider,
|
||||||
|
initial_messages=[],
|
||||||
|
tools=tools,
|
||||||
|
model="test-model",
|
||||||
|
max_iterations=1,
|
||||||
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
|
concurrent_tools=True,
|
||||||
|
),
|
||||||
[
|
[
|
||||||
ToolCallRequest(id="ro1", name="read_a", arguments={}),
|
ToolCallRequest(id="ro1", name="read_a", arguments={}),
|
||||||
ToolCallRequest(id="ddg1", name="ddg_like", arguments={}),
|
ToolCallRequest(id="ddg1", name="ddg_like", arguments={}),
|
||||||
ToolCallRequest(id="ro2", name="read_b", arguments={}),
|
ToolCallRequest(id="ro2", name="read_b", arguments={}),
|
||||||
],
|
],
|
||||||
concurrent=True,
|
{},
|
||||||
external_lookup_counts={},
|
{},
|
||||||
workspace_violation_counts={},
|
|
||||||
hook=AgentHook(),
|
|
||||||
context=AgentHookContext(iteration=0, messages=[]),
|
|
||||||
)
|
)
|
||||||
|
|
||||||
assert shared_events[0] == "start:read_a"
|
assert shared_events[0] == "start:read_a"
|
||||||
|
|||||||
@@ -148,28 +148,28 @@ def test_retain_recent_legal_suffix_keeps_recent_messages():
|
|||||||
assert session.messages[-1]["content"] == "msg9"
|
assert session.messages[-1]["content"] == "msg9"
|
||||||
|
|
||||||
|
|
||||||
def test_retain_recent_legal_suffix_adjusts_last_archived():
|
def test_retain_recent_legal_suffix_adjusts_last_consolidated():
|
||||||
session = Session(key="test:trim-cons")
|
session = Session(key="test:trim-cons")
|
||||||
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}"})
|
||||||
session.last_archived = 7
|
session.last_consolidated = 7
|
||||||
|
|
||||||
session.retain_recent_legal_suffix(4)
|
session.retain_recent_legal_suffix(4)
|
||||||
|
|
||||||
assert len(session.messages) == 4
|
assert len(session.messages) == 4
|
||||||
assert session.last_archived == 1
|
assert session.last_consolidated == 1
|
||||||
|
|
||||||
|
|
||||||
def test_retain_recent_legal_suffix_zero_clears_session():
|
def test_retain_recent_legal_suffix_zero_clears_session():
|
||||||
session = Session(key="test:trim-zero")
|
session = Session(key="test:trim-zero")
|
||||||
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}"})
|
||||||
session.last_archived = 5
|
session.last_consolidated = 5
|
||||||
|
|
||||||
session.retain_recent_legal_suffix(0)
|
session.retain_recent_legal_suffix(0)
|
||||||
|
|
||||||
assert session.messages == []
|
assert session.messages == []
|
||||||
assert session.last_archived == 0
|
assert session.last_consolidated == 0
|
||||||
|
|
||||||
|
|
||||||
def test_retain_recent_legal_suffix_keeps_legal_tool_boundary():
|
def test_retain_recent_legal_suffix_keeps_legal_tool_boundary():
|
||||||
@@ -188,15 +188,15 @@ def test_retain_recent_legal_suffix_keeps_legal_tool_boundary():
|
|||||||
assert history[0]["content"] == "keep"
|
assert history[0]["content"] == "keep"
|
||||||
|
|
||||||
|
|
||||||
# --- last_archived > 0 ---
|
# --- last_consolidated > 0 ---
|
||||||
|
|
||||||
def test_orphan_trim_with_last_archived():
|
def test_orphan_trim_with_last_consolidated():
|
||||||
"""Orphan trimming works correctly when a session is partially archived."""
|
"""Orphan trimming works correctly when session is partially consolidated."""
|
||||||
session = Session(key="test:consolidated")
|
session = Session(key="test:consolidated")
|
||||||
for i in range(10):
|
for i in range(10):
|
||||||
session.messages.append({"role": "user", "content": f"old {i}"})
|
session.messages.append({"role": "user", "content": f"old {i}"})
|
||||||
session.messages.extend(_tool_turn("cons", i))
|
session.messages.extend(_tool_turn("cons", i))
|
||||||
session.last_archived = 30
|
session.last_consolidated = 30
|
||||||
|
|
||||||
session.messages.append({"role": "user", "content": "recent"})
|
session.messages.append({"role": "user", "content": "recent"})
|
||||||
for i in range(15):
|
for i in range(15):
|
||||||
@@ -213,7 +213,7 @@ def test_get_history_replays_recent_messages_after_full_archive():
|
|||||||
for i in range(10):
|
for i in range(10):
|
||||||
session.messages.append({"role": "user", "content": f"u{i}"})
|
session.messages.append({"role": "user", "content": f"u{i}"})
|
||||||
session.messages.append({"role": "assistant", "content": f"a{i}"})
|
session.messages.append({"role": "assistant", "content": f"a{i}"})
|
||||||
session.last_archived = len(session.messages)
|
session.last_consolidated = len(session.messages)
|
||||||
|
|
||||||
history = session.get_history(max_messages=100)
|
history = session.get_history(max_messages=100)
|
||||||
|
|
||||||
@@ -229,8 +229,8 @@ def test_get_history_replays_recent_messages_after_full_archive():
|
|||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
def test_get_history_extends_archived_replay_to_preceding_user():
|
def test_get_history_extends_compacted_replay_to_preceding_user():
|
||||||
session = Session(key="test:archived-tool-turn")
|
session = Session(key="test:compacted-tool-turn")
|
||||||
session.messages.extend(
|
session.messages.extend(
|
||||||
[
|
[
|
||||||
{"role": "user", "content": "old"},
|
{"role": "user", "content": "old"},
|
||||||
@@ -242,7 +242,7 @@ def test_get_history_extends_archived_replay_to_preceding_user():
|
|||||||
{"role": "assistant", "content": "done"},
|
{"role": "assistant", "content": "done"},
|
||||||
]
|
]
|
||||||
)
|
)
|
||||||
session.last_archived = len(session.messages)
|
session.last_consolidated = len(session.messages)
|
||||||
|
|
||||||
history = session.get_history(max_messages=100)
|
history = session.get_history(max_messages=100)
|
||||||
|
|
||||||
@@ -251,8 +251,8 @@ def test_get_history_extends_archived_replay_to_preceding_user():
|
|||||||
_assert_no_orphans(history)
|
_assert_no_orphans(history)
|
||||||
|
|
||||||
|
|
||||||
def test_archived_tool_turn_can_extend_past_message_cap():
|
def test_compacted_tool_turn_can_extend_past_message_cap():
|
||||||
session = Session(key="test:long-archived-tool-turn")
|
session = Session(key="test:long-compacted-tool-turn")
|
||||||
session.messages.extend(
|
session.messages.extend(
|
||||||
[
|
[
|
||||||
{"role": "user", "content": "old"},
|
{"role": "user", "content": "old"},
|
||||||
@@ -263,7 +263,7 @@ def test_archived_tool_turn_can_extend_past_message_cap():
|
|||||||
for i in range(50):
|
for i in range(50):
|
||||||
session.messages.extend(_tool_turn("keep", i))
|
session.messages.extend(_tool_turn("keep", i))
|
||||||
session.messages.append({"role": "assistant", "content": "done"})
|
session.messages.append({"role": "assistant", "content": "done"})
|
||||||
session.last_archived = len(session.messages)
|
session.last_consolidated = len(session.messages)
|
||||||
|
|
||||||
history = session.get_history(max_messages=120)
|
history = session.get_history(max_messages=120)
|
||||||
|
|
||||||
@@ -635,7 +635,7 @@ def test_fork_session_allows_index_equal_to_user_count(tmp_path):
|
|||||||
assert [m["content"] for m in forked.messages] == ["round1", "answer1"]
|
assert [m["content"] for m in forked.messages] == ["round1", "answer1"]
|
||||||
|
|
||||||
|
|
||||||
def test_fork_session_drops_summary_when_fork_point_is_inside_archived_prefix(tmp_path):
|
def test_fork_session_drops_summary_when_fork_point_is_inside_consolidated_prefix(tmp_path):
|
||||||
manager = SessionManager(tmp_path)
|
manager = SessionManager(tmp_path)
|
||||||
source = manager.get_or_create("websocket:source")
|
source = manager.get_or_create("websocket:source")
|
||||||
source.messages = [
|
source.messages = [
|
||||||
@@ -644,7 +644,7 @@ def test_fork_session_drops_summary_when_fork_point_is_inside_archived_prefix(tm
|
|||||||
{"role": "user", "content": "round2 fork me"},
|
{"role": "user", "content": "round2 fork me"},
|
||||||
{"role": "assistant", "content": "answer2"},
|
{"role": "assistant", "content": "answer2"},
|
||||||
]
|
]
|
||||||
source.last_archived = 4
|
source.last_consolidated = 4
|
||||||
source.metadata["_last_summary"] = {"text": "round2 fork me and answer2"}
|
source.metadata["_last_summary"] = {"text": "round2 fork me and answer2"}
|
||||||
manager.save(source)
|
manager.save(source)
|
||||||
|
|
||||||
@@ -656,7 +656,7 @@ def test_fork_session_drops_summary_when_fork_point_is_inside_archived_prefix(tm
|
|||||||
|
|
||||||
assert forked is not None
|
assert forked is not None
|
||||||
assert [m["content"] for m in forked.messages] == ["round1", "answer1"]
|
assert [m["content"] for m in forked.messages] == ["round1", "answer1"]
|
||||||
assert forked.last_archived == 0
|
assert forked.last_consolidated == 0
|
||||||
assert "_last_summary" not in forked.metadata
|
assert "_last_summary" not in forked.metadata
|
||||||
|
|
||||||
|
|
||||||
@@ -880,7 +880,7 @@ def test_retain_recent_legal_suffix_returns_all_on_zero():
|
|||||||
session = Session(key="test:zero-return")
|
session = Session(key="test:zero-return")
|
||||||
for i in range(5):
|
for i in range(5):
|
||||||
session.messages.append({"role": "user", "content": f"msg{i}"})
|
session.messages.append({"role": "user", "content": f"msg{i}"})
|
||||||
session.last_archived = 3
|
session.last_consolidated = 3
|
||||||
|
|
||||||
result = session.retain_recent_legal_suffix(0)
|
result = session.retain_recent_legal_suffix(0)
|
||||||
|
|
||||||
@@ -889,21 +889,22 @@ def test_retain_recent_legal_suffix_returns_all_on_zero():
|
|||||||
assert session.messages == []
|
assert session.messages == []
|
||||||
|
|
||||||
|
|
||||||
def test_retain_recent_legal_suffix_last_archived_correct_in_else_branch():
|
def test_retain_recent_legal_suffix_last_consolidated_correct_in_else_branch():
|
||||||
"""last_archived should count retained messages from the old archived prefix."""
|
"""last_consolidated after retain_recent_legal_suffix should reflect how
|
||||||
|
many retained messages were inside the old consolidated prefix."""
|
||||||
session = Session(key="test:else-lc-correct")
|
session = Session(key="test:else-lc-correct")
|
||||||
# 20 messages: u0..u9, a0..a9
|
# 20 messages: u0..u9, a0..a9
|
||||||
for i in range(10):
|
for i in range(10):
|
||||||
session.messages.append({"role": "user", "content": f"u{i}"})
|
session.messages.append({"role": "user", "content": f"u{i}"})
|
||||||
for i in range(10):
|
for i in range(10):
|
||||||
session.messages.append({"role": "assistant", "content": f"a{i}"})
|
session.messages.append({"role": "assistant", "content": f"a{i}"})
|
||||||
session.last_archived = 12 # u0..u9, a0, a1 archived
|
session.last_consolidated = 12 # u0..u9, a0, a1 consolidated
|
||||||
|
|
||||||
result = session.retain_recent_legal_suffix(4)
|
result = session.retain_recent_legal_suffix(4)
|
||||||
|
|
||||||
# Retained messages start from latest user (u9) + max_messages forward
|
# Retained messages start from latest user (u9) + max_messages forward
|
||||||
# so retained = [u9, a0..a9][:4] → but these are from original indices 9..12
|
# so retained = [u9, a0..a9][:4] → but these are from original indices 9..12
|
||||||
# Of those, indices 9,10,11 are < 12 (before_lc), so new_lc = 3
|
# Of those, indices 9,10,11 are < 12 (before_lc), so new_lc = 3
|
||||||
assert session.last_archived == 3
|
assert session.last_consolidated == 3
|
||||||
# already_cons should count dropped messages with original index < 12
|
# already_cons should count dropped messages with original index < 12
|
||||||
assert result.already_consolidated_count == 9
|
assert result.already_consolidated_count == 9
|
||||||
|
|||||||
@@ -179,7 +179,7 @@ def test_compact_probe_keeps_delivery_in_visible_suffix():
|
|||||||
{"role": "assistant", "content": "a2"},
|
{"role": "assistant", "content": "a2"},
|
||||||
{"role": "assistant", "content": "a3"},
|
{"role": "assistant", "content": "a3"},
|
||||||
]
|
]
|
||||||
probe = Session(key="test:probe", messages=tail)
|
probe = Session(key="test:probe", messages=tail, last_consolidated=0)
|
||||||
|
|
||||||
probe.retain_recent_legal_suffix(3, extend_to_user=True)
|
probe.retain_recent_legal_suffix(3, extend_to_user=True)
|
||||||
|
|
||||||
|
|||||||
@@ -520,35 +520,6 @@ class TestCancelBySession:
|
|||||||
count = await sm.cancel_by_session("nonexistent")
|
count = await sm.cancel_by_session("nonexistent")
|
||||||
assert count == 0
|
assert count == 0
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_cancels_active_and_queued_tasks(self, tmp_path):
|
|
||||||
sm = _manager(tmp_path, max_concurrent_subagents=1)
|
|
||||||
active_entered = asyncio.Event()
|
|
||||||
queued_entered = asyncio.Event()
|
|
||||||
|
|
||||||
async def _blocked_run(spec):
|
|
||||||
task = spec.initial_messages[-1]["content"]
|
|
||||||
if task == "active":
|
|
||||||
active_entered.set()
|
|
||||||
else:
|
|
||||||
queued_entered.set()
|
|
||||||
await asyncio.Event().wait()
|
|
||||||
|
|
||||||
sm.runner.run = _blocked_run
|
|
||||||
runtime = _runtime()
|
|
||||||
await sm.spawn("active", runtime=runtime, session_key="s1")
|
|
||||||
await asyncio.wait_for(active_entered.wait(), timeout=1.0)
|
|
||||||
await sm.spawn("queued", runtime=runtime, session_key="s1")
|
|
||||||
await asyncio.sleep(0)
|
|
||||||
|
|
||||||
assert not queued_entered.is_set()
|
|
||||||
assert await sm.cancel_by_session("s1") == 2
|
|
||||||
await asyncio.sleep(0)
|
|
||||||
|
|
||||||
assert not queued_entered.is_set()
|
|
||||||
assert sm._running_tasks == {}
|
|
||||||
assert sm._session_tasks == {}
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_already_done_not_counted(self, tmp_path):
|
async def test_already_done_not_counted(self, tmp_path):
|
||||||
sm = _manager(tmp_path)
|
sm = _manager(tmp_path)
|
||||||
|
|||||||
@@ -254,7 +254,7 @@ class TestDispatch:
|
|||||||
assert isinstance(second.event, StreamEndEvent)
|
assert isinstance(second.event, StreamEndEvent)
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_same_session_dispatches_serialize(self):
|
async def test_processing_lock_serializes(self):
|
||||||
from nanobot.bus.events import InboundMessage, OutboundMessage
|
from nanobot.bus.events import InboundMessage, OutboundMessage
|
||||||
|
|
||||||
loop, bus = _make_loop()
|
loop, bus = _make_loop()
|
||||||
|
|||||||
@@ -291,7 +291,7 @@ class TestCmdNewUnifiedSession:
|
|||||||
archived = loop.consolidator.archive_session.call_args.args[0]
|
archived = loop.consolidator.archive_session.call_args.args[0]
|
||||||
assert archived.key == "unified:default"
|
assert archived.key == "unified:default"
|
||||||
assert archived.messages == expected_snapshot
|
assert archived.messages == expected_snapshot
|
||||||
assert archived.last_archived == 0
|
assert archived.last_consolidated == 0
|
||||||
loop.consolidator.archive_session.assert_called_once_with(
|
loop.consolidator.archive_session.assert_called_once_with(
|
||||||
archived,
|
archived,
|
||||||
archive_end=len(expected_snapshot),
|
archive_end=len(expected_snapshot),
|
||||||
|
|||||||
@@ -211,8 +211,8 @@ async def test_spawn_forwards_temperature_to_run_spec(tmp_path):
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_background_spawn_waits_for_concurrency_capacity(tmp_path):
|
async def test_spawn_tool_rejects_when_at_concurrency_limit(tmp_path):
|
||||||
"""Background tasks should be accepted and start when capacity becomes available."""
|
"""SpawnTool should return an error string when the concurrency limit is reached."""
|
||||||
from nanobot.agent.subagent import SubagentManager
|
from nanobot.agent.subagent import SubagentManager
|
||||||
from nanobot.agent.tools.spawn import SpawnTool
|
from nanobot.agent.tools.spawn import SpawnTool
|
||||||
from nanobot.bus.queue import MessageBus
|
from nanobot.bus.queue import MessageBus
|
||||||
@@ -224,23 +224,14 @@ async def test_background_spawn_waits_for_concurrency_capacity(tmp_path):
|
|||||||
workspace=tmp_path,
|
workspace=tmp_path,
|
||||||
bus=bus,
|
bus=bus,
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
max_concurrent_subagents=1,
|
|
||||||
)
|
)
|
||||||
mgr._announce_result = AsyncMock()
|
mgr._announce_result = AsyncMock()
|
||||||
|
|
||||||
first_entered = asyncio.Event()
|
# Block the first subagent so it stays "running"
|
||||||
second_entered = asyncio.Event()
|
release = asyncio.Event()
|
||||||
release_first = asyncio.Event()
|
|
||||||
release_second = asyncio.Event()
|
|
||||||
|
|
||||||
async def fake_run(spec):
|
async def fake_run(spec):
|
||||||
task = spec.initial_messages[-1]["content"]
|
await release.wait()
|
||||||
if task == "first task":
|
|
||||||
first_entered.set()
|
|
||||||
await release_first.wait()
|
|
||||||
else:
|
|
||||||
second_entered.set()
|
|
||||||
await release_second.wait()
|
|
||||||
return SimpleNamespace(
|
return SimpleNamespace(
|
||||||
stop_reason="done",
|
stop_reason="done",
|
||||||
final_content="done",
|
final_content="done",
|
||||||
@@ -259,24 +250,19 @@ async def test_background_spawn_waits_for_concurrency_capacity(tmp_path):
|
|||||||
session_key="test:c1",
|
session_key="test:c1",
|
||||||
runtime=_runtime(provider),
|
runtime=_runtime(provider),
|
||||||
)):
|
)):
|
||||||
first_result = await tool.execute(task="first task")
|
# First spawn succeeds
|
||||||
assert "started" in first_result
|
result = await tool.execute(task="first task")
|
||||||
await asyncio.wait_for(first_entered.wait(), timeout=1.0)
|
assert "started" in result
|
||||||
|
|
||||||
second_result = await tool.execute(task="second task")
|
# Second spawn should be rejected (default limit is 1)
|
||||||
assert "started" in second_result
|
result = await tool.execute(task="second task")
|
||||||
tasks = list(mgr._running_tasks.values())
|
assert "Cannot spawn subagent" in result
|
||||||
await asyncio.sleep(0)
|
assert "concurrency limit reached" in result
|
||||||
assert not second_entered.is_set()
|
|
||||||
phases = {status.task_description: status.phase for status in mgr._task_statuses.values()}
|
|
||||||
assert phases == {"first task": "initializing", "second task": "queued"}
|
|
||||||
|
|
||||||
release_first.set()
|
# Release the first subagent
|
||||||
await asyncio.wait_for(second_entered.wait(), timeout=1.0)
|
release.set()
|
||||||
release_second.set()
|
# Allow cleanup
|
||||||
await asyncio.gather(*tasks, return_exceptions=True)
|
await asyncio.gather(*mgr._running_tasks.values(), return_exceptions=True)
|
||||||
await asyncio.sleep(0)
|
|
||||||
assert mgr._running_tasks == {}
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@@ -314,7 +300,7 @@ async def test_spawn_tool_waits_for_inline_result():
|
|||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_inline_spawn_waits_for_concurrency_capacity(tmp_path):
|
async def test_inline_spawn_counts_toward_concurrency_limit(tmp_path):
|
||||||
from nanobot.agent.subagent import SubagentManager
|
from nanobot.agent.subagent import SubagentManager
|
||||||
from nanobot.agent.tools.context import RequestContext, request_context
|
from nanobot.agent.tools.context import RequestContext, request_context
|
||||||
from nanobot.agent.tools.spawn import SpawnTool
|
from nanobot.agent.tools.spawn import SpawnTool
|
||||||
@@ -326,19 +312,12 @@ async def test_inline_spawn_waits_for_concurrency_capacity(tmp_path):
|
|||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
max_concurrent_subagents=1,
|
max_concurrent_subagents=1,
|
||||||
)
|
)
|
||||||
first_entered = asyncio.Event()
|
release = asyncio.Event()
|
||||||
second_entered = asyncio.Event()
|
entered = asyncio.Event()
|
||||||
release_first = asyncio.Event()
|
|
||||||
release_second = asyncio.Event()
|
|
||||||
|
|
||||||
async def fake_run(spec):
|
async def fake_run(spec):
|
||||||
task = spec.initial_messages[-1]["content"]
|
entered.set()
|
||||||
if task == "first":
|
await release.wait()
|
||||||
first_entered.set()
|
|
||||||
await release_first.wait()
|
|
||||||
else:
|
|
||||||
second_entered.set()
|
|
||||||
await release_second.wait()
|
|
||||||
return SimpleNamespace(
|
return SimpleNamespace(
|
||||||
stop_reason="done",
|
stop_reason="done",
|
||||||
final_content="done",
|
final_content="done",
|
||||||
@@ -355,100 +334,19 @@ async def test_inline_spawn_waits_for_concurrency_capacity(tmp_path):
|
|||||||
runtime=_runtime(MagicMock()),
|
runtime=_runtime(MagicMock()),
|
||||||
)):
|
)):
|
||||||
first = asyncio.create_task(tool.execute(task="first", wait=True))
|
first = asyncio.create_task(tool.execute(task="first", wait=True))
|
||||||
await asyncio.wait_for(first_entered.wait(), timeout=1.0)
|
await asyncio.wait_for(entered.wait(), timeout=1.0)
|
||||||
|
|
||||||
second = asyncio.create_task(tool.execute(task="second", wait=True))
|
second = await tool.execute(task="second", wait=True)
|
||||||
await asyncio.sleep(0)
|
|
||||||
|
|
||||||
assert not second.done()
|
assert "concurrency limit reached" in second
|
||||||
assert not second_entered.is_set()
|
assert manager.get_running_count() == 1
|
||||||
assert manager.get_running_count() == 2
|
release.set()
|
||||||
release_first.set()
|
|
||||||
assert await first == "done"
|
assert await first == "done"
|
||||||
await asyncio.wait_for(second_entered.wait(), timeout=1.0)
|
|
||||||
release_second.set()
|
|
||||||
assert await second == "done"
|
|
||||||
|
|
||||||
assert manager.get_running_count() == 0
|
assert manager.get_running_count() == 0
|
||||||
assert manager._session_tasks == {}
|
assert manager._session_tasks == {}
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_runner_executes_inline_spawn_batch_concurrently(tmp_path):
|
|
||||||
"""Adjacent blocking consultations should share one concurrent tool batch."""
|
|
||||||
from nanobot.agent.hook import AgentHook, AgentHookContext
|
|
||||||
from nanobot.agent.subagent import SubagentManager
|
|
||||||
from nanobot.agent.tools.context import RequestContext, request_context
|
|
||||||
from nanobot.agent.tools.execution import execute_tool_calls
|
|
||||||
from nanobot.agent.tools.registry import ToolRegistry
|
|
||||||
from nanobot.agent.tools.spawn import SpawnTool
|
|
||||||
from nanobot.bus.queue import MessageBus
|
|
||||||
from nanobot.providers.base import ToolCallRequest
|
|
||||||
|
|
||||||
manager = SubagentManager(
|
|
||||||
workspace=tmp_path,
|
|
||||||
bus=MessageBus(),
|
|
||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
|
||||||
max_concurrent_subagents=2,
|
|
||||||
)
|
|
||||||
both_entered = asyncio.Event()
|
|
||||||
release = asyncio.Event()
|
|
||||||
entered: list[str] = []
|
|
||||||
|
|
||||||
async def fake_run(spec):
|
|
||||||
entered.append(spec.initial_messages[-1]["content"])
|
|
||||||
if len(entered) == 2:
|
|
||||||
both_entered.set()
|
|
||||||
await release.wait()
|
|
||||||
return SimpleNamespace(
|
|
||||||
stop_reason="done",
|
|
||||||
final_content=spec.initial_messages[-1]["content"],
|
|
||||||
error=None,
|
|
||||||
tool_events=[],
|
|
||||||
)
|
|
||||||
|
|
||||||
manager.runner.run = AsyncMock(side_effect=fake_run)
|
|
||||||
tools = ToolRegistry()
|
|
||||||
tools.register(SpawnTool(manager))
|
|
||||||
runtime = _runtime(MagicMock())
|
|
||||||
calls = [
|
|
||||||
ToolCallRequest(
|
|
||||||
id="spawn-1",
|
|
||||||
name="spawn",
|
|
||||||
arguments={"task": "first", "wait": True},
|
|
||||||
),
|
|
||||||
ToolCallRequest(
|
|
||||||
id="spawn-2",
|
|
||||||
name="spawn",
|
|
||||||
arguments={"task": "second", "wait": True},
|
|
||||||
),
|
|
||||||
]
|
|
||||||
|
|
||||||
with request_context(RequestContext(
|
|
||||||
channel="test",
|
|
||||||
chat_id="c1",
|
|
||||||
session_key="test:c1",
|
|
||||||
runtime=runtime,
|
|
||||||
)):
|
|
||||||
execution = asyncio.create_task(execute_tool_calls(
|
|
||||||
tools,
|
|
||||||
calls,
|
|
||||||
concurrent=True,
|
|
||||||
external_lookup_counts={},
|
|
||||||
workspace_violation_counts={},
|
|
||||||
hook=AgentHook(),
|
|
||||||
context=AgentHookContext(iteration=0, messages=[], session_key="test:c1"),
|
|
||||||
))
|
|
||||||
await asyncio.wait_for(both_entered.wait(), timeout=1.0)
|
|
||||||
release.set()
|
|
||||||
results, events = await execution
|
|
||||||
|
|
||||||
assert set(entered) == {"first", "second"}
|
|
||||||
assert results == ["first", "second"]
|
|
||||||
assert [event["status"] for event in events] == ["ok", "ok"]
|
|
||||||
assert manager._running_tasks == {}
|
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_cancel_by_session_cancels_inline_subagent(tmp_path):
|
async def test_cancel_by_session_cancels_inline_subagent(tmp_path):
|
||||||
from nanobot.agent.subagent import SubagentManager
|
from nanobot.agent.subagent import SubagentManager
|
||||||
@@ -493,7 +391,6 @@ def test_subagent_default_max_concurrent_matches_agent_defaults(tmp_path):
|
|||||||
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
max_tool_result_chars=_MAX_TOOL_RESULT_CHARS,
|
||||||
)
|
)
|
||||||
|
|
||||||
assert AgentDefaults().max_concurrent_subagents == 4
|
|
||||||
assert mgr.max_concurrent_subagents == AgentDefaults().max_concurrent_subagents
|
assert mgr.max_concurrent_subagents == AgentDefaults().max_concurrent_subagents
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -142,7 +142,6 @@ def test_launcher_passes_the_canonical_model_preset_to_the_tui(
|
|||||||
assert captured["NANOBOT_TUI_BOOTSTRAP_URL"] == (
|
assert captured["NANOBOT_TUI_BOOTSTRAP_URL"] == (
|
||||||
"http://127.0.0.1:8765/webui/bootstrap"
|
"http://127.0.0.1:8765/webui/bootstrap"
|
||||||
)
|
)
|
||||||
assert captured["NANOBOT_TUI_HEALTH_URL"] == "http://127.0.0.1:18790/health"
|
|
||||||
assert captured["NANOBOT_TUI_BOOTSTRAP_SECRET"] == "bootstrap-secret"
|
assert captured["NANOBOT_TUI_BOOTSTRAP_SECRET"] == "bootstrap-secret"
|
||||||
assert "NANOBOT_TUI_WS_URL" not in captured
|
assert "NANOBOT_TUI_WS_URL" not in captured
|
||||||
assert "NANOBOT_TUI_API_TOKEN" not in captured
|
assert "NANOBOT_TUI_API_TOKEN" not in captured
|
||||||
|
|||||||
@@ -110,7 +110,7 @@ class TestMidTurnCommandDispatchedDirectly:
|
|||||||
loop = MagicMock()
|
loop = MagicMock()
|
||||||
loop.sessions = MagicMock()
|
loop.sessions = MagicMock()
|
||||||
loop.sessions.get_or_create = MagicMock(return_value=MagicMock(
|
loop.sessions.get_or_create = MagicMock(return_value=MagicMock(
|
||||||
messages=[], last_archived=0, clear=MagicMock(),
|
messages=[], last_consolidated=0, clear=MagicMock(),
|
||||||
))
|
))
|
||||||
loop.sessions.save = MagicMock()
|
loop.sessions.save = MagicMock()
|
||||||
loop.sessions.invalidate = MagicMock()
|
loop.sessions.invalidate = MagicMock()
|
||||||
|
|||||||
@@ -63,3 +63,9 @@ def test_explicit_provider_import_still_works(monkeypatch) -> None:
|
|||||||
finally:
|
finally:
|
||||||
monkeypatch.undo()
|
monkeypatch.undo()
|
||||||
setattr(sys.modules["nanobot"], "providers", original_package)
|
setattr(sys.modules["nanobot"], "providers", original_package)
|
||||||
|
|
||||||
|
|
||||||
|
def test_openai_codex_supports_progress_deltas() -> None:
|
||||||
|
from nanobot.providers.openai_codex_provider import OpenAICodexProvider
|
||||||
|
|
||||||
|
assert OpenAICodexProvider.supports_progress_deltas is True
|
||||||
|
|||||||
@@ -57,41 +57,9 @@ def test_loaded_corrupt_offset_keeps_messages(tmp_path: Path):
|
|||||||
def test_valid_offset_is_preserved():
|
def test_valid_offset_is_preserved():
|
||||||
session = _session(10, 4)
|
session = _session(10, 4)
|
||||||
assert session.last_consolidated == 4
|
assert session.last_consolidated == 4
|
||||||
assert session.last_archived == 4
|
|
||||||
assert len(session.get_history()) == 8
|
assert len(session.get_history()) == 8
|
||||||
|
|
||||||
|
|
||||||
def test_last_archived_field_migrates_with_legacy_alias(tmp_path: Path):
|
|
||||||
manager = SessionManager(tmp_path)
|
|
||||||
path = manager._get_session_path("chan:chat")
|
|
||||||
path.parent.mkdir(parents=True, exist_ok=True)
|
|
||||||
messages = [
|
|
||||||
{"role": "user", "content": "first"},
|
|
||||||
{"role": "assistant", "content": "second"},
|
|
||||||
]
|
|
||||||
path.write_text(
|
|
||||||
"\n".join([
|
|
||||||
json.dumps({
|
|
||||||
"_type": "metadata",
|
|
||||||
"key": "chan:chat",
|
|
||||||
"metadata": {},
|
|
||||||
"last_archived": 1,
|
|
||||||
}),
|
|
||||||
*(json.dumps(message) for message in messages),
|
|
||||||
]) + "\n",
|
|
||||||
encoding="utf-8",
|
|
||||||
)
|
|
||||||
|
|
||||||
session = manager.get_or_create("chan:chat")
|
|
||||||
|
|
||||||
assert session.last_archived == 1
|
|
||||||
assert session.last_consolidated == 1
|
|
||||||
manager.save(session)
|
|
||||||
metadata = json.loads(path.read_text(encoding="utf-8").splitlines()[0])
|
|
||||||
assert metadata["last_archived"] == 1
|
|
||||||
assert metadata["last_consolidated"] == 1
|
|
||||||
|
|
||||||
|
|
||||||
def test_loaded_null_metadata_becomes_empty_dict(tmp_path: Path):
|
def test_loaded_null_metadata_becomes_empty_dict(tmp_path: Path):
|
||||||
"""Session jsonl metadata:null must load as {} so agent .pop/.get work."""
|
"""Session jsonl metadata:null must load as {} so agent .pop/.get work."""
|
||||||
manager = SessionManager(tmp_path)
|
manager = SessionManager(tmp_path)
|
||||||
|
|||||||
@@ -24,12 +24,7 @@ class TestMessageToolSuppressLogic:
|
|||||||
"""Final reply suppressed only when message tool sends to the same target."""
|
"""Final reply suppressed only when message tool sends to the same target."""
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
@pytest.mark.parametrize("ephemeral", [False, True])
|
async def test_suppress_when_sent_to_same_target(self, tmp_path: Path) -> None:
|
||||||
async def test_suppress_when_sent_to_same_target(
|
|
||||||
self,
|
|
||||||
tmp_path: Path,
|
|
||||||
ephemeral: bool,
|
|
||||||
) -> None:
|
|
||||||
loop = _make_loop(tmp_path)
|
loop = _make_loop(tmp_path)
|
||||||
tool_call = ToolCallRequest(
|
tool_call = ToolCallRequest(
|
||||||
id="call1", name="message",
|
id="call1", name="message",
|
||||||
@@ -48,7 +43,7 @@ class TestMessageToolSuppressLogic:
|
|||||||
mt.set_send_callback(AsyncMock(side_effect=lambda m: sent.append(m)))
|
mt.set_send_callback(AsyncMock(side_effect=lambda m: sent.append(m)))
|
||||||
|
|
||||||
msg = InboundMessage(channel="feishu", sender_id="user1", chat_id="chat123", content="Send")
|
msg = InboundMessage(channel="feishu", sender_id="user1", chat_id="chat123", content="Send")
|
||||||
result = await loop._process_message(msg, ephemeral=ephemeral)
|
result = await loop._process_message(msg)
|
||||||
|
|
||||||
assert len(sent) == 1
|
assert len(sent) == 1
|
||||||
assert result is None # suppressed
|
assert result is None # suppressed
|
||||||
@@ -92,34 +87,6 @@ class TestMessageToolSuppressLogic:
|
|||||||
assert result is not None
|
assert result is not None
|
||||||
assert "Hello" in result.content
|
assert "Hello" in result.content
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
|
||||||
async def test_internal_message_check_keeps_final_response(self, tmp_path: Path) -> None:
|
|
||||||
loop = _make_loop(tmp_path)
|
|
||||||
tool_call = ToolCallRequest(
|
|
||||||
id="call1", name="message",
|
|
||||||
arguments={"content": "all clear", "channel": "feishu", "chat_id": "chat123"},
|
|
||||||
)
|
|
||||||
calls = iter([
|
|
||||||
LLMResponse(content="", tool_calls=[tool_call]),
|
|
||||||
LLMResponse(content="Heartbeat summary", tool_calls=[]),
|
|
||||||
])
|
|
||||||
loop.provider.chat_with_retry = AsyncMock(side_effect=lambda *a, **kw: next(calls))
|
|
||||||
loop.tools.get_definitions = MagicMock(return_value=[])
|
|
||||||
|
|
||||||
mt = loop.tools.get("message")
|
|
||||||
assert isinstance(mt, MessageTool)
|
|
||||||
token = mt.set_suppress_delivery(True)
|
|
||||||
try:
|
|
||||||
msg = InboundMessage(
|
|
||||||
channel="feishu", sender_id="user1", chat_id="chat123", content="Check",
|
|
||||||
)
|
|
||||||
result = await loop._process_message(msg)
|
|
||||||
finally:
|
|
||||||
mt.reset_suppress_delivery(token)
|
|
||||||
|
|
||||||
assert result is not None
|
|
||||||
assert result.content == "Heartbeat summary"
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_injected_followup_with_message_tool_does_not_emit_empty_fallback(
|
async def test_injected_followup_with_message_tool_does_not_emit_empty_fallback(
|
||||||
self, tmp_path: Path
|
self, tmp_path: Path
|
||||||
@@ -187,7 +154,22 @@ class TestMessageToolSuppressLogic:
|
|||||||
('read foo.txt', True),
|
('read foo.txt', True),
|
||||||
]
|
]
|
||||||
|
|
||||||
class TestMessageToolSchema:
|
class TestMessageToolTurnTracking:
|
||||||
|
|
||||||
|
def test_sent_in_turn_tracks_same_target(self) -> None:
|
||||||
|
tool = MessageTool()
|
||||||
|
from nanobot.agent.tools.context import RequestContext, request_context
|
||||||
|
|
||||||
|
with request_context(RequestContext(channel="feishu", chat_id="chat1")):
|
||||||
|
assert not tool._sent_in_turn
|
||||||
|
tool._sent_in_turn = True
|
||||||
|
assert tool._sent_in_turn
|
||||||
|
|
||||||
|
def test_start_turn_resets(self) -> None:
|
||||||
|
tool = MessageTool()
|
||||||
|
tool._sent_in_turn = True
|
||||||
|
tool.start_turn()
|
||||||
|
assert not tool._sent_in_turn
|
||||||
|
|
||||||
def test_schema_discourages_current_chat_replies(self) -> None:
|
def test_schema_discourages_current_chat_replies(self) -> None:
|
||||||
tool = MessageTool()
|
tool = MessageTool()
|
||||||
|
|||||||
@@ -59,7 +59,6 @@ def test_tool_context_has_required_fields():
|
|||||||
"config", "workspace", "bus", "subagent_manager",
|
"config", "workspace", "bus", "subagent_manager",
|
||||||
"cron_service", "exec_session_manager", "file_state_store",
|
"cron_service", "exec_session_manager", "file_state_store",
|
||||||
"provider_snapshot_loader", "image_generation_provider_configs", "timezone",
|
"provider_snapshot_loader", "image_generation_provider_configs", "timezone",
|
||||||
"runtime_control",
|
|
||||||
}
|
}
|
||||||
assert required <= field_names
|
assert required <= field_names
|
||||||
|
|
||||||
@@ -72,7 +71,6 @@ def test_tool_context_defaults():
|
|||||||
assert ctx.exec_session_manager is None
|
assert ctx.exec_session_manager is None
|
||||||
assert ctx.provider_snapshot_loader is None
|
assert ctx.provider_snapshot_loader is None
|
||||||
assert ctx.image_generation_provider_configs is None
|
assert ctx.image_generation_provider_configs is None
|
||||||
assert ctx.runtime_control is None
|
|
||||||
assert ctx.timezone == "UTC"
|
assert ctx.timezone == "UTC"
|
||||||
|
|
||||||
|
|
||||||
@@ -93,7 +91,6 @@ def test_discover_finds_concrete_tools():
|
|||||||
assert "ExecTool" in class_names
|
assert "ExecTool" in class_names
|
||||||
assert "CliAppsTool" in class_names
|
assert "CliAppsTool" in class_names
|
||||||
assert "MessageTool" in class_names
|
assert "MessageTool" in class_names
|
||||||
assert "MyTool" in class_names
|
|
||||||
assert "SpawnTool" in class_names
|
assert "SpawnTool" in class_names
|
||||||
assert "ExecSessionTool" in class_names
|
assert "ExecSessionTool" in class_names
|
||||||
|
|
||||||
@@ -376,26 +373,12 @@ def test_my_tool_enabled():
|
|||||||
from nanobot.agent.tools.self import MyTool
|
from nanobot.agent.tools.self import MyTool
|
||||||
mock_config = MagicMock()
|
mock_config = MagicMock()
|
||||||
mock_config.my.enable = True
|
mock_config.my.enable = True
|
||||||
ctx = ToolContext(
|
ctx = ToolContext(config=mock_config, workspace="/tmp")
|
||||||
config=mock_config,
|
|
||||||
workspace="/tmp",
|
|
||||||
runtime_control=MagicMock(),
|
|
||||||
)
|
|
||||||
assert MyTool.enabled(ctx) is True
|
assert MyTool.enabled(ctx) is True
|
||||||
mock_config.my.enable = False
|
mock_config.my.enable = False
|
||||||
assert MyTool.enabled(ctx) is False
|
assert MyTool.enabled(ctx) is False
|
||||||
|
|
||||||
|
|
||||||
def test_my_tool_requires_runtime_control():
|
|
||||||
from nanobot.agent.tools.self import MyTool
|
|
||||||
|
|
||||||
mock_config = MagicMock()
|
|
||||||
mock_config.my.enable = True
|
|
||||||
ctx = ToolContext(config=mock_config, workspace="/tmp")
|
|
||||||
|
|
||||||
assert MyTool.enabled(ctx) is False
|
|
||||||
|
|
||||||
|
|
||||||
def test_mcp_wrappers_not_discoverable():
|
def test_mcp_wrappers_not_discoverable():
|
||||||
from nanobot.agent.tools.mcp import MCPPromptWrapper, MCPResourceWrapper, MCPToolWrapper
|
from nanobot.agent.tools.mcp import MCPPromptWrapper, MCPResourceWrapper, MCPToolWrapper
|
||||||
assert MCPToolWrapper._plugin_discoverable is False
|
assert MCPToolWrapper._plugin_discoverable is False
|
||||||
@@ -428,7 +411,6 @@ def test_loader_registers_same_tools_as_old_hardcoded():
|
|||||||
mock_config.web.user_agent = None
|
mock_config.web.user_agent = None
|
||||||
mock_config.image_generation.enabled = False
|
mock_config.image_generation.enabled = False
|
||||||
mock_config.my.enable = True
|
mock_config.my.enable = True
|
||||||
mock_config.my.allow_set = False
|
|
||||||
|
|
||||||
ctx = ToolContext(
|
ctx = ToolContext(
|
||||||
config=mock_config,
|
config=mock_config,
|
||||||
@@ -437,7 +419,6 @@ def test_loader_registers_same_tools_as_old_hardcoded():
|
|||||||
subagent_manager=MagicMock(),
|
subagent_manager=MagicMock(),
|
||||||
cron_service=MagicMock(),
|
cron_service=MagicMock(),
|
||||||
timezone="UTC",
|
timezone="UTC",
|
||||||
runtime_control=MagicMock(),
|
|
||||||
)
|
)
|
||||||
registry = ToolRegistry()
|
registry = ToolRegistry()
|
||||||
loader = ToolLoader()
|
loader = ToolLoader()
|
||||||
@@ -448,7 +429,6 @@ def test_loader_registers_same_tools_as_old_hardcoded():
|
|||||||
"find_files", "grep", "exec", "exec_session", "list_exec_sessions",
|
"find_files", "grep", "exec", "exec_session", "list_exec_sessions",
|
||||||
"web_search", "web_fetch",
|
"web_search", "web_fetch",
|
||||||
"message", "spawn", "cron",
|
"message", "spawn", "cron",
|
||||||
"my",
|
|
||||||
}
|
}
|
||||||
actual = set(registered)
|
actual = set(registered)
|
||||||
assert expected <= actual, f"Missing tools: {expected - actual}"
|
assert expected <= actual, f"Missing tools: {expected - actual}"
|
||||||
|
|||||||
@@ -14,6 +14,7 @@ def test_session_context_separates_archive_progress_from_replay() -> None:
|
|||||||
session = Session(
|
session = Session(
|
||||||
key="websocket:context",
|
key="websocket:context",
|
||||||
messages=messages,
|
messages=messages,
|
||||||
|
last_consolidated=2,
|
||||||
metadata={
|
metadata={
|
||||||
"_last_summary": {
|
"_last_summary": {
|
||||||
"text": "The archived conversation settled the old question.",
|
"text": "The archived conversation settled the old question.",
|
||||||
@@ -22,7 +23,6 @@ def test_session_context_separates_archive_progress_from_replay() -> None:
|
|||||||
},
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
session.last_archived = 2
|
|
||||||
replay = session.get_history(max_messages=0, include_runtime_context=False)
|
replay = session.get_history(max_messages=0, include_runtime_context=False)
|
||||||
replay_tokens = sum(estimate_message_tokens(message) for message in replay)
|
replay_tokens = sum(estimate_message_tokens(message) for message in replay)
|
||||||
summary_tokens = estimate_message_tokens(
|
summary_tokens = estimate_message_tokens(
|
||||||
|
|||||||
+1
-4
@@ -29,10 +29,7 @@ Changes reuse the gateway's normal model command and workspace policy checks.
|
|||||||
|
|
||||||
When you scroll away from the latest output, the scrollbar and `Ctrl+End` hint appear only until
|
When you scroll away from the latest output, the scrollbar and `Ctrl+End` hint appear only until
|
||||||
you return to the bottom. Large pastes are represented by a short editable placeholder in the
|
you return to the bottom. Large pastes are represented by a short editable placeholder in the
|
||||||
composer; nanobot sends the original text unchanged. Press `Ctrl+V` or `Alt+V` while the composer
|
composer; nanobot sends the original text unchanged.
|
||||||
is focused to attach an image from the system clipboard. Image bytes stay behind removable
|
|
||||||
`[Image #n]` placeholders until the message is sent; each placeholder behaves as one unit, and
|
|
||||||
deleting it removes its image.
|
|
||||||
|
|
||||||
While nanobot is working, the composer prompt becomes
|
While nanobot is working, the composer prompt becomes
|
||||||
`Enter send now · Tab send next`; narrow terminals shorten it to `Enter now · Tab next`.
|
`Enter send now · Tab send next`; narrow terminals shorten it to `Enter now · Tab next`.
|
||||||
|
|||||||
+10
-425
@@ -1,12 +1,5 @@
|
|||||||
import { afterEach, describe, expect, test } from "bun:test"
|
import { afterEach, describe, expect, test } from "bun:test"
|
||||||
import {
|
import { BoxRenderable, CliRenderEvents, TextareaRenderable, TextRenderable } from "@opentui/core"
|
||||||
BoxRenderable,
|
|
||||||
CliRenderEvents,
|
|
||||||
StyledText,
|
|
||||||
TextareaRenderable,
|
|
||||||
TextAttributes,
|
|
||||||
TextRenderable,
|
|
||||||
} from "@opentui/core"
|
|
||||||
import {
|
import {
|
||||||
MockTreeSitterClient,
|
MockTreeSitterClient,
|
||||||
createTestRenderer,
|
createTestRenderer,
|
||||||
@@ -22,8 +15,7 @@ import type {
|
|||||||
WorkspaceScopePayload,
|
WorkspaceScopePayload,
|
||||||
} from "./protocol"
|
} from "./protocol"
|
||||||
import type { HostAgentState, HostMetadata, TuiHost } from "./host"
|
import type { HostAgentState, HostMetadata, TuiHost } from "./host"
|
||||||
import type { ClipboardImageReader } from "./clipboard-image"
|
import type { Transcript } from "./transcript"
|
||||||
import { userMessageText, type Transcript } from "./transcript"
|
|
||||||
|
|
||||||
const options: AppOptions = {
|
const options: AppOptions = {
|
||||||
wsUrl: "ws://localhost.invalid/ws",
|
wsUrl: "ws://localhost.invalid/ws",
|
||||||
@@ -59,20 +51,6 @@ test("formats a reusable session ID after exit", () => {
|
|||||||
)
|
)
|
||||||
})
|
})
|
||||||
|
|
||||||
test("projects image media as stable placeholders without exposing filenames", () => {
|
|
||||||
expect(userMessageText("What is this?", [
|
|
||||||
{ name: "clipboard-image-2.png" },
|
|
||||||
{ kind: "image", name: "screenshot.png" },
|
|
||||||
{ kind: "file", name: "report.pdf" },
|
|
||||||
])).toBe([
|
|
||||||
"What is this? [Image #2] [Image #1]",
|
|
||||||
"Attachments: report.pdf",
|
|
||||||
].join("\n"))
|
|
||||||
expect(userMessageText("What is this?", [
|
|
||||||
{ name: "clipboard-image-1.png" },
|
|
||||||
], "What is this? [Image #1]")).toBe("What is this? [Image #1]")
|
|
||||||
})
|
|
||||||
|
|
||||||
function contrastRatio(foreground: string, background: string): number {
|
function contrastRatio(foreground: string, background: string): number {
|
||||||
const luminance = (color: string) => {
|
const luminance = (color: string) => {
|
||||||
const channel = (offset: number) => {
|
const channel = (offset: number) => {
|
||||||
@@ -204,7 +182,7 @@ describe("NanobotTui layout", () => {
|
|||||||
expect(setup.renderer.height).toBe(height)
|
expect(setup.renderer.height).toBe(height)
|
||||||
expect(occurrences(frame, "Ask nanobot anything")).toBe(1)
|
expect(occurrences(frame, "Ask nanobot anything")).toBe(1)
|
||||||
expect(occurrences(frame, "Ready")).toBe(0)
|
expect(occurrences(frame, "Ready")).toBe(0)
|
||||||
expect(occurrences(frame, "Getting ready…")).toBe(1)
|
expect(occurrences(frame, "Connecting…")).toBe(1)
|
||||||
expect(occurrences(frame, "nanobot · test/model")).toBe(1)
|
expect(occurrences(frame, "nanobot · test/model")).toBe(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -328,278 +306,6 @@ describe("NanobotTui layout", () => {
|
|||||||
expect(ui.composer.plainText).toBe("")
|
expect(ui.composer.plainText).toBe("")
|
||||||
})
|
})
|
||||||
|
|
||||||
test("pastes clipboard images into removable placeholders and sends their data", async () => {
|
|
||||||
const sent: string[] = []
|
|
||||||
const sentOptions: MessageOptions[] = []
|
|
||||||
let disposed = false
|
|
||||||
const clipboard: ClipboardImageReader = {
|
|
||||||
read: async () => ({
|
|
||||||
mimeType: "image/png",
|
|
||||||
dataUrl: "data:image/png;base64,AAEC/w==",
|
|
||||||
}),
|
|
||||||
dispose: async () => { disposed = true },
|
|
||||||
}
|
|
||||||
setup = await createRenderer({ width: 72, height: 20, screenMode: "alternate-screen" })
|
|
||||||
const transport = client(sent, [], [], sentOptions)
|
|
||||||
const recordSend = transport.send
|
|
||||||
transport.send = (content, messageOptions) => {
|
|
||||||
recordSend(content, messageOptions)
|
|
||||||
return `image-turn-${sent.length}`
|
|
||||||
}
|
|
||||||
const app = NanobotTui.mount(
|
|
||||||
setup.renderer,
|
|
||||||
options,
|
|
||||||
transport,
|
|
||||||
new MockTreeSitterClient({ autoResolveTimeout: 0 }),
|
|
||||||
undefined,
|
|
||||||
clipboard,
|
|
||||||
)
|
|
||||||
app.accept({ event: "attached", chat_id: "chat" })
|
|
||||||
await waitUntil(() => (app as unknown as { ready: boolean }).ready)
|
|
||||||
const ui = app as unknown as {
|
|
||||||
composer: TextareaRenderable
|
|
||||||
draft: { imageCount: number }
|
|
||||||
promptHistory: string[]
|
|
||||||
status: { plainText: string }
|
|
||||||
transcript: {
|
|
||||||
userMessages: Set<{ renderable: TextRenderable }>
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
setup.mockInput.pressKey("v", { ctrl: true })
|
|
||||||
await waitUntil(() => ui.composer.plainText === "[Image #1] ")
|
|
||||||
expect(ui.status.plainText).toContain("Pasted Image #1")
|
|
||||||
const placeholderStyle = ui.composer.syntaxStyle?.getStyle("image.placeholder")
|
|
||||||
expect(placeholderStyle?.bold).toBeTrue()
|
|
||||||
expect(placeholderStyle?.fg?.toInts().slice(0, 3)).toEqual([239, 142, 48])
|
|
||||||
const placeholderStyleId = ui.composer.syntaxStyle?.getStyleId("image.placeholder")
|
|
||||||
if (placeholderStyleId === null || placeholderStyleId === undefined) {
|
|
||||||
throw new Error("image placeholder style was not registered")
|
|
||||||
}
|
|
||||||
expect(ui.composer.getLineHighlights(0)).toEqual([{
|
|
||||||
start: 0,
|
|
||||||
end: 10,
|
|
||||||
styleId: placeholderStyleId,
|
|
||||||
priority: 100,
|
|
||||||
hlRef: 0,
|
|
||||||
}])
|
|
||||||
ui.composer.setText("")
|
|
||||||
await waitUntil(() => ui.draft.imageCount === 0)
|
|
||||||
expect(ui.composer.getLineHighlights(0)).toEqual([])
|
|
||||||
|
|
||||||
setup.mockInput.pressKey("v", { ctrl: true })
|
|
||||||
await waitUntil(() => ui.composer.plainText === "[Image #1] ")
|
|
||||||
await setup.mockInput.typeText("[Image #1]")
|
|
||||||
ui.composer.submit()
|
|
||||||
await waitUntil(() => ui.status.plainText.includes("Duplicate image placeholder"))
|
|
||||||
expect(sent).toEqual([])
|
|
||||||
ui.composer.setText("[Image #1]")
|
|
||||||
ui.composer.submit()
|
|
||||||
await waitUntil(() => sent.length === 1)
|
|
||||||
expect(sent).toEqual([""])
|
|
||||||
expect(ui.promptHistory).toEqual([])
|
|
||||||
expect(sentOptions[0]?.media).toEqual([{
|
|
||||||
data_url: "data:image/png;base64,AAEC/w==",
|
|
||||||
name: "clipboard-image-1.png",
|
|
||||||
}])
|
|
||||||
expect(sentOptions[0]).not.toHaveProperty("displayContent")
|
|
||||||
await setup.flush()
|
|
||||||
const frame = setup.captureCharFrame()
|
|
||||||
expect(frame).toContain("[Image #1]")
|
|
||||||
expect(frame).not.toContain("clipboard-image-1.png")
|
|
||||||
const userContent = [...ui.transcript.userMessages].at(-1)?.renderable.content
|
|
||||||
expect(userContent).toBeInstanceOf(StyledText)
|
|
||||||
const imageChunk = (userContent as StyledText).chunks.find(({ text }) => text === "[Image #1]")
|
|
||||||
expect(imageChunk?.attributes).toBe(TextAttributes.BOLD)
|
|
||||||
expect(imageChunk?.fg?.toInts().slice(0, 3)).toEqual([239, 142, 48])
|
|
||||||
|
|
||||||
await setup.mockInput.typeText("这是什么? ")
|
|
||||||
setup.mockInput.pressKey("v", { ctrl: true })
|
|
||||||
await waitUntil(() => ui.status.plainText.includes("Pasted Image #1"), 3_000)
|
|
||||||
expect(ui.composer.plainText).toBe("这是什么? [Image #1] ")
|
|
||||||
setup.mockInput.pressTab()
|
|
||||||
expect(ui.status.plainText).toContain("Images cannot be queued")
|
|
||||||
expect(ui.composer.plainText).toBe("这是什么? [Image #1] ")
|
|
||||||
ui.composer.submit()
|
|
||||||
await waitUntil(() => sent.length === 2)
|
|
||||||
expect(sent[1]).toBe("这是什么?")
|
|
||||||
expect(sentOptions[1]?.media).toHaveLength(1)
|
|
||||||
expect(sentOptions[1]).not.toHaveProperty("displayContent")
|
|
||||||
await setup.flush()
|
|
||||||
expect(setup.captureCharFrame()).toContain("这是什么? [Image #1]")
|
|
||||||
|
|
||||||
setup.renderer.destroy()
|
|
||||||
expect(disposed).toBeTrue()
|
|
||||||
})
|
|
||||||
|
|
||||||
test("keeps image placeholders atomic for cursor movement and deletion", async () => {
|
|
||||||
const clipboard: ClipboardImageReader = {
|
|
||||||
read: async () => ({
|
|
||||||
mimeType: "image/png",
|
|
||||||
dataUrl: "data:image/png;base64,AAEC/w==",
|
|
||||||
}),
|
|
||||||
dispose: async () => undefined,
|
|
||||||
}
|
|
||||||
setup = await createRenderer({ width: 72, height: 20, screenMode: "alternate-screen" })
|
|
||||||
const app = NanobotTui.mount(
|
|
||||||
setup.renderer,
|
|
||||||
options,
|
|
||||||
client(),
|
|
||||||
new MockTreeSitterClient({ autoResolveTimeout: 0 }),
|
|
||||||
undefined,
|
|
||||||
clipboard,
|
|
||||||
)
|
|
||||||
app.accept({ event: "attached", chat_id: "chat" })
|
|
||||||
await waitUntil(() => (app as unknown as { ready: boolean }).ready)
|
|
||||||
const ui = app as unknown as {
|
|
||||||
composer: TextareaRenderable
|
|
||||||
draft: { imageCount: number }
|
|
||||||
status: { plainText: string }
|
|
||||||
}
|
|
||||||
|
|
||||||
setup.mockInput.pressKey("v", { ctrl: true })
|
|
||||||
await waitUntil(() => ui.composer.plainText === "[Image #1] ")
|
|
||||||
await setup.flush()
|
|
||||||
await setup.mockMouse.click(ui.composer.x + 5, ui.composer.y)
|
|
||||||
expect(ui.composer.cursorOffset > 0 && ui.composer.cursorOffset < 10).toBeFalse()
|
|
||||||
ui.composer.cursorOffset = 0
|
|
||||||
setup.mockInput.pressArrow("right")
|
|
||||||
await waitUntil(() => ui.composer.cursorOffset === 10)
|
|
||||||
setup.mockInput.pressArrow("left")
|
|
||||||
await waitUntil(() => ui.composer.cursorOffset === 0)
|
|
||||||
|
|
||||||
setup.mockInput.pressArrow("right", { shift: true })
|
|
||||||
await waitUntil(() => ui.composer.cursorOffset === 10)
|
|
||||||
await setup.mockInput.typeText("replacement")
|
|
||||||
await waitUntil(() => ui.draft.imageCount === 0)
|
|
||||||
expect(ui.composer.plainText).toContain("replacement")
|
|
||||||
expect(ui.composer.plainText).not.toContain("Image #1")
|
|
||||||
|
|
||||||
ui.composer.setText("")
|
|
||||||
setup.mockInput.pressKey("v", { ctrl: true })
|
|
||||||
await waitUntil(() => ui.composer.plainText === "[Image #1] ")
|
|
||||||
ui.composer.cursorOffset = 0
|
|
||||||
setup.mockInput.pressKey("DELETE")
|
|
||||||
await waitUntil(() => ui.draft.imageCount === 0)
|
|
||||||
expect(ui.composer.plainText.trim()).toBe("")
|
|
||||||
expect(ui.status.plainText).toContain("Removed Image #1")
|
|
||||||
|
|
||||||
ui.composer.setText("")
|
|
||||||
setup.mockInput.pressKey("v", { ctrl: true })
|
|
||||||
await waitUntil(() => ui.composer.plainText === "[Image #1] ")
|
|
||||||
ui.composer.cursorOffset = 10
|
|
||||||
setup.mockInput.pressBackspace()
|
|
||||||
await waitUntil(() => ui.draft.imageCount === 0)
|
|
||||||
expect(ui.composer.plainText.trim()).toBe("")
|
|
||||||
|
|
||||||
ui.composer.setText("")
|
|
||||||
setup.mockInput.pressKey("v", { ctrl: true })
|
|
||||||
await waitUntil(() => ui.composer.plainText === "[Image #1] ")
|
|
||||||
ui.composer.setText("Image #1] ")
|
|
||||||
await waitUntil(() => ui.draft.imageCount === 0)
|
|
||||||
expect(ui.composer.plainText.trim()).toBe("")
|
|
||||||
})
|
|
||||||
|
|
||||||
test("keeps clipboard failures visible while an agent turn is active", async () => {
|
|
||||||
const sent: string[] = []
|
|
||||||
const clipboard: ClipboardImageReader = {
|
|
||||||
read: async () => { throw new Error("No image in clipboard") },
|
|
||||||
dispose: async () => undefined,
|
|
||||||
}
|
|
||||||
setup = await createRenderer({ width: 72, height: 20, screenMode: "alternate-screen" })
|
|
||||||
const app = NanobotTui.mount(
|
|
||||||
setup.renderer,
|
|
||||||
options,
|
|
||||||
client(sent),
|
|
||||||
new MockTreeSitterClient({ autoResolveTimeout: 0 }),
|
|
||||||
undefined,
|
|
||||||
clipboard,
|
|
||||||
)
|
|
||||||
app.accept({ event: "attached", chat_id: "chat" })
|
|
||||||
await waitUntil(() => (app as unknown as { ready: boolean }).ready)
|
|
||||||
const composer = (app as unknown as { composer: TextareaRenderable }).composer
|
|
||||||
composer.setText("start")
|
|
||||||
composer.submit()
|
|
||||||
await waitUntil(() => sent.length === 1)
|
|
||||||
|
|
||||||
setup.mockInput.pressKey("v", { ctrl: true })
|
|
||||||
await waitUntil(() => setup?.captureCharFrame().includes("No image in clipboard") === true)
|
|
||||||
})
|
|
||||||
|
|
||||||
test("keeps image placeholders out of command arguments", async () => {
|
|
||||||
const sent: string[] = []
|
|
||||||
const clipboard: ClipboardImageReader = {
|
|
||||||
read: async () => ({
|
|
||||||
mimeType: "image/png",
|
|
||||||
dataUrl: "data:image/png;base64,AAEC/w==",
|
|
||||||
}),
|
|
||||||
dispose: async () => undefined,
|
|
||||||
}
|
|
||||||
setup = await createRenderer({ width: 72, height: 20, screenMode: "alternate-screen" })
|
|
||||||
const app = NanobotTui.mount(
|
|
||||||
setup.renderer,
|
|
||||||
options,
|
|
||||||
client(sent),
|
|
||||||
new MockTreeSitterClient({ autoResolveTimeout: 0 }),
|
|
||||||
undefined,
|
|
||||||
clipboard,
|
|
||||||
)
|
|
||||||
app.accept({ event: "attached", chat_id: "chat" })
|
|
||||||
await waitUntil(() => (app as unknown as { ready: boolean }).ready)
|
|
||||||
const ui = app as unknown as {
|
|
||||||
composer: TextareaRenderable
|
|
||||||
status: { plainText: string }
|
|
||||||
commandMenu: { setCommands(commands: SlashCommand[]): void }
|
|
||||||
}
|
|
||||||
ui.commandMenu.setCommands([{
|
|
||||||
command: "/model",
|
|
||||||
title: "Model",
|
|
||||||
description: "Show or switch model presets",
|
|
||||||
argHint: "[preset]",
|
|
||||||
lifecycle: "side_channel",
|
|
||||||
acceptsArgs: true,
|
|
||||||
}])
|
|
||||||
|
|
||||||
await setup.mockInput.typeText("/model ")
|
|
||||||
setup.mockInput.pressKey("v", { ctrl: true })
|
|
||||||
await waitUntil(() => ui.composer.plainText === "/model [Image #1] ")
|
|
||||||
ui.composer.submit()
|
|
||||||
await waitUntil(() => ui.status.plainText.includes("Images cannot be used with commands"))
|
|
||||||
|
|
||||||
expect(sent).toEqual([])
|
|
||||||
expect(ui.composer.plainText).toBe("/model [Image #1] ")
|
|
||||||
})
|
|
||||||
|
|
||||||
test("ignores a clipboard result that finishes after the renderer is destroyed", async () => {
|
|
||||||
let resolveRead: ((image: {
|
|
||||||
mimeType: "image/png"
|
|
||||||
dataUrl: string
|
|
||||||
}) => void) | undefined
|
|
||||||
let disposed = false
|
|
||||||
const clipboard: ClipboardImageReader = {
|
|
||||||
read: () => new Promise((resolve) => { resolveRead = resolve }),
|
|
||||||
dispose: async () => { disposed = true },
|
|
||||||
}
|
|
||||||
setup = await createRenderer({ width: 72, height: 20, screenMode: "alternate-screen" })
|
|
||||||
const app = NanobotTui.mount(
|
|
||||||
setup.renderer,
|
|
||||||
options,
|
|
||||||
client(),
|
|
||||||
new MockTreeSitterClient({ autoResolveTimeout: 0 }),
|
|
||||||
undefined,
|
|
||||||
clipboard,
|
|
||||||
)
|
|
||||||
app.accept({ event: "attached", chat_id: "chat" })
|
|
||||||
await waitUntil(() => (app as unknown as { ready: boolean }).ready)
|
|
||||||
|
|
||||||
setup.mockInput.pressKey("v", { ctrl: true })
|
|
||||||
await waitUntil(() => resolveRead !== undefined)
|
|
||||||
setup.renderer.destroy()
|
|
||||||
resolveRead?.({ mimeType: "image/png", dataUrl: "data:image/png;base64,AAAA" })
|
|
||||||
await Bun.sleep(10)
|
|
||||||
expect(disposed).toBeTrue()
|
|
||||||
})
|
|
||||||
|
|
||||||
test("steers with Enter, queues with Tab, and restores queued text with Alt+Up", async () => {
|
test("steers with Enter, queues with Tab, and restores queued text with Alt+Up", async () => {
|
||||||
const sent: string[] = []
|
const sent: string[] = []
|
||||||
const sentOptions: MessageOptions[] = []
|
const sentOptions: MessageOptions[] = []
|
||||||
@@ -706,11 +412,6 @@ describe("NanobotTui layout", () => {
|
|||||||
turn_id: "remote-steer",
|
turn_id: "remote-steer",
|
||||||
active_turn_id: "remote-turn",
|
active_turn_id: "remote-turn",
|
||||||
starts_turn: false,
|
starts_turn: false,
|
||||||
media_urls: [{
|
|
||||||
kind: "image",
|
|
||||||
url: "/api/media/sig/image",
|
|
||||||
name: "clipboard-image-2.png",
|
|
||||||
}],
|
|
||||||
})
|
})
|
||||||
await setup.flush()
|
await setup.flush()
|
||||||
|
|
||||||
@@ -719,9 +420,6 @@ describe("NanobotTui layout", () => {
|
|||||||
expect(occurrences(frame, "hello from terminal A")).toBe(1)
|
expect(occurrences(frame, "hello from terminal A")).toBe(1)
|
||||||
expect(occurrences(frame, "Attachments: report.pdf")).toBe(1)
|
expect(occurrences(frame, "Attachments: report.pdf")).toBe(1)
|
||||||
expect(occurrences(frame, "one more remote detail")).toBe(1)
|
expect(occurrences(frame, "one more remote detail")).toBe(1)
|
||||||
expect(occurrences(frame, "[Image #2]")).toBe(1)
|
|
||||||
expect(frame).toContain("one more remote detail [Image #2]")
|
|
||||||
expect(frame).not.toContain("clipboard-image-2.png")
|
|
||||||
expect(state.activeTurn).toBeTrue()
|
expect(state.activeTurn).toBeTrue()
|
||||||
expect(state.activeTurnId).toBe("remote-turn")
|
expect(state.activeTurnId).toBe("remote-turn")
|
||||||
|
|
||||||
@@ -2092,26 +1790,19 @@ describe("NanobotTui layout", () => {
|
|||||||
composer: {
|
composer: {
|
||||||
backgroundColor: { intent: string; toInts(): number[] }
|
backgroundColor: { intent: string; toInts(): number[] }
|
||||||
textColor: { toInts(): number[] }
|
textColor: { toInts(): number[] }
|
||||||
syntaxStyle: { getStyle(name: string): { fg?: { toInts(): number[] } } | undefined } | null
|
|
||||||
}
|
}
|
||||||
transcript: {
|
transcript: {
|
||||||
markdown: Set<{ syntaxStyle: object }>
|
markdown: Set<{ syntaxStyle: object }>
|
||||||
frames: Set<{ borderColor: { toInts(): number[] } }>
|
frames: Set<{ borderColor: { toInts(): number[] } }>
|
||||||
userRows: Set<{ backgroundColor: { intent: string; toInts(): number[] } }>
|
userRows: Set<{ backgroundColor: { intent: string; toInts(): number[] } }>
|
||||||
userMessages: Set<{ renderable: TextRenderable }>
|
user(content: string): void
|
||||||
user(content: string, turnId?: string, media?: Array<{ kind: "image"; name: string }>): void
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
internals.transcript.user("Existing question", undefined, [{
|
internals.transcript.user("Existing question")
|
||||||
kind: "image",
|
|
||||||
name: "clipboard-image-1.png",
|
|
||||||
}])
|
|
||||||
const userRow = [...internals.transcript.userRows][0]
|
const userRow = [...internals.transcript.userRows][0]
|
||||||
const userMessage = [...internals.transcript.userMessages][0]
|
|
||||||
const markdown = [...internals.transcript.markdown][0]
|
const markdown = [...internals.transcript.markdown][0]
|
||||||
const sessionFrame = [...internals.transcript.frames][0]
|
const sessionFrame = [...internals.transcript.frames][0]
|
||||||
const darkSyntax = markdown?.syntaxStyle
|
const darkSyntax = markdown?.syntaxStyle
|
||||||
const darkComposerSyntax = internals.composer.syntaxStyle
|
|
||||||
|
|
||||||
expect(userRow?.backgroundColor.intent).toBe("default")
|
expect(userRow?.backgroundColor.intent).toBe("default")
|
||||||
|
|
||||||
@@ -2130,12 +1821,6 @@ describe("NanobotTui layout", () => {
|
|||||||
expect(sessionFrame?.borderColor.toInts().slice(0, 3)).toEqual([212, 212, 216])
|
expect(sessionFrame?.borderColor.toInts().slice(0, 3)).toEqual([212, 212, 216])
|
||||||
expect(userRow?.backgroundColor.toInts().slice(0, 3)).toEqual([240, 240, 240])
|
expect(userRow?.backgroundColor.toInts().slice(0, 3)).toEqual([240, 240, 240])
|
||||||
expect(markdown?.syntaxStyle).not.toBe(darkSyntax)
|
expect(markdown?.syntaxStyle).not.toBe(darkSyntax)
|
||||||
expect(internals.composer.syntaxStyle).not.toBe(darkComposerSyntax)
|
|
||||||
expect(internals.composer.syntaxStyle?.getStyle("image.placeholder")?.fg?.toInts().slice(0, 3))
|
|
||||||
.toEqual([185, 77, 11])
|
|
||||||
const recolored = userMessage?.renderable.content as StyledText
|
|
||||||
expect(recolored.chunks.find(({ text }) => text === "[Image #1]")?.fg?.toInts().slice(0, 3))
|
|
||||||
.toEqual([185, 77, 11])
|
|
||||||
})
|
})
|
||||||
|
|
||||||
test("distinguishes the composer with a quiet focus edge", async () => {
|
test("distinguishes the composer with a quiet focus edge", async () => {
|
||||||
@@ -2699,85 +2384,6 @@ describe("NanobotTui layout", () => {
|
|||||||
expect(state()).toBe(false)
|
expect(state()).toBe(false)
|
||||||
})
|
})
|
||||||
|
|
||||||
test("shows actionable connection states without implementation details", async () => {
|
|
||||||
setup = await createRenderer({ width: 100, height: 20, screenMode: "alternate-screen" })
|
|
||||||
const app = mount(setup)
|
|
||||||
const ui = app as unknown as {
|
|
||||||
status: TextRenderable
|
|
||||||
handleStatus(
|
|
||||||
status: "starting" | "connecting" | "connected" | "reconnecting" | "unavailable" | "error",
|
|
||||||
detail?: string,
|
|
||||||
info?: {
|
|
||||||
endpoint: string
|
|
||||||
attempt: number
|
|
||||||
elapsedMs: number
|
|
||||||
health?: "ready" | "degraded" | "unreachable"
|
|
||||||
},
|
|
||||||
): void
|
|
||||||
}
|
|
||||||
|
|
||||||
ui.handleStatus("starting", undefined, {
|
|
||||||
endpoint: "127.0.0.1:8769",
|
|
||||||
attempt: 1,
|
|
||||||
elapsedMs: 0,
|
|
||||||
})
|
|
||||||
expect(ui.status.plainText).toBe("Getting ready…")
|
|
||||||
|
|
||||||
ui.handleStatus("connecting")
|
|
||||||
expect(ui.status.plainText).toBe("Getting ready…")
|
|
||||||
|
|
||||||
ui.handleStatus("connected")
|
|
||||||
expect(ui.status.plainText).toBe("Getting ready…")
|
|
||||||
|
|
||||||
ui.handleStatus("error", "gateway sent an invalid event")
|
|
||||||
expect(ui.status.plainText).toBe("Getting ready…")
|
|
||||||
expect(ui.status.plainText).not.toContain("Unable")
|
|
||||||
|
|
||||||
ui.handleStatus("reconnecting", "connection closed", {
|
|
||||||
endpoint: "127.0.0.1:8769",
|
|
||||||
attempt: 2,
|
|
||||||
elapsedMs: 800,
|
|
||||||
})
|
|
||||||
expect(ui.status.plainText).toBe("Resuming…")
|
|
||||||
|
|
||||||
ui.handleStatus("reconnecting", "connection closed", {
|
|
||||||
endpoint: "127.0.0.1:8769",
|
|
||||||
attempt: 2,
|
|
||||||
elapsedMs: 900,
|
|
||||||
health: "degraded",
|
|
||||||
})
|
|
||||||
expect(ui.status.plainText).toBe("Resuming…")
|
|
||||||
|
|
||||||
ui.handleStatus("unavailable", "connection refused", {
|
|
||||||
endpoint: "127.0.0.1:8769",
|
|
||||||
attempt: 7,
|
|
||||||
elapsedMs: 3_200,
|
|
||||||
health: "degraded",
|
|
||||||
})
|
|
||||||
expect(ui.status.plainText).toBe("Still getting ready…")
|
|
||||||
expect(ui.status.plainText).not.toContain("Unable")
|
|
||||||
|
|
||||||
ui.handleStatus("unavailable", "connection refused", {
|
|
||||||
endpoint: "127.0.0.1:8769",
|
|
||||||
attempt: 8,
|
|
||||||
elapsedMs: 3_500,
|
|
||||||
health: "unreachable",
|
|
||||||
})
|
|
||||||
expect(ui.status.plainText).toBe("Nanobot is taking longer to respond…")
|
|
||||||
expect(ui.status.plainText).not.toContain("Unable")
|
|
||||||
|
|
||||||
ui.handleStatus("error", "gateway bootstrap failed: HTTP 401", {
|
|
||||||
endpoint: "127.0.0.1:8769",
|
|
||||||
attempt: 9,
|
|
||||||
elapsedMs: 3_800,
|
|
||||||
})
|
|
||||||
expect(ui.status.plainText).toBe("Nanobot unavailable · restart nanobot")
|
|
||||||
expect(ui.status.plainText).not.toContain("gateway")
|
|
||||||
expect(ui.status.plainText).not.toContain("127.0.0.1")
|
|
||||||
expect(ui.status.plainText).not.toContain("HTTP")
|
|
||||||
expect(ui.status.plainText).not.toContain("attempt")
|
|
||||||
})
|
|
||||||
|
|
||||||
test("replays events after asynchronous history hydration", async () => {
|
test("replays events after asynchronous history hydration", async () => {
|
||||||
setup = await createRenderer({ width: 80, height: 22, screenMode: "alternate-screen" })
|
setup = await createRenderer({ width: 80, height: 22, screenMode: "alternate-screen" })
|
||||||
const original = globalThis.fetch
|
const original = globalThis.fetch
|
||||||
@@ -2838,12 +2444,7 @@ describe("NanobotTui layout", () => {
|
|||||||
client(sent),
|
client(sent),
|
||||||
new MockTreeSitterClient({ autoResolveTimeout: 0 }),
|
new MockTreeSitterClient({ autoResolveTimeout: 0 }),
|
||||||
)
|
)
|
||||||
const ui = app as unknown as {
|
const composer = (app as unknown as { composer: TextareaRenderable }).composer
|
||||||
composer: TextareaRenderable
|
|
||||||
ready: boolean
|
|
||||||
status: TextRenderable
|
|
||||||
}
|
|
||||||
const composer = ui.composer
|
|
||||||
|
|
||||||
try {
|
try {
|
||||||
app.accept({ event: "attached", chat_id: "chat" })
|
app.accept({ event: "attached", chat_id: "chat" })
|
||||||
@@ -2851,18 +2452,16 @@ describe("NanobotTui layout", () => {
|
|||||||
app.accept({ event: "attached", chat_id: "chat" })
|
app.accept({ event: "attached", chat_id: "chat" })
|
||||||
composer.setText("sent during reconnect")
|
composer.setText("sent during reconnect")
|
||||||
composer.submit()
|
composer.submit()
|
||||||
await waitUntil(() => ui.status.plainText.includes("Not sent"))
|
await Bun.sleep(5)
|
||||||
|
|
||||||
expect(sent).toEqual([])
|
expect(sent).toEqual([])
|
||||||
expect(composer.plainText).toBe("sent during reconnect")
|
expect(composer.plainText).toBe("sent during reconnect")
|
||||||
expect(ui.status.plainText).toContain("Not sent · press Enter to retry when ready")
|
|
||||||
|
|
||||||
resolveReconnect(new Response(JSON.stringify({
|
resolveReconnect(new Response(JSON.stringify({
|
||||||
messages: [{ role: "assistant", content: "restored history" }],
|
messages: [{ role: "assistant", content: "restored history" }],
|
||||||
page: { has_more_before: false },
|
page: { has_more_before: false },
|
||||||
})))
|
})))
|
||||||
await waitUntil(() => ui.ready)
|
await waitUntil(() => (app as unknown as { ready: boolean }).ready)
|
||||||
expect(ui.status.plainText).toBe("Not sent · press Enter to retry")
|
|
||||||
composer.submit()
|
composer.submit()
|
||||||
await waitUntil(() => sent.length === 1)
|
await waitUntil(() => sent.length === 1)
|
||||||
await setup.flush()
|
await setup.flush()
|
||||||
@@ -2881,38 +2480,24 @@ describe("NanobotTui layout", () => {
|
|||||||
const app = mount(setup, sent)
|
const app = mount(setup, sent)
|
||||||
const composer = (app as unknown as { composer: TextareaRenderable }).composer
|
const composer = (app as unknown as { composer: TextareaRenderable }).composer
|
||||||
const connection = app as unknown as {
|
const connection = app as unknown as {
|
||||||
handleStatus(
|
handleStatus(status: "connecting" | "connected", detail?: string): void
|
||||||
status: "reconnecting" | "connected",
|
|
||||||
detail?: string,
|
|
||||||
info?: { endpoint: string; attempt: number; elapsedMs: number },
|
|
||||||
): void
|
|
||||||
}
|
}
|
||||||
|
|
||||||
app.accept({ event: "attached", chat_id: "chat" })
|
app.accept({ event: "attached", chat_id: "chat" })
|
||||||
await Bun.sleep(1)
|
await Bun.sleep(1)
|
||||||
connection.handleStatus("reconnecting", "connection closed", {
|
connection.handleStatus("connecting", "reconnecting")
|
||||||
endpoint: "127.0.0.1:8769",
|
|
||||||
attempt: 1,
|
|
||||||
elapsedMs: 0,
|
|
||||||
})
|
|
||||||
connection.handleStatus("connected")
|
connection.handleStatus("connected")
|
||||||
composer.setText("draft before attach")
|
composer.setText("draft before attach")
|
||||||
composer.submit()
|
composer.submit()
|
||||||
await Bun.sleep(5)
|
await Bun.sleep(5)
|
||||||
composer.submit()
|
|
||||||
await Bun.sleep(5)
|
|
||||||
|
|
||||||
expect(sent).toEqual([])
|
expect(sent).toEqual([])
|
||||||
expect(composer.plainText).toBe("draft before attach")
|
expect(composer.plainText).toBe("draft before attach")
|
||||||
|
|
||||||
app.accept({ event: "attached", chat_id: "chat" })
|
|
||||||
app.accept({ event: "attached", chat_id: "chat" })
|
app.accept({ event: "attached", chat_id: "chat" })
|
||||||
await waitUntil(() => (app as unknown as { ready: boolean }).ready)
|
await waitUntil(() => (app as unknown as { ready: boolean }).ready)
|
||||||
expect(sent).toEqual([])
|
|
||||||
composer.submit()
|
composer.submit()
|
||||||
await waitUntil(() => sent.length === 1)
|
await waitUntil(() => sent.length === 1)
|
||||||
app.accept({ event: "attached", chat_id: "chat" })
|
|
||||||
await Bun.sleep(5)
|
|
||||||
|
|
||||||
expect(sent).toEqual(["draft before attach"])
|
expect(sent).toEqual(["draft before attach"])
|
||||||
})
|
})
|
||||||
|
|||||||
+47
-307
@@ -20,9 +20,7 @@ import {
|
|||||||
|
|
||||||
import {
|
import {
|
||||||
NanobotClient,
|
NanobotClient,
|
||||||
connectionEndpoint,
|
|
||||||
fetchAvailableSkills,
|
fetchAvailableSkills,
|
||||||
fetchGatewayHealth,
|
|
||||||
fetchHistory,
|
fetchHistory,
|
||||||
fetchGatewayConnection,
|
fetchGatewayConnection,
|
||||||
fetchMentionCandidates,
|
fetchMentionCandidates,
|
||||||
@@ -31,7 +29,6 @@ import {
|
|||||||
fetchSlashCommands,
|
fetchSlashCommands,
|
||||||
type ApiReauthenticator,
|
type ApiReauthenticator,
|
||||||
type ConnectionStatus,
|
type ConnectionStatus,
|
||||||
type ConnectionStatusInfo,
|
|
||||||
type FileEditEvent,
|
type FileEditEvent,
|
||||||
type GatewayApiConnection,
|
type GatewayApiConnection,
|
||||||
type HistoryMessage,
|
type HistoryMessage,
|
||||||
@@ -65,11 +62,7 @@ import {
|
|||||||
type TranscriptNavigation,
|
type TranscriptNavigation,
|
||||||
type TranscriptTheme,
|
type TranscriptTheme,
|
||||||
} from "./transcript"
|
} from "./transcript"
|
||||||
import { ComposerDraft, MAX_DRAFT_IMAGES } from "./composer-draft"
|
import { ComposerDraft } from "./composer-draft"
|
||||||
import {
|
|
||||||
createClipboardImageReader,
|
|
||||||
type ClipboardImageReader,
|
|
||||||
} from "./clipboard-image"
|
|
||||||
import { BranchMenu, branchPoints } from "./branch-menu"
|
import { BranchMenu, branchPoints } from "./branch-menu"
|
||||||
import {
|
import {
|
||||||
MentionMenu,
|
MentionMenu,
|
||||||
@@ -100,7 +93,6 @@ interface AppOptions {
|
|||||||
wsUrl?: string
|
wsUrl?: string
|
||||||
bootstrapUrl?: string
|
bootstrapUrl?: string
|
||||||
bootstrapSecret?: string
|
bootstrapSecret?: string
|
||||||
healthUrl?: string
|
|
||||||
apiUrl: string
|
apiUrl: string
|
||||||
apiToken: string
|
apiToken: string
|
||||||
chatId?: string
|
chatId?: string
|
||||||
@@ -188,7 +180,6 @@ const LIGHT: Palette = {
|
|||||||
const COMPOSER_PLACEHOLDER = "Ask nanobot anything"
|
const COMPOSER_PLACEHOLDER = "Ask nanobot anything"
|
||||||
const ACTIVE_COMPOSER_PLACEHOLDER = "Enter send now · Tab send next"
|
const ACTIVE_COMPOSER_PLACEHOLDER = "Enter send now · Tab send next"
|
||||||
const COMPACT_ACTIVE_COMPOSER_PLACEHOLDER = "Enter now · Tab next"
|
const COMPACT_ACTIVE_COMPOSER_PLACEHOLDER = "Enter now · Tab next"
|
||||||
const IMAGE_PLACEHOLDER_STYLE = "image.placeholder"
|
|
||||||
const SHIMMER_PAUSE = 16
|
const SHIMMER_PAUSE = 16
|
||||||
const SHIMMER_BAND = 4
|
const SHIMMER_BAND = 4
|
||||||
const SHIMMER_INTERVAL_MS = 80
|
const SHIMMER_INTERVAL_MS = 80
|
||||||
@@ -264,12 +255,6 @@ function syntaxStyle(palette: Palette): SyntaxStyle {
|
|||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
function composerSyntaxStyle(palette: Palette): SyntaxStyle {
|
|
||||||
return SyntaxStyle.fromStyles({
|
|
||||||
[IMAGE_PLACEHOLDER_STYLE]: { fg: RGBA.fromHex(palette.accent), bold: true },
|
|
||||||
})
|
|
||||||
}
|
|
||||||
|
|
||||||
function transcriptTheme(palette: Palette, backgroundKnown: boolean): TranscriptTheme {
|
function transcriptTheme(palette: Palette, backgroundKnown: boolean): TranscriptTheme {
|
||||||
return {
|
return {
|
||||||
text: palette.text,
|
text: palette.text,
|
||||||
@@ -389,21 +374,6 @@ function formatElapsed(milliseconds: number): string {
|
|||||||
return `${Math.floor(seconds / 60)}m ${String(seconds % 60).padStart(2, "0")}s`
|
return `${Math.floor(seconds / 60)}m ${String(seconds % 60).padStart(2, "0")}s`
|
||||||
}
|
}
|
||||||
|
|
||||||
function connectionStatusText(
|
|
||||||
status: ConnectionStatus,
|
|
||||||
info?: ConnectionStatusInfo,
|
|
||||||
): string {
|
|
||||||
if (["starting", "connecting", "connected"].includes(status)) return "Getting ready…"
|
|
||||||
if (status === "reconnecting") return "Resuming…"
|
|
||||||
if (status === "unavailable") {
|
|
||||||
return info?.health === "degraded"
|
|
||||||
? "Still getting ready…"
|
|
||||||
: "Nanobot is taking longer to respond…"
|
|
||||||
}
|
|
||||||
if (status === "error") return "Nanobot unavailable · restart nanobot"
|
|
||||||
return "Session ended"
|
|
||||||
}
|
|
||||||
|
|
||||||
function singleLine(value: string, limit = 120): string {
|
function singleLine(value: string, limit = 120): string {
|
||||||
return value.replace(/\s+/gu, " ").trim().slice(0, limit)
|
return value.replace(/\s+/gu, " ").trim().slice(0, limit)
|
||||||
}
|
}
|
||||||
@@ -451,7 +421,6 @@ export class NanobotTui {
|
|||||||
private readonly titleText: TextRenderable
|
private readonly titleText: TextRenderable
|
||||||
private readonly composerFrame: BoxRenderable
|
private readonly composerFrame: BoxRenderable
|
||||||
private readonly composer: TextareaRenderable
|
private readonly composer: TextareaRenderable
|
||||||
private composerSyntax: SyntaxStyle
|
|
||||||
private readonly status: TextRenderable
|
private readonly status: TextRenderable
|
||||||
private readonly meta: TextRenderable
|
private readonly meta: TextRenderable
|
||||||
private readonly host: TuiHost
|
private readonly host: TuiHost
|
||||||
@@ -480,8 +449,6 @@ export class NanobotTui {
|
|||||||
private shimmerTimer: ReturnType<typeof setInterval> | null = null
|
private shimmerTimer: ReturnType<typeof setInterval> | null = null
|
||||||
private submitPending = false
|
private submitPending = false
|
||||||
private submitGeneration = 0
|
private submitGeneration = 0
|
||||||
private unsentSubmit = false
|
|
||||||
private connectionMessage = "Getting ready…"
|
|
||||||
private readonly promptHistory: string[] = []
|
private readonly promptHistory: string[] = []
|
||||||
private historyCursor = 0
|
private historyCursor = 0
|
||||||
private historyDraft = ""
|
private historyDraft = ""
|
||||||
@@ -522,14 +489,8 @@ export class NanobotTui {
|
|||||||
private hostWorkspace: string
|
private hostWorkspace: string
|
||||||
private hostBranch: string
|
private hostBranch: string
|
||||||
private readonly apiReauthenticator: ApiReauthenticator | undefined
|
private readonly apiReauthenticator: ApiReauthenticator | undefined
|
||||||
private readonly clipboardImageReader: ClipboardImageReader
|
|
||||||
private apiRefreshPromise: Promise<GatewayApiConnection> | null = null
|
private apiRefreshPromise: Promise<GatewayApiConnection> | null = null
|
||||||
private skillLoadId = 0
|
private skillLoadId = 0
|
||||||
private clipboardImagePending = false
|
|
||||||
private clipboardPasteGeneration = 0
|
|
||||||
private composerValue = ""
|
|
||||||
private composerCursor = 0
|
|
||||||
private reconcilingComposer = false
|
|
||||||
|
|
||||||
private constructor(
|
private constructor(
|
||||||
renderer: CliRenderer,
|
renderer: CliRenderer,
|
||||||
@@ -537,10 +498,8 @@ export class NanobotTui {
|
|||||||
client?: ChatClient,
|
client?: ChatClient,
|
||||||
treeSitterClient = getTreeSitterClient(),
|
treeSitterClient = getTreeSitterClient(),
|
||||||
host: TuiHost = createTuiHost({}),
|
host: TuiHost = createTuiHost({}),
|
||||||
clipboardImageReader: ClipboardImageReader = createClipboardImageReader(),
|
|
||||||
) {
|
) {
|
||||||
this.renderer = renderer
|
this.renderer = renderer
|
||||||
this.clipboardImageReader = clipboardImageReader
|
|
||||||
this.defaultModelName = options.model
|
this.defaultModelName = options.model
|
||||||
this.defaultModelPreset = options.modelPreset
|
this.defaultModelPreset = options.modelPreset
|
||||||
this.modelName = options.model
|
this.modelName = options.model
|
||||||
@@ -554,7 +513,6 @@ export class NanobotTui {
|
|||||||
this.backgroundKnown = options.theme !== "auto" || renderer.themeMode !== null
|
this.backgroundKnown = options.theme !== "auto" || renderer.themeMode !== null
|
||||||
this.activeThemeMode = this.resolveThemeMode(renderer.themeMode)
|
this.activeThemeMode = this.resolveThemeMode(renderer.themeMode)
|
||||||
this.palette = this.activeThemeMode === "light" ? LIGHT : DARK
|
this.palette = this.activeThemeMode === "light" ? LIGHT : DARK
|
||||||
this.composerSyntax = composerSyntaxStyle(this.palette)
|
|
||||||
this.host = host
|
this.host = host
|
||||||
this.transcript = new Transcript(
|
this.transcript = new Transcript(
|
||||||
renderer,
|
renderer,
|
||||||
@@ -598,14 +556,11 @@ export class NanobotTui {
|
|||||||
options.apiUrl,
|
options.apiUrl,
|
||||||
`tui-${process.pid}`,
|
`tui-${process.pid}`,
|
||||||
),
|
),
|
||||||
...(options.healthUrl
|
|
||||||
? { checkHealth: () => fetchGatewayHealth(options.healthUrl || "") }
|
|
||||||
: {}),
|
|
||||||
onConnection: (connection) => this.useGatewayConnection(
|
onConnection: (connection) => this.useGatewayConnection(
|
||||||
connection.apiUrl,
|
connection.apiUrl,
|
||||||
connection.apiToken,
|
connection.apiToken,
|
||||||
),
|
),
|
||||||
targetEndpoint: connectionEndpoint(options.bootstrapUrl),
|
connectionRetryLabel: "Starting local gateway",
|
||||||
reconnectDelayMs: 100,
|
reconnectDelayMs: 100,
|
||||||
startupRetryMaxDelayMs: 250,
|
startupRetryMaxDelayMs: 250,
|
||||||
}
|
}
|
||||||
@@ -616,7 +571,7 @@ export class NanobotTui {
|
|||||||
access_mode: options.access.toLocaleLowerCase().includes("full") ? "full" : "restricted",
|
access_mode: options.access.toLocaleLowerCase().includes("full") ? "full" : "restricted",
|
||||||
},
|
},
|
||||||
onEvent: (event) => this.accept(event),
|
onEvent: (event) => this.accept(event),
|
||||||
onStatus: (status, detail, info) => this.handleStatus(status, detail, info),
|
onStatus: (status, detail) => this.handleStatus(status, detail),
|
||||||
})
|
})
|
||||||
|
|
||||||
// The terminal owns its canvas. Keeping the default-background intent is
|
// The terminal owns its canvas. Keeping the default-background intent is
|
||||||
@@ -751,7 +706,6 @@ export class NanobotTui {
|
|||||||
backgroundColor: composerSurface,
|
backgroundColor: composerSurface,
|
||||||
focusedBackgroundColor: composerSurface,
|
focusedBackgroundColor: composerSurface,
|
||||||
cursorColor: this.palette.accent,
|
cursorColor: this.palette.accent,
|
||||||
syntaxStyle: this.composerSyntax,
|
|
||||||
// A steady line cursor avoids the block-cell trails produced by some
|
// A steady line cursor avoids the block-cell trails produced by some
|
||||||
// terminals when a retained full-screen UI redraws around the composer.
|
// terminals when a retained full-screen UI redraws around the composer.
|
||||||
cursorStyle: { style: "line", blinking: false },
|
cursorStyle: { style: "line", blinking: false },
|
||||||
@@ -765,14 +719,18 @@ export class NanobotTui {
|
|||||||
{ name: "return", action: "submit" },
|
{ name: "return", action: "submit" },
|
||||||
],
|
],
|
||||||
onCursorChange: () => {
|
onCursorChange: () => {
|
||||||
this.keepComposerCursorOutsideImages()
|
|
||||||
if (!this.sessionMenu.visible && !this.branchMenu.visible) this.syncComposerMenus()
|
if (!this.sessionMenu.visible && !this.branchMenu.visible) this.syncComposerMenus()
|
||||||
},
|
},
|
||||||
onContentChange: () => this.handleComposerContentChange(),
|
onContentChange: () => {
|
||||||
onMouseDown: () => queueMicrotask(() => this.keepComposerCursorOutsideImages()),
|
this.draft.prune(this.composer.plainText)
|
||||||
onMouseUp: () => queueMicrotask(() => this.keepComposerCursorOutsideImages()),
|
this.runtimeControls.hide()
|
||||||
onMouseDrag: () => queueMicrotask(() => this.keepComposerCursorOutsideImages()),
|
if (this.contextPanel.visible && this.composer.plainText) this.contextPanel.hide()
|
||||||
onMouseDragEnd: () => queueMicrotask(() => this.keepComposerCursorOutsideImages()),
|
this.syncComposerPlaceholder()
|
||||||
|
if (this.sessionMenu.visible) this.syncSessionMenu()
|
||||||
|
else if (this.branchMenu.visible) this.syncBranchMenu()
|
||||||
|
else this.syncComposerMenus()
|
||||||
|
this.resizeComposer()
|
||||||
|
},
|
||||||
// IMEs may commit their final composed glyph after Enter. Matching the
|
// IMEs may commit their final composed glyph after Enter. Matching the
|
||||||
// OpenCode/OpenTUI integration, defer twice before reading plainText.
|
// OpenCode/OpenTUI integration, defer twice before reading plainText.
|
||||||
onSubmit: () => this.deferSubmit(),
|
onSubmit: () => this.deferSubmit(),
|
||||||
@@ -780,7 +738,7 @@ export class NanobotTui {
|
|||||||
})
|
})
|
||||||
this.status = new TextRenderable(renderer, {
|
this.status = new TextRenderable(renderer, {
|
||||||
id: "nanobot-tui-status",
|
id: "nanobot-tui-status",
|
||||||
content: "Getting ready…",
|
content: "Connecting…",
|
||||||
fg: this.palette.muted,
|
fg: this.palette.muted,
|
||||||
height: 1,
|
height: 1,
|
||||||
width: "auto",
|
width: "auto",
|
||||||
@@ -858,23 +816,15 @@ export class NanobotTui {
|
|||||||
client?: ChatClient,
|
client?: ChatClient,
|
||||||
treeSitterClient?: TreeSitterClient,
|
treeSitterClient?: TreeSitterClient,
|
||||||
host?: TuiHost,
|
host?: TuiHost,
|
||||||
clipboardImageReader?: ClipboardImageReader,
|
|
||||||
): NanobotTui {
|
): NanobotTui {
|
||||||
return new NanobotTui(
|
return new NanobotTui(renderer, options, client, treeSitterClient, host)
|
||||||
renderer,
|
|
||||||
options,
|
|
||||||
client,
|
|
||||||
treeSitterClient,
|
|
||||||
host,
|
|
||||||
clipboardImageReader,
|
|
||||||
)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
async start(): Promise<void> {
|
async start(): Promise<void> {
|
||||||
// Network setup and small menu payloads do not depend on terminal colors.
|
// Network setup and small menu payloads do not depend on terminal colors.
|
||||||
// Start them while OSC theme detection is in flight instead of serializing
|
// Start them while OSC theme detection is in flight instead of serializing
|
||||||
// up to one second of otherwise independent startup work.
|
// up to one second of otherwise independent startup work.
|
||||||
this.host.reportState("unknown", "Getting ready")
|
this.host.reportState("unknown", "Connecting")
|
||||||
this.client.connect()
|
this.client.connect()
|
||||||
void this.loadCommands()
|
void this.loadCommands()
|
||||||
void this.loadMentions()
|
void this.loadMentions()
|
||||||
@@ -910,6 +860,7 @@ export class NanobotTui {
|
|||||||
private submit(): void {
|
private submit(): void {
|
||||||
if (this.quitting || this.composer.isDestroyed) return
|
if (this.quitting || this.composer.isDestroyed) return
|
||||||
const visibleContent = this.composer.plainText.trim()
|
const visibleContent = this.composer.plainText.trim()
|
||||||
|
const content = this.draft.expand(visibleContent).trim()
|
||||||
if (this.sessionLoading) {
|
if (this.sessionLoading) {
|
||||||
this.status.content = "Loading sessions…"
|
this.status.content = "Loading sessions…"
|
||||||
return
|
return
|
||||||
@@ -948,10 +899,6 @@ export class NanobotTui {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
if (["/continue", "/dismiss"].includes(visibleContent.toLowerCase())) {
|
if (["/continue", "/dismiss"].includes(visibleContent.toLowerCase())) {
|
||||||
if (!this.ready) {
|
|
||||||
this.markSubmitUnsent()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
this.clearComposer()
|
this.clearComposer()
|
||||||
this.commandMenu.hide()
|
this.commandMenu.hide()
|
||||||
void this.updateRecovery(visibleContent.toLowerCase() === "/continue" ? "continue" : "dismiss")
|
void this.updateRecovery(visibleContent.toLowerCase() === "/continue" ? "continue" : "dismiss")
|
||||||
@@ -965,10 +912,6 @@ export class NanobotTui {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
const command = this.commandMenu.resolve(visibleContent)
|
const command = this.commandMenu.resolve(visibleContent)
|
||||||
if ((command || visibleContent.startsWith("!")) && this.draft.media(visibleContent).length) {
|
|
||||||
this.status.content = "Images cannot be used with commands · remove the image first"
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if (command?.source === "tui") {
|
if (command?.source === "tui") {
|
||||||
if (command.command.action === "sessions") void this.openSessions()
|
if (command.command.action === "sessions") void this.openSessions()
|
||||||
else if (command.command.action === "context") void this.openContext()
|
else if (command.command.action === "context") void this.openContext()
|
||||||
@@ -989,11 +932,10 @@ export class NanobotTui {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
if (!this.ready) {
|
if (!this.ready) {
|
||||||
this.markSubmitUnsent()
|
this.status.content = "Preparing chat…"
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
const prompt = this.composerPrompt()
|
const prompt = { content, options: mentionOptions(content, this.availableMentions()) }
|
||||||
if (!this.canSendPrompt(prompt)) return
|
|
||||||
if (this.activeTurn) {
|
if (this.activeTurn) {
|
||||||
this.sendPrompt(prompt, true)
|
this.sendPrompt(prompt, true)
|
||||||
return
|
return
|
||||||
@@ -1005,22 +947,16 @@ export class NanobotTui {
|
|||||||
let turnId: string
|
let turnId: string
|
||||||
try {
|
try {
|
||||||
turnId = this.client.send(prompt.content, prompt.options)
|
turnId = this.client.send(prompt.content, prompt.options)
|
||||||
} catch {
|
} catch (error) {
|
||||||
this.markSubmitUnsent(true)
|
this.status.content = error instanceof Error ? error.message : String(error)
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
this.unsentSubmit = false
|
|
||||||
this.clearComposer()
|
this.clearComposer()
|
||||||
this.commandMenu.hide()
|
this.commandMenu.hide()
|
||||||
this.mentionMenu.hide()
|
this.mentionMenu.hide()
|
||||||
this.skillMenu.hide()
|
this.skillMenu.hide()
|
||||||
this.recordPrompt(prompt.content)
|
this.recordPrompt(prompt.content)
|
||||||
this.transcript.user(
|
this.transcript.user(prompt.content, turnId)
|
||||||
prompt.content,
|
|
||||||
turnId,
|
|
||||||
prompt.options.media,
|
|
||||||
prompt.displayContent,
|
|
||||||
)
|
|
||||||
this.hostBlocked = false
|
this.hostBlocked = false
|
||||||
this.setCurrentTask(prompt.content)
|
this.setCurrentTask(prompt.content)
|
||||||
if (steering) {
|
if (steering) {
|
||||||
@@ -1110,13 +1046,12 @@ export class NanobotTui {
|
|||||||
this.reconcileTurnOwnership(event)
|
this.reconcileTurnOwnership(event)
|
||||||
return
|
return
|
||||||
case "user_message": {
|
case "user_message": {
|
||||||
if (this.transcript.user(
|
const attachments = event.media_urls?.map((media) => media.name).filter(Boolean) || []
|
||||||
|
const content = [
|
||||||
event.text,
|
event.text,
|
||||||
event.turn_id,
|
attachments.length ? `Attachments: ${attachments.join(", ")}` : "",
|
||||||
event.media_urls,
|
].filter(Boolean).join("\n")
|
||||||
)) {
|
if (this.transcript.user(content, event.turn_id)) this.recordPrompt(event.text)
|
||||||
this.recordPrompt(event.text)
|
|
||||||
}
|
|
||||||
this.hostBlocked = false
|
this.hostBlocked = false
|
||||||
this.setCurrentTask(event.text)
|
this.setCurrentTask(event.text)
|
||||||
this.reconcileTurnOwnership(event)
|
this.reconcileTurnOwnership(event)
|
||||||
@@ -1457,58 +1392,36 @@ export class NanobotTui {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
private handleStatus(
|
private handleStatus(status: ConnectionStatus, detail?: string): void {
|
||||||
status: ConnectionStatus,
|
|
||||||
_detail?: string,
|
|
||||||
info?: ConnectionStatusInfo,
|
|
||||||
): void {
|
|
||||||
// Invalid frames do not mean the transport is unavailable. Keep the last
|
|
||||||
// accurate user-facing state unless the protocol supplied connection diagnostics.
|
|
||||||
if (status === "error" && !info) return
|
|
||||||
this.connectionMessage = connectionStatusText(status, info)
|
|
||||||
if (status === "connected") {
|
if (status === "connected") {
|
||||||
this.ready = false
|
this.ready = false
|
||||||
this.host.reportState("unknown", "Getting ready")
|
this.host.reportState("unknown", "Connecting")
|
||||||
this.renderConnectionMessage()
|
this.status.content = "Connected · preparing chat…"
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if (["starting", "connecting", "reconnecting", "unavailable"].includes(status)) {
|
if (status === "connecting") {
|
||||||
this.ready = false
|
this.ready = false
|
||||||
this.host.reportState("unknown", this.connectionMessage)
|
const label = detail === "Starting local gateway"
|
||||||
if (status === "reconnecting" || status === "unavailable") this.setActive(false)
|
? detail
|
||||||
this.renderConnectionMessage()
|
: detail ? "Reconnecting" : "Connecting"
|
||||||
|
this.host.reportState("unknown", label)
|
||||||
|
if (detail) this.setActive(false)
|
||||||
|
this.status.content = `${label}…`
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if (status === "error") {
|
if (status === "error") {
|
||||||
if (info) this.ready = false
|
|
||||||
this.setActive(false)
|
this.setActive(false)
|
||||||
this.host.reportState("unknown", this.connectionMessage)
|
this.host.reportState("unknown", detail || "Connection error")
|
||||||
this.renderConnectionMessage()
|
this.status.content = detail || "Connection error"
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if (!this.quitting) {
|
if (!this.quitting) {
|
||||||
this.ready = false
|
|
||||||
this.setActive(false)
|
this.setActive(false)
|
||||||
this.host.reportState("unknown", "Disconnected")
|
this.host.reportState("unknown", "Disconnected")
|
||||||
this.renderConnectionMessage()
|
this.status.content = "Disconnected"
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
private renderConnectionMessage(): void {
|
|
||||||
this.status.content = this.unsentSubmit
|
|
||||||
? `Not sent · press Enter to retry when ready · ${this.connectionMessage}`
|
|
||||||
: this.connectionMessage
|
|
||||||
}
|
|
||||||
|
|
||||||
private markSubmitUnsent(sendFailed = false): void {
|
|
||||||
this.unsentSubmit = true
|
|
||||||
if (sendFailed) {
|
|
||||||
this.status.content = "Not sent · send failed; press Enter to retry when ready"
|
|
||||||
return
|
|
||||||
}
|
|
||||||
this.renderConnectionMessage()
|
|
||||||
}
|
|
||||||
|
|
||||||
private setActive(active: boolean, startedAt?: number): void {
|
private setActive(active: boolean, startedAt?: number): void {
|
||||||
if (this.activeTurn === active) {
|
if (this.activeTurn === active) {
|
||||||
if (active && startedAt !== undefined) this.activeStartedAt = startedAt
|
if (active && startedAt !== undefined) this.activeStartedAt = startedAt
|
||||||
@@ -1547,7 +1460,6 @@ export class NanobotTui {
|
|||||||
}
|
}
|
||||||
|
|
||||||
private readyStatus(detail = this.readyDetail): string {
|
private readyStatus(detail = this.readyDetail): string {
|
||||||
if (this.unsentSubmit) return "Not sent · press Enter to retry"
|
|
||||||
if (this.transcriptNavigation.awayFromBottom) {
|
if (this.transcriptNavigation.awayFromBottom) {
|
||||||
return this.transcriptNavigation.unseenOutput
|
return this.transcriptNavigation.unseenOutput
|
||||||
? "New output · Ctrl+End latest"
|
? "New output · Ctrl+End latest"
|
||||||
@@ -1575,36 +1487,6 @@ export class NanobotTui {
|
|||||||
return queue
|
return queue
|
||||||
}
|
}
|
||||||
|
|
||||||
private composerPrompt(): QueuedPrompt {
|
|
||||||
const visible = this.composer.plainText.trim()
|
|
||||||
const content = this.draft.expand(visible).trim()
|
|
||||||
const media = this.draft.media(visible)
|
|
||||||
const displayContent = this.draft.display(visible).trim()
|
|
||||||
return {
|
|
||||||
content,
|
|
||||||
...(media.length ? { displayContent } : {}),
|
|
||||||
options: {
|
|
||||||
...mentionOptions(content, this.availableMentions()),
|
|
||||||
...(media.length ? { media } : {}),
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
private hasPrompt(prompt: QueuedPrompt): boolean {
|
|
||||||
return Boolean(prompt.content || prompt.options.media?.length)
|
|
||||||
}
|
|
||||||
|
|
||||||
private canSendPrompt(prompt: QueuedPrompt): boolean {
|
|
||||||
if (this.draft.hasImageLabelConflict(this.composer.plainText)) {
|
|
||||||
this.status.content = "Duplicate image placeholder text · rename or remove it before sending"
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
if (!this.hasPrompt(prompt)) return false
|
|
||||||
if ((prompt.options.media?.length || 0) <= MAX_DRAFT_IMAGES) return true
|
|
||||||
this.status.content = `Remove images until ${MAX_DRAFT_IMAGES} or fewer remain`
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
private restoreQueuedPrompts(): void {
|
private restoreQueuedPrompts(): void {
|
||||||
const queued = this.promptQueue.restore()
|
const queued = this.promptQueue.restore()
|
||||||
if (!queued.length) return
|
if (!queued.length) return
|
||||||
@@ -1616,10 +1498,6 @@ export class NanobotTui {
|
|||||||
private queueFollowUp(): void {
|
private queueFollowUp(): void {
|
||||||
if (!this.activeTurn || !this.ready) return
|
if (!this.activeTurn || !this.ready) return
|
||||||
const visibleContent = this.composer.plainText.trim()
|
const visibleContent = this.composer.plainText.trim()
|
||||||
if (this.draft.media(visibleContent).length) {
|
|
||||||
this.status.content = "Images cannot be queued · press Enter to send now"
|
|
||||||
return
|
|
||||||
}
|
|
||||||
const content = this.draft.expand(visibleContent).trim()
|
const content = this.draft.expand(visibleContent).trim()
|
||||||
if (!content) return
|
if (!content) return
|
||||||
this.promptQueue.enqueue({
|
this.promptQueue.enqueue({
|
||||||
@@ -1784,18 +1662,6 @@ export class NanobotTui {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if (
|
|
||||||
(key.ctrl || key.meta)
|
|
||||||
&& key.name.toLocaleLowerCase() === "v"
|
|
||||||
&& !this.sessionLoading
|
|
||||||
&& !this.sessionMenu.visible
|
|
||||||
&& !this.branchMenu.visible
|
|
||||||
&& !this.contextPanel.visible
|
|
||||||
) {
|
|
||||||
key.preventDefault()
|
|
||||||
void this.pasteClipboardImage()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if (this.activeTurn && !key.ctrl && !key.meta && key.name === "tab") {
|
if (this.activeTurn && !key.ctrl && !key.meta && key.name === "tab") {
|
||||||
this.queueFollowUp()
|
this.queueFollowUp()
|
||||||
key.preventDefault()
|
key.preventDefault()
|
||||||
@@ -1812,20 +1678,6 @@ export class NanobotTui {
|
|||||||
key.preventDefault()
|
key.preventDefault()
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if (!key.ctrl && !key.meta && !key.shift && (key.name === "left" || key.name === "right")) {
|
|
||||||
const direction = key.name === "left" ? -1 : 1
|
|
||||||
const target = this.draft.moveImageCursor(
|
|
||||||
this.composer.plainText,
|
|
||||||
this.composerStringCursor(),
|
|
||||||
direction,
|
|
||||||
)
|
|
||||||
if (target !== null) {
|
|
||||||
this.composerCursor = target
|
|
||||||
this.setComposerStringCursor(this.composer.plainText, target)
|
|
||||||
key.preventDefault()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if (!key.ctrl && !key.meta && (key.name === "up" || key.name === "down")) {
|
if (!key.ctrl && !key.meta && (key.name === "up" || key.name === "down")) {
|
||||||
const direction = key.name === "up" ? -1 : 1
|
const direction = key.name === "up" ? -1 : 1
|
||||||
const boundary = direction < 0 ? 0 : this.composer.plainText.length
|
const boundary = direction < 0 ? 0 : this.composer.plainText.length
|
||||||
@@ -1886,7 +1738,7 @@ export class NanobotTui {
|
|||||||
}
|
}
|
||||||
|
|
||||||
private navigateHistory(direction: -1 | 1): boolean {
|
private navigateHistory(direction: -1 | 1): boolean {
|
||||||
if (this.promptHistory.length === 0 || this.draft.imageCount) return false
|
if (this.promptHistory.length === 0) return false
|
||||||
if (direction < 0) {
|
if (direction < 0) {
|
||||||
if (this.historyCursor === this.promptHistory.length) this.historyDraft = this.composer.plainText
|
if (this.historyCursor === this.promptHistory.length) this.historyDraft = this.composer.plainText
|
||||||
if (this.historyCursor === 0) return false
|
if (this.historyCursor === 0) return false
|
||||||
@@ -1937,11 +1789,6 @@ export class NanobotTui {
|
|||||||
this.composer.textColor = this.palette.text
|
this.composer.textColor = this.palette.text
|
||||||
this.composer.focusedTextColor = this.palette.text
|
this.composer.focusedTextColor = this.palette.text
|
||||||
this.composer.cursorColor = this.palette.accent
|
this.composer.cursorColor = this.palette.accent
|
||||||
const previousComposerSyntax = this.composerSyntax
|
|
||||||
this.composerSyntax = composerSyntaxStyle(this.palette)
|
|
||||||
this.composer.syntaxStyle = this.composerSyntax
|
|
||||||
this.syncComposerImageHighlights(this.composer.plainText)
|
|
||||||
void this.renderer.idle().catch(() => {}).finally(() => previousComposerSyntax.destroy())
|
|
||||||
this.renderTitleColor()
|
this.renderTitleColor()
|
||||||
this.status.fg = this.palette.muted
|
this.status.fg = this.palette.muted
|
||||||
this.meta.fg = this.palette.faint
|
this.meta.fg = this.palette.faint
|
||||||
@@ -2212,58 +2059,7 @@ export class NanobotTui {
|
|||||||
return this.composer.editBuffer.getTextRange(0, this.composer.cursorOffset).length
|
return this.composer.editBuffer.getTextRange(0, this.composer.cursorOffset).length
|
||||||
}
|
}
|
||||||
|
|
||||||
private keepComposerCursorOutsideImages(): void {
|
|
||||||
if (this.reconcilingComposer) return
|
|
||||||
const value = this.composer.plainText
|
|
||||||
const cursor = this.composerStringCursor()
|
|
||||||
const target = this.draft.snapImageCursor(value, cursor, this.composerCursor)
|
|
||||||
this.composerCursor = target
|
|
||||||
if (target !== cursor) this.setComposerStringCursor(value, target)
|
|
||||||
}
|
|
||||||
|
|
||||||
private handleComposerContentChange(): void {
|
|
||||||
if (this.reconcilingComposer) return
|
|
||||||
let value = this.composer.plainText
|
|
||||||
let cursor = this.composerStringCursor()
|
|
||||||
const edit = this.draft.reconcileImageEdit(this.composerValue, value, cursor)
|
|
||||||
if (edit.value !== value) {
|
|
||||||
this.reconcilingComposer = true
|
|
||||||
try {
|
|
||||||
this.composer.replaceText(edit.value)
|
|
||||||
this.composer.clearSelection()
|
|
||||||
this.setComposerStringCursor(edit.value, edit.cursor)
|
|
||||||
} finally {
|
|
||||||
this.reconcilingComposer = false
|
|
||||||
}
|
|
||||||
value = edit.value
|
|
||||||
cursor = edit.cursor
|
|
||||||
}
|
|
||||||
this.composerValue = value
|
|
||||||
this.composerCursor = cursor
|
|
||||||
this.draft.prune(value)
|
|
||||||
this.syncComposerImageHighlights(value)
|
|
||||||
const clearedUnsent = this.unsentSubmit && !value.trim()
|
|
||||||
if (clearedUnsent) this.unsentSubmit = false
|
|
||||||
this.runtimeControls.hide()
|
|
||||||
if (this.contextPanel.visible && value) this.contextPanel.hide()
|
|
||||||
this.syncComposerPlaceholder()
|
|
||||||
if (this.sessionMenu.visible) this.syncSessionMenu()
|
|
||||||
else if (this.branchMenu.visible) this.syncBranchMenu()
|
|
||||||
else this.syncComposerMenus()
|
|
||||||
this.resizeComposer()
|
|
||||||
if (clearedUnsent && !this.activeTurn) {
|
|
||||||
this.status.content = this.ready ? this.readyStatus() : this.connectionMessage
|
|
||||||
}
|
|
||||||
if (edit.removedImages.length) {
|
|
||||||
this.status.content = `Removed ${edit.removedImages.join(", ")}`
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
private setComposerStringCursor(value: string, cursor: number): void {
|
private setComposerStringCursor(value: string, cursor: number): void {
|
||||||
this.composer.cursorOffset = this.composerOffsetForStringIndex(value, cursor)
|
|
||||||
}
|
|
||||||
|
|
||||||
private composerOffsetForStringIndex(value: string, cursor: number): number {
|
|
||||||
const target = Math.min(Math.max(cursor, 0), value.length)
|
const target = Math.min(Math.max(cursor, 0), value.length)
|
||||||
const before = value.slice(0, target)
|
const before = value.slice(0, target)
|
||||||
const row = before.split("\n").length - 1
|
const row = before.split("\n").length - 1
|
||||||
@@ -2277,68 +2073,18 @@ export class NanobotTui {
|
|||||||
if (candidateLength > target) break
|
if (candidateLength > target) break
|
||||||
if (candidateLength === target) offset = candidate
|
if (candidateLength === target) offset = candidate
|
||||||
}
|
}
|
||||||
return offset
|
this.composer.cursorOffset = offset
|
||||||
}
|
|
||||||
|
|
||||||
private syncComposerImageHighlights(value: string): void {
|
|
||||||
this.composer.clearAllHighlights()
|
|
||||||
const styleId = this.composerSyntax.getStyleId(IMAGE_PLACEHOLDER_STYLE)
|
|
||||||
if (styleId === null) return
|
|
||||||
for (const range of this.draft.imagePlaceholderRanges(value)) {
|
|
||||||
this.composer.addHighlightByCharRange({
|
|
||||||
start: this.composerOffsetForStringIndex(value, range.start),
|
|
||||||
end: this.composerOffsetForStringIndex(value, range.end),
|
|
||||||
styleId,
|
|
||||||
priority: 100,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
private setComposer(content: string): void {
|
private setComposer(content: string): void {
|
||||||
this.clipboardPasteGeneration += 1
|
|
||||||
this.draft.clear()
|
this.draft.clear()
|
||||||
this.composer.setText(content)
|
this.composer.setText(content)
|
||||||
this.composer.cursorOffset = content.length
|
this.composer.cursorOffset = content.length
|
||||||
}
|
}
|
||||||
|
|
||||||
private clearComposer(): void {
|
private clearComposer(): void {
|
||||||
this.unsentSubmit = false
|
this.draft.clear()
|
||||||
this.setComposer("")
|
this.composer.setText("")
|
||||||
}
|
|
||||||
|
|
||||||
private async pasteClipboardImage(): Promise<void> {
|
|
||||||
if (this.clipboardImagePending) return
|
|
||||||
if (this.draft.imageCount >= MAX_DRAFT_IMAGES) {
|
|
||||||
this.status.content = `A message can include up to ${MAX_DRAFT_IMAGES} images`
|
|
||||||
return
|
|
||||||
}
|
|
||||||
const generation = this.clipboardPasteGeneration
|
|
||||||
this.clipboardImagePending = true
|
|
||||||
this.status.content = "Reading clipboard image…"
|
|
||||||
try {
|
|
||||||
const image = await this.clipboardImageReader.read()
|
|
||||||
if (this.quitting || generation !== this.clipboardPasteGeneration) return
|
|
||||||
const insertion = this.draft.image(image, this.composer.plainText)
|
|
||||||
if (!insertion) {
|
|
||||||
this.status.content = `A message can include up to ${MAX_DRAFT_IMAGES} images`
|
|
||||||
return
|
|
||||||
}
|
|
||||||
this.composer.insertText(insertion.text)
|
|
||||||
this.status.content = `Pasted ${insertion.description} · review before sending`
|
|
||||||
} catch (error) {
|
|
||||||
if (
|
|
||||||
this.quitting
|
|
||||||
|| this.composer.isDestroyed
|
|
||||||
|| generation !== this.clipboardPasteGeneration
|
|
||||||
) return
|
|
||||||
const message = error instanceof Error
|
|
||||||
? error.message
|
|
||||||
: "Clipboard image paste is unavailable"
|
|
||||||
this.status.content = message
|
|
||||||
this.transcript.notice(message, true)
|
|
||||||
} finally {
|
|
||||||
this.clipboardImagePending = false
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
private handlePaste(event: PasteEvent): void {
|
private handlePaste(event: PasteEvent): void {
|
||||||
@@ -2637,7 +2383,7 @@ export class NanobotTui {
|
|||||||
options: MessageOptions = {},
|
options: MessageOptions = {},
|
||||||
): void {
|
): void {
|
||||||
if (!this.ready) {
|
if (!this.ready) {
|
||||||
this.markSubmitUnsent()
|
this.status.content = "Preparing chat…"
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if (this.activeTurn && lifecycle === "agent_turn") {
|
if (this.activeTurn && lifecycle === "agent_turn") {
|
||||||
@@ -2647,11 +2393,10 @@ export class NanobotTui {
|
|||||||
let turnId: string
|
let turnId: string
|
||||||
try {
|
try {
|
||||||
turnId = this.client.send(content, options)
|
turnId = this.client.send(content, options)
|
||||||
} catch {
|
} catch (error) {
|
||||||
this.markSubmitUnsent(true)
|
this.status.content = error instanceof Error ? error.message : String(error)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
this.unsentSubmit = false
|
|
||||||
this.commandTurns.set(turnId, lifecycle)
|
this.commandTurns.set(turnId, lifecycle)
|
||||||
if (silent) this.silentCommandTurns.add(turnId)
|
if (silent) this.silentCommandTurns.add(turnId)
|
||||||
if (/^\/model(?:\s|$)/iu.test(content)) this.modelCommandTurns.add(turnId)
|
if (/^\/model(?:\s|$)/iu.test(content)) this.modelCommandTurns.add(turnId)
|
||||||
@@ -2692,7 +2437,6 @@ export class NanobotTui {
|
|||||||
}
|
}
|
||||||
|
|
||||||
private recordPrompt(content: string): void {
|
private recordPrompt(content: string): void {
|
||||||
if (!content) return
|
|
||||||
if (this.promptHistory.at(-1) !== content) this.promptHistory.push(content)
|
if (this.promptHistory.at(-1) !== content) this.promptHistory.push(content)
|
||||||
if (this.promptHistory.length > 50) this.promptHistory.shift()
|
if (this.promptHistory.length > 50) this.promptHistory.shift()
|
||||||
this.historyCursor = this.promptHistory.length
|
this.historyCursor = this.promptHistory.length
|
||||||
@@ -2920,14 +2664,10 @@ export class NanobotTui {
|
|||||||
}
|
}
|
||||||
|
|
||||||
private handleDestroy = (): void => {
|
private handleDestroy = (): void => {
|
||||||
this.quitting = true
|
|
||||||
this.clipboardPasteGeneration += 1
|
|
||||||
if (this.shimmerTimer) clearInterval(this.shimmerTimer)
|
if (this.shimmerTimer) clearInterval(this.shimmerTimer)
|
||||||
this.stopSessionRefresh()
|
this.stopSessionRefresh()
|
||||||
this.composerSyntax.destroy()
|
|
||||||
this.transcript.destroy()
|
this.transcript.destroy()
|
||||||
this.diffViewer.destroy()
|
this.diffViewer.destroy()
|
||||||
void this.clipboardImageReader.dispose().catch(() => {})
|
|
||||||
this.host.release()
|
this.host.release()
|
||||||
this.client.close()
|
this.client.close()
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,46 +0,0 @@
|
|||||||
import { describe, expect, test } from "bun:test"
|
|
||||||
import type { ClipboardReadResult, HostClipboardService } from "@opentui/core"
|
|
||||||
|
|
||||||
import { createClipboardImageReader } from "./clipboard-image"
|
|
||||||
|
|
||||||
function clipboard(result: ClipboardReadResult) {
|
|
||||||
let disposed = false
|
|
||||||
const service = {
|
|
||||||
maxWriteBytes: 1,
|
|
||||||
read: async () => result,
|
|
||||||
writeText: async () => ({ status: "unsupported" as const }),
|
|
||||||
clear: async () => ({ status: "unsupported" as const }),
|
|
||||||
dispose: async () => { disposed = true },
|
|
||||||
} satisfies HostClipboardService
|
|
||||||
return { service, disposed: () => disposed }
|
|
||||||
}
|
|
||||||
|
|
||||||
describe("clipboard image reader", () => {
|
|
||||||
test("encodes supported native clipboard bytes as a data URL", async () => {
|
|
||||||
const fake = clipboard({
|
|
||||||
status: "read",
|
|
||||||
representation: { mimeType: "image/png", bytes: Uint8Array.from([0, 1, 2, 255]) },
|
|
||||||
})
|
|
||||||
const reader = createClipboardImageReader(() => fake.service)
|
|
||||||
|
|
||||||
expect(await reader.read()).toEqual({
|
|
||||||
mimeType: "image/png",
|
|
||||||
dataUrl: "data:image/png;base64,AAEC/w==",
|
|
||||||
})
|
|
||||||
await reader.dispose()
|
|
||||||
expect(fake.disposed()).toBeTrue()
|
|
||||||
})
|
|
||||||
|
|
||||||
test.each([
|
|
||||||
["empty", "No image in clipboard"],
|
|
||||||
["limit-exceeded", "Clipboard image is larger than 6 MB"],
|
|
||||||
["timed-out", "Clipboard image read timed out"],
|
|
||||||
["unsupported", "Clipboard image paste is unavailable"],
|
|
||||||
] as const)("reports %s without exposing native details", async (status, message) => {
|
|
||||||
const fake = clipboard({ status })
|
|
||||||
const reader = createClipboardImageReader(() => fake.service)
|
|
||||||
|
|
||||||
await expect(reader.read()).rejects.toThrow(message)
|
|
||||||
await reader.dispose()
|
|
||||||
})
|
|
||||||
})
|
|
||||||
@@ -1,64 +0,0 @@
|
|||||||
import {
|
|
||||||
createHostClipboard,
|
|
||||||
type HostClipboardService,
|
|
||||||
} from "@opentui/core"
|
|
||||||
|
|
||||||
const IMAGE_MIME_TYPES = [
|
|
||||||
"image/png",
|
|
||||||
"image/jpeg",
|
|
||||||
"image/webp",
|
|
||||||
"image/gif",
|
|
||||||
] as const
|
|
||||||
const MAX_IMAGE_BYTES = 6 * 1024 * 1024
|
|
||||||
|
|
||||||
export interface ClipboardImage {
|
|
||||||
dataUrl: string
|
|
||||||
mimeType: typeof IMAGE_MIME_TYPES[number]
|
|
||||||
}
|
|
||||||
|
|
||||||
export interface ClipboardImageReader {
|
|
||||||
read(): Promise<ClipboardImage>
|
|
||||||
dispose(): Promise<void>
|
|
||||||
}
|
|
||||||
|
|
||||||
type ClipboardFactory = () => HostClipboardService
|
|
||||||
|
|
||||||
function readFailure(status: string): Error {
|
|
||||||
if (status === "empty") return new Error("No image in clipboard")
|
|
||||||
if (status === "limit-exceeded") return new Error("Clipboard image is larger than 6 MB")
|
|
||||||
if (status === "timed-out") return new Error("Clipboard image read timed out")
|
|
||||||
return new Error("Clipboard image paste is unavailable")
|
|
||||||
}
|
|
||||||
|
|
||||||
/** Lazily owns OpenTUI's native host clipboard so ordinary TUI startup does no clipboard work. */
|
|
||||||
export function createClipboardImageReader(
|
|
||||||
createClipboard: ClipboardFactory = () => createHostClipboard({ maxReadBytes: MAX_IMAGE_BYTES }),
|
|
||||||
): ClipboardImageReader {
|
|
||||||
let clipboard: HostClipboardService | null = null
|
|
||||||
let disposed = false
|
|
||||||
|
|
||||||
return {
|
|
||||||
async read(): Promise<ClipboardImage> {
|
|
||||||
if (disposed) throw new Error("Clipboard image paste is unavailable")
|
|
||||||
clipboard ||= createClipboard()
|
|
||||||
const result = await clipboard.read({ preferredTypes: IMAGE_MIME_TYPES })
|
|
||||||
if (result.status !== "read") throw readFailure(result.status)
|
|
||||||
const normalizedMime = result.representation.mimeType.toLowerCase()
|
|
||||||
const mimeType = IMAGE_MIME_TYPES.find((candidate) => candidate === normalizedMime)
|
|
||||||
if (!mimeType) throw new Error("Clipboard does not contain a supported image")
|
|
||||||
const bytes = result.representation.bytes
|
|
||||||
if (!bytes.length) throw new Error("Clipboard image is empty")
|
|
||||||
if (bytes.length > MAX_IMAGE_BYTES) throw new Error("Clipboard image is larger than 6 MB")
|
|
||||||
return {
|
|
||||||
mimeType,
|
|
||||||
dataUrl: `data:${mimeType};base64,${Buffer.from(bytes).toString("base64")}`,
|
|
||||||
}
|
|
||||||
},
|
|
||||||
async dispose(): Promise<void> {
|
|
||||||
if (disposed) return
|
|
||||||
disposed = true
|
|
||||||
await clipboard?.dispose()
|
|
||||||
clipboard = null
|
|
||||||
},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -1,6 +1,6 @@
|
|||||||
import { describe, expect, test } from "bun:test"
|
import { describe, expect, test } from "bun:test"
|
||||||
|
|
||||||
import { ComposerDraft, MAX_DRAFT_IMAGES } from "./composer-draft"
|
import { ComposerDraft } from "./composer-draft"
|
||||||
|
|
||||||
describe("ComposerDraft", () => {
|
describe("ComposerDraft", () => {
|
||||||
test("keeps ordinary pastes editable as ordinary text", () => {
|
test("keeps ordinary pastes editable as ordinary text", () => {
|
||||||
@@ -25,127 +25,4 @@ describe("ComposerDraft", () => {
|
|||||||
expect(draft.expand(first.text.trim())).toBe(first.text.trim())
|
expect(draft.expand(first.text.trim())).toBe(first.text.trim())
|
||||||
expect(draft.expand(second.text.trim())).toBe(content)
|
expect(draft.expand(second.text.trim())).toBe(content)
|
||||||
})
|
})
|
||||||
|
|
||||||
test("keeps image bytes outside the editor and drops attachments with deleted placeholders", () => {
|
|
||||||
const draft = new ComposerDraft()
|
|
||||||
const first = draft.image({ mimeType: "image/png", dataUrl: "data:image/png;base64,AAAA" })
|
|
||||||
const second = draft.image({ mimeType: "image/jpeg", dataUrl: "data:image/jpeg;base64,BBBB" })
|
|
||||||
|
|
||||||
expect(first?.text).toBe("[Image #1] ")
|
|
||||||
expect(second?.text).toBe("[Image #2] ")
|
|
||||||
const visible = `compare ${second?.text}${first?.text}`
|
|
||||||
expect(draft.expand(visible)).toBe("compare ")
|
|
||||||
expect(draft.display(visible)).toBe(visible)
|
|
||||||
expect(draft.media(visible)).toEqual([
|
|
||||||
{ data_url: "data:image/jpeg;base64,BBBB", name: "clipboard-image-2.jpg" },
|
|
||||||
{ data_url: "data:image/png;base64,AAAA", name: "clipboard-image-1.png" },
|
|
||||||
])
|
|
||||||
|
|
||||||
draft.prune(first?.text || "")
|
|
||||||
expect(draft.imageCount).toBe(1)
|
|
||||||
expect(draft.media(second?.text || "")).toEqual([])
|
|
||||||
})
|
|
||||||
|
|
||||||
test("removes a partially edited image placeholder as one atomic unit", () => {
|
|
||||||
const draft = new ComposerDraft()
|
|
||||||
const image = draft.image({ mimeType: "image/png", dataUrl: "data:image/png;base64,AAAA" })
|
|
||||||
const previous = `before ${image?.text}after`
|
|
||||||
const value = previous.replace("[Image #1]", "Image #1]")
|
|
||||||
|
|
||||||
expect(draft.reconcileImageEdit(previous, value, 7)).toEqual({
|
|
||||||
value: "before after",
|
|
||||||
cursor: 7,
|
|
||||||
removedImages: ["Image #1"],
|
|
||||||
})
|
|
||||||
expect(draft.imageCount).toBe(0)
|
|
||||||
expect(draft.media(value)).toEqual([])
|
|
||||||
})
|
|
||||||
|
|
||||||
test("removes an edited duplicate occurrence without leaving a placeholder fragment", () => {
|
|
||||||
const draft = new ComposerDraft()
|
|
||||||
const image = draft.image({ mimeType: "image/png", dataUrl: "data:image/png;base64,AAAA" })
|
|
||||||
const label = image?.text.trim() || ""
|
|
||||||
const previous = `${label} ${label}`
|
|
||||||
|
|
||||||
expect(draft.reconcileImageEdit(previous, previous.slice(1), 0)).toEqual({
|
|
||||||
value: ` ${label}`,
|
|
||||||
cursor: 0,
|
|
||||||
removedImages: [],
|
|
||||||
})
|
|
||||||
expect(draft.imageCount).toBe(1)
|
|
||||||
})
|
|
||||||
|
|
||||||
test("snaps cursor movement across complete image placeholders", () => {
|
|
||||||
const draft = new ComposerDraft()
|
|
||||||
const image = draft.image({ mimeType: "image/png", dataUrl: "data:image/png;base64,AAAA" })
|
|
||||||
const visible = `a ${image?.text}b`
|
|
||||||
|
|
||||||
expect(draft.snapImageCursor(visible, 3, 2)).toBe(12)
|
|
||||||
expect(draft.snapImageCursor(visible, 11, 12)).toBe(2)
|
|
||||||
expect(draft.snapImageCursor(visible, 2, 0)).toBe(2)
|
|
||||||
expect(draft.snapImageCursor(visible, 12, 13)).toBe(12)
|
|
||||||
expect(draft.moveImageCursor(visible, 2, 1)).toBe(12)
|
|
||||||
expect(draft.moveImageCursor(visible, 12, -1)).toBe(2)
|
|
||||||
})
|
|
||||||
|
|
||||||
test("allocates image labels around literal composer text", () => {
|
|
||||||
const draft = new ComposerDraft()
|
|
||||||
const content = "Explain [Image #1]"
|
|
||||||
const insertion = draft.image(
|
|
||||||
{ mimeType: "image/png", dataUrl: "data:image/png;base64,AAAA" },
|
|
||||||
content,
|
|
||||||
)
|
|
||||||
|
|
||||||
expect(insertion?.text).toBe("[Image #2] ")
|
|
||||||
expect(draft.expand(`${content} ${insertion?.text}`.trim())).toBe(`${content} `)
|
|
||||||
})
|
|
||||||
|
|
||||||
test("detects image labels duplicated after insertion without deleting text", () => {
|
|
||||||
const draft = new ComposerDraft()
|
|
||||||
const image = draft.image({ mimeType: "image/png", dataUrl: "data:image/png;base64,AAAA" })
|
|
||||||
const visible = `${image?.text}Explain [Image #1]`
|
|
||||||
|
|
||||||
expect(draft.hasImageLabelConflict(visible)).toBeTrue()
|
|
||||||
expect(draft.expand(visible)).toBe(visible)
|
|
||||||
})
|
|
||||||
|
|
||||||
test("detects image labels inside compacted paste text added afterward", () => {
|
|
||||||
const draft = new ComposerDraft()
|
|
||||||
const image = draft.image({ mimeType: "image/png", dataUrl: "data:image/png;base64,AAAA" })
|
|
||||||
const content = ["Explain [Image #1]", ...Array.from({ length: 11 }, () => "detail")].join("\n")
|
|
||||||
const paste = draft.paste(content)
|
|
||||||
const visible = `${image?.text}${paste.text}`
|
|
||||||
|
|
||||||
expect(draft.hasImageLabelConflict(visible)).toBeTrue()
|
|
||||||
expect(draft.expand(visible)).toContain("Explain [Image #1]")
|
|
||||||
})
|
|
||||||
|
|
||||||
test("allocates image labels around hidden compacted paste text", () => {
|
|
||||||
const draft = new ComposerDraft()
|
|
||||||
const content = ["Explain [Image #1]", ...Array.from({ length: 11 }, () => "detail")].join("\n")
|
|
||||||
const paste = draft.paste(content)
|
|
||||||
const image = draft.image(
|
|
||||||
{ mimeType: "image/png", dataUrl: "data:image/png;base64,AAAA" },
|
|
||||||
paste.text,
|
|
||||||
)
|
|
||||||
|
|
||||||
expect(image?.text).toBe("[Image #2] ")
|
|
||||||
expect(draft.expand(`${paste.text}${image?.text}`)).toContain("Explain [Image #1]")
|
|
||||||
})
|
|
||||||
|
|
||||||
test("matches the gateway image count before accepting another placeholder", () => {
|
|
||||||
const draft = new ComposerDraft()
|
|
||||||
for (let index = 0; index < MAX_DRAFT_IMAGES; index += 1) {
|
|
||||||
expect(draft.image({
|
|
||||||
mimeType: "image/png",
|
|
||||||
dataUrl: `data:image/png;base64,${index}`,
|
|
||||||
})).not.toBeNull()
|
|
||||||
}
|
|
||||||
|
|
||||||
expect(draft.image({
|
|
||||||
mimeType: "image/png",
|
|
||||||
dataUrl: "data:image/png;base64,overflow",
|
|
||||||
})).toBeNull()
|
|
||||||
expect(draft.imageCount).toBe(MAX_DRAFT_IMAGES)
|
|
||||||
})
|
|
||||||
})
|
})
|
||||||
|
|||||||
+2
-157
@@ -1,8 +1,5 @@
|
|||||||
import type { OutboundMedia } from "./protocol"
|
|
||||||
|
|
||||||
const LARGE_PASTE_CHARS = 1_000
|
const LARGE_PASTE_CHARS = 1_000
|
||||||
const LARGE_PASTE_LINES = 10
|
const LARGE_PASTE_LINES = 10
|
||||||
export const MAX_DRAFT_IMAGES = 4
|
|
||||||
|
|
||||||
export interface PasteInsertion {
|
export interface PasteInsertion {
|
||||||
text: string
|
text: string
|
||||||
@@ -10,32 +7,9 @@ export interface PasteInsertion {
|
|||||||
description: string
|
description: string
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface DraftEditReconciliation {
|
/** Keeps large pasted text out of the editor without changing what is sent. */
|
||||||
value: string
|
|
||||||
cursor: number
|
|
||||||
removedImages: string[]
|
|
||||||
}
|
|
||||||
|
|
||||||
const IMAGE_EXTENSIONS = {
|
|
||||||
"image/png": "png",
|
|
||||||
"image/jpeg": "jpg",
|
|
||||||
"image/webp": "webp",
|
|
||||||
"image/gif": "gif",
|
|
||||||
} as const
|
|
||||||
|
|
||||||
interface DraftImage {
|
|
||||||
dataUrl: string
|
|
||||||
mimeType: keyof typeof IMAGE_EXTENSIONS
|
|
||||||
}
|
|
||||||
|
|
||||||
/** Keeps large pasted text and image payloads out of the editable composer surface. */
|
|
||||||
export class ComposerDraft {
|
export class ComposerDraft {
|
||||||
private readonly pastes = new Map<string, string>()
|
private readonly pastes = new Map<string, string>()
|
||||||
private readonly images = new Map<string, OutboundMedia>()
|
|
||||||
|
|
||||||
get imageCount(): number {
|
|
||||||
return this.images.size
|
|
||||||
}
|
|
||||||
|
|
||||||
paste(value: string): PasteInsertion {
|
paste(value: string): PasteInsertion {
|
||||||
const text = value.replace(/\r\n/gu, "\n").replace(/\r/gu, "\n")
|
const text = value.replace(/\r\n/gu, "\n").replace(/\r/gu, "\n")
|
||||||
@@ -52,148 +26,19 @@ export class ComposerDraft {
|
|||||||
return { text: `${label} `, compacted: true, description }
|
return { text: `${label} `, compacted: true, description }
|
||||||
}
|
}
|
||||||
|
|
||||||
private imageLabelInUse(label: string, visible: string): boolean {
|
expand(visible: string): string {
|
||||||
if (this.images.has(label) || visible.includes(label)) return true
|
|
||||||
for (const content of this.pastes.values()) {
|
|
||||||
if (content.includes(label)) return true
|
|
||||||
}
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
|
|
||||||
private nextImageIndex(visible: string): number {
|
|
||||||
let index = 1
|
|
||||||
while (this.imageLabelInUse(`[Image #${index}]`, visible)) index += 1
|
|
||||||
return index
|
|
||||||
}
|
|
||||||
|
|
||||||
image(image: DraftImage, visible = ""): PasteInsertion | null {
|
|
||||||
if (this.images.size >= MAX_DRAFT_IMAGES) return null
|
|
||||||
const index = this.nextImageIndex(visible)
|
|
||||||
const label = `[Image #${index}]`
|
|
||||||
this.images.set(label, {
|
|
||||||
data_url: image.dataUrl,
|
|
||||||
name: `clipboard-image-${index}.${IMAGE_EXTENSIONS[image.mimeType]}`,
|
|
||||||
})
|
|
||||||
return { text: `${label} `, compacted: true, description: label.slice(1, -1) }
|
|
||||||
}
|
|
||||||
|
|
||||||
private expandPastes(visible: string): string {
|
|
||||||
let expanded = visible
|
let expanded = visible
|
||||||
for (const [label, content] of this.pastes) expanded = expanded.split(label).join(content)
|
for (const [label, content] of this.pastes) expanded = expanded.split(label).join(content)
|
||||||
return expanded
|
return expanded
|
||||||
}
|
}
|
||||||
|
|
||||||
private labelOccurrences(content: string, label: string): number {
|
|
||||||
return content.split(label).length - 1
|
|
||||||
}
|
|
||||||
|
|
||||||
imagePlaceholderRanges(visible: string): Array<{ start: number; end: number }> {
|
|
||||||
const ranges: Array<{ start: number; end: number }> = []
|
|
||||||
for (const label of this.images.keys()) {
|
|
||||||
let start = visible.indexOf(label)
|
|
||||||
while (start >= 0) {
|
|
||||||
ranges.push({ start, end: start + label.length })
|
|
||||||
start = visible.indexOf(label, start + label.length)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return ranges.sort((left, right) => left.start - right.start)
|
|
||||||
}
|
|
||||||
|
|
||||||
snapImageCursor(visible: string, cursor: number, previousCursor: number): number {
|
|
||||||
const range = this.imagePlaceholderRanges(visible)
|
|
||||||
.find(({ start, end }) => cursor > start && cursor < end)
|
|
||||||
if (!range) return cursor
|
|
||||||
if (previousCursor <= range.start) return range.end
|
|
||||||
if (previousCursor >= range.end) return range.start
|
|
||||||
return cursor - range.start < range.end - cursor ? range.start : range.end
|
|
||||||
}
|
|
||||||
|
|
||||||
moveImageCursor(visible: string, cursor: number, direction: -1 | 1): number | null {
|
|
||||||
const range = this.imagePlaceholderRanges(visible).find(({ start, end }) => (
|
|
||||||
direction < 0
|
|
||||||
? cursor > start && cursor <= end
|
|
||||||
: cursor >= start && cursor < end
|
|
||||||
))
|
|
||||||
if (!range) return null
|
|
||||||
return direction < 0 ? range.start : range.end
|
|
||||||
}
|
|
||||||
|
|
||||||
reconcileImageEdit(
|
|
||||||
previous: string,
|
|
||||||
value: string,
|
|
||||||
cursor: number,
|
|
||||||
): DraftEditReconciliation {
|
|
||||||
let oldStart = 0
|
|
||||||
const sharedLength = Math.min(previous.length, value.length)
|
|
||||||
while (oldStart < sharedLength && previous[oldStart] === value[oldStart]) oldStart += 1
|
|
||||||
|
|
||||||
let oldEnd = previous.length
|
|
||||||
let newEnd = value.length
|
|
||||||
while (
|
|
||||||
oldEnd > oldStart
|
|
||||||
&& newEnd > oldStart
|
|
||||||
&& previous[oldEnd - 1] === value[newEnd - 1]
|
|
||||||
) {
|
|
||||||
oldEnd -= 1
|
|
||||||
newEnd -= 1
|
|
||||||
}
|
|
||||||
|
|
||||||
const ranges = this.imagePlaceholderRanges(previous).filter(({ start, end }) => (
|
|
||||||
oldStart === oldEnd
|
|
||||||
? oldStart > start && oldStart < end
|
|
||||||
: oldStart < end && oldEnd > start
|
|
||||||
))
|
|
||||||
if (!ranges.length) return { value, cursor, removedImages: [] }
|
|
||||||
|
|
||||||
const replaceStart = Math.min(oldStart, ...ranges.map((range) => range.start))
|
|
||||||
const replaceEnd = Math.max(oldEnd, ...ranges.map((range) => range.end))
|
|
||||||
const inserted = value.slice(oldStart, newEnd)
|
|
||||||
const reconciled = previous.slice(0, replaceStart) + inserted + previous.slice(replaceEnd)
|
|
||||||
const missing = [...this.images.keys()].filter((label) => !reconciled.includes(label))
|
|
||||||
for (const label of missing) this.images.delete(label)
|
|
||||||
return {
|
|
||||||
value: reconciled,
|
|
||||||
cursor: replaceStart + inserted.length,
|
|
||||||
removedImages: missing.map((label) => label.slice(1, -1)),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
hasImageLabelConflict(visible: string): boolean {
|
|
||||||
const expanded = this.expandPastes(visible)
|
|
||||||
return [...this.images.keys()]
|
|
||||||
.some((label) => this.labelOccurrences(expanded, label) !== 1)
|
|
||||||
}
|
|
||||||
|
|
||||||
expand(visible: string): string {
|
|
||||||
let expanded = this.expandPastes(visible)
|
|
||||||
for (const label of this.images.keys()) {
|
|
||||||
if (this.labelOccurrences(expanded, label) === 1) expanded = expanded.replace(label, "")
|
|
||||||
}
|
|
||||||
return expanded
|
|
||||||
}
|
|
||||||
|
|
||||||
display(visible: string): string {
|
|
||||||
return this.expandPastes(visible)
|
|
||||||
}
|
|
||||||
|
|
||||||
media(visible: string): OutboundMedia[] {
|
|
||||||
return [...this.images]
|
|
||||||
.filter(([label]) => visible.includes(label))
|
|
||||||
.sort(([left], [right]) => visible.indexOf(left) - visible.indexOf(right))
|
|
||||||
.map(([, media]) => media)
|
|
||||||
}
|
|
||||||
|
|
||||||
prune(visible: string): void {
|
prune(visible: string): void {
|
||||||
for (const label of this.pastes.keys()) {
|
for (const label of this.pastes.keys()) {
|
||||||
if (!visible.includes(label)) this.pastes.delete(label)
|
if (!visible.includes(label)) this.pastes.delete(label)
|
||||||
}
|
}
|
||||||
for (const label of this.images.keys()) {
|
|
||||||
if (!visible.includes(label)) this.images.delete(label)
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
clear(): void {
|
clear(): void {
|
||||||
this.pastes.clear()
|
this.pastes.clear()
|
||||||
this.images.clear()
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -14,7 +14,6 @@ const workspace = process.env.NANOBOT_TUI_WORKSPACE?.trim() || ""
|
|||||||
const hostWorkspace = process.cwd()
|
const hostWorkspace = process.cwd()
|
||||||
const bootstrapUrl = process.env.NANOBOT_TUI_BOOTSTRAP_URL?.trim() || ""
|
const bootstrapUrl = process.env.NANOBOT_TUI_BOOTSTRAP_URL?.trim() || ""
|
||||||
const wsUrl = process.env.NANOBOT_TUI_WS_URL?.trim() || ""
|
const wsUrl = process.env.NANOBOT_TUI_WS_URL?.trim() || ""
|
||||||
const healthUrl = process.env.NANOBOT_TUI_HEALTH_URL?.trim() || ""
|
|
||||||
const gatewayStopCommand = process.env.NANOBOT_TUI_GATEWAY_STOP_COMMAND?.trim()
|
const gatewayStopCommand = process.env.NANOBOT_TUI_GATEWAY_STOP_COMMAND?.trim()
|
||||||
|| "nanobot gateway stop"
|
|| "nanobot gateway stop"
|
||||||
if (!bootstrapUrl && !wsUrl) {
|
if (!bootstrapUrl && !wsUrl) {
|
||||||
@@ -25,7 +24,6 @@ const options: AppOptions = {
|
|||||||
? {
|
? {
|
||||||
bootstrapUrl,
|
bootstrapUrl,
|
||||||
bootstrapSecret: process.env.NANOBOT_TUI_BOOTSTRAP_SECRET?.trim() || "",
|
bootstrapSecret: process.env.NANOBOT_TUI_BOOTSTRAP_SECRET?.trim() || "",
|
||||||
healthUrl: healthUrl || undefined,
|
|
||||||
}
|
}
|
||||||
: { wsUrl }),
|
: { wsUrl }),
|
||||||
apiUrl: process.env.NANOBOT_TUI_API_URL?.trim() || "",
|
apiUrl: process.env.NANOBOT_TUI_API_URL?.trim() || "",
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ import type { MessageOptions } from "./protocol"
|
|||||||
|
|
||||||
export interface QueuedPrompt {
|
export interface QueuedPrompt {
|
||||||
content: string
|
content: string
|
||||||
displayContent?: string
|
|
||||||
options: MessageOptions
|
options: MessageOptions
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+3
-214
@@ -3,19 +3,14 @@ import { describe, expect, test } from "bun:test"
|
|||||||
import {
|
import {
|
||||||
NanobotClient,
|
NanobotClient,
|
||||||
GatewayConnectionError,
|
GatewayConnectionError,
|
||||||
connectionEndpoint,
|
|
||||||
fetchAvailableSkills,
|
fetchAvailableSkills,
|
||||||
fetchGatewayConnection,
|
fetchGatewayConnection,
|
||||||
fetchGatewayHealth,
|
|
||||||
fetchHistory,
|
fetchHistory,
|
||||||
fetchMentionCandidates,
|
fetchMentionCandidates,
|
||||||
fetchRuntimeControls,
|
fetchRuntimeControls,
|
||||||
fetchSessionContext,
|
fetchSessionContext,
|
||||||
fetchSessions,
|
fetchSessions,
|
||||||
fetchSlashCommands,
|
fetchSlashCommands,
|
||||||
sanitizeConnectionFailure,
|
|
||||||
type ConnectionStatus,
|
|
||||||
type ConnectionStatusInfo,
|
|
||||||
type InboundEvent,
|
type InboundEvent,
|
||||||
} from "./protocol"
|
} from "./protocol"
|
||||||
|
|
||||||
@@ -44,12 +39,6 @@ class FakeSocket {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
async function waitUntil(predicate: () => boolean, timeout = 1_000): Promise<void> {
|
|
||||||
const deadline = Date.now() + timeout
|
|
||||||
while (!predicate() && Date.now() < deadline) await Bun.sleep(2)
|
|
||||||
if (!predicate()) throw new Error(`condition was not met within ${timeout}ms`)
|
|
||||||
}
|
|
||||||
|
|
||||||
describe("gateway protocol", () => {
|
describe("gateway protocol", () => {
|
||||||
test("bootstraps fresh websocket and API credentials", async () => {
|
test("bootstraps fresh websocket and API credentials", async () => {
|
||||||
const original = globalThis.fetch
|
const original = globalThis.fetch
|
||||||
@@ -81,41 +70,6 @@ describe("gateway protocol", () => {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
test("classifies gateway health without sending credentials", async () => {
|
|
||||||
const original = globalThis.fetch
|
|
||||||
const requests: Array<{ url: string; headers: Headers }> = []
|
|
||||||
const responses = [
|
|
||||||
new Response(JSON.stringify({
|
|
||||||
status: "degraded",
|
|
||||||
process: "alive",
|
|
||||||
ready: false,
|
|
||||||
websocket: "unavailable",
|
|
||||||
}), { status: 503 }),
|
|
||||||
new Response(JSON.stringify({
|
|
||||||
status: "ok",
|
|
||||||
process: "alive",
|
|
||||||
ready: true,
|
|
||||||
websocket: "running",
|
|
||||||
})),
|
|
||||||
new Response("not json"),
|
|
||||||
]
|
|
||||||
globalThis.fetch = ((input: string | URL | Request, init?: RequestInit) => {
|
|
||||||
requests.push({ url: String(input), headers: new Headers(init?.headers) })
|
|
||||||
return Promise.resolve(responses.shift() || new Response("missing", { status: 500 }))
|
|
||||||
}) as typeof fetch
|
|
||||||
|
|
||||||
try {
|
|
||||||
const healthUrl = "http://127.0.0.1:18790/health"
|
|
||||||
expect(await fetchGatewayHealth(healthUrl)).toBe("degraded")
|
|
||||||
expect(await fetchGatewayHealth(healthUrl)).toBe("ready")
|
|
||||||
expect(await fetchGatewayHealth(healthUrl)).toBe("unreachable")
|
|
||||||
expect(requests.map(({ url }) => url)).toEqual([healthUrl, healthUrl, healthUrl])
|
|
||||||
expect(requests.every(({ headers }) => [...headers].length === 0)).toBe(true)
|
|
||||||
} finally {
|
|
||||||
globalThis.fetch = original
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
test("rejects malformed bootstrap responses without retrying", async () => {
|
test("rejects malformed bootstrap responses without retrying", async () => {
|
||||||
const original = globalThis.fetch
|
const original = globalThis.fetch
|
||||||
globalThis.fetch = (() => Promise.resolve(new Response("not json"))) as unknown as typeof fetch
|
globalThis.fetch = (() => Promise.resolve(new Response("not json"))) as unknown as typeof fetch
|
||||||
@@ -184,13 +138,8 @@ describe("gateway protocol", () => {
|
|||||||
|
|
||||||
try {
|
try {
|
||||||
const connections: string[] = []
|
const connections: string[] = []
|
||||||
let healthChecks = 0
|
|
||||||
const client = new NanobotClient({
|
const client = new NanobotClient({
|
||||||
resolveConnection: () => new Promise((resolve) => { resolveConnection = resolve }),
|
resolveConnection: () => new Promise((resolve) => { resolveConnection = resolve }),
|
||||||
checkHealth: async () => {
|
|
||||||
healthChecks += 1
|
|
||||||
return "ready"
|
|
||||||
},
|
|
||||||
onConnection: (connection) => connections.push(connection.apiToken),
|
onConnection: (connection) => connections.push(connection.apiToken),
|
||||||
onEvent: () => undefined,
|
onEvent: () => undefined,
|
||||||
onStatus: () => undefined,
|
onStatus: () => undefined,
|
||||||
@@ -206,7 +155,6 @@ describe("gateway protocol", () => {
|
|||||||
await Bun.sleep(1)
|
await Bun.sleep(1)
|
||||||
expect(requestedUrl).toBe("ws://nanobot.test/ws?token=fresh")
|
expect(requestedUrl).toBe("ws://nanobot.test/ws?token=fresh")
|
||||||
expect(connections).toEqual(["fresh-api-token"])
|
expect(connections).toEqual(["fresh-api-token"])
|
||||||
expect(healthChecks).toBe(0)
|
|
||||||
client.close()
|
client.close()
|
||||||
} finally {
|
} finally {
|
||||||
Object.defineProperty(globalThis, "WebSocket", { configurable: true, value: original })
|
Object.defineProperty(globalThis, "WebSocket", { configurable: true, value: original })
|
||||||
@@ -252,133 +200,6 @@ describe("gateway protocol", () => {
|
|||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|
||||||
test("escalates a refused bootstrap with safe endpoint and retry diagnostics", async () => {
|
|
||||||
const original = globalThis.fetch
|
|
||||||
const bootstrapUrl = "http://bootstrap-user:bootstrap-pass@127.0.0.1:8769"
|
|
||||||
+ "/webui/bootstrap?token=socket-secret"
|
|
||||||
const bootstrapSecret = "bootstrap-secret"
|
|
||||||
const statuses: Array<{
|
|
||||||
status: ConnectionStatus
|
|
||||||
detail?: string
|
|
||||||
info?: ConnectionStatusInfo
|
|
||||||
}> = []
|
|
||||||
const refused = new TypeError(
|
|
||||||
`fetch failed for ${bootstrapUrl}&api_token=api-secret`,
|
|
||||||
{
|
|
||||||
cause: Object.assign(new Error("connect ECONNREFUSED 127.0.0.1:8769"), {
|
|
||||||
code: "ECONNREFUSED",
|
|
||||||
}),
|
|
||||||
},
|
|
||||||
)
|
|
||||||
globalThis.fetch = (() => Promise.reject(refused)) as unknown as typeof fetch
|
|
||||||
const client = new NanobotClient({
|
|
||||||
resolveConnection: () => fetchGatewayConnection(
|
|
||||||
bootstrapUrl,
|
|
||||||
bootstrapSecret,
|
|
||||||
"http://127.0.0.1:8769",
|
|
||||||
"tui-42",
|
|
||||||
),
|
|
||||||
targetEndpoint: connectionEndpoint(bootstrapUrl),
|
|
||||||
checkHealth: async () => "degraded",
|
|
||||||
startupFailureDelayMs: 8,
|
|
||||||
reconnectDelayMs: 100,
|
|
||||||
onEvent: () => undefined,
|
|
||||||
onStatus: (status, detail, info) => statuses.push({ status, detail, info }),
|
|
||||||
})
|
|
||||||
|
|
||||||
try {
|
|
||||||
client.connect()
|
|
||||||
await waitUntil(() => statuses.some(({ status }) => status === "unavailable"))
|
|
||||||
const failure = [...statuses].reverse().find(({ status }) => status === "unavailable")
|
|
||||||
|
|
||||||
expect(statuses[0]?.status).toBe("starting")
|
|
||||||
expect(failure?.detail).toBe("connection refused")
|
|
||||||
expect(failure?.info).toMatchObject({
|
|
||||||
endpoint: "127.0.0.1:8769",
|
|
||||||
attempt: 1,
|
|
||||||
elapsedMs: expect.any(Number),
|
|
||||||
health: "degraded",
|
|
||||||
})
|
|
||||||
const visible = JSON.stringify(statuses)
|
|
||||||
expect(visible).not.toContain("bootstrap-user")
|
|
||||||
expect(visible).not.toContain("bootstrap-pass")
|
|
||||||
expect(visible).not.toContain(bootstrapSecret)
|
|
||||||
expect(visible).not.toContain("socket-secret")
|
|
||||||
expect(visible).not.toContain("api-secret")
|
|
||||||
expect(visible).not.toContain("/webui/bootstrap")
|
|
||||||
} finally {
|
|
||||||
client.close()
|
|
||||||
globalThis.fetch = original
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
test("recovers after sustained bootstrap failures without hiding the outage", async () => {
|
|
||||||
const original = globalThis.WebSocket
|
|
||||||
const sockets: FakeSocket[] = []
|
|
||||||
let available = false
|
|
||||||
let attempts = 0
|
|
||||||
const statuses: ConnectionStatus[] = []
|
|
||||||
Object.defineProperty(globalThis, "WebSocket", {
|
|
||||||
configurable: true,
|
|
||||||
value: class extends FakeSocket {
|
|
||||||
constructor() {
|
|
||||||
super()
|
|
||||||
sockets.push(this)
|
|
||||||
}
|
|
||||||
},
|
|
||||||
})
|
|
||||||
const client = new NanobotClient({
|
|
||||||
resolveConnection: async () => {
|
|
||||||
attempts += 1
|
|
||||||
if (!available) {
|
|
||||||
throw Object.assign(new Error("connect ECONNREFUSED"), { code: "ECONNREFUSED" })
|
|
||||||
}
|
|
||||||
return {
|
|
||||||
wsUrl: "ws://127.0.0.1:8769/ws?token=fresh",
|
|
||||||
apiUrl: "http://127.0.0.1:8769",
|
|
||||||
apiToken: "fresh-api-token",
|
|
||||||
}
|
|
||||||
},
|
|
||||||
targetEndpoint: "127.0.0.1:8769",
|
|
||||||
checkHealth: async () => available ? "ready" : "degraded",
|
|
||||||
startupFailureDelayMs: 8,
|
|
||||||
reconnectDelayMs: 2,
|
|
||||||
startupRetryMaxDelayMs: 2,
|
|
||||||
onEvent: () => undefined,
|
|
||||||
onStatus: (status) => statuses.push(status),
|
|
||||||
})
|
|
||||||
|
|
||||||
try {
|
|
||||||
client.connect()
|
|
||||||
await waitUntil(() => statuses.includes("unavailable"))
|
|
||||||
available = true
|
|
||||||
await waitUntil(() => sockets.length === 1)
|
|
||||||
sockets[0]?.emit("open")
|
|
||||||
await waitUntil(() => statuses.at(-1) === "connected")
|
|
||||||
|
|
||||||
expect(attempts).toBeGreaterThan(1)
|
|
||||||
expect(statuses.indexOf("starting")).toBeLessThan(statuses.indexOf("unavailable"))
|
|
||||||
expect(statuses.indexOf("unavailable")).toBeLessThan(statuses.lastIndexOf("connected"))
|
|
||||||
} finally {
|
|
||||||
client.close()
|
|
||||||
Object.defineProperty(globalThis, "WebSocket", { configurable: true, value: original })
|
|
||||||
}
|
|
||||||
})
|
|
||||||
|
|
||||||
test("sanitizes arbitrary connection errors and authenticated URLs", () => {
|
|
||||||
const authenticated = "wss://user:password@127.0.0.1:8769/ws"
|
|
||||||
+ "?token=socket-secret&api_token=api-secret"
|
|
||||||
const unknown = new Error(`could not reach ${authenticated}`)
|
|
||||||
const refused = Object.assign(new Error(`ECONNREFUSED ${authenticated}`), {
|
|
||||||
code: "ECONNREFUSED",
|
|
||||||
})
|
|
||||||
|
|
||||||
expect(connectionEndpoint(authenticated)).toBe("127.0.0.1:8769")
|
|
||||||
expect(sanitizeConnectionFailure(unknown)).toBe("connection failed")
|
|
||||||
expect(sanitizeConnectionFailure(refused)).toBe("connection refused")
|
|
||||||
expect(sanitizeConnectionFailure(unknown)).not.toContain("socket-secret")
|
|
||||||
})
|
|
||||||
|
|
||||||
test("reports a permanent bootstrap rejection without retrying", async () => {
|
test("reports a permanent bootstrap rejection without retrying", async () => {
|
||||||
let attempts = 0
|
let attempts = 0
|
||||||
const statuses: string[] = []
|
const statuses: string[] = []
|
||||||
@@ -459,7 +280,6 @@ describe("gateway protocol", () => {
|
|||||||
}),
|
}),
|
||||||
})
|
})
|
||||||
client.send("hello", {
|
client.send("hello", {
|
||||||
media: [{ data_url: "data:image/png;base64,AAAA", name: "clipboard-image-1.png" }],
|
|
||||||
cliApps: [{ name: "github" }],
|
cliApps: [{ name: "github" }],
|
||||||
sessionMentions: [{ name: "plan", session_key: "websocket:plan" }],
|
sessionMentions: [{ name: "plan", session_key: "websocket:plan" }],
|
||||||
userShell: true,
|
userShell: true,
|
||||||
@@ -478,9 +298,6 @@ describe("gateway protocol", () => {
|
|||||||
expect(outbound[1]?.chat_id).toBe("terminal")
|
expect(outbound[1]?.chat_id).toBe("terminal")
|
||||||
expect(outbound[1]?.content).toBe("hello")
|
expect(outbound[1]?.content).toBe("hello")
|
||||||
expect(outbound[1]?.user_shell).toBe(true)
|
expect(outbound[1]?.user_shell).toBe(true)
|
||||||
expect(outbound[1]?.media).toEqual([
|
|
||||||
{ data_url: "data:image/png;base64,AAAA", name: "clipboard-image-1.png" },
|
|
||||||
])
|
|
||||||
expect(outbound[1]?.cli_apps).toEqual([{ name: "github" }])
|
expect(outbound[1]?.cli_apps).toEqual([{ name: "github" }])
|
||||||
expect(outbound[1]?.session_mentions).toEqual([
|
expect(outbound[1]?.session_mentions).toEqual([
|
||||||
{ name: "plan", session_key: "websocket:plan" },
|
{ name: "plan", session_key: "websocket:plan" },
|
||||||
@@ -771,19 +588,13 @@ describe("gateway protocol", () => {
|
|||||||
})
|
})
|
||||||
|
|
||||||
try {
|
try {
|
||||||
const statuses: Array<{
|
|
||||||
status: ConnectionStatus
|
|
||||||
detail?: string
|
|
||||||
info?: ConnectionStatusInfo
|
|
||||||
}> = []
|
|
||||||
const client = new NanobotClient({
|
const client = new NanobotClient({
|
||||||
url: "ws://nanobot.test/ws",
|
url: "ws://nanobot.test/ws",
|
||||||
reconnectDelayMs: 1,
|
reconnectDelayMs: 1,
|
||||||
onEvent: () => undefined,
|
onEvent: () => undefined,
|
||||||
onStatus: (status, detail, info) => statuses.push({ status, detail, info }),
|
onStatus: () => undefined,
|
||||||
})
|
})
|
||||||
client.connect()
|
client.connect()
|
||||||
sockets[0]?.emit("open")
|
|
||||||
sockets[0]?.emit("message", {
|
sockets[0]?.emit("message", {
|
||||||
data: JSON.stringify({ event: "ready", chat_id: "", client_id: "client" }),
|
data: JSON.stringify({ event: "ready", chat_id: "", client_id: "client" }),
|
||||||
})
|
})
|
||||||
@@ -794,21 +605,11 @@ describe("gateway protocol", () => {
|
|||||||
await Bun.sleep(5)
|
await Bun.sleep(5)
|
||||||
|
|
||||||
expect(sockets).toHaveLength(2)
|
expect(sockets).toHaveLength(2)
|
||||||
const reconnecting = [...statuses].reverse().find(
|
|
||||||
({ status }) => status === "reconnecting",
|
|
||||||
)
|
|
||||||
expect(reconnecting).toMatchObject({
|
|
||||||
status: "reconnecting",
|
|
||||||
detail: "connection closed",
|
|
||||||
info: { endpoint: "nanobot.test", attempt: 1 },
|
|
||||||
})
|
|
||||||
sockets[1]?.emit("open")
|
|
||||||
sockets[1]?.emit("message", {
|
sockets[1]?.emit("message", {
|
||||||
data: JSON.stringify({ event: "ready", chat_id: "", client_id: "client-2" }),
|
data: JSON.stringify({ event: "ready", chat_id: "", client_id: "client-2" }),
|
||||||
})
|
})
|
||||||
const outbound = sockets[1]?.sent.map((value) => JSON.parse(value)) || []
|
const outbound = sockets[1]?.sent.map((value) => JSON.parse(value)) || []
|
||||||
expect(outbound).toEqual([{ type: "attach", chat_id: "generated-chat" }])
|
expect(outbound).toEqual([{ type: "attach", chat_id: "generated-chat" }])
|
||||||
expect(statuses.at(-1)?.status).toBe("connected")
|
|
||||||
client.close()
|
client.close()
|
||||||
} finally {
|
} finally {
|
||||||
Object.defineProperty(globalThis, "WebSocket", { configurable: true, value: original })
|
Object.defineProperty(globalThis, "WebSocket", { configurable: true, value: original })
|
||||||
@@ -879,12 +680,6 @@ describe("gateway protocol", () => {
|
|||||||
return Promise.resolve(new Response(JSON.stringify({
|
return Promise.resolve(new Response(JSON.stringify({
|
||||||
messages: [
|
messages: [
|
||||||
{ role: "user", content: "hello", turnId: "turn-1" },
|
{ role: "user", content: "hello", turnId: "turn-1" },
|
||||||
{
|
|
||||||
role: "user",
|
|
||||||
content: "",
|
|
||||||
turnId: "turn-image",
|
|
||||||
media: [{ kind: "image", url: "/api/media/sig/image", name: "shot.png" }],
|
|
||||||
},
|
|
||||||
{
|
{
|
||||||
role: "tool",
|
role: "tool",
|
||||||
kind: "trace",
|
kind: "trace",
|
||||||
@@ -893,7 +688,7 @@ describe("gateway protocol", () => {
|
|||||||
toolEvents: [{ phase: "end", call_id: "read-1", name: "read_file" }],
|
toolEvents: [{ phase: "end", call_id: "read-1", name: "read_file" }],
|
||||||
},
|
},
|
||||||
{ role: "assistant", kind: "reasoning", content: "private thought" },
|
{ role: "assistant", kind: "reasoning", content: "private thought" },
|
||||||
{ role: "assistant", content: "hi", forkIndex: 2 },
|
{ role: "assistant", content: "hi", forkIndex: 1 },
|
||||||
],
|
],
|
||||||
page: { has_more_before: true, before_cursor: "older-1" },
|
page: { has_more_before: true, before_cursor: "older-1" },
|
||||||
})))
|
})))
|
||||||
@@ -904,18 +699,12 @@ describe("gateway protocol", () => {
|
|||||||
expect(history).toEqual({
|
expect(history).toEqual({
|
||||||
messages: [
|
messages: [
|
||||||
{ role: "user", content: "hello", turnId: "turn-1" },
|
{ role: "user", content: "hello", turnId: "turn-1" },
|
||||||
{
|
|
||||||
role: "user",
|
|
||||||
content: "",
|
|
||||||
turnId: "turn-image",
|
|
||||||
media: [{ kind: "image", url: "/api/media/sig/image", name: "shot.png" }],
|
|
||||||
},
|
|
||||||
{
|
{
|
||||||
role: "activity",
|
role: "activity",
|
||||||
content: "read_file",
|
content: "read_file",
|
||||||
toolEvents: [{ phase: "end", call_id: "read-1", name: "read_file" }],
|
toolEvents: [{ phase: "end", call_id: "read-1", name: "read_file" }],
|
||||||
},
|
},
|
||||||
{ role: "assistant", content: "hi", forkIndex: 2 },
|
{ role: "assistant", content: "hi", forkIndex: 1 },
|
||||||
],
|
],
|
||||||
hasMoreBefore: true,
|
hasMoreBefore: true,
|
||||||
beforeCursor: "older-1",
|
beforeCursor: "older-1",
|
||||||
|
|||||||
+21
-237
@@ -1,21 +1,4 @@
|
|||||||
export type ConnectionStatus =
|
export type ConnectionStatus = "connecting" | "connected" | "closed" | "error"
|
||||||
| "starting"
|
|
||||||
| "connecting"
|
|
||||||
| "connected"
|
|
||||||
| "reconnecting"
|
|
||||||
| "unavailable"
|
|
||||||
| "closed"
|
|
||||||
| "error"
|
|
||||||
|
|
||||||
export interface ConnectionStatusInfo {
|
|
||||||
endpoint: string
|
|
||||||
attempt: number
|
|
||||||
elapsedMs: number
|
|
||||||
retryInMs?: number
|
|
||||||
health?: GatewayHealthStatus
|
|
||||||
}
|
|
||||||
|
|
||||||
export type GatewayHealthStatus = "ready" | "degraded" | "unreachable"
|
|
||||||
|
|
||||||
export interface ToolProgressEvent {
|
export interface ToolProgressEvent {
|
||||||
version?: number
|
version?: number
|
||||||
@@ -53,17 +36,12 @@ interface FileDiff {
|
|||||||
text?: string
|
text?: string
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface MediaAttachment {
|
interface MediaAttachment {
|
||||||
kind: "image" | "video" | "file"
|
kind: "image" | "video" | "file"
|
||||||
url: string
|
url: string
|
||||||
name?: string
|
name?: string
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface OutboundMedia {
|
|
||||||
data_url: string
|
|
||||||
name?: string
|
|
||||||
}
|
|
||||||
|
|
||||||
export interface WorkspaceScopePayload {
|
export interface WorkspaceScopePayload {
|
||||||
project_path: string
|
project_path: string
|
||||||
project_name?: string
|
project_name?: string
|
||||||
@@ -186,7 +164,6 @@ type OutboundEvent =
|
|||||||
turn_id: string
|
turn_id: string
|
||||||
webui: true
|
webui: true
|
||||||
workspace_scope?: WorkspaceScopePayload
|
workspace_scope?: WorkspaceScopePayload
|
||||||
media?: OutboundMedia[]
|
|
||||||
cli_apps?: Array<{ name: string }>
|
cli_apps?: Array<{ name: string }>
|
||||||
mcp_presets?: Array<{ name: string }>
|
mcp_presets?: Array<{ name: string }>
|
||||||
session_mentions?: SessionMention[]
|
session_mentions?: SessionMention[]
|
||||||
@@ -195,16 +172,14 @@ type OutboundEvent =
|
|||||||
export interface ClientOptions {
|
export interface ClientOptions {
|
||||||
url?: string
|
url?: string
|
||||||
resolveConnection?: () => Promise<GatewayConnection>
|
resolveConnection?: () => Promise<GatewayConnection>
|
||||||
checkHealth?: () => Promise<GatewayHealthStatus>
|
|
||||||
onConnection?: (connection: GatewayConnection) => void
|
onConnection?: (connection: GatewayConnection) => void
|
||||||
targetEndpoint?: string
|
connectionRetryLabel?: string
|
||||||
startupFailureDelayMs?: number
|
|
||||||
startupRetryMaxDelayMs?: number
|
startupRetryMaxDelayMs?: number
|
||||||
chatId?: string
|
chatId?: string
|
||||||
initialWorkspaceScope?: WorkspaceScopePayload
|
initialWorkspaceScope?: WorkspaceScopePayload
|
||||||
reconnectDelayMs?: number
|
reconnectDelayMs?: number
|
||||||
onEvent: (event: InboundEvent) => void
|
onEvent: (event: InboundEvent) => void
|
||||||
onStatus: (status: ConnectionStatus, detail?: string, info?: ConnectionStatusInfo) => void
|
onStatus: (status: ConnectionStatus, detail?: string) => void
|
||||||
}
|
}
|
||||||
|
|
||||||
export interface GatewayApiConnection {
|
export interface GatewayApiConnection {
|
||||||
@@ -231,7 +206,6 @@ export interface HistoryMessage {
|
|||||||
role: "user" | "assistant" | "activity"
|
role: "user" | "assistant" | "activity"
|
||||||
content: string
|
content: string
|
||||||
turnId?: string
|
turnId?: string
|
||||||
media?: MediaAttachment[]
|
|
||||||
toolEvents?: ToolProgressEvent[]
|
toolEvents?: ToolProgressEvent[]
|
||||||
fileEdits?: FileEditEvent[]
|
fileEdits?: FileEditEvent[]
|
||||||
forkIndex?: number
|
forkIndex?: number
|
||||||
@@ -295,7 +269,6 @@ export interface SkillCandidate {
|
|||||||
}
|
}
|
||||||
|
|
||||||
export interface MessageOptions {
|
export interface MessageOptions {
|
||||||
media?: OutboundMedia[]
|
|
||||||
cliApps?: Array<{ name: string }>
|
cliApps?: Array<{ name: string }>
|
||||||
mcpPresets?: Array<{ name: string }>
|
mcpPresets?: Array<{ name: string }>
|
||||||
sessionMentions?: SessionMention[]
|
sessionMentions?: SessionMention[]
|
||||||
@@ -631,20 +604,18 @@ export async function fetchHistory(
|
|||||||
(role !== "user" && role !== "assistant")
|
(role !== "user" && role !== "assistant")
|
||||||
|| message.kind === "reasoning"
|
|| message.kind === "reasoning"
|
||||||
|| typeof content !== "string"
|
|| typeof content !== "string"
|
||||||
|
|| !content.trim()
|
||||||
) {
|
) {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
const media = Array.isArray(message.media) ? message.media.filter(isMediaAttachment) : []
|
|
||||||
if (role === "user") {
|
if (role === "user") {
|
||||||
if (!content.trim() && !media.length) continue
|
|
||||||
userIndex += 1
|
userIndex += 1
|
||||||
messages.push({
|
messages.push({
|
||||||
role: "user",
|
role: "user",
|
||||||
content,
|
content,
|
||||||
...(media.length ? { media } : {}),
|
|
||||||
...(typeof message.turnId === "string" ? { turnId: message.turnId } : {}),
|
...(typeof message.turnId === "string" ? { turnId: message.turnId } : {}),
|
||||||
})
|
})
|
||||||
} else if (content.trim()) {
|
} else {
|
||||||
messages.push({ role: "assistant", content, forkIndex: userIndex })
|
messages.push({ role: "assistant", content, forkIndex: userIndex })
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -969,92 +940,6 @@ export async function fetchGatewayConnection(
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
/** Read gateway readiness without sending bootstrap or API credentials. */
|
|
||||||
export async function fetchGatewayHealth(
|
|
||||||
healthUrl: string,
|
|
||||||
timeoutMs = 400,
|
|
||||||
): Promise<GatewayHealthStatus> {
|
|
||||||
const controller = new AbortController()
|
|
||||||
const timer = setTimeout(() => controller.abort(), timeoutMs)
|
|
||||||
try {
|
|
||||||
const response = await fetch(healthUrl, { signal: controller.signal })
|
|
||||||
if (response.status !== 200 && response.status !== 503) return "unreachable"
|
|
||||||
const payload: unknown = await response.json()
|
|
||||||
if (!isRecord(payload)) return "unreachable"
|
|
||||||
if (
|
|
||||||
response.status === 503
|
|
||||||
&& payload.status === "degraded"
|
|
||||||
&& payload.ready === false
|
|
||||||
&& payload.process === "alive"
|
|
||||||
) return "degraded"
|
|
||||||
if (response.status === 200 && payload.status === "ok" && payload.ready !== false) {
|
|
||||||
return "ready"
|
|
||||||
}
|
|
||||||
return "unreachable"
|
|
||||||
} catch {
|
|
||||||
return "unreachable"
|
|
||||||
} finally {
|
|
||||||
clearTimeout(timer)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/** Return only the authority users can act on, never credentials or an authenticated path. */
|
|
||||||
export function connectionEndpoint(value: string | undefined): string {
|
|
||||||
if (!value) return "local gateway"
|
|
||||||
try {
|
|
||||||
return new URL(value).host || "local gateway"
|
|
||||||
} catch {
|
|
||||||
return "local gateway"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
/** Reduce arbitrary fetch/WebSocket errors to a small set of credential-safe reasons. */
|
|
||||||
export function sanitizeConnectionFailure(error: unknown): string {
|
|
||||||
const signals: string[] = []
|
|
||||||
const seen = new Set<unknown>()
|
|
||||||
const collect = (value: unknown): void => {
|
|
||||||
if (value === null || value === undefined || seen.has(value)) return
|
|
||||||
if (typeof value === "object") seen.add(value)
|
|
||||||
if (typeof value === "string") {
|
|
||||||
signals.push(value)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if (value instanceof Error) {
|
|
||||||
signals.push(value.name, value.message)
|
|
||||||
collect(value.cause)
|
|
||||||
if (value instanceof AggregateError) {
|
|
||||||
for (const nested of value.errors) collect(nested)
|
|
||||||
}
|
|
||||||
return
|
|
||||||
}
|
|
||||||
if (!isRecord(value)) return
|
|
||||||
if (typeof value.code === "string") signals.push(value.code)
|
|
||||||
collect(value.cause)
|
|
||||||
if (Array.isArray(value.errors)) {
|
|
||||||
for (const nested of value.errors) collect(nested)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
collect(error)
|
|
||||||
const signal = signals.join(" ")
|
|
||||||
if (/ECONNREFUSED|connection refused/iu.test(signal)) return "connection refused"
|
|
||||||
if (/ETIMEDOUT|timed? out|timeout/iu.test(signal)) return "connection timed out"
|
|
||||||
if (/ENOTFOUND|EAI_AGAIN|name not resolved|host not found/iu.test(signal)) {
|
|
||||||
return "host not found"
|
|
||||||
}
|
|
||||||
if (/certificate|TLS|SSL/iu.test(signal)) return "secure connection failed"
|
|
||||||
const bootstrapStatus = signal.match(/gateway bootstrap failed:\s*HTTP\s*(\d{3})/iu)
|
|
||||||
if (bootstrapStatus?.[1]) return `gateway bootstrap failed: HTTP ${bootstrapStatus[1]}`
|
|
||||||
if (/bootstrap response is missing ws_url/iu.test(signal)) {
|
|
||||||
return "gateway bootstrap response is missing ws_url"
|
|
||||||
}
|
|
||||||
if (/bootstrap response (?:has an invalid ws_url|is invalid)/iu.test(signal)) {
|
|
||||||
return "gateway bootstrap response is invalid"
|
|
||||||
}
|
|
||||||
if (/gateway is still starting/iu.test(signal)) return "gateway is still starting"
|
|
||||||
if (/fetch failed|failed to fetch|network error/iu.test(signal)) return "network request failed"
|
|
||||||
return "connection failed"
|
|
||||||
}
|
|
||||||
|
|
||||||
export class NanobotClient {
|
export class NanobotClient {
|
||||||
private socket: WebSocket | null = null
|
private socket: WebSocket | null = null
|
||||||
private chatId = ""
|
private chatId = ""
|
||||||
@@ -1064,22 +949,13 @@ export class NanobotClient {
|
|||||||
private closedByClient = false
|
private closedByClient = false
|
||||||
private opening = false
|
private opening = false
|
||||||
private connectedOnce = false
|
private connectedOnce = false
|
||||||
private connectionAttempt = 0
|
|
||||||
private retryStartedAt = 0
|
|
||||||
private nextRetryAt = 0
|
|
||||||
private lastFailure = ""
|
|
||||||
private healthStatus: GatewayHealthStatus | undefined
|
|
||||||
private failureEscalationTimer: ReturnType<typeof setTimeout> | null = null
|
|
||||||
private readonly endpoint: string
|
|
||||||
private readonly pendingMutations = new Map<string, {
|
private readonly pendingMutations = new Map<string, {
|
||||||
resolve: (value: unknown) => void
|
resolve: (value: unknown) => void
|
||||||
reject: (error: Error) => void
|
reject: (error: Error) => void
|
||||||
timer: ReturnType<typeof setTimeout>
|
timer: ReturnType<typeof setTimeout>
|
||||||
}>()
|
}>()
|
||||||
|
|
||||||
constructor(private readonly options: ClientOptions) {
|
constructor(private readonly options: ClientOptions) {}
|
||||||
this.endpoint = options.targetEndpoint || connectionEndpoint(options.url)
|
|
||||||
}
|
|
||||||
|
|
||||||
get activeChatId(): string {
|
get activeChatId(): string {
|
||||||
return this.chatId
|
return this.chatId
|
||||||
@@ -1087,21 +963,13 @@ export class NanobotClient {
|
|||||||
|
|
||||||
connect(): void {
|
connect(): void {
|
||||||
this.closedByClient = false
|
this.closedByClient = false
|
||||||
this.connectionAttempt = 0
|
|
||||||
this.reconnectAttempt = 0
|
|
||||||
this.retryStartedAt = Date.now()
|
|
||||||
this.nextRetryAt = 0
|
|
||||||
this.lastFailure = ""
|
|
||||||
this.healthStatus = undefined
|
|
||||||
void this.open()
|
void this.open()
|
||||||
}
|
}
|
||||||
|
|
||||||
private async open(): Promise<void> {
|
private async open(): Promise<void> {
|
||||||
if (this.socket || this.opening || this.closedByClient) return
|
if (this.socket || this.opening || this.closedByClient) return
|
||||||
this.opening = true
|
this.opening = true
|
||||||
this.nextRetryAt = 0
|
this.options.onStatus("connecting")
|
||||||
this.connectionAttempt += 1
|
|
||||||
this.reportConnectionProgress()
|
|
||||||
let url = this.options.url
|
let url = this.options.url
|
||||||
try {
|
try {
|
||||||
if (this.options.resolveConnection) {
|
if (this.options.resolveConnection) {
|
||||||
@@ -1112,52 +980,37 @@ export class NanobotClient {
|
|||||||
}
|
}
|
||||||
} catch (error) {
|
} catch (error) {
|
||||||
if (!this.closedByClient) {
|
if (!this.closedByClient) {
|
||||||
this.lastFailure = sanitizeConnectionFailure(error)
|
|
||||||
if (error instanceof GatewayConnectionError && !error.retryable) {
|
if (error instanceof GatewayConnectionError && !error.retryable) {
|
||||||
this.clearFailureEscalation()
|
this.options.onStatus("error", error.message)
|
||||||
this.options.onStatus("error", this.lastFailure, this.connectionInfo())
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
await this.checkHealthAndScheduleReconnect()
|
this.options.onStatus(
|
||||||
|
"connecting",
|
||||||
|
this.options.connectionRetryLabel || "gateway unavailable",
|
||||||
|
)
|
||||||
|
this.scheduleReconnect(false)
|
||||||
}
|
}
|
||||||
return
|
return
|
||||||
} finally {
|
} finally {
|
||||||
this.opening = false
|
this.opening = false
|
||||||
}
|
}
|
||||||
if (!url) {
|
if (!url) {
|
||||||
this.options.onStatus("error", "gateway URL is not configured", this.connectionInfo())
|
this.options.onStatus("error", "gateway URL is not configured")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
let socket: WebSocket
|
const socket = new WebSocket(url)
|
||||||
try {
|
|
||||||
socket = new WebSocket(url)
|
|
||||||
} catch (error) {
|
|
||||||
this.lastFailure = sanitizeConnectionFailure(error)
|
|
||||||
await this.checkHealthAndScheduleReconnect()
|
|
||||||
return
|
|
||||||
}
|
|
||||||
let opened = false
|
|
||||||
this.socket = socket
|
this.socket = socket
|
||||||
socket.addEventListener("open", () => {
|
socket.addEventListener("open", () => {
|
||||||
if (this.socket !== socket) return
|
if (this.socket !== socket) return
|
||||||
opened = true
|
|
||||||
this.connectedOnce = true
|
this.connectedOnce = true
|
||||||
this.connectionAttempt = 0
|
|
||||||
this.reconnectAttempt = 0
|
this.reconnectAttempt = 0
|
||||||
this.retryStartedAt = 0
|
this.options.onStatus("connected")
|
||||||
this.nextRetryAt = 0
|
|
||||||
this.lastFailure = ""
|
|
||||||
this.healthStatus = "ready"
|
|
||||||
this.clearFailureEscalation()
|
|
||||||
this.options.onStatus("connected", undefined, this.connectionInfo())
|
|
||||||
})
|
})
|
||||||
socket.addEventListener("message", (message) => {
|
socket.addEventListener("message", (message) => {
|
||||||
if (this.socket === socket) this.handleMessage(String(message.data))
|
if (this.socket === socket) this.handleMessage(String(message.data))
|
||||||
})
|
})
|
||||||
socket.addEventListener("error", () => {
|
socket.addEventListener("error", () => {
|
||||||
if (this.socket !== socket) return
|
if (this.socket === socket) this.options.onStatus("error", "connection failed")
|
||||||
this.lastFailure = "connection failed"
|
|
||||||
this.reportRetryState()
|
|
||||||
})
|
})
|
||||||
socket.addEventListener("close", () => {
|
socket.addEventListener("close", () => {
|
||||||
if (this.socket !== socket) return
|
if (this.socket !== socket) return
|
||||||
@@ -1167,14 +1020,7 @@ export class NanobotClient {
|
|||||||
this.options.onStatus("closed")
|
this.options.onStatus("closed")
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if (opened) {
|
this.scheduleReconnect()
|
||||||
this.connectionAttempt = 0
|
|
||||||
this.reconnectAttempt = 0
|
|
||||||
this.retryStartedAt = Date.now()
|
|
||||||
}
|
|
||||||
if (!this.lastFailure) this.lastFailure = "connection closed"
|
|
||||||
this.reportRetryState()
|
|
||||||
void this.checkHealthAndScheduleReconnect()
|
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1182,7 +1028,6 @@ export class NanobotClient {
|
|||||||
this.closedByClient = true
|
this.closedByClient = true
|
||||||
if (this.reconnectTimer) clearTimeout(this.reconnectTimer)
|
if (this.reconnectTimer) clearTimeout(this.reconnectTimer)
|
||||||
this.reconnectTimer = null
|
this.reconnectTimer = null
|
||||||
this.clearFailureEscalation()
|
|
||||||
const socket = this.socket
|
const socket = this.socket
|
||||||
this.socket = null
|
this.socket = null
|
||||||
socket?.close()
|
socket?.close()
|
||||||
@@ -1200,7 +1045,6 @@ export class NanobotClient {
|
|||||||
webui: true,
|
webui: true,
|
||||||
...(this.workspaceScope ? { workspace_scope: this.workspaceScope } : {}),
|
...(this.workspaceScope ? { workspace_scope: this.workspaceScope } : {}),
|
||||||
...(options.userShell ? { user_shell: true } : {}),
|
...(options.userShell ? { user_shell: true } : {}),
|
||||||
...(options.media?.length ? { media: options.media } : {}),
|
|
||||||
...(options.cliApps?.length ? { cli_apps: options.cliApps } : {}),
|
...(options.cliApps?.length ? { cli_apps: options.cliApps } : {}),
|
||||||
...(options.mcpPresets?.length ? { mcp_presets: options.mcpPresets } : {}),
|
...(options.mcpPresets?.length ? { mcp_presets: options.mcpPresets } : {}),
|
||||||
...(options.sessionMentions?.length
|
...(options.sessionMentions?.length
|
||||||
@@ -1338,80 +1182,20 @@ export class NanobotClient {
|
|||||||
this.options.onEvent(event)
|
this.options.onEvent(event)
|
||||||
}
|
}
|
||||||
|
|
||||||
private scheduleReconnect(): void {
|
private scheduleReconnect(announce = true): void {
|
||||||
if (this.reconnectTimer || this.closedByClient) return
|
if (this.reconnectTimer || this.closedByClient) return
|
||||||
if (!this.retryStartedAt) this.retryStartedAt = Date.now()
|
|
||||||
const base = this.options.reconnectDelayMs ?? 500
|
const base = this.options.reconnectDelayMs ?? 500
|
||||||
const maxDelay = this.connectedOnce
|
const maxDelay = this.connectedOnce
|
||||||
? 8_000
|
? 8_000
|
||||||
: this.options.startupRetryMaxDelayMs ?? 8_000
|
: this.options.startupRetryMaxDelayMs ?? 8_000
|
||||||
const delay = Math.min(maxDelay, base * 2 ** Math.min(this.reconnectAttempt++, 4))
|
const delay = Math.min(maxDelay, base * 2 ** Math.min(this.reconnectAttempt++, 4))
|
||||||
this.nextRetryAt = Date.now() + delay
|
if (announce) this.options.onStatus("connecting", `reconnecting in ${delay}ms`)
|
||||||
this.reportRetryState()
|
|
||||||
this.reconnectTimer = setTimeout(() => {
|
this.reconnectTimer = setTimeout(() => {
|
||||||
this.reconnectTimer = null
|
this.reconnectTimer = null
|
||||||
void this.open()
|
void this.open()
|
||||||
}, delay)
|
}, delay)
|
||||||
}
|
}
|
||||||
|
|
||||||
private async checkHealthAndScheduleReconnect(): Promise<void> {
|
|
||||||
if (this.options.checkHealth) {
|
|
||||||
try {
|
|
||||||
this.healthStatus = await this.options.checkHealth()
|
|
||||||
} catch {
|
|
||||||
this.healthStatus = "unreachable"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if (!this.closedByClient) this.scheduleReconnect()
|
|
||||||
}
|
|
||||||
|
|
||||||
private connectionInfo(): ConnectionStatusInfo {
|
|
||||||
return {
|
|
||||||
endpoint: this.endpoint,
|
|
||||||
attempt: Math.max(1, this.connectionAttempt),
|
|
||||||
elapsedMs: this.retryStartedAt ? Math.max(0, Date.now() - this.retryStartedAt) : 0,
|
|
||||||
...(this.nextRetryAt
|
|
||||||
? { retryInMs: Math.max(0, this.nextRetryAt - Date.now()) }
|
|
||||||
: {}),
|
|
||||||
...(this.healthStatus ? { health: this.healthStatus } : {}),
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
private reportConnectionProgress(): void {
|
|
||||||
if (this.connectedOnce) {
|
|
||||||
this.options.onStatus("reconnecting", this.lastFailure || undefined, this.connectionInfo())
|
|
||||||
return
|
|
||||||
}
|
|
||||||
const phase = this.options.resolveConnection ? "starting" : "connecting"
|
|
||||||
this.options.onStatus(phase, undefined, this.connectionInfo())
|
|
||||||
}
|
|
||||||
|
|
||||||
private reportRetryState(): void {
|
|
||||||
const info = this.connectionInfo()
|
|
||||||
if (this.connectedOnce) {
|
|
||||||
this.options.onStatus("reconnecting", this.lastFailure, info)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
const failureDelay = this.options.startupFailureDelayMs ?? 3_000
|
|
||||||
if (info.elapsedMs >= failureDelay) {
|
|
||||||
this.clearFailureEscalation()
|
|
||||||
this.options.onStatus("unavailable", this.lastFailure, info)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
this.reportConnectionProgress()
|
|
||||||
if (this.failureEscalationTimer) return
|
|
||||||
this.failureEscalationTimer = setTimeout(() => {
|
|
||||||
this.failureEscalationTimer = null
|
|
||||||
if (this.closedByClient || this.connectedOnce || !this.lastFailure) return
|
|
||||||
this.options.onStatus("unavailable", this.lastFailure, this.connectionInfo())
|
|
||||||
}, Math.max(0, failureDelay - info.elapsedMs))
|
|
||||||
}
|
|
||||||
|
|
||||||
private clearFailureEscalation(): void {
|
|
||||||
if (this.failureEscalationTimer) clearTimeout(this.failureEscalationTimer)
|
|
||||||
this.failureEscalationTimer = null
|
|
||||||
}
|
|
||||||
|
|
||||||
private write(event: OutboundEvent): void {
|
private write(event: OutboundEvent): void {
|
||||||
if (!this.socket || this.socket.readyState !== WebSocket.OPEN) {
|
if (!this.socket || this.socket.readyState !== WebSocket.OPEN) {
|
||||||
throw new Error("gateway connection is not open")
|
throw new Error("gateway connection is not open")
|
||||||
|
|||||||
+8
-155
@@ -3,21 +3,14 @@ import {
|
|||||||
MarkdownRenderable,
|
MarkdownRenderable,
|
||||||
RGBA,
|
RGBA,
|
||||||
ScrollBoxRenderable,
|
ScrollBoxRenderable,
|
||||||
StyledText,
|
|
||||||
SyntaxStyle,
|
SyntaxStyle,
|
||||||
TextAttributes,
|
TextAttributes,
|
||||||
TextRenderable,
|
TextRenderable,
|
||||||
type CliRenderer,
|
type CliRenderer,
|
||||||
type TextChunk,
|
|
||||||
type TreeSitterClient,
|
type TreeSitterClient,
|
||||||
} from "@opentui/core"
|
} from "@opentui/core"
|
||||||
|
|
||||||
import type {
|
import type { FileEditEvent, HistoryMessage, ToolProgressEvent } from "./protocol"
|
||||||
FileEditEvent,
|
|
||||||
HistoryMessage,
|
|
||||||
MediaAttachment,
|
|
||||||
ToolProgressEvent,
|
|
||||||
} from "./protocol"
|
|
||||||
import { renderLatexAsUnicode } from "./latex"
|
import { renderLatexAsUnicode } from "./latex"
|
||||||
import { hideScrollbars } from "./scrollbox"
|
import { hideScrollbars } from "./scrollbox"
|
||||||
import { mergeToolEvent, renderToolEvent } from "./tool-renderers"
|
import { mergeToolEvent, renderToolEvent } from "./tool-renderers"
|
||||||
@@ -65,54 +58,6 @@ const ACTIVITY_PREVIEW_LINES = 4
|
|||||||
// subsequent deltas to the renderer cadence.
|
// subsequent deltas to the renderer cadence.
|
||||||
const STREAM_FLUSH_MS = 32
|
const STREAM_FLUSH_MS = 32
|
||||||
|
|
||||||
export interface UserMessageMedia {
|
|
||||||
kind?: MediaAttachment["kind"]
|
|
||||||
name?: string
|
|
||||||
}
|
|
||||||
|
|
||||||
interface UserMessageProjection {
|
|
||||||
imageLabels: string[]
|
|
||||||
attachmentNames: string[]
|
|
||||||
}
|
|
||||||
|
|
||||||
function projectUserMessage(media: readonly UserMessageMedia[]): UserMessageProjection {
|
|
||||||
const imageNames: Array<string | undefined> = []
|
|
||||||
const attachmentNames: string[] = []
|
|
||||||
for (const item of media) {
|
|
||||||
// Outbound TUI media has no explicit kind because this path currently only
|
|
||||||
// sends clipboard images. Gateway and history media carry the kind.
|
|
||||||
if (item.kind === undefined || item.kind === "image") imageNames.push(item.name)
|
|
||||||
else if (item.name) attachmentNames.push(item.name)
|
|
||||||
}
|
|
||||||
|
|
||||||
const used = new Set<number>()
|
|
||||||
let next = 1
|
|
||||||
const imageLabels = imageNames.map((name) => {
|
|
||||||
const match = name?.match(/^clipboard-image-(\d+)\.[^.]+$/iu)
|
|
||||||
const preferred = match ? Number(match[1]) : 0
|
|
||||||
let index = Number.isSafeInteger(preferred) && preferred > 0 && !used.has(preferred)
|
|
||||||
? preferred
|
|
||||||
: next
|
|
||||||
while (used.has(index)) index += 1
|
|
||||||
used.add(index)
|
|
||||||
while (used.has(next)) next += 1
|
|
||||||
return `[Image #${index}]`
|
|
||||||
})
|
|
||||||
return { imageLabels, attachmentNames }
|
|
||||||
}
|
|
||||||
|
|
||||||
export function userMessageText(
|
|
||||||
content: string,
|
|
||||||
media: readonly UserMessageMedia[] = [],
|
|
||||||
displayContent?: string,
|
|
||||||
): string {
|
|
||||||
const { imageLabels, attachmentNames } = projectUserMessage(media)
|
|
||||||
return [
|
|
||||||
displayContent ?? [content, imageLabels.join(" ")].filter(Boolean).join(" "),
|
|
||||||
attachmentNames.length ? `Attachments: ${attachmentNames.join(", ")}` : "",
|
|
||||||
].filter(Boolean).join("\n")
|
|
||||||
}
|
|
||||||
|
|
||||||
/** Projects gateway events into retained, reflowable conversation cells. */
|
/** Projects gateway events into retained, reflowable conversation cells. */
|
||||||
export class Transcript {
|
export class Transcript {
|
||||||
readonly root: ScrollBoxRenderable
|
readonly root: ScrollBoxRenderable
|
||||||
@@ -126,12 +71,6 @@ export class Transcript {
|
|||||||
private readonly activities = new Set<Activity>()
|
private readonly activities = new Set<Activity>()
|
||||||
private readonly frames = new Set<BoxRenderable>()
|
private readonly frames = new Set<BoxRenderable>()
|
||||||
private readonly userRows = new Set<BoxRenderable>()
|
private readonly userRows = new Set<BoxRenderable>()
|
||||||
private readonly userMessages = new Set<{
|
|
||||||
renderable: TextRenderable
|
|
||||||
content: string
|
|
||||||
media: UserMessageMedia[]
|
|
||||||
displayContent?: string
|
|
||||||
}>()
|
|
||||||
private readonly userTurnIds = new Set<string>()
|
private readonly userTurnIds = new Set<string>()
|
||||||
private wrote = false
|
private wrote = false
|
||||||
private nextId = 0
|
private nextId = 0
|
||||||
@@ -177,13 +116,6 @@ export class Transcript {
|
|||||||
const previousSyntax = this.theme.syntax
|
const previousSyntax = this.theme.syntax
|
||||||
this.theme = theme
|
this.theme = theme
|
||||||
for (const { renderable, tone } of this.styledText) renderable.fg = theme[tone]
|
for (const { renderable, tone } of this.styledText) renderable.fg = theme[tone]
|
||||||
for (const message of this.userMessages) {
|
|
||||||
message.renderable.content = this.userMessageContent(
|
|
||||||
message.content,
|
|
||||||
message.media,
|
|
||||||
message.displayContent,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
for (const renderable of this.markdown) renderable.syntaxStyle = theme.syntax
|
for (const renderable of this.markdown) renderable.syntaxStyle = theme.syntax
|
||||||
for (const frame of this.frames) frame.borderColor = theme.border
|
for (const frame of this.frames) frame.borderColor = theme.border
|
||||||
for (const row of this.userRows) {
|
for (const row of this.userRows) {
|
||||||
@@ -239,7 +171,6 @@ export class Transcript {
|
|||||||
this.activities.clear()
|
this.activities.clear()
|
||||||
this.frames.clear()
|
this.frames.clear()
|
||||||
this.userRows.clear()
|
this.userRows.clear()
|
||||||
this.userMessages.clear()
|
|
||||||
this.userTurnIds.clear()
|
this.userTurnIds.clear()
|
||||||
this.wrote = false
|
this.wrote = false
|
||||||
this.nextId = 0
|
this.nextId = 0
|
||||||
@@ -251,9 +182,7 @@ export class Transcript {
|
|||||||
|
|
||||||
history(messages: HistoryMessage[]): void {
|
history(messages: HistoryMessage[]): void {
|
||||||
for (const message of messages) {
|
for (const message of messages) {
|
||||||
if (message.role === "user") {
|
if (message.role === "user") this.user(message.content, message.turnId)
|
||||||
this.user(message.content, message.turnId, message.media)
|
|
||||||
}
|
|
||||||
else if (message.role === "assistant") this.assistant(message.content)
|
else if (message.role === "assistant") this.assistant(message.content)
|
||||||
else if (message.fileEdits?.length) this.fileEdits(message.fileEdits)
|
else if (message.fileEdits?.length) this.fileEdits(message.fileEdits)
|
||||||
else this.progress(message.content, message.toolEvents)
|
else this.progress(message.content, message.toolEvents)
|
||||||
@@ -269,7 +198,7 @@ export class Transcript {
|
|||||||
for (const message of messages) {
|
for (const message of messages) {
|
||||||
if (message.role === "user") {
|
if (message.role === "user") {
|
||||||
if (message.turnId && this.userTurnIds.has(message.turnId)) continue
|
if (message.turnId && this.userTurnIds.has(message.turnId)) continue
|
||||||
this.writeUser(message.content, message.media, index++)
|
this.writeRole("›", message.content, "user", index++)
|
||||||
if (message.turnId) this.userTurnIds.add(message.turnId)
|
if (message.turnId) this.userTurnIds.add(message.turnId)
|
||||||
} else if (message.role === "assistant") {
|
} else if (message.role === "assistant") {
|
||||||
this.writeMarkdown(message.content, false, index++)
|
this.writeMarkdown(message.content, false, index++)
|
||||||
@@ -295,16 +224,11 @@ export class Transcript {
|
|||||||
return this.root.scrollTop <= 0
|
return this.root.scrollTop <= 0
|
||||||
}
|
}
|
||||||
|
|
||||||
user(
|
user(content: string, turnId?: string): boolean {
|
||||||
content: string,
|
|
||||||
turnId?: string,
|
|
||||||
media: readonly UserMessageMedia[] = [],
|
|
||||||
displayContent?: string,
|
|
||||||
): boolean {
|
|
||||||
if (turnId && this.userTurnIds.has(turnId)) return false
|
if (turnId && this.userTurnIds.has(turnId)) return false
|
||||||
this.noteOutput()
|
this.noteOutput()
|
||||||
this.finishActivity()
|
this.finishActivity()
|
||||||
this.writeUser(content, media, undefined, displayContent)
|
this.writeRole("›", content, "user")
|
||||||
if (turnId) this.userTurnIds.add(turnId)
|
if (turnId) this.userTurnIds.add(turnId)
|
||||||
return true
|
return true
|
||||||
}
|
}
|
||||||
@@ -412,7 +336,6 @@ export class Transcript {
|
|||||||
this.activity = null
|
this.activity = null
|
||||||
this.frames.clear()
|
this.frames.clear()
|
||||||
this.userRows.clear()
|
this.userRows.clear()
|
||||||
this.userMessages.clear()
|
|
||||||
this.theme.syntax.destroy()
|
this.theme.syntax.destroy()
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -545,7 +468,7 @@ export class Transcript {
|
|||||||
}
|
}
|
||||||
|
|
||||||
private createText(
|
private createText(
|
||||||
content: string | StyledText,
|
content: string,
|
||||||
tone: "text" | "muted" | "error" | "user",
|
tone: "text" | "muted" | "error" | "user",
|
||||||
bold = false,
|
bold = false,
|
||||||
id = "text",
|
id = "text",
|
||||||
@@ -564,10 +487,10 @@ export class Transcript {
|
|||||||
|
|
||||||
private writeRole(
|
private writeRole(
|
||||||
marker: string,
|
marker: string,
|
||||||
content: string | StyledText,
|
content: string,
|
||||||
tone: "muted" | "error" | "user",
|
tone: "muted" | "error" | "user",
|
||||||
index?: number,
|
index?: number,
|
||||||
): TextRenderable {
|
): void {
|
||||||
const row = this.createRow(tone === "user" ? "user" : "notice", "row")
|
const row = this.createRow(tone === "user" ? "user" : "notice", "row")
|
||||||
if (tone === "user") {
|
if (tone === "user") {
|
||||||
row.backgroundColor = this.theme.userBackground
|
row.backgroundColor = this.theme.userBackground
|
||||||
@@ -586,76 +509,6 @@ export class Transcript {
|
|||||||
row.add(text)
|
row.add(text)
|
||||||
this.root.add(row, index)
|
this.root.add(row, index)
|
||||||
this.wrote = true
|
this.wrote = true
|
||||||
return text
|
|
||||||
}
|
|
||||||
|
|
||||||
private writeUser(
|
|
||||||
content: string,
|
|
||||||
media: readonly UserMessageMedia[] = [],
|
|
||||||
index?: number,
|
|
||||||
displayContent?: string,
|
|
||||||
): void {
|
|
||||||
const retainedMedia = [...media]
|
|
||||||
const renderable = this.writeRole(
|
|
||||||
"›",
|
|
||||||
this.userMessageContent(content, retainedMedia, displayContent),
|
|
||||||
"user",
|
|
||||||
index,
|
|
||||||
)
|
|
||||||
this.userMessages.add({ renderable, content, media: retainedMedia, displayContent })
|
|
||||||
}
|
|
||||||
|
|
||||||
private userMessageContent(
|
|
||||||
content: string,
|
|
||||||
media: readonly UserMessageMedia[],
|
|
||||||
displayContent?: string,
|
|
||||||
): StyledText {
|
|
||||||
const { imageLabels, attachmentNames } = projectUserMessage(media)
|
|
||||||
const chunks: TextChunk[] = []
|
|
||||||
const append = (text: string) => {
|
|
||||||
if (text) chunks.push({ __isChunk: true, text })
|
|
||||||
}
|
|
||||||
const nextLine = () => {
|
|
||||||
if (chunks.length) append("\n")
|
|
||||||
}
|
|
||||||
|
|
||||||
if (displayContent !== undefined) {
|
|
||||||
const ranges = imageLabels
|
|
||||||
.map((label) => ({ label, start: displayContent.indexOf(label) }))
|
|
||||||
.filter(({ start }) => start >= 0)
|
|
||||||
.sort((left, right) => left.start - right.start)
|
|
||||||
let cursor = 0
|
|
||||||
for (const { label, start } of ranges) {
|
|
||||||
append(displayContent.slice(cursor, start))
|
|
||||||
chunks.push({
|
|
||||||
__isChunk: true,
|
|
||||||
text: label,
|
|
||||||
fg: RGBA.fromHex(this.theme.user),
|
|
||||||
attributes: TextAttributes.BOLD,
|
|
||||||
})
|
|
||||||
cursor = start + label.length
|
|
||||||
}
|
|
||||||
append(displayContent.slice(cursor))
|
|
||||||
} else {
|
|
||||||
append(content)
|
|
||||||
}
|
|
||||||
if (displayContent === undefined && imageLabels.length) {
|
|
||||||
if (chunks.length) append(" ")
|
|
||||||
for (const [index, label] of imageLabels.entries()) {
|
|
||||||
if (index > 0) append(" ")
|
|
||||||
chunks.push({
|
|
||||||
__isChunk: true,
|
|
||||||
text: label,
|
|
||||||
fg: RGBA.fromHex(this.theme.user),
|
|
||||||
attributes: TextAttributes.BOLD,
|
|
||||||
})
|
|
||||||
}
|
|
||||||
}
|
|
||||||
if (attachmentNames.length) {
|
|
||||||
nextLine()
|
|
||||||
append(`Attachments: ${attachmentNames.join(", ")}`)
|
|
||||||
}
|
|
||||||
return new StyledText(chunks)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
private createMarkdown(content: string, streaming: boolean, id = "markdown"): MarkdownRenderable {
|
private createMarkdown(content: string, streaming: boolean, id = "markdown"): MarkdownRenderable {
|
||||||
|
|||||||
Reference in New Issue
Block a user