#include "sha1.h"
#include <string.h>

#define ROTLEFT(a,b) (((a) << (b)) | ((a) >> (32-(b))))

static void SHA1_Transform(SHA1_CTX* ctx, const uint8_t data[64]) {
    uint32_t a, b, c, d, e, i, j, t, m[80];

    for (i = 0, j = 0; i < 16; ++i, j += 4)
        m[i] = (data[j] << 24) | (data[j+1] << 16) | (data[j+2] << 8) | (data[j+3]);
    for ( ; i < 80; ++i)
        m[i] = ROTLEFT(m[i-3] ^ m[i-8] ^ m[i-14] ^ m[i-16], 1);

    a = ctx->state[0];
    b = ctx->state[1];
    c = ctx->state[2];
    d = ctx->state[3];
    e = ctx->state[4];

    for (i = 0; i < 80; ++i) {
        if (i < 20)
            t = ROTLEFT(a,5) + ((b & c) | (~b & d)) + e + m[i] + 0x5A827999;
        else if (i < 40)
            t = ROTLEFT(a,5) + (b ^ c ^ d) + e + m[i] + 0x6ED9EBA1;
        else if (i < 60)
            t = ROTLEFT(a,5) + ((b & c) | (b & d) | (c & d)) + e + m[i] + 0x8F1BBCDC;
        else
            t = ROTLEFT(a,5) + (b ^ c ^ d) + e + m[i] + 0xCA62C1D6;

        e = d;
        d = c;
        c = ROTLEFT(b,30);
        b = a;
        a = t;
    }

    ctx->state[0] += a;
    ctx->state[1] += b;
    ctx->state[2] += c;
    ctx->state[3] += d;
    ctx->state[4] += e;
}

void SHA1_Init(SHA1_CTX* ctx) {
    ctx->state[0] = 0x67452301;
    ctx->state[1] = 0xEFCDAB89;
    ctx->state[2] = 0x98BADCFE;
    ctx->state[3] = 0x10325476;
    ctx->state[4] = 0xC3D2E1F0;
    ctx->count[0] = ctx->count[1] = 0;
}

void SHA1_Update(SHA1_CTX* ctx, const uint8_t* data, uint32_t len) {
    uint32_t i, j;

    j = (ctx->count[0] >> 3) & 63;
    if ((ctx->count[0] += len << 3) < (len << 3))
        ctx->count[1]++;
    ctx->count[1] += (len >> 29);

    if ((j + len) > 63) {
        memcpy(&ctx->buffer[j], data, (i = 64 - j));
        SHA1_Transform(ctx, ctx->buffer);
        for ( ; i + 63 < len; i += 64)
            SHA1_Transform(ctx, &data[i]);
        j = 0;
    } else {
        i = 0;
    }
    memcpy(&ctx->buffer[j], &data[i], len - i);
}

void SHA1_Final(uint8_t digest[20], SHA1_CTX* ctx) {
    uint8_t finalcount[8];
    uint8_t c;
    uint32_t i;

    for (i = 0; i < 8; i++)
        finalcount[i] = (uint8_t)((ctx->count[(i >= 4 ? 0 : 1)]
                                  >> ((3 - (i & 3)) * 8)) & 255);

    c = 0x80;
    SHA1_Update(ctx, &c, 1);
    while ((ctx->count[0] & 504) != 448) {
        c = 0x00;
        SHA1_Update(ctx, &c, 1);
    }

    SHA1_Update(ctx, finalcount, 8);

    for (i = 0; i < 20; i++)
        digest[i] = (uint8_t)((ctx->state[i>>2] >> ((3 - (i & 3)) * 8)) & 255);

    memset(ctx, 0, sizeof(*ctx));
}
