use anyhow::Result;
use std::sync::Arc;
use arrow_array::{Array, Int64Array, RecordBatch, StringArray};
use arrow_schema::{DataType, Field, Schema};
use futures::TryStreamExt;
use lancedb::{
query::{ExecutableQuery, QueryBase, Select},
Connection,
};
use crate::store::sql::{escape_single_quotes, in_list_predicate};
use crate::store::table_ops::{TableOperations, DELETE_PATHS_CHUNK};
use crate::store::tables;
pub struct MetadataOperations<'a> {
pub db: &'a Connection,
pub table_ops: TableOperations<'a>,
}
impl<'a> MetadataOperations<'a> {
pub fn new(db: &'a Connection) -> Self {
Self {
db,
table_ops: TableOperations::new(db),
}
}
fn commit_marker_schema() -> Arc<Schema> {
Arc::new(Schema::new(vec![
Field::new("commit_hash", DataType::Utf8, false),
Field::new("indexed_at", DataType::Int64, false),
]))
}
async fn get_commit_marker(&self, table_name: &str) -> Result<Option<String>> {
if !self.table_ops.table_exists(table_name).await? {
return Ok(None);
}
let table = self.db.open_table(table_name).execute().await?;
let mut results = table
.query()
.select(Select::Columns(vec!["commit_hash".to_string()]))
.limit(1)
.execute()
.await?;
while let Some(batch) = results.try_next().await? {
if batch.num_rows() > 0 {
if let Some(column) = batch.column_by_name("commit_hash") {
if let Some(hash_array) = column.as_any().downcast_ref::<StringArray>() {
if let Some(hash) = hash_array.iter().next() {
return Ok(hash.map(|s| s.to_string()));
}
}
}
}
}
Ok(None)
}
async fn store_commit_marker(&self, table_name: &str, commit_hash: &str) -> Result<()> {
if let Ok(Some(existing_hash)) = self.get_commit_marker(table_name).await {
if existing_hash == commit_hash {
return Ok(());
}
}
let batch = RecordBatch::try_new(
Self::commit_marker_schema(),
vec![
Arc::new(StringArray::from(vec![commit_hash])),
Arc::new(Int64Array::from(vec![chrono::Utc::now().timestamp()])),
],
)?;
self.table_ops.clear_table(table_name).await?;
self.table_ops.store_batch(table_name, batch).await?;
Ok(())
}
pub async fn store_git_metadata(&self, commit_hash: &str) -> Result<()> {
self.store_commit_marker(tables::GIT_METADATA, commit_hash)
.await
}
pub async fn get_last_commit_hash(&self) -> Result<Option<String>> {
self.get_commit_marker(tables::GIT_METADATA).await
}
pub async fn clear_git_metadata(&self) -> Result<()> {
self.table_ops.clear_table(tables::GIT_METADATA).await
}
pub async fn get_graphrag_last_commit_hash(&self) -> Result<Option<String>> {
self.get_commit_marker(tables::GRAPHRAG_GIT_METADATA).await
}
pub async fn store_graphrag_commit_hash(&self, commit_hash: &str) -> Result<()> {
self.store_commit_marker(tables::GRAPHRAG_GIT_METADATA, commit_hash)
.await
}
pub async fn get_commits_last_commit_hash(&self) -> Result<Option<String>> {
self.get_commit_marker(tables::COMMITS_GIT_METADATA).await
}
pub async fn store_commits_last_commit_hash(&self, commit_hash: &str) -> Result<()> {
self.store_commit_marker(tables::COMMITS_GIT_METADATA, commit_hash)
.await
}
fn file_metadata_schema() -> Arc<Schema> {
Arc::new(Schema::new(vec![
Field::new("path", DataType::Utf8, false),
Field::new("mtime", DataType::Int64, false),
Field::new("indexed_at", DataType::Int64, false),
]))
}
fn file_metadata_record_batch(
entries: &[(String, u64)],
indexed_at: i64,
) -> Result<RecordBatch> {
let paths = entries
.iter()
.map(|(path, _)| path.as_str())
.collect::<Vec<_>>();
let mtimes = entries
.iter()
.map(|(_, mtime)| *mtime as i64)
.collect::<Vec<_>>();
Ok(RecordBatch::try_new(
Self::file_metadata_schema(),
vec![
Arc::new(StringArray::from(paths)),
Arc::new(Int64Array::from(mtimes)),
Arc::new(Int64Array::from(vec![indexed_at; entries.len()])),
],
)?)
}
async fn create_file_metadata_table(&self) -> Result<()> {
self.table_ops
.create_table_with_schema(tables::FILE_METADATA, Self::file_metadata_schema())
.await
}
pub async fn store_file_metadata_batch(&self, entries: &[(String, u64)]) -> Result<()> {
if entries.is_empty() {
return Ok(());
}
if !self.table_ops.table_exists(tables::FILE_METADATA).await? {
self.create_file_metadata_table().await?;
}
let table = self.db.open_table(tables::FILE_METADATA).execute().await?;
let mut existing = Vec::new();
for chunk in entries.chunks(DELETE_PATHS_CHUNK) {
let paths = chunk
.iter()
.map(|(path, _)| path.clone())
.collect::<Vec<_>>();
let mut results = table
.query()
.only_if(in_list_predicate("path", &paths))
.select(Select::Columns(vec!["path".to_string()]))
.execute()
.await?;
while let Some(batch) = results.try_next().await? {
if let Some(path_array) = batch
.column_by_name("path")
.and_then(|column| column.as_any().downcast_ref::<StringArray>())
{
existing.extend(path_array.iter().flatten().map(str::to_string));
}
}
}
if !existing.is_empty() {
self.table_ops
.remove_blocks_by_paths(&existing, tables::FILE_METADATA)
.await?;
}
let batch = Self::file_metadata_record_batch(entries, chrono::Utc::now().timestamp())?;
self.table_ops
.store_batch(tables::FILE_METADATA, batch)
.await
}
pub async fn store_file_metadata(&self, file_path: &str, mtime: u64) -> Result<()> {
self.store_file_metadata_batch(std::slice::from_ref(&(file_path.to_string(), mtime)))
.await
}
pub async fn get_file_mtime(&self, file_path: &str) -> Result<Option<u64>> {
if !self.table_ops.table_exists(tables::FILE_METADATA).await? {
return Ok(None);
}
let table = self.db.open_table(tables::FILE_METADATA).execute().await?;
let mut results = table
.query()
.only_if(format!("path = '{}'", escape_single_quotes(file_path)))
.select(Select::Columns(vec!["mtime".to_string()]))
.limit(1)
.execute()
.await?;
while let Some(batch) = results.try_next().await? {
if batch.num_rows() > 0 {
if let Some(column) = batch.column_by_name("mtime") {
if let Some(mtime_array) = column.as_any().downcast_ref::<Int64Array>() {
if let Some(mtime) = mtime_array.iter().next() {
return Ok(mtime.map(|t| t as u64));
}
}
}
}
}
Ok(None)
}
pub async fn get_all_file_metadata(&self) -> Result<std::collections::HashMap<String, u64>> {
let mut metadata_map = std::collections::HashMap::new();
if !self.table_ops.table_exists(tables::FILE_METADATA).await? {
return Ok(metadata_map);
}
let table = self.db.open_table(tables::FILE_METADATA).execute().await?;
let mut results = table
.query()
.select(Select::Columns(vec![
"path".to_string(),
"mtime".to_string(),
]))
.execute()
.await?;
while let Some(batch) = results.try_next().await? {
if batch.num_rows() > 0 {
if let (Some(path_column), Some(mtime_column)) =
(batch.column_by_name("path"), batch.column_by_name("mtime"))
{
if let (Some(path_array), Some(mtime_array)) = (
path_column.as_any().downcast_ref::<StringArray>(),
mtime_column.as_any().downcast_ref::<Int64Array>(),
) {
for i in 0..path_array.len() {
if path_array.is_null(i) || mtime_array.is_null(i) {
continue;
}
metadata_map.insert(
path_array.value(i).to_string(),
mtime_array.value(i) as u64,
);
}
}
}
}
}
Ok(metadata_map)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn file_metadata_record_batch_round_trips_values() {
let entries = vec![
("src/main.rs".to_string(), 101),
("src/store/mod.rs".to_string(), 202),
("README.md".to_string(), 303),
];
let indexed_at = 1_234_567_890;
let batch = MetadataOperations::file_metadata_record_batch(&entries, indexed_at).unwrap();
assert_eq!(batch.num_rows(), 3);
let schema = batch.schema();
assert_eq!(schema.field(0).name(), "path");
assert_eq!(schema.field(1).name(), "mtime");
assert_eq!(schema.field(2).name(), "indexed_at");
let paths = batch
.column_by_name("path")
.unwrap()
.as_any()
.downcast_ref::<StringArray>()
.unwrap();
let mtimes = batch
.column_by_name("mtime")
.unwrap()
.as_any()
.downcast_ref::<Int64Array>()
.unwrap();
let indexed_ats = batch
.column_by_name("indexed_at")
.unwrap()
.as_any()
.downcast_ref::<Int64Array>()
.unwrap();
assert_eq!(
paths.iter().flatten().collect::<Vec<_>>(),
vec!["src/main.rs", "src/store/mod.rs", "README.md",]
);
assert_eq!(mtimes.values(), &[101, 202, 303]);
assert_eq!(indexed_ats.values(), &[indexed_at; 3]);
}
}