mirror of
https://github.com/rustfs/rustfs.git
synced 2026-08-06 13:27:43 +00:00
refactor(tls): centralize runtime foundation (#3065)
* refactor(targets): move notify net helpers from utils * refactor(tls): centralize runtime foundation * refactor(targets): move notify net helpers from utils * refactor(tls): centralize runtime foundation * feat(tls-runtime): add TLS debug state and admin handler * refactor(tls-runtime): unify TLS debug consumer status view * fix(tls): address PR3065 review feedback * refactor(tls): align debug status payload types * refactor(targets): harden TLS hot reload paths * fix(targets): resolve review-4348251652 findings * fix(targets): finalize tls runtime review follow-ups * fix(targets): harden tls reload and review follow-ups * fix(targets): align tls reload handling across targets * fix(targets): finalize tls reload state and metrics updates * chore(deps): trim unused TLS deps * style(targets): normalize TLS reload formatting * refactor(targets): introduce tls runtime adapter path * chore: update workspace manifests for tls refactor * fix(tls): stabilize material reload and audit workflow * fix(targets): refresh tls fingerprint flow across sinks * fix(tls): align runtime coordinator and http reader updates * fix(sftp): simplify protocol error mapping * fix(tls): harmonize material loading behavior * fix(server): finalize tls material wiring in startup flow * fix(protos): tighten tls generation cache and deps
This commit is contained in:
@@ -1,767 +0,0 @@
|
||||
// Copyright 2024 RustFS Team
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
use rustls::RootCertStore;
|
||||
use rustls::server::{
|
||||
ClientHello, ResolvesServerCert, ResolvesServerCertUsingSni, WebPkiClientVerifier, danger::ClientCertVerifier,
|
||||
};
|
||||
use rustls::sign::CertifiedKey;
|
||||
use rustls_pki_types::{CertificateDer, PrivateKeyDer, pem::PemObject};
|
||||
use std::collections::HashMap;
|
||||
use std::io::Error;
|
||||
use std::path::PathBuf;
|
||||
use std::sync::Arc;
|
||||
use std::{fs, io};
|
||||
use tracing::{debug, warn};
|
||||
|
||||
/// Options for loading certificate/key pairs from a directory tree.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct CertDirectoryLoadOptions {
|
||||
dir_path: PathBuf,
|
||||
cert_filename: String,
|
||||
key_filename: String,
|
||||
}
|
||||
|
||||
impl CertDirectoryLoadOptions {
|
||||
/// Create a builder with explicit certificate and private key filenames.
|
||||
pub fn builder(
|
||||
dir_path: impl Into<PathBuf>,
|
||||
cert_filename: impl Into<String>,
|
||||
key_filename: impl Into<String>,
|
||||
) -> CertDirectoryLoadOptionsBuilder {
|
||||
CertDirectoryLoadOptionsBuilder {
|
||||
dir_path: dir_path.into(),
|
||||
cert_filename: cert_filename.into(),
|
||||
key_filename: key_filename.into(),
|
||||
}
|
||||
}
|
||||
|
||||
fn validate(&self) -> io::Result<()> {
|
||||
if self.cert_filename.is_empty() {
|
||||
return Err(certs_error("certificate filename cannot be empty".to_string()));
|
||||
}
|
||||
if self.key_filename.is_empty() {
|
||||
return Err(certs_error("private key filename cannot be empty".to_string()));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// Builder for [`CertDirectoryLoadOptions`].
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct CertDirectoryLoadOptionsBuilder {
|
||||
dir_path: PathBuf,
|
||||
cert_filename: String,
|
||||
key_filename: String,
|
||||
}
|
||||
|
||||
impl CertDirectoryLoadOptionsBuilder {
|
||||
/// Override the certificate filename searched in the directory.
|
||||
pub fn cert_filename(mut self, cert_filename: impl Into<String>) -> Self {
|
||||
self.cert_filename = cert_filename.into();
|
||||
self
|
||||
}
|
||||
|
||||
/// Override the private key filename searched in the directory.
|
||||
pub fn key_filename(mut self, key_filename: impl Into<String>) -> Self {
|
||||
self.key_filename = key_filename.into();
|
||||
self
|
||||
}
|
||||
|
||||
/// Build the load options value.
|
||||
pub fn build(self) -> CertDirectoryLoadOptions {
|
||||
CertDirectoryLoadOptions {
|
||||
dir_path: self.dir_path,
|
||||
cert_filename: self.cert_filename,
|
||||
key_filename: self.key_filename,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Options for building an mTLS WebPki client verifier.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct WebPkiClientVerifierOptions {
|
||||
tls_path: PathBuf,
|
||||
enabled: bool,
|
||||
client_ca_cert_filename: String,
|
||||
fallback_ca_cert_filename: String,
|
||||
}
|
||||
|
||||
impl WebPkiClientVerifierOptions {
|
||||
/// Create a builder with explicit CA bundle filenames.
|
||||
pub fn builder(
|
||||
tls_path: impl Into<PathBuf>,
|
||||
client_ca_cert_filename: impl Into<String>,
|
||||
fallback_ca_cert_filename: impl Into<String>,
|
||||
) -> WebPkiClientVerifierOptionsBuilder {
|
||||
WebPkiClientVerifierOptionsBuilder {
|
||||
tls_path: tls_path.into(),
|
||||
enabled: false,
|
||||
client_ca_cert_filename: client_ca_cert_filename.into(),
|
||||
fallback_ca_cert_filename: fallback_ca_cert_filename.into(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Builder for [`WebPkiClientVerifierOptions`].
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct WebPkiClientVerifierOptionsBuilder {
|
||||
tls_path: PathBuf,
|
||||
enabled: bool,
|
||||
client_ca_cert_filename: String,
|
||||
fallback_ca_cert_filename: String,
|
||||
}
|
||||
|
||||
impl WebPkiClientVerifierOptionsBuilder {
|
||||
/// Set whether mTLS verification should be enabled.
|
||||
pub fn enabled(mut self, enabled: bool) -> Self {
|
||||
self.enabled = enabled;
|
||||
self
|
||||
}
|
||||
|
||||
/// Override the preferred client CA bundle filename.
|
||||
pub fn client_ca_cert_filename(mut self, client_ca_cert_filename: impl Into<String>) -> Self {
|
||||
self.client_ca_cert_filename = client_ca_cert_filename.into();
|
||||
self
|
||||
}
|
||||
|
||||
/// Override the fallback CA bundle filename.
|
||||
pub fn fallback_ca_cert_filename(mut self, fallback_ca_cert_filename: impl Into<String>) -> Self {
|
||||
self.fallback_ca_cert_filename = fallback_ca_cert_filename.into();
|
||||
self
|
||||
}
|
||||
|
||||
/// Build the verifier options value.
|
||||
pub fn build(self) -> WebPkiClientVerifierOptions {
|
||||
WebPkiClientVerifierOptions {
|
||||
tls_path: self.tls_path,
|
||||
enabled: self.enabled,
|
||||
client_ca_cert_filename: self.client_ca_cert_filename,
|
||||
fallback_ca_cert_filename: self.fallback_ca_cert_filename,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Load public certificate from file.
|
||||
/// This function loads a public certificate from the specified file.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `filename` - A string slice that holds the name of the file containing the public certificate.
|
||||
///
|
||||
/// # Returns
|
||||
/// * An io::Result containing a vector of CertificateDer if successful, or an io::Error if an error occurs during loading.
|
||||
///
|
||||
pub fn load_certs(filename: &str) -> io::Result<Vec<CertificateDer<'static>>> {
|
||||
// Open certificate file.
|
||||
let cert_file = fs::File::open(filename).map_err(|e| certs_error(format!("failed to open {filename}: {e}")))?;
|
||||
let mut reader = io::BufReader::new(cert_file);
|
||||
|
||||
// Load and return certificate.
|
||||
let certs = CertificateDer::pem_reader_iter(&mut reader)
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.map_err(|e| certs_error(format!("certificate file {filename} format error:{e:?}")))?;
|
||||
if certs.is_empty() {
|
||||
return Err(certs_error(format!("No valid certificate was found in the certificate file {filename}")));
|
||||
}
|
||||
Ok(certs)
|
||||
}
|
||||
|
||||
/// Load a PEM certificate bundle and return each certificate as DER bytes.
|
||||
///
|
||||
/// This is a low-level helper intended for TLS clients (reqwest/hyper-rustls) that
|
||||
/// need to add root certificates one-by-one.
|
||||
///
|
||||
/// - Input: a PEM file that may contain multiple cert blocks.
|
||||
/// - Output: Vec of DER-encoded cert bytes, one per cert.
|
||||
///
|
||||
/// NOTE: This intentionally returns raw bytes to avoid forcing downstream crates
|
||||
/// to depend on rustls types.
|
||||
pub fn load_cert_bundle_der_bytes(path: &str) -> io::Result<Vec<Vec<u8>>> {
|
||||
let pem = fs::read(path)?;
|
||||
let mut reader = io::BufReader::new(&pem[..]);
|
||||
|
||||
let certs = CertificateDer::pem_reader_iter(&mut reader)
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.map_err(|e| certs_error(format!("Failed to parse PEM certs from {path}: {e}")))?;
|
||||
|
||||
Ok(certs.into_iter().map(|c| c.to_vec()).collect())
|
||||
}
|
||||
|
||||
/// Builds a WebPkiClientVerifier for mTLS when enabled by the caller.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `options` - mTLS verifier options, including the TLS directory and CA bundle filenames
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(Some(verifier))` if mTLS is enabled and CA certs are found
|
||||
/// * `Ok(None)` if mTLS is disabled
|
||||
/// * `Err` if mTLS is enabled but configuration is invalid
|
||||
pub fn build_webpki_client_verifier(options: WebPkiClientVerifierOptions) -> io::Result<Option<Arc<dyn ClientCertVerifier>>> {
|
||||
if !options.enabled {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let tls_path = &options.tls_path;
|
||||
let ca_path = mtls_ca_bundle_path(&options).ok_or_else(|| {
|
||||
Error::other(format!(
|
||||
"mTLS is enabled but missing {}/{} (or fallback {}/{})",
|
||||
tls_path.display(),
|
||||
options.client_ca_cert_filename,
|
||||
tls_path.display(),
|
||||
options.fallback_ca_cert_filename
|
||||
))
|
||||
})?;
|
||||
|
||||
let ca_path = ca_path
|
||||
.to_str()
|
||||
.ok_or_else(|| Error::other(format!("Invalid UTF-8 in mTLS CA path: {ca_path:?}")))?;
|
||||
|
||||
let der_list = load_cert_bundle_der_bytes(ca_path)?;
|
||||
|
||||
let mut store = RootCertStore::empty();
|
||||
for der in der_list {
|
||||
store
|
||||
.add(der.into())
|
||||
.map_err(|e| Error::other(format!("Invalid client CA cert: {e}")))?;
|
||||
}
|
||||
|
||||
let verifier = WebPkiClientVerifier::builder(Arc::new(store))
|
||||
.build()
|
||||
.map_err(|e| Error::other(format!("Build client cert verifier failed: {e}")))?;
|
||||
|
||||
Ok(Some(verifier))
|
||||
}
|
||||
|
||||
/// Locate the mTLS client CA bundle in the specified TLS path
|
||||
fn mtls_ca_bundle_path(options: &WebPkiClientVerifierOptions) -> Option<PathBuf> {
|
||||
let p1 = options.tls_path.join(&options.client_ca_cert_filename);
|
||||
if p1.exists() {
|
||||
return Some(p1);
|
||||
}
|
||||
let p2 = options.tls_path.join(&options.fallback_ca_cert_filename);
|
||||
if p2.exists() {
|
||||
return Some(p2);
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
/// Load private key from file.
|
||||
/// This function loads a private key from the specified file.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `filename` - A string slice that holds the name of the file containing the private key.
|
||||
///
|
||||
/// # Returns
|
||||
/// * An io::Result containing the PrivateKeyDer if successful, or an io::Error if an error occurs during loading.
|
||||
///
|
||||
pub fn load_private_key(filename: &str) -> io::Result<PrivateKeyDer<'static>> {
|
||||
// Open keyfile.
|
||||
let keyfile = fs::File::open(filename).map_err(|e| certs_error(format!("failed to open {filename}: {e}")))?;
|
||||
let mut reader = io::BufReader::new(keyfile);
|
||||
|
||||
// Load and return a single private key.
|
||||
PrivateKeyDer::from_pem_reader(&mut reader)
|
||||
.map_err(|e| certs_error(format!("failed to parse private key in {filename}: {e}")))
|
||||
}
|
||||
|
||||
/// error function
|
||||
/// This function creates a new io::Error with the provided error message.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `err` - A string containing the error message.
|
||||
///
|
||||
/// # Returns
|
||||
/// * An io::Error instance with the specified error message.
|
||||
///
|
||||
pub fn certs_error(err: String) -> Error {
|
||||
Error::other(err)
|
||||
}
|
||||
|
||||
fn is_discoverable_cert_domain_dir(domain_name: &str) -> bool {
|
||||
!domain_name.starts_with('.')
|
||||
}
|
||||
|
||||
/// Load all certificates and private keys in the directory
|
||||
/// This function loads all certificate and private key pairs from the specified directory.
|
||||
/// It looks for files named `options.cert_filename` and `options.key_filename` in each subdirectory.
|
||||
/// The root directory can also contain a default certificate/private key pair.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `options` - Directory and filename options for discovering certificates and private keys.
|
||||
///
|
||||
/// # Returns
|
||||
/// * An io::Result containing a HashMap where the keys are domain names (or "default" for the root certificate) and the values are tuples of (Vec<CertificateDer>, PrivateKeyDer). If no valid certificate/private key pairs are found, an io::Error is returned.
|
||||
///
|
||||
pub fn load_all_certs_from_directory(
|
||||
options: CertDirectoryLoadOptions,
|
||||
) -> io::Result<HashMap<String, (Vec<CertificateDer<'static>>, PrivateKeyDer<'static>)>> {
|
||||
options.validate()?;
|
||||
|
||||
let mut cert_key_pairs = HashMap::new();
|
||||
let dir = options.dir_path.as_path();
|
||||
|
||||
if !dir.exists() || !dir.is_dir() {
|
||||
return Err(certs_error(format!(
|
||||
"The certificate directory does not exist or is not a directory: {}",
|
||||
dir.display()
|
||||
)));
|
||||
}
|
||||
|
||||
// 1. First check whether there is a certificate/private key pair in the root directory
|
||||
let root_cert_path = dir.join(&options.cert_filename);
|
||||
let root_key_path = dir.join(&options.key_filename);
|
||||
|
||||
if root_cert_path.exists() && root_key_path.exists() {
|
||||
debug!("find the root directory certificate: {:?}", root_cert_path);
|
||||
let root_cert_str = root_cert_path
|
||||
.to_str()
|
||||
.ok_or_else(|| certs_error(format!("Invalid UTF-8 in root certificate path: {root_cert_path:?}")))?;
|
||||
let root_key_str = root_key_path
|
||||
.to_str()
|
||||
.ok_or_else(|| certs_error(format!("Invalid UTF-8 in root key path: {root_key_path:?}")))?;
|
||||
match load_cert_key_pair(root_cert_str, root_key_str) {
|
||||
Ok((certs, key)) => {
|
||||
// The root directory certificate is used as the default certificate and is stored using special keys.
|
||||
cert_key_pairs.insert("default".to_string(), (certs, key));
|
||||
}
|
||||
Err(e) => {
|
||||
warn!("unable to load root directory certificate: {}", e);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 2.iterate through all folders in the directory
|
||||
for entry in fs::read_dir(dir)? {
|
||||
let entry = entry?;
|
||||
let path = entry.path();
|
||||
|
||||
if path.is_dir() {
|
||||
let domain_name: &str = path
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.ok_or_else(|| certs_error(format!("invalid domain name directory:{path:?}")))?;
|
||||
if !is_discoverable_cert_domain_dir(domain_name) {
|
||||
debug!("skip internal certificate directory: {:?}", path);
|
||||
continue;
|
||||
}
|
||||
|
||||
// find certificate and private key files
|
||||
let cert_path = path.join(&options.cert_filename); // e.g., rustfs_cert.pem
|
||||
let key_path = path.join(&options.key_filename); // e.g., rustfs_key.pem
|
||||
|
||||
if cert_path.exists() && key_path.exists() {
|
||||
debug!("find the domain name certificate: {} in {:?}", domain_name, cert_path);
|
||||
let cert_path = match cert_path.to_str() {
|
||||
Some(path) => path,
|
||||
None => {
|
||||
warn!("skip domain certificate load, invalid UTF-8 path: {:?}", cert_path);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
let key_path = match key_path.to_str() {
|
||||
Some(path) => path,
|
||||
None => {
|
||||
warn!("skip domain key load, invalid UTF-8 path: {:?}", key_path);
|
||||
continue;
|
||||
}
|
||||
};
|
||||
|
||||
match load_cert_key_pair(cert_path, key_path) {
|
||||
Ok((certs, key)) => {
|
||||
cert_key_pairs.insert(domain_name.to_string(), (certs, key));
|
||||
}
|
||||
Err(e) => {
|
||||
warn!("unable to load the certificate for {} domain name: {}", domain_name, e);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if cert_key_pairs.is_empty() {
|
||||
return Err(certs_error(format!(
|
||||
"No valid certificate/private key pair found in directory {}",
|
||||
dir.display()
|
||||
)));
|
||||
}
|
||||
|
||||
Ok(cert_key_pairs)
|
||||
}
|
||||
|
||||
/// loading a single certificate private key pair
|
||||
/// This function loads a certificate and private key from the specified paths.
|
||||
/// It returns a tuple containing the certificate and private key.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `cert_path` - A string slice that holds the path to the certificate file.
|
||||
/// * `key_path` - A string slice that holds the path to the private key file
|
||||
///
|
||||
/// # Returns
|
||||
/// * An io::Result containing a tuple of (Vec<CertificateDer>, PrivateKeyDer) if successful, or an io::Error if an error occurs during loading.
|
||||
///
|
||||
fn load_cert_key_pair(cert_path: &str, key_path: &str) -> io::Result<(Vec<CertificateDer<'static>>, PrivateKeyDer<'static>)> {
|
||||
let certs = load_certs(cert_path)?;
|
||||
let key = load_private_key(key_path)?;
|
||||
Ok((certs, key))
|
||||
}
|
||||
|
||||
/// Create a multi-cert resolver
|
||||
/// This function loads all certificates and private keys from the specified directory.
|
||||
/// It uses the first certificate/private key pair found in the root directory as the default certificate.
|
||||
/// The rest of the certificates/private keys are used for SNI resolution.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `cert_key_pairs` - A HashMap where the keys are domain names (or "default" for the root certificate) and the values are tuples of (Vec<CertificateDer>, PrivateKeyDer).
|
||||
///
|
||||
/// # Returns
|
||||
/// * An io::Result containing an implementation of ResolvesServerCert if successful, or an io::Error if an error occurs during loading.
|
||||
///
|
||||
pub fn create_multi_cert_resolver(
|
||||
cert_key_pairs: HashMap<String, (Vec<CertificateDer<'static>>, PrivateKeyDer<'static>)>,
|
||||
) -> io::Result<impl ResolvesServerCert> {
|
||||
#[derive(Debug)]
|
||||
struct MultiCertResolver {
|
||||
cert_resolver: ResolvesServerCertUsingSni,
|
||||
default_cert: Option<Arc<CertifiedKey>>,
|
||||
}
|
||||
impl ResolvesServerCert for MultiCertResolver {
|
||||
fn resolve(&self, client_hello: ClientHello) -> Option<Arc<CertifiedKey>> {
|
||||
// try matching certificates with sni
|
||||
if let Some(cert) = self.cert_resolver.resolve(client_hello) {
|
||||
return Some(cert);
|
||||
}
|
||||
|
||||
// If there is no matching SNI certificate, use the default certificate
|
||||
self.default_cert.clone()
|
||||
}
|
||||
}
|
||||
|
||||
let mut resolver = ResolvesServerCertUsingSni::new();
|
||||
let mut default_cert = None;
|
||||
|
||||
for (domain, (certs, key)) in cert_key_pairs {
|
||||
// create a signature
|
||||
let signing_key = rustls::crypto::aws_lc_rs::sign::any_supported_type(&key)
|
||||
.map_err(|e| certs_error(format!("unsupported private key types:{domain}, err:{e:?}")))?;
|
||||
|
||||
// create a CertifiedKey
|
||||
let certified_key = CertifiedKey::new(certs, signing_key);
|
||||
if domain == "default" {
|
||||
default_cert = Some(Arc::new(certified_key.clone()));
|
||||
} else {
|
||||
// add certificate to resolver
|
||||
resolver
|
||||
.add(&domain, certified_key)
|
||||
.map_err(|e| certs_error(format!("failed to add a domain name certificate:{domain},err: {e:?}")))?;
|
||||
}
|
||||
}
|
||||
|
||||
Ok(MultiCertResolver {
|
||||
cert_resolver: resolver,
|
||||
default_cert,
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::fs;
|
||||
use std::io::ErrorKind;
|
||||
use tempfile::TempDir;
|
||||
|
||||
fn default_load_options(path: impl Into<PathBuf>) -> CertDirectoryLoadOptions {
|
||||
CertDirectoryLoadOptions::builder(path, "rustfs_cert.pem", "rustfs_key.pem").build()
|
||||
}
|
||||
|
||||
fn write_test_cert_pair(dir: &std::path::Path) {
|
||||
let rcgen::CertifiedKey { cert, signing_key } =
|
||||
rcgen::generate_simple_self_signed(vec!["example.com".to_string()]).unwrap();
|
||||
fs::write(dir.join("rustfs_cert.pem"), cert.pem()).unwrap();
|
||||
fs::write(dir.join("rustfs_key.pem"), signing_key.serialize_pem()).unwrap();
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_certs_error_function() {
|
||||
let error_msg = "Test error message";
|
||||
let error = certs_error(error_msg.to_string());
|
||||
|
||||
assert_eq!(error.kind(), ErrorKind::Other);
|
||||
assert_eq!(error.to_string(), error_msg);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_certs_file_not_found() {
|
||||
let result = load_certs("non_existent_file.pem");
|
||||
assert!(result.is_err());
|
||||
|
||||
let error = result.unwrap_err();
|
||||
assert_eq!(error.kind(), ErrorKind::Other);
|
||||
assert!(error.to_string().contains("failed to open"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_private_key_file_not_found() {
|
||||
let result = load_private_key("non_existent_key.pem");
|
||||
assert!(result.is_err());
|
||||
|
||||
let error = result.unwrap_err();
|
||||
assert_eq!(error.kind(), ErrorKind::Other);
|
||||
assert!(error.to_string().contains("failed to open"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_certs_empty_file() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let cert_path = temp_dir.path().join("empty.pem");
|
||||
fs::write(&cert_path, "").unwrap();
|
||||
|
||||
let result = load_certs(cert_path.to_str().unwrap());
|
||||
assert!(result.is_err());
|
||||
|
||||
let error = result.unwrap_err();
|
||||
assert!(error.to_string().contains("No valid certificate was found"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_certs_invalid_format() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let cert_path = temp_dir.path().join("invalid.pem");
|
||||
fs::write(&cert_path, "invalid certificate content").unwrap();
|
||||
|
||||
let result = load_certs(cert_path.to_str().unwrap());
|
||||
assert!(result.is_err());
|
||||
|
||||
let error = result.unwrap_err();
|
||||
assert!(error.to_string().contains("No valid certificate was found"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_private_key_empty_file() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let key_path = temp_dir.path().join("empty_key.pem");
|
||||
fs::write(&key_path, "").unwrap();
|
||||
|
||||
let result = load_private_key(key_path.to_str().unwrap());
|
||||
assert!(result.is_err());
|
||||
|
||||
let error = result.unwrap_err();
|
||||
assert!(error.to_string().contains("failed to parse private key in"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_private_key_invalid_format() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let key_path = temp_dir.path().join("invalid_key.pem");
|
||||
fs::write(&key_path, "invalid private key content").unwrap();
|
||||
|
||||
let result = load_private_key(key_path.to_str().unwrap());
|
||||
assert!(result.is_err());
|
||||
|
||||
let error = result.unwrap_err();
|
||||
assert!(error.to_string().contains("failed to parse private key in"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_all_certs_from_directory_not_exists() {
|
||||
let result = load_all_certs_from_directory(default_load_options("/non/existent/directory"));
|
||||
assert!(result.is_err());
|
||||
|
||||
let error = result.unwrap_err();
|
||||
assert!(error.to_string().contains("does not exist or is not a directory"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_all_certs_from_directory_empty() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
|
||||
let result = load_all_certs_from_directory(default_load_options(temp_dir.path()));
|
||||
assert!(result.is_err());
|
||||
|
||||
let error = result.unwrap_err();
|
||||
assert!(error.to_string().contains("No valid certificate/private key pair found"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_all_certs_from_directory_file_instead_of_dir() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let file_path = temp_dir.path().join("not_a_directory.txt");
|
||||
fs::write(&file_path, "content").unwrap();
|
||||
|
||||
let result = load_all_certs_from_directory(default_load_options(&file_path));
|
||||
assert!(result.is_err());
|
||||
|
||||
let error = result.unwrap_err();
|
||||
assert!(error.to_string().contains("does not exist or is not a directory"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_cert_key_pair_missing_cert() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let key_path = temp_dir.path().join("test_key.pem");
|
||||
fs::write(&key_path, "dummy key content").unwrap();
|
||||
|
||||
let result = load_cert_key_pair("non_existent_cert.pem", key_path.to_str().unwrap());
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_cert_key_pair_missing_key() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let cert_path = temp_dir.path().join("test_cert.pem");
|
||||
fs::write(&cert_path, "dummy cert content").unwrap();
|
||||
|
||||
let result = load_cert_key_pair(cert_path.to_str().unwrap(), "non_existent_key.pem");
|
||||
assert!(result.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_create_multi_cert_resolver_empty_map() {
|
||||
let empty_map = HashMap::new();
|
||||
let result = create_multi_cert_resolver(empty_map);
|
||||
|
||||
// Should succeed even with empty map
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_error_message_formatting() {
|
||||
let test_cases = vec![
|
||||
("file not found", "failed to open test.pem: file not found"),
|
||||
("permission denied", "failed to open key.pem: permission denied"),
|
||||
("invalid format", "certificate file cert.pem format error:invalid format"),
|
||||
];
|
||||
|
||||
for (input, _expected_pattern) in test_cases {
|
||||
let error1 = certs_error(format!("failed to open test.pem: {input}"));
|
||||
assert!(error1.to_string().contains(input));
|
||||
|
||||
let error2 = certs_error(format!("failed to open key.pem: {input}"));
|
||||
assert!(error2.to_string().contains(input));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_path_handling_edge_cases() {
|
||||
// Test with various path formats
|
||||
let path_cases = vec![
|
||||
"", // Empty path
|
||||
".", // Current directory
|
||||
"..", // Parent directory
|
||||
"/", // Root directory (Unix)
|
||||
"relative/path", // Relative path
|
||||
"/absolute/path", // Absolute path
|
||||
];
|
||||
|
||||
for path in path_cases {
|
||||
let result = load_all_certs_from_directory(default_load_options(path));
|
||||
// All should fail since these are not valid cert directories
|
||||
assert!(result.is_err());
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_directory_structure_validation() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
|
||||
// Create a subdirectory without certificates
|
||||
let sub_dir = temp_dir.path().join("example.com");
|
||||
fs::create_dir(&sub_dir).unwrap();
|
||||
|
||||
// Should fail because no certificates found
|
||||
let result = load_all_certs_from_directory(default_load_options(temp_dir.path()));
|
||||
assert!(result.is_err());
|
||||
assert!(
|
||||
result
|
||||
.unwrap_err()
|
||||
.to_string()
|
||||
.contains("No valid certificate/private key pair found")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_load_all_certs_skips_kubernetes_secret_projection_dirs() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
write_test_cert_pair(temp_dir.path());
|
||||
|
||||
let domain_dir = temp_dir.path().join("example.com");
|
||||
fs::create_dir(&domain_dir).unwrap();
|
||||
write_test_cert_pair(&domain_dir);
|
||||
|
||||
for internal_dir_name in ["..data", "..2026_04_28_18_33_53.4209048473"] {
|
||||
let internal_dir = temp_dir.path().join(internal_dir_name);
|
||||
fs::create_dir(&internal_dir).unwrap();
|
||||
write_test_cert_pair(&internal_dir);
|
||||
}
|
||||
|
||||
let certs = load_all_certs_from_directory(default_load_options(temp_dir.path())).unwrap();
|
||||
|
||||
assert!(certs.contains_key("default"));
|
||||
assert!(certs.contains_key("example.com"));
|
||||
assert!(!certs.contains_key("..data"));
|
||||
assert!(!certs.contains_key("..2026_04_28_18_33_53.4209048473"));
|
||||
assert_eq!(certs.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_unicode_path_handling() {
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
|
||||
// Create directory with Unicode characters
|
||||
let unicode_dir = temp_dir.path().join("test_directory");
|
||||
fs::create_dir(&unicode_dir).unwrap();
|
||||
|
||||
let result = load_all_certs_from_directory(default_load_options(&unicode_dir));
|
||||
assert!(result.is_err());
|
||||
assert!(
|
||||
result
|
||||
.unwrap_err()
|
||||
.to_string()
|
||||
.contains("No valid certificate/private key pair found")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_concurrent_access_safety() {
|
||||
use std::sync::Arc;
|
||||
use std::thread;
|
||||
|
||||
let temp_dir = TempDir::new().unwrap();
|
||||
let dir_path = Arc::new(temp_dir.path().to_string_lossy().to_string());
|
||||
|
||||
let handles: Vec<_> = (0..5)
|
||||
.map(|_| {
|
||||
let path = Arc::clone(&dir_path);
|
||||
thread::spawn(move || {
|
||||
let result = load_all_certs_from_directory(default_load_options(path.as_str()));
|
||||
// All should fail since directory is empty
|
||||
assert!(result.is_err());
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
|
||||
for handle in handles {
|
||||
handle.join().expect("Thread should complete successfully");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_memory_efficiency() {
|
||||
let error = certs_error("test".to_string());
|
||||
let error_size = std::mem::size_of_val(&error);
|
||||
|
||||
// Error should not be excessively large
|
||||
assert!(error_size < 1024, "Error size should be reasonable, got {error_size} bytes");
|
||||
}
|
||||
}
|
||||
@@ -12,8 +12,6 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
#[cfg(feature = "tls")]
|
||||
pub mod certs;
|
||||
#[cfg(feature = "ip")]
|
||||
pub mod ip;
|
||||
#[cfg(feature = "net")]
|
||||
@@ -52,9 +50,6 @@ pub mod compress;
|
||||
#[cfg(feature = "path")]
|
||||
pub mod dirs;
|
||||
|
||||
#[cfg(feature = "tls")]
|
||||
pub use certs::*;
|
||||
|
||||
#[cfg(feature = "hash")]
|
||||
pub use hash::*;
|
||||
|
||||
@@ -70,12 +65,6 @@ pub use crypto::*;
|
||||
#[cfg(feature = "compress")]
|
||||
pub use compress::*;
|
||||
|
||||
#[cfg(feature = "notify")]
|
||||
mod notify;
|
||||
|
||||
#[cfg(feature = "notify")]
|
||||
pub use notify::*;
|
||||
|
||||
#[cfg(feature = "obj")]
|
||||
pub mod obj;
|
||||
|
||||
|
||||
@@ -1,205 +0,0 @@
|
||||
// Copyright 2024 RustFS Team
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
mod net;
|
||||
|
||||
use hashbrown::HashMap;
|
||||
use hyper::HeaderMap;
|
||||
use s3s::{S3Request, S3Response};
|
||||
|
||||
pub use net::*;
|
||||
|
||||
/// Extract request parameters from S3Request, mainly header information.
|
||||
pub fn extract_req_params<T>(req: &S3Request<T>) -> HashMap<String, String> {
|
||||
extract_params_header(&req.headers)
|
||||
}
|
||||
|
||||
/// Extract request parameters from hyper::HeaderMap, mainly header information.
|
||||
/// This function is useful when you have a raw HTTP request and need to extract parameters.
|
||||
#[deprecated(since = "0.1.0", note = "Use extract_params_header instead")]
|
||||
pub fn extract_req_params_header(head: &HeaderMap) -> HashMap<String, String> {
|
||||
extract_params_header(head)
|
||||
}
|
||||
|
||||
/// Extract parameters from hyper::HeaderMap, mainly header information.
|
||||
/// This function is useful when you have a raw HTTP request and need to extract parameters.
|
||||
pub fn extract_params_header(head: &HeaderMap) -> HashMap<String, String> {
|
||||
let mut params = HashMap::new();
|
||||
for (key, value) in head.iter() {
|
||||
if let Ok(val_str) = value.to_str() {
|
||||
params.insert(key.as_str().to_string(), val_str.to_string());
|
||||
}
|
||||
}
|
||||
params
|
||||
}
|
||||
|
||||
/// Extract response elements from S3Response, mainly header information.
|
||||
pub fn extract_resp_elements<T>(resp: &S3Response<T>) -> HashMap<String, String> {
|
||||
extract_params_header(&resp.headers)
|
||||
}
|
||||
|
||||
/// Get host from header information.
|
||||
pub fn get_request_host(headers: &HeaderMap) -> String {
|
||||
headers
|
||||
.get("host")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string()
|
||||
}
|
||||
|
||||
/// Get Port from header information.
|
||||
/// Priority:
|
||||
/// 1. x-forwarded-port
|
||||
/// 2. host header (parse port)
|
||||
/// If host has no port, try to deduce from x-forwarded-proto (http->80, https->443)
|
||||
/// 3. port header
|
||||
///
|
||||
/// If the port cannot be determined, returns 0.
|
||||
pub fn get_request_port(headers: &HeaderMap) -> u16 {
|
||||
// 1. Try x-forwarded-port
|
||||
if let Some(port) = headers
|
||||
.get("x-forwarded-port")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|s| s.parse::<u16>().ok())
|
||||
{
|
||||
return port;
|
||||
}
|
||||
|
||||
// 2. Try host header
|
||||
if let Some(host) = headers.get("host").and_then(|v| v.to_str().ok()) {
|
||||
if let Some(idx) = host.rfind(':') {
|
||||
// Check if it's an IPv6 address with port, e.g., [::1]:8080
|
||||
// If ']' is present, the colon must be after it.
|
||||
let valid_colon = match host.rfind(']') {
|
||||
Some(close_bracket_idx) => idx > close_bracket_idx,
|
||||
None => true,
|
||||
};
|
||||
|
||||
if valid_colon
|
||||
&& let Ok(port) = host[idx + 1..].parse::<u16>()
|
||||
&& port > 0
|
||||
{
|
||||
return port;
|
||||
}
|
||||
}
|
||||
|
||||
// If host is present but no port found (or parsing failed, or port is 0),
|
||||
// try to deduce from x-forwarded-proto
|
||||
if let Some(proto) = headers.get("x-forwarded-proto").and_then(|v| v.to_str().ok()) {
|
||||
match proto {
|
||||
"http" => return 80,
|
||||
"https" => return 443,
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// 3. Fallback to "port" header
|
||||
headers
|
||||
.get("port")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|s| s.parse::<u16>().ok())
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
/// Get content-length from header information.
|
||||
pub fn get_request_content_length(headers: &HeaderMap) -> u64 {
|
||||
headers
|
||||
.get("content-length")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.and_then(|s| s.parse::<u64>().ok())
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
/// Get referer from header information.
|
||||
/// If the referer header is not present, returns an empty string.
|
||||
pub fn get_request_referer(headers: &HeaderMap) -> String {
|
||||
headers
|
||||
.get("referer")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string()
|
||||
}
|
||||
|
||||
/// Get user-agent from header information.
|
||||
pub fn get_request_user_agent(headers: &HeaderMap) -> String {
|
||||
headers
|
||||
.get("user-agent")
|
||||
.and_then(|v| v.to_str().ok())
|
||||
.unwrap_or_default()
|
||||
.to_string()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use hyper::header::HeaderValue;
|
||||
|
||||
#[test]
|
||||
fn test_get_request_port() {
|
||||
let mut headers = HeaderMap::new();
|
||||
|
||||
// Case 1: No port info
|
||||
assert_eq!(get_request_port(&headers), 0);
|
||||
|
||||
// Case 2: port header
|
||||
headers.insert("port", HeaderValue::from_static("8080"));
|
||||
assert_eq!(get_request_port(&headers), 8080);
|
||||
|
||||
// Case 3: host header with port
|
||||
headers.remove("port");
|
||||
headers.insert("host", HeaderValue::from_static("example.com:9000"));
|
||||
assert_eq!(get_request_port(&headers), 9000);
|
||||
|
||||
// Case 4: host header without port, no proto
|
||||
headers.insert("host", HeaderValue::from_static("example.com"));
|
||||
assert_eq!(get_request_port(&headers), 0);
|
||||
|
||||
// Case 5: IPv6 host with port
|
||||
headers.insert("host", HeaderValue::from_static("[::1]:9001"));
|
||||
assert_eq!(get_request_port(&headers), 9001);
|
||||
|
||||
// Case 6: IPv6 host without port
|
||||
headers.insert("host", HeaderValue::from_static("[::1]"));
|
||||
assert_eq!(get_request_port(&headers), 0);
|
||||
|
||||
// Case 7: x-forwarded-port
|
||||
headers.insert("x-forwarded-port", HeaderValue::from_static("7000"));
|
||||
// Even if host is present, x-forwarded-port takes precedence
|
||||
assert_eq!(get_request_port(&headers), 7000);
|
||||
|
||||
// Case 8: host without port, but x-forwarded-proto is http
|
||||
headers.remove("x-forwarded-port");
|
||||
headers.insert("host", HeaderValue::from_static("example.com"));
|
||||
headers.insert("x-forwarded-proto", HeaderValue::from_static("http"));
|
||||
assert_eq!(get_request_port(&headers), 80);
|
||||
|
||||
// Case 9: host without port, but x-forwarded-proto is https
|
||||
headers.insert("x-forwarded-proto", HeaderValue::from_static("https"));
|
||||
assert_eq!(get_request_port(&headers), 443);
|
||||
|
||||
// Case 10: host without port, unknown proto
|
||||
headers.insert("x-forwarded-proto", HeaderValue::from_static("ftp"));
|
||||
assert_eq!(get_request_port(&headers), 0);
|
||||
|
||||
// Case 11: host with port 0, should fallback to proto
|
||||
headers.insert("host", HeaderValue::from_static("example.com:0"));
|
||||
headers.insert("x-forwarded-proto", HeaderValue::from_static("https"));
|
||||
assert_eq!(get_request_port(&headers), 443);
|
||||
|
||||
// Case 12: host with port 0, no proto
|
||||
headers.remove("x-forwarded-proto");
|
||||
assert_eq!(get_request_port(&headers), 0);
|
||||
}
|
||||
}
|
||||
@@ -1,664 +0,0 @@
|
||||
// Copyright 2024 RustFS Team
|
||||
//
|
||||
// Licensed under the Apache License, Version 2.0 (the "License");
|
||||
// you may not use this file except in compliance with the License.
|
||||
// You may obtain a copy of the License at
|
||||
//
|
||||
// http://www.apache.org/licenses/LICENSE-2.0
|
||||
//
|
||||
// Unless required by applicable law or agreed to in writing, software
|
||||
// distributed under the License is distributed on an "AS IS" BASIS,
|
||||
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
use regex::Regex;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::net::IpAddr;
|
||||
use std::path::Path;
|
||||
use std::sync::LazyLock;
|
||||
use thiserror::Error;
|
||||
use url::Url;
|
||||
|
||||
// Lazy static for the host label regex.
|
||||
static HOST_LABEL_REGEX: LazyLock<Regex> = LazyLock::new(|| Regex::new(r"^[a-zA-Z0-9]([a-zA-Z0-9-]*[a-zA-Z0-9])?$").unwrap());
|
||||
|
||||
/// NetError represents errors that can occur in network operations.
|
||||
#[derive(Error, Debug)]
|
||||
pub enum NetError {
|
||||
#[error("invalid argument")]
|
||||
InvalidArgument,
|
||||
#[error("invalid hostname")]
|
||||
InvalidHost,
|
||||
#[error("missing '[' in host")]
|
||||
MissingBracket,
|
||||
#[error("parse error: {0}")]
|
||||
ParseError(String),
|
||||
#[error("unexpected scheme: {0}")]
|
||||
UnexpectedScheme(String),
|
||||
#[error("scheme appears with empty host")]
|
||||
SchemeWithEmptyHost,
|
||||
}
|
||||
|
||||
/// Host represents a network host with IP/name and port.
|
||||
/// Similar to Go's net.Host structure.
|
||||
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
|
||||
pub struct Host {
|
||||
pub name: String,
|
||||
pub port: Option<u16>, // Using Option<u16> to represent if port is set, similar to IsPortSet.
|
||||
}
|
||||
|
||||
// Implementation of Host methods.
|
||||
impl Host {
|
||||
// is_empty returns true if the host name is empty.
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.name.is_empty()
|
||||
}
|
||||
|
||||
// equal checks if two hosts are equal by comparing their string representations.
|
||||
pub fn equal(&self, other: &Host) -> bool {
|
||||
self.to_string() == other.to_string()
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for Host {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
match self.port {
|
||||
Some(p) => write!(f, "{}:{}", self.name, p),
|
||||
None => write!(f, "{}", self.name),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// parse_host parses a string into a Host, with validation similar to Go's ParseHost.
|
||||
pub fn parse_host(s: &str) -> Result<Host, NetError> {
|
||||
if s.is_empty() {
|
||||
return Err(NetError::InvalidArgument);
|
||||
}
|
||||
|
||||
// is_valid_host validates the host string, checking for IP or hostname validity.
|
||||
let is_valid_host = |host: &str| -> bool {
|
||||
if host.is_empty() {
|
||||
return true;
|
||||
}
|
||||
if host.parse::<IpAddr>().is_ok() {
|
||||
return true;
|
||||
}
|
||||
if !(1..=253).contains(&host.len()) {
|
||||
return false;
|
||||
}
|
||||
for (i, label) in host.split('.').enumerate() {
|
||||
if i + 1 == host.split('.').count() && label.is_empty() {
|
||||
continue;
|
||||
}
|
||||
if !(1..=63).contains(&label.len()) || !HOST_LABEL_REGEX.is_match(label) {
|
||||
return false;
|
||||
}
|
||||
}
|
||||
true
|
||||
};
|
||||
|
||||
let (host, port) = if let Some(rest) = s.strip_prefix('[') {
|
||||
let Some(end) = rest.find(']') else {
|
||||
return Err(NetError::MissingBracket);
|
||||
};
|
||||
let host = rest[..end].to_string();
|
||||
let port_str = &rest[end + 1..];
|
||||
let port = if let Some(port_str) = port_str.strip_prefix(':') {
|
||||
if port_str.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(port_str.parse().map_err(|_| NetError::ParseError(port_str.to_string()))?)
|
||||
}
|
||||
} else if port_str.is_empty() {
|
||||
None
|
||||
} else {
|
||||
return Err(NetError::InvalidHost);
|
||||
};
|
||||
|
||||
(host, port)
|
||||
} else {
|
||||
if s.contains(']') {
|
||||
return Err(NetError::MissingBracket);
|
||||
}
|
||||
|
||||
// A host with multiple colons is an IPv6 literal, optionally with a
|
||||
// zone identifier. Unbracketed IPv6 with port is ambiguous, so callers
|
||||
// must use the standard bracketed form when they need a port.
|
||||
let (host_str, port_str) = if s.matches(':').count() > 1 {
|
||||
(s, "")
|
||||
} else {
|
||||
s.rsplit_once(':').map_or((s, ""), |(h, p)| (h, p))
|
||||
};
|
||||
let port = if !port_str.is_empty() {
|
||||
Some(port_str.parse().map_err(|_| NetError::ParseError(port_str.to_string()))?)
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
(trim_ipv6(host_str)?, port)
|
||||
};
|
||||
|
||||
// Handle IPv6 zone identifier.
|
||||
let trimmed_host = host.split('%').next().unwrap_or(&host);
|
||||
|
||||
if !is_valid_host(trimmed_host) {
|
||||
return Err(NetError::InvalidHost);
|
||||
}
|
||||
|
||||
Ok(Host { name: host, port })
|
||||
}
|
||||
|
||||
// trim_ipv6 removes square brackets from IPv6 addresses, similar to Go's trimIPv6.
|
||||
fn trim_ipv6(host: &str) -> Result<String, NetError> {
|
||||
if host.ends_with(']') {
|
||||
if !host.starts_with('[') {
|
||||
return Err(NetError::MissingBracket);
|
||||
}
|
||||
Ok(host[1..host.len() - 1].to_string())
|
||||
} else {
|
||||
Ok(host.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
/// URL is a wrapper around url::Url for custom handling.
|
||||
/// Provides methods similar to Go's URL struct.
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ParsedURL(pub Url);
|
||||
|
||||
impl ParsedURL {
|
||||
/// is_empty returns true if the URL is empty or "about:blank".
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `&self` - Reference to the ParsedURL instance.
|
||||
///
|
||||
/// # Returns
|
||||
/// * `bool` - True if the URL is empty or "about:blank", false otherwise.
|
||||
///
|
||||
pub fn is_empty(&self) -> bool {
|
||||
self.0.as_str() == "" || (self.0.scheme() == "about" && self.0.path() == "blank")
|
||||
}
|
||||
|
||||
/// hostname returns the hostname of the URL.
|
||||
///
|
||||
/// # Returns
|
||||
/// * `String` - The hostname of the URL, or an empty string if not set.
|
||||
///
|
||||
pub fn hostname(&self) -> String {
|
||||
self.0.host_str().unwrap_or("").to_string()
|
||||
}
|
||||
|
||||
/// port returns the port of the URL as a string, defaulting to "80" for http and "443" for https if not set.
|
||||
///
|
||||
/// # Returns
|
||||
/// * `String` - The port of the URL as a string.
|
||||
///
|
||||
pub fn port(&self) -> String {
|
||||
match self.0.port() {
|
||||
Some(p) => p.to_string(),
|
||||
None => match self.0.scheme() {
|
||||
"http" => "80".to_string(),
|
||||
"https" => "443".to_string(),
|
||||
_ => "".to_string(),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
/// scheme returns the scheme of the URL.
|
||||
///
|
||||
/// # Returns
|
||||
/// * `&str` - The scheme of the URL.
|
||||
///
|
||||
pub fn scheme(&self) -> &str {
|
||||
self.0.scheme()
|
||||
}
|
||||
|
||||
/// url returns a reference to the underlying Url.
|
||||
///
|
||||
/// # Returns
|
||||
/// * `&Url` - Reference to the underlying Url.
|
||||
///
|
||||
pub fn url(&self) -> &Url {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Display for ParsedURL {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
let mut url = self.0.clone();
|
||||
if let Some(host) = url.host_str().map(|h| h.to_string())
|
||||
&& let Some(port) = url.port()
|
||||
&& ((url.scheme() == "http" && port == 80) || (url.scheme() == "https" && port == 443))
|
||||
{
|
||||
let _ = url.set_host(Some(&host));
|
||||
let _ = url.set_port(None);
|
||||
}
|
||||
let mut s = url.to_string();
|
||||
|
||||
// If the URL ends with a slash and the path is just "/", remove the trailing slash.
|
||||
if s.ends_with('/') && url.path() == "/" {
|
||||
s.pop();
|
||||
}
|
||||
|
||||
write!(f, "{s}")
|
||||
}
|
||||
}
|
||||
|
||||
impl serde::Serialize for ParsedURL {
|
||||
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
|
||||
where
|
||||
S: serde::Serializer,
|
||||
{
|
||||
serializer.serialize_str(&self.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
impl<'de> serde::Deserialize<'de> for ParsedURL {
|
||||
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
|
||||
where
|
||||
D: serde::Deserializer<'de>,
|
||||
{
|
||||
let s: String = serde::Deserialize::deserialize(deserializer)?;
|
||||
if s.is_empty() {
|
||||
Ok(ParsedURL(Url::parse("about:blank").unwrap()))
|
||||
} else {
|
||||
parse_url(&s).map_err(serde::de::Error::custom)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// parse_url parses a string into a ParsedURL, with host validation and path cleaning.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `s` - The URL string to parse.
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(ParsedURL)` - If parsing is successful.
|
||||
/// * `Err(NetError)` - If parsing fails or host is invalid.
|
||||
///
|
||||
/// # Errors
|
||||
/// Returns NetError if parsing fails or host is invalid.
|
||||
///
|
||||
pub fn parse_url(s: &str) -> Result<ParsedURL, NetError> {
|
||||
if let Some(scheme_end) = s.find("://")
|
||||
&& s[scheme_end + 3..].starts_with('/')
|
||||
{
|
||||
let scheme = &s[..scheme_end];
|
||||
if !scheme.is_empty() {
|
||||
return Err(NetError::SchemeWithEmptyHost);
|
||||
}
|
||||
}
|
||||
|
||||
let mut uu = Url::parse(s).map_err(|e| NetError::ParseError(e.to_string()))?;
|
||||
if uu.host_str().is_none_or(|h| h.is_empty()) {
|
||||
if uu.scheme() != "" {
|
||||
return Err(NetError::SchemeWithEmptyHost);
|
||||
}
|
||||
} else {
|
||||
let port_str = uu.port().map(|p| p.to_string()).unwrap_or_else(|| match uu.scheme() {
|
||||
"http" => "80".to_string(),
|
||||
"https" => "443".to_string(),
|
||||
_ => "".to_string(),
|
||||
});
|
||||
|
||||
if !port_str.is_empty() {
|
||||
let host_port = format!("{}:{}", uu.host_str().unwrap(), port_str);
|
||||
parse_host(&host_port)?; // Validate host.
|
||||
}
|
||||
}
|
||||
|
||||
// Clean path: Use Url's path_segments to normalize.
|
||||
if !uu.path().is_empty() {
|
||||
// Url automatically cleans paths, but we ensure trailing slash if original had it.
|
||||
let mut cleaned_path = String::new();
|
||||
for comp in Path::new(uu.path()).components() {
|
||||
use std::path::Component;
|
||||
match comp {
|
||||
Component::RootDir => cleaned_path.push('/'),
|
||||
Component::Normal(s) => {
|
||||
if !cleaned_path.ends_with('/') {
|
||||
cleaned_path.push('/');
|
||||
}
|
||||
cleaned_path.push_str(&s.to_string_lossy());
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
if s.ends_with('/') && !cleaned_path.ends_with('/') {
|
||||
cleaned_path.push('/');
|
||||
}
|
||||
if cleaned_path.is_empty() {
|
||||
cleaned_path.push('/');
|
||||
}
|
||||
uu.set_path(&cleaned_path);
|
||||
}
|
||||
|
||||
Ok(ParsedURL(uu))
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
/// parse_http_url parses a string into a ParsedURL, ensuring the scheme is http or https.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `s` - The URL string to parse.
|
||||
///
|
||||
/// # Returns
|
||||
/// * `Ok(ParsedURL)` - If parsing is successful and scheme is http/https.
|
||||
/// * `Err(NetError)` - If parsing fails or scheme is not http/https.
|
||||
///
|
||||
pub fn parse_http_url(s: &str) -> Result<ParsedURL, NetError> {
|
||||
let u = parse_url(s)?;
|
||||
match u.0.scheme() {
|
||||
"http" | "https" => Ok(u),
|
||||
_ => Err(NetError::UnexpectedScheme(u.0.scheme().to_string())),
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
/// is_network_or_host_down checks if an error indicates network or host down, considering timeouts.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `err` - The std::io::Error to check.
|
||||
/// * `expect_timeouts` - Whether timeouts are expected.
|
||||
///
|
||||
/// # Returns
|
||||
/// * `bool` - True if the error indicates network or host down, false otherwise.
|
||||
///
|
||||
pub fn is_network_or_host_down(err: &std::io::Error, expect_timeouts: bool) -> bool {
|
||||
if err.kind() == std::io::ErrorKind::TimedOut {
|
||||
return !expect_timeouts;
|
||||
}
|
||||
// Simplified checks based on Go logic; adapt for Rust as needed
|
||||
let err_str = err.to_string().to_lowercase();
|
||||
err_str.contains("connection reset by peer")
|
||||
|| err_str.contains("connection timed out")
|
||||
|| err_str.contains("broken pipe")
|
||||
|| err_str.contains("use of closed network connection")
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
/// is_conn_reset_err checks if an error indicates a connection reset by peer.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `err` - The std::io::Error to check.
|
||||
///
|
||||
/// # Returns
|
||||
/// * `bool` - True if the error indicates connection reset, false otherwise.
|
||||
///
|
||||
pub fn is_conn_reset_err(err: &std::io::Error) -> bool {
|
||||
err.to_string().contains("connection reset by peer") || matches!(err.raw_os_error(), Some(libc::ECONNRESET))
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
/// is_conn_refused_err checks if an error indicates a connection refused.
|
||||
///
|
||||
/// # Arguments
|
||||
/// * `err` - The std::io::Error to check.
|
||||
///
|
||||
/// # Returns
|
||||
/// * `bool` - True if the error indicates connection refused, false otherwise.
|
||||
///
|
||||
pub fn is_conn_refused_err(err: &std::io::Error) -> bool {
|
||||
err.to_string().contains("connection refused") || matches!(err.raw_os_error(), Some(libc::ECONNREFUSED))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn parse_host_with_empty_string_returns_error() {
|
||||
let result = parse_host("");
|
||||
assert!(matches!(result, Err(NetError::InvalidArgument)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_host_with_valid_ipv4() {
|
||||
let result = parse_host("192.168.1.1:8080");
|
||||
assert!(result.is_ok());
|
||||
let host = result.unwrap();
|
||||
assert_eq!(host.name, "192.168.1.1");
|
||||
assert_eq!(host.port, Some(8080));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_host_with_valid_hostname() {
|
||||
let result = parse_host("example.com:443");
|
||||
assert!(result.is_ok());
|
||||
let host = result.unwrap();
|
||||
assert_eq!(host.name, "example.com");
|
||||
assert_eq!(host.port, Some(443));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_host_with_ipv6_brackets() {
|
||||
let result = parse_host("[::1]:8080");
|
||||
assert!(result.is_ok());
|
||||
let host = result.unwrap();
|
||||
assert_eq!(host.name, "::1");
|
||||
assert_eq!(host.port, Some(8080));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_host_with_bare_ipv6_without_port() {
|
||||
let result = parse_host("::1");
|
||||
assert!(result.is_ok());
|
||||
let host = result.unwrap();
|
||||
assert_eq!(host.name, "::1");
|
||||
assert_eq!(host.port, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_host_with_ipv6_zone_without_port() {
|
||||
let result = parse_host("fe80::1%eth0");
|
||||
assert!(result.is_ok());
|
||||
let host = result.unwrap();
|
||||
assert_eq!(host.name, "fe80::1%eth0");
|
||||
assert_eq!(host.port, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_host_with_bracketed_ipv6_zone_and_port() {
|
||||
let result = parse_host("[fe80::1%eth0]:9000");
|
||||
assert!(result.is_ok());
|
||||
let host = result.unwrap();
|
||||
assert_eq!(host.name, "fe80::1%eth0");
|
||||
assert_eq!(host.port, Some(9000));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_host_with_bracketed_ipv6_without_port() {
|
||||
let result = parse_host("[::1]");
|
||||
assert!(result.is_ok());
|
||||
let host = result.unwrap();
|
||||
assert_eq!(host.name, "::1");
|
||||
assert_eq!(host.port, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_host_with_invalid_ipv6_missing_bracket() {
|
||||
let result = parse_host("::1]:8080");
|
||||
assert!(matches!(result, Err(NetError::MissingBracket)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_host_with_invalid_hostname() {
|
||||
let result = parse_host("invalid..host:80");
|
||||
assert!(matches!(result, Err(NetError::InvalidHost)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_host_without_port() {
|
||||
let result = parse_host("example.com");
|
||||
assert!(result.is_ok());
|
||||
let host = result.unwrap();
|
||||
assert_eq!(host.name, "example.com");
|
||||
assert_eq!(host.port, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn host_is_empty_when_name_is_empty() {
|
||||
let host = Host {
|
||||
name: "".to_string(),
|
||||
port: None,
|
||||
};
|
||||
assert!(host.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn host_is_not_empty_when_name_present() {
|
||||
let host = Host {
|
||||
name: "example.com".to_string(),
|
||||
port: Some(80),
|
||||
};
|
||||
assert!(!host.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn host_to_string_with_port() {
|
||||
let host = Host {
|
||||
name: "example.com".to_string(),
|
||||
port: Some(80),
|
||||
};
|
||||
assert_eq!(host.to_string(), "example.com:80");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn host_to_string_without_port() {
|
||||
let host = Host {
|
||||
name: "example.com".to_string(),
|
||||
port: None,
|
||||
};
|
||||
assert_eq!(host.to_string(), "example.com");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn host_equal_when_same() {
|
||||
let host1 = Host {
|
||||
name: "example.com".to_string(),
|
||||
port: Some(80),
|
||||
};
|
||||
let host2 = Host {
|
||||
name: "example.com".to_string(),
|
||||
port: Some(80),
|
||||
};
|
||||
assert!(host1.equal(&host2));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn host_not_equal_when_different() {
|
||||
let host1 = Host {
|
||||
name: "example.com".to_string(),
|
||||
port: Some(80),
|
||||
};
|
||||
let host2 = Host {
|
||||
name: "example.com".to_string(),
|
||||
port: Some(443),
|
||||
};
|
||||
assert!(!host1.equal(&host2));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_url_with_valid_http_url() {
|
||||
let result = parse_url("http://example.com/path");
|
||||
assert!(result.is_ok());
|
||||
let parsed = result.unwrap();
|
||||
assert_eq!(parsed.hostname(), "example.com");
|
||||
assert_eq!(parsed.port(), "80");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_url_with_valid_https_url() {
|
||||
let result = parse_url("https://example.com:443/path");
|
||||
assert!(result.is_ok());
|
||||
let parsed = result.unwrap();
|
||||
assert_eq!(parsed.hostname(), "example.com");
|
||||
assert_eq!(parsed.port(), "443");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_url_with_scheme_but_empty_host() {
|
||||
let result = parse_url("http:///path");
|
||||
assert!(matches!(result, Err(NetError::SchemeWithEmptyHost)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_url_with_invalid_host() {
|
||||
let result = parse_url("http://invalid..host/path");
|
||||
assert!(matches!(result, Err(NetError::InvalidHost)));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_url_with_path_cleaning() {
|
||||
let result = parse_url("http://example.com//path/../path/");
|
||||
assert!(result.is_ok());
|
||||
let parsed = result.unwrap();
|
||||
assert_eq!(parsed.0.path(), "/path/");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_http_url_with_http_scheme() {
|
||||
let result = parse_http_url("http://example.com");
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_http_url_with_https_scheme() {
|
||||
let result = parse_http_url("https://example.com");
|
||||
assert!(result.is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_http_url_with_invalid_scheme() {
|
||||
let result = parse_http_url("ftp://example.com");
|
||||
assert!(matches!(result, Err(NetError::UnexpectedScheme(_))));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parsed_url_is_empty_when_url_is_empty() {
|
||||
let url = ParsedURL(Url::parse("about:blank").unwrap());
|
||||
assert!(url.is_empty());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parsed_url_hostname() {
|
||||
let url = ParsedURL(Url::parse("http://example.com:8080").unwrap());
|
||||
assert_eq!(url.hostname(), "example.com");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parsed_url_port() {
|
||||
let url = ParsedURL(Url::parse("http://example.com:8080").unwrap());
|
||||
assert_eq!(url.port(), "8080");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parsed_url_to_string_removes_default_ports() {
|
||||
let url = ParsedURL(Url::parse("http://example.com:80").unwrap());
|
||||
assert_eq!(url.to_string(), "http://example.com");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn is_network_or_host_down_with_timeout() {
|
||||
let err = std::io::Error::new(std::io::ErrorKind::TimedOut, "timeout");
|
||||
assert!(is_network_or_host_down(&err, false));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn is_network_or_host_down_with_expected_timeout() {
|
||||
let err = std::io::Error::new(std::io::ErrorKind::TimedOut, "timeout");
|
||||
assert!(!is_network_or_host_down(&err, true));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn is_conn_reset_err_with_reset_message() {
|
||||
let err = std::io::Error::other("connection reset by peer");
|
||||
assert!(is_conn_reset_err(&err));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn is_conn_refused_err_with_refused_message() {
|
||||
let err = std::io::Error::other("connection refused");
|
||||
assert!(is_conn_refused_err(&err));
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user