Przeglądaj źródła

Use iteration instead of tail recursion in AVL::lookup

so that a large TTL doesn't smash the stack
Ian Goldberg 1 miesiąc temu
rodzic
commit
01ee882483
1 zmienionych plików z 40 dodań i 40 usunięć
  1. 40 40
      avl.cpp

+ 40 - 40
avl.cpp

@@ -941,48 +941,48 @@ void AVL::insert(MPCTIO &tio, yield_t &yield, const Node &node) {
 
 bool AVL::lookup(MPCTIO &tio, yield_t &yield, RegXS ptr, RegAS key, Duoram<Node>::Flat &A,
     int TTL, RegBS isDummy, Node *ret_node) {
-    if(TTL==0) {
-        // Reconstruct and return isDummy
-        // If we found the key, then isDummy will be true
-        bool found = reconstruct_RegBS(tio, yield, isDummy);
-        return found;
-    }
-
-    RegBS isNotDummy = isDummy ^ (!tio.player());
-    Node cnode = A[ptr];
-    // Compare key
-    CDPF cdpf = tio.cdpf(yield);
-    auto [lt, eq, gt] = cdpf.compare(tio, yield, key - cnode.key, tio.aes_ops());
 
-    // Depending on [lteq, gt] select the next ptr/index as
-    // first AVL_PTR_SIZE bits of cnode.pointers if lteq
-    // next AVL_PTR_SIZE bits of cnode.pointers if gt
-    // (the last 2 bits are balance bits)
-    RegXS left = getAVLLeftPtr(cnode.pointers);
-    RegXS right = getAVLRightPtr(cnode.pointers);
-
-    RegXS next_ptr;
-    mpc_select(tio, yield, next_ptr, gt, left, right, 32);
-
-    RegBS F_found;
-    // If we haven't found the key yet, and the lookup matches the current node key,
-    // then we found the node to return
-
-    // If multiple keys in the tree match the lookup key, this returns the last match.
-    // Extracting the first match would add an extra round here, since the
-    // F_found flag will have to be computed first, then the next two based on F_found
-    // instead of eq
-    run_coroutines(tio,
-        [&tio, &F_found, isNotDummy, eq](yield_t &yield)
-        { mpc_and(tio, yield, F_found, isNotDummy, eq);},
-        [&tio, &ret_node, F_found, &cnode](yield_t &yield)
-        { mpc_select(tio, yield, ret_node->key, F_found, ret_node->key, cnode.key);},
-        [&tio, &ret_node, F_found, &cnode](yield_t &yield)
-        { mpc_select(tio, yield, ret_node->value, F_found, ret_node->value, cnode.value);});
-
-    isDummy^=F_found;
-    bool found = lookup(tio, yield, next_ptr, key, A, TTL-1, isDummy, ret_node);
+    while(TTL > 0) {
+        RegBS isNotDummy = isDummy ^ (!tio.player());
+        Node cnode = A[ptr];
+        // Compare key
+        CDPF cdpf = tio.cdpf(yield);
+        auto [lt, eq, gt] = cdpf.compare(tio, yield, key - cnode.key, tio.aes_ops());
+
+        // Depending on [lteq, gt] select the next ptr/index as
+        // first AVL_PTR_SIZE bits of cnode.pointers if lteq
+        // next AVL_PTR_SIZE bits of cnode.pointers if gt
+        // (the last 2 bits are balance bits)
+        RegXS left = getAVLLeftPtr(cnode.pointers);
+        RegXS right = getAVLRightPtr(cnode.pointers);
+
+        RegXS next_ptr;
+        mpc_select(tio, yield, next_ptr, gt, left, right, 32);
+
+        RegBS F_found;
+        // If we haven't found the key yet, and the lookup matches the current node key,
+        // then we found the node to return
+
+        // If multiple keys in the tree match the lookup key, this returns the last match.
+        // Extracting the first match would add an extra round here, since the
+        // F_found flag will have to be computed first, then the next two based on F_found
+        // instead of eq
+        run_coroutines(tio,
+            [&tio, &F_found, isNotDummy, eq](yield_t &yield)
+            { mpc_and(tio, yield, F_found, isNotDummy, eq);},
+            [&tio, &ret_node, F_found, &cnode](yield_t &yield)
+            { mpc_select(tio, yield, ret_node->key, F_found, ret_node->key, cnode.key);},
+            [&tio, &ret_node, F_found, &cnode](yield_t &yield)
+            { mpc_select(tio, yield, ret_node->value, F_found, ret_node->value, cnode.value);});
+
+        isDummy^=F_found;
+        ptr = next_ptr;
+        --TTL;
+    }
 
+    // Reconstruct and return isDummy
+    // If we found the key, then isDummy will be true
+    bool found = reconstruct_RegBS(tio, yield, isDummy);
     return found;
 }