3 ms·
Lists are recursive, right. Better than double pointer indirection is to have the function return the new head. Note that the following code is not functional,
by scscsc 14y ago
Lists are recursive, right. Better than double pointer indirection is to have the function return the new head. Note that the following code is not functional, but returning the new head makes it cleaner.
node *remove_if(node *head, remove_fn rm)
{
if (rm(head)) {
return remove_if(head->next);
}
head->next = remove_if(head->next);
return head;
}
A good use of this technique (returning pointers to the resulting node) is the following code which very cleanly implements AVL trees in C (from http://infoarena.ro/multe-smenuri-de-programare-in-cc-si-nu-numai http://infoarena.ro/multe-smenuri-de-programare-in-cc-si-nu-... -- scroll to the end (sorry, main article in Romanian)):
#define max(a, b) ((a) > (b) ? (a) : (b))
#define geth(n) (n->h = 1 + max(n->l->h, n->r->h))
struct node
{
int key, h;
struct node *l, *r;
} *R, *NIL;
typedef struct node node;
void init(void)
{
R = NIL = (node *) malloc(sizeof(node));
NIL->key = NIL->h = 0,
NIL->l = NIL->r = NULL;
}
node* rotleft(node *n)
{
node *t = n->l;
n->l = t->r, t->r = n,
geth(n), geth(t);
return t;
}
node* rotright(node *n)
{
node *t = n->r;
n->r = t->l, t->l = n,
geth(n), geth(t);
return t;
}
node* balance(node *n)
{
geth(n);
if (n->l->h > n->r->h + 1)
{
if (n->l->r->h > n->l->l->h)
n->l = rotright(n->l);
n = rotleft(n);
}
else
if (n->r->h > n->l->h + 1)
{
if (n->r->l->h > n->r->r->h)
n->r = rotleft(n->r);
n = rotright(n);
}
return n;
}
node* insert(node *n, int key)
{
if (n == NIL)
{
n = (node *) malloc(sizeof(node));
n->key = key, n->h = 1, n->l = n->r = NIL;
return n;
}
if (key < n->key)
n->l = insert(n->l, key);
else
n->r = insert(n->r, key);
return balance(n);
}
node* erase(node *n, int key)
{
node *t;
if (n == NIL) return n;
if (n->key == key)
{
if (n->l == NIL || n->r == NIL)
{
t = n->l == NIL ? n->r : n->l;
free(n); return t;
}
else
{
for (t = n->r; t->l != NIL; t = t->l);
n->key = t->key,
n->r = erase(n->r, t->key);
return balance(n);
}
}
if (key < n->key)
n->l = erase(n->l, key);
else
n->r = erase(n->r, key);
return balance(n);
}
int search(node *n, int key)
{
if (n == NIL) return 0;
if (n->key == key) return 1;
if (key < n->key)
return search(n->l, key);
else
return search(n->r, key);
}