1
use crate::{ Result, Error };
2
use crate::fetcher::Fetcher;
3
use crate::tftp::SequenceId;
4

            
5
use super::Datagram;
6

            
7
enum Data<'a> {
8
    Owned(Vec<u8>),
9
    Ref(Option<&'a[u8]>),
10
}
11

            
12
impl Data<'_>
13
{
14
    pub fn alloc(sz: usize) -> Self {
15
	let mut data = Vec::with_capacity(sz);
16

            
17
	#[allow(clippy::uninit_vec)]
18
	unsafe { data.set_len(sz); }
19

            
20
	Self::Owned(data)
21
    }
22
}
23

            
24
struct Block<'a> {
25
    data:	Data<'a>,
26
    len:	u16,
27
    blksz:	u16,
28
}
29

            
30
impl <'a> Block<'a> {
31
    pub fn new_owned(size: u16) -> Self {
32
	Self {
33
	    data:	Data::alloc(size as usize),
34
	    blksz:	size,
35
	    len:	0,
36
	}
37
    }
38

            
39
6
    pub fn new_ref(size: u16) -> Self {
40
6
	Self {
41
6
	    data:	Data::Ref(None),
42
6
	    blksz:	size,
43
6
	    len:	0,
44
6
	}
45
6
    }
46

            
47
21
    pub fn get_blksize(&self) -> u16 {
48
21
	self.blksz
49
21
    }
50

            
51
11
    pub fn init(&mut self) -> &mut Self {
52
11
	self.len = 0;
53
11
	self
54
11
    }
55

            
56
11
    pub fn set_len(&mut self, sz: usize) {
57
11
	assert!(sz <= self.blksz as usize);
58
11
	self.len = u16::try_from(sz).unwrap();
59
11
    }
60

            
61
40
    pub fn get_data(&self) -> &[u8] {
62
40
	match &self.data {
63
	    Data::Owned(d)	=> &d[0..self.len as usize],
64
40
	    Data::Ref(Some(d))	=> &d[0..self.len as usize],
65
	    Data::Ref(None)	=> panic!("Data::Ref is None"),
66
	}
67
40
    }
68

            
69
10
    pub async fn fill<'b>(&mut self, fetcher: &'b mut Fetcher) -> Result<usize>
70
10
//    where
71
10
//	'b: 'a
72
10
    {
73
10
	let sz = match &mut self.data {
74
	    Data::Owned(d)	=> fetcher.read(d).await?,
75
	    Data::Ref(_)	=> {
76
10
		let data = fetcher.read_mmap(self.get_blksize() as usize)?;
77

            
78
		// TODO: this should be solved by better lifetime specifications...
79
10
		self.data = unsafe { std::mem::transmute(Data::Ref(Some(data))) };
80
10
		data.len()
81
	    }
82
	};
83

            
84
10
	self.set_len(sz);
85
10

            
86
10
	Ok(sz)
87
10
    }
88
}
89

            
90
2
#[derive(Default)]
91
struct BlockInfo {
92
    seq:	SequenceId,
93
    idx:	u16,
94
}
95

            
96
pub struct Xfer<'a> {
97
    start:	BlockInfo,
98
    active_sz:	u16,
99
    blocks:	Vec<Block<'a>>,
100
    is_eof:	bool,
101
}
102

            
103
impl <'a> Xfer<'a> {
104
2
    pub fn new<'b>(fetcher: &'b Fetcher, blk_size: u16, window_size: u16) -> Self
105
2
    where
106
2
	'a: 'b
107
2
    {
108
2
	assert!(window_size > 0);
109
2
	assert!(window_size < u16::MAX);
110

            
111
2
	let window_size = window_size as usize;
112
2

            
113
2
	let mut blocks = Vec::with_capacity(window_size);
114
2

            
115
2
	for _ in 0..window_size {
116
6
	    match fetcher.is_mmaped() {
117
6
		true	=> blocks.push(Block::new_ref(blk_size)),
118
		false	=> blocks.push(Block::new_owned(blk_size)),
119
	    }
120
	}
121

            
122
2
	Self {
123
2
	    start:	BlockInfo::default(),
124
2
	    active_sz:	0,
125
2
	    blocks:	blocks,
126
2
	    is_eof:	false,
127
2
	}
128
2
    }
129

            
130
111
    fn window_size(&self) -> u16
131
111
    {
132
111
	self.blocks.len() as u16
133
111
    }
134

            
135
60
    fn get_rel_block(&self, idx: u16) -> Option<(SequenceId, &Block)>
136
60
    {
137
60
	if idx >= self.active_sz {
138
20
	    return None;
139
40
	}
140
40

            
141
40
	let mut p = self.start.idx + idx;
142
40

            
143
40
	if p >= self.window_size() {
144
14
	    p -= self.window_size();
145
26
	}
146

            
147
40
	Some((self.start.seq + idx, &self.blocks[p as usize]))
148
60
    }
149

            
150
11
    fn alloc_block(&mut self) -> Option<&mut Block<'a>>
151
11
    {
152
11
	if self.active_sz >= self.window_size() {
153
	    return None;
154
11
	}
155
11

            
156
11
	let mut p = self.start.idx + self.active_sz;
157
11

            
158
11
	if p >= self.window_size() {
159
3
	    p -= self.window_size();
160
8
	}
161

            
162
11
	self.active_sz += 1;
163
11

            
164
11
	let block = &mut self.blocks[p as usize];
165
11

            
166
11
	block.init();
167
11

            
168
11
	Some(block)
169
11
    }
170

            
171
10
    fn free_blocks(&mut self, blk_id: SequenceId) -> Result<()>
172
10
    {
173
10
	let delta = if self.active_sz == 0 {
174
2
	    0_u16
175
	} else {
176
8
	    blk_id.delta(self.start.seq)
177
	};
178

            
179
	#[allow(clippy::comparison_chain)]
180
10
	if delta == self.active_sz {
181
5
	    trace!("all active blocks consumed");
182
5
	    self.start.idx = 0;
183
5
	    self.start.seq = blk_id;
184
5
	    self.active_sz = 0;
185
5
	} else if delta > self.active_sz {
186
2
	    return Err(Error::Protocol("blk-id out of window"));
187
	} else {
188
3
	    trace!("freeing {} blocks", delta);
189
3
	    self.start.idx = (self.start.idx + delta) % self.window_size();
190
3
	    self.start.seq += delta;
191
3
	    self.active_sz -= delta;
192
	}
193

            
194
8
	Ok(())
195
10
    }
196

            
197
10
    pub async fn fill_window<'b>(&mut self, blk_id: SequenceId, fetcher: &'b mut Fetcher) -> Result<usize>
198
10
//    where
199
10
//	'b: 'a
200
10
    {
201
10
	assert!(self.active_sz <= self.window_size());
202

            
203
	trace!("filling {:?} in {:?}@{}+{}", blk_id, self.start.seq, self.start.idx, self.active_sz);
204

            
205
10
	self.free_blocks(blk_id)?;
206

            
207
8
	let res = self.active_sz as usize;
208

            
209
19
	while self.active_sz < self.window_size() && !self.is_eof {
210
11
	    let block = self.alloc_block().unwrap();
211

            
212
11
	    let sz = if fetcher.is_eof() {
213
1
		block.set_len(0);
214
1
		0
215
	    } else {
216
10
		block.fill(fetcher).await?
217
	    };
218

            
219
11
	    if sz < block.get_blksize() as usize {
220
2
		self.is_eof = true;
221
9
	    }
222

            
223
	    debug!("read {}; active_sz={}", sz, self.active_sz);
224
	}
225

            
226
8
	Ok(res)
227
10
    }
228

            
229
12
    pub fn is_eof(&self) -> bool
230
12
    {
231
12
	self.is_eof && self.active_sz == 0
232
12
    }
233

            
234
20
    pub fn iter(&'a self) -> XferIterator<'a>
235
20
    {
236
20
	XferIterator {
237
20
	    xfer: self,
238
20
	    pos: 0,
239
20
	}
240
20
    }
241
}
242

            
243
pub struct XferIterator<'a>
244
{
245
    xfer: &'a Xfer<'a>,
246
    pos: u16,
247
}
248

            
249
impl <'a> Iterator for XferIterator<'a> {
250
    type Item = Datagram<'a>;
251

            
252
60
    fn next(&mut self) -> Option<Self::Item>
253
60
    {
254
60
	let (seq, block) = self.xfer.get_rel_block(self.pos)?;
255

            
256
40
	self.pos += 1;
257
40

            
258
40
	Some(Datagram::Data(seq, block.get_data()))
259
60
    }
260
}
261

            
262
#[cfg(test)]
263
mod test {
264
    use super::*;
265

            
266
10
    fn verify_data(xfer: &Xfer, start_idx: SequenceId, cnt: u16)
267
10
    {
268
20
	for (idx, d) in xfer.iter().enumerate() {
269
20
	    println!("idx={}, d={:?}", idx, d);
270
20
	    match d {
271
20
		Datagram::Data(id, data)	=> {
272
20
		    assert_eq!(id, start_idx + idx as u16);
273

            
274
20
		    match id.as_u16() {
275
1
			23	=> assert_eq!(data, &[ 0,  1]),
276
1
			24	=> assert_eq!(data, &[ 2,  3]),
277
4
			25	=> assert_eq!(data, &[ 4,  5]),
278
3
			26	=> assert_eq!(data, &[ 6,  7]),
279
3
			27	=> assert_eq!(data, &[ 8,  9]),
280
1
			28	=> assert_eq!(data, &[10, 11]),
281
1
			29	=> assert_eq!(data, &[12, 13]),
282
2
			30	=> assert_eq!(data, &[14, 15]),
283
1
			31	=> assert_eq!(data, &[]),
284

            
285
1
			50	=> assert_eq!(data, &[ 0,  1]),
286
2
			51	=> assert_eq!(data, &[ 2 ]),
287
			_	=> unreachable!(),
288
		    }
289
		},
290

            
291
		_	=> unreachable!(),
292
	    }
293
	}
294

            
295
10
	assert_eq!(xfer.iter().count(), cnt as usize);
296
10
    }
297

            
298
1
    #[tokio::test]
299
1
    async fn test_0() {
300
1
	tracing_subscriber::fmt::init();
301
1

            
302
1
	let mut f = Fetcher::new_memory(&[0, 1,   2,  3,   4,  5,   6,  7,
303
1
					  8, 9,  10, 11,  12, 13,  14, 15]);
304
1

            
305
1
	let mut xfer = Xfer::new(&f, 2, 3);
306
1

            
307
1
	assert!(!xfer.is_eof());
308

            
309
1
	let mut seq = SequenceId::new(23);
310
1

            
311
1
	info!("23/0 + 3");
312
1
	xfer.fill_window(seq, &mut f).await.expect("fill_window(0) failed");
313
1
	verify_data(&xfer, seq, 3);
314
1
	assert!(!xfer.is_eof());
315

            
316
1
	info!("25/2 + 3; last buffer of previous transfer was lost");
317
1
	seq += 2;		// 25
318
1
	xfer.fill_window(seq, &mut f).await.expect("fill_window(+2) failed");
319
1
	verify_data(&xfer, seq, 3);
320
1
	assert!(!xfer.is_eof());
321

            
322
1
	info!("24; error");
323
1
	seq -= 1;		// 24
324
1
	xfer.fill_window(seq, &mut f).await.expect_err("out-of-window blkid succeeded");
325
1
	seq += 1;		// 25
326
1
	verify_data(&xfer, seq, 3);
327
1
	assert!(!xfer.is_eof());
328

            
329
1
	info!("29; error");
330
1
	seq += 4;		// 29
331
1
	xfer.fill_window(seq, &mut f).await.expect_err("out-of-window blkid succeeded");
332
1
	seq -= 4;		// 25
333
1
	verify_data(&xfer, seq, 3);
334
1
	assert!(!xfer.is_eof());
335

            
336
1
	seq += 3;		// 28
337
1
	xfer.fill_window(seq, &mut f).await.expect("fill_window(+3) failed");
338
1
	verify_data(&xfer, seq, 3);
339
1
	assert!(!xfer.is_eof());
340

            
341
1
	seq += 2;		// 30
342
1
	xfer.fill_window(seq, &mut f).await.expect("fill_window(+3) failed");
343
1
	verify_data(&xfer, seq, 2);
344
1
	assert!(!xfer.is_eof());
345

            
346
1
	seq += 2;		// 32
347
1
	xfer.fill_window(seq, &mut f).await.expect("fill_window(+3) failed");
348
1
	verify_data(&xfer, seq, 0);
349
1
	assert!(xfer.is_eof());
350
1
    }
351

            
352
1
    #[tokio::test]
353
1
    async fn test_1() {
354
1
	let mut f = Fetcher::new_memory(&[0, 1, 2]);
355
1

            
356
1
	let mut xfer = Xfer::new(&f, 2, 3);
357
1

            
358
1
	assert!(!xfer.is_eof());
359

            
360
1
	let mut seq = SequenceId::new(50);
361
1
	xfer.fill_window(seq, &mut f).await.expect("fill_window(0) failed");
362
1
	verify_data(&xfer, seq, 2);
363
1
	assert!(!xfer.is_eof());
364

            
365
1
	seq += 1;		// 51
366
1
	xfer.fill_window(seq, &mut f).await.expect("fill_window(0) failed");
367
1
	verify_data(&xfer, seq, 1);
368
1
	assert!(!xfer.is_eof());
369

            
370
1
	seq += 1;		// 52
371
1
	xfer.fill_window(seq, &mut f).await.expect("fill_window(0) failed");
372
1
	verify_data(&xfer, seq, 0);
373
1
	assert!(xfer.is_eof());
374
1
    }
375
}