pactbench / pact /model_identity.py
BBoran's picture
Publish current portable PACTBench release
f1fc3a0 verified
Raw
History Blame Contribute Delete
1.77 kB
"""Identity checks for OpenAI-compatible local model servers."""
from __future__ import annotations
import json
import re
import urllib.request
from typing import Any, Dict
def normalize_model_name(value: str) -> str:
return re.sub(r"[^a-z0-9]+", "", str(value).lower())
def model_name_matches(requested: str, actual: str) -> bool:
"""Allow an HF cache path/revision while still requiring the requested model family/size."""
req = normalize_model_name(requested.split("/")[-1])
act = normalize_model_name(actual)
return bool(req) and req in act
def query_model_identity(base_url: str, timeout: float = 10.0) -> Dict[str, Any]:
url = base_url.rstrip("/") + "/models"
with urllib.request.urlopen(url, timeout=timeout) as resp:
payload = json.loads(resp.read().decode("utf-8"))
actual = payload.get("model_id") or payload.get("model")
if not actual:
raise RuntimeError(f"model endpoint {url} did not return model/model_id")
revision = payload.get("model_revision")
if not revision:
match = re.search(r"/snapshots/([^/]+)", str(actual))
revision = match.group(1) if match else None
return {
"endpoint": base_url,
"model_id": str(actual),
"model_revision": revision,
"model_type": payload.get("model_type"),
}
def require_model_identity(base_url: str, requested: str, timeout: float = 10.0) -> Dict[str, Any]:
identity = query_model_identity(base_url, timeout=timeout)
if not model_name_matches(requested, identity["model_id"]):
raise RuntimeError(
f"requested model {requested!r}, but endpoint serves {identity['model_id']!r}; "
"refusing to create a mislabeled experiment"
)
return identity