Skip to main content

lancedb_git_vdb_profile/
lancedb_git_vdb_profile.rs

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    // Fill the immutable snapshot cache, construct an approximate lookup over an
264    // already-warm exact view, or warm the unchanged approximate ODB path without
265    // including that one-time work in the samples.
266    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}