1use crate::{
4 Collection, CollectionConfig, Database, DeleteSelector, Error, Filter, GetRequest, GetResult,
5 MutationResult, ObjectId, Point, PointId, Query, QueryResult, Record, Result, ScoredPoint,
6 SnapshotMutation, WriteResult,
7};
8use git2::ErrorCode;
9use std::fs;
10use std::path::Path;
11use std::thread;
12
13const WRITE_ATTEMPTS: usize = 8;
14
15pub fn open(path: impl AsRef<Path>) -> Result<Store> {
21 Store::open(path)
22}
23
24#[derive(Clone, Debug)]
26pub struct Store {
27 database: Database,
28}
29
30#[derive(Clone, Debug)]
32pub struct CollectionHandle {
33 database: Database,
34 name: String,
35}
36
37impl Store {
38 pub fn open(path: impl AsRef<Path>) -> Result<Self> {
40 let path = path.as_ref();
41 let database = match Database::open(path) {
42 Ok(database) => database,
43 Err(_) if !path.exists() => Database::init_bare(path)?,
44 Err(_) if path.is_dir() && directory_is_empty(path)? => Database::init_bare(path)?,
45 Err(open_error) if path.is_dir() => {
46 return Err(Error::Invalid(format!(
47 "database path {} is a nonempty directory but not a Git repository: {open_error}",
48 path.display()
49 )))
50 }
51 Err(open_error) => return Err(open_error),
52 };
53 Ok(Self { database })
54 }
55
56 pub fn collection(&self, name: impl Into<String>) -> CollectionHandle {
61 CollectionHandle {
62 database: self.database.clone(),
63 name: name.into(),
64 }
65 }
66
67 pub fn list_collections(&self) -> Result<Vec<String>> {
69 self.database.list_collections()
70 }
71
72 pub fn advanced(&self) -> &Database {
74 &self.database
75 }
76}
77
78impl CollectionHandle {
79 pub fn upsert(&self, points: impl IntoIterator<Item = Point>) -> Result<WriteResult> {
84 self.upsert_with_vector_space(points, None)
85 }
86
87 pub(crate) fn upsert_with_vector_space(
88 &self,
89 points: impl IntoIterator<Item = Point>,
90 vector_space: Option<&str>,
91 ) -> Result<WriteResult> {
92 let points: Vec<Point> = points.into_iter().collect();
93 let dimension = inferred_dimension(&points)?;
94
95 for _ in 0..WRITE_ATTEMPTS {
96 let collection = match self.database.collection(&self.name) {
97 Ok(collection) => collection,
98 Err(Error::CollectionNotFound(_)) => {
99 let mut config = CollectionConfig::new(dimension);
100 config.vector_space = vector_space.map(str::to_owned);
101 match self.database.create_collection(&self.name, config) {
102 Ok(collection) => collection,
103 Err(Error::CollectionExists(_)) => {
104 thread::yield_now();
105 continue;
106 }
107 Err(Error::Git(error)) if retryable_ref_error(&error) => {
108 thread::yield_now();
109 continue;
110 }
111 Err(error) => return Err(error),
112 }
113 }
114 Err(error) => return Err(error),
115 };
116
117 if let Some(vector_space) = vector_space {
118 let actual = collection.info()?.config.vector_space;
119 if actual.as_deref() != Some(vector_space) {
120 return Err(Error::Invalid(format!(
121 "collection {:?} uses vector space {:?}, expected {:?}",
122 self.name, actual, vector_space
123 )));
124 }
125 }
126
127 match collection.upsert(points.clone()) {
128 Ok(result) => return Ok(result),
129 Err(Error::StaleRoot { .. }) => thread::yield_now(),
130 Err(Error::Git(error)) if retryable_ref_error(&error) => thread::yield_now(),
131 Err(error) => return Err(error),
132 }
133 }
134
135 Err(Error::Invalid(format!(
136 "collection {:?} changed repeatedly while applying the first write; retry the upsert",
137 self.name
138 )))
139 }
140
141 pub fn search(
147 &self,
148 vector: impl IntoIterator<Item = f32>,
149 limit: usize,
150 ) -> Result<Vec<ScoredPoint>> {
151 Ok(self
152 .advanced()?
153 .query(Query::new(vector, limit).with_payload())?
154 .points)
155 }
156
157 pub fn query(&self, query: Query) -> Result<QueryResult> {
159 self.advanced()?.query(query)
160 }
161
162 pub fn query_batch(
164 &self,
165 queries: impl IntoIterator<Item = Query>,
166 ) -> Result<Vec<QueryResult>> {
167 let collection = self.advanced()?;
168 queries
169 .into_iter()
170 .map(|query| collection.query(query))
171 .collect()
172 }
173
174 pub fn get(&self, request: GetRequest) -> Result<GetResult> {
176 self.advanced()?.get(request)
177 }
178
179 pub fn get_ids(&self, ids: impl IntoIterator<Item = PointId>) -> Result<Vec<Record>> {
181 Ok(self
182 .advanced()?
183 .get(GetRequest {
184 ids: ids.into_iter().collect(),
185 with_payload: true,
186 ..GetRequest::default()
187 })?
188 .points)
189 }
190
191 pub fn delete_ids(&self, ids: impl IntoIterator<Item = PointId>) -> Result<WriteResult> {
193 self.advanced()?.delete(DeleteSelector {
194 ids: ids.into_iter().collect(),
195 ..DeleteSelector::default()
196 })
197 }
198
199 pub fn delete(&self, selector: DeleteSelector) -> Result<WriteResult> {
201 self.advanced()?.delete(selector)
202 }
203
204 pub fn apply(
206 &self,
207 mutations: impl IntoIterator<Item = SnapshotMutation>,
208 ) -> Result<MutationResult> {
209 self.advanced()?.apply(mutations.into_iter().collect())
210 }
211
212 pub fn count(&self) -> Result<usize> {
214 Ok(self.advanced()?.count(None)?.count)
215 }
216
217 pub fn count_where(&self, filter: Filter) -> Result<usize> {
219 Ok(self.advanced()?.count(Some(filter))?.count)
220 }
221
222 pub fn root(&self) -> Result<ObjectId> {
224 self.advanced()?.root()
225 }
226
227 pub fn restore(&self, revision: impl AsRef<str>) -> Result<WriteResult> {
229 self.advanced()?.restore(revision)
230 }
231
232 pub fn peek(&self, limit: usize) -> Result<Vec<Record>> {
234 Ok(self
235 .advanced()?
236 .get(GetRequest {
237 limit: Some(limit),
238 with_payload: true,
239 ..GetRequest::default()
240 })?
241 .points)
242 }
243
244 pub fn advanced(&self) -> Result<Collection> {
246 self.database.collection(&self.name)
247 }
248}
249
250fn directory_is_empty(path: &Path) -> Result<bool> {
251 Ok(fs::read_dir(path)?.next().is_none())
252}
253
254fn inferred_dimension(points: &[Point]) -> Result<usize> {
255 let Some(first) = points.first() else {
256 return Err(Error::Invalid("upsert batch must not be empty".into()));
257 };
258 let dimension = first.vector.len();
259 if dimension == 0 {
260 return Err(Error::Invalid("point vector must not be empty".into()));
261 }
262 if let Some(point) = points.iter().find(|point| point.vector.len() != dimension) {
263 return Err(Error::Invalid(format!(
264 "point {} has dimension {}, expected {dimension}",
265 point.id,
266 point.vector.len()
267 )));
268 }
269 Ok(dimension)
270}
271
272fn retryable_ref_error(error: &git2::Error) -> bool {
273 matches!(
274 error.code(),
275 ErrorCode::Exists | ErrorCode::Locked | ErrorCode::Modified
276 )
277}