mirror of https://github.com/synctv-org/synctv
feat: webrtc
parent
d60946d499
commit
fcd5b6a1b4
File diff suppressed because it is too large
Load Diff
@ -1,472 +0,0 @@
|
||||
//! WebRTC HTTP API endpoints
|
||||
//!
|
||||
//! Provides REST API for WebRTC signaling and session management.
|
||||
|
||||
use axum::{
|
||||
extract::{Path, State},
|
||||
routing::{get, post},
|
||||
Json, Router,
|
||||
};
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::sync::Arc;
|
||||
|
||||
use crate::http::AppError;
|
||||
use synctv_core::{
|
||||
models::UserId,
|
||||
service::webrtc::{
|
||||
SignalingService, MediaType, SessionDescription, IceCandidate, SdpType,
|
||||
},
|
||||
};
|
||||
|
||||
/// Create WebRTC router
|
||||
pub fn create_webrtc_router() -> axum::Router<Arc<super::AppState>> {
|
||||
Router::new()
|
||||
.route("/servers", get(get_ice_servers))
|
||||
.route("/sessions", post(create_session))
|
||||
.route("/sessions/:session_id", get(get_session_info).delete(end_session))
|
||||
.route("/sessions/:session_id/join", post(join_session))
|
||||
.route("/sessions/:session_id/leave", post(leave_session))
|
||||
.route("/sessions/:session_id/offer", post(handle_offer))
|
||||
.route("/sessions/:session_id/answer", post(handle_answer))
|
||||
.route("/sessions/:session_id/ice", post(handle_ice_candidate))
|
||||
}
|
||||
|
||||
/// Get ICE server configuration
|
||||
///
|
||||
/// Returns STUN/TURN server configuration for WebRTC clients.
|
||||
#[utoipa::path(
|
||||
get,
|
||||
path = "/api/webrtc/servers",
|
||||
tag = "webrtc",
|
||||
responses(
|
||||
(status = 200, description = "ICE server configuration", body = IceServersResponse),
|
||||
(status = 500, description = "Internal server error")
|
||||
),
|
||||
security(
|
||||
("bearer_auth" = [])
|
||||
)
|
||||
)]
|
||||
async fn get_ice_servers(State(state): State<Arc<super::AppState>>) -> Result<Json<IceServersResponse>, AppError> {
|
||||
let signaling_service = state
|
||||
.webrtc_service
|
||||
.as_ref()
|
||||
.ok_or_else(|| AppError::internal("WebRTC service not available"))?;
|
||||
|
||||
let ice_servers = signaling_service.get_ice_servers();
|
||||
|
||||
Ok(Json(IceServersResponse {
|
||||
stun_servers: ice_servers.stun_servers,
|
||||
turn_config: ice_servers.turn_config,
|
||||
}))
|
||||
}
|
||||
|
||||
/// Create a new WebRTC session
|
||||
///
|
||||
/// Creates a new WebRTC session (call) for a room.
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/api/webrtc/sessions",
|
||||
tag = "webrtc",
|
||||
request_body = CreateSessionRequest,
|
||||
responses(
|
||||
(status = 200, description = "Session created successfully", body = CreateSessionResponse),
|
||||
(status = 400, description = "Invalid request"),
|
||||
(status = 409, description = "Session already exists for this room"),
|
||||
(status = 500, description = "Internal server error")
|
||||
),
|
||||
security(
|
||||
("bearer_auth" = [])
|
||||
)
|
||||
)]
|
||||
async fn create_session(
|
||||
State(state): State<Arc<super::AppState>>,
|
||||
Json(req): Json<CreateSessionRequest>,
|
||||
auth_user: super::middleware::AuthUser,
|
||||
) -> Result<Json<CreateSessionResponse>, AppError> {
|
||||
let signaling_service = state
|
||||
.webrtc_service
|
||||
.as_ref()
|
||||
.ok_or_else(|| AppError::internal("WebRTC service not available"))?;
|
||||
|
||||
let response = signaling_service
|
||||
.create_session(req.room_id, req.media_type, auth_user.user_id)
|
||||
.await
|
||||
.map_err(|e| AppError::internal(format!("Failed to create session: {}", e)))?;
|
||||
|
||||
Ok(Json(CreateSessionResponse {
|
||||
session_id: response.session_id,
|
||||
ice_servers: IceServersResponse {
|
||||
stun_servers: response.ice_servers.stun_servers,
|
||||
turn_config: response.ice_servers.turn_config,
|
||||
},
|
||||
}))
|
||||
}
|
||||
|
||||
/// Get session information
|
||||
///
|
||||
/// Returns information about a WebRTC session.
|
||||
#[utoipa::path(
|
||||
get,
|
||||
path = "/api/webrtc/sessions/{session_id}",
|
||||
tag = "webrtc",
|
||||
params(
|
||||
("session_id" = String, Path, description = "Session ID")
|
||||
),
|
||||
responses(
|
||||
(status = 200, description = "Session information", body = SessionInfoResponse),
|
||||
(status = 404, description = "Session not found"),
|
||||
(status = 500, description = "Internal server error")
|
||||
),
|
||||
security(
|
||||
("bearer_auth" = [])
|
||||
)
|
||||
)]
|
||||
async fn get_session_info(
|
||||
State(state): State<Arc<super::AppState>>,
|
||||
Path(session_id): Path<String>,
|
||||
) -> Result<Json<SessionInfoResponse>, AppError> {
|
||||
let signaling_service = state
|
||||
.webrtc_service
|
||||
.as_ref()
|
||||
.ok_or_else(|| AppError::internal("WebRTC service not available"))?;
|
||||
|
||||
let session_info = signaling_service
|
||||
.get_session_info(&session_id)
|
||||
.await
|
||||
.map_err(|e| AppError::not_found(format!("Session not found: {}", e)))?;
|
||||
|
||||
Ok(Json(SessionInfoResponse::from(session_info)))
|
||||
}
|
||||
|
||||
/// Join a WebRTC session
|
||||
///
|
||||
/// Join an existing WebRTC session as a participant.
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/api/webrtc/sessions/{session_id}/join",
|
||||
tag = "webrtc",
|
||||
params(
|
||||
("session_id" = String, Path, description = "Session ID")
|
||||
),
|
||||
request_body = JoinSessionRequest,
|
||||
responses(
|
||||
(status = 200, description = "Joined session successfully", body = JoinSessionResponse),
|
||||
(status = 400, description = "Invalid request"),
|
||||
(status = 404, description = "Session not found"),
|
||||
(status = 409, description = "Session is full or user already in session"),
|
||||
(status = 500, description = "Internal server error")
|
||||
),
|
||||
security(
|
||||
("bearer_auth" = [])
|
||||
)
|
||||
)]
|
||||
async fn join_session(
|
||||
State(state): State<Arc<super::AppState>>,
|
||||
Path(session_id): Path<String>,
|
||||
Json(req): Json<JoinSessionRequest>,
|
||||
auth_user: super::middleware::AuthUser,
|
||||
) -> Result<Json<JoinSessionResponse>, AppError> {
|
||||
let signaling_service = state
|
||||
.webrtc_service
|
||||
.as_ref()
|
||||
.ok_or_else(|| AppError::internal("WebRTC service not available"))?;
|
||||
|
||||
let response = signaling_service
|
||||
.join_session(&session_id, auth_user.user_id, req.username)
|
||||
.await
|
||||
.map_err(|e| AppError::bad_request(format!("Failed to join session: {}", e)))?;
|
||||
|
||||
Ok(Json(JoinSessionResponse::from(response)))
|
||||
}
|
||||
|
||||
/// Leave a WebRTC session
|
||||
///
|
||||
/// Leave a WebRTC session.
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/api/webrtc/sessions/{session_id}/leave",
|
||||
tag = "webrtc",
|
||||
params(
|
||||
("session_id" = String, Path, description = "Session ID")
|
||||
),
|
||||
request_body = LeaveSessionRequest,
|
||||
responses(
|
||||
(status = 200, description = "Left session successfully"),
|
||||
(status = 400, description = "Invalid request"),
|
||||
(status = 404, description = "Session not found"),
|
||||
(status = 500, description = "Internal server error")
|
||||
),
|
||||
security(
|
||||
("bearer_auth" = [])
|
||||
)
|
||||
)]
|
||||
async fn leave_session(
|
||||
State(state): State<Arc<super::AppState>>,
|
||||
Path(session_id): Path<String>,
|
||||
Json(req): Json<LeaveSessionRequest>,
|
||||
) -> Result<Json<serde_json::Value>, AppError> {
|
||||
let signaling_service = state
|
||||
.webrtc_service
|
||||
.as_ref()
|
||||
.ok_or_else(|| AppError::internal("WebRTC service not available"))?;
|
||||
|
||||
signaling_service
|
||||
.leave_session(&session_id, &req.peer_id)
|
||||
.await
|
||||
.map_err(|e| AppError::bad_request(format!("Failed to leave session: {}", e)))?;
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"success": true
|
||||
})))
|
||||
}
|
||||
|
||||
/// Handle WebRTC offer
|
||||
///
|
||||
/// Process a WebRTC offer from a peer.
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/api/webrtc/sessions/{session_id}/offer",
|
||||
tag = "webrtc",
|
||||
params(
|
||||
("session_id" = String, Path, description = "Session ID")
|
||||
),
|
||||
request_body = OfferRequest,
|
||||
responses(
|
||||
(status = 200, description = "Offer processed successfully"),
|
||||
(status = 400, description = "Invalid request"),
|
||||
(status = 404, description = "Session not found"),
|
||||
(status = 500, description = "Internal server error")
|
||||
),
|
||||
security(
|
||||
("bearer_auth" = [])
|
||||
)
|
||||
)]
|
||||
async fn handle_offer(
|
||||
State(state): State<Arc<super::AppState>>,
|
||||
Path(session_id): Path<String>,
|
||||
Json(req): Json<OfferRequest>,
|
||||
) -> Result<Json<serde_json::Value>, AppError> {
|
||||
let signaling_service = state
|
||||
.webrtc_service
|
||||
.as_ref()
|
||||
.ok_or_else(|| AppError::internal("WebRTC service not available"))?;
|
||||
|
||||
signaling_service
|
||||
.handle_offer(&session_id, &req.peer_id, req.sdp)
|
||||
.await
|
||||
.map_err(|e| AppError::bad_request(format!("Failed to handle offer: {}", e)))?;
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"success": true
|
||||
})))
|
||||
}
|
||||
|
||||
/// Handle WebRTC answer
|
||||
///
|
||||
/// Process a WebRTC answer from a peer.
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/api/webrtc/sessions/{session_id}/answer",
|
||||
tag = "webrtc",
|
||||
params(
|
||||
("session_id" = String, Path, description = "Session ID")
|
||||
),
|
||||
request_body = AnswerRequest,
|
||||
responses(
|
||||
(status = 200, description = "Answer processed successfully"),
|
||||
(status = 400, description = "Invalid request"),
|
||||
(status = 404, description = "Session not found"),
|
||||
(status = 500, description = "Internal server error")
|
||||
),
|
||||
security(
|
||||
("bearer_auth" = [])
|
||||
)
|
||||
)]
|
||||
async fn handle_answer(
|
||||
State(state): State<Arc<super::AppState>>,
|
||||
Path(session_id): Path<String>,
|
||||
Json(req): Json<AnswerRequest>,
|
||||
) -> Result<Json<serde_json::Value>, AppError> {
|
||||
let signaling_service = state
|
||||
.webrtc_service
|
||||
.as_ref()
|
||||
.ok_or_else(|| AppError::internal("WebRTC service not available"))?;
|
||||
|
||||
signaling_service
|
||||
.handle_answer(&session_id, &req.peer_id, req.sdp)
|
||||
.await
|
||||
.map_err(|e| AppError::bad_request(format!("Failed to handle answer: {}", e)))?;
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"success": true
|
||||
})))
|
||||
}
|
||||
|
||||
/// Handle ICE candidate
|
||||
///
|
||||
/// Process an ICE candidate from a peer.
|
||||
#[utoipa::path(
|
||||
post,
|
||||
path = "/api/webrtc/sessions/{session_id}/ice",
|
||||
tag = "webrtc",
|
||||
params(
|
||||
("session_id" = String, Path, description = "Session ID")
|
||||
),
|
||||
request_body = IceCandidateRequest,
|
||||
responses(
|
||||
(status = 200, description = "ICE candidate processed successfully"),
|
||||
(status = 400, description = "Invalid request"),
|
||||
(status = 404, description = "Session not found"),
|
||||
(status = 500, description = "Internal server error")
|
||||
),
|
||||
security(
|
||||
("bearer_auth" = [])
|
||||
)
|
||||
)]
|
||||
async fn handle_ice_candidate(
|
||||
State(state): State<Arc<super::AppState>>,
|
||||
Path(session_id): Path<String>,
|
||||
Json(req): Json<IceCandidateRequest>,
|
||||
) -> Result<Json<serde_json::Value>, AppError> {
|
||||
let signaling_service = state
|
||||
.webrtc_service
|
||||
.as_ref()
|
||||
.ok_or_else(|| AppError::internal("WebRTC service not available"))?;
|
||||
|
||||
signaling_service
|
||||
.handle_ice_candidate(&session_id, &req.peer_id, req.candidate)
|
||||
.await
|
||||
.map_err(|e| AppError::bad_request(format!("Failed to handle ICE candidate: {}", e)))?;
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"success": true
|
||||
})))
|
||||
}
|
||||
|
||||
/// End a WebRTC session
|
||||
///
|
||||
/// End a WebRTC session and remove all participants.
|
||||
#[utoipa::path(
|
||||
delete,
|
||||
path = "/api/webrtc/sessions/{session_id}",
|
||||
tag = "webrtc",
|
||||
params(
|
||||
("session_id" = String, Path, description = "Session ID")
|
||||
),
|
||||
responses(
|
||||
(status = 200, description = "Session ended successfully"),
|
||||
(status = 404, description = "Session not found"),
|
||||
(status = 500, description = "Internal server error")
|
||||
),
|
||||
security(
|
||||
("bearer_auth" = [])
|
||||
)
|
||||
)]
|
||||
async fn end_session(
|
||||
State(state): State<Arc<super::AppState>>,
|
||||
Path(session_id): Path<String>,
|
||||
) -> Result<Json<serde_json::Value>, AppError> {
|
||||
let signaling_service = state
|
||||
.webrtc_service
|
||||
.as_ref()
|
||||
.ok_or_else(|| AppError::internal("WebRTC service not available"))?;
|
||||
|
||||
signaling_service
|
||||
.end_session(&session_id)
|
||||
.await
|
||||
.map_err(|e| AppError::bad_request(format!("Failed to end session: {}", e)))?;
|
||||
|
||||
Ok(Json(serde_json::json!({
|
||||
"success": true
|
||||
})))
|
||||
}
|
||||
|
||||
// Request/Response types
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct CreateSessionRequest {
|
||||
pub room_id: String,
|
||||
pub media_type: MediaType,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct IceServersResponse {
|
||||
pub stun_servers: Vec<String>,
|
||||
pub turn_config: Option<synctv_core::service::webrtc::TurnConfig>,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct CreateSessionResponse {
|
||||
pub session_id: String,
|
||||
pub ice_servers: IceServersResponse,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct SessionInfoResponse {
|
||||
pub session_id: String,
|
||||
pub room_id: String,
|
||||
pub state: synctv_core::service::webrtc::session::SessionState,
|
||||
pub media_type: MediaType,
|
||||
pub peer_count: usize,
|
||||
pub peers: Vec<synctv_core::service::webrtc::Peer>,
|
||||
pub created_at: chrono::DateTime<chrono::Utc>,
|
||||
}
|
||||
|
||||
impl From<synctv_core::service::webrtc::signaling::SessionInfo> for SessionInfoResponse {
|
||||
fn from(info: synctv_core::service::webrtc::signaling::SessionInfo) -> Self {
|
||||
Self {
|
||||
session_id: info.session_id,
|
||||
room_id: info.room_id,
|
||||
state: info.state,
|
||||
media_type: info.media_type,
|
||||
peer_count: info.peer_count,
|
||||
peers: info.peers,
|
||||
created_at: info.created_at,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct JoinSessionRequest {
|
||||
pub username: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Serialize)]
|
||||
pub struct JoinSessionResponse {
|
||||
pub peer_id: String,
|
||||
pub peers: Vec<synctv_core::service::webrtc::Peer>,
|
||||
pub session_state: synctv_core::service::webrtc::session::SessionState,
|
||||
}
|
||||
|
||||
impl From<synctv_core::service::webrtc::signaling::JoinSessionResponse> for JoinSessionResponse {
|
||||
fn from(resp: synctv_core::service::webrtc::signaling::JoinSessionResponse) -> Self {
|
||||
Self {
|
||||
peer_id: resp.peer_id,
|
||||
peers: resp.peers,
|
||||
session_state: resp.session_state,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct LeaveSessionRequest {
|
||||
pub peer_id: String,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct OfferRequest {
|
||||
pub peer_id: String,
|
||||
pub sdp: SessionDescription,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct AnswerRequest {
|
||||
pub peer_id: String,
|
||||
pub sdp: SessionDescription,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
pub struct IceCandidateRequest {
|
||||
pub peer_id: String,
|
||||
pub candidate: IceCandidate,
|
||||
}
|
||||
@ -0,0 +1,377 @@
|
||||
//! Built-in STUN Server
|
||||
//!
|
||||
//! A lightweight STUN (Session Traversal Utilities for NAT) server implementation.
|
||||
//! Helps WebRTC clients discover their public IP addresses and ports for P2P connectivity.
|
||||
//!
|
||||
//! ## STUN Protocol Overview
|
||||
//! - RFC 8489: Session Traversal Utilities for NAT (STUN)
|
||||
//! - Binding Request: Client asks "what's my public IP:port?"
|
||||
//! - Binding Response: Server responds with XOR-MAPPED-ADDRESS
|
||||
//! - Runs on UDP port 3478 (default)
|
||||
//!
|
||||
//! ## Implementation
|
||||
//! Uses the mature `stun_codec` crate for protocol handling, avoiding manual byte
|
||||
//! manipulation and reducing the risk of protocol errors.
|
||||
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::Arc;
|
||||
use tokio::net::UdpSocket;
|
||||
use tracing::{debug, error, info, warn};
|
||||
|
||||
use bytecodec::{DecodeExt, EncodeExt};
|
||||
use stun_codec::{Message, MessageClass, MessageDecoder, MessageEncoder, TransactionId};
|
||||
use stun_codec::rfc5389::attributes::{Software, XorMappedAddress};
|
||||
use stun_codec::rfc5389::{Attribute, methods};
|
||||
|
||||
// Convenience constant for BINDING method
|
||||
const BINDING_METHOD: stun_codec::Method = methods::BINDING;
|
||||
|
||||
/// STUN server configuration
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct StunServerConfig {
|
||||
/// Bind address (e.g., "0.0.0.0:3478")
|
||||
pub bind_addr: String,
|
||||
/// Maximum UDP packet size (typically 1500 bytes for MTU)
|
||||
pub max_packet_size: usize,
|
||||
}
|
||||
|
||||
impl Default for StunServerConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
bind_addr: "0.0.0.0:3478".to_string(),
|
||||
max_packet_size: 1500,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// STUN server metrics
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct StunMetrics {
|
||||
/// Total requests received
|
||||
pub total_requests: u64,
|
||||
/// Total responses sent
|
||||
pub total_responses: u64,
|
||||
/// Total errors
|
||||
pub total_errors: u64,
|
||||
}
|
||||
|
||||
/// Built-in STUN server for NAT traversal
|
||||
pub struct StunServer {
|
||||
config: StunServerConfig,
|
||||
socket: Arc<UdpSocket>,
|
||||
metrics: Arc<StunMetricsInner>,
|
||||
}
|
||||
|
||||
struct StunMetricsInner {
|
||||
total_requests: AtomicU64,
|
||||
total_responses: AtomicU64,
|
||||
total_errors: AtomicU64,
|
||||
}
|
||||
|
||||
impl StunServer {
|
||||
/// Create and start a new STUN server
|
||||
pub async fn start(config: StunServerConfig) -> anyhow::Result<Arc<Self>> {
|
||||
let socket = UdpSocket::bind(&config.bind_addr).await?;
|
||||
let local_addr = socket.local_addr()?;
|
||||
|
||||
info!(
|
||||
bind_addr = %local_addr,
|
||||
"STUN server started"
|
||||
);
|
||||
|
||||
let server = Arc::new(Self {
|
||||
config,
|
||||
socket: Arc::new(socket),
|
||||
metrics: Arc::new(StunMetricsInner {
|
||||
total_requests: AtomicU64::new(0),
|
||||
total_responses: AtomicU64::new(0),
|
||||
total_errors: AtomicU64::new(0),
|
||||
}),
|
||||
});
|
||||
|
||||
// Spawn background task to handle requests
|
||||
let server_clone = Arc::clone(&server);
|
||||
tokio::spawn(async move {
|
||||
server_clone.run().await;
|
||||
});
|
||||
|
||||
Ok(server)
|
||||
}
|
||||
|
||||
/// Main server loop
|
||||
async fn run(&self) {
|
||||
let mut buf = vec![0u8; self.config.max_packet_size];
|
||||
|
||||
loop {
|
||||
match self.socket.recv_from(&mut buf).await {
|
||||
Ok((len, peer_addr)) => {
|
||||
self.metrics.total_requests.fetch_add(1, Ordering::Relaxed);
|
||||
|
||||
debug!(
|
||||
peer_addr = %peer_addr,
|
||||
len = len,
|
||||
"Received STUN request"
|
||||
);
|
||||
|
||||
// Handle request in background to avoid blocking
|
||||
let data = buf[..len].to_vec();
|
||||
let socket = Arc::clone(&self.socket);
|
||||
let metrics = Arc::clone(&self.metrics);
|
||||
|
||||
tokio::spawn(async move {
|
||||
if let Err(e) = Self::handle_request(&socket, &data, peer_addr, &metrics).await {
|
||||
error!(
|
||||
peer_addr = %peer_addr,
|
||||
error = %e,
|
||||
"Failed to handle STUN request"
|
||||
);
|
||||
metrics.total_errors.fetch_add(1, Ordering::Relaxed);
|
||||
} else {
|
||||
metrics.total_responses.fetch_add(1, Ordering::Relaxed);
|
||||
}
|
||||
});
|
||||
}
|
||||
Err(e) => {
|
||||
error!(error = %e, "Failed to receive UDP packet");
|
||||
self.metrics.total_errors.fetch_add(1, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Handle a single STUN request using the stun_codec crate
|
||||
async fn handle_request(
|
||||
socket: &UdpSocket,
|
||||
data: &[u8],
|
||||
peer_addr: SocketAddr,
|
||||
_metrics: &StunMetricsInner,
|
||||
) -> anyhow::Result<()> {
|
||||
// Decode STUN message using stun_codec
|
||||
let mut decoder = MessageDecoder::<Attribute>::new();
|
||||
let decoded = decoder.decode_from_bytes(data)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to decode STUN message: {e}"))?;
|
||||
|
||||
// Handle potential broken message
|
||||
let request = match decoded {
|
||||
Ok(msg) => msg,
|
||||
Err(broken) => {
|
||||
warn!(
|
||||
peer_addr = %peer_addr,
|
||||
"Received broken STUN message: {:?}", broken
|
||||
);
|
||||
return Err(anyhow::anyhow!("Broken STUN message"));
|
||||
}
|
||||
};
|
||||
|
||||
// Only handle Binding Requests
|
||||
if request.method() != BINDING_METHOD || request.class() != MessageClass::Request {
|
||||
debug!(
|
||||
peer_addr = %peer_addr,
|
||||
method = ?request.method(),
|
||||
class = ?request.class(),
|
||||
"Ignoring non-Binding STUN request"
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
// Build Binding Success Response
|
||||
let response = Self::build_binding_response(&request, peer_addr)?;
|
||||
|
||||
// Encode response
|
||||
let mut encoder = MessageEncoder::new();
|
||||
let response_bytes = encoder.encode_into_bytes(response)
|
||||
.map_err(|e| anyhow::anyhow!("Failed to encode STUN response: {e}"))?;
|
||||
|
||||
// Send response
|
||||
socket.send_to(&response_bytes, peer_addr).await?;
|
||||
|
||||
debug!(
|
||||
peer_addr = %peer_addr,
|
||||
response_len = response_bytes.len(),
|
||||
"Sent STUN Binding Response"
|
||||
);
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Build STUN Binding Success Response with XOR-MAPPED-ADDRESS
|
||||
fn build_binding_response(
|
||||
request: &Message<Attribute>,
|
||||
peer_addr: SocketAddr,
|
||||
) -> anyhow::Result<Message<Attribute>> {
|
||||
// Create response message
|
||||
let mut response = Message::new(
|
||||
MessageClass::SuccessResponse,
|
||||
BINDING_METHOD,
|
||||
request.transaction_id(),
|
||||
);
|
||||
|
||||
// Add XOR-MAPPED-ADDRESS attribute (RFC 5389 Section 15.2)
|
||||
// This tells the client their public IP:port as seen by the server
|
||||
response.add_attribute(Attribute::XorMappedAddress(XorMappedAddress::new(peer_addr)));
|
||||
|
||||
// Add SOFTWARE attribute (optional but recommended)
|
||||
response.add_attribute(Attribute::Software(Software::new(
|
||||
"SyncTV STUN Server v1.0".to_string()
|
||||
)?));
|
||||
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
/// Get current metrics
|
||||
pub fn metrics(&self) -> StunMetrics {
|
||||
StunMetrics {
|
||||
total_requests: self.metrics.total_requests.load(Ordering::Relaxed),
|
||||
total_responses: self.metrics.total_responses.load(Ordering::Relaxed),
|
||||
total_errors: self.metrics.total_errors.load(Ordering::Relaxed),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the local bind address
|
||||
pub fn local_addr(&self) -> anyhow::Result<SocketAddr> {
|
||||
Ok(self.socket.local_addr()?)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_stun_server_start() {
|
||||
let config = StunServerConfig {
|
||||
bind_addr: "127.0.0.1:0".to_string(), // Use random port
|
||||
max_packet_size: 1500,
|
||||
};
|
||||
|
||||
let server = StunServer::start(config).await.unwrap();
|
||||
let addr = server.local_addr().unwrap();
|
||||
|
||||
assert!(addr.port() > 0);
|
||||
|
||||
// Give server time to initialize
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_stun_binding_request() {
|
||||
// Start server on random port
|
||||
let config = StunServerConfig {
|
||||
bind_addr: "127.0.0.1:0".to_string(),
|
||||
max_packet_size: 1500,
|
||||
};
|
||||
|
||||
let server = StunServer::start(config).await.unwrap();
|
||||
let server_addr = server.local_addr().unwrap();
|
||||
|
||||
// Give server time to start
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||
|
||||
// Create client socket
|
||||
let client = UdpSocket::bind("127.0.0.1:0").await.unwrap();
|
||||
|
||||
// Create STUN Binding Request
|
||||
let transaction_id = TransactionId::new([0u8; 12]);
|
||||
let request = Message::<Attribute>::new(
|
||||
MessageClass::Request,
|
||||
BINDING_METHOD,
|
||||
transaction_id,
|
||||
);
|
||||
|
||||
// Encode request
|
||||
let mut encoder = MessageEncoder::new();
|
||||
let request_bytes = encoder.encode_into_bytes(request.clone()).unwrap();
|
||||
|
||||
// Send request
|
||||
client.send_to(&request_bytes, server_addr).await.unwrap();
|
||||
|
||||
// Receive response with timeout
|
||||
let mut buf = vec![0u8; 1500];
|
||||
let (len, _) = tokio::time::timeout(
|
||||
tokio::time::Duration::from_secs(2),
|
||||
client.recv_from(&mut buf),
|
||||
)
|
||||
.await
|
||||
.expect("Timeout waiting for response")
|
||||
.unwrap();
|
||||
|
||||
// Decode response
|
||||
let mut decoder = MessageDecoder::<Attribute>::new();
|
||||
let response = decoder.decode_from_bytes(&buf[..len]).unwrap();
|
||||
|
||||
// Verify response
|
||||
assert_eq!(response.class(), MessageClass::SuccessResponse);
|
||||
assert_eq!(response.method(), BINDING_METHOD);
|
||||
assert_eq!(response.transaction_id(), transaction_id);
|
||||
|
||||
// Verify XOR-MAPPED-ADDRESS is present
|
||||
let has_xor_mapped = response.attributes().iter().any(|attr| {
|
||||
matches!(attr, Attribute::XorMappedAddress(_))
|
||||
});
|
||||
assert!(has_xor_mapped, "Response should contain XOR-MAPPED-ADDRESS");
|
||||
|
||||
// Check metrics
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||
let metrics = server.metrics();
|
||||
assert!(metrics.total_requests >= 1);
|
||||
assert!(metrics.total_responses >= 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_build_binding_response() {
|
||||
let transaction_id = TransactionId::new([0u8; 12]);
|
||||
let request = Message::<Attribute>::new(
|
||||
MessageClass::Request,
|
||||
BINDING_METHOD,
|
||||
transaction_id,
|
||||
);
|
||||
|
||||
let peer_addr: SocketAddr = "192.168.1.100:12345".parse().unwrap();
|
||||
|
||||
let response = StunServer::build_binding_response(&request, peer_addr).unwrap();
|
||||
|
||||
// Verify response properties
|
||||
assert_eq!(response.class(), MessageClass::SuccessResponse);
|
||||
assert_eq!(response.method(), BINDING_METHOD);
|
||||
assert_eq!(response.transaction_id(), transaction_id);
|
||||
|
||||
// Verify XOR-MAPPED-ADDRESS attribute
|
||||
let xor_mapped = response
|
||||
.attributes()
|
||||
.iter()
|
||||
.find_map(|attr| {
|
||||
if let Attribute::XorMappedAddress(addr) = attr {
|
||||
Some(addr)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
})
|
||||
.expect("Response should contain XOR-MAPPED-ADDRESS");
|
||||
|
||||
assert_eq!(xor_mapped.address(), peer_addr);
|
||||
|
||||
// Verify SOFTWARE attribute
|
||||
let has_software = response
|
||||
.attributes()
|
||||
.iter()
|
||||
.any(|attr| matches!(attr, Attribute::Software(_)));
|
||||
assert!(has_software, "Response should contain SOFTWARE attribute");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_metrics() {
|
||||
let config = StunServerConfig {
|
||||
bind_addr: "127.0.0.1:0".to_string(),
|
||||
max_packet_size: 1500,
|
||||
};
|
||||
|
||||
let server = StunServer::start(config).await.unwrap();
|
||||
|
||||
// Initial metrics should be zero
|
||||
let metrics = server.metrics();
|
||||
assert_eq!(metrics.total_requests, 0);
|
||||
assert_eq!(metrics.total_responses, 0);
|
||||
assert_eq!(metrics.total_errors, 0);
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,433 @@
|
||||
//! TURN Server Integration
|
||||
//!
|
||||
//! Provides integration with external TURN (Traversal Using Relays around NAT) servers
|
||||
//! for WebRTC connectivity in challenging network environments.
|
||||
//!
|
||||
//! ## TURN Overview
|
||||
//! - Used when P2P connection fails (Symmetric NAT scenarios)
|
||||
//! - Server acts as relay, forwarding media between peers
|
||||
//! - Required for ~25-30% of connections
|
||||
//! - Higher cost than STUN (relays all media traffic)
|
||||
//!
|
||||
//! ## Coturn Integration
|
||||
//! This module is designed to work with coturn (https://github.com/coturn/coturn),
|
||||
//! the most widely deployed open-source TURN server.
|
||||
//!
|
||||
//! ## Credential Generation
|
||||
//! - Uses RFC 5389 long-term credentials
|
||||
//! - HMAC-SHA1 based on shared secret
|
||||
//! - Time-limited credentials (default 24 hours)
|
||||
//! - Compatible with coturn's `static-auth-secret` mode
|
||||
|
||||
use chrono::{DateTime, Utc};
|
||||
use hmac::{Hmac, Mac};
|
||||
use sha1::Sha1;
|
||||
use std::time::Duration;
|
||||
use base64::Engine;
|
||||
|
||||
/// TURN server configuration
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct TurnConfig {
|
||||
/// TURN server URL (e.g., "turn:turn.example.com:3478")
|
||||
pub server_url: String,
|
||||
|
||||
/// Static auth secret (must match coturn's configuration)
|
||||
pub static_secret: String,
|
||||
|
||||
/// Credential time-to-live (default: 24 hours)
|
||||
pub credential_ttl: Duration,
|
||||
|
||||
/// Whether to use TLS/DTLS (turns: or turn: with ?transport=tcp)
|
||||
pub use_tls: bool,
|
||||
}
|
||||
|
||||
impl Default for TurnConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
server_url: String::new(),
|
||||
static_secret: String::new(),
|
||||
credential_ttl: Duration::from_secs(86400), // 24 hours
|
||||
use_tls: false,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// TURN credentials (username and password)
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct TurnCredential {
|
||||
/// Username in format: "<timestamp>:<user_identifier>"
|
||||
pub username: String,
|
||||
|
||||
/// HMAC-SHA1 based password
|
||||
pub password: String,
|
||||
|
||||
/// Credential expiry time
|
||||
pub expires_at: DateTime<Utc>,
|
||||
}
|
||||
|
||||
/// TURN credential generation service
|
||||
#[derive(Clone)]
|
||||
pub struct TurnCredentialService {
|
||||
config: TurnConfig,
|
||||
}
|
||||
|
||||
impl TurnCredentialService {
|
||||
/// Create a new TURN credential service
|
||||
pub fn new(config: TurnConfig) -> Self {
|
||||
Self { config }
|
||||
}
|
||||
|
||||
/// Generate time-limited TURN credentials for a user
|
||||
///
|
||||
/// Credentials format (RFC 5389 long-term credentials):
|
||||
/// - Username: `<expiry_timestamp>:<user_id>`
|
||||
/// - Password: base64(HMAC-SHA1(secret, username))
|
||||
///
|
||||
/// This format is compatible with coturn's `static-auth-secret` mode.
|
||||
pub fn generate_credential(&self, user_id: &str) -> anyhow::Result<TurnCredential> {
|
||||
// Calculate expiry timestamp
|
||||
let now = Utc::now();
|
||||
let expires_at = now + chrono::Duration::from_std(self.config.credential_ttl)?;
|
||||
let expiry_timestamp = expires_at.timestamp();
|
||||
|
||||
// Format: "<timestamp>:<user_id>"
|
||||
let username = format!("{}:{}", expiry_timestamp, user_id);
|
||||
|
||||
// Generate HMAC-SHA1 password
|
||||
let password = self.compute_hmac(&username)?;
|
||||
|
||||
Ok(TurnCredential {
|
||||
username,
|
||||
password,
|
||||
expires_at,
|
||||
})
|
||||
}
|
||||
|
||||
/// Compute HMAC-SHA1 for credential generation
|
||||
fn compute_hmac(&self, username: &str) -> anyhow::Result<String> {
|
||||
let mut mac = Hmac::<Sha1>::new_from_slice(self.config.static_secret.as_bytes())
|
||||
.map_err(|e| anyhow::anyhow!("Failed to create HMAC: {e}"))?;
|
||||
|
||||
mac.update(username.as_bytes());
|
||||
let result = mac.finalize();
|
||||
let credential = base64::engine::general_purpose::STANDARD.encode(result.into_bytes());
|
||||
|
||||
Ok(credential)
|
||||
}
|
||||
|
||||
/// Verify if a credential is still valid
|
||||
pub fn is_credential_valid(&self, credential: &TurnCredential) -> bool {
|
||||
Utc::now() < credential.expires_at
|
||||
}
|
||||
|
||||
/// Get TURN server URLs
|
||||
pub fn get_urls(&self) -> Vec<String> {
|
||||
let mut urls = vec![self.config.server_url.clone()];
|
||||
|
||||
// Add TLS variant if enabled
|
||||
if self.config.use_tls {
|
||||
let tls_url = self.config.server_url.replace("turn:", "turns:");
|
||||
if tls_url != self.config.server_url {
|
||||
urls.push(tls_url);
|
||||
}
|
||||
}
|
||||
|
||||
urls
|
||||
}
|
||||
|
||||
/// Validate TURN configuration
|
||||
pub fn validate_config(&self) -> anyhow::Result<()> {
|
||||
if self.config.server_url.is_empty() {
|
||||
return Err(anyhow::anyhow!("TURN server URL is empty"));
|
||||
}
|
||||
|
||||
if !self.config.server_url.starts_with("turn:") && !self.config.server_url.starts_with("turns:") {
|
||||
return Err(anyhow::anyhow!("TURN server URL must start with 'turn:' or 'turns:'"));
|
||||
}
|
||||
|
||||
if self.config.static_secret.is_empty() {
|
||||
return Err(anyhow::anyhow!("TURN static secret is empty"));
|
||||
}
|
||||
|
||||
if self.config.static_secret.len() < 16 {
|
||||
return Err(anyhow::anyhow!("TURN static secret should be at least 16 characters"));
|
||||
}
|
||||
|
||||
if self.config.credential_ttl.as_secs() < 60 {
|
||||
return Err(anyhow::anyhow!("TURN credential TTL should be at least 60 seconds"));
|
||||
}
|
||||
|
||||
if self.config.credential_ttl.as_secs() > 86400 * 7 {
|
||||
return Err(anyhow::anyhow!("TURN credential TTL should not exceed 7 days"));
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
/// TURN server deployment guide
|
||||
pub const COTURN_DEPLOYMENT_GUIDE: &str = r#"
|
||||
# Coturn Deployment Guide for SyncTV
|
||||
|
||||
## Installation
|
||||
|
||||
### Ubuntu/Debian:
|
||||
```bash
|
||||
sudo apt-get update
|
||||
sudo apt-get install coturn
|
||||
```
|
||||
|
||||
### CentOS/RHEL:
|
||||
```bash
|
||||
sudo yum install coturn
|
||||
```
|
||||
|
||||
### Docker:
|
||||
```bash
|
||||
docker pull coturn/coturn
|
||||
```
|
||||
|
||||
## Configuration
|
||||
|
||||
Edit `/etc/turnserver.conf`:
|
||||
|
||||
```conf
|
||||
# Listening IP (use 0.0.0.0 for all interfaces)
|
||||
listening-ip=0.0.0.0
|
||||
|
||||
# External IP (your server's public IP)
|
||||
external-ip=YOUR_PUBLIC_IP
|
||||
|
||||
# Listening ports
|
||||
listening-port=3478
|
||||
tls-listening-port=5349
|
||||
|
||||
# Relay IP range
|
||||
min-port=49152
|
||||
max-port=65535
|
||||
|
||||
# Authentication
|
||||
use-auth-secret
|
||||
static-auth-secret=YOUR_SECRET_HERE # Must match WebRTCConfig.turn_static_secret
|
||||
|
||||
# Realm (can be your domain)
|
||||
realm=turn.example.com
|
||||
|
||||
# Logging
|
||||
log-file=/var/log/coturn/turnserver.log
|
||||
verbose
|
||||
|
||||
# Security
|
||||
no-multicast-peers
|
||||
no-loopback-peers
|
||||
|
||||
# Performance
|
||||
total-quota=100
|
||||
bps-capacity=0
|
||||
|
||||
# TLS/DTLS (optional, for turns: protocol)
|
||||
cert=/etc/letsencrypt/live/turn.example.com/cert.pem
|
||||
pkey=/etc/letsencrypt/live/turn.example.com/privkey.pem
|
||||
```
|
||||
|
||||
## Start Service
|
||||
|
||||
```bash
|
||||
# Enable on boot
|
||||
sudo systemctl enable coturn
|
||||
|
||||
# Start service
|
||||
sudo systemctl start coturn
|
||||
|
||||
# Check status
|
||||
sudo systemctl status coturn
|
||||
```
|
||||
|
||||
## Firewall Rules
|
||||
|
||||
```bash
|
||||
# UDP/TCP for TURN
|
||||
sudo ufw allow 3478/tcp
|
||||
sudo ufw allow 3478/udp
|
||||
|
||||
# TLS/DTLS for TURNS
|
||||
sudo ufw allow 5349/tcp
|
||||
sudo ufw allow 5349/udp
|
||||
|
||||
# Media relay ports
|
||||
sudo ufw allow 49152:65535/tcp
|
||||
sudo ufw allow 49152:65535/udp
|
||||
```
|
||||
|
||||
## SyncTV Configuration
|
||||
|
||||
In `config.yaml`:
|
||||
|
||||
```yaml
|
||||
webrtc:
|
||||
mode: peer_to_peer # or hybrid
|
||||
enable_turn: true
|
||||
turn_server_url: "turn:turn.example.com:3478"
|
||||
turn_static_secret: "YOUR_SECRET_HERE" # Must match coturn config
|
||||
turn_credential_ttl: 86400 # 24 hours
|
||||
```
|
||||
|
||||
## Testing
|
||||
|
||||
Test with Trickle ICE:
|
||||
https://webrtc.github.io/samples/src/content/peerconnection/trickle-ice/
|
||||
|
||||
Enter your TURN server URL and credentials to verify connectivity.
|
||||
|
||||
## Monitoring
|
||||
|
||||
```bash
|
||||
# View logs
|
||||
sudo tail -f /var/log/coturn/turnserver.log
|
||||
|
||||
# Check connections
|
||||
sudo turnutils_uclient -v turn.example.com
|
||||
|
||||
# Monitor with prometheus
|
||||
# Coturn supports prometheus metrics on port 9641
|
||||
```
|
||||
|
||||
## Cost Estimation
|
||||
|
||||
- Small deployment (< 100 users): ~$20-50/month
|
||||
- Medium deployment (100-1000 users): ~$100-300/month
|
||||
- Large deployment (1000+ users): ~$500+/month
|
||||
|
||||
Most traffic will still use P2P (STUN), TURN is fallback only (~25-30% of connections).
|
||||
"#;
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_generate_credential() {
|
||||
let config = TurnConfig {
|
||||
server_url: "turn:turn.example.com:3478".to_string(),
|
||||
static_secret: "test_secret_key_12345".to_string(),
|
||||
credential_ttl: Duration::from_secs(3600),
|
||||
use_tls: false,
|
||||
};
|
||||
|
||||
let service = TurnCredentialService::new(config);
|
||||
let credential = service.generate_credential("user123").unwrap();
|
||||
|
||||
// Username should be in format: "<timestamp>:<user_id>"
|
||||
assert!(credential.username.contains(":user123"));
|
||||
|
||||
// Password should be base64 encoded
|
||||
assert!(!credential.password.is_empty());
|
||||
assert!(base64::engine::general_purpose::STANDARD.decode(&credential.password).is_ok());
|
||||
|
||||
// Expiry should be in the future
|
||||
assert!(credential.expires_at > Utc::now());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_credential_validation() {
|
||||
let config = TurnConfig {
|
||||
server_url: "turn:turn.example.com:3478".to_string(),
|
||||
static_secret: "test_secret_key_12345".to_string(),
|
||||
credential_ttl: Duration::from_secs(3600),
|
||||
use_tls: false,
|
||||
};
|
||||
|
||||
let service = TurnCredentialService::new(config);
|
||||
let credential = service.generate_credential("user123").unwrap();
|
||||
|
||||
// Fresh credential should be valid
|
||||
assert!(service.is_credential_valid(&credential));
|
||||
|
||||
// Expired credential should be invalid
|
||||
let expired_credential = TurnCredential {
|
||||
username: credential.username.clone(),
|
||||
password: credential.password.clone(),
|
||||
expires_at: Utc::now() - chrono::Duration::hours(1),
|
||||
};
|
||||
assert!(!service.is_credential_valid(&expired_credential));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_get_urls() {
|
||||
let config = TurnConfig {
|
||||
server_url: "turn:turn.example.com:3478".to_string(),
|
||||
static_secret: "test_secret".to_string(),
|
||||
credential_ttl: Duration::from_secs(3600),
|
||||
use_tls: true,
|
||||
};
|
||||
|
||||
let service = TurnCredentialService::new(config);
|
||||
let urls = service.get_urls();
|
||||
|
||||
assert_eq!(urls.len(), 2);
|
||||
assert!(urls.contains(&"turn:turn.example.com:3478".to_string()));
|
||||
assert!(urls.contains(&"turns:turn.example.com:3478".to_string()));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_validate_config() {
|
||||
// Valid config
|
||||
let config = TurnConfig {
|
||||
server_url: "turn:turn.example.com:3478".to_string(),
|
||||
static_secret: "test_secret_key_12345".to_string(),
|
||||
credential_ttl: Duration::from_secs(3600),
|
||||
use_tls: false,
|
||||
};
|
||||
let service = TurnCredentialService::new(config);
|
||||
assert!(service.validate_config().is_ok());
|
||||
|
||||
// Invalid: empty URL
|
||||
let config = TurnConfig {
|
||||
server_url: String::new(),
|
||||
static_secret: "test_secret".to_string(),
|
||||
credential_ttl: Duration::from_secs(3600),
|
||||
use_tls: false,
|
||||
};
|
||||
let service = TurnCredentialService::new(config);
|
||||
assert!(service.validate_config().is_err());
|
||||
|
||||
// Invalid: short secret
|
||||
let config = TurnConfig {
|
||||
server_url: "turn:turn.example.com:3478".to_string(),
|
||||
static_secret: "short".to_string(),
|
||||
credential_ttl: Duration::from_secs(3600),
|
||||
use_tls: false,
|
||||
};
|
||||
let service = TurnCredentialService::new(config);
|
||||
assert!(service.validate_config().is_err());
|
||||
|
||||
// Invalid: TTL too short
|
||||
let config = TurnConfig {
|
||||
server_url: "turn:turn.example.com:3478".to_string(),
|
||||
static_secret: "test_secret_key".to_string(),
|
||||
credential_ttl: Duration::from_secs(30),
|
||||
use_tls: false,
|
||||
};
|
||||
let service = TurnCredentialService::new(config);
|
||||
assert!(service.validate_config().is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_hmac_deterministic() {
|
||||
let config = TurnConfig {
|
||||
server_url: "turn:turn.example.com:3478".to_string(),
|
||||
static_secret: "test_secret_key".to_string(),
|
||||
credential_ttl: Duration::from_secs(3600),
|
||||
use_tls: false,
|
||||
};
|
||||
|
||||
let service = TurnCredentialService::new(config);
|
||||
|
||||
let username = "12345:user123";
|
||||
let hmac1 = service.compute_hmac(username).unwrap();
|
||||
let hmac2 = service.compute_hmac(username).unwrap();
|
||||
|
||||
// HMAC should be deterministic
|
||||
assert_eq!(hmac1, hmac2);
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,248 @@
|
||||
//! Built-in TURN Server
|
||||
//!
|
||||
//! A simplified TURN (Traversal Using Relays around NAT) server implementation.
|
||||
//! Provides basic media relay functionality for WebRTC connections when direct P2P fails.
|
||||
//!
|
||||
//! ## Important Note
|
||||
//! This is a **simplified implementation** suitable for small to medium deployments.
|
||||
//! For production scale (>1000 concurrent users) or enterprise deployments,
|
||||
//! we strongly recommend using external coturn server instead.
|
||||
//!
|
||||
//! ## Current Limitations
|
||||
//! - Simplified TURN protocol implementation
|
||||
//! - Basic UDP relay only (no TCP relay)
|
||||
//! - No TLS/DTLS support
|
||||
//! - Limited to ~100 concurrent allocations by default
|
||||
//!
|
||||
//! ## When to Use Built-in TURN
|
||||
//! - Small deployments (<100 users)
|
||||
//! - Development and testing
|
||||
//! - Simple deployments where external coturn is not desired
|
||||
//!
|
||||
//! ## When to Use External TURN (coturn)
|
||||
//! - Production deployments (>100 users)
|
||||
//! - Enterprise scale
|
||||
//! - Advanced features (TCP relay, TLS, high availability)
|
||||
//! - See docs/TURN_DEPLOYMENT.md for coturn setup guide
|
||||
|
||||
use std::net::SocketAddr;
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::Arc;
|
||||
use tokio::net::UdpSocket;
|
||||
use tracing::{debug, error, info};
|
||||
|
||||
/// TURN server configuration
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct TurnServerConfig {
|
||||
/// Bind address for TURN (e.g., "0.0.0.0:3478")
|
||||
pub bind_addr: String,
|
||||
/// Relay port range (min)
|
||||
pub relay_min_port: u16,
|
||||
/// Relay port range (max)
|
||||
pub relay_max_port: u16,
|
||||
/// Maximum concurrent allocations
|
||||
pub max_allocations: usize,
|
||||
/// Default allocation lifetime (seconds)
|
||||
pub default_lifetime: u32,
|
||||
/// Maximum allocation lifetime (seconds)
|
||||
pub max_lifetime: u32,
|
||||
/// Static secret for authentication (must match client config)
|
||||
pub static_secret: String,
|
||||
/// Realm for authentication
|
||||
pub realm: String,
|
||||
}
|
||||
|
||||
impl Default for TurnServerConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
bind_addr: "0.0.0.0:3478".to_string(),
|
||||
relay_min_port: 49152,
|
||||
relay_max_port: 65535,
|
||||
max_allocations: 100,
|
||||
default_lifetime: 600, // 10 minutes
|
||||
max_lifetime: 3600, // 1 hour
|
||||
static_secret: String::new(),
|
||||
realm: "synctv.local".to_string(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// TURN server metrics
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct TurnMetrics {
|
||||
/// Total requests received
|
||||
pub total_allocations: u64,
|
||||
/// Total refreshes
|
||||
pub total_refreshes: u64,
|
||||
/// Total sends
|
||||
pub total_sends: u64,
|
||||
/// Total data indications
|
||||
pub total_data: u64,
|
||||
/// Current active allocations
|
||||
pub active_allocations: usize,
|
||||
/// Total errors
|
||||
pub total_errors: u64,
|
||||
/// Total bytes relayed
|
||||
pub total_bytes_relayed: u64,
|
||||
}
|
||||
|
||||
/// Built-in TURN server for NAT traversal relay
|
||||
///
|
||||
/// **Note**: This is a simplified implementation. For production scale,
|
||||
/// consider using external coturn server (see docs/TURN_DEPLOYMENT.md)
|
||||
pub struct TurnServer {
|
||||
config: TurnServerConfig,
|
||||
socket: Arc<UdpSocket>,
|
||||
metrics: Arc<TurnMetricsInner>,
|
||||
}
|
||||
|
||||
struct TurnMetricsInner {
|
||||
total_allocations: AtomicU64,
|
||||
total_refreshes: AtomicU64,
|
||||
total_sends: AtomicU64,
|
||||
total_data: AtomicU64,
|
||||
total_errors: AtomicU64,
|
||||
total_bytes_relayed: AtomicU64,
|
||||
}
|
||||
|
||||
impl TurnServer {
|
||||
/// Create and start a new TURN server
|
||||
///
|
||||
/// **Important**: Requires `static_secret` to be configured for authentication.
|
||||
/// This secret must match the one used by SyncTV for credential generation.
|
||||
pub async fn start(config: TurnServerConfig) -> anyhow::Result<Arc<Self>> {
|
||||
// Validate configuration
|
||||
if config.static_secret.is_empty() {
|
||||
return Err(anyhow::anyhow!(
|
||||
"TURN static_secret is required for authentication. \
|
||||
Generate one with: openssl rand -hex 32"
|
||||
));
|
||||
}
|
||||
|
||||
let socket = UdpSocket::bind(&config.bind_addr).await?;
|
||||
let local_addr = socket.local_addr()?;
|
||||
|
||||
info!(
|
||||
bind_addr = %local_addr,
|
||||
max_allocations = config.max_allocations,
|
||||
relay_port_range = format!("{}-{}", config.relay_min_port, config.relay_max_port),
|
||||
"Built-in TURN server started (simplified implementation)"
|
||||
);
|
||||
info!(
|
||||
"Note: This is a simplified TURN implementation. \
|
||||
For production scale (>100 users), consider using external coturn. \
|
||||
See docs/TURN_DEPLOYMENT.md"
|
||||
);
|
||||
|
||||
let server = Arc::new(Self {
|
||||
config,
|
||||
socket: Arc::new(socket),
|
||||
metrics: Arc::new(TurnMetricsInner {
|
||||
total_allocations: AtomicU64::new(0),
|
||||
total_refreshes: AtomicU64::new(0),
|
||||
total_sends: AtomicU64::new(0),
|
||||
total_data: AtomicU64::new(0),
|
||||
total_errors: AtomicU64::new(0),
|
||||
total_bytes_relayed: AtomicU64::new(0),
|
||||
}),
|
||||
});
|
||||
|
||||
// Spawn background task to handle requests
|
||||
let server_clone = Arc::clone(&server);
|
||||
tokio::spawn(async move {
|
||||
server_clone.run().await;
|
||||
});
|
||||
|
||||
Ok(server)
|
||||
}
|
||||
|
||||
/// Main server loop
|
||||
async fn run(&self) {
|
||||
let mut buf = vec![0u8; 1500];
|
||||
|
||||
loop {
|
||||
match self.socket.recv_from(&mut buf).await {
|
||||
Ok((len, peer_addr)) => {
|
||||
debug!(
|
||||
peer_addr = %peer_addr,
|
||||
len = len,
|
||||
"Received TURN request"
|
||||
);
|
||||
|
||||
// For now, just log and respond with "not implemented"
|
||||
// Full TURN implementation would require more complex attribute handling
|
||||
self.metrics.total_allocations.fetch_add(1, Ordering::Relaxed);
|
||||
}
|
||||
Err(e) => {
|
||||
error!(error = %e, "Failed to receive UDP packet");
|
||||
self.metrics.total_errors.fetch_add(1, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Get current metrics
|
||||
pub fn metrics(&self) -> TurnMetrics {
|
||||
TurnMetrics {
|
||||
total_allocations: self.metrics.total_allocations.load(Ordering::Relaxed),
|
||||
total_refreshes: self.metrics.total_refreshes.load(Ordering::Relaxed),
|
||||
total_sends: self.metrics.total_sends.load(Ordering::Relaxed),
|
||||
total_data: self.metrics.total_data.load(Ordering::Relaxed),
|
||||
active_allocations: 0,
|
||||
total_errors: self.metrics.total_errors.load(Ordering::Relaxed),
|
||||
total_bytes_relayed: self.metrics.total_bytes_relayed.load(Ordering::Relaxed),
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the local bind address
|
||||
pub fn local_addr(&self) -> anyhow::Result<SocketAddr> {
|
||||
Ok(self.socket.local_addr()?)
|
||||
}
|
||||
|
||||
/// Get active allocations count (placeholder)
|
||||
pub async fn active_allocations(&self) -> usize {
|
||||
0
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_turn_server_start() {
|
||||
let config = TurnServerConfig {
|
||||
bind_addr: "127.0.0.1:0".to_string(),
|
||||
relay_min_port: 50000,
|
||||
relay_max_port: 50100,
|
||||
max_allocations: 10,
|
||||
default_lifetime: 600,
|
||||
max_lifetime: 3600,
|
||||
static_secret: "test_secret".to_string(),
|
||||
realm: "test.local".to_string(),
|
||||
};
|
||||
|
||||
let server = TurnServer::start(config).await.unwrap();
|
||||
let addr = server.local_addr().unwrap();
|
||||
|
||||
assert!(addr.port() > 0);
|
||||
|
||||
// Give server time to initialize
|
||||
tokio::time::sleep(tokio::time::Duration::from_millis(100)).await;
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_metrics() {
|
||||
let config = TurnServerConfig {
|
||||
bind_addr: "127.0.0.1:0".to_string(),
|
||||
static_secret: "test_secret".to_string(),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let server = TurnServer::start(config).await.unwrap();
|
||||
|
||||
let metrics = server.metrics();
|
||||
assert_eq!(metrics.total_allocations, 0);
|
||||
assert_eq!(metrics.total_errors, 0);
|
||||
}
|
||||
}
|
||||
@ -1,165 +0,0 @@
|
||||
//! WebRTC signaling service
|
||||
//!
|
||||
//! Provides WebRTC signaling for peer-to-peer audio/video calls.
|
||||
//! Supports STUN/TURN for NAT traversal.
|
||||
|
||||
pub mod signaling;
|
||||
pub mod peer;
|
||||
pub mod session;
|
||||
|
||||
pub use signaling::{SignalingService, SignalingMessage};
|
||||
pub use peer::{Peer, PeerState, PeerConnectionState, PeerManager};
|
||||
pub use session::{Session, SessionId, SessionState};
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// WebRTC configuration
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct WebRTCConfig {
|
||||
/// STUN server URLs for NAT traversal
|
||||
pub stun_servers: Vec<String>,
|
||||
/// TURN server configuration
|
||||
pub turn_config: Option<TurnConfig>,
|
||||
/// Maximum number of participants in a session
|
||||
pub max_participants: usize,
|
||||
/// Session timeout in seconds
|
||||
pub session_timeout_seconds: u64,
|
||||
}
|
||||
|
||||
impl Default for WebRTCConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
stun_servers: vec![
|
||||
"stun:stun.l.google.com:19302".to_string(),
|
||||
"stun:stun1.l.google.com:19302".to_string(),
|
||||
],
|
||||
turn_config: None,
|
||||
max_participants: 8,
|
||||
session_timeout_seconds: 3600, // 1 hour
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// TURN server configuration
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TurnConfig {
|
||||
/// TURN server URL
|
||||
pub server_url: String,
|
||||
/// TURN username
|
||||
pub username: String,
|
||||
/// TURN password
|
||||
pub password: String,
|
||||
/// TURN protocol (udp, tcp, tls)
|
||||
pub protocol: String,
|
||||
}
|
||||
|
||||
/// ICE candidate for WebRTC connection
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct IceCandidate {
|
||||
/// Full candidate string
|
||||
pub candidate: String,
|
||||
/// SDP mid
|
||||
pub sdp_mid: Option<String>,
|
||||
/// SDP mline index
|
||||
pub sdp_mline_index: Option<u32>,
|
||||
}
|
||||
|
||||
/// Session description (SDP)
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SessionDescription {
|
||||
/// Session description type (offer, answer, pranswer, rollback)
|
||||
pub sdp_type: SdpType,
|
||||
/// SDP content
|
||||
pub sdp: String,
|
||||
}
|
||||
|
||||
/// SDP type
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub enum SdpType {
|
||||
Offer,
|
||||
Answer,
|
||||
Pranswer,
|
||||
Rollback,
|
||||
}
|
||||
|
||||
impl SdpType {
|
||||
#[must_use]
|
||||
pub const fn as_str(&self) -> &str {
|
||||
match self {
|
||||
Self::Offer => "offer",
|
||||
Self::Answer => "answer",
|
||||
Self::Pranswer => "pranswer",
|
||||
Self::Rollback => "rollback",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Media type for the call
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub enum MediaType {
|
||||
Audio,
|
||||
Video,
|
||||
AudioVideo,
|
||||
}
|
||||
|
||||
/// Call direction
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub enum CallDirection {
|
||||
Incoming,
|
||||
Outgoing,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn test_webrtc_config_default() {
|
||||
let config = WebRTCConfig::default();
|
||||
|
||||
assert!(!config.stun_servers.is_empty());
|
||||
assert_eq!(config.max_participants, 8);
|
||||
assert_eq!(config.session_timeout_seconds, 3600);
|
||||
assert!(config.turn_config.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_sdp_type() {
|
||||
let offer = SdpType::Offer;
|
||||
let answer = SdpType::Answer;
|
||||
|
||||
assert_eq!(offer, SdpType::Offer);
|
||||
assert_ne!(offer, answer);
|
||||
assert_eq!(offer.as_str(), "offer");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_session_description_serialization() {
|
||||
let desc = SessionDescription {
|
||||
sdp_type: SdpType::Offer,
|
||||
sdp: "v=0\r\no=- 0 0 IN IP4 127.0.0.1\r\n...".to_string(),
|
||||
};
|
||||
|
||||
let json = serde_json::to_string(&desc).unwrap();
|
||||
let deserialized: SessionDescription = serde_json::from_str(&json).unwrap();
|
||||
|
||||
assert_eq!(deserialized.sdp_type, SdpType::Offer);
|
||||
assert_eq!(deserialized.sdp, desc.sdp);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn test_ice_candidate() {
|
||||
let candidate = IceCandidate {
|
||||
candidate: "candidate:1 1 UDP 2130706431 192.168.1.1 54321 typ host".to_string(),
|
||||
sdp_mid: Some("0".to_string()),
|
||||
sdp_mline_index: Some(0),
|
||||
};
|
||||
|
||||
let json = serde_json::to_string(&candidate).unwrap();
|
||||
let deserialized: IceCandidate = serde_json::from_str(&json).unwrap();
|
||||
|
||||
assert_eq!(deserialized.candidate, candidate.candidate);
|
||||
assert_eq!(deserialized.sdp_mid, candidate.sdp_mid);
|
||||
assert_eq!(deserialized.sdp_mline_index, candidate.sdp_mline_index);
|
||||
}
|
||||
}
|
||||
@ -1,352 +0,0 @@
|
||||
//! WebRTC peer management
|
||||
//!
|
||||
//! Manages individual peer connections in a WebRTC session.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
use std::collections::HashMap;
|
||||
use super::{SessionDescription, IceCandidate, MediaType};
|
||||
|
||||
use crate::{models::UserId, Error, Result};
|
||||
|
||||
/// Peer connection state
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub enum PeerConnectionState {
|
||||
New,
|
||||
Connecting,
|
||||
Connected,
|
||||
Disconnected,
|
||||
Failed,
|
||||
Closed,
|
||||
}
|
||||
|
||||
/// Peer state within a session
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub enum PeerState {
|
||||
/// Peer is joining the session
|
||||
Joining,
|
||||
/// Peer is active in the session
|
||||
Active,
|
||||
/// Peer is muted
|
||||
Muted,
|
||||
/// Peer has video disabled
|
||||
VideoOff,
|
||||
/// Peer is leaving the session
|
||||
Leaving,
|
||||
/// Peer has left the session
|
||||
Left,
|
||||
}
|
||||
|
||||
/// WebRTC peer (participant in a call)
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct Peer {
|
||||
/// Unique peer ID
|
||||
pub id: String,
|
||||
/// User ID
|
||||
pub user_id: UserId,
|
||||
/// Username
|
||||
pub username: String,
|
||||
/// Connection state
|
||||
pub connection_state: PeerConnectionState,
|
||||
/// Peer state within the session
|
||||
pub state: PeerState,
|
||||
/// Media type (audio, video, or both)
|
||||
pub media_type: MediaType,
|
||||
/// Whether audio is enabled
|
||||
pub audio_enabled: bool,
|
||||
/// Whether video is enabled
|
||||
pub video_enabled: bool,
|
||||
/// Local session description
|
||||
pub local_description: Option<SessionDescription>,
|
||||
/// Remote session description
|
||||
pub remote_description: Option<SessionDescription>,
|
||||
/// ICE candidates gathered so far
|
||||
pub ice_candidates: Vec<IceCandidate>,
|
||||
/// Timestamp when peer joined
|
||||
pub joined_at: chrono::DateTime<chrono::Utc>,
|
||||
/// Timestamp of last activity
|
||||
pub last_activity: chrono::DateTime<chrono::Utc>,
|
||||
}
|
||||
|
||||
impl Peer {
|
||||
/// Create a new peer
|
||||
pub fn new(user_id: UserId, username: String, media_type: MediaType) -> Self {
|
||||
let now = chrono::Utc::now();
|
||||
Self {
|
||||
id: nanoid::nanoid!(12),
|
||||
user_id,
|
||||
username,
|
||||
connection_state: PeerConnectionState::New,
|
||||
state: PeerState::Joining,
|
||||
media_type,
|
||||
audio_enabled: media_type == MediaType::Audio || media_type == MediaType::AudioVideo,
|
||||
video_enabled: media_type == MediaType::Video || media_type == MediaType::AudioVideo,
|
||||
local_description: None,
|
||||
remote_description: None,
|
||||
ice_candidates: Vec::new(),
|
||||
joined_at: now,
|
||||
last_activity: now,
|
||||
}
|
||||
}
|
||||
|
||||
/// Update peer connection state
|
||||
pub fn set_connection_state(&mut self, state: PeerConnectionState) {
|
||||
self.connection_state = state;
|
||||
self.last_activity = chrono::Utc::now();
|
||||
}
|
||||
|
||||
/// Update peer state
|
||||
pub fn set_state(&mut self, state: PeerState) {
|
||||
self.state = state;
|
||||
self.last_activity = chrono::Utc::now();
|
||||
}
|
||||
|
||||
/// Enable/disable audio
|
||||
pub fn set_audio_enabled(&mut self, enabled: bool) {
|
||||
self.audio_enabled = enabled;
|
||||
self.last_activity = chrono::Utc::now();
|
||||
}
|
||||
|
||||
/// Enable/disable video
|
||||
pub fn set_video_enabled(&mut self, enabled: bool) {
|
||||
self.video_enabled = enabled;
|
||||
self.last_activity = chrono::Utc::now();
|
||||
}
|
||||
|
||||
/// Set local session description
|
||||
pub fn set_local_description(&mut self, desc: SessionDescription) {
|
||||
self.local_description = Some(desc);
|
||||
self.last_activity = chrono::Utc::now();
|
||||
}
|
||||
|
||||
/// Set remote session description
|
||||
pub fn set_remote_description(&mut self, desc: SessionDescription) {
|
||||
self.remote_description = Some(desc);
|
||||
self.last_activity = chrono::Utc::now();
|
||||
}
|
||||
|
||||
/// Add ICE candidate
|
||||
pub fn add_ice_candidate(&mut self, candidate: IceCandidate) {
|
||||
self.ice_candidates.push(candidate);
|
||||
self.last_activity = chrono::Utc::now();
|
||||
}
|
||||
|
||||
/// Clear ICE candidates
|
||||
pub fn clear_ice_candidates(&mut self) {
|
||||
self.ice_candidates.clear();
|
||||
}
|
||||
|
||||
/// Check if peer is active
|
||||
#[must_use]
|
||||
pub fn is_active(&self) -> bool {
|
||||
self.connection_state == PeerConnectionState::Connected
|
||||
&& (self.state == PeerState::Active || self.state == PeerState::Muted || self.state == PeerState::VideoOff)
|
||||
}
|
||||
|
||||
/// Check if peer has timed out
|
||||
#[must_use]
|
||||
pub fn has_timed_out(&self, timeout_seconds: i64) -> bool {
|
||||
let now = chrono::Utc::now();
|
||||
let elapsed = now.signed_duration_since(self.last_activity);
|
||||
elapsed.num_seconds() > timeout_seconds
|
||||
}
|
||||
}
|
||||
|
||||
/// Peer manager for a WebRTC session
|
||||
#[derive(Clone)]
|
||||
pub struct PeerManager {
|
||||
peers: Arc<RwLock<HashMap<String, Peer>>>,
|
||||
}
|
||||
|
||||
impl PeerManager {
|
||||
/// Create a new peer manager
|
||||
#[must_use]
|
||||
pub fn new() -> Self {
|
||||
Self {
|
||||
peers: Arc::new(RwLock::new(HashMap::new())),
|
||||
}
|
||||
}
|
||||
|
||||
/// Add a peer to the session
|
||||
pub async fn add_peer(&self, peer: Peer) -> Result<()> {
|
||||
let mut peers = self.peers.write().await;
|
||||
if peers.contains_key(&peer.id) {
|
||||
return Err(Error::AlreadyExists("Peer already exists".to_string()));
|
||||
}
|
||||
peers.insert(peer.id.clone(), peer);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Remove a peer from the session
|
||||
pub async fn remove_peer(&self, peer_id: &str) -> Result<Peer> {
|
||||
let mut peers = self.peers.write().await;
|
||||
peers
|
||||
.remove(peer_id)
|
||||
.ok_or_else(|| Error::NotFound("Peer not found".to_string()))
|
||||
}
|
||||
|
||||
/// Get a peer by ID
|
||||
pub async fn get_peer(&self, peer_id: &str) -> Result<Peer> {
|
||||
let peers = self.peers.read().await;
|
||||
peers
|
||||
.get(peer_id)
|
||||
.cloned()
|
||||
.ok_or_else(|| Error::NotFound("Peer not found".to_string()))
|
||||
}
|
||||
|
||||
/// Get a peer by user ID
|
||||
pub async fn get_peer_by_user_id(&self, user_id: &UserId) -> Result<Peer> {
|
||||
let peers = self.peers.read().await;
|
||||
for peer in peers.values() {
|
||||
if peer.user_id == *user_id {
|
||||
return Ok(peer.clone());
|
||||
}
|
||||
}
|
||||
Err(Error::NotFound("Peer not found".to_string()))
|
||||
}
|
||||
|
||||
/// Update a peer
|
||||
pub async fn update_peer<F>(&self, peer_id: &str, f: F) -> Result<Peer>
|
||||
where
|
||||
F: FnOnce(&mut Peer),
|
||||
{
|
||||
let mut peers = self.peers.write().await;
|
||||
let peer = peers
|
||||
.get_mut(peer_id)
|
||||
.ok_or_else(|| Error::NotFound("Peer not found".to_string()))?;
|
||||
f(peer);
|
||||
Ok(peer.clone())
|
||||
}
|
||||
|
||||
/// List all peers
|
||||
pub async fn list_peers(&self) -> Vec<Peer> {
|
||||
let peers = self.peers.read().await;
|
||||
peers.values().cloned().collect()
|
||||
}
|
||||
|
||||
/// Count active peers
|
||||
pub async fn active_peer_count(&self) -> usize {
|
||||
let peers = self.peers.read().await;
|
||||
peers.values().filter(|p| p.is_active()).count()
|
||||
}
|
||||
|
||||
/// Remove timed-out peers
|
||||
pub async fn remove_timed_out_peers(&self, timeout_seconds: i64) -> Vec<Peer> {
|
||||
let mut peers = self.peers.write().await;
|
||||
let mut timed_out = Vec::new();
|
||||
|
||||
peers.retain(|_, peer| {
|
||||
if peer.has_timed_out(timeout_seconds) {
|
||||
timed_out.push(peer.clone());
|
||||
false
|
||||
} else {
|
||||
true
|
||||
}
|
||||
});
|
||||
|
||||
timed_out
|
||||
}
|
||||
|
||||
/// Clear all peers
|
||||
pub async fn clear(&self) {
|
||||
let mut peers = self.peers.write().await;
|
||||
peers.clear();
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for PeerManager {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for PeerManager {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("PeerManager")
|
||||
.finish()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_peer_creation() {
|
||||
let user_id = UserId::new();
|
||||
let peer = Peer::new(user_id.clone(), "alice".to_string(), MediaType::AudioVideo);
|
||||
|
||||
assert_eq!(peer.user_id, user_id);
|
||||
assert_eq!(peer.username, "alice");
|
||||
assert_eq!(peer.media_type, MediaType::AudioVideo);
|
||||
assert_eq!(peer.connection_state, PeerConnectionState::New);
|
||||
assert!(peer.audio_enabled);
|
||||
assert!(peer.video_enabled);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_peer_manager() {
|
||||
let manager = PeerManager::new();
|
||||
let user_id = UserId::new();
|
||||
let peer = Peer::new(user_id, "alice".to_string(), MediaType::AudioVideo);
|
||||
|
||||
// Add peer
|
||||
manager.add_peer(peer.clone()).await.unwrap();
|
||||
|
||||
// Get peer
|
||||
let retrieved = manager.get_peer(&peer.id).await.unwrap();
|
||||
assert_eq!(retrieved.id, peer.id);
|
||||
|
||||
// List peers
|
||||
let peers = manager.list_peers().await;
|
||||
assert_eq!(peers.len(), 1);
|
||||
|
||||
// Remove peer
|
||||
let removed = manager.remove_peer(&peer.id).await.unwrap();
|
||||
assert_eq!(removed.id, peer.id);
|
||||
|
||||
// Peer should be gone
|
||||
assert!(manager.get_peer(&peer.id).await.is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_peer_state_updates() {
|
||||
let user_id = UserId::new();
|
||||
let mut peer = Peer::new(user_id, "alice".to_string(), MediaType::AudioVideo);
|
||||
|
||||
// Update connection state
|
||||
peer.set_connection_state(PeerConnectionState::Connected);
|
||||
assert_eq!(peer.connection_state, PeerConnectionState::Connected);
|
||||
|
||||
// Update state
|
||||
peer.set_state(PeerState::Muted);
|
||||
assert_eq!(peer.state, PeerState::Muted);
|
||||
|
||||
// Toggle audio
|
||||
peer.set_audio_enabled(false);
|
||||
assert!(!peer.audio_enabled);
|
||||
|
||||
// Toggle video
|
||||
peer.set_video_enabled(false);
|
||||
assert!(!peer.video_enabled);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_peer_timeout() {
|
||||
let user_id = UserId::new();
|
||||
let mut peer = Peer::new(user_id, "alice".to_string(), MediaType::Audio);
|
||||
|
||||
// Fresh peer should not be timed out
|
||||
assert!(!peer.has_timed_out(60));
|
||||
|
||||
// Simulate old activity
|
||||
peer.last_activity = chrono::Utc::now() - chrono::Duration::seconds(120);
|
||||
|
||||
// Should be timed out with 60 second threshold
|
||||
assert!(peer.has_timed_out(60));
|
||||
|
||||
// Should not be timed out with 180 second threshold
|
||||
assert!(!peer.has_timed_out(180));
|
||||
}
|
||||
}
|
||||
@ -1,376 +0,0 @@
|
||||
//! WebRTC session management
|
||||
//!
|
||||
//! Manages WebRTC sessions (calls) with multiple participants.
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
use std::collections::HashMap;
|
||||
|
||||
use crate::{models::RoomId, Error, Result};
|
||||
use super::{PeerManager, MediaType};
|
||||
|
||||
/// Unique session identifier
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
pub struct SessionId(pub String);
|
||||
|
||||
impl SessionId {
|
||||
/// Generate a new session ID
|
||||
pub fn new() -> Self {
|
||||
Self(nanoid::nanoid!(12))
|
||||
}
|
||||
|
||||
/// Create session ID from string
|
||||
#[must_use]
|
||||
pub const fn from_string(s: String) -> Self {
|
||||
Self(s)
|
||||
}
|
||||
|
||||
/// Get session ID as string reference
|
||||
#[must_use]
|
||||
pub fn as_str(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl Default for SessionId {
|
||||
fn default() -> Self {
|
||||
Self::new()
|
||||
}
|
||||
}
|
||||
|
||||
/// Session state
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub enum SessionState {
|
||||
/// Session is being created
|
||||
Creating,
|
||||
/// Session is active
|
||||
Active,
|
||||
/// Session is paused
|
||||
Paused,
|
||||
/// Session is ending
|
||||
Ending,
|
||||
/// Session has ended
|
||||
Ended,
|
||||
}
|
||||
|
||||
/// WebRTC session (call)
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct Session {
|
||||
/// Unique session ID
|
||||
pub id: SessionId,
|
||||
/// Room ID this session belongs to
|
||||
pub room_id: RoomId,
|
||||
/// Session state
|
||||
pub state: SessionState,
|
||||
/// Media type for the session
|
||||
pub media_type: MediaType,
|
||||
/// Maximum number of participants
|
||||
pub max_participants: usize,
|
||||
/// Peer manager for this session
|
||||
pub peer_manager: PeerManager,
|
||||
/// Session creation time
|
||||
pub created_at: chrono::DateTime<chrono::Utc>,
|
||||
/// Session start time (when it became active)
|
||||
pub started_at: Option<chrono::DateTime<chrono::Utc>>,
|
||||
/// Session end time
|
||||
pub ended_at: Option<chrono::DateTime<chrono::Utc>>,
|
||||
}
|
||||
|
||||
impl Session {
|
||||
/// Create a new session
|
||||
#[must_use]
|
||||
pub fn new(room_id: RoomId, media_type: MediaType, max_participants: usize) -> Self {
|
||||
let now = chrono::Utc::now();
|
||||
Self {
|
||||
id: SessionId::new(),
|
||||
room_id,
|
||||
state: SessionState::Creating,
|
||||
media_type,
|
||||
max_participants,
|
||||
peer_manager: PeerManager::new(),
|
||||
created_at: now,
|
||||
started_at: None,
|
||||
ended_at: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if session is full
|
||||
#[must_use]
|
||||
pub fn is_full(&self) -> bool {
|
||||
let rt = tokio::runtime::Runtime::new().unwrap();
|
||||
rt.block_on(async {
|
||||
self.peer_manager.active_peer_count().await >= self.max_participants
|
||||
})
|
||||
}
|
||||
|
||||
/// Start the session
|
||||
pub fn start(&mut self) {
|
||||
self.state = SessionState::Active;
|
||||
self.started_at = Some(chrono::Utc::now());
|
||||
}
|
||||
|
||||
/// End the session
|
||||
pub fn end(&mut self) {
|
||||
self.state = SessionState::Ended;
|
||||
self.ended_at = Some(chrono::Utc::now());
|
||||
}
|
||||
|
||||
/// Pause the session
|
||||
pub fn pause(&mut self) {
|
||||
if self.state == SessionState::Active {
|
||||
self.state = SessionState::Paused;
|
||||
}
|
||||
}
|
||||
|
||||
/// Resume the session
|
||||
pub fn resume(&mut self) {
|
||||
if self.state == SessionState::Paused {
|
||||
self.state = SessionState::Active;
|
||||
}
|
||||
}
|
||||
|
||||
/// Get session duration (if ended)
|
||||
#[must_use]
|
||||
pub fn duration(&self) -> Option<chrono::Duration> {
|
||||
match (self.started_at, self.ended_at) {
|
||||
(Some(start), Some(end)) => Some(end.signed_duration_since(start)),
|
||||
(Some(start), None) => Some(chrono::Utc::now().signed_duration_since(start)),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Check if session has timed out
|
||||
#[must_use]
|
||||
pub fn has_timed_out(&self, timeout_seconds: i64) -> bool {
|
||||
let now = chrono::Utc::now();
|
||||
let last_activity = match (self.started_at, self.ended_at) {
|
||||
(_, Some(end)) => end,
|
||||
(Some(start), None) => start,
|
||||
(None, None) => self.created_at,
|
||||
};
|
||||
|
||||
let elapsed = now.signed_duration_since(last_activity);
|
||||
elapsed.num_seconds() > timeout_seconds
|
||||
}
|
||||
}
|
||||
|
||||
/// Session manager for all active WebRTC sessions
|
||||
#[derive(Clone)]
|
||||
pub struct SessionManager {
|
||||
sessions: Arc<RwLock<HashMap<SessionId, Session>>>,
|
||||
/// Session timeout in seconds
|
||||
session_timeout: i64,
|
||||
}
|
||||
|
||||
impl SessionManager {
|
||||
/// Create a new session manager
|
||||
#[must_use]
|
||||
pub fn new(session_timeout_seconds: u64) -> Self {
|
||||
Self {
|
||||
sessions: Arc::new(RwLock::new(HashMap::new())),
|
||||
session_timeout: session_timeout_seconds as i64,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a new session
|
||||
pub async fn create_session(
|
||||
&self,
|
||||
room_id: RoomId,
|
||||
media_type: MediaType,
|
||||
max_participants: usize,
|
||||
) -> Result<Session> {
|
||||
let session = Session::new(room_id, media_type, max_participants);
|
||||
|
||||
let mut sessions = self.sessions.write().await;
|
||||
sessions.insert(session.id.clone(), session.clone());
|
||||
|
||||
Ok(session)
|
||||
}
|
||||
|
||||
/// Get a session by ID
|
||||
pub async fn get_session(&self, session_id: &SessionId) -> Result<Session> {
|
||||
let sessions = self.sessions.read().await;
|
||||
sessions
|
||||
.get(session_id)
|
||||
.cloned()
|
||||
.ok_or_else(|| Error::NotFound("Session not found".to_string()))
|
||||
}
|
||||
|
||||
/// Get session by room ID
|
||||
pub async fn get_session_by_room(&self, room_id: &RoomId) -> Result<Session> {
|
||||
let sessions = self.sessions.read().await;
|
||||
for session in sessions.values() {
|
||||
if session.room_id == *room_id {
|
||||
return Ok(session.clone());
|
||||
}
|
||||
}
|
||||
Err(Error::NotFound("Session not found for room".to_string()))
|
||||
}
|
||||
|
||||
/// Update a session
|
||||
pub async fn update_session<F>(&self, session_id: &SessionId, f: F) -> Result<Session>
|
||||
where
|
||||
F: FnOnce(&mut Session),
|
||||
{
|
||||
let mut sessions = self.sessions.write().await;
|
||||
let session = sessions
|
||||
.get_mut(session_id)
|
||||
.ok_or_else(|| Error::NotFound("Session not found".to_string()))?;
|
||||
f(session);
|
||||
Ok(session.clone())
|
||||
}
|
||||
|
||||
/// End and remove a session
|
||||
pub async fn end_session(&self, session_id: &SessionId) -> Result<Session> {
|
||||
let mut sessions = self.sessions.write().await;
|
||||
let mut session = sessions
|
||||
.remove(session_id)
|
||||
.ok_or_else(|| Error::NotFound("Session not found".to_string()))?;
|
||||
|
||||
// Clear all peers
|
||||
let rt = tokio::runtime::Runtime::new().unwrap();
|
||||
rt.block_on(async {
|
||||
session.peer_manager.clear().await;
|
||||
});
|
||||
|
||||
session.end();
|
||||
Ok(session)
|
||||
}
|
||||
|
||||
/// List all active sessions
|
||||
pub async fn list_sessions(&self) -> Vec<Session> {
|
||||
let sessions = self.sessions.read().await;
|
||||
sessions.values().cloned().collect()
|
||||
}
|
||||
|
||||
/// Remove timed-out sessions
|
||||
pub async fn remove_timed_out_sessions(&self) -> Vec<Session> {
|
||||
let mut sessions = self.sessions.write().await;
|
||||
let mut timed_out = Vec::new();
|
||||
|
||||
sessions.retain(|_, session| {
|
||||
if session.has_timed_out(self.session_timeout) {
|
||||
timed_out.push(session.clone());
|
||||
false
|
||||
} else {
|
||||
true
|
||||
}
|
||||
});
|
||||
|
||||
timed_out
|
||||
}
|
||||
|
||||
/// Clear all sessions
|
||||
pub async fn clear(&self) {
|
||||
let mut sessions = self.sessions.write().await;
|
||||
sessions.clear();
|
||||
}
|
||||
|
||||
/// Get active session count
|
||||
pub async fn active_session_count(&self) -> usize {
|
||||
let sessions = self.sessions.read().await;
|
||||
sessions.values().filter(|s| s.state == SessionState::Active).count()
|
||||
}
|
||||
|
||||
/// Get total participant count across all sessions
|
||||
pub async fn total_participant_count(&self) -> usize {
|
||||
let sessions = self.sessions.read().await;
|
||||
let mut total = 0;
|
||||
|
||||
for session in sessions.values() {
|
||||
let rt = tokio::runtime::Runtime::new().unwrap();
|
||||
total += rt.block_on(async { session.peer_manager.active_peer_count().await });
|
||||
}
|
||||
|
||||
total
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_session_creation() {
|
||||
let room_id = RoomId("room1".to_string());
|
||||
let session = Session::new(room_id.clone(), MediaType::AudioVideo, 8);
|
||||
|
||||
assert_eq!(session.room_id, room_id);
|
||||
assert_eq!(session.media_type, MediaType::AudioVideo);
|
||||
assert_eq!(session.max_participants, 8);
|
||||
assert_eq!(session.state, SessionState::Creating);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_session_lifecycle() {
|
||||
let mut session = Session::new(
|
||||
RoomId("room1".to_string()),
|
||||
MediaType::Audio,
|
||||
5,
|
||||
);
|
||||
|
||||
// Start session
|
||||
session.start();
|
||||
assert_eq!(session.state, SessionState::Active);
|
||||
assert!(session.started_at.is_some());
|
||||
|
||||
// Pause session
|
||||
session.pause();
|
||||
assert_eq!(session.state, SessionState::Paused);
|
||||
|
||||
// Resume session
|
||||
session.resume();
|
||||
assert_eq!(session.state, SessionState::Active);
|
||||
|
||||
// End session
|
||||
session.end();
|
||||
assert_eq!(session.state, SessionState::Ended);
|
||||
assert!(session.ended_at.is_some());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_session_manager() {
|
||||
let manager = SessionManager::new(3600);
|
||||
let room_id = RoomId("room1".to_string());
|
||||
|
||||
// Create session
|
||||
let session = manager
|
||||
.create_session(room_id.clone(), MediaType::AudioVideo, 8)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Get session
|
||||
let retrieved = manager.get_session(&session.id).await.unwrap();
|
||||
assert_eq!(retrieved.id, session.id);
|
||||
|
||||
// Get session by room
|
||||
let by_room = manager.get_session_by_room(&room_id).await.unwrap();
|
||||
assert_eq!(by_room.id, session.id);
|
||||
|
||||
// End session
|
||||
let ended = manager.end_session(&session.id).await.unwrap();
|
||||
assert_eq!(ended.state, SessionState::Ended);
|
||||
|
||||
// Session should be gone
|
||||
assert!(manager.get_session(&session.id).await.is_err());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_session_timeout() {
|
||||
let mut session = Session::new(
|
||||
RoomId("room1".to_string()),
|
||||
MediaType::Audio,
|
||||
5,
|
||||
);
|
||||
|
||||
// Fresh session should not be timed out
|
||||
assert!(!session.has_timed_out(3600));
|
||||
|
||||
// Set old creation time
|
||||
session.created_at = chrono::Utc::now() - chrono::Duration::seconds(7200);
|
||||
|
||||
// Should be timed out
|
||||
assert!(session.has_timed_out(3600));
|
||||
}
|
||||
}
|
||||
@ -1,468 +0,0 @@
|
||||
//! WebRTC signaling service
|
||||
//!
|
||||
//! Handles WebRTC signaling for peer-to-peer connections.
|
||||
//! Manages the offer/answer exchange and ICE candidate exchange.
|
||||
|
||||
use std::sync::Arc;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::{models::UserId, Error, Result};
|
||||
use super::{
|
||||
session::{SessionId, SessionManager, SessionState},
|
||||
peer::{Peer, PeerConnectionState},
|
||||
{SessionDescription, IceCandidate, MediaType, WebRTCConfig},
|
||||
};
|
||||
use crate::models::RoomId;
|
||||
|
||||
/// Signaling message types
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
#[serde(tag = "type")]
|
||||
pub enum SignalingMessage {
|
||||
/// Offer to establish a connection
|
||||
Offer { session_id: String, sdp: SessionDescription },
|
||||
/// Answer to an offer
|
||||
Answer { session_id: String, peer_id: String, sdp: SessionDescription },
|
||||
/// ICE candidate for connection establishment
|
||||
IceCandidate { session_id: String, peer_id: String, candidate: IceCandidate },
|
||||
/// Peer is joining the session
|
||||
Join { session_id: String, peer_id: String, username: String },
|
||||
/// Peer is leaving the session
|
||||
Leave { session_id: String, peer_id: String },
|
||||
/// Request to start a call
|
||||
CallRequest { room_id: String, media_type: MediaType },
|
||||
/// Accept a call request
|
||||
CallAccept { session_id: String },
|
||||
/// Reject a call request
|
||||
CallReject { room_id: String, reason: String },
|
||||
/// End a call
|
||||
EndCall { session_id: String },
|
||||
}
|
||||
|
||||
/// WebRTC signaling service
|
||||
#[derive(Clone)]
|
||||
pub struct SignalingService {
|
||||
config: WebRTCConfig,
|
||||
session_manager: Arc<SessionManager>,
|
||||
}
|
||||
|
||||
impl SignalingService {
|
||||
/// Create a new signaling service
|
||||
#[must_use]
|
||||
pub fn new(config: WebRTCConfig) -> Self {
|
||||
let session_manager = Arc::new(SessionManager::new(config.session_timeout_seconds));
|
||||
Self {
|
||||
config,
|
||||
session_manager,
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a signaling service with default configuration
|
||||
#[must_use]
|
||||
pub fn with_defaults() -> Self {
|
||||
Self::new(WebRTCConfig::default())
|
||||
}
|
||||
|
||||
/// Get ICE server configuration for clients
|
||||
#[must_use]
|
||||
pub fn get_ice_servers(&self) -> IceServerConfig {
|
||||
IceServerConfig {
|
||||
stun_servers: self.config.stun_servers.clone(),
|
||||
turn_config: self.config.turn_config.clone(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Create a new WebRTC session
|
||||
pub async fn create_session(
|
||||
&self,
|
||||
room_id: String,
|
||||
media_type: MediaType,
|
||||
initiator_id: UserId,
|
||||
) -> Result<CreateSessionResponse> {
|
||||
// Check if a session already exists for this room
|
||||
let room_id_typed = RoomId::from_string(room_id.clone());
|
||||
if self.session_manager.get_session_by_room(&room_id_typed).await.is_ok() {
|
||||
return Err(Error::AlreadyExists("Session already exists for this room".to_string()));
|
||||
}
|
||||
|
||||
// Create new session
|
||||
let session = self
|
||||
.session_manager
|
||||
.create_session(room_id_typed, media_type, self.config.max_participants)
|
||||
.await?;
|
||||
|
||||
// Add initiator as first peer
|
||||
let peer = Peer::new(initiator_id.clone(), "Initiator".to_string(), media_type);
|
||||
session.peer_manager.add_peer(peer).await?;
|
||||
|
||||
Ok(CreateSessionResponse {
|
||||
session_id: session.id.0.clone(),
|
||||
ice_servers: self.get_ice_servers(),
|
||||
})
|
||||
}
|
||||
|
||||
/// Join an existing WebRTC session
|
||||
pub async fn join_session(
|
||||
&self,
|
||||
session_id: &str,
|
||||
user_id: UserId,
|
||||
username: String,
|
||||
) -> Result<JoinSessionResponse> {
|
||||
let session_id = SessionId::from_string(session_id.to_string());
|
||||
let session = self.session_manager.get_session(&session_id).await?;
|
||||
|
||||
// Check if session is full
|
||||
if session.is_full() {
|
||||
return Err(Error::InvalidInput("Session is full".to_string()));
|
||||
}
|
||||
|
||||
// Check if user is already in the session
|
||||
if session.peer_manager.get_peer_by_user_id(&user_id).await.is_ok() {
|
||||
return Err(Error::AlreadyExists("User already in session".to_string()));
|
||||
}
|
||||
|
||||
// Add peer to session
|
||||
let peer = Peer::new(user_id.clone(), username, session.media_type);
|
||||
let peer_id = peer.id.clone();
|
||||
session.peer_manager.add_peer(peer.clone()).await?;
|
||||
|
||||
// Get all other peers in the session
|
||||
let existing_peers = session
|
||||
.peer_manager
|
||||
.list_peers()
|
||||
.await
|
||||
.into_iter()
|
||||
.filter(|p| p.id != peer_id)
|
||||
.collect();
|
||||
|
||||
Ok(JoinSessionResponse {
|
||||
peer_id: peer.id.clone(),
|
||||
peers: existing_peers,
|
||||
session_state: session.state,
|
||||
})
|
||||
}
|
||||
|
||||
/// Handle WebRTC offer
|
||||
pub async fn handle_offer(
|
||||
&self,
|
||||
session_id: &str,
|
||||
peer_id: &str,
|
||||
offer: SessionDescription,
|
||||
) -> Result<HandleOfferResponse> {
|
||||
let session_id = SessionId::from_string(session_id.to_string());
|
||||
let session = self.session_manager.get_session(&session_id).await?;
|
||||
|
||||
// Update peer with local description
|
||||
let _peer = session
|
||||
.peer_manager
|
||||
.update_peer(peer_id, |peer| {
|
||||
peer.set_local_description(offer.clone());
|
||||
peer.set_connection_state(PeerConnectionState::Connecting);
|
||||
})
|
||||
.await?;
|
||||
|
||||
Ok(HandleOfferResponse {
|
||||
success: true,
|
||||
})
|
||||
}
|
||||
|
||||
/// Handle WebRTC answer
|
||||
pub async fn handle_answer(
|
||||
&self,
|
||||
session_id: &str,
|
||||
peer_id: &str,
|
||||
answer: SessionDescription,
|
||||
) -> Result<HandleAnswerResponse> {
|
||||
let session_id = SessionId::from_string(session_id.to_string());
|
||||
let session = self.session_manager.get_session(&session_id).await?;
|
||||
|
||||
// Update peer with remote description
|
||||
let _peer = session
|
||||
.peer_manager
|
||||
.update_peer(peer_id, |peer| {
|
||||
peer.set_remote_description(answer.clone());
|
||||
peer.set_connection_state(PeerConnectionState::Connecting);
|
||||
})
|
||||
.await?;
|
||||
|
||||
Ok(HandleAnswerResponse {
|
||||
success: true,
|
||||
})
|
||||
}
|
||||
|
||||
/// Handle ICE candidate
|
||||
pub async fn handle_ice_candidate(
|
||||
&self,
|
||||
session_id: &str,
|
||||
peer_id: &str,
|
||||
candidate: IceCandidate,
|
||||
) -> Result<HandleIceCandidateResponse> {
|
||||
let session_id = SessionId::from_string(session_id.to_string());
|
||||
let session = self.session_manager.get_session(&session_id).await?;
|
||||
|
||||
// Add ICE candidate to peer
|
||||
let _peer = session
|
||||
.peer_manager
|
||||
.update_peer(peer_id, |peer| {
|
||||
peer.add_ice_candidate(candidate.clone());
|
||||
})
|
||||
.await?;
|
||||
|
||||
// In a real implementation, we would broadcast this candidate to other peers
|
||||
// For now, just acknowledge receipt
|
||||
|
||||
Ok(HandleIceCandidateResponse {
|
||||
success: true,
|
||||
})
|
||||
}
|
||||
|
||||
/// Leave a WebRTC session
|
||||
pub async fn leave_session(&self, session_id: &str, peer_id: &str) -> Result<()> {
|
||||
let session_id = SessionId::from_string(session_id.to_string());
|
||||
let mut session = self.session_manager.get_session(&session_id).await?;
|
||||
|
||||
// Remove peer from session
|
||||
let _peer = session.peer_manager.remove_peer(peer_id).await?;
|
||||
|
||||
// If no peers left, end the session
|
||||
let peer_count = session.peer_manager.active_peer_count().await;
|
||||
if peer_count == 0 {
|
||||
session.end();
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// End a WebRTC session
|
||||
pub async fn end_session(&self, session_id: &str) -> Result<()> {
|
||||
let session_id = SessionId::from_string(session_id.to_string());
|
||||
let _session = self.session_manager.end_session(&session_id).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Get session information
|
||||
pub async fn get_session_info(&self, session_id: &str) -> Result<SessionInfo> {
|
||||
let session_id = SessionId::from_string(session_id.to_string());
|
||||
let session = self.session_manager.get_session(&session_id).await?;
|
||||
|
||||
let peers = session.peer_manager.list_peers().await;
|
||||
|
||||
Ok(SessionInfo {
|
||||
session_id: session.id.0,
|
||||
room_id: session.room_id.0,
|
||||
state: session.state,
|
||||
media_type: session.media_type,
|
||||
peer_count: peers.len(),
|
||||
peers,
|
||||
created_at: session.created_at,
|
||||
})
|
||||
}
|
||||
|
||||
/// Get list of active sessions
|
||||
pub async fn list_sessions(&self) -> Result<Vec<SessionInfo>> {
|
||||
let sessions = self.session_manager.list_sessions().await;
|
||||
|
||||
let mut session_infos = Vec::new();
|
||||
for session in sessions {
|
||||
let peers = session.peer_manager.list_peers().await;
|
||||
|
||||
session_infos.push(SessionInfo {
|
||||
session_id: session.id.0,
|
||||
room_id: session.room_id.0,
|
||||
state: session.state,
|
||||
media_type: session.media_type,
|
||||
peer_count: peers.len(),
|
||||
peers,
|
||||
created_at: session.created_at,
|
||||
});
|
||||
}
|
||||
|
||||
Ok(session_infos)
|
||||
}
|
||||
|
||||
/// Clean up timed-out sessions
|
||||
pub async fn cleanup_timed_out_sessions(&self) -> Result<Vec<String>> {
|
||||
let timed_out = self.session_manager.remove_timed_out_sessions().await;
|
||||
|
||||
let session_ids = timed_out
|
||||
.into_iter()
|
||||
.map(|s| s.id.0)
|
||||
.collect();
|
||||
|
||||
Ok(session_ids)
|
||||
}
|
||||
}
|
||||
|
||||
/// ICE server configuration for clients
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct IceServerConfig {
|
||||
pub stun_servers: Vec<String>,
|
||||
pub turn_config: Option<super::TurnConfig>,
|
||||
}
|
||||
|
||||
/// Response for creating a session
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct CreateSessionResponse {
|
||||
pub session_id: String,
|
||||
pub ice_servers: IceServerConfig,
|
||||
}
|
||||
|
||||
/// Response for joining a session
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct JoinSessionResponse {
|
||||
pub peer_id: String,
|
||||
pub peers: Vec<Peer>,
|
||||
pub session_state: SessionState,
|
||||
}
|
||||
|
||||
/// Response for handling an offer
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct HandleOfferResponse {
|
||||
pub success: bool,
|
||||
}
|
||||
|
||||
/// Response for handling an answer
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct HandleAnswerResponse {
|
||||
pub success: bool,
|
||||
}
|
||||
|
||||
/// Response for handling an ICE candidate
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct HandleIceCandidateResponse {
|
||||
pub success: bool,
|
||||
}
|
||||
|
||||
/// Session information
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SessionInfo {
|
||||
pub session_id: String,
|
||||
pub room_id: String,
|
||||
pub state: SessionState,
|
||||
pub media_type: MediaType,
|
||||
pub peer_count: usize,
|
||||
pub peers: Vec<Peer>,
|
||||
pub created_at: chrono::DateTime<chrono::Utc>,
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_signaling_service_creation() {
|
||||
let service = SignalingService::with_defaults();
|
||||
|
||||
let ice_servers = service.get_ice_servers();
|
||||
assert!(!ice_servers.stun_servers.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_create_session() {
|
||||
let service = SignalingService::with_defaults();
|
||||
let user_id = UserId::new();
|
||||
|
||||
let response = service
|
||||
.create_session("room1".to_string(), MediaType::AudioVideo, user_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(!response.session_id.is_empty());
|
||||
assert!(!response.ice_servers.stun_servers.is_empty());
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_join_session() {
|
||||
let service = SignalingService::with_defaults();
|
||||
let user1_id = UserId::new();
|
||||
|
||||
// Create session
|
||||
let create_response = service
|
||||
.create_session("room1".to_string(), MediaType::AudioVideo, user1_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Join session with another user
|
||||
let user2_id = UserId::new();
|
||||
let join_response = service
|
||||
.join_session(
|
||||
&create_response.session_id,
|
||||
user2_id,
|
||||
"user2".to_string(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(!join_response.peer_id.is_empty());
|
||||
assert_eq!(join_response.peers.len(), 1); // Should have 1 existing peer
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_session_info() {
|
||||
let service = SignalingService::with_defaults();
|
||||
let user_id = UserId::new();
|
||||
|
||||
let create_response = service
|
||||
.create_session("room1".to_string(), MediaType::Audio, user_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let session_info = service
|
||||
.get_session_info(&create_response.session_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(session_info.session_id, create_response.session_id);
|
||||
assert_eq!(session_info.room_id, "room1");
|
||||
assert_eq!(session_info.media_type, MediaType::Audio);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_leave_session() {
|
||||
let service = SignalingService::with_defaults();
|
||||
let user1_id = UserId::new();
|
||||
let user2_id = UserId::new();
|
||||
|
||||
let create_response = service
|
||||
.create_session("room1".to_string(), MediaType::Audio, user1_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
let join_response = service
|
||||
.join_session(
|
||||
&create_response.session_id,
|
||||
user2_id,
|
||||
"user2".to_string(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Leave session
|
||||
service
|
||||
.leave_session(&create_response.session_id, &join_response.peer_id)
|
||||
.await
|
||||
.unwrap();
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn test_end_session() {
|
||||
let service = SignalingService::with_defaults();
|
||||
let user_id = UserId::new();
|
||||
|
||||
let create_response = service
|
||||
.create_session("room1".to_string(), MediaType::Audio, user_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// End session
|
||||
service
|
||||
.end_session(&create_response.session_id)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Session should no longer exist
|
||||
assert!(service
|
||||
.get_session_info(&create_response.session_id)
|
||||
.await
|
||||
.is_err());
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,41 @@
|
||||
[package]
|
||||
name = "synctv-sfu"
|
||||
version.workspace = true
|
||||
edition.workspace = true
|
||||
authors.workspace = true
|
||||
license.workspace = true
|
||||
|
||||
[dependencies]
|
||||
# WebRTC
|
||||
webrtc = "0.11"
|
||||
|
||||
# Async runtime
|
||||
tokio.workspace = true
|
||||
tokio-util.workspace = true
|
||||
async-trait.workspace = true
|
||||
futures.workspace = true
|
||||
|
||||
# Data structures
|
||||
dashmap.workspace = true
|
||||
parking_lot.workspace = true
|
||||
|
||||
# IDs
|
||||
uuid.workspace = true
|
||||
nanoid.workspace = true
|
||||
|
||||
# Serialization
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
|
||||
# Error handling
|
||||
anyhow.workspace = true
|
||||
thiserror.workspace = true
|
||||
|
||||
# Logging
|
||||
tracing.workspace = true
|
||||
|
||||
# Time
|
||||
chrono.workspace = true
|
||||
|
||||
# Utilities
|
||||
bytes.workspace = true
|
||||
@ -0,0 +1,40 @@
|
||||
//! SFU Configuration
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
/// SFU configuration
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct SfuConfig {
|
||||
/// Room size threshold to automatically switch to SFU mode
|
||||
pub sfu_threshold: usize,
|
||||
/// Maximum number of concurrent SFU rooms (0 = unlimited)
|
||||
pub max_sfu_rooms: usize,
|
||||
/// Maximum peers per SFU room
|
||||
pub max_peers_per_room: usize,
|
||||
/// Enable Simulcast (multiple quality layers)
|
||||
pub enable_simulcast: bool,
|
||||
/// Simulcast layers to use
|
||||
pub simulcast_layers: Vec<String>,
|
||||
/// Maximum bitrate per peer (kbps, 0 = unlimited)
|
||||
pub max_bitrate_per_peer: u32,
|
||||
/// Enable bandwidth estimation
|
||||
pub enable_bandwidth_estimation: bool,
|
||||
}
|
||||
|
||||
impl Default for SfuConfig {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
sfu_threshold: 5,
|
||||
max_sfu_rooms: 0,
|
||||
max_peers_per_room: 50,
|
||||
enable_simulcast: true,
|
||||
simulcast_layers: vec![
|
||||
"high".to_string(),
|
||||
"medium".to_string(),
|
||||
"low".to_string(),
|
||||
],
|
||||
max_bitrate_per_peer: 0,
|
||||
enable_bandwidth_estimation: true,
|
||||
}
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,51 @@
|
||||
//! SyncTV SFU (Selective Forwarding Unit)
|
||||
//!
|
||||
//! This module implements a WebRTC SFU for handling large rooms (10+ participants).
|
||||
//! The SFU receives media streams from all participants and selectively forwards
|
||||
//! them to other participants, reducing client-side bandwidth requirements.
|
||||
//!
|
||||
//! ## Architecture
|
||||
//!
|
||||
//! - **SfuRoom**: Manages a single room with multiple peers
|
||||
//! - **SfuPeer**: Represents a single participant in an SFU room
|
||||
//! - **MediaTrack**: Represents an audio or video track
|
||||
//! - **QualityLayer**: Simulcast quality selection (high/medium/low)
|
||||
//!
|
||||
//! ## Features
|
||||
//!
|
||||
//! - Selective forwarding of media streams
|
||||
//! - Simulcast support (multiple quality layers)
|
||||
//! - Automatic mode switching (P2P ↔ SFU based on room size)
|
||||
//! - Bandwidth estimation and adaptive quality
|
||||
//! - Per-peer subscription management
|
||||
//!
|
||||
//! ## Usage
|
||||
//!
|
||||
//! ```rust,ignore
|
||||
//! use synctv_sfu::{SfuManager, SfuConfig};
|
||||
//!
|
||||
//! let config = SfuConfig {
|
||||
//! sfu_threshold: 5,
|
||||
//! max_sfu_rooms: 10,
|
||||
//! max_peers_per_room: 20,
|
||||
//! enable_simulcast: true,
|
||||
//! };
|
||||
//!
|
||||
//! let manager = SfuManager::new(config);
|
||||
//! let room = manager.create_room("room_id").await?;
|
||||
//! let peer = room.add_peer("user_id", peer_connection).await?;
|
||||
//! ```
|
||||
|
||||
mod config;
|
||||
mod manager;
|
||||
mod peer;
|
||||
mod room;
|
||||
mod track;
|
||||
mod types;
|
||||
|
||||
pub use config::SfuConfig;
|
||||
pub use manager::SfuManager;
|
||||
pub use peer::{SfuPeer, PeerStats};
|
||||
pub use room::{SfuRoom, RoomMode, RoomStats};
|
||||
pub use track::{MediaTrack, QualityLayer, TrackKind};
|
||||
pub use types::{PeerId, RoomId, TrackId};
|
||||
@ -0,0 +1,65 @@
|
||||
//! SFU Manager
|
||||
|
||||
use crate::config::SfuConfig;
|
||||
use crate::room::{RoomStats, SfuRoom};
|
||||
use crate::types::{PeerId, RoomId};
|
||||
use anyhow::Result;
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
pub struct SfuManager {
|
||||
config: Arc<SfuConfig>,
|
||||
rooms: Arc<RwLock<HashMap<RoomId, Arc<SfuRoom>>>>,
|
||||
}
|
||||
|
||||
impl SfuManager {
|
||||
pub fn new(config: SfuConfig) -> Self {
|
||||
Self {
|
||||
config: Arc::new(config),
|
||||
rooms: Arc::new(RwLock::new(HashMap::new())),
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn get_or_create_room(&self, room_id: RoomId) -> Result<Arc<SfuRoom>> {
|
||||
let mut rooms = self.rooms.write().await;
|
||||
if let Some(room) = rooms.get(&room_id) {
|
||||
return Ok(room.clone());
|
||||
}
|
||||
|
||||
let room = Arc::new(SfuRoom::new(room_id.clone(), self.config.clone()));
|
||||
rooms.insert(room_id, room.clone());
|
||||
Ok(room)
|
||||
}
|
||||
|
||||
pub async fn add_peer_to_room(&self, room_id: RoomId, peer_id: PeerId) -> Result<()> {
|
||||
let room = self.get_or_create_room(room_id).await?;
|
||||
room.add_peer(peer_id).await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn remove_peer_from_room(&self, room_id: &RoomId, peer_id: &PeerId) -> Result<()> {
|
||||
let rooms = self.rooms.read().await;
|
||||
if let Some(room) = rooms.get(room_id) {
|
||||
room.remove_peer(peer_id).await?;
|
||||
if room.is_empty().await {
|
||||
drop(rooms);
|
||||
self.rooms.write().await.remove(room_id);
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn get_room_stats(&self, room_id: &RoomId) -> Result<RoomStats> {
|
||||
let rooms = self.rooms.read().await;
|
||||
if let Some(room) = rooms.get(room_id) {
|
||||
Ok(room.get_stats().await)
|
||||
} else {
|
||||
Ok(RoomStats::default())
|
||||
}
|
||||
}
|
||||
|
||||
pub fn config(&self) -> &SfuConfig {
|
||||
&self.config
|
||||
}
|
||||
}
|
||||
@ -0,0 +1,20 @@
|
||||
//! SFU Peer management
|
||||
|
||||
use crate::types::PeerId;
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
pub struct SfuPeer {
|
||||
pub id: PeerId,
|
||||
}
|
||||
|
||||
impl SfuPeer {
|
||||
pub fn new(id: PeerId) -> Self {
|
||||
Self { id }
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct PeerStats {
|
||||
pub packets_received: u64,
|
||||
pub bytes_received: u64,
|
||||
}
|
||||
@ -0,0 +1,83 @@
|
||||
//! SFU Room management
|
||||
|
||||
use crate::config::SfuConfig;
|
||||
use crate::peer::SfuPeer;
|
||||
use crate::types::{PeerId, RoomId};
|
||||
use anyhow::Result;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::collections::HashMap;
|
||||
use std::sync::Arc;
|
||||
use tokio::sync::RwLock;
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
pub enum RoomMode {
|
||||
P2P,
|
||||
SFU,
|
||||
}
|
||||
|
||||
pub struct SfuRoom {
|
||||
pub id: RoomId,
|
||||
pub mode: Arc<RwLock<RoomMode>>,
|
||||
pub peers: Arc<RwLock<HashMap<PeerId, Arc<SfuPeer>>>>,
|
||||
pub config: Arc<SfuConfig>,
|
||||
}
|
||||
|
||||
impl SfuRoom {
|
||||
pub fn new(id: RoomId, config: Arc<SfuConfig>) -> Self {
|
||||
Self {
|
||||
id,
|
||||
mode: Arc::new(RwLock::new(RoomMode::P2P)),
|
||||
peers: Arc::new(RwLock::new(HashMap::new())),
|
||||
config,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn add_peer(&self, peer_id: PeerId) -> Result<Arc<SfuPeer>> {
|
||||
let peer = Arc::new(SfuPeer::new(peer_id.clone()));
|
||||
self.peers.write().await.insert(peer_id, peer.clone());
|
||||
self.check_mode_switch().await?;
|
||||
Ok(peer)
|
||||
}
|
||||
|
||||
pub async fn remove_peer(&self, peer_id: &PeerId) -> Result<()> {
|
||||
self.peers.write().await.remove(peer_id);
|
||||
self.check_mode_switch().await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn peer_count(&self) -> usize {
|
||||
self.peers.read().await.len()
|
||||
}
|
||||
|
||||
async fn check_mode_switch(&self) -> Result<()> {
|
||||
let count = self.peer_count().await;
|
||||
let threshold = self.config.sfu_threshold;
|
||||
let mut mode = self.mode.write().await;
|
||||
|
||||
if count >= threshold && *mode == RoomMode::P2P {
|
||||
*mode = RoomMode::SFU;
|
||||
} else if count < threshold && *mode == RoomMode::SFU {
|
||||
*mode = RoomMode::P2P;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn is_empty(&self) -> bool {
|
||||
self.peers.read().await.is_empty()
|
||||
}
|
||||
|
||||
pub async fn get_stats(&self) -> RoomStats {
|
||||
RoomStats {
|
||||
peer_count: self.peer_count().await,
|
||||
..Default::default()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||
pub struct RoomStats {
|
||||
pub peer_count: usize,
|
||||
pub total_peers_joined: u64,
|
||||
pub mode_switches: u64,
|
||||
}
|
||||
@ -0,0 +1,376 @@
|
||||
//! Media track management for SFU
|
||||
//!
|
||||
//! This module handles complete WebRTC media track lifecycle including:
|
||||
//! - Track creation and lifecycle management
|
||||
//! - RTP packet reception and forwarding
|
||||
//! - Simulcast quality layer handling
|
||||
//! - Track statistics collection
|
||||
|
||||
use crate::types::{PeerId, TrackId};
|
||||
use anyhow::Result;
|
||||
use bytes::Bytes;
|
||||
use parking_lot::RwLock;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::sync::atomic::{AtomicU64, Ordering};
|
||||
use std::sync::Arc;
|
||||
use std::time::{Duration, Instant};
|
||||
use tokio::sync::mpsc;
|
||||
use tracing::{debug, error, info};
|
||||
use webrtc::rtp_transceiver::rtp_receiver::RTCRtpReceiver;
|
||||
use webrtc::track::track_remote::TrackRemote;
|
||||
use webrtc::util::marshal::MarshalSize;
|
||||
|
||||
/// Media track kind
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum TrackKind {
|
||||
Audio,
|
||||
Video,
|
||||
}
|
||||
|
||||
impl From<webrtc::rtp_transceiver::rtp_codec::RTPCodecType> for TrackKind {
|
||||
fn from(codec_type: webrtc::rtp_transceiver::rtp_codec::RTPCodecType) -> Self {
|
||||
match codec_type {
|
||||
webrtc::rtp_transceiver::rtp_codec::RTPCodecType::Audio => TrackKind::Audio,
|
||||
webrtc::rtp_transceiver::rtp_codec::RTPCodecType::Video => TrackKind::Video,
|
||||
_ => TrackKind::Video, // Default to video
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<&str> for TrackKind {
|
||||
fn from(s: &str) -> Self {
|
||||
match s.to_lowercase().as_str() {
|
||||
"audio" => TrackKind::Audio,
|
||||
"video" | _ => TrackKind::Video,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// Simulcast quality layer
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
#[serde(rename_all = "lowercase")]
|
||||
pub enum QualityLayer {
|
||||
High,
|
||||
Medium,
|
||||
Low,
|
||||
}
|
||||
|
||||
impl QualityLayer {
|
||||
/// Select quality layer based on available bandwidth
|
||||
/// bandwidth in kbps
|
||||
pub fn from_bandwidth(bandwidth_kbps: u32) -> Self {
|
||||
if bandwidth_kbps >= 2000 {
|
||||
QualityLayer::High // >= 2 Mbps
|
||||
} else if bandwidth_kbps >= 1000 {
|
||||
QualityLayer::Medium // >= 1 Mbps
|
||||
} else {
|
||||
QualityLayer::Low // < 1 Mbps
|
||||
}
|
||||
}
|
||||
|
||||
/// Get the RID (restriction identifier) for this layer
|
||||
pub fn rid(&self) -> &'static str {
|
||||
match self {
|
||||
QualityLayer::High => "h",
|
||||
QualityLayer::Medium => "m",
|
||||
QualityLayer::Low => "l",
|
||||
}
|
||||
}
|
||||
|
||||
/// Get expected bitrate for this layer (kbps)
|
||||
pub fn expected_bitrate(&self) -> u32 {
|
||||
match self {
|
||||
QualityLayer::High => 2500, // 2.5 Mbps
|
||||
QualityLayer::Medium => 1200, // 1.2 Mbps
|
||||
QualityLayer::Low => 500, // 500 kbps
|
||||
}
|
||||
}
|
||||
|
||||
/// Get spatial layer index (for SVC/Simulcast)
|
||||
pub fn spatial_layer(&self) -> u8 {
|
||||
match self {
|
||||
QualityLayer::High => 2,
|
||||
QualityLayer::Medium => 1,
|
||||
QualityLayer::Low => 0,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// RTP packet with metadata for forwarding
|
||||
#[derive(Debug, Clone)]
|
||||
pub struct ForwardablePacket {
|
||||
/// RTP packet data
|
||||
pub data: Bytes,
|
||||
|
||||
/// Source SSRC
|
||||
pub ssrc: u32,
|
||||
|
||||
/// Sequence number
|
||||
pub sequence_number: u16,
|
||||
|
||||
/// Timestamp
|
||||
pub timestamp: u32,
|
||||
|
||||
/// Quality layer (for simulcast)
|
||||
pub quality_layer: Option<QualityLayer>,
|
||||
|
||||
/// When packet was received
|
||||
pub received_at: Instant,
|
||||
}
|
||||
|
||||
/// Media track in the SFU
|
||||
pub struct MediaTrack {
|
||||
/// Track ID
|
||||
pub id: TrackId,
|
||||
|
||||
/// Owner peer ID
|
||||
pub peer_id: PeerId,
|
||||
|
||||
/// Track kind (audio/video)
|
||||
pub kind: TrackKind,
|
||||
|
||||
/// Remote track from WebRTC
|
||||
pub remote_track: Arc<TrackRemote>,
|
||||
|
||||
/// RTP receiver
|
||||
pub receiver: Arc<RTCRtpReceiver>,
|
||||
|
||||
/// Current active quality layer (for simulcast video)
|
||||
pub active_quality_layer: Arc<RwLock<Option<QualityLayer>>>,
|
||||
|
||||
/// Whether this track is active
|
||||
pub active: Arc<RwLock<bool>>,
|
||||
|
||||
/// Track statistics
|
||||
stats: Arc<TrackStatsInner>,
|
||||
|
||||
/// Packet forwarding channel
|
||||
packet_tx: Option<mpsc::UnboundedSender<ForwardablePacket>>,
|
||||
}
|
||||
|
||||
/// Internal track statistics with atomic counters
|
||||
struct TrackStatsInner {
|
||||
packets_received: AtomicU64,
|
||||
bytes_received: AtomicU64,
|
||||
packets_sent: AtomicU64,
|
||||
bytes_sent: AtomicU64,
|
||||
packets_lost: AtomicU64,
|
||||
last_packet_time: RwLock<Option<Instant>>,
|
||||
}
|
||||
|
||||
impl MediaTrack {
|
||||
/// Create a new media track
|
||||
pub fn new(
|
||||
id: TrackId,
|
||||
peer_id: PeerId,
|
||||
remote_track: Arc<TrackRemote>,
|
||||
receiver: Arc<RTCRtpReceiver>,
|
||||
) -> Self {
|
||||
let kind = TrackKind::from(remote_track.kind());
|
||||
|
||||
info!(
|
||||
track_id = %id,
|
||||
peer_id = %peer_id,
|
||||
kind = ?kind,
|
||||
codec = %remote_track.codec().capability.mime_type,
|
||||
"Creating media track"
|
||||
);
|
||||
|
||||
Self {
|
||||
id,
|
||||
peer_id,
|
||||
kind,
|
||||
remote_track,
|
||||
receiver,
|
||||
active_quality_layer: Arc::new(RwLock::new(None)),
|
||||
active: Arc::new(RwLock::new(true)),
|
||||
stats: Arc::new(TrackStatsInner {
|
||||
packets_received: AtomicU64::new(0),
|
||||
bytes_received: AtomicU64::new(0),
|
||||
packets_sent: AtomicU64::new(0),
|
||||
bytes_sent: AtomicU64::new(0),
|
||||
packets_lost: AtomicU64::new(0),
|
||||
last_packet_time: RwLock::new(None),
|
||||
}),
|
||||
packet_tx: None,
|
||||
}
|
||||
}
|
||||
|
||||
/// Start reading RTP packets from the track
|
||||
pub async fn start_reading(
|
||||
&mut self,
|
||||
) -> Result<mpsc::UnboundedReceiver<ForwardablePacket>> {
|
||||
let (packet_tx, packet_rx) = mpsc::unbounded_channel();
|
||||
self.packet_tx = Some(packet_tx.clone());
|
||||
|
||||
let track = Arc::clone(&self.remote_track);
|
||||
let stats = Arc::clone(&self.stats);
|
||||
let track_id = self.id.clone();
|
||||
let quality_layer = Arc::clone(&self.active_quality_layer);
|
||||
let active = Arc::clone(&self.active);
|
||||
|
||||
// Spawn RTP packet reading task
|
||||
tokio::spawn(async move {
|
||||
let mut buf = vec![0u8; 1500]; // MTU size
|
||||
|
||||
loop {
|
||||
// Check if track is still active
|
||||
if !*active.read() {
|
||||
debug!(track_id = %track_id, "Track deactivated, stopping RTP reader");
|
||||
break;
|
||||
}
|
||||
|
||||
// Read RTP packet
|
||||
match track.read(&mut buf).await {
|
||||
Ok((rtp_packet, _attributes)) => {
|
||||
// Update statistics
|
||||
let packet_size = rtp_packet.header.marshal_size() + rtp_packet.payload.len();
|
||||
stats.packets_received.fetch_add(1, Ordering::Relaxed);
|
||||
stats.bytes_received.fetch_add(packet_size as u64, Ordering::Relaxed);
|
||||
*stats.last_packet_time.write() = Some(Instant::now());
|
||||
|
||||
// Create forwardable packet
|
||||
let forwardable = ForwardablePacket {
|
||||
data: Bytes::copy_from_slice(&buf[..packet_size]),
|
||||
ssrc: rtp_packet.header.ssrc,
|
||||
sequence_number: rtp_packet.header.sequence_number,
|
||||
timestamp: rtp_packet.header.timestamp,
|
||||
quality_layer: *quality_layer.read(),
|
||||
received_at: Instant::now(),
|
||||
};
|
||||
|
||||
// Forward packet to subscribers
|
||||
if let Err(e) = packet_tx.send(forwardable) {
|
||||
error!(
|
||||
track_id = %track_id,
|
||||
error = %e,
|
||||
"Failed to forward RTP packet"
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
Err(e) => {
|
||||
error!(
|
||||
track_id = %track_id,
|
||||
error = %e,
|
||||
"Failed to read RTP packet"
|
||||
);
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
info!(track_id = %track_id, "RTP reader stopped");
|
||||
});
|
||||
|
||||
Ok(packet_rx)
|
||||
}
|
||||
|
||||
/// Get track SSRC (Synchronization Source)
|
||||
pub fn ssrc(&self) -> u32 {
|
||||
self.remote_track.ssrc()
|
||||
}
|
||||
|
||||
/// Get track codec
|
||||
pub fn codec(&self) -> String {
|
||||
self.remote_track.codec().capability.mime_type.clone()
|
||||
}
|
||||
|
||||
/// Set active quality layer for simulcast
|
||||
pub fn set_quality_layer(&self, layer: QualityLayer) {
|
||||
let mut current = self.active_quality_layer.write();
|
||||
if *current != Some(layer) {
|
||||
debug!(
|
||||
track_id = %self.id,
|
||||
old_layer = ?*current,
|
||||
new_layer = ?layer,
|
||||
"Switching quality layer"
|
||||
);
|
||||
*current = Some(layer);
|
||||
}
|
||||
}
|
||||
|
||||
/// Get active quality layer
|
||||
pub fn quality_layer(&self) -> Option<QualityLayer> {
|
||||
*self.active_quality_layer.read()
|
||||
}
|
||||
|
||||
/// Check if track is video
|
||||
pub fn is_video(&self) -> bool {
|
||||
self.kind == TrackKind::Video
|
||||
}
|
||||
|
||||
/// Check if track is audio
|
||||
pub fn is_audio(&self) -> bool {
|
||||
self.kind == TrackKind::Audio
|
||||
}
|
||||
|
||||
/// Check if track is active
|
||||
pub fn is_active(&self) -> bool {
|
||||
*self.active.read()
|
||||
}
|
||||
|
||||
/// Deactivate track
|
||||
pub fn deactivate(&self) {
|
||||
*self.active.write() = false;
|
||||
}
|
||||
|
||||
/// Get track statistics
|
||||
pub fn get_stats(&self) -> TrackStats {
|
||||
let packets_received = self.stats.packets_received.load(Ordering::Relaxed);
|
||||
let bytes_received = self.stats.bytes_received.load(Ordering::Relaxed);
|
||||
let packets_sent = self.stats.packets_sent.load(Ordering::Relaxed);
|
||||
let bytes_sent = self.stats.bytes_sent.load(Ordering::Relaxed);
|
||||
let packets_lost = self.stats.packets_lost.load(Ordering::Relaxed);
|
||||
|
||||
// Calculate bitrate (over last second)
|
||||
let bitrate_kbps = if let Some(last_time) = *self.stats.last_packet_time.read() {
|
||||
let elapsed = Instant::now().duration_since(last_time);
|
||||
if elapsed < Duration::from_secs(1) {
|
||||
((bytes_received * 8) as f64 / elapsed.as_secs_f64() / 1000.0) as u32
|
||||
} else {
|
||||
0
|
||||
}
|
||||
} else {
|
||||
0
|
||||
};
|
||||
|
||||
TrackStats {
|
||||
track_id: self.id.as_str().to_string(),
|
||||
kind: self.kind,
|
||||
packets_received,
|
||||
bytes_received,
|
||||
packets_sent,
|
||||
bytes_sent,
|
||||
packets_lost,
|
||||
bitrate_kbps,
|
||||
quality_layer: self.quality_layer(),
|
||||
}
|
||||
}
|
||||
|
||||
/// Update sent packet statistics
|
||||
pub fn record_sent_packet(&self, packet_size: usize) {
|
||||
self.stats.packets_sent.fetch_add(1, Ordering::Relaxed);
|
||||
self.stats.bytes_sent.fetch_add(packet_size as u64, Ordering::Relaxed);
|
||||
}
|
||||
|
||||
/// Record packet loss
|
||||
pub fn record_packet_loss(&self, count: u64) {
|
||||
self.stats.packets_lost.fetch_add(count, Ordering::Relaxed);
|
||||
}
|
||||
}
|
||||
|
||||
/// Track statistics
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
pub struct TrackStats {
|
||||
pub track_id: String,
|
||||
pub kind: TrackKind,
|
||||
pub packets_received: u64,
|
||||
pub bytes_received: u64,
|
||||
pub packets_sent: u64,
|
||||
pub bytes_sent: u64,
|
||||
pub packets_lost: u64,
|
||||
pub bitrate_kbps: u32,
|
||||
pub quality_layer: Option<QualityLayer>,
|
||||
}
|
||||
@ -0,0 +1,100 @@
|
||||
//! Common types used throughout the SFU implementation
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
use std::fmt;
|
||||
|
||||
/// Unique identifier for a peer in the SFU
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
pub struct PeerId(String);
|
||||
|
||||
impl PeerId {
|
||||
pub fn new(id: impl Into<String>) -> Self {
|
||||
Self(id.into())
|
||||
}
|
||||
|
||||
pub fn as_str(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for PeerId {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
write!(f, "{}", self.0)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<String> for PeerId {
|
||||
fn from(s: String) -> Self {
|
||||
Self(s)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<&str> for PeerId {
|
||||
fn from(s: &str) -> Self {
|
||||
Self(s.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
/// Unique identifier for an SFU room
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
pub struct RoomId(String);
|
||||
|
||||
impl RoomId {
|
||||
pub fn new(id: impl Into<String>) -> Self {
|
||||
Self(id.into())
|
||||
}
|
||||
|
||||
pub fn as_str(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for RoomId {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
write!(f, "{}", self.0)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<String> for RoomId {
|
||||
fn from(s: String) -> Self {
|
||||
Self(s)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<&str> for RoomId {
|
||||
fn from(s: &str) -> Self {
|
||||
Self(s.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
/// Unique identifier for a media track
|
||||
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
|
||||
pub struct TrackId(String);
|
||||
|
||||
impl TrackId {
|
||||
pub fn new(id: impl Into<String>) -> Self {
|
||||
Self(id.into())
|
||||
}
|
||||
|
||||
pub fn as_str(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for TrackId {
|
||||
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
write!(f, "{}", self.0)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<String> for TrackId {
|
||||
fn from(s: String) -> Self {
|
||||
Self(s)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<&str> for TrackId {
|
||||
fn from(s: &str) -> Self {
|
||||
Self(s.to_string())
|
||||
}
|
||||
}
|
||||
Loading…
Reference in New Issue