1use std::{
29 io::{BufRead, BufReader, Read, Write},
30 net::{SocketAddr, TcpListener, TcpStream, ToSocketAddrs},
31 sync::{
32 Arc,
33 atomic::{AtomicBool, Ordering},
34 },
35 time::Duration,
36};
37
38#[derive(thiserror::Error, Debug)]
40pub enum PeerError {
41 #[error("peer request was not authorized")]
43 Unauthorized,
44 #[error("peer handle not found")]
46 NotFound,
47 #[error("peer returned HTTP {0}")]
49 Status(u16),
50 #[error("peer protocol error: {0}")]
52 Protocol(String),
53 #[error("{0}")]
55 Io(String),
56}
57
58impl From<std::io::Error> for PeerError {
59 fn from(error: std::io::Error) -> Self {
60 PeerError::Io(error.to_string())
61 }
62}
63
64pub trait ByteSource: Send + Sync {
70 fn len(&self) -> Option<u64>;
72
73 fn read_at(&self, offset: u64, buf: &mut [u8]) -> std::io::Result<usize>;
76
77 fn is_empty(&self) -> bool {
79 self.len() == Some(0)
80 }
81}
82
83pub struct BytesSource {
86 bytes: Vec<u8>,
87}
88
89impl BytesSource {
90 pub fn new(bytes: Vec<u8>) -> Self {
91 Self { bytes }
92 }
93}
94
95impl ByteSource for BytesSource {
96 fn len(&self) -> Option<u64> {
97 Some(self.bytes.len() as u64)
98 }
99
100 fn read_at(&self, offset: u64, buf: &mut [u8]) -> std::io::Result<usize> {
101 let offset = offset.min(self.bytes.len() as u64) as usize;
102 let available = &self.bytes[offset..];
103 let n = available.len().min(buf.len());
104 buf[..n].copy_from_slice(&available[..n]);
105 Ok(n)
106 }
107}
108
109pub type SourceResolver = Arc<dyn Fn(&str) -> Option<Arc<dyn ByteSource>> + Send + Sync>;
111
112pub struct PeerServer {
114 addr: SocketAddr,
115 running: Arc<AtomicBool>,
116}
117
118impl PeerServer {
119 pub fn start(
122 bind_addr: impl ToSocketAddrs,
123 token: impl Into<String>,
124 resolver: SourceResolver,
125 ) -> Result<PeerServer, PeerError> {
126 let listener = TcpListener::bind(bind_addr)?;
127 let addr = listener.local_addr()?;
128 let running = Arc::new(AtomicBool::new(true));
129 let token = token.into();
130
131 let loop_running = running.clone();
132 std::thread::Builder::new()
133 .name("cranpose-peer".to_string())
134 .spawn(move || {
135 for stream in listener.incoming() {
136 if !loop_running.load(Ordering::SeqCst) {
137 break;
138 }
139 let Ok(stream) = stream else { continue };
140 let token = token.clone();
141 let resolver = resolver.clone();
142 let _ = std::thread::Builder::new()
143 .name("cranpose-peer-conn".to_string())
144 .spawn(move || {
145 let _ = handle_connection(stream, &token, &resolver);
146 });
147 }
148 })
149 .map_err(|error| PeerError::Io(error.to_string()))?;
150
151 Ok(PeerServer { addr, running })
152 }
153
154 pub fn local_addr(&self) -> SocketAddr {
156 self.addr
157 }
158
159 pub fn port(&self) -> u16 {
161 self.addr.port()
162 }
163}
164
165impl Drop for PeerServer {
166 fn drop(&mut self) {
167 self.running.store(false, Ordering::SeqCst);
168 let _ = TcpStream::connect(self.addr);
169 }
170}
171
172fn handle_connection(
173 mut stream: TcpStream,
174 token: &str,
175 resolver: &SourceResolver,
176) -> Result<(), PeerError> {
177 stream.set_read_timeout(Some(Duration::from_secs(30)))?;
178 let mut reader = BufReader::new(stream.try_clone()?);
179
180 let mut request_line = String::new();
181 if reader.read_line(&mut request_line)? == 0 {
182 return Ok(());
183 }
184 let mut parts = request_line.split_whitespace();
185 let method = parts.next().unwrap_or("");
186 let path = parts.next().unwrap_or("");
187
188 let mut authorization = None;
189 let mut range = None;
190 loop {
191 let mut line = String::new();
192 if reader.read_line(&mut line)? == 0 {
193 break;
194 }
195 let line = line.trim_end();
196 if line.is_empty() {
197 break;
198 }
199 if let Some((name, value)) = line.split_once(':') {
200 let value = value.trim();
201 match name.trim().to_ascii_lowercase().as_str() {
202 "authorization" => authorization = Some(value.to_string()),
203 "range" => range = parse_range_header(value),
204 _ => {}
205 }
206 }
207 }
208
209 if method != "GET" {
210 return write_status(&mut stream, 405, "Method Not Allowed");
211 }
212 if authorization.as_deref() != Some(&format!("Bearer {token}")) {
213 return write_status(&mut stream, 401, "Unauthorized");
214 }
215 let Some(handle) = path.strip_prefix("/track/") else {
216 return write_status(&mut stream, 404, "Not Found");
217 };
218 let handle = crate::content::percent_decode_lossy(handle);
219 let Some(source) = resolver(&handle) else {
220 return write_status(&mut stream, 404, "Not Found");
221 };
222
223 serve_source(&mut stream, source.as_ref(), range)
224}
225
226fn serve_source(
227 stream: &mut TcpStream,
228 source: &dyn ByteSource,
229 range: Option<(u64, Option<u64>)>,
230) -> Result<(), PeerError> {
231 let total = source.len();
232
233 let (status, reason, start, length) = match (range, total) {
234 (Some((start, end)), Some(total)) if start < total => {
235 let last = end.unwrap_or(total - 1).min(total - 1);
236 if last < start {
237 return write_status(stream, 416, "Range Not Satisfiable");
238 }
239 (206, "Partial Content", start, last - start + 1)
240 }
241 (Some((start, _)), Some(total)) if start >= total => {
242 return write_status(stream, 416, "Range Not Satisfiable");
243 }
244 (_, Some(total)) => (200, "OK", 0, total),
245 (Some(_), None) => return write_status(stream, 416, "Range Not Satisfiable"),
246 (None, None) => {
247 return serve_unknown_length(stream, source);
248 }
249 };
250
251 let mut header = format!(
252 "HTTP/1.1 {status} {reason}\r\nContent-Length: {length}\r\nAccept-Ranges: bytes\r\nContent-Type: application/octet-stream\r\nConnection: close\r\n"
253 );
254 if status == 206
255 && let Some(total) = total
256 {
257 let end = start + length - 1;
258 header.push_str(&format!("Content-Range: bytes {start}-{end}/{total}\r\n"));
259 }
260 header.push_str("\r\n");
261 stream.write_all(header.as_bytes())?;
262
263 stream_bytes(stream, source, start, length)
264}
265
266fn serve_unknown_length(stream: &mut TcpStream, source: &dyn ByteSource) -> Result<(), PeerError> {
267 let header =
268 "HTTP/1.1 200 OK\r\nContent-Type: application/octet-stream\r\nConnection: close\r\n\r\n";
269 stream.write_all(header.as_bytes())?;
270 let mut buf = vec![0u8; 64 * 1024];
271 let mut offset = 0u64;
272 loop {
273 let n = source.read_at(offset, &mut buf)?;
274 if n == 0 {
275 break;
276 }
277 stream.write_all(&buf[..n])?;
278 offset += n as u64;
279 }
280 Ok(())
281}
282
283fn stream_bytes(
284 stream: &mut TcpStream,
285 source: &dyn ByteSource,
286 start: u64,
287 length: u64,
288) -> Result<(), PeerError> {
289 let mut buf = vec![0u8; 64 * 1024];
290 let mut sent = 0u64;
291 while sent < length {
292 let want = ((length - sent) as usize).min(buf.len());
293 let n = source.read_at(start + sent, &mut buf[..want])?;
294 if n == 0 {
295 break;
296 }
297 stream.write_all(&buf[..n])?;
298 sent += n as u64;
299 }
300 Ok(())
301}
302
303fn write_status(stream: &mut TcpStream, code: u16, reason: &str) -> Result<(), PeerError> {
304 let response =
305 format!("HTTP/1.1 {code} {reason}\r\nContent-Length: 0\r\nConnection: close\r\n\r\n");
306 stream.write_all(response.as_bytes())?;
307 Ok(())
308}
309
310fn parse_range_header(value: &str) -> Option<(u64, Option<u64>)> {
311 let spec = value.trim().strip_prefix("bytes=")?;
312 let (start, end) = spec.split_once('-')?;
313 let start = start.trim().parse::<u64>().ok()?;
314 let end = end.trim();
315 let end = if end.is_empty() {
316 None
317 } else {
318 Some(end.parse::<u64>().ok()?)
319 };
320 Some((start, end))
321}
322
323pub struct FetchResult {
325 pub total_len: Option<u64>,
327 pub bytes: Vec<u8>,
329}
330
331struct ResponseHead {
332 total_len: Option<u64>,
333 content_length: Option<u64>,
334 reader: BufReader<TcpStream>,
335}
336
337fn open_request(
338 base: &str,
339 token: &str,
340 handle: &str,
341 start: u64,
342 len: Option<u64>,
343) -> Result<ResponseHead, PeerError> {
344 let mut stream = TcpStream::connect(base)?;
345 stream.set_read_timeout(Some(Duration::from_secs(30)))?;
346
347 let range = match len {
348 Some(len) if len > 0 => format!("bytes={start}-{}", start + len - 1),
349 Some(_) => format!("bytes={start}-{start}"),
350 None => format!("bytes={start}-"),
351 };
352 let request = format!(
353 "GET /track/{} HTTP/1.1\r\nHost: {base}\r\nAuthorization: Bearer {token}\r\nRange: {range}\r\nConnection: close\r\n\r\n",
354 encode_handle(handle)
355 );
356 stream.write_all(request.as_bytes())?;
357
358 let mut reader = BufReader::new(stream);
359 let mut status_line = String::new();
360 reader.read_line(&mut status_line)?;
361 let status = parse_status(&status_line)?;
362
363 let mut total_len = None;
364 let mut content_length = None;
365 loop {
366 let mut line = String::new();
367 if reader.read_line(&mut line)? == 0 {
368 break;
369 }
370 let line = line.trim_end();
371 if line.is_empty() {
372 break;
373 }
374 if let Some((name, value)) = line.split_once(':') {
375 match name.trim().to_ascii_lowercase().as_str() {
376 "content-length" => content_length = value.trim().parse::<u64>().ok(),
377 "content-range" => total_len = parse_content_range_total(value.trim()),
378 _ => {}
379 }
380 }
381 }
382
383 match status {
384 401 => Err(PeerError::Unauthorized),
385 404 => Err(PeerError::NotFound),
386 200 | 206 => Ok(ResponseHead {
387 total_len,
388 content_length,
389 reader,
390 }),
391 other => Err(PeerError::Status(other)),
392 }
393}
394
395pub fn fetch_range(
398 base: &str,
399 token: &str,
400 handle: &str,
401 start: u64,
402 len: Option<u64>,
403) -> Result<FetchResult, PeerError> {
404 let mut head = open_request(base, token, handle, start, len)?;
405 let mut bytes = Vec::new();
406 match head.content_length {
407 Some(length) => {
408 bytes.resize(length as usize, 0);
409 head.reader.read_exact(&mut bytes)?;
410 }
411 None => {
412 head.reader.read_to_end(&mut bytes)?;
413 }
414 }
415 Ok(FetchResult {
416 total_len: head.total_len,
417 bytes,
418 })
419}
420
421pub fn fetch_to_writer(
425 base: &str,
426 token: &str,
427 handle: &str,
428 start: u64,
429 len: Option<u64>,
430 writer: &mut dyn Write,
431) -> Result<Option<u64>, PeerError> {
432 let mut head = open_request(base, token, handle, start, len)?;
433 let mut buf = vec![0u8; 64 * 1024];
434 let mut remaining = head.content_length;
435 loop {
436 let want = match remaining {
437 Some(0) => break,
438 Some(r) => (r as usize).min(buf.len()),
439 None => buf.len(),
440 };
441 let n = head.reader.read(&mut buf[..want])?;
442 if n == 0 {
443 break;
444 }
445 writer.write_all(&buf[..n])?;
446 if let Some(r) = remaining.as_mut() {
447 *r -= n as u64;
448 }
449 }
450 Ok(head.total_len)
451}
452
453pub fn content_length(base: &str, token: &str, handle: &str) -> Result<Option<u64>, PeerError> {
455 Ok(fetch_range(base, token, handle, 0, Some(1))?.total_len)
456}
457
458fn parse_status(line: &str) -> Result<u16, PeerError> {
459 line.split_whitespace()
460 .nth(1)
461 .and_then(|code| code.parse::<u16>().ok())
462 .ok_or_else(|| PeerError::Protocol(format!("bad status line: {line:?}")))
463}
464
465fn parse_content_range_total(value: &str) -> Option<u64> {
466 value.rsplit('/').next()?.trim().parse::<u64>().ok()
467}
468
469fn encode_handle(handle: &str) -> String {
470 let mut out = String::with_capacity(handle.len());
471 for byte in handle.as_bytes() {
472 match byte {
473 b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => {
474 out.push(*byte as char);
475 }
476 other => out.push_str(&format!("%{other:02X}")),
477 }
478 }
479 out
480}
481
482#[cfg(test)]
483mod tests {
484 use super::*;
485
486 fn resolver_for(handle: &'static str, bytes: Vec<u8>) -> SourceResolver {
487 Arc::new(move |requested: &str| {
488 if requested == handle {
489 Some(Arc::new(BytesSource::new(bytes.clone())) as Arc<dyn ByteSource>)
490 } else {
491 None
492 }
493 })
494 }
495
496 #[test]
497 fn round_trips_full_and_partial() {
498 let data: Vec<u8> = (0..=255u8).cycle().take(5000).collect();
499 let server = PeerServer::start("127.0.0.1:0", "secret", resolver_for("song", data.clone()))
500 .expect("start");
501 let base = format!("127.0.0.1:{}", server.port());
502
503 let full = fetch_range(&base, "secret", "song", 0, None).expect("full");
504 assert_eq!(full.bytes, data);
505 assert_eq!(full.total_len, Some(5000));
506
507 let part = fetch_range(&base, "secret", "song", 1000, Some(256)).expect("part");
508 assert_eq!(part.bytes, data[1000..1256]);
509 assert_eq!(part.total_len, Some(5000));
510
511 assert_eq!(content_length(&base, "secret", "song").unwrap(), Some(5000));
512 }
513
514 #[test]
515 fn streams_to_writer_without_buffering() {
516 let data: Vec<u8> = (0..2000u32).map(|i| i as u8).collect();
517 let server =
518 PeerServer::start("127.0.0.1:0", "k", resolver_for("s", data.clone())).expect("start");
519 let base = format!("127.0.0.1:{}", server.port());
520 let mut out = Vec::new();
521 let total = fetch_to_writer(&base, "k", "s", 0, None, &mut out).expect("stream");
522 assert_eq!(out, data);
523 assert_eq!(total, Some(2000));
524 }
525
526 #[test]
527 fn rejects_wrong_token() {
528 let server =
529 PeerServer::start("127.0.0.1:0", "right", resolver_for("a", vec![1, 2, 3])).expect("s");
530 let base = format!("127.0.0.1:{}", server.port());
531 assert!(matches!(
532 fetch_range(&base, "wrong", "a", 0, None),
533 Err(PeerError::Unauthorized)
534 ));
535 }
536
537 #[test]
538 fn unknown_handle_is_not_found() {
539 let server =
540 PeerServer::start("127.0.0.1:0", "t", resolver_for("a", vec![1, 2, 3])).expect("s");
541 let base = format!("127.0.0.1:{}", server.port());
542 assert!(matches!(
543 fetch_range(&base, "t", "missing", 0, None),
544 Err(PeerError::NotFound)
545 ));
546 }
547
548 #[test]
549 fn handle_is_percent_encoded_round_trip() {
550 let server =
551 PeerServer::start("127.0.0.1:0", "t", resolver_for("a b/c.mp3", vec![9, 8, 7]))
552 .expect("s");
553 let base = format!("127.0.0.1:{}", server.port());
554 let got = fetch_range(&base, "t", "a b/c.mp3", 0, None).expect("fetch");
555 assert_eq!(got.bytes, vec![9, 8, 7]);
556 }
557}