gradio-forge / agent /memory.py
Sokheng's picture
update final
344419c
Raw
History Blame Contribute Delete
4.26 kB
"""
Conversation history management.
Backed by LangChain's InMemoryChatMessageHistory for structured message
storage. The agent state dict still uses plain {"role", "content"} dicts
(so the rest of the app stays unchanged), but summarization now uses
LangChain message types internally.
When the raw character count exceeds SUMMARIZE_THRESHOLD, the oldest
turns (beyond the last KEEP_RECENT pairs) are collapsed into a single
summary message so the context never explodes on long sessions.
"""
from __future__ import annotations
import logging
from langchain_core.chat_history import InMemoryChatMessageHistory
from langchain_core.messages import (
SystemMessage,
HumanMessage,
AIMessage,
BaseMessage,
)
logger = logging.getLogger(__name__)
SUMMARIZE_THRESHOLD = 24_000 # chars — rough proxy for ~8 k tokens
KEEP_RECENT = 4 # verbatim turn-pairs to keep after summarization
def _to_lc_message(msg: dict) -> BaseMessage:
"""Convert a {"role": ..., "content": ...} dict to a LangChain message."""
role = msg.get("role", "user")
content = msg.get("content", "")
if role == "system":
return SystemMessage(content=content)
if role == "assistant":
return AIMessage(content=content)
return HumanMessage(content=content)
def _from_lc_message(msg: BaseMessage) -> dict:
"""Convert a LangChain message back to a plain dict."""
if isinstance(msg, SystemMessage):
return {"role": "system", "content": msg.content}
if isinstance(msg, AIMessage):
return {"role": "assistant", "content": msg.content}
return {"role": "user", "content": msg.content}
def _history_from_messages(messages: list[dict]) -> InMemoryChatMessageHistory:
"""Build a LangChain chat history object from a plain messages list."""
history = InMemoryChatMessageHistory()
for m in messages:
history.add_message(_to_lc_message(m))
return history
def make_state(system_prompt: str) -> dict:
"""Return a fresh agent state dict."""
return {
"messages": [{"role": "system", "content": system_prompt}],
"current_spec": None,
"current_code": "",
"iteration": 0,
"heal_attempts": 0,
"last_error": None,
}
def append_user(state: dict, content: str) -> dict:
"""Return new state with user message appended."""
return {**state, "messages": state["messages"] + [{"role": "user", "content": content}]}
def append_assistant(state: dict, content: str) -> dict:
"""Return new state with assistant message appended."""
return {**state, "messages": state["messages"] + [{"role": "assistant", "content": content}]}
def total_chars(messages: list[dict]) -> int:
return sum(len(m.get("content", "")) for m in messages)
def get_lc_history(state: dict) -> InMemoryChatMessageHistory:
"""
Return a LangChain InMemoryChatMessageHistory built from state["messages"].
Useful for passing to LangChain chains that expect a BaseChatMessageHistory.
"""
return _history_from_messages(state["messages"])
def maybe_summarize(state: dict, summarize_fn) -> dict:
"""
If total chars exceed SUMMARIZE_THRESHOLD, summarize the oldest turns
using summarize_fn, then replace them with a single summary message.
summarize_fn(messages_to_summarize: list[dict]) -> str
"""
messages = state["messages"]
if total_chars(messages) <= SUMMARIZE_THRESHOLD:
return state
# Always keep system [0] + last KEEP_RECENT*2 messages verbatim
cutoff = max(1, len(messages) - KEEP_RECENT * 2)
old = messages[1:cutoff] # skip system message
keep = messages[cutoff:]
if not old:
return state
logger.info("Summarizing %d old messages via LangChain history", len(old))
# Build a temporary LangChain history just for the messages being summarized
temp_history = _history_from_messages(old)
summary_text = summarize_fn(
[_from_lc_message(m) for m in temp_history.messages]
)
new_messages = (
[messages[0]]
+ [{"role": "system", "content": f"[Earlier conversation summary]\n{summary_text}"}]
+ keep
)
return {**state, "messages": new_messages}