1use axum::{
4 Json, Router,
5 body::Body,
6 extract::{Path, Query, State},
7 http::{HeaderMap, StatusCode, header},
8 routing::{post, put},
9};
10use futures_util::StreamExt as _;
11use serde::Deserialize;
12use worker::{Env, FixedLengthStream, HttpMetadata, UploadedPart, send::SendFuture};
13
14use super::{Ctx, storage::bucket};
15use crate::{
16 ApiError, Error, Result,
17 config::Environment,
18 storage::{
19 CompletedPart, DirectUploadRequest, FinishRequest, MAX_PART, MAX_WORKER_PART, MultipartUpload, PartUrls,
20 PartsRequest, Rules, check_multipart, check_part_numbers, essence, new_key, presign_parts, r2_endpoint,
21 sign_key, upload_secret, verify_key,
22 },
23};
24
25const PART_URL_EXPIRES_IN: u64 = 24 * 3600;
27
28pub fn multipart_uploads(path: &str, prefix: &'static str, field: &'static str, rules: &'static Rules) -> Router<Ctx> {
69 let base = path.trim_end_matches('/').to_owned();
70 let parts_path = base.clone();
71 Router::new()
72 .route(
73 &base,
74 post(move |State(ctx): State<Ctx>, Json(request): Json<DirectUploadRequest>| {
75 SendFuture::new(async move {
76 start(&ctx, prefix, field, rules, &request).await.map(Json).map_err(ApiError::from)
77 })
78 }),
79 )
80 .route(
81 &format!("{base}/parts"),
82 post(move |State(ctx): State<Ctx>, Json(request): Json<PartsRequest>| {
83 let result = part_urls(&ctx, &parts_path, field, &request);
84 async move { result.map(Json).map_err(ApiError::from) }
85 }),
86 )
87 .route(
88 &format!("{base}/parts/{{part}}"),
89 put(
90 move |State(ctx): State<Ctx>,
91 Path(part): Path<u16>,
92 Query(target): Query<PartTarget>,
93 headers: HeaderMap,
94 body: Body| {
95 SendFuture::new(async move {
96 upload_part(&ctx, field, part, &target, &headers, body).await.map(Json).map_err(ApiError::from)
97 })
98 },
99 ),
100 )
101 .route(
102 &format!("{base}/complete"),
103 post(move |State(ctx): State<Ctx>, Json(request): Json<FinishRequest>| {
104 SendFuture::new(async move { finish(&ctx, field, &request, true).await.map_err(ApiError::from) })
105 }),
106 )
107 .route(
108 &format!("{base}/abort"),
109 post(move |State(ctx): State<Ctx>, Json(request): Json<FinishRequest>| {
110 SendFuture::new(async move { finish(&ctx, field, &request, false).await.map_err(ApiError::from) })
111 }),
112 )
113}
114
115#[derive(Deserialize)]
117struct PartTarget {
118 key: String,
119 upload_id: String,
120}
121
122fn secret(env: &Env) -> Result<String> {
123 upload_secret(&|name| env.var(name).ok().map(|value| value.to_string()))
124}
125
126fn direct(env: &Env) -> bool {
128 !Environment::current().is_development() && r2_endpoint(&|name| env.var(name).ok().map(|v| v.to_string())).is_ok()
129}
130
131async fn start(
132 ctx: &Ctx,
133 prefix: &str,
134 field: &str,
135 rules: &Rules,
136 request: &DirectUploadRequest,
137) -> Result<MultipartUpload> {
138 let env = ctx.env();
139 let max_part = if direct(env) { MAX_PART } else { MAX_WORKER_PART };
140 let content_type = essence(&request.content_type);
141 let (part_size, part_count) = check_multipart(field, request.size, &content_type, rules, max_part)?;
142 let key = new_key(prefix);
143 let metadata = HttpMetadata { content_type: Some(content_type), ..Default::default() };
144 let upload = bucket(env)?
145 .create_multipart_upload(key.clone())
146 .http_metadata(metadata)
147 .execute()
148 .await
149 .map_err(|err| Error::internal(format!("R2 could not start the upload of `{key}`: {err}")))?;
150 Ok(MultipartUpload {
151 signed_key: sign_key(&secret(env)?, &key),
152 upload_id: upload.upload_id().await,
153 part_size,
154 part_count,
155 })
156}
157
158fn part_urls(ctx: &Ctx, path: &str, field: &str, request: &PartsRequest) -> Result<PartUrls> {
159 let env = ctx.env();
160 let key = verify_key(field, &secret(env)?, &request.signed_key)?;
161 if direct(env) {
162 let endpoint = r2_endpoint(&|name| env.var(name).ok().map(|v| v.to_string()))?;
163 return presign_parts(&endpoint, &key, &request.upload_id, &request.parts, crate::now(), PART_URL_EXPIRES_IN);
164 }
165 check_part_numbers(&request.parts)?;
166 let query = serde_urlencoded::to_string([("key", request.signed_key.as_str()), ("upload_id", &request.upload_id)])
167 .map_err(|err| Error::internal(err.to_string()))?;
168 let urls = request.parts.iter().map(|&part| (part, format!("{path}/parts/{part}?{query}"))).collect();
169 Ok(PartUrls { urls })
170}
171
172async fn upload_part(
173 ctx: &Ctx,
174 field: &str,
175 part: u16,
176 target: &PartTarget,
177 headers: &HeaderMap,
178 body: Body,
179) -> Result<CompletedPart> {
180 check_part_numbers(&[part])?;
181 let env = ctx.env();
182 let key = verify_key(field, &secret(env)?, &target.key)?;
183 let size: u64 = headers
184 .get(header::CONTENT_LENGTH)
185 .and_then(|value| value.to_str().ok()?.parse().ok())
186 .ok_or_else(|| Error::bad_request("a part needs a Content-Length"))?;
187 if size > MAX_WORKER_PART {
188 return Err(Error::PayloadTooLarge(format!(
189 "a part sent through the Worker is at most {MAX_WORKER_PART} bytes"
190 )));
191 }
192 let chunks = body
193 .into_data_stream()
194 .map(|chunk| chunk.map(|bytes| bytes.to_vec()).map_err(|err| worker::Error::RustError(err.to_string())));
195 let upload = bucket(env)?.resume_multipart_upload(key.clone(), target.upload_id.clone())?;
196 let stored =
197 upload.upload_part(part, FixedLengthStream::wrap(chunks, size)).await.map_err(|err| match err.to_string() {
198 text if text.contains("does not exist") => Error::NotFound,
200 text => Error::bad_request(format!("R2 refused part {part} of `{key}`: {text}")),
201 })?;
202 Ok(CompletedPart { part_number: stored.part_number(), etag: stored.etag() })
203}
204
205async fn finish(ctx: &Ctx, field: &str, request: &FinishRequest, complete: bool) -> Result<StatusCode> {
206 let env = ctx.env();
207 let key = verify_key(field, &secret(env)?, &request.signed_key)?;
208 let upload = bucket(env)?.resume_multipart_upload(key.clone(), request.upload_id.clone())?;
209 if complete {
210 if request.parts.is_empty() {
211 return Err(Error::bad_request("complete needs the parts"));
212 }
213 let parts = request
215 .parts
216 .iter()
217 .map(|part| UploadedPart::new(part.part_number, part.etag.trim_matches('"').to_owned()));
218 upload
219 .complete(parts)
220 .await
221 .map_err(|err| Error::bad_request(format!("R2 could not assemble `{key}`: {err}")))?;
222 } else {
223 upload.abort().await.map_err(|err| Error::bad_request(format!("R2 could not abort `{key}`: {err}")))?;
224 }
225 Ok(StatusCode::NO_CONTENT)
226}