From 85fae1cb2f2a4d8ce0265e4f88d76f1dc52383b4 Mon Sep 17 00:00:00 2001 From: worldmozara Date: Tue, 14 Jul 2026 17:45:47 +0800 Subject: [PATCH] core/agent: ChatAgent wrapper with optional memory bind --- src/plyngent/__init__.py | 1 + src/plyngent/agent/__init__.py | 14 ++++++ src/plyngent/agent/chat.py | 92 ++++++++++++++++++++++++++++++++++ 3 files changed, 107 insertions(+) create mode 100644 src/plyngent/agent/__init__.py create mode 100644 src/plyngent/agent/chat.py diff --git a/src/plyngent/__init__.py b/src/plyngent/__init__.py index 65a5492..b89dcea 100644 --- a/src/plyngent/__init__.py +++ b/src/plyngent/__init__.py @@ -1,3 +1,4 @@ +from . import agent as agent from . import config as config from . import memory as memory from . import runtime as runtime diff --git a/src/plyngent/agent/__init__.py b/src/plyngent/agent/__init__.py new file mode 100644 index 0000000..3b6fd6d --- /dev/null +++ b/src/plyngent/agent/__init__.py @@ -0,0 +1,14 @@ +from .chat import ChatAgent as ChatAgent +from .client import ChatClient as ChatClient +from .events import AgentEvent as AgentEvent +from .events import AssistantMessageEvent as AssistantMessageEvent +from .events import MaxRoundsEvent as MaxRoundsEvent +from .events import TextDeltaEvent as TextDeltaEvent +from .events import ToolCallEvent as ToolCallEvent +from .events import ToolResultEvent as ToolResultEvent +from .loop import DEFAULT_MAX_ROUNDS as DEFAULT_MAX_ROUNDS +from .loop import run_chat_loop as run_chat_loop +from .tools import ToolDefinition as ToolDefinition +from .tools import ToolRegistry as ToolRegistry +from .tools import schema_from_callable as schema_from_callable +from .tools import tool as tool diff --git a/src/plyngent/agent/chat.py b/src/plyngent/agent/chat.py new file mode 100644 index 0000000..2c5a2c5 --- /dev/null +++ b/src/plyngent/agent/chat.py @@ -0,0 +1,92 @@ +from __future__ import annotations + +from typing import TYPE_CHECKING + +from plyngent.lmproto.openai_compatible.model import UserChatMessage + +from .loop import DEFAULT_MAX_ROUNDS, run_chat_loop + +if TYPE_CHECKING: + from collections.abc import AsyncIterator, Sequence + + from plyngent.lmproto.openai_compatible.model import AnyChatMessage + from plyngent.memory import MemoryStore + + from .client import ChatClient + from .events import AgentEvent + from .tools import ToolRegistry + + +class ChatAgent: + """Thin wrapper: chat client + optional tools + optional memory bind.""" + + client: ChatClient + model: str + tools: ToolRegistry | None + memory: MemoryStore | None + session_id: int | None + max_rounds: int + temperature: float | None + messages: list[AnyChatMessage] + + def __init__( # noqa: PLR0913 + self, + client: ChatClient, + *, + model: str, + tools: ToolRegistry | None = None, + memory: MemoryStore | None = None, + session_id: int | None = None, + max_rounds: int = DEFAULT_MAX_ROUNDS, + temperature: float | None = None, + messages: Sequence[AnyChatMessage] | None = None, + ) -> None: + self.client = client + self.model = model + self.tools = tools + self.memory = memory + self.session_id = session_id + self.max_rounds = max_rounds + self.temperature = temperature + self.messages = list(messages) if messages is not None else [] + + async def load_history(self) -> None: + """Replace in-memory messages from the bound memory session.""" + if self.memory is None or self.session_id is None: + msg = "load_history requires memory and session_id" + raise RuntimeError(msg) + self.messages = await self.memory.list_messages(self.session_id) + + async def bind_session(self, session_id: int, *, load: bool = True) -> None: + """Attach a memory session id; optionally load existing messages.""" + if self.memory is None: + msg = "bind_session requires a MemoryStore" + raise RuntimeError(msg) + self.session_id = session_id + if load: + await self.load_history() + + async def _persist(self, message: AnyChatMessage) -> None: + if self.memory is not None and self.session_id is not None: + _ = await self.memory.append_message(self.session_id, message) + + async def run(self, user_text: str) -> AsyncIterator[AgentEvent]: + """Append a user message, run the tool loop, yield events, persist new messages.""" + user_msg = UserChatMessage(content=user_text) + self.messages.append(user_msg) + await self._persist(user_msg) + + start_len = len(self.messages) + async for event in run_chat_loop( + self.client, + self.messages, + model=self.model, + tools=self.tools, + max_rounds=self.max_rounds, + temperature=self.temperature, + ): + yield event + + # Persist messages appended by the loop after the user message. + for message in self.messages[start_len:]: + await self._persist(message)