Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion src/services/storage/providers/azure_blob/mod.rs
Original file line number Diff line number Diff line change
@@ -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;
Expand Down
53 changes: 48 additions & 5 deletions src/services/storage/providers/azure_blob/models.rs
Original file line number Diff line number Diff line change
@@ -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<String>,
#[serde(default)]
pub connection_string: String,
#[serde(default)]
pub endpoint_url: Option<String>,
Expand All @@ -25,9 +30,36 @@ fn parse_connection_string(cs: &str) -> std::collections::HashMap<String, String
.collect()
}


pub(crate) fn ensure_account_in_endpoint(endpoint: &str, account: &str) -> 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<ResolvedAzure> {
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")
Expand All @@ -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(),
Expand Down
109 changes: 83 additions & 26 deletions src/tests/storage/azure_blob.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
Expand All @@ -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,
Expand All @@ -86,21 +80,12 @@ const AZURITE_KEY: &str =
async fn start_azurite() -> (testcontainers::ContainerAsync<GenericImage>, 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 {
Expand All @@ -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 =
Expand All @@ -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<Vec<Vec<u8>>>` 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();

Expand All @@ -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();
Expand All @@ -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)
Expand All @@ -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<Result<Bytes, std::io::Error>> = data
.chunks(1024)
Expand All @@ -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());
Expand Down