selene-core 0.9.0-rc.1

Backend for selene-server
Documentation
use std::{collections::VecDeque, io};

use ogg::{PacketWriteEndInfo, PacketWriter};
use opusic_c::{Application, Channels, ErrorCode, SampleRate};
use rubato::{
    Fft, Resampler, ResamplerConstructionError, audioadapter_buffers::direct::InterleavedSlice,
};
use selene_common::api::{ApiError, Bitrate};

use crate::transcoding::decoder::{Decoder, DecodingError};

#[derive(Debug, thiserror::Error)]
pub enum OpusTranscodeError {
    #[error("{0}")]
    Io(#[from] io::Error),

    #[error("One or more invalid/out of range arguments")]
    BadArg,
    #[error("Memory allocation has failed")]
    AllocFail,
    #[error("An encoder or decoder structure is invalid or already freed")]
    InvalidState,
    #[error("The compressed data passed is corrupted")]
    InvalidPacket,
    #[error("Not enough bytes allocated in the buffer")]
    BufferTooSmall,
    #[error("An internal error was detected")]
    Internal,
    #[error("Invalid/unsupported request number")]
    Unimplemented,

    #[error("The libopus encoder only supports mono or stereo audio")]
    UnsupportedChannelCount,

    #[error("{0}")]
    DecodeError(#[from] DecodingError),

    #[error("{0}")]
    ResamplerConstructionError(#[from] ResamplerConstructionError),
}

impl From<OpusTranscodeError> for ApiError {
    fn from(_: OpusTranscodeError) -> Self {
        ApiError::INTERNAL_SERVER_ERROR
    }
}

#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OpusEncoderBitrate {
    Hz320000,
    Hz192000,
    Hz128000,
}

impl From<Bitrate> for OpusEncoderBitrate {
    fn from(value: Bitrate) -> Self {
        match value {
            Bitrate::B128 => Self::Hz128000,
            Bitrate::B192 => Self::Hz192000,
            Bitrate::B320 => Self::Hz320000,
        }
    }
}

impl From<u32> for OpusEncoderBitrate {
    fn from(value: u32) -> Self {
        if value == 0 {
            Self::Hz320000
        } else if value > 320_000 {
            Self::Hz320000
        } else if value > 192_000 {
            Self::Hz192000
        } else {
            Self::Hz128000
        }
    }
}

impl OpusEncoderBitrate {
    #[must_use]
    pub fn bitrate(&self) -> u32 {
        match self {
            OpusEncoderBitrate::Hz320000 => 320000,
            OpusEncoderBitrate::Hz192000 => 192000,
            OpusEncoderBitrate::Hz128000 => 128000,
        }
    }
}

impl From<ErrorCode> for OpusTranscodeError {
    fn from(value: ErrorCode) -> Self {
        match value {
            ErrorCode::Ok => panic!("Not an error"),
            ErrorCode::BadArg => OpusTranscodeError::BadArg,
            ErrorCode::AllocFail => OpusTranscodeError::AllocFail,
            ErrorCode::InvalidState => OpusTranscodeError::InvalidState,
            ErrorCode::InvalidPacket => OpusTranscodeError::InvalidPacket,
            ErrorCode::BufferTooSmall => OpusTranscodeError::BufferTooSmall,
            ErrorCode::Internal => OpusTranscodeError::Internal,
            ErrorCode::Unimplemented => OpusTranscodeError::Unimplemented,
            ErrorCode::Unknown => unreachable!("Marked as 'should be impossible'"),
        }
    }
}

pub struct OpusTranscoder {
    decoder: Decoder,
    resampler: Option<(Fft<f32>, VecDeque<f32>)>,
    encoder: opusic_c::Encoder,

    channels: usize,

    pcm_buffer: VecDeque<f32>,
    output_buf: [u8; 4000],
}

impl OpusTranscoder {
    const FRAME_SIZE_SAMPLES: usize = 960;
    const OUTPUT_SAMPLE_RATE: usize = SampleRate::Hz48000 as usize;

    pub fn new(
        decoder: Decoder,
        bitrate: OpusEncoderBitrate,
        resampler_chunk_size: Option<usize>,
    ) -> Result<OpusTranscoder, OpusTranscodeError> {
        let channel_count = decoder.stream.codec_params.channels;
        let channels = match channel_count {
            1 => Channels::Mono,
            2 => Channels::Stereo,
            _ => return Err(OpusTranscodeError::UnsupportedChannelCount),
        };
        let lsb_depth = decoder.stream.codec_params.bits_per_sample;

        let resampler = if decoder.stream.codec_params.sample_rate == SampleRate::Hz48000 as u32 {
            None
        } else {
            let fft = Fft::<f32>::new(
                decoder.stream.codec_params.sample_rate as usize,
                Self::OUTPUT_SAMPLE_RATE,
                resampler_chunk_size.unwrap_or(1024),
                channel_count,
                rubato::FixedSync::Both,
            )?;

            Some((fft, VecDeque::new()))
        };

        let mut encoder =
            opusic_c::Encoder::new(channels, SampleRate::Hz48000, Application::Audio)?;
        encoder.set_complexity(10)?;
        encoder.set_packet_loss(0)?;
        encoder.set_inband_fec(opusic_c::InbandFec::Off)?;
        encoder.set_signal(opusic_c::Signal::Music)?;
        encoder.set_max_bandwidth(opusic_c::Bandwidth::Full)?;

        let lsb_depth = lsb_depth.unwrap_or(24).clamp(8, 24) as u8;
        encoder.set_lsb_depth(lsb_depth)?;

        encoder.set_vbr_constraint(false)?;
        encoder.set_vbr(true)?;

        encoder.set_bitrate(opusic_c::Bitrate::Value(bitrate.bitrate()))?;

        Ok(OpusTranscoder {
            encoder,
            decoder,
            resampler,
            pcm_buffer: VecDeque::new(),

            channels: channel_count,
            output_buf: [0u8; 4000],
        })
    }

    pub fn encode_next_packet(&mut self) -> Result<(Option<&[u8]>, bool), OpusTranscodeError> {
        let needed = Self::FRAME_SIZE_SAMPLES * self.channels;

        loop {
            if self.pcm_buffer.len() >= needed {
                let frame = &self.pcm_buffer.make_contiguous()[..needed];
                let encoded_bytes = self
                    .encoder
                    .encode_float_to_slice(frame, &mut self.output_buf)?;
                self.pcm_buffer.drain(..needed);
                return Ok((
                    Some(&self.output_buf[..encoded_bytes]),
                    self.decoder.at_eof && self.pcm_buffer.is_empty(),
                ));
            }

            let Some(packet) = self.decoder.decode_next_packet()? else {
                if self.pcm_buffer.is_empty() {
                    return Ok((None, true));
                }
                self.pcm_buffer.resize(needed, 0.0);
                let encoded_bytes = self.encoder.encode_float_to_slice(
                    self.pcm_buffer.make_contiguous(),
                    &mut self.output_buf,
                )?;
                self.pcm_buffer.clear();
                return Ok((Some(&self.output_buf[..encoded_bytes]), true));
            };

            let mut frames = vec![0.0; packet.samples_interleaved()];
            packet.copy_to_slice_interleaved(&mut frames);

            if let Some((resampler, resample_buffer)) = &mut self.resampler {
                let next_frame_size = resampler.input_frames_next();
                let chunk_size = next_frame_size * self.channels;

                resample_buffer.extend(frames);

                while resample_buffer.len() >= chunk_size {
                    let buffer_in = InterleavedSlice::new(
                        &resample_buffer.make_contiguous()[..chunk_size],
                        self.channels,
                        next_frame_size,
                    )
                    .expect("Resample buffer verified to contain correct capacity");

                    let resampled = resampler
                        .process(&buffer_in, None)
                        .expect("Resampler input and output have the same number of channels");

                    resample_buffer.drain(..chunk_size);
                    self.pcm_buffer.extend(resampled.take_data());
                }
            } else {
                self.pcm_buffer.extend(frames);
            }
        }
    }
}

/// Encoding to an ogg/opus file
impl OpusTranscoder {
    fn build_opus_head(&self) -> Vec<u8> {
        const VERSION: u8 = 1;
        const PRE_SKIP: &[u8] = &3124_u16.to_le_bytes();
        const OUTPUT_GAIN: &[u8] = &0_i16.to_le_bytes();

        let mut head = Vec::with_capacity(19);

        // Signature
        head.extend_from_slice(b"OpusHead");

        head.push(VERSION);
        head.push(self.channels as u8);
        head.extend_from_slice(PRE_SKIP);

        let input_sample_rate = self.decoder.stream.codec_params.sample_rate.to_le_bytes();
        head.extend_from_slice(&input_sample_rate);

        head.extend_from_slice(OUTPUT_GAIN);

        // Mapping family
        head.push(0);

        head
    }

    fn build_opus_tags() -> Vec<u8> {
        let mut tags = Vec::new();

        // Signature
        tags.extend_from_slice(b"OpusTags");

        // Vendor
        let vendor = format!("Selene v{}", env!("CARGO_PKG_VERSION")).into_bytes();
        tags.extend_from_slice(&(vendor.len() as u32).to_le_bytes());
        tags.extend(vendor);

        // Comment
        const COMMENT_COUNT: &[u8] = &0_u32.to_le_bytes();
        tags.extend_from_slice(COMMENT_COUNT);

        tags
    }

    pub fn encode_to_io<W: io::Write>(&mut self, writer: W) -> Result<(), OpusTranscodeError> {
        const SERIAL: u32 = 1;
        let mut ogg = PacketWriter::new(writer);

        let head = self.build_opus_head();
        ogg.write_packet(head, SERIAL, PacketWriteEndInfo::EndPage, 0)?;

        let tags = Self::build_opus_tags();
        ogg.write_packet(tags, SERIAL, PacketWriteEndInfo::EndPage, 0)?;

        let mut total_samples: u64 = 0;

        while let (Some(packet), at_eof) = self.encode_next_packet()? {
            total_samples += Self::FRAME_SIZE_SAMPLES as u64;

            let end_info = if at_eof {
                PacketWriteEndInfo::EndStream
            } else {
                PacketWriteEndInfo::NormalPacket
            };

            ogg.write_packet(packet.to_vec(), SERIAL, end_info, total_samples)?;

            if at_eof {
                debug_assert!(
                    self.pcm_buffer.is_empty(),
                    "Expected PCM buffer to always be empty when encoder is EOF"
                );
                break;
            }
        }

        Ok(())
    }
}