diff --git a/src/services/storage/providers/azure_blob/mod.rs b/src/services/storage/providers/azure_blob/mod.rs index 1f9a575..3945e26 100644 --- a/src/services/storage/providers/azure_blob/mod.rs +++ b/src/services/storage/providers/azure_blob/mod.rs @@ -1,5 +1,5 @@ pub mod helpers; -mod models; +pub(crate) mod models; use crate::core::context::Context; use crate::services::api::models::agent::status::DatabaseStorage; diff --git a/src/services/storage/providers/azure_blob/models.rs b/src/services/storage/providers/azure_blob/models.rs index 617ce6e..070e462 100644 --- a/src/services/storage/providers/azure_blob/models.rs +++ b/src/services/storage/providers/azure_blob/models.rs @@ -1,13 +1,18 @@ use crate::services::storage::providers::azure_blob::helpers::ResolvedAzure; use anyhow::{Result, anyhow}; use serde::{Deserialize, Serialize}; +use url::Url; #[derive(Debug, Deserialize, Serialize)] pub struct AzureBlobProviderConfig { + #[serde(default)] pub account_name: String, + #[serde(default)] pub account_key: String, pub container_name: String, #[serde(default)] + pub auth_mode: Option, + #[serde(default)] pub connection_string: String, #[serde(default)] pub endpoint_url: Option, @@ -25,9 +30,36 @@ fn parse_connection_string(cs: &str) -> std::collections::HashMap String { + let trimmed = endpoint.trim_end_matches('/'); + if account.is_empty() { + return trimmed.to_string(); + } + if let Ok(url) = Url::parse(trimmed) { + let host = url.host_str().unwrap_or(""); + if host.contains(account) { + return trimmed.to_string(); + } + let path = url.path().trim_matches('/'); + if path == account || path.starts_with(&format!("{account}/")) { + return trimmed.to_string(); + } + } + format!("{trimmed}/{account}") +} + impl AzureBlobProviderConfig { pub fn resolve(&self) -> Result { - if !self.connection_string.trim().is_empty() { + let mode = self.auth_mode.as_deref().unwrap_or("").trim(); + let has_connection_string = !self.connection_string.trim().is_empty(); + + if mode == "connectionString" || (mode.is_empty() && has_connection_string) { + if !has_connection_string { + return Err(anyhow!( + "authMode is connectionString but connectionString is empty" + )); + } let map = parse_connection_string(&self.connection_string); let account_name = map .get("AccountName") @@ -48,11 +80,22 @@ impl AzureBlobProviderConfig { }); } - let blob_endpoint = self + if self.account_name.trim().is_empty() { + return Err(anyhow!("accountName required for accountKey auth")); + } + if self.account_key.trim().is_empty() { + return Err(anyhow!("accountKey required for accountKey auth")); + } + + let blob_endpoint = match self .endpoint_url - .clone() - .filter(|s| !s.trim().is_empty()) - .ok_or_else(|| anyhow!("endpointUrl required when connectionString is empty"))?; + .as_deref() + .map(str::trim) + .filter(|s| !s.is_empty()) + { + Some(endpoint) => ensure_account_in_endpoint(endpoint, &self.account_name), + None => format!("https://{}.blob.core.windows.net", self.account_name), + }; Ok(ResolvedAzure { account_name: self.account_name.clone(), diff --git a/src/tests/storage/azure_blob.rs b/src/tests/storage/azure_blob.rs index 7409858..0085d49 100644 --- a/src/tests/storage/azure_blob.rs +++ b/src/tests/storage/azure_blob.rs @@ -16,9 +16,7 @@ use testcontainers::runners::AsyncRunner; use testcontainers::{GenericImage, ImageExt}; use url::Url; -/// Build an Account SAS query set (test-only). Azurite cannot authorize container-create with -/// a container-scoped Service SAS, so tests create the target container with an Account SAS. -/// Reuses the production HMAC primitive (`hmac_sha256_b64`) to avoid duplicating signing logic. + fn build_account_sas( resolved: &ResolvedAzure, services: &str, @@ -33,12 +31,9 @@ fn build_account_sas( let signed_expiry = (Utc::now() + Duration::hours(1)) .format("%Y-%m-%dT%H:%M:%SZ") .to_string(); - let signed_protocol = "https,http"; // Azurite is http + let signed_protocol = "https,http"; let signed_ip = String::new(); let encryption_scope = String::new(); - - // Account SAS string-to-sign for sv >= 2020-12-06: - // account \n sp \n ss \n srt \n st \n se \n sip \n spr \n sv \n ses \n (trailing newline) let string_to_sign = format!( "{acc}\n{sp}\n{ss}\n{srt}\n{st}\n{se}\n{sip}\n{spr}\n{sv}\n{ses}\n", acc = resolved.account_name, sp = permissions, ss = services, srt = resource_types, @@ -59,7 +54,6 @@ fn build_account_sas( ]) } -/// Build an Account-SAS-scoped URL for a container (test-only container creation). fn build_account_sas_container_url( resolved: &ResolvedAzure, container: &str, @@ -86,21 +80,12 @@ const AZURITE_KEY: &str = async fn start_azurite() -> (testcontainers::ContainerAsync, ResolvedAzure) { let container = GenericImage::new("mcr.microsoft.com/azure-storage/azurite", "latest") .with_exposed_port(10000.tcp()) - // The current `latest` image logs (on stdout): - // "Azurite Blob service successfully listens on http://0.0.0.0:10000" - // Older builds phrased it "...is successfully listening"; this substring matches - // the wording the pulled image actually emits. .with_wait_for(WaitFor::message_on_stdout( "Azurite Blob service successfully listens on", )) - // The GA SDK sends a very recent `x-ms-version`; Azurite 3.35 rejects unknown - // versions unless we tell it to skip that check. .with_cmd(["azurite-blob", "--blobHost", "0.0.0.0", "--skipApiVersionCheck"]) .start().await.unwrap(); - // Use the testcontainers-resolved host (not a hardcoded 127.0.0.1): under - // docker-out-of-docker / remote daemons the published port is not on the test - // process's loopback. All other container tests (mssql, valkey, postgres, ...) - // already do this; azure_blob was the only one hardcoding the host. + let host = container.get_host().await.unwrap().to_string(); let port = container.get_host_port_ipv4(10000).await.unwrap(); let resolved = ResolvedAzure { @@ -118,8 +103,7 @@ async fn spike_sas_block_roundtrip_against_azurite() { let container = "portabase"; let blob = "spike/hello.txt"; - // Container creation must use an Account SAS (service=blob, resource-type=container, - // perms=create+write). Azurite cannot authorize container-create with a Service SAS. + let container_url = build_account_sas_container_url(&resolved, container, "b", "c", "cw").unwrap(); let container_client = @@ -135,9 +119,6 @@ async fn spike_sas_block_roundtrip_against_azurite() { bbc.stage_block(&raw_id, payload.len() as u64, RequestContent::from(payload.to_vec()), None) .await.unwrap(); - // `BlockLookupList.latest` is `Option>>` and base64-encodes each entry - // internally during XML serialization, exactly as `stage_block` base64-encodes the - // `blockid` query. So `latest` must hold the SAME RAW id bytes passed to `stage_block`. let block_list = BlockLookupList { latest: Some(vec![raw_id.clone()]), ..Default::default() }; bbc.commit_block_list(block_list.try_into().unwrap(), None).await.unwrap(); @@ -146,6 +127,85 @@ async fn spike_sas_block_roundtrip_against_azurite() { assert!(read_client.exists().await.unwrap()); } +mod resolve { + use crate::services::storage::providers::azure_blob::models::{ + AzureBlobProviderConfig, ensure_account_in_endpoint, + }; + + const AZURITE_KEY: &str = + "Eby8vdM02xNOcqFlqUwJPLlmEtlCDXJ1OUzFT50uSRZ6IFsuFq2UVErCz4I6tq/K1SZFPTOtr/KBHBeksoGMGw=="; + const CONNECTION_STRING: &str = "DefaultEndpointsProtocol=http;AccountName=devstoreaccount1;AccountKey=Eby8vdM02xNOcqFlqUwJPLlmEtlCDXJ1OUzFT50uSRZ6IFsuFq2UVErCz4I6tq/K1SZFPTOtr/KBHBeksoGMGw==;BlobEndpoint=http://localhost:10000/devstoreaccount1;QueueEndpoint=http://localhost:10001/devstoreaccount1;TableEndpoint=http://localhost:10002/devstoreaccount1;"; + + + #[test] + fn resolve_connection_string_mode() { + let cfg = AzureBlobProviderConfig { + account_name: String::new(), + account_key: String::new(), + container_name: "portabase".into(), + auth_mode: Some("connectionString".into()), + connection_string: CONNECTION_STRING.into(), + endpoint_url: None, + }; + let r = cfg.resolve().unwrap(); + assert_eq!(r.account_name, "devstoreaccount1"); + assert_eq!(r.account_key, AZURITE_KEY); + assert_eq!(r.blob_endpoint, "http://localhost:10000/devstoreaccount1"); + } + + + #[test] + fn resolve_account_key_mode_injects_account_path() { + let cfg = AzureBlobProviderConfig { + account_name: "devstoreaccount1".into(), + account_key: AZURITE_KEY.into(), + container_name: "portabase".into(), + auth_mode: Some("accountKey".into()), + connection_string: CONNECTION_STRING.into(), + endpoint_url: Some("http://localhost:10000".into()), + }; + let r = cfg.resolve().unwrap(); + assert_eq!(r.account_name, "devstoreaccount1"); + assert_eq!(r.account_key, AZURITE_KEY); + assert_eq!(r.blob_endpoint, "http://localhost:10000/devstoreaccount1"); + } + + #[test] + fn resolve_implicit_connection_string() { + let cfg = AzureBlobProviderConfig { + account_name: String::new(), + account_key: String::new(), + container_name: "portabase".into(), + auth_mode: None, + connection_string: CONNECTION_STRING.into(), + endpoint_url: None, + }; + let r = cfg.resolve().unwrap(); + assert_eq!(r.blob_endpoint, "http://localhost:10000/devstoreaccount1"); + } + + #[test] + fn resolve_account_key_default_endpoint() { + let cfg = AzureBlobProviderConfig { + account_name: "myaccount".into(), + account_key: AZURITE_KEY.into(), + container_name: "portabase".into(), + auth_mode: Some("accountKey".into()), + connection_string: String::new(), + endpoint_url: None, + }; + let r = cfg.resolve().unwrap(); + assert_eq!(r.blob_endpoint, "https://myaccount.blob.core.windows.net"); + } + + #[test] + fn ensure_account_keeps_host_style_endpoint() { + let got = + ensure_account_in_endpoint("https://myaccount.blob.core.windows.net", "myaccount"); + assert_eq!(got, "https://myaccount.blob.core.windows.net"); + } +} + #[tokio::test] async fn upload_stream_multi_block_roundtrip() { init_tracing_for_test(); @@ -156,7 +216,6 @@ async fn upload_stream_multi_block_roundtrip() { let container = "portabase"; let blob = "backups/multi.bin"; - // Container setup (provider itself never creates it): Account SAS create. let container_url = build_account_sas_container_url(&resolved, container, "b", "c", "cw").unwrap(); azure_storage_blob::clients::BlobContainerClient::new(container_url, None, None) @@ -165,7 +224,6 @@ async fn upload_stream_multi_block_roundtrip() { .await .unwrap(); - // 10 KiB fed as 1 KiB chunks, forced into 4 KiB blocks => 3 blocks (multi-block path). let data = vec![7u8; 10 * 1024]; let chunks: Vec> = data .chunks(1024) @@ -177,7 +235,6 @@ async fn upload_stream_multi_block_roundtrip() { .await .unwrap(); - // Verify the committed blob reassembles to the exact source bytes via a read-SAS GET. let read_url = build_sas_url(&resolved, container, blob, SasResource::Blob, "r").unwrap(); let got = reqwest::get(read_url).await.unwrap().bytes().await.unwrap(); assert_eq!(got.as_ref(), data.as_slice());