use axum::extract::ws::{Message, WebSocket};
use axum::extract::{Path, State, WebSocketUpgrade};
use axum::http::{HeaderMap, StatusCode, header};
use axum::response::{IntoResponse, Response};
use futures::{SinkExt, StreamExt};
use kanade_shared::feature::Feature;
use kanade_shared::subject;
use kanade_shared::wire::{
FrameKind, FrameMeta, RemoteCtrl, RemoteCtrlReply, TileEncoding, frame_kind,
};
use serde::Serialize;
use std::time::Duration;
use tracing::{info, warn};
use super::AppState;
use crate::auth::{Claims, Role, verify_bearer};
pub const SUBPROTOCOL: &str = "kanade.remote.v1";
const BEARER_PREFIX: &str = "bearer.";
const START_TIMEOUT: Duration = Duration::from_secs(15);
const STOP_TIMEOUT: Duration = Duration::from_secs(5);
#[derive(Serialize, Debug, PartialEq)]
#[serde(tag = "kind", rename_all = "snake_case")]
enum SocketMeta {
Started {
session_id: String,
screen_w: Option<u32>,
screen_h: Option<u32>,
allow_input: bool,
},
Tile {
#[serde(flatten)]
meta: FrameMeta,
encoding: TileEncoding,
},
Gap { reason: String },
Resumed,
Ended { reason: String },
}
fn frame(meta: &SocketMeta, payload: &[u8]) -> Vec<u8> {
let json = serde_json::to_vec(meta).unwrap_or_else(|e| {
warn!(error = %e, "serialising socket meta");
br#"{"kind":"ended","reason":"backend could not encode this message"}"#.to_vec()
});
let mut out = Vec::with_capacity(4 + json.len() + payload.len());
out.extend_from_slice(&(json.len() as u32).to_le_bytes());
out.extend_from_slice(&json);
out.extend_from_slice(payload);
out
}
fn bearer_from_protocols(headers: &HeaderMap) -> Option<String> {
headers
.get(header::SEC_WEBSOCKET_PROTOCOL)?
.to_str()
.ok()?
.split(',')
.map(str::trim)
.find_map(|p| p.strip_prefix(BEARER_PREFIX))
.filter(|t| !t.is_empty())
.map(str::to_owned)
}
fn feature_allowed(claims: &Claims, feature: Feature) -> bool {
claims
.allowed_features
.as_ref()
.is_none_or(|allowed| allowed.contains(&feature))
}
pub async fn ws(
State(state): State<AppState>,
Path(pc_id): Path<String>,
headers: HeaderMap,
upgrade: WebSocketUpgrade,
) -> Response {
let token = bearer_from_protocols(&headers);
let claims = match verify_bearer(&state.pool, token.as_deref()).await {
Ok(c) => c,
Err(reason) => {
warn!(pc_id, reason, "remote ws: auth rejected");
return (StatusCode::UNAUTHORIZED, reason).into_response();
}
};
if !claims.role().allows(Role::Operator) {
warn!(pc_id, sub = %claims.sub, role = claims.role().as_str(), "remote ws: role denied");
return (
StatusCode::FORBIDDEN,
"operator role required to view a remote screen",
)
.into_response();
}
if !feature_allowed(&claims, Feature::Remote) {
warn!(pc_id, sub = %claims.sub, "remote ws: feature denied");
return (
StatusCode::FORBIDDEN,
"account not permitted to access this page (requires remote)",
)
.into_response();
}
let operator = claims.sub.clone();
upgrade
.protocols([SUBPROTOCOL])
.on_upgrade(move |socket| relay(socket, state, pc_id, operator))
}
async fn relay(mut socket: WebSocket, state: AppState, pc_id: String, operator: String) {
let session_id = format!("sess-{}", uuid::Uuid::new_v4());
let frames = match state
.nats
.subscribe(subject::remote_frame(&session_id))
.await
{
Ok(s) => s,
Err(e) => {
warn!(pc_id, session_id, error = %e, "remote ws: frame subscribe failed");
end(&mut socket, format!("backend could not subscribe: {e}")).await;
return;
}
};
let start = RemoteCtrl::Start {
session_id: session_id.clone(),
output_index: 0,
quality: 75,
max_fps: 10,
allow_input: false,
};
let reply = match request_ctrl(&state, &pc_id, &start, START_TIMEOUT).await {
Ok(r) => r,
Err(e) => {
let live = e.may_have_started();
let reason = e.into_reason();
warn!(pc_id, session_id, reason, live, "remote ws: start failed");
end(&mut socket, reason).await;
if live {
stop_session(&state, &pc_id, &session_id).await;
}
return;
}
};
if !reply.accepted {
let reason = reply
.reason
.unwrap_or_else(|| "the endpoint refused the session".to_string());
info!(pc_id, session_id, operator, reason, "remote ws: refused");
end(&mut socket, reason).await;
return;
}
info!(pc_id, session_id, operator, "remote ws: streaming");
let opened = frame(
&SocketMeta::Started {
session_id: session_id.clone(),
screen_w: reply.screen_w,
screen_h: reply.screen_h,
allow_input: false,
},
&[],
);
if socket.send(Message::Binary(opened.into())).await.is_ok() {
pump(&mut socket, frames).await;
}
stop_session(&state, &pc_id, &session_id).await;
info!(pc_id, session_id, operator, "remote ws: closed");
}
async fn stop_session(state: &AppState, pc_id: &str, session_id: &str) {
let stop = RemoteCtrl::Stop {
session_id: session_id.to_owned(),
};
if let Err(e) = request_ctrl(state, pc_id, &stop, STOP_TIMEOUT).await {
let reason = e.into_reason();
warn!(
pc_id,
session_id, reason, "remote ws: stop not acknowledged"
);
}
}
async fn pump(socket: &mut WebSocket, mut frames: async_nats::Subscriber) {
loop {
tokio::select! {
msg = frames.next() => {
let Some(msg) = msg else {
end(socket, "the frame stream ended".to_string()).await;
return;
};
let Some((meta, payload)) = translate(&msg) else { continue };
if socket.send(Message::Binary(frame(&meta, payload).into())).await.is_err() {
return; }
}
incoming = socket.recv() => {
match incoming {
Some(Ok(Message::Close(_))) => {
let _ = socket.send(Message::Close(None)).await;
return;
}
None | Some(Err(_)) => return,
Some(Ok(_)) => {}
}
}
}
}
}
fn translate(msg: &async_nats::Message) -> Option<(SocketMeta, &[u8])> {
let headers = msg.headers.as_ref()?;
let kind = match frame_kind(headers) {
Ok(k) => k,
Err(e) => {
warn!(error = %e, "remote ws: undecodable frame kind");
return None;
}
};
match kind {
FrameKind::Tile => match FrameMeta::from_headers(headers) {
Ok((meta, encoding)) => Some((SocketMeta::Tile { meta, encoding }, &msg.payload[..])),
Err(e) => {
warn!(error = %e, "remote ws: undecodable tile meta");
None
}
},
FrameKind::Gap => Some((
SocketMeta::Gap {
reason: String::from_utf8_lossy(&msg.payload).into_owned(),
},
&[][..],
)),
FrameKind::Resumed => Some((SocketMeta::Resumed, &[][..])),
}
}
#[derive(Debug)]
enum CtrlError {
NotDelivered(String),
Indeterminate(String),
}
impl CtrlError {
fn may_have_started(&self) -> bool {
matches!(self, CtrlError::Indeterminate(_))
}
fn into_reason(self) -> String {
match self {
CtrlError::NotDelivered(r) | CtrlError::Indeterminate(r) => r,
}
}
}
async fn request_ctrl(
state: &AppState,
pc_id: &str,
ctrl: &RemoteCtrl,
timeout: Duration,
) -> Result<RemoteCtrlReply, CtrlError> {
let payload = serde_json::to_vec(ctrl)
.map_err(|e| CtrlError::NotDelivered(format!("backend encode failed: {e}")))?;
let request = state
.nats
.request(subject::remote_ctrl(pc_id), payload.into());
let msg = match tokio::time::timeout(timeout, request).await {
Ok(Ok(msg)) => msg,
Ok(Err(e)) => {
return Err(match e.kind() {
async_nats::client::RequestErrorKind::NoResponders => {
CtrlError::NotDelivered(format!("{pc_id} is not reachable: no agent listening"))
}
_ => CtrlError::Indeterminate(format!("{pc_id} did not answer: {e}")),
});
}
Err(_) => {
return Err(CtrlError::Indeterminate(format!(
"{pc_id} did not answer within {}s",
timeout.as_secs()
)));
}
};
serde_json::from_slice(&msg.payload)
.map_err(|e| CtrlError::Indeterminate(format!("{pc_id} sent a bad reply: {e}")))
}
async fn end(socket: &mut WebSocket, reason: String) {
let bytes = frame(&SocketMeta::Ended { reason }, &[]);
let _ = socket.send(Message::Binary(bytes.into())).await;
let _ = socket.close().await;
}
#[cfg(test)]
mod tests {
use super::*;
use async_nats::HeaderMap as NatsHeaders;
fn parse_frame(bytes: &[u8]) -> (serde_json::Value, &[u8]) {
let len = u32::from_le_bytes(bytes[..4].try_into().unwrap()) as usize;
let meta = serde_json::from_slice(&bytes[4..4 + len]).unwrap();
(meta, &bytes[4 + len..])
}
#[test]
fn frame_prefixes_meta_length_only() {
let bytes = frame(&SocketMeta::Resumed, b"ignored-payload");
let (meta, payload) = parse_frame(&bytes);
assert_eq!(meta["kind"], "resumed");
assert_eq!(payload, b"ignored-payload");
}
#[test]
fn tile_meta_is_flat_so_the_spa_reads_one_object() {
let meta = FrameMeta {
frame_seq: 7,
tile_index: 1,
tile_count: 3,
x: 10,
y: 20,
w: 100,
h: 50,
screen_w: 3840,
screen_h: 1600,
captured_at_ms: 1_700_000_000_000,
};
let bytes = frame(
&SocketMeta::Tile {
meta,
encoding: TileEncoding::Jpeg,
},
&[0xFF, 0xD8, 0xFF],
);
let (json, payload) = parse_frame(&bytes);
assert_eq!(json["kind"], "tile");
assert_eq!(json["frame_seq"], 7);
assert_eq!(json["tile_count"], 3);
assert_eq!(json["screen_w"], 3840);
assert_eq!(json["encoding"], "jpeg");
assert_eq!(payload, &[0xFF, 0xD8, 0xFF]);
}
#[test]
fn socket_kinds_match_the_wire_vocabulary() {
for (meta, want) in [
(
SocketMeta::Tile {
meta: FrameMeta {
frame_seq: 0,
tile_index: 0,
tile_count: 1,
x: 0,
y: 0,
w: 1,
h: 1,
screen_w: 1,
screen_h: 1,
captured_at_ms: 0,
},
encoding: TileEncoding::Jpeg,
},
FrameKind::Tile.as_str(),
),
(
SocketMeta::Gap {
reason: "locked".into(),
},
FrameKind::Gap.as_str(),
),
(SocketMeta::Resumed, FrameKind::Resumed.as_str()),
] {
let (json, _) = parse_frame(&frame(&meta, &[]));
assert_eq!(json["kind"], want);
}
}
#[test]
fn gap_reason_moves_into_the_meta() {
let mut headers = NatsHeaders::new();
headers.insert(kanade_shared::wire::remote_header::KIND, "gap");
let msg = message(headers, b"the workstation is locked".to_vec());
let (meta, payload) = translate(&msg).expect("translated");
assert_eq!(
meta,
SocketMeta::Gap {
reason: "the workstation is locked".into()
}
);
assert!(payload.is_empty());
}
#[test]
fn undecodable_frames_are_dropped_not_fatal() {
assert!(translate(&message(NatsHeaders::new(), vec![])).is_none());
let mut unknown = NatsHeaders::new();
unknown.insert(kanade_shared::wire::remote_header::KIND, "hologram");
assert!(translate(&message(unknown, vec![])).is_none());
let mut headerless_tile = NatsHeaders::new();
headerless_tile.insert(kanade_shared::wire::remote_header::KIND, "tile");
assert!(translate(&message(headerless_tile, vec![1, 2, 3])).is_none());
}
fn message(headers: NatsHeaders, payload: Vec<u8>) -> async_nats::Message {
async_nats::Message {
subject: "remote.frame.sess-test".into(),
reply: None,
payload: payload.into(),
headers: Some(headers),
status: None,
description: None,
length: 0,
}
}
fn protocols(value: &str) -> HeaderMap {
let mut h = HeaderMap::new();
h.insert(header::SEC_WEBSOCKET_PROTOCOL, value.parse().unwrap());
h
}
#[test]
fn bearer_is_read_out_of_the_protocol_list() {
assert_eq!(
bearer_from_protocols(&protocols("kanade.remote.v1, bearer.abc.def.ghi")).as_deref(),
Some("abc.def.ghi")
);
assert_eq!(
bearer_from_protocols(&protocols("bearer.tok,kanade.remote.v1")).as_deref(),
Some("tok")
);
}
#[test]
fn a_missing_or_empty_credential_is_absent_not_blank() {
assert!(bearer_from_protocols(&HeaderMap::new()).is_none());
assert!(bearer_from_protocols(&protocols("kanade.remote.v1")).is_none());
assert!(bearer_from_protocols(&protocols("bearer.")).is_none());
}
#[test]
fn only_proven_non_delivery_skips_the_teardown() {
assert!(CtrlError::Indeterminate("timed out".into()).may_have_started());
assert!(CtrlError::Indeterminate("bad reply".into()).may_have_started());
assert!(!CtrlError::NotDelivered("no agent listening".into()).may_have_started());
assert_eq!(
CtrlError::Indeterminate("PC1 did not answer".into()).into_reason(),
"PC1 did not answer"
);
}
#[test]
fn feature_gate_matches_the_middleware() {
let claims = |allowed: Option<Vec<Feature>>| Claims {
sub: "op".into(),
exp: 0,
aud: None,
roles: vec!["operator".into()],
allowed_features: allowed,
};
assert!(feature_allowed(&claims(None), Feature::Remote));
assert!(feature_allowed(
&claims(Some(vec![Feature::Remote, Feature::Audit])),
Feature::Remote
));
assert!(!feature_allowed(
&claims(Some(vec![Feature::Audit])),
Feature::Remote
));
assert!(!feature_allowed(&claims(Some(vec![])), Feature::Remote));
}
}