use crate::rpc::error::RpcError;
use llama_cpp_sys_4 as sys;
use std::ffi::CString;
use std::path::Path;
use std::ptr::NonNull;
pub fn serve(
endpoint: &str,
cache_dir: Option<&Path>,
n_threads: usize,
devices: &[NonNull<sys::ggml_backend_device>],
) -> Result<(), RpcError> {
if devices.is_empty() {
return Err(RpcError::ServerError {
message: "at least one device must be exposed".to_owned(),
});
}
if devices.len() > sys::GGML_RPC_MAX_SERVERS as usize {
return Err(RpcError::ServerError {
message: format!(
"{} devices requested but llama.cpp serves at most {}",
devices.len(),
sys::GGML_RPC_MAX_SERVERS
),
});
}
let c_endpoint = CString::new(endpoint)?;
let c_cache_dir = cache_dir
.map(|dir| {
let dir = dir.to_str().ok_or_else(|| RpcError::InvalidEndpoint {
endpoint: dir.display().to_string(),
})?;
CString::new(dir).map_err(RpcError::from)
})
.transpose()?;
let mut device_ptrs: Vec<sys::ggml_backend_dev_t> =
devices.iter().map(|d| d.as_ptr()).collect();
unsafe {
sys::ggml_backend_rpc_start_server(
c_endpoint.as_ptr(),
c_cache_dir
.as_ref()
.map_or(std::ptr::null(), |d| d.as_ptr()),
n_threads,
device_ptrs.len(),
device_ptrs.as_mut_ptr(),
);
}
Ok(())
}
pub fn add_rpc_server(endpoint: &str) -> Result<NonNull<sys::ggml_backend_reg>, RpcError> {
let c_endpoint = CString::new(endpoint)?;
let reg = unsafe { sys::ggml_backend_rpc_add_server(c_endpoint.as_ptr()) };
NonNull::new(reg).ok_or_else(|| RpcError::InitializationFailed {
endpoint: endpoint.to_string(),
})
}