#![doc = include_str!("../Documentation.md")]
#![forbid(unsafe_code)]
use kcode_k1_peering::K1Peering;
use kcode_k1_txn_ordering::{K1TxnOrdering, Subsystem, SubsystemId, TxId};
use semver::Version;
use std::sync::{Arc, Mutex, MutexGuard, Weak};
const SUBSYSTEM: &str = "k1-loom-bootstrap";
const WIRE_VERSION: u8 = 2;
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct BeginRecord {
username: String,
full_name: String,
public_key: [u8; 32],
}
impl BeginRecord {
pub fn new(username: String, full_name: String, public_key: [u8; 32]) -> Self {
Self {
username,
full_name,
public_key,
}
}
pub fn username(&self) -> &str {
&self.username
}
pub fn full_name(&self) -> &str {
&self.full_name
}
pub const fn public_key(&self) -> [u8; 32] {
self.public_key
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct ImportedPackage {
logical_name: String,
version: Version,
}
impl ImportedPackage {
pub fn new(logical_name: String, version: Version) -> Result<Self, String> {
if logical_name.is_empty() {
return Err("imported package logical name is empty".to_owned());
}
if !version.pre.is_empty() || !version.build.is_empty() {
return Err("imported package version must be stable".to_owned());
}
Ok(Self {
logical_name,
version,
})
}
pub fn logical_name(&self) -> &str {
&self.logical_name
}
pub fn version(&self) -> &Version {
&self.version
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub struct CompleteRecord {
begin: BeginRecord,
user_id: [u8; 12],
group_id: [u8; 12],
profile_id: [u8; 12],
root_id: [u8; 12],
packages: Vec<ImportedPackage>,
}
impl CompleteRecord {
pub fn new(
begin: BeginRecord,
user_id: [u8; 12],
group_id: [u8; 12],
profile_id: [u8; 12],
root_id: [u8; 12],
mut packages: Vec<ImportedPackage>,
) -> Result<Self, String> {
packages.sort_by(|left, right| {
left.logical_name
.cmp(&right.logical_name)
.then_with(|| left.version.cmp(&right.version))
});
if packages.windows(2).any(|pair| {
pair[0].logical_name == pair[1].logical_name && pair[0].version == pair[1].version
}) {
return Err("duplicate imported package coordinate".to_owned());
}
Ok(Self {
begin,
user_id,
group_id,
profile_id,
root_id,
packages,
})
}
pub fn begin(&self) -> &BeginRecord {
&self.begin
}
pub const fn user_id(&self) -> [u8; 12] {
self.user_id
}
pub const fn group_id(&self) -> [u8; 12] {
self.group_id
}
pub const fn profile_id(&self) -> [u8; 12] {
self.profile_id
}
pub const fn root_id(&self) -> [u8; 12] {
self.root_id
}
pub fn packages(&self) -> &[ImportedPackage] {
&self.packages
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum BootstrapStatus {
Empty,
Begun {
record: BeginRecord,
transaction: TxId,
},
Complete {
record: CompleteRecord,
transaction: TxId,
},
}
#[derive(Clone)]
pub struct K1BootstrapState {
inner: Arc<Inner>,
}
struct Inner {
peering: Arc<K1Peering>,
state: Mutex<State>,
}
struct State {
status: BootstrapStatus,
available: bool,
}
struct Handler(Weak<Inner>);
#[derive(Debug)]
enum Event {
Begin(BeginRecord),
Complete(CompleteRecord),
}
impl K1BootstrapState {
pub fn open(ordering: Arc<K1TxnOrdering>, peering: Arc<K1Peering>) -> Result<Self, String> {
let inner = Arc::new(Inner {
peering,
state: Mutex::new(State {
status: BootstrapStatus::Empty,
available: true,
}),
});
ordering.register_subsystem(
SubsystemId::from_str(SUBSYSTEM)?,
None,
Arc::new(Handler(Arc::downgrade(&inner))),
)?;
Ok(Self { inner })
}
pub fn status(&self) -> Result<BootstrapStatus, String> {
Ok(ready(&self.inner)?.status.clone())
}
pub fn begin(&self, record: BeginRecord) -> Result<TxId, String> {
match self.status()? {
BootstrapStatus::Begun {
record: existing,
transaction,
} if existing == record => {
return Ok(transaction);
}
BootstrapStatus::Complete {
record: existing,
transaction,
} if existing.begin == record => return Ok(transaction),
BootstrapStatus::Empty => {}
_ => return Err("bootstrap Begin conflicts with canonical state".to_owned()),
}
self.submit(Event::Begin(record.clone()))?;
match self.status()? {
BootstrapStatus::Begun {
record: existing,
transaction,
} if existing == record => Ok(transaction),
BootstrapStatus::Complete {
record: existing,
transaction,
} if existing.begin == record => Ok(transaction),
_ => Err("bootstrap Begin lost to conflicting canonical state".to_owned()),
}
}
pub fn complete(&self, record: CompleteRecord) -> Result<TxId, String> {
match self.status()? {
BootstrapStatus::Begun {
record: existing, ..
} if existing == record.begin => {}
BootstrapStatus::Complete {
record: existing,
transaction,
} if existing == record => {
return Ok(transaction);
}
BootstrapStatus::Empty => return Err("bootstrap has not begun".to_owned()),
_ => return Err("bootstrap Complete conflicts with canonical state".to_owned()),
}
self.submit(Event::Complete(record.clone()))?;
match self.status()? {
BootstrapStatus::Complete {
record: existing,
transaction,
} if existing == record => Ok(transaction),
_ => Err("bootstrap Complete lost to conflicting canonical state".to_owned()),
}
}
fn submit(&self, event: Event) -> Result<TxId, String> {
self.inner
.peering
.submit_txn(SubsystemId::from_str(SUBSYSTEM)?, &encode(&event)?)
}
}
impl Subsystem for Handler {
fn submit_txn(&self, id: TxId, payload: &[u8]) -> Result<(), String> {
let inner = self
.0
.upgrade()
.ok_or_else(|| "bootstrap state is unavailable".to_owned())?;
let event = decode(payload).map_err(|error| fail(&inner, error))?;
let mut state = ready(&inner)?;
match (&state.status, event) {
(BootstrapStatus::Empty, Event::Begin(record)) => {
state.status = BootstrapStatus::Begun {
record,
transaction: id,
};
}
(BootstrapStatus::Begun { record: begun, .. }, Event::Complete(record))
if *begun == record.begin =>
{
state.status = BootstrapStatus::Complete {
record,
transaction: id,
};
}
_ => {}
}
Ok(())
}
fn reorg(&self) -> Result<(), String> {
let inner = self
.0
.upgrade()
.ok_or_else(|| "bootstrap state is unavailable".to_owned())?;
lock(&inner)?.available = false;
Ok(())
}
}
fn lock(inner: &Inner) -> Result<MutexGuard<'_, State>, String> {
inner
.state
.lock()
.map_err(|_| "bootstrap state lock poisoned".to_owned())
}
fn ready(inner: &Inner) -> Result<MutexGuard<'_, State>, String> {
let state = lock(inner)?;
if state.available {
Ok(state)
} else {
Err("bootstrap state is unavailable".to_owned())
}
}
fn fail(inner: &Inner, error: String) -> String {
if let Ok(mut state) = lock(inner) {
state.available = false;
}
error
}
fn encode(event: &Event) -> Result<Vec<u8>, String> {
let mut out = vec![
WIRE_VERSION,
match event {
Event::Begin(_) => 1,
Event::Complete(_) => 2,
},
];
match event {
Event::Begin(record) => put_begin(&mut out, record)?,
Event::Complete(record) => {
put_begin(&mut out, &record.begin)?;
out.extend_from_slice(&record.user_id);
out.extend_from_slice(&record.group_id);
out.extend_from_slice(&record.profile_id);
out.extend_from_slice(&record.root_id);
put_u64(&mut out, record.packages.len())?;
for package in &record.packages {
put_string(&mut out, &package.logical_name)?;
put_string(&mut out, &package.version.to_string())?;
}
}
}
Ok(out)
}
fn put_begin(out: &mut Vec<u8>, record: &BeginRecord) -> Result<(), String> {
put_string(out, &record.username)?;
put_string(out, &record.full_name)?;
out.extend_from_slice(&record.public_key);
Ok(())
}
fn put_string(out: &mut Vec<u8>, value: &str) -> Result<(), String> {
put_u64(out, value.len())?;
out.try_reserve(value.len())
.map_err(|_| "bootstrap encoding allocation failed".to_owned())?;
out.extend_from_slice(value.as_bytes());
Ok(())
}
fn put_u64(out: &mut Vec<u8>, value: usize) -> Result<(), String> {
out.extend_from_slice(
&u64::try_from(value)
.map_err(|_| "bootstrap encoding length overflow".to_owned())?
.to_le_bytes(),
);
Ok(())
}
fn decode(bytes: &[u8]) -> Result<Event, String> {
let mut input = Input { bytes, at: 0 };
if input.byte()? != WIRE_VERSION {
return Err("unsupported bootstrap state wire version".to_owned());
}
let kind = input.byte()?;
let begin = input.begin()?;
let event = match kind {
1 => Event::Begin(begin),
2 => {
let user_id = input.array()?;
let group_id = input.array()?;
let profile_id = input.array()?;
let root_id = input.array()?;
let count = input.length()?;
let mut packages = Vec::new();
packages
.try_reserve(count)
.map_err(|_| "bootstrap package allocation failed".to_owned())?;
for _ in 0..count {
packages.push(ImportedPackage::new(
input.string()?,
input
.string()?
.parse()
.map_err(|_| "invalid bootstrap package version".to_owned())?,
)?);
}
Event::Complete(CompleteRecord::new(
begin, user_id, group_id, profile_id, root_id, packages,
)?)
}
_ => return Err("unknown bootstrap state event kind".to_owned()),
};
if input.at != bytes.len() {
return Err("trailing bootstrap state bytes".to_owned());
}
Ok(event)
}
struct Input<'a> {
bytes: &'a [u8],
at: usize,
}
impl Input<'_> {
fn take(&mut self, count: usize) -> Result<&[u8], String> {
let end = self
.at
.checked_add(count)
.ok_or_else(|| "bootstrap state length overflow".to_owned())?;
let value = self
.bytes
.get(self.at..end)
.ok_or_else(|| "truncated bootstrap state event".to_owned())?;
self.at = end;
Ok(value)
}
fn byte(&mut self) -> Result<u8, String> {
Ok(self.take(1)?[0])
}
fn array<const N: usize>(&mut self) -> Result<[u8; N], String> {
self.take(N)?
.try_into()
.map_err(|_| "invalid bootstrap fixed field".to_owned())
}
fn length(&mut self) -> Result<usize, String> {
usize::try_from(u64::from_le_bytes(self.array()?))
.map_err(|_| "bootstrap state length overflow".to_owned())
}
fn string(&mut self) -> Result<String, String> {
let length = self.length()?;
String::from_utf8(self.take(length)?.to_vec())
.map_err(|_| "bootstrap state string is not UTF-8".to_owned())
}
fn begin(&mut self) -> Result<BeginRecord, String> {
Ok(BeginRecord::new(
self.string()?,
self.string()?,
self.array()?,
))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn version_two_wire_round_trips_without_code_hashes() {
let begin = BeginRecord::new("admin".into(), "Admin".into(), [2; 32]);
let packages = vec![ImportedPackage::new("loom".into(), Version::new(0, 1, 0)).unwrap()];
let complete =
CompleteRecord::new(begin, [4; 12], [5; 12], [6; 12], [7; 12], packages).unwrap();
let encoded = encode(&Event::Complete(complete.clone())).unwrap();
assert_eq!(encoded[0], 2);
assert!(matches!(decode(&encoded).unwrap(), Event::Complete(found) if found == complete));
assert!(decode(&[1, 1]).unwrap_err().contains("wire version"));
}
}