use std::ffi::{CString, NulError};
use std::path::{Path, PathBuf};
use std::ptr::NonNull;
use ik_llama_cpp_sys as sys;
#[derive(Debug)]
#[repr(transparent)]
#[allow(clippy::module_name_repetitions)]
pub struct LlamaLoraAdapter {
pub(crate) lora_adapter: NonNull<sys::llama_lora_adapter>,
}
#[derive(Debug, Eq, PartialEq, thiserror::Error)]
pub enum LlamaLoraAdapterInitError {
#[error("null byte in string {0}")]
NullError(#[from] NulError),
#[error("null result from llama cpp")]
NullResult,
#[error("failed to convert path {0} to str")]
PathToStrError(PathBuf),
}
#[derive(Debug, Eq, PartialEq, thiserror::Error)]
pub enum LlamaLoraAdapterSetError {
#[error("error code from llama cpp")]
ErrorResult(i32),
}
#[derive(Debug, Eq, PartialEq, thiserror::Error)]
pub enum LlamaLoraAdapterRemoveError {
#[error("error code from llama cpp")]
ErrorResult(i32),
}
impl crate::model::LlamaModel {
pub fn lora_adapter_init(
&self,
path: &Path,
) -> Result<LlamaLoraAdapter, LlamaLoraAdapterInitError> {
debug_assert!(path.exists(), "{path:?} does not exist");
let path_str = path
.to_str()
.ok_or_else(|| LlamaLoraAdapterInitError::PathToStrError(path.to_path_buf()))?;
let cstr = CString::new(path_str)?;
let adapter = unsafe { sys::llama_lora_adapter_init(self.model.as_ptr(), cstr.as_ptr()) };
let adapter = NonNull::new(adapter).ok_or(LlamaLoraAdapterInitError::NullResult)?;
tracing::debug!(?path, "Initialized lora adapter");
Ok(LlamaLoraAdapter {
lora_adapter: adapter,
})
}
}
impl crate::context::LlamaContext<'_> {
pub fn lora_adapter_set(
&mut self,
adapter: &LlamaLoraAdapter,
scale: f32,
) -> Result<(), LlamaLoraAdapterSetError> {
let err_code = unsafe {
sys::llama_lora_adapter_set(self.context.as_ptr(), adapter.lora_adapter.as_ptr(), scale)
};
if err_code != 0 {
return Err(LlamaLoraAdapterSetError::ErrorResult(err_code));
}
tracing::debug!(scale, "Set lora adapter");
Ok(())
}
pub fn lora_adapter_remove(
&mut self,
adapter: &LlamaLoraAdapter,
) -> Result<(), LlamaLoraAdapterRemoveError> {
let err_code = unsafe {
sys::llama_lora_adapter_remove(self.context.as_ptr(), adapter.lora_adapter.as_ptr())
};
if err_code != 0 {
return Err(LlamaLoraAdapterRemoveError::ErrorResult(err_code));
}
tracing::debug!("Removed lora adapter");
Ok(())
}
pub fn lora_adapter_clear(&mut self) {
unsafe { sys::llama_lora_adapter_clear(self.context.as_ptr()) }
tracing::debug!("Cleared lora adapters");
}
}