#!/usr/bin/env catnip
# Stéganographie LSB (Least Significant Bit) avec pyvips
#
# On cache un message texte dans le bit de poids faible de chaque canal R/G/B.
# Modifier ce bit décale une valeur de canal de ±1 au plus : invisible à l'œil,
# mais suffisant pour transporter un bit par canal. pyvips assure le load/save,
# numpy fait le bit-twiddling en bloc (jamais canal par canal depuis Catnip).
#
# Le porteur de sortie est un PNG : la stégano LSB ne survit pas à la
# recompression avec perte du JPEG, il faut un format sans perte.
#
# Ref : https://en.wikipedia.org/wiki/Steganography
#
# DEPS: pyvips[binary], numpy

pyvips = import('pyvips')
numpy = import('numpy')
import('pathlib', 'Path')

script_dir = Path(META.file).parent
input_path = script_dir / 'data/catnip_test.jpg'
output_dir = script_dir / 'output'
output_dir.mkdir(exist_ok=True)

# En-tête de 32 bits (4 octets) = longueur du message en octets, big-endian.
# C'est le marqueur de fin : le décodeur lit exactement ce nombre d'octets après
# l'en-tête, donc pas besoin de sentinelle dans le flux lui-même.
HEADER_BITS = 32

struct StegoResult {
    carrier: str; message_length: int; bits_used: int; capacity: int;
}

# Longueur -> 4 octets big-endian. n < 2^32, chaque octet est une tranche de 8
# bits extraite par division entière puis modulo 256.
length_header = (n: int) => {
    numpy.array(
        list(int(n / 16777216) % 256, int(n / 65536) % 256, int(n / 256) % 256, n % 256),
        dtype='uint8',
    )
}

# Flux binaire complet : en-tête de longueur suivi des octets UTF-8 du message,
# le tout déplié en un bit par position.
message_bits = (message: str) => {
    payload = numpy.frombuffer(message.encode('utf-8'), dtype='uint8')
    stream = numpy.concatenate(list(length_header(payload.size), payload))
    numpy.unpackbits(stream)
}

# Écrase le LSB des premiers canaux par les bits du flux, laisse le reste intact.
# bitwise_and(x, 254) met le LSB à 0, bitwise_or réinjecte le bit voulu.
hide = (image, message: str) => {
    flat = image.numpy().reshape(-1)
    bits = message_bits(message)
    nbits = bits.size
    head = numpy.bitwise_or(numpy.bitwise_and(flat[:nbits], 254), bits)
    stego = numpy.concatenate(list(head, flat[nbits:]))
    pyvips.Image.new_from_array(stego.reshape(image.height, image.width, image.bands))
}

# Extrait le LSB de chaque canal, relit l'en-tête pour connaître la longueur,
# puis recompose les octets du message et les décode en UTF-8.
reveal = (image) => {
    lsb = numpy.bitwise_and(image.numpy().reshape(-1), 1)
    header = numpy.packbits(lsb[:HEADER_BITS])
    length = int(header[0]) * 16777216 + int(header[1]) * 65536 + int(header[2]) * 256 + int(header[3])
    payload_bits = lsb[HEADER_BITS:HEADER_BITS + length * 8]
    numpy.packbits(payload_bits).tobytes().decode('utf-8')
}

message = "Catnip cache un secret dans les bits de poids faible des pixels."

print("⇒ Chargement du porteur")
source = pyvips.Image.new_from_file(str(input_path)).colourspace('srgb')
capacity = source.width * source.height * source.bands
print(f"  {source.width}×{source.height}, {source.bands} bandes → {capacity} emplacements LSB")

message_length = len(message.encode('utf-8'))
bits_used = HEADER_BITS + message_length * 8
print()
print("⇒ Encodage du message")
print(f"  message : {message_length} octets, {bits_used} bits (en-tête + charge utile)")
if bits_used > capacity {
    print("  ERREUR : message trop long pour ce porteur")
} else {
    print(f"  occupation : {round(100.0 * bits_used / capacity, 4)} % de la capacité")
}

stego = hide(source, message)
carrier_path = output_dir / 'pyvips_steganography.png'
stego.write_to_file(str(carrier_path))
print(f"  porteur → {carrier_path.name}")

result = StegoResult(carrier_path.name, message_length, bits_used, capacity)
print()
print("⇒ StegoResult")
print(f"  carrier={result.carrier} message_length={result.message_length}")
print(f"  bits_used={result.bits_used} capacity={result.capacity}")

# Vérification roundtrip : on relit le PNG écrit sur disque (chemin indépendant
# de l'image en mémoire) et on décode. decode(encode(m)) doit rendre m.
print()
print("⇒ Vérification roundtrip")
reloaded = pyvips.Image.new_from_file(str(carrier_path))
recovered = reveal(reloaded)
print(f"  message décodé : {recovered}")
print(f"  roundtrip == message original : {recovered == message}")