diff --git a/README.md b/README.md index b0c2d08e..a49e7564 100644 --- a/README.md +++ b/README.md @@ -54,6 +54,15 @@ assert b"Hello world!" == recipient.decode(encoded, mac_key) ## You can get decoded protected/unprotected headers with the payload as follows: # protected, unprotected, payload = recipient.decode_with_headers(encoded, mac_key) # assert b"Hello world!" == payload + +## Note that to pass header parameters with tstr labels, or tstr values, and avoid +# clashes with short-string names such as "alg" or value encoding to bstr, you can +# resolve the headers yourself, and pass a cwt.utils.ResolvedHeader({...}). +# +# For example: +# protected=cwt.utils.ResolvedHeader({ +# "string label": "value" +# }) ``` **CWT API** diff --git a/cwt/cbor_processor.py b/cwt/cbor_processor.py index 6de4c723..67118a46 100644 --- a/cwt/cbor_processor.py +++ b/cwt/cbor_processor.py @@ -1,4 +1,4 @@ -from typing import Any, Dict +from typing import Any, Dict, Union from cbor2 import dumps, loads @@ -12,7 +12,7 @@ def _dumps(self, obj: Any) -> bytes: except Exception as err: raise EncodeError("Failed to encode.") from err - def _loads(self, s: bytes) -> Dict[int, Any]: + def _loads(self, s: bytes) -> Dict[Union[str, int], Any]: try: return loads(s) except Exception as err: diff --git a/cwt/cose.py b/cwt/cose.py index 2890b3ca..704102df 100644 --- a/cwt/cose.py +++ b/cwt/cose.py @@ -21,7 +21,7 @@ from .recipient_interface import RecipientInterface from .recipients import Recipients from .signer import Signer -from .utils import sort_keys_for_deterministic_encoding, to_cose_header +from .utils import ResolvedHeader, sort_keys_for_deterministic_encoding, to_cose_header class COSE(CBORProcessor): @@ -132,8 +132,8 @@ def encode( self, payload: bytes, key: Optional[COSEKeyInterface] = None, - protected: Optional[dict] = None, - unprotected: Optional[dict] = None, + protected: Optional[Union[dict, ResolvedHeader]] = None, + unprotected: Optional[Union[dict, ResolvedHeader]] = None, recipients: List[RecipientInterface] = [], signers: List[Signer] = [], external_aad: bytes = b"", @@ -146,8 +146,8 @@ def encode( Args: payload (bytes): A content to be MACed, signed or encrypted. key (Optional[COSEKeyInterface]): A content encryption key as COSEKey. - protected (Optional[dict]): Parameters that are to be cryptographically protected. - unprotected (Optional[dict]): Parameters that are not cryptographically protected. + protected (Optional[Union[dict, ResolvedHeader]]): Parameters that are to be cryptographically protected. + unprotected (Optional[Union[dict, ResolvedHeader]]): Parameters that are not cryptographically protected. recipients (List[RecipientInterface]): A list of recipient information structures. signers (List[Signer]): A list of signer information objects for multiple signer cases. @@ -351,7 +351,7 @@ def decode_with_headers( external_aad: bytes = b"", detached_payload: Optional[bytes] = None, enable_non_aead: bool = False, - ) -> Tuple[Dict[int, Any], Dict[int, Any], bytes]: + ) -> Tuple[Dict[Union[str, int], Any], Dict[Union[str, int], Any], bytes]: """ Verifies and decodes COSE data, and returns protected headers, unprotected headers and payload. @@ -371,7 +371,7 @@ def decode_with_headers( Since non-AEAD ciphers DO NOT provide neither authentication nor integrity of decrypted message, make sure to validate them outside of this library. Returns: - Tuple[Dict[int, Any], Dict[int, Any], bytes]: A dictionary data of decoded protected headers, and a dictionary data of unprotected headers, and a byte string of decoded payload. + Tuple[Dict[Union[str, int], Any], Dict[Union[str, int], Any], bytes]: A dictionary data of decoded protected headers, and a dictionary data of unprotected headers, and a byte string of decoded payload. Raises: ValueError: Invalid arguments. DecodeError: Failed to decode data. @@ -582,10 +582,10 @@ def decode_with_headers( def _encode_headers( self, key: Optional[COSEKeyInterface], - protected: Optional[dict], - unprotected: Optional[dict], + protected: Optional[Union[dict, ResolvedHeader]], + unprotected: Optional[Union[dict, ResolvedHeader]], enable_non_aead: bool, - ) -> Tuple[Dict[int, Any], Dict[int, Any]]: + ) -> Tuple[Dict[Union[str, int], Any], Dict[Union[str, int], Any]]: p = to_cose_header(protected) u = to_cose_header(unprotected) if key is not None: @@ -612,14 +612,16 @@ def _encode_headers( raise ValueError("protected header MUST be zero-length") return p, u - def _decode_headers(self, protected: Any, unprotected: Any) -> Tuple[Dict[int, Any], Dict[int, Any]]: - p: Union[Dict[int, Any], bytes] + def _decode_headers( + self, protected: Any, unprotected: Any + ) -> Tuple[Dict[Union[str, int], Any], Dict[Union[str, int], Any]]: + p: Union[Dict[Union[str, int], Any], bytes] p = self._loads(protected) if protected else {} if isinstance(p, bytes): if len(p) > 0: raise ValueError("Invalid protected header.") p = {} - u: Dict[int, Any] = unprotected + u: Dict[Union[str, int], Any] = unprotected if not isinstance(u, dict): raise ValueError("unprotected header should be dict.") return p, u @@ -627,15 +629,15 @@ def _decode_headers(self, protected: Any, unprotected: Any) -> Tuple[Dict[int, A def _validate_cose_message( self, key: Optional[COSEKeyInterface], - p: Dict[int, Any], - u: Dict[int, Any], + p: Dict[Union[str, int], Any], + u: Dict[Union[str, int], Any], recipients: List[RecipientInterface], signers: List[Signer], ) -> int: if len(recipients) > 0 and len(signers) > 0: raise ValueError("Both recipients and signers are specified.") - h: Dict[int, Any] = {} + h: Dict[Union[str, int], Any] = {} iv_count: int = 0 for k, v in p.items(): if k == 2: # crit @@ -745,8 +747,8 @@ def _encode_and_encrypt( self, payload: bytes, key: Optional[COSEKeyInterface], - p: Dict[int, Any], - u: Dict[int, Any], + p: Dict[Union[str, int], Any], + u: Dict[Union[str, int], Any], recipients: List[RecipientInterface], external_aad: bytes, out: str, @@ -806,8 +808,8 @@ def _encode_and_mac( self, payload: bytes, key: Optional[COSEKeyInterface], - p: Dict[int, Any], - u: Dict[int, Any], + p: Dict[Union[str, int], Any], + u: Dict[Union[str, int], Any], recipients: List[RecipientInterface], external_aad: bytes, out: str, @@ -849,8 +851,8 @@ def _encode_and_sign( self, payload: bytes, key: Optional[COSEKeyInterface], - p: Dict[int, Any], - u: Dict[int, Any], + p: Dict[Union[str, int], Any], + u: Dict[Union[str, int], Any], signers: List[Signer], external_aad: bytes, out: str, diff --git a/cwt/cose_message.py b/cwt/cose_message.py index 79c759f7..bffa4ac1 100644 --- a/cwt/cose_message.py +++ b/cwt/cose_message.py @@ -1,6 +1,6 @@ from __future__ import annotations -from typing import Any, Dict, List, Optional, Tuple +from typing import Any, Dict, List, Optional, Tuple, Union from cbor2 import CBORTag, loads @@ -132,14 +132,14 @@ def type(self) -> COSETypes: return self._type @property - def protected(self) -> Dict[int, Any]: + def protected(self) -> Dict[Union[str, int], Any]: """ The protected headers as a CBOR object. """ return self._loads(self._protected) @property - def unprotected(self) -> Dict[int, Any]: + def unprotected(self) -> Dict[Union[str, int], Any]: """ The unprotected headers as a CBOR object. """ diff --git a/cwt/cwt.py b/cwt/cwt.py index 5d7d755d..cf5d7e0b 100644 --- a/cwt/cwt.py +++ b/cwt/cwt.py @@ -316,7 +316,7 @@ def decode( data: bytes, keys: Union[COSEKeyInterface, List[COSEKeyInterface]], no_verify: bool = False, - ) -> Union[Dict[int, Any], bytes]: + ) -> Union[Dict[Union[str, int], Any], bytes]: """ Verifies and decodes CWT. @@ -333,11 +333,11 @@ def decode( DecodeError: Failed to decode the CWT. VerifyError: Failed to verify the CWT. """ - cwt: Union[bytes, CBORTag, Dict[int, Any]] = self._loads(data) + cwt: Union[bytes, CBORTag, Dict[Union[str, int], Any]] = self._loads(data) if isinstance(cwt, CBORTag) and cwt.tag == CWT.CBOR_TAG: cwt = cwt.value keys = [keys] if isinstance(keys, COSEKeyInterface) else keys - p: Dict[int, Any] = {} + p: Dict[Union[str, int], Any] = {} while isinstance(cwt, CBORTag): p, u, cwt = self._cose.decode_with_headers(cwt, keys) cwt = self._loads(cwt) @@ -399,7 +399,7 @@ def _validate(self, claims: Union[Dict[int, Any], bytes]): Claims.validate(claims) return - def _verify(self, claims: Union[Dict[int, Any], bytes], protected: Dict[int, Any] = {}): + def _verify(self, claims: Union[Dict[Union[str, int], Any], bytes], protected: Dict[Union[str, int], Any] = {}): if not isinstance(claims, dict): raise DecodeError("Failed to decode.") @@ -484,7 +484,7 @@ def decode( data: bytes, keys: Union[COSEKeyInterface, List[COSEKeyInterface]], no_verify: bool = False, -) -> Union[Dict[int, Any], bytes]: +) -> Union[Dict[Union[str, int], Any], bytes]: return _cwt.decode(data, keys, no_verify) diff --git a/cwt/recipient.py b/cwt/recipient.py index 0694677a..2d4bcb0d 100644 --- a/cwt/recipient.py +++ b/cwt/recipient.py @@ -19,7 +19,7 @@ from .recipient_algs.ecdh_direct_hkdf import ECDH_DirectHKDF from .recipient_algs.hpke import HPKE from .recipient_interface import RecipientInterface -from .utils import to_cose_header, to_recipient_context +from .utils import ResolvedHeader, to_cose_header, to_recipient_context class Recipient: @@ -30,8 +30,8 @@ class Recipient: @classmethod def new( cls, - protected: dict = {}, - unprotected: dict = {}, + protected: Union[dict, ResolvedHeader] = {}, + unprotected: Union[dict, ResolvedHeader] = {}, ciphertext: bytes = b"", recipients: List[Any] = [], sender_key: Optional[COSEKeyInterface] = None, @@ -42,8 +42,8 @@ def new( Creates a recipient from a CBOR-like dictionary with numeric keys. Args: - protected (dict): Parameters that are to be cryptographically protected. - unprotected (dict): Parameters that are not cryptographically protected. + protected (Union[dict, ResolvedHeader]): Parameters that are to be cryptographically protected. + unprotected (Union[dict, ResolvedHeader]): Parameters that are not cryptographically protected. ciphertext (List[Any]): A cipher text. sender_key (Optional[COSEKeyInterface]): A sender private key as COSEKey. recipient_key (Optional[COSEKeyInterface]): A recipient public key as COSEKey. @@ -74,7 +74,7 @@ def new( if alg == -6: return DirectKey(p, u) if alg in COSE_ALGORITHMS_KEY_WRAP.values(): - if len(protected) > 0: + if len(p) > 0: raise ValueError("The protected header must be a zero-length string in key wrap mode with an AE algorithm.") if not sender_key: sender_key = COSEKey.from_symmetric_key(alg=alg) diff --git a/cwt/recipient_algs/aes_key_wrap.py b/cwt/recipient_algs/aes_key_wrap.py index b5e9a223..3a6733a8 100644 --- a/cwt/recipient_algs/aes_key_wrap.py +++ b/cwt/recipient_algs/aes_key_wrap.py @@ -15,7 +15,7 @@ class AESKeyWrap(RecipientInterface): def __init__( self, - unprotected: Dict[int, Any], + unprotected: Dict[Union[str, int], Any], ciphertext: bytes = b"", recipients: List[Any] = [], sender_key: Optional[COSEKeyInterface] = None, diff --git a/cwt/recipient_algs/direct.py b/cwt/recipient_algs/direct.py index 286cc3dd..65df3e94 100644 --- a/cwt/recipient_algs/direct.py +++ b/cwt/recipient_algs/direct.py @@ -1,4 +1,4 @@ -from typing import Any, Dict, List +from typing import Any, Dict, List, Union from ..recipient_interface import RecipientInterface @@ -6,8 +6,8 @@ class Direct(RecipientInterface): def __init__( self, - protected: Dict[int, Any], - unprotected: Dict[int, Any], + protected: Dict[Union[str, int], Any], + unprotected: Dict[Union[str, int], Any], ciphertext: bytes = b"", recipients: List[Any] = [], ): diff --git a/cwt/recipient_algs/direct_hkdf.py b/cwt/recipient_algs/direct_hkdf.py index 0ed01a90..2759285e 100644 --- a/cwt/recipient_algs/direct_hkdf.py +++ b/cwt/recipient_algs/direct_hkdf.py @@ -19,8 +19,8 @@ class DirectHKDF(Direct): def __init__( self, - protected: Dict[int, Any] = {}, - unprotected: Dict[int, Any] = {}, + protected: Dict[Union[str, int], Any] = {}, + unprotected: Dict[Union[str, int], Any] = {}, context: List[Any] = [], ): super().__init__(protected, unprotected, b"", []) diff --git a/cwt/recipient_algs/direct_key.py b/cwt/recipient_algs/direct_key.py index 8aafd2cb..4afd1018 100644 --- a/cwt/recipient_algs/direct_key.py +++ b/cwt/recipient_algs/direct_key.py @@ -5,7 +5,7 @@ class DirectKey(Direct): - def __init__(self, protected: Dict[int, Any] = {}, unprotected: Dict[int, Any] = {}): + def __init__(self, protected: Dict[Union[str, int], Any] = {}, unprotected: Dict[Union[str, int], Any] = {}): super().__init__(protected, unprotected, b"", []) if self._alg != -6: diff --git a/cwt/recipient_algs/ecdh_aes_key_wrap.py b/cwt/recipient_algs/ecdh_aes_key_wrap.py index 9abda07c..2849be65 100644 --- a/cwt/recipient_algs/ecdh_aes_key_wrap.py +++ b/cwt/recipient_algs/ecdh_aes_key_wrap.py @@ -18,8 +18,8 @@ class ECDH_AESKeyWrap(RecipientInterface): def __init__( self, - protected: Dict[int, Any], - unprotected: Dict[int, Any], + protected: Dict[Union[str, int], Any], + unprotected: Dict[Union[str, int], Any], ciphertext: bytes = b"", recipients: List[Any] = [], sender_key: Optional[COSEKeyInterface] = None, diff --git a/cwt/recipient_algs/ecdh_direct_hkdf.py b/cwt/recipient_algs/ecdh_direct_hkdf.py index 9332d1bb..cad79290 100644 --- a/cwt/recipient_algs/ecdh_direct_hkdf.py +++ b/cwt/recipient_algs/ecdh_direct_hkdf.py @@ -18,8 +18,8 @@ class ECDH_DirectHKDF(Direct): def __init__( self, - protected: Dict[int, Any], - unprotected: Dict[int, Any], + protected: Dict[Union[str, int], Any], + unprotected: Dict[Union[str, int], Any], ciphertext: bytes = b"", recipients: List[Any] = [], sender_key: Optional[COSEKeyInterface] = None, diff --git a/cwt/recipient_algs/hpke.py b/cwt/recipient_algs/hpke.py index 41f3e42a..977eb1aa 100644 --- a/cwt/recipient_algs/hpke.py +++ b/cwt/recipient_algs/hpke.py @@ -36,8 +36,8 @@ def to_hpke_ciphersuites(alg: int) -> Tuple[int, int, int]: class HPKE(RecipientInterface): def __init__( self, - protected: Dict[int, Any], - unprotected: Dict[int, Any], + protected: Dict[Union[str, int], Any], + unprotected: Dict[Union[str, int], Any], ciphertext: bytes = b"", recipients: List[Any] = [], recipient_key: Optional[COSEKeyInterface] = None, diff --git a/cwt/recipient_interface.py b/cwt/recipient_interface.py index db8d2f22..ae58f2ff 100644 --- a/cwt/recipient_interface.py +++ b/cwt/recipient_interface.py @@ -17,8 +17,8 @@ class RecipientInterface(CBORProcessor): def __init__( self, - protected: Optional[Dict[int, Any]] = None, - unprotected: Optional[Dict[int, Any]] = None, + protected: Optional[Dict[Union[str, int], Any]] = None, + unprotected: Optional[Dict[Union[str, int], Any]] = None, ciphertext: bytes = b"", recipients: List[Any] = [], key_ops: List[int] = [], @@ -106,7 +106,7 @@ def alg(self) -> int: return self._alg @property - def protected(self) -> Dict[int, Any]: + def protected(self) -> Dict[Union[str, int], Any]: """ The parameters that are to be cryptographically protected. """ @@ -122,7 +122,7 @@ def b_protected(self) -> bytes: return self._b_protected @property - def unprotected(self) -> Dict[int, Any]: + def unprotected(self) -> Dict[Union[str, int], Any]: """ The parameters that are not cryptographically protected. """ diff --git a/cwt/signer.py b/cwt/signer.py index 1965ddaa..dc95519e 100644 --- a/cwt/signer.py +++ b/cwt/signer.py @@ -4,7 +4,7 @@ from .const import COSE_ALGORITHMS_SIGNATURE from .cose_key import COSEKey from .cose_key_interface import COSEKeyInterface -from .utils import to_cose_header +from .utils import ResolvedHeader, to_cose_header class Signer(CBORProcessor): @@ -15,8 +15,8 @@ class Signer(CBORProcessor): def __init__( self, cose_key: COSEKeyInterface, - protected: Union[Dict[int, Any], bytes], - unprotected: Dict[int, Any], + protected: Union[Dict[Union[str, int], Any], bytes], + unprotected: Dict[Union[str, int], Any], signature: bytes = b"", ): self._cose_key = cose_key @@ -43,7 +43,7 @@ def protected(self) -> bytes: return self._protected @property - def unprotected(self) -> Dict[int, Any]: + def unprotected(self) -> Dict[Union[str, int], Any]: """ The parameters that are not cryptographically protected. """ @@ -60,8 +60,8 @@ def signature(self) -> bytes: def new( cls, cose_key: COSEKeyInterface, - protected: Union[dict, bytes] = {}, - unprotected: dict = {}, + protected: Union[dict, bytes, ResolvedHeader] = {}, + unprotected: Union[dict, ResolvedHeader] = {}, signature: bytes = b"", ): """ @@ -69,19 +69,19 @@ def new( Args: cose_key (COSEKey): A signature key for the signer. - protected (Union[dict, bytes]): Parameters that are to be cryptographically + protected (Union[dict, bytes, ResolvedHeader]): Parameters that are to be cryptographically protected. - unprotected (dict): Parameters that are not cryptographically protected. + unprotected (Union[dict, bytes, ResolvedHeader]): Parameters that are not cryptographically protected. signature (bytes): A signature as bytes. Returns: Signer: A signer information object. Raises: ValueError: Invalid arguments. """ - p: Union[Dict[int, Any], bytes] = ( - to_cose_header(protected, algs=COSE_ALGORITHMS_SIGNATURE) if isinstance(protected, dict) else protected + p: Union[Dict[Union[str, int], Any], bytes] = ( + protected if isinstance(protected, bytes) else to_cose_header(protected, algs=COSE_ALGORITHMS_SIGNATURE) ) - u = to_cose_header(unprotected, algs=COSE_ALGORITHMS_SIGNATURE) + u: Dict[Union[str, int], Any] = to_cose_header(unprotected, algs=COSE_ALGORITHMS_SIGNATURE) return cls(cose_key, p, u, signature) @classmethod @@ -101,8 +101,8 @@ def from_jwk(cls, data: Union[str, bytes, Dict[str, Any]]): ValueError: Invalid arguments. DecodeError: Failed to decode the key data. """ - protected: Dict[int, Any] = {} - unprotected: Dict[int, Any] = {} + protected: Dict[Union[str, int], Any] = {} + unprotected: Dict[Union[str, int], Any] = {} cose_key = COSEKey.from_jwk(data) @@ -142,8 +142,8 @@ def from_pem( ValueError: Invalid arguments. DecodeError: Failed to decode the key data. """ - protected: Dict[int, Any] = {} - unprotected: Dict[int, Any] = {} + protected: Dict[Union[str, int], Any] = {} + unprotected: Dict[Union[str, int], Any] = {} cose_key = COSEKey.from_pem(data, alg=alg, kid=kid) diff --git a/cwt/utils.py b/cwt/utils.py index 45ab4079..02e113f1 100644 --- a/cwt/utils.py +++ b/cwt/utils.py @@ -167,10 +167,23 @@ def to_cis(context: Dict[str, Any], recipient_alg: Optional[int] = None) -> List return res -def to_cose_header(data: Optional[dict] = None, algs: Dict[str, int] = {}) -> Dict[int, Any]: +class ResolvedHeader: + """ + Wrapping COSE header parameters in a ResolvedHeader is a way to signal + to library calls that they do not require resolution against const.COSE_HEADER_PARAMETERS, + nor value encoding to bstr, and can be passed directly to the CBOR encoding logic. + """ + + def __init__(self, params: Dict[Union[str, int], Any]): + self.params = params + + +def to_cose_header(data: Optional[Union[dict, ResolvedHeader]] = None, algs: Dict[str, int] = {}) -> Dict[Union[str, int], Any]: if data is None: return {} - res: Dict[int, Any] = {} + res: Dict[Union[str, int], Any] = {} + if isinstance(data, ResolvedHeader): + return data.params if len(data) == 0 or not isinstance(list(data.keys())[0], str): return data if not algs: @@ -317,7 +330,7 @@ def _validate_context(context: List[Any]) -> List[Any]: return context -def to_recipient_context(alg: int, u: Dict[int, Any], context: Union[List[Any], Dict[str, Any]]) -> List[Any]: +def to_recipient_context(alg: int, u: Dict[Union[str, int], Any], context: Union[List[Any], Dict[str, Any]]) -> List[Any]: ctx: List[Any] = [ None, [ @@ -347,5 +360,5 @@ def to_recipient_context(alg: int, u: Dict[int, Any], context: Union[List[Any], return ctx -def sort_keys_for_deterministic_encoding(d: Dict[int, Any]) -> Dict[int, Any]: +def sort_keys_for_deterministic_encoding(d: Dict[Union[str, int], Any]) -> Dict[Union[str, int], Any]: return {k: v for k, v in sorted(d.items(), key=lambda kv: cbor2.dumps(kv[0]))} diff --git a/docs/conf.py b/docs/conf.py index 5ce66463..5a413209 100644 --- a/docs/conf.py +++ b/docs/conf.py @@ -42,6 +42,8 @@ ("py:class", "T"), ("py:class", "cwt.COSEKeyInterface"), ("py:class", "cwt.RecipientInterface"), + ("py:class", "cwt.utils.ResolvedHeader"), + ("py:class", "ResolvedHeader"), ("py:class", "cwt.claims.T"), ("py:class", "cwt.cbor_processor.CBORProcessor"), ("py:class", "_cbor2.CBORTag"), diff --git a/samples/eudcc/swedish_verifier.py b/samples/eudcc/swedish_verifier.py index 9e9b2a73..9627541c 100644 --- a/samples/eudcc/swedish_verifier.py +++ b/samples/eudcc/swedish_verifier.py @@ -59,7 +59,7 @@ def refresh_trustlist(self): json.dump(self._trustlist, f, indent=4) return - def verify_and_decode(self, eudcc: bytes) -> Union[Dict[int, Any], bytes]: + def verify_and_decode(self, eudcc: bytes) -> Union[Dict[Union[str, int], Any], bytes]: if eudcc.startswith(b"HC1:"): # Decode Base45 data. eudcc = b45decode(eudcc[4:]) diff --git a/samples/eudcc/verifier.py b/samples/eudcc/verifier.py index 1681528a..71e87241 100644 --- a/samples/eudcc/verifier.py +++ b/samples/eudcc/verifier.py @@ -64,7 +64,7 @@ def refresh_trustlist(self): json.dump([v for v in self._trustlist if v["x_kid"] in active_kids], f, indent=4) return - def verify_and_decode(self, eudcc: bytes) -> Union[Dict[int, Any], bytes]: + def verify_and_decode(self, eudcc: bytes) -> Union[Dict[Union[str, int], Any], bytes]: if eudcc.startswith(b"HC1:"): # Decode Base45 data. eudcc = b45decode(eudcc[4:]) diff --git a/tests/test_cose_sample.py b/tests/test_cose_sample.py index 991d6514..ab0e8276 100644 --- a/tests/test_cose_sample.py +++ b/tests/test_cose_sample.py @@ -2,7 +2,7 @@ import pytest -from cwt import COSE, COSEAlgs, COSEHeaders, COSEKey, Recipient, Signer +from cwt import COSE, COSEAlgs, COSEHeaders, COSEKey, Recipient, Signer, utils class TestCOSESample: @@ -768,3 +768,59 @@ def test_cose_usage_examples_cose_signature(self): ) encoded3 = sender.encode_and_sign(b"Hello world!", signers=[signer]) assert b"Hello world!" == recipient.decode(encoded3, pub_key) + + def test_cose_usage_with_resolved_header(self): + cose_key = COSEKey.from_jwk( + { + "kty": "EC", + "kid": "01", + "crv": "P-256", + "x": "usWxHK2PmfnHKwXPS54m0kTcGJ90UiglWiGahtagnv8", + "y": "IBOL-C3BttVivg-lSreASjpkttcsz-1rb7btKLv8EX4", + "d": "V8kgd2ZBRuh2dgyVINBUqpPDr7BOMGcF22CQMIUHtNM", + } + ) + + recipient = COSE.new() + pub_key = COSEKey.from_jwk( + { + "kty": "EC", + "kid": "01", + "crv": "P-256", + "x": "usWxHK2PmfnHKwXPS54m0kTcGJ90UiglWiGahtagnv8", + "y": "IBOL-C3BttVivg-lSreASjpkttcsz-1rb7btKLv8EX4", + } + ) + + ctx = COSE.new(alg_auto_inclusion=True) + + payload = b"Hello world!" + protected_first_not_str = {COSEHeaders.KID: b"01", COSEHeaders.ALG: COSEAlgs.ES256, "string key": "value"} + + # This works because the first key is not a string, and so the early exit in to_cose_header() happens + encoded = ctx.encode_and_sign(payload, cose_key, protected=protected_first_not_str) + phdr, _, payload_ = recipient.decode_with_headers(encoded, pub_key) + assert payload_ == payload + assert phdr[COSEHeaders.ALG] == COSEAlgs.ES256 + assert phdr["string key"] == "value" + + # This works because ResolvedHeader always skips string key encoding + encoded = ctx.encode_and_sign(payload, cose_key, protected=utils.ResolvedHeader(protected_first_not_str)) + phdr, _, payload_ = recipient.decode_with_headers(encoded, pub_key) + assert payload_ == payload + assert phdr[COSEHeaders.ALG] == COSEAlgs.ES256 + assert phdr["string key"] == "value" + + protected_first_str = {"string key": "value", COSEHeaders.KID: b"01", COSEHeaders.ALG: COSEAlgs.ES256} + + with pytest.raises(ValueError): + # Raises ValueError: Unsupported or unknown COSE header parameter: 4. + # This fails because the first key is a string, and to_cose_header attempts to resolve the params + ctx.encode_and_sign(b"Hello world!", cose_key, protected=protected_first_str) + + # This works because ResolvedHeader always skips string key encoding + encoded = ctx.encode_and_sign(payload, cose_key, protected=utils.ResolvedHeader(protected_first_str)) + phdr, _, payload_ = recipient.decode_with_headers(encoded, pub_key) + assert payload_ == payload + assert phdr[COSEHeaders.ALG] == COSEAlgs.ES256 + assert phdr["string key"] == "value"