feat: more test

pull/370/head
zijiren233 6 months ago
parent 95bdc7aea9
commit 42ca2f56c7
No known key found for this signature in database
GPG Key ID: 534E082AAA9B39DC

@ -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://<server.host>/api/oauth2/<instance_name>/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/<instance_name>/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: ""

@ -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(

@ -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 ==========

@ -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");

@ -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!(

@ -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<AppState> {
#[cfg(test)]
fn register_all_routes_for_test(state: AppState) -> Router<AppState> {
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<AppState>, Router<AppState>, Router<AppState>) {
@ -683,7 +686,7 @@ fn register_all_routes(state: AppState) -> (Router<AppState>, Router<AppState>,
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"),
)

@ -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(

@ -226,7 +226,9 @@ impl From<synctv_core::Error> 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!(

@ -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 _<provider>, _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

@ -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"

@ -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;

@ -11,9 +11,8 @@ use tokio::sync::{broadcast, mpsc};
use tracing::{debug, info, warn};
#[cfg(test)]
type AsyncTestHook = Arc<
dyn Fn() -> std::pin::Pin<Box<dyn std::future::Future<Output = ()> + Send>> + Send + Sync,
>;
type AsyncTestHook =
Arc<dyn Fn() -> std::pin::Pin<Box<dyn std::future::Future<Output = ()> + 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<tokio::sync::Mutex<()>> {
fn connection_lifecycle_lock(&self, connection_id: &str) -> Arc<tokio::sync::Mutex<()>> {
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 = {

@ -664,7 +664,10 @@ async fn init_oauth2_service(
let state_store: Arc<dyn crate::service::OAuthStateStore> = 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();

@ -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<T, F>(future: F) -> std::result::Result<T, L2RedisAttemptError>
where
F: Future<Output = std::result::Result<T, redis::RedisError>>,
{
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<T, F>(operation: impl Into<String>, future: F) -> Result<T>
where
F: Future<Output = std::result::Result<T, redis::RedisError>>,
{
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<String>>(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<String>>(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<Option<String>> =
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<String>) = tokio::time::timeout(
REDIS_OPERATION_TIMEOUT,
let scan_result: (u64, Vec<String>) = 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() {

@ -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;

@ -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<BasicClient<EndpointSet, EndpointNotSet, EndpointNotSet, EndpointNotSet, EndpointSet>>,
oauth2_http_client: Arc<oauth2::reqwest::Client>,
http_client: Arc<Client>,
}
@ -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")?;

@ -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<BasicClient<EndpointSet, EndpointNotSet, EndpointNotSet, EndpointNotSet, EndpointSet>>,
oauth2_http_client: Arc<oauth2::reqwest::Client>,
http_client: Arc<Client>,
}
@ -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")?;

@ -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<BasicClient<EndpointSet, EndpointNotSet, EndpointNotSet, EndpointNotSet, EndpointSet>>,
endpoint: String,
oauth2_http_client: Arc<oauth2::reqwest::Client>,
http_client: Arc<Client>,
}
@ -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")?;

@ -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, Error> {
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<Arc<Client>, 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::Client, Error> {
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<Arc<oauth2::reqwest::Client>, Error> {
Ok(Arc::new(build_oauth2_http_client_with_timeout(
HTTP_REQUEST_TIMEOUT,
)?))
}
pub(super) fn map_provider_http_error<E>(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")
));
}
}

@ -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<ResolvedOidc>,
/// Stored config for lazy initialization (only used in issuer-only mode)
init_config: OidcInitConfig,
oauth2_http_client: Arc<oauth2::reqwest::Client>,
http_client: Arc<Client>,
}
@ -85,13 +87,6 @@ impl OidcProvider {
redirect_url: String,
issuer: &str,
) -> Result<Self, Error> {
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<String>,
) -> Result<Self, Error> {
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")?;

@ -346,6 +346,11 @@ impl AlistProvider {
#[async_trait]
impl MediaProvider for AlistProvider {
#[cfg(test)]
fn test_client_manager_marker(&self) -> Option<usize> {
Some(self.client_manager.marker())
}
fn name(&self) -> &'static str {
Self::NAME
}

@ -292,6 +292,11 @@ impl TryFrom<&Value> for BilibiliSourceConfig {
#[async_trait]
impl MediaProvider for BilibiliProvider {
#[cfg(test)]
fn test_client_manager_marker(&self) -> Option<usize> {
Some(self.client_manager.marker())
}
fn name(&self) -> &'static str {
Self::NAME
}

@ -405,6 +405,11 @@ impl TryFrom<&Value> for EmbySourceConfig {
#[async_trait]
impl MediaProvider for EmbyProvider {
#[cfg(test)]
fn test_client_manager_marker(&self) -> Option<usize> {
Some(self.client_manager.marker())
}
fn name(&self) -> &'static str {
Self::NAME
}

@ -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<impl Into<String>>,
) -> Self {
pub fn new(channel: tonic::transport::Channel, auth_secret: Option<impl Into<String>>) -> Self {
Self {
channel,
auth_secret: auth_secret.map(|secret| Arc::<str>::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<T>(
payload: T,
) -> Result<Request<T>, 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<T>(
Ok(request)
}
pub(crate) fn validate_auth_secret(auth_secret: Option<&str>) -> Result<Option<&str>, ProviderError> {
pub(crate) fn validate_auth_secret(
auth_secret: Option<&str>,
) -> Result<Option<&str>, 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<reqwest::StatusCode> {
}
}
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"

@ -183,6 +183,11 @@ pub trait MediaProvider: Send + Sync {
None
}
#[cfg(test)]
fn test_client_manager_marker(&self) -> Option<usize> {
None
}
// ========== Validation ==========
/// Validate `source_config` before saving to database

@ -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(())
}

@ -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::<std::io::Error>() {
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::<std::io::Error>() {
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(&not_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));
}
}

@ -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;
}

@ -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());

@ -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<dyn TokenBlacklistStore>;
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]

@ -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;
}

@ -61,6 +61,30 @@ use redis::aio::ConnectionManager as RedisConnectionManager;
use redis::Script;
use std::future::Future;
async fn run_distributed_lock_redis_op<T, F>(operation: impl Into<String>, future: F) -> Result<T>
where
F: Future<Output = std::result::Result<T, redis::RedisError>>,
{
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<T, F>(
key: &str,
timeout: std::time::Duration,
future: F,
) -> Result<T>
where
F: Future<Output = Result<T>>,
{
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::<u64>(&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<String> = tokio::time::timeout(
crate::resilience::timeout::REDIS_OPERATION_TIMEOUT,
let result: Option<String> = 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::<i32>(&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::<i32>(&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:

@ -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;
}

@ -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<tokio::sync::RwLock<redis::aio::ConnectionManager>>,
key_prefix: String,
}
impl RedisOAuthStateStore {
async fn run_redis_op<T, F>(&self, operation: &'static str, future: F) -> Result<T>
where
F: Future<Output = std::result::Result<T, redis::RedisError>>,
{
run_oauth_state_redis_op(operation, future).await
}
fn normalize_key_prefix(prefix: impl Into<String>) -> 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<RwLock<ConnectionManager>>`.
#[must_use]
pub const fn new(
pub fn new(
conn: std::sync::Arc<tokio::sync::RwLock<redis::aio::ConnectionManager>>,
key_prefix: impl Into<String>,
) -> 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<T, F>(operation: &'static str, future: F) -> Result<T>
where
F: Future<Output = std::result::Result<T, redis::RedisError>>,
{
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<Option<OAuth2State>> {
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<String> = 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<String> = 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"
));
}
}

@ -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<RemoteProviderManager>,
/// Default injected local provider clients used by provider instances
/// when they do not specify a per-instance HTTP transport override.
default_client_manager: Arc<ProviderClientManager>,
/// 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<RemoteProviderManager>,
_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<String> {
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();

@ -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<RemoteProviderConnection> {
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<RemoteProviderConnection> {
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",
"<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", "<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", "<bilibili>", client.user_info(request).await)
Self::probe_reports_authenticated_health("emby", "<emby>", client.me(request).await)
}
fn build_authenticated_request<T>(
@ -1257,18 +1238,47 @@ impl RemoteProviderManager {
provider: &str,
instance_name: &str,
result: Result<tonic::Response<T>, 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;
}
}

@ -97,6 +97,51 @@ impl UserService {
}
}
pub(crate) fn oauth2_username_candidates(
&self,
provider_user_id: &str,
username: &str,
) -> Result<(String, Vec<String>)> {
let sanitized_username = username
.chars()
.filter(|c| c.is_alphanumeric() || *c == '_' || *c == '-')
.collect::<String>()
.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::<String>()
} 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<User> {
// Sanitize OAuth2 username: remove invalid characters and trim
let sanitized_username = username
.chars()
.filter(|c| c.is_alphanumeric() || *c == '_' || *c == '-')
.collect::<String>()
.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::<String>()
} 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()
)))
}

@ -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<T, F>(&self, operation: &'static str, future: F) -> Result<T>
where
F: Future<Output = std::result::Result<T, redis::RedisError>>,
{
run_ws_ticket_redis_op(operation, future).await
}
fn normalize_key_prefix(prefix: impl Into<String>) -> 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<T, F>(operation: &'static str, future: F) -> Result<T>
where
F: Future<Output = std::result::Result<T, redis::RedisError>>,
{
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<String> = 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<String> = 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<String> = 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<String> = 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);

@ -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`"),

@ -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]

@ -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"
);
}

File diff suppressed because it is too large Load Diff

@ -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
// ============================================================================

@ -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"));
}

@ -113,11 +113,13 @@ impl StreamRegistryTrait for InMemoryStreamRegistry {
_user_id: &str,
) -> Result<PublisherRefreshOutcome> {
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<()> {

@ -29,11 +29,9 @@ pub struct MockStreamRegistry {
fail_refresh_publisher_ttl: std::sync::Arc<std::sync::atomic::AtomicBool>,
/// 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<std::sync::atomic::AtomicBool>,
fail_refresh_publisher_ttl_with_wrapped_timeout: std::sync::Arc<std::sync::atomic::AtomicBool>,
/// When true, `refresh_publisher_ttl` returns a persistent non-I/O registry error.
fail_refresh_publisher_ttl_with_response_error:
std::sync::Arc<std::sync::atomic::AtomicBool>,
fail_refresh_publisher_ttl_with_response_error: std::sync::Arc<std::sync::atomic::AtomicBool>,
/// Number of upcoming epoch-fenced unregister calls that should fail.
fail_unregister_if_epoch_matches_times: std::sync::Arc<std::sync::atomic::AtomicUsize>,
}
@ -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",

@ -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<Vec<(String, String)>> {
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<Vec<(String, String)>> {
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!(

@ -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
}

Loading…
Cancel
Save