use std::os::raw::c_int;
use std::panic::{catch_unwind, AssertUnwindSafe};
use std::ptr;
use std::sync::Arc;
use tokio::sync::{
mpsc::{UnboundedReceiver, UnboundedSender},
oneshot,
};
use super::api::{Api, Kvps};
use super::ffi::*;
use super::items::{
item_from_native, item_to_native, make_bytes_item, make_openai_json_item,
read_speech_result_text, read_text_item,
};
use super::manager::NativeManager;
use super::native::NativeModel;
use crate::error::{FoundryLocalError, Result};
use crate::item::Item;
pub(crate) type StreamTransform = Box<dyn Fn(String) -> Option<String> + Send>;
pub(crate) struct NativeRequest {
api: Arc<Api>,
ptr: *mut flRequest,
}
unsafe impl Send for NativeRequest {}
unsafe impl Sync for NativeRequest {}
impl NativeRequest {
pub(crate) fn new(api: Arc<Api>) -> Result<Self> {
let mut ptr: *mut flRequest = ptr::null_mut();
api.check(unsafe { (api.inference_api().Request_Create)(&mut ptr) })?;
Ok(Self { api, ptr })
}
pub(crate) fn add_item(&self, item: *mut flItem, take_ownership: bool) -> Result<()> {
let status =
unsafe { (self.api.inference_api().Request_AddItem)(self.ptr, item, take_ownership) };
self.api.check(status)
}
pub(crate) fn add_item_value(&self, item: &Item) -> Result<()> {
let native = item_to_native(&self.api, item)?;
if let Err(e) = self.add_item(native, true) {
unsafe { (self.api.item_api().Item_Release)(native) };
return Err(e);
}
Ok(())
}
pub(crate) fn add_input_queue(&self, queue: &NativeItemQueue) -> Result<()> {
self.add_item(queue.as_item_ptr(), false)
}
pub(crate) fn set_options(&self, options: *const flKeyValuePairs) -> Result<()> {
let status = unsafe { (self.api.inference_api().Request_SetOptions)(self.ptr, options) };
self.api.check(status)
}
pub(crate) fn cancel(&self) {
let status = unsafe { (self.api.inference_api().Request_Cancel)(self.ptr) };
let _ = self.api.check(status);
}
}
impl Drop for NativeRequest {
fn drop(&mut self) {
if !self.ptr.is_null() {
unsafe { (self.api.inference_api().Request_Release)(self.ptr) };
self.ptr = ptr::null_mut();
}
}
}
pub(crate) struct NativeResponse {
api: Arc<Api>,
ptr: *mut flResponse,
}
impl NativeResponse {
pub(crate) fn item_count(&self) -> usize {
unsafe { (self.api.inference_api().Response_GetItemCount)(self.ptr) }
}
pub(crate) fn item_text(&self, idx: usize) -> Result<Option<String>> {
let mut item: *const flItem = ptr::null();
let status =
unsafe { (self.api.inference_api().Response_GetItem)(self.ptr, idx, &mut item) };
self.api.check(status)?;
unsafe { read_text_item(&self.api, item) }
}
pub(crate) fn item_speech_result_text(&self, idx: usize) -> Result<Option<String>> {
let mut item: *const flItem = ptr::null();
let status =
unsafe { (self.api.inference_api().Response_GetItem)(self.ptr, idx, &mut item) };
self.api.check(status)?;
unsafe { read_speech_result_text(&self.api, item) }
}
pub(crate) fn item(&self, idx: usize) -> Result<Item> {
let mut item: *const flItem = ptr::null();
let status =
unsafe { (self.api.inference_api().Response_GetItem)(self.ptr, idx, &mut item) };
self.api.check(status)?;
unsafe { item_from_native(&self.api, item) }
}
pub(crate) fn items(&self) -> Result<Vec<Item>> {
let count = self.item_count();
let mut out = Vec::with_capacity(count);
for i in 0..count {
out.push(self.item(i)?);
}
Ok(out)
}
pub(crate) fn finish_reason(&self) -> flFinishReason {
unsafe { (self.api.inference_api().Response_GetFinishReason)(self.ptr) }
}
pub(crate) fn usage(&self) -> Result<(i64, i64, i64)> {
let mut usage = flUsage {
version: FOUNDRY_LOCAL_API_VERSION,
prompt_tokens: 0,
completion_tokens: 0,
total_tokens: 0,
};
let status = unsafe { (self.api.inference_api().Response_GetUsage)(self.ptr, &mut usage) };
self.api.check(status)?;
Ok((
usage.prompt_tokens,
usage.completion_tokens,
usage.total_tokens,
))
}
}
impl Drop for NativeResponse {
fn drop(&mut self) {
if !self.ptr.is_null() {
unsafe { (self.api.inference_api().Response_Release)(self.ptr) };
self.ptr = ptr::null_mut();
}
}
}
pub(crate) struct NativeItemQueue {
api: Arc<Api>,
ptr: *mut flItemQueue,
}
unsafe impl Send for NativeItemQueue {}
unsafe impl Sync for NativeItemQueue {}
impl NativeItemQueue {
pub(crate) fn new(api: Arc<Api>) -> Result<Self> {
let mut ptr: *mut flItemQueue = ptr::null_mut();
api.check(unsafe { (api.item_api().ItemQueue_Create)(&mut ptr) })?;
Ok(Self { api, ptr })
}
pub(crate) fn as_item_ptr(&self) -> *mut flItem {
self.ptr as *mut flItem
}
pub(crate) fn push_item(&self, item: *mut flItem) -> Result<()> {
self.api
.check(unsafe { (self.api.item_api().ItemQueue_Push)(self.ptr, item) })
}
pub(crate) fn push_bytes(&self, data: &[u8], item_type: flItemType) -> Result<()> {
let item = make_bytes_item(&self.api, data, item_type)?;
self.push_item(item)
}
pub(crate) fn push_value(&self, item: &Item) -> Result<()> {
let native = item_to_native(&self.api, item)?;
self.push_item(native)
}
pub(crate) fn try_pop_value(&self) -> Result<Option<Item>> {
let mut item: *mut flItem = ptr::null_mut();
let popped = unsafe { (self.api.item_api().ItemQueue_TryPop)(self.ptr, &mut item) };
if !popped || item.is_null() {
return Ok(None);
}
let decoded = unsafe { item_from_native(&self.api, item) };
unsafe { (self.api.item_api().Item_Release)(item) };
decoded.map(Some)
}
pub(crate) fn size(&self) -> usize {
unsafe { (self.api.item_api().ItemQueue_Size)(self.ptr) }
}
pub(crate) fn is_finished(&self) -> bool {
unsafe { (self.api.item_api().ItemQueue_IsFinished)(self.ptr) }
}
pub(crate) fn mark_finished(&self) {
unsafe { (self.api.item_api().ItemQueue_MarkFinished)(self.ptr) };
}
}
impl Drop for NativeItemQueue {
fn drop(&mut self) {
if !self.ptr.is_null() {
unsafe { (self.api.item_api().Item_Release)(self.ptr as *mut flItem) };
self.ptr = ptr::null_mut();
}
}
}
pub(crate) struct NativeSession {
pub(crate) api: Arc<Api>,
ptr: *mut flSession,
op_lock: std::sync::Mutex<()>,
_manager: Arc<NativeManager>,
}
unsafe impl Send for NativeSession {}
unsafe impl Sync for NativeSession {}
impl NativeSession {
pub(crate) fn create(model: &NativeModel) -> Result<Self> {
let api = Arc::clone(&model.api);
let mut ptr: *mut flSession = ptr::null_mut();
api.check(unsafe { (api.inference_api().Session_Create)(model.ptr, &mut ptr) })?;
Ok(Self {
api,
ptr,
op_lock: std::sync::Mutex::new(()),
_manager: model.manager(),
})
}
pub(crate) fn lock_ops(&self) -> std::sync::MutexGuard<'_, ()> {
self.op_lock
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
}
pub(crate) fn set_streaming_callback(
&self,
callback: flStreamingCallback,
user_data: *mut std::ffi::c_void,
) -> Result<()> {
let status = unsafe {
(self.api.inference_api().Session_SetStreamingCallback)(self.ptr, callback, user_data)
};
self.api.check(status)
}
pub(crate) fn set_options(&self, options: *const flKeyValuePairs) -> Result<()> {
let _guard = self.lock_ops();
let status = unsafe { (self.api.inference_api().Session_SetOptions)(self.ptr, options) };
self.api.check(status)
}
pub(crate) fn add_tool_definition(
&self,
name: &str,
description: Option<&str>,
json_schema: &str,
) -> Result<()> {
let _guard = self.lock_ops();
let name_c = super::api::to_cstring(name)?;
let desc_c = super::api::to_cstring(description.unwrap_or(""))?;
let schema_c = super::api::to_cstring(json_schema)?;
let def = flToolDefinition {
version: FOUNDRY_LOCAL_API_VERSION,
name: name_c.as_ptr(),
description: desc_c.as_ptr(),
json_schema: schema_c.as_ptr(),
};
let status =
unsafe { (self.api.inference_api().Session_AddToolDefinition)(self.ptr, &def) };
self.api.check(status)
}
pub(crate) fn remove_tool_definition(&self, name: &str) -> Result<bool> {
let _guard = self.lock_ops();
let name_c = super::api::to_cstring(name)?;
let mut removed = false;
let status = unsafe {
(self.api.inference_api().Session_RemoveToolDefinition)(
self.ptr,
name_c.as_ptr(),
&mut removed,
)
};
self.api.check(status)?;
Ok(removed)
}
pub(crate) fn turn_count(&self) -> usize {
let _guard = self.lock_ops();
unsafe { (self.api.inference_api().Session_GetTurnCount)(self.ptr) }
}
pub(crate) fn undo_turns(&self, count: usize) -> Result<()> {
let _guard = self.lock_ops();
let status = unsafe { (self.api.inference_api().Session_UndoTurns)(self.ptr, count) };
self.api.check(status)
}
pub(crate) fn process_request(&self, request: &NativeRequest) -> Result<NativeResponse> {
let mut resp: *mut flResponse = ptr::null_mut();
let status = unsafe {
(self.api.inference_api().Session_ProcessRequest)(self.ptr, request.ptr, &mut resp)
};
self.api.check(status)?;
Ok(NativeResponse {
api: Arc::clone(&self.api),
ptr: resp,
})
}
pub(crate) fn run_openai_json(&self, request_json: &str) -> Result<String> {
let _guard = self.lock_ops();
let request = NativeRequest::new(Arc::clone(&self.api))?;
let item = make_openai_json_item(&self.api, request_json)?;
request.add_item(item, true)?;
let response = self.process_request(&request)?;
if response.item_count() == 0 {
return Err(FoundryLocalError::CommandExecution {
reason: "Native response contained no items".into(),
});
}
response
.item_text(0)?
.ok_or_else(|| FoundryLocalError::CommandExecution {
reason: "Native response item was not readable text".into(),
})
}
}
impl Drop for NativeSession {
fn drop(&mut self) {
if !self.ptr.is_null() {
unsafe { (self.api.inference_api().Session_Release)(self.ptr) };
self.ptr = ptr::null_mut();
}
}
}
struct StreamCtx {
api: Arc<Api>,
tx: UnboundedSender<Result<String>>,
transform: StreamTransform,
}
unsafe extern "C" fn stream_trampoline(
data: flStreamingCallbackData,
user_data: *mut std::ffi::c_void,
) -> c_int {
if user_data.is_null() {
return 0;
}
let result = catch_unwind(AssertUnwindSafe(|| {
let ctx = &*(user_data as *const StreamCtx);
let queue = data.item_queue;
if queue.is_null() {
return 0;
}
let item_api = ctx.api.item_api();
loop {
let mut item: *mut flItem = ptr::null_mut();
let popped = (item_api.ItemQueue_TryPop)(queue, &mut item);
if !popped {
break;
}
if item.is_null() {
continue;
}
let text = read_text_item(&ctx.api, item);
(item_api.Item_Release)(item);
let text = match text {
Ok(text) => text,
Err(error) => {
let _ = ctx.tx.send(Err(error));
return 1; }
};
if let Some(text) = text {
if let Some(transformed) = (ctx.transform)(text) {
if ctx.tx.send(Ok(transformed)).is_err() {
return 1; }
}
}
}
0
}));
result.unwrap_or(1)
}
pub(crate) fn run_openai_json_streaming(
session: NativeSession,
request_json: String,
transform: StreamTransform,
) -> UnboundedReceiver<Result<String>> {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel::<Result<String>>();
tokio::task::spawn_blocking(move || {
let ctx = Box::new(StreamCtx {
api: Arc::clone(&session.api),
tx: tx.clone(),
transform,
});
let ctx_ptr = &*ctx as *const StreamCtx as *mut std::ffi::c_void;
let guard = session.lock_ops();
if let Err(e) = session.set_streaming_callback(Some(stream_trampoline), ctx_ptr) {
let _ = tx.send(Err(e));
return;
}
let run = (|| -> Result<()> {
let request = NativeRequest::new(Arc::clone(&session.api))?;
let item = make_openai_json_item(&session.api, &request_json)?;
request.add_item(item, true)?;
let _response = session.process_request(&request)?;
Ok(())
})();
if let Err(e) = run {
let _ = tx.send(Err(e));
}
let _ = session.set_streaming_callback(None, ptr::null_mut());
drop(guard);
drop(ctx);
drop(session);
});
rx
}
struct ItemStreamCtx {
api: Arc<Api>,
tx: UnboundedSender<Result<Item>>,
}
fn duplicate_stream_error(error: &FoundryLocalError) -> FoundryLocalError {
match error {
FoundryLocalError::Native { code, message } => FoundryLocalError::Native {
code: *code,
message: message.clone(),
},
FoundryLocalError::LibraryLoad { reason } => FoundryLocalError::LibraryLoad {
reason: reason.clone(),
},
FoundryLocalError::CommandExecution { reason } => FoundryLocalError::CommandExecution {
reason: reason.clone(),
},
FoundryLocalError::InvalidConfiguration { reason } => {
FoundryLocalError::InvalidConfiguration {
reason: reason.clone(),
}
}
FoundryLocalError::ModelOperation { reason } => FoundryLocalError::ModelOperation {
reason: reason.clone(),
},
FoundryLocalError::Validation { reason } => FoundryLocalError::Validation {
reason: reason.clone(),
},
FoundryLocalError::Internal { reason } => FoundryLocalError::Internal {
reason: reason.clone(),
},
FoundryLocalError::HttpRequest(_)
| FoundryLocalError::Serialization(_)
| FoundryLocalError::Io(_) => FoundryLocalError::Internal {
reason: error.to_string(),
},
}
}
unsafe extern "C" fn item_stream_trampoline(
data: flStreamingCallbackData,
user_data: *mut std::ffi::c_void,
) -> c_int {
if user_data.is_null() {
return 0;
}
let result = catch_unwind(AssertUnwindSafe(|| {
let ctx = &*(user_data as *const ItemStreamCtx);
let queue = data.item_queue;
if queue.is_null() {
return 0;
}
let item_api = ctx.api.item_api();
loop {
let mut item: *mut flItem = ptr::null_mut();
let popped = (item_api.ItemQueue_TryPop)(queue, &mut item);
if !popped {
break;
}
if item.is_null() {
continue;
}
let decoded = item_from_native(&ctx.api, item);
(item_api.Item_Release)(item);
match decoded {
Ok(decoded) => {
if ctx.tx.send(Ok(decoded)).is_err() {
return 1; }
}
Err(error) => {
let _ = ctx.tx.send(Err(error));
return 1;
}
}
}
0
}));
result.unwrap_or(1)
}
pub(crate) fn run_item_streaming(
session: Arc<NativeSession>,
items: Vec<Item>,
input_queue: Option<Arc<NativeItemQueue>>,
option_pairs: Vec<(String, String)>,
) -> (
UnboundedReceiver<Result<Item>>,
oneshot::Receiver<Result<crate::response::Response>>,
) {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel::<Result<Item>>();
let (response_tx, response_rx) = oneshot::channel();
tokio::task::spawn_blocking(move || {
let ctx = Box::new(ItemStreamCtx {
api: Arc::clone(&session.api),
tx: tx.clone(),
});
let ctx_ptr = &*ctx as *const ItemStreamCtx as *mut std::ffi::c_void;
let guard = session.lock_ops();
if let Err(e) = session.set_streaming_callback(Some(item_stream_trampoline), ctx_ptr) {
let _ = tx.send(Err(duplicate_stream_error(&e)));
let _ = response_tx.send(Err(e));
return;
}
let run = (|| -> Result<crate::response::Response> {
let request = NativeRequest::new(Arc::clone(&session.api))?;
for item in &items {
request.add_item_value(item)?;
}
if let Some(queue) = &input_queue {
request.add_input_queue(queue)?;
}
if !option_pairs.is_empty() {
let kvps = Kvps::from_pairs(
Arc::clone(&session.api),
option_pairs.iter().map(|(k, v)| (k.as_str(), v.as_str())),
)?;
request.set_options(kvps.as_ptr())?;
}
let response = session.process_request(&request)?;
crate::response::Response::from_native(&response)
})();
if let Err(e) = &run {
let _ = tx.send(Err(duplicate_stream_error(e)));
}
let _ = session.set_streaming_callback(None, ptr::null_mut());
drop(guard);
drop(ctx);
drop(tx);
drop(session);
let _ = response_tx.send(run);
});
(rx, response_rx)
}