1
use std::{
2
    collections::{HashMap, HashSet},
3
    os::fd::{AsRawFd, FromRawFd, OwnedFd},
4
    path::PathBuf,
5
    sync::Arc,
6
    time::Duration,
7
};
8

            
9
use anyhow::Context;
10
use clap::Parser;
11
use tokio::{net::UdpSocket, sync::RwLock};
12
use tokio_util::sync::CancellationToken;
13
use tracing::level_filters::LevelFilter;
14
use tracing_subscriber::layer::SubscriberExt;
15

            
16
use roowho2_lib::{
17
    server::{
18
        config::{DEFAULT_CONFIG_PATH, LogLevel},
19
        ignore_list::IgnoreList,
20
        rwhod::{
21
            RwhodStatusRegistry, RwhodStatusStore, rwhod_packet_receiver_task,
22
            rwhod_packet_sender_task,
23
        },
24
        varlink_api::varlink_client_server_task,
25
    },
26
    version,
27
};
28

            
29
#[derive(Parser)]
30
#[command(
31
  author = "Programvareverkstedet <projects@pvv.ntnu.no>",
32
  about,
33
  version,
34
  long_version = version::LONG_VERSION
35
)]
36
struct Args {
37
    /// Path to configuration file
38
    #[arg(
39
        short = 'c',
40
        long = "config",
41
        default_value = DEFAULT_CONFIG_PATH,
42
        value_name = "PATH"
43
    )]
44
    config_path: PathBuf,
45
}
46

            
47
#[tokio::main]
48
async fn main() -> anyhow::Result<()> {
49
    let args = Args::parse();
50

            
51
    let config = toml::from_str::<roowho2_lib::server::config::Config>(
52
        &std::fs::read_to_string(&args.config_path).context(format!(
53
            "Failed to read configuration file {:?}",
54
            args.config_path
55
        ))?,
56
    )?;
57

            
58
    let log_filter = match config.log_level.unwrap_or(LogLevel::Info) {
59
        LogLevel::Info => LevelFilter::INFO,
60
        LogLevel::Debug => LevelFilter::DEBUG,
61
        LogLevel::Trace => LevelFilter::TRACE,
62
    };
63

            
64
    let subscriber = tracing_subscriber::registry()
65
        .with(log_filter)
66
        .with(tracing_journald::layer()?);
67

            
68
    tracing::subscriber::set_global_default(subscriber)
69
        .context("Failed to set global default tracing subscriber")?;
70

            
71
    let rwhod_ignore_list = IgnoreList::load_optional(config.rwhod.ignore_list_path.as_deref())?;
72
    let finger_ignore_list = IgnoreList::load_optional(config.fingerd.ignore_list_path.as_deref())?;
73

            
74
    let fd_map: HashMap<String, OwnedFd> =
75
        HashMap::from_iter(sd_notify::listen_fds_with_names()?.map(|(fd_num, name)| {
76
            (
77
                name.clone(),
78
                // SAFETY: please don't mess around with file descriptors in random places
79
                //         around the codebase lol
80
                unsafe { std::os::fd::OwnedFd::from_raw_fd(fd_num) },
81
            )
82
        }));
83

            
84
    let mut join_set = tokio::task::JoinSet::new();
85

            
86
    let whod_status_store: RwhodStatusStore = Arc::new(RwLock::new(RwhodStatusRegistry::new(
87
        config.rwhod.max_status_entries(),
88
    )));
89

            
90
    let client_server_token = CancellationToken::new();
91
    let client_server_token_ = client_server_token.clone();
92
    tokio::spawn(async move {
93
        client_server_token_.cancelled().await;
94
        tracing::info!("RWHOD client-server is now accepting connections");
95
        #[cfg(feature = "systemd")]
96
        sd_notify::notify(&[sd_notify::NotifyState::Ready]).ok();
97
        Ok::<(), anyhow::Error>(())
98
    });
99

            
100
    if config.rwhod.enable {
101
        tracing::info!("Starting RWHOD server");
102

            
103
        let socket = fd_map
104
            .get("rwhod_socket")
105
            .map(|fd| {
106
                // SAFETY: see above
107
                let std_socket = unsafe { std::net::UdpSocket::from_raw_fd(fd.as_raw_fd()) };
108
                std_socket.set_nonblocking(true)?;
109
                UdpSocket::from_std(std_socket)
110
            })
111
            .context("RWHOD server is enabled, but socket fd not provided by systemd")??;
112

            
113
        let audit_socket: Option<OwnedFd> = fd_map
114
            .get("audit_socket")
115
            .map(|fd| fd.try_clone())
116
            .transpose()
117
            .context("Failed to clone audit socket fd")?;
118

            
119
        join_set.spawn(rwhod_server(
120
            socket,
121
            whod_status_store.clone(),
122
            rwhod_ignore_list.clone(),
123
            config.rwhod.interfaces.clone(),
124
            config.rwhod.send_interval(),
125
            config.rwhod.realtime_updates_enabled(),
126
            audit_socket,
127
        ));
128
    } else {
129
        tracing::debug!("RWHOD server is disabled in configuration");
130
    }
131

            
132
    join_set.spawn(client_server(
133
        fd_map
134
            .get("client_socket")
135
            .context("RWHOD client-server socket fd not provided by systemd")?
136
            .try_clone()
137
            .context("Failed to clone RWHOD client-server socket fd")?,
138
        whod_status_store.clone(),
139
        config.rwhod.enable,
140
        config.fingerd.enable,
141
        config.walld.enable,
142
        finger_ignore_list,
143
        client_server_token,
144
    ));
145

            
146
    join_set.spawn(ctrl_c_handler());
147

            
148
    join_set.join_next().await.unwrap()??;
149

            
150
    Ok(())
151
}
152

            
153
async fn ctrl_c_handler() -> anyhow::Result<()> {
154
    tokio::signal::ctrl_c()
155
        .await
156
        .map_err(|e| anyhow::anyhow!("Failed to listen for Ctrl-C: {}", e))
157
}
158

            
159
async fn rwhod_server(
160
    socket: UdpSocket,
161
    whod_status_store: RwhodStatusStore,
162
    ignore_list: Option<IgnoreList>,
163
    allowed_interfaces: Option<HashSet<String>>,
164
    send_interval: Duration,
165
    realtime_updates_enabled: bool,
166
    audit_socket: Option<OwnedFd>,
167
) -> anyhow::Result<()> {
168
    let socket = Arc::new(socket);
169

            
170
    let interfaces =
171
        roowho2_lib::server::rwhod::determine_relevant_interfaces(allowed_interfaces.as_ref())?;
172

            
173
    let (tx, rx) = tokio::sync::mpsc::channel(1);
174
    match (realtime_updates_enabled, audit_socket) {
175
        (true, Some(fd)) => {
176
            tokio::spawn(roowho2_lib::server::rwhod::audit_change_notifier(tx, fd));
177
        }
178
        (true, None) => {
179
            tracing::warn!(
180
                "rwhod.realtime_updates is enabled, but no audit socket fd was provided by \
181
                 systemd (e.g. the kernel might not have been booted with `audit=1`); \
182
                 falling back to interval-only updates"
183
            );
184
        }
185
        (false, _) => {}
186
    }
187

            
188
    let sender_task =
189
        rwhod_packet_sender_task(socket.clone(), interfaces, ignore_list, send_interval, rx);
190
    let receiver_task = rwhod_packet_receiver_task(socket.clone(), whod_status_store);
191

            
192
    tokio::select! {
193
        res = sender_task => res?,
194
        res = receiver_task => res?,
195
    }
196

            
197
    Ok(())
198
}
199

            
200
async fn client_server(
201
    socket_fd: OwnedFd,
202
    whod_status_store: RwhodStatusStore,
203
    rwhod_enabled: bool,
204
    fingerd_enabled: bool,
205
    walld_enabled: bool,
206
    finger_ignore_list: Option<IgnoreList>,
207
    startup_token: CancellationToken,
208
) -> anyhow::Result<()> {
209
    // SAFETY: see above
210
    let std_socket =
211
        unsafe { std::os::unix::net::UnixListener::from_raw_fd(socket_fd.as_raw_fd()) };
212
    std_socket.set_nonblocking(true)?;
213
    let zlink_listener = zlink::tokio::unix::Listener::try_from(OwnedFd::from(std_socket))?;
214
    let client_server_task = varlink_client_server_task(
215
        zlink_listener,
216
        whod_status_store,
217
        rwhod_enabled,
218
        fingerd_enabled,
219
        walld_enabled,
220
        finger_ignore_list,
221
    );
222

            
223
    startup_token.cancel();
224

            
225
    client_server_task.await?;
226

            
227
    Ok(())
228
}