diff --git a/packages/pyshared/pyproject.toml b/packages/pyshared/pyproject.toml new file mode 100644 index 0000000..4b6a0ca --- /dev/null +++ b/packages/pyshared/pyproject.toml @@ -0,0 +1,19 @@ +[project] +name = "xcloudify-shared" +version = "0.1.0" +description = "Code shared across the xcloudify service planes" +requires-python = ">=3.11" +dependencies = [ + "flask", + "requests", + "pyjwt[crypto]", + "colorlog", + "cryptography", +] + +[build-system] +requires = ["setuptools>=68"] +build-backend = "setuptools.build_meta" + +[tool.setuptools.packages.find] +include = ["xcloudify_shared*"] diff --git a/packages/pyshared/xcloudify_shared/__init__.py b/packages/pyshared/xcloudify_shared/__init__.py new file mode 100644 index 0000000..d8bbb05 --- /dev/null +++ b/packages/pyshared/xcloudify_shared/__init__.py @@ -0,0 +1,4 @@ +from xcloudify_shared.logging import logger, set_request_id, clear_request_id +from xcloudify_shared.responses import api_response + +__all__ = ["logger", "set_request_id", "clear_request_id", "api_response"] diff --git a/packages/pyshared/xcloudify_shared/cloudflare.py b/packages/pyshared/xcloudify_shared/cloudflare.py new file mode 100644 index 0000000..6f7882b --- /dev/null +++ b/packages/pyshared/xcloudify_shared/cloudflare.py @@ -0,0 +1,314 @@ +import requests +import json +import time +import secrets +import logging +from typing import List, Dict, Optional + +CLOUDFLARE_API_BASE = "https://api.cloudflare.com/client/v4" + +class CloudflareTunnelManager: + def __init__(self, api_token: str, account_id: str, zone_id: str, logger: logging.Logger): + self.api_token = api_token + self.account_id = account_id + self.zone_id = zone_id + self.headers = { + "Authorization": f"Bearer {self.api_token}", + "Content-Type": "application/json" + } + self.logger = logger + + def _make_request(self, method: str, endpoint: str, payload: Optional[dict] = None) -> dict: + """Helper method to make API requests with consistent error handling.""" + url = f"{CLOUDFLARE_API_BASE}{endpoint}" + self.logger.debug(f"Making {method} request to {url} with payload: {payload}") + + try: + response = requests.request( + method, + url, + headers=self.headers, + json=payload + ) + response.raise_for_status() + return response.json() + except requests.exceptions.RequestException as e: + self.logger.error(f"API request failed: {method} {url} - Error: {str(e)}") + if hasattr(e, 'response') and e.response is not None: + self.logger.error(f"Response content: {e.response.text}") + raise + + def verify_credentials(self) -> str: + """Check that api_token/account_id/zone_id are valid and consistent, + by making call against each. Used when a project updates its Cloudflare credentials, before saving them. + + """ + self._make_request("GET", f"/accounts/{self.account_id}/cfd_tunnel") + zone = self._make_request("GET", f"/zones/{self.zone_id}") + return zone["result"]["name"] + + def get_tunnel(self, tunnel_name: str) -> Optional[dict]: + """Get tunnel by name if it exists.""" + self.logger.info(f"Checking for existing tunnel '{tunnel_name}'...") + try: + response = self._make_request( + "GET", + f"/accounts/{self.account_id}/cfd_tunnel" + ) + # print(response) + for tunnel in response.get('result', []): + if tunnel['name'] == tunnel_name: # and tunnel['status']!='down': + + self.logger.info(f"Found existing tunnel '{tunnel_name}' (ID: {tunnel['id']}) and status {tunnel['status']}") + return tunnel + + self.logger.info(f"No existing tunnel found with name '{tunnel_name}'") + return None + except Exception as e: + self.logger.error(f"Failed to fetch tunnels list: {e}") + raise + + def create_tunnel(self, tunnel_name: str) -> dict: + """Create a new tunnel with a random secret.""" + self.logger.info(f"Creating new tunnel '{tunnel_name}'...") + payload = { + "name": tunnel_name, + "tunnel_secret": secrets.token_hex(32) + } + + try: + response = self._make_request( + "POST", + f"/accounts/{self.account_id}/cfd_tunnel", + payload + ) + tunnel = response['result'] + self.logger.info(f"Successfully created tunnel '{tunnel_name}' (ID: {tunnel['id']})") + return tunnel + except Exception as e: + self.logger.error(f"Failed to create tunnel '{tunnel_name}': {e}") + raise + + def get_tunnel_config(self, tunnel_id: str) -> dict: + """Get current tunnel configuration.""" + self.logger.info(f"Fetching current configuration for tunnel ID '{tunnel_id}'...") + try: + response = self._make_request( + "GET", + f"/accounts/{self.account_id}/cfd_tunnel/{tunnel_id}/configurations" + ) + config = response.get('result', {}) + self.logger.debug(f"Current tunnel configuration: {config}") + return config + except Exception as e: + self.logger.error(f"Failed to fetch tunnel configuration: {e}") + raise + + def update_tunnel_config(self, tunnel_id: str, ingress_rules: List[dict]) -> dict: + """Update tunnel configuration with new ingress rules.""" + self.logger.info(f"Updating configuration for tunnel ID '{tunnel_id}'...") + + # Always ensure we have a catch-all 404 rule at the end + if not any(rule.get('service') == 'http_status:404' for rule in ingress_rules): + ingress_rules.append({"service": "http_status:404"}) + + payload = { + "config": { + "ingress": ingress_rules + } + } + + self.logger.debug(f"New tunnel configuration payload: {payload}") + + try: + response = self._make_request( + "PUT", + f"/accounts/{self.account_id}/cfd_tunnel/{tunnel_id}/configurations", + payload + ) + self.logger.info(f"Successfully updated configuration for tunnel ID '{tunnel_id}'") + return response['result'] + except Exception as e: + self.logger.error(f"Failed to update tunnel configuration: {e}") + raise + + def create_dns_record(self, hostname: str, tunnel_id: str) -> dict: + """Create DNS CNAME record pointing to the tunnel.""" + self.logger.info(f"Creating DNS record for '{hostname}' pointing to tunnel '{tunnel_id}'...") + payload = { + "type": "CNAME", + "name": hostname, + "content": f"{tunnel_id}.cfargotunnel.com", + "ttl": 120, + "proxied": True + } + + try: + response = self._make_request( + "POST", + f"/zones/{self.zone_id}/dns_records", + payload + ) + self.logger.info(f"Successfully created DNS record for '{hostname}'") + return response['result'] + except Exception as e: + self.logger.error(f"Failed to create DNS record for '{hostname}': {e}") + raise + + def get_dns_record(self, hostname: str) -> Optional[dict]: + """Check if DNS record exists for given hostname.""" + self.logger.info(f"Checking for existing DNS record for '{hostname}'...") + try: + response = self._make_request( + "GET", + f"/zones/{self.zone_id}/dns_records?type=CNAME&name={hostname}" + ) + records = response.get('result', []) + if records: + self.logger.info(f"Found existing DNS record for '{hostname}'") + return records[0] + return None + except Exception as e: + self.logger.error(f"Failed to check DNS records for '{hostname}': {e}") + raise + + def delete_dns_record(self, record_id: str) -> bool: + """Delete a DNS record by ID.""" + self.logger.info(f"Deleting DNS record ID '{record_id}'...") + try: + self._make_request( + "DELETE", + f"/zones/{self.zone_id}/dns_records/{record_id}" + ) + self.logger.info(f"Successfully deleted DNS record ID '{record_id}'") + return True + except Exception as e: + self.logger.error(f"Failed to delete DNS record ID '{record_id}': {e}") + raise + + def setup_tunnel(self, tunnel_name: str, ingress_mappings: List[Dict[str, str]]) -> dict: + """ + Ensure tunnel exists with all specified ingress mappings. + + Args: + tunnel_name: Name of the tunnel to create or manage + ingress_mappings: List of dicts with: + - dns_hostname: Public hostname (e.g., 'subdomain.example.com') + - local_ip: Local service IP (e.g., 'localhost' or '192.168.1.100') + - local_port: Local service port (e.g., '8080') + """ + self.logger.info(f"Starting tunnel setup for '{tunnel_name}' with {len(ingress_mappings)} mappings...") + current_rules=[] + # Get or create tunnel + tunnel = self.get_tunnel(tunnel_name) + if not tunnel: + tunnel = self.create_tunnel(tunnel_name) + else: + current_config = self.get_tunnel_config(tunnel['id']) + current_rules = current_config.get('config', {}).get('ingress', []) + + # Filter out the catch-all 404 rule from current rules + current_rules = [rule for rule in current_rules if rule.get('service') != 'http_status:404'] + + # Prepare desired ingress rules + desired_rules = [] + dns_records_to_create = [] + + for mapping in ingress_mappings: + hostname = mapping['dns_hostname'] + service_url = f"{mapping['local_ip']}:{mapping['local_port']}" + + # Prepare ingress rule + desired_rule = { + "hostname": hostname, + "service": f"http://{service_url}", + "originRequest": {} + } + desired_rules.append(desired_rule) + + # Check if DNS record needs to be created + dns_record = self.get_dns_record(hostname) + if not dns_record: + dns_records_to_create.append(hostname) + elif dns_record['content'] != f"{tunnel['id']}.cfargotunnel.com": + self.logger.warning(f"DNS record for '{hostname}' exists but points to different target. Will recreate.") + self.delete_dns_record(dns_record['id']) + dns_records_to_create.append(hostname) + + # Update tunnel configuration if needed + if current_rules != desired_rules: + self.logger.info("Current tunnel configuration differs from desired state. Updating...") + self.update_tunnel_config(tunnel['id'], desired_rules) + else: + self.logger.info("Tunnel configuration already matches desired state.") + + # Create any missing DNS records + dns_responses = [] + for hostname in dns_records_to_create: + try: + dns_response = self.create_dns_record(hostname, tunnel['id']) + dns_responses.append({ + 'hostname': hostname, + 'response': dns_response + }) + except Exception as e: + self.logger.error(f"Failed to create DNS record for '{hostname}': {e}") + continue + + return { + 'tunnel_id': tunnel['id'], + 'tunnel_name': tunnel_name, + 'tunnel_secret': tunnel.get('credentials_file', {}).get('TunnelSecret'), + 'token': tunnel.get('token'), + 'ingress_mappings': ingress_mappings, + 'dns_records_created': dns_responses, + 'config_updated': current_rules != desired_rules + } + + def cleanup_tunnel(self, tunnel_name: str, delete_tunnel: bool = False) -> dict: + """ + Clean up all resources associated with a tunnel. + + Args: + tunnel_name: Name of the tunnel to clean up + delete_tunnel: Whether to delete the tunnel itself + """ + self.logger.info(f"Starting cleanup for tunnel '{tunnel_name}'...") + result = { + 'dns_records_deleted': [], + 'tunnel_deleted': False + } + + tunnel = self.get_tunnel(tunnel_name) + if not tunnel: + self.logger.info(f"No tunnel found with name '{tunnel_name}'") + return result + + # Get current config to find all hostnames + try: + config = self.get_tunnel_config(tunnel['id']) + ingress_rules = config.get('config', {}).get('ingress', []) + hostnames = [rule['hostname'] for rule in ingress_rules if 'hostname' in rule] + + # Delete DNS records for all hostnames + for hostname in hostnames: + dns_record = self.get_dns_record(hostname) + if dns_record: + self.delete_dns_record(dns_record['id']) + result['dns_records_deleted'].append(hostname) + except Exception as e: + self.logger.error(f"Error during cleanup of DNS records: {e}") + + # Delete tunnel if requested + if delete_tunnel: + try: + self._make_request( + "DELETE", + f"/accounts/{self.account_id}/cfd_tunnel/{tunnel['id']}?cascade=true" + ) + result['tunnel_deleted'] = True + self.logger.info(f"Successfully deleted tunnel '{tunnel_name}'") + except Exception as e: + self.logger.error(f"Failed to delete tunnel '{tunnel_name}': {e}") + + return result \ No newline at end of file diff --git a/packages/pyshared/xcloudify_shared/core_client.py b/packages/pyshared/xcloudify_shared/core_client.py new file mode 100644 index 0000000..be0f2ce --- /dev/null +++ b/packages/pyshared/xcloudify_shared/core_client.py @@ -0,0 +1,25 @@ +import requests + +SERVICE_KEY_HEADER = "X-Internal-Xcloudify-API" +_TIMEOUT = 30 + + +def core_request(method: str, path: str, base_url: str, api_key: str, **kwargs) -> requests.Response: + """Forward a request to core. `path` is relative to core's /api root, + e.g. 'workloads/virtual_machines'.""" + url = f"{base_url.rstrip('/')}/api/{path.lstrip('/')}" + headers = {SERVICE_KEY_HEADER: api_key, **kwargs.pop("headers", {})} + return requests.request(method, url, headers=headers, timeout=_TIMEOUT, **kwargs) + + +def app_core_request(method: str, path: str, **kwargs) -> requests.Response: + """core_request with the base URL and key taken from the running Flask + app's CORE_API_BASE_URL / CORE_API_KEY.""" + from flask import current_app + + return core_request( + method, path, + base_url=current_app.config["CORE_API_BASE_URL"], + api_key=current_app.config["CORE_API_KEY"], + **kwargs, + ) diff --git a/packages/pyshared/xcloudify_shared/crypto.py b/packages/pyshared/xcloudify_shared/crypto.py new file mode 100644 index 0000000..0cef52d --- /dev/null +++ b/packages/pyshared/xcloudify_shared/crypto.py @@ -0,0 +1,27 @@ +"""Symmetric encryption for secrets a product stores on a user's behalf +(a Cloudflare API token, e.g.). Each product holds its own key.""" +from cryptography.fernet import Fernet, InvalidToken + +from xcloudify_shared.logging import logger + + +class SecretBox: + def __init__(self, key: str, key_name: str): + self._key_name = key_name + self._fernet = Fernet(key) if key else None + if not key: + logger.error("%s is empty -- storing any encrypted secret will fail", key_name) + + def _require(self) -> Fernet: + if not self._fernet: + raise RuntimeError(f"{self._key_name} is not configured") + return self._fernet + + def encrypt(self, plaintext: str) -> str: + return self._require().encrypt(plaintext.encode()).decode() + + def decrypt(self, ciphertext: str) -> str: + try: + return self._require().decrypt(ciphertext.encode()).decode() + except InvalidToken: + raise ValueError(f"Stored secret could not be decrypted -- {self._key_name} may have changed") diff --git a/packages/pyshared/xcloudify_shared/env.py b/packages/pyshared/xcloudify_shared/env.py new file mode 100644 index 0000000..fa54a87 --- /dev/null +++ b/packages/pyshared/xcloudify_shared/env.py @@ -0,0 +1,30 @@ +"""Deployment mode, read the same way by every product. + +APP_ENV is the one switch. Anything that trades safety for convenience (the +auth dev bypass, today) is only possible when it names a development +environment, and an unset APP_ENV is production -- forgetting to configure a +deployment must never be what turns authentication off. +""" +from xcloudify_shared.logging import logger + +DEV_ENVS = {"dev", "development", "local"} + + +def is_dev_env(app_env: str | None) -> bool: + return (app_env or "").strip().lower() in DEV_ENVS + + +def dev_bypass_enabled(app_env: str | None, requested: bool | None) -> bool: + """Whether a product may resolve a request to a local dev user without a token. + + On by default in a dev environment, where `requested=False` (e.g. + AUTH_DEV_BYPASS=false) turns it off to exercise real sign-in locally. + Never on anywhere else, whatever `requested` says. + """ + if not is_dev_env(app_env): + if requested: + logger.error( + "AUTH_DEV_BYPASS=true ignored: APP_ENV=%r is not a development environment", app_env, + ) + return False + return requested is not False diff --git a/packages/pyshared/xcloudify_shared/gateway.py b/packages/pyshared/xcloudify_shared/gateway.py new file mode 100644 index 0000000..4bb9149 --- /dev/null +++ b/packages/pyshared/xcloudify_shared/gateway.py @@ -0,0 +1,160 @@ +"""What a product's gateway may forward to core, and the bookkeeping around it. + +Each product mounts its own gateway route (cloud: /vdcs//core/..., community: +/leases//core/...) and decides who the caller is, which tenant they act +in, and what gets force-stamped on the request. What is shared is the +allowlist of core routes a tenant may reach at all, and the hooks that read a +forwarded request back afterwards. +""" +import re +from typing import Iterable, Optional + + +class CoreRouteAllowlist: + """Which (method, subpath) pairs are exposed, and whether each needs write. + + `needs_write` returns True/False for an exposed route and None for one + that is not exposed -- callers refuse on None. + """ + + def __init__(self, *, exact: dict, by_id_collections: Iterable[str], + subresources: dict, named_listings: dict, actions: list): + self.exact = dict(exact) + self.by_id_collections = frozenset(by_id_collections) + self.subresources = {k: frozenset(v) for k, v in subresources.items()} + self.named_listings = dict(named_listings) + self.actions = [(re.compile(p) if isinstance(p, str) else p, m, w) for p, m, w in actions] + + def extend(self, *, exact: Optional[dict] = None) -> "CoreRouteAllowlist": + """A copy with extra exact routes, for a product that exposes more.""" + return CoreRouteAllowlist( + exact={**self.exact, **(exact or {})}, + by_id_collections=self.by_id_collections, + subresources=self.subresources, + named_listings=self.named_listings, + actions=self.actions, + ) + + def needs_write(self, method: str, subpath: str) -> Optional[bool]: + key = (method, subpath) + if key in self.exact: + return self.exact[key] + + # / + if method in ("GET", "PUT", "DELETE") and "/" in subpath: + collection, _, obj_id = subpath.rpartition("/") + if collection in self.by_id_collections and obj_id: + return method in ("PUT", "DELETE") + + parts = subpath.split("/") + if method == "GET" and len(parts) == 3: + # // + collection, obj_id, subresource = parts + if obj_id and subresource in self.subresources.get(collection, ()): + return False + # // + collection, verb, obj_id = parts + if obj_id and (collection, verb) in self.named_listings: + return self.named_listings[(collection, verb)] + + for pattern, m, write in self.actions: + if method == m and pattern.match(subpath): + return write + return None + + +WORKLOAD_ROUTES = CoreRouteAllowlist( + exact={ + ("POST", "workloads/virtual_machines"): True, + ("GET", "workloads/virtual_machines"): False, + ("POST", "workloads/containers"): True, + ("GET", "workloads/containers"): False, + ("GET", "workloads/pods"): False, + ("POST", "networks"): True, + ("GET", "networks"): False, + ("POST", "volumes"): True, + ("GET", "volumes"): False, + ("GET", "gpu_models"): False, + ("POST", "network_ports"): True, + ("GET", "audit/actions"): False, + ("GET", "workload_hosts"): False, + }, + by_id_collections={ + "workloads/virtual_machines", + "workloads/containers", + "workloads/pods", + "networks", + "volumes", + "network_ports", + }, + subresources={ + "volumes": {"workloads"}, + "workloads": {"volumes"}, + }, + named_listings={ + ("network_ports", "by_network"): False, + ("audit", "container"): False, + ("audit", "pod"): False, + }, + actions=[ + (r"^volumes/[^/]+/attach$", "POST", True), + (r"^volumes/[^/]+/detach/[^/]+$", "DELETE", True), + (r"^workloads/containers/[^/]+/lifecycle/[^/]+$", "POST", True), + (r"^workloads/pods/[^/]+/lifecycle/[^/]+$", "POST", True), + (r"^workloads/virtual_machines/[^/]+/action$", "POST", True), + (r"^networks/[^/]+/nscontroller$", "POST", True), + ], +) + + +def pending_dns_exposures(subpath: str, method: str, json_body, resp): + """After a container create, the ports it asked to expose publicly. + + Core only stores `use_dns` as an intent flag and takes no action on it, so + the product above core provisions the Cloudflare side. Returns + (pod_id, [{"container_workload_id", "container_name", "internal_port"}]) + or None when there is nothing to expose. + """ + if method != "POST" or subpath != "workloads/containers" or not json_body: + return None + + containers = json_body.get("containers") or [] + exposures = [ + {"container_name": c.get("container_name"), "internal_port": p.get("internal")} + for c in containers + for p in (c.get("ports") or []) + if p.get("use_dns") + ] + if not exposures: + return None + + try: + data = resp.json().get("data") or {} + except ValueError: + return None + pod_id = data.get("pod_id") + container_ids = data.get("container_ids") or [] + if not pod_id or len(container_ids) != len(containers): + from xcloudify_shared.logging import logger + logger.error("Cannot map use_dns ports to container ids for pod %s -- shape mismatch", pod_id) + return None + + name_to_id = {c.get("container_name"): cid for c, cid in zip(containers, container_ids)} + for e in exposures: + e["container_workload_id"] = name_to_id.get(e["container_name"]) + return pod_id, exposures + + +def deleted_workload(subpath: str, method: str): + """After a delete, ("container" | "pod", id) if it removed one, else None -- + so any public exposure pointing at it can be torn down.""" + if method != "DELETE" or "/" not in subpath: + return None + collection, _, obj_id = subpath.rpartition("/") + if not obj_id: + return None + if collection == "workloads/containers": + return "container", obj_id + if collection == "workloads/pods": + return "pod", obj_id + return None diff --git a/packages/pyshared/xcloudify_shared/logging.py b/packages/pyshared/xcloudify_shared/logging.py new file mode 100644 index 0000000..95b8c5c --- /dev/null +++ b/packages/pyshared/xcloudify_shared/logging.py @@ -0,0 +1,69 @@ +import logging +from colorlog import ColoredFormatter +import os +import threading +from contextvars import ContextVar + +# Context for non-Flask usage (scripts, background jobs, etc.) +_request_id_ctx = ContextVar("request_id", default=None) + +def set_request_id(request_id: str) -> None: + _request_id_ctx.set(request_id) + +def clear_request_id() -> None: + _request_id_ctx.set(None) + +class RequestFilter(logging.Filter): + """Populate record.request_id from ContextVar, Flask, or a process/thread fallback.""" + def filter(self, record): + record.funcName = record.funcName if hasattr(record, 'funcName') else '' + record.request_id = 'no-context' + + try: + ctx_id = _request_id_ctx.get() + if ctx_id: + record.request_id = ctx_id + return True + except Exception: + pass + + try: + from flask import g, has_request_context + if has_request_context() and hasattr(g, 'request_id'): + record.request_id = g.request_id + return True + except Exception: + pass + + try: + record.request_id = f"pid:{os.getpid()}|thr:{threading.current_thread().name}" + except Exception: + pass + + return True + +log_format = ( + "%(log_color)s%(asctime)s - %(levelname)s - RequestID:[%(request_id)s] - %(funcName)s - %(message)s" +) +date_format = "%Y-%m-%d %H:%M:%S" + +formatter = ColoredFormatter( + log_format, + datefmt=date_format, + log_colors={ + "DEBUG": "cyan", + "INFO": "green", + "WARNING": "yellow", + "ERROR": "red", + "CRITICAL": "bold_red", + }, +) + +handler = logging.StreamHandler() +handler.setFormatter(formatter) + +logger = logging.getLogger("xcloudify") +logger.setLevel(logging.DEBUG) +logger.addFilter(RequestFilter()) +logger.addHandler(handler) +logger.propagate = False diff --git a/packages/pyshared/xcloudify_shared/oidc.py b/packages/pyshared/xcloudify_shared/oidc.py new file mode 100644 index 0000000..df31e48 --- /dev/null +++ b/packages/pyshared/xcloudify_shared/oidc.py @@ -0,0 +1,87 @@ +"""OIDC access-token verification. + +Only the part that is identical for any relying party: fetch the issuer's +signing keys, check signature, issuer and audience, return the claims. What a +product does with those claims -- which user row they map to, who is an admin, +what they may reach -- stays in that product. +""" +import threading +from typing import Optional + +import jwt +import requests + +_HTTP_TIMEOUT = 5 +_ALGORITHMS = ["RS256", "RS512", "ES256"] + + +def bearer_token(flask_request) -> Optional[str]: + """The access token on a request: an Authorization bearer, or the header + oauth2-proxy forwards when it holds the session.""" + header = flask_request.headers.get("Authorization", "") + if header.startswith("Bearer "): + token = header.split(" ", 1)[1].strip() + if token: + return token + return flask_request.headers.get("X-Forwarded-Access-Token") or None + + +class OIDCVerifier: + def __init__(self, issuer: str, jwks_url: str = "", audience: str = "", + additional_audiences: str = "", user_agent: str = "xcloudify"): + self.issuer = issuer + self._jwks_url = jwks_url + self._headers = {"User-Agent": user_agent} + auds = [audience] + [a.strip() for a in (additional_audiences or "").split(",")] + self.audiences = [a for a in auds if a] + self._lock = threading.Lock() + self._jwks = None + self._discovery = None + + @classmethod + def from_config(cls, config, user_agent: str) -> "OIDCVerifier": + return cls( + issuer=config.get("OIDC_ISSUER") or "", + jwks_url=config.get("OIDC_JWKS_URL") or "", + audience=config.get("OIDC_AUDIENCE") or "", + additional_audiences=config.get("OIDC_ADDITIONAL_AUDIENCES") or "", + user_agent=user_agent, + ) + + def _jwks_uri(self) -> str: + if self._jwks_url: + return self._jwks_url + if self._discovery is None: + url = self.issuer.rstrip("/") + "/.well-known/openid-configuration" + resp = requests.get(url, timeout=_HTTP_TIMEOUT, headers=self._headers) + resp.raise_for_status() + self._discovery = resp.json() + return self._discovery["jwks_uri"] + + def _signing_key(self, token: str): + kid = jwt.get_unverified_header(token).get("kid") + with self._lock: + if self._jwks is None or kid not in {k.key_id for k in self._jwks.keys}: + resp = requests.get(self._jwks_uri(), timeout=_HTTP_TIMEOUT, headers=self._headers) + resp.raise_for_status() + self._jwks = jwt.PyJWKSet.from_dict(resp.json()) + jwks = self._jwks + for k in jwks.keys: + if k.key_id == kid: + return k.key + raise jwt.exceptions.PyJWKClientError(f"no signing key for kid {kid}") + + def verify(self, token: str) -> dict: + if not self.issuer: + raise jwt.InvalidIssuerError("OIDC_ISSUER is not configured") + claims = jwt.decode( + token, + self._signing_key(token), + algorithms=_ALGORITHMS, + audience=self.audiences or None, + options={"verify_iss": False, "verify_aud": bool(self.audiences)}, + ) + valid = {self.issuer, self.issuer.rstrip("/"), self.issuer.rstrip("/") + "/"} + if claims.get("iss") not in valid: + raise jwt.InvalidIssuerError(claims.get("iss")) + return claims diff --git a/packages/pyshared/xcloudify_shared/responses.py b/packages/pyshared/xcloudify_shared/responses.py new file mode 100644 index 0000000..adc0948 --- /dev/null +++ b/packages/pyshared/xcloudify_shared/responses.py @@ -0,0 +1,35 @@ +from flask import jsonify, g, has_request_context + +ENVELOPE_VERSION = 1 + + +def api_response(*, data=None, success=True, + message: str = "", + status: int = 200, + error_type: str | None = None, + error_details: dict | None = None, + meta: dict | None = None): + """Same envelope shape as core's, so the frontend handles both identically.""" + request_id = getattr(g, "request_id", "n/a") if has_request_context() else "n/a" + + payload: dict = { + "version": ENVELOPE_VERSION, + "success": success, + "code": status, + "message": message, + "request_id": request_id, + } + + if meta: + payload["meta"] = meta + + if success: + if data is not None: + payload["data"] = data + else: + payload["error"] = { + "type": error_type or "UNKNOWN", + "details": error_details or {}, + } + + return jsonify(payload), status diff --git a/packages/pyshared/xcloudify_shared/ssh_keys.py b/packages/pyshared/xcloudify_shared/ssh_keys.py new file mode 100644 index 0000000..2132381 --- /dev/null +++ b/packages/pyshared/xcloudify_shared/ssh_keys.py @@ -0,0 +1,164 @@ +"""The saved SSH key vault every product above core offers. + +Core has no key vault: it accepts raw OpenSSH public key text on VM create +and keeps no record of it. A product that wants a saved-key picker keeps an +`ssh_keys` table of its own (id, user_id, key_name, public_key_data, +key_fingerprint, is_default, created_at, updated_at, deleted, deleted_at, +plus to_json / to_json_with_key) and mounts these routes on it. Its gateway +then turns the ids a VM create carries into key text with resolve_ssh_keys. +""" +import base64 +import hashlib +from datetime import datetime +from typing import Callable, Iterable, Optional + +from flask import request + +from xcloudify_shared.logging import logger +from xcloudify_shared.responses import api_response + +ALLOWED_SSH_KEY_PREFIXES = ( + "ssh-rsa", "ssh-ed25519", "ecdsa-sha2-nistp256", + "ecdsa-sha2-nistp384", "ecdsa-sha2-nistp521", + "sk-ssh-ed25519@openssh.com", "sk-ecdsa-sha2-nistp256@openssh.com", +) + + +def compute_fingerprint(public_key_data: str) -> Optional[str]: + """SHA-256 fingerprint of an OpenSSH public key, e.g. 'SHA256:AbCdEf...'.""" + try: + parts = public_key_data.strip().split() + if len(parts) < 2: + return None + digest = hashlib.sha256(base64.b64decode(parts[1])).digest() + return "SHA256:" + base64.b64encode(digest).decode().rstrip("=") + except Exception: + return None + + +def resolve_ssh_keys(json_body: dict, SSHKey, user_id: Optional[str]) -> None: + """Replace each VM's `public_keys` ids with that key's OpenSSH text, in place. + + Scoped to `user_id`'s own keys, so an id that isn't theirs is silently + dropped rather than resolved -- guessing another user's key id must not + pull its material in. + """ + if user_id is None: + return + for vm in json_body.get("virtual-machines") or []: + key_ids = vm.get("public_keys") + if not key_ids: + continue + if isinstance(key_ids, str): + key_ids = [key_ids] + keys = SSHKey.query.filter( + SSHKey.id.in_(key_ids), SSHKey.user_id == user_id, SSHKey.deleted == False # noqa: E712 + ).all() + by_id = {k.id: k.public_key_data for k in keys} + vm["public_keys"] = [by_id[k] for k in key_ids if k in by_id] + + +def register_ssh_key_routes( + bp, + *, + db, + SSHKey, + current_user_id: Callable[[], str], + decorators: Iterable[Callable] = (), + on_event: Optional[Callable[[str, object, str], None]] = None, +) -> None: + """Mount /ssh-keys on `bp`. + + `decorators` wrap every view (the product's own auth check); `on_event` + is called as (action, key, description) after a create or delete, for a + product that keeps an audit trail. + """ + decorators = list(decorators) + + def _audit(action: str, key, description: str) -> None: + if on_event is None: + return + try: + on_event(action, key, description) + except Exception as exc: + logger.error("Audit logging failed for %s %s: %s", action, key.id, exc) + + def _route(rule: str, methods: list): + def register(fn): + view = fn + for decorator in reversed(decorators): + view = decorator(view) + bp.add_url_rule(rule, endpoint=fn.__name__, view_func=view, methods=methods) + return fn + return register + + def _not_found(): + return api_response(success=False, status=404, message="SSH key not found.", error_type="NOT_FOUND") + + @_route("/ssh-keys", ["GET"]) + def list_ssh_keys(): + keys = ( + SSHKey.query.filter_by(user_id=current_user_id(), deleted=False) + .order_by(SSHKey.created_at.desc()) + .all() + ) + return api_response(data=[k.to_json() for k in keys]) + + @_route("/ssh-keys", ["POST"]) + def create_ssh_key(): + data = request.get_json(silent=True) or {} + key_name = (data.get("key_name") or "").strip() + public_key_data = (data.get("public_key_data") or "").strip() + is_default = bool(data.get("is_default", False)) + + if not key_name: + return api_response(success=False, status=400, message="'key_name' is required.", error_type="VALIDATION_ERROR") + if not public_key_data: + return api_response(success=False, status=400, message="'public_key_data' is required.", error_type="VALIDATION_ERROR") + if not any(public_key_data.startswith(p) for p in ALLOWED_SSH_KEY_PREFIXES): + return api_response( + success=False, status=400, + message="'public_key_data' does not appear to be a valid OpenSSH public key.", + error_type="VALIDATION_ERROR", + ) + + user_id = current_user_id() + fingerprint = compute_fingerprint(public_key_data) + if is_default: + SSHKey.query.filter_by(user_id=user_id, is_default=True, deleted=False).update({"is_default": False}) + + key = SSHKey( + user_id=user_id, + key_name=key_name, + public_key_data=public_key_data, + key_fingerprint=fingerprint, + is_default=is_default, + ) + db.session.add(key) + db.session.commit() + + logger.info("Created SSH key %s (user=%s name=%s)", key.id, user_id, key_name) + _audit("ssh_key_created", key, f"key_name={key_name} fingerprint={fingerprint}") + return api_response(data=key.to_json_with_key(), status=201, message="SSH key added successfully.") + + @_route("/ssh-keys/", ["GET"]) + def get_ssh_key(key_id: str): + key = SSHKey.query.filter_by(id=key_id, user_id=current_user_id(), deleted=False).first() + if key is None: + return _not_found() + return api_response(data=key.to_json_with_key()) + + @_route("/ssh-keys/", ["DELETE"]) + def delete_ssh_key(key_id: str): + user_id = current_user_id() + key = SSHKey.query.filter_by(id=key_id, user_id=user_id, deleted=False).first() + if key is None: + return _not_found() + + key.deleted = True + key.deleted_at = datetime.utcnow() + db.session.commit() + + logger.info("Deleted SSH key %s (user=%s)", key_id, user_id) + _audit("ssh_key_deleted", key, f"key_name={key.key_name}") + return api_response(message="SSH key deleted.", status=200)