Implement an LRUMemoryStorage for context
This commit is contained in:
parent
c102d4ba15
commit
48cdb8a1db
4 changed files with 153 additions and 2 deletions
13
docs/a2a.md
13
docs/a2a.md
|
|
@ -81,3 +81,16 @@ Each conversation is identified by a `context_id`. All messages within the same
|
||||||
- Remember what was discussed
|
- Remember what was discussed
|
||||||
- Track which documents were already found
|
- Track which documents were already found
|
||||||
- Provide contextual follow-up answers
|
- Provide contextual follow-up answers
|
||||||
|
|
||||||
|
### Memory Management
|
||||||
|
|
||||||
|
To prevent memory growth, the server uses LRU (Least Recently Used) eviction:
|
||||||
|
|
||||||
|
- Maximum 1000 contexts kept in memory (configurable via `A2A_MAX_CONTEXTS`)
|
||||||
|
- When limit exceeded, least recently used contexts are automatically evicted
|
||||||
|
- No periodic cleanup needed - eviction happens on-demand
|
||||||
|
|
||||||
|
Configure via environment variable:
|
||||||
|
```bash
|
||||||
|
export A2A_MAX_CONTEXTS=1000
|
||||||
|
```
|
||||||
|
|
|
||||||
|
|
@ -1,5 +1,6 @@
|
||||||
import logging
|
import logging
|
||||||
import uuid
|
import uuid
|
||||||
|
from collections import OrderedDict
|
||||||
from contextlib import asynccontextmanager
|
from contextlib import asynccontextmanager
|
||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
|
|
@ -24,9 +25,10 @@ try:
|
||||||
Message,
|
Message,
|
||||||
TaskIdParams,
|
TaskIdParams,
|
||||||
TaskSendParams,
|
TaskSendParams,
|
||||||
|
TaskState,
|
||||||
TextPart,
|
TextPart,
|
||||||
)
|
)
|
||||||
from fasta2a.storage import InMemoryStorage # type: ignore
|
from fasta2a.storage import InMemoryStorage, Storage # type: ignore
|
||||||
except ImportError as e:
|
except ImportError as e:
|
||||||
raise ImportError(
|
raise ImportError(
|
||||||
"A2A support requires the 'a2a' extra. "
|
"A2A support requires the 'a2a' extra. "
|
||||||
|
|
@ -148,6 +150,64 @@ def save_message_history(message_history: list[ModelMessage]) -> Message:
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
class LRUMemoryStorage(Storage[list["Message"]]): # type: ignore
|
||||||
|
"""Storage wrapper with LRU eviction for contexts.
|
||||||
|
|
||||||
|
Enforces a maximum context limit using LRU (Least Recently Used) eviction.
|
||||||
|
"""
|
||||||
|
|
||||||
|
def __init__(self, storage: InMemoryStorage, max_contexts: int):
|
||||||
|
self.storage = storage
|
||||||
|
self.max_contexts = max_contexts
|
||||||
|
# Track context access order (LRU cache)
|
||||||
|
self.context_order: OrderedDict[str, None] = OrderedDict()
|
||||||
|
|
||||||
|
async def load_context(self, context_id: str) -> list["Message"] | None:
|
||||||
|
"""Load context and update access order."""
|
||||||
|
result = await self.storage.load_context(context_id)
|
||||||
|
if result is not None:
|
||||||
|
# Move to end (most recently used)
|
||||||
|
self.context_order.pop(context_id, None)
|
||||||
|
self.context_order[context_id] = None
|
||||||
|
return result
|
||||||
|
|
||||||
|
async def update_context(self, context_id: str, context: list["Message"]) -> None:
|
||||||
|
"""Update context and enforce LRU limit."""
|
||||||
|
await self.storage.update_context(context_id, context)
|
||||||
|
# Move to end (most recently used)
|
||||||
|
self.context_order.pop(context_id, None)
|
||||||
|
self.context_order[context_id] = None
|
||||||
|
|
||||||
|
# Enforce max contexts limit (LRU eviction)
|
||||||
|
while len(self.context_order) > self.max_contexts:
|
||||||
|
# Remove oldest (first item in OrderedDict)
|
||||||
|
oldest_context_id = next(iter(self.context_order))
|
||||||
|
self.context_order.pop(oldest_context_id)
|
||||||
|
logger.debug(
|
||||||
|
f"Evicted context {oldest_context_id} (LRU, limit={self.max_contexts})"
|
||||||
|
)
|
||||||
|
|
||||||
|
async def load_task(self, task_id: str, history_length: int | None = None):
|
||||||
|
"""Delegate to underlying storage."""
|
||||||
|
return await self.storage.load_task(task_id, history_length)
|
||||||
|
|
||||||
|
async def update_task(
|
||||||
|
self,
|
||||||
|
task_id: str,
|
||||||
|
state: TaskState,
|
||||||
|
new_artifacts: list["Artifact"] | None = None,
|
||||||
|
new_messages: list["Message"] | None = None,
|
||||||
|
):
|
||||||
|
"""Delegate to underlying storage."""
|
||||||
|
return await self.storage.update_task(
|
||||||
|
task_id, state, new_artifacts, new_messages
|
||||||
|
)
|
||||||
|
|
||||||
|
async def submit_task(self, context_id: str, message: "Message"):
|
||||||
|
"""Delegate to underlying storage."""
|
||||||
|
return await self.storage.submit_task(context_id, message)
|
||||||
|
|
||||||
|
|
||||||
def extract_question_from_task(task_history: list[Message]) -> str | None:
|
def extract_question_from_task(task_history: list[Message]) -> str | None:
|
||||||
"""Extract the user's question from task history.
|
"""Extract the user's question from task history.
|
||||||
|
|
||||||
|
|
@ -176,7 +236,10 @@ def create_a2a_app(db_path: Path):
|
||||||
Returns:
|
Returns:
|
||||||
A FastA2A ASGI application
|
A FastA2A ASGI application
|
||||||
"""
|
"""
|
||||||
storage = InMemoryStorage()
|
base_storage = InMemoryStorage()
|
||||||
|
storage = LRUMemoryStorage(
|
||||||
|
storage=base_storage, max_contexts=Config.A2A_MAX_CONTEXTS
|
||||||
|
)
|
||||||
broker = InMemoryBroker()
|
broker = InMemoryBroker()
|
||||||
|
|
||||||
# Create the agent with native search tool
|
# Create the agent with native search tool
|
||||||
|
|
@ -341,6 +404,7 @@ def create_a2a_app(db_path: Path):
|
||||||
# Create FastA2A app with custom worker lifecycle
|
# Create FastA2A app with custom worker lifecycle
|
||||||
@asynccontextmanager
|
@asynccontextmanager
|
||||||
async def lifespan(app):
|
async def lifespan(app):
|
||||||
|
logger.info(f"Started A2A server (max contexts: {Config.A2A_MAX_CONTEXTS})")
|
||||||
async with app.task_manager:
|
async with app.task_manager:
|
||||||
async with worker.run():
|
async with worker.run():
|
||||||
yield
|
yield
|
||||||
|
|
|
||||||
|
|
@ -62,6 +62,10 @@ class AppConfig(BaseModel):
|
||||||
# to allow concurrent connections to safely use recent versions.
|
# to allow concurrent connections to safely use recent versions.
|
||||||
VACUUM_RETENTION_SECONDS: int = 60
|
VACUUM_RETENTION_SECONDS: int = 60
|
||||||
|
|
||||||
|
# Maximum number of A2A contexts to keep in memory. When exceeded, least
|
||||||
|
# recently used contexts will be evicted. Default is 1000.
|
||||||
|
A2A_MAX_CONTEXTS: int = 1000
|
||||||
|
|
||||||
@field_validator("MONITOR_DIRECTORIES", mode="before")
|
@field_validator("MONITOR_DIRECTORIES", mode="before")
|
||||||
@classmethod
|
@classmethod
|
||||||
def parse_monitor_directories(cls, v):
|
def parse_monitor_directories(cls, v):
|
||||||
|
|
|
||||||
|
|
@ -138,6 +138,76 @@ async def test_extract_question_from_task_no_text():
|
||||||
assert question is None
|
assert question is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_lru_memory_storage_lru_eviction():
|
||||||
|
"""Test that LRUMemoryStorage evicts least recently used contexts."""
|
||||||
|
from fasta2a.storage import InMemoryStorage
|
||||||
|
|
||||||
|
from haiku.rag.a2a import LRUMemoryStorage
|
||||||
|
|
||||||
|
base_storage = InMemoryStorage()
|
||||||
|
storage = LRUMemoryStorage(storage=base_storage, max_contexts=3)
|
||||||
|
|
||||||
|
# Add 3 contexts (at limit)
|
||||||
|
await storage.update_context("ctx1", [])
|
||||||
|
await storage.update_context("ctx2", [])
|
||||||
|
await storage.update_context("ctx3", [])
|
||||||
|
|
||||||
|
# All 3 should be tracked
|
||||||
|
assert len(storage.context_order) == 3
|
||||||
|
assert "ctx1" in storage.context_order
|
||||||
|
assert "ctx2" in storage.context_order
|
||||||
|
assert "ctx3" in storage.context_order
|
||||||
|
|
||||||
|
# Add 4th context - should evict ctx1 (oldest)
|
||||||
|
await storage.update_context("ctx4", [])
|
||||||
|
assert len(storage.context_order) == 3
|
||||||
|
assert "ctx1" not in storage.context_order
|
||||||
|
assert "ctx2" in storage.context_order
|
||||||
|
assert "ctx3" in storage.context_order
|
||||||
|
assert "ctx4" in storage.context_order
|
||||||
|
|
||||||
|
# Access ctx2 (moves it to end)
|
||||||
|
await storage.load_context("ctx2")
|
||||||
|
|
||||||
|
# Add 5th context - should evict ctx3 (now oldest since ctx2 was accessed)
|
||||||
|
await storage.update_context("ctx5", [])
|
||||||
|
assert len(storage.context_order) == 3
|
||||||
|
assert "ctx3" not in storage.context_order
|
||||||
|
assert "ctx2" in storage.context_order # Still present (was accessed)
|
||||||
|
assert "ctx4" in storage.context_order
|
||||||
|
assert "ctx5" in storage.context_order
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_lru_memory_storage_access_order():
|
||||||
|
"""Test that accessing contexts updates their order."""
|
||||||
|
from fasta2a.storage import InMemoryStorage
|
||||||
|
|
||||||
|
from haiku.rag.a2a import LRUMemoryStorage
|
||||||
|
|
||||||
|
base_storage = InMemoryStorage()
|
||||||
|
storage = LRUMemoryStorage(storage=base_storage, max_contexts=2)
|
||||||
|
|
||||||
|
# Add 2 contexts
|
||||||
|
await storage.update_context("ctx1", [])
|
||||||
|
await storage.update_context("ctx2", [])
|
||||||
|
|
||||||
|
# Order should be: ctx1, ctx2
|
||||||
|
assert list(storage.context_order.keys()) == ["ctx1", "ctx2"]
|
||||||
|
|
||||||
|
# Load ctx1 (moves to end)
|
||||||
|
await storage.load_context("ctx1")
|
||||||
|
# Order should be: ctx2, ctx1
|
||||||
|
assert list(storage.context_order.keys()) == ["ctx2", "ctx1"]
|
||||||
|
|
||||||
|
# Add ctx3 - should evict ctx2 (oldest)
|
||||||
|
await storage.update_context("ctx3", [])
|
||||||
|
assert "ctx2" not in storage.context_order
|
||||||
|
assert "ctx1" in storage.context_order
|
||||||
|
assert "ctx3" in storage.context_order
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_a2a_app_creation(temp_db_path):
|
async def test_a2a_app_creation(temp_db_path):
|
||||||
"""Test that A2A app can be created successfully."""
|
"""Test that A2A app can be created successfully."""
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue