/*
Discrete Fourier Transform on complex numbers
compilation: cc -o cdft cdft.c -lm
*/

#include <stdio.h>
#include <math.h>

typedef struct cnum
{
    float r; // real part
    float i; // imaginary part
} cnum;

/*
Compute forward DFT for complex numbers

in - signal in time domain
out - signal in frequency domain
N - number of complex numbers in in, out

out[k] = sum(n=0,N-1): in[n] * e^(-j*2pi*k/N*n)
*/
void cdft(cnum* in, cnum* out, int N)
{
    for(int n=0; n<N; ++n)
    {
        out[n].r = 0.0;
        out[n].i = 0.0;
    }

    // for every frequency
    for(int k=0; k<N; ++k)
    {
        // for every point
        for(int n=0; n<N; ++n)
        {
            float ar = in[n].r;
            float ai = in[n].i;
            float br = cosf(2*M_PI*k*n/N);
            float bi = -sinf(2*M_PI*k*n/N); // notice minus
            // (ar + i * ai) * (br + i* bi) = (ar*br - ai*bi) + i (ar*bi + ai*br)
            out[k].r += ar*br - ai*bi;
            out[k].i += ar*bi + ai*br;
        }
    }
}

/*
Compute inverse DFT for complex numbers

in - signal in frequency domain
out - signal in time domain
N - number of points in in, out

out[n] = 1/N * sum(k=0,N-1): in[k] * e^(j*2pi*k/N*n)
*/
void icdft(cnum *in, cnum *out, int N)
{
    for(int n=0; n<N; ++n)
    {
        out[n].r = 0.0;
        out[n].i = 0.0;
    }

    // for every point
    for(int n=0; n<N; ++n)
    {
        // for every frequency
        for(int k=0; k<N; ++k)
        {
            float ar = in[k].r;
            float ai = in[k].i;
            float br = cosf(2*M_PI*k*n/N);
            float bi = sinf(2*M_PI*k*n/N);
            out[n].r += ar*br - ai*bi;
            out[n].i += ar*bi + ai*br;
        }
        // divide out[n] by N
        out[n].r /= N;
        out[n].i /= N;
    }
}

int main(void)
{
    cnum in[5], out[5], in2[5];

    // create some input signal
    in[0].r = 1.0; // 1+2i
    in[0].i = 2.0;
    in[1].r = 3.0; // 3+4i
    in[1].i = 4.0;
    in[2].r = 2.0; // 2+4i
    in[2].i = 4.0;
    in[3].r = 3.0; // 3+7i
    in[3].i = 7.0;
    in[4].r = 2.0; // 2+5i
    in[4].i = 5.0;

    // compute forward dft
    cdft(in, out, 5);

    // display result of forward dft
    printf("%f%+fi\n", out[0].r, out[0].i);
    printf("%f%+fi\n", out[1].r, out[1].i);
    printf("%f%+fi\n", out[2].r, out[2].i);
    printf("%f%+fi\n", out[3].r, out[3].i);
    printf("%f%+fi\n", out[4].r, out[4].i);

    // compute inverse dft = restore signal from frequencies
    icdft(out, in2, 5);

    // display restored signal
    printf("\nidft:\n");
    printf("%f%+fi\n", in2[0].r, in2[0].i);
    printf("%f%+fi\n", in2[1].r, in2[1].i);
    printf("%f%+fi\n", in2[2].r, in2[2].i);
    printf("%f%+fi\n", in2[3].r, in2[3].i);
    printf("%f%+fi\n", in2[4].r, in2[4].i);

    return 0;
}
