use std::time::Duration;
use pomelo_data::error::DataError;
use pomelo_data::{ObjectLister, ObjectSink, ObjectSource};
use rusty_s3::actions::ListObjectsV2;
use rusty_s3::{Bucket, Credentials, S3Action, UrlStyle};
const LIST_RESPONSE_LIMIT: u64 = 16 * 1024 * 1024;
const SIGN_TTL: Duration = Duration::from_secs(300);
pub struct S3Source {
bucket: Bucket,
creds: Credentials,
}
impl S3Source {
pub fn new(
endpoint: &str,
bucket: &str,
access_key: &str,
secret_key: &str,
session_token: Option<&str>,
region: &str,
) -> Result<Self, DataError> {
let url = endpoint
.parse()
.map_err(|e| DataError::Io(format!("bad S3 endpoint {endpoint:?}: {e}")))?;
let bucket = Bucket::new(url, UrlStyle::Path, bucket.to_string(), region.to_string())
.map_err(|e| DataError::Io(format!("bad S3 bucket: {e}")))?;
let creds = match session_token {
Some(t) => Credentials::new_with_token(access_key, secret_key, t),
None => Credentials::new(access_key, secret_key),
};
Ok(Self { bucket, creds })
}
pub fn from_env() -> Result<Self, DataError> {
let var = |k: &str| std::env::var(k).map_err(|_| DataError::Io(format!("missing env {k}")));
let region = std::env::var("S3_REGION").unwrap_or_else(|_| "auto".to_string());
let token = std::env::var("S3_SESSION_TOKEN").ok();
Self::new(
&var("S3_ENDPOINT")?,
&var("S3_BUCKET")?,
&var("S3_ACCESS_KEY_ID")?,
&var("S3_SECRET_ACCESS_KEY")?,
token.as_deref(),
®ion,
)
}
}
impl ObjectSource for S3Source {
fn get(&self, key: &str) -> Result<Option<Vec<u8>>, DataError> {
let url = self
.bucket
.get_object(Some(&self.creds), key)
.sign(SIGN_TTL);
match ureq::get(url.as_str()).call() {
Ok(resp) => {
let bytes = resp
.into_body()
.with_config()
.limit(256 * 1024 * 1024)
.read_to_vec()
.map_err(|e| DataError::Io(format!("read {key}: {e}")))?;
Ok(Some(bytes))
}
Err(ureq::Error::StatusCode(404)) => Ok(None),
Err(e) => Err(DataError::Io(format!("GET {key}: {e}"))),
}
}
}
impl ObjectSink for S3Source {
fn put(&self, key: &str, bytes: &[u8]) -> Result<(), DataError> {
let url = self
.bucket
.put_object(Some(&self.creds), key)
.sign(SIGN_TTL);
match ureq::put(url.as_str()).send(bytes) {
Ok(_) => Ok(()),
Err(e) => Err(DataError::Io(format!("PUT {key}: {e}"))),
}
}
}
impl ObjectLister for S3Source {
fn list(&self, prefix: &str) -> Result<Vec<String>, DataError> {
let mut keys = Vec::new();
let mut continuation_token: Option<String> = None;
loop {
let mut action = ListObjectsV2::new(&self.bucket, Some(&self.creds));
action.with_prefix(prefix);
if let Some(tok) = &continuation_token {
action.with_continuation_token(tok.as_str());
}
let url = action.sign(SIGN_TTL);
let body = match ureq::get(url.as_str()).call() {
Ok(resp) => resp
.into_body()
.with_config()
.limit(LIST_RESPONSE_LIMIT)
.read_to_string()
.map_err(|e| DataError::Io(format!("read LIST {prefix}: {e}")))?,
Err(e) => return Err(DataError::Io(format!("LIST {prefix}: {e}"))),
};
let parsed = ListObjectsV2::parse_response(&body)
.map_err(|e| DataError::Io(format!("parse LIST {prefix} response: {e}")))?;
keys.extend(parsed.contents.into_iter().map(|c| c.key));
continuation_token = parsed.next_continuation_token;
if continuation_token.is_none() {
break;
}
}
Ok(keys)
}
}
pub struct S3Conn {
pub endpoint: String,
pub access_key: String,
pub secret_key: String,
pub session_token: Option<String>,
pub region: String,
}
pub fn resolve_s3_conn(get: impl Fn(&str) -> Option<String>) -> Result<S3Conn, String> {
for p in ["S3_", "AWS_"] {
let Some(access_key) = get(&format!("{p}ACCESS_KEY_ID")) else {
continue;
};
let secret_key = get(&format!("{p}SECRET_ACCESS_KEY")).ok_or_else(|| {
format!("{p}ACCESS_KEY_ID is set but {p}SECRET_ACCESS_KEY is missing")
})?;
let session_token = get(&format!("{p}SESSION_TOKEN"));
let region = get(&format!("{p}REGION")).unwrap_or_else(|| "auto".to_string());
let endpoint = get(&format!("{p}ENDPOINT"))
.or_else(|| get(&format!("{p}ENDPOINT_URL")))
.or_else(|| {
(p == "AWS_" && region != "auto")
.then(|| format!("https://s3.{region}.amazonaws.com"))
})
.ok_or_else(|| {
format!("set {p}ENDPOINT (R2: https://<acct>.r2.cloudflarestorage.com) or {p}REGION (AWS)")
})?;
return Ok(S3Conn {
endpoint,
access_key,
secret_key,
session_token,
region,
});
}
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())
}
pub enum OutStore {
Local(pomelo_data::LocalSource),
S3 { src: Box<S3Source>, prefix: String },
}
impl OutStore {
pub fn parse(out: &str) -> Result<Self, String> {
let Some(rest) = out.strip_prefix("s3://") else {
return Ok(OutStore::Local(pomelo_data::LocalSource::new(out)));
};
let (bucket, prefix) = match rest.split_once('/') {
Some((b, p)) => (b, p.trim_matches('/')),
None => (rest, ""),
};
if bucket.is_empty() {
return Err("s3:// URL needs a bucket: s3://bucket[/prefix]".into());
}
let conn = resolve_s3_conn(|k| std::env::var(k).ok())?;
let src = S3Source::new(
&conn.endpoint,
bucket,
&conn.access_key,
&conn.secret_key,
conn.session_token.as_deref(),
&conn.region,
)
.map_err(|e| e.to_string())?;
Ok(OutStore::S3 {
src: Box::new(src),
prefix: prefix.to_string(),
})
}
pub fn is_s3(&self) -> bool {
matches!(self, OutStore::S3 { .. })
}
fn prefixed(prefix: &str, key: &str) -> String {
if prefix.is_empty() {
key.to_string()
} else {
format!("{prefix}/{key}")
}
}
}
impl ObjectSource for OutStore {
fn get(&self, key: &str) -> Result<Option<Vec<u8>>, DataError> {
match self {
OutStore::Local(l) => l.get(key),
OutStore::S3 { src, prefix } => src.get(&OutStore::prefixed(prefix, key)),
}
}
}
impl ObjectSink for OutStore {
fn put(&self, key: &str, bytes: &[u8]) -> Result<(), DataError> {
match self {
OutStore::Local(l) => l.put(key, bytes),
OutStore::S3 { src, prefix } => src.put(&OutStore::prefixed(prefix, key), bytes),
}
}
}
impl ObjectLister for OutStore {
fn list(&self, prefix: &str) -> Result<Vec<String>, DataError> {
match self {
OutStore::Local(l) => l.list(prefix),
OutStore::S3 {
src,
prefix: store_prefix,
} => {
let full_prefix = OutStore::prefixed(store_prefix, prefix);
let keys = src.list(&full_prefix)?;
let strip_from = if store_prefix.is_empty() {
0
} else {
store_prefix.len() + 1
};
Ok(keys
.into_iter()
.map(|k| k.get(strip_from..).unwrap_or(&k).to_string())
.collect())
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::fs;
use std::io::{Read, Write};
use std::net::TcpListener;
use std::thread;
fn spawn_stub(n: usize) -> String {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
thread::spawn(move || {
for _ in 0..n {
let (mut sock, _) = listener.accept().unwrap();
let mut buf = [0u8; 2048];
let read = sock.read(&mut buf).unwrap();
let line = String::from_utf8_lossy(&buf[..read]);
let first = line.lines().next().unwrap_or("");
let resp = if first.contains("missing") {
"HTTP/1.1 404 Not Found\r\nContent-Length: 0\r\nConnection: close\r\n\r\n"
.to_string()
} else {
"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nhi"
.to_string()
};
sock.write_all(resp.as_bytes()).unwrap();
}
});
format!("http://{addr}")
}
fn spawn_encoded_stub(body: Vec<u8>) -> String {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
thread::spawn(move || {
let (mut sock, _) = listener.accept().unwrap();
let mut buf = [0u8; 2048];
let _ = sock.read(&mut buf).unwrap();
let head = format!(
"HTTP/1.1 200 OK\r\nContent-Encoding: gzip\r\nContent-Length: {}\r\nConnection: close\r\n\r\n",
body.len()
);
sock.write_all(head.as_bytes()).unwrap();
sock.write_all(&body).unwrap();
});
format!("http://{addr}")
}
fn list_objects_v2_xml(keys: &[String], next_token: Option<&str>) -> String {
let contents: String = keys
.iter()
.map(|k| {
format!(
"<Contents><Key>{k}</Key><LastModified>2020-01-01T00:00:00.000Z</LastModified>\
<ETag>\"e\"</ETag><Size>1</Size><StorageClass>STANDARD</StorageClass></Contents>"
)
})
.collect();
let token = next_token
.map(|t| format!("<NextContinuationToken>{t}</NextContinuationToken>"))
.unwrap_or_default();
format!(
"<?xml version=\"1.0\" encoding=\"UTF-8\"?>\
<ListBucketResult xmlns=\"http://s3.amazonaws.com/doc/2006-03-01/\">\
<Name>bucket</Name><Prefix>prices/</Prefix><KeyCount>{}</KeyCount>\
<MaxKeys>1000</MaxKeys><IsTruncated>{}</IsTruncated>{contents}{token}\
<EncodingType>url</EncodingType></ListBucketResult>",
keys.len(),
next_token.is_some(),
)
}
fn spawn_body_pages_stub(bodies: Vec<String>) -> String {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
thread::spawn(move || {
for body in bodies {
let (mut sock, _) = listener.accept().unwrap();
let mut buf = [0u8; 4096];
let _ = sock.read(&mut buf).unwrap();
let resp = format!(
"HTTP/1.1 200 OK\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{body}",
body.len()
);
sock.write_all(resp.as_bytes()).unwrap();
}
});
format!("http://{addr}")
}
#[test]
fn get_returns_bytes_and_none_on_404() {
let endpoint = spawn_stub(2);
let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
assert_eq!(src.get("prices/AAPL.csv.gz").unwrap(), Some(b"hi".to_vec()));
assert_eq!(src.get("prices/missing.csv.gz").unwrap(), None);
}
#[test]
fn get_does_not_decompress_content_encoding_gzip() {
use flate2::write::GzEncoder;
use flate2::Compression;
let mut enc = GzEncoder::new(Vec::new(), Compression::default());
enc.write_all(b"hello,world\n").unwrap();
let gz = enc.finish().unwrap();
let endpoint = spawn_encoded_stub(gz.clone());
let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
let got = src.get("fundamentals/AAA.csv.gz").unwrap().unwrap();
assert_eq!(got, gz, "S3Source must return raw stored bytes");
assert_eq!(&got[..2], &[0x1f, 0x8b]);
}
#[test]
fn put_returns_ok_on_2xx() {
let endpoint = spawn_stub(1); let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
src.put("panels/close.csv.gz", b"gzip-bytes").unwrap();
}
#[test]
fn new_rejects_a_malformed_endpoint() {
let err = match S3Source::new("not a url", "bucket", "ak", "sk", None, "auto") {
Err(e) => e,
Ok(_) => panic!("expected a malformed-endpoint error"),
};
assert!(matches!(err, DataError::Io(_)));
}
fn dead_endpoint() -> String {
let listener = TcpListener::bind("127.0.0.1:0").unwrap();
let addr = listener.local_addr().unwrap();
drop(listener); format!("http://{addr}")
}
#[test]
fn get_and_put_surface_transport_errors() {
let src = S3Source::new(&dead_endpoint(), "bucket", "ak", "sk", None, "auto").unwrap();
assert!(matches!(
src.get("prices/AAA.csv.gz"),
Err(DataError::Io(_))
));
assert!(matches!(
src.put("panels/close.csv.gz", b"x"),
Err(DataError::Io(_))
));
}
#[test]
fn from_env_reads_vars_and_reports_missing() {
for k in [
"S3_ENDPOINT",
"S3_BUCKET",
"S3_ACCESS_KEY_ID",
"S3_SECRET_ACCESS_KEY",
"S3_REGION",
] {
std::env::remove_var(k);
}
assert!(matches!(
S3Source::from_env(),
Err(DataError::Io(ref m)) if m.contains("S3_ENDPOINT")
));
std::env::set_var("S3_ENDPOINT", "https://example.r2.cloudflarestorage.com");
std::env::set_var("S3_BUCKET", "bucket");
std::env::set_var("S3_ACCESS_KEY_ID", "ak");
std::env::set_var("S3_SECRET_ACCESS_KEY", "sk");
S3Source::from_env().expect("from_env builds when all required vars are present");
for k in [
"S3_ENDPOINT",
"S3_BUCKET",
"S3_ACCESS_KEY_ID",
"S3_SECRET_ACCESS_KEY",
] {
std::env::remove_var(k);
}
}
#[test]
fn parse_local_vs_s3_and_key_prefixing() {
assert!(!OutStore::parse("./mydata").unwrap().is_s3());
assert!(!OutStore::parse("/tmp/x").unwrap().is_s3());
assert_eq!(
OutStore::prefixed("", "prices/AAPL.csv.gz"),
"prices/AAPL.csv.gz"
);
assert_eq!(
OutStore::prefixed("mirror/v1", "panels/piotroski_score.csv.gz"),
"mirror/v1/panels/piotroski_score.csv.gz"
);
}
fn env_of(pairs: &[(&str, &str)]) -> impl Fn(&str) -> Option<String> {
let owned: Vec<(String, String)> = pairs
.iter()
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect();
move |k: &str| owned.iter().find(|(kk, _)| kk == k).map(|(_, v)| v.clone())
}
#[test]
fn resolve_s3_conn_prefers_s3_then_falls_back_to_aws() {
let c = resolve_s3_conn(env_of(&[
("S3_ENDPOINT", "https://acct.r2.cloudflarestorage.com"),
("S3_ACCESS_KEY_ID", "r2ak"),
("S3_SECRET_ACCESS_KEY", "r2sk"),
("AWS_ACCESS_KEY_ID", "awsak"),
("AWS_SECRET_ACCESS_KEY", "awssk"),
]))
.unwrap();
assert_eq!(c.access_key, "r2ak");
assert_eq!(c.endpoint, "https://acct.r2.cloudflarestorage.com");
assert_eq!(c.region, "auto");
assert!(c.session_token.is_none());
let c = resolve_s3_conn(env_of(&[
("AWS_ACCESS_KEY_ID", "ASIAEXAMPLE"),
("AWS_SECRET_ACCESS_KEY", "sk"),
("AWS_SESSION_TOKEN", "tok"),
("AWS_REGION", "us-east-1"),
]))
.unwrap();
assert_eq!(c.access_key, "ASIAEXAMPLE");
assert_eq!(c.session_token.as_deref(), Some("tok"));
assert_eq!(c.endpoint, "https://s3.us-east-1.amazonaws.com");
assert!(resolve_s3_conn(env_of(&[("S3_ACCESS_KEY_ID", "x")])).is_err());
assert!(resolve_s3_conn(env_of(&[])).is_err());
}
#[test]
fn list_paginates_past_a_thousand_keys_via_continuation_token() {
let page1: Vec<String> = (0..1000)
.map(|i| format!("prices/SYM{i:04}.csv.gz"))
.collect();
let page2: Vec<String> = (1000..1200)
.map(|i| format!("prices/SYM{i:04}.csv.gz"))
.collect();
let endpoint = spawn_body_pages_stub(vec![
list_objects_v2_xml(&page1, Some("tok1")),
list_objects_v2_xml(&page2, None),
]);
let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
let keys = src.list("prices").unwrap();
assert_eq!(keys.len(), 1200);
assert_eq!(keys[0], "prices/SYM0000.csv.gz");
assert_eq!(keys[999], "prices/SYM0999.csv.gz");
assert_eq!(keys[1199], "prices/SYM1199.csv.gz");
}
#[test]
fn list_returns_empty_on_no_contents() {
let endpoint = spawn_body_pages_stub(vec![list_objects_v2_xml(&[], None)]);
let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
assert_eq!(src.list("panels").unwrap(), Vec::<String>::new());
}
#[test]
fn list_surfaces_transport_errors() {
let src = S3Source::new(&dead_endpoint(), "bucket", "ak", "sk", None, "auto").unwrap();
assert!(matches!(src.list("prices"), Err(DataError::Io(_))));
}
#[test]
fn outstore_list_strips_the_store_prefix_from_s3_keys() {
let body = list_objects_v2_xml(&["mirror/v1/prices/AAPL.csv.gz".to_string()], None);
let endpoint = spawn_body_pages_stub(vec![body]);
let src = S3Source::new(&endpoint, "bucket", "ak", "sk", None, "auto").unwrap();
let out = OutStore::S3 {
src: Box::new(src),
prefix: "mirror/v1".to_string(),
};
assert_eq!(
out.list("prices").unwrap(),
vec!["prices/AAPL.csv.gz".to_string()]
);
}
#[test]
fn outstore_list_delegates_to_local_source() {
let dir = std::env::temp_dir().join("pomelo_s3_outstore_list_test");
let _ = fs::remove_dir_all(&dir);
fs::create_dir_all(dir.join("prices")).unwrap();
fs::write(dir.join("prices/AAPL.csv.gz"), b"x").unwrap();
let out = OutStore::parse(dir.to_str().unwrap()).unwrap();
assert_eq!(
out.list("prices").unwrap(),
vec!["prices/AAPL.csv.gz".to_string()]
);
}
}