LiveHouse-TS / scripts /community_endpoint_protocol.py
ziyuzhou02's picture
Deploy GitHub main 3feb6cda1511
e317359 verified
Raw
History Blame Contribute Delete
9.7 kB
"""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