|
11 | 11 |
|
12 | 12 | from __future__ import annotations |
13 | 13 |
|
| 14 | +import threading |
| 15 | + |
14 | 16 | import pytest |
15 | 17 |
|
16 | 18 | from headroom.cache.token_count_memo import TokenCountMemo, count_messages_memoized |
@@ -133,7 +135,7 @@ def test_memoized_count_handles_append_only_delta() -> None: |
133 | 135 |
|
134 | 136 | count_messages_memoized(memo, tokenizer, turn1) |
135 | 137 | prefix_key = TokenCountMemo.message_hash(turn1[0]) |
136 | | - assert memo.get(prefix_key) is not None |
| 138 | + assert memo.get(prefix_key, tokenizer=tokenizer) is not None |
137 | 139 |
|
138 | 140 | memoized_turn2 = count_messages_memoized(memo, tokenizer, turn2) |
139 | 141 | assert memoized_turn2 == tokenizer.count_messages(turn2) |
@@ -174,21 +176,25 @@ def count_messages(self, messages) -> int: # noqa: ANN001 |
174 | 176 | class TestTokenCountMemoEviction: |
175 | 177 | def test_eviction_at_max_entries(self) -> None: |
176 | 178 | memo = TokenCountMemo(max_entries=3) |
177 | | - memo.put("a", 1) |
178 | | - memo.put("b", 2) |
179 | | - memo.put("c", 3) |
180 | | - memo.get("a") # touch "a" so it's not the least-recently-used |
181 | | - memo.put("d", 4) # should evict "b" (oldest untouched) |
182 | | - |
183 | | - assert memo.get("a") == 1 |
184 | | - assert memo.get("b") is None |
185 | | - assert memo.get("c") == 3 |
186 | | - assert memo.get("d") == 4 |
| 179 | + tok = object() # get/put require the bound tokenizer's identity |
| 180 | + memo.bind_or_reset(tok) |
| 181 | + memo.put("a", 1, tokenizer=tok) |
| 182 | + memo.put("b", 2, tokenizer=tok) |
| 183 | + memo.put("c", 3, tokenizer=tok) |
| 184 | + memo.get("a", tokenizer=tok) # touch "a" so it's not the least-recently-used |
| 185 | + memo.put("d", 4, tokenizer=tok) # should evict "b" (oldest untouched) |
| 186 | + |
| 187 | + assert memo.get("a", tokenizer=tok) == 1 |
| 188 | + assert memo.get("b", tokenizer=tok) is None |
| 189 | + assert memo.get("c", tokenizer=tok) == 3 |
| 190 | + assert memo.get("d", tokenizer=tok) == 4 |
187 | 191 |
|
188 | 192 | def test_get_stats_reports_entry_count(self) -> None: |
189 | 193 | memo = TokenCountMemo() |
190 | | - memo.put("a", 1) |
191 | | - memo.put("b", 2) |
| 194 | + tok = object() |
| 195 | + memo.bind_or_reset(tok) |
| 196 | + memo.put("a", 1, tokenizer=tok) |
| 197 | + memo.put("b", 2, tokenizer=tok) |
192 | 198 | assert memo.get_stats()["entries"] == 2 |
193 | 199 |
|
194 | 200 |
|
@@ -258,3 +264,94 @@ def test_same_tokenizer_keeps_memo_bound(self) -> None: |
258 | 264 | entries = memo.get_stats()["entries"] |
259 | 265 | count_messages_memoized(memo, tokenizer, messages) |
260 | 266 | assert memo.get_stats()["entries"] == entries # no clear, warm hits |
| 267 | + |
| 268 | + |
| 269 | +class _GatedTokenizer: |
| 270 | + """Additive tokenizer with a fixed per-message count and an interleaving hook. |
| 271 | +
|
| 272 | + ``per_message`` differs between the two instances so a cross-tokenizer leak |
| 273 | + shows up as a wrong number rather than a coincidence. |
| 274 | + """ |
| 275 | + |
| 276 | + ADDITIVE_COUNTS = True |
| 277 | + REPLY_OVERHEAD = 0 |
| 278 | + |
| 279 | + def __init__(self, per_message: int, before_count: object = None) -> None: |
| 280 | + self.per_message = per_message |
| 281 | + self._before_count = before_count |
| 282 | + |
| 283 | + def count_message(self, message: dict) -> int: # noqa: ARG002 |
| 284 | + if self._before_count is not None: |
| 285 | + self._before_count() |
| 286 | + return self.per_message |
| 287 | + |
| 288 | + def count_messages(self, messages: list[dict]) -> int: |
| 289 | + return sum(self.count_message(m) for m in messages) + self.REPLY_OVERHEAD |
| 290 | + |
| 291 | + |
| 292 | +class TestCrossTokenizerRebindRace: |
| 293 | + """The per-session memo is shared by concurrent requests, and the count |
| 294 | + fail-open path (``proxy/token_counting._count_offloaded``) hands each |
| 295 | + request a *fresh* ``EstimatingTokenCounter``. So request A can be mid-count |
| 296 | + under tokenizer A while request B rebinds the memo to tokenizer B. A's |
| 297 | + counts must never reach B. |
| 298 | + """ |
| 299 | + |
| 300 | + def test_put_after_concurrent_rebind_never_leaks_to_new_binding(self) -> None: |
| 301 | + messages = [{"role": "user", "content": "shared prefix message"}] |
| 302 | + key = TokenCountMemo.message_hash(messages[0]) |
| 303 | + canonical = TokenCountMemo.canonical_message(messages[0]) |
| 304 | + |
| 305 | + a_is_counting = threading.Event() |
| 306 | + b_has_rebound = threading.Event() |
| 307 | + |
| 308 | + def _pause_a() -> None: |
| 309 | + # A has already called bind_or_reset and missed the cache; hold it |
| 310 | + # here so B's whole rebind+count lands before A's put. |
| 311 | + a_is_counting.set() |
| 312 | + assert b_has_rebound.wait(timeout=10), "B never rebound" |
| 313 | + |
| 314 | + a = _GatedTokenizer(per_message=7, before_count=_pause_a) |
| 315 | + b = _GatedTokenizer(per_message=31) |
| 316 | + memo = TokenCountMemo() |
| 317 | + |
| 318 | + a_total: list[int] = [] |
| 319 | + a_thread = threading.Thread( |
| 320 | + target=lambda: a_total.append(count_messages_memoized(memo, a, messages)), |
| 321 | + daemon=True, |
| 322 | + ) |
| 323 | + a_thread.start() |
| 324 | + assert a_is_counting.wait(timeout=10), "A never started counting" |
| 325 | + |
| 326 | + # B rebinds (clearing A's era) and populates its own count. |
| 327 | + b_total = count_messages_memoized(memo, b, messages) |
| 328 | + b_has_rebound.set() |
| 329 | + a_thread.join(timeout=10) |
| 330 | + assert not a_thread.is_alive() |
| 331 | + |
| 332 | + # A's own total stays exact under tokenizer A — the fix degrades A to |
| 333 | + # uncached counting, it does not hand A B's numbers. |
| 334 | + assert a_total == [7] |
| 335 | + assert b_total == 31 |
| 336 | + |
| 337 | + # The proof: A's post-rebind put was dropped, so the memo still holds |
| 338 | + # only B's count for the shared message and a fresh B read agrees. |
| 339 | + assert memo.get(key, canonical, tokenizer=b) == 31 |
| 340 | + assert count_messages_memoized(memo, b, messages) == 31 |
| 341 | + |
| 342 | + def test_get_from_obsolete_binding_misses_instead_of_reading_new_counts(self) -> None: |
| 343 | + """Mirror direction: a stale binding must not read the rebinder's |
| 344 | + counts either (the swap guarantee stays symmetric).""" |
| 345 | + messages = [{"role": "user", "content": "shared prefix message"}] |
| 346 | + key = TokenCountMemo.message_hash(messages[0]) |
| 347 | + canonical = TokenCountMemo.canonical_message(messages[0]) |
| 348 | + |
| 349 | + a = _GatedTokenizer(per_message=7) |
| 350 | + b = _GatedTokenizer(per_message=31) |
| 351 | + memo = TokenCountMemo() |
| 352 | + |
| 353 | + count_messages_memoized(memo, a, messages) |
| 354 | + count_messages_memoized(memo, b, messages) # rebinds to b |
| 355 | + |
| 356 | + assert memo.get(key, canonical, tokenizer=a) is None |
| 357 | + assert memo.get(key, canonical, tokenizer=b) == 31 |
0 commit comments