use std::{collections::VecDeque, sync::Arc};
use futures::StreamExt;
use futures::stream::FuturesUnordered;
use rustc_hash::FxHashSet;
use tracing::trace;
use uv_configuration::{Constraints, Overrides};
use uv_distribution::{DistributionDatabase, Reporter};
use uv_distribution_types::{Dist, Identifier, Requirement, RequirementSource};
use uv_resolver::{InMemoryIndex, MetadataResponse, ResolverEnvironment};
use uv_types::{BuildContext, HashStrategy, RequestedRequirements};
use crate::{Error, required_dist};
pub struct LookaheadResolver<'a, Context: BuildContext> {
requirements: &'a [Requirement],
constraints: &'a Constraints,
overrides: &'a Overrides,
hasher: &'a HashStrategy,
index: &'a InMemoryIndex,
database: DistributionDatabase<'a, Context>,
}
impl<'a, Context: BuildContext> LookaheadResolver<'a, Context> {
pub fn new(
requirements: &'a [Requirement],
constraints: &'a Constraints,
overrides: &'a Overrides,
hasher: &'a HashStrategy,
index: &'a InMemoryIndex,
database: DistributionDatabase<'a, Context>,
) -> Self {
Self {
requirements,
constraints,
overrides,
hasher,
index,
database,
}
}
#[must_use]
pub fn with_reporter(self, reporter: Arc<dyn Reporter>) -> Self {
Self {
database: self.database.with_reporter(reporter),
..self
}
}
pub async fn resolve(
self,
env: &ResolverEnvironment,
) -> Result<Vec<RequestedRequirements>, Error> {
let mut results = Vec::new();
let mut futures = FuturesUnordered::new();
let mut seen = FxHashSet::default();
let mut queue: VecDeque<_> = self
.constraints
.apply(self.overrides.apply(self.requirements))
.filter(|requirement| requirement.evaluate_markers(env.marker_environment(), &[]))
.map(|requirement| (*requirement).clone())
.collect();
while !queue.is_empty() || !futures.is_empty() {
while let Some(requirement) = queue.pop_front() {
if !matches!(requirement.source, RequirementSource::Registry { .. }) {
if seen.insert(requirement.clone()) {
futures.push(self.lookahead(requirement));
}
}
}
while let Some(result) = futures.next().await {
if let Some(lookahead) = result? {
for requirement in self
.constraints
.apply(self.overrides.apply(lookahead.requirements()))
{
if requirement
.evaluate_markers(env.marker_environment(), lookahead.extras())
{
queue.push_back((*requirement).clone());
}
}
results.push(lookahead);
}
}
}
Ok(results)
}
async fn lookahead(
&self,
requirement: Requirement,
) -> Result<Option<RequestedRequirements>, Error> {
trace!("Performing lookahead for {requirement}");
let Some(dist) = required_dist(&requirement)? else {
return Ok(None);
};
let direct = if let Dist::Source(source_dist) = &dist {
source_dist.as_path().is_some_and(std::path::Path::is_dir)
} else {
false
};
let metadata = {
let id = dist.distribution_id();
if self.index.distributions().register(id.clone()) {
let archive = self
.database
.get_or_build_wheel_metadata(&dist, self.hasher.get(&dist))
.await
.map_err(|err| Error::from_dist(dist, err))?;
let metadata = archive.metadata.clone();
self.index
.distributions()
.done(id, Arc::new(MetadataResponse::Found(archive)));
metadata
} else {
let response = self
.index
.distributions()
.wait(&id)
.await
.expect("missing value for registered task");
let MetadataResponse::Found(archive) = &*response else {
panic!("Failed to find metadata for: {requirement}");
};
archive.metadata.clone()
}
};
let requires_dist = Box::into_iter(metadata.requires_dist)
.chain(
metadata
.dependency_groups
.into_iter()
.filter_map(|(group, dependencies)| {
if requirement.groups.contains(&group) {
Some(dependencies)
} else {
None
}
})
.flatten(),
)
.map(|dependency| {
if dependency.name == requirement.name {
Requirement {
source: requirement.source.clone(),
..dependency
}
} else {
dependency
}
})
.collect();
Ok(Some(RequestedRequirements::new(
requirement.extras,
requires_dist,
direct,
)))
}
}