shiftreg-spi 0.1.0

SPI-based driver for shift registers such as 74HC595 with embedded-hal API
Documentation
use core::cell::RefCell;

use critical_section::Mutex;
use embedded_hal::{
    digital::{ErrorType, OutputPin, PinState},
    spi::SpiDevice,
};

use crate::ShiftRegError;

// The design of this struct is inspired by `port-expander` clate ( https://crates.io/crates/port-expander ).
pub struct SipoShiftReg<Spi, const BITS: usize, const BYTES: usize>(
    Mutex<RefCell<SipoShiftRegInner<Spi, BITS, BYTES>>>,
);

impl<Spi, const BITS: usize, const BYTES: usize> SipoShiftReg<Spi, BITS, BYTES>
where
    Spi: SpiDevice,
{
    pub fn new(spi: Spi) -> Self {
        Self(Mutex::new(RefCell::new(SipoShiftRegInner {
            spi,
            state: [0; BYTES],
        })))
    }

    pub fn split<'a>(&'a self) -> [SipoShiftRegPin<'a, Spi, BITS, BYTES>; BITS] {
        core::array::from_fn(|i| SipoShiftRegPin {
            shift_reg: &self.0,
            idx: i,
        })
    }
}

struct SipoShiftRegInner<Spi, const BITS: usize, const BYTES: usize> {
    spi: Spi,
    state: [u8; BYTES],
}

impl<Spi, const BITS: usize, const BYTES: usize> SipoShiftRegInner<Spi, BITS, BYTES>
where
    Spi: SpiDevice,
{
    fn set_pin_state(
        &mut self,
        idx: usize,
        state: PinState,
    ) -> Result<(), ShiftRegError<Spi::Error>> {
        let byte_idx = idx / 8;
        let bit_idx = idx % 8;

        match state {
            PinState::High => self.state[byte_idx] |= 1 << bit_idx,
            PinState::Low => self.state[byte_idx] &= !(1 << bit_idx),
        };

        self.update()
    }

    fn update(&mut self) -> Result<(), ShiftRegError<Spi::Error>> {
        self.spi.write(&self.state).map_err(ShiftRegError::Spi)
    }
}

pub struct SipoShiftRegPin<'a, Spi, const BITS: usize, const BYTES: usize> {
    shift_reg: &'a Mutex<RefCell<SipoShiftRegInner<Spi, BITS, BYTES>>>,
    idx: usize,
}

impl<'a, Spi, const BITS: usize, const BYTES: usize> ErrorType
    for SipoShiftRegPin<'a, Spi, BITS, BYTES>
where
    Spi: SpiDevice,
{
    type Error = ShiftRegError<Spi::Error>;
}

impl<'a, Spi, const BITS: usize, const BYTES: usize> OutputPin
    for SipoShiftRegPin<'a, Spi, BITS, BYTES>
where
    Spi: SpiDevice,
{
    fn set_high(&mut self) -> Result<(), Self::Error> {
        self.set_state(PinState::High)
    }

    fn set_low(&mut self) -> Result<(), Self::Error> {
        self.set_state(PinState::Low)
    }

    fn set_state(&mut self, state: PinState) -> Result<(), Self::Error> {
        critical_section::with(|cs| {
            self.shift_reg
                .borrow_ref_mut(cs)
                .set_pin_state(self.idx, state)
        })
    }
}