From a6166d11b94e230bd40c21d32e94ec94510cff8d Mon Sep 17 00:00:00 2001 From: "mustafa.ulukaya" Date: Tue, 4 Aug 2026 21:23:57 +0300 Subject: [PATCH] feat(ssl): add CSR generation and signed-certificate import (backend) New /api/ssl/csrs endpoint group: generate a private key + CSR server-side (RSA 2048/4096, ECDSA P-256/P-384; full subject + DNS SANs with wildcard support), list/detail/delete CSRs, and import the CA-signed certificate. - New ssl_csrs table (SCHEMA_VERSION 9 -> 10, additive + idempotent); the migration re-raises on failure so a failed run is retried instead of being stamped as applied. - Import verifies the certificate against the stored key as a hard gate (match=None is treated as an integrity error, not a lenient pass), rejects malformed and expired certificates with 400, warns on SAN drift, and creates a normal ssl_certificates row (source=csr, cluster_id=NULL, last_config_status=PENDING) so it flows through the standard Apply Management -> agent pull pipeline. - Concurrency: FOR UPDATE row lock serialises double-import and delete-during-import; a partial unique index reserves pending CSR names; soft-deleted same-name certs are reactivated preserving the row id. - Security: no CSR endpoint ever returns the private key (explicit column lists, enforced by a static test); the key copy on the CSR row is NULLed after import; ssl.create/read/delete permissions enforced on every endpoint incl. reads; per-user rate limit on key generation, which runs in a worker thread; csr_id and cluster_ids are int32-guarded. - ssl_service: extract _prepare_cert_fields from create_cert_row (behaviour unchanged, extraction tests untouched) and add stage_ssl_config_versions reusing the exact ssl-{id}-create-{ts} version-name scheme. - Tests: crypto round-trip for all four algorithms, model validation, import-flow unit tests, endpoint auth/permission pinning, migration and key-non-exposure static assertions. --- backend/database/migrations.py | 81 +++- backend/main.py | 2 + backend/models/csr.py | 251 ++++++++++++ backend/routers/csr.py | 375 ++++++++++++++++++ backend/services/csr_service.py | 432 +++++++++++++++++++++ backend/services/ssl_service.py | 152 +++++++- backend/tests/test_csr_import.py | 469 +++++++++++++++++++++++ backend/tests/test_csr_migration.py | 104 +++++ backend/tests/test_csr_models.py | 172 +++++++++ backend/tests/test_csr_router_auth.py | 66 ++++ backend/tests/test_csr_service_crypto.py | 171 +++++++++ 11 files changed, 2260 insertions(+), 15 deletions(-) create mode 100644 backend/models/csr.py create mode 100644 backend/routers/csr.py create mode 100644 backend/services/csr_service.py create mode 100644 backend/tests/test_csr_import.py create mode 100644 backend/tests/test_csr_migration.py create mode 100644 backend/tests/test_csr_models.py create mode 100644 backend/tests/test_csr_router_auth.py create mode 100644 backend/tests/test_csr_service_crypto.py diff --git a/backend/database/migrations.py b/backend/database/migrations.py index c37dfaf..de0f437 100644 --- a/backend/database/migrations.py +++ b/backend/database/migrations.py @@ -1753,7 +1753,12 @@ async def ensure_agent_activity_logs_table(): # bump, already-deployed databases (version >= 8) skip the whole migration run and never gain # the columns, so the frontends SELECT/INSERT would fail. Additive + idempotent + nullable; # existing rows stay NULL and render byte-identical. -SCHEMA_VERSION = 9 +# v1.9.0 (CSR creation): bumped 9 -> 10 for the brand-new `ssl_csrs` table +# (ensure_ssl_csrs_table step). Holds a locally generated private key + CSR PEM +# until the operator imports the CA-signed certificate; the import creates a +# normal ssl_certificates row and NULLs the key copy here. Additive + idempotent; +# no existing table is altered, agents never read this table. +SCHEMA_VERSION = 10 async def run_all_migrations(): @@ -1890,12 +1895,84 @@ async def _run_all_migrations_inner(): await ensure_mfa_columns() # Issue #27 — HA/VIP Keepalived management (v1.7.0): two brand-new tables. - # MUST stay last: FK-references haproxy_cluster_pools/agents/users, all created above. + # MUST run after its FK targets (haproxy_cluster_pools/agents/users), all created above. await ensure_vip_tables() + # v1.9.0 — CSR creation: brand-new ssl_csrs table. FK-references + # ssl_certificates/users, both created above. + await ensure_ssl_csrs_table() + logger.info("Database migrations completed successfully.") +async def ensure_ssl_csrs_table(): + """v1.9.0 — CSR (Certificate Signing Request) creation. Additive only: + one brand-new table (ssl_csrs) + indexes. No ALTER of any existing table, + so the entire current fleet is byte-identical. Fully idempotent + (CREATE TABLE/INDEX IF NOT EXISTS). FK targets (ssl_certificates, users) + are created earlier in the sequence. + + A CSR row holds a locally generated private key + CSR PEM until the + operator imports the CA-signed certificate. The import creates a normal + ssl_certificates row (source='csr', last_config_status='PENDING') and + NULLs the private_key_pem copy here — the key then lives only on the + certificate row, like every other key in the system. Agents never read + this table: the agent SSL delivery endpoint selects from + ssl_certificates only, so a pending CSR can never leak to an agent. + """ + conn = None + try: + conn = await get_database_connection() + + await conn.execute(""" + CREATE TABLE IF NOT EXISTS ssl_csrs ( + id SERIAL PRIMARY KEY, + name VARCHAR(100) NOT NULL, + common_name VARCHAR(253) NOT NULL, + subject JSONB NOT NULL DEFAULT '{}'::jsonb, + sans JSONB NOT NULL DEFAULT '[]'::jsonb, + key_algorithm VARCHAR(20) NOT NULL DEFAULT 'rsa-2048', + csr_pem TEXT NOT NULL, + private_key_pem TEXT, + status VARCHAR(20) NOT NULL DEFAULT 'pending', + ssl_certificate_id INTEGER REFERENCES ssl_certificates(id) ON DELETE SET NULL, + completed_at TIMESTAMP, + created_by INTEGER REFERENCES users(id) ON DELETE SET NULL, + created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, + CONSTRAINT ssl_csrs_status_check CHECK (status IN ('pending', 'completed')) + ); + """) + + # Only PENDING CSRs reserve their name: the name becomes the + # ssl_certificates.name (and thus /etc/ssl/haproxy/{name}.pem on every + # agent) at import time, so two open CSRs must not target the same + # cert name. Completed CSRs are history and may share a name across + # reissues — mirrors the uq_vip_name_active partial-index rationale. + await conn.execute( + "CREATE UNIQUE INDEX IF NOT EXISTS uq_ssl_csrs_name_pending ON ssl_csrs(name) WHERE status = 'pending';" + ) + await conn.execute( + "CREATE INDEX IF NOT EXISTS idx_ssl_csrs_status ON ssl_csrs(status);" + ) + await conn.execute( + "CREATE INDEX IF NOT EXISTS idx_ssl_csrs_cert ON ssl_csrs(ssl_certificate_id);" + ) + + logger.info("ssl_csrs table ensured (v1.9.0 CSR creation)") + except Exception as e: + logger.error(f"Error ensuring ssl_csrs table: {e}") + # Re-raise (ensure_ssl_cluster_junction_table precedent): this step is + # part of the SCHEMA_VERSION=10 bump, and run_all_migrations() records + # the marker only after the inner sequence completes cleanly. Swallowing + # a failure here would stamp version 10 with no ssl_csrs table, and the + # version gate would then skip every future retry — permanently. + raise + finally: + if conn: + await close_database_connection(conn) + + async def ensure_mfa_columns(): """Issue #18 — TOTP MFA (v1.6.0): additive columns on users + 3 new tables. diff --git a/backend/main.py b/backend/main.py index c31e9f5..7c87d9a 100644 --- a/backend/main.py +++ b/backend/main.py @@ -47,6 +47,7 @@ from routers.acme_diagnostics import router as acme_diagnostics_router from routers.site_wizard import router as site_wizard_router from routers.mfa import router as mfa_router from routers.vip import router as vip_router # Issue #27 — HA/VIP (Keepalived) management +from routers.csr import router as csr_router # v1.9.0 — CSR creation (in-app key+CSR generation, signed-cert import) # Production logging configuration from utils.logging_config import setup_production_logging @@ -892,6 +893,7 @@ app.include_router(dashboard_stats_router) # HAProxy stats dashboard app.include_router(agent_router) app.include_router(waf_router) app.include_router(ssl_router) +app.include_router(csr_router) # v1.9.0: CSR creation (in-app key+CSR generation, signed-cert import) app.include_router(security_router) app.include_router(configuration_router) app.include_router(settings_router) diff --git a/backend/models/csr.py b/backend/models/csr.py new file mode 100644 index 0000000..3e81cb2 --- /dev/null +++ b/backend/models/csr.py @@ -0,0 +1,251 @@ +""" +Pydantic models for the CSR (Certificate Signing Request) feature (v1.9.0). + +A CSR row is the precursor of an ssl_certificates row: the backend generates +the private key + CSR locally, the operator has the CSR signed by an external +CA and then imports the signed certificate. The CSR `name` therefore obeys the +exact same path-traversal contract as the SSL certificate name (Bulgu #21) — +at import time it becomes /etc/ssl/haproxy/{name}.pem on every agent and is +shell-processed by the agent script as root. + +The import model deliberately has NO private key field: the key never leaves +the server. It is stored on the ssl_csrs row at generation time and paired +with the signed certificate server-side. +""" + +import re +from typing import List, Optional + +from pydantic import BaseModel, field_validator, model_validator + +KEY_ALGORITHMS = ('rsa-2048', 'rsa-4096', 'ecdsa-p256', 'ecdsa-p384') + +# RFC 1035 LDH hostname, lowercase, optional single leftmost wildcard label. +# Single-label names are allowed (internal CAs routinely sign bare hostnames). +_DNS_NAME_PATTERN = re.compile( + r'^(\*\.)?[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?' + r'(\.[a-z0-9]([a-z0-9-]{0,61}[a-z0-9])?)*$' +) + +# Reject control characters in free-text subject fields: they would be +# persisted, echoed into the UI / issuer column, and printed into agent logs +# via `openssl -subject` output. +_CONTROL_CHARS_PATTERN = re.compile(r'[\x00-\x1f\x7f]') + +_MAX_SANS = 100 +_MAX_CERT_PEM_BYTES = 64 * 1024 # a leaf certificate is ~2 KB; 64 KB is generous +_MAX_CHAIN_PEM_BYTES = 256 * 1024 # agents re-download all cert content every poll + + +def _validate_dns_name(value: str, field_label: str) -> str: + v = (value or '').strip().lower() + if not v: + raise ValueError(f'{field_label} must not be empty') + if len(v) > 253: + raise ValueError(f'{field_label} must be 253 characters or fewer') + if not _DNS_NAME_PATTERN.match(v): + raise ValueError( + f'{field_label} {value!r} is not a valid DNS name — lowercase ' + 'letters, digits, hyphens and dots only; a wildcard is allowed ' + 'only as the leftmost label (e.g. *.example.com).' + ) + return v + + +def _validate_subject_text(value: Optional[str], field_label: str, max_len: int = 64) -> Optional[str]: + if value is None: + return None + v = value.strip() + if not v: + return None + if len(v) > max_len: + raise ValueError(f'{field_label} must be {max_len} characters or fewer') + if _CONTROL_CHARS_PATTERN.search(v): + raise ValueError(f'{field_label} must not contain control characters') + return v + + +def _validate_csr_name(v: str) -> str: + """Mirror of SSLCertificateCreate.validate_name_no_path_traversal (Bulgu #21) + with one deliberate tightening: max length 100, matching the + ssl_certificates.name VARCHAR(100) column (the historical 200-char limit + overflows the column and 500s — not replicated here).""" + if v is None: + raise ValueError('CSR name is required') + stripped = v.strip() + if not stripped: + raise ValueError('CSR name must not be empty') + if stripped != v: + raise ValueError('CSR name must not contain leading/trailing whitespace') + if len(stripped) > 100: + raise ValueError('CSR name must be 100 characters or fewer') + if not re.match(r'^[A-Za-z0-9_.-]+$', stripped): + raise ValueError( + f'CSR name={v!r} contains forbidden characters — only letters, ' + 'digits, underscore, hyphen, and dot are allowed (the name becomes ' + 'a filename component under /etc/ssl/haproxy/ at import).' + ) + if '..' in stripped: + raise ValueError(f'CSR name={v!r} must not contain ".." (path traversal)') + if stripped.startswith('.'): + raise ValueError(f'CSR name={v!r} must not start with "." (hidden filename)') + if stripped.startswith('-'): + raise ValueError(f'CSR name={v!r} must not start with "-" (CLI flag confusion)') + return stripped + + +class SSLCSRCreate(BaseModel): + name: str # becomes the certificate name at import + common_name: str + organization: Optional[str] = None # O + organizational_unit: Optional[str] = None # OU + locality: Optional[str] = None # L + state: Optional[str] = None # ST + country: Optional[str] = None # C — exactly 2 letters + email: Optional[str] = None # emailAddress + sans: List[str] = [] # DNS names; CN is auto-added server-side + key_algorithm: str = 'rsa-2048' + + @field_validator('name') + @classmethod + def validate_name(cls, v): + return _validate_csr_name(v) + + @field_validator('common_name') + @classmethod + def validate_common_name(cls, v): + v = _validate_dns_name(v, 'Common Name') + # RFC 5280 ub-common-name — many CAs reject CNs longer than 64 chars. + if len(v) > 64: + raise ValueError( + 'Common Name must be 64 characters or fewer (RFC 5280 upper ' + 'bound) — put longer names in the SAN list instead.' + ) + return v + + @field_validator('sans') + @classmethod + def validate_sans(cls, v): + if not v: + return [] + if len(v) > _MAX_SANS: + raise ValueError(f'At most {_MAX_SANS} SAN entries are allowed') + seen = set() + result = [] + for entry in v: + normalised = _validate_dns_name(entry, 'SAN entry') + if normalised not in seen: + seen.add(normalised) + result.append(normalised) + return result + + @field_validator('organization') + @classmethod + def validate_organization(cls, v): + return _validate_subject_text(v, 'Organization (O)') + + @field_validator('organizational_unit') + @classmethod + def validate_organizational_unit(cls, v): + return _validate_subject_text(v, 'Organizational Unit (OU)') + + @field_validator('locality') + @classmethod + def validate_locality(cls, v): + return _validate_subject_text(v, 'Locality (L)') + + @field_validator('state') + @classmethod + def validate_state(cls, v): + return _validate_subject_text(v, 'State/Province (ST)') + + @field_validator('country') + @classmethod + def validate_country(cls, v): + # cryptography raises a bare ValueError for a non-2-char COUNTRY_NAME; + # pre-validate so the operator gets a friendly 422 instead of a 500. + if v is None: + return None + v = v.strip() + if not v: + return None + if not re.match(r'^[A-Za-z]{2}$', v): + raise ValueError('Country (C) must be exactly 2 letters (ISO 3166-1 alpha-2, e.g. TR, US)') + return v.upper() + + @field_validator('email') + @classmethod + def validate_email(cls, v): + v = _validate_subject_text(v, 'Email', max_len=254) + if v is not None and ('@' not in v or v.startswith('@') or v.endswith('@')): + raise ValueError('Email must be a valid address (missing or misplaced "@")') + return v + + @field_validator('key_algorithm') + @classmethod + def validate_key_algorithm(cls, v): + if v not in KEY_ALGORITHMS: + raise ValueError( + f'key_algorithm must be one of: {", ".join(KEY_ALGORITHMS)}' + ) + return v + + +class SSLCSRImport(BaseModel): + """Import the CA-signed certificate for a pending CSR. The private key is + NOT part of the request — it is already stored on the CSR row.""" + certificate_content: str # PEM + chain_content: Optional[str] = None # PEM, optional + usage_type: str = 'frontend' # "frontend" or "server" + is_global: bool = False + cluster_ids: Optional[List[int]] = None + # Escape hatch for name collisions that appeared AFTER the CSR was + # created: overrides the CSR's reserved name for the certificate row. + name: Optional[str] = None + + @field_validator('certificate_content') + @classmethod + def validate_certificate(cls, v): + if not v or not v.strip(): + raise ValueError('Certificate content is required') + v = v.strip() + if len(v.encode('utf-8', errors='ignore')) > _MAX_CERT_PEM_BYTES: + raise ValueError('Certificate content exceeds the 64 KB limit') + if '-----BEGIN CERTIFICATE-----' not in v or '-----END CERTIFICATE-----' not in v: + raise ValueError('Certificate must be in PEM format') + return v + + @field_validator('chain_content') + @classmethod + def validate_chain(cls, v): + if v and v.strip(): + v = v.strip() + if len(v.encode('utf-8', errors='ignore')) > _MAX_CHAIN_PEM_BYTES: + raise ValueError('Certificate chain exceeds the 256 KB limit') + if '-----BEGIN CERTIFICATE-----' not in v or '-----END CERTIFICATE-----' not in v: + raise ValueError('Certificate chain must be in PEM format') + return v + return None + + @field_validator('usage_type') + @classmethod + def validate_usage_type(cls, v): + if v not in ['frontend', 'server']: + raise ValueError('usage_type must be either "frontend" or "server"') + return v + + @field_validator('name') + @classmethod + def validate_name(cls, v): + if v is None or not str(v).strip(): + return None + return _validate_csr_name(v) + + @model_validator(mode='after') + def validate_cluster_selection(self): + if not self.is_global and not self.cluster_ids: + raise ValueError( + 'cluster_ids is required when is_global is false — pick at ' + 'least one cluster or import the certificate as global.' + ) + return self diff --git a/backend/routers/csr.py b/backend/routers/csr.py new file mode 100644 index 0000000..a4296d6 --- /dev/null +++ b/backend/routers/csr.py @@ -0,0 +1,375 @@ +""" +CSR (Certificate Signing Request) endpoints (v1.9.0). + +Generate a private key + CSR in-app, download the CSR PEM, have it signed by +an external CA, then import the signed certificate — which creates a normal +ssl_certificates row that flows through the existing pipeline +(config version → Apply Management → agent pull). + +Security posture: +- All endpoints enforce ssl.* permissions explicitly (including the read + endpoints — deliberately stricter than the legacy cert detail route). +- The private key is NEVER returned by any endpoint here; after import it is + reachable only via the existing certificate detail route. +- Key generation is offloaded to a thread (RSA-4096 takes seconds; the + backend runs a single-worker event loop by default) and rate-limited + per user via the user_activity_logs COUNT pattern (acme_diagnostics + precedent — slowapi is not registered on the app). +""" + +import asyncio +import logging +from typing import Optional + +from fastapi import APIRouter, HTTPException, Request, Header + +from database.connection import get_database_connection, close_database_connection +from auth_middleware import get_current_user_from_token, check_user_permission +from models.csr import SSLCSRCreate, SSLCSRImport +from services import csr_service, ssl_service +from routers.ssl import _assert_safe_cert_name, validate_user_cluster_access +from utils.activity_log import log_user_activity + +router = APIRouter(prefix="/api/ssl/csrs", tags=["SSL CSRs"]) +logger = logging.getLogger(__name__) + +_RATE_LIMIT_CREATE_PER_MIN = 10 + +# Columns exposed to the API — private_key_pem is deliberately absent so a +# future `SELECT *` refactor cannot silently start leaking it. +_CSR_LIST_COLUMNS = """ + c.id, c.name, c.common_name, c.subject, c.sans, c.key_algorithm, + c.status, c.ssl_certificate_id, c.completed_at, c.created_at, c.updated_at, + s.name AS certificate_name, u.username AS created_by_username +""" + +_INT32_MAX = 2_147_483_647 + + +def _client_ip(request: Optional[Request]) -> Optional[str]: + try: + return str(request.client.host) if request and request.client else None + except Exception: + return None + + +def _user_agent(request: Optional[Request]) -> Optional[str]: + try: + return request.headers.get("user-agent") if request else None + except Exception: + return None + + +async def _require(authorization: Optional[str], action: str): + """Authenticate + enforce ssl.; returns current_user or raises 401/403.""" + current_user = await get_current_user_from_token(authorization) + ok = await check_user_permission(current_user["id"], "ssl", action, current_user=current_user) + if not ok: + raise HTTPException(status_code=403, detail=f"Insufficient permissions: ssl.{action} required") + return current_user + + +def _assert_int32_id(csr_id: int) -> None: + """ssl_csrs.id is int4 — an out-of-range path param would surface as an + asyncpg DataError 500 (Bulgu #96 precedent); return a clean 404 instead.""" + if csr_id < 1 or csr_id > _INT32_MAX: + raise HTTPException(status_code=404, detail="CSR not found") + + +def _assert_valid_cluster_id(cluster_id: int) -> None: + """Same int4 guard for body-supplied cluster ids: haproxy_clusters.id is + SERIAL/int4, so an out-of-range value would raise asyncpg DataError inside + validate_user_cluster_access and surface as a 500 with the raw driver + error. Fail with the same clean 404 the cluster lookup itself produces.""" + if not isinstance(cluster_id, int) or cluster_id < 1 or cluster_id > _INT32_MAX: + raise HTTPException(status_code=404, detail="Cluster not found") + + +async def _enforce_create_rate_limit(conn, user_id: int) -> None: + """Per-user per-minute limit on key generation, counted against the + csr_create audit-log action (acme_diagnostics _enforce_rate_limit pattern, + backed by the (user_id, action, created_at DESC) composite index).""" + cnt = await conn.fetchval( + """ + SELECT COUNT(*) + FROM user_activity_logs + WHERE user_id = $1 + AND action = 'csr_create' + AND created_at >= NOW() - INTERVAL '60 seconds' + """, + user_id, + ) + if cnt is not None and cnt >= _RATE_LIMIT_CREATE_PER_MIN: + raise HTTPException( + status_code=429, + detail=( + f"Rate limit exceeded: at most {_RATE_LIMIT_CREATE_PER_MIN} " + "CSRs may be created per minute" + ), + ) + + +@router.post("") +async def create_csr(payload: SSLCSRCreate, request: Request, authorization: Optional[str] = Header(None)): + """Generate a private key + CSR. Returns the CSR PEM immediately (so the + UI can show copy/download in one round trip) — never the private key.""" + current_user = await _require(authorization, "create") + conn = None + try: + # Belt and braces on top of the model validator — same duplication + # convention as the certificate create route. + _assert_safe_cert_name(payload.name) + + conn = await get_database_connection() + await _enforce_create_rate_limit(conn, current_user["id"]) + + # Fail fast on a taken name BEFORE burning CPU on key generation; + # insert_csr_row re-checks and the partial unique index closes the race. + await csr_service.assert_csr_name_available(conn, payload.name) + + bundle = await asyncio.to_thread(csr_service.generate_csr_bundle, payload) + csr_id = await csr_service.insert_csr_row(conn, payload, bundle, current_user["id"]) + + row = await conn.fetchrow( + f""" + SELECT {_CSR_LIST_COLUMNS}, c.csr_pem + FROM ssl_csrs c + LEFT JOIN ssl_certificates s ON c.ssl_certificate_id = s.id + LEFT JOIN users u ON c.created_by = u.id + WHERE c.id = $1 + """, + csr_id, + ) + + await log_user_activity( + user_id=current_user["id"], + action='csr_create', + resource_type='ssl_csr', + resource_id=str(csr_id), + details={ + 'csr_name': payload.name, + 'common_name': payload.common_name, + 'sans': bundle['sans'], + 'key_algorithm': payload.key_algorithm, + }, + ip_address=_client_ip(request), + user_agent=_user_agent(request), + ) + + return { + "message": f"CSR '{payload.name}' created successfully", + "csr": csr_service.csr_row_to_dict(row, include_pem=True), + } + except HTTPException: + raise + except Exception as e: + logger.error(f"Error creating CSR: {e}") + raise HTTPException(status_code=500, detail=str(e)) + finally: + if conn: + await close_database_connection(conn) + + +@router.get("") +async def list_csrs(authorization: Optional[str] = Header(None)): + """List CSRs (no PEM payloads — fetch the detail route for the CSR PEM). + Cluster-agnostic: a CSR binds to clusters only at import time.""" + await _require(authorization, "read") + conn = None + try: + conn = await get_database_connection() + rows = await conn.fetch( + f""" + SELECT {_CSR_LIST_COLUMNS} + FROM ssl_csrs c + LEFT JOIN ssl_certificates s ON c.ssl_certificate_id = s.id + LEFT JOIN users u ON c.created_by = u.id + ORDER BY c.created_at DESC + """ + ) + return [csr_service.csr_row_to_dict(r) for r in rows] + except HTTPException: + raise + except Exception as e: + logger.error(f"Error listing CSRs: {e}") + raise HTTPException(status_code=500, detail=str(e)) + finally: + if conn: + await close_database_connection(conn) + + +@router.get("/{csr_id}") +async def get_csr(csr_id: int, authorization: Optional[str] = Header(None)): + """CSR detail including the CSR PEM. The private key is never included.""" + await _require(authorization, "read") + _assert_int32_id(csr_id) + conn = None + try: + conn = await get_database_connection() + row = await conn.fetchrow( + f""" + SELECT {_CSR_LIST_COLUMNS}, c.csr_pem + FROM ssl_csrs c + LEFT JOIN ssl_certificates s ON c.ssl_certificate_id = s.id + LEFT JOIN users u ON c.created_by = u.id + WHERE c.id = $1 + """, + csr_id, + ) + if not row: + raise HTTPException(status_code=404, detail="CSR not found") + return csr_service.csr_row_to_dict(row, include_pem=True) + except HTTPException: + raise + except Exception as e: + logger.error(f"Error fetching CSR {csr_id}: {e}") + raise HTTPException(status_code=500, detail=str(e)) + finally: + if conn: + await close_database_connection(conn) + + +@router.post("/{csr_id}/import") +async def import_csr_certificate( + csr_id: int, + payload: SSLCSRImport, + request: Request, + authorization: Optional[str] = Header(None), +): + """Import the CA-signed certificate for a pending CSR. Creates an + ssl_certificates row (source='csr', PENDING) and stages one config + version per affected cluster — the operator applies manually.""" + current_user = await _require(authorization, "create") + _assert_int32_id(csr_id) + conn = None + try: + if payload.name: + _assert_safe_cert_name(payload.name) + + conn = await get_database_connection() + + if not payload.is_global: + for cluster_id in payload.cluster_ids or []: + _assert_valid_cluster_id(cluster_id) + await validate_user_cluster_access(current_user["id"], cluster_id, conn) + + result = await csr_service.import_signed_certificate( + conn, csr_id, payload, current_user["id"] + ) + cert_id = result["certificate_id"] + + if payload.is_global: + cluster_rows = await conn.fetch( + "SELECT id FROM haproxy_clusters WHERE is_active = TRUE" + ) + affected_clusters = [r['id'] for r in cluster_rows] + else: + affected_clusters = payload.cluster_ids or [] + + # Post-commit staging — a config-generation failure never rolls back + # the certificate (same semantics as the manual create flow). + sync_results = await ssl_service.stage_ssl_config_versions( + conn, cert_id, affected_clusters, action='create', + created_by=current_user["id"], + ) + + await log_user_activity( + user_id=current_user["id"], + action='create', + resource_type='ssl_certificate', + resource_id=str(cert_id), + details={ + 'certificate_name': result['certificate_name'], + 'domain': result.get('primary_domain', 'unknown'), + 'via': 'csr', + 'csr_id': csr_id, + 'usage_type': payload.usage_type, + 'is_global': payload.is_global, + 'cluster_ids': payload.cluster_ids, + 'warnings': result['warnings'], + }, + ip_address=_client_ip(request), + user_agent=_user_agent(request), + ) + await log_user_activity( + user_id=current_user["id"], + action='csr_import', + resource_type='ssl_csr', + resource_id=str(csr_id), + details={ + 'certificate_id': cert_id, + 'certificate_name': result['certificate_name'], + }, + ip_address=_client_ip(request), + user_agent=_user_agent(request), + ) + + return { + "message": ( + f"Certificate '{result['certificate_name']}' imported " + "successfully. Go to Apply Management to deploy." + ), + "certificate_id": cert_id, + "warnings": result["warnings"], + "sync_results": sync_results, + } + except HTTPException: + raise + except Exception as e: + logger.error(f"Error importing signed certificate for CSR {csr_id}: {e}") + raise HTTPException(status_code=500, detail=str(e)) + finally: + if conn: + await close_database_connection(conn) + + +@router.delete("/{csr_id}") +async def delete_csr(csr_id: int, request: Request, authorization: Optional[str] = Header(None)): + """Hard delete. For a pending CSR this permanently destroys the private + key (any certificate later signed from that CSR becomes unusable); for a + completed CSR it only removes history — the imported certificate is not + affected (the FK points csr → cert).""" + current_user = await _require(authorization, "delete") + _assert_int32_id(csr_id) + conn = None + try: + conn = await get_database_connection() + async with conn.transaction(): + # FOR UPDATE serialises against an in-flight import of the same CSR. + row = await conn.fetchrow( + "SELECT id, name, status FROM ssl_csrs WHERE id = $1 FOR UPDATE", + csr_id, + ) + if not row: + raise HTTPException(status_code=404, detail="CSR not found") + await conn.execute("DELETE FROM ssl_csrs WHERE id = $1", csr_id) + + await log_user_activity( + user_id=current_user["id"], + action='delete', + resource_type='ssl_csr', + resource_id=str(csr_id), + details={'csr_name': row['name'], 'status': row['status']}, + ip_address=_client_ip(request), + user_agent=_user_agent(request), + ) + + if row['status'] == 'pending': + message = ( + f"CSR '{row['name']}' deleted — its private key has been " + "permanently destroyed." + ) + else: + message = ( + f"CSR '{row['name']}' deleted (history only) — the imported " + "certificate is not affected." + ) + return {"message": message} + except HTTPException: + raise + except Exception as e: + logger.error(f"Error deleting CSR {csr_id}: {e}") + raise HTTPException(status_code=500, detail=str(e)) + finally: + if conn: + await close_database_connection(conn) diff --git a/backend/services/csr_service.py b/backend/services/csr_service.py new file mode 100644 index 0000000..844f431 --- /dev/null +++ b/backend/services/csr_service.py @@ -0,0 +1,432 @@ +""" +csr_service: CSR (Certificate Signing Request) generation + signed-certificate +import (v1.9.0). + +Flow: + 1. `generate_csr_bundle` builds a private key + CSR locally (pure crypto, + no DB/IO — callers MUST run it via `asyncio.to_thread`: RSA-4096 + generation takes seconds and would stall the single-worker event loop). + 2. The bundle is persisted to `ssl_csrs` (`insert_csr_row`); the operator + downloads the CSR PEM and has it signed by an external CA. + 3. `import_signed_certificate` pairs the CA response with the stored key, + creates a normal `ssl_certificates` row (source='csr', + last_config_status='PENDING' — agents never see it before Apply) and + NULLs the key copy on the CSR row. + +The CSR builder generalises the in-repo ACME reference +(services/acme_service.py finalize_order): PEM output instead of DER, full +subject instead of CN-only, ECDSA support, same PKCS8/NoEncryption key +serialisation (the agent concatenates cert+key+chain into one PEM and HAProxy +cannot read passphrase-protected keys). + +Private keys are stored PLAINTEXT, consistent with every other key in the +system (ssl_certificates.private_key_content, letsencrypt_orders.cert_private_key). +The key is NEVER returned by any CSR API response — `csr_row_to_dict` strips +it unconditionally. +""" + +import json +import logging +from typing import Any, Dict, List, Optional +from types import SimpleNamespace + +import asyncpg +from fastapi import HTTPException + +from cryptography import x509 +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import ec, rsa +from cryptography.x509.oid import NameOID + +from services import ssl_service + +logger = logging.getLogger(__name__) + + +_KEY_FACTORIES = { + 'rsa-2048': lambda: rsa.generate_private_key(public_exponent=65537, key_size=2048), + 'rsa-4096': lambda: rsa.generate_private_key(public_exponent=65537, key_size=4096), + 'ecdsa-p256': lambda: ec.generate_private_key(ec.SECP256R1()), + 'ecdsa-p384': lambda: ec.generate_private_key(ec.SECP384R1()), +} + +# (payload attribute, x509 OID, subject-JSON key) +_SUBJECT_OID_MAP = [ + ('organization', NameOID.ORGANIZATION_NAME, 'O'), + ('organizational_unit', NameOID.ORGANIZATIONAL_UNIT_NAME, 'OU'), + ('locality', NameOID.LOCALITY_NAME, 'L'), + ('state', NameOID.STATE_OR_PROVINCE_NAME, 'ST'), + ('country', NameOID.COUNTRY_NAME, 'C'), + ('email', NameOID.EMAIL_ADDRESS, 'emailAddress'), +] + + +def generate_csr_bundle(payload: Any) -> Dict[str, Any]: + """Generate a private key + CSR for a validated SSLCSRCreate payload. + + Pure CPU-bound crypto — no DB, no network. Callers must offload via + `asyncio.to_thread` (see module docstring). + + Returns {'csr_pem', 'private_key_pem', 'sans', 'subject'}. + """ + key = _KEY_FACTORIES[payload.key_algorithm]() + + attrs = [x509.NameAttribute(NameOID.COMMON_NAME, payload.common_name)] + subject_json: Dict[str, str] = {} + for attr_name, oid, json_key in _SUBJECT_OID_MAP: + value = getattr(payload, attr_name, None) + if value and str(value).strip(): + cleaned = str(value).strip() + attrs.append(x509.NameAttribute(oid, cleaned)) + subject_json[json_key] = cleaned + + # CN always first in the SAN list, then the extra names, deduped with + # order preserved (mirrors the ACME flow where domains[0] is the CN). + sans = list(dict.fromkeys([payload.common_name, *(payload.sans or [])])) + + builder = ( + x509.CertificateSigningRequestBuilder() + .subject_name(x509.Name(attrs)) + .add_extension( + x509.SubjectAlternativeName([x509.DNSName(d) for d in sans]), + critical=False, + ) + ) + csr = builder.sign(key, hashes.SHA256()) + + return { + 'csr_pem': csr.public_bytes(serialization.Encoding.PEM).decode('utf-8'), + 'private_key_pem': key.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ).decode('utf-8'), + 'sans': sans, + 'subject': subject_json, + } + + +def diff_domains(csr_sans: Optional[List[str]], cert_domains: Optional[List[str]]) -> List[str]: + """Human-readable warnings for SAN drift between the CSR and the signed + certificate (case-insensitive set diff). CAs legitimately add/normalise + SANs, so drift is WARN-only — the hard gate is the key match.""" + csr_set = {d.lower() for d in (csr_sans or []) if d} + cert_set = {d.lower() for d in (cert_domains or []) if d} + warnings: List[str] = [] + added = sorted(cert_set - csr_set) + dropped = sorted(csr_set - cert_set) + if added: + warnings.append( + f"The CA added domains that were not in the CSR: {', '.join(added)}" + ) + if dropped: + warnings.append( + f"The CA dropped domains that were requested in the CSR: {', '.join(dropped)}" + ) + return warnings + + +def _maybe_json_list(value: Any) -> List[str]: + """asyncpg returns JSONB columns as str unless a codec is registered.""" + if isinstance(value, str): + try: + parsed = json.loads(value) + return parsed if isinstance(parsed, list) else [] + except Exception: + return [] + return list(value) if value else [] + + +def csr_row_to_dict(row: Any, include_pem: bool = False) -> Dict[str, Any]: + """Row → API dict. ALWAYS strips private_key_pem — the key never leaves + the server via a CSR endpoint. csr_pem included only on demand + (detail/create responses, not lists).""" + d = dict(row) + d.pop('private_key_pem', None) + if not include_pem: + d.pop('csr_pem', None) + for key in ('subject', 'sans'): + if key in d and isinstance(d[key], str): + try: + d[key] = json.loads(d[key]) + except Exception: + pass + return d + + +async def assert_csr_name_available(conn, name: str) -> None: + """Reject a CSR name that is already taken by an ACTIVE certificate or + another PENDING CSR. Called BEFORE key generation (cheap fail-fast) and + re-run inside `insert_csr_row` (the unique index closes the race).""" + existing_cert = await conn.fetchval( + "SELECT id FROM ssl_certificates WHERE name = $1 AND is_active = TRUE", + name, + ) + if existing_cert: + raise HTTPException( + status_code=400, + detail=( + f"An active SSL certificate named '{name}' already exists. " + "The CSR name becomes the certificate name at import — choose " + "a different name or remove the existing certificate first." + ), + ) + existing_csr = await conn.fetchval( + "SELECT id FROM ssl_csrs WHERE name = $1 AND status = 'pending'", + name, + ) + if existing_csr: + raise HTTPException( + status_code=400, + detail=( + f"A pending CSR named '{name}' already exists (id={existing_csr}). " + "Import or delete it first, or choose a different name." + ), + ) + + +async def insert_csr_row(conn, payload: Any, bundle: Dict[str, Any], user_id: Optional[int]) -> int: + """Persist a freshly generated CSR bundle. Returns the new csr id.""" + await assert_csr_name_available(conn, payload.name) + try: + csr_id = await conn.fetchval( + """ + INSERT INTO ssl_csrs + (name, common_name, subject, sans, key_algorithm, csr_pem, + private_key_pem, status, created_by) + VALUES ($1, $2, $3::jsonb, $4::jsonb, $5, $6, $7, 'pending', $8) + RETURNING id + """, + payload.name, + payload.common_name, + json.dumps(bundle['subject']), + json.dumps(bundle['sans']), + payload.key_algorithm, + bundle['csr_pem'], + bundle['private_key_pem'], + user_id, + ) + except asyncpg.exceptions.UniqueViolationError: + # uq_ssl_csrs_name_pending — a concurrent request won the name. + raise HTTPException( + status_code=400, + detail=( + f"A pending CSR named '{payload.name}' was just created by a " + "concurrent request — choose a different name." + ), + ) + return csr_id + + +async def import_signed_certificate(conn, csr_id: int, imp: Any, user_id: Optional[int]) -> Dict[str, Any]: + """Pair the CA-signed certificate with the stored CSR key and create the + ssl_certificates row. Atomic: cert row + CSR state change commit together. + + Returns {'certificate_id', 'certificate_name', 'primary_domain', + 'warnings', 'reactivated'}. Raises HTTPException on every failure + (404 missing, 409 already completed, 400 validation). + """ + async with conn.transaction(): + # Row lock serialises concurrent imports AND a concurrent DELETE of + # the same CSR; works across multiple uvicorn workers (DB-level lock). + row = await conn.fetchrow( + "SELECT * FROM ssl_csrs WHERE id = $1 FOR UPDATE", csr_id + ) + if not row: + raise HTTPException(status_code=404, detail="CSR not found") + if row['status'] == 'completed': + raise HTTPException( + status_code=409, + detail=( + f"CSR '{row['name']}' is already completed — certificate " + f"id {row['ssl_certificate_id']} was imported from it. " + "Create a new CSR to reissue." + ), + ) + stored_key = row['private_key_pem'] + if not stored_key: + raise HTTPException( + status_code=500, + detail=( + "Stored CSR private key is missing — the CSR row is " + "corrupt. Delete it and create a new CSR." + ), + ) + + effective_name = getattr(imp, 'name', None) or row['name'] + + # Parse the pasted certificate FIRST so a malformed/truncated CA + # response gets the manual flow's 400, not a 500 from the key-match + # step below (verify_certificate_key_match reports an unparseable + # cert as match=None, which we treat as an integrity failure). + from utils.ssl_parser import parse_ssl_certificate, verify_certificate_key_match + precheck = parse_ssl_certificate(imp.certificate_content) + if precheck.get('error'): + raise HTTPException( + status_code=400, + detail=f"Invalid SSL certificate: {precheck['error']}", + ) + + # THE defining check of this feature: the CA response must match the + # key we generated. Deliberately stricter than create_cert_row's + # lenient fallback — we generated this key ourselves, so an + # unverifiable pair is an integrity failure, not operator input. + match_result = verify_certificate_key_match(imp.certificate_content, stored_key) + if match_result.get('match') is False: + raise HTTPException( + status_code=400, + detail=( + "The signed certificate does not match this CSR's private " + "key — the CA response likely belongs to a different " + "CSR/key. Verify you pasted the certificate that was " + "issued for this exact CSR." + ), + ) + if match_result.get('match') is not True: + raise HTTPException( + status_code=500, + detail=( + "Could not verify the certificate/key pair: " + f"{match_result.get('reason', 'unknown')}" + ), + ) + + # Full parse/validation pipeline shared with the manual + wizard + # flows: invalid PEM, bad chain and already-expired certs all 400. + payload = SimpleNamespace( + name=effective_name, + certificate_content=imp.certificate_content, + private_key_content=stored_key, + chain_content=getattr(imp, 'chain_content', None), + usage_type=getattr(imp, 'usage_type', 'frontend') or 'frontend', + ) + fields = ssl_service._prepare_cert_fields(payload) + + # Global name uniqueness (ssl_certificates.cluster_id is always NULL + # under the R38 schema, so name is effectively a global namespace). + existing = await conn.fetchrow( + "SELECT id, is_active FROM ssl_certificates WHERE name = $1 LIMIT 1", + effective_name, + ) + if existing and existing['is_active']: + raise HTTPException( + status_code=400, + detail=( + f"An active SSL certificate named '{effective_name}' " + "already exists (created after this CSR). Delete or " + "rename it, or pass a different `name` in the import " + "request — the CSR stays pending and can be re-imported." + ), + ) + + reactivated = False + if existing and not existing['is_active']: + # Reactivate the soft-deleted row (mirrors create_cert_row): + # preserves the row id so historical references keep working. + await conn.execute( + "DELETE FROM ssl_certificate_clusters WHERE ssl_certificate_id = $1", + existing['id'], + ) + await conn.execute( + """ + UPDATE ssl_certificates + SET is_active = TRUE, + last_config_status = 'PENDING', + certificate_content = $2, + private_key_content = $3, + chain_content = $4, + primary_domain = $5, + all_domains = $6::jsonb, + expiry_date = $7, + usage_type = $8, + issuer = $9, + fingerprint = $10, + status = $11, + days_until_expiry = $12, + source = 'csr', + updated_at = CURRENT_TIMESTAMP + WHERE id = $1 + """, + existing['id'], + fields['cert_content'], + fields['private_key_content'], + fields['chain_content'], + fields['primary_domain'], + json.dumps(fields['all_domains']), + fields['expiry_date'], + fields['usage_type'], + fields['issuer'], + fields['fingerprint'], + fields['status'], + fields['days_until_expiry'], + ) + cert_id = existing['id'] + reactivated = True + logger.info( + f"csr_service.import_signed_certificate: reactivated " + f"soft-deleted cert '{effective_name}' (id={cert_id}) for CSR {csr_id}" + ) + else: + cert_id = await conn.fetchval( + """ + INSERT INTO ssl_certificates ( + name, primary_domain, certificate_content, private_key_content, + chain_content, expiry_date, issuer, fingerprint, status, + days_until_expiry, all_domains, is_active, cluster_id, + last_config_status, usage_type, source + ) VALUES ( + $1, $2, $3, $4, $5, $6, $7, $8, $9, $10, $11::jsonb, + TRUE, NULL, 'PENDING', $12, 'csr' + ) + RETURNING id + """, + effective_name, + fields['primary_domain'], + fields['cert_content'], + fields['private_key_content'], + fields['chain_content'], + fields['expiry_date'], + fields['issuer'], + fields['fingerprint'], + fields['status'], + fields['days_until_expiry'], + json.dumps(fields['all_domains']), + fields['usage_type'], + ) + + # Cluster bindings: global = zero junction rows (existing convention). + if not getattr(imp, 'is_global', False): + for cluster_id in (getattr(imp, 'cluster_ids', None) or []): + await ssl_service.ensure_cluster_junction(conn, cert_id, cluster_id) + + # Complete the CSR and destroy the key copy — the key now lives on + # the certificate row only, like every other key in the system. + await conn.execute( + """ + UPDATE ssl_csrs + SET status = 'completed', + ssl_certificate_id = $2, + private_key_pem = NULL, + completed_at = CURRENT_TIMESTAMP, + updated_at = CURRENT_TIMESTAMP + WHERE id = $1 + """, + csr_id, + cert_id, + ) + + warnings = diff_domains(_maybe_json_list(row['sans']), fields['all_domains']) + if reactivated: + warnings.append( + f"A soft-deleted certificate named '{effective_name}' was " + f"reactivated (row id {cert_id}) — existing entities that still " + "reference that id now serve the newly imported certificate." + ) + + return { + 'certificate_id': cert_id, + 'certificate_name': effective_name, + 'primary_domain': fields['primary_domain'], + 'warnings': warnings, + 'reactivated': reactivated, + } diff --git a/backend/services/ssl_service.py b/backend/services/ssl_service.py index 1a74cec..2cfd7fe 100644 --- a/backend/services/ssl_service.py +++ b/backend/services/ssl_service.py @@ -40,10 +40,12 @@ flow. (callers translate to wizard step-jumpback toasts). """ +import hashlib import json import logging +import time from datetime import datetime, timezone -from typing import Any, Optional +from typing import Any, List, Optional from fastapi import HTTPException @@ -106,23 +108,22 @@ def _recompute_status_from_expiry( return cert_info_status or "valid", cert_info_days or 0 -async def create_cert_row( - conn, - payload: Any, - cluster_id: int, -) -> int: - """Insert a row into ssl_certificates (always cluster_id=NULL) + junction - binding to the given cluster_id. Returns new ssl_certificate_id. +def _prepare_cert_fields(payload: Any) -> dict: + """Parse + validate the PEM material on `payload` and derive every + ssl_certificates column value from it (v1.9.0 extraction — shared by + `create_cert_row` and the CSR import flow in services/csr_service.py, + byte-identical to the former inline body of `create_cert_row`). payload is expected to expose: name, certificate_content, private_key_content, chain_content, usage_type (optional, default 'frontend'). - All cert metadata (primary_domain, all_domains, expiry_date, - issuer, fingerprint, status, days_until_expiry) is now parsed - FROM the PEM content via `parse_ssl_certificate` — operator- - supplied values on the payload are accepted as a graceful - fallback only when parsing fails (which itself raises 400). + Raises HTTPException(400) on any parse/validation failure (invalid PEM, + bad private key, cert/key mismatch, bad chain, already-expired cert). + + Returns a dict with keys: cert_content, private_key_content, + chain_content, cert_info, primary_domain, all_domains, expiry_date, + issuer, fingerprint, status, days_until_expiry, usage_type. """ cert_content = getattr(payload, "certificate_content", None) or "" if not cert_content.strip(): @@ -213,6 +214,53 @@ async def create_cert_row( ) usage_type = getattr(payload, "usage_type", "frontend") or "frontend" + return { + "cert_content": cert_content, + "private_key_content": private_key_content, + "chain_content": chain_content, + "cert_info": cert_info, + "primary_domain": primary_domain, + "all_domains": all_domains, + "expiry_date": expiry_date, + "issuer": issuer, + "fingerprint": fingerprint, + "status": status, + "days_until_expiry": days_until_expiry, + "usage_type": usage_type, + } + + +async def create_cert_row( + conn, + payload: Any, + cluster_id: int, +) -> int: + """Insert a row into ssl_certificates (always cluster_id=NULL) + junction + binding to the given cluster_id. Returns new ssl_certificate_id. + + payload is expected to expose: + name, certificate_content, private_key_content, chain_content, + usage_type (optional, default 'frontend'). + + All cert metadata (primary_domain, all_domains, expiry_date, + issuer, fingerprint, status, days_until_expiry) is now parsed + FROM the PEM content via `parse_ssl_certificate` — operator- + supplied values on the payload are accepted as a graceful + fallback only when parsing fails (which itself raises 400). + """ + fields = _prepare_cert_fields(payload) + cert_content = fields["cert_content"] + private_key_content = fields["private_key_content"] + chain_content = fields["chain_content"] + expiry_date = fields["expiry_date"] + primary_domain = fields["primary_domain"] + all_domains = fields["all_domains"] + issuer = fields["issuer"] + fingerprint = fields["fingerprint"] + status = fields["status"] + days_until_expiry = fields["days_until_expiry"] + usage_type = fields["usage_type"] + existing = await conn.fetchrow( """ SELECT s.id, s.is_active @@ -408,3 +456,81 @@ async def validate_server_ca_bundle_eligibility( cluster_id, ) return row is not None + + +async def stage_ssl_config_versions( + conn, + cert_id: int, + cluster_ids: List[int], + action: str = "create", + created_by: Optional[int] = None, +) -> List[dict]: + """Stage one PENDING config version per affected cluster after an SSL + certificate mutation (v1.9.0 — distilled from the routers/ssl.py POST + /certificates staging loop; used by the CSR import flow). + + Uses the EXACT `ssl-{cert_id}-{action}-{timestamp}` version-name scheme of + the manual SSL flow so Apply Management, the `has_pending_config` + LIKE-filter ('ssl-' || id || '-%'), and the agent delivery predicates + treat CSR-imported certificates identically to manually uploaded ones. + Agents are NOT notified here — the operator applies manually. + + Per-cluster failures are caught and reported in the returned + sync_results list (the DB save has already succeeded — same semantics as + the manual flow, where a config-generation failure never rolls back the + certificate row). + """ + # Local import: keeps services/haproxy_config free to import ssl helpers + # without a module-level cycle. + from services.haproxy_config import generate_haproxy_config_for_cluster + + sync_results: List[dict] = [] + for cluster_id in cluster_ids: + try: + config_content = await generate_haproxy_config_for_cluster(cluster_id) + config_hash = hashlib.sha256(config_content.encode()).hexdigest() + version_name = f"ssl-{cert_id}-{action}-{int(time.time())}" + + version_created_by = created_by + if version_created_by is None: + version_created_by = await conn.fetchval( + "SELECT id FROM users WHERE username = 'admin' LIMIT 1" + ) or 1 + + await conn.fetchval( + """ + INSERT INTO config_versions + (cluster_id, version_name, config_content, checksum, created_by, is_active, status) + VALUES ($1, $2, $3, $4, $5, FALSE, 'PENDING') + RETURNING id + """, + cluster_id, + version_name, + config_content, + config_hash, + version_created_by, + ) + logger.info( + f"APPLY WORKFLOW: Created PENDING config version {version_name} " + f"for cluster {cluster_id} (ssl_service.stage_ssl_config_versions)" + ) + sync_results.append({ + 'node': 'pending', + 'success': True, + 'cluster_id': cluster_id, + 'version': version_name, + 'status': 'PENDING', + 'message': 'SSL certificate staged. Click Apply to activate.', + }) + except Exception as e: + logger.error( + f"Cluster config staging failed for SSL certificate {cert_id} " + f"on cluster {cluster_id}: {e}" + ) + sync_results.append({ + 'node': 'cluster', + 'success': False, + 'cluster_id': cluster_id, + 'error': str(e), + }) + return sync_results diff --git a/backend/tests/test_csr_import.py b/backend/tests/test_csr_import.py new file mode 100644 index 0000000..57003af --- /dev/null +++ b/backend/tests/test_csr_import.py @@ -0,0 +1,469 @@ +""" +v1.9.0 CSR creation — unit tests for the signed-certificate import flow and +config-version staging (pattern: test_ssl_service_extraction.py, AsyncMock conn). + +Pins the security-relevant invariants: +- key match is a HARD gate: match=False → 400 before any INSERT, and + match=None (unverifiable) → 500, never a lenient pass (we generated the + key ourselves — deliberate divergence from create_cert_row's fallback). +- the new cert row is cluster_id=NULL / last_config_status='PENDING' / + source='csr' (PENDING keeps it invisible to agents until Apply). +- completing the CSR NULLs the private key copy. +- staging reuses the exact `ssl-{id}-create-{ts}` version-name scheme. +""" +import json +from contextlib import contextmanager +from datetime import datetime, timezone +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest +from fastapi import HTTPException + +from models.csr import SSLCSRImport +from services.csr_service import ( + assert_csr_name_available, + import_signed_certificate, + insert_csr_row, +) +from services.ssl_service import stage_ssl_config_versions + + +_VALID_PARSE = { + "primary_domain": "www.example.com", + "all_domains": ["www.example.com"], + "expiry_date": datetime(2099, 1, 1, tzinfo=timezone.utc), + "issuer": "CN=Test CA", + "fingerprint": "AA:BB:CC", + "status": "valid", + "days_until_expiry": 365, +} + +_FAKE_CERT = "-----BEGIN CERTIFICATE-----\nX\n-----END CERTIFICATE-----" +_FAKE_KEY = "-----BEGIN PRIVATE KEY-----\nY\n-----END PRIVATE KEY-----" + + +def _csr_row(**overrides): + row = { + "id": 5, + "name": "csr-www", + "common_name": "www.example.com", + "subject": "{}", + "sans": json.dumps(["www.example.com"]), + "key_algorithm": "rsa-2048", + "csr_pem": "-----BEGIN CERTIFICATE REQUEST-----\nZ\n-----END CERTIFICATE REQUEST-----", + "private_key_pem": _FAKE_KEY, + "status": "pending", + "ssl_certificate_id": None, + } + row.update(overrides) + return row + + +def _mk_conn(): + conn = AsyncMock() + # asyncpg's conn.transaction() is a SYNC call returning an async CM. + conn.transaction = MagicMock() + return conn + + +def _import_payload(**overrides): + base = dict( + certificate_content=_FAKE_CERT, + chain_content=None, + usage_type="frontend", + is_global=False, + cluster_ids=[1, 2], + name=None, + ) + base.update(overrides) + return SSLCSRImport(**base) + + +@contextmanager +def _patched(match=None, parse=None): + """Patch every parser touchpoint of the import path: the function-local + imports in csr_service (utils.ssl_parser.*) and the module-level imports + in ssl_service._prepare_cert_fields (services.ssl_service.*).""" + match_result = match if match is not None else {"match": True} + parse_result = dict(parse or _VALID_PARSE) + with patch("utils.ssl_parser.verify_certificate_key_match", return_value=match_result), \ + patch("utils.ssl_parser.parse_ssl_certificate", return_value=dict(parse_result)), \ + patch("services.ssl_service.parse_ssl_certificate", return_value=dict(parse_result)), \ + patch("services.ssl_service.validate_private_key", return_value=True), \ + patch("services.ssl_service.validate_certificate_chain", return_value=True): + yield + + +# ---------------------------------------------------------------------------- +# import_signed_certificate +# ---------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_import_happy_path_inserts_pending_csr_sourced_cert(): + conn = _mk_conn() + conn.fetchrow.side_effect = [_csr_row(), None] # FOR UPDATE row, no name clash + conn.fetchval.return_value = 42 # INSERT ... RETURNING id + + with _patched(): + result = await import_signed_certificate(conn, 5, _import_payload(), user_id=7) + + assert result["certificate_id"] == 42 + assert result["reactivated"] is False + + # Concurrency invariants: everything runs inside a transaction and the + # CSR row is locked FOR UPDATE (serialises double-import and delete-races). + assert conn.transaction.call_count == 1 + lock_sql = conn.fetchrow.call_args_list[0].args[0] + assert "FOR UPDATE" in lock_sql + + insert_sql, *insert_args = conn.fetchval.call_args.args + assert "INSERT INTO ssl_certificates" in insert_sql + assert "NULL, 'PENDING'" in insert_sql, "cert must stay invisible to agents until Apply" + assert "'csr'" in insert_sql, "source column must record the CSR origin" + # The stored CSR key — not any request-supplied key — must be persisted. + assert _FAKE_KEY in insert_args + + # One junction row per requested cluster. + junction_calls = [ + c for c in conn.execute.call_args_list + if c.args and "ssl_certificate_clusters" in c.args[0] and "INSERT" in c.args[0] + ] + assert len(junction_calls) == 2 + assert {c.args[2] for c in junction_calls} == {1, 2} + + # CSR completion must destroy the key copy. + completion_calls = [ + c for c in conn.execute.call_args_list + if c.args and "UPDATE ssl_csrs" in c.args[0] + ] + assert len(completion_calls) == 1 + assert "private_key_pem = NULL" in completion_calls[0].args[0] + assert "status = 'completed'" in completion_calls[0].args[0] + assert completion_calls[0].args[1] == 5 # csr_id + assert completion_calls[0].args[2] == 42 # cert_id + + +@pytest.mark.asyncio +async def test_import_global_creates_zero_junction_rows(): + conn = _mk_conn() + conn.fetchrow.side_effect = [_csr_row(), None] + conn.fetchval.return_value = 42 + + with _patched(): + await import_signed_certificate( + conn, 5, _import_payload(is_global=True, cluster_ids=None), user_id=7 + ) + + junction_calls = [ + c for c in conn.execute.call_args_list + if c.args and "ssl_certificate_clusters" in c.args[0] and "INSERT" in c.args[0] + ] + assert junction_calls == [], "global cert = zero junction rows (existing convention)" + + +@pytest.mark.asyncio +async def test_import_key_mismatch_rejected_400_before_any_write(): + conn = _mk_conn() + conn.fetchrow.side_effect = [_csr_row()] + + with _patched(match={"match": False, "reason": "public key mismatch"}): + with pytest.raises(HTTPException) as exc_info: + await import_signed_certificate(conn, 5, _import_payload(), user_id=7) + + assert exc_info.value.status_code == 400 + assert "does not match" in exc_info.value.detail + assert not conn.fetchval.await_count, "nothing must be inserted on mismatch" + assert not conn.execute.await_count + + +@pytest.mark.asyncio +async def test_import_unverifiable_key_match_is_hard_error_not_lenient(): + """match=None means OUR stored key is unreadable — integrity failure, + never the lenient pass create_cert_row historically allows.""" + conn = _mk_conn() + conn.fetchrow.side_effect = [_csr_row()] + + with _patched(match={"match": None, "reason": "key could not be parsed"}): + with pytest.raises(HTTPException) as exc_info: + await import_signed_certificate(conn, 5, _import_payload(), user_id=7) + + assert exc_info.value.status_code == 500 + assert not conn.fetchval.await_count + + +@pytest.mark.asyncio +async def test_import_expired_certificate_rejected_400(): + conn = _mk_conn() + conn.fetchrow.side_effect = [_csr_row()] + + expired = dict(_VALID_PARSE) + expired["status"] = "expired" + expired["days_until_expiry"] = -10 + with _patched(parse=expired): + with pytest.raises(HTTPException) as exc_info: + await import_signed_certificate(conn, 5, _import_payload(), user_id=7) + + assert exc_info.value.status_code == 400 + assert "expired" in exc_info.value.detail.lower() + assert not conn.fetchval.await_count + + +@pytest.mark.asyncio +async def test_import_malformed_certificate_rejected_400_not_500(): + """A cert with PEM markers but unparseable content (truncated CA response) + is OPERATOR INPUT — it must get the manual flow's 400, not the 500 that + the strict key-match branch reserves for a corrupt STORED key.""" + conn = _mk_conn() + conn.fetchrow.side_effect = [_csr_row()] + + with _patched(parse={"error": "Could not parse certificate"}): + with pytest.raises(HTTPException) as exc_info: + await import_signed_certificate(conn, 5, _import_payload(), user_id=7) + + assert exc_info.value.status_code == 400 + assert "Invalid SSL certificate" in exc_info.value.detail + assert not conn.fetchval.await_count + assert not conn.execute.await_count + + +@pytest.mark.asyncio +async def test_import_completed_csr_conflicts_409(): + conn = _mk_conn() + conn.fetchrow.side_effect = [_csr_row(status="completed", ssl_certificate_id=42)] + + with pytest.raises(HTTPException) as exc_info: + await import_signed_certificate(conn, 5, _import_payload(), user_id=7) + + assert exc_info.value.status_code == 409 + assert "already completed" in exc_info.value.detail + + +@pytest.mark.asyncio +async def test_import_missing_csr_404(): + conn = _mk_conn() + conn.fetchrow.side_effect = [None] + + with pytest.raises(HTTPException) as exc_info: + await import_signed_certificate(conn, 999, _import_payload(), user_id=7) + + assert exc_info.value.status_code == 404 + + +@pytest.mark.asyncio +async def test_import_active_name_collision_rejected_with_hint(): + conn = _mk_conn() + conn.fetchrow.side_effect = [_csr_row(), {"id": 9, "is_active": True}] + + with _patched(): + with pytest.raises(HTTPException) as exc_info: + await import_signed_certificate(conn, 5, _import_payload(), user_id=7) + + assert exc_info.value.status_code == 400 + assert "already exists" in exc_info.value.detail + assert "name" in exc_info.value.detail # points at the override escape hatch + assert not conn.fetchval.await_count + + +@pytest.mark.asyncio +async def test_import_name_override_is_used_for_the_cert_row(): + conn = _mk_conn() + conn.fetchrow.side_effect = [_csr_row(), None] + conn.fetchval.return_value = 42 + + with _patched(): + result = await import_signed_certificate( + conn, 5, _import_payload(name="renamed-cert"), user_id=7 + ) + + assert result["certificate_name"] == "renamed-cert" + _, *insert_args = conn.fetchval.call_args.args + assert "renamed-cert" in insert_args + # And the collision check must have run against the override, not csr.name. + name_lookup = conn.fetchrow.call_args_list[1] + assert name_lookup.args[1] == "renamed-cert" + + +@pytest.mark.asyncio +async def test_import_reactivates_soft_deleted_name_and_warns(): + conn = _mk_conn() + conn.fetchrow.side_effect = [_csr_row(), {"id": 77, "is_active": False}] + + with _patched(): + result = await import_signed_certificate(conn, 5, _import_payload(), user_id=7) + + assert result["certificate_id"] == 77 + assert result["reactivated"] is True + assert any("reactivated" in w for w in result["warnings"]) + assert not conn.fetchval.await_count, "reactivation must UPDATE, not INSERT" + update_calls = [ + c for c in conn.execute.call_args_list + if c.args and "UPDATE ssl_certificates" in c.args[0] + ] + assert len(update_calls) == 1 + update_sql = update_calls[0].args[0] + assert "source = 'csr'" in update_sql + # The reactivated row must come back to life invisible to agents until + # Apply, with the row itself active again. + assert "last_config_status = 'PENDING'" in update_sql + assert "is_active = TRUE" in update_sql + # Old cluster bindings must be wiped before re-binding to the new scope. + junction_deletes = [ + c for c in conn.execute.call_args_list + if c.args and "DELETE FROM ssl_certificate_clusters" in c.args[0] + ] + assert len(junction_deletes) == 1 + assert junction_deletes[0].args[1] == 77 + # …and the importer's requested clusters re-bound via the junction. + junction_inserts = [ + c for c in conn.execute.call_args_list + if c.args and "INSERT INTO ssl_certificate_clusters" in c.args[0] + ] + assert {c.args[2] for c in junction_inserts} == {1, 2} + + +@pytest.mark.asyncio +async def test_import_san_drift_warns_but_succeeds(): + conn = _mk_conn() + conn.fetchrow.side_effect = [ + _csr_row(sans=json.dumps(["www.example.com", "api.example.com"])), + None, + ] + conn.fetchval.return_value = 42 + + drifted = dict(_VALID_PARSE) + drifted["all_domains"] = ["www.example.com", "cdn.example.com"] + with _patched(parse=drifted): + result = await import_signed_certificate(conn, 5, _import_payload(), user_id=7) + + assert result["certificate_id"] == 42 + assert any("added" in w and "cdn.example.com" in w for w in result["warnings"]) + assert any("dropped" in w and "api.example.com" in w for w in result["warnings"]) + + +# ---------------------------------------------------------------------------- +# insert_csr_row / assert_csr_name_available +# ---------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_csr_name_taken_by_active_cert_rejected(): + conn = _mk_conn() + conn.fetchval.side_effect = [11] # active cert with the name exists + + with pytest.raises(HTTPException) as exc_info: + await assert_csr_name_available(conn, "taken") + assert exc_info.value.status_code == 400 + assert "certificate" in exc_info.value.detail.lower() + + +@pytest.mark.asyncio +async def test_csr_name_taken_by_pending_csr_rejected(): + conn = _mk_conn() + conn.fetchval.side_effect = [None, 12] # no cert, but a pending CSR + + with pytest.raises(HTTPException) as exc_info: + await assert_csr_name_available(conn, "taken") + assert exc_info.value.status_code == 400 + assert "pending CSR" in exc_info.value.detail + + +@pytest.mark.asyncio +async def test_insert_csr_row_translates_unique_violation_to_400(): + """The uq_ssl_csrs_name_pending partial index closes the create/create + race — the loser must get a clean 400, not a 500.""" + import asyncpg as _asyncpg + + conn = _mk_conn() + # availability checks pass, INSERT hits the unique index + conn.fetchval.side_effect = [ + None, None, _asyncpg.exceptions.UniqueViolationError("dup"), + ] + payload = SimpleNamespace( + name="raced", common_name="www.example.com", key_algorithm="rsa-2048" + ) + bundle = {"subject": {}, "sans": ["www.example.com"], "csr_pem": "PEM", "private_key_pem": "KEY"} + + with pytest.raises(HTTPException) as exc_info: + await insert_csr_row(conn, payload, bundle, user_id=1) + assert exc_info.value.status_code == 400 + assert "concurrent" in exc_info.value.detail + + +# ---------------------------------------------------------------------------- +# router-level guards +# ---------------------------------------------------------------------------- + + +def test_cluster_id_int32_guard_rejects_out_of_range_with_404(): + """Body-supplied cluster ids must never reach asyncpg out of int4 range + (DataError → raw 500) — same Bulgu #96 hygiene as the csr_id path param.""" + from routers.csr import _assert_valid_cluster_id + + _assert_valid_cluster_id(1) + _assert_valid_cluster_id(2_147_483_647) + for bad in (0, -1, 2_147_483_648, 99_999_999_999): + with pytest.raises(HTTPException) as exc_info: + _assert_valid_cluster_id(bad) + assert exc_info.value.status_code == 404 + assert "Cluster not found" in exc_info.value.detail + + +# ---------------------------------------------------------------------------- +# stage_ssl_config_versions +# ---------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_stage_creates_one_pending_version_per_cluster_with_ssl_naming(): + import re + + conn = _mk_conn() + conn.fetchval.return_value = 1001 # config_versions INSERT RETURNING id + + with patch( + "services.haproxy_config.generate_haproxy_config_for_cluster", + new=AsyncMock(return_value="# cfg"), + ): + results = await stage_ssl_config_versions(conn, 42, [1, 2], created_by=7) + + assert len(results) == 2 + assert all(r["success"] for r in results) + assert [r["cluster_id"] for r in results] == [1, 2] + + insert_calls = [ + c for c in conn.fetchval.call_args_list + if c.args and "INSERT INTO config_versions" in c.args[0] + ] + assert len(insert_calls) == 2 + for call in insert_calls: + sql = call.args[0] + assert "FALSE, 'PENDING'" in sql, "staged versions must be inactive + PENDING" + version_name = call.args[2] + # EXACT manual-flow scheme: Apply Management + has_pending_config + # LIKE-filters key off 'ssl-{id}-...'. + assert re.match(r"^ssl-42-create-\d+$", version_name), version_name + assert call.args[5] == 7 # created_by honours the importing user + + +@pytest.mark.asyncio +async def test_stage_reports_per_cluster_failure_without_raising(): + conn = _mk_conn() + conn.fetchval.return_value = 1001 + + async def _gen(cluster_id): + if cluster_id == 2: + raise RuntimeError("config generation exploded") + return "# cfg" + + with patch( + "services.haproxy_config.generate_haproxy_config_for_cluster", + new=AsyncMock(side_effect=_gen), + ): + results = await stage_ssl_config_versions(conn, 42, [1, 2], created_by=7) + + assert len(results) == 2 + assert results[0]["success"] is True + assert results[1]["success"] is False + assert "exploded" in results[1]["error"] diff --git a/backend/tests/test_csr_migration.py b/backend/tests/test_csr_migration.py new file mode 100644 index 0000000..3ae2a84 --- /dev/null +++ b/backend/tests/test_csr_migration.py @@ -0,0 +1,104 @@ +""" +v1.9.0 CSR creation — static source assertions (pattern: test_vip_purge.py). + +Guards the migration wiring that a unit test cannot exercise without a real +database: the SCHEMA_VERSION bump (without it, deployed installs skip the +whole migration run and the ssl_csrs table never appears), the migration +registration, the security-relevant DDL, and the router registration. +""" +import os +import re + +_BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + + +def _read(rel_path: str) -> str: + with open(os.path.join(_BACKEND_DIR, rel_path), encoding="utf-8") as f: + return f.read() + + +def test_schema_version_bumped_to_10(): + src = _read(os.path.join("database", "migrations.py")) + m = re.search(r"^SCHEMA_VERSION\s*=\s*(\d+)", src, re.MULTILINE) + assert m, "SCHEMA_VERSION constant not found in migrations.py" + assert int(m.group(1)) >= 10, ( + "SCHEMA_VERSION must be >= 10 for the v1.9.0 ssl_csrs table — " + "without the bump, existing installs (version >= 9) skip the whole " + "migration run and never gain the table." + ) + + +def test_ssl_csrs_migration_defined_and_registered(): + src = _read(os.path.join("database", "migrations.py")) + assert "async def ensure_ssl_csrs_table" in src + + inner = src.split("async def _run_all_migrations_inner", 1)[1] + inner = inner.split("\nasync def ", 1)[0] # body of the runner only + assert "await ensure_ssl_csrs_table()" in inner, ( + "ensure_ssl_csrs_table must be invoked from _run_all_migrations_inner" + ) + + +def test_ssl_csrs_ddl_essentials(): + src = _read(os.path.join("database", "migrations.py")) + ddl_start = src.index("CREATE TABLE IF NOT EXISTS ssl_csrs") + ddl = src[ddl_start:ddl_start + 2500] + + assert "private_key_pem TEXT" in ddl + assert "name VARCHAR(100) NOT NULL" in ddl, ( + "ssl_csrs.name must align with ssl_certificates.name VARCHAR(100)" + ) + assert "ssl_certificate_id INTEGER REFERENCES ssl_certificates(id) ON DELETE SET NULL" in ddl, ( + "deleting the imported cert must not cascade into CSR history" + ) + # Partial unique index: only PENDING CSRs reserve their target cert name. + assert "uq_ssl_csrs_name_pending" in src + assert re.search( + r"uq_ssl_csrs_name_pending\s+ON\s+ssl_csrs\(name\)\s+WHERE\s+status\s*=\s*'pending'", + src, + ), "name uniqueness must be scoped to pending CSRs (partial index)" + + +def test_csr_router_registered_in_main(): + src = _read("main.py") + assert "from routers.csr import router as csr_router" in src + assert "app.include_router(csr_router)" in src + + +def test_csr_endpoint_permission_mapping(): + """Pin which ssl. permission each endpoint enforces: a regression + that dropped or weakened a _require() call would otherwise pass the + auth-rejection tests (they only assert 401/403 for unauthenticated calls).""" + src = _read(os.path.join("routers", "csr.py")) + + def _handler_body(decorator): + start = src.index(decorator) + nxt = src.find("@router.", start + 1) + return src[start:nxt if nxt != -1 else len(src)] + + expectations = [ + ('@router.post("")', '"create"'), + ('@router.get("")', '"read"'), + ('@router.get("/{csr_id}")', '"read"'), + ('@router.post("/{csr_id}/import")', '"create"'), + ('@router.delete("/{csr_id}")', '"delete"'), + ] + for decorator, action in expectations: + body = _handler_body(decorator) + assert f"_require(authorization, {action})" in body, ( + f"endpoint {decorator} must enforce ssl.{action.strip(chr(34))}" + ) + + +def test_csr_router_never_selects_private_key(): + """The CSR endpoints must use the explicit column list — a bare + `SELECT *` into an API response is how the key would leak. The one place + SELECT * is allowed is the service-layer FOR UPDATE row (it needs the key + to pair with the cert); the router itself must not touch the column.""" + src = _read(os.path.join("routers", "csr.py")) + code_only = re.sub(r"#.*", "", src) # strip comments; the column name may + # legitimately appear there as documentation + assert "private_key_pem" not in code_only, ( + "routers/csr.py must never reference private_key_pem in code" + ) + assert "SELECT *" not in code_only, "routers/csr.py must use explicit column lists" diff --git a/backend/tests/test_csr_models.py b/backend/tests/test_csr_models.py new file mode 100644 index 0000000..13fb0a3 --- /dev/null +++ b/backend/tests/test_csr_models.py @@ -0,0 +1,172 @@ +""" +v1.9.0 CSR creation — Pydantic model validation tests (models/csr.py). + +The CSR name shares the SSL certificate name's path-traversal contract +(Bulgu #21) with one deliberate tightening: max 100 chars, matching the +ssl_certificates.name VARCHAR(100) column. +""" +import pytest +from pydantic import ValidationError + +from models.csr import SSLCSRCreate, SSLCSRImport + +_CERT_PEM = "-----BEGIN CERTIFICATE-----\nX\n-----END CERTIFICATE-----" + + +def _create(**overrides): + base = dict(name="my-csr", common_name="www.example.com") + base.update(overrides) + return SSLCSRCreate(**base) + + +# ---------------------------------------------------------------------------- +# SSLCSRCreate +# ---------------------------------------------------------------------------- + + +def test_minimal_valid_create(): + m = _create() + assert m.name == "my-csr" + assert m.common_name == "www.example.com" + assert m.key_algorithm == "rsa-2048" + assert m.sans == [] + + +@pytest.mark.parametrize("bad_name", [ + "../../etc/cron.d/evil", # path traversal + "a..b", # embedded .. + ".hidden", # hidden filename + "-flag", # CLI flag confusion + "has space", + "wild*card", + "", + "x" * 101, # VARCHAR(100) alignment — 200 is NOT allowed here +]) +def test_name_rejects_unsafe_values(bad_name): + with pytest.raises(ValidationError): + _create(name=bad_name) + + +def test_name_accepts_100_chars(): + assert _create(name="x" * 100).name == "x" * 100 + + +def test_common_name_wildcard_accepted_and_lowercased(): + m = _create(common_name="*.Example.COM") + assert m.common_name == "*.example.com" + + +@pytest.mark.parametrize("bad_cn", [ + "", + "under_score.example.com", # _ is not LDH + "*.*.example.com", # wildcard only as leftmost single label + "-leading.example.com", + "a" * 70 + ".example.com", # label > 63 + "cn-longer-than-64-chars-" + "x" * 45 + ".example.com", # CN > 64 total +]) +def test_common_name_rejects_invalid(bad_cn): + with pytest.raises(ValidationError): + _create(common_name=bad_cn) + + +def test_sans_normalised_deduped_and_capped(): + m = _create(sans=["API.example.com", "api.example.com", "cdn.example.com"]) + assert m.sans == ["api.example.com", "cdn.example.com"] + + with pytest.raises(ValidationError): + _create(sans=[f"h{i}.example.com" for i in range(101)]) + + +def test_country_normalised_or_rejected(): + assert _create(country="tr").country == "TR" + assert _create(country=None).country is None + for bad in ("TUR", "T", "1A"): + with pytest.raises(ValidationError): + _create(country=bad) + + +def test_subject_fields_reject_control_characters(): + with pytest.raises(ValidationError): + _create(organization="Evil\x00Corp") + with pytest.raises(ValidationError): + _create(locality="line\nbreak") + + +def test_subject_fields_reject_overlength(): + with pytest.raises(ValidationError): + _create(organization="x" * 65) + + +def test_key_algorithm_strict_enum(): + for good in ("rsa-2048", "rsa-4096", "ecdsa-p256", "ecdsa-p384"): + assert _create(key_algorithm=good).key_algorithm == good + for bad in ("rsa-1024", "rsa-8192", "ed25519", "2048", ""): + with pytest.raises(ValidationError): + _create(key_algorithm=bad) + + +def test_email_basic_validation(): + assert _create(email="ops@example.com").email == "ops@example.com" + with pytest.raises(ValidationError): + _create(email="not-an-email") + + +# ---------------------------------------------------------------------------- +# SSLCSRImport +# ---------------------------------------------------------------------------- + + +def test_import_minimal_global(): + m = SSLCSRImport(certificate_content=_CERT_PEM, is_global=True) + assert m.usage_type == "frontend" + assert m.name is None + + +def test_import_requires_clusters_when_not_global(): + with pytest.raises(ValidationError): + SSLCSRImport(certificate_content=_CERT_PEM, is_global=False) + with pytest.raises(ValidationError): + SSLCSRImport(certificate_content=_CERT_PEM, is_global=False, cluster_ids=[]) + m = SSLCSRImport(certificate_content=_CERT_PEM, is_global=False, cluster_ids=[1]) + assert m.cluster_ids == [1] + + +def test_import_certificate_must_be_pem(): + with pytest.raises(ValidationError): + SSLCSRImport(certificate_content="not a pem", is_global=True) + with pytest.raises(ValidationError): + SSLCSRImport(certificate_content="", is_global=True) + + +def test_import_certificate_size_capped(): + huge = _CERT_PEM + "A" * (64 * 1024 + 1) + with pytest.raises(ValidationError): + SSLCSRImport(certificate_content=huge, is_global=True) + + +def test_import_chain_optional_but_validated(): + m = SSLCSRImport(certificate_content=_CERT_PEM, is_global=True, chain_content=" ") + assert m.chain_content is None + with pytest.raises(ValidationError): + SSLCSRImport( + certificate_content=_CERT_PEM, is_global=True, chain_content="garbage" + ) + + +def test_import_name_override_shares_the_name_contract(): + m = SSLCSRImport(certificate_content=_CERT_PEM, is_global=True, name="renamed") + assert m.name == "renamed" + with pytest.raises(ValidationError): + SSLCSRImport(certificate_content=_CERT_PEM, is_global=True, name="../evil") + # Empty override collapses to None (falls back to the CSR's own name). + m2 = SSLCSRImport(certificate_content=_CERT_PEM, is_global=True, name=" ") + assert m2.name is None + + +def test_import_usage_type_enum(): + for good in ("frontend", "server"): + assert SSLCSRImport( + certificate_content=_CERT_PEM, is_global=True, usage_type=good + ).usage_type == good + with pytest.raises(ValidationError): + SSLCSRImport(certificate_content=_CERT_PEM, is_global=True, usage_type="both") diff --git a/backend/tests/test_csr_router_auth.py b/backend/tests/test_csr_router_auth.py new file mode 100644 index 0000000..d453ad5 --- /dev/null +++ b/backend/tests/test_csr_router_auth.py @@ -0,0 +1,66 @@ +""" +v1.9.0 CSR creation — behavioral auth tests for /api/ssl/csrs endpoints +(pattern: test_ssl_list_endpoint_auth.py). + +Every CSR endpoint must refuse unauthenticated / garbage-token requests. +The CSR detail route additionally must never 200 without auth because it +returns the CSR PEM; no endpoint ever returns the private key, but auth is +the first line regardless. +""" +import pytest + +_VALID_CREATE_BODY = { + "name": "auth-test-csr", + "common_name": "www.example.com", +} + +_VALID_IMPORT_BODY = { + "certificate_content": ( + "-----BEGIN CERTIFICATE-----\nX\n-----END CERTIFICATE-----" + ), + "is_global": True, +} + +_ENDPOINTS = [ + ("get", "/api/ssl/csrs", None), + ("get", "/api/ssl/csrs/1", None), + ("post", "/api/ssl/csrs", _VALID_CREATE_BODY), + ("post", "/api/ssl/csrs/1/import", _VALID_IMPORT_BODY), + ("delete", "/api/ssl/csrs/1", None), +] + + +@pytest.mark.parametrize("method,path,body", _ENDPOINTS) +def test_csr_endpoint_unauthenticated_rejected(client, method, path, body): + """No Authorization header → endpoint must refuse the request.""" + res = getattr(client, method)(path, json=body) if body is not None else getattr(client, method)(path) + assert res.status_code in (401, 403, 422), ( + f"{method.upper()} {path} without Authorization returned " + f"{res.status_code} — anonymous access to CSR data must not be " + f"possible. Body: {res.text[:200]}" + ) + if res.status_code == 200: # defensive, mirrors the R18 test style + data = res.json() + assert not data, "CSR endpoint returned data without auth" + + +@pytest.mark.parametrize("method,path,body", _ENDPOINTS) +def test_csr_endpoint_invalid_token_rejected(client, method, path, body): + """Garbage token → endpoint must refuse the request.""" + headers = {"Authorization": "Bearer not-a-valid-jwt"} + if body is not None: + res = getattr(client, method)(path, json=body, headers=headers) + else: + res = getattr(client, method)(path, headers=headers) + assert res.status_code in (401, 403, 422), ( + f"{method.upper()} {path} with an invalid token returned {res.status_code}" + ) + + +def test_csr_routes_are_registered(client): + """The router must actually be mounted — a 404 would make the auth tests + above pass vacuously.""" + res = client.get("/api/ssl/csrs") + assert res.status_code != 404, ( + "GET /api/ssl/csrs returned 404 — csr_router is not registered in main.py" + ) diff --git a/backend/tests/test_csr_service_crypto.py b/backend/tests/test_csr_service_crypto.py new file mode 100644 index 0000000..7816ed2 --- /dev/null +++ b/backend/tests/test_csr_service_crypto.py @@ -0,0 +1,171 @@ +""" +v1.9.0 CSR creation — pure-crypto tests for services/csr_service.py. + +No mocks: every algorithm's output must parse with `cryptography` and the +CSR's public key must match the generated private key (the property the +whole import flow depends on). +""" +from types import SimpleNamespace + +import pytest + +from cryptography import x509 +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec, rsa +from cryptography.x509.oid import ExtensionOID, NameOID + +from services.csr_service import csr_row_to_dict, diff_domains, generate_csr_bundle + + +def _payload(**overrides): + base = dict( + name="test-csr", + common_name="www.example.com", + organization=None, + organizational_unit=None, + locality=None, + state=None, + country=None, + email=None, + sans=[], + key_algorithm="rsa-2048", + ) + base.update(overrides) + return SimpleNamespace(**base) + + +def _spki(key): + return key.public_key().public_bytes( + serialization.Encoding.DER, + serialization.PublicFormat.SubjectPublicKeyInfo, + ) + + +@pytest.mark.parametrize( + "algo,key_cls,key_check", + [ + ("rsa-2048", rsa.RSAPrivateKey, lambda k: k.key_size == 2048), + ("rsa-4096", rsa.RSAPrivateKey, lambda k: k.key_size == 4096), + ("ecdsa-p256", ec.EllipticCurvePrivateKey, lambda k: k.curve.name == "secp256r1"), + ("ecdsa-p384", ec.EllipticCurvePrivateKey, lambda k: k.curve.name == "secp384r1"), + ], +) +def test_generate_bundle_all_algorithms(algo, key_cls, key_check): + bundle = generate_csr_bundle(_payload(key_algorithm=algo)) + + csr = x509.load_pem_x509_csr(bundle["csr_pem"].encode()) + key = serialization.load_pem_private_key( + bundle["private_key_pem"].encode(), password=None + ) + + assert isinstance(key, key_cls) + assert key_check(key) + # The CSR must be signed by exactly this key. + csr_spki = csr.public_key().public_bytes( + serialization.Encoding.DER, + serialization.PublicFormat.SubjectPublicKeyInfo, + ) + assert csr_spki == _spki(key) + assert csr.is_signature_valid + # PKCS8, unencrypted — the agent concatenates cert+key into one PEM and + # HAProxy cannot read passphrase-protected keys. + assert bundle["private_key_pem"].startswith("-----BEGIN PRIVATE KEY-----") + + +def test_subject_contains_all_provided_fields(): + bundle = generate_csr_bundle(_payload( + organization="Example Corp", + organizational_unit="IT", + locality="Istanbul", + state="Marmara", + country="TR", + email="ops@example.com", + )) + csr = x509.load_pem_x509_csr(bundle["csr_pem"].encode()) + + def _one(oid): + attrs = csr.subject.get_attributes_for_oid(oid) + return attrs[0].value if attrs else None + + assert _one(NameOID.COMMON_NAME) == "www.example.com" + assert _one(NameOID.ORGANIZATION_NAME) == "Example Corp" + assert _one(NameOID.ORGANIZATIONAL_UNIT_NAME) == "IT" + assert _one(NameOID.LOCALITY_NAME) == "Istanbul" + assert _one(NameOID.STATE_OR_PROVINCE_NAME) == "Marmara" + assert _one(NameOID.COUNTRY_NAME) == "TR" + assert _one(NameOID.EMAIL_ADDRESS) == "ops@example.com" + assert bundle["subject"] == { + "O": "Example Corp", "OU": "IT", "L": "Istanbul", + "ST": "Marmara", "C": "TR", "emailAddress": "ops@example.com", + } + + +def test_subject_omits_empty_fields(): + bundle = generate_csr_bundle(_payload()) + csr = x509.load_pem_x509_csr(bundle["csr_pem"].encode()) + assert not csr.subject.get_attributes_for_oid(NameOID.ORGANIZATION_NAME) + assert bundle["subject"] == {} + + +def test_sans_cn_first_and_deduped(): + bundle = generate_csr_bundle(_payload( + common_name="www.example.com", + sans=["api.example.com", "www.example.com", "api.example.com", "cdn.example.com"], + )) + assert bundle["sans"] == ["www.example.com", "api.example.com", "cdn.example.com"] + + csr = x509.load_pem_x509_csr(bundle["csr_pem"].encode()) + san_ext = csr.extensions.get_extension_for_oid( + ExtensionOID.SUBJECT_ALTERNATIVE_NAME + ) + dns_names = san_ext.value.get_values_for_type(x509.DNSName) + assert dns_names == ["www.example.com", "api.example.com", "cdn.example.com"] + + +def test_wildcard_common_name_flows_into_san(): + bundle = generate_csr_bundle(_payload(common_name="*.example.com")) + csr = x509.load_pem_x509_csr(bundle["csr_pem"].encode()) + san_ext = csr.extensions.get_extension_for_oid( + ExtensionOID.SUBJECT_ALTERNATIVE_NAME + ) + assert san_ext.value.get_values_for_type(x509.DNSName) == ["*.example.com"] + + +def test_diff_domains_reports_added_and_dropped(): + warnings = diff_domains( + ["www.example.com", "api.example.com"], + ["WWW.example.com", "cdn.example.com"], + ) + assert len(warnings) == 2 + added = next(w for w in warnings if "added" in w) + dropped = next(w for w in warnings if "dropped" in w) + assert "cdn.example.com" in added + assert "api.example.com" in dropped + # Case-insensitive: www must NOT be reported in either direction. + assert "www.example.com" not in added + assert "www.example.com" not in dropped + + +def test_diff_domains_identical_sets_yield_no_warnings(): + assert diff_domains(["a.example.com"], ["A.EXAMPLE.COM"]) == [] + assert diff_domains([], []) == [] + + +def test_csr_row_to_dict_never_exposes_private_key(): + row = { + "id": 1, + "name": "x", + "private_key_pem": "-----BEGIN PRIVATE KEY-----\nSECRET\n-----END PRIVATE KEY-----", + "csr_pem": "-----BEGIN CERTIFICATE REQUEST-----\nX\n-----END CERTIFICATE REQUEST-----", + "subject": '{"O": "Example"}', + "sans": '["a.example.com"]', + } + out = csr_row_to_dict(row) + assert "private_key_pem" not in out + assert "csr_pem" not in out # lists exclude the PEM + assert out["subject"] == {"O": "Example"} + assert out["sans"] == ["a.example.com"] + + detail = csr_row_to_dict(row, include_pem=True) + assert "private_key_pem" not in detail # NEVER, even on detail + assert detail["csr_pem"].startswith("-----BEGIN CERTIFICATE REQUEST-----")