diff --git a/packages/bun-types/bun.d.ts b/packages/bun-types/bun.d.ts index 221aa2fea1f6..be1ca03d1394 100644 --- a/packages/bun-types/bun.d.ts +++ b/packages/bun-types/bun.d.ts @@ -4729,6 +4729,11 @@ declare module "bun" { */ serverName?: string; + /** + * Alias of {@link serverName} (the `node:tls` spelling). `serverName` wins if both are set. + */ + servername?: string | undefined; + /** * Sets `OPENSSL_RELEASE_BUFFERS` to 1. * Reduces overall performance but saves some memory. diff --git a/packages/bun-types/redis.d.ts b/packages/bun-types/redis.d.ts index 79b394fd2fb2..33da1370141b 100644 --- a/packages/bun-types/redis.d.ts +++ b/packages/bun-types/redis.d.ts @@ -35,9 +35,26 @@ declare module "bun" { enableOfflineQueue?: boolean; /** - * TLS options - */ - tls?: boolean | Bun.TLSOptions; + * TLS options. `true` (or a `rediss://` URL) enables TLS with the default + * options. `serverName` sets the SNI and the name the certificate is + * verified against (by default the URL host). + */ + tls?: + | boolean + | (Bun.TLSOptions & { + /** + * Replaces the built-in check of the server certificate against + * `serverName`, as in `tls.connect()`. It runs once the certificate + * chain is verified, and not at all with `rejectUnauthorized: false`. + * @param hostname The name the certificate is expected to match + * @param cert The server's certificate, with its issuers in `issuerCertificate` + * @returns An error if the server is unauthorized, otherwise undefined. Any + * other truthy value also fails the connection, and a `Promise` is one: the + * function cannot be `async`. Every falsy value approves the certificate, as + * in Node, so `false` does not reject it. + */ + checkServerIdentity?: NonNullable | undefined; + }); /** * Whether to enable auto-pipelining diff --git a/packages/bun-types/sql.d.ts b/packages/bun-types/sql.d.ts index 86c7ec2fb7fa..ee2f638b864e 100644 --- a/packages/bun-types/sql.d.ts +++ b/packages/bun-types/sql.d.ts @@ -201,6 +201,27 @@ declare module "bun" { onclose?: ((err: Error | null) => void) | undefined; } + /** + * TLS options for a PostgreSQL or MySQL connection + */ + interface TLSOptions extends Bun.TLSOptions { + /** + * Replaces the built-in check of the server certificate against + * `serverName` (by default the connection's hostname), as in + * `tls.connect()`. It runs once the certificate chain is verified, under + * the `verify-ca` and `verify-full` SSL modes. Like `ca`, setting it + * turns certificate verification on unless an SSL mode or + * `rejectUnauthorized: false` says otherwise. + * @param hostname The name the certificate is expected to match + * @param cert The server's certificate, with its issuers in `issuerCertificate` + * @returns An error if the server is unauthorized, otherwise undefined. Any + * other truthy value also fails the connection, and a `Promise` is one: the + * function cannot be `async`. Every falsy value approves the certificate, as + * in Node, so `false` does not reject it. + */ + checkServerIdentity?: NonNullable | undefined; + } + interface PostgresOrMySQLOptions { /** * Connection URL, for example `postgres://user:pass@localhost:5432/mydb` diff --git a/src/boringssl_sys/boringssl.rs b/src/boringssl_sys/boringssl.rs index fcf5ac41b303..16fcd29554d5 100644 --- a/src/boringssl_sys/boringssl.rs +++ b/src/boringssl_sys/boringssl.rs @@ -514,6 +514,12 @@ impl SSL { } } + /// Client side: the SNI host name to send. Set it before the handshake; `false` if BoringSSL rejects it. + pub fn set_servername(&mut self, hostname: &core::ffi::CStr) -> bool { + // SAFETY: `self` is a live SSL; BoringSSL copies `hostname`. + unsafe { SSL_set_tlsext_host_name(self, hostname.as_ptr()) == 1 } + } + /// The peer's leaf certificate, borrowed from this SSL's cert chain. pub fn peer_leaf_certificate(&mut self) -> Option<&mut X509> { // SAFETY: the chain and its entries are owned by this SSL and outlive diff --git a/src/js/internal/sql/shared.ts b/src/js/internal/sql/shared.ts index 8f9cd3d94274..bc99acba97fd 100644 --- a/src/js/internal/sql/shared.ts +++ b/src/js/internal/sql/shared.ts @@ -1874,7 +1874,7 @@ function parseOptions( let username: string | null | undefined; let password: string | (() => Bun.MaybePromise) | undefined | null; let database: string | undefined; - let tls: Bun.TLSOptions | (Bun.BunFile & Bun.TLSOptions) | boolean | undefined; + let tls: Bun.SQL.TLSOptions | (Bun.BunFile & Bun.SQL.TLSOptions) | boolean | undefined; let query: string = ""; let idleTimeout: number | null | undefined; let connectionTimeout: number | null | undefined; @@ -2148,15 +2148,26 @@ function parseOptions( } } - if ($isObject(tls) && sslMode < SSLMode.verify_ca) { - if (tls.rejectUnauthorized === true || (tls.rejectUnauthorized !== false && (tls.ca || tls.caFile))) { + const checkServerIdentity = $isObject(tls) ? tls.checkServerIdentity : undefined; + if ($isObject(tls)) { + const { rejectUnauthorized } = tls; + if (checkServerIdentity !== undefined && !$isCallable(checkServerIdentity)) { + throw $ERR_INVALID_ARG_TYPE("tls.checkServerIdentity", "function", checkServerIdentity); + } + // These options imply certificate verification unless it is explicitly turned off. + if ( + sslMode < SSLMode.verify_ca && + (rejectUnauthorized === true || (rejectUnauthorized !== false && (tls.ca || tls.caFile || checkServerIdentity))) + ) { sslMode = SSLMode.verify_full; } } - if (sslMode !== SSLMode.disable && !(tls as Exclude)?.serverName) { + const tlsObject = tls as Exclude; + if (sslMode !== SSLMode.disable && !tlsObject?.serverName && !tlsObject?.servername) { if (hostname) { - tls = { ...(tls as Exclude), serverName: hostname }; + // The spread copies own enumerable properties only: a callback that is a class method would be lost. + tls = { ...tlsObject, checkServerIdentity, serverName: hostname }; } else if (tls) { tls = true; } diff --git a/src/jsc/TlsServerIdentity.rs b/src/jsc/TlsServerIdentity.rs new file mode 100644 index 000000000000..fcbf896bce33 --- /dev/null +++ b/src/jsc/TlsServerIdentity.rs @@ -0,0 +1,252 @@ +//! The peer certificate as JS, and `tls.checkServerIdentity`, for every TLS client whose handshake events run on the JS thread. + +use core::ffi::c_int; + +use bun_boringssl::c::{SSL, SSL_CTX, X509, X509_STORE, X509_STORE_CTX, struct_stack_st_X509}; + +use crate::{ErrorCode, JSGlobalObject, JSValue, JsResult}; + +unsafe extern "C" { + // JSX509Certificate.cpp. Opaque `repr(C)` handles: `&mut`/`&` are ABI-identical to non-null pointers. + safe fn Bun__X509__toJSLegacyEncoding( + cert: &mut X509, + global_object: &JSGlobalObject, + ) -> JSValue; + + safe fn SSL_is_server(ssl: &SSL) -> c_int; + // The process-wide default root store; up-refs before returning, so + // the caller owns a reference it must release with X509_STORE_free. + fn us_get_shared_default_ca_store() -> *mut X509_STORE; + fn us_ssl_ctx_has_user_ca(ctx: *mut SSL_CTX) -> c_int; + fn X509_STORE_free(store: *mut X509_STORE); + // X509_STORE_CTX lifecycle for issuer lookups; `new` allocates, + // `init` borrows the store, `free` releases. Used to extend the peer + // certificate chain through the local trust store. + fn X509_STORE_CTX_new() -> *mut X509_STORE_CTX; + fn X509_STORE_CTX_init( + ctx: *mut X509_STORE_CTX, + store: *mut X509_STORE, + x509: *mut X509, + chain: *mut struct_stack_st_X509, + ) -> c_int; + fn X509_STORE_CTX_free(ctx: *mut X509_STORE_CTX); + // Writes a +1 X509 reference to `*issuer` on success (> 0). + fn X509_STORE_CTX_get1_issuer( + issuer: *mut *mut X509, + ctx: *mut X509_STORE_CTX, + x: *mut X509, + ) -> c_int; + // Returns X509_V_OK (0) when `issuer` could have issued `subject`. + fn X509_check_issued(issuer: *mut X509, subject: *mut X509) -> c_int; +} + +/// The `tls.getPeerCertificate()`-style object for `cert` (borrowed, not adopted). +pub fn x509_to_legacy_object(cert: &mut X509, global: &JSGlobalObject) -> JsResult { + crate::from_js_host_call(global, || Bun__X509__toJSLegacyEncoding(cert, global)) +} + +/// `getPeerCertificate(true)`: the peer's certificate, or `undefined` when it sent none. +pub fn peer_certificate_chain(ssl: &mut SSL, global: &JSGlobalObject) -> JsResult { + let ssl_ptr: *mut SSL = ssl; + let mut cert: *mut X509 = core::ptr::null_mut(); + if SSL_is_server(ssl) != 0 { + // SSL_get_peer_certificate returns a +1 reference; we must free it. + // SAFETY: `ssl_ptr` is the live SSL behind `ssl`. + cert = unsafe { bun_boringssl::c::SSL_get_peer_certificate(ssl_ptr) }; + } + let _guard = scopeguard::guard(cert, |c| { + if !c.is_null() { + // SAFETY: `c` is the +1 X509 reference returned by SSL_get_peer_certificate; we own it. + unsafe { bun_boringssl::c::X509_free(c) }; + } + }); + + // SAFETY: `ssl_ptr` is the live SSL behind `ssl`; the chain and its entries are borrowed from it. + let cert_chain = unsafe { bun_boringssl::c::SSL_get_peer_cert_chain(ssl_ptr) }; + let first_cert: *mut X509 = if !cert.is_null() { + cert + } else if !cert_chain.is_null() { + // SAFETY: `cert_chain` is non-null and owned by the SSL. + unsafe { bun_boringssl::c::sk_X509_value(cert_chain, 0) } + } else { + core::ptr::null_mut() + }; + + if first_cert.is_null() { + return Ok(JSValue::UNDEFINED); + } + + // The detailed form returns the whole chain the peer presented, each + // certificate linking to its issuer through `issuerCertificate`, the way + // Node's getPeerCertificate(true) does. SSL_get_peer_cert_chain includes + // the leaf on the client side but not on the server side, where the +1 + // peer certificate above is the leaf instead. + let first_obj = x509_to_legacy_object(X509::opaque_mut(first_cert), global)?; + // Link each certificate to its predecessor immediately so every object in + // the chain is reachable from the stack-rooted `first_obj` before the next + // `x509_to_legacy_object` allocation can trigger a GC - a heap-backed Vec + // is not stack-scanned. + let mut prev_obj: JSValue = first_obj; + let mut last_cert: *mut X509 = first_cert; + if !cert_chain.is_null() { + let mut i: usize = if cert.is_null() { 1 } else { 0 }; + loop { + // SAFETY: `cert_chain` is non-null and owned by the SSL; out of range yields null. + let next = unsafe { bun_boringssl::c::sk_X509_value(cert_chain, i) }; + if next.is_null() { + break; + } + let obj = x509_to_legacy_object(X509::opaque_mut(next), global)?; + prev_obj.put(global, b"issuerCertificate", obj); + prev_obj = obj; + last_cert = next; + i += 1; + } + } + + // Extend the chain through the local trust store until a self-issued + // certificate is reached, the way Node's getPeerCertificate(true) walks + // X509_STORE_CTX_get1_issuer to surface the root that completed + // verification even though the peer never sent it. + let mut last_is_self_issued = false; + // SAFETY: the store ctx is created, initialized against the live SSL_CTX's + // store, used only within this scope and freed before returning; every + // issuer returned by get1_issuer is a +1 reference collected in `extras` + // and released after its fields have been copied into JS values and the + // terminal self-issued check has run. + unsafe { + let ssl_ctx = bun_boringssl::c::SSL_get_SSL_CTX(ssl_ptr); + let mut store = bun_boringssl::c::SSL_CTX_get_cert_store(ssl_ctx); + // A context built without an explicit `ca` (and without requestCert, + // which installs the shared roots) carries an empty store and the + // issuer walk would stop at whatever the peer sent. Fall back to the + // process-wide default roots the way Node's per-context store always + // contains the bundled roots. The getter up-refs, so the temporary + // reference is released after the walk. + let mut shared_store: *mut X509_STORE = core::ptr::null_mut(); + if store.is_null() || us_ssl_ctx_has_user_ca(ssl_ctx) == 0 { + shared_store = us_get_shared_default_ca_store(); + if !shared_store.is_null() { + store = shared_store; + } + } + let store_ctx = X509_STORE_CTX_new(); + if !store_ctx.is_null() { + if !store.is_null() + && X509_STORE_CTX_init( + store_ctx, + store, + core::ptr::null_mut(), + core::ptr::null_mut(), + ) == 1 + { + let mut extras: Vec<*mut X509> = Vec::new(); + // Cap the walk so a cyclic store cannot loop forever. + while extras.len() < 16 && X509_check_issued(last_cert, last_cert) != 0 { + let mut issuer: *mut X509 = core::ptr::null_mut(); + if X509_STORE_CTX_get1_issuer(&raw mut issuer, store_ctx, last_cert) <= 0 + || issuer.is_null() + { + break; + } + match x509_to_legacy_object(X509::opaque_mut(issuer), global) { + Ok(obj) => { + prev_obj.put(global, b"issuerCertificate", obj); + prev_obj = obj; + } + Err(e) => { + bun_boringssl::c::X509_free(issuer); + for extra in extras { + bun_boringssl::c::X509_free(extra); + } + X509_STORE_CTX_free(store_ctx); + if !shared_store.is_null() { + X509_STORE_free(shared_store); + } + return Err(e); + } + } + extras.push(issuer); + last_cert = issuer; + } + last_is_self_issued = X509_check_issued(last_cert, last_cert) == 0; + for extra in extras { + bun_boringssl::c::X509_free(extra); + } + } + X509_STORE_CTX_free(store_ctx); + } + if !shared_store.is_null() { + X509_STORE_free(shared_store); + } + } + + // A self-issued terminal certificate references itself, like Node. + if last_is_self_issued { + prev_obj.put(global, b"issuerCertificate", prev_obj); + } + Ok(first_obj) +} + +/// Runs `callback(hostname, getPeerCertificate(true))` as Node does. `Err` is what the connection fails with. +pub fn check_with_callback( + global: &JSGlobalObject, + callback: JSValue, + ssl: Option<&mut SSL>, + hostname: &[u8], +) -> Result<(), JSValue> { + let js_cert = match ssl { + Some(ssl) => peer_certificate_chain(ssl, global).map_err(|e| global.take_exception(e))?, + None => JSValue::UNDEFINED, + }; + if js_cert.is_undefined() { + return Err(global + .err( + ErrorCode::TLS_CERT_ALTNAME_INVALID, + format_args!("The server did not present a certificate"), + ) + .to_js()); + } + let js_hostname = crate::bun_string_jsc::create_utf8_for_js(global, hostname) + .map_err(|e| global.take_exception(e))?; + let result = { + let _scope = global.bun_vm().enter_event_loop_scope(); + callback.call(global, JSValue::UNDEFINED, &[js_hostname, js_cert]) + }; + let result = result.map_err(|e| global.take_exception(e))?; + // A VM that is stopping calls nobody and answers `undefined`: nobody approved this certificate. + let vm = global.bun_vm(); + if !vm.script_allowed() || global.vm().execution_forbidden() || vm.calls_nobody() { + return Err(global + .err( + ErrorCode::TLS_CERT_ALTNAME_INVALID, + format_args!("\"tls.checkServerIdentity\" did not run"), + ) + .to_js()); + } + verdict_of(global, result) +} + +/// What a `tls.checkServerIdentity` function decided, from the value it returned. `Err` is the reason it refused the server. +pub fn verdict_of(global: &JSGlobalObject, returned: JSValue) -> Result<(), JSValue> { + // > Returns object [...] on failure + // Any object counts: a DOMException or a util.inherits() error is not an ErrorInstance cell. + if returned.is_object() && returned.as_any_promise().is_none() { + return Err(returned); + } + // Like Node, fail on any other truthy value, a Promise included: https://github.com/nodejs/node/blob/v26.3.0/lib/internal/tls/wrap.js#L1671-L1688 + if returned.to_boolean() { + let received = JSGlobalObject::determine_specific_type(global, returned) + .map_err(|e| global.take_exception(e))?; + return Err(global + .err( + ErrorCode::INVALID_RETURN_VALUE, + format_args!( + "Expected undefined or an Error to be returned from the \"tls.checkServerIdentity\" function but got {received}." + ), + ) + .to_js()); + } + // > On success, returns + Ok(()) +} diff --git a/src/jsc/lib.rs b/src/jsc/lib.rs index 9bdfc2592295..6e2b8f06ace9 100644 --- a/src/jsc/lib.rs +++ b/src/jsc/lib.rs @@ -447,6 +447,8 @@ pub mod js_property_iterator; pub mod node_compile_cache; #[path = "SystemError.rs"] pub mod system_error; +#[path = "TlsServerIdentity.rs"] +pub mod tls_server_identity; #[path = "URL.rs"] pub mod url; #[path = "VM.rs"] diff --git a/src/runtime/api/BunObject.rs b/src/runtime/api/BunObject.rs index b216ec4eb039..8594d52d580f 100644 --- a/src/runtime/api/BunObject.rs +++ b/src/runtime/api/BunObject.rs @@ -1806,7 +1806,8 @@ fn get_valkey_default_client(global_this: &JSGlobalObject, _: &JSObject) -> JSVa &global_this.js_thread(vm.root_context()), &[JSValue::UNDEFINED], ) { - Ok(p) => p, + // No options object, so no `tls.checkServerIdentity` to store. + Ok((p, _)) => p, Err(jsc::JsError::Thrown) => return JSValue::ZERO, Err(err) => { let _ = diff --git a/src/runtime/api/bun/x509.rs b/src/runtime/api/bun/x509.rs index 95120b3d7526..ee9e0d671f62 100644 --- a/src/runtime/api/bun/x509.rs +++ b/src/runtime/api/bun/x509.rs @@ -1,11 +1,7 @@ use bun_boringssl_sys::X509; use bun_jsc::{JSGlobalObject, JSValue, JsResult}; -pub(crate) fn to_js(cert: &mut X509, global_object: &JSGlobalObject) -> JsResult { - bun_jsc::from_js_host_call(global_object, || { - Bun__X509__toJSLegacyEncoding(cert, global_object) - }) -} +pub(crate) use bun_jsc::tls_server_identity::x509_to_legacy_object as to_js; pub(crate) fn to_js_object(cert: &mut X509, global_object: &JSGlobalObject) -> JsResult { Ok(Bun__X509__toJS(cert, global_object)) @@ -14,9 +10,5 @@ pub(crate) fn to_js_object(cert: &mut X509, global_object: &JSGlobalObject) -> J // `X509`/`JSGlobalObject` are opaque `repr(C)` handles; `&mut`/`&` are // ABI-identical to non-null pointers, so the validity proof is in the type. unsafe extern "C" { - safe fn Bun__X509__toJSLegacyEncoding( - cert: &mut X509, - global_object: &JSGlobalObject, - ) -> JSValue; safe fn Bun__X509__toJS(cert: &mut X509, global_object: &JSGlobalObject) -> JSValue; } diff --git a/src/runtime/api/sql.classes.ts b/src/runtime/api/sql.classes.ts index c4c7ca8ef009..8f81bf418c5b 100644 --- a/src/runtime/api/sql.classes.ts +++ b/src/runtime/api/sql.classes.ts @@ -53,8 +53,8 @@ for (const type of types) { }, values: type === "PostgresSQL" - ? ["onconnect", "onclose", "queries", "onnotification"] - : ["onconnect", "onclose", "queries"], + ? ["onconnect", "onclose", "queries", "checkServerIdentity", "onnotification"] + : ["onconnect", "onclose", "queries", "checkServerIdentity"], }), ); diff --git a/src/runtime/socket/tls_socket_functions.rs b/src/runtime/socket/tls_socket_functions.rs index 543fbed1d181..68061474b29f 100644 --- a/src/runtime/socket/tls_socket_functions.rs +++ b/src/runtime/socket/tls_socket_functions.rs @@ -18,7 +18,7 @@ use crate::api::bun_x509 as X509; // ────────────────────────────────────────────────────────────────────────── #[allow(non_camel_case_types, non_upper_case_globals)] pub(super) mod ffi { - use super::boringssl::{SSL, SSL_CTX, X509, X509_STORE, X509_STORE_CTX, struct_stack_st_X509}; + use super::boringssl::{SSL, SSL_CTX, X509, struct_stack_st_X509}; use core::ffi::{c_char, c_int, c_long, c_uint, c_void}; // Re-export the one decl whose `*const c_char` NUL-terminated arg keeps a @@ -243,32 +243,6 @@ pub(super) mod ffi { >, arg: *mut c_void, ); - // Returns the borrowed cert store of a live `SSL_CTX*`. - pub(crate) safe fn SSL_CTX_get_cert_store(ctx: &SSL_CTX) -> *mut X509_STORE; - // The process-wide default root store; up-refs before returning, so - // the caller owns a reference it must release with X509_STORE_free. - pub(crate) fn us_get_shared_default_ca_store() -> *mut X509_STORE; - pub(crate) fn us_ssl_ctx_has_user_ca(ctx: *mut SSL_CTX) -> c_int; - pub(crate) fn X509_STORE_free(store: *mut X509_STORE); - // X509_STORE_CTX lifecycle for issuer lookups; `new` allocates, - // `init` borrows the store, `free` releases. Used to extend the peer - // certificate chain through the local trust store. - pub(crate) fn X509_STORE_CTX_new() -> *mut X509_STORE_CTX; - pub(crate) fn X509_STORE_CTX_init( - ctx: *mut X509_STORE_CTX, - store: *mut X509_STORE, - x509: *mut X509, - chain: *mut struct_stack_st_X509, - ) -> c_int; - pub(crate) fn X509_STORE_CTX_free(ctx: *mut X509_STORE_CTX); - // Writes a +1 X509 reference to `*issuer` on success (> 0). - pub(crate) fn X509_STORE_CTX_get1_issuer( - issuer: *mut *mut X509, - ctx: *mut X509_STORE_CTX, - x: *mut X509, - ) -> c_int; - // Returns X509_V_OK (0) when `issuer` could have issued `subject`. - pub(crate) fn X509_check_issued(issuer: *mut X509, subject: *mut X509) -> c_int; } } use crate::node::StringOrBuffer; @@ -478,143 +452,7 @@ pub(super) fn get_peer_certificate( return X509::to_js(boringssl::X509::opaque_mut(cert), global); } - let mut cert: *mut boringssl::X509 = core::ptr::null_mut(); - if is_server_ssl { - // SSL_get_peer_certificate returns a +1 reference; we must free it. - cert = ffi::SSL_get_peer_certificate(boringssl::SSL::opaque_ref(ssl_ptr)); - } - let _guard = scopeguard::guard(cert, |c| { - if !c.is_null() { - // SAFETY: `c` is the +1 X509 reference returned by SSL_get_peer_certificate; we own it. - unsafe { boringssl::X509_free(c) }; - } - }); - - let cert_chain = ffi::SSL_get_peer_cert_chain(boringssl::SSL::opaque_ref(ssl_ptr)); - let first_cert: *mut boringssl::X509 = if !cert.is_null() { - cert - } else if !cert_chain.is_null() { - ffi::sk_X509_value(boringssl::struct_stack_st_X509::opaque_ref(cert_chain), 0) - } else { - core::ptr::null_mut() - }; - - if first_cert.is_null() { - return Ok(JSValue::UNDEFINED); - } - - // The detailed form returns the whole chain the peer presented, each - // certificate linking to its issuer through `issuerCertificate`, the way - // Node's getPeerCertificate(true) does. SSL_get_peer_cert_chain includes - // the leaf on the client side but not on the server side, where the +1 - // peer certificate above is the leaf instead. - let first_obj = X509::to_js(boringssl::X509::opaque_mut(first_cert), global)?; - // Link each certificate to its predecessor immediately so every object in - // the chain is reachable from the stack-rooted `first_obj` before the next - // `X509::to_js` allocation can trigger a GC - a heap-backed Vec - // is not stack-scanned. - let mut prev_obj: JSValue = first_obj; - let mut last_cert: *mut boringssl::X509 = first_cert; - if !cert_chain.is_null() { - let mut i: usize = if cert.is_null() { 1 } else { 0 }; - loop { - let next = - ffi::sk_X509_value(boringssl::struct_stack_st_X509::opaque_ref(cert_chain), i); - if next.is_null() { - break; - } - let obj = X509::to_js(boringssl::X509::opaque_mut(next), global)?; - prev_obj.put(global, b"issuerCertificate", obj); - prev_obj = obj; - last_cert = next; - i += 1; - } - } - - // Extend the chain through the local trust store until a self-issued - // certificate is reached, the way Node's getPeerCertificate(true) walks - // X509_STORE_CTX_get1_issuer to surface the root that completed - // verification even though the peer never sent it. - let mut last_is_self_issued = false; - // SAFETY: the store ctx is created, initialized against the live SSL_CTX's - // store, used only within this scope and freed before returning; every - // issuer returned by get1_issuer is a +1 reference collected in `extras` - // and released after its fields have been copied into JS values and the - // terminal self-issued check has run. - unsafe { - let mut store = ffi::SSL_CTX_get_cert_store(boringssl::SSL_CTX::opaque_ref( - ffi::SSL_get_SSL_CTX(boringssl::SSL::opaque_ref(ssl_ptr)), - )); - // A context built without an explicit `ca` (and without requestCert, - // which installs the shared roots) carries an empty store and the - // issuer walk would stop at whatever the peer sent. Fall back to the - // process-wide default roots the way Node's per-context store always - // contains the bundled roots. The getter up-refs, so the temporary - // reference is released after the walk. - let mut shared_store: *mut boringssl::X509_STORE = core::ptr::null_mut(); - let ssl_ctx = ffi::SSL_get_SSL_CTX(boringssl::SSL::opaque_ref(ssl_ptr)); - if store.is_null() || ffi::us_ssl_ctx_has_user_ca(ssl_ctx) == 0 { - shared_store = ffi::us_get_shared_default_ca_store(); - if !shared_store.is_null() { - store = shared_store; - } - } - let store_ctx = ffi::X509_STORE_CTX_new(); - if !store_ctx.is_null() { - if !store.is_null() - && ffi::X509_STORE_CTX_init( - store_ctx, - store, - core::ptr::null_mut(), - core::ptr::null_mut(), - ) == 1 - { - let mut extras: Vec<*mut boringssl::X509> = Vec::new(); - // Cap the walk so a cyclic store cannot loop forever. - while extras.len() < 16 && ffi::X509_check_issued(last_cert, last_cert) != 0 { - let mut issuer: *mut boringssl::X509 = core::ptr::null_mut(); - if ffi::X509_STORE_CTX_get1_issuer(&raw mut issuer, store_ctx, last_cert) <= 0 - || issuer.is_null() - { - break; - } - match X509::to_js(boringssl::X509::opaque_mut(issuer), global) { - Ok(obj) => { - prev_obj.put(global, b"issuerCertificate", obj); - prev_obj = obj; - } - Err(e) => { - boringssl::X509_free(issuer); - for extra in extras { - boringssl::X509_free(extra); - } - ffi::X509_STORE_CTX_free(store_ctx); - if !shared_store.is_null() { - ffi::X509_STORE_free(shared_store); - } - return Err(e); - } - } - extras.push(issuer); - last_cert = issuer; - } - last_is_self_issued = ffi::X509_check_issued(last_cert, last_cert) == 0; - for extra in extras { - boringssl::X509_free(extra); - } - } - ffi::X509_STORE_CTX_free(store_ctx); - } - if !shared_store.is_null() { - ffi::X509_STORE_free(shared_store); - } - } - - // A self-issued terminal certificate references itself, like Node. - if last_is_self_issued { - prev_obj.put(global, b"issuerCertificate", prev_obj); - } - Ok(first_obj) + jsc::tls_server_identity::peer_certificate_chain(boringssl::SSL::opaque_mut(ssl_ptr), global) } pub(super) fn get_certificate( diff --git a/src/runtime/valkey_jsc/js_valkey.rs b/src/runtime/valkey_jsc/js_valkey.rs index 3f863f0a8838..4d71ec565ae8 100644 --- a/src/runtime/valkey_jsc/js_valkey.rs +++ b/src/runtime/valkey_jsc/js_valkey.rs @@ -8,7 +8,7 @@ use bun_io::KeepAlive; use bun_jsc::virtual_machine::VirtualMachine; use bun_jsc::{ self as jsc, CallFrame, GlobalRef, JSArray, JSGlobalObject, JSMap, JSPromise, JSValue, JsCell, - JsRef, JsResult, + JsRef, JsResult, tls_server_identity, }; use bun_ptr::{AsCtxPtr, BackRef, RefPtr}; use bun_uws as uws; @@ -439,6 +439,13 @@ impl JSValkeyClient { // SAFETY: `self` is the live heap allocation. unsafe { RefPtr::init_ref(self.as_ctx_ptr()) } } + /// `tls.checkServerIdentity`: replaces the native name check, and runs after the handshake. + pub(crate) fn check_server_identity_callback(&self) -> Option { + self.this_value + .get() + .try_get() + .and_then(Js::check_server_identity_get_cached) + } #[inline] pub(crate) fn new(init: JSValkeyClient) -> *mut JSValkeyClient { // bun.TrivialNew(@This()) → heap::alloc(Box::new(init)) @@ -488,13 +495,13 @@ impl JSValkeyClient { ) } - /// Create a Valkey client that does not have an associated JS object nor a SubscriptionCtx. + /// Create a client with no JS object and no SubscriptionCtx; also returns `tls.checkServerIdentity` for the JS object. /// /// This whole client needs a refactor. pub(crate) fn create_no_js_no_pubsub( cx: &bun_jsc::JsThread<'_>, arguments: &[JSValue], - ) -> JsResult<*mut JSValkeyClient> { + ) -> JsResult<(*mut JSValkeyClient, JSValue)> { let vm: &'static VirtualMachine = cx.global().bun_vm(); let vm_ref = vm; @@ -698,8 +705,9 @@ impl JSValkeyClient { bun_core::analytics::Features::VALKEY.fetch_add(1, core::sync::atomic::Ordering::Relaxed); + let check_server_identity = options.check_server_identity; // `_subscription_ctx` is a placeholder here; properly initialized later by `create()`. - Ok(JSValkeyClient::new(JSValkeyClient { + let client = JSValkeyClient::new(JSValkeyClient { ref_count: bun_ptr::RefCount::init(), _subscription_ctx: JsCell::new(SubscriptionCtx::default()), client: JsCell::new(valkey::ValkeyClient { @@ -753,7 +761,8 @@ impl JSValkeyClient { timer: RefCountedTimer::new(Timer::Tag::ValkeyConnectionTimeout), reconnect_timer: RefCountedTimer::new(Timer::Tag::ValkeyConnectionReconnect), context: cx.context().id(), - })) + }); + Ok((client, check_server_identity)) } pub(crate) fn create( @@ -761,12 +770,16 @@ impl JSValkeyClient { arguments: &[JSValue], js_this: JSValue, ) -> JsResult<*mut JSValkeyClient> { - let new_client_ptr = JSValkeyClient::create_no_js_no_pubsub(cx, arguments)?; + let (new_client_ptr, check_server_identity) = + JSValkeyClient::create_no_js_no_pubsub(cx, arguments)?; // SAFETY: just allocated above let new_client = unsafe { &*new_client_ptr }; // Initially, we only need to hold a weak reference to the JS object. new_client.this_value.set(JsRef::init_weak(js_this)); + if check_server_identity.is_callable() { + Js::check_server_identity_set_cached(js_this, cx.global(), check_server_identity); + } // Need to associate the subscription context, after the JS ref has been populated. new_client @@ -1749,6 +1762,14 @@ impl SocketHandler { if client.tls.reject_unauthorized(client.vm) { socket.set_inline_reject(); } + // RFC 6066: an IP literal is never sent as SNI. + let sni = Self::configured_hostname(this); + if !sni.is_empty() + && !bun_core::ip_address::is_ip_address(sni) + && let Some(ssl) = socket.ssl_mut() + { + ssl.set_servername(bun_core::ZBox::from_bytes(sni).as_cstr()); + } } this.client_mut().socket = Self::socket(socket); this.client_mut().on_open(Self::socket(socket)) @@ -1760,12 +1781,24 @@ impl SocketHandler { ssl: &mut boringssl::c::SSL, ) -> boringssl::ServerIdentity { let client = this.client.get(); - let rejects = client.tls.reject_unauthorized(client.vm); - let hostname = rejects.then(|| Self::identity_hostname(this, ssl)); + let native = client.tls.reject_unauthorized(client.vm) + && this.check_server_identity_callback().is_none(); + let hostname = native.then(|| Self::identity_hostname(this, ssl)); boringssl::server_identity(ssl, hostname.as_deref()) } - /// The name to match: the SNI servername, else the URL host. Empty for a unix socket, which has none. + /// `tls.serverName`, else the URL host without the brackets of an IPv6 literal. Empty for a unix socket with neither. + fn configured_hostname(this: &JSValkeyClient) -> &[u8] { + let client = this.client.get(); + let hostname = match (client.tls.server_name(), &client.address) { + (Some(server_name), _) => server_name, + (None, valkey::Address::Host { host, .. }) => &host[..], + (None, valkey::Address::Unix(_)) => b"", + }; + bun_core::ip_address::strip_ipv6_brackets(hostname) + } + + /// The name to match: the SNI servername, else the configured one. Empty for a unix socket, which has none. fn identity_hostname( this: &JSValkeyClient, ssl_ptr: *mut boringssl::c::SSL, @@ -1785,12 +1818,7 @@ impl SocketHandler { .to_vec() .into() } else { - match &this.client.get().address { - valkey::Address::Host { host, .. } => { - bun_core::ip_address::strip_ipv6_brackets(&host[..]).into() - } - valkey::Address::Unix(_) => (&b""[..]).into(), - } + Self::configured_hostname(this).into() } } @@ -1849,6 +1877,26 @@ impl SocketHandler { // Certificate chain is valid; verify the hostname matches the // certificate. let hostname = Self::identity_hostname(this, ssl_ptr); + if let Some(callback) = this.check_server_identity_callback() { + let verdict = tls_server_identity::check_with_callback( + &this.global_object, + callback, + socket.ssl_mut(), + &hostname, + ); + // User JS ran: the verdict is for `socket`, and the client may have closed it or dialed again. + let client = this.client.get(); + if client.status != valkey::Status::Connecting + || socket.is_closed() + || *client.socket.socket() != socket.socket + { + return Ok(()); + } + return match verdict { + Ok(()) => this.client_mut().start(), + Err(err) => Self::fail_handshake(this, vm, err), + }; + } // With no `SSL*` there is no certificate to match: fail closed. let identity_ok = hostname.is_empty() || (!ssl_ptr.is_null() @@ -2039,14 +2087,23 @@ impl Options { valkey::TLS::None }; } else if tls.is_object() { - // SAFETY: `bun_vm()` returns the live per-global VM pointer. - if let Some(ssl_config) = - SSLConfig::from_js(global_object.bun_vm(), global_object, tls)? + if let Some(callback) = tls.get(global_object, "checkServerIdentity")? + && !callback.is_undefined() { - this.tls = valkey::TLS::Custom(Box::new(ssl_config)); - } else { - return Err(global_object.throw_invalid_argument_type("tls", "tls", "object")); + if !callback.is_callable() { + return Err(global_object.throw_invalid_argument_type( + "tls", + "tls.checkServerIdentity", + "function", + )); + } + this.check_server_identity = callback; } + // An object with no recognized option still enables TLS, with defaults (as `tls: true`). + this.tls = match SSLConfig::from_js(global_object.bun_vm(), global_object, tls)? { + Some(ssl_config) => valkey::TLS::Custom(Box::new(ssl_config)), + None => valkey::TLS::Enabled, + }; } else { return Err(global_object.throw_invalid_argument_type( "tls", diff --git a/src/runtime/valkey_jsc/js_valkey_functions.rs b/src/runtime/valkey_jsc/js_valkey_functions.rs index aedb635a630e..890b7a30eda5 100644 --- a/src/runtime/valkey_jsc/js_valkey_functions.rs +++ b/src/runtime/valkey_jsc/js_valkey_functions.rs @@ -5,7 +5,7 @@ use bun_jsc::{ JsRef, JsResult, }; -use super::js_valkey::{JSValkeyClient, SubscriptionCtx}; +use super::js_valkey::{JSValkeyClient, Js, SubscriptionCtx}; use super::protocol_jsc as protocol; use super::valkey; use super::valkey_command_body::{Args as CommandArgs, Command, Meta as CommandMeta}; @@ -2093,6 +2093,9 @@ impl JSValkeyClient { new_client ._subscription_ctx .set(SubscriptionCtx::init(new_client)?); + if let Some(check_server_identity) = this.check_server_identity_callback() { + Js::check_server_identity_set_cached(new_client_js, global, check_server_identity); + } // If the original client is already connected and not manually closed, start connecting the new client. if this.client.get().status == valkey::Status::Connected && !this.client.get().flags.is_manually_closed diff --git a/src/runtime/valkey_jsc/valkey.classes.ts b/src/runtime/valkey_jsc/valkey.classes.ts index 7a3d5e311219..ac031e83a939 100644 --- a/src/runtime/valkey_jsc/valkey.classes.ts +++ b/src/runtime/valkey_jsc/valkey.classes.ts @@ -638,6 +638,6 @@ export default [ xgroup: { fn: "xgroup", length: 2 }, xsetid: { fn: "xsetid", length: 2 }, }, - values: ["onconnect", "onclose", "connectionPromise", "hello", "subscriptionCallbackMap"], + values: ["onconnect", "onclose", "connectionPromise", "hello", "subscriptionCallbackMap", "checkServerIdentity"], }), ]; diff --git a/src/runtime/valkey_jsc/valkey.rs b/src/runtime/valkey_jsc/valkey.rs index 0562b03f1a36..b54c58659c90 100644 --- a/src/runtime/valkey_jsc/valkey.rs +++ b/src/runtime/valkey_jsc/valkey.rs @@ -135,6 +135,16 @@ impl TLS { _ => false, } } + + /// `tls.serverName`: replaces the URL host as SNI and as the name the certificate must match. + pub(crate) fn server_name(&self) -> Option<&[u8]> { + match self { + TLS::Custom(ssl_config) => ssl_config + .server_name_bytes() + .filter(|name| !name.is_empty()), + _ => None, + } + } } // Call sites only ever compare against `TLS::None` / `TLS::Enabled`; `SSLConfig` @@ -155,6 +165,8 @@ pub(crate) struct Options { pub(crate) enable_auto_pipelining: bool, pub(crate) tls: TLS, + /// `tls.checkServerIdentity`, or `undefined`. Stored on the JS wrapper. + pub(crate) check_server_identity: JSValue, } impl Default for Options { @@ -167,6 +179,7 @@ impl Default for Options { enable_offline_queue: true, enable_auto_pipelining: true, tls: TLS::None, + check_server_identity: JSValue::UNDEFINED, } } } diff --git a/src/runtime/webcore/fetch/FetchTasklet.rs b/src/runtime/webcore/fetch/FetchTasklet.rs index a43dea07cf17..ad77ba6d777a 100644 --- a/src/runtime/webcore/fetch/FetchTasklet.rs +++ b/src/runtime/webcore/fetch/FetchTasklet.rs @@ -1249,28 +1249,7 @@ impl FetchTasklet { Err(e) => return Err(Some(global_object.take_exception(e))), }; - // > Returns object [...] on failure - // Any object counts: a DOMException or a util.inherits() error is not an ErrorInstance cell. - if check_result.is_object() && check_result.as_any_promise().is_none() { - return Err(Some(check_result)); - } - // Like Node, fail on any other truthy value, a Promise included: https://github.com/nodejs/node/blob/v26.3.0/lib/internal/tls/wrap.js#L1671-L1688 - if check_result.to_boolean() { - let received = JSGlobalObject::determine_specific_type(&global_object, check_result) - .map_err(|e| Some(global_object.take_exception(e)))?; - return Err(Some( - global_object - .err( - jsc::ErrorCode::INVALID_RETURN_VALUE, - format_args!( - "Expected undefined or an Error to be returned from the \"tls.checkServerIdentity\" function but got {received}." - ), - ) - .to_js(), - )); - } - // > On success, returns - Ok(()) + jsc::tls_server_identity::verdict_of(&global_object, check_result).map_err(Some) } /// Fail the request for a rejected certificate. Returns whether it may proceed. diff --git a/src/sql_jsc/jsc.rs b/src/sql_jsc/jsc.rs index 378017437dcb..d9b9a73cdc65 100644 --- a/src/sql_jsc/jsc.rs +++ b/src/sql_jsc/jsc.rs @@ -654,7 +654,7 @@ pub use bun_jsc::JsClass; pub mod codegen { ::bun_jsc::js_class_module!(JSPostgresSQLConnection = "PostgresSQLConnection" - as crate::postgres::PostgresSQLConnection { queries, onconnect, onclose, onnotification }); + as crate::postgres::PostgresSQLConnection { queries, onconnect, onclose, checkServerIdentity, onnotification }); ::bun_jsc::js_class_module!( JSPostgresSQLQuery = "PostgresSQLQuery" as crate::postgres::PostgresSQLQuery, impl_js_class { @@ -666,7 +666,7 @@ pub mod codegen { ); ::bun_jsc::js_class_module!(js_mysql_connection = "MySQLConnection" - as crate::mysql::js_my_sql_connection::JSMySQLConnection { queries, onconnect, onclose }); + as crate::mysql::js_my_sql_connection::JSMySQLConnection { queries, onconnect, onclose, checkServerIdentity }); pub use js_mysql_connection as JSMySQLConnection; ::bun_jsc::js_class_module!( diff --git a/src/sql_jsc/mysql/JSMySQLConnection.rs b/src/sql_jsc/mysql/JSMySQLConnection.rs index 9ada804b5188..53a54b16f00e 100644 --- a/src/sql_jsc/mysql/JSMySQLConnection.rs +++ b/src/sql_jsc/mysql/JSMySQLConnection.rs @@ -30,6 +30,7 @@ use crate::mysql::protocol::error_packet_jsc::ErrorPacketJsc; use super::my_sql_connection::{self as my_sql_connection}; use super::my_sql_statement::MySQLStatement; use super::protocol::result_set::{self as ResultSet}; +use bun_jsc::tls_server_identity; bun_core::declare_scope!(MySQLConnection, visible); @@ -106,6 +107,13 @@ impl JSMySQLConnection { self.connection.get().server_identity(ssl) } + fn check_server_identity_callback(&self) -> Option { + self.js_value + .get() + .try_get() + .and_then(js::check_server_identity_get_cached) + } + /// Hold a ref on `self` for the guard's lifetime (across re-entrant calls). #[inline] fn ref_guard(&self) -> RefPtr { @@ -484,6 +492,7 @@ impl JSMySQLConnection { tls_config, secure, args.ssl_mode, + args.check_server_identity.is_callable(), allow_public_key_retrieval, )), auto_flusher: JsCell::new(AutoFlusher::default()), @@ -553,6 +562,13 @@ impl JSMySQLConnection { .with_mut(|r| r.set_strong(js_value, global_object)); js::onconnect_set_cached(js_value, global_object, on_connect); js::onclose_set_cached(js_value, global_object, on_close); + if args.check_server_identity.is_callable() { + js::check_server_identity_set_cached( + js_value, + global_object, + args.check_server_identity, + ); + } Ok(js_value) } @@ -894,10 +910,36 @@ impl SocketHandler { fn on_handshake( this: &JSMySQLConnection, - _: NewSocketHandler, + socket: NewSocketHandler, success: i32, ssl_error: uws::us_bun_verify_error_t, ) { + let callback_hostname = this + .connection + .get() + .callback_identity_hostname() + .map(<[u8]>::to_vec); + if success == 1 + && ssl_error.error_no == 0 + && let Some(hostname) = callback_hostname + && let Some(callback) = this.check_server_identity_callback() + { + // User JS: it can close this connection. + let _guard = this.ref_guard(); + let verdict = tls_server_identity::check_with_callback( + &this.global_object, + callback, + socket.ssl_mut(), + &hostname, + ); + if let Err(err) = verdict { + this.connection_mut().reject_server_identity(); + return this.fail_with_js_value(err); + } + if !this.connection.get().is_active() { + return; + } + } let handshake_was_successful = match this.connection_mut().do_handshake(success, ssl_error) { Ok(v) => v, diff --git a/src/sql_jsc/mysql/MySQLConnection.rs b/src/sql_jsc/mysql/MySQLConnection.rs index f15ebf7ed83e..13492e430fcb 100644 --- a/src/sql_jsc/mysql/MySQLConnection.rs +++ b/src/sql_jsc/mysql/MySQLConnection.rs @@ -85,6 +85,8 @@ pub struct MySQLConnection { tls_config: SSLConfig, tls_status: TLSStatus, ssl_mode: SSLMode, + /// `tls.checkServerIdentity` is set: it replaces the native name check, and runs after the handshake. + has_check_server_identity: bool, allow_public_key_retrieval: bool, flags: ConnectionFlags, } @@ -117,6 +119,7 @@ impl Default for MySQLConnection { tls_config: SSLConfig::default(), tls_status: TLSStatus::None, ssl_mode: SSLMode::Disable, + has_check_server_identity: false, allow_public_key_retrieval: false, flags: ConnectionFlags::default(), } @@ -135,6 +138,7 @@ impl MySQLConnection { tls_config: SSLConfig, secure: Option, ssl_mode: SSLMode, + has_check_server_identity: bool, allow_public_key_retrieval: bool, ) -> Self { Self { @@ -147,6 +151,7 @@ impl MySQLConnection { tls_config, secure, ssl_mode, + has_check_server_identity, allow_public_key_retrieval, tls_status: if ssl_mode != SSLMode::Disable { TLSStatus::Pending @@ -396,10 +401,24 @@ impl MySQLConnection { /// The name verify-full matches, in and after the handshake. Empty (none configured) matches no certificate. fn native_identity_hostname(&self) -> Option<&[u8]> { - (self.tls_config.reject_unauthorized() != 0 && self.ssl_mode == SSLMode::VerifyFull) + (self.tls_config.reject_unauthorized() != 0 + && self.ssl_mode == SSLMode::VerifyFull + && !self.has_check_server_identity) .then(|| self.tls_config.server_name_bytes()) } + /// The name `tls.checkServerIdentity` is asked about, when it decides: it is set and the ssl mode verifies the peer. + pub(crate) fn callback_identity_hostname(&self) -> Option<&[u8]> { + (self.has_check_server_identity + && self.tls_config.reject_unauthorized() != 0 + && matches!(self.ssl_mode, SSLMode::VerifyCa | SSLMode::VerifyFull)) + .then(|| self.tls_config.server_name_bytes()) + } + + pub(crate) fn reject_server_identity(&mut self) { + self.tls_status = TLSStatus::SslFailed; + } + pub(crate) fn do_handshake( &mut self, success: i32, diff --git a/src/sql_jsc/postgres/PostgresSQLConnection.rs b/src/sql_jsc/postgres/PostgresSQLConnection.rs index 39624b865b0d..16ed5d22ec22 100644 --- a/src/sql_jsc/postgres/PostgresSQLConnection.rs +++ b/src/sql_jsc/postgres/PostgresSQLConnection.rs @@ -1,5 +1,6 @@ use bun_collections::VecExt; use bun_jsc::JsCell; +use bun_jsc::tls_server_identity; use core::cell::Cell; use core::ffi::c_void; use core::sync::atomic::{AtomicU32, Ordering}; @@ -886,7 +887,10 @@ impl PostgresSQLConnection { /// verify-full's name check, asked inside the handshake. pub fn server_identity(&self, ssl: &mut bun_boringssl_sys::SSL) -> BoringSSL::ServerIdentity { - BoringSSL::server_identity(ssl, self.native_identity_hostname()) + let hostname = self + .native_identity_hostname() + .filter(|_| self.check_server_identity_callback().is_none()); + BoringSSL::server_identity(ssl, hostname) } /// The name verify-full matches, in and after the handshake. Empty (none configured) matches no certificate. @@ -895,6 +899,14 @@ impl PostgresSQLConnection { .then(|| self.tls_config.server_name_bytes()) } + /// `tls.checkServerIdentity`: replaces the native name check, and runs after the handshake. + fn check_server_identity_callback(&self) -> Option { + self.js_value + .get() + .try_get() + .and_then(js::check_server_identity_get_cached) + } + pub(crate) fn on_handshake(&self, success: i32, ssl_error: uws::us_bun_verify_error_t) { debug!("onHandshake: {} {}", success, ssl_error.error_no); let handshake_success = success == 1; @@ -910,7 +922,19 @@ impl PostgresSQLConnection { return; } - if let Some(hostname) = self.native_identity_hostname() { + if let Some(callback) = self.check_server_identity_callback() { + // User JS: it can close this connection. + let _guard = self.ref_guard(); + let verdict = tls_server_identity::check_with_callback( + self.global(), + callback, + self.socket.get().ssl_mut(), + self.tls_config.server_name_bytes(), + ); + if let Err(err) = verdict { + self.fail_with_js_value(err); + } + } else if let Some(hostname) = self.native_identity_hostname() { // SAFETY: native handle of a connected TLS socket is `SSL*`. let ssl_ptr: *mut BoringSSL::c::SSL = self .socket @@ -1288,6 +1312,9 @@ pub(crate) fn call(global_object: &JSGlobalObject, callframe: &CallFrame) -> JsR this.js_value.set(crate::jsc::JsRef::init_weak(js_value)); js::onconnect_set_cached(js_value, global_object, on_connect); js::onclose_set_cached(js_value, global_object, on_close); + if args.check_server_identity.is_callable() { + js::check_server_identity_set_cached(js_value, global_object, args.check_server_identity); + } bun_analytics::features::postgres_connections.fetch_add(1, Ordering::Relaxed); Ok(js_value) } diff --git a/src/sql_jsc/shared/ConnectionCtorArgs.rs b/src/sql_jsc/shared/ConnectionCtorArgs.rs index d41985923f73..e1c945cb83c3 100644 --- a/src/sql_jsc/shared/ConnectionCtorArgs.rs +++ b/src/sql_jsc/shared/ConnectionCtorArgs.rs @@ -42,6 +42,8 @@ pub(crate) struct ConnectionCtorArgs { pub database_str: bun_core::String, pub ssl_mode: M, pub tls_config: SSLConfig, + /// `tls.checkServerIdentity`, or `undefined`. Kept alive by `arguments`. + pub check_server_identity: JSValue, /// Moves into the connection. pub secure: Option, } @@ -68,11 +70,18 @@ impl ConnectionCtorArgs { let tls_object = arguments[6]; let mut tls_config = SSLConfig::default(); + let mut check_server_identity = JSValue::UNDEFINED; let mut secure: Option = None; if ssl_mode != modes[0] { tls_config = if tls_object.is_boolean() && tls_object.to_boolean() { SSLConfig::default() } else if tls_object.is_object() { + // Validated as a function (or absent) by the JS layer. + if let Some(callback) = tls_object.get(global_object, "checkServerIdentity")? + && callback.is_callable() + { + check_server_identity = callback; + } match SSLConfig::from_js(&mut *vm, global_object, tls_object) { Ok(opt) => opt.unwrap_or_default(), Err(_) => return Ok(None), @@ -114,6 +123,7 @@ impl ConnectionCtorArgs { database_str, ssl_mode, tls_config, + check_server_identity, secure, })) } diff --git a/src/uws_sys/socket.rs b/src/uws_sys/socket.rs index 933188762a57..053d59d25c57 100644 --- a/src/uws_sys/socket.rs +++ b/src/uws_sys/socket.rs @@ -957,6 +957,15 @@ impl AnySocket { } } + /// The `SSL` handle of a connected TLS socket; `None` for TCP. + #[inline] + pub fn ssl_mut(&self) -> Option<&mut bun_boringssl_sys::SSL> { + match self { + AnySocket::SocketTcp(_) => None, + AnySocket::SocketTls(s) => s.ssl_mut(), + } + } + any_socket_forward! { fn is_closed(&self) -> bool; fn is_shutdown(&self) -> bool; diff --git a/test/js/sql/sql-tls-server-identity-worker-fixture.ts b/test/js/sql/sql-tls-server-identity-worker-fixture.ts new file mode 100644 index 000000000000..f7aa0b698e5b --- /dev/null +++ b/test/js/sql/sql-tls-server-identity-worker-fixture.ts @@ -0,0 +1,25 @@ +// A worker that never leaves `tls.checkServerIdentity`: the parent terminates it there. +import { SQL } from "bun"; +import { workerData } from "node:worker_threads"; + +const { url, ca, counters } = workerData as { url: string; ca: string; counters: SharedArrayBuffer }; +const count = new Int32Array(counters); // [0] callback entered, [1] onclose ran, [2] connect() settled + +const sql = new SQL({ + url, + max: 1, + onclose: () => void Atomics.add(count, 1, 1), + tls: { + ca, + checkServerIdentity: () => { + Atomics.add(count, 0, 1); + Atomics.notify(count, 0); + for (;;) {} + }, + }, +}); +const settled = () => { + Atomics.add(count, 2, 1); + Atomics.notify(count, 0); +}; +sql.connect().then(settled, settled); diff --git a/test/js/sql/sql-tls-server-identity.test.ts b/test/js/sql/sql-tls-server-identity.test.ts new file mode 100644 index 000000000000..105bca62df1e --- /dev/null +++ b/test/js/sql/sql-tls-server-identity.test.ts @@ -0,0 +1,422 @@ +// `tls.checkServerIdentity`, `tls.serverName` / `tls.servername` and the +// built-in hostname check for Bun.SQL over TLS (PostgreSQL and MySQL). +// +// These need a server that presents a certificate for a name of the test's +// choosing, which the shared containers cannot do, so both adapters talk to a +// minimal mock that upgrades to TLS and accepts the login. All wire-protocol +// bytes come from test/js/sql/wire-frames.ts. + +import { SQL } from "bun"; +import { describe, expect, test } from "bun:test"; +import { tls as localhostTls } from "harness"; +import { X509Certificate } from "node:crypto"; +import { once } from "node:events"; +import fs from "node:fs"; +import type net from "node:net"; +import path from "node:path"; +import tls from "node:tls"; +import { Worker } from "node:worker_threads"; +import { + MYSQL_CLIENT_SSL, + MYSQL_DEFAULT_CAPABILITIES, + listeningServer, + mysqlAckSessionSetup, + mysqlHandshakeV10, + mysqlOkPacket, + mysqlReadPackets, + pgAuthenticationOk, + pgReadyForQuery, + pgSSLResponse, +} from "./wire-frames"; + +// CN=agent1 (no SAN), signed by ca1: trusted through `ca1` but valid for no +// host the tests dial, so only `serverName: "agent1"` or a custom +// `checkServerIdentity` can accept it. +const fixturesDir = path.join(import.meta.dirname, "..", "node", "tls", "fixtures"); +const agent1 = { + key: fs.readFileSync(path.join(fixturesDir, "agent1-key.pem"), "utf8"), + cert: fs.readFileSync(path.join(fixturesDir, "agent1-cert.pem"), "utf8"), + ca: fs.readFileSync(path.join(fixturesDir, "ca1-cert.pem"), "utf8"), +}; +// The harness certificate: CN=server-bun, SAN localhost / 127.0.0.1 / ::1, self-signed. +const localhost = { key: localhostTls.key, cert: localhostTls.cert, ca: localhostTls.cert }; + +type ServerCert = { key: string; cert: string }; +type MockServer = { + url: string; + /** SNI of every completed TLS handshake, in order. */ + servernames: (string | false)[]; + /** Set by a test that needs them: called with the bytes a connection sent inside TLS, when it closes. */ + onTlsClose?: (bytesFromClient: number) => void; + close(): void; +}; + +/** + * Wraps `rawSocket` in a server-side TLSSocket once the plaintext prelude is + * done. Bytes already buffered past the prelude are TLS records: hand them to + * the TLS engine instead of the plaintext parser. + */ +function upgrade(rawSocket: net.Socket, cert: ServerCert, leftover: Buffer, mock: MockServer) { + rawSocket.pause(); + if (leftover.length) rawSocket.unshift(leftover); + const socket = new tls.TLSSocket(rawSocket, { isServer: true, ...cert }); + let bytesFromClient = 0; + socket.on("secure", () => mock.servernames.push(socket.servername)); + socket.on("data", (chunk: Buffer) => (bytesFromClient += chunk.length)); + socket.on("close", () => mock.onTlsClose?.(bytesFromClient)); + socket.on("error", () => {}); + return socket; +} + +/** Answers SSLRequest with 'S', upgrades, then accepts any StartupMessage. */ +async function postgresServer(cert: ServerCert): Promise { + const mock: MockServer = { url: "", servernames: [], close: () => {} }; + const { server, port } = await listeningServer(rawSocket => { + rawSocket.on("error", () => {}); + let buffered = Buffer.alloc(0); + const onPlainData = (chunk: Buffer) => { + // SSLRequest is Int32(8) Int32(80877103). The client sends nothing else + // until it has the one-byte answer, then its ClientHello. + buffered = Buffer.concat([buffered, chunk]); + if (buffered.length < 8) return; + rawSocket.removeListener("data", onPlainData); + rawSocket.write(pgSSLResponse("S")); + const socket = upgrade(rawSocket, cert, buffered.subarray(8), mock); + let startup = true; + socket.on("data", () => { + if (startup) { + startup = false; + socket.write(Buffer.concat([pgAuthenticationOk(), pgReadyForQuery()])); + } + }); + }; + rawSocket.on("data", onPlainData); + }); + mock.url = `postgres://u@127.0.0.1:${port}/db`; + mock.close = () => server.close(); + return mock; +} + +/** Advertises CLIENT_SSL, upgrades after the SSLRequest packet, then accepts the login. */ +async function mysqlServer(cert: ServerCert): Promise { + const mock: MockServer = { url: "", servernames: [], close: () => {} }; + const { server, port } = await listeningServer(rawSocket => { + rawSocket.on("error", () => {}); + rawSocket.write(mysqlHandshakeV10({ capabilities: MYSQL_DEFAULT_CAPABILITIES | MYSQL_CLIENT_SSL })); + let buffered = Buffer.alloc(0); + const onPlainData = (chunk: Buffer) => { + buffered = Buffer.concat([buffered, chunk]); + if (buffered.length < 4) return; + const length = buffered[0] | (buffered[1] << 8) | (buffered[2] << 16); + if (buffered.length < 4 + length) return; + // The SSLRequest packet; the ClientHello may already follow it. + const leftover = buffered.subarray(4 + length); + buffered = Buffer.alloc(0); + rawSocket.removeListener("data", onPlainData); + const socket = upgrade(rawSocket, cert, leftover, mock); + let authed = false; + socket.on("data", (chunk: Buffer) => { + buffered = mysqlReadPackets(Buffer.concat([buffered, chunk]), (seq, payload) => { + if (!authed) { + authed = true; + socket.write(mysqlOkPacket(seq + 1)); + return; + } + if (!mysqlAckSessionSetup(socket, payload)) socket.end(); + }); + }); + }; + rawSocket.on("data", onPlainData); + }); + mock.url = `mysql://u@127.0.0.1:${port}/db`; + mock.close = () => server.close(); + return mock; +} + +async function connect(url: string, tlsOptions: Bun.SQL.Options["tls"], sslmode = "verify-full"): Promise { + await using sql = new SQL({ url: `${url}?sslmode=${sslmode}`, tls: tlsOptions, max: 1, idleTimeout: 1 }); + return await sql.connect().then( + () => "CONNECTED", + e => e, + ); +} + +describe.each([ + ["PostgreSQL", "postgres", postgresServer], + ["MySQL", "mysql", mysqlServer], +] as const)("%s TLS server identity", (_, scheme, startServer) => { + async function withServer(cert: ServerCert, fn: (server: MockServer) => Promise): Promise { + const server = await startServer(cert); + try { + return await fn(server); + } finally { + server.close(); + } + } + + test("without tls.checkServerIdentity, verify-full rejects a trusted certificate issued for another host", async () => { + await withServer(agent1, async server => { + const err: any = await connect(server.url, { ca: agent1.ca }); + expect(err?.code).toBe("ERR_TLS_CERT_ALTNAME_INVALID"); + }); + }); + + // `servername` is the node:tls spelling of the same option. + test.each(["serverName", "servername"] as const)( + "tls.%s sets the SNI and the name the certificate is verified against", + async key => { + await withServer(agent1, async server => { + expect(await connect(server.url, { ca: agent1.ca, [key]: "agent1" })).toBe("CONNECTED"); + expect(server.servernames).toEqual(["agent1"]); + }); + }, + ); + + test("tls.checkServerIdentity is called with the hostname and certificate, and the Error it returns fails the connection", async () => { + await withServer(localhost, async server => { + const calls: [string, string, string][] = []; + const pin = new Error("PIN_MISMATCH"); + const err = await connect(server.url, { + ca: localhost.ca, + checkServerIdentity: (hostname: string, cert: tls.PeerCertificate) => { + calls.push([hostname, cert.subject.CN, cert.fingerprint256]); + return pin; + }, + }); + expect(err).toBe(pin); + const fingerprint256 = new X509Certificate(localhost.cert).fingerprint256; + expect(calls).toEqual([["127.0.0.1", "server-bun", fingerprint256]]); + }); + }); + + test("tls.checkServerIdentity replaces the built-in hostname check when it returns undefined", async () => { + await withServer(agent1, async server => { + const calls: [string, string][] = []; + const result = await connect(server.url, { + ca: agent1.ca, + checkServerIdentity: (hostname: string, cert: tls.PeerCertificate) => { + calls.push([hostname, cert.subject.CN]); + return undefined; + }, + }); + expect(result).toBe("CONNECTED"); + expect(calls).toEqual([["127.0.0.1", "agent1"]]); + }); + }); + + test("tls.checkServerIdentity receives tls.serverName as the hostname", async () => { + await withServer(agent1, async server => { + const hostnames: string[] = []; + const result = await connect(server.url, { + ca: agent1.ca, + serverName: "agent1", + checkServerIdentity: (hostname: string) => { + hostnames.push(hostname); + return undefined; + }, + }); + expect(result).toBe("CONNECTED"); + expect(hostnames).toEqual(["agent1"]); + }); + }); + + test("an exception thrown by tls.checkServerIdentity fails the connection", async () => { + await withServer(localhost, async server => { + const thrown = new TypeError("from checkServerIdentity"); + const err = await connect(server.url, { + ca: localhost.ca, + checkServerIdentity: () => { + throw thrown; + }, + }); + expect(err).toBe(thrown); + }); + }); + + test("tls.checkServerIdentity that is a method of a class is called", async () => { + await withServer(localhost, async server => { + const pin = new Error("PIN_MISMATCH"); + class PinnedTls { + ca = localhost.ca; + checkServerIdentity() { + return pin; + } + } + expect(await connect(server.url, new PinnedTls())).toBe(pin); + }); + }); + + test("tls.checkServerIdentity also runs under sslmode=verify-ca", async () => { + await withServer(agent1, async server => { + const pin = new Error("PIN_MISMATCH"); + expect(await connect(server.url, { ca: agent1.ca, checkServerIdentity: () => pin }, "verify-ca")).toBe(pin); + // Without a callback verify-ca does not check the hostname. + expect(await connect(server.url, { ca: agent1.ca }, "verify-ca")).toBe("CONNECTED"); + }); + }); + + test("tls.checkServerIdentity requests certificate verification like tls.ca does", async () => { + // No `ca` and no verify-* sslmode: the callback alone turns verification + // on, so the untrusted self-signed chain fails before the callback runs. + await withServer(localhost, async server => { + let calls = 0; + const err: any = await connect( + server.url, + { + checkServerIdentity: () => { + calls++; + return undefined; + }, + }, + "prefer", + ); + expect(err?.code).toBe("DEPTH_ZERO_SELF_SIGNED_CERT"); + expect(calls).toBe(0); + }); + }); + + test("tls.checkServerIdentity is not called when rejectUnauthorized is false", async () => { + await withServer(agent1, async server => { + let calls = 0; + const result = await connect( + server.url, + { + ca: agent1.ca, + rejectUnauthorized: false, + checkServerIdentity: () => { + calls++; + return new Error("unreachable"); + }, + }, + "require", + ); + expect(result).toBe("CONNECTED"); + expect(calls).toBe(0); + }); + }); + + test("a worker terminated inside tls.checkServerIdentity stops, and sends nothing to the server", async () => { + await withServer(localhost, async server => { + const { promise: tlsClosed, resolve } = Promise.withResolvers(); + server.onTlsClose = resolve; + const counters = new SharedArrayBuffer(12); + const count = new Int32Array(counters); + const worker = new Worker(new URL("./sql-tls-server-identity-worker-fixture.ts", import.meta.url), { + workerData: { url: `${server.url}?sslmode=verify-full`, ca: localhost.ca, counters }, + }); + const exited = once(worker, "exit"); + await Promise.race([ + Atomics.waitAsync(count, 0, 0).value, + once(worker, "error").then(([error]) => Promise.reject(error)), + exited.then(([code]) => Promise.reject(new Error(`the worker exited with code ${code} before the callback`))), + ]); + // The fixture also wakes this wait when connect() settles: then the callback was never called. + expect({ callbackEntered: count[0], connectSettled: count[2] }).toEqual({ + callbackEntered: 1, + connectSettled: 0, + }); + await worker.terminate(); + const [bytesFromClient] = await Promise.all([tlsClosed, exited]); + expect({ callbackEntered: count[0], oncloseRan: count[1], connectSettled: count[2], bytesFromClient }).toEqual({ + callbackEntered: 1, + oncloseRan: 0, + connectSettled: 0, + bytesFromClient: 0, + }); + }); + }); + + test("a tls.checkServerIdentity that closes the client leaves it closed", async () => { + await withServer(localhost, async server => { + let calls = 0; + const sql = new SQL({ + url: `${server.url}?sslmode=verify-full`, + max: 1, + idleTimeout: 1, + tls: { + ca: localhost.ca, + checkServerIdentity: () => { + calls++; + void sql.close(); + return undefined; + }, + }, + }); + // The pool decides how a connect in flight settles; the client must end up closed. + await sql.connect().catch(() => {}); + const err: any = await sql`select 1`.then( + () => null, + e => e, + ); + expect({ calls, code: err?.code }).toEqual({ calls: 1, code: `ERR_${scheme.toUpperCase()}_CONNECTION_CLOSED` }); + }); + }); + + // Node refuses the server for every truthy return value, so an `async` callback cannot accept it by accident. + test.each([ + ["a Promise (async callback)", async () => undefined, "an instance of Promise"], + ["a string", () => "PIN_MISMATCH", "type string ('PIN_MISMATCH')"], + ["true", () => true, "type boolean (true)"], + ] as const)("tls.checkServerIdentity that returns %s fails the connection", async (_, callback, received) => { + await withServer(localhost, async server => { + const err: any = await connect(server.url, { ca: localhost.ca, checkServerIdentity: callback as any }); + expect({ name: err?.name, code: err?.code, message: err?.message }).toEqual({ + name: "TypeError", + code: "ERR_INVALID_RETURN_VALUE", + message: `Expected undefined or an Error to be returned from the "tls.checkServerIdentity" function but got ${received}.`, + }); + }); + }); + + test("an object that tls.checkServerIdentity returns fails the connection", async () => { + await withServer(localhost, async server => { + // Not an Error instance. The adapter reports every failure as its own Error class and keeps the fields. + const err: any = await connect(server.url, { + ca: localhost.ca, + checkServerIdentity: (() => ({ code: "PIN_MISMATCH" })) as any, + }); + expect({ isError: err instanceof Error, code: err?.code }).toEqual({ isError: true, code: "PIN_MISMATCH" }); + }); + }); + + test.each([null, false, 0, ""])("tls.checkServerIdentity that returns %p accepts the server", async value => { + await withServer(localhost, async server => { + expect(await connect(server.url, { ca: localhost.ca, checkServerIdentity: (() => value) as any })).toBe( + "CONNECTED", + ); + }); + }); + + test("tls.checkServerIdentity receives the chain of getPeerCertificate(true)", async () => { + await withServer(agent1, async server => { + const chain: string[] = []; + let rootIsItsOwnIssuer = false; + const result = await connect(server.url, { + ca: agent1.ca, + checkServerIdentity: (_hostname: string, cert: tls.PeerCertificate) => { + // The walk of Node's certificate pinning example. + let current = cert as tls.DetailedPeerCertificate; + let last: string; + do { + chain.push(current.subject.CN); + last = current.fingerprint256; + current = current.issuerCertificate; + } while (current.fingerprint256 !== last); + rootIsItsOwnIssuer = current.issuerCertificate === current; + return undefined; + }, + }); + expect({ result, chain, rootIsItsOwnIssuer }).toEqual({ + result: "CONNECTED", + chain: ["agent1", "ca1"], + rootIsItsOwnIssuer: true, + }); + }); + }); + + test("tls.checkServerIdentity must be a function", () => { + expect(() => new SQL({ url: `${scheme}://u@127.0.0.1:1/db`, tls: { checkServerIdentity: "nope" as any } })).toThrow( + expect.objectContaining({ code: "ERR_INVALID_ARG_TYPE" }), + ); + }); +}); diff --git a/test/js/valkey/valkey-tls-verify-worker-fixture.ts b/test/js/valkey/valkey-tls-verify-worker-fixture.ts new file mode 100644 index 000000000000..236910a18dcc --- /dev/null +++ b/test/js/valkey/valkey-tls-verify-worker-fixture.ts @@ -0,0 +1,24 @@ +// A worker that never leaves `tls.checkServerIdentity`: the parent terminates it there. +import { RedisClient } from "bun"; +import { workerData } from "node:worker_threads"; + +const { port, ca, counters } = workerData as { port: number; ca: string; counters: SharedArrayBuffer }; +const count = new Int32Array(counters); // [0] callback entered, [1] onclose ran, [2] command settled + +const client = new RedisClient(`rediss://localhost:${port}`, { + autoReconnect: false, + tls: { + ca, + checkServerIdentity: () => { + Atomics.add(count, 0, 1); + Atomics.notify(count, 0); + for (;;) {} + }, + }, +}); +client.onclose = () => void Atomics.add(count, 1, 1); +const settled = () => { + Atomics.add(count, 2, 1); + Atomics.notify(count, 0); +}; +client.send("PING", []).then(settled, settled); diff --git a/test/js/valkey/valkey-tls-verify.test.ts b/test/js/valkey/valkey-tls-verify.test.ts index a1df1d287b4e..f9505e7957ef 100644 --- a/test/js/valkey/valkey-tls-verify.test.ts +++ b/test/js/valkey/valkey-tls-verify.test.ts @@ -6,6 +6,7 @@ import fs from "node:fs"; import type { AddressInfo } from "node:net"; import path from "node:path"; import tls from "node:tls"; +import { Worker } from "node:worker_threads"; // Server presents a cert for CN=agent1 (no SAN), signed by ca1. // A client connecting to host "localhost" with ca1 trusted will pass chain @@ -56,17 +57,37 @@ function fakeServer(serverOpts: tls.TlsOptions): tls.Server { return server; } -async function withServer(serverOpts: tls.TlsOptions, fn: (port: number) => Promise, host?: string): Promise { +async function withServer( + serverOpts: tls.TlsOptions, + fn: (port: number, server: tls.Server) => Promise, + host?: string, +): Promise { const server = fakeServer(serverOpts); server.listen(0, host); await once(server, "listening"); try { - return await fn((server.address() as AddressInfo).port); + return await fn((server.address() as AddressInfo).port, server); } finally { server.close(); } } +/** The SNI names the server saw, one entry per completed handshake. */ +function recordServernames(server: tls.Server): (string | false)[] { + const names: (string | false)[] = []; + server.on("secureConnection", socket => names.push(socket.servername)); + return names; +} + +async function ping(url: string, tlsOptions: Bun.RedisOptions["tls"]): Promise { + const client = new RedisClient(url, { autoReconnect: false, connectionTimeout: 5000, tls: tlsOptions }); + try { + return await client.send("PING", []); + } finally { + client.close(); + } +} + async function withUnixServer(serverOpts: tls.TlsOptions, fn: (socketPath: string) => Promise): Promise { using dir = tempDir("valkey-tls-unix", {}); const socketPath = path.join(String(dir), "r.sock"); @@ -240,3 +261,355 @@ describe("RedisClient TLS hostname verification", () => { }); }); }); + +describe("RedisClient tls.serverName", () => { + test("sends the URL host as SNI by default", async () => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async (port, server) => { + const servernames = recordServernames(server); + expect(await ping(`rediss://localhost:${port}`, { ca: localhostTls.cert })).toBe("PONG"); + expect(servernames).toEqual(["localhost"]); + }); + }); + + test("does not send an IP literal as SNI", async () => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async (port, server) => { + const servernames = recordServernames(server); + expect(await ping(`rediss://127.0.0.1:${port}`, { ca: localhostTls.cert })).toBe("PONG"); + // One handshake, no SNI. + expect(servernames.map(Boolean)).toEqual([false]); + }); + }); + + // `servername` is the node:tls spelling of the same option. + test.each(["serverName", "servername"] as const)( + "%s sets the SNI and the name the certificate is verified against", + async key => { + // The server presents CN=agent1; dialing 127.0.0.1 only verifies when the + // client checks the certificate against "agent1" instead of the URL host. + await withServer({ key: serverKey, cert: serverCert }, async (port, server) => { + const servernames = recordServernames(server); + expect(await ping(`rediss://127.0.0.1:${port}`, { ca, [key]: "agent1" })).toBe("PONG"); + expect(servernames).toEqual(["agent1"]); + }); + }, + ); + + test("serverName that does not match the certificate is rejected", async () => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async port => { + const err: any = await ping(`rediss://localhost:${port}`, { ca: localhostTls.cert, serverName: "agent1" }).then( + () => null, + e => e, + ); + expect(err?.code).toBe("ERR_TLS_CERT_ALTNAME_INVALID"); + expect(err.message).toContain("agent1"); + }); + }); +}); + +describe("RedisClient tls.checkServerIdentity", () => { + test("is called with the hostname and certificate, and the Error it returns rejects the connection", async () => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async port => { + const calls: [string, string][] = []; + const pin = new Error("PIN_MISMATCH"); + const err = await ping(`rediss://localhost:${port}`, { + ca: localhostTls.cert, + checkServerIdentity: (hostname: string, cert: tls.PeerCertificate) => { + calls.push([hostname, cert.subject.CN]); + return pin; + }, + }).then( + () => null, + e => e, + ); + expect(err).toBe(pin); + // The harness certificate: CN=server-bun, SAN localhost/127.0.0.1/::1. + expect(calls).toEqual([["localhost", "server-bun"]]); + }); + }); + + test("replaces the built-in hostname check when it returns undefined", async () => { + // CN=agent1 does not match "localhost": only the callback can accept it. + await withServer({ key: serverKey, cert: serverCert }, async port => { + const calls: [string, string][] = []; + const result = await ping(`rediss://localhost:${port}`, { + ca, + checkServerIdentity: (hostname: string, cert: tls.PeerCertificate) => { + calls.push([hostname, cert.subject.CN]); + return undefined; + }, + }); + expect(result).toBe("PONG"); + expect(calls).toEqual([["localhost", "agent1"]]); + }); + }); + + test("receives tls.serverName as the hostname", async () => { + await withServer({ key: serverKey, cert: serverCert }, async port => { + const hostnames: string[] = []; + const result = await ping(`rediss://127.0.0.1:${port}`, { + ca, + serverName: "agent1", + checkServerIdentity: (hostname: string) => { + hostnames.push(hostname); + return undefined; + }, + }); + expect(result).toBe("PONG"); + expect(hostnames).toEqual(["agent1"]); + }); + }); + + test("an exception thrown by the callback rejects the connection", async () => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async port => { + const thrown = new TypeError("from checkServerIdentity"); + const err = await ping(`rediss://localhost:${port}`, { + ca: localhostTls.cert, + checkServerIdentity: () => { + throw thrown; + }, + }).then( + () => null, + e => e, + ); + expect(err).toBe(thrown); + }); + }); + + test("is not called when the certificate chain fails to verify", async () => { + // No `ca`: the self-signed localhost certificate is untrusted. A tls object + // whose only member is the callback still enables TLS with default options. + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async port => { + let calls = 0; + const err: any = await ping(`rediss://localhost:${port}`, { + checkServerIdentity: () => { + calls++; + return undefined; + }, + }).then( + () => null, + e => e, + ); + expect(err?.code).toBe("DEPTH_ZERO_SELF_SIGNED_CERT"); + expect(calls).toBe(0); + }); + }); + + test("is not called when rejectUnauthorized is false", async () => { + await withServer({ key: serverKey, cert: serverCert }, async port => { + let calls = 0; + const result = await ping(`rediss://localhost:${port}`, { + ca, + rejectUnauthorized: false, + checkServerIdentity: () => { + calls++; + return new Error("unreachable"); + }, + }); + expect(result).toBe("PONG"); + expect(calls).toBe(0); + }); + }); + + test("must be a function", () => { + expect(() => new RedisClient("rediss://localhost:6379", { tls: { checkServerIdentity: "nope" as any } })).toThrow( + expect.objectContaining({ code: "ERR_INVALID_ARG_TYPE" }), + ); + }); + + // Node refuses the server for every truthy return value, so an `async` callback cannot accept it by accident. + test.each([ + ["a Promise (async callback)", async () => undefined, "an instance of Promise"], + ["a string", () => "PIN_MISMATCH", "type string ('PIN_MISMATCH')"], + ] as const)("a callback that returns %s rejects the connection", async (_, callback, received) => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async port => { + const err: any = await ping(`rediss://localhost:${port}`, { + ca: localhostTls.cert, + checkServerIdentity: callback as any, + }).then( + () => null, + e => e, + ); + expect({ name: err?.name, code: err?.code, message: err?.message }).toEqual({ + name: "TypeError", + code: "ERR_INVALID_RETURN_VALUE", + message: `Expected undefined or an Error to be returned from the "tls.checkServerIdentity" function but got ${received}.`, + }); + }); + }); + + test("an object that the callback returns is the reason the connection fails", async () => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async port => { + // Not an Error instance: the same rule as fetch() and node:tls. + const reason = { code: "PIN_MISMATCH" }; + const err = await ping(`rediss://localhost:${port}`, { + ca: localhostTls.cert, + checkServerIdentity: (() => reason) as any, + }).then( + () => null, + e => e, + ); + expect(err).toBe(reason); + }); + }); + + test("receives the chain of getPeerCertificate(true)", async () => { + await withServer({ key: serverKey, cert: serverCert }, async port => { + const chain: string[] = []; + const result = await ping(`rediss://localhost:${port}`, { + ca, + checkServerIdentity: (_hostname: string, cert: tls.PeerCertificate) => { + let current = cert as tls.DetailedPeerCertificate; + let last: string; + do { + chain.push(current.subject.CN); + last = current.fingerprint256; + current = current.issuerCertificate; + } while (current.fingerprint256 !== last); + return undefined; + }, + }); + expect({ result, chain }).toEqual({ result: "PONG", chain: ["agent1", "ca1"] }); + }); + }); + + test("duplicate() keeps the callback", async () => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async port => { + let calls = 0; + const client = new RedisClient(`rediss://localhost:${port}`, { + autoReconnect: false, + connectionTimeout: 5000, + tls: { + ca: localhostTls.cert, + // Accepts the first connection only. + checkServerIdentity: () => (calls++ === 0 ? undefined : new Error("PIN_MISMATCH")), + }, + }); + let duplicate: unknown; + try { + expect(await client.send("PING", [])).toBe("PONG"); + // A connected client's duplicate dials at once: the callback refuses that connection. + duplicate = await client.duplicate().then( + d => d, + e => e, + ); + expect({ calls, duplicateConnected: duplicate instanceof RedisClient && duplicate.connected }).toEqual({ + calls: 2, + duplicateConnected: false, + }); + } finally { + if (duplicate instanceof RedisClient) duplicate.close(); + client.close(); + } + }); + }); + + test("a callback that is a method of a class is called", async () => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async port => { + const pin = new Error("PIN_MISMATCH"); + class PinnedTls { + ca = localhostTls.cert; + checkServerIdentity() { + return pin; + } + } + const err = await ping(`rediss://localhost:${port}`, new PinnedTls()).then( + () => null, + e => e, + ); + expect(err).toBe(pin); + }); + }); + + test("a worker terminated inside the callback stops, and sends nothing to the server", async () => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async (port, server) => { + const { promise: serverSocketClosed, resolve } = Promise.withResolvers(); + let bytesFromClient = 0; + server.on("secureConnection", socket => { + socket.on("data", chunk => (bytesFromClient += chunk.length)); + socket.on("close", () => resolve()); + }); + const counters = new SharedArrayBuffer(12); + const count = new Int32Array(counters); + const worker = new Worker(new URL("./valkey-tls-verify-worker-fixture.ts", import.meta.url), { + workerData: { port, ca: localhostTls.cert, counters }, + }); + const exited = once(worker, "exit"); + await Promise.race([ + Atomics.waitAsync(count, 0, 0).value, + once(worker, "error").then(([error]) => Promise.reject(error)), + exited.then(([code]) => Promise.reject(new Error(`the worker exited with code ${code} before the callback`))), + ]); + // The fixture also wakes this wait when the command settles: then the callback was never called. + expect({ callbackEntered: count[0], commandSettled: count[2] }).toEqual({ + callbackEntered: 1, + commandSettled: 0, + }); + await worker.terminate(); + await Promise.all([exited, serverSocketClosed]); + expect({ callbackEntered: count[0], oncloseRan: count[1], commandSettled: count[2], bytesFromClient }).toEqual({ + callbackEntered: 1, + oncloseRan: 0, + commandSettled: 0, + bytesFromClient: 0, + }); + }); + }); + + test("a callback that closes the client rejects the pending command", async () => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async port => { + const client = new RedisClient(`rediss://localhost:${port}`, { + autoReconnect: false, + connectionTimeout: 5000, + tls: { + ca: localhostTls.cert, + checkServerIdentity: () => { + client.close(); + return undefined; + }, + }, + }); + const err: any = await client.send("PING", []).then( + () => null, + e => e, + ); + expect(err?.code).toBe("ERR_REDIS_CONNECTION_CLOSED"); + expect(client.connected).toBe(false); + }); + }); + + test("a callback that closes the client and dials again leaves the new connection to its own handshake", async () => { + await withServer({ key: localhostTls.key, cert: localhostTls.cert }, async (port, server) => { + const servernames = recordServernames(server); + let calls = 0; + let redial: Promise | undefined; + const client = new RedisClient(`rediss://localhost:${port}`, { + autoReconnect: false, + connectionTimeout: 5000, + tls: { + ca: localhostTls.cert, + checkServerIdentity: () => { + if (calls++ === 0) { + client.close(); + redial = client.connect(); + // The first connection is refused; only the second one may be used. + return new Error("FIRST_CONNECTION_REFUSED"); + } + return undefined; + }, + }, + }); + try { + const first: any = await client.send("PING", []).then( + () => null, + e => e, + ); + expect(first?.code).toBe("ERR_REDIS_CONNECTION_CLOSED"); + await redial; + expect(await client.send("PING", [])).toBe("PONG"); + expect({ calls, handshakes: servernames.length }).toEqual({ calls: 2, handshakes: 2 }); + } finally { + client.close(); + } + }); + }); +});