Spaces:
Running
Running
| """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 | |