use crate::error::{Error, Result};
use crate::ffi;
use std::ffi::CString;
use std::ptr::NonNull;
use std::sync::{Arc, Condvar, Mutex};
pub struct LanguageModelSession {
ptr: NonNull<std::ffi::c_void>,
}
unsafe impl Send for LanguageModelSession {}
unsafe impl Sync for LanguageModelSession {}
impl LanguageModelSession {
pub fn new() -> Result<Self> {
Self::with_instructions_opt(None)
}
pub fn with_instructions(instructions: &str) -> Result<Self> {
Self::with_instructions_opt(Some(instructions))
}
fn with_instructions_opt(instructions: Option<&str>) -> Result<Self> {
if !unsafe { ffi::fm_check_availability() } {
return Err(Error::ModelNotAvailable);
}
let c_instructions = match instructions {
Some(s) => Some(
CString::new(s)
.map_err(|_| Error::InvalidInput("Instructions contain null byte".into()))?,
),
None => None,
};
let ptr = unsafe {
ffi::fm_create_session(
c_instructions
.as_ref()
.map_or(std::ptr::null(), |s| s.as_ptr()),
)
};
NonNull::new(ptr)
.map(|ptr| Self { ptr })
.ok_or_else(|| Error::InternalError("Failed to create session".into()))
}
pub fn from_transcript_json(transcript_json: &str) -> Result<Self> {
if !unsafe { ffi::fm_check_availability() } {
return Err(Error::ModelNotAvailable);
}
let c_json = CString::new(transcript_json)
.map_err(|_| Error::InvalidInput("Transcript JSON contains null byte".into()))?;
let ptr = unsafe { ffi::fm_create_session_from_transcript(c_json.as_ptr()) };
NonNull::new(ptr)
.map(|ptr| Self { ptr })
.ok_or_else(|| Error::InternalError("Failed to restore session from transcript".into()))
}
pub fn transcript_json(&self) -> Result<String> {
let json_ptr = unsafe { ffi::fm_get_transcript_json(self.ptr.as_ptr()) };
if json_ptr.is_null() {
return Ok("[]".to_string());
}
let json = unsafe {
let s = std::ffi::CStr::from_ptr(json_ptr)
.to_string_lossy()
.into_owned();
ffi::fm_free_string(json_ptr);
s
};
Ok(json)
}
pub fn response(&self, prompt: &str) -> Result<String> {
if prompt.is_empty() {
return Err(Error::InvalidInput("Prompt cannot be empty".into()));
}
let c_prompt = CString::new(prompt)
.map_err(|_| Error::InvalidInput("Prompt contains null byte".into()))?;
let state = Arc::new((Mutex::new(ResponseState::default()), Condvar::new()));
let state_ptr = Box::into_raw(Box::new(Arc::clone(&state)));
unsafe {
ffi::fm_session_response(
self.ptr.as_ptr(),
c_prompt.as_ptr(),
state_ptr as *mut _,
response_chunk_callback,
response_done_callback,
response_error_callback,
);
}
let (mutex, cvar) = &*state;
let mut response_state = mutex.lock().map_err(|_| Error::PoisonError)?;
while !response_state.finished {
response_state = cvar.wait(response_state).map_err(|_| Error::PoisonError)?;
}
if let Some(error) = &response_state.error {
if error.contains("not available") {
return Err(Error::ModelNotAvailable);
}
return Err(Error::GenerationError(error.clone()));
}
Ok(response_state.text.clone())
}
pub fn stream_response<F>(&self, prompt: &str, on_chunk: F) -> Result<()>
where
F: FnMut(&str),
{
if prompt.is_empty() {
return Err(Error::InvalidInput("Prompt cannot be empty".into()));
}
let c_prompt = CString::new(prompt)
.map_err(|_| Error::InvalidInput("Prompt contains null byte".into()))?;
let state = Arc::new((Mutex::new(StreamState::default()), Condvar::new()));
let user_data = Box::into_raw(Box::new((
Arc::clone(&state),
Box::new(on_chunk) as Box<dyn FnMut(&str)>,
)));
unsafe {
ffi::fm_session_stream(
self.ptr.as_ptr(),
c_prompt.as_ptr(),
user_data as *mut _,
stream_chunk_callback,
stream_done_callback,
stream_error_callback,
);
}
let (mutex, cvar) = &*state;
let mut stream_state = mutex.lock().map_err(|_| Error::PoisonError)?;
while !stream_state.finished {
stream_state = cvar.wait(stream_state).map_err(|_| Error::PoisonError)?;
}
if let Some(error) = &stream_state.error {
if error.contains("not available") {
return Err(Error::ModelNotAvailable);
}
return Err(Error::GenerationError(error.clone()));
}
Ok(())
}
pub fn cancel_stream(&self) {
unsafe {
ffi::fm_session_cancel_stream(self.ptr.as_ptr());
}
}
}
impl Drop for LanguageModelSession {
fn drop(&mut self) {
unsafe {
ffi::fm_destroy_session(self.ptr.as_ptr());
}
}
}
#[derive(Default)]
struct ResponseState {
text: String,
finished: bool,
error: Option<String>,
}
#[derive(Default)]
struct StreamState {
finished: bool,
error: Option<String>,
}
extern "C" fn response_chunk_callback(
chunk: *const std::os::raw::c_char,
user_data: *mut std::os::raw::c_void,
) {
if chunk.is_null() || user_data.is_null() {
return;
}
unsafe {
let state = &*(user_data as *const Arc<(Mutex<ResponseState>, Condvar)>);
let chunk_str = std::ffi::CStr::from_ptr(chunk).to_string_lossy();
let (mutex, _) = &**state;
if let Ok(mut response_state) = mutex.lock() {
response_state.text.push_str(&chunk_str);
}
}
}
extern "C" fn response_done_callback(user_data: *mut std::os::raw::c_void) {
if user_data.is_null() {
return;
}
unsafe {
let state = Box::from_raw(user_data as *mut Arc<(Mutex<ResponseState>, Condvar)>);
let (mutex, cvar) = &**state;
if let Ok(mut response_state) = mutex.lock() {
response_state.finished = true;
cvar.notify_all();
}
}
}
extern "C" fn response_error_callback(
error: *const std::os::raw::c_char,
user_data: *mut std::os::raw::c_void,
) {
if user_data.is_null() {
return;
}
unsafe {
let state = Box::from_raw(user_data as *mut Arc<(Mutex<ResponseState>, Condvar)>);
let (mutex, cvar) = &**state;
if let Ok(mut response_state) = mutex.lock() {
if !error.is_null() {
let error_str = std::ffi::CStr::from_ptr(error)
.to_string_lossy()
.into_owned();
response_state.error = Some(error_str);
}
response_state.finished = true;
cvar.notify_all();
}
}
}
type StreamCallback = Box<dyn FnMut(&str)>;
type StreamUserData = (Arc<(Mutex<StreamState>, Condvar)>, StreamCallback);
extern "C" fn stream_chunk_callback(
chunk: *const std::os::raw::c_char,
user_data: *mut std::os::raw::c_void,
) {
if chunk.is_null() || user_data.is_null() {
return;
}
unsafe {
let data = &mut *(user_data as *mut StreamUserData);
let chunk_str = std::ffi::CStr::from_ptr(chunk).to_string_lossy();
(data.1)(&chunk_str);
}
}
extern "C" fn stream_done_callback(user_data: *mut std::os::raw::c_void) {
if user_data.is_null() {
return;
}
unsafe {
let data = Box::from_raw(user_data as *mut StreamUserData);
let (mutex, cvar) = &*data.0;
if let Ok(mut stream_state) = mutex.lock() {
stream_state.finished = true;
cvar.notify_all();
}
}
}
extern "C" fn stream_error_callback(
error: *const std::os::raw::c_char,
user_data: *mut std::os::raw::c_void,
) {
if user_data.is_null() {
return;
}
unsafe {
let data = Box::from_raw(user_data as *mut StreamUserData);
let (mutex, cvar) = &*data.0;
if let Ok(mut stream_state) = mutex.lock() {
if !error.is_null() {
let error_str = std::ffi::CStr::from_ptr(error)
.to_string_lossy()
.into_owned();
stream_state.error = Some(error_str);
}
stream_state.finished = true;
cvar.notify_all();
}
}
}