@@ -1,15 +1,29 @@
|
||||
"""
|
||||
智能体记忆系统
|
||||
使用 SQLite + 简单的文本相似度匹配
|
||||
Agent memory system (local-only).
|
||||
|
||||
This module stores agent experiences in SQLite and retrieves relevant past cases
|
||||
to inject into prompts (RAG-style). It does NOT finetune model weights.
|
||||
|
||||
Retrieval (configurable):
|
||||
- Vector similarity via deterministic local embeddings (default)
|
||||
- Fallback to difflib text similarity when embeddings are missing
|
||||
|
||||
Ranking combines:
|
||||
- similarity
|
||||
- recency decay (half-life)
|
||||
- optional returns weight
|
||||
"""
|
||||
|
||||
import sqlite3
|
||||
import json
|
||||
import os
|
||||
import math
|
||||
from typing import List, Dict, Any, Optional
|
||||
from datetime import datetime
|
||||
from datetime import datetime, timezone
|
||||
import difflib
|
||||
|
||||
from app.utils.logger import get_logger
|
||||
from .embedding import EmbeddingService, cosine_sim
|
||||
|
||||
logger = get_logger(__name__)
|
||||
|
||||
@@ -34,6 +48,13 @@ class AgentMemory:
|
||||
db_path = os.path.join(db_dir, f'{agent_name}_memory.db')
|
||||
|
||||
self.db_path = db_path
|
||||
self.embedder = EmbeddingService()
|
||||
self.enable_vector = os.getenv("AGENT_MEMORY_ENABLE_VECTOR", "true").lower() == "true"
|
||||
self.candidate_limit = int(os.getenv("AGENT_MEMORY_CANDIDATE_LIMIT", "500") or 500)
|
||||
self.half_life_days = float(os.getenv("AGENT_MEMORY_HALF_LIFE_DAYS", "30") or 30)
|
||||
self.w_sim = float(os.getenv("AGENT_MEMORY_W_SIM", "0.75") or 0.75)
|
||||
self.w_recency = float(os.getenv("AGENT_MEMORY_W_RECENCY", "0.20") or 0.20)
|
||||
self.w_returns = float(os.getenv("AGENT_MEMORY_W_RETURNS", "0.05") or 0.05)
|
||||
self._init_database()
|
||||
|
||||
def _init_database(self):
|
||||
@@ -49,22 +70,91 @@ class AgentMemory:
|
||||
recommendation TEXT NOT NULL,
|
||||
result TEXT,
|
||||
returns REAL,
|
||||
market TEXT,
|
||||
symbol TEXT,
|
||||
timeframe TEXT,
|
||||
features_json TEXT,
|
||||
embedding BLOB,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
)
|
||||
''')
|
||||
|
||||
# 创建索引
|
||||
|
||||
# Best-effort migration for older DBs
|
||||
# NOTE: For existing tables, we must add missing columns BEFORE creating indexes
|
||||
# that reference them (otherwise we'll hit: "no such column: market").
|
||||
cursor.execute("PRAGMA table_info(memories)")
|
||||
existing_cols = {row[1] for row in cursor.fetchall() or []}
|
||||
for col, ddl in {
|
||||
"market": "TEXT",
|
||||
"symbol": "TEXT",
|
||||
"timeframe": "TEXT",
|
||||
"features_json": "TEXT",
|
||||
"embedding": "BLOB",
|
||||
}.items():
|
||||
if col not in existing_cols:
|
||||
cursor.execute(f"ALTER TABLE memories ADD COLUMN {col} {ddl}")
|
||||
|
||||
# 创建索引(放在迁移之后,兼容旧库)
|
||||
cursor.execute('''
|
||||
CREATE INDEX IF NOT EXISTS idx_created_at ON memories(created_at)
|
||||
''')
|
||||
cursor.execute('''
|
||||
CREATE INDEX IF NOT EXISTS idx_market_symbol ON memories(market, symbol)
|
||||
''')
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
except Exception as e:
|
||||
logger.error(f"初始化记忆数据库失败: {e}")
|
||||
|
||||
def add_memory(self, situation: str, recommendation: str, result: Optional[str] = None, returns: Optional[float] = None):
|
||||
|
||||
def _now_utc(self) -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
def _parse_ts(self, ts_val: Any) -> Optional[datetime]:
|
||||
if ts_val is None:
|
||||
return None
|
||||
if isinstance(ts_val, datetime):
|
||||
return ts_val
|
||||
s = str(ts_val)
|
||||
try:
|
||||
return datetime.fromisoformat(s.replace("Z", ""))
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def _recency_score(self, created_at: Any) -> float:
|
||||
dt = self._parse_ts(created_at)
|
||||
if not dt:
|
||||
return 0.0
|
||||
if dt.tzinfo is None:
|
||||
dt = dt.replace(tzinfo=timezone.utc)
|
||||
age_days = max(0.0, (self._now_utc() - dt).total_seconds() / 86400.0)
|
||||
hl = max(0.1, float(self.half_life_days or 30.0))
|
||||
return float(math.exp(-math.log(2.0) * (age_days / hl)))
|
||||
|
||||
def _returns_score(self, returns: Any) -> float:
|
||||
try:
|
||||
r = float(returns)
|
||||
except Exception:
|
||||
return 0.0
|
||||
return float(math.tanh(r / 10.0))
|
||||
|
||||
def _build_embed_text(self, situation: str, recommendation: str, result: Optional[str], features_json: Optional[str]) -> str:
|
||||
return "\n".join([
|
||||
f"situation: {situation or ''}",
|
||||
f"recommendation: {recommendation or ''}",
|
||||
f"result: {result or ''}",
|
||||
f"features: {features_json or ''}",
|
||||
])
|
||||
|
||||
def add_memory(
|
||||
self,
|
||||
situation: str,
|
||||
recommendation: str,
|
||||
result: Optional[str] = None,
|
||||
returns: Optional[float] = None,
|
||||
metadata: Optional[Dict[str, Any]] = None,
|
||||
):
|
||||
"""
|
||||
添加记忆
|
||||
|
||||
@@ -73,15 +163,32 @@ class AgentMemory:
|
||||
recommendation: 建议/决策
|
||||
result: 结果描述(可选)
|
||||
returns: 收益(可选)
|
||||
metadata: Optional structured metadata (market/symbol/timeframe/features...)
|
||||
"""
|
||||
try:
|
||||
meta = metadata or {}
|
||||
market = (meta.get("market") or "").strip() or None
|
||||
symbol = (meta.get("symbol") or "").strip() or None
|
||||
timeframe = (meta.get("timeframe") or "").strip() or None
|
||||
features = meta.get("features") if isinstance(meta, dict) else None
|
||||
try:
|
||||
features_json = json.dumps(features, ensure_ascii=False) if features is not None else None
|
||||
except Exception:
|
||||
features_json = None
|
||||
|
||||
embedding_blob = None
|
||||
if self.enable_vector:
|
||||
text = self._build_embed_text(situation, recommendation, result, features_json)
|
||||
vec = self.embedder.embed(text)
|
||||
embedding_blob = self.embedder.to_bytes(vec)
|
||||
|
||||
conn = sqlite3.connect(self.db_path)
|
||||
cursor = conn.cursor()
|
||||
|
||||
cursor.execute('''
|
||||
INSERT INTO memories (situation, recommendation, result, returns)
|
||||
VALUES (?, ?, ?, ?)
|
||||
''', (situation, recommendation, result, returns))
|
||||
INSERT INTO memories (situation, recommendation, result, returns, market, symbol, timeframe, features_json, embedding)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
''', (situation, recommendation, result, returns, market, symbol, timeframe, features_json, embedding_blob))
|
||||
|
||||
conn.commit()
|
||||
conn.close()
|
||||
@@ -89,7 +196,7 @@ class AgentMemory:
|
||||
except Exception as e:
|
||||
logger.error(f"添加记忆失败: {e}")
|
||||
|
||||
def get_memories(self, current_situation: str, n_matches: int = 2) -> List[Dict[str, Any]]:
|
||||
def get_memories(self, current_situation: str, n_matches: int = 5, metadata: Optional[Dict[str, Any]] = None) -> List[Dict[str, Any]]:
|
||||
"""
|
||||
检索相似记忆
|
||||
|
||||
@@ -106,11 +213,11 @@ class AgentMemory:
|
||||
|
||||
# 获取所有记忆
|
||||
cursor.execute('''
|
||||
SELECT id, situation, recommendation, result, returns, created_at
|
||||
SELECT id, situation, recommendation, result, returns, created_at, market, symbol, timeframe, features_json, embedding
|
||||
FROM memories
|
||||
ORDER BY created_at DESC
|
||||
LIMIT 100
|
||||
''')
|
||||
LIMIT ?
|
||||
''', (int(self.candidate_limit),))
|
||||
|
||||
all_memories = cursor.fetchall()
|
||||
conn.close()
|
||||
@@ -118,33 +225,71 @@ class AgentMemory:
|
||||
if not all_memories:
|
||||
return []
|
||||
|
||||
# 计算相似度
|
||||
scored_memories = []
|
||||
for mem in all_memories:
|
||||
mem_id, situation, recommendation, result, returns, created_at = mem
|
||||
|
||||
# 使用简单的文本相似度
|
||||
similarity = difflib.SequenceMatcher(
|
||||
None,
|
||||
current_situation.lower(),
|
||||
situation.lower()
|
||||
).ratio()
|
||||
|
||||
scored_memories.append({
|
||||
meta = metadata or {}
|
||||
tf = (meta.get("timeframe") or "").strip()
|
||||
features = meta.get("features") if isinstance(meta, dict) else None
|
||||
try:
|
||||
q_features_json = json.dumps(features, ensure_ascii=False) if features is not None else None
|
||||
except Exception:
|
||||
q_features_json = None
|
||||
|
||||
query_vec = []
|
||||
if self.enable_vector:
|
||||
query_text = self._build_embed_text(current_situation, "", "", q_features_json)
|
||||
query_vec = self.embedder.embed(query_text)
|
||||
|
||||
ranked = []
|
||||
for row in all_memories:
|
||||
(
|
||||
mem_id,
|
||||
situation,
|
||||
recommendation,
|
||||
result,
|
||||
returns,
|
||||
created_at,
|
||||
market,
|
||||
symbol,
|
||||
timeframe,
|
||||
features_json,
|
||||
embedding_blob,
|
||||
) = row
|
||||
|
||||
sim = 0.0
|
||||
if self.enable_vector and embedding_blob:
|
||||
try:
|
||||
mem_vec = self.embedder.from_bytes(embedding_blob)
|
||||
sim = cosine_sim(query_vec, mem_vec)
|
||||
except Exception:
|
||||
sim = 0.0
|
||||
else:
|
||||
sim = difflib.SequenceMatcher(None, (current_situation or "").lower(), (situation or "").lower()).ratio()
|
||||
|
||||
rec = self._recency_score(created_at)
|
||||
ret = self._returns_score(returns)
|
||||
|
||||
score = (self.w_sim * sim) + (self.w_recency * rec) + (self.w_returns * ret)
|
||||
|
||||
if tf and timeframe and str(timeframe).strip() != tf:
|
||||
score -= 0.15
|
||||
|
||||
ranked.append({
|
||||
'id': mem_id,
|
||||
'matched_situation': situation,
|
||||
'recommendation': recommendation,
|
||||
'result': result,
|
||||
'returns': returns,
|
||||
'similarity_score': similarity,
|
||||
'created_at': created_at
|
||||
'created_at': created_at,
|
||||
'market': market,
|
||||
'symbol': symbol,
|
||||
'timeframe': timeframe,
|
||||
'features_json': features_json,
|
||||
'score': float(score),
|
||||
'sim': float(sim),
|
||||
'recency': float(rec),
|
||||
})
|
||||
|
||||
# 按相似度排序
|
||||
scored_memories.sort(key=lambda x: x['similarity_score'], reverse=True)
|
||||
|
||||
# 返回前 n_matches 个
|
||||
return scored_memories[:n_matches]
|
||||
|
||||
ranked.sort(key=lambda x: x.get('score', 0.0), reverse=True)
|
||||
return ranked[: max(0, int(n_matches or 0))]
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"检索记忆失败: {e}")
|
||||
|
||||
Reference in New Issue
Block a user