use super::*;
pub(crate) const CAPABILITY_DEADLINE: std::time::Duration = std::time::Duration::from_secs(2);
#[derive(Deserialize)]
pub(crate) struct Hello {
pub(crate) capability: String,
}
#[derive(Deserialize)]
#[serde(tag = "t")]
pub(crate) enum Said {
#[serde(rename = "watch")]
Watch { panes: Vec<String> },
#[serde(rename = "in")]
In { p: String, d: String },
#[serde(rename = "size")]
Size { p: String, c: u16, r: u16 },
#[serde(rename = "pace")]
Pace { slow: Vec<String> },
#[serde(rename = "more")]
More {
p: String,
before: usize,
n: usize,
#[serde(default)]
old: Option<u64>,
},
}
pub(crate) async fn forward(
live: Arc<crate::pane::Live>,
first: Vec<String>,
mut rx: broadcast::Receiver<Arc<str>>,
out: tokio::sync::mpsc::Sender<Arc<str>>,
) {
for f in first {
if out.send(f.into()).await.is_err() {
return;
}
}
loop {
match rx.recv().await {
Ok(m) => {
if out.send(m).await.is_err() {
return;
}
}
Err(broadcast::error::RecvError::Lagged(_)) => {
let (first, fresh) = live.attach();
rx = fresh;
for f in first {
if out.send(f.into()).await.is_err() {
return;
}
}
}
Err(broadcast::error::RecvError::Closed) => return,
}
}
}
pub(crate) async fn desk_socket(
State(app): S,
headers: HeaderMap,
Query(q): Query<std::collections::HashMap<String, String>>,
ws: WebSocketUpgrade,
) -> Response {
if q.contains_key(crate::desktop::CAPABILITY_KEY) || q.contains_key("capability") {
return (
StatusCode::FORBIDDEN,
Json(json!({ "error": "the capability is not a query parameter" })),
)
.into_response();
}
if !from_this_page(&headers) {
return (
StatusCode::FORBIDDEN,
Json(json!({ "error": "not from this page" })),
)
.into_response();
}
ws.on_upgrade(move |socket| desk_session(app, socket))
}
pub(crate) fn hello_allows(caps: &crate::capability::Capabilities, frame: Option<&str>) -> bool {
frame
.and_then(|f| serde_json::from_str::<Hello>(f).ok())
.is_some_and(|h| caps.verify(&h.capability))
}
pub(crate) async fn desk_session(app: Arc<App>, mut socket: WebSocket) {
let first = tokio::time::timeout(CAPABILITY_DEADLINE, socket.recv()).await;
let frame = match &first {
Ok(Some(Ok(Message::Text(t)))) => Some(t.as_str()),
_ => None,
};
if !hello_allows(&app.capabilities, frame) {
let _ = socket
.send(Message::Text(
json!({ "error": "no capability" }).to_string().into(),
))
.await;
let _ = socket.send(Message::Close(None)).await;
return;
}
let _ = socket
.send(Message::Text(json!({ "ok": true }).to_string().into()))
.await;
app.panes.arm_marks();
let (out, mut frames) = tokio::sync::mpsc::channel::<Arc<str>>(64);
let mut going = app.shutdown.subscribe();
let mut watching: std::collections::HashMap<
String,
(Arc<crate::pane::Live>, tokio::task::JoinHandle<()>),
> = Default::default();
let mut slowed = Slowed::default();
loop {
tokio::select! {
_ = going.recv() => break,
f = frames.recv() => {
let Some(f) = f else { break };
if socket.send(Message::Text(f.to_string().into())).await.is_err() {
break;
}
}
msg = socket.recv() => {
let text = match msg {
Some(Ok(Message::Text(t))) => t,
Some(Ok(Message::Close(_))) | None | Some(Err(_)) => break,
Some(Ok(_)) => continue,
};
let Ok(said) = serde_json::from_str::<Said>(text.as_str()) else { continue };
match said {
Said::Watch { panes } => {
let wanted: std::collections::HashSet<String> = panes
.into_iter()
.filter(|id| crate::pane::valid_id(id))
.filter(|id| matches!(app.store.pane(id), Ok(Some(_))))
.take(crate::desk::PER_DESK as usize)
.collect();
watching.retain(|id, (live, task)| {
let keep = wanted.contains(id);
if !keep {
task.abort();
slowed.set(id, live, false);
}
keep
});
for id in wanted {
if watching.contains_key(&id) {
continue;
}
let live = app.panes.get(&id);
let (first, rx) = live.attach();
let task = tokio::spawn(forward(live.clone(), first, rx, out.clone()));
watching.insert(id, (live, task));
}
}
Said::In { p, d } => {
if let Some((live, _)) = watching.get(&p) {
live.input(d.as_bytes(), &app.panes);
}
}
Said::Size { p, c, r } => {
if let Some((live, _)) = watching.get(&p) {
live.resize(c, r);
}
}
Said::Pace { slow } => {
for (id, (live, _)) in &watching {
slowed.set(id, live, slow.contains(id));
}
}
Said::More { p, before, n, old } => {
let Some((live, _)) = watching.get(&p) else { continue };
let Some(f) = live.more(old, before, n.min(crate::screen::KEEP_LINES)) else { continue };
if socket.send(Message::Text(f.into())).await.is_err() {
break;
}
}
}
}
}
}
for (_, (_, task)) in watching {
task.abort();
}
}
#[derive(Default)]
pub(crate) struct Slowed(pub(crate) std::collections::HashMap<String, Arc<crate::pane::Live>>);
impl Slowed {
pub(crate) fn set(&mut self, id: &str, live: &Arc<crate::pane::Live>, slow: bool) {
if slow == self.0.contains_key(id) {
return;
}
live.pace(slow);
if slow {
self.0.insert(id.to_string(), live.clone());
} else {
self.0.remove(id);
}
}
}
impl Drop for Slowed {
fn drop(&mut self) {
for live in self.0.values() {
live.pace(false);
}
}
}