use std::collections::{BTreeSet, HashMap};
use std::io::{Read, Write};
use std::net::{TcpListener, TcpStream};
use std::path::PathBuf;
use std::process::ExitCode;
use std::sync::mpsc::{self, RecvTimeoutError};
use std::time::{Duration, Instant, SystemTime, UNIX_EPOCH};
use ch32rv_contract::{ErrorKind, ResultEnvelope};
use ch32rv_oep::codec::{
LengthDeframer, Request, Resolution, encode_result, length_frame, parse_tlvs,
};
use ch32rv_oep::link::{Call, Link, Reply};
use ch32rv_oep::registry::{self, core as oep_core, outcomes, reject_reasons, wire_rvswd};
use ch32rv_oep::session::Probe;
use serde_json::{Value, json};
use crate::args::Cli;
use crate::cmd_probe::fail;
use crate::oep::OepAddr;
const FIRST_CLIENT_WAIT: Duration = Duration::from_secs(10);
const START_WAIT: Duration = Duration::from_secs(5);
const LEASE_MS: u32 = 3000;
const KEEPALIVE_EVERY: Duration = Duration::from_millis(1000);
fn key_for(path: &str) -> String {
ch32rv_usb::sanitize_key(&format!("oep-{}", ch32rv_usb::normalize_port(path)))
}
fn endpoint_file(key: &str) -> PathBuf {
ch32rv_usb::runtime_dir().join(format!("{key}.oep"))
}
fn now_ms() -> u64 {
SystemTime::now()
.duration_since(UNIX_EPOCH)
.map(|d| d.as_millis() as u64)
.unwrap_or(0)
}
fn read_endpoint(key: &str) -> Option<Value> {
let text = std::fs::read_to_string(endpoint_file(key)).ok()?;
serde_json::from_str(&text).ok()
}
fn write_endpoint(key: &str, v: &Value) -> std::io::Result<()> {
let dir = ch32rv_usb::runtime_dir();
std::fs::create_dir_all(&dir)?;
let tmp = dir.join(format!("{key}.oep.{}", std::process::id()));
{
let mut o = std::fs::OpenOptions::new();
o.write(true).create(true).truncate(true);
#[cfg(unix)]
{
use std::os::unix::fs::OpenOptionsExt;
o.mode(0o600);
}
let mut f = o.open(&tmp)?;
f.write_all(v.to_string().as_bytes())?;
}
std::fs::rename(&tmp, endpoint_file(key))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum BrokerTarget {
Serial(String),
Wch { id: String, selector: String },
}
impl BrokerTarget {
pub(crate) fn wch(entry: &crate::cmd_probe::Entry) -> Self {
match entry.dev.serial() {
Some(sn) => BrokerTarget::Wch {
id: sn.to_owned(),
selector: format!("serial:{sn}"),
},
None => {
let t = entry.dev.topology();
BrokerTarget::Wch {
id: t.clone(),
selector: format!("usb:{t}"),
}
}
}
}
fn key(&self) -> String {
match self {
BrokerTarget::Serial(p) => key_for(p),
BrokerTarget::Wch { id, .. } => ch32rv_usb::sanitize_key(&format!("wch-{id}")),
}
}
fn selector(&self) -> String {
match self {
BrokerTarget::Serial(p) => format!("port:{p}"),
BrokerTarget::Wch { selector, .. } => selector.clone(),
}
}
}
fn spawn_broker(target: &BrokerTarget) -> Result<(), String> {
let exe = std::env::current_exe().map_err(|e| format!("current exe: {e}"))?;
let mut cmd = std::process::Command::new(exe);
cmd.args(["broker", "serve", "--probe", &target.selector()])
.stdin(std::process::Stdio::null())
.stdout(std::process::Stdio::null())
.stderr(std::process::Stdio::null());
#[cfg(unix)]
{
use std::os::unix::process::CommandExt;
cmd.process_group(0);
}
#[cfg(windows)]
{
use std::os::windows::process::CommandExt;
const DETACHED_PROCESS: u32 = 0x0000_0008;
const CREATE_NEW_PROCESS_GROUP: u32 = 0x0000_0200;
const CREATE_BREAKAWAY_FROM_JOB: u32 = 0x0100_0000;
cmd.creation_flags(DETACHED_PROCESS | CREATE_NEW_PROCESS_GROUP | CREATE_BREAKAWAY_FROM_JOB);
}
cmd.spawn()
.map(|_| ())
.map_err(|e| format!("start the broker: {e}"))
}
pub(crate) fn client_link_for(target: &BrokerTarget) -> Result<Link, String> {
let key = target.key();
let try_connect = |v: &Value| -> Option<Link> {
let port = v.get("port")?.as_u64()?;
ch32rv_oep::link::open_tcp(&format!("127.0.0.1:{port}")).ok()
};
if let Some(v) = read_endpoint(&key)
&& let Some(l) = try_connect(&v)
{
return Ok(l);
}
let t0 = now_ms();
spawn_broker(target)?;
let deadline = Instant::now() + START_WAIT;
while Instant::now() < deadline {
if let Some(v) = read_endpoint(&key) {
let fresh = v.get("time").and_then(Value::as_u64).unwrap_or(0) + 50 >= t0;
if let Some(l) = try_connect(&v) {
return Ok(l);
}
if fresh && let Some(e) = v.get("error").and_then(Value::as_str) {
return Err(e.to_owned());
}
}
std::thread::sleep(Duration::from_millis(30));
}
Err(format!(
"the broker for {} did not start within {} s",
target.selector(),
START_WAIT.as_secs()
))
}
pub(crate) fn client_link(path: &str) -> Result<Link, String> {
client_link_for(&BrokerTarget::Serial(path.to_owned()))
}
pub(crate) fn existing_link_for(target: &BrokerTarget) -> Option<Link> {
let v = read_endpoint(&target.key())?;
let port = v.get("port")?.as_u64()?;
ch32rv_oep::link::open_tcp(&format!("127.0.0.1:{port}")).ok()
}
pub(crate) fn existing_link(path: &str) -> Option<Link> {
existing_link_for(&BrokerTarget::Serial(path.to_owned()))
}
pub(crate) fn listing_cache(path: &str) -> PathBuf {
ch32rv_usb::runtime_dir().join(format!("{}.slots.json", key_for(path)))
}
pub(crate) fn endpoint(cli: &Cli) -> ExitCode {
const CMD: &str = "broker.endpoint";
let target = match crate::oep::addr(cli, CMD) {
Ok(Some(OepAddr::Serial(p) | OepAddr::Slot { path: p, .. })) => BrokerTarget::Serial(p),
Ok(Some(OepAddr::Wch(t))) => t,
Ok(Some(OepAddr::Tcp(_))) => {
return fail(
cli,
CMD,
ErrorKind::Usage,
"a tcp: endpoint has no broker",
None,
);
}
Ok(None) => match crate::cmd_probe::select_entry(cli, CMD) {
Ok(e) => BrokerTarget::wch(&e),
Err(c) => return c,
},
Err(c) => return c,
};
let v = read_endpoint(&target.key());
let live = v.as_ref().and_then(|v| {
let port = v.get("port")?.as_u64()?;
TcpStream::connect(("127.0.0.1", port as u16)).ok()?;
Some(format!("127.0.0.1:{port}"))
});
if cli.json {
let mut env = ResultEnvelope::success(CMD);
env.result = Some(json!({
"endpoint": live,
"pid": live.as_ref().and(v.as_ref().and_then(|v| v.get("pid").cloned())),
"transport": live.as_ref().and(v.as_ref().and_then(|v| v.get("transport").cloned())),
}));
crate::print_envelope(&env)
} else {
println!(
"{}",
live.as_deref()
.unwrap_or("no broker is running for this probe")
);
ExitCode::SUCCESS
}
}
enum Ev {
Connected(u64, TcpStream),
Msg(u64, Vec<u8>),
Gone(u64),
}
#[derive(Default)]
struct Ledger {
conns: HashMap<u16, (u16, BTreeSet<u64>)>,
plans: HashMap<u64, BTreeSet<u16>>,
}
enum Pending {
Local(u64, Vec<u8>),
Forward(u64, Request),
}
pub(crate) fn serve(cli: &Cli) -> ExitCode {
const CMD: &str = "broker.serve";
let target = match crate::oep::addr(cli, CMD) {
Ok(Some(OepAddr::Serial(p) | OepAddr::Slot { path: p, .. })) => BrokerTarget::Serial(p),
Ok(Some(OepAddr::Wch(_))) => return ExitCode::SUCCESS,
Ok(Some(OepAddr::Tcp(_))) => {
return fail(
cli,
CMD,
ErrorKind::Usage,
"a tcp: endpoint needs no broker",
None,
);
}
Ok(None) => match crate::cmd_probe::select_entry(cli, CMD) {
Ok(e) => {
let t = BrokerTarget::wch(&e);
return serve_target(cli, t, Some(e));
}
Err(c) => return c,
},
Err(c) => return c,
};
serve_target(cli, target, None)
}
fn serve_target(
cli: &Cli,
target: BrokerTarget,
wch_entry: Option<crate::cmd_probe::Entry>,
) -> ExitCode {
let key = target.key();
let Ok(_guard) = ch32rv_usb::DeviceLock::acquire(&format!("{key}.broker"), Duration::ZERO)
else {
return ExitCode::SUCCESS;
};
let report_error = |m: String| {
let _ = write_endpoint(
&key,
&json!({"error": m, "pid": std::process::id(), "time": now_ms()}),
);
ExitCode::from(ErrorKind::DeviceOpenFailed.exit_code())
};
let mut transport = "wchlink";
let (up, sid) = match (&target, wch_entry) {
(BrokerTarget::Serial(path), _) => {
let mut probe = match crate::oep::connect_upstream(path, None) {
Ok((p, t)) => {
transport = t;
p
}
Err(m) => return report_error(format!("{path}: {m}")),
};
let serial = crate::oep::single_serial(&mut probe);
let owner = format!("ch32rv broker pid {}", std::process::id());
let sid = match crate::oep::open_with_lock_rule(&mut probe, serial, &owner, LEASE_MS) {
Ok(s) => s,
Err(e) => return report_error(e.to_string()),
};
(Upstream::Oep(Box::new(probe)), sid)
}
(BrokerTarget::Wch { .. }, Some(entry)) => (
Upstream::Wch(Box::new(crate::broker_wch::WchUpstream::new(
entry,
Duration::from_secs(cli.lock_timeout),
))),
0,
),
(BrokerTarget::Wch { .. }, None) => return report_error("no WCH-Link entry".to_owned()),
};
let mut up = up;
let wires: BTreeSet<u16> = match up.wires() {
Ok(w) => w,
Err(e) => return report_error(e),
};
let listener = match TcpListener::bind(("127.0.0.1", 0)) {
Ok(l) => l,
Err(e) => return report_error(format!("listen: {e}")),
};
let port = listener.local_addr().map(|a| a.port()).unwrap_or(0);
if write_endpoint(
&key,
&json!({"port": port, "pid": std::process::id(), "time": now_ms(), "transport": transport}),
)
.is_err()
{
return ExitCode::from(ErrorKind::DeviceOpenFailed.exit_code());
}
let (tx, rx) = mpsc::channel::<Ev>();
std::thread::spawn(move || accept_loop(listener, tx));
let mut b = Broker {
up,
sid,
wires,
clients: HashMap::new(),
ledger: Ledger::default(),
};
let r = b.run(&rx);
if read_endpoint(&key).and_then(|v| v.get("pid")?.as_u64())
== Some(u64::from(std::process::id()))
{
let _ = std::fs::remove_file(endpoint_file(&key));
}
b.up.end();
match r {
Ok(()) => ExitCode::SUCCESS,
Err(_) => ExitCode::from(ErrorKind::TransferFailed.exit_code()),
}
}
enum Upstream {
Oep(Box<Probe>),
Wch(Box<crate::broker_wch::WchUpstream>),
}
impl Upstream {
fn limits(&self) -> ch32rv_oep::link::Limits {
match self {
Upstream::Oep(p) => p.limits(),
Upstream::Wch(w) => w.limits(),
}
}
fn boot_id(&self) -> u32 {
match self {
Upstream::Oep(p) => p.boot_id().unwrap_or(0),
Upstream::Wch(_) => 0,
}
}
fn wires(&mut self) -> Result<BTreeSet<u16>, String> {
match self {
Upstream::Oep(p) => p
.list("oep.wire")
.map(|l| l.into_iter().map(|i| i.func).collect())
.map_err(|e| e.to_string()),
Upstream::Wch(w) => Ok(w.wires().into_iter().collect()),
}
}
fn exchange(&mut self, calls: Vec<(u64, Call)>) -> Result<Vec<Reply>, String> {
match self {
Upstream::Oep(p) => p
.link()
.exchange(calls.into_iter().map(|(_, c)| c).collect())
.map_err(|e| e.to_string()),
Upstream::Wch(w) => Ok(calls.iter().map(|(id, c)| w.handle(*id, c)).collect()),
}
}
fn client_gone(&mut self, id: u64) {
if let Upstream::Wch(w) = self {
w.client_gone(id);
}
}
fn keepalive(&mut self) -> Result<(), String> {
match self {
Upstream::Oep(p) => p.keepalive().map_err(|e| e.to_string()),
Upstream::Wch(_) => Ok(()),
}
}
fn tick(&mut self) -> Duration {
match self {
Upstream::Oep(_) => Duration::from_millis(250),
Upstream::Wch(w) => w.tick(),
}
}
fn end(&mut self) {
if let Upstream::Oep(p) = self {
let _ = p.end();
}
}
}
fn accept_loop(listener: TcpListener, tx: mpsc::Sender<Ev>) {
for (id, stream) in listener.incoming().flatten().enumerate() {
let id = id as u64 + 1;
let _ = stream.set_nodelay(true);
let Ok(reader) = stream.try_clone() else {
continue;
};
if tx.send(Ev::Connected(id, stream)).is_err() {
return;
}
let tx = tx.clone();
std::thread::spawn(move || read_loop(id, reader, tx));
}
}
fn read_loop(id: u64, mut s: TcpStream, tx: mpsc::Sender<Ev>) {
let mut d = LengthDeframer::new(0xFFFF);
let mut buf = [0u8; 8192];
loop {
match s.read(&mut buf) {
Ok(0) | Err(_) => break,
Ok(n) => match d.push(&buf[..n]) {
Ok(msgs) => {
for m in msgs {
if tx.send(Ev::Msg(id, m)).is_err() {
return;
}
}
}
Err(_) => break,
},
}
}
let _ = tx.send(Ev::Gone(id));
}
struct Broker {
up: Upstream,
sid: u32,
wires: BTreeSet<u16>,
clients: HashMap<u64, TcpStream>,
ledger: Ledger,
}
impl Broker {
fn run(&mut self, rx: &mpsc::Receiver<Ev>) -> Result<(), String> {
let started = Instant::now();
let mut had_client = false;
let mut last_upstream = Instant::now();
loop {
let wait = self.up.tick();
let ev = match rx.recv_timeout(wait) {
Ok(ev) => ev,
Err(RecvTimeoutError::Timeout) => {
if self.clients.is_empty()
&& (had_client || started.elapsed() > FIRST_CLIENT_WAIT)
{
return Ok(());
}
if last_upstream.elapsed() >= KEEPALIVE_EVERY {
self.up.keepalive()?;
last_upstream = Instant::now();
}
continue;
}
Err(RecvTimeoutError::Disconnected) => return Ok(()),
};
let mut batch = vec![ev];
while let Ok(ev) = rx.try_recv() {
batch.push(ev);
}
let mut msgs = Vec::new();
for ev in batch {
match ev {
Ev::Connected(id, s) => {
self.clients.insert(id, s);
}
Ev::Msg(id, m) => {
had_client = true;
msgs.push((id, m));
}
Ev::Gone(id) => {
self.serve_batch(std::mem::take(&mut msgs))?;
self.release(id)?;
self.up.client_gone(id);
self.clients.remove(&id);
}
}
}
if !msgs.is_empty() {
self.serve_batch(msgs)?;
last_upstream = Instant::now();
}
if had_client && self.clients.is_empty() {
return Ok(());
}
}
}
fn serve_batch(&mut self, msgs: Vec<(u64, Vec<u8>)>) -> Result<(), String> {
let mut pending = Vec::with_capacity(msgs.len());
let mut calls = Vec::new();
for (id, m) in msgs {
let Some(req) = Request::decode(&m) else {
continue; };
if let Some(answer) = self.local(id, &req) {
pending.push(Pending::Local(id, answer));
continue;
}
calls.push((
id,
Call {
func: req.func,
op: req.op,
session: req.session.map(|_| self.sid),
payload: req.payload.clone(),
},
));
pending.push(Pending::Forward(id, req));
}
let mut replies = if calls.is_empty() {
Vec::new()
} else {
self.up.exchange(calls)?
}
.into_iter();
for p in pending {
match p {
Pending::Local(id, answer) => self.send(id, &answer),
Pending::Forward(id, req) => {
let Some(r) = replies.next() else { break };
self.note(id, &req, &r);
self.send(id, &encode_result(req.corr, r.resolution, &r.payload));
}
}
}
Ok(())
}
fn send(&mut self, id: u64, msg: &[u8]) {
if let Some(s) = self.clients.get_mut(&id) {
let _ = s.write_all(&length_frame(msg));
}
}
fn local(&mut self, id: u64, req: &Request) -> Option<Vec<u8>> {
let ok = |p: &[u8]| encode_result(req.corr, Resolution::Completed(outcomes::SUCCESS), p);
if req.func == oep_core::FN {
let op = req.op;
return match op {
o if o == oep_core::op::CONFIRM => {
let l = self.up.limits();
let mut p = registry::constants::CONFIRM_RESULT_MAGIC
.as_bytes()
.to_vec();
p.extend_from_slice(&[l.revision, 0]);
p.extend_from_slice(&l.max_frame.to_le_bytes());
p.extend_from_slice(&l.window.to_le_bytes());
p.push(l.max_inflight);
Some(ok(&p))
}
o if o == oep_core::op::OPEN => {
let lease = req
.payload
.get(4..8)
.map(|b| u32::from_le_bytes([b[0], b[1], b[2], b[3]]))
.filter(|&l| l != 0)
.unwrap_or(LEASE_MS)
.clamp(1000, 60_000);
let mut p = lease.to_le_bytes().to_vec();
p.extend_from_slice(&self.up.boot_id().to_le_bytes());
p.push(0);
Some(ok(&p))
}
o if o == oep_core::op::END || o == oep_core::op::KEEPALIVE => Some(ok(&[])),
o if o == oep_core::op::LOCK_STATE => {
let mut p = vec![1u8];
p.extend_from_slice(&LEASE_MS.to_le_bytes());
Some(ok(&p))
}
o if o == oep_core::op::SUBSCRIBE || o == oep_core::op::UNSUBSCRIBE => {
Some(encode_result(
req.corr,
Resolution::Rejected(reject_reasons::UNAVAILABLE),
&[],
))
}
_ => None,
};
}
if self.wires.contains(&req.func) && req.op == wire_rvswd::op::DETACH {
let conn = u16::from_le_bytes([*req.payload.first()?, *req.payload.get(1)?]);
let (_, users) = self.ledger.conns.get_mut(&conn)?;
if users.len() > 1 || !users.contains(&id) {
users.remove(&id);
return Some(ok(&[]));
}
}
None
}
fn note(&mut self, id: u64, req: &Request, r: &Reply) {
if !r.succeeded() {
return;
}
if self.wires.contains(&req.func) {
match req.op {
o if o == wire_rvswd::op::ATTACH || o == wire_rvswd::op::ATTACH_UNDER_RESET => {
if r.payload.len() >= 2 {
let conn = u16::from_le_bytes([r.payload[0], r.payload[1]]);
self.ledger
.conns
.entry(conn)
.or_insert((req.func, BTreeSet::new()))
.1
.insert(id);
}
}
o if o == wire_rvswd::op::DETACH && req.payload.len() >= 2 => {
let conn = u16::from_le_bytes([req.payload[0], req.payload[1]]);
self.ledger.conns.remove(&conn);
}
_ => {}
}
} else if req.func == oep_core::FN && req.op == oep_core::op::PLAN_APPLY {
if let Ok(tlvs) = parse_tlvs(&req.payload) {
let fns = self.ledger.plans.entry(id).or_default();
for t in tlvs.iter().filter(|t| t.value.len() >= 2) {
fns.insert(u16::from_le_bytes([t.value[0], t.value[1]]));
}
}
} else if req.func == oep_core::FN && req.op == oep_core::op::PLAN_RELEASE {
let fns = self.ledger.plans.entry(id).or_default();
match req.payload.first() {
Some(0) | None => fns.clear(),
Some(&n) => {
for i in 0..usize::from(n) {
if let Some(b) = req.payload.get(1 + 2 * i..3 + 2 * i) {
fns.remove(&u16::from_le_bytes([b[0], b[1]]));
}
}
}
}
}
}
fn release(&mut self, id: u64) -> Result<(), String> {
let mut calls = Vec::new();
let mut dead = Vec::new();
for (&conn, (func, users)) in self.ledger.conns.iter_mut() {
if users.remove(&id) && users.is_empty() {
calls.push((
id,
Call {
func: *func,
op: wire_rvswd::op::DETACH,
session: Some(self.sid),
payload: conn.to_le_bytes().to_vec(),
},
));
dead.push(conn);
}
}
for c in dead {
self.ledger.conns.remove(&c);
}
if let Some(fns) = self.ledger.plans.remove(&id)
&& !fns.is_empty()
{
let mut p = vec![fns.len() as u8];
for f in fns {
p.extend_from_slice(&f.to_le_bytes());
}
calls.push((
id,
Call {
func: oep_core::FN,
op: oep_core::op::PLAN_RELEASE,
session: Some(self.sid),
payload: p,
},
));
}
if !calls.is_empty() {
self.up.exchange(calls)?;
}
Ok(())
}
}
pub(crate) struct Lend {
probe: Probe,
func: u16,
}
pub(crate) fn lend(target: &BrokerTarget) -> Result<Lend, String> {
let link = existing_link_for(target).ok_or("the Link's broker is not running")?;
let mut probe = Probe::connect(link).map_err(|e| e.to_string())?;
let owner = format!("ch32rv flash pid {}", std::process::id());
probe
.open(
ch32rv_oep::session::random_session_id(),
60_000,
false,
Some(&owner),
)
.map_err(|e| e.to_string())?;
let func = probe
.interface(crate::broker_wch::WCHLINK_IF)
.map_err(|e| e.to_string())?
.func;
let r = probe
.call(func, crate::broker_wch::OP_LEND, Vec::new())
.map_err(|e| e.to_string())?;
ch32rv_oep::session::check(r)
.map_err(|e| format!("the broker would not lend the Link: {e}"))?;
Ok(Lend { probe, func })
}
impl Lend {
fn hand_back_with_reset(mut self) -> Option<bool> {
let r = self
.probe
.call(self.func, crate::broker_wch::OP_RECLAIM, vec![1])
.ok()?;
let running = ch32rv_oep::session::check(r)
.ok()
.and_then(|p| p.first().map(|&b| b != 0));
let _ = self.probe.end();
self.func = 0; running
}
}
impl Drop for Lend {
fn drop(&mut self) {
if self.func == 0 {
return;
}
let _ = self
.probe
.call(self.func, crate::broker_wch::OP_RECLAIM, Vec::new());
let _ = self.probe.end();
}
}
thread_local! {
static LENT: std::cell::RefCell<Option<Lend>> = const { std::cell::RefCell::new(None) };
}
pub(crate) struct LendScope;
impl LendScope {
pub(crate) fn new(l: Lend) -> Self {
LENT.with(|c| *c.borrow_mut() = Some(l));
LendScope
}
}
impl Drop for LendScope {
fn drop(&mut self) {
LENT.with(|c| drop(c.borrow_mut().take()));
}
}
pub(crate) fn lend_active() -> bool {
LENT.with(|c| c.borrow().is_some())
}
pub(crate) fn hand_back_with_reset() -> Option<bool> {
LENT.with(|c| c.borrow_mut().take())?.hand_back_with_reset()
}