Skip to main content

tvm_ffi/
function.rs

1/*
2 * Licensed to the Apache Software Foundation (ASF) under one
3 * or more contributor license agreements.  See the NOTICE file
4 * distributed with this work for additional information
5 * regarding copyright ownership.  The ASF licenses this file
6 * to you under the Apache License, Version 2.0 (the
7 * "License"); you may not use this file except in compliance
8 * with the License.  You may obtain a copy of the License at
9 *
10 *   http://www.apache.org/licenses/LICENSE-2.0
11 *
12 * Unless required by applicable law or agreed to in writing,
13 * software distributed under the License is distributed on an
14 * "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
15 * KIND, either express or implied.  See the License for the
16 * specific language governing permissions and limitations
17 * under the License.
18 */
19use 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 crate::type_traits::AnyCompatible;
25use tvm_ffi_sys::{
26    TVMFFIAny, TVMFFIByteArray, TVMFFIFunctionCell, TVMFFIFunctionCreate, TVMFFIFunctionGetGlobal,
27    TVMFFIFunctionSetGlobal, TVMFFIGetTypeInfo, TVMFFIObjectHandle, TVMFFISafeCallType,
28    TVMFFITypeIndex, TVMFFITypeKeyToIndex,
29};
30
31/// function object
32#[repr(C)]
33#[derive(Object)]
34#[type_key = "ffi.Function"]
35#[type_index(TVMFFITypeIndex::kTVMFFIFunction)]
36pub struct FunctionObj {
37    object: Object,
38    cell: TVMFFIFunctionCell,
39}
40
41/// Error reference class
42#[derive(Clone, ObjectRef)]
43pub struct Function {
44    data: ObjectArc<FunctionObj>,
45}
46
47//------------------------------------------------------------------------
48// CallbackFunctionObjImpl
49//------------------------------------------------------------------------
50/// Special helper class to hold a generic callback state as Object
51/// Logically this Impl can be viewed as a FunctionObj
52/// We can create an ObjectArc<CallbackFunctionObjImpl<F>> so the deleter
53/// can correctly delete the entire object including callback part
54/// then we will convert to ObjectArc<FunctionObj> to be used as function
55#[repr(C)]
56struct CallbackFunctionObjImpl<F: Fn(&[AnyView]) -> Result<Any> + 'static> {
57    function: FunctionObj,
58    callback: F,
59}
60
61impl<F: Fn(&[AnyView]) -> Result<Any> + 'static> CallbackFunctionObjImpl<F> {
62    pub fn from_callback(callback: F) -> Self {
63        Self {
64            function: FunctionObj {
65                object: Object::new(),
66                cell: TVMFFIFunctionCell {
67                    // specfic callback for F
68                    safe_call: Self::invoke_callback,
69                    cxx_call: std::ptr::null_mut(),
70                },
71            },
72            callback,
73        }
74    }
75
76    unsafe extern "C" fn invoke_callback(
77        handle: *mut std::ffi::c_void,
78        args: *const TVMFFIAny,
79        num_args: i32,
80        result: *mut TVMFFIAny,
81    ) -> i32 {
82        let this = &*(handle as *mut Self);
83        let packed_args = std::slice::from_raw_parts(args as *const AnyView, num_args as usize);
84        let ret_value = match std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
85            (this.callback)(packed_args)
86        })) {
87            Ok(ret_value) => ret_value,
88            Err(payload) => Err(crate::function_internal::panic_to_error(payload)),
89        };
90        match ret_value {
91            Ok(value) => {
92                *result = Any::into_raw_ffi_any(value);
93                0
94            }
95            Err(error) => {
96                Error::set_raised(&error);
97                -1
98            }
99        }
100    }
101}
102
103unsafe impl<F: Fn(&[AnyView]) -> Result<Any> + 'static> ObjectCore for CallbackFunctionObjImpl<F> {
104    const TYPE_KEY: &'static str = FunctionObj::TYPE_KEY;
105    const TYPE_DEPTH: i32 = FunctionObj::TYPE_DEPTH;
106    fn type_index() -> i32 {
107        FunctionObj::type_index()
108    }
109    unsafe fn object_header_mut(this: &mut Self) -> &mut tvm_ffi_sys::TVMFFIObject {
110        FunctionObj::object_header_mut(&mut this.function)
111    }
112}
113
114impl Function {
115    /// Call the function in packed format.
116    pub fn call_packed(&self, packed_args: &[AnyView]) -> Result<Any> {
117        unsafe {
118            let packed_args_ptr = packed_args.as_ptr() as *const TVMFFIAny;
119            let mut result = Any::new();
120            let ret_code = (self.data.cell.safe_call)(
121                ObjectArc::as_raw(&self.data) as *mut FunctionObj as *mut std::ffi::c_void,
122                packed_args_ptr,
123                packed_args.len() as i32,
124                Any::as_data_ptr(&mut result),
125            );
126            if ret_code == 0 {
127                Ok(result)
128            } else {
129                Err(Error::from_raised())
130            }
131        }
132    }
133
134    pub fn call_tuple<TupleType>(&self, tuple_args: TupleType) -> Result<Any>
135    where
136        TupleType: TupleAsPackedArgs,
137    {
138        // This is a workaround for Rust's requirement that stack allocation size
139        // must be known at compile time for generic types.
140        // While we know args_len is a constant, Rust doesn't allow us to directly
141        // declare [AnyView::new(); args_len] in generic contexts.
142        //
143        // We use a small vector optimization pattern:
144        // 1. First allocate a small stack buffer (stack_args)
145        // 2. If args_len exceeds STACK_LEN, allocate a heap buffer (heap_args)
146        // 3. Use the appropriate buffer based on size
147        //
148        // Since args_len is a compile-time constant, the compiler should optimize
149        // away the unused branch, making this approach efficient.
150        const STACK_LEN: usize = 4;
151        let mut stack_args = [AnyView::new(); STACK_LEN];
152        let mut heap_args = Vec::<AnyView>::new();
153        let args_len = <TupleType as TupleAsPackedArgs>::LEN;
154        // get packed arguments
155        let packed_args: &mut [AnyView] = if args_len <= STACK_LEN {
156            &mut stack_args[..args_len]
157        } else {
158            heap_args.resize(args_len, AnyView::new());
159            &mut heap_args[..args_len]
160        };
161        (&tuple_args).fill_any_view(packed_args);
162        self.call_packed(packed_args)
163    }
164    /// Call function with compile-time known argument count
165    /// This is an optimized version of call_tuple for when the argument count
166    /// is known at compile time, avoiding the small vector optimization overhead.
167    ///
168    /// # Arguments
169    /// * `tuple_args` - The tuple arguments
170    ///
171    /// # Returns
172    /// * `Any` - The result
173    pub fn call_tuple_with_len<const LEN: usize, TupleType>(
174        &self,
175        tuple_args: TupleType,
176    ) -> Result<Any>
177    where
178        TupleType: TupleAsPackedArgs,
179    {
180        let mut packed_args = [AnyView::new(); LEN];
181        (&tuple_args).fill_any_view(&mut packed_args);
182        self.call_packed(&packed_args)
183    }
184    /// Get global function by name
185    /// This function will throw an error if the function is not found.
186    ///
187    /// # Arguments
188    /// * `name` - The name of the function
189    ///
190    /// # Returns
191    /// * `Function` - The global function
192    pub fn get_global(name: &str) -> Result<Function> {
193        unsafe {
194            let name_arg = TVMFFIByteArray::from_str(name);
195            let mut result: TVMFFIObjectHandle = ::std::ptr::null_mut();
196            crate::check_safe_call!(TVMFFIFunctionGetGlobal(&name_arg, &mut result))?;
197            if result.is_null() {
198                crate::bail!(crate::error::RUNTIME_ERROR, "Function {} not found", name);
199            }
200            Ok(Self {
201                data: ObjectArc::<FunctionObj>::from_raw(result as *mut FunctionObj),
202            })
203        }
204    }
205
206    /// Look up a reflected method of a type by type index and method name
207    ///
208    /// Methods registered through the C++ reflection registry
209    /// (`refl::ObjectDef<T>().def(...)`) live in the per-type method table
210    /// rather than the global function table. Constructors registered via
211    /// `refl::init` are reachable under the reserved name `__ffi_init__`.
212    /// For instance methods, the first packed argument is the object itself.
213    ///
214    /// `type_index` must be a registered type index (e.g. obtained from a
215    /// live object via `Any::type_index` or from a type key); the underlying
216    /// C API treats an unregistered index as a fatal error.
217    ///
218    /// # Arguments
219    /// * `type_index` - The type index of the type that owns the method
220    /// * `method_name` - The name of the method
221    ///
222    /// # Returns
223    /// * `Function` - The reflected method
224    pub fn from_type_method(type_index: i32, method_name: &str) -> Result<Function> {
225        unsafe {
226            let type_info = TVMFFIGetTypeInfo(type_index);
227            if type_info.is_null() {
228                crate::bail!(
229                    crate::error::TYPE_ERROR,
230                    "Cannot find type info for type_index={}",
231                    type_index
232                );
233            }
234            let type_info = &*type_info;
235            for i in 0..type_info.num_methods as usize {
236                let method_info = &*type_info.methods.add(i);
237                if method_info.name.as_str() != method_name {
238                    continue;
239                }
240                if !<Function as AnyCompatible>::check_any_strict(&method_info.method) {
241                    crate::bail!(
242                        crate::error::TYPE_ERROR,
243                        "Method `{}` of type `{}` is not a Function",
244                        method_name,
245                        type_info.type_key.as_str()
246                    );
247                }
248                // the table entry stores the method as a non-owning AnyView;
249                // copy out a strong reference
250                return Ok(<Function as AnyCompatible>::copy_from_any_view_after_check(
251                    &method_info.method,
252                ));
253            }
254            crate::bail!(
255                crate::error::TYPE_ERROR,
256                "Cannot find method `{}` of type `{}`",
257                method_name,
258                type_info.type_key.as_str()
259            );
260        }
261    }
262
263    /// Look up a function-valued attribute for a concrete runtime type.
264    ///
265    /// Type attributes are not inherited from base types.
266    pub fn from_type_attr(type_index: i32, attr_name: &str) -> Result<Function> {
267        let value = crate::reflection::get_type_attr(type_index, attr_name).ok_or_else(|| {
268            crate::error::Error::new(
269                crate::error::TYPE_ERROR,
270                &format!(
271                    "Cannot find type attribute `{}` for type_index={}",
272                    attr_name, type_index
273                ),
274                "",
275            )
276        })?;
277        Function::try_from(value)
278    }
279
280    /// Look up a reflected method of a type by type key and method name
281    ///
282    /// Same as [`Function::from_type_method`], but resolves `type_key` to a
283    /// type index first.
284    ///
285    /// # Arguments
286    /// * `type_key` - The type key of the type that owns the method
287    /// * `method_name` - The name of the method
288    ///
289    /// # Returns
290    /// * `Function` - The reflected method
291    pub fn from_type_key_method(type_key: &str, method_name: &str) -> Result<Function> {
292        unsafe {
293            let type_key_arg = TVMFFIByteArray::from_str(type_key);
294            let mut type_index: i32 = 0;
295            crate::check_safe_call!(TVMFFITypeKeyToIndex(&type_key_arg, &mut type_index))?;
296            Self::from_type_method(type_index, method_name)
297        }
298    }
299
300    /// Register a function as a global function
301    /// # Arguments
302    /// * `name` - The name of the function
303    /// * `func` - The function to register
304    ///
305    /// # Returns
306    /// * `Result<()>` - The result of the registration
307    pub fn register_global(name: &str, func: Function) -> Result<()> {
308        unsafe {
309            let name_arg = TVMFFIByteArray::from_str(name);
310            let can_override = 0;
311            crate::check_safe_call!(TVMFFIFunctionSetGlobal(
312                &name_arg,
313                ObjectArc::as_raw(&func.data) as *mut FunctionObj as TVMFFIObjectHandle,
314                can_override
315            ))?;
316            Ok(())
317        }
318    }
319    /// Construct a function from a packed function
320    /// # Arguments
321    /// * `func` - The packed function in signature of `Fn(&[AnyView]) -> Result<Any>`
322    ///
323    /// # Returns
324    /// * `Function` - The function
325    ///
326    /// Report errors by returning them. A panic in `func` is caught and raised
327    /// to the caller as an `InternalError`, but panicking is discouraged.
328    pub fn from_packed<F>(func: F) -> Self
329    where
330        F: Fn(&[AnyView]) -> Result<Any> + 'static,
331    {
332        unsafe {
333            let callback_arc = ObjectArc::new(CallbackFunctionObjImpl::from_callback(func));
334            let func_arc = ObjectArc::<FunctionObj>::from_raw(
335                ObjectArc::into_raw(callback_arc) as *mut FunctionObj
336            );
337            Self { data: func_arc }
338        }
339    }
340
341    /// Construct a function from a typed function
342    /// # Arguments
343    /// * `func` - The typed function with function signature of `F(T0, T1, ...) -> Result<O>`
344    ///
345    /// # Returns
346    /// * `Function` - The function
347    ///
348    /// Report errors by returning them. A panic in `func` is caught and raised
349    /// to the caller as an `InternalError`, but panicking is discouraged.
350    pub fn from_typed<F, I, O>(func: F) -> Self
351    where
352        F: AsPackedCallable<I, O> + 'static,
353    {
354        let closure = move |packed_args: &[AnyView]| -> Result<Any> {
355            let ret_value = func.call_packed(packed_args)?;
356            Ok(ret_value)
357        };
358        Self::from_packed(closure)
359    }
360
361    /// # Safety
362    ///
363    /// `handle` must be a valid pointer (or null) that is compatible with
364    /// `safe_call` and `deleter`. The caller must ensure the handle outlives
365    /// the returned `Function` (or that `deleter` properly frees it).
366    pub unsafe fn from_extern_c(
367        handle: *mut std::ffi::c_void,
368        safe_call: TVMFFISafeCallType,
369        deleter: Option<unsafe extern "C" fn(*mut std::ffi::c_void)>,
370    ) -> Self {
371        unsafe {
372            let mut out_handle: TVMFFIObjectHandle = std::ptr::null_mut();
373            crate::check_safe_call!(TVMFFIFunctionCreate(
374                handle,
375                safe_call,
376                deleter,
377                &mut out_handle
378            ))
379            .unwrap();
380            Self {
381                data: ObjectArc::<FunctionObj>::from_raw(out_handle as *mut FunctionObj),
382            }
383        }
384    }
385}