Skip to main content

ocl_core/
extension_functions.rs

1#![allow(non_snake_case)]
2
3use crate::{
4    get_extension_function_address_for_platform, get_platform_info, Error, PlatformId,
5    PlatformInfo, PlatformInfoResult, Result,
6};
7use cl_sys::*;
8use std::ffi::c_void;
9use std::mem::transmute;
10
11#[derive(Default, Clone)]
12pub struct ExtensionFunctions {
13    // OpenGL
14    pub clGetGLContextInfoKHR: Option<clGetGLContextInfoKHR_fn>,
15
16    // D3D11
17    pub clGetDeviceIDsFromD3D11: Option<clGetDeviceIDsFromD3D11_fn>,
18    pub clCreateFromD3D11Buffer: Option<clCreateFromD3D11Buffer_fn>,
19    pub clCreateFromD3D11Texture2D: Option<clCreateFromD3D11Texture2D_fn>,
20    pub clCreateFromD3D11Texture3D: Option<clCreateFromD3D11Texture3D_fn>,
21    pub clEnqueueAcquireD3D11Objects: Option<clEnqueueAcquireD3D11Objects_fn>,
22    pub clEnqueueReleaseD3D11Objects: Option<clEnqueueReleaseD3D11Objects_fn>,
23}
24
25impl ExtensionFunctions {
26    pub fn resolve_all(platform: PlatformId) -> Result<Self> {
27        let extensions = match get_platform_info(platform, PlatformInfo::Extensions) {
28            Ok(PlatformInfoResult::Extensions(s)) => s,
29            Ok(_) => {
30                return Err(Error::EmptyInfoResult(
31                    crate::EmptyInfoResultError::Platform,
32                ));
33            }
34            Err(e) => {
35                return Err(e.into());
36            }
37        };
38
39        let supports_khr_d3d11 = extensions.contains("cl_khr_d3d11_sharing");
40        let supports_nv_d3d11 = extensions.contains("cl_nv_d3d11_sharing");
41
42        let mut functions = Self::default();
43        functions.clGetGLContextInfoKHR =
44            get_pointer(&platform, "clGetGLContextInfoKHR", "")?.map(|p| unsafe { transmute(p) });
45
46        if supports_nv_d3d11 || supports_khr_d3d11 {
47            let suffix = if supports_nv_d3d11 { "NV" } else { "KHR" };
48            functions.clGetDeviceIDsFromD3D11 =
49                get_pointer(&platform, "clGetDeviceIDsFromD3D11", suffix)?
50                    .map(|p| unsafe { transmute(p) });
51            functions.clCreateFromD3D11Buffer =
52                get_pointer(&platform, "clCreateFromD3D11Buffer", suffix)?
53                    .map(|p| unsafe { transmute(p) });
54            functions.clCreateFromD3D11Texture2D =
55                get_pointer(&platform, "clCreateFromD3D11Texture2D", suffix)?
56                    .map(|p| unsafe { transmute(p) });
57            functions.clCreateFromD3D11Texture3D =
58                get_pointer(&platform, "clCreateFromD3D11Texture3D", suffix)?
59                    .map(|p| unsafe { transmute(p) });
60            functions.clEnqueueAcquireD3D11Objects =
61                get_pointer(&platform, "clEnqueueAcquireD3D11Objects", suffix)?
62                    .map(|p| unsafe { transmute(p) });
63            functions.clEnqueueReleaseD3D11Objects =
64                get_pointer(&platform, "clEnqueueReleaseD3D11Objects", suffix)?
65                    .map(|p| unsafe { transmute(p) });
66        }
67        Ok(functions)
68    }
69}
70
71fn get_pointer(
72    platform: &PlatformId,
73    func_name: &str,
74    suffix: &str,
75) -> Result<Option<*mut c_void>> {
76    unsafe {
77        match get_extension_function_address_for_platform(
78            platform,
79            &format!("{}{}", func_name, suffix),
80            None,
81        ) {
82            Ok(pointer) => Ok(Some(pointer)),
83            Err(Error::ApiWrapper(_)) => Ok(None),
84            Err(e) => Err(e.into()),
85        }
86    }
87}
88
89impl std::fmt::Debug for ExtensionFunctions {
90    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
91        f.debug_struct("ExtensionFunctions")
92            .field(
93                "clGetGLContextInfoKHR",
94                &self.clGetGLContextInfoKHR.map(|x| x as *mut c_void),
95            )
96            .field(
97                "clGetDeviceIDsFromD3D11",
98                &self.clGetDeviceIDsFromD3D11.map(|x| x as *mut c_void),
99            )
100            .field(
101                "clCreateFromD3D11Buffer",
102                &self.clCreateFromD3D11Buffer.map(|x| x as *mut c_void),
103            )
104            .field(
105                "clCreateFromD3D11Texture2D",
106                &self.clCreateFromD3D11Texture2D.map(|x| x as *mut c_void),
107            )
108            .field(
109                "clCreateFromD3D11Texture3D",
110                &self.clCreateFromD3D11Texture3D.map(|x| x as *mut c_void),
111            )
112            .field(
113                "clEnqueueAcquireD3D11Objects",
114                &self.clEnqueueAcquireD3D11Objects.map(|x| x as *mut c_void),
115            )
116            .field(
117                "clEnqueueReleaseD3D11Objects",
118                &self.clEnqueueReleaseD3D11Objects.map(|x| x as *mut c_void),
119            )
120            .finish()
121    }
122}