use std::net::{Ipv4Addr, TcpListener, TcpStream};
use std::thread;
use tungstenite::handshake::server::{ErrorResponse, Request, Response};
use tungstenite::http::{HeaderValue, StatusCode};
use tungstenite::protocol::WebSocketConfig;
use tungstenite::{Message, WebSocket, accept_hdr_with_config};
use nerve_protocol::codec::decode;
use nerve_protocol::constants::{HEADER_SIZE, MAX_PAYLOAD_SIZE};
use nerve_protocol::frame::OwnedFrame;
use crate::config::Config;
use crate::dispatch::{DispatchAction, dispatch_frame};
use crate::request_table::RequestTable;
fn nerve_ws_config() -> WebSocketConfig {
let limit = MAX_PAYLOAD_SIZE + HEADER_SIZE;
WebSocketConfig {
max_frame_size: Some(limit),
max_message_size: Some(limit),
..WebSocketConfig::default()
}
}
pub fn run_ws(config: Config, token: String) -> std::io::Result<()> {
let listener = TcpListener::bind((config.bind_addr, config.ws_port))?;
tracing::info!(addr = %listener.local_addr()?, "WebSocket server listening");
run_ws_on_listener(listener, config, token)
}
pub fn run_ws_on_listener(
listener: TcpListener,
config: Config,
token: String,
) -> std::io::Result<()> {
if let Ok(addr) = listener.local_addr() {
debug_assert!(
addr.ip() == std::net::IpAddr::V4(Ipv4Addr::LOCALHOST),
"WebSocket server must bind only to 127.0.0.1, got {addr}"
);
}
for stream in listener.incoming() {
match stream {
Ok(tcp) => {
let cfg = config.clone();
let tok = token.clone();
thread::spawn(move || {
handle_ws_connection(tcp, &cfg, &tok);
});
}
Err(e) => {
tracing::warn!(error = %e, "WebSocket accept error");
}
}
}
Ok(())
}
#[allow(clippy::result_large_err)]
fn handle_ws_connection(stream: TcpStream, config: &Config, token: &str) {
let config_c = config.clone();
let token_c = token.to_string();
let mut ws: WebSocket<TcpStream> = match accept_hdr_with_config(
stream,
move |req: &Request, resp: Response| validate_handshake(req, resp, &config_c, &token_c),
Some(nerve_ws_config()),
) {
Ok(ws) => ws,
Err(e) => {
tracing::debug!(error = %e, "WebSocket handshake rejected");
return;
}
};
tracing::debug!("WebSocket connection authenticated");
let mut requests = RequestTable::new();
loop {
let msg = match ws.read() {
Ok(m) => m,
Err(e) => {
tracing::debug!(error = %e, "WebSocket read error");
break;
}
};
let data: Vec<u8> = match msg {
Message::Binary(b) => b,
Message::Close(_) => break,
Message::Ping(_) | Message::Pong(_) => continue,
Message::Text(_) | Message::Frame(_) => {
tracing::warn!("received non-binary WebSocket frame; closing connection");
break;
}
};
let owned: OwnedFrame = match decode_nerve_frame(&data) {
Ok(f) => f,
Err(e) => {
tracing::warn!(error = e, "invalid NERVE frame; closing connection");
break;
}
};
let action = match dispatch_frame(owned.as_borrowed(), &mut requests) {
Ok(a) => a,
Err(e) => {
tracing::warn!(error = %e, "dispatch error");
break;
}
};
match action {
DispatchAction::Handled => {}
DispatchAction::Reply(bytes) => {
if let Err(e) = ws.send(Message::Binary(bytes)) {
tracing::warn!(error = %e, "WebSocket write error");
break;
}
}
DispatchAction::ForwardToAiDaemon(req_id) => {
tracing::debug!(
request_id = req_id.0,
"SearchQuery pending AI daemon (WebSocket)"
);
}
}
}
tracing::debug!("WebSocket connection closed");
}
#[allow(clippy::result_large_err)]
fn validate_handshake(
request: &Request,
mut response: Response,
config: &Config,
stored_token: &str,
) -> Result<Response, ErrorResponse> {
if let Some(allowed_origin) = config.allowed_origin() {
let origin = request
.headers()
.get("Origin")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
if origin != allowed_origin {
tracing::warn!("WebSocket rejected: invalid Origin");
return Err(forbidden());
}
}
let protocol_header = request
.headers()
.get("Sec-WebSocket-Protocol")
.and_then(|v| v.to_str().ok())
.unwrap_or("");
let matched = match find_auth_protocol(protocol_header, stored_token) {
Some(p) => p,
None => {
tracing::warn!("WebSocket rejected: invalid or missing token");
return Err(forbidden());
}
};
match HeaderValue::from_str(matched) {
Ok(hv) => {
response.headers_mut().insert("Sec-WebSocket-Protocol", hv);
}
Err(_) => {
return Err(forbidden());
}
}
Ok(response)
}
fn find_auth_protocol<'a>(protocol_header: &'a str, stored_token: &str) -> Option<&'a str> {
for entry in protocol_header.split(',') {
let p = entry.trim();
if let Some(provided_token) = p.strip_prefix("anvesha-v1.")
&& provided_token == stored_token
{
return Some(p);
}
}
None
}
fn forbidden() -> ErrorResponse {
Response::builder()
.status(StatusCode::FORBIDDEN)
.body(None)
.unwrap()
}
fn decode_nerve_frame(data: &[u8]) -> Result<OwnedFrame, &'static str> {
let frame = decode(data).map_err(|_| "NERVE frame decode failed")?;
Ok(OwnedFrame {
header: frame.header,
payload: frame.payload.to_vec(),
})
}