Skip to main content

uv_configuration/
overrides.rs

1use std::borrow::Cow;
2
3use either::Either;
4use rustc_hash::{FxBuildHasher, FxHashMap, FxHashSet};
5use serde::de::IntoDeserializer;
6
7use uv_distribution_types::{Requirement, RequirementSource};
8use uv_normalize::PackageName;
9use uv_pep440::Version;
10use uv_pep508::MarkerTree;
11
12/// An override that applies to the dependencies of a specific package version.
13#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize)]
14#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
15#[serde(
16    rename_all = "kebab-case",
17    deny_unknown_fields,
18    bound(
19        serialize = "T: serde::Serialize",
20        deserialize = "T: serde::Deserialize<'de>"
21    )
22)]
23pub struct PackageOverride<T> {
24    pub package: PackageOverrideTarget,
25    pub dependencies: Box<[T]>,
26}
27
28/// The package and optional version selected by a [`PackageOverride`].
29#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize)]
30#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema))]
31#[serde(rename_all = "kebab-case", deny_unknown_fields)]
32pub struct PackageOverrideTarget {
33    name: PackageName,
34    #[cfg_attr(
35        feature = "schemars",
36        schemars(
37            with = "Option<String>",
38            description = "PEP 440-style package version, e.g., `1.2.3`"
39        )
40    )]
41    version: Option<Version>,
42}
43
44/// An override, either global or scoped to a specific package version.
45#[derive(Debug, Clone, PartialEq, Eq, PartialOrd, Ord, serde::Serialize)]
46#[cfg_attr(feature = "schemars", derive(schemars::JsonSchema), schemars(untagged))]
47#[serde(untagged, bound(serialize = "T: serde::Serialize"))]
48pub enum Override<T> {
49    Package(PackageOverride<T>),
50    Requirement(T),
51}
52
53// A derived `#[serde(untagged)]` implementation collapses detailed requirement parse errors into
54// "data did not match any variant", so use a type-directed visitor for string requirements.
55impl<'de, T> serde::Deserialize<'de> for Override<T>
56where
57    T: serde::Deserialize<'de>,
58{
59    fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
60    where
61        D: serde::Deserializer<'de>,
62    {
63        #[derive(serde::Deserialize)]
64        #[serde(untagged)]
65        enum MapOverride<T> {
66            Package(PackageOverride<T>),
67            Requirement(T),
68        }
69
70        serde_untagged::UntaggedEnumVisitor::new()
71            .string(|string| T::deserialize(string.into_deserializer()).map(Self::Requirement))
72            .map(|map| {
73                map.deserialize::<MapOverride<T>>()
74                    .map(|entry| match entry {
75                        MapOverride::Package(package) => Self::Package(package),
76                        MapOverride::Requirement(requirement) => Self::Requirement(requirement),
77                    })
78            })
79            .deserialize(deserializer)
80    }
81}
82
83/// A set of overrides for a set of requirements.
84#[derive(Debug, Default, Clone)]
85pub struct Overrides {
86    global: FxHashMap<PackageName, Vec<Requirement>>,
87    scoped: FxHashMap<PackageName, Vec<ScopedOverrides>>,
88}
89
90#[derive(Debug, Clone)]
91struct ScopedOverrides {
92    version: Option<Version>,
93    overrides: FxHashMap<PackageName, Vec<Requirement>>,
94}
95
96/// An unsupported source in a scoped dependency override.
97#[derive(Debug, thiserror::Error)]
98pub enum ScopedOverrideSourceError {
99    #[error(
100        "Scoped override for `{package}` cannot use a URL or path source for `{dependency}`; scoped overrides currently support version specifiers only"
101    )]
102    Url {
103        package: PackageName,
104        dependency: PackageName,
105    },
106    #[error(
107        "Scoped override for `{package}` cannot use an explicit index for `{dependency}`; scoped overrides currently support version specifiers only"
108    )]
109    Index {
110        package: PackageName,
111        dependency: PackageName,
112    },
113}
114
115impl Overrides {
116    /// Create a new set of overrides from a set of requirements.
117    pub fn from_requirements(requirements: Vec<Requirement>) -> Self {
118        let mut global: FxHashMap<PackageName, Vec<Requirement>> =
119            FxHashMap::with_capacity_and_hasher(requirements.len(), FxBuildHasher);
120        for requirement in requirements {
121            global
122                .entry(requirement.name.clone())
123                .or_default()
124                .push(requirement);
125        }
126        Self {
127            global,
128            scoped: FxHashMap::default(),
129        }
130    }
131
132    /// Create an indexed set of overrides.
133    pub fn from_entries(
134        entries: Vec<Override<Requirement>>,
135    ) -> Result<Self, ScopedOverrideSourceError> {
136        let mut global: FxHashMap<PackageName, Vec<Requirement>> =
137            FxHashMap::with_capacity_and_hasher(entries.len(), FxBuildHasher);
138        let mut scoped: FxHashMap<PackageName, Vec<ScopedOverrides>> = FxHashMap::default();
139
140        for entry in entries {
141            match entry {
142                Override::Requirement(requirement) => {
143                    global
144                        .entry(requirement.name.clone())
145                        .or_default()
146                        .push(requirement);
147                }
148                Override::Package(package) => {
149                    for requirement in &package.dependencies {
150                        match &requirement.source {
151                            RequirementSource::Registry { index: Some(_), .. } => {
152                                return Err(ScopedOverrideSourceError::Index {
153                                    package: package.package.name.clone(),
154                                    dependency: requirement.name.clone(),
155                                });
156                            }
157                            RequirementSource::Registry { index: None, .. } => {}
158                            RequirementSource::Url { .. }
159                            | RequirementSource::GitDirectory { .. }
160                            | RequirementSource::GitPath { .. }
161                            | RequirementSource::Path { .. }
162                            | RequirementSource::Directory { .. } => {
163                                return Err(ScopedOverrideSourceError::Url {
164                                    package: package.package.name.clone(),
165                                    dependency: requirement.name.clone(),
166                                });
167                            }
168                        }
169                    }
170                    let packages = scoped.entry(package.package.name.clone()).or_default();
171                    let position = packages
172                        .iter()
173                        .position(|overrides| overrides.version == package.package.version)
174                        .unwrap_or_else(|| {
175                            let position = packages.len();
176                            packages.push(ScopedOverrides {
177                                version: package.package.version,
178                                overrides: FxHashMap::default(),
179                            });
180                            position
181                        });
182                    let overrides = &mut packages[position].overrides;
183                    for requirement in package.dependencies {
184                        overrides
185                            .entry(requirement.name.clone())
186                            .or_default()
187                            .push(requirement);
188                    }
189                }
190            }
191        }
192
193        Ok(Self { global, scoped })
194    }
195
196    /// Return an iterator over all global [`Requirement`]s in the override set.
197    pub fn global_requirements(&self) -> impl Iterator<Item = &Requirement> {
198        self.global
199            .values()
200            .flat_map(|requirements| requirements.iter())
201    }
202
203    /// Return all scoped [`Requirement`]s with the package and version they apply to.
204    pub fn scoped_requirements(
205        &self,
206    ) -> impl Iterator<Item = (&PackageName, Option<&Version>, &Requirement)> {
207        self.scoped.iter().flat_map(|(package, entries)| {
208            entries.iter().flat_map(move |entry| {
209                entry
210                    .overrides
211                    .values()
212                    .flatten()
213                    .map(move |requirement| (package, entry.version.as_ref(), requirement))
214            })
215        })
216    }
217
218    /// Return the scoped [`Requirement`]s that apply to a specific package version.
219    pub fn scoped_requirements_for(
220        &self,
221        package: &PackageName,
222        version: &Version,
223    ) -> impl Iterator<Item = &Requirement> {
224        self.scoped_for(package, version)
225            .into_iter()
226            .flat_map(|scoped| scoped.overrides.values().flatten())
227    }
228
229    /// Return whether a package has overrides for an exact version.
230    pub(crate) fn has_exact_scope(&self, package: &PackageName, version: &Version) -> bool {
231        self.scoped.get(package).is_some_and(|entries| {
232            entries
233                .iter()
234                .any(|entry| entry.version.as_ref() == Some(version))
235        })
236    }
237
238    /// Get the overrides for a package.
239    fn get(&self, name: &PackageName) -> Option<&Vec<Requirement>> {
240        self.global.get(name)
241    }
242
243    /// Get the overrides for a specific package version.
244    fn scoped_for(&self, package: &PackageName, version: &Version) -> Option<&ScopedOverrides> {
245        self.scoped.get(package).and_then(|entries| {
246            entries
247                .iter()
248                .find(|entry| entry.version.as_ref() == Some(version))
249                .or_else(|| entries.iter().find(|entry| entry.version.is_none()))
250        })
251    }
252
253    /// Apply the overrides to a set of requirements.
254    ///
255    /// NB: Change this method together with [`Constraints::apply`].
256    pub fn apply<'a, I>(
257        &'a self,
258        requirements: I,
259    ) -> impl Iterator<Item = Cow<'a, Requirement>> + use<'a, I>
260    where
261        I: IntoIterator<Item = &'a Requirement>,
262    {
263        self.apply_inner(requirements, None)
264    }
265
266    /// Apply the overrides to the dependencies of a specific package version.
267    pub fn apply_for<'a, I>(
268        &'a self,
269        package: &PackageName,
270        version: &Version,
271        requirements: I,
272    ) -> impl Iterator<Item = Cow<'a, Requirement>> + use<'a, I>
273    where
274        I: IntoIterator<Item = &'a Requirement>,
275    {
276        self.apply_inner(requirements, Some((package, version)))
277    }
278
279    /// Apply overrides with optional package-version context.
280    pub fn apply_for_package<'a, I>(
281        &'a self,
282        package: Option<(&PackageName, &Version)>,
283        requirements: I,
284    ) -> impl Iterator<Item = Cow<'a, Requirement>> + use<'a, I>
285    where
286        I: IntoIterator<Item = &'a Requirement>,
287    {
288        self.apply_inner(requirements, package)
289    }
290
291    fn apply_inner<'a, I>(
292        &'a self,
293        requirements: I,
294        package: Option<(&PackageName, &Version)>,
295    ) -> impl Iterator<Item = Cow<'a, Requirement>> + use<'a, I>
296    where
297        I: IntoIterator<Item = &'a Requirement>,
298    {
299        let scoped = package.and_then(|(package, version)| self.scoped_for(package, version));
300        if let Some(scoped) = scoped {
301            let requirements = requirements.into_iter().collect::<Vec<_>>();
302            let names = requirements
303                .iter()
304                .map(|requirement| requirement.name.clone())
305                .collect::<FxHashSet<_>>();
306            let mut additions = scoped
307                .overrides
308                .iter()
309                .filter(|(name, _)| !names.contains(*name))
310                .flat_map(|(_, requirements)| requirements)
311                .collect::<Vec<_>>();
312            additions.sort_unstable();
313
314            return Either::Left(
315                requirements
316                    .into_iter()
317                    .flat_map(move |requirement| self.apply_requirement(requirement, Some(scoped)))
318                    .chain(additions.into_iter().map(Cow::Borrowed)),
319            );
320        }
321
322        if self.global.is_empty() {
323            // Fast path: There are no overrides.
324            return Either::Right(Either::Left(requirements.into_iter().map(Cow::Borrowed)));
325        }
326
327        Either::Right(Either::Right(requirements.into_iter().flat_map(
328            move |requirement| self.apply_requirement(requirement, None),
329        )))
330    }
331
332    fn apply_requirement<'a>(
333        &'a self,
334        requirement: &'a Requirement,
335        scoped: Option<&'a ScopedOverrides>,
336    ) -> impl Iterator<Item = Cow<'a, Requirement>> {
337        let overrides = scoped
338            .and_then(|scoped| scoped.overrides.get(&requirement.name))
339            .or_else(|| self.get(&requirement.name));
340        let Some(overrides) = overrides else {
341            // Case 1: No override(s).
342            return Either::Left(std::iter::once(Cow::Borrowed(requirement)));
343        };
344
345        // ASSUMPTION: There is one `extra = "..."`, and it's either the only marker or part
346        // of the main conjunction.
347        let Some(extra_expression) = requirement.marker.top_level_extra() else {
348            // Case 2: A non-optional dependency with override(s).
349            return Either::Right(Either::Right(overrides.iter().map(Cow::Borrowed)));
350        };
351
352        // Case 3: An optional dependency with override(s).
353        //
354        // When the original requirement is an optional dependency, the override(s) need to
355        // be optional for the same extra, otherwise we activate extras that should be inactive.
356        Either::Right(Either::Left(overrides.iter().map(
357            move |override_requirement| {
358                // Add the extra to the override marker.
359                let joint_marker = MarkerTree::expression(extra_expression.clone())
360                    .and(override_requirement.marker);
361                Cow::Owned(Requirement {
362                    marker: joint_marker,
363                    ..override_requirement.clone()
364                })
365            },
366        )))
367    }
368}