1
1
#![allow(clippy::redundant_field_names)]
2
#![allow(dead_code)]
3
//#![allow(unused_variables)]
4

            
5
#[macro_use]
6
extern crate tracing;
7

            
8
mod tftp;
9
pub mod errors;
10
pub mod util;
11
pub mod fetcher;
12

            
13
use std::{sync::Arc, os::unix::prelude::RawFd};
14
use util::{ UdpSocket, UdpRecvInfo, SocketAddr, Bucket };
15

            
16
use tftp::{ Session, SessionStats };
17

            
18
pub use errors::{ Error, Result };
19

            
20

            
21
pub struct Environment {
22
    dir:		std::path::PathBuf,
23
    fallback_uri:	Option<std::ffi::OsString>,
24
    max_block_size:	u16,
25
    max_window_size:	u16,
26
    max_connections:	u32,
27
    timeout:		std::time::Duration,
28
}
29

            
30
struct SpeedInfo {
31
    duration:		std::time::Duration,
32
    stats:		SessionStats,
33
}
34

            
35
impl SpeedInfo {
36
    pub fn new(now: std::time::Instant, stats: SessionStats) -> Self {
37
	Self {
38
	    duration:	now.elapsed(),
39
	    stats:	stats,
40
	}
41
    }
42
}
43

            
44
impl std::fmt::Display for SpeedInfo {
45
    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
46
	use num_format::{ ToFormattedString, SystemLocale };
47
	let locale = SystemLocale::default().unwrap();
48

            
49
	match self.stats.speed_bit_per_s(self.duration) {
50
	    None			=> write!(f, "n/a"),
51
	    Some((speed_f, speed_n)) if speed_f == speed_n	=>
52
		write!(f, "total={} bps",
53
		       (speed_f as u64).to_formatted_string(&locale)),
54

            
55
	    Some((speed_f, speed_n))	=>
56
		write!(f, "file={} bps, net={} bps",
57
		       (speed_f as u64).to_formatted_string(&locale),
58
		       (speed_n as u64).to_formatted_string(&locale)),
59
	}
60
    }
61
}
62

            
63
use tracing::field::Empty;
64

            
65
#[instrument(skip_all,
66
	     fields(id = id,
67
		    remote = Empty,
68
		    local = Empty,
69
		    filename = Empty,
70
		    filesize = Empty,
71
		    op = Empty))]
72
async fn handle_request(env: std::sync::Arc<Environment>,
73
			id: u64,
74
			info: UdpRecvInfo,
75
			req: Vec<u8>,
76
			bucket: Arc<Bucket>)
77
{
78
    let instant = std::time::Instant::now();
79
    let session = Session::new(&env, info.remote, info.local).await;
80

            
81
    if let Err(e) = session {
82
	warn!("failed to create tftp session: {:?}", e);
83
	return;
84
    }
85

            
86
    let session = session.unwrap();
87

            
88
    let b = bucket.acquire();
89

            
90
    let res = match b.is_ok() {
91
	false	=> session.do_reject().await,
92
	true	=> session.run(req).await
93
    };
94

            
95
    match res {
96
	Ok(stats)	=> {
97
	    info!("stats: {}", stats);
98
	    info!("speed: {}", SpeedInfo::new(instant, stats))
99
	},
100
	Err(e)	=> error!("request failed: {:?}", e),
101
    };
102
}
103

            
104
async fn run_tftpd_loop(env: std::sync::Arc<Environment>, sock: UdpSocket) -> Result<()> {
105
    let mut buf = vec![0u8; 1500];
106

            
107
    let bucket = Arc::new(Bucket::new(env.max_connections));
108
    let mut num = 0;
109

            
110
    loop {
111
	let info = sock.recvmsg(&mut buf).await?;
112
	let request = Vec::from(&buf[..info.size]);
113

            
114
	tokio::task::spawn(handle_request(env.clone(), num, info,
115
					  request, bucket.clone()));
116

            
117
	num += 1;
118
    }
119
}
120

            
121
enum Either<T: Sized, U: Sized> {
122
    A(T),
123
    B(U),
124
}
125

            
126
#[tokio::main(flavor = "current_thread")]
127
async fn run(env: Environment, info: Either<SocketAddr, RawFd>) -> Result<()> {
128
    // UdpSocket creation must happen with active Tokio runtime
129
    let mut sock = match info {
130
	Either::A(addr)	=> UdpSocket::bind(addr),
131
	Either::B(fd)	=> UdpSocket::from_raw(fd),
132
    }?;
133

            
134
    sock.set_nonblocking()?;
135
    sock.set_request_pktinfo()?;
136

            
137
    run_tftpd_loop(std::sync::Arc::new(env), sock).await?;
138

            
139
    Ok(())
140
}
141

            
142
use clap::Parser;
143

            
144
#[derive(clap::Parser, Debug)]
145
struct CliOpts {
146
    #[clap(short, long, help("use systemd fd propagation"), value_parser)]
147
    systemd:		bool,
148

            
149
    #[clap(short, long, value_parser, help("port to listen on"), default_value("69"))]
150
    port:		u16,
151

            
152
    #[clap(short, long, value_parser, help("ip address to listen on"),
153
	   value_name("IP"), default_value("::"))]
154
    listen:		std::net::IpAddr,
155

            
156
    #[clap(short, long, value_parser, help("maximum number of connections"),
157
	   value_name("NUM"), default_value("64"))]
158
    max_connections:	u32,
159

            
160
    #[clap(short, long, value_parser, help("timeout in seconds during tftp transfers"),
161
	   default_value("3"))]
162
    timeout:		f32,
163

            
164
    #[clap(short, long, value_parser, value_name("URI"), help("fallback uri"))]
165
    fallback:		Option<String>,
166
}
167

            
168
fn main() {
169
    tracing_subscriber::fmt::init();
170

            
171
    let args = CliOpts::parse();
172

            
173
    let env = Environment {
174
	dir:			".".into(),
175
	fallback_uri:		args.fallback.map(|s| s.into()),
176
	max_block_size:		1500,
177
	max_window_size:	64,
178
	max_connections:	args.max_connections,
179
	timeout:		std::time::Duration::from_secs_f32(args.timeout),
180
    };
181

            
182
    let fd = match args.systemd {
183
	true	=> listenfd::ListenFd::from_env()
184
	    .take_raw_fd(0)
185
	    .unwrap(),
186
	false	=> None
187
    };
188

            
189
    let info = match fd {
190
	None		=> Either::A(SocketAddr::new(args.listen, args.port)),
191
	Some(fd)	=> Either::B(fd),
192
    };
193

            
194
    run(env, info).unwrap();
195
}