1use git_vdb::{
2 CollectionConfig, Point, PointId, Query, QueryParams, SnapshotEngine, SnapshotMutation,
3};
4use serde::Deserialize;
5use serde_json::{json, Map};
6use sha2::{Digest, Sha256};
7use std::collections::HashSet;
8use std::env;
9use std::fs;
10use std::path::{Path, PathBuf};
11use std::thread;
12use std::time::Duration;
13use std::time::Instant;
14
15#[derive(Deserialize)]
16struct RunSpec {
17 schema_version: u32,
18 dimension: usize,
19 point_count: usize,
20 query_count: usize,
21 points_path: PathBuf,
22 queries_path: PathBuf,
23 k: Vec<usize>,
24}
25
26fn main() -> Result<(), Box<dyn std::error::Error>> {
27 let args = env::args_os().skip(1).collect::<Vec<_>>();
28 match args.as_slice() {
29 [command, input, repository, output] if command == "build" => {
30 build(Path::new(input), Path::new(repository), Path::new(output))
31 }
32 [command, input, repository, build_report, mode, output] if command == "query" => query(
33 Path::new(input),
34 Path::new(repository),
35 Path::new(build_report),
36 mode.to_str().ok_or("query mode is not UTF-8")?,
37 Path::new(output),
38 ),
39 [command, input, repository, build_report, fraction, output] if command == "mutate" => {
40 mutate(
41 Path::new(input),
42 Path::new(repository),
43 Path::new(build_report),
44 fraction
45 .to_str()
46 .ok_or("mutation fraction is not UTF-8")?
47 .parse()?,
48 Path::new(output),
49 false,
50 )
51 }
52 [command, input, repository, build_report, fraction, output]
53 if command == "mutate-sample-stable" =>
54 {
55 mutate(
56 Path::new(input),
57 Path::new(repository),
58 Path::new(build_report),
59 fraction
60 .to_str()
61 .ok_or("mutation fraction is not UTF-8")?
62 .parse()?,
63 Path::new(output),
64 true,
65 )
66 }
67 [command, repository, build_report, output] if command == "validate" => validate(
68 Path::new(repository),
69 Path::new(build_report),
70 Path::new(output),
71 ),
72 _ => Err(
73 "usage: lancedb_git_vdb_profile build INPUT.json REPOSITORY OUTPUT.json\n lancedb_git_vdb_profile query INPUT.json REPOSITORY BUILD.json exact|approximate|approximate-after-exact OUTPUT.json\n lancedb_git_vdb_profile mutate INPUT.json REPOSITORY BUILD.json FRACTION OUTPUT.json\n lancedb_git_vdb_profile mutate-sample-stable INPUT.json REPOSITORY BUILD.json FRACTION OUTPUT.json\n lancedb_git_vdb_profile validate REPOSITORY BUILD.json OUTPUT.json"
74 .into(),
75 ),
76 }
77}
78
79fn mutate(
80 input: &Path,
81 repository: &Path,
82 build_report: &Path,
83 fraction: f64,
84 output: &Path,
85 sample_stable: bool,
86) -> Result<(), Box<dyn std::error::Error>> {
87 let spec = read_spec(input)?;
88 if !(0.0..=1.0).contains(&fraction) || fraction == 0.0 {
89 return Err("mutation fraction must be greater than zero and at most one".into());
90 }
91 let vectors = read_vectors(&spec.points_path, spec.point_count, spec.dimension)?;
92 let points = make_points(&vectors);
93 let count = ((spec.point_count as f64 * fraction).round() as usize).clamp(1, spec.point_count);
94 let mut changed = if sample_stable {
95 sample_stable_points(points, count)?
96 } else {
97 points[..count].to_vec()
98 };
99 for point in &mut changed {
100 point.vector[0] += 0.001;
101 }
102 let build: serde_json::Value = serde_json::from_slice(&fs::read(build_report)?)?;
103 let root = build
104 .get("root")
105 .and_then(serde_json::Value::as_str)
106 .ok_or("build report root is missing")?;
107 let engine = SnapshotEngine::open(repository)?;
108
109 let upsert_started = Instant::now();
110 let upserted = engine.apply(
111 root,
112 changed.into_iter().map(SnapshotMutation::upsert).collect(),
113 )?;
114 let upsert_us = micros(upsert_started);
115 let delete_started = Instant::now();
116 let deleted = engine.apply(
117 root,
118 vec![SnapshotMutation::delete_ids(
119 (0..count).map(|id| PointId::from(id as u64)),
120 )],
121 )?;
122 let delete_us = micros(delete_started);
123 fs::write(
124 output,
125 serde_json::to_vec_pretty(&json!({
126 "schema_version": 1,
127 "root": root,
128 "fraction": fraction,
129 "points": count,
130 "sample_stable": sample_stable,
131 "upsert_us": upsert_us,
132 "delete_us": delete_us,
133 "upsert_root": upserted.root(),
134 "delete_root": deleted.root(),
135 "on_disk_bytes_after": directory_bytes(repository)?,
136 }))?,
137 )?;
138 Ok(())
139}
140
141fn sample_stable_points(
142 points: Vec<Point>,
143 count: usize,
144) -> Result<Vec<Point>, Box<dyn std::error::Error>> {
145 let mut sample_order = points
146 .iter()
147 .map(|point| Ok((uint_id_digest(&point.id)?, point.id.clone())))
148 .collect::<Result<Vec<_>, Box<dyn std::error::Error>>>()?;
149 sample_order.sort();
150 let sample_ids = sample_order
151 .into_iter()
152 .take(8_192.min(points.len()))
153 .map(|(_, id)| id)
154 .collect::<HashSet<_>>();
155 let selected = points
156 .into_iter()
157 .filter(|point| !sample_ids.contains(&point.id))
158 .take(count)
159 .collect::<Vec<_>>();
160 if selected.len() != count {
161 return Err(format!(
162 "sample-stable mutation requested {count} points but only {} are outside the training sample",
163 selected.len()
164 )
165 .into());
166 }
167 Ok(selected)
168}
169
170fn uint_id_digest(id: &PointId) -> Result<[u8; 32], Box<dyn std::error::Error>> {
171 let PointId::UInt(value) = id else {
172 return Err("sample-stable profile expects generated unsigned IDs".into());
173 };
174 let mut bytes = [0_u8; 10];
175 bytes[..2].copy_from_slice(b"u\0");
176 bytes[2..].copy_from_slice(&value.to_be_bytes());
177 Ok(Sha256::digest(bytes).into())
178}
179
180fn validate(
181 repository: &Path,
182 build_report: &Path,
183 output: &Path,
184) -> Result<(), Box<dyn std::error::Error>> {
185 let build: serde_json::Value = serde_json::from_slice(&fs::read(build_report)?)?;
186 let root = build
187 .get("root")
188 .and_then(serde_json::Value::as_str)
189 .ok_or("build report root is missing")?;
190 let engine = SnapshotEngine::open(repository)?;
191 let started = Instant::now();
192 let report = engine.validate(root, true)?;
193 fs::write(
194 output,
195 serde_json::to_vec_pretty(&json!({
196 "schema_version": 1,
197 "root": root,
198 "validation_us": micros(started),
199 "report": report,
200 }))?,
201 )?;
202 Ok(())
203}
204
205fn build(input: &Path, repository: &Path, output: &Path) -> Result<(), Box<dyn std::error::Error>> {
206 if repository.exists() {
207 return Err(format!("repository already exists: {}", repository.display()).into());
208 }
209 let spec = read_spec(input)?;
210 let vectors = read_vectors(&spec.points_path, spec.point_count, spec.dimension)?;
211 let points = make_points(&vectors);
212 let config = CollectionConfig {
213 dimension: spec.dimension,
214 ..CollectionConfig::default()
215 };
216 let engine = SnapshotEngine::init(repository)?;
217 let started = Instant::now();
218 let snapshot = engine.build(config, points)?;
219 let build_us = micros(started);
220 fs::write(
221 output,
222 serde_json::to_vec_pretty(&json!({
223 "schema_version": 1,
224 "root": snapshot.root(),
225 "build_us": build_us,
226 "on_disk_bytes": directory_bytes(repository)?,
227 }))?,
228 )?;
229 Ok(())
230}
231
232fn query(
233 input: &Path,
234 repository: &Path,
235 build_report: &Path,
236 mode: &str,
237 output: &Path,
238) -> Result<(), Box<dyn std::error::Error>> {
239 let (exact, warm_exact) = match mode {
240 "exact" => (true, false),
241 "approximate" => (false, false),
242 "approximate-after-exact" => (false, true),
243 _ => return Err(format!("unsupported query mode: {mode}").into()),
244 };
245 let spec = read_spec(input)?;
246 let queries = read_vectors(&spec.queries_path, spec.query_count, spec.dimension)?;
247 let maximum_k = *spec.k.iter().max().ok_or("k must not be empty")?;
248 let build: serde_json::Value = serde_json::from_slice(&fs::read(build_report)?)?;
249 let root = build
250 .get("root")
251 .and_then(serde_json::Value::as_str)
252 .ok_or("build report root is missing")?;
253 let engine = SnapshotEngine::open(repository)?;
254 let snapshot = engine.open_snapshot(root)?;
255
256 let cache_build_us = if warm_exact {
257 let started = Instant::now();
258 snapshot.query(make_query(&queries[0], maximum_k, true))?;
259 Some(micros(started))
260 } else {
261 None
262 };
263 let warmup_started = Instant::now();
267 snapshot.query(make_query(&queries[0], maximum_k, exact))?;
268 let warmup_us = micros(warmup_started);
269 wait_for_profiler()?;
270
271 let mut query_us = Vec::with_capacity(queries.len());
272 let mut results = Vec::with_capacity(queries.len());
273 let mut vectors_scored = Vec::with_capacity(queries.len());
274 let batch_started = Instant::now();
275 for vector in &queries {
276 let started = Instant::now();
277 let result = snapshot.query(make_query(vector, maximum_k, exact))?;
278 query_us.push(micros(started));
279 vectors_scored.push(result.stats.vectors_scored);
280 results.push(result.points);
281 }
282 let batch_us = micros(batch_started);
283 fs::write(
284 output,
285 serde_json::to_vec_pretty(&json!({
286 "schema_version": 1,
287 "root": root,
288 "mode": mode,
289 "cache_build_us": cache_build_us,
290 "warmup_us": warmup_us,
291 "query_us": query_us,
292 "batch_us": batch_us,
293 "vectors_scored": vectors_scored,
294 "results": results,
295 }))?,
296 )?;
297 Ok(())
298}
299
300fn wait_for_profiler() -> Result<(), Box<dyn std::error::Error>> {
301 let Ok(ready_path) = env::var("GIT_VDB_PROFILE_READY") else {
302 return Ok(());
303 };
304 let go_path = env::var("GIT_VDB_PROFILE_GO")
305 .map_err(|_| "GIT_VDB_PROFILE_GO is required when profiler waiting is enabled")?;
306 fs::write(ready_path, std::process::id().to_string())?;
307 while !Path::new(&go_path).exists() {
308 thread::sleep(Duration::from_millis(10));
309 }
310 Ok(())
311}
312
313fn read_spec(path: &Path) -> Result<RunSpec, Box<dyn std::error::Error>> {
314 let spec: RunSpec = serde_json::from_slice(&fs::read(path)?)?;
315 if spec.schema_version != 1 {
316 return Err(format!("unsupported harness schema version {}", spec.schema_version).into());
317 }
318 Ok(spec)
319}
320
321fn make_query(vector: &[f32], limit: usize, exact: bool) -> Query {
322 Query {
323 vector: vector.to_vec(),
324 limit,
325 params: QueryParams {
326 exact: Some(exact),
327 ..QueryParams::default()
328 },
329 ..Query::default()
330 }
331}
332
333fn read_vectors(
334 path: &Path,
335 count: usize,
336 dimension: usize,
337) -> Result<Vec<Vec<f32>>, Box<dyn std::error::Error>> {
338 let bytes = fs::read(path)?;
339 let expected = count
340 .checked_mul(dimension)
341 .and_then(|components| components.checked_mul(4))
342 .ok_or("dataset size overflow")?;
343 if bytes.len() != expected {
344 return Err(format!(
345 "{} has {} bytes, expected {expected}",
346 path.display(),
347 bytes.len()
348 )
349 .into());
350 }
351 Ok(bytes
352 .chunks_exact(dimension * 4)
353 .map(|row| {
354 row.chunks_exact(4)
355 .map(|chunk| f32::from_le_bytes(chunk.try_into().unwrap()))
356 .collect()
357 })
358 .collect())
359}
360
361fn make_points(vectors: &[Vec<f32>]) -> Vec<Point> {
362 vectors
363 .iter()
364 .enumerate()
365 .map(|(id, vector)| {
366 let mut payload = Map::new();
367 payload.insert("selectivity_bucket".into(), json!(id % 1000));
368 Point {
369 id: (id as u64).into(),
370 vector: vector.clone(),
371 payload,
372 }
373 })
374 .collect()
375}
376
377fn directory_bytes(path: &Path) -> Result<u64, std::io::Error> {
378 let mut total = 0;
379 for entry in fs::read_dir(path)? {
380 let entry = entry?;
381 let metadata = entry.metadata()?;
382 total += if metadata.is_dir() {
383 directory_bytes(&entry.path())?
384 } else {
385 metadata.len()
386 };
387 }
388 Ok(total)
389}
390
391fn micros(started: Instant) -> u64 {
392 started.elapsed().as_micros().try_into().unwrap_or(u64::MAX)
393}