mirror of
https://github.com/outbackdingo/optimclaw.git
synced 2026-08-26 15:40:18 +00:00
refactor: replace LanceDB Database decorator with VectorStore composition
Instead of wrapping all ~80 Database trait methods in a 664-line decorator (lancedb_wrapper.rs), introduce a 4-method VectorStore trait that any vector backend can implement. Workspace composes FTS from the database with vector search from the external store via RRF fusion. - Add src/workspace/vector_store.rs with VectorStore trait - Rewrite lancedb_store.rs to implement VectorStore (not wrap Database) - Delete src/db/lancedb_wrapper.rs (664 lines removed) - Remove get_chunk_by_id from Database trait and all backends - Workspace gains with_vector_store() builder for optional composition - Fix LanceDB tests: bypass_vector_index() for brute-force search - Fix integration tests: use temp file DB (libSQL :memory: is per-connection) - Merge duplicate mod tests in config.rs Net: -724 lines. Adding a new vector backend requires 4 methods, not 80. Co-Authored-By: Claude Opus 4.6 <[email protected]>
This commit is contained in:
Generated
+2
@@ -3926,6 +3926,8 @@ dependencies = [
|
||||
"aes-gcm",
|
||||
"aho-corasick",
|
||||
"anyhow",
|
||||
"arrow-array",
|
||||
"arrow-schema",
|
||||
"async-trait",
|
||||
"axum 0.8.8",
|
||||
"base64 0.22.1",
|
||||
|
||||
+3
-1
@@ -110,6 +110,8 @@ pgvector = { version = "0.4", features = ["postgres"], optional = true }
|
||||
|
||||
# LanceDB vector store (optional alternative to pgvector/libsql for workspace search)
|
||||
lancedb = { version = "0.26", optional = true }
|
||||
arrow-array = { version = "57", optional = true }
|
||||
arrow-schema = { version = "57", optional = true }
|
||||
|
||||
# WASM sandbox for untrusted tool execution
|
||||
wasmtime = { version = "28", features = ["component-model"] }
|
||||
@@ -156,7 +158,7 @@ tempfile = "3"
|
||||
|
||||
[features]
|
||||
default = ["postgres", "libsql"]
|
||||
lancedb = ["dep:lancedb"]
|
||||
lancedb = ["dep:lancedb", "dep:arrow-array", "dep:arrow-schema"]
|
||||
postgres = [
|
||||
"dep:deadpool-postgres",
|
||||
"dep:tokio-postgres",
|
||||
|
||||
+50
-36
@@ -409,18 +409,19 @@ impl DatabaseConfig {
|
||||
});
|
||||
}
|
||||
|
||||
let vector_backend: VectorBackend =
|
||||
optional_env("VECTOR_BACKEND")?
|
||||
.and_then(|s| s.parse().ok())
|
||||
.unwrap_or_default();
|
||||
let vector_backend: VectorBackend = optional_env("VECTOR_BACKEND")?
|
||||
.and_then(|s| s.parse().ok())
|
||||
.unwrap_or_default();
|
||||
|
||||
let lancedb_path = optional_env("LANCEDB_PATH")?.map(PathBuf::from).or_else(|| {
|
||||
if vector_backend == VectorBackend::LanceDb {
|
||||
Some(default_lancedb_path())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
});
|
||||
let lancedb_path = optional_env("LANCEDB_PATH")?
|
||||
.map(PathBuf::from)
|
||||
.or_else(|| {
|
||||
if vector_backend == VectorBackend::LanceDb {
|
||||
Some(default_lancedb_path())
|
||||
} else {
|
||||
None
|
||||
}
|
||||
});
|
||||
|
||||
Ok(Self {
|
||||
backend,
|
||||
@@ -456,31 +457,6 @@ pub fn default_lancedb_path() -> PathBuf {
|
||||
.join("lancedb")
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::VectorBackend;
|
||||
|
||||
#[test]
|
||||
fn test_vector_backend_parse() {
|
||||
assert_eq!("builtin".parse::<VectorBackend>().unwrap(), VectorBackend::Builtin);
|
||||
assert_eq!("pgvector".parse::<VectorBackend>().unwrap(), VectorBackend::Builtin);
|
||||
assert_eq!("libsql".parse::<VectorBackend>().unwrap(), VectorBackend::Builtin);
|
||||
assert_eq!("".parse::<VectorBackend>().unwrap(), VectorBackend::Builtin);
|
||||
|
||||
assert_eq!("lancedb".parse::<VectorBackend>().unwrap(), VectorBackend::LanceDb);
|
||||
assert_eq!("lance".parse::<VectorBackend>().unwrap(), VectorBackend::LanceDb);
|
||||
assert_eq!("Lancedb".parse::<VectorBackend>().unwrap(), VectorBackend::LanceDb);
|
||||
|
||||
assert!("invalid".parse::<VectorBackend>().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_default_lancedb_path() {
|
||||
let path = super::default_lancedb_path();
|
||||
assert!(path.to_string_lossy().ends_with(".ironclaw/lancedb"));
|
||||
}
|
||||
}
|
||||
|
||||
/// Which LLM backend to use.
|
||||
///
|
||||
/// Defaults to `NearAi` to keep IronClaw close to the NEAR ecosystem.
|
||||
@@ -2022,4 +1998,42 @@ mod tests {
|
||||
std::env::remove_var("LLM_MODEL");
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_vector_backend_parse() {
|
||||
assert_eq!(
|
||||
"builtin".parse::<VectorBackend>().unwrap(),
|
||||
VectorBackend::Builtin
|
||||
);
|
||||
assert_eq!(
|
||||
"pgvector".parse::<VectorBackend>().unwrap(),
|
||||
VectorBackend::Builtin
|
||||
);
|
||||
assert_eq!(
|
||||
"libsql".parse::<VectorBackend>().unwrap(),
|
||||
VectorBackend::Builtin
|
||||
);
|
||||
assert_eq!("".parse::<VectorBackend>().unwrap(), VectorBackend::Builtin);
|
||||
|
||||
assert_eq!(
|
||||
"lancedb".parse::<VectorBackend>().unwrap(),
|
||||
VectorBackend::LanceDb
|
||||
);
|
||||
assert_eq!(
|
||||
"lance".parse::<VectorBackend>().unwrap(),
|
||||
VectorBackend::LanceDb
|
||||
);
|
||||
assert_eq!(
|
||||
"Lancedb".parse::<VectorBackend>().unwrap(),
|
||||
VectorBackend::LanceDb
|
||||
);
|
||||
|
||||
assert!("invalid".parse::<VectorBackend>().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_default_lancedb_path() {
|
||||
let path = super::default_lancedb_path();
|
||||
assert!(path.to_string_lossy().ends_with(".ironclaw/lancedb"));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,664 +0,0 @@
|
||||
//! Database wrapper that uses LanceDB for vector search.
|
||||
//!
|
||||
//! When `VECTOR_BACKEND=lancedb`, this wraps the main database (Postgres or libSQL)
|
||||
//! and delegates vector search to LanceDB while using the inner DB for FTS and all
|
||||
//! other operations.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::agent::BrokenTool;
|
||||
use crate::agent::routine::{Routine, RoutineRun, RunStatus};
|
||||
use crate::context::{ActionRecord, JobContext, JobState};
|
||||
use crate::db::Database;
|
||||
use crate::error::{DatabaseError, WorkspaceError};
|
||||
use crate::history::{
|
||||
ConversationMessage, ConversationSummary, JobEventRecord, LlmCallRecord, SandboxJobRecord,
|
||||
SandboxJobSummary, SettingRow,
|
||||
};
|
||||
use crate::workspace::{
|
||||
MemoryChunk, MemoryDocument, SearchConfig, SearchResult, WorkspaceEntry,
|
||||
};
|
||||
use crate::workspace::search::{RankedResult, reciprocal_rank_fusion};
|
||||
use crate::workspace::lancedb_store::VectorStore;
|
||||
|
||||
/// Wraps a Database with a LanceDB VectorStore for hybrid search.
|
||||
///
|
||||
/// Documents and chunks stay in the inner DB. Vector search uses LanceDB.
|
||||
pub struct DbWithLanceVectorStore {
|
||||
inner: Arc<dyn Database>,
|
||||
vector_store: Arc<dyn VectorStore>,
|
||||
}
|
||||
|
||||
impl DbWithLanceVectorStore {
|
||||
pub fn new(inner: Arc<dyn Database>, vector_store: Arc<dyn VectorStore>) -> Self {
|
||||
Self { inner, vector_store }
|
||||
}
|
||||
}
|
||||
|
||||
#[async_trait]
|
||||
impl Database for DbWithLanceVectorStore {
|
||||
async fn run_migrations(&self) -> Result<(), DatabaseError> {
|
||||
self.inner.run_migrations().await
|
||||
}
|
||||
|
||||
async fn create_conversation(
|
||||
&self,
|
||||
channel: &str,
|
||||
user_id: &str,
|
||||
thread_id: Option<&str>,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.create_conversation(channel, user_id, thread_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn touch_conversation(&self, id: Uuid) -> Result<(), DatabaseError> {
|
||||
self.inner.touch_conversation(id).await
|
||||
}
|
||||
|
||||
async fn add_conversation_message(
|
||||
&self,
|
||||
conversation_id: Uuid,
|
||||
role: &str,
|
||||
content: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.add_conversation_message(conversation_id, role, content)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn ensure_conversation(
|
||||
&self,
|
||||
id: Uuid,
|
||||
channel: &str,
|
||||
user_id: &str,
|
||||
thread_id: Option<&str>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.ensure_conversation(id, channel, user_id, thread_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_conversations_with_preview(
|
||||
&self,
|
||||
user_id: &str,
|
||||
channel: &str,
|
||||
limit: i64,
|
||||
) -> Result<Vec<ConversationSummary>, DatabaseError> {
|
||||
self.inner
|
||||
.list_conversations_with_preview(user_id, channel, limit)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_or_create_assistant_conversation(
|
||||
&self,
|
||||
user_id: &str,
|
||||
channel: &str,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.get_or_create_assistant_conversation(user_id, channel)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn create_conversation_with_metadata(
|
||||
&self,
|
||||
user_id: &str,
|
||||
channel: &str,
|
||||
metadata: &serde_json::Value,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.create_conversation_with_metadata(user_id, channel, metadata)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_conversation_messages_paginated(
|
||||
&self,
|
||||
conversation_id: Uuid,
|
||||
before: Option<chrono::DateTime<chrono::Utc>>,
|
||||
limit: i64,
|
||||
) -> Result<(Vec<ConversationMessage>, bool), DatabaseError> {
|
||||
self.inner
|
||||
.list_conversation_messages_paginated(conversation_id, before, limit)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn update_conversation_metadata_field(
|
||||
&self,
|
||||
id: Uuid,
|
||||
key: &str,
|
||||
value: &serde_json::Value,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.update_conversation_metadata_field(id, key, value)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_conversation_metadata(
|
||||
&self,
|
||||
id: Uuid,
|
||||
) -> Result<Option<serde_json::Value>, DatabaseError> {
|
||||
self.inner.get_conversation_metadata(id).await
|
||||
}
|
||||
|
||||
async fn list_conversation_messages(
|
||||
&self,
|
||||
conversation_id: Uuid,
|
||||
) -> Result<Vec<ConversationMessage>, DatabaseError> {
|
||||
self.inner.list_conversation_messages(conversation_id).await
|
||||
}
|
||||
|
||||
async fn conversation_belongs_to_user(
|
||||
&self,
|
||||
conversation_id: Uuid,
|
||||
user_id: &str,
|
||||
) -> Result<bool, DatabaseError> {
|
||||
self.inner
|
||||
.conversation_belongs_to_user(conversation_id, user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn save_job(&self, ctx: &JobContext) -> Result<(), DatabaseError> {
|
||||
self.inner.save_job(ctx).await
|
||||
}
|
||||
|
||||
async fn get_job(&self, id: Uuid) -> Result<Option<JobContext>, DatabaseError> {
|
||||
self.inner.get_job(id).await
|
||||
}
|
||||
|
||||
async fn update_job_status(
|
||||
&self,
|
||||
id: Uuid,
|
||||
status: crate::context::JobState,
|
||||
failure_reason: Option<&str>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner.update_job_status(id, status, failure_reason).await
|
||||
}
|
||||
|
||||
async fn mark_job_stuck(&self, id: Uuid) -> Result<(), DatabaseError> {
|
||||
self.inner.mark_job_stuck(id).await
|
||||
}
|
||||
|
||||
async fn get_stuck_jobs(&self) -> Result<Vec<Uuid>, DatabaseError> {
|
||||
self.inner.get_stuck_jobs().await
|
||||
}
|
||||
|
||||
async fn save_action(&self, job_id: Uuid, action: &ActionRecord) -> Result<(), DatabaseError> {
|
||||
self.inner.save_action(job_id, action).await
|
||||
}
|
||||
|
||||
async fn get_job_actions(&self, job_id: Uuid) -> Result<Vec<ActionRecord>, DatabaseError> {
|
||||
self.inner.get_job_actions(job_id).await
|
||||
}
|
||||
|
||||
async fn record_llm_call(&self, record: &LlmCallRecord<'_>) -> Result<Uuid, DatabaseError> {
|
||||
self.inner.record_llm_call(record).await
|
||||
}
|
||||
|
||||
async fn save_estimation_snapshot(
|
||||
&self,
|
||||
job_id: Uuid,
|
||||
category: &str,
|
||||
tool_names: &[String],
|
||||
estimated_cost: rust_decimal::Decimal,
|
||||
estimated_time_secs: i32,
|
||||
estimated_value: rust_decimal::Decimal,
|
||||
) -> Result<Uuid, DatabaseError> {
|
||||
self.inner
|
||||
.save_estimation_snapshot(
|
||||
job_id,
|
||||
category,
|
||||
tool_names,
|
||||
estimated_cost,
|
||||
estimated_time_secs,
|
||||
estimated_value,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn update_estimation_actuals(
|
||||
&self,
|
||||
id: Uuid,
|
||||
actual_cost: rust_decimal::Decimal,
|
||||
actual_time_secs: i32,
|
||||
actual_value: Option<rust_decimal::Decimal>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.update_estimation_actuals(id, actual_cost, actual_time_secs, actual_value)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn save_sandbox_job(&self, job: &SandboxJobRecord) -> Result<(), DatabaseError> {
|
||||
self.inner.save_sandbox_job(job).await
|
||||
}
|
||||
|
||||
async fn get_sandbox_job(&self, id: Uuid) -> Result<Option<SandboxJobRecord>, DatabaseError> {
|
||||
self.inner.get_sandbox_job(id).await
|
||||
}
|
||||
|
||||
async fn list_sandbox_jobs(&self) -> Result<Vec<SandboxJobRecord>, DatabaseError> {
|
||||
self.inner.list_sandbox_jobs().await
|
||||
}
|
||||
|
||||
async fn update_sandbox_job_status(
|
||||
&self,
|
||||
id: Uuid,
|
||||
status: &str,
|
||||
success: Option<bool>,
|
||||
message: Option<&str>,
|
||||
started_at: Option<chrono::DateTime<chrono::Utc>>,
|
||||
completed_at: Option<chrono::DateTime<chrono::Utc>>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.update_sandbox_job_status(id, status, success, message, started_at, completed_at)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn cleanup_stale_sandbox_jobs(&self) -> Result<u64, DatabaseError> {
|
||||
self.inner.cleanup_stale_sandbox_jobs().await
|
||||
}
|
||||
|
||||
async fn sandbox_job_summary(&self) -> Result<SandboxJobSummary, DatabaseError> {
|
||||
self.inner.sandbox_job_summary().await
|
||||
}
|
||||
|
||||
async fn list_sandbox_jobs_for_user(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> Result<Vec<SandboxJobRecord>, DatabaseError> {
|
||||
self.inner.list_sandbox_jobs_for_user(user_id).await
|
||||
}
|
||||
|
||||
async fn sandbox_job_summary_for_user(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> Result<SandboxJobSummary, DatabaseError> {
|
||||
self.inner.sandbox_job_summary_for_user(user_id).await
|
||||
}
|
||||
|
||||
async fn sandbox_job_belongs_to_user(
|
||||
&self,
|
||||
job_id: Uuid,
|
||||
user_id: &str,
|
||||
) -> Result<bool, DatabaseError> {
|
||||
self.inner
|
||||
.sandbox_job_belongs_to_user(job_id, user_id)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn update_sandbox_job_mode(&self, id: Uuid, mode: &str) -> Result<(), DatabaseError> {
|
||||
self.inner.update_sandbox_job_mode(id, mode).await
|
||||
}
|
||||
|
||||
async fn get_sandbox_job_mode(&self, id: Uuid) -> Result<Option<String>, DatabaseError> {
|
||||
self.inner.get_sandbox_job_mode(id).await
|
||||
}
|
||||
|
||||
async fn save_job_event(
|
||||
&self,
|
||||
job_id: Uuid,
|
||||
event_type: &str,
|
||||
data: &serde_json::Value,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner.save_job_event(job_id, event_type, data).await
|
||||
}
|
||||
|
||||
async fn list_job_events(
|
||||
&self,
|
||||
job_id: Uuid,
|
||||
limit: Option<i64>,
|
||||
) -> Result<Vec<crate::history::JobEventRecord>, DatabaseError> {
|
||||
self.inner.list_job_events(job_id, limit).await
|
||||
}
|
||||
|
||||
async fn create_routine(&self, routine: &crate::agent::routine::Routine) -> Result<(), DatabaseError> {
|
||||
self.inner.create_routine(routine).await
|
||||
}
|
||||
|
||||
async fn get_routine(&self, id: Uuid) -> Result<Option<crate::agent::routine::Routine>, DatabaseError> {
|
||||
self.inner.get_routine(id).await
|
||||
}
|
||||
|
||||
async fn get_routine_by_name(
|
||||
&self,
|
||||
user_id: &str,
|
||||
name: &str,
|
||||
) -> Result<Option<crate::agent::routine::Routine>, DatabaseError> {
|
||||
self.inner.get_routine_by_name(user_id, name).await
|
||||
}
|
||||
|
||||
async fn list_routines(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> Result<Vec<crate::agent::routine::Routine>, DatabaseError> {
|
||||
self.inner.list_routines(user_id).await
|
||||
}
|
||||
|
||||
async fn list_event_routines(&self) -> Result<Vec<crate::agent::routine::Routine>, DatabaseError> {
|
||||
self.inner.list_event_routines().await
|
||||
}
|
||||
|
||||
async fn list_due_cron_routines(&self) -> Result<Vec<crate::agent::routine::Routine>, DatabaseError> {
|
||||
self.inner.list_due_cron_routines().await
|
||||
}
|
||||
|
||||
async fn update_routine(&self, routine: &crate::agent::routine::Routine) -> Result<(), DatabaseError> {
|
||||
self.inner.update_routine(routine).await
|
||||
}
|
||||
|
||||
async fn update_routine_runtime(
|
||||
&self,
|
||||
id: Uuid,
|
||||
last_run_at: chrono::DateTime<chrono::Utc>,
|
||||
next_fire_at: Option<chrono::DateTime<chrono::Utc>>,
|
||||
run_count: u64,
|
||||
consecutive_failures: u32,
|
||||
state: &serde_json::Value,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.update_routine_runtime(
|
||||
id,
|
||||
last_run_at,
|
||||
next_fire_at,
|
||||
run_count,
|
||||
consecutive_failures,
|
||||
state,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn delete_routine(&self, id: Uuid) -> Result<bool, DatabaseError> {
|
||||
self.inner.delete_routine(id).await
|
||||
}
|
||||
|
||||
async fn create_routine_run(&self, run: &RoutineRun) -> Result<Uuid, DatabaseError> {
|
||||
self.inner.create_routine_run(run).await
|
||||
}
|
||||
|
||||
async fn complete_routine_run(
|
||||
&self,
|
||||
id: Uuid,
|
||||
status: RunStatus,
|
||||
result_summary: Option<&str>,
|
||||
tokens_used: Option<i32>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner
|
||||
.complete_routine_run(id, status, result_summary, tokens_used)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_routine_runs(
|
||||
&self,
|
||||
routine_id: Uuid,
|
||||
limit: i64,
|
||||
) -> Result<Vec<RoutineRun>, DatabaseError> {
|
||||
self.inner.list_routine_runs(routine_id, limit).await
|
||||
}
|
||||
|
||||
async fn count_running_routine_runs(&self, routine_id: Uuid) -> Result<i64, DatabaseError> {
|
||||
self.inner.count_running_routine_runs(routine_id).await
|
||||
}
|
||||
|
||||
async fn record_tool_failure(
|
||||
&self,
|
||||
tool_name: &str,
|
||||
error_message: &str,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner.record_tool_failure(tool_name, error_message).await
|
||||
}
|
||||
|
||||
async fn get_broken_tools(&self, threshold: i32) -> Result<Vec<BrokenTool>, DatabaseError> {
|
||||
self.inner.get_broken_tools(threshold).await
|
||||
}
|
||||
|
||||
async fn mark_tool_repaired(&self, tool_name: &str) -> Result<(), DatabaseError> {
|
||||
self.inner.mark_tool_repaired(tool_name).await
|
||||
}
|
||||
|
||||
async fn increment_repair_attempts(&self, tool_name: &str) -> Result<(), DatabaseError> {
|
||||
self.inner.increment_repair_attempts(tool_name).await
|
||||
}
|
||||
|
||||
async fn get_setting(
|
||||
&self,
|
||||
user_id: &str,
|
||||
key: &str,
|
||||
) -> Result<Option<serde_json::Value>, DatabaseError> {
|
||||
self.inner.get_setting(user_id, key).await
|
||||
}
|
||||
|
||||
async fn get_setting_full(
|
||||
&self,
|
||||
user_id: &str,
|
||||
key: &str,
|
||||
) -> Result<Option<SettingRow>, DatabaseError> {
|
||||
self.inner.get_setting_full(user_id, key).await
|
||||
}
|
||||
|
||||
async fn set_setting(
|
||||
&self,
|
||||
user_id: &str,
|
||||
key: &str,
|
||||
value: &serde_json::Value,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner.set_setting(user_id, key, value).await
|
||||
}
|
||||
|
||||
async fn delete_setting(&self, user_id: &str, key: &str) -> Result<bool, DatabaseError> {
|
||||
self.inner.delete_setting(user_id, key).await
|
||||
}
|
||||
|
||||
async fn list_settings(&self, user_id: &str) -> Result<Vec<SettingRow>, DatabaseError> {
|
||||
self.inner.list_settings(user_id).await
|
||||
}
|
||||
|
||||
async fn get_all_settings(
|
||||
&self,
|
||||
user_id: &str,
|
||||
) -> Result<std::collections::HashMap<String, serde_json::Value>, DatabaseError> {
|
||||
self.inner.get_all_settings(user_id).await
|
||||
}
|
||||
|
||||
async fn set_all_settings(
|
||||
&self,
|
||||
user_id: &str,
|
||||
settings: &std::collections::HashMap<String, serde_json::Value>,
|
||||
) -> Result<(), DatabaseError> {
|
||||
self.inner.set_all_settings(user_id, settings).await
|
||||
}
|
||||
|
||||
async fn has_settings(&self, user_id: &str) -> Result<bool, DatabaseError> {
|
||||
self.inner.has_settings(user_id).await
|
||||
}
|
||||
|
||||
async fn get_document_by_path(
|
||||
&self,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
path: &str,
|
||||
) -> Result<MemoryDocument, WorkspaceError> {
|
||||
self.inner
|
||||
.get_document_by_path(user_id, agent_id, path)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_document_by_id(&self, id: Uuid) -> Result<MemoryDocument, WorkspaceError> {
|
||||
self.inner.get_document_by_id(id).await
|
||||
}
|
||||
|
||||
async fn get_or_create_document_by_path(
|
||||
&self,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
path: &str,
|
||||
) -> Result<MemoryDocument, WorkspaceError> {
|
||||
self.inner
|
||||
.get_or_create_document_by_path(user_id, agent_id, path)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn update_document(&self, id: Uuid, content: &str) -> Result<(), WorkspaceError> {
|
||||
self.inner.update_document(id, content).await
|
||||
}
|
||||
|
||||
async fn delete_document_by_path(
|
||||
&self,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
path: &str,
|
||||
) -> Result<(), WorkspaceError> {
|
||||
self.inner
|
||||
.delete_document_by_path(user_id, agent_id, path)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_directory(
|
||||
&self,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
directory: &str,
|
||||
) -> Result<Vec<WorkspaceEntry>, WorkspaceError> {
|
||||
self.inner
|
||||
.list_directory(user_id, agent_id, directory)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn list_all_paths(
|
||||
&self,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
) -> Result<Vec<String>, WorkspaceError> {
|
||||
self.inner.list_all_paths(user_id, agent_id).await
|
||||
}
|
||||
|
||||
async fn list_documents(
|
||||
&self,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
) -> Result<Vec<MemoryDocument>, WorkspaceError> {
|
||||
self.inner.list_documents(user_id, agent_id).await
|
||||
}
|
||||
|
||||
async fn delete_chunks(&self, document_id: Uuid) -> Result<(), WorkspaceError> {
|
||||
self.inner.delete_chunks(document_id).await?;
|
||||
self.vector_store.delete_chunks(document_id).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn insert_chunk(
|
||||
&self,
|
||||
document_id: Uuid,
|
||||
chunk_index: i32,
|
||||
content: &str,
|
||||
embedding: Option<&[f32]>,
|
||||
) -> Result<Uuid, WorkspaceError> {
|
||||
let chunk_id = self
|
||||
.inner
|
||||
.insert_chunk(document_id, chunk_index, content, embedding)
|
||||
.await?;
|
||||
|
||||
if let Some(emb) = embedding {
|
||||
let doc = self.inner.get_document_by_id(document_id).await?;
|
||||
self.vector_store
|
||||
.insert_chunk(
|
||||
chunk_id,
|
||||
document_id,
|
||||
&doc.user_id,
|
||||
doc.agent_id,
|
||||
content,
|
||||
emb,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
|
||||
Ok(chunk_id)
|
||||
}
|
||||
|
||||
async fn update_chunk_embedding(
|
||||
&self,
|
||||
chunk_id: Uuid,
|
||||
embedding: &[f32],
|
||||
) -> Result<(), WorkspaceError> {
|
||||
self.inner.update_chunk_embedding(chunk_id, embedding).await?;
|
||||
|
||||
if let Some(chunk) = self.inner.get_chunk_by_id(chunk_id).await? {
|
||||
let doc = self.inner.get_document_by_id(chunk.document_id).await?;
|
||||
self.vector_store
|
||||
.update_chunk_embedding(
|
||||
chunk_id,
|
||||
chunk.document_id,
|
||||
&doc.user_id,
|
||||
doc.agent_id,
|
||||
&chunk.content,
|
||||
embedding,
|
||||
)
|
||||
.await?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn get_chunks_without_embeddings(
|
||||
&self,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
limit: usize,
|
||||
) -> Result<Vec<MemoryChunk>, WorkspaceError> {
|
||||
self.inner
|
||||
.get_chunks_without_embeddings(user_id, agent_id, limit)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_chunk_by_id(
|
||||
&self,
|
||||
chunk_id: Uuid,
|
||||
) -> Result<Option<MemoryChunk>, WorkspaceError> {
|
||||
self.inner.get_chunk_by_id(chunk_id).await
|
||||
}
|
||||
|
||||
async fn hybrid_search(
|
||||
&self,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
query: &str,
|
||||
embedding: Option<&[f32]>,
|
||||
config: &SearchConfig,
|
||||
) -> Result<Vec<SearchResult>, WorkspaceError> {
|
||||
let fts_config = SearchConfig {
|
||||
use_fts: true,
|
||||
use_vector: false,
|
||||
..*config
|
||||
};
|
||||
let fts_search_results = self
|
||||
.inner
|
||||
.hybrid_search(user_id, agent_id, query, None, &fts_config)
|
||||
.await?;
|
||||
|
||||
let fts_results: Vec<RankedResult> = fts_search_results
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(i, r)| RankedResult {
|
||||
chunk_id: r.chunk_id,
|
||||
document_id: r.document_id,
|
||||
content: r.content.clone(),
|
||||
rank: (i + 1) as u32,
|
||||
})
|
||||
.collect();
|
||||
|
||||
let vector_results = if config.use_vector {
|
||||
if let Some(emb) = embedding {
|
||||
self.vector_store
|
||||
.vector_search(user_id, agent_id, emb, config.pre_fusion_limit)
|
||||
.await?
|
||||
} else {
|
||||
Vec::new()
|
||||
}
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
|
||||
Ok(reciprocal_rank_fusion(fts_results, vector_results, config))
|
||||
}
|
||||
}
|
||||
@@ -2486,46 +2486,6 @@ impl Database for LibSqlBackend {
|
||||
Ok(chunks)
|
||||
}
|
||||
|
||||
async fn get_chunk_by_id(
|
||||
&self,
|
||||
chunk_id: Uuid,
|
||||
) -> Result<Option<MemoryChunk>, WorkspaceError> {
|
||||
let conn = self
|
||||
.connect()
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: e.to_string(),
|
||||
})?;
|
||||
let mut rows = conn
|
||||
.query(
|
||||
r#"
|
||||
SELECT id, document_id, chunk_index, content, created_at
|
||||
FROM memory_chunks WHERE id = ?1
|
||||
"#,
|
||||
params![chunk_id.to_string()],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!("Query failed: {}", e),
|
||||
})?;
|
||||
|
||||
let row = rows
|
||||
.next()
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!("Row fetch failed: {}", e),
|
||||
})?;
|
||||
|
||||
Ok(row.map(|r| MemoryChunk {
|
||||
id: get_text(&r, 0).parse().unwrap_or_default(),
|
||||
document_id: get_text(&r, 1).parse().unwrap_or_default(),
|
||||
chunk_index: get_i64(&r, 2) as i32,
|
||||
content: get_text(&r, 3),
|
||||
embedding: None,
|
||||
created_at: get_ts(&r, 4),
|
||||
}))
|
||||
}
|
||||
|
||||
// ==================== Workspace: Search ====================
|
||||
|
||||
async fn hybrid_search(
|
||||
|
||||
+6
-32
@@ -18,9 +18,6 @@ pub mod libsql_backend;
|
||||
#[cfg(feature = "libsql")]
|
||||
pub mod libsql_migrations;
|
||||
|
||||
#[cfg(all(feature = "lancedb", any(feature = "postgres", feature = "libsql")))]
|
||||
pub mod lancedb_wrapper;
|
||||
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -51,7 +48,7 @@ use crate::workspace::{SearchConfig, SearchResult};
|
||||
pub async fn connect_from_config(
|
||||
config: &crate::config::DatabaseConfig,
|
||||
) -> Result<Arc<dyn Database>, DatabaseError> {
|
||||
let inner: Arc<dyn Database> = match config.backend {
|
||||
match config.backend {
|
||||
#[cfg(feature = "libsql")]
|
||||
crate::config::DatabaseBackend::LibSql => {
|
||||
use secrecy::ExposeSecret as _;
|
||||
@@ -78,7 +75,7 @@ pub async fn connect_from_config(
|
||||
.map_err(|e| DatabaseError::Pool(e.to_string()))?
|
||||
};
|
||||
backend.run_migrations().await?;
|
||||
Arc::new(backend)
|
||||
Ok(Arc::new(backend))
|
||||
}
|
||||
#[cfg(feature = "postgres")]
|
||||
_ => {
|
||||
@@ -86,33 +83,13 @@ pub async fn connect_from_config(
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Pool(e.to_string()))?;
|
||||
pg.run_migrations().await?;
|
||||
Arc::new(pg)
|
||||
Ok(Arc::new(pg))
|
||||
}
|
||||
#[cfg(not(feature = "postgres"))]
|
||||
_ => {
|
||||
return Err(DatabaseError::Pool(
|
||||
"No database backend available. Enable 'postgres' or 'libsql' feature."
|
||||
.to_string(),
|
||||
));
|
||||
}
|
||||
};
|
||||
|
||||
#[cfg(feature = "lancedb")]
|
||||
if config.vector_backend == crate::config::VectorBackend::LanceDb {
|
||||
let path = config
|
||||
.lancedb_path
|
||||
.clone()
|
||||
.unwrap_or_else(crate::config::default_lancedb_path);
|
||||
let store = crate::workspace::LanceDbVectorStore::new(path)
|
||||
.await
|
||||
.map_err(|e| DatabaseError::Pool(format!("LanceDB: {}", e)))?;
|
||||
return Ok(Arc::new(lancedb_wrapper::DbWithLanceVectorStore::new(
|
||||
inner,
|
||||
Arc::new(store),
|
||||
)) as Arc<dyn Database>);
|
||||
_ => Err(DatabaseError::Pool(
|
||||
"No database backend available. Enable 'postgres' or 'libsql' feature.".to_string(),
|
||||
)),
|
||||
}
|
||||
|
||||
Ok(inner)
|
||||
}
|
||||
|
||||
/// Backend-agnostic database trait.
|
||||
@@ -551,9 +528,6 @@ pub trait Database: Send + Sync {
|
||||
limit: usize,
|
||||
) -> Result<Vec<MemoryChunk>, WorkspaceError>;
|
||||
|
||||
/// Get a chunk by ID (for LanceDB update flow).
|
||||
async fn get_chunk_by_id(&self, chunk_id: Uuid) -> Result<Option<MemoryChunk>, WorkspaceError>;
|
||||
|
||||
// ==================== Workspace: Search ====================
|
||||
|
||||
/// Perform hybrid search combining FTS and vector similarity.
|
||||
|
||||
@@ -614,13 +614,6 @@ impl Database for PgBackend {
|
||||
.await
|
||||
}
|
||||
|
||||
async fn get_chunk_by_id(
|
||||
&self,
|
||||
chunk_id: Uuid,
|
||||
) -> Result<Option<MemoryChunk>, WorkspaceError> {
|
||||
self.repo.get_chunk_by_id(chunk_id).await
|
||||
}
|
||||
|
||||
// ==================== Workspace: Search ====================
|
||||
|
||||
async fn hybrid_search(
|
||||
|
||||
+12
-12
@@ -378,7 +378,7 @@ async fn main() -> anyhow::Result<()> {
|
||||
#[cfg(feature = "libsql")]
|
||||
let mut libsql_db: Option<std::sync::Arc<libsql::Database>> = None;
|
||||
|
||||
let mut inner_db: Option<Arc<dyn ironclaw::db::Database>> = if cli.no_db {
|
||||
let db: Option<Arc<dyn ironclaw::db::Database>> = if cli.no_db {
|
||||
tracing::warn!("Running without database connection");
|
||||
None
|
||||
} else {
|
||||
@@ -435,8 +435,8 @@ async fn main() -> anyhow::Result<()> {
|
||||
}
|
||||
};
|
||||
|
||||
// Optionally wrap with LanceDB vector store for workspace semantic search
|
||||
let db: Option<Arc<dyn ironclaw::db::Database>> = if let Some(inner) = inner_db.take() {
|
||||
// Create optional external vector store for workspace semantic search
|
||||
let vector_store: Option<Arc<dyn ironclaw::workspace::VectorStore>> = {
|
||||
#[cfg(feature = "lancedb")]
|
||||
{
|
||||
if config.database.vector_backend == ironclaw::config::VectorBackend::LanceDb {
|
||||
@@ -449,12 +449,9 @@ async fn main() -> anyhow::Result<()> {
|
||||
.await
|
||||
.map_err(|e| anyhow::anyhow!("LanceDB: {}", e))?;
|
||||
tracing::info!("LanceDB vector store connected for workspace search");
|
||||
Some(Arc::new(ironclaw::db::lancedb_wrapper::DbWithLanceVectorStore::new(
|
||||
inner,
|
||||
Arc::new(store),
|
||||
)) as Arc<dyn ironclaw::db::Database>)
|
||||
Some(Arc::new(store) as Arc<dyn ironclaw::workspace::VectorStore>)
|
||||
} else {
|
||||
Some(inner)
|
||||
None
|
||||
}
|
||||
}
|
||||
#[cfg(not(feature = "lancedb"))]
|
||||
@@ -463,12 +460,9 @@ async fn main() -> anyhow::Result<()> {
|
||||
anyhow::bail!(
|
||||
"VECTOR_BACKEND=lancedb requires the 'lancedb' feature. Build with: cargo build --features lancedb"
|
||||
);
|
||||
} else {
|
||||
Some(inner)
|
||||
}
|
||||
None
|
||||
}
|
||||
} else {
|
||||
None
|
||||
};
|
||||
|
||||
// Post-init operations using the database
|
||||
@@ -761,6 +755,9 @@ async fn main() -> anyhow::Result<()> {
|
||||
if let Some(ref emb) = embeddings {
|
||||
workspace = workspace.with_embeddings(emb.clone());
|
||||
}
|
||||
if let Some(ref vs) = vector_store {
|
||||
workspace = workspace.with_vector_store(vs.clone());
|
||||
}
|
||||
let workspace = Arc::new(workspace);
|
||||
tools.register_memory_tools(workspace);
|
||||
}
|
||||
@@ -1265,6 +1262,9 @@ async fn main() -> anyhow::Result<()> {
|
||||
if let Some(ref emb) = embeddings {
|
||||
ws = ws.with_embeddings(emb.clone());
|
||||
}
|
||||
if let Some(ref vs) = vector_store {
|
||||
ws = ws.with_vector_store(vs.clone());
|
||||
}
|
||||
Some(Arc::new(ws))
|
||||
} else {
|
||||
None
|
||||
|
||||
+110
-171
@@ -8,77 +8,26 @@
|
||||
//! LANCEDB_PATH=~/.ironclaw/lancedb # Default
|
||||
//! VECTOR_BACKEND=lancedb # Use LanceDB for vector search
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use async_trait::async_trait;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::error::WorkspaceError;
|
||||
use crate::workspace::search::RankedResult;
|
||||
|
||||
/// Embedding dimension (text-embedding-3-small default).
|
||||
/// Must match the embedding model used.
|
||||
pub const DEFAULT_EMBEDDING_DIM: i32 = 1536;
|
||||
|
||||
/// Vector store abstraction for semantic search.
|
||||
///
|
||||
/// Implementations: pgvector/libsql (embedded in Database), LanceDB (this module).
|
||||
#[async_trait]
|
||||
pub trait VectorStore: Send + Sync {
|
||||
/// Insert a chunk with its embedding.
|
||||
async fn insert_chunk(
|
||||
&self,
|
||||
chunk_id: Uuid,
|
||||
document_id: Uuid,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
content: &str,
|
||||
embedding: &[f32],
|
||||
) -> Result<(), WorkspaceError>;
|
||||
|
||||
/// Update an existing chunk's embedding.
|
||||
///
|
||||
/// For LanceDB, this performs delete+insert since LanceDB has limited update
|
||||
/// support. Caller must provide full chunk metadata for the re-insert.
|
||||
async fn update_chunk_embedding(
|
||||
&self,
|
||||
chunk_id: Uuid,
|
||||
document_id: Uuid,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
content: &str,
|
||||
embedding: &[f32],
|
||||
) -> Result<(), WorkspaceError>;
|
||||
|
||||
/// Delete all chunks for a document.
|
||||
async fn delete_chunks(&self, document_id: Uuid) -> Result<(), WorkspaceError>;
|
||||
|
||||
/// Vector similarity search, filtered by user and agent.
|
||||
async fn vector_search(
|
||||
&self,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
embedding: &[f32],
|
||||
limit: usize,
|
||||
) -> Result<Vec<RankedResult>, WorkspaceError>;
|
||||
}
|
||||
|
||||
#[cfg(feature = "lancedb")]
|
||||
mod impl_lancedb {
|
||||
use std::sync::Arc;
|
||||
|
||||
use arrow_array::types::Float32Type;
|
||||
use arrow_array::{
|
||||
Array, FixedSizeListArray, RecordBatch, RecordBatchIterator, StringArray,
|
||||
};
|
||||
use arrow_array::{Array, FixedSizeListArray, RecordBatch, RecordBatchIterator, StringArray};
|
||||
use arrow_schema::{DataType, Field, Schema};
|
||||
use async_trait::async_trait;
|
||||
use futures::StreamExt;
|
||||
use lancedb::index::Index;
|
||||
use lancedb::query::{ExecutableQuery, QueryBase};
|
||||
use uuid::Uuid;
|
||||
|
||||
use super::{RankedResult, VectorStore, DEFAULT_EMBEDDING_DIM};
|
||||
use super::DEFAULT_EMBEDDING_DIM;
|
||||
use crate::error::WorkspaceError;
|
||||
use crate::workspace::search::RankedResult;
|
||||
use crate::workspace::vector_store::VectorStore;
|
||||
|
||||
const TABLE_NAME: &str = "memory_chunks";
|
||||
|
||||
@@ -105,12 +54,11 @@ mod impl_lancedb {
|
||||
reason: "Invalid LanceDB path".to_string(),
|
||||
})?;
|
||||
|
||||
let db = lancedb::connect(path_str)
|
||||
.execute()
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
let db = lancedb::connect(path_str).execute().await.map_err(|e| {
|
||||
WorkspaceError::SearchFailed {
|
||||
reason: format!("Failed to connect to LanceDB: {}", e),
|
||||
})?;
|
||||
}
|
||||
})?;
|
||||
|
||||
let store = Self {
|
||||
db: Arc::new(db),
|
||||
@@ -142,20 +90,9 @@ mod impl_lancedb {
|
||||
reason: format!("Failed to create table: {}", e),
|
||||
})?;
|
||||
|
||||
let table = self.db.open_table(&self.table_name).execute().await.map_err(|e| {
|
||||
WorkspaceError::SearchFailed {
|
||||
reason: format!("Failed to open table: {}", e),
|
||||
}
|
||||
})?;
|
||||
|
||||
table
|
||||
.create_index(&["vector"], Index::Auto)
|
||||
.execute()
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!("Failed to create vector index: {}", e),
|
||||
})?;
|
||||
|
||||
// Index creation is deferred — brute-force search via
|
||||
// bypass_vector_index() works without a pre-built index and is
|
||||
// sufficient for personal workspace sizes.
|
||||
Ok(())
|
||||
}
|
||||
|
||||
@@ -180,7 +117,7 @@ mod impl_lancedb {
|
||||
|
||||
#[async_trait]
|
||||
impl VectorStore for LanceDbVectorStore {
|
||||
async fn insert_chunk(
|
||||
async fn store_embedding(
|
||||
&self,
|
||||
chunk_id: Uuid,
|
||||
document_id: Uuid,
|
||||
@@ -199,21 +136,23 @@ mod impl_lancedb {
|
||||
});
|
||||
}
|
||||
|
||||
let table = self.db.open_table(&self.table_name).execute().await.map_err(|e| {
|
||||
WorkspaceError::SearchFailed {
|
||||
let table = self
|
||||
.db
|
||||
.open_table(&self.table_name)
|
||||
.execute()
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!("Failed to open table: {}", e),
|
||||
}
|
||||
})?;
|
||||
})?;
|
||||
|
||||
let chunk_ids = StringArray::from(vec![chunk_id.to_string()]);
|
||||
let document_ids = StringArray::from(vec![document_id.to_string()]);
|
||||
let user_ids = StringArray::from(vec![user_id]);
|
||||
let agent_ids = StringArray::from(vec![agent_id.map(|a| a.to_string())]);
|
||||
let contents = StringArray::from(vec![content]);
|
||||
let vec_values: Vec<Option<f32>> =
|
||||
embedding.iter().map(|&x| Some(x)).collect();
|
||||
let vec_values: Vec<Option<f32>> = embedding.iter().map(|&x| Some(x)).collect();
|
||||
let vectors = FixedSizeListArray::from_iter_primitive::<Float32Type, _, _>(
|
||||
vec![Some(vec_values)].into_iter(),
|
||||
vec![Some(vec_values)],
|
||||
self.embedding_dim,
|
||||
);
|
||||
|
||||
@@ -232,13 +171,11 @@ mod impl_lancedb {
|
||||
reason: format!("Failed to create record batch: {}", e),
|
||||
})?;
|
||||
|
||||
let batches = RecordBatchIterator::new(
|
||||
vec![Ok(batch)].into_iter(),
|
||||
Arc::new(self.schema()),
|
||||
);
|
||||
let batches =
|
||||
RecordBatchIterator::new(vec![Ok(batch)].into_iter(), Arc::new(self.schema()));
|
||||
|
||||
table
|
||||
.add(Box::new(batches))
|
||||
.add(Box::new(batches) as Box<dyn arrow_array::RecordBatchReader + Send>)
|
||||
.execute()
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::ChunkingFailed {
|
||||
@@ -248,7 +185,7 @@ mod impl_lancedb {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn update_chunk_embedding(
|
||||
async fn update_embedding(
|
||||
&self,
|
||||
chunk_id: Uuid,
|
||||
document_id: Uuid,
|
||||
@@ -257,11 +194,14 @@ mod impl_lancedb {
|
||||
content: &str,
|
||||
embedding: &[f32],
|
||||
) -> Result<(), WorkspaceError> {
|
||||
let table = self.db.open_table(&self.table_name).execute().await.map_err(|e| {
|
||||
WorkspaceError::SearchFailed {
|
||||
let table = self
|
||||
.db
|
||||
.open_table(&self.table_name)
|
||||
.execute()
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!("Failed to open table: {}", e),
|
||||
}
|
||||
})?;
|
||||
})?;
|
||||
|
||||
table
|
||||
.delete(&format!(
|
||||
@@ -273,16 +213,19 @@ mod impl_lancedb {
|
||||
reason: format!("Failed to delete chunk for update: {}", e),
|
||||
})?;
|
||||
|
||||
self.insert_chunk(chunk_id, document_id, user_id, agent_id, content, embedding)
|
||||
self.store_embedding(chunk_id, document_id, user_id, agent_id, content, embedding)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn delete_chunks(&self, document_id: Uuid) -> Result<(), WorkspaceError> {
|
||||
let table = self.db.open_table(&self.table_name).execute().await.map_err(|e| {
|
||||
WorkspaceError::SearchFailed {
|
||||
async fn delete_embeddings(&self, document_id: Uuid) -> Result<(), WorkspaceError> {
|
||||
let table = self
|
||||
.db
|
||||
.open_table(&self.table_name)
|
||||
.execute()
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!("Failed to open table: {}", e),
|
||||
}
|
||||
})?;
|
||||
})?;
|
||||
|
||||
table
|
||||
.delete(&format!(
|
||||
@@ -304,11 +247,14 @@ mod impl_lancedb {
|
||||
embedding: &[f32],
|
||||
limit: usize,
|
||||
) -> Result<Vec<RankedResult>, WorkspaceError> {
|
||||
let table = self.db.open_table(&self.table_name).execute().await.map_err(|e| {
|
||||
WorkspaceError::SearchFailed {
|
||||
let table = self
|
||||
.db
|
||||
.open_table(&self.table_name)
|
||||
.execute()
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!("Failed to open table: {}", e),
|
||||
}
|
||||
})?;
|
||||
})?;
|
||||
|
||||
let filter = if let Some(aid) = agent_id {
|
||||
format!(
|
||||
@@ -323,79 +269,81 @@ mod impl_lancedb {
|
||||
)
|
||||
};
|
||||
|
||||
let mut stream = table
|
||||
let query = table
|
||||
.query()
|
||||
.nearest_to(embedding)
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!("Invalid query vector: {}", e),
|
||||
})?
|
||||
.only_if(&filter)
|
||||
.limit(limit as u32)
|
||||
.execute()
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
.bypass_vector_index()
|
||||
.limit(limit);
|
||||
let mut stream = ExecutableQuery::execute(&query).await.map_err(|e| {
|
||||
WorkspaceError::SearchFailed {
|
||||
reason: format!("Vector search failed: {}", e),
|
||||
})?;
|
||||
}
|
||||
})?;
|
||||
|
||||
let mut results = Vec::new();
|
||||
let mut rank: u32 = 1;
|
||||
while let Some(batch) = stream.next().await {
|
||||
let batch = batch.map_err(|e| WorkspaceError::SearchFailed {
|
||||
while let Some(batch_result) = stream.next().await {
|
||||
let batch = batch_result.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!("Stream error: {}", e),
|
||||
})?;
|
||||
|
||||
let chunk_id_col = batch
|
||||
.column_by_name("chunk_id")
|
||||
.ok_or_else(|| WorkspaceError::SearchFailed {
|
||||
let chunk_id_col = batch.column_by_name("chunk_id").ok_or_else(|| {
|
||||
WorkspaceError::SearchFailed {
|
||||
reason: "chunk_id column missing".to_string(),
|
||||
})?;
|
||||
let document_id_col = batch
|
||||
.column_by_name("document_id")
|
||||
.ok_or_else(|| WorkspaceError::SearchFailed {
|
||||
}
|
||||
})?;
|
||||
let document_id_col = batch.column_by_name("document_id").ok_or_else(|| {
|
||||
WorkspaceError::SearchFailed {
|
||||
reason: "document_id column missing".to_string(),
|
||||
})?;
|
||||
let content_col = batch
|
||||
.column_by_name("content")
|
||||
.ok_or_else(|| WorkspaceError::SearchFailed {
|
||||
}
|
||||
})?;
|
||||
let content_col = batch.column_by_name("content").ok_or_else(|| {
|
||||
WorkspaceError::SearchFailed {
|
||||
reason: "content column missing".to_string(),
|
||||
})?;
|
||||
}
|
||||
})?;
|
||||
|
||||
let chunk_ids = chunk_id_col.as_any().downcast_ref::<StringArray>().ok_or_else(|| {
|
||||
WorkspaceError::SearchFailed {
|
||||
let chunk_ids = chunk_id_col
|
||||
.as_any()
|
||||
.downcast_ref::<StringArray>()
|
||||
.ok_or_else(|| WorkspaceError::SearchFailed {
|
||||
reason: "chunk_id wrong type".to_string(),
|
||||
}
|
||||
})?;
|
||||
let document_ids = document_id_col.as_any().downcast_ref::<StringArray>().ok_or_else(|| {
|
||||
WorkspaceError::SearchFailed {
|
||||
})?;
|
||||
let document_ids = document_id_col
|
||||
.as_any()
|
||||
.downcast_ref::<StringArray>()
|
||||
.ok_or_else(|| WorkspaceError::SearchFailed {
|
||||
reason: "document_id wrong type".to_string(),
|
||||
}
|
||||
})?;
|
||||
let contents = content_col.as_any().downcast_ref::<StringArray>().ok_or_else(|| {
|
||||
WorkspaceError::SearchFailed {
|
||||
})?;
|
||||
let contents = content_col
|
||||
.as_any()
|
||||
.downcast_ref::<StringArray>()
|
||||
.ok_or_else(|| WorkspaceError::SearchFailed {
|
||||
reason: "content wrong type".to_string(),
|
||||
}
|
||||
})?;
|
||||
})?;
|
||||
|
||||
for i in 0..batch.num_rows() {
|
||||
let raw_chunk_id = chunk_ids.value(i);
|
||||
let chunk_id = raw_chunk_id
|
||||
.parse::<Uuid>()
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!(
|
||||
"Invalid chunk_id UUID '{}': {}",
|
||||
raw_chunk_id, e
|
||||
),
|
||||
})?;
|
||||
let chunk_id =
|
||||
raw_chunk_id
|
||||
.parse::<Uuid>()
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!("Invalid chunk_id UUID '{}': {}", raw_chunk_id, e),
|
||||
})?;
|
||||
|
||||
let raw_document_id = document_ids.value(i);
|
||||
let document_id = raw_document_id
|
||||
.parse::<Uuid>()
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
let document_id = raw_document_id.parse::<Uuid>().map_err(|e| {
|
||||
WorkspaceError::SearchFailed {
|
||||
reason: format!(
|
||||
"Invalid document_id UUID '{}': {}",
|
||||
raw_document_id, e
|
||||
),
|
||||
})?;
|
||||
}
|
||||
})?;
|
||||
let content = contents.value(i).to_string();
|
||||
|
||||
results.push(RankedResult {
|
||||
@@ -418,12 +366,11 @@ pub use impl_lancedb::LanceDbVectorStore;
|
||||
|
||||
#[cfg(all(test, feature = "lancedb"))]
|
||||
mod tests {
|
||||
use std::sync::Arc;
|
||||
|
||||
use tempfile::TempDir;
|
||||
use uuid::Uuid;
|
||||
|
||||
use super::{LanceDbVectorStore, VectorStore, DEFAULT_EMBEDDING_DIM};
|
||||
use super::{DEFAULT_EMBEDDING_DIM, LanceDbVectorStore};
|
||||
use crate::workspace::vector_store::VectorStore;
|
||||
|
||||
fn make_embedding(seed: f32) -> Vec<f32> {
|
||||
(0..DEFAULT_EMBEDDING_DIM as usize)
|
||||
@@ -443,14 +390,7 @@ mod tests {
|
||||
let embedding = make_embedding(1.0);
|
||||
|
||||
store
|
||||
.insert_chunk(
|
||||
chunk_id,
|
||||
document_id,
|
||||
user_id,
|
||||
None,
|
||||
content,
|
||||
&embedding,
|
||||
)
|
||||
.store_embedding(chunk_id, document_id, user_id, None, content, &embedding)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
@@ -474,10 +414,9 @@ mod tests {
|
||||
let doc_id = Uuid::new_v4();
|
||||
let user_id = "user1";
|
||||
|
||||
// Insert 3 chunks with different embeddings
|
||||
for (i, seed) in [1.0, 2.0, 3.0].iter().enumerate() {
|
||||
store
|
||||
.insert_chunk(
|
||||
.store_embedding(
|
||||
Uuid::new_v4(),
|
||||
doc_id,
|
||||
user_id,
|
||||
@@ -489,7 +428,6 @@ mod tests {
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
// Search returns all 3, ordered by similarity
|
||||
let query_emb = make_embedding(2.0);
|
||||
let results = store
|
||||
.vector_search(user_id, None, &query_emb, 5)
|
||||
@@ -512,7 +450,7 @@ mod tests {
|
||||
let user_id = "user1";
|
||||
|
||||
store
|
||||
.insert_chunk(
|
||||
.store_embedding(
|
||||
Uuid::new_v4(),
|
||||
doc_id,
|
||||
user_id,
|
||||
@@ -529,7 +467,7 @@ mod tests {
|
||||
.unwrap();
|
||||
assert_eq!(results.len(), 1);
|
||||
|
||||
store.delete_chunks(doc_id).await.unwrap();
|
||||
store.delete_embeddings(doc_id).await.unwrap();
|
||||
|
||||
let results_after = store
|
||||
.vector_search(user_id, None, &make_embedding(1.0), 5)
|
||||
@@ -549,7 +487,7 @@ mod tests {
|
||||
let content = "original content";
|
||||
|
||||
store
|
||||
.insert_chunk(
|
||||
.store_embedding(
|
||||
chunk_id,
|
||||
doc_id,
|
||||
user_id,
|
||||
@@ -560,14 +498,12 @@ mod tests {
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Update with new embedding
|
||||
let new_embedding = make_embedding(5.0);
|
||||
store
|
||||
.update_chunk_embedding(chunk_id, doc_id, user_id, None, content, &new_embedding)
|
||||
.update_embedding(chunk_id, doc_id, user_id, None, content, &new_embedding)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Search with new embedding should find it
|
||||
let results = store
|
||||
.vector_search(user_id, None, &new_embedding, 5)
|
||||
.await
|
||||
@@ -585,7 +521,7 @@ mod tests {
|
||||
let embedding = make_embedding(1.0);
|
||||
|
||||
store
|
||||
.insert_chunk(
|
||||
.store_embedding(
|
||||
Uuid::new_v4(),
|
||||
doc_id,
|
||||
"user1",
|
||||
@@ -597,7 +533,7 @@ mod tests {
|
||||
.unwrap();
|
||||
|
||||
store
|
||||
.insert_chunk(
|
||||
.store_embedding(
|
||||
Uuid::new_v4(),
|
||||
doc_id,
|
||||
"user2",
|
||||
@@ -637,7 +573,7 @@ mod tests {
|
||||
let wrong_dim: Vec<f32> = vec![1.0; 100];
|
||||
|
||||
let err = store
|
||||
.insert_chunk(
|
||||
.store_embedding(
|
||||
Uuid::new_v4(),
|
||||
Uuid::new_v4(),
|
||||
"user1",
|
||||
@@ -648,6 +584,9 @@ mod tests {
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(err, crate::error::WorkspaceError::EmbeddingFailed { .. }));
|
||||
assert!(matches!(
|
||||
err,
|
||||
crate::error::WorkspaceError::EmbeddingFailed { .. }
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
+128
-6
@@ -42,22 +42,24 @@
|
||||
|
||||
mod chunker;
|
||||
mod document;
|
||||
mod embeddings;
|
||||
pub mod embeddings;
|
||||
pub mod hygiene;
|
||||
#[cfg(feature = "lancedb")]
|
||||
pub mod lancedb_store;
|
||||
#[cfg(feature = "postgres")]
|
||||
mod repository;
|
||||
mod search;
|
||||
pub mod vector_store;
|
||||
|
||||
pub use chunker::{ChunkConfig, chunk_document};
|
||||
pub use document::{MemoryChunk, MemoryDocument, WorkspaceEntry, paths};
|
||||
pub use embeddings::{EmbeddingProvider, MockEmbeddings, NearAiEmbeddings, OpenAiEmbeddings};
|
||||
#[cfg(feature = "lancedb")]
|
||||
pub use lancedb_store::{DEFAULT_EMBEDDING_DIM, LanceDbVectorStore};
|
||||
#[cfg(feature = "postgres")]
|
||||
pub use repository::Repository;
|
||||
#[cfg(feature = "lancedb")]
|
||||
pub use lancedb_store::{LanceDbVectorStore, VectorStore, DEFAULT_EMBEDDING_DIM};
|
||||
pub use search::{RankedResult, SearchConfig, SearchResult, reciprocal_rank_fusion};
|
||||
pub use vector_store::VectorStore;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
@@ -279,6 +281,12 @@ pub struct Workspace {
|
||||
storage: WorkspaceStorage,
|
||||
/// Embedding provider for semantic search.
|
||||
embeddings: Option<Arc<dyn EmbeddingProvider>>,
|
||||
/// Optional external vector store for semantic search.
|
||||
///
|
||||
/// When set, embeddings are stored here instead of (or in addition to)
|
||||
/// the database's built-in vector support, and hybrid search uses this
|
||||
/// for the vector component while FTS comes from the database.
|
||||
vector_store: Option<Arc<dyn VectorStore>>,
|
||||
}
|
||||
|
||||
impl Workspace {
|
||||
@@ -290,6 +298,7 @@ impl Workspace {
|
||||
agent_id: None,
|
||||
storage: WorkspaceStorage::Repo(Repository::new(pool)),
|
||||
embeddings: None,
|
||||
vector_store: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -302,6 +311,7 @@ impl Workspace {
|
||||
agent_id: None,
|
||||
storage: WorkspaceStorage::Db(db),
|
||||
embeddings: None,
|
||||
vector_store: None,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -317,6 +327,17 @@ impl Workspace {
|
||||
self
|
||||
}
|
||||
|
||||
/// Set an external vector store for semantic search.
|
||||
///
|
||||
/// When set, vector operations (store/search/delete embeddings) use this
|
||||
/// store instead of the database's built-in vector support. FTS continues
|
||||
/// to use the database. Hybrid search combines FTS from the database with
|
||||
/// vector results from this store via RRF.
|
||||
pub fn with_vector_store(mut self, store: Arc<dyn VectorStore>) -> Self {
|
||||
self.vector_store = Some(store);
|
||||
self
|
||||
}
|
||||
|
||||
/// Get the user ID.
|
||||
pub fn user_id(&self) -> &str {
|
||||
&self.user_id
|
||||
@@ -405,9 +426,21 @@ impl Workspace {
|
||||
|
||||
/// Delete a file.
|
||||
///
|
||||
/// Also deletes associated chunks.
|
||||
/// Also deletes associated chunks (from both DB and external vector store).
|
||||
pub async fn delete(&self, path: &str) -> Result<(), WorkspaceError> {
|
||||
let path = normalize_path(path);
|
||||
|
||||
// Clean up external vector store before DB cascade deletes chunks
|
||||
if let Some(ref vs) = self.vector_store
|
||||
&& let Ok(doc) = self
|
||||
.storage
|
||||
.get_document_by_path(&self.user_id, self.agent_id, &path)
|
||||
.await
|
||||
&& let Err(e) = vs.delete_embeddings(doc.id).await
|
||||
{
|
||||
tracing::warn!("Failed to delete embeddings from vector store: {}", e);
|
||||
}
|
||||
|
||||
self.storage
|
||||
.delete_document_by_path(&self.user_id, self.agent_id, &path)
|
||||
.await
|
||||
@@ -599,6 +632,50 @@ impl Workspace {
|
||||
None
|
||||
};
|
||||
|
||||
// When an external vector store is configured, do FTS from the
|
||||
// database and vector search from the store, then fuse with RRF.
|
||||
if let Some(ref vs) = self.vector_store {
|
||||
// FTS from database (disable vector to avoid double-searching)
|
||||
let fts_results = if config.use_fts {
|
||||
let fts_config = SearchConfig {
|
||||
use_fts: true,
|
||||
use_vector: false,
|
||||
..config.clone()
|
||||
};
|
||||
let fts_search = self
|
||||
.storage
|
||||
.hybrid_search(&self.user_id, self.agent_id, query, None, &fts_config)
|
||||
.await?;
|
||||
fts_search
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.map(|(i, r)| RankedResult {
|
||||
chunk_id: r.chunk_id,
|
||||
document_id: r.document_id,
|
||||
content: r.content,
|
||||
rank: (i + 1) as u32,
|
||||
})
|
||||
.collect()
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
|
||||
// Vector search from external store
|
||||
let vector_results = if config.use_vector {
|
||||
if let Some(ref emb) = embedding {
|
||||
vs.vector_search(&self.user_id, self.agent_id, emb, config.pre_fusion_limit)
|
||||
.await?
|
||||
} else {
|
||||
Vec::new()
|
||||
}
|
||||
} else {
|
||||
Vec::new()
|
||||
};
|
||||
|
||||
return Ok(reciprocal_rank_fusion(fts_results, vector_results, &config));
|
||||
}
|
||||
|
||||
// No external vector store — use database's built-in hybrid search
|
||||
self.storage
|
||||
.hybrid_search(
|
||||
&self.user_id,
|
||||
@@ -620,9 +697,14 @@ impl Workspace {
|
||||
// Chunk the content
|
||||
let chunks = chunk_document(&doc.content, ChunkConfig::default());
|
||||
|
||||
// Delete old chunks
|
||||
// Delete old chunks from database
|
||||
self.storage.delete_chunks(document_id).await?;
|
||||
|
||||
// Delete old embeddings from external vector store
|
||||
if let Some(ref vs) = self.vector_store {
|
||||
vs.delete_embeddings(document_id).await?;
|
||||
}
|
||||
|
||||
// Insert new chunks
|
||||
for (index, content) in chunks.into_iter().enumerate() {
|
||||
// Generate embedding if provider available
|
||||
@@ -638,9 +720,26 @@ impl Workspace {
|
||||
None
|
||||
};
|
||||
|
||||
self.storage
|
||||
let chunk_id = self
|
||||
.storage
|
||||
.insert_chunk(document_id, index as i32, &content, embedding.as_deref())
|
||||
.await?;
|
||||
|
||||
// Sync embedding to external vector store
|
||||
if let (Some(vs), Some(emb)) = (&self.vector_store, &embedding)
|
||||
&& let Err(e) = vs
|
||||
.store_embedding(
|
||||
chunk_id,
|
||||
document_id,
|
||||
&doc.user_id,
|
||||
doc.agent_id,
|
||||
&content,
|
||||
emb,
|
||||
)
|
||||
.await
|
||||
{
|
||||
tracing::warn!("Failed to store embedding in vector store: {}", e);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
@@ -757,6 +856,29 @@ impl Workspace {
|
||||
self.storage
|
||||
.update_chunk_embedding(chunk.id, &embedding)
|
||||
.await?;
|
||||
|
||||
// Sync to external vector store
|
||||
if let Some(ref vs) = self.vector_store {
|
||||
let doc = self.storage.get_document_by_id(chunk.document_id).await?;
|
||||
if let Err(e) = vs
|
||||
.update_embedding(
|
||||
chunk.id,
|
||||
chunk.document_id,
|
||||
&doc.user_id,
|
||||
doc.agent_id,
|
||||
&chunk.content,
|
||||
&embedding,
|
||||
)
|
||||
.await
|
||||
{
|
||||
tracing::warn!(
|
||||
"Failed to sync embedding to vector store for chunk {}: {}",
|
||||
chunk.id,
|
||||
e
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
count += 1;
|
||||
}
|
||||
Err(e) => {
|
||||
|
||||
@@ -347,36 +347,6 @@ impl Repository {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Get a chunk by ID.
|
||||
pub async fn get_chunk_by_id(
|
||||
&self,
|
||||
chunk_id: Uuid,
|
||||
) -> Result<Option<MemoryChunk>, WorkspaceError> {
|
||||
let conn = self.conn().await?;
|
||||
|
||||
let row = conn
|
||||
.query_opt(
|
||||
r#"
|
||||
SELECT id, document_id, chunk_index, content, created_at
|
||||
FROM memory_chunks WHERE id = $1
|
||||
"#,
|
||||
&[&chunk_id],
|
||||
)
|
||||
.await
|
||||
.map_err(|e| WorkspaceError::SearchFailed {
|
||||
reason: format!("Query failed: {}", e),
|
||||
})?;
|
||||
|
||||
Ok(row.map(|r| MemoryChunk {
|
||||
id: r.get("id"),
|
||||
document_id: r.get("document_id"),
|
||||
chunk_index: r.get("chunk_index"),
|
||||
content: r.get("content"),
|
||||
embedding: None,
|
||||
created_at: r.get("created_at"),
|
||||
}))
|
||||
}
|
||||
|
||||
/// Get chunks without embeddings for backfilling.
|
||||
pub async fn get_chunks_without_embeddings(
|
||||
&self,
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
//! Vector store abstraction for workspace semantic search.
|
||||
//!
|
||||
//! Separates vector search from the main `Database` trait so that
|
||||
//! third-party vector backends (LanceDB, Qdrant, Pinecone, etc.) can
|
||||
//! be added by implementing a 4-method trait instead of wrapping the
|
||||
//! entire ~80-method `Database` trait.
|
||||
//!
|
||||
//! When no external vector store is configured, the built-in database
|
||||
//! vector support (pgvector / libsql_vector_idx) is used via the
|
||||
//! `Database::hybrid_search` method directly.
|
||||
|
||||
use async_trait::async_trait;
|
||||
use uuid::Uuid;
|
||||
|
||||
use crate::error::WorkspaceError;
|
||||
use crate::workspace::search::RankedResult;
|
||||
|
||||
/// External vector store for semantic search.
|
||||
///
|
||||
/// Implementations hold chunk embeddings and perform vector similarity
|
||||
/// queries. Document/chunk metadata and FTS stay in the main database;
|
||||
/// only embeddings live here.
|
||||
///
|
||||
/// # Adding a new backend
|
||||
///
|
||||
/// 1. Implement this trait for your backend (4 methods).
|
||||
/// 2. Feature-gate the module (`#[cfg(feature = "mybackend")]`).
|
||||
/// 3. Pass `Arc<dyn VectorStore>` to `Workspace::with_vector_store()`.
|
||||
///
|
||||
/// That's it — no Database wrapper, no delegation boilerplate.
|
||||
#[async_trait]
|
||||
pub trait VectorStore: Send + Sync {
|
||||
/// Store an embedding for a chunk.
|
||||
async fn store_embedding(
|
||||
&self,
|
||||
chunk_id: Uuid,
|
||||
document_id: Uuid,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
content: &str,
|
||||
embedding: &[f32],
|
||||
) -> Result<(), WorkspaceError>;
|
||||
|
||||
/// Update an existing chunk's embedding (delete + re-insert is fine).
|
||||
async fn update_embedding(
|
||||
&self,
|
||||
chunk_id: Uuid,
|
||||
document_id: Uuid,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
content: &str,
|
||||
embedding: &[f32],
|
||||
) -> Result<(), WorkspaceError>;
|
||||
|
||||
/// Delete all embeddings for a document.
|
||||
async fn delete_embeddings(&self, document_id: Uuid) -> Result<(), WorkspaceError>;
|
||||
|
||||
/// Vector similarity search, filtered by user and optional agent.
|
||||
///
|
||||
/// Returns results ranked by similarity (rank 1 = most similar).
|
||||
async fn vector_search(
|
||||
&self,
|
||||
user_id: &str,
|
||||
agent_id: Option<Uuid>,
|
||||
embedding: &[f32],
|
||||
limit: usize,
|
||||
) -> Result<Vec<RankedResult>, WorkspaceError>;
|
||||
}
|
||||
+66
-102
@@ -1,20 +1,18 @@
|
||||
//! Integration tests for LanceDB vector store with Database wrapper.
|
||||
//! Integration tests for LanceDB vector store with Workspace composition.
|
||||
//!
|
||||
//! Requires: cargo test --features "libsql,lancedb"
|
||||
//!
|
||||
//! Verifies DbWithLanceVectorStore: document + chunk insert, hybrid search
|
||||
//! (FTS from libSQL, vector from LanceDB), delete_chunks sync.
|
||||
//! Verifies that Workspace correctly composes FTS from libSQL with vector
|
||||
//! search from LanceDB via the VectorStore trait.
|
||||
|
||||
#![cfg(all(feature = "libsql", feature = "lancedb"))]
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use ironclaw::db::lancedb_wrapper::DbWithLanceVectorStore;
|
||||
use ironclaw::db::Database;
|
||||
use ironclaw::db::libsql_backend::LibSqlBackend;
|
||||
use ironclaw::workspace::{LanceDbVectorStore, SearchConfig};
|
||||
use ironclaw::workspace::{LanceDbVectorStore, SearchConfig, Workspace};
|
||||
use tempfile::TempDir;
|
||||
use uuid::Uuid;
|
||||
|
||||
const EMBEDDING_DIM: usize = 1536;
|
||||
|
||||
@@ -24,136 +22,102 @@ fn make_embedding(seed: f32) -> Vec<f32> {
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn setup_wrapped_db() -> (Arc<dyn Database>, TempDir) {
|
||||
let libsql = LibSqlBackend::new_memory().await.unwrap();
|
||||
/// Mock embedding provider that returns deterministic embeddings.
|
||||
struct FixedEmbeddings {
|
||||
embedding: Vec<f32>,
|
||||
}
|
||||
|
||||
#[async_trait::async_trait]
|
||||
impl ironclaw::workspace::EmbeddingProvider for FixedEmbeddings {
|
||||
fn dimension(&self) -> usize {
|
||||
EMBEDDING_DIM
|
||||
}
|
||||
|
||||
fn model_name(&self) -> &str {
|
||||
"fixed-test"
|
||||
}
|
||||
|
||||
fn max_input_length(&self) -> usize {
|
||||
8192
|
||||
}
|
||||
|
||||
async fn embed(
|
||||
&self,
|
||||
_text: &str,
|
||||
) -> Result<Vec<f32>, ironclaw::workspace::embeddings::EmbeddingError> {
|
||||
Ok(self.embedding.clone())
|
||||
}
|
||||
}
|
||||
|
||||
async fn setup_workspace() -> (Workspace, TempDir, TempDir) {
|
||||
// Use a temp file (not :memory:) because libSQL in-memory DBs are connection-local
|
||||
let db_dir = TempDir::new().unwrap();
|
||||
let db_path = db_dir.path().join("test.db");
|
||||
let libsql = LibSqlBackend::new_local(&db_path).await.unwrap();
|
||||
libsql.run_migrations().await.unwrap();
|
||||
|
||||
let lancedb_dir = TempDir::new().unwrap();
|
||||
let store = LanceDbVectorStore::new(lancedb_dir.path())
|
||||
.await
|
||||
.unwrap();
|
||||
let store = LanceDbVectorStore::new(lancedb_dir.path()).await.unwrap();
|
||||
|
||||
let db = Arc::new(DbWithLanceVectorStore::new(
|
||||
Arc::new(libsql) as Arc<dyn Database>,
|
||||
Arc::new(store),
|
||||
)) as Arc<dyn Database>;
|
||||
let embedding = make_embedding(1.0);
|
||||
let ws = Workspace::new_with_db("test_user", Arc::new(libsql) as Arc<dyn Database>)
|
||||
.with_vector_store(Arc::new(store))
|
||||
.with_embeddings(Arc::new(FixedEmbeddings { embedding }));
|
||||
|
||||
(db, lancedb_dir)
|
||||
(ws, lancedb_dir, db_dir)
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_wrapper_hybrid_search_combines_fts_and_vector() {
|
||||
let (db, _) = setup_wrapped_db().await;
|
||||
async fn test_workspace_hybrid_search_with_lancedb() {
|
||||
let (ws, _keep_lance, _keep_db) = setup_workspace().await;
|
||||
|
||||
let user_id = "test_user";
|
||||
let agent_id: Option<Uuid> = None;
|
||||
// Write a document — this triggers chunking + embedding + LanceDB sync
|
||||
ws.write(
|
||||
"context/rust.md",
|
||||
"Rust is a systems programming language focused on safety.",
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Create document
|
||||
let doc = db
|
||||
.get_or_create_document_by_path(user_id, agent_id, "context/rust.md")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Write content for FTS
|
||||
db.update_document(doc.id, "Rust is a systems programming language focused on safety and performance.")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Chunk and insert with embedding (triggers sync to LanceDB)
|
||||
let content = "Rust is a systems programming language focused on safety.";
|
||||
let embedding = make_embedding(1.0);
|
||||
|
||||
let chunk_id = db
|
||||
.insert_chunk(doc.id, 0, content, Some(&embedding))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Hybrid search: FTS for "Rust" + vector for semantic
|
||||
let config = SearchConfig::default().with_limit(5);
|
||||
let results = db
|
||||
.hybrid_search(
|
||||
user_id,
|
||||
agent_id,
|
||||
"Rust",
|
||||
Some(&embedding),
|
||||
&config,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
// Hybrid search: FTS for "Rust" + vector from LanceDB
|
||||
let results = ws.search("Rust", 5).await.unwrap();
|
||||
|
||||
assert!(!results.is_empty(), "hybrid search should return results");
|
||||
assert_eq!(results[0].chunk_id, chunk_id);
|
||||
assert!(results[0].content.contains("Rust"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_wrapper_delete_chunks_removes_from_both() {
|
||||
let (db, _) = setup_wrapped_db().await;
|
||||
async fn test_workspace_delete_removes_from_lancedb() {
|
||||
let (ws, _keep_lance, _keep_db) = setup_workspace().await;
|
||||
|
||||
let user_id = "test_user";
|
||||
let agent_id: Option<Uuid> = None;
|
||||
|
||||
let doc = db
|
||||
.get_or_create_document_by_path(user_id, agent_id, "notes/deleted.md")
|
||||
ws.write("notes/deleted.md", "Content to be deleted.")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
db.update_document(doc.id, "Content to be deleted.").await.unwrap();
|
||||
|
||||
let embedding = make_embedding(2.0);
|
||||
db.insert_chunk(doc.id, 0, "Content to be deleted.", Some(&embedding))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let before = db
|
||||
.hybrid_search(user_id, agent_id, "deleted", Some(&embedding), &SearchConfig::default())
|
||||
.await
|
||||
.unwrap();
|
||||
let before = ws.search("deleted", 5).await.unwrap();
|
||||
assert_eq!(before.len(), 1);
|
||||
|
||||
db.delete_chunks(doc.id).await.unwrap();
|
||||
ws.delete("notes/deleted.md").await.unwrap();
|
||||
|
||||
let after = db
|
||||
.hybrid_search(user_id, agent_id, "deleted", Some(&embedding), &SearchConfig::default())
|
||||
.await
|
||||
.unwrap();
|
||||
let after = ws.search("deleted", 5).await.unwrap();
|
||||
assert!(after.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_wrapper_insert_chunk_syncs_to_lancedb() {
|
||||
let (db, _) = setup_wrapped_db().await;
|
||||
async fn test_workspace_vector_only_search_uses_lancedb() {
|
||||
let (ws, _keep_lance, _keep_db) = setup_workspace().await;
|
||||
|
||||
let user_id = "sync_user";
|
||||
let agent_id: Option<Uuid> = None;
|
||||
|
||||
let doc = db
|
||||
.get_or_create_document_by_path(user_id, agent_id, "sync/test.md")
|
||||
ws.write("sync/test.md", "Semantic content for vector search")
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let content = "Semantic content for vector search";
|
||||
let embedding = make_embedding(3.0);
|
||||
|
||||
let chunk_id = db
|
||||
.insert_chunk(doc.id, 0, content, Some(&embedding))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Vector-only search (no FTS query match) - should still find via LanceDB
|
||||
// Vector-only search should find via LanceDB even with non-matching FTS query
|
||||
let config = SearchConfig::default().vector_only().with_limit(5);
|
||||
let results = db
|
||||
.hybrid_search(
|
||||
user_id,
|
||||
agent_id,
|
||||
"nonexistent_fts_term",
|
||||
Some(&embedding),
|
||||
&config,
|
||||
)
|
||||
let results = ws
|
||||
.search_with_config("nonexistent_fts_term", config)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(results.len(), 1);
|
||||
assert_eq!(results[0].chunk_id, chunk_id);
|
||||
assert_eq!(results[0].content, content);
|
||||
assert!(results[0].content.contains("Semantic content"));
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user