diff --git a/Cargo.lock b/Cargo.lock index 9acff23d..6bddcaf3 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -483,6 +483,7 @@ version = "0.8.0" dependencies = [ "async-trait", "futures", + "hickory-proto", "hickory-server", "thiserror 2.0.18", "tokio", @@ -584,6 +585,7 @@ dependencies = [ "h3", "h3-quinn", "hickory-proto", + "hickory-resolver", "http", "http-body-util", "hyper", @@ -591,6 +593,7 @@ dependencies = [ "ipnet", "libc", "log", + "lru_time_cache", "maxminddb", "memchr", "memory-stats", @@ -607,7 +610,7 @@ dependencies = [ "serde_yaml", "sha2", "smoltcp", - "socket2", + "socket2 0.6.2", "subtle", "thiserror 2.0.18", "time", @@ -626,7 +629,7 @@ dependencies = [ "url", "uuid", "watfaq-netstack", - "webpki-roots", + "webpki-roots 1.0.5", "windows 0.62.2", "x509-parser 0.16.0", ] @@ -720,6 +723,15 @@ dependencies = [ "crossbeam-utils", ] +[[package]] +name = "crossbeam-epoch" +version = "0.9.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5b82ac4a3c2ca9c3460964f020e1402edd5753411d7737aa39c3714ad1b5420e" +dependencies = [ + "crossbeam-utils", +] + [[package]] name = "crossbeam-utils" version = "0.8.21" @@ -1410,23 +1422,53 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8a6fe56c0038198998a6f217ca4e7ef3a5e51f46163bd6dd60b5c71ca6c6502" dependencies = [ "async-trait", + "bytes", "cfg-if", "data-encoding", "enum-as-inner", "futures-channel", "futures-io", "futures-util", + "h2", + "http", "idna", "ipnet", "once_cell", "rand 0.9.2", "ring", + "rustls", "serde", "thiserror 2.0.18", "tinyvec", "tokio", + "tokio-rustls", "tracing", "url", + "webpki-roots 0.26.11", +] + +[[package]] +name = "hickory-resolver" +version = "0.25.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc62a9a99b0bfb44d2ab95a7208ac952d31060efc16241c87eaf36406fecf87a" +dependencies = [ + "cfg-if", + "futures-util", + "hickory-proto", + "ipconfig", + "moka", + "once_cell", + "parking_lot", + "rand 0.9.2", + "resolv-conf", + "rustls", + "smallvec", + "thiserror 2.0.18", + "tokio", + "tokio-rustls", + "tracing", + "webpki-roots 0.26.11", ] [[package]] @@ -1543,7 +1585,7 @@ dependencies = [ "libc", "percent-encoding", "pin-project-lite", - "socket2", + "socket2 0.6.2", "system-configuration", "tokio", "tower-layer", @@ -1696,6 +1738,18 @@ dependencies = [ "serde_core", ] +[[package]] +name = "ipconfig" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b58db92f96b720de98181bbbe63c831e87005ab460c1bf306eb2622b4707997f" +dependencies = [ + "socket2 0.5.10", + "widestring", + "windows-sys 0.48.0", + "winreg 0.50.0", +] + [[package]] name = "ipnet" version = "2.11.0" @@ -1841,6 +1895,12 @@ version = "0.1.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "112b39cec0b298b6c1999fee3e31427f74f676e4cb9879ed1a121b43661a4154" +[[package]] +name = "lru_time_cache" +version = "0.11.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9106e1d747ffd48e6be5bb2d97fa706ed25b144fbee4d5c02eae110cd8d6badd" + [[package]] name = "managed" version = "0.8.0" @@ -1933,6 +1993,23 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "moka" +version = "0.12.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85f8024e1c8e71c778968af91d43700ce1d11b219d127d79fb2934153b82b42b" +dependencies = [ + "crossbeam-channel", + "crossbeam-epoch", + "crossbeam-utils", + "equivalent", + "parking_lot", + "portable-atomic", + "smallvec", + "tagptr", + "uuid", +] + [[package]] name = "netconfig-rs" version = "0.1.5" @@ -2407,7 +2484,7 @@ dependencies = [ "quinn-udp", "rustc-hash 2.1.1", "rustls", - "socket2", + "socket2 0.6.2", "thiserror 2.0.18", "tokio", "tracing", @@ -2445,7 +2522,7 @@ dependencies = [ "cfg_aliases", "libc", "once_cell", - "socket2", + "socket2 0.6.2", "tracing", "windows-sys 0.60.2", ] @@ -2593,6 +2670,12 @@ version = "0.8.8" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7a2d987857b319362043e95f5353c0535c1f58eec5336fdfcf626430af7def58" +[[package]] +name = "resolv-conf" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1e061d1b48cb8d38042de4ae0a7a6401009d6143dc80d2e2d6f31f0bdd6470c7" + [[package]] name = "ring" version = "0.17.14" @@ -2906,6 +2989,16 @@ dependencies = [ "managed", ] +[[package]] +name = "socket2" +version = "0.5.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e22376abed350d73dd1cd119b57ffccad95b4e585a7cda43e286245ce23c0678" +dependencies = [ + "libc", + "windows-sys 0.52.0", +] + [[package]] name = "socket2" version = "0.6.2" @@ -3012,6 +3105,12 @@ dependencies = [ "libc", ] +[[package]] +name = "tagptr" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417" + [[package]] name = "thiserror" version = "1.0.69" @@ -3151,7 +3250,7 @@ dependencies = [ "parking_lot", "pin-project-lite", "signal-hook-registry", - "socket2", + "socket2 0.6.2", "tokio-macros", "tracing", "windows-sys 0.61.2", @@ -3374,7 +3473,7 @@ dependencies = [ "tokio", "widestring", "windows-sys 0.61.2", - "winreg", + "winreg 0.55.0", ] [[package]] @@ -3619,7 +3718,7 @@ dependencies = [ "netstack-smoltcp", "rand 0.10.0", "smoltcp", - "socket2", + "socket2 0.6.2", "tokio", "tracing", "tracing-subscriber", @@ -3636,6 +3735,15 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "webpki-roots" +version = "0.26.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9" +dependencies = [ + "webpki-roots 1.0.5", +] + [[package]] name = "webpki-roots" version = "1.0.5" @@ -3865,6 +3973,15 @@ dependencies = [ "windows-link 0.2.1", ] +[[package]] +name = "windows-sys" +version = "0.48.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "677d2418bec65e3338edb076e806bc1ec15693c5d0104683f2efe857f61056a9" +dependencies = [ + "windows-targets 0.48.5", +] + [[package]] name = "windows-sys" version = "0.52.0" @@ -3901,6 +4018,21 @@ dependencies = [ "windows-link 0.2.1", ] +[[package]] +name = "windows-targets" +version = "0.48.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a2fa6e2155d7247be68c096456083145c183cbbbc2764150dda45a87197940c" +dependencies = [ + "windows_aarch64_gnullvm 0.48.5", + "windows_aarch64_msvc 0.48.5", + "windows_i686_gnu 0.48.5", + "windows_i686_msvc 0.48.5", + "windows_x86_64_gnu 0.48.5", + "windows_x86_64_gnullvm 0.48.5", + "windows_x86_64_msvc 0.48.5", +] + [[package]] name = "windows-targets" version = "0.52.6" @@ -3952,6 +4084,12 @@ dependencies = [ "windows-link 0.2.1", ] +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.48.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2b38e32f0abccf9987a4e3079dfb67dcd799fb61361e53e2882c3cbaf0d905d8" + [[package]] name = "windows_aarch64_gnullvm" version = "0.52.6" @@ -3964,6 +4102,12 @@ version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a9d8416fa8b42f5c947f8482c43e7d89e73a173cead56d044f6a56104a6d1b53" +[[package]] +name = "windows_aarch64_msvc" +version = "0.48.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc35310971f3b2dbbf3f0690a219f40e2d9afcf64f9ab7cc1be722937c26b4bc" + [[package]] name = "windows_aarch64_msvc" version = "0.52.6" @@ -3976,6 +4120,12 @@ version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b9d782e804c2f632e395708e99a94275910eb9100b2114651e04744e9b125006" +[[package]] +name = "windows_i686_gnu" +version = "0.48.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a75915e7def60c94dcef72200b9a8e58e5091744960da64ec734a6c6e9b3743e" + [[package]] name = "windows_i686_gnu" version = "0.52.6" @@ -4000,6 +4150,12 @@ version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fa7359d10048f68ab8b09fa71c3daccfb0e9b559aed648a8f95469c27057180c" +[[package]] +name = "windows_i686_msvc" +version = "0.48.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f55c233f70c4b27f66c523580f78f1004e8b5a8b659e05a4eb49d4166cca406" + [[package]] name = "windows_i686_msvc" version = "0.52.6" @@ -4012,6 +4168,12 @@ version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1e7ac75179f18232fe9c285163565a57ef8d3c89254a30685b57d83a38d326c2" +[[package]] +name = "windows_x86_64_gnu" +version = "0.48.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "53d40abd2583d23e4718fddf1ebec84dbff8381c07cae67ff7768bbf19c6718e" + [[package]] name = "windows_x86_64_gnu" version = "0.52.6" @@ -4024,6 +4186,12 @@ version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9c3842cdd74a865a8066ab39c8a7a473c0778a3f29370b5fd6b4b9aa7df4a499" +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.48.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b7b52767868a23d5bab768e390dc5f5c55825b6d30b86c844ff2dc7414044cc" + [[package]] name = "windows_x86_64_gnullvm" version = "0.52.6" @@ -4036,6 +4204,12 @@ version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0ffa179e2d07eee8ad8f57493436566c7cc30ac536a3379fdf008f47f6bb7ae1" +[[package]] +name = "windows_x86_64_msvc" +version = "0.48.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed94fce61571a4006852b7389a063ab983c02eb1bb37b47f8272ce92d06d9538" + [[package]] name = "windows_x86_64_msvc" version = "0.52.6" @@ -4048,6 +4222,16 @@ version = "0.53.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d6bbff5f0aada427a1e5a6da5f1f98158182f26556f345ac9e04d36d0ebed650" +[[package]] +name = "winreg" +version = "0.50.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "524e57b2c537c0f9b1e69f1965311ec12182b4122e45035b1508cd24d2adadb1" +dependencies = [ + "cfg-if", + "windows-sys 0.48.0", +] + [[package]] name = "winreg" version = "0.55.0" diff --git a/clash-dns/Cargo.toml b/clash-dns/Cargo.toml index 542b753d..40b60ee2 100644 --- a/clash-dns/Cargo.toml +++ b/clash-dns/Cargo.toml @@ -13,5 +13,6 @@ futures = "0.3" async-trait = "0.1" hickory-server = { version = "0.25", default-features = false } +hickory-proto = "0.25" tokio = { version = "1", features = ["full"] } tracing = "0.1" diff --git a/clash-dns/src/handler.rs b/clash-dns/src/handler.rs index 8bff362f..82940e0e 100644 --- a/clash-dns/src/handler.rs +++ b/clash-dns/src/handler.rs @@ -3,9 +3,11 @@ use std::time::Duration; use crate::utils::new_io_error; use crate::{DNSListenAddr, DnsMessageExchanger}; use async_trait::async_trait; +use hickory_proto::op::{Header, Message, ResponseCode}; use hickory_server::server::Request; use hickory_server::{ ServerFuture, + authority::MessageResponseBuilder, server::{RequestHandler, ResponseHandler, ResponseInfo}, }; use thiserror::Error; @@ -24,10 +26,10 @@ struct DnsHandler { pub enum DNSError { #[error(transparent)] Io(#[from] std::io::Error), - /* #[error("invalid OP code: {0}")] + #[error("invalid OP query: {0}")] InvalidOpQuery(String), #[error("query failed: {0}")] - QueryFailed(String), */ + QueryFailed(String), } #[async_trait] @@ -38,9 +40,44 @@ where async fn handle_request( &self, request: &Request, - response_handle: H, + mut response_handle: H, ) -> ResponseInfo { - todo!() + let req = match to_dns_message(request) { + Ok(req) => req, + Err(err) => { + error!("failed to parse dns request: {}", err); + return servfail_info(); + } + }; + + let resp = match self.exchanger.exchange(&req).await { + Ok(resp) => resp, + Err(err) => { + warn!("dns exchange failed: {}", err); + build_servfail_message(&req) + } + }; + + let mut builder = MessageResponseBuilder::from_message_request(request); + if let Some(edns) = resp.extensions().clone() { + builder.edns(edns); + } + + let response = builder.build( + resp.header().clone(), + resp.answers(), + resp.name_servers(), + std::iter::empty::<&hickory_proto::rr::Record>(), + resp.additionals(), + ); + + match response_handle.send_response(response).await { + Ok(info) => info, + Err(err) => { + error!("failed to send dns response: {}", err); + servfail_info() + } + } } } @@ -84,15 +121,18 @@ where .is_ok(); } if let Some(c) = listen.doh { - todo!() + let _ = c; + warn!("DoH listener is not implemented yet"); } if let Some(c) = listen.dot { - todo!() + let _ = c; + warn!("DoT listener is not implemented yet"); } if let Some(c) = listen.doh3 { - todo!() + let _ = c; + warn!("DoH3 listener is not implemented yet"); } if !has_server { @@ -109,3 +149,53 @@ where }) })) } + +fn to_dns_message(request: &Request) -> Result { + let mut message = Message::new(); + message.set_id(request.id()); + message.set_op_code(request.op_code()); + message.set_message_type(request.message_type()); + message.set_authoritative(request.authoritative()); + message.set_truncated(request.truncated()); + message.set_recursion_desired(request.recursion_desired()); + message.set_recursion_available(request.recursion_available()); + message.set_authentic_data(request.authentic_data()); + message.set_checking_disabled(request.checking_disabled()); + message.set_response_code(request.response_code()); + message.add_queries(request.queries().iter().map(|q| q.original().clone())); + message.add_answers(request.answers().iter().cloned()); + message.add_name_servers(request.name_servers().iter().cloned()); + message.add_additionals(request.additionals().iter().cloned()); + if let Some(edns) = request.edns().cloned() { + message.set_edns(edns); + } + Ok(message) +} + +fn build_servfail_message(req: &Message) -> Message { + let mut header = Header::response_from_request(req.header()); + header.set_response_code(ResponseCode::ServFail); + + let mut message = Message::new(); + message.set_id(header.id()); + message.set_message_type(header.message_type()); + message.set_op_code(header.op_code()); + message.set_authoritative(header.authoritative()); + message.set_truncated(header.truncated()); + message.set_recursion_desired(header.recursion_desired()); + message.set_recursion_available(header.recursion_available()); + message.set_authentic_data(header.authentic_data()); + message.set_checking_disabled(header.checking_disabled()); + message.set_response_code(header.response_code()); + message.add_queries(req.queries().iter().cloned()); + if let Some(edns) = req.extensions().clone() { + message.set_edns(edns); + } + message +} + +fn servfail_info() -> ResponseInfo { + let mut header = Header::new(); + header.set_response_code(ResponseCode::ServFail); + header.into() +} diff --git a/clash-dns/src/lib.rs b/clash-dns/src/lib.rs index b3834547..24af7042 100644 --- a/clash-dns/src/lib.rs +++ b/clash-dns/src/lib.rs @@ -1,5 +1,8 @@ use std::{net::SocketAddr, path::Path}; +use async_trait::async_trait; +use hickory_proto::op::Message; + mod handler; mod utils; @@ -26,7 +29,8 @@ pub struct DNSListenAddr { pub doh3: Option, } +#[async_trait] pub trait DnsMessageExchanger: Send + Sync { fn ipv6(&self) -> bool; - // async fn exchange(&self, message: &Message) -> Result; + async fn exchange(&self, message: &Message) -> Result; } diff --git a/clash-lib/Cargo.toml b/clash-lib/Cargo.toml index a00cafcf..bea71ecb 100644 --- a/clash-lib/Cargo.toml +++ b/clash-lib/Cargo.toml @@ -37,8 +37,10 @@ h3-quinn = { version = "0.0.10", optional = true } quinn-proto = { version = "0.11.13", default-features = false, optional = true } maxminddb = "0.27" hickory-proto = "0.25" -url = { version = "2", optional = true } +hickory-resolver = { version = "0.25", features = ["tokio", "system-config", "webpki-roots", "tls-aws-lc-rs", "https-aws-lc-rs"] } +url = { version = "2" } ipnet = { version = "2" } +lru_time_cache = "0.11" network-interface = { version = "2", optional = true } serde = { version = "1", features = ["derive"] } @@ -154,7 +156,7 @@ tun = [ "dep:watfaq-netstack", "dep:smoltcp", "dep:network-interface", - "dep:url", + # "dep:url", ] tproxy = ["dep:etherparse"] redir = [] diff --git a/clash-lib/src/app/api/handlers/config.rs b/clash-lib/src/app/api/handlers/config.rs index d2c18b42..717af98e 100644 --- a/clash-lib/src/app/api/handlers/config.rs +++ b/clash-lib/src/app/api/handlers/config.rs @@ -17,8 +17,12 @@ use crate::{ dispatcher::Dispatcher, dns::ThreadSafeDNSResolver, inbound::manager::{InboundManager, Ports}, + logging, + }, + config::{ + def::{self}, + internal::config::BindAddress, }, - config::def::{self, LogLevel, RunMode}, }; #[derive(Clone)] @@ -36,7 +40,10 @@ pub fn routes( dns_resolver: ThreadSafeDNSResolver, ) -> Router> { Router::new() - .route("/", get(get_configs).put(update_configs)) + .route( + "/", + get(get_configs).put(update_configs).patch(patch_configs), + ) .with_state(ConfigState { inbound_manager, dispatcher, @@ -45,7 +52,7 @@ pub fn routes( }) } -#[derive(Serialize)] +#[derive(Serialize, Deserialize)] #[serde(rename_all = "kebab-case")] struct PatchConfigRequest { port: Option, @@ -60,6 +67,26 @@ struct PatchConfigRequest { allow_lan: Option, } +impl PatchConfigRequest { + fn rebuild_listeners(&self) -> bool { + self.port.is_some() + || self.socks_port.is_some() + || self.redir_port.is_some() + || self.tproxy_port.is_some() + || self.mixed_port.is_some() + || self.bind_address.is_some() + } +} + +fn parse_bind_address(value: &str) -> Result { + match value { + "*" => Ok(BindAddress::all_v4()), + "localhost" => Ok(BindAddress::local()), + "[::]" | "::" => Ok(BindAddress::dual_stack()), + _ => value.parse().map(BindAddress).map_err(|_| ()), + } +} + async fn get_configs(State(state): State) -> impl IntoResponse { let run_mode = state.dispatcher.get_mode().await; let global_state = state.global_state.lock().await; @@ -156,3 +183,73 @@ async fn update_configs( } } } + +async fn patch_configs( + State(state): State, + Json(payload): Json, +) -> impl IntoResponse { + let inbound_manager = state.inbound_manager.clone(); + let mut need_restart = false; + + if let Some(bind_address) = payload.bind_address.clone() { + match parse_bind_address(&bind_address) { + Ok(bind_address) => { + inbound_manager.set_bind_address(bind_address).await; + need_restart = true; + } + Err(_) => { + return ( + StatusCode::BAD_REQUEST, + format!("invalid bind address: {bind_address}"), + ) + .into_response(); + } + } + } + + let mut global_state = state.global_state.lock().await; + + if payload.rebuild_listeners() { + let ports = Ports { + port: payload.port, + socks_port: payload.socks_port, + redir_port: payload.redir_port, + tproxy_port: payload.tproxy_port, + mixed_port: payload.mixed_port, + }; + inbound_manager.change_ports(ports).await; + need_restart = true; + } + + if let Some(allow_lan) = payload.allow_lan + && allow_lan != inbound_manager.get_allow_lan().await + { + inbound_manager.set_allow_lan(allow_lan).await; + need_restart = true; + } + + if need_restart { + inbound_manager.restart().await; + } + + if let Some(mode) = payload.mode { + state.dispatcher.set_mode(mode).await; + } + + if let Some(log_level) = payload.log_level { + global_state.log_level = log_level; + if let Err(err) = logging::set_log_level(log_level) { + return ( + StatusCode::INTERNAL_SERVER_ERROR, + format!("failed to update log level: {err}"), + ) + .into_response(); + } + } + + if let Some(ipv6) = payload.ipv6 { + state.dns_resolver.set_ipv6(ipv6); + } + + StatusCode::ACCEPTED.into_response() +} diff --git a/clash-lib/src/app/api/handlers/traffic.rs b/clash-lib/src/app/api/handlers/traffic.rs index 8278d122..f2d867ef 100644 --- a/clash-lib/src/app/api/handlers/traffic.rs +++ b/clash-lib/src/app/api/handlers/traffic.rs @@ -1,7 +1,9 @@ -use std::{net::SocketAddr, sync::Arc}; +use std::sync::Arc; use axum::{ - extract::{ConnectInfo, State, WebSocketUpgrade, ws::Message}, + body::Body, + extract::{FromRequest, Request, State, WebSocketUpgrade, ws::Message}, + http::StatusCode, response::IntoResponse, }; use serde::Serialize; @@ -16,12 +18,22 @@ struct TrafficResponse { } pub async fn handle( - ws: WebSocketUpgrade, - ConnectInfo(addr): ConnectInfo, State(state): State>, + req: Request, ) -> impl IntoResponse { + let ws = match WebSocketUpgrade::from_request(req, &state).await { + Ok(ws) => ws, + Err(_) => { + return ( + StatusCode::BAD_REQUEST, + "the /traffic endpoint requires websocket upgrade", + ) + .into_response(); + } + }; + ws.on_failed_upgrade(move |e| { - warn!("ws upgrade error: {} with {}", e, addr); + warn!("ws upgrade error: {}", e); }) .on_upgrade(move |mut socket| async move { let mgr = state.statistics_manager.clone(); diff --git a/clash-lib/src/app/api/ipc.rs b/clash-lib/src/app/api/ipc.rs new file mode 100644 index 00000000..944b1412 --- /dev/null +++ b/clash-lib/src/app/api/ipc.rs @@ -0,0 +1,109 @@ +use tracing::error; + +#[cfg(windows)] +pub async fn serve_ipc(router: axum::Router, path: &str) -> crate::Result<()> { + use hyper_util::rt::TokioIo; + use tokio::net::windows::named_pipe; + use tower::Service as _; + use tracing::info; + + info!("Starting API server on NamedPipe {path}"); + + let server = named_pipe::ServerOptions::new() + .first_pipe_instance(true) + .create(path) + .map_err(|e| crate::Error::Operation(format!("Cannot create pipe {e}")))?; + + let mut server = server; + loop { + server + .connect() + .await + .map_err(|e| crate::Error::Operation(format!("NamedPipe error: {e}")))?; + let connected_client = server; + server = named_pipe::ServerOptions::new().create(path).map_err(|e| { + crate::Error::Operation(format!("Cannot create NamedPipe: {e}")) + })?; + let router = router.clone(); + tokio::spawn(async move { + let io = TokioIo::new(connected_client); + let hyper_service = hyper::service::service_fn(move |request: _| { + router.clone().call(request) + }); + + if let Err(e) = hyper::server::conn::http1::Builder::new() + .serve_connection(io, hyper_service) + .await + { + error!("NamedPipe error: {}", e); + } + }); + } +} + +#[cfg(unix)] +pub async fn serve_ipc(router: axum::Router, path: &str) -> crate::Result<()> { + use std::{path::PathBuf, sync::Arc}; + + use axum::{extract::connect_info::Connected, serve::IncomingStream}; + use tokio::net::UnixListener; + use tracing::info; + + let path = PathBuf::from(path); + + info!("Start API server on IPC address {:?}", path); + + if let Err(e) = tokio::fs::remove_file(&path).await + && e.kind() != std::io::ErrorKind::NotFound + { + return Err(crate::Error::Operation(format!( + "Cannot remove existing IPC file: {e}", + ))); + } + + if let Some(parent) = path.parent() { + tokio::fs::create_dir_all(parent).await.map_err(|e| { + crate::Error::Operation(format!("Cannot create IPC dir: {e}")) + })?; + } + + let uds = UnixListener::bind(&path).map_err(|e| { + crate::Error::Operation(format!("Cannot bind on IPC address: {e}")) + })?; + + #[derive(Clone, Debug)] + #[allow(dead_code)] + struct UdsConnectInfo { + peer_addr: Arc, + peer_cred: tokio::net::unix::UCred, + } + + impl Connected> for UdsConnectInfo { + fn connect_info(stream: IncomingStream<'_, UnixListener>) -> Self { + let peer_addr = stream.io().peer_addr().unwrap(); + let peer_cred = stream.io().peer_cred().unwrap(); + Self { + peer_addr: Arc::new(peer_addr), + peer_cred, + } + } + } + + axum::serve( + uds, + router.into_make_service_with_connect_info::(), + ) + .await + .map_err(|e| { + error!("IPC API server error: {}", e); + crate::Error::Operation(format!("IPC API server error: {e}")) + }) +} + +#[cfg(all(not(unix), not(windows)))] +pub async fn serve_ipc(_router: axum::Router, _path: &str) -> crate::Result<()> { + error!("IPC only get supported on Unix and Windows"); + Err(crate::Error::Operation( + "IPC only get supported on Unix and Windows".to_string(), + )) +} diff --git a/clash-lib/src/app/api/mod.rs b/clash-lib/src/app/api/mod.rs index ffa4f947..8e88fdef 100644 --- a/clash-lib/src/app/api/mod.rs +++ b/clash-lib/src/app/api/mod.rs @@ -27,6 +27,7 @@ use crate::{ }; mod handlers; +mod ipc; pub struct AppState { log_source_tx: Sender, @@ -47,8 +48,11 @@ pub fn get_api_runner( _cwd: String, ) -> Option { tracing::debug!("API controller configuration: {:?}", controller_cfg); - let tcp_addr = controller_cfg.external_controller; - let ipc_addr = controller_cfg.external_controller_ipc; + let tcp_addr = controller_cfg + .external_controller + .clone() + .filter(|value| !value.is_empty()); + let ipc_addr = controller_cfg.external_controller_ipc.clone(); if tcp_addr.is_none() && ipc_addr.is_none() { return None; @@ -176,13 +180,19 @@ pub fn get_api_runner( None }; - if ipc_addr.is_some() { - warn!("IPC API listeners are not wired yet"); - } + let ipc_fut = ipc_addr + .map(|ipc_path| async move { ipc::serve_ipc(router, &ipc_path).await }); - match tcp_fut { - Some(tcp) => tcp.await, - None => Ok(()), + match (tcp_fut, ipc_fut) { + (Some(tcp), Some(ipc)) => { + tokio::select! { + result = tcp => result, + result = ipc => result, + } + } + (Some(tcp), None) => tcp.await, + (None, Some(ipc)) => ipc.await, + (None, None) => Ok(()), } }; diff --git a/clash-lib/src/app/dispatcher/dispatcher_impl.rs b/clash-lib/src/app/dispatcher/dispatcher_impl.rs index 25f76ab8..1dabf4e5 100644 --- a/clash-lib/src/app/dispatcher/dispatcher_impl.rs +++ b/clash-lib/src/app/dispatcher/dispatcher_impl.rs @@ -1,7 +1,7 @@ use std::{fmt, sync::Arc, time::Duration}; use tokio::{io::AsyncWriteExt, sync::RwLock}; -use tracing::{Instrument, debug, info_span, instrument, trace, warn}; +use tracing::{Instrument, debug, error, info_span, instrument, trace, warn}; use tracing_log::log; use crate::{ @@ -228,7 +228,31 @@ async fn reverse_lookup( ) -> Option { let dst = match dst { crate::session::SocksAddr::Ip(socket_addr) => { - todo!() + if resolver.fake_ip_enabled() { + let ip = socket_addr.ip(); + if resolver.is_fake_ip(ip).await { + trace!("looking up fake ip: {}", socket_addr.ip()); + match resolver.reverse_lookup(ip).await { + Some(host) => (host, socket_addr.port()) + .try_into() + .expect("must be valid domain"), + None => { + error!("failed to reverse lookup fake ip: {}", ip); + return None; + } + } + } else { + (*socket_addr).into() + } + } else { + trace!("looking up resolve cache ip: {}", socket_addr.ip()); + match resolver.cached_for(socket_addr.ip()).await { + Some(host) => (host, socket_addr.port()) + .try_into() + .expect("must be valid domain"), + None => (*socket_addr).into(), + } + } } crate::session::SocksAddr::Domain(host, port) => (host.to_owned(), *port) .try_into() diff --git a/clash-lib/src/app/dns/config.rs b/clash-lib/src/app/dns/config.rs index d41b0560..05b311c2 100644 --- a/clash-lib/src/app/dns/config.rs +++ b/clash-lib/src/app/dns/config.rs @@ -1,16 +1,243 @@ -use std::net::SocketAddr; +pub use super::dns_client::DNSNetMode; +use std::collections::HashMap; +use std::net::{IpAddr, SocketAddr}; use chimera_dns::DNSListenAddr; +use ipnet::{IpNet, Ipv4Net, Ipv6Net}; +use std::fmt::Display; +use url::Url; -use crate::{Error, config::def::DNSListen}; +use crate::{ + Error, + config::def::{ + DNSListen, DNSMode, EdnsClientSubnet as DefEdnsClientSubnet, + FallbackFilter as DefFallbackFilter, + }, +}; + +#[derive(Clone, Debug)] +pub struct NameServer { + pub net: DNSNetMode, + pub host: url::Host, + pub port: u16, + pub interface: Option, + pub proxy: Option, +} + +impl Display for NameServer { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!( + f, + "{}://{}:{}#{:?}", + self.net, self.host, self.port, self.interface, + ) + } +} + +impl NameServer { + pub async fn to_socket_addr(&self) -> anyhow::Result { + match &self.host { + url::Host::Ipv4(ip) => { + return Ok(SocketAddr::new((*ip).into(), self.port)); + } + url::Host::Ipv6(ip) => { + return Ok(SocketAddr::new((*ip).into(), self.port)); + } + url::Host::Domain(host) => { + return tokio::net::lookup_host((host.as_str(), self.port)) + .await? + .next() + .ok_or_else(|| { + anyhow::anyhow!("no ip resolved for dns server {}", host) + }); + } + } + } +} + +#[derive(Default)] +pub struct FallbackFilter { + pub geo_ip: bool, + pub geo_ip_code: String, + pub ip_cidr: Vec, + pub domain: Vec, +} + +#[derive(Clone, Debug, Default, PartialEq)] +pub struct EdnsClientSubnet { + pub ipv4: Option, + pub ipv6: Option, +} #[derive(Default)] pub struct DNSConfig { pub listen: DNSListenAddr, - /// 2 + pub nameserver: Vec, + pub fallback: Vec, + pub fallback_filter: FallbackFilter, + pub default_nameserver: Vec, + pub nameserver_policy: HashMap, + pub hosts: HashMap, + pub enhance_mode: DNSMode, + pub fake_ip_range: IpNet, + pub fake_ip_filter: Vec, + pub store_fake_ip: bool, pub ipv6: bool, - /// 3 pub enable: bool, + pub edns_client_subnet: Option, + pub fw_mark: Option, +} + +impl DNSConfig { + fn parse_nameserver(servers: &[String]) -> Result, Error> { + let mut nameservers = Vec::new(); + + for server in servers { + let mut server = server.clone(); + + if !server.contains("://") { + if server.contains(':') && !server.starts_with('[') { + server = format!("udp://[{server}]"); + } else { + server = format!("udp://{server}"); + } + } + + let url = Url::parse(&server).map_err(|_| { + Error::InvalidConfig(format!("invalid dns server: {}", server)) + })?; + + let host = match url.host() { + Some(host) => host, + None if url.scheme() == "dhcp" => url::Host::Domain("system"), + None => { + return Err(Error::InvalidConfig(format!( + "invalid dns server: no host found in {}", + server + ))); + } + }; + + let host = match host { + url::Host::Domain(value) => { + match value.parse::() { + Ok(ipv4) => url::Host::Ipv4(ipv4), + Err(_) => url::Host::Domain(value.to_string()), + } + } + value => value.to_owned(), + }; + + let (net, port) = match url.scheme() { + "udp" => (DNSNetMode::Udp, url.port().unwrap_or(53)), + "tcp" => (DNSNetMode::Tcp, url.port().unwrap_or(53)), + "tls" => (DNSNetMode::DoT, url.port().unwrap_or(853)), + "https" => (DNSNetMode::DoH, url.port().unwrap_or(443)), + "dhcp" => (DNSNetMode::Dhcp, url.port().unwrap_or(0)), + scheme => { + return Err(Error::InvalidConfig(format!( + "unsupported dns server scheme: {scheme}" + ))); + } + }; + + nameservers.push(NameServer { + net, + host, + port, + interface: None, + proxy: None, + }); + } + + Ok(nameservers) + } + + fn parse_nameserver_policy( + policy: &HashMap, + ) -> Result, Error> { + let mut out = HashMap::new(); + + for (domain, server) in policy { + let parsed = DNSConfig::parse_nameserver(std::slice::from_ref(server))?; + let ns = parsed.into_iter().next().ok_or_else(|| { + Error::InvalidConfig(format!( + "invalid dns nameserver policy for domain {domain}" + )) + })?; + out.insert(domain.to_ascii_lowercase(), ns); + } + + Ok(out) + } + + fn parse_hosts( + hosts: &HashMap, + ) -> Result, Error> { + let mut out = HashMap::from([( + "localhost".to_string(), + "127.0.0.1" + .parse::() + .expect("localhost ip should be valid"), + )]); + + for (host, ip) in hosts { + let ip = ip.parse::().map_err(|e| { + Error::InvalidConfig(format!("invalid hosts entry {host}: {e}")) + })?; + out.insert(host.trim_end_matches('.').to_ascii_lowercase(), ip); + } + + Ok(out) + } + + fn parse_fallback_filter( + filter: &DefFallbackFilter, + ) -> Result { + let mut ip_cidr = Vec::with_capacity(filter.ip_cidr.len()); + for cidr in &filter.ip_cidr { + ip_cidr.push(cidr.parse::().map_err(|e| { + Error::InvalidConfig(format!("invalid fallback ipcidr {cidr}: {e}")) + })?); + } + + Ok(FallbackFilter { + geo_ip: filter.geo_ip, + geo_ip_code: filter.geo_ip_code.clone(), + ip_cidr, + domain: filter.domain.clone(), + }) + } + + fn parse_edns_client_subnet( + ecs: &DefEdnsClientSubnet, + ) -> Result { + let ipv4 = ecs + .ipv4 + .as_ref() + .map(|value| { + value.parse::().map_err(|_| { + Error::InvalidConfig(format!( + "invalid edns-client-subnet ipv4 network: {value}" + )) + }) + }) + .transpose()?; + + let ipv6 = ecs + .ipv6 + .as_ref() + .map(|value| { + value.parse::().map_err(|_| { + Error::InvalidConfig(format!( + "invalid edns-client-subnet ipv6 network: {value}" + )) + }) + }) + .transpose()?; + + Ok(EdnsClientSubnet { ipv4, ipv6 }) + } } impl TryFrom for DNSConfig { @@ -38,17 +265,85 @@ impl TryFrom<&crate::config::def::Config> for DNSConfig { "invalid dns udp listen address: {u}" )) })?; - // future: will delete Ok::(DNSListenAddr { udp: Some(addr), ..Default::default() }) } + DNSListen::Multiple(map) => { + let mut udp = None; + let mut tcp = None; + + for (key, value) in map { + match key.as_str() { + "udp" => { + let addr = value + .as_str() + .ok_or_else(|| { + Error::InvalidConfig(format!( + "invalid udp dns listen address: {value:?}" + )) + })? + .parse::() + .map_err(|_| { + Error::InvalidConfig(format!( + "invalid dns udp listen address: {value:?}" + )) + })?; + udp = Some(addr); + } + "tcp" => { + let addr = value + .as_str() + .ok_or_else(|| { + Error::InvalidConfig(format!( + "invalid tcp dns listen address: {value:?}" + )) + })? + .parse::() + .map_err(|_| { + Error::InvalidConfig(format!( + "invalid dns tcp listen address: {value:?}" + )) + })?; + tcp = Some(addr); + } + _ => {} + } + } + + Ok::(DNSListenAddr { + udp, + tcp, + ..Default::default() + }) + } }) .transpose()? .unwrap_or_default(), + nameserver: DNSConfig::parse_nameserver(&dc.nameserver)?, + fallback: DNSConfig::parse_nameserver(&dc.fallback)?, + fallback_filter: DNSConfig::parse_fallback_filter(&dc.fallback_filter)?, + default_nameserver: DNSConfig::parse_nameserver(&dc.default_nameserver)?, + nameserver_policy: DNSConfig::parse_nameserver_policy( + &dc.nameserver_policy, + )?, + hosts: DNSConfig::parse_hosts(&c.hosts)?, + enhance_mode: dc.enhanced_mode.clone(), + fake_ip_range: dc + .fake_ip_range + .parse::() + .map_err(|e| Error::InvalidConfig(format!("invalid fake-ip-range: {e}")))?, + fake_ip_filter: dc.fake_ip_filter.clone(), + store_fake_ip: c.profile.store_fake_ip, ipv6: dc.ipv6, enable: dc.enable, + edns_client_subnet: dc + .edns_client_subnet + .as_ref() + .map(DNSConfig::parse_edns_client_subnet) + .transpose()?, + fw_mark: None, }) } } diff --git a/clash-lib/src/app/dns/dns_client.rs b/clash-lib/src/app/dns/dns_client.rs new file mode 100644 index 00000000..9a7b92aa --- /dev/null +++ b/clash-lib/src/app/dns/dns_client.rs @@ -0,0 +1,830 @@ +use std::{ + fmt::{Debug, Display, Formatter}, + net::{IpAddr, SocketAddr}, + sync::Arc, +}; + +use async_trait::async_trait; +use futures::{SinkExt, StreamExt}; +use hickory_proto::{ + op::Message, + op::ResponseCode, + rr::{ + RecordType, + rdata::opt::{ClientSubnet, EdnsCode, EdnsOption}, + }, + runtime::{Time, iocompat::AsyncIoTokioAsStd}, + rustls::{client_config, tls_client_stream::tls_client_connect_with_future}, + tcp::{TcpClientStream, TcpStream}, + xfer::{ + DnsExchange, DnsHandle, DnsMultiplexer, DnsRequest, DnsRequestOptions, + Protocol, + }, +}; +use hickory_resolver::{ + TokioResolver, + config::{NameServerConfig, ResolverConfig, ResolverOpts}, + name_server::TokioConnectionProvider, +}; +use tokio::sync::RwLock; +use tracing::warn; + +use crate::{ + Error, + app::dispatcher::{BoxedChainedDatagram, BoxedChainedStream}, + app::dns::{ + ClashResolver, config::EdnsClientSubnet, helper::build_dns_response_message, + }, + app::net::OutboundInterface, + proxy::{ + OutboundHandler, + datagram::UdpPacket, + utils::{new_tcp_stream, new_udp_socket}, + }, + session::{Network, Session, SocksAddr, Type}, +}; + +use super::{Client, ThreadSafeDNSClient, resolver::SystemResolver}; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum DNSNetMode { + Udp, + Tcp, + DoT, + DoH, + Dhcp, +} + +impl Display for DNSNetMode { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + match self { + Self::Udp => write!(f, "UDP"), + Self::Tcp => write!(f, "TCP"), + Self::DoT => write!(f, "DoT"), + Self::DoH => write!(f, "DoH"), + Self::Dhcp => write!(f, "DHCP"), + } + } +} + +#[derive(Clone)] +pub struct Opts { + pub father: Option>, + pub host: url::Host, + pub port: u16, + pub net: DNSNetMode, + pub iface: Option, + pub proxy: Arc, + pub ecs: Option, + pub fw_mark: Option, + pub ipv6: bool, +} + +enum DnsConfig { + Udp(SocketAddr), + Tcp(SocketAddr), + Tls(SocketAddr, url::Host), + Https(SocketAddr, url::Host), + Dhcp(Option), +} + +impl Display for DnsConfig { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + match self { + DnsConfig::Udp(addr) => write!(f, "UDP: {}:{}", addr.ip(), addr.port()), + DnsConfig::Tcp(addr) => write!(f, "TCP: {}:{}", addr.ip(), addr.port()), + DnsConfig::Tls(addr, host) => { + write!(f, "TLS: {}:{} host: {}", addr.ip(), addr.port(), host) + } + DnsConfig::Https(addr, host) => { + write!(f, "HTTPS: {}:{} host: {}", addr.ip(), addr.port(), host) + } + DnsConfig::Dhcp(iface) => write!(f, "DHCP: {:?}", iface), + } + } +} + +struct Inner { + resolver: Option, +} + +pub struct DnsClient { + inner: Arc>, + cfg: DnsConfig, + proxy: Arc, + host: url::Host, + port: u16, + net: DNSNetMode, + iface: Option, + ecs: Option, + fw_mark: Option, + ipv6: bool, + bind_addr: Option, +} + +impl DnsClient { + pub async fn new_client(opts: Opts) -> anyhow::Result { + if opts.net == DNSNetMode::Dhcp { + let iface = match &opts.host { + url::Host::Domain(iface) + if iface != "system" && !iface.is_empty() => + { + Some(iface.clone()) + } + _ => opts.iface.clone(), + }; + + return Ok(Arc::new(Self { + inner: Arc::new(RwLock::new(Inner { resolver: None })), + cfg: DnsConfig::Dhcp(iface.clone()), + proxy: opts.proxy, + host: opts.host, + port: opts.port, + net: opts.net, + iface, + ecs: opts.ecs, + fw_mark: opts.fw_mark, + ipv6: opts.ipv6, + bind_addr: None, + })); + } + + let resolved_ip = match &opts.host { + url::Host::Ipv4(ip) => Some(IpAddr::V4(*ip)), + url::Host::Ipv6(ip) => Some(IpAddr::V6(*ip)), + url::Host::Domain(domain) => match &opts.father { + Some(father) => father.resolve(domain, false).await?, + None => tokio::net::lookup_host((domain.as_str(), opts.port)) + .await? + .next() + .map(|addr| addr.ip()), + }, + }; + let ip = resolved_ip.ok_or_else(|| { + anyhow::anyhow!("no ip resolved for dns server {}", opts.host) + })?; + let socket_addr = SocketAddr::new(ip, opts.port); + let bind_addr = resolve_bind_addr(opts.iface.as_deref(), socket_addr); + + let cfg = match opts.net { + DNSNetMode::Udp => DnsConfig::Udp(socket_addr), + DNSNetMode::Tcp => DnsConfig::Tcp(socket_addr), + DNSNetMode::DoT => DnsConfig::Tls(socket_addr, opts.host.clone()), + DNSNetMode::DoH => DnsConfig::Https(socket_addr, opts.host.clone()), + DNSNetMode::Dhcp => unreachable!("dhcp handled before resolving host"), + }; + + Ok(Arc::new(Self { + inner: Arc::new(RwLock::new(Inner { resolver: None })), + cfg, + proxy: opts.proxy, + host: opts.host, + port: opts.port, + net: opts.net, + iface: opts.iface, + ecs: opts.ecs, + fw_mark: opts.fw_mark, + ipv6: opts.ipv6, + bind_addr, + })) + } + + fn apply_edns_client_subnet(&self, message: &mut Message) { + let Some(ecs) = &self.ecs else { + return; + }; + + if ecs.ipv4.is_none() && ecs.ipv6.is_none() { + return; + } + + if message + .extensions() + .as_ref() + .is_some_and(|edns| edns.option(EdnsCode::Subnet).is_some()) + { + return; + } + + let prefer_ipv6 = matches!( + message.query().map(|q| q.query_type()), + Some(RecordType::AAAA) + ); + + let candidate = if prefer_ipv6 { + ecs.ipv6 + .map(|ipv6| (IpAddr::from(ipv6.network()), ipv6.prefix_len())) + .or_else(|| { + ecs.ipv4.map(|ipv4| { + (IpAddr::from(ipv4.network()), ipv4.prefix_len()) + }) + }) + } else { + ecs.ipv4 + .map(|ipv4| (IpAddr::from(ipv4.network()), ipv4.prefix_len())) + .or_else(|| { + ecs.ipv6.map(|ipv6| { + (IpAddr::from(ipv6.network()), ipv6.prefix_len()) + }) + }) + }; + + let Some((addr, prefix)) = candidate else { + return; + }; + + let edns = message + .extensions_mut() + .get_or_insert_with(hickory_proto::op::Edns::new); + + let options = edns.options_mut(); + options.remove(EdnsCode::Subnet); + options.insert(EdnsOption::Subnet(ClientSubnet::new(addr, prefix, prefix))); + } + + async fn ensure_resolver(&self) -> anyhow::Result { + if let Some(resolver) = self.inner.read().await.resolver.clone() { + return Ok(resolver); + } + + if self.proxy.name() != "DIRECT" { + warn!( + proxy = self.proxy.name(), + dns = %self.id(), + "dns upstream proxy dialing is not implemented yet, falling back to direct connect" + ); + } + + if self.fw_mark.is_some() { + warn!( + fw_mark = self.fw_mark, + dns = %self.id(), + "dns upstream fw_mark is not implemented yet" + ); + } + + if let DnsConfig::Dhcp(iface) = &self.cfg { + if iface.is_some() { + warn!( + iface = ?iface, + dns = %self.id(), + "dhcp dns interface selection is not implemented yet; using system dns configuration" + ); + } + + let (config, mut resolver_opts) = + hickory_resolver::system_conf::read_system_conf().map_err(|e| { + anyhow::anyhow!( + "failed to read system dns config for dhcp upstream: {e}" + ) + })?; + resolver_opts.ip_strategy = if self.ipv6 { + hickory_resolver::config::LookupIpStrategy::Ipv4AndIpv6 + } else { + hickory_resolver::config::LookupIpStrategy::Ipv4Only + }; + + let resolver = TokioResolver::builder_with_config( + config, + TokioConnectionProvider::default(), + ) + .with_options(resolver_opts) + .build(); + + self.inner.write().await.resolver = Some(resolver.clone()); + return Ok(resolver); + } + + let mut config = ResolverConfig::new(); + let mut name_server = match &self.cfg { + DnsConfig::Udp(addr) => NameServerConfig::new(*addr, Protocol::Udp), + DnsConfig::Tcp(addr) => NameServerConfig::new(*addr, Protocol::Tcp), + DnsConfig::Tls(addr, host) => { + let mut ns = NameServerConfig::new(*addr, Protocol::Tls); + ns.tls_dns_name = Some(host.to_string()); + ns + } + DnsConfig::Https(addr, host) => { + let mut ns = NameServerConfig::new(*addr, Protocol::Https); + ns.tls_dns_name = Some(host.to_string()); + ns.http_endpoint = Some("/dns-query".to_string()); + ns + } + DnsConfig::Dhcp(_) => unreachable!("dhcp handled above"), + }; + name_server.bind_addr = self.bind_addr; + config.add_name_server(name_server); + + let mut resolver_opts = ResolverOpts::default(); + resolver_opts.ip_strategy = if self.ipv6 { + hickory_resolver::config::LookupIpStrategy::Ipv4AndIpv6 + } else { + hickory_resolver::config::LookupIpStrategy::Ipv4Only + }; + + let resolver = TokioResolver::builder_with_config( + config, + TokioConnectionProvider::default(), + ) + .with_options(resolver_opts) + .build(); + + self.inner.write().await.resolver = Some(resolver.clone()); + Ok(resolver) + } + + async fn exchange_via_resolver( + &self, + message: &Message, + ) -> anyhow::Result { + let query = message + .query() + .ok_or_else(|| anyhow::anyhow!("invalid query message"))?; + + let lookup = self + .ensure_resolver() + .await? + .lookup(query.name().clone(), query.query_type()) + .await?; + + let records: Vec<_> = lookup.record_iter().cloned().collect(); + let mut response = build_dns_response_message(message, true, false); + + if records.is_empty() { + response.set_response_code(ResponseCode::NXDomain); + return Ok(response); + } + + response.set_response_code(ResponseCode::NoError); + response.set_answer_count(records.len() as u16); + response.add_answers(records); + Ok(response) + } + + async fn exchange_direct_with_mark( + &self, + message: &Message, + ) -> anyhow::Result { + match &self.cfg { + DnsConfig::Udp(addr) => { + self.exchange_via_direct_udp(*addr, message).await + } + DnsConfig::Tcp(addr) => { + self.exchange_via_direct_tcp(*addr, message).await + } + DnsConfig::Tls(addr, host) => { + self.exchange_via_direct_tls(*addr, host.to_string(), message) + .await + } + DnsConfig::Https(addr, host) => { + self.exchange_via_direct_https(*addr, host.to_string(), message) + .await + } + DnsConfig::Dhcp(_) => self.exchange_via_resolver(message).await, + } + } + + async fn exchange_via_proxy( + &self, + message: &Message, + ) -> anyhow::Result { + match &self.cfg { + DnsConfig::Udp(addr) => { + self.exchange_via_proxy_udp(*addr, message).await + } + DnsConfig::Tcp(addr) => { + self.exchange_via_proxy_tcp(*addr, message).await + } + DnsConfig::Tls(addr, host) => { + self.exchange_via_proxy_tls(*addr, host.to_string(), message) + .await + } + DnsConfig::Https(addr, host) => { + self.exchange_via_proxy_https(*addr, host.to_string(), message) + .await + } + DnsConfig::Dhcp(_) => self.exchange_via_resolver(message).await, + } + } + + async fn exchange_via_proxy_udp( + &self, + socket_addr: SocketAddr, + message: &Message, + ) -> anyhow::Result { + let mut datagram = self.connect_proxy_datagram(socket_addr).await?; + let request = UdpPacket::new( + message.to_vec()?, + SocksAddr::any_ipv4(), + self.proxy_destination(socket_addr), + ); + + datagram.send(request).await?; + let response = + tokio::time::timeout(std::time::Duration::from_secs(5), datagram.next()) + .await + .map_err(|_| anyhow::anyhow!("dns udp upstream timeout"))? + .ok_or_else(|| { + anyhow::anyhow!("dns udp upstream returned no response") + })?; + + Ok(Message::from_vec(&response.data)?) + } + + async fn exchange_via_direct_udp( + &self, + socket_addr: SocketAddr, + message: &Message, + ) -> anyhow::Result { + let socket = new_udp_socket( + socket_addr, + resolve_outbound_interface(self.iface.as_deref()).as_ref(), + #[cfg(target_os = "linux")] + self.fw_mark, + )?; + + socket.send_to(&message.to_vec()?, socket_addr).await?; + let mut buf = vec![0u8; 65535]; + let (len, _) = tokio::time::timeout( + std::time::Duration::from_secs(5), + socket.recv_from(&mut buf), + ) + .await + .map_err(|_| anyhow::anyhow!("dns udp upstream timeout"))??; + + buf.truncate(len); + Ok(Message::from_vec(&buf)?) + } + + async fn exchange_via_proxy_tcp( + &self, + socket_addr: SocketAddr, + message: &Message, + ) -> anyhow::Result { + let future = self.connect_proxy_stream(socket_addr); + let (stream, handle) = TcpStream::with_future( + future, + socket_addr, + std::time::Duration::from_secs(5), + ); + let stream = Box::pin(async move { + let stream = stream.await?; + Ok::<_, hickory_proto::ProtoError>(TcpClientStream::from_stream(stream)) + }); + + self.run_multiplexed_exchange(stream, handle, message).await + } + + async fn exchange_via_direct_tcp( + &self, + socket_addr: SocketAddr, + message: &Message, + ) -> anyhow::Result { + let iface = resolve_outbound_interface(self.iface.as_deref()); + #[cfg(target_os = "linux")] + let so_mark = self.fw_mark; + let future = Box::pin(async move { + Ok::<_, std::io::Error>(AsyncIoTokioAsStd( + new_tcp_stream( + socket_addr, + iface.as_ref(), + #[cfg(target_os = "linux")] + so_mark, + ) + .await?, + )) + }); + let (stream, handle) = TcpStream::with_future( + future, + socket_addr, + std::time::Duration::from_secs(5), + ); + let stream = Box::pin(async move { + let stream = stream.await?; + Ok::<_, hickory_proto::ProtoError>(TcpClientStream::from_stream(stream)) + }); + self.run_multiplexed_exchange(stream, handle, message).await + } + + async fn exchange_via_direct_tls( + &self, + socket_addr: SocketAddr, + dns_name: String, + message: &Message, + ) -> anyhow::Result { + let iface = resolve_outbound_interface(self.iface.as_deref()); + #[cfg(target_os = "linux")] + let so_mark = self.fw_mark; + let mut tls_config = client_config(); + tls_config.enable_sni = false; + let (stream, handle) = tls_client_connect_with_future( + Box::pin(async move { + let stream = new_tcp_stream( + socket_addr, + iface.as_ref(), + #[cfg(target_os = "linux")] + so_mark, + ) + .await?; + Ok(AsyncIoTokioAsStd(stream)) + }), + socket_addr, + dns_name, + Arc::new(tls_config), + ); + + self.run_multiplexed_exchange(stream, handle, message).await + } + + async fn exchange_via_direct_https( + &self, + socket_addr: SocketAddr, + dns_name: String, + message: &Message, + ) -> anyhow::Result { + let iface = resolve_outbound_interface(self.iface.as_deref()); + #[cfg(target_os = "linux")] + let so_mark = self.fw_mark; + let exchange: hickory_proto::xfer::DnsExchangeConnect< + _, + _, + hickory_proto::runtime::TokioTime, + > = DnsExchange::connect(hickory_proto::h2::HttpsClientConnect::new( + Box::pin(async move { + let stream = new_tcp_stream( + socket_addr, + iface.as_ref(), + #[cfg(target_os = "linux")] + so_mark, + ) + .await?; + Ok(AsyncIoTokioAsStd(stream)) + }), + Arc::new(client_config()), + socket_addr, + dns_name, + "/dns-query".to_string(), + )); + + self.run_exchange_connect(exchange, message).await + } + + async fn exchange_via_proxy_tls( + &self, + socket_addr: SocketAddr, + dns_name: String, + message: &Message, + ) -> anyhow::Result { + let mut tls_config = client_config(); + tls_config.enable_sni = false; + + let (stream, handle) = tls_client_connect_with_future( + self.connect_proxy_stream(socket_addr), + socket_addr, + dns_name, + Arc::new(tls_config), + ); + + self.run_multiplexed_exchange(stream, handle, message).await + } + + async fn exchange_via_proxy_https( + &self, + socket_addr: SocketAddr, + dns_name: String, + message: &Message, + ) -> anyhow::Result { + let exchange: hickory_proto::xfer::DnsExchangeConnect< + _, + _, + hickory_proto::runtime::TokioTime, + > = DnsExchange::connect(hickory_proto::h2::HttpsClientConnect::new( + self.connect_proxy_stream(socket_addr), + Arc::new(client_config()), + socket_addr, + dns_name, + "/dns-query".to_string(), + )); + + self.run_exchange_connect(exchange, message).await + } + + async fn run_multiplexed_exchange( + &self, + stream: F, + handle: hickory_proto::BufDnsStreamHandle, + message: &Message, + ) -> anyhow::Result + where + F: std::future::Future> + + Send + + Unpin + + 'static, + S: hickory_proto::xfer::DnsClientStream + Unpin + 'static, + { + let exchange: hickory_proto::xfer::DnsExchangeConnect< + _, + _, + hickory_proto::runtime::TokioTime, + > = DnsExchange::connect(DnsMultiplexer::with_timeout( + stream, + handle, + std::time::Duration::from_secs(5), + None, + )); + + self.run_exchange_connect(exchange, message).await + } + + async fn run_exchange_connect( + &self, + exchange: hickory_proto::xfer::DnsExchangeConnect, + message: &Message, + ) -> anyhow::Result + where + F: std::future::Future> + + Send + + Unpin + + 'static, + S: hickory_proto::xfer::DnsRequestSender, + TE: Time + Unpin + Send + 'static, + { + let (exchange, background) = exchange.await?; + tokio::spawn(background); + + let mut options = DnsRequestOptions::default(); + options.use_edns = true; + options.recursion_desired = message.recursion_desired(); + + let response = exchange + .send(DnsRequest::new(message.clone(), options)) + .next() + .await + .ok_or_else(|| anyhow::anyhow!("dns upstream returned no response"))??; + + Ok(response.into_message()) + } + + fn connect_proxy_stream( + &self, + socket_addr: SocketAddr, + ) -> std::pin::Pin< + Box< + dyn std::future::Future< + Output = std::io::Result>, + > + Send + + 'static, + >, + > { + let proxy = self.proxy.clone(); + let resolver = Arc::new( + SystemResolver::new(self.ipv6) + .expect("failed to create system resolver for proxied dns upstream"), + ); + let destination = self.proxy_destination(socket_addr); + let iface = resolve_outbound_interface(self.iface.as_deref()); + let so_mark = self.fw_mark; + + Box::pin(async move { + let session = Session { + network: Network::Tcp, + typ: Type::Ignore, + destination, + iface, + #[cfg(target_os = "linux")] + so_mark, + ..Default::default() + }; + + let stream = proxy.connect_stream(&session, resolver).await?; + Ok(AsyncIoTokioAsStd(stream)) + }) + } + + async fn connect_proxy_datagram( + &self, + socket_addr: SocketAddr, + ) -> std::io::Result { + let resolver = Arc::new( + SystemResolver::new(self.ipv6) + .expect("failed to create system resolver for proxied dns upstream"), + ); + let session = Session { + network: Network::Udp, + typ: Type::Ignore, + destination: self.proxy_destination(socket_addr), + iface: resolve_outbound_interface(self.iface.as_deref()), + #[cfg(target_os = "linux")] + so_mark: self.fw_mark, + ..Default::default() + }; + + self.proxy.connect_datagram(&session, resolver).await + } + + fn proxy_destination(&self, socket_addr: SocketAddr) -> SocksAddr { + match &self.host { + url::Host::Domain(domain) => { + SocksAddr::Domain(domain.clone(), self.port) + } + url::Host::Ipv4(_) | url::Host::Ipv6(_) => SocksAddr::Ip(socket_addr), + } + } +} + +fn interface_bind_addr( + iface: &OutboundInterface, + remote: SocketAddr, +) -> Option { + match remote { + SocketAddr::V4(_) => { + iface.addr_v4.map(|ip| SocketAddr::new(IpAddr::V4(ip), 0)) + } + SocketAddr::V6(_) => { + iface.addr_v6.map(|ip| SocketAddr::new(IpAddr::V6(ip), 0)) + } + } +} + +#[cfg(feature = "tun")] +fn resolve_bind_addr( + iface_name: Option<&str>, + remote: SocketAddr, +) -> Option { + let iface_name = iface_name?; + let iface = crate::app::net::get_interface_by_name(iface_name); + let bind_addr = iface + .as_ref() + .and_then(|iface| interface_bind_addr(iface, remote)); + + if bind_addr.is_none() { + warn!( + iface = iface_name, + remote = %remote, + "dns upstream interface requested but no compatible address was found" + ); + } + + bind_addr +} + +#[cfg(not(feature = "tun"))] +fn resolve_bind_addr( + iface_name: Option<&str>, + _remote: SocketAddr, +) -> Option { + if let Some(iface_name) = iface_name { + warn!( + iface = iface_name, + "dns upstream interface binding requires the `tun` feature; ignoring interface" + ); + } + + None +} + +impl Debug for DnsClient { + fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result { + f.debug_struct("DnsClient") + .field("host", &self.host) + .field("port", &self.port) + .field("net", &self.net) + .field("iface", &self.iface) + .field("proxy", &self.proxy.name()) + .finish() + } +} + +#[async_trait] +impl Client for DnsClient { + fn id(&self) -> String { + format!("{}#{}:{}", &self.net, &self.host, &self.port) + } + + async fn exchange(&self, msg: &Message) -> anyhow::Result { + let mut outbound = msg.clone(); + self.apply_edns_client_subnet(&mut outbound); + + if self.proxy.name() == "DIRECT" && self.fw_mark.is_some() { + self.exchange_direct_with_mark(&outbound).await + } else if self.proxy.name() == "DIRECT" { + self.exchange_via_resolver(&outbound).await + } else { + self.exchange_via_proxy(&outbound).await + } + } +} + +#[cfg(feature = "tun")] +fn resolve_outbound_interface( + iface_name: Option<&str>, +) -> Option { + iface_name.and_then(crate::app::net::get_interface_by_name) +} + +#[cfg(not(feature = "tun"))] +fn resolve_outbound_interface( + _iface_name: Option<&str>, +) -> Option { + None +} diff --git a/clash-lib/src/app/dns/fakeip/file_store.rs b/clash-lib/src/app/dns/fakeip/file_store.rs new file mode 100644 index 00000000..c297c4c4 --- /dev/null +++ b/clash-lib/src/app/dns/fakeip/file_store.rs @@ -0,0 +1,44 @@ +use async_trait::async_trait; + +use crate::app::profile::ThreadSafeCacheFile; + +use super::Store; + +pub struct FileStore(ThreadSafeCacheFile); + +impl FileStore { + pub fn new(store: ThreadSafeCacheFile) -> Self { + Self(store) + } +} + +#[async_trait] +impl Store for FileStore { + async fn get_by_host(&mut self, host: &str) -> Option { + self.0 + .get_fake_ip(host) + .await + .and_then(|ip| ip.parse::().ok()) + } + + async fn pub_by_host(&mut self, host: &str, ip: std::net::IpAddr) { + self.0.set_host_to_ip(host, &ip.to_string()).await; + } + + async fn get_by_ip(&mut self, ip: std::net::IpAddr) -> Option { + self.0.get_fake_ip(&ip.to_string()).await + } + + async fn put_by_ip(&mut self, ip: std::net::IpAddr, host: &str) { + self.0.set_ip_to_host(&ip.to_string(), host).await; + } + + async fn del_by_ip(&mut self, ip: std::net::IpAddr) { + let host = self.get_by_ip(ip).await.unwrap_or_default(); + self.0.delete_fake_ip_pair(&ip.to_string(), &host).await; + } + + async fn exist(&mut self, ip: std::net::IpAddr) -> bool { + self.0.get_fake_ip(&ip.to_string()).await.is_some() + } +} diff --git a/clash-lib/src/app/dns/fakeip/mem_store.rs b/clash-lib/src/app/dns/fakeip/mem_store.rs new file mode 100644 index 00000000..b7176b30 --- /dev/null +++ b/clash-lib/src/app/dns/fakeip/mem_store.rs @@ -0,0 +1,85 @@ +use std::collections::{HashMap, VecDeque}; +use std::net::IpAddr; + +use async_trait::async_trait; + +use super::Store; + +pub struct InMemStore { + capacity: usize, + ip_to_host: HashMap, + host_to_ip: HashMap, + order: VecDeque, +} + +impl InMemStore { + pub fn new(capacity: usize) -> Self { + Self { + capacity, + ip_to_host: HashMap::new(), + host_to_ip: HashMap::new(), + order: VecDeque::new(), + } + } + + fn touch(&mut self, host: &str) { + if let Some(index) = self.order.iter().position(|item| item == host) { + self.order.remove(index); + } + self.order.push_back(host.to_string()); + self.evict_if_needed(); + } + + fn evict_if_needed(&mut self) { + while self.host_to_ip.len() > self.capacity { + let Some(host) = self.order.pop_front() else { + break; + }; + if let Some(ip) = self.host_to_ip.remove(&host) { + self.ip_to_host.remove(&ip); + } + } + } +} + +#[async_trait] +impl Store for InMemStore { + async fn get_by_host(&mut self, host: &str) -> Option { + let ip = self.host_to_ip.get(host).copied(); + if ip.is_some() { + self.touch(host); + } + ip + } + + async fn pub_by_host(&mut self, host: &str, ip: IpAddr) { + self.host_to_ip.insert(host.to_string(), ip); + self.touch(host); + } + + async fn get_by_ip(&mut self, ip: IpAddr) -> Option { + let host = self.ip_to_host.get(&ip).cloned(); + if let Some(host) = &host { + self.touch(host); + } + host + } + + async fn put_by_ip(&mut self, ip: IpAddr, host: &str) { + self.ip_to_host.insert(ip, host.to_string()); + self.touch(host); + } + + async fn del_by_ip(&mut self, ip: IpAddr) { + if let Some(host) = self.ip_to_host.remove(&ip) { + self.host_to_ip.remove(&host); + if let Some(index) = self.order.iter().position(|item| item == &host) { + self.order.remove(index); + } + } + } + + async fn exist(&mut self, ip: IpAddr) -> bool { + self.ip_to_host.contains_key(&ip) + } +} diff --git a/clash-lib/src/app/dns/fakeip/mod.rs b/clash-lib/src/app/dns/fakeip/mod.rs new file mode 100644 index 00000000..3f362b9b --- /dev/null +++ b/clash-lib/src/app/dns/fakeip/mod.rs @@ -0,0 +1,218 @@ +use std::net::{IpAddr, Ipv4Addr}; + +use async_trait::async_trait; +use tokio::sync::RwLock; + +use crate::{Error, common::trie::StringTrie}; + +mod file_store; +mod mem_store; + +pub use file_store::FileStore; +pub use mem_store::InMemStore; +pub type ThreadSafeFakeDns = std::sync::Arc>; + +pub struct Opts { + pub ipnet: ipnet::IpNet, + pub skipped_hostnames: Option>, + pub store: Box, +} + +#[async_trait] +pub trait Store: Sync + Send { + async fn get_by_host(&mut self, host: &str) -> Option; + async fn pub_by_host(&mut self, host: &str, ip: IpAddr); + async fn get_by_ip(&mut self, ip: IpAddr) -> Option; + async fn put_by_ip(&mut self, ip: IpAddr, host: &str); + async fn del_by_ip(&mut self, ip: IpAddr); + async fn exist(&mut self, ip: IpAddr) -> bool; +} + +pub struct FakeDns { + max: u32, + min: u32, + offset: u32, + skipped_hostnames: Option>, + ipnet: ipnet::IpNet, + store: Box, +} + +impl FakeDns { + pub fn new(opt: Opts) -> Result { + let ip = match opt.ipnet.network() { + IpAddr::V4(ip) => ip, + _ => { + return Err(Error::InvalidConfig( + "fake-ip-range must be valid ipv4 subnet".to_string(), + )); + } + }; + + let min = Self::ip_to_uint(&ip) + 2; + let prefix_len = opt.ipnet.prefix_len(); + let max_prefix_len = opt.ipnet.max_prefix_len(); + let total = (1 << (max_prefix_len - prefix_len)) - 2; + let max = min + total - 1; + + Ok(Self { + max, + min, + offset: 0, + skipped_hostnames: opt.skipped_hostnames, + ipnet: opt.ipnet, + store: opt.store, + }) + } + + pub async fn lookup(&mut self, host: &str) -> IpAddr { + if let Some(ip) = self.store.get_by_host(host).await { + return ip; + } + + let ip = self.next_ip(host).await; + self.store.pub_by_host(host, ip).await; + ip + } + + pub async fn reverse_lookup(&mut self, ip: IpAddr) -> Option { + if ip.is_ipv4() { + self.store.get_by_ip(ip).await + } else { + None + } + } + + pub fn should_skip(&self, domain: &str) -> bool { + self.skipped_hostnames.as_ref().is_some_and(|hostnames| { + hostnames + .search(&domain.trim_end_matches('.').to_ascii_lowercase()) + .is_some() + }) + } + + pub async fn is_fake_ip(&mut self, ip: IpAddr) -> bool { + ip.is_ipv4() && self.ipnet.contains(&ip) + } + + async fn next_ip(&mut self, host: &str) -> IpAddr { + let current = self.offset; + + loop { + self.offset = (self.offset + 1) % (self.max - self.min); + + if self.offset == current { + self.offset = (self.offset + 1) % (self.max - self.min); + let ip = Ipv4Addr::from(self.min + self.offset - 1); + self.store.del_by_ip(IpAddr::V4(ip)).await; + break; + } + + let ip = Ipv4Addr::from(self.min + self.offset - 1); + if !self.store.exist(IpAddr::V4(ip)).await { + break; + } + } + + let ip = Ipv4Addr::from(self.min + self.offset - 1); + self.store.put_by_ip(IpAddr::V4(ip), host).await; + IpAddr::V4(ip) + } + + fn ip_to_uint(ip: &Ipv4Addr) -> u32 { + u32::from_be_bytes(ip.octets()) + } +} + +#[cfg(test)] +mod tests { + use super::{FakeDns, InMemStore, Opts}; + use crate::common::trie::StringTrie; + use std::{net::IpAddr, sync::Arc}; + + #[tokio::test] + async fn allocates_and_reuses_addresses() { + let mut fake_dns = FakeDns::new(Opts { + ipnet: "192.168.0.0/29".parse().expect("valid fake-ip range"), + skipped_hostnames: None, + store: Box::new(InMemStore::new(16)), + }) + .expect("fake dns should build"); + + let first = fake_dns.lookup("foo.com").await; + let second = fake_dns.lookup("bar.com").await; + + assert_eq!(first, IpAddr::from([192, 168, 0, 2])); + assert_eq!(fake_dns.lookup("foo.com").await, first); + assert_eq!(second, IpAddr::from([192, 168, 0, 3])); + assert_eq!( + fake_dns.reverse_lookup(second).await, + Some("bar.com".into()) + ); + assert!(fake_dns.is_fake_ip(second).await); + assert!( + !fake_dns + .is_fake_ip("::1".parse().expect("valid ipv6")) + .await + ); + } + + #[tokio::test] + async fn cycles_when_pool_is_exhausted() { + let mut fake_dns = FakeDns::new(Opts { + ipnet: "192.168.0.0/29".parse().expect("valid fake-ip range"), + skipped_hostnames: None, + store: Box::new(InMemStore::new(16)), + }) + .expect("fake dns should build"); + + let first = fake_dns.lookup("foo.com").await; + let second = fake_dns.lookup("bar.com").await; + + for index in 0..3 { + fake_dns.lookup(&format!("{index}.com")).await; + } + + let recycled = fake_dns.lookup("baz.com").await; + let next = fake_dns.lookup("foo.com").await; + + assert_eq!(recycled, first); + assert_eq!(next, second); + } + + #[tokio::test] + async fn reassigns_when_store_capacity_evicts_entry() { + let mut fake_dns = FakeDns::new(Opts { + ipnet: "192.168.0.0/24".parse().expect("valid fake-ip range"), + skipped_hostnames: None, + store: Box::new(InMemStore::new(2)), + }) + .expect("fake dns should build"); + + let first = fake_dns.lookup("foo.com").await; + fake_dns.lookup("bar.com").await; + fake_dns.lookup("baz.com").await; + let next = fake_dns.lookup("foo.com").await; + + assert_ne!(first, next); + } + + #[tokio::test] + async fn skips_hosts_via_trie() { + let mut skipped = StringTrie::new(); + skipped.insert("*.example.com", Arc::new(true)); + + let fake_dns = FakeDns::new(Opts { + ipnet: "198.18.0.0/16".parse().expect("valid fake-ip range"), + skipped_hostnames: Some(skipped), + store: Box::new(InMemStore::new(16)), + }) + .expect("fake dns should build"); + + assert!(fake_dns.should_skip("foo.example.com")); + assert!(!fake_dns.should_skip("example.com")); + assert!(!fake_dns.should_skip("foo.example.net")); + assert!( + !fake_dns.should_skip(IpAddr::from([127, 0, 0, 1]).to_string().as_str()) + ); + } +} diff --git a/clash-lib/src/app/dns/filters.rs b/clash-lib/src/app/dns/filters.rs new file mode 100644 index 00000000..ef15b057 --- /dev/null +++ b/clash-lib/src/app/dns/filters.rs @@ -0,0 +1,96 @@ +use std::{net::IpAddr, sync::Arc}; + +use crate::common::{mmdb::MmdbLookup, trie::StringTrie}; + +pub trait FallbackIpFilter: Sync + Send { + fn apply(&self, ip: &IpAddr) -> bool; +} +pub use FallbackIpFilter as FallbackIPFilter; + +pub trait FallbackDomainFilter: Sync + Send { + fn apply(&self, domain: &str) -> bool; +} + +pub struct GeoIpFilter { + code: String, + mmdb: Option, +} +pub use GeoIpFilter as GeoIPFilter; + +impl GeoIpFilter { + pub fn new(code: &str, mmdb: Option) -> Self { + Self { + code: code.to_string(), + mmdb, + } + } +} + +impl FallbackIpFilter for GeoIpFilter { + fn apply(&self, ip: &IpAddr) -> bool { + !self.mmdb.as_ref().is_some_and(|mmdb| { + mmdb.lookup_country(*ip) + .map(|country| country.country_code == self.code) + .unwrap_or(false) + }) + } +} + +pub struct IpNetFilter(ipnet::IpNet); + +impl IpNetFilter { + pub fn new(ipnet: ipnet::IpNet) -> Self { + Self(ipnet) + } +} +pub use IpNetFilter as IPNetFilter; + +impl FallbackIpFilter for IpNetFilter { + fn apply(&self, ip: &IpAddr) -> bool { + self.0.contains(ip) + } +} + +pub struct DomainFilter(StringTrie>); + +impl DomainFilter { + pub fn new(domains: &[String]) -> Self { + let mut filter = Self(StringTrie::new()); + + for domain in domains { + let domain = domain.trim_end_matches('.').to_ascii_lowercase(); + filter.0.insert(&domain, Arc::new(None)); + } + + filter + } +} + +impl FallbackDomainFilter for DomainFilter { + fn apply(&self, domain: &str) -> bool { + let domain = domain.trim_end_matches('.').to_ascii_lowercase(); + self.0.search(&domain).is_some() + } +} + +#[cfg(test)] +mod tests { + use super::{DomainFilter, FallbackDomainFilter}; + + #[test] + fn domain_filter_matches_trie_patterns() { + let filter = DomainFilter::new(&[ + "*.example.com".to_string(), + ".apple.*".to_string(), + "+.foo.com".to_string(), + ]); + + assert!(filter.apply("sub.example.com")); + assert!(filter.apply("test.apple.com")); + assert!(filter.apply("foo.com")); + assert!(filter.apply("bar.foo.com")); + + assert!(!filter.apply("example.com")); + assert!(!filter.apply("foo.example.net")); + } +} diff --git a/clash-lib/src/app/dns/helper.rs b/clash-lib/src/app/dns/helper.rs new file mode 100644 index 00000000..af301040 --- /dev/null +++ b/clash-lib/src/app/dns/helper.rs @@ -0,0 +1,116 @@ +use tracing::{debug, warn}; + +use crate::{ + app::dns::{ + ClashResolver, ThreadSafeDNSClient, + config::{EdnsClientSubnet, NameServer}, + dns_client::DnsClient, + dns_client::Opts, + }, + proxy, +}; + +use hickory_proto::{ + op::{Message, MessageType}, + rr::{ + RData, Record, RecordType, + rdata::{A, AAAA}, + }, +}; + +pub async fn make_clients( + servers: &[NameServer], + resolver: Option>, + outbounds: std::collections::HashMap< + String, + std::sync::Arc, + >, + edns_client_subnet: Option, + fw_mark: Option, + ipv6: bool, +) -> Vec { + let mut rv = Vec::new(); + + for server in servers { + debug!( + host = %server.host, + port = server.port, + "building nameserver" + ); + + match DnsClient::new_client(Opts { + father: resolver.as_ref().cloned(), + host: server.host.clone(), + port: server.port, + net: server.net.clone(), + iface: server.interface.clone(), + proxy: outbounds + .get(server.proxy.as_deref().unwrap_or("DIRECT")) + .cloned() + .unwrap_or_else(|| { + std::sync::Arc::new(proxy::direct::Handler::new("DIRECT")) + }), + ecs: edns_client_subnet.clone(), + fw_mark, + ipv6, + }) + .await + { + Ok(client) => rv.push(client), + Err(err) => warn!( + host = %server.host, + port = server.port, + err = ?err, + "initializing dns client failed" + ), + } + } + + rv +} + +pub fn build_dns_response_message( + req: &Message, + recursive_available: bool, + authoritative: bool, +) -> Message { + let mut res = Message::new(); + + res.set_id(req.id()); + res.set_op_code(req.op_code()); + res.set_message_type(MessageType::Response); + res.add_queries(req.queries().iter().cloned()); + res.set_recursion_available(recursive_available); + res.set_authoritative(authoritative); + res.set_recursion_desired(req.recursion_desired()); + res.set_checking_disabled(req.checking_disabled()); + if let Some(edns) = req.extensions().clone() { + res.set_edns(edns); + } + + if let Some(edns) = res.extensions_mut() { + edns.options_mut() + .remove(hickory_proto::rr::rdata::opt::EdnsCode::Padding); + } + + res +} + +pub fn ip_records( + name: hickory_proto::rr::Name, + ttl: u32, + query_type: RecordType, + ips: &[std::net::IpAddr], +) -> Vec { + ips.iter() + .filter_map(|ip| match (query_type, ip) { + (RecordType::A, std::net::IpAddr::V4(ip)) => { + Some(Record::from_rdata(name.clone(), ttl, RData::A(A(*ip)))) + } + (RecordType::AAAA, std::net::IpAddr::V6(ip)) => Some( + Record::from_rdata(name.clone(), ttl, RData::AAAA(AAAA(*ip))), + ), + _ => None, + }) + .collect() +} diff --git a/clash-lib/src/app/dns/mod.rs b/clash-lib/src/app/dns/mod.rs index 42e966a3..de2b3374 100644 --- a/clash-lib/src/app/dns/mod.rs +++ b/clash-lib/src/app/dns/mod.rs @@ -2,21 +2,32 @@ use async_trait::async_trait; use hickory_proto::op; -use std::sync::Arc; +use std::{ + fmt::Debug, + net::{IpAddr, Ipv4Addr, Ipv6Addr}, + sync::Arc, +}; /// 2 mod config; +mod dns_client; +mod fakeip; +mod filters; +mod helper; /// 3 pub mod resolver; /// 1 mod server; pub use config::DNSConfig; +pub use dns_client::DNSNetMode; +pub use server::exchange_with_resolver; pub use server::get_dns_listener; pub use resolver::new as new_resolver; pub type ThreadSafeDNSResolver = Arc; +pub type ThreadSafeDNSClient = Arc; #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum ResolverKind { @@ -24,6 +35,12 @@ pub enum ResolverKind { System, } +#[async_trait] +pub trait Client: Sync + Send + Debug { + fn id(&self) -> String; + async fn exchange(&self, msg: &op::Message) -> anyhow::Result; +} + /// A implementation of "anti-poisoning" Resolver /// it can hold multiple clients in different protocols /// each client can also hold a "default_resolver" @@ -35,10 +52,39 @@ pub trait ClashResolver: Sync + Send { async fn exchange(&self, message: &op::Message) -> anyhow::Result; fn ipv6(&self) -> bool; + fn set_ipv6(&self, enable: bool); + fn kind(&self) -> ResolverKind; + fn fake_ip_enabled(&self) -> bool; + + async fn reverse_lookup(&self, ip: IpAddr) -> Option; + async fn is_fake_ip(&self, ip: IpAddr) -> bool; + async fn cached_for(&self, ip: IpAddr) -> Option; async fn resolve( &self, host: &str, enhanced: bool, ) -> anyhow::Result>; + + async fn resolve_v4( + &self, + host: &str, + enhanced: bool, + ) -> anyhow::Result> { + Ok(match self.resolve(host, enhanced).await? { + Some(IpAddr::V4(ip)) => Some(ip), + _ => None, + }) + } + + async fn resolve_v6( + &self, + host: &str, + enhanced: bool, + ) -> anyhow::Result> { + Ok(match self.resolve(host, enhanced).await? { + Some(IpAddr::V6(ip)) => Some(ip), + _ => None, + }) + } } diff --git a/clash-lib/src/app/dns/resolver/enhanced.rs b/clash-lib/src/app/dns/resolver/enhanced.rs new file mode 100644 index 00000000..2791b986 --- /dev/null +++ b/clash-lib/src/app/dns/resolver/enhanced.rs @@ -0,0 +1,727 @@ +use std::{ + collections::HashMap, + net::IpAddr, + sync::{ + Arc, + atomic::{AtomicBool, Ordering}, + }, + time::{Duration, Instant}, +}; + +use async_trait::async_trait; +use futures::{FutureExt, TryFutureExt, future}; +use hickory_proto::{ + op::Message, + rr::RecordType, + rr::{RData, Record}, +}; +use hickory_resolver::dns_lru::{DnsLru, TtlConfig}; +use lru_time_cache::LruCache; +use tokio::sync::RwLock; +use tracing::{debug, error, instrument, trace, warn}; + +use super::SystemResolver; + +use crate::{ + app::{ + dns::helper::build_dns_response_message, + dns::{ + ClashResolver, DNSConfig, ThreadSafeDNSClient, + fakeip::{ + FakeDns, FileStore, InMemStore, Opts as FakeDnsOpts, + ThreadSafeFakeDns, + }, + filters::{ + DomainFilter, FallbackDomainFilter, FallbackIPFilter, GeoIPFilter, + IPNetFilter, + }, + helper::make_clients, + }, + profile::ThreadSafeCacheFile, + }, + common::{mmdb::MmdbLookup, trie::StringTrie}, + config::def::DNSMode, + proxy::OutboundHandler, +}; + +pub struct EnhancedResolver { + ipv6: AtomicBool, + store: ThreadSafeCacheFile, + hosts: Option>, + main: Vec, + fallback: Option>, + fallback_domain_filters: Option>>, + fallback_ip_filters: Option>>, + lru_cache: Option>>, + policy: Option>>, + fake_dns: Option, + reverse_lookup_cache: Option>>>, + _mmdb: Option, + _outbounds: HashMap>, +} + +impl EnhancedResolver { + pub async fn new( + cfg: DNSConfig, + store: ThreadSafeCacheFile, + mmdb: Option, + outbounds: HashMap>, + ) -> Self { + debug!( + ipv6 = cfg.ipv6, + nameservers = cfg.nameserver.len(), + "creating enhanced resolver" + ); + + let (fallback_domain_filters, fallback_ip_filters) = + build_fallback_filters(&cfg, mmdb.clone()); + let hosts = build_hosts_trie(&cfg.hosts); + let fake_dns = + build_fake_dns(&cfg, store.clone()).expect("failed to create fake dns"); + let default_resolver = Arc::new(EnhancedResolver { + ipv6: AtomicBool::new(false), + store: store.clone(), + hosts: None, + main: make_clients( + &cfg.default_nameserver, + None, + outbounds.clone(), + cfg.edns_client_subnet.clone(), + cfg.fw_mark, + false, + ) + .await, + fallback: None, + fallback_domain_filters: None, + fallback_ip_filters: None, + lru_cache: None, + policy: None, + fake_dns: None, + reverse_lookup_cache: None, + _mmdb: None, + _outbounds: outbounds.clone(), + }); + let main = make_clients( + &cfg.nameserver, + Some(default_resolver.clone()), + outbounds.clone(), + cfg.edns_client_subnet.clone(), + cfg.fw_mark, + cfg.ipv6, + ) + .await; + let fallback = if cfg.fallback.is_empty() { + None + } else { + Some( + make_clients( + &cfg.fallback, + Some(default_resolver.clone()), + outbounds.clone(), + cfg.edns_client_subnet.clone(), + cfg.fw_mark, + cfg.ipv6, + ) + .await, + ) + }; + let policy = build_policy_resolvers( + &cfg, + Some(default_resolver.clone()), + outbounds.clone(), + ) + .await; + + Self { + ipv6: AtomicBool::new(cfg.ipv6), + store, + hosts, + main, + fallback, + fallback_domain_filters, + fallback_ip_filters, + lru_cache: Some(Arc::new(RwLock::new(DnsLru::new( + 4096, + TtlConfig::new( + Some(Duration::from_secs(1)), + Some(Duration::from_secs(1)), + Some(Duration::from_secs(60)), + Some(Duration::from_secs(10)), + ), + )))), + policy, + fake_dns, + reverse_lookup_cache: Some(Arc::new(RwLock::new( + LruCache::with_expiry_duration_and_capacity( + Duration::from_secs(3), + 4096, + ), + ))), + _mmdb: mmdb, + _outbounds: outbounds, + } + } +} + +#[async_trait] +impl ClashResolver for EnhancedResolver { + async fn exchange(&self, message: &Message) -> anyhow::Result { + let query = message + .query() + .ok_or_else(|| anyhow::anyhow!("invalid query message"))?; + + if let Some(lru) = &self.lru_cache + && let Some(cached) = lru.read().await.get(query, Instant::now()) + { + if !message.recursion_desired() { + trace!(query = %query.name(), "dns cache hit"); + if let Ok(cached) = cached { + let mut response = + build_dns_response_message(message, true, false); + response.add_answers(cached.records().iter().cloned()); + return Ok(response); + } + } else { + trace!(query = %query.name(), "dns cache present but bypassed"); + } + } + + self.exchange_no_cache(message).await.map(|mut response| { + if let Some(edns) = response.extensions_mut() { + edns.options_mut() + .remove(hickory_proto::rr::rdata::opt::EdnsCode::Padding); + } + response + }) + } + + fn ipv6(&self) -> bool { + self.ipv6.load(Ordering::Relaxed) + } + + fn set_ipv6(&self, enable: bool) { + self.ipv6.store(enable, Ordering::Relaxed); + } + + fn kind(&self) -> crate::app::dns::ResolverKind { + crate::app::dns::ResolverKind::Clash + } + + fn fake_ip_enabled(&self) -> bool { + self.fake_dns.is_some() + } + + async fn is_fake_ip(&self, ip: std::net::IpAddr) -> bool { + let Some(fake_dns) = &self.fake_dns else { + return false; + }; + + fake_dns.write().await.is_fake_ip(ip).await + } + + async fn resolve( + &self, + host: &str, + enhanced: bool, + ) -> anyhow::Result> { + let resolved = if self.ipv6() { + let v6 = self + .resolve_v6(host, enhanced) + .map(|result| result.map(|ip| ip.map(IpAddr::from))); + let v4 = self + .resolve_v4(host, enhanced) + .map(|result| result.map(|ip| ip.map(IpAddr::from))); + + let (first, remaining) = + future::select_ok(vec![v6.boxed(), v4.boxed()]).await?; + if first.is_some() { + first + } else { + future::select_all(remaining).await.0? + } + } else { + self.resolve_v4(host, enhanced).await?.map(IpAddr::from) + }; + + if let Some(ip) = resolved { + let ip = ip.to_string(); + self.store.set_host_to_ip(host, &ip).await; + self.store.set_ip_to_host(&ip, host).await; + } + Ok(resolved) + } + + async fn resolve_v4( + &self, + host: &str, + enhanced: bool, + ) -> anyhow::Result> { + Ok(self + .resolve_ip_by_type(host, RecordType::A, enhanced) + .await? + .and_then(|ip| match ip { + IpAddr::V4(ip) => Some(ip), + IpAddr::V6(_) => None, + })) + } + + async fn resolve_v6( + &self, + host: &str, + enhanced: bool, + ) -> anyhow::Result> { + if !self.ipv6() { + return Ok(None); + } + + Ok(self + .resolve_ip_by_type(host, RecordType::AAAA, enhanced) + .await? + .and_then(|ip| match ip { + IpAddr::V6(ip) => Some(ip), + IpAddr::V4(_) => None, + })) + } + + async fn reverse_lookup(&self, ip: std::net::IpAddr) -> Option { + if let Some(lru) = &self.reverse_lookup_cache + && let Some(cached) = lru.read().await.peek(&ip).cloned() + { + trace!(%ip, host = cached, "reverse lookup cache hit"); + return Some(cached); + } + + let Some(fake_dns) = &self.fake_dns else { + return None; + }; + + fake_dns.write().await.reverse_lookup(ip).await + } + + async fn cached_for(&self, ip: std::net::IpAddr) -> Option { + if let Some(lru) = &self.reverse_lookup_cache + && let Some(cached) = lru.read().await.peek(&ip).cloned() + { + return Some(cached); + } + self.store.get_fake_ip(&ip.to_string()).await + } +} + +impl EnhancedResolver { + async fn lookup_ip( + &self, + host: &str, + query_type: RecordType, + ) -> anyhow::Result> { + let mut message = Message::new(); + let mut query = hickory_proto::op::Query::new(); + let name = hickory_proto::rr::Name::from_str_relaxed(host) + .map_err(|_| anyhow::anyhow!("invalid domain: {host}"))? + .append_domain(&hickory_proto::rr::Name::root())?; + query.set_name(name); + query.set_query_type(query_type); + message.add_query(query); + message.set_recursion_desired(true); + + let response = self.exchange(&message).await?; + let ips = Self::ip_list_of_message(&response); + if ips.is_empty() { + Err(anyhow::anyhow!("no record for hostname: {host}")) + } else { + Ok(ips) + } + } + + async fn exchange_no_cache(&self, message: &Message) -> anyhow::Result { + let q = message + .query() + .ok_or_else(|| anyhow::anyhow!("invalid query message"))?; + + let query = async move { + if Self::is_ip_request(q) { + return self.ip_exchange(message).await; + } + + if let Some(matched) = self.match_policy(message) { + return Self::batch_exchange(matched, message).await; + } + + self.exchange_non_ip_query(message).await + }; + + let rv = query.await; + + if let Ok(msg) = &rv { + self.maybe_cache_response(q, msg).await; + } + + rv + } + + #[instrument(skip_all, level = "trace")] + async fn ip_exchange(&self, message: &Message) -> anyhow::Result { + let host = Self::domain_name_of_message(message) + .ok_or_else(|| anyhow::anyhow!("invalid query message"))?; + + let response = if let Some(policy) = self.match_policy(message) { + Self::batch_exchange(policy, message).await? + } else if self.should_only_query_fallback(message) { + Self::batch_exchange( + self.fallback.as_ref().ok_or_else(|| { + anyhow::anyhow!("no fallback resolver available") + })?, + message, + ) + .await? + } else { + self.exchange_with_main_then_fallback(message).await? + }; + + for ip in Self::ip_list_of_message(&response) { + self.save_reverse_lookup(ip, host.clone()).await; + } + + Ok(response) + } + + async fn maybe_cache_response( + &self, + query: &hickory_proto::op::Query, + response: &Message, + ) { + if query.query_type() == RecordType::TXT + && query.name().to_ascii().starts_with("_acme-challenge.") + { + return; + } + + if let Some(lru) = &self.lru_cache { + lru.write().await.insert_records( + query.clone(), + response.answers().iter().cloned(), + Instant::now(), + ); + } + } + + fn domain_name_of_message(message: &Message) -> Option { + message + .query() + .map(|query| query.name().to_ascii().trim_end_matches('.').to_owned()) + } + + fn is_ip_request(query: &hickory_proto::op::Query) -> bool { + query.query_class() == hickory_proto::rr::DNSClass::IN + && matches!(query.query_type(), RecordType::A | RecordType::AAAA) + } + + fn ip_list_of_message(message: &Message) -> Vec { + Self::ip_list_of_records(message.answers()) + } + + fn ip_list_of_records(records: &[Record]) -> Vec { + records + .iter() + .filter_map(|record| match record.data() { + RData::A(v4) => Some(IpAddr::V4(**v4)), + RData::AAAA(v6) => Some(IpAddr::V6(**v6)), + _ => None, + }) + .collect() + } + + fn should_only_query_fallback(&self, message: &Message) -> bool { + if let (Some(_), Some(fallback_domain_filters)) = + (&self.fallback, &self.fallback_domain_filters) + && let Some(domain) = Self::domain_name_of_message(message) + { + for filter in fallback_domain_filters.iter() { + if filter.apply(domain.as_str()) { + return true; + } + } + } + + false + } + + fn match_policy(&self, message: &Message) -> Option<&Vec> { + if let (Some(_fallback), Some(_fallback_domain_filters), Some(policy)) = + (&self.fallback, &self.fallback_domain_filters, &self.policy) + && let Some(host) = Self::domain_name_of_message(message) + { + return policy + .search(&host.trim_end_matches('.').to_ascii_lowercase()) + .and_then(|node| node.get_data()); + } + + None + } + + async fn exchange_with_main_then_fallback( + &self, + message: &Message, + ) -> anyhow::Result { + if self.main.is_empty() { + return Self::batch_exchange( + self.fallback + .as_ref() + .ok_or_else(|| anyhow::anyhow!("no resolver available"))?, + message, + ) + .await; + } + + let main_result = Self::batch_exchange(&self.main, message).await; + + if self.fallback.is_none() { + return main_result; + } + + if let Ok(response) = main_result { + let ips = Self::ip_list_of_message(&response); + if ips.first().is_some_and(|ip| self.should_ip_fallback(ip)) { + return Self::batch_exchange( + self.fallback.as_ref().expect("checked above"), + message, + ) + .await; + } + return Ok(response); + } + + Self::batch_exchange(self.fallback.as_ref().expect("checked above"), message) + .await + } + + async fn exchange_non_ip_query( + &self, + message: &Message, + ) -> anyhow::Result { + if !self.main.is_empty() { + if let Ok(response) = Self::batch_exchange(&self.main, message).await { + return Ok(response); + } + } + + if let Some(fallback) = &self.fallback { + return Self::batch_exchange(fallback, message).await; + } + + Err(anyhow::anyhow!("no resolver available for dns query")) + } + + fn should_ip_fallback(&self, ip: &IpAddr) -> bool { + self.fallback_ip_filters + .as_ref() + .is_some_and(|filters| filters.iter().any(|filter| filter.apply(ip))) + } + + #[instrument(skip(message), level = "trace")] + async fn batch_exchange( + resolvers: &Vec, + message: &Message, + ) -> anyhow::Result { + if resolvers.is_empty() { + return Err(anyhow::anyhow!("no resolver available")); + } + + let mut queries = Vec::new(); + for resolver in resolvers { + queries.push( + async move { + resolver + .exchange(message) + .inspect_err(|err| { + error!(err = ?err, "resolve error"); + }) + .await + } + .boxed(), + ); + } + + let timeout = tokio::time::sleep(Duration::from_secs(10)); + tokio::select! { + result = future::select_ok(queries) => match result { + Ok((response, _)) => Ok(response), + Err(err) => Err(err), + }, + _ = timeout => Err(anyhow::anyhow!("dns query timeout")), + } + } + + async fn resolve_ip_by_type( + &self, + host: &str, + query_type: RecordType, + enhanced: bool, + ) -> anyhow::Result> { + let normalized_host = host.trim_end_matches('.').to_ascii_lowercase(); + + if enhanced + && let Some(hosts) = &self.hosts + && let Some(node) = hosts.search(&normalized_host) + && let Some(ip) = node.get_data() + { + return Ok(match query_type { + RecordType::A if ip.is_ipv4() => Some(*ip), + RecordType::AAAA if ip.is_ipv6() => Some(*ip), + _ => None, + }); + } + + if let Ok(ip) = host.parse::() { + return Ok(match query_type { + RecordType::A if ip.is_ipv4() => Some(ip), + RecordType::AAAA if ip.is_ipv6() => Some(ip), + _ => None, + }); + } + + if enhanced + && query_type == RecordType::A + && let Some(fake_dns) = &self.fake_dns + { + let mut fake_dns = fake_dns.write().await; + if !fake_dns.should_skip(host) { + return Ok(Some(fake_dns.lookup(host).await)); + } + } + + match self.lookup_ip(host, query_type).await { + Ok(ips) => Ok(ips.into_iter().find(|ip| match query_type { + RecordType::A => ip.is_ipv4(), + RecordType::AAAA => self.ipv6() && ip.is_ipv6(), + _ => false, + })), + Err(_) => { + let response = tokio::net::lookup_host(format!("{host}:0")).await?; + Ok(response.map(|addr| addr.ip()).find(|ip| match query_type { + RecordType::A => ip.is_ipv4(), + RecordType::AAAA => self.ipv6() && ip.is_ipv6(), + _ => false, + })) + } + } + } + + async fn save_reverse_lookup(&self, ip: IpAddr, host: String) { + trace!(%ip, host = %host, "reverse lookup cache insert"); + if let Some(lru) = &self.reverse_lookup_cache { + lru.write().await.insert(ip, host); + } + } +} + +fn build_fake_dns( + cfg: &DNSConfig, + store: ThreadSafeCacheFile, +) -> Result, crate::Error> { + match cfg.enhance_mode { + DNSMode::FakeIp => { + let store: Box = if cfg.store_fake_ip + { + Box::new(FileStore::new(store)) + } else { + Box::new(InMemStore::new(1000)) + }; + + Ok(Some(Arc::new(RwLock::new(FakeDns::new(FakeDnsOpts { + ipnet: cfg.fake_ip_range, + skipped_hostnames: build_skipped_hostnames_trie(&cfg.fake_ip_filter), + store, + })?)))) + } + DNSMode::RedirHost => { + warn!("dns redir-host is not supported and will not do anything"); + Ok(None) + } + DNSMode::Normal => Ok(None), + } +} + +fn build_fallback_filters( + cfg: &DNSConfig, + mmdb: Option, +) -> ( + Option>>, + Option>>, +) { + let mut domain_filters: Vec> = Vec::new(); + let mut ip_filters: Vec> = Vec::new(); + + if !cfg.fallback_filter.domain.is_empty() { + domain_filters + .push(Box::new(DomainFilter::new(&cfg.fallback_filter.domain))); + } + + if cfg.fallback_filter.geo_ip || !cfg.fallback_filter.ip_cidr.is_empty() { + if cfg.fallback_filter.geo_ip { + ip_filters.push(Box::new(GeoIPFilter::new( + &cfg.fallback_filter.geo_ip_code, + mmdb, + ))); + } + + for cidr in &cfg.fallback_filter.ip_cidr { + ip_filters.push(Box::new(IPNetFilter::new(*cidr))); + } + } + + ( + (!domain_filters.is_empty()).then_some(domain_filters), + (!ip_filters.is_empty()).then_some(ip_filters), + ) +} + +async fn build_policy_resolvers( + cfg: &DNSConfig, + resolver: Option>, + outbounds: HashMap>, +) -> Option>> { + let mut out = StringTrie::new(); + let mut has_entries = false; + for (domain, nameserver) in &cfg.nameserver_policy { + let resolvers = make_clients( + std::slice::from_ref(nameserver), + resolver.clone(), + outbounds.clone(), + cfg.edns_client_subnet.clone(), + cfg.fw_mark, + cfg.ipv6, + ) + .await; + if !resolvers.is_empty() { + has_entries = true; + out.insert(domain, Arc::new(resolvers)); + } + } + has_entries.then_some(out) +} + +fn build_hosts_trie(hosts: &HashMap) -> Option> { + let mut out = StringTrie::new(); + let mut has_entries = false; + + for (host, ip) in hosts { + has_entries = true; + out.insert(host, Arc::new(*ip)); + } + + has_entries.then_some(out) +} + +fn build_skipped_hostnames_trie(hosts: &[String]) -> Option> { + let mut out = StringTrie::new(); + let mut has_entries = false; + + for host in hosts { + let host = host.trim_end_matches('.').to_ascii_lowercase(); + has_entries = true; + out.insert(&host, Arc::new(true)); + } + + has_entries.then_some(out) +} diff --git a/clash-lib/src/app/dns/resolver/mod.rs b/clash-lib/src/app/dns/resolver/mod.rs index aa9c6f22..f339ba4d 100644 --- a/clash-lib/src/app/dns/resolver/mod.rs +++ b/clash-lib/src/app/dns/resolver/mod.rs @@ -1,3 +1,5 @@ +mod enhanced; + use crate::{ app::{ dns::{DNSConfig, ThreadSafeDNSResolver}, @@ -14,6 +16,7 @@ use std::{collections::HashMap, sync::Arc}; #[path = "system.rs"] mod system; +pub use enhanced::EnhancedResolver; pub use system::SystemResolver; pub async fn new( @@ -25,8 +28,7 @@ pub async fn new( if cfg.enable { match store { Some(store) => { - /* Arc::new(EnhancedResolver::new(cfg, store, mmdb, outbounds).await) */ - todo!() + Arc::new(EnhancedResolver::new(cfg, store, mmdb, outbounds).await) } _ => print_and_exit!("enhanced resolver requires cache store"), } diff --git a/clash-lib/src/app/dns/resolver/system.rs b/clash-lib/src/app/dns/resolver/system.rs index 7e30c5d9..359e0d76 100644 --- a/clash-lib/src/app/dns/resolver/system.rs +++ b/clash-lib/src/app/dns/resolver/system.rs @@ -1,7 +1,4 @@ -use std::{ - net::{IpAddr, Ipv4Addr, Ipv6Addr}, - sync::atomic::{AtomicBool, Ordering}, -}; +use std::sync::atomic::{AtomicBool, Ordering}; use async_trait::async_trait; use rand::seq::IteratorRandom; @@ -36,10 +33,6 @@ impl ClashResolver for SystemResolver { )) } - fn ipv6(&self) -> bool { - self.ipv6.load(std::sync::atomic::Ordering::Relaxed) - } - async fn resolve( &self, host: &str, @@ -63,4 +56,32 @@ impl ClashResolver for SystemResolver { // todo: use the first address for now, we may want to randomize it later Ok(response.into_iter().choose(&mut rand::rng())) } + + fn ipv6(&self) -> bool { + self.ipv6.load(Ordering::Relaxed) + } + + fn set_ipv6(&self, enable: bool) { + self.ipv6.store(enable, Ordering::Relaxed); + } + + fn kind(&self) -> ResolverKind { + ResolverKind::System + } + + fn fake_ip_enabled(&self) -> bool { + false + } + + async fn reverse_lookup(&self, _: std::net::IpAddr) -> Option { + None + } + + async fn is_fake_ip(&self, _: std::net::IpAddr) -> bool { + false + } + + async fn cached_for(&self, _: std::net::IpAddr) -> Option { + None + } } diff --git a/clash-lib/src/app/dns/server/handler.rs b/clash-lib/src/app/dns/server/handler.rs index ad522821..3588a093 100644 --- a/clash-lib/src/app/dns/server/handler.rs +++ b/clash-lib/src/app/dns/server/handler.rs @@ -8,13 +8,14 @@ pub async fn exchange_with_resolver( req: &Message, _enhanced: bool, ) -> Result { - tracing::debug!("todo: enhanced dns resolve: {}", _enhanced); + tracing::debug!("dns resolve request, enhanced={}", _enhanced); match resolver.exchange(req).await { Ok(m) => Ok(m), Err(e) => { debug!("dns resolve error: {}", e); - todo!() - // Err(chimera_dns::DNSError::QueryFailed(e.to_string())) + Err(chimera_dns::DNSError::Io(std::io::Error::other( + e.to_string(), + ))) } } } diff --git a/clash-lib/src/app/dns/server/mod.rs b/clash-lib/src/app/dns/server/mod.rs index 0c6e3782..4951d578 100644 --- a/clash-lib/src/app/dns/server/mod.rs +++ b/clash-lib/src/app/dns/server/mod.rs @@ -8,7 +8,7 @@ use crate::{Runner, app::dns::ThreadSafeDNSResolver}; mod handler; pub use handler::exchange_with_resolver; -static DEFAULT_DNS_SERVER_TTL: u32 = 60; +pub(crate) static DEFAULT_DNS_SERVER_TTL: u32 = 60; struct DnsMessageExchanger { resolver: ThreadSafeDNSResolver, @@ -20,9 +20,12 @@ impl chimera_dns::DnsMessageExchanger for DnsMessageExchanger { self.resolver.ipv6() } - /* async fn exchange(&self, message: &Message) -> Result { + async fn exchange( + &self, + message: &Message, + ) -> Result { exchange_with_resolver(&self.resolver, message, true).await - } */ + } } pub async fn get_dns_listener( @@ -38,8 +41,7 @@ pub async fn get_dns_listener( Ok(()) => Ok(()), Err(err) => { error!("dns listener error: {}", err); - todo!() - // Err(err.into()) + Err(crate::Error::Io(std::io::Error::other(err.to_string()))) } } })), diff --git a/clash-lib/src/app/logging.rs b/clash-lib/src/app/logging.rs index 6708056f..4310c14c 100644 --- a/clash-lib/src/app/logging.rs +++ b/clash-lib/src/app/logging.rs @@ -1,11 +1,15 @@ -use std::{io::IsTerminal, sync::Once}; +use std::{ + io::IsTerminal, + sync::{Once, OnceLock}, +}; use anyhow::anyhow; use serde::Serialize; use tokio::sync::broadcast::Sender; use tracing_log::LogTracer; use tracing_subscriber::{ - EnvFilter, Layer, filter::filter_fn, fmt::time::LocalTime, prelude::*, + EnvFilter, Layer, Registry, filter::filter_fn, fmt::time::LocalTime, prelude::*, + reload, }; use crate::config::def::LogLevel; @@ -61,6 +65,8 @@ struct LoggingGuard { static SETUP_LOGGING: Once = Once::new(); static mut LOGGING_GUARD: Option = None; +static LOG_FILTER_RELOAD: OnceLock> = + OnceLock::new(); pub fn setup_logging( level: LogLevel, @@ -76,11 +82,17 @@ pub fn setup_logging( have been initialized" ); }); - LOGGING_GUARD = setup_logging_inner(level, collector, cwd, log_file) - .unwrap_or_else(|e| { - eprintln!("Failed to setup logging: {e}"); - None - }); + let (guard, reload_handle) = setup_logging_inner( + level, collector, cwd, log_file, + ) + .unwrap_or_else(|e| { + eprintln!("Failed to setup logging: {e}"); + (None, None) + }); + LOGGING_GUARD = guard; + if let Some(reload_handle) = reload_handle { + _ = LOG_FILTER_RELOAD.set(reload_handle); + } }); } } @@ -90,7 +102,10 @@ fn setup_logging_inner( collector: EventCollector, cwd: &str, log_file: Option, -) -> anyhow::Result> { +) -> anyhow::Result<( + Option, + Option>, +)> { let default_log_level = format!("warn,clash={level}"); let filter = EnvFilter::try_from_default_env() .inspect(|f| { @@ -103,6 +118,7 @@ fn setup_logging_inner( } }) .unwrap_or(EnvFilter::new(default_log_level)); + let (filter_layer, reload_handle) = reload::Layer::new(filter); let (appender, guard) = if let Some(log_file) = log_file { let path_buf = std::path::PathBuf::from(&log_file); @@ -161,7 +177,7 @@ fn setup_logging_inner( #[cfg(feature = "tracing")] { subscriber - .with(filter) + .with(filter_layer) .with(collector.with_filter(exclude.clone())) .with(log_to_file_layer) .with(log_stdout_layer) @@ -169,7 +185,7 @@ fn setup_logging_inner( #[cfg(not(feature = "tracing"))] { subscriber - .with(filter) // Global filter + .with(filter_layer) .with(collector.with_filter(exclude.clone())) .with(log_to_file_layer) .with(log_stdout_layer) @@ -179,9 +195,23 @@ fn setup_logging_inner( tracing::subscriber::set_global_default(subscriber) .map_err(|x| anyhow!("setup logging error: {}", x))?; - Ok(Some(LoggingGuard { - _file_appender: guard, - })) + Ok(( + Some(LoggingGuard { + _file_appender: guard, + }), + Some(reload_handle), + )) +} + +pub fn set_log_level(level: LogLevel) -> anyhow::Result<()> { + let default_log_level = format!("warn,clash={level}"); + let filter = EnvFilter::new(default_log_level); + let handle = LOG_FILTER_RELOAD + .get() + .ok_or_else(|| anyhow!("logging reload handle not initialized"))?; + handle + .reload(filter) + .map_err(|e| anyhow!("failed to reload log level: {e}")) } struct EventVisitor<'a>(&'a mut Vec); diff --git a/clash-lib/src/app/profile/mod.rs b/clash-lib/src/app/profile/mod.rs index 3dd55c0c..53d48f99 100644 --- a/clash-lib/src/app/profile/mod.rs +++ b/clash-lib/src/app/profile/mod.rs @@ -77,6 +77,22 @@ impl ThreadSafeCacheFile { None } } + + pub async fn set_ip_to_host(&self, ip: &str, host: &str) { + self.0.write().await.set_ip_to_host(ip, host); + } + + pub async fn set_host_to_ip(&self, host: &str, ip: &str) { + self.0.write().await.set_host_to_ip(host, ip); + } + + pub async fn get_fake_ip(&self, ip_or_host: &str) -> Option { + self.0.read().await.get_fake_ip(ip_or_host) + } + + pub async fn delete_fake_ip_pair(&self, ip: &str, host: &str) { + self.0.write().await.delete_fake_ip_pair(ip, host); + } } struct CacheFile { @@ -128,4 +144,25 @@ impl CacheFile { .selected .insert(group.to_string(), server.to_string()); } + + pub fn set_ip_to_host(&mut self, ip: &str, host: &str) { + self.db.ip_to_host.insert(ip.to_string(), host.to_string()); + } + + pub fn set_host_to_ip(&mut self, host: &str, ip: &str) { + self.db.host_to_ip.insert(host.to_string(), ip.to_string()); + } + + pub fn get_fake_ip(&self, ip_or_host: &str) -> Option { + self.db + .ip_to_host + .get(ip_or_host) + .or_else(|| self.db.host_to_ip.get(ip_or_host)) + .cloned() + } + + pub fn delete_fake_ip_pair(&mut self, ip: &str, host: &str) { + self.db.ip_to_host.remove(ip); + self.db.host_to_ip.remove(host); + } } diff --git a/clash-lib/src/common/http/client.rs b/clash-lib/src/common/http/client.rs index 9b4da4b0..e7b5d16f 100644 --- a/clash-lib/src/common/http/client.rs +++ b/clash-lib/src/common/http/client.rs @@ -2,12 +2,13 @@ use std::{collections::HashMap, io, sync::Arc}; use futures::FutureExt; use hyper_util::rt::TokioIo; -use tokio::net::TcpStream; use tracing::{trace, warn}; use crate::{ - app::dns::ThreadSafeDNSResolver, config::internal::proxy::PROXY_DIRECT, - proxy::AnyOutboundHandler, + app::dns::ThreadSafeDNSResolver, + config::internal::proxy::PROXY_DIRECT, + proxy::{AnyOutboundHandler, direct}, + session::Session, }; #[cfg(feature = "tls")] @@ -59,15 +60,6 @@ impl HttpClient { }) } - async fn connect_tcp(&self, host: &str, port: u16) -> io::Result { - let connect = TcpStream::connect((host, port)); - tokio::time::timeout(self.timeout, connect) - .await - .map_err(|_| { - io::Error::new(io::ErrorKind::TimedOut, "tcp connect timeout") - })? - } - pub async fn request( &self, mut req: http::Request, @@ -112,16 +104,36 @@ impl HttpClient { } let req_ext = req.extensions().get::(); - let outbound_name = req_ext - .and_then(|ext| ext.outbound.as_deref()) - .unwrap_or(PROXY_DIRECT); - if let Some(outbounds) = self.outbounds.as_ref() { - if let Some(outbound) = outbounds.get(outbound_name) { - trace!(outbound = %outbound.name(), "using outbound for http client"); - } - } - - let stream = self.connect_tcp(&host, port).await?; + let outbound = req_ext + .and_then(|ext| ext.outbound.clone()) + .as_ref() + .and_then(|name| { + self.outbounds + .as_ref() + .and_then(|outbounds| outbounds.get(name).cloned()) + }) + .unwrap_or_else(|| Arc::new(direct::Handler::new(PROXY_DIRECT)) as _); + + trace!(outbound = %outbound.name(), "using outbound for http client"); + let session = Session { + network: crate::session::Network::Tcp, + typ: crate::session::Type::Ignore, + destination: crate::session::SocksAddr::Domain(host.clone(), port), + ..Default::default() + }; + let stream = tokio::time::timeout( + self.timeout, + outbound.connect_stream(&session, self.dns_resolver.clone()), + ) + .await + .map_err(|_| io::Error::new(io::ErrorKind::TimedOut, "tcp connect timeout"))? + .inspect_err(|error| { + warn!( + outbound = outbound.name(), + err = ?error, + "http client outbound connect failed" + ); + })?; #[cfg(not(feature = "tls"))] if uri.scheme() == Some(&http::uri::Scheme::HTTPS) { diff --git a/clash-lib/src/common/mod.rs b/clash-lib/src/common/mod.rs index 41d06cfc..c1da3c29 100644 --- a/clash-lib/src/common/mod.rs +++ b/clash-lib/src/common/mod.rs @@ -6,4 +6,5 @@ pub mod io; pub mod mmdb; #[cfg(feature = "tls")] pub mod tls; +pub mod trie; pub mod utils; diff --git a/clash-lib/src/common/trie.rs b/clash-lib/src/common/trie.rs new file mode 100644 index 00000000..fcd4825e --- /dev/null +++ b/clash-lib/src/common/trie.rs @@ -0,0 +1,271 @@ +use std::{collections::HashMap, sync::Arc}; + +static DOMAIN_STEP: &str = "."; +static COMPLEX_WILDCARD: &str = "+"; +static DOT_WILDCARD: &str = ""; +static WILDCARD: &str = "*"; + +pub struct Node { + children: HashMap>, + data: Option>, +} + +impl Default for Node { + fn default() -> Self { + Self::new() + } +} + +impl Node { + pub fn new() -> Self { + Self { + children: HashMap::new(), + data: None, + } + } + + pub fn get_data(&self) -> Option<&T> { + self.data.as_deref() + } + + pub fn get_child(&self, key: &str) -> Option<&Self> { + self.children.get(key) + } + + pub fn get_child_mut(&mut self, key: &str) -> Option<&mut Self> { + self.children.get_mut(key) + } + + pub fn has_child(&self, key: &str) -> bool { + self.get_child(key).is_some() + } + + pub fn add_child(&mut self, key: &str, child: Node) { + self.children.insert(key.to_string(), child); + } +} + +pub struct StringTrie { + root: Node, +} + +impl Default for StringTrie { + fn default() -> Self { + Self::new() + } +} + +impl StringTrie { + pub fn new() -> Self { + Self { root: Node::new() } + } + + pub fn insert(&mut self, domain: &str, data: Arc) -> bool { + let (parts, valid) = valid_and_split_domain(domain); + if !valid { + return false; + } + + let mut parts = parts.unwrap(); + match parts[0] { + part if part == COMPLEX_WILDCARD => { + self.insert_inner(&parts[1..], data.clone()); + parts[0] = DOT_WILDCARD; + self.insert_inner(&parts, data); + } + _ => self.insert_inner(&parts, data), + } + + true + } + + pub fn search(&self, domain: &str) -> Option<&Node> { + let (parts, valid) = valid_and_split_domain(domain); + if !valid { + return None; + } + + let parts = parts.unwrap(); + if parts[0].is_empty() { + return None; + } + + if let Some(node) = Self::search_inner(&self.root, parts) + && node.data.is_some() + { + return Some(node); + } + + None + } + + fn insert_inner(&mut self, parts: &[&str], data: Arc) { + let mut node = &mut self.root; + + for index in (0..parts.len()).rev() { + let part = parts[index]; + if !node.has_child(part) { + node.add_child(part, Node::new()); + } + + node = node.get_child_mut(part).expect("child just inserted"); + } + + node.data = Some(data); + } + + fn search_inner<'a>(node: &'a Node, parts: Vec<&str>) -> Option<&'a Node> { + if parts.is_empty() { + return Some(node); + } + + if let Some(child) = node.get_child(parts.last().expect("non-empty parts")) + && let Some(found) = + Self::search_inner(child, parts[..parts.len() - 1].into()) + && found.data.is_some() + { + return Some(found); + } + + if let Some(child) = node.get_child(WILDCARD) + && let Some(found) = + Self::search_inner(child, parts[..parts.len() - 1].into()) + && found.data.is_some() + { + return Some(found); + } + + node.get_child(DOT_WILDCARD) + } +} + +pub fn valid_and_split_domain(domain: &str) -> (Option>, bool) { + if !domain.is_empty() && domain.ends_with('.') { + return (None, false); + } + + let parts: Vec<&str> = domain.split(DOMAIN_STEP).collect(); + if parts.len() == 1 { + if parts[0].is_empty() { + return (None, false); + } + return (Some(parts), true); + } + + for part in parts.iter().skip(1) { + if part.is_empty() { + return (None, false); + } + } + + (Some(parts), true) +} + +#[cfg(test)] +mod tests { + use std::{net::Ipv4Addr, sync::Arc}; + + use super::StringTrie; + + static LOCAL_IP: Ipv4Addr = Ipv4Addr::new(127, 0, 0, 1); + + #[test] + fn test_basic() { + let mut tree = StringTrie::new(); + + let domains = ["example.com", "google.com", "localhost"]; + + for domain in domains { + assert!(tree.insert(domain, Arc::new(LOCAL_IP))); + } + + let node = tree + .search("example.com") + .expect("should match example.com"); + assert_eq!(node.get_data(), Some(&LOCAL_IP)); + assert!(!tree.insert("", Arc::new(LOCAL_IP))); + assert!(tree.search("").is_none()); + assert!(tree.search("localhost").is_some()); + assert!(tree.search("www.google.com").is_none()); + } + + #[test] + fn test_wildcard() { + let mut tree = StringTrie::new(); + + let domains = [ + "*.example.com", + "sub.*.example.com", + "*.dev", + ".org", + ".example.net", + ".apple.*", + "+.foo.com", + "+.stun.*.*", + "+.stun.*.*.*", + "+.stun.*.*.*.*", + "stun.l.google.com", + ]; + + for domain in domains { + assert!(tree.insert(domain, Arc::new(LOCAL_IP))); + } + + assert!(tree.search("sub.example.com").is_some()); + assert!(tree.search("sub.foo.example.com").is_some()); + assert!(tree.search("test.org").is_some()); + assert!(tree.search("test.example.net").is_some()); + assert!(tree.search("test.apple.com").is_some()); + assert!(tree.search("foo.com").is_some()); + assert!(tree.search("global.stun.website.com").is_some()); + + assert!(tree.search("foo.sub.example.com").is_none()); + assert!(tree.search("foo.example.dev").is_none()); + assert!(tree.search("example.com").is_none()); + } + + #[test] + fn test_priority() { + let mut tree = StringTrie::new(); + + let domains = [".dev", "example.dev", "*.example.dev", "test.example.dev"]; + + for (index, domain) in domains.iter().enumerate() { + assert!(tree.insert(domain, Arc::new(index))); + } + + let assert_match = |domain: &str| -> Arc { + tree.search(domain) + .expect("domain should match") + .data + .clone() + .expect("node should have data") + }; + + assert_eq!(assert_match("test.dev"), Arc::new(0)); + assert_eq!(assert_match("foo.bar.dev"), Arc::new(0)); + assert_eq!(assert_match("example.dev"), Arc::new(1)); + assert_eq!(assert_match("foo.example.dev"), Arc::new(2)); + assert_eq!(assert_match("test.example.dev"), Arc::new(3)); + } + + #[test] + fn test_boundary() { + let mut tree = StringTrie::new(); + + assert!(tree.insert("*.dev", Arc::new(LOCAL_IP))); + assert!(!tree.insert(".", Arc::new(LOCAL_IP))); + assert!(!tree.insert("..dev", Arc::new(LOCAL_IP))); + assert!(tree.search("dev").is_none()); + } + + #[test] + fn test_wildcard_boundary() { + let mut tree = StringTrie::new(); + + assert!(tree.insert("+.*", Arc::new(LOCAL_IP))); + assert!(tree.insert("stun.*.*.*", Arc::new(LOCAL_IP))); + + assert!(tree.search("example.com").is_some()); + } +} diff --git a/clash-lib/src/config/def.rs b/clash-lib/src/config/def.rs index 326fb400..b29d2281 100644 --- a/clash-lib/src/config/def.rs +++ b/clash-lib/src/config/def.rs @@ -99,6 +99,8 @@ pub struct Config { /// - "https://example.com" #[serde(rename = "cors-allow-origins")] pub cors_allow_origins: Option>, + #[serde(default)] + pub hosts: HashMap, #[cfg_attr(not(unix), serde(alias = "external-controller-pipe"))] #[cfg_attr(unix, serde(alias = "external-controller-unix"))] pub external_controller_ipc: Option, @@ -213,8 +215,32 @@ impl Display for LogLevel { #[serde(untagged)] pub enum DNSListen { Udp(String), - // todo - // Multiple(HashMap), + Multiple(HashMap), +} + +#[derive(Serialize, Deserialize, Default, Clone, Debug, PartialEq, Eq)] +#[serde(rename_all = "kebab-case")] +pub enum DNSMode { + #[default] + Normal, + FakeIp, + RedirHost, +} + +#[derive(Serialize, Deserialize, Clone, Educe)] +#[serde(default)] +#[serde(rename_all = "kebab-case")] +#[educe(Default)] +pub struct FallbackFilter { + #[serde(rename = "geoip")] + #[educe(Default = true)] + pub geo_ip: bool, + #[serde(rename = "geoip-code")] + #[educe(Default = "CN")] + pub geo_ip_code: String, + #[serde(rename = "ipcidr")] + pub ip_cidr: Vec, + pub domain: Vec, } /// DNS client/server settings @@ -245,6 +271,24 @@ pub enum DNSListen { #[serde(rename_all = "kebab-case", default)] #[educe(Default)] pub struct DNS { + /// Whether to use fake IP addresses + pub enhanced_mode: DNSMode, + /// DNS upstream servers + pub nameserver: Vec, + /// Fallback DNS upstream servers + pub fallback: Vec, + /// Fallback DNS filter + pub fallback_filter: FallbackFilter, + /// Default nameservers used for resolving DNS upstream hostnames later + pub default_nameserver: Vec, + pub edns_client_subnet: Option, + /// Lookup domains via specific nameservers + pub nameserver_policy: HashMap, + /// Fake IP addresses pool CIDR + #[educe(Default = "198.18.0.1/16")] + pub fake_ip_range: String, + /// Fake IP addresses filter + pub fake_ip_filter: Vec, /// Enable IPv6 DNS responses (AAAA) pub ipv6: bool, /// DNS server listening address. If not present, the DNS server will be @@ -255,6 +299,13 @@ pub struct DNS { pub enable: bool, } +#[derive(Serialize, Deserialize, Default, Clone, Debug, PartialEq, Eq)] +#[serde(rename_all = "kebab-case")] +pub struct EdnsClientSubnet { + pub ipv4: Option, + pub ipv6: Option, +} + #[derive(Serialize, Deserialize)] #[serde(default)] #[serde(rename_all = "kebab-case")] diff --git a/clash-lib/src/lib.rs b/clash-lib/src/lib.rs index 9451469c..35e24966 100644 --- a/clash-lib/src/lib.rs +++ b/clash-lib/src/lib.rs @@ -52,6 +52,10 @@ mod session; pub use session::Session; +pub use proxy::utils::{ + SocketProtector, clear_socket_protector, set_socket_protector, +}; + #[derive(Error, Debug)] pub enum Error { #[error(transparent)] diff --git a/clash-lib/src/proxy/transport/xhttp/mod.rs b/clash-lib/src/proxy/transport/xhttp/mod.rs index b2f11df7..60a10eaa 100644 --- a/clash-lib/src/proxy/transport/xhttp/mod.rs +++ b/clash-lib/src/proxy/transport/xhttp/mod.rs @@ -322,9 +322,7 @@ fn build_request( fn build_xhttp_padding_referer(uri: &str) -> String { let separator = if uri.contains('?') { '&' } else { '?' }; let padding = "X".repeat(DEFAULT_XHTTP_PADDING_BYTES); - format!( - "{uri}{separator}{DEFAULT_XHTTP_PADDING_QUERY_KEY}={padding}" - ) + format!("{uri}{separator}{DEFAULT_XHTTP_PADDING_QUERY_KEY}={padding}") } #[async_trait] @@ -552,8 +550,7 @@ async fn forward_response_body( #[cfg(test)] mod tests { use super::{ - Client, XhttpDownloadConfig, XhttpMode, XhttpSecurity, - build_request, + Client, XhttpDownloadConfig, XhttpMode, XhttpSecurity, build_request, }; use crate::proxy::transport::Transport; use bytes::Bytes; diff --git a/clash-lib/src/proxy/tun/datagram.rs b/clash-lib/src/proxy/tun/datagram.rs new file mode 100644 index 00000000..3e2c9338 --- /dev/null +++ b/clash-lib/src/proxy/tun/datagram.rs @@ -0,0 +1,71 @@ +use crate::app::dns::{ThreadSafeDNSResolver, exchange_with_resolver}; +use tracing::{debug, trace, warn}; + +pub(crate) async fn handle_inbound_datagram( + socket: watfaq_netstack::UdpSocket, + resolver: ThreadSafeDNSResolver, + dns_hijack: bool, +) { + let (mut rx, mut tx) = socket.split(); + + debug!("tun UDP ready"); + + while let Some(watfaq_netstack::UdpPacket { + data, + local_addr, + remote_addr, + }) = rx.recv().await + { + if remote_addr.ip().is_multicast() { + continue; + } + + if dns_hijack && remote_addr.port() == 53 { + trace!( + "hijack dns request: {} -> {} ({} bytes)", + local_addr, + remote_addr, + data.data().len() + ); + + let msg = match hickory_proto::op::Message::from_vec(data.data()) { + Ok(msg) => msg, + Err(error) => { + warn!("failed to parse dns packet: {}", error); + continue; + } + }; + + let mut resp = match exchange_with_resolver(&resolver, &msg, true).await + { + Ok(resp) => resp, + Err(error) => { + warn!("failed to exchange dns message: {}", error); + continue; + } + }; + + resp.set_id(msg.id()); + + let data = match resp.to_vec() { + Ok(data) => data, + Err(error) => { + warn!("failed to serialize dns response: {}", error); + continue; + } + }; + + if let Err(error) = tx.send((data, remote_addr, local_addr).into()).await + { + warn!("failed to send dns response to netstack: {}", error); + } + + continue; + } + + trace!( + "dropping tun UDP packet: {} -> {} (dns_hijack={})", + local_addr, remote_addr, dns_hijack + ); + } +} diff --git a/clash-lib/src/proxy/tun/inbound.rs b/clash-lib/src/proxy/tun/inbound.rs index fc3e9881..c848c52e 100644 --- a/clash-lib/src/proxy/tun/inbound.rs +++ b/clash-lib/src/proxy/tun/inbound.rs @@ -8,7 +8,9 @@ use crate::{ Error, Result, Runner, app::{dispatcher::Dispatcher, dns::ThreadSafeDNSResolver}, config::internal::config::TunConfig, - proxy::tun::{routes, stream::handle_inbound_stream}, + proxy::tun::{ + datagram::handle_inbound_datagram, routes, stream::handle_inbound_stream, + }, }; #[derive(Default)] @@ -39,7 +41,7 @@ impl Drop for RouteCleanupGuard { pub fn get_runner( cfg: TunConfig, dispatcher: Arc, - _resolver: ThreadSafeDNSResolver, + resolver: ThreadSafeDNSResolver, ) -> Result> { if !cfg.enable { trace!("tun is disabled"); @@ -153,10 +155,11 @@ pub fn get_runner( } }; - let (stack, mut tcp_listener, _udp_socket) = watfaq_netstack::NetStack::new(); + let (stack, mut tcp_listener, udp_socket) = watfaq_netstack::NetStack::new(); Ok(Some(Box::pin(async move { let so_mark = cfg.so_mark; + let dns_hijack = cfg.dns_hijack; let _route_cleanup_guard = RouteCleanupGuard::new(cfg); let framed = tun_rs::async_framed::DeviceFramed::new( @@ -230,6 +233,12 @@ pub fn get_runner( Err(Error::Operation("tun stopped unexpectedly 2".to_string())) })); + futs.push(Box::pin(async move { + handle_inbound_datagram(udp_socket, resolver, dns_hijack).await; + + Err(Error::Operation("tun stopped unexpectedly 3".to_string())) + })); + futures::future::select_all(futs).await.0.map_err(|e| { error!("tun error: {}. stopped", e); e diff --git a/clash-lib/src/proxy/tun/mod.rs b/clash-lib/src/proxy/tun/mod.rs index 422f993b..f840ddca 100644 --- a/clash-lib/src/proxy/tun/mod.rs +++ b/clash-lib/src/proxy/tun/mod.rs @@ -1,3 +1,4 @@ +mod datagram; pub mod inbound; pub use inbound::get_runner as get_tun_runner; mod routes; diff --git a/clash-lib/src/proxy/utils/mod.rs b/clash-lib/src/proxy/utils/mod.rs index 42595b6c..1359e42f 100644 --- a/clash-lib/src/proxy/utils/mod.rs +++ b/clash-lib/src/proxy/utils/mod.rs @@ -1,5 +1,7 @@ pub mod socket_helpers; +#[allow(unused_imports)] +pub use platform::{SocketProtector, clear_socket_protector, set_socket_protector}; pub use socket_helpers::*; /// 2 diff --git a/clash-lib/src/proxy/utils/platform/android.rs b/clash-lib/src/proxy/utils/platform/android.rs new file mode 100644 index 00000000..70162a7a --- /dev/null +++ b/clash-lib/src/proxy/utils/platform/android.rs @@ -0,0 +1,55 @@ +use std::{ + io, + os::fd::AsRawFd, + sync::{Arc, LazyLock, RwLock}, +}; + +use tracing::trace; + +use crate::app::net::OutboundInterface; + +pub trait SocketProtector: Send + Sync { + fn protect_socket_fd(&self, fd: i32) -> io::Result<()>; +} + +static SOCKET_PROTECTOR: LazyLock>>> = + LazyLock::new(|| RwLock::new(None)); + +pub fn set_socket_protector(protector: Arc) { + if let Ok(mut guard) = SOCKET_PROTECTOR.write() { + *guard = Some(protector); + } +} + +pub fn clear_socket_protector() { + if let Ok(mut guard) = SOCKET_PROTECTOR.write() { + *guard = None; + } +} + +pub(crate) fn maybe_protect_socket(socket: &socket2::Socket) -> io::Result<()> { + let protector = SOCKET_PROTECTOR + .read() + .ok() + .and_then(|guard| guard.as_ref().cloned()); + + let Some(protector) = protector else { + return Ok(()); + }; + + let fd = socket.as_raw_fd(); + trace!(fd, "protecting android socket before connect"); + protector.protect_socket_fd(fd) +} + +pub(crate) fn must_bind_socket_on_interface( + _socket: &socket2::Socket, + iface: &OutboundInterface, + _family: socket2::Domain, +) -> io::Result<()> { + trace!( + iface = %iface.name, + "android outbound interface binding is handled by socket protection" + ); + Ok(()) +} diff --git a/clash-lib/src/proxy/utils/platform/mod.rs b/clash-lib/src/proxy/utils/platform/mod.rs index 085dcec9..4870e3d2 100644 --- a/clash-lib/src/proxy/utils/platform/mod.rs +++ b/clash-lib/src/proxy/utils/platform/mod.rs @@ -1,22 +1,41 @@ +#[cfg(target_os = "android")] +mod android; +#[cfg(target_os = "android")] +pub use android::{SocketProtector, clear_socket_protector, set_socket_protector}; +#[cfg(target_os = "android")] +pub(crate) use android::{maybe_protect_socket, must_bind_socket_on_interface}; + #[cfg(target_vendor = "apple")] mod apple; #[cfg(target_vendor = "apple")] pub(crate) use apple::must_bind_socket_on_interface; -#[cfg(any( - target_os = "fuchsia", - target_os = "linux", - target_os = "freebsd", - target_os = "android" -))] +#[cfg(any(target_os = "fuchsia", target_os = "linux", target_os = "freebsd"))] pub(crate) mod unix; -#[cfg(any( - target_os = "fuchsia", - target_os = "linux", - target_os = "freebsd", - target_os = "android" -))] +#[cfg(any(target_os = "fuchsia", target_os = "linux", target_os = "freebsd"))] pub(crate) use unix::must_bind_socket_on_interface; #[cfg(windows)] pub(crate) mod win; #[cfg(windows)] pub(crate) use win::must_bind_socket_on_interface; + +#[cfg(not(target_os = "android"))] +use std::{io, sync::Arc}; + +#[cfg(not(target_os = "android"))] +#[allow(dead_code)] +pub trait SocketProtector: Send + Sync { + fn protect_socket_fd(&self, fd: i32) -> io::Result<()>; +} + +#[cfg(not(target_os = "android"))] +#[allow(dead_code)] +pub fn set_socket_protector(_protector: Arc) {} + +#[cfg(not(target_os = "android"))] +#[allow(dead_code)] +pub fn clear_socket_protector() {} + +#[cfg(not(target_os = "android"))] +pub(crate) fn maybe_protect_socket(_socket: &socket2::Socket) -> io::Result<()> { + Ok(()) +} diff --git a/clash-lib/src/proxy/utils/socket_helpers.rs b/clash-lib/src/proxy/utils/socket_helpers.rs index 4178a34c..86b19287 100644 --- a/clash-lib/src/proxy/utils/socket_helpers.rs +++ b/clash-lib/src/proxy/utils/socket_helpers.rs @@ -4,12 +4,14 @@ use std::time::Duration; use socket2::TcpKeepalive; -use tokio::net::{TcpListener, TcpSocket, TcpStream}; +use tokio::net::{TcpListener, TcpSocket, TcpStream, UdpSocket}; use tokio::time::timeout; use tracing::{debug, trace}; use crate::app::net::OutboundInterface; -use crate::proxy::utils::platform::must_bind_socket_on_interface; +use crate::proxy::utils::platform::{ + maybe_protect_socket, must_bind_socket_on_interface, +}; pub fn apply_tcp_options(s: &TcpStream) -> std::io::Result<()> { #[cfg(not(target_os = "windows"))] @@ -114,12 +116,11 @@ pub async fn new_tcp_stream( }; debug!("created tcp socket"); - if !cfg!(target_os = "android") - && let Some(iface) = iface - { + if let Some(iface) = iface { must_bind_socket_on_interface(&socket, iface, family)?; - trace!("tcp socket bound to interface: {socket:?}"); + trace!(iface = ?iface, "tcp socket prepared for outbound interface"); } + maybe_protect_socket(&socket)?; #[cfg(not(target_os = "android"))] #[cfg(target_os = "linux")] @@ -137,3 +138,44 @@ pub async fn new_tcp_stream( ) .await? } + +pub fn new_udp_socket( + endpoint: SocketAddr, + iface: Option<&OutboundInterface>, + #[cfg(target_os = "linux")] so_mark: Option, +) -> std::io::Result { + let (socket, family) = match endpoint { + SocketAddr::V4(_) => ( + socket2::Socket::new(socket2::Domain::IPV4, socket2::Type::DGRAM, None)?, + socket2::Domain::IPV4, + ), + SocketAddr::V6(_) => ( + socket2::Socket::new(socket2::Domain::IPV6, socket2::Type::DGRAM, None)?, + socket2::Domain::IPV6, + ), + }; + maybe_protect_socket(&socket)?; + + if let Some(iface) = iface + && !cfg!(target_os = "android") + { + must_bind_socket_on_interface(&socket, iface, family)?; + trace!("udp socket bound to interface: {socket:?}"); + } + + #[cfg(not(target_os = "android"))] + #[cfg(target_os = "linux")] + if let Some(so_mark) = so_mark { + socket.set_mark(so_mark)?; + } + + socket.set_nonblocking(true)?; + + let bind_addr = match endpoint { + SocketAddr::V4(_) => SocketAddr::from(([0, 0, 0, 0], 0)), + SocketAddr::V6(_) => SocketAddr::from(([0u16; 8], 0)), + }; + socket.bind(&bind_addr.into())?; + + UdpSocket::from_std(socket.into()) +}