use car_memgine::MemgineEngine;
use car_server_core::{run_dispatch, ServerState, ServerStateConfig};
use futures::{SinkExt, StreamExt};
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4};
use std::sync::Arc;
use tempfile::TempDir;
use tokio::net::TcpListener;
use tokio::sync::Mutex;
use tokio_tungstenite::{accept_async, connect_async, tungstenite::Message};
type Ws =
tokio_tungstenite::WebSocketStream<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>>;
fn state(journal_dir: std::path::PathBuf) -> Arc<ServerState> {
let engine = Arc::new(Mutex::new(MemgineEngine::new(None)));
let config = ServerStateConfig::new(journal_dir).with_shared_memgine(engine);
Arc::new(ServerState::with_config(config))
}
async fn spawn_dispatcher(state: Arc<ServerState>, connections: usize) -> SocketAddr {
let listener = TcpListener::bind(SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::LOCALHOST, 0)))
.await
.expect("bind loopback");
let address = listener.local_addr().expect("local address");
tokio::spawn(async move {
for _ in 0..connections {
let (stream, peer) = listener.accept().await.expect("accept");
let socket = accept_async(stream).await.expect("WebSocket handshake");
let (write, read) = socket.split();
let state = state.clone();
tokio::spawn(async move {
let _ = run_dispatch(read, Box::pin(write), peer.to_string(), state).await;
});
}
});
address
}
async fn call(
socket: &mut Ws,
id: &str,
method: &str,
params: serde_json::Value,
) -> serde_json::Value {
socket
.send(Message::Text(
serde_json::json!({
"jsonrpc": "2.0",
"id": id,
"method": method,
"params": params,
})
.to_string()
.into(),
))
.await
.expect("send request");
let text = socket
.next()
.await
.expect("response frame")
.expect("response frame ok")
.into_text()
.expect("text response");
serde_json::from_str(&text).expect("parse response")
}
async fn negotiate(socket: &mut Ws, id: &str) -> serde_json::Value {
call(
socket,
id,
"server.handshake",
serde_json::json!({ "protocol_version": car_proto::PROTOCOL_VERSION }),
)
.await
}
fn assert_handshake_required(response: &serde_json::Value, method: &str) {
assert_eq!(
response["error"]["code"],
car_proto::PROTOCOL_HANDSHAKE_REQUIRED_ERROR_CODE,
"{method} should fail with the typed handshake-required code: {response}"
);
assert!(
response["error"]["message"]
.as_str()
.unwrap_or_default()
.starts_with(car_proto::PROTOCOL_HANDSHAKE_REQUIRED_MESSAGE_PREFIX),
"{method} should carry the stable handshake-required prefix: {response}"
);
assert!(
response.get("result").is_none(),
"{method} must not dispatch before negotiation: {response}"
);
}
#[tokio::test]
async fn auth_and_host_surfaces_require_exact_v2_before_dispatch() {
let journal = TempDir::new().expect("journal tempdir");
let address = spawn_dispatcher(state(journal.path().to_path_buf()), 2).await;
let (mut legacy, _) = connect_async(format!("ws://{address}"))
.await
.expect("connect legacy client");
for (index, (method, params)) in [
("auth.start", serde_json::json!({})),
(
"auth.complete",
serde_json::json!({
"redirect_uri": "http://127.0.0.1/callback",
"code": "must-not-be-consumed",
"verifier": "legacy-verifier",
"attempt_id": "legacy-attempt",
}),
),
(
"auth.completion_status",
serde_json::json!({ "attempt_id": "legacy-attempt" }),
),
("auth.status", serde_json::json!({})),
("host.subscribe", serde_json::json!({})),
]
.into_iter()
.enumerate()
{
let response = call(&mut legacy, &format!("legacy-{index}"), method, params).await;
assert_handshake_required(&response, method);
}
for (id, params) in [
("missing-version", serde_json::json!({})),
(
"string-version",
serde_json::json!({ "protocol_version": "2" }),
),
(
"mismatch",
serde_json::json!({ "protocol_version": car_proto::PROTOCOL_VERSION - 1 }),
),
] {
let mismatch = call(&mut legacy, id, "server.handshake", params).await;
assert_eq!(
mismatch["error"]["code"],
car_proto::PROTOCOL_VERSION_MISMATCH_ERROR_CODE
);
assert!(mismatch["error"]["message"]
.as_str()
.unwrap_or_default()
.starts_with(car_proto::PROTOCOL_VERSION_MISMATCH_MESSAGE_PREFIX));
}
let after_mismatch = call(
&mut legacy,
"after-mismatch",
"auth.start",
serde_json::json!({}),
)
.await;
assert_handshake_required(&after_mismatch, "auth.start");
let (mut compatible, _) = connect_async(format!("ws://{address}"))
.await
.expect("connect v2 client");
let handshake = negotiate(&mut compatible, "v2").await;
assert_eq!(
handshake["result"]["protocol_version"],
car_proto::PROTOCOL_VERSION
);
assert_eq!(
handshake["result"]["client_protocol_version"],
car_proto::PROTOCOL_VERSION
);
let repeated = negotiate(&mut compatible, "v2-again").await;
assert_eq!(
repeated["result"]["protocol_version"],
car_proto::PROTOCOL_VERSION
);
let subscribed = call(
&mut compatible,
"subscribe",
"host.subscribe",
serde_json::json!({}),
)
.await;
assert_eq!(subscribed["result"]["subscribed"], true, "{subscribed}");
}
#[tokio::test]
async fn reconnect_starts_unnegotiated_and_must_handshake_again() {
let journal = TempDir::new().expect("journal tempdir");
let address = spawn_dispatcher(state(journal.path().to_path_buf()), 2).await;
let (mut first, _) = connect_async(format!("ws://{address}"))
.await
.expect("connect first session");
assert!(negotiate(&mut first, "first-handshake")
.await
.get("error")
.is_none());
let first_subscribe = call(
&mut first,
"first-subscribe",
"host.subscribe",
serde_json::json!({}),
)
.await;
assert_eq!(first_subscribe["result"]["subscribed"], true);
first.close(None).await.expect("close first session");
let (mut second, _) = connect_async(format!("ws://{address}"))
.await
.expect("connect second session");
let before_handshake = call(
&mut second,
"second-subscribe-early",
"host.subscribe",
serde_json::json!({}),
)
.await;
assert_handshake_required(&before_handshake, "host.subscribe");
assert!(negotiate(&mut second, "second-handshake")
.await
.get("error")
.is_none());
let after_handshake = call(
&mut second,
"second-subscribe",
"host.subscribe",
serde_json::json!({}),
)
.await;
assert_eq!(after_handshake["result"]["subscribed"], true);
}