"""Tiered (layered) cache backend for QDK/Chemistry.
Chains multiple cache backends in priority order so that reads hit the
fastest / closest tier first and writes propagate to every tier.
Typical setup pairs a fast local cache with a shared network cache::
cache = TieredCache([
FolderCache("./cache"), # L1 — fast local
FolderCache("/shared/team_cache"), # L2 — shared
])
Read path
Each tier is checked in order. On a hit in a slower tier the result
is *backfilled* into all faster tiers so subsequent reads are local.
Write path
Data is written to **every** tier (write-through).
"""
# --------------------------------------------------------------------------------------------
# Copyright (c) Microsoft Corporation. All rights reserved.
# Licensed under the MIT License. See LICENSE.txt in the project root for license information.
# --------------------------------------------------------------------------------------------
from __future__ import annotations
from typing import TYPE_CHECKING, Any
from qdk_chemistry.remote.cache.base import CacheBackend
if TYPE_CHECKING:
from qdk_chemistry.remote.job import Job
[docs]
class TieredCache(CacheBackend):
"""Composite cache that layers multiple backends.
Args:
tiers: Ordered list of cache backends, fastest first.
Raises:
ValueError: If *tiers* is empty.
Example::
from qdk_chemistry.remote.cache import FolderCache, TieredCache
cache = TieredCache([
FolderCache("./local_cache"),
FolderCache("/shared/team_cache"),
])
energy, wfn = scf.run(mol, 0, 1, "cc-pvdz", cache=cache)
"""
name = "tiered"
[docs]
def __init__(self, tiers: list[CacheBackend | dict[str, Any]], **_kwargs: Any):
"""Initialise with an ordered list of cache tiers.
Args:
tiers: Cache instances or serialized cache configurations, fastest first.
**_kwargs: Ignored compatibility configuration.
"""
super().__init__()
if not tiers:
raise ValueError("TieredCache requires at least one tier")
self._tiers = [self._resolve_tier(tier) for tier in tiers]
@staticmethod
def _resolve_tier(tier: CacheBackend | dict[str, Any]) -> CacheBackend:
"""Resolve a cache instance or serialized cache configuration.
Args:
tier: Cache instance or configuration containing its registry name.
"""
if isinstance(tier, CacheBackend):
return tier
if not isinstance(tier, dict) or "name" not in tier:
raise TypeError("TieredCache tiers must be CacheBackend instances or cache configurations")
from qdk_chemistry.remote.cache import get_cache # noqa: PLC0415
config = dict(tier)
name = config.pop("name")
return get_cache(name, **config)
[docs]
@property
def is_shared(self) -> bool:
"""A tiered cache is shared if any of its tiers is shared."""
return any(tier.is_shared for tier in self._tiers)
[docs]
@property
def tiers(self) -> list[CacheBackend]:
"""Return a copy of the tier list."""
return list(self._tiers)
[docs]
def for_remote(self) -> CacheBackend | None:
"""Return a cache containing only tiers reachable from remote nodes."""
remote_tiers = [remote for tier in self._tiers if (remote := tier.for_remote()) is not None]
if not remote_tiers:
return None
if len(remote_tiers) == 1:
return remote_tiers[0]
return TieredCache(remote_tiers)
[docs]
def to_config(self) -> dict:
"""Return kwargs to reconstruct this TieredCache."""
return {"tiers": [{"name": tier.name, **tier.to_config()} for tier in self._tiers]}
# ── Job metadata ─────────────────────────────────────────────────────
[docs]
def get_job(self, run_hash: str) -> Job | None:
"""Check each tier in order; backfill faster tiers on a hit."""
for i, tier in enumerate(self._tiers):
job = tier.get_job(run_hash)
if job is not None:
# Backfill all faster tiers that missed
for faster in self._tiers[:i]:
faster.put_job(run_hash, job)
return job
return None
[docs]
def put_job(self, run_hash: str, job: Job) -> None:
"""Write-through to every tier."""
for tier in self._tiers:
tier.put_job(run_hash, job)
# ── Data blobs ───────────────────────────────────────────────────────
[docs]
def get_data(self, content_hash: str) -> Any | None:
"""Check each tier in order; backfill faster tiers on a hit."""
for i, tier in enumerate(self._tiers):
data = tier.get_data(content_hash)
if data is not None:
for faster in self._tiers[:i]:
faster.put_data(content_hash, data)
return data
return None
[docs]
def put_data(self, content_hash: str, data: Any, *, shared_only: bool = False) -> None:
"""Write through to every eligible tier."""
for tier in self._tiers:
tier.put_data(content_hash, data, shared_only=shared_only)
[docs]
def has_data(self, content_hash: str, *, shared_only: bool = False) -> bool:
"""Return ``True`` if any eligible tier contains the blob."""
return any(tier.has_data(content_hash, shared_only=shared_only) for tier in self._tiers)
# ── Deletion ─────────────────────────────────────────────────────────
[docs]
def delete_job(self, run_hash: str) -> bool:
"""Remove from every tier. Returns ``True`` if any had it."""
existed = False
for tier in self._tiers:
if tier.delete_job(run_hash):
existed = True
return existed
[docs]
def delete_data(self, content_hash: str) -> bool:
"""Remove from every tier. Returns ``True`` if any had it."""
existed = False
for tier in self._tiers:
if tier.delete_data(content_hash):
existed = True
return existed
[docs]
def clear(self) -> None:
"""Clear every tier."""
for tier in self._tiers:
tier.clear()