| 1 | #!/usr/bin/env python3 |
| 2 | import json |
| 3 | import grp |
| 4 | import os |
| 5 | import re |
| 6 | import subprocess |
| 7 | import sys |
| 8 | import time |
| 9 | import urllib.parse |
| 10 | import urllib.request |
| 11 | |
| 12 | ROUTES = "/var/lib/caddy/routes.caddy" |
| 13 | TOKEN = "/var/lib/studio/router.token" |
| 14 | ROUTE_DIR = "/var/lib/studio/routes" |
| 15 | HOST = re.compile(r"[a-z0-9](?:[a-z0-9.-]*[a-z0-9])?\Z") |
| 16 | |
| 17 | |
| 18 | def nomad(path, token): |
| 19 | request = urllib.request.Request( |
| 20 | "http://127.0.0.1:4646" + path, |
| 21 | headers={"X-Nomad-Token": token}, |
| 22 | ) |
| 23 | with urllib.request.urlopen(request, timeout=5) as response: |
| 24 | return json.load(response) |
| 25 | |
| 26 | |
| 27 | def proxy(upstreams, indent, uncompressed=False, upstream_host=False): |
| 28 | lines = [f"{indent}reverse_proxy {upstreams} {{", f"{indent} lb_try_duration 5s", |
| 29 | f"{indent} fail_duration 30s"] |
| 30 | if uncompressed: |
| 31 | lines.append(f"{indent} header_up Accept-Encoding identity") |
| 32 | if upstream_host: |
| 33 | lines.append(f"{indent} header_up Host {{upstream_hostport}}") |
| 34 | return [*lines, f"{indent}}}"] |
| 35 | |
| 36 | |
| 37 | def shale_mcp_routes(port): |
| 38 | return [" @shale_mcp_link path_regexp shale_mcp_link ^/-/studio-mcp/([A-Za-z0-9_-]{43})$", |
| 39 | " handle @shale_mcp_link {", " rewrite * /oauth/shale/link/{re.shale_mcp_link.1}", |
| 40 | " request_header -User-Name", " request_header -User-Groups", " request_header -Studio-Proxy-Token", |
| 41 | *proxy(f"127.0.0.1:{port}", " "), " }", |
| 42 | " @shale_mcp_callback {", " path /-/callback", " header Cookie *studio_mcp_shale_link=*", " }", |
| 43 | " handle @shale_mcp_callback {", " rewrite * /oauth/shale/callback", |
| 44 | " request_header -User-Name", " request_header -User-Groups", " request_header -Studio-Proxy-Token", |
| 45 | *proxy(f"127.0.0.1:{port}", " "), " }"] |
| 46 | |
| 47 | |
| 48 | def render(token): |
| 49 | with open(os.environ["STUDIO_PROXY_TOKEN_FILE"]) as file: |
| 50 | dashboard_proof = file.read().strip() |
| 51 | if not re.fullmatch(r"[0-9a-fA-F]{64}", dashboard_proof): |
| 52 | raise ValueError("invalid dashboard proxy token") |
| 53 | routes = {} |
| 54 | internal_services = {} |
| 55 | auth_upstreams = set() |
| 56 | allocation_checks = {} |
| 57 | for namespace in nomad("/v1/services", token): |
| 58 | if namespace["Namespace"] != "default": |
| 59 | continue |
| 60 | for service in namespace["Services"]: |
| 61 | name = urllib.parse.quote(service["ServiceName"], safe="") |
| 62 | for instance in nomad(f"/v1/service/{name}", token): |
| 63 | tags = instance.get("Tags") or [] |
| 64 | if instance["ServiceName"] != "forward-auth" and not any(tag.startswith("caddy-") for tag in tags): |
| 65 | continue |
| 66 | address = instance["Address"] |
| 67 | port = instance["Port"] |
| 68 | if not re.fullmatch(r"[0-9a-fA-F:.]+", address) or type(port) is not int or not 1 <= port <= 65535: |
| 69 | continue |
| 70 | upstream = f"[{address}]:{port}" if ":" in address else f"{address}:{port}" |
| 71 | if not re.fullmatch(r"[a-z][a-z0-9-]*", instance["ServiceName"]): |
| 72 | raise ValueError("invalid internal service name") |
| 73 | internal_services.setdefault(instance["ServiceName"], set()).add(upstream) |
| 74 | alloc = instance["AllocID"] |
| 75 | if alloc not in allocation_checks: |
| 76 | allocation_checks[alloc] = nomad(f"/v1/allocation/{alloc}/checks", token) |
| 77 | service_checks = list(allocation_checks[alloc].values()) |
| 78 | if not service_checks or any(check["Status"] != "success" for check in service_checks): |
| 79 | continue |
| 80 | if instance["ServiceName"] == "forward-auth": |
| 81 | auth_upstreams.add(upstream) |
| 82 | access = [tag.split("=", 1)[1] for tag in tags if tag.startswith("caddy-auth-role=")] |
| 83 | user_headers = [tag.split("=", 1)[1] for tag in tags if tag.startswith("caddy-user-header=")] |
| 84 | if len(access) > 1 or (access and not re.fullmatch(r"[a-z][a-z0-9-]*", access[0])): |
| 85 | raise ValueError("invalid route access role") |
| 86 | if len(user_headers) > 1 or (user_headers and (not access or not re.fullmatch(r"[A-Za-z][A-Za-z0-9-]*", user_headers[0]))): |
| 87 | raise ValueError("invalid authenticated user header") |
| 88 | if access and not any(tag.startswith("caddy-host=") for tag in tags): |
| 89 | raise ValueError("route access role has no host") |
| 90 | for tag in tags: |
| 91 | if tag.startswith("caddy-host="): |
| 92 | host = tag.split("=", 1)[1] |
| 93 | if not HOST.fullmatch(host): |
| 94 | raise ValueError(f"invalid Caddy host: {host!r}") |
| 95 | route = routes.setdefault((80, host), {"upstreams": set(), "services": set(), "authRole": None, "userHeader": None, "internal": False, "realIp": False}) |
| 96 | if access and route["authRole"] not in (None, access[0]): |
| 97 | raise ValueError("conflicting route access roles") |
| 98 | if user_headers and route["userHeader"] not in (None, user_headers[0]): |
| 99 | raise ValueError("conflicting authenticated user headers") |
| 100 | if access: |
| 101 | route["authRole"] = access[0] |
| 102 | if user_headers: |
| 103 | route["userHeader"] = user_headers[0] |
| 104 | if "caddy-internal=true" in tags: |
| 105 | route["internal"] = True |
| 106 | if "caddy-real-ip=true" in tags: |
| 107 | route["realIp"] = True |
| 108 | route["upstreams"].add(upstream) |
| 109 | route["services"].add(instance.get("JobID") or instance["ServiceName"]) |
| 110 | elif tag.startswith("caddy-port="): |
| 111 | listener = int(tag.split("=", 1)[1]) |
| 112 | if not 1 <= listener <= 65535: |
| 113 | raise ValueError(f"invalid Caddy port: {listener}") |
| 114 | route = routes.setdefault((listener, ""), {"upstreams": set(), "services": set()}) |
| 115 | route["upstreams"].add(upstream) |
| 116 | |
| 117 | lines = [] |
| 118 | for (listener, host), route in sorted(routes.items()): |
| 119 | if host: |
| 120 | role = route["authRole"] |
| 121 | if role and not auth_upstreams: |
| 122 | continue |
| 123 | lines.append(f"{host if listener == 80 else f'{host}:{listener}'} {{") |
| 124 | if host.endswith(".test") or route["internal"]: |
| 125 | lines.append(" tls internal") |
| 126 | if route["realIp"]: |
| 127 | lines.append(" request_header X-Real-Ip {remote_host}") |
| 128 | if len(route["services"]) != 1: |
| 129 | raise ValueError(f"multiple services claim {host}") |
| 130 | service = next(iter(route["services"])) |
| 131 | if not re.fullmatch(r"[a-z][a-z0-9-]*", service): |
| 132 | raise ValueError(f"invalid service name: {service}") |
| 133 | if service == "shale": |
| 134 | port = int(os.environ["STUDIO_DASHBOARD_PORT"]) |
| 135 | if not 1 <= port <= 65535: |
| 136 | raise ValueError("invalid dashboard port for Shale linking") |
| 137 | lines += shale_mcp_routes(port) |
| 138 | lines += [" tracing {", f" span {service}", " span_attributes {", |
| 139 | f" studio.service {service}", " studio.kind edge", " }", " }"] |
| 140 | manifest = os.path.join(ROUTE_DIR, service + ".json") |
| 141 | assets = {} |
| 142 | headers = {} |
| 143 | head_html = {} |
| 144 | proof = None |
| 145 | scrub = [] |
| 146 | if os.path.exists(manifest): |
| 147 | with open(manifest) as file: |
| 148 | assets = json.load(file) |
| 149 | identity = assets.get("identity", {}) |
| 150 | headers = identity.get(host, {}) |
| 151 | head_html = assets.get("headHtml", {}).get(host, {}) |
| 152 | proof = assets.get("identityProof", {}).get(host) |
| 153 | scrub = sorted({name for configured in identity.values() for name in configured} | set(assets.get("identityProof", {}).values())) |
| 154 | if any(not re.fullmatch(r"[A-Za-z][A-Za-z0-9-]*", name) for name in scrub): |
| 155 | raise ValueError(f"invalid identity header: {host}") |
| 156 | if proof and not headers: |
| 157 | raise ValueError(f"identity proof has no identity headers: {host}") |
| 158 | if any(source not in {"X-Auth-Request-User", "X-Auth-Request-Groups", "X-Auth-Request-Preferred-Username"} for source in headers.values()): |
| 159 | raise ValueError(f"invalid identity claim: {host}") |
| 160 | if role and headers: |
| 161 | raise ValueError(f"route has both required and optional auth: {host}") |
| 162 | if role and (assets.get("files") or assets.get("dirs")): |
| 163 | raise ValueError(f"authenticated route has static overrides: {host}") |
| 164 | if head_html and (role or headers): |
| 165 | raise ValueError(f"authenticated route has HTML injection: {host}") |
| 166 | for request, markup in head_html.items(): |
| 167 | if not re.fullmatch(r"/[A-Za-z0-9/._-]*", request) or not markup: |
| 168 | raise ValueError(f"invalid HTML injection: {host} {request}") |
| 169 | for request, path in assets.get("files", {}).items(): |
| 170 | if not request.startswith("/") or " " in request or "\n" in request: |
| 171 | raise ValueError(f"invalid static path: {request}") |
| 172 | lines += [f" handle {request} {{", f" root * {json.dumps(os.path.dirname(path))}", |
| 173 | f" rewrite * /{os.path.basename(path)}", ' header Cache-Control "no-store"', |
| 174 | " file_server", " }"] |
| 175 | for request, path in assets.get("dirs", {}).items(): |
| 176 | if not request.startswith("/") or " " in request or "\n" in request: |
| 177 | raise ValueError(f"invalid static prefix: {request}") |
| 178 | lines += [f" handle_path {request}* {{", f" root * {json.dumps(path)}", |
| 179 | " file_server", " }"] |
| 180 | metrics_path = assets.get("metricsPaths", {}).get(host) |
| 181 | if metrics_path: |
| 182 | if not re.fullmatch(r"/[A-Za-z0-9/._-]+", metrics_path): |
| 183 | raise ValueError(f"invalid metrics path: {host}") |
| 184 | lines += [f" handle {metrics_path} {{", " respond 404", " }"] |
| 185 | upstreams = " ".join(sorted(route["upstreams"])) |
| 186 | if role: |
| 187 | auth = " ".join(sorted(auth_upstreams)) |
| 188 | lines += [" handle /snow.oauth2/* {", *proxy(auth, " "), " }", " handle {"] |
| 189 | if route["userHeader"]: |
| 190 | lines.append(f" request_header -{route['userHeader']}") |
| 191 | lines += [ |
| 192 | f" reverse_proxy {auth} {{", " lb_try_duration 5s", " fail_duration 30s", " method GET", |
| 193 | " rewrite /snow.oauth2/auth", " header_up X-Forwarded-Method {method}", |
| 194 | " header_up X-Forwarded-Uri {uri}", " @unauthorized status 401", |
| 195 | " handle_response @unauthorized {", |
| 196 | " redir * /snow.oauth2/sign_in?rd={scheme}://{host}{uri}", " }", |
| 197 | f" @allowed header X-Auth-Request-Groups *role:{role}*", |
| 198 | " handle_response @allowed {", " method {method}", " rewrite {uri}", |
| 199 | ] |
| 200 | if route["userHeader"]: |
| 201 | lines.append(f" request_header {route['userHeader']} {{rp.header.X-Auth-Request-Preferred-Username}}") |
| 202 | lines += [ |
| 203 | *proxy(upstreams, " "), " }", |
| 204 | " handle_response {", " respond 403", " }", " }", " }", "}", |
| 205 | ] |
| 206 | elif headers and auth_upstreams: |
| 207 | auth = " ".join(sorted(auth_upstreams)) |
| 208 | lines += [" handle /snow.oauth2/* {", *proxy(auth, " "), " }", " handle {"] |
| 209 | lines += [f" request_header -{name}" for name in scrub] |
| 210 | lines += [f" reverse_proxy {auth} {{", " lb_try_duration 5s", " fail_duration 30s", " method GET", " rewrite /snow.oauth2/auth", |
| 211 | " header_up X-Forwarded-Method {method}", " header_up X-Forwarded-Uri {uri}", |
| 212 | " @authenticated status 2xx", " handle_response @authenticated {"] |
| 213 | lines += [f" request_header {name} {{rp.header.{source}}}" for name, source in headers.items()] |
| 214 | if proof: |
| 215 | lines.append(f" request_header {proof} 1") |
| 216 | lines += [" }", " @anonymous status 4xx", " handle_response @anonymous {", |
| 217 | f" request_header -{scrub[0]}", " }", " }", |
| 218 | *proxy(upstreams, " "), " }", "}"] |
| 219 | else: |
| 220 | for index, (request, markup) in enumerate(head_html.items()): |
| 221 | name = f"@studio_head_{index}" |
| 222 | lines += [f" {name} path {request}", f" handle {name} {{", " route {", |
| 223 | f" replace </head> {json.dumps(markup + '</head>')} {{", |
| 224 | " match {", " header Content-Type text/html*", " }", " }", |
| 225 | *proxy(upstreams, " ", uncompressed=True), " }", " }"] |
| 226 | lines += [" handle {", *(f" request_header -{name}" for name in scrub), |
| 227 | *proxy(upstreams, " "), " }", "}"] |
| 228 | else: |
| 229 | lines += [f":{listener} {{", *proxy(" ".join(sorted(route["upstreams"])), " "), "}"] |
| 230 | traces = next((route for route in routes.values() if "victoria-traces" in route["services"]), None) |
| 231 | if traces: |
| 232 | lines += ["http://127.0.0.1:10428 {", " bind 127.0.0.1", |
| 233 | *proxy(" ".join(sorted(traces["upstreams"])), " "), "}"] |
| 234 | dashboard_host = "snowglobe." + os.environ["STUDIO_DOMAIN"] |
| 235 | dashboard_port = int(os.environ["STUDIO_DASHBOARD_PORT"]) |
| 236 | if not HOST.fullmatch(dashboard_host) or not 1 <= dashboard_port <= 65535: |
| 237 | raise ValueError("invalid dashboard route") |
| 238 | if (80, dashboard_host) in routes: |
| 239 | raise ValueError(f"dashboard route conflicts with a Nomad service: {dashboard_host}") |
| 240 | lines.append(f"{dashboard_host} {{") |
| 241 | if dashboard_host.endswith(".test"): |
| 242 | lines.append(" tls internal") |
| 243 | lines += [" encode zstd gzip", " tracing {", " span globe", " span_attributes {", " studio.kind edge", " }", " }"] |
| 244 | lines += [" @mcp_public path /oauth/* /mcp/* /.well-known/oauth-* /pairing /agent/connect /api/v1/*", |
| 245 | " handle @mcp_public {", " request_header -User-Name", " request_header -User-Groups", |
| 246 | " request_header -Studio-Proxy-Token", *proxy(f"127.0.0.1:{dashboard_port}", " "), " }"] |
| 247 | if auth_upstreams: |
| 248 | auth = " ".join(sorted(auth_upstreams)) |
| 249 | lines += [ |
| 250 | " handle /snow.oauth2/* {", *proxy(auth, " "), " }", |
| 251 | " handle {", " request_header -User-Name", " request_header -User-Groups", " request_header -Studio-Proxy-Token", |
| 252 | f" reverse_proxy {auth} {{", " lb_try_duration 5s", " fail_duration 30s", |
| 253 | " method GET", " rewrite /snow.oauth2/auth", |
| 254 | " header_up X-Forwarded-Method {method}", " header_up X-Forwarded-Uri {uri}", |
| 255 | " @unauthorized status 401", " handle_response @unauthorized {", |
| 256 | " redir * /snow.oauth2/sign_in?rd={scheme}://{host}{uri}", " }", |
| 257 | " @authenticated status 2xx", " handle_response @authenticated {", |
| 258 | " method {method}", " rewrite {uri}", |
| 259 | " request_header User-Name {rp.header.X-Auth-Request-Preferred-Username}", |
| 260 | " request_header User-Groups {rp.header.X-Auth-Request-Groups}", |
| 261 | f" request_header Studio-Proxy-Token {dashboard_proof}", |
| 262 | *proxy(f"127.0.0.1:{dashboard_port}", " "), |
| 263 | " }", " handle_response {", " respond 403", " }", |
| 264 | " }", " }", "}", |
| 265 | ] |
| 266 | else: |
| 267 | lines += [" header Retry-After 5", ' respond "Sign-in is unavailable. Try again in a moment." 503', "}"] |
| 268 | internal_port = os.environ.get("STUDIO_INTERNAL_PORT") |
| 269 | if internal_port is not None: |
| 270 | internal_host = "dashboard.internal." + os.environ["STUDIO_DOMAIN"] |
| 271 | if not HOST.fullmatch(internal_host) or not internal_port.isdecimal() or not 1 <= int(internal_port) <= 65535: |
| 272 | raise ValueError("invalid internal dashboard route") |
| 273 | if (80, internal_host) in routes or (int(internal_port), "") in routes: |
| 274 | raise ValueError("internal dashboard route conflicts with a Nomad service") |
| 275 | lines += [f"{internal_host}:{internal_port} {{", " tls internal", |
| 276 | f" @dashboard header Studio-Proxy-Token {dashboard_proof}", |
| 277 | " handle @dashboard {", " request_header -Studio-Proxy-Token", |
| 278 | " request_header -User-Name", " request_header -User-Groups", |
| 279 | " handle_path /nomad/* {", *proxy("127.0.0.1:4646", " ", upstream_host=True), " }"] |
| 280 | for service, upstreams in sorted(internal_services.items()): |
| 281 | lines += [f" handle_path /services/{service}/* {{", |
| 282 | *proxy(" ".join(sorted(upstreams)), " ", upstream_host=True), " }"] |
| 283 | lines += [" handle {", " respond 404", " }", " }", " handle {", " respond 403", " }", "}"] |
| 284 | return "\n".join(lines) + "\n" |
| 285 | |
| 286 | |
| 287 | def update(content): |
| 288 | try: |
| 289 | with open(ROUTES) as file: |
| 290 | previous = file.read() |
| 291 | except FileNotFoundError: |
| 292 | previous = "" |
| 293 | if content == previous: |
| 294 | return |
| 295 | pending = ROUTES + ".pending" |
| 296 | def replace(value): |
| 297 | with open(pending, "w") as file: |
| 298 | file.write(value) |
| 299 | os.chown(pending, -1, grp.getgrnam("caddy").gr_gid) |
| 300 | os.chmod(pending, 0o640) |
| 301 | os.replace(pending, ROUTES) |
| 302 | |
| 303 | replace(content) |
| 304 | result = subprocess.run(["systemctl", "reload", "caddy"], capture_output=True, text=True) |
| 305 | if result.returncode: |
| 306 | replace(previous) |
| 307 | raise RuntimeError(result.stderr.strip()) |
| 308 | print("Caddy routes updated", flush=True) |
| 309 | |
| 310 | |
| 311 | def main(): |
| 312 | with open(TOKEN) as file: |
| 313 | token = file.read().strip() |
| 314 | while True: |
| 315 | try: |
| 316 | update(render(token)) |
| 317 | except Exception as error: |
| 318 | print(f"Caddy route update failed: {error}", file=sys.stderr, flush=True) |
| 319 | time.sleep(5) |
| 320 | |
| 321 | |
| 322 | if __name__ == "__main__": |
| 323 | main() |