use crate::{
Collection, CollectionConfig, Database, DeleteSelector, Error, Filter, GetRequest, GetResult,
MutationResult, ObjectId, Point, PointId, Query, QueryResult, Record, Result, ScoredPoint,
SnapshotMutation, WriteResult,
};
use git2::ErrorCode;
use std::fs;
use std::path::Path;
use std::thread;
const WRITE_ATTEMPTS: usize = 8;
pub fn open(path: impl AsRef<Path>) -> Result<Store> {
Store::open(path)
}
#[derive(Clone, Debug)]
pub struct Store {
database: Database,
}
#[derive(Clone, Debug)]
pub struct CollectionHandle {
database: Database,
name: String,
}
impl Store {
pub fn open(path: impl AsRef<Path>) -> Result<Self> {
let path = path.as_ref();
let database = match Database::open(path) {
Ok(database) => database,
Err(_) if !path.exists() => Database::init_bare(path)?,
Err(_) if path.is_dir() && directory_is_empty(path)? => Database::init_bare(path)?,
Err(open_error) if path.is_dir() => {
return Err(Error::Invalid(format!(
"database path {} is a nonempty directory but not a Git repository: {open_error}",
path.display()
)))
}
Err(open_error) => return Err(open_error),
};
Ok(Self { database })
}
pub fn collection(&self, name: impl Into<String>) -> CollectionHandle {
CollectionHandle {
database: self.database.clone(),
name: name.into(),
}
}
pub fn list_collections(&self) -> Result<Vec<String>> {
self.database.list_collections()
}
pub fn advanced(&self) -> &Database {
&self.database
}
}
impl CollectionHandle {
pub fn upsert(&self, points: impl IntoIterator<Item = Point>) -> Result<WriteResult> {
self.upsert_with_vector_space(points, None)
}
pub(crate) fn upsert_with_vector_space(
&self,
points: impl IntoIterator<Item = Point>,
vector_space: Option<&str>,
) -> Result<WriteResult> {
let points: Vec<Point> = points.into_iter().collect();
let dimension = inferred_dimension(&points)?;
for _ in 0..WRITE_ATTEMPTS {
let collection = match self.database.collection(&self.name) {
Ok(collection) => collection,
Err(Error::CollectionNotFound(_)) => {
let mut config = CollectionConfig::new(dimension);
config.vector_space = vector_space.map(str::to_owned);
match self.database.create_collection(&self.name, config) {
Ok(collection) => collection,
Err(Error::CollectionExists(_)) => {
thread::yield_now();
continue;
}
Err(Error::Git(error)) if retryable_ref_error(&error) => {
thread::yield_now();
continue;
}
Err(error) => return Err(error),
}
}
Err(error) => return Err(error),
};
if let Some(vector_space) = vector_space {
let actual = collection.info()?.config.vector_space;
if actual.as_deref() != Some(vector_space) {
return Err(Error::Invalid(format!(
"collection {:?} uses vector space {:?}, expected {:?}",
self.name, actual, vector_space
)));
}
}
match collection.upsert(points.clone()) {
Ok(result) => return Ok(result),
Err(Error::StaleRoot { .. }) => thread::yield_now(),
Err(Error::Git(error)) if retryable_ref_error(&error) => thread::yield_now(),
Err(error) => return Err(error),
}
}
Err(Error::Invalid(format!(
"collection {:?} changed repeatedly while applying the first write; retry the upsert",
self.name
)))
}
pub fn search(
&self,
vector: impl IntoIterator<Item = f32>,
limit: usize,
) -> Result<Vec<ScoredPoint>> {
Ok(self
.advanced()?
.query(Query::new(vector, limit).with_payload())?
.points)
}
pub fn query(&self, query: Query) -> Result<QueryResult> {
self.advanced()?.query(query)
}
pub fn query_batch(
&self,
queries: impl IntoIterator<Item = Query>,
) -> Result<Vec<QueryResult>> {
let collection = self.advanced()?;
queries
.into_iter()
.map(|query| collection.query(query))
.collect()
}
pub fn get(&self, request: GetRequest) -> Result<GetResult> {
self.advanced()?.get(request)
}
pub fn get_ids(&self, ids: impl IntoIterator<Item = PointId>) -> Result<Vec<Record>> {
Ok(self
.advanced()?
.get(GetRequest {
ids: ids.into_iter().collect(),
with_payload: true,
..GetRequest::default()
})?
.points)
}
pub fn delete_ids(&self, ids: impl IntoIterator<Item = PointId>) -> Result<WriteResult> {
self.advanced()?.delete(DeleteSelector {
ids: ids.into_iter().collect(),
..DeleteSelector::default()
})
}
pub fn delete(&self, selector: DeleteSelector) -> Result<WriteResult> {
self.advanced()?.delete(selector)
}
pub fn apply(
&self,
mutations: impl IntoIterator<Item = SnapshotMutation>,
) -> Result<MutationResult> {
self.advanced()?.apply(mutations.into_iter().collect())
}
pub fn count(&self) -> Result<usize> {
Ok(self.advanced()?.count(None)?.count)
}
pub fn count_where(&self, filter: Filter) -> Result<usize> {
Ok(self.advanced()?.count(Some(filter))?.count)
}
pub fn root(&self) -> Result<ObjectId> {
self.advanced()?.root()
}
pub fn restore(&self, revision: impl AsRef<str>) -> Result<WriteResult> {
self.advanced()?.restore(revision)
}
pub fn peek(&self, limit: usize) -> Result<Vec<Record>> {
Ok(self
.advanced()?
.get(GetRequest {
limit: Some(limit),
with_payload: true,
..GetRequest::default()
})?
.points)
}
pub fn advanced(&self) -> Result<Collection> {
self.database.collection(&self.name)
}
}
fn directory_is_empty(path: &Path) -> Result<bool> {
Ok(fs::read_dir(path)?.next().is_none())
}
fn inferred_dimension(points: &[Point]) -> Result<usize> {
let Some(first) = points.first() else {
return Err(Error::Invalid("upsert batch must not be empty".into()));
};
let dimension = first.vector.len();
if dimension == 0 {
return Err(Error::Invalid("point vector must not be empty".into()));
}
if let Some(point) = points.iter().find(|point| point.vector.len() != dimension) {
return Err(Error::Invalid(format!(
"point {} has dimension {}, expected {dimension}",
point.id,
point.vector.len()
)));
}
Ok(dimension)
}
fn retryable_ref_error(error: &git2::Error) -> bool {
matches!(
error.code(),
ErrorCode::Exists | ErrorCode::Locked | ErrorCode::Modified
)
}