feat: native embeddings module with Ollama and vector store (#521)

* feat: replace fastembed with candle for local embeddings

- Introduced a new `CandleEmbedding` provider using the `candle` ML framework, eliminating C++ dependencies and enhancing performance.
- Updated configuration to default to `candle` for embedding provider settings, replacing the previous `fastembed` references.
- Added new modules for embedding providers, including `candle_embed`, `noop`, and `openai`, to support various embedding strategies.
- Enhanced the `EmbeddingProvider` interface to accommodate the new `CandleEmbedding` implementation.
- Refactored related tests to ensure comprehensive coverage of the new embedding functionalities and maintain backward compatibility with existing configurations.

* feat: update embedding provider to Ollama

- Replaced the default embedding provider from Candle to Ollama, enhancing local embedding capabilities with improved model management and GPU acceleration.
- Updated configuration defaults for embedding model and dimensions to align with Ollama specifications.
- Introduced a new Ollama embedding module, including necessary constants and functionality for embedding requests.
- Refactored related code and tests to ensure compatibility with the new provider, maintaining backward compatibility with existing configurations.

* refactor: remove Candle embedding provider and related dependencies

- Deleted the `CandleEmbedding` module and its associated files, streamlining the embedding provider architecture.
- Updated `Cargo.toml` and `Cargo.lock` to remove references to Candle-related packages, ensuring a cleaner dependency tree.
- Refactored the `EmbeddingProvider` interface to eliminate support for the Candle provider, maintaining compatibility with existing providers.
- Adjusted tests and documentation to reflect the removal of the Candle embedding functionality, ensuring clarity and consistency across the codebase.

* feat: add SQLite-backed vector store for embeddings

- Introduced a new `store` module to implement a local vector store backed by SQLite, enabling efficient storage and retrieval of text embeddings.
- Added functionality for inserting, updating, and searching embeddings using cosine similarity, enhancing the embedding management capabilities.
- Updated the `mod.rs` file to include the new `store` module and expose relevant functions for cosine similarity and vector operations.
- Enhanced documentation to provide usage examples and clarify the integration of the vector store with existing embedding providers.

* Add fs2 dependency for improved file locking in ComposeIO trigger history

* test: boost embeddings module coverage to 97%+

Add tests for batch insert mismatch error path, invalid metadata
JSON handling, disk store parent directory creation, and fake
embedding dimensions accessor. 95 tests total across the module,
all files at 97-100% line coverage.

* style: apply cargo fmt to embeddings module

* feat: enhance embedding provider creation with error handling

- Updated the `create_embedding_provider` function to return an `anyhow::Result`, allowing for immediate error reporting on unrecognized provider names.
- Added support for a "none" provider that returns a no-op embedding.
- Enhanced tests to cover new error handling paths and validate behavior for known and unknown providers, improving overall test coverage.

* feat: enhance Ollama and OpenAI embedding providers with improved handling for blank inputs and response validation

- Updated the `embed` method in both `OllamaEmbedding` and `OpenAiEmbedding` to skip blank inputs while preserving their positions in the output as zero-vectors.
- Added validation to ensure the response count matches the input count, with appropriate error handling for dimension mismatches.
- Enhanced tests to cover new behavior for blank inputs and response validation, ensuring robustness in embedding functionality.

* fix: address PR review findings for embeddings module

- Factory returns Result for unknown providers instead of silent noop
- Ollama: preserve positional alignment when blank texts are filtered
- Ollama: validate response count and dimensions before returning
- OpenAI: strict parsing errors on non-numeric embedding values
- OpenAI: skip Authorization header when api_key is empty
- OpenAI: validate response count and dimensions
- OpenAI/Ollama: add tracing at entry, success, and error paths
- Store: add store_meta table to persist and validate embedding
  provider/dimensions on open — errors on dimension mismatch
- Store: propagate row decode errors instead of filter_map(ok)
- Store: add tracing to search, insert, delete, and count paths
- Update factories.rs caller for Result-returning factory

101 tests pass, all files at 97%+ coverage.
This commit is contained in:
Steven Enamakel
2026-04-12 18:54:38 -07:00
committed by GitHub
parent 3a20599f45
commit 4cf608c2be
12 changed files with 2241 additions and 1197 deletions
Generated
+31 -596
View File
File diff suppressed because it is too large Load Diff
-2
View File
@@ -35,8 +35,6 @@ anyhow = "1.0"
async-trait = "0.1"
chacha20poly1305 = "0.10"
hex = "0.4"
fastembed = { version = "5.13", default-features = false, features = ["hf-hub-native-tls", "ort-load-dynamic"] }
ort = { version = "=2.0.0-rc.11", default-features = false, features = ["std", "ndarray", "load-dynamic"] }
tokio-util = { version = "0.7", features = ["rt"] }
tokio-tungstenite = { version = "0.24", features = ["rustls-tls-webpki-roots"] }
futures = "0.3"
@@ -109,7 +109,7 @@ fn default_true() -> bool {
}
fn default_embedding_provider() -> String {
"fastembed".into()
"ollama".into()
}
fn default_hygiene_enabled() -> bool {
true
@@ -124,10 +124,10 @@ fn default_conversation_retention_days() -> u32 {
30
}
fn default_embedding_model() -> String {
"BGESmallENV15".into()
"nomic-embed-text:latest".into()
}
fn default_embedding_dims() -> usize {
384
768
}
fn default_vector_weight() -> f64 {
0.7
+214
View File
@@ -0,0 +1,214 @@
//! Embedding providers for the OpenHuman memory system.
//!
//! Converts text into numerical vectors for semantic search. Providers:
//!
//! - **Ollama** (default): Delegates to a local Ollama server — handles model
//! management, quantization, and GPU acceleration out of the box.
//! - **OpenAI**: Cloud-based embeddings via the OpenAI API or compatible endpoints.
//! - **Noop**: A fallback provider for keyword-only search.
pub mod noop;
pub mod ollama;
pub mod openai;
pub mod store;
use std::sync::Arc;
use async_trait::async_trait;
pub use noop::NoopEmbedding;
pub use ollama::{OllamaEmbedding, DEFAULT_OLLAMA_DIMENSIONS, DEFAULT_OLLAMA_MODEL};
pub use openai::OpenAiEmbedding;
pub use store::{bytes_to_vec, cosine_similarity, vec_to_bytes, SearchResult, VectorStore};
/// Interface for embedding providers that convert text into numerical vectors.
#[async_trait]
pub trait EmbeddingProvider: Send + Sync {
/// Returns the name of the provider (e.g., "ollama", "openai").
fn name(&self) -> &str;
/// Returns the number of dimensions in the generated embeddings.
fn dimensions(&self) -> usize;
/// Generates embeddings for a batch of strings.
async fn embed(&self, texts: &[&str]) -> anyhow::Result<Vec<Vec<f32>>>;
/// Generates an embedding for a single string.
async fn embed_one(&self, text: &str) -> anyhow::Result<Vec<f32>> {
let mut results = self.embed(&[text]).await?;
results
.pop()
.ok_or_else(|| anyhow::anyhow!("Empty embedding result"))
}
}
// ── Factory ──────────────────────────────────────────────────
/// Creates an embedding provider based on the specified name and configuration.
///
/// Supported provider names:
/// - `"ollama"` → local Ollama server (default, preferred)
/// - `"openai"` → OpenAI API
/// - `"custom:<url>"` → OpenAI-compatible endpoint
/// - `"none"` → no-op (keyword-only search, no embeddings)
///
/// Returns an error for unrecognised provider names so configuration
/// mistakes surface immediately rather than silently degrading to
/// keyword-only search.
pub fn create_embedding_provider(
provider: &str,
api_key: Option<&str>,
model: &str,
dims: usize,
) -> anyhow::Result<Box<dyn EmbeddingProvider>> {
match provider {
"ollama" => Ok(Box::new(OllamaEmbedding::new("", model, dims))),
"openai" => {
let key = api_key.unwrap_or("");
Ok(Box::new(OpenAiEmbedding::new(
"https://api.openai.com",
key,
model,
dims,
)))
}
name if name.starts_with("custom:") => {
let base_url = name.strip_prefix("custom:").unwrap_or("");
let key = api_key.unwrap_or("");
Ok(Box::new(OpenAiEmbedding::new(base_url, key, model, dims)))
}
"none" => Ok(Box::new(NoopEmbedding)),
unknown => Err(anyhow::anyhow!(
"unknown embedding provider: \"{unknown}\". \
Supported: \"ollama\", \"openai\", \"custom:<url>\", \"none\""
)),
}
}
/// Returns the default local embedding provider (Ollama-backed).
pub fn default_local_embedding_provider() -> Arc<dyn EmbeddingProvider> {
Arc::new(OllamaEmbedding::default())
}
#[cfg(test)]
mod tests {
use super::*;
// ── Trait default method ─────────────────────────────────
#[test]
fn noop_name_and_dims() {
let p = NoopEmbedding;
assert_eq!(p.name(), "none");
assert_eq!(p.dimensions(), 0);
}
#[tokio::test]
async fn noop_embed_returns_empty() {
let p = NoopEmbedding;
let result = p.embed(&["hello"]).await.unwrap();
assert!(result.is_empty());
}
#[tokio::test]
async fn noop_embed_one_returns_error() {
// embed returns empty vec → pop() returns None → error from default impl
let p = NoopEmbedding;
let err = p.embed_one("hello").await.unwrap_err();
assert!(err.to_string().contains("Empty embedding result"));
}
#[tokio::test]
async fn noop_embed_empty_batch() {
let p = NoopEmbedding;
let result = p.embed(&[]).await.unwrap();
assert!(result.is_empty());
}
// ── Factory — success ────────────────────────────────────
#[test]
fn factory_ollama() {
let p = create_embedding_provider("ollama", None, DEFAULT_OLLAMA_MODEL, 768).unwrap();
assert_eq!(p.name(), "ollama");
assert_eq!(p.dimensions(), 768);
}
#[test]
fn factory_openai() {
let p = create_embedding_provider("openai", Some("key"), "text-embedding-3-small", 1536)
.unwrap();
assert_eq!(p.name(), "openai");
assert_eq!(p.dimensions(), 1536);
}
#[test]
fn factory_openai_no_api_key() {
let p = create_embedding_provider("openai", None, "text-embedding-3-small", 1536).unwrap();
assert_eq!(p.name(), "openai");
assert_eq!(p.dimensions(), 1536);
}
#[test]
fn factory_custom_url() {
let p =
create_embedding_provider("custom:http://localhost:1234", None, "model", 768).unwrap();
assert_eq!(p.name(), "openai"); // OpenAI-compatible under the hood
assert_eq!(p.dimensions(), 768);
}
#[test]
fn factory_custom_empty_url() {
let p = create_embedding_provider("custom:", None, "model", 768).unwrap();
assert_eq!(p.name(), "openai");
}
#[test]
fn factory_none() {
let p = create_embedding_provider("none", None, "", 0).unwrap();
assert_eq!(p.name(), "none");
assert_eq!(p.dimensions(), 0);
}
// ── Factory — errors ─────────────────────────────────────
#[test]
fn factory_unknown_provider_errors() {
let result = create_embedding_provider("cohere", None, "model", 1536);
let msg = result.err().expect("should be an error").to_string();
assert!(
msg.contains("cohere"),
"should include provider name: {msg}"
);
assert!(msg.contains("unknown"), "should say unknown: {msg}");
}
#[test]
fn factory_empty_string_errors() {
let result = create_embedding_provider("", None, "model", 1536);
assert!(result
.err()
.expect("should error")
.to_string()
.contains("unknown"));
}
#[test]
fn factory_fastembed_errors() {
let result = create_embedding_provider("fastembed", None, "BGESmallENV15", 384);
assert!(result
.err()
.expect("should error")
.to_string()
.contains("fastembed"));
}
// ── Default provider ─────────────────────────────────────
#[test]
fn default_local_provider_uses_ollama() {
let p = default_local_embedding_provider();
assert_eq!(p.name(), "ollama");
assert_eq!(p.dimensions(), DEFAULT_OLLAMA_DIMENSIONS);
}
}
+24
View File
@@ -0,0 +1,24 @@
//! No-op embedding provider for keyword-only search fallback.
use async_trait::async_trait;
use super::EmbeddingProvider;
/// A "no-op" embedding provider used when semantic search is disabled.
/// Returns empty vectors.
pub struct NoopEmbedding;
#[async_trait]
impl EmbeddingProvider for NoopEmbedding {
fn name(&self) -> &str {
"none"
}
fn dimensions(&self) -> usize {
0
}
async fn embed(&self, _texts: &[&str]) -> anyhow::Result<Vec<Vec<f32>>> {
Ok(Vec::new())
}
}
+527
View File
@@ -0,0 +1,527 @@
//! Ollama-based embedding provider.
//!
//! Calls the local Ollama server's `/api/embed` endpoint for embeddings.
//! This is the preferred local provider: Ollama handles model management,
//! quantization, and GPU acceleration (Metal on macOS, CUDA on Linux/Windows).
//!
//! Default model: `nomic-embed-text:latest` (768 dimensions).
use async_trait::async_trait;
use super::EmbeddingProvider;
/// Default Ollama base URL.
pub const DEFAULT_OLLAMA_URL: &str = "http://localhost:11434";
/// Default embedding model for Ollama.
pub const DEFAULT_OLLAMA_MODEL: &str = "nomic-embed-text:latest";
/// Default dimensions for nomic-embed-text.
pub const DEFAULT_OLLAMA_DIMENSIONS: usize = 768;
/// Embedding provider backed by a local Ollama instance.
///
/// Ollama must be running and have the configured model pulled.
/// On first embed call, if the model isn't available, Ollama will
/// auto-pull it (this may take a moment on first use).
pub struct OllamaEmbedding {
base_url: String,
model: String,
dims: usize,
}
impl OllamaEmbedding {
/// Creates a new Ollama embedding provider.
///
/// - `base_url`: Ollama server URL (default: `http://localhost:11434`)
/// - `model`: Model name (default: `nomic-embed-text:latest`)
/// - `dims`: Expected embedding dimensions (default: 768)
pub fn new(base_url: &str, model: &str, dims: usize) -> Self {
let base_url = if base_url.trim().is_empty() {
DEFAULT_OLLAMA_URL.to_string()
} else {
base_url.trim_end_matches('/').to_string()
};
let model = if model.trim().is_empty() {
DEFAULT_OLLAMA_MODEL.to_string()
} else {
model.trim().to_string()
};
let dims = if dims == 0 {
DEFAULT_OLLAMA_DIMENSIONS
} else {
dims
};
tracing::debug!(
target: "embeddings.ollama",
"[embeddings] OllamaEmbedding created: url={base_url}, model={model}, dims={dims}"
);
Self {
base_url,
model,
dims,
}
}
/// Creates a provider with all defaults.
pub fn default() -> Self {
Self::new(
DEFAULT_OLLAMA_URL,
DEFAULT_OLLAMA_MODEL,
DEFAULT_OLLAMA_DIMENSIONS,
)
}
/// Returns the configured base URL.
pub fn base_url(&self) -> &str {
&self.base_url
}
/// Returns the configured model name.
pub fn model(&self) -> &str {
&self.model
}
/// Build an HTTP client with proxy support.
fn http_client(&self) -> reqwest::Client {
crate::openhuman::config::build_runtime_proxy_client("embeddings.ollama")
}
/// The embed endpoint URL.
fn embed_url(&self) -> String {
format!("{}/api/embed", self.base_url)
}
}
/// Ollama `/api/embed` request body.
#[derive(serde::Serialize)]
struct OllamaEmbedRequest {
model: String,
input: Vec<String>,
}
/// Ollama `/api/embed` response body.
#[derive(serde::Deserialize)]
struct OllamaEmbedResponse {
#[serde(default)]
embeddings: Vec<Vec<f32>>,
}
#[async_trait]
impl EmbeddingProvider for OllamaEmbedding {
fn name(&self) -> &str {
"ollama"
}
fn dimensions(&self) -> usize {
self.dims
}
/// Sends texts to Ollama's embed API.
///
/// Blank/whitespace-only entries are skipped for the remote call but their
/// positions in the result are preserved as zero-vectors so the returned
/// `Vec` always has the same length as `texts`.
async fn embed(&self, texts: &[&str]) -> anyhow::Result<Vec<Vec<f32>>> {
if texts.is_empty() {
return Ok(Vec::new());
}
// Build a list of (original_index, trimmed_text) for non-blank entries.
let live: Vec<(usize, String)> = texts
.iter()
.enumerate()
.filter_map(|(i, t)| {
let trimmed = t.trim().to_string();
if trimmed.is_empty() {
None
} else {
Some((i, trimmed))
}
})
.collect();
if live.is_empty() {
// All entries were blank — return zero-vectors.
return Ok(vec![Vec::new(); texts.len()]);
}
let input: Vec<String> = live.iter().map(|(_, t)| t.clone()).collect();
tracing::debug!(
target: "embeddings.ollama",
"[embeddings] sending {} text(s) to ollama model={} ({} blank skipped)",
input.len(), self.model, texts.len() - input.len()
);
let resp = self
.http_client()
.post(self.embed_url())
.json(&OllamaEmbedRequest {
model: self.model.clone(),
input: input.clone(),
})
.send()
.await
.map_err(|e| {
anyhow::anyhow!(
"ollama embed request failed (is Ollama running at {}?): {e}",
self.base_url
)
})?;
if !resp.status().is_success() {
let status = resp.status();
let body = resp.text().await.unwrap_or_default();
let detail = body.trim();
anyhow::bail!(
"ollama embed failed with status {status}{}",
if detail.is_empty() {
String::new()
} else {
format!(": {detail}")
}
);
}
let payload: OllamaEmbedResponse = resp
.json()
.await
.map_err(|e| anyhow::anyhow!("ollama embed response parse failed: {e}"))?;
// Validate response count matches what we sent.
if payload.embeddings.len() != input.len() {
anyhow::bail!(
"ollama embed count mismatch: sent {} texts, got {} embeddings",
input.len(),
payload.embeddings.len()
);
}
// Validate dimensions on every returned vector.
for (i, vec) in payload.embeddings.iter().enumerate() {
if vec.len() != self.dims {
anyhow::bail!(
"ollama embed dimension mismatch at index {i}: expected {}, got {}",
self.dims,
vec.len()
);
}
}
tracing::debug!(
target: "embeddings.ollama",
"[embeddings] received {} embeddings, dims={}",
payload.embeddings.len(),
self.dims
);
// Reconstruct full-length result with zero-vectors for blank positions.
let mut result = vec![Vec::new(); texts.len()];
for ((orig_idx, _), embedding) in live.iter().zip(payload.embeddings.into_iter()) {
result[*orig_idx] = embedding;
}
Ok(result)
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::{extract::Json, http::StatusCode, routing::post, Router};
use std::net::SocketAddr;
/// Spin up a local axum server and return its base URL.
async fn start_mock(app: Router) -> String {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr: SocketAddr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
format!("http://127.0.0.1:{}", addr.port())
}
// ── Constructor ──────────────────────────────────────────
#[test]
fn defaults() {
let p = OllamaEmbedding::default();
assert_eq!(p.base_url, DEFAULT_OLLAMA_URL);
assert_eq!(p.model, DEFAULT_OLLAMA_MODEL);
assert_eq!(p.dims, DEFAULT_OLLAMA_DIMENSIONS);
}
#[test]
fn name_is_ollama() {
let p = OllamaEmbedding::default();
assert_eq!(p.name(), "ollama");
}
#[test]
fn custom_values() {
let p = OllamaEmbedding::new("http://gpu-box:11434/", "mxbai-embed-large", 1024);
assert_eq!(p.base_url, "http://gpu-box:11434");
assert_eq!(p.model, "mxbai-embed-large");
assert_eq!(p.dims, 1024);
}
#[test]
fn empty_values_use_defaults() {
let p = OllamaEmbedding::new("", "", 0);
assert_eq!(p.base_url, DEFAULT_OLLAMA_URL);
assert_eq!(p.model, DEFAULT_OLLAMA_MODEL);
assert_eq!(p.dims, DEFAULT_OLLAMA_DIMENSIONS);
}
#[test]
fn whitespace_only_values_use_defaults() {
let p = OllamaEmbedding::new(" ", " ", 0);
assert_eq!(p.base_url, DEFAULT_OLLAMA_URL);
assert_eq!(p.model, DEFAULT_OLLAMA_MODEL);
}
#[test]
fn trailing_slash_stripped() {
let p = OllamaEmbedding::new("http://host:1234/", "m", 1);
assert_eq!(p.base_url, "http://host:1234");
}
#[test]
fn model_trimmed() {
let p = OllamaEmbedding::new("", " nomic-embed-text ", 768);
assert_eq!(p.model, "nomic-embed-text");
}
#[test]
fn embed_url_format() {
let p = OllamaEmbedding::default();
assert_eq!(p.embed_url(), "http://localhost:11434/api/embed");
}
#[test]
fn accessor_methods() {
let p = OllamaEmbedding::new("http://x:1", "m", 42);
assert_eq!(p.base_url(), "http://x:1");
assert_eq!(p.model(), "m");
assert_eq!(p.dimensions(), 42);
}
// ── embed — empty / whitespace ──────────────────────────
#[tokio::test]
async fn empty_input_returns_empty() {
let p = OllamaEmbedding::default();
let result = p.embed(&[]).await.unwrap();
assert!(result.is_empty());
}
#[tokio::test]
async fn whitespace_only_input_returns_zero_vecs() {
let p = OllamaEmbedding::default();
let result = p.embed(&[" ", "\t", "\n"]).await.unwrap();
// Length preserved, all entries are empty zero-vectors.
assert_eq!(result.len(), 3);
assert!(result.iter().all(|v| v.is_empty()));
}
// ── embed — positional alignment ────────────────────────
#[tokio::test]
async fn embed_preserves_positions_for_blanks() {
let app = Router::new().route(
"/api/embed",
post(|Json(body): Json<serde_json::Value>| async move {
let inputs = body["input"].as_array().unwrap();
// Server receives only non-blank texts.
let embeddings: Vec<Vec<f32>> = inputs.iter().map(|_| vec![1.0, 2.0]).collect();
Json(serde_json::json!({ "embeddings": embeddings }))
}),
);
let url = start_mock(app).await;
let p = OllamaEmbedding::new(&url, "m", 2);
// Mix of blank and real texts.
let result = p.embed(&["hello", "", " ", "world"]).await.unwrap();
assert_eq!(result.len(), 4);
assert_eq!(result[0], vec![1.0, 2.0]); // real
assert!(result[1].is_empty()); // blank
assert!(result[2].is_empty()); // blank
assert_eq!(result[3], vec![1.0, 2.0]); // real
}
// ── embed — successful response ─────────────────────────
#[tokio::test]
async fn embed_success_single() {
let app = Router::new().route(
"/api/embed",
post(|Json(_body): Json<serde_json::Value>| async {
Json(serde_json::json!({
"embeddings": [[0.1, 0.2, 0.3]]
}))
}),
);
let url = start_mock(app).await;
let p = OllamaEmbedding::new(&url, "test-model", 3);
let result = p.embed(&["hello"]).await.unwrap();
assert_eq!(result.len(), 1);
assert_eq!(result[0], vec![0.1, 0.2, 0.3]);
}
#[tokio::test]
async fn embed_success_batch() {
let app = Router::new().route(
"/api/embed",
post(|Json(_body): Json<serde_json::Value>| async {
Json(serde_json::json!({
"embeddings": [[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]]
}))
}),
);
let url = start_mock(app).await;
let p = OllamaEmbedding::new(&url, "test-model", 2);
let result = p.embed(&["a", "b", "c"]).await.unwrap();
assert_eq!(result.len(), 3);
assert_eq!(result[2], vec![5.0, 6.0]);
}
#[tokio::test]
async fn embed_verifies_request_body() {
let app = Router::new().route(
"/api/embed",
post(|Json(body): Json<serde_json::Value>| async move {
assert_eq!(body["model"], "my-model");
let inputs = body["input"].as_array().unwrap();
assert_eq!(inputs.len(), 1);
assert_eq!(inputs[0], "test text");
Json(serde_json::json!({ "embeddings": [[1.0]] }))
}),
);
let url = start_mock(app).await;
let p = OllamaEmbedding::new(&url, "my-model", 1);
p.embed(&["test text"]).await.unwrap();
}
// ── embed — error paths ─────────────────────────────────
#[tokio::test]
async fn embed_server_error_with_body() {
let app = Router::new().route(
"/api/embed",
post(|| async { (StatusCode::INTERNAL_SERVER_ERROR, "model crashed") }),
);
let url = start_mock(app).await;
let p = OllamaEmbedding::new(&url, "m", 1);
let err = p.embed(&["hi"]).await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("500"), "should contain status code: {msg}");
assert!(msg.contains("model crashed"), "should contain body: {msg}");
}
#[tokio::test]
async fn embed_server_error_empty_body() {
let app = Router::new().route(
"/api/embed",
post(|| async { (StatusCode::BAD_REQUEST, "") }),
);
let url = start_mock(app).await;
let p = OllamaEmbedding::new(&url, "m", 1);
let err = p.embed(&["hi"]).await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("400"), "should contain status code: {msg}");
}
#[tokio::test]
async fn embed_count_mismatch() {
let app = Router::new().route(
"/api/embed",
post(|| async {
// Return 1 embedding even though 2 texts were sent.
Json(serde_json::json!({ "embeddings": [[1.0]] }))
}),
);
let url = start_mock(app).await;
let p = OllamaEmbedding::new(&url, "m", 1);
let err = p.embed(&["a", "b"]).await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("count mismatch"), "msg: {msg}");
}
#[tokio::test]
async fn embed_dimension_mismatch() {
let app = Router::new().route(
"/api/embed",
post(|| async {
// Return 3-dim vector when provider expects 2.
Json(serde_json::json!({ "embeddings": [[1.0, 2.0, 3.0]] }))
}),
);
let url = start_mock(app).await;
let p = OllamaEmbedding::new(&url, "m", 2);
let err = p.embed(&["hi"]).await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("dimension mismatch"), "msg: {msg}");
}
#[tokio::test]
async fn embed_empty_embeddings_array() {
let app = Router::new().route(
"/api/embed",
post(|| async { Json(serde_json::json!({ "embeddings": [] })) }),
);
let url = start_mock(app).await;
let p = OllamaEmbedding::new(&url, "m", 1);
let err = p.embed(&["hi"]).await.unwrap_err();
assert!(err.to_string().contains("count mismatch"));
}
#[tokio::test]
async fn embed_malformed_json_response() {
let app = Router::new().route(
"/api/embed",
post(|| async { (StatusCode::OK, "not json at all") }),
);
let url = start_mock(app).await;
let p = OllamaEmbedding::new(&url, "m", 1);
let err = p.embed(&["hi"]).await.unwrap_err();
assert!(err.to_string().contains("parse failed"));
}
#[tokio::test]
async fn embed_connection_refused() {
let p = OllamaEmbedding::new("http://127.0.0.1:1", "m", 1);
let err = p.embed(&["hi"]).await.unwrap_err();
assert!(
err.to_string().contains("is Ollama running"),
"should mention Ollama: {}",
err
);
}
// ── embed_one (trait default) ───────────────────────────
#[tokio::test]
async fn embed_one_success() {
let app = Router::new().route(
"/api/embed",
post(|| async { Json(serde_json::json!({ "embeddings": [[7.0, 8.0]] })) }),
);
let url = start_mock(app).await;
let p = OllamaEmbedding::new(&url, "m", 2);
let vec = p.embed_one("test").await.unwrap();
assert_eq!(vec, vec![7.0, 8.0]);
}
}
+550
View File
@@ -0,0 +1,550 @@
//! OpenAI-compatible embedding provider.
//!
//! Works with OpenAI, LocalAI, Ollama, and any endpoint that implements the
//! `POST /v1/embeddings` contract.
use async_trait::async_trait;
use super::EmbeddingProvider;
/// Embedding provider for OpenAI and compatible APIs (e.g., LocalAI, Ollama).
pub struct OpenAiEmbedding {
base_url: String,
api_key: String,
model: String,
dims: usize,
}
impl OpenAiEmbedding {
/// Creates a new OpenAI-style provider.
pub fn new(base_url: &str, api_key: &str, model: &str, dims: usize) -> Self {
Self {
base_url: base_url.trim_end_matches('/').to_string(),
api_key: api_key.to_string(),
model: model.to_string(),
dims,
}
}
/// Returns the configured base URL.
pub fn base_url(&self) -> &str {
&self.base_url
}
/// Returns the configured model name.
pub fn model(&self) -> &str {
&self.model
}
/// Internal helper to build an HTTP client with proxy support.
fn http_client(&self) -> reqwest::Client {
crate::openhuman::config::build_runtime_proxy_client("memory.embeddings")
}
/// Checks if the base URL includes a specific path (e.g., /api/v1).
fn has_explicit_api_path(&self) -> bool {
let Ok(url) = reqwest::Url::parse(&self.base_url) else {
return false;
};
let path = url.path().trim_end_matches('/');
!path.is_empty() && path != "/"
}
/// Checks if the URL already ends with /embeddings.
fn has_embeddings_endpoint(&self) -> bool {
let Ok(url) = reqwest::Url::parse(&self.base_url) else {
return false;
};
url.path().trim_end_matches('/').ends_with("/embeddings")
}
/// Constructs the final URL for the embeddings endpoint.
pub fn embeddings_url(&self) -> String {
if self.has_embeddings_endpoint() {
return self.base_url.clone();
}
if self.has_explicit_api_path() {
format!("{}/embeddings", self.base_url)
} else {
format!("{}/v1/embeddings", self.base_url)
}
}
}
#[async_trait]
impl EmbeddingProvider for OpenAiEmbedding {
fn name(&self) -> &str {
"openai"
}
fn dimensions(&self) -> usize {
self.dims
}
/// Sends a POST request to the embedding API.
async fn embed(&self, texts: &[&str]) -> anyhow::Result<Vec<Vec<f32>>> {
if texts.is_empty() {
return Ok(Vec::new());
}
let url = self.embeddings_url();
tracing::debug!(
target: "openai::embed",
"[openai] embed: model={}, count={}, url={}",
self.model, texts.len(), url
);
let body = serde_json::json!({
"model": self.model,
"input": texts,
});
let mut req = self
.http_client()
.post(&url)
.header("Content-Type", "application/json")
.json(&body);
// Only set Authorization header when an API key is configured.
if !self.api_key.is_empty() {
req = req.header("Authorization", format!("Bearer {}", self.api_key));
}
let resp = req.send().await?;
if !resp.status().is_success() {
let status = resp.status();
let text = resp.text().await.unwrap_or_default();
tracing::debug!(
target: "openai::embed",
"[openai] embed error: status={status}, body={text}"
);
anyhow::bail!("Embedding API error {status}: {text}");
}
let json: serde_json::Value = resp.json().await?;
let data = json
.get("data")
.and_then(|d| d.as_array())
.ok_or_else(|| anyhow::anyhow!("Invalid embedding response: missing 'data'"))?;
// Validate that the response count matches the input count.
if data.len() != texts.len() {
anyhow::bail!(
"openai embed count mismatch: sent {} texts, got {} items in 'data'",
texts.len(),
data.len()
);
}
let mut embeddings = Vec::with_capacity(data.len());
for (i, item) in data.iter().enumerate() {
let embedding = item
.get("embedding")
.and_then(|e| e.as_array())
.ok_or_else(|| {
anyhow::anyhow!("Invalid embedding item at index {i}: missing 'embedding'")
})?;
let mut vec = Vec::with_capacity(embedding.len());
for (j, v) in embedding.iter().enumerate() {
#[allow(clippy::cast_possible_truncation)]
let f = v.as_f64().ok_or_else(|| {
anyhow::anyhow!("non-numeric value at data[{i}].embedding[{j}]: {v}")
})? as f32;
vec.push(f);
}
// Validate dimensions.
if self.dims > 0 && vec.len() != self.dims {
anyhow::bail!(
"openai embed dimension mismatch at index {i}: expected {}, got {}",
self.dims,
vec.len()
);
}
embeddings.push(vec);
}
tracing::debug!(
target: "openai::embed",
"[openai] embed success: model={}, count={}, dims={}",
self.model, embeddings.len(),
embeddings.first().map(|v| v.len()).unwrap_or(0)
);
Ok(embeddings)
}
}
#[cfg(test)]
mod tests {
use super::*;
use axum::{
extract::Json,
http::{HeaderMap, StatusCode},
routing::post,
Router,
};
use std::net::SocketAddr;
async fn start_mock(app: Router) -> String {
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr: SocketAddr = listener.local_addr().unwrap();
tokio::spawn(async move {
axum::serve(listener, app).await.unwrap();
});
format!("http://127.0.0.1:{}", addr.port())
}
// ── Constructor & URL building ──────────────────────────
#[test]
fn trailing_slash_stripped() {
let p = OpenAiEmbedding::new("https://api.openai.com/", "key", "model", 1536);
assert_eq!(p.base_url, "https://api.openai.com");
}
#[test]
fn dimensions_custom() {
let p = OpenAiEmbedding::new("http://localhost", "k", "m", 384);
assert_eq!(p.dimensions(), 384);
}
#[test]
fn accessors() {
let p = OpenAiEmbedding::new("http://x", "k", "m", 1);
assert_eq!(p.base_url(), "http://x");
assert_eq!(p.model(), "m");
assert_eq!(p.name(), "openai");
}
#[test]
fn url_standard_openai() {
let p = OpenAiEmbedding::new("https://api.openai.com", "key", "model", 1536);
assert_eq!(p.embeddings_url(), "https://api.openai.com/v1/embeddings");
}
#[test]
fn url_base_with_v1_no_duplicate() {
let p = OpenAiEmbedding::new("https://api.example.com/v1", "key", "model", 1536);
assert_eq!(p.embeddings_url(), "https://api.example.com/v1/embeddings");
}
#[test]
fn url_non_v1_api_path() {
let p = OpenAiEmbedding::new(
"https://api.example.com/api/coding/v3",
"key",
"model",
1536,
);
assert_eq!(
p.embeddings_url(),
"https://api.example.com/api/coding/v3/embeddings"
);
}
#[test]
fn url_already_ends_with_embeddings() {
let p = OpenAiEmbedding::new(
"https://my-api.example.com/api/v2/embeddings",
"key",
"model",
1536,
);
assert_eq!(
p.embeddings_url(),
"https://my-api.example.com/api/v2/embeddings"
);
}
#[test]
fn url_already_ends_with_embeddings_trailing_slash() {
let p = OpenAiEmbedding::new(
"https://api.example.com/v1/embeddings/",
"key",
"model",
1536,
);
assert_eq!(p.embeddings_url(), "https://api.example.com/v1/embeddings");
}
#[test]
fn url_root_only() {
let p = OpenAiEmbedding::new("http://localhost:8080", "k", "m", 1);
assert_eq!(p.embeddings_url(), "http://localhost:8080/v1/embeddings");
}
#[test]
fn url_root_with_trailing_slash() {
let p = OpenAiEmbedding::new("http://localhost:8080/", "k", "m", 1);
assert_eq!(p.embeddings_url(), "http://localhost:8080/v1/embeddings");
}
#[test]
fn has_explicit_api_path_invalid_url() {
let p = OpenAiEmbedding::new("not-a-url", "k", "m", 1);
assert!(!p.has_explicit_api_path());
}
#[test]
fn has_embeddings_endpoint_invalid_url() {
let p = OpenAiEmbedding::new("not-a-url", "k", "m", 1);
assert!(!p.has_embeddings_endpoint());
}
// ── embed — empty input ─────────────────────────────────
#[tokio::test]
async fn empty_input_returns_empty() {
let p = OpenAiEmbedding::new("http://unused", "k", "m", 1);
let result = p.embed(&[]).await.unwrap();
assert!(result.is_empty());
}
// ── embed — success ─────────────────────────────────────
#[tokio::test]
async fn embed_success_single() {
let app = Router::new().route(
"/v1/embeddings",
post(|| async {
Json(serde_json::json!({
"data": [{ "embedding": [0.1, 0.2, 0.3] }]
}))
}),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "test-key", "test-model", 3);
let result = p.embed(&["hello"]).await.unwrap();
assert_eq!(result.len(), 1);
assert_eq!(result[0], vec![0.1_f32, 0.2, 0.3]);
}
#[tokio::test]
async fn embed_success_batch() {
let app = Router::new().route(
"/v1/embeddings",
post(|| async {
Json(serde_json::json!({
"data": [
{ "embedding": [1.0, 2.0] },
{ "embedding": [3.0, 4.0] }
]
}))
}),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "k", "m", 2);
let result = p.embed(&["a", "b"]).await.unwrap();
assert_eq!(result.len(), 2);
assert_eq!(result[1], vec![3.0_f32, 4.0]);
}
#[tokio::test]
async fn embed_sends_auth_header() {
let app = Router::new().route(
"/v1/embeddings",
post(
|headers: HeaderMap, Json(body): Json<serde_json::Value>| async move {
let auth = headers.get("Authorization").unwrap().to_str().unwrap();
assert_eq!(auth, "Bearer my-secret-key");
assert_eq!(body["model"], "text-embedding-3-small");
Json(serde_json::json!({
"data": [{ "embedding": [1.0] }]
}))
},
),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "my-secret-key", "text-embedding-3-small", 1);
p.embed(&["test"]).await.unwrap();
}
#[tokio::test]
async fn embed_skips_auth_header_when_key_empty() {
let app = Router::new().route(
"/v1/embeddings",
post(|headers: HeaderMap| async move {
// No Authorization header should be present.
assert!(
headers.get("Authorization").is_none(),
"should not send auth header when key is empty"
);
Json(serde_json::json!({
"data": [{ "embedding": [1.0] }]
}))
}),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "", "m", 1);
p.embed(&["test"]).await.unwrap();
}
// ── embed — error paths ─────────────────────────────────
#[tokio::test]
async fn embed_server_error() {
let app = Router::new().route(
"/v1/embeddings",
post(|| async { (StatusCode::INTERNAL_SERVER_ERROR, "rate limited") }),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "k", "m", 1);
let err = p.embed(&["hi"]).await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("500"), "status: {msg}");
assert!(msg.contains("rate limited"), "body: {msg}");
}
#[tokio::test]
async fn embed_missing_data_field() {
let app = Router::new().route(
"/v1/embeddings",
post(|| async { Json(serde_json::json!({ "result": "ok" })) }),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "k", "m", 1);
let err = p.embed(&["hi"]).await.unwrap_err();
assert!(err.to_string().contains("missing 'data'"));
}
#[tokio::test]
async fn embed_missing_embedding_field_in_item() {
let app = Router::new().route(
"/v1/embeddings",
post(|| async {
Json(serde_json::json!({
"data": [{ "index": 0 }]
}))
}),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "k", "m", 1);
let err = p.embed(&["hi"]).await.unwrap_err();
assert!(err.to_string().contains("missing 'embedding'"));
}
#[tokio::test]
async fn embed_non_numeric_value_errors() {
let app = Router::new().route(
"/v1/embeddings",
post(|| async {
Json(serde_json::json!({
"data": [{ "embedding": [1.0, "not_a_number", 3.0] }]
}))
}),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "k", "m", 3);
let err = p.embed(&["hi"]).await.unwrap_err();
let msg = err.to_string();
assert!(msg.contains("non-numeric"), "msg: {msg}");
}
#[tokio::test]
async fn embed_count_mismatch() {
let app = Router::new().route(
"/v1/embeddings",
post(|| async {
Json(serde_json::json!({
"data": [{ "embedding": [1.0] }]
}))
}),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "k", "m", 1);
let err = p.embed(&["a", "b"]).await.unwrap_err();
assert!(err.to_string().contains("count mismatch"));
}
#[tokio::test]
async fn embed_dimension_mismatch() {
let app = Router::new().route(
"/v1/embeddings",
post(|| async {
Json(serde_json::json!({
"data": [{ "embedding": [1.0, 2.0, 3.0] }]
}))
}),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "k", "m", 2);
let err = p.embed(&["hi"]).await.unwrap_err();
assert!(err.to_string().contains("dimension mismatch"));
}
#[tokio::test]
async fn embed_malformed_json() {
let app = Router::new().route(
"/v1/embeddings",
post(|| async { (StatusCode::OK, "not json") }),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "k", "m", 1);
let err = p.embed(&["hi"]).await.unwrap_err();
assert!(err.is::<reqwest::Error>());
}
#[tokio::test]
async fn embed_connection_refused() {
let p = OpenAiEmbedding::new("http://127.0.0.1:1", "k", "m", 1);
let err = p.embed(&["hi"]).await.unwrap_err();
assert!(err.is::<reqwest::Error>());
}
// ── embed_one (trait default) ───────────────────────────
#[tokio::test]
async fn embed_one_success() {
let app = Router::new().route(
"/v1/embeddings",
post(|| async {
Json(serde_json::json!({
"data": [{ "embedding": [9.0, 8.0, 7.0] }]
}))
}),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&url, "k", "m", 3);
let vec = p.embed_one("test").await.unwrap();
assert_eq!(vec, vec![9.0_f32, 8.0, 7.0]);
}
// ── URL building — custom endpoint ──────────────────────
#[tokio::test]
async fn embed_with_explicit_api_path() {
let app = Router::new().route(
"/custom/api/embeddings",
post(|| async {
Json(serde_json::json!({
"data": [{ "embedding": [1.0] }]
}))
}),
);
let url = start_mock(app).await;
let p = OpenAiEmbedding::new(&format!("{url}/custom/api"), "k", "m", 1);
let result = p.embed(&["test"]).await.unwrap();
assert_eq!(result.len(), 1);
}
}
+884
View File
@@ -0,0 +1,884 @@
//! Local vector store backed by SQLite.
//!
//! Provides a self-contained vector database for storing, searching, and
//! managing text embeddings. Uses SQLite for persistence and brute-force
//! cosine similarity for retrieval (fast enough for on-device workloads up
//! to ~100K vectors).
//!
//! # Usage
//!
//! ```ignore
//! let embedder = Arc::new(OllamaEmbedding::default());
//! let store = VectorStore::open(db_path, embedder)?;
//!
//! store.insert("doc-1", "notes", "The quick brown fox", json!({})).await?;
//! let results = store.search("notes", "fast animal", 5).await?;
//! ```
use std::path::Path;
use std::sync::Arc;
use parking_lot::Mutex;
use rusqlite::Connection;
use super::EmbeddingProvider;
/// SQL to create the vector store schema.
const INIT_SQL: &str = "
PRAGMA journal_mode = WAL;
PRAGMA synchronous = NORMAL;
CREATE TABLE IF NOT EXISTS vectors (
id TEXT NOT NULL,
namespace TEXT NOT NULL,
text TEXT NOT NULL,
embedding BLOB NOT NULL,
metadata TEXT NOT NULL DEFAULT '{}',
created_at REAL NOT NULL,
updated_at REAL NOT NULL,
PRIMARY KEY (namespace, id)
);
CREATE INDEX IF NOT EXISTS idx_vectors_ns ON vectors(namespace);
CREATE TABLE IF NOT EXISTS store_meta (
key TEXT PRIMARY KEY,
value TEXT NOT NULL,
updated_at REAL NOT NULL
);
";
/// A single search result from the vector store.
#[derive(Debug, Clone)]
pub struct SearchResult {
/// The stored document ID.
pub id: String,
/// The namespace.
pub namespace: String,
/// The original text.
pub text: String,
/// Cosine similarity score (0.0 1.0).
pub score: f64,
/// Arbitrary JSON metadata attached at insert time.
pub metadata: serde_json::Value,
}
/// SQLite-backed local vector store.
///
/// Thread-safe: the inner connection is behind a `parking_lot::Mutex` and
/// the struct is `Send + Sync`. Embedding calls are async and run through
/// the configured [`EmbeddingProvider`].
pub struct VectorStore {
conn: Arc<Mutex<Connection>>,
embedder: Arc<dyn EmbeddingProvider>,
}
impl VectorStore {
/// Opens (or creates) a vector store at the given SQLite database path.
///
/// On first open the embedding provider name, model-name-hint, and
/// dimensions are persisted to a `store_meta` table. On subsequent opens
/// the stored dimensions are compared against the runtime embedder and an
/// error is returned if they mismatch (prevents silent cosine-similarity
/// corruption from mixed-dimension vectors).
pub fn open(db_path: &Path, embedder: Arc<dyn EmbeddingProvider>) -> anyhow::Result<Self> {
if let Some(parent) = db_path.parent() {
std::fs::create_dir_all(parent)?;
}
let conn = Connection::open(db_path)?;
conn.execute_batch(INIT_SQL)?;
Self::check_or_store_meta(&conn, &*embedder)?;
tracing::debug!(
target: "embeddings.store",
"[vector-store] opened at {}, embedder={}, dims={}",
db_path.display(),
embedder.name(),
embedder.dimensions()
);
Ok(Self {
conn: Arc::new(Mutex::new(conn)),
embedder,
})
}
/// Opens an in-memory vector store (useful for tests).
pub fn open_in_memory(embedder: Arc<dyn EmbeddingProvider>) -> anyhow::Result<Self> {
let conn = Connection::open_in_memory()?;
conn.execute_batch(INIT_SQL)?;
Self::check_or_store_meta(&conn, &*embedder)?;
Ok(Self {
conn: Arc::new(Mutex::new(conn)),
embedder,
})
}
/// Returns a reference to the embedding provider.
pub fn embedder(&self) -> &dyn EmbeddingProvider {
self.embedder.as_ref()
}
/// Persist or validate the embedding configuration in `store_meta`.
fn check_or_store_meta(
conn: &Connection,
embedder: &dyn EmbeddingProvider,
) -> anyhow::Result<()> {
let now = now_ts();
let stored_dims: Option<String> = conn
.query_row(
"SELECT value FROM store_meta WHERE key = 'embed_dims'",
[],
|row| row.get(0),
)
.ok();
match stored_dims {
None => {
// First open — persist metadata.
let stmts: &[(&str, &str)] = &[
("embed_provider", embedder.name()),
("embed_dims", &embedder.dimensions().to_string()),
];
for (key, value) in stmts {
conn.execute(
"INSERT OR REPLACE INTO store_meta (key, value, updated_at) VALUES (?1, ?2, ?3)",
rusqlite::params![key, value, now],
)?;
}
tracing::debug!(
target: "embeddings.store",
"[vector-store] stored meta: provider={}, dims={}",
embedder.name(),
embedder.dimensions()
);
}
Some(dims_str) => {
let stored: usize = dims_str.parse().unwrap_or(0);
let runtime = embedder.dimensions();
if stored != 0 && runtime != 0 && stored != runtime {
anyhow::bail!(
"vector store dimension mismatch: database was created with \
{stored}-dim embeddings but the current provider ({}) uses \
{runtime} dims. Delete the database or reconfigure the provider.",
embedder.name()
);
}
}
}
Ok(())
}
// ── Write operations ─────────────────────────────────────
/// Inserts or updates a text entry. The text is embedded automatically.
///
/// If an entry with the same `(namespace, id)` already exists it is replaced.
pub async fn insert(
&self,
id: &str,
namespace: &str,
text: &str,
metadata: serde_json::Value,
) -> anyhow::Result<()> {
tracing::trace!(
target: "embeddings.store",
"[vector-store] insert: id={id}, ns={namespace}, text_len={}",
text.len()
);
let embedding = self.embedder.embed_one(text).await?;
self.insert_with_vector(id, namespace, text, &embedding, metadata)
}
/// Inserts with a pre-computed embedding vector (skips the embed call).
pub fn insert_with_vector(
&self,
id: &str,
namespace: &str,
text: &str,
embedding: &[f32],
metadata: serde_json::Value,
) -> anyhow::Result<()> {
let blob = vec_to_bytes(embedding);
let meta_str = serde_json::to_string(&metadata)?;
let now = now_ts();
let conn = self.conn.lock();
conn.execute(
"INSERT OR REPLACE INTO vectors (id, namespace, text, embedding, metadata, created_at, updated_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)",
rusqlite::params![id, namespace, text, blob, meta_str, now, now],
)?;
tracing::trace!(
target: "embeddings.store",
"[vector-store] inserted id={id} ns={namespace} dims={}",
embedding.len()
);
Ok(())
}
/// Bulk-insert multiple entries. Each text is embedded automatically.
pub async fn insert_batch(
&self,
namespace: &str,
entries: &[(&str, &str, serde_json::Value)], // (id, text, metadata)
) -> anyhow::Result<()> {
if entries.is_empty() {
return Ok(());
}
tracing::debug!(
target: "embeddings.store",
"[vector-store] insert_batch: ns={namespace}, count={}",
entries.len()
);
let texts: Vec<&str> = entries.iter().map(|(_, text, _)| *text).collect();
let embeddings = self.embedder.embed(&texts).await?;
if embeddings.len() != entries.len() {
anyhow::bail!(
"embedding count mismatch: got {} embeddings for {} entries",
embeddings.len(),
entries.len()
);
}
let now = now_ts();
let conn = self.conn.lock();
let tx = conn.unchecked_transaction()?;
for ((id, text, metadata), embedding) in entries.iter().zip(embeddings.iter()) {
let blob = vec_to_bytes(embedding);
let meta_str = serde_json::to_string(metadata)?;
tx.execute(
"INSERT OR REPLACE INTO vectors (id, namespace, text, embedding, metadata, created_at, updated_at)
VALUES (?1, ?2, ?3, ?4, ?5, ?6, ?7)",
rusqlite::params![id, namespace, text, blob, meta_str, now, now],
)?;
}
tx.commit()?;
tracing::debug!(
target: "embeddings.store",
"[vector-store] batch inserted {} entries in ns={namespace}",
entries.len()
);
Ok(())
}
// ── Search ───────────────────────────────────────────────
/// Searches for the `limit` most similar entries to `query` within a namespace.
///
/// The query is embedded via the configured provider and compared against
/// all stored vectors using cosine similarity.
pub async fn search(
&self,
namespace: &str,
query: &str,
limit: usize,
) -> anyhow::Result<Vec<SearchResult>> {
tracing::trace!(
target: "embeddings.store",
"[vector-store] search: ns={namespace}, limit={limit}, query_len={}",
query.len()
);
let query_vec = self.embedder.embed_one(query).await?;
self.search_by_vector(namespace, &query_vec, limit)
}
/// Searches using a pre-computed query vector.
pub fn search_by_vector(
&self,
namespace: &str,
query_vec: &[f32],
limit: usize,
) -> anyhow::Result<Vec<SearchResult>> {
if limit == 0 {
tracing::trace!(
target: "embeddings.store",
"[vector-store] search_by_vector: limit=0, returning empty"
);
return Ok(Vec::new());
}
let conn = self.conn.lock();
let mut stmt = conn.prepare(
"SELECT id, namespace, text, embedding, metadata FROM vectors WHERE namespace = ?1",
)?;
let rows: Vec<(String, String, String, Vec<u8>, String)> = stmt
.query_map(rusqlite::params![namespace], |row| {
Ok((
row.get::<_, String>(0)?,
row.get::<_, String>(1)?,
row.get::<_, String>(2)?,
row.get::<_, Vec<u8>>(3)?,
row.get::<_, String>(4)?,
))
})?
.collect::<rusqlite::Result<Vec<_>>>()?;
let mut scored: Vec<SearchResult> = rows
.into_iter()
.map(|(id, ns, text, blob, meta_str)| {
let stored_vec = bytes_to_vec(&blob);
let score = cosine_similarity(query_vec, &stored_vec);
let metadata = serde_json::from_str(&meta_str).unwrap_or(serde_json::Value::Null);
SearchResult {
id,
namespace: ns,
text,
score,
metadata,
}
})
.collect();
// Sort descending by score.
scored.sort_by(|a, b| {
b.score
.partial_cmp(&a.score)
.unwrap_or(std::cmp::Ordering::Equal)
});
scored.truncate(limit);
tracing::trace!(
target: "embeddings.store",
"[vector-store] search_by_vector: ns={namespace}, scanned={}, returned={}",
scored.len() + scored.capacity() - scored.len(), // approximate total before truncate
scored.len()
);
Ok(scored)
}
// ── Delete / management ──────────────────────────────────
/// Deletes a single entry by ID within a namespace.
///
/// Returns `true` if a row was actually deleted.
pub fn delete(&self, namespace: &str, id: &str) -> anyhow::Result<bool> {
let conn = self.conn.lock();
let affected = conn.execute(
"DELETE FROM vectors WHERE namespace = ?1 AND id = ?2",
rusqlite::params![namespace, id],
)?;
tracing::trace!(
target: "embeddings.store",
"[vector-store] delete: ns={namespace}, id={id}, affected={affected}"
);
Ok(affected > 0)
}
/// Deletes all entries in a namespace.
///
/// Returns the number of deleted rows.
pub fn clear_namespace(&self, namespace: &str) -> anyhow::Result<usize> {
let conn = self.conn.lock();
let affected = conn.execute(
"DELETE FROM vectors WHERE namespace = ?1",
rusqlite::params![namespace],
)?;
tracing::debug!(
target: "embeddings.store",
"[vector-store] cleared namespace={namespace}, deleted={affected}"
);
Ok(affected)
}
/// Returns the number of entries in a namespace (or all if `None`).
pub fn count(&self, namespace: Option<&str>) -> anyhow::Result<usize> {
let conn = self.conn.lock();
let count: usize = match namespace {
Some(ns) => conn.query_row(
"SELECT COUNT(*) FROM vectors WHERE namespace = ?1",
rusqlite::params![ns],
|row| row.get(0),
)?,
None => conn.query_row("SELECT COUNT(*) FROM vectors", [], |row| row.get(0))?,
};
Ok(count)
}
/// Lists all distinct namespaces.
pub fn list_namespaces(&self) -> anyhow::Result<Vec<String>> {
let conn = self.conn.lock();
let mut stmt = conn.prepare("SELECT DISTINCT namespace FROM vectors ORDER BY namespace")?;
let namespaces: Vec<String> = stmt
.query_map([], |row| row.get(0))?
.collect::<rusqlite::Result<Vec<_>>>()?;
Ok(namespaces)
}
}
// ── Vector math utilities ────────────────────────────────────
/// Serializes a float vector to little-endian bytes for SQLite BLOB storage.
pub fn vec_to_bytes(v: &[f32]) -> Vec<u8> {
let mut bytes = Vec::with_capacity(v.len() * 4);
for &f in v {
bytes.extend_from_slice(&f.to_le_bytes());
}
bytes
}
/// Deserializes little-endian bytes back to a float vector.
pub fn bytes_to_vec(bytes: &[u8]) -> Vec<f32> {
bytes
.chunks_exact(4)
.map(|chunk| {
let arr: [u8; 4] = chunk.try_into().unwrap_or([0; 4]);
f32::from_le_bytes(arr)
})
.collect()
}
/// Computes cosine similarity between two vectors. Returns 0.0 for
/// mismatched lengths, empty vectors, or zero-magnitude vectors.
pub fn cosine_similarity(a: &[f32], b: &[f32]) -> f64 {
if a.len() != b.len() || a.is_empty() {
return 0.0;
}
let mut dot = 0.0_f64;
let mut norm_a = 0.0_f64;
let mut norm_b = 0.0_f64;
for (x, y) in a.iter().zip(b.iter()) {
let x = f64::from(*x);
let y = f64::from(*y);
dot += x * y;
norm_a += x * x;
norm_b += y * y;
}
let denom = norm_a.sqrt() * norm_b.sqrt();
if denom <= f64::EPSILON {
return 0.0;
}
(dot / denom).clamp(0.0, 1.0)
}
fn now_ts() -> f64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_secs_f64())
.unwrap_or(0.0)
}
// ── Tests ────────────────────────────────────────────────────
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
/// A test embedding provider that returns deterministic vectors.
struct FakeEmbedding {
dims: usize,
}
#[async_trait::async_trait]
impl EmbeddingProvider for FakeEmbedding {
fn name(&self) -> &str {
"fake"
}
fn dimensions(&self) -> usize {
self.dims
}
async fn embed(&self, texts: &[&str]) -> anyhow::Result<Vec<Vec<f32>>> {
Ok(texts.iter().map(|t| text_to_vec(t, self.dims)).collect())
}
}
fn text_to_vec(text: &str, dims: usize) -> Vec<f32> {
let mut vec = vec![0.0_f32; dims];
for (i, byte) in text.bytes().enumerate() {
vec[i % dims] += byte as f32 / 255.0;
}
let norm: f32 = vec.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
for x in &mut vec {
*x /= norm;
}
}
vec
}
struct MismatchEmbedding;
#[async_trait::async_trait]
impl EmbeddingProvider for MismatchEmbedding {
fn name(&self) -> &str {
"mismatch"
}
fn dimensions(&self) -> usize {
2
}
async fn embed(&self, _texts: &[&str]) -> anyhow::Result<Vec<Vec<f32>>> {
Ok(vec![vec![1.0, 0.0]])
}
}
fn fake_store(dims: usize) -> VectorStore {
VectorStore::open_in_memory(Arc::new(FakeEmbedding { dims })).unwrap()
}
// ── vec_to_bytes / bytes_to_vec ─────────────────────────
#[test]
fn roundtrip_vec_bytes() {
let original = vec![1.0_f32, -2.5, 3.14, 0.0, f32::MAX, f32::MIN];
let bytes = vec_to_bytes(&original);
assert_eq!(bytes.len(), original.len() * 4);
assert_eq!(original, bytes_to_vec(&bytes));
}
#[test]
fn empty_vec_roundtrip() {
assert!(bytes_to_vec(&vec_to_bytes(&[])).is_empty());
}
#[test]
fn bytes_to_vec_truncates_partial_bytes() {
assert_eq!(bytes_to_vec(&[0u8; 5]).len(), 1);
}
// ── cosine_similarity ───────────────────────────────────
#[test]
fn cosine_identical() {
let v = vec![1.0_f32, 2.0, 3.0];
assert!((cosine_similarity(&v, &v) - 1.0).abs() < 1e-6);
}
#[test]
fn cosine_orthogonal() {
assert!(cosine_similarity(&[1.0, 0.0], &[0.0, 1.0]).abs() < 1e-6);
}
#[test]
fn cosine_opposite() {
assert!(cosine_similarity(&[1.0, 0.0], &[-1.0, 0.0]).abs() < 1e-6);
}
#[test]
fn cosine_mismatched_lengths() {
assert_eq!(cosine_similarity(&[1.0, 2.0], &[1.0, 2.0, 3.0]), 0.0);
}
#[test]
fn cosine_empty() {
assert_eq!(cosine_similarity(&[], &[]), 0.0);
}
#[test]
fn cosine_zero_vector() {
assert_eq!(cosine_similarity(&[0.0, 0.0], &[1.0, 0.0]), 0.0);
}
#[test]
fn cosine_similar_high() {
assert!(cosine_similarity(&[1.0, 2.0, 3.0], &[1.1, 2.1, 3.1]) > 0.99);
}
// ── VectorStore: open / metadata ────────────────────────
#[test]
fn open_in_memory_succeeds() {
let store = fake_store(3);
assert_eq!(store.count(None).unwrap(), 0);
}
#[test]
fn open_on_disk() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("sub/dir/vectors.db");
let store = VectorStore::open(&db_path, Arc::new(FakeEmbedding { dims: 3 })).unwrap();
assert_eq!(store.count(None).unwrap(), 0);
assert!(db_path.exists());
}
#[test]
fn open_reopen_same_dims_succeeds() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("v.db");
VectorStore::open(&db_path, Arc::new(FakeEmbedding { dims: 4 })).unwrap();
// Reopen with same dims — should work.
VectorStore::open(&db_path, Arc::new(FakeEmbedding { dims: 4 })).unwrap();
}
#[test]
fn open_reopen_different_dims_errors() {
let dir = tempfile::tempdir().unwrap();
let db_path = dir.path().join("v.db");
VectorStore::open(&db_path, Arc::new(FakeEmbedding { dims: 4 })).unwrap();
let result = VectorStore::open(&db_path, Arc::new(FakeEmbedding { dims: 8 }));
let msg = result.err().expect("should be an error").to_string();
assert!(msg.contains("dimension mismatch"), "msg: {msg}");
assert!(msg.contains("4"), "should mention stored dims: {msg}");
assert!(msg.contains("8"), "should mention runtime dims: {msg}");
}
#[test]
fn embedder_accessor() {
let store = fake_store(3);
assert_eq!(store.embedder().name(), "fake");
assert_eq!(store.embedder().dimensions(), 3);
}
// ── insert + count ──────────────────────────────────────
#[tokio::test]
async fn insert_and_count() {
let store = fake_store(4);
store.insert("a", "ns1", "hello", json!({})).await.unwrap();
store.insert("b", "ns1", "world", json!({})).await.unwrap();
store.insert("c", "ns2", "other", json!({})).await.unwrap();
assert_eq!(store.count(Some("ns1")).unwrap(), 2);
assert_eq!(store.count(Some("ns2")).unwrap(), 1);
assert_eq!(store.count(None).unwrap(), 3);
}
#[tokio::test]
async fn insert_upsert_replaces() {
let store = fake_store(4);
store
.insert("a", "ns", "original", json!({"v": 1}))
.await
.unwrap();
store
.insert("a", "ns", "updated", json!({"v": 2}))
.await
.unwrap();
assert_eq!(store.count(Some("ns")).unwrap(), 1);
let results = store
.search_by_vector("ns", &text_to_vec("updated", 4), 10)
.unwrap();
assert_eq!(results[0].text, "updated");
assert_eq!(results[0].metadata["v"], 2);
}
#[test]
fn insert_with_vector_sync() {
let store = fake_store(3);
store
.insert_with_vector("id1", "ns", "text", &[1.0, 0.0, 0.0], json!({"k": "v"}))
.unwrap();
assert_eq!(store.count(Some("ns")).unwrap(), 1);
}
// ── insert_batch ────────────────────────────────────────
#[tokio::test]
async fn insert_batch_multiple() {
let store = fake_store(4);
let entries = vec![
("a", "alpha", json!({})),
("b", "beta", json!({})),
("c", "gamma", json!({})),
];
store.insert_batch("ns", &entries).await.unwrap();
assert_eq!(store.count(Some("ns")).unwrap(), 3);
}
#[tokio::test]
async fn insert_batch_empty() {
let store = fake_store(4);
store.insert_batch("ns", &[]).await.unwrap();
assert_eq!(store.count(None).unwrap(), 0);
}
#[tokio::test]
async fn insert_batch_mismatch_error() {
let store = VectorStore::open_in_memory(Arc::new(MismatchEmbedding)).unwrap();
let entries = vec![("a", "alpha", json!({})), ("b", "beta", json!({}))];
let err = store.insert_batch("ns", &entries).await.unwrap_err();
assert!(err.to_string().contains("mismatch"));
}
// ── search ──────────────────────────────────────────────
#[tokio::test]
async fn search_returns_ranked_results() {
let store = fake_store(8);
store
.insert("a", "ns", "the quick brown fox", json!({}))
.await
.unwrap();
store
.insert("b", "ns", "a lazy dog sleeps", json!({}))
.await
.unwrap();
store
.insert("c", "ns", "the quick brown fox jumps", json!({}))
.await
.unwrap();
let results = store.search("ns", "the quick brown fox", 2).await.unwrap();
assert_eq!(results.len(), 2);
assert!(results[0].score >= results[1].score);
}
#[tokio::test]
async fn search_respects_limit() {
let store = fake_store(4);
for i in 0..10 {
store
.insert(&format!("id-{i}"), "ns", &format!("text {i}"), json!({}))
.await
.unwrap();
}
assert_eq!(store.search("ns", "text", 3).await.unwrap().len(), 3);
}
#[tokio::test]
async fn search_empty_namespace() {
let store = fake_store(4);
assert!(store.search("empty", "query", 10).await.unwrap().is_empty());
}
#[tokio::test]
async fn search_namespace_isolation() {
let store = fake_store(4);
store.insert("a", "ns1", "hello", json!({})).await.unwrap();
store.insert("b", "ns2", "hello", json!({})).await.unwrap();
assert_eq!(store.search("ns1", "hello", 10).await.unwrap()[0].id, "a");
assert_eq!(store.search("ns2", "hello", 10).await.unwrap()[0].id, "b");
}
// ── search_by_vector ────────────────────────────────────
#[test]
fn search_by_vector_limit_zero() {
let store = fake_store(3);
store
.insert_with_vector("a", "ns", "t", &[1.0, 0.0, 0.0], json!({}))
.unwrap();
assert!(store
.search_by_vector("ns", &[1.0, 0.0, 0.0], 0)
.unwrap()
.is_empty());
}
#[test]
fn search_by_vector_scores_correct() {
let store = fake_store(3);
store
.insert_with_vector("x", "ns", "x", &[1.0, 0.0, 0.0], json!({}))
.unwrap();
store
.insert_with_vector("y", "ns", "y", &[0.0, 1.0, 0.0], json!({}))
.unwrap();
let results = store.search_by_vector("ns", &[1.0, 0.0, 0.0], 2).unwrap();
assert_eq!(results[0].id, "x");
assert!((results[0].score - 1.0).abs() < 1e-6);
assert!(results[1].score < 1e-6);
}
#[test]
fn search_by_vector_preserves_metadata() {
let store = fake_store(2);
store
.insert_with_vector("a", "ns", "t", &[1.0, 0.0], json!({"key": "value"}))
.unwrap();
assert_eq!(
store.search_by_vector("ns", &[1.0, 0.0], 1).unwrap()[0].metadata["key"],
"value"
);
}
#[test]
fn search_handles_invalid_metadata_json() {
let store = fake_store(2);
{
let conn = store.conn.lock();
conn.execute(
"INSERT INTO vectors (id, namespace, text, embedding, metadata, created_at, updated_at)
VALUES ('bad', 'ns', 'text', ?1, 'not-json', 0.0, 0.0)",
rusqlite::params![vec_to_bytes(&[1.0, 0.0])],
).unwrap();
}
let results = store.search_by_vector("ns", &[1.0, 0.0], 1).unwrap();
assert_eq!(results[0].id, "bad");
assert!(results[0].metadata.is_null());
}
// ── delete ──────────────────────────────────────────────
#[tokio::test]
async fn delete_existing() {
let store = fake_store(4);
store.insert("a", "ns", "text", json!({})).await.unwrap();
assert!(store.delete("ns", "a").unwrap());
assert_eq!(store.count(Some("ns")).unwrap(), 0);
}
#[test]
fn delete_nonexistent() {
assert!(!fake_store(3).delete("ns", "no-such-id").unwrap());
}
#[tokio::test]
async fn delete_wrong_namespace() {
let store = fake_store(4);
store.insert("a", "ns1", "text", json!({})).await.unwrap();
assert!(!store.delete("ns2", "a").unwrap());
assert_eq!(store.count(Some("ns1")).unwrap(), 1);
}
// ── clear_namespace ─────────────────────────────────────
#[tokio::test]
async fn clear_namespace_removes_all() {
let store = fake_store(4);
store.insert("a", "ns", "one", json!({})).await.unwrap();
store.insert("b", "ns", "two", json!({})).await.unwrap();
store
.insert("c", "other", "three", json!({}))
.await
.unwrap();
assert_eq!(store.clear_namespace("ns").unwrap(), 2);
assert_eq!(store.count(Some("ns")).unwrap(), 0);
assert_eq!(store.count(Some("other")).unwrap(), 1);
}
#[test]
fn clear_empty_namespace() {
assert_eq!(fake_store(3).clear_namespace("empty").unwrap(), 0);
}
// ── list_namespaces ─────────────────────────────────────
#[tokio::test]
async fn list_namespaces_empty() {
assert!(fake_store(3).list_namespaces().unwrap().is_empty());
}
#[tokio::test]
async fn list_namespaces_populated() {
let store = fake_store(4);
store.insert("a", "beta", "t", json!({})).await.unwrap();
store.insert("b", "alpha", "t", json!({})).await.unwrap();
store.insert("c", "beta", "t", json!({})).await.unwrap();
assert_eq!(store.list_namespaces().unwrap(), vec!["alpha", "beta"]);
}
// ── count ───────────────────────────────────────────────
#[test]
fn count_empty() {
let store = fake_store(3);
assert_eq!(store.count(None).unwrap(), 0);
assert_eq!(store.count(Some("ns")).unwrap(), 0);
}
}
+5 -594
View File
@@ -1,596 +1,7 @@
//! Embedding providers for the OpenHuman memory system.
//! Re-exports from the top-level `openhuman::embeddings` module.
//!
//! This module provides a unified interface for converting text into vector
//! embeddings. It supports multiple providers:
//! - **Fastembed**: Local, high-performance embeddings using ONNX runtime.
//! - **OpenAI**: Cloud-based embeddings via the OpenAI API or compatible endpoints.
//! - **Noop**: A fallback provider for keyword-only search.
//! The canonical embedding logic now lives in `src/openhuman/embeddings/`.
//! This file keeps the old `memory::embeddings::*` import paths working so
//! that existing call sites do not need to change immediately.
use async_trait::async_trait;
use parking_lot::Mutex;
use std::env;
use std::path::PathBuf;
use std::str::FromStr;
use std::sync::Arc;
/// Default model name for Fastembed.
pub const DEFAULT_FASTEMBED_MODEL: &str = "BGESmallENV15";
/// Default dimensions for the BGESmallENV15 model.
pub const DEFAULT_FASTEMBED_DIMENSIONS: usize = 384;
/// Interface for embedding providers that convert text into numerical vectors.
#[async_trait]
pub trait EmbeddingProvider: Send + Sync {
/// Returns the name of the provider (e.g., "fastembed", "openai").
fn name(&self) -> &str;
/// Returns the number of dimensions in the generated embeddings.
fn dimensions(&self) -> usize;
/// Generates embeddings for a batch of strings.
async fn embed(&self, texts: &[&str]) -> anyhow::Result<Vec<Vec<f32>>>;
/// Generates an embedding for a single string.
async fn embed_one(&self, text: &str) -> anyhow::Result<Vec<f32>> {
let mut results = self.embed(&[text]).await?;
results
.pop()
.ok_or_else(|| anyhow::anyhow!("Empty embedding result"))
}
}
// ── Noop provider (keyword-only fallback) ────────────────────
/// A "no-op" embedding provider used when semantic search is disabled.
/// Returns empty vectors.
pub struct NoopEmbedding;
#[async_trait]
impl EmbeddingProvider for NoopEmbedding {
fn name(&self) -> &str {
"none"
}
fn dimensions(&self) -> usize {
0
}
async fn embed(&self, _texts: &[&str]) -> anyhow::Result<Vec<Vec<f32>>> {
Ok(Vec::new())
}
}
/// Represents the initialization state of the local Fastembed model.
enum FastembedState {
/// Initial state before the model is loaded.
Uninitialized,
/// Model is loaded into memory and ready for inference.
Ready(Box<fastembed::TextEmbedding>),
/// An error occurred during model loading.
Failed(String),
}
/// Local embedding provider using the `fastembed-rs` library.
/// Executes in a dedicated blocking thread to avoid stalling the async runtime.
pub struct FastembedEmbedding {
model: String,
dims: usize,
state: Arc<Mutex<FastembedState>>,
}
impl FastembedEmbedding {
/// Creates a new Fastembed provider with the specified model and dimensions.
pub fn new(model: &str, dims: usize) -> Self {
Self {
model: if model.trim().is_empty() {
DEFAULT_FASTEMBED_MODEL.to_string()
} else {
model.trim().to_string()
},
dims: if dims == 0 {
DEFAULT_FASTEMBED_DIMENSIONS
} else {
dims
},
state: Arc::new(Mutex::new(FastembedState::Uninitialized)),
}
}
/// Maps a string model name to a `fastembed::EmbeddingModel` enum.
fn resolve_model(&self) -> fastembed::EmbeddingModel {
fastembed::EmbeddingModel::from_str(&self.model)
.unwrap_or(fastembed::EmbeddingModel::BGESmallENV15)
}
/// Internal helper to initialize the model on first use.
fn init_model(&self) -> anyhow::Result<fastembed::TextEmbedding> {
ensure_fastembed_ort_dylib_path();
fastembed::TextEmbedding::try_new(
fastembed::InitOptions::new(self.resolve_model()).with_show_download_progress(false),
)
.map_err(|e| anyhow::anyhow!("fastembed init failed for {}: {e}", self.model))
}
}
/// Configures the search path for the ONNX Runtime dynamic library.
///
/// This is critical for Fastembed to function across different platforms and
/// installation methods (e.g., local dev, bundled app). It checks several
/// locations in order of priority:
/// 1. `ORT_DYLIB_PATH` environment variable.
/// 2. `ORT_LIB_LOCATION` environment variable.
/// 3. OpenHuman-specific cache directories.
/// 4. Standard system library paths (Linux only).
fn ensure_fastembed_ort_dylib_path() {
if env::var_os("ORT_DYLIB_PATH").is_some() {
return;
}
// Check for explicit library location override.
if let Some(lib_path) = env::var_os("ORT_LIB_LOCATION") {
let candidate = PathBuf::from(lib_path);
if candidate.is_file() {
env::set_var("ORT_DYLIB_PATH", candidate);
return;
}
#[cfg(target_os = "windows")]
let runtime_lib = candidate.join("onnxruntime.dll");
#[cfg(target_os = "macos")]
let runtime_lib = candidate.join("libonnxruntime.dylib");
#[cfg(target_os = "linux")]
let runtime_lib = candidate.join("libonnxruntime.so");
if runtime_lib.exists() {
env::set_var("ORT_DYLIB_PATH", runtime_lib);
}
}
// Fallback to system-wide paths on Linux.
#[cfg(target_os = "linux")]
{
for candidate in [
"/usr/lib/x86_64-linux-gnu/libonnxruntime.so",
"/usr/local/lib/libonnxruntime.so",
"/usr/lib/libonnxruntime.so",
] {
let candidate = PathBuf::from(candidate);
if candidate.exists() {
env::set_var("ORT_DYLIB_PATH", candidate);
return;
}
}
}
}
#[async_trait]
impl EmbeddingProvider for FastembedEmbedding {
fn name(&self) -> &str {
"fastembed"
}
fn dimensions(&self) -> usize {
self.dims
}
/// Performs embedding using a blocking task to prevent executor starvation.
async fn embed(&self, texts: &[&str]) -> anyhow::Result<Vec<Vec<f32>>> {
if texts.is_empty() {
return Ok(Vec::new());
}
let items = texts
.iter()
.map(|text| (*text).to_string())
.collect::<Vec<_>>();
let state = Arc::clone(&self.state);
let provider = self.model.clone();
let join_result = tokio::task::spawn_blocking(move || -> anyhow::Result<Vec<Vec<f32>>> {
ensure_fastembed_ort_dylib_path();
let mut guard = state.lock();
// Lazy initialization of the model on the first request.
//
// `fastembed::TextEmbedding::try_new` reaches into the `ort`
// crate's global environment, which uses a `std::sync::Mutex`.
// If any previous caller panicked while that mutex was held
// (common when the ONNX Runtime dylib path is wrong or a
// background init failed), every subsequent call panics with
// `"Mutex poisoned"`. Without `catch_unwind`, that panic
// propagates out of this `spawn_blocking` closure, kills the
// tokio blocking worker, and surfaces as a process-level
// stack trace — even though the caller only wanted an error.
//
// We trap the panic here, flip our own state to `Failed`, and
// return a regular `anyhow::Error` so every later call short-
// circuits on the cached failure without touching `ort` again.
if matches!(*guard, FastembedState::Uninitialized) {
let provider_for_init = provider.clone();
let init_result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
fastembed::TextEmbedding::try_new(
fastembed::InitOptions::new(
fastembed::EmbeddingModel::from_str(&provider_for_init)
.unwrap_or(fastembed::EmbeddingModel::BGESmallENV15),
)
.with_show_download_progress(false),
)
}));
match init_result {
Ok(Ok(model)) => *guard = FastembedState::Ready(Box::new(model)),
Ok(Err(err)) => {
let message = format!("fastembed init failed for {provider}: {err}");
tracing::error!(target: "memory.embeddings", "[embeddings] {message}");
*guard = FastembedState::Failed(message);
}
Err(panic_payload) => {
let panic_msg = extract_panic_message(&panic_payload);
let message = format!(
"fastembed init panicked for {provider}: {panic_msg}\
the ONNX Runtime global environment is in a poisoned state. \
Check ORT_DYLIB_PATH / ORT_LIB_LOCATION and restart the \
process to retry."
);
tracing::error!(target: "memory.embeddings", "[embeddings] {message}");
*guard = FastembedState::Failed(message);
}
}
}
match &mut *guard {
FastembedState::Ready(model) => {
// Also guard the actual embed call — fastembed / ort
// can panic on certain inputs or runtime errors, and
// we want to surface those as regular errors too.
let embed_result =
std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| {
model.embed(items, None)
}));
match embed_result {
Ok(Ok(vectors)) => Ok(vectors),
Ok(Err(e)) => Err(anyhow::anyhow!("fastembed embed failed: {e}")),
Err(panic_payload) => {
let panic_msg = extract_panic_message(&panic_payload);
Err(anyhow::anyhow!("fastembed embed panicked: {panic_msg}"))
}
}
}
FastembedState::Failed(message) => Err(anyhow::anyhow!(message.clone())),
FastembedState::Uninitialized => {
Err(anyhow::anyhow!("fastembed provider did not initialize"))
}
}
})
.await;
join_result.map_err(|e| anyhow::anyhow!("fastembed task join failed: {e}"))?
}
}
/// Best-effort extraction of a readable message from a `catch_unwind` payload.
/// Panics produced by `panic!("...")` downcast to `&'static str` or `String`;
/// everything else falls back to a generic label.
fn extract_panic_message(panic: &Box<dyn std::any::Any + Send>) -> String {
if let Some(s) = panic.downcast_ref::<&'static str>() {
(*s).to_string()
} else if let Some(s) = panic.downcast_ref::<String>() {
s.clone()
} else {
"unknown panic payload".to_string()
}
}
// ── OpenAI-compatible embedding provider ─────────────────────
/// Embedding provider for OpenAI and compatible APIs (e.g., LocalAI, Ollama).
pub struct OpenAiEmbedding {
base_url: String,
api_key: String,
model: String,
dims: usize,
}
impl OpenAiEmbedding {
/// Creates a new OpenAI-style provider.
pub fn new(base_url: &str, api_key: &str, model: &str, dims: usize) -> Self {
Self {
base_url: base_url.trim_end_matches('/').to_string(),
api_key: api_key.to_string(),
model: model.to_string(),
dims,
}
}
/// Internal helper to build an HTTP client with proxy support.
fn http_client(&self) -> reqwest::Client {
crate::openhuman::config::build_runtime_proxy_client("memory.embeddings")
}
/// Checks if the base URL includes a specific path (e.g., /api/v1).
fn has_explicit_api_path(&self) -> bool {
let Ok(url) = reqwest::Url::parse(&self.base_url) else {
return false;
};
let path = url.path().trim_end_matches('/');
!path.is_empty() && path != "/"
}
/// Checks if the URL already ends with /embeddings.
fn has_embeddings_endpoint(&self) -> bool {
let Ok(url) = reqwest::Url::parse(&self.base_url) else {
return false;
};
url.path().trim_end_matches('/').ends_with("/embeddings")
}
/// Constructs the final URL for the embeddings endpoint.
fn embeddings_url(&self) -> String {
if self.has_embeddings_endpoint() {
return self.base_url.clone();
}
if self.has_explicit_api_path() {
format!("{}/embeddings", self.base_url)
} else {
format!("{}/v1/embeddings", self.base_url)
}
}
}
#[async_trait]
impl EmbeddingProvider for OpenAiEmbedding {
fn name(&self) -> &str {
"openai"
}
fn dimensions(&self) -> usize {
self.dims
}
/// Sends a POST request to the embedding API.
async fn embed(&self, texts: &[&str]) -> anyhow::Result<Vec<Vec<f32>>> {
if texts.is_empty() {
return Ok(Vec::new());
}
let body = serde_json::json!({
"model": self.model,
"input": texts,
});
let resp = self
.http_client()
.post(self.embeddings_url())
.header("Authorization", format!("Bearer {}", self.api_key))
.header("Content-Type", "application/json")
.json(&body)
.send()
.await?;
if !resp.status().is_success() {
let status = resp.status();
let text = resp.text().await.unwrap_or_default();
anyhow::bail!("Embedding API error {status}: {text}");
}
let json: serde_json::Value = resp.json().await?;
let data = json
.get("data")
.and_then(|d| d.as_array())
.ok_or_else(|| anyhow::anyhow!("Invalid embedding response: missing 'data'"))?;
let mut embeddings = Vec::with_capacity(data.len());
for item in data {
let embedding = item
.get("embedding")
.and_then(|e| e.as_array())
.ok_or_else(|| anyhow::anyhow!("Invalid embedding item"))?;
#[allow(clippy::cast_possible_truncation)]
let vec: Vec<f32> = embedding
.iter()
.filter_map(|v| v.as_f64().map(|f| f as f32))
.collect();
embeddings.push(vec);
}
Ok(embeddings)
}
}
// ── Factory ──────────────────────────────────────────────────
/// Creates an embedding provider based on the specified name and configuration.
///
/// Supports "fastembed", "openai", and "custom:<url>".
pub fn create_embedding_provider(
provider: &str,
api_key: Option<&str>,
model: &str,
dims: usize,
) -> Box<dyn EmbeddingProvider> {
match provider {
"fastembed" => Box::new(FastembedEmbedding::new(model, dims)),
"openai" => {
let key = api_key.unwrap_or("");
Box::new(OpenAiEmbedding::new(
"https://api.openai.com",
key,
model,
dims,
))
}
name if name.starts_with("custom:") => {
let base_url = name.strip_prefix("custom:").unwrap_or("");
let key = api_key.unwrap_or("");
Box::new(OpenAiEmbedding::new(base_url, key, model, dims))
}
_ => Box::new(NoopEmbedding),
}
}
/// Returns the default local embedding provider (Fastembed).
pub fn default_local_embedding_provider() -> Arc<dyn EmbeddingProvider> {
Arc::new(FastembedEmbedding::new(
DEFAULT_FASTEMBED_MODEL,
DEFAULT_FASTEMBED_DIMENSIONS,
))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn noop_name() {
let p = NoopEmbedding;
assert_eq!(p.name(), "none");
assert_eq!(p.dimensions(), 0);
}
#[tokio::test]
async fn noop_embed_returns_empty() {
let p = NoopEmbedding;
let result = p.embed(&["hello"]).await.unwrap();
assert!(result.is_empty());
}
#[test]
fn factory_none() {
let p = create_embedding_provider("none", None, "model", 1536);
assert_eq!(p.name(), "none");
}
#[test]
fn factory_openai() {
let p = create_embedding_provider("openai", Some("key"), "text-embedding-3-small", 1536);
assert_eq!(p.name(), "openai");
assert_eq!(p.dimensions(), 1536);
}
#[test]
fn factory_fastembed() {
let p = create_embedding_provider("fastembed", None, DEFAULT_FASTEMBED_MODEL, 384);
assert_eq!(p.name(), "fastembed");
assert_eq!(p.dimensions(), 384);
}
#[test]
fn factory_custom_url() {
let p = create_embedding_provider("custom:http://localhost:1234", None, "model", 768);
assert_eq!(p.name(), "openai"); // uses OpenAiEmbedding internally
assert_eq!(p.dimensions(), 768);
}
// ── Edge cases ───────────────────────────────────────────────
#[tokio::test]
async fn noop_embed_one_returns_error() {
let p = NoopEmbedding;
// embed returns empty vec → pop() returns None → error
let result = p.embed_one("hello").await;
assert!(result.is_err());
}
#[tokio::test]
async fn noop_embed_empty_batch() {
let p = NoopEmbedding;
let result = p.embed(&[]).await.unwrap();
assert!(result.is_empty());
}
#[tokio::test]
async fn noop_embed_multiple_texts() {
let p = NoopEmbedding;
let result = p.embed(&["a", "b", "c"]).await.unwrap();
assert!(result.is_empty());
}
#[test]
fn factory_empty_string_returns_noop() {
let p = create_embedding_provider("", None, "model", 1536);
assert_eq!(p.name(), "none");
}
#[test]
fn factory_unknown_provider_returns_noop() {
let p = create_embedding_provider("cohere", None, "model", 1536);
assert_eq!(p.name(), "none");
}
#[test]
fn default_local_provider_uses_fastembed_defaults() {
let p = default_local_embedding_provider();
assert_eq!(p.name(), "fastembed");
assert_eq!(p.dimensions(), DEFAULT_FASTEMBED_DIMENSIONS);
}
#[test]
fn factory_custom_empty_url() {
// "custom:" with no URL — should still construct without panic
let p = create_embedding_provider("custom:", None, "model", 768);
assert_eq!(p.name(), "openai");
}
#[test]
fn factory_openai_no_api_key() {
let p = create_embedding_provider("openai", None, "text-embedding-3-small", 1536);
assert_eq!(p.name(), "openai");
assert_eq!(p.dimensions(), 1536);
}
#[test]
fn openai_trailing_slash_stripped() {
let p = OpenAiEmbedding::new("https://api.openai.com/", "key", "model", 1536);
assert_eq!(p.base_url, "https://api.openai.com");
}
#[test]
fn openai_dimensions_custom() {
let p = OpenAiEmbedding::new("http://localhost", "k", "m", 384);
assert_eq!(p.dimensions(), 384);
}
#[test]
fn embeddings_url_standard_openai() {
let p = OpenAiEmbedding::new("https://api.openai.com", "key", "model", 1536);
assert_eq!(p.embeddings_url(), "https://api.openai.com/v1/embeddings");
}
#[test]
fn embeddings_url_base_with_v1_no_duplicate() {
let p = OpenAiEmbedding::new("https://api.example.com/v1", "key", "model", 1536);
assert_eq!(p.embeddings_url(), "https://api.example.com/v1/embeddings");
}
#[test]
fn embeddings_url_non_v1_api_path_uses_raw_suffix() {
let p = OpenAiEmbedding::new(
"https://api.example.com/api/coding/v3",
"key",
"model",
1536,
);
assert_eq!(
p.embeddings_url(),
"https://api.example.com/api/coding/v3/embeddings"
);
}
#[test]
fn embeddings_url_custom_full_endpoint() {
let p = OpenAiEmbedding::new(
"https://my-api.example.com/api/v2/embeddings",
"key",
"model",
1536,
);
assert_eq!(
p.embeddings_url(),
"https://my-api.example.com/api/v2/embeddings"
);
}
}
pub use crate::openhuman::embeddings::*;
+1 -1
View File
@@ -72,7 +72,7 @@ impl MemoryClient {
std::fs::create_dir_all(&workspace_dir)
.map_err(|e| format!("Create workspace dir {}: {e}", workspace_dir.display()))?;
// Initialize the default local embedding provider (e.g., FastEmbed).
// Initialize the default local embedding provider (Ollama).
let embedder: Arc<dyn EmbeddingProvider> = embeddings::default_local_embedding_provider();
// Create the underlying UnifiedMemory instance.
+1 -1
View File
@@ -77,7 +77,7 @@ pub fn create_memory_with_storage_and_routes(
api_key,
&config.embedding_model,
config.embedding_dimensions,
));
)?);
// 2. Instantiate UnifiedMemory which handles SQLite and vector storage.
let mem = UnifiedMemory::new(workspace_dir, embedder, config.sqlite_open_timeout_secs)?;
+1
View File
@@ -30,6 +30,7 @@ pub mod credentials;
pub mod cron;
pub mod dev_paths;
pub mod doctor;
pub mod embeddings;
pub mod encryption;
pub mod health;
pub mod heartbeat;