WebSocket tests passing

This commit is contained in:
Mauro D
2023-05-24 17:39:09 +00:00
parent 4ff2158783
commit 86a8a5f7d5
18 changed files with 1004 additions and 331 deletions

View File

@@ -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"

View File

@@ -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);

View File

@@ -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

View File

@@ -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() {

View File

@@ -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(),

View File

@@ -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);

View File

@@ -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)

View File

@@ -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,

View File

@@ -0,0 +1,2 @@
pub mod stream;
pub mod upgrade;

View 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;
}
}
}
}

View 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()
}