Enhance memory management and compaction features

- Implement global memory layer for cross-project patterns.
- Improve memory entry scoring with recency and tokenization.
- Adjust MAX_TEXT limit from 400 to 800 for better context retention.
- Add TTL pruning for stale memory entries.
- Capture and summarize dropped content during compaction.
- Update tests to reflect changes in memory behavior and compaction logic.
This commit is contained in:
asepharyana
2026-09-08 21:27:47 +07:00
parent b1ff00d840
commit 7db69f19da
8 changed files with 522 additions and 36 deletions
+6 -1
View File
@@ -232,7 +232,12 @@ const externalTools = await loadExternalTools(process.cwd(), async (command) =>
);
const memory = has('--no-memory') ? undefined : new Memory(process.cwd(), languageModel);
if (memory) await memory.load();
if (memory) {
await memory.load();
try { await memory.loadGlobal(); } catch {}
// Best-effort TTL cleanup so lifelong store doesn't bloat with stale 0-hit notes.
try { await memory.pruneExpired(); } catch {}
}
/** Placeholder until /provider supplies a key; it never gets called because the UI gates input. */
const unconfiguredModel: LanguageModel = {
+248 -21
View File
@@ -14,19 +14,25 @@ export type MemoryEntry = {
createdAt: string;
/** Bumped on each recall so summarisation can keep what gets used. */
hits: number;
/** Last time this entry was recalled; falls back to createdAt when absent. */
lastHitAt?: string;
};
const MAX_ENTRIES = 300;
const MAX_TEXT = 400;
const MAX_TEXT = 800;
const BOOT_ENTRIES = 20;
const SEARCH_HITS = 15;
/** Summarise once the store passes this, so the boot block stays small. */
const SUMMARISE_AT = 60;
/** Entries with 0 hits older than this are pruned (lifelong: drop stale cruft). */
const TTL_DAYS = 90;
const TTL_MS = TTL_DAYS * 24 * 60 * 60 * 1000;
const root = () => join(process.env['SHIRO_HOME'] ?? homedir(), '.shiro-neko', 'memory');
/** One file per project directory; the path is hashed because it is not filename-safe. */
const fileFor = (cwd: string) => join(root(), `${createHash('sha256').update(cwd).digest('hex').slice(0, 16)}.json`);
const globalFile = () => join(root(), '_global.json');
const KIND_LABEL: Record<MemoryKind, string> = {
fact: 'fact',
@@ -35,6 +41,28 @@ const KIND_LABEL: Record<MemoryKind, string> = {
command: 'command',
};
/** Normalize for dedup and tolerant matching: lower, collapse whitespace, strip leading/trailing punct. */
function normalize(text: string): string {
return text.trim().toLowerCase().replace(/\s+/g, ' ').replace(/^[\p{P}\s]+|[\p{P}\s]+$/gu, '');
}
/** Tokenize query/entry for scoring: split on non-alnum, drop empties. Underscore/hyphen are separators. */
function tokensOf(s: string): string[] {
return s
.toLowerCase()
.replace(/[_-]+/g, ' ')
.replace(/[^\p{L}\p{N}\s]/gu, ' ')
.split(/\s+/)
.filter(Boolean);
}
function isExpired(e: MemoryEntry, now: number): boolean {
if (e.hits > 0) return false;
const ts = Date.parse(e.createdAt);
if (Number.isNaN(ts)) return false;
return now - ts > TTL_MS;
}
/**
* Durable per-project memory, separate from the session transcript.
*
@@ -43,7 +71,9 @@ const KIND_LABEL: Record<MemoryKind, string> = {
*/
export class Memory {
private entries: MemoryEntry[] = [];
private globalEntries: MemoryEntry[] = [];
private loaded = false;
private globalLoaded = false;
constructor(
private readonly cwd = process.cwd(),
@@ -62,23 +92,53 @@ export class Memory {
this.entries = [];
}
}
// TTL pruning on load (only drops stale 0-hit cruft)
const now = Date.now();
const beforeLen = this.entries.length;
this.entries = this.entries.filter((e) => !isExpired(e, now));
if (this.entries.length !== beforeLen) await this.persist();
return this.entries;
}
async loadGlobal(): Promise<MemoryEntry[]> {
if (this.globalLoaded) return this.globalEntries;
this.globalLoaded = true;
const f = Bun.file(globalFile());
if (await f.exists()) {
try {
const parsed: unknown = await f.json();
if (Array.isArray(parsed)) this.globalEntries = parsed.filter(isEntry);
} catch {
this.globalEntries = [];
}
}
return this.globalEntries;
}
all(): MemoryEntry[] {
return [...this.entries];
}
allWithGlobal(): MemoryEntry[] {
return [...this.globalEntries, ...this.entries];
}
private async persist(): Promise<void> {
this.entries = this.entries.slice(-MAX_ENTRIES);
await Bun.write(fileFor(this.cwd), JSON.stringify(this.entries, null, 2));
}
private async persistGlobal(): Promise<void> {
this.globalEntries = this.globalEntries.slice(-MAX_ENTRIES);
await Bun.write(globalFile(), JSON.stringify(this.globalEntries, null, 2));
}
async add(kind: MemoryKind, text: string): Promise<MemoryEntry | undefined> {
await this.load();
const clean = text.trim().slice(0, MAX_TEXT);
if (!clean) throw new Error('memory text is empty');
if (this.entries.some((e) => e.text === clean)) return undefined;
const norm = normalize(clean);
if (this.entries.some((e) => normalize(e.text) === norm)) return undefined;
const entry: MemoryEntry = {
id: Bun.randomUUIDv7(),
@@ -92,6 +152,25 @@ export class Memory {
return entry;
}
/** Add to global layer (cross-project pattern). */
async addGlobal(kind: MemoryKind, text: string): Promise<MemoryEntry | undefined> {
await this.loadGlobal();
const clean = text.trim().slice(0, MAX_TEXT);
if (!clean) throw new Error('memory text is empty');
const norm = normalize(clean);
if (this.globalEntries.some((e) => normalize(e.text) === norm)) return undefined;
const entry: MemoryEntry = {
id: Bun.randomUUIDv7(),
kind,
text: clean,
createdAt: new Date().toISOString(),
hits: 0,
};
this.globalEntries.push(entry);
await this.persistGlobal();
return entry;
}
async forget(idOrPrefix: string): Promise<number> {
await this.load();
const before = this.entries.length;
@@ -106,32 +185,136 @@ export class Memory {
await this.persist();
}
/** Every term must appear. Matching entries get a hit, which protects them from summarisation. */
async search(query: string): Promise<MemoryEntry[]> {
/** Prune stale 0-hit entries older than TTL. Returns count removed. */
async pruneExpired(): Promise<number> {
await this.load();
const terms = query.toLowerCase().split(/\s+/).filter(Boolean);
if (terms.length === 0) throw new Error('query is empty');
const found = this.entries.filter((e) => {
const lower = e.text.toLowerCase();
return terms.every((t) => lower.includes(t));
});
for (const e of found) e.hits += 1;
if (found.length > 0) await this.persist();
return found.slice(-SEARCH_HITS).reverse();
const before = this.entries.length;
const now = Date.now();
this.entries = this.entries.filter((e) => !isExpired(e, now));
const removed = before - this.entries.length;
if (removed > 0) await this.persist();
return removed;
}
/** The block injected at boot: most-used first, then most recent. */
/**
* Every term must appear (compat). Scoring adds hits + recency + exact-phrase bonus,
* so the best entries surface first instead of arbitrary file order.
*/
async search(query: string): Promise<MemoryEntry[]> {
await this.load();
const terms = tokensOf(query);
if (terms.length === 0) throw new Error('query is empty');
const now = Date.now();
const scored: { e: MemoryEntry; score: number }[] = [];
for (const e of this.entries) {
const lower = e.text.toLowerCase();
const entryTokens = new Set(tokensOf(e.text));
// every query token must appear as substring or token (keeps compat with "snake_case" queries)
const allPresent = terms.every((t) => lower.includes(t) || entryTokens.has(t));
if (!allPresent) continue;
// scoring: hits weight + recency + phrase bonus
const lastHit = e.lastHitAt ? Date.parse(e.lastHitAt) : Date.parse(e.createdAt);
const ageDays = Number.isNaN(lastHit) ? 999 : (now - lastHit) / (24 * 60 * 60 * 1000);
const recency = Math.max(0, 10 - ageDays * 0.1); // ~10 pts fresh, decays over 100d
const phraseBonus = lower.includes(query.toLowerCase().trim()) ? 5 : 0;
const score = e.hits * 3 + recency + phraseBonus;
scored.push({ e, score });
}
scored.sort((a, b) => b.score - a.score || b.e.hits - a.e.hits || b.e.createdAt.localeCompare(a.e.createdAt));
const found = scored.map((s) => s.e);
const nowIso = new Date().toISOString();
for (const e of found) {
e.hits += 1;
e.lastHitAt = nowIso;
}
if (found.length > 0) await this.persist();
return found.slice(0, SEARCH_HITS);
}
/**
* The block injected at boot: diverse mix of most-used + most-recent, so a few
* stale high-hit entries cannot evict fresh decisions.
*/
render(limit = BOOT_ENTRIES): string {
if (this.entries.length === 0) return '';
const ranked = [...this.entries]
.sort((a, b) => b.hits - a.hits || b.createdAt.localeCompare(a.createdAt))
.slice(0, limit);
// ensure global is loaded synchronously if already loaded; otherwise project-only
const pool = this.entries;
const byHits = [...pool].sort((a, b) => b.hits - a.hits || b.createdAt.localeCompare(a.createdAt));
const byRecent = [...pool].sort((a, b) => b.createdAt.localeCompare(a.createdAt));
const seen = new Set<string>();
const ranked: MemoryEntry[] = [];
// interleave: take from hits, then recent, until limit
let hi = 0;
let ri = 0;
while (ranked.length < limit && (hi < byHits.length || ri < byRecent.length)) {
if (hi < byHits.length) {
const e = byHits[hi++]!;
if (!seen.has(e.id)) {
seen.add(e.id);
ranked.push(e);
}
if (ranked.length >= limit) break;
}
if (ri < byRecent.length) {
const e = byRecent[ri++]!;
if (!seen.has(e.id)) {
seen.add(e.id);
ranked.push(e);
}
}
}
// recency bonus visible in ordering already via interleave; keep hits bias slightly by stable sort
// final sort by composite score for determinism: hits*2 + recency
const now = Date.now();
ranked.sort((a, b) => {
const aRec = Math.max(0, 10 - (now - Date.parse(a.lastHitAt ?? a.createdAt)) / (24 * 60 * 60 * 1000) * 0.1);
const bRec = Math.max(0, 10 - (now - Date.parse(b.lastHitAt ?? b.createdAt)) / (24 * 60 * 60 * 1000) * 0.1);
return b.hits * 2 + bRec - (a.hits * 2 + aRec) || b.createdAt.localeCompare(a.createdAt);
});
return [
'',
'What you learned about this project in earlier sessions. Trust it, but verify anything',
'that contradicts what you can see in the code now:',
...ranked.map((e) => `- (${KIND_LABEL[e.kind]}) ${e.text}`),
...ranked.slice(0, limit).map((e) => `- (${KIND_LABEL[e.kind]}) ${e.text}`),
].join('\n');
}
/** Render including global entries (for prompt). Falls back to project-only when global empty. */
renderWithGlobal(limit = BOOT_ENTRIES): string {
const hasGlobal = this.globalEntries.length > 0;
if (!hasGlobal) return this.render(limit);
const merged = [...this.globalEntries, ...this.entries];
if (merged.length === 0) return '';
const byHits = [...merged].sort((a, b) => b.hits - a.hits || b.createdAt.localeCompare(a.createdAt));
const byRecent = [...merged].sort((a, b) => b.createdAt.localeCompare(a.createdAt));
const seen = new Set<string>();
const ranked: MemoryEntry[] = [];
let hi = 0;
let ri = 0;
while (ranked.length < limit && (hi < byHits.length || ri < byRecent.length)) {
if (hi < byHits.length) {
const e = byHits[hi++]!;
if (!seen.has(e.id)) {
seen.add(e.id);
ranked.push(e);
}
if (ranked.length >= limit) break;
}
if (ri < byRecent.length) {
const e = byRecent[ri++]!;
if (!seen.has(e.id)) {
seen.add(e.id);
ranked.push(e);
}
}
}
return [
'',
'What you learned about this project in earlier sessions (project + global). Trust but verify:',
...ranked.slice(0, limit).map((e) => `- (${KIND_LABEL[e.kind]}) ${e.text}`),
].join('\n');
}
@@ -184,6 +367,44 @@ export class Memory {
return { before, after: this.entries.length };
}
/**
* Suggest 1-3 memory candidates from a transcript. Used as afterTurn hook for
* lifelong learning; caller decides whether to persist (via add/addGlobal).
*/
async suggestFromTranscript(
messages: { role: string; content: unknown }[],
): Promise<{ kind: MemoryKind; text: string }[]> {
if (!this.model) return [];
if (messages.length < 4) return [];
try {
const slice = messages.slice(-20);
const transcript = slice
.map((m) => {
const c = typeof m.content === 'string' ? m.content : JSON.stringify(m.content).slice(0, 2000);
return `${m.role}: ${c}`;
})
.join('\n')
.slice(0, 8000);
const { text } = await generateText({
model: this.model,
system:
'Extract 0-3 durable learnings from this coding session that will still be true next session. ' +
'Only decisions with reason, traps, or working commands — not narration. ' +
'Output one per line as [fact|decision|gotcha|command] text, or empty if nothing durable.',
prompt: transcript,
maxRetries: 1,
});
return text
.split('\n')
.map((line) => /^\s*\[(fact|decision|gotcha|command)\]\s*(.+?)\s*$/i.exec(line))
.filter((m): m is RegExpExecArray => m !== null)
.map((m) => ({ kind: m[1]!.toLowerCase() as MemoryKind, text: m[2]!.trim().slice(0, MAX_TEXT) }))
.slice(0, 3);
} catch {
return [];
}
}
tools() {
return {
remember: tool({
@@ -196,8 +417,14 @@ export class Memory {
.enum(['fact', 'decision', 'gotcha', 'command'])
.describe('fact: how it is. decision: what was chosen and why. gotcha: a trap. command: an invocation that works'),
text: z.string().describe('One self-contained line, understandable with no other context'),
scope: z.enum(['project', 'global']).optional().describe('project (default) or global (cross-project pattern)'),
}),
execute: async ({ kind, text }) => {
execute: async ({ kind, text, scope }) => {
if (scope === 'global') {
const entry = await this.addGlobal(kind, text);
if (!entry) return `Already recorded globally: ${text.trim()}`;
return `Remembered globally as ${entry.kind} (${this.globalEntries.length} global): ${entry.text}`;
}
const entry = await this.add(kind, text);
if (!entry) return `Already recorded: ${text.trim()}`;
return `Remembered as ${entry.kind} (${this.entries.length} stored): ${entry.text}`;
@@ -238,7 +465,7 @@ export class Memory {
}
}
export { fileFor as memoryFileFor, root as memoryDir, KIND_LABEL };
export { fileFor as memoryFileFor, root as memoryDir, KIND_LABEL, globalFile as globalMemoryFile };
function isEntry(value: unknown): value is MemoryEntry {
if (!value || typeof value !== 'object') return false;
+64 -9
View File
@@ -173,8 +173,37 @@ export function prunePreservingItems(options: PruneOptions): ModelMessage[] {
* One agent step is two messages — the assistant's tool call and the tool message
* answering it — so 64 is about 32 steps of memory.
*/
export function droppedSpan(before: ModelMessage[], after: ModelMessage[]): ModelMessage[] {
const norm = (m: ModelMessage) => {
const c = (m as { content?: unknown }).content;
if (typeof c === 'string') return `${m.role}:${c}`;
try {
// Ignore reasoning parts and all providerOptions: a kept-but-detached
// message (reasoning stripped, itemId removed) is not considered dropped.
const filtered = Array.isArray(c)
? c.filter((p: unknown) => (p as { type?: string }).type !== 'reasoning')
: c;
const stripped = JSON.stringify(filtered, (k, v) => (k === 'providerOptions' ? undefined : k === 'itemId' ? undefined : v));
return `${m.role}:${stripped}`;
} catch {
return `${m.role}:${String(c)}`;
}
};
const keptNorm = new Set(after.map(norm));
return before.filter((m) => !keptNorm.has(norm(m)));
}
const KEEP_LADDER = [64, 32, 16, 8, 4] as const;
/**
* Token estimate used by the session harness. `len/4` undercounts tool envelopes
* (role + toolCallId + providerOptions); `len/3.6 + 8*msgs` tracks cl100k closer
* without pulling a tokenizer. Exported so session and tests share it.
*/
export function estimateTokens(messages: ModelMessage[]): number {
return Math.round(JSON.stringify(messages).length / 3.6 + messages.length * 8);
}
export type FitOptions = {
messages: ModelMessage[];
/** Estimated tokens the wire history must come in under. */
@@ -202,21 +231,47 @@ export type FitOptions = {
* not, the narrowest is returned, because sending something is better than sending a
* request that will be rejected for size.
*/
/** Head that must never be pruned: the initial user goal and first assistant ack. */
function headOf(messages: ModelMessage[]): ModelMessage[] {
if (messages.length === 0) return [];
const firstUser = messages.find((m) => m.role === 'user');
if (!firstUser) return [];
// Keep first user message; if an assistant immediately follows, keep it too (goal ack).
const idx = messages.indexOf(firstUser);
const next = messages[idx + 1];
if (next && next.role === 'assistant' && idx === 0) return messages.slice(0, 2);
return [firstUser];
}
function withHeadPreserved(all: ModelMessage[], pruned: ModelMessage[]): ModelMessage[] {
const head = headOf(all);
if (head.length === 0) return pruned;
// If head already in pruned (by identity of content's first 80 chars), leave it.
const headText = JSON.stringify(head[0]!.content).slice(0, 80);
if (pruned.some((m) => JSON.stringify(m.content).slice(0, 80) === headText)) return pruned;
// Prepend head; dedupe if head was partially kept.
return [...head, ...pruned];
}
export function pruneToFit({ messages, threshold, estimate }: FitOptions): ModelMessage[] {
const withoutReasoning = detachProviderItems(
prunePreservingItems({ messages, reasoning: 'all', emptyMessages: 'remove' }),
const withoutReasoning = withHeadPreserved(
messages,
detachProviderItems(prunePreservingItems({ messages, reasoning: 'all', emptyMessages: 'remove' })),
);
if (estimate(withoutReasoning) <= threshold) return withoutReasoning;
let narrowest = withoutReasoning;
for (const keep of KEEP_LADDER) {
narrowest = detachProviderItems(
prunePreservingItems({
messages,
reasoning: 'all',
toolCalls: `before-last-${keep}-messages`,
emptyMessages: 'remove',
}),
narrowest = withHeadPreserved(
messages,
detachProviderItems(
prunePreservingItems({
messages,
reasoning: 'all',
toolCalls: `before-last-${keep}-messages`,
emptyMessages: 'remove',
}),
),
);
if (estimate(narrowest) <= threshold) return narrowest;
}
+45 -3
View File
@@ -17,7 +17,7 @@ import { Permissions, type PermissionConfig } from './permission';
import type { PluginHost } from './plugins';
import { costOf, formatUsd } from './pricing';
import { systemPrompt } from './prompt';
import { detachProviderItems, pruneToFit } from './prune';
import { detachProviderItems, droppedSpan, estimateTokens as pruneEstimateTokens, pruneToFit } from './prune';
import { createSkillTool, renderSkills, type Skill } from './skills';
import { disabledToolNames, onBashOutput, tools as builtinTools, type ToolSetName } from './tools';
@@ -94,7 +94,7 @@ export type SessionOptions = {
onNotebookChange?: (state: NotebookState) => void;
};
const estimateTokens = (messages: ModelMessage[]) => Math.round(JSON.stringify(messages).length / 4);
const estimateTokens = pruneEstimateTokens;
/** Estimated tokens at which the wire history is pruned. */
const DEFAULT_COMPACT_THRESHOLD = 120_000;
@@ -115,6 +115,30 @@ const isStaleItemError = (error: unknown): boolean =>
const STALE_ITEM_NOTICE =
'The provider no longer had part of this session stored. Re-sent the history inline and carried on.';
async function summarizeDiscarded(span: ModelMessage[], model: LanguageModel): Promise<string | undefined> {
if (span.length === 0) return undefined;
const excerpt = span
.map((m) => {
const c = typeof m.content === 'string' ? m.content : JSON.stringify(m.content).slice(0, 2000);
return `${m.role}: ${c}`;
})
.join('\n')
.slice(0, 6000);
if (!excerpt.trim()) return undefined;
try {
const { text } = await generateText({
model,
system: 'Summarize the dropped part of a long coding session so nothing important is lost. Keep: user goals, files touched with paths, decisions and why, tool results that matter, and what remains. One short handover note, 3-6 lines, no preamble.',
prompt: excerpt,
maxRetries: 1,
});
const t = text.trim();
return t.length > 0 ? t : undefined;
} catch {
return undefined;
}
}
type ApprovalContext = Pick<ApprovalRequest, 'matchedPattern' | 'suggestedPattern' | 'repeated'>;
export class Session {
@@ -279,11 +303,14 @@ export class Session {
}
private systemFor(): string {
const mem = this.opts.memory;
// Render project+global when available; fall back to project-only.
const memoryBlock = mem ? (typeof (mem as unknown as { renderWithGlobal?: (n?: number) => string }).renderWithGlobal === 'function' ? (mem as unknown as { renderWithGlobal: (n?: number) => string }).renderWithGlobal() || mem.render() : mem.render()) : '';
return systemPrompt({
cwd: this.opts.cwd ?? process.cwd(),
instructions: this.opts.instructions ?? [],
notebook: this.notebook.render(),
memory: this.opts.memory?.render() ?? '',
memory: memoryBlock,
skills: renderSkills(this.opts.skills ?? []),
agent: renderAgent(this.variant),
plugins: this.opts.plugins?.appendix ?? '',
@@ -439,6 +466,7 @@ export class Session {
// Each iteration is one model run. A run ends either finished, or suspended
// on tool approvals, in which case we collect decisions and run again.
let compactionReported = false;
let compactionSpan: ModelMessage[] | undefined;
while (true) {
const pending: ApprovalRequest[] = [];
const compactions: Extract<AgentEvent, { type: 'compacted' }>[] = [];
@@ -465,6 +493,10 @@ export class Session {
const instructions = this.systemFor();
if (estimateTokens(messages) <= threshold) return { instructions };
const pruned = pruneToFit({ messages, threshold, estimate: estimateTokens });
// Capture what was dropped by reference identity — the lossless note is built after the stream.
if (!compactionReported && pruned.length < messages.length) {
compactionSpan = droppedSpan(messages, pruned);
}
// prepareStep cannot yield, so queue the notice and drain it in the loop.
if (!compactionReported) {
compactions.push({ type: 'compacted', before: messages.length, after: pruned.length });
@@ -581,6 +613,16 @@ export class Session {
this.messages.push(...(await result.responseMessages));
this.opts.onChange?.(this.messages);
// Lossless compaction: summarize what the wire pruned so future turns keep it.
if (compactionSpan && compactionSpan.length > 0) {
const retained = await summarizeDiscarded(compactionSpan, this.model);
if (retained) {
this.messages.push({ role: 'user', content: `Note (retained from compacted history):\n${retained}` });
this.opts.onChange?.(this.messages);
}
compactionSpan = undefined;
}
if (pending.length === 0) {
const usage = await result.usage;
this.inputTokens += usage.inputTokens ?? 0;