feat(server): tracker ws em tempo real com registro e broadcast de peers
This commit is contained in:
Generated
+834
-21
File diff suppressed because it is too large
Load Diff
+5
-1
@@ -5,7 +5,7 @@ edition = "2021"
|
||||
publish = false
|
||||
|
||||
[dependencies]
|
||||
axum = "0.8"
|
||||
axum = { version = "0.8", features = ["ws"] }
|
||||
tokio = { version = "1", features = ["full", "macros", "rt-multi-thread"] }
|
||||
serde = { version = "1", features = ["derive"] }
|
||||
serde_json = "1"
|
||||
@@ -15,7 +15,11 @@ tower-http = { version = "0.6", features = ["cors", "trace"] }
|
||||
rusqlite = { version = "0.32", features = ["bundled"] }
|
||||
thiserror = "2"
|
||||
chrono = { version = "0.4", default-features = false, features = ["clock"] }
|
||||
futures-util = "0.3"
|
||||
tokio-stream = { version = "0.1", features = ["sync"] }
|
||||
|
||||
[dev-dependencies]
|
||||
axum-test = "21"
|
||||
tempfile = "3"
|
||||
tokio-tungstenite = "0.24"
|
||||
reqwest = { version = "0.12", features = ["json"] }
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
pub mod announce;
|
||||
pub mod files;
|
||||
pub mod health;
|
||||
pub mod peers;
|
||||
pub mod search;
|
||||
pub mod ws;
|
||||
|
||||
use axum::{routing::{get, post}, Router};
|
||||
use tower_http::cors::CorsLayer;
|
||||
@@ -13,6 +15,8 @@ pub fn router(state: SharedState) -> Router {
|
||||
.route("/search", get(search::search))
|
||||
.route("/files/{file_id}", get(files::get_file))
|
||||
.route("/announce", post(announce::announce))
|
||||
.route("/peers/{file_id}", get(peers::peers))
|
||||
.route("/tracker/ws", get(ws::ws_handler))
|
||||
.with_state(state);
|
||||
|
||||
Router::new()
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
use axum::extract::{Path, State};
|
||||
use axum::Json;
|
||||
use crate::error::AppResult;
|
||||
use crate::state::SharedState;
|
||||
|
||||
pub async fn peers(
|
||||
State(state): State<SharedState>,
|
||||
Path(file_id): Path<String>,
|
||||
) -> AppResult<Json<serde_json::Value>> {
|
||||
let t = state.tracker.lock().unwrap();
|
||||
let peers = t.peers_of(&file_id);
|
||||
Ok(Json(serde_json::json!({ "file_id": file_id, "peers": peers })))
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade};
|
||||
use axum::extract::State;
|
||||
use axum::response::Response;
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use tokio::sync::mpsc;
|
||||
use tokio_stream::wrappers::UnboundedReceiverStream;
|
||||
use tracing::{debug, warn};
|
||||
|
||||
use crate::state::SharedState;
|
||||
|
||||
#[derive(serde::Deserialize, Debug)]
|
||||
struct RegisterMsg {
|
||||
peer_id: String,
|
||||
file_id: String,
|
||||
}
|
||||
|
||||
pub async fn ws_handler(
|
||||
ws: WebSocketUpgrade,
|
||||
State(state): State<SharedState>,
|
||||
) -> Response {
|
||||
ws.on_upgrade(move |socket| handle_socket(socket, state))
|
||||
}
|
||||
|
||||
async fn handle_socket(socket: WebSocket, state: SharedState) {
|
||||
let (mut sender, mut receiver) = socket.split();
|
||||
let (tx, rx) = mpsc::unbounded_channel::<Message>();
|
||||
let mut recv_stream = UnboundedReceiverStream::new(rx);
|
||||
|
||||
let registration = wait_for_registration(&mut receiver).await;
|
||||
let (file_id, peer_id) = match registration {
|
||||
Some((file_id, peer_id)) => {
|
||||
let mut t = state.tracker.lock().unwrap();
|
||||
t.register(file_id.clone(), peer_id.clone(), tx.clone());
|
||||
(file_id, peer_id)
|
||||
}
|
||||
None => return,
|
||||
};
|
||||
|
||||
// notifica os demais pares do mesmo arquivo (exceto o próprio)
|
||||
{
|
||||
let t = state.tracker.lock().unwrap();
|
||||
t.broadcast_except(
|
||||
&file_id,
|
||||
&peer_id,
|
||||
&format!(
|
||||
r#"{{"type":"peer_online","file_id":"{}","peer_id":"{}"}}"#,
|
||||
file_id, peer_id
|
||||
),
|
||||
);
|
||||
}
|
||||
|
||||
loop {
|
||||
tokio::select! {
|
||||
msg = receiver.next() => {
|
||||
match msg {
|
||||
Some(Ok(Message::Close(_))) | None => break,
|
||||
Some(Ok(Message::Ping(p))) => {
|
||||
if sender.send(Message::Pong(p)).await.is_err() { break; }
|
||||
}
|
||||
Some(Ok(_)) => {} // ignoramos frames não relevantes
|
||||
Some(Err(e)) => {
|
||||
warn!("ws error: {e}");
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
msg = recv_stream.next() => {
|
||||
if let Some(m) = msg {
|
||||
if sender.send(m).await.is_err() { break; }
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let mut t = state.tracker.lock().unwrap();
|
||||
t.unregister(&file_id, &peer_id);
|
||||
debug!("peer {peer_id} removed from {file_id}");
|
||||
}
|
||||
|
||||
/// Awaits the first text frame carrying `peer_id` + `file_id`.
|
||||
async fn wait_for_registration(receiver: &mut futures_util::stream::SplitStream<WebSocket>) -> Option<(String, String)> {
|
||||
loop {
|
||||
match receiver.next().await {
|
||||
Some(Ok(Message::Text(text))) => {
|
||||
if let Ok(reg) = serde_json::from_str::<RegisterMsg>(&text) {
|
||||
return Some((reg.file_id, reg.peer_id));
|
||||
}
|
||||
}
|
||||
Some(Ok(Message::Close(_))) | None => return None,
|
||||
Some(Err(_)) => return None,
|
||||
_ => continue,
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -6,6 +6,7 @@ pub mod search;
|
||||
pub mod state;
|
||||
pub mod store;
|
||||
pub mod validate;
|
||||
pub mod ws;
|
||||
|
||||
use axum::Router;
|
||||
use crate::config::Config;
|
||||
|
||||
@@ -1,12 +1,14 @@
|
||||
use std::sync::{Arc, Mutex};
|
||||
use rusqlite::Connection;
|
||||
use crate::config::Config;
|
||||
use crate::ws::tracker::PeerTracker;
|
||||
|
||||
pub type DbConn = Mutex<Connection>;
|
||||
|
||||
pub struct AppState {
|
||||
pub config: Config,
|
||||
pub db: DbConn,
|
||||
pub tracker: Mutex<PeerTracker>,
|
||||
}
|
||||
|
||||
impl AppState {
|
||||
@@ -14,6 +16,7 @@ impl AppState {
|
||||
Arc::new(Self {
|
||||
config,
|
||||
db: Mutex::new(db),
|
||||
tracker: Mutex::new(PeerTracker::default()),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
pub mod tracker;
|
||||
@@ -0,0 +1,50 @@
|
||||
use std::collections::HashMap;
|
||||
use tokio::sync::mpsc;
|
||||
use axum::extract::ws::Message;
|
||||
|
||||
pub type WsSender = mpsc::UnboundedSender<Message>;
|
||||
|
||||
#[derive(Default)]
|
||||
pub struct PeerTracker {
|
||||
pub by_file: HashMap<String, HashMap<String, WsSender>>,
|
||||
}
|
||||
|
||||
impl PeerTracker {
|
||||
pub fn register(&mut self, file_id: String, peer_id: String, tx: WsSender) {
|
||||
self.by_file.entry(file_id).or_default().insert(peer_id, tx);
|
||||
}
|
||||
|
||||
pub fn unregister(&mut self, file_id: &str, peer_id: &str) {
|
||||
if let Some(peers) = self.by_file.get_mut(file_id) {
|
||||
peers.remove(peer_id);
|
||||
if peers.is_empty() {
|
||||
self.by_file.remove(file_id);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn peers_of(&self, file_id: &str) -> Vec<String> {
|
||||
self.by_file
|
||||
.get(file_id)
|
||||
.map(|peers| peers.keys().cloned().collect())
|
||||
.unwrap_or_default()
|
||||
}
|
||||
|
||||
pub fn online_total(&self) -> usize {
|
||||
self.by_file
|
||||
.values()
|
||||
.flat_map(|peers| peers.keys())
|
||||
.collect::<std::collections::HashSet<_>>()
|
||||
.len()
|
||||
}
|
||||
|
||||
pub fn broadcast_except(&self, file_id: &str, except_peer: &str, payload: &str) {
|
||||
if let Some(peers) = self.by_file.get(file_id) {
|
||||
for (pid, tx) in peers {
|
||||
if pid != except_peer {
|
||||
let _ = tx.send(Message::Text(payload.to_string().into()));
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,72 @@
|
||||
use ares_server::config::Config;
|
||||
use ares_server::{build_app, router};
|
||||
use futures_util::{SinkExt, StreamExt};
|
||||
use serde_json::json;
|
||||
use tempfile::tempdir;
|
||||
|
||||
async fn spawn_server() -> (String, ares_server::state::SharedState) {
|
||||
let dir = tempdir().unwrap();
|
||||
let cfg = Config {
|
||||
db_path: dir.path().join("t.db").to_string_lossy().into(),
|
||||
..Default::default()
|
||||
};
|
||||
let state = build_app(cfg).await;
|
||||
let app = router(state.clone());
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let addr = listener.local_addr().unwrap();
|
||||
tokio::spawn(async move {
|
||||
let _ = axum::serve(listener, app).await;
|
||||
});
|
||||
(format!("http://{addr}"), state)
|
||||
}
|
||||
|
||||
fn ws_url(base: &str) -> String {
|
||||
base.replace("http", "ws")
|
||||
}
|
||||
|
||||
async fn connect_ws(base: &str, peer_id: &str, file_id: &str) -> tokio_tungstenite::WebSocketStream<tokio_tungstenite::MaybeTlsStream<tokio::net::TcpStream>> {
|
||||
let (mut socket, _) =
|
||||
tokio_tungstenite::connect_async(&format!("{}/api/v1/tracker/ws", ws_url(base))).await.unwrap();
|
||||
socket
|
||||
.send(tokio_tungstenite::tungstenite::Message::Text(
|
||||
json!({"peer_id": peer_id, "file_id": file_id}).to_string(),
|
||||
)).await.unwrap();
|
||||
socket
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ws_registers_and_http_peers_lists_it() {
|
||||
let (base, _state) = spawn_server().await;
|
||||
let _socket = connect_ws(&base, "peerA", "abc123").await;
|
||||
|
||||
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
let res = client
|
||||
.get(format!("{base}/api/v1/peers/abc123"))
|
||||
.send().await.unwrap();
|
||||
let body: serde_json::Value = res.json().await.unwrap();
|
||||
assert_eq!(body["file_id"], "abc123");
|
||||
assert!(body["peers"].as_array().unwrap().contains(&serde_json::Value::String("peerA".into())));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn ws_broadcasts_peer_online() {
|
||||
let (base, _state) = spawn_server().await;
|
||||
let mut a = connect_ws(&base, "peerA", "f1").await;
|
||||
tokio::time::sleep(std::time::Duration::from_millis(150)).await;
|
||||
|
||||
let _b = connect_ws(&base, "peerB", "f1").await;
|
||||
|
||||
// peerA deve receber a notificação de que peerB ficou online
|
||||
let timeout = tokio::time::timeout(std::time::Duration::from_secs(2), a.next());
|
||||
let got = timeout.await.unwrap().expect("peerA fechou").expect("erro ws");
|
||||
let text = match got {
|
||||
tokio_tungstenite::tungstenite::Message::Text(t) => t.to_string(),
|
||||
other => panic!("esperado Text, recebido: {other:?}"),
|
||||
};
|
||||
let parsed: serde_json::Value = serde_json::from_str(&text).unwrap();
|
||||
eprintln!("received: {parsed}");
|
||||
assert_eq!(parsed["type"], "peer_online");
|
||||
assert_eq!(parsed["peer_id"], "peerB");
|
||||
}
|
||||
Reference in New Issue
Block a user