feat: initial release - Hunzi agent framework v0.1.0
This commit is contained in:
@@ -0,0 +1,231 @@
|
||||
"""Memory system — conversation history and long-term storage."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from abc import ABC, abstractmethod
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Message model
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class Message:
|
||||
"""A single message in structured format.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
role : str
|
||||
``"user"``, ``"assistant"``, ``"system"``, or ``"tool"``.
|
||||
content : str
|
||||
The message body.
|
||||
tool_call_id : str, optional
|
||||
ID of the tool call this message responds to (for ``role="tool"``).
|
||||
tool_calls : list, optional
|
||||
List of tool-call dicts (for ``role="assistant"``).
|
||||
metadata : dict, optional
|
||||
Extra bookkeeping data (timestamp, etc.).
|
||||
"""
|
||||
|
||||
role: str
|
||||
content: str
|
||||
tool_call_id: Optional[str] = None
|
||||
tool_calls: Optional[List[Dict[str, Any]]] = None
|
||||
metadata: Dict[str, Any] = field(default_factory=dict)
|
||||
|
||||
def to_openai(self) -> Dict[str, Any]:
|
||||
"""Convert to an OpenAI-compatible message dict."""
|
||||
msg: Dict[str, Any] = {"role": self.role, "content": self.content}
|
||||
if self.tool_call_id:
|
||||
msg["tool_call_id"] = self.tool_call_id
|
||||
if self.tool_calls:
|
||||
msg["tool_calls"] = self.tool_calls
|
||||
return msg
|
||||
|
||||
@classmethod
|
||||
def from_openai(cls, data: Dict[str, Any]) -> "Message":
|
||||
"""Create a Message from an OpenAI-format dict."""
|
||||
return cls(
|
||||
role=data.get("role", "user"),
|
||||
content=data.get("content", ""),
|
||||
tool_call_id=data.get("tool_call_id"),
|
||||
tool_calls=data.get("tool_calls"),
|
||||
metadata={
|
||||
"timestamp": data.get("timestamp", datetime.now(timezone.utc).isoformat()),
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Base memory
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class BaseMemory(ABC):
|
||||
"""Abstract base for all memory backends."""
|
||||
|
||||
@abstractmethod
|
||||
def add(self, message: Message) -> None:
|
||||
"""Append a message."""
|
||||
|
||||
@abstractmethod
|
||||
def get_recent(self, n: int = 10) -> List[Message]:
|
||||
"""Return the most recent *n* messages."""
|
||||
|
||||
@abstractmethod
|
||||
def clear(self) -> None:
|
||||
"""Remove all stored messages."""
|
||||
|
||||
# -- conveniences -----------------------------------------------------------
|
||||
|
||||
def to_openai_messages(self) -> List[Dict[str, Any]]:
|
||||
"""Return the full history as OpenAI-format dicts."""
|
||||
return [m.to_openai() for m in self.get_recent(n=1000)]
|
||||
|
||||
def __len__(self) -> int:
|
||||
return len(self.get_recent(n=100000))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Conversation memory (sliding window)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ConversationMemory(BaseMemory):
|
||||
"""In-memory message buffer with a sliding-window recall.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
window_size : int
|
||||
Maximum number of messages kept. Excess messages are
|
||||
trimmed from the front on ``add``.
|
||||
"""
|
||||
|
||||
def __init__(self, window_size: int = 50) -> None:
|
||||
self.window_size = window_size
|
||||
self._messages: List[Message] = []
|
||||
|
||||
# -- ABC -------------------------------------------------------------------
|
||||
|
||||
def add(self, message: Message) -> None:
|
||||
"""Append a message, trimming the oldest if the window is full."""
|
||||
self._messages.append(message)
|
||||
if len(self._messages) > self.window_size:
|
||||
overflow = len(self._messages) - self.window_size
|
||||
self._messages = self._messages[overflow:]
|
||||
logger.debug(
|
||||
"ConversationMemory trimmed %d old message(s)", overflow
|
||||
)
|
||||
|
||||
def get_recent(self, n: int = 10) -> List[Message]:
|
||||
"""Return the last *n* messages (up to what is stored)."""
|
||||
return list(self._messages[-n:]) if n > 0 else []
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Erase the entire buffer."""
|
||||
self._messages.clear()
|
||||
|
||||
# -- helpers ----------------------------------------------------------------
|
||||
|
||||
def summary(self) -> str:
|
||||
"""Return a concise text summary of the conversation."""
|
||||
if not self._messages:
|
||||
return "(empty conversation)"
|
||||
parts = []
|
||||
for msg in self._messages:
|
||||
snippet = msg.content[:120].replace("\n", " ")
|
||||
parts.append(f"[{msg.role}] {snippet}")
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Long-term memory (keyword → value store)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class LongTermMemory(BaseMemory):
|
||||
"""Simple keyword-based persistent memory using a dict.
|
||||
|
||||
Keys are normalised (lower-cased, stripped). Each key maps to a list
|
||||
of values that were stored under it over time, allowing the agent to
|
||||
recall past facts.
|
||||
|
||||
Parameters
|
||||
----------
|
||||
max_entries_per_key : int
|
||||
Cap on stored values per keyword.
|
||||
"""
|
||||
|
||||
def __init__(self, max_entries_per_key: int = 20) -> None:
|
||||
self.max_entries_per_key = max_entries_per_key
|
||||
self._store: Dict[str, List[str]] = {}
|
||||
|
||||
# -- ABC -------------------------------------------------------------------
|
||||
|
||||
def add(self, message: Message) -> None:
|
||||
"""Extract keywords from the message and store the content."""
|
||||
# Simple keyword extraction: split on whitespace / punctuation, keep
|
||||
# tokens >= 3 chars, lower-case.
|
||||
words = set(
|
||||
w.strip().lower()
|
||||
for w in message.content.replace("-", " ").split()
|
||||
if len(w.strip()) >= 3
|
||||
)
|
||||
for kw in words:
|
||||
bucket = self._store.setdefault(kw, [])
|
||||
# Avoid exact duplicates
|
||||
if bucket and bucket[-1] == message.content:
|
||||
continue
|
||||
bucket.append(message.content)
|
||||
if len(bucket) > self.max_entries_per_key:
|
||||
self._store[kw] = bucket[-self.max_entries_per_key :]
|
||||
|
||||
def get_recent(self, n: int = 10) -> List[Message]:
|
||||
"""Return *n* messages reconstructed from recent keyword inserts."""
|
||||
# Flatten store in insertion order (most recent first)
|
||||
all_items: List[tuple] = []
|
||||
for values in self._store.values():
|
||||
all_items.extend(enumerate(values))
|
||||
all_items.sort(reverse=True)
|
||||
|
||||
seen: set = set()
|
||||
results: List[Message] = []
|
||||
for _idx, val in all_items:
|
||||
if val not in seen and len(results) < n:
|
||||
seen.add(val)
|
||||
results.append(Message(role="system", content=f"[memory] {val}"))
|
||||
return results
|
||||
|
||||
def clear(self) -> None:
|
||||
self._store.clear()
|
||||
|
||||
# -- keyword search ---------------------------------------------------------
|
||||
|
||||
def recall(self, keyword: str, limit: int = 5) -> List[str]:
|
||||
"""Retrieve stored values for a given keyword."""
|
||||
key = keyword.strip().lower()
|
||||
bucket = self._store.get(key, [])
|
||||
return bucket[-limit:]
|
||||
|
||||
def search(self, query: str) -> List[str]:
|
||||
"""Search across all stored values for *query*.
|
||||
|
||||
Returns unique content strings that contain the query substring.
|
||||
"""
|
||||
needle = query.lower()
|
||||
hits: set = set()
|
||||
for values in self._store.values():
|
||||
for val in values:
|
||||
if needle in val.lower():
|
||||
hits.add(val)
|
||||
return list(hits)[:10]
|
||||
|
||||
def keys(self) -> List[str]:
|
||||
"""Return all registered keywords."""
|
||||
return list(self._store.keys())
|
||||
Reference in New Issue
Block a user