#include <stdio.h>
#include <stdlib.h>
#include <string.h>

#include "align_standard.h"
#include "ndb_misclib.h"

int check_target_sequence(FILE *report, const char *entity_sequence, const char *target_filename,
        const char *entity_id, const char *target_id)
{
       int  i, length_b, length_a, *index_s, index_t[50000], count = 0, similarity = 0;
       char line[100], sequence[50000];
       ALIGN *sa = NULL, *sb = NULL, *ss = NULL;
       FILE *fp = NULL;

       fp = fopen(target_filename, "r");
       if (fp == NULL) return (-1);

       length_a = 0;
       memset(sequence, 0, 50000);
       while (!feof(fp)) {
            fgets(line, 99, fp);
            if (feof(fp)) break;
	    ndb_clean_string(line);
	    if ((length_a + strlen(line)) >= 50000) {
                 ndb_log_message(NDB_MSG_ERR, "Target sequence larger than 50000 residues.");
		 ndb_close_log();
		 exit (1);
            }
	    length_a += strlen(line);
	    strcat(sequence, line);
       }
       fclose (fp);

       length_b = strlen(entity_sequence);
       index_s = new int[length_b];

       for (i = 0; i < length_b; i++) {
            index_s[i] = find_index(entity_sequence[i]);
       }

       length_a = strlen(sequence);
       for (i = 0; i < length_a; i++) {
            index_t[i] = find_index(sequence[i]);
       }

       _indexa = index_t;
       _indexb = index_s;

       sa = new ALIGN;
       sa->AlignLength = length_a;
       sa->aln = new ALN[sa->AlignLength];
       for (i = 0; i < sa->AlignLength; i++) sa->aln[i].pos[0] = i;
       sb = new ALIGN;
       sb->AlignLength = length_b;
       sb->aln = new ALN[sb->AlignLength];
       for (i = 0; i < sb->AlignLength; i++) sb->aln[i].pos[0] = i;
       ss = alignment_standard(sa, sb, 10, 1, score);

       count = 0;
       for (i = 0; i < ss->AlignLength; i++) {
            if (ss->aln[i].pos[0] >= 0 && ss->aln[i].pos[1] >= 0) {
                 if (sequence[ss->aln[i].pos[0]] == entity_sequence[ss->aln[i].pos[1]])
                      count++;
            }
       }

       similarity = 100 * count / length_b;
       if (count < length_b) {
            fprintf(report, "Sequence Identity between entity %s and target %s = %d%%\n\n",
                        entity_id, target_id, similarity);
	    if (similarity < 80) {
                 print_alignment(report, ss, sequence, entity_sequence);
            }
       }

       delete [] sa->aln; delete sa;
       delete [] sb->aln; delete sb;
       delete [] ss->aln; delete ss;

       delete [] index_s;
       return similarity;
}

static void print_alignment(FILE *fp, ALIGN *ss, const char *_seqa, const char *_seqb)
{
       int i, l, m, n, num_per_line = 60, x, y, u, v;

       l = ss->AlignLength / num_per_line;
       x = ss->AlignLength % num_per_line;
       m = l;
       if (x == 0) m = l - 1;

       for (i = 0; i <= m; i++) {
            n = num_per_line;
            if (i == l) n = x;

	    fprintf(fp, "Target: ");
            for (y = 0; y < n; y++) {
                 u = ss->aln[i * num_per_line + y].pos[0];
                 if (u < 0)
                      fprintf(fp, "-");
                 else fprintf(fp, "%c", _seqa[u]);
            }
            fprintf(fp, "\n");

	    fprintf(fp, "        ");
            for (y = 0; y < n; y++) {
                 u = ss->aln[i * num_per_line + y].pos[0];
                 v = ss->aln[i * num_per_line + y].pos[1];
                 if (u >= 0 && v >= 0 && _seqa[u] == _seqb[v])
                      fprintf(fp, "|");
                 else fprintf(fp, " ");
            }
            fprintf(fp, "\n");

	    fprintf(fp, "Entity: ");
            for (y = 0; y < n; y++) {
                 v = ss->aln[i * num_per_line + y].pos[1];
                 if (v < 0)
                      fprintf(fp, "-");
                 else fprintf(fp, "%c", _seqb[v]);
            }
            fprintf(fp, "\n\n\n");
       }
}

static float score(const int k, const int l)
{
       int index_k, index_l;

       index_k = _indexa[k];
       index_l = _indexb[l];
       if (index_k >= 0 && index_l >= 0)
            return matrix[index_k][index_l];
       else return (-4);
}

static int find_index(char aa)
{
       int i;

       for (i = 0; i < 21; i++) {
            if (aa == AA[i]) return i;
       }
       return (-1);
}

ALIGN *alignment_standard(ALIGN *sa, ALIGN *sb, const float vv, const float uu,
          float (*score)(const int, const int))
{
       int origin = dynamic(sa, sb, vv, uu, score);
       return (mu_trcback(origin, sa, sb));
}

static int dynamic(ALIGN *sa, ALIGN *sb, const float vv, const float uu,
             float (*score)(const int, const int))
{
       int   i, j, k = 0, m, n, r, origin;
       char  *dir = NULL;
       float x, max_val = 0, maximum = -1.0e38;
       RECD  nd, f, *d, *g, dt, dtt;
       RECD  *dd = NULL, *gg = NULL;

       m = sa->AlignLength;
       n = sb->AlignLength;

       dd = new RECD[n + 1];
       gg = new RECD[n + 1];
       dir = new char[m + n + 1];

       i_record = -1;
       origin = adr(0, 0, -1);
       for (i = 0; i < (m + n + 1); i++) dir[i] = 0;
       for (i = 0; i <= n; i++) {
            dd[i].val = 0; dd[i].ptr = origin;
            gg[i].val = 0; gg[i].ptr = origin;
       }

       for (i = 0; i < m; i++) {
            j  = 0;
            d = dd + j;
            g = gg + j;
            d->val = f.val = 0;
            d->ptr = f.ptr = origin;
            r = j - i + m;
            dt = *d;
            for (d++, g++; j < n; d++, g++, r++, j++) {
                 if ((x = (d-1)->val - vv -uu) >= (f.val -= uu)) {
                      f.val = x; f.ptr = (d-1)->ptr;
                 }
                 if ((x = d->val - vv - uu) >= (g->val -= uu)) {
                      g->val = x; g->ptr = d->ptr;
                 }
                 nd.val = maximum;
                 if (f.val > nd.val) nd = f;
                 if (g->val > nd.val) nd = *g;
                 dt.val += (*score)(sa->aln[i].pos[0], sb->aln[j].pos[0]);
                 dtt = *d;
                 *d = dt;
                 if (nd.val > d->val) {
                      *d = nd;
                      dir[r] = 0;
                 } else if (!dir[r]) {
                      d->ptr = adr(i, j, d->ptr);
                      dir[r] = 1;
                 }
                 if (d->val > max_val) {
                      max_val = d->val;
                      k = d->ptr;
                 }
                 dt = dtt;
            }
       }
       origin = adr(m, n, k);

       delete [] dd;
       delete [] gg;
       delete [] dir;

       return origin;
}

static ALIGN *mu_trcback(const int origin, ALIGN *a, ALIGN *b)
{
       LIST *top = NULL, *t = NULL, *p = NULL, *q = NULL;
       int  n1, n2, i, j, m, n, r;

       top = new LIST;
       top->tm = dr[origin].m;
       top->tn = dr[origin].n;
       top->next = NULL;
       i = dr[origin].pp;
       while (i >= 0) {
            t = new LIST;
            t->tm = dr[i].m;
            t->tn = dr[i].n;
            t->next = top;
            top = t;
            i = dr[i].pp;
       }

       p = top;
       n1 = n2 = 0;
       while (p && (q = p->next)) {
            if ((i = p->tm) - (j = p->tn) > (m = q->tm) - (n = q->tn) &&
                (r = m - i + j) != j) {
                 t = new LIST;
                 t->tm = m;
                 t->tn = r;
                 n1 += (n -r );
                 p = mid_insert(p, t, q);
            }

            if ((i = p->tm) - (j = p->tn) < (m = q->tm) - (n = q->tn) &&
                (r = n + i - j) != i) {
                 t = new LIST;
                 t->tm = r;
                 t->tn = n;
                 n2 += (m - r);
                 p = mid_insert(p, t, q);
            }
            p = q;
       }

       m = a->AlignLength + n1;
       n = b->AlignLength + n2;
       ALIGN *ss = new ALIGN;
       ss->AlignLength = MAX(m, n);
       ss->aln = (ALN *) new ALN[ss->AlignLength];

       Size_dr = 0;
       free ((void *) dr);
       dr = NULL;

       track(ss, top, a, b);
       delete_list(top);
       return ss;
}

static void track(ALIGN *aa, LIST *top, ALIGN *sa, ALIGN *sb)
{
       int  i, x, y, m, n, r;
       LIST *p = NULL, *q = NULL;

       p = top; r = 0;
       while (p && (q = p->next)) {
            if ((x = p->tm) - (y = p->tn) == (m = q->tm) - (n = q->tn)) {
                 for (i = 0; i < m - x; i++) {
                      aa->aln[r + i].pos[0] = sa->aln[x + i].pos[0];
                      aa->aln[r + i].pos[1] = sb->aln[y + i].pos[0];
                 }
                 r += (m - x);
            } else if(x - y < m - n) {
                 for (i = 0; i < m - x; i++) {
                      aa->aln[r + i].pos[0] = sa->aln[x + i].pos[0];
                      aa->aln[r + i].pos[1] = -1;
                 }
                 r += (m - x);
            } else {
                 for (i = 0; i < n - y; i++) {
                      aa->aln[r + i].pos[0] = -1;
                      aa->aln[r + i].pos[1] = sb->aln[y + i].pos[0];
                 }
                 r += (n - y);
            }
            p = q;
       }
}

static int adr(int m, int n, int pp)
{
       i_record++;

       if (Size_dr == 0) {
           dr = (SAVE *) calloc(1000, sizeof(SAVE));
           Size_dr = 1000;
       } else if ((i_record + 1) > Size_dr) {
           Size_dr += 200;
           dr = (SAVE *) realloc(dr, Size_dr * sizeof(SAVE));
       }

       dr[i_record].m = m; dr[i_record].n = n; dr[i_record].pp = pp;
       return (i_record);
}

static LIST *mid_insert(LIST *top, LIST *insert, LIST *last)
{
       top->next=insert;
       insert->next=last;
       return top;
}

static void delete_list(LIST *node)
{
       if (node == NULL) return;

       delete_list(node->next);
       delete node;
       node = NULL;
}

