use std::collections::BTreeMap;
use std::fs;
use std::io::{Read, Write};
use std::path::{Path, PathBuf};
use std::time::{Duration, Instant};
use anyhow::{bail, Context, Result};
use rpi_loader_ota::{encode, Bundle, Entry, Format, Role};
use serde::Deserialize;
const DEFAULT_MAX_ENTRIES: usize = 32;
const UPLOAD_TIMEOUT: Duration = Duration::from_secs(300);
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct Manifest {
magic: String,
name: String,
max_entries: Option<usize>,
kernel: Option<KernelSpec>,
#[serde(default)]
files: Vec<FileSpec>,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct KernelSpec {
source: PathBuf,
dest: Option<String>,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct FileSpec {
source: PathBuf,
dest: Option<String>,
#[serde(default)]
role: FileRole,
}
#[derive(Deserialize, Clone, Copy, Default, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
enum FileRole {
#[default]
File,
Firmware,
Config,
Kernel,
}
struct Loaded {
role: Role,
path: String,
data: Vec<u8>,
}
pub fn run(
manifest: &Path,
output: Option<PathBuf>,
sdcard: Option<&Path>,
upload: Option<&str>,
) -> Result<()> {
let text = fs::read_to_string(manifest)
.with_context(|| format!("reading the manifest {}", manifest.display()))?;
let parsed: Manifest = toml::from_str(&text)
.with_context(|| format!("parsing the manifest {}", manifest.display()))?;
let root = manifest.parent().unwrap_or(Path::new("."));
let format = Format {
magic: parse_magic(&parsed.magic)?,
max_entries: parsed.max_entries.unwrap_or(DEFAULT_MAX_ENTRIES),
};
let loaded = load(&parsed, root)?;
let entries: Vec<Entry<'_>> = loaded
.iter()
.map(|file| Entry {
role: file.role,
path: &file.path,
data: &file.data,
})
.collect();
let bytes = encode(&format, &entries)?;
let out = output.unwrap_or_else(|| root.join("target").join(format!("{}.bundle", parsed.name)));
if let Some(parent) = out.parent() {
fs::create_dir_all(parent).with_context(|| format!("creating {}", parent.display()))?;
}
fs::write(&out, &bytes).with_context(|| format!("writing {}", out.display()))?;
report(&out, &bytes, &loaded);
if let Some(directory) = sdcard {
unpack(&format, &bytes, directory)?;
}
match upload {
Some(url) => post(url, &bytes),
None => Ok(()),
}
}
fn unpack(format: &Format, bytes: &[u8], directory: &Path) -> Result<()> {
if !directory.is_dir() {
bail!(
"{} is not a directory — is the card mounted?",
directory.display()
);
}
let bundle = Bundle::parse(format, bytes)?;
println!(
"\nwriting {} entries to {}",
bundle.count(),
directory.display()
);
for entry in bundle.iter() {
let mut destination = directory.to_path_buf();
for component in entry.path.split('/') {
destination.push(component);
}
if let Some(parent) = destination.parent() {
fs::create_dir_all(parent).with_context(|| format!("creating {}", parent.display()))?;
}
let mut file = fs::File::create(&destination)
.with_context(|| format!("creating {}", destination.display()))?;
file.write_all(entry.data)
.with_context(|| format!("writing {}", destination.display()))?;
file.sync_all()
.with_context(|| format!("flushing {}", destination.display()))?;
println!(" {:>9} {}", entry.data.len(), entry.path);
}
Ok(())
}
fn load(manifest: &Manifest, root: &Path) -> Result<Vec<Loaded>> {
let mut loaded = Vec::new();
if let Some(kernel) = &manifest.kernel {
let source = root.join(&kernel.source);
let path = match &kernel.dest {
Some(dest) => dest.clone(),
None => file_name(&source)?,
};
loaded.push(Loaded {
role: Role::Kernel,
path,
data: read(&source)?,
});
}
for spec in &manifest.files {
let role = match spec.role {
FileRole::File => Role::File,
FileRole::Firmware => Role::Firmware,
FileRole::Config => Role::Config,
FileRole::Kernel => bail!(
"{}: role = \"kernel\" is not allowed in [[files]]; \
name the boot image in the [kernel] table instead",
spec.source.display()
),
};
let source = root.join(&spec.source);
let dest = match &spec.dest {
Some(dest) => dest.clone(),
None => file_name(&source)?,
};
if source.is_dir() {
for (relative, data) in walk(&source)? {
loaded.push(Loaded {
role,
path: format!("{dest}/{relative}"),
data,
});
}
} else {
loaded.push(Loaded {
role,
path: dest,
data: read(&source)?,
});
}
}
Ok(loaded)
}
fn walk(directory: &Path) -> Result<BTreeMap<String, Vec<u8>>> {
let mut found = BTreeMap::new();
collect(directory, "", &mut found)?;
Ok(found)
}
fn collect(directory: &Path, prefix: &str, into: &mut BTreeMap<String, Vec<u8>>) -> Result<()> {
let listing =
fs::read_dir(directory).with_context(|| format!("reading {}", directory.display()))?;
for entry in listing {
let entry = entry.with_context(|| format!("reading {}", directory.display()))?;
let name = entry.file_name().to_string_lossy().into_owned();
if name.starts_with('.') {
continue;
}
let relative = if prefix.is_empty() {
name
} else {
format!("{prefix}/{name}")
};
let path = entry.path();
if path.is_dir() {
collect(&path, &relative, into)?;
} else {
into.insert(relative, read(&path)?);
}
}
Ok(())
}
fn read(path: &Path) -> Result<Vec<u8>> {
fs::read(path).with_context(|| format!("reading {}", path.display()))
}
fn file_name(path: &Path) -> Result<String> {
path.file_name()
.map(|name| name.to_string_lossy().into_owned())
.with_context(|| {
format!(
"{} has no file name to use as a destination",
path.display()
)
})
}
fn parse_magic(magic: &str) -> Result<[u8; 4]> {
let bytes = magic.as_bytes();
if bytes.len() != 4 || !magic.is_ascii() {
bail!("magic must be exactly four ASCII characters, not {magic:?}");
}
Ok([bytes[0], bytes[1], bytes[2], bytes[3]])
}
fn report(out: &Path, bytes: &[u8], loaded: &[Loaded]) {
println!(
"{} ({} bytes, {} entries)",
out.display(),
bytes.len(),
loaded.len()
);
let width = loaded.iter().map(|file| file.path.len()).max().unwrap_or(0);
for file in loaded {
let role = match file.role {
Role::File => "",
Role::Kernel => " kernel",
Role::Firmware => " firmware",
Role::Config => " config",
};
println!(
" {:<width$} {:>9}{role}",
file.path,
file.data.len(),
width = width
);
}
}
fn post(url: &str, bytes: &[u8]) -> Result<()> {
println!("\nuploading {} bytes to {url}", bytes.len());
let agent = ureq::Agent::config_builder()
.timeout_global(Some(UPLOAD_TIMEOUT))
.http_status_as_error(false)
.build()
.new_agent();
let started = Instant::now();
let mut progress = Progress::new(bytes, started);
let sent = agent
.post(url)
.content_type("application/octet-stream")
.header("Content-Length", bytes.len().to_string())
.send(ureq::SendBody::from_reader(&mut progress));
progress.finish();
let mut response = sent.with_context(|| format!("posting to {url}"))?;
let status = response.status();
let mut body = String::new();
response
.body_mut()
.as_reader()
.read_to_string(&mut body)
.with_context(|| format!("reading the reply from {url}"))?;
println!(
"{status}, {:.1}s round trip",
started.elapsed().as_secs_f64()
);
match serde_json::from_str::<serde_json::Value>(&body) {
Ok(json) => println!("{}", serde_json::to_string_pretty(&json)?),
Err(_) => println!("{}", body.trim_end()),
}
if !status.is_success() {
bail!("the board rejected the bundle ({status})");
}
Ok(())
}
const PROGRESS_INTERVAL: Duration = Duration::from_millis(100);
struct Progress<'a> {
bytes: &'a [u8],
sent: usize,
started: Instant,
drawn: Option<Instant>,
terminal: bool,
}
impl<'a> Progress<'a> {
fn new(bytes: &'a [u8], started: Instant) -> Self {
Self {
bytes,
sent: 0,
started,
drawn: None,
terminal: std::io::IsTerminal::is_terminal(&std::io::stdout()),
}
}
fn draw(&mut self, force: bool) {
if !self.terminal {
return;
}
let now = Instant::now();
if !force && self.drawn.is_some_and(|at| now - at < PROGRESS_INTERVAL) {
return;
}
self.drawn = Some(now);
let total = self.bytes.len();
let percent = (self.sent * 100).checked_div(total).unwrap_or(100);
let elapsed = (now - self.started).as_secs_f64();
let rate = if elapsed > 0.0 {
self.sent as f64 / 1024.0 / elapsed
} else {
0.0
};
let mut out = std::io::stdout();
let _ = write!(
out,
"\r {:.1} / {:.1} MB {percent:>3}% {rate:.0} KiB/s ",
self.sent as f64 / 1e6,
total as f64 / 1e6,
);
let _ = out.flush();
}
fn finish(&mut self) {
if !self.terminal {
return;
}
self.draw(true);
if self.sent == self.bytes.len() {
println!("\n sent; waiting for the board to install it");
} else {
println!();
}
}
}
impl Read for Progress<'_> {
fn read(&mut self, buf: &mut [u8]) -> std::io::Result<usize> {
let n = (&self.bytes[self.sent..]).read(buf)?;
self.sent += n;
self.draw(false);
Ok(n)
}
}