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}