use std::{error::Error as StdError, io::Error, num::NonZeroUsize};
use kithara_assets::AssetStore;
use kithara_audio::AudioObserver;
use kithara_bufpool::HasPool;
use kithara_download::DownloaderEvent;
use kithara_events::{Envelope, EventBus, RecvError, ScopeLabel, TrackId};
use kithara_net::NetError;
use kithara_platform::{
CancelGroup, CancelToken,
maybe_send::MaybeSend,
sync::Arc,
time::Duration,
tokio,
tokio::{
runtime::Handle as RuntimeHandle,
sync::Semaphore,
task::{JoinHandle, spawn, spawn_on},
},
};
use kithara_play::{Resource, ResourceConfig, ResourceSrc, player::PlayerControl};
use kithara_test_utils::kithara;
use tracing::debug;
use crate::{
attempts::{LoadClass, Ticket},
error::QueueError,
event::TrackStatus,
track::{TrackSource, Tracks},
};
pub(crate) struct Loader<S>
where
S: HasPool<u8> + HasPool<f32> + Send + Sync + 'static,
{
interactive_lane: Arc<Semaphore>,
prefetch_lane: Arc<Semaphore>,
tracks: Arc<Tracks<S>>,
store: AssetStore<S>,
cancel: CancelToken,
runtime: Option<RuntimeHandle>,
player: PlayerControl<S>,
}
impl<S> Loader<S>
where
S: HasPool<u8> + HasPool<f32> + Send + Sync + 'static,
{
const HANG_TIMEOUT: Duration = Duration::from_secs(60);
pub(crate) fn new(
player: PlayerControl<S>,
runtime: Option<RuntimeHandle>,
store: AssetStore<S>,
max_concurrent_loads: NonZeroUsize,
tracks: Arc<Tracks<S>>,
cancel: CancelToken,
) -> Self {
Self {
cancel,
player,
runtime,
tracks,
store,
interactive_lane: Arc::new(Semaphore::new(1)),
prefetch_lane: Arc::new(Semaphore::new(max_concurrent_loads.get())),
}
}
pub(crate) fn attach_observer<O: AudioObserver>(&self, id: TrackId, observer: O) {
self.tracks.attach_observer(id, Box::new(observer));
}
fn attempt_config(
&self,
id: TrackId,
source: TrackSource<S>,
) -> Result<(ResourceConfig<S>, CancelToken), QueueError> {
let config = self.build_config(id, source)?;
let Some(cancel) = config.cancel().cloned() else {
return Err(QueueError::Resource(format!(
"track {id:?}: resource config missing per-track cancel"
)));
};
Ok((config, cancel))
}
pub(crate) fn build_config(
&self,
id: TrackId,
source: TrackSource<S>,
) -> Result<ResourceConfig<S>, QueueError> {
let mut config = match source {
TrackSource::Uri(url) => {
let src = ResourceSrc::parse(&url)
.map_err(|e| QueueError::InvalidUrl(format!("{url}: {e}")))?;
ResourceConfig::for_src(src)
.store(self.store.clone())
.build()
}
TrackSource::Config(boxed) => *boxed,
};
if config.bus().is_none() {
config.set_bus(self.player.bus().scoped_labeled(ScopeLabel {
track: Some(id),
..ScopeLabel::default()
}));
}
self.player.prepare_config(config).map_err(QueueError::from)
}
#[kithara::hang_watchdog(timeout = Self::HANG_TIMEOUT)]
async fn load(&self, id: TrackId, config: ResourceConfig<S>) -> Result<Resource, QueueError> {
let slow_watcher =
Self::watch_for_slow_status(id, config.bus().cloned(), Arc::clone(&self.tracks));
tokio::pin!(slow_watcher);
loop {
let observer = self.tracks.observer_relay(id);
let attempt =
async { Resource::new_observed(config.clone(), Box::new(observer)).await };
let result = tokio::select! {
biased;
result = attempt => result,
never = &mut slow_watcher => match never {},
};
let err = match result {
Ok(resource) => return Ok(resource),
Err(err) => err,
};
if !can_answer_later(&err, self.tracks.attempt_selected(id)) {
return Err(QueueError::Resource(format!("{err}")));
}
hang_tick!();
debug!(?id, error = %err, "load failed on a cause a later ask can answer; asking again");
}
}
pub(crate) fn promote_load(
self: &Arc<Self>,
id: TrackId,
source: TrackSource<S>,
) -> Option<JoinHandle<Result<Resource, QueueError>>> {
let (config, cancel) = match self.attempt_config(id, source) {
Ok(pair) => pair,
Err(err) => {
self.tracks
.set_status(id, TrackStatus::Failed(err.to_string()));
return None;
}
};
let ticket = self.tracks.promote_attempt(id, cancel.clone())?;
Some(self.spawn_attempt(ticket, config, cancel, LoadClass::Interactive))
}
fn spawn_attempt(
self: &Arc<Self>,
ticket: Ticket,
config: ResourceConfig<S>,
track_cancel: CancelToken,
class: LoadClass,
) -> JoinHandle<Result<Resource, QueueError>> {
let this = Arc::clone(self);
self.spawn(async move {
let id = ticket.id;
let cancel = CancelGroup::new(vec![track_cancel.clone(), this.cancel.clone()]);
let lane = match class {
LoadClass::Interactive => &this.interactive_lane,
LoadClass::Prefetch => &this.prefetch_lane,
};
kithara::probe_event!(admission_started, track_id = id.as_u64());
let permit = tokio::select! {
biased;
_ = Self::wait_and_cancel_track(&cancel, &track_cancel) => {
this.tracks.finish_attempt(&ticket, None);
return Err(QueueError::Cancelled(id));
}
permit = Arc::clone(lane).acquire_owned() => permit
.map_err(|e| QueueError::Resource(format!("semaphore closed: {e}")))?,
};
if !this.tracks.mark_loading(&ticket) {
drop(permit);
return Err(QueueError::Cancelled(id));
}
let result = tokio::select! {
biased;
_ = Self::wait_and_cancel_track(&cancel, &track_cancel) =>
Err(QueueError::Cancelled(id)),
result = this.load(id, config) => result,
};
drop(permit);
let failure = match &result {
Ok(_) | Err(QueueError::Cancelled(_)) => None,
Err(e) => Some(format!("{e}")),
};
this.tracks.finish_attempt(&ticket, failure);
result
})
}
#[track_caller]
pub(crate) fn spawn<F>(&self, future: F) -> JoinHandle<F::Output>
where
F: Future + MaybeSend + 'static,
F::Output: MaybeSend + 'static,
{
match &self.runtime {
Some(runtime) => spawn_on(runtime, future),
None => spawn(future),
}
}
pub(crate) fn spawn_load(
self: &Arc<Self>,
id: TrackId,
source: TrackSource<S>,
class: LoadClass,
) -> Option<JoinHandle<Result<Resource, QueueError>>> {
let (config, cancel) = match self.attempt_config(id, source) {
Ok(pair) => pair,
Err(err) => {
self.tracks
.set_status(id, TrackStatus::Failed(err.to_string()));
return None;
}
};
let ticket =
self.tracks
.begin_attempt(id, cancel.clone(), class == LoadClass::Interactive)?;
Some(self.spawn_attempt(ticket, config, cancel, class))
}
async fn wait_and_cancel_track(cancel: &CancelGroup, track_cancel: &CancelToken) {
cancel.cancelled().await;
track_cancel.cancel();
}
async fn watch_for_slow_status(
id: TrackId,
bus: Option<EventBus>,
tracks: Arc<Tracks<S>>,
) -> std::convert::Infallible {
let mut rx = match bus {
Some(b) => b.subscribe::<DownloaderEvent>(),
None => return std::future::pending().await,
};
let mut marked = false;
loop {
match rx.recv().await {
Ok(Envelope { event: ev, .. }) => {
if !marked && matches!(ev, DownloaderEvent::LoadSlow { .. }) {
tracks.set_status(id, TrackStatus::Slow);
marked = true;
}
}
Err(RecvError::Lagged(_)) => {}
Err(RecvError::Closed) => break,
}
}
std::future::pending().await
}
}
fn can_answer_later(error: &(dyn StdError + 'static), selected: bool) -> bool {
if !selected {
return false;
}
net_cause(error).is_some_and(NetError::can_answer_later)
}
fn net_cause<'e>(error: &'e (dyn StdError + 'static)) -> Option<&'e NetError> {
let mut current = Some(error);
while let Some(err) = current {
if let Some(net) = err.downcast_ref::<NetError>() {
return Some(net);
}
if let Some(net) = err
.downcast_ref::<Error>()
.and_then(Error::get_ref)
.and_then(|payload| net_cause(payload))
{
return Some(net);
}
current = err.source();
}
None
}
#[cfg(test)]
mod tests {
use std::{
future::{self, Future},
num::{NonZeroU16, NonZeroU32, NonZeroU64},
pin::pin,
sync::atomic::{AtomicUsize, Ordering},
task::{Context, Waker},
};
use kithara_assets::{AssetStore, StorageBackend};
use kithara_download::RequestId;
use kithara_events::EventBus;
use kithara_platform::{time::Duration, tokio::sync::oneshot};
use kithara_play::{
ArtifactSource, PlayWorker, PlayWorkerConfig, PlayerConfig, PlayerImpl, StreamShape, mock,
player::PlayerControlSource,
};
use kithara_test_utils::kithara;
use kithara_warp::WarpConfig;
use kithara_waveform::Waveform;
use super::*;
use crate::{
event::QueueEvent,
test_pools::{TestPools, pools},
track::TrackRecord,
};
struct CancelDropProbe {
state: Arc<AtomicUsize>,
cancel: CancelToken,
}
impl Drop for CancelDropProbe {
fn drop(&mut self) {
self.state
.store(usize::from(self.cancel.is_cancelled()), Ordering::SeqCst);
}
}
#[kithara::test]
fn a_refused_host_can_answer_later() {
let refused = NetError::RetryExhausted {
max_retries: 3,
source: Box::new(NetError::Status {
status: NonZeroU16::new(503).expect("503 is not zero"),
url: None,
body: Some("network offline".to_string()),
}),
};
assert!(can_answer_later(&Error::other(refused), true));
}
#[kithara::test]
fn a_vanished_host_can_answer_later() {
let gone = NetError::Network("connection closed".to_string());
assert!(can_answer_later(&Error::other(gone), true));
}
#[kithara::test]
fn a_stalled_transfer_is_not_asked_again() {
let stalled = NetError::RetryExhausted {
max_retries: 1,
source: Box::new(NetError::Timeout),
};
assert!(!can_answer_later(&Error::other(stalled), true));
}
#[kithara::test]
fn a_missing_resource_is_not_asked_again() {
let missing = NetError::Status {
status: NonZeroU16::new(404).expect("404 is not zero"),
url: None,
body: None,
};
assert!(!can_answer_later(&Error::other(missing), true));
}
#[kithara::test]
fn a_failure_with_no_network_cause_is_not_asked_again() {
let local = Error::other("unsupported container");
assert!(!can_answer_later(&local, true));
}
#[kithara::test]
fn an_unselected_attempt_is_not_asked_again() {
let refused = NetError::Network("connection refused".to_string());
assert!(!can_answer_later(&Error::other(refused), false));
}
#[derive(fieldwork::Fieldwork)]
#[fieldwork(with, vis = "")]
struct LoaderFixtureSpec {
cap: NonZeroUsize,
}
impl Default for LoaderFixtureSpec {
fn default() -> Self {
const CAP_3: NonZeroUsize = match NonZeroUsize::new(3) {
Some(n) => n,
None => unreachable!(),
};
Self { cap: CAP_3 }
}
}
#[kithara::test(tokio)]
async fn cancellation_precedes_in_flight_future_drop() {
let owner = CancelToken::root();
let queue_cancel = owner.child();
let track_cancel = owner.child();
let group = CancelGroup::new(vec![queue_cancel.clone(), track_cancel.clone()]);
let state = Arc::new(AtomicUsize::new(0));
let probe_state = Arc::clone(&state);
let probe_cancel = track_cancel.clone();
let (started_tx, started_rx) = oneshot::channel();
let in_flight = async move {
let _probe = CancelDropProbe {
cancel: probe_cancel,
state: probe_state,
};
let _ = started_tx.send(());
future::pending::<()>().await;
};
let canceller = spawn(async move {
started_rx.await.expect("in-flight future must start");
queue_cancel.cancel();
});
tokio::select! {
biased;
_ = Loader::<TestPools>::wait_and_cancel_track(&group, &track_cancel) => {}
() = in_flight => panic!("in-flight future must stay pending"),
}
canceller.await.expect("canceller task must not panic");
assert_eq!(state.load(Ordering::SeqCst), 1);
}
#[kithara::test(tokio)]
async fn cancellation_wakes_an_attempt_waiting_for_admission() {
let fixture = LoaderFixtureSpec::default()
.with_cap(NonZeroUsize::MIN)
.build();
let permit = Arc::clone(&fixture.loader.prefetch_lane)
.acquire_owned()
.await
.expect("loader keeps the prefetch semaphore open");
let id = TrackId::allocate();
let source = TrackSource::Uri("https://example.com/pending.mp3".into());
fixture
.tracks
.lock()
.push(TrackRecord::new(id, "pending".into(), source.clone()));
let handle = fixture
.loader
.spawn_load(id, source, LoadClass::Prefetch)
.expect("fresh track starts one load attempt");
assert!(fixture.tracks.lock().iter().any(|track| {
track.id == id && track.load.as_ref().is_some_and(|attempt| attempt.waiting)
}));
fixture.loader.cancel.cancel();
let result = kithara_platform::tokio::time::timeout(Duration::from_secs(1), handle)
.await
.expect("cancellation must wake the pending loader")
.expect("loader task must not panic");
assert!(matches!(
result,
Err(QueueError::Cancelled(cancelled)) if cancelled == id
));
drop(permit);
}
#[kithara::test(native)]
fn a_slow_watch_survives_a_bus_that_dropped_a_burst() {
const CAPACITY: usize = 4;
let bus = EventBus::new(CAPACITY);
let tracks = Arc::new(Tracks::<TestPools>::new(bus.clone()));
let id = TrackId::allocate();
tracks.lock().push(TrackRecord::new(
id,
"slow".into(),
TrackSource::Uri("https://example.com/slow.mp3".into()),
));
let mut watch = pin!(Loader::watch_for_slow_status(
id,
Some(bus.clone()),
Arc::clone(&tracks)
));
let mut cx = Context::from_waker(Waker::noop());
assert!(
watch.as_mut().poll(&mut cx).is_pending(),
"the watch must subscribe before the burst it has to survive"
);
let request_id = RequestId::new(NonZeroU64::MIN);
for _ in 0..=CAPACITY {
bus.publish(DownloaderEvent::RequestStarted {
request_id,
wait_in_queue: Duration::ZERO,
});
}
bus.publish(DownloaderEvent::LoadSlow {
request_id,
elapsed: Duration::ZERO,
});
assert!(
watch.as_mut().poll(&mut cx).is_pending(),
"the watch never completes: it ends only with the resource it races"
);
assert_eq!(
tracks.lock()[0].status,
TrackStatus::Slow,
"a dropped burst must not deafen the watch to the `LoadSlow` behind it"
);
}
struct LoaderFixture {
loader: Arc<Loader<TestPools>>,
tracks: Arc<Tracks<TestPools>>,
bus: EventBus,
_player: PlayerImpl<TestPools>,
}
impl LoaderFixtureSpec {
fn build(self) -> LoaderFixture {
let worker = PlayWorker::new(PlayWorkerConfig::builder(pools()).build());
let player = PlayerImpl::new(
PlayerConfig::builder()
.sample_rate(crate::queue::TEST_SAMPLE_RATE)
.worker(worker)
.session(crate::queue::test_session())
.build(),
);
let bus = player.bus().clone();
let tracks = Arc::new(Tracks::new(bus.clone()));
let store = AssetStore::builder(player.pools().clone()).build();
let loader = Arc::new(Loader::new(
player.control(),
RuntimeHandle::try_current().ok(),
store,
self.cap,
Arc::clone(&tracks),
CancelToken::root(),
));
LoaderFixture {
loader,
tracks,
bus,
_player: player,
}
}
}
#[kithara::test(tokio)]
async fn build_config_preserves_caller_supplied_config() {
let fixture = LoaderFixtureSpec::default().build();
let loader = &fixture.loader;
let supplied_store = AssetStore::builder(pools())
.backend(StorageBackend::Memory)
.build();
let Ok(src) = ResourceSrc::parse("https://example.com/a.mp3") else {
panic!("valid url");
};
let given = ResourceConfig::for_src(src)
.store(supplied_store.clone())
.preferred_peak_bitrate(321.0)
.build();
let Ok(returned) = loader.build_config(TrackId(1), TrackSource::Config(Box::new(given)))
else {
panic!("build_config should succeed");
};
assert!(
(returned.preferred_peak_bitrate() - 321.0).abs() < f64::EPSILON,
"caller-set fields must be preserved"
);
assert!(returned.store().is_same(&supplied_store));
assert!(!returned.store().is_same(&loader.store));
}
#[kithara::test(tokio)]
async fn build_config_forwards_a_prepared_artifact() {
let fixture = LoaderFixtureSpec::default().build();
let Ok(src) = ResourceSrc::parse("https://example.com/a.mp3") else {
panic!("valid url");
};
let Ok(grid) = ResourceSrc::parse("https://example.com/a.grid") else {
panic!("valid artifact url");
};
let given = ResourceConfig::for_src(src)
.store(AssetStore::builder(pools()).build())
.beat_grid(ArtifactSource::from(grid.clone()))
.waveform(ArtifactSource::Value(Arc::new(Waveform::default())))
.build();
let Ok(returned) = fixture
.loader
.build_config(TrackId(7), TrackSource::Config(Box::new(given)))
else {
panic!("build_config should succeed");
};
assert!(
matches!(returned.beat_grid(), Some(ArtifactSource::Source(src)) if *src == grid),
"a grid source must reach the resource untouched"
);
assert!(
matches!(returned.waveform(), Some(ArtifactSource::Value(_))),
"a caller-held waveform must reach the resource untouched"
);
}
#[kithara::test(tokio)]
async fn build_config_labels_default_bus_with_track_id() {
let fixture = LoaderFixtureSpec::default().build();
let mut rx = fixture.bus.subscribe::<QueueEvent>();
let Ok(config) = fixture.loader.build_config(
TrackId(42),
TrackSource::Uri("https://example.com/a.mp3".into()),
) else {
panic!("build_config should succeed");
};
let Some(bus) = config.bus() else {
panic!("build_config must inject a per-track bus");
};
assert!(config.store().is_same(&fixture.loader.store));
bus.publish(QueueEvent::QueueEnded);
let Ok(envelope) = rx.try_recv() else {
panic!("scoped publish must reach the root subscriber");
};
assert_eq!(envelope.meta.track, Some(TrackId(42)));
}
#[kithara::test(tokio)]
async fn build_config_invalid_uri_errors() {
let fixture = LoaderFixtureSpec::default().build();
let loader = &fixture.loader;
let Err(err) = loader.build_config(TrackId(1), TrackSource::Uri("not-a-url".into())) else {
panic!("should reject relative path");
};
assert!(matches!(err, QueueError::InvalidUrl(_)));
}
#[kithara::test(tokio, multi_thread)]
async fn prefetch_lane_caps_concurrent_loads() {
let cap = NonZeroUsize::new(2).expect("BUG: 2 > 0 is mathematically guaranteed");
let fixture = LoaderFixtureSpec::default().with_cap(cap).build();
let loader = &fixture.loader;
let in_flight = Arc::new(AtomicUsize::new(0));
let max_seen = Arc::new(AtomicUsize::new(0));
let mut handles = Vec::new();
for _ in 0..6 {
let sem = Arc::clone(&loader.prefetch_lane);
let in_flight = Arc::clone(&in_flight);
let max_seen = Arc::clone(&max_seen);
handles.push(spawn(async move {
let _permit = sem
.acquire_owned()
.await
.expect("BUG: semaphore not closed in test");
let cur = in_flight.fetch_add(1, Ordering::SeqCst) + 1;
max_seen.fetch_max(cur, Ordering::SeqCst);
time::sleep(Duration::from_millis(50)).await;
in_flight.fetch_sub(1, Ordering::SeqCst);
}));
}
for h in handles {
h.await.expect("BUG: spawned task panicked");
}
assert!(
max_seen.load(Ordering::SeqCst) <= 2,
"concurrency exceeded cap: {}",
max_seen.load(Ordering::SeqCst)
);
}
#[kithara::test(tokio, multi_thread)]
async fn spawn_load_bad_url_emits_failed_status() {
let fx = LoaderFixtureSpec::default().build();
fx.tracks.lock().push(TrackRecord::new(
TrackId(42),
String::new(),
TrackSource::Uri("not-a-url".into()),
));
let mut rx = fx.bus.subscribe();
let loader = fx.loader;
assert!(
loader
.spawn_load(
TrackId(42),
TrackSource::Uri("not-a-url".into()),
LoadClass::Prefetch,
)
.is_none()
);
let status = fx.tracks.lock()[0].status.clone();
assert!(matches!(&status, TrackStatus::Failed(_)));
let mut saw_failed = false;
for _ in 0..8 {
match time::timeout(Duration::from_millis(200), rx.recv()).await {
Ok(Ok(Envelope {
event:
QueueEvent::TrackStatusChanged {
id: TrackId(42),
status: TrackStatus::Loading,
},
..
})) => panic!("invalid config must not emit Loading"),
Ok(Ok(Envelope {
event:
QueueEvent::TrackStatusChanged {
id: TrackId(42),
status: event_status,
},
..
})) if event_status == status => saw_failed = true,
Ok(Ok(_)) => {}
Ok(Err(_)) | Err(_) => break,
}
}
assert!(saw_failed, "Failed status event missing");
}
#[kithara::test]
fn config_failure_without_runtime_updates_tracks_synchronously() {
let worker = PlayWorker::new(PlayWorkerConfig::builder(pools()).build());
let player = PlayerImpl::new(
PlayerConfig::builder()
.sample_rate(mock::SAMPLE_RATE)
.worker(worker)
.session(mock::session_with_shape(Some(StreamShape::new(
NonZeroU32::new(128).expect("fixture output block is non-zero"),
mock::SAMPLE_RATE,
))))
.warp(
WarpConfig::builder()
.render_quantum_frames(
NonZeroUsize::new(64).expect("fixture quantum is non-zero"),
)
.build(),
)
.build(),
);
let bus = player.bus().clone();
let tracks = Arc::new(Tracks::new(bus.clone()));
let loader = Arc::new(Loader::new(
player.control(),
None,
AssetStore::builder(player.pools().clone()).build(),
NonZeroUsize::MIN,
Arc::clone(&tracks),
CancelToken::root(),
));
let source = TrackSource::Uri("not a url".into());
let spawn_id = TrackId(42);
let promote_id = TrackId(43);
tracks.lock().extend([
TrackRecord::new(spawn_id, String::new(), source.clone()),
TrackRecord::new(promote_id, String::new(), source.clone()),
]);
let Err(expected) = loader.build_config(spawn_id, source.clone()) else {
panic!("fixture source must be rejected");
};
assert!(matches!(expected, QueueError::InvalidUrl(_)));
let reason = expected.to_string();
let mut rx = bus.subscribe::<QueueEvent>();
assert!(
loader
.spawn_load(spawn_id, source.clone(), LoadClass::Prefetch)
.is_none()
);
assert_eq!(tracks.lock()[0].status, TrackStatus::Failed(reason.clone()));
assert!(matches!(
rx.try_recv(),
Ok(Envelope {
event: QueueEvent::TrackStatusChanged { id, status },
..
}) if id == spawn_id && status == TrackStatus::Failed(reason.clone())
));
assert!(loader.promote_load(promote_id, source).is_none());
assert_eq!(tracks.lock()[1].status, TrackStatus::Failed(reason.clone()));
assert!(matches!(
rx.try_recv(),
Ok(Envelope {
event: QueueEvent::TrackStatusChanged { id, status },
..
}) if id == promote_id && status == TrackStatus::Failed(reason)
));
assert!(
rx.try_recv().is_err(),
"config failure must not emit Loading"
);
}
}