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:
@@ -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")
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user