use std::time::Duration;
use daimon_core::{EmbeddingModel, Result};
use crate::openai_compat::{EmbedRequest, Http, api_error, parse_embed_response};
const DEFAULT_BASE_URL: &str = "http://localhost:8080";
#[derive(Debug)]
pub struct LlamaCppEmbedding {
http: Http,
model: Option<String>,
dimensions: usize,
}
impl Default for LlamaCppEmbedding {
fn default() -> Self {
Self::new()
}
}
impl LlamaCppEmbedding {
pub fn new() -> Self {
Self {
http: Http::new(DEFAULT_BASE_URL),
model: None,
dimensions: 768,
}
}
pub fn with_base_url(mut self, url: impl Into<String>) -> Self {
self.http.set_base_url(url);
self
}
pub fn with_model(mut self, name: impl Into<String>) -> Self {
self.model = Some(name.into());
self
}
pub fn with_api_key(mut self, key: impl Into<String>) -> Self {
self.http.set_api_key(key);
self
}
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.http.set_timeout(timeout);
self
}
pub fn with_max_retries(mut self, retries: u32) -> Self {
self.http.set_max_retries(retries);
self
}
pub fn allow_plaintext_api_key(mut self) -> Self {
self.http.set_allow_plaintext_api_key(true);
self
}
pub fn with_dimensions(mut self, dims: usize) -> Self {
self.dimensions = dims;
self
}
}
impl EmbeddingModel for LlamaCppEmbedding {
async fn embed(&self, texts: &[&str]) -> Result<Vec<Vec<f32>>> {
let body = EmbedRequest {
model: self.model.as_deref(),
input: texts,
};
let resp = self.http.post("/v1/embeddings", &body).await?;
let status = resp.status();
if !status.is_success() {
let text = resp.text().await.unwrap_or_default();
return Err(api_error(status, &text, "llama.cpp"));
}
let bytes = resp.bytes().await.map_err(|e| {
daimon_core::DaimonError::Model(format!("llama.cpp embedding read error: {e}"))
})?;
parse_embed_response(&bytes, "llama.cpp")
}
fn dimensions(&self) -> usize {
self.dimensions
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_builder_defaults() {
let embed = LlamaCppEmbedding::new();
assert_eq!(embed.dimensions, 768);
assert!(embed.model.is_none());
}
#[test]
fn test_builder_chain() {
let embed = LlamaCppEmbedding::new()
.with_base_url("http://gpu-box:8080")
.with_model("nomic-embed")
.with_dimensions(1024)
.with_max_retries(5);
assert_eq!(embed.model.as_deref(), Some("nomic-embed"));
assert_eq!(embed.dimensions, 1024);
assert_eq!(embed.http.max_retries(), 5);
}
#[tokio::test]
async fn test_plaintext_api_key_over_http_is_blocked_by_default() {
let embed = LlamaCppEmbedding::new().with_api_key("secret");
let err = embed.embed(&["hello"]).await.unwrap_err();
assert!(matches!(err, daimon_core::DaimonError::Builder(_)));
}
#[tokio::test]
async fn test_plaintext_api_key_allowed_when_opted_in() {
let embed = LlamaCppEmbedding::new()
.with_base_url("http://localhost:1")
.with_api_key("secret")
.with_max_retries(0)
.allow_plaintext_api_key();
let err = embed.embed(&["hello"]).await.unwrap_err();
assert!(!matches!(err, daimon_core::DaimonError::Builder(_)));
}
}