fix: tree hash type discriminant and leaf count from rightmost node
RFC 9420 Section 7.9 tree hash input includes a type discriminant byte: - Leaf: H(uint8(1) || uint32(leaf_index) || optional<LeafNode>) - Parent: H(uint8(2) || optional<ParentNode> || opaque left<V> || opaque right<V>) Also fixed RatchetTree leaf count computation to use the rightmost non-blank node position instead of total serialized node count, since trees may be serialized with trailing blank nodes. Test results: 35/41 passing (85%). https://claude.ai/code/session_01NocQDWj2Y92FugjfgazzL3
This commit is contained in:
+19
-4
@@ -146,13 +146,16 @@ class RatchetTree(
|
|||||||
/**
|
/**
|
||||||
* Recursive tree hash computation per RFC 9420 Section 7.9.
|
* Recursive tree hash computation per RFC 9420 Section 7.9.
|
||||||
*
|
*
|
||||||
* Leaf: H(uint32(leaf_index) || optional<LeafNode>)
|
* Leaf: H(uint8(1) || uint32(leaf_index) || optional<LeafNode>)
|
||||||
* Parent: H(optional<ParentNode> || opaque left_hash<V> || opaque right_hash<V>)
|
* Parent: H(uint8(2) || optional<ParentNode> || opaque left_hash<V> || opaque right_hash<V>)
|
||||||
|
*
|
||||||
|
* The uint8 type discriminant (1=leaf, 2=parent) is part of the hash input.
|
||||||
*/
|
*/
|
||||||
private fun treeHashNode(nodeIndex: Int): ByteArray {
|
private fun treeHashNode(nodeIndex: Int): ByteArray {
|
||||||
if (BinaryTree.isLeaf(nodeIndex)) {
|
if (BinaryTree.isLeaf(nodeIndex)) {
|
||||||
val leafIndex = BinaryTree.nodeToLeaf(nodeIndex)
|
val leafIndex = BinaryTree.nodeToLeaf(nodeIndex)
|
||||||
val writer = TlsWriter()
|
val writer = TlsWriter()
|
||||||
|
writer.putUint8(1) // TreeHashInput::Leaf discriminant
|
||||||
writer.putUint32(leafIndex.toLong())
|
writer.putUint32(leafIndex.toLong())
|
||||||
val leaf = getNode(nodeIndex)
|
val leaf = getNode(nodeIndex)
|
||||||
if (leaf != null) {
|
if (leaf != null) {
|
||||||
@@ -168,6 +171,7 @@ class RatchetTree(
|
|||||||
val rightHash = treeHashNode(BinaryTree.right(nodeIndex))
|
val rightHash = treeHashNode(BinaryTree.right(nodeIndex))
|
||||||
|
|
||||||
val writer = TlsWriter()
|
val writer = TlsWriter()
|
||||||
|
writer.putUint8(2) // TreeHashInput::Parent discriminant
|
||||||
val parent = getNode(nodeIndex)
|
val parent = getNode(nodeIndex)
|
||||||
if (parent != null) {
|
if (parent != null) {
|
||||||
writer.putUint8(1) // present
|
writer.putUint8(1) // present
|
||||||
@@ -319,8 +323,19 @@ class RatchetTree(
|
|||||||
|
|
||||||
val tree = RatchetTree()
|
val tree = RatchetTree()
|
||||||
tree.nodes.addAll(nodesList)
|
tree.nodes.addAll(nodesList)
|
||||||
// Compute leaf count from node count: nodes = 2*leaves - 1
|
// Compute leaf count from the rightmost non-blank node.
|
||||||
tree._leafCount = (nodesList.size + 1) / 2
|
// Trees may be serialized with trailing blank nodes that should be trimmed.
|
||||||
|
var rightmost = nodesList.size - 1
|
||||||
|
while (rightmost >= 0 && nodesList[rightmost] == null) {
|
||||||
|
rightmost--
|
||||||
|
}
|
||||||
|
// The rightmost non-blank node determines the minimum tree size
|
||||||
|
tree._leafCount =
|
||||||
|
if (rightmost < 0) {
|
||||||
|
0
|
||||||
|
} else {
|
||||||
|
(rightmost / 2) + 1
|
||||||
|
}
|
||||||
return tree
|
return tree
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user