use libloading::{Library, Symbol};
use serde_derive::{Deserialize, Serialize};
use serde_with::{serde_as, DisplayFromStr};
use std::collections::HashMap;
use std::ffi::{c_char, CStr, CString};
use std::path::{Path, PathBuf};
use std::sync::{Arc, RwLock};
use crate::layout::{Layout, Struct};
use crate::{Context, Error};
pub type ExtensionInit = unsafe extern "C" fn() -> *const c_char;
pub const EXTENSION_INIT_SYMBOL: &[u8] = b"extension_init\0";
#[repr(transparent)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Outcome(pub(crate) *mut ());
#[repr(transparent)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct RawResource(pub(crate) *mut ());
#[repr(transparent)]
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct Dumped(pub(crate) *mut ());
#[derive(Debug, Serialize, Deserialize)]
#[serde(untagged)]
pub enum LoadOutcome {
Failed {
error: String,
},
Loaded(Box<ExtensionManifest>),
}
#[derive(Debug, Serialize, Deserialize)]
pub struct ExtensionManifest {
pub metadata: ExtensionMetadata,
pub outcome: OutcomeManifest,
pub dumped: DumpedManifest,
pub string: StringManifest,
pub resources: HashMap<String, ResourceManifest>,
}
#[serde_as]
#[derive(Debug, Serialize, Deserialize)]
pub struct ExtensionMetadata {
pub name: String,
#[serde_as(as = "DisplayFromStr")]
pub version: semver::Version,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct OutcomeManifest {
pub fn_get_err: String,
pub fn_get_ok: String,
pub fn_drop: String,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct DumpedManifest {
pub fn_get_ptr: String,
pub fn_get_len: String,
pub fn_drop: String,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct StringManifest {
pub fn_drop: String,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct ResourceManifest {
pub fn_from_bytes: String,
pub fn_dump: String,
pub fn_size: String,
pub fn_get_method_def: String,
pub fn_drop: String,
}
#[derive(Debug, Serialize, Deserialize)]
pub struct ExternalMethod {
pub fn_ptr: usize,
pub input_layout: Struct,
pub output_layout: Layout,
}
fn str_to_symbol_name(s: &str) -> Result<Vec<u8>, Error> {
Ok(CString::new(s)
.map_err(|err| err.to_string())?
.into_bytes_with_nul())
}
unsafe fn get_symbol<T: Copy>(library: &Library, name: &str) -> Result<T, Error> {
Ok(*library.get::<T>(&str_to_symbol_name(name)?)?)
}
#[derive(Debug)]
pub struct OutcomeSymbols {
pub fn_get_err: unsafe extern "C" fn(Outcome) -> *const c_char,
pub fn_get_ok: unsafe extern "C" fn(Outcome) -> *mut (),
pub fn_drop: unsafe extern "C" fn(Outcome),
}
impl OutcomeSymbols {
unsafe fn load(library: &Library, manifest: &OutcomeManifest) -> Result<OutcomeSymbols, Error> {
macro_rules! symbol {
($($sym:ident),*) => { Self {$(
$sym: get_symbol(library, &manifest.$sym).context(
concat!("getting symbol for ", stringify!($sym)
)
)?,
)*}}
}
Ok(symbol!(fn_get_err, fn_get_ok, fn_drop))
}
}
#[derive(Debug)]
pub struct DumpedSymbols {
pub fn_get_len: unsafe extern "C" fn(Dumped) -> usize,
pub fn_get_ptr: unsafe extern "C" fn(Dumped) -> *const u8,
pub fn_drop: unsafe extern "C" fn(Dumped),
}
impl DumpedSymbols {
unsafe fn load(library: &Library, manifest: &DumpedManifest) -> Result<DumpedSymbols, Error> {
macro_rules! symbol {
($($sym:ident),*) => { Self {$(
$sym: get_symbol(library, &manifest.$sym).context(
concat!("getting symbol for ", stringify!($sym)
)
)?,
)*}}
}
Ok(symbol!(fn_get_len, fn_get_ptr, fn_drop))
}
}
#[derive(Debug, Clone)]
pub struct StringSymbols {
pub fn_drop: unsafe extern "C" fn(*mut c_char),
}
impl StringSymbols {
unsafe fn load(library: &Library, manifest: &StringManifest) -> Result<StringSymbols, Error> {
macro_rules! symbol {
($($sym:ident),*) => { Self {$(
$sym: get_symbol(library, &manifest.$sym).context(
concat!("getting symbol for ", stringify!($sym)
)
)?,
)*}}
}
Ok(symbol!(fn_drop))
}
}
#[derive(Debug, Clone)]
pub struct ResourceSymbols {
pub fn_from_bytes: unsafe extern "C" fn(*const u8, usize) -> Outcome,
pub fn_dump: unsafe extern "C" fn(RawResource) -> Outcome,
pub fn_size: unsafe extern "C" fn(RawResource) -> usize,
pub fn_get_method_def: unsafe extern "C" fn(RawResource, *const c_char) -> *mut c_char,
pub fn_drop: unsafe extern "C" fn(RawResource),
}
impl ResourceSymbols {
unsafe fn load(
library: &Library,
manifest: &ResourceManifest,
) -> Result<ResourceSymbols, Error> {
macro_rules! symbol {
($($sym:ident),*) => { Self {$(
$sym: get_symbol(library, &manifest.$sym).context(
concat!("getting symbol for ", stringify!($sym)
)
)?,
)*}}
}
Ok(symbol!(
fn_from_bytes,
fn_dump,
fn_size,
fn_get_method_def,
fn_drop
))
}
}
type LoadedExtensionVersions = HashMap<semver::Version, Arc<Extension>>;
lazy_static::lazy_static! {
static ref EXTENSIONS: RwLock<HashMap<String, LoadedExtensionVersions>> =
RwLock::default();
}
#[derive(Debug)]
pub struct Extension {
_library: Library,
metadata: ExtensionMetadata,
outcome: OutcomeSymbols,
dumped: DumpedSymbols,
pub string: StringSymbols,
resources: HashMap<String, ResourceSymbols>,
}
impl Extension {
pub(crate) fn load(path: &Path) -> Result<Extension, Error> {
unsafe {
let library = Library::new(path)?;
let extension_init: Symbol<ExtensionInit> = library.get(EXTENSION_INIT_SYMBOL)?;
let outcome = extension_init();
if outcome.is_null() {
return Err(format!("library {path:?} failed to load").into());
}
let parsed: LoadOutcome = serde_json::from_slice(CStr::from_ptr(outcome).to_bytes())
.map_err(|err| err.to_string())?;
let manifest = match parsed {
LoadOutcome::Loaded(manifest) => manifest,
LoadOutcome::Failed { error } => return Err(error.into()),
};
let string = StringSymbols::load(&library, &manifest.string)
.with_context(|| format!("loading `string` symbols from {path:?}"))?;
let fn_drop = string.fn_drop;
scopeguard::defer! {
(fn_drop)(outcome as *mut i8);
}
let outcome = OutcomeSymbols::load(&library, &manifest.outcome)
.with_context(|| format!("loading `outcome` symbols from {path:?}"))?;
let dumped = DumpedSymbols::load(&library, &manifest.dumped)
.with_context(|| format!("loading `dumped` symbols from {path:?}"))?;
let resources = manifest
.resources
.iter()
.map(|(name, resource)| {
Ok((
name.clone(),
ResourceSymbols::load(&library, resource)
.with_context(|| format!("loading resource {name:?} from {path:?}"))?,
))
})
.collect::<Result<_, Error>>()?;
Ok(Extension {
_library: library,
metadata: manifest.metadata,
outcome,
dumped,
string,
resources,
})
}
}
pub(crate) unsafe fn outcome_to_result(&self, outcome: Outcome) -> Result<*mut (), Error> {
unsafe {
let maybe_err = (self.outcome.fn_get_err)(outcome);
let result = if !maybe_err.is_null() {
Err(CStr::from_ptr(maybe_err)
.to_string_lossy()
.to_string()
.into())
} else {
Ok((self.outcome.fn_get_ok)(outcome))
};
scopeguard::defer! {
(self.outcome.fn_drop)(outcome);
}
result
}
}
pub(crate) unsafe fn dumped_to_vec(&self, dumped: Dumped) -> Result<Vec<u8>, Error> {
unsafe {
scopeguard::defer! {
(self.dumped.fn_drop)(dumped)
}
let dump_ptr = (self.dumped.fn_get_ptr)(dumped);
if dump_ptr.is_null() {
return Err("dump location was null".to_string().into());
}
let dump_len = (self.dumped.fn_get_len)(dumped);
Ok(std::slice::from_raw_parts(dump_ptr, dump_len).to_vec())
}
}
pub(crate) fn get_resource(&self, name: &str) -> Option<ResourceSymbols> {
self.resources.get(name).cloned()
}
pub fn name(&self) -> &str {
&self.metadata.name
}
pub fn version(&self) -> &semver::Version {
&self.metadata.version
}
pub fn resources(&self) -> impl Iterator<Item = &str> {
self.resources.keys().map(|key| key.as_str())
}
}
#[cfg(target_os = "linux")]
const SO_EXTENSION: &str = "so";
#[cfg(target_os = "macos")]
const SO_EXTENSION: &str = "dylib";
#[cfg(target_os = "windows")]
const SO_EXTENSION: &str = "dll";
fn test_valid_name(name: &str) -> Result<(), Error> {
let is_valid = name.starts_with(|ch: char| ch.is_ascii_lowercase())
&& name
.chars()
.all(|ch| ch.is_ascii_lowercase() || ch.is_ascii_digit() || ch == '_');
if !is_valid {
return Err(format!("extension name {name:?} is invalid").into());
}
Ok(())
}
fn resolve_name(
name: &str,
version_req: &semver::VersionReq,
) -> Result<(semver::Version, PathBuf), Error> {
let full_path = std::env::var("JYAFN_PATH").unwrap_or_else(|_| {
home::home_dir()
.map(|home| home.join(".jyafn/extensions").to_string_lossy().to_string())
.unwrap_or_default()
});
let mut tried = vec![];
for alternative in full_path.split(',') {
let alternative = alternative.trim();
let mut candidates = vec![];
let glob = format!("{alternative}/{name}-*.{SO_EXTENSION}");
for path in glob::glob(&glob).map_err(|err| err.to_string())? {
let path = path.map_err(glob::GlobError::into_error)?;
if path.extension() != Some(SO_EXTENSION.as_ref()) {
continue;
}
let Some(filename_os) = path.file_stem() else {
continue;
};
let filename = filename_os.to_string_lossy();
let Some(version) = filename.split('-').last() else {
tried.push(format!("{path:?}"));
continue;
};
let Ok(semver) = version.parse::<semver::Version>() else {
tried.push(format!("{path:?}"));
continue;
};
if version_req.matches(&semver) {
candidates.push((semver, path));
} else {
tried.push(format!("{path:?}"));
}
}
if let Some(best_candidate) = candidates
.into_iter()
.max_by_key(|(semver, _)| semver.clone())
{
return Ok(best_candidate);
}
}
Err(format!(
"failed to resolve extension {name:?} (tried {})",
tried.join(", ")
)
.into())
}
pub fn try_get(name: &str, version_req: &semver::VersionReq) -> Result<Arc<Extension>, Error> {
test_valid_name(name)?;
let (version, path) = resolve_name(name, version_req)?;
let mut lock = EXTENSIONS.write().expect("poisoned");
let loaded_extensions = lock.entry(name.to_owned()).or_default();
if let Some(extension) = loaded_extensions.get(&version) {
return Ok(extension.clone());
}
let extension =
Arc::new(Extension::load(&path).with_context(|| format!("loading extension {name:?}"))?);
if extension.metadata.name != name {
return Err(format!(
"file {path:?} should provide {name:?} but provides {:?}",
extension.metadata.name
)
.into());
}
if extension.metadata.version != version {
return Err(format!(
"file {path:?} should provide version {version} but provides {}",
extension.metadata.version
)
.into());
}
loaded_extensions.insert(version, extension.clone());
Ok(extension)
}
pub fn get(name: &str, version_req: &semver::VersionReq) -> Arc<Extension> {
try_get(name, version_req).expect("extension not loaded")
}
pub fn list() -> HashMap<String, Vec<semver::Version>> {
EXTENSIONS
.read()
.expect("poisoned")
.iter()
.map(|(name, versions)| (name.clone(), versions.keys().cloned().collect::<Vec<_>>()))
.collect()
}
#[cfg(test)]
mod test {
use super::*;
#[test]
fn test_load_extension() {
get("dummy", &"*".parse().unwrap());
}
}