Chore: Add xcloudify_shared package

API response envelope, logging, core client, Cloudflare tunnel manager, secret encryption, OIDC token verification, the APP_ENV dev-bypass switch, SSH key routes, and the core gateway allowlist are reusable components
This commit is contained in:
2026-09-01 13:18:08 +05:45
parent efda0f4e8f
commit ca29ba96e3
11 changed files with 934 additions and 0 deletions
+19
View File
@@ -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*"]
@@ -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"]
@@ -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
@@ -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,
)
@@ -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")
+30
View File
@@ -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
@@ -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/<id>/core/..., community:
/leases/<id>/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]
# <collection>/<id>
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>/<id>/<subresource>
collection, obj_id, subresource = parts
if obj_id and subresource in self.subresources.get(collection, ()):
return False
# <collection>/<verb>/<id>
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
@@ -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 '<unknown>'
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
@@ -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
@@ -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
@@ -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/<key_id>", ["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/<key_id>", ["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)