use std::ffi::{c_char, CStr, NulError};
use llama_cpp_sys_4 as sys;
use crate::token::LlamaToken;
#[derive(Debug, thiserror::Error)]
pub enum ShimError {
#[error("argument contained an interior NUL byte")]
Nul(#[from] NulError),
#[error("invalid argument passed to the shim")]
InvalidArg,
#[error("input was not valid JSON: {0}")]
BadJson(String),
#[error("llama.cpp call failed: {0}")]
Failed(String),
#[error("could not construct: {0}")]
Init(String),
#[error("llama.cpp returned non-UTF-8 output")]
Utf8(#[from] std::string::FromUtf8Error),
#[error("shim reported an out-of-range offset")]
CorruptResult,
}
pub(crate) type Result<T> = std::result::Result<T, ShimError>;
pub(crate) fn last_error() -> String {
let ptr = unsafe { sys::llama_shim_last_error() };
if ptr.is_null() {
return String::new();
}
unsafe { CStr::from_ptr(ptr) }
.to_string_lossy()
.into_owned()
}
pub(crate) fn check_status(status: i32) -> Result<()> {
match status {
sys::LLAMA_SHIM_OK => Ok(()),
sys::LLAMA_SHIM_INVALID_ARG => Err(ShimError::InvalidArg),
sys::LLAMA_SHIM_BAD_JSON => Err(ShimError::BadJson(last_error())),
_ => Err(ShimError::Failed(last_error())),
}
}
pub(crate) fn read_string<F>(mut call: F) -> Result<String>
where
F: FnMut(*mut c_char, usize, *mut usize) -> i32,
{
let mut needed: usize = 0;
let status = call(std::ptr::null_mut(), 0, &raw mut needed);
if status != sys::LLAMA_SHIM_BUFFER_TOO_SMALL {
check_status(status)?;
}
if needed == 0 {
return Ok(String::new());
}
let mut buf = vec![0u8; needed];
let status = call(buf.as_mut_ptr().cast::<c_char>(), buf.len(), &raw mut needed);
check_status(status)?;
let end = buf.iter().position(|b| *b == 0).unwrap_or(buf.len());
buf.truncate(end);
String::from_utf8(buf).map_err(ShimError::from)
}
pub(crate) fn read_tokens<F>(mut call: F) -> Result<Vec<LlamaToken>>
where
F: FnMut(*mut i32, usize, *mut usize) -> i32,
{
let mut needed: usize = 0;
let status = call(std::ptr::null_mut(), 0, &raw mut needed);
if status != sys::LLAMA_SHIM_BUFFER_TOO_SMALL {
check_status(status)?;
}
if needed == 0 {
return Ok(Vec::new());
}
let mut buf = vec![0i32; needed];
let status = call(buf.as_mut_ptr(), buf.len(), &raw mut needed);
check_status(status)?;
buf.truncate(needed);
Ok(buf.into_iter().map(LlamaToken).collect())
}
pub(crate) fn read_i32s<F>(mut call: F) -> Result<Vec<i32>>
where
F: FnMut(*mut i32, usize, *mut usize) -> i32,
{
let mut needed: usize = 0;
let status = call(std::ptr::null_mut(), 0, &raw mut needed);
if status != sys::LLAMA_SHIM_BUFFER_TOO_SMALL {
check_status(status)?;
}
if needed == 0 {
return Ok(Vec::new());
}
let mut buf = vec![0i32; needed];
let status = call(buf.as_mut_ptr(), buf.len(), &raw mut needed);
check_status(status)?;
buf.truncate(needed);
Ok(buf)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn status_codes_are_distinct_and_signed_as_documented() {
assert_eq!(sys::LLAMA_SHIM_OK, 0);
assert_eq!(
sys::LLAMA_SHIM_BUFFER_TOO_SMALL, 1,
"a size query is not a failure, so it must be positive"
);
assert_eq!(sys::LLAMA_SHIM_INVALID_ARG, -1);
assert_eq!(sys::LLAMA_SHIM_BAD_JSON, -2);
assert_eq!(sys::LLAMA_SHIM_THROWN, -3);
}
#[test]
fn check_status_maps_each_code() {
assert!(check_status(sys::LLAMA_SHIM_OK).is_ok());
assert!(matches!(
check_status(sys::LLAMA_SHIM_INVALID_ARG),
Err(ShimError::InvalidArg)
));
assert!(matches!(
check_status(sys::LLAMA_SHIM_BAD_JSON),
Err(ShimError::BadJson(_))
));
assert!(matches!(
check_status(sys::LLAMA_SHIM_THROWN),
Err(ShimError::Failed(_))
));
}
#[test]
fn last_error_is_empty_before_any_failure() {
let _ = last_error();
}
}