mirror of
https://github.com/rustdesk/rustdesk-server.git
synced 2026-08-27 04:28:08 +00:00
bind interface and refactor doc
This commit is contained in:
+116
-6
@@ -7,10 +7,33 @@ use sodiumoxide::crypto::sign;
|
||||
use std::{
|
||||
io::prelude::*,
|
||||
io::Read,
|
||||
net::SocketAddr,
|
||||
net::{IpAddr, SocketAddr},
|
||||
time::{Instant, SystemTime},
|
||||
};
|
||||
|
||||
pub fn parse_bind_address(value: &str) -> Result<Option<IpAddr>> {
|
||||
let value = value.trim();
|
||||
if value.is_empty() {
|
||||
Ok(None)
|
||||
} else {
|
||||
value
|
||||
.parse()
|
||||
.with_context(|| format!("Invalid bind address: {value}"))
|
||||
.map(Some)
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn listen_tcp(
|
||||
bind_addr: Option<IpAddr>,
|
||||
port: u16,
|
||||
) -> ResultType<hbb_common::tokio::net::TcpListener> {
|
||||
if let Some(bind_addr) = bind_addr {
|
||||
hbb_common::tcp::new_listener(SocketAddr::new(bind_addr, port), true).await
|
||||
} else {
|
||||
hbb_common::tcp::listen_any(port).await
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub(crate) fn get_expired_time() -> Instant {
|
||||
let now = Instant::now();
|
||||
@@ -52,6 +75,12 @@ fn arg_name(name: &str) -> String {
|
||||
name.to_uppercase().replace('_', "-")
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
#[inline]
|
||||
pub fn set_arg(name: &str, value: &str) {
|
||||
std::env::set_var(arg_name(name), value);
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn init_args(args: &str, name: &str, about: &str) {
|
||||
let matches = App::new(name)
|
||||
@@ -64,7 +93,7 @@ pub fn init_args(args: &str, name: &str, about: &str) {
|
||||
if let Some(section) = v.section(None::<String>) {
|
||||
section
|
||||
.iter()
|
||||
.for_each(|(k, v)| std::env::set_var(arg_name(k), v));
|
||||
.for_each(|(k, v)| set_arg(k, v));
|
||||
}
|
||||
}
|
||||
if let Some(config) = matches.value_of("config") {
|
||||
@@ -72,17 +101,42 @@ pub fn init_args(args: &str, name: &str, about: &str) {
|
||||
if let Some(section) = v.section(None::<String>) {
|
||||
section
|
||||
.iter()
|
||||
.for_each(|(k, v)| std::env::set_var(arg_name(k), v));
|
||||
.for_each(|(k, v)| set_arg(k, v));
|
||||
}
|
||||
}
|
||||
}
|
||||
for (k, v) in matches.args {
|
||||
if let Some(v) = v.vals.first() {
|
||||
std::env::set_var(arg_name(k), v.to_string_lossy().to_string());
|
||||
set_arg(k, &v.to_string_lossy());
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
pub fn get_arg_opt(name: &str) -> Option<String> {
|
||||
let dashed = arg_name(name);
|
||||
let underscored = dashed.replace('-', "_");
|
||||
let lower_dashed = dashed.to_lowercase();
|
||||
let lower_underscored = underscored.to_lowercase();
|
||||
for alias in [&dashed, &underscored, &lower_dashed, &lower_underscored] {
|
||||
if let Ok(value) = std::env::var(alias) {
|
||||
return Some(value);
|
||||
}
|
||||
}
|
||||
let mut aliases = std::env::vars_os()
|
||||
.filter_map(|(key, value)| {
|
||||
let key = key.into_string().ok()?;
|
||||
if arg_name(&key) == dashed {
|
||||
Some((key, value.into_string().ok()?))
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
aliases.sort_by(|a, b| a.0.cmp(&b.0));
|
||||
aliases.into_iter().next().map(|(_, value)| value)
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
#[inline]
|
||||
pub fn get_arg(name: &str) -> String {
|
||||
@@ -92,7 +146,7 @@ pub fn get_arg(name: &str) -> String {
|
||||
#[allow(dead_code)]
|
||||
#[inline]
|
||||
pub fn get_arg_or(name: &str, default: String) -> String {
|
||||
std::env::var(arg_name(name)).unwrap_or(default)
|
||||
get_arg_opt(name).unwrap_or(default)
|
||||
}
|
||||
|
||||
#[allow(dead_code)]
|
||||
@@ -215,4 +269,60 @@ async fn check_software_update_() -> hbb_common::ResultType<()> {
|
||||
log::info!("new version is available: {}", latest_release_version);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use std::net::{Ipv4Addr, Ipv6Addr};
|
||||
|
||||
#[test]
|
||||
fn argument_names_ignore_case_and_separator() {
|
||||
let aliases = [
|
||||
"RUSTDESK-CONFIG-ALIAS-TEST",
|
||||
"RUSTDESK_CONFIG_ALIAS_TEST",
|
||||
"rustdesk-config-alias-test",
|
||||
"rustdesk_config_alias_test",
|
||||
"RustDesk_Config-Alias_Test",
|
||||
];
|
||||
for alias in aliases {
|
||||
std::env::remove_var(alias);
|
||||
}
|
||||
for alias in aliases {
|
||||
std::env::set_var(alias, alias);
|
||||
assert_eq!(get_arg("RUSTDESK_CONFIG_ALIAS_TEST"), alias);
|
||||
std::env::remove_var(alias);
|
||||
}
|
||||
set_arg("rustdesk_config_alias_test", "normalized");
|
||||
assert_eq!(
|
||||
std::env::var("RUSTDESK-CONFIG-ALIAS-TEST").unwrap(),
|
||||
"normalized"
|
||||
);
|
||||
std::env::set_var("RUSTDESK_CONFIG_ALIAS_TEST", "inherited");
|
||||
set_arg("rustdesk-config-alias-test", "higher-priority");
|
||||
assert_eq!(get_arg("rustdesk_config_alias_test"), "higher-priority");
|
||||
std::env::remove_var("RUSTDESK-CONFIG-ALIAS-TEST");
|
||||
std::env::remove_var("RUSTDESK_CONFIG_ALIAS_TEST");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parses_bind_address() {
|
||||
assert_eq!(parse_bind_address("").unwrap(), None);
|
||||
assert_eq!(
|
||||
parse_bind_address("127.0.0.1").unwrap(),
|
||||
Some(IpAddr::V4(Ipv4Addr::LOCALHOST))
|
||||
);
|
||||
assert_eq!(
|
||||
parse_bind_address("::1").unwrap(),
|
||||
Some(IpAddr::V6(Ipv6Addr::LOCALHOST))
|
||||
);
|
||||
assert!(parse_bind_address("not-an-ip").is_err());
|
||||
}
|
||||
|
||||
#[hbb_common::tokio::test]
|
||||
async fn tcp_listener_uses_bind_address() {
|
||||
let bind_addr = IpAddr::V4(Ipv4Addr::LOCALHOST);
|
||||
let listener = listen_tcp(Some(bind_addr), 0).await.unwrap();
|
||||
assert_eq!(listener.local_addr().unwrap().ip(), bind_addr);
|
||||
}
|
||||
}
|
||||
|
||||
+1
-2
@@ -51,8 +51,7 @@ impl Database {
|
||||
if !std::path::Path::new(url).exists() {
|
||||
std::fs::File::create(url).ok();
|
||||
}
|
||||
let n: usize = std::env::var("MAX_DATABASE_CONNECTIONS")
|
||||
.unwrap_or_else(|_| "1".to_owned())
|
||||
let n: usize = crate::common::get_arg_or("MAX_DATABASE_CONNECTIONS", "1".to_owned())
|
||||
.parse()
|
||||
.unwrap_or(1);
|
||||
log::debug!("MAX_DATABASE_CONNECTIONS={}", n);
|
||||
|
||||
+16
-7
@@ -13,7 +13,8 @@ fn main() -> ResultType<()> {
|
||||
.write_mode(WriteMode::Async)
|
||||
.start()?;
|
||||
let args = format!(
|
||||
"-p, --port=[NUMBER(default={RELAY_PORT})] 'Sets the listening port'
|
||||
"-b, --bind=[IP] 'Sets the IP address to bind to (default: all interfaces)'
|
||||
-p, --port=[NUMBER(default={RELAY_PORT})] 'Sets the listening port'
|
||||
-k, --key=[KEY] 'Only allow the client with the same key'
|
||||
",
|
||||
);
|
||||
@@ -25,21 +26,29 @@ fn main() -> ResultType<()> {
|
||||
.get_matches();
|
||||
if let Ok(v) = ini::Ini::load_from_file(".env") {
|
||||
if let Some(section) = v.section(None::<String>) {
|
||||
section.iter().for_each(|(k, v)| std::env::set_var(k, v));
|
||||
section.iter().for_each(|(k, v)| common::set_arg(k, v));
|
||||
}
|
||||
}
|
||||
let mut port = RELAY_PORT;
|
||||
if let Ok(v) = std::env::var("PORT") {
|
||||
if let Some(v) = common::get_arg_opt("PORT") {
|
||||
let v: i32 = v.parse().unwrap_or_default();
|
||||
if v > 0 {
|
||||
port = v + 1;
|
||||
}
|
||||
}
|
||||
start(
|
||||
let bind = matches
|
||||
.value_of("bind")
|
||||
.map(str::to_owned)
|
||||
.unwrap_or_else(|| common::get_arg("BIND"));
|
||||
let bind_addr = common::parse_bind_address(&bind)?;
|
||||
let key = matches
|
||||
.value_of("key")
|
||||
.map(str::to_owned)
|
||||
.unwrap_or_else(|| common::get_arg("KEY"));
|
||||
start_with_bind(
|
||||
bind_addr,
|
||||
matches.value_of("port").unwrap_or(&port.to_string()),
|
||||
matches
|
||||
.value_of("key")
|
||||
.unwrap_or(&std::env::var("KEY").unwrap_or_default()),
|
||||
&key,
|
||||
)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
+13
-5
@@ -15,13 +15,14 @@ fn main() -> ResultType<()> {
|
||||
.start()?;
|
||||
let args = format!(
|
||||
"-c --config=[FILE] +takes_value 'Sets a custom config file'
|
||||
-b, --bind=[IP] 'Sets the IP address to bind to (default: all interfaces)'
|
||||
-p, --port=[NUMBER(default={RENDEZVOUS_PORT})] 'Sets the listening port'
|
||||
-s, --serial=[NUMBER(default=0)] 'Sets configure update serial number'
|
||||
-R, --rendezvous-servers=[HOSTS] 'Sets rendezvous servers, separated by comma'
|
||||
-u, --software-url=[URL] 'Sets download url of RustDesk software of newest version'
|
||||
-s, --serial=[NUMBER(default=0)] '[DEPRECATED] Sets configure update serial number'
|
||||
-R, --rendezvous-servers=[HOSTS] '[DEPRECATED] Sets rendezvous servers, separated by comma'
|
||||
-u, --software-url=[URL] '[DEPRECATED] Sets download url of RustDesk software of newest version'
|
||||
-r, --relay-servers=[HOST] 'Sets the default relay servers, separated by comma'
|
||||
-M, --rmem=[NUMBER(default={RMEM})] 'Sets UDP recv buffer size, set system rmem_max first, e.g., sudo sysctl -w net.core.rmem_max=52428800. vi /etc/sysctl.conf, net.core.rmem_max=52428800, sudo sysctl –p'
|
||||
, --mask=[MASK] 'Determine if the connection comes from LAN, e.g. 192.168.0.0/16'
|
||||
, --mask=[MASK] '[DEPRECATED] Determine if the connection comes from LAN, e.g. 192.168.0.0/16'
|
||||
-k, --key=[KEY] 'Only allow the client with the same key'",
|
||||
);
|
||||
init_args(&args, "hbbs", "RustDesk ID/Rendezvous Server");
|
||||
@@ -29,9 +30,16 @@ fn main() -> ResultType<()> {
|
||||
if port < 3 {
|
||||
bail!("Invalid port");
|
||||
}
|
||||
let bind_addr = parse_bind_address(&get_arg("bind"))?;
|
||||
let rmem = get_arg("rmem").parse::<usize>().unwrap_or(RMEM);
|
||||
let serial: i32 = get_arg("serial").parse().unwrap_or(0);
|
||||
crate::common::check_software_update();
|
||||
RendezvousServer::start(port, serial, &get_arg_or("key", "-".to_owned()), rmem)?;
|
||||
RendezvousServer::start_with_bind(
|
||||
bind_addr,
|
||||
port,
|
||||
serial,
|
||||
&get_arg_or("key", "-".to_owned()),
|
||||
rmem,
|
||||
)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
+1
-1
@@ -67,7 +67,7 @@ pub(crate) struct PeerMap {
|
||||
|
||||
impl PeerMap {
|
||||
pub(crate) async fn new() -> ResultType<Self> {
|
||||
let db = std::env::var("DB_URL").unwrap_or({
|
||||
let db = get_arg_opt("DB_URL").unwrap_or_else(|| {
|
||||
let mut db = "db_v2.sqlite3".to_owned();
|
||||
#[cfg(all(windows, not(debug_assertions)))]
|
||||
{
|
||||
|
||||
+23
-14
@@ -8,7 +8,7 @@ use hbb_common::{
|
||||
protobuf::Message as _,
|
||||
rendezvous_proto::*,
|
||||
sleep,
|
||||
tcp::{listen_any, FramedStream},
|
||||
tcp::FramedStream,
|
||||
timeout,
|
||||
tokio::{
|
||||
self,
|
||||
@@ -24,7 +24,7 @@ use std::{
|
||||
collections::{HashMap, HashSet},
|
||||
io::prelude::*,
|
||||
io::Error,
|
||||
net::SocketAddr,
|
||||
net::{IpAddr, SocketAddr},
|
||||
sync::atomic::{AtomicUsize, Ordering},
|
||||
};
|
||||
|
||||
@@ -46,7 +46,11 @@ const BLACKLIST_FILE: &str = "blacklist.txt";
|
||||
const BLOCKLIST_FILE: &str = "blocklist.txt";
|
||||
|
||||
#[tokio::main(flavor = "multi_thread")]
|
||||
pub async fn start(port: &str, key: &str) -> ResultType<()> {
|
||||
pub async fn start_with_bind(
|
||||
bind_addr: Option<IpAddr>,
|
||||
port: &str,
|
||||
key: &str,
|
||||
) -> ResultType<()> {
|
||||
let key = get_server_sk(key);
|
||||
if let Ok(mut file) = std::fs::File::open(BLACKLIST_FILE) {
|
||||
let mut contents = String::new();
|
||||
@@ -85,7 +89,12 @@ pub async fn start(port: &str, key: &str) -> ResultType<()> {
|
||||
let main_task = async move {
|
||||
loop {
|
||||
log::info!("Start");
|
||||
io_loop(listen_any(port).await?, listen_any(port2).await?, &key).await;
|
||||
io_loop(
|
||||
crate::common::listen_tcp(bind_addr, port).await?,
|
||||
crate::common::listen_tcp(bind_addr, port2).await?,
|
||||
&key,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
};
|
||||
let listen_signal = crate::common::listen_signal();
|
||||
@@ -96,8 +105,8 @@ pub async fn start(port: &str, key: &str) -> ResultType<()> {
|
||||
}
|
||||
|
||||
fn check_params() {
|
||||
let tmp = std::env::var("DOWNGRADE_THRESHOLD")
|
||||
.map(|x| x.parse::<f64>().unwrap_or(0.))
|
||||
let tmp = crate::common::get_arg("DOWNGRADE_THRESHOLD")
|
||||
.parse::<f64>()
|
||||
.unwrap_or(0.);
|
||||
if tmp > 0. {
|
||||
DOWNGRADE_THRESHOLD_100.store((tmp * 100.) as _, Ordering::SeqCst);
|
||||
@@ -106,8 +115,8 @@ fn check_params() {
|
||||
"DOWNGRADE_THRESHOLD: {}",
|
||||
DOWNGRADE_THRESHOLD_100.load(Ordering::SeqCst) as f64 / 100.
|
||||
);
|
||||
let tmp = std::env::var("DOWNGRADE_START_CHECK")
|
||||
.map(|x| x.parse::<usize>().unwrap_or(0))
|
||||
let tmp = crate::common::get_arg("DOWNGRADE_START_CHECK")
|
||||
.parse::<usize>()
|
||||
.unwrap_or(0);
|
||||
if tmp > 0 {
|
||||
DOWNGRADE_START_CHECK.store(tmp * 1000, Ordering::SeqCst);
|
||||
@@ -116,8 +125,8 @@ fn check_params() {
|
||||
"DOWNGRADE_START_CHECK: {}s",
|
||||
DOWNGRADE_START_CHECK.load(Ordering::SeqCst) / 1000
|
||||
);
|
||||
let tmp = std::env::var("LIMIT_SPEED")
|
||||
.map(|x| x.parse::<f64>().unwrap_or(0.))
|
||||
let tmp = crate::common::get_arg("LIMIT_SPEED")
|
||||
.parse::<f64>()
|
||||
.unwrap_or(0.);
|
||||
if tmp > 0. {
|
||||
LIMIT_SPEED.store((tmp * 1024. * 1024.) as usize, Ordering::SeqCst);
|
||||
@@ -126,8 +135,8 @@ fn check_params() {
|
||||
"LIMIT_SPEED: {}Mb/s",
|
||||
LIMIT_SPEED.load(Ordering::SeqCst) as f64 / 1024. / 1024.
|
||||
);
|
||||
let tmp = std::env::var("TOTAL_BANDWIDTH")
|
||||
.map(|x| x.parse::<f64>().unwrap_or(0.))
|
||||
let tmp = crate::common::get_arg("TOTAL_BANDWIDTH")
|
||||
.parse::<f64>()
|
||||
.unwrap_or(0.);
|
||||
if tmp > 0. {
|
||||
TOTAL_BANDWIDTH.store((tmp * 1024. * 1024.) as usize, Ordering::SeqCst);
|
||||
@@ -137,8 +146,8 @@ fn check_params() {
|
||||
"TOTAL_BANDWIDTH: {}Mb/s",
|
||||
TOTAL_BANDWIDTH.load(Ordering::SeqCst) as f64 / 1024. / 1024.
|
||||
);
|
||||
let tmp = std::env::var("SINGLE_BANDWIDTH")
|
||||
.map(|x| x.parse::<f64>().unwrap_or(0.))
|
||||
let tmp = crate::common::get_arg("SINGLE_BANDWIDTH")
|
||||
.parse::<f64>()
|
||||
.unwrap_or(0.);
|
||||
if tmp > 0. {
|
||||
SINGLE_BANDWIDTH.store((tmp * 1024. * 1024.) as usize, Ordering::SeqCst);
|
||||
|
||||
+51
-22
@@ -16,7 +16,7 @@ use hbb_common::{
|
||||
register_pk_response::Result::{TOO_FREQUENT, UUID_MISMATCH},
|
||||
*,
|
||||
},
|
||||
tcp::{listen_any, FramedStream},
|
||||
tcp::FramedStream,
|
||||
timeout,
|
||||
tokio::{
|
||||
self,
|
||||
@@ -98,18 +98,25 @@ enum LoopFailure {
|
||||
}
|
||||
|
||||
impl RendezvousServer {
|
||||
pub fn start(port: i32, serial: i32, key: &str, rmem: usize) -> ResultType<()> {
|
||||
Self::start_with_bind(None, port, serial, key, rmem)
|
||||
}
|
||||
|
||||
#[tokio::main(flavor = "multi_thread")]
|
||||
pub async fn start(port: i32, serial: i32, key: &str, rmem: usize) -> ResultType<()> {
|
||||
pub async fn start_with_bind(
|
||||
bind_addr: Option<IpAddr>,
|
||||
port: i32,
|
||||
serial: i32,
|
||||
key: &str,
|
||||
rmem: usize,
|
||||
) -> ResultType<()> {
|
||||
let (key, sk) = Self::get_server_sk(key);
|
||||
let nat_port = port - 1;
|
||||
let ws_port = port + 2;
|
||||
let pm = PeerMap::new().await?;
|
||||
log::info!("serial={}", serial);
|
||||
let rendezvous_servers = get_servers(&get_arg("rendezvous-servers"), "rendezvous-servers");
|
||||
log::info!("Listening on tcp/udp :{}", port);
|
||||
log::info!("Listening on tcp :{}, extra port for NAT test", nat_port);
|
||||
log::info!("Listening on websocket :{}", ws_port);
|
||||
let mut socket = create_udp_listener(port, rmem).await?;
|
||||
let mut socket = create_udp_listener(bind_addr, port, rmem).await?;
|
||||
let (tx, mut rx) = mpsc::unbounded_channel::<Data>();
|
||||
let software_url = get_arg("software-url");
|
||||
let version = hbb_common::get_version_from_url(&software_url);
|
||||
@@ -147,15 +154,17 @@ impl RendezvousServer {
|
||||
log::info!("local-ip: {:?}", rs.inner.local_ip);
|
||||
std::env::set_var("PORT_FOR_API", port.to_string());
|
||||
rs.parse_relay_servers(&get_arg("relay-servers"));
|
||||
let mut listener = create_tcp_listener(port).await?;
|
||||
let mut listener2 = create_tcp_listener(nat_port).await?;
|
||||
let mut listener3 = create_tcp_listener(ws_port).await?;
|
||||
let test_addr = std::env::var("TEST_HBBS").unwrap_or_default();
|
||||
if std::env::var("ALWAYS_USE_RELAY")
|
||||
.unwrap_or_default()
|
||||
.to_uppercase()
|
||||
== "Y"
|
||||
{
|
||||
let mut listener = create_tcp_listener(bind_addr, port).await?;
|
||||
let mut listener2 = create_tcp_listener(bind_addr, nat_port).await?;
|
||||
let mut listener3 = create_tcp_listener(bind_addr, ws_port).await?;
|
||||
log::info!("Listening on tcp/udp {}", listener.local_addr()?);
|
||||
log::info!(
|
||||
"Listening on tcp {}, extra port for NAT test",
|
||||
listener2.local_addr()?
|
||||
);
|
||||
log::info!("Listening on websocket {}", listener3.local_addr()?);
|
||||
let test_addr = get_arg("TEST_HBBS");
|
||||
if get_arg("ALWAYS_USE_RELAY").to_uppercase() == "Y" {
|
||||
ALWAYS_USE_RELAY.store(true, Ordering::SeqCst);
|
||||
}
|
||||
log::info!(
|
||||
@@ -204,19 +213,19 @@ impl RendezvousServer {
|
||||
{
|
||||
LoopFailure::UdpSocket => {
|
||||
drop(socket);
|
||||
socket = create_udp_listener(port, rmem).await?;
|
||||
socket = create_udp_listener(bind_addr, port, rmem).await?;
|
||||
}
|
||||
LoopFailure::Listener => {
|
||||
drop(listener);
|
||||
listener = create_tcp_listener(port).await?;
|
||||
listener = create_tcp_listener(bind_addr, port).await?;
|
||||
}
|
||||
LoopFailure::Listener2 => {
|
||||
drop(listener2);
|
||||
listener2 = create_tcp_listener(nat_port).await?;
|
||||
listener2 = create_tcp_listener(bind_addr, nat_port).await?;
|
||||
}
|
||||
LoopFailure::Listener3 => {
|
||||
drop(listener3);
|
||||
listener3 = create_tcp_listener(ws_port).await?;
|
||||
listener3 = create_tcp_listener(bind_addr, ws_port).await?;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -1352,7 +1361,15 @@ async fn send_rk_res(
|
||||
socket.send(&msg_out, addr).await
|
||||
}
|
||||
|
||||
async fn create_udp_listener(port: i32, rmem: usize) -> ResultType<FramedSocket> {
|
||||
async fn create_udp_listener(
|
||||
bind_addr: Option<IpAddr>,
|
||||
port: i32,
|
||||
rmem: usize,
|
||||
) -> ResultType<FramedSocket> {
|
||||
if let Some(bind_addr) = bind_addr {
|
||||
let addr = SocketAddr::new(bind_addr, port as _);
|
||||
return FramedSocket::new_reuse(&addr, true, rmem).await;
|
||||
}
|
||||
let addr = SocketAddr::new(IpAddr::V6(Ipv6Addr::UNSPECIFIED), port as _);
|
||||
if let Ok(s) = FramedSocket::new_reuse(&addr, true, rmem).await {
|
||||
log::debug!("listen on udp {:?}", s.local_addr());
|
||||
@@ -1365,8 +1382,20 @@ async fn create_udp_listener(port: i32, rmem: usize) -> ResultType<FramedSocket>
|
||||
}
|
||||
|
||||
#[inline]
|
||||
async fn create_tcp_listener(port: i32) -> ResultType<TcpListener> {
|
||||
let s = listen_any(port as _).await?;
|
||||
async fn create_tcp_listener(bind_addr: Option<IpAddr>, port: i32) -> ResultType<TcpListener> {
|
||||
let s = listen_tcp(bind_addr, port as _).await?;
|
||||
log::debug!("listen on tcp {:?}", s.local_addr());
|
||||
Ok(s)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[hbb_common::tokio::test]
|
||||
async fn udp_listener_uses_bind_address() {
|
||||
let bind_addr = IpAddr::V4(Ipv4Addr::LOCALHOST);
|
||||
let socket = create_udp_listener(Some(bind_addr), 0, 0).await.unwrap();
|
||||
assert_eq!(socket.local_addr().unwrap().ip(), bind_addr);
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user