diff --git a/nanobot/agent/loop.py b/nanobot/agent/loop.py index 2e6185e7..91deccba 100644 --- a/nanobot/agent/loop.py +++ b/nanobot/agent/loop.py @@ -202,6 +202,7 @@ class AgentLoop: timezone: str | None = None, session_ttl_minutes: int = 0, consolidation_ratio: float = 0.5, + max_messages: int = 0, hooks: list[AgentHook] | None = None, unified_session: bool = False, disabled_skills: list[str] | None = None, @@ -259,6 +260,7 @@ class AgentLoop: disabled_skills=disabled_skills, ) self._unified_session = unified_session + self._max_messages = max_messages if max_messages > 0 else 0 self._running = False self._mcp_servers = mcp_servers or {} self._mcp_stacks: dict[str, AsyncExitStack] = {} @@ -884,10 +886,13 @@ class AgentLoop: channel, chat_id, msg.metadata.get("message_id"), msg.metadata, session_key=key, ) - history = session.get_history( - max_tokens=self._replay_token_budget(), - include_timestamps=True, - ) + _hist_kwargs: dict[str, Any] = { + "max_tokens": self._replay_token_budget(), + "include_timestamps": True, + } + if self._max_messages > 0: + _hist_kwargs["max_messages"] = self._max_messages + history = session.get_history(**_hist_kwargs) current_role = "assistant" if is_subagent else "user" # Subagent content is already in `history` above; passing it again @@ -971,10 +976,13 @@ class AgentLoop: if isinstance(message_tool, MessageTool): message_tool.start_turn() - history = session.get_history( - max_tokens=self._replay_token_budget(), - include_timestamps=True, - ) + _hist_kwargs: dict[str, Any] = { + "max_tokens": self._replay_token_budget(), + "include_timestamps": True, + } + if self._max_messages > 0: + _hist_kwargs["max_messages"] = self._max_messages + history = session.get_history(**_hist_kwargs) pending_ask_id = pending_ask_user_id(history) if pending_ask_id: diff --git a/nanobot/cli/commands.py b/nanobot/cli/commands.py index 62b698ce..2fe39746 100644 --- a/nanobot/cli/commands.py +++ b/nanobot/cli/commands.py @@ -538,6 +538,7 @@ def serve( disabled_skills=runtime_config.agents.defaults.disabled_skills, session_ttl_minutes=runtime_config.agents.defaults.session_ttl_minutes, consolidation_ratio=runtime_config.agents.defaults.consolidation_ratio, + max_messages=runtime_config.agents.defaults.max_messages, tools_config=runtime_config.tools, ) @@ -651,6 +652,7 @@ def _run_gateway( disabled_skills=config.agents.defaults.disabled_skills, session_ttl_minutes=config.agents.defaults.session_ttl_minutes, consolidation_ratio=config.agents.defaults.consolidation_ratio, + max_messages=config.agents.defaults.max_messages, tools_config=config.tools, provider_snapshot_loader=load_provider_snapshot, provider_signature=provider_snapshot.signature, @@ -1043,6 +1045,7 @@ def agent( disabled_skills=config.agents.defaults.disabled_skills, session_ttl_minutes=config.agents.defaults.session_ttl_minutes, consolidation_ratio=config.agents.defaults.consolidation_ratio, + max_messages=config.agents.defaults.max_messages, tools_config=config.tools, ) restart_notice = consume_restart_notice_from_env() diff --git a/nanobot/config/schema.py b/nanobot/config/schema.py index e1f91aeb..9eb8a464 100644 --- a/nanobot/config/schema.py +++ b/nanobot/config/schema.py @@ -90,6 +90,10 @@ class AgentDefaults(Base): validation_alias=AliasChoices("idleCompactAfterMinutes", "sessionTtlMinutes"), serialization_alias="idleCompactAfterMinutes", ) # Auto-compact idle threshold in minutes (0 = disabled) + max_messages: int = Field( + default=0, + ge=0, + ) # Max messages to replay from session history (0 = use default, respects token budget) consolidation_ratio: float = Field( default=0.5, ge=0.1, diff --git a/tests/agent/test_max_messages_config.py b/tests/agent/test_max_messages_config.py new file mode 100644 index 00000000..bf9f096b --- /dev/null +++ b/tests/agent/test_max_messages_config.py @@ -0,0 +1,153 @@ +"""Tests for max_messages config wiring into session history replay.""" + +from __future__ import annotations + +from pathlib import Path +from typing import Any +from unittest.mock import MagicMock, patch + +import pytest + +from nanobot.agent.loop import AgentLoop +from nanobot.bus.queue import MessageBus +from nanobot.session.manager import HISTORY_MAX_MESSAGES, Session + + +def _make_loop(tmp_path: Path, max_messages: int = 0) -> AgentLoop: + provider = MagicMock() + provider.get_default_model.return_value = "test-model" + return AgentLoop( + bus=MessageBus(), + provider=provider, + workspace=tmp_path, + model="test-model", + max_messages=max_messages, + ) + + +def _populated_session(n: int) -> Session: + """Create a session with *n* user/assistant turn pairs.""" + session = Session(key="test:populated") + for i in range(n): + session.add_message("user", f"msg-{i}") + session.add_message("assistant", f"reply-{i}") + return session + + +class TestMaxMessagesInit: + """Verify AgentLoop stores the config value correctly.""" + + def test_default_is_zero(self, tmp_path: Path) -> None: + loop = _make_loop(tmp_path) + assert loop._max_messages == 0 + + def test_positive_value_stored(self, tmp_path: Path) -> None: + loop = _make_loop(tmp_path, max_messages=25) + assert loop._max_messages == 25 + + def test_zero_means_unlimited(self, tmp_path: Path) -> None: + """max_messages=0 should not constrain get_history (uses default).""" + loop = _make_loop(tmp_path, max_messages=0) + assert loop._max_messages == 0 + + def test_negative_treated_as_zero(self, tmp_path: Path) -> None: + """Negative values should not produce negative slicing.""" + loop = _make_loop(tmp_path, max_messages=-5) + assert loop._max_messages == 0 + + +class TestGetHistoryWithMaxMessages: + """Verify get_history respects max_messages parameter.""" + + def test_default_uses_constant(self) -> None: + session = _populated_session(80) + history = session.get_history() + # Default HISTORY_MAX_MESSAGES=120, 80 pairs = 160 msgs, sliced to 120 + assert len(history) <= HISTORY_MAX_MESSAGES + + def test_explicit_max_messages_limits_output(self) -> None: + session = _populated_session(40) # 80 messages total + history = session.get_history(max_messages=20) + assert len(history) <= 20 + + def test_max_messages_starts_at_user_turn(self) -> None: + """Sliced history should start with a user message, not mid-turn.""" + session = _populated_session(30) # 60 messages + history = session.get_history(max_messages=25) + assert history[0]["role"] == "user" + + def test_max_messages_zero_returns_all(self) -> None: + """max_messages=0 with the default constant returns up to the constant.""" + session = _populated_session(10) # 20 messages + # When we pass 0 explicitly, unconsolidated[-0:] returns everything + # but the default is HISTORY_MAX_MESSAGES so this tests the default path + history = session.get_history() + assert len(history) == 20 + + def test_small_session_unaffected(self) -> None: + """When session has fewer messages than max_messages, all are returned.""" + session = _populated_session(5) # 10 messages + history = session.get_history(max_messages=25) + assert len(history) == 10 + + +class TestMaxMessagesIntegration: + """Verify the config flows from AgentLoop into get_history calls.""" + + def test_config_wired_to_history_call(self, tmp_path: Path) -> None: + """When max_messages > 0, get_history should receive it.""" + loop = _make_loop(tmp_path, max_messages=25) + session = _populated_session(40) # 80 messages + + with patch.object(session, "get_history", wraps=session.get_history) as mock_hist: + # Call the internal method that builds history kwargs + kwargs: dict[str, Any] = { + "max_tokens": loop._replay_token_budget(), + "include_timestamps": True, + } + if loop._max_messages > 0: + kwargs["max_messages"] = loop._max_messages + session.get_history(**kwargs) + + assert mock_hist.call_count == 1 + call_kwargs = mock_hist.call_args + # max_messages is positional arg (first) or keyword + if call_kwargs.args: + assert call_kwargs.args[0] == 25 + else: + assert call_kwargs.kwargs.get("max_messages") == 25 + + def test_zero_config_omits_max_messages_kwarg(self, tmp_path: Path) -> None: + """When max_messages=0, get_history should use its default.""" + loop = _make_loop(tmp_path, max_messages=0) + + kwargs: dict[str, Any] = { + "max_tokens": loop._replay_token_budget(), + "include_timestamps": True, + } + if loop._max_messages > 0: + kwargs["max_messages"] = loop._max_messages + + assert "max_messages" not in kwargs + + +class TestSchemaConfig: + """Verify the config schema accepts max_messages.""" + + def test_schema_default(self) -> None: + from nanobot.config.schema import AgentDefaults + + defaults = AgentDefaults() + assert defaults.max_messages == 0 + + def test_schema_accepts_positive(self) -> None: + from nanobot.config.schema import AgentDefaults + + defaults = AgentDefaults(max_messages=25) + assert defaults.max_messages == 25 + + def test_schema_rejects_negative(self) -> None: + from nanobot.config.schema import AgentDefaults + + with pytest.raises(Exception): # Pydantic validation error + AgentDefaults(max_messages=-1)