1mod gpu_encoders_manager;
4pub mod metrics;
5pub mod wgpu;
6
7use crate::plotter::gpu::gpu_encoders_manager::GpuRecordsEncoderManager;
8use crate::plotter::gpu::metrics::GpuPlotterMetrics;
9use crate::plotter::{Plotter, SectorPlottingProgress};
10use async_lock::{Mutex as AsyncMutex, Semaphore, SemaphoreGuardArc};
11use async_trait::async_trait;
12use bytes::Bytes;
13use event_listener_primitives::{Bag, HandlerId};
14use futures::channel::mpsc;
15use futures::stream::FuturesUnordered;
16use futures::{FutureExt, Sink, SinkExt, StreamExt, select, stream};
17use prometheus_client::registry::Registry;
18use std::error::Error;
19use std::fmt;
20use std::future::pending;
21use std::num::TryFromIntError;
22use std::pin::pin;
23use std::sync::Arc;
24use std::sync::atomic::{AtomicBool, Ordering};
25use std::task::Poll;
26use std::time::Instant;
27use subspace_core_primitives::PublicKey;
28use subspace_core_primitives::sectors::SectorIndex;
29use subspace_data_retrieval::piece_getter::PieceGetter;
30use subspace_erasure_coding::ErasureCoding;
31use subspace_farmer_components::FarmerProtocolInfo;
32use subspace_farmer_components::plotting::{
33 DownloadSectorOptions, EncodeSectorOptions, PlottingError, RecordsEncoder, download_sector,
34 encode_sector, write_sector,
35};
36use subspace_kzg::Kzg;
37use subspace_process::AsyncJoinOnDrop;
38use tokio::task::yield_now;
39use tracing::{Instrument, warn};
40
41pub type HandlerFn3<A, B, C> = Arc<dyn Fn(&A, &B, &C) + Send + Sync + 'static>;
43type Handler3<A, B, C> = Bag<HandlerFn3<A, B, C>, A, B, C>;
44
45#[derive(Default, Debug)]
46struct Handlers {
47 plotting_progress: Handler3<PublicKey, SectorIndex, SectorPlottingProgress>,
48}
49
50pub trait GpuRecordsEncoder: RecordsEncoder + fmt::Debug + Send {
52 const TYPE: &'static str;
54}
55
56pub struct GpuPlotter<PG, GRE> {
58 piece_getter: PG,
59 downloading_semaphore: Arc<Semaphore>,
60 gpu_records_encoders_manager: GpuRecordsEncoderManager<GRE>,
61 global_mutex: Arc<AsyncMutex<()>>,
62 kzg: Kzg,
63 erasure_coding: ErasureCoding,
64 handlers: Arc<Handlers>,
65 tasks_sender: mpsc::Sender<AsyncJoinOnDrop<()>>,
66 _background_tasks: AsyncJoinOnDrop<()>,
67 abort_early: Arc<AtomicBool>,
68 metrics: Option<Arc<GpuPlotterMetrics>>,
69}
70
71impl<PG, GRE> fmt::Debug for GpuPlotter<PG, GRE>
72where
73 GRE: GpuRecordsEncoder + 'static,
74{
75 #[inline]
76 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
77 f.debug_struct(&format!("GpuPlotter[type = {}]", GRE::TYPE))
78 .finish_non_exhaustive()
79 }
80}
81
82impl<PG, RE> Drop for GpuPlotter<PG, RE> {
83 #[inline]
84 fn drop(&mut self) {
85 self.abort_early.store(true, Ordering::Release);
86 self.tasks_sender.close_channel();
87 }
88}
89
90#[async_trait]
91impl<PG, GRE> Plotter for GpuPlotter<PG, GRE>
92where
93 PG: PieceGetter + Clone + Send + Sync + 'static,
94 GRE: GpuRecordsEncoder + 'static,
95{
96 async fn has_free_capacity(&self) -> Result<bool, String> {
97 Ok(self.downloading_semaphore.try_acquire().is_some())
98 }
99
100 async fn plot_sector(
101 &self,
102 public_key: PublicKey,
103 sector_index: SectorIndex,
104 farmer_protocol_info: FarmerProtocolInfo,
105 pieces_in_sector: u16,
106 _replotting: bool,
107 progress_sender: mpsc::Sender<SectorPlottingProgress>,
108 ) {
109 let start = Instant::now();
110
111 let downloading_permit = self.downloading_semaphore.acquire_arc().await;
114
115 self.plot_sector_internal(
116 start,
117 downloading_permit,
118 public_key,
119 sector_index,
120 farmer_protocol_info,
121 pieces_in_sector,
122 progress_sender,
123 )
124 .await
125 }
126
127 async fn try_plot_sector(
128 &self,
129 public_key: PublicKey,
130 sector_index: SectorIndex,
131 farmer_protocol_info: FarmerProtocolInfo,
132 pieces_in_sector: u16,
133 _replotting: bool,
134 progress_sender: mpsc::Sender<SectorPlottingProgress>,
135 ) -> bool {
136 let start = Instant::now();
137
138 let Some(downloading_permit) = self.downloading_semaphore.try_acquire_arc() else {
139 return false;
140 };
141
142 self.plot_sector_internal(
143 start,
144 downloading_permit,
145 public_key,
146 sector_index,
147 farmer_protocol_info,
148 pieces_in_sector,
149 progress_sender,
150 )
151 .await;
152
153 true
154 }
155}
156
157impl<PG, GRE> GpuPlotter<PG, GRE>
158where
159 PG: PieceGetter + Clone + Send + Sync + 'static,
160 GRE: GpuRecordsEncoder + 'static,
161{
162 pub fn new(
166 piece_getter: PG,
167 downloading_semaphore: Arc<Semaphore>,
168 gpu_records_encoders: Vec<GRE>,
169 global_mutex: Arc<AsyncMutex<()>>,
170 kzg: Kzg,
171 erasure_coding: ErasureCoding,
172 registry: Option<&mut Registry>,
173 ) -> Result<Self, TryFromIntError> {
174 let (tasks_sender, mut tasks_receiver) = mpsc::channel(1);
175
176 let background_tasks = AsyncJoinOnDrop::new(
178 tokio::spawn(async move {
179 let background_tasks = FuturesUnordered::new();
180 let mut background_tasks = pin!(background_tasks);
181 background_tasks.push(AsyncJoinOnDrop::new(tokio::spawn(pending::<()>()), true));
183
184 loop {
185 select! {
186 maybe_background_task = tasks_receiver.next().fuse() => {
187 let Some(background_task) = maybe_background_task else {
188 break;
189 };
190
191 background_tasks.push(background_task);
192 },
193 _ = background_tasks.select_next_some() => {
194 }
196 }
197 }
198 }),
199 true,
200 );
201
202 let abort_early = Arc::new(AtomicBool::new(false));
203 let gpu_records_encoders_manager = GpuRecordsEncoderManager::new(gpu_records_encoders)?;
204 let metrics = registry.map(|registry| {
205 Arc::new(GpuPlotterMetrics::new(
206 registry,
207 GRE::TYPE,
208 gpu_records_encoders_manager.gpu_records_encoders(),
209 ))
210 });
211
212 Ok(Self {
213 piece_getter,
214 downloading_semaphore,
215 gpu_records_encoders_manager,
216 global_mutex,
217 kzg,
218 erasure_coding,
219 handlers: Arc::default(),
220 tasks_sender,
221 _background_tasks: background_tasks,
222 abort_early,
223 metrics,
224 })
225 }
226
227 pub fn on_plotting_progress(
229 &self,
230 callback: HandlerFn3<PublicKey, SectorIndex, SectorPlottingProgress>,
231 ) -> HandlerId {
232 self.handlers.plotting_progress.add(callback)
233 }
234
235 #[allow(clippy::too_many_arguments)]
236 async fn plot_sector_internal<PS>(
237 &self,
238 start: Instant,
239 downloading_permit: SemaphoreGuardArc,
240 public_key: PublicKey,
241 sector_index: SectorIndex,
242 farmer_protocol_info: FarmerProtocolInfo,
243 pieces_in_sector: u16,
244 mut progress_sender: PS,
245 ) where
246 PS: Sink<SectorPlottingProgress> + Unpin + Send + 'static,
247 PS::Error: Error,
248 {
249 if let Some(metrics) = &self.metrics {
250 metrics.sector_plotting.inc();
251 }
252
253 let progress_updater = ProgressUpdater {
254 public_key,
255 sector_index,
256 handlers: Arc::clone(&self.handlers),
257 metrics: self.metrics.clone(),
258 };
259
260 let plotting_fut = {
261 let piece_getter = self.piece_getter.clone();
262 let gpu_records_encoders_manager = self.gpu_records_encoders_manager.clone();
263 let global_mutex = Arc::clone(&self.global_mutex);
264 let kzg = self.kzg.clone();
265 let erasure_coding = self.erasure_coding.clone();
266 let abort_early = Arc::clone(&self.abort_early);
267 let metrics = self.metrics.clone();
268
269 async move {
270 let downloaded_sector = {
272 if !progress_updater
273 .update_progress_and_events(
274 &mut progress_sender,
275 SectorPlottingProgress::Downloading,
276 )
277 .await
278 {
279 return;
280 }
281
282 global_mutex.lock().await;
284
285 let downloading_start = Instant::now();
286
287 let downloaded_sector_fut = download_sector(DownloadSectorOptions {
288 public_key: &public_key,
289 sector_index,
290 piece_getter: &piece_getter,
291 farmer_protocol_info,
292 kzg: &kzg,
293 erasure_coding: &erasure_coding,
294 pieces_in_sector,
295 });
296
297 let downloaded_sector = match downloaded_sector_fut.await {
298 Ok(downloaded_sector) => downloaded_sector,
299 Err(error) => {
300 warn!(%error, "Failed to download sector");
301
302 progress_updater
303 .update_progress_and_events(
304 &mut progress_sender,
305 SectorPlottingProgress::Error {
306 error: format!("Failed to download sector: {error}"),
307 },
308 )
309 .await;
310
311 return;
312 }
313 };
314
315 if !progress_updater
316 .update_progress_and_events(
317 &mut progress_sender,
318 SectorPlottingProgress::Downloaded(downloading_start.elapsed()),
319 )
320 .await
321 {
322 return;
323 }
324
325 downloaded_sector
326 };
327
328 let (sector, plotted_sector) = {
330 let mut records_encoder = gpu_records_encoders_manager.get_encoder().await;
331 if let Some(metrics) = &metrics {
332 metrics.plotting_capacity_used.inc();
333 }
334
335 yield_now().await;
337
338 if !progress_updater
339 .update_progress_and_events(
340 &mut progress_sender,
341 SectorPlottingProgress::Encoding,
342 )
343 .await
344 {
345 if let Some(metrics) = &metrics {
346 metrics.plotting_capacity_used.dec();
347 }
348 return;
349 }
350
351 let encoding_start = Instant::now();
352
353 let plotting_result = tokio::task::block_in_place(move || {
354 let encoded_sector = encode_sector(
355 downloaded_sector,
356 EncodeSectorOptions {
357 sector_index,
358 records_encoder: &mut *records_encoder,
359 abort_early: &abort_early,
360 },
361 )?;
362
363 if abort_early.load(Ordering::Acquire) {
364 return Err(PlottingError::AbortEarly);
365 }
366
367 drop(records_encoder);
368
369 let mut sector = Vec::new();
370
371 write_sector(&encoded_sector, &mut sector)?;
372
373 Ok((sector, encoded_sector.plotted_sector))
374 });
375
376 if let Some(metrics) = &metrics {
377 metrics.plotting_capacity_used.dec();
378 }
379
380 match plotting_result {
381 Ok(plotting_result) => {
382 if !progress_updater
383 .update_progress_and_events(
384 &mut progress_sender,
385 SectorPlottingProgress::Encoded(encoding_start.elapsed()),
386 )
387 .await
388 {
389 return;
390 }
391
392 plotting_result
393 }
394 Err(PlottingError::AbortEarly) => {
395 return;
396 }
397 Err(error) => {
398 progress_updater
399 .update_progress_and_events(
400 &mut progress_sender,
401 SectorPlottingProgress::Error {
402 error: format!("Failed to encode sector: {error}"),
403 },
404 )
405 .await;
406
407 return;
408 }
409 }
410 };
411
412 progress_updater
413 .update_progress_and_events(
414 &mut progress_sender,
415 SectorPlottingProgress::Finished {
416 plotted_sector,
417 time: start.elapsed(),
418 sector: Box::pin({
419 let mut sector = Some(Ok(Bytes::from(sector)));
420
421 stream::poll_fn(move |_cx| {
422 let _downloading_permit = &downloading_permit;
424
425 Poll::Ready(sector.take())
426 })
427 }),
428 },
429 )
430 .await;
431 }
432 };
433
434 let plotting_task =
436 AsyncJoinOnDrop::new(tokio::spawn(plotting_fut.in_current_span()), true);
437 if let Err(error) = self.tasks_sender.clone().send(plotting_task).await {
438 warn!(%error, "Failed to send plotting task");
439
440 let progress = SectorPlottingProgress::Error {
441 error: format!("Failed to send plotting task: {error}"),
442 };
443
444 self.handlers
445 .plotting_progress
446 .call_simple(&public_key, §or_index, &progress);
447 }
448 }
449}
450
451struct ProgressUpdater {
452 public_key: PublicKey,
453 sector_index: SectorIndex,
454 handlers: Arc<Handlers>,
455 metrics: Option<Arc<GpuPlotterMetrics>>,
456}
457
458impl ProgressUpdater {
459 async fn update_progress_and_events<PS>(
461 &self,
462 progress_sender: &mut PS,
463 progress: SectorPlottingProgress,
464 ) -> bool
465 where
466 PS: Sink<SectorPlottingProgress> + Unpin,
467 PS::Error: Error,
468 {
469 if let Some(metrics) = &self.metrics {
470 match &progress {
471 SectorPlottingProgress::Downloading => {
472 metrics.sector_downloading.inc();
473 }
474 SectorPlottingProgress::Downloaded(time) => {
475 metrics.sector_downloading_time.observe(time.as_secs_f64());
476 metrics.sector_downloaded.inc();
477 }
478 SectorPlottingProgress::Encoding => {
479 metrics.sector_encoding.inc();
480 }
481 SectorPlottingProgress::Encoded(time) => {
482 metrics.sector_encoding_time.observe(time.as_secs_f64());
483 metrics.sector_encoded.inc();
484 }
485 SectorPlottingProgress::Finished { time, .. } => {
486 metrics.sector_plotting_time.observe(time.as_secs_f64());
487 metrics.sector_plotted.inc();
488 }
489 SectorPlottingProgress::Error { .. } => {
490 metrics.sector_plotting_error.inc();
491 }
492 }
493 }
494 self.handlers.plotting_progress.call_simple(
495 &self.public_key,
496 &self.sector_index,
497 &progress,
498 );
499
500 if let Err(error) = progress_sender.send(progress).await {
501 warn!(%error, "Failed to send progress update");
502
503 false
504 } else {
505 true
506 }
507 }
508}