diff --git a/src/plyngent/agent/chat.py b/src/plyngent/agent/chat.py index 1ffe532..02571bd 100644 --- a/src/plyngent/agent/chat.py +++ b/src/plyngent/agent/chat.py @@ -2,8 +2,9 @@ from __future__ import annotations from typing import TYPE_CHECKING -from plyngent.lmproto.openai_compatible.model import UserChatMessage +from plyngent.lmproto.openai_compatible.model import SystemChatMessage, UserChatMessage +from .budget import DEFAULT_TOOL_RESULT_MAX_CHARS from .loop import DEFAULT_MAX_ROUNDS, run_chat_loop if TYPE_CHECKING: @@ -31,6 +32,9 @@ class ChatAgent: temperature: float | None on_limit: LimitContinueHook | None stream: bool + system_prompt: str | None + max_tool_result_chars: int + parallel_tools: bool messages: list[AnyChatMessage] pending_retry_text: str | None @@ -47,6 +51,9 @@ class ChatAgent: messages: Sequence[AnyChatMessage] | None = None, on_limit: LimitContinueHook | None = None, stream: bool = True, + system_prompt: str | None = None, + max_tool_result_chars: int = DEFAULT_TOOL_RESULT_MAX_CHARS, + parallel_tools: bool = True, ) -> None: self.client = client self.model = model @@ -57,8 +64,20 @@ class ChatAgent: self.temperature = temperature self.on_limit = on_limit self.stream = stream + self.system_prompt = system_prompt + self.max_tool_result_chars = max_tool_result_chars + self.parallel_tools = parallel_tools self.messages = list(messages) if messages is not None else [] self.pending_retry_text = None + self._ensure_system_prompt() + + def _ensure_system_prompt(self) -> None: + """Prepend system prompt once when configured and history has none.""" + if not self.system_prompt: + return + if self.messages and isinstance(self.messages[0], SystemChatMessage): + return + self.messages.insert(0, SystemChatMessage(content=self.system_prompt)) async def load_history(self) -> None: """Replace in-memory messages from the bound memory session.""" @@ -67,6 +86,7 @@ class ChatAgent: raise RuntimeError(msg) self.messages = await self.memory.list_messages(self.session_id) self.pending_retry_text = None + self._ensure_system_prompt() async def bind_session(self, session_id: int, *, load: bool = True) -> None: """Attach a memory session id; optionally load existing messages.""" @@ -103,11 +123,12 @@ class ChatAgent: temperature=self.temperature, on_limit=self.on_limit, stream=self.stream, + max_tool_result_chars=self.max_tool_result_chars, + parallel_tools=self.parallel_tools, ): yield event completed = True except BaseException: - # Includes CancelledError / KeyboardInterrupt paths via task cancel. if not completed: self._rollback_turn(pre_len, user_msg.content) raise @@ -118,6 +139,7 @@ class ChatAgent: async def run(self, user_text: str) -> AsyncIterator[AgentEvent]: """Append a user message, run the tool loop, yield events, persist only on success.""" + self._ensure_system_prompt() user_msg = UserChatMessage(content=user_text) self.messages.append(user_msg) async for event in self._run_from_user_message(user_msg): @@ -129,6 +151,7 @@ class ChatAgent: if text is None: msg = "nothing to retry" raise RuntimeError(msg) + self._ensure_system_prompt() user_msg = UserChatMessage(content=text) self.messages.append(user_msg) async for event in self._run_from_user_message(user_msg): diff --git a/src/plyngent/cli/editor.py b/src/plyngent/cli/editor.py index 1bb1879..97692aa 100644 --- a/src/plyngent/cli/editor.py +++ b/src/plyngent/cli/editor.py @@ -22,6 +22,11 @@ _MINIMAL_CONFIG = """\ # implementation = "sqlite" # url = "/path/to/chat.db" +# [agent] +# system_prompt = "You are a careful coding assistant." +# max_tool_result_chars = 32000 +# parallel_tools = true + # [providers.example] # preset = "openai-compatible" # url = "https://api.openai.com/v1" diff --git a/src/plyngent/cli/state.py b/src/plyngent/cli/state.py index 4660edb..489b843 100644 --- a/src/plyngent/cli/state.py +++ b/src/plyngent/cli/state.py @@ -45,6 +45,8 @@ class ReplState: def _make_agent(self) -> ChatAgent: from plyngent.cli.limits import prompt_continue_limit + agent_cfg = self.config.agent_config + system_prompt = agent_cfg.system_prompt or None return ChatAgent( self.client, model=self.model, @@ -54,6 +56,9 @@ class ReplState: max_rounds=self.max_rounds, on_limit=prompt_continue_limit, stream=True, + system_prompt=system_prompt, + max_tool_result_chars=agent_cfg.max_tool_result_chars, + parallel_tools=agent_cfg.parallel_tools, ) def rebuild_client(self) -> None: diff --git a/tests/test_agent/test_loop.py b/tests/test_agent/test_loop.py index c6d4d85..43451b3 100644 --- a/tests/test_agent/test_loop.py +++ b/tests/test_agent/test_loop.py @@ -391,3 +391,57 @@ async def test_chat_agent_retry_after_failure() -> None: assert isinstance(loaded[0], UserChatMessage) assert loaded[0].content == "ping" await store.close() + + +async def test_chat_agent_system_prompt_prepended() -> None: + from plyngent.lmproto.openai_compatible.model import SystemChatMessage + + client = ScriptedClient([_response(AssistantChatMessage(content="ok"))]) + agent = ChatAgent(client, model="m", system_prompt="Be brief.", stream=False) + _ = [e async for e in agent.run("hi")] + assert isinstance(agent.messages[0], SystemChatMessage) + assert agent.messages[0].content == "Be brief." + assert isinstance(client.calls[0].messages[0], SystemChatMessage) + + +async def test_tool_result_char_budget() -> None: + @tool + def big() -> str: + return "x" * 100 + + registry = ToolRegistry([big]) + client = ScriptedClient( + [ + _response( + AssistantChatMessage( + content="", + tool_calls=[ + AssistantFunctionToolCall( + id="1", + function=AssistantFunctionTool(name="big", arguments="{}"), + ) + ], + ) + ), + _response(AssistantChatMessage(content="done")), + ] + ) + messages: list[AnyChatMessage] = [UserChatMessage(content="go")] + _ = [ + e + async for e in run_chat_loop( + client, + messages, + model="m", + tools=registry, + stream=False, + max_tool_result_chars=20, + parallel_tools=False, + ) + ] + from plyngent.lmproto.openai_compatible.model import ToolChatMessage + + tool_msgs = [m for m in messages if isinstance(m, ToolChatMessage)] + assert len(tool_msgs) == 1 + assert tool_msgs[0].content.startswith("x" * 20) + assert "truncated" in tool_msgs[0].content