File size: 13,836 Bytes
a04f9ea 6b117e7 a04f9ea c344089 a04f9ea 6b117e7 c344089 a04f9ea c344089 6b117e7 c344089 a04f9ea c344089 a04f9ea c344089 a04f9ea 6b117e7 a04f9ea 6b117e7 a04f9ea c344089 6b117e7 c344089 a04f9ea | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 | from __future__ import annotations
import unittest
from uuid import uuid4
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
from langchain_core.outputs import ChatGeneration, LLMResult
from app.chat_service import CLEARED_TOOL_OUTPUT_PLACEHOLDER
from app import telemetry
from app.telemetry import (
TurnUsageHandler,
context_window_stats,
estimate_cost_usd,
aggregate_cost_breakdown,
pop_turn_events,
pop_turn_signals,
pricing_for_model,
record_turn_signal,
record_turn_event,
record_turn_signal_max,
reset_turn_signals,
usage_totals,
)
class UsageTotalsTests(unittest.TestCase):
def test_sums_across_models_including_cache_details(self) -> None:
usage = {
"gemini-3.5-flash": {
"input_tokens": 1_000,
"output_tokens": 200,
"total_tokens": 1_200,
"input_token_details": {"cache_read": 400},
},
"claude-haiku-4-5-20251001": {
"input_tokens": 500,
"output_tokens": 100,
"total_tokens": 600,
"input_token_details": {"cache_read": 50, "cache_creation": 25},
},
}
totals = usage_totals(usage)
self.assertEqual(totals["input_tokens"], 1_500)
self.assertEqual(totals["output_tokens"], 300)
self.assertEqual(totals["total_tokens"], 1_800)
self.assertEqual(totals["cache_read_tokens"], 450)
self.assertEqual(totals["cache_creation_tokens"], 25)
def test_missing_details_default_to_zero(self) -> None:
totals = usage_totals({"m": {"input_tokens": 10, "output_tokens": 5}})
self.assertEqual(totals["cache_read_tokens"], 0)
self.assertEqual(totals["total_tokens"], 0)
class CostEstimateTests(unittest.TestCase):
def test_dated_model_name_matches_family_pricing(self) -> None:
self.assertIsNotNone(pricing_for_model("claude-haiku-4-5-20251001"))
self.assertIsNone(pricing_for_model("some-unknown-model"))
def test_cache_tokens_priced_separately(self) -> None:
# claude-haiku-4-5: input $1, output $5, cache read $0.10, write $1.25
# per MTok. input_tokens includes the cached buckets, so the plain
# bucket is 1M - 400k - 100k = 500k.
usage = {
"claude-haiku-4-5-20251001": {
"input_tokens": 1_000_000,
"output_tokens": 200_000,
"input_token_details": {
"cache_read": 400_000,
"cache_creation": 100_000,
},
}
}
expected = 0.5 * 1.00 + 0.4 * 0.10 + 0.1 * 1.25 + 0.2 * 5.00
self.assertAlmostEqual(estimate_cost_usd(usage), expected)
def test_cache_creation_without_write_rate_bills_as_input(self) -> None:
# gemini-3.5-flash: input $1.50, cache creation defaults to input rate
# because Gemini's hourly explicit-cache storage is not represented in
# per-turn usage metadata.
usage = {
"gemini-3.5-flash": {
"input_tokens": 1_000_000,
"output_tokens": 0,
"input_token_details": {"cache_creation": 200_000},
}
}
expected = 0.8 * 1.50 + 0.2 * 1.50
self.assertAlmostEqual(estimate_cost_usd(usage), expected)
def test_unknown_model_yields_none_not_zero(self) -> None:
self.assertIsNone(estimate_cost_usd({"mystery-model": {"input_tokens": 10}}))
mixed = {
"gemini-3.5-flash": {"input_tokens": 10, "output_tokens": 1},
"mystery-model": {"input_tokens": 10, "output_tokens": 1},
}
self.assertIsNone(estimate_cost_usd(mixed))
def test_deepseek_cost_breakdown_is_mutually_exclusive(self) -> None:
usage = {
"deepseek-v4-flash": {
"input_tokens": 1_000_000,
"output_tokens": 100_000,
"input_token_details": {"cache_read": 900_000},
}
}
breakdown = aggregate_cost_breakdown(usage)
self.assertIsNotNone(breakdown)
self.assertAlmostEqual(breakdown["cache_miss_input_usd"], 0.014)
self.assertAlmostEqual(breakdown["cache_read_input_usd"], 0.00252)
self.assertAlmostEqual(breakdown["output_usd"], 0.028)
self.assertAlmostEqual(
breakdown["total_usd"],
breakdown["cache_miss_input_usd"]
+ breakdown["cache_read_input_usd"]
+ breakdown["cache_creation_input_usd"]
+ breakdown["output_usd"],
)
class ContextWindowStatsTests(unittest.TestCase):
def test_counts_summaries_and_cleared_tool_outputs(self) -> None:
messages = [
HumanMessage(
content="Here is a summary of the conversation to date: ...",
additional_kwargs={"lc_source": "summarization"},
),
HumanMessage(content="What is RAG?"),
AIMessage(content="RAG is retrieval-augmented generation."),
ToolMessage(
content=CLEARED_TOOL_OUTPUT_PLACEHOLDER,
tool_call_id="call_1",
),
ToolMessage(content="$ ls\nwiki", tool_call_id="call_2"),
]
stats = context_window_stats(messages, CLEARED_TOOL_OUTPUT_PLACEHOLDER)
self.assertEqual(stats["context_messages"], 5)
self.assertEqual(stats["summary_messages"], 1)
self.assertEqual(stats["cleared_tool_outputs"], 1)
self.assertGreater(stats["context_tokens_approx"], 0)
def test_empty_context(self) -> None:
stats = context_window_stats([], CLEARED_TOOL_OUTPUT_PLACEHOLDER)
self.assertEqual(stats["context_messages"], 0)
self.assertEqual(stats["context_tokens_approx"], 0)
class TurnUsageHandlerTests(unittest.TestCase):
def test_counts_calls_and_aggregates_usage(self) -> None:
handler = TurnUsageHandler()
def result(input_tokens: int, output_tokens: int) -> LLMResult:
message = AIMessage(
content="ok",
usage_metadata={
"input_tokens": input_tokens,
"output_tokens": output_tokens,
"total_tokens": input_tokens + output_tokens,
},
response_metadata={"model_name": "gemini-3.5-flash"},
)
return LLMResult(generations=[[ChatGeneration(message=message)]])
handler.on_llm_end(result(100, 20))
handler.on_llm_end(result(50, 10))
self.assertEqual(handler.llm_calls, 2)
usage = handler.usage_metadata["gemini-3.5-flash"]
self.assertEqual(usage["input_tokens"], 150)
self.assertEqual(usage["output_tokens"], 30)
def test_records_one_explanatory_row_per_call(self) -> None:
handler = TurnUsageHandler()
run_id = uuid4()
handler.on_chat_model_start(
{},
[[HumanMessage(content="hello")]],
run_id=run_id,
metadata={"lc_source": "summarization"},
)
message = AIMessage(
content="summary",
usage_metadata={
"input_tokens": 100,
"output_tokens": 20,
"total_tokens": 120,
"input_token_details": {"cache_read": 80},
},
response_metadata={"model_name": "deepseek-v4-flash"},
)
handler.on_llm_end(
LLMResult(generations=[[ChatGeneration(message=message)]]),
run_id=run_id,
)
call = handler.model_calls[0]
self.assertEqual(call["source"], "summarization")
self.assertEqual(call["cache_read_tokens"], 80)
self.assertEqual(call["cache_miss_tokens"], 20)
self.assertTrue(call["cache_details_reported"])
self.assertGreater(call["request_context_tokens_approx"], 0)
self.assertAlmostEqual(
call["cost"]["total_usd"],
estimate_cost_usd({"deepseek-v4-flash": message.usage_metadata}),
)
class LangchainOpenAICacheFieldContractTests(unittest.TestCase):
"""Pin the langchain-openai usage conversion our cache accounting rides on.
DeepSeek (and OpenAI) report cached prompt tokens as
``prompt_tokens_details.cached_tokens``; langchain-openai must surface that
as ``usage_metadata.input_token_details.cache_read`` or TurnUsageHandler's
cache buckets (and the ~50x DeepSeek cache-read discount) silently read 0.
"""
def _convert(self, payload: dict) -> dict:
try:
from langchain_openai.chat_models.base import _create_usage_metadata
except ImportError as exc:
self.fail(
"langchain_openai.chat_models.base._create_usage_metadata is no "
f"longer importable ({exc}). A langchain-openai upgrade moved "
"the usage conversion; re-pin the prompt_tokens_details."
"cached_tokens -> input_token_details.cache_read mapping "
"against its new location."
)
return _create_usage_metadata(payload)
def test_cached_tokens_map_to_cache_read_details(self) -> None:
usage = self._convert(
{
"prompt_tokens": 1_000,
"completion_tokens": 100,
"total_tokens": 1_100,
"prompt_tokens_details": {"cached_tokens": 900},
}
)
# LangChain convention: input_tokens INCLUDES the cached bucket; the
# cost code carves cache_read out instead of adding it on top.
self.assertEqual(usage["input_tokens"], 1_000)
self.assertEqual(usage["output_tokens"], 100)
self.assertEqual(usage["input_token_details"]["cache_read"], 900)
def test_converted_usage_flows_through_turn_usage_handler(self) -> None:
usage = self._convert(
{
"prompt_tokens": 1_000,
"completion_tokens": 100,
"total_tokens": 1_100,
"prompt_tokens_details": {"cached_tokens": 900},
}
)
handler = TurnUsageHandler()
run_id = uuid4()
handler.on_chat_model_start({}, [[HumanMessage(content="q")]], run_id=run_id)
message = AIMessage(
content="answer",
usage_metadata=usage,
response_metadata={"model_name": "deepseek-v4-flash"},
)
handler.on_llm_end(
LLMResult(generations=[[ChatGeneration(message=message)]]),
run_id=run_id,
)
call = handler.model_calls[0]
self.assertEqual(call["cache_read_tokens"], 900)
self.assertEqual(call["cache_miss_tokens"], 100)
self.assertTrue(call["cache_details_reported"])
# deepseek-v4-flash: $0.14 miss / $0.0028 cache-read / $0.28 output
# per MTok, so the cache discount must show up in the billed cost.
expected = (100 * 0.14 + 900 * 0.0028 + 100 * 0.28) / 1_000_000
self.assertAlmostEqual(call["cost"]["total_usd"], expected)
class TurnSignalRegistryTests(unittest.TestCase):
def test_accumulates_and_pops_per_turn(self) -> None:
reset_turn_signals("turn-a")
record_turn_signal("turn-a", "dropped_messages", 3)
record_turn_signal("turn-a", "dropped_messages", 2)
record_turn_signal("turn-a", "truncated_tool_outputs", 1)
signals = pop_turn_signals("turn-a")
self.assertEqual(signals["dropped_messages"], 5)
self.assertEqual(signals["truncated_tool_outputs"], 1)
# Popping clears the entry: a second pop is empty.
self.assertEqual(pop_turn_signals("turn-a"), {})
def test_structured_events_are_isolated_and_popped(self) -> None:
reset_turn_signals("turn-events")
record_turn_event("turn-events", {"event": "summarization", "tokens": 9})
self.assertEqual(
pop_turn_events("turn-events"),
[{"event": "summarization", "tokens": 9}],
)
self.assertEqual(pop_turn_events("turn-events"), [])
def test_turns_are_isolated_and_noops_are_ignored(self) -> None:
reset_turn_signals("turn-x")
reset_turn_signals("turn-y")
record_turn_signal("turn-x", "dropped_messages", 4)
record_turn_signal("turn-y", "dropped_messages", 0) # no-op
record_turn_signal("", "dropped_messages", 9) # no turn id -> ignored
self.assertEqual(pop_turn_signals("turn-x"), {"dropped_messages": 4})
self.assertEqual(pop_turn_signals("turn-y"), {})
def test_max_records_peak_not_sum(self) -> None:
# A middleware fires once per model call within a turn; the max across
# calls is the real per-turn figure, not the (overlapping) sum.
reset_turn_signals("turn-m")
record_turn_signal_max("turn-m", "dropped_messages", 5)
record_turn_signal_max("turn-m", "dropped_messages", 8)
record_turn_signal_max("turn-m", "dropped_messages", 3)
self.assertEqual(pop_turn_signals("turn-m"), {"dropped_messages": 8})
def test_overflow_evicts_oldest_keeps_recent(self) -> None:
# The cap must evict the OLDEST turns, never wipe the freshest in-flight
# ones (the concurrency-safety fix).
reset_turn_signals("oldest-turn")
record_turn_signal_max("oldest-turn", "dropped_messages", 1)
for i in range(telemetry._MAX_TRACKED_TURNS + 5):
reset_turn_signals(f"filler-{i}")
reset_turn_signals("recent-turn")
record_turn_signal_max("recent-turn", "dropped_messages", 7)
self.assertEqual(pop_turn_signals("recent-turn"), {"dropped_messages": 7})
self.assertEqual(pop_turn_signals("oldest-turn"), {}) # evicted
if __name__ == "__main__":
unittest.main()
|