chopratejas commited on
Commit
e51d8c6
·
1 Parent(s): 628157f

fix: propagate HEADROOM_MODE through wrap CLI and drop cache write premium penalty

Browse files

Two fixes:

1. `headroom wrap claude` ignored HEADROOM_MODE env var — the Click-based
proxy CLI never passed `mode=` to ProxyConfig, so it always defaulted
to cost_savings. Added --mode flag to `headroom proxy` and forwarded
HEADROOM_MODE from wrap's _start_proxy().

2. Cache write premium (1.25x) was subtracted from savings as a penalty,
but Claude Code already pays this baseline cost regardless of Headroom.
Renamed bust_penalty_usd → write_premium_usd (observability only) and
stopped deducting it from net_savings_usd.

headroom/cli/proxy.py CHANGED
@@ -10,6 +10,12 @@ from .main import main
10
  @main.command()
11
  @click.option("--host", default="127.0.0.1", help="Host to bind to (default: 127.0.0.1)")
12
  @click.option("--port", "-p", default=8787, type=int, help="Port to bind to (default: 8787)")
 
 
 
 
 
 
13
  @click.option("--no-optimize", is_flag=True, help="Disable optimization (passthrough mode)")
14
  @click.option("--no-cache", is_flag=True, help="Disable semantic caching")
15
  @click.option("--no-rate-limit", is_flag=True, help="Disable rate limiting")
@@ -117,6 +123,7 @@ from .main import main
117
  @click.pass_context
118
  def proxy(
119
  ctx: click.Context,
 
120
  host: str,
121
  port: int,
122
  no_optimize: bool,
@@ -176,11 +183,15 @@ def proxy(
176
  # Resolve anyllm provider: env var takes precedence over CLI default (matches argparse path)
177
  effective_anyllm_provider = os.environ.get("HEADROOM_ANYLLM_PROVIDER") or anyllm_provider
178
 
 
 
 
179
  config = ProxyConfig(
180
  host=host,
181
  port=port,
182
  openai_api_url=effective_openai_api_url,
183
  gemini_api_url=effective_gemini_api_url,
 
184
  optimize=not no_optimize,
185
  cache_enabled=not no_cache,
186
  rate_limit_enabled=not no_rate_limit,
@@ -280,6 +291,7 @@ Starting proxy server...
280
 
281
  URL: http://{config.host}:{config.port}
282
  Backend: {backend_status}
 
283
  Optimization: {"ENABLED" if config.optimize else "DISABLED"}
284
  Caching: {"ENABLED" if config.cache_enabled else "DISABLED"}
285
  Rate Limit: {"ENABLED" if config.rate_limit_enabled else "DISABLED"}
 
10
  @main.command()
11
  @click.option("--host", default="127.0.0.1", help="Host to bind to (default: 127.0.0.1)")
12
  @click.option("--port", "-p", default=8787, type=int, help="Port to bind to (default: 8787)")
13
+ @click.option(
14
+ "--mode",
15
+ default=None,
16
+ type=click.Choice(["cost_savings", "token_headroom"]),
17
+ help="Optimization mode: cost_savings (preserve prefix cache) or token_headroom (compress for session extension). Default: cost_savings. Env: HEADROOM_MODE",
18
+ )
19
  @click.option("--no-optimize", is_flag=True, help="Disable optimization (passthrough mode)")
20
  @click.option("--no-cache", is_flag=True, help="Disable semantic caching")
21
  @click.option("--no-rate-limit", is_flag=True, help="Disable rate limiting")
 
123
  @click.pass_context
124
  def proxy(
125
  ctx: click.Context,
126
+ mode: str | None,
127
  host: str,
128
  port: int,
129
  no_optimize: bool,
 
183
  # Resolve anyllm provider: env var takes precedence over CLI default (matches argparse path)
184
  effective_anyllm_provider = os.environ.get("HEADROOM_ANYLLM_PROVIDER") or anyllm_provider
185
 
186
+ # Resolve mode: CLI flag > env var > default
187
+ effective_mode = mode or os.environ.get("HEADROOM_MODE", "cost_savings")
188
+
189
  config = ProxyConfig(
190
  host=host,
191
  port=port,
192
  openai_api_url=effective_openai_api_url,
193
  gemini_api_url=effective_gemini_api_url,
194
+ mode=effective_mode,
195
  optimize=not no_optimize,
196
  cache_enabled=not no_cache,
197
  rate_limit_enabled=not no_rate_limit,
 
291
 
292
  URL: http://{config.host}:{config.port}
293
  Backend: {backend_status}
294
+ Mode: {config.mode}
295
  Optimization: {"ENABLED" if config.optimize else "DISABLED"}
296
  Caching: {"ENABLED" if config.cache_enabled else "DISABLED"}
297
  Rate Limit: {"ENABLED" if config.rate_limit_enabled else "DISABLED"}
headroom/cli/wrap.py CHANGED
@@ -56,6 +56,11 @@ def _start_proxy(port: int) -> subprocess.Popen:
56
  """
57
  cmd = [sys.executable, "-m", "headroom.cli", "proxy", "--port", str(port)]
58
 
 
 
 
 
 
59
  log_path = _get_log_path()
60
  log_file = open(log_path, "a") # noqa: SIM115
61
 
 
56
  """
57
  cmd = [sys.executable, "-m", "headroom.cli", "proxy", "--port", str(port)]
58
 
59
+ # Forward HEADROOM_MODE env var so the proxy respects the user's mode choice
60
+ headroom_mode = os.environ.get("HEADROOM_MODE")
61
+ if headroom_mode:
62
+ cmd.extend(["--mode", headroom_mode])
63
+
64
  log_path = _get_log_path()
65
  log_file = open(log_path, "a") # noqa: SIM115
66
 
headroom/dashboard/templates/dashboard.html CHANGED
@@ -156,7 +156,7 @@
156
  <div>
157
  <div class="text-xs text-gray-500 uppercase tracking-wide mb-1">Cache Writes</div>
158
  <div class="text-2xl font-light tabular-nums text-amber-400" x-text="formatNumber(stats.prefix_cache?.totals?.cache_write_tokens || 0)"></div>
159
- <div class="text-xs text-amber-400/70" x-text="'$' + formatCurrency(stats.prefix_cache?.totals?.bust_penalty_usd || 0) + ' penalty'"></div>
160
  </div>
161
  <div>
162
  <div class="text-xs text-gray-500 uppercase tracking-wide mb-1">Hit Rate</div>
 
156
  <div>
157
  <div class="text-xs text-gray-500 uppercase tracking-wide mb-1">Cache Writes</div>
158
  <div class="text-2xl font-light tabular-nums text-amber-400" x-text="formatNumber(stats.prefix_cache?.totals?.cache_write_tokens || 0)"></div>
159
+ <div class="text-xs text-amber-400/70" x-text="'$' + formatCurrency(stats.prefix_cache?.totals?.write_premium_usd || 0) + ' write premium'"></div>
160
  </div>
161
  <div>
162
  <div class="text-xs text-gray-500 uppercase tracking-wide mb-1">Hit Rate</div>
headroom/proxy/server.py CHANGED
@@ -272,7 +272,7 @@ def _build_prefix_cache_stats(
272
  "bust_count": 0,
273
  "bust_write_tokens": 0,
274
  "savings_usd": 0.0,
275
- "bust_penalty_usd": 0.0,
276
  }
277
 
278
  for provider, pc in metrics.cache_by_provider.items():
@@ -302,19 +302,21 @@ def _build_prefix_cache_stats(
302
  break
303
 
304
  # Calculate savings:
305
- # Without cache: all tokens at 1.0x input price
306
- # With cache: read tokens at read_mult, write tokens at write_mult
 
 
307
  read_tokens: int = pc["cache_read_tokens"] # type: ignore[assignment]
308
  write_tokens: int = pc["cache_write_tokens"] # type: ignore[assignment]
309
  savings_usd = 0.0
310
- bust_penalty_usd = 0.0
311
 
312
  if input_price_per_token:
313
  # Savings from reads: tokens * price * (1.0 - read_multiplier)
314
  savings_usd = read_tokens * input_price_per_token * (1.0 - read_mult)
315
- # Penalty from writes: tokens * price * (write_multiplier - 1.0)
316
  if write_mult > 1.0:
317
- bust_penalty_usd = write_tokens * input_price_per_token * (write_mult - 1.0)
318
 
319
  hit_rate = round(pc["hit_requests"] / pc["requests"] * 100, 1) if pc["requests"] > 0 else 0
320
 
@@ -329,8 +331,8 @@ def _build_prefix_cache_stats(
329
  "read_discount": f"{(1.0 - read_mult) * 100:.0f}%",
330
  "write_premium": f"{(write_mult - 1.0) * 100:.0f}%" if write_mult > 1.0 else "none",
331
  "savings_usd": round(savings_usd, 4),
332
- "bust_penalty_usd": round(bust_penalty_usd, 4),
333
- "net_savings_usd": round(savings_usd - bust_penalty_usd, 4),
334
  "label": str(econ["label"]),
335
  }
336
  by_provider[provider] = provider_stats
@@ -343,11 +345,11 @@ def _build_prefix_cache_stats(
343
  totals["bust_count"] += pc["bust_count"]
344
  totals["bust_write_tokens"] += pc["bust_write_tokens"]
345
  totals["savings_usd"] += savings_usd
346
- totals["bust_penalty_usd"] += bust_penalty_usd
347
 
348
- totals["net_savings_usd"] = round(totals["savings_usd"] - totals["bust_penalty_usd"], 4)
349
  totals["savings_usd"] = round(totals["savings_usd"], 4)
350
- totals["bust_penalty_usd"] = round(totals["bust_penalty_usd"], 4)
351
  totals["hit_rate"] = (
352
  round(totals["hit_requests"] / totals["requests"] * 100, 1) if totals["requests"] > 0 else 0
353
  )
 
272
  "bust_count": 0,
273
  "bust_write_tokens": 0,
274
  "savings_usd": 0.0,
275
+ "write_premium_usd": 0.0,
276
  }
277
 
278
  for provider, pc in metrics.cache_by_provider.items():
 
302
  break
303
 
304
  # Calculate savings:
305
+ # Cache reads save (1.0 - read_mult) per token vs uncached input price.
306
+ # Cache write premium is NOT deducted — it's baseline cost that the
307
+ # client (e.g. Claude Code) pays regardless of Headroom. We track it
308
+ # for observability but don't penalise our savings number.
309
  read_tokens: int = pc["cache_read_tokens"] # type: ignore[assignment]
310
  write_tokens: int = pc["cache_write_tokens"] # type: ignore[assignment]
311
  savings_usd = 0.0
312
+ write_premium_usd = 0.0
313
 
314
  if input_price_per_token:
315
  # Savings from reads: tokens * price * (1.0 - read_multiplier)
316
  savings_usd = read_tokens * input_price_per_token * (1.0 - read_mult)
317
+ # Write premium (observability only — not subtracted from savings)
318
  if write_mult > 1.0:
319
+ write_premium_usd = write_tokens * input_price_per_token * (write_mult - 1.0)
320
 
321
  hit_rate = round(pc["hit_requests"] / pc["requests"] * 100, 1) if pc["requests"] > 0 else 0
322
 
 
331
  "read_discount": f"{(1.0 - read_mult) * 100:.0f}%",
332
  "write_premium": f"{(write_mult - 1.0) * 100:.0f}%" if write_mult > 1.0 else "none",
333
  "savings_usd": round(savings_usd, 4),
334
+ "write_premium_usd": round(write_premium_usd, 4),
335
+ "net_savings_usd": round(savings_usd, 4),
336
  "label": str(econ["label"]),
337
  }
338
  by_provider[provider] = provider_stats
 
345
  totals["bust_count"] += pc["bust_count"]
346
  totals["bust_write_tokens"] += pc["bust_write_tokens"]
347
  totals["savings_usd"] += savings_usd
348
+ totals["write_premium_usd"] += write_premium_usd
349
 
350
+ totals["net_savings_usd"] = round(totals["savings_usd"], 4)
351
  totals["savings_usd"] = round(totals["savings_usd"], 4)
352
+ totals["write_premium_usd"] = round(totals["write_premium_usd"], 4)
353
  totals["hit_rate"] = (
354
  round(totals["hit_requests"] / totals["requests"] * 100, 1) if totals["requests"] > 0 else 0
355
  )