Files
zk-data-agent/src/personal_memory.py
T
2026-05-13 14:04:42 +08:00

738 lines
26 KiB
Python

"""账号级个性化记忆。
这个模块只做轻量、可替换的第一版实现:
- Markdown 是人可编辑的记忆正文。
- SQLite 是后台事件队列、状态和 revision 账本。
- 主对话链路只追加事件和读取已有记忆,不在同步路径里整理记忆。
"""
from __future__ import annotations
import json
import re
import sqlite3
import threading
import time
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
from uuid import uuid4
from .agent_types import ModelConfig
from .openai_compat import OpenAICompatClient
MEMORY_DB_FILENAME = 'memory.db'
USER_MEMORY_FILENAME = 'user.md'
SKILL_MEMORY_DIRNAME = 'skills'
MAX_MEMORY_LINES_BEFORE_COMPACT = 300
EXPLICIT_MEMORY_PATTERNS = (
'记住',
'以后',
'下次',
'默认',
'总是',
'不要',
'应该',
'固定',
)
CORRECTION_PATTERNS = (
'不对',
'不是这样',
'格式错',
'之前说过',
'还是不行',
'这个不对',
'需要改成',
)
@dataclass(frozen=True)
class MemoryEvent:
account_id: str
session_id: str
interaction_id: str
user_prompt: str
assistant_output: str
skills: tuple[str, ...]
model: str
priority: int
signals: tuple[str, ...]
created_at: str
class PersonalMemoryManager:
"""管理用户记忆、Skill 记忆和异步整理队列。"""
def __init__(self, accounts_root: Path, model_config_getter: Any) -> None:
self.accounts_root = accounts_root
self._model_config_getter = model_config_getter
self._stop_event = threading.Event()
self._worker_thread: threading.Thread | None = None
self._init_lock = threading.Lock()
def start(self) -> None:
if self._worker_thread and self._worker_thread.is_alive():
return
self._worker_thread = threading.Thread(
target=self._worker_loop,
name='personal-memory-worker',
daemon=True,
)
self._worker_thread.start()
def stop(self) -> None:
self._stop_event.set()
def ensure_account(self, account_id: str) -> Path:
base = self._account_memory_root(account_id)
(base / SKILL_MEMORY_DIRNAME).mkdir(parents=True, exist_ok=True)
self._ensure_db(account_id)
user_path = base / USER_MEMORY_FILENAME
if not user_path.exists():
self._atomic_write_text(
user_path,
'# 用户记忆\n\n暂无用户记忆。\n',
)
return base
def render_injection(
self,
account_id: str | None,
enabled_skill_names: tuple[str, ...] | None,
) -> str:
"""渲染给模型的个性化记忆上下文。"""
if not account_id:
return ''
base = self.ensure_account(account_id)
sections: list[str] = []
user_text = self._read_memory_file(base / USER_MEMORY_FILENAME)
if user_text:
sections.extend(['## 用户记忆', user_text])
skill_root = base / SKILL_MEMORY_DIRNAME
skill_sections: list[str] = []
skill_names = enabled_skill_names or ()
for skill_name in skill_names:
safe_name = _safe_memory_name(skill_name)
if not safe_name:
continue
text = self._read_memory_file(skill_root / f'{safe_name}.md')
if text:
skill_sections.extend([f'### {skill_name}', text])
if skill_sections:
sections.extend(['## Skill 使用记忆', *skill_sections])
if not sections:
return ''
return '\n'.join(
[
'# 个性化记忆',
'以下是该账号长期保存的偏好和 Skill 使用经验。若与用户本轮明确要求冲突,以用户本轮要求为准。',
*sections,
]
).strip()
def enqueue_interaction(
self,
*,
account_id: str | None,
session_id: str | None,
user_prompt: str,
assistant_output: str,
skills: tuple[str, ...],
model: str,
) -> None:
if not account_id or not session_id:
return
signals = detect_memory_signals(user_prompt, assistant_output, skills)
if not signals:
return
event = MemoryEvent(
account_id=account_id,
session_id=session_id,
interaction_id=uuid4().hex,
user_prompt=user_prompt[-6000:],
assistant_output=assistant_output[-6000:],
skills=skills,
model=model,
priority=10 if 'explicit' in signals else 3,
signals=signals,
created_at=_now_iso(),
)
self._insert_event(event)
def list_user_memory(self, account_id: str) -> dict[str, Any]:
base = self.ensure_account(account_id)
return {
'account_id': account_id,
'kind': 'user',
'content': (base / USER_MEMORY_FILENAME).read_text(encoding='utf-8'),
'line_count': _line_count(base / USER_MEMORY_FILENAME),
'updated_at': _mtime_iso(base / USER_MEMORY_FILENAME),
}
def update_user_memory(self, account_id: str, content: str) -> dict[str, Any]:
base = self.ensure_account(account_id)
path = base / USER_MEMORY_FILENAME
self._atomic_write_text(path, _normalize_memory_markdown(content, '用户记忆'))
self._record_revision(account_id, 'user', None)
return self.list_user_memory(account_id)
def list_skill_memories(self, account_id: str) -> list[dict[str, Any]]:
base = self.ensure_account(account_id)
skill_root = base / SKILL_MEMORY_DIRNAME
result: list[dict[str, Any]] = []
for path in sorted(skill_root.glob('*.md')):
result.append(
{
'skill': path.stem,
'content': path.read_text(encoding='utf-8'),
'line_count': _line_count(path),
'updated_at': _mtime_iso(path),
}
)
return result
def read_skill_memory(self, account_id: str, skill_name: str) -> dict[str, Any]:
base = self.ensure_account(account_id)
safe_name = _safe_memory_name(skill_name)
if not safe_name:
raise ValueError('skill name must not be empty')
path = base / SKILL_MEMORY_DIRNAME / f'{safe_name}.md'
if not path.exists():
self._atomic_write_text(path, f'# {safe_name} 使用记忆\n\n暂无 Skill 使用记忆。\n')
return {
'account_id': account_id,
'skill': safe_name,
'content': path.read_text(encoding='utf-8'),
'line_count': _line_count(path),
'updated_at': _mtime_iso(path),
}
def update_skill_memory(
self,
account_id: str,
skill_name: str,
content: str,
) -> dict[str, Any]:
base = self.ensure_account(account_id)
safe_name = _safe_memory_name(skill_name)
if not safe_name:
raise ValueError('skill name must not be empty')
path = base / SKILL_MEMORY_DIRNAME / f'{safe_name}.md'
self._atomic_write_text(
path,
_normalize_memory_markdown(content, f'{safe_name} 使用记忆'),
)
self._record_revision(account_id, 'skill', safe_name)
return self.read_skill_memory(account_id, safe_name)
def delete_skill_memory(self, account_id: str, skill_name: str) -> dict[str, Any]:
base = self.ensure_account(account_id)
safe_name = _safe_memory_name(skill_name)
if not safe_name:
raise ValueError('skill name must not be empty')
path = base / SKILL_MEMORY_DIRNAME / f'{safe_name}.md'
if path.exists():
path.unlink()
self._record_revision(account_id, 'skill', safe_name)
return {'deleted': True, 'skill': safe_name}
def queue_snapshot(self, account_id: str | None = None) -> dict[str, Any]:
accounts = [account_id] if account_id else self._list_account_ids()
totals = {
'pending': 0,
'processing': 0,
'done': 0,
'failed': 0,
'events': 0,
}
account_rows: list[dict[str, Any]] = []
for item in accounts:
self.ensure_account(item)
conn = self._connect(item)
try:
rows = conn.execute(
'select status, count(*) from memory_events group by status'
).fetchall()
counts = {str(status): int(count) for status, count in rows}
events = sum(counts.values())
account_rows.append(
{
'account_id': item,
'events': events,
'pending': counts.get('pending', 0),
'processing': counts.get('processing', 0),
'done': counts.get('done', 0),
'failed': counts.get('failed', 0),
}
)
totals['events'] += events
for key in ('pending', 'processing', 'done', 'failed'):
totals[key] += counts.get(key, 0)
finally:
conn.close()
return {'totals': totals, 'accounts': account_rows}
def list_recent_events(self, account_id: str, limit: int = 50) -> list[dict[str, Any]]:
self.ensure_account(account_id)
conn = self._connect(account_id)
try:
rows = conn.execute(
'''
select id, session_id, skills_json, signals_json, priority, status,
error, created_at, updated_at
from memory_events
order by created_at desc
limit ?
''',
(max(1, min(limit, 200)),),
).fetchall()
return [
{
'id': row['id'],
'session_id': row['session_id'],
'skills': _loads_json_list(row['skills_json']),
'signals': _loads_json_list(row['signals_json']),
'priority': row['priority'],
'status': row['status'],
'error': row['error'],
'created_at': row['created_at'],
'updated_at': row['updated_at'],
}
for row in rows
]
finally:
conn.close()
def _worker_loop(self) -> None:
while not self._stop_event.wait(5.0):
for account_id in self._list_account_ids():
try:
self._process_account_events(account_id)
except Exception:
continue
def _process_account_events(self, account_id: str) -> None:
self.ensure_account(account_id)
conn = self._connect(account_id)
try:
conn.execute('begin immediate')
rows = conn.execute(
'''
select * from memory_events
where status = 'pending'
order by priority desc, created_at asc
limit 8
'''
).fetchall()
if not rows:
conn.commit()
return
now = _now_iso()
ids = [row['id'] for row in rows]
conn.executemany(
'update memory_events set status = ?, updated_at = ? where id = ?',
[('processing', now, row_id) for row_id in ids],
)
conn.commit()
except Exception:
conn.rollback()
raise
finally:
conn.close()
try:
self._consolidate_events(account_id, rows)
except Exception as exc:
conn = self._connect(account_id)
try:
now = _now_iso()
conn.executemany(
'''
update memory_events
set status = 'failed', error = ?, updated_at = ?
where id = ?
''',
[(str(exc), now, row_id) for row_id in ids],
)
conn.commit()
finally:
conn.close()
return
conn = self._connect(account_id)
try:
now = _now_iso()
conn.executemany(
"update memory_events set status = 'done', updated_at = ? where id = ?",
[(now, row_id) for row_id in ids],
)
conn.commit()
finally:
conn.close()
def _consolidate_events(self, account_id: str, rows: list[sqlite3.Row]) -> None:
events = [
{
'user_prompt': row['user_prompt'],
'assistant_output': row['assistant_output'],
'skills': _loads_json_list(row['skills_json']),
'signals': _loads_json_list(row['signals_json']),
}
for row in rows
]
skills = sorted(
{
skill
for event in events
for skill in event['skills']
if isinstance(skill, str) and skill.strip()
}
)
base = self.ensure_account(account_id)
existing_user = self._read_memory_file(base / USER_MEMORY_FILENAME)
existing_skills = {
skill: self._read_memory_file(
base / SKILL_MEMORY_DIRNAME / f'{_safe_memory_name(skill)}.md'
)
for skill in skills
if _safe_memory_name(skill)
}
updates = self._generate_memory_updates(
account_id=account_id,
events=events,
existing_user=existing_user,
existing_skills=existing_skills,
)
if updates.get('user_memory'):
self.update_user_memory(account_id, str(updates['user_memory']))
skill_updates = updates.get('skill_memories')
if isinstance(skill_updates, dict):
for skill, content in skill_updates.items():
if isinstance(skill, str) and isinstance(content, str) and content.strip():
self.update_skill_memory(account_id, skill, content)
def _generate_memory_updates(
self,
*,
account_id: str,
events: list[dict[str, Any]],
existing_user: str,
existing_skills: dict[str, str],
) -> dict[str, Any]:
config = self._model_config_getter(account_id)
messages = [
{
'role': 'system',
'content': (
'你是 ZK Data Agent 的记忆整理后台。'
'请只沉淀长期稳定的用户偏好和 Skill 使用经验。'
'不要记录一次性任务目标、临时文件名、普通聊天内容。'
'输出必须是 JSON 对象。'
),
},
{
'role': 'user',
'content': json.dumps(
{
'existing_user_memory': existing_user,
'existing_skill_memories': existing_skills,
'new_events': events,
'requirements': [
'保留 Markdown,可整理合并,不要追加流水账。',
'如果没有值得更新的用户记忆,user_memory 返回空字符串。',
'skill_memories 只返回需要更新的 skill。',
],
'output_schema': {
'user_memory': '完整的用户记忆 Markdown,或空字符串',
'skill_memories': {
'skill-name': '完整的 Skill 记忆 Markdown'
},
},
},
ensure_ascii=False,
),
},
]
turn = OpenAICompatClient(config).complete(messages, tools=[])
content = turn.content.strip()
try:
parsed = json.loads(_extract_json_object(content))
except Exception:
parsed = self._fallback_memory_updates(events, existing_user, existing_skills)
if not isinstance(parsed, dict):
return {}
return parsed
def _fallback_memory_updates(
self,
events: list[dict[str, Any]],
existing_user: str,
existing_skills: dict[str, str],
) -> dict[str, Any]:
user_lines = _memory_body_lines(existing_user)
skill_updates: dict[str, str] = {}
for event in events:
prompt = str(event.get('user_prompt') or '').strip()
if any(pattern in prompt for pattern in EXPLICIT_MEMORY_PATTERNS):
candidate = _shorten_memory_line(prompt)
if candidate and candidate not in user_lines:
user_lines.append(candidate)
for skill in event.get('skills') or []:
if not isinstance(skill, str) or not skill.strip():
continue
body = _memory_body_lines(existing_skills.get(skill, ''))
if any(pattern in prompt for pattern in CORRECTION_PATTERNS):
candidate = _shorten_memory_line(prompt)
if candidate and candidate not in body:
body.append(candidate)
skill_updates[skill] = _render_memory_doc(
f'{skill} 使用记忆',
body,
)
result: dict[str, Any] = {'skill_memories': skill_updates}
if user_lines != _memory_body_lines(existing_user):
result['user_memory'] = _render_memory_doc('用户记忆', user_lines)
return result
def _insert_event(self, event: MemoryEvent) -> None:
self.ensure_account(event.account_id)
conn = self._connect(event.account_id)
try:
conn.execute(
'''
insert into memory_events (
id, account_id, session_id, interaction_id, user_prompt,
assistant_output, skills_json, signals_json, model, priority,
status, created_at, updated_at
)
values (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'pending', ?, ?)
''',
(
uuid4().hex,
event.account_id,
event.session_id,
event.interaction_id,
event.user_prompt,
event.assistant_output,
json.dumps(list(event.skills), ensure_ascii=False),
json.dumps(list(event.signals), ensure_ascii=False),
event.model,
event.priority,
event.created_at,
event.created_at,
),
)
conn.commit()
finally:
conn.close()
def _record_revision(
self,
account_id: str,
kind: str,
skill_name: str | None,
) -> None:
self.ensure_account(account_id)
key = f'{kind}:{skill_name or ""}'
conn = self._connect(account_id)
try:
now = _now_iso()
conn.execute(
'''
insert into memory_revisions (key, kind, skill_name, revision, updated_at)
values (?, ?, ?, 1, ?)
on conflict(key) do update set
revision = revision + 1,
updated_at = excluded.updated_at
''',
(key, kind, skill_name, now),
)
conn.commit()
finally:
conn.close()
def _ensure_db(self, account_id: str) -> None:
with self._init_lock:
conn = self._connect(account_id)
try:
conn.executescript(
'''
create table if not exists memory_events (
id text primary key,
account_id text not null,
session_id text not null,
interaction_id text not null,
user_prompt text not null,
assistant_output text not null,
skills_json text not null,
signals_json text not null,
model text not null,
priority integer not null default 0,
status text not null default 'pending',
error text,
created_at text not null,
updated_at text not null
);
create index if not exists idx_memory_events_status
on memory_events(status, priority, created_at);
create table if not exists memory_revisions (
key text primary key,
kind text not null,
skill_name text,
revision integer not null default 1,
updated_at text not null
);
'''
)
conn.commit()
finally:
conn.close()
def _connect(self, account_id: str) -> sqlite3.Connection:
path = self._account_memory_root(account_id) / MEMORY_DB_FILENAME
path.parent.mkdir(parents=True, exist_ok=True)
conn = sqlite3.connect(path, timeout=10)
conn.row_factory = sqlite3.Row
return conn
def _account_memory_root(self, account_id: str) -> Path:
return self.accounts_root / _safe_account_id(account_id) / 'memory'
def _list_account_ids(self) -> list[str]:
if not self.accounts_root.exists():
return []
result: list[str] = []
for path in sorted(self.accounts_root.iterdir()):
if not path.is_dir():
continue
if (path / 'sessions').exists() or (path / 'memory').exists():
result.append(path.name)
return result
@staticmethod
def _read_memory_file(path: Path) -> str:
try:
text = path.read_text(encoding='utf-8')
except OSError:
return ''
if '暂无' in text and len(_memory_body_lines(text)) == 0:
return ''
return text.strip()
@staticmethod
def _atomic_write_text(path: Path, content: str) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
tmp = path.with_suffix(path.suffix + '.tmp')
tmp.write_text(content.rstrip() + '\n', encoding='utf-8')
tmp.replace(path)
def detect_memory_signals(
user_prompt: str,
assistant_output: str,
skills: tuple[str, ...],
) -> tuple[str, ...]:
signals: list[str] = []
prompt = user_prompt or ''
if any(pattern in prompt for pattern in EXPLICIT_MEMORY_PATTERNS):
signals.append('explicit')
if any(pattern in prompt for pattern in CORRECTION_PATTERNS):
signals.append('correction')
if skills and any(keyword in prompt for keyword in ('skill', 'Skill', '工具', '流程', '格式')):
signals.append('skill')
if '模型返回的工具参数不是合法 JSON' in assistant_output:
signals.append('tool_experience')
return tuple(dict.fromkeys(signals))
def _extract_json_object(text: str) -> str:
text = text.strip()
if text.startswith('```'):
text = re.sub(r'^```(?:json)?\s*', '', text)
text = re.sub(r'\s*```$', '', text)
start = text.find('{')
end = text.rfind('}')
if start == -1 or end == -1 or end < start:
raise ValueError('no JSON object found')
return text[start : end + 1]
def _normalize_memory_markdown(content: str, title: str) -> str:
text = content.strip()
if not text:
return f'# {title}\n\n暂无记忆。\n'
if not text.lstrip().startswith('#'):
text = f'# {title}\n\n{text}'
return text
def _render_memory_doc(title: str, lines: list[str]) -> str:
clean_lines = [line.strip() for line in lines if line.strip()]
if not clean_lines:
return f'# {title}\n\n暂无记忆。\n'
return '\n'.join([f'# {title}', '', *[f'- {line}' for line in clean_lines]])
def _memory_body_lines(text: str) -> list[str]:
lines: list[str] = []
for raw in text.splitlines():
line = raw.strip()
if not line or line.startswith('#') or '暂无' in line:
continue
if line.startswith('- '):
line = line[2:].strip()
lines.append(line)
return lines
def _shorten_memory_line(text: str) -> str:
text = re.sub(r'\s+', ' ', text).strip()
if not text:
return ''
return text[:240]
def _safe_account_id(value: str) -> str:
clean = re.sub(r'[^a-zA-Z0-9._-]+', '_', value.strip())
return clean[:80] or 'unknown'
def _safe_memory_name(value: str) -> str:
clean = re.sub(r'[^a-zA-Z0-9._-]+', '-', value.strip().lower())
return clean[:100]
def _line_count(path: Path) -> int:
try:
return len(path.read_text(encoding='utf-8').splitlines())
except OSError:
return 0
def _mtime_iso(path: Path) -> str | None:
try:
return datetime.fromtimestamp(path.stat().st_mtime, tz=timezone.utc).isoformat()
except OSError:
return None
def _loads_json_list(value: str | None) -> list[str]:
if not value:
return []
try:
parsed = json.loads(value)
except json.JSONDecodeError:
return []
if not isinstance(parsed, list):
return []
return [str(item) for item in parsed if str(item).strip()]
def _now_iso() -> str:
return datetime.now(timezone.utc).isoformat()