aether-context 0.3.0__py3-none-any.whl

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
@@ -0,0 +1,501 @@
1
+ # aether-context (Unlimited Context)
2
+ # Copyright (c) 2026 Aether AI
3
+ # SPDX-License-Identifier: Apache-2.0
4
+ """B3 slice loader — the prefetch **pager** for virtual-memory-for-attention.
5
+
6
+ Make context-slice retrieval fast by PRE-LOADING the slices the session is about to need,
7
+ instead of retrieving cold on every turn. In a coding run the next slice is highly
8
+ predictable from what the model is reasoning about *now*: embed the current reasoning text,
9
+ search the pool, and keep the nearest slices in a small **warm set**. When the model needs
10
+ that context, the slice is an O(1) memory lookup, not an ANN search.
11
+
12
+ Expected retrieval latency is then ``E[t] = h·t_warm + (1−h)·t_cold`` where ``h`` is the hit
13
+ rate; the pager's whole job is to push ``h → 1`` by prefetching from current state. ``t_warm``
14
+ is a dict lookup (~µs); ``t_cold`` is the classical pool search (~ms). The hit rate is
15
+ **measured**, not assumed.
16
+
17
+ Design
18
+ ------
19
+ A single-threaded, LRU-budgeted warm cache (``prefetch``/``get``/``invalidate`` + hit-rate),
20
+ plus two small disciplines:
21
+
22
+ * the **idle-aware ε re-probe** (:func:`reprobe_probability` / :func:`should_reprobe`) — a
23
+ key that has gone idle gets a rising probability of being re-checked, so a
24
+ stale-but-recoverable region never stays dark forever; and
25
+ * the **depth-cap-1 provenance grounding verdict** (:func:`grounding_verdict`, capped at
26
+ :data:`MAX_CORRECTION_DEPTH`) — a paged-back slice is flagged only if it has *no
27
+ provenance* or contradicts a *hard fact*; merely disagreeing with recent (possibly
28
+ stale) context is **not** a flag.
29
+
30
+ The key
31
+ -------
32
+ A :class:`SliceKey(session, topic)` is a plain discrete coordinate: ``session`` is the
33
+ namespace and ``topic`` is a coarse phase/subject label. The cold path is injected as
34
+ ``retrieve_fn`` defaulting to ``context_pool.search`` (scoped to ``key.session``). Keys are
35
+ **discrete strings only** — :class:`SliceKey` rejects non-string coordinates with
36
+ ``TypeError`` (vectors address slices only *inside* :meth:`Pager.get`, never as a key).
37
+
38
+ Single-threaded by design
39
+ -------------------------
40
+ The pager core is **single-threaded and pure**. Concurrency belongs to the caller: the session
41
+ runs :meth:`Pager.prefetch_from` on a background thread *while the model generates* (the backend
42
+ HTTP/subprocess call releases the GIL, so a prefetch thread genuinely overlaps generation). No
43
+ thread, lock, or queue lives in this module.
44
+
45
+ Fail-soft
46
+ ---------
47
+ The pager is an *optimization*, never a correctness dependency. A cold-search error or an
48
+ encoder hiccup is logged and degrades to an empty window — it never raises into a long run.
49
+ """
50
+ from __future__ import annotations
51
+
52
+ from dataclasses import dataclass
53
+ from enum import Enum
54
+ from typing import Callable, Iterable, Protocol
55
+
56
+ import numpy as np
57
+
58
+ from aether_context._log import get_logger
59
+ from aether_context.context_pool import Slice
60
+
61
+ logger = get_logger(__name__)
62
+
63
+ #: Number of pre-assembled slice *keys* kept warm at once (the pager's working-set budget).
64
+ #: Mirrors the upstream ``DEFAULT_WARM_BUDGET``; 16 keys is plenty of reach for one turn while
65
+ #: staying tiny in RAM (the slice payloads themselves live in the pool, not here).
66
+ DEFAULT_WARM_BUDGET: int = 16
67
+
68
+ # --- idle-aware re-probe constants (ported from exploration.py) --------------
69
+ #: Floor re-probe probability right after a key was accessed (never 0 — nothing stays dark).
70
+ BASE_EPS: float = 0.02
71
+ #: Idle periods at which the gap-to-certain halves (so probability rises smoothly toward 1).
72
+ HALF_LIFE_PERIODS: float = 20.0
73
+
74
+ # --- depth-cap grounding constants (ported from latency_budget.py) -----------
75
+ #: One correction attempt, then abstain — the grounding check never recurses unboundedly.
76
+ MAX_CORRECTION_DEPTH: int = 1
77
+
78
+
79
+ class Grounding(str, Enum):
80
+ """Provenance-first grounding verdict for a paged-back slice."""
81
+
82
+ PASS = "pass"
83
+ FLAG = "flag"
84
+
85
+
86
+ # --- the key (generalized; discrete; no trading/8-dim coordinate) ------------
87
+ @dataclass(frozen=True)
88
+ class SliceKey:
89
+ """A discrete, hashable coordinate that addresses a region of the pool.
90
+
91
+ A plain ``(session, topic)`` pair: ``session`` is the namespace (scopes the cold search)
92
+ and ``topic`` is a coarse phase/subject label the session's own state machine assigns.
93
+
94
+ The key is **discrete strings only**. A vector (the 256-dim retrieval embedding) is
95
+ rejected with ``TypeError`` — vectors address slices only *inside* :meth:`Pager.get`,
96
+ never as a key.
97
+ """
98
+
99
+ session: str
100
+ topic: str
101
+
102
+ def __post_init__(self) -> None:
103
+ if not isinstance(self.session, str):
104
+ raise TypeError(
105
+ f"SliceKey.session must be a str, got {type(self.session).__name__}; "
106
+ "a SliceKey is a discrete (session, topic) coordinate, not a vector."
107
+ )
108
+ if not isinstance(self.topic, str):
109
+ raise TypeError(
110
+ f"SliceKey.topic must be a str, got {type(self.topic).__name__}; a key is a "
111
+ "discrete string coordinate, not a vector. Use a topic label; query vectors "
112
+ "address slices only inside Pager.get/search."
113
+ )
114
+
115
+
116
+ # --- minimal structural contracts the pager depends on -----------------------
117
+ class _PoolLike(Protocol):
118
+ """The slice of :class:`~aether_context.context_pool.ContextPool` the pager uses."""
119
+
120
+ def search(
121
+ self, query_vec: np.ndarray, k: int, session: str | None = ...
122
+ ) -> list[Slice]: ...
123
+
124
+
125
+ class _EncoderLike(Protocol):
126
+ """The slice of :class:`~aether_context.encoder.StaticEncoder` the pager uses."""
127
+
128
+ def encode(self, text: str) -> np.ndarray: ...
129
+
130
+
131
+ #: Cold path signature: a key + query vector + k -> the slices for that region.
132
+ RetrieveFn = Callable[["SliceKey", np.ndarray, int], list[Slice]]
133
+
134
+
135
+ # --- idle-aware ε re-probe (ported from exploration.py) ----------------------
136
+ def reprobe_probability(
137
+ periods_since_probe: float,
138
+ base_eps: float = BASE_EPS,
139
+ half_life: float = HALF_LIFE_PERIODS,
140
+ ) -> float:
141
+ """Re-probe probability for a region idle ``periods_since_probe`` periods.
142
+
143
+ Rises from ``base_eps`` (just accessed) toward ``1.0`` (long idle), so no suppressed /
144
+ stale region stays dark forever. The gap-to-1 halves every ``half_life`` idle periods.
145
+ Monotone non-decreasing in idle time and clamped to ``[base_eps, 1.0]``.
146
+ """
147
+ if half_life <= 0:
148
+ return 1.0
149
+ gap = (1.0 - base_eps) * (0.5 ** (max(0.0, periods_since_probe) / half_life))
150
+ return max(base_eps, min(1.0, 1.0 - gap))
151
+
152
+
153
+ def should_reprobe(
154
+ periods_since_probe: float,
155
+ rng_uniform: float,
156
+ base_eps: float = BASE_EPS,
157
+ half_life: float = HALF_LIFE_PERIODS,
158
+ ) -> bool:
159
+ """Decide whether to re-probe an idle region; ``rng_uniform`` is a draw in ``[0,1)``."""
160
+ return rng_uniform < reprobe_probability(periods_since_probe, base_eps, half_life)
161
+
162
+
163
+ # --- depth-cap-1 provenance grounding (ported from latency_budget.py) --------
164
+ def grounding_verdict(has_provenance: bool, contradicts_hard_fact: bool) -> Grounding:
165
+ """Provenance-first grounding for a paged-back slice (recursion capped at depth 1).
166
+
167
+ A claim is flagged only if it **contradicts a hard fact** or has **no provenance** at all.
168
+ Disagreeing with recent (possibly stale) *context* is explicitly NOT a flag — that is how
169
+ a correct new decision survives the catcher during a shift. A slice that came from the pool
170
+ inherently has provenance (it was encoded and externalized from real prior context), so the
171
+ common pager case (resident slice, no hard-fact contradiction) is :attr:`Grounding.PASS`.
172
+ """
173
+ if contradicts_hard_fact:
174
+ return Grounding.FLAG
175
+ if not has_provenance:
176
+ return Grounding.FLAG
177
+ return Grounding.PASS
178
+
179
+
180
+ # --- the pager ---------------------------------------------------------------
181
+ class Pager:
182
+ """Single-threaded, budget-bounded warm cache of pool slices — the B3 pager.
183
+
184
+ Wraps a :class:`~aether_context.context_pool.ContextPool` and a
185
+ :class:`~aether_context.encoder.StaticEncoder`. Construct with a warm-key budget; the
186
+ pager keeps at most ``budget`` :class:`SliceKey` regions warm, evicting the least-recently
187
+ used when full. The slice payloads live in the pool — the warm set holds only the small
188
+ ``key -> [Slice]`` mapping and per-key LRU / idle bookkeeping.
189
+
190
+ Public surface:
191
+ * :meth:`prefetch_from` — embed reasoning text, search the pool, warm the result (the
192
+ method the session runs on a side thread while the model generates).
193
+ * :meth:`prefetch` — warm a key with explicit text and/or an explicit query vector.
194
+ * :meth:`get` — hot path: warm ``O(1)`` hit, else a cold search that warms opportunistically.
195
+ * :meth:`window` — the resident slices (the working set the model can be handed this turn).
196
+ * :meth:`hit_rate` — measured hits / (hits + misses).
197
+ * :meth:`invalidate` — drop warm keys matching a predicate (stale entry / topic change).
198
+ * :meth:`reprobe_probability` — idle-aware re-probe probability for a warm key.
199
+ * :meth:`ground` — depth-cap-1 grounding verdict for a paged-back slice.
200
+ """
201
+
202
+ def __init__(
203
+ self,
204
+ pool: _PoolLike,
205
+ encoder: _EncoderLike,
206
+ budget: int = DEFAULT_WARM_BUDGET,
207
+ *,
208
+ retrieve_fn: RetrieveFn | None = None,
209
+ default_k: int = 8,
210
+ ) -> None:
211
+ self._pool = pool
212
+ self._encoder = encoder
213
+ self.budget = max(1, int(budget))
214
+ self._default_k = max(1, int(default_k))
215
+ # Cold path: SliceKey + query vector -> slices. Defaults to a session-scoped pool
216
+ # search so namespace isolation rides through the pager for free.
217
+ self._retrieve: RetrieveFn = (
218
+ retrieve_fn if retrieve_fn is not None else self._pool_retrieve
219
+ )
220
+ # Warm state. The slice payloads stay in the pool; here we keep only the mapping
221
+ # and the LRU / idle counters (pure bookkeeping, single-threaded).
222
+ self._warm: dict[SliceKey, list[Slice]] = {}
223
+ self._lastused: dict[SliceKey, int] = {}
224
+ self._seq = 0
225
+ self.hits = 0
226
+ self.misses = 0
227
+ self.prefetched = 0
228
+
229
+ # -- cold path (injected; defaults to a session-scoped pool search) --------
230
+ def _pool_retrieve(self, key: SliceKey, query_vec: np.ndarray, k: int) -> list[Slice]:
231
+ """Default cold path: a session-scoped ``pool.search`` for ``key``'s region."""
232
+ return self._pool.search(query_vec, k, session=key.session)
233
+
234
+ # -- LRU bookkeeping (single-threaded) ------------------------------------
235
+ def _touch(self, key: SliceKey) -> None:
236
+ """Mark ``key`` as most-recently-used (monotone sequence counter)."""
237
+ self._seq += 1
238
+ self._lastused[key] = self._seq
239
+
240
+ def _evict_one(self, protect: set[SliceKey]) -> None:
241
+ """Evict the least-recently-used warm key not in ``protect`` (no-op if all protected)."""
242
+ cand = [k for k in self._warm if k not in protect]
243
+ if not cand:
244
+ return
245
+ victim = min(cand, key=lambda k: self._lastused.get(k, 0))
246
+ self._warm.pop(victim, None)
247
+ self._lastused.pop(victim, None)
248
+
249
+ def _store_warm(self, key: SliceKey, slices: list[Slice], protect: set[SliceKey]) -> None:
250
+ """Insert/refresh ``key`` in the warm set, evicting LRU to stay within budget."""
251
+ if key not in self._warm and len(self._warm) >= self.budget:
252
+ self._evict_one(protect)
253
+ if key not in self._warm and len(self._warm) >= self.budget:
254
+ # Every other warm key is protected this pass; drop the incoming one.
255
+ return
256
+ self._warm[key] = slices
257
+ self._touch(key)
258
+
259
+ # -- embedding (fail-soft) ------------------------------------------------
260
+ def _embed(self, text: str) -> np.ndarray | None:
261
+ """Encode ``text`` to a query vector; on any encoder error degrade to ``None``."""
262
+ try:
263
+ return np.asarray(self._encoder.encode(text), dtype=np.float32)
264
+ except Exception as exc: # noqa: BLE001 - fail-soft: pager is an optimization
265
+ logger.warning("encoder failed during prefetch (%s); skipping warm", exc)
266
+ return None
267
+
268
+ def _retrieve_safe(
269
+ self, key: SliceKey, query_vec: np.ndarray, k: int
270
+ ) -> list[Slice] | None:
271
+ """Run the injected cold path; on any error degrade to ``None`` (never raise)."""
272
+ try:
273
+ return list(self._retrieve(key, query_vec, k))
274
+ except Exception as exc: # noqa: BLE001 - fail-soft: never crash a long run
275
+ logger.warning("cold retrieve failed for %r (%s); degrading", key, exc)
276
+ return None
277
+
278
+ # -- write: warm the predicted-next region --------------------------------
279
+ def prefetch_from(
280
+ self, key: SliceKey, reasoning_text: str, *, k: int | None = None
281
+ ) -> list[Slice]:
282
+ """Embed ``reasoning_text``, search the pool, and warm the nearest slices under ``key``.
283
+
284
+ This is the pager's headline move and the one the session runs on a side thread *while
285
+ the model generates*: from what the model is reasoning about *now*, predict and warm the
286
+ slices it is about to need. Returns the warmed slices (``[]`` on an encoder/search hiccup
287
+ — fail-soft, never raises). Off the hot path; protects the prefetched key from eviction.
288
+ """
289
+ query = self._embed(reasoning_text)
290
+ if query is None:
291
+ return []
292
+ return self.prefetch(key, reasoning_text, query_vec=query, k=k)
293
+
294
+ def prefetch(
295
+ self,
296
+ key: SliceKey,
297
+ reasoning_text: str | None = None,
298
+ *,
299
+ query_vec: np.ndarray | None = None,
300
+ k: int | None = None,
301
+ ) -> list[Slice]:
302
+ """Warm ``key`` from an explicit ``query_vec`` (or embed ``reasoning_text``).
303
+
304
+ Idempotent-ish: an already-warm key is left in place and returned. Bounded by the warm
305
+ budget (LRU eviction of an unprotected key when full). Returns the slices now warm for
306
+ ``key`` (``[]`` on a degraded embed/search). Re-warming refreshes the key's idle clock.
307
+ """
308
+ if key in self._warm:
309
+ self._touch(key)
310
+ return self._warm[key]
311
+ query = query_vec
312
+ if query is None:
313
+ if reasoning_text is None:
314
+ return []
315
+ query = self._embed(reasoning_text)
316
+ if query is None:
317
+ return []
318
+ query = np.asarray(query, dtype=np.float32)
319
+ slices = self._retrieve_safe(key, query, k if k is not None else self._default_k)
320
+ if slices is None:
321
+ return []
322
+ self._store_warm(key, slices, protect={key})
323
+ self.prefetched += 1
324
+ return self._warm.get(key, [])
325
+
326
+ # -- read: the hot path ----------------------------------------------------
327
+ def get(
328
+ self,
329
+ key: SliceKey,
330
+ reasoning_text: str | None = None,
331
+ *,
332
+ query_vec: np.ndarray | None = None,
333
+ k: int | None = None,
334
+ ) -> list[Slice]:
335
+ """Hot path. Warm ``key`` -> ``O(1)`` hit; cold -> a search that warms opportunistically.
336
+
337
+ On a warm hit the slices come straight from the warm set (no cold call) and ``hits`` is
338
+ incremented. On a miss the cold path runs (embedding ``reasoning_text`` or using
339
+ ``query_vec``), the result warms the key (within budget), and ``misses`` is incremented —
340
+ so :meth:`hit_rate` reflects reality. A degraded cold path returns ``[]`` (still a miss).
341
+ """
342
+ warm = self._warm.get(key)
343
+ if warm is not None:
344
+ self.hits += 1
345
+ self._touch(key)
346
+ return warm
347
+ self.misses += 1
348
+ query = query_vec
349
+ if query is None and reasoning_text is not None:
350
+ query = self._embed(reasoning_text)
351
+ if query is None:
352
+ # No way to address the region without a query vector -> empty window (fail-soft).
353
+ return []
354
+ query = np.asarray(query, dtype=np.float32)
355
+ slices = self._retrieve_safe(key, query, k if k is not None else self._default_k)
356
+ if slices is None:
357
+ return []
358
+ self._store_warm(key, slices, protect={key})
359
+ return self._warm.get(key, [])
360
+
361
+ # -- invalidation ----------------------------------------------------------
362
+ def invalidate(self, predicate: Callable[[SliceKey], bool]) -> int:
363
+ """Drop warm keys matching ``predicate`` (stale entry / topic / session change).
364
+
365
+ Returns how many warm keys were dropped. The pool is untouched — invalidation only
366
+ cools the warm set so the next :meth:`get` re-reads fresh from the pool.
367
+ """
368
+ drop = [k for k in self._warm if predicate(k)]
369
+ for k in drop:
370
+ self._warm.pop(k, None)
371
+ self._lastused.pop(k, None)
372
+ return len(drop)
373
+
374
+ @property
375
+ def default_k(self) -> int:
376
+ """How many slices the cold path pulls per region by default (the resident width)."""
377
+ return self._default_k
378
+
379
+ @default_k.setter
380
+ def default_k(self, value: int) -> None:
381
+ """Set the resident width (floored at 1). Used by the Extended-Thinking toggle."""
382
+ self._default_k = max(1, int(value))
383
+
384
+ def reset(self) -> int:
385
+ """Drop the entire resident window (every warm key); return how many were dropped.
386
+
387
+ This is the *resident* half of the engine's clear semantics: it empties the working
388
+ set the model would be handed this turn, leaving the pool (the reachable slices on
389
+ disk) completely untouched. The next :meth:`get` / :meth:`prefetch_from` re-reads
390
+ fresh from the pool. Hit/miss counters are left intact so a measured hit rate is not
391
+ forged by a clear. Returns the number of warm keys evicted (``0`` when already empty).
392
+ """
393
+ dropped = len(self._warm)
394
+ self._warm.clear()
395
+ self._lastused.clear()
396
+ return dropped
397
+
398
+ def is_warm(self, key: SliceKey) -> bool:
399
+ """Whether ``key`` currently has a resident warm entry."""
400
+ return key in self._warm
401
+
402
+ # -- resident window -------------------------------------------------------
403
+ def window(self) -> list[Slice]:
404
+ """The resident slices across all warm keys — the working set for this turn.
405
+
406
+ De-duplicated by slice id (a slice may be warmed under more than one key), ordered by
407
+ warm-key recency (most-recently-used keys first) so the freshest context leads.
408
+ """
409
+ seen: set[str] = set()
410
+ out: list[Slice] = []
411
+ for key in sorted(self._warm, key=lambda k: self._lastused.get(k, 0), reverse=True):
412
+ for sl in self._warm[key]:
413
+ if sl.id not in seen:
414
+ seen.add(sl.id)
415
+ out.append(sl)
416
+ return out
417
+
418
+ @property
419
+ def warm_count(self) -> int:
420
+ """Number of warm keys currently resident (``<= budget``)."""
421
+ return len(self._warm)
422
+
423
+ # -- measured hit rate + latency math -------------------------------------
424
+ def hit_rate(self) -> float:
425
+ """Measured hit rate ``hits / (hits + misses)`` (``0.0`` before any access)."""
426
+ n = self.hits + self.misses
427
+ return self.hits / n if n else 0.0
428
+
429
+ def expected_latency(self, t_warm: float, t_cold: float) -> float:
430
+ """``E[t] = h·t_warm + (1−h)·t_cold`` at the current measured hit rate ``h``."""
431
+ h = self.hit_rate()
432
+ return h * t_warm + (1.0 - h) * t_cold
433
+
434
+ def speedup(self, t_warm: float, t_cold: float) -> float:
435
+ """How many times faster than always-cold at the current hit rate (``inf`` if free)."""
436
+ e = self.expected_latency(t_warm, t_cold)
437
+ return (t_cold / e) if e > 0 else float("inf")
438
+
439
+ # -- idle-aware re-probe ---------------------------------------------------
440
+ def reprobe_probability(
441
+ self, key: SliceKey, base_eps: float = BASE_EPS, half_life: float = HALF_LIFE_PERIODS
442
+ ) -> float:
443
+ """Idle-aware re-probe probability for warm ``key``.
444
+
445
+ ``periods_since_probe`` is the number of pager accesses since ``key`` was last touched
446
+ (a never-warmed key reads as maximally idle). Rises toward ``1.0`` the longer ``key`` has
447
+ gone unaccessed, so a stale-but-recoverable region gets re-checked. See
448
+ :func:`reprobe_probability`.
449
+ """
450
+ last = self._lastused.get(key)
451
+ idle = float(self._seq - last) if last is not None else float(self._seq)
452
+ return reprobe_probability(idle, base_eps=base_eps, half_life=half_life)
453
+
454
+ def should_reprobe(
455
+ self,
456
+ key: SliceKey,
457
+ rng_uniform: float,
458
+ base_eps: float = BASE_EPS,
459
+ half_life: float = HALF_LIFE_PERIODS,
460
+ ) -> bool:
461
+ """Whether to re-probe warm ``key`` now; ``rng_uniform`` is a draw in ``[0,1)``."""
462
+ return rng_uniform < self.reprobe_probability(
463
+ key, base_eps=base_eps, half_life=half_life
464
+ )
465
+
466
+ # -- depth-cap-1 grounding -------------------------------------------------
467
+ def ground(self, slice_: Slice, *, contradicts_hard_fact: bool = False) -> Grounding:
468
+ """Depth-cap-1 grounding verdict for a paged-back ``slice_``.
469
+
470
+ A slice that came from the pool has provenance (it was encoded and externalized from real
471
+ prior context), so it PASSes unless it contradicts a *hard fact*. Merely disagreeing with
472
+ recent context is not a flag. See :func:`grounding_verdict`.
473
+ """
474
+ has_provenance = bool(slice_.text) or bool(slice_.meta)
475
+ return grounding_verdict(
476
+ has_provenance=has_provenance, contradicts_hard_fact=contradicts_hard_fact
477
+ )
478
+
479
+ # -- convenience -----------------------------------------------------------
480
+ def prefetch_many(self, items: Iterable[tuple[SliceKey, str]]) -> int:
481
+ """Warm several ``(key, reasoning_text)`` pairs; return how many keys ended up warm."""
482
+ warmed = 0
483
+ for key, text in items:
484
+ if self.prefetch_from(key, text):
485
+ warmed += 1
486
+ return warmed
487
+
488
+
489
+ __all__ = [
490
+ "Pager",
491
+ "SliceKey",
492
+ "Grounding",
493
+ "grounding_verdict",
494
+ "reprobe_probability",
495
+ "should_reprobe",
496
+ "RetrieveFn",
497
+ "DEFAULT_WARM_BUDGET",
498
+ "BASE_EPS",
499
+ "HALF_LIFE_PERIODS",
500
+ "MAX_CORRECTION_DEPTH",
501
+ ]
@@ -0,0 +1,64 @@
1
+ # aether-context (Unlimited Context)
2
+ # Copyright (c) 2026 Aether AI
3
+ # SPDX-License-Identifier: Apache-2.0
4
+ """Token-count seam.
5
+
6
+ Budget math across the engine is backend-agnostic, so by default we estimate token
7
+ counts with the ``CHARS_PER_TOKEN = 4`` rule (``len(text) // 4``). When a
8
+ backend exposes a real tokenizer (llama.cpp, HF), ``from_backend`` prefers it for more
9
+ accurate budgeting and falls back — fail-soft — to the estimate if the backend has no
10
+ counter or its counter raises.
11
+ """
12
+ from __future__ import annotations
13
+
14
+ from typing import Callable, Optional, Protocol, runtime_checkable
15
+
16
+ from aether_context._log import get_logger
17
+
18
+ _log = get_logger(__name__)
19
+
20
+ #: Average characters per token, used for backend-agnostic budgeting.
21
+ CHARS_PER_TOKEN = 4
22
+
23
+
24
+ def estimate(text: str) -> int:
25
+ """Estimate token count as ``len(text) // CHARS_PER_TOKEN``.
26
+
27
+ Empty string → 0. Any non-empty string costs at least 1 token so short fragments are
28
+ never undercounted to zero in budget math. Monotonic non-decreasing in length.
29
+ """
30
+ if not isinstance(text, str):
31
+ raise TypeError(f"estimate() expects str, got {type(text).__name__}")
32
+ if not text:
33
+ return 0
34
+ return max(1, len(text) // CHARS_PER_TOKEN)
35
+
36
+
37
+ @runtime_checkable
38
+ class _Counter(Protocol):
39
+ def count_tokens(self, text: str) -> int: ...
40
+
41
+
42
+ def from_backend(model: Optional[object]) -> Callable[[str], int]:
43
+ """Return a ``count(text) -> int`` callable, preferring the backend's tokenizer.
44
+
45
+ If ``model`` exposes a working ``count_tokens(text)`` method, that is used. Otherwise
46
+ (no method, ``None`` model, or the backend counter raising) we fall back to
47
+ :func:`estimate` so budget math always works.
48
+ """
49
+ if isinstance(model, _Counter):
50
+ counter = model # narrowed by the Protocol check
51
+
52
+ def _count(text: str) -> int:
53
+ try:
54
+ return int(counter.count_tokens(text))
55
+ except Exception as exc: # fail-soft: never let budgeting break the run
56
+ _log.debug("backend count_tokens failed, using estimate: %s", exc)
57
+ return estimate(text)
58
+
59
+ return _count
60
+
61
+ return estimate
62
+
63
+
64
+ __all__ = ["estimate", "from_backend", "CHARS_PER_TOKEN"]