client_mgr: add

format

client_mgr
This commit is contained in:
Luna Yao
2026-04-06 11:54:01 +02:00
parent b7073504d8
commit dc1e1929c8
5 changed files with 270 additions and 153 deletions
+5 -5
View File
@@ -1,7 +1,7 @@
use crate::dns::config::DNS_SERVER_RPC_ADDR; use crate::dns::config::DNS_SERVER_RPC_ADDR;
use crate::dns::peer_mgr::DnsPeerMgr; use crate::dns::peer_mgr::DnsPeerMgr;
use crate::peers::peer_manager::PeerManager; use crate::peers::peer_manager::PeerManager;
use crate::proto::dns::{DnsPeerMgrRpcServer, DnsServerRpcClientFactory, HeartbeatRequest}; use crate::proto::dns::{DnsClientMgrRpcClientFactory, DnsPeerMgrRpcServer, HeartbeatRequest};
use crate::proto::peer_rpc::RoutePeerInfo; use crate::proto::peer_rpc::RoutePeerInfo;
use crate::proto::rpc_impl::standalone::StandAloneClient; use crate::proto::rpc_impl::standalone::StandAloneClient;
use crate::proto::rpc_types::controller::BaseController; use crate::proto::rpc_types::controller::BaseController;
@@ -63,22 +63,22 @@ impl DnsClient {
) -> anyhow::Result<()> { ) -> anyhow::Result<()> {
let request = if heartbeat.snapshot.is_none() || self.mgr.dirty.peers.reset() { let request = if heartbeat.snapshot.is_none() || self.mgr.dirty.peers.reset() {
heartbeat.update(self.mgr.snapshot()); heartbeat.update(self.mgr.snapshot());
heartbeat.clone().into() heartbeat.clone()
} else { } else {
let snapshot = heartbeat.snapshot.take(); let snapshot = heartbeat.snapshot.take();
let request = heartbeat.clone().into(); let request = heartbeat.clone();
heartbeat.snapshot = snapshot; heartbeat.snapshot = snapshot;
request request
}; };
let client = rpc let client = rpc
.scoped_client::<DnsServerRpcClientFactory<BaseController>>("".to_string()) .scoped_client::<DnsClientMgrRpcClientFactory<BaseController>>("".to_string())
.await?; .await?;
let response = client.heartbeat(BaseController::default(), request).await?; let response = client.heartbeat(BaseController::default(), request).await?;
if response.resync { if response.resync {
client client
.heartbeat(BaseController::default(), heartbeat.clone().into()) .heartbeat(BaseController::default(), heartbeat.clone())
.await?; .await?;
} }
+165
View File
@@ -0,0 +1,165 @@
use crate::dns::utils::{DirtyFlag, NameServerAddr};
use crate::dns::zone::{Zone, ZoneGroup};
use crate::proto::dns::DnsClientMgrRpc;
use crate::proto::dns::{DnsSnapshot, HeartbeatRequest, HeartbeatResponse};
use crate::proto::rpc_types;
use crate::proto::rpc_types::controller::BaseController;
use crate::utils::{DeterministicDigest, MapTryInto};
use anyhow::Error;
use derive_more::{Deref, DerefMut};
use hickory_server::authority::Catalog;
use itertools::Itertools;
use moka::future::Cache;
use std::collections::HashSet;
use std::time::Duration;
use tokio::sync::Notify;
use uuid::Uuid;
#[derive(Debug, Clone, Default)]
pub struct DnsClientInfo {
digest: Vec<u8>,
zones: ZoneGroup,
addresses: HashSet<NameServerAddr>,
listeners: HashSet<NameServerAddr>,
}
impl TryFrom<&DnsSnapshot> for DnsClientInfo {
type Error = Error;
fn try_from(value: &DnsSnapshot) -> Result<Self, Self::Error> {
Ok(Self {
digest: value.digest(),
zones: (&value.zones).try_into()?,
addresses: value.addresses.iter().map_try_into().try_collect()?,
listeners: value.listeners.iter().map_try_into().try_collect()?,
})
}
}
const DNS_CLIENT_TTL: Duration = Duration::from_secs(5);
// TODO: same as DnsPeerMgrDirtyState
#[derive(Debug, Default, Deref, DerefMut)]
pub struct DnsClientMgrDirtyState {
pub(super) zones: DirtyFlag,
pub(super) addresses: DirtyFlag,
pub(super) listeners: DirtyFlag,
#[deref]
#[deref_mut]
notify: Notify,
}
#[derive(Debug)]
pub struct DnsClientMgr {
clients: Cache<Uuid, DnsClientInfo>,
pub(super) dirty: DnsClientMgrDirtyState,
}
impl DnsClientMgr {
pub fn new() -> Self {
Self {
clients: Cache::builder().time_to_live(DNS_CLIENT_TTL).build(),
dirty: Default::default(),
}
}
pub fn catalog(&self) -> Catalog {
let zones = self.collect_zones();
let mut catalog = Catalog::new();
for zone in zones.iter() {
catalog.upsert(
zone.origin.clone(),
zone.create_memory_authority().into_iter().collect(),
);
}
for zone in zones.iter() {
catalog.upsert(
zone.origin.clone(),
zone.create_forward_authority().into_iter().collect(),
);
}
catalog
}
pub fn collect_zones(&self) -> ZoneGroup {
let mut zones = vec![Zone::system()];
let mut local = HashSet::<NameServerAddr>::new();
for (_, info) in self.clients.iter() {
zones.extend(info.zones);
local.extend(info.addresses);
local.extend(info.listeners);
}
for zone in zones.iter_mut() {
if let Some(forward) = zone.forward.as_mut() {
forward
.name_servers
.retain(|ns| !local.contains(&ns.clone().into()));
}
}
zones.into()
}
pub fn iter_addresses(&self) -> impl Iterator<Item = NameServerAddr> + use<'_> {
self.clients
.iter()
.flat_map(|(_, info)| info.addresses)
.unique()
}
pub fn iter_listeners(&self) -> impl Iterator<Item = NameServerAddr> + use<'_> {
self.clients
.iter()
.flat_map(|(_, info)| info.listeners)
.unique()
}
}
#[async_trait::async_trait]
impl DnsClientMgrRpc for DnsClientMgr {
type Controller = BaseController;
async fn heartbeat(
&self,
_: BaseController,
input: HeartbeatRequest,
) -> rpc_types::error::Result<HeartbeatResponse> {
let id = input
.id
.ok_or(anyhow::anyhow!(
"missing id in heartbeat request: {:?}",
input
))?
.into();
let resync = if let Some(snapshot) = input.snapshot.as_ref() {
let new = DnsClientInfo::try_from(snapshot)?;
let old = self.clients.get(&id).await.unwrap_or_default();
if new.digest != old.digest {
self.dirty.zones.mark();
if new.addresses != old.addresses {
self.dirty.addresses.mark();
}
if new.listeners != old.listeners {
self.dirty.listeners.mark();
}
self.clients.insert(id, new).await;
self.dirty.notify_one();
}
false
} else {
self.clients
.get(&id)
.await
.is_none_or(|info| info.digest != input.digest)
};
Ok(HeartbeatResponse { resync })
}
}
+1
View File
@@ -1,4 +1,5 @@
mod client; mod client;
mod client_mgr;
pub mod config; pub mod config;
mod peer_mgr; mod peer_mgr;
mod server; mod server;
+98 -147
View File
@@ -1,52 +1,25 @@
use super::{utils::NameServerAddr, zone::Zone}; use super::{utils::NameServerAddr, zone::Zone};
use crate::dns::utils::DirtyFlag; use crate::common::PeerId;
use crate::dns::zone::ZoneGroup; use crate::dns::client_mgr::DnsClientMgr;
use crate::proto::dns::DnsSnapshot; use cidr::Ipv4Inet;
use crate::proto::rpc_types; use derivative::Derivative;
use crate::proto::{ use derive_more::{Deref, DerefMut, From, Into};
dns::{DnsServerRpc, HeartbeatRequest, HeartbeatResponse}, use hickory_proto::rr::Record;
rpc_types::controller::BaseController, use hickory_proto::serialize::binary::BinEncoder;
};
use crate::utils::{DeterministicDigest, MapTryInto};
use anyhow::Error;
use hickory_proto::xfer::Protocol; use hickory_proto::xfer::Protocol;
use hickory_server::{ use hickory_server::{
authority::Catalog, authority::{Catalog, MessageResponse},
server::{Request, RequestHandler, ResponseHandler, ResponseInfo}, server::{Request, RequestHandler, ResponseHandler, ResponseInfo},
ServerFuture, ServerFuture,
}; };
use itertools::Itertools; use itertools::Itertools;
use moka::future::Cache; use parking_lot::Mutex;
use std::collections::HashSet; use std::collections::HashSet;
use std::io;
use std::{sync::Arc, time::Duration}; use std::{sync::Arc, time::Duration};
use derivative::Derivative;
use derive_more::{Deref, DerefMut};
use tokio::net::{TcpListener, UdpSocket}; use tokio::net::{TcpListener, UdpSocket};
use tokio::{sync::RwLock, task::JoinHandle}; use tokio::{sync::RwLock, task::JoinHandle};
use tokio::sync::Notify;
use tokio_util::sync::CancellationToken; use tokio_util::sync::CancellationToken;
use uuid::Uuid;
#[derive(Debug, Clone, Default)]
pub struct DnsClientInfo {
digest: Vec<u8>,
zones: ZoneGroup,
addresses: HashSet<NameServerAddr>,
listeners: HashSet<NameServerAddr>,
}
impl TryFrom<&DnsSnapshot> for DnsClientInfo {
type Error = Error;
fn try_from(value: &DnsSnapshot) -> Result<Self, Self::Error> {
Ok(Self {
digest: value.digest(),
zones: (&value.zones).try_into()?,
addresses: value.addresses.iter().map_try_into().try_collect()?,
listeners: value.listeners.iter().map_try_into().try_collect()?,
})
}
}
#[derive(Clone)] #[derive(Clone)]
pub struct DynamicCatalog { pub struct DynamicCatalog {
@@ -80,17 +53,6 @@ impl RequestHandler for DynamicCatalog {
} }
} }
// TODO: same as DnsPeerMgrDirtyState
#[derive(Debug, Default, Deref, DerefMut)]
pub struct DnsServerDirtyState {
zones: DirtyFlag,
addresses: DirtyFlag,
listeners: DirtyFlag,
#[deref]
#[deref_mut]
notify: Notify,
}
struct DnsServerRuntime { struct DnsServerRuntime {
token: CancellationToken, token: CancellationToken,
task: JoinHandle<()>, task: JoinHandle<()>,
@@ -116,70 +78,102 @@ impl DnsServerRuntime {
} }
} }
// ResponseWrapper for serializing DNS responses into a byte buffer.
// Used by the address hijacking NIC packet filter to produce DNS replies in-place.
#[derive(Debug, Clone, From, Into, Deref, DerefMut)]
struct Response(Arc<Mutex<Vec<u8>>>);
impl Response {
pub fn new(capacity: usize) -> Self {
Self(Arc::new(Mutex::new(Vec::with_capacity(capacity))))
}
pub fn into_inner(self) -> Option<Vec<u8>> {
Arc::into_inner(self.0).map(Mutex::into_inner)
}
}
trait RecordIter<'r>: Iterator<Item = &'r Record> + Send + 'r {}
impl<'r, T> RecordIter<'r> for T where T: Iterator<Item = &'r Record> + Send + 'r {}
#[async_trait::async_trait]
impl ResponseHandler for Response {
async fn send_response<'r>(
&mut self,
response: MessageResponse<
'_,
'r,
impl RecordIter<'r>,
impl RecordIter<'r>,
impl RecordIter<'r>,
impl RecordIter<'r>,
>,
) -> io::Result<ResponseInfo> {
let max_size = if let Some(edns) = response.get_edns() {
edns.max_payload()
} else {
hickory_proto::udp::MAX_RECEIVE_BUFFER_SIZE as u16
};
let mut this = self.lock();
let mut encoder = BinEncoder::new(this.as_mut());
encoder.set_max_size(max_size);
response
.destructive_emit(&mut encoder)
.map_err(io::Error::other)
}
}
#[derive(Derivative)] #[derive(Derivative)]
#[derivative(Debug)] #[derivative(Debug)]
pub struct DnsServer { pub struct DnsServer {
clients: Cache<Uuid, DnsClientInfo>, mgr: DnsClientMgr,
dirty: DnsServerDirtyState,
#[derivative(Debug = "ignore")] #[derivative(Debug = "ignore")]
catalog: DynamicCatalog, catalog: DynamicCatalog,
/// Current set of hijacked addresses (only UDP protocol addresses).
hijacked: RwLock<HashSet<NameServerAddr>>,
/// Tun device name, needed for adding/removing routes.
tun_dev: RwLock<Option<String>>,
/// Tun device IP inet, used to check if an address is within the tun subnet.
tun_inet: RwLock<Option<Ipv4Inet>>,
/// Our peer ID, used to set the to_peer_id on response packets.
my_peer_id: RwLock<PeerId>,
} }
const DNS_CLIENT_TTL: Duration = Duration::from_secs(5);
const DNS_SERVER_LISTENER_TCP_TIMEOUT: Duration = Duration::from_secs(5); const DNS_SERVER_LISTENER_TCP_TIMEOUT: Duration = Duration::from_secs(5);
impl DnsServer { impl DnsServer {
async fn reload_zones(&self) { async fn reload_zones(&self) {
let mut zones = vec![Zone::system()]; self.catalog.replace(self.mgr.catalog()).await;
let mut local = HashSet::<NameServerAddr>::new();
for (_, info) in self.clients.iter() {
zones.extend(info.zones.iter().cloned());
local.extend(info.addresses.iter());
local.extend(info.listeners.iter());
}
let mut catalog = Catalog::new();
for zone in zones.iter_mut() {
if let Some(forward) = zone.forward.as_mut() {
forward
.name_servers
.retain(|ns| !local.contains(&ns.clone().into()));
}
}
for zone in zones.iter() {
catalog.upsert(
zone.origin.clone(),
zone.create_memory_authority().into_iter().collect(),
);
}
for zone in zones.iter() {
catalog.upsert(
zone.origin.clone(),
zone.create_forward_authority().into_iter().collect(),
);
}
self.catalog.replace(catalog).await;
} }
async fn reload_addresses(&self) { async fn reload_addresses(&self, addresses: impl IntoIterator<Item = NameServerAddr>) {
todo!() let addresses = addresses.into_iter().collect::<HashSet<_>>();
let mut active = self.hijacked.write().await;
let added = addresses.difference(&active).cloned().collect_vec();
let removed = active.difference(&addresses).cloned().collect_vec();
if added.is_empty() && removed.is_empty() {
return;
}
*active = addresses;
// TODO
} }
async fn reload_listeners(&self, runtime: &mut Option<DnsServerRuntime>) -> anyhow::Result<()> { async fn reload_listeners(
let listeners = self &self,
.clients listeners: impl IntoIterator<Item = NameServerAddr>,
.iter() runtime: &mut Option<DnsServerRuntime>,
.map(|(_, info)| info.listeners.into_iter()) ) -> anyhow::Result<()> {
.flatten()
.collect_vec();
if let Some(old) = runtime.take() { if let Some(old) = runtime.take() {
old.stop().await?; old.stop().await?;
} }
@@ -209,24 +203,27 @@ impl DnsServer {
} }
pub async fn run(&self) { pub async fn run(&self) {
let dirty = &self.dirty; let dirty = &self.mgr.dirty;
let mut runtime = None; let mut runtime = None;
loop { loop {
dirty.notified().await; dirty.notified().await;
if dirty.zones.reset() { if dirty.zones.reset() {
self.reload_zones().await; self.reload_zones(&self.mgr.collect_zones()).await;
} }
if dirty.addresses.reset() { if dirty.addresses.reset() {
self.reload_addresses().await; self.reload_addresses(self.mgr.iter_addresses()).await;
} }
if dirty.listeners.reset() { if dirty.listeners.reset() {
if let Err(e) = self.reload_listeners(&mut runtime).await { if let Err(e) = self
.reload_listeners(self.mgr.iter_listeners(), &mut runtime)
.await
{
tracing::error!("failed to reload listeners: {:?}", e); tracing::error!("failed to reload listeners: {:?}", e);
self.dirty.listeners.mark(); dirty.listeners.mark();
self.dirty.notify_one(); dirty.notify_one();
} }
} }
@@ -234,49 +231,3 @@ impl DnsServer {
} }
} }
} }
#[async_trait::async_trait]
impl DnsServerRpc for DnsServer {
type Controller = BaseController;
async fn heartbeat(
&self,
_: BaseController,
input: HeartbeatRequest,
) -> rpc_types::error::Result<HeartbeatResponse> {
let id = input
.id
.ok_or(anyhow::anyhow!(
"missing id in heartbeat request: {:?}",
input
))?
.into();
let resync = if let Some(snapshot) = input.snapshot.as_ref() {
let new = DnsClientInfo::try_from(snapshot)?;
let old = self.clients.get(&id).await.unwrap_or_default();
if new.digest != old.digest {
if new.zones != old.zones {
self.dirty.zones.mark();
}
if new.addresses != old.addresses {
self.dirty.addresses.mark();
}
if new.listeners != old.listeners {
self.dirty.listeners.mark();
}
self.clients.insert(id, new).await;
self.dirty.notify_one();
}
false
} else {
self.clients
.get(&id)
.await
.is_none_or(|info| info.digest != input.digest)
};
Ok(HeartbeatResponse { resync })
}
}
+1 -1
View File
@@ -38,6 +38,6 @@ message HeartbeatResponse {
bool resync = 1; bool resync = 1;
} }
service DnsServerRpc { service DnsClientMgrRpc {
rpc Heartbeat(HeartbeatRequest) returns (HeartbeatResponse) {} rpc Heartbeat(HeartbeatRequest) returns (HeartbeatResponse) {}
} }