优化加密状态下的重连

This commit is contained in:
lubeilin
2024-04-13 11:13:37 +08:00
parent e9ec6e8903
commit a657eae599
8 changed files with 112 additions and 63 deletions
+1
View File
@@ -3,6 +3,7 @@ syntax = "proto3";
message HandshakeRequest { message HandshakeRequest {
string version = 1; string version = 1;
bool secret = 2; bool secret = 2;
string key_finger = 3;
} }
message HandshakeResponse { message HandshakeResponse {
string version = 1; string version = 1;
+5 -1
View File
@@ -74,6 +74,7 @@ impl Deref for Context {
/// 对称网络增加的udp socket数目,有助于增加打洞成功率 /// 对称网络增加的udp socket数目,有助于增加打洞成功率
pub const SYMMETRIC_CHANNEL_NUM: usize = 100; pub const SYMMETRIC_CHANNEL_NUM: usize = 100;
const PACKET_LOSS_RATE_DENOMINATOR: u32 = 100_0000; const PACKET_LOSS_RATE_DENOMINATOR: u32 = 100_0000;
pub struct ContextInner { pub struct ContextInner {
// 核心udp socket // 核心udp socket
pub(crate) main_udp_socket: Vec<UdpSocket>, pub(crate) main_udp_socket: Vec<UdpSocket>,
@@ -198,6 +199,9 @@ impl ContextInner {
self.send_main_udp(self.main_index.load(Ordering::Relaxed), buf, addr) self.send_main_udp(self.main_index.load(Ordering::Relaxed), buf, addr)
} }
} }
pub fn is_default_route(&self, route_key: RouteKey) -> bool {
self.is_tcp == route_key.is_tcp && self.main_index.load(Ordering::Relaxed) == route_key.index
}
pub fn change_main_index(&self) { pub fn change_main_index(&self) {
let index = (self.main_index.load(Ordering::Relaxed) + 1) % self.main_udp_socket.len(); let index = (self.main_index.load(Ordering::Relaxed) + 1) % self.main_udp_socket.len();
self.main_index.store(index, Ordering::Relaxed); self.main_index.store(index, Ordering::Relaxed);
@@ -295,7 +299,7 @@ impl ContextInner {
pub struct RouteTable { pub struct RouteTable {
pub(crate) route_table: pub(crate) route_table:
RwLock<HashMap<Ipv4Addr, (AtomicUsize, Vec<(Route, AtomicCell<Instant>)>)>>, RwLock<HashMap<Ipv4Addr, (AtomicUsize, Vec<(Route, AtomicCell<Instant>)>)>>,
first_latency: bool, first_latency: bool,
channel_num: usize, channel_num: usize,
use_channel_type: UseChannelType, use_channel_type: UseChannelType,
+26 -20
View File
@@ -1,7 +1,7 @@
use crate::protocol::NetPacket;
use std::io; use std::io;
use { use {
crate::protocol::body::{RsaSecretBody, RSA_ENCRYPTION_RESERVED}, crate::protocol::body::{RSA_ENCRYPTION_RESERVED, RsaSecretBody},
rand::Rng, rand::Rng,
rsa::pkcs8::der::Decode, rsa::pkcs8::der::Decode,
rsa::RsaPublicKey, rsa::RsaPublicKey,
@@ -9,6 +9,8 @@ use {
spki::{DecodePublicKey, EncodePublicKey}, spki::{DecodePublicKey, EncodePublicKey},
}; };
use crate::protocol::NetPacket;
#[derive(Clone)] #[derive(Clone)]
pub struct RsaCipher { pub struct RsaCipher {
inner: Inner, inner: Inner,
@@ -16,13 +18,15 @@ pub struct RsaCipher {
#[derive(Clone)] #[derive(Clone)]
struct Inner { struct Inner {
public_key: RsaPublicKey, public_key: RsaPublicKey,
finger:String,
} }
impl RsaCipher { impl RsaCipher {
pub fn new(der: &[u8]) -> io::Result<Self> { pub fn new(der: &[u8]) -> io::Result<Self> {
match RsaPublicKey::from_public_key_der(der) { match RsaPublicKey::from_public_key_der(der) {
Ok(public_key) => { Ok(public_key) => {
let inner = Inner { public_key }; let finger = finger(&public_key)?;
let inner = Inner { public_key,finger };
Ok(Self { inner }) Ok(Self { inner })
} }
Err(e) => Err(io::Error::new( Err(e) => Err(io::Error::new(
@@ -31,30 +35,32 @@ impl RsaCipher {
)), )),
} }
} }
pub fn finger(&self) ->&String{
pub fn finger(&self) -> io::Result<String> { &self.inner.finger
match self.inner.public_key.to_public_key_der() { }
Ok(der) => match rsa::pkcs8::SubjectPublicKeyInfoRef::from_der(der.as_bytes()) { pub fn public_key(&self) -> io::Result<&RsaPublicKey> {
Ok(spki) => match spki.fingerprint_base64() { return Ok(&self.inner.public_key);
Ok(finger) => Ok(finger), }
Err(e) => Err(io::Error::new( }
io::ErrorKind::Other, pub fn finger(public_key: &RsaPublicKey) -> io::Result<String> {
format!("fingerprint_base64 error {}", e), match public_key.to_public_key_der() {
)), Ok(der) => match rsa::pkcs8::SubjectPublicKeyInfoRef::from_der(der.as_bytes()) {
}, Ok(spki) => match spki.fingerprint_base64() {
Ok(finger) => Ok(finger),
Err(e) => Err(io::Error::new( Err(e) => Err(io::Error::new(
io::ErrorKind::Other, io::ErrorKind::Other,
format!("from_der error {}", e), format!("fingerprint_base64 error {}", e),
)), )),
}, },
Err(e) => Err(io::Error::new( Err(e) => Err(io::Error::new(
io::ErrorKind::Other, io::ErrorKind::Other,
format!("to_public_key_der error {}", e), format!("from_der error {}", e),
)), )),
} },
} Err(e) => Err(io::Error::new(
pub fn public_key(&self) -> io::Result<&RsaPublicKey> { io::ErrorKind::Other,
return Ok(&self.inner.public_key); format!("to_public_key_der error {}", e),
)),
} }
} }
+2 -1
View File
@@ -79,6 +79,7 @@ impl Vnt {
config.token.clone(), config.token.clone(),
config.ip, config.ip,
config.password.is_some(), config.password.is_some(),
config.server_encrypt,
config.device_id.clone(), config.device_id.clone(),
config.server_address_str.clone(), config.server_address_str.clone(),
); );
@@ -144,7 +145,7 @@ impl Vnt {
let down_counter = let down_counter =
U64Adder::with_capacity(config.ports.as_ref().map(|v| v.len()).unwrap_or_default() + 8); U64Adder::with_capacity(config.ports.as_ref().map(|v| v.len()).unwrap_or_default() + 8);
let down_count_watcher = down_counter.watch(); let down_count_watcher = down_counter.watch();
let handshake = Handshake::new(); let handshake = Handshake::new(rsa_cipher.clone());
let handler = RecvDataHandler::new( let handler = RecvDataHandler::new(
#[cfg(feature = "server_encrypt")] #[cfg(feature = "server_encrypt")]
rsa_cipher, rsa_cipher,
+33 -25
View File
@@ -4,6 +4,7 @@ use std::sync::Arc;
use std::time::{Duration, Instant}; use std::time::{Duration, Instant};
use crossbeam_utils::atomic::AtomicCell; use crossbeam_utils::atomic::AtomicCell;
use parking_lot::Mutex;
use protobuf::Message; use protobuf::Message;
use crate::channel::context::Context; use crate::channel::context::Context;
@@ -27,11 +28,13 @@ pub enum HandshakeEnum {
#[derive(Clone)] #[derive(Clone)]
pub struct Handshake { pub struct Handshake {
time: Arc<AtomicCell<Instant>>, time: Arc<AtomicCell<Instant>>,
rsa_cipher: Arc<Mutex<Option<RsaCipher>>>
} }
impl Handshake { impl Handshake {
pub fn new() -> Self { pub fn new( rsa_cipher: Arc<Mutex<Option<RsaCipher>>>) -> Self {
Handshake { Handshake {
time: Arc::new(AtomicCell::new(Instant::now() - Duration::from_secs(60))), time: Arc::new(AtomicCell::new(Instant::now() - Duration::from_secs(60))),
rsa_cipher
} }
} }
pub fn send(&self, context: &Context, secret: bool, addr: SocketAddr) -> io::Result<()> { pub fn send(&self, context: &Context, secret: bool, addr: SocketAddr) -> io::Result<()> {
@@ -40,37 +43,42 @@ impl Handshake {
if last.elapsed() < Duration::from_secs(3) { if last.elapsed() < Duration::from_secs(3) {
return Ok(()); return Ok(());
} }
let request_packet = handshake_request_packet(secret)?; let request_packet = self.handshake_request_packet(secret)?;
log::info!("发送握手请求,secret={},{:?}", secret, addr); log::info!("发送握手请求,secret={},{:?}", secret, addr);
context.send_default(request_packet.buffer(), addr)?; context.send_default(request_packet.buffer(), addr)?;
self.time.store(Instant::now()); self.time.store(Instant::now());
Ok(()) Ok(())
} }
/// 第一次握手数据
pub fn handshake_request_packet(&self,secret: bool) -> io::Result<NetPacket<Vec<u8>>> {
let mut request = HandshakeRequest::new();
request.secret = secret;
request.version = crate::VNT_VERSION.to_string();
if let Some(finger) = self.rsa_cipher.lock().as_ref().map(|v| v.finger().clone()){
request.key_finger = finger;
}
let bytes = request.write_to_bytes().map_err(|e| {
io::Error::new(
io::ErrorKind::Other,
format!("handshake_request_packet {:?}", e),
)
})?;
let buf = vec![0u8; 12 + bytes.len()];
let mut net_packet = NetPacket::new(buf)?;
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_destination(GATEWAY_IP);
net_packet.set_source(SELF_IP);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol(service_packet::Protocol::HandshakeRequest.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_payload(&bytes)?;
Ok(net_packet)
}
} }
/// 第一次握手数据
pub fn handshake_request_packet(secret: bool) -> io::Result<NetPacket<Vec<u8>>> {
let mut request = HandshakeRequest::new();
request.secret = secret;
request.version = crate::VNT_VERSION.to_string();
let bytes = request.write_to_bytes().map_err(|e| {
io::Error::new(
io::ErrorKind::Other,
format!("handshake_request_packet {:?}", e),
)
})?;
let buf = vec![0u8; 12 + bytes.len()];
let mut net_packet = NetPacket::new(buf)?;
net_packet.set_version(Version::V1);
net_packet.set_gateway_flag(true);
net_packet.set_destination(GATEWAY_IP);
net_packet.set_source(SELF_IP);
net_packet.set_protocol(Protocol::Service);
net_packet.set_transport_protocol(service_packet::Protocol::HandshakeRequest.into());
net_packet.first_set_ttl(MAX_TTL);
net_packet.set_payload(&bytes)?;
Ok(net_packet)
}
/// 第二次加密握手 /// 第二次加密握手
#[cfg(feature = "server_encrypt")] #[cfg(feature = "server_encrypt")]
+4 -4
View File
@@ -6,14 +6,14 @@ use std::time::Duration;
use crossbeam_utils::atomic::AtomicCell; use crossbeam_utils::atomic::AtomicCell;
use mio::net::TcpStream; use mio::net::TcpStream;
use crate::{ErrorInfo, VntCallback};
use crate::channel::context::Context; use crate::channel::context::Context;
use crate::channel::idle::{Idle, IdleType}; use crate::channel::idle::{Idle, IdleType};
use crate::channel::sender::AcceptSocketSender; use crate::channel::sender::AcceptSocketSender;
use crate::handle::{BaseConfigInfo, ConnectStatus, CurrentDeviceInfo};
use crate::handle::callback::{ConnectInfo, ErrorType}; use crate::handle::callback::{ConnectInfo, ErrorType};
use crate::handle::handshaker::Handshake; use crate::handle::handshaker::Handshake;
use crate::handle::{handshaker, BaseConfigInfo, ConnectStatus, CurrentDeviceInfo};
use crate::util::Scheduler; use crate::util::Scheduler;
use crate::{ErrorInfo, VntCallback};
pub fn idle_route<Call: VntCallback>( pub fn idle_route<Call: VntCallback>(
scheduler: &Scheduler, scheduler: &Scheduler,
@@ -133,11 +133,11 @@ fn check_gateway_channel<Call: VntCallback>(
//需要重连 //需要重连
call.connect(ConnectInfo::new(*count, current_device.connect_server)); call.connect(ConnectInfo::new(*count, current_device.connect_server));
log::info!("发送握手请求,{:?}", config); log::info!("发送握手请求,{:?}", config);
if let Err(e) = handshake.send(context, config.client_secret, current_device.connect_server) if let Err(e) = handshake.send(context, config.server_secret, current_device.connect_server)
{ {
log::warn!("{:?}", e); log::warn!("{:?}", e);
if context.is_main_tcp() { if context.is_main_tcp() {
let request_packet = handshaker::handshake_request_packet(config.client_secret)?; let request_packet = handshake.handshake_request_packet(config.server_secret)?;
//tcp需要重连 //tcp需要重连
let tcp_stream = std::net::TcpStream::connect_timeout( let tcp_stream = std::net::TcpStream::connect_timeout(
&current_device.connect_server, &current_device.connect_server,
+3
View File
@@ -51,6 +51,7 @@ pub struct BaseConfigInfo {
pub token: String, pub token: String,
pub ip: Option<Ipv4Addr>, pub ip: Option<Ipv4Addr>,
pub client_secret: bool, pub client_secret: bool,
pub server_secret: bool,
pub device_id: String, pub device_id: String,
pub server_addr: String, pub server_addr: String,
} }
@@ -61,6 +62,7 @@ impl BaseConfigInfo {
token: String, token: String,
ip: Option<Ipv4Addr>, ip: Option<Ipv4Addr>,
client_secret: bool, client_secret: bool,
server_secret: bool,
device_id: String, device_id: String,
server_addr: String, server_addr: String,
) -> Self { ) -> Self {
@@ -69,6 +71,7 @@ impl BaseConfigInfo {
token, token,
ip, ip,
client_secret, client_secret,
server_secret,
device_id, device_id,
server_addr, server_addr,
} }
+38 -12
View File
@@ -11,30 +11,30 @@ use protobuf::Message;
use packet::icmp::{icmp, Kind}; use packet::icmp::{icmp, Kind};
use packet::ip::ipv4; use packet::ip::ipv4;
use packet::ip::ipv4::packet::IpV4Packet; use packet::ip::ipv4::packet::IpV4Packet;
use tun::device::IFace;
use tun::Device; use tun::Device;
use tun::device::IFace;
use crate::channel::context::Context;
use crate::channel::{Route, RouteKey}; use crate::channel::{Route, RouteKey};
use crate::channel::context::Context;
use crate::cipher::Cipher; use crate::cipher::Cipher;
#[cfg(feature = "server_encrypt")] #[cfg(feature = "server_encrypt")]
use crate::cipher::RsaCipher; use crate::cipher::RsaCipher;
use crate::external_route::ExternalRoute; use crate::external_route::ExternalRoute;
use crate::handle::{
BaseConfigInfo, ConnectStatus, CurrentDeviceInfo, GATEWAY_IP, PeerDeviceInfo, registrar,
};
use crate::handle::callback::{ErrorInfo, ErrorType, HandshakeInfo, RegisterInfo, VntCallback}; use crate::handle::callback::{ErrorInfo, ErrorType, HandshakeInfo, RegisterInfo, VntCallback};
#[cfg(feature = "server_encrypt")] #[cfg(feature = "server_encrypt")]
use crate::handle::handshaker; use crate::handle::handshaker;
use crate::handle::handshaker::Handshake; use crate::handle::handshaker::Handshake;
use crate::handle::recv_data::PacketHandler; use crate::handle::recv_data::PacketHandler;
use crate::handle::{
registrar, BaseConfigInfo, ConnectStatus, CurrentDeviceInfo, PeerDeviceInfo, GATEWAY_IP,
};
use crate::nat::NatTest; use crate::nat::NatTest;
use crate::proto; use crate::proto;
use crate::proto::message::{DeviceList, HandshakeResponse, RegistrationResponse}; use crate::proto::message::{DeviceList, HandshakeResponse, RegistrationResponse};
use crate::protocol::{ip_turn_packet, MAX_TTL, NetPacket, Protocol, service_packet, Version};
use crate::protocol::body::ENCRYPTION_RESERVED; use crate::protocol::body::ENCRYPTION_RESERVED;
use crate::protocol::control_packet::ControlPacket; use crate::protocol::control_packet::ControlPacket;
use crate::protocol::error_packet::InErrorPacket; use crate::protocol::error_packet::InErrorPacket;
use crate::protocol::{ip_turn_packet, service_packet, NetPacket, Protocol, Version, MAX_TTL};
/// 处理来源于服务端的包 /// 处理来源于服务端的包
#[derive(Clone)] #[derive(Clone)]
@@ -136,13 +136,35 @@ impl<Call: VntCallback> PacketHandler for ServerPacketHandler<Call> {
HandshakeResponse::parse_from_bytes(net_packet.payload()).map_err(|e| { HandshakeResponse::parse_from_bytes(net_packet.payload()).map_err(|e| {
io::Error::new(io::ErrorKind::Other, format!("HandshakeResponse {:?}", e)) io::Error::new(io::ErrorKind::Other, format!("HandshakeResponse {:?}", e))
})?; })?;
log::info!("握手响应:{:?},{}",route_key, response);
//如果开启了加密,则发送加密握手请求 //如果开启了加密,则发送加密握手请求
#[cfg(feature = "server_encrypt")] #[cfg(feature = "server_encrypt")]
if let Some(key) = self.server_cipher.key() { if let Some(key) = self.server_cipher.key() {
{
let guard = self.rsa_cipher.lock();
if let Some(rsa_cipher) = guard.as_ref(){
if rsa_cipher.finger()==&response.key_finger{
let packet = handshaker::secret_handshake_request_packet(
rsa_cipher,
self.config_info.token.clone(),
key,
)?;
drop(guard);
context.send_by_key(packet.buffer(), route_key)?;
return Ok(());
}
log::info!("服务端密钥对变化,原指纹:{:?},新指纹:{:?}", rsa_cipher.finger(),response.key_finger);
}
drop(guard);
}
let rsa_cipher = RsaCipher::new(&response.public_key)?; let rsa_cipher = RsaCipher::new(&response.public_key)?;
if rsa_cipher.finger() != &response.key_finger {
log::info!("服务端密钥和指纹不匹 配拒绝握手,指纹1:{:?},指纹2:{:?}", rsa_cipher.finger(),response.key_finger);
return Ok(());
}
let handshake_info = HandshakeInfo::new( let handshake_info = HandshakeInfo::new(
rsa_cipher.public_key()?.clone(), rsa_cipher.public_key()?.clone(),
rsa_cipher.finger()?, response.key_finger,
response.version, response.version,
); );
log::info!("加密握手请求:{:?}", handshake_info); log::info!("加密握手请求:{:?}", handshake_info);
@@ -158,7 +180,9 @@ impl<Call: VntCallback> PacketHandler for ServerPacketHandler<Call> {
} }
return Ok(()); return Ok(());
} }
if let Ok(rsa_cipher) = RsaCipher::new(&response.public_key){
self.rsa_cipher.lock().replace(rsa_cipher);
}
let handshake_info = HandshakeInfo::new_no_secret(response.version); let handshake_info = HandshakeInfo::new_no_secret(response.version);
if self.callback.handshake(handshake_info) { if self.callback.handshake(handshake_info) {
//没有加密,则发送注册请求 //没有加密,则发送注册请求
@@ -334,9 +358,11 @@ impl<Call: VntCallback> ServerPacketHandler<Call> {
self.set_device_info_list(response.device_info_list, response.epoch as _); self.set_device_info_list(response.device_info_list, response.epoch as _);
} }
service_packet::Protocol::SecretHandshakeResponse => { service_packet::Protocol::SecretHandshakeResponse => {
log::info!("SecretHandshakeResponse"); if context.is_default_route(route_key){
//加密握手结束,发送注册数据 log::info!("SecretHandshakeResponse");
self.register(current_device, context)?; //加密握手结束,发送注册数据
self.register(current_device, context)?;
}
} }
_ => { _ => {
log::warn!( log::warn!(
@@ -415,7 +441,7 @@ impl<Call: VntCallback> ServerPacketHandler<Call> {
drop(dev); drop(dev);
} }
self.handshake self.handshake
.send(context, self.config_info.client_secret, route_key.addr)?; .send(context, self.config_info.server_secret, route_key.addr)?;
// self.register(current_device, context, route_key)?; // self.register(current_device, context, route_key)?;
} }
InErrorPacket::AddressExhausted => { InErrorPacket::AddressExhausted => {