File size: 2,954 Bytes
79e8e52
e383cb5
79e8e52
 
e383cb5
 
 
79e8e52
 
e383cb5
 
79e8e52
 
e383cb5
 
79e8e52
e383cb5
79e8e52
 
 
e383cb5
 
 
 
79e8e52
 
e383cb5
 
79e8e52
 
 
 
e383cb5
 
79e8e52
 
 
 
 
 
e383cb5
79e8e52
 
e383cb5
79e8e52
 
 
 
 
 
 
 
 
 
 
 
e383cb5
 
79e8e52
 
e383cb5
79e8e52
 
 
 
 
 
 
 
e383cb5
 
79e8e52
 
 
 
 
e383cb5
79e8e52
 
 
 
 
 
 
 
 
 
 
 
 
e383cb5
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
"""tool_cache.py — V6: tool cache for repeated invocations.

Cache LRU para resultados de ferramentas. Reduz computação redundante
quando a mesma entrada é processada múltiplas vezes.
"""
from __future__ import annotations
import hashlib
import time
from typing import Optional, Dict, Any, Callable
from collections import OrderedDict

import torch


class ToolCache:
    """V6: cache LRU para ferramentas.

    Usage:
        cache = ToolCache(max_size=1024, ttl_seconds=3600)
        result = cache.get_or_compute("tool_name", input_tensor, compute_fn)
    """
    def __init__(
        self,
        max_size: int = 1024,
        ttl_seconds: float = 3600.0,
        hash_fn: Optional[Callable] = None,
    ):
        self.max_size = max_size
        self.ttl_seconds = ttl_seconds
        self.hash_fn = hash_fn or self._default_hash
        self._cache: OrderedDict[str, tuple] = OrderedDict()  # key -> (value, timestamp)
        self._stats = {"hits": 0, "misses": 0, "evictions": 0}

    @staticmethod
    def _default_hash(key: Any) -> str:
        if isinstance(key, torch.Tensor):
            key = key.detach().cpu().numpy().tobytes()
        elif isinstance(key, (list, tuple)):
            key = str(key)
        return hashlib.sha256(str(key).encode("utf-8")).hexdigest()

    def _make_key(self, tool_name: str, input_key: Any) -> str:
        return f"{tool_name}:{self.hash_fn(input_key)}"

    def get_or_compute(
        self,
        tool_name: str,
        input_key: Any,
        compute_fn: Callable,
    ) -> Any:
        key = self._make_key(tool_name, input_key)
        now = time.time()
        # Check cache
        if key in self._cache:
            value, ts = self._cache[key]
            if now - ts < self.ttl_seconds:
                self._cache.move_to_end(key)
                self._stats["hits"] += 1
                return value
            else:
                del self._cache[key]
        # Compute
        value = compute_fn()
        self._cache[key] = (value, now)
        self._stats["misses"] += 1
        # Evict LRU
        while len(self._cache) > self.max_size:
            self._cache.popitem(last=False)
            self._stats["evictions"] += 1
        return value

    def invalidate(self, tool_name: str) -> None:
        """Remove todas as entradas de uma ferramenta."""
        keys_to_remove = [k for k in self._cache if k.startswith(f"{tool_name}:")]
        for k in keys_to_remove:
            del self._cache[k]

    def clear(self) -> None:
        self._cache.clear()
        self._stats = {"hits": 0, "misses": 0, "evictions": 0}

    def get_stats(self) -> Dict[str, Any]:
        total = self._stats["hits"] + self._stats["misses"]
        hit_rate = self._stats["hits"] / max(1, total)
        return {
            **self._stats,
            "hit_rate": hit_rate,
            "size": len(self._cache),
            "max_size": self.max_size,
        }


__all__ = ["ToolCache"]