mirror of
https://github.com/strands-agents/harness-sdk.git
synced 2026-10-02 02:44:48 +08:00
feat(context-offloader): add turn-based eviction to InMemoryStorage (#2648)
This commit is contained in:
@@ -36,12 +36,12 @@ import json
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ...hooks.events import AfterToolCallEvent
|
||||
from ...hooks.events import AfterToolCallEvent, BeforeModelCallEvent
|
||||
from ...plugins import Plugin, hook
|
||||
from ...tools.decorator import tool
|
||||
from ...types.content import Message
|
||||
from ...types.tools import ToolContext, ToolResult, ToolResultContent
|
||||
from .storage import Storage
|
||||
from .storage import InMemoryStorage, Storage
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ...agent.agent import Agent
|
||||
@@ -140,11 +140,19 @@ class ContextOffloader(Plugin):
|
||||
super().__init__()
|
||||
|
||||
def init_agent(self, agent: Agent) -> None:
|
||||
"""Conditionally register the retrieval tool."""
|
||||
"""Conditionally register the retrieval tool and bind storage."""
|
||||
if isinstance(self._storage, InMemoryStorage):
|
||||
self._storage._bind(id(agent))
|
||||
if not self._include_retrieval_tool:
|
||||
# Remove the auto-discovered retrieval tool
|
||||
self._tools = [t for t in self._tools if t.tool_name != "retrieve_offloaded_content"]
|
||||
|
||||
@hook
|
||||
def _on_before_model_call(self, event: BeforeModelCallEvent) -> None:
|
||||
"""Trigger eviction of stale entries based on the agent's cycle count."""
|
||||
if isinstance(self._storage, InMemoryStorage):
|
||||
self._storage._evict(event.agent.event_loop_metrics.cycle_count)
|
||||
|
||||
@tool(context=True)
|
||||
def retrieve_offloaded_content(
|
||||
self,
|
||||
|
||||
@@ -26,6 +26,7 @@ Example:
|
||||
"""
|
||||
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import threading
|
||||
import time
|
||||
@@ -36,6 +37,8 @@ import boto3
|
||||
from botocore.config import Config as BotocoreConfig
|
||||
from botocore.exceptions import ClientError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _sanitize_id(raw_id: str) -> str:
|
||||
"""Sanitize an ID for safe use in filenames and object keys.
|
||||
@@ -213,17 +216,47 @@ class InMemoryStorage:
|
||||
Useful for testing and serverless environments where disk access
|
||||
is not available or not desired. Thread-safe.
|
||||
|
||||
Supports turn-based eviction: entries not accessed (stored or retrieved)
|
||||
within ``evict_after_turns`` agent loop cycles are automatically removed.
|
||||
The ``ContextOffloader`` plugin triggers eviction on each model invocation
|
||||
cycle. Eviction is enabled by default (20 cycles). Pass ``None`` to disable.
|
||||
|
||||
Note:
|
||||
Content accumulates for the lifetime of this instance. For long-running
|
||||
agents, consider creating a new instance per session or switching to
|
||||
``FileStorage`` or ``S3Storage`` for persistent storage with external
|
||||
lifecycle management.
|
||||
Content does not survive process restarts. For multi-session
|
||||
persistence, use ``FileStorage`` or ``S3Storage``. Each agent should
|
||||
use its own ``InMemoryStorage`` instance — sharing one across multiple
|
||||
agents is not supported when eviction is enabled.
|
||||
|
||||
Evicted entries are permanently deleted from memory. The agent will
|
||||
receive an error if it attempts to retrieve evicted content. The
|
||||
original tool result is not preserved in the conversation history
|
||||
after offloading — only the preview and references remain in context.
|
||||
|
||||
Args:
|
||||
evict_after_turns: Number of cycles of inactivity before an entry is
|
||||
evicted. Defaults to 20. ``None`` disables eviction.
|
||||
"""
|
||||
|
||||
def __init__(self) -> None:
|
||||
"""Initialize in-memory storage."""
|
||||
self._store: dict[str, tuple[bytes, str]] = {}
|
||||
_DEFAULT_EVICT_AFTER_TURNS = 20
|
||||
|
||||
def __init__(self, evict_after_turns: int | None = _DEFAULT_EVICT_AFTER_TURNS) -> None:
|
||||
"""Initialize in-memory storage.
|
||||
|
||||
Args:
|
||||
evict_after_turns: Number of cycles of inactivity before an entry is
|
||||
evicted. Defaults to 20. ``None`` disables eviction.
|
||||
|
||||
Raises:
|
||||
ValueError: If evict_after_turns is not a positive integer.
|
||||
"""
|
||||
if evict_after_turns is not None and evict_after_turns < 1:
|
||||
raise ValueError("evict_after_turns must be a positive integer")
|
||||
|
||||
self._store: dict[str, tuple[bytes, str, int]] = {}
|
||||
self._counter: int = 0
|
||||
self._current_cycle: int = 0
|
||||
self._evict_after_turns: int | None = evict_after_turns
|
||||
self._bound_agent_id: int | None = None
|
||||
self._lock = threading.Lock()
|
||||
|
||||
def store(self, key: str, content: bytes, content_type: str = "text/plain") -> str:
|
||||
@@ -240,12 +273,15 @@ class InMemoryStorage:
|
||||
with self._lock:
|
||||
self._counter += 1
|
||||
reference = f"mem_{self._counter}_{key}"
|
||||
self._store[reference] = (content, content_type)
|
||||
self._store[reference] = (content, content_type, self._current_cycle)
|
||||
return reference
|
||||
|
||||
def retrieve(self, reference: str) -> tuple[bytes, str]:
|
||||
"""Retrieve content from memory.
|
||||
|
||||
Refreshes the last-accessed turn so the entry stays alive longer
|
||||
when eviction is enabled.
|
||||
|
||||
Args:
|
||||
reference: The reference returned by store().
|
||||
|
||||
@@ -253,12 +289,52 @@ class InMemoryStorage:
|
||||
A tuple of (content bytes, content type).
|
||||
|
||||
Raises:
|
||||
KeyError: If the reference is not found.
|
||||
KeyError: If the reference is not found (or was evicted).
|
||||
"""
|
||||
with self._lock:
|
||||
if reference not in self._store:
|
||||
raise KeyError(f"Reference not found: {reference}")
|
||||
return self._store[reference]
|
||||
content, content_type, _ = self._store[reference]
|
||||
self._store[reference] = (content, content_type, self._current_cycle)
|
||||
return content, content_type
|
||||
|
||||
def _bind(self, agent_id: int) -> None:
|
||||
"""Claim this storage for a single agent.
|
||||
|
||||
Raises:
|
||||
ValueError: If already bound to a different agent.
|
||||
"""
|
||||
with self._lock:
|
||||
if self._bound_agent_id is None:
|
||||
self._bound_agent_id = agent_id
|
||||
elif self._bound_agent_id != agent_id:
|
||||
raise ValueError(
|
||||
"InMemoryStorage cannot be shared across multiple agents. "
|
||||
"Use a separate InMemoryStorage instance per agent."
|
||||
)
|
||||
|
||||
def _evict(self, cycle: int) -> None:
|
||||
"""Update current cycle and evict stale entries.
|
||||
|
||||
Called by the ContextOffloader plugin on each ``BeforeModelCallEvent``.
|
||||
Entries whose last-accessed cycle is more than ``evict_after_turns``
|
||||
behind the current cycle are removed.
|
||||
|
||||
Args:
|
||||
cycle: The agent's current event loop cycle count.
|
||||
"""
|
||||
with self._lock:
|
||||
self._current_cycle = cycle
|
||||
if self._evict_after_turns is None:
|
||||
return
|
||||
threshold = cycle - self._evict_after_turns
|
||||
stale_refs = [
|
||||
ref for ref, (_, _, last_cycle) in self._store.items() if last_cycle < threshold
|
||||
]
|
||||
for ref in stale_refs:
|
||||
del self._store[ref]
|
||||
if stale_refs:
|
||||
logger.debug("evicted=<%d>, cycle=<%d> | stale entries removed", len(stale_refs), cycle)
|
||||
|
||||
def clear(self) -> None:
|
||||
"""Remove all stored content.
|
||||
|
||||
@@ -7,7 +7,7 @@ from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from strands.hooks.events import AfterToolCallEvent
|
||||
from strands.hooks.events import AfterToolCallEvent, BeforeModelCallEvent
|
||||
from strands.types.tools import ToolContext, ToolUse
|
||||
from strands.vended_plugins.context_offloader import (
|
||||
ContextOffloader,
|
||||
@@ -84,8 +84,10 @@ class TestContextOffloader:
|
||||
assert plugin.name == "context_offloader"
|
||||
|
||||
def test_hooks_auto_discovered(self, plugin):
|
||||
assert len(plugin.hooks) == 1
|
||||
assert plugin.hooks[0].__name__ == "_handle_tool_result"
|
||||
assert len(plugin.hooks) == 2
|
||||
hook_names = {h.__name__ for h in plugin.hooks}
|
||||
assert "_handle_tool_result" in hook_names
|
||||
assert "_on_before_model_call" in hook_names
|
||||
|
||||
def test_raises_on_non_positive_max_result_tokens(self):
|
||||
with pytest.raises(ValueError, match="max_result_tokens must be positive"):
|
||||
@@ -589,3 +591,37 @@ class TestActionableReferences:
|
||||
|
||||
result_text = event.result["content"][0]["text"]
|
||||
assert "mem_" in result_text
|
||||
|
||||
|
||||
class TestBeforeModelCallHook:
|
||||
@staticmethod
|
||||
def _make_event(cycle_count):
|
||||
agent = MagicMock()
|
||||
agent.event_loop_metrics.cycle_count = cycle_count
|
||||
return BeforeModelCallEvent(agent=agent, invocation_state={})
|
||||
|
||||
def test_calls_evict_with_cycle_count(self):
|
||||
storage = InMemoryStorage(evict_after_turns=5)
|
||||
plugin = ContextOffloader(storage=storage, max_result_tokens=25, preview_tokens=10)
|
||||
|
||||
plugin._on_before_model_call(self._make_event(7))
|
||||
|
||||
assert storage._current_cycle == 7
|
||||
|
||||
def test_does_not_crash_on_storage_without_evict(self):
|
||||
storage = MagicMock(spec=["store", "retrieve"])
|
||||
plugin = ContextOffloader(storage=storage, max_result_tokens=25, preview_tokens=10)
|
||||
|
||||
plugin._on_before_model_call(self._make_event(1))
|
||||
|
||||
def test_eviction_triggered_via_hook(self):
|
||||
storage = InMemoryStorage(evict_after_turns=2)
|
||||
plugin = ContextOffloader(storage=storage, max_result_tokens=25, preview_tokens=10)
|
||||
|
||||
ref = storage.store("key_1", b"content")
|
||||
|
||||
# stored at cycle 0, evict at cycle 3: threshold = 3 - 2 = 1, 0 < 1 → evicted
|
||||
plugin._on_before_model_call(self._make_event(3))
|
||||
with pytest.raises(KeyError):
|
||||
storage.retrieve(ref)
|
||||
|
||||
|
||||
@@ -94,6 +94,125 @@ class TestInMemoryStorage:
|
||||
storage.clear()
|
||||
|
||||
|
||||
class TestInMemoryStorageEviction:
|
||||
def test_evict_after_turns_validation(self):
|
||||
with pytest.raises(ValueError, match="evict_after_turns must be a positive integer"):
|
||||
InMemoryStorage(evict_after_turns=0)
|
||||
with pytest.raises(ValueError, match="evict_after_turns must be a positive integer"):
|
||||
InMemoryStorage(evict_after_turns=-1)
|
||||
|
||||
|
||||
def test_eviction_enabled_by_default(self):
|
||||
storage = InMemoryStorage()
|
||||
assert storage._evict_after_turns == 20
|
||||
ref = storage.store("key_1", b"content")
|
||||
for _ in range(21):
|
||||
storage._evict(storage._current_cycle + 1)
|
||||
with pytest.raises(KeyError):
|
||||
storage.retrieve(ref)
|
||||
|
||||
def test_eviction_disabled_with_none(self):
|
||||
storage = InMemoryStorage(evict_after_turns=None)
|
||||
ref = storage.store("key_1", b"content")
|
||||
for _ in range(100):
|
||||
storage._evict(storage._current_cycle + 1)
|
||||
assert storage.retrieve(ref) == (b"content", "text/plain")
|
||||
|
||||
def test_entry_evicted_after_n_turns(self):
|
||||
storage = InMemoryStorage(evict_after_turns=3)
|
||||
ref = storage.store("key_1", b"content")
|
||||
|
||||
storage._evict(storage._current_cycle + 1) # turn 1
|
||||
storage._evict(storage._current_cycle + 1) # turn 2
|
||||
storage._evict(storage._current_cycle + 1) # turn 3 — threshold = 3 - 3 = 0, stored at 0, 0 < 0 is false
|
||||
assert storage.retrieve(ref) == (b"content", "text/plain")
|
||||
|
||||
storage._evict(storage._current_cycle + 1) # turn 4 — refreshed by retrieve to 3, 3 < 1 is false
|
||||
assert storage.retrieve(ref) == (b"content", "text/plain")
|
||||
|
||||
def test_entry_evicted_without_access(self):
|
||||
storage = InMemoryStorage(evict_after_turns=2)
|
||||
ref = storage.store("key_1", b"content")
|
||||
|
||||
storage._evict(storage._current_cycle + 1) # turn 1
|
||||
storage._evict(storage._current_cycle + 1) # turn 2 — threshold = 2 - 2 = 0, stored at 0, 0 < 0 is false
|
||||
# Entry still alive at exact boundary
|
||||
content, _ = storage.retrieve(ref)
|
||||
assert content == b"content"
|
||||
|
||||
def test_entry_evicted_past_boundary(self):
|
||||
storage = InMemoryStorage(evict_after_turns=2)
|
||||
ref = storage.store("key_1", b"content")
|
||||
|
||||
storage._evict(storage._current_cycle + 1) # turn 1
|
||||
storage._evict(storage._current_cycle + 1) # turn 2
|
||||
storage._evict(storage._current_cycle + 1) # turn 3 — threshold = 1, 0 < 1 → evicted
|
||||
with pytest.raises(KeyError):
|
||||
storage.retrieve(ref)
|
||||
|
||||
def test_retrieve_refreshes_last_accessed(self):
|
||||
storage = InMemoryStorage(evict_after_turns=2)
|
||||
ref = storage.store("key_1", b"content")
|
||||
|
||||
storage._evict(storage._current_cycle + 1) # turn 1
|
||||
storage.retrieve(ref) # refreshes last_accessed to turn 1
|
||||
|
||||
storage._evict(storage._current_cycle + 1) # turn 2
|
||||
storage._evict(storage._current_cycle + 1) # turn 3 — threshold = 1, last_accessed = 1, 1 < 1 is false
|
||||
assert storage.retrieve(ref) == (b"content", "text/plain")
|
||||
|
||||
def test_multiple_entries_evicted_independently(self):
|
||||
storage = InMemoryStorage(evict_after_turns=2)
|
||||
ref1 = storage.store("key_1", b"first")
|
||||
|
||||
storage._evict(storage._current_cycle + 1) # turn 1
|
||||
ref2 = storage.store("key_2", b"second")
|
||||
|
||||
storage._evict(storage._current_cycle + 1) # turn 2
|
||||
storage._evict(storage._current_cycle + 1) # turn 3 — threshold = 1. ref1 evicted, ref2 survives
|
||||
with pytest.raises(KeyError):
|
||||
storage.retrieve(ref1)
|
||||
assert storage.retrieve(ref2) == (b"second", "text/plain")
|
||||
|
||||
def test_evict_updates_current_cycle(self):
|
||||
storage = InMemoryStorage()
|
||||
assert storage._current_cycle == 0
|
||||
storage._evict(5)
|
||||
assert storage._current_cycle == 5
|
||||
|
||||
def test_rejects_shared_storage_across_agents(self):
|
||||
storage = InMemoryStorage()
|
||||
storage._bind(1)
|
||||
with pytest.raises(ValueError, match="cannot be shared"):
|
||||
storage._bind(2)
|
||||
|
||||
def test_allows_same_agent_repeated_bind(self):
|
||||
storage = InMemoryStorage()
|
||||
storage._bind(42)
|
||||
storage._bind(42)
|
||||
|
||||
def test_thread_safety_with_eviction(self):
|
||||
storage = InMemoryStorage(evict_after_turns=5)
|
||||
refs: list[str] = []
|
||||
errors: list[Exception] = []
|
||||
|
||||
def store_and_tick(i: int):
|
||||
try:
|
||||
ref = storage.store(f"key_{i}", f"content_{i}".encode())
|
||||
refs.append(ref)
|
||||
storage._evict(storage._current_cycle + 1)
|
||||
except Exception as e:
|
||||
errors.append(e)
|
||||
|
||||
threads = [threading.Thread(target=store_and_tick, args=(i,)) for i in range(50)]
|
||||
for t in threads:
|
||||
t.start()
|
||||
for t in threads:
|
||||
t.join()
|
||||
|
||||
assert not errors
|
||||
|
||||
|
||||
class TestFileStorage:
|
||||
def test_round_trip(self, tmp_path):
|
||||
storage = FileStorage(artifact_dir=str(tmp_path / "artifacts"))
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import { describe, it, expect, vi } from 'vitest'
|
||||
import { ContextOffloader } from '../plugin.js'
|
||||
import { InMemoryStorage } from '../storage.js'
|
||||
import { AfterToolCallEvent } from '../../../hooks/events.js'
|
||||
import { AfterToolCallEvent, BeforeModelCallEvent } from '../../../hooks/events.js'
|
||||
import { TextBlock, JsonBlock, ToolResultBlock } from '../../../types/messages.js'
|
||||
import { ImageBlock, VideoBlock, DocumentBlock } from '../../../types/media.js'
|
||||
import { createMockAgent, invokeTrackedHook } from '../../../__fixtures__/agent-helpers.js'
|
||||
@@ -61,12 +61,14 @@ describe('ContextOffloader', () => {
|
||||
expect(plugin.name).toBe('strands:context-offloader')
|
||||
})
|
||||
|
||||
it('registers AfterToolCallEvent hook', () => {
|
||||
it('registers hooks', () => {
|
||||
const plugin = new ContextOffloader({ storage: new InMemoryStorage() })
|
||||
const agent = createMockAgent()
|
||||
plugin.initAgent(agent)
|
||||
expect(agent.trackedHooks).toHaveLength(1)
|
||||
expect(agent.trackedHooks[0]!.eventType).toBe(AfterToolCallEvent)
|
||||
expect(agent.trackedHooks).toHaveLength(2)
|
||||
const eventTypes = agent.trackedHooks.map((h) => h.eventType)
|
||||
expect(eventTypes).toContain(AfterToolCallEvent)
|
||||
expect(eventTypes).toContain(BeforeModelCallEvent)
|
||||
})
|
||||
|
||||
it('returns retrieval tool by default', () => {
|
||||
@@ -608,4 +610,53 @@ describe('ContextOffloader', () => {
|
||||
expect(result).not.toContain('line 2')
|
||||
})
|
||||
})
|
||||
|
||||
describe('eviction via BeforeModelCallEvent', () => {
|
||||
it('calls _evict on storage with incrementing cycle count', () => {
|
||||
const storage = new InMemoryStorage(5)
|
||||
const plugin = new ContextOffloader({ storage, maxResultTokens: 10, previewTokens: 5 })
|
||||
const agent = createMockAgent()
|
||||
plugin.initAgent(agent)
|
||||
|
||||
const hook = agent.trackedHooks.find((h) => h.eventType === BeforeModelCallEvent)!
|
||||
const event = new BeforeModelCallEvent({ agent, model: mockModel, invocationState: {} })
|
||||
|
||||
hook.callback(event)
|
||||
expect((storage as unknown as { _currentCycle: number })._currentCycle).toBe(1)
|
||||
|
||||
hook.callback(event)
|
||||
expect((storage as unknown as { _currentCycle: number })._currentCycle).toBe(2)
|
||||
})
|
||||
|
||||
it('evicts stale entries on BeforeModelCallEvent', async () => {
|
||||
const storage = new InMemoryStorage(2)
|
||||
const plugin = new ContextOffloader({ storage, maxResultTokens: 10, previewTokens: 5 })
|
||||
const agent = createMockAgent()
|
||||
plugin.initAgent(agent)
|
||||
|
||||
const ref = await storage.store('key1', new TextEncoder().encode('test'))
|
||||
|
||||
const hook = agent.trackedHooks.find((h) => h.eventType === BeforeModelCallEvent)!
|
||||
const event = new BeforeModelCallEvent({ agent, model: mockModel, invocationState: {} })
|
||||
|
||||
// Advance 3 cycles: threshold = 3 - 2 = 1, stored at 0, 0 < 1 → evicted
|
||||
hook.callback(event)
|
||||
hook.callback(event)
|
||||
hook.callback(event)
|
||||
|
||||
await expect(storage.retrieve(ref)).rejects.toThrow('Reference not found')
|
||||
})
|
||||
|
||||
it('does not crash when storage lacks _evict', () => {
|
||||
const storage = { store: vi.fn(), retrieve: vi.fn() }
|
||||
const plugin = new ContextOffloader({ storage, maxResultTokens: 10, previewTokens: 5 })
|
||||
const agent = createMockAgent()
|
||||
plugin.initAgent(agent)
|
||||
|
||||
const hook = agent.trackedHooks.find((h) => h.eventType === BeforeModelCallEvent)!
|
||||
const event = new BeforeModelCallEvent({ agent, model: mockModel, invocationState: {} })
|
||||
|
||||
expect(() => hook.callback(event)).not.toThrow()
|
||||
})
|
||||
})
|
||||
})
|
||||
|
||||
@@ -56,6 +56,83 @@ describe('InMemoryStorage', () => {
|
||||
})
|
||||
})
|
||||
|
||||
describe('InMemoryStorage eviction', () => {
|
||||
it('eviction enabled by default (20 cycles)', async () => {
|
||||
const storage = new InMemoryStorage()
|
||||
const ref = await storage.store('key1', new TextEncoder().encode('test'))
|
||||
storage._evict(21)
|
||||
await expect(storage.retrieve(ref)).rejects.toThrow('Reference not found')
|
||||
})
|
||||
|
||||
it('eviction disabled with null', async () => {
|
||||
const storage = new InMemoryStorage(null)
|
||||
const ref = await storage.store('key1', new TextEncoder().encode('test'))
|
||||
storage._evict(100)
|
||||
const result = await storage.retrieve(ref)
|
||||
expect(new TextDecoder().decode(result.content)).toBe('test')
|
||||
})
|
||||
|
||||
it('throws on invalid evictAfterTurns', () => {
|
||||
expect(() => new InMemoryStorage(0)).toThrow('evictAfterTurns must be a positive integer')
|
||||
expect(() => new InMemoryStorage(-1)).toThrow('evictAfterTurns must be a positive integer')
|
||||
})
|
||||
|
||||
it('entry survives within TTL', async () => {
|
||||
const storage = new InMemoryStorage(3)
|
||||
const ref = await storage.store('key1', new TextEncoder().encode('test'))
|
||||
// stored at cycle 0, evict at cycle 3: threshold = 3 - 3 = 0, 0 < 0 is false
|
||||
storage._evict(3)
|
||||
const result = await storage.retrieve(ref)
|
||||
expect(new TextDecoder().decode(result.content)).toBe('test')
|
||||
})
|
||||
|
||||
it('entry evicted past TTL', async () => {
|
||||
const storage = new InMemoryStorage(2)
|
||||
const ref = await storage.store('key1', new TextEncoder().encode('test'))
|
||||
// stored at cycle 0, evict at cycle 3: threshold = 3 - 2 = 1, 0 < 1 → evicted
|
||||
storage._evict(3)
|
||||
await expect(storage.retrieve(ref)).rejects.toThrow('Reference not found')
|
||||
})
|
||||
|
||||
it('retrieve refreshes last accessed cycle', async () => {
|
||||
const storage = new InMemoryStorage(2)
|
||||
const ref = await storage.store('key1', new TextEncoder().encode('test'))
|
||||
storage._evict(1)
|
||||
await storage.retrieve(ref) // refreshes to cycle 1
|
||||
// threshold at cycle 3 = 3 - 2 = 1, last_accessed = 1, 1 < 1 is false
|
||||
storage._evict(3)
|
||||
const result = await storage.retrieve(ref)
|
||||
expect(new TextDecoder().decode(result.content)).toBe('test')
|
||||
})
|
||||
|
||||
it('multiple entries evicted independently', async () => {
|
||||
const storage = new InMemoryStorage(2)
|
||||
const ref1 = await storage.store('key1', new TextEncoder().encode('first'))
|
||||
storage._evict(1)
|
||||
const ref2 = await storage.store('key2', new TextEncoder().encode('second'))
|
||||
// threshold at cycle 3 = 3 - 2 = 1. ref1 at 0 (evicted), ref2 at 1 (survives)
|
||||
storage._evict(3)
|
||||
await expect(storage.retrieve(ref1)).rejects.toThrow('Reference not found')
|
||||
const result = await storage.retrieve(ref2)
|
||||
expect(new TextDecoder().decode(result.content)).toBe('second')
|
||||
})
|
||||
|
||||
it('rejects shared storage across agents', () => {
|
||||
const storage = new InMemoryStorage()
|
||||
const agentA = {}
|
||||
const agentB = {}
|
||||
storage._bind(agentA)
|
||||
expect(() => storage._bind(agentB)).toThrow('cannot be shared')
|
||||
})
|
||||
|
||||
it('allows same agent repeated bind', () => {
|
||||
const storage = new InMemoryStorage()
|
||||
const agent = {}
|
||||
storage._bind(agent)
|
||||
storage._bind(agent)
|
||||
})
|
||||
})
|
||||
|
||||
describe('S3Storage', () => {
|
||||
let mockSend: ReturnType<typeof vi.fn>
|
||||
let mockS3Client: { send: ReturnType<typeof vi.fn> }
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import type { Plugin } from '../../plugins/plugin.js'
|
||||
import type { Tool, ToolContext } from '../../tools/tool.js'
|
||||
import type { LocalAgent } from '../../types/agent.js'
|
||||
import { AfterToolCallEvent } from '../../hooks/events.js'
|
||||
import { AfterToolCallEvent, BeforeModelCallEvent } from '../../hooks/events.js'
|
||||
import { TextBlock, JsonBlock, ToolResultBlock, Message } from '../../types/messages.js'
|
||||
import type { ToolResultContent } from '../../types/messages.js'
|
||||
import { ImageBlock, VideoBlock, DocumentBlock } from '../../types/media.js'
|
||||
@@ -10,7 +10,7 @@ import { tool } from '../../tools/tool-factory.js'
|
||||
import { z } from 'zod'
|
||||
import { logger } from '../../logging/logger.js'
|
||||
import type { JSONValue } from '../../types/json.js'
|
||||
import { FileStorage, type Storage } from './storage.js'
|
||||
import { FileStorage, InMemoryStorage, type Storage } from './storage.js'
|
||||
import { isSearchableContent, searchContent } from './search.js'
|
||||
|
||||
const CHARS_PER_TOKEN = 4
|
||||
@@ -155,8 +155,18 @@ export class ContextOffloader implements Plugin {
|
||||
}
|
||||
|
||||
initAgent(agent: LocalAgent): void {
|
||||
if (this._storage instanceof InMemoryStorage) {
|
||||
this._storage._bind(agent)
|
||||
}
|
||||
this._storageForAgent(agent)
|
||||
agent.addHook(AfterToolCallEvent, (event) => this._handleToolResult(event))
|
||||
let cycleCount = 0
|
||||
agent.addHook(BeforeModelCallEvent, () => {
|
||||
cycleCount++
|
||||
if (this._storage instanceof InMemoryStorage) {
|
||||
this._storage._evict(cycleCount)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
getTools(): Tool[] {
|
||||
|
||||
@@ -7,6 +7,7 @@
|
||||
*/
|
||||
|
||||
import type { Sandbox } from '../../sandbox/base.js'
|
||||
import { logger } from '../../logging/logger.js'
|
||||
|
||||
/**
|
||||
* Backend for storing and retrieving offloaded content blocks.
|
||||
@@ -47,17 +48,41 @@ function sanitizeId(rawId: string): string {
|
||||
* In-memory storage backend.
|
||||
*
|
||||
* Useful for testing and serverless environments where disk access is not available.
|
||||
* Content accumulates for the lifetime of this instance; call {@link clear} to free memory.
|
||||
* Supports turn-based eviction: entries not accessed (stored or retrieved) within
|
||||
* `evictAfterTurns` agent loop cycles are automatically removed when the plugin
|
||||
* calls `_evict`. Eviction is enabled by default (20 cycles). Pass `null` to disable.
|
||||
*
|
||||
* Note: content does not survive process restarts. For multi-session persistence,
|
||||
* use {@link FileStorage} or {@link S3Storage}. Each agent should use its own
|
||||
* `InMemoryStorage` instance — sharing one across multiple agents is not supported
|
||||
* when eviction is enabled.
|
||||
*
|
||||
* Evicted entries are permanently deleted from memory. The agent will receive
|
||||
* an error if it attempts to retrieve evicted content.
|
||||
*
|
||||
* @param evictAfterTurns - Cycles of inactivity before eviction. Defaults to 20. `null` disables.
|
||||
*/
|
||||
export class InMemoryStorage implements Storage {
|
||||
private _store = new Map<string, { content: Uint8Array; contentType: string }>()
|
||||
private _store = new Map<string, { content: Uint8Array; contentType: string; lastAccessedCycle: number }>()
|
||||
private _counter = 0
|
||||
private _currentCycle = 0
|
||||
private readonly _evictAfterTurns: number | null
|
||||
private _boundAgent: WeakRef<object> | null = null
|
||||
|
||||
static readonly DEFAULT_EVICT_AFTER_TURNS = 20
|
||||
|
||||
constructor(evictAfterTurns: number | null = InMemoryStorage.DEFAULT_EVICT_AFTER_TURNS) {
|
||||
if (evictAfterTurns !== null && evictAfterTurns < 1) {
|
||||
throw new Error('evictAfterTurns must be a positive integer')
|
||||
}
|
||||
this._evictAfterTurns = evictAfterTurns
|
||||
}
|
||||
|
||||
/** {@inheritdoc} */
|
||||
async store(key: string, content: Uint8Array, contentType: string = 'text/plain'): Promise<string> {
|
||||
this._counter++
|
||||
const reference = `mem_${this._counter}_${key}`
|
||||
this._store.set(reference, { content, contentType })
|
||||
this._store.set(reference, { content, contentType, lastAccessedCycle: this._currentCycle })
|
||||
return reference
|
||||
}
|
||||
|
||||
@@ -67,7 +92,44 @@ export class InMemoryStorage implements Storage {
|
||||
if (!entry) {
|
||||
throw new Error(`Reference not found: ${reference}`)
|
||||
}
|
||||
return entry
|
||||
entry.lastAccessedCycle = this._currentCycle
|
||||
return { content: entry.content, contentType: entry.contentType }
|
||||
}
|
||||
|
||||
/**
|
||||
* Claim this storage for a single agent. Throws if already bound to a different agent.
|
||||
* @internal
|
||||
*/
|
||||
_bind(agent: object): void {
|
||||
if (this._boundAgent === null) {
|
||||
this._boundAgent = new WeakRef(agent)
|
||||
} else if (this._boundAgent.deref() !== agent) {
|
||||
throw new Error(
|
||||
'InMemoryStorage cannot be shared across multiple agents. ' +
|
||||
'Use a separate InMemoryStorage instance per agent.'
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
/**
|
||||
* Update current cycle and evict stale entries.
|
||||
* Called by the ContextOffloader plugin on each BeforeModelCallEvent.
|
||||
* @internal
|
||||
*/
|
||||
_evict(cycle: number): void {
|
||||
this._currentCycle = cycle
|
||||
if (this._evictAfterTurns === null) return
|
||||
const threshold = cycle - this._evictAfterTurns
|
||||
let evicted = 0
|
||||
for (const [ref, entry] of this._store) {
|
||||
if (entry.lastAccessedCycle < threshold) {
|
||||
this._store.delete(ref)
|
||||
evicted++
|
||||
}
|
||||
}
|
||||
if (evicted > 0) {
|
||||
logger.debug(`evicted=<${evicted}>, cycle=<${cycle}> | stale entries removed`)
|
||||
}
|
||||
}
|
||||
|
||||
/** Remove all stored content. */
|
||||
|
||||
Reference in New Issue
Block a user