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#[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#[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#[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
53impl<'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#[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#[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 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 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 pub fn global_requirements(&self) -> impl Iterator<Item = &Requirement> {
198 self.global
199 .values()
200 .flat_map(|requirements| requirements.iter())
201 }
202
203 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 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 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 fn get(&self, name: &PackageName) -> Option<&Vec<Requirement>> {
240 self.global.get(name)
241 }
242
243 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 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 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 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 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 return Either::Left(std::iter::once(Cow::Borrowed(requirement)));
343 };
344
345 let Some(extra_expression) = requirement.marker.top_level_extra() else {
348 return Either::Right(Either::Right(overrides.iter().map(Cow::Borrowed)));
350 };
351
352 Either::Right(Either::Left(overrides.iter().map(
357 move |override_requirement| {
358 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}