use std::io::{ErrorKind, IsTerminal, Read, Write};
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::Arc;
use std::thread::sleep;
use std::time::{Duration, Instant};
use anyhow::{anyhow, bail, Context, Result};
use crossterm::terminal::{disable_raw_mode, enable_raw_mode};
use serialport::{ClearBuffer, SerialPort};
const HELLO: &[u8; 4] = b"RPIL";
const ACK: &[u8; 4] = b"LIPR";
const PROTOCOL_VERSION: u8 = 2;
const OK: u8 = 1;
const FAIL: u8 = 0;
const CHUNK_SIZE: usize = 4096;
const CMD_MEM_WRITE: u8 = 1;
const CMD_SET_BAUD: u8 = 2;
const CMD_EXEC: u8 = 3;
const CMD_SD_LIST: u8 = 4;
const CMD_SD_READ: u8 = 5;
const CMD_SD_WRITE: u8 = 6;
const CMD_SD_DELETE: u8 = 7;
const CMD_SD_MKDIR: u8 = 8;
const CMD_EEPROM_READ: u8 = 9;
const CMD_EEPROM_WRITE: u8 = 10;
pub const BASE_BAUD: u32 = 115_200;
pub const DEFAULT_BAUD: u32 = 1_500_000;
const RESPONSE_TIMEOUT: Duration = Duration::from_secs(2);
const MAX_CHUNK_RETRIES: u32 = 5;
const SD_STATUS_ATTEMPTS: u32 = 3;
const EEPROM_STATUS_ATTEMPTS: u32 = 4;
const DRAIN_LIMIT: Duration = Duration::from_millis(500);
fn err_name(code: u8) -> String {
match code {
1 => "SD bring-up failed".into(),
2 => "no such file or directory".into(),
3 => "filesystem error".into(),
4 => "directory listing too large".into(),
5 => "bad path".into(),
6 => "write failed".into(),
7 => "I2C transfer failed (nothing answering at that address, or the bus is held)".into(),
8 => "read-back did not match what was written (is the EEPROM write-protected?)".into(),
9 => "the device will not address that range (offset + length past 64 KiB, \
or an implausible page size)"
.into(),
10 => "a page was written, but the part never answered the read that checks it".into(),
other => format!("error code {other}"),
}
}
pub struct Entry {
pub is_dir: bool,
pub size: u64,
pub name: String,
}
pub struct Link {
port: Box<dyn SerialPort>,
interrupted: Arc<AtomicBool>,
}
impl Link {
pub fn open(device: &str, interrupted: Arc<AtomicBool>) -> Result<Self> {
let port = serialport::new(device, BASE_BAUD)
.timeout(RESPONSE_TIMEOUT)
.open()
.with_context(|| format!("opening {device}"))?;
Ok(Self { port, interrupted })
}
pub fn interrupted(&self) -> bool {
self.interrupted.load(Ordering::SeqCst)
}
fn check_interrupt(&self) -> Result<()> {
if self.interrupted() {
bail!("interrupted");
}
Ok(())
}
fn read_exact_or_timeout(&mut self, n: usize) -> Result<Option<Vec<u8>>> {
let mut buf = vec![0u8; n];
let mut filled = 0;
while filled < n {
match self.port.read(&mut buf[filled..]) {
Ok(0) => return Ok(None),
Ok(got) => filled += got,
Err(e) if e.kind() == ErrorKind::TimedOut => return Ok(None),
Err(e) if e.kind() == ErrorKind::Interrupted => continue,
Err(e) => return Err(e).context("reading from the serial port"),
}
}
Ok(Some(buf))
}
fn read_status(&mut self, attempts: u32) -> Result<Option<u8>> {
for _ in 0..attempts {
if let Some(b) = self.read_exact_or_timeout(1)? {
return Ok(Some(b[0]));
}
}
Ok(None)
}
fn fail_reason(&mut self) -> Result<String> {
Ok(match self.read_status(1)? {
Some(code) => err_name(code),
None => "timed out reading error code".into(),
})
}
fn wait_for_ack_and_version(&mut self) -> Result<Option<u8>> {
let mut matched = 0;
while matched < ACK.len() {
let Some(b) = self.read_exact_or_timeout(1)? else {
return Ok(None);
};
let value = b[0];
if value == ACK[matched] {
matched += 1;
} else if value == ACK[0] {
matched = 1;
} else {
matched = 0;
if (0x20..=0x7E).contains(&value) || value == b'\n' || value == b'\r' {
let mut out = std::io::stdout();
out.write_all(&b)?;
out.flush()?;
}
}
}
Ok(self.read_exact_or_timeout(1)?.map(|b| b[0]))
}
pub fn handshake(&mut self) -> Result<u8> {
self.port.clear(ClearBuffer::All).ok();
eprintln!("Connecting to rpi-loader (sending HELLO)...");
let version = loop {
self.check_interrupt()?;
self.write_all(HELLO)?;
if let Some(version) = self.wait_for_ack_and_version()? {
break version;
}
};
if version != PROTOCOL_VERSION {
eprintln!("Warning: device protocol version {version}, expected {PROTOCOL_VERSION}");
}
sleep(Duration::from_millis(20));
self.port.clear(ClearBuffer::Input).ok();
Ok(version)
}
pub fn negotiate_baud(&mut self, baud: u32) -> Result<u32> {
let current = self.port.baud_rate().context("reading the current baud")?;
if current == baud {
return Ok(baud);
}
let mut packet = vec![CMD_SET_BAUD];
packet.extend_from_slice(&baud.to_le_bytes());
self.write_all(&packet)?;
if self.read_status(1)? != Some(OK) {
eprintln!("Device declined {baud} baud; staying at {current}");
return Ok(current);
}
self.port
.set_baud_rate(baud)
.with_context(|| format!("switching the host to {baud} baud"))?;
sleep(Duration::from_millis(50));
self.port.clear(ClearBuffer::Input).ok();
Ok(baud)
}
fn send_chunked(&mut self, data: &[u8], ack_attempts: u32) -> Result<()> {
let total = data.len();
for (offset, chunk) in (0..).step_by(CHUNK_SIZE).zip(data.chunks(CHUNK_SIZE)) {
let mut packet = crc32(chunk).to_le_bytes().to_vec();
packet.extend_from_slice(chunk);
let mut sent = false;
for attempt in 1..=MAX_CHUNK_RETRIES {
self.check_interrupt()?;
self.write_all(&packet)?;
if self.read_status(ack_attempts)? == Some(OK) {
sent = true;
break;
}
eprintln!(
" chunk at {offset} failed/timed out, retry {attempt}/{MAX_CHUNK_RETRIES}"
);
}
if !sent {
bail!("giving up on chunk at {offset} after {MAX_CHUNK_RETRIES} retries");
}
eprintln!(" {}/{total}", (offset + CHUNK_SIZE).min(total));
}
Ok(())
}
fn recv_chunked(&mut self) -> Result<Vec<u8>> {
let header = self
.read_exact_or_timeout(8)?
.ok_or_else(|| anyhow!("timed out waiting for stream header"))?;
let total_len = u32::from_le_bytes(header[0..4].try_into().unwrap()) as usize;
let chunk_size = u32::from_le_bytes(header[4..8].try_into().unwrap()) as usize;
let mut out = Vec::with_capacity(total_len);
while out.len() < total_len {
let this = chunk_size.min(total_len - out.len());
let mut received = false;
for attempt in 1..=MAX_CHUNK_RETRIES {
self.check_interrupt()?;
let crc_bytes = self.read_exact_or_timeout(4)?;
let data = self.read_exact_or_timeout(this)?;
let (Some(crc_bytes), Some(data)) = (crc_bytes, data) else {
bail!("timed out mid-stream");
};
if crc32(&data) == u32::from_le_bytes(crc_bytes.try_into().unwrap()) {
self.write_all(&[OK])?;
out.extend_from_slice(&data);
received = true;
break;
}
self.write_all(&[FAIL])?;
eprintln!(
" chunk at {} bad CRC, retry {attempt}/{MAX_CHUNK_RETRIES}",
out.len()
);
}
if !received {
bail!(
"giving up on chunk at {} after {MAX_CHUNK_RETRIES} retries",
out.len()
);
}
}
Ok(out)
}
fn send_path(&mut self, path: &str) -> Result<()> {
let encoded = path.as_bytes();
let len = u16::try_from(encoded.len()).context("path is too long for the protocol")?;
let mut packet = len.to_le_bytes().to_vec();
packet.extend_from_slice(encoded);
self.write_all(&packet)
}
fn write_all(&mut self, data: &[u8]) -> Result<()> {
self.port
.write_all(data)
.context("writing to the serial port")
}
pub fn mem_write(&mut self, addr: u32, data: &[u8]) -> Result<()> {
let mut header = vec![CMD_MEM_WRITE];
header.extend_from_slice(&(data.len() as u32).to_le_bytes());
header.extend_from_slice(&(CHUNK_SIZE as u32).to_le_bytes());
header.extend_from_slice(&addr.to_le_bytes());
header.extend_from_slice(&crc32(data).to_le_bytes());
self.write_all(&header)?;
if self.read_status(1)? != Some(OK) {
bail!("device rejected header (bad size/address?)");
}
eprintln!("Sending {} bytes to {addr:#x}...", data.len());
self.send_chunked(data, 1)?;
if self.read_status(1)? != Some(OK) {
bail!("device reported overall checksum mismatch");
}
Ok(())
}
pub fn exec(&mut self, addr: u32) -> Result<()> {
let mut packet = vec![CMD_EXEC];
packet.extend_from_slice(&addr.to_le_bytes());
self.write_all(&packet)?;
if self.read_status(1)? != Some(OK) {
bail!("device refused to exec {addr:#x} (bad address?)");
}
Ok(())
}
fn start_sd_command(&mut self, command: u8, path: &str, what: &str) -> Result<()> {
self.write_all(&[command])?;
self.send_path(path)?;
if self.read_status(SD_STATUS_ATTEMPTS)? != Some(OK) {
let reason = self.fail_reason()?;
bail!("{what} failed: {reason}");
}
Ok(())
}
pub fn sd_list(&mut self, path: &str) -> Result<Vec<Entry>> {
self.start_sd_command(CMD_SD_LIST, path, "sd-list")?;
let listing = self.recv_chunked()?;
let listing = String::from_utf8_lossy(&listing);
listing
.lines()
.map(|line| {
let mut fields = line.splitn(3, '\t');
let (Some(kind), Some(size), Some(name)) =
(fields.next(), fields.next(), fields.next())
else {
bail!("malformed listing line from the device: {line:?}");
};
Ok(Entry {
is_dir: kind == "D",
size: size
.parse()
.context("listing entry has a non-numeric size")?,
name: name.to_string(),
})
})
.collect()
}
pub fn sd_read(&mut self, remote: &str) -> Result<Vec<u8>> {
self.start_sd_command(CMD_SD_READ, remote, "sd-read")?;
self.recv_chunked()
}
pub fn sd_write(&mut self, remote: &str, data: &[u8]) -> Result<()> {
self.write_all(&[CMD_SD_WRITE])?;
self.send_path(remote)?;
let mut header = (data.len() as u32).to_le_bytes().to_vec();
header.extend_from_slice(&(CHUNK_SIZE as u32).to_le_bytes());
self.write_all(&header)?;
if self.read_status(SD_STATUS_ATTEMPTS)? != Some(OK) {
let reason = self.fail_reason()?;
bail!("sd-write failed: {reason}");
}
eprintln!("Sending {} bytes -> {remote}...", data.len());
self.send_chunked(data, 1)?;
if self.read_status(SD_STATUS_ATTEMPTS)? != Some(OK) {
let reason = self.fail_reason()?;
bail!("sd-write did not commit: {reason}");
}
Ok(())
}
pub fn sd_delete(&mut self, remote: &str) -> Result<()> {
self.start_sd_command(CMD_SD_DELETE, remote, "sd-delete")
}
pub fn sd_mkdir(&mut self, remote: &str) -> Result<()> {
self.start_sd_command(CMD_SD_MKDIR, remote, "sd-mkdir")
}
pub fn eeprom_read(&mut self, address: u8, offset: u32, length: u32) -> Result<Vec<u8>> {
let mut packet = vec![CMD_EEPROM_READ, address];
packet.extend_from_slice(&offset.to_le_bytes());
packet.extend_from_slice(&length.to_le_bytes());
self.write_all(&packet)?;
if self.read_status(EEPROM_STATUS_ATTEMPTS)? != Some(OK) {
let reason = self.fail_reason()?;
bail!("eeprom-read failed: {reason}");
}
self.recv_chunked()
}
pub fn eeprom_write(
&mut self,
address: u8,
offset: u32,
page_size: u32,
data: &[u8],
) -> Result<()> {
let mut packet = vec![CMD_EEPROM_WRITE, address];
packet.extend_from_slice(&offset.to_le_bytes());
packet.extend_from_slice(&(data.len() as u32).to_le_bytes());
packet.extend_from_slice(&(CHUNK_SIZE as u32).to_le_bytes());
packet.extend_from_slice(&page_size.to_le_bytes());
self.write_all(&packet)?;
if self.read_status(EEPROM_STATUS_ATTEMPTS)? != Some(OK) {
let reason = self.fail_reason()?;
bail!("eeprom-write failed: {reason}");
}
eprintln!(
"Programming {} bytes at offset {offset} of 0x{address:02x}...",
data.len()
);
self.send_chunked(data, EEPROM_STATUS_ATTEMPTS)?;
if self.read_status(EEPROM_STATUS_ATTEMPTS)? != Some(OK) {
let reason = self.fail_reason()?;
bail!("eeprom-write did not commit: {reason}");
}
Ok(())
}
pub fn terminal(&mut self) -> Result<()> {
let mut from_device = self
.port
.try_clone()
.context("cloning the serial port for the terminal")?;
from_device
.set_timeout(Duration::from_millis(100))
.context("setting the terminal read timeout")?;
let raw = RawMode::enable();
if raw.active {
eprintln!(
"Entering terminal mode. Ctrl-] exits; everything else, \
Ctrl-C included, goes to the device.\r"
);
} else {
eprintln!("Entering terminal mode (stdin is not a terminal; Ctrl-C exits).");
}
let finished = Arc::new(AtomicBool::new(false));
let reader_finished = Arc::clone(&finished);
let reader_interrupted = Arc::clone(&self.interrupted);
let reader = std::thread::spawn(move || {
let mut buf = [0u8; 256];
let mut out = std::io::stdout();
while !reader_finished.load(Ordering::SeqCst)
&& !reader_interrupted.load(Ordering::SeqCst)
{
match from_device.read(&mut buf) {
Ok(0) => {}
Ok(n) => {
if out.write_all(&buf[..n]).is_err() || out.flush().is_err() {
break;
}
}
Err(e) if e.kind() == ErrorKind::TimedOut => {}
Err(e) if e.kind() == ErrorKind::Interrupted => {}
Err(_) => break,
}
}
let deadline = Instant::now() + DRAIN_LIMIT;
while Instant::now() < deadline {
match from_device.read(&mut buf) {
Ok(0) => break,
Ok(n) => {
if out.write_all(&buf[..n]).is_err() || out.flush().is_err() {
break;
}
}
_ => break,
}
}
reader_finished.store(true, Ordering::SeqCst);
});
let result = self.forward_keyboard(&finished);
finished.store(true, Ordering::SeqCst);
let _ = reader.join();
drop(raw);
eprintln!();
eprintln!("Exiting.");
result
}
fn forward_keyboard(&mut self, finished: &AtomicBool) -> Result<()> {
let mut stdin = std::io::stdin().lock();
let mut buf = [0u8; 256];
loop {
if self.interrupted() || finished.load(Ordering::SeqCst) {
return Ok(());
}
let n = match stdin.read(&mut buf) {
Ok(0) => break,
Ok(n) => n,
Err(e) if e.kind() == ErrorKind::Interrupted => continue,
Err(e) => return Err(e).context("reading the keyboard"),
};
if let Some(at) = buf[..n].iter().position(|&b| b == ESCAPE) {
self.write_all(&buf[..at])?;
return Ok(());
}
self.write_all(&buf[..n])?;
}
while !self.interrupted() && !finished.load(Ordering::SeqCst) {
sleep(Duration::from_millis(100));
}
Ok(())
}
}
const ESCAPE: u8 = 0x1D;
struct RawMode {
active: bool,
}
impl RawMode {
fn enable() -> Self {
if !std::io::stdin().is_terminal() {
return Self { active: false };
}
Self {
active: enable_raw_mode().is_ok(),
}
}
}
impl Drop for RawMode {
fn drop(&mut self) {
if self.active {
let _ = disable_raw_mode();
}
}
}
fn crc32(data: &[u8]) -> u32 {
let mut hasher = crc32fast::Hasher::new();
hasher.update(data);
hasher.finalize()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn raw_mode_stays_off_without_a_terminal_on_stdin() {
if std::io::stdin().is_terminal() {
return;
}
assert!(!RawMode::enable().active);
}
}