use std::sync::Arc;
use axiolid_contracts::{BackendId, ExecutionOptions, GeomError, GeomResult, Operation};
use crate::device::matches_device;
use axiolid_pointcloud::PointCloud;
use axiolid_pointcloud_reconstruction_contract::{
conformance, PointcloudReconstruction, Reconstruction, ReconstructionRequest,
};
#[derive(Clone)]
struct RegisteredReconstruction {
priority: i32,
provider: Arc<dyn PointcloudReconstruction>,
}
impl core::fmt::Debug for RegisteredReconstruction {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("RegisteredReconstruction")
.field("priority", &self.priority)
.field("backend", &self.provider.descriptor().id)
.finish()
}
}
#[derive(Debug, Clone, Default)]
pub struct PointcloudReconstructionRegistry {
providers: Vec<RegisteredReconstruction>,
}
impl PointcloudReconstructionRegistry {
pub const fn new() -> Self {
Self {
providers: Vec::new(),
}
}
pub fn register<B>(&mut self, priority: i32, provider: B)
where
B: PointcloudReconstruction + 'static,
{
self.register_arc(priority, Arc::new(provider));
}
pub fn register_conformant<B>(
&mut self,
priority: i32,
provider: B,
) -> Result<(), Box<conformance::ConformanceReport>>
where
B: PointcloudReconstruction + 'static,
{
let report = conformance::run(&provider);
if !report.is_conformant() {
return Err(Box::new(report));
}
self.register_arc(priority, Arc::new(provider));
Ok(())
}
pub fn register_arc(&mut self, priority: i32, provider: Arc<dyn PointcloudReconstruction>) {
self.providers
.push(RegisteredReconstruction { priority, provider });
self.providers
.sort_by_key(|entry| std::cmp::Reverse(entry.priority));
}
pub fn providers(&self) -> impl Iterator<Item = &dyn PointcloudReconstruction> {
self.providers.iter().map(|entry| entry.provider.as_ref())
}
pub fn is_empty(&self) -> bool {
self.providers.is_empty()
}
pub fn reconstruct(
&self,
cloud: &PointCloud,
request: &ReconstructionRequest,
options: &ExecutionOptions,
) -> GeomResult<Reconstruction> {
let mut last_retryable = None;
let mut over_budget = None;
for entry in &self.providers {
let descriptor = entry.provider.descriptor();
if !matches_device(options.device(), descriptor.id, descriptor.target) {
continue;
}
if !entry
.provider
.scratch_requirement()
.fits_budget(options, cloud.len())
{
over_budget = Some(GeomError::BudgetExceeded { resource: "memory" });
continue;
}
match entry.provider.reconstruct(cloud, request, options) {
Ok(result) => return Ok(result),
Err(error @ (GeomError::Unsupported { .. } | GeomError::Unavailable { .. })) => {
last_retryable = Some(error);
}
Err(error) => return Err(error),
}
}
Err(last_retryable
.or(over_budget)
.unwrap_or(GeomError::Unsupported {
backend: BackendId::new("pointcloud-reconstruction-registry"),
operation: Operation::PointcloudReconstruction,
}))
}
}