File size: 12,027 Bytes
406a5e6 | 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 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 | """Read-through disk cache for corpus clients.
``CachedClient`` wraps ANY object that duck-types
``iris_asta.asta_client.AstaClient`` and caches the eight corpus-read
methods on disk (one JSON file per call signature) so repeated solver /
scoring iterations cost zero network calls. Stdlib only; the wrapped
client is never imported, only called.
Contract:
* Cache key: sha256 over ``(method_name, canonical-JSON of args/kwargs)``.
Canonical means: dict keys sorted, tuples as lists, sets as sorted
lists, compact separators — so semantically-equal kwargs always map to
the same key. NOTE the key is built from the args as PASSED: calling
``get_paper("1")`` and ``get_paper(corpus_id="1")`` are different keys
(both correct, just a duplicate fetch); pfbmax callers should pick one
calling convention per site.
* Cached methods: ``snippet_search``, ``paper_search``,
``search_paper_by_title``, ``get_paper``, ``get_paper_batch``,
``get_citations``, ``search_authors_by_name``, ``get_author_papers``.
Every other attribute (method or plain attr) delegates to the inner
client uncached and uncounted.
* Serialization: results may contain iris_asta dataclasses
(``Paper``/``Snippet``). Object nodes are converted to dicts via
``dataclasses.asdict`` (fallback ``vars()``) and tagged with a marker
key; on load ONLY tagged nodes become ``types.SimpleNamespace`` —
plain dicts (``Paper.extra``, author records, ``ref_mentions``) stay
plain dicts, because consumers do ``extra.get("citationCount")`` /
``isinstance(record, dict)`` (see iris_asta/solvers/pfb.py). The
round-trip preserves corpusId/corpus_id/title/abstract/text/score/
year/venue/authors/citationCount under whichever attribute names the
inner client produced.
* Both hits AND misses return the round-tripped (SimpleNamespace) form,
so consumer behavior is identical on cold and warm cache.
* ``None`` results (dead corpus id, unresolvable title) are cached too —
negative caching; file existence discriminates hit from miss.
* Counters: ``.calls`` (inner-client invocations made), ``.hits``,
``.misses``. For cached methods every invocation increments exactly
one of hits/misses, and every miss increments calls (calls == misses
unless you reset counters). Uncached delegation is not counted.
Counters are per-instance and not thread-synchronized; the disk layer
IS safe across processes (atomic ``os.replace`` writes).
* ``bypass=True`` disables cache READS (every call goes to the inner
client and counts as a miss) but still WRITES results — use it to
refresh entries while seeding the shared cache.
* Failure semantics: if the inner client raises, nothing is written and
the exception propagates unchanged. Corrupt/unreadable cache files are
treated as misses and rewritten. Cache-write failures (disk issues)
are swallowed — caching must never break a solve.
Default cache directory: the literal default ``"pfbmax/cache"`` is
anchored at THIS file's directory (``<repo>/pfbmax/cache``) regardless
of the caller's cwd, so every stage shares one physical cache. Any
other explicit path is respected as given.
"""
from __future__ import annotations
import dataclasses
import hashlib
import json
import os
import time
import uuid
from pathlib import Path
from types import SimpleNamespace
__all__ = ["CachedClient", "CACHED_METHODS", "cache_key"]
#: Corpus-read methods served from disk; everything else delegates raw.
CACHED_METHODS = frozenset(
{
"snippet_search",
"paper_search",
"search_paper_by_title",
"get_paper",
"get_paper_batch",
"get_citations",
"search_authors_by_name",
"get_author_papers",
}
)
#: Default cache location; this exact value is anchored at the
#: pfbmax package directory so the shared cache location is cwd-independent.
_DEFAULT_CACHE_DIR = "pfbmax/cache"
#: Marker key tagging dict nodes that were attribute-objects (dataclasses /
#: vars()-able) at serialization time; only these become SimpleNamespace on
#: load. No corpus payload uses this key (S2/MCP JSON never dunders).
_OBJ_KEY = "__pfbmax_obj__"
# ------------------------------------------------------------------ keys --
def _canon(value):
"""Deterministic JSON-ready form of an args/kwargs value.
tuples -> lists, sets -> sorted lists (sorted by their canonical JSON
so mixed types cannot raise), dict keys stringified (json.dumps sorts
them), exotic objects -> str. PYTHONHASHSEED can reorder set iteration
between processes — sorting keeps keys stable across runs.
"""
if value is None or isinstance(value, (bool, int, float, str)):
return value
if isinstance(value, (list, tuple)):
return [_canon(v) for v in value]
if isinstance(value, (set, frozenset)):
return sorted(
(_canon(v) for v in value),
key=lambda v: json.dumps(v, sort_keys=True, default=str),
)
if isinstance(value, dict):
return {str(k): _canon(v) for k, v in value.items()}
return str(value)
def cache_key(method: str, args=(), kwargs=None) -> str:
"""sha256 hex key for one call: (method_name, canonical args/kwargs)."""
payload = json.dumps(
{
"method": str(method),
"args": _canon(list(args)),
"kwargs": _canon(dict(kwargs or {})),
},
sort_keys=True,
separators=(",", ":"),
ensure_ascii=True,
default=str,
)
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
# --------------------------------------------------------- serialization --
def _to_jsonable(value):
"""Normalize a result tree to pure JSON types.
Object nodes (dataclasses via ``dataclasses.asdict``, other objects
via ``vars()``) become dicts tagged with ``_OBJ_KEY``; plain dicts and
lists pass through untagged so they round-trip as themselves.
"""
if value is None or isinstance(value, (bool, int, float, str)):
return value
if isinstance(value, (list, tuple)):
return [_to_jsonable(v) for v in value]
if isinstance(value, dict):
return {str(k): _to_jsonable(v) for k, v in value.items()}
if dataclasses.is_dataclass(value) and not isinstance(value, type):
fields = dataclasses.asdict(value)
else:
try:
fields = vars(value)
except TypeError:
return str(value) # opaque scalar-ish object: best-effort string
out = {_OBJ_KEY: True}
for k, v in fields.items():
out[str(k)] = _to_jsonable(v)
return out
def _from_jsonable(value):
"""Inverse of :func:`_to_jsonable`.
Tagged dicts -> ``SimpleNamespace`` (attribute access for consumers'
getattr-tolerant accessors); untagged dicts stay dicts (consumers do
``extra.get(...)`` / ``isinstance(record, dict)``); lists stay lists.
"""
if isinstance(value, list):
return [_from_jsonable(v) for v in value]
if isinstance(value, dict):
if value.get(_OBJ_KEY) is True:
ns = SimpleNamespace()
for k, v in value.items():
if k != _OBJ_KEY:
ns.__dict__[k] = _from_jsonable(v)
return ns
return {k: _from_jsonable(v) for k, v in value.items()}
return value
# ----------------------------------------------------------------- client --
class CachedClient:
"""Duck-typed AstaClient wrapper with a read-through disk cache.
``CachedClient(inner)`` is a drop-in replacement for ``inner``
anywhere a corpus client is passed (router/solvers/semantic
retrieval): cached methods are intercepted, everything else —
including plain attributes like ``cfg`` — resolves on the inner
client via ``__getattr__``.
"""
def __init__(self, inner, cache_dir: str = _DEFAULT_CACHE_DIR, bypass: bool = False):
self._inner = inner
if cache_dir == _DEFAULT_CACHE_DIR:
self.cache_dir = Path(__file__).resolve().parent / "cache"
else:
self.cache_dir = Path(cache_dir)
self.cache_dir.mkdir(parents=True, exist_ok=True)
self.bypass = bool(bypass)
self.calls = 0 # inner-client invocations actually made
self.hits = 0 # served from disk
self.misses = 0 # went to the inner client
# -------------------------------------------------------- delegation --
def __getattr__(self, name):
# Only reached for names not defined on CachedClient itself.
if name == "_inner": # guard: no recursion before __init__ ran
raise AttributeError(name)
attr = getattr(self._inner, name)
if name in CACHED_METHODS and callable(attr):
def cached_method(*args, _method=name, _fn=attr, **kwargs):
return self._cached_call(_method, _fn, args, kwargs)
cached_method.__name__ = name
cached_method.__qualname__ = f"CachedClient.{name}"
cached_method.__doc__ = getattr(attr, "__doc__", None)
return cached_method
return attr
def __repr__(self): # pragma: no cover - debugging aid
return (
f"CachedClient(inner={type(self._inner).__name__}, "
f"dir={str(self.cache_dir)!r}, bypass={self.bypass}, "
f"calls={self.calls}, hits={self.hits}, misses={self.misses})"
)
# ------------------------------------------------------------- cache --
def cache_path(self, method: str, args=(), kwargs=None) -> Path:
"""Disk path that would serve this exact call (test/tooling aid)."""
return self.cache_dir / (cache_key(method, args, kwargs) + ".json")
def _cached_call(self, method, fn, args, kwargs):
path = self.cache_path(method, args, kwargs)
if not self.bypass:
entry = self._read_entry(path)
if entry is not None:
self.hits += 1
return _from_jsonable(entry["result"])
self.misses += 1
# PFBMAX_CACHE_ONLY=1: never call the remote on a miss -- return the
# method's empty shape instead. For offline pool REBUILDS while the
# corpus API is rate-dead: a missed channel costs a few candidates,
# not a 2-minute retry ladder per call.
if (os.environ.get("PFBMAX_CACHE_ONLY") or "").strip():
return []
self.calls += 1
result = fn(*args, **kwargs) # inner exception: nothing cached, propagates
jsonable = _to_jsonable(result)
self._write_entry(path, method, args, kwargs, jsonable)
# Return the round-tripped form so hits and misses behave identically.
return _from_jsonable(jsonable)
@staticmethod
def _read_entry(path: Path):
"""Load one cache entry; any unreadable/corrupt file is a miss."""
try:
with open(path, "r", encoding="utf-8") as f:
data = json.load(f)
except (OSError, ValueError):
return None
if not isinstance(data, dict) or "result" not in data:
return None
return data
def _write_entry(self, path: Path, method, args, kwargs, jsonable) -> None:
"""Atomically write one entry; write failures never break the call."""
now = time.time()
entry = {
"method": method,
"args": _canon(list(args)),
"kwargs": _canon(dict(kwargs)),
"timestamp": now,
"timestamp_iso": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime(now)),
"result": jsonable,
}
tmp = path.with_name(f"{path.name}.tmp-{os.getpid()}-{uuid.uuid4().hex[:8]}")
try:
with open(tmp, "w", encoding="utf-8") as f:
json.dump(entry, f, ensure_ascii=False, separators=(",", ":"))
os.replace(tmp, path) # atomic on POSIX and Windows (same volume)
except OSError:
try:
os.unlink(tmp)
except OSError:
pass
|