use std::io::{self, BufRead, BufReader};
use std::ops::ControlFlow;
use std::os::unix::net::UnixStream;
use std::path::Path;
use std::time::Duration;
use serde::Deserialize;
use crate::mxc::wire::Message;
use crate::mxc::PROTOCOL_VERSION;
const KNOWN_TAGS: &[&str] = &["theme", "bye"];
const BACKOFF_START: Duration = Duration::from_millis(100);
const BACKOFF_CAP: Duration = Duration::from_secs(5);
#[derive(Deserialize)]
struct Envelope {
t: String,
#[serde(default)]
v: u32,
}
#[derive(Debug)]
pub struct Subscriber {
reader: BufReader<UnixStream>,
line: String,
poisoned: bool,
}
impl Subscriber {
pub fn connect(path: &Path) -> io::Result<Subscriber> {
let stream = UnixStream::connect(path)?;
Ok(Subscriber {
reader: BufReader::new(stream),
line: String::new(),
poisoned: false,
})
}
pub fn next_message(&mut self) -> io::Result<Option<Message>> {
if self.poisoned {
return Ok(None);
}
loop {
self.line.clear();
if self.reader.read_line(&mut self.line)? == 0 {
return Ok(None);
}
let raw = self.line.trim();
if raw.is_empty() {
continue;
}
let Ok(env) = serde_json::from_str::<Envelope>(raw) else {
continue; };
if env.v > PROTOCOL_VERSION {
self.poisoned = true;
return Err(io::Error::new(
io::ErrorKind::InvalidData,
format!(
"MXC protocol version {} is newer than this client's v{PROTOCOL_VERSION}; \
refusing to guess at message {:?}",
env.v, env.t
),
));
}
if !KNOWN_TAGS.contains(&env.t.as_str()) {
continue; }
match serde_json::from_str::<Message>(raw) {
Ok(msg) => return Ok(Some(msg)),
Err(_) => continue,
}
}
}
}
impl Iterator for Subscriber {
type Item = io::Result<Message>;
fn next(&mut self) -> Option<Self::Item> {
match self.next_message() {
Ok(Some(msg)) => Some(Ok(msg)),
Ok(None) => None,
Err(e) => Some(Err(e)),
}
}
}
pub fn watch<F>(path: &Path, mut on_message: F) -> io::Result<()>
where
F: FnMut(Message) -> ControlFlow<()>,
{
let mut backoff = BACKOFF_START;
loop {
if let Ok(mut sub) = Subscriber::connect(path) {
backoff = BACKOFF_START;
while let Ok(Some(msg)) = sub.next_message() {
if on_message(msg).is_break() {
return Ok(());
}
}
}
std::thread::sleep(backoff);
backoff = (backoff * 2).min(BACKOFF_CAP);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::gradient::Rgb;
use crate::mxc::contrast::Contrast;
use crate::mxc::wire::{ByeEvent, ByeReason, Colors, Hex, Origin, OriginKind, ThemeEvent};
use std::io::Write;
use std::os::unix::net::UnixListener;
use std::path::PathBuf;
use std::sync::atomic::{AtomicU32, Ordering};
use std::sync::mpsc;
use std::thread::JoinHandle;
const GENEROUS: Duration = Duration::from_secs(10);
static COUNTER: AtomicU32 = AtomicU32::new(0);
fn temp_sock() -> PathBuf {
let n = COUNTER.fetch_add(1, Ordering::Relaxed);
let pid = std::process::id();
std::env::temp_dir().join(format!("mxcs{pid}-{n}.sock"))
}
struct Fixture {
path: PathBuf,
server: Option<JoinHandle<()>>,
}
impl Fixture {
fn spawn<F>(conns: usize, serve: F) -> Fixture
where
F: Fn(usize, &mut UnixStream) + Send + 'static,
{
let path = temp_sock();
let _ = std::fs::remove_file(&path);
let listener = UnixListener::bind(&path).expect("bind test socket");
let server = std::thread::spawn(move || {
for i in 0..conns {
match listener.accept() {
Ok((mut stream, _)) => {
serve(i, &mut stream);
let _ = stream.flush();
}
Err(_) => return,
}
}
});
Fixture {
path,
server: Some(server),
}
}
fn path(&self) -> &Path {
&self.path
}
}
impl Drop for Fixture {
fn drop(&mut self) {
let _ = std::fs::remove_file(&self.path);
if let Some(h) = self.server.take() {
if h.is_finished() {
let _ = h.join();
}
}
}
}
fn sample_colors() -> Colors {
Colors {
primary: Hex(Rgb::new(0x64, 0xe0, 0xd0)),
secondary: Hex(Rgb::new(0x4a, 0x9f, 0xd8)),
accent: Hex(Rgb::new(0xf4, 0xaa, 0x48)),
error: Hex(Rgb::new(0xe0, 0x55, 0x61)),
warning: Hex(Rgb::new(0xd9, 0xa4, 0x41)),
success: Hex(Rgb::new(0x61, 0xc7, 0x66)),
info: Hex(Rgb::new(0x64, 0xe0, 0xd0)),
text: Hex(Rgb::new(0xd8, 0xef, 0xff)),
text_muted: Hex(Rgb::new(0x7a, 0x90, 0xa4)),
background: Hex(Rgb::new(0x08, 0x10, 0x18)),
background_panel: Hex(Rgb::new(0x10, 0x1d, 0x2a)),
background_element: Hex(Rgb::new(0x18, 0x29, 0x3a)),
border: Hex(Rgb::new(0x22, 0x37, 0x4a)),
border_active: Hex(Rgb::new(0x42, 0xd9, 0xd0)),
border_subtle: Hex(Rgb::new(0x18, 0x28, 0x38)),
border_dimmest: Hex(Rgb::new(0x10, 0x1c, 0x28)),
}
}
fn theme_msg(seq: u64) -> Message {
Message::Theme(ThemeEvent {
v: PROTOCOL_VERSION,
seq,
ts: 1_785_616_484_123,
origin: Origin::named(OriginKind::AlbumArt, "Blue Monday"),
fade_ms: 600,
is_dark: true,
colors: sample_colors(),
contrast: Contrast::compute(&sample_colors()),
})
}
fn theme_line(seq: u64) -> String {
theme_msg(seq).to_ndjson().unwrap()
}
fn seq_of(msg: &Message) -> u64 {
match msg {
Message::Theme(t) => t.seq,
Message::Bye(b) => b.seq,
}
}
#[test]
fn known_tags_covers_every_message_variant() {
let variants = [
theme_msg(0),
Message::Bye(ByeEvent {
v: PROTOCOL_VERSION,
seq: 1,
ts: 1,
reason: ByeReason::Shutdown,
}),
];
for m in variants {
let v: serde_json::Value = serde_json::from_str(&m.to_ndjson().unwrap()).unwrap();
let tag = v["t"].as_str().unwrap().to_string();
assert!(
KNOWN_TAGS.contains(&tag.as_str()),
"{tag} missing from KNOWN_TAGS"
);
}
}
#[test]
fn well_formed_theme_line_parses() {
let fx = Fixture::spawn(1, |_, s| {
let _ = s.write_all(theme_line(7).as_bytes());
});
let mut sub = Subscriber::connect(fx.path()).unwrap();
let msg = sub.next_message().unwrap().expect("a message");
match msg {
Message::Theme(t) => {
assert_eq!(t.seq, 7);
assert_eq!(t.colors.primary.to_string(), "#64e0d0");
assert_eq!(t.origin.kind, OriginKind::AlbumArt);
}
other => panic!("expected theme, got {other:?}"),
}
}
#[test]
fn message_split_across_two_writes_still_parses() {
let fx = Fixture::spawn(1, |_, s| {
let line = theme_line(42);
let half = line.len() / 2;
let _ = s.write_all(&line.as_bytes()[..half]);
let _ = s.flush();
std::thread::sleep(Duration::from_millis(150));
let _ = s.write_all(&line.as_bytes()[half..]);
});
let mut sub = Subscriber::connect(fx.path()).unwrap();
let msg = sub.next_message().unwrap().expect("a message");
assert_eq!(seq_of(&msg), 42, "split message must reassemble intact");
}
#[test]
fn unknown_tag_is_skipped_and_next_message_still_arrives() {
let fx = Fixture::spawn(1, |_, s| {
let _ =
s.write_all(b"{\"t\":\"future_thing\",\"v\":1,\"seq\":1,\"ts\":1,\"wat\":[1,2]}\n");
let _ = s.write_all(theme_line(2).as_bytes());
});
let mut sub = Subscriber::connect(fx.path()).unwrap();
let msg = sub.next_message().unwrap().expect("a message");
assert_eq!(seq_of(&msg), 2, "the future-tagged line must be skipped");
}
#[test]
fn unknown_fields_do_not_break_a_theme() {
let fx = Fixture::spawn(1, |_, s| {
let mut v: serde_json::Value = serde_json::from_str(theme_line(3).trim()).unwrap();
v["invented_in_v2"] = serde_json::json!({"nested": true});
let _ = s.write_all(format!("{v}\n").as_bytes());
});
let mut sub = Subscriber::connect(fx.path()).unwrap();
let msg = sub.next_message().unwrap().expect("a message");
assert_eq!(seq_of(&msg), 3);
}
#[test]
fn newer_protocol_version_errors_instead_of_being_misread() {
let fx = Fixture::spawn(1, |_, s| {
let mut v: serde_json::Value = serde_json::from_str(theme_line(4).trim()).unwrap();
v["v"] = serde_json::json!(PROTOCOL_VERSION + 1);
let _ = s.write_all(format!("{v}\n").as_bytes());
});
let mut sub = Subscriber::connect(fx.path()).unwrap();
let err = sub
.next_message()
.expect_err("version skew must be an error");
assert_eq!(err.kind(), io::ErrorKind::InvalidData);
assert!(
err.to_string().contains("newer than this client"),
"error must be legible, got: {err}"
);
assert!(sub.next_message().unwrap().is_none());
}
#[test]
fn blank_and_malformed_lines_are_skipped() {
let fx = Fixture::spawn(1, |_, s| {
let _ = s.write_all(b"\n");
let _ = s.write_all(b" \n");
let _ = s.write_all(b"{not json at all,,,\n");
let _ = s.write_all(b"[1,2,3]\n");
let _ = s.write_all(b"{\"t\":\"theme\",\"v\":1}\n"); let _ = s.write_all(theme_line(5).as_bytes());
});
let mut sub = Subscriber::connect(fx.path()).unwrap();
let msg = sub.next_message().unwrap().expect("a message");
assert_eq!(seq_of(&msg), 5);
}
#[test]
fn clean_eof_yields_none_and_ends_the_iterator() {
let fx = Fixture::spawn(1, |_, s| {
let _ = s.write_all(theme_line(1).as_bytes());
});
let mut sub = Subscriber::connect(fx.path()).unwrap();
assert_eq!(seq_of(&sub.next_message().unwrap().unwrap()), 1);
assert!(
sub.next_message().unwrap().is_none(),
"clean EOF is Ok(None)"
);
let fx2 = Fixture::spawn(1, |_, s| {
let _ = s.write_all(theme_line(1).as_bytes());
});
let collected: Vec<_> = Subscriber::connect(fx2.path())
.unwrap()
.collect::<io::Result<Vec<_>>>()
.unwrap();
assert_eq!(collected.len(), 1, "iterator terminates at EOF");
}
#[test]
fn connect_fails_immediately_when_nobody_is_listening() {
let err = Subscriber::connect(&temp_sock()).expect_err("must not block");
assert_eq!(err.kind(), io::ErrorKind::NotFound);
}
#[test]
fn watch_reconnects_after_the_publisher_drops_the_connection() {
let fx = Fixture::spawn(2, |i, s| {
let seq = if i == 0 { 100 } else { 200 };
let _ = s.write_all(theme_line(seq).as_bytes());
});
let path = fx.path().to_path_buf();
let (tx, rx) = mpsc::channel();
let worker = std::thread::spawn(move || {
let mut seen = 0usize;
watch(&path, |msg| {
let _ = tx.send(seq_of(&msg));
seen += 1;
if seen == 2 {
ControlFlow::Break(())
} else {
ControlFlow::Continue(())
}
})
});
let first = rx.recv_timeout(GENEROUS).expect("first message");
let second = rx.recv_timeout(GENEROUS).expect("message after reconnect");
assert_eq!((first, second), (100, 200));
worker.join().expect("watch thread").expect("watch ok");
}
#[test]
fn watch_waits_for_a_publisher_that_is_not_up_yet() {
let path = temp_sock();
let watch_path = path.clone();
let (tx, rx) = mpsc::channel();
let worker = std::thread::spawn(move || {
watch(&watch_path, |msg| {
let _ = tx.send(seq_of(&msg));
ControlFlow::Break(())
})
});
std::thread::sleep(Duration::from_millis(250));
let listener = UnixListener::bind(&path).expect("late bind");
let server = std::thread::spawn(move || {
if let Ok((mut s, _)) = listener.accept() {
let _ = s.write_all(theme_line(9).as_bytes());
let _ = s.flush();
}
});
assert_eq!(rx.recv_timeout(GENEROUS).expect("late publisher"), 9);
worker.join().expect("watch thread").expect("watch ok");
let _ = server.join();
let _ = std::fs::remove_file(&path);
}
}