Avoid URL parsing when deserializing wheels (#5235)

## Summary

This PR slots in `UrlString` for `WheelWire`, which IIUC should avoid
parsing URLs during deserialization?
This commit is contained in:
Charlie Marsh
2024-07-19 22:11:26 -04:00
committed by GitHub
parent 833097b93f
commit 1243c5e28c
7 changed files with 71 additions and 104 deletions
+1 -1
View File
@@ -14,7 +14,7 @@ pub enum Error {
WheelFilename(#[from] distribution_filename::WheelFilenameError),
#[error("Could not extract path segments from URL: {0}")]
MissingPathSegments(Url),
MissingPathSegments(String),
#[error("Distribution not found at: {0}")]
NotFound(Url),
+43 -2
View File
@@ -743,7 +743,7 @@ impl RemoteSource for Url {
// Identify the last segment of the URL as the filename.
let path_segments = self
.path_segments()
.ok_or_else(|| Error::MissingPathSegments(self.clone()))?;
.ok_or_else(|| Error::MissingPathSegments(self.to_string()))?;
// This is guaranteed by the contract of `Url::path_segments`.
let last = path_segments.last().expect("path segments is non-empty");
@@ -759,6 +759,29 @@ impl RemoteSource for Url {
}
}
impl RemoteSource for UrlString {
fn filename(&self) -> Result<Cow<'_, str>, Error> {
// Take the last segment, stripping any query or fragment.
let last = self
.as_ref()
.split_once(['#', '?'])
.map(|(path, _)| path)
.unwrap_or(self.as_ref())
.split('/')
.last()
.ok_or_else(|| Error::MissingPathSegments(self.to_string()))?;
// Decode the filename, which may be percent-encoded.
let filename = urlencoding::decode(last)?;
Ok(filename)
}
fn size(&self) -> Option<u64> {
None
}
}
impl RemoteSource for RegistryBuiltWheel {
fn filename(&self) -> Result<Cow<'_, str>, Error> {
self.file.filename()
@@ -1215,7 +1238,8 @@ impl Identifier for BuildableSource<'_> {
#[cfg(test)]
mod test {
use crate::{BuiltDist, Dist, SourceDist};
use crate::{BuiltDist, Dist, RemoteSource, SourceDist, UrlString};
use url::Url;
/// Ensure that we don't accidentally grow the `Dist` sizes.
#[test]
@@ -1236,4 +1260,21 @@ mod test {
std::mem::size_of::<SourceDist>()
);
}
#[test]
fn remote_source() {
for url in [
"https://example.com/foo-0.1.0.tar.gz",
"https://example.com/foo-0.1.0.tar.gz#fragment",
"https://example.com/foo-0.1.0.tar.gz?query",
"https://example.com/foo-0.1.0.tar.gz?query#fragment",
"https://example.com/foo-0.1.0.tar.gz?query=1/2#fragment",
"https://example.com/foo-0.1.0.tar.gz?query=1/2#fragment/3",
] {
let url = Url::parse(url).unwrap();
assert_eq!(url.filename().unwrap(), "foo-0.1.0.tar.gz", "{url}");
let url = UrlString::from(url.clone());
assert_eq!(url.filename().unwrap(), "foo-0.1.0.tar.gz", "{url}");
}
}
}
+1 -3
View File
@@ -33,8 +33,6 @@ use pyo3::{
create_exception, exceptions::PyNotImplementedError, pyclass, pyclass::CompareOp, pymethods,
pymodule, types::PyModule, IntoPy, PyObject, PyResult, Python,
};
use schemars::gen::SchemaGenerator;
use schemars::schema::Schema;
use serde::{de, Deserialize, Deserializer, Serialize, Serializer};
use thiserror::Error;
use url::Url;
@@ -497,7 +495,7 @@ impl<T: Pep508Url> schemars::JsonSchema for Requirement<T> {
"Requirement".to_string()
}
fn json_schema(_gen: &mut SchemaGenerator) -> Schema {
fn json_schema(_gen: &mut schemars::gen::SchemaGenerator) -> schemars::schema::Schema {
schemars::schema::SchemaObject {
instance_type: Some(schemars::schema::InstanceType::String.into()),
metadata: Some(Box::new(schemars::schema::Metadata {
+8 -8
View File
@@ -801,7 +801,7 @@ impl Distribution {
requires_python: None,
size: sdist.size(),
upload_time_utc_ms: None,
url: FileLocation::AbsoluteUrl(file_url.clone().into()),
url: FileLocation::AbsoluteUrl(file_url.clone()),
yanked: None,
});
let index = IndexUrl::Url(VerbatimUrl::from_url(url.clone()));
@@ -1599,7 +1599,7 @@ struct SourceDistMetadata {
#[serde(untagged)]
enum SourceDist {
Url {
url: Url,
url: UrlString,
#[serde(flatten)]
metadata: SourceDistMetadata,
},
@@ -1634,7 +1634,7 @@ impl SourceDist {
}
}
fn url(&self) -> Option<&Url> {
fn url(&self) -> Option<&UrlString> {
match &self {
SourceDist::Url { url, .. } => Some(url),
SourceDist::Path { .. } => None,
@@ -1662,7 +1662,7 @@ impl SourceDist {
let mut table = InlineTable::new();
match &self {
SourceDist::Url { url, .. } => {
table.insert("url", Value::from(url.as_str()));
table.insert("url", Value::from(url.as_ref()));
}
SourceDist::Path { path, .. } => {
table.insert("path", Value::from(serialize_path_with_dot(path).as_ref()));
@@ -1733,7 +1733,7 @@ impl SourceDist {
let url = reg_dist
.file
.url
.to_url()
.to_url_string()
.map_err(LockErrorKind::InvalidFileUrl)
.map_err(LockError::from)?;
let hash = reg_dist.file.hashes.iter().max().cloned().map(Hash::from);
@@ -1758,7 +1758,7 @@ impl SourceDist {
return Err(kind.into());
};
Ok(SourceDist::Url {
url: direct_dist.url.to_url(),
url: UrlString::from(direct_dist.url.to_url()),
metadata: SourceDistMetadata {
hash: Some(hash),
size: None,
@@ -1975,7 +1975,7 @@ struct WheelWire {
/// against was found. The location does not need to exist in the future,
/// so this should be treated as only a hint to where to look and/or
/// recording where the wheel file originally came from.
url: Url,
url: UrlString,
/// A hash of the built distribution.
///
/// This is only present for wheels that come from registries and direct
@@ -2016,7 +2016,7 @@ impl TryFrom<WheelWire> for Wheel {
.map_err(|err| format!("failed to parse `{filename}` as wheel filename: {err}"))?;
Ok(Wheel {
url: wire.url.into(),
url: wire.url,
hash: wire.hash,
size: wire.size,
filename,
@@ -32,21 +32,9 @@ Ok(
},
sdist: Some(
Url {
url: Url {
scheme: "https",
cannot_be_a_base: false,
username: "",
password: None,
host: Some(
Domain(
"example.com",
),
),
port: None,
path: "/",
query: None,
fragment: None,
},
url: UrlString(
"https://example.com",
),
metadata: SourceDistMetadata {
hash: Some(
Hash(
@@ -93,21 +81,9 @@ Ok(
},
sdist: Some(
Url {
url: Url {
scheme: "https",
cannot_be_a_base: false,
username: "",
password: None,
host: Some(
Domain(
"example.com",
),
),
port: None,
path: "/",
query: None,
fragment: None,
},
url: UrlString(
"https://example.com",
),
metadata: SourceDistMetadata {
hash: Some(
Hash(
@@ -32,21 +32,9 @@ Ok(
},
sdist: Some(
Url {
url: Url {
scheme: "https",
cannot_be_a_base: false,
username: "",
password: None,
host: Some(
Domain(
"example.com",
),
),
port: None,
path: "/",
query: None,
fragment: None,
},
url: UrlString(
"https://example.com",
),
metadata: SourceDistMetadata {
hash: Some(
Hash(
@@ -93,21 +81,9 @@ Ok(
},
sdist: Some(
Url {
url: Url {
scheme: "https",
cannot_be_a_base: false,
username: "",
password: None,
host: Some(
Domain(
"example.com",
),
),
port: None,
path: "/",
query: None,
fragment: None,
},
url: UrlString(
"https://example.com",
),
metadata: SourceDistMetadata {
hash: Some(
Hash(
@@ -32,21 +32,9 @@ Ok(
},
sdist: Some(
Url {
url: Url {
scheme: "https",
cannot_be_a_base: false,
username: "",
password: None,
host: Some(
Domain(
"example.com",
),
),
port: None,
path: "/",
query: None,
fragment: None,
},
url: UrlString(
"https://example.com",
),
metadata: SourceDistMetadata {
hash: Some(
Hash(
@@ -93,21 +81,9 @@ Ok(
},
sdist: Some(
Url {
url: Url {
scheme: "https",
cannot_be_a_base: false,
username: "",
password: None,
host: Some(
Domain(
"example.com",
),
),
port: None,
path: "/",
query: None,
fragment: None,
},
url: UrlString(
"https://example.com",
),
metadata: SourceDistMetadata {
hash: Some(
Hash(