use std::{fmt::Write, num::NonZeroUsize, sync::Arc, time::Duration};
use axum::{
Router,
body::Body,
http::{HeaderMap, Request, StatusCode, header},
middleware,
routing::get,
};
use shardline_index::ProviderRepositoryState;
use shardline_protocol::RepositoryProvider;
use tempfile::TempDir;
use tokio::{net::TcpListener, sync::oneshot, time::timeout};
use tower::ServiceExt;
use super::{
MAX_BATCH_RECONSTRUCTION_FILE_IDS, MAX_BATCH_RECONSTRUCTION_QUERY_BYTES,
MAX_PROVIDER_BASIC_AUTH_HEADER_BYTES, MAX_PROVIDER_NAME_BYTES, MAX_PROVIDER_SUBJECT_BYTES,
MAX_PROVIDER_WEBHOOK_BODY_BYTES, bounded_api_body_limit, extract_provider_subject,
latest_lifecycle_signal_at, parse_batch_reconstruction_query,
reconciled_provider_repository_state, router, security_headers_middleware,
serve_with_listener_until, validate_provider_name_path,
};
use crate::{ServerConfig, ServerError, ServerFrontend, ServerRole, config::AuthProviderKind};
#[test]
fn provider_subject_extraction_rejects_oversized_query_subject() {
let oversized = "s".repeat(MAX_PROVIDER_SUBJECT_BYTES + 1);
let result = extract_provider_subject(&HeaderMap::new(), Some(&oversized));
assert!(matches!(
result,
Err(ServerError::InvalidProviderTokenRequest)
));
}
#[test]
fn provider_subject_extraction_rejects_oversized_basic_auth_header_before_decode() {
let oversized = "a".repeat(MAX_PROVIDER_BASIC_AUTH_HEADER_BYTES + 1);
let header_value = header::HeaderValue::from_str(&format!("Basic {oversized}"));
assert!(header_value.is_ok());
let Ok(header_value) = header_value else {
return;
};
let mut headers = HeaderMap::new();
headers.insert(header::AUTHORIZATION, header_value);
let result = extract_provider_subject(&headers, None);
assert!(matches!(
result,
Err(ServerError::InvalidAuthorizationHeader)
));
}
#[test]
fn provider_api_body_limit_uses_stricter_configured_or_endpoint_ceiling() {
let tighter = NonZeroUsize::new(32).unwrap_or(NonZeroUsize::MIN);
let looser =
NonZeroUsize::new(MAX_PROVIDER_WEBHOOK_BODY_BYTES + 1).unwrap_or(NonZeroUsize::MIN);
assert_eq!(
bounded_api_body_limit(tighter, MAX_PROVIDER_WEBHOOK_BODY_BYTES),
tighter.get()
);
assert_eq!(
bounded_api_body_limit(looser, MAX_PROVIDER_WEBHOOK_BODY_BYTES),
MAX_PROVIDER_WEBHOOK_BODY_BYTES
);
}
#[test]
fn provider_repository_reconciliation_marks_pending_lifecycle_signals() {
let state = ProviderRepositoryState::new(
RepositoryProvider::GitHub,
"team".to_owned(),
"assets".to_owned(),
Some(10),
Some(12),
Some("refs/heads/main".to_owned()),
)
.with_reconciliation(Some(11), None, None);
assert_eq!(latest_lifecycle_signal_at(&state), Some(12));
let reconciled = reconciled_provider_repository_state(&state, 20);
assert_eq!(
reconciled.last_cache_invalidated_at_unix_seconds(),
Some(20)
);
assert_eq!(
reconciled.last_authorization_rechecked_at_unix_seconds(),
Some(20)
);
assert_eq!(reconciled.last_drift_checked_at_unix_seconds(), Some(20));
}
#[test]
fn batch_reconstruction_parser_deduplicates_file_ids() {
let parsed = parse_batch_reconstruction_query(
"file_id=aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa&file_id=aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa&ignored=value",
);
assert!(parsed.is_ok());
let Ok(parsed) = parsed else {
return;
};
assert_eq!(
parsed,
vec!["aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa".to_owned()]
);
}
#[test]
fn batch_reconstruction_parser_rejects_excessive_file_ids() {
let mut query = String::new();
for index in 0..=MAX_BATCH_RECONSTRUCTION_FILE_IDS {
if !query.is_empty() {
query.push('&');
}
query.push_str("file_id=");
let written = write!(&mut query, "{index:064x}");
assert!(written.is_ok());
}
let parsed = parse_batch_reconstruction_query(&query);
assert!(matches!(
parsed,
Err(ServerError::TooManyBatchReconstructionFileIds)
));
}
#[test]
fn batch_reconstruction_parser_rejects_oversized_query_before_scanning() {
let mut query = String::from("ignored=");
query.push_str(&"a".repeat(MAX_BATCH_RECONSTRUCTION_QUERY_BYTES + 1));
let parsed = parse_batch_reconstruction_query(&query);
assert!(matches!(parsed, Err(ServerError::RequestQueryTooLarge)));
}
#[test]
fn provider_path_name_rejects_empty_or_oversized_values() {
let empty = validate_provider_name_path("");
let oversized = validate_provider_name_path(&"p".repeat(MAX_PROVIDER_NAME_BYTES + 1));
let valid = validate_provider_name_path("github");
assert!(matches!(
empty,
Err(ServerError::InvalidProviderTokenRequest)
));
assert!(matches!(
oversized,
Err(ServerError::InvalidProviderTokenRequest)
));
assert!(valid.is_ok());
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn security_headers_middleware_adds_xss_protection_headers() {
async fn handler() -> &'static str {
"ok"
}
let app = Router::new()
.route("/test", get(handler))
.layer(middleware::from_fn(security_headers_middleware));
let response = app
.oneshot(Request::builder().uri("/test").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let headers = response.headers();
assert_eq!(
headers.get(header::X_CONTENT_TYPE_OPTIONS).unwrap(),
"nosniff"
);
assert_eq!(headers.get(header::X_FRAME_OPTIONS).unwrap(), "DENY");
assert_eq!(
headers.get(header::STRICT_TRANSPORT_SECURITY).unwrap(),
"max-age=31536000"
);
assert_eq!(
headers.get(header::REFERRER_POLICY).unwrap(),
"strict-origin-when-cross-origin"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn security_headers_middleware_does_not_overwrite_existing_headers() {
async fn handler() -> &'static str {
"ok"
}
let app = Router::new()
.route("/test", get(handler))
.layer(middleware::from_fn(security_headers_middleware));
let response = app
.oneshot(
Request::builder()
.uri("/test")
.header(header::X_CONTENT_TYPE_OPTIONS, "custom")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
let hdr = response.headers().get(header::X_CONTENT_TYPE_OPTIONS);
assert!(hdr.is_some(), "header should be present");
let val = hdr.unwrap().to_str().unwrap_or("");
assert!(
val == "custom" || val == "nosniff",
"expected 'custom' or 'nosniff', got '{val}'"
);
}
async fn build_test_router(frontends: &[ServerFrontend], role: ServerRole) -> (Router, TempDir) {
let tmp = TempDir::new().unwrap();
let chunk_size = NonZeroUsize::new(65536).unwrap();
let config = ServerConfig::new(
"127.0.0.1:0".parse().unwrap(),
"http://127.0.0.1:8080".to_owned(),
tmp.path().to_path_buf(),
chunk_size,
)
.with_server_frontends(frontends.to_vec())
.unwrap()
.with_server_role(role)
.with_token_signing_key(vec![0u8; 32])
.unwrap();
let app = router(config).await;
(app.unwrap(), tmp)
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_builds_with_xet_frontend() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Xet], ServerRole::All).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/healthz")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/readyz")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/metrics")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert!(resp.status() == StatusCode::OK || resp.status() == StatusCode::NOT_FOUND);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_xet_api_routes_are_registered() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Xet], ServerRole::All).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/reconstructions")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert!(
resp.status() == StatusCode::UNAUTHORIZED
|| resp.status() == StatusCode::BAD_REQUEST
|| resp.status() == StatusCode::OK
|| resp.status() == StatusCode::METHOD_NOT_ALLOWED
);
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/stats")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_xet_transfer_routes_are_registered() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Xet], ServerRole::All).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/chunks/default/aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert!(resp.status() == StatusCode::NOT_FOUND || resp.status() == StatusCode::UNAUTHORIZED);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_lfs_routes_are_registered() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Lfs], ServerRole::All).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/lfs/objects/batch")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_ne!(resp.status(), StatusCode::NOT_FOUND);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_bazel_routes_are_registered() {
let (app, _tmp) = build_test_router(&[ServerFrontend::BazelHttp], ServerRole::All).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/bazel/cache/ac/0000000000000000000000000000000000000000000000000000000000000000")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert!(
resp.status() == StatusCode::NOT_FOUND
|| resp.status() == StatusCode::FORBIDDEN
|| resp.status() == StatusCode::UNAUTHORIZED
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_oci_routes_are_registered() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Oci], ServerRole::All).await;
let resp = app
.clone()
.oneshot(Request::builder().uri("/v2/").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_oci_api_role_has_v2_token_and_root() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Oci], ServerRole::Api).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v2/token")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert!(
resp.status() == StatusCode::UNAUTHORIZED
|| resp.status() == StatusCode::METHOD_NOT_ALLOWED
);
let resp = app
.clone()
.oneshot(Request::builder().uri("/v2/").body(Body::empty()).unwrap())
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::UNAUTHORIZED);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_oci_transfer_role_has_v2_catch_all() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Oci], ServerRole::Transfer).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v2/some/path")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::NOT_FOUND);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_api_role_excludes_transfer_routes() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Xet], ServerRole::Api).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/chunks/default/aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::NOT_FOUND);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_transfer_role_excludes_api_routes() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Xet], ServerRole::Transfer).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/stats")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::NOT_FOUND);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_responses_include_security_headers() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Xet], ServerRole::All).await;
let response = app
.oneshot(
Request::builder()
.uri("/healthz")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
let headers = response.headers();
assert_eq!(
headers.get(header::X_CONTENT_TYPE_OPTIONS).unwrap(),
"nosniff"
);
assert_eq!(headers.get(header::X_FRAME_OPTIONS).unwrap(), "DENY");
}
#[test]
fn bounded_api_body_limit_with_zero_endpoint_limit() {
let configured = NonZeroUsize::new(1024).unwrap();
let result = bounded_api_body_limit(configured, 0);
assert_eq!(result, 0);
}
#[test]
fn bounded_api_body_limit_with_equal_values() {
let val = NonZeroUsize::new(8192).unwrap();
let result = bounded_api_body_limit(val, 8192);
assert_eq!(result, 8192);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn auth_provider_local_without_signing_key_returns_none() {
let tmp = TempDir::new().unwrap();
let chunk_size = NonZeroUsize::new(65536).unwrap();
let config = ServerConfig::new(
"127.0.0.1:0".parse().unwrap(),
"http://127.0.0.1:8080".to_owned(),
tmp.path().to_path_buf(),
chunk_size,
)
.with_auth_provider(AuthProviderKind::Local);
let app = router(config).await;
assert!(
matches!(
app.err().unwrap(),
ServerError::Config(
crate::config::ServerConfigError::MissingTokenSigningKeyForServedRoutes
)
),
"should fail with MissingTokenSigningKeyForServedRoutes"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn auth_provider_passthrough_builds_successfully() {
let tmp = TempDir::new().unwrap();
let chunk_size = NonZeroUsize::new(65536).unwrap();
let config = ServerConfig::new(
"127.0.0.1:0".parse().unwrap(),
"http://127.0.0.1:8080".to_owned(),
tmp.path().to_path_buf(),
chunk_size,
)
.with_auth_provider(AuthProviderKind::Passthrough)
.with_token_signing_key(vec![0u8; 32])
.unwrap();
let app = router(config).await;
assert!(
app.is_ok(),
"router should build with Passthrough auth, got: {:?}",
app.err()
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn auth_provider_oidc_with_unreachable_url_errors() {
let tmp = TempDir::new().unwrap();
let chunk_size = NonZeroUsize::new(65536).unwrap();
let config = ServerConfig::new(
"127.0.0.1:0".parse().unwrap(),
"http://127.0.0.1:8080".to_owned(),
tmp.path().to_path_buf(),
chunk_size,
)
.with_auth_provider(AuthProviderKind::Oidc)
.with_token_signing_key(vec![0u8; 32])
.unwrap()
.with_auth_oidc_issuer("http://127.0.0.1:1/not-exist".to_owned());
let app = router(config).await;
assert!(
app.is_err(),
"Oidc with unreachable issuer should fail to build router"
);
assert!(
matches!(app.err().unwrap(), ServerError::Config(_)),
"error should be ServerError::Config"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn auth_provider_jwks_with_unreachable_url_errors() {
let tmp = TempDir::new().unwrap();
let chunk_size = NonZeroUsize::new(65536).unwrap();
let config = ServerConfig::new(
"127.0.0.1:0".parse().unwrap(),
"http://127.0.0.1:8080".to_owned(),
tmp.path().to_path_buf(),
chunk_size,
)
.with_auth_provider(AuthProviderKind::Jwks)
.with_token_signing_key(vec![0u8; 32])
.unwrap()
.with_auth_jwks_url("http://127.0.0.1:1/not-exist".to_owned());
let app = router(config).await;
assert!(
app.is_err(),
"Jwks with unreachable URL should fail to build router"
);
assert!(
matches!(app.err().unwrap(), ServerError::Config(_)),
"error should be ServerError::Config"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn hub_frontend_builds_router_successfully() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Hub], ServerRole::All).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/healthz")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn hub_state_is_none_without_hub_frontend() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Xet], ServerRole::All).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/healthz")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn serve_accepts_valid_config_and_fails_on_bind_conflict() {
let tmp = TempDir::new().unwrap();
let chunk_size = NonZeroUsize::new(65536).unwrap();
let config = ServerConfig::new(
"127.0.0.1:0".parse().unwrap(),
"http://127.0.0.1:8080".to_owned(),
tmp.path().to_path_buf(),
chunk_size,
)
.with_auth_provider(AuthProviderKind::Local)
.with_token_signing_key(vec![0u8; 32])
.unwrap();
let app = router(config).await;
assert!(
app.is_ok(),
"router should build successfully, got: {:?}",
app.err()
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn shutdown_timeout_starts_after_the_shutdown_signal() {
let tmp = TempDir::new().unwrap();
let listener = TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let config = ServerConfig::new(
addr,
format!("http://{addr}"),
tmp.path().to_path_buf(),
NonZeroUsize::new(4096).unwrap(),
)
.with_token_signing_key(vec![0_u8; 32])
.unwrap()
.with_shutdown_timeout(Duration::from_millis(40));
let (shutdown_tx, shutdown_rx) = oneshot::channel();
let server = tokio::spawn(serve_with_listener_until(config, listener, async move {
let _ignored = shutdown_rx.await;
}));
let client = reqwest::Client::new();
let mut became_healthy = false;
for _attempt in 0..20 {
if let Ok(response) = client.get(format!("http://{addr}/healthz")).send().await
&& response.status() == StatusCode::OK
{
became_healthy = true;
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
assert!(
became_healthy,
"server should become healthy before shutdown"
);
tokio::time::sleep(Duration::from_millis(80)).await;
let response = client.get(format!("http://{addr}/healthz")).send().await;
assert!(response.is_ok(), "server must not time out before shutdown");
let Ok(response) = response else {
return;
};
assert_eq!(response.status(), StatusCode::OK);
let _ignored = shutdown_tx.send(());
let result = timeout(Duration::from_secs(1), server).await;
assert!(
result.is_ok(),
"server should drain after the shutdown signal"
);
let Ok(result) = result else {
return;
};
assert!(result.is_ok(), "server task should not panic");
let Ok(result) = result else {
return;
};
assert!(result.is_ok(), "server should exit cleanly: {result:?}");
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_xet_api_only_role_registers_only_api_routes() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Xet], ServerRole::Api).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/reconstructions")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_ne!(resp.status(), StatusCode::NOT_FOUND);
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/chunks/default/aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::NOT_FOUND);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_xet_transfer_only_role_registers_only_transfer_routes() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Xet], ServerRole::Transfer).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/chunks/default/aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_ne!(resp.status(), StatusCode::NOT_FOUND);
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/stats")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::NOT_FOUND);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_lfs_api_only_role_excludes_transfer_routes() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Lfs], ServerRole::Api).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/lfs/objects/abc")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::NOT_FOUND);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_lfs_transfer_only_role_excludes_api_routes() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Lfs], ServerRole::Transfer).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.method("POST")
.uri("/v1/lfs/objects/batch")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert!(
resp.status() == StatusCode::NOT_FOUND || resp.status() == StatusCode::METHOD_NOT_ALLOWED,
"expected 404 or 405, got {}",
resp.status()
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_bazel_transfer_only_role_registers_transfer_routes() {
let (app, _tmp) = build_test_router(&[ServerFrontend::BazelHttp], ServerRole::Transfer).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/bazel/cache/ac/0000000000000000000000000000000000000000000000000000000000000000")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_ne!(resp.status(), StatusCode::NOT_FOUND);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_bazel_api_only_role_has_no_routes() {
let (app, _tmp) = build_test_router(&[ServerFrontend::BazelHttp], ServerRole::Api).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v1/bazel/cache/ac/0000000000000000000000000000000000000000000000000000000000000000")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::NOT_FOUND);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn hub_frontend_builds_router_with_xet_frontend() {
let (app, _tmp) =
build_test_router(&[ServerFrontend::Hub, ServerFrontend::Xet], ServerRole::All).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/healthz")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_eq!(resp.status(), StatusCode::OK);
}
#[test]
fn endpoint_body_limit_with_zero_config_returns_overflow() {
use super::endpoint_body_limit;
use std::num::NonZeroUsize;
let result = endpoint_body_limit(NonZeroUsize::new(0).unwrap_or(NonZeroUsize::MIN), 0);
assert!(matches!(result, Err(ServerError::Overflow)));
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn router_oci_all_role_has_all_routes() {
let (app, _tmp) = build_test_router(&[ServerFrontend::Oci], ServerRole::All).await;
let resp = app
.clone()
.oneshot(
Request::builder()
.uri("/v2/token")
.body(Body::empty())
.unwrap(),
)
.await
.unwrap();
assert_ne!(resp.status(), StatusCode::NOT_FOUND);
let resp = app
.clone()
.oneshot(Request::builder().uri("/v2/").body(Body::empty()).unwrap())
.await
.unwrap();
assert_ne!(resp.status(), StatusCode::NOT_FOUND);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn authorize_with_no_auth_returns_ok_none() {
use crate::ServerConfig;
use crate::config::AuthProviderKind;
use axum::http::HeaderMap;
use shardline_protocol::TokenScope;
let tmp = TempDir::new().unwrap();
let config = ServerConfig::new(
"127.0.0.1:0".parse().unwrap(),
"http://127.0.0.1:8080".to_owned(),
tmp.path().to_path_buf(),
NonZeroUsize::new(65536).unwrap(),
)
.with_auth_provider(AuthProviderKind::Local);
let state = Arc::new(crate::AppState {
config,
role: ServerRole::All,
backend: crate::ServerBackend::Local(
crate::LocalBackend::new(
tmp.path().to_path_buf(),
"http://127.0.0.1:8080".to_owned(),
NonZeroUsize::new(65536).unwrap(),
)
.await
.unwrap(),
),
auth: None,
provider_tokens: None,
reconstruction_cache: crate::ReconstructionCacheService::disabled(),
transfer_limiter: crate::TransferLimiter::new(
NonZeroUsize::new(65536).unwrap(),
NonZeroUsize::new(4).unwrap(),
),
oci_registry_token_limiter: Arc::new(tokio::sync::Semaphore::new(8)),
protocol_metrics: crate::ProtocolMetrics::default(),
});
let result = super::authorize(&state, &HeaderMap::new(), TokenScope::Read);
assert!(result.is_ok());
assert!(result.unwrap().is_none());
}
#[tokio::test]
async fn acquire_chunk_transfer_permit_times_out_when_permits_exhausted() {
tokio::time::pause();
let tmp = TempDir::new().unwrap();
let chunk_size = NonZeroUsize::new(65536).unwrap();
let hash = "aa".repeat(32);
let prefix = &hash[..2];
let chunk_dir = tmp.path().join("chunks").join(prefix);
std::fs::create_dir_all(&chunk_dir).unwrap();
std::fs::write(chunk_dir.join(&hash), b"some chunk data").unwrap();
let backend = crate::LocalBackend::new(
tmp.path().to_path_buf(),
"http://127.0.0.1:8080".to_owned(),
chunk_size,
)
.await
.unwrap();
let max_in_flight = NonZeroUsize::new(1).unwrap();
let transfer_limiter = crate::TransferLimiter::new(chunk_size, max_in_flight)
.with_acquire_timeout(std::time::Duration::from_millis(50));
let state = Arc::new(crate::AppState {
config: crate::ServerConfig::new(
"127.0.0.1:0".parse().unwrap(),
"http://127.0.0.1:8080".to_owned(),
tmp.path().to_path_buf(),
chunk_size,
),
role: ServerRole::All,
backend: crate::ServerBackend::Local(backend),
auth: None,
provider_tokens: None,
reconstruction_cache: crate::ReconstructionCacheService::disabled(),
transfer_limiter,
oci_registry_token_limiter: Arc::new(tokio::sync::Semaphore::new(8)),
protocol_metrics: crate::ProtocolMetrics::default(),
});
let _permit = state.transfer_limiter.acquire_bytes(4).await.unwrap();
let result = super::acquire_chunk_transfer_permit(&state, &hash).await;
assert!(
matches!(result, Err(ServerError::TransferLimiterTimedOut)),
"expected TransferLimiterTimedOut, got {result:?}"
);
}