diff --git a/native/csrc/catalog/bindings_store.cpp b/native/csrc/catalog/bindings_store.cpp index 5866d9f79..11d8ebdec 100644 --- a/native/csrc/catalog/bindings_store.cpp +++ b/native/csrc/catalog/bindings_store.cpp @@ -36,12 +36,19 @@ dmi_store::S3Config s3_config(const py::dict& d) { c.secret_key = get(d, "s3_secret_key", ""); c.session_token = get(d, "s3_session_token", ""); c.allow_insecure_http = get(d, "s3_allow_insecure_http", false); + c.ca_file = get(d, "s3_ca_file", ""); + c.ca_path = get(d, "s3_ca_path", ""); c.connect_timeout_s = get(d, "s3_connect_timeout_s", c.connect_timeout_s); c.read_timeout_s = get(d, "s3_read_timeout_s", c.read_timeout_s); c.max_attempts = get(d, "s3_max_attempts", c.max_attempts); if (c.endpoint.empty() || c.bucket.empty()) { throw py::value_error("s3_endpoint and s3_bucket are required"); } + // The client would refuse every request of an invalid config (a CA on + // http://, a missing CA file, a sub-5 MiB part); raising here makes the + // refusal land when the service or reader is built instead. + const std::string invalid = dmi_store::S3Client::ValidateConfig(c); + if (!invalid.empty()) throw py::value_error("s3: " + invalid); return c; } diff --git a/native/csrc/catalog/conformance_catalog.cpp b/native/csrc/catalog/conformance_catalog.cpp index a22f5ed37..02a7223c3 100644 --- a/native/csrc/catalog/conformance_catalog.cpp +++ b/native/csrc/catalog/conformance_catalog.cpp @@ -479,6 +479,8 @@ std::string respond(const std::string& line, Session* session) { s3_config.access_key = jc::FindString(line, "access"); s3_config.secret_key = jc::FindString(line, "secret"); s3_config.allow_insecure_http = jc::FindBool(line, "insecure"); + s3_config.ca_file = jc::FindString(line, "ca_file"); + s3_config.ca_path = jc::FindString(line, "ca_path"); dmi_store::S3Client s3(s3_config); dmi_catalog::NativeCaptureReader reader( &s3, s3_config.bucket, session->client, rc); @@ -511,6 +513,8 @@ std::string respond(const std::string& line, Session* session) { s3_config.access_key = jc::FindString(line, "access"); s3_config.secret_key = jc::FindString(line, "secret"); s3_config.allow_insecure_http = jc::FindBool(line, "insecure"); + s3_config.ca_file = jc::FindString(line, "ca_file"); + s3_config.ca_path = jc::FindString(line, "ca_path"); dmi_store::S3Client s3(s3_config); dmi_catalog::NativeCaptureReader reader( &s3, s3_config.bucket, session->client, rc); @@ -547,6 +551,8 @@ std::string respond(const std::string& line, Session* session) { s3_config.access_key = jc::FindString(line, "access"); s3_config.secret_key = jc::FindString(line, "secret"); s3_config.allow_insecure_http = jc::FindBool(line, "insecure"); + s3_config.ca_file = jc::FindString(line, "ca_file"); + s3_config.ca_path = jc::FindString(line, "ca_path"); dmi_store::S3Client s3(s3_config); dmi_catalog::NativeCaptureReader reader( &s3, s3_config.bucket, session->client, rc); @@ -954,6 +960,8 @@ std::string respond(const std::string& line, Session* session) { s3_config.access_key = jc::FindString(line, "access"); s3_config.secret_key = jc::FindString(line, "secret"); s3_config.allow_insecure_http = jc::FindBool(line, "insecure"); + s3_config.ca_file = jc::FindString(line, "ca_file"); + s3_config.ca_path = jc::FindString(line, "ca_path"); dmi_store::S3Client s3(s3_config); std::vector refs; for (const std::string& element : jc::SplitElements( diff --git a/native/csrc/store/conformance_store.cpp b/native/csrc/store/conformance_store.cpp index 091926856..3456bf043 100644 --- a/native/csrc/store/conformance_store.cpp +++ b/native/csrc/store/conformance_store.cpp @@ -6,7 +6,7 @@ // "access":"...","secret":"...","token":null,"insecure":true, // "key":"...","data_b64":"...","metadata":{...},"content_type":"...", // "multipart_threshold":N,"multipart_chunk":N,"max_attempts":N, -// "connect_timeout":N,"read_timeout":N} +// "connect_timeout":N,"read_timeout":N,"ca_file":"...","ca_path":"..."} // -> {"ok":true,"etag":"...","attempts":N} // {"op":"get",...,"offset":N,"length":N} -> {"ok":true,"data_b64":"...","attempts":N} // {"op":"head",...} -> {"ok":true,"found":bool,"size":N,"metadata":{...}, @@ -92,6 +92,8 @@ dmi_store::S3Config ReadConfig(const std::string& line) { } } config.allow_insecure_http = jc::FindBool(line, "insecure"); + config.ca_file = jc::FindString(line, "ca_file"); + config.ca_path = jc::FindString(line, "ca_path"); const int64_t connect_timeout = Integer(line, "connect_timeout"); config.connect_timeout_s = static_cast(connect_timeout > 0 ? connect_timeout : 5); diff --git a/native/csrc/store/s3_client.cpp b/native/csrc/store/s3_client.cpp index 7024e5bf6..370f08562 100644 --- a/native/csrc/store/s3_client.cpp +++ b/native/csrc/store/s3_client.cpp @@ -1,6 +1,8 @@ #include "s3_client.h" #include +#include +#include #include #include @@ -109,8 +111,50 @@ std::string XmlTag(const std::string& xml, const std::string& tag) { return xml.substr(start, end - start); } +// Checked when the client is built, so a mistyped path fails with its name +// rather than as libcurl's "problem with the SSL CA cert" at the first +// request. +bool IsReadableFile(const std::string& path) { + struct stat st {}; + return ::stat(path.c_str(), &st) == 0 && S_ISREG(st.st_mode) && + ::access(path.c_str(), R_OK) == 0; +} + +bool IsDirectory(const std::string& path) { + struct stat st {}; + return ::stat(path.c_str(), &st) == 0 && S_ISDIR(st.st_mode); +} + } // namespace +std::string S3Client::ValidateConfig(const S3Config& config) { + const bool https = config.endpoint.compare(0, 8, "https://") == 0; + if (https && config.allow_insecure_http) { + // Silently downgrading TLS was never the intent of the flag -- it gates + // plain http:// endpoints for local Garage. + return "https endpoint with allow_insecure_http is refused"; + } + if (!https && (!config.ca_file.empty() || !config.ca_path.empty())) { + // A CA on a plain-http endpoint would read as "this is TLS" while every + // byte, credentials included, goes in the clear. + return "ca_file and ca_path apply only to an https endpoint"; + } + if (!config.ca_file.empty() && !IsReadableFile(config.ca_file)) { + return "ca_file is not a readable file: " + config.ca_file; + } + if (!config.ca_path.empty() && !IsDirectory(config.ca_path)) { + return "ca_path is not a directory: " + config.ca_path; + } + if (config.multipart_chunk_bytes < kMinMultipartPartBytes) { + // The Python store refuses the same (s3.py _MIN_MULTIPART_BYTES). Left + // to the server, the upload fails only at CompleteMultipartUpload, after + // every part was sent. + return "multipart_chunk_bytes must be at least " + + std::to_string(kMinMultipartPartBytes) + " (S3's minimum part size)"; + } + return ""; +} + S3Client::S3Client(S3Config config) : config_(std::move(config)) { // Explicit, rather than leaning on the implicit init inside // curl_easy_init: that implicit path carries libcurl's thread-safety @@ -132,11 +176,7 @@ S3Client::S3Client(S3Config config) : config_(std::move(config)) { } while (!rest.empty() && rest.back() == '/') rest.pop_back(); host_ = rest; - if (is_https_ && config_.allow_insecure_http) { - // Refused at construction: silently downgrading TLS was never the intent - // of the flag — it gates plain http:// endpoints for local Garage. - host_.clear(); - } + config_error_ = ValidateConfig(config_); } S3Client::~S3Client() = default; @@ -147,19 +187,17 @@ S3Response S3Client::Exchange( const std::map& extra_headers, const uint8_t* body, size_t body_len, const std::string& body_hash_hex) { S3Response response; - if (host_.empty()) { - response.error = "https endpoint with allow_insecure_http is refused"; + if (!config_error_.empty()) { + last_attempts_ = 0; + response.error = config_error_; return response; } - const std::string amz_date = AmzDate(std::time(nullptr)); - const std::string datestamp = Datestamp(amz_date); - const std::string encoded_resource = "/" + config_.bucket + "/" + key; const std::string encoded_path = UriEncode(encoded_resource, true); + // Every signed header except x-amz-date, which each attempt stamps. std::map headers; headers["host"] = host_; - headers["x-amz-date"] = amz_date; headers["x-amz-content-sha256"] = body_hash_hex; if (!config_.session_token.empty()) { headers["x-amz-security-token"] = config_.session_token; @@ -167,10 +205,6 @@ S3Response S3Client::Exchange( for (const auto& [name, value] : extra_headers) { headers[LowerHeader(name)] = value; } - const std::string authz = AuthorizationHeader( - config_.access_key, config_.secret_key, datestamp, amz_date, - config_.region, "s3", method, encoded_path, query, headers, - body_hash_hex); // Canonical query string for the URL (same encoding the signer used). std::string query_text; @@ -196,6 +230,16 @@ S3Response S3Client::Exchange( last_attempts_ = 0; for (int attempt = 0; attempt < config_.max_attempts; ++attempt) { ++last_attempts_; + // Signed per attempt, not once before the loop: SigV4 binds the + // signature to x-amz-date, and S3 refuses a date more than 15 minutes + // off. Replaying the first attempt's date let a slow first attempt (up + // to read_timeout_s each, plus backoff) age every retry after it. + const std::string amz_date = AmzDate(std::time(nullptr)); + headers["x-amz-date"] = amz_date; + const std::string authz = AuthorizationHeader( + config_.access_key, config_.secret_key, Datestamp(amz_date), amz_date, + config_.region, "s3", method, encoded_path, query, headers, + body_hash_hex); CURL* curl = curl_easy_init(); if (!curl) { response.error = "curl_easy_init failed"; @@ -218,15 +262,24 @@ S3Response S3Client::Exchange( curl_easy_setopt(curl, CURLOPT_CONNECTTIMEOUT, config_.connect_timeout_s); curl_easy_setopt(curl, CURLOPT_TIMEOUT, config_.read_timeout_s); curl_easy_setopt(curl, CURLOPT_NOSIGNAL, 1L); - if (!is_https_) { - // Plain http only by explicit opt-in (local Garage); https always - // verifies (no CURLOPT_SSL_VERIFYPEER toggle exists anywhere here). - if (!config_.allow_insecure_http) { - response.error = "plain http endpoint requires allow_insecure_http"; - curl_slist_free_all(chunk); - curl_easy_cleanup(curl); - return response; + if (is_https_) { + // https always verifies: peer and host name, stated explicitly rather + // than left to libcurl's defaults, and never switched off. A private + // CA adds trust; it does not relax the check. + curl_easy_setopt(curl, CURLOPT_SSL_VERIFYPEER, 1L); + curl_easy_setopt(curl, CURLOPT_SSL_VERIFYHOST, 2L); + if (!config_.ca_file.empty()) { + curl_easy_setopt(curl, CURLOPT_CAINFO, config_.ca_file.c_str()); + } + if (!config_.ca_path.empty()) { + curl_easy_setopt(curl, CURLOPT_CAPATH, config_.ca_path.c_str()); } + } else if (!config_.allow_insecure_http) { + // Plain http only by explicit opt-in (local Garage). + response.error = "plain http endpoint requires allow_insecure_http"; + curl_slist_free_all(chunk); + curl_easy_cleanup(curl); + return response; } std::string response_body; std::map response_headers; diff --git a/native/csrc/store/s3_client.h b/native/csrc/store/s3_client.h index 4707c8064..84e62c4af 100644 --- a/native/csrc/store/s3_client.h +++ b/native/csrc/store/s3_client.h @@ -32,10 +32,18 @@ struct S3Config { std::string secret_key; std::string session_token; // empty when unused bool allow_insecure_http = false; + // https only: trust a private CA. ca_file is a PEM bundle + // (CURLOPT_CAINFO), ca_path an OpenSSL-hashed certificate directory + // (CURLOPT_CAPATH). Both empty uses libcurl's default trust store. Either + // way https always verifies the peer and the host name. + std::string ca_file; + std::string ca_path; int connect_timeout_s = 5; int read_timeout_s = 120; int max_attempts = 4; uint64_t multipart_threshold_bytes = 64ull * 1024 * 1024; + // The size of every part but the last. S3 refuses a smaller non-final + // part (EntityTooSmall), so the client refuses one under kMinMultipartPartBytes. uint64_t multipart_chunk_bytes = 16ull * 1024 * 1024; std::string user_agent = "dmi-native-store/1"; }; @@ -70,9 +78,18 @@ struct ListResult { std::vector objects; }; +// S3's minimum size for every part of a multipart upload but the last. +inline constexpr uint64_t kMinMultipartPartBytes = 5ull * 1024 * 1024; + class S3Client { public: + // An invalid config (see ValidateConfig) does not throw: the client + // refuses every request with the reason, before anything goes out. explicit S3Client(S3Config config); + + // Empty when `config` is usable; otherwise why not. Callers with an error + // channel of their own (the Python bindings) check it at construction. + static std::string ValidateConfig(const S3Config& config); ~S3Client(); S3Client(const S3Client&) = delete; @@ -116,6 +133,7 @@ class S3Client { private: S3Config config_; + std::string config_error_; // non-empty: every request is refused std::string host_; // endpoint host (with :port when non-default) std::string scheme_; bool is_https_ = false; diff --git a/src/dmi/storage/native_capture.py b/src/dmi/storage/native_capture.py index db2c90b6e..3adde5b5c 100644 --- a/src/dmi/storage/native_capture.py +++ b/src/dmi/storage/native_capture.py @@ -81,6 +81,11 @@ class NativeCaptureStorageConfig: s3_session_token: str = field(default="", repr=False) # Plain-HTTP endpoints (a local Garage or MinIO) must be opted into. s3_allow_insecure_http: bool = False + # https only: trust a private CA, as a PEM bundle (s3_ca_file) or an + # OpenSSL-hashed certificate directory (s3_ca_path). Empty uses the + # system trust store. https always verifies the peer either way. + s3_ca_file: str = "" + s3_ca_path: str = "" # The name packs are indexed under; readers resolve it to this store. store_id: str = "s3" @@ -131,6 +136,16 @@ def __post_init__(self) -> None: "never downgrades TLS; leave it False for https://") else: raise ValueError("s3_endpoint must start with http:// or https://") + for name in ("s3_ca_file", "s3_ca_path"): + value = getattr(self, name) + if type(value) is not str: + raise TypeError(f"{name} must be a str (empty for the system " + "trust store)") + # The native client refuses the same pairing: a CA on http:// + # reads as "this is TLS" while credentials go in the clear. + if value and not self.s3_endpoint.startswith("https://"): + raise ValueError(f"{name} applies only to an https:// " + "s3_endpoint") if type(self.clickhouse_port) is not int or not 0 < self.clickhouse_port < 65536: raise ValueError("clickhouse_port must be in 1..65535") _positive("poll_interval_s", self.poll_interval_s, float) @@ -158,6 +173,8 @@ def _native_dict(self) -> dict[str, Any]: "s3_secret_key": self.s3_secret_key, "s3_session_token": self.s3_session_token, "s3_allow_insecure_http": self.s3_allow_insecure_http, + "s3_ca_file": self.s3_ca_file, + "s3_ca_path": self.s3_ca_path, "store_id": self.store_id, "clickhouse_host": self.clickhouse_host, "clickhouse_port": self.clickhouse_port, diff --git a/tests/test_native_capture_storage_live.py b/tests/test_native_capture_storage_live.py index 390af9945..43d2aabca 100644 --- a/tests/test_native_capture_storage_live.py +++ b/tests/test_native_capture_storage_live.py @@ -35,7 +35,7 @@ # Module-level so the fake-S3 fixture registers in this module. from tests.test_native_s3_client import ( # noqa: E402 - ACCESS, BUCKET, REGION, SECRET, STATE, fake_s3, + ACCESS, BUCKET, REGION, SECRET, STATE, fake_s3, fake_s3_tls, private_ca, ) REPO = Path(__file__).resolve().parents[1] @@ -768,3 +768,88 @@ def _native(spool, holder): assert snapshot["upload_failures"] >= 1, snapshot assert snapshot["lease_renewals"] >= 3, snapshot + + +# --- https with a private CA ---------------------------------------------------- + + +def test_a_64_mib_pack_over_https_with_a_private_ca_hydrates_exactly( + fake_s3_tls, private_ca, tmp_path): + """s3_ca_file reaches both halves: the service's upload and the reader. + + 17 x 4 MiB records stage as one pack over the client's 64 MiB multipart + threshold. The service uploads it over TLS to a store whose certificate + only the private CA vouches for, indexes it, and the reader hydrates + every capture byte-equal over the same TLS. A reader without the CA is + refused by the store's certificate. + """ + import random + + from dmi.storage.capture import CaptureMetadata + + ca_file, _ca_path, _cert, _key = private_ca + records, record_bytes = 17, 4 << 20 + spool_root = tmp_path / "spool" + payloads = {} + sink = _Driver(SINK_DRIVER) + try: + assert sink.call( + op="open", root=str(spool_root), max_bytes=1 << 40, + max_queue_records=records, max_queue_bytes=2 * records * record_bytes, + max_pack_bytes=2 * records * record_bytes, + max_pack_records=records, max_linger_ns=60_000_000_000, + overload="drop_newest", admission_timeout=-1)["ok"] + for index in range(records): + metadata = CaptureMetadata( + capture_id=f"tls-{index:04d}", tenant_id="t", + experiment_id="e", run_id="r", session_id="s", + request_id=f"q{index}", sequence_id=f"n{index}", + model_id="m", model_revision="mr", adapter_revision=None, + capture_policy_version="v", hook_name="resid_post", + layer_number=0, producer_rank=0, step_number=index, + token_start=index, token_end=index + 1, batch_position=0, + dtype="uint8", shape=(record_bytes,), + captured_at_ns=1_700_000_000_000_000_000 + index) + payload = random.Random(index).randbytes(record_bytes) + response = sink.call(op="submit", metadata=metadata.to_mapping(), + payload_b64=base64.b64encode(payload).decode()) + assert response["admission"] == "accepted", response + payloads[metadata.capture_id] = payload + assert sink.call(op="flush", timeout=60)["ok"] + assert sink.call(op="close", timeout=60)["snapshot"][ + "persisted_records"] == records + finally: + sink.close() + [pack] = _ready(spool_root) + assert pack.stat().st_size >= 64 << 20 + + with _catalog() as (_client, catalog): + config = _storage_config(fake_s3_tls, catalog.table_prefix, + s3_allow_insecure_http=False, + s3_ca_file=ca_file) + service = _service(config, spool_root) + service.start() + try: + service.flush(60.0) + snapshot = service.snapshot() + finally: + service.stop() + assert snapshot["uploaded_packs"] == 1, snapshot + assert snapshot["indexed_rows"] == records, snapshot + parts = [call for call in STATE.calls + if call["method"] == "PUT" and "partNumber=" in call["path"]] + assert len(parts) >= 5, len(parts) # multipart, 16 MiB parts + + reader = _reader(config) + selection = reader.select(tenant_id="t") + captures = {capture.descriptor["capture_id"]: capture.payload + for capture in reader.read(selection, byte_limit=1 << 30)} + assert sorted(captures) == sorted(payloads) + for capture_id, payload in payloads.items(): + assert captures[capture_id] == payload, capture_id + + untrusted = _reader(_storage_config( + fake_s3_tls, catalog.table_prefix, s3_allow_insecure_http=False)) + with pytest.raises(Exception, match="(?i)certificate"): + untrusted.read(untrusted.select(tenant_id="t"), + byte_limit=1 << 30) diff --git a/tests/test_native_capture_storage_wiring.py b/tests/test_native_capture_storage_wiring.py index ce59315b4..061afa3f7 100644 --- a/tests/test_native_capture_storage_wiring.py +++ b/tests/test_native_capture_storage_wiring.py @@ -113,6 +113,75 @@ def test_https_with_the_insecure_flag_is_refused(): s3_allow_insecure_http=True) +# --- a private CA for an https object store ---------------------------------- + + +def test_a_private_ca_reaches_the_native_service_and_reader(monkeypatch): + from dmi.storage import native_capture + + seen = [] + + class _Native: + def __init__(self, config): + seen.append(dict(config)) + + monkeypatch.setattr( + native_capture, "_load_native_store_extension", + lambda: SimpleNamespace(StorageService=_Native, CaptureReader=_Native, + SEARCH_ITEM_COLUMNS=())) + config = _storage_config(s3_ca_file="/etc/dmi/ca.pem", + s3_ca_path="/etc/dmi/ca.d") + native_capture.NativeCaptureStorage(config, spool_root="/tmp/spool", + spool_max_bytes=1 << 30, + sweep_spool=True) + native_capture.NativeCaptureReader(config) + assert [(d["s3_ca_file"], d["s3_ca_path"]) for d in seen] == \ + [("/etc/dmi/ca.pem", "/etc/dmi/ca.d")] * 2 + # Unset means libcurl's default trust store: empty, not absent. + default = _storage_config()._native_dict() + assert (default["s3_ca_file"], default["s3_ca_path"]) == ("", "") + + +@pytest.mark.parametrize("name", ["s3_ca_file", "s3_ca_path"]) +def test_a_ca_on_a_plain_http_endpoint_is_refused(name): + # It would read as "this is TLS" while credentials go in the clear; the + # native client refuses the same pairing. + with pytest.raises(ValueError, match=f"{name}.*https"): + _storage_config(s3_endpoint="http://127.0.0.1:3900", + s3_allow_insecure_http=True, **{name: "/etc/dmi/ca"}) + + +@pytest.mark.parametrize("name", ["s3_ca_file", "s3_ca_path"]) +def test_a_ca_option_must_be_a_string(name): + with pytest.raises(TypeError, match=name): + _storage_config(**{name: None}) + + +def _native_store_or_skip(): + from dmi.storage import native_capture + + try: + return native_capture._load_native_store_extension() + except ImportError: + pytest.skip("_dmi_native_store is not built") + + +@pytest.mark.parametrize("name", ["s3_ca_file", "s3_ca_path"]) +def test_the_native_module_names_a_missing_ca_at_construction(tmp_path, name): + """The real bindings read the CA fields into the client's config. + + No store or catalog is contacted: the client's own validation refuses + the path when the service or reader is built, not at the first upload. + """ + module = _native_store_or_skip() + missing = str(tmp_path / "no-such-ca") + native = _storage_config(**{name: missing})._native_dict() + with pytest.raises(ValueError, match="no-such-ca"): + module.CaptureReader(native) + with pytest.raises(ValueError, match="no-such-ca"): + module.StorageService({**native, "spool_root": str(tmp_path / "spool")}) + + def _fake_reader(monkeypatch): from dmi.storage import native_capture diff --git a/tests/test_native_s3_client.py b/tests/test_native_s3_client.py index f7901231a..f47e81589 100644 --- a/tests/test_native_s3_client.py +++ b/tests/test_native_s3_client.py @@ -17,6 +17,8 @@ import base64 import hashlib import json +import shutil +import ssl import subprocess import sys import threading @@ -176,6 +178,11 @@ def under(name: str) -> bool: STATE.fault_counts[key] = n + 1 if under("fault/once-500") and n == 0: return 500, b"boom" + if under("fault/slow-once-500") and n == 0: + # Outlive the one-second resolution of x-amz-date, so a retry + # that reuses the first attempt's signature is visible. + time.sleep(1.2) + return 500, b"boom" if under("fault/always-500"): return 500, b"boom" if under("fault/forbidden"): @@ -354,7 +361,13 @@ def _handle_complete(self, key: str, query: dict, body: bytes): if not numbers: self._send(400, {}, b"invalid xml") return - assembled = b"".join(upload["parts"][n] for n in sorted(numbers)) + ordered = sorted(numbers) + # S3's rule: every part but the last is at least 5 MiB. + if any(len(upload["parts"][n]) < 5 * 1024 * 1024 + for n in ordered[:-1]): + self._send(400, {}, b"EntityTooSmall") + return + assembled = b"".join(upload["parts"][n] for n in ordered) meta = upload["meta"] with STATE.lock: STATE.objects[key] = { @@ -371,12 +384,16 @@ def _handle_abort(self, key: str, query: dict): self._send(204, {}) -@pytest.fixture() -def fake_s3(): +def _reset_state(): STATE.objects.clear() STATE.uploads.clear() STATE.calls.clear() STATE.fault_counts.clear() + + +@pytest.fixture() +def fake_s3(): + _reset_state() server = ThreadingHTTPServer(("127.0.0.1", 0), FakeS3Handler) thread = threading.Thread(target=server.serve_forever, daemon=True) thread.start() @@ -384,6 +401,122 @@ def fake_s3(): server.shutdown() +# --- TLS --------------------------------------------------------------------- +# +# The same signature-verifying fake, behind TLS with a server certificate +# issued by a private CA generated here. Nothing about the CA is installed +# anywhere: a client trusts it only when told to (ca_file / ca_path). + + +def _openssl(*args: str, cwd: Path) -> str: + return subprocess.run(["openssl", *args], cwd=cwd, check=True, + capture_output=True, text=True).stdout + + +@pytest.fixture(scope="session") +def private_ca(tmp_path_factory): + """(ca_file, ca_path, server_cert, server_key) for 127.0.0.1.""" + if shutil.which("openssl") is None: + pytest.skip("the openssl CLI is needed to mint the test CA") + root = tmp_path_factory.mktemp("private-ca") + # Own config files, not the system openssl.cnf: its v3_ca section adds a + # basicConstraints of its own, and a duplicated extension makes OpenSSL + # reject the CA ("unable to get local issuer certificate"). + (root / "ca.cnf").write_text( + "[req]\nprompt=no\ndistinguished_name=dn\nx509_extensions=v3_ca\n" + "[dn]\nCN=DMI test private CA\n" + "[v3_ca]\nbasicConstraints=critical,CA:TRUE\n" + "keyUsage=critical,keyCertSign,cRLSign\n" + "subjectKeyIdentifier=hash\n") + (root / "server.cnf").write_text( + "[req]\nprompt=no\ndistinguished_name=dn\n[dn]\nCN=127.0.0.1\n") + (root / "server.ext").write_text( + "basicConstraints=CA:FALSE\n" + "keyUsage=critical,digitalSignature,keyEncipherment\n" + "extendedKeyUsage=serverAuth\n" + "subjectAltName=IP:127.0.0.1,DNS:localhost\n" + "authorityKeyIdentifier=keyid\n") + _openssl("req", "-config", "ca.cnf", "-x509", "-newkey", "rsa:2048", + "-nodes", "-keyout", "ca.key", "-out", "ca.pem", "-days", "2", + cwd=root) + _openssl("req", "-config", "server.cnf", "-new", "-newkey", "rsa:2048", + "-nodes", "-keyout", "server.key", "-out", "server.csr", + cwd=root) + _openssl("x509", "-req", "-in", "server.csr", "-CA", "ca.pem", + "-CAkey", "ca.key", "-CAcreateserial", "-out", "server.pem", + "-days", "2", "-extfile", "server.ext", cwd=root) + _openssl("verify", "-CAfile", "ca.pem", "server.pem", cwd=root) + # CURLOPT_CAPATH reads an OpenSSL-hashed directory: .0. + ca_dir = root / "ca-dir" + ca_dir.mkdir() + subject_hash = _openssl("x509", "-hash", "-noout", "-in", "ca.pem", + cwd=root).strip() + shutil.copy(root / "ca.pem", ca_dir / f"{subject_hash}.0") + return (str(root / "ca.pem"), str(ca_dir), str(root / "server.pem"), + str(root / "server.key")) + + +class _QuietTLSServer(ThreadingHTTPServer): + def handle_error(self, request, client_address): + # A client that refuses the certificate aborts the handshake; that is + # the outcome under test, not a server fault worth a traceback. + if isinstance(sys.exc_info()[1], (ssl.SSLError, ConnectionError)): + return + super().handle_error(request, client_address) + + +@pytest.fixture() +def fake_s3_tls(private_ca): + """https://127.0.0.1: served with the private CA's certificate.""" + _reset_state() + _ca_file, _ca_path, cert, key = private_ca + context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + context.load_cert_chain(cert, key) + server = _QuietTLSServer(("127.0.0.1", 0), FakeS3Handler) + # The handshake runs lazily, on the handler's thread, so a client that + # hangs or aborts it cannot stall the accept loop. + server.socket = context.wrap_socket(server.socket, server_side=True, + do_handshake_on_connect=False) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + yield f"https://127.0.0.1:{server.server_port}" + server.shutdown() + + +@pytest.fixture() +def fake_s3_tls_wrong_name(private_ca): + """https://127.0.0.1: served with a certificate the private CA + issued for ANOTHER name: trusted chain, mismatched host.""" + _reset_state() + ca_file, _ca_path, _cert, _key = private_ca + root = Path(ca_file).parent + (root / "other.ext").write_text( + "basicConstraints=CA:FALSE\n" + "keyUsage=critical,digitalSignature,keyEncipherment\n" + "extendedKeyUsage=serverAuth\n" + "subjectAltName=DNS:other.example\n" + "authorityKeyIdentifier=keyid\n") + (root / "other.cnf").write_text( + "[req]\nprompt=no\ndistinguished_name=dn\n[dn]\nCN=other.example\n") + if not (root / "other.pem").exists(): + _openssl("req", "-config", "other.cnf", "-new", "-newkey", "rsa:2048", + "-nodes", "-keyout", "other.key", "-out", "other.csr", + cwd=root) + _openssl("x509", "-req", "-in", "other.csr", "-CA", "ca.pem", + "-CAkey", "ca.key", "-CAcreateserial", "-out", "other.pem", + "-days", "2", "-extfile", "other.ext", cwd=root) + _openssl("verify", "-CAfile", "ca.pem", "other.pem", cwd=root) + context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + context.load_cert_chain(str(root / "other.pem"), str(root / "other.key")) + server = _QuietTLSServer(("127.0.0.1", 0), FakeS3Handler) + server.socket = context.wrap_socket(server.socket, server_side=True, + do_handshake_on_connect=False) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + yield f"https://127.0.0.1:{server.server_port}" + server.shutdown() + + def _base(endpoint: str, **overrides) -> dict: request = { "endpoint": endpoint, @@ -447,18 +580,54 @@ def test_put_get_head_delete_round_trip(fake_s3): assert missing["ok"] and not missing["found"] +MIB = 1024 * 1024 + + def test_put_multipart_round_trip(fake_s3): - payload = bytes((i * 7) & 0xFF for i in range(3 * 1024 * 1024)) + """Real part sizes: S3 refuses a part under 5 MiB unless it is the last. + + 12 MiB in 5 MiB parts is two full parts and a 2 MiB tail -- the + smallest shape with both a minimum-size part and a short final one. + """ + payload = bytes((i * 7) & 0xFF for i in range(12 * MIB)) put = _call("put", **_base(fake_s3), key="packs/big.dmi-pack", data_b64=base64.b64encode(payload).decode(), metadata={}, content_type="application/vnd.dmi.pack", - multipart_threshold=1024 * 1024, multipart_chunk=1024 * 1024) + multipart_threshold=5 * MIB, multipart_chunk=5 * MIB) assert put["ok"], put + parts = [call["body_len"] for call in STATE.calls + if call["method"] == "PUT" and "partNumber=" in call["path"]] + assert parts == [5 * MIB, 5 * MIB, 2 * MIB] echo = _call("get", **_base(fake_s3), key="packs/big.dmi-pack", offset=0, length=len(payload)) assert echo["ok"] and base64.b64decode(echo["data_b64"]) == payload +def test_multipart_part_under_5_mib_is_refused_before_any_request(fake_s3): + """The client refuses the part size, as the Python store does (s3.py). + + Left to the server, a sub-5 MiB part fails only at + CompleteMultipartUpload, after every part was sent -- and a payload + under the threshold never notices. The refusal is the configuration's, + so it lands on every operation, before any request. + """ + payload = bytes(3 * MIB) + put = _call("put", **_base(fake_s3), key="packs/small-parts.dmi-pack", + data_b64=base64.b64encode(payload).decode(), metadata={}, + content_type="application/vnd.dmi.pack", + multipart_threshold=MIB, multipart_chunk=MIB) + assert not put["ok"], put + assert "multipart_chunk_bytes" in put["what"], put + assert str(5 * MIB) in put["what"], put + head = _call("head", **_base(fake_s3), key="anything", + multipart_chunk=5 * MIB - 1) + assert not head["ok"] and "multipart_chunk_bytes" in head["what"], head + assert STATE.calls == [] + # Exactly 5 MiB is S3's minimum, and is accepted. + assert _call("head", **_base(fake_s3), key="anything", + multipart_chunk=5 * MIB)["ok"] + + def test_list_pagination(fake_s3): for i in range(5): ack = _call("put", **_base(fake_s3), key=f"v1/pack-{i}", @@ -484,6 +653,37 @@ def test_retry_then_success_on_500(fake_s3): assert put["attempts"] == 2 +def _header(call: dict, name: str) -> str: + return next(v for k, v in call["headers"].items() if k.lower() == name) + + +def test_every_attempt_is_signed_afresh(fake_s3): + """A retry carries its own x-amz-date and signature, both still valid. + + SigV4 binds the signature to x-amz-date, and S3 refuses a request whose + date is more than 15 minutes off. Signing once before the loop replayed + the first attempt's date on every retry, so a slow first attempt (up to + read_timeout_s each, plus backoff) aged every later one. The first + attempt here outlives a second before failing with a 500: a fresh + signature must carry a later date. Both attempts passed the server's + botocore re-signing check -- the 500 is only reached after it -- and + the second one stored the object. + """ + put = _call("put", **_base(fake_s3), key="fault/slow-once-500", + data_b64=base64.b64encode(b"data").decode(), metadata={}, + content_type="application/octet-stream") + assert put["ok"], put + assert put["attempts"] == 2 + attempts = [call for call in STATE.calls + if call["path"].endswith("/fault/slow-once-500")] + assert len(attempts) == 2, attempts + first, second = attempts + assert _header(second, "x-amz-date") > _header(first, "x-amz-date") + assert _header(second, "authorization") != \ + _header(first, "authorization") + assert STATE.objects["fault/slow-once-500"]["body"] == b"data" + + def test_no_retry_on_403(fake_s3): put = _call("put", **_base(fake_s3), key="fault/forbidden", data_b64=base64.b64encode(b"data").decode(), metadata={}, @@ -596,3 +796,75 @@ def test_the_64_bit_boundaries_still_reach_the_transport(fake_s3): length=len(payload)) assert get["ok"], get assert base64.b64decode(get["data_b64"]) == payload + + +# --- https with a private CA -------------------------------------------------- + + +def _tls(endpoint: str, **overrides) -> dict: + return _base(endpoint, insecure=False, **overrides) + + +@pytest.mark.parametrize("trust", ["ca_file", "ca_path"]) +def test_https_round_trip_trusts_a_private_ca(fake_s3_tls, private_ca, trust): + ca_file, ca_path, _cert, _key = private_ca + anchor = {"ca_file": ca_file} if trust == "ca_file" else {"ca_path": ca_path} + payload = bytes(range(256)) * 16 + meta = {"dmi-format": "dmi-pack-v1"} + put = _call("put", **_tls(fake_s3_tls, **anchor), key="tls/a", + data_b64=base64.b64encode(payload).decode(), metadata=meta, + content_type="application/octet-stream") + assert put["ok"], put + head = _call("head", **_tls(fake_s3_tls, **anchor), key="tls/a") + assert head["ok"] and head["found"] and head["metadata"] == meta, head + get = _call("get", **_tls(fake_s3_tls, **anchor), key="tls/a", + offset=0, length=len(payload)) + assert get["ok"], get + assert base64.b64decode(get["data_b64"]) == payload + + +def test_https_without_the_private_ca_is_refused(fake_s3_tls): + """libcurl's default trust store does not know the CA: no request lands. + + A certificate failure is not transient, so it is not retried. + """ + put = _call("put", **_tls(fake_s3_tls), key="tls/untrusted", + data_b64=base64.b64encode(b"data").decode(), metadata={}, + content_type="application/octet-stream") + assert not put["ok"], put + assert "certificate" in put["what"].lower(), put + assert put["attempts"] == 1, put + assert STATE.calls == [] and STATE.objects == {} + + +def test_https_refuses_a_certificate_for_another_host( + fake_s3_tls_wrong_name, private_ca): + """The chain is trusted (the private CA issued it), but it names + other.example, not 127.0.0.1: host-name verification must refuse it + before any request, with no retry.""" + ca_file, _ca_path, _cert, _key = private_ca + head = _call("head", **_tls(fake_s3_tls_wrong_name, ca_file=ca_file), + key="tls/wrong-name") + assert not head["ok"], head + assert head["attempts"] == 1, head + assert STATE.calls == [] + + +def test_ca_options_on_plain_http_are_refused(fake_s3, private_ca): + ca_file, ca_path, _cert, _key = private_ca + for anchor in ({"ca_file": ca_file}, {"ca_path": ca_path}): + head = _call("head", **_base(fake_s3, **anchor), key="anything") + assert not head["ok"], head + assert "https" in head["what"], head + assert STATE.calls == [] + + +def test_a_missing_ca_is_named_before_any_request(fake_s3_tls, tmp_path): + missing = tmp_path / "no-such-ca.pem" + head = _call("head", **_tls(fake_s3_tls, ca_file=str(missing)), + key="anything") + assert not head["ok"] and str(missing) in head["what"], head + head = _call("head", **_tls(fake_s3_tls, ca_path=str(missing)), + key="anything") + assert not head["ok"] and str(missing) in head["what"], head + assert STATE.calls == [] diff --git a/tests/test_native_uploader.py b/tests/test_native_uploader.py index c17d9f809..36ccebfb5 100644 --- a/tests/test_native_uploader.py +++ b/tests/test_native_uploader.py @@ -13,6 +13,7 @@ import base64 import json +import random import subprocess import sys from pathlib import Path @@ -53,6 +54,8 @@ _base as _client_base, _call as _store_call, fake_s3, + fake_s3_tls, + private_ca, ) @@ -592,3 +595,106 @@ def test_upload_pending_keeps_the_64_bit_limits(fake_s3, tmp_path): finally: sink.close() store.close() + + +# --- a large pack over https with a private CA ---------------------------------- + +MIB = 1024 * 1024 + + +def _stage_large_pack(root: Path, records: int, record_bytes: int) -> dict: + """One pack of `records` random payloads, staged by the native sink.""" + sink = DriverSession(SINK_DRIVER) + try: + assert sink.call( + op="open", root=str(root), max_bytes=1 << 40, + max_queue_records=records, max_queue_bytes=2 * records * record_bytes, + max_pack_bytes=2 * records * record_bytes, + max_pack_records=records, max_linger_ns=60_000_000_000, + overload="drop_newest", admission_timeout=-1, + )["ok"] + for index in range(records): + meta = CaptureMetadata( + capture_id=f"large-{index:04d}", tenant_id="t", + experiment_id="e", run_id="r", session_id="s", + request_id=f"q{index}", sequence_id=f"n{index}", model_id="m", + model_revision="mr", adapter_revision=None, + capture_policy_version="v", hook_name="h", layer_number=0, + producer_rank=0, step_number=index, token_start=index, + token_end=index + 1, batch_position=0, dtype="uint8", + shape=(record_bytes,), + captured_at_ns=1_700_000_000_000_000_000 + index, + ) + # Random bytes: nothing in the path can shrink the pack below + # the multipart threshold. + payload = random.Random(index).randbytes(record_bytes) + response = sink.call(op="submit", metadata=meta.to_mapping(), + payload_b64=base64.b64encode(payload).decode()) + assert response["admission"] == "accepted", response + assert sink.call(op="flush", timeout=60)["ok"] + snapshot = sink.call(op="close", timeout=60)["snapshot"] + assert snapshot["persisted_records"] == records, snapshot + finally: + sink.close() + recover = subprocess.run( + [str(STORE_DRIVER.parent / "conformance_spool")], + input=json.dumps({"op": "recover", "root": str(root), + "max_bytes": 1 << 40}) + "\n", + capture_output=True, text=True, timeout=60, + ) + entries = json.loads(recover.stdout.strip())["staged"] + assert len(entries) == 1, entries + return entries[0] + + +def test_a_64_mib_pack_uploads_multipart_over_https_with_a_private_ca( + fake_s3_tls, private_ca, tmp_path): + """The uploader's default multipart shape, over TLS to a private CA. + + 17 x 4 MiB records make one pack over the client's 64 MiB threshold, so + it goes up as 16 MiB parts, verified end to end: the server checks every + part's signature, and the object reads back over the same TLS byte for + byte. Without the CA the upload is refused and the pack stays staged. + """ + ca_file, _ca_path, _cert, _key = private_ca + root = tmp_path / "spool" + staged = _stage_large_pack(root, records=17, record_bytes=4 * MIB) + assert staged["object_bytes"] >= 64 * MIB, staged + staged_bytes = Path(staged["path"]).read_bytes() + assert len(staged_bytes) == staged["object_bytes"] + + store = DriverSession(STORE_DRIVER) + try: + untrusted = _upload_pending(store, fake_s3_tls, root, insecure=False, + upload_max_attempts=1) + # refs[i] pairs with failures[i]; exactly one of them is set. + assert untrusted["ok"], untrusted + [failure] = untrusted["failures"] + assert untrusted["refs"][0]["pack_id"] == "", untrusted + assert failure["pack_id"] == staged["pack_id"], failure + assert "certificate" in failure["error"].lower(), failure + assert STATE.calls == [] and STATE.objects == {} + assert Path(staged["path"]).exists() + + trusted = _upload_pending(store, fake_s3_tls, root, insecure=False, + ca_file=ca_file) + assert trusted["ok"], trusted + assert trusted["failures"][0]["pack_id"] == "", trusted + [ref] = trusted["refs"] + assert ref["pack_id"] == staged["pack_id"], ref + assert ref["checksum"] == staged["checksum"] + assert ref["object_bytes"] == staged["object_bytes"] + assert not Path(staged["path"]).exists() + finally: + store.close() + + parts = [call["body_len"] for call in STATE.calls + if call["method"] == "PUT" and "partNumber=" in call["path"]] + full, tail = divmod(staged["object_bytes"], 16 * MIB) + assert parts == [16 * MIB] * full + ([tail] if tail else []), parts + + fetched = _store_call( + "get", **_client_base(fake_s3_tls, insecure=False, ca_file=ca_file), + key=staged["object_key"], offset=0, length=staged["object_bytes"]) + assert fetched["ok"], {k: v for k, v in fetched.items() if k != "data_b64"} + assert base64.b64decode(fetched["data_b64"]) == staged_bytes