Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

3 changes: 3 additions & 0 deletions livekit-api/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -103,3 +103,6 @@ isahc = { version = "1.7.2", default-features = false, features = [ "json", "tex

scopeguard = "1.2.0"
rand = { workspace = true }

[dev-dependencies]
tokio = { workspace = true, features = ["rt", "rt-multi-thread", "net", "time", "macros", "io-util"] }
142 changes: 130 additions & 12 deletions livekit-api/src/signal_client/mod.rs
Original file line number Diff line number Diff line change
Expand Up @@ -49,6 +49,9 @@ pub type SignalEvents = mpsc::UnboundedReceiver<SignalEvent>;
pub type SignalResult<T> = Result<T, SignalError>;

pub const JOIN_RESPONSE_TIMEOUT: Duration = Duration::from_secs(5);
pub const SIGNAL_CONNECT_TIMEOUT: Duration = Duration::from_secs(5);
const REGION_FETCH_TIMEOUT: Duration = Duration::from_secs(3);
const VALIDATE_TIMEOUT: Duration = Duration::from_secs(3);
Comment thread
davidzhao marked this conversation as resolved.
pub const PROTOCOL_VERSION: u32 = 16;

#[derive(Error, Debug)]
Expand Down Expand Up @@ -96,6 +99,8 @@ pub struct SignalOptions {
pub sdk_options: SignalSdkOptions,
/// Enable single peer connection mode
pub single_peer_connection: bool,
/// Timeout for each individual signal connection attempt
pub connect_timeout: Duration,
}

impl Default for SignalOptions {
Expand All @@ -105,6 +110,7 @@ impl Default for SignalOptions {
adaptive_stream: false,
sdk_options: SignalSdkOptions::default(),
single_peer_connection: true,
connect_timeout: SIGNAL_CONNECT_TIMEOUT,
}
}
}
Expand Down Expand Up @@ -268,7 +274,7 @@ impl SignalInner {
let lk_url = get_livekit_url(url, &options, use_v1_path, false, None, "")?;
// Try to connect to the SignalClient
let (stream, mut events, single_pc_mode_active) =
match SignalStream::connect(lk_url.clone(), token).await {
match SignalStream::connect(lk_url.clone(), token, options.connect_timeout).await {
Ok((new_stream, stream_events)) => {
log::debug!(
"signal connection successful: path={}, single_pc_mode={}",
Expand Down Expand Up @@ -297,7 +303,13 @@ impl SignalInner {
if use_v1_path && is_not_found {
let lk_url_v0 = get_livekit_url(url, &options, false, false, None, "")?;
log::warn!("v1 path not found (404), falling back to v0 path");
match SignalStream::connect(lk_url_v0.clone(), token).await {
match SignalStream::connect(
lk_url_v0.clone(),
token,
options.connect_timeout,
)
.await
{
Ok((new_stream, stream_events)) => (new_stream, stream_events, false),
Err(err) => {
log::error!("v0 fallback also failed: {:?}", err);
Expand Down Expand Up @@ -338,18 +350,24 @@ impl SignalInner {
async fn validate(ws_url: url::Url) -> SignalResult<()> {
let validate_url = get_validate_url(ws_url);

if let Ok(res) = http_client::get(validate_url.as_str()).await {
let status = res.status();
let body = res.text().await.ok().unwrap_or_default();
let validate_fut = async {
if let Ok(res) = http_client::get(validate_url.as_str()).await {
let status = res.status();
let body = res.text().await.ok().unwrap_or_default();

if status.is_client_error() {
return Err(SignalError::Client(status, body));
} else if status.is_server_error() {
return Err(SignalError::Server(status, body));
if status.is_client_error() {
return Err(SignalError::Client(status, body));
} else if status.is_server_error() {
return Err(SignalError::Server(status, body));
}
}
}

Ok(())
Ok(())
};

livekit_runtime::timeout(VALIDATE_TIMEOUT, validate_fut)
.await
.map_err(|_| SignalError::Timeout("validate request timed out".into()))?
}

/// Returns whether single peer connection mode is active
Expand Down Expand Up @@ -383,7 +401,8 @@ impl SignalInner {
get_livekit_url(&self.url, &self.options, self.single_pc_mode_active, true, None, sid)
.unwrap();

let (new_stream, mut events) = SignalStream::connect(lk_url, &token).await?;
let (new_stream, mut events) =
SignalStream::connect(lk_url, &token, self.options.connect_timeout).await?;
let reconnect_response = get_reconnect_response(&mut events).await?;
*stream = Some(new_stream);

Expand Down Expand Up @@ -755,4 +774,103 @@ mod tests {
assert_eq!(validate_url.path(), "/rtc/validate");
assert_eq!(validate_url.scheme(), "https");
}

#[cfg(feature = "signal-client-tokio")]
#[tokio::test]
async fn signal_stream_connect_timeout() {
use tokio::net::TcpListener;

// Bind a TCP listener that accepts connections but never sends data
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();

// Spawn a task that accepts connections but does nothing (simulates a hanging server)
let _accept_task = tokio::spawn(async move {
loop {
let Ok((_socket, _)) = listener.accept().await else {
break;
};
// Hold the connection open but never write anything
tokio::time::sleep(Duration::from_secs(60)).await;
}
});

let url = url::Url::parse(&format!("ws://127.0.0.1:{}", addr.port())).unwrap();
let result = SignalStream::connect(url, "fake-token", Duration::from_millis(500)).await;

assert!(result.is_err());
let err = result.unwrap_err();
assert!(matches!(err, SignalError::Timeout(_)), "expected Timeout error, got: {:?}", err);
}

#[cfg(feature = "signal-client-tokio")]
#[tokio::test]
async fn region_fetch_parses_response() {
use tokio::io::AsyncWriteExt;
use tokio::net::TcpListener;

let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();

// Spawn a task that serves a hand-crafted HTTP response with region JSON
tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();

// Read the request (consume it so the connection doesn't stall)
let mut buf = [0u8; 4096];
let _ = tokio::io::AsyncReadExt::read(&mut socket, &mut buf).await;

let body = r#"{"regions":[{"region":"us-east-1","url":"wss://us-east.livekit.cloud","distance":"100"},{"region":"eu-west-1","url":"wss://eu-west.livekit.cloud","distance":"200"}]}"#;

let response = format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\n\r\n{}",
body.len(),
body
);
socket.write_all(response.as_bytes()).await.unwrap();
});

let endpoint = format!("http://127.0.0.1:{}/settings/regions", addr.port());
let result = region::fetch_from_endpoint(&endpoint, "fake-token").await;

let urls = result.unwrap();
assert_eq!(
urls,
vec![
"wss://us-east.livekit.cloud".to_string(),
"wss://eu-west.livekit.cloud".to_string(),
]
);
}

#[cfg(feature = "signal-client-tokio")]
#[tokio::test]
async fn region_fetch_timeout() {
use tokio::net::TcpListener;

// Bind a listener that accepts but never responds
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();

tokio::spawn(async move {
loop {
let Ok((_socket, _)) = listener.accept().await else {
break;
};
// Hold connection open, never write
tokio::time::sleep(Duration::from_secs(60)).await;
}
});

let endpoint = format!("http://127.0.0.1:{}/settings/regions", addr.port());
let result = region::fetch_from_endpoint(&endpoint, "fake-token").await;

assert!(result.is_err());
let err = result.unwrap_err();
assert!(
matches!(err, SignalError::RegionError(ref msg) if msg.contains("timed out")),
"expected RegionError with 'timed out', got: {:?}",
err
);
}
}
97 changes: 72 additions & 25 deletions livekit-api/src/signal_client/region.rs
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,7 @@ use serde::Deserialize;

use crate::http_client;

use super::{get_livekit_url, SignalError, SignalResult};
use super::{SignalError, SignalResult, REGION_FETCH_TIMEOUT};

pub struct RegionUrlProvider;

Expand All @@ -36,36 +36,44 @@ pub struct RegionUrlInfo {
impl RegionUrlProvider {
pub async fn fetch_region_urls(url: &str, token: &str) -> SignalResult<Vec<String>> {
if is_cloud_url(url)? {
let client = http_client::Client::new();
let mut headers = HeaderMap::new();
headers.insert(
AUTHORIZATION,
HeaderValue::from_str(&format!("Bearer {}", token)).unwrap(),
);
let res = client
.get(region_endpoint(url)?)
.headers(headers)
.send()
.await
.map_err(|e| SignalError::RegionError(e.to_string()))?;

if !res.status().is_success() {
return Err(SignalError::Client(
res.status(),
res.text().await.unwrap_or_default(),
));
}
let res = res
.json::<RegionUrlResponse>()
.await
.map_err(|e| SignalError::RegionError(e.to_string()))?;
Ok(res.regions.into_iter().map(|i| i.url).collect())
let endpoint = region_endpoint(url)?;
fetch_from_endpoint(&endpoint, token).await
} else {
Ok(vec![])
}
}
}

pub(crate) async fn fetch_from_endpoint(
endpoint_url: &str,
token: &str,
) -> SignalResult<Vec<String>> {
let fetch_fut = async {
let client = http_client::Client::new();
let mut headers = HeaderMap::new();
headers.insert(AUTHORIZATION, HeaderValue::from_str(&format!("Bearer {}", token)).unwrap());
let res = client
.get(endpoint_url)
.headers(headers)
.send()
.await
.map_err(|e| SignalError::RegionError(e.to_string()))?;

if !res.status().is_success() {
return Err(SignalError::Client(res.status(), res.text().await.unwrap_or_default()));
}
let res = res
.json::<RegionUrlResponse>()
.await
.map_err(|e| SignalError::RegionError(e.to_string()))?;
Ok(res.regions.into_iter().map(|i| i.url).collect())
};

livekit_runtime::timeout(REGION_FETCH_TIMEOUT, fetch_fut)
.await
.map_err(|_| SignalError::RegionError("region fetch timed out".into()))?
}

fn is_cloud_url(url: &str) -> SignalResult<bool> {
let url = url::Url::parse(url).map_err(|err| SignalError::UrlParse(err.to_string()))?;
let host = match url.host() {
Expand All @@ -89,3 +97,42 @@ fn region_endpoint(url: &str) -> SignalResult<String> {

Ok(url.to_string())
}

#[cfg(test)]
mod tests {
use super::*;

#[test]
fn test_is_cloud_url() {
assert!(is_cloud_url("wss://myapp.livekit.cloud").unwrap());
assert!(is_cloud_url("wss://myapp.livekit.run").unwrap());
assert!(is_cloud_url("https://myapp.livekit.cloud").unwrap());

assert!(!is_cloud_url("wss://localhost:7880").unwrap());
assert!(!is_cloud_url("wss://example.com").unwrap());
assert!(!is_cloud_url("wss://livekit.cloud.example.com").unwrap());
}

#[test]
fn test_region_endpoint() {
assert_eq!(
region_endpoint("wss://myapp.livekit.cloud").unwrap(),
"https://myapp.livekit.cloud/settings/regions"
);
assert_eq!(
region_endpoint("ws://myapp.livekit.run").unwrap(),
"http://myapp.livekit.run/settings/regions"
);
assert_eq!(
region_endpoint("https://myapp.livekit.cloud").unwrap(),
"https://myapp.livekit.cloud/settings/regions"
);
}

#[tokio::test]
async fn test_fetch_non_cloud_url_returns_empty() {
let result =
RegionUrlProvider::fetch_region_urls("wss://localhost:7880", "fake-token").await;
assert_eq!(result.unwrap(), Vec::<String>::new());
}
}
13 changes: 12 additions & 1 deletion livekit-api/src/signal_client/signal_stream.rs
Original file line number Diff line number Diff line change
Expand Up @@ -19,7 +19,7 @@ use futures_util::{
use livekit_protocol as proto;
use livekit_runtime::{JoinHandle, TcpStream};
use prost::Message as ProtoMessage;
use std::{env, io};
use std::{env, io, time::Duration};

use tokio::sync::{mpsc, oneshot};

Expand Down Expand Up @@ -88,6 +88,17 @@ impl SignalStream {
pub async fn connect(
url: url::Url,
token: &str,
connect_timeout: Duration,
) -> SignalResult<(Self, mpsc::UnboundedReceiver<Box<proto::signal_response::Message>>)> {
let connect_fut = Self::connect_inner(url, token);
livekit_runtime::timeout(connect_timeout, connect_fut)
.await
.map_err(|_| SignalError::Timeout("signal connection timed out".into()))?
}

async fn connect_inner(
url: url::Url,
token: &str,
) -> SignalResult<(Self, mpsc::UnboundedReceiver<Box<proto::signal_response::Message>>)> {
log::info!("connecting to {}", url);
let mut request = url.clone().into_client_request()?;
Expand Down
8 changes: 8 additions & 0 deletions livekit-ffi-node-bindings/src/proto/room_pb.ts
Original file line number Diff line number Diff line change
Expand Up @@ -2478,6 +2478,13 @@ export class RoomOptions extends Message<RoomOptions> {
*/
singlePeerConnection?: boolean;

/**
* timeout in milliseconds for each signal connection attempt (default: 5000)
*
* @generated from field: optional uint64 connect_timeout_ms = 9;
*/
connectTimeoutMs?: bigint;

constructor(data?: PartialMessage<RoomOptions>) {
super();
proto2.util.initPartial(data, this);
Expand All @@ -2494,6 +2501,7 @@ export class RoomOptions extends Message<RoomOptions> {
{ no: 6, name: "join_retries", kind: "scalar", T: 13 /* ScalarType.UINT32 */, opt: true },
{ no: 7, name: "encryption", kind: "message", T: E2eeOptions, opt: true },
{ no: 8, name: "single_peer_connection", kind: "scalar", T: 8 /* ScalarType.BOOL */, opt: true },
{ no: 9, name: "connect_timeout_ms", kind: "scalar", T: 4 /* ScalarType.UINT64 */, opt: true },
]);

static fromBinary(bytes: Uint8Array, options?: Partial<BinaryReadOptions>): RoomOptions {
Expand Down
1 change: 1 addition & 0 deletions livekit-ffi/protocol/room.proto
Original file line number Diff line number Diff line change
Expand Up @@ -311,6 +311,7 @@ message RoomOptions {
optional uint32 join_retries = 6;
optional E2eeOptions encryption = 7;
optional bool single_peer_connection = 8; // use single peer connection for both publish/subscribe (default: true)
optional uint64 connect_timeout_ms = 9; // timeout in milliseconds for each signal connection attempt (default: 5000)
}

//
Expand Down
Loading