|
3 | 3 | from __future__ import annotations |
4 | 4 |
|
5 | 5 | from collections import OrderedDict |
6 | | -from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar, runtime_checkable |
| 6 | +from typing import TYPE_CHECKING, Any, Generic, Protocol, TypeVar, cast, runtime_checkable |
7 | 7 |
|
8 | 8 | if TYPE_CHECKING: |
9 | 9 | from collections.abc import Sequence |
@@ -137,10 +137,20 @@ def _tier_has_save(tier: CacheTier) -> bool: |
137 | 137 | return callable(getattr(tier, "save", None)) |
138 | 138 |
|
139 | 139 |
|
| 140 | +def _tier_save(tier: CacheTier, key: str, value: Any) -> None: |
| 141 | + """Call ``tier.save(key, value)`` — caller must guard with ``_tier_has_save`` first.""" |
| 142 | + cast("Any", tier).save(key, value) |
| 143 | + |
| 144 | + |
140 | 145 | def _tier_has_clear(tier: CacheTier) -> bool: |
141 | 146 | return callable(getattr(tier, "clear", None)) |
142 | 147 |
|
143 | 148 |
|
| 149 | +def _tier_clear(tier: CacheTier, key: str) -> None: |
| 150 | + """Call ``tier.clear(key)`` — caller must guard with ``_tier_has_clear`` first.""" |
| 151 | + cast("Any", tier).clear(key) |
| 152 | + |
| 153 | + |
144 | 154 | class CascadingCache(Generic[V]): # noqa: UP046 |
145 | 155 | """Multi-tier cache where each entry is a ``state()`` node. |
146 | 156 |
|
@@ -174,7 +184,7 @@ def _promote(self, key: str, value: V, hit_tier: int) -> None: |
174 | 184 | for i in range(hit_tier): |
175 | 185 | tier = self._tiers[i] |
176 | 186 | if _tier_has_save(tier): |
177 | | - tier.save(key, value) |
| 187 | + _tier_save(tier, key, value) |
178 | 188 |
|
179 | 189 | def _cascade(self, key: str, nd: Node[Any]) -> None: |
180 | 190 | for tier_index, tier in enumerate(self._tiers): |
@@ -202,10 +212,10 @@ def _evict_if_needed(self) -> None: |
202 | 212 | # Demote to deepest tier with save before evicting |
203 | 213 | for i in range(len(self._tiers) - 1, -1, -1): |
204 | 214 | if _tier_has_save(self._tiers[i]): |
205 | | - self._tiers[i].save(victim, value) |
| 215 | + _tier_save(self._tiers[i], victim, value) |
206 | 216 | for j in range(i): |
207 | 217 | if _tier_has_clear(self._tiers[j]): |
208 | | - self._tiers[j].clear(victim) |
| 218 | + _tier_clear(self._tiers[j], victim) |
209 | 219 | break |
210 | 220 | nd.down([(MessageType.TEARDOWN,)]) |
211 | 221 | del self._entries[victim] |
@@ -239,9 +249,9 @@ def save(self, key: str, value: V) -> None: |
239 | 249 | if self._write_through: |
240 | 250 | for tier in self._tiers: |
241 | 251 | if _tier_has_save(tier): |
242 | | - tier.save(key, value) |
| 252 | + _tier_save(tier, key, value) |
243 | 253 | elif self._tiers and _tier_has_save(self._tiers[0]): |
244 | | - self._tiers[0].save(key, value) |
| 254 | + _tier_save(self._tiers[0], key, value) |
245 | 255 | if key in self._entries: |
246 | 256 | self._entries[key].down([(MessageType.DATA, value)]) |
247 | 257 | if self._eviction is not None: |
@@ -273,7 +283,7 @@ def delete(self, key: str) -> None: |
273 | 283 | self._eviction.delete(key) |
274 | 284 | for tier in self._tiers: |
275 | 285 | if _tier_has_clear(tier): |
276 | | - tier.clear(key) |
| 286 | + _tier_clear(tier, key) |
277 | 287 |
|
278 | 288 | def has(self, key: str) -> bool: |
279 | 289 | """Check if a key is in the in-memory entries.""" |
|
0 commit comments