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

bool_type alloc_bool(var_name)
int_type alloc_int(var_name,bits)

void int_equal_const(i,c)
void int_not_equal_const(i,c)
void ints_equal(i1,i2)
void int_non_negative(i)
void ints_add(sum,i1,i2)
void ints_mult(prod,i1,i2)

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

#include <stdio.h>

#include "wnlib.h"
#include "wnswap.h"

#include "compile.h"


local int next_var=1;


local int alloc_var()
{
  int ret;

  ret = next_var;
  ++next_var;

  return(ret);
}


local int alloc_var_vect(len)

int len;

{
  int ret;

  ret = next_var;
  next_var += len;

  return(ret);
}


bool_type alloc_bool(var_name)

char *var_name;

{
  bool_type ret;
  int var;

  var = alloc_var();

  fprintf(stderr,"%s: bool %d\n",var_name,var);

  ret = (bool_type)malloc(sizeof(struct bool_type_struct));

  ret->var_name = var_name;
  ret->var = var;

  return(ret);
}


int_type alloc_int(var_name,bits)

char *var_name;
int bits;

{
  int_type ret;
  int low_var,hi_var;

  low_var = alloc_var_vect(bits);
  hi_var = low_var+bits;

  fprintf(stderr,"%s: int %d-%d\n",var_name,low_var,hi_var-1);

  ret = (int_type)malloc(sizeof(struct int_type_struct));

  ret->var_name = var_name;
  ret->low_var = low_var;
  ret->hi_var = hi_var;

  return(ret);
}


void bit_equal_const(in,c)

int in;
bool c;

{
  if(c)
  {
    printf("%d\n",in);
  }
  else
  {
    printf("%d\n",-in);
  }
}


void bits_equal(in1,in2)

int in1,in2;

{
  printf("%d %d\n",in1,-in2);
  printf("%d %d\n",-in1,in2);
}


void bits_not_equal(in1,in2)

int in1,in2;

{
  printf("%d %d\n",in1,in2);
  printf("%d %d\n",-in1,-in2);
}


void bit_and(out,in1,in2)

int out,in1,in2;

{
  printf("%d %d\n",in1,-out);
  printf("%d %d\n",in2,-out);
  printf("%d %d %d\n",-in1,-in2,out);
}


void bit_and3(out,in1,in2,in3)

int out,in1,in2,in3;

{
  printf("%d %d\n",in1,-out);
  printf("%d %d\n",in2,-out);
  printf("%d %d\n",in3,-out);
  printf("%d %d %d %d\n",-in1,-in2,-in3,out);
}


void bit_and4(out,in1,in2,in3,in4)

int out,in1,in2,in3,in4;

{
  printf("%d %d\n",in1,-out);
  printf("%d %d\n",in2,-out);
  printf("%d %d\n",in3,-out);
  printf("%d %d\n",in4,-out);
  printf("%d %d %d %d %d\n",-in1,-in2,-in3,-in4,out);
}


void bit_or(out,in1,in2)

int out,in1,in2;

{
  printf("%d %d\n",-in1,out);
  printf("%d %d\n",-in2,out);
  printf("%d %d %d\n",in1,in2,-out);
}


void bit_or3(out,in1,in2,in3)

int out,in1,in2,in3;

{
  printf("%d %d\n",-in1,out);
  printf("%d %d\n",-in2,out);
  printf("%d %d\n",-in3,out);
  printf("%d %d %d %d\n",in1,in2,in3,-out);
}


void bit_or4(out,in1,in2,in3,in4)

int out,in1,in2,in3,in4;

{
  printf("%d %d\n",-in1,out);
  printf("%d %d\n",-in2,out);
  printf("%d %d\n",-in3,out);
  printf("%d %d\n",-in4,out);
  printf("%d %d %d %d %d\n",in1,in2,in3,in4,-out);
}


void bit_xor(in1,in2,in3)

int in1,in2,in3;

{
  printf("%d %d %d\n",in1,in2,-in3);
  printf("%d %d %d\n",in1,-in2,in3);
  printf("%d %d %d\n",-in1,in2,in3);
  printf("%d %d %d\n",-in1,-in2,-in3);
}


void bit_add2(sum,carry,in1,in2)

int sum,carry,in1,in2;

{
  bit_xor(sum,in1,in2);
  bit_and(carry,in1,in2);
}


void bit_add3(sum,carry,in1,in2,in3)

int sum,carry,in1,in2;

{
  int carry1,sum1,carry2;

  carry1 = alloc_var();
  sum1 = alloc_var();
  carry2 = alloc_var();

  bit_add2(sum1,carry1,in1,in2);
  bit_add2(sum,carry2,sum1,in3);
  bit_or(carry,carry1,carry2);
}


void int_equal_const(i,c)

int_type i;
int c;

{
  int v;

  for(v=i->low_var;v<i->hi_var;++v)
  {
    bit_equal_const(v,(c&1));

    c >>= 1;
  }
}


void int_not_equal_const(i,c)

int_type i;
int c;

{
  int v;

  for(v=i->low_var;v<i->hi_var;++v)
  {
    if((c&1) == 0)
    {
      printf("%d ",v);
    }
    else
    {
      printf("%d ",-v);
    }

    c >>= 1;
  }

  printf("\n");
}


void ints_equal(i1,i2)

int_type i1,i2;

{
  int i,len1,len2;

  len1 = i1->hi_var - i1->low_var;
  len2 = i2->hi_var - i2->low_var;

  if(len2 < len1)
  {
    wn_swap(i1,i2,int_type);
    wn_swap(len1,len2,int);
  }

  for(i=0;i<len1;++i)
  {
    bits_equal(i1->low_var+i,i2->low_var+i);
  }
  for(i=len1;i<len2;++i)
  {
    bit_equal_const(i2->low_var+i,0);
  }
}


void int_non_negative(i)

int_type i;

{
  bit_equal_const(i->hi_var-1,0);
}


void ints_add(sum,i1,i2)

int_type sum,i1,i2;

{
  int i,len1,len2,len_sum;
  int last_carry,this_carry;

  len1 = i1->hi_var - i1->low_var;
  len2 = i2->hi_var - i2->low_var;
  len_sum = sum->hi_var - sum->low_var;

  wn_assert(len1 == len2);
  wn_assert(len1 == len_sum);

  last_carry = alloc_var();
  bit_add2(sum->low_var,last_carry,i1->low_var,i2->low_var);

  for(i=1;i<len1;++i)
  {
    this_carry = alloc_var();
    bit_add3(sum->low_var+i,this_carry,last_carry,i1->low_var+i,i2->low_var+i);
    last_carry = this_carry;
  }
}


void ints_mult(prod,i1,i2)

int_type prod,i1,i2;

{
  int *sum_vect;
  int len1,len2,len_prod;
  int j1,j2,new_carry,last_carry,sum,bit_prod;

  len1 = i1->hi_var - i1->low_var;
  len2 = i2->hi_var - i2->low_var;
  len_prod = prod->hi_var - prod->low_var;

  wn_assert(len1+len2 == len_prod);

  if(len1 > len2)
  {
    wn_swap(len1,len2,int);
    wn_swap(i1,i2,int_type);
  }

  sum_vect = (int *)malloc(len1*sizeof(int));

  bit_and(prod->low_var,i1->low_var,i2->low_var);

  for(j1=1;j1<len1;++j1)
  {
    sum_vect[j1-1] = alloc_var();
    bit_and(sum_vect[j1-1],i1->low_var+j1,i2->low_var);
  }
  sum_vect[len1-1] = alloc_var();
  bit_equal_const(sum_vect[len1-1],0);

  for(j2=1;j2<len2-1;++j2)
  {
    bit_prod = alloc_var();
    bit_and(bit_prod,i1->low_var,i2->low_var+j2);
    new_carry = alloc_var();
    bit_add2(prod->low_var+j2,new_carry,bit_prod,sum_vect[0]);

    for(j1=1;j1<len1;++j1)
    {
      last_carry = new_carry;
      bit_prod = alloc_var();
      bit_and(bit_prod,i1->low_var+j1,i2->low_var+j2);
      new_carry = alloc_var();
      sum = alloc_var();
      bit_add3(sum,new_carry,last_carry,bit_prod,sum_vect[j1]);
      sum_vect[j1-1] = sum;
    }

    sum_vect[len1-1] = new_carry;
  }

  j2 = len2-1;

  bit_prod = alloc_var();
  bit_and(bit_prod,i1->low_var,i2->low_var+j2);
  new_carry = alloc_var();
  bit_add2(prod->low_var+j2,new_carry,bit_prod,sum_vect[0]);

  for(j1=1;j1<len1;++j1)
  {
    last_carry = new_carry;
    bit_prod = alloc_var();
    bit_and(bit_prod,i1->low_var+j1,i2->low_var+j2);
    new_carry = alloc_var();
    sum = prod->low_var+j2+j1;
    bit_add3(sum,new_carry,last_carry,bit_prod,sum_vect[j1]);
  }

  bits_equal(prod->low_var+j2+len1,new_carry);

  free(sum_vect);
}




