diff --git a/src/gateway/api/deps.py b/src/gateway/api/deps.py index 65ba1e7366..8b3e56faf5 100644 --- a/src/gateway/api/deps.py +++ b/src/gateway/api/deps.py @@ -15,7 +15,7 @@ from gateway.core.database import DATABASE_ERRORS, create_session, get_db from gateway.core.feature import CoreFeature from gateway.log_config import logger -from gateway.metrics import record_auth_failure +from gateway.metrics import REGISTRY, Counter from gateway.models.entities import APIKey from gateway.models.tenancy import User as TenancyUser from gateway.ports.billing_port import BillingPort @@ -39,6 +39,18 @@ _config: GatewayConfig | None = None _LAST_USED_UPDATE_INTERVAL_SECONDS = 300 +AUTH_FAILURES = Counter( + "gateway_auth_failures", + "Total number of authentication failures", + ["reason"], + registry=REGISTRY, +) + + +def record_auth_failure(reason: str) -> None: + """Record an authentication failure.""" + AUTH_FAILURES.labels(reason=reason).inc() + def _as_utc(value: datetime | None) -> datetime | None: """Return ``value`` as a timezone-aware datetime in UTC. diff --git a/src/gateway/api/routes/_attempts.py b/src/gateway/api/routes/_attempts.py index 452895e954..dd43c2cb44 100644 --- a/src/gateway/api/routes/_attempts.py +++ b/src/gateway/api/routes/_attempts.py @@ -38,10 +38,10 @@ from gateway.api.routes._platform import ( _provider_failure_http_exc, is_provider_billing_error, + record_abandoned_attempt, upstream_exception_shape, ) from gateway.log_config import logger -from gateway.metrics import record_abandoned_attempt from gateway.services.mcp_loop import MaxToolIterationsExceeded from gateway.services.sandbox_backend import SandboxNotReachableError from gateway.services.web_retrieval_backend import WebSearchNotReachableError diff --git a/src/gateway/api/routes/_pipeline.py b/src/gateway/api/routes/_pipeline.py index 54f9cbaf71..44761c29d8 100644 --- a/src/gateway/api/routes/_pipeline.py +++ b/src/gateway/api/routes/_pipeline.py @@ -85,6 +85,7 @@ _resolve_platform_mcp_servers, _resolve_platform_web_search, is_provider_billing_error, + record_abandoned_attempt, run_platform_attempts, upstream_error_message, upstream_exception_chain, @@ -117,7 +118,8 @@ ) from gateway.inflight import track_request from gateway.log_config import logger -from gateway.metrics import record_abandoned_attempt, record_cost, record_inline_cost_settlement, record_tokens +from gateway.metrics import REGISTRY, Histogram +from gateway.metrics import Counter as PrometheusCounter from gateway.model_labeling import relabel_model from gateway.models.entities import APIKey, ModelPricing, UsageLog from gateway.models.guardrails import GuardrailConfig @@ -220,6 +222,46 @@ ResultT = TypeVar("ResultT") ChunkT = TypeVar("ChunkT") +TOKENS = PrometheusCounter( + "gateway_tokens", + "Total number of tokens processed", + ["provider", "model", "type"], + registry=REGISTRY, +) + +REQUEST_COST_DOLLARS = Histogram( + "gateway_request_cost_dollars", + "Request cost in USD", + ["provider", "model"], + registry=REGISTRY, +) + +INLINE_COST_SETTLEMENTS = PrometheusCounter( + "gateway_inline_cost_settlements", + "Inline platform cost settlement outcomes on the hybrid response path", + ["outcome"], + registry=REGISTRY, +) + + +def record_tokens(provider: str, model: str, prompt_tokens: int, completion_tokens: int) -> None: + """Record token usage metrics.""" + if prompt_tokens: + TOKENS.labels(provider=provider, model=model, type="input").inc(prompt_tokens) + if completion_tokens: + TOKENS.labels(provider=provider, model=model, type="output").inc(completion_tokens) + + +def record_cost(provider: str, model: str, cost: float) -> None: + """Record request cost.""" + REQUEST_COST_DOLLARS.labels(provider=provider, model=model).observe(cost) + + +def record_inline_cost_settlement(outcome: str) -> None: + """Record an attached, unattached, or timed-out inline settlement.""" + INLINE_COST_SETTLEMENTS.labels(outcome=outcome).inc() + + # --------------------------------------------------------------------------- # Shared wire-level detail strings. These are client-visible API contract # values; do not edit them without a deprecation plan. diff --git a/src/gateway/api/routes/_platform.py b/src/gateway/api/routes/_platform.py index fe1bbd2b9f..c6fad67d32 100644 --- a/src/gateway/api/routes/_platform.py +++ b/src/gateway/api/routes/_platform.py @@ -34,7 +34,7 @@ cache_write_tokens_of, ) from gateway.log_config import logger -from gateway.metrics import record_abandoned_attempt +from gateway.metrics import REGISTRY, Counter from gateway.models.mcp import McpServerConfig, ResolvedMcpServer from gateway.services.bedrock_gateway_auth import build_bedrock_client_args from gateway.services.mcp_loop import MaxToolIterationsExceeded @@ -49,6 +49,26 @@ T = TypeVar("T") +ABANDONED_ATTEMPTS = Counter( + "gateway_abandoned_attempts", + "Total upstream attempts abandoned before their first chunk (provider fallback / timeout waste)", + ["provider", "model", "reason", "position"], + registry=REGISTRY, +) + + +def record_abandoned_attempt(provider: str, model: str, reason: str, position: int) -> None: + """Record an upstream attempt abandoned before it produced its first chunk. + + ``reason`` is one of ``timeout`` (the first-chunk wait elapsed), + ``build_error`` (opening the upstream stream failed), or ``upstream_error`` + (the upstream raised before yielding a chunk). ``position`` is the attempt's + index in the resolved routing plan; label cardinality stays bounded by the + plan length. + """ + ABANDONED_ATTEMPTS.labels(provider=provider, model=model, reason=reason, position=str(position)).inc() + + # Status codes returned by the platform's usage-report endpoint that the # gateway should NOT retry. Auth, payment-required, not-found, conflict, gone, # and unprocessable are all permanent rejection signals: retrying would just diff --git a/src/gateway/api/routes/auth_oauth.py b/src/gateway/api/routes/auth_oauth.py index e55c7e8515..8b95eb671c 100644 --- a/src/gateway/api/routes/auth_oauth.py +++ b/src/gateway/api/routes/auth_oauth.py @@ -46,7 +46,7 @@ from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession -from gateway.api.deps import IdentityProviderPortDep, get_config, get_db +from gateway.api.deps import IdentityProviderPortDep, get_config, get_db, record_auth_failure from gateway.api.routes._public_auth import throttle_public_auth # The same refusal the password and passkey sign-ins carry, imported rather than @@ -55,7 +55,6 @@ from gateway.api.routes.auth_session import MAINTENANCE_MODE_REFUSAL from gateway.core.config import OAUTH_PROVIDERS, GatewayConfig from gateway.log_config import logger -from gateway.metrics import record_auth_failure from gateway.services.dashboard_session_service import ( apply_session_cookie, create_dashboard_session, diff --git a/src/gateway/api/routes/auth_session.py b/src/gateway/api/routes/auth_session.py index 161e90d19c..7258a86c7c 100644 --- a/src/gateway/api/routes/auth_session.py +++ b/src/gateway/api/routes/auth_session.py @@ -48,10 +48,9 @@ from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession -from gateway.api.deps import get_config, get_db, is_valid_master_key +from gateway.api.deps import get_config, get_db, is_valid_master_key, record_auth_failure from gateway.core.config import GatewayConfig from gateway.log_config import logger -from gateway.metrics import record_auth_failure from gateway.models.tenancy import User as TenancyUser from gateway.rate_limit import RateLimiter from gateway.services.dashboard_session_service import ( diff --git a/src/gateway/api/routes/auth_webauthn.py b/src/gateway/api/routes/auth_webauthn.py index 6d5621f7e1..e03152a22b 100644 --- a/src/gateway/api/routes/auth_webauthn.py +++ b/src/gateway/api/routes/auth_webauthn.py @@ -36,7 +36,7 @@ from sqlalchemy.exc import SQLAlchemyError from sqlalchemy.ext.asyncio import AsyncSession -from gateway.api.deps import CurrentIdentity, get_config, get_db +from gateway.api.deps import CurrentIdentity, get_config, get_db, record_auth_failure from gateway.api.routes._public_auth import throttle_public_auth # The refusal is the same refusal, so it is imported rather than restated: the @@ -45,7 +45,6 @@ from gateway.api.routes.auth_session import MAINTENANCE_MODE_REFUSAL from gateway.core.config import GatewayConfig from gateway.log_config import logger -from gateway.metrics import record_auth_failure from gateway.models.tenancy import ( MAX_WEBAUTHN_CREDENTIAL_NAME, WebAuthnCredentialPublic, diff --git a/src/gateway/metrics.py b/src/gateway/metrics.py index 8181e78b66..882cda1882 100644 --- a/src/gateway/metrics.py +++ b/src/gateway/metrics.py @@ -1,4 +1,7 @@ -"""Prometheus metrics for the gateway.""" +"""Prometheus registry, metric types, and HTTP request instrumentation for the gateway. + +The metric types are re-exported so that code declaring a metric need not depend on ``prometheus_client`` directly. +""" from __future__ import annotations @@ -21,6 +24,8 @@ from starlette.requests import Request from starlette.types import ASGIApp, Message, Receive, Scope, Send +__all__ = ["Counter", "Gauge", "Histogram", "MetricsMiddleware", "REGISTRY", "metrics_endpoint"] + REGISTRY = CollectorRegistry() # process_resident_memory_bytes and friends. The gateway keeps its own registry @@ -49,80 +54,6 @@ registry=REGISTRY, ) -TOKENS = Counter( - "gateway_tokens", - "Total number of tokens processed", - ["provider", "model", "type"], - registry=REGISTRY, -) - -REQUEST_COST_DOLLARS = Histogram( - "gateway_request_cost_dollars", - "Request cost in USD", - ["provider", "model"], - registry=REGISTRY, -) - -ABANDONED_ATTEMPTS = Counter( - "gateway_abandoned_attempts", - "Total upstream attempts abandoned before their first chunk (provider fallback / timeout waste)", - ["provider", "model", "reason", "position"], - registry=REGISTRY, -) - -INLINE_COST_SETTLEMENTS = Counter( - "gateway_inline_cost_settlements", - "Inline platform cost settlement outcomes on the hybrid response path", - ["outcome"], - registry=REGISTRY, -) - -RATE_LIMIT_HITS = Counter( - "gateway_rate_limit_hits", - "Total number of rate limit hits", - registry=REGISTRY, -) - -BUDGET_EXCEEDED = Counter( - "gateway_budget_exceeded", - "Total number of budget exceeded events", - registry=REGISTRY, -) - -AUTH_FAILURES = Counter( - "gateway_auth_failures", - "Total number of authentication failures", - ["reason"], - registry=REGISTRY, -) - -LOG_WRITER_QUEUE_DEPTH = Gauge( - "gateway_usage_log_queue_depth", - "Number of usage log entries waiting to be written", - registry=REGISTRY, -) - -LOG_WRITER_BATCH_SIZE = Histogram( - "gateway_usage_log_batch_size", - "Number of rows per flush batch", - ["writer"], - registry=REGISTRY, -) - -LOG_WRITER_FLUSH_DURATION = Histogram( - "gateway_usage_log_flush_duration_seconds", - "Time spent flushing usage log batches", - ["writer", "result"], - registry=REGISTRY, -) - -LOG_WRITER_ROWS = Counter( - "gateway_usage_log_rows", - "Total usage log rows by outcome", - ["writer", "result"], - registry=REGISTRY, -) - _PROMETHEUS_CONTENT_TYPE = "text/plain; version=0.0.4; charset=utf-8" @@ -202,54 +133,3 @@ async def send_wrapper(message: Message) -> None: REQUEST_DURATION_SECONDS.labels(method=method, endpoint=endpoint, api_version=api_version).observe( duration ) - - -def record_tokens(provider: str, model: str, prompt_tokens: int, completion_tokens: int) -> None: - """Record token usage metrics.""" - if prompt_tokens: - TOKENS.labels(provider=provider, model=model, type="input").inc(prompt_tokens) - if completion_tokens: - TOKENS.labels(provider=provider, model=model, type="output").inc(completion_tokens) - - -def record_cost(provider: str, model: str, cost: float) -> None: - """Record request cost.""" - REQUEST_COST_DOLLARS.labels(provider=provider, model=model).observe(cost) - - -def record_abandoned_attempt(provider: str, model: str, reason: str, position: int) -> None: - """Record an upstream attempt abandoned before it produced its first chunk. - - ``reason`` is one of ``timeout`` (the first-chunk wait elapsed), - ``build_error`` (opening the upstream stream failed), or ``upstream_error`` - (the upstream raised before yielding a chunk). ``position`` is the attempt's - index in the resolved routing plan; label cardinality stays bounded by the - plan length. - """ - ABANDONED_ATTEMPTS.labels(provider=provider, model=model, reason=reason, position=str(position)).inc() - - -def record_inline_cost_settlement(outcome: str) -> None: - """Record an attached, unattached, or timed-out inline settlement.""" - INLINE_COST_SETTLEMENTS.labels(outcome=outcome).inc() - - -def record_rate_limit_hit() -> None: - """Record a rate limit hit.""" - RATE_LIMIT_HITS.inc() - - -def record_budget_exceeded() -> None: - """Record a budget exceeded event.""" - BUDGET_EXCEEDED.inc() - - -def record_auth_failure(reason: str) -> None: - """Record an authentication failure.""" - AUTH_FAILURES.labels(reason=reason).inc() - - -log_writer_queue_depth = LOG_WRITER_QUEUE_DEPTH -log_writer_batch_size = LOG_WRITER_BATCH_SIZE -log_writer_flush_duration = LOG_WRITER_FLUSH_DURATION -log_writer_rows = LOG_WRITER_ROWS diff --git a/src/gateway/rate_limit.py b/src/gateway/rate_limit.py index 775849cc04..0eb55235dc 100644 --- a/src/gateway/rate_limit.py +++ b/src/gateway/rate_limit.py @@ -7,7 +7,13 @@ from fastapi import HTTPException, Request, status -from gateway.metrics import record_rate_limit_hit +from gateway.metrics import REGISTRY, Counter + +RATE_LIMIT_HITS = Counter( + "gateway_rate_limit_hits", + "Total number of rate limit hits", + registry=REGISTRY, +) @dataclass @@ -59,7 +65,7 @@ def check(self, user_id: str) -> RateLimitInfo: if len(timestamps) >= self._rpm: oldest = timestamps[0] retry_after = math.ceil(oldest - cutoff) - record_rate_limit_hit() + RATE_LIMIT_HITS.inc() raise HTTPException( status_code=status.HTTP_429_TOO_MANY_REQUESTS, detail="Rate limit exceeded", diff --git a/src/gateway/services/budget_service.py b/src/gateway/services/budget_service.py index 2d3cf6dac2..d3deb70ef2 100644 --- a/src/gateway/services/budget_service.py +++ b/src/gateway/services/budget_service.py @@ -16,7 +16,7 @@ from gateway.core.metered_pricing import estimate_metered_cost from gateway.log_config import logger -from gateway.metrics import record_budget_exceeded +from gateway.metrics import REGISTRY, Counter from gateway.models.entities import MAX_COUNT_LIMIT, Budget, BudgetResetLog, ModelPricing, User from gateway.models.money import to_usd from gateway.repositories.users_repository import get_active_user @@ -35,6 +35,12 @@ from gateway.services.scoped_budget_service import settle as settle_scoped from gateway.types.budget_state import BudgetState +BUDGET_EXCEEDED = Counter( + "gateway_budget_exceeded", + "Total number of budget exceeded events", + registry=REGISTRY, +) + # Every counter in this module is a ``NUMERIC(18, 6)`` column, so the constants # the SQL is built from are ``Decimal`` too. A bare ``0.0`` in a CASE arm or a # ``.values()`` would make PostgreSQL resolve the whole expression as double @@ -613,7 +619,7 @@ async def reserve_budget( db, scoped, usd, tokens=held_tokens, requests=held_requests, new_request=new_request ) if refused is not None: - record_budget_exceeded() + BUDGET_EXCEEDED.inc() axis = await blocked_axis( db, refused, @@ -715,7 +721,7 @@ def guards_for(committed: Any, held_amount: Any, cap: Any) -> list[Any]: await db.commit() if not getattr(result, "rowcount", 0): - record_budget_exceeded() + BUDGET_EXCEEDED.inc() # The scoped ceilings admitted this request and are already holding it, so # give every axis back before rejecting. Without this the holds would leak # on every per-user refusal and permanently shrink each ceiling. @@ -1026,7 +1032,7 @@ async def increase_reservation( db, handle.scoped_budgets, additional, tokens=grown_tokens, requests=0, new_request=False ) if refused is not None: - record_budget_exceeded() + BUDGET_EXCEEDED.inc() axis = await blocked_axis( db, refused, diff --git a/src/gateway/services/log_writer.py b/src/gateway/services/log_writer.py index 48b813589c..baf0db1b42 100644 --- a/src/gateway/services/log_writer.py +++ b/src/gateway/services/log_writer.py @@ -8,14 +8,36 @@ from gateway.core.database import DATABASE_ERRORS, create_log_session from gateway.log_config import logger -from gateway.metrics import ( - log_writer_batch_size, - log_writer_flush_duration, - log_writer_queue_depth, - log_writer_rows, -) +from gateway.metrics import REGISTRY, Counter, Gauge, Histogram from gateway.models.entities import UsageLog +QUEUE_DEPTH = Gauge( + "gateway_usage_log_queue_depth", + "Number of usage log entries waiting to be written", + registry=REGISTRY, +) + +BATCH_SIZE = Histogram( + "gateway_usage_log_batch_size", + "Number of rows per flush batch", + ["writer"], + registry=REGISTRY, +) + +FLUSH_DURATION = Histogram( + "gateway_usage_log_flush_duration_seconds", + "Time spent flushing usage log batches", + ["writer", "result"], + registry=REGISTRY, +) + +ROWS = Counter( + "gateway_usage_log_rows", + "Total usage log rows by outcome", + ["writer", "result"], + registry=REGISTRY, +) + class LogWriter(Protocol): async def put(self, log: UsageLog) -> None: ... @@ -37,11 +59,11 @@ async def put(self, log: UsageLog) -> None: # and, under the batch writer, lag the budget gate. db.add(log) await db.commit() - log_writer_rows.labels(writer="single", result="written").inc() + ROWS.labels(writer="single", result="written").inc() except DATABASE_ERRORS as e: # pragma: no cover - defensive logging await db.rollback() logger.error("SingleLogWriter failed: %s", e) - log_writer_rows.labels(writer="single", result="dropped").inc() + ROWS.labels(writer="single", result="dropped").inc() async def start(self) -> None: pass @@ -61,7 +83,7 @@ def __init__(self, max_batch: int = 100, flush_interval: float = 1.0) -> None: async def put(self, log: UsageLog) -> None: await self._queue.put(log) - log_writer_queue_depth.set(self._queue.qsize()) + QUEUE_DEPTH.set(self._queue.qsize()) async def start(self) -> None: self._task = asyncio.create_task(self._run()) @@ -106,7 +128,7 @@ async def _collect_batch(self) -> list[UsageLog]: async def _flush(self, batch: list[UsageLog]) -> None: start = time.monotonic() - log_writer_batch_size.labels(writer="batch").observe(len(batch)) + BATCH_SIZE.labels(writer="batch").observe(len(batch)) try: async with create_log_session() as db: # Spend is reconciled inline via the budget reservation path, not @@ -114,12 +136,12 @@ async def _flush(self, batch: list[UsageLog]) -> None: for log in batch: db.add(log) await db.commit() - log_writer_rows.labels(writer="batch", result="written").inc(len(batch)) - log_writer_flush_duration.labels(writer="batch", result="ok").observe(time.monotonic() - start) + ROWS.labels(writer="batch", result="written").inc(len(batch)) + FLUSH_DURATION.labels(writer="batch", result="ok").observe(time.monotonic() - start) except DATABASE_ERRORS as e: # pragma: no cover - defensive logging logger.error("BatchLogWriter flush failed, dropping %d rows: %s", len(batch), e) - log_writer_rows.labels(writer="batch", result="dropped").inc(len(batch)) - log_writer_flush_duration.labels(writer="batch", result="error").observe(time.monotonic() - start) + ROWS.labels(writer="batch", result="dropped").inc(len(batch)) + FLUSH_DURATION.labels(writer="batch", result="error").observe(time.monotonic() - start) async def _flush_all(self) -> None: batch: list[UsageLog] = [] diff --git a/tests/unit/test_gateway_metrics.py b/tests/unit/test_gateway_metrics.py index ec3760ce55..d6529c8f2b 100644 --- a/tests/unit/test_gateway_metrics.py +++ b/tests/unit/test_gateway_metrics.py @@ -13,13 +13,6 @@ MetricsMiddleware, _endpoint_label, metrics_endpoint, - record_abandoned_attempt, - record_auth_failure, - record_budget_exceeded, - record_cost, - record_inline_cost_settlement, - record_rate_limit_hit, - record_tokens, ) @@ -28,97 +21,44 @@ def _sample(name: str, labels: dict[str, str] | None = None) -> float: return REGISTRY.get_sample_value(name, labels or {}) or 0.0 -def test_record_tokens_increments_counters() -> None: - input_labels = {"provider": "test-prov", "model": "test-model", "type": "input"} - output_labels = {"provider": "test-prov", "model": "test-model", "type": "output"} - before_in = _sample("gateway_tokens_total", input_labels) - before_out = _sample("gateway_tokens_total", output_labels) - - record_tokens("test-prov", "test-model", 100, 50) - - assert _sample("gateway_tokens_total", input_labels) - before_in == 100.0 - assert _sample("gateway_tokens_total", output_labels) - before_out == 50.0 - - -def test_record_tokens_skips_zero_values() -> None: - labels_in = {"provider": "zero-prov", "model": "zero-model", "type": "input"} - labels_out = {"provider": "zero-prov", "model": "zero-model", "type": "output"} - before_in = _sample("gateway_tokens_total", labels_in) - before_out = _sample("gateway_tokens_total", labels_out) - - record_tokens("zero-prov", "zero-model", 0, 0) - - assert _sample("gateway_tokens_total", labels_in) == before_in - assert _sample("gateway_tokens_total", labels_out) == before_out - - -def test_record_cost_observes_histogram() -> None: - labels = {"provider": "cost-prov", "model": "cost-model"} - before_count = _sample("gateway_request_cost_dollars_count", labels) - - record_cost("cost-prov", "cost-model", 1.23) - - assert _sample("gateway_request_cost_dollars_count", labels) - before_count == 1.0 - assert _sample("gateway_request_cost_dollars_sum", labels) >= 1.23 - - -@pytest.mark.parametrize("outcome", ["attached", "unattached", "timeout"]) -def test_record_inline_cost_settlement_increments_counter(outcome: str) -> None: - labels = {"outcome": outcome} - before = _sample("gateway_inline_cost_settlements_total", labels) - - record_inline_cost_settlement(outcome) - - assert _sample("gateway_inline_cost_settlements_total", labels) - before == 1.0 - - -def test_record_rate_limit_hit_increments_counter() -> None: - before = _sample("gateway_rate_limit_hits_total") - - record_rate_limit_hit() - - assert _sample("gateway_rate_limit_hits_total") - before == 1.0 - - -def test_record_budget_exceeded_increments_counter() -> None: - before = _sample("gateway_budget_exceeded_total") - - record_budget_exceeded() - - assert _sample("gateway_budget_exceeded_total") - before == 1.0 - - -def test_record_auth_failure_increments_counter() -> None: - labels = {"reason": "unit-test-reason"} - before = _sample("gateway_auth_failures_total", labels) - - record_auth_failure("unit-test-reason") - - assert _sample("gateway_auth_failures_total", labels) - before == 1.0 - - -def test_record_abandoned_attempt_increments_counter() -> None: - labels = {"provider": "ab-prov", "model": "ab-model", "reason": "timeout", "position": "0"} - before = _sample("gateway_abandoned_attempts_total", labels) - - record_abandoned_attempt("ab-prov", "ab-model", "timeout", 0) - - assert _sample("gateway_abandoned_attempts_total", labels) - before == 1.0 - - -def test_record_abandoned_attempt_labels_by_reason_and_position() -> None: - """Each (reason, position) pair is its own series so operators can spot which - plan entry and failure phase dominates the fallback waste.""" - build_labels = {"provider": "ab-prov2", "model": "ab-model2", "reason": "build_error", "position": "1"} - upstream_labels = {"provider": "ab-prov2", "model": "ab-model2", "reason": "upstream_error", "position": "2"} - before_build = _sample("gateway_abandoned_attempts_total", build_labels) - before_upstream = _sample("gateway_abandoned_attempts_total", upstream_labels) +# Every family the gateway registers, as (name, type, label names). A metric +# may move to the module that increments it, but the scrape is an external +# contract: dashboards, recording rules and alerts outside this repository +# key on these names and labels. +_EXPOSED_FAMILIES: set[tuple[str, str, tuple[str, ...]]] = { + ("gateway_abandoned_attempts", "counter", ("provider", "model", "reason", "position")), + ("gateway_active_requests", "gauge", ()), + ("gateway_auth_failures", "counter", ("reason",)), + ("gateway_budget_exceeded", "counter", ()), + ("gateway_inline_cost_settlements", "counter", ("outcome",)), + ("gateway_rate_limit_hits", "counter", ()), + ("gateway_request_cost_dollars", "histogram", ("provider", "model")), + ("gateway_request_duration_seconds", "histogram", ("method", "endpoint", "api_version")), + ("gateway_requests", "counter", ("method", "endpoint", "api_version", "status")), + ("gateway_tokens", "counter", ("provider", "model", "type")), + ("gateway_usage_log_batch_size", "histogram", ("writer",)), + ("gateway_usage_log_flush_duration_seconds", "histogram", ("writer", "result")), + ("gateway_usage_log_queue_depth", "gauge", ()), + ("gateway_usage_log_rows", "counter", ("writer", "result")), +} + + +def test_scrape_exposes_the_pinned_families() -> None: + """The set of gateway metric families, with their types and label names, is fixed. + + A labeled family with no series yet shows only its HELP and TYPE lines in a + scrape, so the label names are read off the collectors rather than the text. + """ + import gateway.main # noqa: F401 # imports every module that registers a metric - record_abandoned_attempt("ab-prov2", "ab-model2", "build_error", 1) - record_abandoned_attempt("ab-prov2", "ab-model2", "upstream_error", 2) + families: set[tuple[str, str, tuple[str, ...]]] = set() + for collector in REGISTRY._collector_to_names: + describe = getattr(collector, "describe", collector.collect) + for metric in describe(): + if metric.name.startswith("gateway_"): + families.add((metric.name, metric.type, tuple(getattr(collector, "_labelnames", ())))) - assert _sample("gateway_abandoned_attempts_total", build_labels) - before_build == 1.0 - assert _sample("gateway_abandoned_attempts_total", upstream_labels) - before_upstream == 1.0 + assert families == _EXPOSED_FAMILIES def test_config_enable_metrics_defaults_to_false() -> None: @@ -291,22 +231,6 @@ def test_active_requests_returns_to_zero() -> None: assert _sample("gateway_active_requests") == before -def test_rate_limiter_records_metric_on_429() -> None: - """RateLimiter.check() records a metric before raising 429.""" - from gateway.rate_limit import RateLimiter - - limiter = RateLimiter(rpm=1) - limiter.check("metric-rl-user") - - before = _sample("gateway_rate_limit_hits_total") - - with pytest.raises(HTTPException) as exc_info: - limiter.check("metric-rl-user") - - assert exc_info.value.status_code == 429 - assert _sample("gateway_rate_limit_hits_total") - before == 1.0 - - @pytest.mark.skipif(not os.path.exists("/proc/stat"), reason="ProcessCollector needs /proc") def test_metrics_expose_process_memory() -> None: """Without this the only way to read the resident set is a shell in the container.""" diff --git a/tests/unit/test_pipeline_metrics.py b/tests/unit/test_pipeline_metrics.py new file mode 100644 index 0000000000..bb23acd5d2 --- /dev/null +++ b/tests/unit/test_pipeline_metrics.py @@ -0,0 +1,55 @@ +"""Unit tests for the token, cost and inline settlement metrics the pipeline records.""" + +import pytest + +from gateway.api.routes._pipeline import record_cost, record_inline_cost_settlement, record_tokens +from gateway.metrics import REGISTRY + + +def _sample(name: str, labels: dict[str, str] | None = None) -> float: + """Read a metric sample value from the registry, returning 0.0 if not found.""" + return REGISTRY.get_sample_value(name, labels or {}) or 0.0 + + +def test_record_tokens_increments_counters() -> None: + input_labels = {"provider": "test-prov", "model": "test-model", "type": "input"} + output_labels = {"provider": "test-prov", "model": "test-model", "type": "output"} + before_in = _sample("gateway_tokens_total", input_labels) + before_out = _sample("gateway_tokens_total", output_labels) + + record_tokens("test-prov", "test-model", 100, 50) + + assert _sample("gateway_tokens_total", input_labels) - before_in == 100.0 + assert _sample("gateway_tokens_total", output_labels) - before_out == 50.0 + + +def test_record_tokens_skips_zero_values() -> None: + labels_in = {"provider": "zero-prov", "model": "zero-model", "type": "input"} + labels_out = {"provider": "zero-prov", "model": "zero-model", "type": "output"} + before_in = _sample("gateway_tokens_total", labels_in) + before_out = _sample("gateway_tokens_total", labels_out) + + record_tokens("zero-prov", "zero-model", 0, 0) + + assert _sample("gateway_tokens_total", labels_in) == before_in + assert _sample("gateway_tokens_total", labels_out) == before_out + + +def test_record_cost_observes_histogram() -> None: + labels = {"provider": "cost-prov", "model": "cost-model"} + before_count = _sample("gateway_request_cost_dollars_count", labels) + + record_cost("cost-prov", "cost-model", 1.23) + + assert _sample("gateway_request_cost_dollars_count", labels) - before_count == 1.0 + assert _sample("gateway_request_cost_dollars_sum", labels) >= 1.23 + + +@pytest.mark.parametrize("outcome", ["attached", "unattached", "timeout"]) +def test_record_inline_cost_settlement_increments_counter(outcome: str) -> None: + labels = {"outcome": outcome} + before = _sample("gateway_inline_cost_settlements_total", labels) + + record_inline_cost_settlement(outcome) + + assert _sample("gateway_inline_cost_settlements_total", labels) - before == 1.0 diff --git a/tests/unit/test_rate_limiter_core.py b/tests/unit/test_rate_limiter_core.py index ffb5a02fb1..7570a76e0c 100644 --- a/tests/unit/test_rate_limiter_core.py +++ b/tests/unit/test_rate_limiter_core.py @@ -8,6 +8,7 @@ from pydantic import ValidationError from gateway.core.config import GatewayConfig +from gateway.metrics import REGISTRY from gateway.rate_limit import RateLimiter, RateLimitInfo @@ -140,3 +141,17 @@ def test_config_accepts_positive_rate_limit() -> None: def test_config_accepts_none_rate_limit() -> None: config = GatewayConfig(rate_limit_rpm=None) assert config.rate_limit_rpm is None + + +def test_check_records_metric_on_429() -> None: + """RateLimiter.check() records a metric before raising 429.""" + limiter = RateLimiter(rpm=1) + limiter.check("metric-rl-user") + + before = REGISTRY.get_sample_value("gateway_rate_limit_hits_total") or 0.0 + + with pytest.raises(HTTPException) as exc_info: + limiter.check("metric-rl-user") + + assert exc_info.value.status_code == 429 + assert (REGISTRY.get_sample_value("gateway_rate_limit_hits_total") or 0.0) - before == 1.0 diff --git a/tests/unit/test_run_platform_attempts.py b/tests/unit/test_run_platform_attempts.py index 780de99cef..e4f8cae7d1 100644 --- a/tests/unit/test_run_platform_attempts.py +++ b/tests/unit/test_run_platform_attempts.py @@ -18,7 +18,13 @@ from fastapi import HTTPException from gateway.api.routes import _platform -from gateway.api.routes._platform import ResolvedAttempt, ResolvedRoute, default_attempt_kwargs, run_platform_attempts +from gateway.api.routes._platform import ( + ResolvedAttempt, + ResolvedRoute, + default_attempt_kwargs, + record_abandoned_attempt, + run_platform_attempts, +) from gateway.core.config import GatewayConfig from gateway.metrics import REGISTRY from gateway.services.mcp_loop import MaxToolIterationsExceeded @@ -36,6 +42,27 @@ def _abandoned_sample(provider: str, model: str, reason: str, position: int) -> ) +def test_record_abandoned_attempt_increments_counter() -> None: + before = _abandoned_sample("ab-prov", "ab-model", "timeout", 0) + + record_abandoned_attempt("ab-prov", "ab-model", "timeout", 0) + + assert _abandoned_sample("ab-prov", "ab-model", "timeout", 0) - before == 1.0 + + +def test_record_abandoned_attempt_labels_by_reason_and_position() -> None: + """Each (reason, position) pair is its own series so operators can spot which + plan entry and failure phase dominates the fallback waste.""" + before_build = _abandoned_sample("ab-prov2", "ab-model2", "build_error", 1) + before_upstream = _abandoned_sample("ab-prov2", "ab-model2", "upstream_error", 2) + + record_abandoned_attempt("ab-prov2", "ab-model2", "build_error", 1) + record_abandoned_attempt("ab-prov2", "ab-model2", "upstream_error", 2) + + assert _abandoned_sample("ab-prov2", "ab-model2", "build_error", 1) - before_build == 1.0 + assert _abandoned_sample("ab-prov2", "ab-model2", "upstream_error", 2) - before_upstream == 1.0 + + def _single_attempt(provider: str, model: str) -> ResolvedAttempt: return ResolvedAttempt( attempt_id="a0", position=0, provider=provider, model=model, api_key="k", managed=False