From f9462162a565f2465774f4e195910ab8bb0ff5ec Mon Sep 17 00:00:00 2001 From: "shiro.lee" Date: Wed, 3 Jul 2024 23:00:32 +0800 Subject: [PATCH] =?UTF-8?q?fix:=20=E4=BC=98=E5=8C=96net?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- ecstore/src/endpoint.rs | 169 +++++++++------------------------------ ecstore/src/utils/net.rs | 83 +++++++++---------- 2 files changed, 79 insertions(+), 173 deletions(-) diff --git a/ecstore/src/endpoint.rs b/ecstore/src/endpoint.rs index dd9f67f35..5796df563 100644 --- a/ecstore/src/endpoint.rs +++ b/ecstore/src/endpoint.rs @@ -1,17 +1,12 @@ use super::disks_layout::{DisksLayout, PoolDisksLayout}; use super::error::{Error, Result}; use super::utils::{ - net::{is_local_host, split_host_port}, + net::{is_local_host, is_socket_addr, split_host_port}, string::new_string_set, }; use path_absolutize::Absolutize; use std::fmt::Display; -use std::{ - collections::HashMap, - net::{IpAddr, SocketAddr}, - path::Path, - usize, -}; +use std::{collections::HashMap, path::Path, usize}; use url::{ParseError, Url}; pub const DEFAULT_PORT: u16 = 9000; @@ -45,7 +40,7 @@ pub struct Node { // } /// any type of endpoint. -#[derive(Debug, Clone, Default)] +#[derive(Debug, Clone)] pub struct Endpoint { pub url: url::Url, pub is_local: bool, @@ -57,10 +52,10 @@ pub struct Endpoint { impl Display for Endpoint { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { - if self.url.has_host() { - write!(f, "{}", self.url) - } else { + if self.url.scheme() == "file" { write!(f, "{}", self.url.path()) + } else { + write!(f, "{}", self.url) } } } @@ -77,7 +72,7 @@ impl TryFrom<&str> for Endpoint { let mut is_local = false; let url = match Url::parse(value) { - Ok(url) => { + Ok(mut url) => { // URL style of endpoint. // Valid URL style endpoint is // - Scheme field must contain "http" or "https" @@ -89,6 +84,7 @@ impl TryFrom<&str> for Endpoint { { return Err(Error::from_string("invalid URL endpoint format")); } + if is_empty_path(url.path()) { return Err(Error::from_string("empty or root endpoint is not supported")); } @@ -108,7 +104,12 @@ impl TryFrom<&str> for Endpoint { // Another additional benefit here is that this style also // supports providing \\host\share support as well. #[cfg(windows)] - {} + { + let path = url.path().to_owned(); + if Path::new(&path[1..]).has_root() { + url.set_path(&path[1..]); + } + } url } @@ -135,7 +136,7 @@ impl TryFrom<&str> for Endpoint { is_local = true; url } - Err(err) => return Err(Error::from_string("Convert a file path into an URL failed")), + Err(_) => return Err(Error::from_string("Convert a file path into an URL failed")), } } _ => return Err(Error::from_string(format!("invalid URL endpoint format: {}", e))), @@ -145,7 +146,9 @@ impl TryFrom<&str> for Endpoint { Ok(Endpoint { url, is_local, - ..Default::default() + pool_idx: None, + set_idx: None, + disk_idx: None, }) } } @@ -155,128 +158,34 @@ fn is_empty_path(path: &str) -> bool { ["", "/", "\\"].iter().any(|&v| v.eq(path)) } -// helper for validating if the provided arg is an ip address. -fn is_socket_addr(host: &str) -> bool { - host.parse::().is_ok() || host.parse::().is_ok() -} - impl Endpoint { - pub fn new(arg: &str) -> Result { - if is_empty_path(arg) { - return Err(Error::from_string("不支持空或根endpoint")); - } - - let url = Url::parse(arg).or_else(|e| match e { - ParseError::EmptyHost => Err(Error::from_string("远程地址,域名不能为空")), - ParseError::IdnaError => Err(Error::from_string("域名格式不正确")), - ParseError::InvalidPort => Err(Error::from_string("端口格式不正确")), - ParseError::InvalidIpv4Address => Err(Error::from_string("IP格式不正确")), - ParseError::InvalidIpv6Address => Err(Error::from_string("IP格式不正确")), - ParseError::InvalidDomainCharacter => Err(Error::from_string("域名字符格式不正确")), - // url::ParseError::RelativeUrlWithoutBase => todo!(), - // url::ParseError::RelativeUrlWithCannotBeABaseBase => todo!(), - // url::ParseError::SetHostOnCannotBeABaseUrl => todo!(), - ParseError::Overflow => Err(Error::from_string("长度过长")), - _ => { - if is_host_ip(arg) { - return Err(Error::from_string("无效的URL endpoint格式: 缺少 http 或 https")); - } - - let abs_arg = Path::new(arg).canonicalize()?; - - let abs = abs_arg.to_str().ok_or(Error::from_string("绝对路径错误"))?; - let url = Url::from_file_path(abs).unwrap(); - Ok(url) - } - })?; - - if url.scheme() == "file" { - return Ok(Endpoint { - url: url, - is_local: true, - pool_idx: -1, - set_idx: -1, - disk_idx: -1, - }); - } - - if url.port().is_none() { - return Err(Error::from_string("必须提供端口号")); - } - - if !(url.scheme() == "http" || url.scheme() == "https") { - return Err(Error::from_string("URL endpoint格式无效: Scheme字段必须包含'http'或'https'")); - } - - // 检查路径 - let path = url.path(); - if is_empty_path(path) { - return Err(Error::from_string("URL endpoint不支持空或根路径")); - } - - // TODO: Windows 系统上的路径处理 - #[cfg(windows)] - { - use std::env; - if env::consts::OS == "windows" { - // 处理 Windows 路径的特殊逻辑 - } - } - - Ok(Endpoint { - url: url, - is_local: false, - pool_idx: -1, - set_idx: -1, - disk_idx: -1, - }) - } - - // pub fn host_port_str(&self) -> String { - // if self.url.has_host() && self.port() > 0 { - // return format!("{}:{}", self.hostname(), self.port()); - // } else if self.url.has_host() && self.port() == 0 { - // return self.hostname().to_string(); - // } else if !self.url.has_host() && self.port() > 0 { - // return format!(":{}", self.port()); - // } else { - // return String::new(); - // } - // } - - // pub fn port(&self) -> u16 { - // self.url.port().unwrap_or(0) - // } - // pub fn hostname(&self) -> &str { - // self.url.host_str().unwrap_or("") - // } - + /// returns type of endpoint. pub fn get_type(&self) -> EndpointType { - if self.url.has_host() { - EndpointType::Url - } else { + if self.url.scheme() == "file" { EndpointType::Path + } else { + EndpointType::Url } } - // pub fn get_scheme(&self) -> String { - // self.url.scheme().to_string() - // } - - pub fn set_pool_index(&mut self, idx: i32) { - self.pool_idx = idx + /// sets a specific pool number to this node + pub fn set_pool_index(&mut self, idx: usize) { + self.pool_idx = Some(idx) } - pub fn set_set_index(&mut self, idx: i32) { - self.set_idx = idx + /// sets a specific set number to this node + pub fn set_set_index(&mut self, idx: usize) { + self.set_idx = Some(idx) } - pub fn set_disk_index(&mut self, idx: i32) { - self.disk_idx = idx + /// sets a specific disk number to this node + pub fn set_disk_index(&mut self, idx: usize) { + self.disk_idx = Some(idx) } - fn update_islocal(&mut self) -> Result<()> { - if self.url.has_host() { + /// resolves the host and updates if it is local or not. + fn update_is_local(&mut self) -> Result<()> { + if self.url.scheme() != "file" { self.is_local = is_local_host(self.url.host().unwrap(), self.url.port().unwrap(), DEFAULT_PORT); } @@ -323,7 +232,7 @@ impl Endpoints { let mut eps = Vec::new(); let mut uniq_args = new_string_set(); for (i, arg) in args.iter().enumerate() { - let endpoint = Endpoint::new(arg)?; + let endpoint = Endpoint::try_from(arg.as_str())?; if i == 0 { ep_type = endpoint.get_type(); scheme = endpoint.url.scheme().to_string(); @@ -365,7 +274,7 @@ impl PoolEndpointList { for eps in self.0.iter_mut() { for ep in eps.iter_mut() { // TODO: - ep.update_islocal()? + ep.update_is_local()? } } @@ -573,9 +482,9 @@ pub fn create_pool_endpoints(server_addr: String, pools: &Vec) let mut eps = Endpoints::from_args(set_layout.to_owned())?; // TODO: checkCrossDeviceMounts for (disk_idx, ep) in eps.0.iter_mut().enumerate() { - ep.set_pool_index(pool_idx as i32); - ep.set_set_index(set_idx as i32); - ep.set_disk_index(disk_idx as i32); + ep.set_pool_index(pool_idx); + ep.set_set_index(set_idx); + ep.set_disk_index(disk_idx); endpoints.0.push(ep.to_owned()); } diff --git a/ecstore/src/utils/net.rs b/ecstore/src/utils/net.rs index 734e6a3ad..6c91c4f75 100644 --- a/ecstore/src/utils/net.rs +++ b/ecstore/src/utils/net.rs @@ -1,12 +1,17 @@ use std::{ - collections::HashMap, - net::{IpAddr, ToSocketAddrs}, + collections::HashSet, + net::{IpAddr, SocketAddr, ToSocketAddrs}, }; use anyhow::{Error, Result}; use netif; use url::Host; +// helper for validating if the provided arg is an ip address. +pub fn is_socket_addr(host: &str) -> bool { + host.parse::().is_ok() || host.parse::().is_ok() +} + pub fn split_host_port(s: &str) -> Result<(String, u16)> { let parts: Vec<&str> = s.split(':').collect(); if parts.len() == 2 { @@ -17,38 +22,23 @@ pub fn split_host_port(s: &str) -> Result<(String, u16)> { Err(Error::msg("Invalid address format or port number")) } -// is_local_host 判断是否是本地ip +/// checks if the given parameter correspond to one of +/// the local IP of the current machine pub fn is_local_host(host: Host<&str>, port: u16, local_port: u16) -> bool { let local_ips = must_get_local_ips(); - let local_map = - local_ips - .iter() - .map(|ip| ip.to_string()) - .fold(HashMap::new(), |mut acc, item| { - *acc.entry(item).or_insert(true) = true; - acc - }); - + let local_set: HashSet = local_ips.into_iter().collect(); let is_local_host = match host { Host::Domain(domain) => { - let ips: Vec = (domain, 0) - .to_socket_addrs() - .unwrap_or(Vec::new().into_iter()) - .map(|addr| addr.ip().to_string()) - .collect(); + let ips = match (domain, 0).to_socket_addrs().map(|v| v.map(|v| v.ip()).collect::>()) { + Ok(ips) => ips, + Err(_) => return false, + }; - let mut isok = false; - for ip in ips.iter() { - if local_map.contains_key(ip) { - isok = true; - break; - } - } - isok + ips.iter().any(|ip| local_set.contains(ip)) } - Host::Ipv4(ip) => local_map.contains_key(&ip.to_string()), - Host::Ipv6(ip) => local_map.contains_key(&ip.to_string()), + Host::Ipv4(ip) => local_set.contains(&IpAddr::V4(ip)), + Host::Ipv6(ip) => local_set.contains(&IpAddr::V6(ip)), }; if port > 0 { @@ -58,37 +48,44 @@ pub fn is_local_host(host: Host<&str>, port: u16, local_port: u16) -> bool { is_local_host } +/// returns IPs of local interface pub fn must_get_local_ips() -> Vec { - let mut v: Vec = Vec::new(); - if let Some(up) = netif::up().ok() { - v = up.map(|x| x.address().to_owned()).collect(); + match netif::up() { + Ok(up) => up.map(|x| x.address().to_owned()).collect(), + Err(_) => vec![], } - - v } #[cfg(test)] mod test { - use std::net::Ipv4Addr; + use std::net::{Ipv4Addr, Ipv6Addr}; use super::*; #[test] fn test_must_get_local_ips() { - let ips = must_get_local_ips(); - for ip in ips.iter() { - println!("{:?}", ip) - } + let local_ips = must_get_local_ips(); + let local_set: HashSet = local_ips.into_iter().collect(); + + assert!(local_set.contains(&IpAddr::V4(Ipv4Addr::new(127, 0, 0, 1)))); } #[test] fn test_is_local_host() { - // let host = Host::Ipv4(Ipv4Addr::new(192, 168, 0, 233)); + let host = Host::Domain("localhost"); + let is = is_local_host(host, 0, 9000); + assert!(is); + + let host = Host::Ipv6(Ipv6Addr::new(0, 0, 0, 0, 0, 0, 0, 1)); + let is = is_local_host(host, 0, 9000); + assert!(is); + + let host = Host::Ipv4(Ipv4Addr::new(8, 8, 8, 8)); + let is = is_local_host(host, 0, 9000); + assert!(!is); + let host = Host::Ipv4(Ipv4Addr::new(127, 0, 0, 1)); - // let host = Host::Domain("localhost"); - let port = 0; - let local_port = 9000; - let is = is_local_host(host, port, local_port); - assert!(is) + let is = is_local_host(host, 8000, 9000); + assert!(!is); } }