diff --git a/crates/at-mst/src/tree.rs b/crates/at-mst/src/tree.rs
index 2e8d708..30b0bfc 100644
--- a/crates/at-mst/src/tree.rs
+++ b/crates/at-mst/src/tree.rs
@@ -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> {
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 > {
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 {
- Some(sub) => self.get_in_tree(sub, key_hash),
- None => Ok(None),
- },
- Ordering::Equal => Ok(Some(entries[0].value)),
+ 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 => {
+ 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 {
- Some(sub) => return self.get_in_tree(sub, key_hash),
- None => return Ok(None),
- },
+ 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,
fanout: usize,
) -> Result {
- 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>,
new_blocks: &mut HashMap>,
@@ -650,31 +702,36 @@ impl Mst {
entries: &[MstEntry],
raw_key: &str,
fanout: usize,
- ) -> Result<(Option, Option)> {
- let key_bytes = raw_key.as_bytes();
- let pos = find_position(entries, key_bytes)?;
+ ) -> Result<(Option, Option, Option)> {
+ 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,
raw_key: &str,
fanout: usize,
- ) -> Result<(Option, Option)> {
+ ) -> Result<(Option, Option, Option)> {
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> {
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,50 +1032,38 @@ 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 {
- if entries.is_empty() {
- return Ok(Pos::AfterLast);
- }
- let first_key = decode_key(&entries[0].key)?;
- if key_bytes < 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() {
- return Ok(Pos::Between(i));
+ fn find_position(entries: &[MstEntry], key_hash: &[u8]) -> Result {
+ if entries.is_empty() {
+ return Ok(Pos::AfterLast);
}
- }
- Ok(Pos::AfterLast)
+ let first_key = decode_key(&entries[0].key)?;
+ 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_hash < ek.as_slice() {
+ return Ok(Pos::Between(i));
+ }
+ }
+ Ok(Pos::AfterLast)
}
/// 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, 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, 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"
diff --git a/crates/at-mst/src/util.rs b/crates/at-mst/src/util.rs
index e646b83..d27bb43 100644
--- a/crates/at-mst/src/util.rs
+++ b/crates/at-mst/src/util.rs
@@ -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
-/// () 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
-///
-/// 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
+/// () 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` 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> {
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)]