1use axum::extract::{FromRequestParts, OptionalFromRequestParts, Query};
4use axum::http::{HeaderName, HeaderValue, request::Parts};
5use axum_extra::headers::{ContentLength, Error as HeaderError, Header};
6use futures_util::TryStreamExt;
7use objectstore_service::error::Error as ServiceError;
8use objectstore_service::stream::ClientStream;
9use objectstore_types::resumable::{
10 HEADER_UPLOAD_LENGTH, HEADER_UPLOAD_OFFSET, SessionToken, UploadOffset,
11};
12use serde::Deserialize;
13use serde::de::IgnoredAny;
14
15use crate::endpoints::common::{ApiError, ApiResult};
16
17#[derive(Clone, Copy, Debug, Deserialize, PartialEq, Eq)]
19#[serde(rename_all = "snake_case")]
20enum UploadType {
21 Resumable,
23}
24
25#[derive(Debug, Deserialize)]
27struct ResumableQuery {
28 upload_type: Option<UploadType>,
30 session: Option<IgnoredAny>,
33}
34
35impl ResumableQuery {
36 fn classify(self) -> Option<ResumableTarget> {
38 match (self.upload_type, self.session) {
39 (_, Some(_)) => Some(ResumableTarget::ExistingSession),
42 (Some(UploadType::Resumable), None) => Some(ResumableTarget::NewSession),
43 (None, None) => None,
44 }
45 }
46}
47
48#[derive(Debug)]
50pub(crate) enum ResumableTarget {
51 NewSession,
53 ExistingSession,
55}
56
57impl<S> OptionalFromRequestParts<S> for ResumableTarget
58where
59 S: Send + Sync,
60{
61 type Rejection = ApiError;
62
63 async fn from_request_parts(
64 parts: &mut Parts,
65 _state: &S,
66 ) -> ApiResult<Option<ResumableTarget>> {
67 if parts.uri.query().is_none() {
68 return Ok(None);
69 }
70
71 let Query(query) = Query::<ResumableQuery>::try_from_uri(&parts.uri)
72 .map_err(|error| ApiError::map_client("invalid query parameters", error))?;
73
74 Ok(query.classify())
75 }
76}
77
78#[derive(Debug)]
80pub(crate) struct Session(pub(crate) SessionToken);
81
82#[derive(Debug, Deserialize)]
83struct SessionQuery {
84 session: String,
85}
86
87impl<S> FromRequestParts<S> for Session
88where
89 S: Send + Sync,
90{
91 type Rejection = ApiError;
92
93 async fn from_request_parts(parts: &mut Parts, _state: &S) -> ApiResult<Session> {
94 let Query(SessionQuery { session }) = Query::<SessionQuery>::try_from_uri(&parts.uri)
95 .map_err(|error| ApiError::map_client("invalid query parameters", error))?;
96 Ok(Session(SessionToken::from_base64url(&session).map_err(
97 |error| ApiError::map_client("invalid session token", error),
98 )?))
99 }
100}
101
102#[derive(Clone, Copy, Debug)]
104pub(crate) struct UploadLengthHeader(pub(crate) u64);
105
106impl Header for UploadLengthHeader {
107 fn name() -> &'static HeaderName {
108 static NAME: HeaderName = HeaderName::from_static(HEADER_UPLOAD_LENGTH);
109 &NAME
110 }
111
112 fn decode<'i, I>(values: &mut I) -> Result<Self, HeaderError>
113 where
114 I: Iterator<Item = &'i HeaderValue>,
115 {
116 let value = decode_single_value(values)?;
117 let value = value.to_str().map_err(|_| HeaderError::invalid())?;
118
119 if !value.bytes().all(|byte| byte.is_ascii_digit()) {
120 return Err(HeaderError::invalid());
121 }
122
123 value.parse().map(Self).map_err(|_| HeaderError::invalid())
124 }
125
126 fn encode<E>(&self, values: &mut E)
127 where
128 E: Extend<HeaderValue>,
129 {
130 values.extend(std::iter::once(HeaderValue::from(self.0)));
131 }
132}
133
134#[derive(Clone, Copy, Debug)]
136pub(crate) struct UploadOffsetHeader(pub(crate) UploadOffset);
137
138impl Header for UploadOffsetHeader {
139 fn name() -> &'static HeaderName {
140 static NAME: HeaderName = HeaderName::from_static(HEADER_UPLOAD_OFFSET);
141 &NAME
142 }
143
144 fn decode<'i, I>(values: &mut I) -> Result<Self, HeaderError>
145 where
146 I: Iterator<Item = &'i HeaderValue>,
147 {
148 decode_single_value(values)?
149 .to_str()
150 .map_err(|_| HeaderError::invalid())?
151 .parse()
152 .map(Self)
153 .map_err(|_| HeaderError::invalid())
154 }
155
156 fn encode<E>(&self, values: &mut E)
157 where
158 E: Extend<HeaderValue>,
159 {
160 let value = match self.0 {
161 UploadOffset::At(offset) => HeaderValue::from(offset),
162 UploadOffset::Unknown => HeaderValue::from_static("*"),
163 };
164 values.extend(std::iter::once(value));
165 }
166}
167
168fn decode_single_value<'i, I>(values: &mut I) -> Result<&'i HeaderValue, HeaderError>
169where
170 I: Iterator<Item = &'i HeaderValue>,
171{
172 let value = values.next().ok_or_else(HeaderError::invalid)?;
173 if values.next().is_some() {
174 return Err(HeaderError::invalid());
175 }
176 Ok(value)
177}
178
179pub(crate) async fn require_empty_body(
181 content_length: Option<ContentLength>,
182 mut body: ClientStream,
183 request: &str,
184) -> ApiResult<()> {
185 if content_length.is_some_and(|ContentLength(length)| length > 0) {
186 return Err(ApiError::client(format!(
187 "{request} must be sent with an empty body"
188 )));
189 }
190
191 while let Some(chunk) = body.try_next().await.map_err(ServiceError::from)? {
192 if !chunk.is_empty() {
193 return Err(ApiError::client(format!(
194 "{request} must be sent with an empty body"
195 )));
196 }
197 }
198
199 Ok(())
200}
201
202#[cfg(test)]
203mod tests {
204 use super::*;
205 use axum::http::HeaderMap;
206
207 fn query(upload_type: Option<UploadType>, has_session: bool) -> ResumableQuery {
208 ResumableQuery {
209 upload_type,
210 session: has_session.then_some(IgnoredAny),
211 }
212 }
213
214 fn decode_header<H: Header>(headers: &HeaderMap) -> Result<H, HeaderError> {
215 H::decode(&mut headers.get_all(H::name()).iter())
216 }
217
218 #[test]
219 fn classify_recognizes_each_operation() {
220 assert!(matches!(
221 query(Some(UploadType::Resumable), false).classify(),
222 Some(ResumableTarget::NewSession)
223 ));
224 assert!(matches!(
225 query(None, true).classify(),
226 Some(ResumableTarget::ExistingSession)
227 ));
228 assert!(query(None, false).classify().is_none());
229 }
230
231 #[test]
232 fn classify_prefers_session_over_upload_type() {
233 assert!(matches!(
234 query(Some(UploadType::Resumable), true).classify(),
235 Some(ResumableTarget::ExistingSession)
236 ));
237 }
238
239 #[test]
240 fn session_token_decodes_from_unpadded_base64url() {
241 assert_eq!(
242 SessionToken::from_base64url("Li4vZXNjYXBl")
243 .unwrap()
244 .as_bytes(),
245 b"../escape"
246 );
247 }
248
249 #[test]
250 fn session_token_rejects_invalid_query_encodings() {
251 for invalid in ["%%%", "dG9rM24="] {
252 assert!(
253 SessionToken::from_base64url(invalid).is_err(),
254 "accepted {invalid:?}"
255 );
256 }
257 }
258
259 #[test]
260 fn upload_length_requires_a_byte_count() {
261 let mut headers = HeaderMap::new();
262 assert!(
263 decode_header::<UploadLengthHeader>(&headers).is_err(),
264 "missing header"
265 );
266
267 for invalid in ["", "-1", "+1", "1.5", "abc", " 1"] {
268 headers.insert(HEADER_UPLOAD_LENGTH, invalid.parse().unwrap());
269 assert!(
270 decode_header::<UploadLengthHeader>(&headers).is_err(),
271 "accepted {invalid:?}"
272 );
273 }
274
275 headers.insert(HEADER_UPLOAD_LENGTH, "1048576".parse().unwrap());
276 assert_eq!(
277 decode_header::<UploadLengthHeader>(&headers).unwrap().0,
278 1_048_576
279 );
280 }
281
282 #[test]
283 fn upload_offset_parses_chunk_and_wildcard() {
284 let mut headers = HeaderMap::new();
285 assert!(
286 decode_header::<UploadOffsetHeader>(&headers).is_err(),
287 "missing header"
288 );
289
290 headers.insert(HEADER_UPLOAD_OFFSET, "*".parse().unwrap());
291 assert_eq!(
292 decode_header::<UploadOffsetHeader>(&headers).unwrap().0,
293 UploadOffset::Unknown
294 );
295
296 headers.insert(HEADER_UPLOAD_OFFSET, "262144".parse().unwrap());
297 assert_eq!(
298 decode_header::<UploadOffsetHeader>(&headers).unwrap().0,
299 UploadOffset::At(262_144)
300 );
301
302 headers.insert(HEADER_UPLOAD_OFFSET, "nope".parse().unwrap());
303 assert!(decode_header::<UploadOffsetHeader>(&headers).is_err());
304 }
305}