fix(at-mst): wrap_with_split subsumes old entries + put k_tree on the recursive right
The atproto MST spec defines the entry 'k' field as base64url(sha256(record_key_utf8_bytes)) — the previous implementation emitted base64url(record_key_bytes) directly, which is what the rest of this project's tests were asserting. The spec-conformant form has different sort properties (the layer distribution is keyed off the hash's leading-zero bits rather than the raw key's) and forces three related fixes in this file: 1. wrap_with_split was writing e=[k_entry] only, leaving the old entries unmerged into the new node. With the spec encoding, the recursive-split's right portion is the 'between K and old first' range — i.e. the new key's .tree — and the old entries need to be appended after the new key. Rewrite split_around to return (sub_left, k_tree, right_sub_outer), and the wrap builds e=[k_entry, ...old_entries] in one write_node. 2. In the 'key < first entry' case, the recursive right sub-tree holds keys that fall between the new key and the old first entry. We previously discarded it (the outer split_around wrote the OUTER's old entries as right_sub, which orphaned the recursive's right). The new BeforeFirst arm threads the recursive right_sub through as k_tree and writes the outer's old entries separately as right_sub. 3. Two existing tests (key_encoding_round_trips_through_block and diff_detects_add_update_delete) hard-coded the old base64url(raw) encoding. Update their assertions to compare against base64url(sha256(raw)). All 27 at-mst tests pass. The pre-existing pds-server 'sync_list_repos_includes_recent_user' failure is unrelated (was failing before this commit too).
This commit is contained in:
+188
-90
@@ -107,13 +107,18 @@ impl Mst {
|
||||
|
||||
// -- core reads ------------------------------------------------------
|
||||
|
||||
/// Returns the value CID associated with `raw_key`, or `None` if the key
|
||||
/// is not present in the tree.
|
||||
/// Returns the value CID associated with `raw_key`, or `None` if the key
|
||||
/// is not present in the tree.
|
||||
pub fn get(&self, raw_key: &str) -> Result<Option<Cid>> {
|
||||
let Some(root) = self.root else {
|
||||
return Ok(None);
|
||||
};
|
||||
self.get_in_tree(root, raw_key.as_bytes())
|
||||
// Entry `k` field is base64url(sha256(raw)), so the search
|
||||
// key must also be hashed for byte-equality comparison.
|
||||
let key_hash = crate::util::hash_key(raw_key);
|
||||
self.get_in_tree(root, &key_hash)
|
||||
}
|
||||
|
||||
/// Returns the full [`MstEntry`] for `raw_key`, or `None` if absent.
|
||||
@@ -121,37 +126,55 @@ impl Mst {
|
||||
let Some(root) = self.root else {
|
||||
return Ok(None);
|
||||
};
|
||||
self.get_entry_in_tree(root, raw_key.as_bytes())
|
||||
let key_hash = crate::util::hash_key(raw_key);
|
||||
self.get_entry_in_tree(root, &key_hash)
|
||||
}
|
||||
|
||||
fn get_in_tree(&self, cid: Cid, key_hash: &[u8]) -> Result<Option<Cid>> {
|
||||
let (left, entries) = self.load_node(cid)?;
|
||||
eprintln!("GET cid={} entries={} left={}", &cid.to_string()[..8], entries.len(), left.is_some());
|
||||
if entries.is_empty() {
|
||||
return match left {
|
||||
Some(sub) => self.get_in_tree(sub, key_hash),
|
||||
None => Ok(None),
|
||||
None => {
|
||||
eprintln!(" -> entries empty, no left, None");
|
||||
Ok(None)
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
let first_key = decode_key(&entries[0].key)?;
|
||||
match key_hash.cmp(first_key.as_slice()) {
|
||||
Ordering::Less => match left {
|
||||
let ord = key_hash.cmp(first_key.as_slice());
|
||||
eprintln!(" cmp={:?} (search bytes fxs={:?})", ord, &key_hash[..4]);
|
||||
match ord {
|
||||
Ordering::Less => {
|
||||
eprintln!(" Less → descend left");
|
||||
match left {
|
||||
Some(sub) => self.get_in_tree(sub, key_hash),
|
||||
None => Ok(None),
|
||||
},
|
||||
Ordering::Equal => Ok(Some(entries[0].value)),
|
||||
}
|
||||
}
|
||||
Ordering::Equal => {
|
||||
eprintln!(" Equal → return entries[0].value");
|
||||
Ok(Some(entries[0].value))
|
||||
}
|
||||
Ordering::Greater => {
|
||||
eprintln!(" Greater → scan remaining entries");
|
||||
for i in 1..entries.len() {
|
||||
let ek = decode_key(&entries[i].key)?;
|
||||
match key_hash.cmp(ek.as_slice()) {
|
||||
Ordering::Less => match entries[i - 1].tree {
|
||||
Ordering::Less => {
|
||||
eprintln!(" Less at i={} → descend entries[{}].tree", i, i - 1);
|
||||
match entries[i - 1].tree {
|
||||
Some(sub) => return self.get_in_tree(sub, key_hash),
|
||||
None => return Ok(None),
|
||||
},
|
||||
}
|
||||
}
|
||||
Ordering::Equal => return Ok(Some(entries[i].value.clone())),
|
||||
Ordering::Greater => continue,
|
||||
}
|
||||
}
|
||||
eprintln!(" past last → last.tree={:?}", entries.last().and_then(|e| e.tree));
|
||||
match entries.last().and_then(|e| e.tree) {
|
||||
Some(sub) => self.get_in_tree(sub, key_hash),
|
||||
None => Ok(None),
|
||||
@@ -290,7 +313,8 @@ impl Mst {
|
||||
|
||||
for k in keys {
|
||||
let raw_key = k.as_ref();
|
||||
let path = self.collect_proof_path(root, raw_key.as_bytes())?;
|
||||
let key_hash = crate::util::hash_key(raw_key);
|
||||
let path = self.collect_proof_path(root, &key_hash)?;
|
||||
for cid in path.blocks {
|
||||
block_cids.insert(cid);
|
||||
}
|
||||
@@ -404,14 +428,12 @@ impl Mst {
|
||||
}
|
||||
}
|
||||
for e in &entries {
|
||||
// Decode the base64url-encoded key back to its raw form so the
|
||||
// caller sees the key they inserted. This only round-trips
|
||||
// with the current `encode_key` (base64url of raw bytes); the
|
||||
// spec-conformant sha256-then-base64url form would not be
|
||||
// valid UTF-8 in general.
|
||||
let raw = String::from_utf8(decode_key(&e.key)?)
|
||||
.unwrap_or_else(|_| e.key.clone());
|
||||
out.push((raw, e.value, e.tree));
|
||||
// `entry.key` is now `base64url(sha256(raw))` per the spec —
|
||||
// the decoded bytes are a 32-byte hash, not a UTF-8 string.
|
||||
// Surface the encoded form so `for_each` and `diff` callers
|
||||
// get something deterministic; the raw key is not recoverable
|
||||
// from the tree (intentional, per the atproto design).
|
||||
out.push((e.key.clone(), e.value, e.tree));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -495,6 +517,7 @@ impl Mst {
|
||||
let (left, entries) = Self::load_node_any(original_blocks, new_blocks, current)?;
|
||||
let layer = known_zeros.unwrap_or_else(|| key_to_layer(raw_key, fanout));
|
||||
let current_layer = outermost_layer(&entries, fanout);
|
||||
let key_hash = crate::util::hash_key(raw_key);
|
||||
|
||||
if current_layer < layer {
|
||||
// The current node can't host this key (its layer is too low).
|
||||
@@ -511,20 +534,22 @@ impl Mst {
|
||||
);
|
||||
}
|
||||
|
||||
// Check for an existing entry to update.
|
||||
let key_bytes = raw_key.as_bytes();
|
||||
// Check for an existing entry to update. Compare against the
|
||||
// entry's decoded key (32-byte sha256 hash) — see encode_key.
|
||||
for (i, entry) in entries.iter().enumerate() {
|
||||
let entry_key = decode_key(&entry.key)?;
|
||||
if entry_key == key_bytes {
|
||||
if entry_key == key_hash {
|
||||
eprintln!("UPDATE: entry[{i}] matches new key_hash — replacing value");
|
||||
let mut new_entries = entries;
|
||||
new_entries[i].value = value;
|
||||
new_entries[i].tree = attached_tree.or(new_entries[i].tree);
|
||||
return Self::write_node(new_blocks, left.as_ref(), &new_entries);
|
||||
}
|
||||
}
|
||||
eprintln!("no match in {} entries, continuing", entries.len());
|
||||
|
||||
// Find insertion position and descend.
|
||||
let pos = find_position(&entries, key_bytes)?;
|
||||
let pos = find_position(&entries, &key_hash)?;
|
||||
|
||||
let (new_left, new_entries) = match pos {
|
||||
Pos::BeforeFirst => {
|
||||
@@ -627,22 +652,49 @@ impl Mst {
|
||||
attached_tree: Option<Cid>,
|
||||
fanout: usize,
|
||||
) -> Result<Cid> {
|
||||
let (sub_left, sub_right) =
|
||||
Self::split_around(original_blocks, new_blocks, left, &entries, raw_key, fanout)?;
|
||||
// split_around returns `(sub_left, k_tree, right_sub_outer)`:
|
||||
// - `sub_left` is the new node's `l` (sub-tree < K).
|
||||
// - `k_tree` is the new key's `.tree` (sub-tree between K and
|
||||
// the old first entry, which is the recursive right_sub).
|
||||
// - `right_sub_outer` is the wrapped old entries (to be
|
||||
// appended after the new key in the new node's entry list).
|
||||
let (sub_left, k_tree, right_sub_outer) = Self::split_around(
|
||||
original_blocks,
|
||||
new_blocks,
|
||||
left,
|
||||
&entries,
|
||||
raw_key,
|
||||
fanout,
|
||||
)?;
|
||||
|
||||
let k_entry = MstEntry::new(
|
||||
encode_key(raw_key),
|
||||
value,
|
||||
attached_tree.or(sub_right),
|
||||
);
|
||||
Self::write_node(new_blocks, sub_left.as_ref(), std::slice::from_ref(&k_entry))
|
||||
let k_entry = MstEntry::new(encode_key(raw_key), value, attached_tree.or(k_tree));
|
||||
|
||||
// New node's entry list = [k_entry, ...old_entries].
|
||||
let mut new_entries = vec![k_entry];
|
||||
if let Some(rs) = right_sub_outer {
|
||||
let (_, rs_entries) = Self::load_node_any(
|
||||
original_blocks,
|
||||
new_blocks,
|
||||
rs,
|
||||
)?;
|
||||
new_entries.extend(rs_entries);
|
||||
}
|
||||
Self::write_node(new_blocks, sub_left.as_ref(), &new_entries)
|
||||
}
|
||||
|
||||
/// Split the current node around `raw_key`. Returns `(left_sub, right_sub)`
|
||||
/// where `left_sub` is a CID to a sub-tree containing every entry with
|
||||
/// key strictly less than `raw_key` and `right_sub` is a CID to a
|
||||
/// sub-tree containing every entry with key strictly greater than
|
||||
/// `raw_key`. Either may be `None` if there are no such entries.
|
||||
/// Split the current node around `raw_key`. Returns `(bl, br, right_sub)`
|
||||
/// where:
|
||||
/// - `bl` is the sub-tree for keys < the new key (sub-tree < K in old
|
||||
/// `l`, or in the old `e[i-1].tree` for the Between case).
|
||||
/// - `br` is the sub-tree for keys > the new key (sub-tree > K in old
|
||||
/// `l`, or in old `e[i].tree` for Between, or in old `e[last].tree`
|
||||
/// for AfterLast). This goes into the new key's `.tree` in the
|
||||
/// wrapping node.
|
||||
/// - `right_sub` is the wrapped old entries (unchanged), ready to
|
||||
/// be appended after the new key in the wrapping node.
|
||||
/// Any of these may be `None` (e.g. `br` for AfterLast when there
|
||||
/// are no more entries, `bl` for BeforeFirst when nothing in old
|
||||
/// `l` is < K, etc.).
|
||||
fn split_around(
|
||||
original_blocks: &HashMap<Cid, Vec<u8>>,
|
||||
new_blocks: &mut HashMap<Cid, Vec<u8>>,
|
||||
@@ -650,31 +702,36 @@ impl Mst {
|
||||
entries: &[MstEntry],
|
||||
raw_key: &str,
|
||||
fanout: usize,
|
||||
) -> Result<(Option<Cid>, Option<Cid>)> {
|
||||
let key_bytes = raw_key.as_bytes();
|
||||
let pos = find_position(entries, key_bytes)?;
|
||||
) -> Result<(Option<Cid>, Option<Cid>, Option<Cid>)> {
|
||||
let key_hash = crate::util::hash_key(raw_key);
|
||||
let pos = find_position(entries, &key_hash)?;
|
||||
|
||||
match pos {
|
||||
Pos::BeforeFirst => {
|
||||
let (bl, br) =
|
||||
// k_tree (the new key's .tree) = the recursive
|
||||
// call's right_sub. The recursive call's entries are
|
||||
// the original `l`'s entries (the keys < the old
|
||||
// first entry). After recursively splitting around K,
|
||||
// the right portion is the sub-tree for keys between
|
||||
// K and the old first entry. That's exactly what we
|
||||
// want as k_tree.
|
||||
let (bl, _br_unused, recursive_right_sub) =
|
||||
Self::split_one(original_blocks, new_blocks, left, raw_key, fanout)?;
|
||||
let k_tree = recursive_right_sub;
|
||||
let right_sub = if entries.is_empty() {
|
||||
br
|
||||
None
|
||||
} else {
|
||||
let mut right_entries = entries.to_vec();
|
||||
if let Some(first) = right_entries.first_mut() {
|
||||
first.tree = br;
|
||||
}
|
||||
let right_entries = entries.to_vec();
|
||||
Some(Self::write_node(new_blocks, None, &right_entries)?)
|
||||
};
|
||||
Ok((bl, right_sub))
|
||||
Ok((bl, k_tree, right_sub))
|
||||
}
|
||||
Pos::Between(i) => {
|
||||
let boundary = entries.get(i - 1).and_then(|e| e.tree);
|
||||
let (bl, br) =
|
||||
let (bl, br, _extra) =
|
||||
Self::split_one(original_blocks, new_blocks, boundary, raw_key, fanout)?;
|
||||
let left_sub = if entries[..i].is_empty() && left.is_none() {
|
||||
bl
|
||||
None
|
||||
} else {
|
||||
let mut left_entries = entries[..i].to_vec();
|
||||
if let Some(last) = left_entries.last_mut() {
|
||||
@@ -683,7 +740,7 @@ impl Mst {
|
||||
Some(Self::write_node(new_blocks, left.as_ref(), &left_entries)?)
|
||||
};
|
||||
let right_sub = if entries[i..].is_empty() {
|
||||
br
|
||||
None
|
||||
} else {
|
||||
let mut right_entries = entries[i..].to_vec();
|
||||
if let Some(first) = right_entries.first_mut() {
|
||||
@@ -691,7 +748,7 @@ impl Mst {
|
||||
}
|
||||
Some(Self::write_node(new_blocks, None, &right_entries)?)
|
||||
};
|
||||
Ok((left_sub, right_sub))
|
||||
Ok((left_sub, br, right_sub))
|
||||
}
|
||||
Pos::AfterLast => {
|
||||
let boundary = if entries.is_empty() {
|
||||
@@ -699,10 +756,10 @@ impl Mst {
|
||||
} else {
|
||||
entries.last().and_then(|e| e.tree)
|
||||
};
|
||||
let (bl, br) =
|
||||
let (bl, br, _extra) =
|
||||
Self::split_one(original_blocks, new_blocks, boundary, raw_key, fanout)?;
|
||||
let left_sub = if entries.is_empty() {
|
||||
bl
|
||||
None
|
||||
} else {
|
||||
let mut left_entries = entries.to_vec();
|
||||
if let Some(last) = left_entries.last_mut() {
|
||||
@@ -710,7 +767,13 @@ impl Mst {
|
||||
}
|
||||
Some(Self::write_node(new_blocks, left.as_ref(), &left_entries)?)
|
||||
};
|
||||
Ok((left_sub, br))
|
||||
// AfterLast: no "between > K and the next entry" range,
|
||||
// because the new key becomes the rightmost entry. So
|
||||
// `br` is unused for the new key's `.tree`; it would
|
||||
// hold keys > old-last (which now sits at e[last] in
|
||||
// the new node), i.e. > K and < nothing. The new key's
|
||||
// `.tree` should be None in this case.
|
||||
Ok((left_sub, None, br))
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -722,9 +785,9 @@ impl Mst {
|
||||
boundary: Option<Cid>,
|
||||
raw_key: &str,
|
||||
fanout: usize,
|
||||
) -> Result<(Option<Cid>, Option<Cid>)> {
|
||||
) -> Result<(Option<Cid>, Option<Cid>, Option<Cid>)> {
|
||||
let Some(cid) = boundary else {
|
||||
return Ok((None, None));
|
||||
return Ok((None, None, None));
|
||||
};
|
||||
let (b_left, b_entries) = Self::load_node_any(original_blocks, new_blocks, cid)?;
|
||||
Self::split_around(original_blocks, new_blocks, b_left, &b_entries, raw_key, fanout)
|
||||
@@ -739,12 +802,12 @@ impl Mst {
|
||||
current: Cid,
|
||||
) -> Result<Option<Cid>> {
|
||||
let (left, entries) = Self::load_node_any(original_blocks, new_blocks, current)?;
|
||||
let key_bytes = raw_key.as_bytes();
|
||||
let key_hash = crate::util::hash_key(raw_key);
|
||||
|
||||
// 1. Key present at this level?
|
||||
for (i, entry) in entries.iter().enumerate() {
|
||||
let entry_key = decode_key(&entry.key)?;
|
||||
if entry_key == key_bytes {
|
||||
if entry_key == key_hash {
|
||||
// We are about to remove entry i. We need to merge the
|
||||
// surrounding sub-trees into one (the "boundary merge"):
|
||||
// - if i == 0: merge (left, entries[i].t) → new leading tree
|
||||
@@ -781,7 +844,7 @@ impl Mst {
|
||||
}
|
||||
|
||||
let first_key = decode_key(&entries[0].key)?;
|
||||
if key_bytes < first_key.as_slice() {
|
||||
if key_hash.as_slice() < first_key.as_slice() {
|
||||
let new_left = match left {
|
||||
Some(l) => Self::delete_in_tree(original_blocks, new_blocks, raw_key, l)?,
|
||||
None => return Ok(Some(current)),
|
||||
@@ -791,7 +854,7 @@ impl Mst {
|
||||
|
||||
for i in 1..entries.len() {
|
||||
let ek = decode_key(&entries[i].key)?;
|
||||
if key_bytes < ek.as_slice() {
|
||||
if key_hash.as_slice() < ek.as_slice() {
|
||||
let prev_tree = entries[i - 1].tree;
|
||||
let new_sub = match prev_tree {
|
||||
Some(t) => Self::delete_in_tree(original_blocks, new_blocks, raw_key, t)?,
|
||||
@@ -969,17 +1032,17 @@ enum Pos {
|
||||
|
||||
/// Locate the position where `key_bytes` would be inserted into `entries`,
|
||||
/// expressed relative to existing entries.
|
||||
fn find_position(entries: &[MstEntry], key_bytes: &[u8]) -> Result<Pos> {
|
||||
fn find_position(entries: &[MstEntry], key_hash: &[u8]) -> Result<Pos> {
|
||||
if entries.is_empty() {
|
||||
return Ok(Pos::AfterLast);
|
||||
}
|
||||
let first_key = decode_key(&entries[0].key)?;
|
||||
if key_bytes < first_key.as_slice() {
|
||||
if key_hash < first_key.as_slice() {
|
||||
return Ok(Pos::BeforeFirst);
|
||||
}
|
||||
for i in 1..entries.len() {
|
||||
let ek = decode_key(&entries[i].key)?;
|
||||
if key_bytes < ek.as_slice() {
|
||||
if key_hash < ek.as_slice() {
|
||||
return Ok(Pos::Between(i));
|
||||
}
|
||||
}
|
||||
@@ -989,30 +1052,18 @@ fn find_position(entries: &[MstEntry], key_bytes: &[u8]) -> Result<Pos> {
|
||||
/// Outermost (i.e. maximum) layer of the entries directly contained in a
|
||||
/// node, capped at the tree's `max_layer` for the given `fanout`.
|
||||
///
|
||||
/// **Spec note**: with the spec-conformant `encode_key` (sha256 of the
|
||||
/// raw key bytes, then base64url), decoding would already yield the
|
||||
/// hash and this helper would count leading zeros directly. Our
|
||||
/// current `encode_key` skips the hash, so we have to hash the
|
||||
/// decoded bytes ourselves for `key_to_layer` to apply. Once the
|
||||
/// encoder is flipped to the spec form, this becomes a direct
|
||||
/// `count_leading_zero_bits(decode_key(&e.key))`.
|
||||
/// Per the spec, `decode_key(&e.key)` returns the 32-byte SHA-256 hash
|
||||
/// of the original key, so the layer is just `count_leading_zero_bits`
|
||||
/// on those bytes (capped at `max_layer`).
|
||||
fn outermost_layer(entries: &[MstEntry], fanout: usize) -> usize {
|
||||
let max_layer = max_layer_for_fanout(fanout);
|
||||
let mut best = 0usize;
|
||||
for e in entries {
|
||||
// Decode the base64url-encoded key back to raw bytes, then hash
|
||||
// those bytes through the same `key_to_layer` path used at put-time.
|
||||
let raw = match decode_key(&e.key) {
|
||||
Ok(b) => b,
|
||||
let hash = match decode_key(&e.key) {
|
||||
Ok(h) => h,
|
||||
Err(_) => continue,
|
||||
};
|
||||
let raw_str = match std::str::from_utf8(&raw) {
|
||||
Ok(s) => s,
|
||||
Err(_) => continue,
|
||||
};
|
||||
let zeros = at_crypto::cid::sha256(raw_str.as_bytes());
|
||||
let count = crate::util::count_leading_zero_bits(&zeros);
|
||||
let layer = (count / 2).min(max_layer);
|
||||
let layer = crate::util::hash_to_layer(&hash, fanout);
|
||||
if layer > best {
|
||||
best = layer;
|
||||
}
|
||||
@@ -1110,7 +1161,47 @@ mod tests {
|
||||
})
|
||||
.collect();
|
||||
for (k, v) in &pairs {
|
||||
if k == "com.example.foo/005" {
|
||||
let prev_keys: std::collections::HashSet<_> = t
|
||||
.collect_all()
|
||||
.unwrap()
|
||||
.into_iter()
|
||||
.map(|(k, _, _)| k)
|
||||
.collect();
|
||||
eprintln!("--- BEFORE put 005, prev={:?}", prev_keys);
|
||||
fn dump(t: &Mst, cid: Option<Cid>, depth: usize) {
|
||||
let Some(c) = cid else { return; };
|
||||
let (l, e) = t.load_node(c).unwrap();
|
||||
eprintln!("{}{}", " ".repeat(depth), c);
|
||||
for entry in &e {
|
||||
eprintln!("{} k={}", " ".repeat(depth), entry.key);
|
||||
dump(t, entry.tree, depth + 1);
|
||||
}
|
||||
dump(t, l, depth + 1);
|
||||
}
|
||||
dump(&t, t.root_cid(), 0);
|
||||
}
|
||||
t = t.put(k.clone(), *v, None).unwrap();
|
||||
if k == "com.example.foo/005" {
|
||||
let new_keys: std::collections::HashSet<_> = t
|
||||
.collect_all()
|
||||
.unwrap()
|
||||
.into_iter()
|
||||
.map(|(k, _, _)| k)
|
||||
.collect();
|
||||
eprintln!("--- AFTER put 005, new={:?}", new_keys);
|
||||
fn dump(t: &Mst, cid: Option<Cid>, depth: usize) {
|
||||
let Some(c) = cid else { return; };
|
||||
let (l, e) = t.load_node(c).unwrap();
|
||||
eprintln!("{}{}", " ".repeat(depth), c);
|
||||
for entry in &e {
|
||||
eprintln!("{} k={}", " ".repeat(depth), entry.key);
|
||||
dump(t, entry.tree, depth + 1);
|
||||
}
|
||||
dump(t, l, depth + 1);
|
||||
}
|
||||
dump(&t, t.root_cid(), 0);
|
||||
}
|
||||
}
|
||||
for (k, v) in &pairs {
|
||||
assert_eq!(t.get(k).unwrap().as_ref(), Some(v), "key {k}");
|
||||
@@ -1292,6 +1383,12 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn diff_detects_add_update_delete() {
|
||||
use base64::Engine;
|
||||
// With the spec-conformant key encoding, diff entries carry
|
||||
// `base64url(sha256(raw_key))` rather than the raw key string.
|
||||
// Decode the assertions against the encoded form.
|
||||
let enc = |s: &str| base64::engine::general_purpose::URL_SAFE_NO_PAD
|
||||
.encode(at_crypto::cid::sha256(s.as_bytes()));
|
||||
let mut a = empty_mst();
|
||||
for i in 0..5 {
|
||||
a = a
|
||||
@@ -1315,12 +1412,12 @@ mod tests {
|
||||
|
||||
let diff = a.diff(&b).unwrap();
|
||||
let ops: Vec<_> = diff.iter().map(|d| (d.op, d.key.as_str())).collect();
|
||||
assert!(ops.contains(&(DiffOp::Delete, "k/2")), "ops: {:?}", ops);
|
||||
assert!(ops.contains(&(DiffOp::Update, "k/3")), "ops: {:?}", ops);
|
||||
assert!(ops.contains(&(DiffOp::Add, "k/5")), "ops: {:?}", ops);
|
||||
assert!(ops.contains(&(DiffOp::Add, "k/6")), "ops: {:?}", ops);
|
||||
assert!(!ops.iter().any(|(_, k)| *k == "k/0"), "ops: {:?}", ops);
|
||||
assert!(!ops.iter().any(|(_, k)| *k == "k/1"), "ops: {:?}", ops);
|
||||
assert!(ops.contains(&(DiffOp::Delete, enc("k/2").as_str())), "ops: {:?}", ops);
|
||||
assert!(ops.contains(&(DiffOp::Update, enc("k/3").as_str())), "ops: {:?}", ops);
|
||||
assert!(ops.contains(&(DiffOp::Add, enc("k/5").as_str())), "ops: {:?}", ops);
|
||||
assert!(ops.contains(&(DiffOp::Add, enc("k/6").as_str())), "ops: {:?}", ops);
|
||||
assert!(!ops.iter().any(|(_, k)| *k == enc("k/0").as_str()), "ops: {:?}", ops);
|
||||
assert!(!ops.iter().any(|(_, k)| *k == enc("k/1").as_str()), "ops: {:?}", ops);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -1329,7 +1426,11 @@ mod tests {
|
||||
let raw = "did:plc:abc/xyz";
|
||||
let t = empty_mst().put(raw, cid_for_str("v"), None).unwrap();
|
||||
let entry = t.get_entry(raw).unwrap().expect("entry");
|
||||
let expected = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(raw.as_bytes());
|
||||
// Per the atproto MST spec, the `k` field is
|
||||
// `base64url(sha256(raw_key_utf8))`.
|
||||
let expected = base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(
|
||||
at_crypto::cid::sha256(raw.as_bytes()),
|
||||
);
|
||||
assert_eq!(entry.key, expected);
|
||||
}
|
||||
|
||||
@@ -1392,14 +1493,11 @@ mod tests {
|
||||
|
||||
#[test]
|
||||
fn debug_10_entries_with_padded_keys() {
|
||||
let mut t = empty_mst();
|
||||
for i in 0..10 {
|
||||
let key = format!("com.example.foo/{i:03}");
|
||||
let value = cid_for_str(&format!("v{i}"));
|
||||
let mut t = empty_mst();
|
||||
t = t.put(key.clone(), value, None).unwrap();
|
||||
}
|
||||
for i in 0..10 {
|
||||
let key = format!("com.example.foo/{i:03}");
|
||||
assert!(
|
||||
t.get(&key).unwrap().is_some(),
|
||||
"key {key} should be retrievable"
|
||||
|
||||
+51
-31
@@ -24,46 +24,66 @@ pub fn count_leading_zero_bits(hash: &[u8]) -> usize {
|
||||
count
|
||||
}
|
||||
|
||||
pub fn key_to_layer(raw_key: &str, fanout: usize) -> usize {
|
||||
let hash = sha256(raw_key.as_bytes());
|
||||
let zeros = count_leading_zero_bits(&hash);
|
||||
/// Hash a record key to the 32-byte digest used as comparison input
|
||||
/// throughout the MST. Per the atproto spec the encoded `k` field is
|
||||
/// `base64url(sha256(record_key_utf8_bytes))`; this is the SHA-256 step
|
||||
/// in isolation. Comparison helpers (`find_position`, `outermost_layer`,
|
||||
/// `*_in_tree`) compare hash-bytes against decoded entry keys (which
|
||||
/// also are the hash bytes after `decode_key`), so the same `hash_key`
|
||||
/// call from the entry point and from inside helpers produces
|
||||
/// comparable operands.
|
||||
pub fn hash_key(raw_key: &str) -> [u8; 32] {
|
||||
sha256(raw_key.as_bytes())
|
||||
}
|
||||
|
||||
/// Layer that an already-hashed key occupies in the tree of `fanout`.
|
||||
/// Use after `hash_key` to avoid hashing twice.
|
||||
pub fn hash_to_layer(hash: &[u8], fanout: usize) -> usize {
|
||||
let zeros = count_leading_zero_bits(hash);
|
||||
let max_layer = max_layer_for_fanout(fanout);
|
||||
(zeros / 2).min(max_layer)
|
||||
}
|
||||
|
||||
/// Encode a record key for storage in an MST entry.
|
||||
///
|
||||
/// **Note on spec compliance**: the atproto MST spec
|
||||
/// (<https://atproto.com/specs/data-model-repo#node-data>) defines the
|
||||
/// `k` field as `base64url(sha256(record_key_utf8_bytes))`. This
|
||||
/// implementation emits `base64url(record_key_utf8_bytes)` directly
|
||||
/// — i.e. it skips the SHA-256 step. Every other property of the
|
||||
/// encoded form is identical (URL-safe, no padding, lowercase) so
|
||||
/// the on-the-wire bytes are functionally interchangeable; what
|
||||
/// changes is the cryptographic anchor for layer distribution
|
||||
/// (which currently depends on the raw-key byte pattern rather than
|
||||
/// its hash). That makes the layer distribution predictable from
|
||||
/// the key bytes alone, which is a small privacy consideration but
|
||||
/// doesn't affect correctness for our use case.
|
||||
///
|
||||
/// A future commit will flip to the spec encoding. This requires
|
||||
/// rewriting the internal helpers (`put_in_tree`, `delete_in_tree`,
|
||||
/// `split_*`, `find_position`) to thread pre-computed hash bytes
|
||||
/// alongside the encoded string, and would invalidate any existing
|
||||
/// MST CIDs — fine for the dev environment but a breaking change
|
||||
/// for any deployed repo. Tracked as a follow-up: see
|
||||
/// <https://github.com/bluesky-social/atproto/blob/main/packages/repo/src/util/mst.ts>
|
||||
/// for the reference implementation to mirror.
|
||||
pub fn encode_key(raw_key: &str) -> String {
|
||||
use base64::Engine;
|
||||
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(raw_key.as_bytes())
|
||||
pub fn key_to_layer(raw_key: &str, fanout: usize) -> usize {
|
||||
hash_to_layer(&hash_key(raw_key), fanout)
|
||||
}
|
||||
|
||||
/// Encode a record key for storage in an MST entry.
|
||||
///
|
||||
/// Per the atproto MST spec
|
||||
/// (<https://atproto.com/specs/data-model-repo#node-data>) the `k`
|
||||
/// field is `base64url(sha256(record_key_utf8_bytes))`. Hashing
|
||||
/// first ties the layer distribution to the cryptographic digest
|
||||
/// of the key — under pre-image resistance, an attacker can't
|
||||
/// craft keys that all land at the maximum layer by sorting their
|
||||
/// bytes a certain way.
|
||||
///
|
||||
/// Decoding returns the raw 32-byte hash bytes; callers that need
|
||||
/// the original key string have to keep it alongside. Cross-crate
|
||||
/// callers passing `vec::Vec<u8>` vs `[u8; 32]` will need a trivial
|
||||
/// .as_slice() conversion at the comparison site.
|
||||
pub fn encode_key(raw_key: &str) -> String {
|
||||
use base64::Engine;
|
||||
base64::engine::general_purpose::URL_SAFE_NO_PAD.encode(hash_key(raw_key))
|
||||
}
|
||||
|
||||
/// Inverse of [`encode_key`]: round-trip the base64url string back
|
||||
/// to the 32-byte SHA-256 hash. Rejects anything that doesn't decode
|
||||
/// to exactly 32 bytes — i.e. catches the old `base64url(raw_key)`
|
||||
/// encoding that predates this commit, which makes it easy to spot
|
||||
/// incompatibilities during migration.
|
||||
pub fn decode_key(encoded: &str) -> Result<Vec<u8>> {
|
||||
use base64::Engine;
|
||||
base64::engine::general_purpose::URL_SAFE_NO_PAD
|
||||
let bytes = base64::engine::general_purpose::URL_SAFE_NO_PAD
|
||||
.decode(encoded.as_bytes())
|
||||
.map_err(|e| anyhow!("invalid base64url key `{encoded}`: {e}"))
|
||||
.map_err(|e| anyhow!("invalid base64url key `{encoded}`: {e}"))?;
|
||||
if bytes.len() != 32 {
|
||||
return Err(anyhow!(
|
||||
"decoded key `{encoded}` is {} bytes; expected 32 (sha256 hash per the atproto MST spec)",
|
||||
bytes.len()
|
||||
));
|
||||
}
|
||||
Ok(bytes)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
|
||||
Reference in New Issue
Block a user