use std::{
collections::HashMap,
fs::{File, OpenOptions},
io::{Read, Seek, SeekFrom, Write},
path::Path,
sync::Mutex,
};
pub use kcode_k1_access_types::{Authority, GroupId, TxId, UserId};
pub use kcode_k1_kmap_format::NodeId;
const HEADER: &[u8; 8] = b"K1LNV2\0\0";
const USER_TAG: u8 = 1;
const GROUP_TAG: u8 = 2;
const MAX_TARGET_BYTES: usize = 320;
const RECORD_PREFIX: usize = 15;
const RECORD_SUFFIX: usize = 16;
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
pub struct TargetName(String);
impl TargetName {
pub fn new(value: String) -> Result<Self, String> {
let characters = value.chars().count();
if !(1..=80).contains(&characters) {
return Err("target name must contain 1 through 80 Unicode characters".into());
}
if value.chars().any(char::is_control) {
return Err("target name must not contain control characters".into());
}
if !value.chars().any(|character| !character.is_whitespace()) {
return Err("target name must contain a non-whitespace character".into());
}
if value.len() > MAX_TARGET_BYTES {
return Err("target name exceeds its canonical UTF-8 bound".into());
}
Ok(Self(value))
}
pub fn as_str(&self) -> &str {
&self.0
}
pub fn into_string(self) -> String {
self.0
}
}
#[derive(Clone, Debug, Eq, Hash, PartialEq)]
pub struct TargetId {
authority: Authority,
target: TargetName,
}
impl TargetId {
pub fn new(authority: Authority, target: TargetName) -> Self {
Self { authority, target }
}
pub fn authority(&self) -> Authority {
self.authority
}
pub fn target(&self) -> &TargetName {
&self.target
}
}
struct Inner {
file: File,
bindings: HashMap<TargetId, NodeId>,
failed: bool,
}
pub struct LaunchNodes {
inner: Mutex<Inner>,
}
impl LaunchNodes {
pub fn open(path: &Path) -> Result<Self, String> {
let created = !path.exists();
let mut file = OpenOptions::new()
.create(true)
.truncate(false)
.read(true)
.write(true)
.open(path)
.map_err(|error| format!("open launch nodes file: {error}"))?;
if created {
file.write_all(HEADER)
.and_then(|()| file.sync_data())
.map_err(|error| format!("initialize launch nodes file: {error}"))?;
sync_parent(path)?;
}
file.seek(SeekFrom::Start(0))
.map_err(|error| format!("seek launch nodes file: {error}"))?;
let mut bytes = Vec::new();
file.read_to_end(&mut bytes)
.map_err(|error| format!("read launch nodes file: {error}"))?;
if bytes.len() < HEADER.len() || &bytes[..HEADER.len()] != HEADER {
return Err("unsupported launch nodes file format".into());
}
let mut bindings = HashMap::new();
let mut offset = HEADER.len();
while offset < bytes.len() {
let remaining = bytes.len() - offset;
if remaining < RECORD_PREFIX {
repair_tail(&mut file, offset)?;
break;
}
let tag = bytes[offset];
if tag != USER_TAG && tag != GROUP_TAG {
return Err("launch nodes record has an invalid authority tag".into());
}
let name_len = u16::from_le_bytes([bytes[offset + 13], bytes[offset + 14]]) as usize;
if name_len > MAX_TARGET_BYTES {
return Err("launch nodes record has an invalid target length".into());
}
let record_len = RECORD_PREFIX + name_len + RECORD_SUFFIX;
if remaining < record_len {
repair_tail(&mut file, offset)?;
break;
}
let record = &bytes[offset..offset + record_len];
let checksum_offset = record_len - 4;
let expected = u32::from_le_bytes(
record[checksum_offset..]
.try_into()
.map_err(|_| "launch nodes checksum is unavailable".to_owned())?,
);
if checksum(&record[..checksum_offset]) != expected {
return Err("launch nodes record checksum is invalid".into());
}
let authority_bytes: [u8; 12] = record[1..13]
.try_into()
.map_err(|_| "launch nodes authority is invalid".to_owned())?;
let authority_txid = TxId::from_bytes(authority_bytes);
let authority = match tag {
USER_TAG => Authority::User(UserId::from_tx_id(authority_txid)),
GROUP_TAG => Authority::Group(GroupId::new(authority_txid)),
_ => unreachable!(),
};
let name_start = RECORD_PREFIX;
let name_end = name_start + name_len;
let name = String::from_utf8(record[name_start..name_end].to_vec())
.map_err(|_| "launch nodes target is not valid UTF-8".to_owned())?;
let target = TargetName::new(name)
.map_err(|error| format!("launch nodes target is invalid: {error}"))?;
let node_bytes: [u8; 12] = record[name_end..name_end + 12]
.try_into()
.map_err(|_| "launch nodes node ID is invalid".to_owned())?;
bindings.insert(TargetId::new(authority, target), NodeId(node_bytes));
offset += record_len;
}
file.seek(SeekFrom::End(0))
.map_err(|error| format!("seek launch nodes append position: {error}"))?;
Ok(Self {
inner: Mutex::new(Inner {
file,
bindings,
failed: false,
}),
})
}
pub fn get(&self, target: &TargetId) -> Result<Option<NodeId>, String> {
let inner = self
.inner
.lock()
.map_err(|_| "launch nodes lock poisoned".to_owned())?;
if inner.failed {
return Err("launch nodes store is unavailable".into());
}
Ok(inner.bindings.get(target).copied())
}
pub fn set(&self, target: TargetId, node: NodeId) -> Result<(), String> {
let mut inner = self
.inner
.lock()
.map_err(|_| "launch nodes lock poisoned".to_owned())?;
if inner.failed {
return Err("launch nodes store is unavailable".into());
}
let record = encode_record(&target, node);
let result = inner
.file
.seek(SeekFrom::End(0))
.and_then(|_| inner.file.write_all(&record))
.and_then(|()| inner.file.sync_data());
if let Err(error) = result {
inner.failed = true;
return Err(format!("append launch nodes record: {error}"));
}
inner.bindings.insert(target, node);
Ok(())
}
}
fn encode_record(target: &TargetId, node: NodeId) -> Vec<u8> {
let (tag, authority) = match target.authority {
Authority::User(user) => (USER_TAG, *user.as_tx_id().as_bytes()),
Authority::Group(group) => (GROUP_TAG, *group.txid().as_bytes()),
};
let name = target.target.as_str().as_bytes();
let mut record = Vec::with_capacity(RECORD_PREFIX + name.len() + RECORD_SUFFIX);
record.push(tag);
record.extend_from_slice(&authority);
record.extend_from_slice(&(name.len() as u16).to_le_bytes());
record.extend_from_slice(name);
record.extend_from_slice(&node.0);
let checksum = checksum(&record);
record.extend_from_slice(&checksum.to_le_bytes());
record
}
fn checksum(bytes: &[u8]) -> u32 {
bytes.iter().fold(0x811c_9dc5, |value, byte| {
(value ^ u32::from(*byte)).wrapping_mul(0x0100_0193)
})
}
fn repair_tail(file: &mut File, offset: usize) -> Result<(), String> {
file.set_len(offset as u64)
.and_then(|()| file.sync_data())
.map_err(|error| format!("repair launch nodes incomplete tail: {error}"))
}
fn sync_parent(path: &Path) -> Result<(), String> {
let parent = path
.parent()
.filter(|parent| !parent.as_os_str().is_empty())
.unwrap_or_else(|| Path::new("."));
File::open(parent)
.and_then(|directory| directory.sync_all())
.map_err(|error| format!("synchronize launch nodes parent directory: {error}"))
}
#[cfg(test)]
mod tests {
use super::*;
use std::{
fs,
path::PathBuf,
time::{SystemTime, UNIX_EPOCH},
};
fn path(label: &str) -> PathBuf {
let nonce = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_nanos();
std::env::temp_dir().join(format!(
"k1-launch-nodes-{label}-{}-{nonce}",
std::process::id()
))
}
fn tx(byte: u8) -> TxId {
TxId::from_bytes([byte; 12])
}
fn user(byte: u8) -> Authority {
Authority::User(UserId::from_tx_id(tx(byte)))
}
fn group(byte: u8) -> Authority {
Authority::Group(GroupId::new(tx(byte)))
}
fn target(authority: Authority, name: &str) -> TargetId {
TargetId::new(authority, TargetName::new(name.to_owned()).unwrap())
}
#[test]
fn target_names_preserve_exact_slash_and_unicode_boundaries() {
let exact = "界".repeat(80);
assert_eq!(
TargetName::new("model/chatgpt/5.6-sol/date/xhigh".to_owned())
.unwrap()
.as_str(),
"model/chatgpt/5.6-sol/date/xhigh"
);
assert_eq!(TargetName::new(exact.clone()).unwrap().into_string(), exact);
assert!(TargetName::new("界".repeat(81)).is_err());
assert!(TargetName::new(String::new()).is_err());
assert!(TargetName::new(" ".to_owned()).is_err());
assert!(TargetName::new("line\nbreak".to_owned()).is_err());
assert_eq!(
TargetName::new(" Mixed/Case ".to_owned())
.unwrap()
.into_string(),
" Mixed/Case "
);
}
#[test]
fn authorities_have_overlapping_names_and_replay_last_write_wins() {
let file = path("replay");
let user_target = target(user(1), "default/chat");
let group_target = target(group(1), "default/chat");
{
let store = LaunchNodes::open(&file).unwrap();
store.set(user_target.clone(), NodeId([1; 12])).unwrap();
store.set(group_target.clone(), NodeId([2; 12])).unwrap();
store.set(user_target.clone(), NodeId([3; 12])).unwrap();
}
let reopened = LaunchNodes::open(&file).unwrap();
assert_eq!(reopened.get(&user_target).unwrap(), Some(NodeId([3; 12])));
assert_eq!(reopened.get(&group_target).unwrap(), Some(NodeId([2; 12])));
fs::remove_file(file).unwrap();
}
#[test]
fn incomplete_tail_repairs_but_complete_corruption_fails() {
let partial = path("partial");
let key = target(user(2), "model/x");
{
let store = LaunchNodes::open(&partial).unwrap();
store.set(key.clone(), NodeId([4; 12])).unwrap();
}
let clean_len = fs::metadata(&partial).unwrap().len();
{
let mut file = OpenOptions::new().append(true).open(&partial).unwrap();
file.write_all(&[USER_TAG, 1, 2]).unwrap();
file.sync_data().unwrap();
}
let reopened = LaunchNodes::open(&partial).unwrap();
assert_eq!(reopened.get(&key).unwrap(), Some(NodeId([4; 12])));
assert_eq!(fs::metadata(&partial).unwrap().len(), clean_len);
fs::remove_file(partial).unwrap();
let corrupt = path("corrupt");
{
let store = LaunchNodes::open(&corrupt).unwrap();
store
.set(target(group(3), "default"), NodeId([5; 12]))
.unwrap();
}
let length = fs::metadata(&corrupt).unwrap().len();
{
let mut file = OpenOptions::new()
.read(true)
.write(true)
.open(&corrupt)
.unwrap();
file.seek(SeekFrom::Start(length - 1)).unwrap();
let mut byte = [0];
file.read_exact(&mut byte).unwrap();
file.seek(SeekFrom::Start(length - 1)).unwrap();
file.write_all(&[byte[0] ^ 0xff]).unwrap();
file.sync_data().unwrap();
}
assert!(LaunchNodes::open(&corrupt).is_err());
fs::remove_file(corrupt).unwrap();
}
#[test]
fn headerless_v1_is_rejected() {
let file = path("v1");
fs::write(&file, [0_u8; 28]).unwrap();
assert_eq!(
LaunchNodes::open(&file).err().unwrap(),
"unsupported launch nodes file format"
);
fs::remove_file(file).unwrap();
}
#[test]
fn complete_package_stays_below_the_managed_limit() {
let files = [
include_str!("../Cargo.toml"),
include_str!("../Documentation.md"),
include_str!("lib.rs"),
];
let count = files
.iter()
.flat_map(|file| file.lines())
.filter(|line| !line.trim().is_empty())
.count();
assert!(count < 500, "complete package has {count} nonblank lines");
}
}