Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
prompt_cache.py732 linesDownload Raw Back to langsmith
1"""Prompt caching module for LangSmith SDK.2 3This module provides thread-safe LRU caches with background refresh4for prompt caching. Includes both sync and async implementations.5"""6 7from __future__ import annotations8 9import asyncio10import json11import logging12import threading13import time14import warnings15from abc import ABC16from collections import OrderedDict17from collections.abc import Awaitable18from dataclasses import dataclass19from pathlib import Path20from typing import TYPE_CHECKING, Any, Callable, Optional, Union21 22if TYPE_CHECKING:23    pass24 25logger = logging.getLogger("langsmith.cache")26 27 28DEFAULT_PROMPT_CACHE_TTL_SECONDS = 5 * 60  # 5 minutes29DEFAULT_PROMPT_CACHE_MAX_SIZE = 10030DEFAULT_PROMPT_CACHE_REFRESH_INTERVAL_SECONDS = 60  # 1 minute31 32 33@dataclass34class CacheEntry:35    """A single cache entry with metadata for TTL tracking."""36 37    value: Any  # The cached value (e.g., PromptCommit)38    created_at: float  # time.time() when entry was created/refreshed39    refresh_func: Optional[Callable[[], Any]] = None  # Function to refresh this entry40 41    def is_stale(self, ttl_seconds: Optional[float]) -> bool:42        """Check if entry is past its TTL (needs refresh)."""43        if ttl_seconds is None:44            return False  # Infinite TTL, never stale45        return (time.time() - self.created_at) > ttl_seconds46 47 48@dataclass49class CacheMetrics:50    """Cache performance metrics."""51 52    hits: int = 053    misses: int = 054    refreshes: int = 055    refresh_errors: int = 056 57    @property58    def total_requests(self) -> int:59        """Total cache requests (hits + misses)."""60        return self.hits + self.misses61 62    @property63    def hit_rate(self) -> float:64        """Cache hit rate (0.0 to 1.0)."""65        total = self.total_requests66        return self.hits / total if total > 0 else 0.067 68 69class _BasePromptCache(ABC):70    """Base class for prompt caches with shared LRU logic.71 72    Provides thread-safe in-memory LRU cache operations.73    Subclasses implement the background refresh mechanism.74    """75 76    __slots__ = [77        "_cache",78        "_lock",79        "_max_size",80        "_ttl_seconds",81        "_refresh_interval",82        "_metrics",83    ]84 85    def __init__(86        self,87        max_size: int = DEFAULT_PROMPT_CACHE_MAX_SIZE,88        ttl_seconds: Optional[float] = DEFAULT_PROMPT_CACHE_TTL_SECONDS,89        refresh_interval_seconds: float = DEFAULT_PROMPT_CACHE_REFRESH_INTERVAL_SECONDS,90    ) -> None:91        """Initialize the base cache.92 93        Args:94            max_size: Maximum entries in cache (LRU eviction when exceeded).95            ttl_seconds: Time before entry is considered stale. Set to None for96                infinite TTL (entries never expire, no background refresh).97            refresh_interval_seconds: How often to check for stale entries.98        """99        self._cache: OrderedDict[str, CacheEntry] = OrderedDict()100        self._lock = threading.RLock()101        self._metrics = CacheMetrics()102        self._configure(103            max_size=max_size,104            ttl_seconds=ttl_seconds,105            refresh_interval_seconds=refresh_interval_seconds,106        )107 108    @property109    def metrics(self) -> CacheMetrics:110        """Get cache performance metrics."""111        return self._metrics112 113    def reset_metrics(self) -> None:114        """Reset all metrics to zero."""115        self._metrics = CacheMetrics()116 117    def get(self, key: str, refresh_func: Callable[[], Any]) -> Optional[Any]:118        """Get a value from cache.119 120        Args:121            key: The cache key (prompt identifier like "owner/name:hash").122            refresh_func: Function to refresh this cache entry when stale.123 124        Returns:125            The cached value or None if not found.126            Stale entries are still returned (background refresh handles updates).127        """128        # If max_size is 0, cache is disabled129        if self._max_size == 0:130            return None131 132        with self._lock:133            if key not in self._cache:134                self._metrics.misses += 1135                return None136 137            entry = self._cache[key]138 139            # Update refresh function140            entry.refresh_func = refresh_func141 142            # Move to end for LRU143            self._cache.move_to_end(key)144 145            self._metrics.hits += 1146            return entry.value147 148    def _set(self, key: str, value: Any, refresh_func: Callable[[], Any]) -> None:149        """Set a value in the cache.150 151        Args:152            key: The cache key (prompt identifier).153            value: The value to cache.154            refresh_func: Function to refresh this cache entry when stale.155        """156        # If max_size is 0, cache is disabled - do nothing157        if self._max_size == 0:158            return159 160        with self._lock:161            now = time.time()162            entry = CacheEntry(value=value, created_at=now, refresh_func=refresh_func)163 164            # Check if we need to evict165            if key not in self._cache and len(self._cache) >= self._max_size:166                # Evict oldest (first item in OrderedDict)167                oldest_key = next(iter(self._cache))168                self._cache.pop(oldest_key)169                logger.debug(f"Evicted oldest cache entry: {oldest_key}")170 171            self._cache[key] = entry172            self._cache.move_to_end(key)173 174    def invalidate(self, key: str) -> None:175        """Remove a specific entry from cache.176 177        Args:178            key: The cache key to invalidate.179        """180        with self._lock:181            self._cache.pop(key, None)182 183    def clear(self) -> None:184        """Clear all cache entries from memory."""185        with self._lock:186            self._cache.clear()187 188    def _get_stale_entries(self) -> list[tuple[str, CacheEntry]]:189        """Get list of stale cache entries (thread-safe)."""190        with self._lock:191            return [192                (key, entry)193                for key, entry in self._cache.items()194                if entry.is_stale(self._ttl_seconds)195            ]196 197    def dump(self, path: Union[str, Path]) -> None:198        """Dump cache contents to a JSON file for offline use.199 200        Args:201            path: Path to the output JSON file.202        """203        from langsmith import schemas as ls_schemas204 205        path = Path(path)206        path.parent.mkdir(parents=True, exist_ok=True)207 208        with self._lock:209            entries = {}210            for key, entry in self._cache.items():211                # Serialize PromptCommit using Pydantic212                if isinstance(entry.value, ls_schemas.PromptCommit):213                    # Handle both pydantic v1 and v2214                    if hasattr(entry.value, "model_dump"):215                        value_data = entry.value.model_dump(mode="json")216                    else:217                        value_data = entry.value.dict()218                else:219                    # Fallback for other types220                    value_data = entry.value221 222                entries[key] = value_data223 224            data = {"entries": entries}225 226        # Atomic write: write to temp file then rename227        temp_path = path.with_suffix(".tmp")228        try:229            with open(temp_path, "w") as f:230                json.dump(data, f, indent=2)231            temp_path.replace(path)232            logger.debug(f"Dumped {len(entries)} cache entries to {path}")233        except Exception as e:234            # Clean up temp file on failure235            if temp_path.exists():236                temp_path.unlink()237            raise e238 239    def load(self, path: Union[str, Path]) -> int:240        """Load cache contents from a JSON file.241 242        Args:243            path: Path to the JSON file to load.244 245        Returns:246            Number of entries loaded.247 248        Loaded entries get a fresh TTL starting from load time.249        If the file doesn't exist or is corrupted, returns 0.250        """251        from langsmith import schemas as ls_schemas252 253        path = Path(path)254 255        if not path.exists():256            logger.debug(f"Cache file not found: {path}")257            return 0258 259        try:260            with open(path) as f:261                data = json.load(f)262        except (json.JSONDecodeError, OSError) as e:263            logger.warning(f"Failed to load cache file {path}: {e}")264            return 0265 266        entries = data.get("entries", {})267        loaded = 0268        now = time.time()269 270        with self._lock:271            for key, value_data in entries.items():272                if len(self._cache) >= self._max_size:273                    logger.debug(f"Reached max cache size, stopping load at {loaded}")274                    break275 276                try:277                    # Deserialize PromptCommit using Pydantic (v1 and v2 compatible)278                    if hasattr(ls_schemas.PromptCommit, "model_validate"):279                        value = ls_schemas.PromptCommit.model_validate(value_data)280                    else:281                        value = ls_schemas.PromptCommit.parse_obj(value_data)282 283                    # Fresh TTL from load time284                    entry = CacheEntry(value=value, created_at=now)285                    self._cache[key] = entry286                    loaded += 1287                except Exception as e:288                    logger.warning(f"Failed to load cache entry {key}: {e}")289                    continue290 291        logger.debug(f"Loaded {loaded} cache entries from {path}")292        return loaded293 294    def _configure(295        self,296        max_size: int,297        ttl_seconds: Optional[float],298        refresh_interval_seconds: float,299    ) -> None:300        self._max_size = max_size301        self._ttl_seconds = ttl_seconds302        self._refresh_interval = refresh_interval_seconds303 304 305class PromptCache(_BasePromptCache):306    """Thread-safe LRU cache with background thread refresh.307 308    For use with the synchronous Client.309 310    Features:311    - In-memory LRU cache with configurable max size312    - Background thread for refreshing stale entries313    - Stale-while-revalidate: returns stale data while refresh happens314    - Thread-safe for concurrent access315 316    Example:317        >>> def fetch_prompt(key: str) -> PromptCommit:318        ...     return client._fetch_prompt_from_api(key)319        >>> cache = PromptCache(320        ...     max_size=100,321        ...     ttl_seconds=3600,322        ...     fetch_func=fetch_prompt,323        ... )324        >>> cache.set("my-prompt:latest", prompt_commit)325        >>> cached = cache.get("my-prompt:latest")326        >>> cache.shutdown()327    """328 329    __slots__ = ["_shutdown_event", "_refresh_thread"]330 331    def __init__(332        self,333        *,334        max_size: int = DEFAULT_PROMPT_CACHE_MAX_SIZE,335        ttl_seconds: Optional[float] = DEFAULT_PROMPT_CACHE_TTL_SECONDS,336        refresh_interval_seconds: float = DEFAULT_PROMPT_CACHE_REFRESH_INTERVAL_SECONDS,337    ) -> None:338        """Initialize the sync prompt cache.339 340        Args:341            max_size: Maximum entries in cache (LRU eviction when exceeded).342            ttl_seconds: Time before entry is considered stale. Set to None for343                infinite TTL (offline mode - entries never expire).344                Default: 300 (5 minutes).345            refresh_interval_seconds: How often to check for stale entries.346        """347        super().__init__(348            max_size=max_size,349            ttl_seconds=ttl_seconds,350            refresh_interval_seconds=refresh_interval_seconds,351        )352        self._shutdown_event = threading.Event()353        self._refresh_thread: Optional[threading.Thread] = None354 355        # Background refresh will be started lazily on first set() operation356 357    def set(self, key: str, value: Any, refresh_func: Callable[[], Any]) -> None:358        """Set a value in the cache.359 360        Args:361            key: The cache key (prompt identifier).362            value: The value to cache.363            refresh_func: Function to refresh this cache entry when stale.364        """365        # Start background refresh on first set (lazy initialization)366        if self._refresh_thread is None:367            self._start_refresh_thread()368        self._set(key, value, refresh_func)369 370    def stop(self) -> None:371        """Stop background refresh thread.372 373        Should be called when the client is being cleaned up.374        """375        self.shutdown()376 377    def shutdown(self) -> None:378        """Stop background refresh thread.379 380        Should be called when the client is being cleaned up.381        """382        if self._shutdown_event is not None:383            self._shutdown_event.set()384        if self._refresh_thread is not None:385            self._refresh_thread.join(timeout=5.0)386            self._refresh_thread = None387 388    def _start_refresh_thread(self) -> None:389        """Start background thread for refreshing stale entries."""390        if self._ttl_seconds is not None:391            self._shutdown_event.clear()392            self._refresh_thread = threading.Thread(393                target=self._refresh_loop,394                daemon=True,395                name="PromptCache-refresh",396            )397            self._refresh_thread.start()398            logger.debug("Started cache refresh thread")399 400    def _refresh_loop(self) -> None:401        """Background loop to refresh stale entries."""402        while not self._shutdown_event.wait(self._refresh_interval):403            try:404                self._refresh_stale_entries()405            except Exception as e:406                # Log but don't die - keep the refresh loop running407                logger.exception(f"Unexpected error in cache refresh loop: {e}")408 409    def _refresh_stale_entries(self) -> None:410        """Check for stale entries and refresh them."""411        stale_entries = self._get_stale_entries()412 413        if not stale_entries:414            return415 416        logger.debug(f"Refreshing {len(stale_entries)} stale cache entries")417 418        for key, entry in stale_entries:419            if self._shutdown_event.is_set():420                break421            if entry.refresh_func is not None:422                try:423                    new_value = entry.refresh_func()424                    self.set(key, new_value, entry.refresh_func)425                    self._metrics.refreshes += 1426                    logger.debug(f"Refreshed cache entry: {key}")427                except Exception as e:428                    # Keep stale data on refresh failure429                    self._metrics.refresh_errors += 1430                    logger.warning(f"Failed to refresh cache entry {key}: {e}")431 432    def configure(433        self,434        *,435        max_size: int = DEFAULT_PROMPT_CACHE_MAX_SIZE,436        ttl_seconds: Optional[float] = DEFAULT_PROMPT_CACHE_TTL_SECONDS,437        refresh_interval_seconds: float = DEFAULT_PROMPT_CACHE_REFRESH_INTERVAL_SECONDS,438    ) -> None:439        """Reconfigure the cache parameters.440 441        Args:442            max_size: Maximum entries in cache (LRU eviction when exceeded).443            ttl_seconds: Time before entry is considered stale.444            refresh_interval_seconds: How often to check for stale entries.445        """446        self.stop()447        self._configure(448            max_size=max_size,449            ttl_seconds=ttl_seconds,450            refresh_interval_seconds=refresh_interval_seconds,451        )452 453 454class AsyncPromptCache(_BasePromptCache):455    """Thread-safe LRU cache with asyncio task refresh.456 457    For use with the asynchronous AsyncClient.458 459    Features:460    - In-memory LRU cache with configurable max size461    - Asyncio task for refreshing stale entries462    - Stale-while-revalidate: returns stale data while refresh happens463    - Thread-safe for concurrent access464 465    Example:466        >>> async def fetch_prompt(key: str) -> PromptCommit:467        ...     return await client._afetch_prompt_from_api(key)468        >>> cache = AsyncPromptCache(469        ...     max_size=100,470        ...     ttl_seconds=3600,471        ...     fetch_func=fetch_prompt,472        ... )473        >>> await cache.start()474        >>> cache.set("my-prompt:latest", prompt_commit)475        >>> cached = cache.get("my-prompt:latest")476        >>> await cache.stop()477    """478 479    __slots__ = ["_refresh_task"]480 481    def __init__(482        self,483        *,484        max_size: int = DEFAULT_PROMPT_CACHE_MAX_SIZE,485        ttl_seconds: Optional[float] = DEFAULT_PROMPT_CACHE_TTL_SECONDS,486        refresh_interval_seconds: float = DEFAULT_PROMPT_CACHE_REFRESH_INTERVAL_SECONDS,487    ) -> None:488        """Initialize the async prompt cache.489 490        Args:491            max_size: Maximum entries in cache (LRU eviction when exceeded).492            ttl_seconds: Time before entry is considered stale. Set to None for493                infinite TTL (offline mode - entries never expire).494            refresh_interval_seconds: How often to check for stale entries.495        """496        super().__init__(497            max_size=max_size,498            ttl_seconds=ttl_seconds,499            refresh_interval_seconds=refresh_interval_seconds,500        )501        self._refresh_task: Optional[asyncio.Task[None]] = None502 503    async def aset(504        self, key: str, value: Any, refresh_func: Callable[[], Awaitable[Any]]505    ) -> None:506        """Set a value in the cache.507 508        Args:509            key: The cache key (prompt identifier).510            value: The value to cache.511            refresh_func: Async function to refresh this cache entry when stale.512        """513        # Start background refresh on first set (lazy initialization)514        if self._refresh_task is None:515            await self.start()516        self._set(key, value, refresh_func)517 518    async def start(self) -> None:519        """Start async background refresh loop.520 521        Must be called from an async context. Creates an asyncio task that522        periodically checks for stale entries and refreshes them.523        Does nothing if ttl_seconds is None (infinite TTL mode).524        """525        if self._ttl_seconds is None:526            return527 528        if self._refresh_task is not None:529            # Already running530            return531 532        self._refresh_task = asyncio.create_task(533            self._refresh_loop(),534            name="AsyncPromptCache-refresh",535        )536        logger.debug("Started async cache refresh task")537 538    def shutdown(self) -> None:539        """Stop background refresh task.540 541        Synchronous wrapper that cancels the refresh task.542        For proper cleanup in async context, use stop() instead.543        """544        if self._refresh_task is not None:545            self._refresh_task.cancel()546            self._refresh_task = None547 548    async def stop(self) -> None:549        """Stop async background refresh loop.550 551        Cancels the refresh task and waits for it to complete.552        """553        if self._refresh_task is None:554            return555 556        self._refresh_task.cancel()557        try:558            await self._refresh_task559        except asyncio.CancelledError:560            pass561        self._refresh_task = None562        logger.debug("Stopped async cache refresh task")563 564    async def _refresh_loop(self) -> None:565        """Async background loop to refresh stale entries."""566        while True:567            try:568                await asyncio.sleep(self._refresh_interval)569                await self._refresh_stale_entries()570            except asyncio.CancelledError:571                raise572            except Exception as e:573                # Log but don't die - keep the refresh loop running574                logger.exception(f"Unexpected error in async cache refresh loop: {e}")575 576    async def _refresh_stale_entries(self) -> None:577        """Check for stale entries and refresh them asynchronously."""578        stale_entries = self._get_stale_entries()579 580        if not stale_entries:581            return582 583        logger.debug(f"Async refreshing {len(stale_entries)} stale cache entries")584 585        for key, entry in stale_entries:586            if entry.refresh_func is not None:587                try:588                    new_value = await entry.refresh_func()589                    await self.aset(key, new_value, entry.refresh_func)590                    self._metrics.refreshes += 1591                    logger.debug(f"Async refreshed cache entry: {key}")592                except Exception as e:593                    # Keep stale data on refresh failure594                    self._metrics.refresh_errors += 1595                    logger.warning(f"Failed to async refresh cache entry {key}: {e}")596 597    async def configure(598        self,599        *,600        max_size: int = DEFAULT_PROMPT_CACHE_MAX_SIZE,601        ttl_seconds: Optional[float] = DEFAULT_PROMPT_CACHE_TTL_SECONDS,602        refresh_interval_seconds: float = DEFAULT_PROMPT_CACHE_REFRESH_INTERVAL_SECONDS,603    ) -> None:604        """Reconfigure the cache parameters.605 606        Args:607            max_size: Maximum entries in cache (LRU eviction when exceeded).608            ttl_seconds: Time before entry is considered stale.609            refresh_interval_seconds: How often to check for stale entries.610        """611        await self.stop()612        self._configure(max_size, ttl_seconds, refresh_interval_seconds)613 614 615# Global singleton instances for prompt caching616prompt_cache_singleton = PromptCache()617async_prompt_cache_singleton = AsyncPromptCache()618 619 620def configure_global_prompt_cache(621    *,622    max_size: int = DEFAULT_PROMPT_CACHE_MAX_SIZE,623    ttl_seconds: Optional[float] = DEFAULT_PROMPT_CACHE_TTL_SECONDS,624    refresh_interval_seconds: float = DEFAULT_PROMPT_CACHE_REFRESH_INTERVAL_SECONDS,625) -> None:626    """Configure the global prompt cache.627 628    This should be called before any cache instances are created or used.629 630    Args:631        max_size: Maximum entries in cache (LRU eviction when exceeded).632        ttl_seconds: Time before entry is considered stale.633        refresh_interval_seconds: How often to check for stale entries.634 635    Example:636        >>> from langsmith import configure_global_prompt_cache637        >>> configure_global_prompt_cache(max_size=200, ttl_seconds=7200)638    """639    prompt_cache_singleton.configure(640        max_size=max_size,641        ttl_seconds=ttl_seconds,642        refresh_interval_seconds=refresh_interval_seconds,643    )644 645 646async def configure_global_async_prompt_cache(647    *,648    max_size: int = DEFAULT_PROMPT_CACHE_MAX_SIZE,649    ttl_seconds: Optional[float] = DEFAULT_PROMPT_CACHE_TTL_SECONDS,650    refresh_interval_seconds: float = DEFAULT_PROMPT_CACHE_REFRESH_INTERVAL_SECONDS,651) -> None:652    """Configure the global prompt cache.653 654    This should be called before any cache instances are created or used.655 656    Args:657        max_size: Maximum entries in cache (LRU eviction when exceeded).658        ttl_seconds: Time before entry is considered stale.659        refresh_interval_seconds: How often to check for stale entries.660 661    Example:662        >>> from langsmith import configure_global_prompt_cache663        >>> configure_global_prompt_cache(max_size=200, ttl_seconds=7200)664    """665    await async_prompt_cache_singleton.configure(666        max_size=max_size,667        ttl_seconds=ttl_seconds,668        refresh_interval_seconds=refresh_interval_seconds,669    )670 671 672# Deprecated alias for backwards compatibility673 674 675def _deprecated_cache_class_warning() -> None:676    warnings.warn(677        "The 'Cache' class is deprecated and will be removed in a future version. "678        "Use 'PromptCache' instead.",679        DeprecationWarning,680        stacklevel=3,681    )682 683 684class Cache(PromptCache):685    """Deprecated alias for PromptCache. Use PromptCache instead."""686 687    def __init__(688        self,689        *,690        max_size: int = DEFAULT_PROMPT_CACHE_MAX_SIZE,691        ttl_seconds: Optional[float] = DEFAULT_PROMPT_CACHE_TTL_SECONDS,692        refresh_interval_seconds: float = DEFAULT_PROMPT_CACHE_REFRESH_INTERVAL_SECONDS,693    ) -> None:694        """Initialize the deprecated Cache class.695 696        Args:697            max_size: Maximum entries in cache (LRU eviction when exceeded).698            ttl_seconds: Time before entry is considered stale.699            refresh_interval_seconds: How often to check for stale entries.700        """701        _deprecated_cache_class_warning()702        super().__init__(703            max_size=max_size,704            ttl_seconds=ttl_seconds,705            refresh_interval_seconds=refresh_interval_seconds,706        )707 708 709class AsyncCache(AsyncPromptCache):710    """Deprecated alias for AsyncPromptCache. Use AsyncPromptCache instead."""711 712    def __init__(713        self,714        *,715        max_size: int = DEFAULT_PROMPT_CACHE_MAX_SIZE,716        ttl_seconds: Optional[float] = DEFAULT_PROMPT_CACHE_TTL_SECONDS,717        refresh_interval_seconds: float = DEFAULT_PROMPT_CACHE_REFRESH_INTERVAL_SECONDS,718    ) -> None:719        """Initialize the deprecated AsyncCache class.720 721        Args:722            max_size: Maximum entries in cache (LRU eviction when exceeded).723            ttl_seconds: Time before entry is considered stale.724            refresh_interval_seconds: How often to check for stale entries.725        """726        _deprecated_cache_class_warning()727        super().__init__(728            max_size=max_size,729            ttl_seconds=ttl_seconds,730            refresh_interval_seconds=refresh_interval_seconds,731        )732 
codekingpro/portable-devtools · Team Ai