|
10 | 10 | import base64 |
11 | 11 | import hashlib |
12 | 12 | import logging |
| 13 | +import textwrap |
13 | 14 | from cryptography.hazmat.primitives import serialization |
14 | 15 | from cryptography import x509 |
15 | 16 | from nrfcloud_utils.cli_helpers import write_file, setup_logging |
@@ -70,10 +71,124 @@ def base64_decode(string): |
70 | 71 | """ |
71 | 72 | add padding before decoding. |
72 | 73 | """ |
73 | | - padding = 4 - (len(string) % 4) |
| 74 | + # Base64url data can omit padding; only add what is required. |
| 75 | + padding = (-len(string)) % 4 |
74 | 76 | string = string + ("=" * padding) |
75 | 77 | return base64.urlsafe_b64decode(string) |
76 | 78 |
|
| 79 | +def _decode_tlv_length(data, offset): |
| 80 | + """Decode DER length and return (value_length, header_length).""" |
| 81 | + first = data[offset] |
| 82 | + if first < 0x80: |
| 83 | + return first, 1 |
| 84 | + |
| 85 | + nbytes = first & 0x7F |
| 86 | + if nbytes == 0: |
| 87 | + raise ValueError("Indefinite DER length is not supported") |
| 88 | + |
| 89 | + value_len = 0 |
| 90 | + for i in range(nbytes): |
| 91 | + value_len = (value_len << 8) | data[offset + 1 + i] |
| 92 | + return value_len, 1 + nbytes |
| 93 | + |
| 94 | +def _encode_tlv_length(value_len): |
| 95 | + """Encode DER length bytes for value_len.""" |
| 96 | + if value_len < 0x80: |
| 97 | + return bytes([value_len]) |
| 98 | + |
| 99 | + out = [] |
| 100 | + value = value_len |
| 101 | + while value > 0: |
| 102 | + out.append(value & 0xFF) |
| 103 | + value >>= 8 |
| 104 | + out.reverse() |
| 105 | + return bytes([0x80 | len(out), *out]) |
| 106 | + |
| 107 | +def _normalize_ecdsa_csr_der(der_bytes): |
| 108 | + """ |
| 109 | + Normalize CSR DER that encodes ecdsa-with-SHA256 with NULL params. |
| 110 | +
|
| 111 | + Older modem output may include signatureAlgorithm = SEQUENCE(OID, NULL). |
| 112 | + For ECDSA-with-SHA* this must be encoded without parameters. |
| 113 | + """ |
| 114 | + # ecdsa-with-SHA256 OID (1.2.840.10045.4.3.2) |
| 115 | + ecdsa_sha256_oid_tlv = b"\x06\x08\x2A\x86\x48\xCE\x3D\x04\x03\x02" |
| 116 | + |
| 117 | + if len(der_bytes) < 8 or der_bytes[0] != 0x30: |
| 118 | + return None |
| 119 | + |
| 120 | + try: |
| 121 | + outer_len, outer_len_hdr = _decode_tlv_length(der_bytes, 1) |
| 122 | + except (IndexError, ValueError): |
| 123 | + return None |
| 124 | + |
| 125 | + outer_start = 1 + outer_len_hdr |
| 126 | + outer_end = outer_start + outer_len |
| 127 | + if outer_end != len(der_bytes): |
| 128 | + return None |
| 129 | + |
| 130 | + # certificationRequestInfo |
| 131 | + cri_tag_idx = outer_start |
| 132 | + if der_bytes[cri_tag_idx] != 0x30: |
| 133 | + return None |
| 134 | + try: |
| 135 | + cri_len, cri_len_hdr = _decode_tlv_length(der_bytes, cri_tag_idx + 1) |
| 136 | + except (IndexError, ValueError): |
| 137 | + return None |
| 138 | + cri_start = cri_tag_idx |
| 139 | + cri_value_start = cri_tag_idx + 1 + cri_len_hdr |
| 140 | + cri_end = cri_value_start + cri_len |
| 141 | + |
| 142 | + # signatureAlgorithm |
| 143 | + sig_alg_tag_idx = cri_end |
| 144 | + if sig_alg_tag_idx >= outer_end or der_bytes[sig_alg_tag_idx] != 0x30: |
| 145 | + return None |
| 146 | + try: |
| 147 | + sig_alg_len, sig_alg_len_hdr = _decode_tlv_length(der_bytes, sig_alg_tag_idx + 1) |
| 148 | + except (IndexError, ValueError): |
| 149 | + return None |
| 150 | + sig_alg_value_start = sig_alg_tag_idx + 1 + sig_alg_len_hdr |
| 151 | + sig_alg_end = sig_alg_value_start + sig_alg_len |
| 152 | + if sig_alg_end > outer_end: |
| 153 | + return None |
| 154 | + |
| 155 | + sig_alg_value = der_bytes[sig_alg_value_start:sig_alg_end] |
| 156 | + # Only normalize the exact legacy ECDSA+NULL form. |
| 157 | + if not (sig_alg_value.startswith(ecdsa_sha256_oid_tlv) and sig_alg_value.endswith(b"\x05\x00")): |
| 158 | + return None |
| 159 | + if len(sig_alg_value) != len(ecdsa_sha256_oid_tlv) + 2: |
| 160 | + return None |
| 161 | + |
| 162 | + # signature BIT STRING (keep as-is) |
| 163 | + sig_tag_idx = sig_alg_end |
| 164 | + if sig_tag_idx >= outer_end or der_bytes[sig_tag_idx] != 0x03: |
| 165 | + return None |
| 166 | + try: |
| 167 | + sig_len, sig_len_hdr = _decode_tlv_length(der_bytes, sig_tag_idx + 1) |
| 168 | + except (IndexError, ValueError): |
| 169 | + return None |
| 170 | + sig_value_start = sig_tag_idx + 1 + sig_len_hdr |
| 171 | + sig_end = sig_value_start + sig_len |
| 172 | + if sig_end != outer_end: |
| 173 | + return None |
| 174 | + |
| 175 | + cri_tlv = der_bytes[cri_start:cri_end] |
| 176 | + normalized_sig_alg = b"\x30" + _encode_tlv_length(len(ecdsa_sha256_oid_tlv)) + ecdsa_sha256_oid_tlv |
| 177 | + sig_tlv = der_bytes[sig_tag_idx:sig_end] |
| 178 | + |
| 179 | + new_outer_value = cri_tlv + normalized_sig_alg + sig_tlv |
| 180 | + return b"\x30" + _encode_tlv_length(len(new_outer_value)) + new_outer_value |
| 181 | + |
| 182 | +def _csr_der_to_pem(der_bytes): |
| 183 | + """Convert CSR DER bytes to PEM without changing the original DER payload.""" |
| 184 | + b64 = base64.b64encode(der_bytes).decode("ascii") |
| 185 | + wrapped = "\n".join(textwrap.wrap(b64, 64)) |
| 186 | + return ( |
| 187 | + "-----BEGIN CERTIFICATE REQUEST-----\n" |
| 188 | + + wrapped |
| 189 | + + "\n-----END CERTIFICATE REQUEST-----\n" |
| 190 | + ).encode("ascii") |
| 191 | + |
77 | 192 | def format_uuid(hex_str): |
78 | 193 | return '{0}-{1}-{2}-{3}-{4}'.format(hex_str[:8], hex_str[8:12], |
79 | 194 | hex_str[12:16], hex_str[16:20], |
@@ -179,30 +294,26 @@ def parse_keygen_output(keygen_str): |
179 | 294 | body = body_cose[0] |
180 | 295 |
|
181 | 296 | # Decode base64url to binary |
182 | | - body_bytes = base64_decode(body) |
| 297 | + payload_body_bytes = base64_decode(body) |
| 298 | + body_bytes = payload_body_bytes |
183 | 299 |
|
184 | | - # This can be either a CSR or device public key |
185 | 300 | try: |
186 | | - # Try to load CSR, if it fails, assume public key |
187 | 301 | csr = x509.load_der_x509_csr(body_bytes) |
| 302 | + except: |
| 303 | + logger.warning("normalizing CSR") |
| 304 | + normalized = _normalize_ecdsa_csr_der(body_bytes) |
| 305 | + logger.warning(f"keygen_str: {keygen_str}, normalized: {base64.urlsafe_b64encode(normalized)}") |
| 306 | + csr = x509.load_der_x509_csr(normalized) |
188 | 307 |
|
189 | | - except ValueError: |
190 | | - # Handle public key only |
191 | | - pub_key = serialization.load_der_public_key(body_bytes) |
192 | | - pub_key_bytes = pub_key.public_bytes(serialization.Encoding.PEM, serialization.PublicFormat.SubjectPublicKeyInfo) |
193 | | - |
194 | | - else: |
195 | | - # CSR loaded, logger.info it |
196 | | - csr_pem_bytes = csr.public_bytes(serialization.Encoding.PEM) |
197 | | - csr_pem_list = str(csr_pem_bytes.decode()).split('\n') |
198 | | - logger.info(csr_pem_bytes.decode().replace('\n', '\\n')) |
| 308 | + csr_pem_bytes = csr.public_bytes(serialization.Encoding.PEM) |
| 309 | + logger.info(csr_pem_bytes.decode().replace('\n', '\\n')) |
199 | 310 |
|
200 | | - # Extract public key |
201 | | - pub_key_bytes = csr.public_key().public_bytes(serialization.Encoding.PEM, serialization.PublicFormat.SubjectPublicKeyInfo) |
| 311 | + # Extract public key |
| 312 | + pub_key_bytes = csr.public_key().public_bytes(serialization.Encoding.PEM, serialization.PublicFormat.SubjectPublicKeyInfo) |
202 | 313 |
|
203 | 314 | logger.info("Device public key: {}".format(pub_key_bytes.decode().replace('\n', '\\n'))) |
204 | 315 |
|
205 | | - payload_digest = hashlib.sha256(body_bytes).hexdigest() |
| 316 | + payload_digest = hashlib.sha256(payload_body_bytes).hexdigest() |
206 | 317 | logger.info(f"SHA256 Digest: {payload_digest}") |
207 | 318 |
|
208 | 319 | # Get optional cose |
|
0 commit comments