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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
77 changes: 32 additions & 45 deletions lychee-bin/src/commands/check.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,7 +27,7 @@ use crate::formatters::suggestion::Suggestion;
use crate::progress::Progress;
use crate::{ExitCode, cache::Cache};

type RecursiveRequest = (WaitGuard, Request, usize);
type RecursiveRequest = (WaitGuard, Result<Request, RequestError>, usize);

#[derive(Clone)]
struct Recursion {
Expand Down Expand Up @@ -142,31 +142,12 @@ pub(crate) async fn check(

let (waiter, wait_guard) = WaitGroup::new();

// Split initial requests into: valid requests and request errors. Note that
// this stream closure *owns* a wait guard, so we must drop the closure after
// it's finished to avoid deadlock. This is done using the `.chain()` combinator.
let (valid_requests, request_errors) = requests
// The initial stream owns a wait guard. Drop its closure when it finishes
// so that the recursive stream can terminate once all work is complete.
let initial_requests = requests
.inspect(|_| progress.inc_length(1))
.map(move |request| (request, wait_guard.clone()))
.chain(futures::stream::empty())
.map(|(request, guard)| match request {
Ok(request) => Ok((guard, request, 0)),
Err(request_error) => Err((guard, request_error)),
})
.partition_result::<(WaitGuard, Request, usize), (WaitGuard, RequestError)>();

// Further partition the request errors into request building errors (like
// unresolved relative URLs) and fatal errors when fetching a user input fails.
let (request_building_errors, mut fatal_errors) = request_errors
.map(
|(guard, request_error)| match request_error.into_response() {
Ok(request_building_error) => Ok((guard, request_building_error)),
Err(fatal_user_input_error) => Err((guard, fatal_user_input_error)),
},
)
.partition_result::<(WaitGuard, Response), (WaitGuard, ErrorKind)>();
let request_building_errors =
request_building_errors.map(|(guard, response)| (guard, response, 0));
.map(move |request| (wait_guard.clone(), request, 0))
.chain(futures::stream::empty());

let (recursive_channel_send, recursive_channel_recv) = mpsc::channel(max_concurrency);

Expand All @@ -176,17 +157,34 @@ pub(crate) async fn check(
async move { send_recursive_request(recursive_channel_send, &progress, request) }
};

// Combine recursive requests and input requests.
// Initial and recursive collection results share the same error handling.
let requests = futures::stream::select_with_strategy(
valid_requests,
initial_requests,
ReceiverStream::new(recursive_channel_recv).take_until(waiter.wait()),
|()| futures::stream::PollNext::Right, // Recursive requests consume memory, prefer those.
);
let (valid_requests, request_errors) = requests
.map(|(guard, request, depth)| match request {
Ok(request) => Ok((guard, request, depth)),
Err(request_error) => Err((guard, request_error, depth)),
})
.partition_result::<(WaitGuard, Request, usize), (WaitGuard, RequestError, usize)>();

// Further partition the request errors into request building errors (like
// unresolved relative URLs) and fatal errors when fetching input content fails.
let (request_building_errors, mut fatal_errors) = request_errors
.map(
|(guard, request_error, depth)| match request_error.into_response() {
Ok(request_building_error) => Ok((guard, request_building_error, depth)),
Err(fatal_input_error) => Err((guard, fatal_input_error)),
},
)
.partition_result::<(WaitGuard, Response, usize), (WaitGuard, ErrorKind)>();

/* Main link checking pipeline */

// Perform requests. This is the only part of the main pipeline that happens concurrently.
let check_responses = requests
let check_responses = valid_requests
.map(
async |(guard, request, depth)| -> (WaitGuard, Response, usize) {
let check_url = |r| check_url(&client, r);
Expand Down Expand Up @@ -229,12 +227,12 @@ pub(crate) async fn check(
// WARNING: Before changing the `.await` structure, be aware of the
// requirements imposed by [`lychee_lib::async_lib::stream::partition_result`].
// Partitioned streams must be concurrently polled. At the moment, this is
// achieved because all partitioned streams are within `all_done`.
// achieved by polling `all_done` and `fatal_errors` concurrently below.
match futures::future::select(pin!(all_done), fatal_errors.next()).await {
Either::Left(((), _fatal_errors)) => (),
Either::Right((None, remaining)) => remaining.await,
Either::Right((Some((_guard, fatal_error)), _remaining)) => {
progress.finish("Error while fetching initial inputs");
progress.finish("Error while fetching inputs");
return Err(fatal_error);
}
}
Expand Down Expand Up @@ -324,23 +322,12 @@ async fn collect_recursive_requests(
collector: Collector,
url: Url,
extensions: FileExtensions,
) -> Vec<Request> {
let input = Input::from_input_source(InputSource::RemoteUrl(Box::new(url.clone())));
let results = collector
) -> Vec<Result<Request, RequestError>> {
let input = Input::from_input_source(InputSource::RemoteUrl(Box::new(url)));
collector
.collect_links_from_file_types(HashSet::from([input]), extensions)
.collect::<Vec<_>>()
.await;

results
.into_iter()
.filter_map(|result| match result {
Ok(request) => Some(request),
Err(error) => {
warn!("unable to collect links recursively from {url}: {error}");
None
}
})
.collect()
.await
}

async fn suggest_archived_links(
Expand Down
95 changes: 95 additions & 0 deletions lychee-bin/tests/cli.rs
Original file line number Diff line number Diff line change
Expand Up @@ -2738,6 +2738,101 @@ The config file should contain every possible key for documentation purposes."
.stdout(contains("0 Errors"));
}

#[tokio::test]
async fn test_recursive_link_checking_reports_collection_failure() {
use std::sync::{
Arc,
atomic::{AtomicUsize, Ordering},
};

let mock_server = wiremock::MockServer::start().await;
Mock::given(method("GET"))
.and(path("/index.html"))
.respond_with(
ResponseTemplate::new(200)
.set_body_raw(r#"<a href="/child.html">child</a>"#, "text/html"),
)
.mount(&mock_server)
.await;

let requests = Arc::new(AtomicUsize::new(0));
let child_requests = requests.clone();
Mock::given(method("GET"))
.and(path("/child.html"))
.respond_with(move |_: &wiremock::Request| {
// The link check succeeds, but fetching its contents to recurse fails.
let status = if child_requests.fetch_add(1, Ordering::SeqCst) == 0 {
200
} else {
503
};
ResponseTemplate::new(status).set_body_raw("", "text/html")
})
.mount(&mock_server)
.await;

let output = cargo_bin_cmd!()
.args([
"--recursive",
"--max-depth=1",
"--max-concurrency=1",
"--max-retries=0",
])
.arg(format!("{}/index.html", mock_server.uri()))
.timeout(Duration::from_secs(10))
.assert();

assert_eq!(requests.load(Ordering::SeqCst), 2);
output.code(1).stderr(contains("503"));
}

#[tokio::test]
async fn test_recursive_link_checking_reports_invalid_links() {
let mock_server = wiremock::MockServer::start().await;
// Exceed the error partition's buffer while also checking a valid link.
let invalid_links = (0..17)
.map(|index| format!(r#"<a href="http://[invalid{index}]/">invalid</a>"#))
.collect::<String>();
Mock::given(method("GET"))
.and(path("/index.html"))
.respond_with(
ResponseTemplate::new(200)
.set_body_raw(r#"<a href="/child.html">child</a>"#, "text/html"),
)
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/child.html"))
.respond_with(ResponseTemplate::new(200).set_body_raw(
format!(r#"{invalid_links}<a href="/ok.html">ok</a>"#),
"text/html",
))
.mount(&mock_server)
.await;
Mock::given(method("GET"))
.and(path("/ok.html"))
.respond_with(ResponseTemplate::new(200))
.expect(1)
.mount(&mock_server)
.await;

cargo_bin_cmd!()
.args([
"--recursive",
"--max-depth=1",
"--max-concurrency=1",
"--max-retries=0",
"--verbose",
])
.arg(format!("{}/index.html", mock_server.uri()))
.timeout(Duration::from_secs(10))
.assert()
.code(2)
.stdout(contains("Cannot parse 'http://[invalid").count(17))
.stdout(contains("19 Total"))
.stdout(contains("17 Errors"));
}

async fn external_recursive_test_servers()
-> (wiremock::MockServer, wiremock::MockServer, String) {
let root_server = wiremock::MockServer::start().await;
Expand Down