From 8fbbe4b15e3980bf4dbbd9a93f611bf4858c9fcd Mon Sep 17 00:00:00 2001 From: euzu Date: Mon, 28 Apr 2025 19:52:59 +0200 Subject: [PATCH] provider and user connection management --- src/api/api_utils.rs | 82 ++++++++++++------- src/api/endpoints/hls_api.rs | 7 +- src/api/endpoints/m3u_api.rs | 2 +- src/api/endpoints/v1_api.rs | 1 - src/api/endpoints/xtream_api.rs | 2 +- src/api/model/active_provider_manager.rs | 21 ++--- src/api/model/active_user_manager.rs | 38 ++++++++- src/api/model/streams/active_client_stream.rs | 2 +- src/processing/parser/hls.rs | 5 +- src/repository/bplustree.rs | 17 +--- src/utils/string_utils.rs | 31 +++++++ 11 files changed, 139 insertions(+), 69 deletions(-) diff --git a/src/api/api_utils.rs b/src/api/api_utils.rs index f6413fcac..b53f176eb 100644 --- a/src/api/api_utils.rs +++ b/src/api/api_utils.rs @@ -482,16 +482,16 @@ fn create_session_cookie(cookie: &str) -> String { format!("{SESSION_COOKIE_NAME}{cookie}; Max-Age=10800; HttpOnly; SameSite=Strict") } -pub fn create_token_for_provider(secret: &[u8], virtual_id: u32, provider_name: &str, stream_url: &str) -> Option { - if let Ok(cookie_value) = obfuscate_text(secret, &format!("{virtual_id}:{provider_name}@{stream_url}")) { +pub fn create_token_for_provider(secret: &[u8], token: &str, virtual_id: u32, provider_name: &str, stream_url: &str) -> Option { + if let Ok(cookie_value) = obfuscate_text(secret, &format!("{virtual_id}:{token}:{provider_name}@{stream_url}")) { return Some(cookie_value); } None } -pub fn create_session_cookie_for_provider(secret: &[u8], virtual_id: u32, provider_name: &str, stream_url: &str) -> Option { - if let Some(cookie_value) = create_token_for_provider(secret, virtual_id, provider_name, stream_url) { +pub fn create_session_cookie_for_provider(secret: &[u8], token: &str, virtual_id: u32, provider_name: &str, stream_url: &str) -> Option { + if let Some(cookie_value) = create_token_for_provider(secret, token, virtual_id, provider_name, stream_url) { return Some(create_session_cookie(&cookie_value)); } None @@ -512,27 +512,30 @@ pub fn read_session_cookie(headers: &HeaderMap) -> Option { None } -pub fn get_stream_info_from_crypted_cookie(secret: &[u8], cookie: &str) -> Option<(u32, String, String)> { +/// # Panics +pub fn get_stream_info_from_crypted_cookie(secret: &[u8], cookie: &str) -> Option<(String, u32, String, String)> { if let Ok(decrypted) = deobfuscate_text(secret, cookie) { let (virtual_id_and_provider, stream_url) = decrypted.split_once('@')?; - let (virtual_id, provider_name) = virtual_id_and_provider.split_once(':')?; - - if virtual_id.is_empty() || provider_name.is_empty() || stream_url.is_empty() { + let mut items: Vec = virtual_id_and_provider.split(':').filter(|s| !s.is_empty()).map(ToString::to_string).collect(); + if items.len() != 3 { return None; } - - if let Ok(vid) = virtual_id.parse::() { - return Some(( - vid, - provider_name.to_string(), - stream_url.to_string(), - )); - } + items.reverse(); + let virtual_id = items.pop().unwrap(); + let vid = virtual_id.parse::().ok()?; + let token = items.pop().unwrap(); + let provider_name = items.pop().unwrap(); + return Some(( + token, + vid, + provider_name.to_string(), + stream_url.to_string(), + )); } None } -pub fn get_stream_info_from_cookie(secret: &[u8], headers: &HeaderMap) -> Option<(u32, String, String)> { +pub fn get_stream_info_from_cookie(secret: &[u8], headers: &HeaderMap) -> Option<(String, u32, String, String)> { let encrypted_cookie = read_session_cookie(headers)?; get_stream_info_from_crypted_cookie(secret, &encrypted_cookie) } @@ -564,8 +567,8 @@ pub async fn force_provider_stream_response(app_state: &AppState, req_headers: &HeaderMap, input: &ConfigInput, user: &ProxyUserCredentials) -> impl axum::response::IntoResponse + Send { - if let Some((stream_virtual_id, provider_name, stream_url)) = get_stream_info_from_crypted_cookie(&app_state.config.t_encrypt_secret, cookie) { - if stream_virtual_id == virtual_id { + if let Some((stream_token, stream_virtual_id, provider_name, stream_url)) = get_stream_info_from_crypted_cookie(&app_state.config.t_encrypt_secret, cookie) { + if stream_virtual_id == virtual_id && app_state.active_users.has_token(&user.username, &stream_token).await { let stream_options = get_stream_options(app_state); let share_stream = false; let connection_permission = UserConnectionPermission::Allowed; @@ -654,8 +657,9 @@ pub async fn stream_response(app_state: &AppState, } if let Some(provider) = provider_name { - if matches!(item_type, PlaylistItemType::LiveHls | PlaylistItemType::LiveDash) { - if let Some(cookie_value) = create_session_cookie_for_provider(&app_state.config.t_encrypt_secret, virtual_id, &provider, stream_url) { + if matches!(item_type, PlaylistItemType::LiveHls | PlaylistItemType::LiveDash | PlaylistItemType::Video | PlaylistItemType::Series) { + let grace_token = app_state.active_users.create_token(&user.username).await; + if let Some(cookie_value) = create_session_cookie_for_provider(&app_state.config.t_encrypt_secret, &grace_token, virtual_id, &provider, stream_url) { response = response.header(axum::http::header::SET_COOKIE, &cookie_value); } } @@ -663,7 +667,6 @@ pub async fn stream_response(app_state: &AppState, let body_stream = prepare_body_stream(app_state, item_type, stream); response.body(body_stream).unwrap().into_response() - // if content_length > 0 { response_builder.body(SizedStream::new(content_length, stream)) } else { response_builder.streaming(stream) } }; return stream_resp.into_response(); @@ -824,11 +827,12 @@ pub fn redirect(url: &str) -> impl IntoResponse { .unwrap() } -pub fn is_seek_response( +pub async fn is_seek_response( + app_state: &AppState, cluster: XtreamCluster, virtual_id: u32, - secret: &[u8], req_headers: &HeaderMap, + username: &str, ) -> Option { // seek only for non-live streams if cluster == XtreamCluster::Live { @@ -836,8 +840,12 @@ pub fn is_seek_response( } let cookie = read_session_cookie(req_headers)?; - match get_stream_info_from_crypted_cookie(secret, &cookie) { - Some((vid, _, _)) if vid == virtual_id => {} + match get_stream_info_from_crypted_cookie(&app_state.config.t_encrypt_secret, &cookie) { + Some((token, vid, _, _)) if vid == virtual_id => { + if !app_state.active_users.has_token(username, &token).await { + return None; + } + } _ => return None, } @@ -852,22 +860,34 @@ pub fn is_seek_response( pub async fn check_force_provider(app_state: &AppState, virtual_id: u32, item_type: PlaylistItemType, req_headers: &HeaderMap, user: &ProxyUserCredentials) -> (Option, UserConnectionPermission) { - if ! matches!(item_type, PlaylistItemType::LiveHls | PlaylistItemType::LiveDash) { + if ! matches!(item_type, PlaylistItemType::LiveHls | PlaylistItemType::LiveDash | PlaylistItemType::Series | PlaylistItemType::Video) { return (None, user.connection_permission(app_state).await); } // if you have multi provider setup you need to delegate the same hls requests // to the same provider. Hls has alternating m3u8 and stream requests. let mut provider_name = None; - if let Some((stream_virtual_id, stream_provider_name, _stream_url)) = get_stream_info_from_cookie(&app_state.config.t_encrypt_secret, req_headers) { - if stream_virtual_id == virtual_id { + if let Some((stream_token, stream_virtual_id, stream_provider_name, _stream_url)) = get_stream_info_from_cookie(&app_state.config.t_encrypt_secret, req_headers) { + if stream_virtual_id == virtual_id && app_state.active_users.has_token(&user.username, &stream_token).await { provider_name = Some(stream_provider_name); } } let connection_permission = match provider_name { - None => user.connection_permission(app_state).await, - Some(_) => UserConnectionPermission::Allowed + Some(_) => UserConnectionPermission::Allowed, + None => { + let permission = user.connection_permission(app_state).await; + match permission { + UserConnectionPermission::GracePeriod => { + if app_state.active_users.get_token(&user.username).await.is_some() { + UserConnectionPermission::Exhausted + } else { + UserConnectionPermission::GracePeriod + } + } + _ => permission, + } + } }; (provider_name, connection_permission) diff --git a/src/api/endpoints/hls_api.rs b/src/api/endpoints/hls_api.rs index df9b6fd76..bd3ed33ee 100644 --- a/src/api/endpoints/hls_api.rs +++ b/src/api/endpoints/hls_api.rs @@ -45,10 +45,12 @@ pub(in crate::api) async fn handle_hls_stream_request(app_state: &Arc, let url = replace_url_extension(hls_url, HLS_EXT); let server_info = app_state.config.get_user_server_info(user).await; + let grace_token = app_state.active_users.create_token(&user.username).await; let create_stream_and_cookie = |provider_cfg: &Arc| { let stream_url = get_stream_alternative_url(&url, input, provider_cfg); let cookie = create_session_cookie_for_provider( &app_state.config.t_encrypt_secret, + &grace_token, virtual_id, &provider_cfg.name, &stream_url, @@ -77,6 +79,7 @@ pub(in crate::api) async fn handle_hls_stream_request(app_state: &Arc, virtual_id, input_id: input.id, provider_name: provider.unwrap_or_default(), // this should not happen + user_token: grace_token }; let hls_content = rewrite_hls(user, &rewrite_hls_props); hls_response(hls_content, cookie).into_response() @@ -109,12 +112,12 @@ async fn hls_api_stream( return create_custom_video_stream_response(&app_state.config, &CustomVideoStreamType::UserConnectionsExhausted).into_response(); } - let Some((stream_virtual_id, stream_provider_name, hls_url)) = get_stream_info_from_crypted_cookie(&app_state.config.t_encrypt_secret, ¶ms.token) + let Some((stream_token, stream_virtual_id, stream_provider_name, hls_url)) = get_stream_info_from_crypted_cookie(&app_state.config.t_encrypt_secret, ¶ms.token) else { return bad_response_with_delete_cookie().into_response(); }; - if stream_virtual_id != virtual_id { + if stream_virtual_id != virtual_id || app_state.active_users.get_token(&user.username).await.is_some_and(|t| ! t.eq(&stream_token)) { return bad_response_with_delete_cookie().into_response(); } diff --git a/src/api/endpoints/m3u_api.rs b/src/api/endpoints/m3u_api.rs index 4ccef7a08..a87e89703 100644 --- a/src/api/endpoints/m3u_api.rs +++ b/src/api/endpoints/m3u_api.rs @@ -85,7 +85,7 @@ async fn m3u_api_stream( let cluster = XtreamCluster::try_from(pli.item_type).unwrap_or(XtreamCluster::Live); - if let Some(cookie) = is_seek_response(cluster, pli.virtual_id, &app_state.config.t_encrypt_secret, &req_headers) { + if let Some(cookie) = is_seek_response(&app_state, cluster, pli.virtual_id, &req_headers, &user.username).await { // partial request means we are in reverse proxy mode, seek happened return force_provider_stream_response(&app_state, &cookie, pli.virtual_id, pli.item_type, &req_headers, input, &user).await.into_response() } diff --git a/src/api/endpoints/v1_api.rs b/src/api/endpoints/v1_api.rs index 881c67523..960c7af24 100644 --- a/src/api/endpoints/v1_api.rs +++ b/src/api/endpoints/v1_api.rs @@ -325,7 +325,6 @@ async fn create_status_check(app_state: &Arc) -> StatusCheck { cache, } } -#[axum::debug_handler] async fn status(axum::extract::State(app_state): axum::extract::State>) -> axum::response::Response { let status = create_status_check(&app_state).await; match serde_json::to_string_pretty(&status) { diff --git a/src/api/endpoints/xtream_api.rs b/src/api/endpoints/xtream_api.rs index c36c363fc..40e806f24 100644 --- a/src/api/endpoints/xtream_api.rs +++ b/src/api/endpoints/xtream_api.rs @@ -197,7 +197,7 @@ async fn xtream_player_api_stream( let input = try_option_bad_request!(app_state.config.get_input_by_name(pli.input_name.as_str()), true, format!("Cant find input for target {target_name}, context {}, stream_id {virtual_id}", stream_req.context)); let cluster = pli.xtream_cluster; - if let Some(cookie) = is_seek_response(cluster, pli.virtual_id, &app_state.config.t_encrypt_secret, req_headers) { + if let Some(cookie) = is_seek_response(app_state, cluster, pli.virtual_id, req_headers, &user.username).await { // partial request means we are in reverse proxy mode, seek happened return force_provider_stream_response(app_state, &cookie, pli.virtual_id, pli.item_type, req_headers, input, &user).await.into_response() } diff --git a/src/api/model/active_provider_manager.rs b/src/api/model/active_provider_manager.rs index 2eb69f242..332573b23 100644 --- a/src/api/model/active_provider_manager.rs +++ b/src/api/model/active_provider_manager.rs @@ -14,6 +14,13 @@ pub struct ProviderConnectionGuard { } impl ProviderConnectionGuard { + pub fn new(manager: Arc, allocation: ProviderAllocation) -> Self { + Self { + manager, + allocation, + } + } + pub fn get_provider_name(&self) -> Option { match self.allocation { ProviderAllocation::Exhausted => None, @@ -466,10 +473,7 @@ impl ActiveProviderManager { Some((_lineup, config)) => config.force_allocate().await, }; - ProviderConnectionGuard { - manager: Arc::new(self.clone_inner()), - allocation, - } + ProviderConnectionGuard::new(Arc::new(self.clone_inner()), allocation) } // Returns the next available provider connection @@ -490,10 +494,7 @@ impl ActiveProviderManager { } } - ProviderConnectionGuard { - manager: Arc::new(self.clone_inner()), - allocation, - } + ProviderConnectionGuard::new(Arc::new(self.clone_inner()), allocation) } // This method is used for redirects to cycle through provider @@ -658,7 +659,7 @@ mod tests { // Create MultiProviderLineup with the provider and alias let lineup = MultiProviderLineup::new(&input); - let rt = tokio::runtime::Runtime::new().unwrap(); + let rt = tokio::runtime::Runtime::new().unwrap(); rt.block_on(async move { // Test that the alias provider is available should_available!(lineup, 1, 5); @@ -732,7 +733,7 @@ mod tests { input.aliases = Some(vec![alias1, alias2]); let lineup = MultiProviderLineup::new(&input); - let rt = tokio::runtime::Runtime::new().unwrap(); + let rt = tokio::runtime::Runtime::new().unwrap(); rt.block_on(async move { // Acquire connection from alias2 should_available!(lineup, 3, 5); diff --git a/src/api/model/active_user_manager.rs b/src/api/model/active_user_manager.rs index bd7a9f0af..7373d89f4 100644 --- a/src/api/model/active_user_manager.rs +++ b/src/api/model/active_user_manager.rs @@ -22,17 +22,21 @@ impl Drop for UserConnectionGuard { } struct UserConnectionData { + max_connections: u32, connections: u32, granted_grace: bool, grace_ts: u64, + token: Option, } impl UserConnectionData { - fn new() -> Self { + fn new(max_connections: u32) -> Self { Self { connections: 1, granted_grace: false, grace_ts: 0, + token: None, + max_connections, } } } @@ -130,12 +134,12 @@ impl ActiveUserManager { self.user.read().await.values().map(|c| c.connections as usize).sum() } - pub async fn add_connection(&self, username: &str) -> UserConnectionGuard { + pub async fn add_connection(&self, username: &str, max_connections: u32) -> UserConnectionGuard { let mut lock = self.user.write().await; if let Some(connection_data) = lock.get_mut(username) { connection_data.connections += 1; } else { - lock.insert(username.to_string(), UserConnectionData::new()); + lock.insert(username.to_string(), UserConnectionData::new(max_connections)); } drop(lock); @@ -159,6 +163,10 @@ impl ActiveUserManager { if connection_data.connections == 0 { lock.remove(username); + } else { + if connection_data.connections < connection_data.max_connections { + connection_data.token = None; + } } } drop(lock); @@ -166,7 +174,29 @@ impl ActiveUserManager { self.log_active_user().await; } - async fn log_active_user(&self) { + pub async fn create_token(&self, username: &str) -> String { + let result = crate::utils::string_utils::generate_random_string(6); + let mut lock = self.user.write().await; + if let Some(connection_data) = lock.get_mut(username) { + connection_data.token = Some(result.to_string()); + } + drop(lock); + result + } + + pub async fn get_token(&self, username: &str) -> Option { + let mut lock = self.user.write().await; + if let Some(connection_data) = lock.get_mut(username) { + connection_data.token.clone() + } else { + None + } + } + pub async fn has_token(&self, username: &str, token: &str) -> bool { + self.get_token(username).await.is_some_and(|t| token == t) + } + + async fn log_active_user(&self) { if self.log_active_user { let user_count = self.active_users().await; let user_connection_count = self.active_connections().await; diff --git a/src/api/model/streams/active_client_stream.rs b/src/api/model/streams/active_client_stream.rs index 970f3d60a..a283c5a1b 100644 --- a/src/api/model/streams/active_client_stream.rs +++ b/src/api/model/streams/active_client_stream.rs @@ -41,7 +41,7 @@ impl ActiveClientStream { } let grant_user_grace_period = connection_permission == UserConnectionPermission::GracePeriod; let username = user.username.as_str(); - let user_connection_guard = Some(active_user.add_connection(username).await); + let user_connection_guard = Some(active_user.add_connection(username, user.max_connections).await); let grace_stop_flag = Self::stream_grace_period(&stream_details, &active_provider, grant_user_grace_period, user, &active_user); Self { inner: stream_details.stream.take().unwrap(), diff --git a/src/processing/parser/hls.rs b/src/processing/parser/hls.rs index 2a474c458..d501d2f0e 100644 --- a/src/processing/parser/hls.rs +++ b/src/processing/parser/hls.rs @@ -11,6 +11,7 @@ pub struct RewriteHlsProps<'a> { pub virtual_id: u32, pub input_id: u16, pub provider_name: String, + pub user_token: String, } fn rewrite_hls_url(input: &str, replacement: &str) -> String { @@ -33,7 +34,7 @@ fn rewrite_uri_attrib(line: &str, props: &RewriteHlsProps) -> String { if let Some(caps) = CONSTANTS.re_hls_uri.captures(line) { let uri = &caps[1]; let target_url = &rewrite_hls_url(&props.hls_url, uri); - if let Some(token) = create_token_for_provider(props.secret, props.virtual_id, &props.provider_name, target_url) { + if let Some(token) = create_token_for_provider(props.secret, &props.user_token, props.virtual_id, &props.provider_name, target_url) { return CONSTANTS.re_hls_uri.replace(line, format!(r#"URI="{token}""#)).to_string(); } } @@ -58,7 +59,7 @@ pub fn rewrite_hls(user: &ProxyUserCredentials, props: &RewriteHlsProps) -> Stri } else { rewrite_hls_url(&props.hls_url, line) }; - if let Some(token) = create_token_for_provider(props.secret, props.virtual_id, &props.provider_name, &target_url) { + if let Some(token) = create_token_for_provider(props.secret, &props.user_token, props.virtual_id, &props.provider_name, &target_url) { let url = format!( "{}/{HLS_PREFIX}/{}/{}/{}/{}/{}", props.base_url, diff --git a/src/repository/bplustree.rs b/src/repository/bplustree.rs index f725b315f..469c43c06 100644 --- a/src/repository/bplustree.rs +++ b/src/repository/bplustree.rs @@ -780,26 +780,11 @@ mod tests { data: String, } - use crate::utils::time_utils::current_time_secs; - - fn generate_random_string(length: usize) -> String { - let mut rng = current_time_secs(); - let charset = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"; - - let mut random_string = String::with_capacity(length); - for _ in 0..length { - rng = rng.wrapping_add(1); - let idx = (rng % charset.len() as u64) as usize; - random_string.push(charset[idx] as char); - } - - random_string - } #[test] fn insert_test() -> io::Result<()> { let test_size = 500; - let content = generate_random_string(1024); + let content = crate::utils::string_utils::generate_random_string(1024); let mut tree = BPlusTree::::new(); for i in 0u32..=test_size { tree.insert(i, Record { diff --git a/src/utils/string_utils.rs b/src/utils/string_utils.rs index 76e964e76..8f4be50e2 100644 --- a/src/utils/string_utils.rs +++ b/src/utils/string_utils.rs @@ -1,3 +1,5 @@ +use rand::Rng; + // other implementations like calculating text_distance on all titles took too much time // we keep it now as simple as possible and less memory intensive. pub fn get_title_group(text: &str) -> String { @@ -41,3 +43,32 @@ pub fn get_trimmed_string(value: &Option) -> Option { } None } + +pub fn generate_random_string(length: usize) -> String { + let charset = b"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"; + let mut rng = rand::rng(); + + let random_string: String = (0..length) + .map(|_| { + let idx = rng.random_range(0..charset.len()); + charset[idx] as char + }) + .collect(); + + random_string +} + +#[cfg(test)] +mod test { + use std::collections::HashSet; + use crate::utils::string_utils::generate_random_string; + + #[test] + fn test_generate_random_string() { + let mut strings = HashSet::new(); + for _i in 0..100 { + strings.insert(generate_random_string(5)); + } + assert_eq!(strings.len(), 100); + } +} \ No newline at end of file