use crate::physical::types::{OperatorResult, PhysicalOperatorExec};
use akar_common::types::{PhysicalTypeID, Value};
use akar_common::vector::{DataChunk, ValueVector};
use akar_storage::table::TableCatalog;
use std::sync::Arc;
pub struct PhysicalCountRelTable {
pub table_name: String,
pub table_id: u64,
pub table_catalog: Option<Arc<TableCatalog>>,
}
impl PhysicalOperatorExec for PhysicalCountRelTable {
fn operator_type(&self) -> &str {
"count_rel_table"
}
fn execute(&self, _input: Vec<DataChunk>) -> OperatorResult {
let tc = self
.table_catalog
.as_ref()
.ok_or_else(|| "No table catalog for CountRelTable".to_string())?;
let count = if let Some(table) = tc.get_rel_table(self.table_id) {
table.num_rows as i64
} else {
0
};
let mut v = ValueVector::new(PhysicalTypeID::Int64, 1);
v.resize(1);
v.set_i64(0, count);
let arr = akar_common::arrow_vector::ArrowVector::from_legacy(&v).array;
Ok(vec![DataChunk::new(vec![arr], vec![PhysicalTypeID::Int64])])
}
}
pub struct PhysicalCreateFtsIndex {
pub index_name: String,
pub table_name: String,
pub column_name: String,
pub docs_table: String,
pub terms_table: String,
pub posting_table: String,
pub table_catalog: Arc<TableCatalog>,
}
impl PhysicalOperatorExec for PhysicalCreateFtsIndex {
fn operator_type(&self) -> &str {
"create_fts_index"
}
fn execute(&self, _input: Vec<DataChunk>) -> OperatorResult {
let source_table = match self.table_catalog.get_node_table_by_name(&self.table_name) {
Some(t) => t,
None => return Err(format!("Table '{}' not found", self.table_name).into()),
};
let col_idx = source_table
.columns
.iter()
.position(|c| c.name == self.column_name)
.ok_or_else(|| format!("Column '{}' not found in '{}'", self.column_name, self.table_name))?;
if self.table_catalog.get_node_table_by_name(&self.docs_table).is_none() {
let docs_cols = vec![
akar_storage::table::ColumnDefinition {
name: "doc_id".into(),
logical_type: akar_common::types::LogicalTypeID::Int64,
is_primary_key: true,
compression: akar_common::enums::CompressionType::Uncompressed,
},
akar_storage::table::ColumnDefinition {
name: "text".into(),
logical_type: akar_common::types::LogicalTypeID::String,
is_primary_key: false,
compression: akar_common::enums::CompressionType::Uncompressed,
},
];
self.table_catalog.create_node_table(self.docs_table.clone(), docs_cols);
}
if self.table_catalog.get_node_table_by_name(&self.terms_table).is_none() {
let terms_cols = vec![
akar_storage::table::ColumnDefinition {
name: "term_id".into(),
logical_type: akar_common::types::LogicalTypeID::Int64,
is_primary_key: true,
compression: akar_common::enums::CompressionType::Uncompressed,
},
akar_storage::table::ColumnDefinition {
name: "term".into(),
logical_type: akar_common::types::LogicalTypeID::String,
is_primary_key: false,
compression: akar_common::enums::CompressionType::Uncompressed,
},
akar_storage::table::ColumnDefinition {
name: "doc_freq".into(),
logical_type: akar_common::types::LogicalTypeID::Int64,
is_primary_key: false,
compression: akar_common::enums::CompressionType::Uncompressed,
},
];
self.table_catalog
.create_node_table(self.terms_table.clone(), terms_cols);
}
let source_data = source_table.to_column_major_data();
let num_rows = source_table.num_rows as usize;
let mut term_map: std::collections::HashMap<String, (i64, i64)> = std::collections::HashMap::new();
let mut doc_rows: Vec<Vec<Value>> = Vec::new();
let mut postings: Vec<(i64, i64, i64)> = Vec::new();
for row_idx in 0..num_rows {
let text = if let Some(col_data) = source_data.get(col_idx) {
if let Some(Value::String(s)) = col_data.get(row_idx) {
s.clone()
} else {
String::new()
}
} else {
String::new()
};
let doc_id = row_idx as i64;
doc_rows.push(vec![Value::Int64(doc_id), Value::String(text.clone())]);
let tokens = akar_fts::tokenize(&text);
let mut freq_map: std::collections::HashMap<String, i64> = std::collections::HashMap::new();
for token in tokens {
let stemmed = akar_fts::stem_word(&token);
if !akar_fts::STOP_WORDS.contains(&stemmed.as_str()) {
*freq_map.entry(stemmed).or_insert(0) += 1;
}
}
for (term, freq) in freq_map {
let next_id = term_map.len() as i64;
let (term_id, doc_freq) = term_map.entry(term).or_insert((next_id, 0));
*doc_freq += 1;
postings.push((*term_id, doc_id, freq));
}
}
{
let mut docs_table = self.table_catalog.get_node_table_by_name_mut(&self.docs_table).unwrap();
for row in doc_rows {
docs_table.insert_row(row)?;
}
}
if self.table_catalog.get_node_table_by_name(&self.terms_table).is_some() {
let mut terms_table = self
.table_catalog
.get_node_table_by_name_mut(&self.terms_table)
.unwrap();
let mut term_list: Vec<(String, i64, i64)> =
term_map.into_iter().map(|(t, (id, df))| (t, id, df)).collect();
term_list.sort_by_key(|(_, id, _)| *id);
for (term, term_id, doc_freq) in term_list {
terms_table.insert_row(vec![Value::Int64(term_id), Value::String(term), Value::Int64(doc_freq)])?;
}
}
let docs_table_id = self
.table_catalog
.get_node_table_by_name(&self.docs_table)
.unwrap()
.table_id;
let terms_table_id = self
.table_catalog
.get_node_table_by_name(&self.terms_table)
.unwrap()
.table_id;
if self.table_catalog.get_rel_table_by_name(&self.posting_table).is_none() {
let posting_cols = vec![akar_storage::table::ColumnDefinition {
name: "term_freq".into(),
logical_type: akar_common::types::LogicalTypeID::Int64,
is_primary_key: false,
compression: akar_common::enums::CompressionType::Uncompressed,
}];
self.table_catalog.create_rel_table(
self.posting_table.clone(),
terms_table_id,
docs_table_id,
posting_cols,
);
}
{
let mut posting_table = self
.table_catalog
.get_rel_table_by_name_mut(&self.posting_table)
.unwrap();
for (term_id, doc_id, freq) in postings {
posting_table.insert_rel(term_id as u64, doc_id as u64, vec![Value::Int64(freq)])?;
}
}
let mut result_vec = akar_common::vector::ValueVector::new(akar_common::types::PhysicalTypeID::String, 1);
result_vec.resize(1);
result_vec
.set_value(
0,
&Value::String(format!("FTS index '{}' built successfully.", self.index_name)),
)
.unwrap();
let arr = akar_common::arrow_vector::ArrowVector::from_legacy(&result_vec).array;
let mut result = DataChunk::new(vec![arr], vec![akar_common::types::PhysicalTypeID::String]);
result.size = 1;
result.field_names = vec!["result".to_string()];
Ok(vec![result])
}
}
#[derive(Debug, Clone)]
pub struct PhysicalFtsScan {
pub index_name: String,
pub query_string: String,
pub docs_table: String,
pub terms_table: String,
pub posting_table: String,
pub table_catalog: Arc<TableCatalog>,
}
impl PhysicalOperatorExec for PhysicalFtsScan {
fn operator_type(&self) -> &str {
"fts_scan"
}
fn execute(&self, _input: Vec<DataChunk>) -> OperatorResult {
let query_tokens: Vec<String> = akar_fts::tokenize(&self.query_string)
.into_iter()
.map(|t| akar_fts::stem_word(&t))
.filter(|t| !akar_fts::STOP_WORDS.contains(&t.as_str()))
.collect();
let terms_table = match self.table_catalog.get_node_table_by_name(&self.terms_table) {
Some(t) => t,
None => {
return Err(format!(
"FTS terms table '{}' not found. Has the index been created?",
self.terms_table
)
.into());
}
};
let num_docs = self
.table_catalog
.get_node_table_by_name(&self.docs_table)
.map(|t| t.num_rows as f64)
.unwrap_or(1.0);
let terms_data = terms_table.to_column_major_data();
let num_terms = terms_table.num_rows as usize;
let mut matching_terms: Vec<(i64, i64)> = Vec::new();
for row_idx in 0..num_terms {
let term_val = terms_data.get(1).and_then(|d| d.get(row_idx));
let term_str = if let Some(Value::String(s)) = term_val {
s.clone()
} else {
continue;
};
if query_tokens.contains(&term_str) {
let term_id = if let Some(Value::Int64(id)) = terms_data.first().and_then(|d| d.get(row_idx)) {
*id
} else {
continue;
};
let doc_freq = if let Some(Value::Int64(df)) = terms_data.get(2).and_then(|d| d.get(row_idx)) {
*df
} else {
0
};
matching_terms.push((term_id, doc_freq));
}
}
let mut doc_scores: std::collections::HashMap<i64, f64> = std::collections::HashMap::new();
if let Some(posting_table) = self.table_catalog.get_rel_table_by_name(&self.posting_table) {
for &(term_id, doc_freq) in &matching_terms {
let idf = ((num_docs - doc_freq as f64 + 0.5) / (doc_freq as f64 + 0.5) + 1.0).ln();
let posting_rels = posting_table.get_outgoing_edges(term_id as u64);
for (doc_id, rel_vals) in posting_rels {
let tf = if let Some(Value::Int64(freq)) = rel_vals.first() {
*freq as f64
} else {
1.0
};
let k1 = 1.5_f64;
let score = idf * (tf * (k1 + 1.0)) / (tf + k1);
*doc_scores.entry(doc_id as i64).or_insert(0.0) += score;
}
}
}
let mut ranked: Vec<(i64, f64)> = doc_scores.into_iter().collect();
ranked.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
let n = ranked.len();
let mut id_vec = akar_common::vector::ValueVector::new(akar_common::types::PhysicalTypeID::Int64, n);
let mut score_vec = akar_common::vector::ValueVector::new(akar_common::types::PhysicalTypeID::Double, n);
id_vec.resize(n);
score_vec.resize(n);
for (i, (doc_id, score)) in ranked.into_iter().enumerate() {
id_vec.set_i64(i, doc_id);
score_vec.set_double(i, score);
}
let arr1 = akar_common::arrow_vector::ArrowVector::from_legacy(&id_vec).array;
let arr2 = akar_common::arrow_vector::ArrowVector::from_legacy(&score_vec).array;
let mut chunk = DataChunk::new(
vec![arr1, arr2],
vec![
akar_common::types::PhysicalTypeID::Int64,
akar_common::types::PhysicalTypeID::Double,
],
);
chunk.size = n;
chunk.field_names = vec!["doc_id".to_string(), "score".to_string()];
Ok(vec![chunk])
}
}