aboutsummaryrefslogtreecommitdiffstats
diff options
context:
space:
mode:
authorJan Tuomi <jan@jantuomi.fi>2024-09-29 13:14:16 +0300
committerJan Tuomi <jan@jantuomi.fi>2024-10-01 21:01:37 +0300
commit1570c6e33e100701f568c57c4944c02bd43cca14 (patch)
tree096439160ad3a5a6783a640d6342b64fd4e0f6f3
parent559c9a7346a263db2fde5abb4820fed4af25c385 (diff)
Implement basic schema functionality, upsert
-rw-r--r--src/lib.rs281
-rw-r--r--tests/integration.rs120
2 files changed, 356 insertions, 45 deletions
diff --git a/src/lib.rs b/src/lib.rs
index 6a1e1d0..72d216b 100644
--- a/src/lib.rs
+++ b/src/lib.rs
@@ -3,7 +3,7 @@ extern crate log;
pub mod log_db {
use fs2::FileExt;
use priority_queue::PriorityQueue;
- use std::collections::BTreeMap;
+ use std::collections::{BTreeMap, HashSet};
use std::fs;
use std::io;
use std::io::Write;
@@ -11,12 +11,105 @@ pub mod log_db {
const ACTIVE_LOG_FILENAME: &str = "db";
- pub trait Record {
- fn serialize(&self) -> Vec<u8>;
- fn deserialize(data: Vec<u8>) -> Self;
+ #[derive(Debug, Clone, Ord, PartialOrd, Eq, PartialEq)]
+ pub enum IndexableValue {
+ Int(i64),
+ String(String),
}
- pub struct Config {
+ #[derive(Debug, Clone)]
+ pub enum RecordFieldType {
+ Int,
+ Float,
+ String,
+ Bytes,
+ }
+
+ #[derive(Debug, Clone)]
+ pub enum RecordValue {
+ Null,
+ Int(i64),
+ Float(f64),
+ String(String),
+ Bytes(Vec<u8>),
+ }
+
+ impl RecordValue {
+ fn serialize(&self) -> Vec<u8> {
+ match self {
+ RecordValue::Null => {
+ vec![0] // Tag for Null
+ }
+ RecordValue::Int(i) => {
+ let mut bytes = vec![1]; // Tag for Int
+ bytes.extend(&i.to_be_bytes());
+ bytes
+ }
+ RecordValue::Float(f) => {
+ let mut bytes = vec![2]; // Tag for Float
+ bytes.extend(&f.to_be_bytes());
+ bytes
+ }
+ RecordValue::String(s) => {
+ let mut bytes = vec![3]; // Tag for String
+ let length = s.len() as u64;
+ bytes.extend(&length.to_be_bytes());
+ bytes.extend(s.as_bytes());
+ bytes
+ }
+ RecordValue::Bytes(b) => {
+ let mut bytes = vec![3]; // Tag for Bytes
+ let length = b.len() as u64;
+ bytes.extend(&length.to_be_bytes());
+ bytes.extend(b);
+ bytes
+ }
+ }
+ }
+
+ fn deserialize(bytes: &[u8]) -> RecordValue {
+ match bytes[0] {
+ 0 => RecordValue::Null,
+ 1 => {
+ let mut int_bytes = [0; 8];
+ int_bytes.copy_from_slice(&bytes[1..9]);
+ RecordValue::Int(i64::from_be_bytes(int_bytes))
+ }
+ 2 => {
+ let mut float_bytes = [0; 8];
+ float_bytes.copy_from_slice(&bytes[1..9]);
+ RecordValue::Float(f64::from_be_bytes(float_bytes))
+ }
+ 3 => {
+ let length_bytes = &bytes[1..9];
+ let length = u64::from_be_bytes(length_bytes.try_into().unwrap()) as usize;
+ RecordValue::String(String::from_utf8(bytes[9..9 + length].to_vec()).unwrap())
+ }
+ 4 => {
+ let length_bytes = &bytes[1..9];
+ let length = u64::from_be_bytes(length_bytes.try_into().unwrap()) as usize;
+ RecordValue::Bytes(bytes[9..9 + length].to_vec())
+ }
+ _ => panic!("Invalid tag"),
+ }
+ }
+
+ fn as_indexable(&self) -> Option<IndexableValue> {
+ match self {
+ RecordValue::Int(i) => Some(IndexableValue::Int(*i)),
+ RecordValue::String(s) => Some(IndexableValue::String(s.clone())),
+ _ => None,
+ }
+ }
+ }
+
+ #[derive(Debug, Clone)]
+ pub struct Record {
+ pub values: Vec<RecordValue>,
+ }
+
+ #[derive(Clone)]
+ pub struct Config<Field: Eq + Clone> {
/// Directory where the database will store its data.
pub data_dir: String,
/// The maximum size of a segment file in bytes.
@@ -26,19 +119,26 @@ pub mod log_db {
/// The maximum size of a single memtable in bytes.
/// Note that each secondary index will have its own memtable.
pub memtable_size: u64,
+ /// The field schema of the database.
+ pub fields: Vec<(Field, RecordFieldType)>,
+ /// The primary key of the database, used to construct
+ /// the primary memtable index. This should be the field
+ /// that is most frequently queried.
+ pub primary_key: Field,
+ /// The secondary keys of the database, used to construct
+ /// the secondary memtable indexes.
+ pub secondary_keys: Vec<Field>,
}
- pub struct DB<T: Record> {
- data_dir: String,
- segment_size: u64,
- memtable_size: u64,
-
- pub log_path: PathBuf,
- primary_memtable: BTreeMap<u64, T>,
+ pub struct DB<Field: Eq + Clone> {
+ config: Config<Field>,
+ log_path: PathBuf,
+ primary_memtable: BTreeMap<IndexableValue, Record>,
+ secondary_memtables: Vec<BTreeMap<IndexableValue, HashSet<Record>>>,
}
- impl<T: Record> DB<T> {
- pub fn initialize(config: &Config) -> Result<DB<T>, io::Error> {
+ impl<Field: Eq + Clone> DB<Field> {
+ pub fn initialize(config: &Config<Field>) -> Result<DB<Field>, io::Error> {
info!("Initializing DB");
// If data_dir does not exist, create it
if !fs::exists(&config.data_dir)? {
@@ -47,18 +147,89 @@ pub mod log_db {
let log_path = Path::new(&config.data_dir).join(ACTIVE_LOG_FILENAME);
- let db = DB::<T> {
- data_dir: config.data_dir.clone(),
- segment_size: config.segment_size,
- memtable_size: config.memtable_size,
+ // Create the log file if it does not exist
+ let _file = fs::OpenOptions::new()
+ .create(true)
+ .append(true)
+ .open(&log_path)?;
+
+ // Join primary key and secondary keys vec into a single vec
+ let mut all_keys = vec![&config.primary_key];
+ all_keys.extend(&config.secondary_keys);
+
+ // If any of the keys is not in the schema or
+ // is not an IndexableValue, return an error
+ for &key in &all_keys {
+ let (_, field_type) =
+ config
+ .fields
+ .iter()
+ .find(|(field, _)| field == key)
+ .ok_or(io::Error::new(
+ io::ErrorKind::InvalidInput,
+ "Secondary key must be present in the field schema",
+ ))?;
+ match field_type {
+ RecordFieldType::Int | RecordFieldType::String => {}
+ _ => {
+ return Err(io::Error::new(
+ io::ErrorKind::InvalidInput,
+ "Secondary key must be an IndexableValue",
+ ))
+ }
+ }
+ }
+ let primary_memtable = BTreeMap::<IndexableValue, Record>::new();
+ let secondary_memtables = config
+ .secondary_keys
+ .iter()
+ .map(|_| BTreeMap::<IndexableValue, HashSet<Record>>::new())
+ .collect();
+
+ let db = DB::<Field> {
+ config: config.clone(),
log_path,
- primary_memtable: BTreeMap::new(),
+ primary_memtable,
+ secondary_memtables,
};
Ok(db)
}
- pub fn upsert(&self, record: &T) -> Result<(), io::Error> {
+ pub fn upsert(&mut self, record: &Record) -> Result<(), io::Error> {
+ // Validate the record length
+ if record.values.len() != self.config.fields.len() {
+ return Err(io::Error::new(
+ io::ErrorKind::InvalidInput,
+ format!(
+ "Record has an incorrect number of fields: {}, expected {}",
+ record.values.len(),
+ self.config.fields.len()
+ ),
+ ));
+ }
+
+ // Validate that record fields match schema types
+ // TODO: handle Null
+ for (i, (_, field_type)) in self.config.fields.iter().enumerate() {
+ match (&record.values[i], field_type) {
+ (RecordValue::Int(_), RecordFieldType::Int) => {}
+ (RecordValue::Float(_), RecordFieldType::Float) => {}
+ (RecordValue::String(_), RecordFieldType::String) => {}
+ (RecordValue::Bytes(_), RecordFieldType::Bytes) => {}
+ _ => {
+ return Err(io::Error::new(
+ io::ErrorKind::InvalidInput,
+ format!(
+ "Record field {} has incorrect type: {:?}, expected {:?}",
+ i, record.values[i], field_type
+ ),
+ ))
+ }
+ }
+ }
+
+ // Open the log file in append mode
let mut file = fs::OpenOptions::new()
.create(true)
.append(true)
@@ -68,7 +239,13 @@ pub mod log_db {
file.lock_exclusive()?;
// Write the record to the log
- let serialized = record.serialize();
+ // Each serialized row is suffixed with the field separator character
+ let mut serialized = record
+ .values
+ .iter()
+ .flat_map(|field| field.serialize())
+ .collect::<Vec<u8>>();
+ serialized.extend(vec![b'\x1C']);
file.write_all(&serialized)?;
// Sync to disk
@@ -77,7 +254,69 @@ pub mod log_db {
file.unlock()?;
+ // Update the primary memtable
+ let primary_key_index = self
+ .config
+ .fields
+ .iter()
+ .position(|(field, _)| field == &self.config.primary_key)
+ .ok_or(io::Error::new(
+ io::ErrorKind::InvalidInput,
+ "Primary key not found in schema after initialize",
+ ))?;
+ let primary_value =
+ &record.values[primary_key_index]
+ .as_indexable()
+ .ok_or(io::Error::new(
+ io::ErrorKind::InvalidInput,
+ "Primary key must be an IndexableValue",
+ ))?;
+
+ self.primary_memtable
+ .insert(primary_value.clone(), record.clone());
+
+ // TODO: Update secondary memtables
+
Ok(())
}
+
+ pub fn get(&self, field: Field, key: RecordValue) -> Result<Option<Record>, io::Error> {
+ // If the requested field is the primary key, look up the value in the primary memtable
+ if field == self.config.primary_key {
+ let primary_value = key.as_indexable().ok_or(io::Error::new(
+ io::ErrorKind::InvalidInput,
+ "Primary key must be an IndexableValue",
+ ))?;
+
+ let found = self.primary_memtable.get(&primary_value);
+ if let Some(record) = found {
+ return Ok(Some(record.clone()));
+ }
+ }
+
+ // TODO: query secondary memtables
+
+ // Get the index of the requested field
+ let key_index = self
+ .config
+ .fields
+ .iter()
+ .position(|(schema_field, _)| schema_field == &field)
+ .ok_or(io::Error::new(
+ io::ErrorKind::InvalidInput,
+ "Key not found in schema after initialize",
+ ))?;
+
+ // Open the file and acquire a shared lock for reading
+ let file = fs::OpenOptions::new().read(true).open(&self.log_path)?;
+ file.lock_shared()?;
+
+ // Do a log file scan from the bottom up to find the record
+ // TODO
+
+ file.unlock()?;
+
+ Ok(None)
+ }
}
}
diff --git a/tests/integration.rs b/tests/integration.rs
index e500219..2c0a958 100644
--- a/tests/integration.rs
+++ b/tests/integration.rs
@@ -1,37 +1,31 @@
-use log_db::log_db;
+use log_db::log_db::{Config, Record, RecordFieldType, RecordValue, DB};
use serial_test::serial;
const TEST_DATA_DIR: &str = "test_db_data";
const TEST_SEGMENT_SIZE: u64 = 1024 * 1024; // 1 MB
const TEST_MEMTABLE_SIZE: u64 = 1024 * 1024; // 1 MB
-struct TestRecord {
- id: u64,
- data: String,
-}
-
-impl log_db::Record for TestRecord {
- fn serialize(&self) -> Vec<u8> {
- format!("{}:{}", self.id, self.data).into_bytes()
- }
-
- fn deserialize(data: Vec<u8>) -> Self {
- let data = String::from_utf8(data).unwrap();
- let parts: Vec<&str> = data.split(':').collect();
- TestRecord {
- id: parts[0].parse().unwrap(),
- data: parts[1].to_string(),
- }
- }
+#[derive(Eq, PartialEq, Clone)]
+enum Field {
+ Id,
+ Name,
+ Data,
}
#[test]
#[serial]
fn test_initialize() {
- let _db: log_db::DB<TestRecord> = log_db::DB::initialize(&log_db::Config {
+ let _db = DB::initialize(&Config {
data_dir: TEST_DATA_DIR.to_string(),
segment_size: TEST_SEGMENT_SIZE,
memtable_size: TEST_MEMTABLE_SIZE,
+ fields: vec![
+ (Field::Id, RecordFieldType::Int),
+ (Field::Name, RecordFieldType::String),
+ (Field::Data, RecordFieldType::Bytes),
+ ],
+ primary_key: Field::Id,
+ secondary_keys: vec![],
})
.unwrap();
@@ -42,19 +36,97 @@ fn test_initialize() {
#[test]
#[serial]
fn test_upsert_to_empty_db() {
- let db: log_db::DB<TestRecord> = log_db::DB::initialize(&log_db::Config {
+ let mut db = DB::initialize(&Config {
data_dir: TEST_DATA_DIR.to_string(),
segment_size: TEST_SEGMENT_SIZE,
memtable_size: TEST_MEMTABLE_SIZE,
+ fields: vec![
+ (Field::Id, RecordFieldType::Int),
+ (Field::Name, RecordFieldType::String),
+ (Field::Data, RecordFieldType::Bytes),
+ ],
+ primary_key: Field::Id,
+ secondary_keys: vec![],
})
.unwrap();
- let record = TestRecord {
- id: 1,
- data: "hello".to_string(),
+ let record = Record {
+ values: vec![
+ RecordValue::Int(1),
+ RecordValue::String("Alice".to_string()),
+ RecordValue::Bytes(vec![0, 1, 2]),
+ ],
};
db.upsert(&record).unwrap();
+ let result = db.get(Field::Id, RecordValue::Int(1)).unwrap().unwrap();
+
+ // Check that the IDs match
+ assert!(match (&result.values[0], &record.values[0]) {
+ (RecordValue::Int(a), RecordValue::Int(b)) => a == b,
+ _ => false,
+ });
+
+ // Clean up
+ std::fs::remove_dir_all(TEST_DATA_DIR.to_string()).unwrap();
+}
+
+#[test]
+#[serial]
+fn test_upsert_fails_on_invalid_number_of_values() {
+ let mut db = DB::initialize(&Config {
+ data_dir: TEST_DATA_DIR.to_string(),
+ segment_size: TEST_SEGMENT_SIZE,
+ memtable_size: TEST_MEMTABLE_SIZE,
+ fields: vec![
+ (Field::Id, RecordFieldType::Int),
+ (Field::Name, RecordFieldType::String),
+ (Field::Data, RecordFieldType::Bytes),
+ ],
+ primary_key: Field::Id,
+ secondary_keys: vec![],
+ })
+ .unwrap();
+
+ let record = Record {
+ // Missing primary key
+ values: vec![
+ RecordValue::String("Alice".to_string()),
+ RecordValue::Bytes(vec![0, 1, 2]),
+ ],
+ };
+ assert!(db.upsert(&record).is_err());
+
+ // Clean up
+ std::fs::remove_dir_all(TEST_DATA_DIR.to_string()).unwrap();
+}
+
+#[test]
+#[serial]
+fn test_upsert_fails_on_invalid_value_type() {
+ let mut db = DB::initialize(&Config {
+ data_dir: TEST_DATA_DIR.to_string(),
+ segment_size: TEST_SEGMENT_SIZE,
+ memtable_size: TEST_MEMTABLE_SIZE,
+ fields: vec![
+ (Field::Id, RecordFieldType::Int),
+ (Field::Name, RecordFieldType::String),
+ (Field::Data, RecordFieldType::Bytes),
+ ],
+ primary_key: Field::Id,
+ secondary_keys: vec![],
+ })
+ .unwrap();
+
+ let record = Record {
+ values: vec![
+ RecordValue::String("foo".to_string()),
+ RecordValue::String("bar".to_string()),
+ RecordValue::String("baz".to_string()),
+ ],
+ };
+ assert!(db.upsert(&record).is_err());
+
// Clean up
std::fs::remove_dir_all(TEST_DATA_DIR.to_string()).unwrap();
}