diff --git a/agentlightning/store/memory.py b/agentlightning/store/memory.py index 2c64f2822..e14226af8 100644 --- a/agentlightning/store/memory.py +++ b/agentlightning/store/memory.py @@ -6,13 +6,30 @@ import functools import hashlib import logging +import sys import threading import time import uuid from collections import deque -from typing import Any, Callable, Counter, Dict, List, Literal, Optional, Sequence, TypeVar, cast +from collections.abc import Iterable +from collections.abc import Mapping as MappingABC +from typing import ( + Any, + Callable, + Counter, + Dict, + List, + Literal, + Mapping, + Optional, + Sequence, + Set, + TypeVar, + cast, +) from opentelemetry.sdk.trace import ReadableSpan +from pydantic import BaseModel from agentlightning.types import ( Attempt, @@ -35,6 +52,21 @@ logger = logging.getLogger(__name__) +def estimate_model_size(obj: Any) -> int: + """Rough recursive size estimate for Pydantic BaseModel instances.""" + + if isinstance(obj, BaseModel): + values = cast(Iterable[Any], obj.__dict__.values()) + return sum(estimate_model_size(value) for value in values) + sys.getsizeof(cast(object, obj)) + if isinstance(obj, MappingABC): + mapping = cast(Mapping[Any, Any], obj) + return sum(estimate_model_size(value) for value in mapping.values()) + sys.getsizeof(cast(object, obj)) + if isinstance(obj, (list, tuple, set)): + iterable = cast(Iterable[Any], obj) + return sum(estimate_model_size(value) for value in iterable) + sys.getsizeof(cast(object, obj)) + return sys.getsizeof(cast(object, obj)) + + def _healthcheck_wrapper(func: T_callable) -> T_callable: """ Decorator to run the watchdog healthcheck **before** executing the decorated method. @@ -82,6 +114,19 @@ def _generate_attempt_id() -> str: return "at-" + short_id +def _detect_total_memory_bytes() -> int: + """Best-effort detection of the total available system memory in bytes.""" + + try: + import psutil + + return int(psutil.virtual_memory().total) + except ImportError: + # Fallback to 8GB if memory cannot be detected. + logger.error("psutil is not installed. Falling back to 8GB of memory in total.") + return 8 * 1024**3 + + class InMemoryLightningStore(LightningStore): """ In-memory implementation of LightningStore using Python data structures. @@ -89,9 +134,23 @@ class InMemoryLightningStore(LightningStore): The methods in this class should generally not call each other, especially those that are locked. + + Args: + eviction_memory_threshold: The threshold for evicting spans in bytes. + By default, it's 70% of the total VRAM available. + safe_memory_threshold: The threshold for safe memory usage in bytes. + By default, it's 80% of the eviction threshold. + span_size_estimator: A function to estimate the size of a span in bytes. + By default, it's a simple size estimator that uses sys.getsizeof. """ - def __init__(self): + def __init__( + self, + *, + eviction_memory_threshold: float | int | None = None, + safe_memory_threshold: float | int | None = None, + span_size_estimator: Callable[[Span], int] | None = None, + ): self._lock = asyncio.Lock() # Task queue and rollouts storage @@ -105,6 +164,36 @@ def __init__(self): # Spans storage self._spans: Dict[str, List[Span]] = {} # rollout_id -> list of spans self._span_sequence_ids: Dict[str, int] = Counter() # rollout_id -> sequence_id + self._span_bytes_by_rollout: Dict[str, int] = Counter() + self._total_span_bytes: int = 0 + self._evicted_rollout_span_sets: Set[str] = set() + + self._memory_capacity_bytes = _detect_total_memory_bytes() + if self._memory_capacity_bytes <= 0: + raise ValueError("Detected memory capacity must be positive") + + self._eviction_threshold_bytes = self._resolve_memory_threshold( + eviction_memory_threshold, + default_ratio=0.7, + capacity_bytes=self._memory_capacity_bytes, + name="eviction_memory_threshold", + minimum=1, + ) + + if safe_memory_threshold is None: + safe_memory_threshold = max(int(self._eviction_threshold_bytes * 0.8), 0) + + self._safe_threshold_bytes = self._resolve_memory_threshold( + safe_memory_threshold, + default_ratio=self._eviction_threshold_bytes / self._memory_capacity_bytes, + capacity_bytes=self._memory_capacity_bytes, + name="safe_memory_threshold", + minimum=0, + ) + + if not (0 <= self._safe_threshold_bytes < self._eviction_threshold_bytes): + raise ValueError("safe_memory_threshold must be smaller than eviction_memory_threshold") + self._custom_span_size_estimator = span_size_estimator # Attempt tracking self._attempts: Dict[str, List[Attempt]] = {} # rollout_id -> list of attempts @@ -426,6 +515,8 @@ async def _add_span_unlocked(self, span: Span) -> Span: if span.rollout_id not in self._spans: self._spans[span.rollout_id] = [] self._spans[span.rollout_id].append(span) + self._account_span_size(span) + self._maybe_evict_spans() # Update attempt heartbeat current_attempt.last_heartbeat_time = time.time() @@ -449,6 +540,77 @@ async def _add_span_unlocked(self, span: Span) -> Span: return span + @staticmethod + def _resolve_memory_threshold( + value: float | int | None, + *, + default_ratio: float, + capacity_bytes: int, + name: str, + minimum: int, + ) -> int: + if value is None: + resolved = int(capacity_bytes * default_ratio) + elif isinstance(value, float): + if minimum == 0: + if not (0 <= value <= 1): + raise ValueError(f"{name} ratio must be between 0 and 1 inclusive") + else: + if not (0 < value <= 1): + raise ValueError(f"{name} ratio must be greater than 0 and at most 1") + resolved = int(capacity_bytes * value) + else: + value_int = value + if value_int < 0: + raise ValueError(f"{name} must be non-negative") + resolved = value_int + + if resolved < minimum: + raise ValueError(f"{name} must be at least {minimum} bytes") + + return resolved + + def _account_span_size(self, span: Span) -> int: + if self._custom_span_size_estimator is not None: + size = max(int(self._custom_span_size_estimator(span)), 0) + else: + size = estimate_model_size(span) + + self._span_bytes_by_rollout[span.rollout_id] += size + self._total_span_bytes += size + return size + + def _maybe_evict_spans(self) -> None: + if self._total_span_bytes <= self._eviction_threshold_bytes: + return + + candidates: List[tuple[float, str]] = [] + for rollout_id, spans in self._spans.items(): + if not spans: + continue + rollout = self._rollouts.get(rollout_id) + start_time = rollout.start_time if rollout is not None else (spans[0].start_time or 0.0) + candidates.append((start_time, rollout_id)) + + candidates.sort(key=lambda item: item[0]) + + logger.info(f"Evicting spans for {len(candidates)} rollouts to free up memory...") + memory_consumed_before = self._total_span_bytes + for _, rollout_id in candidates: + if self._total_span_bytes <= self._safe_threshold_bytes: + break + logger.debug(f"Evicting spans for rollout {rollout_id} to free up memory...") + self._evict_spans_for_rollout(rollout_id) + logger.info(f"Freed up {memory_consumed_before - self._total_span_bytes} bytes of memory") + + def _evict_spans_for_rollout(self, rollout_id: str) -> None: + spans = self._spans.pop(rollout_id, []) + if not spans: + return + removed_bytes = self._span_bytes_by_rollout.pop(rollout_id, 0) + self._total_span_bytes = max(self._total_span_bytes - removed_bytes, 0) + self._evicted_rollout_span_sets.add(rollout_id) + @_healthcheck_wrapper async def wait_for_rollouts(self, *, rollout_ids: List[str], timeout: Optional[float] = None) -> List[Rollout]: """ @@ -510,6 +672,8 @@ async def query_spans(self, rollout_id: str, attempt_id: str | Literal["latest"] Returns an empty list if no spans are found. """ async with self._lock: + if rollout_id in self._evicted_rollout_span_sets: + raise RuntimeError(f"Spans for rollout {rollout_id} have been evicted") spans = self._spans.get(rollout_id, []) if attempt_id is None: return spans diff --git a/tests/store/test_memory.py b/tests/store/test_memory.py index 2c7ffa5ae..4dbf908fb 100644 --- a/tests/store/test_memory.py +++ b/tests/store/test_memory.py @@ -14,20 +14,28 @@ """ import asyncio +import sys import time -from typing import List +from typing import List, Optional, cast from unittest.mock import Mock import pytest +from pydantic import BaseModel -from agentlightning.store.memory import InMemoryLightningStore +from agentlightning.store.memory import InMemoryLightningStore, estimate_model_size from agentlightning.types import ( LLM, + AttemptedRollout, + Event, + Link, PromptTemplate, + Resource, ResourcesUpdate, Rollout, RolloutConfig, Span, + SpanContext, + TraceStatus, ) # Core CRUD Operations Tests @@ -608,6 +616,172 @@ async def test_query_spans_by_attempt(inmemory_store: InMemoryLightningStore, mo assert len(no_spans) == 0 +@pytest.mark.asyncio +async def test_span_eviction_removes_oldest_rollouts(mock_readable_span: Mock, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("agentlightning.store.memory._detect_total_memory_bytes", lambda: 100) + store = InMemoryLightningStore( + eviction_memory_threshold=0.5, + safe_memory_threshold=0.05, + span_size_estimator=lambda span: 20, + ) + + attempted_rollouts: List[AttemptedRollout] = [] + for index in range(4): + attempted = await store.start_rollout(input={"index": index}) + attempted_rollouts.append(attempted) + await store.add_otel_span(attempted.rollout_id, attempted.attempt.attempt_id, mock_readable_span) + + for attempted in attempted_rollouts[:3]: + with pytest.raises(RuntimeError): + await store.query_spans(attempted.rollout_id) + + remaining_spans = await store.query_spans(attempted_rollouts[3].rollout_id) + assert len(remaining_spans) == 1 + assert remaining_spans[0].rollout_id == attempted_rollouts[3].rollout_id + + +def test_memory_threshold_accepts_byte_values() -> None: + store = InMemoryLightningStore( + eviction_memory_threshold=150, + safe_memory_threshold=20, + ) + + assert store._eviction_threshold_bytes == 150 # pyright: ignore[reportPrivateUsage] + assert store._safe_threshold_bytes == 20 # pyright: ignore[reportPrivateUsage] + + +def test_memory_threshold_accepts_ratios_with_zero_safe(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr("agentlightning.store.memory._detect_total_memory_bytes", lambda: 200) + store = InMemoryLightningStore( + eviction_memory_threshold=0.6, + safe_memory_threshold=0.0, + ) + + assert store._eviction_threshold_bytes == int(200 * 0.6) # pyright: ignore[reportPrivateUsage] + assert store._safe_threshold_bytes == 0 # pyright: ignore[reportPrivateUsage] + + +def test_invalid_safe_threshold_raises_value_error() -> None: + with pytest.raises(ValueError): + InMemoryLightningStore( + eviction_memory_threshold=50, + safe_memory_threshold=100, + ) + + +def test_estimate_model_size_counts_nested_models() -> None: + class Inner(BaseModel): + value: int + data: List[int] + + class Outer(BaseModel): + inner: Inner + mapping: dict[str, str] + tags: List[str] + + inner = Inner(value=7, data=[1, 2, 3]) + outer = Outer(inner=inner, mapping={"alpha": "beta"}, tags=["x", "yz"]) + + inner_expected = ( + sys.getsizeof(inner) + + sys.getsizeof(inner.value) + + sys.getsizeof(inner.data) + + sum(sys.getsizeof(item) for item in inner.data) + ) + assert estimate_model_size(inner) == inner_expected + + mapping_expected = sys.getsizeof(outer.mapping) + sum(sys.getsizeof(v) for v in outer.mapping.values()) + tags_expected = sys.getsizeof(outer.tags) + sum(sys.getsizeof(tag) for tag in outer.tags) + outer_expected = sys.getsizeof(outer) + inner_expected + mapping_expected + tags_expected + assert estimate_model_size(outer) == outer_expected + + +def test_estimate_model_size_handles_span_objects() -> None: + status = TraceStatus(status_code="OK", description="fine") + context = SpanContext(trace_id="trace", span_id="parent", is_remote=False, trace_state={"foo": "bar"}) + event = Event(name="step", attributes={"detail": "value"}, timestamp=1.0) + link = Link(context=context, attributes=None) + resource = Resource(attributes={"service.name": "unit"}, schema_url="schema") + + span = Span( + rollout_id="ro-1", + attempt_id="at-1", + sequence_id=1, + trace_id="trace", + span_id="span", + parent_id=None, + name="operation", + status=status, + attributes={"foo": "bar", "answer": 42}, + events=[event], + links=[link], + start_time=1.0, + end_time=2.0, + context=None, + parent=None, + resource=resource, + ) + + status_expected = sys.getsizeof(status) + sys.getsizeof(status.status_code) + sys.getsizeof(status.description) + + trace_state_values = context.trace_state.values() + context_expected = ( + sys.getsizeof(context) + + sys.getsizeof(context.trace_id) + + sys.getsizeof(context.span_id) + + sys.getsizeof(context.is_remote) + + sys.getsizeof(context.trace_state) + + sum(sys.getsizeof(v) for v in trace_state_values) + ) + + event_attributes_expected = sys.getsizeof(event.attributes) + sys.getsizeof("value") + event_expected = ( + sys.getsizeof(event) + sys.getsizeof(event.name) + event_attributes_expected + sys.getsizeof(event.timestamp) + ) + events_expected = sys.getsizeof(span.events) + event_expected + + link_attributes = cast(Optional[dict[str, str]], link.attributes) + link_attribute_values = link_attributes.values() if link_attributes is not None else () + link_attributes_expected = sys.getsizeof(link_attributes if link_attributes is not None else None) + sum( + sys.getsizeof(v) for v in link_attribute_values + ) + link_expected = sys.getsizeof(link) + context_expected + link_attributes_expected + links_expected = sys.getsizeof(span.links) + link_expected + + attributes_expected = ( + sys.getsizeof(span.attributes) + sys.getsizeof("bar") + sys.getsizeof(span.attributes["answer"]) + ) + + resource_expected = ( + sys.getsizeof(resource) + + sys.getsizeof(resource.attributes) + + sum(sys.getsizeof(v) for v in resource.attributes.values()) + + sys.getsizeof(resource.schema_url) + ) + + expected_size = ( + sys.getsizeof(span) + + sys.getsizeof(span.rollout_id) + + sys.getsizeof(span.attempt_id) + + sys.getsizeof(span.sequence_id) + + sys.getsizeof(span.trace_id) + + sys.getsizeof(span.span_id) + + sys.getsizeof(span.parent_id) + + sys.getsizeof(span.name) + + status_expected + + attributes_expected + + events_expected + + links_expected + + sys.getsizeof(span.start_time) + + sys.getsizeof(span.end_time) + + sys.getsizeof(span.context) + + sys.getsizeof(span.parent) + + resource_expected + ) + + assert estimate_model_size(span) == expected_size + + @pytest.mark.asyncio async def test_span_triggers_status_transition( inmemory_store: InMemoryLightningStore, mock_readable_span: Mock