Simple mini clone of torrra

This commit is contained in:
Alexander
2026-07-02 17:02:49 +02:00
commit 80ebf1cb63
24 changed files with 6424 additions and 0 deletions
+122
View File
@@ -0,0 +1,122 @@
use anyhow::Result;
use sqlx::postgres::{PgPool, PgPoolOptions};
use uuid::Uuid;
pub async fn connect(database_url: &str) -> Result<PgPool> {
let pool = PgPoolOptions::new()
.max_connections(5)
.connect(database_url)
.await?;
Ok(pool)
}
pub async fn ping(pool: &PgPool) -> Result<()> {
sqlx::query("SELECT 1").execute(pool).await?;
Ok(())
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, sqlx::Type)]
#[sqlx(type_name = "torrent_state", rename_all = "lowercase")]
pub enum TorrentState {
Pending,
Downloading,
Paused,
Finished,
Error,
}
#[derive(Debug, Clone, sqlx::FromRow)]
pub struct TorrentRow {
pub id: Uuid,
pub info_hash: String,
pub name: Option<String>,
pub source: String,
pub output_path: String,
pub total_bytes: Option<i64>,
pub downloaded_bytes: i64,
pub state: TorrentState,
pub error_message: Option<String>,
}
/// Insert a new pending torrent, or return the existing row if the info_hash is already tracked.
pub async fn insert_pending(
pool: &PgPool,
info_hash: &str,
source: &str,
output_path: &str,
) -> Result<TorrentRow> {
let row = sqlx::query_as::<_, TorrentRow>(
"INSERT INTO torrents (info_hash, source, output_path)
VALUES ($1, $2, $3)
ON CONFLICT (info_hash) DO UPDATE SET updated_at = now()
RETURNING id, info_hash, name, source, output_path, total_bytes, downloaded_bytes, state, error_message",
)
.bind(info_hash)
.bind(source)
.bind(output_path)
.fetch_one(pool)
.await?;
Ok(row)
}
pub async fn get(pool: &PgPool, id: Uuid) -> Result<Option<TorrentRow>> {
let row = sqlx::query_as::<_, TorrentRow>(
"SELECT id, info_hash, name, source, output_path, total_bytes, downloaded_bytes, state, error_message
FROM torrents WHERE id = $1",
)
.bind(id)
.fetch_optional(pool)
.await?;
Ok(row)
}
pub async fn list(pool: &PgPool) -> Result<Vec<TorrentRow>> {
let rows = sqlx::query_as::<_, TorrentRow>(
"SELECT id, info_hash, name, source, output_path, total_bytes, downloaded_bytes, state, error_message
FROM torrents ORDER BY added_at DESC",
)
.fetch_all(pool)
.await?;
Ok(rows)
}
pub async fn list_pending(pool: &PgPool) -> Result<Vec<TorrentRow>> {
let rows = sqlx::query_as::<_, TorrentRow>(
"SELECT id, info_hash, name, source, output_path, total_bytes, downloaded_bytes, state, error_message
FROM torrents WHERE state = 'pending' ORDER BY added_at",
)
.fetch_all(pool)
.await?;
Ok(rows)
}
pub struct Progress<'a> {
pub name: Option<&'a str>,
pub total_bytes: i64,
pub downloaded_bytes: i64,
pub state: TorrentState,
pub error_message: Option<&'a str>,
}
pub async fn update_progress(pool: &PgPool, id: Uuid, progress: Progress<'_>) -> Result<()> {
sqlx::query(
"UPDATE torrents
SET name = COALESCE($2, name),
total_bytes = $3,
downloaded_bytes = $4,
state = $5,
error_message = $6,
updated_at = now(),
completed_at = CASE WHEN $5 = 'finished' THEN now() ELSE completed_at END
WHERE id = $1",
)
.bind(id)
.bind(progress.name)
.bind(progress.total_bytes)
.bind(progress.downloaded_bytes)
.bind(progress.state)
.bind(progress.error_message)
.execute(pool)
.await?;
Ok(())
}
+11
View File
@@ -0,0 +1,11 @@
use anyhow::{Context, Result};
use librqbit::Magnet;
/// Extract the BTIH info hash from a magnet link, normalized to lowercase hex.
pub fn info_hash(magnet: &str) -> Result<String> {
let parsed = Magnet::parse(magnet).context("failed to parse magnet link")?;
let id20 = parsed
.as_id20()
.context("magnet link has no v1 (BTIH) info hash")?;
Ok(id20.as_string())
}
+137
View File
@@ -0,0 +1,137 @@
mod db;
mod magnet;
mod torrents;
use std::path::PathBuf;
use std::time::Duration;
use anyhow::{Context, Result};
use clap::Parser;
use tokio::net::UnixListener;
use tokio_stream::wrappers::UnixListenerStream;
use tonic::transport::Server;
use tonic_health::ServingStatus;
use tonic_health::server::HealthReporter;
use tora_proto::TorrentsServer;
use tracing::{error, info};
use torrents::{GrpcTorrents, TorrentManager};
const HEALTH_POLL_INTERVAL: Duration = Duration::from_secs(3);
const HEALTH_QUERY_TIMEOUT: Duration = Duration::from_secs(2);
#[derive(Parser)]
#[command(name = "torad", about = "tora background daemon")]
struct Args {
/// Path to the Unix domain socket to serve on.
#[arg(long, env = "TORAD_SOCKET")]
socket: Option<PathBuf>,
/// Postgres connection string.
#[arg(long, env = "DATABASE_URL")]
database_url: String,
/// Directory torrents are downloaded into by default.
#[arg(long)]
download_dir: Option<PathBuf>,
/// Path to write torad's PID to, so it can be killed without hunting for the process
/// (e.g. `kill $(cat /tmp/torad.pid)`).
#[arg(long, env = "TORAD_PID_FILE")]
pid_file: Option<PathBuf>,
}
fn default_socket_path() -> PathBuf {
std::env::temp_dir().join("torad.sock")
}
fn default_pid_file_path() -> PathBuf {
std::env::temp_dir().join("torad.pid")
}
fn default_download_dir() -> PathBuf {
std::env::var_os("HOME")
.map(PathBuf::from)
.unwrap_or_else(|| PathBuf::from("."))
.join("Downloads")
}
#[tokio::main]
async fn main() -> Result<()> {
tracing_subscriber::fmt::init();
let args = Args::parse();
let socket = args.socket.unwrap_or_else(default_socket_path);
let pid_file = args.pid_file.unwrap_or_else(default_pid_file_path);
let download_dir = args.download_dir.unwrap_or_else(default_download_dir);
let pool = db::connect(&args.database_url)
.await
.context("failed to connect to Postgres")?;
info!("connected to Postgres");
let manager = TorrentManager::new(pool.clone(), download_dir)
.await
.context("failed to start torrent manager")?;
manager.clone().spawn_poller();
std::fs::write(&pid_file, std::process::id().to_string())
.with_context(|| format!("failed to write pid file at {}", pid_file.display()))?;
if socket.exists() {
std::fs::remove_file(&socket)
.with_context(|| format!("failed to remove stale socket at {}", socket.display()))?;
}
let listener = UnixListener::bind(&socket)
.with_context(|| format!("failed to bind socket at {}", socket.display()))?;
let uds_stream = UnixListenerStream::new(listener);
let (health_reporter, health_service) = tonic_health::server::health_reporter();
tokio::spawn(poll_db_health(pool, health_reporter));
info!(socket = %socket.display(), "torad listening");
Server::builder()
.add_service(health_service)
.add_service(TorrentsServer::new(GrpcTorrents::new(manager)))
.serve_with_incoming_shutdown(uds_stream, shutdown_signal())
.await
.context("gRPC server error")?;
info!("shutting down");
let _ = std::fs::remove_file(&socket);
let _ = std::fs::remove_file(&pid_file);
Ok(())
}
async fn poll_db_health(pool: sqlx::PgPool, reporter: HealthReporter) {
let mut interval = tokio::time::interval(HEALTH_POLL_INTERVAL);
loop {
interval.tick().await;
let status = match tokio::time::timeout(HEALTH_QUERY_TIMEOUT, db::ping(&pool)).await {
Ok(Ok(())) => ServingStatus::Serving,
Ok(Err(err)) => {
error!(error = %err, "postgres health check query failed");
ServingStatus::NotServing
}
Err(_) => {
error!("postgres health check timed out");
ServingStatus::NotServing
}
};
reporter.set_service_status("", status).await;
}
}
async fn shutdown_signal() {
use tokio::signal::unix::{SignalKind, signal};
let mut sigint = signal(SignalKind::interrupt()).expect("failed to install SIGINT handler");
let mut sigterm = signal(SignalKind::terminate()).expect("failed to install SIGTERM handler");
tokio::select! {
_ = sigint.recv() => {},
_ = sigterm.recv() => {},
}
}
+234
View File
@@ -0,0 +1,234 @@
use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use anyhow::{Context, Result};
use librqbit::{AddTorrent, AddTorrentOptions, ManagedTorrent, Session, TorrentStatsState};
use sqlx::PgPool;
use tokio::sync::Mutex;
use tonic::{Request, Response, Status};
use tora_proto::{
AddRequest, AddResponse, ListRequest, ListResponse, State as ProtoState, StatusRequest,
TorrentStatus, Torrents,
};
use tracing::{error, info, warn};
use uuid::Uuid;
use crate::db::{self, TorrentRow, TorrentState};
use crate::magnet;
const POLL_INTERVAL: Duration = Duration::from_secs(2);
pub struct TorrentManager {
pool: PgPool,
session: Arc<Session>,
default_output_dir: PathBuf,
tracked: Mutex<HashMap<Uuid, Arc<ManagedTorrent>>>,
}
impl TorrentManager {
pub async fn new(pool: PgPool, download_dir: PathBuf) -> Result<Arc<Self>> {
std::fs::create_dir_all(&download_dir).with_context(|| {
format!(
"failed to create download directory {}",
download_dir.display()
)
})?;
let session = Session::new(download_dir.clone())
.await
.context("failed to create librqbit session")?;
Ok(Arc::new(Self {
pool,
session,
default_output_dir: download_dir,
tracked: Mutex::new(HashMap::new()),
}))
}
pub async fn add(&self, magnet: &str, output_dir: Option<&str>) -> Result<TorrentRow> {
let info_hash = magnet::info_hash(magnet)?;
let output_path = output_dir
.filter(|s| !s.is_empty())
.map(str::to_string)
.unwrap_or_else(|| self.default_output_dir.display().to_string());
db::insert_pending(&self.pool, &info_hash, magnet, &output_path).await
}
pub async fn get(&self, id: Uuid) -> Result<Option<TorrentRow>> {
db::get(&self.pool, id).await
}
pub async fn list(&self) -> Result<Vec<TorrentRow>> {
db::list(&self.pool).await
}
pub fn spawn_poller(self: Arc<Self>) {
tokio::spawn(async move {
let mut interval = tokio::time::interval(POLL_INTERVAL);
loop {
interval.tick().await;
self.pick_up_pending().await;
self.report_progress().await;
}
});
}
async fn pick_up_pending(&self) {
let pending = match db::list_pending(&self.pool).await {
Ok(rows) => rows,
Err(err) => {
error!(error = %err, "failed to list pending torrents");
return;
}
};
if pending.is_empty() {
return;
}
let mut tracked = self.tracked.lock().await;
for row in pending {
if tracked.contains_key(&row.id) {
continue;
}
let options = AddTorrentOptions {
output_folder: Some(row.output_path.clone()),
overwrite: true,
..Default::default()
};
match self
.session
.add_torrent(AddTorrent::from_url(row.source.clone()), Some(options))
.await
{
Ok(response) => match response.into_handle() {
Some(handle) => {
info!(id = %row.id, info_hash = %row.info_hash, "torrent added to session");
tracked.insert(row.id, handle);
}
None => warn!(id = %row.id, "add_torrent returned no handle"),
},
Err(err) => {
error!(id = %row.id, error = %err, "failed to add torrent to session");
let message = err.to_string();
let progress = db::Progress {
name: None,
total_bytes: 0,
downloaded_bytes: 0,
state: TorrentState::Error,
error_message: Some(message.as_str()),
};
if let Err(err) = db::update_progress(&self.pool, row.id, progress).await {
error!(error = %err, "failed to persist torrent add error");
}
}
}
}
}
async fn report_progress(&self) {
let tracked = self.tracked.lock().await;
for (id, handle) in tracked.iter() {
let stats = handle.stats();
let name = handle.name();
let state = if stats.finished {
TorrentState::Finished
} else {
match stats.state {
TorrentStatsState::Initializing => TorrentState::Pending,
TorrentStatsState::Live => TorrentState::Downloading,
TorrentStatsState::Paused => TorrentState::Paused,
TorrentStatsState::Error => TorrentState::Error,
}
};
let progress = db::Progress {
name: name.as_deref(),
total_bytes: i64::try_from(stats.total_bytes).unwrap_or(i64::MAX),
downloaded_bytes: i64::try_from(stats.progress_bytes).unwrap_or(i64::MAX),
state,
error_message: stats.error.as_deref(),
};
if let Err(err) = db::update_progress(&self.pool, *id, progress).await {
error!(id = %id, error = %err, "failed to persist torrent progress");
}
}
}
}
fn proto_state(state: TorrentState) -> ProtoState {
match state {
TorrentState::Pending => ProtoState::Pending,
TorrentState::Downloading => ProtoState::Downloading,
TorrentState::Paused => ProtoState::Paused,
TorrentState::Finished => ProtoState::Finished,
TorrentState::Error => ProtoState::Error,
}
}
fn row_to_status(row: TorrentRow) -> TorrentStatus {
TorrentStatus {
id: row.id.to_string(),
info_hash: row.info_hash,
name: row.name.unwrap_or_default(),
source: row.source,
output_path: row.output_path,
total_bytes: row.total_bytes.unwrap_or(0) as u64,
downloaded_bytes: row.downloaded_bytes as u64,
state: proto_state(row.state) as i32,
error_message: row.error_message.unwrap_or_default(),
}
}
// ---- gRPC Torrents service. This is the only interface `tora` talks to. ----
pub struct GrpcTorrents {
manager: Arc<TorrentManager>,
}
impl GrpcTorrents {
pub fn new(manager: Arc<TorrentManager>) -> Self {
Self { manager }
}
}
#[tonic::async_trait]
impl Torrents for GrpcTorrents {
async fn add(&self, request: Request<AddRequest>) -> Result<Response<AddResponse>, Status> {
let req = request.into_inner();
let output_dir = (!req.output_dir.is_empty()).then_some(req.output_dir.as_str());
let row = self
.manager
.add(&req.magnet, output_dir)
.await
.map_err(|err| Status::internal(err.to_string()))?;
Ok(Response::new(AddResponse {
id: row.id.to_string(),
}))
}
async fn status(
&self,
request: Request<StatusRequest>,
) -> Result<Response<TorrentStatus>, Status> {
let id = Uuid::parse_str(&request.into_inner().id)
.map_err(|_| Status::invalid_argument("invalid id"))?;
let row = self
.manager
.get(id)
.await
.map_err(|err| Status::internal(err.to_string()))?
.ok_or_else(|| Status::not_found("torrent not found"))?;
Ok(Response::new(row_to_status(row)))
}
async fn list(&self, _request: Request<ListRequest>) -> Result<Response<ListResponse>, Status> {
let rows = self
.manager
.list()
.await
.map_err(|err| Status::internal(err.to_string()))?;
Ok(Response::new(ListResponse {
torrents: rows.into_iter().map(row_to_status).collect(),
}))
}
}