mirror of
https://github.com/EasyTier/EasyTier.git
synced 2026-09-01 00:39:24 +00:00
248 lines
8.1 KiB
Rust
248 lines
8.1 KiB
Rust
use std::{net::SocketAddr, sync::Arc};
|
|
|
|
use super::{create_connector_by_url, http_connector::TunnelWithInfo};
|
|
use crate::utils::dns::{srv_lookup, txt_resolve};
|
|
use crate::{
|
|
common::{error::Error, global_ctx::ArcGlobalCtx, log},
|
|
proto::common::TunnelInfo,
|
|
tunnel::{IpScheme, IpVersion, Tunnel, TunnelConnector, TunnelError, TunnelScheme},
|
|
};
|
|
use anyhow::Context;
|
|
use dashmap::DashSet;
|
|
use hickory_proto::rr::rdata::SRV;
|
|
use rand::{Rng as _, seq::SliceRandom};
|
|
use strum::VariantArray;
|
|
|
|
fn weighted_choice<T>(options: &[(T, u64)]) -> Option<&T> {
|
|
let total_weight = options.iter().map(|(_, weight)| *weight).sum();
|
|
let mut rng = rand::thread_rng();
|
|
let rand_value = rng.gen_range(0..total_weight);
|
|
let mut accumulated_weight = 0;
|
|
|
|
for (item, weight) in options {
|
|
accumulated_weight += *weight;
|
|
if rand_value < accumulated_weight {
|
|
return Some(item);
|
|
}
|
|
}
|
|
|
|
None
|
|
}
|
|
|
|
#[derive(Debug)]
|
|
pub struct DnsTunnelConnector {
|
|
scheme: TunnelScheme,
|
|
addr: url::Url,
|
|
bind_addrs: Vec<SocketAddr>,
|
|
global_ctx: ArcGlobalCtx,
|
|
ip_version: IpVersion,
|
|
}
|
|
|
|
impl DnsTunnelConnector {
|
|
pub fn new(addr: url::Url, global_ctx: ArcGlobalCtx) -> Self {
|
|
Self {
|
|
scheme: (&addr).try_into().unwrap(),
|
|
addr,
|
|
bind_addrs: Vec::new(),
|
|
global_ctx,
|
|
ip_version: IpVersion::Both,
|
|
}
|
|
}
|
|
|
|
#[tracing::instrument(ret, err)]
|
|
pub async fn handle_txt_record(
|
|
&self,
|
|
domain_name: &str,
|
|
) -> Result<Box<dyn TunnelConnector>, Error> {
|
|
let txt_data = txt_resolve(domain_name)
|
|
.await
|
|
.with_context(|| format!("resolve txt record failed, domain_name: {}", domain_name))?;
|
|
|
|
let candidate_urls = txt_data
|
|
.iter()
|
|
.filter_map(|s| url::Url::parse(s).ok())
|
|
.collect::<Vec<_>>();
|
|
|
|
// shuffle candidate_urls and get the first one
|
|
let url = candidate_urls
|
|
.choose(&mut rand::thread_rng())
|
|
.with_context(|| {
|
|
format!(
|
|
"no valid url found, txt_data: {:?}, expecting an url list split by space",
|
|
txt_data
|
|
)
|
|
})?;
|
|
|
|
let connector =
|
|
create_connector_by_url(url.as_str(), &self.global_ctx, self.ip_version).await?;
|
|
Ok(connector)
|
|
}
|
|
|
|
fn handle_one_srv_record(record: SRV, protocol: IpScheme) -> Result<(url::Url, u64), Error> {
|
|
// port must be non-zero
|
|
if record.port == 0 {
|
|
return Err(anyhow::anyhow!("port must be non-zero").into());
|
|
}
|
|
|
|
let connector_dst = record.target.to_utf8();
|
|
let dst_url = format!("{}://{}:{}", protocol, connector_dst, record.port);
|
|
|
|
Ok((
|
|
dst_url.parse().with_context(|| {
|
|
format!(
|
|
"parse dst_url failed, protocol: {}, connector_dst: {}, port: {}, dst_url: {}",
|
|
protocol, connector_dst, record.port, dst_url
|
|
)
|
|
})?,
|
|
record.priority as _,
|
|
))
|
|
}
|
|
|
|
#[tracing::instrument(ret, err)]
|
|
pub async fn handle_srv_record(
|
|
&self,
|
|
domain_name: &str,
|
|
) -> Result<Box<dyn TunnelConnector>, Error> {
|
|
tracing::info!("handle_srv_record: {}", domain_name);
|
|
|
|
let srv_domains = IpScheme::VARIANTS
|
|
.iter()
|
|
.map(|s| (s, format!("_easytier._{}.{}", s, domain_name)))
|
|
.collect::<Vec<_>>();
|
|
tracing::info!("build srv_domains: {:?}", srv_domains);
|
|
let responses = Arc::new(DashSet::new());
|
|
let srv_lookup_tasks = srv_domains
|
|
.iter()
|
|
.map(|(protocol, srv_domain)| {
|
|
let responses = responses.clone();
|
|
async move {
|
|
for record in srv_lookup(srv_domain).await? {
|
|
let parsed_record = Self::handle_one_srv_record(record, **protocol);
|
|
tracing::info!(?parsed_record, ?srv_domain, "parsed_record");
|
|
if let Err(e) = &parsed_record {
|
|
log::warn!("got invalid srv record {:?}", e);
|
|
continue;
|
|
}
|
|
responses.insert(parsed_record.unwrap());
|
|
}
|
|
Ok::<_, Error>(())
|
|
}
|
|
})
|
|
.collect::<Vec<_>>();
|
|
let _ = futures::future::join_all(srv_lookup_tasks).await;
|
|
|
|
let srv_records = responses.iter().map(|r| r.clone()).collect::<Vec<_>>();
|
|
if srv_records.is_empty() {
|
|
return Err(anyhow::anyhow!("no srv record found").into());
|
|
}
|
|
|
|
let url = weighted_choice(srv_records.as_slice()).with_context(|| {
|
|
format!(
|
|
"failed to choose a srv record, domain_name: {}, srv_records: {:?}",
|
|
domain_name, srv_records
|
|
)
|
|
})?;
|
|
|
|
let connector =
|
|
create_connector_by_url(url.as_str(), &self.global_ctx, self.ip_version).await?;
|
|
Ok(connector)
|
|
}
|
|
}
|
|
|
|
#[async_trait::async_trait]
|
|
impl super::TunnelConnector for DnsTunnelConnector {
|
|
async fn connect(&mut self) -> Result<Box<dyn Tunnel>, TunnelError> {
|
|
let mut conn = match self.scheme {
|
|
TunnelScheme::Txt => self
|
|
.handle_txt_record(
|
|
self.addr
|
|
.host_str()
|
|
.as_ref()
|
|
.ok_or(anyhow::anyhow!("host should not be empty in txt url"))?,
|
|
)
|
|
.await
|
|
.with_context(|| "get txt record url failed")?,
|
|
TunnelScheme::Srv => self
|
|
.handle_srv_record(
|
|
self.addr
|
|
.host_str()
|
|
.as_ref()
|
|
.ok_or(anyhow::anyhow!("host should not be empty in srv url"))?,
|
|
)
|
|
.await
|
|
.with_context(|| "get srv record url failed")?,
|
|
_ => return Err(anyhow::anyhow!("unsupported dns scheme: {:?}", self.scheme).into()),
|
|
};
|
|
let t = conn.connect().await?;
|
|
let info = t.info().unwrap_or_default();
|
|
Ok(Box::new(TunnelWithInfo::new(
|
|
t,
|
|
TunnelInfo {
|
|
local_addr: info.local_addr.clone(),
|
|
remote_addr: Some(self.addr.clone().into()),
|
|
resolved_remote_addr: info
|
|
.resolved_remote_addr
|
|
.clone()
|
|
.or(info.remote_addr.clone()),
|
|
tunnel_type: format!("{}-{}", self.addr.scheme(), info.tunnel_type),
|
|
},
|
|
)))
|
|
}
|
|
|
|
fn remote_url(&self) -> url::Url {
|
|
self.addr.clone()
|
|
}
|
|
|
|
fn set_bind_addrs(&mut self, addrs: Vec<SocketAddr>) {
|
|
self.bind_addrs = addrs;
|
|
}
|
|
|
|
fn set_ip_version(&mut self, ip_version: IpVersion) {
|
|
self.ip_version = ip_version;
|
|
}
|
|
}
|
|
|
|
#[cfg(test)]
|
|
mod tests {
|
|
use super::*;
|
|
use crate::common::global_ctx::tests::get_mock_global_ctx;
|
|
|
|
#[tokio::test]
|
|
async fn test_txt() {
|
|
let url = "txt://txt.easytier.cn";
|
|
let global_ctx = get_mock_global_ctx();
|
|
let mut connector = DnsTunnelConnector::new(url.parse().unwrap(), global_ctx);
|
|
connector.set_ip_version(IpVersion::V4);
|
|
for _ in 0..5 {
|
|
match connector.connect().await {
|
|
Ok(ret) => {
|
|
println!("{:?}", ret.info());
|
|
return;
|
|
}
|
|
Err(e) => {
|
|
println!("{:?}", e);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn test_srv() {
|
|
let url = "srv://easytier.cn";
|
|
let global_ctx = get_mock_global_ctx();
|
|
let mut connector = DnsTunnelConnector::new(url.parse().unwrap(), global_ctx);
|
|
connector.set_ip_version(IpVersion::V4);
|
|
for _ in 0..5 {
|
|
match connector.connect().await {
|
|
Ok(ret) => {
|
|
println!("{:?}", ret.info());
|
|
return;
|
|
}
|
|
Err(e) => {
|
|
println!("{:?}", e);
|
|
}
|
|
}
|
|
}
|
|
}
|
|
}
|