mirror of
https://github.com/NCBM/plyngent.git
synced 2026-07-23 05:55:16 +08:00
core/agent+cli: wire system prompt and agent budgets
ChatAgent accepts system_prompt, result budget, and parallel_tools; CLI loads them from config [agent].
This commit is contained in:
@@ -2,8 +2,9 @@ from __future__ import annotations
|
|||||||
|
|
||||||
from typing import TYPE_CHECKING
|
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
|
from .loop import DEFAULT_MAX_ROUNDS, run_chat_loop
|
||||||
|
|
||||||
if TYPE_CHECKING:
|
if TYPE_CHECKING:
|
||||||
@@ -31,6 +32,9 @@ class ChatAgent:
|
|||||||
temperature: float | None
|
temperature: float | None
|
||||||
on_limit: LimitContinueHook | None
|
on_limit: LimitContinueHook | None
|
||||||
stream: bool
|
stream: bool
|
||||||
|
system_prompt: str | None
|
||||||
|
max_tool_result_chars: int
|
||||||
|
parallel_tools: bool
|
||||||
messages: list[AnyChatMessage]
|
messages: list[AnyChatMessage]
|
||||||
pending_retry_text: str | None
|
pending_retry_text: str | None
|
||||||
|
|
||||||
@@ -47,6 +51,9 @@ class ChatAgent:
|
|||||||
messages: Sequence[AnyChatMessage] | None = None,
|
messages: Sequence[AnyChatMessage] | None = None,
|
||||||
on_limit: LimitContinueHook | None = None,
|
on_limit: LimitContinueHook | None = None,
|
||||||
stream: bool = True,
|
stream: bool = True,
|
||||||
|
system_prompt: str | None = None,
|
||||||
|
max_tool_result_chars: int = DEFAULT_TOOL_RESULT_MAX_CHARS,
|
||||||
|
parallel_tools: bool = True,
|
||||||
) -> None:
|
) -> None:
|
||||||
self.client = client
|
self.client = client
|
||||||
self.model = model
|
self.model = model
|
||||||
@@ -57,8 +64,20 @@ class ChatAgent:
|
|||||||
self.temperature = temperature
|
self.temperature = temperature
|
||||||
self.on_limit = on_limit
|
self.on_limit = on_limit
|
||||||
self.stream = stream
|
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.messages = list(messages) if messages is not None else []
|
||||||
self.pending_retry_text = None
|
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:
|
async def load_history(self) -> None:
|
||||||
"""Replace in-memory messages from the bound memory session."""
|
"""Replace in-memory messages from the bound memory session."""
|
||||||
@@ -67,6 +86,7 @@ class ChatAgent:
|
|||||||
raise RuntimeError(msg)
|
raise RuntimeError(msg)
|
||||||
self.messages = await self.memory.list_messages(self.session_id)
|
self.messages = await self.memory.list_messages(self.session_id)
|
||||||
self.pending_retry_text = None
|
self.pending_retry_text = None
|
||||||
|
self._ensure_system_prompt()
|
||||||
|
|
||||||
async def bind_session(self, session_id: int, *, load: bool = True) -> None:
|
async def bind_session(self, session_id: int, *, load: bool = True) -> None:
|
||||||
"""Attach a memory session id; optionally load existing messages."""
|
"""Attach a memory session id; optionally load existing messages."""
|
||||||
@@ -103,11 +123,12 @@ class ChatAgent:
|
|||||||
temperature=self.temperature,
|
temperature=self.temperature,
|
||||||
on_limit=self.on_limit,
|
on_limit=self.on_limit,
|
||||||
stream=self.stream,
|
stream=self.stream,
|
||||||
|
max_tool_result_chars=self.max_tool_result_chars,
|
||||||
|
parallel_tools=self.parallel_tools,
|
||||||
):
|
):
|
||||||
yield event
|
yield event
|
||||||
completed = True
|
completed = True
|
||||||
except BaseException:
|
except BaseException:
|
||||||
# Includes CancelledError / KeyboardInterrupt paths via task cancel.
|
|
||||||
if not completed:
|
if not completed:
|
||||||
self._rollback_turn(pre_len, user_msg.content)
|
self._rollback_turn(pre_len, user_msg.content)
|
||||||
raise
|
raise
|
||||||
@@ -118,6 +139,7 @@ class ChatAgent:
|
|||||||
|
|
||||||
async def run(self, user_text: str) -> AsyncIterator[AgentEvent]:
|
async def run(self, user_text: str) -> AsyncIterator[AgentEvent]:
|
||||||
"""Append a user message, run the tool loop, yield events, persist only on success."""
|
"""Append a user message, run the tool loop, yield events, persist only on success."""
|
||||||
|
self._ensure_system_prompt()
|
||||||
user_msg = UserChatMessage(content=user_text)
|
user_msg = UserChatMessage(content=user_text)
|
||||||
self.messages.append(user_msg)
|
self.messages.append(user_msg)
|
||||||
async for event in self._run_from_user_message(user_msg):
|
async for event in self._run_from_user_message(user_msg):
|
||||||
@@ -129,6 +151,7 @@ class ChatAgent:
|
|||||||
if text is None:
|
if text is None:
|
||||||
msg = "nothing to retry"
|
msg = "nothing to retry"
|
||||||
raise RuntimeError(msg)
|
raise RuntimeError(msg)
|
||||||
|
self._ensure_system_prompt()
|
||||||
user_msg = UserChatMessage(content=text)
|
user_msg = UserChatMessage(content=text)
|
||||||
self.messages.append(user_msg)
|
self.messages.append(user_msg)
|
||||||
async for event in self._run_from_user_message(user_msg):
|
async for event in self._run_from_user_message(user_msg):
|
||||||
|
|||||||
@@ -22,6 +22,11 @@ _MINIMAL_CONFIG = """\
|
|||||||
# implementation = "sqlite"
|
# implementation = "sqlite"
|
||||||
# url = "/path/to/chat.db"
|
# url = "/path/to/chat.db"
|
||||||
|
|
||||||
|
# [agent]
|
||||||
|
# system_prompt = "You are a careful coding assistant."
|
||||||
|
# max_tool_result_chars = 32000
|
||||||
|
# parallel_tools = true
|
||||||
|
|
||||||
# [providers.example]
|
# [providers.example]
|
||||||
# preset = "openai-compatible"
|
# preset = "openai-compatible"
|
||||||
# url = "https://api.openai.com/v1"
|
# url = "https://api.openai.com/v1"
|
||||||
|
|||||||
@@ -45,6 +45,8 @@ class ReplState:
|
|||||||
def _make_agent(self) -> ChatAgent:
|
def _make_agent(self) -> ChatAgent:
|
||||||
from plyngent.cli.limits import prompt_continue_limit
|
from plyngent.cli.limits import prompt_continue_limit
|
||||||
|
|
||||||
|
agent_cfg = self.config.agent_config
|
||||||
|
system_prompt = agent_cfg.system_prompt or None
|
||||||
return ChatAgent(
|
return ChatAgent(
|
||||||
self.client,
|
self.client,
|
||||||
model=self.model,
|
model=self.model,
|
||||||
@@ -54,6 +56,9 @@ class ReplState:
|
|||||||
max_rounds=self.max_rounds,
|
max_rounds=self.max_rounds,
|
||||||
on_limit=prompt_continue_limit,
|
on_limit=prompt_continue_limit,
|
||||||
stream=True,
|
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:
|
def rebuild_client(self) -> None:
|
||||||
|
|||||||
@@ -391,3 +391,57 @@ async def test_chat_agent_retry_after_failure() -> None:
|
|||||||
assert isinstance(loaded[0], UserChatMessage)
|
assert isinstance(loaded[0], UserChatMessage)
|
||||||
assert loaded[0].content == "ping"
|
assert loaded[0].content == "ping"
|
||||||
await store.close()
|
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
|
||||||
|
|||||||
Reference in New Issue
Block a user