"""Portable validation helpers for the TS-Live community endpoint protocol. This module intentionally depends only on ``httpx`` and the Python standard library so that the onboarding wizard works after installing the repository's small root ``requirements.txt``. It is kept separate from the evaluator's runtime adapter, which has substantially heavier benchmark dependencies. """ from __future__ import annotations import math import os import time from datetime import datetime, timezone from typing import Any, Sequence from urllib.parse import urlparse import httpx PROTOCOL_VERSION = "tsfm-realworld-v1" DEFAULT_QUANTILES = tuple(index / 10 for index in range(1, 10)) MAX_RESPONSE_BYTES = 5 * 1024 * 1024 def validate_endpoint_url(endpoint_url: str, *, require_https: bool) -> str: """Validate and normalize a public or loopback ``/forecast`` URL.""" endpoint_url = endpoint_url.strip() parsed = urlparse(endpoint_url) allowed_schemes = {"https"} if require_https else {"http", "https"} if parsed.scheme not in allowed_schemes or not parsed.netloc: scheme = "HTTPS" if require_https else "HTTP(S)" raise ValueError(f"endpoint URL must be an absolute {scheme} URL") if parsed.scheme == "http" and parsed.hostname not in { "localhost", "127.0.0.1", "::1", }: raise ValueError("plain HTTP is allowed only for a loopback endpoint") if parsed.path.rstrip("/") != "/forecast": raise ValueError("endpoint URL path must be /forecast") if parsed.username or parsed.password: raise ValueError("endpoint URL must not contain embedded credentials") if parsed.query or parsed.fragment: raise ValueError("endpoint URL must not contain a query string or fragment") return endpoint_url def health_url_for(endpoint_url: str) -> str: return endpoint_url.rsplit("/", 1)[0] + "/health" def _finite_vector(value: Any, *, label: str, expected_length: int) -> list[float]: if not isinstance(value, list) or len(value) < expected_length: raise ValueError(f"{label} must contain at least {expected_length} values") result = [] for item in value: if isinstance(item, bool): raise ValueError(f"{label} contains a non-numeric value") try: number = float(item) except (TypeError, ValueError) as exc: raise ValueError(f"{label} contains a non-numeric value") from exc if not math.isfinite(number): raise ValueError(f"{label} contains a non-finite value") result.append(number) return result[:expected_length] def _quantile_map(output: dict[str, Any]) -> dict[str, Any]: """Accept both supported response encodings for quantile forecasts.""" quantiles = output.get("quantiles") if isinstance(quantiles, dict): return {str(key): value for key, value in quantiles.items()} predictions = output.get("quantile_predictions") if not isinstance(predictions, list): raise ValueError( "forecast output must contain a quantiles object or " "quantile_predictions list" ) result: dict[str, Any] = {} for index, item in enumerate(predictions): if not isinstance(item, dict) or "values" not in item: raise ValueError( f"quantile_predictions[{index}] must contain level and values" ) raw_level = item.get("level", item.get("quantile")) if raw_level is None: raise ValueError( f"quantile_predictions[{index}] must contain level and values" ) result[f"{float(raw_level):g}"] = item["values"] return result def validate_forecast_response( payload: Any, *, prediction_length: int, quantiles: Sequence[float] = DEFAULT_QUANTILES, ) -> dict[str, Any]: """Validate a one-series response and return a compact receipt summary.""" if not isinstance(payload, dict): raise ValueError("response body must be a JSON object") outputs = payload.get("outputs", payload.get("forecasts")) if not isinstance(outputs, list) or len(outputs) != 1: raise ValueError("response must contain exactly one forecast output") output = outputs[0] if not isinstance(output, dict): raise ValueError("forecast output must be a JSON object") q_map = _quantile_map(output) normalized_q_map = {} for key, value in q_map.items(): try: normalized_key = key[1:] if key.lower().startswith("q") else key normalized_q_map[f"{float(normalized_key):g}"] = value except (TypeError, ValueError) as exc: raise ValueError(f"invalid quantile key: {key!r}") from exc validated_quantiles = {} for level in quantiles: key = f"{float(level):g}" if key not in normalized_q_map: raise ValueError(f"response is missing requested quantile {key}") validated_quantiles[key] = _finite_vector( normalized_q_map[key], label=f"quantile {key}", expected_length=prediction_length, ) mean_value = output.get("mean") if mean_value is None: mean_value = output.get("prediction") if mean_value is None: mean_value = validated_quantiles.get("0.5") if mean_value is None: raise ValueError("response must contain mean or the 0.5 quantile") mean = _finite_vector( mean_value, label="mean", expected_length=prediction_length, ) return { "forecast_keys": ["mean", *validated_quantiles.keys()], "mean_preview": mean[: min(3, len(mean))], } def build_validation_payload( *, model_id: str, prediction_length: int, context_length: int, quantiles: Sequence[float], ) -> dict[str, Any]: if prediction_length < 1: raise ValueError("prediction length must be positive") if context_length < 1: raise ValueError("context length must be positive") target = [ 10.0 + 0.05 * index + math.sin(index / 4.0) for index in range(context_length) ] return { "protocol_version": PROTOCOL_VERSION, "model": model_id, "inputs": [{"series_id": "series-validation", "target": target}], "parameters": { "prediction_length": prediction_length, "freq": "h", "quantiles": [float(level) for level in quantiles], }, } def validate_endpoint( *, endpoint_url: str, model_id: str, prediction_length: int = 8, context_length: int = 64, quantiles: Sequence[float] = DEFAULT_QUANTILES, timeout: float = 90.0, wait_seconds: float = 0.0, retry_interval: float = 15.0, require_https: bool = True, auth_token_env: str | None = None, transport: httpx.BaseTransport | None = None, ) -> dict[str, Any]: """Check health and a complete forecast request, retrying until ready.""" endpoint_url = validate_endpoint_url(endpoint_url, require_https=require_https) health_url = health_url_for(endpoint_url) if timeout <= 0: raise ValueError("timeout must be positive") if wait_seconds < 0: raise ValueError("wait seconds must not be negative") if retry_interval <= 0: raise ValueError("retry interval must be positive") headers = {} if auth_token_env: token = os.environ.get(auth_token_env) if not token: raise ValueError( f"authentication environment variable {auth_token_env!r} is not set" ) headers["Authorization"] = f"Bearer {token}" request_payload = build_validation_payload( model_id=model_id, prediction_length=prediction_length, context_length=context_length, quantiles=quantiles, ) deadline = time.monotonic() + wait_seconds last_error: Exception | None = None with httpx.Client( timeout=timeout, headers=headers, follow_redirects=False, trust_env=False, transport=transport, ) as client: while True: try: health_response = client.get(health_url) health_response.raise_for_status() forecast_response = client.post(endpoint_url, json=request_payload) forecast_response.raise_for_status() if len(forecast_response.content) > MAX_RESPONSE_BYTES: raise ValueError( f"response exceeds {MAX_RESPONSE_BYTES} byte limit" ) summary = validate_forecast_response( forecast_response.json(), prediction_length=prediction_length, quantiles=quantiles, ) return { "status": "ok", "protocol_version": PROTOCOL_VERSION, "model_id": model_id, "endpoint_url": endpoint_url, "health_url": health_url, "health_status_code": health_response.status_code, "checked_at_utc": datetime.now(timezone.utc).isoformat(), "prediction_length": prediction_length, **summary, } except (httpx.HTTPError, ValueError) as exc: last_error = exc remaining = deadline - time.monotonic() if remaining <= 0: break time.sleep(min(retry_interval, remaining)) raise RuntimeError( f"endpoint validation failed for {endpoint_url}: {last_error}" ) from last_error