Implement notification stream rpc

This commit is contained in:
Alexander
2026-07-22 20:32:44 +02:00
parent e1d35f81bc
commit 9c2499a4a3
13 changed files with 2338 additions and 264 deletions
+11 -5
View File
@@ -23,6 +23,7 @@ pub enum TorrentState {
Paused,
Finished,
Error,
Stale,
}
#[derive(Debug, Clone, sqlx::FromRow)]
@@ -126,8 +127,12 @@ pub struct Progress<'a> {
pub error_message: Option<&'a str>,
}
pub async fn update_progress(pool: &PgPool, id: Uuid, progress: Progress<'_>) -> Result<()> {
sqlx::query(
pub async fn update_progress(
pool: &PgPool,
id: Uuid,
progress: Progress<'_>,
) -> Result<TorrentRow> {
let row = sqlx::query_as::<_, TorrentRow>(
"UPDATE torrents
SET name = COALESCE($2, name),
total_bytes = $3,
@@ -136,7 +141,8 @@ pub async fn update_progress(pool: &PgPool, id: Uuid, progress: Progress<'_>) ->
error_message = $6,
updated_at = now(),
completed_at = CASE WHEN $5 = 'finished' THEN now() ELSE completed_at END
WHERE id = $1",
WHERE id = $1
RETURNING id, info_hash, name, source, output_path, total_bytes, downloaded_bytes, state, error_message",
)
.bind(id)
.bind(progress.name)
@@ -144,7 +150,7 @@ pub async fn update_progress(pool: &PgPool, id: Uuid, progress: Progress<'_>) ->
.bind(progress.downloaded_bytes)
.bind(progress.state)
.bind(progress.error_message)
.execute(pool)
.fetch_one(pool)
.await?;
Ok(())
Ok(row)
}
+491 -27
View File
@@ -1,19 +1,20 @@
use std::collections::{HashMap, HashSet};
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use std::time::{Duration, SystemTime, UNIX_EPOCH};
use anyhow::{Context, Result};
use librqbit::api::TorrentIdOrHash;
use librqbit::{AddTorrent, AddTorrentOptions, ManagedTorrent, Session, TorrentStatsState};
use librqbit_core::Id20;
use sqlx::PgPool;
use tokio::sync::Mutex;
use tokio::sync::{Mutex, broadcast};
use tonic::{Request, Response, Status};
use tora_proto::{
AddRequest, AddResponse, ListRequest, ListResponse, PauseRequest, PauseResponse, RemoveRequest,
RemoveResponse, ResumeRequest, ResumeResponse, State as ProtoState, StatusRequest,
TorrentStatus, Torrents,
AddRequest, AddResponse, AddedInfo, ListRequest, ListResponse, Notification, NotificationKind,
NotificationsRequest, PauseRequest, PauseResponse, ProgressMark, RemoveRequest, RemoveResponse,
ResumeRequest, ResumeResponse, State as ProtoState, StateChanged, StatusRequest, TorrentStatus,
Torrents,
};
use tracing::{error, info, warn};
use uuid::Uuid;
@@ -24,6 +25,14 @@ use crate::source::SourceResolver;
const POLL_INTERVAL: Duration = Duration::from_secs(2);
const ADD_TIMEOUT: Duration = Duration::from_secs(60);
const BYTES_PER_MIB: f64 = 1_048_576.0;
/// `broadcast` channel capacity. Subscribers that fall more than this many
/// events behind receive `Lagged` errors and continue. Picked for ~5 s of
/// slack at the maximum realistic emit rate (~50 events/sec).
const EVENT_CHANNEL_CAPACITY: usize = 256;
/// Consecutive poller ticks a `Live && !finished` torrent must show no byte
/// progress AND zero connected peers before being flipped to `Stale`.
/// At a 2 s `POLL_INTERVAL`, this is 6 s of inactivity.
const STALE_TICKS_THRESHOLD: u32 = 3;
pub enum ResolveError {
NotFound,
@@ -38,6 +47,8 @@ pub struct TorrentManager {
source: SourceResolver,
tracked: Mutex<HashMap<Uuid, Arc<ManagedTorrent>>>,
adding: Mutex<HashSet<Uuid>>,
events: broadcast::Sender<Notification>,
poll_cache: Mutex<HashMap<Uuid, PollCache>>,
}
impl TorrentManager {
@@ -52,6 +63,7 @@ impl TorrentManager {
.await
.context("failed to create librqbit session")?;
let source = SourceResolver::new().context("failed to build source resolver")?;
let (events, _) = broadcast::channel(EVENT_CHANNEL_CAPACITY);
Ok(Arc::new(Self {
pool,
session,
@@ -59,9 +71,22 @@ impl TorrentManager {
source,
tracked: Mutex::new(HashMap::new()),
adding: Mutex::new(HashSet::new()),
events,
poll_cache: Mutex::new(HashMap::new()),
}))
}
pub fn subscribe(&self) -> broadcast::Receiver<Notification> {
self.events.subscribe()
}
pub(crate) fn notify(&self, mut notification: Notification) {
if notification.observed_at_unix_millis == 0 {
notification.observed_at_unix_millis = now_millis();
}
let _ = self.events.send(notification);
}
pub async fn add(&self, source: &str, output_dir: Option<&str>) -> Result<TorrentRow> {
let resolved = self.source.resolve(source).await?;
let stored_source = resolved.rewritten_source.as_deref().unwrap_or(source);
@@ -263,30 +288,158 @@ impl TorrentManager {
}
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,
}
};
// Snapshot stats under a brief lock; do all DB work and emission outside.
let snapshots: Vec<(Uuid, librqbit::TorrentStats, Option<String>)> = {
let tracked = self.tracked.lock().await;
tracked
.iter()
.map(|(id, handle)| (*id, handle.stats(), handle.name()))
.collect()
};
for (id, stats, name) in snapshots {
let raw_state = raw_state_from_stats(&stats);
let live = live_stats_from_stats(&stats);
let (final_state, prev_state_proto, percent, bytes_delta, state_changed) = self
.diff_and_update_cache(
id,
raw_state,
stats.progress_bytes,
live.peers,
stats.total_bytes,
)
.await;
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,
state: final_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");
let row = match db::update_progress(&self.pool, id, progress).await {
Ok(row) => row,
Err(err) => {
error!(id = %id, error = %err, "failed to persist torrent progress");
continue;
}
};
let status = row_to_status(row, live);
if state_changed {
self.notify(Notification {
observed_at_unix_millis: 0,
kind: NotificationKind::StateChanged as i32,
torrent: Some(status.clone()),
detail: Some(tora_proto::notification_detail::Detail::StateChanged(
StateChanged {
previous: prev_state_proto as i32,
current: status.state,
},
)),
});
}
if let Some(delta) = bytes_delta {
self.notify(Notification {
observed_at_unix_millis: 0,
kind: NotificationKind::Progress as i32,
torrent: Some(status),
detail: Some(tora_proto::notification_detail::Detail::Progress(
ProgressMark {
percent,
bytes_delta: delta,
},
)),
});
}
}
}
/// Updates `poll_cache` for `id` from this tick's observation and returns
/// the data needed to emit notifications. Returns
/// `(final_state, previous_proto_state, percent, Option<bytes_delta>, state_changed)`.
async fn diff_and_update_cache(
&self,
id: Uuid,
raw_state: TorrentState,
progress_bytes: u64,
live_peers: u32,
total_bytes: u64,
) -> (TorrentState, ProtoState, u32, Option<u64>, bool) {
let mut poll_cache = self.poll_cache.lock().await;
let cache = poll_cache
.entry(id)
.or_insert_with(|| PollCache::new(proto_state(raw_state)));
compute_diff(cache, raw_state, progress_bytes, live_peers, total_bytes)
}
}
/// Pure diff/staleness logic. Updates `cache` in place and returns
/// `(final_state, previous_proto_state, percent, Option<bytes_delta>, state_changed)`.
/// Extracted from `diff_and_update_cache` so it can be unit-tested without a
/// live TorrentManager (which needs Postgres + librqbit session).
fn compute_diff(
cache: &mut PollCache,
raw_state: TorrentState,
progress_bytes: u64,
live_peers: u32,
total_bytes: u64,
) -> (TorrentState, ProtoState, u32, Option<u64>, bool) {
if raw_state == TorrentState::Downloading {
if progress_bytes == cache.last_bytes {
cache.stale_ticks = cache.stale_ticks.saturating_add(1);
} else {
cache.stale_ticks = 0;
}
} else {
cache.stale_ticks = 0;
}
cache.last_bytes = progress_bytes;
let final_state = if raw_state == TorrentState::Downloading
&& cache.stale_ticks >= STALE_TICKS_THRESHOLD
&& live_peers == 0
{
TorrentState::Stale
} else {
raw_state
};
let prev_proto = cache.last_state;
let final_proto = proto_state(final_state);
let state_changed = prev_proto != final_proto;
cache.last_state = final_proto;
// u128 mul to avoid overflow on multi-exabyte totals; cap at 100.
let percent = if total_bytes == 0 {
0
} else {
((progress_bytes as u128 * 100 / total_bytes as u128) as u32).min(100)
};
let bytes_delta = if percent != cache.last_percent {
let delta = progress_bytes.saturating_sub(cache.bytes_at_last_progress);
cache.last_percent = percent;
cache.bytes_at_last_progress = progress_bytes;
Some(delta)
} else {
None
};
(final_state, prev_proto, percent, bytes_delta, state_changed)
}
fn raw_state_from_stats(stats: &librqbit::TorrentStats) -> TorrentState {
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,
}
}
}
@@ -301,6 +454,34 @@ struct LiveTorrentStats {
seeds: u32,
}
#[derive(Clone)]
struct PollCache {
last_state: ProtoState,
last_percent: u32,
last_bytes: u64,
bytes_at_last_progress: u64,
stale_ticks: u32,
}
impl PollCache {
fn new(state: ProtoState) -> Self {
Self {
last_state: state,
last_percent: 0,
last_bytes: 0,
bytes_at_last_progress: 0,
stale_ticks: 0,
}
}
}
fn now_millis() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_millis() as u64)
.unwrap_or(0)
}
async fn persist_error(pool: &PgPool, id: Uuid, message: &str) {
let progress = db::Progress {
name: None,
@@ -318,10 +499,13 @@ fn extract_live_stats(tracked: &HashMap<Uuid, Arc<ManagedTorrent>>, id: &Uuid) -
let Some(handle) = tracked.get(id) else {
return LiveTorrentStats::default();
};
let stats = handle.stats();
live_stats_from_stats(&handle.stats())
}
fn live_stats_from_stats(stats: &librqbit::TorrentStats) -> LiveTorrentStats {
let uploaded_bytes = stats.uploaded_bytes;
let Some(live) = stats.live else {
let Some(live) = stats.live.as_ref() else {
return LiveTorrentStats {
uploaded_bytes,
..Default::default()
@@ -364,6 +548,7 @@ fn proto_state(state: TorrentState) -> ProtoState {
TorrentState::Paused => ProtoState::Paused,
TorrentState::Finished => ProtoState::Finished,
TorrentState::Error => ProtoState::Error,
TorrentState::Stale => ProtoState::Stale,
}
}
@@ -419,9 +604,18 @@ impl Torrents for GrpcTorrents {
.add(&req.magnet, output_dir)
.await
.map_err(|err| Status::internal(err.to_string()))?;
Ok(Response::new(AddResponse {
id: row.id.to_string(),
}))
let id_str = row.id.to_string();
let status = row_to_status(row, LiveTorrentStats::default());
let source = status.source.clone();
self.manager.notify(Notification {
observed_at_unix_millis: 0,
kind: NotificationKind::TorrentAdded as i32,
torrent: Some(status),
detail: Some(tora_proto::notification_detail::Detail::Added(AddedInfo {
source,
})),
});
Ok(Response::new(AddResponse { id: id_str }))
}
async fn status(
@@ -468,6 +662,17 @@ impl Torrents for GrpcTorrents {
.map_err(|err| Status::internal(err.to_string()))?
.ok_or_else(|| Status::not_found("torrent not found"))?;
info!(id = %id, name = ?row.name, "torrent removed");
let status = row_to_status(row, LiveTorrentStats::default());
self.manager.notify(Notification {
observed_at_unix_millis: 0,
kind: NotificationKind::TorrentRemoved as i32,
torrent: Some(status),
detail: Some(tora_proto::notification_detail::Detail::Removed(
tora_proto::RemovedInfo {
files_deleted: req.delete_files,
},
)),
});
Ok(Response::new(RemoveResponse {}))
}
@@ -508,4 +713,263 @@ impl Torrents for GrpcTorrents {
info!(id = %id, name = ?row.name, "torrent resumed");
Ok(Response::new(ResumeResponse {}))
}
type NotificationsStream =
tokio_stream::wrappers::ReceiverStream<std::result::Result<Notification, Status>>;
async fn notifications(
&self,
request: Request<NotificationsRequest>,
) -> Result<Response<Self::NotificationsStream>, Status> {
use tokio::sync::mpsc;
use tokio_stream::wrappers::ReceiverStream;
let req = request.into_inner();
// Resolve id-prefix filter once at subscribe time. If any prefix fails
// to resolve, the whole subscription is rejected — partial filters
// would surprise the client.
let allowed_ids: Option<HashSet<Uuid>> = if req.ids.is_empty() {
None
} else {
let mut set = HashSet::new();
for id_str in &req.ids {
let uuid = self.manager.resolve_id(id_str).await.map_err(resolve_err)?;
set.insert(uuid);
}
Some(set)
};
let allowed_kinds: Option<HashSet<i32>> = if req.kinds.is_empty() {
None
} else {
Some(req.kinds.iter().copied().collect())
};
let mut rx = self.manager.subscribe();
let (tx, rx_stream) = mpsc::channel::<std::result::Result<Notification, Status>>(64);
tokio::spawn(async move {
loop {
match rx.recv().await {
Ok(notification) => {
if let Some(ref allowed) = allowed_ids
&& let Some(ref torrent) = notification.torrent
&& Uuid::parse_str(&torrent.id)
.ok()
.filter(|id| allowed.contains(id))
.is_none()
{
continue;
}
if let Some(ref kinds) = allowed_kinds
&& !kinds.contains(&notification.kind)
{
continue;
}
if tx.send(Ok(notification)).await.is_err() {
break;
}
}
Err(broadcast::error::RecvError::Lagged(n)) => {
warn!(skipped = n, "notifications subscriber lagged");
continue;
}
Err(broadcast::error::RecvError::Closed) => break,
}
}
});
Ok(Response::new(ReceiverStream::new(rx_stream)))
}
}
#[cfg(test)]
mod tests {
use super::*;
/// A `PollCache` primed as though a prior tick already recorded
/// `bytes` at `total` size, in the Downloading state. The first
/// `compute_diff` call from a test then represents the first
/// *observation* after that baseline (not an artificial advance from 0).
fn primed_cache(bytes: u64, total: u64) -> PollCache {
let percent = if total == 0 {
0
} else {
((bytes as u128 * 100 / total as u128) as u32).min(100)
};
PollCache {
last_state: ProtoState::Downloading,
last_percent: percent,
last_bytes: bytes,
bytes_at_last_progress: bytes,
stale_ticks: 0,
}
}
/// `(final_state, state_changed, percent, bytes_delta)`
fn simplify(
out: (TorrentState, ProtoState, u32, Option<u64>, bool),
) -> (TorrentState, bool, u32, Option<u64>) {
(out.0, out.4, out.2, out.3)
}
#[test]
fn below_stale_threshold_stays_downloading() {
let mut cache = primed_cache(100_000, 1_000_000);
// Tick 1: stale_ticks=1.
let out = compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
assert_eq!(simplify(out), (TorrentState::Downloading, false, 10, None));
// Tick 2: stale_ticks=2.
let out = compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
assert_eq!(simplify(out), (TorrentState::Downloading, false, 10, None));
}
#[test]
fn stale_triggers_at_threshold_with_zero_peers() {
let mut cache = primed_cache(100_000, 1_000_000);
// Tick 1: stale_ticks=1.
compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
// Tick 2: stale_ticks=2.
compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
// Tick 3: stale_ticks=3 → STALE.
let out = compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
assert_eq!(out.0, TorrentState::Stale);
assert_eq!(out.1, ProtoState::Downloading);
assert!(out.4);
}
#[test]
fn peers_above_zero_prevents_stale() {
let mut cache = primed_cache(100_000, 1_000_000);
for _ in 0..10 {
let out = compute_diff(&mut cache, TorrentState::Downloading, 100_000, 1, 1_000_000);
assert_eq!(out.0, TorrentState::Downloading);
}
}
#[test]
fn bytes_advance_resets_stale_counter() {
let mut cache = primed_cache(100_000, 1_000_000);
// Two no-progress ticks bring stale_ticks to 2.
compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
// Tick 3: bytes advance to 200 — stale_ticks resets to 0, stays Downloading.
let out = compute_diff(&mut cache, TorrentState::Downloading, 200_000, 0, 1_000_000);
assert_eq!(out.0, TorrentState::Downloading);
assert!(!out.4);
// Tick 4: no progress again, stale_ticks=1 (not threshold).
let out = compute_diff(&mut cache, TorrentState::Downloading, 200_000, 0, 1_000_000);
assert_eq!(out.0, TorrentState::Downloading);
}
#[test]
fn stale_recovers_on_bytes_advance() {
let mut cache = primed_cache(100_000, 1_000_000);
// Reach STALE state.
for _ in 0..3 {
compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
}
assert_eq!(cache.last_state, ProtoState::Stale);
// Bytes advance — should recover to Downloading and emit state_changed.
let out = compute_diff(&mut cache, TorrentState::Downloading, 200_000, 0, 1_000_000);
assert_eq!(out.0, TorrentState::Downloading);
assert_eq!(out.1, ProtoState::Stale);
assert!(out.4);
}
#[test]
fn stale_recovers_on_peer_connect() {
let mut cache = primed_cache(100_000, 1_000_000);
// Reach STALE state.
for _ in 0..3 {
compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
}
assert_eq!(cache.last_state, ProtoState::Stale);
// Peer connects (bytes unchanged).
let out = compute_diff(&mut cache, TorrentState::Downloading, 100_000, 1, 1_000_000);
assert_eq!(out.0, TorrentState::Downloading);
assert_eq!(out.1, ProtoState::Stale);
assert!(out.4);
}
#[test]
fn stale_persists_when_nothing_changes() {
let mut cache = primed_cache(100_000, 1_000_000);
for _ in 0..3 {
compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
}
assert_eq!(cache.last_state, ProtoState::Stale);
// Subsequent tick: still no progress, still no peers.
let out = compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
assert_eq!(out.0, TorrentState::Stale);
assert!(!out.4);
}
#[test]
fn pause_and_resume_emit_state_change() {
let mut cache = primed_cache(100_000, 1_000_000);
// Pause: raw_state changes from Downloading to Paused.
let out = compute_diff(&mut cache, TorrentState::Paused, 100_000, 0, 1_000_000);
assert_eq!(out.0, TorrentState::Paused);
assert!(out.4);
// Resume: back to Downloading.
let out = compute_diff(&mut cache, TorrentState::Downloading, 100_000, 0, 1_000_000);
assert_eq!(out.0, TorrentState::Downloading);
assert!(out.4);
}
#[test]
fn finished_takes_precedence_over_stale() {
let mut cache = primed_cache(100_000, 1_000_000);
let out = compute_diff(&mut cache, TorrentState::Finished, 1_000_000, 0, 1_000_000);
assert_eq!(out.0, TorrentState::Finished);
assert!(out.4);
}
#[test]
fn percent_change_emits_progress_with_delta() {
let mut cache = primed_cache(0, 1000);
// 0% → 10% (bytes from 0 to 100 of 1000).
let out = compute_diff(&mut cache, TorrentState::Downloading, 100, 1, 1000);
assert_eq!(out.2, 10);
assert_eq!(out.3, Some(100));
// 10% → 50% (bytes from 100 to 500 of 1000).
let out = compute_diff(&mut cache, TorrentState::Downloading, 500, 1, 1000);
assert_eq!(out.2, 50);
assert_eq!(out.3, Some(400));
}
#[test]
fn no_progress_emit_when_percent_unchanged() {
let mut cache = primed_cache(0, 1000);
// 0% → 10%.
compute_diff(&mut cache, TorrentState::Downloading, 100, 1, 1000);
// 10% again (bytes from 100 to 109, still 10%).
let out = compute_diff(&mut cache, TorrentState::Downloading, 109, 1, 1000);
assert_eq!(out.2, 10);
assert_eq!(out.3, None);
}
#[test]
fn zero_total_bytes_gives_zero_percent() {
let mut cache = primed_cache(0, 0);
let out = compute_diff(&mut cache, TorrentState::Downloading, 0, 0, 0);
assert_eq!(out.2, 0);
assert_eq!(out.3, None);
}
#[test]
fn first_tick_initializes_state_without_false_event() {
let mut cache = PollCache::new(ProtoState::Pending);
let out = compute_diff(&mut cache, TorrentState::Pending, 0, 0, 0);
assert!(!out.4);
}
#[test]
fn percent_caps_at_100_on_full_download() {
let mut cache = primed_cache(0, 1000);
let out = compute_diff(&mut cache, TorrentState::Downloading, 1000, 1, 1000);
assert_eq!(out.2, 100);
}
}