#include "rc4.h"

void RC4Init(TRC4Context* rc4, const uint8_t* key, uint32_t keyLen) {
    uint8_t R, T, K;
    uint32_t U;
    uint32_t L = keyLen;
    int S; // Changed from uint8_t to int to prevent infinite loop

    rc4->I = 0;
    rc4->J = 0;

    // First loop: Initialize state
    for (S = 0; S < 256; S++)
        rc4->D[S] = (uint8_t)S;

    R = 0;
    U = 0;

    // Second loop: Key scheduling
    for (S = 0; S < 256; S++) {
        if (L > 0) K = key[U]; // Added check for safety
        else       K = 0;
        
        U++;
        if (U >= L) U = 0;

        R = (uint8_t)(R + rc4->D[S] + K);
        T = rc4->D[S];
        rc4->D[S] = rc4->D[R];
        rc4->D[R] = T;
    }
}

void rc4Decrypt(const void* InData, void* OutData, uint32_t Size, TRC4Context* ctx) {
    uint32_t i = 0, j = 0, t, k;
    const uint8_t* in  = (const uint8_t*)InData;
    uint8_t* out       = (uint8_t*)OutData;

    for (k = 0; k < Size; k++) {
        i = (i + 1) & 0xFF;
        t = ctx->D[i];
        j = (j + t) & 0xFF;
        ctx->D[i] = ctx->D[j];
        ctx->D[j] = (uint8_t)t;
        t = (t + ctx->D[i]) & 0xFF;
        out[k] = in[k] ^ ctx->D[t];
    }
}
