From 42ca2f56c72a0e7fa26731de25dc653281361564 Mon Sep 17 00:00:00 2001 From: zijiren233 Date: Sun, 22 Mar 2026 02:18:51 +0800 Subject: [PATCH] feat: more test --- config.example.yaml | 19 +- synctv-api/src/client_ip.rs | 5 +- synctv-api/src/grpc/blacklist_layer.rs | 68 +- synctv-api/src/grpc/interceptors.rs | 5 +- synctv-api/src/http/error.rs | 5 +- synctv-api/src/http/mod.rs | 32 +- synctv-api/src/http/oauth2.rs | 7 +- synctv-api/src/impls/mod.rs | 16 +- synctv-api/src/impls/oauth2.rs | 121 +- synctv-cluster/src/grpc/mod.rs | 5 +- synctv-cluster/src/grpc/server.rs | 2 +- synctv-cluster/src/sync/connection_manager.rs | 52 +- synctv-core/src/bootstrap/services.rs | 80 +- synctv-core/src/cache/l2_backend.rs | 134 +- synctv-core/src/lib.rs | 2 +- synctv-core/src/oauth2/providers/github.rs | 18 +- synctv-core/src/oauth2/providers/google.rs | 16 +- synctv-core/src/oauth2/providers/logto.rs | 16 +- synctv-core/src/oauth2/providers/mod.rs | 160 +++ synctv-core/src/oauth2/providers/oidc.rs | 39 +- synctv-core/src/provider/alist.rs | 5 + synctv-core/src/provider/bilibili.rs | 5 + synctv-core/src/provider/emby.rs | 5 + synctv-core/src/provider/provider_client.rs | 89 +- synctv-core/src/provider/traits.rs | 5 + .../src/repository/user_oauth_provider.rs | 10 +- synctv-core/src/resilience.rs | 34 +- .../src/service/audit_partition_manager.rs | 9 +- .../src/service/auth/security_pipeline.rs | 15 +- .../src/service/auth/token_blacklist.rs | 23 + .../src/service/chat_partition_manager.rs | 4 +- synctv-core/src/service/distributed_lock.rs | 111 +- .../service/notification_partition_manager.rs | 4 +- synctv-core/src/service/oauth2.rs | 242 +++- synctv-core/src/service/providers_manager.rs | 106 +- .../src/service/remote_provider_manager.rs | 360 +++--- synctv-core/src/service/user.rs | 108 +- synctv-core/src/service/ws_ticket.rs | 89 +- synctv-core/testing/src/postgres.rs | 23 +- synctv-core/testing/src/redis.rs | 15 +- synctv-core/tests/oauth2_state_store_tests.rs | 47 +- .../tests/remote_provider_manager_tests.rs | 1145 ++++++++++++++--- synctv-core/tests/user_auth_service_tests.rs | 148 ++- .../user_oauth_provider_repository_tests.rs | 141 +- .../src/relay/in_memory_registry.rs | 12 +- synctv-livestream/src/relay/mock_registry.rs | 20 +- .../src/relay/publisher_manager.rs | 130 +- synctv/src/migrations.rs | 11 +- 48 files changed, 2751 insertions(+), 967 deletions(-) diff --git a/config.example.yaml b/config.example.yaml index 4734a038..c5e59793 100644 --- a/config.example.yaml +++ b/config.example.yaml @@ -350,7 +350,7 @@ media_providers: # Common fields (all providers): # client_id: OAuth2 client ID # client_secret: OAuth2 client secret -# redirect_url: (optional, auto-generated if not provided) +# redirect_url: (required) frontend/client callback URI # # Provider-specific fields: # logto: @@ -365,11 +365,13 @@ media_providers: # Use different instance names (e.g., logto1, logto2) and add 'type' field # to specify the provider type. # -# Callback URLs are automatically generated based on instance name: -# http:///api/oauth2//callback -# Examples: -# http://localhost:8080/api/oauth2/github/callback -# http://localhost:8080/api/oauth2/logto1/callback +# Callback URLs must point to your frontend or native client, because SyncTV +# uses a frontend-driven OAuth2 flow: +# 1. Provider redirects to your frontend/client callback URI with code+state +# 2. Frontend/client calls POST /api/oauth2//exchange +# Examples: +# https://app.example.com/oauth2/callback +# synctv://oauth2/callback # # Security Notes: # - OAuth2 tokens are only used temporarily during login and then discarded @@ -386,10 +388,12 @@ oauth2: # github: # client_id: "your_github_client_id" # client_secret: "your_github_client_secret" + # redirect_url: "https://app.example.com/oauth2/callback" # # google: # client_id: "your_google_client_id" # client_secret: "your_google_client_secret" + # redirect_url: "https://app.example.com/oauth2/callback" # # # Multiple instances of same provider (requires 'type' field) # logto1: @@ -397,17 +401,20 @@ oauth2: # client_id: "logto1_client_id" # client_secret: "logto1_client_secret" # endpoint: "https://logto1.example.com" + # redirect_url: "https://app.example.com/oauth2/callback" # # logto2: # type: logto # client_id: "logto2_client_id" # client_secret: "logto2_client_secret" # endpoint: "https://logto2.example.com" + # redirect_url: "https://app.example.com/oauth2/callback" # # # Generic OIDC provider # custom_oidc: # type: oidc # client_id: "custom_client_id" + # redirect_url: "https://app.example.com/oauth2/callback" # client_secret: "custom_client_secret" # issuer: "https://custom.oidc.provider.com" providers: "" diff --git a/synctv-api/src/client_ip.rs b/synctv-api/src/client_ip.rs index 9dcf97ce..501887cb 100644 --- a/synctv-api/src/client_ip.rs +++ b/synctv-api/src/client_ip.rs @@ -46,7 +46,10 @@ mod tests { config.server.trusted_proxies = vec!["127.0.0.1".to_string()]; let mut headers = HeaderMap::new(); - headers.insert("x-forwarded-for", "203.0.113.50, 70.41.3.18".parse().unwrap()); + headers.insert( + "x-forwarded-for", + "203.0.113.50, 70.41.3.18".parse().unwrap(), + ); assert_eq!( extract_client_ip_from_headers( diff --git a/synctv-api/src/grpc/blacklist_layer.rs b/synctv-api/src/grpc/blacklist_layer.rs index ed74b827..558696a7 100644 --- a/synctv-api/src/grpc/blacklist_layer.rs +++ b/synctv-api/src/grpc/blacklist_layer.rs @@ -114,9 +114,13 @@ fn extract_bearer_token(headers: &http::HeaderMap) -> BearerTokenState { fn security_pipeline_error_status(err: &CoreError) -> tonic::Status { match SecurityPipeline::classify_auth_error(err) { - AuthErrorCategory::Authentication => tonic::Status::unauthenticated("Authentication failed"), + AuthErrorCategory::Authentication => { + tonic::Status::unauthenticated("Authentication failed") + } AuthErrorCategory::Authorization => tonic::Status::permission_denied("Permission denied"), - AuthErrorCategory::Unavailable => tonic::Status::unavailable("Authentication service unavailable"), + AuthErrorCategory::Unavailable => { + tonic::Status::unavailable("Authentication service unavailable") + } AuthErrorCategory::Internal => tonic::Status::internal("Internal error"), } } @@ -158,41 +162,42 @@ where tracing::warn!( "gRPC request rejected: malformed bearer authorization metadata" ); - let response = tonic::Status::unauthenticated("Invalid authorization header") - .into_http(); + let response = + tonic::Status::unauthenticated("Invalid authorization header").into_http(); return Ok(response); } BearerTokenState::Present(token) => { - // Security check order (matches HTTP AuthUser extractor): - // 1. JWT verification 2. Password invalidation 3. Banned/deleted user - - // Step 1: Verify JWT and extract claims - let claims = match jwt_service.verify_access_token(&token) { - Ok(claims) => claims, - Err(e) => { - tracing::warn!("gRPC request rejected: JWT validation failed: {e}"); - let response = - tonic::Status::unauthenticated("Invalid or expired token").into_http(); - return Ok(response); - } - }; - - // Steps 2-3: Shared security pipeline (password invalidation, user status) - let authenticated_token: AuthenticatedToken = - match security_pipeline.check(&claims).await { - Ok(authenticated_token) => authenticated_token, + // Security check order (matches HTTP AuthUser extractor): + // 1. JWT verification 2. Password invalidation 3. Banned/deleted user + + // Step 1: Verify JWT and extract claims + let claims = match jwt_service.verify_access_token(&token) { + Ok(claims) => claims, Err(e) => { - tracing::warn!("gRPC request rejected by security pipeline: {e}"); - let response = security_pipeline_error_status(&e).into_http(); + tracing::warn!("gRPC request rejected: JWT validation failed: {e}"); + let response = + tonic::Status::unauthenticated("Invalid or expired token") + .into_http(); return Ok(response); } }; - // Preserve the authenticated identity so downstream gRPC - // interceptors and handlers can reuse it without re-running - // JWT verification or the security pipeline. - req.extensions_mut().insert(authenticated_token); - } + // Steps 2-3: Shared security pipeline (password invalidation, user status) + let authenticated_token: AuthenticatedToken = + match security_pipeline.check(&claims).await { + Ok(authenticated_token) => authenticated_token, + Err(e) => { + tracing::warn!("gRPC request rejected by security pipeline: {e}"); + let response = security_pipeline_error_status(&e).into_http(); + return Ok(response); + } + }; + + // Preserve the authenticated identity so downstream gRPC + // interceptors and handlers can reuse it without re-running + // JWT verification or the security pipeline. + req.extensions_mut().insert(authenticated_token); + } } // Inject SecurityCheckPassed marker into request extensions. @@ -330,7 +335,10 @@ mod tests { // No auth header (public endpoint -- both layers should pass through) let empty_headers = http::HeaderMap::new(); - assert_eq!(extract_bearer_token(&empty_headers), BearerTokenState::Missing); + assert_eq!( + extract_bearer_token(&empty_headers), + BearerTokenState::Missing + ); } // ========== SecurityCheckPassed Marker Tests ========== diff --git a/synctv-api/src/grpc/interceptors.rs b/synctv-api/src/grpc/interceptors.rs index b4c03bda..0bb5bd55 100644 --- a/synctv-api/src/grpc/interceptors.rs +++ b/synctv-api/src/grpc/interceptors.rs @@ -538,7 +538,10 @@ mod tests { let interceptor = AuthInterceptor::new(jwt_service); let result = interceptor.inject_user(request); - assert!(result.is_err(), "Non-bearer authorization should not look authenticated"); + assert!( + result.is_err(), + "Non-bearer authorization should not look authenticated" + ); let err = result.unwrap_err(); assert_eq!(err.code(), tonic::Code::Unauthenticated); assert_eq!(err.message(), "Missing authorization header"); diff --git a/synctv-api/src/http/error.rs b/synctv-api/src/http/error.rs index b4f63d0f..44e73a5b 100644 --- a/synctv-api/src/http/error.rs +++ b/synctv-api/src/http/error.rs @@ -703,8 +703,9 @@ mod tests { #[test] fn test_from_core_auth_service_unavailable_stays_service_unavailable() { - let core_err = - synctv_core::Error::ServiceUnavailable("Authentication service unavailable".to_string()); + let core_err = synctv_core::Error::ServiceUnavailable( + "Authentication service unavailable".to_string(), + ); let app_err = AppError::from(core_err); assert_eq!(app_err.status, StatusCode::SERVICE_UNAVAILABLE); assert!( diff --git a/synctv-api/src/http/mod.rs b/synctv-api/src/http/mod.rs index ddf2aedd..3935dc7a 100644 --- a/synctv-api/src/http/mod.rs +++ b/synctv-api/src/http/mod.rs @@ -48,7 +48,8 @@ use tower_http::trace::TraceLayer; pub use error::{AppError, AppResult}; -const HTTP_REQUEST_TIMEOUT: std::time::Duration = synctv_core::resilience::timeout::HTTP_REQUEST_TIMEOUT; +const HTTP_REQUEST_TIMEOUT: std::time::Duration = + synctv_core::resilience::timeout::HTTP_REQUEST_TIMEOUT; /// Configuration for creating the HTTP router #[derive(Clone)] @@ -631,7 +632,9 @@ fn register_websocket_routes(state: &AppState) -> Router { #[cfg(test)] fn register_all_routes_for_test(state: AppState) -> Router { let (timeout_router, no_timeout_router, upgrade_router) = register_all_routes(state); - timeout_router.merge(no_timeout_router).merge(upgrade_router) + timeout_router + .merge(no_timeout_router) + .merge(upgrade_router) } fn register_all_routes(state: AppState) -> (Router, Router, Router) { @@ -683,7 +686,7 @@ fn register_all_routes(state: AppState) -> (Router, Router, state.clone(), middleware::read_rate_limit, )), - ) + ), ); timeout_router = timeout_router.merge(register_provider_proxy_routes(&state)); @@ -957,11 +960,10 @@ fn apply_global_layers_with_timeout( move |request: axum::extract::Request, next: axum::middleware::Next| async move { match tokio::time::timeout(request_timeout, next.run(request)).await { Ok(response) => response, - Err(_) => AppError::new( - axum::http::StatusCode::REQUEST_TIMEOUT, - "Request timed out", - ) - .into_response(), + Err(_) => { + AppError::new(axum::http::StatusCode::REQUEST_TIMEOUT, "Request timed out") + .into_response() + } } }, )), @@ -994,8 +996,8 @@ fn apply_global_layers_with_timeout( mod tests { use super::{ apply_global_layers_with_timeout, build_app_state, build_cors_layer, - register_all_routes_for_test, - start_proxy_cache_lifecycle, RouterConfig, HTTP_REQUEST_TIMEOUT, + register_all_routes_for_test, start_proxy_cache_lifecycle, RouterConfig, + HTTP_REQUEST_TIMEOUT, }; use axum::body::Body; use axum::http::{Request, StatusCode}; @@ -1609,8 +1611,9 @@ mod tests { "DENY", "timed out responses must still include security headers" ); - let timeout_body = - axum::body::to_bytes(timeout_response.into_body(), usize::MAX).await.unwrap(); + let timeout_body = axum::body::to_bytes(timeout_response.into_body(), usize::MAX) + .await + .unwrap(); let timeout_json: serde_json::Value = serde_json::from_slice(&timeout_body).expect("timeout response should be JSON"); assert_eq!(timeout_json["status"], 408); @@ -1751,10 +1754,7 @@ mod tests { .method("OPTIONS") .uri("/test") .header(axum::http::header::ORIGIN, "https://example.com") - .header( - axum::http::header::ACCESS_CONTROL_REQUEST_METHOD, - "GET", - ) + .header(axum::http::header::ACCESS_CONTROL_REQUEST_METHOD, "GET") .body(Body::empty()) .expect("request"), ) diff --git a/synctv-api/src/http/oauth2.rs b/synctv-api/src/http/oauth2.rs index f1b8f923..6f3cdff0 100644 --- a/synctv-api/src/http/oauth2.rs +++ b/synctv-api/src/http/oauth2.rs @@ -117,8 +117,11 @@ pub async fn exchange_authorization_code( let current_user_id = maybe_auth.as_ref().map(|a| &a.user_id); // Extract client IP for brute-force protection (Issue #24). - let client_ip = - crate::client_ip::extract_client_ip_from_headers(&state.config, connect_info.0.ip(), &headers); + let client_ip = crate::client_ip::extract_client_ip_from_headers( + &state.config, + connect_info.0.ip(), + &headers, + ); let result = oauth2_api .exchange_authorization_code( diff --git a/synctv-api/src/impls/mod.rs b/synctv-api/src/impls/mod.rs index d56baf11..a82d8dc2 100644 --- a/synctv-api/src/impls/mod.rs +++ b/synctv-api/src/impls/mod.rs @@ -226,7 +226,9 @@ impl From for ApiError { synctv_core::Error::AlreadyExists(msg) => Self::AlreadyExists(msg), synctv_core::Error::InvalidInput(msg) => Self::InvalidInput(msg), synctv_core::Error::RateLimited(msg) => Self::RateLimited(msg), - synctv_core::Error::ServiceUnavailable(msg) => Self::ServiceUnavailable(msg), + synctv_core::Error::ServiceUnavailable(msg) | synctv_core::Error::Timeout(msg) => { + Self::ServiceUnavailable(msg) + } other => Self::Internal(other.to_string()), } } @@ -874,6 +876,18 @@ mod tests { assert_eq!(api_err.code(), error_codes::SERVICE_UNAVAILABLE); } + #[test] + fn test_api_error_from_core_timeout_maps_to_service_unavailable() { + let core_err = synctv_core::Error::Timeout("oauth2 provider timed out".to_string()); + let api_err = ApiError::from(core_err); + assert!(matches!( + api_err, + ApiError::ServiceUnavailable(ref msg) if msg == "oauth2 provider timed out" + )); + assert!(matches!(api_err.classify(), ErrorKind::ServiceUnavailable)); + assert_eq!(api_err.code(), error_codes::SERVICE_UNAVAILABLE); + } + #[test] fn test_classify_by_prefix_rate_limited() { assert!(matches!( diff --git a/synctv-api/src/impls/oauth2.rs b/synctv-api/src/impls/oauth2.rs index 59d360d1..229eaf83 100644 --- a/synctv-api/src/impls/oauth2.rs +++ b/synctv-api/src/impls/oauth2.rs @@ -20,7 +20,7 @@ //! 8. Backend returns JWT token to frontend use std::sync::Arc; -use synctv_core::models::{SignupMethod, User, UserId, UserRole, UserStatus}; +use synctv_core::models::{User, UserId, UserRole, UserStatus}; use synctv_core::service::{OAuth2Service, UserService}; use synctv_proto::client::{LinkedProvider, OAuth2ProviderInstance, OAuth2UserInfo}; @@ -204,123 +204,16 @@ impl OAuth2ApiImpl { .await .map_err(ApiError::from)? } else { - // User doesn't exist, create new account and link OAuth2 provider - // atomically within a single transaction to prevent orphaned users. - // - // Generate a random password for the OAuth2 user. This password is never - // used for login (OAuth2 users authenticate via their provider), but it's - // required by the user model. The password is hashed with Argon2, which - // is a CPU-intensive operation. - let random_password = nanoid::nanoid!(32); - tracing::debug!( - username = %user_info.username, - provider = %provider_type.as_str(), - "Creating new OAuth2 user with random password" - ); - - let pool = self.user_service.pool(); - - let mut tx = pool.begin().await.map_err(|e| { - tracing::error!(error = %e, "Failed to begin transaction for OAuth2 user creation"); - ApiError::Internal(format!("Failed to begin transaction: {e}")) - })?; - - // Try registering with the provider username, retrying with suffixed - // usernames if there's a collision (AlreadyExists). Try up to 4 times: - // original, then with _, _1, _2 suffixes. - let suffixes = [ - String::new(), - format!("_{}", provider_type.as_str()), - "_1".to_string(), - "_2".to_string(), - ]; - let mut new_user = None; - let mut last_err = None; - for suffix in &suffixes { - let candidate = format!("{}{}", user_info.username, suffix); - match self - .user_service - .register_with_executor( - candidate.clone(), - user_info.email.clone(), - random_password.clone(), - SignupMethod::OAuth2, - &mut *tx, - ) - .await - { - Ok(user) => { - new_user = Some(user); - break; - } - Err(synctv_core::Error::AlreadyExists(_)) => { - tracing::debug!( - username = %candidate, - provider = %provider_type.as_str(), - "OAuth2 username collision, trying next suffix" - ); - last_err = Some(ApiError::AlreadyExists( - "Username already taken".to_string(), - )); - continue; - } - Err(e) => { - tracing::error!( - error = %e, - username = %candidate, - provider = %provider_type.as_str(), - "Failed to create OAuth2 user" - ); - return Err(ApiError::from(e)); - } - } - } - let new_user = match new_user { - Some(u) => u, - None => { - return Err(last_err.unwrap_or_else(|| { - ApiError::Internal( - "Failed to create OAuth2 user after all retries".to_string(), - ) - })) - } - }; - - self.oauth2_service - .upsert_user_provider_with_executor( - &new_user.id, - &provider_type, - &user_info.provider_user_id, - &user_info, - &mut *tx, - ) + let (user_id, _is_new) = self + .oauth2_service + .find_or_create_and_link(&self.user_service, &provider_type, &user_info) .await .map_err(ApiError::from)?; - // Set email_verified inside the transaction if the OAuth2 provider confirmed the email - if user_info.email_verified && user_info.email.is_some() { - sqlx::query( - "UPDATE users SET email_verified = true, updated_at = NOW() WHERE id = $1", - ) - .bind(new_user.id.as_str()) - .execute(&mut *tx) - .await - .map_err(|e| { - ApiError::Internal(format!("Failed to set email_verified in transaction: {e}")) - })?; - } - - tx.commit() - .await - .map_err(|e| ApiError::Internal(format!("Failed to commit transaction: {e}")))?; - - let (access_token, refresh_token) = self - .user_service - .finalize_registration(&new_user) + self.user_service + .login_oauth2(&user_id, &user_info.provider_user_id, client_ip) .await - .map_err(ApiError::from)?; - - (new_user, access_token, refresh_token) + .map_err(ApiError::from)? }; // Get the actual access token duration from the JWT service diff --git a/synctv-cluster/src/grpc/mod.rs b/synctv-cluster/src/grpc/mod.rs index 52fc5a0b..70ee0f5a 100644 --- a/synctv-cluster/src/grpc/mod.rs +++ b/synctv-cluster/src/grpc/mod.rs @@ -52,10 +52,7 @@ impl ClusterAuthInterceptor { /// Validate the shared secret directly from request metadata. #[allow(clippy::result_large_err)] - pub fn validate_metadata( - &self, - metadata: &tonic::metadata::MetadataMap, - ) -> Result<(), Status> { + pub fn validate_metadata(&self, metadata: &tonic::metadata::MetadataMap) -> Result<(), Status> { if self.secret.is_empty() { tracing::error!( "Cluster gRPC auth interceptor is misconfigured: shared secret is empty" diff --git a/synctv-cluster/src/grpc/server.rs b/synctv-cluster/src/grpc/server.rs index 588488b5..12172354 100644 --- a/synctv-cluster/src/grpc/server.rs +++ b/synctv-cluster/src/grpc/server.rs @@ -5,13 +5,13 @@ use std::sync::Arc; use tonic::{Request, Response, Status}; -use super::ClusterAuthInterceptor; use super::synctv::cluster::cluster_service_server::ClusterService; use super::synctv::cluster::{ DeregisterNodeRequest, DeregisterNodeResponse, GetNodesRequest, GetNodesResponse, GetRoomConnectionsRequest, GetRoomConnectionsResponse, GetUserOnlineStatusRequest, GetUserOnlineStatusResponse, NodeInfo, NodeStatus, RoomConnection, UserOnlineStatus, }; +use super::ClusterAuthInterceptor; use crate::discovery::{NodeInfo as DiscoveryNodeInfo, NodeRegistry}; use crate::sync::connection_manager::ConnectionManager; diff --git a/synctv-cluster/src/sync/connection_manager.rs b/synctv-cluster/src/sync/connection_manager.rs index 4c79c67f..29b7b650 100644 --- a/synctv-cluster/src/sync/connection_manager.rs +++ b/synctv-cluster/src/sync/connection_manager.rs @@ -11,9 +11,8 @@ use tokio::sync::{broadcast, mpsc}; use tracing::{debug, info, warn}; #[cfg(test)] -type AsyncTestHook = Arc< - dyn Fn() -> std::pin::Pin + Send>> + Send + Sync, ->; +type AsyncTestHook = + Arc std::pin::Pin + Send>> + Send + Sync>; /// Disconnect signal for forcing connections to close #[derive(Debug, Clone)] @@ -459,10 +458,7 @@ impl ConnectionManager { self } - fn connection_lifecycle_lock( - &self, - connection_id: &str, - ) -> Arc> { + fn connection_lifecycle_lock(&self, connection_id: &str) -> Arc> { let mut hasher = std::collections::hash_map::DefaultHasher::new(); connection_id.hash(&mut hasher); let index = (hasher.finish() as usize) % self.connection_lifecycle_locks.len(); @@ -748,7 +744,11 @@ impl ConnectionManager { }; let conn_key = format!("{}conn_mgr:conn:{}", self.redis_key_prefix, connection_id); - let user_index_key = format!("{}conn_mgr:user:{}", self.redis_key_prefix, user_id.as_str()); + let user_index_key = format!( + "{}conn_mgr:user:{}", + self.redis_key_prefix, + user_id.as_str() + ); let persistent = ConnectionInfoPersistent::from(&conn_info); match serde_json::to_string(&persistent) { @@ -801,7 +801,11 @@ impl ConnectionManager { transition.room_id.as_str() ); let previous_room_index_key = transition.previous_room_id.as_ref().map(|room_id| { - format!("{}conn_mgr:room:{}", self.redis_key_prefix, room_id.as_str()) + format!( + "{}conn_mgr:room:{}", + self.redis_key_prefix, + room_id.as_str() + ) }); let persistent = ConnectionInfoPersistent::from(&conn_info); @@ -1018,17 +1022,15 @@ impl ConnectionManager { .take(); if let Some(handle) = disconnect_retry_handle { report.disconnect_retry = Some( - Self::await_shutdown_task( - "disconnect retry", - Duration::from_secs(5), - handle, - ) - .await, + Self::await_shutdown_task("disconnect retry", Duration::from_secs(5), handle).await, ); } if !report.all_clean() { - warn!(?report, "ConnectionManager shutdown observed background task failures"); + warn!( + ?report, + "ConnectionManager shutdown observed background task failures" + ); } report @@ -1073,11 +1075,17 @@ impl ConnectionManager { ) -> ShutdownTaskOutcome { match tokio::time::timeout(timeout_budget, handle).await { Ok(Ok(())) => { - debug!(task = task_name, "ConnectionManager background task stopped"); + debug!( + task = task_name, + "ConnectionManager background task stopped" + ); ShutdownTaskOutcome::Completed } Ok(Err(error)) if error.is_cancelled() => { - debug!(task = task_name, "ConnectionManager background task cancelled"); + debug!( + task = task_name, + "ConnectionManager background task cancelled" + ); ShutdownTaskOutcome::Cancelled } Ok(Err(error)) => { @@ -3469,8 +3477,8 @@ mod tests { async fn test_duplicate_register_fails_fast_while_first_attempt_holds_lifecycle_lock() { let first_entered = Arc::new(tokio::sync::Notify::new()); let release_first = Arc::new(tokio::sync::Notify::new()); - let manager = Arc::new(ConnectionManager::default().with_register_after_lifecycle_lock_hook( - { + let manager = Arc::new( + ConnectionManager::default().with_register_after_lifecycle_lock_hook({ let first_entered = Arc::clone(&first_entered); let release_first = Arc::clone(&release_first); Arc::new(move || { @@ -3481,8 +3489,8 @@ mod tests { release_first.notified().await; }) }) - }, - )); + }), + ); let user_id = UserId::from_string("dup-fast-user".to_string()); let first = { diff --git a/synctv-core/src/bootstrap/services.rs b/synctv-core/src/bootstrap/services.rs index cf5c7b57..bd31952f 100644 --- a/synctv-core/src/bootstrap/services.rs +++ b/synctv-core/src/bootstrap/services.rs @@ -664,7 +664,10 @@ async fn init_oauth2_service( let state_store: Arc = match (cluster_mode, redis_conn) { (true, Some(conn)) => { info!("OAuth2 state store: Redis (cluster mode)"); - Arc::new(crate::service::RedisOAuthStateStore::new(conn)) + Arc::new(crate::service::RedisOAuthStateStore::new( + conn, + config.redis.key_prefix.clone(), + )) } (true, None) => { return Err(anyhow::anyhow!( @@ -674,7 +677,10 @@ async fn init_oauth2_service( } (false, Some(conn)) => { info!("OAuth2 state store: Redis"); - Arc::new(crate::service::RedisOAuthStateStore::new(conn)) + Arc::new(crate::service::RedisOAuthStateStore::new( + conn, + config.redis.key_prefix.clone(), + )) } (false, None) => { info!("OAuth2 state store: in-memory (standalone mode)"); @@ -704,23 +710,18 @@ async fn init_oauth2_service( .unwrap_or(&instance_name) .to_string(); - // Create a mutable config for adding redirect_url - let mut full_config = full_config.clone(); - - // Add redirect_url to config (merge it in) - // Use configured scheme (http/https) to support reverse proxy TLS termination - let scheme = &config.oauth2.redirect_scheme; - let redirect_url = format!( - "{}://{}/api/oauth2/{}/callback", - scheme, - config.advertise_host(), - instance_name - ); - if let Some(mapping) = full_config.as_object_mut() { - mapping.insert( - "redirect_url".to_string(), - serde_json::Value::String(redirect_url.clone()), - ); + let full_config = full_config.clone(); + if full_config + .get("redirect_url") + .and_then(serde_json::Value::as_str) + .is_none_or(str::is_empty) + { + return Err(anyhow::anyhow!( + "OAuth2 provider '{}' is missing redirect_url. \ + Frontend-driven OAuth2 requires an explicit frontend/client callback URI \ + (for example https://app.example.com/oauth2/callback or myapp://oauth2/callback).", + instance_name + )); } // Use factory to create provider with full config @@ -1435,6 +1436,47 @@ mod tests { ); } + #[tokio::test] + async fn test_init_oauth2_service_requires_explicit_redirect_url_for_frontend_flow() { + let pool = PgPool::connect_lazy("postgresql://test").expect("lazy pool should build"); + let mut config = Config::default(); + config.oauth2.providers = serde_json::json!({ + "github": { + "type": "github", + "client_id": "test-client-id", + "client_secret": "test-client-secret" + } + }); + + let error = init_oauth2_service(pool, &config, None, false) + .await + .expect_err("frontend-driven OAuth2 must require an explicit redirect_url"); + + assert!( + error.to_string().contains("redirect_url"), + "unexpected error: {error}" + ); + } + + #[tokio::test] + async fn test_init_oauth2_service_preserves_explicit_redirect_url() { + let pool = PgPool::connect_lazy("postgresql://test").expect("lazy pool should build"); + let mut config = Config::default(); + let redirect_url = "synctv://oauth2/callback".to_string(); + config.oauth2.providers = serde_json::json!({ + "github": { + "type": "github", + "client_id": "test-client-id", + "client_secret": "test-client-secret", + "redirect_url": redirect_url + } + }); + + init_oauth2_service(pool, &config, None, false) + .await + .expect("explicit frontend/client redirect_url should be accepted"); + } + #[test] fn test_init_services_uses_configured_messaging_rate_limits() { let mut config = Config::default(); diff --git a/synctv-core/src/cache/l2_backend.rs b/synctv-core/src/cache/l2_backend.rs index 22822754..67168e57 100644 --- a/synctv-core/src/cache/l2_backend.rs +++ b/synctv-core/src/cache/l2_backend.rs @@ -10,6 +10,39 @@ use crate::resilience::timeout::REDIS_OPERATION_TIMEOUT; use crate::{Error, Result}; use async_trait::async_trait; +use std::future::Future; + +enum L2RedisAttemptError { + Redis(redis::RedisError), + Timeout, +} + +async fn run_l2_redis_attempt(future: F) -> std::result::Result +where + F: Future>, +{ + match tokio::time::timeout(REDIS_OPERATION_TIMEOUT, future).await { + Ok(Ok(value)) => Ok(value), + Ok(Err(err)) => Err(L2RedisAttemptError::Redis(err)), + Err(_) => Err(L2RedisAttemptError::Timeout), + } +} + +async fn run_l2_redis_op(operation: impl Into, future: F) -> Result +where + F: Future>, +{ + let operation = operation.into(); + match run_l2_redis_attempt(future).await { + Ok(value) => Ok(value), + Err(L2RedisAttemptError::Timeout) => { + Err(Error::Timeout(format!("L2 cache timeout: {operation}"))) + } + Err(L2RedisAttemptError::Redis(err)) => { + Err(Error::Internal(format!("Failed to {operation}: {err}"))) + } + } +} /// Backend for the L2 (remote) cache layer in `TieredCache`. /// @@ -108,10 +141,7 @@ impl CacheL2Backend for RedisCacheL2 { let mut conn = self.conn().await; let result = - tokio::time::timeout(REDIS_OPERATION_TIMEOUT, conn.get::<_, Option>(key)) - .await - .map_err(|_| Error::Internal("L2 cache get operation timed out".to_string()))? - .map_err(|e| Error::Internal(format!("Failed to get from L2 cache: {e}")))?; + run_l2_redis_op("get from L2 cache", conn.get::<_, Option>(key)).await?; Ok(result) } @@ -119,13 +149,11 @@ impl CacheL2Backend for RedisCacheL2 { use redis::AsyncCommands; let mut conn = self.conn().await; - tokio::time::timeout( - REDIS_OPERATION_TIMEOUT, + run_l2_redis_op( + "set in L2 cache", conn.set_ex::<_, _, ()>(key, json, ttl_secs), ) - .await - .map_err(|_| Error::Internal("L2 cache set operation timed out".to_string()))? - .map_err(|e| Error::Internal(format!("Failed to set in L2 cache: {e}")))?; + .await?; Ok(()) } @@ -133,10 +161,7 @@ impl CacheL2Backend for RedisCacheL2 { use redis::AsyncCommands; let mut conn = self.conn().await; - tokio::time::timeout(REDIS_OPERATION_TIMEOUT, conn.del::<_, ()>(key)) - .await - .map_err(|_| Error::Internal("L2 cache delete operation timed out".to_string()))? - .map_err(|e| Error::Internal(format!("Failed to delete from L2 cache: {e}")))?; + run_l2_redis_op("delete from L2 cache", conn.del::<_, ()>(key)).await?; Ok(()) } @@ -144,11 +169,9 @@ impl CacheL2Backend for RedisCacheL2 { use redis::AsyncCommands; for attempt in 0..max_retries { let mut conn = self.conn().await; - let result = - tokio::time::timeout(REDIS_OPERATION_TIMEOUT, conn.del::<_, ()>(key)).await; - match result { - Ok(Ok(())) => return Ok(()), - Ok(Err(e)) => { + match run_l2_redis_attempt(conn.del::<_, ()>(key)).await { + Ok(()) => return Ok(()), + Err(L2RedisAttemptError::Redis(e)) => { let is_last_attempt = attempt == max_retries - 1; if is_last_attempt { crate::metrics::cache::CACHE_ERRORS @@ -178,7 +201,7 @@ impl CacheL2Backend for RedisCacheL2 { tokio::time::sleep(tokio::time::Duration::from_millis(backoff_ms)).await; } } - Err(_) => { + Err(L2RedisAttemptError::Timeout) => { let is_last_attempt = attempt == max_retries - 1; if is_last_attempt { crate::metrics::cache::CACHE_ERRORS @@ -190,8 +213,8 @@ impl CacheL2Backend for RedisCacheL2 { cache_type = %cache_type, "Redis L2 cache delete timed out after retries" ); - return Err(Error::Internal( - "Failed to delete from Redis cache: operation timed out".to_string(), + return Err(Error::Timeout( + "L2 cache timeout: delete from Redis cache".to_string(), )); } else { let backoff_ms = 10 * u64::pow(5, attempt); @@ -219,10 +242,7 @@ impl CacheL2Backend for RedisCacheL2 { } let results: Vec> = - tokio::time::timeout(REDIS_OPERATION_TIMEOUT, pipe.query_async(&mut conn)) - .await - .map_err(|_| Error::Internal("L2 cache batch get operation timed out".to_string()))? - .map_err(|e| Error::Internal(format!("Failed to batch get from L2: {e}")))?; + run_l2_redis_op("batch get from L2 cache", pipe.query_async(&mut conn)).await?; Ok(results) } @@ -255,8 +275,8 @@ impl CacheL2Backend for RedisCacheL2 { ", ); - let result: i64 = tokio::time::timeout( - REDIS_OPERATION_TIMEOUT, + let result: i64 = run_l2_redis_op( + "run set_if_newer Lua script", script .key(key) .arg(json) @@ -264,9 +284,7 @@ impl CacheL2Backend for RedisCacheL2 { .arg(new_ts_iso) .invoke_async(&mut conn), ) - .await - .map_err(|_| Error::Internal("L2 cache set_if_newer operation timed out".to_string()))? - .map_err(|e| Error::Internal(format!("Failed to run set_if_newer Lua script: {e}")))?; + .await?; Ok(result == 1) } @@ -280,8 +298,8 @@ impl CacheL2Backend for RedisCacheL2 { let pattern = format!("{prefix}*"); let mut cursor: u64 = 0; loop { - let scan_result: (u64, Vec) = tokio::time::timeout( - REDIS_OPERATION_TIMEOUT, + let scan_result: (u64, Vec) = run_l2_redis_op( + format!("scan L2 cache keys for prefix '{prefix}'"), redis::cmd("SCAN") .arg(cursor) .arg("MATCH") @@ -290,19 +308,16 @@ impl CacheL2Backend for RedisCacheL2 { .arg(100u64) .query_async(&mut conn), ) - .await - .map_err(|_| Error::Internal(format!("SCAN timed out for prefix '{prefix}'")))? - .map_err(|e| Error::Internal(format!("SCAN failed for prefix '{prefix}': {e}")))?; + .await?; let (next_cursor, keys) = scan_result; if !keys.is_empty() { - tokio::time::timeout(REDIS_OPERATION_TIMEOUT, conn.del::<_, ()>(keys.as_slice())) - .await - .map_err(|_| Error::Internal(format!("DEL timed out for prefix '{prefix}'")))? - .map_err(|e| { - Error::Internal(format!("DEL failed for prefix '{prefix}': {e}")) - })?; + run_l2_redis_op( + format!("delete L2 cache keys for prefix '{prefix}'"), + conn.del::<_, ()>(keys.as_slice()), + ) + .await?; } cursor = next_cursor; @@ -474,6 +489,43 @@ mod tests { /// Short timeout for tests — just needs to prove the timeout mechanism works. const TEST_TIMEOUT: Duration = Duration::from_millis(200); + #[tokio::test(start_paused = true)] + async fn test_l2_redis_timeout_maps_to_timeout_error() { + let timeout_future = run_l2_redis_op("get from L2 cache", async { + std::future::pending::<()>().await; + #[allow(unreachable_code)] + Ok::<(), redis::RedisError>(()) + }); + + tokio::pin!(timeout_future); + tokio::task::yield_now().await; + tokio::time::advance(REDIS_OPERATION_TIMEOUT).await; + + let err = timeout_future.await.expect_err("operation should time out"); + assert!(matches!( + err, + Error::Timeout(ref msg) if msg == "L2 cache timeout: get from L2 cache" + )); + } + + #[tokio::test(start_paused = true)] + async fn test_l2_redis_retry_attempt_reports_timeout() { + let timeout_future = run_l2_redis_attempt(async { + std::future::pending::<()>().await; + #[allow(unreachable_code)] + Ok::<(), redis::RedisError>(()) + }); + + tokio::pin!(timeout_future); + tokio::task::yield_now().await; + tokio::time::advance(REDIS_OPERATION_TIMEOUT).await; + + let err = timeout_future + .await + .expect_err("retryable redis operation should time out"); + assert!(matches!(err, L2RedisAttemptError::Timeout)); + } + /// Test that get() times out when operation takes too long #[tokio::test] async fn test_get_timeout() { diff --git a/synctv-core/src/lib.rs b/synctv-core/src/lib.rs index fe900b72..5e150e8a 100644 --- a/synctv-core/src/lib.rs +++ b/synctv-core/src/lib.rs @@ -1,7 +1,7 @@ #![cfg_attr(test, allow(clippy::unwrap_used))] -pub mod bootstrap; pub mod bench_support; +pub mod bootstrap; pub mod cache; pub mod config; pub mod error; diff --git a/synctv-core/src/oauth2/providers/github.rs b/synctv-core/src/oauth2/providers/github.rs index f7115498..1b5e4d14 100644 --- a/synctv-core/src/oauth2/providers/github.rs +++ b/synctv-core/src/oauth2/providers/github.rs @@ -1,5 +1,6 @@ //! GitHub `OAuth2` provider +use super::{build_oauth2_http_client, build_provider_http_client, map_provider_http_error}; use crate::oauth2::{OAuth2UserInfo, Provider}; use crate::{Error, InternalExt}; use async_trait::async_trait; @@ -23,6 +24,7 @@ pub struct GitHubConfig { pub struct GitHubProvider { client: Arc>, + oauth2_http_client: Arc, http_client: Arc, } @@ -55,12 +57,8 @@ impl GitHubProvider { Ok(Self { client, - http_client: Arc::new( - Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .internal_with_err("Failed to build HTTP client")?, - ), + oauth2_http_client: build_oauth2_http_client()?, + http_client: build_provider_http_client()?, }) } } @@ -89,7 +87,7 @@ impl GitHubProvider { .header("User-Agent", "synctv-rs") .send() .await - .internal_with_err("Failed to fetch user emails")? + .map_err(|err| map_provider_http_error("Failed to fetch user emails", err))? .error_for_status() .internal_with_err("GitHub emails API error")?; @@ -140,9 +138,9 @@ impl Provider for GitHubProvider { .client .exchange_code(oauth2::AuthorizationCode::new(code.to_string())) .set_pkce_verifier(verifier) - .request_async(&oauth2::reqwest::Client::new()) + .request_async(self.oauth2_http_client.as_ref()) .await - .internal_with_err("Failed to exchange code")?; + .map_err(|err| map_provider_http_error("Failed to exchange code", err))?; // Fetch user info let resp = self @@ -155,7 +153,7 @@ impl Provider for GitHubProvider { .header("User-Agent", "synctv-rs") .send() .await - .internal_with_err("Failed to fetch user info")? + .map_err(|err| map_provider_http_error("Failed to fetch user info", err))? .error_for_status() .internal_with_err("GitHub API error")?; diff --git a/synctv-core/src/oauth2/providers/google.rs b/synctv-core/src/oauth2/providers/google.rs index 311dd765..622bf884 100644 --- a/synctv-core/src/oauth2/providers/google.rs +++ b/synctv-core/src/oauth2/providers/google.rs @@ -1,5 +1,6 @@ //! Google `OAuth2` provider +use super::{build_oauth2_http_client, build_provider_http_client, map_provider_http_error}; use crate::oauth2::{OAuth2UserInfo, Provider}; use crate::{Error, InternalExt}; use async_trait::async_trait; @@ -23,6 +24,7 @@ pub struct GoogleConfig { pub struct GoogleProvider { client: Arc>, + oauth2_http_client: Arc, http_client: Arc, } @@ -54,12 +56,8 @@ impl GoogleProvider { Ok(Self { client, - http_client: Arc::new( - Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .internal_with_err("Failed to build HTTP client")?, - ), + oauth2_http_client: build_oauth2_http_client()?, + http_client: build_provider_http_client()?, }) } } @@ -91,9 +89,9 @@ impl Provider for GoogleProvider { .client .exchange_code(oauth2::AuthorizationCode::new(code.to_string())) .set_pkce_verifier(verifier) - .request_async(&oauth2::reqwest::Client::new()) + .request_async(self.oauth2_http_client.as_ref()) .await - .internal_with_err("Failed to exchange code")?; + .map_err(|err| map_provider_http_error("Failed to exchange code", err))?; // Fetch user info let resp = self @@ -105,7 +103,7 @@ impl Provider for GoogleProvider { ) .send() .await - .internal_with_err("Failed to fetch user info")? + .map_err(|err| map_provider_http_error("Failed to fetch user info", err))? .error_for_status() .internal_with_err("Google API error")?; diff --git a/synctv-core/src/oauth2/providers/logto.rs b/synctv-core/src/oauth2/providers/logto.rs index f4a5cfcf..0f4c2e0d 100644 --- a/synctv-core/src/oauth2/providers/logto.rs +++ b/synctv-core/src/oauth2/providers/logto.rs @@ -1,5 +1,6 @@ //! Logto `OAuth2` provider +use super::{build_oauth2_http_client, build_provider_http_client, map_provider_http_error}; use crate::oauth2::{OAuth2UserInfo, Provider}; use crate::{Error, InternalExt}; use async_trait::async_trait; @@ -28,6 +29,7 @@ pub struct LogtoProvider { client: Arc>, endpoint: String, + oauth2_http_client: Arc, http_client: Arc, } @@ -60,12 +62,8 @@ impl LogtoProvider { Ok(Self { client, endpoint: endpoint.to_string(), - http_client: Arc::new( - Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .internal_with_err("Failed to build HTTP client")?, - ), + oauth2_http_client: build_oauth2_http_client()?, + http_client: build_provider_http_client()?, }) } } @@ -97,9 +95,9 @@ impl Provider for LogtoProvider { .client .exchange_code(oauth2::AuthorizationCode::new(code.to_string())) .set_pkce_verifier(verifier) - .request_async(&oauth2::reqwest::Client::new()) + .request_async(self.oauth2_http_client.as_ref()) .await - .internal_with_err("Failed to exchange code")?; + .map_err(|err| map_provider_http_error("Failed to exchange code", err))?; // Fetch user info from Logto let resp = self @@ -111,7 +109,7 @@ impl Provider for LogtoProvider { ) .send() .await - .internal_with_err("Failed to fetch user info")? + .map_err(|err| map_provider_http_error("Failed to fetch user info", err))? .error_for_status() .internal_with_err("Logto API error")?; diff --git a/synctv-core/src/oauth2/providers/mod.rs b/synctv-core/src/oauth2/providers/mod.rs index b57865b1..de5e2314 100644 --- a/synctv-core/src/oauth2/providers/mod.rs +++ b/synctv-core/src/oauth2/providers/mod.rs @@ -18,6 +18,58 @@ pub use google::{GoogleConfig, GoogleProvider}; pub use logto::{LogtoConfig, LogtoProvider}; pub use oidc::{OidcConfig, OidcProvider}; +use crate::{resilience::timeout::HTTP_REQUEST_TIMEOUT, Error, InternalExt}; +use reqwest::Client; +use std::{sync::Arc, time::Duration}; + +pub(super) fn build_provider_http_client_with_timeout(timeout: Duration) -> Result { + Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .timeout(timeout) + .build() + .internal_with_err("Failed to build HTTP client") +} + +pub(super) fn build_provider_http_client() -> Result, Error> { + Ok(Arc::new(build_provider_http_client_with_timeout( + HTTP_REQUEST_TIMEOUT, + )?)) +} + +pub(super) fn build_oauth2_http_client_with_timeout( + timeout: Duration, +) -> Result { + oauth2::reqwest::ClientBuilder::new() + .redirect(oauth2::reqwest::redirect::Policy::none()) + .timeout(timeout) + .build() + .internal_with_err("Failed to build OAuth2 HTTP client") +} + +pub(super) fn build_oauth2_http_client() -> Result, Error> { + Ok(Arc::new(build_oauth2_http_client_with_timeout( + HTTP_REQUEST_TIMEOUT, + )?)) +} + +pub(super) fn map_provider_http_error(context: &str, err: E) -> Error +where + E: std::error::Error + 'static, +{ + let err_debug = format!("{err:?}").to_lowercase(); + let err_display = err.to_string().to_lowercase(); + if crate::resilience::retry::should_retry_error(&err) + || err_debug.contains("timedout") + || err_debug.contains("timeout") + || err_display.contains("timed out") + || err_display.contains("timeout") + { + Error::Timeout(format!("{context}: {err}")) + } else { + Error::Internal(format!("{context}: {err}")) + } +} + /// Build a registry populated with all built-in `OAuth2` providers. #[must_use] pub fn provider_registry() -> crate::oauth2::ProviderRegistry { @@ -28,3 +80,111 @@ pub fn provider_registry() -> crate::oauth2::ProviderRegistry { registry.register("oidc", oidc::oidc_factory); registry } + +#[cfg(test)] +mod tests { + use super::*; + use oauth2::{ + basic::BasicClient, AuthUrl, AuthorizationCode, ClientId, ClientSecret, RedirectUrl, + TokenUrl, + }; + use std::future::pending; + use tokio::{ + io::AsyncReadExt, + net::TcpListener, + task::JoinHandle, + time::{timeout, Duration}, + }; + + async fn spawn_hanging_http_server() -> (String, JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let addr = listener.local_addr().unwrap(); + + let handle = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.unwrap(); + let mut buf = [0_u8; 1024]; + let _ = stream.read(&mut buf).await; + pending::<()>().await; + }); + + (format!("http://{addr}"), handle) + } + + #[tokio::test] + async fn provider_http_client_times_out_hanging_userinfo_requests() { + let client = build_provider_http_client_with_timeout(Duration::from_millis(50)).unwrap(); + let (base_url, server_handle) = spawn_hanging_http_server().await; + + let result = timeout( + Duration::from_millis(250), + client.get(format!("{base_url}/userinfo")).send(), + ) + .await; + server_handle.abort(); + + let err = match result { + Ok(Err(err)) => err, + Ok(Ok(_)) => panic!("expected request to fail with timeout"), + Err(_) => panic!("request client did not enforce its own timeout"), + }; + let mapped = map_provider_http_error("Failed to fetch user info", err); + + assert!(matches!( + mapped, + Error::Timeout(ref msg) if msg.contains("Failed to fetch user info") + )); + } + + #[tokio::test] + async fn provider_http_timeout_maps_to_core_timeout_error() { + let client = build_provider_http_client_with_timeout(Duration::from_millis(50)).unwrap(); + let (base_url, server_handle) = spawn_hanging_http_server().await; + + let err = client + .get(format!("{base_url}/userinfo")) + .send() + .await + .expect_err("request should time out"); + server_handle.abort(); + + let mapped = map_provider_http_error("Failed to fetch user info", err); + assert!(matches!( + mapped, + Error::Timeout(ref msg) if msg.contains("Failed to fetch user info") + )); + } + + #[tokio::test] + async fn token_exchange_client_times_out_hanging_token_endpoint() { + let http_client = build_oauth2_http_client_with_timeout(Duration::from_millis(50)).unwrap(); + let (base_url, server_handle) = spawn_hanging_http_server().await; + + let client = BasicClient::new(ClientId::new("client_id".to_string())) + .set_client_secret(ClientSecret::new("client_secret".to_string())) + .set_auth_uri(AuthUrl::new("https://example.com/auth".to_string()).unwrap()) + .set_token_uri(TokenUrl::new(format!("{base_url}/token")).unwrap()) + .set_redirect_uri( + RedirectUrl::new("https://example.com/callback".to_string()).unwrap(), + ); + + let result = timeout( + Duration::from_millis(250), + client + .exchange_code(AuthorizationCode::new("code".to_string())) + .request_async(&http_client), + ) + .await; + server_handle.abort(); + + let err = match result { + Ok(Err(err)) => err, + Ok(Ok(_)) => panic!("expected token exchange to fail with timeout"), + Err(_) => panic!("token exchange client did not enforce its own timeout"), + }; + let mapped = map_provider_http_error("Failed to exchange code", err); + assert!(matches!( + mapped, + Error::Timeout(ref msg) if msg.contains("Failed to exchange code") + )); + } +} diff --git a/synctv-core/src/oauth2/providers/oidc.rs b/synctv-core/src/oauth2/providers/oidc.rs index f22ae0da..88966b5b 100644 --- a/synctv-core/src/oauth2/providers/oidc.rs +++ b/synctv-core/src/oauth2/providers/oidc.rs @@ -1,5 +1,6 @@ //! Generic OIDC provider +use super::{build_oauth2_http_client, build_provider_http_client, map_provider_http_error}; use crate::oauth2::{OAuth2UserInfo, Provider}; use crate::{Error, InternalExt}; use async_trait::async_trait; @@ -52,6 +53,7 @@ pub struct OidcProvider { resolved: OnceCell, /// Stored config for lazy initialization (only used in issuer-only mode) init_config: OidcInitConfig, + oauth2_http_client: Arc, http_client: Arc, } @@ -85,13 +87,6 @@ impl OidcProvider { redirect_url: String, issuer: &str, ) -> Result { - let http_client = Arc::new( - Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .internal_with_err("Failed to build HTTP client")?, - ); - Ok(Self { resolved: OnceCell::new(), init_config: OidcInitConfig { @@ -101,7 +96,8 @@ impl OidcProvider { issuer: issuer.trim_end_matches('/').to_string(), static_endpoints: None, }, - http_client, + oauth2_http_client: build_oauth2_http_client()?, + http_client: build_provider_http_client()?, }) } @@ -119,13 +115,6 @@ impl OidcProvider { userinfo_url: Option, ) -> Result { let issuer_trimmed = issuer.trim_end_matches('/'); - let http_client = Arc::new( - Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .internal_with_err("Failed to build HTTP client")?, - ); - Ok(Self { resolved: OnceCell::new(), init_config: OidcInitConfig { @@ -139,7 +128,8 @@ impl OidcProvider { userinfo_url, }), }, - http_client, + oauth2_http_client: build_oauth2_http_client()?, + http_client: build_provider_http_client()?, }) } @@ -168,10 +158,13 @@ impl OidcProvider { .get(&discovery_url) .send() .await - .map_err(|e| { - Error::Internal(format!( - "Failed to fetch OIDC discovery document from {discovery_url}: {e}" - )) + .map_err(|err| { + map_provider_http_error( + &format!( + "Failed to fetch OIDC discovery document from {discovery_url}" + ), + err, + ) })? .error_for_status() .map_err(|e| { @@ -248,9 +241,9 @@ impl Provider for OidcProvider { .client .exchange_code(oauth2::AuthorizationCode::new(code.to_string())) .set_pkce_verifier(verifier) - .request_async(&oauth2::reqwest::Client::new()) + .request_async(self.oauth2_http_client.as_ref()) .await - .internal_with_err("Failed to exchange code")?; + .map_err(|err| map_provider_http_error("Failed to exchange code", err))?; // Fetch user info from userinfo endpoint let userinfo_url = resolved.userinfo_url.as_ref().ok_or_else(|| { @@ -268,7 +261,7 @@ impl Provider for OidcProvider { ) .send() .await - .internal_with_err("Failed to fetch user info")? + .map_err(|err| map_provider_http_error("Failed to fetch user info", err))? .error_for_status() .internal_with_err("OIDC API error")?; diff --git a/synctv-core/src/provider/alist.rs b/synctv-core/src/provider/alist.rs index fac2deb8..6c705795 100644 --- a/synctv-core/src/provider/alist.rs +++ b/synctv-core/src/provider/alist.rs @@ -346,6 +346,11 @@ impl AlistProvider { #[async_trait] impl MediaProvider for AlistProvider { + #[cfg(test)] + fn test_client_manager_marker(&self) -> Option { + Some(self.client_manager.marker()) + } + fn name(&self) -> &'static str { Self::NAME } diff --git a/synctv-core/src/provider/bilibili.rs b/synctv-core/src/provider/bilibili.rs index 43140c2c..44463d1d 100644 --- a/synctv-core/src/provider/bilibili.rs +++ b/synctv-core/src/provider/bilibili.rs @@ -292,6 +292,11 @@ impl TryFrom<&Value> for BilibiliSourceConfig { #[async_trait] impl MediaProvider for BilibiliProvider { + #[cfg(test)] + fn test_client_manager_marker(&self) -> Option { + Some(self.client_manager.marker()) + } + fn name(&self) -> &'static str { Self::NAME } diff --git a/synctv-core/src/provider/emby.rs b/synctv-core/src/provider/emby.rs index ef2e5733..f39a3cfb 100644 --- a/synctv-core/src/provider/emby.rs +++ b/synctv-core/src/provider/emby.rs @@ -405,6 +405,11 @@ impl TryFrom<&Value> for EmbySourceConfig { #[async_trait] impl MediaProvider for EmbyProvider { + #[cfg(test)] + fn test_client_manager_marker(&self) -> Option { + Some(self.client_manager.marker()) + } + fn name(&self) -> &'static str { Self::NAME } diff --git a/synctv-core/src/provider/provider_client.rs b/synctv-core/src/provider/provider_client.rs index ff2c091f..0ea15343 100644 --- a/synctv-core/src/provider/provider_client.rs +++ b/synctv-core/src/provider/provider_client.rs @@ -21,12 +21,17 @@ use super::ProviderError; use async_trait::async_trait; +#[cfg(test)] +use std::sync::atomic::{AtomicUsize, Ordering as AtomicOrdering}; use std::sync::Arc; use std::time::Duration; use synctv_media_providers::alist::{AlistError, AlistInterface}; use synctv_media_providers::grpc::alist::{FsGetResp, FsListResp, FsOtherResp}; use tonic::{Code, Request, Status}; +#[cfg(test)] +static PROVIDER_CLIENT_MANAGER_MARKER_SEQ: AtomicUsize = AtomicUsize::new(1); + #[derive(Clone, Debug)] pub struct RemoteProviderConnection { channel: tonic::transport::Channel, @@ -35,10 +40,7 @@ pub struct RemoteProviderConnection { impl RemoteProviderConnection { #[must_use] - pub fn new( - channel: tonic::transport::Channel, - auth_secret: Option>, - ) -> Self { + pub fn new(channel: tonic::transport::Channel, auth_secret: Option>) -> Self { Self { channel, auth_secret: auth_secret.map(|secret| Arc::::from(secret.into())), @@ -92,19 +94,16 @@ macro_rules! impl_grpc_method { let mut client = _client_mod::$client_name::new(self.connection.channel()); let request = build_grpc_request(self.connection.auth_secret(), request) .map_err(<$error>::from)?; - let response = tokio::time::timeout( - GRPC_REQUEST_TIMEOUT, - client.$method(request), - ) - .await - .map_err(|_| { - <$error>::Network(format!( - "gRPC request timeout ({}s) for {}", - GRPC_REQUEST_TIMEOUT.as_secs(), - stringify!($method), - )) - })? - .map_err(|e| <$error>::from(map_grpc_status(stringify!($method), e)))?; + let response = tokio::time::timeout(GRPC_REQUEST_TIMEOUT, client.$method(request)) + .await + .map_err(|_| { + <$error>::Network(format!( + "gRPC request timeout ({}s) for {}", + GRPC_REQUEST_TIMEOUT.as_secs(), + stringify!($method), + )) + })? + .map_err(|e| <$error>::from(map_grpc_status(stringify!($method), e)))?; Ok(response.into_inner()) }) } @@ -116,7 +115,10 @@ fn build_grpc_request( payload: T, ) -> Result, synctv_media_providers::ProviderClientError> { let mut request = Request::new(payload); - let Some(auth_secret) = auth_secret.map(str::trim).filter(|secret| !secret.is_empty()) else { + let Some(auth_secret) = auth_secret + .map(str::trim) + .filter(|secret| !secret.is_empty()) + else { return Ok(request); }; @@ -132,7 +134,9 @@ fn build_grpc_request( Ok(request) } -pub(crate) fn validate_auth_secret(auth_secret: Option<&str>) -> Result, ProviderError> { +pub(crate) fn validate_auth_secret( + auth_secret: Option<&str>, +) -> Result, ProviderError> { match auth_secret.map(str::trim) { Some("") => Err(ProviderError::InvalidConfig( "remote provider auth secret must not be empty".to_string(), @@ -164,10 +168,7 @@ fn grpc_status_to_http_status(code: Code) -> Option { } } -fn map_grpc_status( - context: &str, - status: Status, -) -> synctv_media_providers::ProviderClientError { +fn map_grpc_status(context: &str, status: Status) -> synctv_media_providers::ProviderClientError { let message = status.message().to_string(); match status.code() { Code::Unauthenticated => synctv_media_providers::ProviderClientError::Auth(message), @@ -238,6 +239,8 @@ pub struct ProviderClientManager { local_bilibili: BilibiliClientArc, /// Local Emby client (singleton within this manager) local_emby: EmbyClientArc, + #[cfg(test)] + marker: usize, } impl std::fmt::Debug for ProviderClientManager { @@ -278,6 +281,8 @@ impl ProviderClientManager { local_emby: Arc::new(synctv_media_providers::emby::EmbyService::with_client( client, )), + #[cfg(test)] + marker: PROVIDER_CLIENT_MANAGER_MARKER_SEQ.fetch_add(1, AtomicOrdering::Relaxed), } } @@ -292,6 +297,8 @@ impl ProviderClientManager { local_alist: Arc::new(alist), local_bilibili: Arc::new(bilibili), local_emby: Arc::new(emby), + #[cfg(test)] + marker: PROVIDER_CLIENT_MANAGER_MARKER_SEQ.fetch_add(1, AtomicOrdering::Relaxed), } } @@ -308,6 +315,8 @@ impl ProviderClientManager { local_alist, local_bilibili, local_emby, + #[cfg(test)] + marker: PROVIDER_CLIENT_MANAGER_MARKER_SEQ.fetch_add(1, AtomicOrdering::Relaxed), } } @@ -370,6 +379,11 @@ impl ProviderClientManager { None => self.local_emby_client(), } } + + #[cfg(test)] + pub(crate) fn marker(&self) -> usize { + self.marker + } } // ============================================================================ @@ -450,18 +464,15 @@ impl AlistInterface for GrpcAlistClient { use synctv_media_providers::grpc::alist::alist_client::AlistClient; let mut client = AlistClient::new(self.connection.channel()); let request = build_grpc_request(self.connection.auth_secret(), request)?; - let response = tokio::time::timeout( - GRPC_REQUEST_TIMEOUT, - client.login(request), - ) - .await - .map_err(|_| { - AlistError::Network(format!( - "gRPC request timeout ({}s) for login", - GRPC_REQUEST_TIMEOUT.as_secs(), - )) - })? - .map_err(|e| AlistError::from(map_grpc_status("login", e)))?; + let response = tokio::time::timeout(GRPC_REQUEST_TIMEOUT, client.login(request)) + .await + .map_err(|_| { + AlistError::Network(format!( + "gRPC request timeout ({}s) for login", + GRPC_REQUEST_TIMEOUT.as_secs(), + )) + })? + .map_err(|e| AlistError::from(map_grpc_status("login", e)))?; Ok(response.into_inner().token) } } @@ -1068,10 +1079,7 @@ mod tests { #[test] fn test_map_grpc_status_unauthenticated_to_auth() { - let error = map_grpc_status( - "login", - Status::unauthenticated("Invalid provider secret"), - ); + let error = map_grpc_status("login", Status::unauthenticated("Invalid provider secret")); assert!(matches!( error, ProviderClientError::Auth(message) if message == "Invalid provider secret" @@ -1080,8 +1088,7 @@ mod tests { #[test] fn test_map_grpc_status_invalid_argument_to_invalid_config() { - let error = - map_grpc_status("fs_get", Status::invalid_argument("missing host parameter")); + let error = map_grpc_status("fs_get", Status::invalid_argument("missing host parameter")); assert!(matches!( error, ProviderClientError::InvalidConfig(message) if message == "missing host parameter" diff --git a/synctv-core/src/provider/traits.rs b/synctv-core/src/provider/traits.rs index 06238a0a..0d710969 100644 --- a/synctv-core/src/provider/traits.rs +++ b/synctv-core/src/provider/traits.rs @@ -183,6 +183,11 @@ pub trait MediaProvider: Send + Sync { None } + #[cfg(test)] + fn test_client_manager_marker(&self) -> Option { + None + } + // ========== Validation ========== /// Validate `source_config` before saving to database diff --git a/synctv-core/src/repository/user_oauth_provider.rs b/synctv-core/src/repository/user_oauth_provider.rs index 94378162..8dc31a0f 100644 --- a/synctv-core/src/repository/user_oauth_provider.rs +++ b/synctv-core/src/repository/user_oauth_provider.rs @@ -59,17 +59,17 @@ impl UserOAuthProviderRepository { { let id = nanoid::nanoid!(12); - sqlx::query( + let result = sqlx::query( r" INSERT INTO oauth2_clients (id, provider, provider_user_id, user_id, username, email, avatar_url) VALUES ($1, $2, $3, $4, $5, $6, $7) ON CONFLICT (provider, provider_user_id) DO UPDATE SET - user_id = EXCLUDED.user_id, username = EXCLUDED.username, email = EXCLUDED.email, avatar_url = EXCLUDED.avatar_url, updated_at = CURRENT_TIMESTAMP + WHERE oauth2_clients.user_id = EXCLUDED.user_id " ) .bind(&id) @@ -82,6 +82,12 @@ impl UserOAuthProviderRepository { .execute(executor) .await?; + if result.rows_affected() == 0 { + return Err(crate::Error::AlreadyExists( + "OAuth2 provider identity is already linked to another user".to_string(), + )); + } + Ok(()) } diff --git a/synctv-core/src/resilience.rs b/synctv-core/src/resilience.rs index 0ffb1204..5020678c 100644 --- a/synctv-core/src/resilience.rs +++ b/synctv-core/src/resilience.rs @@ -99,9 +99,16 @@ pub mod retry { /// Checks the error for known transient I/O error kinds, then falls back to /// string matching for errors that don't expose `std::io::Error` directly. pub fn should_retry_error(err: &(dyn std::error::Error + 'static)) -> bool { - // Check top-level error for std::io::Error with transient kinds - if let Some(io_err) = err.downcast_ref::() { - return is_transient_io_error(io_err); + // Walk the full error chain because reqwest/oauth2 often wrap the + // underlying timeout several layers deep. + let mut current: Option<&(dyn std::error::Error + 'static)> = Some(err); + while let Some(candidate) = current { + if let Some(io_err) = candidate.downcast_ref::() { + if is_transient_io_error(io_err) { + return true; + } + } + current = candidate.source(); } // Fallback: check the display message for transient indicators. @@ -227,4 +234,25 @@ mod tests { let not_found = std::io::Error::new(std::io::ErrorKind::NotFound, "not found"); assert!(!retry::should_retry_error(¬_found)); } + + #[test] + fn test_should_retry_error_walks_source_chain() { + #[derive(Debug)] + struct Wrapper(std::io::Error); + + impl std::fmt::Display for Wrapper { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + write!(f, "wrapped transport error") + } + } + + impl std::error::Error for Wrapper { + fn source(&self) -> Option<&(dyn std::error::Error + 'static)> { + Some(&self.0) + } + } + + let wrapped = Wrapper(std::io::Error::new(std::io::ErrorKind::TimedOut, "timeout")); + assert!(retry::should_retry_error(&wrapped)); + } } diff --git a/synctv-core/src/service/audit_partition_manager.rs b/synctv-core/src/service/audit_partition_manager.rs index 0c1e3487..50e27f72 100644 --- a/synctv-core/src/service/audit_partition_manager.rs +++ b/synctv-core/src/service/audit_partition_manager.rs @@ -380,7 +380,10 @@ async fn run_audit_partition_maintenance(manager: &AuditPartitionManager) { match manager.check_health().await { Ok(health) => { if health.missing_count > 0 { - warn!("Found {} missing partitions, creating now", health.missing_count); + warn!( + "Found {} missing partitions, creating now", + health.missing_count + ); if let Err(e) = manager.ensure_future_partitions_with_retry(6).await { tracing::error!( error = %e, @@ -413,9 +416,7 @@ async fn wait_for_initial_leader( } if !logged_wait { - info!( - "Delaying initial {task_name} run until cluster leadership is established" - ); + info!("Delaying initial {task_name} run until cluster leadership is established"); logged_wait = true; } diff --git a/synctv-core/src/service/auth/security_pipeline.rs b/synctv-core/src/service/auth/security_pipeline.rs index 811fd9d8..e72788f0 100644 --- a/synctv-core/src/service/auth/security_pipeline.rs +++ b/synctv-core/src/service/auth/security_pipeline.rs @@ -104,9 +104,10 @@ impl SecurityPipeline { match err { Error::Authentication(_) => AuthErrorCategory::Authentication, Error::Authorization(_) | Error::EmailNotVerified => AuthErrorCategory::Authorization, - Error::ServiceUnavailable(_) | Error::Database(_) | Error::Redis(_) | Error::Timeout(_) => { - AuthErrorCategory::Unavailable - } + Error::ServiceUnavailable(_) + | Error::Database(_) + | Error::Redis(_) + | Error::Timeout(_) => AuthErrorCategory::Unavailable, _ => AuthErrorCategory::Internal, } } @@ -581,12 +582,8 @@ mod tests { let jwt_service = JwtService::new("test-secret-key-for-security-pipeline-unit-tests-min-32") .expect("failed to create jwt service"); - let username_cache = UsernameCache::new( - Arc::new(NoopCacheL2), - "test:username:".to_string(), - 100, - 0, - ); + let username_cache = + UsernameCache::new(Arc::new(NoopCacheL2), "test:username:".to_string(), 100, 0); let token_blacklist = Arc::new(InMemoryTokenBlacklistStore::new(1000, 3600, 86400)); let key_builder = KeyBuilder::new("test"); let brute_force = BruteForceProtection::in_memory("test".to_string()); diff --git a/synctv-core/src/service/auth/token_blacklist.rs b/synctv-core/src/service/auth/token_blacklist.rs index 76f74d57..768067d4 100644 --- a/synctv-core/src/service/auth/token_blacklist.rs +++ b/synctv-core/src/service/auth/token_blacklist.rs @@ -1491,6 +1491,7 @@ impl RedisSyncableTokenBlacklistStore { stats.family_synced += 1; } Err(e) => { + stats.family_failed += 1; tracing::warn!( key = %key, error = %e, @@ -2193,6 +2194,28 @@ mod tests { ); } + #[tokio::test] + async fn test_sync_pending_writes_counts_failed_family_syncs() { + let primary = Arc::new(AlwaysFailTokenBlacklistStore) as Arc; + let store = RedisSyncableTokenBlacklistStore::with_defaults(primary); + + let key = "family:pending_sync_failure"; + let timestamp = 1_700_000_000_i64; + + let result = store.set_family_revoked(key, timestamp, 3600).await; + assert!(result.is_err(), "initial primary write should fail"); + assert_eq!(store.pending_family_count(), 1); + + let stats = store + .sync_pending_writes() + .await + .expect("sync should report stats even when primary write fails"); + + assert_eq!(stats.family_synced, 0); + assert_eq!(stats.family_failed, 1); + assert_eq!(store.pending_family_count(), 1); + } + // ---- Hybrid/Tiered fallback scenario tests ---- #[tokio::test] diff --git a/synctv-core/src/service/chat_partition_manager.rs b/synctv-core/src/service/chat_partition_manager.rs index dd71f703..e25d4978 100644 --- a/synctv-core/src/service/chat_partition_manager.rs +++ b/synctv-core/src/service/chat_partition_manager.rs @@ -244,9 +244,7 @@ async fn wait_for_initial_leader( } if !logged_wait { - info!( - "Delaying initial {task_name} run until cluster leadership is established" - ); + info!("Delaying initial {task_name} run until cluster leadership is established"); logged_wait = true; } diff --git a/synctv-core/src/service/distributed_lock.rs b/synctv-core/src/service/distributed_lock.rs index be455260..449afb9a 100644 --- a/synctv-core/src/service/distributed_lock.rs +++ b/synctv-core/src/service/distributed_lock.rs @@ -61,6 +61,30 @@ use redis::aio::ConnectionManager as RedisConnectionManager; use redis::Script; use std::future::Future; +async fn run_distributed_lock_redis_op(operation: impl Into, future: F) -> Result +where + F: Future>, +{ + let operation = operation.into(); + tokio::time::timeout(crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, future) + .await + .map_err(|_| Error::Timeout(format!("Redis timeout: {operation}")))? + .internal_with_err(&format!("Failed to {operation}")) +} + +async fn run_distributed_lock_client_op( + key: &str, + timeout: std::time::Duration, + future: F, +) -> Result +where + F: Future>, +{ + tokio::time::timeout(timeout, future) + .await + .map_err(|_| Error::Timeout(format!("Lock operation timed out for key: {key}")))? +} + /// Abstraction over a distributed migration lock. /// /// Consumers that only need acquire/release semantics (e.g. `run_migrations`) @@ -306,19 +330,11 @@ impl DistributedLock { ", ); - tokio::time::timeout( - crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, + run_distributed_lock_redis_op( + format!("generate fencing token for lock '{key}'"), script.key(&token_key).invoke_async::(&mut conn), ) .await - .map_err(|_| { - Error::Internal(format!( - "Redis timeout: generate fencing token for lock '{key}'" - )) - })? - .internal_with_err(&format!( - "Failed to generate fencing token for lock '{key}'" - )) } /// Acquire a lock (using SET NX EX atomic operation) @@ -396,8 +412,8 @@ impl DistributedLock { // SET key value NX EX ttl // NX: Only set if not exists // EX: Set expiration time - let result: Option = tokio::time::timeout( - crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, + let result: Option = run_distributed_lock_redis_op( + "acquire lock", redis::cmd("SET") .arg(&lock_key) .arg(&lock_value) @@ -406,9 +422,7 @@ impl DistributedLock { .arg(ttl_seconds) .query_async(&mut conn), ) - .await - .map_err(|_| Error::Internal("Redis timeout: acquire lock".to_string()))? - .internal_with_err("Failed to acquire lock")?; + .await?; if result.is_some() { // Generate fencing token only if requested (saves Redis round-trip) @@ -463,16 +477,14 @@ impl DistributedLock { let mut conn = self.conn().await; - let result: i32 = tokio::time::timeout( - crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, + let result: i32 = run_distributed_lock_redis_op( + "release lock", script .key(&lock_key) .arg(lock_value) .invoke_async::(&mut conn), ) - .await - .map_err(|_| Error::Internal("Redis timeout: release lock".to_string()))? - .internal_with_err("Failed to release lock")?; + .await?; let released = result == 1; if released { @@ -521,7 +533,7 @@ impl DistributedLock { // ttl_seconds, so we allow ttl + 5s for network round-trips. let client_timeout = std::time::Duration::from_secs(ttl_seconds + 5); - tokio::time::timeout(client_timeout, async { + run_distributed_lock_client_op(key, client_timeout, async { // Try to acquire lock let lock_value = self .acquire(key, ttl_seconds) @@ -543,7 +555,6 @@ impl DistributedLock { result }) .await - .map_err(|_| Error::Internal(format!("Lock operation timed out for key: {key}")))? } /// Try to acquire a lock and execute an operation @@ -571,7 +582,7 @@ impl DistributedLock { { let client_timeout = std::time::Duration::from_secs(ttl_seconds + 5); - tokio::time::timeout(client_timeout, async { + run_distributed_lock_client_op(key, client_timeout, async { // Try to acquire lock let lock_value = match self.acquire(key, ttl_seconds).await? { Some(value) => value, @@ -593,7 +604,6 @@ impl DistributedLock { result.map(Some) }) .await - .map_err(|_| Error::Internal(format!("Lock operation timed out for key: {key}")))? } /// Execute an operation with automatic lock acquisition and release (with fencing token) @@ -625,7 +635,7 @@ impl DistributedLock { { let client_timeout = std::time::Duration::from_secs(ttl_seconds + 5); - tokio::time::timeout(client_timeout, async { + run_distributed_lock_client_op(key, client_timeout, async { // Try to acquire lock with token let (lock_value, fencing_token) = self .acquire_with_token(key, ttl_seconds) @@ -647,7 +657,6 @@ impl DistributedLock { result }) .await - .map_err(|_| Error::Internal(format!("Lock operation timed out for key: {key}")))? } /// Try to acquire a lock and execute an operation (with fencing token) @@ -675,7 +684,7 @@ impl DistributedLock { { let client_timeout = std::time::Duration::from_secs(ttl_seconds + 5); - tokio::time::timeout(client_timeout, async { + run_distributed_lock_client_op(key, client_timeout, async { // Try to acquire lock with token let (lock_value, fencing_token) = match self.acquire_with_token(key, ttl_seconds).await? { @@ -698,7 +707,6 @@ impl DistributedLock { result.map(Some) }) .await - .map_err(|_| Error::Internal(format!("Lock operation timed out for key: {key}")))? } /// Extend lock TTL (refresh expiration) @@ -724,17 +732,15 @@ impl DistributedLock { let mut conn = self.conn().await; - let result: i32 = tokio::time::timeout( - crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, + let result: i32 = run_distributed_lock_redis_op( + "extend lock", script .key(&lock_key) .arg(lock_value) .arg(ttl_seconds) .invoke_async::(&mut conn), ) - .await - .map_err(|_| Error::Internal("Redis timeout: extend lock".to_string()))? - .internal_with_err("Failed to extend lock")?; + .await?; Ok(result == 1) } @@ -1327,6 +1333,45 @@ mod tests { assert_eq!(client_timeout, std::time::Duration::from_secs(15)); } + #[tokio::test(start_paused = true)] + async fn test_distributed_lock_redis_timeout_maps_to_timeout_error() { + let timeout_future = run_distributed_lock_redis_op("acquire lock", async { + std::future::pending::<()>().await; + #[allow(unreachable_code)] + Ok::<(), redis::RedisError>(()) + }); + + tokio::pin!(timeout_future); + tokio::task::yield_now().await; + tokio::time::advance(crate::resilience::timeout::REDIS_OPERATION_TIMEOUT).await; + + let err = timeout_future.await.expect_err("operation should time out"); + assert!(matches!( + err, + Error::Timeout(ref msg) if msg == "Redis timeout: acquire lock" + )); + } + + #[tokio::test(start_paused = true)] + async fn test_distributed_lock_client_timeout_maps_to_timeout_error() { + let timeout_future = + run_distributed_lock_client_op("test-key", std::time::Duration::from_secs(15), async { + std::future::pending::<()>().await; + #[allow(unreachable_code)] + Ok::<(), Error>(()) + }); + + tokio::pin!(timeout_future); + tokio::task::yield_now().await; + tokio::time::advance(std::time::Duration::from_secs(15)).await; + + let err = timeout_future.await.expect_err("operation should time out"); + assert!(matches!( + err, + Error::Timeout(ref msg) if msg == "Lock operation timed out for key: test-key" + )); + } + #[test] fn test_lua_script_release_logic() { // The release Lua script logic: diff --git a/synctv-core/src/service/notification_partition_manager.rs b/synctv-core/src/service/notification_partition_manager.rs index a0b68f09..7f9296fa 100644 --- a/synctv-core/src/service/notification_partition_manager.rs +++ b/synctv-core/src/service/notification_partition_manager.rs @@ -183,9 +183,7 @@ async fn wait_for_initial_leader( } if !logged_wait { - info!( - "Delaying initial {task_name} run until cluster leadership is established" - ); + info!("Delaying initial {task_name} run until cluster leadership is established"); logged_wait = true; } diff --git a/synctv-core/src/service/oauth2.rs b/synctv-core/src/service/oauth2.rs index 2be72cf5..525fc0e6 100644 --- a/synctv-core/src/service/oauth2.rs +++ b/synctv-core/src/service/oauth2.rs @@ -13,6 +13,7 @@ use serde::{Deserialize, Serialize}; use std::collections::HashMap; +use std::future::Future; use std::sync::Arc; use tokio::sync::RwLock; use tracing::{debug, info}; @@ -70,25 +71,61 @@ pub trait OAuthStateStore: Send + Sync { pub struct RedisOAuthStateStore { /// Shared Redis connection handle that follows Sentinel failover. conn: std::sync::Arc>, + key_prefix: String, } impl RedisOAuthStateStore { + async fn run_redis_op(&self, operation: &'static str, future: F) -> Result + where + F: Future>, + { + run_oauth_state_redis_op(operation, future).await + } + + fn normalize_key_prefix(prefix: impl Into) -> String { + let key_prefix = prefix.into(); + if key_prefix.is_empty() || key_prefix.ends_with(':') { + key_prefix + } else { + format!("{key_prefix}:") + } + } + /// Create from the shared `Arc>`. #[must_use] - pub const fn new( + pub fn new( conn: std::sync::Arc>, + key_prefix: impl Into, ) -> Self { - Self { conn } + Self { + conn, + key_prefix: Self::normalize_key_prefix(key_prefix), + } } /// Acquire a fresh ConnectionManager clone from the shared handle. async fn get_conn(&self) -> redis::aio::ConnectionManager { self.conn.read().await.clone() } + + fn redis_key(&self, token_id: &str) -> String { + format!("{}{}", self.key_prefix, Self::state_key_suffix(token_id)) + } + + fn state_key_suffix(token_id: &str) -> String { + format!("oauth2:state:{token_id}") + } } -/// Redis key prefix for `OAuth2` state tokens -const OAUTH2_STATE_KEY_PREFIX: &str = "oauth2:state:"; +async fn run_oauth_state_redis_op(operation: &'static str, future: F) -> Result +where + F: Future>, +{ + tokio::time::timeout(crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, future) + .await + .map_err(|_| Error::Timeout(format!("Redis timeout: {operation}")))? + .internal_with_err(&format!("Failed to {operation}")) +} #[async_trait::async_trait] impl OAuthStateStore for RedisOAuthStateStore { @@ -102,19 +139,18 @@ impl OAuthStateStore for RedisOAuthStateStore { state: &OAuth2State, ttl: std::time::Duration, ) -> Result<()> { - let key = format!("{OAUTH2_STATE_KEY_PREFIX}{token_id}"); + let key = self.redis_key(token_id); let value = serde_json::to_string(state).internal_with_err("Failed to serialize OAuth2 state")?; let mut conn = self.get_conn().await; use redis::AsyncCommands; - let _: () = tokio::time::timeout( - crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, - conn.set_ex(&key, value, ttl.as_secs()), - ) - .await - .map_err(|_| Error::Internal("Redis timeout: store OAuth2 state".to_string()))? - .internal_with_err("Failed to store OAuth2 state in Redis")?; + let _: () = self + .run_redis_op( + "store OAuth2 state in Redis", + conn.set_ex(&key, value, ttl.as_secs()), + ) + .await?; debug!( "Stored OAuth2 state in Redis for token {}", @@ -124,7 +160,7 @@ impl OAuthStateStore for RedisOAuthStateStore { } async fn consume(&self, token_id: &str) -> Result> { - let key = format!("{OAUTH2_STATE_KEY_PREFIX}{token_id}"); + let key = self.redis_key(token_id); let mut conn = self.get_conn().await; // Atomic GET + DEL via Lua script (same pattern as WsTicketService) @@ -138,13 +174,12 @@ impl OAuthStateStore for RedisOAuthStateStore { "#, ); - let value: Option = tokio::time::timeout( - crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, - lua_script.key(&key).invoke_async(&mut conn), - ) - .await - .map_err(|_| Error::Internal("Redis timeout: consume OAuth2 state".to_string()))? - .internal_with_err("Failed to consume OAuth2 state from Redis")?; + let value: Option = self + .run_redis_op( + "consume OAuth2 state from Redis", + lua_script.key(&key).invoke_async(&mut conn), + ) + .await?; match value { Some(json) => { @@ -802,6 +837,20 @@ impl OAuth2Service { let pool = self.repository.pool(); let mut tx = pool.begin().await?; + let advisory_lock_key = format!( + "oauth2:{}:{}", + provider.as_str(), + user_info.provider_user_id + ); + // Serialize creation for a single external identity so concurrent logins + // cannot race on local username/email creation before the winning mapping + // becomes visible. + sqlx::query("SELECT pg_advisory_xact_lock(hashtextextended($1, 0))") + .bind(&advisory_lock_key) + .execute(&mut *tx) + .await + .internal_with_err("Failed to acquire OAuth2 identity advisory lock")?; + // Re-check inside the transaction to guard against the race where another // concurrent request created the user between our initial lookup and here. let existing = self @@ -815,29 +864,127 @@ impl OAuth2Service { return Ok((mapping.user_id, false)); } - // Generate a random password (OAuth2 users authenticate via provider, not password). + let (base_username, candidates) = user_service + .oauth2_username_candidates(&user_info.provider_user_id, &user_info.username)?; let random_password = nanoid::nanoid!(32); + let password_hash = crate::service::auth::hash_password(&random_password).await?; + let user_email = user_info.email.clone(); + + let mut new_user = None; + for (attempt, candidate) in candidates.iter().enumerate() { + let savepoint = format!("oauth2_user_create_{attempt}"); + sqlx::query(&format!("SAVEPOINT {savepoint}")) + .execute(&mut *tx) + .await + .internal_with_err("Failed to create OAuth2 user savepoint")?; - // Create the user record inside the transaction. - let new_user: User = user_service - .register_with_executor( - user_info.username.clone(), - user_info.email.clone(), - random_password, + let user = User::new_with_status( + candidate.clone(), + user_email.clone(), + password_hash.clone(), SignupMethod::OAuth2, - &mut *tx, - ) - .await?; + crate::models::UserStatus::Active, + ); + match user_service + .repository + .create_with_executor(&user, &mut *tx) + .await + { + Ok(created_user) => { + sqlx::query(&format!("RELEASE SAVEPOINT {savepoint}")) + .execute(&mut *tx) + .await + .internal_with_err("Failed to release OAuth2 user savepoint")?; + + user_service + .cache_oauth2_username_best_effort(&created_user.id, candidate) + .await; + + if candidate == &base_username { + tracing::info!( + "Created new user {} (username='{}', sanitized from '{}') via OAuth2 provider {} (provider_user_id={})", + created_user.id.as_str(), + candidate, + user_info.username, + provider.as_str(), + user_info.provider_user_id + ); + } else { + tracing::info!( + "Username '{}' was taken; created user {} as '{}' (original '{}') via OAuth2 provider {} (provider_user_id={})", + base_username, + created_user.id.as_str(), + candidate, + user_info.username, + provider.as_str(), + user_info.provider_user_id + ); + } + + new_user = Some(created_user); + break; + } + Err(Error::AlreadyExists(ref msg)) + if msg.contains("username") || msg.contains("Username") => + { + sqlx::query(&format!("ROLLBACK TO SAVEPOINT {savepoint}")) + .execute(&mut *tx) + .await + .internal_with_err( + "Failed to roll back OAuth2 user savepoint after username collision", + )?; + continue; + } + Err(err) => { + sqlx::query(&format!("ROLLBACK TO SAVEPOINT {savepoint}")) + .execute(&mut *tx) + .await + .internal_with_err( + "Failed to roll back OAuth2 user savepoint after create error", + )?; + return Err(err); + } + } + } + + let new_user: User = new_user.ok_or_else(|| { + Error::Internal(format!( + "Could not generate a unique username for base '{}' after {} attempts", + user_info.username, + candidates.len() + )) + })?; // Link the OAuth2 provider mapping inside the same transaction. - self.upsert_user_provider_with_executor( - &new_user.id, - provider, - &user_info.provider_user_id, - user_info, - &mut *tx, - ) - .await?; + match self + .upsert_user_provider_with_executor( + &new_user.id, + provider, + &user_info.provider_user_id, + user_info, + &mut *tx, + ) + .await + { + Ok(()) => {} + Err(Error::AlreadyExists(_)) => { + // Another concurrent request bound this provider identity first. + // Roll back the provisional user so we do not commit an orphan row, + // then return the winning mapping. + tx.rollback().await?; + let existing = self + .repository + .find_by_provider(provider, &user_info.provider_user_id) + .await? + .ok_or_else(|| { + Error::Internal( + "OAuth2 mapping conflicted but could not be reloaded".to_string(), + ) + })?; + return Ok((existing.user_id, false)); + } + Err(err) => return Err(err), + } // Set email_verified if the provider confirmed the email. if user_info.email_verified && user_info.email.is_some() { @@ -2431,4 +2578,23 @@ mod tests { "Exactly one consume should succeed (single-use guarantee)" ); } + + #[tokio::test(start_paused = true)] + async fn test_redis_state_store_timeout_maps_to_timeout_error() { + let timeout_future = run_oauth_state_redis_op("store OAuth2 state in Redis", async { + std::future::pending::<()>().await; + #[allow(unreachable_code)] + Ok::<(), redis::RedisError>(()) + }); + + tokio::pin!(timeout_future); + tokio::task::yield_now().await; + tokio::time::advance(crate::resilience::timeout::REDIS_OPERATION_TIMEOUT).await; + + let err = timeout_future.await.expect_err("operation should time out"); + assert!(matches!( + err, + Error::Timeout(ref msg) if msg == "Redis timeout: store OAuth2 state in Redis" + )); + } } diff --git a/synctv-core/src/service/providers_manager.rs b/synctv-core/src/service/providers_manager.rs index de61de44..a6b48519 100644 --- a/synctv-core/src/service/providers_manager.rs +++ b/synctv-core/src/service/providers_manager.rs @@ -5,7 +5,7 @@ use crate::provider::{ AlistProvider, BilibiliProvider, DirectUrlProvider, EmbyProvider, LiveProxyProvider, - MediaProvider, RtmpProvider, + MediaProvider, ProviderClientManager, RtmpProvider, }; use crate::service::RemoteProviderManager; use crate::Config; @@ -78,6 +78,9 @@ pub struct ProvidersManager { /// Provider instance manager (for local/remote dispatch) instance_manager: Arc, + /// Default injected local provider clients used by provider instances + /// when they do not specify a per-instance HTTP transport override. + default_client_manager: Arc, /// Default connect timeout used when building per-instance override clients. default_provider_connect_timeout: std::time::Duration, } @@ -99,13 +102,17 @@ impl ProvidersManager { #[must_use] pub fn new_with_provider_http_client( instance_manager: Arc, - _default_provider_http_client: reqwest::Client, + default_provider_http_client: reqwest::Client, default_provider_connect_timeout: std::time::Duration, ) -> Self { + let default_client_manager = Arc::new( + ProviderClientManager::new_with_provider_http_client(default_provider_http_client), + ); let mut manager = Self { factories: HashMap::new(), instances: Arc::new(RwLock::new(HashMap::new())), instance_manager, + default_client_manager, default_provider_connect_timeout, }; @@ -123,6 +130,7 @@ impl ProvidersManager { /// Register all built-in provider factories fn register_builtin_providers(&mut self) { + let default_client_manager = Arc::clone(&self.default_client_manager); let default_provider_connect_timeout = self.default_provider_connect_timeout; // Alist factory - reads optional timeout from config self.register_factory( @@ -140,13 +148,17 @@ impl ProvidersManager { ), ) } else { - AlistProvider::new(instance_manager) + AlistProvider::with_client_manager( + instance_manager, + Arc::clone(&default_client_manager), + ) }; Ok(Arc::new(provider)) }), ); // Bilibili factory - reads optional timeout from config + let default_client_manager = Arc::clone(&self.default_client_manager); let default_provider_connect_timeout = self.default_provider_connect_timeout; self.register_factory( "bilibili", @@ -163,13 +175,17 @@ impl ProvidersManager { ), ) } else { - BilibiliProvider::new(instance_manager) + BilibiliProvider::with_client_manager( + instance_manager, + Arc::clone(&default_client_manager), + ) }; Ok(Arc::new(provider)) }), ); // Emby factory - reads optional timeout from config + let default_client_manager = Arc::clone(&self.default_client_manager); let default_provider_connect_timeout = self.default_provider_connect_timeout; self.register_factory( "emby", @@ -186,7 +202,10 @@ impl ProvidersManager { ), ) } else { - EmbyProvider::new(instance_manager) + EmbyProvider::with_client_manager( + instance_manager, + Arc::clone(&default_client_manager), + ) }; Ok(Arc::new(provider)) }), @@ -369,6 +388,11 @@ impl ProvidersManager { pub fn list_types(&self) -> Vec { self.factories.keys().cloned().collect() } + + #[cfg(test)] + fn default_client_manager_marker(&self) -> usize { + self.default_client_manager.marker() + } } impl std::fmt::Debug for ProvidersManager { @@ -549,6 +573,78 @@ mod tests { assert!(manager.has_factory("emby")); } + #[tokio::test] + async fn test_default_provider_instances_use_injected_default_client_manager() { + let pool = PgPool::connect_lazy("postgresql://test").unwrap(); + let repo = Arc::new(ProviderInstanceRepository::new(pool)); + let instance_manager = Arc::new(RemoteProviderManager::new_with_invalidation(repo, None)); + let client = synctv_common::http::SsrfSafeClientBuilder::provider() + .connect_timeout(std::time::Duration::from_secs(4)) + .request_timeout(std::time::Duration::from_secs(12)) + .build() + .unwrap(); + + let manager = ProvidersManager::new_with_provider_http_client( + instance_manager, + client, + std::time::Duration::from_secs(4), + ); + let expected_marker = manager.default_client_manager_marker(); + + for provider_type in ["alist", "bilibili", "emby"] { + let provider = manager + .create_provider( + provider_type, + &format!("{provider_type}_default"), + &serde_json::json!({}), + ) + .await + .unwrap(); + + assert_eq!( + provider.test_client_manager_marker(), + Some(expected_marker), + "default {provider_type} provider should reuse the injected default client manager", + ); + } + } + + #[tokio::test] + async fn test_per_instance_timeout_override_keeps_dedicated_client_manager() { + let pool = PgPool::connect_lazy("postgresql://test").unwrap(); + let repo = Arc::new(ProviderInstanceRepository::new(pool)); + let instance_manager = Arc::new(RemoteProviderManager::new_with_invalidation(repo, None)); + let client = synctv_common::http::SsrfSafeClientBuilder::provider() + .connect_timeout(std::time::Duration::from_secs(4)) + .request_timeout(std::time::Duration::from_secs(12)) + .build() + .unwrap(); + + let manager = ProvidersManager::new_with_provider_http_client( + instance_manager, + client, + std::time::Duration::from_secs(4), + ); + let default_marker = manager.default_client_manager_marker(); + + let provider = manager + .create_provider( + "alist", + "alist_override", + &serde_json::json!({"timeout_seconds": 30}), + ) + .await + .unwrap(); + + let actual_marker = provider + .test_client_manager_marker() + .expect("test provider should expose its client manager marker"); + assert_ne!( + actual_marker, default_marker, + "per-instance timeout overrides should build a dedicated client manager" + ); + } + #[tokio::test] async fn test_rtmp_provider_no_longer_requires_base_url() { let pool = PgPool::connect_lazy("postgresql://test").unwrap(); diff --git a/synctv-core/src/service/remote_provider_manager.rs b/synctv-core/src/service/remote_provider_manager.rs index 346319f8..1d0de458 100644 --- a/synctv-core/src/service/remote_provider_manager.rs +++ b/synctv-core/src/service/remote_provider_manager.rs @@ -29,8 +29,8 @@ use synctv_media_providers::grpc::{ emby::{emby_client::EmbyClient, MeReq as EmbyMeReq}, }; use tokio::task::JoinHandle; -use tonic::{Request, Status}; use tonic::transport::{Certificate, Channel, ClientTlsConfig, Endpoint, Uri}; +use tonic::{Request, Status}; use tonic_health::pb::{health_client::HealthClient, HealthCheckRequest}; /// Default channel cache TTL (5 minutes) @@ -39,9 +39,6 @@ const CHANNEL_CACHE_TTL_SECS: u64 = 300; /// Maximum number of cached channels const MAX_CACHED_CHANNELS: u64 = 1_000; -/// Maximum time to wait for the first durable invalidation subscription to become active. -const INVALIDATION_LISTENER_READY_TIMEOUT: Duration = Duration::from_secs(5); - #[derive(Clone, Copy, Debug, Eq, PartialEq)] enum RemoteConfigValidationMode { RequireAuthSecret, @@ -171,25 +168,23 @@ impl RemoteProviderManager { } match self.create_grpc_channel(&config).await { - Ok(channel) => { - match Self::build_remote_connection(&config, channel) { - Ok(connection) => { - self.channel_cache - .insert(config.name.clone(), connection) - .await; - tracing::info!("Pre-warmed provider instance cache: {}", config.name); - success_count += 1; - } - Err(e) => { - tracing::error!( - "Failed to pre-warm provider instance {}: {}", - config.name, - e - ); - error_count += 1; - } + Ok(channel) => match Self::build_remote_connection(&config, channel) { + Ok(connection) => { + self.channel_cache + .insert(config.name.clone(), connection) + .await; + tracing::info!("Pre-warmed provider instance cache: {}", config.name); + success_count += 1; } - } + Err(e) => { + tracing::error!( + "Failed to pre-warm provider instance {}: {}", + config.name, + e + ); + error_count += 1; + } + }, Err(e) => { tracing::error!( "Failed to pre-warm provider instance {}: {}", @@ -270,9 +265,6 @@ impl RemoteProviderManager { *guard = Some(handle); drop(guard); - tokio::time::sleep(INVALIDATION_LISTENER_READY_TIMEOUT.min(Duration::from_millis(10))) - .await; - tracing::info!("Provider instance cache invalidation listener started (durable stream)"); Ok(()) } @@ -415,10 +407,89 @@ impl RemoteProviderManager { ) -> crate::Result { let auth_secret = validate_auth_secret(Some(Self::required_auth_secret(config)?)) .map_err(|e| crate::Error::InvalidInput(e.to_string()))?; - Ok(RemoteProviderConnection::new( - channel, - auth_secret, - )) + Ok(RemoteProviderConnection::new(channel, auth_secret)) + } + + async fn build_validated_remote_connection( + &self, + config: &ProviderInstance, + ) -> crate::Result { + let channel = self.create_grpc_channel(config).await?; + let connection = Self::build_remote_connection(config, channel)?; + self.validate_remote_connection(config, &connection).await?; + Ok(connection) + } + + async fn validate_remote_connection( + &self, + config: &ProviderInstance, + connection: &RemoteProviderConnection, + ) -> crate::Result<()> { + let mut client = HealthClient::new(connection.channel()); + let request = tonic::Request::new(HealthCheckRequest { + service: String::new(), + }); + let timeout = Duration::from_secs(5); + + let response = tokio::time::timeout(timeout, client.check(request)) + .await + .map_err(|_| { + crate::Error::InvalidInput(format!( + "Remote provider instance '{}' connectivity validation timed out after {}s", + config.name, + timeout.as_secs() + )) + })? + .map_err(|status| { + crate::Error::InvalidInput(format!( + "Remote provider instance '{}' health check failed: {status}", + config.name + )) + })?; + + let status = response.into_inner().status; + if status != 1 { + return Err(crate::Error::InvalidInput(format!( + "Remote provider instance '{}' is not serving (health status: {status})", + config.name + ))); + } + + self.validate_authenticated_provider_health(config, connection) + .await + } + + async fn validate_authenticated_provider_health( + &self, + config: &ProviderInstance, + connection: &RemoteProviderConnection, + ) -> crate::Result<()> { + let timeout = Duration::from_secs(5); + let probe = async { + if config + .providers + .iter() + .any(|provider| provider == "bilibili") + { + Self::probe_bilibili_auth(connection).await?; + } + if config.providers.iter().any(|provider| provider == "alist") { + Self::probe_alist_auth(connection).await?; + } + if config.providers.iter().any(|provider| provider == "emby") { + Self::probe_emby_auth(connection).await?; + } + + Ok(()) + }; + + tokio::time::timeout(timeout, probe).await.map_err(|_| { + crate::Error::InvalidInput(format!( + "Authenticated provider probe timed out for instance '{}' after {}s", + config.name, + timeout.as_secs() + )) + })? } fn resolve_ssrf_validated_address( @@ -812,10 +883,7 @@ impl RemoteProviderManager { Self::validate_config(&config, RemoteConfigValidationMode::RequireAuthSecret)?; let connection = if config.enabled && Self::requires_remote_connection(&config) { - Some(Self::build_remote_connection( - &config, - self.create_grpc_channel(&config).await?, - )?) + Some(self.build_validated_remote_connection(&config).await?) } else { None }; @@ -848,7 +916,9 @@ impl RemoteProviderManager { self.repository .get_by_name(&config.name) .await? - .ok_or_else(|| crate::Error::NotFound(format!("Instance '{}' not found", config.name)))?; + .ok_or_else(|| { + crate::Error::NotFound(format!("Instance '{}' not found", config.name)) + })?; let validation_mode = if Self::requires_remote_connection(&config) { RemoteConfigValidationMode::RequireAuthSecret @@ -859,10 +929,7 @@ impl RemoteProviderManager { Self::validate_config(&config, validation_mode)?; let connection = if config.enabled && Self::requires_remote_connection(&config) { - Some(Self::build_remote_connection( - &config, - self.create_grpc_channel(&config).await?, - )?) + Some(self.build_validated_remote_connection(&config).await?) } else { None }; @@ -916,17 +983,10 @@ impl RemoteProviderManager { if config.enabled { if Self::requires_remote_connection(&config) { - if let Some(connection) = self.get(&config.name).await { - self.channel_cache - .insert(config.name.clone(), connection) - .await; - } else { - let channel = self.create_grpc_channel(&config).await?; - let connection = Self::build_remote_connection(&config, channel)?; - self.channel_cache - .insert(config.name.clone(), connection) - .await; - } + let connection = self.build_validated_remote_connection(&config).await?; + self.channel_cache + .insert(config.name.clone(), connection) + .await; } else { self.channel_cache.invalidate(&config.name).await; } @@ -939,8 +999,7 @@ impl RemoteProviderManager { config.enabled = true; if Self::requires_remote_connection(&config) { - let channel = self.create_grpc_channel(&config).await?; - let connection = Self::build_remote_connection(&config, channel)?; + let connection = self.build_validated_remote_connection(&config).await?; // Persist only after a valid channel can be constructed. self.repository.enable(name).await?; @@ -1004,8 +1063,7 @@ impl RemoteProviderManager { ))); } - let channel = self.create_grpc_channel(&config).await?; - let connection = Self::build_remote_connection(&config, channel)?; + let connection = self.build_validated_remote_connection(&config).await?; self.channel_cache .insert(config.name.clone(), connection) .await; @@ -1076,160 +1134,83 @@ impl RemoteProviderManager { config: &ProviderInstance, connection: &RemoteProviderConnection, ) -> bool { - let mut client = HealthClient::new(connection.channel()); - - let request = tonic::Request::new(HealthCheckRequest { - service: String::new(), - }); - - // Set timeout for health check (5 seconds) - let timeout = Duration::from_secs(5); - - match tokio::time::timeout(timeout, client.check(request)).await { - Ok(Ok(response)) => { - let status = response.into_inner().status; - let is_serving = status == 1; - - if is_serving { - let auth_ok = self.check_authenticated_provider_health(config, connection).await; - if auth_ok { - tracing::debug!("Provider instance '{}' is healthy", name); - } - auth_ok - } else { - tracing::warn!( - "Provider instance '{}' is not serving (status: {})", - name, - status - ); - false - } - } - Ok(Err(e)) => { - tracing::error!("Health check failed for instance '{}': {}", name, e); - false - } - Err(_) => { - tracing::error!("Health check timeout for instance '{}' (5s)", name); - false - } - } - } - - async fn check_authenticated_provider_health( - &self, - config: &ProviderInstance, - connection: &RemoteProviderConnection, - ) -> bool { - let timeout = Duration::from_secs(5); - let probe = async { - if config.providers.iter().any(|provider| provider == "alist") { - return Self::probe_alist_auth(config, connection).await; - } - if config.providers.iter().any(|provider| provider == "emby") { - return Self::probe_emby_auth(config, connection).await; - } - if config.providers.iter().any(|provider| provider == "bilibili") { - return Self::probe_bilibili_auth(connection).await; + match self.validate_remote_connection(config, connection).await { + Ok(()) => { + tracing::debug!("Provider instance '{}' is healthy", name); + true } - false - }; - - match tokio::time::timeout(timeout, probe).await { - Ok(result) => result, - Err(_) => { - tracing::error!( - "Authenticated provider health probe timeout for instance '{}' (5s)", - config.name - ); + Err(error) => { + tracing::error!("Health check failed for instance '{}': {}", name, error); false } } } - async fn probe_alist_auth( - config: &ProviderInstance, - connection: &RemoteProviderConnection, - ) -> bool { - let mut client = AlistClient::new(connection.channel()); + async fn probe_bilibili_auth(connection: &RemoteProviderConnection) -> crate::Result<()> { + let mut client = BilibiliClient::new(connection.channel()); let request = match Self::build_authenticated_request( connection, - AlistMeReq { - host: config.endpoint.clone(), - token: "health-check-token".to_string(), + UserInfoReq { + cookies: HashMap::from([("SESSDATA".to_string(), "health-check".to_string())]), }, ) { Ok(request) => request, Err(error) => { - tracing::warn!( - "Authenticated Alist health probe request build failed for '{}': {}", - config.name, + return Err(crate::Error::InvalidInput(format!( + "Authenticated Bilibili probe request build failed: {}", error - ); - return false; + ))); } }; Self::probe_reports_authenticated_health( - "alist", - &config.name, - client.me(request).await, + "bilibili", + "", + client.user_info(request).await, ) } - async fn probe_emby_auth( - config: &ProviderInstance, - connection: &RemoteProviderConnection, - ) -> bool { - let mut client = EmbyClient::new(connection.channel()); + async fn probe_alist_auth(connection: &RemoteProviderConnection) -> crate::Result<()> { + let mut client = AlistClient::new(connection.channel()); let request = match Self::build_authenticated_request( connection, - EmbyMeReq { - host: config.endpoint.clone(), + AlistMeReq { + host: "http://health-check.invalid".to_string(), token: "health-check-token".to_string(), - user_id: "health-check-user".to_string(), }, ) { Ok(request) => request, Err(error) => { - tracing::warn!( - "Authenticated Emby health probe request build failed for '{}': {}", - config.name, + return Err(crate::Error::InvalidInput(format!( + "Authenticated Alist probe request build failed: {}", error - ); - return false; + ))); } }; - Self::probe_reports_authenticated_health( - "emby", - &config.name, - client.me(request).await, - ) + Self::probe_reports_authenticated_health("alist", "", client.me(request).await) } - async fn probe_bilibili_auth(connection: &RemoteProviderConnection) -> bool { - let mut client = BilibiliClient::new(connection.channel()); + async fn probe_emby_auth(connection: &RemoteProviderConnection) -> crate::Result<()> { + let mut client = EmbyClient::new(connection.channel()); let request = match Self::build_authenticated_request( connection, - UserInfoReq { - cookies: HashMap::from([( - "SESSDATA".to_string(), - "health-check".to_string(), - )]), + EmbyMeReq { + host: "http://health-check.invalid".to_string(), + token: "health-check-token".to_string(), + user_id: "health-check-user".to_string(), }, ) { Ok(request) => request, Err(error) => { - tracing::warn!( - "Authenticated Bilibili health probe request build failed: {}", + return Err(crate::Error::InvalidInput(format!( + "Authenticated Emby probe request build failed: {}", error - ); - return false; + ))); } }; - Self::probe_reports_authenticated_health("bilibili", "", client.user_info(request).await) + Self::probe_reports_authenticated_health("emby", "", client.me(request).await) } fn build_authenticated_request( @@ -1257,18 +1238,47 @@ impl RemoteProviderManager { provider: &str, instance_name: &str, result: Result, Status>, - ) -> bool { + ) -> crate::Result<()> { match result { - Ok(_) => true, - Err(status) => { - tracing::warn!( - "Authenticated {} health probe failed for '{}': {}", - provider, - instance_name, - status - ); - false - } + Ok(_) => Ok(()), + Err(status) => Err(crate::Error::InvalidInput(format!( + "Authenticated {} probe failed for '{}': {}", + provider, instance_name, status + ))), } } } + +#[cfg(test)] +mod tests { + use super::*; + use crate::cache::CacheInvalidationService; + use crate::repository::ProviderInstanceRepository; + + #[tokio::test(start_paused = true)] + async fn start_invalidation_listener_does_not_wait_for_fake_readiness() { + let pool = sqlx::PgPool::connect_lazy("postgresql://test") + .expect("lazy pool should build without a live database"); + let repository = Arc::new(ProviderInstanceRepository::new(pool)); + let invalidation = CacheInvalidationService::new( + None, + "test-node".to_string(), + "test:provider:invalidate".to_string(), + ); + let manager = RemoteProviderManager::new_with_invalidation(repository, Some(invalidation)); + + let start = tokio::time::Instant::now(); + manager + .start_invalidation_listener() + .await + .expect("listener should start"); + + assert_eq!( + tokio::time::Instant::now().duration_since(start), + Duration::ZERO, + "listener startup should not advance time via a fake readiness sleep" + ); + + manager.shutdown().await; + } +} diff --git a/synctv-core/src/service/user.rs b/synctv-core/src/service/user.rs index c18e63c8..8ca09cff 100644 --- a/synctv-core/src/service/user.rs +++ b/synctv-core/src/service/user.rs @@ -97,6 +97,51 @@ impl UserService { } } + pub(crate) fn oauth2_username_candidates( + &self, + provider_user_id: &str, + username: &str, + ) -> Result<(String, Vec)> { + let sanitized_username = username + .chars() + .filter(|c| c.is_alphanumeric() || *c == '_' || *c == '-') + .collect::() + .trim() + .to_string(); + + let base_username = if sanitized_username.is_empty() { + format!( + "user_{}", + &provider_user_id[..provider_user_id.len().min(20)] + ) + } else { + sanitized_username + }; + + self.validate_username(&base_username)?; + + let max_attempts = 10; + let mut candidates = Vec::with_capacity(max_attempts); + candidates.push(base_username.clone()); + for _ in 1..max_attempts { + let max_base_len = 42; + let base = if base_username.chars().count() > max_base_len { + base_username.chars().take(max_base_len).collect::() + } else { + base_username.clone() + }; + let suffix = nanoid::nanoid!(6); + candidates.push(format!("{base}_{suffix}")); + } + + Ok((base_username, candidates)) + } + + pub(crate) async fn cache_oauth2_username_best_effort(&self, user_id: &UserId, username: &str) { + self.cache_username_best_effort(user_id, username, "create_or_load_by_oauth2") + .await; + } + async fn invalidate_username_cache_best_effort( &self, user_id: &UserId, @@ -1249,73 +1294,26 @@ impl UserService { username: &str, email: Option<&str>, ) -> Result { - // Sanitize OAuth2 username: remove invalid characters and trim - let sanitized_username = username - .chars() - .filter(|c| c.is_alphanumeric() || *c == '_' || *c == '-') - .collect::() - .trim() - .to_string(); - - // If sanitization resulted in empty username, use provider user ID - let base_username = if sanitized_username.is_empty() { - format!( - "user_{}", - &provider_user_id[..provider_user_id.len().min(20)] - ) - } else { - sanitized_username - }; - - // Validate the sanitized username - self.validate_username(&base_username)?; - + let (base_username, candidates) = + self.oauth2_username_candidates(provider_user_id, username)?; // Generate a random password (OAuth2 users don't need password login) let random_password = nanoid::nanoid!(32); - - // Use provided email, or None if not provided let user_email = email.map(std::string::ToString::to_string); - // Hash password let password_hash = hash_password(&random_password).await?; - // Try to create user with the desired username first. If the DB UNIQUE - // constraint rejects it, fall back to random-suffixed variants. Using - // random suffixes (instead of sequential) avoids thundering herd under - // concurrent OAuth2 signups with the same base username. - let max_attempts = 10; - let mut candidates = Vec::with_capacity(max_attempts); - candidates.push(base_username.clone()); - for _ in 1..max_attempts { - // Cap the base to leave room for the suffix within the 50-char limit - let max_base_len = 42; - // Use character count instead of byte length to avoid panics on multi-byte UTF-8 - let base = if base_username.chars().count() > max_base_len { - base_username.chars().take(max_base_len).collect::() - } else { - base_username.clone() - }; - // Random 6-char alphanumeric suffix - let suffix = nanoid::nanoid!(6); - candidates.push(format!("{base}_{suffix}")); - } - for candidate in &candidates { - let user = User::new( + let user = User::new_with_status( candidate.clone(), user_email.clone(), password_hash.clone(), SignupMethod::OAuth2, + crate::models::UserStatus::Active, ); match self.repository.create(&user).await { Ok(created_user) => { - // Populate username cache - self.cache_username_best_effort( - &created_user.id, - candidate, - "create_or_load_by_oauth2", - ) - .await; + self.cache_oauth2_username_best_effort(&created_user.id, candidate) + .await; if candidate == &base_username { tracing::info!( @@ -1343,7 +1341,6 @@ impl UserService { Err(Error::AlreadyExists(ref msg)) if msg.contains("username") || msg.contains("Username") => { - // Username conflict -- try next candidate continue; } Err(e) => return Err(e), @@ -1351,7 +1348,8 @@ impl UserService { } Err(Error::Internal(format!( - "Could not generate a unique username for base '{username}' after {max_attempts} attempts" + "Could not generate a unique username for base '{username}' after {} attempts", + candidates.len() ))) } diff --git a/synctv-core/src/service/ws_ticket.rs b/synctv-core/src/service/ws_ticket.rs index bd4f4e4e..d063dce8 100644 --- a/synctv-core/src/service/ws_ticket.rs +++ b/synctv-core/src/service/ws_ticket.rs @@ -24,6 +24,7 @@ use async_trait::async_trait; use base64::Engine; use rand::RngExt; use serde::{Deserialize, Serialize}; +use std::future::Future; use std::sync::Arc; use tracing::{debug, warn}; @@ -156,6 +157,13 @@ pub struct RedisTicketStore { } impl RedisTicketStore { + async fn run_redis_op(&self, operation: &'static str, future: F) -> Result + where + F: Future>, + { + run_ws_ticket_redis_op(operation, future).await + } + fn normalize_key_prefix(prefix: impl Into) -> String { let key_prefix = prefix.into(); if key_prefix.is_empty() || key_prefix.ends_with(':') { @@ -190,6 +198,16 @@ impl RedisTicketStore { } } +async fn run_ws_ticket_redis_op(operation: &'static str, future: F) -> Result +where + F: Future>, +{ + tokio::time::timeout(crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, future) + .await + .map_err(|_| Error::Timeout(format!("Redis timeout: {operation}")))? + .map_err(|e| Error::Internal(format!("Failed to {operation}: {e}"))) +} + #[async_trait] impl TicketStore for RedisTicketStore { async fn store(&self, ticket: &str, data: &WsTicketData, ttl_secs: u64) -> Result<()> { @@ -201,13 +219,9 @@ impl TicketStore for RedisTicketStore { .map_err(|e| Error::Internal(format!("Failed to serialize ticket data: {e}")))?; let mut conn = self.conn().await; - let _: () = tokio::time::timeout( - crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, - conn.set_ex(&key, json, ttl_secs), - ) - .await - .map_err(|_| Error::Internal("Redis timeout: store ticket".to_string()))? - .map_err(|e| Error::Internal(format!("Failed to store ticket: {e}")))?; + let _: () = self + .run_redis_op("store ticket", conn.set_ex(&key, json, ttl_secs)) + .await?; Ok(()) } @@ -218,13 +232,7 @@ impl TicketStore for RedisTicketStore { let key = self.redis_key(ticket, expected_room_id); let mut conn = self.conn().await; - let json: Option = tokio::time::timeout( - crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, - conn.get(&key), - ) - .await - .map_err(|_| Error::Internal("Redis timeout: load ticket".to_string()))? - .map_err(|e| Error::Internal(format!("Failed to load ticket: {e}")))?; + let json: Option = self.run_redis_op("load ticket", conn.get(&key)).await?; let Some(json) = json else { return Ok(None); @@ -247,10 +255,11 @@ impl TicketStore for RedisTicketStore { let expected_json = serde_json::to_string(expected_ticket) .map_err(|e| Error::Internal(format!("Failed to serialize ticket data: {e}")))?; - let deleted: i64 = tokio::time::timeout( - crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, - redis::Script::new( - r#" + let deleted: i64 = self + .run_redis_op( + "claim ticket", + redis::Script::new( + r#" local value = redis.call("GET", KEYS[1]) if not value then return 0 @@ -261,14 +270,12 @@ impl TicketStore for RedisTicketStore { redis.call("DEL", KEYS[1]) return 1 "#, + ) + .key(&key) + .arg(&expected_json) + .invoke_async(&mut conn), ) - .key(&key) - .arg(&expected_json) - .invoke_async(&mut conn), - ) - .await - .map_err(|_| Error::Internal("Redis timeout: claim ticket".to_string()))? - .map_err(|e| Error::Internal(format!("Failed to claim ticket: {e}")))?; + .await?; Ok(deleted > 0) } @@ -292,13 +299,12 @@ impl TicketStore for RedisTicketStore { "#, ); - let json: Option = tokio::time::timeout( - crate::resilience::timeout::REDIS_OPERATION_TIMEOUT, - lua_script.key(&key).invoke_async(&mut conn), - ) - .await - .map_err(|_| Error::Internal("Redis timeout: validate ticket".to_string()))? - .map_err(|e| Error::Internal(format!("Failed to validate ticket: {e}")))?; + let json: Option = self + .run_redis_op( + "validate ticket", + lua_script.key(&key).invoke_async(&mut conn), + ) + .await?; let Some(json) = json else { return Ok(None); @@ -1132,6 +1138,25 @@ mod tests { assert!(debug_str.contains("memory")); } + #[tokio::test(start_paused = true)] + async fn test_ws_ticket_redis_timeout_maps_to_timeout_error() { + let timeout_future = run_ws_ticket_redis_op("store ticket", async { + std::future::pending::<()>().await; + #[allow(unreachable_code)] + Ok::<(), redis::RedisError>(()) + }); + + tokio::pin!(timeout_future); + tokio::task::yield_now().await; + tokio::time::advance(crate::resilience::timeout::REDIS_OPERATION_TIMEOUT).await; + + let err = timeout_future.await.expect_err("operation should time out"); + assert!(matches!( + err, + Error::Timeout(ref msg) if msg == "Redis timeout: store ticket" + )); + } + #[test] fn test_non_cluster_mode_allows_memory() { let service = WsTicketService::new(None, "synctv:", None); diff --git a/synctv-core/testing/src/postgres.rs b/synctv-core/testing/src/postgres.rs index bbe8535b..de7dd969 100644 --- a/synctv-core/testing/src/postgres.rs +++ b/synctv-core/testing/src/postgres.rs @@ -886,8 +886,14 @@ mod tests { warning.contains("fallback `docker rm -f` succeeded"), "warning should explain that cleanup fell back to force remove: {warning}" ); - assert!(fallback_called, "explicit cleanup failure must try fallback removal"); - assert!(cleaned_up, "successful fallback should mark the container as cleaned up"); + assert!( + fallback_called, + "explicit cleanup failure must try fallback removal" + ); + assert!( + cleaned_up, + "successful fallback should mark the container as cleaned up" + ); } #[test] @@ -907,7 +913,10 @@ mod tests { warning.contains("already removed"), "warning should explain that the container was already gone: {warning}" ); - assert!(cleaned_up, "missing container should still be treated as cleaned up"); + assert!( + cleaned_up, + "missing container should still be treated as cleaned up" + ); } #[test] @@ -945,11 +954,9 @@ mod tests { #[test] fn docker_rm_force_reports_spawn_failure() { - let err = docker_rm_force_with_program( - "synctv-command-that-should-not-exist", - "synctv-pg-test", - ) - .expect_err("spawn failure must surface as an error"); + let err = + docker_rm_force_with_program("synctv-command-that-should-not-exist", "synctv-pg-test") + .expect_err("spawn failure must surface as an error"); assert!( err.contains("failed to spawn `synctv-command-that-should-not-exist`"), diff --git a/synctv-core/testing/src/redis.rs b/synctv-core/testing/src/redis.rs index beecdea8..60c7c109 100644 --- a/synctv-core/testing/src/redis.rs +++ b/synctv-core/testing/src/redis.rs @@ -696,8 +696,14 @@ mod tests { warning.contains("fallback `docker rm -f` succeeded"), "warning should explain that cleanup fell back to force remove: {warning}" ); - assert!(fallback_called, "explicit cleanup failure must try fallback removal"); - assert!(cleaned_up, "successful fallback should mark the container as cleaned up"); + assert!( + fallback_called, + "explicit cleanup failure must try fallback removal" + ); + assert!( + cleaned_up, + "successful fallback should mark the container as cleaned up" + ); } #[test] @@ -717,7 +723,10 @@ mod tests { warning.contains("already removed"), "warning should explain that the container was already gone: {warning}" ); - assert!(cleaned_up, "missing container should still be treated as cleaned up"); + assert!( + cleaned_up, + "missing container should still be treated as cleaned up" + ); } #[test] diff --git a/synctv-core/tests/oauth2_state_store_tests.rs b/synctv-core/tests/oauth2_state_store_tests.rs index 4c49ff82..77aed8b6 100644 --- a/synctv-core/tests/oauth2_state_store_tests.rs +++ b/synctv-core/tests/oauth2_state_store_tests.rs @@ -32,7 +32,7 @@ fn make_state(instance_name: &str) -> OAuth2State { #[ignore = "Requires Docker"] async fn test_redis_oauth_state_store_and_consume() { let (_container, conn) = start_redis().await; - let store = RedisOAuthStateStore::new(Arc::new(RwLock::new(conn))); + let store = RedisOAuthStateStore::new(Arc::new(RwLock::new(conn)), ""); let state = make_state("github"); let ttl = std::time::Duration::from_mins(1); @@ -60,7 +60,7 @@ async fn test_redis_oauth_state_store_and_consume() { #[ignore = "Requires Docker"] async fn test_redis_oauth_state_consume_is_atomic() { let (_container, conn) = start_redis().await; - let store = Arc::new(RedisOAuthStateStore::new(Arc::new(RwLock::new(conn)))); + let store = Arc::new(RedisOAuthStateStore::new(Arc::new(RwLock::new(conn)), "")); let state = make_state("atomic_test"); let ttl = std::time::Duration::from_mins(1); @@ -99,7 +99,7 @@ async fn test_redis_oauth_state_consume_is_atomic() { #[ignore = "Requires Docker"] async fn test_redis_oauth_state_ttl_expiry() { let (_container, conn) = start_redis().await; - let store = RedisOAuthStateStore::new(Arc::new(RwLock::new(conn))); + let store = RedisOAuthStateStore::new(Arc::new(RwLock::new(conn)), ""); let state = make_state("ttl_test"); let ttl = std::time::Duration::from_secs(1); @@ -130,7 +130,7 @@ async fn test_redis_oauth_state_concurrent_with_barrier() { use tokio::sync::Barrier; let (_container, conn) = start_redis().await; - let store = Arc::new(RedisOAuthStateStore::new(Arc::new(RwLock::new(conn)))); + let store = Arc::new(RedisOAuthStateStore::new(Arc::new(RwLock::new(conn)), "")); let state = make_state("barrier_test"); let ttl = std::time::Duration::from_mins(1); @@ -190,7 +190,7 @@ async fn test_redis_oauth_state_concurrent_with_barrier() { #[ignore = "Requires Docker"] async fn test_redis_oauth_state_multiple_tokens_isolated() { let (_container, conn) = start_redis().await; - let store = RedisOAuthStateStore::new(Arc::new(RwLock::new(conn))); + let store = RedisOAuthStateStore::new(Arc::new(RwLock::new(conn)), ""); let ttl = std::time::Duration::from_mins(1); @@ -232,7 +232,7 @@ async fn test_redis_oauth_state_multiple_tokens_isolated() { #[ignore = "Requires Docker"] async fn test_redis_oauth_state_created_at_expiry_check() { let (_container, conn) = start_redis().await; - let store = RedisOAuthStateStore::new(Arc::new(RwLock::new(conn))); + let store = RedisOAuthStateStore::new(Arc::new(RwLock::new(conn)), ""); // Create a state that's already expired (6 minutes ago, exceeding 5-minute TTL) let expired_time = chrono::Utc::now() - chrono::Duration::seconds(360); @@ -258,3 +258,38 @@ async fn test_redis_oauth_state_created_at_expiry_check() { // But the service layer (consume_state) should reject it based on created_at // Note: This test validates the store layer; service layer expiry is tested elsewhere } + +#[tokio::test] +#[ignore = "Requires Docker"] +async fn test_redis_oauth_state_store_uses_configured_key_prefix() { + let (_container, conn) = start_redis().await; + let shared_conn = Arc::new(RwLock::new(conn)); + let store = RedisOAuthStateStore::new(shared_conn.clone(), "tenant-a"); + + let state = make_state("prefixed"); + store + .store("prefixed_token", &state, std::time::Duration::from_mins(1)) + .await + .expect("storing state should succeed"); + + let mut raw_conn = shared_conn.read().await.clone(); + use redis::AsyncCommands; + + let prefixed_exists: bool = raw_conn + .exists("tenant-a:oauth2:state:prefixed_token") + .await + .expect("prefixed key existence check should succeed"); + assert!( + prefixed_exists, + "state must be stored under configured key prefix" + ); + + let unprefixed_exists: bool = raw_conn + .exists("oauth2:state:prefixed_token") + .await + .expect("unprefixed key existence check should succeed"); + assert!( + !unprefixed_exists, + "state must not leak into the global unprefixed namespace" + ); +} diff --git a/synctv-core/tests/remote_provider_manager_tests.rs b/synctv-core/tests/remote_provider_manager_tests.rs index 06d134f8..cffbdbdf 100644 --- a/synctv-core/tests/remote_provider_manager_tests.rs +++ b/synctv-core/tests/remote_provider_manager_tests.rs @@ -14,15 +14,17 @@ use std::collections::HashMap; use std::net::SocketAddr; use std::sync::Arc; use std::time::Duration; -use synctv_media_providers::grpc::{ - alist::alist_server::AlistServer, alist_server::AlistService as AlistGrpcService, -}; use synctv_core::{ - cache::CacheInvalidationService, models::ProviderInstance, + cache::CacheInvalidationService, + models::ProviderInstance, repository::ProviderInstanceRepository, service::{remote_provider_manager::RemoteProviderManager, CredentialEncryption}, }; use synctv_core_testing::{create_test_pool_with_options_and_label, start_redis_with_client}; +use synctv_media_providers::grpc::{ + alist::alist_server::AlistServer, alist_server::AlistService as AlistGrpcService, + emby::emby_server::EmbyServer, emby_server::EmbyService as EmbyGrpcService, +}; use tokio::sync::{Barrier, RwLock}; use tonic::transport::Server; use tonic_health::ServingStatus; @@ -118,6 +120,20 @@ fn make_test_instance(name: &str) -> ProviderInstance { } } +fn make_reachable_remote_instance(name: &str, host: &str, port: u16) -> ProviderInstance { + let mut instance = make_test_instance(name); + instance.endpoint = format!("http://{host}:{port}"); + instance.providers = vec!["alist".to_string()]; + instance +} + +fn make_test_address_overrides(host: &str, port: u16) -> HashMap { + HashMap::from([( + host.to_string(), + SocketAddr::from((std::net::Ipv4Addr::LOCALHOST, port)), + )]) +} + async fn spawn_authenticated_provider_server( auth_secret: &str, ) -> (SocketAddr, tokio::task::JoinHandle<()>) { @@ -140,7 +156,9 @@ async fn spawn_authenticated_provider_server( let handle = tokio::spawn(async move { Server::builder() .add_service(health_service) - .add_service(AlistServer::new(GrpcAuthProbeAlistService::new(auth_secret))) + .add_service(AlistServer::new(GrpcAuthProbeAlistService::new( + auth_secret, + ))) .serve_with_incoming(tokio_stream::wrappers::TcpListenerStream::new(listener)) .await .expect("provider auth test server should run"); @@ -153,6 +171,7 @@ async fn spawn_authenticated_provider_server( struct GrpcAuthProbeAlistService { expected_secret: Arc, fail_after_auth: bool, + require_real_upstream_token: bool, } impl GrpcAuthProbeAlistService { @@ -160,6 +179,7 @@ impl GrpcAuthProbeAlistService { Self { expected_secret: Arc::::from(expected_secret), fail_after_auth: false, + require_real_upstream_token: false, } } @@ -167,6 +187,15 @@ impl GrpcAuthProbeAlistService { Self { expected_secret: Arc::::from(expected_secret), fail_after_auth: true, + require_real_upstream_token: false, + } + } + + fn requiring_real_upstream_token(expected_secret: String) -> Self { + Self { + expected_secret: Arc::::from(expected_secret), + fail_after_auth: false, + require_real_upstream_token: true, } } @@ -192,7 +221,9 @@ impl synctv_media_providers::grpc::alist::alist_server::Alist for GrpcAuthProbeA _request: tonic::Request, ) -> Result, tonic::Status> { - Err(tonic::Status::unimplemented("login not needed for health probe")) + Err(tonic::Status::unimplemented( + "login not needed for health probe", + )) } async fn fs_get( @@ -200,7 +231,9 @@ impl synctv_media_providers::grpc::alist::alist_server::Alist for GrpcAuthProbeA _request: tonic::Request, ) -> Result, tonic::Status> { - Err(tonic::Status::unimplemented("fs_get not needed for health probe")) + Err(tonic::Status::unimplemented( + "fs_get not needed for health probe", + )) } async fn fs_list( @@ -208,7 +241,9 @@ impl synctv_media_providers::grpc::alist::alist_server::Alist for GrpcAuthProbeA _request: tonic::Request, ) -> Result, tonic::Status> { - Err(tonic::Status::unimplemented("fs_list not needed for health probe")) + Err(tonic::Status::unimplemented( + "fs_list not needed for health probe", + )) } async fn fs_other( @@ -216,25 +251,31 @@ impl synctv_media_providers::grpc::alist::alist_server::Alist for GrpcAuthProbeA _request: tonic::Request, ) -> Result, tonic::Status> { - Err(tonic::Status::unimplemented("fs_other not needed for health probe")) + Err(tonic::Status::unimplemented( + "fs_other not needed for health probe", + )) } async fn fs_search( &self, _request: tonic::Request, - ) -> Result< - tonic::Response, - tonic::Status, - > { - Err(tonic::Status::unimplemented("fs_search not needed for health probe")) + ) -> Result, tonic::Status> + { + Err(tonic::Status::unimplemented( + "fs_search not needed for health probe", + )) } async fn me( &self, request: tonic::Request, - ) -> Result, tonic::Status> - { + ) -> Result, tonic::Status> { self.validate_secret(&request)?; + if self.require_real_upstream_token && request.get_ref().token == "health-check-token" { + return Err(tonic::Status::unauthenticated( + "upstream provider rejected placeholder token", + )); + } if self.fail_after_auth { return Err(tonic::Status::internal( "authenticated provider handler failure", @@ -255,6 +296,161 @@ impl synctv_media_providers::grpc::alist::alist_server::Alist for GrpcAuthProbeA } } +#[derive(Clone)] +struct GrpcAuthProbeEmbyService { + expected_secret: Arc, + fail_after_auth: bool, +} + +impl GrpcAuthProbeEmbyService { + fn failing_after_auth(expected_secret: String) -> Self { + Self { + expected_secret: Arc::::from(expected_secret), + fail_after_auth: true, + } + } + + fn validate_secret(&self, request: &tonic::Request) -> Result<(), tonic::Status> { + let value = request + .metadata() + .get("x-provider-secret") + .ok_or_else(|| tonic::Status::unauthenticated("Missing x-provider-secret header"))?; + let provided = value + .to_str() + .map_err(|_| tonic::Status::unauthenticated("Invalid x-provider-secret header"))?; + if provided != self.expected_secret.as_ref() { + return Err(tonic::Status::unauthenticated("Invalid provider secret")); + } + Ok(()) + } +} + +#[tonic::async_trait] +impl synctv_media_providers::grpc::emby::emby_server::Emby for GrpcAuthProbeEmbyService { + async fn login( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented( + "login not needed for health probe", + )) + } + + async fn me( + &self, + request: tonic::Request, + ) -> Result, tonic::Status> { + self.validate_secret(&request)?; + if self.fail_after_auth { + return Err(tonic::Status::internal( + "authenticated emby provider handler failure", + )); + } + Ok(tonic::Response::new( + synctv_media_providers::grpc::emby::MeResp { + id: "health-check-user".to_string(), + name: "health-check".to_string(), + server_id: "health-check-server".to_string(), + policy: None, + }, + )) + } + + async fn get_items( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> + { + Err(tonic::Status::unimplemented( + "get_items not needed for health probe", + )) + } + + async fn get_item( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented( + "get_item not needed for health probe", + )) + } + + async fn get_system_info( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> + { + Err(tonic::Status::unimplemented( + "get_system_info not needed for health probe", + )) + } + + async fn fs_list( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> + { + Err(tonic::Status::unimplemented( + "fs_list not needed for health probe", + )) + } + + async fn logout( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented( + "logout not needed for health probe", + )) + } + + async fn playback_info( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> + { + Err(tonic::Status::unimplemented( + "playback_info not needed for health probe", + )) + } + + async fn delete_active_encodings( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented( + "delete_active_encodings not needed for health probe", + )) + } + + async fn report_playback_start( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented( + "report_playback_start not needed for health probe", + )) + } + + async fn report_playback_stop( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented( + "report_playback_stop not needed for health probe", + )) + } + + async fn report_playback_progress( + &self, + _request: tonic::Request, + ) -> Result, tonic::Status> { + Err(tonic::Status::unimplemented( + "report_playback_progress not needed for health probe", + )) + } +} + async fn spawn_authenticated_provider_server_with_handler_failure( auth_secret: &str, ) -> (SocketAddr, tokio::task::JoinHandle<()>) { @@ -288,6 +484,94 @@ async fn spawn_authenticated_provider_server_with_handler_failure( (addr, handle) } +async fn spawn_authenticated_provider_server_rejecting_placeholder_upstream_auth( + auth_secret: &str, +) -> (SocketAddr, tokio::task::JoinHandle<()>) { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("provider auth placeholder rejection server should bind"); + let addr = listener + .local_addr() + .expect("provider auth placeholder rejection server should expose local address"); + + let (reporter, health_service) = tonic_health::server::health_reporter(); + reporter + .set_service_status("", ServingStatus::Serving) + .await; + reporter + .set_serving::>() + .await; + + let auth_secret = auth_secret.to_string(); + let handle = tokio::spawn(async move { + Server::builder() + .add_service(health_service) + .add_service(AlistServer::new( + GrpcAuthProbeAlistService::requiring_real_upstream_token(auth_secret), + )) + .serve_with_incoming(tokio_stream::wrappers::TcpListenerStream::new(listener)) + .await + .expect("provider auth placeholder rejection server should run"); + }); + + (addr, handle) +} + +async fn spawn_authenticated_emby_provider_server_with_handler_failure( + auth_secret: &str, +) -> (SocketAddr, tokio::task::JoinHandle<()>) { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("emby auth failure test server should bind to an ephemeral port"); + let addr = listener + .local_addr() + .expect("emby auth failure test server should expose a local address"); + + let (reporter, health_service) = tonic_health::server::health_reporter(); + reporter + .set_service_status("", ServingStatus::Serving) + .await; + reporter.set_serving::>().await; + + let auth_secret = auth_secret.to_string(); + let handle = tokio::spawn(async move { + Server::builder() + .add_service(health_service) + .add_service(EmbyServer::new( + GrpcAuthProbeEmbyService::failing_after_auth(auth_secret), + )) + .serve_with_incoming(tokio_stream::wrappers::TcpListenerStream::new(listener)) + .await + .expect("emby auth failure test server should run"); + }); + + (addr, handle) +} + +async fn spawn_stalling_tcp_server() -> (SocketAddr, tokio::task::JoinHandle<()>) { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("stalling test server should bind to an ephemeral port"); + let addr = listener + .local_addr() + .expect("stalling test server should expose a local address"); + + let handle = tokio::spawn(async move { + loop { + let (stream, _) = listener + .accept() + .await + .expect("stalling test server should accept connections"); + tokio::spawn(async move { + let _stream = stream; + std::future::pending::<()>().await; + }); + } + }); + + (addr, handle) +} + /// Create a test provider instance with TLS fn make_test_instance_tls(name: &str, insecure: bool) -> ProviderInstance { let now = Utc::now(); @@ -312,28 +596,27 @@ fn make_test_instance_tls(name: &str, insecure: bool) -> ProviderInstance { async fn scenario_channel_creation_from_db_config() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "channel-create.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); - // Create instance in DB - let instance = make_test_instance("test-instance-1"); + // Create instance in DB through the validated management path. + let instance = make_reachable_remote_instance("test-instance-1", host, health_addr.port()); manager.add(instance.clone()).await.unwrap(); - // Get channel - should create from DB config - // Note: The channel creation will attempt to connect. Even though there's - // no actual gRPC server, tonic creates lazy channels that don't connect - // until the first RPC call. So we expect Some(channel) here. + // Get channel - should load the validated/cached remote connection. let channel = manager.get("test-instance-1").await; - // Channel should be Some (lazy channel created, even if it will fail later) assert!( channel.is_some(), - "Channel should be created (lazy connection)" + "validated remote channel should be available" ); // Verify instance exists in DB @@ -341,6 +624,9 @@ async fn scenario_channel_creation_from_db_config() { let fetched = repo.get_by_name("test-instance-1").await.unwrap(); assert!(fetched.is_some()); assert_eq!(fetched.unwrap().name, "test-instance-1"); + + health_handle.abort(); + let _ = health_handle.await; } // ─── Test 2: Channel cache hit (cached channel returned) ───────────────────── @@ -348,16 +634,19 @@ async fn scenario_channel_creation_from_db_config() { async fn scenario_channel_cache_hit() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "cache-hit.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); - // Create instance in DB - let instance = make_test_instance("test-instance-2"); + // Create reachable instance in DB. + let instance = make_reachable_remote_instance("test-instance-2", host, health_addr.port()); manager.add(instance.clone()).await.unwrap(); // First get - cache miss, attempts DB lookup @@ -369,6 +658,9 @@ async fn scenario_channel_cache_hit() { // Verify DB was only queried once (cache working) // This is implicit - if cache wasn't working, we'd see multiple DB queries // in logs. For now, we just verify no panics occur. + + health_handle.abort(); + let _ = health_handle.await; } // ─── Test 3: Channel cache TTL expiration ─────────────────────────────────── @@ -376,17 +668,19 @@ async fn scenario_channel_cache_hit() { async fn scenario_channel_cache_ttl_expiration() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "cache-ttl.test.localhost"; - // Create a manager with a very short TTL for testing + // Create a manager and back it with a reachable remote instance. let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); - // Create instance in DB - let instance = make_test_instance("test-instance-3"); + let instance = make_reachable_remote_instance("test-instance-3", host, health_addr.port()); manager.add(instance.clone()).await.unwrap(); // First get - populates cache @@ -405,6 +699,9 @@ async fn scenario_channel_cache_ttl_expiration() { // 2. Setting a very short TTL (e.g., 100ms) // 3. Waiting and verifying cache miss // This is left as an exercise for future enhancement + + health_handle.abort(); + let _ = health_handle.await; } // ─── Test 4: Redis invalidation on delete ─────────────────────────────────── @@ -412,6 +709,9 @@ async fn scenario_channel_cache_ttl_expiration() { async fn scenario_redis_invalidation_on_delete() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "redis-delete.test.localhost"; let repo = provider_repo(&infra.pool); let stream_key = format!("test:provider:invalidate:{}", nanoid::nanoid!(8)); let invalidation1 = CacheInvalidationService::new( @@ -433,18 +733,23 @@ async fn scenario_redis_invalidation_on_delete() { invalidation1.start().await.unwrap(); invalidation2.start().await.unwrap(); - let manager1 = RemoteProviderManager::new_with_invalidation( + let address_overrides = make_test_address_overrides(host, health_addr.port()); + let manager1 = RemoteProviderManager::new_with_test_address_overrides( Arc::new(provider_repo(&infra.pool)), Some(invalidation1.clone()), + address_overrides.clone(), + ); + let manager2 = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + Some(invalidation2.clone()), + address_overrides, ); - let manager2 = - RemoteProviderManager::new_with_invalidation(Arc::new(repo), Some(invalidation2.clone())); // Start invalidation listener manager2.start_invalidation_listener().await.unwrap(); // Create instance via manager1 - let instance = make_test_instance("test-instance-5"); + let instance = make_reachable_remote_instance("test-instance-5", host, health_addr.port()); manager1.add(instance.clone()).await.unwrap(); // Pre-warm manager2's cache @@ -479,6 +784,8 @@ async fn scenario_redis_invalidation_on_delete() { manager2.shutdown().await; invalidation1.stop().await; invalidation2.stop().await; + health_handle.abort(); + let _ = health_handle.await; } // ─── Test 5: Health check integration ─────────────────────────────────────── @@ -528,16 +835,20 @@ async fn scenario_health_check_integration() { async fn scenario_health_check_respects_enabled_flag() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "health-enabled.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); // Create enabled instance - let instance_enabled = make_test_instance("test-instance-7a"); + let instance_enabled = + make_reachable_remote_instance("test-instance-7a", host, health_addr.port()); manager.add(instance_enabled).await.unwrap(); // Create disabled instance @@ -560,6 +871,9 @@ async fn scenario_health_check_respects_enabled_flag() { !health_results.contains_key("test-instance-7b"), "Health check should skip disabled instance" ); + + health_handle.abort(); + let _ = health_handle.await; } async fn scenario_health_check_reports_enabled_instance_with_invalid_secret_as_unhealthy() { @@ -676,10 +990,8 @@ async fn scenario_health_check_reports_authenticated_provider_failure_as_unhealt let infra = TestInfra::new().await; flush_provider_instances(&infra).await; let (health_addr, health_handle) = - spawn_authenticated_provider_server_with_handler_failure( - "remote-provider-test-secret", - ) - .await; + spawn_authenticated_provider_server_with_handler_failure("remote-provider-test-secret") + .await; let repo = provider_repo(&infra.pool); let manager = RemoteProviderManager::new_with_test_address_overrides( @@ -717,6 +1029,126 @@ async fn scenario_health_check_reports_authenticated_provider_failure_as_unhealt let _ = health_handle.await; } +async fn scenario_add_alist_instance_does_not_require_fake_upstream_auth_for_management_validation() +{ + let infra = TestInfra::new().await; + flush_provider_instances(&infra).await; + let (health_addr, health_handle) = + spawn_authenticated_provider_server_rejecting_placeholder_upstream_auth( + "remote-provider-test-secret", + ) + .await; + let host = "alist-management-validation.test.localhost"; + + let repo = provider_repo(&infra.pool); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); + + let instance = make_reachable_remote_instance( + "test-instance-alist-management-validation", + host, + health_addr.port(), + ); + + manager + .add(instance) + .await + .expect("management validation should not depend on fake upstream Alist credentials"); + + health_handle.abort(); + let _ = health_handle.await; +} + +async fn scenario_health_check_reports_emby_authenticated_provider_failure_as_unhealthy() { + let infra = TestInfra::new().await; + flush_provider_instances(&infra).await; + let (health_addr, health_handle) = + spawn_authenticated_emby_provider_server_with_handler_failure( + "remote-provider-test-secret", + ) + .await; + + let repo = provider_repo(&infra.pool); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + HashMap::from([( + "emby-handler-failure-health.test.localhost".to_string(), + SocketAddr::from((std::net::Ipv4Addr::LOCALHOST, health_addr.port())), + )]), + ); + + let mut broken = make_test_instance("test-instance-7h-emby-handler-failure"); + broken.endpoint = format!( + "http://emby-handler-failure-health.test.localhost:{}", + health_addr.port() + ); + broken.providers = vec!["emby".to_string()]; + + provider_repo(&infra.pool) + .create(&broken) + .await + .expect("emby handler-failure row should persist for health-check coverage"); + + let health_results = manager.health_check().await; + assert!( + health_results.contains_key(&broken.name), + "emby instances with authenticated provider failures should appear in health results" + ); + assert!( + !health_results[&broken.name], + "authenticated emby provider handler failures must be reported unhealthy" + ); + + health_handle.abort(); + let _ = health_handle.await; +} + +async fn scenario_add_emby_instance_rejects_authenticated_handler_failure() { + let infra = TestInfra::new().await; + flush_provider_instances(&infra).await; + let (health_addr, health_handle) = + spawn_authenticated_emby_provider_server_with_handler_failure( + "remote-provider-test-secret", + ) + .await; + let host = "emby-management-validation.test.localhost"; + + let repo = provider_repo(&infra.pool); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); + + let mut instance = make_reachable_remote_instance( + "test-instance-emby-management-validation", + host, + health_addr.port(), + ); + instance.providers = vec!["emby".to_string()]; + + let result = manager.add(instance.clone()).await; + assert!( + result.is_err(), + "management validation must reject emby instances when authenticated RPCs fail" + ); + assert!( + provider_repo(&infra.pool) + .get_by_name(&instance.name) + .await + .expect("lookup should succeed") + .is_none(), + "failed add must not persist an emby instance with broken authenticated handlers" + ); + + health_handle.abort(); + let _ = health_handle.await; +} + // ─── Test 8: TLS configuration (non-insecure) ─────────────────────────────── async fn scenario_tls_configuration_secure() { @@ -922,16 +1354,19 @@ async fn scenario_fallback_when_channel_creation_fails() { async fn scenario_enable_disable_instance() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "enable-disable.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); // Create enabled instance - let instance = make_test_instance("test-instance-13"); + let instance = make_reachable_remote_instance("test-instance-13", host, health_addr.port()); manager.add(instance.clone()).await.unwrap(); // Verify it's enabled and gettable @@ -959,6 +1394,9 @@ async fn scenario_enable_disable_instance() { let fetched = repo.get_by_name("test-instance-13").await.unwrap(); assert!(fetched.is_some()); assert!(fetched.unwrap().enabled); + + health_handle.abort(); + let _ = health_handle.await; } async fn scenario_enable_with_invalid_endpoint_preserves_disabled_state() { @@ -1017,7 +1455,9 @@ async fn scenario_enable_remote_instance_requires_jwt_secret() { let repo = provider_repo(&infra.pool); repo.create(&instance).await.unwrap(); - let result = manager.enable("test-instance-13-missing-secret-enable").await; + let result = manager + .enable("test-instance-13-missing-secret-enable") + .await; assert!( result.is_err(), "enabling a remote instance without jwt_secret must fail" @@ -1055,7 +1495,9 @@ async fn scenario_enable_already_enabled_legacy_remote_instance_without_jwt_secr .await .expect("legacy enabled row should persist"); - let result = manager.enable("test-instance-13-legacy-already-enabled").await; + let result = manager + .enable("test-instance-13-legacy-already-enabled") + .await; assert!( result.is_err(), "re-enabling an already-enabled legacy row without jwt_secret must fail" @@ -1087,27 +1529,25 @@ async fn scenario_enable_already_enabled_legacy_remote_instance_without_jwt_secr async fn scenario_reconnect_instance() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "reconnect.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); // Create instance - let instance = make_test_instance("test-instance-14"); + let instance = make_reachable_remote_instance("test-instance-14", host, health_addr.port()); manager.add(instance.clone()).await.unwrap(); - // Try to reconnect - since tonic creates lazy channels, this will succeed - // even though there's no actual server let result = manager.reconnect("test-instance-14").await; - - // Reconnect will succeed because tonic creates lazy channels - // (connection isn't established until first RPC call) assert!( result.is_ok(), - "Reconnect should succeed with lazy channel (even without server)" + "Reconnect should succeed for a reachable remote instance" ); // Disable the instance @@ -1119,6 +1559,9 @@ async fn scenario_reconnect_instance() { result.is_err(), "Reconnect should fail for disabled instance" ); + + health_handle.abort(); + let _ = health_handle.await; } // ─── Test 15: Add duplicate instance fails ─────────────────────────────────── @@ -1126,16 +1569,19 @@ async fn scenario_reconnect_instance() { async fn scenario_add_duplicate_instance_fails() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "duplicate-add.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); // Create instance - let instance = make_test_instance("test-instance-15"); + let instance = make_reachable_remote_instance("test-instance-15", host, health_addr.port()); manager.add(instance.clone()).await.unwrap(); // Try to add duplicate - should fail @@ -1148,6 +1594,9 @@ async fn scenario_add_duplicate_instance_fails() { "Error should be AlreadyExists variant" ); } + + health_handle.abort(); + let _ = health_handle.await; } async fn scenario_add_disabled_instance_is_not_retrievable_via_get() { @@ -1182,15 +1631,22 @@ async fn scenario_add_disabled_instance_is_not_retrievable_via_get() { async fn scenario_update_to_disabled_invalidates_cached_channel() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "update-disable.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); - let instance = make_test_instance("test-instance-15-update-disabled"); + let instance = make_reachable_remote_instance( + "test-instance-15-update-disabled", + host, + health_addr.port(), + ); manager.add(instance.clone()).await.unwrap(); let initial = manager.get("test-instance-15-update-disabled").await; @@ -1213,24 +1669,26 @@ async fn scenario_update_to_disabled_invalidates_cached_channel() { channel.is_none(), "update(enabled=false) must evict any cached channel" ); + + health_handle.abort(); + let _ = health_handle.await; } async fn scenario_concurrent_duplicate_add_returns_one_success_and_one_already_exists() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "concurrent-dup.test.localhost"; - let manager = Arc::new(RemoteProviderManager::new( + let manager = Arc::new(RemoteProviderManager::new_with_test_address_overrides( Arc::new(provider_repo(&infra.pool)), - redis_conn, - redis_client, - "", + None, + make_test_address_overrides(host, health_addr.port()), )); let barrier = Arc::new(Barrier::new(3)); - let instance = make_test_instance("test-instance-15-concurrent-dup"); + let instance = + make_reachable_remote_instance("test-instance-15-concurrent-dup", host, health_addr.port()); let task1 = { let manager = Arc::clone(&manager); @@ -1285,6 +1743,9 @@ async fn scenario_concurrent_duplicate_add_returns_one_success_and_one_already_e .await .unwrap(); assert_eq!(stored_count, 1, "only one DB row should be persisted"); + + health_handle.abort(); + let _ = health_handle.await; } // ─── Test 16: Update non-existent instance fails ───────────────────────────── @@ -1436,8 +1897,7 @@ async fn scenario_update_existing_remote_instance_requires_jwt_secret() { .expect("lookup should succeed") .expect("instance should still exist"); assert_ne!( - persisted.comment, - instance.comment, + persisted.comment, instance.comment, "failed update must not persist other field changes" ); assert_eq!( @@ -1525,17 +1985,24 @@ async fn scenario_delete_nonexistent_instance_fails() { async fn scenario_get_all_instances() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "get-all.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); // Create multiple instances for i in 1..=3 { - let instance = make_test_instance(&format!("test-instance-18a-{i}")); + let instance = make_reachable_remote_instance( + &format!("test-instance-18a-{i}"), + host, + health_addr.port(), + ); manager.add(instance).await.unwrap(); } @@ -1564,6 +2031,9 @@ async fn scenario_get_all_instances() { disabled_count >= 1, "Should have at least 1 disabled instance" ); + + health_handle.abort(); + let _ = health_handle.await; } // ─── Test 19: Manager without Redis (local-only invalidation) ─────────────── @@ -1571,13 +2041,15 @@ async fn scenario_get_all_instances() { async fn scenario_manager_without_redis() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "manager-no-redis.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new( + let manager = RemoteProviderManager::new_with_test_address_overrides( Arc::new(repo), - None, // No Redis - None, // No Redis client - "", + None, + make_test_address_overrides(host, health_addr.port()), ); // Start invalidation listener - should return Ok without starting @@ -1588,7 +2060,7 @@ async fn scenario_manager_without_redis() { ); // Create instance - let instance = make_test_instance("test-instance-19"); + let instance = make_reachable_remote_instance("test-instance-19", host, health_addr.port()); manager.add(instance.clone()).await.unwrap(); // Get should still work @@ -1600,6 +2072,9 @@ async fn scenario_manager_without_redis() { instances.contains(&"test-instance-19".to_string()), "Should list the instance even without Redis" ); + + health_handle.abort(); + let _ = health_handle.await; } // ─── Test 20: Init pre-warms cache ────────────────────────────────────────── @@ -1607,17 +2082,24 @@ async fn scenario_manager_without_redis() { async fn scenario_init_pre_warms_cache() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "init-prewarm.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); // Create instances before init for i in 1..=3 { - let instance = make_test_instance(&format!("test-instance-20-{i}")); + let instance = make_reachable_remote_instance( + &format!("test-instance-20-{i}"), + host, + health_addr.port(), + ); manager.add(instance).await.unwrap(); } @@ -1631,18 +2113,24 @@ async fn scenario_init_pre_warms_cache() { instances.len() >= 3, "Should list at least 3 instances after init" ); + + health_handle.abort(); + let _ = health_handle.await; } async fn scenario_init_skips_invalid_secret_and_continues_prewarming() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "init-invalid-secret.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); let mut invalid = make_test_instance("test-instance-20-invalid-secret"); invalid.jwt_secret = Some("shared\nsecret".to_string()); @@ -1651,7 +2139,8 @@ async fn scenario_init_skips_invalid_secret_and_continues_prewarming() { .await .expect("invalid legacy row should persist for compatibility coverage"); - let healthy = make_test_instance("test-instance-20-healthy"); + let healthy = + make_reachable_remote_instance("test-instance-20-healthy", host, health_addr.port()); manager .add(healthy.clone()) .await @@ -1676,6 +2165,9 @@ async fn scenario_init_skips_invalid_secret_and_continues_prewarming() { invalid_connection.is_none(), "instance with invalid secret should be skipped instead of poisoning the whole prewarm pass" ); + + health_handle.abort(); + let _ = health_handle.await; } // ─── Test 21: SSRF validation prevents internal endpoints ─────────────────── @@ -1733,17 +2225,28 @@ async fn scenario_ssrf_validation_allows_public_endpoints() { let result = manager.add(instance.clone()).await; - // SSRF validation should pass (connection may fail, but that's different) - // The instance should be added to DB + // Public endpoints should pass static SSRF validation, but the management + // path now also requires real connectivity before persisting the instance. assert!( - result.is_ok(), - "Adding instance with public endpoint should pass SSRF validation" + result.is_err(), + "Adding an unreachable public endpoint must fail connectivity validation" + ); + + let error_message = result + .expect_err("unreachable public endpoint should fail") + .to_string(); + assert!( + !error_message.contains("SSRF validation: host"), + "public host should not be rejected by static SSRF policy: {error_message}" + ); + assert!( + provider_repo(&infra.pool) + .get_by_name("test-instance-22") + .await + .expect("lookup should succeed") + .is_none(), + "failed connectivity validation must not persist the instance" ); - - // Verify it's in the DB - let repo = provider_repo(&infra.pool); - let fetched = repo.get_by_name("test-instance-22").await.unwrap(); - assert!(fetched.is_some()); } // ─── Test 23: resolve_client with remote instance ─────────────────────────── @@ -1751,31 +2254,28 @@ async fn scenario_ssrf_validation_allows_public_endpoints() { async fn scenario_resolve_client_uses_remote_when_available() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "resolve-remote.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = Arc::new(RemoteProviderManager::new( + let manager = Arc::new(RemoteProviderManager::new_with_test_address_overrides( Arc::new(repo), - redis_conn, - redis_client, - "", + None, + make_test_address_overrides(host, health_addr.port()), )); - // Create instance (tonic creates lazy channel) - let instance = make_test_instance("test-instance-23"); + let instance = make_reachable_remote_instance("test-instance-23", host, health_addr.port()); manager.add(instance).await.unwrap(); - // Since tonic creates lazy channels, the remote path will be taken - // even though there's no actual server let result = manager .resolve_client(Some("test-instance-23"), |_channel| "remote", || "local") .await; - // Should be "remote" because the lazy channel was created successfully assert_eq!(result, "remote"); + + health_handle.abort(); + let _ = health_handle.await; } // ─── Test 24: Cache respects max capacity ─────────────────────────────────── @@ -1783,29 +2283,42 @@ async fn scenario_resolve_client_uses_remote_when_available() { async fn scenario_cache_respects_max_capacity() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; - let redis_conn = Some(Arc::new(RwLock::new( - infra.redis_connection_manager().await, - ))); - let redis_client = Some(infra.redis_client.clone()); + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "cache-capacity.test.localhost"; let repo = provider_repo(&infra.pool); - let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, health_addr.port()), + ); // Create instances (default max is 1000, so this won't test eviction) // This is more of a sanity check that the cache doesn't panic for i in 1..=10 { - let instance = make_test_instance(&format!("test-instance-24-{i}")); + let instance = make_reachable_remote_instance( + &format!("test-instance-24-{i}"), + host, + health_addr.port(), + ); manager.add(instance).await.unwrap(); } // All should be listable let instances = manager.list().await.unwrap(); assert!(instances.len() >= 10, "Should list at least 10 instances"); + + health_handle.abort(); + let _ = health_handle.await; } async fn scenario_redis_invalidation_respects_key_prefix() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "redis-prefix.test.localhost"; let stream_key = format!("tenant-a:test:provider:invalidate:{}", nanoid::nanoid!(8)); let invalidation1 = CacheInvalidationService::new( Some(infra.redis_client.clone()), @@ -1826,18 +2339,21 @@ async fn scenario_redis_invalidation_respects_key_prefix() { invalidation1.start().await.unwrap(); invalidation2.start().await.unwrap(); - let manager1 = RemoteProviderManager::new_with_invalidation( + let address_overrides = make_test_address_overrides(host, health_addr.port()); + let manager1 = RemoteProviderManager::new_with_test_address_overrides( Arc::new(provider_repo(&infra.pool)), Some(invalidation1.clone()), + address_overrides.clone(), ); - let manager2 = RemoteProviderManager::new_with_invalidation( + let manager2 = RemoteProviderManager::new_with_test_address_overrides( Arc::new(provider_repo(&infra.pool)), Some(invalidation2.clone()), + address_overrides, ); manager2.start_invalidation_listener().await.unwrap(); - let instance = make_test_instance("test-instance-prefix"); + let instance = make_reachable_remote_instance("test-instance-prefix", host, health_addr.port()); manager1.add(instance).await.unwrap(); let _ = manager2.get("test-instance-prefix").await; tokio::time::sleep(Duration::from_millis(100)).await; @@ -1858,6 +2374,8 @@ async fn scenario_redis_invalidation_respects_key_prefix() { manager2.shutdown().await; invalidation1.stop().await; invalidation2.stop().await; + health_handle.abort(); + let _ = health_handle.await; } async fn scenario_invalidation_listener_shutdown_is_idempotent() { @@ -1896,6 +2414,9 @@ async fn scenario_invalidation_listener_shutdown_is_idempotent() { async fn scenario_durable_invalidation_catches_up_after_listener_starts_late() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + let host = "durable-invalidation.test.localhost"; let stream_key = format!("test:provider:durable:{}", nanoid::nanoid!(8)); let invalidation1 = CacheInvalidationService::new( @@ -1917,16 +2438,20 @@ async fn scenario_durable_invalidation_catches_up_after_listener_starts_late() { invalidation1.start().await.unwrap(); invalidation2.start().await.unwrap(); - let manager1 = RemoteProviderManager::new_with_invalidation( + let address_overrides = make_test_address_overrides(host, health_addr.port()); + let manager1 = RemoteProviderManager::new_with_test_address_overrides( Arc::new(provider_repo(&infra.pool)), Some(invalidation1.clone()), + address_overrides.clone(), ); - let manager2 = RemoteProviderManager::new_with_invalidation( + let manager2 = RemoteProviderManager::new_with_test_address_overrides( Arc::new(provider_repo(&infra.pool)), Some(invalidation2.clone()), + address_overrides, ); - let instance = make_test_instance("durable-provider-instance"); + let instance = + make_reachable_remote_instance("durable-provider-instance", host, health_addr.port()); manager1.add(instance).await.unwrap(); let channel = manager2.get("durable-provider-instance").await; @@ -1956,6 +2481,8 @@ async fn scenario_durable_invalidation_catches_up_after_listener_starts_late() { manager2.shutdown().await; invalidation1.stop().await; invalidation2.stop().await; + health_handle.abort(); + let _ = health_handle.await; } // ─── Test 25: Provider instance supports_provider ─────────────────────────── @@ -2007,7 +2534,10 @@ async fn scenario_add_remote_instance_requires_jwt_secret() { instance.jwt_secret = None; let result = manager.add(instance).await; - assert!(result.is_err(), "remote instance without jwt_secret must be rejected"); + assert!( + result.is_err(), + "remote instance without jwt_secret must be rejected" + ); let error = result.expect_err("missing secret should fail"); let error_message = error.to_string(); @@ -2101,6 +2631,239 @@ async fn scenario_add_local_only_instance_allows_empty_jwt_secret() { assert_eq!(fetched.jwt_secret, None); } +async fn scenario_add_unreachable_remote_instance_fails_connectivity_validation() { + let infra = TestInfra::new().await; + flush_provider_instances(&infra).await; + let redis_conn = Some(Arc::new(RwLock::new( + infra.redis_connection_manager().await, + ))); + let redis_client = Some(infra.redis_client.clone()); + + let repo = provider_repo(&infra.pool); + let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + + let mut instance = make_test_instance("test-instance-unreachable-add"); + instance.endpoint = "http://unreachable-provider.example.invalid:50051".to_string(); + instance.providers = vec!["alist".to_string()]; + instance.timeout = "1s".to_string(); + + let result = manager.add(instance.clone()).await; + assert!( + result.is_err(), + "remote instance add must fail when the configured endpoint is unreachable" + ); + + assert!( + provider_repo(&infra.pool) + .get_by_name(&instance.name) + .await + .expect("lookup should succeed") + .is_none(), + "failed add must not persist an unreachable remote instance" + ); +} + +async fn scenario_add_reachable_remote_instance_succeeds_with_connectivity_validation() { + let infra = TestInfra::new().await; + flush_provider_instances(&infra).await; + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + + let repo = provider_repo(&infra.pool); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + HashMap::from([( + "reachable-provider.test.localhost".to_string(), + SocketAddr::from((std::net::Ipv4Addr::LOCALHOST, health_addr.port())), + )]), + ); + + let mut instance = make_test_instance("test-instance-reachable-add"); + instance.endpoint = format!( + "http://reachable-provider.test.localhost:{}", + health_addr.port() + ); + instance.providers = vec!["alist".to_string()]; + + manager + .add(instance.clone()) + .await + .expect("reachable remote instance should pass connectivity validation"); + + let connection = manager + .get(&instance.name) + .await + .expect("reachable remote instance should be cached"); + assert_eq!( + connection.auth_secret(), + Some("remote-provider-test-secret"), + "validated remote instance should retain its auth secret" + ); + + let stored = provider_repo(&infra.pool) + .get_by_name(&instance.name) + .await + .expect("lookup should succeed") + .expect("reachable instance should be persisted"); + assert_eq!(stored.endpoint, instance.endpoint); + + health_handle.abort(); + let _ = health_handle.await; +} + +async fn scenario_enable_unreachable_remote_instance_preserves_disabled_state() { + let infra = TestInfra::new().await; + flush_provider_instances(&infra).await; + let redis_conn = Some(Arc::new(RwLock::new( + infra.redis_connection_manager().await, + ))); + let redis_client = Some(infra.redis_client.clone()); + + let repo = provider_repo(&infra.pool); + let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + + let mut instance = make_test_instance("test-instance-unreachable-enable"); + instance.enabled = false; + instance.endpoint = "http://unreachable-provider.example.invalid:50051".to_string(); + instance.providers = vec!["alist".to_string()]; + instance.timeout = "1s".to_string(); + + provider_repo(&infra.pool) + .create(&instance) + .await + .expect("disabled instance should persist"); + + let result = manager.enable(&instance.name).await; + assert!( + result.is_err(), + "enable must fail when the remote endpoint is unreachable" + ); + + let stored = provider_repo(&infra.pool) + .get_by_name(&instance.name) + .await + .expect("lookup should succeed") + .expect("instance should still exist"); + assert!( + !stored.enabled, + "failed enable must leave the instance disabled in the database" + ); +} + +async fn scenario_reconnect_unreachable_remote_instance_fails_connectivity_validation() { + let infra = TestInfra::new().await; + flush_provider_instances(&infra).await; + let redis_conn = Some(Arc::new(RwLock::new( + infra.redis_connection_manager().await, + ))); + let redis_client = Some(infra.redis_client.clone()); + + let repo = provider_repo(&infra.pool); + let manager = RemoteProviderManager::new(Arc::new(repo), redis_conn, redis_client, ""); + + let mut instance = make_test_instance("test-instance-unreachable-reconnect"); + instance.endpoint = "http://unreachable-provider.example.invalid:50051".to_string(); + instance.providers = vec!["alist".to_string()]; + instance.timeout = "1s".to_string(); + + provider_repo(&infra.pool) + .create(&instance) + .await + .expect("enabled instance should persist for reconnect coverage"); + + let result = manager.reconnect(&instance.name).await; + assert!( + result.is_err(), + "reconnect must fail when the remote endpoint is unreachable" + ); +} + +async fn scenario_update_unreachable_remote_instance_preserves_existing_configuration() { + let infra = TestInfra::new().await; + flush_provider_instances(&infra).await; + let (health_addr, health_handle) = + spawn_authenticated_provider_server("remote-provider-test-secret").await; + + let repo = provider_repo(&infra.pool); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + HashMap::from([( + "update-provider.test.localhost".to_string(), + SocketAddr::from((std::net::Ipv4Addr::LOCALHOST, health_addr.port())), + )]), + ); + + let mut instance = make_test_instance("test-instance-unreachable-update"); + instance.endpoint = format!( + "http://update-provider.test.localhost:{}", + health_addr.port() + ); + instance.providers = vec!["alist".to_string()]; + manager + .add(instance.clone()) + .await + .expect("reachable instance should be added before update"); + + let mut updated = instance.clone(); + updated.endpoint = "http://unreachable-provider.example.invalid:50051".to_string(); + updated.timeout = "1s".to_string(); + + let result = manager.update(updated.clone()).await; + assert!( + result.is_err(), + "update must fail when the new remote endpoint is unreachable" + ); + + let stored = provider_repo(&infra.pool) + .get_by_name(&instance.name) + .await + .expect("lookup should succeed") + .expect("instance should still exist"); + assert_eq!( + stored.endpoint, instance.endpoint, + "failed update must preserve the last known-good endpoint" + ); + + health_handle.abort(); + let _ = health_handle.await; +} + +async fn scenario_add_stalling_remote_instance_honors_connect_timeout() { + let infra = TestInfra::new().await; + flush_provider_instances(&infra).await; + let (stall_addr, stall_handle) = spawn_stalling_tcp_server().await; + let host = "connect-timeout.test.localhost"; + + let repo = provider_repo(&infra.pool); + let manager = RemoteProviderManager::new_with_test_address_overrides( + Arc::new(repo), + None, + make_test_address_overrides(host, stall_addr.port()), + ); + + let mut instance = + make_reachable_remote_instance("test-instance-connect-timeout", host, stall_addr.port()); + instance.timeout = "500ms".to_string(); + + let start = tokio::time::Instant::now(); + let result = manager.add(instance.clone()).await; + let elapsed = start.elapsed(); + + assert!( + result.is_err(), + "stalling remote endpoint must fail connectivity validation" + ); + assert!( + elapsed < Duration::from_millis(1500), + "configured timeout should bound management-path validation latency, elapsed: {elapsed:?}" + ); + + stall_handle.abort(); + let _ = stall_handle.await; +} + async fn scenario_legacy_remote_instance_without_jwt_secret_is_rejected_at_runtime() { let infra = TestInfra::new().await; flush_provider_instances(&infra).await; @@ -2248,6 +3011,28 @@ async fn test_health_check_reports_authenticated_provider_failure_as_unhealthy() scenario_health_check_reports_authenticated_provider_failure_as_unhealthy().await; } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[ignore = "Requires Docker"] +async fn test_health_check_reports_emby_authenticated_provider_failure_as_unhealthy() { + install_rustls_provider_once(); + scenario_health_check_reports_emby_authenticated_provider_failure_as_unhealthy().await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[ignore = "Requires Docker"] +async fn test_add_alist_instance_does_not_require_fake_upstream_auth_for_management_validation() { + install_rustls_provider_once(); + scenario_add_alist_instance_does_not_require_fake_upstream_auth_for_management_validation() + .await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[ignore = "Requires Docker"] +async fn test_add_emby_instance_rejects_authenticated_handler_failure() { + install_rustls_provider_once(); + scenario_add_emby_instance_rejects_authenticated_handler_failure().await; +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[ignore = "Requires Docker"] async fn test_tls_configuration_secure() { @@ -2444,6 +3229,48 @@ async fn test_add_local_only_instance_allows_empty_jwt_secret() { scenario_add_local_only_instance_allows_empty_jwt_secret().await; } +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[ignore = "Requires Docker"] +async fn test_add_unreachable_remote_instance_fails_connectivity_validation() { + install_rustls_provider_once(); + scenario_add_unreachable_remote_instance_fails_connectivity_validation().await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[ignore = "Requires Docker"] +async fn test_add_reachable_remote_instance_succeeds_with_connectivity_validation() { + install_rustls_provider_once(); + scenario_add_reachable_remote_instance_succeeds_with_connectivity_validation().await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[ignore = "Requires Docker"] +async fn test_enable_unreachable_remote_instance_preserves_disabled_state() { + install_rustls_provider_once(); + scenario_enable_unreachable_remote_instance_preserves_disabled_state().await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[ignore = "Requires Docker"] +async fn test_reconnect_unreachable_remote_instance_fails_connectivity_validation() { + install_rustls_provider_once(); + scenario_reconnect_unreachable_remote_instance_fails_connectivity_validation().await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[ignore = "Requires Docker"] +async fn test_update_unreachable_remote_instance_preserves_existing_configuration() { + install_rustls_provider_once(); + scenario_update_unreachable_remote_instance_preserves_existing_configuration().await; +} + +#[tokio::test(flavor = "multi_thread", worker_threads = 2)] +#[ignore = "Requires Docker"] +async fn test_add_stalling_remote_instance_honors_connect_timeout() { + install_rustls_provider_once(); + scenario_add_stalling_remote_instance_honors_connect_timeout().await; +} + #[tokio::test(flavor = "multi_thread", worker_threads = 2)] #[ignore = "Requires Docker"] async fn test_legacy_remote_instance_without_jwt_secret_is_rejected_at_runtime() { diff --git a/synctv-core/tests/user_auth_service_tests.rs b/synctv-core/tests/user_auth_service_tests.rs index 0e758a0c..a354b82f 100644 --- a/synctv-core/tests/user_auth_service_tests.rs +++ b/synctv-core/tests/user_auth_service_tests.rs @@ -15,10 +15,11 @@ use sqlx::PgPool; use synctv_core::{ cache::{CacheL2Backend, KeyBuilder, NoopCacheL2, UsernameCache}, config::PasswordComplexityConfig, - models::UserId, + models::{OAuth2Provider, UserId}, + repository::{UserOAuthProviderRepository, UserRepository}, service::{ - auth::jwt::JwtService, BruteForceProtection, InMemoryTokenBlacklistStore, RateLimiter, - TokenBlacklistStore, UserService, + auth::jwt::JwtService, BruteForceProtection, InMemoryOAuthStateStore, + InMemoryTokenBlacklistStore, OAuth2Service, RateLimiter, TokenBlacklistStore, UserService, }, Error, }; @@ -1446,6 +1447,11 @@ async fn test_create_or_load_by_oauth2_username_sanitization() { !result.username.contains('!'), "! should be stripped from username" ); + assert_eq!( + result.status, + synctv_core::models::UserStatus::Active, + "OAuth2-created users should start active so first login succeeds" + ); } #[tokio::test] @@ -1536,6 +1542,142 @@ async fn test_create_or_load_by_oauth2_empty_username_uses_provider_id() { ); } +#[tokio::test] +#[ignore = "Requires Docker"] +async fn test_find_or_create_and_link_concurrent_requests_do_not_commit_orphan_oauth2_users() { + let (_container, pool) = create_test_pool().await; + let user_service = create_user_service(pool.clone()); + let oauth_service = OAuth2Service::new( + UserOAuthProviderRepository::new(pool.clone()), + Arc::new(InMemoryOAuthStateStore::new()), + synctv_core::oauth2::ProviderRegistry::new(), + false, + ) + .expect("OAuth2 service should initialize"); + + let provider = OAuth2Provider::Google; + let user_info = synctv_core::service::OAuth2UserInfo { + provider: provider.clone(), + provider_user_id: format!("oauth_concurrent_{}", nanoid::nanoid!(8)), + username: format!("oauth_concurrent_user_{}", nanoid::nanoid!(6)), + email: Some(format!("oauth_concurrent_{}@test.com", nanoid::nanoid!(6))), + avatar: None, + email_verified: true, + }; + + let first = oauth_service.find_or_create_and_link(&user_service, &provider, &user_info); + let second = oauth_service.find_or_create_and_link(&user_service, &provider, &user_info); + let (first_result, second_result) = tokio::join!(first, second); + + let (first_user_id, _) = first_result.expect("first concurrent login must succeed"); + let (second_user_id, _) = second_result.expect("second concurrent login must succeed"); + assert_eq!( + first_user_id, second_user_id, + "Concurrent logins for the same provider identity must converge to one user" + ); + + let oauth_repo = UserOAuthProviderRepository::new(pool.clone()); + let mapping = oauth_repo + .find_by_provider(&provider, &user_info.provider_user_id) + .await + .expect("mapping lookup must succeed") + .expect("mapping must exist"); + assert_eq!(mapping.user_id, first_user_id); + + let user_repo = UserRepository::new(pool.clone()); + let oauth2_user_count: i64 = sqlx::query_scalar( + r" + SELECT COUNT(*) + FROM users u + JOIN oauth2_clients oc ON oc.user_id = u.id + WHERE oc.provider = $1 + AND oc.provider_user_id = $2 + AND u.deleted_at IS NULL + ", + ) + .bind(provider.as_str()) + .bind(&user_info.provider_user_id) + .fetch_one(&pool) + .await + .expect("user count query must succeed"); + assert_eq!( + oauth2_user_count, 1, + "Concurrent OAuth2 signups must not commit an extra orphan user row" + ); + + let persisted_user = user_repo + .get_by_id(&first_user_id) + .await + .expect("user lookup must succeed") + .expect("winning user must exist"); + assert_eq!(persisted_user.email.as_deref(), user_info.email.as_deref()); + assert!(persisted_user.email_verified); + assert_eq!( + persisted_user.status, + synctv_core::models::UserStatus::Active, + "OAuth2-created users must be active immediately" + ); +} + +#[tokio::test] +#[ignore = "Requires Docker"] +async fn test_find_or_create_and_link_retries_with_suffixed_username_on_collision() { + let (_container, pool) = create_test_pool().await; + let user_service = create_user_service(pool.clone()); + let oauth_service = OAuth2Service::new( + UserOAuthProviderRepository::new(pool.clone()), + Arc::new(InMemoryOAuthStateStore::new()), + synctv_core::oauth2::ProviderRegistry::new(), + false, + ) + .expect("OAuth2 service should initialize"); + + user_service + .register( + "oauth_collision_user".to_string(), + Some("local_collision@test.com".to_string()), + "StrongPass1".to_string(), + None, + ) + .await + .expect("seed local user should be created"); + + let provider = OAuth2Provider::Google; + let user_info = synctv_core::service::OAuth2UserInfo { + provider: provider.clone(), + provider_user_id: format!("oauth_collision_{}", nanoid::nanoid!(8)), + username: "oauth_collision_user".to_string(), + email: Some(format!("oauth_collision_{}@test.com", nanoid::nanoid!(6))), + avatar: None, + email_verified: true, + }; + + let (created_user_id, is_new) = oauth_service + .find_or_create_and_link(&user_service, &provider, &user_info) + .await + .expect("OAuth2 signup should succeed by choosing a suffixed username"); + + assert!(is_new, "first OAuth2 login should create a new user"); + + let user_repo = UserRepository::new(pool.clone()); + let created_user = user_repo + .get_by_id(&created_user_id) + .await + .expect("user lookup should succeed") + .expect("created OAuth2 user should exist"); + + assert_ne!(created_user.username, "oauth_collision_user"); + assert!( + created_user.username.starts_with("oauth_collision_user_"), + "expected suffixed username, got {}", + created_user.username + ); + assert_eq!( + created_user.signup_method, + synctv_core::models::SignupMethod::OAuth2 + ); +} + // ============================================================================ // S1 additional: refresh_token with email verification re-check // ============================================================================ diff --git a/synctv-core/tests/user_oauth_provider_repository_tests.rs b/synctv-core/tests/user_oauth_provider_repository_tests.rs index 75729b27..9eb7885f 100644 --- a/synctv-core/tests/user_oauth_provider_repository_tests.rs +++ b/synctv-core/tests/user_oauth_provider_repository_tests.rs @@ -1,6 +1,6 @@ //! `UserOAuthProviderRepository` integration tests //! -//! Tests: upsert with different `user_id`, transaction executor path, +//! Tests: upsert conflict handling, transaction executor path, //! `delete_all_for_user_with_executor`. //! //! Run with: cargo test -p synctv-core --test `user_oauth_provider_repository_tests` @@ -38,11 +38,11 @@ async fn create_user(pool: &PgPool, username: &str) -> User { user_repo.create(&make_user(username)).await.unwrap() } -// ─── upsert with different user_id (OAuth identity update) ─────────── +// ─── upsert conflict handling ──────────────────────────────────────── #[tokio::test] #[ignore = "Requires Docker"] -async fn test_upsert_different_user_id_updates_mapping() { +async fn test_upsert_different_user_id_rejects_rebinding_and_preserves_mapping() { let (_container, pool) = create_test_pool().await; let oauth_repo = UserOAuthProviderRepository::new(pool.clone()); @@ -72,11 +72,21 @@ async fn test_upsert_different_user_id_updates_mapping() { .unwrap(); assert_eq!(mapping.user_id, user_a.id); - // Upsert again with user_b (re-linking the OAuth identity) - oauth_repo + // Upsert again with user_b must be rejected: external identities are stable + // and must never be silently reassigned to another local user. + let err = oauth_repo .upsert(&user_b.id, &provider, provider_user_id, &user_info) .await - .unwrap(); + .expect_err("OAuth identity rebinding must be rejected"); + + assert!( + matches!( + err, + synctv_core::Error::AlreadyExists(ref msg) + if msg.contains("already linked to another user") + ), + "Unexpected error: {err}" + ); let mapping = oauth_repo .find_by_provider(&provider, provider_user_id) @@ -84,19 +94,124 @@ async fn test_upsert_different_user_id_updates_mapping() { .unwrap() .unwrap(); assert_eq!( - mapping.user_id, user_b.id, - "OAuth identity should now be linked to user_b" + mapping.user_id, user_a.id, + "Original OAuth identity binding must be preserved" ); - // user_a should no longer have this mapping + // user_a must still own the mapping let user_a_mappings = oauth_repo.find_by_user(&user_a.id).await.unwrap(); + assert_eq!(user_a_mappings.len(), 1); + + // user_b must not gain the mapping + let user_b_mappings = oauth_repo.find_by_user(&user_b.id).await.unwrap(); assert!( - user_a_mappings.is_empty(), - "user_a should have no OAuth mappings after re-link" + user_b_mappings.is_empty(), + "user_b should not receive another user's OAuth mapping" + ); +} + +#[tokio::test] +#[ignore = "Requires Docker"] +async fn test_upsert_same_user_id_updates_profile_fields_without_rebinding() { + let (_container, pool) = create_test_pool().await; + let oauth_repo = UserOAuthProviderRepository::new(pool.clone()); + + let user = create_user(&pool, "oauth_profile_user").await; + + let provider = OAuth2Provider::GitHub; + let provider_user_id = "gh_profile_001"; + let initial_info = OAuth2UserInfo { + provider: provider.clone(), + provider_user_id: provider_user_id.to_string(), + username: "oldname".to_string(), + email: None, + avatar: None, + }; + + oauth_repo + .upsert(&user.id, &provider, provider_user_id, &initial_info) + .await + .unwrap(); + + let updated_info = OAuth2UserInfo { + provider: provider.clone(), + provider_user_id: provider_user_id.to_string(), + username: "newname".to_string(), + email: Some("new@example.com".to_string()), + avatar: Some("https://avatar.example/new.png".to_string()), + }; + + oauth_repo + .upsert(&user.id, &provider, provider_user_id, &updated_info) + .await + .unwrap(); + + let mapping = oauth_repo + .find_by_provider(&provider, provider_user_id) + .await + .unwrap() + .unwrap(); + assert_eq!(mapping.user_id, user.id); + assert_eq!(mapping.username, "newname"); + assert_eq!(mapping.email.as_deref(), Some("new@example.com")); + assert_eq!( + mapping.avatar_url.as_deref(), + Some("https://avatar.example/new.png") ); } -// ─── transaction executor path ─────────────────────────────────────── +#[tokio::test] +#[ignore = "Requires Docker"] +async fn test_upsert_with_executor_rejects_rebinding_inside_transaction() { + let (_container, pool) = create_test_pool().await; + let oauth_repo = UserOAuthProviderRepository::new(pool.clone()); + + let user_a = create_user(&pool, "oauth_tx_owner").await; + let user_b = create_user(&pool, "oauth_tx_conflict").await; + let provider = OAuth2Provider::Google; + let provider_user_id = "google_tx_conflict_001"; + let user_info = OAuth2UserInfo { + provider: provider.clone(), + provider_user_id: provider_user_id.to_string(), + username: "googleuser".to_string(), + email: Some("tx@google.com".to_string()), + avatar: None, + }; + + oauth_repo + .upsert(&user_a.id, &provider, provider_user_id, &user_info) + .await + .unwrap(); + + let mut tx = pool.begin().await.unwrap(); + let err = oauth_repo + .upsert_with_executor( + &user_b.id, + &provider, + provider_user_id, + &user_info, + &mut *tx, + ) + .await + .expect_err("Rebinding in transaction must be rejected"); + tx.rollback().await.unwrap(); + + assert!( + matches!( + err, + synctv_core::Error::AlreadyExists(ref msg) + if msg.contains("already linked to another user") + ), + "Unexpected error: {err}" + ); + + let mapping = oauth_repo + .find_by_provider(&provider, provider_user_id) + .await + .unwrap() + .unwrap(); + assert_eq!(mapping.user_id, user_a.id); +} #[tokio::test] #[ignore = "Requires Docker"] @@ -124,11 +239,13 @@ async fn test_upsert_with_executor_in_transaction() { tx.commit().await.unwrap(); // Verify it was persisted + let user_mappings = oauth_repo.find_by_user(&user.id).await.unwrap(); let mapping = oauth_repo .find_by_provider(&provider, provider_user_id) .await .unwrap() .unwrap(); + assert_eq!(user_mappings.len(), 1); assert_eq!(mapping.user_id, user.id); assert_eq!(mapping.email.as_deref(), Some("tx@google.com")); } diff --git a/synctv-livestream/src/relay/in_memory_registry.rs b/synctv-livestream/src/relay/in_memory_registry.rs index 6ef68773..93253734 100644 --- a/synctv-livestream/src/relay/in_memory_registry.rs +++ b/synctv-livestream/src/relay/in_memory_registry.rs @@ -113,11 +113,13 @@ impl StreamRegistryTrait for InMemoryStreamRegistry { _user_id: &str, ) -> Result { let publishers = self.publishers.lock().await; - Ok(if publishers.contains_key(&(room_id.to_string(), media_id.to_string())) { - PublisherRefreshOutcome::Refreshed - } else { - PublisherRefreshOutcome::Missing - }) + Ok( + if publishers.contains_key(&(room_id.to_string(), media_id.to_string())) { + PublisherRefreshOutcome::Refreshed + } else { + PublisherRefreshOutcome::Missing + }, + ) } async fn unregister_publisher(&self, room_id: &str, media_id: &str) -> Result<()> { diff --git a/synctv-livestream/src/relay/mock_registry.rs b/synctv-livestream/src/relay/mock_registry.rs index dbf76d1b..80afbe72 100644 --- a/synctv-livestream/src/relay/mock_registry.rs +++ b/synctv-livestream/src/relay/mock_registry.rs @@ -29,11 +29,9 @@ pub struct MockStreamRegistry { fail_refresh_publisher_ttl: std::sync::Arc, /// When true, `refresh_publisher_ttl` returns the wrapped timeout error shape /// used by the real Redis-backed registry helper. - fail_refresh_publisher_ttl_with_wrapped_timeout: - std::sync::Arc, + fail_refresh_publisher_ttl_with_wrapped_timeout: std::sync::Arc, /// When true, `refresh_publisher_ttl` returns a persistent non-I/O registry error. - fail_refresh_publisher_ttl_with_response_error: - std::sync::Arc, + fail_refresh_publisher_ttl_with_response_error: std::sync::Arc, /// Number of upcoming epoch-fenced unregister calls that should fail. fail_unregister_if_epoch_matches_times: std::sync::Arc, } @@ -57,7 +55,9 @@ impl MockStreamRegistry { std::sync::atomic::AtomicUsize::new(0), ), fail_get_publisher: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)), - fail_refresh_publisher_ttl: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)), + fail_refresh_publisher_ttl: std::sync::Arc::new(std::sync::atomic::AtomicBool::new( + false, + )), fail_refresh_publisher_ttl_with_wrapped_timeout: std::sync::Arc::new( std::sync::atomic::AtomicBool::new(false), ), @@ -88,7 +88,9 @@ impl MockStreamRegistry { std::sync::atomic::AtomicUsize::new(0), ), fail_get_publisher: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)), - fail_refresh_publisher_ttl: std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)), + fail_refresh_publisher_ttl: std::sync::Arc::new(std::sync::atomic::AtomicBool::new( + false, + )), fail_refresh_publisher_ttl_with_wrapped_timeout: std::sync::Arc::new( std::sync::atomic::AtomicBool::new(false), ), @@ -297,10 +299,8 @@ impl StreamRegistryTrait for MockStreamRegistry { .fail_unregister_if_epoch_matches_times .load(std::sync::atomic::Ordering::SeqCst); if remaining_failures > 0 { - self.fail_unregister_if_epoch_matches_times.fetch_sub( - 1, - std::sync::atomic::Ordering::SeqCst, - ); + self.fail_unregister_if_epoch_matches_times + .fetch_sub(1, std::sync::atomic::Ordering::SeqCst); let redis_error = redis::RedisError::from(( redis::ErrorKind::Io, "simulated Redis failure in unregister_publisher_if_epoch_matches", diff --git a/synctv-livestream/src/relay/publisher_manager.rs b/synctv-livestream/src/relay/publisher_manager.rs index 262a0bce..709f1896 100644 --- a/synctv-livestream/src/relay/publisher_manager.rs +++ b/synctv-livestream/src/relay/publisher_manager.rs @@ -656,7 +656,8 @@ impl PublisherManager { "Reconcile (reverse): adding missing publisher room={} media={} to local tracking", room_id, media_id ); - let entry = Arc::new(PublisherEntry::with_registration(info.user_id, info.epoch)); + let entry = + Arc::new(PublisherEntry::with_registration(info.user_id, info.epoch)); self.active_publishers.insert(publisher_key, entry); added += 1; } @@ -970,9 +971,12 @@ impl PublisherManager { } let tracked_entry = Arc::clone(entry.value()); drop(entry); - let Some((_, removed_entry)) = self.active_publishers.remove_if(&publisher_key, |_, current| { - Arc::ptr_eq(current, &tracked_entry) && current.epoch == expected_epoch - }) else { + let Some((_, removed_entry)) = self + .active_publishers + .remove_if(&publisher_key, |_, current| { + Arc::ptr_eq(current, &tracked_entry) && current.epoch == expected_epoch + }) + else { debug!( "Cleanup skipped for room={} media={} because local owner changed before removal", room_id, media_id @@ -1193,9 +1197,8 @@ impl PublisherManager { // // Use structured redis::ErrorKind matching instead of string comparison // to avoid brittle matching against error message text. - let is_redis_unreachable = last_error - .as_ref() - .is_some_and(is_redis_unreachable_error); + let is_redis_unreachable = + last_error.as_ref().is_some_and(is_redis_unreachable_error); if is_redis_unreachable { // Redis itself is unreachable — do NOT count toward publisher cleanup threshold. @@ -1217,9 +1220,7 @@ impl PublisherManager { room_id.to_string(), media_id.to_string(), entry.epoch, - format!( - "Redis unreachable for {redis_failures} consecutive cycles" - ), + format!("Redis unreachable for {redis_failures} consecutive cycles"), ); } else { warn!( @@ -1301,8 +1302,8 @@ impl PublisherManager { mod tests { use super::super::MockStreamRegistry; use super::*; - use anyhow::Result; use crate::relay::PublisherInfo; + use anyhow::Result; use chrono::Utc; /// Create a test `PublisherManager` with a dummy `StreamHubEventSender`. @@ -1563,7 +1564,10 @@ mod tests { async fn list_active_streams(&self) -> Result> { Ok(if self.publisher.lock().await.is_some() { - vec![("room-reregister".to_string(), "media-reregister".to_string())] + vec![( + "room-reregister".to_string(), + "media-reregister".to_string(), + )] } else { Vec::new() }) @@ -1571,14 +1575,19 @@ mod tests { async fn get_user_publishers(&self, user_id: &str) -> Result> { let publisher = self.publisher.lock().await; - Ok(if publisher - .as_ref() - .is_some_and(|current| current.user_id == user_id) - { - vec![("room-reregister".to_string(), "media-reregister".to_string())] - } else { - Vec::new() - }) + Ok( + if publisher + .as_ref() + .is_some_and(|current| current.user_id == user_id) + { + vec![( + "room-reregister".to_string(), + "media-reregister".to_string(), + )] + } else { + Vec::new() + }, + ) } async fn unregister_all_user_publishers(&self, user_id: &str) -> Result<()> { @@ -1956,7 +1965,8 @@ mod tests { }) .expect("fill channel to create backpressure"); - let cleanup = manager.cleanup_publisher("room-backpressure", "media-backpressure", 1, "test"); + let cleanup = + manager.cleanup_publisher("room-backpressure", "media-backpressure", 1, "test"); let delayed_recv = async move { tokio::time::sleep(Duration::from_millis(50)).await; let _ = rx.recv().await; @@ -2077,7 +2087,9 @@ mod tests { .await; assert!( - !manager.active_publishers.contains_key("room-retry:media-retry"), + !manager + .active_publishers + .contains_key("room-retry:media-retry"), "cleanup should remove the local publisher entry" ); @@ -2110,7 +2122,13 @@ mod tests { let (manager, mut rx) = test_manager(registry.clone(), "test-node"); registry - .try_register_publisher("room-recover", "media-recover", "test-node", "user1", "addr1") + .try_register_publisher( + "room-recover", + "media-recover", + "test-node", + "user1", + "addr1", + ) .await .unwrap(); let original = registry @@ -2211,7 +2229,10 @@ mod tests { tokio::time::advance(Duration::from_millis(UNREGISTER_RETRY_DELAYS_MS[1] + 1)).await; tokio::task::yield_now().await; tokio::task::yield_now().await; - let event = rx.recv().await.expect("heartbeat cleanup should emit unpublish"); + let event = rx + .recv() + .await + .expect("heartbeat cleanup should emit unpublish"); let StreamHubEvent::UnPublish { identifier } = event else { panic!("expected unpublish event"); }; @@ -2230,7 +2251,9 @@ mod tests { }; observe.await; - heartbeat.await.expect("heartbeat task should complete successfully"); + heartbeat + .await + .expect("heartbeat task should complete successfully"); } #[tokio::test(start_paused = true)] @@ -2263,7 +2286,9 @@ mod tests { } assert!( - manager.active_publishers.contains_key("room-redis:media-redis"), + manager + .active_publishers + .contains_key("room-redis:media-redis"), "Redis connectivity failures must not clear active publishers" ); assert!( @@ -2276,7 +2301,13 @@ mod tests { async fn test_wrapped_redis_timeout_counts_as_unreachable_cycle() { let registry = Arc::new(MockStreamRegistry::new()); registry - .try_register_publisher("room-timeout", "media-timeout", "test-node", "user1", "addr1") + .try_register_publisher( + "room-timeout", + "media-timeout", + "test-node", + "user1", + "addr1", + ) .await .unwrap(); let current = registry @@ -2315,7 +2346,9 @@ mod tests { "wrapped Redis timeouts must not count as publisher-missing failures" ); assert!( - manager.active_publishers.contains_key("room-timeout:media-timeout"), + manager + .active_publishers + .contains_key("room-timeout:media-timeout"), "last-resort cleanup must not trigger before the redis-unreachable threshold" ); assert!( @@ -2364,7 +2397,10 @@ mod tests { tokio::task::yield_now().await; tokio::time::advance(Duration::from_millis(HEARTBEAT_RETRY_BASE_DELAY_MS + 1)).await; tokio::task::yield_now().await; - tokio::time::advance(Duration::from_millis((HEARTBEAT_RETRY_BASE_DELAY_MS * 2) + 1)).await; + tokio::time::advance(Duration::from_millis( + (HEARTBEAT_RETRY_BASE_DELAY_MS * 2) + 1, + )) + .await; tokio::task::yield_now().await; assert!( @@ -2396,7 +2432,9 @@ mod tests { } ); assert!( - !manager.active_publishers.contains_key("room-outage:media-outage"), + !manager + .active_publishers + .contains_key("room-outage:media-outage"), "background cleanup should eventually remove local state after redis outage threshold" ); assert_eq!(registry.unregister_if_epoch_matches_call_count(), 3); @@ -2406,7 +2444,13 @@ mod tests { async fn test_persistent_non_io_registry_failures_still_trigger_cleanup() { let registry = Arc::new(MockStreamRegistry::new()); registry - .try_register_publisher("room-response", "media-response", "test-node", "user1", "addr1") + .try_register_publisher( + "room-response", + "media-response", + "test-node", + "user1", + "addr1", + ) .await .unwrap(); let current = registry @@ -2420,9 +2464,10 @@ mod tests { "user1".to_string(), current.epoch, )); - manager - .active_publishers - .insert("room-response:media-response".to_string(), Arc::clone(&entry)); + manager.active_publishers.insert( + "room-response:media-response".to_string(), + Arc::clone(&entry), + ); registry.set_fail_refresh_publisher_ttl_with_response_error(true); @@ -2441,7 +2486,9 @@ mod tests { "persistent non-I/O registry failures must not increment redis_unreachable_cycles" ); assert!( - manager.active_publishers.contains_key("room-response:media-response"), + manager + .active_publishers + .contains_key("room-response:media-response"), "cleanup should wait until the heartbeat-failure threshold is reached" ); assert!( @@ -2454,10 +2501,15 @@ mod tests { tokio::task::yield_now().await; assert!( - !manager.active_publishers.contains_key("room-response:media-response"), + !manager + .active_publishers + .contains_key("room-response:media-response"), "persistent non-I/O registry failures should eventually trigger cleanup" ); - let event = rx.recv().await.expect("cleanup should emit unpublish at threshold"); + let event = rx + .recv() + .await + .expect("cleanup should emit unpublish at threshold"); let StreamHubEvent::UnPublish { identifier } = event else { panic!("expected unpublish event"); }; @@ -2537,7 +2589,9 @@ mod tests { "publisher-missing failure must reset redis-unreachable streak" ); assert!( - manager.active_publishers.contains_key("room-switch:media-switch"), + manager + .active_publishers + .contains_key("room-switch:media-switch"), "mixed failure classes must not trigger premature cleanup" ); assert!( diff --git a/synctv/src/migrations.rs b/synctv/src/migrations.rs index 5ef580ed..6ae234ac 100644 --- a/synctv/src/migrations.rs +++ b/synctv/src/migrations.rs @@ -46,14 +46,9 @@ async fn run_migrations_with_mode( key_prefix: &str, cluster_mode: bool, ) -> Result<()> { - run_migrations_with_runner( - pool, - lock, - key_prefix, - cluster_mode, - true, - &|pool| Box::pin(run_migrate(pool)), - ) + run_migrations_with_runner(pool, lock, key_prefix, cluster_mode, true, &|pool| { + Box::pin(run_migrate(pool)) + }) .await }