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