"""Three tool-call-memory arms. An arm exposes two operations: - draft(query, functions, user_id) -> canonical draft string - observe(query, functions, user_id, target_name, target_args) -> None record the genuine target so future drafts can reuse it. Arms: 1. NoMemory -- schema-only draft (best you can do with zero history). 2. StaticGlobal -- ToolSpec-style: one global datastore, built during warmup and then FROZEN (no growth, no eviction, no per-user view). 3. PersonalMemory (ours) -- per-user store that grows across sessions, with an eviction policy (LRU or LFU) capping per-user size, and personalized retrieval (query only the user's own store, backing off to a small shared store when empty). Retrieval is top-1 cosine similarity over sentence-transformer embeddings of the query text. The drafted call is the canonicalized (name, arguments) of the most similar past observation. """ from __future__ import annotations import re from collections import OrderedDict from dataclasses import dataclass, field from typing import Any import numpy as np from .metrics import canonical_call_str # --------------------------------------------------------------------------- # # Embedding backend # --------------------------------------------------------------------------- # class Embedder: def __init__(self, name: str = "sentence-transformers/all-MiniLM-L6-v2"): from sentence_transformers import SentenceTransformer self.model = SentenceTransformer(name, device="cpu") self._cache: dict[str, np.ndarray] = {} def embed(self, text: str) -> np.ndarray: v = self._cache.get(text) if v is None: v = self.model.encode(text, normalize_embeddings=True) v = np.asarray(v, dtype=np.float32) self._cache[text] = v return v def schema_draft(functions: list[dict[str, Any]]) -> str: """Zero-history draft: the first tool's signature with placeholder args.""" if not functions: return "" f = functions[0] props = (f.get("parameters", {}) or {}).get("properties", {}) or {} required = (f.get("parameters", {}) or {}).get("required", []) or list(props) args = {k: None for k in required} return canonical_call_str(f["name"], args) # --------------------------------------------------------------------------- # # Datastore entry # --------------------------------------------------------------------------- # @dataclass class Entry: emb: np.ndarray call: str # canonical target call string freq: int = 1 # times this (query-region) has been reinforced def _best_match(emb: np.ndarray, entries: list[Entry]) -> tuple[int, float]: if not entries: return -1, -1.0 mat = np.stack([e.emb for e in entries]) # (N, d), rows unit-norm sims = mat @ emb # cosine (emb unit-norm) i = int(np.argmax(sims)) return i, float(sims[i]) # --------------------------------------------------------------------------- # # Arms # --------------------------------------------------------------------------- # class NoMemory: name = "no_memory" def draft(self, query, functions, user_id, embedder) -> str: return schema_draft(functions) def observe(self, *a, **k): return None class StaticGlobal: """ToolSpec-style frozen global datastore.""" name = "static_global" def __init__(self): self.entries: list[Entry] = [] self.frozen = False def freeze(self): self.frozen = True def draft(self, query, functions, user_id, embedder) -> str: emb = embedder.embed(query) i, _ = _best_match(emb, self.entries) if i < 0: return schema_draft(functions) return self.entries[i].call def observe(self, query, functions, user_id, name, args, embedder) -> None: if self.frozen: return self.entries.append(Entry(embedder.embed(query), canonical_call_str(name, args))) class PersonalMemory: """Ours: per-user growing store + eviction + personalization.""" name = "personal_memory" def __init__(self, capacity: int = 32, eviction: str = "lru", sim_threshold: float = 0.35): self.capacity = capacity self.eviction = eviction # "lru" | "lfu" self.sim_threshold = sim_threshold # per-user ordered store (insertion/most-recent-use order for LRU) self.stores: dict[str, "OrderedDict[int, Entry]"] = {} self.shared: list[Entry] = [] # cold-start backoff self._next_id = 0 def _store(self, user_id) -> "OrderedDict[int, Entry]": return self.stores.setdefault(user_id, OrderedDict()) def draft(self, query, functions, user_id, embedder) -> str: emb = embedder.embed(query) store = self._store(user_id) entries = list(store.values()) i, sim = _best_match(emb, entries) if i >= 0 and sim >= self.sim_threshold: key = list(store.keys())[i] entry = store[key] if self.eviction == "lru": # mark most-recently-used store.move_to_end(key) entry.freq += 1 return entry.call # personal store empty / too dissimilar -> shared backoff j, sj = _best_match(emb, self.shared) if j >= 0 and sj >= self.sim_threshold: return self.shared[j].call return schema_draft(functions) def observe(self, query, functions, user_id, name, args, embedder) -> None: emb = embedder.embed(query) call = canonical_call_str(name, args) store = self._store(user_id) eid = self._next_id self._next_id += 1 store[eid] = Entry(emb, call) store.move_to_end(eid) self._evict(store) def seed_shared(self, query, name, args, embedder) -> None: self.shared.append(Entry(embedder.embed(query), canonical_call_str(name, args))) def _evict(self, store: "OrderedDict[int, Entry]") -> None: while len(store) > self.capacity: if self.eviction == "lru": store.popitem(last=False) # drop least-recently-used elif self.eviction == "lfu": k = min(store, key=lambda x: store[x].freq) del store[k] else: store.popitem(last=False) def total_entries(self) -> int: return sum(len(s) for s in self.stores.values()) # --------------------------------------------------------------------------- # # Ablation arms (2x2: personalized? x evicting?) -- see run_ablation.py. # These decompose PersonalMemory's gain over StaticGlobal into the contribution # of per-user partitioning vs. the contribution of online-growth+eviction. # --------------------------------------------------------------------------- # class PersonalNoEvict(PersonalMemory): """[+personalization, -eviction]: per-user store that grows online but is never bounded/evicted. Isolates how much of ours' gain is eviction.""" name = "personal_noevict" def __init__(self, sim_threshold: float = 0.35): # capacity = +inf so _evict() never fires. super().__init__(capacity=10**9, eviction="lru", sim_threshold=sim_threshold) class GlobalEvict: """[-personalization, +online-growth+eviction]: a single GLOBAL store (not per-user) that keeps ingesting after warmup and LRU-evicts at a total capacity. Isolates how much of ours' gain is per-user partitioning: it differs from PersonalMemory only in that retrieval ignores user id.""" name = "global_evict" def __init__(self, capacity: int = 1920, sim_threshold: float = 0.35): self.capacity = capacity self.sim_threshold = sim_threshold self.store: "OrderedDict[int, Entry]" = OrderedDict() self._next_id = 0 def draft(self, query, functions, user_id, embedder) -> str: emb = embedder.embed(query) entries = list(self.store.values()) i, sim = _best_match(emb, entries) if i >= 0 and sim >= self.sim_threshold: key = list(self.store.keys())[i] self.store.move_to_end(key) # LRU touch self.store[key].freq += 1 return self.store[key].call return schema_draft(functions) def observe(self, query, functions, user_id, name, args, embedder) -> None: emb = embedder.embed(query) eid = self._next_id self._next_id += 1 self.store[eid] = Entry(emb, canonical_call_str(name, args)) self.store.move_to_end(eid) while len(self.store) > self.capacity: self.store.popitem(last=False) # drop least-recently-used def seed_shared(self, *a, **k): return None # --------------------------------------------------------------------------- # # Phase 4.4 — faithful stronger ToolSpec-style baseline # --------------------------------------------------------------------------- # def _schema_scaffold(functions, name) -> str: """Schema-aware structural draft for a NAMED function: the function's required argument keys in canonical (sorted) order with placeholder values. This is the structurally-valid fallback a schema-aware FSM emits when retrieval is not confident enough to commit a concrete prior call.""" fn = next((f for f in functions if f.get("name") == name), None) if fn is None: return schema_draft(functions) params = fn.get("parameters", {}) or {} props = params.get("properties", {}) or {} required = params.get("required", []) or list(props) return canonical_call_str(name, {k: None for k in required}) class ToolSpecBaseline: """Faithful ToolSpec reproduction (Xia et al., 2026): a *frozen global* retrieval store (no eviction, no personalization — the ToolSpec regime) with the two ToolSpec mechanisms the simple ``StaticGlobal`` proxy omits: 1. **Confidence-gated retrieval.** Return the nearest stored call verbatim only while its similarity clears ``sim_lo``; ``StaticGlobal`` instead returns its single nearest neighbour unconditionally, so on a cold / far query it drafts a wholly unrelated call. 2. **Schema-aware fallback (FSM surrogate).** On a cold miss, rather than emitting a random far neighbour we emit a *structurally valid* draft for the nearest neighbour's function (its required-arg scaffold in canonical order) — the acceptance a schema-constrained decoder guarantees on the call's structural tokens even without a value hit. This makes the arm **strictly at least as strong as ``StaticGlobal``**: identical on confident hits, better on cold misses. ToolSpec has no public code, so the FSM is approximated by this schema-aware scaffold. We also tested a ``k``-NN summed-similarity vote on the target *function* (retrieval-augmented denoising); it *degraded* MAT on these traces because the highly skewed telecom workload (one diagnostic call dominates) lets the majority function override correct top-1 picks — reported honestly in the write-up, and NOT used here. Everything else (frozen, global, eviction-free) matches ToolSpec and is deliberately NOT personalized — the property under test. """ name = "toolspec" def __init__(self, sim_lo: float = 0.30): self.entries: list[Entry] = [] self.frozen = False self.sim_lo = sim_lo self._names: list[str] = [] # parallel function name per entry def freeze(self): self.frozen = True def draft(self, query, functions, user_id, embedder) -> str: if not self.entries: return schema_draft(functions) emb = embedder.embed(query) i, sim = _best_match(emb, self.entries) if sim >= self.sim_lo: return self.entries[i].call # confident retrieval hit return _schema_scaffold(functions, self._names[i]) # schema-aware miss def observe(self, query, functions, user_id, name, args, embedder) -> None: if self.frozen: return self.entries.append(Entry(embedder.embed(query), canonical_call_str(name, args))) self._names.append(name) def seed_shared(self, *a, **k): return None # --------------------------------------------------------------------------- # # SuffixDecoding baseline (Oliaro et al., NeurIPS 2025; arXiv 2411.04975) # --------------------------------------------------------------------------- # _TOK_RE = re.compile(r"\w+|[^\w\s]") _SEP = " " def _tokenize(text: str) -> tuple[str, ...]: """Word/punctuation-level tokens (lowercased). A deliberate approximation of SuffixDecoding's model-BPE tokens: the mechanistic contrast under test is *exact token matching vs. embedding similarity*, which this preserves; the exact subword vocabulary is not what distinguishes the two arms.""" return tuple(_TOK_RE.findall(text.lower())) def _suffix_key(tokens: tuple[str, ...]) -> str: """Separator-delimited form so whole-token substring tests never match across partial tokens (every boundary is a _SEP).""" return _SEP + _SEP.join(tokens) + _SEP def _longest_suffix_match(q: tuple[str, ...], stored_key: str, floor: int) -> int: """Longest k>floor such that the k-token *suffix* of the current query q is a contiguous whole-token substring of a stored sequence (its suffix key). Returns 0 if no suffix longer than `floor` matches. This is SuffixDecoding's 'walk the tree to the node matching the context suffix' step, adapted to our per-request query context. Uses binary search: the predicate ``q[-k:] is a substring of stored`` is monotonic in k (if the k-token suffix matches, every shorter suffix does), so we find the largest matching k in O(log|q|) containment tests rather than O(|q|) — essential because dialogue-context queries run to hundreds of tokens.""" best = 0 lo, hi = floor + 1, len(q) while lo <= hi: mid = (lo + hi) // 2 cand = _SEP + _SEP.join(q[-mid:]) + _SEP if cand in stored_key: best = mid lo = mid + 1 else: hi = mid - 1 return best class SuffixDecodingBaseline: """Faithful adaptation of **SuffixDecoding** (Oliaro et al., *SuffixDecoding: Extreme Speculative Decoding for Emerging AI Applications*, NeurIPS 2025 Spotlight; arXiv 2411.04975) as a retrieval arm. SuffixDecoding keeps a **global suffix tree accumulated from previous requests' token streams** (live/growing across the deployment, so request N benefits from request N-1), matches the **suffix of the current context** against the tree at each step, and speculates the highest-frequency continuation with adaptive length. It is **token-level (exact match), global, and NOT personalized**; the tree is **size-capped** (~10.75 B/token, ~31 days on a 144 GB host), not unbounded. We reproduce that regime and swap **exactly one mechanism** vs. our own ``GlobalEvict`` arm: retrieval is **longest-token-suffix match** on the query instead of embedding cosine similarity. Everything else — global live write-back, canonicalization, a size cap, no per-user partitioning — is held identical, so the comparison isolates *retrieval mechanism* (exact token match vs. semantic embedding), not confounds like different data pools or personalization. Frequency and recency break ties in match length, mirroring SuffixDecoding's frequency-ranked tree paths. Faithful vs. adapted: - Faithful: global, live-growing, size-capped, non-personalized store; exact token-suffix matching; frequency-ranked selection. - Adapted: we match the *query* token context (which selects a tool call in our one-call-per-request setting) rather than a running generation, and our token-LCP acceptance already truncates the speculated call at the first mismatch, subsuming SuffixDecoding's adaptive speculation length. Tokenization is word/punct level, not the model BPE (approximation). """ name = "suffixdecoding" def __init__(self, capacity: int = 1920, min_match: int = 1): # capacity matches GlobalEvict's total footprint (U*C = 40*48) for a # same-size comparison; min_match=1 lets any shared suffix token yield a # (frequency-ranked) speculation, as SuffixDecoding's tree always does. self.capacity = capacity self.min_match = min_match # eid -> [q_tokens, suffix_key, call, freq] self.store: "OrderedDict[int, list]" = OrderedDict() self._next_id = 0 def draft(self, query, functions, user_id, embedder) -> str: q = _tokenize(query) if not q: return schema_draft(functions) floor = self.min_match - 1 best_key, best_eid = None, None # rank by (match_len, freq), then recency for eid, rec in self.store.items(): m = _longest_suffix_match(q, rec[1], floor) # true length (fixed floor) if m < self.min_match: continue key = (m, rec[3]) if best_key is None or key >= best_key: # >= => later (more recent) wins ties best_key, best_eid = key, eid if best_eid is None: return schema_draft(functions) rec = self.store[best_eid] rec[3] += 1 # frequency reinforcement self.store.move_to_end(best_eid) # recency touch return rec[2] def observe(self, query, functions, user_id, name, args, embedder) -> None: q = _tokenize(query) eid = self._next_id self._next_id += 1 self.store[eid] = [q, _suffix_key(q), canonical_call_str(name, args), 1] self.store.move_to_end(eid) while len(self.store) > self.capacity: self.store.popitem(last=False) # drop oldest (size cap) def seed_shared(self, *a, **k): return None def total_entries(self) -> int: return len(self.store)