use super::api::types as cl;
use super::memory::*;
use super::{Device, Error, Platform, Queue, API};
use device::Error as DeviceError;
use device::{IDevice, MemorySync};
#[cfg(feature = "native")]
use frameworks::native::device::Cpu;
#[cfg(feature = "native")]
use frameworks::native::flatbox::FlatBox;
use std::any::Any;
use std::hash::{Hash, Hasher};
use std::{mem, ptr};
#[derive(Debug, Clone)]
pub struct Context {
id: isize,
devices: Vec<Device>,
queue: Option<Queue>,
}
#[derive(PartialEq, Debug, Clone, Copy)]
pub enum ContextProperties {
Platform(Platform),
InteropUserSync(bool),
}
#[derive(PartialEq, Debug, Copy, Clone)]
pub enum ContextInfoQuery {
ReferenceCount,
NumDevices,
Properties,
Devices,
}
#[derive(Clone, Debug)]
pub enum ContextInfo {
ReferenceCount(u32),
NumDevices(u32),
Properties(Vec<ContextProperties>),
Devices(Vec<Device>),
}
impl Context {
pub fn new(devices: Vec<Device>) -> Result<Context, Error> {
let callback = unsafe { mem::transmute(ptr::null::<fn()>()) };
let mut context = Context::from_c(
API::create_context(devices.clone(), ptr::null(), callback, ptr::null_mut())?,
devices.clone(),
);
context.queue_mut();
Ok(context)
}
pub fn from_c(id: cl::context_id, devices: Vec<Device>) -> Context {
Context {
id: id as isize,
devices: devices,
queue: None,
}
}
pub fn queue(&self) -> Option<&Queue> {
self.queue.as_ref()
}
pub fn queue_mut(&mut self) -> &mut Queue {
if self.queue.is_some() {
self.queue.as_mut().unwrap()
} else {
self.queue = Some(Queue::new(self, &self.hardwares()[0], None).unwrap());
self.queue.as_mut().unwrap()
}
}
pub fn get_context_info(&self, query: ContextInfoQuery) -> Result<ContextInfo, Error> {
API::get_context_info(self.id as cl::context_id, query)
}
pub fn id(&self) -> isize {
self.id
}
pub fn id_c(&self) -> cl::context_id {
self.id as cl::context_id
}
}
impl IDevice for Context {
type H = Device;
type M = Memory;
fn id(&self) -> &isize {
&self.id
}
fn hardwares(&self) -> &Vec<Device> {
&self.devices
}
fn alloc_memory(&self, size: usize) -> Result<Memory, DeviceError> {
Ok(Memory::new(self, size)?)
}
}
impl MemorySync for Context {
fn sync_in(
&self,
my_memory: &mut Any,
src_device: &Any,
src_memory: &Any,
) -> Result<(), DeviceError> {
if let Some(_) = src_device.downcast_ref::<Cpu>() {
let mut my_mem = my_memory.downcast_mut::<Memory>().unwrap();
let src_mem = src_memory.downcast_ref::<FlatBox>().unwrap();
API::write_to_memory(
self.queue().unwrap(),
my_mem,
true,
0,
src_mem.byte_size(),
src_mem.as_slice().as_ptr(),
&[],
)?;
Ok(())
} else {
Err(DeviceError::NoMemorySyncRoute)
}
}
fn sync_out(
&self,
my_memory: &Any,
dst_device: &Any,
dst_memory: &mut Any,
) -> Result<(), DeviceError> {
if let Some(_) = dst_device.downcast_ref::<Cpu>() {
let my_mem = my_memory.downcast_ref::<Memory>().unwrap();
let mut dst_mem = dst_memory.downcast_mut::<FlatBox>().unwrap();
API::read_from_memory(
self.queue().unwrap(),
my_mem,
true,
0,
dst_mem.byte_size(),
dst_mem.as_mut_slice().as_mut_ptr(),
&[],
)?;
Ok(())
} else {
Err(DeviceError::NoMemorySyncRoute)
}
}
}
impl PartialEq for Context {
fn eq(&self, other: &Self) -> bool {
self.id() == other.id()
}
}
impl Eq for Context {}
impl Hash for Context {
fn hash<H: Hasher>(&self, state: &mut H) {
self.id().hash(state);
}
}