use crate::abi::{FfiReturn, FfiStatus, FfiStr, GenericValue, HostApi};
use crate::{PluginError, ValueKind};
macro_rules! host_error_constructors {
($($(#[$doc:meta])* $name:ident => $class_name:literal),* $(,)?) => {
$(
$(#[$doc])*
#[must_use]
pub fn $name(&self, message: &str) -> PluginError {
self.error($class_name, message)
}
)*
};
}
pub struct Host<'a> {
api: &'a HostApi,
}
#[derive(Debug, Clone, Copy)]
pub enum ArgValue<'h> {
Nil,
Bool(bool),
Int(i64),
BigInt(GenericValue),
Float(f64),
Rational(GenericValue),
Str(&'h str),
List(GenericValue),
Tuple(GenericValue),
Dict(GenericValue),
Set(GenericValue),
Range(GenericValue),
StopIteration,
Instance(GenericValue),
Class(GenericValue),
Function(GenericValue),
Module(GenericValue),
Exception(GenericValue),
Generator(GenericValue),
Iterator(GenericValue),
Other(GenericValue),
}
impl<'a> Host<'a> {
#[doc(hidden)]
#[must_use]
pub const fn new(api: &'a HostApi) -> Self {
Self { api }
}
#[must_use]
pub fn kind(&self, value: GenericValue) -> ValueKind {
ValueKind::from_u32((self.api.value_kind)(self.api.ctx, value))
}
#[must_use]
pub fn decode(&self, value: GenericValue) -> ArgValue<'_> {
match self.kind(value) {
ValueKind::Nil => ArgValue::Nil,
ValueKind::Bool => ArgValue::Bool(self.as_bool(value).unwrap_or_default()),
ValueKind::Int => ArgValue::Int(self.as_int(value).unwrap_or_default()),
ValueKind::BigInt => ArgValue::BigInt(value),
ValueKind::Float => ArgValue::Float(self.as_float(value).unwrap_or_default()),
ValueKind::Rational => ArgValue::Rational(value),
ValueKind::String => ArgValue::Str(self.as_str(value).unwrap_or_default()),
ValueKind::List => ArgValue::List(value),
ValueKind::Tuple => ArgValue::Tuple(value),
ValueKind::Dict => ArgValue::Dict(value),
ValueKind::Set => ArgValue::Set(value),
ValueKind::Range => ArgValue::Range(value),
ValueKind::StopIteration => ArgValue::StopIteration,
ValueKind::Instance => ArgValue::Instance(value),
ValueKind::Class => ArgValue::Class(value),
ValueKind::Function => ArgValue::Function(value),
ValueKind::Module => ArgValue::Module(value),
ValueKind::Exception => ArgValue::Exception(value),
ValueKind::Generator => ArgValue::Generator(value),
ValueKind::Iterator => ArgValue::Iterator(value),
ValueKind::Other => ArgValue::Other(value),
}
}
#[must_use]
pub fn as_bool(&self, value: GenericValue) -> Option<bool> {
let mut out = false;
(self.api.bool_get)(self.api.ctx, value, &raw mut out).then_some(out)
}
#[must_use]
pub fn as_int(&self, value: GenericValue) -> Option<i64> {
let mut out = 0i64;
(self.api.int_get)(self.api.ctx, value, &raw mut out).then_some(out)
}
#[must_use]
pub fn as_float(&self, value: GenericValue) -> Option<f64> {
let mut out = 0f64;
(self.api.float_get)(self.api.ctx, value, &raw mut out).then_some(out)
}
#[must_use]
pub fn as_str(&self, value: GenericValue) -> Option<&str> {
let mut out = FfiStr::null();
if !(self.api.string_get)(self.api.ctx, value, &raw mut out) {
return None;
}
if out.ptr.is_null() {
return None;
}
let bytes = unsafe { core::slice::from_raw_parts(out.ptr, out.len) };
core::str::from_utf8(bytes).ok()
}
#[must_use]
pub fn list_len(&self, value: GenericValue) -> Option<usize> {
let mut out = 0usize;
(self.api.list_len)(self.api.ctx, value, &raw mut out).then_some(out)
}
pub fn list_get(&self, value: GenericValue, index: usize) -> Result<GenericValue, PluginError> {
self.ffi_result((self.api.list_get)(self.api.ctx, value, index))
}
#[must_use]
pub fn tuple_len(&self, value: GenericValue) -> Option<usize> {
let mut out = 0usize;
(self.api.tuple_len)(self.api.ctx, value, &raw mut out).then_some(out)
}
pub fn tuple_get(
&self,
value: GenericValue,
index: usize,
) -> Result<GenericValue, PluginError> {
self.ffi_result((self.api.tuple_get)(self.api.ctx, value, index))
}
#[must_use]
pub fn dict_len(&self, value: GenericValue) -> Option<usize> {
let mut out = 0usize;
(self.api.dict_len)(self.api.ctx, value, &raw mut out).then_some(out)
}
#[must_use]
pub fn set_len(&self, value: GenericValue) -> Option<usize> {
let mut out = 0usize;
(self.api.set_len)(self.api.ctx, value, &raw mut out).then_some(out)
}
pub fn builtin(&self, name: &str) -> Result<GenericValue, PluginError> {
self.ffi_result((self.api.builtin_get)(self.api.ctx, Self::ffi_str(name)))
}
pub fn is_instance(
&self,
value: GenericValue,
class: GenericValue,
) -> Result<bool, PluginError> {
let result = self.ffi_result((self.api.is_instance)(self.api.ctx, value, class))?;
Ok(self.as_bool(result).unwrap_or_default())
}
pub fn class_of(&self, value: GenericValue) -> Result<GenericValue, PluginError> {
self.ffi_result((self.api.class_of)(self.api.ctx, value))
}
pub fn attr_get(
&self,
receiver: GenericValue,
name: &str,
) -> Result<GenericValue, PluginError> {
self.ffi_result((self.api.attr_get)(
self.api.ctx,
receiver,
Self::ffi_str(name),
))
}
pub fn attr_set(
&self,
receiver: GenericValue,
name: &str,
value: GenericValue,
) -> Result<(), PluginError> {
self.ffi_result((self.api.attr_set)(
self.api.ctx,
receiver,
Self::ffi_str(name),
value,
))
.map(|_| ())
}
pub fn attr_has(&self, receiver: GenericValue, name: &str) -> Result<bool, PluginError> {
let result = self.ffi_result((self.api.attr_has)(
self.api.ctx,
receiver,
Self::ffi_str(name),
))?;
Ok(self.as_bool(result).unwrap_or_default())
}
#[must_use]
pub fn make_nil(&self) -> GenericValue {
(self.api.nil_new)(self.api.ctx)
}
#[must_use]
pub fn make_bool(&self, value: bool) -> GenericValue {
(self.api.bool_new)(self.api.ctx, value)
}
#[must_use]
pub fn make_int(&self, value: i64) -> GenericValue {
(self.api.int_new)(self.api.ctx, value)
}
#[must_use]
pub fn make_float(&self, value: f64) -> GenericValue {
(self.api.float_new)(self.api.ctx, value)
}
#[must_use]
pub fn make_str(&self, value: &str) -> GenericValue {
let ffi = FfiStr {
ptr: value.as_ptr(),
len: value.len(),
};
self.ffi_result((self.api.string_new)(self.api.ctx, ffi))
.expect("host rejected a valid UTF-8 string")
}
#[must_use]
pub fn make_list(&self) -> GenericValue {
(self.api.list_new)(self.api.ctx)
}
pub fn list_push(&self, list: GenericValue, item: GenericValue) -> Result<(), PluginError> {
self.ffi_result((self.api.list_push)(self.api.ctx, list, item))
.map(|_| ())
}
pub fn list_set(
&self,
list: GenericValue,
index: usize,
value: GenericValue,
) -> Result<(), PluginError> {
self.ffi_result((self.api.list_set)(self.api.ctx, list, index, value))
.map(|_| ())
}
pub fn make_exception(
&self,
class: GenericValue,
message: &str,
) -> Result<GenericValue, PluginError> {
self.ffi_result((self.api.exception_new)(
self.api.ctx,
class,
Self::ffi_str(message),
))
}
fn error(&self, class_name: &str, message: &str) -> PluginError {
let result = self
.builtin(class_name)
.and_then(|class| self.make_exception(class, message))
.or_else(|error| {
if matches!(error, PluginError::Fatal) {
return Err(error);
}
let class = self.builtin("Exception")?;
self.make_exception(class, message)
});
match result {
Ok(exception) => PluginError::Exception(exception),
Err(PluginError::Fatal) => PluginError::Fatal,
Err(_) => PluginError::Exception(self.make_nil()),
}
}
host_error_constructors!(
exception => "Exception",
type_error => "TypeError",
value_error => "ValueError",
name_error => "NameError",
const_reassignment_error => "ConstReassignmentError",
attribute_error => "AttributeError",
import_error => "ImportError",
assertion_error => "AssertionError",
io_error => "IoError",
key_error => "KeyError",
index_error => "IndexError",
);
#[must_use]
pub fn display(&self, value: GenericValue) -> GenericValue {
(self.api.value_display)(self.api.ctx, value)
}
#[must_use]
pub fn display_string(&self, value: GenericValue) -> String {
let displayed = self.display(value);
self.as_str(displayed).unwrap_or_default().to_owned()
}
pub fn call(
&mut self,
callee: GenericValue,
args: &[GenericValue],
) -> Result<GenericValue, PluginError> {
self.ffi_result((self.api.call_value)(
self.api.ctx,
callee,
args.as_ptr(),
args.len(),
))
}
pub fn invoke(
&mut self,
receiver: GenericValue,
name: &str,
args: &[GenericValue],
) -> Result<GenericValue, PluginError> {
let name = FfiStr {
ptr: name.as_ptr(),
len: name.len(),
};
self.ffi_result((self.api.invoke_method)(
self.api.ctx,
receiver,
name,
args.as_ptr(),
args.len(),
))
}
pub fn to_str(&mut self, value: GenericValue) -> Result<GenericValue, PluginError> {
self.ffi_result((self.api.value_str)(self.api.ctx, value))
}
pub fn dict_get(
&mut self,
dict: GenericValue,
key: GenericValue,
) -> Result<GenericValue, PluginError> {
self.ffi_result((self.api.dict_get)(self.api.ctx, dict, key))
}
pub fn dict_set(
&mut self,
dict: GenericValue,
key: GenericValue,
value: GenericValue,
) -> Result<(), PluginError> {
self.ffi_result((self.api.dict_set)(self.api.ctx, dict, key, value))
.map(|_| ())
}
pub fn dict_contains(
&mut self,
dict: GenericValue,
key: GenericValue,
) -> Result<bool, PluginError> {
let value = self.ffi_result((self.api.dict_contains)(self.api.ctx, dict, key))?;
Ok(self.as_bool(value).unwrap_or_default())
}
pub fn set_add(&mut self, set: GenericValue, item: GenericValue) -> Result<(), PluginError> {
self.ffi_result((self.api.set_add)(self.api.ctx, set, item))
.map(|_| ())
}
pub fn set_contains(
&mut self,
set: GenericValue,
item: GenericValue,
) -> Result<bool, PluginError> {
let value = self.ffi_result((self.api.set_contains)(self.api.ctx, set, item))?;
Ok(self.as_bool(value).unwrap_or_default())
}
pub fn truthy(&mut self, value: GenericValue) -> Result<bool, PluginError> {
let result = self.ffi_result((self.api.value_truthy)(self.api.ctx, value))?;
Ok(self.as_bool(result).unwrap_or_default())
}
pub fn equals(&mut self, a: GenericValue, b: GenericValue) -> Result<bool, PluginError> {
let result = self.ffi_result((self.api.value_equals)(self.api.ctx, a, b))?;
Ok(self.as_bool(result).unwrap_or_default())
}
pub fn hash(&mut self, value: GenericValue) -> Result<i64, PluginError> {
let result = self.ffi_result((self.api.value_hash)(self.api.ctx, value))?;
Ok(self.as_int(result).unwrap_or_default())
}
pub fn root(&self, value: GenericValue) {
(self.api.root)(self.api.ctx, value);
}
pub fn unroot(&self, n: usize) {
(self.api.unroot)(self.api.ctx, n);
}
#[must_use]
pub fn rooted(&self, value: GenericValue) -> Rooted<'a> {
(self.api.root)(self.api.ctx, value);
Rooted {
api: self.api,
value,
}
}
pub fn set_opaque(
&self,
receiver: GenericValue,
ptr: *mut core::ffi::c_void,
) -> Result<(), PluginError> {
self.ffi_result((self.api.instance_set_opaque)(self.api.ctx, receiver, ptr))
.map(|_| ())
}
#[must_use]
pub fn get_opaque(&self, receiver: GenericValue) -> *mut core::ffi::c_void {
(self.api.instance_get_opaque)(self.api.ctx, receiver)
}
#[allow(clippy::mut_from_ref)]
#[must_use]
pub unsafe fn opaque_ref<T>(&self, receiver: GenericValue) -> Option<&mut T> {
let ptr = self.get_opaque(receiver).cast::<T>();
unsafe { ptr.as_mut() }
}
const fn ffi_str(s: &str) -> FfiStr {
FfiStr {
ptr: s.as_ptr(),
len: s.len(),
}
}
fn ffi_result(&self, ret: FfiReturn) -> Result<GenericValue, PluginError> {
match FfiStatus::from_u32(ret.status) {
Some(FfiStatus::Ok) => Ok(ret.value),
Some(FfiStatus::Exception) => Err(PluginError::Exception(ret.value)),
Some(FfiStatus::Fatal) => Err(PluginError::Fatal),
None => Err(self.protocol_violation(&format!(
"host callback returned unknown status {}",
ret.status
))),
}
}
fn protocol_violation(&self, message: &str) -> PluginError {
let class = (self.api.builtin_get)(self.api.ctx, Self::ffi_str("Exception"));
if FfiStatus::from_u32(class.status) == Some(FfiStatus::Ok) {
let exception =
(self.api.exception_new)(self.api.ctx, class.value, Self::ffi_str(message));
if FfiStatus::from_u32(exception.status) == Some(FfiStatus::Ok) {
return PluginError::Exception(exception.value);
}
}
PluginError::Exception(self.make_nil())
}
}
pub struct Rooted<'a> {
api: &'a HostApi,
value: GenericValue,
}
impl Rooted<'_> {
#[must_use]
pub const fn get(&self) -> GenericValue {
self.value
}
}
impl Drop for Rooted<'_> {
fn drop(&mut self) {
(self.api.unroot)(self.api.ctx, 1);
}
}
pub type RustPluginFn = fn(&mut Host, &[GenericValue]) -> Result<GenericValue, PluginError>;
#[doc(hidden)]
pub unsafe fn __invoke_plugin_fn(
fun: RustPluginFn,
host: *const HostApi,
args: *const GenericValue,
nargs: usize,
) -> FfiReturn {
let api = unsafe { &*host };
let args: &[GenericValue] = if nargs == 0 {
&[]
} else {
unsafe { core::slice::from_raw_parts(args, nargs) }
};
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let mut host = Host::new(api);
fun(&mut host, args)
}));
finish_plugin_invoke(api, result)
}
pub type RustPluginMethodFn =
fn(&mut Host, GenericValue, &[GenericValue]) -> Result<GenericValue, PluginError>;
#[doc(hidden)]
pub unsafe fn __invoke_plugin_method_fn(
fun: RustPluginMethodFn,
host: *const HostApi,
receiver: GenericValue,
args: *const GenericValue,
nargs: usize,
) -> FfiReturn {
let api = unsafe { &*host };
let args: &[GenericValue] = if nargs == 0 {
&[]
} else {
unsafe { core::slice::from_raw_parts(args, nargs) }
};
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let mut host = Host::new(api);
fun(&mut host, receiver, args)
}));
finish_plugin_invoke(api, result)
}
pub type RustPluginValueFn = fn(&mut Host) -> Result<GenericValue, PluginError>;
#[doc(hidden)]
pub unsafe fn __invoke_plugin_value_fn(fun: RustPluginValueFn, host: *const HostApi) -> FfiReturn {
let api = unsafe { &*host };
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
let mut host = Host::new(api);
fun(&mut host)
}));
finish_plugin_invoke(api, result)
}
fn finish_plugin_invoke(
api: &HostApi,
result: std::thread::Result<Result<GenericValue, PluginError>>,
) -> FfiReturn {
let host = Host::new(api);
match result {
Ok(Ok(value)) => FfiReturn {
status: FfiStatus::Ok as u32,
value,
},
Ok(Err(error)) => error_return(&host, error),
Err(panic) => {
let message = panic
.downcast_ref::<&str>()
.map(ToString::to_string)
.or_else(|| panic.downcast_ref::<String>().cloned())
.unwrap_or_else(|| "plugin function panicked".to_owned());
error_return(&host, host.exception(&format!("panic: {message}")))
}
}
}
fn error_return(host: &Host, error: PluginError) -> FfiReturn {
match error {
PluginError::Exception(value) => FfiReturn {
status: FfiStatus::Exception as u32,
value,
},
PluginError::Fatal => FfiReturn {
status: FfiStatus::Fatal as u32,
value: host.make_nil(),
},
}
}