Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
145 changes: 143 additions & 2 deletions ainode/api/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down Expand Up @@ -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={
Expand All @@ -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:
Expand All @@ -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},
Expand Down
56 changes: 42 additions & 14 deletions ainode/embeddings/api_routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,9 @@
from __future__ import annotations

import asyncio
import json
import logging
import time
from typing import List, Optional

import aiohttp
Expand Down Expand Up @@ -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.

Expand Down Expand Up @@ -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.
Expand All @@ -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,
Expand All @@ -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:
Expand Down
9 changes: 1 addition & 8 deletions ainode/web/static/js/metrics-data.js
Original file line number Diff line number Diff line change
Expand Up @@ -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',
Expand All @@ -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',
Expand Down
24 changes: 17 additions & 7 deletions ainode/web/static/js/metrics.js
Original file line number Diff line number Diff line change
Expand Up @@ -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',
Expand All @@ -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)
// ======================================================================
Expand Down Expand Up @@ -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,
Expand Down Expand Up @@ -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.'];
Expand Down
Loading
Loading