Skip to main content

sim_platform_ubuntu_pc/
compute.rs

1// conformance: native compute probing stays capsule-owned and fails closed.
2
3//! Native compute discovery owned by the Ubuntu platform capsule.
4
5use std::{process::Command, sync::mpsc};
6
7use sim_lib_compute_auto::{ComputeDeviceIdentity, ComputeEvidenceKind};
8use sim_lib_compute_cuda::{
9    CudaLoadError, CudaProbePort, CudaRuntimeLoader, CudaRuntimeProbe, DynamicCudaLoader,
10};
11use sim_lib_compute_rocm::{
12    DynamicRocmLoader, RocmLoadError, RocmProbePort, RocmRuntimeLoader, RocmRuntimeProbe,
13};
14use sim_lib_compute_wgpu::{
15    AllocationAttempt, ProbeEvidence, ProbePolicy, RequestedWgpuProfile, TransferEvidence,
16    WgpuAdapterEvidence, WgpuAdapterProbe, WgpuAdapterRuntime, WgpuCapabilityEvidence,
17    WgpuDiscoveryError, WgpuLimitEvidence, WgpuProbePort,
18};
19use wgpu::{
20    BufferDescriptor, BufferUsages, DeviceDescriptor, ExperimentalFeatures, Features, Instance,
21    Limits, MapMode, MemoryHints, PollType,
22};
23
24/// Ubuntu implementation of the compute probe membrane.
25#[derive(Clone, Debug)]
26pub struct UbuntuComputeProbe {
27    cuda: CudaRuntimeLoader,
28    rocm: RocmRuntimeLoader,
29}
30
31impl UbuntuComputeProbe {
32    /// Uses Ubuntu's normal native library and adapter search policy.
33    #[must_use]
34    pub fn new() -> Self {
35        Self {
36            cuda: CudaRuntimeLoader::new(),
37            rocm: RocmRuntimeLoader::new().with_observed_gfx_targets(observed_gfx_targets()),
38        }
39    }
40}
41
42impl Default for UbuntuComputeProbe {
43    fn default() -> Self {
44        Self::new()
45    }
46}
47
48fn observed_gfx_targets() -> Vec<String> {
49    let Ok(output) = Command::new("rocm_agent_enumerator").output() else {
50        return Vec::new();
51    };
52    if !output.status.success() {
53        return Vec::new();
54    }
55    let Ok(stdout) = String::from_utf8(output.stdout) else {
56        return Vec::new();
57    };
58    parse_gfx_targets(&stdout)
59}
60
61fn parse_gfx_targets(stdout: &str) -> Vec<String> {
62    let mut targets = stdout
63        .lines()
64        .map(str::trim)
65        .filter(|line| line.starts_with("gfx") && *line != "gfx000")
66        .map(ToOwned::to_owned)
67        .collect::<Vec<_>>();
68    targets.sort();
69    targets.dedup();
70    targets
71}
72
73impl CudaProbePort for UbuntuComputeProbe {
74    fn probe_cuda(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
75        self.cuda.discover()
76    }
77}
78
79impl RocmProbePort for UbuntuComputeProbe {
80    fn probe_rocm(&self) -> Result<RocmRuntimeProbe, RocmLoadError> {
81        self.rocm.discover()
82    }
83}
84
85impl WgpuProbePort for UbuntuComputeProbe {
86    fn probe_wgpu(
87        &self,
88        policy: &ProbePolicy,
89    ) -> Result<Vec<WgpuAdapterRuntime>, WgpuDiscoveryError> {
90        let instance = Instance::default();
91        let adapters = pollster::block_on(instance.enumerate_adapters(policy.backends));
92        let mut runtimes = adapters
93            .into_iter()
94            .filter_map(|adapter| probe_adapter(&adapter, policy).ok())
95            .collect::<Vec<_>>();
96        runtimes.sort_by(|left, right| {
97            left.probe()
98                .adapter
99                .sort_key()
100                .cmp(&right.probe().adapter.sort_key())
101        });
102        Ok(runtimes)
103    }
104}
105
106fn probe_adapter(
107    adapter: &wgpu::Adapter,
108    policy: &ProbePolicy,
109) -> Result<WgpuAdapterRuntime, WgpuDiscoveryError> {
110    let info = adapter.get_info();
111    let backend = format!("{:?}", info.backend);
112    let identity = ComputeDeviceIdentity::new(info.name.clone(), "wgpu", backend.clone());
113    let required_features = adapter.features() & (Features::TIMESTAMP_QUERY | Features::SHADER_F16);
114    let required_limits = Limits::downlevel_defaults().using_resolution(adapter.limits());
115    let descriptor = DeviceDescriptor {
116        label: Some("sim-platform-ubuntu-compute-probe"),
117        required_features,
118        required_limits: required_limits.clone(),
119        experimental_features: ExperimentalFeatures::disabled(),
120        memory_hints: MemoryHints::Performance,
121        trace: wgpu::Trace::default(),
122    };
123    let (device, queue) = pollster::block_on(adapter.request_device(&descriptor))
124        .map_err(|error| WgpuDiscoveryError::new(format!("wgpu request_device failed: {error}")))?;
125    let bytes = policy.transfer_bytes.max(4).next_multiple_of(4);
126    let payload = (0..bytes)
127        .map(|index| (index % 251) as u8)
128        .collect::<Vec<_>>();
129    let source = device.create_buffer(&BufferDescriptor {
130        label: Some("sim-platform-wgpu-probe"),
131        size: bytes,
132        usage: BufferUsages::COPY_SRC | BufferUsages::COPY_DST,
133        mapped_at_creation: false,
134    });
135    let readback = device.create_buffer(&BufferDescriptor {
136        label: Some("sim-platform-wgpu-readback-probe"),
137        size: bytes,
138        usage: BufferUsages::COPY_DST | BufferUsages::MAP_READ,
139        mapped_at_creation: false,
140    });
141    queue.write_buffer(&source, 0, &payload);
142    let mut encoder = device.create_command_encoder(&wgpu::CommandEncoderDescriptor {
143        label: Some("sim-platform-wgpu-transfer-probe"),
144    });
145    encoder.copy_buffer_to_buffer(&source, 0, &readback, 0, bytes);
146    queue.submit([encoder.finish()]);
147    let (sender, receiver) = mpsc::channel();
148    readback.slice(..).map_async(MapMode::Read, move |result| {
149        let _ = sender.send(result);
150    });
151    device
152        .poll(PollType::wait_indefinitely())
153        .map_err(|error| WgpuDiscoveryError::new(format!("wgpu poll failed: {error}")))?;
154    receiver
155        .recv()
156        .map_err(|error| WgpuDiscoveryError::new(format!("wgpu map callback failed: {error}")))?
157        .map_err(|error| WgpuDiscoveryError::new(format!("wgpu map failed: {error}")))?;
158    let mapping_ok = readback
159        .slice(..)
160        .get_mapped_range()
161        .map_err(|error| WgpuDiscoveryError::new(format!("wgpu mapped range failed: {error}")))?
162        .as_ref()
163        == payload.as_slice();
164    readback.unmap();
165    let ceiling = device
166        .limits()
167        .max_buffer_size
168        .min(policy.max_allocation_probe_bytes)
169        .max(4);
170    let allocation = device.create_buffer(&BufferDescriptor {
171        label: Some("sim-platform-wgpu-allocation-probe"),
172        size: ceiling,
173        usage: BufferUsages::COPY_DST,
174        mapped_at_creation: false,
175    });
176    drop(allocation);
177    let probe = WgpuAdapterProbe {
178        evidence_kind: ComputeEvidenceKind::PhysicalDevice,
179        claimed_identity: Some(identity.clone()),
180        observed_identity: Some(identity),
181        adapter: WgpuAdapterEvidence {
182            ordinal: 0,
183            name: info.name,
184            backend,
185            adapter_type: format!("{:?}", info.device_type),
186            vendor: info.vendor,
187            device: info.device,
188            requested: RequestedWgpuProfile::from_parts(required_limits, required_features),
189            granted_limits: WgpuLimitEvidence::from_limits(&device.limits()),
190            granted_features: WgpuCapabilityEvidence::from_features(device.features()),
191        },
192        probe: ProbeEvidence {
193            transfer: TransferEvidence {
194                bytes,
195                transfer_ok: true,
196                mapping_ok,
197            },
198            allocation_attempts: vec![AllocationAttempt {
199                bytes: ceiling,
200                success: true,
201            }],
202        },
203    };
204    Ok(WgpuAdapterRuntime::new(probe, device, queue))
205}
206
207#[cfg(test)]
208mod tests {
209    use super::*;
210    use sim_lib_compute_cuda::ComputeCudaLib;
211    use sim_lib_compute_rocm::ComputeRocmLib;
212    use sim_lib_compute_wgpu::ComputeWgpuLib;
213
214    struct EmptyWgpuProbe;
215
216    impl WgpuProbePort for EmptyWgpuProbe {
217        fn probe_wgpu(
218            &self,
219            _policy: &ProbePolicy,
220        ) -> Result<Vec<WgpuAdapterRuntime>, WgpuDiscoveryError> {
221            Ok(Vec::new())
222        }
223    }
224
225    #[test]
226    fn amd_target_observation_is_filtered_sorted_and_deduplicated() {
227        assert_eq!(
228            parse_gfx_targets("gfx1151\ngfx000\nnoise\ngfx1103\ngfx1151\n"),
229            vec!["gfx1103", "gfx1151"]
230        );
231    }
232
233    #[test]
234    fn absent_native_runtimes_export_no_provider() {
235        let probe = UbuntuComputeProbe {
236            cuda: CudaRuntimeLoader::with_search_dirs_only(Vec::new()),
237            rocm: RocmRuntimeLoader::with_search_dirs_only(Vec::new()),
238        };
239        assert!(ComputeCudaLib::from_probe_port(&probe).is_err());
240        assert!(ComputeRocmLib::from_probe_port(&probe).is_err());
241        let wgpu =
242            ComputeWgpuLib::from_probe_port(&EmptyWgpuProbe, &ProbePolicy::default()).unwrap();
243        assert!(sim_kernel::Lib::manifest(&wgpu).exports.is_empty());
244    }
245}