/********************************************************************
 *
 * Implementation of the adaptive algorithm for determining
 * strings from substrings.
 * Suffix tree version.
 *
 * Suffix tree-related functions.
 *
 * (C) February - December 1994, Dimitris Margaritis, Steven Skiena
 *
 * $Id: suffix_tree.c,v 2.2 1995/04/18 21:34:53 dmarg Exp dmarg $
 *
 ********************************************************************/

#include "includes.h"

/********************************************************************
 *
 * Construct a new suffix-tree node.  Return a pointer to it.
 *
 ********************************************************************/

NODE *create_node()
{
NODE *newnode;
ARC **tmp_ptr;
int i;
extern unsigned q;     /* The alphabet size. */
/* static unsigned numnodes = 0; */

    newnode = (NODE *) getmem(sizeof(NODE));
    newnode -> suff_link = (NODE *) NULL;
    newnode -> arcarray = (ARC **) getmem((q + 1) * sizeof(ARC *));
    for (i = 0, tmp_ptr = newnode -> arcarray; i <= q; i++)
        *tmp_ptr++ = (ARC *) NULL;
/*
    printf("numnodes = %u\n", ++numnodes);
*/
    return newnode;
}

/********************************************************************
 *
 * Count how many strings of length "newlen" are consistent with
 * strings of length "limit" in the prefix tree.
 *
 * Recursive version.
 *
 ********************************************************************/

int stree_count_strings_aux(NODE *subroot, unsigned skirtstart,
                                 unsigned skirtlength, unsigned currlen,
                                 unsigned newlen, unsigned limit,
                                 int sumlimit)
{
int i, arclength;
int sum, count;
unsigned subrootdepth;
extern unsigned q, sequence_len;     /* The alphabet size. */
extern char *sequence;
ARC *arc;
extern NODE *root;
BOOL done;
#ifdef DEBUG
static char *str = NULL;
#endif

#ifdef DEBUG
    if (str == NULL)
        str = getmem(sequence_len * sizeof(char));
#endif

    sum = 0;

    done = FALSE;
    while ( ! done ) {

        subrootdepth = subroot -> depth;

        if (currlen >= newlen) {
            sum = 1;
            done = TRUE;
#ifdef MOREDEBUG
            str[newlen] = '\0';
            printf("string: %s\n", str);
#endif
        } else if (skirtlength != 0) {
            arc = (subroot -> arcarray)[sequence[skirtstart] - '0'];
            arclength = (arc -> last) - (arc -> first) + 1;

            if (subrootdepth + arclength <= limit) {
                if (skirtlength >= arclength) {
                    subroot = arc -> pointsto;
                    skirtstart += arclength;
                    skirtlength -= arclength;
                } else {  /* skirtlength < arclength */
#ifdef MOREDEBUG
                    strncpy(str + currlen,
                            sequence + skirtstart + skirtlength,
                            arclength - skirtlength);
#endif
                    subroot = arc -> pointsto;
                    skirtstart = arc -> last + 1;
                    currlen += arclength - skirtlength;
                    skirtlength = 0;
                }
            } else {    /* subrootdepth + arclength > limit */
                if (subroot == root) {
#ifdef MOREDEBUG
                    strncpy(str + currlen,
                            sequence + skirtstart + skirtlength,
                            limit - skirtlength);
#endif
                    skirtstart++;
                    currlen += limit - skirtlength;
                    skirtlength = limit - 1;
                } else {  /* subroot != root */
#ifdef MOREDEBUG
                    strncpy(str + currlen, sequence + skirtstart + skirtlength,
                            limit - (subrootdepth + skirtlength));
#endif
                    subroot = subroot -> suff_link;
                    currlen += limit - (subrootdepth + skirtlength);
                    skirtlength = limit - subrootdepth;
                }
            }
        } else {   /* skirtlength == 0 */
            if (subrootdepth == limit)
                subroot = subroot -> suff_link;
            for (i = 0; i < q; i++) {
                if ((arc = (subroot -> arcarray)[i]) != NULL) {
                    if (sumlimit == STREE_COUNT_INFINITY) {
#ifdef MOREDEBUG
                        strncpy(str + currlen, sequence + (arc -> first), 1);
#endif
                        sum += stree_count_strings_aux(subroot,
                                arc -> first, 1, currlen + 1,
                                newlen, limit, sumlimit);
                    } else {
#ifdef MOREDEBUG
                        strncpy(str + currlen, sequence + (arc -> first), 1);
#endif
                        count = stree_count_strings_aux(subroot,
                                arc -> first, 1, currlen + 1,
                                newlen, limit, sumlimit - sum);
                        if (count == -1 || (sum += count) > sumlimit) {
                            sum = -1;
                            break;
                        }
                    }
                }
            }
            done = TRUE;
        }
    }

    return sum;
}

/* ---------------------------------------------------------------- */

int stree_count_strings(unsigned newlen, unsigned limit,
                             int sumlimit)
{
extern NODE *root;

    return stree_count_strings_aux(root, 0, 0, 0, newlen, limit,
                                   sumlimit);
}

/********************************************************************/

void print_suffix_tree(NODE *subroot, unsigned ident)
{
ARC *arc_ptr;
NODE *node_ptr;
extern NODE *root;
static char *tmpstr = NULL;
extern unsigned sequence_len, q;
extern char *sequence;
int i, j, length;


    if (tmpstr == NULL) {   /* First time called, allocate memory. */
        tmpstr = getmem((sequence_len + 1) * sizeof(char));
        printf("root = %p\n", root);
    }

    for (i = 0; i <= q; i++)
        if ((arc_ptr = (subroot -> arcarray)[i]) != NULL) {
            for (j = 0; j < ident; j++)
                sprintf(tmpstr + j, " ");
            length = (arc_ptr -> last) - (arc_ptr -> first) + 1;
            strncpy(tmpstr + j, sequence + (arc_ptr -> first), length);
            (tmpstr + j)[length] = '\0';
            printf(tmpstr);
            printf(" \t(node = %p, ", arc_ptr -> pointsto);
            printf("depth = %d, ", (arc_ptr -> pointsto) -> depth);
            printf("suff_link = %p)\n", (arc_ptr -> pointsto) -> suff_link);
            if ((node_ptr = arc_ptr -> pointsto) != NULL)
                print_suffix_tree(node_ptr, ident + 4);  /* Recursion. */
        }
}


/********************************************************************
 *
 * Print out the strings of length "newlen" that are consistent
 * with strings of length "len" in the prefix tree.
 *
 ********************************************************************/
/*  commented out
void ptree_print_strings(NODE *subroot, unsigned currlen,
                         unsigned newlen, char *str, unsigned len)
{
int i;
extern unsigned q;     /* The alphabet size. * /

    if (currlen == newlen + 1) {
        str[currlen] = '\0';
        printf("%s\n", str);
    } else if (subroot -> depth < len) {
        str[currlen] = '\0';
        printf("%s\n", str);
        for (i = 0; i <= q; i++)
            if ((subroot -> nodearray)[i] != (NODE *) NULL) {
                str[currlen] = '0' + i;
                ptree_print_strings(subroot -> nodearray[i],
                                      currlen + 1, newlen, str, len);
            }
    } else
        ptree_print_strings(subroot -> suff_link, currlen, newlen, str, len);

}
*/
/********************************************************************/

BOOL end_point(NODE *s, int k, int l, char a)
{
char b;
ARC *arc;
int k1;
extern char *sequence;

    b = (k <= l) ? sequence[k] : a;

    if ((arc = (s -> arcarray)[b - '0']) == NULL)
        return FALSE;
    else {
        k1 = arc -> first;
        if (sequence[k1 + (l - k) + 1] == a)
            return TRUE;
        else
            return FALSE;
    }
}

/********************************************************************/

ARC *create_arc(int start, int end, NODE *from, NODE *to)
{
ARC *arc;
extern char *sequence;

    arc = (ARC *) getmem(sizeof(ARC));
    arc -> first = start;
    arc -> last = end;
    arc -> comesfrom = from;
    arc -> pointsto = to;
    (from -> arcarray)[sequence[start] - '0'] = arc;

    return arc;
}

/********************************************************************/

NODE *stree_create_node(NODE *s, int k, int l)
{
ARC *arc;
NODE *r, *s1;
int k1, l1;
extern char *sequence;

    if (l < k)
        return s;
    else {
        if ((arc = (s -> arcarray)[sequence[k] - '0']) == NULL) {
            fprintf(stderr, "stree_create_node(): error: ");
            fprintf(stderr, "no arc found\n");
            exit(1);
        }
        k1 = arc -> first;
        l1 = arc -> last;
        if ((s1 = arc -> pointsto) == NULL) {
            fprintf(stderr, "stree_create_node(): error: ");
            fprintf(stderr, "arc points to NULL node\n");
            exit(1);
        }

        free(arc);  /* Destroy arc. */

        r = create_node();
        (void) create_arc(k1, k1 + (l - k), s, r);

        (void) create_arc(k1 + (l - k + 1), l1, r, s1);

        return r;
    }

}
        

/********************************************************************/

/* Call-by-reference function because it returns two values. */

void canonise(NODE *s, int k, int l, NODE **s_result,
              int *k_result)
{
ARC *arc;
int k1, l1;
NODE *s1;
extern char *sequence;

    if (l < k) {
        *s_result = s;
        *k_result = k;
    } else {
        if ((arc = (s -> arcarray)[sequence[k] - '0']) == NULL) {
            fprintf(stderr, "canonise(): error: arc not found for ");
            fprintf(stderr, " char '%c'\n", sequence[k]);
            exit(1);
        }
        if ((s1 = arc -> pointsto) == NULL) {
            fprintf(stderr, "stree_create_node(): error: ");
            fprintf(stderr, "arc points to NULL node\n");
            exit(1);
        }
        k1 = arc -> first;
        l1 = arc -> last;
        while (l1 - k1 <= l - k) {
            k += l1 - k1 + 1;
            s = s1;
            if (k <= l) {
                if ((arc = (s -> arcarray)[sequence[k] - '0']) == NULL) {
                    fprintf(stderr, "canonise(): error: arc not found for ");
                    fprintf(stderr, "char '%c'\n", sequence[k]);
                    exit(1);
                }
                if ((s1 = arc -> pointsto) == NULL) {
                    fprintf(stderr, "stree_create_node(): error: ");
                    fprintf(stderr, "arc points to NULL node\n");
                    exit(1);
                }
                k1 = arc -> first;
                l1 = arc -> last;
            }
        }

        *s_result = s;
        *k_result = k;
    }
}

/********************************************************************/

void update(NODE *s, int k, int i, NODE **s_result, int *k_result)
{
NODE *oldr, *r, *r1, *s_result1;
extern NODE *root;
ARC *arc;
extern char *sequence;
int k_result1;

    oldr = root;
    while ( ! end_point(s, k, i - 1, sequence[i]) ) {
        r = stree_create_node(s, k, i - 1);
        r1 = create_node();
        arc = create_arc(i, STREE_INFINITY, r, r1);
        if (oldr != root)
            oldr -> suff_link = r;
        oldr = r;
        s = s -> suff_link;

        if (s == NULL) printf("error!\n");

        canonise(s, k, i - 1, &s_result1, &k_result1);
        s = s_result1;
        k = k_result1;
    }

    if (oldr != root)
        oldr -> suff_link = s;

    *s_result = s;
    *k_result = k;

}
   
/********************************************************************/

NODE *find_suffix(int start, int end)
{
extern NODE *root;
NODE *s_result;
int k_result_dummy;

    canonise(root, start + 1, end, &s_result, &k_result_dummy);
    return s_result;

}

/********************************************************************/

void fix_leaves(NODE *s, int length, int end)
{
extern unsigned q;
BOOL is_leaf;
int i;
ARC *arc;

    is_leaf = TRUE;
    for (i = 0; i <= q; i++)
        if ((arc = (s -> arcarray)[i]) != NULL) {
            is_leaf = FALSE;
            fix_leaves(arc -> pointsto,
                       length + (arc -> last) - (arc -> first) + 1,
                       arc -> last);
        }

    if (is_leaf)
        s -> suff_link = find_suffix(end - length + 1, end);
}

/********************************************************************/

void remove_end_markers(ARC **incoming, NODE *s)
{
extern unsigned q;
BOOL is_leaf;
int i, arclength;
ARC *arc;

    is_leaf = TRUE;
    for (i = 0; i <= q; i++)
        if ((arc = (s -> arcarray)[i]) != NULL) {
            is_leaf = FALSE;
            remove_end_markers(&((s -> arcarray)[i]), arc -> pointsto);
        }

    if (is_leaf) {
        arclength = ((*incoming) -> last) - ((*incoming) -> first) + 1;
        if (arclength == 1) { /* This is a single end-marker node. */
            free(s);          /* Destroy node. */
            free(*incoming);  /* Destroy arc, */
            *incoming = NULL; /* and pointer to it. */
        } else {              /* Else remone the end marker from the arc. */
            ((*incoming) -> last)--;
            (s -> depth)--;
        }
    }
}

/********************************************************************/

void fix_tree()
{
extern NODE *root;
extern unsigned q;
int i;
ARC *arc;

    for (i = 0; i <= q; i++)
        if ((arc = (root -> arcarray)[i]) != NULL)
            remove_end_markers(&((root -> arcarray)[i]), arc -> pointsto);

    for (i = 0; i <= q; i++)
        if ((arc = (root -> arcarray)[i]) != NULL)
            fix_leaves(arc -> pointsto,
                       (arc -> last) - (arc -> first) + 1, arc -> last);

}

/********************************************************************/

void fix_infinities_and_depths(NODE *subroot, int dep)
{
int i;
ARC *arc;
NODE *node;
extern unsigned sequence_len, q;

    subroot -> depth = dep;
    for (i = 0; i <= q; i++)
        if ((arc = (subroot -> arcarray)[i]) != NULL) {
            if (arc -> last == STREE_INFINITY)
                arc -> last = sequence_len - 1;  /* "-2" to exclude */
                                                 /* ending marker */
            if ((node = arc -> pointsto) != NULL)
                fix_infinities_and_depths(node,
                               dep + (arc -> last) - (arc -> first) + 1);
        }
}

/********************************************************************/

NODE *create_suffix_tree()
{
extern NODE *root;
NODE *star, *s, *s_result;
int i, k, j, k_result;
extern unsigned sequence_len;
ARC *arc;

    root = create_node();
    star = create_node();

    for (j = 0; j < sequence_len; j++)
        arc = create_arc(j, j, star, root);

    root -> suff_link = star;

    s = root;
    k = 0;
    for (i = 0; i < sequence_len; i++) {
        update(s, k, i, &s_result, &k_result);
        canonise(s_result, k_result, i, &s, &k);
    }

  /* Fix all "infinite" numbers to sequence_len-1. */

    fix_infinities_and_depths(root, 0);

  /* Fix all the leaves' suffix links (normally NULLed by the algorithm). */

    fix_tree();

    return root;
}

/********************************************************************/
