use std::path::PathBuf;
use crate::InfinoError;
#[derive(Debug, Clone)]
pub(crate) enum Backend {
LocalFs { root: PathBuf },
S3 { bucket: String, prefix: String },
Azure { container: String, prefix: String },
Gcs { bucket: String, prefix: String },
#[cfg_attr(not(feature = "remote"), allow(dead_code))]
Remote { base_url: String, database: String },
Memory,
}
impl Backend {
pub(crate) fn join(&self, segment: &str) -> Backend {
match self {
Backend::LocalFs { root } => Backend::LocalFs {
root: root.join(segment),
},
Backend::S3 { bucket, prefix } => Backend::S3 {
bucket: bucket.clone(),
prefix: join_prefix(prefix, segment),
},
Backend::Azure { container, prefix } => Backend::Azure {
container: container.clone(),
prefix: join_prefix(prefix, segment),
},
Backend::Gcs { bucket, prefix } => Backend::Gcs {
bucket: bucket.clone(),
prefix: join_prefix(prefix, segment),
},
Backend::Memory => Backend::Memory,
Backend::Remote { .. } => self.clone(),
}
}
}
fn join_prefix(prefix: &str, segment: &str) -> String {
let p = prefix.trim_matches('/');
if p.is_empty() {
segment.to_string()
} else {
format!("{p}/{segment}")
}
}
pub(crate) fn parse_uri(uri: &str) -> Result<Backend, InfinoError> {
if uri == "memory://" || uri == "memory:" || uri == "memory" {
return Ok(Backend::Memory);
}
if let Some(rest) = uri.strip_prefix("s3://") {
let (bucket, prefix) = split_bucket_prefix(rest);
if bucket.is_empty() {
return Err(InfinoError::Backend(format!(
"s3 URI missing bucket: {uri}"
)));
}
return Ok(Backend::S3 { bucket, prefix });
}
if let Some(rest) = uri
.strip_prefix("az://")
.or_else(|| uri.strip_prefix("azure://"))
{
let (container, prefix) = split_bucket_prefix(rest);
if container.is_empty() {
return Err(InfinoError::Backend(format!(
"azure URI missing container: {uri}"
)));
}
return Ok(Backend::Azure { container, prefix });
}
if let Some(rest) = uri
.strip_prefix("gs://")
.or_else(|| uri.strip_prefix("gcs://"))
{
let (bucket, prefix) = split_bucket_prefix(rest);
if bucket.is_empty() {
return Err(InfinoError::Backend(format!(
"gcs URI missing bucket: {uri}"
)));
}
return Ok(Backend::Gcs { bucket, prefix });
}
if let Some(rest) = uri.strip_prefix("file://") {
return Ok(Backend::LocalFs {
root: PathBuf::from(rest),
});
}
if let Some(rest) = uri
.strip_prefix("https://")
.or_else(|| uri.strip_prefix("http://"))
{
let is_http = uri.starts_with("http://");
let (host, database) = match rest.split_once('/') {
Some((host, db)) => (host, db.trim_matches('/')),
None => (rest, ""),
};
if host.is_empty() {
return Err(InfinoError::Backend(format!(
"remote URI missing host: {uri}"
)));
}
if is_http && !is_localhost(host) {
return Err(InfinoError::Backend(format!(
"http:// is only allowed for localhost; use https:// for a remote host: {uri}"
)));
}
if database.is_empty() {
return Err(InfinoError::Backend(format!(
"remote URI missing database path (expected https://host/<database>): {uri}"
)));
}
let scheme = if is_http { "http://" } else { "https://" };
return Ok(Backend::Remote {
base_url: format!("{scheme}{host}"),
database: database.to_string(),
});
}
if uri.contains("://") {
return Err(InfinoError::Backend(format!(
"unsupported catalog URI scheme: {uri}"
)));
}
Ok(Backend::LocalFs {
root: PathBuf::from(uri),
})
}
fn is_localhost(host: &str) -> bool {
let bare = host.split(':').next().unwrap_or(host);
bare == "localhost" || bare == "127.0.0.1"
}
fn split_bucket_prefix(rest: &str) -> (String, String) {
match rest.split_once('/') {
Some((bucket, prefix)) => (bucket.to_string(), prefix.trim_matches('/').to_string()),
None => (rest.to_string(), String::new()),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parses_memory() {
assert!(matches!(parse_uri("memory://"), Ok(Backend::Memory)));
}
#[test]
fn parses_bare_path_as_localfs() {
match parse_uri("./data").expect("parse") {
Backend::LocalFs { root } => assert_eq!(root, PathBuf::from("./data")),
other => panic!("expected LocalFs, got {other:?}"),
}
}
#[test]
fn parses_s3_bucket_and_prefix() {
match parse_uri("s3://my-bucket/some/prefix").expect("parse") {
Backend::S3 { bucket, prefix } => {
assert_eq!(bucket, "my-bucket");
assert_eq!(prefix, "some/prefix");
}
other => panic!("expected S3, got {other:?}"),
}
}
#[test]
fn join_appends_table_segment() {
let b = parse_uri("s3://b/root").expect("parse").join("users");
match b {
Backend::S3 { prefix, .. } => assert_eq!(prefix, "root/users"),
other => panic!("expected S3, got {other:?}"),
}
}
#[test]
fn rejects_unknown_scheme() {
assert!(parse_uri("gdrive://bucket/x").is_err());
}
#[test]
fn parses_gcs_bucket_and_prefix() {
match parse_uri("gs://my-bucket/some/prefix").expect("parse") {
Backend::Gcs { bucket, prefix } => {
assert_eq!(bucket, "my-bucket");
assert_eq!(prefix, "some/prefix");
}
other => panic!("expected Gcs, got {other:?}"),
}
}
#[test]
fn parses_gcs_alias_scheme() {
assert!(matches!(
parse_uri("gcs://b/p").expect("parse"),
Backend::Gcs { .. }
));
}
#[test]
fn gcs_join_appends_table_segment() {
match parse_uri("gs://b/root").expect("parse").join("users") {
Backend::Gcs { prefix, .. } => assert_eq!(prefix, "root/users"),
other => panic!("expected Gcs, got {other:?}"),
}
}
#[test]
fn rejects_gcs_uri_without_bucket() {
assert!(parse_uri("gs://").is_err());
}
#[test]
fn parses_https_remote() {
match parse_uri("https://base.example.ai/my-db").expect("parse") {
Backend::Remote { base_url, database } => {
assert_eq!(base_url, "https://base.example.ai");
assert_eq!(database, "my-db");
}
other => panic!("expected Remote, got {other:?}"),
}
}
#[test]
fn http_allowed_only_for_localhost() {
assert!(matches!(
parse_uri("http://localhost:8080/db").expect("parse"),
Backend::Remote { .. }
));
assert!(matches!(
parse_uri("http://127.0.0.1:9000/db").expect("parse"),
Backend::Remote { .. }
));
assert!(parse_uri("http://example.com/db").is_err());
}
#[test]
fn remote_requires_db_path() {
assert!(parse_uri("https://base.example.ai/").is_err());
assert!(parse_uri("https://base.example.ai").is_err());
}
#[test]
fn remote_root_join_is_noop() {
let root = parse_uri("https://base.example.ai/my-db").expect("parse");
assert!(
matches!(root.join("users"), Backend::Remote { database, .. } if database == "my-db")
);
}
}