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

static double getarg(const char *s)
{
    double a;
    char *end;
    a = strtod(s, &end);
    if (!*s || *end)
    {
        fprintf(stderr, "Unrecognized floating point number `%s'.\n", s);
        exit(1);
    }
    return a;
}

int main(int argc, char *argv[])
{
    if (argc != 5)
    {
        fprintf(stderr, "Usage: %s a b c d\nCalculate all solutions of a cubic with these real coefficients.\n", argv[0]);
        return 1;
    }

    double a, b, c, d;
    a = getarg(argv[1]);
    b = getarg(argv[2]);
    c = getarg(argv[3]);
    d = getarg(argv[4]);

    printf("Equation: %g x³ %+g x² %+g x %+g = 0\n", a, b, c, d);
    if (a == 0.0)
    {
        /* quadratic equation */
        if (b == 0)
        {
            /* linear equation */
            if (c == 0)
            {
                /* constant equation? */
                if (d == 0)
                {
                    printf("True statement\n");
                }
                else
                {
                    printf("False statement\n");
                }
            }
            else
            {
                printf("x = %g\n", -d/c);
            }
        }
        else
        {
            c /= b;
            d /= b;
            printf("Normalized: x² %+g x %+g = 0\n", c, d);
            double m = -c / 2;
            double disc = m * m - d;
            if (disc < 0)
            {
                printf("Two complex roots: x = %g ± %g i\n", m, sqrt(-disc));
            }
            else if (disc == 0)
            {
                printf("A double root: x = %g\n", m);
            }
            else
            {
                printf("Two real roots: x₁ = %g, x₂ = %g\n", m - sqrt(disc), m + sqrt(disc));
            }
        }
    }
    else
    {
        b /= a;
        c /= a;
        d /= a;
        printf("Normalized: x³ %+g x² %+g x %+g = 0\n", b, c, d);
        double move = -b / 3;
        /* substitute x = t - b/3a. But a is already normalized to 1. So then we get (t - b/3)³ + b (t - b/3)² + c (t - b/3) + d = 0
         * t³ - 3 b/3 t² + 3 b²/9 t - b³/27 + b (t² - 2/3 b t + b²/9) + ct - 1/3 bc + d = 0
         * t³ - b t² + b²/3 t - b³/27 + bt² - 2/3 b² t + b³/9 + ct - 1/3 bc + d = 0
         * t³ + (-b + b)t² + (b²/3 - 2/3 b² + c + d)
         */
        double p = -b * b / 3 + c;
        double q = 2/27.0 * b * b * b - b * c / 3 + d;

        printf("Depressed: t³ %+g t %+g = 0\n", p, q);
        if (p == 0)
        {
            double t = cbrt(-q);
            /* t³ + q = 0 -> t³ = -q -> t = -³√1 ³√q */
            printf("x₁ = %g\n", t + move );
            double x = t + move;
            double complex z = d + x * (c + x * (b + x));
            printf("Test expression: %g %+g i\n", creal(z), cimag(z));
            /* ³√1 = cos(2π/3) ± i sin(2π/3) = -1/2 ± i √3/2 */
            z = -t/2 + move + I * sqrt(3)*t / 2;
            printf("x₂ = %g %+g i\n", creal(z), cimag(z));
            z = d + z * (c + z * (b + z));
            printf("Test expression: %g %+g i\n", creal(z), cimag(z));
            z = -t/2 + move + I * (-sqrt(3)*t/2);
            printf("x₃ = %g %+g i\n", creal(z), cimag(z));
            z = d + z * (c + z * (b + z));
            printf("Test expression: %g %+g i \n", creal(z), cimag(z));
        }
        else if (q == 0)
        {
            /* t³ + pt = 0 -> t(t² + p) = 0 -> t₁ = 0, t₂ = -√p, t₃ = √p */
            printf("x₁ = %g\n", move);
            double z = d + move * (c + move * (b + move));
            printf("Test expression: %g\n", z);
            if (p > 0)
            {
                double complex z = move - I * sqrt(p);
                printf("x₂ = %g %+g i\n", creal(z), cimag(z));
                z = d + z * (c + z * (b + z));
                printf("Test expression: %g %+g i\n", creal(z), cimag(z));
                z = move + I * sqrt(p);
                printf("x₃ = %g %+g i\n", creal(z), cimag(z));
                z = d + z * (c + z * (b + z));
                printf("Test expression: %g %+g i\n", creal(z), cimag(z));
            }
            else
            {
                z = move - sqrt(-p);
                printf("x₂ = %g\n", z);
                z = d + z * (c + z * (b + z));
                printf("Test expression: %g\n", z);
                z = move + sqrt(-p);
                printf("x₃ = %g\n", z);
                z = d + z * (c + z * (b + z));
                printf("Test expression: %g\n", z);
            }
        }
        else
        {
            /* let u + v = x. Then x³ = u³ + 3u²v + 3uv² + v³ = -px - q
             * So u³ + v³ + 3uv(u + v) = -px - q.
             * So u³ + v³ + 3uvx = -px - q. So
             * I. u³ + v³ = -q
             * II. 3uv = -p
             * So v = -p/3u
             * Plug into I.: u³ - p³/(27u³) = -q
             * u⁶ - p³/27 = -qu³
             * u⁶ + qu³ - p³/27 = 0
             * So a quadratic in u³.
             * u³ = -q/2 ± √(q²/4 + p³/27)
             * Due to symmetry, any choice for u³ means v³ is the other choice. So that doesn't matter.
             * Now, solve for u by taking the cube root, then multiplying with a cube root of unity, then calculating the corresponding v and getting all the result back out.
             */
            double complex u, v;
            u = cpow(-q/2 + csqrt(q * q / 4 + p * p * p / 27), 1.0/3.0);
            v = -p/(3 * u);

            double complex x = u + v + move;
            printf("x₁ = %g %+g i\n", creal(x), cimag(x));
            double complex z = d + x * (c + x * (b + x));
            printf("Test expression: %g %+g i\n", creal(z), cimag(z));
            u *= cexp(atan2(0, -1) * 2 * I/3.0);
            v = -p/(3 * u);
            x = u + v + move;
            printf("x₂ = %g %+g i\n", creal(x), cimag(x));
            z = d + x * (c + x * (b + x));
            printf("Test expression: %g %+g i\n", creal(z), cimag(z));
            u *= cexp(atan2(0, -1) * 2 * I/3.0);
            v = -p/(3 * u);
            x = u + v + move;
            printf("x₃ = %g %+g i\n", creal(x), cimag(x));
            z = d + x * (c + x * (b + x));
            printf("Test expression: %g %+g i\n", creal(z), cimag(z));
        }
    }
    return 0;
}
/* (x + 1)(x² + x + 1) = 0
 * x³ + x² + x + x² + x + 1 = 0
 * x³ + 2x² + 2x + 1 = 0
 *
 * It's normal.
 * let x = t - 2/3
 *
 * (t - 2/3)³ + 2(t - 2/3)² + 2(t - 2/3) + 1 = 0
 * t³ - 2t² + 4/3 t - 8/27 + 2(t² - 4/3 t + 4/9) + 2t - 4/3 + 1 = 0
 * t³ - 2t² + 4/3 t - 8/27 + 2t² - 8/3 t + 8/9 + 2t - 4/3 + 1 = 0
 * t³ + (4/3 - 8/3 + 2)t + (-8/27 + 8/9 - 4/3 + 1) = 0
 * t³ + 2/3t + 7/27 = 0
 * t³ = -2/3 t - 7/27
 * Let u + v = t
 * I. 3uv = -2/3 -> v = -2/(9u)
 * II. u³ + v³ = -7/27 -> u³ - 8/(729u³) = -2/3 -> u⁶ + 2/3u³ -8/729 = 0
 * -> u³ = -1/3 ± √(1/9 + 8/729) = -1/3 ± √(89/729) = -1/3 ± √89/(9√9) = -1/3 ± √89/27 -> u₁³ = (-9-√89)/27, u₂³ = (-9 + √89)/27
 * 
 * u₁ = ³√((-9 - √89)/27) = ³√(-9 - √89)/3 -> v₁ = -2/(3 ³√(-9 - √89)) = (-2 ³√(-9 - √89)²)/(3 (-9 - √89))
 *  = (-2 ³√(-9 - √89)²)/(-27 - 3 √89) = (-2 ³√(81 + 18 √89 + 89))(-27 - 3 √89)/(729 + 801)
 *  = (2 ³√(170 + 18 √89))(27 + 3 √89)/1530
 *  = (54 ³√(170 + 18 √89) + 6 √89 ³√(170 + 18 √89)/1530
 *  = (9 ³√(170 + 18 √89) + √89 ³√(170 + 18 √89)/255
 */
