From fc7238163614b6fe24c3d3481d3349b5d5fef9d8 Mon Sep 17 00:00:00 2001 From: webdevtodayjason Date: Tue, 29 Sep 2026 22:57:13 -0500 Subject: [PATCH] Count proxied inference tokens in metrics --- ainode/api/server.py | 145 ++++++++++++++++++++++++++- ainode/embeddings/api_routes.py | 56 ++++++++--- ainode/web/static/js/metrics-data.js | 9 +- ainode/web/static/js/metrics.js | 24 +++-- tests/test_embeddings_routing.py | 9 +- tests/test_messages_proxy.py | 32 +++++- tests/test_metrics_chart.py | 30 ++---- tests/test_routing_caps.py | 2 +- tests/test_ui_truth.py | 2 +- 9 files changed, 249 insertions(+), 60 deletions(-) diff --git a/ainode/api/server.py b/ainode/api/server.py index e99678ab..020d0668 100644 --- a/ainode/api/server.py +++ b/ainode/api/server.py @@ -1944,6 +1944,139 @@ def _missing_form_model() -> web.Response: ) +def _token_count(value) -> Optional[int]: + """A non-negative integer token count, or None for an absent/bad value.""" + if isinstance(value, bool) or not isinstance(value, (int, float)): + return None + if value < 0 or int(value) != value: + return None + return int(value) + + +def response_output_tokens(payload) -> Optional[int]: + """Read generated-token usage from an OpenAI, Responses or Messages reply. + + vLLM exposes more than one protocol through this proxy. OpenAI completion + replies call the figure ``completion_tokens``; the Responses and Anthropic + Messages APIs call it ``output_tokens``. A Responses streaming completion + nests the final usage block under ``response``. + """ + if not isinstance(payload, dict): + return None + candidates = [payload] + response = payload.get("response") + if isinstance(response, dict): + candidates.append(response) + for candidate in candidates: + usage = candidate.get("usage") + if not isinstance(usage, dict): + continue + for key in ("completion_tokens", "output_tokens"): + count = _token_count(usage.get(key)) + if count is not None: + return count + return None + + +def response_body_output_tokens(body: bytes) -> int: + """Generated tokens reported by one non-streamed JSON response.""" + try: + payload = json.loads(body) + except (TypeError, ValueError, UnicodeDecodeError): + return 0 + return response_output_tokens(payload) or 0 + + +def _sse_has_output_delta(payload) -> bool: + """Whether one SSE event represents generated output when usage is absent.""" + if not isinstance(payload, dict): + return False + event_type = payload.get("type") + if isinstance(event_type, str) and event_type.endswith(".delta"): + return True + delta = payload.get("delta") + if isinstance(delta, dict): + return any(delta.get(key) not in (None, "", [], {}) for key in ( + "content", "reasoning_content", "text", "thinking", + "partial_json", "tool_calls", + )) + choices = payload.get("choices") + if not isinstance(choices, list): + return False + for choice in choices: + if not isinstance(choice, dict): + continue + delta = choice.get("delta") + if isinstance(delta, dict) and any( + delta.get(key) not in (None, "", [], {}) + for key in ("content", "reasoning_content", "tool_calls") + ): + return True + if choice.get("text") not in (None, ""): + return True + return False + + +class SSETokenCounter: + """Incrementally count an SSE reply without changing the forwarded bytes. + + A protocol-reported usage total wins. When a client did not request a final + usage block, each output delta is counted as one streamed token, matching + vLLM's token-at-a-time completion events. Framing is buffered because an + aiohttp chunk may split anywhere inside an SSE event. + """ + + def __init__(self) -> None: + self._buffer = bytearray() + self._reported: Optional[int] = None + self._deltas = 0 + + def feed(self, chunk: bytes) -> None: + self._buffer.extend(chunk) + while True: + raw = self._pop_event() + if raw is None: + return + self._consume(raw) + + def finish(self) -> int: + if self._buffer: + self._consume(bytes(self._buffer)) + self._buffer.clear() + return self._reported if self._reported is not None else self._deltas + + def _pop_event(self) -> Optional[bytes]: + positions = [ + (self._buffer.find(marker), marker) + for marker in (b"\n\n", b"\r\n\r\n") + ] + positions = [(position, marker) for position, marker in positions if position >= 0] + if not positions: + return None + position, marker = min(positions, key=lambda item: item[0]) + raw = bytes(self._buffer[:position]) + del self._buffer[:position + len(marker)] + return raw + + def _consume(self, event: bytes) -> None: + data = [] + for line in event.splitlines(): + if line.startswith(b"data:"): + data.append(line[5:].lstrip()) + raw = b"\n".join(data) + if not raw or raw == b"[DONE]": + return + try: + payload = json.loads(raw) + except (ValueError, UnicodeDecodeError): + return + reported = response_output_tokens(payload) + if reported is not None: + self._reported = reported + elif _sse_has_output_delta(payload): + self._deltas += 1 + + async def proxy_to_vllm(request: web.Request) -> web.StreamResponse: """Forward the request to the node serving the requested model (F1 federation).""" config: NodeConfig = request.app["config"] @@ -2082,6 +2215,7 @@ async def proxy_to_vllm(request: web.Request) -> web.StreamResponse: async with session.request(request.method, vllm_url, **kwargs) as upstream: is_sse = "text/event-stream" in upstream.headers.get("Content-Type", "") if is_sse: + token_counter = SSETokenCounter() resp = web.StreamResponse( status=upstream.status, headers={ @@ -2093,9 +2227,13 @@ async def proxy_to_vllm(request: web.Request) -> web.StreamResponse: ) await resp.prepare(request) async for chunk in upstream.content.iter_any(): + token_counter.feed(chunk) await resp.write(chunk) await resp.write_eof() - collector.record_request(model, (time.time() - start_time) * 1000, error=False) + collector.record_request( + model, (time.time() - start_time) * 1000, + tokens_generated=token_counter.finish(), error=False, + ) return resp body = await upstream.read() # A multimodal-limit 400 is a ROUTING miss, not a bad request: @@ -2113,7 +2251,10 @@ async def proxy_to_vllm(request: web.Request) -> web.StreamResponse: refused_labels.append(label) last_err = f"{label} serves '{model}' without images" continue - collector.record_request(model, (time.time() - start_time) * 1000, error=False) + collector.record_request( + model, (time.time() - start_time) * 1000, + tokens_generated=response_body_output_tokens(body), error=False, + ) return web.Response( status=upstream.status, body=body, headers={SERVED_BY_HEADER: served_by}, diff --git a/ainode/embeddings/api_routes.py b/ainode/embeddings/api_routes.py index d15813f6..1bb6b5d8 100644 --- a/ainode/embeddings/api_routes.py +++ b/ainode/embeddings/api_routes.py @@ -20,7 +20,9 @@ from __future__ import annotations import asyncio +import json import logging +import time from typing import List, Optional import aiohttp @@ -85,6 +87,21 @@ def _error(message: str, *, code: str = "invalid_request_error", status: int = 4 ) +def _prompt_tokens(body: bytes) -> int: + """Prompt tokens reported by an OpenAI-compatible embeddings reply.""" + try: + payload = json.loads(body) + except (TypeError, ValueError, UnicodeDecodeError): + return 0 + usage = payload.get("usage") if isinstance(payload, dict) else None + value = usage.get("prompt_tokens") if isinstance(usage, dict) else None + if isinstance(value, bool) or not isinstance(value, (int, float)): + return 0 + if value < 0 or int(value) != value: + return 0 + return int(value) + + def fleet_candidates(app, model: str) -> list: """Every ``(host, port)`` in the fleet serving ``model`` right now. @@ -173,21 +190,31 @@ async def handle_v1_embeddings(request: web.Request) -> web.Response: takes a list of strings, not a body. """ manager: EmbeddingManager = request.app["embedding_manager"] + collector = request.app.get("metrics_collector") + started = time.time() + metric_model = "unknown" + + def finish(response: web.Response, tokens: int = 0) -> web.Response: + if collector is not None: + collector.record_request( + metric_model, (time.time() - started) * 1000, + tokens_generated=tokens, error=response.status >= 400, + ) + return response body_bytes = await request.read() try: - import json as _json - - body = _json.loads(body_bytes) if body_bytes else None + body = json.loads(body_bytes) if body_bytes else None except Exception: - return _error("Invalid JSON body") + return finish(_error("Invalid JSON body")) if not isinstance(body, dict): - return _error("Body must be a JSON object") + return finish(_error("Body must be a JSON object")) model_id = body.get("model") if not model_id or not isinstance(model_id, str): - return _error("'model' is required and must be a string") + return finish(_error("'model' is required and must be a string")) + metric_model = model_id # Tag the request so the server-view log shows the embedding model, whichever # of the two paths answers it. @@ -199,36 +226,37 @@ async def handle_v1_embeddings(request: web.Request) -> web.Response: # --- the fleet, if anything in it serves this model id -------------------- candidates = fleet_candidates(request.app, model_id) if candidates: - return await forward_to_fleet(request, model_id, body_bytes, candidates) + response = await forward_to_fleet(request, model_id, body_bytes, candidates) + return finish(response, _prompt_tokens(response.body)) # --- otherwise this process, on the CPU ---------------------------------- raw_input = body.get("input") if raw_input is None: - return _error("'input' is required (string or array of strings)") + return finish(_error("'input' is required (string or array of strings)")) if isinstance(raw_input, str): texts: List[str] = [raw_input] elif isinstance(raw_input, list): if not all(isinstance(x, str) for x in raw_input): - return _error("'input' array must contain only strings") + return finish(_error("'input' array must contain only strings")) texts = raw_input else: - return _error("'input' must be a string or array of strings") + return finish(_error("'input' must be a string or array of strings")) try: vectors = await manager.aembed(model_id, texts) except RuntimeError as exc: - return _error(str(exc), code="dependency_missing", status=503) + return finish(_error(str(exc), code="dependency_missing", status=503)) except Exception as exc: # pragma: no cover - defensive logger.exception("embedding failure for %s", model_id) - return _error(f"embedding failed: {exc}", code="server_error", status=500) + return finish(_error(f"embedding failed: {exc}", code="server_error", status=500)) total_tokens = sum(_approx_tokens(t) for t in texts) data = [ {"object": "embedding", "embedding": vec, "index": idx} for idx, vec in enumerate(vectors) ] - return web.json_response( + return finish(web.json_response( { "object": "list", "data": data, @@ -238,7 +266,7 @@ async def handle_v1_embeddings(request: web.Request) -> web.Response: "total_tokens": total_tokens, }, } - ) + ), total_tokens) async def handle_list_embedding_models(request: web.Request) -> web.Response: diff --git a/ainode/web/static/js/metrics-data.js b/ainode/web/static/js/metrics-data.js index 471127f8..87f6239d 100644 --- a/ainode/web/static/js/metrics-data.js +++ b/ainode/web/static/js/metrics-data.js @@ -56,14 +56,6 @@ // sawtooths back to zero the moment the process came up, so a gap in the other // series can be told apart from a dead GPU. // - // NOT here: requests.tokens_generated and requests.tokens_per_second. The - // store keeps both and the collector reports both, but nothing in the product - // ever passes `tokens_generated` to `MetricsCollector.record_request`, so both - // are 0 on every node forever. Drawn, that is a flat line at zero saying "this - // node generated no tokens", when the truth is "no code path counts tokens". - // A chart may not say the first when it means the second. Wire the proxy to - // pass the usage block through (and tally the SSE path), and then a tokens - // panel is one entry here and one in metrics.js::PANELS. var SERIES = [ 'gpu.memory_used_mb', 'gpu.memory_total_mb', @@ -72,6 +64,7 @@ 'gpu.temperature_c', 'requests.total', 'requests.errors', + 'requests.tokens_generated', 'requests.latency_ms.p50', 'requests.latency_ms.p95', 'requests.latency_ms.p99', diff --git a/ainode/web/static/js/metrics.js b/ainode/web/static/js/metrics.js index ce26b53a..55691953 100644 --- a/ainode/web/static/js/metrics.js +++ b/ainode/web/static/js/metrics.js @@ -75,6 +75,13 @@ const AINodeMetrics = { unit: ' ms', axis: { zeroFloor: true, minSpan: 20, pad: 0.12 }, }, + { + key: 'tokens', + title: 'Token throughput', + hint: 'output tokens, or input tokens for embeddings, reported by the serving engine', + unit: ' tok/s', + axis: { zeroFloor: true, minSpan: 1, pad: 0.1 }, + }, { key: 'uptime', title: 'Process uptime', @@ -84,13 +91,6 @@ const AINodeMetrics = { }, ], - // There is no tokens per second panel, and that is deliberate. See the note on - // SERIES in metrics-data.js: the counter behind it is never incremented by - // anything in the product, so the panel would be a flat zero reading "this node - // generated no tokens" when the truth is "nothing counts tokens". Uptime took - // the slot because it is a series that is actually measured, and because it is - // what tells a reader whether a gap in the charts above was a restart. - // ====================================================================== // ENTRY POINT (called from app.js's refresh switch, once per poll tick) // ====================================================================== @@ -446,6 +446,12 @@ const AINodeMetrics = { { name: 'p99', color: palette.line3, points: series('requests.latency_ms.p99'), digits: 0 }, ]; } + if (key === 'tokens') { + return [{ + name: 'tokens', color: palette.line2, + points: D.rate(series('requests.tokens_generated'), 1, gap), digits: 1, + }]; + } if (key === 'uptime') { return [{ name: 'uptime', color: palette.line2, @@ -488,6 +494,10 @@ const AINodeMetrics = { return ['No request has been counted in this window. A rate needs two samples of ' + 'the counter, so the first point of the window is always absent.']; } + if (key === 'tokens' && !measured) { + return ['No token rate is available in this window. A rate needs two ' + + 'samples of the counter, so the first point of the window is always absent.']; + } if (key === 'temp' && !measured) { return ['This node reports no temperature. Nothing is drawn rather than a zero, which ' + 'would read as a cold GPU.']; diff --git a/tests/test_embeddings_routing.py b/tests/test_embeddings_routing.py index bb315633..96191b9b 100644 --- a/tests/test_embeddings_routing.py +++ b/tests/test_embeddings_routing.py @@ -37,6 +37,7 @@ handle_v1_embeddings, ) from ainode.embeddings.manager import EmbeddingManager +from ainode.metrics.collector import MetricsCollector EMBED = "Qwen/Qwen3-Embedding-0.6B" CHAT = "fraserprice/DeepSeek-V4-Flash-DSpark" @@ -163,6 +164,7 @@ def _app(cluster, session=None, node_id="spark1", api_port=8000): "cluster_state": cluster, "client_session": session if session is not None else _Session(), "embedding_manager": EmbeddingManager(), + "metrics_collector": MetricsCollector(), } @@ -229,6 +231,7 @@ def test_a_fleet_instance_gets_the_body_verbatim(): # caller's exact float formatting) reaches the engine untouched. assert json.loads(session.bodies[0]) == body assert request.tags["_log_model"] == EMBED + assert app["metrics_collector"].get_request_stats()["tokens_generated"] == 3 def test_the_upstream_status_and_content_type_come_back(): @@ -306,14 +309,16 @@ def test_a_model_no_node_serves_is_answered_in_process(): """Unchanged behaviour, which is the point: the CPU path is the fallback, not the thing that was removed.""" session = _Session() - response, payload, request = _post(_app(_fleet_with_stacked_embedder(), session), - {"model": MINILM, "input": ["a", "b"]}) + app = _app(_fleet_with_stacked_embedder(), session) + response, payload, request = _post(app, {"model": MINILM, "input": ["a", "b"]}) assert response.status == 200 assert session.tried == [] # nothing left this process assert payload["model"] == MINILM assert len(payload["data"]) == 2 assert len(payload["data"][0]["embedding"]) == 8 # the fake's width assert request.tags["_log_model"] == MINILM + assert payload["usage"]["prompt_tokens"] == 2 + assert app["metrics_collector"].get_request_stats()["tokens_generated"] == 2 def test_the_body_is_still_validated_on_the_in_process_path(): diff --git a/tests/test_messages_proxy.py b/tests/test_messages_proxy.py index 7c7d4240..042c8a4a 100644 --- a/tests/test_messages_proxy.py +++ b/tests/test_messages_proxy.py @@ -18,7 +18,7 @@ from aiohttp import web from aiohttp.test_utils import TestClient, TestServer -from ainode.api.server import create_app +from ainode.api.server import SSETokenCounter, create_app, response_body_output_tokens from ainode.core.config import NodeConfig from ainode.discovery.broadcast import NodeStatus from ainode.discovery.cluster import ClusterNode @@ -61,6 +61,8 @@ async def messages(self, request): ("message_start", {"type": "message_start", "message": {"id": "msg_01"}}), ("content_block_delta", {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": ANSWER}}), + ("message_delta", {"type": "message_delta", + "usage": {"output_tokens": 9}}), ("message_stop", {"type": "message_stop"}), ): await asyncio.sleep(0.002) @@ -132,6 +134,7 @@ async def test_messages_is_routed_by_the_model_in_the_body(client, engine_fake): # It reached the node that advertises this model, at the path it was sent to. assert [s["path"] for s in engine_fake.seen] == ["/v1/messages"] assert engine_fake.seen[0]["body"]["model"] == MODEL + assert client.server.app["metrics_collector"].get_request_stats()["tokens_generated"] == 9 @pytest.mark.asyncio @@ -176,6 +179,33 @@ async def test_a_streamed_answer_passes_through_as_sse(client, engine_fake): < text.index("message_stop") assert ANSWER in text assert engine_fake.seen[0]["body"]["stream"] is True + assert client.server.app["metrics_collector"].get_request_stats()["tokens_generated"] == 9 + + +def test_openai_json_usage_uses_completion_tokens(): + body = json.dumps({"usage": {"prompt_tokens": 11, "completion_tokens": 7}}).encode() + assert response_body_output_tokens(body) == 7 + + +def test_responses_json_usage_uses_output_tokens(): + body = json.dumps({"usage": {"input_tokens": 11, "output_tokens": 6}}).encode() + assert response_body_output_tokens(body) == 6 + + +def test_stream_counter_handles_split_events_and_counts_deltas_without_usage(): + counter = SSETokenCounter() + counter.feed(b'data: {"choices":[{"delta":{"content":"a"}}]}\n') + counter.feed(b'\ndata: {"choices":[{"delta":{"content":"b"}}]}\n\n') + counter.feed(b'data: [DONE]\n\n') + assert counter.finish() == 2 + + +def test_stream_reported_usage_wins_over_delta_counting(): + counter = SSETokenCounter() + counter.feed(b'data: {"type":"response.output_text.delta","delta":"many"}\n\n') + counter.feed(b'data: {"type":"response.completed","response":' + b'{"usage":{"output_tokens":12}}}\n\n') + assert counter.finish() == 12 # --------------------------------------------------------------- count_tokens diff --git a/tests/test_metrics_chart.py b/tests/test_metrics_chart.py index 9bedd8de..51b31af5 100644 --- a/tests/test_metrics_chart.py +++ b/tests/test_metrics_chart.py @@ -159,32 +159,14 @@ def test_a_series_that_measured_nothing_is_said_in_words(): assert "n/a" in view # the legend, for the same series -def test_there_is_no_tokens_per_second_panel_while_the_counter_is_dead(): - """Nothing in the product ever increments the token counter. - - ``MetricsCollector.record_request`` takes ``tokens_generated`` and not one of - its eight call sites passes it, so ``requests.tokens_generated`` and - ``requests.tokens_per_second`` are 0 on every node forever. A panel over that - draws a flat line at zero saying "this node generated no tokens", when the - truth is "no code path counts tokens", and a chart may not say the first when - it means the second. When the proxy passes the usage block through, the panel - is one entry in SERIES and one in PANELS, and this test goes away. - """ - import inspect - - from ainode.metrics.collector import MetricsCollector - - source = inspect.getsource(MetricsCollector) - assert "tokens_generated: int = 0" in source, "the signature changed; recheck" - +def test_token_throughput_is_derived_from_the_generated_token_counter(): + """The proxy now fills the counter, so the view may state its interval rate.""" view = METRICS_JS.read_text() - assert "key: 'tokens'" not in view - # And the reason is written down where the next person will look. - assert "nothing counts tokens" in view - # The series is not even asked for: the store answers it, and a payload full - # of zeros invites exactly the panel this test exists to prevent. + assert "key: 'tokens'" in view + assert "D.rate(series('requests.tokens_generated'), 1, gap)" in view asked = _js_ranges()["series"] - assert "requests.tokens_generated" not in asked + assert "requests.tokens_generated" in asked + # The collector's process-lifetime average is not the chart's interval rate. assert "requests.tokens_per_second" not in asked diff --git a/tests/test_routing_caps.py b/tests/test_routing_caps.py index f2aa6237..af7cfedd 100644 --- a/tests/test_routing_caps.py +++ b/tests/test_routing_caps.py @@ -56,7 +56,7 @@ class _Collector: def __init__(self): self.calls: list = [] - def record_request(self, model, ms, error=False): + def record_request(self, model, ms, tokens_generated=0, error=False): self.calls.append((model, error)) diff --git a/tests/test_ui_truth.py b/tests/test_ui_truth.py index b94f0ac9..7e796f23 100644 --- a/tests/test_ui_truth.py +++ b/tests/test_ui_truth.py @@ -61,7 +61,7 @@ class _Collector: def __init__(self): self.calls: list = [] - def record_request(self, model, ms, error=False): + def record_request(self, model, ms, tokens_generated=0, error=False): self.calls.append((model, error))