File size: 4,402 Bytes
e317359
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
#!/usr/bin/env python3
"""Generate validated YAML metadata for a TS-Live community model request."""

from __future__ import annotations

import argparse
import re
from pathlib import Path
from urllib.parse import urlparse

import yaml


HUB_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]*/[A-Za-z0-9][A-Za-z0-9._-]*$")


def _hub_id(value: str, label: str) -> str:
    value = value.strip()
    if not HUB_ID.fullmatch(value):
        raise ValueError(f"{label} must have the form owner/repository")
    return value


def _endpoint_url(value: str) -> str:
    value = value.strip()
    parsed = urlparse(value)
    if parsed.scheme != "https" or not parsed.netloc:
        raise ValueError("endpoint URL must be an absolute HTTPS URL")
    if parsed.path.rstrip("/") != "/forecast":
        raise ValueError("endpoint URL path must be /forecast")
    if parsed.query or parsed.fragment:
        raise ValueError("endpoint URL must not contain a query string or fragment")
    return value


def _public_https_url(value: str, label: str) -> str:
    value = value.strip()
    parsed = urlparse(value)
    if parsed.scheme != "https" or not parsed.netloc:
        raise ValueError(f"{label} must be an absolute HTTPS URL")
    if parsed.username or parsed.password:
        raise ValueError(f"{label} must not contain embedded credentials")
    if parsed.fragment:
        raise ValueError(f"{label} must not contain a fragment")
    return value


def validate_submission_identity(
    *, model_id: str, display_name: str, code_url: str
) -> tuple[str, str, str]:
    """Validate fields that are known before a public endpoint is created."""

    model_id = _hub_id(model_id, "model ID")
    display_name = display_name.strip()
    if not display_name:
        raise ValueError("display name must not be empty")
    code_url = _public_https_url(code_url, "code URL")
    return model_id, display_name, code_url


def build_metadata(
    *, model_id: str, display_name: str, code_url: str, endpoint_url: str
) -> dict[str, object]:
    model_id, display_name, code_url = validate_submission_identity(
        model_id=model_id,
        display_name=display_name,
        code_url=code_url,
    )
    endpoint_url = _endpoint_url(endpoint_url)
    return {
        "models": [
            {
                "model_id": model_id,
                "display_name": display_name,
                "enabled": False,
                # The maintainer replaces this with the acceptance timestamp.
                # Models without an admission time cannot enter prequential tasks.
                "admitted_at": None,
                "model_type": "external_api",
                "model_link": f"https://huggingface.co/{model_id}",
                "code_link": code_url,
                "endpoint_url": endpoint_url,
                "timeout": 90,
                "max_retries": 2,
                "max_context_points": 4096,
                "max_response_bytes": 5 * 1024 * 1024,
                "require_https": True,
                "send_item_metadata": False,
            }
        ]
    }


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--model-id", required=True)
    parser.add_argument("--display-name", required=True)
    parser.add_argument(
        "--code-url",
        required=True,
        help="public HTTPS URL for the endpoint source code or service",
    )
    parser.add_argument("--endpoint-url", required=True)
    parser.add_argument("--output", type=Path, default=Path("community_model.yaml"))
    return parser.parse_args()


def main() -> int:
    args = parse_args()
    metadata = build_metadata(
        model_id=args.model_id,
        display_name=args.display_name,
        code_url=args.code_url,
        endpoint_url=args.endpoint_url,
    )
    rendered = yaml.safe_dump(metadata, sort_keys=False, allow_unicode=True)
    if yaml.safe_load(rendered) != metadata:
        raise RuntimeError("generated YAML did not pass its round-trip schema check")
    args.output.parent.mkdir(parents=True, exist_ok=True)
    args.output.write_text(rendered, encoding="utf-8")
    print(f"Wrote validated metadata to {args.output}")
    return 0


if __name__ == "__main__":
    try:
        raise SystemExit(main())
    except ValueError as exc:
        raise SystemExit(f"metadata validation failed: {exc}") from exc