1
use std::time::Duration;
2

            
3
use crate::{ Error, Result };
4
use crate::util::{ UdpSocket, SocketAddr };
5
use super::{ Request, RequestError as E, RequestResult, SequenceId };
6

            
7

            
8
20
#[derive(Debug)]
9
pub enum Datagram<'a> {
10
    Read(Request<'a>),
11
    Write(Request<'a>),
12
    Data(SequenceId, &'a[u8]),
13
    Ack(SequenceId),
14
    Error(u16, &'a[u8]),
15
    OAck
16
}
17

            
18
trait TftpSlice {
19
    fn assert_len(&self, sz: usize) -> RequestResult<()>;
20
    fn get_u16(&self, idx: usize) -> u16;
21
    fn get_sequence_id(&self, idx: usize) -> SequenceId;
22
}
23

            
24
impl TftpSlice for &[u8]
25
{
26
    fn assert_len(&self, sz: usize) -> RequestResult<()>
27
    {
28
	if self.len() < sz {
29
	    return Err(E::TooShort);
30
	}
31

            
32
	Ok(())
33
    }
34

            
35
    fn get_u16(&self, idx: usize) -> u16
36
    {
37
	(self[idx] as u16) << 8 | (self[idx + 1] as u16)
38
    }
39

            
40
    fn get_sequence_id(&self, idx: usize) -> SequenceId
41
    {
42
	SequenceId::new(self.get_u16(idx))
43
    }
44
}
45

            
46
impl <'a> TryFrom<&'a[u8]> for Datagram<'a> {
47
    type Error = Error;
48

            
49
    #[instrument(level = "trace", skip(v), ret)]
50
    fn try_from(v: &'a [u8]) -> Result<Self> {
51
	use super::request::Dir;
52

            
53
	v.assert_len(2)?;
54
	let op = v.get_u16(0);
55

            
56
	Ok(match op {
57
	    1	=> {
58
		v.assert_len(2 + 1)?;
59
		Datagram::Read(Request::from_slice(&v[2..], Dir::Read)?)
60
	    },
61
	    2	=> {
62
		v.assert_len(2 + 1)?;
63
		Datagram::Write(Request::from_slice(&v[2..], Dir::Write)?)
64
	    },
65
	    3	=> {
66
		v.assert_len(2 + 2)?;
67
		Datagram::Data(v.get_sequence_id(2), &v[4..])
68
	    },
69
	    4	=> {
70
		v.assert_len(2 + 2)?;
71
		Datagram::Ack(v.get_sequence_id(2))
72
	    },
73
	    5	=> {
74
		v.assert_len(2 + 2)?;
75
		Datagram::Error(v.get_u16(2), &v[4..])
76
	    },
77
	    6	=> Datagram::OAck,
78
	    _	=> Err(E::BadOpCode(op))?,
79
	})
80
    }
81
}
82

            
83
impl <'a> Datagram<'a> {
84
    async fn recv_inner(sock: &UdpSocket,
85
			buf: &'a mut [u8], exp_addr: &SocketAddr) -> Result<Datagram<'a>>
86
    {
87
	loop {
88
	    let (len, addr) = sock.recvfrom(buf).await?;
89

            
90
	    if &addr != exp_addr {
91
		error!("unexpected address: {} vs {}", addr, exp_addr);
92
		// TODO: audit this event?
93
		continue;
94
	    }
95

            
96
	    return Self::try_from(&buf[0..len])
97
	}
98
    }
99

            
100
    pub async fn recv(sock: &UdpSocket,
101
		      buf: &'a mut [u8], exp_addr: &SocketAddr, to: Duration) -> Result<Datagram<'a>>
102
    {
103
	use tokio::time::timeout;
104

            
105
	timeout(to, Self::recv_inner(sock, buf, exp_addr)).await
106
	    .map_err(|_| Error::Timeout)
107
	    .and_then(|v| v)
108
    }
109

            
110
    pub fn is_ack(&self) -> bool {
111
	matches!(self, Self::Ack(_))
112
    }
113

            
114
    pub fn get_data_len(&self) -> usize {
115
	match self {
116
	    &Self::Data(_, d)	=> d.len(),
117
	    _			=> 0,
118
	}
119
    }
120
}