From d8d8b54284123a65f1a28a8b09b0894c874c1054 Mon Sep 17 00:00:00 2001 From: BoKeum Date: Mon, 10 Aug 2026 19:21:41 +0900 Subject: [PATCH 1/8] feat(BA-7317): render prometheus query preset templates with Jinja Replace the str.format-based query_template engine with a sandboxed Jinja environment. Templates now use {{ labels }}, {{ window }}, {{ group_by }}; the legacy {placeholder} syntax is rejected at the API boundary with a guidance message, and a data migration rewrites all stored presets (including seeded defaults) to the Jinja form. Validation is parser-based: an AST whitelist permits only literal text and variable substitution, and StrictUndefined rejects unknown variables. Co-Authored-By: Claude Fable 5 --- .../example-prometheus-query-presets.json | 44 +++--- .../backend/common/data/idle_checker/types.py | 4 +- .../v2/prometheus_query_preset/request.py | 12 +- .../v2/prometheus_query_preset/validators.py | 72 ++++----- .../prometheus_query_preset/types/inputs.py | 4 +- .../clients/prometheus/fixed_query_builder.py | 20 +-- .../manager/clients/prometheus/preset.py | 24 ++- ...metheus_query_preset_templates_to_jinja.py | 97 ++++++++++++ .../prometheus/test_client_integration.py | 2 +- .../clients/prometheus/test_sd_relabel.py | 2 +- .../test_prometheus_query_preset_preview.py | 6 +- .../prometheus_query_preset/test_request.py | 17 ++- .../manager/clients/prometheus/test_client.py | 12 +- .../manager/clients/prometheus/test_preset.py | 139 ++++++------------ .../metric/test_session_utilization.py | 4 +- .../test_prometheus_query_preset_options.py | 2 +- ...test_prometheus_query_preset_repository.py | 8 +- .../services/idle_checker/test_service.py | 2 +- .../test_prometheus_query_preset_service.py | 6 +- .../test_container_metric.py | 34 ++--- 20 files changed, 299 insertions(+), 212 deletions(-) create mode 100644 src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py diff --git a/fixtures/manager/example-prometheus-query-presets.json b/fixtures/manager/example-prometheus-query-presets.json index a1360631856..392726e772f 100644 --- a/fixtures/manager/example-prometheus-query-presets.json +++ b/fixtures/manager/example-prometheus-query-presets.json @@ -7,7 +7,7 @@ "rank": 100, "category_name": "container", "metric_name": "backendai_container_utilization", - "query_template": "sum by ({group_by})(backendai_container_utilization{{{labels}}})", + "query_template": "sum by ({{ group_by }})(backendai_container_utilization{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -36,7 +36,7 @@ "rank": 110, "category_name": "container", "metric_name": "backendai_container_utilization", - "query_template": "avg by ({group_by})(backendai_container_utilization{{{labels}}})", + "query_template": "avg by ({{ group_by }})(backendai_container_utilization{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -65,7 +65,7 @@ "rank": 120, "category_name": "container", "metric_name": "backendai_container_utilization", - "query_template": "max by ({group_by})(backendai_container_utilization{{{labels}}})", + "query_template": "max by ({{ group_by }})(backendai_container_utilization{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -94,7 +94,7 @@ "rank": 130, "category_name": "container", "metric_name": "backendai_container_utilization", - "query_template": "min by ({group_by})(backendai_container_utilization{{{labels}}})", + "query_template": "min by ({{ group_by }})(backendai_container_utilization{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -123,7 +123,7 @@ "rank": 200, "category_name": "container", "metric_name": "backendai_container_utilization", - "query_template": "sum by ({group_by})(rate(backendai_container_utilization{{{labels}}}[{window}]))", + "query_template": "sum by ({{ group_by }})(rate(backendai_container_utilization{ {{ labels }}}[{{ window }}]))", "time_window": "5m", "options": { "filter_labels": [ @@ -152,7 +152,7 @@ "rank": 210, "category_name": "container", "metric_name": "backendai_container_utilization", - "query_template": "avg by ({group_by})(rate(backendai_container_utilization{{{labels}}}[{window}]))", + "query_template": "avg by ({{ group_by }})(rate(backendai_container_utilization{ {{ labels }}}[{{ window }}]))", "time_window": "5m", "options": { "filter_labels": [ @@ -181,7 +181,7 @@ "rank": 220, "category_name": "container", "metric_name": "backendai_container_utilization", - "query_template": "max by ({group_by})(rate(backendai_container_utilization{{{labels}}}[{window}]))", + "query_template": "max by ({{ group_by }})(rate(backendai_container_utilization{ {{ labels }}}[{{ window }}]))", "time_window": "5m", "options": { "filter_labels": [ @@ -210,7 +210,7 @@ "rank": 230, "category_name": "container", "metric_name": "backendai_container_utilization", - "query_template": "min by ({group_by})(rate(backendai_container_utilization{{{labels}}}[{window}]))", + "query_template": "min by ({{ group_by }})(rate(backendai_container_utilization{ {{ labels }}}[{{ window }}]))", "time_window": "5m", "options": { "filter_labels": [ @@ -239,7 +239,7 @@ "rank": 300, "category_name": "vllm-inference", "metric_name": "vllm:num_requests_running", - "query_template": "sum by ({group_by})(vllm:num_requests_running{{{labels}}})", + "query_template": "sum by ({{ group_by }})(vllm:num_requests_running{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -257,7 +257,7 @@ "rank": 310, "category_name": "vllm-inference", "metric_name": "vllm:num_requests_running", - "query_template": "avg by ({group_by})(vllm:num_requests_running{{{labels}}})", + "query_template": "avg by ({{ group_by }})(vllm:num_requests_running{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -275,7 +275,7 @@ "rank": 320, "category_name": "vllm-inference", "metric_name": "vllm:num_requests_running", - "query_template": "max by ({group_by})(vllm:num_requests_running{{{labels}}})", + "query_template": "max by ({{ group_by }})(vllm:num_requests_running{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -293,7 +293,7 @@ "rank": 400, "category_name": "vllm-inference", "metric_name": "vllm:num_requests_waiting", - "query_template": "sum by ({group_by})(vllm:num_requests_waiting{{{labels}}})", + "query_template": "sum by ({{ group_by }})(vllm:num_requests_waiting{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -311,7 +311,7 @@ "rank": 410, "category_name": "vllm-inference", "metric_name": "vllm:num_requests_waiting", - "query_template": "avg by ({group_by})(vllm:num_requests_waiting{{{labels}}})", + "query_template": "avg by ({{ group_by }})(vllm:num_requests_waiting{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -329,7 +329,7 @@ "rank": 420, "category_name": "vllm-inference", "metric_name": "vllm:num_requests_waiting", - "query_template": "max by ({group_by})(vllm:num_requests_waiting{{{labels}}})", + "query_template": "max by ({{ group_by }})(vllm:num_requests_waiting{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -347,7 +347,7 @@ "rank": 500, "category_name": "vllm-inference", "metric_name": "vllm:gpu_cache_usage_perc", - "query_template": "avg by ({group_by})(vllm:gpu_cache_usage_perc{{{labels}}})", + "query_template": "avg by ({{ group_by }})(vllm:gpu_cache_usage_perc{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -365,7 +365,7 @@ "rank": 510, "category_name": "vllm-inference", "metric_name": "vllm:gpu_cache_usage_perc", - "query_template": "max by ({group_by})(vllm:gpu_cache_usage_perc{{{labels}}})", + "query_template": "max by ({{ group_by }})(vllm:gpu_cache_usage_perc{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -383,7 +383,7 @@ "rank": 520, "category_name": "vllm-inference", "metric_name": "vllm:gpu_cache_usage_perc", - "query_template": "min by ({group_by})(vllm:gpu_cache_usage_perc{{{labels}}})", + "query_template": "min by ({{ group_by }})(vllm:gpu_cache_usage_perc{ {{ labels }}})", "time_window": null, "options": { "filter_labels": [ @@ -401,7 +401,7 @@ "rank": 600, "category_name": "vllm-inference", "metric_name": "vllm:request_success_total", - "query_template": "sum by ({group_by})(rate(vllm:request_success_total{{{labels}}}[{window}]))", + "query_template": "sum by ({{ group_by }})(rate(vllm:request_success_total{ {{ labels }}}[{{ window }}]))", "time_window": "5m", "options": { "filter_labels": [ @@ -419,7 +419,7 @@ "rank": 610, "category_name": "vllm-inference", "metric_name": "vllm:request_success_total", - "query_template": "avg by ({group_by})(rate(vllm:request_success_total{{{labels}}}[{window}]))", + "query_template": "avg by ({{ group_by }})(rate(vllm:request_success_total{ {{ labels }}}[{{ window }}]))", "time_window": "5m", "options": { "filter_labels": [ @@ -437,7 +437,7 @@ "rank": 620, "category_name": "vllm-inference", "metric_name": "vllm:request_success_total", - "query_template": "max by ({group_by})(rate(vllm:request_success_total{{{labels}}}[{window}]))", + "query_template": "max by ({{ group_by }})(rate(vllm:request_success_total{ {{ labels }}}[{{ window }}]))", "time_window": "5m", "options": { "filter_labels": [ @@ -455,7 +455,7 @@ "rank": 700, "category_name": "vllm-inference", "metric_name": "vllm:e2e_request_latency_seconds", - "query_template": "avg by ({group_by})(rate(vllm:e2e_request_latency_seconds_sum{{{labels}}}[{window}]) / rate(vllm:e2e_request_latency_seconds_count{{{labels}}}[{window}]))", + "query_template": "avg by ({{ group_by }})(rate(vllm:e2e_request_latency_seconds_sum{ {{ labels }}}[{{ window }}]) / rate(vllm:e2e_request_latency_seconds_count{ {{ labels }}}[{{ window }}]))", "time_window": "5m", "options": { "filter_labels": [ @@ -473,7 +473,7 @@ "rank": 710, "category_name": "vllm-inference", "metric_name": "vllm:e2e_request_latency_seconds", - "query_template": "max by ({group_by})(rate(vllm:e2e_request_latency_seconds_sum{{{labels}}}[{window}]) / rate(vllm:e2e_request_latency_seconds_count{{{labels}}}[{window}]))", + "query_template": "max by ({{ group_by }})(rate(vllm:e2e_request_latency_seconds_sum{ {{ labels }}}[{{ window }}]) / rate(vllm:e2e_request_latency_seconds_count{ {{ labels }}}[{{ window }}]))", "time_window": "5m", "options": { "filter_labels": [ diff --git a/src/ai/backend/common/data/idle_checker/types.py b/src/ai/backend/common/data/idle_checker/types.py index 0ba2408593f..0c38560c005 100644 --- a/src/ai/backend/common/data/idle_checker/types.py +++ b/src/ai/backend/common/data/idle_checker/types.py @@ -48,12 +48,12 @@ class UtilizationThresholdEntry(BackendAISchema): ) filter_labels: list[MetricLabel] = Field( default_factory=list, - description="Label filters injected into the preset's {labels} placeholder.", + description="Label filters injected into the preset's {{ labels }} placeholder.", ) group_labels: list[str] = Field( default_factory=lambda: [SESSION_ID_LABEL], description=( - "Labels injected into the preset's {group_by} placeholder. " + "Labels injected into the preset's {{ group_by }} placeholder. " "Must include 'session_id' for per-session values to be mapped. " "When 'session_id' is grouped, the checker adds a session_id filter that " "limits the query to the sessions being evaluated; a user-provided " diff --git a/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/request.py b/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/request.py index 10fe0c8a6d2..536bce67b0d 100644 --- a/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/request.py +++ b/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/request.py @@ -65,7 +65,11 @@ class CreateQueryDefinitionInput(BaseRequestModel): rank: int = Field(default=0, ge=0, description="Sort rank (lower = higher priority)") category_id: UUID | None = Field(default=None, description="Category ID") metric_name: str = Field(description="Prometheus metric name") - query_template: str = Field(description="PromQL template with placeholders") + query_template: str = Field( + description=( + "PromQL template with Jinja placeholders ({{ labels }}, {{ window }}, {{ group_by }})" + ) + ) time_window: str | None = Field( default=None, pattern=PROMETHEUS_DURATION_PATTERN, @@ -122,7 +126,11 @@ class ModifyQueryDefinitionInput(BaseRequestModel): ) metric_name: str | None = Field(default=None, description="Updated Prometheus metric name") query_template: str | None = Field( - default=None, description="Updated PromQL template with placeholders" + default=None, + description=( + "Updated PromQL template with Jinja placeholders" + " ({{ labels }}, {{ window }}, {{ group_by }})" + ), ) time_window: str | Sentinel | None = Field( default=SENTINEL, diff --git a/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/validators.py b/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/validators.py index 3d75d29931b..fa402e35488 100644 --- a/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/validators.py +++ b/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/validators.py @@ -1,66 +1,70 @@ -"""Request DTO validators for prometheus_query_preset templates.""" +"""Validators for prometheus_query_preset templates.""" from __future__ import annotations import re +from collections.abc import Iterator + +from jinja2 import StrictUndefined, TemplateError, TemplateSyntaxError, nodes +from jinja2.sandbox import ImmutableSandboxedEnvironment from ai.backend.common.exception import InvalidMetricPresetTemplate __all__ = ( "PLACEHOLDER_NAMES", - "escape_non_placeholders", + "PROMQL_TEMPLATE_ENV", "validate_query_template", ) PLACEHOLDER_NAMES = frozenset({"labels", "window", "group_by"}) -_BRACE_BLOCK_RE = re.compile(r"\{([^{}]*)\}") -_UNSUPPORTED_TEMPLATE_VAR_RE = re.compile(r"\$\{[^}]+\}|\$[A-Za-z_][A-Za-z0-9_]*") +# Sandboxed: templates are user input from the admin API. +PROMQL_TEMPLATE_ENV = ImmutableSandboxedEnvironment(undefined=StrictUndefined) +# Literal text and `{{ placeholder }}` substitution only. +_ALLOWED_NODE_TYPES = (nodes.Template, nodes.Output, nodes.TemplateData, nodes.Name) -def escape_non_placeholders(template: str) -> str: - """Normalize each ``{X}`` so ``str.format`` produces a single PromQL ``{value}`` - regardless of how many braces the user wrote. - """ +_UNSUPPORTED_TEMPLATE_VAR_RE = re.compile(r"\$\{[^}]+\}|\$[A-Za-z_][A-Za-z0-9_]*") +# Bare `{placeholder}` or any `{{{`: the pre-Jinja str.format syntax. +_LEGACY_TEMPLATE_RE = re.compile(r"(? str: - name = match.group(1) - start, end = match.span() - text = match.string - already_wrapped = ( - start > 0 and text[start - 1] == "{" and end < len(text) and text[end] == "}" - ) - inside_escaped_braces = ( - text.rfind("{{", 0, start) > text.rfind("}}", 0, start) and text.find("}}", end) != -1 - ) - if name not in PLACEHOLDER_NAMES: - return match.group(0) if already_wrapped else "{{" + name + "}}" - if name != "labels": - return match.group(0) - return ( - match.group(0) - if already_wrapped or inside_escaped_braces - else "{{" + match.group(0) + "}}" - ) - return _BRACE_BLOCK_RE.sub(repl, template) +def _walk(node: nodes.Node) -> Iterator[nodes.Node]: + yield node + for child in node.iter_child_nodes(): + yield from _walk(child) -def validate_query_template(template: str) -> str: - """Reject empty templates, foreign variables, or malformed braces.""" +def validate_query_template(template: str) -> None: + """Validate a Jinja PromQL template; raises ``InvalidMetricPresetTemplate``.""" if not template.strip(): raise InvalidMetricPresetTemplate("Template must not be empty.") unsupported_vars = _UNSUPPORTED_TEMPLATE_VAR_RE.findall(template) if unsupported_vars: - placeholders = ", ".join(f"{{{name}}}" for name in sorted(PLACEHOLDER_NAMES)) + placeholders = ", ".join(f"{{{{ {name} }}}}" for name in sorted(PLACEHOLDER_NAMES)) raise InvalidMetricPresetTemplate( f"Unsupported template variables: {unsupported_vars}. " f"Use placeholders {placeholders} or literal PromQL values." ) + if _LEGACY_TEMPLATE_RE.search(template): + raise InvalidMetricPresetTemplate( + "Legacy str.format template syntax is no longer supported; " + f"use {{{{ labels }}}}, {{{{ window }}}}, {{{{ group_by }}}}: {template!r}" + ) + try: + ast = PROMQL_TEMPLATE_ENV.parse(template) + except TemplateSyntaxError as e: + raise InvalidMetricPresetTemplate(f"Invalid template syntax ({e}): {template!r}") from e + for node in _walk(ast): + if not isinstance(node, _ALLOWED_NODE_TYPES): + raise InvalidMetricPresetTemplate( + f"Only {{{{ placeholder }}}} substitution is allowed; " + f"found {type(node).__name__}: {template!r}" + ) try: - escape_non_placeholders(template).format(labels="", window="", group_by="") - except (ValueError, KeyError, IndexError) as e: + # Smoke-render with empty values; StrictUndefined rejects unknown variables. + PROMQL_TEMPLATE_ENV.from_string(template).render(labels="", window="", group_by="") + except TemplateError as e: raise InvalidMetricPresetTemplate( f"Failed to render PromQL template ({type(e).__name__}: {e}): {template!r}" ) from e - return template diff --git a/src/ai/backend/manager/api/gql/prometheus_query_preset/types/inputs.py b/src/ai/backend/manager/api/gql/prometheus_query_preset/types/inputs.py index b3ef766d318..9e998b84321 100644 --- a/src/ai/backend/manager/api/gql/prometheus_query_preset/types/inputs.py +++ b/src/ai/backend/manager/api/gql/prometheus_query_preset/types/inputs.py @@ -61,7 +61,9 @@ class CreateQueryDefinitionInput(PydanticInputMixin[CreateQueryDefinitionInputDT category_id: UUID | None = gql_field(description="Category UUID.", default=None) metric_name: str = gql_field(description="Prometheus metric name.") query_template: str = gql_field( - description="PromQL template with {labels}, {window}, {group_by} placeholders." + description=( + "PromQL template with Jinja placeholders ({{ labels }}, {{ window }}, {{ group_by }})." + ) ) time_window: str | None = gql_field(description="Default time window.", default=None) options: QueryDefinitionOptionsInput = gql_field( diff --git a/src/ai/backend/manager/clients/prometheus/fixed_query_builder.py b/src/ai/backend/manager/clients/prometheus/fixed_query_builder.py index a2978d6a6db..5df571f7b0a 100644 --- a/src/ai/backend/manager/clients/prometheus/fixed_query_builder.py +++ b/src/ai/backend/manager/clients/prometheus/fixed_query_builder.py @@ -19,20 +19,22 @@ from ai.backend.manager.clients.prometheus.types import ValueType _GAUGE_TEMPLATE: Final[str] = ( - f"sum by ({{group_by}})({CONTAINER_UTILIZATION_METRIC_NAME}{{{{{{labels}}}}}})" + "sum by ({{ group_by }})(" + CONTAINER_UTILIZATION_METRIC_NAME + "{ {{ labels }} })" ) _RATE_TEMPLATE: Final[str] = ( - "sum by ({group_by})(rate(" - f"{CONTAINER_UTILIZATION_METRIC_NAME}{{{{{{labels}}}}}}[{{window}}]))" + "sum by ({{ group_by }})(rate(" + + CONTAINER_UTILIZATION_METRIC_NAME + + "{ {{ labels }} }[{{ window }}]))" ) _DIFF_TEMPLATE: Final[str] = ( - "sum by ({group_by})(rate(" - f"{CONTAINER_UTILIZATION_METRIC_NAME}{{{{{{labels}}}}}}[{{window}}]))" + "sum by ({{ group_by }})(rate(" + + CONTAINER_UTILIZATION_METRIC_NAME + + "{ {{ labels }} }[{{ window }}]))" ) -_LIVE_STAT_MAX_TEMPLATE: Final[str] = f"max_over_time(({_GAUGE_TEMPLATE})[{{window}}:])" -_LIVE_STAT_AVG_TEMPLATE: Final[str] = f"avg_over_time(({_GAUGE_TEMPLATE})[{{window}}:])" -_LIVE_STAT_RATE_MAX_TEMPLATE: Final[str] = f"max_over_time(({_RATE_TEMPLATE})[{{window}}:])" -_LIVE_STAT_RATE_AVG_TEMPLATE: Final[str] = f"avg_over_time(({_RATE_TEMPLATE})[{{window}}:])" +_LIVE_STAT_MAX_TEMPLATE: Final[str] = "max_over_time((" + _GAUGE_TEMPLATE + ")[{{ window }}:])" +_LIVE_STAT_AVG_TEMPLATE: Final[str] = "avg_over_time((" + _GAUGE_TEMPLATE + ")[{{ window }}:])" +_LIVE_STAT_RATE_MAX_TEMPLATE: Final[str] = "max_over_time((" + _RATE_TEMPLATE + ")[{{ window }}:])" +_LIVE_STAT_RATE_AVG_TEMPLATE: Final[str] = "avg_over_time((" + _RATE_TEMPLATE + ")[{{ window }}:])" _INSTANT_GROUP_BY: Final[frozenset[str]] = frozenset({ "kernel_id", "container_metric_name", diff --git a/src/ai/backend/manager/clients/prometheus/preset.py b/src/ai/backend/manager/clients/prometheus/preset.py index 014b7888cdd..77c510f1d57 100644 --- a/src/ai/backend/manager/clients/prometheus/preset.py +++ b/src/ai/backend/manager/clients/prometheus/preset.py @@ -2,10 +2,13 @@ from collections.abc import Mapping, Sequence, Set from dataclasses import dataclass, field from enum import StrEnum +from functools import lru_cache from typing import Self +from jinja2 import Template, TemplateError + from ai.backend.common.dto.manager.v2.prometheus_query_preset.validators import ( - escape_non_placeholders, + PROMQL_TEMPLATE_ENV, ) from ai.backend.common.exception import InvalidMetricPresetTemplate @@ -42,20 +45,25 @@ def _escape_label_value(value: str) -> str: return value.replace("\\", "\\\\").replace('"', '\\"').replace("\n", "\\n").replace("\r", "\\r") +@lru_cache(maxsize=256) +def _compiled_template(template: str) -> Template: + return PROMQL_TEMPLATE_ENV.from_string(template) + + @dataclass(frozen=True) class MetricPreset: - """PromQL query preset with template (placeholders: {labels}, {window}, {group_by}).""" + """PromQL query preset with a Jinja template + (placeholders: ``{{ labels }}``, ``{{ window }}``, ``{{ group_by }}``).""" - # PromQL template (placeholders: {labels}, {window}, {group_by}) template: str - # Query labels (injected into {labels} placeholder) + # Injected into {{ labels }} labels: Mapping[str, LabelMatcher] = field(default_factory=dict) - # Group by labels (injected into {group_by} placeholder) + # Injected into {{ group_by }} group_by: Set[str] = field(default_factory=frozenset) - # Window (injected into {window} placeholder) + # Injected into {{ window }} window: str = "" def render(self) -> str: @@ -65,12 +73,12 @@ def render(self) -> str: for key, value in self.labels.items() ) try: - return escape_non_placeholders(self.template).format( + return _compiled_template(self.template).render( labels=label_str, window=self.window, group_by=",".join(sorted(self.group_by)), ) - except (ValueError, KeyError, IndexError) as e: + except TemplateError as e: raise InvalidMetricPresetTemplate( f"Failed to render PromQL template ({type(e).__name__}: {e}): {self.template!r}" ) from e diff --git a/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py b/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py new file mode 100644 index 00000000000..07cc331a754 --- /dev/null +++ b/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py @@ -0,0 +1,97 @@ +"""convert_prometheus_query_preset_templates_to_jinja + +The app renders ``query_template`` with Jinja only; the legacy ``str.format`` +syntax (``{labels}``, ``{{{labels}}}``, escaped braces) is no longer supported. +This migration rewrites all stored templates, including the seeded defaults, to +the Jinja form. The conversion helpers are a frozen copy of the removed legacy +parsing logic. Idempotent: already-Jinja templates are left untouched. + +Revision ID: 4b8e2f7a91d3 +Revises: 37d711158a8c +Create Date: 2026-08-10 00:00:00.000000 + +""" + +import re +import string + +import jinja2 +import sqlalchemy as sa +from alembic import op + +# revision identifiers, used by Alembic. +revision = "4b8e2f7a91d3" +down_revision = "37d711158a8c" +# Part of: NEXT_RELEASE_VERSION +branch_labels = None +depends_on = None + +_PLACEHOLDER_NAMES = frozenset({"labels", "window", "group_by"}) +_BRACE_BLOCK_RE = re.compile(r"\{([^{}]*)\}") + + +def _escape_non_placeholders(template: str) -> str: + """Escape a legacy template so ``str.format`` sees only the three placeholders.""" + + def repl(match: re.Match[str]) -> str: + name = match.group(1) + start, end = match.span() + text = match.string + already_wrapped = ( + start > 0 and text[start - 1] == "{" and end < len(text) and text[end] == "}" + ) + inside_escaped_braces = ( + text.rfind("{{", 0, start) > text.rfind("}}", 0, start) and text.find("}}", end) != -1 + ) + if name not in _PLACEHOLDER_NAMES: + return match.group(0) if already_wrapped else "{{" + name + "}}" + if name != "labels": + return match.group(0) + return ( + match.group(0) + if already_wrapped or inside_escaped_braces + else "{{" + match.group(0) + "}}" + ) + + return _BRACE_BLOCK_RE.sub(repl, template) + + +def _to_jinja(template: str) -> str: + """Rewrite a legacy ``str.format`` template as Jinja; other templates unchanged.""" + try: + parsed = list(string.Formatter().parse(_escape_non_placeholders(template))) + except ValueError: + return template + if not any(field in _PLACEHOLDER_NAMES for _, field, _, _ in parsed): + try: + jinja2.Environment().parse(template) + return template + except jinja2.TemplateSyntaxError: + pass # legacy escaped braces, e.g. `metric{{job="x"}}` — rebuild as literals + out = "" + for literal, field, _spec, _conv in parsed: + out += literal + if field is not None: + if out.endswith("{"): + out += " " # `{` directly before `{{` breaks the Jinja lexer + out += "{{ " + field + " }}" + return out + + +def upgrade() -> None: + conn = op.get_bind() + rows = conn.execute(sa.text("SELECT id, query_template FROM prometheus_query_presets")).all() + for row_id, template in rows: + converted = _to_jinja(template) + if converted != template: + conn.execute( + sa.text( + "UPDATE prometheus_query_presets SET query_template = :template WHERE id = :id" + ), + parameters={"template": converted, "id": row_id}, + ) + + +def downgrade() -> None: + # Data-only migration; the legacy syntax is no longer renderable by the app. + pass diff --git a/tests/component/manager/clients/prometheus/test_client_integration.py b/tests/component/manager/clients/prometheus/test_client_integration.py index 5f70a3fa3b2..4b4fe84fc46 100644 --- a/tests/component/manager/clients/prometheus/test_client_integration.py +++ b/tests/component/manager/clients/prometheus/test_client_integration.py @@ -44,7 +44,7 @@ async def prometheus_client( @pytest.fixture def up_metric_preset() -> MetricPreset: return MetricPreset( - template="up{{{labels}}}", + template="up{ {{ labels }} }", labels={"job": LabelMatcher.exact("prometheus")}, group_by=frozenset(), ) diff --git a/tests/component/manager/clients/prometheus/test_sd_relabel.py b/tests/component/manager/clients/prometheus/test_sd_relabel.py index 7d5bb8e302b..10f14ad88c5 100644 --- a/tests/component/manager/clients/prometheus/test_sd_relabel.py +++ b/tests/component/manager/clients/prometheus/test_sd_relabel.py @@ -194,7 +194,7 @@ async def prometheus_client_with_relabel( @pytest.fixture def up_model_service_preset() -> MetricPreset: return MetricPreset( - template="up{{{labels}}}", + template="up{ {{ labels }} }", labels={"service_group": LabelMatcher.exact(MODEL_SERVICE_GROUP)}, group_by=frozenset(), ) diff --git a/tests/component/prometheus_query_preset/test_prometheus_query_preset_preview.py b/tests/component/prometheus_query_preset/test_prometheus_query_preset_preview.py index 84976faa510..368eb554e67 100644 --- a/tests/component/prometheus_query_preset/test_prometheus_query_preset_preview.py +++ b/tests/component/prometheus_query_preset/test_prometheus_query_preset_preview.py @@ -23,11 +23,11 @@ class TestPrometheusQueryPresetPreview: ("query_template", "result_type"), [ # Instant vector wrapping a range vector (typical preset shape). - ("sum(rate(metric{{{labels}}}[{window}]))", "vector"), + ("sum(rate(metric{ {{ labels }} }[{{ window }}]))", "vector"), # Plain instant vector. - ("metric{{{labels}}}", "vector"), + ("metric{ {{ labels }} }", "vector"), # Raw range vector — accepted by query_instant, returns matrix. - ("metric{{{labels}}}[{window}]", "matrix"), + ("metric{ {{ labels }} }[{{ window }}]", "matrix"), ], ) async def test_returns_prometheus_response( diff --git a/tests/unit/common/dto/manager/v2/prometheus_query_preset/test_request.py b/tests/unit/common/dto/manager/v2/prometheus_query_preset/test_request.py index 744b118a79d..f49608558f3 100644 --- a/tests/unit/common/dto/manager/v2/prometheus_query_preset/test_request.py +++ b/tests/unit/common/dto/manager/v2/prometheus_query_preset/test_request.py @@ -175,12 +175,21 @@ def test_unsupported_template_var_in_query_template_raises(self) -> None: options=_make_create_options(), ) - def test_malformed_query_template_raises(self) -> None: + def test_disallowed_jinja_construct_raises(self) -> None: with pytest.raises(InvalidMetricPresetTemplate): CreateQueryDefinitionInput( name="test", metric_name="metric", - query_template="metric{", + query_template="{% if labels %}metric{% endif %}", + options=_make_create_options(), + ) + + def test_legacy_template_syntax_raises(self) -> None: + with pytest.raises(InvalidMetricPresetTemplate, match="Legacy"): + CreateQueryDefinitionInput( + name="test", + metric_name="metric", + query_template="sum by ({group_by})(metric{{{labels}}})", options=_make_create_options(), ) @@ -268,9 +277,9 @@ def test_unsupported_template_var_in_query_template_raises(self) -> None: query_template='rate(metric{mode!="idle"}[$__rate_interval])', ) - def test_malformed_query_template_raises(self) -> None: + def test_disallowed_jinja_construct_raises(self) -> None: with pytest.raises(InvalidMetricPresetTemplate): - ModifyQueryDefinitionInput(query_template="metric{") + ModifyQueryDefinitionInput(query_template="metric{ {{ unknown_var }} }") def test_round_trip_serialization(self) -> None: inp = ModifyQueryDefinitionInput( diff --git a/tests/unit/manager/clients/prometheus/test_client.py b/tests/unit/manager/clients/prometheus/test_client.py index 776e4f8f23b..21bfca4fd9b 100644 --- a/tests/unit/manager/clients/prometheus/test_client.py +++ b/tests/unit/manager/clients/prometheus/test_client.py @@ -73,7 +73,7 @@ class TestQueryRange: @pytest.fixture def sample_preset(self) -> MetricPreset: return MetricPreset( - template="sum(my_metric{{{labels}}}) by ({group_by})", + template="sum(my_metric{ {{ labels }} }) by ({{ group_by }})", labels={ "container_metric_name": LabelMatcher.exact("mem"), "value_type": LabelMatcher.exact("current"), @@ -185,7 +185,7 @@ class TestQueryInstant: @pytest.fixture def sample_preset(self) -> MetricPreset: return MetricPreset( - template="sum(my_metric{{{labels}}}) by ({group_by})", + template="sum(my_metric{ {{ labels }} }) by ({{ group_by }})", labels={ "container_metric_name": LabelMatcher.exact("mem"), "value_type": LabelMatcher.exact("current"), @@ -330,7 +330,7 @@ async def test_with_time_range_executes_range_query( result = await prometheus_client.execute_preset( MetricPreset( - template="sum(my_metric{{{labels}}}) by ({group_by})", + template="sum(my_metric{ {{ labels }} }) by ({{ group_by }})", labels={"kernel_id": LabelMatcher.exact("kernel-1")}, group_by={"kernel_id"}, window="5m", @@ -344,7 +344,7 @@ async def test_with_time_range_executes_range_query( assert mock_session.post.call_args.args[0] == "query_range" form_data = mock_session.post.call_args.kwargs["data"] field_values = {field[0]["name"]: field[2] for field in form_data._fields} - assert field_values["query"] == 'sum(my_metric{kernel_id="kernel-1"}) by (kernel_id)' + assert field_values["query"] == 'sum(my_metric{ kernel_id="kernel-1" }) by (kernel_id)' assert field_values["start"] == time_range.start assert field_values["end"] == time_range.end assert field_values["step"] == time_range.step @@ -377,7 +377,7 @@ async def test_without_time_range_executes_instant_query( ) -> None: result = await prometheus_client.execute_preset( MetricPreset( - template="sum(my_metric{{{labels}}}) by ({group_by})", + template="sum(my_metric{ {{ labels }} }) by ({{ group_by }})", labels={"kernel_id": LabelMatcher.exact("kernel-1")}, group_by={"kernel_id"}, window="5m", @@ -417,7 +417,7 @@ def client_with_custom_timeout(self, mock_pool: Mock) -> PrometheusClient: @pytest.fixture def sample_preset(self) -> MetricPreset: return MetricPreset( - template="sum(my_metric{{{labels}}})", + template="sum(my_metric{ {{ labels }} })", labels={}, group_by=frozenset(), ) diff --git a/tests/unit/manager/clients/prometheus/test_preset.py b/tests/unit/manager/clients/prometheus/test_preset.py index e9d69d8a83b..0215c13a06d 100644 --- a/tests/unit/manager/clients/prometheus/test_preset.py +++ b/tests/unit/manager/clients/prometheus/test_preset.py @@ -30,101 +30,71 @@ class TestMetricPresetRender: [ RenderTestCase( id="empty_labels", - template="sum(my_metric{{{labels}}}) by ({group_by})", + template="sum(my_metric{ {{ labels }} }) by ({{ group_by }})", labels={}, group_by=frozenset({"value_type"}), window="", - expected="sum(my_metric{}) by (value_type)", + expected="sum(my_metric{ }) by (value_type)", ), RenderTestCase( id="multiple_group_by_sorted", - template="sum(my_metric{{{labels}}}) by ({group_by})", + template="sum(my_metric{ {{ labels }} }) by ({{ group_by }})", labels={"job": LabelMatcher.exact("test")}, group_by=frozenset({"value_type", "kernel_id", "session_id"}), window="", - expected='sum(my_metric{job="test"}) by (kernel_id,session_id,value_type)', - ), - RenderTestCase( - id="group_by_deduplicated", - template="sum(my_metric{{{labels}}}) by ({group_by})", - labels={}, - group_by=frozenset([ - "a", - "b", - "a", - ]), # list allows duplicates, frozenset deduplicates - window="", - expected="sum(my_metric{}) by (a,b)", + expected='sum(my_metric{ job="test" }) by (kernel_id,session_id,value_type)', ), RenderTestCase( id="with_window", - template="sum(rate(my_metric{{{labels}}}[{window}])) by ({group_by})", + template="sum(rate(my_metric{ {{ labels }} }[{{ window }}])) by ({{ group_by }})", labels={"job": LabelMatcher.exact("test")}, group_by=frozenset({"instance"}), window="5m", - expected='sum(rate(my_metric{job="test"}[5m])) by (instance)', + expected='sum(rate(my_metric{ job="test" }[5m])) by (instance)', ), RenderTestCase( id="escapes_double_quotes_in_label_value", - template="my_metric{{{labels}}}", + template="my_metric{ {{ labels }} }", labels={"key": LabelMatcher.exact('value with "quotes"')}, group_by=frozenset(), window="", - expected='my_metric{key="value with \\"quotes\\""}', + expected='my_metric{ key="value with \\"quotes\\"" }', ), RenderTestCase( id="escapes_backslash_in_label_value", - template="my_metric{{{labels}}}", + template="my_metric{ {{ labels }} }", labels={"path": LabelMatcher.exact("C:\\Users\\test")}, group_by=frozenset(), window="", - expected='my_metric{path="C:\\\\Users\\\\test"}', + expected='my_metric{ path="C:\\\\Users\\\\test" }', ), RenderTestCase( id="escapes_newline_in_label_value", - template="my_metric{{{labels}}}", + template="my_metric{ {{ labels }} }", labels={"msg": LabelMatcher.exact("line1\nline2")}, group_by=frozenset(), window="", - expected='my_metric{msg="line1\\nline2"}', - ), - RenderTestCase( - id="escapes_mixed_special_chars", - template="my_metric{{{labels}}}", - labels={"data": LabelMatcher.exact('path\\to\\"file"\nend')}, - group_by=frozenset(), - window="", - expected='my_metric{data="path\\\\to\\\\\\"file\\"\\nend"}', + expected='my_metric{ msg="line1\\nline2" }', ), RenderTestCase( id="regex_matcher", - template="my_metric{{{labels}}}", + template="my_metric{ {{ labels }} }", labels={"kernel_id": LabelMatcher.regex("kernel-1|kernel-2")}, group_by=frozenset(), window="", - expected='my_metric{kernel_id=~"kernel-1|kernel-2"}', - ), - # Regression: original bug — `!=` in label matcher was parsed as - # str.format conversion specifier and raised ValueError. - RenderTestCase( - id="raw_label_matcher_passes_through", - template='rate(node_cpu_seconds_total{mode!="idle"}[5m])', - labels={}, - group_by=frozenset(), - window="", - expected='rate(node_cpu_seconds_total{mode!="idle"}[5m])', + expected='my_metric{ kernel_id=~"kernel-1|kernel-2" }', ), # Static and injected matchers coexist in one selector. RenderTestCase( id="static_matcher_with_all_placeholders", - template='sum by ({group_by})(rate(metric{{mode!="idle",{labels}}}[{window}]))', + template='sum by ({{ group_by }})(rate(metric{mode!="idle",{{ labels }}}[{{ window }}]))', labels={"job": LabelMatcher.exact("api")}, group_by=frozenset({"instance"}), window="5m", expected='sum by (instance)(rate(metric{mode!="idle",job="api"}[5m]))', ), - # Grafana paste with no {labels} placeholder — provided labels must - # be silently ignored, raw matcher must survive. + # Raw PromQL without placeholders — provided values are ignored, + # single braces are literal text. RenderTestCase( id="raw_template_ignores_provided_labels", template='rate(node_cpu_seconds_total{mode!="idle"}[5m])', @@ -133,39 +103,13 @@ class TestMetricPresetRender: window="5m", expected='rate(node_cpu_seconds_total{mode!="idle"}[5m])', ), - # Bare `{labels}` (single-brace) auto-wraps into PromQL `{value}`. RenderTestCase( - id="bare_labels_placeholder_auto_wraps", - template='sum by ({group_by})(rate(metric{mode!="idle"}{labels}[{window}]))', - labels={"job": LabelMatcher.exact("api")}, - group_by=frozenset({"instance"}), - window="5m", - expected='sum by (instance)(rate(metric{mode!="idle"}{job="api"}[5m]))', - ), - RenderTestCase( - id="bare_labels_with_empty_labels", - template="metric{labels}", + id="orphan_open_brace_is_literal", + template="metric{", labels={}, group_by=frozenset(), window="", - expected="metric{}", - ), - # User pre-escaped a raw matcher with `{{...}}` — must not be re-escaped. - RenderTestCase( - id="user_escaped_double_brace_matcher", - template='metric{{job="api"}}', - labels={}, - group_by=frozenset(), - window="", - expected='metric{job="api"}', - ), - RenderTestCase( - id="user_escaped_empty_braces", - template="metric{{}}", - labels={}, - group_by=frozenset(), - window="", - expected="metric{}", + expected="metric{", ), ], ids=lambda c: c.id, @@ -185,12 +129,11 @@ async def test_render(self, case: RenderTestCase) -> None: @pytest.mark.parametrize( "template", [ - pytest.param("metric}", id="orphan_close_brace"), - pytest.param("metric{", id="orphan_open_brace"), - pytest.param("metric{a{b}c}", id="nested_braces"), + pytest.param("sum(metric{{{labels}}}) by ({group_by})", id="legacy_triple_brace"), + pytest.param("metric{ {{ unknown_var }} }", id="unknown_variable"), ], ) - async def test_render_raises_on_malformed_template(self, template: str) -> None: + async def test_render_raises(self, template: str) -> None: preset = MetricPreset(template=template) with pytest.raises(InvalidMetricPresetTemplate): @@ -208,21 +151,33 @@ class TestValidateQueryTemplate: id="raw_promql", ), pytest.param( - "sum by ({group_by})(metric{{{labels}}}[{window}])", - id="with_placeholders", + 'count(metric{a="1",b=~"x|y"})', + id="multiple_matchers", ), pytest.param( - 'sum by (session_id)(metric{{value_type="current",{labels}}})', - id="static_and_dynamic_labels", + "sum by ({{ group_by }})(metric{ {{ labels }} }[{{ window }}])", + id="jinja_placeholders", ), pytest.param( - 'count(metric{a="1",b=~"x|y"})', - id="multiple_matchers", + 'sum by ({{ group_by }})(metric{mode!="idle",{{ labels }}})', + id="static_and_dynamic_labels", ), ], ) def test_accepts_valid_template(self, template: str) -> None: - assert validate_query_template(template) == template + validate_query_template(template) # does not raise + + @pytest.mark.parametrize( + "template", + [ + pytest.param("sum(metric{labels})", id="bare_placeholder"), + pytest.param("sum by ({group_by})(metric[{window}])", id="bare_group_by_and_window"), + pytest.param("sum(metric{{{labels}}})", id="triple_brace"), + ], + ) + def test_rejects_legacy_syntax(self, template: str) -> None: + with pytest.raises(InvalidMetricPresetTemplate, match="Legacy"): + validate_query_template(template) @pytest.mark.parametrize( "template", @@ -239,11 +194,13 @@ def test_rejects_unsupported_template_variables(self, template: str) -> None: @pytest.mark.parametrize( "template", [ - pytest.param("metric}", id="orphan_close_brace"), - pytest.param("metric{", id="orphan_open_brace"), - pytest.param("metric{a{b}c}", id="nested_braces"), + pytest.param("{% if labels %}metric{% endif %}", id="statement_block"), + pytest.param("metric{ {{ unknown_var }} }", id="unknown_variable"), + pytest.param("metric{ {{ labels | upper }} }", id="filter"), + pytest.param("metric{ {{ labels.attr }} }", id="attribute_access"), + pytest.param(" ", id="blank"), ], ) - def test_rejects_malformed_template(self, template: str) -> None: + def test_rejects_disallowed_constructs(self, template: str) -> None: with pytest.raises(InvalidMetricPresetTemplate): validate_query_template(template) diff --git a/tests/unit/manager/repositories/metric/test_session_utilization.py b/tests/unit/manager/repositories/metric/test_session_utilization.py index e81a2a64070..92856cdc4bf 100644 --- a/tests/unit/manager/repositories/metric/test_session_utilization.py +++ b/tests/unit/manager/repositories/metric/test_session_utilization.py @@ -82,8 +82,8 @@ def preset(self) -> PrometheusQueryPresetData: category_id=None, metric_name="cpu_used", query_template=( - 'avg by ({group_by}) (backendai_container_utilization{{value_type="current",' - "{labels}}})" + 'avg by ({{ group_by }}) (backendai_container_utilization{value_type="current",' + "{{ labels }}})" ), time_window="5m", filter_labels=["session_id", "container_metric_name"], diff --git a/tests/unit/manager/repositories/prometheus_query_preset/test_prometheus_query_preset_options.py b/tests/unit/manager/repositories/prometheus_query_preset/test_prometheus_query_preset_options.py index faae7c86794..fead56ddf33 100644 --- a/tests/unit/manager/repositories/prometheus_query_preset/test_prometheus_query_preset_options.py +++ b/tests/unit/manager/repositories/prometheus_query_preset/test_prometheus_query_preset_options.py @@ -37,7 +37,7 @@ class PresetSeed: name: str metric_name: str = "backendai_metric" - query_template: str = "{metric_name}{{{labels}}}" + query_template: str = "{metric_name}{ {{ labels }} }" time_window: str | None = "5m" filter_labels: tuple[str, ...] = () group_labels: tuple[str, ...] = () diff --git a/tests/unit/manager/repositories/prometheus_query_preset/test_prometheus_query_preset_repository.py b/tests/unit/manager/repositories/prometheus_query_preset/test_prometheus_query_preset_repository.py index 2254961998c..f4d9b4f5eb2 100644 --- a/tests/unit/manager/repositories/prometheus_query_preset/test_prometheus_query_preset_repository.py +++ b/tests/unit/manager/repositories/prometheus_query_preset/test_prometheus_query_preset_repository.py @@ -82,7 +82,7 @@ async def sample_preset_id( id=preset_id, name="container_cpu_rate", metric_name="backendai_container_utilization", - query_template="sum by ({group_by})(rate({metric_name}{{{labels}}}[{window}]))", + query_template="sum by ({{ group_by }})(rate({metric_name}{ {{ labels }} }[{{ window }}]))", time_window="5m", options=PresetOptions( filter_labels=["container_metric_name", "kernel_id"], @@ -127,7 +127,7 @@ async def test_create( ) -> None: name = "gpu_memory_usage" metric_name = "backendai_gpu_memory" - query_template = "avg({metric_name}{{{labels}}})" + query_template = "avg({metric_name}{ {{ labels }} })" time_window = "10m" filter_labels = ["kernel_id", "device_id"] group_labels = ["kernel_id"] @@ -302,12 +302,12 @@ async def test_delegates_to_client_with_template_and_window( canned_response: PrometheusResponse, ) -> None: result = await repository.preview_template( - query_template="sum(rate(metric{{{labels}}}[{window}]))", + query_template="sum(rate(metric{ {{ labels }} }[{{ window }}]))", default_window="5m", ) prometheus_client.preview_query_template.assert_called_once_with( - query_template="sum(rate(metric{{{labels}}}[{window}]))", + query_template="sum(rate(metric{ {{ labels }} }[{{ window }}]))", default_window="5m", ) assert result is canned_response diff --git a/tests/unit/manager/services/idle_checker/test_service.py b/tests/unit/manager/services/idle_checker/test_service.py index bce593e62d4..af3ebb53b37 100644 --- a/tests/unit/manager/services/idle_checker/test_service.py +++ b/tests/unit/manager/services/idle_checker/test_service.py @@ -80,7 +80,7 @@ def preset(self) -> PrometheusQueryPresetData: rank=0, category_id=None, metric_name="backendai_container_utilization", - query_template="sum by ({group_by})(backendai_container_utilization{{{labels}}})", + query_template="sum by ({{ group_by }})(backendai_container_utilization{ {{ labels }} })", time_window="5m", filter_labels=["container_metric_name", "session_id"], group_labels=["session_id", "device"], diff --git a/tests/unit/manager/services/prometheus_query_preset/test_prometheus_query_preset_service.py b/tests/unit/manager/services/prometheus_query_preset/test_prometheus_query_preset_service.py index 2bb18ea702d..03c16fece5e 100644 --- a/tests/unit/manager/services/prometheus_query_preset/test_prometheus_query_preset_service.py +++ b/tests/unit/manager/services/prometheus_query_preset/test_prometheus_query_preset_service.py @@ -60,7 +60,7 @@ def preset_data(self) -> PrometheusQueryPresetData: rank=0, category_id=None, metric_name="backendai_container_cpu_util", - query_template="rate(container_cpu_usage_seconds_total{{{labels}}}[{window}])", + query_template="rate(container_cpu_usage_seconds_total{ {{ labels }} }[{{ window }}])", time_window="5m", filter_labels=["kernel_id", "session_id"], group_labels=["kernel_id"], @@ -494,10 +494,10 @@ async def test_preview_uses_server_default_window( ) await service.preview_preset( - PreviewPresetAction(query_template="sum(rate(metric{{{labels}}}[{window}]))") + PreviewPresetAction(query_template="sum(rate(metric{ {{ labels }} }[{{ window }}]))") ) mock_repository.preview_template.assert_called_once_with( - query_template="sum(rate(metric{{{labels}}}[{window}]))", + query_template="sum(rate(metric{ {{ labels }} }[{{ window }}]))", default_window="1m", ) diff --git a/tests/unit/manager/services/utilization_metric/test_container_metric.py b/tests/unit/manager/services/utilization_metric/test_container_metric.py index 7b1f7bf05db..d18ca744f14 100644 --- a/tests/unit/manager/services/utilization_metric/test_container_metric.py +++ b/tests/unit/manager/services/utilization_metric/test_container_metric.py @@ -681,7 +681,7 @@ class TestBuiltinQueryProvider: timewindow="5m", expected_query=( "sum by (value_type)(backendai_container_utilization" - '{container_metric_name="mem",value_type="current"})' + '{ container_metric_name="mem",value_type="current" })' ), ), BuiltinQueryTestCase( @@ -694,8 +694,8 @@ class TestBuiltinQueryProvider: timewindow="5m", expected_query=( "sum by (user_id,value_type)(backendai_container_utilization" - '{container_metric_name="mem",value_type="capacity",' - 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4"})' + '{ container_metric_name="mem",value_type="capacity",' + 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4" })' ), ), BuiltinQueryTestCase( @@ -708,8 +708,8 @@ class TestBuiltinQueryProvider: timewindow="5m", expected_query=( "sum by (user_id,value_type)(backendai_container_utilization" - '{container_metric_name="cuda_util",value_type="current",' - 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4"})' + '{ container_metric_name="cuda_util",value_type="current",' + 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4" })' ), ), BuiltinQueryTestCase( @@ -722,8 +722,8 @@ class TestBuiltinQueryProvider: timewindow="5m", expected_query=( "sum by (user_id,value_type)(backendai_container_utilization" - '{container_metric_name="io_read",value_type="current",' - 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4"})' + '{ container_metric_name="io_read",value_type="current",' + 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4" })' ), ), # RATE - rate() already returns per-second values. @@ -734,7 +734,7 @@ class TestBuiltinQueryProvider: timewindow="5m", expected_query=( "sum by (value_type)(rate(backendai_container_utilization" - '{container_metric_name="net_rx",value_type="current"}[5m]))' + '{ container_metric_name="net_rx",value_type="current" }[5m]))' ), ), BuiltinQueryTestCase( @@ -747,8 +747,8 @@ class TestBuiltinQueryProvider: timewindow="5m", expected_query=( "sum by (user_id,value_type)(rate(backendai_container_utilization" - '{container_metric_name="net_tx",value_type="current",' - 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4"}[5m]))' + '{ container_metric_name="net_tx",value_type="current",' + 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4" }[5m]))' ), ), BuiltinQueryTestCase( @@ -761,8 +761,8 @@ class TestBuiltinQueryProvider: timewindow="5m", expected_query=( "sum by (user_id,value_type)(rate(backendai_container_utilization" - '{container_metric_name="net_rx",value_type="capacity",' - 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4"}[5m]))' + '{ container_metric_name="net_rx",value_type="capacity",' + 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4" }[5m]))' ), ), # DIFF - uses window but no interval divisor @@ -773,7 +773,7 @@ class TestBuiltinQueryProvider: timewindow="5m", expected_query=( "sum by (value_type)(rate(backendai_container_utilization" - '{container_metric_name="cpu_util",value_type="current"}[5m]))' + '{ container_metric_name="cpu_util",value_type="current" }[5m]))' ), ), BuiltinQueryTestCase( @@ -786,8 +786,8 @@ class TestBuiltinQueryProvider: timewindow="5m", expected_query=( "sum by (user_id,value_type)(rate(backendai_container_utilization" - '{container_metric_name="cpu_util",value_type="current",' - 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4"}[5m]))' + '{ container_metric_name="cpu_util",value_type="current",' + 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4" }[5m]))' ), ), # GAUGE for cpu_util with capacity (not DIFF since value_type != current) @@ -801,8 +801,8 @@ class TestBuiltinQueryProvider: timewindow="5m", expected_query=( "sum by (user_id,value_type)(backendai_container_utilization" - '{container_metric_name="cpu_util",value_type="capacity",' - 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4"})' + '{ container_metric_name="cpu_util",value_type="capacity",' + 'user_id="f38dea23-50fa-42a0-b5ae-338f5f4693f4" })' ), ), ], From 34de6e2494b322c8794cb2422ff8d65c7e2c76f2 Mon Sep 17 00:00:00 2001 From: BoKeum Date: Mon, 10 Aug 2026 10:25:57 +0000 Subject: [PATCH 2/8] chore: update api schema dump Co-authored-by: octodog --- docs/manager/graphql-reference/supergraph.graphql | 4 +++- docs/manager/graphql-reference/v2-schema.graphql | 4 +++- docs/manager/rest-reference/openapi.json | 4 ++-- 3 files changed, 8 insertions(+), 4 deletions(-) diff --git a/docs/manager/graphql-reference/supergraph.graphql b/docs/manager/graphql-reference/supergraph.graphql index d3482b87107..b4ba29726d7 100644 --- a/docs/manager/graphql-reference/supergraph.graphql +++ b/docs/manager/graphql-reference/supergraph.graphql @@ -4556,7 +4556,9 @@ input CreateQueryDefinitionInput """Prometheus metric name.""" metricName: String! - """PromQL template with {labels}, {window}, {group_by} placeholders.""" + """ + PromQL template with Jinja placeholders ({{ labels }}, {{ window }}, {{ group_by }}). + """ queryTemplate: String! """Default time window.""" diff --git a/docs/manager/graphql-reference/v2-schema.graphql b/docs/manager/graphql-reference/v2-schema.graphql index 05d479f0a18..5d3a0a9ddd7 100644 --- a/docs/manager/graphql-reference/v2-schema.graphql +++ b/docs/manager/graphql-reference/v2-schema.graphql @@ -3068,7 +3068,9 @@ input CreateQueryDefinitionInput { """Prometheus metric name.""" metricName: String! - """PromQL template with {labels}, {window}, {group_by} placeholders.""" + """ + PromQL template with Jinja placeholders ({{ labels }}, {{ window }}, {{ group_by }}). + """ queryTemplate: String! """Default time window.""" diff --git a/docs/manager/rest-reference/openapi.json b/docs/manager/rest-reference/openapi.json index 45cf94f1c5b..f7d17fdce08 100644 --- a/docs/manager/rest-reference/openapi.json +++ b/docs/manager/rest-reference/openapi.json @@ -17143,7 +17143,7 @@ "type": "string" }, "query_template": { - "description": "PromQL template with placeholders", + "description": "PromQL template with Jinja placeholders ({{ labels }}, {{ window }}, {{ group_by }})", "title": "Query Template", "type": "string" }, @@ -17481,7 +17481,7 @@ } ], "default": null, - "description": "Updated PromQL template with placeholders", + "description": "Updated PromQL template with Jinja placeholders ({{ labels }}, {{ window }}, {{ group_by }})", "title": "Query Template" }, "time_window": { From a4d618ea83433657e96b194056456bc3b4573cef Mon Sep 17 00:00:00 2001 From: BoKeum Date: Tue, 11 Aug 2026 10:12:09 +0900 Subject: [PATCH 3/8] changelog: add news fragment for PR #13675 --- changes/13675.feature.md | 1 + 1 file changed, 1 insertion(+) create mode 100644 changes/13675.feature.md diff --git a/changes/13675.feature.md b/changes/13675.feature.md new file mode 100644 index 00000000000..2e15d8aa55e --- /dev/null +++ b/changes/13675.feature.md @@ -0,0 +1 @@ +Switch Prometheus query preset templates to sandboxed Jinja syntax ({{ labels }}, {{ window }}, {{ group_by }}) with automatic migration of stored presets; the legacy str.format placeholder syntax is no longer accepted From fd27cb7645610de879531d66812613be9c188d0d Mon Sep 17 00:00:00 2001 From: BoKeum Date: Wed, 12 Aug 2026 14:29:38 +0900 Subject: [PATCH 4/8] refactor(BA-7317): prefer explicit loops and inline literals in the migration Co-Authored-By: Claude Fable 5 --- ...metheus_query_preset_templates_to_jinja.py | 30 ++++++++++++------- 1 file changed, 19 insertions(+), 11 deletions(-) diff --git a/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py b/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py index 07cc331a754..98df5c4cca1 100644 --- a/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py +++ b/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py @@ -26,7 +26,6 @@ branch_labels = None depends_on = None -_PLACEHOLDER_NAMES = frozenset({"labels", "window", "group_by"}) _BRACE_BLOCK_RE = re.compile(r"\{([^{}]*)\}") @@ -43,7 +42,7 @@ def repl(match: re.Match[str]) -> str: inside_escaped_braces = ( text.rfind("{{", 0, start) > text.rfind("}}", 0, start) and text.find("}}", end) != -1 ) - if name not in _PLACEHOLDER_NAMES: + if name not in ("labels", "window", "group_by"): return match.group(0) if already_wrapped else "{{" + name + "}}" if name != "labels": return match.group(0) @@ -62,7 +61,12 @@ def _to_jinja(template: str) -> str: parsed = list(string.Formatter().parse(_escape_non_placeholders(template))) except ValueError: return template - if not any(field in _PLACEHOLDER_NAMES for _, field, _, _ in parsed): + has_placeholder = False + for _literal, field, _spec, _conv in parsed: + if field in ("labels", "window", "group_by"): + has_placeholder = True + break + if not has_placeholder: try: jinja2.Environment().parse(template) return template @@ -73,7 +77,9 @@ def _to_jinja(template: str) -> str: out += literal if field is not None: if out.endswith("{"): - out += " " # `{` directly before `{{` breaks the Jinja lexer + # `{` directly before `{{` breaks the Jinja lexer (`{{{` lexes as `{{` + `{`). + # So we add a space: `metric{` + `{{ labels }}` → `metric{ {{ labels }}` + out += " " out += "{{ " + field + " }}" return out @@ -83,13 +89,15 @@ def upgrade() -> None: rows = conn.execute(sa.text("SELECT id, query_template FROM prometheus_query_presets")).all() for row_id, template in rows: converted = _to_jinja(template) - if converted != template: - conn.execute( - sa.text( - "UPDATE prometheus_query_presets SET query_template = :template WHERE id = :id" - ), - parameters={"template": converted, "id": row_id}, - ) + # Skip if the template is already Jinja or otherwise unchanged by the conversion. + if converted == template: + continue + conn.execute( + sa.text( + "UPDATE prometheus_query_presets SET query_template = :template WHERE id = :id" + ), + parameters={"template": converted, "id": row_id}, + ) def downgrade() -> None: From 82177f5c2cf9e141e5a09f52a037dfdc02a81994 Mon Sep 17 00:00:00 2001 From: BoKeum Date: Wed, 12 Aug 2026 14:32:44 +0900 Subject: [PATCH 5/8] fix(BA-7317): repoint the template migration onto the current alembic head Co-Authored-By: Claude Fable 5 --- ...91d3_convert_prometheus_query_preset_templates_to_jinja.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py b/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py index 98df5c4cca1..130af19fa87 100644 --- a/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py +++ b/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py @@ -7,7 +7,7 @@ parsing logic. Idempotent: already-Jinja templates are left untouched. Revision ID: 4b8e2f7a91d3 -Revises: 37d711158a8c +Revises: c8d51e7a3b62 Create Date: 2026-08-10 00:00:00.000000 """ @@ -21,7 +21,7 @@ # revision identifiers, used by Alembic. revision = "4b8e2f7a91d3" -down_revision = "37d711158a8c" +down_revision = "c8d51e7a3b62" # Part of: NEXT_RELEASE_VERSION branch_labels = None depends_on = None From 848646e51ba4a9fbd4d46045cb3f9dd3a3b17f47 Mon Sep 17 00:00:00 2001 From: BoKeum Date: Fri, 14 Aug 2026 10:49:30 +0900 Subject: [PATCH 6/8] changelog: reclassify the PR #13675 news fragment as a breaking change Co-Authored-By: Claude Fable 5 --- changes/{13675.feature.md => 13675.breaking.md} | 0 1 file changed, 0 insertions(+), 0 deletions(-) rename changes/{13675.feature.md => 13675.breaking.md} (100%) diff --git a/changes/13675.feature.md b/changes/13675.breaking.md similarity index 100% rename from changes/13675.feature.md rename to changes/13675.breaking.md From c3d967a58623580529d1716e802135f76fe56a02 Mon Sep 17 00:00:00 2001 From: BoKeum Date: Fri, 14 Aug 2026 11:09:55 +0900 Subject: [PATCH 7/8] fix(BA-7317): repoint the template migration onto the current alembic head Co-Authored-By: Claude Fable 5 --- ...91d3_convert_prometheus_query_preset_templates_to_jinja.py | 4 ++-- 1 file changed, 2 insertions(+), 2 deletions(-) diff --git a/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py b/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py index 130af19fa87..834103155b3 100644 --- a/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py +++ b/src/ai/backend/manager/models/alembic/versions/4b8e2f7a91d3_convert_prometheus_query_preset_templates_to_jinja.py @@ -7,7 +7,7 @@ parsing logic. Idempotent: already-Jinja templates are left untouched. Revision ID: 4b8e2f7a91d3 -Revises: c8d51e7a3b62 +Revises: e7b2c9f04d31 Create Date: 2026-08-10 00:00:00.000000 """ @@ -21,7 +21,7 @@ # revision identifiers, used by Alembic. revision = "4b8e2f7a91d3" -down_revision = "c8d51e7a3b62" +down_revision = "e7b2c9f04d31" # Part of: NEXT_RELEASE_VERSION branch_labels = None depends_on = None From 97a15fe6a5b34b8a87063146f0ed7ccc82eae002 Mon Sep 17 00:00:00 2001 From: BoKeum Date: Fri, 14 Aug 2026 15:39:28 +0900 Subject: [PATCH 8/8] refactor(BA-7317): move template validation from DTO into the service layer --- .../prometheus_query_preset/request.py | 16 ---- .../v2/prometheus_query_preset/request.py | 20 ----- .../v2/prometheus_query_preset/validators.py | 70 --------------- .../manager/clients/prometheus/client.py | 8 +- .../manager/clients/prometheus/preset.py | 84 ++++++++++++++---- src/ai/backend/manager/services/factory.py | 2 + .../prometheus_query_preset/service.py | 23 ++++- .../prometheus_query_preset/conftest.py | 2 + .../prometheus_query_preset/test_request.py | 55 ++---------- .../prometheus/test_fixed_query_builder.py | 57 +++++++----- .../manager/clients/prometheus/test_preset.py | 47 +++++----- .../test_prometheus_query_preset_service.py | 88 ++++++++++++++++++- .../test_container_metric.py | 40 ++++++--- 13 files changed, 279 insertions(+), 233 deletions(-) delete mode 100644 src/ai/backend/common/dto/manager/v2/prometheus_query_preset/validators.py diff --git a/src/ai/backend/common/dto/manager/prometheus_query_preset/request.py b/src/ai/backend/common/dto/manager/prometheus_query_preset/request.py index 49ec55ace9c..ed2810419c4 100644 --- a/src/ai/backend/common/dto/manager/prometheus_query_preset/request.py +++ b/src/ai/backend/common/dto/manager/prometheus_query_preset/request.py @@ -14,9 +14,6 @@ from ai.backend.common.dto.clients.prometheus.request import QueryTimeRange from ai.backend.common.dto.manager.defs import DEFAULT_PAGE_LIMIT, MAX_PAGE_LIMIT from ai.backend.common.dto.manager.query import StringFilter -from ai.backend.common.dto.manager.v2.prometheus_query_preset.validators import ( - validate_query_template, -) from .types import QueryDefinitionOrder @@ -51,12 +48,6 @@ class CreateQueryDefinitionRequest(BaseRequestModel): ) options: CreateQueryDefinitionOptionsRequest = Field(description="Query definition options") - @field_validator("query_template") - @classmethod - def _validate_query_template(cls, v: str) -> str: - validate_query_template(v) - return v - class ModifyQueryDefinitionOptionsRequest(BaseRequestModel): """Options for modifying a prometheus query definition. @@ -92,13 +83,6 @@ def _validate_time_window(cls, v: str | Sentinel | None) -> str | Sentinel | Non raise ValueError(f"Invalid Prometheus duration format: {v!r}") return v - @field_validator("query_template") - @classmethod - def _validate_query_template(cls, v: str | None) -> str | None: - if v is not None: - validate_query_template(v) - return v - class QueryDefinitionFilter(BaseRequestModel): """Filter for prometheus query definition search.""" diff --git a/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/request.py b/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/request.py index 536bce67b0d..98c50397724 100644 --- a/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/request.py +++ b/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/request.py @@ -15,7 +15,6 @@ from ai.backend.common.dto.manager.query import StringFilter, UUIDFilter from .types import OrderDirection, QueryDefinitionOrderField -from .validators import validate_query_template __all__ = ( # Options inputs @@ -85,12 +84,6 @@ def name_must_not_be_blank(cls, v: str) -> str: raise ValueError("name must not be blank or whitespace-only") return stripped - @field_validator("query_template") - @classmethod - def _validate_query_template(cls, v: str) -> str: - validate_query_template(v) - return v - class ModifyQueryDefinitionOptionsInput(BaseRequestModel): """Options for modifying a prometheus query definition. @@ -160,13 +153,6 @@ def _validate_time_window(cls, v: str | Sentinel | None) -> str | Sentinel | Non raise ValueError(f"Invalid Prometheus duration format: {v!r}") return v - @field_validator("query_template") - @classmethod - def _validate_query_template(cls, v: str | None) -> str | None: - if v is not None: - validate_query_template(v) - return v - class DeleteQueryDefinitionInput(BaseRequestModel): """Input for deleting a prometheus query definition.""" @@ -247,12 +233,6 @@ class PreviewQueryDefinitionInput(BaseRequestModel): query_template: str = Field(description="PromQL template to validate") - @field_validator("query_template") - @classmethod - def _validate_query_template(cls, v: str) -> str: - validate_query_template(v) - return v - class ExecuteQueryDefinitionInput(BaseRequestModel): """Input for executing a prometheus query definition.""" diff --git a/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/validators.py b/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/validators.py deleted file mode 100644 index fa402e35488..00000000000 --- a/src/ai/backend/common/dto/manager/v2/prometheus_query_preset/validators.py +++ /dev/null @@ -1,70 +0,0 @@ -"""Validators for prometheus_query_preset templates.""" - -from __future__ import annotations - -import re -from collections.abc import Iterator - -from jinja2 import StrictUndefined, TemplateError, TemplateSyntaxError, nodes -from jinja2.sandbox import ImmutableSandboxedEnvironment - -from ai.backend.common.exception import InvalidMetricPresetTemplate - -__all__ = ( - "PLACEHOLDER_NAMES", - "PROMQL_TEMPLATE_ENV", - "validate_query_template", -) - -PLACEHOLDER_NAMES = frozenset({"labels", "window", "group_by"}) - -# Sandboxed: templates are user input from the admin API. -PROMQL_TEMPLATE_ENV = ImmutableSandboxedEnvironment(undefined=StrictUndefined) - -# Literal text and `{{ placeholder }}` substitution only. -_ALLOWED_NODE_TYPES = (nodes.Template, nodes.Output, nodes.TemplateData, nodes.Name) - -_UNSUPPORTED_TEMPLATE_VAR_RE = re.compile(r"\$\{[^}]+\}|\$[A-Za-z_][A-Za-z0-9_]*") -# Bare `{placeholder}` or any `{{{`: the pre-Jinja str.format syntax. -_LEGACY_TEMPLATE_RE = re.compile(r"(? Iterator[nodes.Node]: - yield node - for child in node.iter_child_nodes(): - yield from _walk(child) - - -def validate_query_template(template: str) -> None: - """Validate a Jinja PromQL template; raises ``InvalidMetricPresetTemplate``.""" - if not template.strip(): - raise InvalidMetricPresetTemplate("Template must not be empty.") - unsupported_vars = _UNSUPPORTED_TEMPLATE_VAR_RE.findall(template) - if unsupported_vars: - placeholders = ", ".join(f"{{{{ {name} }}}}" for name in sorted(PLACEHOLDER_NAMES)) - raise InvalidMetricPresetTemplate( - f"Unsupported template variables: {unsupported_vars}. " - f"Use placeholders {placeholders} or literal PromQL values." - ) - if _LEGACY_TEMPLATE_RE.search(template): - raise InvalidMetricPresetTemplate( - "Legacy str.format template syntax is no longer supported; " - f"use {{{{ labels }}}}, {{{{ window }}}}, {{{{ group_by }}}}: {template!r}" - ) - try: - ast = PROMQL_TEMPLATE_ENV.parse(template) - except TemplateSyntaxError as e: - raise InvalidMetricPresetTemplate(f"Invalid template syntax ({e}): {template!r}") from e - for node in _walk(ast): - if not isinstance(node, _ALLOWED_NODE_TYPES): - raise InvalidMetricPresetTemplate( - f"Only {{{{ placeholder }}}} substitution is allowed; " - f"found {type(node).__name__}: {template!r}" - ) - try: - # Smoke-render with empty values; StrictUndefined rejects unknown variables. - PROMQL_TEMPLATE_ENV.from_string(template).render(labels="", window="", group_by="") - except TemplateError as e: - raise InvalidMetricPresetTemplate( - f"Failed to render PromQL template ({type(e).__name__}: {e}): {template!r}" - ) from e diff --git a/src/ai/backend/manager/clients/prometheus/client.py b/src/ai/backend/manager/clients/prometheus/client.py index 6f581efa603..2ce839c7c2f 100644 --- a/src/ai/backend/manager/clients/prometheus/client.py +++ b/src/ai/backend/manager/clients/prometheus/client.py @@ -29,7 +29,7 @@ KernelLiveStatBatchResult, MetricResultValue, ) -from ai.backend.manager.clients.prometheus.preset import MetricPreset +from ai.backend.manager.clients.prometheus.preset import MetricPreset, PromQLTemplateRenderer DEFAULT_TIMEOUT_SECONDS: float = 30.0 @@ -42,6 +42,7 @@ class PrometheusClient: _timeout: aiohttp.ClientTimeout _container_metric_query_builder: ContainerMetricQueryBuilder _container_live_stat_query_builder: ContainerLiveStatQueryBuilder + _template_renderer: PromQLTemplateRenderer def __init__( self, @@ -57,6 +58,7 @@ def __init__( self._timeout = aiohttp.ClientTimeout(total=timeout) self._container_metric_query_builder = container_metric_query_builder self._container_live_stat_query_builder = container_live_stat_query_builder + self._template_renderer = PromQLTemplateRenderer() async def fetch_available_container_metric_names(self) -> list[str]: query = self._container_metric_query_builder.get_container_metric_metadata_query() @@ -149,7 +151,7 @@ async def _query_range( Returns: PrometheusResponse with query results. """ - query = preset.render() + query = self._template_renderer.render(preset) form_data = aiohttp.FormData({ "query": query, "start": time_range.start, @@ -174,7 +176,7 @@ async def _query_instant( Returns: PrometheusResponse with query results. """ - query = preset.render() + query = self._template_renderer.render(preset) form_fields: dict[str, str] = {"query": query} if time is not None: form_fields["time"] = time diff --git a/src/ai/backend/manager/clients/prometheus/preset.py b/src/ai/backend/manager/clients/prometheus/preset.py index 77c510f1d57..f95cc261391 100644 --- a/src/ai/backend/manager/clients/prometheus/preset.py +++ b/src/ai/backend/manager/clients/prometheus/preset.py @@ -1,17 +1,24 @@ import re -from collections.abc import Mapping, Sequence, Set +from collections.abc import Callable, Iterator, Mapping, Sequence, Set from dataclasses import dataclass, field from enum import StrEnum from functools import lru_cache from typing import Self -from jinja2 import Template, TemplateError +from jinja2 import StrictUndefined, Template, TemplateError, TemplateSyntaxError, nodes +from jinja2.sandbox import ImmutableSandboxedEnvironment -from ai.backend.common.dto.manager.v2.prometheus_query_preset.validators import ( - PROMQL_TEMPLATE_ENV, -) from ai.backend.common.exception import InvalidMetricPresetTemplate +PLACEHOLDER_NAMES = frozenset({"labels", "window", "group_by"}) + +# Literal text and `{{ placeholder }}` substitution only. +_ALLOWED_NODE_TYPES = (nodes.Template, nodes.Output, nodes.TemplateData, nodes.Name) + +_UNSUPPORTED_TEMPLATE_VAR_RE = re.compile(r"\$\{[^}]+\}|\$[A-Za-z_][A-Za-z0-9_]*") +# Bare `{placeholder}` or any `{{{`: the pre-Jinja str.format syntax. +_LEGACY_TEMPLATE_RE = re.compile(r"(? str: return value.replace("\\", "\\\\").replace('"', '\\"').replace("\n", "\\n").replace("\r", "\\r") -@lru_cache(maxsize=256) -def _compiled_template(template: str) -> Template: - return PROMQL_TEMPLATE_ENV.from_string(template) +def _walk(node: nodes.Node) -> Iterator[nodes.Node]: + yield node + for child in node.iter_child_nodes(): + yield from _walk(child) @dataclass(frozen=True) @@ -66,19 +74,65 @@ class MetricPreset: # Injected into {{ window }} window: str = "" - def render(self) -> str: - """Render the PromQL query with all values injected.""" + +class PromQLTemplateRenderer: + """Validates and renders PromQL Jinja templates in a sandboxed environment.""" + + _env: ImmutableSandboxedEnvironment + _compile: Callable[[str], Template] + + def __init__(self) -> None: + # Sandboxed: templates are user input from the admin API. + self._env = ImmutableSandboxedEnvironment(undefined=StrictUndefined) + self._compile = lru_cache(maxsize=256)(self._env.from_string) + + def validate(self, template: str) -> None: + """Validate a Jinja PromQL template; raises ``InvalidMetricPresetTemplate``.""" + if not template.strip(): + raise InvalidMetricPresetTemplate("Template must not be empty.") + unsupported_vars = _UNSUPPORTED_TEMPLATE_VAR_RE.findall(template) + if unsupported_vars: + placeholders = ", ".join(f"{{{{ {name} }}}}" for name in sorted(PLACEHOLDER_NAMES)) + raise InvalidMetricPresetTemplate( + f"Unsupported template variables: {unsupported_vars}. " + f"Use placeholders {placeholders} or literal PromQL values." + ) + if _LEGACY_TEMPLATE_RE.search(template): + raise InvalidMetricPresetTemplate( + "Legacy str.format template syntax is no longer supported; " + f"use {{{{ labels }}}}, {{{{ window }}}}, {{{{ group_by }}}}: {template!r}" + ) + try: + ast = self._env.parse(template) + except TemplateSyntaxError as e: + raise InvalidMetricPresetTemplate(f"Invalid template syntax ({e}): {template!r}") from e + for node in _walk(ast): + if not isinstance(node, _ALLOWED_NODE_TYPES): + raise InvalidMetricPresetTemplate( + f"Only {{{{ placeholder }}}} substitution is allowed; " + f"found {type(node).__name__}: {template!r}" + ) + try: + # Smoke-render with empty values; StrictUndefined rejects unknown variables. + self._compile(template).render(labels="", window="", group_by="") + except TemplateError as e: + raise InvalidMetricPresetTemplate( + f"Failed to render PromQL template ({type(e).__name__}: {e}): {template!r}" + ) from e + + def render(self, preset: MetricPreset) -> str: + """Render the PromQL query with all preset values injected.""" label_str = ",".join( f'{key}{value.operator}"{_escape_label_value(value.value)}"' - for key, value in self.labels.items() + for key, value in preset.labels.items() ) try: - return _compiled_template(self.template).render( + return self._compile(preset.template).render( labels=label_str, - window=self.window, - group_by=",".join(sorted(self.group_by)), + window=preset.window, + group_by=",".join(sorted(preset.group_by)), ) except TemplateError as e: raise InvalidMetricPresetTemplate( - f"Failed to render PromQL template ({type(e).__name__}: {e}): {self.template!r}" + f"Failed to render PromQL template ({type(e).__name__}: {e}): {preset.template!r}" ) from e diff --git a/src/ai/backend/manager/services/factory.py b/src/ai/backend/manager/services/factory.py index a4d0967e6ee..d5e4f1e4854 100644 --- a/src/ai/backend/manager/services/factory.py +++ b/src/ai/backend/manager/services/factory.py @@ -4,6 +4,7 @@ from ai.backend.manager.actions.monitors import ActionMonitors from ai.backend.manager.actions.registry import ProcessorDependencies, ProcessorRegistry from ai.backend.manager.actions.validators import ActionValidators +from ai.backend.manager.clients.prometheus.preset import PromQLTemplateRenderer from ai.backend.manager.repositories.ops.repository import OpsRepository from ai.backend.manager.repositories.resource_allocation.repository import ( ResourceAllocationRepository, @@ -315,6 +316,7 @@ def create_services(args: ServiceArgs) -> Services: repository=repositories.prometheus_query_preset.repository, prometheus_client=args.prometheus_client, default_timewindow=args.config_provider.config.metric.timewindow, + template_renderer=PromQLTemplateRenderer(), ), prometheus_query_preset_category=PrometheusQueryPresetCategoryService( repository=repositories.prometheus_query_preset_category.repository, diff --git a/src/ai/backend/manager/services/prometheus_query_preset/service.py b/src/ai/backend/manager/services/prometheus_query_preset/service.py index e12ef0f0114..e72f2a5a953 100644 --- a/src/ai/backend/manager/services/prometheus_query_preset/service.py +++ b/src/ai/backend/manager/services/prometheus_query_preset/service.py @@ -1,9 +1,14 @@ import logging +from typing import cast from ai.backend.common.exception import PrometheusQueryPresetInvalidLabel from ai.backend.logging.utils import BraceStyleAdapter from ai.backend.manager.clients.prometheus.client import PrometheusClient -from ai.backend.manager.clients.prometheus.preset import LabelMatcher, MetricPreset +from ai.backend.manager.clients.prometheus.preset import ( + LabelMatcher, + MetricPreset, + PromQLTemplateRenderer, +) from ai.backend.manager.data.prometheus_query_preset import ( ExecutePresetOptions, PrometheusQueryPresetData, @@ -11,6 +16,12 @@ from ai.backend.manager.repositories.prometheus_query_preset import ( PrometheusQueryPresetRepository, ) +from ai.backend.manager.repositories.prometheus_query_preset.creators import ( + PrometheusQueryPresetCreatorSpec, +) +from ai.backend.manager.repositories.prometheus_query_preset.updaters import ( + PrometheusQueryPresetUpdaterSpec, +) from ai.backend.manager.services.prometheus_query_preset.actions import ( CreatePresetAction, CreatePresetActionResult, @@ -35,18 +46,23 @@ class PrometheusQueryPresetService: _repository: PrometheusQueryPresetRepository _prometheus_client: PrometheusClient _default_timewindow: str + _template_renderer: PromQLTemplateRenderer def __init__( self, repository: PrometheusQueryPresetRepository, prometheus_client: PrometheusClient, default_timewindow: str, + template_renderer: PromQLTemplateRenderer, ) -> None: self._repository = repository self._prometheus_client = prometheus_client self._default_timewindow = default_timewindow + self._template_renderer = template_renderer async def create_preset(self, action: CreatePresetAction) -> CreatePresetActionResult: + spec = cast(PrometheusQueryPresetCreatorSpec, action.creator.spec) + self._template_renderer.validate(spec.query_template) preset_data = await self._repository.create(action.creator) return CreatePresetActionResult(preset=preset_data) @@ -64,6 +80,10 @@ async def search_presets(self, action: SearchPresetsAction) -> SearchPresetsActi ) async def modify_preset(self, action: ModifyPresetAction) -> ModifyPresetActionResult: + spec = cast(PrometheusQueryPresetUpdaterSpec, action.updater.spec) + template = spec.query_template.optional_value() + if template is not None: + self._template_renderer.validate(template) preset_data = await self._repository.update(action.updater) return ModifyPresetActionResult(preset=preset_data) @@ -92,6 +112,7 @@ def _validate_labels( ) async def preview_preset(self, action: PreviewPresetAction) -> PreviewPresetActionResult: + self._template_renderer.validate(action.query_template) response = await self._repository.preview_template( query_template=action.query_template, default_window=self._default_timewindow, diff --git a/tests/component/prometheus_query_preset/conftest.py b/tests/component/prometheus_query_preset/conftest.py index 0eb90f4fe2a..5e2d300a948 100644 --- a/tests/component/prometheus_query_preset/conftest.py +++ b/tests/component/prometheus_query_preset/conftest.py @@ -29,6 +29,7 @@ register_v2_prometheus_query_preset_routes, ) from ai.backend.manager.clients.prometheus.client import PrometheusClient +from ai.backend.manager.clients.prometheus.preset import PromQLTemplateRenderer from ai.backend.manager.models.prometheus_query_preset import PrometheusQueryPresetRow from ai.backend.manager.models.prometheus_query_preset.row import PresetOptions from ai.backend.manager.models.prometheus_query_preset_category import ( @@ -81,6 +82,7 @@ def prometheus_query_preset_processors( repository=repo, prometheus_client=prometheus_client_mock, default_timewindow="5m", + template_renderer=PromQLTemplateRenderer(), ) return PrometheusQueryPresetProcessors( service=service, diff --git a/tests/unit/common/dto/manager/v2/prometheus_query_preset/test_request.py b/tests/unit/common/dto/manager/v2/prometheus_query_preset/test_request.py index f49608558f3..5d6e430c87f 100644 --- a/tests/unit/common/dto/manager/v2/prometheus_query_preset/test_request.py +++ b/tests/unit/common/dto/manager/v2/prometheus_query_preset/test_request.py @@ -26,7 +26,7 @@ OrderDirection, QueryDefinitionOrderField, ) -from ai.backend.common.exception import BackendAISchemaValidationFailed, InvalidMetricPresetTemplate +from ai.backend.common.exception import BackendAISchemaValidationFailed _SAMPLE_UUID = UUID("550e8400-e29b-41d4-a716-446655440000") @@ -166,33 +166,6 @@ def test_round_trip_serialization(self) -> None: assert restored.name == "cpu_usage" assert restored.time_window == "5m" - def test_unsupported_template_var_in_query_template_raises(self) -> None: - with pytest.raises(InvalidMetricPresetTemplate, match="Unsupported"): - CreateQueryDefinitionInput( - name="test", - metric_name="metric", - query_template='rate(metric{mode!="idle"}[$__rate_interval])', - options=_make_create_options(), - ) - - def test_disallowed_jinja_construct_raises(self) -> None: - with pytest.raises(InvalidMetricPresetTemplate): - CreateQueryDefinitionInput( - name="test", - metric_name="metric", - query_template="{% if labels %}metric{% endif %}", - options=_make_create_options(), - ) - - def test_legacy_template_syntax_raises(self) -> None: - with pytest.raises(InvalidMetricPresetTemplate, match="Legacy"): - CreateQueryDefinitionInput( - name="test", - metric_name="metric", - query_template="sum by ({group_by})(metric{{{labels}}})", - options=_make_create_options(), - ) - class TestModifyQueryDefinitionOptionsInput: """Tests for ModifyQueryDefinitionOptionsInput model.""" @@ -267,20 +240,6 @@ def test_partial_update(self) -> None: assert inp.metric_name == "new_metric" assert inp.query_template is None - def test_query_template_none_skips_validation(self) -> None: - inp = ModifyQueryDefinitionInput(query_template=None) - assert inp.query_template is None - - def test_unsupported_template_var_in_query_template_raises(self) -> None: - with pytest.raises(InvalidMetricPresetTemplate, match="Unsupported"): - ModifyQueryDefinitionInput( - query_template='rate(metric{mode!="idle"}[$__rate_interval])', - ) - - def test_disallowed_jinja_construct_raises(self) -> None: - with pytest.raises(InvalidMetricPresetTemplate): - ModifyQueryDefinitionInput(query_template="metric{ {{ unknown_var }} }") - def test_round_trip_serialization(self) -> None: inp = ModifyQueryDefinitionInput( name="updated", @@ -453,12 +412,8 @@ def test_round_trip_serialization(self) -> None: class TestPreviewQueryDefinitionInput: - """Tests for PreviewQueryDefinitionInput field validation.""" - - def test_rejects_empty_template(self) -> None: - with pytest.raises(InvalidMetricPresetTemplate, match="empty"): - PreviewQueryDefinitionInput(query_template="") + """Tests for PreviewQueryDefinitionInput model.""" - def test_rejects_foreign_var(self) -> None: - with pytest.raises(InvalidMetricPresetTemplate, match="Unsupported"): - PreviewQueryDefinitionInput(query_template="rate(metric[$range])") + def test_valid_creation(self) -> None: + inp = PreviewQueryDefinitionInput(query_template="sum(metric{ {{ labels }} })") + assert inp.query_template == "sum(metric{ {{ labels }} })" diff --git a/tests/unit/manager/clients/prometheus/test_fixed_query_builder.py b/tests/unit/manager/clients/prometheus/test_fixed_query_builder.py index 0ef7f2dee8c..f9fb8120677 100644 --- a/tests/unit/manager/clients/prometheus/test_fixed_query_builder.py +++ b/tests/unit/manager/clients/prometheus/test_fixed_query_builder.py @@ -16,10 +16,20 @@ ContainerMetricOptionalLabel, MetricType, ) -from ai.backend.manager.clients.prometheus.preset import LabelMatcher, MetricPreset, regex_union +from ai.backend.manager.clients.prometheus.preset import ( + LabelMatcher, + MetricPreset, + PromQLTemplateRenderer, + regex_union, +) from ai.backend.manager.clients.prometheus.types import ValueType +@pytest.fixture +def renderer() -> PromQLTemplateRenderer: + return PromQLTemplateRenderer() + + class TestGetContainerMetricType: @pytest.fixture def builder(self) -> ContainerMetricQueryBuilder: @@ -73,11 +83,14 @@ def test_gauge_query_preset(self, builder: ContainerMetricQueryBuilder) -> None: @pytest.mark.parametrize("metric_name", ["net_rx", "cpu_util"]) def test_rate_based_query_uses_rate_function( - self, builder: ContainerMetricQueryBuilder, metric_name: str + self, + builder: ContainerMetricQueryBuilder, + renderer: PromQLTemplateRenderer, + metric_name: str, ) -> None: label = ContainerMetricOptionalLabel(value_type=ValueType.CURRENT) - rendered = builder.get_container_metric_query(metric_name, label).render() + rendered = renderer.render(builder.get_container_metric_query(metric_name, label)) assert "rate(" in rendered assert "[5m]" in rendered @@ -96,14 +109,14 @@ def test_query_with_optional_labels(self, builder: ContainerMetricQueryBuilder) class TestGetContainerLiveStatQueries: - def test_kernel_id_regex_filter(self) -> None: + def test_kernel_id_regex_filter(self, renderer: PromQLTemplateRenderer) -> None: builder = ContainerLiveStatQueryBuilder("1m") kid1 = KernelId(UUID("aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa")) kid2 = KernelId(UUID("bbbbbbbb-bbbb-bbbb-bbbb-bbbbbbbbbbbb")) result = builder.get_container_live_stat_queries([kid1, kid2]) rendered = "\n".join( - query.render() + renderer.render(query) for query in ( result.instant, result.rate_current, @@ -118,37 +131,35 @@ def test_kernel_id_regex_filter(self) -> None: assert str(kid2) in rendered assert "cccccccc-cccc-cccc-cccc-cccccccccccc" not in rendered - def test_window_queries_read_current_series(self) -> None: + def test_window_queries_read_current_series(self, renderer: PromQLTemplateRenderer) -> None: builder = ContainerLiveStatQueryBuilder("1m") kid = KernelId(UUID("aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa")) result = builder.get_container_live_stat_queries([kid]) - assert "sum by (container_metric_name,kernel_id)" in result.max.render() - assert "sum by (container_metric_name,kernel_id)" in result.avg.render() - assert 'value_type="current"' in result.max.render() - assert 'value_type="current"' in result.avg.render() - assert "rate(" not in result.max.render() - assert "rate(" not in result.avg.render() + assert "sum by (container_metric_name,kernel_id)" in renderer.render(result.max) + assert "sum by (container_metric_name,kernel_id)" in renderer.render(result.avg) + assert 'value_type="current"' in renderer.render(result.max) + assert 'value_type="current"' in renderer.render(result.avg) + assert "rate(" not in renderer.render(result.max) + assert "rate(" not in renderer.render(result.avg) - def test_rate_window_queries_read_rate_series(self) -> None: + def test_rate_window_queries_read_rate_series(self, renderer: PromQLTemplateRenderer) -> None: builder = ContainerLiveStatQueryBuilder("1m") kid = KernelId(UUID("aaaaaaaa-aaaa-aaaa-aaaa-aaaaaaaaaaaa")) result = builder.get_container_live_stat_queries([kid]) - assert ( - "max_over_time((sum by (container_metric_name,kernel_id)(rate(" - in result.rate_max.render() + assert "max_over_time((sum by (container_metric_name,kernel_id)(rate(" in renderer.render( + result.rate_max ) - assert ( - "avg_over_time((sum by (container_metric_name,kernel_id)(rate(" - in result.rate_avg.render() + assert "avg_over_time((sum by (container_metric_name,kernel_id)(rate(" in renderer.render( + result.rate_avg ) - assert 'container_metric_name=~"cpu_util|net_rx|net_tx"' in result.rate_max.render() - assert 'container_metric_name=~"cpu_util|net_rx|net_tx"' in result.rate_avg.render() - assert 'value_type="current"' in result.rate_max.render() - assert 'value_type="current"' in result.rate_avg.render() + assert 'container_metric_name=~"cpu_util|net_rx|net_tx"' in renderer.render(result.rate_max) + assert 'container_metric_name=~"cpu_util|net_rx|net_tx"' in renderer.render(result.rate_avg) + assert 'value_type="current"' in renderer.render(result.rate_max) + assert 'value_type="current"' in renderer.render(result.rate_avg) class TestRegexUnion: diff --git a/tests/unit/manager/clients/prometheus/test_preset.py b/tests/unit/manager/clients/prometheus/test_preset.py index 0215c13a06d..28fd1ed50ac 100644 --- a/tests/unit/manager/clients/prometheus/test_preset.py +++ b/tests/unit/manager/clients/prometheus/test_preset.py @@ -2,16 +2,19 @@ import pytest -from ai.backend.common.dto.manager.v2.prometheus_query_preset.validators import ( - validate_query_template, -) from ai.backend.common.exception import InvalidMetricPresetTemplate -from ai.backend.manager.clients.prometheus import ( +from ai.backend.manager.clients.prometheus.preset import ( LabelMatcher, MetricPreset, + PromQLTemplateRenderer, ) +@pytest.fixture +def renderer() -> PromQLTemplateRenderer: + return PromQLTemplateRenderer() + + @dataclass class RenderTestCase: id: str @@ -22,8 +25,8 @@ class RenderTestCase: expected: str -class TestMetricPresetRender: - """Tests for MetricPreset.render() method.""" +class TestPromQLTemplateRendererRender: + """Tests for PromQLTemplateRenderer.render().""" @pytest.mark.parametrize( "case", @@ -114,7 +117,7 @@ class TestMetricPresetRender: ], ids=lambda c: c.id, ) - async def test_render(self, case: RenderTestCase) -> None: + async def test_render(self, renderer: PromQLTemplateRenderer, case: RenderTestCase) -> None: preset = MetricPreset( template=case.template, labels=case.labels, @@ -122,7 +125,7 @@ async def test_render(self, case: RenderTestCase) -> None: window=case.window, ) - result = preset.render() + result = renderer.render(preset) assert result == case.expected @@ -133,15 +136,15 @@ async def test_render(self, case: RenderTestCase) -> None: pytest.param("metric{ {{ unknown_var }} }", id="unknown_variable"), ], ) - async def test_render_raises(self, template: str) -> None: + async def test_render_raises(self, renderer: PromQLTemplateRenderer, template: str) -> None: preset = MetricPreset(template=template) with pytest.raises(InvalidMetricPresetTemplate): - preset.render() + renderer.render(preset) -class TestValidateQueryTemplate: - """Tests for validate_query_template() called from Pydantic field validators.""" +class TestPromQLTemplateRendererValidate: + """Tests for PromQLTemplateRenderer.validate() called from the service layer.""" @pytest.mark.parametrize( "template", @@ -164,8 +167,8 @@ class TestValidateQueryTemplate: ), ], ) - def test_accepts_valid_template(self, template: str) -> None: - validate_query_template(template) # does not raise + def test_accepts_valid_template(self, renderer: PromQLTemplateRenderer, template: str) -> None: + renderer.validate(template) # does not raise @pytest.mark.parametrize( "template", @@ -175,9 +178,9 @@ def test_accepts_valid_template(self, template: str) -> None: pytest.param("sum(metric{{{labels}}})", id="triple_brace"), ], ) - def test_rejects_legacy_syntax(self, template: str) -> None: + def test_rejects_legacy_syntax(self, renderer: PromQLTemplateRenderer, template: str) -> None: with pytest.raises(InvalidMetricPresetTemplate, match="Legacy"): - validate_query_template(template) + renderer.validate(template) @pytest.mark.parametrize( "template", @@ -187,9 +190,11 @@ def test_rejects_legacy_syntax(self, template: str) -> None: pytest.param('metric{region="${region}"}', id="braced_dollar_var"), ], ) - def test_rejects_unsupported_template_variables(self, template: str) -> None: + def test_rejects_unsupported_template_variables( + self, renderer: PromQLTemplateRenderer, template: str + ) -> None: with pytest.raises(InvalidMetricPresetTemplate, match="Unsupported"): - validate_query_template(template) + renderer.validate(template) @pytest.mark.parametrize( "template", @@ -201,6 +206,8 @@ def test_rejects_unsupported_template_variables(self, template: str) -> None: pytest.param(" ", id="blank"), ], ) - def test_rejects_disallowed_constructs(self, template: str) -> None: + def test_rejects_disallowed_constructs( + self, renderer: PromQLTemplateRenderer, template: str + ) -> None: with pytest.raises(InvalidMetricPresetTemplate): - validate_query_template(template) + renderer.validate(template) diff --git a/tests/unit/manager/services/prometheus_query_preset/test_prometheus_query_preset_service.py b/tests/unit/manager/services/prometheus_query_preset/test_prometheus_query_preset_service.py index 03c16fece5e..5b082f0c1fd 100644 --- a/tests/unit/manager/services/prometheus_query_preset/test_prometheus_query_preset_service.py +++ b/tests/unit/manager/services/prometheus_query_preset/test_prometheus_query_preset_service.py @@ -18,11 +18,16 @@ PrometheusResponse, ) from ai.backend.common.exception import ( + InvalidMetricPresetTemplate, PrometheusQueryPresetInvalidLabel, PrometheusQueryPresetNotFound, ) from ai.backend.manager.clients.prometheus.client import PrometheusClient -from ai.backend.manager.clients.prometheus.preset import LabelMatcher, MetricPreset +from ai.backend.manager.clients.prometheus.preset import ( + LabelMatcher, + MetricPreset, + PromQLTemplateRenderer, +) from ai.backend.manager.data.prometheus_query_preset import ( ExecutePresetOptions, PrometheusQueryPresetData, @@ -35,6 +40,12 @@ from ai.backend.manager.repositories.prometheus_query_preset import ( PrometheusQueryPresetRepository, ) +from ai.backend.manager.repositories.prometheus_query_preset.creators import ( + PrometheusQueryPresetCreatorSpec, +) +from ai.backend.manager.repositories.prometheus_query_preset.updaters import ( + PrometheusQueryPresetUpdaterSpec, +) from ai.backend.manager.services.prometheus_query_preset.actions import ( CreatePresetAction, DeletePresetAction, @@ -47,6 +58,7 @@ from ai.backend.manager.services.prometheus_query_preset.service import ( PrometheusQueryPresetService, ) +from ai.backend.manager.types import OptionalState class TestPrometheusQueryPresetService: @@ -83,6 +95,7 @@ def service( mock_prometheus_client: MagicMock, ) -> PrometheusQueryPresetService: return PrometheusQueryPresetService( + template_renderer=PromQLTemplateRenderer(), repository=mock_repository, prometheus_client=mock_prometheus_client, default_timewindow="1m", @@ -96,13 +109,45 @@ async def test_create_preset( ) -> None: mock_repository.create = AsyncMock(return_value=preset_data) - creator = MagicMock(spec=Creator) + creator = Creator( + spec=PrometheusQueryPresetCreatorSpec( + name="cpu_usage", + metric_name="metric", + query_template="sum(metric{ {{ labels }} })", + time_window=None, + filter_labels=[], + group_labels=[], + ) + ) action = CreatePresetAction(creator=creator) result = await service.create_preset(action) assert result.preset == preset_data mock_repository.create.assert_called_once_with(creator) + async def test_create_preset_rejects_invalid_template( + self, + service: PrometheusQueryPresetService, + mock_repository: MagicMock, + ) -> None: + mock_repository.create = AsyncMock() + + creator = Creator( + spec=PrometheusQueryPresetCreatorSpec( + name="test", + metric_name="metric", + query_template="sum by ({group_by})(metric{{{labels}}})", + time_window=None, + filter_labels=[], + group_labels=[], + ) + ) + action = CreatePresetAction(creator=creator) + + with pytest.raises(InvalidMetricPresetTemplate): + await service.create_preset(action) + mock_repository.create.assert_not_called() + async def test_get_preset( self, service: PrometheusQueryPresetService, @@ -193,13 +238,37 @@ async def test_modify_preset( ) -> None: mock_repository.update = AsyncMock(return_value=preset_data) - updater = MagicMock(spec=Updater) + updater = Updater( + spec=PrometheusQueryPresetUpdaterSpec( + query_template=OptionalState[str].update("sum(metric{ {{ labels }} })"), + ), + pk_value=preset_data.id, + ) action = ModifyPresetAction(preset_id=preset_data.id, updater=updater) result = await service.modify_preset(action) assert result.preset == preset_data mock_repository.update.assert_called_once_with(updater) + async def test_modify_preset_rejects_invalid_template( + self, + service: PrometheusQueryPresetService, + mock_repository: MagicMock, + ) -> None: + mock_repository.update = AsyncMock() + + updater = Updater( + spec=PrometheusQueryPresetUpdaterSpec( + query_template=OptionalState[str].update("metric{ {{ unknown_var }} }"), + ), + pk_value=uuid4(), + ) + action = ModifyPresetAction(preset_id=uuid4(), updater=updater) + + with pytest.raises(InvalidMetricPresetTemplate): + await service.modify_preset(action) + mock_repository.update.assert_not_called() + async def test_delete_preset( self, service: PrometheusQueryPresetService, @@ -501,3 +570,16 @@ async def test_preview_uses_server_default_window( query_template="sum(rate(metric{ {{ labels }} }[{{ window }}]))", default_window="1m", ) + + async def test_preview_preset_rejects_invalid_template( + self, + service: PrometheusQueryPresetService, + mock_repository: MagicMock, + ) -> None: + mock_repository.preview_template = AsyncMock() + + action = PreviewPresetAction(query_template="rate(metric[$__rate_interval])") + + with pytest.raises(InvalidMetricPresetTemplate): + await service.preview_preset(action) + mock_repository.preview_template.assert_not_called() diff --git a/tests/unit/manager/services/utilization_metric/test_container_metric.py b/tests/unit/manager/services/utilization_metric/test_container_metric.py index d18ca744f14..ab0bcf4dbe1 100644 --- a/tests/unit/manager/services/utilization_metric/test_container_metric.py +++ b/tests/unit/manager/services/utilization_metric/test_container_metric.py @@ -35,6 +35,7 @@ KernelLiveStatBatchResult, MetricType, ) +from ai.backend.manager.clients.prometheus.preset import PromQLTemplateRenderer from ai.backend.manager.clients.prometheus.types import ValueType from ai.backend.manager.repositories.metric.repository import MetricRepository from ai.backend.manager.services.metric.actions.container import ( @@ -42,6 +43,11 @@ ) +@pytest.fixture +def renderer() -> PromQLTemplateRenderer: + return PromQLTemplateRenderer() + + def _make_query_range_response( metric_data: list[dict[str, Any]], ) -> PrometheusResponse: @@ -808,11 +814,13 @@ class TestBuiltinQueryProvider: ], ids=lambda c: c.id, ) - async def test_build_query_renders_expected_promql(self, case: BuiltinQueryTestCase) -> None: + async def test_build_query_renders_expected_promql( + self, renderer: PromQLTemplateRenderer, case: BuiltinQueryTestCase + ) -> None: query_builder = ContainerMetricQueryBuilder(case.timewindow) query = query_builder.get_container_metric_query(case.metric_name, case.labels) - rendered_query = query.render() + rendered_query = renderer.render(query) assert rendered_query == case.expected_query @@ -827,17 +835,19 @@ def queries(self) -> ContainerLiveStatQueries: return query_builder.get_container_live_stat_queries([kernel_id]) def test_instant_query_fetches_live_stat_fields( - self, queries: ContainerLiveStatQueries + self, renderer: PromQLTemplateRenderer, queries: ContainerLiveStatQueries ) -> None: - rendered = queries.instant.render() + rendered = renderer.render(queries.instant) assert "backendai_container_utilization" in rendered assert "sum by (container_metric_name,kernel_id,value_type)" in rendered assert 'value_type=~"current|capacity"' in rendered assert "pct" not in rendered - def test_max_query_reads_current_series(self, queries: ContainerLiveStatQueries) -> None: - rendered = queries.max.render() + def test_max_query_reads_current_series( + self, renderer: PromQLTemplateRenderer, queries: ContainerLiveStatQueries + ) -> None: + rendered = renderer.render(queries.max) assert "label_replace" not in rendered assert "max_over_time" in rendered @@ -846,8 +856,10 @@ def test_max_query_reads_current_series(self, queries: ContainerLiveStatQueries) assert "rate(" not in rendered assert "backendai_container_utilization" in rendered - def test_rate_max_query_reads_rate_series(self, queries: ContainerLiveStatQueries) -> None: - rendered = queries.rate_max.render() + def test_rate_max_query_reads_rate_series( + self, renderer: PromQLTemplateRenderer, queries: ContainerLiveStatQueries + ) -> None: + rendered = renderer.render(queries.rate_max) assert "label_replace" not in rendered assert "max_over_time" in rendered @@ -856,8 +868,10 @@ def test_rate_max_query_reads_rate_series(self, queries: ContainerLiveStatQuerie assert 'container_metric_name=~"cpu_util|net_rx|net_tx"' in rendered assert 'value_type="current"' in rendered - def test_avg_query_reads_current_series(self, queries: ContainerLiveStatQueries) -> None: - rendered = queries.avg.render() + def test_avg_query_reads_current_series( + self, renderer: PromQLTemplateRenderer, queries: ContainerLiveStatQueries + ) -> None: + rendered = renderer.render(queries.avg) assert "label_replace" not in rendered assert "avg_over_time" in rendered @@ -866,8 +880,10 @@ def test_avg_query_reads_current_series(self, queries: ContainerLiveStatQueries) assert "rate(" not in rendered assert "backendai_container_utilization" in rendered - def test_rate_avg_query_reads_rate_series(self, queries: ContainerLiveStatQueries) -> None: - rendered = queries.rate_avg.render() + def test_rate_avg_query_reads_rate_series( + self, renderer: PromQLTemplateRenderer, queries: ContainerLiveStatQueries + ) -> None: + rendered = renderer.render(queries.rate_avg) assert "label_replace" not in rendered assert "avg_over_time" in rendered