Skip to main content

aria2_core/engine/
bt_upload_session.rs

1use tracing::{debug, warn};
2
3use crate::error::Result;
4use crate::rate_limiter::RateLimiter;
5use crate::rate_limiter::RateLimiterConfig;
6
7use aria2_protocol::bittorrent::message::types::BtMessage;
8use aria2_protocol::bittorrent::peer::connection::PeerConnection;
9
10pub trait PieceDataProvider: Send + Sync {
11    fn get_piece_data(&self, piece_index: u32, offset: u32, length: u32) -> Option<Vec<u8>>;
12    fn has_piece(&self, piece_index: u32) -> bool;
13    fn num_pieces(&self) -> u32;
14}
15
16pub struct BtSeedingConfig {
17    pub max_upload_bytes_per_sec: Option<u64>,
18    pub max_peers_to_unchoke: usize,
19    pub optimistic_unchoke_interval_secs: u64,
20}
21
22impl Default for BtSeedingConfig {
23    fn default() -> Self {
24        Self {
25            max_upload_bytes_per_sec: None,
26            max_peers_to_unchoke: 4,
27            optimistic_unchoke_interval_secs: 30,
28        }
29    }
30}
31
32pub struct BtUploadSession {
33    conn: PeerConnection,
34    am_choke_state: bool,
35    peer_interested: bool,
36    uploaded_bytes: u64,
37    upload_limiter: Option<RateLimiter>,
38    pub(crate) is_dead: bool,
39}
40
41impl BtUploadSession {
42    pub fn new(conn: PeerConnection, config: &BtSeedingConfig) -> Self {
43        let upload_limiter = config
44            .max_upload_bytes_per_sec
45            .filter(|&r| r > 0)
46            .map(|r| RateLimiter::new(&RateLimiterConfig::new(None, Some(r))));
47
48        Self {
49            conn,
50            am_choke_state: false,
51            peer_interested: false,
52            uploaded_bytes: 0,
53            upload_limiter,
54            is_dead: false,
55        }
56    }
57
58    pub async fn handle_incoming_messages(
59        &mut self,
60        provider: &dyn PieceDataProvider,
61    ) -> Result<u64> {
62        if self.is_dead {
63            return Ok(0);
64        }
65
66        let round_uploaded = self.uploaded_bytes;
67
68        loop {
69            match self.conn.read_message().await {
70                Ok(Some(msg)) => match msg {
71                    BtMessage::Request { request } => {
72                        if !self.am_choke_state && self.peer_interested {
73                            debug!(
74                                "Upload request: piece={}, offset={}, len={}",
75                                request.index, request.begin, request.length
76                            );
77
78                            let data = provider.get_piece_data(
79                                request.index,
80                                request.begin,
81                                request.length,
82                            );
83                            if let Some(piece_data) = data {
84                                let data_len = piece_data.len() as u64;
85                                if let Some(ref lim) = self.upload_limiter {
86                                    lim.acquire_upload(data_len).await;
87                                }
88                                self.conn.send_message(&BtMessage::Piece {
89                                        index: request.index,
90                                        begin: request.begin,
91                                        data: piece_data,
92                                    }).await.map_err(|e| crate::error::Aria2Error::Recoverable(
93                                        crate::error::RecoverableError::TemporaryNetworkFailure { message: e }
94                                    ))?;
95                                self.uploaded_bytes += data_len;
96                            } else {
97                                warn!(
98                                    "No data for piece {} at offset {}",
99                                    request.index, request.begin
100                                );
101                            }
102                        } else {
103                            debug!(
104                                "Ignoring request: choked={} interested={}",
105                                self.am_choke_state, self.peer_interested
106                            );
107                        }
108                    }
109                    BtMessage::Interested => {
110                        self.peer_interested = true;
111                        if !self.am_choke_state {
112                            self.conn.send_unchoke().await.ok();
113                        }
114                    }
115                    BtMessage::NotInterested => {
116                        self.peer_interested = false;
117                    }
118                    BtMessage::Choke => {
119                        debug!("Peer choked us");
120                    }
121                    BtMessage::Unchoke => {
122                        debug!("Peer unchoked us");
123                    }
124                    BtMessage::Have { piece_index } => {
125                        debug!("Peer has piece {}", piece_index);
126                    }
127                    BtMessage::Cancel { request } => {
128                        debug!(
129                            "Peer cancelled request for piece {} offset {}",
130                            request.index, request.begin
131                        );
132                    }
133                    BtMessage::Piece { .. } => {
134                        debug!("Unexpected Piece from peer during seeding");
135                    }
136                    BtMessage::Bitfield { .. } => {}
137                    BtMessage::KeepAlive => {}
138                    BtMessage::Port { port: _ } => {}
139                    BtMessage::AllowedFast { index } => {
140                        debug!("Received AllowedFast for piece {}", index);
141                    }
142                    BtMessage::Reject {
143                        index,
144                        offset,
145                        length,
146                    } => {
147                        debug!(
148                            "Received Reject for piece {} offset {} len {}",
149                            index, offset, length
150                        );
151                    }
152                    BtMessage::Suggest { index } => {
153                        debug!("Received Suggest for piece {}", index);
154                    }
155                    BtMessage::HaveAll => {
156                        debug!("Received HaveAll");
157                    }
158                    BtMessage::HaveNone => {
159                        debug!("Received HaveNone");
160                    }
161                },
162                Ok(None) => {
163                    debug!("EOF from peer, marking session as dead");
164                    self.is_dead = true;
165                    break;
166                }
167                Err(e) => {
168                    warn!("Read error from upload peer: {}, marking dead", e);
169                    self.is_dead = true;
170                    break;
171                }
172            }
173        }
174
175        Ok(self.uploaded_bytes - round_uploaded)
176    }
177
178    pub async fn unchoke_peer(&mut self) -> Result<()> {
179        if self.am_choke_state {
180            self.conn.send_unchoke().await.map_err(|e| {
181                crate::error::Aria2Error::Recoverable(
182                    crate::error::RecoverableError::TemporaryNetworkFailure { message: e },
183                )
184            })?;
185            self.am_choke_state = false;
186        }
187        Ok(())
188    }
189
190    pub async fn choke_peer(&mut self) -> Result<()> {
191        if !self.am_choke_state {
192            self.conn.send_choke().await.map_err(|e| {
193                crate::error::Aria2Error::Recoverable(
194                    crate::error::RecoverableError::TemporaryNetworkFailure { message: e },
195                )
196            })?;
197            self.am_choke_state = true;
198        }
199        Ok(())
200    }
201
202    pub fn is_peer_choked(&self) -> bool {
203        self.am_choke_state
204    }
205
206    pub fn is_peer_interested(&self) -> bool {
207        self.peer_interested
208    }
209
210    pub fn is_dead(&self) -> bool {
211        self.is_dead
212    }
213
214    pub fn uploaded_bytes(&self) -> u64 {
215        self.uploaded_bytes
216    }
217
218    pub fn connection_mut(&mut self) -> &mut PeerConnection {
219        &mut self.conn
220    }
221}
222
223pub struct InMemoryPieceProvider {
224    pieces: Vec<Option<Vec<u8>>>,
225    piece_length: u32,
226}
227
228impl InMemoryPieceProvider {
229    pub fn new(piece_length: u32, num_pieces: u32) -> Self {
230        let mut pieces = Vec::with_capacity(num_pieces as usize);
231        for _ in 0..num_pieces {
232            pieces.push(None);
233        }
234        Self {
235            pieces,
236            piece_length,
237        }
238    }
239
240    pub fn set_piece_data(&mut self, index: u32, data: Vec<u8>) {
241        if (index as usize) < self.pieces.len() {
242            self.pieces[index as usize] = Some(data);
243        }
244    }
245
246    pub fn set_all_from_pattern<F>(&mut self, f: F)
247    where
248        F: Fn(u32, u32) -> u8,
249    {
250        for i in 0..self.pieces.len() {
251            let len = if i == self.pieces.len() - 1 {
252                let total = self.piece_length as usize * (self.pieces.len() - 1);
253                1024 * 100 - total
254            } else {
255                self.piece_length as usize
256            };
257            let mut data = Vec::with_capacity(len);
258            for j in 0..len {
259                data.push(f(i as u32, j as u32));
260            }
261            self.pieces[i] = Some(data);
262        }
263    }
264}
265
266impl PieceDataProvider for InMemoryPieceProvider {
267    fn get_piece_data(&self, piece_index: u32, offset: u32, length: u32) -> Option<Vec<u8>> {
268        let piece = self.pieces.get(piece_index as usize)?.as_ref()?;
269        let start = offset as usize;
270        let end = (start + length as usize).min(piece.len());
271        if start >= piece.len() {
272            return None;
273        }
274        Some(piece[start..end].to_vec())
275    }
276
277    fn has_piece(&self, piece_index: u32) -> bool {
278        self.pieces
279            .get(piece_index as usize)
280            .is_some_and(|p| p.is_some())
281    }
282
283    fn num_pieces(&self) -> u32 {
284        self.pieces.len() as u32
285    }
286}
287
288#[cfg(test)]
289mod tests {
290    use super::*;
291
292    #[test]
293    fn test_seeding_config_default() {
294        let cfg = BtSeedingConfig::default();
295        assert!(cfg.max_upload_bytes_per_sec.is_none());
296        assert_eq!(cfg.max_peers_to_unchoke, 4);
297        assert_eq!(cfg.optimistic_unchoke_interval_secs, 30);
298    }
299
300    #[test]
301    fn test_in_memory_provider_creation() {
302        let provider = InMemoryPieceProvider::new(16384, 10);
303        assert_eq!(provider.num_pieces(), 10);
304        assert!(!provider.has_piece(0));
305        assert!(provider.get_piece_data(0, 0, 100).is_none());
306    }
307
308    #[test]
309    fn test_in_memory_provider_set_and_get() {
310        let mut provider = InMemoryPieceProvider::new(256, 4);
311        provider.set_piece_data(0, vec![0xAB; 256]);
312        provider.set_piece_data(2, vec![0xCD; 128]);
313
314        assert!(provider.has_piece(0));
315        assert!(!provider.has_piece(1));
316        assert!(provider.has_piece(2));
317
318        let data = provider.get_piece_data(0, 10, 50).unwrap();
319        assert_eq!(data.len(), 50);
320        assert!(data.iter().all(|&b| b == 0xAB));
321
322        let partial = provider.get_piece_data(2, 100, 28).unwrap();
323        assert_eq!(partial.len(), 28);
324        assert!(partial.iter().all(|&b| b == 0xCD));
325    }
326
327    #[test]
328    fn test_in_memory_provider_set_all_from_pattern() {
329        let mut provider = InMemoryPieceProvider::new(100, 5);
330        provider.set_all_from_pattern(|piece_idx, byte_idx| {
331            ((piece_idx * 37 + byte_idx * 13) % 256) as u8
332        });
333
334        for i in 0..5u32 {
335            assert!(provider.has_piece(i));
336            let data = provider.get_piece_data(i, 0, 100).unwrap();
337            for (j, &byte) in data.iter().enumerate() {
338                assert_eq!(byte, ((i * 37 + j as u32 * 13) % 256) as u8);
339            }
340        }
341    }
342
343    #[test]
344    fn test_in_memory_provider_offset_beyond_piece() {
345        let mut provider = InMemoryPieceProvider::new(50, 2);
346        provider.set_piece_data(0, vec![0x42; 50]);
347
348        assert!(provider.get_piece_data(0, 40, 20).is_some());
349        assert!(provider.get_piece_data(0, 60, 10).is_none());
350        assert!(provider.get_piece_data(99, 0, 10).is_none());
351    }
352
353    #[test]
354    fn test_in_memory_provider_last_piece_smaller() {
355        let total_size = 260u32;
356        let piece_len = 100u32;
357        let num_pieces = total_size.div_ceil(piece_len);
358        let mut provider = InMemoryPieceProvider::new(piece_len, num_pieces);
359
360        provider.set_all_from_pattern(|_, idx| idx as u8);
361
362        assert!(provider.has_piece(0));
363        assert!(provider.has_piece(1));
364        assert!(provider.has_piece(2));
365
366        let last_piece = provider.get_piece_data(2, 0, 60).unwrap();
367        assert_eq!(last_piece.len(), 60);
368    }
369}