1use crate::any::{Any, AnyView};
20use crate::derive::{Object, ObjectRef};
21use crate::error::{Error, Result};
22use crate::function_internal::{AsPackedCallable, TupleAsPackedArgs};
23use crate::object::{Object, ObjectArc, ObjectCore};
24use tvm_ffi_sys::{
25 TVMFFIAny, TVMFFIByteArray, TVMFFIFunctionCell, TVMFFIFunctionCreate, TVMFFIFunctionGetGlobal,
26 TVMFFIFunctionSetGlobal, TVMFFIObjectHandle, TVMFFISafeCallType, TVMFFITypeIndex,
27};
28
29#[repr(C)]
31#[derive(Object)]
32#[type_key = "ffi.Function"]
33#[type_index(TVMFFITypeIndex::kTVMFFIFunction)]
34pub struct FunctionObj {
35 object: Object,
36 cell: TVMFFIFunctionCell,
37}
38
39#[derive(Clone, ObjectRef)]
41pub struct Function {
42 data: ObjectArc<FunctionObj>,
43}
44
45#[repr(C)]
54struct CallbackFunctionObjImpl<F: Fn(&[AnyView]) -> Result<Any> + 'static> {
55 function: FunctionObj,
56 callback: F,
57}
58
59impl<F: Fn(&[AnyView]) -> Result<Any> + 'static> CallbackFunctionObjImpl<F> {
60 pub fn from_callback(callback: F) -> Self {
61 Self {
62 function: FunctionObj {
63 object: Object::new(),
64 cell: TVMFFIFunctionCell {
65 safe_call: Self::invoke_callback,
67 cxx_call: std::ptr::null_mut(),
68 },
69 },
70 callback,
71 }
72 }
73
74 unsafe extern "C" fn invoke_callback(
75 handle: *mut std::ffi::c_void,
76 args: *const TVMFFIAny,
77 num_args: i32,
78 result: *mut TVMFFIAny,
79 ) -> i32 {
80 let this = &*(handle as *mut Self);
81 let packed_args = std::slice::from_raw_parts(args as *const AnyView, num_args as usize);
82 let ret_value = (this.callback)(packed_args);
83 match ret_value {
84 Ok(value) => {
85 *result = Any::into_raw_ffi_any(value);
86 0
87 }
88 Err(error) => {
89 Error::set_raised(&error);
90 -1
91 }
92 }
93 }
94}
95
96unsafe impl<F: Fn(&[AnyView]) -> Result<Any> + 'static> ObjectCore for CallbackFunctionObjImpl<F> {
97 const TYPE_KEY: &'static str = FunctionObj::TYPE_KEY;
98 const TYPE_DEPTH: i32 = FunctionObj::TYPE_DEPTH;
99 fn type_index() -> i32 {
100 FunctionObj::type_index()
101 }
102 unsafe fn object_header_mut(this: &mut Self) -> &mut tvm_ffi_sys::TVMFFIObject {
103 FunctionObj::object_header_mut(&mut this.function)
104 }
105}
106
107impl Function {
108 pub fn call_packed(&self, packed_args: &[AnyView]) -> Result<Any> {
110 unsafe {
111 let packed_args_ptr = packed_args.as_ptr() as *const TVMFFIAny;
112 let mut result = Any::new();
113 let ret_code = (self.data.cell.safe_call)(
114 ObjectArc::as_raw(&self.data) as *mut FunctionObj as *mut std::ffi::c_void,
115 packed_args_ptr,
116 packed_args.len() as i32,
117 Any::as_data_ptr(&mut result),
118 );
119 if ret_code == 0 {
120 Ok(result)
121 } else {
122 Err(Error::from_raised())
123 }
124 }
125 }
126
127 pub fn call_tuple<TupleType>(&self, tuple_args: TupleType) -> Result<Any>
128 where
129 TupleType: TupleAsPackedArgs,
130 {
131 const STACK_LEN: usize = 4;
144 let mut stack_args = [AnyView::new(); STACK_LEN];
145 let mut heap_args = Vec::<AnyView>::new();
146 let args_len = <TupleType as TupleAsPackedArgs>::LEN;
147 let packed_args: &mut [AnyView] = if args_len <= STACK_LEN {
149 &mut stack_args[..args_len]
150 } else {
151 heap_args.resize(args_len, AnyView::new());
152 &mut heap_args[..args_len]
153 };
154 (&tuple_args).fill_any_view(packed_args);
155 self.call_packed(packed_args)
156 }
157 pub fn call_tuple_with_len<const LEN: usize, TupleType>(
167 &self,
168 tuple_args: TupleType,
169 ) -> Result<Any>
170 where
171 TupleType: TupleAsPackedArgs,
172 {
173 let mut packed_args = [AnyView::new(); LEN];
174 (&tuple_args).fill_any_view(&mut packed_args);
175 self.call_packed(&packed_args)
176 }
177 pub fn get_global(name: &str) -> Result<Function> {
186 unsafe {
187 let name_arg = TVMFFIByteArray::from_str(name);
188 let mut result: TVMFFIObjectHandle = ::std::ptr::null_mut();
189 crate::check_safe_call!(TVMFFIFunctionGetGlobal(&name_arg, &mut result))?;
190 if result.is_null() {
191 crate::bail!(crate::error::RUNTIME_ERROR, "Function {} not found", name);
192 }
193 Ok(Self {
194 data: ObjectArc::<FunctionObj>::from_raw(result as *mut FunctionObj),
195 })
196 }
197 }
198
199 pub fn register_global(name: &str, func: Function) -> Result<()> {
207 unsafe {
208 let name_arg = TVMFFIByteArray::from_str(name);
209 let can_override = 0;
210 crate::check_safe_call!(TVMFFIFunctionSetGlobal(
211 &name_arg,
212 ObjectArc::as_raw(&func.data) as *mut FunctionObj as TVMFFIObjectHandle,
213 can_override
214 ))?;
215 Ok(())
216 }
217 }
218 pub fn from_packed<F>(func: F) -> Self
225 where
226 F: Fn(&[AnyView]) -> Result<Any> + 'static,
227 {
228 unsafe {
229 let callback_arc = ObjectArc::new(CallbackFunctionObjImpl::from_callback(func));
230 let func_arc = ObjectArc::<FunctionObj>::from_raw(
231 ObjectArc::into_raw(callback_arc) as *mut FunctionObj
232 );
233 Self { data: func_arc }
234 }
235 }
236
237 pub fn from_typed<F, I, O>(func: F) -> Self
244 where
245 F: AsPackedCallable<I, O> + 'static,
246 {
247 let closure = move |packed_args: &[AnyView]| -> Result<Any> {
248 let ret_value = func.call_packed(packed_args)?;
249 Ok(ret_value)
250 };
251 Self::from_packed(closure)
252 }
253
254 pub unsafe fn from_extern_c(
260 handle: *mut std::ffi::c_void,
261 safe_call: TVMFFISafeCallType,
262 deleter: Option<unsafe extern "C" fn(*mut std::ffi::c_void)>,
263 ) -> Self {
264 unsafe {
265 let mut out_handle: TVMFFIObjectHandle = std::ptr::null_mut();
266 crate::check_safe_call!(TVMFFIFunctionCreate(
267 handle,
268 safe_call,
269 deleter,
270 &mut out_handle
271 ))
272 .unwrap();
273 Self {
274 data: ObjectArc::<FunctionObj>::from_raw(out_handle as *mut FunctionObj),
275 }
276 }
277 }
278}