// A wrapper for LZMA SDK library (C interface)
// (C) 2010-05-08 Adam Sawicki - sawickiap@poczta.onet.pl - http://regedit.gamedev.pl/

#include "LzmUtils.hpp"

extern "C" {
#include "LzmaLib.h"
#include "LzmaEnc.h"
#include "LzmaDec.h"
}

static void * AllocForLzma(void *p, size_t size) { return malloc(size); }
static void FreeForLzma(void *p, void *address) { free(address); }
static ISzAlloc g_AllocForLzma = { &AllocForLzma, &FreeForLzma };

class LzmaError : public Error
{
public:
	static const tchar * ErrorCodeToStr(int errorCode);
	static const tchar * StatusToStr(ELzmaStatus status);
	LzmaError(int errorCode, const tstring &msg = _T(""), const tstring &file = _T(""), int line = 0);
	LzmaError(ELzmaStatus status, const tstring &msg = _T(""), const tstring &file = _T(""), int line = 0);
};

const tchar * LzmaError::ErrorCodeToStr(int errorCode)
{
	switch (errorCode)
	{
	case SZ_ERROR_DATA:        return _T("Data");
	case SZ_ERROR_MEM:         return _T("Mem");
	case SZ_ERROR_CRC:         return _T("CRC");
	case SZ_ERROR_UNSUPPORTED: return _T("Unsupported");
	case SZ_ERROR_PARAM:       return _T("Param");
	case SZ_ERROR_INPUT_EOF:   return _T("Input EOF");
	case SZ_ERROR_OUTPUT_EOF:  return _T("Output EOF");
	case SZ_ERROR_READ:        return _T("Read");
	case SZ_ERROR_WRITE:       return _T("Write");
	case SZ_ERROR_PROGRESS:    return _T("Progress");
	case SZ_ERROR_FAIL:        return _T("Fail");
	case SZ_ERROR_THREAD:      return _T("Thread");
	case SZ_ERROR_ARCHIVE:     return _T("Archive");
	case SZ_ERROR_NO_ARCHIVE:  return _T("No archive");
	default: return _T("Unknown");
	}
}

const tchar * LzmaError::StatusToStr(ELzmaStatus status)
{
	switch (status)
	{
	case LZMA_STATUS_NOT_SPECIFIED:      return _T("Not specified");
	case LZMA_STATUS_FINISHED_WITH_MARK: return _T("Finished with mark");
	case LZMA_STATUS_NOT_FINISHED:       return _T("Not finished");
	case LZMA_STATUS_NEEDS_MORE_INPUT:   return _T("Needs more input");
	case LZMA_STATUS_MAYBE_FINISHED_WITHOUT_MARK: return _T("Maybe finished without mark");
	default: return _T("Unknown");
	}
}

LzmaError::LzmaError(int errorCode, const tstring &msg, const tstring &file, int line)
{
	Push(Format(_T("LZMA SRes #: #")) % errorCode % ErrorCodeToStr(errorCode));
	Push(msg, file, line);
}

LzmaError::LzmaError(ELzmaStatus status, const tstring &msg, const tstring &file, int line)
{
	Push(Format(_T("LZMA Status #: #")) % (unsigned)status % StatusToStr(status));
	Push(msg, file, line);
}

void LzmaCompressData(
	void *dstPtr, uint &inoutDstBytes,
	const void *srcPtr, uint srcBytes,
	LzmaProps &outProps, bool useEndMark, int level, uint dictSize)
{
	assert(0 <= level && level <= 9);

	uint propsSize = LZMA_PROPS_SIZE;
	CLzmaEncProps props;
	LzmaEncProps_Init(&props);
	props.writeEndMark = useEndMark ? 1 : 0;
	props.level = level;
	if (dictSize)
		props.dictSize = dictSize;

	int res = LzmaEncode(
		(uint8*)dstPtr, &inoutDstBytes,
		(const uint8*)srcPtr, srcBytes,
		&props, outProps.Data, &propsSize, props.writeEndMark,
		NULL, &g_AllocForLzma, &g_AllocForLzma);
	if (res != SZ_OK)
		throw LzmaError(res, _T("LzmaCompressData: LzmaEncode failed."), __TFILE__, __LINE__);
}

bool LzmaUncompressData(void *dstPtr, uint &inoutDstBytes,
	const void *srcPtr, uint &inoutSrcBytes,
	const LzmaProps &props)
{
	ELzmaStatus status;
	SRes res = LzmaDecode(
		(uint8*)dstPtr, &inoutDstBytes,
		(const uint8*)srcPtr, &inoutSrcBytes,
		props.Data, LZMA_PROPS_SIZE,
		LZMA_FINISH_END,
		&status,
		&g_AllocForLzma);
	if (res != SZ_OK)
		throw LzmaError(res, _T("LzmaUncompressData: LzmaDecode failed."), __TFILE__, __LINE__);

	if (status == LZMA_STATUS_FINISHED_WITH_MARK)
		return true;
	else if (status == LZMA_STATUS_MAYBE_FINISHED_WITHOUT_MARK)
		return false;
	else
		throw LzmaError(status, _T("LzmaUncompressData: LzmaDecode returned invalid status."), __TFILE__, __LINE__);
}


struct LzmaInStream
{
	ISeqInStream m_SeqInStream;
	Stream *m_Stream;
	uint64 m_BytesLeft;
};

static SRes LzmaInStream_Read(void *p, void *buf, size_t *size)
{
	assert(*size);
	LzmaInStream *ctx = (LzmaInStream*)p;
	
	if (ctx->m_BytesLeft == 0)
	{
		*size = 0; // end_of_stream.
		return SZ_OK;
	}
	
	try
	{
		if (*size > ctx->m_BytesLeft)
			*size = (uint)ctx->m_BytesLeft;
		*size = ctx->m_Stream->Read(buf, *size);
		ctx->m_BytesLeft -= *size;
		return SZ_OK;
	}
	catch (const Error&)
	{
		return SZ_ERROR_FAIL;
	}
}

struct LzmaOutStream
{
	ISeqOutStream m_SeqOutStream;
	Stream *m_Stream;
};

static size_t LzmaOutStream_Write(void *p, const void *buf, size_t size)
{
	assert(size);
	LzmaOutStream *ctx = (LzmaOutStream*)p;
	try
	{
		ctx->m_Stream->Write(buf, size);
		return size;
	}
	catch (const Error&)
	{
		return 0; // Value other than size means error.
	}
}

void LzmaCompressStream(Stream &dst, Stream &src, bool useEndMark,
	uint64 maxBytes, int level, uint dictSize)
{
	assert(0 <= level && level <= 9);

	CLzmaEncHandle enc = LzmaEnc_Create(&g_AllocForLzma);
	assert(enc);

	CLzmaEncProps props;
	LzmaEncProps_Init(&props);
	props.writeEndMark = useEndMark ? 1 : 0;
	props.level = level;
	if (dictSize)
		props.dictSize = dictSize;
	
	SRes res = LzmaEnc_SetProps(enc, &props);
	if (res != SZ_OK)
	{
		LzmaEnc_Destroy(enc, &g_AllocForLzma, &g_AllocForLzma);
		throw LzmaError(res, _T("LzmaCompressStream: LzmaEnc_SetProps failed."), __TFILE__, __LINE__);
	}

	uint propsSize = LZMA_PROPS_SIZE;
	LzmaProps encodedProps;
	res = LzmaEnc_WriteProperties(enc, encodedProps.Data, &propsSize);
	if (res != SZ_OK)
	{
		LzmaEnc_Destroy(enc, &g_AllocForLzma, &g_AllocForLzma);
		throw LzmaError(res, _T("LzmaCompressStream: LzmaEnc_WriteProperties failed."), __TFILE__, __LINE__);
	}
	assert(propsSize == LZMA_PROPS_SIZE);
	dst.Write(encodedProps.Data, 5);

	LzmaInStream lzmaInStream = { &LzmaInStream_Read, &src, maxBytes };
	LzmaOutStream lzmaOutStream = { &LzmaOutStream_Write, &dst };

	res = LzmaEnc_Encode(enc,
		(ISeqOutStream*)&lzmaOutStream, (ISeqInStream*)&lzmaInStream,
		0, &g_AllocForLzma, &g_AllocForLzma);

	LzmaEnc_Destroy(enc, &g_AllocForLzma, &g_AllocForLzma);

	if (res != SZ_OK)
		throw LzmaError(res, _T("LzmaCompressStream: LzmaEnc_Encode failed."), __TFILE__, __LINE__);
}

bool LzmaUncompressStream(Stream &dst, Stream &src, uint64 dstLen, uint bufSize)
{
	assert(bufSize);

	CLzmaDec dec;
	LzmaDec_Construct(&dec);

	uint8 encodedProps[5];
	src.MustRead(encodedProps, sizeof(encodedProps));

	SRes res = LzmaDec_Allocate(&dec, encodedProps, LZMA_PROPS_SIZE, &g_AllocForLzma);
	if (res != SZ_OK)
		throw LzmaError(res, _T("LzmaUncompressStream: LzmaDec_Allocate failed."), __TFILE__, __LINE__);

	LzmaDec_Init(&dec);

	std::vector<uint8> srcBuf(bufSize), dstBuf(bufSize);
	ELzmaFinishMode finishMode = LZMA_FINISH_ANY;
	uint64 dstSum = 0;
	bool finishedWithEndMark = false;
	for (;;)
	{
		uint srcSize2 = src.Read(&srcBuf[0], bufSize);
		// End of src - decoding should have finished and break the loop earlier.
		if (srcSize2 == 0)
		{
			LzmaDec_Free(&dec, &g_AllocForLzma);
			throw Error(_T("LzmaUncompressStream: Unexpected end of input data (1)."), __TFILE__, __LINE__);
		}
		
		uint srcOff = 0;
		bool finished = false;
		while (srcOff < srcSize2)
		{
			uint dstSize2 = bufSize;
			if (dstSum + dstSize2 >= dstLen)
			{
				dstSize2 = (uint)(dstLen - dstSum);
				finishMode = LZMA_FINISH_END;
			}

			uint srcSize3 = srcSize2 - srcOff;
			uint dstSize3 = dstSize2;
			ELzmaStatus status;
			res = LzmaDec_DecodeToBuf(&dec, &dstBuf[0], &dstSize3, &srcBuf[srcOff], &srcSize3, finishMode, &status);
			if (res != SZ_OK)
			{
				LzmaDec_Free(&dec, &g_AllocForLzma);
				throw LzmaError(res, _T("LzmaUncompressStream: LzmaDec_DecodeToBuf failed."), __TFILE__, __LINE__);
			}
			if (srcSize3 != srcSize2 - srcOff && dstSize3 != dstSize2)
			{
				LzmaDec_Free(&dec, &g_AllocForLzma);
				throw Error(_T("LzmaUncompressStream: Internal error (2)."), __TFILE__, __LINE__);
			}
			if (dstSize3)
				dst.Write(&dstBuf[0], dstSize3);

			srcOff += srcSize3;
			dstSum += dstSize3;

			if (status == LZMA_STATUS_FINISHED_WITH_MARK)
			{
				finished = true;
				finishedWithEndMark = true;
				break;
			}
			if (dstSum > dstLen)
			{
				LzmaDec_Free(&dec, &g_AllocForLzma);
				throw Error(_T("LzmaUncompressStream: Internal error (4)."), __TFILE__, __LINE__);
			}
			// Uncompressed length met.
			if (dstSum == dstLen)
			{
				if (status != LZMA_STATUS_MAYBE_FINISHED_WITHOUT_MARK)
				{
					LzmaDec_Free(&dec, &g_AllocForLzma);
					throw LzmaError(status, _T("LzmaUncompressStream: Decompression finished but LzmaDec_DecodeToBuf returned invalid status."), __TFILE__, __LINE__);
				}
				// Finish gracefully.
				finished = true;
				// finishedWithEndMark stays false.
				break;
			}
		}

		if (finished)
			break;

		// End of src - decoding should have finished and break the loop earlier.
		if (srcSize2 != bufSize)
		{
			LzmaDec_Free(&dec, &g_AllocForLzma);
			throw Error(_T("LzmaUncompressStream: Unexpected end of input data (2)."), __TFILE__, __LINE__);
		}
	}

	LzmaDec_Free(&dec, &g_AllocForLzma);
	return finishedWithEndMark;
}


LzmaDecompressionStream::LzmaDecompressionStream(Stream *src, uint64 dstLen, uint bufSize)
: OverlayStream(src)
, m_Dec(new CLzmaDec())
, m_SrcBuf(bufSize)
, m_SrcBufSize(0)
, m_SrcBufOff(0)
, m_DstLen(dstLen)
, m_DstSum(0)
, m_End(false)
, m_FinishedWithEndMark(false)
{
	assert(bufSize);

	CLzmaDec *dec = (CLzmaDec*)m_Dec;
	LzmaDec_Construct(dec);

	uint8 encodedProps[5];
	src->MustRead(encodedProps, sizeof(encodedProps));

	SRes res = LzmaDec_Allocate(dec, encodedProps, LZMA_PROPS_SIZE, &g_AllocForLzma);
	if (res != SZ_OK)
		throw LzmaError(res, _T("LzmaDecompressionStream: LzmaDec_Allocate failed."), __TFILE__, __LINE__);

	LzmaDec_Init(dec);
}

LzmaDecompressionStream::~LzmaDecompressionStream()
{
	CLzmaDec *dec = (CLzmaDec*)m_Dec;
	delete dec;
}

size_t LzmaDecompressionStream::Read(void *Out, size_t MaxLength)
{
	if (m_End)
		return 0;

	CLzmaDec *dec = (CLzmaDec*)m_Dec;

	uint8 *outBytes = (uint8*)Out;
	uint outOff = 0;
	ELzmaFinishMode finishMode = LZMA_FINISH_ANY;
	while (outOff < MaxLength && !m_End)
	{
		// Input buffer is empty: Read some more compressed data.
		if (m_SrcBufOff == m_SrcBufSize)
		{
			m_SrcBufOff = 0;
			m_SrcBufSize = GetStream()->Read(&m_SrcBuf[0], m_SrcBuf.size());
			// End of src - decoding should have finished and break the loop earlier.
			if (m_SrcBufSize == 0)
				throw Error(_T("LzmaDecompressionStream: Unexpected end of input data (1)."), __TFILE__, __LINE__);
		}

		uint dstSize2 = MaxLength - outOff;
		if (m_DstSum + dstSize2 >= m_DstLen)
		{
			dstSize2 = (uint)(m_DstLen - m_DstSum);
			finishMode = LZMA_FINISH_END;
		}

		uint srcSize3 = m_SrcBufSize - m_SrcBufOff;
		uint dstSize3 = dstSize2;
		ELzmaStatus status;
		SRes res = LzmaDec_DecodeToBuf(dec, &outBytes[outOff], &dstSize3, &m_SrcBuf[m_SrcBufOff], &srcSize3, finishMode, &status);
		if (res != SZ_OK)
			throw LzmaError(res, _T("LzmaDecompressionStream: LzmaDec_DecodeToBuf failed."), __TFILE__, __LINE__);
		if (srcSize3 != m_SrcBufSize - m_SrcBufOff && dstSize3 != dstSize2)
			throw Error(_T("LzmaDecompressionStream: Internal error (2)."), __TFILE__, __LINE__);

		m_SrcBufOff += srcSize3;
		outOff += dstSize3;
		m_DstSum += dstSize3;

		// End mark met - finish gracefully.
		if (status == LZMA_STATUS_FINISHED_WITH_MARK)
		{
			m_End = true;
			m_FinishedWithEndMark = true;
			break;
		}

		if (m_DstSum > m_DstLen)
			throw Error(_T("LzmaDecompressionStream: Internal error (4)."), __TFILE__, __LINE__);
		// Uncompressed length met.
		if (m_DstSum == m_DstLen)
		{
			if (status != LZMA_STATUS_MAYBE_FINISHED_WITHOUT_MARK)
				throw LzmaError(status, _T("LzmaDecompressionStream: Decompression finished but LzmaDec_DecodeToBuf returned invalid status."), __TFILE__, __LINE__);
			// Finish gracefully.
			m_End = true;
			// m_FinishedWithEndMark stays false.
			break;
		}

		// End of src - decoding should have finished and break the loop earlier.
		if (m_SrcBufOff == m_SrcBufSize)
		{
			if (m_SrcBufSize != m_SrcBuf.size())
				throw Error(_T("LzmaDecompressionStream: Unexpected end of input data (2)."), __TFILE__, __LINE__);
		}
	}

	return outOff;
}
