Skip to content

Commit 0712e44

Browse files
committed
fix(sleep): single-owner token accounting for provider-usage backends
- _cached_call/_reflect now record provider usage exactly once; a backend that reports usage returns (text, usage) and we charge the exact total, falling back to the len//4 estimate only when usage is unavailable. AzureOpenAI and AzureResponses no longer self-record (no double charge). - OpenCode error path routes prompt-only cost through _record_delta (no manual _tokens/_thread_local.delta update). - Added fake-client regressions: Azure single call, empty-retry+success accumulation, cache-hit no-charge, and OpenCode error-path.
1 parent 8a17aa6 commit 0712e44

3 files changed

Lines changed: 192 additions & 24 deletions

File tree

‎skillopt_sleep/backend.py‎

Lines changed: 29 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -390,9 +390,17 @@ def _cached_call(self, key: str, prompt: str, *, max_tokens: int = 1024) -> str:
390390
# The model call is intentionally outside the lock so parallel workers
391391
# over the same backend overlap; a concurrent miss may duplicate a call,
392392
# but _cache/_tokens reads+writes below are atomic.
393-
out = self._call(prompt, max_tokens=max_tokens)
394-
# Charge every real call's tokens AND record it call-local in one place.
395-
delta = self._record_cost(prompt, out)
393+
result = self._call(prompt, max_tokens=max_tokens)
394+
# Charge every real call's tokens ONCE, in one place. Backends that can
395+
# report provider usage return (text, usage); record the exact total.
396+
# Otherwise fall back to the len//4 length estimate. A backend that
397+
# reports usage must NOT also record it internally, or it is double-charged.
398+
if isinstance(result, tuple):
399+
out, usage = result
400+
delta = self._record_delta(usage) if usage is not None else self._record_cost(prompt, out)
401+
else:
402+
out = result
403+
delta = self._record_cost(prompt, out)
396404
with self._lock:
397405
# The cache dedup below may reuse another worker's success, but the
398406
# model call above still consumed tokens (already charged).
@@ -579,7 +587,15 @@ def _explain(c: str) -> str:
579587
"Reply with ONLY the JSON array, no prose, no markdown fences."
580588
)
581589
raw = self._call(p, max_tokens=1024)
582-
self._record_cost(p, raw)
590+
if isinstance(raw, tuple):
591+
raw_text, usage = raw
592+
if usage is not None:
593+
self._record_delta(usage)
594+
else:
595+
self._record_cost(p, raw_text)
596+
raw = raw_text
597+
else:
598+
self._record_cost(p, raw)
583599
if ev is not None:
584600
ev.log("reflect", "exchange", target=target, attempt=attempt + 1,
585601
backend=self.name, model=self.model,
@@ -1439,10 +1455,7 @@ def attempt_with_tools(
14391455
]
14401456
except OpenCodeError as exc:
14411457
# Prompt-only cost on the error path (no response text).
1442-
delta = exc.prompt_chars // 4
1443-
with self._lock:
1444-
self._tokens += delta
1445-
self._thread_local.delta = delta
1458+
delta = self._record_delta(exc.prompt_chars // 4)
14461459
self.last_call_error = str(exc)
14471460
return "", []
14481461
self._record_cost(prompt, text)
@@ -2455,6 +2468,7 @@ def _call(self, prompt: str, *, max_tokens: int = 1024, retries: int = 5) -> str
24552468
client = self._get_client()
24562469
last_exc = None
24572470
n_attempts = max(1, retries)
2471+
usage_total = 0
24582472
for attempt in range(n_attempts):
24592473
try:
24602474
kwargs: Dict[str, Any] = {
@@ -2477,17 +2491,14 @@ def _call(self, prompt: str, *, max_tokens: int = 1024, retries: int = 5) -> str
24772491
text = (resp.choices[0].message.content or "").strip()
24782492
try:
24792493
u = resp.usage
2480-
self._record_delta(
2481-
(getattr(u, "prompt_tokens", 0) or 0)
2482-
+ (getattr(u, "completion_tokens", 0) or 0)
2483-
)
2494+
usage_total += (getattr(u, "prompt_tokens", 0) or 0) + (getattr(u, "completion_tokens", 0) or 0)
24842495
except Exception:
24852496
pass
24862497
if text:
24872498
# A recovered retry must not leave a stale error behind:
24882499
# last_call_error always reflects the LATEST outcome.
24892500
self.last_call_error = ""
2490-
return text
2501+
return text, usage_total
24912502
# empty but no exception: model genuinely returned nothing — one
24922503
# quick retry can help (reasoning models occasionally yield empty)
24932504
last_exc = "empty-response"
@@ -2506,7 +2517,7 @@ def _call(self, prompt: str, *, max_tokens: int = 1024, retries: int = 5) -> str
25062517
self.last_call_error = (
25072518
f"{self.deployment}: empty response on all {n_attempts} attempts"
25082519
)
2509-
return ""
2520+
return "", usage_total
25102521

25112522

25122523
class AzureResponsesBackend(AzureOpenAIBackend):
@@ -2577,6 +2588,7 @@ def _call(self, prompt: str, *, max_tokens: int = 1024, retries: int = 5) -> str
25772588
last = None
25782589
base_ep = self._next_endpoint() # this call's primary endpoint
25792590
base_idx = self.endpoints.index(base_ep)
2591+
usage_total = 0
25802592
for attempt in range(max(1, retries)):
25812593
# on retry, fail over to the other endpoint(s)
25822594
ep = self.endpoints[(base_idx + attempt) % len(self.endpoints)]
@@ -2589,20 +2601,17 @@ def _call(self, prompt: str, *, max_tokens: int = 1024, retries: int = 5) -> str
25892601
text = (getattr(resp, "output_text", "") or "").strip()
25902602
try:
25912603
u = resp.usage
2592-
self._record_delta(
2593-
(getattr(u, "input_tokens", 0) or 0)
2594-
+ (getattr(u, "output_tokens", 0) or 0)
2595-
)
2604+
usage_total += (getattr(u, "input_tokens", 0) or 0) + (getattr(u, "output_tokens", 0) or 0)
25962605
except Exception:
25972606
pass
25982607
if text:
2599-
return text
2608+
return text, usage_total
26002609
last = "empty-response"
26012610
except Exception as e: # noqa: BLE001
26022611
last = e
26032612
if attempt < retries - 1:
26042613
_t.sleep(min(8.0, (2 ** attempt) * 0.5) + _r.random() * 0.4)
2605-
return ""
2614+
return "", usage_total
26062615

26072616

26082617
def get_backend(

‎tests/test_azure_openai_compat.py‎

Lines changed: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -151,7 +151,7 @@ class TestRequestKwargs(unittest.TestCase):
151151
def test_compat_mode_sends_standard_max_tokens(self):
152152
be = _backend_with(["hi"], COMPAT_ENV)
153153
with mock.patch.dict(os.environ, COMPAT_ENV, clear=True):
154-
out = be._call("p", retries=1)
154+
out, _ = be._call("p", retries=1)
155155
self.assertEqual(out, "hi")
156156
(call,) = be._client.chat.completions.calls
157157
self.assertEqual(call["max_tokens"], 8192)
@@ -215,23 +215,23 @@ def test_recovered_retry_clears_last_call_error(self):
215215
be = _backend_with([RuntimeError("transient boom"), "recovered"], COMPAT_ENV)
216216
with mock.patch.dict(os.environ, COMPAT_ENV, clear=True), \
217217
mock.patch("time.sleep"):
218-
out = be._call("p", retries=2)
218+
out, _ = be._call("p", retries=2)
219219
self.assertEqual(out, "recovered")
220220
self.assertEqual(be.last_call_error, "")
221221

222222
def test_all_empty_responses_set_diagnostic(self):
223223
be = _backend_with(["", ""], COMPAT_ENV)
224224
with mock.patch.dict(os.environ, COMPAT_ENV, clear=True), \
225225
mock.patch("time.sleep"):
226-
out = be._call("p", retries=2)
226+
out, _ = be._call("p", retries=2)
227227
self.assertEqual(out, "")
228228
self.assertIn("empty response on all 2 attempts", be.last_call_error)
229229

230230
def test_persistent_exception_is_surfaced(self):
231231
be = _backend_with([RuntimeError("boom-1"), RuntimeError("boom-2")], COMPAT_ENV)
232232
with mock.patch.dict(os.environ, COMPAT_ENV, clear=True), \
233233
mock.patch("time.sleep"):
234-
out = be._call("p", retries=2)
234+
out, _ = be._call("p", retries=2)
235235
self.assertEqual(out, "")
236236
self.assertIn("boom-2", be.last_call_error)
237237

Lines changed: 159 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,159 @@
1+
"""Azure provider-usage accounting regressions: single-owner charging.
2+
3+
A backend that reports provider usage (AzureOpenAI / AzureResponses) must be
4+
charged exactly once in ``_cached_call`` — with the provider's own token count,
5+
never double-charged by the ``len//4`` length estimate, and never charged on a
6+
cache hit. (The maintainer's reproduction: a 30-token provider usage was being
7+
recorded as a 110-token length estimate because ``_call`` recorded usage and
8+
``_cached_call`` then recorded the length estimate on top.)
9+
10+
Also covers the OpenCode error path routing through ``_record_delta``.
11+
"""
12+
from __future__ import annotations
13+
14+
from types import SimpleNamespace
15+
from unittest import mock
16+
17+
from skillopt_sleep.backend import AzureOpenAIBackend, AzureResponsesBackend, OpenCodeCliBackend
18+
19+
20+
class _ChatResp:
21+
def __init__(self, text, prompt_tokens, completion_tokens):
22+
self.choices = [SimpleNamespace(message=SimpleNamespace(content=text))]
23+
self.usage = SimpleNamespace(prompt_tokens=prompt_tokens, completion_tokens=completion_tokens)
24+
25+
26+
class _FakeChatClient:
27+
"""Scripted chat.completions.create returning a _ChatResp or raising."""
28+
29+
def __init__(self, replies):
30+
self.replies = list(replies)
31+
self.calls = []
32+
33+
def create(self, **kwargs):
34+
self.calls.append(kwargs)
35+
item = self.replies.pop(0)
36+
if isinstance(item, Exception):
37+
raise item
38+
return item
39+
40+
41+
class _ResponsesResp:
42+
def __init__(self, text, input_tokens, output_tokens):
43+
self.output_text = text
44+
self.usage = SimpleNamespace(input_tokens=input_tokens, output_tokens=output_tokens)
45+
46+
47+
class _FakeResponsesClient:
48+
def __init__(self, replies):
49+
self.replies = list(replies)
50+
self.calls = []
51+
52+
def create(self, **kwargs):
53+
self.calls.append(kwargs)
54+
item = self.replies.pop(0)
55+
if isinstance(item, Exception):
56+
raise item
57+
return item
58+
59+
60+
def _azure_chat(replies):
61+
be = AzureOpenAIBackend(deployment="gpt-5.5")
62+
be._client = SimpleNamespace(chat=SimpleNamespace(completions=_FakeChatClient(replies)))
63+
return be
64+
65+
66+
def _azure_responses(replies):
67+
be = AzureResponsesBackend(deployment="gpt-5.5", endpoints=["https://t.openai.azure.com/"])
68+
fake = SimpleNamespace(responses=_FakeResponsesClient(replies))
69+
be._next_endpoint = lambda: be.endpoints[0]
70+
be._client_for = lambda ep: fake
71+
return be
72+
73+
74+
def test_azure_chat_single_call_charges_exact_usage():
75+
"""A single Azure chat call must charge the provider usage, not len//4."""
76+
be = _azure_chat([_ChatResp("ok", 10, 20)])
77+
with mock.patch("time.sleep"):
78+
out = be._cached_call("k:1", "x" * 400)
79+
assert out == "ok"
80+
# provider usage = 30; len//4 of 400+2 would be ~100 — must NOT be that.
81+
assert be._tokens == 30, f"expected exact provider usage 30, got {be._tokens}"
82+
assert be.token_delta() == 30
83+
84+
85+
def test_azure_chat_empty_retry_then_success_accumulates():
86+
"""Empty-response retry + success must accumulate usage across paid attempts."""
87+
be = _azure_chat([_ChatResp("", 7, 0), _ChatResp("ok", 10, 20)])
88+
with mock.patch("time.sleep"):
89+
out = be._cached_call("k:1", "hello")
90+
assert out == "ok"
91+
assert be._tokens == 7 + 30, f"expected accumulated 37, got {be._tokens}"
92+
assert be.token_delta() == 37
93+
94+
95+
def test_azure_chat_cache_hit_does_not_charge():
96+
"""A cache hit must reset the call-local delta and leave the aggregate alone."""
97+
be = _azure_chat([_ChatResp("ok", 10, 20)])
98+
with mock.patch("time.sleep"):
99+
be._cached_call("k:1", "hello")
100+
assert be._tokens == 30
101+
assert be.token_delta() == 30
102+
with mock.patch("time.sleep"):
103+
out2 = be._cached_call("k:1", "hello")
104+
assert out2 == "ok"
105+
assert be._tokens == 30, "cache hit changed the aggregate"
106+
assert be.token_delta() == 0, "cache hit leaked the prior delta"
107+
108+
109+
def test_azure_responses_single_call_charges_exact_usage():
110+
be = _azure_responses([_ResponsesResp("ok", 12, 18)])
111+
with mock.patch("time.sleep"):
112+
out = be._cached_call("k:1", "x" * 400)
113+
assert out == "ok"
114+
assert be._tokens == 30, f"expected exact provider usage 30, got {be._tokens}"
115+
assert be.token_delta() == 30
116+
117+
118+
def test_azure_responses_empty_retry_then_success_accumulates():
119+
be = _azure_responses([_ResponsesResp("", 5, 0), _ResponsesResp("ok", 12, 18)])
120+
with mock.patch("time.sleep"):
121+
out = be._cached_call("k:1", "hello")
122+
assert out == "ok"
123+
assert be._tokens == 5 + 30, f"expected accumulated 35, got {be._tokens}"
124+
assert be.token_delta() == 35
125+
126+
127+
def test_azure_responses_cache_hit_does_not_charge():
128+
be = _azure_responses([_ResponsesResp("ok", 12, 18)])
129+
with mock.patch("time.sleep"):
130+
be._cached_call("k:1", "hello")
131+
assert be._tokens == 30
132+
with mock.patch("time.sleep"):
133+
be._cached_call("k:1", "hello")
134+
assert be._tokens == 30
135+
assert be.token_delta() == 0
136+
137+
138+
def test_opencode_error_path_uses_record_delta(monkeypatch):
139+
"""The OpenCode error path must route prompt-only cost through _record_delta."""
140+
import contextlib
141+
from types import SimpleNamespace
142+
143+
import skillopt_sleep.backend as bm
144+
from skillopt_sleep.backend import OpenCodeCliBackend
145+
146+
b = OpenCodeCliBackend(model="", opencode_path="opencode", tool_replay=True)
147+
monkeypatch.setattr(bm, "_opencode_temporary_workspace", lambda *a, **k: contextlib.nullcontext())
148+
149+
def _fail(*args, **kwargs):
150+
raise bm.OpenCodeError("boom", prompt_chars=100)
151+
152+
monkeypatch.setattr(bm, "_prepare_opencode_replay_project", _fail)
153+
154+
task = SimpleNamespace(intent="intent", context_excerpt="ctx")
155+
out, called = b.attempt_with_tools(task, skill="s", memory="m", tools=["search"])
156+
157+
assert out == "" and called == []
158+
assert b.token_delta() == 100 // 4, "OpenCode error path did not route through _record_delta"
159+
assert b._tokens == 100 // 4, f"expected 25, got {b._tokens}"

0 commit comments

Comments
 (0)