Skip to content

Commit 8a17aa6

Browse files
committed
fix(sleep): record call-local delta on real-usage (Azure/OpenCode) paths
- Add _record_delta(delta) and route the Azure/OpenCode real-usage accounting through it, so those backends also set the call-local _thread_local.delta — the last accounting path that did not (replay_one() saw 0/stale and fell back to a length estimate for these backends).
1 parent c8ff7f8 commit 8a17aa6

1 file changed

Lines changed: 22 additions & 12 deletions

File tree

‎skillopt_sleep/backend.py‎

Lines changed: 22 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -344,22 +344,26 @@ def __init__(self, model: str = "", timeout: int = 180) -> None:
344344
def _call(self, prompt: str, *, max_tokens: int = 1024) -> str:
345345
raise NotImplementedError
346346

347-
def _record_cost(self, prompt: str, response: str) -> int:
348-
"""THE single path to record an inference's token cost.
349-
350-
Computes the ``len//4`` delta, adds it to the aggregate ``_tokens``
351-
(atomically under ``_lock``) and records it as the call-local
352-
``_thread_local.delta`` for ``replay_one()``. Every inference path
353-
(``_cached_call``, ``attempt_with_tools``, ``reflect``) must route cost
354-
here so the aggregate and call-local totals always agree and no path
355-
under- or over-counts.
347+
def _record_delta(self, delta: int) -> int:
348+
"""Add a computed delta to the aggregate ``_tokens`` AND record it
349+
call-local. The one place both totals are updated, so they always agree.
356350
"""
357-
delta = len(prompt or "") // 4 + len(response or "") // 4
358351
with self._lock:
359352
self._tokens += delta
360353
self._thread_local.delta = delta
361354
return delta
362355

356+
def _record_cost(self, prompt: str, response: str) -> int:
357+
"""THE single path to record an inference's token cost (``len//4``).
358+
359+
Computes the ``len//4`` delta and delegates to ``_record_delta``. Every
360+
inference path (``_cached_call``, ``attempt_with_tools``, ``reflect``)
361+
must route cost here so the aggregate and call-local totals always agree
362+
and no path under- or over-counts.
363+
"""
364+
delta = len(prompt or "") // 4 + len(response or "") // 4
365+
return self._record_delta(delta)
366+
363367
def _reset_call_delta(self) -> None:
364368
"""Zero the call-local delta for a NO-CALL path.
365369
@@ -2473,7 +2477,10 @@ def _call(self, prompt: str, *, max_tokens: int = 1024, retries: int = 5) -> str
24732477
text = (resp.choices[0].message.content or "").strip()
24742478
try:
24752479
u = resp.usage
2476-
self._tokens += (getattr(u, "prompt_tokens", 0) or 0) + (getattr(u, "completion_tokens", 0) or 0)
2480+
self._record_delta(
2481+
(getattr(u, "prompt_tokens", 0) or 0)
2482+
+ (getattr(u, "completion_tokens", 0) or 0)
2483+
)
24772484
except Exception:
24782485
pass
24792486
if text:
@@ -2582,7 +2589,10 @@ def _call(self, prompt: str, *, max_tokens: int = 1024, retries: int = 5) -> str
25822589
text = (getattr(resp, "output_text", "") or "").strip()
25832590
try:
25842591
u = resp.usage
2585-
self._tokens += (getattr(u, "input_tokens", 0) or 0) + (getattr(u, "output_tokens", 0) or 0)
2592+
self._record_delta(
2593+
(getattr(u, "input_tokens", 0) or 0)
2594+
+ (getattr(u, "output_tokens", 0) or 0)
2595+
)
25862596
except Exception:
25872597
pass
25882598
if text:

0 commit comments

Comments
 (0)