1use std::time::Duration;
9
10use pomelo_data::error::DataError;
11use pomelo_data::{ObjectLister, ObjectSink, ObjectSource};
12use rusty_s3::actions::ListObjectsV2;
13use rusty_s3::{Bucket, Credentials, S3Action, UrlStyle};
14
15const LIST_RESPONSE_LIMIT: u64 = 16 * 1024 * 1024;
18
19const SIGN_TTL: Duration = Duration::from_secs(300);
21
22pub struct S3Source {
24 bucket: Bucket,
25 creds: Credentials,
26}
27
28impl S3Source {
29 pub fn new(
36 endpoint: &str,
37 bucket: &str,
38 access_key: &str,
39 secret_key: &str,
40 session_token: Option<&str>,
41 region: &str,
42 ) -> Result<Self, DataError> {
43 let url = endpoint
44 .parse()
45 .map_err(|e| DataError::Io(format!("bad S3 endpoint {endpoint:?}: {e}")))?;
46 let bucket = Bucket::new(url, UrlStyle::Path, bucket.to_string(), region.to_string())
47 .map_err(|e| DataError::Io(format!("bad S3 bucket: {e}")))?;
48 let creds = match session_token {
49 Some(t) => Credentials::new_with_token(access_key, secret_key, t),
50 None => Credentials::new(access_key, secret_key),
51 };
52 Ok(Self { bucket, creds })
53 }
54
55 pub fn from_env() -> Result<Self, DataError> {
59 let var = |k: &str| std::env::var(k).map_err(|_| DataError::Io(format!("missing env {k}")));
60 let region = std::env::var("S3_REGION").unwrap_or_else(|_| "auto".to_string());
61 let token = std::env::var("S3_SESSION_TOKEN").ok();
62 Self::new(
63 &var("S3_ENDPOINT")?,
64 &var("S3_BUCKET")?,
65 &var("S3_ACCESS_KEY_ID")?,
66 &var("S3_SECRET_ACCESS_KEY")?,
67 token.as_deref(),
68 ®ion,
69 )
70 }
71}
72
73impl ObjectSource for S3Source {
74 fn get(&self, key: &str) -> Result<Option<Vec<u8>>, DataError> {
75 let url = self
76 .bucket
77 .get_object(Some(&self.creds), key)
78 .sign(SIGN_TTL);
79 match ureq::get(url.as_str()).call() {
80 Ok(resp) => {
81 let bytes = resp
84 .into_body()
85 .with_config()
86 .limit(256 * 1024 * 1024)
87 .read_to_vec()
88 .map_err(|e| DataError::Io(format!("read {key}: {e}")))?;
89 Ok(Some(bytes))
90 }
91 Err(ureq::Error::StatusCode(404)) => Ok(None),
92 Err(e) => Err(DataError::Io(format!("GET {key}: {e}"))),
93 }
94 }
95}
96
97impl ObjectSink for S3Source {
98 fn put(&self, key: &str, bytes: &[u8]) -> Result<(), DataError> {
99 let url = self
100 .bucket
101 .put_object(Some(&self.creds), key)
102 .sign(SIGN_TTL);
103 match ureq::put(url.as_str()).send(bytes) {
104 Ok(_) => Ok(()),
105 Err(e) => Err(DataError::Io(format!("PUT {key}: {e}"))),
106 }
107 }
108}
109
110impl ObjectLister for S3Source {
111 fn list(&self, prefix: &str) -> Result<Vec<String>, DataError> {
115 let mut keys = Vec::new();
116 let mut continuation_token: Option<String> = None;
117 loop {
118 let mut action = ListObjectsV2::new(&self.bucket, Some(&self.creds));
119 action.with_prefix(prefix);
120 if let Some(tok) = &continuation_token {
121 action.with_continuation_token(tok.as_str());
122 }
123 let url = action.sign(SIGN_TTL);
124 let body = match ureq::get(url.as_str()).call() {
125 Ok(resp) => resp
126 .into_body()
127 .with_config()
128 .limit(LIST_RESPONSE_LIMIT)
129 .read_to_string()
130 .map_err(|e| DataError::Io(format!("read LIST {prefix}: {e}")))?,
131 Err(e) => return Err(DataError::Io(format!("LIST {prefix}: {e}"))),
132 };
133 let parsed = ListObjectsV2::parse_response(&body)
134 .map_err(|e| DataError::Io(format!("parse LIST {prefix} response: {e}")))?;
135 keys.extend(parsed.contents.into_iter().map(|c| c.key));
136 continuation_token = parsed.next_continuation_token;
137 if continuation_token.is_none() {
138 break;
139 }
140 }
141 Ok(keys)
142 }
143}
144
145pub struct S3Conn {
147 pub endpoint: String,
148 pub access_key: String,
149 pub secret_key: String,
150 pub session_token: Option<String>,
152 pub region: String,
153}
154
155pub fn resolve_s3_conn(get: impl Fn(&str) -> Option<String>) -> Result<S3Conn, String> {
163 for p in ["S3_", "AWS_"] {
164 let Some(access_key) = get(&format!("{p}ACCESS_KEY_ID")) else {
165 continue;
166 };
167 let secret_key = get(&format!("{p}SECRET_ACCESS_KEY")).ok_or_else(|| {
168 format!("{p}ACCESS_KEY_ID is set but {p}SECRET_ACCESS_KEY is missing")
169 })?;
170 let session_token = get(&format!("{p}SESSION_TOKEN"));
171 let region = get(&format!("{p}REGION")).unwrap_or_else(|| "auto".to_string());
172 let endpoint = get(&format!("{p}ENDPOINT"))
173 .or_else(|| get(&format!("{p}ENDPOINT_URL")))
174 .or_else(|| {
175 (p == "AWS_" && region != "auto")
177 .then(|| format!("https://s3.{region}.amazonaws.com"))
178 })
179 .ok_or_else(|| {
180 format!("set {p}ENDPOINT (R2: https://<acct>.r2.cloudflarestorage.com) or {p}REGION (AWS)")
181 })?;
182 return Ok(S3Conn {
183 endpoint,
184 access_key,
185 secret_key,
186 session_token,
187 region,
188 });
189 }
190 Err("no S3 credentials in env: set S3_ACCESS_KEY_ID + S3_SECRET_ACCESS_KEY (+ S3_ENDPOINT) for R2, or AWS_ACCESS_KEY_ID + AWS_SECRET_ACCESS_KEY (+ AWS_SESSION_TOKEN) for an AWS IAM role".into())
191}
192
193pub enum OutStore {
198 Local(pomelo_data::LocalSource),
199 S3 { src: Box<S3Source>, prefix: String },
201}
202
203impl OutStore {
204 pub fn parse(out: &str) -> Result<Self, String> {
207 let Some(rest) = out.strip_prefix("s3://") else {
208 return Ok(OutStore::Local(pomelo_data::LocalSource::new(out)));
209 };
210 let (bucket, prefix) = match rest.split_once('/') {
211 Some((b, p)) => (b, p.trim_matches('/')),
212 None => (rest, ""),
213 };
214 if bucket.is_empty() {
215 return Err("s3:// URL needs a bucket: s3://bucket[/prefix]".into());
216 }
217 let conn = resolve_s3_conn(|k| std::env::var(k).ok())?;
218 let src = S3Source::new(
219 &conn.endpoint,
220 bucket,
221 &conn.access_key,
222 &conn.secret_key,
223 conn.session_token.as_deref(),
224 &conn.region,
225 )
226 .map_err(|e| e.to_string())?;
227 Ok(OutStore::S3 {
228 src: Box::new(src),
229 prefix: prefix.to_string(),
230 })
231 }
232
233 pub fn is_s3(&self) -> bool {
234 matches!(self, OutStore::S3 { .. })
235 }
236
237 fn prefixed(prefix: &str, key: &str) -> String {
239 if prefix.is_empty() {
240 key.to_string()
241 } else {
242 format!("{prefix}/{key}")
243 }
244 }
245}
246
247impl ObjectSource for OutStore {
248 fn get(&self, key: &str) -> Result<Option<Vec<u8>>, DataError> {
249 match self {
250 OutStore::Local(l) => l.get(key),
251 OutStore::S3 { src, prefix } => src.get(&OutStore::prefixed(prefix, key)),
252 }
253 }
254}
255
256impl ObjectSink for OutStore {
257 fn put(&self, key: &str, bytes: &[u8]) -> Result<(), DataError> {
258 match self {
259 OutStore::Local(l) => l.put(key, bytes),
260 OutStore::S3 { src, prefix } => src.put(&OutStore::prefixed(prefix, key), bytes),
261 }
262 }
263}
264
265impl ObjectLister for OutStore {
266 fn list(&self, prefix: &str) -> Result<Vec<String>, DataError> {
267 match self {
268 OutStore::Local(l) => l.list(prefix),
269 OutStore::S3 {
270 src,
271 prefix: store_prefix,
272 } => {
273 let full_prefix = OutStore::prefixed(store_prefix, prefix);
274 let keys = src.list(&full_prefix)?;
275 let strip_from = if store_prefix.is_empty() {
278 0
279 } else {
280 store_prefix.len() + 1
281 };
282 Ok(keys
283 .into_iter()
284 .map(|k| k.get(strip_from..).unwrap_or(&k).to_string())
285 .collect())
286 }
287 }
288 }
289}
290
291#[cfg(test)]
292mod tests {
293 use super::*;
294 use std::fs;
295 use std::io::{Read, Write};
296 use std::net::TcpListener;
297 use std::thread;
298
299 fn spawn_stub(n: usize) -> String {
302 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
303 let addr = listener.local_addr().unwrap();
304 thread::spawn(move || {
305 for _ in 0..n {
306 let (mut sock, _) = listener.accept().unwrap();
307 let mut buf = [0u8; 2048];
308 let read = sock.read(&mut buf).unwrap();
309 let line = String::from_utf8_lossy(&buf[..read]);
310 let first = line.lines().next().unwrap_or("");
311 let resp = if first.contains("missing") {
312 "HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
313 .to_string()
314 } else {
315 "HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nhi"
316 .to_string()
317 };
318 sock.write_all(resp.as_bytes()).unwrap();
319 }
320 });
321 format!("http://{addr}")
322 }
323
324 fn spawn_encoded_stub(body: Vec<u8>) -> String {
328 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
329 let addr = listener.local_addr().unwrap();
330 thread::spawn(move || {
331 let (mut sock, _) = listener.accept().unwrap();
332 let mut buf = [0u8; 2048];
333 let _ = sock.read(&mut buf).unwrap();
334 let head = format!(
335 "HTTP/1.1 200 OK\r\nContent-Encoding: gzip\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
336 body.len()
337 );
338 sock.write_all(head.as_bytes()).unwrap();
339 sock.write_all(&body).unwrap();
340 });
341 format!("http://{addr}")
342 }
343
344 fn list_objects_v2_xml(keys: &[String], next_token: Option<&str>) -> String {
347 let contents: String = keys
348 .iter()
349 .map(|k| {
350 format!(
351 "<Contents><Key>{k}</Key><LastModified>2020-01-01T00:00:00.000Z</LastModified>\
352 <ETag>\"e\"</ETag><Size>1</Size><StorageClass>STANDARD</StorageClass></Contents>"
353 )
354 })
355 .collect();
356 let token = next_token
357 .map(|t| format!("<NextContinuationToken>{t}</NextContinuationToken>"))
358 .unwrap_or_default();
359 format!(
360 "<?xml version=\"1.0\" encoding=\"UTF-8\"?>\
361 <ListBucketResult xmlns=\"http://s3.amazonaws.com/doc/2006-03-01/\">\
362 <Name>bucket</Name><Prefix>prices/</Prefix><KeyCount>{}</KeyCount>\
363 <MaxKeys>1000</MaxKeys><IsTruncated>{}</IsTruncated>{contents}{token}\
364 <EncodingType>url</EncodingType></ListBucketResult>",
365 keys.len(),
366 next_token.is_some(),
367 )
368 }
369
370 fn spawn_body_pages_stub(bodies: Vec<String>) -> String {
373 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
374 let addr = listener.local_addr().unwrap();
375 thread::spawn(move || {
376 for body in bodies {
377 let (mut sock, _) = listener.accept().unwrap();
378 let mut buf = [0u8; 4096];
379 let _ = sock.read(&mut buf).unwrap();
380 let resp = format!(
381 "HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
382 body.len()
383 );
384 sock.write_all(resp.as_bytes()).unwrap();
385 }
386 });
387 format!("http://{addr}")
388 }
389
390 #[test]
391 fn get_returns_bytes_and_none_on_404() {
392 let endpoint = spawn_stub(2);
393 let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
394 assert_eq!(src.get("prices/AAPL.csv.gz").unwrap(), Some(b"hi".to_vec()));
395 assert_eq!(src.get("prices/missing.csv.gz").unwrap(), None);
396 }
397
398 #[test]
399 fn get_does_not_decompress_content_encoding_gzip() {
400 use flate2::write::GzEncoder;
401 use flate2::Compression;
402 let mut enc = GzEncoder::new(Vec::new(), Compression::default());
403 enc.write_all(b"hello,world\n").unwrap();
404 let gz = enc.finish().unwrap();
405
406 let endpoint = spawn_encoded_stub(gz.clone());
407 let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
408 let got = src.get("fundamentals/AAA.csv.gz").unwrap().unwrap();
409 assert_eq!(got, gz, "S3Source must return raw stored bytes");
411 assert_eq!(&got[..2], &[0x1f, 0x8b]);
412 }
413
414 #[test]
415 fn put_returns_ok_on_2xx() {
416 let endpoint = spawn_stub(1); let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
418 src.put("panels/close.csv.gz", b"gzip-bytes").unwrap();
419 }
420
421 #[test]
422 fn new_rejects_a_malformed_endpoint() {
423 let err = match S3Source::new("not a url", "bucket", "ak", "sk", None, "auto") {
424 Err(e) => e,
425 Ok(_) => panic!("expected a malformed-endpoint error"),
426 };
427 assert!(matches!(err, DataError::Io(_)));
428 }
429
430 fn dead_endpoint() -> String {
433 let listener = TcpListener::bind("127.0.0.1:0").unwrap();
434 let addr = listener.local_addr().unwrap();
435 drop(listener); format!("http://{addr}")
437 }
438
439 #[test]
440 fn get_and_put_surface_transport_errors() {
441 let src = S3Source::new(&dead_endpoint(), "bucket", "ak", "sk", None, "auto").unwrap();
442 assert!(matches!(
443 src.get("prices/AAA.csv.gz"),
444 Err(DataError::Io(_))
445 ));
446 assert!(matches!(
447 src.put("panels/close.csv.gz", b"x"),
448 Err(DataError::Io(_))
449 ));
450 }
451
452 #[test]
453 fn from_env_reads_vars_and_reports_missing() {
454 for k in [
457 "S3_ENDPOINT",
458 "S3_BUCKET",
459 "S3_ACCESS_KEY_ID",
460 "S3_SECRET_ACCESS_KEY",
461 "S3_REGION",
462 ] {
463 std::env::remove_var(k);
464 }
465 assert!(matches!(
467 S3Source::from_env(),
468 Err(DataError::Io(ref m)) if m.contains("S3_ENDPOINT")
469 ));
470
471 std::env::set_var("S3_ENDPOINT", "https://example.r2.cloudflarestorage.com");
472 std::env::set_var("S3_BUCKET", "bucket");
473 std::env::set_var("S3_ACCESS_KEY_ID", "ak");
474 std::env::set_var("S3_SECRET_ACCESS_KEY", "sk");
475 S3Source::from_env().expect("from_env builds when all required vars are present");
477
478 for k in [
479 "S3_ENDPOINT",
480 "S3_BUCKET",
481 "S3_ACCESS_KEY_ID",
482 "S3_SECRET_ACCESS_KEY",
483 ] {
484 std::env::remove_var(k);
485 }
486 }
487
488 #[test]
489 fn parse_local_vs_s3_and_key_prefixing() {
490 assert!(!OutStore::parse("./mydata").unwrap().is_s3());
492 assert!(!OutStore::parse("/tmp/x").unwrap().is_s3());
493 assert_eq!(
495 OutStore::prefixed("", "prices/AAPL.csv.gz"),
496 "prices/AAPL.csv.gz"
497 );
498 assert_eq!(
499 OutStore::prefixed("mirror/v1", "panels/piotroski_score.csv.gz"),
500 "mirror/v1/panels/piotroski_score.csv.gz"
501 );
502 }
503
504 fn env_of(pairs: &[(&str, &str)]) -> impl Fn(&str) -> Option<String> {
507 let owned: Vec<(String, String)> = pairs
508 .iter()
509 .map(|(k, v)| (k.to_string(), v.to_string()))
510 .collect();
511 move |k: &str| owned.iter().find(|(kk, _)| kk == k).map(|(_, v)| v.clone())
512 }
513
514 #[test]
515 fn resolve_s3_conn_prefers_s3_then_falls_back_to_aws() {
516 let c = resolve_s3_conn(env_of(&[
518 ("S3_ENDPOINT", "https://acct.r2.cloudflarestorage.com"),
519 ("S3_ACCESS_KEY_ID", "r2ak"),
520 ("S3_SECRET_ACCESS_KEY", "r2sk"),
521 ("AWS_ACCESS_KEY_ID", "awsak"),
522 ("AWS_SECRET_ACCESS_KEY", "awssk"),
523 ]))
524 .unwrap();
525 assert_eq!(c.access_key, "r2ak");
526 assert_eq!(c.endpoint, "https://acct.r2.cloudflarestorage.com");
527 assert_eq!(c.region, "auto");
528 assert!(c.session_token.is_none());
529
530 let c = resolve_s3_conn(env_of(&[
532 ("AWS_ACCESS_KEY_ID", "ASIAEXAMPLE"),
533 ("AWS_SECRET_ACCESS_KEY", "sk"),
534 ("AWS_SESSION_TOKEN", "tok"),
535 ("AWS_REGION", "us-east-1"),
536 ]))
537 .unwrap();
538 assert_eq!(c.access_key, "ASIAEXAMPLE");
539 assert_eq!(c.session_token.as_deref(), Some("tok"));
540 assert_eq!(c.endpoint, "https://s3.us-east-1.amazonaws.com");
541
542 assert!(resolve_s3_conn(env_of(&[("S3_ACCESS_KEY_ID", "x")])).is_err());
544 assert!(resolve_s3_conn(env_of(&[])).is_err());
546 }
547
548 #[test]
549 fn list_paginates_past_a_thousand_keys_via_continuation_token() {
550 let page1: Vec<String> = (0..1000)
554 .map(|i| format!("prices/SYM{i:04}.csv.gz"))
555 .collect();
556 let page2: Vec<String> = (1000..1200)
557 .map(|i| format!("prices/SYM{i:04}.csv.gz"))
558 .collect();
559 let endpoint = spawn_body_pages_stub(vec![
560 list_objects_v2_xml(&page1, Some("tok1")),
561 list_objects_v2_xml(&page2, None),
562 ]);
563 let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
564 let keys = src.list("prices").unwrap();
565 assert_eq!(keys.len(), 1200);
566 assert_eq!(keys[0], "prices/SYM0000.csv.gz");
567 assert_eq!(keys[999], "prices/SYM0999.csv.gz");
568 assert_eq!(keys[1199], "prices/SYM1199.csv.gz");
569 }
570
571 #[test]
572 fn list_returns_empty_on_no_contents() {
573 let endpoint = spawn_body_pages_stub(vec![list_objects_v2_xml(&[], None)]);
574 let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
575 assert_eq!(src.list("panels").unwrap(), Vec::<String>::new());
576 }
577
578 #[test]
579 fn list_surfaces_transport_errors() {
580 let src = S3Source::new(&dead_endpoint(), "bucket", "ak", "sk", None, "auto").unwrap();
581 assert!(matches!(src.list("prices"), Err(DataError::Io(_))));
582 }
583
584 #[test]
585 fn outstore_list_strips_the_store_prefix_from_s3_keys() {
586 let body = list_objects_v2_xml(&["mirror/v1/prices/AAPL.csv.gz".to_string()], None);
587 let endpoint = spawn_body_pages_stub(vec![body]);
588 let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
589 let out = OutStore::S3 {
590 src: Box::new(src),
591 prefix: "mirror/v1".to_string(),
592 };
593 assert_eq!(
594 out.list("prices").unwrap(),
595 vec!["prices/AAPL.csv.gz".to_string()]
596 );
597 }
598
599 #[test]
600 fn outstore_list_delegates_to_local_source() {
601 let dir = std::env::temp_dir().join("pomelo_s3_outstore_list_test");
602 let _ = fs::remove_dir_all(&dir);
603 fs::create_dir_all(dir.join("prices")).unwrap();
604 fs::write(dir.join("prices/AAPL.csv.gz"), b"x").unwrap();
605 let out = OutStore::parse(dir.to_str().unwrap()).unwrap();
606 assert_eq!(
607 out.list("prices").unwrap(),
608 vec!["prices/AAPL.csv.gz".to_string()]
609 );
610 }
611}