WebSocket tests passing
This commit is contained in:
@@ -34,6 +34,8 @@ p256 = { version = "0.13", features = ["ecdh"] }
|
||||
hkdf = "0.12.3"
|
||||
sha2 = "0.10.1"
|
||||
reqwest = { version = "0.11", default-features = false, features = ["rustls-tls"]}
|
||||
tokio-tungstenite = "0.19.0"
|
||||
tungstenite = "0.19.0"
|
||||
|
||||
[dev-dependencies]
|
||||
ece = "2.2"
|
||||
|
||||
@@ -101,6 +101,9 @@ impl crate::Config {
|
||||
oauth_max_auth_attempts: settings.property_or_static("oauth.max-auth-attempts", "3")?,
|
||||
event_source_throttle: settings
|
||||
.property_or_static("jmap.event-source.throttle", "1s")?,
|
||||
web_socket_throttle: settings.property_or_static("jmap.web-socket.throttle", "1s")?,
|
||||
web_socket_timeout: settings.property_or_static("jmap.web-socket.timeout", "10m")?,
|
||||
web_socket_heartbeat: settings.property_or_static("jmap.web-socket.heartbeat", "1m")?,
|
||||
push_max_total: settings.property_or_static("jmap.push.max-total", "100")?,
|
||||
};
|
||||
config.add_capabilites(settings);
|
||||
|
||||
@@ -24,7 +24,7 @@ struct Ping {
|
||||
impl JMAP {
|
||||
pub async fn handle_event_source(
|
||||
&self,
|
||||
req: &HttpRequest,
|
||||
req: HttpRequest,
|
||||
acl_token: Arc<AclToken>,
|
||||
) -> HttpResponse {
|
||||
// Parse query
|
||||
|
||||
@@ -10,6 +10,7 @@ use hyper::{
|
||||
};
|
||||
use jmap_proto::{
|
||||
error::request::{RequestError, RequestLimitError},
|
||||
request::Request,
|
||||
response::Response,
|
||||
types::{blob::BlobId, id::Id},
|
||||
};
|
||||
@@ -23,39 +24,96 @@ use crate::{
|
||||
auth::oauth::OAuthMetadata,
|
||||
blob::{DownloadResponse, UploadResponse},
|
||||
services::state,
|
||||
websocket::upgrade::upgrade_websocket_connection,
|
||||
JMAP,
|
||||
};
|
||||
|
||||
use super::{session::Session, HtmlResponse, HttpResponse, JmapSessionManager, JsonResponse};
|
||||
use super::{
|
||||
session::Session, HtmlResponse, HttpRequest, HttpResponse, JmapSessionManager, JsonResponse,
|
||||
};
|
||||
|
||||
impl JMAP {
|
||||
pub async fn parse_request(
|
||||
&self,
|
||||
req: &mut hyper::Request<hyper::body::Incoming>,
|
||||
remote_ip: IpAddr,
|
||||
instance: &Arc<ServerInstance>,
|
||||
) -> HttpResponse {
|
||||
let mut path = req.uri().path().split('/');
|
||||
path.next();
|
||||
pub async fn parse_jmap_request(
|
||||
jmap: Arc<JMAP>,
|
||||
mut req: HttpRequest,
|
||||
remote_ip: IpAddr,
|
||||
instance: Arc<ServerInstance>,
|
||||
) -> HttpResponse {
|
||||
let mut path = req.uri().path().split('/');
|
||||
path.next();
|
||||
|
||||
match path.next().unwrap_or("") {
|
||||
"jmap" => {
|
||||
// Authenticate request
|
||||
let (_in_flight, acl_token) = match self.authenticate_headers(req, remote_ip).await
|
||||
{
|
||||
Ok(Some(session)) => session,
|
||||
Ok(None) => return RequestError::unauthorized().into_http_response(),
|
||||
Err(err) => return err.into_http_response(),
|
||||
};
|
||||
match path.next().unwrap_or("") {
|
||||
"jmap" => {
|
||||
// Authenticate request
|
||||
let (_in_flight, acl_token) = match jmap.authenticate_headers(&req, remote_ip).await {
|
||||
Ok(Some(session)) => session,
|
||||
Ok(None) => return RequestError::unauthorized().into_http_response(),
|
||||
Err(err) => return err.into_http_response(),
|
||||
};
|
||||
|
||||
match (path.next().unwrap_or(""), req.method()) {
|
||||
("", &Method::POST) => {
|
||||
return match fetch_body(req, self.config.request_max_size).await {
|
||||
match (path.next().unwrap_or(""), req.method()) {
|
||||
("", &Method::POST) => {
|
||||
return match fetch_body(&mut req, jmap.config.request_max_size)
|
||||
.await
|
||||
.and_then(|bytes| {
|
||||
Request::parse(
|
||||
&bytes,
|
||||
jmap.config.request_max_calls,
|
||||
jmap.config.request_max_size,
|
||||
)
|
||||
}) {
|
||||
Ok(request) => {
|
||||
//let _ = println!("<- {}", String::from_utf8_lossy(&bytes));
|
||||
|
||||
match jmap.handle_request(request, acl_token, &instance).await {
|
||||
Ok(response) => response.into_http_response(),
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
Err(err) => err.into_http_response(),
|
||||
};
|
||||
}
|
||||
("download", &Method::GET) => {
|
||||
if let (Some(_), Some(blob_id), Some(name)) = (
|
||||
path.next().and_then(|p| Id::from_bytes(p.as_bytes())),
|
||||
path.next().and_then(BlobId::from_base32),
|
||||
path.next(),
|
||||
) {
|
||||
return match jmap.blob_download(&blob_id, &acl_token).await {
|
||||
Ok(Some(blob)) => DownloadResponse {
|
||||
filename: name.to_string(),
|
||||
content_type: req
|
||||
.uri()
|
||||
.query()
|
||||
.and_then(|q| {
|
||||
form_urlencoded::parse(q.as_bytes())
|
||||
.find(|(k, _)| k == "accept")
|
||||
.map(|(_, v)| v.into_owned())
|
||||
})
|
||||
.unwrap_or("application/octet-stream".to_string()),
|
||||
blob,
|
||||
}
|
||||
.into_http_response(),
|
||||
Ok(None) => RequestError::not_found().into_http_response(),
|
||||
Err(_) => RequestError::internal_server_error().into_http_response(),
|
||||
};
|
||||
}
|
||||
}
|
||||
("upload", &Method::POST) => {
|
||||
if let Some(account_id) = path.next().and_then(|p| Id::from_bytes(p.as_bytes()))
|
||||
{
|
||||
return match fetch_body(&mut req, jmap.config.upload_max_size).await {
|
||||
Ok(bytes) => {
|
||||
//let delete = "fd";
|
||||
//println!("<- {}", String::from_utf8_lossy(&bytes));
|
||||
|
||||
match self.handle_request(&bytes, acl_token, instance).await {
|
||||
match jmap
|
||||
.blob_upload(
|
||||
account_id,
|
||||
req.headers()
|
||||
.get(CONTENT_TYPE)
|
||||
.and_then(|h| h.to_str().ok())
|
||||
.unwrap_or("application/octet-stream"),
|
||||
&bytes,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(response) => response.into_http_response(),
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
@@ -63,141 +121,90 @@ impl JMAP {
|
||||
Err(err) => err.into_http_response(),
|
||||
};
|
||||
}
|
||||
("download", &Method::GET) => {
|
||||
if let (Some(_), Some(blob_id), Some(name)) = (
|
||||
path.next().and_then(|p| Id::from_bytes(p.as_bytes())),
|
||||
path.next().and_then(BlobId::from_base32),
|
||||
path.next(),
|
||||
) {
|
||||
return match self.blob_download(&blob_id, &acl_token).await {
|
||||
Ok(Some(blob)) => DownloadResponse {
|
||||
filename: name.to_string(),
|
||||
content_type: req
|
||||
.uri()
|
||||
.query()
|
||||
.and_then(|q| {
|
||||
form_urlencoded::parse(q.as_bytes())
|
||||
.find(|(k, _)| k == "accept")
|
||||
.map(|(_, v)| v.into_owned())
|
||||
})
|
||||
.unwrap_or("application/octet-stream".to_string()),
|
||||
blob,
|
||||
}
|
||||
.into_http_response(),
|
||||
Ok(None) => RequestError::not_found().into_http_response(),
|
||||
Err(_) => {
|
||||
RequestError::internal_server_error().into_http_response()
|
||||
}
|
||||
};
|
||||
}
|
||||
}
|
||||
("upload", &Method::POST) => {
|
||||
if let Some(account_id) =
|
||||
path.next().and_then(|p| Id::from_bytes(p.as_bytes()))
|
||||
{
|
||||
return match fetch_body(req, self.config.upload_max_size).await {
|
||||
Ok(bytes) => {
|
||||
match self
|
||||
.blob_upload(
|
||||
account_id,
|
||||
req.headers()
|
||||
.get(CONTENT_TYPE)
|
||||
.and_then(|h| h.to_str().ok())
|
||||
.unwrap_or("application/octet-stream"),
|
||||
&bytes,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(response) => response.into_http_response(),
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
Err(err) => err.into_http_response(),
|
||||
};
|
||||
}
|
||||
}
|
||||
("eventsource", &Method::GET) => {
|
||||
return self.handle_event_source(req, acl_token).await
|
||||
}
|
||||
("ws", &Method::GET) => {
|
||||
todo!()
|
||||
}
|
||||
_ => (),
|
||||
}
|
||||
}
|
||||
".well-known" => match (path.next().unwrap_or(""), req.method()) {
|
||||
("jmap", &Method::GET) => {
|
||||
// Authenticate request
|
||||
let (_in_flight, acl_token) =
|
||||
match self.authenticate_headers(req, remote_ip).await {
|
||||
Ok(Some(session)) => session,
|
||||
Ok(None) => return RequestError::unauthorized().into_http_response(),
|
||||
Err(err) => return err.into_http_response(),
|
||||
};
|
||||
|
||||
return match self.handle_session_resource(instance, acl_token).await {
|
||||
Ok(session) => session.into_http_response(),
|
||||
Err(err) => err.into_http_response(),
|
||||
};
|
||||
("eventsource", &Method::GET) => {
|
||||
return jmap.handle_event_source(req, acl_token).await
|
||||
}
|
||||
("oauth-authorization-server", &Method::GET) => {
|
||||
let remote_addr = self.build_remote_addr(req, remote_ip);
|
||||
// Limit anonymous requests
|
||||
return match self.is_anonymous_allowed(remote_addr) {
|
||||
Ok(_) => JsonResponse::new(OAuthMetadata::new(&instance.data))
|
||||
.into_http_response(),
|
||||
Err(err) => err.into_http_response(),
|
||||
};
|
||||
("ws", &Method::GET) => {
|
||||
return upgrade_websocket_connection(jmap, req, acl_token, instance.clone())
|
||||
.await;
|
||||
}
|
||||
_ => (),
|
||||
},
|
||||
"auth" => {
|
||||
let remote_addr = self.build_remote_addr(req, remote_ip);
|
||||
}
|
||||
}
|
||||
".well-known" => match (path.next().unwrap_or(""), req.method()) {
|
||||
("jmap", &Method::GET) => {
|
||||
// Authenticate request
|
||||
let (_in_flight, acl_token) = match jmap.authenticate_headers(&req, remote_ip).await
|
||||
{
|
||||
Ok(Some(session)) => session,
|
||||
Ok(None) => return RequestError::unauthorized().into_http_response(),
|
||||
Err(err) => return err.into_http_response(),
|
||||
};
|
||||
|
||||
match (path.next().unwrap_or(""), req.method()) {
|
||||
("", &Method::GET) => {
|
||||
return match self.is_anonymous_allowed(remote_addr) {
|
||||
Ok(_) => self.handle_user_device_auth(req).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
return match jmap.handle_session_resource(instance, acl_token).await {
|
||||
Ok(session) => session.into_http_response(),
|
||||
Err(err) => err.into_http_response(),
|
||||
};
|
||||
}
|
||||
("oauth-authorization-server", &Method::GET) => {
|
||||
let remote_addr = jmap.build_remote_addr(&req, remote_ip);
|
||||
// Limit anonymous requests
|
||||
return match jmap.is_anonymous_allowed(remote_addr) {
|
||||
Ok(_) => {
|
||||
JsonResponse::new(OAuthMetadata::new(&instance.data)).into_http_response()
|
||||
}
|
||||
("", &Method::POST) => {
|
||||
return match self.is_auth_allowed(remote_addr) {
|
||||
Ok(_) => self.handle_user_device_auth_post(req).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
("code", &Method::GET) => {
|
||||
return match self.is_anonymous_allowed(remote_addr) {
|
||||
Ok(_) => self.handle_user_code_auth(req).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
("code", &Method::POST) => {
|
||||
return match self.is_auth_allowed(remote_addr) {
|
||||
Ok(_) => self.handle_user_code_auth_post(req).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
("device", &Method::POST) => {
|
||||
return match self.is_anonymous_allowed(remote_addr) {
|
||||
Ok(_) => self.handle_device_auth(req, instance).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
("token", &Method::POST) => {
|
||||
return match self.is_anonymous_allowed(remote_addr) {
|
||||
Ok(_) => self.handle_token_request(req).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
_ => (),
|
||||
}
|
||||
Err(err) => err.into_http_response(),
|
||||
};
|
||||
}
|
||||
_ => (),
|
||||
},
|
||||
"auth" => {
|
||||
let remote_addr = jmap.build_remote_addr(&req, remote_ip);
|
||||
|
||||
match (path.next().unwrap_or(""), req.method()) {
|
||||
("", &Method::GET) => {
|
||||
return match jmap.is_anonymous_allowed(remote_addr) {
|
||||
Ok(_) => jmap.handle_user_device_auth(&mut req).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
("", &Method::POST) => {
|
||||
return match jmap.is_auth_allowed(remote_addr) {
|
||||
Ok(_) => jmap.handle_user_device_auth_post(&mut req).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
("code", &Method::GET) => {
|
||||
return match jmap.is_anonymous_allowed(remote_addr) {
|
||||
Ok(_) => jmap.handle_user_code_auth(&mut req).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
("code", &Method::POST) => {
|
||||
return match jmap.is_auth_allowed(remote_addr) {
|
||||
Ok(_) => jmap.handle_user_code_auth_post(&mut req).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
("device", &Method::POST) => {
|
||||
return match jmap.is_anonymous_allowed(remote_addr) {
|
||||
Ok(_) => jmap.handle_device_auth(&mut req, instance).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
("token", &Method::POST) => {
|
||||
return match jmap.is_anonymous_allowed(remote_addr) {
|
||||
Ok(_) => jmap.handle_token_request(&mut req).await,
|
||||
Err(err) => err.into_http_response(),
|
||||
}
|
||||
}
|
||||
_ => (),
|
||||
}
|
||||
}
|
||||
RequestError::not_found().into_http_response()
|
||||
_ => (),
|
||||
}
|
||||
RequestError::not_found().into_http_response()
|
||||
}
|
||||
|
||||
impl SessionManager for JmapSessionManager {
|
||||
@@ -246,7 +253,7 @@ impl SessionManager for JmapSessionManager {
|
||||
}
|
||||
}
|
||||
|
||||
async fn handle_request<T: AsyncRead + AsyncWrite + Unpin + 'static>(
|
||||
async fn handle_request<T: AsyncRead + AsyncWrite + Unpin + Send + 'static>(
|
||||
jmap: Arc<JMAP>,
|
||||
session: SessionData<T>,
|
||||
) {
|
||||
@@ -257,27 +264,25 @@ async fn handle_request<T: AsyncRead + AsyncWrite + Unpin + 'static>(
|
||||
.keep_alive(true)
|
||||
.serve_connection(
|
||||
session.stream,
|
||||
service_fn(|mut req: hyper::Request<body::Incoming>| {
|
||||
service_fn(|req: hyper::Request<body::Incoming>| {
|
||||
let jmap = jmap.clone();
|
||||
let span = span.clone();
|
||||
let instance = session.instance.clone();
|
||||
|
||||
async move {
|
||||
let response = jmap
|
||||
.parse_request(&mut req, session.remote_ip, &instance)
|
||||
.await;
|
||||
|
||||
tracing::debug!(
|
||||
parent: &span,
|
||||
event = "request",
|
||||
uri = req.uri().to_string(),
|
||||
status = response.status().to_string(),
|
||||
);
|
||||
|
||||
let response = parse_jmap_request(jmap, req, session.remote_ip, instance).await;
|
||||
|
||||
Ok::<_, hyper::Error>(response)
|
||||
}
|
||||
}),
|
||||
)
|
||||
.with_upgrades()
|
||||
.await
|
||||
{
|
||||
tracing::debug!(
|
||||
@@ -289,10 +294,7 @@ async fn handle_request<T: AsyncRead + AsyncWrite + Unpin + 'static>(
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn fetch_body(
|
||||
req: &mut hyper::Request<hyper::body::Incoming>,
|
||||
max_size: usize,
|
||||
) -> Result<Vec<u8>, RequestError> {
|
||||
pub async fn fetch_body(req: &mut HttpRequest, max_size: usize) -> Result<Vec<u8>, RequestError> {
|
||||
let mut bytes = Vec::with_capacity(1024);
|
||||
while let Some(Ok(frame)) = req.frame().await {
|
||||
if let Some(data) = frame.data_ref() {
|
||||
|
||||
@@ -17,15 +17,10 @@ use crate::{auth::AclToken, JMAP};
|
||||
impl JMAP {
|
||||
pub async fn handle_request(
|
||||
&self,
|
||||
bytes: &[u8],
|
||||
request: Request,
|
||||
acl_token: Arc<AclToken>,
|
||||
instance: &Arc<ServerInstance>,
|
||||
) -> Result<Response, RequestError> {
|
||||
let request = Request::parse(
|
||||
bytes,
|
||||
self.config.request_max_calls,
|
||||
self.config.request_max_size,
|
||||
)?;
|
||||
let mut response = Response::new(
|
||||
acl_token.state(),
|
||||
request.created_ids.unwrap_or_default(),
|
||||
|
||||
@@ -143,7 +143,7 @@ pub struct BaseCapabilities {
|
||||
impl JMAP {
|
||||
pub async fn handle_session_resource(
|
||||
&self,
|
||||
instance: &ServerInstance,
|
||||
instance: Arc<ServerInstance>,
|
||||
acl_token: Arc<AclToken>,
|
||||
) -> Result<Session, RequestError> {
|
||||
let mut session = Session::new(&instance.data, &self.config.capabilities);
|
||||
|
||||
@@ -31,7 +31,7 @@ impl JMAP {
|
||||
pub async fn handle_device_auth(
|
||||
&self,
|
||||
req: &mut HttpRequest,
|
||||
instance: &ServerInstance,
|
||||
instance: Arc<ServerInstance>,
|
||||
) -> HttpResponse {
|
||||
// Parse form
|
||||
let client_id = match parse_form_data(req)
|
||||
|
||||
@@ -46,6 +46,7 @@ pub mod sieve;
|
||||
pub mod submission;
|
||||
pub mod thread;
|
||||
pub mod vacation;
|
||||
pub mod websocket;
|
||||
|
||||
pub const SUPERUSER_ID: u32 = 0;
|
||||
pub const LONG_SLUMBER: Duration = Duration::from_secs(60 * 60 * 24);
|
||||
@@ -103,6 +104,10 @@ pub struct Config {
|
||||
pub event_source_throttle: Duration,
|
||||
pub push_max_total: usize,
|
||||
|
||||
pub web_socket_throttle: Duration,
|
||||
pub web_socket_timeout: Duration,
|
||||
pub web_socket_heartbeat: Duration,
|
||||
|
||||
pub oauth_key: String,
|
||||
pub oauth_expiry_user_code: u64,
|
||||
pub oauth_expiry_auth_code: u64,
|
||||
|
||||
2
crates/jmap/src/websocket/mod.rs
Normal file
2
crates/jmap/src/websocket/mod.rs
Normal file
@@ -0,0 +1,2 @@
|
||||
pub mod stream;
|
||||
pub mod upgrade;
|
||||
192
crates/jmap/src/websocket/stream.rs
Normal file
192
crates/jmap/src/websocket/stream.rs
Normal file
@@ -0,0 +1,192 @@
|
||||
use std::{sync::Arc, time::Instant};
|
||||
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use hyper::upgrade::Upgraded;
|
||||
use jmap_proto::{
|
||||
error::request::RequestError,
|
||||
request::websocket::{
|
||||
WebSocketMessage, WebSocketRequestError, WebSocketResponse, WebSocketStateChange,
|
||||
},
|
||||
types::type_state::TypeState,
|
||||
};
|
||||
use tokio_tungstenite::WebSocketStream;
|
||||
use tungstenite::Message;
|
||||
use utils::{listener::ServerInstance, map::bitmap::Bitmap};
|
||||
|
||||
use crate::{auth::AclToken, JMAP};
|
||||
|
||||
impl JMAP {
|
||||
pub async fn handle_websocket_stream(
|
||||
&self,
|
||||
mut stream: WebSocketStream<Upgraded>,
|
||||
acl_token: Arc<AclToken>,
|
||||
instance: Arc<ServerInstance>,
|
||||
) {
|
||||
let span = tracing::info_span!(
|
||||
"WebSocket connection established",
|
||||
"account_id" = acl_token.primary_id(),
|
||||
"url" = instance.data,
|
||||
);
|
||||
|
||||
// Set timeouts
|
||||
let throttle = self.config.web_socket_throttle;
|
||||
let timeout = self.config.web_socket_timeout;
|
||||
let heartbeat = self.config.web_socket_heartbeat;
|
||||
let mut last_request = Instant::now();
|
||||
let mut last_changes_sent = Instant::now() - throttle;
|
||||
let mut last_heartbeat = Instant::now() - heartbeat;
|
||||
let mut next_event = heartbeat;
|
||||
|
||||
// Register with state manager
|
||||
let mut change_rx = if let Some(change_rx) = self
|
||||
.subscribe_state_manager(
|
||||
acl_token.primary_id(),
|
||||
acl_token.primary_id(),
|
||||
Bitmap::all(),
|
||||
)
|
||||
.await
|
||||
{
|
||||
change_rx
|
||||
} else {
|
||||
let _ = stream
|
||||
.send(Message::Text(
|
||||
WebSocketRequestError::from(RequestError::internal_server_error()).to_json(),
|
||||
))
|
||||
.await;
|
||||
return;
|
||||
};
|
||||
let mut changes = WebSocketStateChange::new(None);
|
||||
let mut change_types: Bitmap<TypeState> = Bitmap::new();
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
event = tokio::time::timeout(next_event, stream.next()) => {
|
||||
match event {
|
||||
Ok(Some(Ok(event))) => {
|
||||
match event {
|
||||
Message::Text(text) => {
|
||||
let response = match WebSocketMessage::parse(
|
||||
text.as_bytes(),
|
||||
self.config.request_max_calls,
|
||||
self.config.request_max_size,
|
||||
) {
|
||||
Ok(WebSocketMessage::Request(request)) => {
|
||||
match self
|
||||
.handle_request(
|
||||
request.request,
|
||||
acl_token.clone(),
|
||||
&instance,
|
||||
)
|
||||
.await
|
||||
{
|
||||
Ok(response) => {
|
||||
WebSocketResponse::from_response(response, request.id)
|
||||
.to_json()
|
||||
}
|
||||
Err(err) => {
|
||||
WebSocketRequestError::from_error(err, request.id)
|
||||
.to_json()
|
||||
}
|
||||
}
|
||||
}
|
||||
Ok(WebSocketMessage::PushEnable(push_enable)) => {
|
||||
change_types = if !push_enable.data_types.is_empty() {
|
||||
push_enable.data_types.into()
|
||||
} else {
|
||||
Bitmap::all()
|
||||
};
|
||||
continue;
|
||||
}
|
||||
Ok(WebSocketMessage::PushDisable) => {
|
||||
change_types = Bitmap::new();
|
||||
continue;
|
||||
}
|
||||
Err(err) => err.to_json(),
|
||||
};
|
||||
if let Err(err) = stream.send(Message::Text(response)).await {
|
||||
tracing::debug!(parent: &span, error = ?err, "Failed to send text message");
|
||||
}
|
||||
}
|
||||
Message::Ping(bytes) => {
|
||||
if let Err(err) = stream.send(Message::Pong(bytes)).await {
|
||||
tracing::debug!(parent: &span, error = ?err, "Failed to send pong message");
|
||||
}
|
||||
}
|
||||
Message::Close(frame) => {
|
||||
let _ = stream.close(frame).await;
|
||||
break;
|
||||
}
|
||||
_ => (),
|
||||
}
|
||||
|
||||
last_request = Instant::now();
|
||||
last_heartbeat = Instant::now();
|
||||
}
|
||||
Ok(Some(Err(err))) => {
|
||||
tracing::debug!(parent: &span, error = ?err, "Websocket error");
|
||||
break;
|
||||
}
|
||||
Ok(None) => break,
|
||||
Err(_) => {
|
||||
// Verify timeout
|
||||
if last_request.elapsed() > timeout {
|
||||
tracing::debug!(
|
||||
parent: &span,
|
||||
event = "disconnect",
|
||||
"Disconnecting idle client"
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
state_change = change_rx.recv() => {
|
||||
if let Some(state_change) = state_change {
|
||||
if !change_types.is_empty() && state_change
|
||||
.types
|
||||
.iter()
|
||||
.any(|(t, _)| change_types.contains(*t))
|
||||
{
|
||||
for (type_state, change_id) in state_change.types {
|
||||
changes
|
||||
.changed
|
||||
.get_mut_or_insert(state_change.account_id.into())
|
||||
.set(type_state, change_id.into());
|
||||
}
|
||||
}
|
||||
} else {
|
||||
tracing::debug!(
|
||||
parent: &span,
|
||||
event = "channel-closed",
|
||||
"Disconnecting client, channel closed"
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if !changes.changed.is_empty() {
|
||||
// Send any queued changes
|
||||
let elapsed = last_changes_sent.elapsed();
|
||||
if elapsed >= throttle {
|
||||
if let Err(err) = stream.send(Message::Text(changes.to_json())).await {
|
||||
tracing::debug!(parent: &span, error = ?err, "Failed to send state change message");
|
||||
}
|
||||
changes.changed.clear();
|
||||
last_changes_sent = Instant::now();
|
||||
last_heartbeat = Instant::now();
|
||||
next_event = heartbeat;
|
||||
} else {
|
||||
next_event = throttle - elapsed;
|
||||
}
|
||||
} else if last_heartbeat.elapsed() > heartbeat {
|
||||
if let Err(err) = stream.send(Message::Ping(vec![])).await {
|
||||
tracing::debug!(parent: &span, error = ?err, "Failed to send ping message");
|
||||
break;
|
||||
}
|
||||
last_heartbeat = Instant::now();
|
||||
next_event = heartbeat;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
88
crates/jmap/src/websocket/upgrade.rs
Normal file
88
crates/jmap/src/websocket/upgrade.rs
Normal file
@@ -0,0 +1,88 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use http_body_util::{BodyExt, Full};
|
||||
use hyper::{body::Bytes, Response, StatusCode};
|
||||
use jmap_proto::error::request::RequestError;
|
||||
use tokio_tungstenite::WebSocketStream;
|
||||
use tungstenite::{handshake::derive_accept_key, protocol::Role};
|
||||
use utils::listener::ServerInstance;
|
||||
|
||||
use crate::{
|
||||
api::{http::ToHttpResponse, HttpRequest, HttpResponse},
|
||||
auth::AclToken,
|
||||
JMAP,
|
||||
};
|
||||
|
||||
pub async fn upgrade_websocket_connection(
|
||||
jmap: Arc<JMAP>,
|
||||
req: HttpRequest,
|
||||
acl_token: Arc<AclToken>,
|
||||
instance: Arc<ServerInstance>,
|
||||
) -> HttpResponse {
|
||||
let headers = req.headers();
|
||||
if headers
|
||||
.get(hyper::header::CONNECTION)
|
||||
.and_then(|h| h.to_str().ok())
|
||||
!= Some("Upgrade")
|
||||
|| headers
|
||||
.get(hyper::header::UPGRADE)
|
||||
.and_then(|h| h.to_str().ok())
|
||||
!= Some("websocket")
|
||||
{
|
||||
return RequestError::blank(
|
||||
StatusCode::BAD_REQUEST.as_u16(),
|
||||
"WebSocket upgrade failed",
|
||||
"Missing or Invalid Connection or Upgrade headers.",
|
||||
)
|
||||
.into_http_response();
|
||||
}
|
||||
let derived_key = match (
|
||||
headers
|
||||
.get("Sec-WebSocket-Key")
|
||||
.and_then(|h| h.to_str().ok()),
|
||||
headers
|
||||
.get("Sec-WebSocket-Version")
|
||||
.and_then(|h| h.to_str().ok()),
|
||||
) {
|
||||
(Some(key), Some(version)) if version == "13" => derive_accept_key(key.as_bytes()),
|
||||
_ => {
|
||||
return RequestError::blank(
|
||||
StatusCode::BAD_REQUEST.as_u16(),
|
||||
"WebSocket upgrade failed",
|
||||
"Missing or Invalid Sec-WebSocket-Key headers.",
|
||||
)
|
||||
.into_http_response();
|
||||
}
|
||||
};
|
||||
|
||||
// Spawn WebSocket connection
|
||||
tokio::spawn(async move {
|
||||
// Upgrade connection
|
||||
match hyper::upgrade::on(req).await {
|
||||
Ok(upgraded) => {
|
||||
jmap.handle_websocket_stream(
|
||||
WebSocketStream::from_raw_socket(upgraded, Role::Server, None).await,
|
||||
acl_token,
|
||||
instance,
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Err(e) => {
|
||||
tracing::debug!("WebSocket upgrade failed: {}", e);
|
||||
}
|
||||
}
|
||||
});
|
||||
|
||||
Response::builder()
|
||||
.status(hyper::StatusCode::SWITCHING_PROTOCOLS)
|
||||
.header(hyper::header::CONNECTION, "upgrade")
|
||||
.header(hyper::header::UPGRADE, "websocket")
|
||||
.header("Sec-WebSocket-Accept", &derived_key)
|
||||
.header("Sec-WebSocket-Protocol", "jmap")
|
||||
.body(
|
||||
Full::new(Bytes::from("Switching to WebSocket protocol"))
|
||||
.map_err(|never| match never {})
|
||||
.boxed(),
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
Reference in New Issue
Block a user