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
128 changes: 128 additions & 0 deletions src/datacustomcode/named_credential/direct/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,11 +19,15 @@
from __future__ import annotations

import base64
import datetime
import hashlib
import hmac
from typing import (
TYPE_CHECKING,
Any,
Dict,
)
import urllib.parse

from requests.auth import AuthBase

Expand All @@ -32,6 +36,8 @@
if TYPE_CHECKING:
from requests.models import PreparedRequest

_SIGV4_ALGORITHM = "AWS4-HMAC-SHA256"


class DynamicAuthHandler(AuthBase):
def __init__(self, cred_config: Dict[str, Any]) -> None:
Expand All @@ -57,7 +63,129 @@ def __call__(self, request: PreparedRequest) -> PreparedRequest:
)
request.headers["Authorization"] = f"Bearer {bearer}"

elif self.auth_type == AuthType.AWS_SIG_V4.value:
self._sign_aws_sigv4(request)

else:
raise ValueError(f"Unsupported auth_type '{self.auth_type}'.")

return request

def _sign_aws_sigv4(self, request: PreparedRequest) -> None:
"""Sign ``request`` with AWS Signature Version 4.

Requires ``aws_access_key_id``, ``aws_secret_access_key``, ``aws_region``,
and ``aws_service`` in the credential config; ``aws_session_token`` is
optional (for temporary credentials). The signed date, payload hash, and
(when present) session token are added as ``x-amz-*`` headers so the sent
request matches what was signed.
"""
access_key = self.config.get("aws_access_key_id")
secret_key = self.config.get("aws_secret_access_key")
region = self.config.get("aws_region")
service = self.config.get("aws_service")
session_token = self.config.get("aws_session_token")
missing = [
name
for name, value in (
("aws_access_key_id", access_key),
("aws_secret_access_key", secret_key),
("aws_region", region),
("aws_service", service),
)
if not value
]
if missing:
raise ValueError(f"'{self.auth_type}' auth requires {', '.join(missing)}.")
access_key = str(access_key)
secret_key = str(secret_key)
region = str(region)
service = str(service)

parsed = urllib.parse.urlsplit(str(request.url or ""))
host = parsed.netloc
canonical_uri = urllib.parse.quote(parsed.path or "/", safe="/-_.~")
canonical_query = _canonical_query_string(parsed.query)

body = request.body or b""
if isinstance(body, str):
body = body.encode("utf-8")
payload_hash = hashlib.sha256(body).hexdigest()

now = datetime.datetime.now(datetime.timezone.utc)
amz_date = now.strftime("%Y%m%dT%H%M%SZ")
datestamp = now.strftime("%Y%m%d")

request.headers["x-amz-date"] = amz_date
request.headers["x-amz-content-sha256"] = payload_hash
if session_token:
request.headers["x-amz-security-token"] = session_token

signed = {
"host": host,
"x-amz-content-sha256": payload_hash,
"x-amz-date": amz_date,
}
if session_token:
signed["x-amz-security-token"] = session_token
signed_headers = ";".join(sorted(signed))
canonical_headers = "".join(
f"{name}:{signed[name]}\n" for name in sorted(signed)
)

canonical_request = "\n".join(
[
request.method or "GET",
canonical_uri,
canonical_query,
canonical_headers,
signed_headers,
payload_hash,
]
)
credential_scope = f"{datestamp}/{region}/{service}/aws4_request"
string_to_sign = "\n".join(
[
_SIGV4_ALGORITHM,
amz_date,
credential_scope,
hashlib.sha256(canonical_request.encode("utf-8")).hexdigest(),
]
)
signing_key = _derive_signing_key(secret_key, datestamp, region, service)
signature = hmac.new(
signing_key, string_to_sign.encode("utf-8"), hashlib.sha256
).hexdigest()

request.headers["Authorization"] = (
f"{_SIGV4_ALGORITHM} Credential={access_key}/{credential_scope}, "
f"SignedHeaders={signed_headers}, Signature={signature}"
)


def _canonical_query_string(query: str) -> str:
"""Build the AWS Sig V4 canonical query string from a raw query string."""
pairs = urllib.parse.parse_qsl(query, keep_blank_values=True)
encoded = [
(
urllib.parse.quote(key, safe="-_.~"),
urllib.parse.quote(value, safe="-_.~"),
)
for key, value in pairs
]
encoded.sort()
return "&".join(f"{key}={value}" for key, value in encoded)


def _derive_signing_key(
secret_key: str, datestamp: str, region: str, service: str
) -> bytes:
"""Derive the AWS Sig V4 signing key via the chained HMAC-SHA256 sequence."""

def _hmac(key: bytes, msg: str) -> bytes:
return hmac.new(key, msg.encode("utf-8"), hashlib.sha256).digest()

k_date = _hmac(f"AWS4{secret_key}".encode(), datestamp)
k_region = _hmac(k_date, region)
k_service = _hmac(k_region, service)
return _hmac(k_service, "aws4_request")
1 change: 1 addition & 0 deletions src/datacustomcode/named_credential/direct/credentials.py
Original file line number Diff line number Diff line change
Expand Up @@ -41,6 +41,7 @@
class AuthType(str, Enum):
"""External Credential auth types supported by External Services."""

AWS_SIG_V4 = "AwsSv4"
BASIC = "Basic"
CUSTOM = "Custom"
JWT = "Jwt"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -122,15 +122,19 @@ matching entries in `config.json`) to point at your own DLOs.
`auth_type` selects how auth is injected for local testing. It should mirror the
External Credential your Named Credential uses in the org, so local and deployed
runs behave the same. This example uses `Custom` (Gemini's `X-goog-api-key`);
all four supported types:
all supported types:

| `auth_type` | Fields read | Header sent |
| ----------- | ------------------------------- | --------------------------------------- |
| `Basic` | `username`, `password` | `Authorization: Basic <base64 user:pw>` |
| `Custom` | `custom_headers` (sent verbatim)| the headers you list |
| `OAuth` | `access_token` or `token` | `Authorization: Bearer <token>` |
| `Jwt` | `access_token` or `token` | `Authorization: Bearer <token>` |
| `AwsSv4` | `aws_access_key_id`, `aws_secret_access_key`, `aws_region`, `aws_service`, optional `aws_session_token` | `Authorization: AWS4-HMAC-SHA256 ...` plus `x-amz-date` / `x-amz-content-sha256` (and `x-amz-security-token` when a session token is set) |

`OAuth`/`Jwt` take a token you supply for the local run — the SDK does not fetch
or refresh it. In the Data Cloud runtime the Named Credential handles token
acquisition; this local config only stands in for that during testing.

`AwsSv4` signs the request with AWS Signature Version 4 using the keys you
supply, mirroring an AWS Signature Version 4 External Credential in the org.
136 changes: 136 additions & 0 deletions tests/test_named_credential_direct.py
Original file line number Diff line number Diff line change
Expand Up @@ -87,6 +87,142 @@ def test_unsupported_auth_type_raises_value_error(self):
handler(_prepared_request())


def _reference_sigv4(config, method, url, body):
"""Independent SigV4 reference used to validate DynamicAuthHandler output."""
import datetime
import hashlib
import hmac
import urllib.parse

parsed = urllib.parse.urlsplit(url)
payload = body.encode("utf-8") if isinstance(body, str) else (body or b"")
payload_hash = hashlib.sha256(payload).hexdigest()
now = datetime.datetime(2024, 1, 2, 3, 4, 5, tzinfo=datetime.timezone.utc)
amz_date = now.strftime("%Y%m%dT%H%M%SZ")
datestamp = now.strftime("%Y%m%d")

signed = {
"host": parsed.netloc,
"x-amz-content-sha256": payload_hash,
"x-amz-date": amz_date,
}
token = config.get("aws_session_token")
if token:
signed["x-amz-security-token"] = token
signed_headers = ";".join(sorted(signed))
canonical_headers = "".join(f"{n}:{signed[n]}\n" for n in sorted(signed))

pairs = sorted(
(urllib.parse.quote(k, safe="-_.~"), urllib.parse.quote(v, safe="-_.~"))
for k, v in urllib.parse.parse_qsl(parsed.query, keep_blank_values=True)
)
canonical_query = "&".join(f"{k}={v}" for k, v in pairs)
canonical_uri = urllib.parse.quote(parsed.path or "/", safe="/-_.~")
canonical_request = "\n".join(
[
method,
canonical_uri,
canonical_query,
canonical_headers,
signed_headers,
payload_hash,
]
)
scope = f"{datestamp}/{config['aws_region']}/{config['aws_service']}/aws4_request"
string_to_sign = "\n".join(
[
"AWS4-HMAC-SHA256",
amz_date,
scope,
hashlib.sha256(canonical_request.encode()).hexdigest(),
]
)

def _h(key, msg):
return hmac.new(key, msg.encode(), hashlib.sha256).digest()

k = _h(f"AWS4{config['aws_secret_access_key']}".encode(), datestamp)
k = _h(k, config["aws_region"])
k = _h(k, config["aws_service"])
k = _h(k, "aws4_request")
signature = hmac.new(k, string_to_sign.encode(), hashlib.sha256).hexdigest()
return (
(
f"AWS4-HMAC-SHA256 Credential={config['aws_access_key_id']}/{scope}, "
f"SignedHeaders={signed_headers}, Signature={signature}"
),
amz_date,
payload_hash,
)


class TestAwsSigV4Auth:
_CONFIG: ClassVar[dict] = {
"auth_type": AuthType.AWS_SIG_V4.value,
"aws_access_key_id": "AKIDEXAMPLE",
"aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY",
"aws_region": "us-east-1",
"aws_service": "s3",
}

@staticmethod
def _freeze_clock(monkeypatch):
import datetime as _dt

from datacustomcode.named_credential.direct import auth as auth_mod

frozen = _dt.datetime(2024, 1, 2, 3, 4, 5, tzinfo=_dt.timezone.utc)

class _FrozenDatetime(_dt.datetime):
@classmethod
def now(cls, tz=None):
return frozen if tz is None else frozen.astimezone(tz)

monkeypatch.setattr(auth_mod.datetime, "datetime", _FrozenDatetime)

def _sign(self, monkeypatch, config, method="GET", url=None, body=None):
self._freeze_clock(monkeypatch)
url = url or "https://bucket.s3.amazonaws.com/key"
request = PreparedRequest()
request.prepare(method=method, url=url, data=body)
return DynamicAuthHandler(config)(request)

def test_matches_reference_signature(self, monkeypatch):
url = "https://bucket.s3.amazonaws.com/some/key?b=2&a=1"
request = self._sign(monkeypatch, self._CONFIG, "PUT", url, "payload")
expected, amz_date, payload_hash = _reference_sigv4(
self._CONFIG, "PUT", url, "payload"
)
assert request.headers["Authorization"] == expected
assert request.headers["x-amz-date"] == amz_date
assert request.headers["x-amz-content-sha256"] == payload_hash
assert "x-amz-security-token" not in request.headers

def test_empty_body_hashes_empty_string(self, monkeypatch):
request = self._sign(monkeypatch, self._CONFIG)
assert (
request.headers["x-amz-content-sha256"]
== "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855"
)

def test_session_token_is_signed(self, monkeypatch):
config = {**self._CONFIG, "aws_session_token": "SESSIONTOKEN"}
request = self._sign(monkeypatch, config)
assert request.headers["x-amz-security-token"] == "SESSIONTOKEN"
assert "x-amz-security-token" in request.headers["Authorization"]
expected, _, _ = _reference_sigv4(config, "GET", request.url, None)
assert request.headers["Authorization"] == expected

@pytest.mark.parametrize(
"missing",
["aws_access_key_id", "aws_secret_access_key", "aws_region", "aws_service"],
)
def test_missing_field_raises(self, monkeypatch, missing):
config = {k: v for k, v in self._CONFIG.items() if k != missing}
with pytest.raises(ValueError, match=missing):
self._sign(monkeypatch, config)


class TestCredentialStore:
def test_get_returns_config(self, tmp_path, monkeypatch):
cred_file = tmp_path / "external_callout_config.json"
Expand Down
Loading