diff --git a/strix/config/settings.py b/strix/config/settings.py index e9ed84272..f6c85dd03 100644 --- a/strix/config/settings.py +++ b/strix/config/settings.py @@ -59,9 +59,9 @@ class LlmSettings(BaseSettings): alias="STRIX_PROMPT_CACHE", ) # Providers cache prompts in fixed-size token blocks, so a fully cached prompt - # reads back rounded down to a multiple of this. 64 is what the GLM calls in - # local runs showed; it's a per-deployment setting (vLLM defaults to 16). - cache_block_tokens: int = Field(default=64, ge=1, alias="STRIX_CACHE_BLOCK_TOKENS") + # can read back up to a block short. 128 covers the largest common size + # (OpenAI; DeepSeek and GLM use 64, vLLM defaults to 16). + cache_block_tokens: int = Field(default=128, ge=1, alias="STRIX_CACHE_BLOCK_TOKENS") disable_streaming: bool = Field( default=False, alias="LLM_DISABLE_STREAMING", diff --git a/strix/report/usage.py b/strix/report/usage.py index f84d02e8e..bb20b254d 100644 --- a/strix/report/usage.py +++ b/strix/report/usage.py @@ -99,12 +99,11 @@ class LLMUsageLedger: tally.cached_tokens += cached_tokens if agent_id: previous = self._last_input_tokens.get(agent_id, 0) - # The previous prompt is a prefix of this one, so every full block of it - # should read back cached. A shrinking prompt means compaction rewrote - # it, so a miss is expected. - expected = previous - (previous % cache_block_tokens) - missed = expected - cached_tokens - if input_tokens >= previous and missed > 0: + # The previous prompt is a prefix of this one, so all of it but a + # partial last block should read back cached. A shrinking prompt means + # compaction rewrote it, so a miss is expected. + missed = previous - cached_tokens + if input_tokens >= previous and missed >= cache_block_tokens: tally.cache_misses += 1 tally.missed_tokens += missed self._last_input_tokens[agent_id] = input_tokens diff --git a/tests/test_cost_tracking.py b/tests/test_cost_tracking.py index c9bc283b5..4c0ac9223 100644 --- a/tests/test_cost_tracking.py +++ b/tests/test_cost_tracking.py @@ -307,7 +307,7 @@ def test_provider_tally_survives_run_record_round_trip() -> None: input_tokens=input_tokens, cached_tokens=cached_tokens, cost=cost, - cache_block_tokens=64, + cache_block_tokens=128, ) restored = LLMUsageLedger() @@ -329,8 +329,8 @@ def test_provider_tally_counts_cache_misses_per_agent() -> None: ledger = LLMUsageLedger() calls = [ ("Z.AI", "a1", 1000, 0), # first call: nothing to miss - ("Z.AI", "a1", 1200, 960), # previous 1000 cached, rounded down to 64s - ("DeepInfra", "a1", 1500, 200), # 1152 of the previous 1200 due, 952 lost + ("Z.AI", "a1", 1200, 960), # 40 short of the previous 1000: within a block + ("DeepInfra", "a1", 1500, 200), # 1000 of the previous 1200 lost ("Z.AI", "a2", 800, 0), # another agent's first call ("Z.AI", "a1", 600, 0), # prompt shrank: compaction, not a miss ] @@ -341,12 +341,12 @@ def test_provider_tally_counts_cache_misses_per_agent() -> None: input_tokens=input_tokens, cached_tokens=cached_tokens, cost=0.0, - cache_block_tokens=64, + cache_block_tokens=128, ) providers = ledger.to_record()["providers"] assert providers["DeepInfra"]["cache_misses"] == 1 - assert providers["DeepInfra"]["missed_tokens"] == 952 + assert providers["DeepInfra"]["missed_tokens"] == 1000 assert providers["Z.AI"]["cache_misses"] == 0