package md2_cj

import stdx.encoding.hex.toHexString
import std.fs.File
import std.fs
import std.fs.OpenMode
import std.fs.Path
import std.io.BufferedInputStream
import md2_cj.utils

// The S-table values
let S: Array<UInt8> = [0x29, 0x2E, 0x43, 0xC9, 0xA2, 0xD8, 0x7C, 0x01, 0x3D, 0x36, 0x54, 0xA1, 0xEC, 0xF0, 0x06, 0x13,
    0x62, 0xA7, 0x05, 0xF3, 0xC0, 0xC7, 0x73, 0x8C, 0x98, 0x93, 0x2B, 0xD9, 0xBC, 0x4C, 0x82, 0xCA, 0x1E, 0x9B, 0x57,
    0x3C, 0xFD, 0xD4, 0xE0, 0x16, 0x67, 0x42, 0x6F, 0x18, 0x8A, 0x17, 0xE5, 0x12, 0xBE, 0x4E, 0xC4, 0xD6, 0xDA, 0x9E,
    0xDE, 0x49, 0xA0, 0xFB, 0xF5, 0x8E, 0xBB, 0x2F, 0xEE, 0x7A, 0xA9, 0x68, 0x79, 0x91, 0x15, 0xB2, 0x07, 0x3F, 0x94,
    0xC2, 0x10, 0x89, 0x0B, 0x22, 0x5F, 0x21, 0x80, 0x7F, 0x5D, 0x9A, 0x5A, 0x90, 0x32, 0x27, 0x35, 0x3E, 0xCC, 0xE7,
    0xBF, 0xF7, 0x97, 0x03, 0xFF, 0x19, 0x30, 0xB3, 0x48, 0xA5, 0xB5, 0xD1, 0xD7, 0x5E, 0x92, 0x2A, 0xAC, 0x56, 0xAA,
    0xC6, 0x4F, 0xB8, 0x38, 0xD2, 0x96, 0xA4, 0x7D, 0xB6, 0x76, 0xFC, 0x6B, 0xE2, 0x9C, 0x74, 0x04, 0xF1, 0x45, 0x9D,
    0x70, 0x59, 0x64, 0x71, 0x87, 0x20, 0x86, 0x5B, 0xCF, 0x65, 0xE6, 0x2D, 0xA8, 0x02, 0x1B, 0x60, 0x25, 0xAD, 0xAE,
    0xB0, 0xB9, 0xF6, 0x1C, 0x46, 0x61, 0x69, 0x34, 0x40, 0x7E, 0x0F, 0x55, 0x47, 0xA3, 0x23, 0xDD, 0x51, 0xAF, 0x3A,
    0xC3, 0x5C, 0xF9, 0xCE, 0xBA, 0xC5, 0xEA, 0x26, 0x2C, 0x53, 0x0D, 0x6E, 0x85, 0x28, 0x84, 0x09, 0xD3, 0xDF, 0xCD,
    0xF4, 0x41, 0x81, 0x4D, 0x52, 0x6A, 0xDC, 0x37, 0xC8, 0x6C, 0xC1, 0xAB, 0xFA, 0x24, 0xE1, 0x7B, 0x08, 0x0C, 0xBD,
    0xB1, 0x4A, 0x78, 0x88, 0x95, 0x8B, 0xE3, 0x63, 0xE8, 0x6D, 0xE9, 0xCB, 0xD5, 0xFE, 0x3B, 0x00, 0x1D, 0x39, 0xF2,
    0xEF, 0xB7, 0x0E, 0x66, 0x58, 0xD0, 0xE4, 0xA6, 0x77, 0x72, 0xF8, 0xEB, 0x75, 0x4B, 0x0A, 0x31, 0x44, 0x50, 0xB4,
    0x8F, 0xED, 0x1F, 0x1A, 0xDB, 0x99, 0x8D, 0x33, 0x9F, 0x11, 0x83, 0x14]
let PADDING: Array<Array<Byte>> = [
    [],
    [0x01],
    [0x02, 0x02],
    [0x03, 0x03, 0x03],
    [0x04, 0x04, 0x04, 0x04],
    [0x05, 0x05, 0x05, 0x05, 0x05],
    [0x06, 0x06, 0x06, 0x06, 0x06, 0x06],
    [0x07, 0x07, 0x07, 0x07, 0x07, 0x07, 0x07],
    [0x08, 0x08, 0x08, 0x08, 0x08, 0x08, 0x08, 0x08],
    [0x09, 0x09, 0x09, 0x09, 0x09, 0x09, 0x09, 0x09, 0x09],
    [0x0A, 0x0A, 0x0A, 0x0A, 0x0A, 0x0A, 0x0A, 0x0A, 0x0A, 0x0A],
    [0x0B, 0x0B, 0x0B, 0x0B, 0x0B, 0x0B, 0x0B, 0x0B, 0x0B, 0x0B, 0x0B],
    [0x0C, 0x0C, 0x0C, 0x0C, 0x0C, 0x0C, 0x0C, 0x0C, 0x0C, 0x0C, 0x0C, 0x0C],
    [0x0D, 0x0D, 0x0D, 0x0D, 0x0D, 0x0D, 0x0D, 0x0D, 0x0D, 0x0D, 0x0D, 0x0D, 0x0D],
    [0x0E, 0x0E, 0x0E, 0x0E, 0x0E, 0x0E, 0x0E, 0x0E, 0x0E, 0x0E, 0x0E, 0x0E, 0x0E, 0x0E],
    [0x0F, 0x0F, 0x0F, 0x0F, 0x0F, 0x0F, 0x0F, 0x0F, 0x0F, 0x0F, 0x0F, 0x0F, 0x0F, 0x0F, 0x0F],
    [0x10, 0x10, 0x10, 0x10, 0x10, 0x10, 0x10, 0x10, 0x10, 0x10, 0x10, 0x10, 0x10, 0x10, 0x10, 0x10]
]
// 16字节为一组
const HASH_BLOCK_SIZE: Int64 = 16
// 最后X[]缓冲区的前16个字节就是最终结果
const HASH_DIGEST_SIZE: Int64 = 16
// md2 函数核心步骤执行18轮
const HASH_ROUND_NUM: UInt16 = 18u16
// 在数据末尾追加16字节的校验码
const MD2_CHECKSUM_SIZE: Int64 = 16

// last block
// @Derive[ToString]
internal struct LastBlock {
    var used: Int64 // used bytes
    var buf: Array<UInt8> // block data buffer

    public init() {
        used = 0
        buf = Array<Byte>(HASH_BLOCK_SIZE, repeat: 0u8)
    }
}

// @Derive[ToString]
// internal class Md2Context <: ToString {
internal class Md2Context {
    // 真实数据的字节长度
    // 感觉没用?难道计算的上限要被Int64限制?
    // var total: Int64

    // 48 bytes buffer
    var x: VArray<Byte, $48>

    // last block
    var last: LastBlock

    // checksum
    var checksum: Array<Byte>

    init() {
        // total = 0
        x = VArray<Byte, $48>(repeat: 0u8)
        last = LastBlock()
        checksum = Array<Byte>(MD2_CHECKSUM_SIZE, repeat: 0u8)
    }

    // public func toString(): String {
    //     let ret = StringBuilder()
    //     // 输出total
    //     //ret.append("total:UInt64 = ${total}\n")
    //     // 输出x
    //     ret.append("x:VArray<Byte, $48> =")

    //     for (i in 0..x.size) {
    //         // 每输出16字节换行
    //         if (i % HASH_BLOCK_SIZE == 0) {
    //             ret.append("\n")
    //         }
    //         ret.append("\t${i}-${x[i]} ")
    //     }

    //     ret.append("\n")
    //     // 输出last
    //     ret.append("last: LastBlock =\n")
    //     ret.append("\tused: UInt8 = ${last.used}\n")
    //     ret.append("\tbuf: VArray<UInt8, $16> = ")
    //     for (i in 0..last.buf.size) {
    //         ret.append("${i}-${last.buf[i]} ")
    //     }
    //     ret.append("\n")

    //     // 输出checksum
    //     ret.append("checksum:VArray<Byte, $16> = ")

    //     for (i in 0..checksum.size) {
    //         // 每输出16字节换行
    //         if (i % HASH_BLOCK_SIZE == 0) {
    //             ret.append("\n")
    //         }
    //         ret.append("${i}-${checksum[i]} ")
    //     }

    //     ret.append("\n")

    //     return ret.toString()
    // }
}

/**
 * 给定字符串,返回MD2哈希值
 * @param data 待计算的字符串
 */
public func md2(data: String): String {
    // 初始化
    let ctx: Md2Context = Md2Context()
    //核心 
    md2Update(ctx, data.iterator())

    // 取结果并返回
    return getResult(ctx)
}

/**
 * 给定文件路径,返回MD2哈希值
 * @param filePath 待计算的文件路径
 * @return Option<String> 文件的MD2哈希值,如果不是有效的普通文件或符号链接,返回None
 */
public func md2(filePath: Path): ?String {
    // 检查一下就扔给 md2(file: File)
    if (!utils.fileExists(filePath)) {
        // 判断目标地址是否存在。
        return None<String>
    }
    // 初始化
    let ctx: Md2Context = Md2Context()
    //核心
    try (file = File(filePath, OpenMode.Read)) {
        md2File(ctx, file)
    } catch (_) {
        println("还能到这里来?")
        println("read file ${filePath} failed")
    }

    // 会转成 Option<String>,所以这里不用处理
    // 参见文档:虽然 T 和 Option<T> 是不同的类型,
    // 但是当明确知道某个位置需要的是 Option<T> 类型的值时,
    // 可以直接传一个 T 类型的值,编译器会用 Option<T> 类型的 Some 构造器
    // 将 T 类型的值封装成 Option<T> 类型的值(注意:这里并不是类型转换)
    return getResult(ctx)
}

/**
 * 给定文件,返回MD2哈希值
 * @param file 待计算的文件
 * @return Option<String> 文件的MD2哈希值,如果不是有效的普通文件或符号链接,返回None
 * @deprecated 传递File会导致函数内难以处理是否要close的问题,建议传递Path,使用md2(filePath: Path)
 */
// public func md2(file: File): ?String {
//     // 检查一下就扔给 md2File
//     if (file.info.isDirectory()) {
//         return None<String>
//     }
//     // 已经是合法File排除了目录的可能就是文件了
//     return md2File(file)
// }

/**
 * @param data 此时file就是合法的普通文件或符号链接,不再判断合法性
 */
private func md2File(ctx: Md2Context, file: File): Md2Context {
    let bufferedFileInputStream = BufferedInputStream(file)

    //一次传16字节的流,直到流结束
    while (let num <- bufferedFileInputStream.read(ctx.last.buf)) {
        if (num == 0) {
            break
        }

        ctx.last.used = num
        if (num == HASH_BLOCK_SIZE) {
            md2Block(ctx, ctx.last.buf)
            md2UpdateChecksum(ctx, ctx.last.buf)
        }
    }

    // 原始流读完了,填充数据
    md2PaddingUpdate(ctx)

    // 处理校验码数据库,这是最后一个分组
    md2Block(ctx, ctx.checksum)

    return ctx
}

/**
 * 给一个完整数据流计算md2值
 * TODO 还应该能支持分块计算,最后再合并
 * @param data 完整数据流
 */
private func md2Update(ctx: Md2Context, data: Iterator<Byte>): Md2Context {
    // let blockData: Array<Byte> = Array<Byte>(HASH_BLOCK_SIZE, repeat: 0u8)

    //一次传16字节的流,直到流结束
    var index = 0 //本次要传的16字节的流已经读了多少字节
    for (value in data) {
        // blockData[index] = value
        ctx.last.buf[index] = value

        // ctx.total++
        index++
        ctx.last.used = index
        if (index == HASH_BLOCK_SIZE) {
            // md2Block(ctx, blockData)
            // md2UpdateChecksum(ctx, blockData)
            md2Block(ctx, ctx.last.buf)
            md2UpdateChecksum(ctx, ctx.last.buf)

            index = 0
        }
    }
    // 原始流读完了,填充数据
    md2PaddingUpdate(ctx)

    // 处理校验码数据库,这是最后一个分组
    md2Block(ctx, ctx.checksum)
    return ctx
}

/**
 * 核心,负责填充数据的md2函数计算过程
 */
private func md2PaddingUpdate(ctx: Md2Context): Unit {
    // 原始流读完了,填充数据
    // 长度不够的则将数据填充到16字节的倍数,长度足够则额外追加16字节的数据(MD2实现就是这样滴~)
    // 设数据长度为n字节,则需要填充m=16-(n mod 16)个字节的数据。m个字节的内容均为十进制的m
    // let paddingSize = HASH_BLOCK_SIZE - (ctx.total % HASH_BLOCK_SIZE)
    // 按原始文档应该是total % HASH_BLOCK_SIZE,但是total会越来越大,所以这里改用last.used
    let paddingSize = HASH_BLOCK_SIZE - ctx.last.used % HASH_BLOCK_SIZE
    if (paddingSize != HASH_BLOCK_SIZE) {
        //把blockData补完
        ctx.last.buf[ctx.last.used..] = PADDING[paddingSize]
        md2Block(ctx, ctx.last.buf)
        md2UpdateChecksum(ctx, ctx.last.buf)
    } else {
        // 长度足够则额外追加16字节的数据
        // ctx.last.buf[..] = PADDING[paddingSize] // 取消这么写,改成直接用PADDING数据,减少一次复制
        // 注意,此时 ctx.last.buf 没有更新,后面也不用了,所以就这样了
        md2Block(ctx, PADDING[paddingSize])
        md2UpdateChecksum(ctx, PADDING[paddingSize])
    }
}

/**
 * 核心,只管计算48字节的md2函数计算过程
 * 不负责数据填充,不负责末尾追加
 * 文档里的 3.4 Step 4. Process Message in 16-Byte Blocks
 * @param ctx
 * @param data 16字节长的数据分组
 */
private func md2Block(ctx: Md2Context, data: Iterable<Byte>): Md2Context {

    /****************************************************************/
    //开辟一个48字节的缓冲区,依次存放16字节长的上一次MD函数输出、16字节长的数据分组和二者的异或。

    // 1. 第一个16字节在初始化、上一轮 md2Block 函数计算已做
    // 2. 依次存放16字节长的数据分组和异或
    var index = 0
    var index2 = HASH_BLOCK_SIZE + index
    var index3 = HASH_BLOCK_SIZE + index2
    for (byteItem in data) {
        ctx.x[index2] = byteItem
        ctx.x[index3] = ctx.x[index] ^ ctx.x[index2]
        index++
        index2++
        index3++
    }

    /****************************************************************/
    /* 核心步骤
    初始化字节t为\0,将t经过S盒变换(S-Box Substitution,其实就是查表替换)后和
    缓冲区的一个字节进行位异或并写回缓冲区,结果作为下次S盒变换的输入。依次类推,
    直至处理完缓冲区的所有数据。再让t加上变换轮数-1,对256取余。这个过程进行18轮。
     */
    // 初始化字节t为\0
    var t: Byte = 0
    // 18轮
    for (round in 0u16..HASH_ROUND_NUM) {
        for (index in 0..ctx.x.size) {
            // S盒变换,异或
            t = S[Int64(t)] ^ ctx.x[index]
            // 写回缓冲区
            ctx.x[index] = t
        }
        // t = (t + 已进行的轮数-1) mod 256
        // UInt16(t)是为了避免UInt8溢出
        t = UInt8((UInt16(t) + round) % 256u16)
    }
    /****************************************************************/

    return ctx
}

/**
 * 校验码计算
 * 其实是每一次16字节搜集好时都要计算checksum和 md2Block 函数
 * @param ctx
 * @param data 16字节长的数据分组
 */
private func md2UpdateChecksum(ctx: Md2Context, data: Iterable<Byte>): Md2Context {
    // 文档里(3.2 Step 2. Append Checksum)说"Set L to 0."" 所以这里也大写
    var L: Byte = ctx.checksum[ctx.checksum.size - 1]
    // var L: Byte = ctx.checksum.last //last是Option<Byte>,麻烦

    var index = 0
    for (byteItem in data) {
        L = ctx.checksum[index] ^ S[Int64(byteItem ^ L)]
        ctx.checksum[index] = L
        index++
    }

    return ctx
}

/**
 * 取缓冲区前16字节数据,转成16进制字符串作为结果返回
 * @param ctx 上下文
 * @return String 最终md2计算的结果
 */
private func getResult(ctx: Md2Context): String {
    // 返回结果,取 ctx.x的前 HASH_DIGEST_SIZE 字节,转成Array<UInt8>
    let retArray: Array<Byte> = Array<Byte>(HASH_DIGEST_SIZE, {i => ctx.x[i]})
    // 再转成String
    return toHexString(retArray)
}