Garm commited on
Commit
6de75aa
·
1 Parent(s): d30a2e7

Format display-session tracker changes

Browse files
headroom/proxy/savings_tracker.py CHANGED
@@ -49,7 +49,12 @@ def _utc_now() -> datetime:
49
 
50
 
51
  def _to_utc_iso(dt: datetime) -> str:
52
- return dt.astimezone(timezone.utc).replace(microsecond=0).isoformat().replace("+00:00", "Z")
 
 
 
 
 
53
 
54
 
55
  def _parse_timestamp(value: Any) -> datetime | None:
@@ -117,7 +122,11 @@ def _resolve_litellm_model(model: str) -> str:
117
  if model.startswith(pattern):
118
  candidate = f"{prefix}{model}"
119
  try:
120
- litellm.cost_per_token(model=candidate, prompt_tokens=1, completion_tokens=0)
 
 
 
 
121
  return candidate
122
  except Exception:
123
  break
@@ -170,8 +179,14 @@ def _estimate_input_cost_usd(
170
  return 0.0
171
 
172
  if cache_read + cache_write + uncached > 0:
173
- cache_read_cost = info.get("cache_read_input_token_cost", input_cost_per_token)
174
- cache_write_cost = info.get("cache_creation_input_token_cost", input_cost_per_token)
 
 
 
 
 
 
175
  return (
176
  float(cache_read) * float(cache_read_cost)
177
  + float(cache_write) * float(cache_write_cost)
@@ -247,7 +262,10 @@ def _normalize_display_session(entry: Any) -> dict[str, Any]:
247
  tokens_saved = _coerce_int(entry.get("tokens_saved"))
248
  total_input_tokens = _coerce_int(entry.get("total_input_tokens"))
249
  total_before = tokens_saved + total_input_tokens
250
- savings_percent = round((tokens_saved / total_before * 100) if total_before > 0 else 0.0, 2)
 
 
 
251
 
252
  return {
253
  "requests": _coerce_int(entry.get("requests")),
@@ -257,7 +275,10 @@ def _normalize_display_session(entry: Any) -> dict[str, Any]:
257
  6,
258
  ),
259
  "total_input_tokens": total_input_tokens,
260
- "total_input_cost_usd": round(_coerce_float(entry.get("total_input_cost_usd")), 6),
 
 
 
261
  "savings_percent": savings_percent,
262
  "started_at": _to_utc_iso(started_at),
263
  "last_activity_at": _to_utc_iso(last_activity_at),
@@ -272,13 +293,18 @@ class SavingsTracker:
272
  path: str | None = None,
273
  max_history_points: int = DEFAULT_MAX_HISTORY_POINTS,
274
  max_history_age_days: int = DEFAULT_MAX_HISTORY_AGE_DAYS,
275
- display_session_inactivity_minutes: int = DEFAULT_DISPLAY_SESSION_INACTIVITY_MINUTES,
 
 
276
  ) -> None:
277
  self._path = Path(path or get_default_savings_storage_path())
278
  self._max_history_points = max_history_points
279
  self._max_history_age_days = max_history_age_days
280
  self._display_session_inactivity_minutes = max(
281
- _coerce_int(display_session_inactivity_minutes, DEFAULT_DISPLAY_SESSION_INACTIVITY_MINUTES),
 
 
 
282
  1,
283
  )
284
  self._lock = threading.Lock()
@@ -446,7 +472,9 @@ class SavingsTracker:
446
  )
447
  total_before = session["tokens_saved"] + session["total_input_tokens"]
448
  session["savings_percent"] = round(
449
- (session["tokens_saved"] / total_before * 100) if total_before > 0 else 0.0,
 
 
450
  2,
451
  )
452
  session["last_activity_at"] = _to_utc_iso(timestamp_dt)
@@ -556,7 +584,9 @@ class SavingsTracker:
556
  "lifetime": dict(self._state["lifetime"]),
557
  "display_session": self._display_session_snapshot_locked(),
558
  "display_session_policy": {
559
- "rollover_inactivity_minutes": self._display_session_inactivity_minutes,
 
 
560
  },
561
  "history": history,
562
  "retention": {
@@ -615,13 +645,20 @@ class SavingsTracker:
615
  if isinstance(lifetime_raw, dict):
616
  lifetime_requests = _coerce_int(lifetime_raw.get("requests"))
617
  lifetime_tokens_saved = _coerce_int(lifetime_raw.get("tokens_saved"))
618
- lifetime_savings_usd = _coerce_float(lifetime_raw.get("compression_savings_usd"))
 
 
619
  lifetime_input_tokens = _coerce_int(lifetime_raw.get("total_input_tokens"))
620
- lifetime_input_cost_usd = _coerce_float(lifetime_raw.get("total_input_cost_usd"))
 
 
621
 
622
  if normalized_history:
623
  last = normalized_history[-1]
624
- lifetime_tokens_saved = max(lifetime_tokens_saved, last["total_tokens_saved"])
 
 
 
625
  lifetime_savings_usd = max(
626
  lifetime_savings_usd,
627
  _coerce_float(last["compression_savings_usd"]),
@@ -649,7 +686,9 @@ class SavingsTracker:
649
  }
650
 
651
  if normalized_history:
652
- reference_time = _parse_timestamp(normalized_history[-1]["timestamp"]) or _utc_now()
 
 
653
  original_state = self._state if hasattr(self, "_state") else None
654
  self._state = state
655
  try:
@@ -667,7 +706,9 @@ class SavingsTracker:
667
  return
668
 
669
  if self._max_history_age_days > 0:
670
- cutoff = (reference_time or _utc_now()) - timedelta(days=self._max_history_age_days)
 
 
671
  filtered = [
672
  item
673
  for item in history
@@ -754,7 +795,11 @@ class SavingsTracker:
754
  minutes=self._display_session_inactivity_minutes
755
  )
756
 
757
- def _build_rollup(self, history: list[dict[str, Any]], bucket: str) -> list[dict[str, Any]]:
 
 
 
 
758
  if not history:
759
  return []
760
 
 
49
 
50
 
51
  def _to_utc_iso(dt: datetime) -> str:
52
+ return (
53
+ dt.astimezone(timezone.utc)
54
+ .replace(microsecond=0)
55
+ .isoformat()
56
+ .replace("+00:00", "Z")
57
+ )
58
 
59
 
60
  def _parse_timestamp(value: Any) -> datetime | None:
 
122
  if model.startswith(pattern):
123
  candidate = f"{prefix}{model}"
124
  try:
125
+ litellm.cost_per_token(
126
+ model=candidate,
127
+ prompt_tokens=1,
128
+ completion_tokens=0,
129
+ )
130
  return candidate
131
  except Exception:
132
  break
 
179
  return 0.0
180
 
181
  if cache_read + cache_write + uncached > 0:
182
+ cache_read_cost = info.get(
183
+ "cache_read_input_token_cost",
184
+ input_cost_per_token,
185
+ )
186
+ cache_write_cost = info.get(
187
+ "cache_creation_input_token_cost",
188
+ input_cost_per_token,
189
+ )
190
  return (
191
  float(cache_read) * float(cache_read_cost)
192
  + float(cache_write) * float(cache_write_cost)
 
262
  tokens_saved = _coerce_int(entry.get("tokens_saved"))
263
  total_input_tokens = _coerce_int(entry.get("total_input_tokens"))
264
  total_before = tokens_saved + total_input_tokens
265
+ savings_percent = round(
266
+ (tokens_saved / total_before * 100) if total_before > 0 else 0.0,
267
+ 2,
268
+ )
269
 
270
  return {
271
  "requests": _coerce_int(entry.get("requests")),
 
275
  6,
276
  ),
277
  "total_input_tokens": total_input_tokens,
278
+ "total_input_cost_usd": round(
279
+ _coerce_float(entry.get("total_input_cost_usd")),
280
+ 6,
281
+ ),
282
  "savings_percent": savings_percent,
283
  "started_at": _to_utc_iso(started_at),
284
  "last_activity_at": _to_utc_iso(last_activity_at),
 
293
  path: str | None = None,
294
  max_history_points: int = DEFAULT_MAX_HISTORY_POINTS,
295
  max_history_age_days: int = DEFAULT_MAX_HISTORY_AGE_DAYS,
296
+ display_session_inactivity_minutes: int = (
297
+ DEFAULT_DISPLAY_SESSION_INACTIVITY_MINUTES
298
+ ),
299
  ) -> None:
300
  self._path = Path(path or get_default_savings_storage_path())
301
  self._max_history_points = max_history_points
302
  self._max_history_age_days = max_history_age_days
303
  self._display_session_inactivity_minutes = max(
304
+ _coerce_int(
305
+ display_session_inactivity_minutes,
306
+ DEFAULT_DISPLAY_SESSION_INACTIVITY_MINUTES,
307
+ ),
308
  1,
309
  )
310
  self._lock = threading.Lock()
 
472
  )
473
  total_before = session["tokens_saved"] + session["total_input_tokens"]
474
  session["savings_percent"] = round(
475
+ (session["tokens_saved"] / total_before * 100)
476
+ if total_before > 0
477
+ else 0.0,
478
  2,
479
  )
480
  session["last_activity_at"] = _to_utc_iso(timestamp_dt)
 
584
  "lifetime": dict(self._state["lifetime"]),
585
  "display_session": self._display_session_snapshot_locked(),
586
  "display_session_policy": {
587
+ "rollover_inactivity_minutes": (
588
+ self._display_session_inactivity_minutes
589
+ ),
590
  },
591
  "history": history,
592
  "retention": {
 
645
  if isinstance(lifetime_raw, dict):
646
  lifetime_requests = _coerce_int(lifetime_raw.get("requests"))
647
  lifetime_tokens_saved = _coerce_int(lifetime_raw.get("tokens_saved"))
648
+ lifetime_savings_usd = _coerce_float(
649
+ lifetime_raw.get("compression_savings_usd")
650
+ )
651
  lifetime_input_tokens = _coerce_int(lifetime_raw.get("total_input_tokens"))
652
+ lifetime_input_cost_usd = _coerce_float(
653
+ lifetime_raw.get("total_input_cost_usd")
654
+ )
655
 
656
  if normalized_history:
657
  last = normalized_history[-1]
658
+ lifetime_tokens_saved = max(
659
+ lifetime_tokens_saved,
660
+ last["total_tokens_saved"],
661
+ )
662
  lifetime_savings_usd = max(
663
  lifetime_savings_usd,
664
  _coerce_float(last["compression_savings_usd"]),
 
686
  }
687
 
688
  if normalized_history:
689
+ reference_time = (
690
+ _parse_timestamp(normalized_history[-1]["timestamp"]) or _utc_now()
691
+ )
692
  original_state = self._state if hasattr(self, "_state") else None
693
  self._state = state
694
  try:
 
706
  return
707
 
708
  if self._max_history_age_days > 0:
709
+ cutoff = (reference_time or _utc_now()) - timedelta(
710
+ days=self._max_history_age_days
711
+ )
712
  filtered = [
713
  item
714
  for item in history
 
795
  minutes=self._display_session_inactivity_minutes
796
  )
797
 
798
+ def _build_rollup(
799
+ self,
800
+ history: list[dict[str, Any]],
801
+ bucket: str,
802
+ ) -> list[dict[str, Any]]:
803
  if not history:
804
  return []
805
 
tests/test_proxy_savings_history.py CHANGED
@@ -43,7 +43,9 @@ def _record_request(
43
  def test_savings_tracker_helpers_normalize_inputs_and_paths(tmp_path, monkeypatch):
44
  override_path = tmp_path / "custom-savings.json"
45
  monkeypatch.setenv(HEADROOM_SAVINGS_PATH_ENV_VAR, str(override_path))
46
- assert savings_tracker_module.get_default_savings_storage_path() == str(override_path)
 
 
47
 
48
  monkeypatch.delenv(HEADROOM_SAVINGS_PATH_ENV_VAR, raising=False)
49
  default_path = savings_tracker_module.get_default_savings_storage_path()
@@ -102,7 +104,11 @@ def test_savings_tracker_sanitizes_legacy_state_and_applies_retention(tmp_path):
102
  encoding="utf-8",
103
  )
104
 
105
- tracker = SavingsTracker(path=str(path), max_history_points=1, max_history_age_days=2)
 
 
 
 
106
  snapshot = tracker.snapshot()
107
 
108
  assert snapshot["schema_version"] == 2
@@ -113,7 +119,10 @@ def test_savings_tracker_sanitizes_legacy_state_and_applies_retention(tmp_path):
113
  "total_input_tokens": 0,
114
  "total_input_cost_usd": 0.0,
115
  }
116
- assert snapshot["display_session"] == savings_tracker_module._empty_display_session()
 
 
 
117
  assert snapshot["history"] == [
118
  {
119
  "timestamp": "2026-03-27T09:00:00Z",
@@ -143,7 +152,10 @@ def test_non_dict_savings_state_resets_to_default(tmp_path):
143
  "total_input_tokens": 0,
144
  "total_input_cost_usd": 0.0,
145
  }
146
- assert snapshot["display_session"] == savings_tracker_module._empty_display_session()
 
 
 
147
  assert snapshot["history"] == []
148
 
149
 
@@ -242,7 +254,9 @@ def test_litellm_resolution_and_savings_estimation_fallbacks(monkeypatch):
242
  ) == pytest.approx(0.2)
243
 
244
  fake_litellm.model_cost = {}
245
- assert savings_tracker_module._estimate_compression_savings_usd("gpt-4o", 100) == 0.0
 
 
246
  assert savings_tracker_module._estimate_input_cost_usd("gpt-4o", 100) == 0.0
247
 
248
  monkeypatch.setattr(
@@ -250,11 +264,19 @@ def test_litellm_resolution_and_savings_estimation_fallbacks(monkeypatch):
250
  "cost_per_token",
251
  lambda **kwargs: (_ for _ in ()).throw(RuntimeError("boom")),
252
  )
253
- assert savings_tracker_module._resolve_litellm_model("mystery-model") == "mystery-model"
254
- assert savings_tracker_module._estimate_compression_savings_usd("mystery-model", 100) == 0.0
 
 
 
 
 
 
255
 
256
  monkeypatch.setattr(savings_tracker_module, "LITELLM_AVAILABLE", False)
257
- assert savings_tracker_module._estimate_compression_savings_usd("gpt-4o", 100) == 0.0
 
 
258
  assert savings_tracker_module._estimate_input_cost_usd("gpt-4o", 100) == 0.0
259
 
260
 
@@ -309,7 +331,10 @@ def test_display_session_rolls_after_inactivity_and_counts_zero_savings_requests
309
  "_utc_now",
310
  lambda: datetime(2026, 3, 27, 9, 45, tzinfo=timezone.utc),
311
  )
312
- assert tracker.snapshot()["display_session"] == savings_tracker_module._empty_display_session()
 
 
 
313
 
314
  tracker.record_request(
315
  model="gpt-4o",
@@ -339,7 +364,11 @@ def test_display_session_rolls_after_inactivity_and_counts_zero_savings_requests
339
 
340
  def test_savings_tracker_rollups_preserve_spend_and_input_history(tmp_path, monkeypatch):
341
  path = tmp_path / "proxy_savings.json"
342
- tracker = SavingsTracker(path=str(path), max_history_points=100, max_history_age_days=30)
 
 
 
 
343
  monkeypatch.setattr(
344
  "headroom.proxy.savings_tracker._estimate_compression_savings_usd",
345
  lambda model, tokens_saved: tokens_saved / 1000.0,
@@ -476,7 +505,9 @@ def test_savings_tracker_rollups_preserve_spend_and_input_history(tmp_path, monk
476
  ]
477
 
478
 
479
- def test_stats_history_persists_across_restarts_and_stats_stays_compatible(tmp_path, monkeypatch):
 
 
480
  savings_path = tmp_path / "proxy_savings.json"
481
  monkeypatch.setenv("HEADROOM_SAVINGS_PATH", str(savings_path))
482
  monkeypatch.setattr(
@@ -514,7 +545,12 @@ def test_stats_history_persists_across_restarts_and_stats_stays_compatible(tmp_p
514
  assert history_data["display_session"]["tokens_saved"] == 40
515
  assert history_data["display_session"]["total_input_tokens"] == 120
516
  assert history_data["display_session"]["savings_percent"] == pytest.approx(25.0)
517
- assert list(history_data["series"].keys()) == ["hourly", "daily", "weekly", "monthly"]
 
 
 
 
 
518
  assert history_data["exports"]["available_series"][-2:] == ["weekly", "monthly"]
519
  assert history_data["series"]["hourly"][0]["total_input_tokens_delta"] == 120
520
  assert history_data["series"]["hourly"][0]["total_input_cost_usd_delta"] == pytest.approx(
@@ -522,7 +558,10 @@ def test_stats_history_persists_across_restarts_and_stats_stays_compatible(tmp_p
522
  )
523
 
524
  assert stats_data["display_session"] == history_data["display_session"]
525
- assert stats_data["persistent_savings"]["display_session"] == history_data["display_session"]
 
 
 
526
 
527
  with TestClient(create_app(config)) as client:
528
  history = client.get("/stats-history")
 
43
  def test_savings_tracker_helpers_normalize_inputs_and_paths(tmp_path, monkeypatch):
44
  override_path = tmp_path / "custom-savings.json"
45
  monkeypatch.setenv(HEADROOM_SAVINGS_PATH_ENV_VAR, str(override_path))
46
+ assert savings_tracker_module.get_default_savings_storage_path() == str(
47
+ override_path
48
+ )
49
 
50
  monkeypatch.delenv(HEADROOM_SAVINGS_PATH_ENV_VAR, raising=False)
51
  default_path = savings_tracker_module.get_default_savings_storage_path()
 
104
  encoding="utf-8",
105
  )
106
 
107
+ tracker = SavingsTracker(
108
+ path=str(path),
109
+ max_history_points=1,
110
+ max_history_age_days=2,
111
+ )
112
  snapshot = tracker.snapshot()
113
 
114
  assert snapshot["schema_version"] == 2
 
119
  "total_input_tokens": 0,
120
  "total_input_cost_usd": 0.0,
121
  }
122
+ assert (
123
+ snapshot["display_session"]
124
+ == savings_tracker_module._empty_display_session()
125
+ )
126
  assert snapshot["history"] == [
127
  {
128
  "timestamp": "2026-03-27T09:00:00Z",
 
152
  "total_input_tokens": 0,
153
  "total_input_cost_usd": 0.0,
154
  }
155
+ assert (
156
+ snapshot["display_session"]
157
+ == savings_tracker_module._empty_display_session()
158
+ )
159
  assert snapshot["history"] == []
160
 
161
 
 
254
  ) == pytest.approx(0.2)
255
 
256
  fake_litellm.model_cost = {}
257
+ assert (
258
+ savings_tracker_module._estimate_compression_savings_usd("gpt-4o", 100) == 0.0
259
+ )
260
  assert savings_tracker_module._estimate_input_cost_usd("gpt-4o", 100) == 0.0
261
 
262
  monkeypatch.setattr(
 
264
  "cost_per_token",
265
  lambda **kwargs: (_ for _ in ()).throw(RuntimeError("boom")),
266
  )
267
+ assert (
268
+ savings_tracker_module._resolve_litellm_model("mystery-model")
269
+ == "mystery-model"
270
+ )
271
+ assert (
272
+ savings_tracker_module._estimate_compression_savings_usd("mystery-model", 100)
273
+ == 0.0
274
+ )
275
 
276
  monkeypatch.setattr(savings_tracker_module, "LITELLM_AVAILABLE", False)
277
+ assert (
278
+ savings_tracker_module._estimate_compression_savings_usd("gpt-4o", 100) == 0.0
279
+ )
280
  assert savings_tracker_module._estimate_input_cost_usd("gpt-4o", 100) == 0.0
281
 
282
 
 
331
  "_utc_now",
332
  lambda: datetime(2026, 3, 27, 9, 45, tzinfo=timezone.utc),
333
  )
334
+ assert (
335
+ tracker.snapshot()["display_session"]
336
+ == savings_tracker_module._empty_display_session()
337
+ )
338
 
339
  tracker.record_request(
340
  model="gpt-4o",
 
364
 
365
  def test_savings_tracker_rollups_preserve_spend_and_input_history(tmp_path, monkeypatch):
366
  path = tmp_path / "proxy_savings.json"
367
+ tracker = SavingsTracker(
368
+ path=str(path),
369
+ max_history_points=100,
370
+ max_history_age_days=30,
371
+ )
372
  monkeypatch.setattr(
373
  "headroom.proxy.savings_tracker._estimate_compression_savings_usd",
374
  lambda model, tokens_saved: tokens_saved / 1000.0,
 
505
  ]
506
 
507
 
508
+ def test_stats_history_persists_across_restarts_and_stats_stays_compatible(
509
+ tmp_path, monkeypatch
510
+ ):
511
  savings_path = tmp_path / "proxy_savings.json"
512
  monkeypatch.setenv("HEADROOM_SAVINGS_PATH", str(savings_path))
513
  monkeypatch.setattr(
 
545
  assert history_data["display_session"]["tokens_saved"] == 40
546
  assert history_data["display_session"]["total_input_tokens"] == 120
547
  assert history_data["display_session"]["savings_percent"] == pytest.approx(25.0)
548
+ assert list(history_data["series"].keys()) == [
549
+ "hourly",
550
+ "daily",
551
+ "weekly",
552
+ "monthly",
553
+ ]
554
  assert history_data["exports"]["available_series"][-2:] == ["weekly", "monthly"]
555
  assert history_data["series"]["hourly"][0]["total_input_tokens_delta"] == 120
556
  assert history_data["series"]["hourly"][0]["total_input_cost_usd_delta"] == pytest.approx(
 
558
  )
559
 
560
  assert stats_data["display_session"] == history_data["display_session"]
561
+ assert (
562
+ stats_data["persistent_savings"]["display_session"]
563
+ == history_data["display_session"]
564
+ )
565
 
566
  with TestClient(create_app(config)) as client:
567
  history = client.get("/stats-history")