diff --git a/benchmarks/claude_session_branch_compare.py b/benchmarks/claude_session_branch_compare.py new file mode 100644 index 000000000..f807158d4 --- /dev/null +++ b/benchmarks/claude_session_branch_compare.py @@ -0,0 +1,595 @@ +#!/usr/bin/env python3 +"""Compare Claude session mode simulations across two git refs.""" + +from __future__ import annotations + +import argparse +import json +import os +import shutil +import subprocess +import sys +import tempfile +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Any + +if __package__ in {None, ""}: + sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from benchmarks.claude_session_mode_benchmark import ( + IMPACT_DIRECTION, + OUTPUT_JSON, + PROXY_MODE_CACHE, + PROXY_MODE_TOKEN, + format_currency, +) + +DEFAULT_OUTPUT_DIR = Path("benchmark_results") / "branch_compare" + + +@dataclass +class BranchResult: + ref: str + label: str + commit: str + summary: str + dataset: dict[str, Any] + observed: dict[str, Any] + summaries: dict[str, dict[str, Any]] + winners: dict[str, str] + output_dir: str + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--left-ref", default="upstream/main") + parser.add_argument("--right-ref", default="HEAD") + parser.add_argument("--left-label", default="main") + parser.add_argument("--right-label", default="pr") + parser.add_argument("--root", type=Path, default=Path.home() / ".claude" / "projects") + parser.add_argument("--output-dir", type=Path, default=DEFAULT_OUTPUT_DIR) + parser.add_argument("--max-sessions", type=int, default=None) + parser.add_argument("--recent-turns-per-session", type=int, default=None) + parser.add_argument("--cache-ttl-minutes", type=int, default=5) + parser.add_argument("--cache-write-multiplier", type=float, default=1.25) + parser.add_argument("--workers", type=int, default=1) + parser.add_argument( + "--python", + default=sys.executable, + help="Python executable to use inside each worktree.", + ) + parser.add_argument( + "--keep-worktrees", + action="store_true", + help="Do not remove temporary worktrees after the comparison run.", + ) + return parser.parse_args() + + +def _run_git(args: list[str], cwd: Path) -> str: + completed = subprocess.run( + ["git", *args], + cwd=cwd, + check=True, + capture_output=True, + text=True, + ) + return completed.stdout.strip() + + +def _ref_slug(ref: str) -> str: + return "".join(ch if ch.isalnum() else "-" for ch in ref).strip("-").lower() or "ref" + + +def _branch_output_dir(base: Path, label: str) -> Path: + return base / _ref_slug(label) + + +def _comparison_paths(base: Path) -> tuple[Path, Path, Path]: + return ( + base / "claude_session_branch_compare.md", + base / "claude_session_branch_compare.json", + base / "claude_session_branch_compare.html", + ) + + +def _mode_metric(branch: BranchResult, mode: str, field: str) -> float: + summary = branch.summaries[mode] + if field == "no_cache_total_cost_usd": + if "no_cache_total_cost_usd" in summary: + value = summary["no_cache_total_cost_usd"] + else: + value = ( + float(summary["paid_input_cost_usd"]) + + (float(summary["cache_read_cost_usd"]) * 10.0) + + float(summary["paid_output_cost_usd"]) + ) + elif field == "prompt_window_with_cache": + value = float(summary["forwarded_input_tokens"]) + elif field == "prompt_window_without_cache_reads": + value = float(summary["forwarded_input_tokens"]) - float(summary["cache_read_tokens"]) + else: + value = summary[field] + if isinstance(value, bool): + return float(value) + return float(value) + + +def _delta(left: float, right: float) -> float: + return right - left + + +def _classify_delta(field: str, delta: float) -> str: + direction = IMPACT_DIRECTION.get(field, "same") + tolerance = 1e-9 + if abs(delta) <= tolerance: + return "no_change" + if direction == "lower": + return "assist" if delta < 0 else "harm" + if direction == "higher": + return "assist" if delta > 0 else "harm" + return "harm" + + +def _build_benchmark_command( + python_executable: str, + script_path: Path, + root: Path, + output_dir: Path, + max_sessions: int | None, + recent_turns_per_session: int | None, + cache_ttl_minutes: int, + cache_write_multiplier: float, + workers: int, +) -> list[str]: + command = [ + python_executable, + str(script_path), + "--root", + str(root), + "--output-dir", + str(output_dir), + "--cache-ttl-minutes", + str(cache_ttl_minutes), + "--cache-write-multiplier", + str(cache_write_multiplier), + "--workers", + str(workers), + ] + if max_sessions is not None: + command.extend(["--max-sessions", str(max_sessions)]) + if recent_turns_per_session is not None: + command.extend(["--recent-turns-per-session", str(recent_turns_per_session)]) + return command + + +def _load_branch_result( + repo_root: Path, + ref: str, + label: str, + branch_output_dir: Path, +) -> BranchResult: + payload = json.loads((branch_output_dir / OUTPUT_JSON).read_text(encoding="utf-8")) + commit = _run_git(["rev-parse", ref], repo_root) + summary = _run_git(["show", "-s", "--format=%s", ref], repo_root) + return BranchResult( + ref=ref, + label=label, + commit=commit, + summary=summary, + dataset=payload["dataset"], + observed=payload["observed"], + summaries=payload["summaries"], + winners=payload["winners"], + output_dir=str(branch_output_dir), + ) + + +def _run_branch_benchmark( + repo_root: Path, + ref: str, + label: str, + args: argparse.Namespace, + worktree_root: Path, +) -> BranchResult: + worktree_dir = worktree_root / _ref_slug(label) + branch_output_dir = _branch_output_dir(args.output_dir, label) + branch_output_dir.mkdir(parents=True, exist_ok=True) + if worktree_dir.exists(): + shutil.rmtree(worktree_dir) + _run_git(["worktree", "add", "--detach", str(worktree_dir), ref], repo_root) + try: + command = _build_benchmark_command( + python_executable=args.python, + script_path=repo_root / "benchmarks" / "claude_session_mode_benchmark.py", + root=args.root, + output_dir=branch_output_dir, + max_sessions=args.max_sessions, + recent_turns_per_session=args.recent_turns_per_session, + cache_ttl_minutes=args.cache_ttl_minutes, + cache_write_multiplier=args.cache_write_multiplier, + workers=args.workers, + ) + env = os.environ.copy() + env["PYTHONPATH"] = os.pathsep.join( + [str(worktree_dir), str(repo_root), env.get("PYTHONPATH", "")] + ).rstrip(os.pathsep) + subprocess.run(command, cwd=worktree_dir, check=True, env=env) + return _load_branch_result(repo_root, ref, label, branch_output_dir) + finally: + if not args.keep_worktrees: + subprocess.run( + ["git", "worktree", "remove", "--force", str(worktree_dir)], + cwd=repo_root, + check=True, + ) + + +def _winner_line(metric: str, left: BranchResult, right: BranchResult) -> str: + left_winner = left.winners[metric] + right_winner = right.winners[metric] + if left_winner == right_winner: + return f"- {metric}: both pick `{left_winner}`" + return ( + f"- {metric}: `{left.label}` picks `{left_winner}`, `{right.label}` picks `{right_winner}`" + ) + + +def _build_six_way_rows( + left: BranchResult, right: BranchResult +) -> list[dict[str, str | float | int]]: + rows: list[dict[str, str | float | int]] = [] + for branch in (left, right): + for mode in ("baseline", PROXY_MODE_TOKEN, PROXY_MODE_CACHE): + summary = branch.summaries[mode] + cost_delta = _mode_metric(branch, mode, "total_cost_usd") - _mode_metric( + branch, "baseline", "total_cost_usd" + ) + window_delta = int( + _mode_metric(branch, mode, "prompt_window_with_cache") + - _mode_metric(branch, "baseline", "prompt_window_with_cache") + ) + read_delta = int( + _mode_metric(branch, mode, "cache_read_tokens") + - _mode_metric(branch, "baseline", "cache_read_tokens") + ) + write_delta = int( + _mode_metric(branch, mode, "cache_write_tokens") + - _mode_metric(branch, "baseline", "cache_write_tokens") + ) + paid_input_delta = int( + _mode_metric(branch, mode, "regular_input_tokens") + - _mode_metric(branch, "baseline", "regular_input_tokens") + ) + rows.append( + { + "branch": branch.label, + "mode": mode, + "forwarded_input_tokens": int(summary["forwarded_input_tokens"]), + "cache_read_tokens": int(summary["cache_read_tokens"]), + "cache_write_tokens": int(summary["cache_write_tokens"]), + "regular_input_tokens": int(summary["regular_input_tokens"]), + "output_tokens": int(summary["output_tokens"]), + "total_cost_usd": float(summary["total_cost_usd"]), + "cost_delta_vs_branch_baseline": cost_delta, + "window_delta_vs_branch_baseline": window_delta, + "cache_read_delta_vs_branch_baseline": read_delta, + "cache_write_delta_vs_branch_baseline": write_delta, + "paid_input_delta_vs_branch_baseline": paid_input_delta, + "is_branch_winner": "yes" if branch.winners["total_cost"] == mode else "no", + } + ) + return rows + + +def build_compare_markdown(left: BranchResult, right: BranchResult) -> str: + six_way_rows = _build_six_way_rows(left, right) + lines = [ + "# Claude Session Branch Comparison", + "", + "## Branches", + "", + f"- {left.label}: `{left.ref}` @ `{left.commit[:12]}` - {left.summary}", + f"- {right.label}: `{right.ref}` @ `{right.commit[:12]}` - {right.summary}", + "", + "## Dataset", + "", + f"- Projects: {right.dataset['projects']}", + f"- Sessions: {right.dataset['sessions']}", + f"- Requests: {right.dataset['requests']}", + f"- Sampled requests: {right.dataset.get('sampled_requests', 0)}", + f"- Sampling: {right.dataset.get('sampling_note', 'Full sessions')}", + "", + "## Winner Comparison", + "", + _winner_line("total_cost", left, right), + _winner_line("no_cache_total_cost", left, right), + _winner_line("window_with_cache", left, right), + _winner_line("window_without_cache_reads", left, right), + "", + "## Six-Way Mode Matrix", + "", + "| Branch | Mode | Forwarded Input | Cache Read | Cache Write | Paid Input | Paid Output | Total Cost | Cost Δ vs Branch Baseline | Window Δ vs Branch Baseline | Winner |", + "| --- | --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | --- |", + *[ + "| " + + " | ".join( + [ + str(row["branch"]), + str(row["mode"]), + f"{int(row['forwarded_input_tokens']):,}", + f"{int(row['cache_read_tokens']):,}", + f"{int(row['cache_write_tokens']):,}", + f"{int(row['regular_input_tokens']):,}", + f"{int(row['output_tokens']):,}", + format_currency(float(row["total_cost_usd"])), + format_currency(float(row["cost_delta_vs_branch_baseline"])), + f"{int(row['window_delta_vs_branch_baseline']):,}", + str(row["is_branch_winner"]), + ] + ) + + " |" + for row in six_way_rows + ], + "", + "## Mode Deltas", + "", + f"| Mode | Metric | {left.label} | {right.label} | Delta ({right.label} - {left.label}) | Classification |", + "| --- | --- | ---: | ---: | ---: | --- |", + ] + metrics = [ + ("total_cost_usd", "Total Cost", format_currency), + ("no_cache_total_cost_usd", "No-Cache Total Cost", format_currency), + ("forwarded_input_tokens", "Forwarded Input Tokens", lambda v: f"{int(v):,}"), + ("cache_read_tokens", "Cache Read Tokens", lambda v: f"{int(v):,}"), + ("cache_write_tokens", "Cache Write Tokens", lambda v: f"{int(v):,}"), + ("cache_bust_turns", "Cache Bust Turns", lambda v: f"{int(v):,}"), + ("ttl_expiry_turns", "TTL Expiry Turns", lambda v: f"{int(v):,}"), + ("prompt_window_with_cache", "Window With Cache", lambda v: f"{int(v):,}"), + ( + "prompt_window_without_cache_reads", + "Window Without Cache Reads", + lambda v: f"{int(v):,}", + ), + ] + for mode in ("baseline", PROXY_MODE_TOKEN, PROXY_MODE_CACHE): + for field, label, formatter in metrics: + left_value = _mode_metric(left, mode, field) + right_value = _mode_metric(right, mode, field) + delta = _delta(left_value, right_value) + delta_text = format_currency(delta) if "cost" in field else f"{int(delta):,}" + classification = _classify_delta(field, delta) + lines.append( + f"| {mode} | {label} | {formatter(left_value)} | {formatter(right_value)} | {delta_text} | {classification} |" + ) + return "\n".join(lines) + + +def build_compare_html(left: BranchResult, right: BranchResult) -> str: + six_way_rows = [] + for row in _build_six_way_rows(left, right): + six_way_rows.append( + "" + f"{row['branch']}" + f"{row['mode']}" + f"{int(row['forwarded_input_tokens']):,}" + f"{int(row['cache_read_tokens']):,}" + f"{int(row['cache_write_tokens']):,}" + f"{int(row['regular_input_tokens']):,}" + f"{int(row['output_tokens']):,}" + f"{format_currency(float(row['total_cost_usd']))}" + f"{format_currency(float(row['cost_delta_vs_branch_baseline']))}" + f"{int(row['window_delta_vs_branch_baseline']):,}" + f"{row['is_branch_winner']}" + "" + ) + cards = [] + for branch in (left, right): + cards.append( + "
" + f"
{branch.label}
" + f"

{branch.ref}

" + f"

{branch.commit[:12]}

" + f"

{branch.summary}

" + "
" + f"
Total Cost{branch.winners['total_cost']}
" + f"
No Cache{branch.winners['no_cache_total_cost']}
" + f"
Window + Cache{branch.winners['window_with_cache']}
" + "
Window - Reads" + f"{branch.winners['window_without_cache_reads']}
" + "
" + "
" + ) + rows = [] + for mode in ("baseline", PROXY_MODE_TOKEN, PROXY_MODE_CACHE): + for field, label in ( + ("total_cost_usd", "Total Cost"), + ("no_cache_total_cost_usd", "No-Cache Total Cost"), + ("forwarded_input_tokens", "Forwarded Input Tokens"), + ("cache_read_tokens", "Cache Read Tokens"), + ("cache_write_tokens", "Cache Write Tokens"), + ("cache_bust_turns", "Cache Bust Turns"), + ("prompt_window_with_cache", "Window With Cache"), + ("prompt_window_without_cache_reads", "Window Without Cache Reads"), + ): + left_value = _mode_metric(left, mode, field) + right_value = _mode_metric(right, mode, field) + delta = _delta(left_value, right_value) + is_cost = "cost" in field + formatter = format_currency if is_cost else (lambda v: f"{int(v):,}") + delta_text = format_currency(delta) if is_cost else f"{int(delta):,}" + delta_class = "pos" if delta > 0 else "neg" if delta < 0 else "neutral" + classification = _classify_delta(field, delta) + rows.append( + "" + f"{mode}" + f"{label}" + f"{formatter(left_value)}" + f"{formatter(right_value)}" + f"{delta_text}" + f"{classification}" + "" + ) + return f""" + + + + + Claude Session Branch Comparison + + + +
+
+
Branch Comparison
+

Claude Session Mode Simulation

+

Same local Claude transcript corpus. Same simulation knobs. Two git refs. This report isolates code-level behavior changes between the branches.

+
+ {"".join(cards)} +
+
+
+
+ + + + + + + + + + + + + + + + + + {"".join(six_way_rows)} + +
BranchModeForwarded InputCache ReadCache WritePaid InputPaid OutputTotal CostCost Δ vs Branch BaselineWindow Δ vs Branch BaselineWinner
+
+
+
+
+ + + + + + + + + + + + + {"".join(rows)} + +
ModeMetric{left.label}{right.label}DeltaClassification
+
+
+
+ +""" + + +def write_compare_report( + output_dir: Path, + left: BranchResult, + right: BranchResult, +) -> tuple[Path, Path, Path]: + output_dir.mkdir(parents=True, exist_ok=True) + md_path, json_path, html_path = _comparison_paths(output_dir) + md_path.write_text(build_compare_markdown(left, right), encoding="utf-8") + html_path.write_text(build_compare_html(left, right), encoding="utf-8") + payload = { + "left": asdict(left), + "right": asdict(right), + "left_winners": left.winners, + "right_winners": right.winners, + } + json_path.write_text(json.dumps(payload, indent=2), encoding="utf-8") + return md_path, json_path, html_path + + +def main() -> int: + args = parse_args() + repo_root = Path(__file__).resolve().parents[1] + if not args.output_dir.is_absolute(): + args.output_dir = (repo_root / args.output_dir).resolve() + if not args.root.is_absolute(): + args.root = args.root.resolve() + args.output_dir.mkdir(parents=True, exist_ok=True) + worktree_root = Path(tempfile.mkdtemp(prefix="headroom-branch-compare-")) + try: + left = _run_branch_benchmark(repo_root, args.left_ref, args.left_label, args, worktree_root) + right = _run_branch_benchmark( + repo_root, args.right_ref, args.right_label, args, worktree_root + ) + md_path, json_path, html_path = write_compare_report(args.output_dir, left, right) + print(f"Compared {left.label} ({left.ref}) vs {right.label} ({right.ref})") + print(f"Markdown report: {md_path}") + print(f"JSON report: {json_path}") + print(f"HTML report: {html_path}") + return 0 + finally: + if args.keep_worktrees: + print(f"Retained worktrees under {worktree_root}") + else: + shutil.rmtree(worktree_root, ignore_errors=True) + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/benchmarks/claude_session_mode_benchmark.py b/benchmarks/claude_session_mode_benchmark.py index ccaa64a65..259dc53fd 100644 --- a/benchmarks/claude_session_mode_benchmark.py +++ b/benchmarks/claude_session_mode_benchmark.py @@ -20,11 +20,16 @@ from headroom.cache.prefix_tracker import PrefixCacheTracker from headroom.pricing.litellm_pricing import get_model_pricing from headroom.proxy.handlers.anthropic import AnthropicHandlerMixin from headroom.proxy.models import ProxyConfig -from headroom.proxy.modes import PROXY_MODE_CACHE, PROXY_MODE_TOKEN from headroom.proxy.server import HeadroomProxy from headroom.tokenizers import get_tokenizer from headroom.utils import extract_user_query +try: + from headroom.proxy.modes import PROXY_MODE_CACHE, PROXY_MODE_TOKEN +except ImportError: + PROXY_MODE_CACHE = "cache" + PROXY_MODE_TOKEN = "token" + DEFAULT_ROOT = Path.home() / ".claude" / "projects" DEFAULT_OUTPUT_DIR = Path("benchmark_results") DEFAULT_CACHE_TTL_MINUTES = 5 @@ -96,6 +101,9 @@ class ModeSummary: cache_eligible_turns: int = 0 cache_bust_turns: int = 0 ttl_expiry_turns: int = 0 + rewrite_turns: int = 0 + retroactive_rewrite_turns: int = 0 + latest_turn_only_rewrite_turns: int = 0 turns: list[TurnMetrics] = field(default_factory=list) @property @@ -136,6 +144,24 @@ class DatasetSummary: sampling_note: str = "" +IMPACT_DIRECTION = { + "forwarded_input_tokens": "lower", + "cache_read_tokens": "higher", + "cache_write_tokens": "lower", + "regular_input_tokens": "lower", + "output_tokens": "same", + "total_cost_usd": "lower", + "no_cache_total_cost_usd": "lower", + "prompt_window_with_cache": "lower", + "prompt_window_without_cache_reads": "lower", + "cache_bust_turns": "lower", + "ttl_expiry_turns": "lower", + "rewrite_turns": "lower", + "retroactive_rewrite_turns": "lower", + "latest_turn_only_rewrite_turns": "lower", +} + + @dataclass class ObservedSummary: sessions: int = 0 @@ -221,6 +247,9 @@ def _mode_summary_from_dict(data: dict[str, Any]) -> ModeSummary: cache_eligible_turns=data.get("cache_eligible_turns", 0), cache_bust_turns=data.get("cache_bust_turns", 0), ttl_expiry_turns=data.get("ttl_expiry_turns", 0), + rewrite_turns=data.get("rewrite_turns", 0), + retroactive_rewrite_turns=data.get("retroactive_rewrite_turns", 0), + latest_turn_only_rewrite_turns=data.get("latest_turn_only_rewrite_turns", 0), turns=turns, ) return summary @@ -618,6 +647,133 @@ def _common_prefix_tokens( return common +def _rewrite_scope( + original_messages: list[dict[str, Any]], + forwarded_messages: list[dict[str, Any]], + *, + stable_prefix_message_count: int, +) -> tuple[bool, bool]: + if original_messages == forwarded_messages: + return False, False + stable_count = min( + stable_prefix_message_count, + len(original_messages), + len(forwarded_messages), + ) + retroactive = False + if len(forwarded_messages) < stable_prefix_message_count: + retroactive = True + elif stable_count > 0 and forwarded_messages[:stable_count] != original_messages[:stable_count]: + retroactive = True + return True, retroactive + + +def _extract_cache_stable_delta( + current_messages: list[dict[str, Any]], + previous_original_messages: list[dict[str, Any]] | None, + previous_forwarded_messages: list[dict[str, Any]] | None, +) -> tuple[list[dict[str, Any]], list[dict[str, Any]]] | None: + if previous_original_messages is None or previous_forwarded_messages is None: + return None + if len(current_messages) < len(previous_original_messages): + return None + stable_count = len(previous_original_messages) + if current_messages[:stable_count] != previous_original_messages: + return None + return ( + copy.deepcopy(previous_forwarded_messages), + copy.deepcopy(current_messages[stable_count:]), + ) + + +def _extract_cache_stable_last_message_suffix( + current_messages: list[dict[str, Any]], + previous_original_messages: list[dict[str, Any]] | None, + previous_forwarded_messages: list[dict[str, Any]] | None, +) -> tuple[list[dict[str, Any]], dict[str, Any], list[dict[str, Any]]] | None: + if not previous_original_messages or previous_forwarded_messages is None: + return None + if ( + len(current_messages) != len(previous_original_messages) + or len(previous_forwarded_messages) != len(previous_original_messages) + or not current_messages + ): + return None + prefix_len = len(current_messages) - 1 + if prefix_len > 0 and current_messages[:prefix_len] != previous_original_messages[:prefix_len]: + return None + + current_last = current_messages[-1] + previous_original_last = previous_original_messages[-1] + previous_forwarded_last = previous_forwarded_messages[-1] + if current_last.get("role") != previous_original_last.get("role") or current_last.get( + "role" + ) != previous_forwarded_last.get("role"): + return None + + current_content = current_last.get("content") + previous_original_content = previous_original_last.get("content") + previous_forwarded_content = previous_forwarded_last.get("content") + + if ( + isinstance(current_content, str) + and isinstance(previous_original_content, str) + and isinstance(previous_forwarded_content, str) + and current_content.startswith(previous_original_content) + ): + suffix = current_content[len(previous_original_content) :] + delta_messages = [] + if suffix: + delta_messages = [{**copy.deepcopy(current_last), "content": suffix}] + return ( + copy.deepcopy(previous_forwarded_messages[:-1]), + copy.deepcopy(previous_forwarded_last), + delta_messages, + ) + + if ( + isinstance(current_content, list) + and isinstance(previous_original_content, list) + and isinstance(previous_forwarded_content, list) + and len(current_content) >= len(previous_original_content) + and current_content[: len(previous_original_content)] == previous_original_content + ): + delta_blocks = copy.deepcopy(current_content[len(previous_original_content) :]) + delta_messages = [] + if delta_blocks: + delta_messages = [{**copy.deepcopy(current_last), "content": delta_blocks}] + return ( + copy.deepcopy(previous_forwarded_messages[:-1]), + copy.deepcopy(previous_forwarded_last), + delta_messages, + ) + return None + + +def _merge_appended_message_delta( + previous_forwarded_message: dict[str, Any], + delta_forwarded_message: dict[str, Any] | None, +) -> dict[str, Any] | None: + if delta_forwarded_message is None: + return copy.deepcopy(previous_forwarded_message) + if previous_forwarded_message.get("role") != delta_forwarded_message.get("role"): + return None + + previous_content = previous_forwarded_message.get("content") + delta_content = delta_forwarded_message.get("content") + if isinstance(previous_content, str) and isinstance(delta_content, str): + return { + **copy.deepcopy(previous_forwarded_message), + "content": previous_content + delta_content, + } + if isinstance(previous_content, list) and isinstance(delta_content, list): + return { + **copy.deepcopy(previous_forwarded_message), + "content": copy.deepcopy(previous_content) + copy.deepcopy(delta_content), + } + return None + + def _make_proxy(mode: str) -> HeadroomProxy: cfg = ProxyConfig( mode=mode, @@ -655,25 +811,47 @@ def _apply_mode_to_messages( assert proxy is not None assert prefix_tracker is not None if mode == PROXY_MODE_CACHE: - delta = AnthropicHandlerMixin._extract_cache_stable_delta( + supports_delta_replay = hasattr( + AnthropicHandlerMixin, "_extract_cache_stable_last_message_suffix" + ) + if not supports_delta_replay: + frozen_message_count = prefix_tracker.get_frozen_message_count() + context_limit = proxy.anthropic_provider.get_context_limit(model) + result = proxy.anthropic_pipeline.apply( + messages=copy.deepcopy(messages), + model=model, + model_limit=context_limit, + context=extract_user_query(messages), + frozen_message_count=frozen_message_count, + ) + if hasattr(AnthropicHandlerMixin, "_restore_frozen_prefix"): + result.messages, _ = AnthropicHandlerMixin._restore_frozen_prefix( + messages, + result.messages, + frozen_message_count=frozen_message_count, + ) + return result.messages + + delta = _extract_cache_stable_delta( messages, previous_original_messages, previous_forwarded_messages, ) - if delta is None: - return copy.deepcopy(messages) - stable_forwarded_prefix, delta_messages = delta - if not delta_messages: - return stable_forwarded_prefix - context_limit = proxy.anthropic_provider.get_context_limit(model) - result = proxy.anthropic_pipeline.apply( - messages=delta_messages, - model=model, - model_limit=context_limit, - context=extract_user_query(delta_messages), - frozen_message_count=0, - ) - return stable_forwarded_prefix + result.messages + if delta is not None: + stable_forwarded_prefix, delta_messages = delta + if not delta_messages: + return stable_forwarded_prefix + context_limit = proxy.anthropic_provider.get_context_limit(model) + result = proxy.anthropic_pipeline.apply( + messages=delta_messages, + model=model, + model_limit=context_limit, + context=extract_user_query(delta_messages), + frozen_message_count=0, + ) + return stable_forwarded_prefix + result.messages + + return copy.deepcopy(messages) frozen_message_count = prefix_tracker.get_frozen_message_count() @@ -842,6 +1020,9 @@ def _merge_mode_summary(target: ModeSummary, source: ModeSummary) -> None: target.cache_eligible_turns += source.cache_eligible_turns target.cache_bust_turns += source.cache_bust_turns target.ttl_expiry_turns += source.ttl_expiry_turns + target.rewrite_turns += source.rewrite_turns + target.retroactive_rewrite_turns += source.retroactive_rewrite_turns + target.latest_turn_only_rewrite_turns += source.latest_turn_only_rewrite_turns def _disable_headroom_benchmark_logging() -> None: @@ -914,6 +1095,32 @@ def _write_checkpoint_by_session_id( path.write_text(json.dumps(payload, indent=2), encoding="utf-8") +def _update_prefix_tracker( + prefix_tracker: PrefixCacheTracker, + *, + cache_read_tokens: int, + cache_write_tokens: int, + messages: list[dict[str, Any]], + message_token_counts: list[int], + original_messages: list[dict[str, Any]] | None = None, +) -> None: + try: + prefix_tracker.update_from_response( + cache_read_tokens=cache_read_tokens, + cache_write_tokens=cache_write_tokens, + messages=messages, + message_token_counts=message_token_counts, + original_messages=original_messages, + ) + except TypeError: + prefix_tracker.update_from_response( + cache_read_tokens=cache_read_tokens, + cache_write_tokens=cache_write_tokens, + messages=messages, + message_token_counts=message_token_counts, + ) + + def _simulate_single_replay_mode( replay: SessionReplay, mode: str, @@ -938,6 +1145,7 @@ def _simulate_single_replay_mode( for turn in replay.turns: tokenizer = get_tokenizer(turn.model) turn_input_token_total = sum(tokenizer.count_message(msg) for msg in turn.input_messages) + prior_context_message_count = len(conversation) conversation.extend(turn.input_messages) raw_input_tokens = conversation_token_total + turn_input_token_total forwarded = _apply_mode_to_messages( @@ -950,6 +1158,17 @@ def _simulate_single_replay_mode( previous_original_messages=previous_original_context, previous_forwarded_messages=previous_forwarded_context, ) + rewrite, retroactive_rewrite = _rewrite_scope( + conversation, + forwarded, + stable_prefix_message_count=prior_context_message_count, + ) + if rewrite: + summary.rewrite_turns += 1 + if retroactive_rewrite: + summary.retroactive_rewrite_turns += 1 + else: + summary.latest_turn_only_rewrite_turns += 1 if pending is not None: _apply_turn_metrics( pending.summary, @@ -968,7 +1187,8 @@ def _simulate_single_replay_mode( previous_timestamp = pending.turn.timestamp if prefix_tracker is not None: - prefix_tracker.update_from_response( + _update_prefix_tracker( + prefix_tracker, cache_read_tokens=0, cache_write_tokens=0, messages=forwarded, @@ -1231,12 +1451,60 @@ def determine_winners(summaries: dict[str, ModeSummary]) -> dict[str, str]: } +def _metric_value(summary: ModeSummary, field: str) -> float: + value = getattr(summary, field) + return float(value) + + +def classify_metric_impact( + baseline: ModeSummary, + candidate: ModeSummary, + field: str, +) -> dict[str, float | str]: + baseline_value = _metric_value(baseline, field) + candidate_value = _metric_value(candidate, field) + delta = candidate_value - baseline_value + direction = IMPACT_DIRECTION[field] + tolerance = 1e-9 + + if abs(delta) <= tolerance: + impact = "no_change" + elif direction == "lower": + impact = "assist" if delta < 0 else "harm" + elif direction == "higher": + impact = "assist" if delta > 0 else "harm" + else: + impact = "harm" if abs(delta) > tolerance else "no_change" + + return { + "baseline": baseline_value, + "candidate": candidate_value, + "delta": delta, + "impact": impact, + "direction": direction, + } + + +def summarize_mode_impact_vs_baseline( + summaries: dict[str, ModeSummary], +) -> dict[str, dict[str, dict[str, float | str]]]: + baseline = summaries["baseline"] + result: dict[str, dict[str, dict[str, float | str]]] = {} + for mode in (PROXY_MODE_TOKEN, PROXY_MODE_CACHE): + candidate = summaries[mode] + result[mode] = { + field: classify_metric_impact(baseline, candidate, field) for field in IMPACT_DIRECTION + } + return result + + def format_currency(value: float) -> str: return f"${value:,.2f}" def print_console_report(dataset: DatasetSummary, summaries: dict[str, ModeSummary]) -> None: winners = determine_winners(summaries) + impacts = summarize_mode_impact_vs_baseline(summaries) print("Claude session mode simulation") print( f"Dataset: {dataset.projects} projects, {dataset.sessions} sessions, " @@ -1245,7 +1513,7 @@ def print_console_report(dataset: DatasetSummary, summaries: dict[str, ModeSumma print(f"Sampling: {dataset.sampling_note}") print() print( - "mode raw_tok cache_tok cache_read cache_write paid_in paid_out busts ttl_exp total_cost no_cache" + "mode raw_tok cache_tok cache_read cache_write paid_in paid_out busts ttl_exp rewrite retro_rw total_cost no_cache" ) for mode in ("baseline", PROXY_MODE_TOKEN, PROXY_MODE_CACHE): summary = summaries[mode] @@ -1254,6 +1522,7 @@ def print_console_report(dataset: DatasetSummary, summaries: dict[str, ModeSumma f"{summary.cache_read_tokens:>11,} {summary.cache_write_tokens:>12,} " f"{summary.regular_input_tokens:>10,} {summary.output_tokens:>12,} " f"{summary.cache_bust_turns:>7,} {summary.ttl_expiry_turns:>9,} " + f"{summary.rewrite_turns:>9,} {summary.retroactive_rewrite_turns:>10,} " f"{format_currency(summary.total_cost_usd):>11} " f"{format_currency(summary.no_cache_total_cost_usd):>11}" ) @@ -1265,6 +1534,26 @@ def print_console_report(dataset: DatasetSummary, summaries: dict[str, ModeSumma "Winner if cache read tokens do not count against window: " f"{winners['window_without_cache_reads']}" ) + print() + print("Impact vs baseline") + for mode in (PROXY_MODE_TOKEN, PROXY_MODE_CACHE): + impact = impacts[mode] + print( + f"{mode}: total_cost={impact['total_cost_usd']['impact']} " + f"({format_currency(impact['total_cost_usd']['delta'])}), " + f"cache_read={impact['cache_read_tokens']['impact']} " + f"({int(impact['cache_read_tokens']['delta']):,}), " + f"cache_write={impact['cache_write_tokens']['impact']} " + f"({int(impact['cache_write_tokens']['delta']):,}), " + f"paid_input={impact['regular_input_tokens']['impact']} " + f"({int(impact['regular_input_tokens']['delta']):,}), " + f"rewrite={impact['rewrite_turns']['impact']} " + f"({int(impact['rewrite_turns']['delta']):,}), " + f"retro_rw={impact['retroactive_rewrite_turns']['impact']} " + f"({int(impact['retroactive_rewrite_turns']['delta']):,}), " + f"window={impact['prompt_window_with_cache']['impact']} " + f"({int(impact['prompt_window_with_cache']['delta']):,})" + ) def print_observed_console_report(observed: ObservedSummary) -> None: @@ -1288,6 +1577,7 @@ def build_report_markdown( summaries: dict[str, ModeSummary], ) -> str: winners = determine_winners(summaries) + impacts = summarize_mode_impact_vs_baseline(summaries) model_lines = "\n".join(f"- `{model}`: {count}" for model, count in dataset.models.items()) rows = [] for mode in ("baseline", PROXY_MODE_TOKEN, PROXY_MODE_CACHE): @@ -1311,12 +1601,36 @@ def build_report_markdown( format_currency(summary.no_cache_total_cost_usd), f"{summary.cache_bust_turns:,}", f"{summary.ttl_expiry_turns:,}", + f"{summary.rewrite_turns:,}", + f"{summary.retroactive_rewrite_turns:,}", + f"{summary.latest_turn_only_rewrite_turns:,}", f"{summary.prompt_window_with_cache:,}", f"{summary.prompt_window_without_cache_reads:,}", ] ) + " |" ) + impact_rows = [] + for mode in (PROXY_MODE_TOKEN, PROXY_MODE_CACHE): + for metric_key, label in ( + ("total_cost_usd", "Total Cost"), + ("cache_read_tokens", "Cache Read Tokens"), + ("cache_write_tokens", "Cache Write Tokens"), + ("regular_input_tokens", "Paid Input Tokens"), + ("output_tokens", "Paid Output Tokens"), + ("prompt_window_with_cache", "Window With Cache"), + ("prompt_window_without_cache_reads", "Window Without Cache Reads"), + ("cache_bust_turns", "Cache Bust Turns"), + ("rewrite_turns", "Rewrite Turns"), + ("retroactive_rewrite_turns", "Retroactive Rewrite Turns"), + ("latest_turn_only_rewrite_turns", "Latest-Turn-Only Rewrite Turns"), + ): + impact = impacts[mode][metric_key] + delta = impact["delta"] + delta_text = format_currency(delta) if "cost" in metric_key else f"{int(delta):,}" + impact_rows.append( + f"| {mode} | {label} | {impact['impact']} | {delta_text} | {impact['direction']} |" + ) return "\n".join( [ "# Claude Session Mode Simulation", @@ -1351,10 +1665,16 @@ def build_report_markdown( "", "## Summary", "", - "| Mode | Raw Tokens | Cache Tokens | Cache Read | Cache Write | Paid Input Tokens | Paid Output Tokens | Paid Input Cost | Cache Read Cost | Cache Write Cost | Paid Output Cost | Total Cost | No-Cache Total Cost | Cache Bust Turns | TTL Expiry Turns | Window Tokens (Cache Counted) | Window Tokens (Cache Reads Excluded) |", - "| --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: |", + "| Mode | Raw Tokens | Cache Tokens | Cache Read | Cache Write | Paid Input Tokens | Paid Output Tokens | Paid Input Cost | Cache Read Cost | Cache Write Cost | Paid Output Cost | Total Cost | No-Cache Total Cost | Cache Bust Turns | TTL Expiry Turns | Rewrite Turns | Retroactive Rewrite Turns | Latest-Turn-Only Rewrite Turns | Window Tokens (Cache Counted) | Window Tokens (Cache Reads Excluded) |", + "| --- | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: | ---: |", *rows, "", + "## Impact vs Baseline", + "", + "| Mode | Metric | Classification | Delta | Better Direction |", + "| --- | --- | --- | ---: | --- |", + *impact_rows, + "", "## Winners", "", f"- Total cost winner: `{winners['total_cost']}`", @@ -1372,6 +1692,7 @@ def build_report_html( summaries: dict[str, ModeSummary], ) -> str: winners = determine_winners(summaries) + impacts = summarize_mode_impact_vs_baseline(summaries) model_items = "".join( f"
  • {model}{count:,}
  • " for model, count in dataset.models.items() @@ -1390,12 +1711,42 @@ def build_report_html( f"{summary.output_tokens:,}" f"{summary.cache_bust_turns:,}" f"{summary.ttl_expiry_turns:,}" + f"{summary.rewrite_turns:,}" + f"{summary.retroactive_rewrite_turns:,}" + f"{summary.latest_turn_only_rewrite_turns:,}" f"{format_currency(summary.total_cost_usd)}" f"{format_currency(summary.no_cache_total_cost_usd)}" f"{summary.prompt_window_with_cache:,}" f"{summary.prompt_window_without_cache_reads:,}" "" ) + impact_rows = [] + for mode in (PROXY_MODE_TOKEN, PROXY_MODE_CACHE): + for metric_key, label in ( + ("total_cost_usd", "Total Cost"), + ("cache_read_tokens", "Cache Read Tokens"), + ("cache_write_tokens", "Cache Write Tokens"), + ("regular_input_tokens", "Paid Input Tokens"), + ("output_tokens", "Paid Output Tokens"), + ("prompt_window_with_cache", "Window With Cache"), + ("prompt_window_without_cache_reads", "Window Without Cache Reads"), + ("cache_bust_turns", "Cache Bust Turns"), + ("rewrite_turns", "Rewrite Turns"), + ("retroactive_rewrite_turns", "Retroactive Rewrite Turns"), + ("latest_turn_only_rewrite_turns", "Latest-Turn-Only Rewrite Turns"), + ): + impact = impacts[mode][metric_key] + delta = impact["delta"] + delta_text = format_currency(delta) if "cost" in metric_key else f"{int(delta):,}" + impact_rows.append( + "" + f"{mode}" + f"{label}" + f"{impact['impact']}" + f"{delta_text}" + f"{impact['direction']}" + "" + ) return f""" @@ -1517,7 +1868,7 @@ def build_report_html( - + @@ -1526,6 +1877,21 @@ def build_report_html(
    ModeRaw TokensCache TokensCache ReadCache WritePaid InputPaid OutputCache BustsTTL ExpiryTotal CostNo-Cache CostWindow With CacheWindow Without Cache ReadsModeRaw TokensCache TokensCache ReadCache WritePaid InputPaid OutputCache BustsTTL ExpiryRewrite TurnsRetroactive RewritesLatest-Turn-Only RewritesTotal CostNo-Cache CostWindow With CacheWindow Without Cache Reads
    +
    +

    Impact vs Baseline

    +
    + + + + + + + + {"".join(impact_rows)} + +
    ModeMetricClassificationDeltaBetter Direction
    +
    +
    """ @@ -1548,6 +1914,7 @@ def write_report( "observed": asdict(observed), "summaries": {mode: asdict(summary) for mode, summary in summaries.items()}, "winners": determine_winners(summaries), + "impact_vs_baseline": summarize_mode_impact_vs_baseline(summaries), } json_path.write_text(json.dumps(payload, indent=2), encoding="utf-8") return md_path, json_path, html_path diff --git a/docs/benchmarks.md b/docs/benchmarks.md index 7d78ecf36..2530c6bc5 100644 --- a/docs/benchmarks.md +++ b/docs/benchmarks.md @@ -209,6 +209,9 @@ python benchmarks/proxy_mode_benchmark.py --turns 12 --show-real-harness # Replay local Claude Code transcripts (no API calls) python benchmarks/claude_session_mode_benchmark.py --workers 1 + +# Compare two refs on the same local Claude transcript corpus +python benchmarks/claude_session_branch_compare.py --left-ref upstream/main --right-ref HEAD --recent-turns-per-session 200 --workers 1 ``` This benchmark compares `token` vs `cache` proxy modes on the same synthetic conversation: @@ -218,6 +221,13 @@ This benchmark compares `token` vs `cache` proxy modes on the same synthetic con `--show-real-harness` prints optional steps for running the same comparison with Claude Code, but does not call APIs by default. +`claude_session_branch_compare.py` runs the real local session replay benchmark twice, once per git ref, in isolated worktrees. It writes: + +- per-ref replay outputs under `benchmark_results/branch_compare/