From 96deb1f6af1d3fc7336b1ae2318b97c9f33786fd Mon Sep 17 00:00:00 2001 From: soltrinox Date: Sat, 5 Sep 2026 00:59:03 -0700 Subject: [PATCH] feat(engine): CC-6/7/8/10 tokens, hop_legal, bundle, quantization (M4) --- docs/HOOK_CONTRACT.md | 11 + docs/HOP_LEGAL.md | 27 +++ engine/env.example | 9 + engine/src/chat_compressor/bundle.py | 326 +++++++++++++++++++++++++++ engine/src/chat_compressor/handle.py | 31 +++ engine/src/chat_compressor/store.py | 93 +++++++- engine/src/chat_compressor/tokens.py | 84 +++++++ engine/tests/test_cc10_quant.py | 67 ++++++ engine/tests/test_cc6_tokens.py | 64 ++++++ engine/tests/test_cc7_hop_legal.py | 74 ++++++ engine/tests/test_cc8_bundle.py | 92 ++++++++ 11 files changed, 876 insertions(+), 2 deletions(-) create mode 100644 docs/HOP_LEGAL.md create mode 100644 engine/src/chat_compressor/bundle.py create mode 100644 engine/src/chat_compressor/tokens.py create mode 100644 engine/tests/test_cc10_quant.py create mode 100644 engine/tests/test_cc6_tokens.py create mode 100644 engine/tests/test_cc7_hop_legal.py create mode 100644 engine/tests/test_cc8_bundle.py diff --git a/docs/HOOK_CONTRACT.md b/docs/HOOK_CONTRACT.md index 1c6f7e8..43d9d02 100644 --- a/docs/HOOK_CONTRACT.md +++ b/docs/HOOK_CONTRACT.md @@ -43,3 +43,14 @@ Optional file-based handoff from an external router (comPASS). When present and - **Missing, stale, or corrupt advisory MUST NOT block Agent Chat** — hook still returns the event-safe default (`continue: true` / empty context) - Hook process never loads provider keys to read the advisory file +## Hop legality (CC-7) + +`PersistentAgentHandle.hop_legal()` must be consulted before changing +`recipient_id` mid-session. Returns false with pending tool state. Details: +[`HOP_LEGAL.md`](./HOP_LEGAL.md). + +## Tensor quantization (CC-10) + +Optional env `CHAT_COMPRESSOR_TENSOR_QUANT` (`float32` default, or `float16` / +`int8`). Scheme is recorded on `StateNode.meta.quantization`. Default path is +unchanged float32 mmap behavior. diff --git a/docs/HOP_LEGAL.md b/docs/HOP_LEGAL.md new file mode 100644 index 0000000..e1ac523 --- /dev/null +++ b/docs/HOP_LEGAL.md @@ -0,0 +1,27 @@ +# Hop legality (`hop_legal`) + +CC-7 exposes `PersistentAgentHandle.hop_legal()` as a scheduling constraint for +model hops (comPASS Tier 4). + +## Rules + +- **Legal** only at turn boundaries with **no pending tool state**. +- Returns **False** when: + - `handle.set_pending_tool(True)` was called and not cleared, or + - latest `StateNode.meta["pending_tool"]` is `true`, or + - `meta["tool_status"]` is one of `pending`, `in_flight`, `awaiting`, `tool_pending`. +- Returns **True** when tool state is unknown/`stub` and no pending flag is set + (preserves 0.2.0 hop permissiveness). + +## Usage + +```python +if handle.hop_legal(): + # safe to change recipient_id on the next step + ... +else: + # defer hop until tools settle; clear with handle.clear_pending_tool() + ... +``` + +See also `docs/HOOK_CONTRACT.md` and prototype §14.2 CC-7. diff --git a/engine/env.example b/engine/env.example index 0381a93..23b7669 100644 --- a/engine/env.example +++ b/engine/env.example @@ -51,3 +51,12 @@ PRODUCER_B_DIM=128 # Missing/stale/corrupt files are ignored (fail-open; never blocks Agent Chat). # CHAT_COMPRESSOR_ADVISORY_PATH=advisory/latest.json +# CC-6: packing always uses metrics.estimate_tokens (chars/4). Accurate counters +# for cost are registered in chat_compressor.tokens (per tokenizer_id). + +# CC-7: hop_legal() is false while pending tool state is set on the handle or +# StateNode.meta (pending_tool=true / tool_status in pending|in_flight|awaiting). + +# CC-10: optional tensor quantization for StateStore blobs (default float32). +# float16 | int8 reduce transfer size; reconstruction must meet cosine budget. +# CHAT_COMPRESSOR_TENSOR_QUANT=float32 diff --git a/engine/src/chat_compressor/bundle.py b/engine/src/chat_compressor/bundle.py new file mode 100644 index 0000000..90bfba0 --- /dev/null +++ b/engine/src/chat_compressor/bundle.py @@ -0,0 +1,326 @@ +"""CC-8: portable state bundle export/import (prototype §15). + +Format (bundle.v1/): + manifest.json, graph.json, states/*.safetensors, + inject_ledger.json, lineage.json + +Producer mismatch is reported explicitly — never silent embedding reuse. +""" + +from __future__ import annotations + +import hashlib +import json +import shutil +from dataclasses import dataclass, field +from datetime import datetime, timezone +from pathlib import Path +from typing import Any, Literal + +import numpy as np +from safetensors import safe_open + +from chat_compressor.store import ( + StateNode, + StateStore, + _read_inject_doc, + _write_inject_doc, +) + +BUNDLE_SCHEMA = "bundle.v1" +ImportMode = Literal["full", "graph_only", "reproject_required"] + + +def _now() -> str: + return datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ") + + +def _sha256_file(path: Path) -> str: + h = hashlib.sha256() + with path.open("rb") as fh: + for chunk in iter(lambda: fh.read(65536), b""): + h.update(chunk) + return h.hexdigest() + + +@dataclass +class BundleManifest: + schema: str = BUNDLE_SCHEMA + version: str = "1" + producer: str = "" + d: int = 0 + k_max: int = 0 + tokenizer_id: str = "hashed-ngram" + quantization: str = "float32" + agent_id: str = "" + lineage_head: str | None = None + created_at: str = "" + checksums: dict[str, str] = field(default_factory=dict) + + def to_dict(self) -> dict[str, Any]: + return { + "schema": self.schema, + "version": self.version, + "producer": self.producer, + "d": self.d, + "k_max": self.k_max, + "tokenizer_id": self.tokenizer_id, + "quantization": self.quantization, + "agent_id": self.agent_id, + "lineage_head": self.lineage_head, + "created_at": self.created_at, + "checksums": dict(self.checksums), + } + + @classmethod + def from_dict(cls, data: dict[str, Any]) -> BundleManifest: + return cls( + schema=str(data.get("schema") or BUNDLE_SCHEMA), + version=str(data.get("version") or "1"), + producer=str(data.get("producer") or ""), + d=int(data.get("d") or 0), + k_max=int(data.get("k_max") or 0), + tokenizer_id=str(data.get("tokenizer_id") or "hashed-ngram"), + quantization=str(data.get("quantization") or "float32"), + agent_id=str(data.get("agent_id") or ""), + lineage_head=data.get("lineage_head"), + created_at=str(data.get("created_at") or ""), + checksums=dict(data.get("checksums") or {}), + ) + + +@dataclass +class ImportResult: + mode: ImportMode + agent_id: str + states_imported: int + graph_imported: bool + ledger_imported: bool + producer_matched: bool + notes: str = "" + + +def export_bundle( + store: StateStore, + agent_id: str, + dest: str | Path, + *, + tokenizer_id: str | None = None, +) -> Path: + """Export agent state into a bundle.v1 directory. Returns dest path.""" + out = Path(dest) + if out.exists(): + shutil.rmtree(out) + states_dir = out / "states" + states_dir.mkdir(parents=True, exist_ok=True) + + lineage = store.lineage(agent_id) + if not lineage: + raise ValueError(f"no states for agent_id={agent_id!r}") + + agent_dir = Path(store.root) / agent_id + head = lineage[-1] + tokenizer = tokenizer_id or str((head.meta or {}).get("tokenizer_id") or "hashed-ngram") + quant = str((head.meta or {}).get("quantization") or "float32") + + k_max = head.k + with store._connect() as conn: + row = conn.execute( + "SELECT producer, d, k_max FROM agents WHERE agent_id = ?", (agent_id,) + ).fetchone() + producer = head.producer + d = head.d + if row is not None: + producer = row["producer"] or producer + d = int(row["d"]) + k_max = int(row["k_max"]) + + checksums: dict[str, str] = {} + lineage_rows: list[dict[str, Any]] = [] + + for node in lineage: + src = Path(node.blob_path) + name = f"t{node.t:04d}.safetensors" + dst = states_dir / name + if src.is_file(): + shutil.copy2(src, dst) + checksums[f"states/{name}"] = _sha256_file(dst) + spans = src.with_name(src.stem + ".spans.json") + if spans.is_file(): + spans_dst = states_dir / spans.name + shutil.copy2(spans, spans_dst) + checksums[f"states/{spans.name}"] = _sha256_file(spans_dst) + lineage_rows.append( + { + "state_id": node.state_id, + "parent_id": node.parent_id, + "t": node.t, + "producer": node.producer, + "d": node.d, + "k": node.k, + "blob": name, + "meta": dict(node.meta or {}), + "created_at": node.created_at, + } + ) + + graph_src = agent_dir / "graph.json" + if graph_src.is_file(): + graph_dst = out / "graph.json" + shutil.copy2(graph_src, graph_dst) + checksums["graph.json"] = _sha256_file(graph_dst) + else: + (out / "graph.json").write_text("{}\n", encoding="utf-8") + checksums["graph.json"] = _sha256_file(out / "graph.json") + + ledger_doc = _read_inject_doc(agent_dir) + ledger_path = out / "inject_ledger.json" + ledger_path.write_text( + json.dumps(ledger_doc, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + checksums["inject_ledger.json"] = _sha256_file(ledger_path) + + lineage_path = out / "lineage.json" + lineage_path.write_text( + json.dumps({"agent_id": agent_id, "states": lineage_rows}, indent=2) + "\n", + encoding="utf-8", + ) + checksums["lineage.json"] = _sha256_file(lineage_path) + + manifest = BundleManifest( + producer=str(producer), + d=int(d), + k_max=int(k_max), + tokenizer_id=tokenizer, + quantization=quant, + agent_id=agent_id, + lineage_head=head.state_id, + created_at=_now(), + checksums=checksums, + ) + (out / "manifest.json").write_text( + json.dumps(manifest.to_dict(), indent=2) + "\n", encoding="utf-8" + ) + return out + + +def import_bundle( + src: str | Path, + store: StateStore, + *, + agent_id: str | None = None, + expected_producer: str | None = None, + expected_d: int | None = None, +) -> ImportResult: + """Import a bundle.v1 directory into store. + + When producer/d mismatch, imports graph + ledger but skips tensor blobs + (mode=graph_only) and reports the mismatch. + """ + root = Path(src) + man_path = root / "manifest.json" + if not man_path.is_file(): + raise FileNotFoundError(f"missing manifest.json in {root}") + manifest = BundleManifest.from_dict(json.loads(man_path.read_text(encoding="utf-8"))) + aid = agent_id or manifest.agent_id + if not aid: + raise ValueError("agent_id required") + + producer_ok = True + notes: list[str] = [] + if expected_producer is not None and expected_producer != manifest.producer: + producer_ok = False + notes.append( + f"producer mismatch: bundle={manifest.producer!r} expected={expected_producer!r}" + ) + if expected_d is not None and int(expected_d) != int(manifest.d): + producer_ok = False + notes.append(f"d mismatch: bundle={manifest.d} expected={expected_d}") + + lineage_path = root / "lineage.json" + lineage_doc = json.loads(lineage_path.read_text(encoding="utf-8")) + states_meta = list(lineage_doc.get("states") or []) + + agent_dir = Path(store.root) / aid + agent_dir.mkdir(parents=True, exist_ok=True) + + graph_imported = False + graph_src = root / "graph.json" + if graph_src.is_file(): + shutil.copy2(graph_src, agent_dir / "graph.json") + graph_imported = True + + ledger_imported = False + ledger_src = root / "inject_ledger.json" + if ledger_src.is_file(): + try: + doc = json.loads(ledger_src.read_text(encoding="utf-8")) + if isinstance(doc, dict): + _write_inject_doc(agent_dir, doc) + ledger_imported = True + elif isinstance(doc, list): + _write_inject_doc(agent_dir, {"turns": doc, "recipients": {}}) + ledger_imported = True + except (OSError, json.JSONDecodeError) as exc: + notes.append(f"ledger import skipped: {exc}") + + states_imported = 0 + store.ensure_agent(aid, manifest.producer, manifest.d, manifest.k_max) + if not producer_ok: + notes.append("tensors skipped due to producer mismatch; reproject_required") + return ImportResult( + mode="graph_only", + agent_id=aid, + states_imported=0, + graph_imported=graph_imported, + ledger_imported=ledger_imported, + producer_matched=False, + notes="; ".join(notes), + ) + + parent: StateNode | None = None + states_dir = root / "states" + for row in states_meta: + blob_name = str(row.get("blob") or f"t{int(row['t']):04d}.safetensors") + blob_src = states_dir / blob_name + if not blob_src.is_file(): + notes.append(f"missing blob {blob_name}") + continue + tensors: dict[str, Any] = {} + with safe_open(str(blob_src), framework="np") as handle: + for key in handle.keys(): + tensors[key] = handle.get_tensor(key) + C = np.asarray(tensors["C"]) + M = tensors.get("M") + KV = tensors.get("KV") + meta = dict(row.get("meta") or {}) + node = store.save( + agent_id=aid, + C=C, + M=M, + parent=parent, + producer=str(row.get("producer") or manifest.producer), + graph_path=agent_dir / "graph.json", + KV=KV, + meta=meta, + k_max=manifest.k_max, + ) + spans_src = states_dir / (Path(blob_name).stem + ".spans.json") + if spans_src.is_file(): + shutil.copy2( + spans_src, + Path(node.blob_path).with_name(Path(node.blob_path).stem + ".spans.json"), + ) + parent = node + states_imported += 1 + + return ImportResult( + mode="full", + agent_id=aid, + states_imported=states_imported, + graph_imported=graph_imported, + ledger_imported=ledger_imported, + producer_matched=True, + notes="; ".join(notes), + ) diff --git a/engine/src/chat_compressor/handle.py b/engine/src/chat_compressor/handle.py index 170c309..15253d4 100644 --- a/engine/src/chat_compressor/handle.py +++ b/engine/src/chat_compressor/handle.py @@ -81,6 +81,8 @@ def __init__( self._turn_index = 0 self._last_graph_path: str | None = None self.last_sample_ms: float = 0.0 + # CC-7: in-memory pending-tool flag (hop illegal while set). + self._pending_tool: bool = False def _agent_dir(self) -> Path: return Path(self.store.root) / self.agent_id @@ -315,6 +317,35 @@ def _last_user_query(self) -> str: turns.sort(key=lambda n: (n.valid_start, n.attrs.get("index", 0))) return (turns[-1].summary or "").strip() + + def set_pending_tool(self, pending: bool = True) -> None: + """Mark whether a tool call is in flight (CC-7 hop gate).""" + self._pending_tool = bool(pending) + + def clear_pending_tool(self) -> None: + """Clear pending-tool flag at a clean turn boundary.""" + self._pending_tool = False + + def hop_legal(self) -> bool: + """Return False when a hop would cross pending tool state (CC-7). + + Legal only at turn boundaries with no pending tool state. + Default: tool_status stub/unknown and no pending flag ⇒ True + (do not block hops that 0.2.0 would have allowed). + """ + if self._pending_tool: + return False + node = self.latest() + if node is None: + return True + meta = node.meta or {} + if meta.get("pending_tool") is True: + return False + status = str(meta.get("tool_status") or "").strip().lower() + if status in {"pending", "in_flight", "awaiting", "tool_pending"}: + return False + return True + def expand_spans(self, query: str, k: int = 4) -> list[str]: """Local-only: nearest verbatim chunks from tNNNN.spans.json sidecars.""" texts: list[str] = [] diff --git a/engine/src/chat_compressor/store.py b/engine/src/chat_compressor/store.py index 1c1d4d0..447f66f 100644 --- a/engine/src/chat_compressor/store.py +++ b/engine/src/chat_compressor/store.py @@ -3,6 +3,7 @@ from __future__ import annotations import json +import os import sqlite3 import uuid from dataclasses import dataclass, field @@ -53,6 +54,83 @@ def _new_state_id() -> str: # Optional StateNode.meta keys (CC-1 / comPASS routing attribution). Absent ⇒ 0.2.0 behavior. RECIPIENT_META_KEYS = ("recipient_id", "recipient_version", "route_decision_id") +# --- CC-10 optional tensor quantization --------------------------------- +QUANT_FLOAT32 = "float32" +QUANT_FLOAT16 = "float16" +QUANT_INT8 = "int8" +_VALID_QUANTS = {QUANT_FLOAT32, QUANT_FLOAT16, QUANT_INT8} + +# Behavioral reconstruction budget: mean row cosine similarity vs original. +DEFAULT_RECON_COSINE_BUDGET = 0.99 + + +def tensor_quantization_scheme() -> str: + """Env CHAT_COMPRESSOR_TENSOR_QUANT: float32 (default) | float16 | int8.""" + raw = (os.environ.get("CHAT_COMPRESSOR_TENSOR_QUANT") or QUANT_FLOAT32).strip().lower() + if raw in {"fp16", "float16", "f16"}: + return QUANT_FLOAT16 + if raw in {"int8", "i8", "qint8"}: + return QUANT_INT8 + return QUANT_FLOAT32 + + +def quantize_C(C: np.ndarray, scheme: str) -> tuple[dict[str, np.ndarray], dict[str, Any]]: + """Quantize C for storage. Returns (tensors_extra_or_override, meta_fields). + + float32: no change (caller writes C as float32). + float16: stores C as float16. + int8: symmetric per-row scale; stores C_int8 + C_scale (float32, shape k). + """ + arr = np.asarray(C, dtype=np.float32) + if arr.ndim == 1: + arr = arr[None, :] + scheme = scheme if scheme in _VALID_QUANTS else QUANT_FLOAT32 + meta = {"quantization": scheme} + if scheme == QUANT_FLOAT32: + return {"C": arr}, meta + if scheme == QUANT_FLOAT16: + return {"C": arr.astype(np.float16)}, meta + # int8 symmetric per-row + absmax = np.max(np.abs(arr), axis=1).astype(np.float32) + absmax = np.where(absmax < 1e-12, 1.0, absmax) + scale = (absmax / 127.0).astype(np.float32) + q = np.clip(np.round(arr / scale[:, None]), -127, 127).astype(np.int8) + return {"C": q, "C_scale": scale}, meta + + +def dequantize_C(tensors: dict[str, np.ndarray], meta: dict[str, Any] | None = None) -> np.ndarray: + """Reconstruct float32 C from possibly quantized tensors.""" + meta = meta or {} + scheme = str(meta.get("quantization") or QUANT_FLOAT32).lower() + C = tensors["C"] + if scheme in {QUANT_FLOAT16, "fp16", "f16"} or C.dtype == np.float16: + return np.asarray(C, dtype=np.float32) + if scheme in {QUANT_INT8, "i8", "qint8"} or C.dtype == np.int8: + scale = tensors.get("C_scale") + if scale is None: + return np.asarray(C, dtype=np.float32) + scale = np.asarray(scale, dtype=np.float32) + return (np.asarray(C, dtype=np.float32) * scale[:, None]).astype(np.float32) + return np.asarray(C, dtype=np.float32) + + +def reconstruction_cosine(original: np.ndarray, reconstructed: np.ndarray) -> float: + """Mean per-row cosine similarity; used for CC-10 acceptance budget.""" + a = np.asarray(original, dtype=np.float32) + b = np.asarray(reconstructed, dtype=np.float32) + if a.ndim == 1: + a = a[None, :] + if b.ndim == 1: + b = b[None, :] + if a.shape != b.shape or a.shape[0] == 0: + return 0.0 + dots = np.sum(a * b, axis=1) + na = np.linalg.norm(a, axis=1) + nb = np.linalg.norm(b, axis=1) + denom = np.maximum(na * nb, 1e-12) + return float(np.mean(dots / denom)) + + @dataclass class StateNode: @@ -130,7 +208,18 @@ def save( agent_dir = self.root / agent_id agent_dir.mkdir(parents=True, exist_ok=True) blob_path = agent_dir / f"t{t:04d}.safetensors" - tensors: dict[str, np.ndarray] = {"C": arr, "M": mask} + # CC-10: optional quantization (default float32 — unchanged local mmap). + # Only record quantization in meta when non-default so absent-meta paths + # stay identical to 0.2.0. + meta_out: dict[str, Any] = dict(meta or {}) + scheme = str(meta_out.get("quantization") or tensor_quantization_scheme()) + q_tensors, q_meta = quantize_C(arr, scheme) + if scheme != QUANT_FLOAT32: + meta_out.update(q_meta) + elif "quantization" in meta_out and meta_out.get("quantization") == QUANT_FLOAT32: + # Explicit float32 from caller — keep; otherwise omit for 0.2.0 parity. + pass + tensors: dict[str, np.ndarray] = {"M": mask, **q_tensors} if KV is not None: tensors["KV"] = np.asarray(KV, dtype=np.float32) save_file(tensors, str(blob_path)) @@ -152,7 +241,7 @@ def save( producer, graph_str, created, - json.dumps(meta or {}), + json.dumps(meta_out), ), ) return self.load(state_id) diff --git a/engine/src/chat_compressor/tokens.py b/engine/src/chat_compressor/tokens.py new file mode 100644 index 0000000..2e2af47 --- /dev/null +++ b/engine/src/chat_compressor/tokens.py @@ -0,0 +1,84 @@ +"""CC-6: pluggable token counters. + +Keep ``metrics.estimate_tokens`` (chars/4) for internal packing budgets. +Use ``count_tokens`` for cost decisions — resolved per recipient tokenizer_id +when a counter is registered, otherwise the cheap estimate as fallback. +""" + +from __future__ import annotations + +from typing import Callable + +from chat_compressor.metrics import estimate_tokens + +TokenCounter = Callable[[str], int] + +_REGISTRY: dict[str, TokenCounter] = {} + + +def register_counter(tokenizer_id: str, counter: TokenCounter) -> None: + """Register an accurate counter for a tokenizer identity.""" + tid = str(tokenizer_id).strip() + if not tid: + raise ValueError("tokenizer_id must be non-empty") + _REGISTRY[tid] = counter + + +def unregister_counter(tokenizer_id: str) -> None: + _REGISTRY.pop(str(tokenizer_id).strip(), None) + + +def clear_counters() -> None: + _REGISTRY.clear() + + +def resolve_counter(tokenizer_id: str | None = None) -> TokenCounter: + """Return registered counter for tokenizer_id, else cheap estimate.""" + if tokenizer_id is not None: + tid = str(tokenizer_id).strip() + if tid and tid in _REGISTRY: + return _REGISTRY[tid] + return estimate_tokens + + +def count_tokens(text: str, tokenizer_id: str | None = None) -> int: + """Accurate-when-available token count for cost; never raises on empty text.""" + if not text: + return 0 + counter = resolve_counter(tokenizer_id) + try: + n = int(counter(text)) + except Exception: # noqa: BLE001 — cost path must fail-open to estimate + return estimate_tokens(text) + return max(0, n) + + +def packing_tokens(text: str) -> int: + """Always the cheap chars/4 estimate used by pack.py budgets.""" + return estimate_tokens(text) + + +def try_hf_counter(tokenizer_id: str) -> TokenCounter | None: + """Best-effort HuggingFace AutoTokenizer counter when [hf] extra is installed. + + Returns None when transformers is unavailable or load fails (NOT_RUN path). + """ + tid = str(tokenizer_id).strip() + if not tid: + return None + try: + from transformers import AutoTokenizer # type: ignore[import-not-found] + except Exception: + return None + try: + tok = AutoTokenizer.from_pretrained(tid, use_fast=True) + except Exception: + return None + + def _count(text: str) -> int: + if not text: + return 0 + ids = tok.encode(text, add_special_tokens=False) + return len(ids) + + return _count diff --git a/engine/tests/test_cc10_quant.py b/engine/tests/test_cc10_quant.py new file mode 100644 index 0000000..9c76a19 --- /dev/null +++ b/engine/tests/test_cc10_quant.py @@ -0,0 +1,67 @@ +"""CC-10: optional tensor quantization with reconstruction budget.""" + +from __future__ import annotations + +from pathlib import Path + +import numpy as np +import pytest + +from chat_compressor.store import ( + DEFAULT_RECON_COSINE_BUDGET, + StateStore, + dequantize_C, + quantize_C, + reconstruction_cosine, + tensor_quantization_scheme, +) + + +def _l2_rows(arr: np.ndarray) -> np.ndarray: + a = np.asarray(arr, dtype=np.float32) + n = np.linalg.norm(a, axis=1, keepdims=True) + n = np.where(n < 1e-12, 1.0, n) + return a / n + + +def test_default_scheme_float32(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("CHAT_COMPRESSOR_TENSOR_QUANT", raising=False) + assert tensor_quantization_scheme() == "float32" + + +def test_int8_reconstruction_within_budget(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("CHAT_COMPRESSOR_TENSOR_QUANT", "int8") + rng = np.random.default_rng(0) + original = _l2_rows(rng.normal(size=(8, 32)).astype(np.float32)) + store = StateStore(tmp_path / "state") + node = store.save(agent_id="q", C=original, producer="embed", k_max=8) + assert node.meta.get("quantization") == "int8" + cos = reconstruction_cosine(original, node.C) + assert cos >= DEFAULT_RECON_COSINE_BUDGET + + +def test_float16_roundtrip(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("CHAT_COMPRESSOR_TENSOR_QUANT", "float16") + original = _l2_rows(np.eye(4, 16, dtype=np.float32)) + store = StateStore(tmp_path / "state") + node = store.save(agent_id="f16", C=original, producer="embed", k_max=4) + assert node.meta.get("quantization") == "float16" + assert reconstruction_cosine(original, node.C) >= DEFAULT_RECON_COSINE_BUDGET + + +def test_float32_unchanged_default(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.delenv("CHAT_COMPRESSOR_TENSOR_QUANT", raising=False) + original = np.eye(3, 8, dtype=np.float32) + store = StateStore(tmp_path / "state") + node = store.save(agent_id="f32", C=original, producer="embed", k_max=3) + # Default float32 omits quantization key for 0.2.0 meta parity. + assert "quantization" not in node.meta + np.testing.assert_allclose(node.C, original, atol=1e-6) + + +def test_quantize_dequantize_helpers() -> None: + original = _l2_rows(np.ones((4, 8), dtype=np.float32)) + tensors, meta = quantize_C(original, "int8") + assert meta["quantization"] == "int8" + recon = dequantize_C(tensors, meta) + assert reconstruction_cosine(original, recon) >= DEFAULT_RECON_COSINE_BUDGET diff --git a/engine/tests/test_cc6_tokens.py b/engine/tests/test_cc6_tokens.py new file mode 100644 index 0000000..863865a --- /dev/null +++ b/engine/tests/test_cc6_tokens.py @@ -0,0 +1,64 @@ +"""CC-6: pluggable token counter; packing stays on cheap estimate.""" + +from __future__ import annotations + +import pytest + +from chat_compressor import metrics +from chat_compressor.tokens import ( + clear_counters, + count_tokens, + packing_tokens, + register_counter, + resolve_counter, + try_hf_counter, + unregister_counter, +) + + +@pytest.fixture(autouse=True) +def _clean_registry(): + clear_counters() + yield + clear_counters() + + +def test_packing_tokens_matches_estimate(): + text = "hello world " * 20 + assert packing_tokens(text) == metrics.estimate_tokens(text) + + +def test_count_tokens_falls_back_to_estimate(): + text = "abcdefghi" + assert count_tokens(text) == metrics.estimate_tokens(text) + assert count_tokens("") == 0 + + +def test_registered_counter_used_for_cost_not_packing(): + register_counter("unit-tok", lambda t: 42 if t else 0) + assert count_tokens("anything", tokenizer_id="unit-tok") == 42 + # Packing path ignores registry. + assert packing_tokens("anything") == metrics.estimate_tokens("anything") + assert resolve_counter("unit-tok")("x") == 42 + unregister_counter("unit-tok") + assert count_tokens("anything", tokenizer_id="unit-tok") == metrics.estimate_tokens( + "anything" + ) + + +def test_counter_exception_fails_open_to_estimate(): + def boom(_t: str) -> int: + raise RuntimeError("tok fail") + + register_counter("bad", boom) + text = "zzzz" + assert count_tokens(text, tokenizer_id="bad") == metrics.estimate_tokens(text) + + +def test_hf_counter_optional(): + """When transformers unavailable, try_hf_counter returns None (NOT_RUN).""" + counter = try_hf_counter("gpt2") + if counter is None: + pytest.skip("transformers/HF tokenizer not available (NOT_RUN)") + n = counter("Hello world") + assert isinstance(n, int) and n > 0 diff --git a/engine/tests/test_cc7_hop_legal.py b/engine/tests/test_cc7_hop_legal.py new file mode 100644 index 0000000..6270f4d --- /dev/null +++ b/engine/tests/test_cc7_hop_legal.py @@ -0,0 +1,74 @@ +"""CC-7: hop_legal() false with pending tool state.""" + +from __future__ import annotations + +from pathlib import Path + +import numpy as np + +from chat_compressor.handle import PersistentAgentHandle +from chat_compressor.producer import EmbeddingProducer +from chat_compressor.store import StateStore + + +def _handle(tmp_path: Path) -> PersistentAgentHandle: + store = StateStore(tmp_path / "state") + return PersistentAgentHandle( + agent_id="hop-legal", + store=store, + producer=EmbeddingProducer(d=32, k_max=4), + k_max=4, + ) + + +def test_hop_legal_true_at_clean_boundary(tmp_path: Path) -> None: + h = _handle(tmp_path) + assert h.hop_legal() is True + h.step("hello", recipient_id="model-a") + assert h.hop_legal() is True # stub tool_status, no pending + + +def test_hop_legal_false_when_pending_flag_set(tmp_path: Path) -> None: + h = _handle(tmp_path) + h.step("hello") + h.set_pending_tool(True) + assert h.hop_legal() is False + h.clear_pending_tool() + assert h.hop_legal() is True + + +def test_hop_legal_false_when_meta_pending_tool(tmp_path: Path) -> None: + store = StateStore(tmp_path / "state") + c = np.eye(2, 4, dtype=np.float32) + store.save( + agent_id="m", + C=c, + producer="embed", + meta={"tool_status": "stub", "pending_tool": True}, + k_max=4, + ) + h = PersistentAgentHandle( + agent_id="m", + store=store, + producer=EmbeddingProducer(d=4, k_max=4), + k_max=4, + ) + assert h.hop_legal() is False + + +def test_hop_legal_false_when_tool_status_pending(tmp_path: Path) -> None: + store = StateStore(tmp_path / "state") + store.save( + agent_id="m2", + C=np.ones((2, 4), dtype=np.float32), + producer="embed", + meta={"tool_status": "pending"}, + k_max=4, + ) + h = PersistentAgentHandle( + agent_id="m2", + store=store, + producer=EmbeddingProducer(d=4, k_max=4), + k_max=4, + ) + assert h.hop_legal() is False diff --git a/engine/tests/test_cc8_bundle.py b/engine/tests/test_cc8_bundle.py new file mode 100644 index 0000000..0ddd843 --- /dev/null +++ b/engine/tests/test_cc8_bundle.py @@ -0,0 +1,92 @@ +"""CC-8: export_bundle / import_bundle round-trip.""" + +from __future__ import annotations + +import json +from pathlib import Path + +import numpy as np + +from chat_compressor.bundle import export_bundle, import_bundle +from chat_compressor.handle import PersistentAgentHandle +from chat_compressor.producer import EmbeddingProducer +from chat_compressor.store import ( + StateStore, + append_inject_history, + load_inject_history, +) + + +def test_bundle_round_trip_full(tmp_path: Path) -> None: + store = StateStore(tmp_path / "state") + h = PersistentAgentHandle( + agent_id="agent-a", + store=store, + producer=EmbeddingProducer(d=32, k_max=4), + k_max=4, + ) + h.step("Create todo buy milk", recipient_id="model-a", recipient_version="v1") + h.step("Add bread to list", recipient_id="model-a", recipient_version="v1") + append_inject_history( + h._agent_dir(), + {"t": 1, "hashes": ["abc"], "packed_tokens": 10, "novel_tokens": 10}, + recipient_id="model-a", + ) + # Capture pre-export graph hot_set / typed for behavioral check. + hot_before = h.graph.hot_set() + typed_before = h.graph.typed_projection(None) + + dest = tmp_path / "bundle.v1" + export_bundle(store, "agent-a", dest) + assert (dest / "manifest.json").is_file() + assert (dest / "graph.json").is_file() + assert (dest / "lineage.json").is_file() + assert (dest / "inject_ledger.json").is_file() + assert list((dest / "states").glob("t*.safetensors")) + + store2 = StateStore(tmp_path / "state2") + result = import_bundle( + dest, + store2, + expected_producer="embed", + expected_d=32, + ) + assert result.mode == "full" + assert result.producer_matched is True + assert result.states_imported == 2 + assert result.graph_imported is True + assert result.ledger_imported is True + + chain = store2.lineage("agent-a") + assert len(chain) == 2 + assert chain[0].meta.get("recipient_id") == "model-a" + ledger = load_inject_history(Path(store2.root) / "agent-a", recipient_id="model-a") + assert ledger and ledger[0].get("hashes") == ["abc"] + + # Reload graph and check projections unchanged. + from chat_compressor.graph import CtxGraph + + g2 = CtxGraph.load(Path(store2.root) / "agent-a" / "graph.json") + assert g2.hot_set() == hot_before + assert g2.typed_projection(None) == typed_before + + +def test_bundle_producer_mismatch_graph_only(tmp_path: Path) -> None: + store = StateStore(tmp_path / "state") + c = np.eye(2, 8, dtype=np.float32) + store.save(agent_id="x", C=c, producer="embed", k_max=8) + (Path(store.root) / "x" / "graph.json").write_text( + json.dumps({"schema": "ctx-graph/v1", "nodes": [], "edges": []}) + "\n", + encoding="utf-8", + ) + dest = tmp_path / "b" + export_bundle(store, "x", dest) + store2 = StateStore(tmp_path / "other") + result = import_bundle( + dest, store2, expected_producer="other-producer", expected_d=8 + ) + assert result.mode == "graph_only" + assert result.producer_matched is False + assert result.states_imported == 0 + assert result.graph_imported is True + assert "mismatch" in result.notes.lower()