#![cfg(feature = "async")]
use std::ffi::{CStr, CString, c_void};
use std::os::raw::c_char;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, Condvar, Mutex, OnceLock};
use std::time::Duration;
use tensogram::{
DataObjectDescriptor, DecodeOptions, GlobalMetadata, MessageLayout, TensogramError,
TensogramFile,
};
pub(crate) enum AsyncOutcome {
Ok(TaskResult),
Timeout,
Cancelled,
Inner(TensogramError),
}
use crate::{TgmBytes, TgmError, TgmMessage, TgmMetadata, set_last_error, to_error_code};
struct RuntimeConfig {
workers: u32,
#[allow(dead_code)]
dispatcher_workers: u32,
#[allow(dead_code)]
multipart_part_size_bytes: u64,
}
impl Default for RuntimeConfig {
fn default() -> Self {
Self {
workers: std::cmp::min(num_cpus_or_default(), 8),
dispatcher_workers: std::cmp::min(num_cpus_or_default(), 4),
multipart_part_size_bytes: 8 * 1024 * 1024,
}
}
}
fn num_cpus_or_default() -> u32 {
std::thread::available_parallelism()
.map(|n| n.get() as u32)
.unwrap_or(4)
}
static RUNTIME_CONFIG: OnceLock<RuntimeConfig> = OnceLock::new();
static DISPATCHER: OnceLock<DispatcherPool> = OnceLock::new();
enum RuntimeState {
Uninit,
Built(tokio::runtime::Runtime),
BuildFailed(String),
ShutDown,
}
static RUNTIME: Mutex<RuntimeState> = Mutex::new(RuntimeState::Uninit);
static LIVE_TASKS: AtomicUsize = AtomicUsize::new(0);
struct LiveTaskGuard;
impl LiveTaskGuard {
fn new() -> Self {
LIVE_TASKS.fetch_add(1, Ordering::SeqCst);
LiveTaskGuard
}
}
impl Drop for LiveTaskGuard {
fn drop(&mut self) {
LIVE_TASKS.fetch_sub(1, Ordering::SeqCst);
}
}
struct CompletionGuard {
shared: Arc<TaskShared>,
}
impl Drop for CompletionGuard {
fn drop(&mut self) {
self.shared.complete(AsyncOutcome::Cancelled);
}
}
fn runtime_config() -> &'static RuntimeConfig {
RUNTIME_CONFIG.get_or_init(RuntimeConfig::default)
}
struct DispatcherJob {
cb: extern "C" fn(*mut c_void),
userdata: usize,
}
unsafe impl Send for DispatcherJob {}
struct DispatcherPool {
sender: std::sync::mpsc::SyncSender<DispatcherJob>,
_workers: Vec<std::thread::JoinHandle<()>>,
}
fn dispatcher_pool() -> &'static DispatcherPool {
DISPATCHER.get_or_init(|| {
let workers = runtime_config().dispatcher_workers.max(1) as usize;
let (tx, rx) = std::sync::mpsc::sync_channel::<DispatcherJob>(workers * 4);
let rx = std::sync::Arc::new(std::sync::Mutex::new(rx));
let mut handles = Vec::with_capacity(workers);
for i in 0..workers {
let rx = std::sync::Arc::clone(&rx);
let h = std::thread::Builder::new()
.name(format!("tensogram-dispatch-{i}"))
.spawn(move || {
loop {
let job = {
let guard = rx.lock().expect("dispatcher rx poisoned");
match guard.recv() {
Ok(j) => j,
Err(_) => break, }
};
(job.cb)(job.userdata as *mut c_void);
}
})
.expect("dispatcher worker spawn");
handles.push(h);
}
DispatcherPool {
sender: tx,
_workers: handles,
}
})
}
fn dispatch_to_pool(cb: extern "C" fn(*mut c_void), userdata: *mut c_void) {
let job = DispatcherJob {
cb,
userdata: userdata as usize,
};
if let Err(std::sync::mpsc::TrySendError::Full(job)) = dispatcher_pool().sender.try_send(job) {
(job.cb)(job.userdata as *mut c_void);
}
}
fn runtime() -> Result<tokio::runtime::Handle, String> {
let mut guard = RUNTIME.lock().expect("runtime mutex poisoned");
match &*guard {
RuntimeState::Built(rt) => Ok(rt.handle().clone()),
RuntimeState::BuildFailed(e) => Err(e.clone()),
RuntimeState::ShutDown => Err("runtime has been shut down".to_string()),
RuntimeState::Uninit => {
let cfg = runtime_config();
match tokio::runtime::Builder::new_multi_thread()
.worker_threads(cfg.workers as usize)
.enable_all()
.thread_name("tensogram-async")
.build()
{
Ok(rt) => {
let handle = rt.handle().clone();
*guard = RuntimeState::Built(rt);
Ok(handle)
}
Err(e) => {
let msg = format!("failed to build tokio runtime: {e}");
*guard = RuntimeState::BuildFailed(msg.clone());
Err(msg)
}
}
}
}
}
#[allow(dead_code)] pub(crate) enum TaskResult {
File(Box<crate::TgmFile>),
AsyncFile(Box<TgmAsyncFile>),
AsyncStreamingEncoder(Box<crate::async_streaming::TgmAsyncStreamingEncoder>),
Message(Box<TgmMessage>),
Metadata(Box<TgmMetadata>),
Bytes(Vec<u8>),
MultiBytes(Vec<Vec<u8>>),
Size(u64),
Layouts(Vec<MessageLayout>),
Void,
}
#[derive(Debug, PartialEq)]
enum TaskState {
Pending,
Ready,
Consumed,
}
struct TaskInner {
state: TaskState,
result: Option<AsyncOutcome>,
completion_cb: Option<extern "C" fn(*mut c_void)>,
completion_registered: bool,
completion_userdata: *mut c_void,
}
unsafe impl Send for TaskInner {}
pub struct TgmAsyncTask {
inner: Arc<TaskShared>,
}
struct TaskShared {
state: Mutex<TaskInner>,
ready: Condvar,
cancel: tokio_util::sync::CancellationToken,
external: Option<tokio_util::sync::CancellationToken>,
}
impl TaskShared {
fn new(external: Option<&TgmCancellationToken>) -> Arc<Self> {
Arc::new(Self {
state: Mutex::new(TaskInner {
state: TaskState::Pending,
result: None,
completion_cb: None,
completion_registered: false,
completion_userdata: std::ptr::null_mut(),
}),
ready: Condvar::new(),
cancel: tokio_util::sync::CancellationToken::new(),
external: external.map(|t| t.inner.token.clone()),
})
}
fn complete(&self, result: AsyncOutcome) {
let cb_to_fire = {
let mut inner = self.state.lock().expect("task mutex poisoned");
if inner.state != TaskState::Pending {
return;
}
inner.result = Some(result);
inner.state = TaskState::Ready;
inner.completion_cb.take().map(|cb| {
(
cb,
std::mem::replace(&mut inner.completion_userdata, std::ptr::null_mut()),
)
})
};
self.ready.notify_all();
if let Some((cb, userdata)) = cb_to_fire {
dispatch_to_pool(cb, userdata);
}
}
}
pub(crate) fn spawn_task<F>(
fut: F,
cancel: Option<&TgmCancellationToken>,
timeout_ms: u64,
) -> Result<*mut TgmAsyncTask, String>
where
F: std::future::Future<Output = Result<TaskResult, TensogramError>> + Send + 'static,
{
let rt = runtime()?;
let shared = TaskShared::new(cancel);
let shared_clone = shared.clone();
let internal_token = shared.cancel.clone();
let external_token = shared.external.clone();
let live_guard = LiveTaskGuard::new();
let completion_guard = CompletionGuard {
shared: shared.clone(),
};
rt.spawn(async move {
let _live = live_guard;
let _completion = completion_guard;
let outcome: AsyncOutcome = if timeout_ms > 0 {
let deadline = Duration::from_millis(timeout_ms);
tokio::select! {
_ = internal_token.cancelled() => AsyncOutcome::Cancelled,
_ = async { match external_token { Some(t) => t.cancelled().await, None => std::future::pending().await } } => AsyncOutcome::Cancelled,
res = tokio::time::timeout(deadline, fut) => match res {
Ok(Ok(v)) => AsyncOutcome::Ok(v),
Ok(Err(e)) => AsyncOutcome::Inner(e),
Err(_elapsed) => AsyncOutcome::Timeout,
},
}
} else {
tokio::select! {
_ = internal_token.cancelled() => AsyncOutcome::Cancelled,
_ = async { match external_token { Some(t) => t.cancelled().await, None => std::future::pending().await } } => AsyncOutcome::Cancelled,
res = fut => match res {
Ok(v) => AsyncOutcome::Ok(v),
Err(e) => AsyncOutcome::Inner(e),
},
}
};
shared_clone.complete(outcome);
});
Ok(Box::into_raw(Box::new(TgmAsyncTask { inner: shared })))
}
pub(crate) struct TgmCancellationTokenInner {
token: tokio_util::sync::CancellationToken,
}
pub struct TgmCancellationToken {
inner: Arc<TgmCancellationTokenInner>,
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_cancellation_token_create() -> *mut TgmCancellationToken {
Box::into_raw(Box::new(TgmCancellationToken {
inner: Arc::new(TgmCancellationTokenInner {
token: tokio_util::sync::CancellationToken::new(),
}),
}))
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_cancellation_token_cancel(tok: *mut TgmCancellationToken) {
if tok.is_null() {
return;
}
let t = unsafe { &*tok };
t.inner.token.cancel();
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_cancellation_token_is_cancelled(tok: *const TgmCancellationToken) -> bool {
if tok.is_null() {
return false;
}
let t = unsafe { &*tok };
t.inner.token.is_cancelled()
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_cancellation_token_free(tok: *mut TgmCancellationToken) {
if !tok.is_null() {
unsafe { drop(Box::from_raw(tok)) };
}
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_task_set_completion(
task: *mut TgmAsyncTask,
cb: extern "C" fn(*mut c_void),
userdata: *mut c_void,
) -> TgmError {
if task.is_null() {
set_last_error("null task");
return TgmError::InvalidArg;
}
let t = unsafe { &*task };
let fire_inline = {
let mut inner = t.inner.state.lock().expect("task mutex poisoned");
if inner.completion_registered {
set_last_error("completion callback already registered");
return TgmError::InvalidArg;
}
inner.completion_registered = true;
match inner.state {
TaskState::Pending => {
inner.completion_cb = Some(cb);
inner.completion_userdata = userdata;
false
}
TaskState::Ready => true,
TaskState::Consumed => {
set_last_error("task result already consumed");
return TgmError::InvalidArg;
}
}
};
if fire_inline {
dispatch_to_pool(cb, userdata);
}
TgmError::Ok
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_task_is_ready(task: *const TgmAsyncTask) -> bool {
if task.is_null() {
return false;
}
let t = unsafe { &*task };
let inner = t.inner.state.lock().expect("task mutex poisoned");
inner.state != TaskState::Pending
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_task_cancel(task: *mut TgmAsyncTask) {
if task.is_null() {
return;
}
let t = unsafe { &*task };
t.inner.cancel.cancel();
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_task_free(task: *mut TgmAsyncTask) {
if !task.is_null() {
unsafe { drop(Box::from_raw(task)) };
}
}
pub(crate) fn join_internal(task: *mut TgmAsyncTask) -> Result<TaskResult, TgmError> {
if task.is_null() {
set_last_error("null task");
return Err(TgmError::InvalidArg);
}
let t = unsafe { &*task };
let mut inner = t.inner.state.lock().expect("task mutex poisoned");
while inner.state == TaskState::Pending {
inner = t.inner.ready.wait(inner).expect("task condvar poisoned");
}
if inner.state == TaskState::Consumed {
set_last_error("task result already consumed");
return Err(TgmError::InvalidArg);
}
let res = inner.result.take().expect("ready task missing result");
inner.state = TaskState::Consumed;
drop(inner);
match res {
AsyncOutcome::Ok(v) => Ok(v),
AsyncOutcome::Timeout => {
set_last_error("async task timed out");
Err(TgmError::Timeout)
}
AsyncOutcome::Cancelled => {
set_last_error("async task cancelled");
Err(TgmError::Cancelled)
}
AsyncOutcome::Inner(e) => {
set_last_error(&e.to_string());
Err(to_error_code(&e))
}
}
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_task_join_void(task: *mut TgmAsyncTask) -> TgmError {
match join_internal(task) {
Ok(TaskResult::Void) => TgmError::Ok,
Ok(_) => {
set_last_error("task result type mismatch (expected void)");
TgmError::InvalidArg
}
Err(code) => code,
}
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_task_join_size(task: *mut TgmAsyncTask, out: *mut u64) -> TgmError {
if out.is_null() {
set_last_error("null out pointer");
return TgmError::InvalidArg;
}
match join_internal(task) {
Ok(TaskResult::Size(n)) => {
unsafe { *out = n };
TgmError::Ok
}
Ok(_) => {
set_last_error("task result type mismatch (expected size)");
TgmError::InvalidArg
}
Err(code) => code,
}
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_task_join_bytes(
task: *mut TgmAsyncTask,
out: *mut TgmBytes,
) -> TgmError {
if out.is_null() {
set_last_error("null out pointer");
return TgmError::InvalidArg;
}
match join_internal(task) {
Ok(TaskResult::Bytes(mut v)) => {
v.shrink_to_fit();
let len = v.len();
let ptr = v.as_mut_ptr();
std::mem::forget(v);
unsafe {
(*out).data = ptr;
(*out).len = len;
}
TgmError::Ok
}
Ok(_) => {
set_last_error("task result type mismatch (expected bytes)");
TgmError::InvalidArg
}
Err(code) => code,
}
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_task_join_message(
task: *mut TgmAsyncTask,
out: *mut *mut TgmMessage,
) -> TgmError {
if out.is_null() {
set_last_error("null out pointer");
return TgmError::InvalidArg;
}
match join_internal(task) {
Ok(TaskResult::Message(m)) => {
unsafe { *out = Box::into_raw(m) };
TgmError::Ok
}
Ok(_) => {
set_last_error("task result type mismatch (expected message)");
TgmError::InvalidArg
}
Err(code) => code,
}
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_task_join_metadata(
task: *mut TgmAsyncTask,
out: *mut *mut TgmMetadata,
) -> TgmError {
if out.is_null() {
set_last_error("null out pointer");
return TgmError::InvalidArg;
}
match join_internal(task) {
Ok(TaskResult::Metadata(m)) => {
unsafe { *out = Box::into_raw(m) };
TgmError::Ok
}
Ok(_) => {
set_last_error("task result type mismatch (expected metadata)");
TgmError::InvalidArg
}
Err(code) => code,
}
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_task_join_async_file(
task: *mut TgmAsyncTask,
out: *mut *mut TgmAsyncFile,
) -> TgmError {
if out.is_null() {
set_last_error("null out pointer");
return TgmError::InvalidArg;
}
match join_internal(task) {
Ok(TaskResult::AsyncFile(f)) => {
unsafe { *out = Box::into_raw(f) };
TgmError::Ok
}
Ok(_) => {
set_last_error("task result type mismatch (expected async_file)");
TgmError::InvalidArg
}
Err(code) => code,
}
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_task_join_multi_bytes(
task: *mut TgmAsyncTask,
out_array: *mut *mut TgmBytes,
out_count: *mut usize,
) -> TgmError {
if out_array.is_null() || out_count.is_null() {
set_last_error("null out pointer");
return TgmError::InvalidArg;
}
match join_internal(task) {
Ok(TaskResult::MultiBytes(parts)) => {
let count = parts.len();
let mut entries: Vec<TgmBytes> = parts
.into_iter()
.map(|mut v| {
v.shrink_to_fit();
let len = v.len();
let ptr = v.as_mut_ptr();
std::mem::forget(v);
TgmBytes { data: ptr, len }
})
.collect();
entries.shrink_to_fit();
let ptr = entries.as_mut_ptr();
std::mem::forget(entries);
unsafe {
*out_array = ptr;
*out_count = count;
}
TgmError::Ok
}
Ok(_) => {
set_last_error("task result type mismatch (expected multi_bytes)");
TgmError::InvalidArg
}
Err(code) => code,
}
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_multi_bytes_free(array: *mut TgmBytes, count: usize) {
if array.is_null() {
return;
}
unsafe {
let v = Vec::from_raw_parts(array, count, count);
for entry in v {
crate::tgm_bytes_free(entry);
}
}
}
pub struct TgmAsyncFile {
file: Arc<TensogramFile>,
path_string: CString,
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_file_path(file: *const TgmAsyncFile) -> *const c_char {
if file.is_null() {
return std::ptr::null();
}
unsafe { (*file).path_string.as_ptr() }
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_file_close(file: *mut TgmAsyncFile) {
if !file.is_null() {
unsafe { drop(Box::from_raw(file)) };
}
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_file_open(
path: *const c_char,
cancel: *mut TgmCancellationToken,
timeout_ms: u64,
out_task: *mut *mut TgmAsyncTask,
) -> TgmError {
if path.is_null() || out_task.is_null() {
set_last_error("null argument");
return TgmError::InvalidArg;
}
let path_str = match unsafe { CStr::from_ptr(path) }.to_str() {
Ok(s) => s.to_string(),
Err(e) => {
set_last_error(&format!("invalid UTF-8 in path: {e}"));
return TgmError::InvalidArg;
}
};
let cancel_ref = if cancel.is_null() {
None
} else {
Some(unsafe { &*cancel })
};
let path_for_task = path_str.clone();
let fut = async move {
let f = TensogramFile::open(&path_for_task)?;
let path_string = CString::new(path_for_task.as_str()).unwrap_or_default();
Ok(TaskResult::AsyncFile(Box::new(TgmAsyncFile {
file: Arc::new(f),
path_string,
})))
};
spawn_or_set_error(fut, cancel_ref, timeout_ms, out_task)
}
#[unsafe(no_mangle)]
#[allow(clippy::too_many_arguments)]
pub extern "C" fn tgm_async_file_open_remote(
url: *const c_char,
storage_keys: *const *const c_char,
storage_values: *const *const c_char,
nopts: usize,
bidirectional: bool,
cancel: *mut TgmCancellationToken,
timeout_ms: u64,
out_task: *mut *mut TgmAsyncTask,
) -> TgmError {
#[cfg(not(feature = "async-remote"))]
{
let _ = (
url,
storage_keys,
storage_values,
nopts,
bidirectional,
cancel,
timeout_ms,
out_task,
);
set_last_error(
"tgm_async_file_open_remote: this build of tensogram-ffi was compiled \
without the `async-remote` Cargo feature; rebuild with --features=async-remote \
to enable S3/GCS/Azure/HTTP support",
);
TgmError::Remote
}
#[cfg(feature = "async-remote")]
{
if url.is_null() || out_task.is_null() {
set_last_error("null argument");
return TgmError::InvalidArg;
}
let url_str = match unsafe { CStr::from_ptr(url) }.to_str() {
Ok(s) => s.to_string(),
Err(e) => {
set_last_error(&format!("invalid UTF-8 in url: {e}"));
return TgmError::InvalidArg;
}
};
let mut storage = std::collections::BTreeMap::new();
if nopts > 0 {
if storage_keys.is_null() || storage_values.is_null() {
set_last_error("null storage opts pointer");
return TgmError::InvalidArg;
}
for i in 0..nopts {
let kp = unsafe { *storage_keys.add(i) };
let vp = unsafe { *storage_values.add(i) };
if kp.is_null() || vp.is_null() {
set_last_error("null storage opt entry");
return TgmError::InvalidArg;
}
let k = match unsafe { CStr::from_ptr(kp) }.to_str() {
Ok(s) => s.to_string(),
Err(e) => {
set_last_error(&format!("invalid UTF-8 in storage key: {e}"));
return TgmError::InvalidArg;
}
};
let v = match unsafe { CStr::from_ptr(vp) }.to_str() {
Ok(s) => s.to_string(),
Err(e) => {
set_last_error(&format!("invalid UTF-8 in storage value: {e}"));
return TgmError::InvalidArg;
}
};
storage.insert(k, v);
}
}
let cancel_ref = if cancel.is_null() {
None
} else {
Some(unsafe { &*cancel })
};
let scan_opts = tensogram::RemoteScanOptions { bidirectional };
let url_for_task = url_str.clone();
let fut = async move {
let f =
TensogramFile::open_remote_async(&url_for_task, &storage, Some(scan_opts)).await?;
let path_string = CString::new(url_for_task.as_str()).unwrap_or_default();
Ok(TaskResult::AsyncFile(Box::new(TgmAsyncFile {
file: Arc::new(f),
path_string,
})))
};
spawn_or_set_error(fut, cancel_ref, timeout_ms, out_task)
}
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_file_message_count(
file: *mut TgmAsyncFile,
cancel: *mut TgmCancellationToken,
timeout_ms: u64,
out_task: *mut *mut TgmAsyncTask,
) -> TgmError {
if file.is_null() || out_task.is_null() {
set_last_error("null argument");
return TgmError::InvalidArg;
}
let f = unsafe { &*file }.file.clone();
let cancel_ref = if cancel.is_null() {
None
} else {
Some(unsafe { &*cancel })
};
let fut = async move {
let n = f.message_count_async().await?;
Ok(TaskResult::Size(n as u64))
};
spawn_or_set_error(fut, cancel_ref, timeout_ms, out_task)
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_file_read_message(
file: *mut TgmAsyncFile,
index: usize,
cancel: *mut TgmCancellationToken,
timeout_ms: u64,
out_task: *mut *mut TgmAsyncTask,
) -> TgmError {
if file.is_null() || out_task.is_null() {
set_last_error("null argument");
return TgmError::InvalidArg;
}
let f = unsafe { &*file }.file.clone();
let cancel_ref = if cancel.is_null() {
None
} else {
Some(unsafe { &*cancel })
};
let fut = async move {
let bytes = f.read_message_async(index).await?;
Ok(TaskResult::Bytes(bytes))
};
spawn_or_set_error(fut, cancel_ref, timeout_ms, out_task)
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_file_decode_message(
file: *mut TgmAsyncFile,
index: usize,
native_byte_order: bool,
threads: u32,
restore_non_finite: bool,
verify_hash: bool,
cancel: *mut TgmCancellationToken,
timeout_ms: u64,
out_task: *mut *mut TgmAsyncTask,
) -> TgmError {
if file.is_null() || out_task.is_null() {
set_last_error("null argument");
return TgmError::InvalidArg;
}
let f = unsafe { &*file }.file.clone();
let cancel_ref = if cancel.is_null() {
None
} else {
Some(unsafe { &*cancel })
};
let opts = DecodeOptions {
native_byte_order,
threads,
restore_non_finite,
verify_hash,
..Default::default()
};
let fut = async move {
let (gm, objs) = f.decode_message_async(index, &opts).await?;
Ok(TaskResult::Message(Box::new(build_tgm_message(gm, objs))))
};
spawn_or_set_error(fut, cancel_ref, timeout_ms, out_task)
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_file_decode_metadata(
file: *mut TgmAsyncFile,
index: usize,
cancel: *mut TgmCancellationToken,
timeout_ms: u64,
out_task: *mut *mut TgmAsyncTask,
) -> TgmError {
if file.is_null() || out_task.is_null() {
set_last_error("null argument");
return TgmError::InvalidArg;
}
let f = unsafe { &*file }.file.clone();
let cancel_ref = if cancel.is_null() {
None
} else {
Some(unsafe { &*cancel })
};
let fut = async move {
let gm = f.decode_metadata_async(index).await?;
Ok(TaskResult::Metadata(Box::new(TgmMetadata {
global_metadata: gm,
cache: std::cell::RefCell::new(std::collections::BTreeMap::new()),
})))
};
spawn_or_set_error(fut, cancel_ref, timeout_ms, out_task)
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_file_decode_object(
file: *mut TgmAsyncFile,
msg_index: usize,
obj_index: usize,
native_byte_order: bool,
threads: u32,
restore_non_finite: bool,
verify_hash: bool,
cancel: *mut TgmCancellationToken,
timeout_ms: u64,
out_task: *mut *mut TgmAsyncTask,
) -> TgmError {
if file.is_null() || out_task.is_null() {
set_last_error("null argument");
return TgmError::InvalidArg;
}
let f = unsafe { &*file }.file.clone();
let cancel_ref = if cancel.is_null() {
None
} else {
Some(unsafe { &*cancel })
};
let opts = DecodeOptions {
native_byte_order,
threads,
restore_non_finite,
verify_hash,
..Default::default()
};
let fut = async move {
let (gm, desc, payload) = f.decode_object_async(msg_index, obj_index, &opts).await?;
Ok(TaskResult::Message(Box::new(build_tgm_message(
gm,
vec![(desc, payload)],
))))
};
spawn_or_set_error(fut, cancel_ref, timeout_ms, out_task)
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_async_file_decode_range(
file: *mut TgmAsyncFile,
msg_index: usize,
obj_index: usize,
offsets: *const u64,
counts: *const u64,
n_ranges: usize,
native_byte_order: bool,
threads: u32,
cancel: *mut TgmCancellationToken,
timeout_ms: u64,
out_task: *mut *mut TgmAsyncTask,
) -> TgmError {
if file.is_null() || out_task.is_null() || offsets.is_null() || counts.is_null() {
set_last_error("null argument");
return TgmError::InvalidArg;
}
let f = unsafe { &*file }.file.clone();
let cancel_ref = if cancel.is_null() {
None
} else {
Some(unsafe { &*cancel })
};
let mut ranges: Vec<(u64, u64)> = Vec::with_capacity(n_ranges);
for i in 0..n_ranges {
let o = unsafe { *offsets.add(i) };
let c = unsafe { *counts.add(i) };
ranges.push((o, c));
}
let opts = DecodeOptions {
native_byte_order,
threads,
..Default::default()
};
let fut = async move {
let (_desc, parts) = f
.decode_range_async(msg_index, obj_index, &ranges, &opts)
.await?;
Ok(TaskResult::MultiBytes(parts))
};
spawn_or_set_error(fut, cancel_ref, timeout_ms, out_task)
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_runtime_configure(
workers: u32,
dispatcher_workers: u32,
multipart_part_size_bytes: u64,
) -> TgmError {
let cfg = RuntimeConfig {
workers: if workers == 0 {
std::cmp::min(num_cpus_or_default(), 8)
} else {
workers
},
dispatcher_workers: if dispatcher_workers == 0 {
std::cmp::min(num_cpus_or_default(), 4)
} else {
dispatcher_workers
},
multipart_part_size_bytes: if multipart_part_size_bytes == 0 {
8 * 1024 * 1024
} else {
multipart_part_size_bytes
},
};
if RUNTIME_CONFIG.set(cfg).is_err() {
set_last_error("runtime already configured or built");
return TgmError::InvalidArg;
}
TgmError::Ok
}
#[unsafe(no_mangle)]
pub extern "C" fn tgm_runtime_shutdown_blocking(timeout_ms: u64) -> u64 {
let taken = {
let mut guard = RUNTIME.lock().expect("runtime mutex poisoned");
match std::mem::replace(&mut *guard, RuntimeState::ShutDown) {
RuntimeState::Built(rt) => Some(rt),
_ => None,
}
};
let Some(rt) = taken else {
return 0;
};
let unfinished = drain_until(
|| LIVE_TASKS.load(Ordering::SeqCst),
Duration::from_millis(timeout_ms),
Duration::from_millis(5),
std::thread::sleep,
);
rt.shutdown_timeout(Duration::ZERO);
unfinished
}
fn drain_until<C, S>(live_count: C, timeout: Duration, poll_interval: Duration, mut sleep: S) -> u64
where
C: Fn() -> usize,
S: FnMut(Duration),
{
let start = std::time::Instant::now();
loop {
if live_count() == 0 {
break;
}
let elapsed = start.elapsed();
if elapsed >= timeout {
break;
}
sleep(poll_interval.min(timeout - elapsed));
}
live_count() as u64
}
pub(crate) fn spawn_or_set_error<F>(
fut: F,
cancel: Option<&TgmCancellationToken>,
timeout_ms: u64,
out_task: *mut *mut TgmAsyncTask,
) -> TgmError
where
F: std::future::Future<Output = Result<TaskResult, TensogramError>> + Send + 'static,
{
match spawn_task(fut, cancel, timeout_ms) {
Ok(t) => {
unsafe { *out_task = t };
TgmError::Ok
}
Err(s) => {
set_last_error(&s);
TgmError::Io
}
}
}
fn build_tgm_message(
gm: GlobalMetadata,
objects: Vec<(DataObjectDescriptor, Vec<u8>)>,
) -> TgmMessage {
let mut dtype_strings = Vec::with_capacity(objects.len());
let mut type_strings = Vec::with_capacity(objects.len());
let mut byte_order_strings = Vec::with_capacity(objects.len());
let mut filter_strings = Vec::with_capacity(objects.len());
let mut compression_strings = Vec::with_capacity(objects.len());
let mut encoding_strings = Vec::with_capacity(objects.len());
let hash_type_strings = vec![None; objects.len()];
let hash_value_strings = vec![None; objects.len()];
for (desc, _) in &objects {
dtype_strings.push(CString::new(dtype_name(desc.dtype)).unwrap_or_default());
type_strings.push(CString::new(desc.obj_type.as_str()).unwrap_or_default());
byte_order_strings.push(
CString::new(if desc.byte_order == tensogram::ByteOrder::Big {
"big"
} else {
"little"
})
.unwrap_or_default(),
);
filter_strings.push(CString::new(desc.filter.as_str()).unwrap_or_default());
compression_strings.push(CString::new(desc.compression.as_str()).unwrap_or_default());
encoding_strings.push(CString::new(desc.encoding.as_str()).unwrap_or_default());
}
TgmMessage {
global_metadata: gm,
objects,
dtype_strings,
type_strings,
byte_order_strings,
filter_strings,
compression_strings,
encoding_strings,
hash_type_strings,
hash_value_strings,
}
}
fn dtype_name(d: tensogram::dtype::Dtype) -> &'static str {
use tensogram::dtype::Dtype;
match d {
Dtype::Float16 => "float16",
Dtype::Bfloat16 => "bfloat16",
Dtype::Float32 => "float32",
Dtype::Float64 => "float64",
Dtype::Complex64 => "complex64",
Dtype::Complex128 => "complex128",
Dtype::Int8 => "int8",
Dtype::Int16 => "int16",
Dtype::Int32 => "int32",
Dtype::Int64 => "int64",
Dtype::Uint8 => "uint8",
Dtype::Uint16 => "uint16",
Dtype::Uint32 => "uint32",
Dtype::Uint64 => "uint64",
Dtype::Bitmask => "bitmask",
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::cell::RefCell;
#[test]
fn drain_until_zero_count_breaks_immediately_without_sleeping() {
let sleeps = RefCell::new(0usize);
let residual = drain_until(
|| 0,
Duration::from_secs(3600),
Duration::from_millis(5),
|_| {
*sleeps.borrow_mut() += 1;
panic!("zero count must break before any sleep");
},
);
assert_eq!(residual, 0, "zero count must report zero unfinished");
assert_eq!(
*sleeps.borrow(),
0,
"a zero count must break before any sleep"
);
}
#[test]
fn drain_until_terminates_when_count_reaches_zero() {
let calls = RefCell::new(0usize);
let sleeps = RefCell::new(0usize);
let residual = drain_until(
|| {
let mut c = calls.borrow_mut();
*c += 1;
if *c >= 3 { 0 } else { 1 }
},
Duration::from_secs(3600),
Duration::from_millis(1),
|_| *sleeps.borrow_mut() += 1,
);
assert_eq!(residual, 0);
assert_eq!(*sleeps.borrow(), 2, "should sleep once per busy poll");
}
#[test]
fn drain_until_breaks_on_timeout_and_reports_residual() {
let residual = drain_until(
|| 4,
Duration::ZERO, Duration::from_millis(5),
|_| {},
);
assert_eq!(
residual, 4,
"a never-draining count must report its residual"
);
}
#[test]
fn drain_until_clamps_sleep_to_remaining_budget() {
let recorded: RefCell<Vec<Duration>> = RefCell::new(Vec::new());
let polls = RefCell::new(0usize);
let _ = drain_until(
|| {
let mut p = polls.borrow_mut();
*p += 1;
if *p >= 2 { 0 } else { 1 }
},
Duration::from_millis(20),
Duration::from_secs(10), |d| recorded.borrow_mut().push(d),
);
let recorded = recorded.borrow();
assert_eq!(recorded.len(), 1, "exactly one busy poll → one sleep");
assert!(
recorded[0] <= Duration::from_millis(20),
"sleep must be clamped to the remaining budget, got {:?}",
recorded[0]
);
}
#[test]
fn completion_guard_drop_cancels_pending_task() {
let shared = TaskShared::new(None);
{
let _guard = CompletionGuard {
shared: shared.clone(),
};
assert_eq!(
shared.state.lock().unwrap().state,
TaskState::Pending,
"task should still be Pending before the guard drops"
);
}
let inner = shared.state.lock().unwrap();
assert_eq!(
inner.state,
TaskState::Ready,
"dropping the guard must move the task out of Pending"
);
assert!(
matches!(inner.result, Some(AsyncOutcome::Cancelled)),
"a guard-cancelled task must carry the Cancelled outcome"
);
}
#[test]
fn completion_guard_drop_is_noop_when_already_completed() {
let shared = TaskShared::new(None);
shared.complete(AsyncOutcome::Ok(TaskResult::Void));
{
let _guard = CompletionGuard {
shared: shared.clone(),
};
}
let inner = shared.state.lock().unwrap();
assert!(
matches!(inner.result, Some(AsyncOutcome::Ok(TaskResult::Void))),
"the guard must not overwrite an already-set outcome"
);
}
#[test]
fn live_task_guard_increments_then_decrements() {
let before = LIVE_TASKS.load(Ordering::SeqCst);
{
let _g = LiveTaskGuard::new();
assert_eq!(
LIVE_TASKS.load(Ordering::SeqCst),
before + 1,
"new() must increment LIVE_TASKS"
);
}
assert_eq!(
LIVE_TASKS.load(Ordering::SeqCst),
before,
"drop must decrement LIVE_TASKS back to the starting value"
);
}
#[test]
fn shutdown_blocking_drains_and_reports_unfinished() {
let task = spawn_task(
async {
tokio::time::sleep(Duration::from_secs(3600)).await;
Ok(TaskResult::Void)
},
None,
0,
)
.expect("spawn before shutdown should succeed");
let unfinished = tgm_runtime_shutdown_blocking(50);
assert!(
unfinished >= 1,
"a task sleeping for an hour must be reported unfinished after a 50 ms drain, got {unfinished}"
);
let err = spawn_task(async { Ok(TaskResult::Void) }, None, 0)
.expect_err("spawn after shutdown must fail");
assert!(
err.contains("shut down"),
"post-shutdown spawn error should mention shutdown, got: {err}"
);
assert_eq!(
tgm_runtime_shutdown_blocking(10),
0,
"second shutdown must be a no-op returning 0"
);
let mut ready = false;
for _ in 0..200 {
if tgm_async_task_is_ready(task) {
ready = true;
break;
}
std::thread::sleep(Duration::from_millis(5));
}
assert!(
ready,
"the CompletionGuard backstop must move a shutdown-stranded task out of Pending"
);
let join_tag: Result<(), TgmError> = join_internal(task).map(|_| ());
assert!(
matches!(join_tag, Err(TgmError::Cancelled)),
"a task stranded by shutdown must join as Cancelled; got {join_tag:?}"
);
tgm_async_task_free(task);
}
}