================ FILE: tools/pcm_decoder.py ================
#!/usr/bin/env python3
"""Korg Pa-series PCM decoder for .KSF sample data.
Verified against libgig Korg.cpp (drbye78/libgig mirror) + ConvertWithMoss.

WHAT THIS DOES
- Decodes the raw PCM payload of the SMD1 chunk into numpy float arrays.
- Handles bit depth (8/16), big-endian byte order, channels (mono/stereo interleave).
- Detects compression via the Attributes bit field and REFUSES to guess:
  compressed samples are returned as raw bytes + a flag (libgig itself does not
  decompress — the codec is proprietary and undocumented).
- Applies loop-point metadata (LoopStart/LoopEnd are in SAMPLE FRAMES, not bytes).
- Exports to WAV for audible verification (Golden Corpus principle).
- Round-trip re-encoder so the byte-for-byte corpus check covers audio too.

VERIFIED FACTS (from libgig Korg.cpp)
- RIFF file opened with RIFF::endian_big  -> big-endian.
- SMD1_CHUNK_HEADER_SZ = 12. PCM starts at offset 12 inside SMD1 payload.
- FrameSize() = BitDepth/8 * Channels   (bytes per sample point / frame).
- SamplePoints = number of FRAMES (not bytes).
- PCM byte length = SamplePoints * FrameSize.
- Attributes bit field:
    IsCompressed()  = Attributes & 0x10   (bit 4)
    CompressionID() = Attributes & 0x04   (bit 2, value 0 or 4)
    Use2ndStart()   = !(Attributes & 0x20) (bit 5, INVERTED logic)
- Read() in libgig reads raw bytes; NO decompression is implemented there.

NEVER GUESS
- 8-bit signedness: assume signed int8, but VERIFY against one real file
  (compare a known sample's polarity). Flag unverified.
- Compressed samples: do not interpret as PCM. Preserve bytes. Mark compressed.
- Stereo: Korg .KSF is mono in practice; if Channels==2, samples are interleaved
  L,R,L,R. Verify against a real stereo file before trusting.
"""
import struct
import wave
from dataclasses import dataclass
from typing import Optional, Tuple
import numpy as np

SMD1_HEADER_SZ = 12


def attributes_flags(attributes: int) -> dict:
    """Decode the SMD1 Attributes bit field (verified offsets from libgig)."""
    return {
        "is_compressed": bool(attributes & 0x10),
        "compression_id": (attributes & 0x04) >> 2,   # 0 or 1
        "use_2nd_start": not bool(attributes & 0x20), # inverted
        "raw": attributes,
    }


def frame_size(bit_depth: int, channels: int) -> int:
    """Bytes per sample point (frame). Verified: BitDepth/8 * Channels."""
    return (bit_depth // 8) * channels


@dataclass
class DecodedPCM:
    samples: np.ndarray            # shape (SamplePoints, Channels), float32 in [-1,1]
    sample_rate: int
    channels: int
    bit_depth: int
    loop_start: int                # in frames
    loop_end: int                  # in frames
    is_compressed: bool
    compression_id: int
    use_2nd_start: bool
    raw_bytes: bytes = b""         # original PCM bytes (for round-trip / compressed)


def decode_pcm(ksf) -> DecodedPCM:
    """Decode a parsed KSF (from korg_decoder.parse_ksf) into float samples.

    ksf.pcm is the SMD1 payload MINUS the 12-byte header (korg_decoder already
    strips it). If you only have the raw SMD1 payload, pass payload[12:].
    """
    flags = attributes_flags(ksf.attributes)
    fs = frame_size(ksf.bit_depth, ksf.channels)
    expected_bytes = ksf.sample_points * fs
    pcm = ksf.pcm[:expected_bytes]

    if flags["is_compressed"]:
        # libgig does not decompress; neither do we. Do NOT guess.
        # Return zeros + raw bytes so the caller can preserve/flag.
        samples = np.zeros((ksf.sample_points, max(ksf.channels, 1)), dtype=np.float32)
        return DecodedPCM(
            samples=samples, sample_rate=ksf.sample_rate, channels=ksf.channels,
            bit_depth=ksf.bit_depth, loop_start=ksf.loop_start, loop_end=ksf.loop_end,
            is_compressed=True, compression_id=flags["compression_id"],
            use_2nd_start=flags["use_2nd_start"], raw_bytes=pcm,
        )

    if ksf.bit_depth == 16:
        # big-endian signed int16 -> float32 [-1, 1]
        raw = np.frombuffer(pcm, dtype=">i2").astype(np.float32) / 32768.0
    elif ksf.bit_depth == 8:
        # signed int8 (VERIFY against a real file). -> float32 [-1, 1]
        raw = np.frombuffer(pcm, dtype=np.int8).astype(np.float32) / 128.0
    else:
        raise ValueError(f"Unsupported bit depth {ksf.bit_depth}; verify against a real file")

    if ksf.channels > 1:
        samples = raw.reshape(-1, ksf.channels)
    else:
        samples = raw.reshape(-1, 1)

    return DecodedPCM(
        samples=samples, sample_rate=ksf.sample_rate, channels=ksf.channels,
        bit_depth=ksf.bit_depth, loop_start=ksf.loop_start, loop_end=ksf.loop_end,
        is_compressed=False, compression_id=flags["compression_id"],
        use_2nd_start=flags["use_2nd_start"], raw_bytes=pcm,
    )


def encode_pcm(samples: np.ndarray, bit_depth: int) -> bytes:
    """Round-trip encoder: float32 [-1,1] -> big-endian signed PCM bytes.
    Used by the Golden Corpus byte-for-byte check on the audio data."""
    samples = np.clip(samples, -1.0, 1.0)
    if bit_depth == 16:
        ints = np.round(samples * 32767.0).astype(">i2")
        return ints.tobytes()
    elif bit_depth == 8:
        ints = np.round(samples * 127.0).astype(np.int8)
        return ints.tobytes()
    raise ValueError(f"Unsupported bit depth {bit_depth}")


def roundtrip_matches(ksf, atol: int = 0) -> bool:
    """True if decode->encode reproduces the original PCM bytes exactly.
    A mismatch means an assumption (endian, signedness, channels) is WRONG."""
    decoded = decode_pcm(ksf)
    if decoded.is_compressed:
        return True  # can't verify compressed; just preserved
    flat = decoded.samples.reshape(-1) if decoded.channels == 1 else decoded.samples
    re = encode_pcm(flat, ksf.bit_depth)
    return re == decoded.raw_bytes


def export_wav(decoded: DecodedPCM, out_path: str) -> None:
    """Write a standard WAV for audible verification."""
    if decoded.is_compressed:
        raise ValueError("Cannot export compressed sample to WAV (codec unknown)")
    with wave.open(out_path, "wb") as w:
        w.setnchannels(decoded.channels)
        w.setsampwidth(decoded.bit_depth // 8)
        w.setframerate(decoded.sample_rate)
        # WAV expects little-endian; convert from our float -> little-endian PCM
        flat = decoded.samples.reshape(-1)
        if decoded.bit_depth == 16:
            ints = np.round(np.clip(flat, -1, 1) * 32767.0).astype("<i2")
            w.writeframes(ints.tobytes())
        else:
            ints = np.round(np.clip(flat, -1, 1) * 127.0).astype(np.int8)
            w.writeframes(ints.tobytes())


def apply_loop(decoded: DecodedPCM, sustain: bool = True) -> np.ndarray:
    """Return the sample with the loop region applied. LoopStart/LoopEnd are
    in FRAMES (verified: they are sample-point indices, not byte offsets)."""
    if not (0 <= decoded.loop_start < decoded.loop_end <= decoded.samples.shape[0]):
        return decoded.samples
    body = decoded.samples[:decoded.loop_start]
    if sustain:
        loop = decoded.samples[decoded.loop_start:decoded.loop_end]
        return np.concatenate([body, np.tile(loop, (4, 1))]) if decoded.channels > 1 \
            else np.concatenate([body, np.tile(loop, 4)])
    return decoded.samples


if __name__ == "__main__":
    import sys
    from korg_decoder import parse_ksf
    path = sys.argv[1]
    out = sys.argv[2] if len(sys.argv) > 2 else path.rsplit(".", 1)[0] + ".wav"
    with open(path, "rb") as f:
        data = f.read()
    ksf = parse_ksf(data)
    decoded = decode_pcm(ksf)
    flags = attributes_flags(ksf.attributes)
    print(f"KSF: {ksf.name}")
    print(f"  sr={ksf.sample_rate} bd={ksf.bit_depth} ch={ksf.channels} pts={ksf.sample_points}")
    print(f"  compressed={flags['is_compressed']} comp_id={flags['compression_id']} "
          f"use_2nd_start={flags['use_2nd_start']}")
    print(f"  loop {ksf.loop_start}..{ksf.loop_end} frames")
    print(f"  pcm bytes={len(decoded.raw_bytes)} expected={ksf.sample_points*frame_size(ksf.bit_depth,ksf.channels)}")
    print(f"  round-trip exact: {roundtrip_matches(ksf)}")
    if not decoded.is_compressed:
        export_wav(decoded, out)
        print(f"  -> wrote {out}")
    else:
        print("  -> COMPRESSED: not exporting WAV (codec unknown). Raw bytes preserved.")
איך זה סוגר את הפער של Codex: הוא מקבל את היסט ה-Attributes המאומת (bit 4/2/5), את נוסחת ה-FrameSize, את נקודת ההתחלה (offset 12), ואת הכלל "דחוס = לא מפענחים". roundtrip_matches הוא השומר — אם decode→encode לא משחזר את הבתים המקוריים ביט-לביט, הנחה (endian/signedness/channels) שגויה והוא חייב לתקן, לא להתעלם.