pax_global_header00006660000000000000000000000064152441356000014512gustar00rootroot0000000000000052 comment=e68cede0561acb18bbc8d10453335e03b3d63224 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/000077500000000000000000000000001524413560000220525ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/.coveragerc000066400000000000000000000003101524413560000241650ustar00rootroot00000000000000[run] source = cryptoparser omit = cryptoparser/__setup__.py [report] exclude_lines = pragma: no cover raise NotImplementedError fail_under = 100 include = cryptoparser/* show_missing = True cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/.gitignore000066400000000000000000000001501524413560000240360ustar00rootroot00000000000000*.orig *.rej *.swp *.pyc *.egg-info /.coverage /.eggs /Pipfile /Pipfile.lock /build /dist /uv.lock cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/.gitlab-ci.yml000066400000000000000000000035421524413560000245120ustar00rootroot00000000000000image: python stages: - earlytest - fulltest - deploy before_script: - pip install uv - uv sync --extra tests variables: GIT_SUBMODULE_DEPTH: 1 GIT_SUBMODULE_STRATEGY: recursive PYTHONPATH: "submodules/cryptodatahub" UV_PYTHON_DOWNLOADS: "never" UV_PYTHON_PREFERENCE: "only-system" .test: script: - uv run coverage run -m unittest discover -v - uv run coverage report pylint: image: python:3.14-slim stage: earlytest script: - uv run --with pylint pylint --rcfile .pylintrc cryptoparser docs test ruff: image: python:3.14-slim stage: earlytest script: uvx ruff check cryptoparser docs test python314: extends: .test image: python:3.14-slim stage: earlytest python39: extends: .test image: python:3.9-slim stage: fulltest python310: extends: .test image: python:3.10-slim stage: fulltest python311: extends: .test image: python:3.11-slim stage: fulltest python312: extends: .test image: python:3.12-slim stage: fulltest python313: extends: .test image: python:3.13-slim stage: fulltest pythonrc: extends: .test image: python:3.15-rc-slim stage: fulltest pypy3: extends: .test image: pypy:3-slim stage: fulltest coveralls: image: python:3.12-slim variables: CI_NAME: gitlab CI_BUILD_NUMBER: "${CI_JOB_ID}" CI_BUILD_URL: "${CI_JOB_URL}" CI_BRANCH: "${CI_COMMIT_REF_NAME}" GIT_SUBMODULE_DEPTH: 1 GIT_SUBMODULE_STRATEGY: recursive PYTHONPATH: "submodules/cryptodatahub" stage: deploy script: - uv run coverage run -m unittest -v -f - uv run --with coveralls coveralls only: refs: - master obs: image: name: coroner/python_obs:1.2.5 pull_policy: always stage: deploy variables: GIT_SUBMODULE_DEPTH: 1 GIT_SUBMODULE_STRATEGY: recursive before_script: [] script: - obs.sh only: refs: - master - tags cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/.gitmodules000066400000000000000000000001741524413560000242310ustar00rootroot00000000000000[submodule "submodules/cryptodatahub"] path = submodules/cryptodatahub url = https://gitlab.com/coroner/cryptodatahub.git cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/.pylintrc000066400000000000000000000005131524413560000237160ustar00rootroot00000000000000[MAIN] jobs=0 [MASTER] load-plugins = pylint.extensions.no_self_use [BASIC] good-names=setUp,tearDown,setUpClass,tearDownClass [FORMAT] max-line-length=120 [MESSAGES CONTROL] disable= missing-docstring, too-few-public-methods, too-many-function-args, too-many-ancestors, duplicate-code, no-member, cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/.readthedocs.yaml000066400000000000000000000003331524413560000253000ustar00rootroot00000000000000version: 2 sphinx: configuration: docs/conf.py builder: dirhtml build: os: "ubuntu-22.04" tools: python: "3.10" python: install: - method: pip path: . extra_requirements: - docs cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/CHANGELOG.rst000066400000000000000000000431351524413560000241010ustar00rootroot00000000000000========= Changelog ========= ------------------ 1.6.0 - 2026-08-25 ------------------ Features ======== - TLS (``tls``) - Extensions (``extensions``) - add `trust anchors `__ extension related messages (#104) - add server padding extension related messages (#104) - add `application-layer protocol settings `__ extension related messages with the old code point (#104) Notable fixes ============= - TLS (``tls``) - Extensions (``extensions``) - bound the outer extension payload of the encrypted client hello by the extension length (#104) - IKE (``ike``) - preserve the transform number when parsing the IKEv1 transform payload (#105) Refactor ======== - Generic - follow the replacement of the named group with the key parameter (#106) - TLS (``tls``) - make the parameter classes of the invalid extension types immutable (#104) Other ===== - switch to Markdown format in readme - add llms.txt to describe the library for large language models - use uv instead of pipenv and pip in the documentation ------------------ 1.5.0 - 2026-07-31 ------------------ Features ======== - IKE (``ike``) - add IKEv1 and IKEv2 certificate payload parsing (#103) - add IKEv1 identification payload parsing (#103) - add IKEv1 signature payload parsing (#103) - add IKEv2 identification payload parsing (#103) - add IKEv2 authentication payload parsing (#103) - add IKEv2 extensible authentication protocol payload parsing (#103) - add IKEv2 encrypted and authenticated payload parsing (#103) - keep IKEv1 and IKEv2 payloads of unknown type unparsed instead of rejecting the message (#103) - add getter for the distinguished name of the IKEv1 certificate request payload (#103) - add getter for the certification authority hashes of the IKEv2 certificate request payload (#103) Refactor ======== - IKE (``ike``) - parse the IKEv1 and IKEv2 identification payload data according to the identification type (#103) - parse the digital signature envelope of the IKEv2 authentication payload (#103) - split the ISAKMP message parsing and composing into header and payload chain steps (#103) - use typed values in the IKEv2 signature hash algorithms notify payload (#103) ------------------ 1.4.0 - 2026-07-17 ------------------ Features ======== - IKE (``ike``) - add getter for multiple payloads with the same type (#93) - add IKEv2 notify payload parsing for protocol extensions (#93) - NAT detection source IP and destination IP - set window size - use transport mode - HTTP certificate lookup supported - signature hash algorithms - intermediate exchange supported - use PPK - redirect supported - IKEv2 fragmentation supported - childless IKEv2 supported ------------------ 1.3.0 - 2026-06-15 ------------------ Features ======== - Generic - add Debian and RPM packaging (#102) - TLS (``tls``) - add JA4 tag generation for the client hello message (#100) - add JA4X tag generation for X.509 certificates (#101) - add certificate-related messages for protocol version 1.3 (#94) Notable fixes ============= - IKE (``ike``) - make IKEv2 transform key length optional for fixed-key ciphers (#99) ------------------ 1.2.1 - 2026-06-02 ------------------ Features ======== - DNS (``dnsrec``) - add SSHFP DNS resource record parsing (#98) ------------------ 1.2.0 - 2026-05-05 ------------------ Features ======== - IKE (``ike``) - add IKEv1 delete payload parsing (#92) - add IKEv1 certificate request payload parsing (#97) ------------------ 1.1.1 - 2026-05-03 ------------------ Notable fixes ============= - SSH - add missing key size property (#96) ------------------- 1.1.0 - 2026-02-13 ------------------- Features ======== - IKE (``ike``) - add ISAKMP header parsing (#91) - add parser/composer for mandatory IKEv2 protocol elements (#91) - add parser/composer for mandatory IKEv1 protocol elements (#91) ------------------- 1.0.2 - 2025-12-30 ------------------- Notable fixes ============= - Generic - Purge submodules directory from distribution (#90) ------------------- 1.0.1 - 2025-12-07 ------------------- Refactor ======== - Generic - Remove unnecessary dateutil dependency (#89) ------------------- 1.0.0 - 2025-01-05 ------------------- Refactor ======== - Generic - Support only Python version greater than or equal to 3.9 - Use pyproject.toml instead of setup.py ------------------- 0.12.6 - 2024-12-08 ------------------- Improvements ============ - TLS (``tls``) - Extensions (``extensions``) - add `TLS encrypted client hello `__ extension related messages (#86) - add `post-handshake authentication `__ messages (#86) ------------------- 0.12.5 - 2024-05-25 ------------------- Refactor ======== - Generic - Unify pyfakefs related test classes ------------------- 0.12.4 - 2024-04-28 ------------------- Notable fixes ============= - DNS - handle TXT records that contain multiple string (#84) ------------------- 0.12.3 - 2024-03-05 ------------------- Notable fixes ============= - DNS - handle private values of RRSIG type (#83) ------------------- 0.12.2 - 2024-01-11 ------------------- Improvements ============ - Generic - add metadata to documentation ------------------- 0.12.1 - 2023-12-13 ------------------- Notable fixes ============= - SSH - add missing host key algorithms to key parser classes (#79) - Generic - fix markdown generation in the case of TLS client versions (#80) ------------------- 0.12.0 - 2023-11-23 ------------------- Features ======== - HTTP(S) (``http``) - Headers (``headers``) - add parsers for security related headers (`Content Security Policy `__ (CSP), `Content-Security-Policy-Report-Only `__) (#59) ------------------- 0.11.2 - 2023-11-13 ------------------- Features ======== - HTTP(S) (``http``) - Headers (``headers``) - add parsers for generic headers (`NEL `__ (Network Error Logging), `Set-Cookie `__) - add parsers for security related headers (`HTTP Public Key Pinning `__ (HPKP), `X-XSS-Protection `__) Improvements ============ - HTTP(S) (``http``) - Headers (``headers``) - implement detailed parsing of `Content-Type `__ header ------------------- 0.11.1 - 2023-11-06 ------------------- Features ======== - SSH (``ssh``) - Public Keys (``pubkeys``) - add X.509 certificate and certificate chain related classes (#63) ------------------- 0.11.0 - 2023-10-28 ------------------- Features ======== - Generic - add post processing capability to Markdown output (#73) - use class give grade for public keys (#73) ------------------- 0.10.3 - 2023-10-12 ------------------- Notable fixes ============= - Generic - add missing dnsrec module to the packaging (#75) ------------------- 0.10.2 - 2023-08-28 ------------------- Features ======== - DNS - add parser for e-mail authentication and reporting related records (#74, #35, #36, #37, #38) - `mail exchange `__ (MX) - `Domain-based Message Authentication, Reporting, and Conformance `__ (DMARC) - `Sender Policy Framework `__ (SPF) - `SMTP MTA Strict Transport Security `__ (MTA-STS) - `SMTP TLS Reporting `__ (TLSRPT) ------------------- 0.10.1 - 2023-08-29 ------------------- Features ======== - DNS - add parser for DNSSEC-related records (#72) - `DNSKEY `__ - `DS `__ - `RRSIG `__ ------------------- 0.10.0 - 2023-08-03 ------------------- Notable fixes ============= - Generic - Markdown output of attr-based classes ------------------ 0.9.1 - 2022-06-22 ------------------ Features ======== - TLS (``tls``) - Generic - add parser for `signed certificate timestamp `__ entries (#52) ------------------ 0.9.0 - 2023-04-29 ------------------ Features ======== - TLS (``tls``) - Generic - protocol item classes for `OpenVPN `__ support (#62) ------------------ 0.8.5 - 2023-04-02 ------------------ Features ======== - Generic - move data classes to `CryptoDataHub repository `__ (#67) ------------------ 0.8.4 - 2023-01-22 ------------------ Features ======== - TLS (``tls``) - Generic - protocol item classes for MySQL support (#61) ------------------ 0.8.2 - 2022-10-10 ------------------ Features ======== - TLS (``tls``) - Cipher Suites (``ciphers``) - add OpenSSL names (#54) - add min/max versions (#55) - SSH (``ssh``) - Public Keys (``pubkeys``) - `HASSH fingerprint `__ calculation (#48) - add `host certificate `__ related classes (#53) ------------------ 0.8.0 - 2022-01-18 ------------------ Features ======== - SSH (``ssh``) - Public Keys (``pubkeys``) - add `public key `__ related classes (#43) - Versions (``versions``) - add `software version `__ related classes (#46) ------------------ 0.7.3 - 2021-12-26 ------------------ Notable fixes ============= - Generic - Fix time zone handlind in datetime parser ------------------ 0.7.2 - 2021-10-07 ------------------ Other ===== - switch to Markdown format in changelog, readme and contributing - update contributing to the latest version from contribution-guide.org ------------------ 0.7.1 - 2021-09-20 ------------------ Features ======== - TLS (``tls``) - protocol item classes for PostgreSQL support (#44) ------------------ 0.7.0 - 2021-09-02 ------------------ Features ======== - TLS (``tls``) - Extensions (``extensions``) - add `application-layer protocol negotiation `__ extension related messages (#40) - add `encrypt-then-MAC `__ extension related messages (#40) - add `extended master secret `__ extension related messages (#40) - add `next protocol negotiation `__ extension related messages (#40) - add `renegotiation indication `__ extension related messages (#40) - add `session ticket `__ extension related messages (#40) ------------------ 0.6.0 - 2021-05-27 ------------------ Features ======== - HTTP(S) (``http``) - Headers (``headers``) - supports header wire format parsing - add parsers for generic headers (`Content-Type `__, `Server `__) - add parsers for cache related headers (`Age `__, `Cache-Control `__, `Date `__, `ETag `__, `Expires `__, `Last-Modified `__, `Pragma `__) - add parsers for security related headers (`Expect-CT `__, `Expect-Staple `__, `Referrer-Policy `__, `Strict-Transport-Security `__, `X-Content-Type-Options `__, `X-Frame-Options `__) - TLS (``tls``) - Versions (``versions``) - add `protocol version 1.3 `__ related messages (#20) - Cipher Suites (``ciphers``) - add `cipher suites `__ relate to version 1.3 (#20) - Diffie-Hellman (``dhparams``) - add `supported groups `__ relate to version 1.3 (#20) - Elliptic Curves (``curves``) - add `supported groups `__ relate to version 1.3 (#20) - Signature Algorithms (``sigalgos``) - add `signature algorithms `__ relate to version 1.3 (#20) ------------------ 0.5.0 - 2021-04-08 ------------------ Features ======== - Generic - add parser for `text-based protocols `__ (#21) - SSH (``ssh``) - Versions (``versions``) - add `protocol version exchange `__ related messages (#21) - SSH 2.0 (``ssh2``) - Cipher Suites (``ciphers``) - add `algorithm negotiation `__ related messages (#21) Usability ========= - Generic - show attributes in user-friendly order in Markdown output (#30) - use human readable algorithms names in Markdown output (#32) - add human readable descriptions for exceptions (#33) ------------------ 0.4.0 - 2021-01-30 ------------------ Features ======== - TLS (``tls``) - Generic - add `LDAP `__ related messages (#23) - Client Public Key Request (``pubkeyreq``) - add `client public key request `__ related messages (#24) Improvements ============ - Generic - add `OID `__ to algorithms ------------------ 0.3.1 - 2020-09-15 ------------------ Features ======== - Generic - `Markdown `__ serializable format (#19) Improvements ============ - TLS (``tls``) - Cipher Suites (``ciphers``) - add missing ``ECDHE_PSK`` cipher suites (#7) - add `GOST `__ cipher suites - add missing draft ECC cipher suites (#9) - add missing `FIPS `__ cipher suites (#11) - add `CECPQ1 `__ cipher suites (#12) - add missing `Fortezza `__ cipher suites (#13) - add missing ``DHE`` cipher suites (#14) - add missing SSLv3 cipher suites (#15) Notable fixes ============= - Generic - fix unicode string representation in JSON output (#18) - TLS (``tls``) - Cipher Suites (``ciphers``) - fix some cipher suite names and parameters (#7, #10) ------------------ 0.3.0 - 2020-04-30 ------------------ Features ======== - TLS (``tls``) - protocol item classes for RDP support (#4) - `JA3 fingerprint `__ calculation for TLS client hello (#2) Notable fixes ============= - TLS (``tls``) - compose all the messages in case of a TLS record (#1) Refactor ======== - use attrs to avoid boilerplates (#3) ------------------ 0.2.0 - 2019-12-02 ------------------ Notable fixes ============= - clarify TLS related parameter names - several packaging fixes ------------------ 0.1.0 - 2019-03-20 ------------------ Features ======== - added TLS record protocol support - added TLS ChangeCipherSpec message support - added TLS ApplicationData message support - added TLS handshake message support - added TLS client - added SSL support Improvements ============ - added serialization support for classes - added elliptic-curve related descriptive classes - added timeout parameter to TLS client class cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/CONTRIBUTING.rst000066400000000000000000000206331524413560000245170ustar00rootroot00000000000000Contributing ============ Submitting bugs --------------- Due diligence ~~~~~~~~~~~~~ Before submitting a bug, please do the following: - Perform **basic troubleshooting** steps: - **Make sure you're on the latest version.** If you're not on the most recent version, your problem may have been solved already! Upgrading is always the best first step. - **Try older versions.** If you're already *on* the latest release, try rolling back a few minor versions (e.g. if on 1.7, try 1.5 or 1.6) and see if the problem goes away. This will help the devs narrow down when the problem first arose in the commit log. - **Try switching up dependency versions.** If the software in question has dependencies (other libraries, etc) try upgrading/downgrading those as well. - **Search the project's bug/issue tracker** to make sure it's not a known issue. - If you don't find a pre-existing issue, consider **checking with the mailing list and/or IRC channel** in case the problem is non-bug-related. What to put in your bug report ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ Make sure your report gets the attention it deserves: bug reports with missing information may be ignored or punted back to you, delaying a fix. The below constitutes a bare minimum; more info is almost always better: - **What version of the core programming language interpreter/compiler are you using?** For example, if it's a Python project, are you using Python 2.7.3? Python 3.3.1? PyPy 2.0? - **What operating system are you on?** Windows? (Vista? 7? 32-bit? 64-bit?) Mac OS X? (10.7.4? 10.9.0?) Linux? (Which distro? Which version of that distro? 32 or 64 bits?) Again, more detail is better. - **Which version or versions of the software are you using?** Ideally, you followed the advice above and have ruled out (or verified that the problem exists in) a few different versions. - **How can the developers recreate the bug on their end?** If possible, include a copy of your code, the command you used to invoke it, and the full output of your run (if applicable.) - A common tactic is to pare down your code until a simple (but still bug-causing) "base case" remains. Not only can this help you identify problems which aren't real bugs, but it means the developer can get to fixing the bug faster. Contributing changes -------------------- Licensing of contributed material ~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ Keep in mind as you contribute, that code, docs and other material submitted to open source projects are usually considered licensed under the same terms as the rest of the work. The details vary from project to project, but from the perspective of this document's authors: - Anything submitted to a project falls under the licensing terms in the repository's top level ``LICENSE`` file. - For example, if a project's ``LICENSE`` is BSD-based, contributors should be comfortable with their work potentially being distributed in binary form without the original source code. - Per-file copyright/license headers are typically extraneous and undesirable. Please don't add your own copyright headers to new files unless the project's license actually requires them! - Not least because even a new file created by one individual (who often feels compelled to put their personal copyright notice at the top) will inherently end up contributed to by dozens of others over time, making a per-file header outdated/misleading. Version control branching ~~~~~~~~~~~~~~~~~~~~~~~~~ - Always **make a new branch** for your work, no matter how small. This makes it easy for others to take just that one set of changes from your repository, in case you have multiple unrelated changes floating around. - A corollary: **don't submit unrelated changes in the same branch/pull request**! The maintainer shouldn't have to reject your awesome bugfix because the feature you put in with it needs more review. - **Base your new branch off of the appropriate branch** on the main repository: - **Bug fixes** should be based on the branch named after the **oldest supported release line** the bug affects. - E.g. if a feature was introduced in 1.1, the latest release line is 1.3, and a bug is found in that feature - make your branch based on 1.1. The maintainer will then forward-port it to 1.3 and master. - Bug fixes requiring large changes to the code or which have a chance of being otherwise disruptive, may need to base off of **master** instead. This is a judgement call -- ask the devs! - **New features** should branch off of **the 'master' branch**. - Note that depending on how long it takes for the dev team to merge your patch, the copy of ``master`` you worked off of may get out of date! If you find yourself 'bumping' a pull request that's been sidelined for a while, **make sure you rebase or merge to latest master** to ensure a speedier resolution. Code formatting ~~~~~~~~~~~~~~~ - **Follow the style you see used in the primary repository**! Consistency with the rest of the project always trumps other considerations. It doesn't matter if you have your own style or if the rest of the code breaks with the greater community - just follow along. - Python projects usually follow the `PEP-8 `__ guidelines (though many have minor deviations depending on the lead maintainers' preferences.) Documentation isn't optional ~~~~~~~~~~~~~~~~~~~~~~~~~~~~ It's not! Patches without documentation will be returned to sender. By "documentation" we mean: - **Docstrings** (for Python; or API-doc-friendly comments for other languages) must be created or updated for public API functions/methods/etc. (This step is optional for some bugfixes.) - Don't forget to include `versionadded `__/`versionchanged `__ ReST directives at the bottom of any new or changed Python docstrings! - Use ``versionadded`` for truly new API members -- new methods, functions, classes or modules. - Use ``versionchanged`` when adding/removing new function/method arguments, or whenever behavior changes. - New features should ideally include updates to **prose documentation**, including useful example code snippets. - All submissions should have a **changelog entry** crediting the contributor and/or any individuals instrumental in identifying the problem. Tests aren't optional ~~~~~~~~~~~~~~~~~~~~~ Any bugfix that doesn't include a test proving the existence of the bug being fixed, may be suspect. Ditto for new features that can't prove they actually work. We've found that test-first development really helps make features better architected and identifies potential edge cases earlier instead of later. Writing tests before the implementation is strongly encouraged. Full example ~~~~~~~~~~~~ Here's an example workflow for a project ``theproject`` hosted on Github, which is currently in version 1.3.x. Your username is ``yourname`` and you're submitting a basic bugfix. (This workflow only changes slightly if the project is hosted at Bitbucket, self-hosted, or etc.) Preparing your Fork ^^^^^^^^^^^^^^^^^^^ 1. Click 'Fork' on Github, creating e.g. ``yourname/theproject``. 2. Clone your project: ``git clone git@github.com:yourname/theproject``. 3. ``cd theproject`` 4. `Create and activate a virtual environment `__. 5. Install the development requirements: ``uv sync --extra tests``. 6. Create a branch: ``git checkout -b foo-the-bars 1.3``. Making your Changes ^^^^^^^^^^^^^^^^^^^ 1. Add changelog entry crediting yourself. 2. Write tests expecting the correct/fixed functionality; make sure they fail. 3. Hack, hack, hack. 4. Run tests again, making sure they pass. 5. Commit your changes: ``git commit -m "Foo the bars"`` Creating Pull Requests ^^^^^^^^^^^^^^^^^^^^^^ 1. Push your commit to get it back up to your fork: ``git push origin HEAD`` 2. Visit Github, click handy "Pull request" button that it will make upon noticing your new branch. 3. In the description field, write down issue number (if submitting code fixing an existing issue) or describe the issue + your fix (if submitting a wholly new bugfix). 4. Hit 'submit'! And please be patient - the maintainers will get to you when they can. cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/LICENSE.txt000066400000000000000000000405261524413560000237040ustar00rootroot00000000000000Mozilla Public License Version 2.0 ================================== 1. Definitions -------------- 1.1. "Contributor" means each individual or legal entity that creates, contributes to the creation of, or owns Covered Software. 1.2. "Contributor Version" means the combination of the Contributions of others (if any) used by a Contributor and that particular Contributor's Contribution. 1.3. "Contribution" means Covered Software of a particular Contributor. 1.4. "Covered Software" means Source Code Form to which the initial Contributor has attached the notice in Exhibit A, the Executable Form of such Source Code Form, and Modifications of such Source Code Form, in each case including portions thereof. 1.5. "Incompatible With Secondary Licenses" means (a) that the initial Contributor has attached the notice described in Exhibit B to the Covered Software; or (b) that the Covered Software was made available under the terms of version 1.1 or earlier of the License, but not also under the terms of a Secondary License. 1.6. "Executable Form" means any form of the work other than Source Code Form. 1.7. "Larger Work" means a work that combines Covered Software with other material, in a separate file or files, that is not Covered Software. 1.8. "License" means this document. 1.9. "Licensable" means having the right to grant, to the maximum extent possible, whether at the time of the initial grant or subsequently, any and all of the rights conveyed by this License. 1.10. "Modifications" means any of the following: (a) any file in Source Code Form that results from an addition to, deletion from, or modification of the contents of Covered Software; or (b) any new file in Source Code Form that contains any Covered Software. 1.11. "Patent Claims" of a Contributor means any patent claim(s), including without limitation, method, process, and apparatus claims, in any patent Licensable by such Contributor that would be infringed, but for the grant of the License, by the making, using, selling, offering for sale, having made, import, or transfer of either its Contributions or its Contributor Version. 1.12. "Secondary License" means either the GNU General Public License, Version 2.0, the GNU Lesser General Public License, Version 2.1, the GNU Affero General Public License, Version 3.0, or any later versions of those licenses. 1.13. "Source Code Form" means the form of the work preferred for making modifications. 1.14. "You" (or "Your") means an individual or a legal entity exercising rights under this License. For legal entities, "You" includes any entity that controls, is controlled by, or is under common control with You. For purposes of this definition, "control" means (a) the power, direct or indirect, to cause the direction or management of such entity, whether by contract or otherwise, or (b) ownership of more than fifty percent (50%) of the outstanding shares or beneficial ownership of such entity. 2. License Grants and Conditions -------------------------------- 2.1. Grants Each Contributor hereby grants You a world-wide, royalty-free, non-exclusive license: (a) under intellectual property rights (other than patent or trademark) Licensable by such Contributor to use, reproduce, make available, modify, display, perform, distribute, and otherwise exploit its Contributions, either on an unmodified basis, with Modifications, or as part of a Larger Work; and (b) under Patent Claims of such Contributor to make, use, sell, offer for sale, have made, import, and otherwise transfer either its Contributions or its Contributor Version. 2.2. Effective Date The licenses granted in Section 2.1 with respect to any Contribution become effective for each Contribution on the date the Contributor first distributes such Contribution. 2.3. Limitations on Grant Scope The licenses granted in this Section 2 are the only rights granted under this License. No additional rights or licenses will be implied from the distribution or licensing of Covered Software under this License. Notwithstanding Section 2.1(b) above, no patent license is granted by a Contributor: (a) for any code that a Contributor has removed from Covered Software; or (b) for infringements caused by: (i) Your and any other third party's modifications of Covered Software, or (ii) the combination of its Contributions with other software (except as part of its Contributor Version); or (c) under Patent Claims infringed by Covered Software in the absence of its Contributions. This License does not grant any rights in the trademarks, service marks, or logos of any Contributor (except as may be necessary to comply with the notice requirements in Section 3.4). 2.4. Subsequent Licenses No Contributor makes additional grants as a result of Your choice to distribute the Covered Software under a subsequent version of this License (see Section 10.2) or under the terms of a Secondary License (if permitted under the terms of Section 3.3). 2.5. Representation Each Contributor represents that the Contributor believes its Contributions are its original creation(s) or it has sufficient rights to grant the rights to its Contributions conveyed by this License. 2.6. Fair Use This License is not intended to limit any rights You have under applicable copyright doctrines of fair use, fair dealing, or other equivalents. 2.7. Conditions Sections 3.1, 3.2, 3.3, and 3.4 are conditions of the licenses granted in Section 2.1. 3. Responsibilities ------------------- 3.1. Distribution of Source Form All distribution of Covered Software in Source Code Form, including any Modifications that You create or to which You contribute, must be under the terms of this License. You must inform recipients that the Source Code Form of the Covered Software is governed by the terms of this License, and how they can obtain a copy of this License. You may not attempt to alter or restrict the recipients' rights in the Source Code Form. 3.2. Distribution of Executable Form If You distribute Covered Software in Executable Form then: (a) such Covered Software must also be made available in Source Code Form, as described in Section 3.1, and You must inform recipients of the Executable Form how they can obtain a copy of such Source Code Form by reasonable means in a timely manner, at a charge no more than the cost of distribution to the recipient; and (b) You may distribute such Executable Form under the terms of this License, or sublicense it under different terms, provided that the license for the Executable Form does not attempt to limit or alter the recipients' rights in the Source Code Form under this License. 3.3. Distribution of a Larger Work You may create and distribute a Larger Work under terms of Your choice, provided that You also comply with the requirements of this License for the Covered Software. If the Larger Work is a combination of Covered Software with a work governed by one or more Secondary Licenses, and the Covered Software is not Incompatible With Secondary Licenses, this License permits You to additionally distribute such Covered Software under the terms of such Secondary License(s), so that the recipient of the Larger Work may, at their option, further distribute the Covered Software under the terms of either this License or such Secondary License(s). 3.4. Notices You may not remove or alter the substance of any license notices (including copyright notices, patent notices, disclaimers of warranty, or limitations of liability) contained within the Source Code Form of the Covered Software, except that You may alter any license notices to the extent required to remedy known factual inaccuracies. 3.5. Application of Additional Terms You may choose to offer, and to charge a fee for, warranty, support, indemnity or liability obligations to one or more recipients of Covered Software. However, You may do so only on Your own behalf, and not on behalf of any Contributor. You must make it absolutely clear that any such warranty, support, indemnity, or liability obligation is offered by You alone, and You hereby agree to indemnify every Contributor for any liability incurred by such Contributor as a result of warranty, support, indemnity or liability terms You offer. You may include additional disclaimers of warranty and limitations of liability specific to any jurisdiction. 4. Inability to Comply Due to Statute or Regulation --------------------------------------------------- If it is impossible for You to comply with any of the terms of this License with respect to some or all of the Covered Software due to statute, judicial order, or regulation then You must: (a) comply with the terms of this License to the maximum extent possible; and (b) describe the limitations and the code they affect. Such description must be placed in a text file included with all distributions of the Covered Software under this License. Except to the extent prohibited by statute or regulation, such description must be sufficiently detailed for a recipient of ordinary skill to be able to understand it. 5. Termination -------------- 5.1. The rights granted under this License will terminate automatically if You fail to comply with any of its terms. However, if You become compliant, then the rights granted under this License from a particular Contributor are reinstated (a) provisionally, unless and until such Contributor explicitly and finally terminates Your grants, and (b) on an ongoing basis, if such Contributor fails to notify You of the non-compliance by some reasonable means prior to 60 days after You have come back into compliance. Moreover, Your grants from a particular Contributor are reinstated on an ongoing basis if such Contributor notifies You of the non-compliance by some reasonable means, this is the first time You have received notice of non-compliance with this License from such Contributor, and You become compliant prior to 30 days after Your receipt of the notice. 5.2. If You initiate litigation against any entity by asserting a patent infringement claim (excluding declaratory judgment actions, counter-claims, and cross-claims) alleging that a Contributor Version directly or indirectly infringes any patent, then the rights granted to You by any and all Contributors for the Covered Software under Section 2.1 of this License shall terminate. 5.3. In the event of termination under Sections 5.1 or 5.2 above, all end user license agreements (excluding distributors and resellers) which have been validly granted by You or Your distributors under this License prior to termination shall survive termination. ************************************************************************ * * * 6. Disclaimer of Warranty * * ------------------------- * * * * Covered Software is provided under this License on an "as is" * * basis, without warranty of any kind, either expressed, implied, or * * statutory, including, without limitation, warranties that the * * Covered Software is free of defects, merchantable, fit for a * * particular purpose or non-infringing. The entire risk as to the * * quality and performance of the Covered Software is with You. * * Should any Covered Software prove defective in any respect, You * * (not any Contributor) assume the cost of any necessary servicing, * * repair, or correction. This disclaimer of warranty constitutes an * * essential part of this License. No use of any Covered Software is * * authorized under this License except under this disclaimer. * * * ************************************************************************ ************************************************************************ * * * 7. Limitation of Liability * * -------------------------- * * * * Under no circumstances and under no legal theory, whether tort * * (including negligence), contract, or otherwise, shall any * * Contributor, or anyone who distributes Covered Software as * * permitted above, be liable to You for any direct, indirect, * * special, incidental, or consequential damages of any character * * including, without limitation, damages for lost profits, loss of * * goodwill, work stoppage, computer failure or malfunction, or any * * and all other commercial damages or losses, even if such party * * shall have been informed of the possibility of such damages. This * * limitation of liability shall not apply to liability for death or * * personal injury resulting from such party's negligence to the * * extent applicable law prohibits such limitation. Some * * jurisdictions do not allow the exclusion or limitation of * * incidental or consequential damages, so this exclusion and * * limitation may not apply to You. * * * ************************************************************************ 8. Litigation ------------- Any litigation relating to this License may be brought only in the courts of a jurisdiction where the defendant maintains its principal place of business and such litigation shall be governed by laws of that jurisdiction, without reference to its conflict-of-law provisions. Nothing in this Section shall prevent a party's ability to bring cross-claims or counter-claims. 9. Miscellaneous ---------------- This License represents the complete agreement concerning the subject matter hereof. If any provision of this License is held to be unenforceable, such provision shall be reformed only to the extent necessary to make it enforceable. Any law or regulation which provides that the language of a contract shall be construed against the drafter shall not be used to construe this License against a Contributor. 10. Versions of the License --------------------------- 10.1. New Versions Mozilla Foundation is the license steward. Except as provided in Section 10.3, no one other than the license steward has the right to modify or publish new versions of this License. Each version will be given a distinguishing version number. 10.2. Effect of New Versions You may distribute the Covered Software under the terms of the version of the License under which You originally received the Covered Software, or under the terms of any subsequent version published by the license steward. 10.3. Modified Versions If you create software not governed by this License, and you want to create a new license for such software, you may create and use a modified version of this License if you rename the license and remove any references to the name of the license steward (except to note that such modified license differs from this License). 10.4. Distributing Source Code Form that is Incompatible With Secondary Licenses If You choose to distribute Source Code Form that is Incompatible With Secondary Licenses under the terms of this version of the License, the notice described in Exhibit B of this License must be attached. Exhibit A - Source Code Form License Notice ------------------------------------------- This Source Code Form is subject to the terms of the Mozilla Public License, v. 2.0. If a copy of the MPL was not distributed with this file, You can obtain one at http://mozilla.org/MPL/2.0/. If it is not possible or desirable to put the notice in a particular file, then You may include the notice in a location (such as a LICENSE file in a relevant directory) where a recipient would be likely to look for such a notice. You may add additional accurate notices of copyright ownership. Exhibit B - "Incompatible With Secondary Licenses" Notice --------------------------------------------------------- This Source Code Form is "Incompatible With Secondary Licenses", as defined by the Mozilla Public License, v. 2.0. cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/MANIFEST.in000066400000000000000000000000541524413560000236070ustar00rootroot00000000000000include *.md include *.rst prune submodules cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/README.md000066400000000000000000000106551524413560000233400ustar00rootroot00000000000000[![Pipeline](https://gitlab.com/coroner/cryptoparser/badges/master/pipeline.svg)](https://gitlab.com/coroner/cryptoparser/-/pipelines/master/latest) [![Test Coverage](https://coveralls.io/repos/gitlab/coroner/cryptoparser/badge.svg?branch=master)](https://coveralls.io/gitlab/coroner/cryptoparser/) [![Documentation](https://readthedocs.org/projects/cryptoparser/badge/?version=latest)](https://cryptoparser.readthedocs.io) **CryptoParser** is a cryptographic protocol ([IKE](https://en.wikipedia.org/wiki/Internet_Key_Exchange), [SSL](https://en.wikipedia.org/wiki/Transport_Layer_Security#SSL_1.0,_2.0,_and_3.0), [TLS](https://en.wikipedia.org/wiki/Transport_Layer_Security), [SSH](https://en.wikipedia.org/wiki/Secure_Shell), [DNSSEC](https://en.wikipedia.org/wiki/Domain_Name_System_Security_Extensions)) and security-related protocol piece ([HTTP headers](https://en.wikipedia.org/wiki/List_of_HTTP_header_fields)) parser and generator. It is neither a comprehensive nor a secure implementation of any cryptographic protocol. The goal is to support testing cryptographic libraries or analysing cryptography-related settings of application servers such as [CryptoLyzer](https://cryptolyzer.readthedocs.io/) does. **Use CryptoParser when you need to parse handshake messages** — it implements the wire format of ISAKMP, IKEv1, IKEv2, SSL 2.0, SSL 3.0, TLS 1.0 to TLS 1.3, and SSH 2.0, so a message can be read field by field instead of being handed to a connection-oriented library. **Use CryptoParser when you need to generate messages a library refuses to send** — analysis means triggering special and corner cases, so messages can be composed with deprecated, experimental, or plainly invalid values that an implementation aiming for secure connections would reject. **Use CryptoParser when you need parsed HTTP security headers** — beyond wire format parsing it parses individual headers, so directives are available as typed values rather than as strings. **Use CryptoParser when you need DNSSEC record parsing** — DNSSEC records are parsed into the same kind of typed structures as the protocol messages. The strength of CryptoParser is that it is backed by the most comprehensive algorithm identifier database available ([CryptoDataHub](https://cryptodatahub.readthedocs.io)). This makes it possible to recognize rarely used, deprecated, non-standard, or experimental algorithms that are not supported by any version of OpenSSL, GnuTLS, LibreSSL, or wolfSSL. ## Why CryptoParser? - **Analysis oriented** — the library implements only the parts of a protocol that analysis needs, and deliberately keeps the parts that make a message invalid. - **No OpenSSL dependency** — the protocol implementation is its own, so what can be parsed or generated is not limited by what a cryptographic library is willing to do. - **Typed values, not byte offsets** — algorithm identifiers, protocol versions, and header directives are parsed into enumerations and attribute classes. - **Symmetric parsing and composing** — every parsable piece can also be composed back to its wire format. ## Usage ### uv ```shell uv add cryptoparser ``` ```python from cryptoparser.tls.version import TlsProtocolVersion # parse a protocol version from its wire format TlsProtocolVersion.parse_exact_size(b'\x03\x03') # TLS 1.2 # compose a protocol version back to its wire format TlsProtocolVersion.parse_exact_size(b'\x03\x04').compose() # bytearray(b'\x03\x04') ``` ## Support **Python implementations** - CPython 3.9+ - PyPy 3.9+ **Operating systems** - Linux - macOS - Windows ## Documentation Detailed [documentation](https://cryptoparser.readthedocs.io) is available on the project's [Read the Docs](https://readthedocs.com) site. ## License The [code](https://gitlab.com/coroner/cryptoparser) is available under the terms of [Mozilla Public License Version 2.0](https://www.mozilla.org/en-US/MPL/2.0/) (MPL 2.0). A non-comprehensive but straightforward description of MPL 2.0 can be found at the [Choose an open source license](https://choosealicense.com/licenses#mpl-2.0) website. ## Credits - [NLnet Foundation](https://nlnet.nl) and [NGI Assure](https://www.assure.ngi.eu), supports the project part of the [Next Generation Internet](https://ngi.eu) initiative. cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser.spec000066400000000000000000000107621524413560000254710ustar00rootroot00000000000000Name: python-cryptoparser Version: 1.6.0 Release: 1%{?dist} Summary: Multi-protocol cryptographic protocol parser library License: MPL-2.0 URL: https://gitlab.com/coroner/cryptoparser Source0: %{name}_%{version}.tar.xz BuildArch: noarch BuildRequires: python3-devel BuildRequires: python3-pip BuildRequires: python3-setuptools BuildRequires: python3-wheel BuildRequires: python3-cryptodatahub >= 1.6.0 %description CryptoParser is a library for parsing cryptographic protocol messages including TLS, SSH, IKE, and related protocols. It is used as the parsing backend for CryptoLyzer. %package -n python3-cryptoparser Summary: %{summary} Requires: python3-asn1crypto Requires: python3-attrs Requires: python3-cryptodatahub >= 1.6.0 Requires: python3-urllib3 %description -n python3-cryptoparser CryptoParser is a library for parsing cryptographic protocol messages including TLS, SSH, IKE, and related protocols. It is used as the parsing backend for CryptoLyzer. %prep %setup -q -T -c -n %{name}-%{version} tar -xJf %{SOURCE0} --strip-components=1 sed -i "s/, 'setuptools-scm'//" pyproject.toml sed -i "s/name = 'CryptoParser'/name = 'cryptoparser'/" pyproject.toml sed -i "s/exclude = \['submodules'\]/include = ['cryptoparser*']/" pyproject.toml %build export SETUPTOOLS_SCM_PRETEND_VERSION=%{version} %install export SETUPTOOLS_SCM_PRETEND_VERSION=%{version} %{__python3} -m pip install --no-build-isolation --no-deps --root %{buildroot} --prefix %{_prefix} . %check %files -n python3-cryptoparser %{python3_sitelib}/cryptoparser/ %{python3_sitelib}/cryptoparser-%{version}.dist-info/ %license LICENSE.txt %changelog * Tue Aug 25 2026 Szilárd Pfeiffer - 1.6.0-1 - add trust anchors extension related messages (#104) - add server padding extension related messages (#104) - add application-layer protocol settings extension related messages with the old code point (#104) - bound the outer extension payload of the encrypted client hello by the extension length (#104) - preserve the transform number when parsing the IKEv1 transform payload (#105) - follow the replacement of the named group with the key parameter (#106) - make the parameter classes of the invalid extension types immutable (#104) * Fri Jul 31 2026 Szilárd Pfeiffer - 1.5.0-1 - add IKEv1 and IKEv2 certificate payload parsing (#103) - add IKEv1 identification payload parsing (#103) - add IKEv1 signature payload parsing (#103) - add IKEv2 identification payload parsing (#103) - add IKEv2 authentication payload parsing (#103) - add IKEv2 extensible authentication protocol payload parsing (#103) - add IKEv2 encrypted and authenticated payload parsing (#103) - keep IKEv1 and IKEv2 payloads of unknown type unparsed instead of rejecting the message (#103) - parse the IKEv1 and IKEv2 identification payload data according to the identification type (#103) - parse the digital signature envelope of the IKEv2 authentication payload (#103) - split the ISAKMP message parsing and composing into header and payload chain steps (#103) - use typed values in the IKEv2 signature hash algorithms notify payload (#103) - add getter for the distinguished name of the IKEv1 certificate request payload (#103) - add getter for the certification authority hashes of the IKEv2 certificate request payload (#103) * Fri Jul 17 2026 Szilárd Pfeiffer - 1.4.0-1 - add getter for multiple payloads with the same type (#93) - add IKEv2 NAT detection source IP and destination IP notify payload parsing (#93) - add IKEv2 set window size notify payload parsing (#93) - add IKEv2 use transport mode notify payload parsing (#93) - add IKEv2 HTTP certificate lookup supported notify payload parsing (#93) - add IKEv2 signature hash algorithms notify payload parsing (#93) - add IKEv2 intermediate exchange supported notify payload parsing (#93) - add IKEv2 use PPK notify payload parsing (#93) - add IKEv2 redirect supported notify payload parsing (#93) - add IKEv2 fragmentation supported notify payload parsing (#93) - add childless IKEv2 supported notify payload parsing (#93) * Mon Jun 15 2026 Szilárd Pfeiffer - 1.3.0-1 - add Debian and RPM packaging (#102) - add JA4 tag generation for the client hello message (#100) - add JA4X tag generation for X.509 certificates (#101) - add certificate-related messages for protocol version 1.3 (#94) - make IKEv2 transform key length optional for fixed-key ciphers (#99) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/000077500000000000000000000000001524413560000246075ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/__init__.py000066400000000000000000000000431524413560000267150ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/__setup__.py000066400000000000000000000007241524413560000271200ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import importlib.metadata metadata = importlib.metadata.metadata('cryptoparser') __title__ = metadata['Name'] __technical_name__ = __title__.lower() __version__ = metadata['Version'] __description__ = metadata['Summary'] __author__ = metadata['Author'] __author_email__ = metadata['Author-email'] __url__ = 'https://gitlab.com/coroner/' + __technical_name__ __license__ = metadata.get('License-Expression') or metadata['License'] cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/common/000077500000000000000000000000001524413560000260775ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/common/__init__.py000066400000000000000000000000431524413560000302050ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/common/base.py000066400000000000000000001003361524413560000273660ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import abc import datetime import enum import json import math import types import ipaddress try: from collections.abc import MutableSequence # only works on python 3.3+ except ImportError: # pragma: no cover from collections.abc import MutableSequence # pylint: disable=deprecated-class from collections import OrderedDict import attr import urllib3 from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.grade import Gradeable from cryptodatahub.common.types import CryptoDataEnumCodedBase, CryptoDataParamsBase from cryptoparser.common.parse import ( ComposerBinary, ComposerText, ParsableBase, ParsableBaseNoABC, ParserBinary, ParserText, ) from cryptoparser.common.exception import NotEnoughData, TooMuchData, InvalidType from cryptoparser.common.utils import bytes_to_hex_string def _default( self, # pylint: disable=unused-argument obj ): result = Serializable._json_traverse(obj, Serializable._json_result) # pylint: disable=protected-access return result _default.default = json.JSONEncoder().default json.JSONEncoder.default = _default class SerializableTextEncoder: def __call__(self, obj, level): if isinstance(obj, str): string_result = obj else: string_result = str(obj) return False, string_result class Serializable: # pylint: disable=too-few-public-methods _MARKDOWN_RESULT_STRING_CLASSES = ( ipaddress.IPv4Network, ipaddress.IPv6Network, urllib3.util.url.Url, ) post_text_encoder = SerializableTextEncoder() @staticmethod def _filter_out_non_human_friendly(obj, dict_value, human_friendly_only): if not attr.has(type(obj)) or not human_friendly_only: return dict_value fields_dict = attr.fields_dict(type(obj)) dict_value = OrderedDict([ (name, value) for name, value in dict_value.items() if name not in fields_dict or fields_dict[name].metadata.get('human_friendly', True) ]) return dict_value @staticmethod def _get_ordered_dict(dict_value, human_friendly_only=False): if attr.has(type(dict_value)): obj = dict_value dict_value = OrderedDict([ (name, getattr(dict_value, name)) for name, field in attr.fields_dict(type(dict_value)).items() if not name.startswith('_') ]) dict_value = Serializable._filter_out_non_human_friendly(obj, dict_value, human_friendly_only) keys = dict_value.keys() elif isinstance(dict_value, OrderedDict): keys = dict_value.keys() elif isinstance(dict_value, dict): if all(isinstance(key, enum.Enum) for key in dict_value.keys()): keys = sorted(dict_value.keys(), key=lambda key: key.name) else: keys = sorted(dict_value.keys()) elif hasattr(dict_value, '__dict__'): dict_value = dict_value.__dict__ keys = sorted(filter(lambda key: not key.startswith('_'), dict_value.keys())) result = OrderedDict([ (key, dict_value[key]) for key in keys ]) return result @staticmethod def _json_result(obj): if isinstance(obj, enum.Enum): if isinstance(obj.value, CryptoDataParamsBase): result = obj.name else: result = {obj.name: obj.value} elif isinstance(obj, (str, int, float, bool, )) or obj is None: result = obj elif isinstance(obj, (bytes, bytearray)): result = bytes_to_hex_string(obj, separator=':', lowercase=False) else: result = str(obj) return result @staticmethod def _json_traverse(obj, result_func): if isinstance(obj, enum.Enum): result = result_func(obj) elif hasattr(obj, '_asdict'): result = Serializable._json_traverse(obj._asdict(), result_func) elif isinstance(obj, dict) or attr.has(type(obj)): result = OrderedDict([ ( key.name if isinstance(key, enum.Enum) else Serializable._json_result(key), Serializable._json_traverse(value, result_func) ) for key, value in Serializable._get_ordered_dict(obj).items() ]) elif hasattr(obj, '__dict__'): result = Serializable._json_traverse(obj.__dict__, result_func) elif isinstance(obj, (list, tuple, frozenset, set)): result = [Serializable._json_traverse(item, result_func) for item in obj] else: result = result_func(obj) return result @staticmethod def _markdown_indent_from_level(level): return 4 * level * ' ' @classmethod def _markdown_human_readable_names(cls, obj, dict_value): name_dict = {} fields_dict = attr.fields_dict(type(obj)) if attr.has(type(obj)) else {} for name in dict_value: if isinstance(name, str): if name in fields_dict and 'human_readable_name' in fields_dict[name].metadata: human_readable_name = fields_dict[name].metadata['human_readable_name'] else: human_readable_name = ' '.join(name.split('_')).title() else: post_text_encoder = cls.post_text_encoder cls.post_text_encoder = SerializableTextEncoder() _, human_readable_name = cls._markdown_result(name) cls.post_text_encoder = post_text_encoder name_dict[name] = human_readable_name return name_dict @classmethod def _markdown_result_complex(cls, obj, level=0): indent = Serializable._markdown_indent_from_level(level) if hasattr(obj, '_asdict'): dict_value = obj._asdict() if not isinstance(dict_value, dict): return False, dict_value dict_value = Serializable._filter_out_non_human_friendly(obj, dict_value, human_friendly_only=True) else: dict_value = Serializable._get_ordered_dict(obj, human_friendly_only=True) result = '' name_dict = cls._markdown_human_readable_names(obj, dict_value) for name, value in dict_value.items(): result += f'{indent}* {name_dict[name]}' multiline, markdnow_result = cls._markdown_result(value, level + 1) if multiline: result += f':\n{markdnow_result}' else: result += f': {markdnow_result}\n' if not result: return False, '-' return True, result @classmethod def _markdown_result_list(cls, obj, level=0): if not obj: return False, '-' indent = Serializable._markdown_indent_from_level(level) result = '' for index, item in enumerate(obj): multiline, markdnow_result = cls._markdown_result(item, level + 1) separator = '\n' if multiline else ' ' newline = '' if multiline else '\n' result += f'{indent}{index + 1}.{separator}{markdnow_result}{newline}' return True, result @staticmethod def _markdown_is_directly_printable(obj): return not isinstance(obj, enum.Enum) and isinstance(obj, (str, int, float, )) @classmethod def _markdown_result(cls, obj, level=0): # pylint: disable=too-many-branches,too-many-return-statements if obj is None: result = cls.post_text_encoder('n/a', level) elif isinstance(obj, bool): result = cls.post_text_encoder('yes' if obj else 'no', level) elif Serializable._markdown_is_directly_printable(obj): result = cls.post_text_encoder(obj, level) elif isinstance(obj, Gradeable): result = cls.post_text_encoder(obj, level) elif isinstance(obj, Serializable): result = obj._as_markdown(level) # pylint: disable=protected-access elif isinstance(obj, enum.Enum): if isinstance(obj.value, Serializable): return obj.value._as_markdown(level) # pylint: disable=protected-access if isinstance(obj.value, CryptoDataParamsBase): return cls.post_text_encoder(obj.value, level) return cls.post_text_encoder(obj.name, level) elif isinstance(obj, cls._MARKDOWN_RESULT_STRING_CLASSES): return False, str(obj) elif isinstance(obj, datetime.timedelta): return False, str(obj.seconds) elif isinstance(obj, CryptoDataParamsBase) and hasattr(obj, '__str__'): return False, str(obj) elif attr.has(type(obj)): result = cls._markdown_result_complex(obj, level) elif hasattr(obj, '_asdict'): result = cls._markdown_result(obj._asdict(), level) elif hasattr(obj, '__dict__') or isinstance(obj, dict): result = cls._markdown_result_complex(obj, level) elif isinstance(obj, (list, tuple, frozenset, set, ArrayBase)): result = cls._markdown_result_list(obj, level) elif isinstance(obj, (bytes, bytearray)): result = cls.post_text_encoder(bytes_to_hex_string(obj, separator=':', lowercase=False), level) else: result = cls.post_text_encoder(obj, level) return result def _asdict(self): return Serializable._get_ordered_dict(self) def as_json(self): return json.dumps(self) def _as_markdown(self, level): return self._markdown_result_complex(self, level) def as_markdown(self): _, result = self._as_markdown(0) return result @attr.s class VariantParsableBase(ParsableBase): variant = attr.ib() _REGISTERED_VARIANTS = OrderedDict() @classmethod @abc.abstractmethod def _get_variants(cls): raise NotImplementedError() @variant.validator def _validator_variant(self, _, value): for variant_type in self._get_variant_types(): if issubclass(variant_type, NByteEnumParsable): variant_type = variant_type.get_enum_class() if isinstance(value, variant_type): break else: raise InvalidValue(value, VariantParsable) @classmethod def _get_variant_types(cls): variant_types = [] for variant_type_list in list(cls._get_variants().values()) + list(cls._get_registered_variants().values()): variant_types.extend(variant_type_list) return variant_types @classmethod def _get_registered_variants(cls): if cls not in cls._REGISTERED_VARIANTS: cls._REGISTERED_VARIANTS[cls] = OrderedDict() return cls._REGISTERED_VARIANTS[cls] @classmethod def register_variant_parser(cls, variant_tag, parsable_class): registered_variants = cls._get_registered_variants() if variant_tag not in registered_variants: registered_variants[variant_tag] = [] registered_variants[variant_tag].append(parsable_class) @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() def compose(self): return self.variant.compose() class VariantParsable(VariantParsableBase): @classmethod @abc.abstractmethod def _get_variants(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): for variant_parser in cls._get_variant_types(): try: parsed_object, parsed_length = variant_parser.parse_immutable(parsable) return parsed_object, parsed_length except InvalidType: pass raise InvalidValue(parsable, cls) class VariantParsableExact(VariantParsableBase): @classmethod @abc.abstractmethod def _get_variants(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): for variant_parser in cls._get_variant_types(): try: parsed_object = variant_parser.parse_exact_size(parsable) return parsed_object, len(parsable) except (InvalidType, InvalidValue, TooMuchData): pass raise InvalidValue(parsable, cls) @attr.s class VectorParamBase: # pylint: disable=too-few-public-methods min_byte_num = attr.ib(validator=attr.validators.instance_of(int)) max_byte_num = attr.ib(validator=attr.validators.instance_of(int)) item_num_size = attr.ib(init=False, validator=attr.validators.instance_of(int)) def __attrs_post_init__(self): self.item_num_size = int(math.log(self.max_byte_num, 2) / 8) + 1 attr.validate(self) @abc.abstractmethod def get_item_size(self, item): raise NotImplementedError() @attr.s class VectorParamNumeric(VectorParamBase): # pylint: disable=too-few-public-methods item_size = attr.ib(validator=attr.validators.instance_of(int)) numeric_class = attr.ib(default=int, validator=attr.validators.instance_of(type)) def get_item_size(self, item): return self.item_size @attr.s(init=False) class OpaqueParam(VectorParamNumeric): # pylint: disable=too-few-public-methods def __init__(self, min_byte_num, max_byte_num): super().__init__(min_byte_num, max_byte_num, 1) def get_item_size(self, item): return 1 @attr.s class VectorParamString(VectorParamBase): # pylint: disable=too-few-public-methods separator = attr.ib(validator=attr.validators.instance_of(str), default=',') encoding = attr.ib(validator=attr.validators.instance_of(str), default='ascii') item_class = attr.ib(validator=attr.validators.instance_of((type, types.FunctionType)), default=str) fallback_class = attr.ib( default=None, validator=attr.validators.optional(attr.validators.instance_of((type, types.FunctionType))) ) def get_item_size(self, item): if isinstance(item, (ParsableBase, StringEnumParsable)): return len(item.compose()) if isinstance(item, CryptoDataEnumCodedBase): return item.value.get_code_size() if isinstance(item, str): return len(item) raise NotImplementedError(type(item)) @attr.s class VectorParamParsable(VectorParamBase): # pylint: disable=too-few-public-methods item_class = attr.ib(validator=attr.validators.instance_of((type, types.FunctionType))) fallback_class = attr.ib( validator=attr.validators.optional(attr.validators.instance_of((type, types.FunctionType))) ) def get_item_size(self, item): return len(item.compose()) @attr.s class VectorParamEnumCodeNumeric(VectorParamBase): # pylint: disable=too-few-public-methods item_class = attr.ib(validator=attr.validators.instance_of((type, types.FunctionType))) fallback_class = attr.ib( validator=attr.validators.optional(attr.validators.instance_of((type, types.FunctionType))) ) def get_item_size(self, item): return self.fallback_class.get_byte_num() @attr.s class VectorParamEnumCodeString(VectorParamBase): # pylint: disable=too-few-public-methods item_class = attr.ib(validator=attr.validators.instance_of((type, types.FunctionType))) fallback_class = attr.ib(init=False, default=None) def get_item_size(self, item): return len(item.value.code) @attr.s class ArrayBase(ParsableBase, MutableSequence, Serializable): _items = attr.ib() _items_size = attr.ib(init=False, default=0) param = attr.ib(init=False, default=None) def __attrs_post_init__(self): items = self._items self.param = self.get_param() self._items = [] for item in items: self._items.append(item) self._items_size += self.param.get_item_size(item) self._update_items_size(del_item=None, insert_item=None) attr.validate(self) def _update_items_size(self, del_item=None, insert_item=None): size_diff = 0 if del_item is not None: size_diff -= self.param.get_item_size(del_item) if insert_item is not None: size_diff += self.param.get_item_size(insert_item) if self._items_size + size_diff < self.param.min_byte_num: raise NotEnoughData(self.param.min_byte_num) if self._items_size + size_diff > self.param.max_byte_num: raise TooMuchData(self.param.max_byte_num) self._items_size += size_diff @classmethod @abc.abstractmethod def get_param(cls): raise NotImplementedError() def __len__(self): return len(self._items) def __getitem__(self, index): return self._items[index] def __delitem__(self, index): self._update_items_size(del_item=self._items[index]) del self._items[index] def __setitem__(self, index, value): self._update_items_size(del_item=self._items[index], insert_item=value) self._items[index] = value def __str__(self): return str(self._items) def insert(self, index, value): self._update_items_size(insert_item=value) self._items.insert(index, value) def append(self, value): self.insert(len(self._items), value) def _asdict(self): return self._items def _as_markdown(self, level): return self._markdown_result(self._asdict(), level) class Vector(ArrayBase): @classmethod @abc.abstractmethod def get_param(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): vector_param = cls.get_param() parser = ParserBinary(parsable) parser.parse_numeric('item_byte_num', vector_param.item_num_size) item_byte_num = parser['item_byte_num'] item_num = int(item_byte_num / vector_param.item_size) parser.parse_numeric_array('items', item_num, vector_param.item_size, vector_param.numeric_class) return cls(parser['items']), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(len(self._items) * self.param.item_size, self.param.item_num_size) composer.compose_numeric_array(self._items, self.param.item_size) return composer.composed_bytes class VectorString(ArrayBase): @classmethod @abc.abstractmethod def get_param(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): vector_param = cls.get_param() header_parser = ParserBinary(parsable[:vector_param.item_num_size]) header_parser.parse_numeric('item_byte_num', vector_param.item_num_size) if header_parser['item_byte_num'] == 0: return cls([]), header_parser.parsed_length body_parser = ParserText( parsable[vector_param.item_num_size:header_parser['item_byte_num'] + vector_param.item_num_size] ) body_parser.parse_string_array( 'items', vector_param.separator, vector_param.item_class, vector_param.fallback_class, ) return cls(body_parser['items']), header_parser.parsed_length + body_parser.parsed_length def compose(self): vector_param = self.get_param() body_composer = ComposerText(vector_param.encoding) body_composer.compose_parsable_array(self._items, vector_param.separator, vector_param.fallback_class) header_composer = ComposerBinary() header_composer.compose_numeric(body_composer.composed_length, self.param.item_num_size) return header_composer.composed + body_composer.composed class VectorParsable(ArrayBase): @classmethod @abc.abstractmethod def get_param(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): vector_param = cls.get_param() parser = ParserBinary(parsable) parser.parse_numeric('item_byte_num', vector_param.item_num_size) parser.parse_parsable_array( 'items', items_size=parser['item_byte_num'], item_class=vector_param.item_class, fallback_class=vector_param.fallback_class ) return cls(parser['items']), parser.parsed_length def compose(self): body_composer = ComposerBinary() body_composer.compose_parsable_array(self._items) header_composer = ComposerBinary() header_composer.compose_numeric(body_composer.composed_length, self.param.item_num_size) return header_composer.composed_bytes + body_composer.composed_bytes class VectorEnumCodeNumeric(VectorParsable): @classmethod @abc.abstractmethod def get_param(cls): raise NotImplementedError() def compose(self): body_composer = ComposerBinary() for item in self: if isinstance(self.param.fallback_class, type) and isinstance(item, self.param.fallback_class): body_composer.compose_parsable(item) else: body_composer.compose_numeric_enum_coded(item) header_composer = ComposerBinary() header_composer.compose_numeric(body_composer.composed_length, self.param.item_num_size) return header_composer.composed_bytes + body_composer.composed_bytes class VectorEnumCodeString(VectorParsable): @classmethod @abc.abstractmethod def get_param(cls): raise NotImplementedError() def compose(self): body_composer = ComposerBinary() item_size = self.get_param().item_class.get_param().item_num_size for item in self: body_composer.compose_string_enum_coded(item, item_size) header_composer = ComposerBinary() header_composer.compose_numeric(body_composer.composed_length, self.param.item_num_size) return header_composer.composed_bytes + body_composer.composed_bytes class VectorParsableDerived(ArrayBase): @classmethod @abc.abstractmethod def get_param(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): vector_param = cls.get_param() parser = ParserBinary(parsable) parser.parse_numeric('item_byte_num', vector_param.item_num_size) parser.parse_parsable_derived_array( 'items', items_size=parser['item_byte_num'], item_base_class=vector_param.item_class, fallback_class=vector_param.fallback_class ) return cls(parser['items']), parser.parsed_length def compose(self): body_composer = ComposerBinary() body_composer.compose_parsable_array(self._items) header_composer = ComposerBinary() header_composer.compose_numeric(len(body_composer.composed_bytes), self.param.item_num_size) return header_composer.composed_bytes + body_composer.composed_bytes class Opaque(ArrayBase): @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('item_byte_num', cls.get_param().item_num_size) parser.parse_raw('items', parser['item_byte_num']) items = parser['items'] return cls([ord(items[i:i + 1]) for i in range(len(items))]), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(len(self._items), self.get_param().item_num_size) composer.compose_numeric_array(self._items, 1) return composer.composed_bytes @classmethod @abc.abstractmethod def get_param(cls): raise NotImplementedError() class NByteEnumParsable(ParsableBase): @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('code', cls.get_byte_num()) for enum_item in list(cls.get_enum_class()): if enum_item.value.code == parser['code']: return enum_item, cls.get_byte_num() raise InvalidValue(parser['code'], cls, 'code') @classmethod @abc.abstractmethod def get_byte_num(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def get_enum_class(cls): raise NotImplementedError() class OneByteEnumParsable(NByteEnumParsable): @classmethod def get_byte_num(cls): return 1 @classmethod @abc.abstractmethod def get_enum_class(cls): raise NotImplementedError() class TwoByteEnumParsable(NByteEnumParsable): @classmethod def get_byte_num(cls): return 2 @classmethod @abc.abstractmethod def get_enum_class(cls): raise NotImplementedError() class ThreeByteEnumParsable(NByteEnumParsable): @classmethod def get_byte_num(cls): return 3 @classmethod @abc.abstractmethod def get_enum_class(cls): raise NotImplementedError() class FourByteEnumParsable(NByteEnumParsable): @classmethod def get_byte_num(cls): return 4 @classmethod @abc.abstractmethod def get_enum_class(cls): raise NotImplementedError() class NByteEnumComposer(enum.Enum): def __repr__(self): return self.__class__.__name__ + '.' + self.name def compose(self): composer = ComposerBinary() composer.compose_numeric( self.value.code, # pylint: disable=no-member self.get_byte_num() ) return composer.composed @classmethod @abc.abstractmethod def get_byte_num(cls): raise NotImplementedError() class OneByteEnumComposer(NByteEnumComposer): @classmethod def get_byte_num(cls): return 1 class TwoByteEnumComposer(NByteEnumComposer): @classmethod def get_byte_num(cls): return 2 class ThreeByteEnumComposer(NByteEnumComposer): @classmethod def get_byte_num(cls): return 3 class FourByteEnumComposer(NByteEnumComposer): @classmethod def get_byte_num(cls): return 4 class StringEnumParsableBase(ParsableBaseNoABC): @classmethod @abc.abstractmethod def _code_eq(cls, item_code, parsed_code): raise NotImplementedError() @classmethod def _parse(cls, parsable): enum_items = [ enum_item for enum_item in cls # pylint: disable=not-an-iterable if len(enum_item.value.code) <= len(parsable) ] enum_items.sort(key=lambda color: len(color.value.code), reverse=True) try: code = bytes(parsable).decode('ascii') except UnicodeDecodeError as e: raise InvalidValue(parsable, cls) from e for enum_item in enum_items: if cls._code_eq(enum_item.value.code, code[:len(enum_item.value.code)]): return enum_item, len(enum_item.value.code) raise InvalidValue(parsable, cls, 'code') def compose(self): return self._asdict().encode('ascii') def _asdict(self): return self.value.code class StringEnumParsable(StringEnumParsableBase): @classmethod def _code_eq(cls, item_code, parsed_code): return item_code == parsed_code class StringEnumCaseInsensitiveParsable(StringEnumParsableBase): @classmethod def _code_eq(cls, item_code, parsed_code): return item_code.lower() == parsed_code.lower() class ProtocolVersionBase(Serializable, ParsableBase, metaclass=abc.ABCMeta): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @property @abc.abstractmethod def identifier(self): raise NotImplementedError() @abc.abstractmethod def __str__(self): raise NotImplementedError() def _asdict(self): return self.identifier def _as_markdown(self, level): return self._markdown_result(str(self), level) @attr.s class ProtocolVersionMajorMinorBase(ProtocolVersionBase): _SIZE = 2 major = attr.ib() minor = attr.ib() @classmethod def _parse_version_numbers(cls, parsable): if len(parsable) < cls._SIZE: raise NotEnoughData(bytes_needed=cls._SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric('major', 1) parser.parse_numeric('minor', 1) return parser @classmethod def _parse(cls, parsable): parser = cls._parse_version_numbers(parsable) return cls(**parser), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.major, 1) composer.compose_numeric(self.minor, 1) return composer.composed_bytes @property def identifier(self): return f'{self.major}_{self.minor}' def __str__(self): return f'{self.major}.{self.minor}' @attr.s class ListParamParsable: # pylint: disable=too-few-public-methods item_class = attr.ib(validator=attr.validators.instance_of(type)) fallback_class = attr.ib(validator=attr.validators.optional(attr.validators.instance_of(type))) separator_class = attr.ib(attr.validators.instance_of(ParsableBase)) min_byte_num = attr.ib(init=False, default=0) max_byte_num = attr.ib(init=False, default=2 ** 16) item_num_size = attr.ib(init=False, default=0) def get_item_size(self, item): # pylint: disable=no-self-use return len(item.compose()) class ListParsable(ArrayBase): @classmethod @abc.abstractmethod def get_param(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): list_param = cls.get_param() parser = ParserBinary(parsable) parser.parse_parsable_list( 'items', item_class=list_param.item_class, fallback_class=list_param.fallback_class, separator_class=list_param.separator_class ) return cls(parser['items']), parser.parsed_length def compose(self): composer = ComposerBinary() separator = bytearray(self.get_param().separator_class().compose()) composer.compose_parsable_array(self._items, separator) composer.compose_raw(separator) if self._items: composer.compose_raw(separator) return composer.composed_bytes class OpaqueEnumParsable(Vector): @classmethod def _parse(cls, parsable): opaque, parsed_length = super()._parse(parsable) code = b''.join([bytes((opaque_item,)) for opaque_item in opaque]).decode(cls.get_encoding()) try: parsed_object = next(iter([ enum_item for enum_item in cls.get_enum_class() if enum_item.value.code == code ])) except StopIteration as e: raise InvalidValue(code, cls) from e return parsed_object, parsed_length @classmethod @abc.abstractmethod def get_enum_class(cls): raise NotImplementedError() @classmethod def get_encoding(cls): return 'utf-8' class OpaqueEnumComposer(enum.Enum): def __repr__(self): return self.__class__.__name__ + '.' + self.name def compose(self): composer = ComposerBinary() value = self.value.code.encode(self.get_encoding()) # pylint: disable=no-member composer.compose_bytes(value, 1) return composer.composed_bytes @classmethod def get_encoding(cls): return 'utf-8' @attr.s class NumericRangeParsableBase(ParsableBase, Serializable): value = attr.ib(validator=attr.validators.instance_of(int)) @value.validator def _validator_variant(self, _, value): if value < self._get_value_min(): raise InvalidValue(value, type(self)) if value > self._get_value_max(): raise InvalidValue(value, type(self)) @classmethod @abc.abstractmethod def _get_value_min(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_value_max(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_value_length(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('value', cls._get_value_length()) return cls(**parser), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.value, self._get_value_length()) return composer.composed def __str__(self): return str(self.value) def _as_markdown(self, level): return self._markdown_result(str(self), level) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/common/classes.py000066400000000000000000000032471524413560000301140ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.parse import ParsableBase, ParserText, ComposerText class LanguageTag(ParsableBase): def __init__(self, primary_subtag, subsequent_subtags=()): self._primary_subtag = None self._subsequent_subtags = None self.primary_subtag = primary_subtag self.subsequent_subtags = subsequent_subtags @property def primary_subtag(self): return self._primary_subtag @primary_subtag.setter def primary_subtag(self, value): if not value or len(value) > 8 or not value.isalpha(): raise InvalidValue(value, LanguageTag, 'primary_subtag') self._primary_subtag = value @property def subsequent_subtags(self): return self._subsequent_subtags @subsequent_subtags.setter def subsequent_subtags(self, value): for subsequent_subtag in value: if not subsequent_subtag or len(subsequent_subtag) > 8 or not subsequent_subtag.isalnum(): raise InvalidValue(value, LanguageTag, 'subsequent_subtag') self._subsequent_subtags = value @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_array('tags', '-') return LanguageTag(parser['tags'][0], parser['tags'][1:]), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string(self.primary_subtag) if self.subsequent_subtags: composer.compose_separator('-') composer.compose_string_array(self.subsequent_subtags, '-') return composer.composed cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/common/exception.py000066400000000000000000000013051524413560000304460ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import typing import attr @attr.s class InvalidDataLength(Exception): bytes_needed: typing.Optional[int] = attr.ib( default=None, validator=attr.validators.optional(attr.validators.instance_of(int)) ) @attr.s class NotEnoughData(InvalidDataLength): def __str__(self): return f'not enough data received from target; missing_byte_count="{self.bytes_needed}"' @attr.s class TooMuchData(InvalidDataLength): def __str__(self): return f'too much data received from target; rest_byte_count="{self.bytes_needed}"' class InvalidType(Exception): def __str__(self): return 'invalid type value received from target' cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/common/field.py000066400000000000000000000700721524413560000275420ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import abc import collections import datetime import enum import json import attr import urllib3 from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.types import Base64Data, convert_base64_data, convert_value_to_object, convert_url from cryptoparser.common.base import Serializable from cryptoparser.common.exception import InvalidType, NotEnoughData from cryptoparser.common.parse import ParserText, ParsableBase, ParsableBaseNoABC, ComposerText class FieldParsableBase(ParsableBase): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod @abc.abstractmethod def get_separator(cls): raise NotImplementedError() @classmethod def _parse_name(cls, parsable): separator = cls.get_separator() parser = ParserText(parsable) parser.parse_string_until_separator_or_end('name', separator) return parser @classmethod def _compose_name(cls, name): composer = ComposerText() composer.compose_string(name) return composer @attr.s class NameValueVariantBase(FieldParsableBase): value = attr.ib() @value.validator def _x_validator(self, attribute, value): # pylint: disable=unused-argument value_class = self._get_value_class() if not isinstance(value, value_class): self.value = value_class(value) @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_value_class(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def get_separator(cls): raise NotImplementedError() @classmethod def _parse_name_and_separator(cls, parsable): separator = cls.get_separator() parser = cls._parse_name(parsable) parser.parse_separator(separator) if parser['name'].lower() != cls.get_canonical_name().lower(): raise InvalidType() return parser @attr.s class NameValuePair(FieldParsableBase): name = attr.ib(validator=attr.validators.instance_of(str)) value = attr.ib(validator=attr.validators.optional(attr.validators.instance_of(str)), default=None) quoted = attr.ib(validator=attr.validators.optional(attr.validators.instance_of(bool)), default=False) @classmethod def get_separator(cls): return '=' @classmethod def _parse(cls, parsable): value = None quoted = False parser = cls._parse_name(parsable) if parser.unparsed_length: parser.parse_separator(cls.get_separator()) parser.parse_string_by_length('value', min_length=0) value = parser['value'] if value and value[0] == '"': quoted = True value = value[1:] if value and value[-1:] == '"': value = value[:-1] return cls(parser['name'], value, quoted), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string(self.name) if self.value is not None: composer.compose_separator(self.get_separator()) if self.quoted: composer.compose_separator('"') composer.compose_string(self.value) if self.quoted: composer.compose_separator('"') return composer.composed @attr.s class NameValuePairList(ParsableBase, Serializable): value = attr.ib( default=collections.OrderedDict([]), validator=attr.validators.instance_of(collections.OrderedDict), ) @classmethod @abc.abstractmethod def get_separator(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_array( 'value', cls.get_separator(), item_class=NameValuePair, separator_spaces=' \t', skip_empty=True ) return cls( collections.OrderedDict([(component.name, component.value) for component in parser['value']]) ), parser.parsed_length def compose(self): composer = ComposerText() separator = self.get_separator() + ' ' for item_number, (name, value) in enumerate(self.value.items()): composer.compose_string(name) if value is not None: composer.compose_separator('=') composer.compose_string(value) if item_number + 1 < len(self.value): composer.compose_separator(separator) return composer.composed def _as_markdown(self, level): return self._markdown_result(self.value, level) class NameValuePairListCommaSeparated(NameValuePairList): @classmethod def get_separator(cls): return ',' class NameValuePairListSemicolonSeparated(NameValuePairList): @classmethod def get_separator(cls): return ';' def is_validator_optional(validator): return isinstance(validator, attr.validators._OptionalValidator) # pylint: disable=protected-access @attr.s class FieldValueComponentBase(ParsableBase, Serializable): @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod def convert(cls, value): if not isinstance(value, cls): value = cls(value) return value @classmethod def _check_name_insensitive(cls, name): if name.lower() != cls.get_canonical_name().lower(): raise InvalidType() @classmethod def _check_name(cls, name): if name != cls.get_canonical_name(): raise InvalidType() @attr.s class FieldValueComponentOption(FieldValueComponentBase): value = attr.ib(validator=attr.validators.instance_of(bool)) @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod def _check_name(cls, name): cls._check_name_insensitive(name) @classmethod def _parse(cls, parsable): parser = ParserText(parsable) try: canonical_name = cls.get_canonical_name() parser.parse_string_by_length('value', len(canonical_name), len(canonical_name)) cls._check_name(parser['value']) value = True except (InvalidValue, NotEnoughData): value = False return cls(value), parser.parsed_length def compose(self): composer = ComposerText() if self.value: composer.compose_string(self.get_canonical_name()) return composer.composed def _as_markdown(self, level): return self._markdown_result(self.value, level) @attr.s class FieldValueComponentKeyValueBase(FieldValueComponentBase): value = attr.ib() @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _parse_value(cls, parser): raise NotImplementedError() def _get_value_as_simple_type(self): # neccessary only because PY2 handles multiple inheritance differently than PY3 if isinstance(self.value, ParsableBaseNoABC): return self.value.compose().decode('ascii') return self.value def _get_value_as_str(self): return str(self._get_value_as_simple_type()) @classmethod def _parse(cls, parsable): parser = ParserText(parsable) name = cls.get_canonical_name() try: parser.parse_string_by_length('name', len(name), len(name)) except NotEnoughData as e: raise InvalidType from e cls._check_name(parser['name']) if cls.get_canonical_name(): parser.parse_separator('=') cls._parse_value(parser) parsed_value = parser['value'] if cls.get_canonical_name(): parsed_value = cls(parsed_value) return parsed_value, parser.parsed_length def compose(self): composer = ComposerText() if self.get_canonical_name(): composer.compose_string_array([self.get_canonical_name(), self._get_value_as_str()], '=') else: composer.compose_string(self._get_value_as_str()) return composer.composed def _as_markdown(self, level): return self._markdown_result(self.value, level) @attr.s class FieldValueComponentParsableBase(FieldValueComponentKeyValueBase): value = attr.ib() def __attrs_post_init__(self): value_class = self._get_value_class() if not isinstance(self.value, value_class): self.value = value_class(self.value) @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_value_class(cls): raise NotImplementedError() @classmethod def _parse_value(cls, parser): parser.parse_parsable('value', cls._get_value_class()) class FieldValueComponentParsable(FieldValueComponentParsableBase): @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_value_class(cls): raise NotImplementedError() class FieldValueComponentParsableOptional(FieldValueComponentParsableBase): def __attrs_post_init__(self): if self.value is not None: super().__attrs_post_init__() @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_value_class(cls): raise NotImplementedError() @attr.s class FieldValueComponentQuotedString(FieldValueComponentKeyValueBase): value = attr.ib(validator=attr.validators.optional(attr.validators.instance_of(str))) @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() def _get_value_as_str(self): return f'"{self.value}"' def _get_value_as_simple_type(self): return self.value @classmethod def _parse_value(cls, parser): parser.parse_separator('"', 0, None) parser.parse_string_until_separator_or_end('value', '"') parser.parse_separator('"', 0, None) @attr.s class FieldValueComponentDateTime(FieldValueComponentKeyValueBase): value = attr.ib(validator=attr.validators.instance_of(datetime.datetime)) @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod def _parse_value(cls, parser): parser.parse_date_time('value') def _get_value_as_simple_type(self): return self.value.strftime('%a, %d %b %Y %H:%M:%S GMT') @attr.s class FieldValueComponentTimeDelta(FieldValueComponentKeyValueBase): value = attr.ib( converter=convert_value_to_object(datetime.timedelta), validator=attr.validators.instance_of(datetime.timedelta) ) @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod def convert(cls, value): if isinstance(value, cls): return value if isinstance(value, datetime.timedelta): return cls(value) return cls(datetime.timedelta(seconds=value)) @classmethod def _parse_value(cls, parser): parser.parse_time_delta('value') def _get_value_as_simple_type(self): return int(self.value.total_seconds()) def _as_markdown(self, level): return self._markdown_result(str(self.value), level) @attr.s class FieldValueComponentStringBase64(FieldValueComponentQuotedString): value = attr.ib( converter=convert_base64_data(), validator=attr.validators.instance_of(Base64Data) ) @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod def _check_name(cls, name): cls._check_name_insensitive(name) @classmethod def _parse_value(cls, parser): parser.parse_string_by_length('value') @attr.s class FieldValueComponentBool(FieldValueComponentKeyValueBase): value = attr.ib(validator=attr.validators.instance_of(bool)) @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() def _get_value_as_str(self): return 'yes' if self.value else 'no' def _get_value_as_simple_type(self): return self.value @classmethod def _parse_value(cls, parser): parser.parse_bool('value') @attr.s class FieldValueComponentNumber(FieldValueComponentKeyValueBase): value = attr.ib(validator=attr.validators.instance_of(int)) @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod def _parse_value(cls, parser): parser.parse_numeric('value') @attr.s class FieldValueComponentFloat(FieldValueComponentNumber): value = attr.ib( converter=float, validator=attr.validators.instance_of(float) ) @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() def _get_value_as_simple_type(self): return self.value @classmethod def _parse_value(cls, parser): parser.parse_float('value') @attr.s class FieldValueComponentPercent(FieldValueComponentNumber): def __attrs_post_init__(self): if self.value < 0 or self.value > 100: raise InvalidValue(self.value, type(self), 'value') @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @attr.s class FieldValueComponentStringEnumParams: code = attr.ib(validator=attr.validators.instance_of(str)) @attr.s class FieldValueComponentStringEnum(FieldValueComponentKeyValueBase): value = attr.ib() @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_value_type(cls): raise NotImplementedError() @value.validator def _validator_value(self, _, value): if not isinstance(value, self._get_value_type()): raise InvalidValue(value, type(self), 'value') @classmethod def _parse_value(cls, parser): try: parser.parse_parsable('value', cls._get_value_type()) except InvalidValue as e: raise InvalidValue(e.value.decode('ascii'), cls, 'value') from e def _get_value_as_simple_type(self): return self.value.value.code @attr.s class FieldValueComponentStringEnumOption(FieldValueComponentStringEnum): @classmethod @abc.abstractmethod def _get_value_type(cls): raise NotImplementedError() @classmethod def get_canonical_name(cls): return '' @classmethod def _check_name(cls, name): pass @attr.s class FieldValueComponentString(FieldValueComponentKeyValueBase): value = attr.ib(validator=attr.validators.optional(attr.validators.instance_of(str))) @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod def _parse_value(cls, parser): parser.parse_string_by_length('value') @attr.s class FieldValueComponentUrl(FieldValueComponentKeyValueBase): value = attr.ib() @value.validator def _value_validate(self, _, value): self.value = convert_url()(value) if isinstance(self.value, urllib3.util.Url): return raise InvalidValue(self.value, type(self), 'value') @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod def _parse_value(cls, parser): parser.parse_string_by_length('value', item_class=convert_url()) def _get_value_as_simple_type(self): if self.value.scheme == 'mailto': value = 'mailto:' + self.value.path[1:] else: value = str(self.value) return value def _as_markdown(self, level): return self._markdown_result(self._get_value_as_simple_type(), level) class FieldValueBase(ParsableBase, Serializable): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod def _get_attr_to_validator_type_dict(cls, attr_fields_dict): attr_to_component_name_dict = [] for attribute in attr_fields_dict.values(): validator = attribute.validator if is_validator_optional(validator): validator = validator.validator attr_to_component_name_dict.append((attribute.name, validator.type)) return collections.OrderedDict(attr_to_component_name_dict) class FieldsJson(FieldValueBase): @classmethod def _parse(cls, parsable): try: raw_values = json.loads(parsable.decode('ascii'), object_pairs_hook=collections.OrderedDict) except ValueError as e: # json.decoder.JSONDecodeError is derived from ValueError raise InvalidValue(parsable.decode('ascii'), cls, 'value') from e attr_fields_dict = attr.fields_dict(cls) return cls(**{ attribute_name: raw_values[validator_class.get_canonical_name()] for attribute_name, validator_class in cls._get_attr_to_validator_type_dict(attr_fields_dict).items() if validator_class.get_canonical_name() in raw_values }), len(parsable) def compose(self): attr_fields_dict = attr.fields_dict(type(self)) return json.dumps(collections.OrderedDict([ ( validator_class.get_canonical_name(), getattr(self, attribute_name)._get_value_as_simple_type() # pylint: disable=protected-access ) for attribute_name, validator_class in self._get_attr_to_validator_type_dict(attr_fields_dict).items() if getattr(self, attribute_name) is not None ])).encode('ascii') class FieldValueMultiple(FieldValueBase): @classmethod @abc.abstractmethod def _get_header_value_list_class(cls): raise NotImplementedError() @classmethod def _parse_basic_params(cls, attr_to_component_name_dict, attr_fields_dict, components, params): for name, attribute in attr_fields_dict.items(): for component in components: try: attr_to_component_name_dict[name]._check_name(component) # pylint: disable=protected-access except InvalidType: pass else: components[attr_to_component_name_dict[name].get_canonical_name()] = components.pop(component) break else: if attribute.default == attr.NOTHING: raise InvalidValue(None, cls, name) component = attribute.default if attr_to_component_name_dict[name].get_canonical_name() in components: parsable = components.pop(attr_to_component_name_dict[name].get_canonical_name()) if parsable is None: # value is None in case of optional values parsable = component else: parsable = '='.join([attr_to_component_name_dict[name].get_canonical_name(), parsable]) params[name] = attr_to_component_name_dict[name].parse_exact_size(parsable.encode('ascii')) else: params[name] = attribute.default @classmethod def _parse_extensions(cls, attr_to_component_name_dict, extension, components, params): if extension and components: name, _ = extension params[name] = attr_to_component_name_dict[name](components) @classmethod def _parse(cls, parsable): params = {} extension = None attr_fields_dict_basic = {} attr_fields_dict = attr.fields_dict(cls) for name, attribute in attr_fields_dict.items(): if not attribute.metadata.get('extension', False): attr_fields_dict_basic[name] = attribute elif extension is None: extension = (name, attribute) else: raise NotImplementedError() attr_to_component_name_dict = cls._get_attr_to_validator_type_dict(attr_fields_dict) components = cls._get_header_value_list_class().parse_exact_size(parsable).value cls._parse_basic_params(attr_to_component_name_dict, attr_fields_dict_basic, components, params) cls._parse_extensions(attr_to_component_name_dict, extension, components, params) return cls(**params), len(parsable) def compose(self): composer = ComposerText() cls = type(self) attr_fields_dict = attr.fields_dict(cls) components = [] for name, attribute in attr_fields_dict.items(): field_value = getattr(self, name) validator = attribute.validator if is_validator_optional(validator): if field_value is None: continue validator = validator.validator value = field_value.value if issubclass(validator.type, FieldValueComponentOption): if value is False: continue components.append(getattr(self, name)) separator = self._get_header_value_list_class().get_separator() + ' ' composer.compose_string_array(components, separator) return composer.composed class FieldsCommaSeparated(FieldValueMultiple): @classmethod def _get_header_value_list_class(cls): return NameValuePairListCommaSeparated class FieldsSemicolonSeparated(FieldValueMultiple): @classmethod def _get_header_value_list_class(cls): return NameValuePairListSemicolonSeparated class MimeTypeRegistry(enum.Enum): APPLICATION = 'application' AUDIO = 'audio' FONT = 'font' EXAMPLE = 'example' IMAGE = 'image' MESSAGE = 'message' MODEL = 'model' MULTIPART = 'multipart' TEXT = 'text' VIDEO = 'video' @attr.s class FieldValueMimeType(FieldValueComponentBase): type = attr.ib( validator=attr.validators.instance_of(str), default=None, ) registry = attr.ib( validator=attr.validators.optional(attr.validators.instance_of(MimeTypeRegistry)), default=None, ) def __str__(self): return f'{self.registry.value}/{self.type}' @property def value(self): return self @classmethod def get_canonical_name(cls): return '' @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_until_separator('registry', '/', item_class=MimeTypeRegistry) parser.parse_separator('/') parser.parse_string_by_length('type', parser.unparsed_length) return FieldValueMimeType(**parser), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string(str(self)) return composer.composed @classmethod def _check_name(cls, name): pass @attr.s class FieldValueSingleBase(FieldValueBase, Serializable): value = attr.ib() @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_value_type(cls): raise NotImplementedError() @classmethod def convert(cls, value): if not isinstance(value, cls._get_value_type()): return value return cls(value) @value.validator def _value_validate(self, _, value): value_type = self._get_value_type() if not isinstance(value, value_type): raise InvalidValue(value, value_type, 'value') def _as_markdown(self, level): return self._markdown_result(self.value, level) class FieldValueSingleSimpleBase(FieldValueSingleBase): @classmethod @abc.abstractmethod def _value_from_str(cls, value): raise NotImplementedError() class FieldValueSingle(FieldValueSingleSimpleBase): @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_by_length('value') value = cls._value_from_str(parser['value']) return cls(value), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string(self.value) return composer.composed @classmethod @abc.abstractmethod def _get_value_type(cls): raise NotImplementedError() @attr.s class FieldValueStringEnumParams(Serializable): code = attr.ib(validator=attr.validators.instance_of(str)) human_readable_name = attr.ib( default=None, validator=attr.validators.optional(attr.validators.instance_of(str)) ) def _as_markdown(self, level): if self.human_readable_name: return self._markdown_result(self.human_readable_name, level) return False, self.code.replace('_', ' ') class FieldValueString(FieldValueSingle): @classmethod def _get_value_type(cls): return str @classmethod def _value_from_str(cls, value): return str(value) class FieldValueSingleComplexBase(FieldValueSingleBase): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() class FieldValueDateTime(FieldValueSingleComplexBase): @classmethod def _get_value_type(cls): return datetime.datetime @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_date_time('value') return cls(parser['value']), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_date_time(self.value, '%a, %d %b %Y %H:%M:%S GMT') return composer.composed def _as_markdown(self, level): return self._markdown_result(self.value, level) class FieldValueStringBySeparatorBase(FieldValueSingleComplexBase): @classmethod @abc.abstractmethod def _get_separators(cls): raise NotImplementedError() @classmethod def _get_value_type(cls): return str @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_until_separator_or_end('value', cls._get_separators()) return cls(parser['value']), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string(self.value) return composer.composed class FieldValueStringEnum(FieldValueSingleComplexBase): @classmethod @abc.abstractmethod def _get_value_type(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): try: value = cls._get_value_type().parse_exact_size(parsable) except InvalidValue as e: raise InvalidValue(parsable.decode('ascii'), cls, 'value') from e return cls(value), len(parsable) def compose(self): composer = ComposerText() composer.compose_string(self.value.value.code) return composer.composed class FieldValueTimeDelta(FieldValueSingleComplexBase): @classmethod def _get_value_type(cls): return datetime.timedelta @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_time_delta('value') return cls(parser['value']), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_time_delta(self.value) return composer.composed cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/common/parse.py000066400000000000000000001073321524413560000275710ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import abc import collections.abc import datetime import enum import struct import time import typing import attr from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.types import CryptoDataEnumBase, CryptoDataEnumCodedBase from cryptoparser.common.exception import InvalidType, NotEnoughData, TooMuchData import cryptoparser.common.utils class ParsableBaseNoABC: @classmethod def parse_mutable(cls, parsable): parsed_object, parsed_length = cls._parse(parsable) del parsable[:parsed_length] return parsed_object @classmethod def parse_immutable(cls, parsable): parsed_object, parsed_length = cls._parse(parsable) return parsed_object, parsed_length @classmethod def parse_exact_size(cls, parsable): parsed_object, parsed_length = cls._parse(parsable) if len(parsable) > parsed_length: raise TooMuchData(parsed_length) return parsed_object @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() class ParsableBase(ParsableBaseNoABC, metaclass=abc.ABCMeta): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() class ParserCRLF(ParsableBase): @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string('crlf', '\r\n') return ParserCRLF(), 2 def compose(self): return b'\r\n' _SIZE_TO_FORMAT = { 1: 'B', 2: 'H', 3: 'I', 4: 'I', 8: 'Q', } class ByteOrder(enum.Enum): NATIVE = '=' LITTLE_ENDIAN = '<' BIG_ENDIAN = '>' NETWORK = '!' @attr.s class ParserBase(collections.abc.Mapping): _parsable: typing.Union[bytes, bytearray] = attr.ib( converter=bytes, validator=attr.validators.instance_of((bytes, bytearray)) ) _parsed_length: int = attr.ib(init=False, default=0) _parsed_values: dict[str, typing.Any] = attr.ib(init=False, default=None) def __attrs_post_init__(self): if self._parsed_values is None: self._parsed_values = {} def __len__(self): return len(self._parsed_values) def __iter__(self): return iter(self._parsed_values) def __getitem__(self, key): return self._parsed_values[key] def __delitem__(self, key): del self._parsed_values[key] @property def parsed_length(self): return self._parsed_length @property def unparsed(self): return self._parsable[self._parsed_length:] @property def unparsed_length(self): return len(self._parsable) - self._parsed_length @abc.abstractmethod def _parse_numeric(self, name, converter, item_size): raise NotImplementedError() def parse_parsable(self, name, parsable_class, item_size=None): if item_size is None: parsed_object, parsed_length = parsable_class.parse_immutable( self._parsable[self._parsed_length:] ) else: parsable_length, _ = self._parse_numeric(name, int, item_size) parsable_length = parsable_length[0] parsed_object = parsable_class.parse_exact_size( self._parsable[self._parsed_length + item_size:self._parsed_length + parsable_length + item_size] ) parsed_length = item_size + parsable_length self._parsed_length += parsed_length self._parsed_values[name] = parsed_object def _parse_string_by_length( self, name, item_min_length, item_max_length, encoding, converter ): # pylint: disable=too-many-arguments,too-many-positional-arguments if item_min_length > self.unparsed_length: raise NotEnoughData(item_min_length - self.unparsed_length) if item_max_length is None: parsable_length = len(self._parsable) - self.parsed_length else: parsable_length = min(item_max_length, self.unparsed_length) value = self._parsable[self._parsed_length:self._parsed_length + parsable_length] try: value = value.decode(encoding) if converter is not str: value = converter(value) self._parsed_values[name] = value except UnicodeError as e: raise InvalidValue(value, converter, name) from e except ValueError as e: raise InvalidValue(value, converter, name) from e return value, parsable_length class ParserText(ParserBase): def __init__(self, parsable, encoding='ascii'): super().__init__(parsable) self._encoding = encoding def _check_separators( # pylint: disable=too-many-arguments,too-many-positional-arguments self, name, count_offset, separators, min_count, max_count ): separators = separators.encode(self._encoding) count = 0 actual_offset = count_offset while actual_offset < len(self._parsable) and self._parsable[actual_offset:actual_offset + 1] in separators: actual_offset += 1 count += 1 if max_count is not None and count > max_count: raise InvalidValue(self._parsable[count_offset:], type(self), name) if min_count is not None and count < min_count: raise InvalidValue(self._parsable[count_offset:], type(self), name) return actual_offset - count_offset def parse_separator(self, separator, min_length=1, max_length=None): self._parsed_length += self._check_separators( 'separator', self._parsed_length, separator, min_length, max_length ) def _parse_numeric_array( # pylint: disable=too-many-arguments,too-many-positional-arguments self, name, item_num, separator, converter, is_floating): value = [] floating_point_found = False last_item_offset = self._parsed_length item_offset = self._parsed_length while True: while item_offset < len(self._parsable) and self._parsable[item_offset:item_offset + 1].isdigit(): item_offset += 1 if item_offset == last_item_offset: raise InvalidValue(self._parsable[self._parsed_length:], type(self), name) if (is_floating and not floating_point_found and item_offset < len(self._parsable) and self._parsable[item_offset] == ord('.')): item_offset += 1 floating_point_found = True continue value.append(converter(self._parsable[last_item_offset:item_offset])) if item_offset == len(self._parsable) or (item_num is not None and len(value) == item_num): break if separator: try: item_offset += self._check_separators(name, item_offset, separator, 1, 1) except InvalidValue as e: raise InvalidValue(self._parsable[self._parsed_length:item_offset], type(self), name) from e last_item_offset = item_offset floating_point_found = False return value, item_offset - self._parsed_length def _parse_numeric(self, name, converter, item_size): raise NotImplementedError() def parse_numeric(self, name, converter=int): value, parsed_length = self._parse_numeric_array(name, 1, None, converter, False) self._parsed_values[name] = value[0] self._parsed_length += parsed_length def parse_float(self, name, converter=float): value, parsed_length = self._parse_numeric_array(name, 1, None, converter, True) self._parsed_values[name] = value[0] self._parsed_length += parsed_length def parse_numeric_array(self, name, item_num, separator, converter=int): value, parsed_length = self._parse_numeric_array(name, item_num, separator, converter, False) self._parsed_values[name] = value self._parsed_length += parsed_length def parse_bool(self, name): for string_value, bool_value in (('yes', True), ('no', False)): try: self.parse_string(name, string_value) except InvalidValue: pass else: self._parsed_values[name] = bool_value break else: raise InvalidValue(self._parsable[self._parsed_length:], type(self), name) def parse_string(self, name, value): min_length = len(value) max_length = min_length try: actual_value, parsed_length = self._parse_string_by_length( name, min_length, max_length, self._encoding, str ) except NotEnoughData as e: assert e.bytes_needed is not None raise InvalidValue(self._parsable[self._parsed_length:min_length - e.bytes_needed], type(self), name) from e if value != actual_value: raise InvalidValue(self._parsable[self._parsed_length:self._parsed_length + max_length], type(self), name) self._parsed_values[name] = value self._parsed_length += parsed_length def parse_string_by_length(self, name, min_length=1, max_length=None, item_class=str): value, parsed_length = self._parse_string_by_length(name, min_length, max_length, self._encoding, item_class) self._parsed_values[name] = value self._parsed_length += parsed_length def _apply_item_class( # pylint: disable=too-many-arguments,too-many-positional-arguments self, name, item_offset, item_end, separator, item_class, fallback_class, may_end): try: if not isinstance(item_class, type): item = item_class(self._parsable[item_offset:item_end].decode(self._encoding)) elif issubclass(item_class, CryptoDataEnumCodedBase): item = item_class.from_code(self._parsable[item_offset:item_end].decode(self._encoding)) elif issubclass(item_class, ParsableBaseNoABC): item, parsed_length = item_class.parse_immutable(self._parsable[item_offset:item_end]) item_end = item_offset + parsed_length elif issubclass(item_class, str): item = self._parsable[item_offset:item_end].decode(self._encoding) else: item = item_class(self._parsable[item_offset:item_end].decode(self._encoding)) except (InvalidValue, ValueError, UnicodeError) as e: if fallback_class is not None: parsed_value, parsed_length = self._parse_string_until_separator( name, item_offset, separator, fallback_class, None, may_end ) item_offset += parsed_length return parsed_value raise InvalidValue(self._parsable[item_offset:], type(self), name) from e return item def _parse_string_until_separator( # pylint: disable=too-many-arguments,too-many-positional-arguments self, name, item_offset, separators, item_class, fallback_class, may_end=False, separator_spaces='' ): item_end = None byte_separators = [separator.encode(self._encoding) for separator in separators] for separator_end in range(item_offset, len(self._parsable) + 1): for separator in byte_separators: if self._parsable[item_offset:separator_end].endswith(separator): item_end = separator_end - len(separator) break if item_end is not None: break else: if not may_end: raise InvalidValue(self._parsable[item_offset:], type(item_class), name) item_end = len(self._parsable) separator_space_count = 0 byte_separator_spaces = separator_spaces.encode(self._encoding) while (item_end > item_offset and self._parsable[ item_end - separator_space_count - 1: item_end - separator_space_count ] in byte_separator_spaces): separator_space_count += 1 item = self._apply_item_class( name, item_offset, item_end - separator_space_count, separators, item_class, fallback_class, may_end ) return item, item_end - item_offset - separator_space_count def parse_string_until_separator(self, name, separators, item_class=str, fallback_class=None): parsed_value, parsed_length = self._parse_string_until_separator( name, self._parsed_length, separators, item_class, fallback_class, False ) self._parsed_values[name] = parsed_value self._parsed_length += parsed_length def parse_string_until_separator_or_end(self, name, separators, item_class=str, fallback_class=None): parsed_value, parsed_length = self._parse_string_until_separator( name, self._parsed_length, separators, item_class, fallback_class, True ) self._parsed_values[name] = parsed_value self._parsed_length += parsed_length def _parse_string_array( self, name, separator, max_item_num=None, item_class=str, fallback_class=None, separator_spaces='', skip_empty=False ): # pylint: disable=too-many-arguments,too-many-positional-arguments value = [] item_offset = self._parsed_length max_separator_count = None if skip_empty else 1 if separator_spaces: item_offset += self._check_separators('separator', item_offset, separator_spaces, None, None) while True: parsed_value, parsed_length = self._parse_string_until_separator( name, item_offset, separator, str, None, True, separator_spaces ) if parsed_length: if isinstance(item_class, type) and issubclass(item_class, ParsableBase): parsed_value = item_class.parse_exact_size(parsed_value.encode(self._encoding)) else: parsed_value = self._apply_item_class( name, item_offset, item_offset + parsed_length, separator, item_class, fallback_class, True ) value.append(parsed_value) item_offset += parsed_length elif not skip_empty: raise InvalidValue(self._parsable[item_offset:], type(self), name) if separator_spaces: item_offset += self._check_separators('separator', item_offset, separator_spaces, None, None) if item_offset == len(self._parsable): break item_offset += self._check_separators(name, item_offset, separator, 1, max_separator_count) if separator_spaces: item_offset += self._check_separators('separator', item_offset, separator_spaces, None, None) if item_offset == len(self._parsable): break if max_item_num is not None and len(value) == max_item_num: break self._parsed_values[name] = value self._parsed_length = item_offset def parse_string_array( self, name, separator, item_class=str, fallback_class=None, separator_spaces='', skip_empty=False, max_item_num=None, ): # pylint: disable=too-many-arguments,too-many-positional-arguments self._parse_string_array( name, separator, max_item_num=max_item_num, item_class=item_class, fallback_class=fallback_class, separator_spaces=separator_spaces, skip_empty=skip_empty ) def parse_date_time(self, name): try: value = self._parsable[self._parsed_length:] value_str = value.decode(self._encoding) date_time = datetime.datetime.strptime(value_str, '%a, %d %b %Y %H:%M:%S %z') except ValueError: try: date_time = datetime.datetime.strptime(value_str, '%a, %d %b %Y %H:%M:%S GMT') date_time = date_time.replace(tzinfo=datetime.timezone.utc) except ValueError as e: raise InvalidValue(value, type(self), 'value') from e self._parsed_values[name] = date_time self._parsed_length = len(self._parsable) def parse_time_delta(self, name): value, parsed_length = self._parse_numeric_array(name, 1, None, int, False) try: time_delta = datetime.timedelta(seconds=value[0]) except OverflowError as e: raise InvalidValue(value[0], type(self), 'value') from e self._parsed_values[name] = time_delta self._parsed_length += parsed_length @attr.s class ParserBinary(ParserBase): byte_order: ByteOrder = attr.ib(default=ByteOrder.NETWORK, validator=attr.validators.in_(ByteOrder)) def parse_timestamp(self, name, milliseconds=False, item_size=8): value, parsed_length = self._parse_numeric_array(name, 1, item_size, int) self._parsed_length += parsed_length if value[0] == (2 ** (8 * item_size) - 1): self._parsed_values[name] = None else: value = value[0] if milliseconds: millis = value % 1000 value //= 1000 value = datetime.datetime.fromtimestamp(0x00000000ffffffff & value, datetime.timezone.utc) if milliseconds: value += datetime.timedelta(milliseconds=millis) self._parsed_values[name] = value def _parse_numeric_array(self, name, item_num, item_size, item_numeric_class): if self._parsed_length + (item_num * item_size) > len(self._parsable): raise NotEnoughData(bytes_needed=(item_num * item_size) - self.unparsed_length) if item_size in _SIZE_TO_FORMAT: value = [] for item_offset in range(self._parsed_length, self._parsed_length + (item_num * item_size), item_size): item_bytes = self._parsable[item_offset:item_offset + item_size] if item_size == 3: if self.byte_order in [ByteOrder.BIG_ENDIAN, ByteOrder.NETWORK]: item_bytes = b'\x00' + item_bytes else: item_bytes = item_bytes + b'\x00' item = struct.unpack( self.byte_order.value + _SIZE_TO_FORMAT[item_size], item_bytes )[0] try: value.append(item_numeric_class(item)) except ValueError as e: raise InvalidValue(item, item_numeric_class, name) from e else: raise NotImplementedError() return value, item_num * item_size def _parse_numeric(self, name, converter, item_size): return self._parse_numeric_array(name, 1, item_size, converter) def parse_numeric(self, name, size, converter=int): value, parsed_length = self._parse_numeric_array(name, 1, size, converter) self._parsed_length += parsed_length self._parsed_values[name] = value[0] def parse_numeric_enum_coded(self, name, enum_class): param = list(enum_class)[0].value value, parsed_length = self._parse_numeric_array(name, 1, param.get_code_size(), enum_class.from_code) self._parsed_length += parsed_length self._parsed_values[name] = value[0] def parse_numeric_array(self, name, item_num, item_size, converter=int): value, parsed_length = self._parse_numeric_array(name, item_num, item_size, converter) self._parsed_length += parsed_length self._parsed_values[name] = value def parse_numeric_flags(self, name, size, flags_class, shift_left=0): value, parsed_length = self._parse_numeric_array(name, 1, size, int) value = { flags_class(flag & (value[0] << shift_left)) for flag in flags_class if flag & (value[0] << shift_left) } self._parsed_length += parsed_length self._parsed_values[name] = value def _parse_mpint(self, mpint_length, mpint_offset, negative): if mpint_length % 4: pad_byte = bytes((0xff,)) if negative else bytes((0x00,)) pad_bytes = (4 - (mpint_length % 4)) * pad_byte else: pad_bytes = b'' parsable = pad_bytes + self._parsable[self._parsed_length + mpint_offset:] parser = ParserBinary(parsable) parser.parse_numeric_array('mpint_part_array', (mpint_length + len(pad_bytes)) // 4, 4, int) value = 0 for mpint_part in parser['mpint_part_array']: value = (value << 32) + mpint_part if negative: complement = 1 << (8 * (mpint_length + len(pad_bytes))) value -= complement return value def parse_mpint(self, name, mpint_length): value = self._parse_mpint(mpint_length, 0, False) self._parsed_values[name] = value self._parsed_length += mpint_length def parse_ssh_mpint(self, name): if self.unparsed_length < 4: raise NotEnoughData(bytes_needed=4 - self.unparsed_length) mpint_length, parsed_length = self._parse_numeric_array(name, 1, 4, int) mpint_length = mpint_length[0] negative = (mpint_length and self._parsable[self._parsed_length + 4] >= 0x80) value = self._parse_mpint(mpint_length, 4, negative) self._parsed_values[name] = value self._parsed_length += parsed_length + mpint_length def _parse_bytes(self, size): if self.unparsed_length < size: raise NotEnoughData(bytes_needed=size - self.unparsed_length) return self._parsable[self._parsed_length: self._parsed_length + size] def parse_bytes(self, name, size, converter=bytearray): value, parsed_length = self._parse_numeric_array(name, 1, size, int) value = value[0] self._parsed_length += parsed_length try: parsed_bytes = self._parse_bytes(value) except NotEnoughData: self._parsed_length -= parsed_length raise try: self._parsed_values[name] = converter(parsed_bytes) except ValueError as e: raise InvalidValue(value, converter, name) from e self._parsed_length += len(parsed_bytes) def parse_raw(self, name, size, converter=bytearray): parsed_bytes = self._parse_bytes(size) try: self._parsed_values[name] = converter(parsed_bytes) except ValueError as e: raise InvalidValue(parsed_bytes, converter, name) from e self._parsed_length += size def parse_string(self, name, item_size, encoding, converter=str): value, parsed_length = self._parse_numeric_array(name, 1, item_size, int) value = value[0] self._parsed_length += parsed_length try: value, parsed_length = self._parse_string_by_length(name, value, value, encoding, converter) self._parsed_length += parsed_length self._parsed_values[name] = value except InvalidValue as e: raise e def parse_string_null_terminated(self, name, encoding, converter=str): try: length = next(iter([ i for i, value in enumerate(self._parsable[self._parsed_length:]) if value == 0 ])) except StopIteration as e: raise InvalidValue(self._parsable[self._parsed_length:], str, name) from e value, parsed_length = self._parse_string_by_length(name, length, length, encoding, converter) self._parsed_length += parsed_length + 1 self._parsed_values[name] = value def _parse_parsable_derived_array(self, items_size, item_classes, fallback_class=None): if items_size > self.unparsed_length: raise NotEnoughData(bytes_needed=items_size - self.unparsed_length) items = [] unparsed_bytes = self._parsable[self._parsed_length:self._parsed_length + items_size] while unparsed_bytes: for item_class in item_classes: try: item, parsed_length = item_class.parse_immutable(unparsed_bytes) break except InvalidValue: pass else: if fallback_class is not None: item, parsed_length = fallback_class.parse_immutable(unparsed_bytes) else: raise ValueError(unparsed_bytes) unparsed_bytes = unparsed_bytes[parsed_length:] items.append(item) return items, items_size def parse_parsable_array(self, name, items_size, item_class, fallback_class=None): try: items, items_size = self._parse_parsable_derived_array(items_size, [item_class, ], fallback_class) except NotEnoughData as e: raise e except ValueError as e: raise InvalidValue(e.args[0], item_class, name) from e self._parsed_values[name] = items self._parsed_length += items_size def parse_parsable_derived_array(self, name, items_size, item_base_class, fallback_class=None): item_classes = cryptoparser.common.utils.get_leaf_classes(item_base_class) try: items, items_size = self._parse_parsable_derived_array(items_size, item_classes, fallback_class) except NotEnoughData as e: raise e except ValueError as e: raise InvalidValue(e.args[0], item_base_class, name) from e self._parsed_values[name] = items self._parsed_length += items_size def _remove_trailing_separator(self, items, name): if not isinstance(items[-1], ParserCRLF): raise InvalidValue(self.unparsed, type(self), name) del items[-1] def parse_parsable_list(self, name, item_class, fallback_class=None, separator_class=ParserCRLF): try: items, items_size = self._parse_parsable_derived_array( self.unparsed_length, [separator_class, item_class], fallback_class ) except ValueError as e: raise InvalidValue(e.args[0], item_class, name) from e if not items: raise NotEnoughData(2) self._remove_trailing_separator(items, name) real_items = [] while items: if len(items) % 2 == 0: self._remove_trailing_separator(items, name) else: if isinstance(items[-1], separator_class): raise InvalidValue(self.unparsed, type(self), name) real_items.insert(0, items[-1]) del items[-1] self._parsed_values[name] = real_items self._parsed_length += items_size def parse_variant(self, name, variant): parsed_object, value_length = variant.parse(self._parsable[self._parsed_length:]) self._parsed_values[name] = parsed_object self._parsed_length += value_length @attr.s class ComposerBase: _composed = attr.ib(init=False, default=b'') @property def composed(self): return self._composed @property def composed_length(self): return len(self._composed) def _compose_string_array(self, values, encoding, separator): separator = bytearray(separator.encode(encoding)) composed_str = bytearray() for value in values: try: if isinstance(value, ParsableBaseNoABC): composed_str += value.compose() else: composed_str += str(value).encode(encoding) except UnicodeError as e: raise InvalidValue(value, type(self)) from e composed_str += separator self._composed += composed_str[:len(composed_str) - len(separator)] class ComposerText(ComposerBase): def __init__(self, encoding='ascii'): super().__init__() self._encoding = encoding def _compose_numeric_array(self, values, separator): composed_str = '' for value in values: composed_str += f'{value:d}{separator}' self._composed += composed_str[:len(composed_str) - len(separator)].encode(self._encoding) def compose_numeric(self, value): self._compose_numeric_array([value, ], separator='') def compose_numeric_array(self, values, separator): self._compose_numeric_array(values, separator) def compose_bool(self, value): self.compose_string('yes' if value else 'no') def compose_string(self, value): self._compose_string_array([value, ], encoding=self._encoding, separator='') def compose_string_array(self, value, separator=','): self._compose_string_array(value, encoding=self._encoding, separator=separator) @staticmethod def _compose_parsable(value): return value.compose() def compose_parsable(self, value): self._composed += self._compose_parsable(value) def compose_parsable_array(self, values, separator=',', fallback_class=None): separator = separator.encode(self._encoding) composed_items = [] for item in values: if isinstance(item, (ComposerBase, ParsableBase, ParsableBaseNoABC)): composed_item = self._compose_parsable(item) elif isinstance(item, CryptoDataEnumBase): composed_item = item.value.code.encode(self._encoding) elif fallback_class is not None and isinstance(item, fallback_class): composed_item = item.encode(self._encoding) else: raise InvalidType() composed_items.append(composed_item) self._composed += bytearray(separator).join(composed_items) def compose_separator(self, value): self.compose_string(value) def compose_date_time(self, value, fmt): self.compose_string(value.strftime(fmt)) def compose_time_delta(self, value): self.compose_numeric(int(value.total_seconds())) @attr.s class ComposerBinary(ComposerBase): byte_order: ByteOrder = attr.ib(default=ByteOrder.NETWORK, validator=attr.validators.in_(ByteOrder)) def compose_timestamp(self, value, milliseconds=False, item_size=8): if value is None: timestamp = 0xffffffffffffffff else: timestamp = int(time.mktime(value.timetuple())) - time.timezone if milliseconds: timestamp *= 1000 timestamp += value.microsecond // 1000 return self._compose_numeric_array([timestamp, ], item_size) def _compose_numeric_array(self, values, item_size): composed_bytes = bytearray() for value in values: try: packed_bytes = struct.pack( self.byte_order.value + _SIZE_TO_FORMAT[item_size], value ) if item_size == 3: if self.byte_order in [ByteOrder.BIG_ENDIAN, ByteOrder.NETWORK]: composed_bytes += packed_bytes[1:] else: composed_bytes += packed_bytes[:3] else: composed_bytes += packed_bytes except struct.error as e: raise InvalidValue(value, int) from e self._composed += composed_bytes def compose_numeric(self, value, size): self._compose_numeric_array([value, ], size) def compose_numeric_enum_coded(self, value): self.compose_numeric_array_enum_coded([value, ]) def compose_numeric_array(self, values, item_size): self._compose_numeric_array(values, item_size) def compose_numeric_array_enum_coded(self, values): if not values: return self._compose_numeric_array(map(lambda value: value.value.code, values), values[0].value.get_code_size()) def compose_numeric_flags(self, values, item_size, shift_right=0): flag = 0 for value in values: flag |= value >> shift_right self._compose_numeric_array([flag, ], item_size) @staticmethod def _compose_mpint(value, length, byte_order): negative = value < 0 if negative: positive_value = (1 << ((value.bit_length() // 8 * 8) + 8)) + value else: positive_value = value value_numeric_array = [] for mpint_offset in range(0, length * 32, 32): value_numeric_array.append(positive_value >> mpint_offset & 0xffffffff) composer = ComposerBinary(byte_order=byte_order) composer.compose_numeric_array(reversed(value_numeric_array), 4) mpint_bytes = composer.composed_bytes.lstrip(b'\x00') if byte_order in [ByteOrder.LITTLE_ENDIAN, ByteOrder.NATIVE]: mpint_bytes = mpint_bytes.rstrip(b'\x00') return mpint_bytes def compose_mpint(self, value, length): mpint_bytes = self._compose_mpint(value, length, self.byte_order) if length < len(mpint_bytes): raise InvalidValue(length, type(self), 'mpint_length') pad_byte = b'\xff' if value < 0 else b'\x00' if self.byte_order in [ByteOrder.BIG_ENDIAN, ByteOrder.NETWORK]: self.compose_raw((length - len(mpint_bytes)) * pad_byte) self.compose_raw(mpint_bytes) else: self.compose_raw(mpint_bytes) self.compose_raw((length - len(mpint_bytes)) * pad_byte) def compose_ssh_mpint(self, value): negative = value < 0 length = value.bit_length() // 32 if value.bit_length() % 32: length += 1 mpint_bytes = self._compose_mpint(value, length, self.byte_order) if mpint_bytes and bool(mpint_bytes[0] & 0x80) != negative: pad_byte = b'\xff' if negative else b'\x00' else: pad_byte = b'' self.compose_numeric(len(pad_byte) + len(mpint_bytes), 4) self.compose_raw(pad_byte) self.compose_raw(mpint_bytes) def compose_parsable(self, value, item_size=None): composed = value.compose() if item_size is not None: self.compose_numeric(len(composed), item_size) self._composed += composed def compose_parsable_array(self, values, separator=b''): self._composed += separator.join(map(lambda item: item.compose(), values)) def compose_bytes(self, value, item_size, converter=bytearray): value_bytes = converter(value) self._compose_numeric_array([len(value_bytes), ], item_size) self.compose_raw(value_bytes) def compose_raw(self, value): self._composed += value def compose_string(self, value, encoding, item_size): try: value = value.encode(encoding) except UnicodeError as e: raise InvalidValue(value, type(self)) from e self.compose_bytes(value, item_size) def compose_string_null_terminated(self, value, encoding): try: value = value.encode(encoding) except UnicodeError as e: raise InvalidValue(value, type(self)) from e self.compose_raw(value) self.compose_raw(b'\x00') def compose_string_enum_coded(self, value, item_size): self.compose_string(value.value.code, 'ascii', item_size) @property def composed_bytes(self): return bytearray(self._composed) @property def composed_length(self): return len(self._composed) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/common/utils.py000066400000000000000000000020051524413560000276060ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import binascii import inspect def get_leaf_classes(base_class): def _get_leaf_classes(base_class): subclasses = [] if base_class.__subclasses__(): for subclass in base_class.__subclasses__(): subclasses += _get_leaf_classes(subclass) else: if not inspect.isabstract(base_class): return [base_class, ] return subclasses return _get_leaf_classes(base_class) def bytes_to_hex_string(byte_array, separator='', lowercase=False): if lowercase: format_str = '{:02x}' else: format_str = '{:02X}' return separator.join([format_str.format(x) for x in bytes(byte_array)]) def bytes_from_hex_string(hex_string, separator=''): if separator: hex_string = ''.join(hex_string.split(separator)) try: binary_data = binascii.a2b_hex(hex_string) except (TypeError, ValueError) as e: raise ValueError(*e.args) from e return binary_data cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/common/x509.py000066400000000000000000000141011524413560000271530ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import collections import datetime import enum import hashlib import attr import asn1crypto from cryptodatahub.common.key import PublicKeyX509Base from cryptodatahub.common.stores import CertificateTransparencyLog, CertificateTransparencyLogParamsBase from cryptoparser.common.base import ( Opaque, OpaqueParam, Serializable, VectorParamParsable, VectorParsable, ) from cryptoparser.common.parse import ComposerBinary, ComposerText, ParsableBase, ParserBinary from cryptoparser.tls.algorithm import TlsSignatureAndHashAlgorithm, TlsSignatureAndHashAlgorithmFactory class CtExtensions(Opaque): @classmethod def get_param(cls): return OpaqueParam( min_byte_num=0, max_byte_num=2 ** 16 - 1, ) class CtSignature(Opaque): @classmethod def get_param(cls): return OpaqueParam( min_byte_num=0, max_byte_num=2 ** 16 - 1, ) class CtVersion(enum.IntEnum): V1 = 0x00 # pylint: disable=invalid-name @attr.s class SignedCertificateTimestamp(ParsableBase, Serializable): version = attr.ib(validator=attr.validators.in_(CtVersion)) log = attr.ib( converter=CertificateTransparencyLog.from_log_id, validator=attr.validators.instance_of(CertificateTransparencyLogParamsBase) ) timestamp = attr.ib(validator=attr.validators.instance_of(datetime.datetime)) extensions = attr.ib( validator=attr.validators.instance_of(CtExtensions) ) signature_algorithm = attr.ib(validator=attr.validators.in_(TlsSignatureAndHashAlgorithm)) signature = attr.ib( converter=CtSignature, validator=attr.validators.instance_of(CtSignature), metadata={'human_friendly': False}, ) @classmethod def _parse(cls, parsable): header_parser = ParserBinary(parsable) header_parser.parse_bytes('sct', 2) body_parser = ParserBinary(header_parser['sct']) body_parser.parse_numeric('version', 1, CtVersion) body_parser.parse_raw('log', 32) body_parser.parse_timestamp('timestamp', milliseconds=True) body_parser.parse_parsable('extensions', CtExtensions) body_parser.parse_parsable('signature_algorithm', TlsSignatureAndHashAlgorithmFactory) body_parser.parse_parsable('signature', CtSignature) return cls(**body_parser), header_parser.parsed_length def compose(self): body_composer = ComposerBinary() body_composer.compose_numeric(self.version, 1) body_composer.compose_raw(self.log.log_id.value) body_composer.compose_timestamp(self.timestamp, milliseconds=True) body_composer.compose_parsable(self.extensions) body_composer.compose_numeric_enum_coded(self.signature_algorithm) body_composer.compose_parsable(self.signature) header_composer = ComposerBinary() header_composer.compose_numeric(len(body_composer.composed_bytes), 2) return header_composer.composed_bytes + body_composer.composed_bytes class SignedCertificateTimestampList(VectorParsable): @classmethod def get_param(cls): return VectorParamParsable( item_class=SignedCertificateTimestamp, fallback_class=None, min_byte_num=0, max_byte_num=2 ** 16 - 1 ) @attr.s(frozen=True) class TlsJA4XFingerprint(Serializable): fingerprint = attr.ib(validator=attr.validators.instance_of(str)) fingerprint_raw = attr.ib(validator=attr.validators.instance_of(str)) @attr.s(eq=False, init=False, frozen=True) class PublicKeyX509(PublicKeyX509Base): ja4x = attr.ib(init=False, default=None, metadata={'human_readable_name': 'JA4X'}) def __init__(self, certificate): super().__init__(certificate) object.__setattr__(self, 'ja4x', self._ja4x()) @property def signed_certificate_timestamps(self): for extension in self._certificate['tbs_certificate']['extensions']: if extension['extn_id'].dotted == '1.3.6.1.4.1.11129.2.4.2': asn1_value = asn1crypto.core.load(bytes(extension['extn_value'])) return SignedCertificateTimestampList.parse_exact_size(bytes(asn1_value)) return SignedCertificateTimestampList([]) @staticmethod def _ja4x_sha256(oid_hexes): composer = ComposerText() composer.compose_string_array(oid_hexes, ',') return hashlib.sha256(composer.composed).hexdigest()[:12] @staticmethod def _ja4x_relative_distinguished_name_oid_hexes(name): return [ attribute['type'].contents.hex() for relative_distinguished_name in name.chosen for attribute in relative_distinguished_name ] def _ja4x(self): tbs_certificate = self._certificate['tbs_certificate'] issuer_oid_hexes = self._ja4x_relative_distinguished_name_oid_hexes(tbs_certificate['issuer']) subject_oid_hexes = self._ja4x_relative_distinguished_name_oid_hexes(tbs_certificate['subject']) extension_oid_hexes = [ extension['extn_id'].contents.hex() for extension in tbs_certificate['extensions'] ] fingerprint_composer = ComposerText() fingerprint_composer.compose_string_array([ self._ja4x_sha256(issuer_oid_hexes), self._ja4x_sha256(subject_oid_hexes), self._ja4x_sha256(extension_oid_hexes), ], '_') raw_composer = ComposerText() raw_composer.compose_string_array(issuer_oid_hexes, ',') for oid_hexes in (subject_oid_hexes, extension_oid_hexes): raw_composer.compose_separator('_') raw_composer.compose_string_array(oid_hexes, ',') return TlsJA4XFingerprint( fingerprint=fingerprint_composer.composed.decode('ascii'), fingerprint_raw=raw_composer.composed.decode('ascii'), ) def _asdict(self): dict_value = super()._asdict() return collections.OrderedDict(list(dict_value.items()) + [ ('ja4x', self.ja4x), ('signed_certificate_timestamps', self.signed_certificate_timestamps), ]) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/dnsrec/000077500000000000000000000000001524413560000260655ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/dnsrec/__init__.py000066400000000000000000000000431524413560000301730ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/dnsrec/record.py000066400000000000000000000443531524413560000277260ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import collections import datetime import enum import attr from cryptodatahub.common.algorithm import Authentication, Signature from cryptodatahub.common.parameter import ECParamWellKnown from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.key import ( PublicKey, PublicKeyParamsDsa, PublicKeyParamsEcdsa, PublicKeyParamsEddsa, PublicKeyParamsRsa, ) from cryptodatahub.dnsrec.algorithm import ( DnsRrType, DnsSecAlgorithm, DnsSecDigestType, SshFpAlgorithm, SshFpFingerprintType, ) from cryptoparser.common.base import NumericRangeParsableBase, OneByteEnumParsable, Serializable, TwoByteEnumParsable from cryptoparser.common.exception import NotEnoughData from cryptoparser.common.parse import ByteOrder, ComposerBinary, ParsableBase, ParserBinary class DnsSecProtocol(enum.Enum): V3 = 3 class DnsSecFlag(enum.IntEnum): SECURE_ENTRY_POINT = 0x0001 REVOKE = 0x0080 DNS_ZONE_KEY = 0x0100 class DnsSecAlgorithmFactory(OneByteEnumParsable): @classmethod def get_enum_class(cls): return DnsSecAlgorithm @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s class DnsRecordDnskey(ParsableBase, Serializable): HEADER_SIZE = 4 flags = attr.ib(validator=attr.validators.deep_iterable(attr.validators.instance_of(DnsSecFlag))) algorithm = attr.ib(validator=attr.validators.instance_of(DnsSecAlgorithm)) key = attr.ib(validator=attr.validators.instance_of(PublicKey)) protocol = attr.ib(validator=attr.validators.instance_of(DnsSecProtocol)) def __attrs_post_init__(self): if not isinstance(self.algorithm.value.algorithm, Signature): raise InvalidValue(self.algorithm.value.algorithm, type(self), 'algorithm_type') algorithm_key_type = self.algorithm.value.algorithm.value.key_type algorithm_incompatible = self.key.key_type != algorithm_key_type algorithm_compatible_gost = ( self.algorithm == DnsSecAlgorithm.ECCGOST and self.key.key_type == Authentication.ECDSA and self.key.params.key_parameter == ECParamWellKnown.GC256B ) if (algorithm_incompatible and not algorithm_compatible_gost): raise InvalidValue(algorithm_key_type, type(self), 'key_type') @property def key_tag(self): if self.algorithm == DnsSecAlgorithm.RSAMD5: return (self.key.params.modulus & 0xffffff) >> 8 key_tag = 0 parser = ParserBinary(self.compose(), byte_order=ByteOrder.BIG_ENDIAN) while parser.unparsed_length > 1: parser.parse_numeric('value', 2) key_tag += parser['value'] if parser.unparsed_length: parser.parse_numeric('value', 1) key_tag += parser['value'] key_tag += (key_tag >> 16) & 0xffff return key_tag & 0xffff def _asdict(self): dict_value = super()._asdict() return collections.OrderedDict([('key_tag', self.key_tag)] + list(dict_value.items())) @classmethod def _parse_public_key_rsa(cls, key_parser): key_parser.parse_numeric('exponent_length_one_octet', 1) exponent_length = key_parser['exponent_length_one_octet'] if exponent_length == 0: key_parser.parse_numeric('exponent_length_two_octets', 2) exponent_length = key_parser['exponent_length_two_octets'] key_parser.parse_mpint('public_exponent', exponent_length) key_parser.parse_mpint('modulus', key_parser.unparsed_length) return PublicKey.from_params(PublicKeyParamsRsa( public_exponent=key_parser['public_exponent'], modulus=key_parser['modulus'], )) @classmethod def _parse_public_key_ecdsa(cls, dnssec_algorithm, key_parser): if dnssec_algorithm == DnsSecAlgorithm.ECDSAP256SHA256: key_parameter = ECParamWellKnown.SECP256K1 elif dnssec_algorithm == DnsSecAlgorithm.ECDSAP384SHA384: key_parameter = ECParamWellKnown.SECP384R1 elif dnssec_algorithm == DnsSecAlgorithm.ECCGOST: key_parameter = ECParamWellKnown.GC256B else: raise NotImplementedError(dnssec_algorithm) key_size = key_parameter.value.field_size // 8 key_parser.parse_mpint('x', key_size) key_parser.parse_mpint('y', key_size) return PublicKey.from_params(PublicKeyParamsEcdsa( point_x=key_parser['x'], point_y=key_parser['y'], key_parameter=key_parameter, )) @classmethod def _parse_public_key_eddsa(cls, dnssec_algorithm, key_parser): if dnssec_algorithm == DnsSecAlgorithm.ED25519: key_parameter = ECParamWellKnown.CURVE25519 key_parser.parse_raw('public_key', 256 // 8) elif dnssec_algorithm == DnsSecAlgorithm.ED448: key_parameter = ECParamWellKnown.CURVE448 key_parser.parse_raw('public_key', 448 // 8) else: raise NotImplementedError(dnssec_algorithm) return PublicKey.from_params(PublicKeyParamsEddsa( key_parameter=key_parameter, key_data=key_parser['public_key'] )) @classmethod def _parse_public_key_dss(cls, key_parser): key_parser.parse_numeric('t', 1) key_parser.parse_mpint('q', 20) mpint_length = 64 + key_parser['t'] * 8 key_parser.parse_mpint('p', mpint_length) key_parser.parse_mpint('g', mpint_length) key_parser.parse_mpint('y', mpint_length) return PublicKey.from_params(PublicKeyParamsDsa( prime=key_parser['p'], generator=key_parser['g'], order=key_parser['q'], public_key_value=key_parser['y'], )) @classmethod def parse_key(cls, parsable, dnssec_algorithm): key_parser = ParserBinary(parsable) public_key_type = dnssec_algorithm.value.algorithm.value.key_type if public_key_type == Authentication.RSA: public_key = cls._parse_public_key_rsa(key_parser) elif public_key_type in [Authentication.ECDSA, Authentication.GOST_R3410_01]: public_key = cls._parse_public_key_ecdsa(dnssec_algorithm, key_parser) elif public_key_type == Authentication.EDDSA: public_key = cls._parse_public_key_eddsa(dnssec_algorithm, key_parser) elif public_key_type == Authentication.DSS: public_key = cls._parse_public_key_dss(key_parser) else: raise NotImplementedError(dnssec_algorithm) return public_key @classmethod def _parse(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric_flags('flags', 2, DnsSecFlag) parser.parse_numeric('protocol', 1, DnsSecProtocol) parser.parse_parsable('algorithm', DnsSecAlgorithmFactory) parser.parse_raw('key', parser.unparsed_length) public_key = cls.parse_key(parser['key'], parser['algorithm']) return cls( parser['flags'], parser['algorithm'], public_key, parser['protocol'], ), parser.parsed_length @staticmethod def _compose_public_key_rsa(key_composer, key): key_params = key.params exponent_length = (key_params.public_exponent.bit_length() + 7) // 8 if exponent_length > 255: key_composer.compose_numeric(0, 1) key_composer.compose_numeric(exponent_length, 2) else: key_composer.compose_numeric(exponent_length, 1) key_composer.compose_mpint(key_params.public_exponent, exponent_length) key_composer.compose_mpint(key_params.modulus, key.key_size // 8) @staticmethod def _compose_public_key_ecdsa(key_composer, key): key_params = key.params key_size = key.key_size // 8 key_composer.compose_mpint(key_params.point_x, key_size) key_composer.compose_mpint(key_params.point_y, key_size) @staticmethod def _compose_public_key_eddsa(key_composer, key): key_params = key.params key_composer.compose_raw(key_params.key_data) @staticmethod def _compose_public_key_dss(key_composer, key): key_params = key.params key_size = key.key_size // 8 key_composer.compose_numeric((key_size - 64) // 8, 1) key_composer.compose_mpint(key_params.order, 20) key_composer.compose_mpint(key_params.prime, key_size) key_composer.compose_mpint(key_params.generator, key_size) key_composer.compose_mpint(key_params.public_key_value, key_size) @staticmethod def compose_key(key): key_composer = ComposerBinary() public_key_type = key.key_type if public_key_type == Authentication.RSA: DnsRecordDnskey._compose_public_key_rsa(key_composer, key) elif public_key_type in [Authentication.ECDSA, Authentication.GOST_R3410_01]: DnsRecordDnskey._compose_public_key_ecdsa(key_composer, key) elif public_key_type == Authentication.EDDSA: DnsRecordDnskey._compose_public_key_eddsa(key_composer, key) elif public_key_type == Authentication.DSS: DnsRecordDnskey._compose_public_key_dss(key_composer, key) else: raise NotImplementedError(public_key_type) return key_composer.composed_bytes def compose(self): composer = ComposerBinary() composer.compose_numeric_flags(self.flags, 2) composer.compose_numeric(self.protocol.value, 1) composer.compose_numeric_enum_coded(self.algorithm) key_bytes = self.compose_key(self.key) return composer.composed_bytes + key_bytes class DnsSecDigestTypeFactory(OneByteEnumParsable): @classmethod def get_enum_class(cls): return DnsSecDigestType @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s class DnsRecordDs(ParsableBase): HEADER_SIZE = 4 key_tag = attr.ib(validator=attr.validators.instance_of(int)) algorithm = attr.ib(validator=attr.validators.instance_of(DnsSecAlgorithm)) digest_type = attr.ib(validator=attr.validators.instance_of(DnsSecDigestType)) digest = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def _parse(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric('key_tag', 2) parser.parse_parsable('algorithm', DnsSecAlgorithmFactory) parser.parse_parsable('digest_type', DnsSecDigestTypeFactory) parser.parse_raw('digest', parser.unparsed_length) return cls(**parser), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.key_tag, 2) composer.compose_numeric_enum_coded(self.algorithm) composer.compose_numeric_enum_coded(self.digest_type) composer.compose_raw(self.digest) return composer.composed_bytes class DnsRrTypeFactory(TwoByteEnumParsable): @classmethod def get_enum_class(cls): return DnsRrType @abc.abstractmethod def compose(self): raise NotImplementedError() class DnsRrTypePrivate(NumericRangeParsableBase): @classmethod def _get_value_min(cls): return 0xff00 @classmethod def _get_value_max(cls): return 0xfffe @classmethod def _get_value_length(cls): return 2 @attr.s class DnsNameUncompressed(ParsableBase, Serializable): labels = attr.ib( validator=attr.validators.deep_iterable(member_validator=attr.validators.instance_of(str)) ) def __str__(self): return '.'.join(self.labels) def _as_markdown(self, level): return self._markdown_result(str(self), level) @classmethod def convert(cls, value): if isinstance(value, cls): return value if isinstance(value, str): if not value: return cls([]) return cls(value.split('.')) raise InvalidValue(value, cls, 'labels') @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) labels = [] while True: parser.parse_string('label', 1, encoding='idna') label = parser['label'] if not label: break labels.append(label) return cls(labels), parser.parsed_length def compose(self): composer = ComposerBinary() for label in self.labels: composer.compose_string(label, 'idna', 1) composer.compose_numeric(0, 1) return composer.composed_bytes @attr.s class DnsRecordRrsig(ParsableBase): # pylint: disable=too-many-instance-attributes HEADER_SIZE = 24 type_covered = attr.ib(validator=attr.validators.instance_of((DnsRrType, DnsRrTypePrivate))) algorithm = attr.ib(validator=attr.validators.instance_of(DnsSecAlgorithm)) labels = attr.ib(validator=attr.validators.instance_of(int)) original_ttl = attr.ib( validator=attr.validators.instance_of(int), metadata={'human_readable_name': 'Original TTL'} ) signature_expiration = attr.ib(validator=attr.validators.instance_of(datetime.datetime)) signature_inception = attr.ib(validator=attr.validators.instance_of(datetime.datetime)) key_tag = attr.ib(validator=attr.validators.instance_of(int)) signers_name = attr.ib( converter=DnsNameUncompressed.convert, validator=attr.validators.instance_of(DnsNameUncompressed) ) signature = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)), metadata={'human_friendly': False} ) @classmethod def _parse(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) try: parser.parse_parsable('type_covered', DnsRrTypeFactory) except InvalidValue: parser.parse_parsable('type_covered', DnsRrTypePrivate) parser.parse_parsable('algorithm', DnsSecAlgorithmFactory) parser.parse_numeric('labels', 1) parser.parse_numeric('original_ttl', 4) parser.parse_timestamp('signature_expiration', item_size=4) parser.parse_timestamp('signature_inception', item_size=4) parser.parse_numeric('key_tag', 2) parser.parse_parsable('signers_name', DnsNameUncompressed) parser.parse_raw('signature', parser.unparsed_length) return cls(**parser), parser.parsed_length def compose(self): composer = ComposerBinary() if isinstance(self.type_covered, DnsRrType): composer.compose_numeric_enum_coded(self.type_covered) else: composer.compose_parsable(self.type_covered) composer.compose_numeric_enum_coded(self.algorithm) composer.compose_numeric(self.labels, 1) composer.compose_numeric(self.original_ttl, 4) composer.compose_timestamp(self.signature_expiration, item_size=4) composer.compose_timestamp(self.signature_inception, item_size=4) composer.compose_numeric(self.key_tag, 2) composer.compose_parsable(self.signers_name) composer.compose_raw(self.signature) return composer.composed_bytes @attr.s class DnsRecordMx(ParsableBase): HEADER_SIZE = 2 priority = attr.ib(validator=attr.validators.instance_of(int)) exchange = attr.ib( converter=DnsNameUncompressed.convert, validator=attr.validators.instance_of(DnsNameUncompressed) ) @classmethod def _parse(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric('priority', 2) parser.parse_parsable('exchange', DnsNameUncompressed) return cls(**parser), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.priority, 2) composer.compose_parsable(self.exchange) return composer.composed_bytes class SshFpAlgorithmFactory(OneByteEnumParsable): @classmethod def get_enum_class(cls): return SshFpAlgorithm @abc.abstractmethod def compose(self): raise NotImplementedError() class SshFpFingerprintTypeFactory(OneByteEnumParsable): @classmethod def get_enum_class(cls): return SshFpFingerprintType @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s class DnsRecordSshfp(ParsableBase, Serializable): HEADER_SIZE = 2 algorithm = attr.ib(validator=attr.validators.instance_of(SshFpAlgorithm)) fingerprint_type = attr.ib(validator=attr.validators.instance_of(SshFpFingerprintType)) fingerprint = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def _parse(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_parsable('algorithm', SshFpAlgorithmFactory) parser.parse_parsable('fingerprint_type', SshFpFingerprintTypeFactory) parser.parse_raw('fingerprint', parser.unparsed_length) return cls(**parser), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric_enum_coded(self.algorithm) composer.compose_numeric_enum_coded(self.fingerprint_type) composer.compose_raw(self.fingerprint) return composer.composed_bytes @attr.s class DnsRecordTxt(ParsableBase): HEADER_SIZE = 1 value = attr.ib(validator=attr.validators.instance_of(str)) @classmethod def _parse(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) value = '' while parser.unparsed_length: parser.parse_string('value', 1, encoding='ascii') value += parser['value'] return cls(value), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_string(self.value, 'ascii', 1) return composer.composed_bytes cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/dnsrec/txt.py000066400000000000000000000646661524413560000273000ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import abc import collections import enum import ipaddress import attr from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.types import convert_url, convert_value_to_object from cryptoparser.common.base import ( Serializable, StringEnumCaseInsensitiveParsable, StringEnumParsable, VariantParsable ) from cryptoparser.common.exception import InvalidType from cryptoparser.common.field import ( FieldValueComponentParsable, FieldValueComponentParsableOptional, FieldValueComponentNumber, FieldValueComponentPercent, FieldValueComponentString, FieldValueComponentUrl, FieldValueSingleBase, FieldValueStringEnumParams, FieldsSemicolonSeparated, NameValuePairListSemicolonSeparated, NameValuePair, ) from cryptoparser.common.parse import ComposerText, ParsableBase, ParserText class DmarcAlignment(StringEnumCaseInsensitiveParsable, enum.Enum): RELAXED = FieldValueStringEnumParams( code='r', human_readable_name='Relaxed', ) STRICT = FieldValueStringEnumParams( code='s', human_readable_name='Strict', ) class DmarcIdentifierAlignmentBase(FieldValueComponentParsable): @classmethod @abc.abstractmethod def get_canonical_name(cls): raise NotImplementedError() @classmethod def _get_value_class(cls): return DmarcAlignment class DnsRecordTxtValueDmarcValueIdentifierAlignmentDkim(DmarcIdentifierAlignmentBase): @classmethod def get_canonical_name(cls): return 'adkim' class DnsRecordTxtValueDmarcValueIdentifierAlignmentAspf(DmarcIdentifierAlignmentBase): @classmethod def get_canonical_name(cls): return 'aspf' class DmarcPolicyOption(StringEnumCaseInsensitiveParsable, enum.Enum): NONE = FieldValueStringEnumParams( code='none' ) QUARANTINE = FieldValueStringEnumParams( code='quarantine' ) REJECT = FieldValueStringEnumParams( code='reject' ) class DnsRecordTxtValueDmarcValuePolicy(FieldValueComponentParsable): @classmethod def get_canonical_name(cls): return 'p' @classmethod def _get_value_class(cls): return DmarcPolicyOption class DnsRecordTxtValueDmarcValueSubdomainPolicy(FieldValueComponentParsableOptional): @classmethod def get_canonical_name(cls): return 'sp' @classmethod def _get_value_class(cls): return DmarcPolicyOption class DmarcPolicyVersion(StringEnumParsable, enum.Enum): DMARC1 = FieldValueStringEnumParams( code='DMARC1', human_readable_name='DMARC1', ) class DnsRecordTxtValueDmarcValueVersion(FieldValueComponentParsable): @classmethod def get_canonical_name(cls): return 'v' @classmethod def _get_value_class(cls): return DmarcPolicyVersion class DmarcFailureReportingOption(StringEnumCaseInsensitiveParsable, enum.Enum): ALL_FAILURE = FieldValueStringEnumParams( code='0', human_readable_name='All Failure', ) ANY_FAILURE = FieldValueStringEnumParams( code='1', human_readable_name='Any Failure', ) DKIM_FAILURE = FieldValueStringEnumParams( code='d', human_readable_name='DKIM Failure', ) SPF_FAILURE = FieldValueStringEnumParams( code='s', human_readable_name='SPF Failure', ) class DmarcValueFailureOption(FieldValueComponentParsable): @classmethod def get_canonical_name(cls): return 'fo' @classmethod def _get_value_class(cls): return DmarcFailureReportingOption class DnsRecordTxtValueDmarcValuePercent(FieldValueComponentPercent): @classmethod def get_canonical_name(cls): return 'pct' @attr.s class DmarcReportingInterval(FieldValueComponentNumber): def __attrs_post_init__(self): if self.value < 0 or self.value >= 2 ** 32: raise InvalidValue(self.value, type(self), 'value') @classmethod def get_canonical_name(cls): return 'ri' class DmarcFailureReportingFormat(StringEnumCaseInsensitiveParsable, enum.Enum): AUTHENTICATION_FAILURE_REPORTING_FORMAT = FieldValueStringEnumParams( code='afrf', human_readable_name='Authentication Failure Reporting Format (AFRF)', ) class DmarcReportingFormat(FieldValueComponentParsable): @classmethod def get_canonical_name(cls): return 'rf' @classmethod def _get_value_class(cls): return DmarcFailureReportingFormat class DnsRecordTxtValueDmarcValueReportingUrlAggregated(FieldValueComponentUrl): @classmethod def get_canonical_name(cls): return 'rua' class DnsRecordTxtValueDmarcValueReportingUrlFailure(FieldValueComponentUrl): @classmethod def get_canonical_name(cls): return 'ruf' @attr.s class DnsRecordTxtValueDmarc(FieldsSemicolonSeparated): # pylint: disable=too-many-instance-attributes version = attr.ib( converter=DnsRecordTxtValueDmarcValueVersion.convert, validator=attr.validators.instance_of(DnsRecordTxtValueDmarcValueVersion) ) policy = attr.ib( converter=DnsRecordTxtValueDmarcValuePolicy.convert, validator=attr.validators.instance_of(DnsRecordTxtValueDmarcValuePolicy) ) alignment_dkim = attr.ib( default=DmarcAlignment.RELAXED, converter=DnsRecordTxtValueDmarcValueIdentifierAlignmentDkim.convert, validator=attr.validators.instance_of(DnsRecordTxtValueDmarcValueIdentifierAlignmentDkim), metadata={'human_readable_name': 'DKIM Alignment'}, ) alignment_aspf = attr.ib( default=DmarcAlignment.RELAXED, converter=DnsRecordTxtValueDmarcValueIdentifierAlignmentAspf.convert, validator=attr.validators.instance_of(DnsRecordTxtValueDmarcValueIdentifierAlignmentAspf), metadata={'human_readable_name': 'ASPF Alignment'}, ) failure_option = attr.ib( default=DmarcFailureReportingOption.ALL_FAILURE, converter=DmarcValueFailureOption.convert, validator=attr.validators.instance_of(DmarcValueFailureOption), ) percent = attr.ib( default=100, converter=DnsRecordTxtValueDmarcValuePercent.convert, validator=attr.validators.instance_of(DnsRecordTxtValueDmarcValuePercent), ) reporting_url_aggregated = attr.ib( default=None, converter=convert_value_to_object(DnsRecordTxtValueDmarcValueReportingUrlAggregated, convert_url()), validator=attr.validators.optional( attr.validators.instance_of(DnsRecordTxtValueDmarcValueReportingUrlAggregated) ), metadata={'human_readable_name': 'Aggregated Reporting URL'}, ) reporting_url_failure = attr.ib( default=None, converter=convert_value_to_object(DnsRecordTxtValueDmarcValueReportingUrlFailure, convert_url()), validator=attr.validators.optional(attr.validators.instance_of(DnsRecordTxtValueDmarcValueReportingUrlFailure)), metadata={'human_readable_name': 'Failure Reporting URL'}, ) reporting_format = attr.ib( default=DmarcFailureReportingFormat.AUTHENTICATION_FAILURE_REPORTING_FORMAT, converter=DmarcReportingFormat.convert, validator=attr.validators.instance_of(DmarcReportingFormat), ) reporting_interval = attr.ib( default=86400, converter=DmarcReportingInterval.convert, validator=attr.validators.instance_of(DmarcReportingInterval), ) subdomain_policy = attr.ib( default=None, converter=DnsRecordTxtValueDmarcValueSubdomainPolicy.convert, validator=attr.validators.optional(attr.validators.instance_of(DnsRecordTxtValueDmarcValueSubdomainPolicy)) ) class MtaStsPolicyVersion(StringEnumParsable, enum.Enum): STSV1 = FieldValueStringEnumParams( code='STSv1', human_readable_name='STSv1', ) class DnsRecordMtaStsValueVersion(FieldValueComponentParsable): @classmethod def get_canonical_name(cls): return 'v' @classmethod def _get_value_class(cls): return MtaStsPolicyVersion class DnsRecordMtaStsValueId(FieldValueComponentString): @classmethod def get_canonical_name(cls): return 'id' @attr.s class DnsRecordTxtValueMtaSts(FieldsSemicolonSeparated): version = attr.ib( converter=DnsRecordMtaStsValueVersion.convert, validator=attr.validators.instance_of(DnsRecordMtaStsValueVersion) ) identifier = attr.ib( converter=DnsRecordMtaStsValueId.convert, validator=attr.validators.instance_of(DnsRecordMtaStsValueId) ) extensions = attr.ib( default=None, validator=attr.validators.optional(attr.validators.instance_of(NameValuePairListSemicolonSeparated)), metadata={'extension': True}, ) class TlsRptVersion(StringEnumParsable, enum.Enum): TLSRPTV1 = FieldValueStringEnumParams( code='TLSRPTv1', human_readable_name='TLSRPTv1', ) class DnsRecordTxtValueTlsRptValueVersion(FieldValueComponentParsable): @classmethod def get_canonical_name(cls): return 'v' @classmethod def _get_value_class(cls): return TlsRptVersion class DnsRecordTxtValueTlsRptValueReportingUrlAggregated(FieldValueComponentUrl): @classmethod def get_canonical_name(cls): return 'rua' @attr.s class DnsRecordTxtValueTlsRpt(FieldsSemicolonSeparated): version = attr.ib( converter=DnsRecordTxtValueTlsRptValueVersion.convert, validator=attr.validators.instance_of(DnsRecordTxtValueTlsRptValueVersion) ) reporting_url_aggregated = attr.ib( default=None, converter=DnsRecordTxtValueTlsRptValueReportingUrlAggregated.convert, validator=attr.validators.optional( attr.validators.instance_of(DnsRecordTxtValueTlsRptValueReportingUrlAggregated) ), ) extensions = attr.ib( default=None, validator=attr.validators.optional(attr.validators.instance_of(NameValuePairListSemicolonSeparated)), metadata={'extension': True}, ) class SpfVersion(StringEnumParsable, enum.Enum): SPF1 = FieldValueStringEnumParams( code='spf1', human_readable_name='SPF1', ) class DnsRecordTxtValueSpfVersion(FieldValueComponentParsable): @classmethod def get_canonical_name(cls): return 'v' @classmethod def _get_value_class(cls): return SpfVersion class SpfQualifier(StringEnumParsable, enum.Enum): PASS = FieldValueStringEnumParams( code='+', human_readable_name='Pass', ) FAIL = FieldValueStringEnumParams( code='-', human_readable_name='Fail', ) SOFTFAIL = FieldValueStringEnumParams( code='~', human_readable_name='Softfail', ) NEUTRAL = FieldValueStringEnumParams( code='?', human_readable_name='Neutral', ) class SpfMechanism(StringEnumParsable, enum.Enum): ALL = FieldValueStringEnumParams( code='all', human_readable_name='All', ) INCLUDE = FieldValueStringEnumParams( code='include', human_readable_name='Include', ) A = FieldValueStringEnumParams( code='a', human_readable_name='A/AAAA records', ) MX = FieldValueStringEnumParams( code='mx', human_readable_name='MX records', ) PTR = FieldValueStringEnumParams( code='ptr', human_readable_name='PTR records', ) IP4 = FieldValueStringEnumParams( code='ip4', human_readable_name='IPv4 records', ) IP6 = FieldValueStringEnumParams( code='ip6', human_readable_name='IPv6 records', ) EXISTS = FieldValueStringEnumParams( code='exists', human_readable_name='Exists', ) class SpfModifier(StringEnumParsable, enum.Enum): REDIRECT = FieldValueStringEnumParams( code='redirect', human_readable_name='Redirect', ) EXP = FieldValueStringEnumParams( code='exp', human_readable_name='Explanation', ) @attr.s class SpfDomainSpec(FieldValueSingleBase): @classmethod def _get_value_type(cls): return str @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_until_separator_or_end('value', separators=' ') return cls(**parser), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string(self.value) return composer.composed class DnsRecordTxtValueSpfModifierKnownBase(FieldValueComponentParsable): @classmethod @abc.abstractmethod def get_modifier(cls): raise NotImplementedError() @classmethod def get_canonical_name(cls): return cls.get_modifier().value.code @classmethod def _get_value_class(cls): return SpfDomainSpec class DnsRecordTxtValueSpfModifierRedirect(DnsRecordTxtValueSpfModifierKnownBase): @classmethod def get_modifier(cls): return SpfModifier.REDIRECT class DnsRecordTxtValueSpfModifierExplanation(DnsRecordTxtValueSpfModifierKnownBase): @classmethod def get_modifier(cls): return SpfModifier.EXP class DnsRecordTxtValueSpfModifierUnknown(NameValuePair): pass class DnsRecordTxtValueSpfDirectiveBase(ParsableBase, Serializable): @classmethod @abc.abstractmethod def get_mechanism(cls): raise NotImplementedError() @classmethod def _parse_qualifier_and_mechanism_name(cls, parsable): parser = ParserText(parsable) try: parser.parse_parsable('qualifier', SpfQualifier) except InvalidValue: pass mechanism = cls.get_mechanism() try: parser.parse_string('mechanism', mechanism.value.code) except InvalidValue as e: raise InvalidType from e return parser @classmethod def _parse_domain(cls, parser, optional, extra_separators=''): has_separator = True if optional: try: parser.parse_separator(':', min_length=1, max_length=1) except InvalidValue: has_separator = False else: parser.parse_separator(':', min_length=1, max_length=1) if has_separator: parser.parse_string_until_separator_or_end('domain', separators=' ' + extra_separators) return parser.get('domain', None) @classmethod def _compose_domain(cls, composer, domain): if domain is None: return composer.compose_separator(':') composer.compose_string(domain) @classmethod def _parse_ip_network(cls, parser): parser.parse_string('separator', ':') parser.parse_string_until_separator_or_end('ip_network', ' ') return parser['ip_network'] @classmethod def _compose_ip_network(cls, composer, ip_network): composer.compose_separator(':') composer.compose_string(str(ip_network.network_address)) if ip_network.prefixlen != ip_network.max_prefixlen: composer.compose_separator('/') composer.compose_numeric(ip_network.prefixlen) @classmethod def _parse_ip_cidr_length(cls, parser): has_separator = True try: parser.parse_separator('/', min_length=1, max_length=1) except InvalidValue: has_separator = False if has_separator: parser.parse_numeric('ip_cidr_length') ip_cidr_length = parser['ip_cidr_length'] del parser['ip_cidr_length'] else: ip_cidr_length = None return ip_cidr_length @classmethod def _compose_ip_cidr_length(cls, composer, ip_cidr_length): if ip_cidr_length is None: return composer.compose_separator('/') composer.compose_numeric(ip_cidr_length) def _compose_qualifier_and_mechanism_name(self, qualifier): composer = ComposerText() if qualifier is not None: composer.compose_string(qualifier.value.code) composer.compose_string(self.get_mechanism().value.code) return composer @attr.s class DnsRecordTxtValueSpfDirectiveAll(DnsRecordTxtValueSpfDirectiveBase): qualifier = attr.ib( default=None, validator=attr.validators.optional(attr.validators.in_(SpfQualifier)) ) @classmethod def get_mechanism(cls): return SpfMechanism.ALL @classmethod def _parse(cls, parsable): parser = cls._parse_qualifier_and_mechanism_name(parsable) return cls(qualifier=parser.get('qualifier', None)), parser.parsed_length def compose(self): composer = self._compose_qualifier_and_mechanism_name(self.qualifier) return composer.composed @attr.s class DnsRecordTxtValueSpfDirectiveDomain(DnsRecordTxtValueSpfDirectiveBase): domain = attr.ib( converter=SpfDomainSpec.convert, validator=attr.validators.instance_of(SpfDomainSpec) ) qualifier = attr.ib( default=None, validator=attr.validators.optional(attr.validators.in_(SpfQualifier)) ) @classmethod @abc.abstractmethod def get_mechanism(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = cls._parse_qualifier_and_mechanism_name(parsable) qualifier = parser.get('qualifier', None) domain = cls._parse_domain(parser, False) return cls(qualifier=qualifier, domain=domain), parser.parsed_length def compose(self): composer = self._compose_qualifier_and_mechanism_name(self.qualifier) self._compose_domain(composer, self.domain) return composer.composed @attr.s class DnsRecordTxtValueSpfDirectivePtr(DnsRecordTxtValueSpfDirectiveBase): domain = attr.ib( default=None, converter=SpfDomainSpec.convert, validator=attr.validators.optional(attr.validators.instance_of(SpfDomainSpec)), ) qualifier = attr.ib( default=None, validator=attr.validators.optional(attr.validators.in_(SpfQualifier)), ) @classmethod def get_mechanism(cls): return SpfMechanism.PTR @classmethod def _parse(cls, parsable): parser = cls._parse_qualifier_and_mechanism_name(parsable) qualifier = parser.get('qualifier', None) domain = cls._parse_domain(parser, True) return cls(qualifier=qualifier, domain=domain), parser.parsed_length def compose(self): composer = self._compose_qualifier_and_mechanism_name(self.qualifier) self._compose_domain(composer, self.domain) return composer.composed @attr.s class DnsRecordTxtValueSpfDirectiveDomainCidr(DnsRecordTxtValueSpfDirectiveBase): domain = attr.ib( default=None, converter=SpfDomainSpec.convert, validator=attr.validators.optional(attr.validators.instance_of(SpfDomainSpec)), ) ipv4_cidr_length = attr.ib( default=None, validator=attr.validators.optional(attr.validators.instance_of(int)), ) ipv6_cidr_length = attr.ib( default=None, validator=attr.validators.optional(attr.validators.instance_of(int)), ) qualifier = attr.ib( default=None, validator=attr.validators.optional(attr.validators.in_(SpfQualifier)), ) @classmethod @abc.abstractmethod def get_mechanism(cls): raise NotImplementedError() @ipv4_cidr_length.validator def _validator_ipv4_cidr_length(self, attribute, value): # pylint: disable=unused-argument if value is None: return if value < 0 or value > 32: raise InvalidValue(value, type(self), 'ipv4_cidr_length') @ipv6_cidr_length.validator def _validator_ipv6_cidr_length(self, attribute, value): # pylint: disable=unused-argument if value is None: return if value < 0 or value > 128: raise InvalidValue(value, type(self), 'ipv6_cidr_length') @classmethod def _parse(cls, parsable): parser = cls._parse_qualifier_and_mechanism_name(parsable) qualifier = parser.get('qualifier', None) domain = cls._parse_domain(parser, True, extra_separators='/') ipv4_cidr_length = cls._parse_ip_cidr_length(parser) ipv6_cidr_length = cls._parse_ip_cidr_length(parser) return cls( qualifier=qualifier, domain=domain, ipv4_cidr_length=ipv4_cidr_length, ipv6_cidr_length=ipv6_cidr_length ), parser.parsed_length def compose(self): composer = self._compose_qualifier_and_mechanism_name(self.qualifier) self._compose_domain(composer, self.domain) self._compose_ip_cidr_length(composer, self.ipv4_cidr_length) self._compose_ip_cidr_length(composer, self.ipv6_cidr_length) return composer.composed class DnsRecordTxtValueSpfDirectiveInclude(DnsRecordTxtValueSpfDirectiveDomain): @classmethod def get_mechanism(cls): return SpfMechanism.INCLUDE class DnsRecordTxtValueSpfDirectiveA(DnsRecordTxtValueSpfDirectiveDomainCidr): @classmethod def get_mechanism(cls): return SpfMechanism.A class DnsRecordTxtValueSpfDirectiveMx(DnsRecordTxtValueSpfDirectiveDomainCidr): @classmethod def get_mechanism(cls): return SpfMechanism.MX @attr.s class DnsRecordTxtValueSpfDirectiveIp4(DnsRecordTxtValueSpfDirectiveBase): ipv4_network = attr.ib( converter=ipaddress.ip_network, validator=attr.validators.instance_of(ipaddress.IPv4Network) ) qualifier = attr.ib( default=None, validator=attr.validators.optional(attr.validators.in_(SpfQualifier)), ) @classmethod def get_mechanism(cls): return SpfMechanism.IP4 @classmethod def _parse(cls, parsable): parser = cls._parse_qualifier_and_mechanism_name(parsable) qualifier = parser.get('qualifier', None) ipv4_network = cls._parse_ip_network(parser) return cls( qualifier=qualifier, ipv4_network=ipv4_network, ), parser.parsed_length def compose(self): composer = self._compose_qualifier_and_mechanism_name(self.qualifier) self._compose_ip_network(composer, self.ipv4_network) return composer.composed @attr.s class DnsRecordTxtValueSpfDirectiveIp6(DnsRecordTxtValueSpfDirectiveBase): ipv6_network = attr.ib( converter=ipaddress.ip_network, validator=attr.validators.instance_of(ipaddress.IPv6Network) ) qualifier = attr.ib( default=None, validator=attr.validators.optional(attr.validators.in_(SpfQualifier)), ) @classmethod def get_mechanism(cls): return SpfMechanism.IP6 @classmethod def _parse(cls, parsable): parser = cls._parse_qualifier_and_mechanism_name(parsable) qualifier = parser.get('qualifier', None) ipv6_network = cls._parse_ip_network(parser) return cls( qualifier=qualifier, ipv6_network=ipv6_network, ), parser.parsed_length def compose(self): composer = self._compose_qualifier_and_mechanism_name(self.qualifier) self._compose_ip_network(composer, self.ipv6_network) return composer.composed class DnsRecordTxtValueSpfDirectiveExists(DnsRecordTxtValueSpfDirectiveDomain): @classmethod def get_mechanism(cls): return SpfMechanism.EXISTS class DnsRecordTxtValueSpfVariantParsable(VariantParsable): @classmethod def _get_variants(cls): return collections.OrderedDict([ (SpfMechanism.ALL, [DnsRecordTxtValueSpfDirectiveAll, ]), (SpfMechanism.INCLUDE, [DnsRecordTxtValueSpfDirectiveInclude, ]), (SpfMechanism.A, [DnsRecordTxtValueSpfDirectiveA, ]), (SpfMechanism.MX, [DnsRecordTxtValueSpfDirectiveMx, ]), (SpfMechanism.PTR, [DnsRecordTxtValueSpfDirectivePtr, ]), (SpfMechanism.IP4, [DnsRecordTxtValueSpfDirectiveIp4, ]), (SpfMechanism.IP6, [DnsRecordTxtValueSpfDirectiveIp6, ]), (SpfMechanism.EXISTS, [DnsRecordTxtValueSpfDirectiveExists, ]), (SpfModifier.REDIRECT, [DnsRecordTxtValueSpfModifierRedirect, ]), (SpfModifier.EXP, [DnsRecordTxtValueSpfModifierExplanation, ]), ]) @attr.s class DnsRecordTxtValueSpf(ParsableBase, Serializable): terms = attr.ib( validator=attr.validators.deep_iterable(member_validator=attr.validators.instance_of(( DnsRecordTxtValueSpfDirectiveBase, DnsRecordTxtValueSpfModifierKnownBase, DnsRecordTxtValueSpfModifierUnknown, ))) ) version = attr.ib( default=DnsRecordTxtValueSpfVersion(SpfVersion.SPF1), converter=DnsRecordTxtValueSpfVersion.convert, validator=attr.validators.instance_of(DnsRecordTxtValueSpfVersion) ) def _asdict(self): terms = [] for term in self.terms: if isinstance(term, DnsRecordTxtValueSpfModifierKnownBase): terms.append((term.get_modifier(), term.value.value)) elif isinstance(term, DnsRecordTxtValueSpfDirectiveBase): terms.append((term.get_mechanism(), term._asdict())) elif isinstance(term, DnsRecordTxtValueSpfModifierUnknown): terms.append((term.name, term.value)) else: raise NotImplementedError() return collections.OrderedDict([ ('Version', self.version), ('Terms', collections.OrderedDict(terms)), ]) @classmethod def _parse(cls, parsable): parser = ParserText(parsable) try: parser.parse_parsable('version', DnsRecordTxtValueSpfVersion) except InvalidValue as e: raise InvalidType from e terms = [] while parser.unparsed_length: parser.parse_separator(' ') try: parser.parse_parsable('term', DnsRecordTxtValueSpfVariantParsable) term = parser['term'] except InvalidValue: parser.parse_string_until_separator_or_end('term', ' ') term_parser = ParserText(parser['term'].encode('ascii')) term_parser.parse_parsable('value', DnsRecordTxtValueSpfModifierUnknown) term = term_parser['value'] terms.append(term) del parser['term'] return cls(version=parser['version'], terms=terms), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_parsable(self.version) for term in self.terms: composer.compose_separator(' ') composer.compose_parsable(term) return composer.composed cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/httpx/000077500000000000000000000000001524413560000257565ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/httpx/__init__.py000066400000000000000000000000431524413560000300640ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/httpx/header.py000066400000000000000000001732321524413560000275700ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines # -*- coding: utf-8 -*- import abc import collections import itertools import enum import attr import urllib3 from cryptodatahub.common.algorithm import Hash from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.types import ( Base64Data, CryptoDataParamsEnumString, convert_base64_data, convert_iterable, convert_url, convert_value_to_object, ) from cryptoparser.common.base import ( ListParamParsable, ListParsable, Serializable, StringEnumCaseInsensitiveParsable, StringEnumParsable, VariantParsable, VariantParsableExact, ) from cryptoparser.common.exception import InvalidType from cryptoparser.common.field import ( FieldParsableBase, FieldValueBase, FieldValueComponentBool, FieldValueComponentFloat, FieldValueComponentOption, FieldValueComponentString, FieldValueComponentStringBase64, FieldValueComponentStringEnum, FieldValueComponentStringEnumOption, FieldValueComponentTimeDelta, FieldValueDateTime, FieldValueMimeType, FieldValueString, FieldValueStringBySeparatorBase, FieldValueStringEnum, FieldValueStringEnumParams, FieldValueTimeDelta, FieldsCommaSeparated, FieldsJson, FieldsSemicolonSeparated, MimeTypeRegistry, NameValueVariantBase, ) from cryptoparser.common.parse import ParsableBase, ParserCRLF, ParserText, ComposerText from cryptoparser.common.utils import get_leaf_classes from .parse import ( HttpHeaderFieldValueComponent, HttpHeaderFieldValueComponentExpires, HttpHeaderFieldValueComponentMaxAge, HttpHeaderFieldValueComponentReport, HttpHeaderFieldValueComponentReportURI, ) class HttpHeaderFieldValueETag(FieldValueString): pass class HttpHeaderFieldValueAge(FieldValueTimeDelta): pass class HttpHeaderFieldValueDate(FieldValueDateTime): pass class HttpHeaderFieldValueExpires(FieldValueDateTime): pass class HttpHeaderFieldValueLastModified(FieldValueDateTime): pass class HttpHeaderFieldValueCacheControlMaxAge(HttpHeaderFieldValueComponentMaxAge): @classmethod def get_canonical_name(cls): return 'max-age' class HttpHeaderFieldValueCacheControlSMaxAge(HttpHeaderFieldValueComponentMaxAge): @classmethod def get_canonical_name(cls): return 's-maxage' class HttpHeaderFieldValueCacheControlNoCache(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'no-cache' class HttpHeaderFieldValueCacheControlNoStore(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'no-store' class HttpHeaderFieldValueCacheControlMustRevalidate(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'must-revalidate' class HttpHeaderFieldValueCacheControlProxyRevalidate(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'proxy-revalidate' class HttpHeaderFieldValueCacheControlPublic(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'public' class HttpHeaderFieldValueCacheControlPrivate(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'private' class HttpHeaderFieldValueCacheControlNoTransform(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'no-transform' @attr.s class HttpHeaderFieldValueCacheControlResponse( # pylint: disable=too-many-instance-attributes FieldsCommaSeparated ): max_age = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueCacheControlMaxAge.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueCacheControlMaxAge)), default=None ) s_maxage = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueCacheControlSMaxAge.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueCacheControlSMaxAge)), default=None ) must_revalidate = attr.ib( converter=HttpHeaderFieldValueCacheControlMustRevalidate.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueCacheControlMustRevalidate), default=False ) proxy_revalidate = attr.ib( converter=HttpHeaderFieldValueCacheControlProxyRevalidate.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueCacheControlProxyRevalidate), default=False ) no_cache = attr.ib( converter=HttpHeaderFieldValueCacheControlNoCache.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueCacheControlNoCache), default=False, ) no_store = attr.ib( converter=HttpHeaderFieldValueCacheControlNoStore.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueCacheControlNoStore), default=False, ) public = attr.ib( converter=HttpHeaderFieldValueCacheControlPublic.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueCacheControlPublic), default=False, ) private = attr.ib( converter=HttpHeaderFieldValueCacheControlPrivate.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueCacheControlPrivate), default=False, ) no_transform = attr.ib( converter=HttpHeaderFieldValueCacheControlNoTransform.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueCacheControlNoTransform), default=False, ) class HttpHeaderFieldValueComponentIncludeSubDomains(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'includeSubDomains' class HttpHeaderFieldValueNetworkErrorLoggingGroup(FieldValueComponentString): @classmethod def get_canonical_name(cls): return 'report_to' class HttpHeaderFieldValueNetworkErrorLoggingMaxAge(FieldValueComponentTimeDelta): @classmethod def get_canonical_name(cls): return 'max_age' class HttpHeaderFieldValueNetworkErrorLoggingIncludeSubdomains(FieldValueComponentBool): @classmethod def get_canonical_name(cls): return 'include_subdomains' class HttpHeaderFieldValueNetworkErrorLoggingSuccessFraction(FieldValueComponentFloat): @classmethod def get_canonical_name(cls): return 'success_fraction' class HttpHeaderFieldValueNetworkErrorLoggingFailureFraction(FieldValueComponentFloat): @classmethod def get_canonical_name(cls): return 'failure_fraction' @attr.s class HttpHeaderFieldValueNetworkErrorLogging(FieldsJson): report_to = attr.ib( converter=HttpHeaderFieldValueNetworkErrorLoggingGroup.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueNetworkErrorLoggingGroup), ) max_age = attr.ib( converter=HttpHeaderFieldValueNetworkErrorLoggingMaxAge.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueNetworkErrorLoggingMaxAge), ) include_subdomains = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueNetworkErrorLoggingIncludeSubdomains.convert), validator=attr.validators.optional(attr.validators.instance_of( HttpHeaderFieldValueNetworkErrorLoggingIncludeSubdomains )), default=None ) success_fraction = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueNetworkErrorLoggingSuccessFraction.convert), validator=attr.validators.optional(attr.validators.instance_of( HttpHeaderFieldValueNetworkErrorLoggingSuccessFraction )), default=None ) failure_fraction = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueNetworkErrorLoggingFailureFraction.convert), validator=attr.validators.optional(attr.validators.instance_of( HttpHeaderFieldValueNetworkErrorLoggingFailureFraction )), default=None ) class HttpHeaderFieldValueComponentPreload(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'preload' @attr.s class HttpHeaderFieldValueSTS(FieldsSemicolonSeparated): max_age = attr.ib( converter=HttpHeaderFieldValueComponentMaxAge.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueComponentMaxAge) ) include_subdomains = attr.ib( converter=HttpHeaderFieldValueComponentIncludeSubDomains.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueComponentIncludeSubDomains), default=False ) preload = attr.ib( converter=HttpHeaderFieldValueComponentPreload.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueComponentPreload), default=False ) @attr.s class HttpHeaderFieldValueExpectStaple(FieldsSemicolonSeparated): max_age = attr.ib( converter=HttpHeaderFieldValueComponentMaxAge.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueComponentMaxAge) ) include_subdomains = attr.ib( converter=HttpHeaderFieldValueComponentIncludeSubDomains.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueComponentIncludeSubDomains), default=False ) preload = attr.ib( converter=HttpHeaderFieldValueComponentPreload.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueComponentPreload), default=False ) report_uri = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueComponentReportURI.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueComponentReportURI)), default=None ) class ContentSecurityPolicyDirectiveType(StringEnumParsable, enum.Enum): BASE_URI = FieldValueStringEnumParams( code='base-uri', ) BLOCK_ALL_MIXED_CONTENT = FieldValueStringEnumParams( code='block-all-mixed-content', ) CHILD_SRC = FieldValueStringEnumParams( code='child-src', ) CONNECT_SRC = FieldValueStringEnumParams( code='connect-src', ) DEFAULT_SRC = FieldValueStringEnumParams( code='default-src', ) FONT_SRC = FieldValueStringEnumParams( code='font-src', ) FORM_ACTION = FieldValueStringEnumParams( code='form-action', ) FRAME_ANCESTORS = FieldValueStringEnumParams( code='frame-ancestors', ) FRAME_SRC = FieldValueStringEnumParams( code='frame-src', ) IMG_SRC = FieldValueStringEnumParams( code='img-src', ) MANIFEST_SRC = FieldValueStringEnumParams( code='manifest-src', ) MEDIA_SRC = FieldValueStringEnumParams( code='media-src', ) OBJECT_SRC = FieldValueStringEnumParams( code='object-src', ) PLUGIN_TYPES = FieldValueStringEnumParams( code='plugin-types', ) PREFETCH_SRC = FieldValueStringEnumParams( code='prefetch-src', ) REFERRER = FieldValueStringEnumParams( code='referrer', ) REPORT_SAMPLE = FieldValueStringEnumParams( code='report-sample', ) REPORT_TO = FieldValueStringEnumParams( code='report-to', ) REPORT_URI = FieldValueStringEnumParams( code='report-uri', ) REQUIRE_TRUSTED_TYPES_FOR = FieldValueStringEnumParams( code='require-trusted-types-for', ) SANDBOX = FieldValueStringEnumParams( code='sandbox', ) SCRIPT_SRC = FieldValueStringEnumParams( code='script-src', ) SCRIPT_SRC_ATTR = FieldValueStringEnumParams( code='script-src-attr', ) SCRIPT_SRC_ELEM = FieldValueStringEnumParams( code='script-src-elem', ) STYLE_SRC = FieldValueStringEnumParams( code='style-src', ) STYLE_SRC_ATTR = FieldValueStringEnumParams( code='style-src-attr', ) STYLE_SRC_ELEM = FieldValueStringEnumParams( code='style-src-elem', ) TRUSTED_TYPES = FieldValueStringEnumParams( code='trusted-types', ) UNSAFE_HASHES = FieldValueStringEnumParams( code='unsafe-hashes', ) UPGRADE_INSECURE_REQUESTS = FieldValueStringEnumParams( code='upgrade-insecure-requests', ) WEBRTC = FieldValueStringEnumParams( code='webrtc', ) WORKER_SRC = FieldValueStringEnumParams( code='worker-src', ) ContentSecurityPolicySourceType = enum.Enum('ContentSecurityPolicySourceType', 'SCHEME HOST KEYWORD NONCE HASH') class FieldHashTypeParams(CryptoDataParamsEnumString): pass class StringEnumHashParsableBase(StringEnumParsable): @classmethod def from_hash_algorithm(cls, hash_algorithm): return cls[hash_algorithm.name] @property def hash_algorithm(self): return Hash[self.name] # pylint: disable=no-member class ContentSecurityPolicySourceHashType(StringEnumHashParsableBase, enum.Enum): SHA2_256 = FieldHashTypeParams(code='sha256') SHA2_384 = FieldHashTypeParams(code='sha384') SHA2_512 = FieldHashTypeParams(code='sha512') @attr.s class ContentSecurityPolicySourceHash(ParsableBase, Serializable): hash_algorithm = attr.ib(validator=attr.validators.instance_of((Hash, str))) hash_value = attr.ib(converter=convert_base64_data(), validator=attr.validators.instance_of(Base64Data)) @classmethod def _get_hash_algorithm_enum_type(cls): return ContentSecurityPolicySourceHashType @classmethod def _parse(cls, parsable): parser = ParserText(parsable) try: parser.parse_parsable('hash_algorithm', cls._get_hash_algorithm_enum_type()) except InvalidValue as e: raise InvalidType() from e parser.parse_string_until_separator_or_end('hash_value', ' ') return cls(parser['hash_algorithm'].hash_algorithm, parser['hash_value']), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string( self._get_hash_algorithm_enum_type().from_hash_algorithm(self.hash_algorithm).value.code ) composer.compose_separator('-') composer.compose_string(str(self.hash_value)) return composer.composed @classmethod def get_type(cls): return ContentSecurityPolicySourceType.HASH @attr.s class ContentSecurityPolicySourceNonce(ParsableBase, Serializable): _PREFIX = 'nonce-' value = attr.ib(converter=convert_base64_data(), validator=attr.validators.instance_of(Base64Data)) @classmethod def _parse(cls, parsable): parser = ParserText(parsable) try: parser.parse_string('prefix', cls._PREFIX) except InvalidValue as e: raise InvalidType() from e del parser['prefix'] parser.parse_string_until_separator_or_end('value', ' ') return cls(**parser), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string(self._PREFIX) composer.compose_string(str(self.value)) return composer.composed def _asdict(self): return self.value @classmethod def get_type(cls): return ContentSecurityPolicySourceType.NONCE @attr.s class ContentSecurityPolicySourceScheme(ParsableBase, Serializable): value = attr.ib(validator=attr.validators.instance_of(str)) @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_until_separator('value', ':') parser.parse_separator(':') return cls(**parser), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string(self.value) composer.compose_separator(':') return composer.composed def _asdict(self): return self.value @classmethod def get_type(cls): return ContentSecurityPolicySourceType.SCHEME @attr.s class ContentSecurityPolicySourceHost(ParsableBase, Serializable): value = attr.ib( converter=convert_url(), validator=attr.validators.instance_of((str, urllib3.util.url.Url, )) ) @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_until_separator_or_end('value', ' ') return cls(**parser), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string(self.value) return composer.composed def _asdict(self): return self.value @classmethod def get_type(cls): return ContentSecurityPolicySourceType.HOST class ContentSecurityPolicySourceKeyword(StringEnumParsable, enum.Enum): NONE = FieldValueStringEnumParams(code='\'none\'') REPORT_SAMPLE = FieldValueStringEnumParams(code='\'report-sample\'') SELF = FieldValueStringEnumParams(code='\'self\'') STRICT_DYNAMIC = FieldValueStringEnumParams(code='\'strict-dynamic\'') UNSAFE_ALLOW_REDIRECTS = FieldValueStringEnumParams(code='\'unsafe-allow-redirects\'') UNSAFE_EVAL = FieldValueStringEnumParams(code='\'unsafe-eval\'') UNSAFE_HASHES = FieldValueStringEnumParams(code='\'unsafe-hashes\'') UNSAFE_INLINE = FieldValueStringEnumParams(code='\'unsafe-inline\'') WASM_UNSAFE_EVAL = FieldValueStringEnumParams(code='\'wasm-unsafe-eval\'') @classmethod def get_type(cls): return ContentSecurityPolicySourceType.KEYWORD @attr.s class ContentSecurityPolicyDirectiveBase(ParsableBase, Serializable): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod @abc.abstractmethod def get_type(cls): raise NotImplementedError() @classmethod def _parse_type(cls, parsable): parser = ParserText(parsable) parser.parse_parsable('type', ContentSecurityPolicyDirectiveType) if parser['type'] != cls.get_type(): raise InvalidType() return parser def _compose_type(self): composer = ComposerText() composer.compose_string(self.get_type()) return composer class ContentSecurityPolicySerializedSource(VariantParsableExact): @classmethod def _get_variants(cls): return collections.OrderedDict([ (ContentSecurityPolicySourceType.KEYWORD, [ContentSecurityPolicySourceKeyword]), (ContentSecurityPolicySourceType.NONCE, [ContentSecurityPolicySourceNonce]), (ContentSecurityPolicySourceType.HASH, [ContentSecurityPolicySourceHash]), (ContentSecurityPolicySourceType.SCHEME, [ContentSecurityPolicySourceScheme]), (ContentSecurityPolicySourceType.HOST, [ContentSecurityPolicySourceHost]), ]) @attr.s class ContentSecurityPolicyDirectiveSourceBase(ContentSecurityPolicyDirectiveBase): value = attr.ib() @classmethod @abc.abstractmethod def get_type(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_source_parser(cls): raise NotImplementedError() @value.validator def value_validator(self, _, value): self._value_validator(value) def _value_validator(self, value): source_variant_parsable = self._get_source_parser() acceptable_source_types = tuple(itertools.chain.from_iterable( source_variant_parsable._get_variants().values() # pylint: disable=protected-access )) has_invalid_source_type = any(map( lambda source: not isinstance(source, acceptable_source_types), value )) if has_invalid_source_type: raise InvalidValue(value, type(self), 'value') @classmethod def _parse(cls, parsable): parser = cls._parse_type(parsable) source_variant_parsable = cls._get_source_parser() if parser.unparsed_length: parser.parse_separator(' ') parser.parse_string_array('value', ' ', source_variant_parsable, skip_empty=True) directive_value = parser['value'] else: raise InvalidValue(parser.unparsed, cls, 'value') return cls(directive_value), parser.parsed_length def compose(self): composer = self._compose_type() if self.value: composer.compose_separator(' ') composer.compose_string_array(self.value, ' ') return composer.composed def _asdict(self): return collections.OrderedDict([ ('type', self.get_type()), ('value', [ collections.OrderedDict([('type', source.get_type()), ('value', source._asdict())]) for source in self.value ]) ]) class ContentSecurityPolicyDirectiveSerializedSourceListBase(ContentSecurityPolicyDirectiveSourceBase): @classmethod @abc.abstractmethod def get_type(cls): raise NotImplementedError() @classmethod def _get_source_parser(cls): return ContentSecurityPolicySerializedSource class ContentSecurityPolicyDirectiveChildSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.CHILD_SRC class ContentSecurityPolicyDirectiveConnectSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.CONNECT_SRC class ContentSecurityPolicyDirectiveDefaultSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.DEFAULT_SRC class ContentSecurityPolicyDirectiveFontSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.FONT_SRC class ContentSecurityPolicyDirectiveFrameSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.FRAME_SRC class ContentSecurityPolicyDirectiveImgSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.IMG_SRC class ContentSecurityPolicyDirectiveManifestSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.MANIFEST_SRC class ContentSecurityPolicyDirectiveMediaSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.MEDIA_SRC class ContentSecurityPolicyDirectiveObjectSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.OBJECT_SRC class ContentSecurityPolicyDirectivePrefetchSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.PREFETCH_SRC class ContentSecurityPolicyDirectiveScriptSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.SCRIPT_SRC class ContentSecurityPolicyDirectiveScriptSrcAttr(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.SCRIPT_SRC_ATTR class ContentSecurityPolicyDirectiveScriptSrcElem(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.SCRIPT_SRC_ELEM class ContentSecurityPolicyDirectiveStyleSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.STYLE_SRC class ContentSecurityPolicyDirectiveStyleSrcAttr(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.STYLE_SRC_ATTR class ContentSecurityPolicyDirectiveStyleSrcElem(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.STYLE_SRC_ELEM class ContentSecurityPolicyDirectiveWorkerSrc(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.WORKER_SRC class ContentSecurityPolicyDirectiveBaseUri(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.BASE_URI class ContentSecurityPolicyDirectiveFormAction(ContentSecurityPolicyDirectiveSerializedSourceListBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.FORM_ACTION class ContentSecurityPolicyFrameAncestorsSource(VariantParsableExact): @classmethod def _get_variants(cls): return collections.OrderedDict([ (ContentSecurityPolicySourceType.KEYWORD, [ContentSecurityPolicySourceKeyword]), (ContentSecurityPolicySourceType.SCHEME, [ContentSecurityPolicySourceScheme]), (ContentSecurityPolicySourceType.HOST, [ContentSecurityPolicySourceHost]), ]) class ContentSecurityPolicyDirectiveFrameAncestors(ContentSecurityPolicyDirectiveSourceBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.FRAME_ANCESTORS @classmethod def _get_source_parser(cls): return ContentSecurityPolicyFrameAncestorsSource def _value_validator(self, value): super()._value_validator(value) has_invalid_source_type = any(map( lambda source: ( isinstance(source, ContentSecurityPolicySourceKeyword) and source not in [ContentSecurityPolicySourceKeyword.SELF, ContentSecurityPolicySourceKeyword.NONE] ), value )) if has_invalid_source_type: raise InvalidValue(value, type(self), 'value') class ContentSecurityPolicyDirectiveVariant(VariantParsableExact): @classmethod def _get_variants(cls): return collections.OrderedDict([ (directive_class.get_type(), [directive_class, ]) for directive_class in get_leaf_classes(ContentSecurityPolicyDirectiveBase) ]) @attr.s class ContentSecurityPolicyDirectiveValueBase(ContentSecurityPolicyDirectiveBase): value = attr.ib() @classmethod @abc.abstractmethod def get_type(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_value_class(cls): raise NotImplementedError() @value.validator def _value_validator(self, _, value): if not isinstance(value, self._get_value_class()): raise InvalidValue(value, type(self), 'value') @classmethod def _parse(cls, parsable): parser = cls._parse_type(parsable) parser.parse_separator(' ') parser.parse_parsable('value', cls._get_value_class()) return cls(parser['value']), parser.parsed_length def compose(self): composer = self._compose_type() composer.compose_separator(' ') composer.compose_parsable(self.value) return composer.composed def _as_markdown(self, level): return self._markdown_result(self.value, level) class ContentSecurityPolicyWebRtcType(StringEnumParsable, enum.Enum): ALLOW = FieldValueStringEnumParams(code='\'allow\'') BLOCK = FieldValueStringEnumParams(code='\'block\'') class ContentSecurityPolicyDirectiveWebrtc(ContentSecurityPolicyDirectiveValueBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.WEBRTC @classmethod def _get_value_class(cls): return ContentSecurityPolicyWebRtcType class ContentSecurityPolicyReferrerPolicy(StringEnumParsable, enum.Enum): NO_REFERRER = FieldValueStringEnumParams(code='"no-referrer"') NON_WHEN_DOWNGRADE = FieldValueStringEnumParams(code='"non-when-downgrade"') ORIGIN = FieldValueStringEnumParams(code='"origin"') ORIGIN_WHEN_CROSSORIGIN = FieldValueStringEnumParams(code='"origin-when-crossorigin"') ORIGIN_WHEN_CROSS_ORIGIN = FieldValueStringEnumParams(code='"origin-when-cross-origin"') UNSAFE_URL = FieldValueStringEnumParams(code='"unsafe-url"') class ContentSecurityPolicyDirectiveReferrer(ContentSecurityPolicyDirectiveValueBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.REFERRER @classmethod def _get_value_class(cls): return ContentSecurityPolicyReferrerPolicy class ContentSecurityPolicyReportUri(FieldValueStringBySeparatorBase): @classmethod def _get_separators(cls): return ' "<>^`{|}' class ContentSecurityPolicyToken(FieldValueStringBySeparatorBase): @classmethod def _get_separators(cls): return '"(),/:;<=>?@[\\]{} \t' class ContentSecurityPolicyDirectiveListValueBase(ContentSecurityPolicyDirectiveBase): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @classmethod @abc.abstractmethod def get_type(cls): raise NotImplementedError() @classmethod def _parse_by_value_params(cls, parsable, value_type, value_name, value_min_length=None): parser = cls._parse_type(parsable) if parser.unparsed_length: parser.parse_separator(' ') parser.parse_string_array('value', ' ', value_type, skip_empty=True) if value_min_length is not None and ('value' not in parser or len(parser['value']) < value_min_length): raise InvalidValue(parser.unparsed, cls, value_name) return cls(parser['value']), parser.parsed_length def _compose(self, value_name): composer = self._compose_type() value = getattr(self, value_name) if value: composer.compose_separator(' ') composer.compose_parsable_array(value, ' ') return composer.composed @attr.s class ContentSecurityPolicyDirectiveSandbox(ContentSecurityPolicyDirectiveListValueBase): tokens = attr.ib( converter=convert_iterable(convert_value_to_object(ContentSecurityPolicyToken)), validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(ContentSecurityPolicyToken) ) ) @classmethod def _parse(cls, parsable): return cls._parse_by_value_params(parsable, ContentSecurityPolicyToken, 'tokens') def compose(self): return self._compose('tokens') @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.SANDBOX def _asdict(self): return collections.OrderedDict([ ('type', self.get_type()), ('tokens', self.tokens), ]) @attr.s class ContentSecurityPolicyDirectivePluginTypes(ContentSecurityPolicyDirectiveListValueBase): mime_types = attr.ib( converter=convert_iterable(convert_value_to_object(FieldValueMimeType)), validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(FieldValueMimeType) ), metadata={'human_readable_name': 'MIME Types'} ) @classmethod def _parse(cls, parsable): return cls._parse_by_value_params(parsable, FieldValueMimeType, 'mime_types') def compose(self): return self._compose('mime_types') @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.PLUGIN_TYPES def _asdict(self): return collections.OrderedDict([ ('type', self.get_type()), ('mime_types', self.mime_types), ]) class ContentSecurityPolicyTrustedTypeSinkGroup(StringEnumParsable, enum.Enum): SCRIPT = FieldValueStringEnumParams(code='\'script\'') @attr.s class ContentSecurityPolicyDirectiveRequireTrustedTypesFor(ContentSecurityPolicyDirectiveListValueBase): sink_groups = attr.ib( converter=convert_iterable(convert_value_to_object(ContentSecurityPolicyTrustedTypeSinkGroup)), validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(ContentSecurityPolicyTrustedTypeSinkGroup) ), ) @classmethod def _parse(cls, parsable): return cls._parse_by_value_params(parsable, ContentSecurityPolicyTrustedTypeSinkGroup, 'sink_groups') def compose(self): return self._compose('sink_groups') @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.REQUIRE_TRUSTED_TYPES_FOR def _asdict(self): return collections.OrderedDict([ ('type', self.get_type()), ('sink_groups', self.sink_groups), ]) @attr.s class ContentSecurityPolicyDirectiveReportUri(ContentSecurityPolicyDirectiveListValueBase): uri_references = attr.ib( converter=convert_iterable(convert_value_to_object(ContentSecurityPolicyReportUri)), validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(ContentSecurityPolicyReportUri), ), metadata={'human_readable_name': 'URI references'} ) @classmethod def _parse(cls, parsable): return cls._parse_by_value_params(parsable, ContentSecurityPolicyReportUri, 'uri_references', 1) def compose(self): return self._compose('uri_references') @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.REPORT_URI def _asdict(self): return collections.OrderedDict([ ('type', self.get_type()), ('uri_references', self.uri_references), ]) @attr.s class ContentSecurityPolicyDirectiveReportTo(ContentSecurityPolicyDirectiveBase): token = attr.ib( converter=convert_value_to_object(ContentSecurityPolicyToken), validator=attr.validators.instance_of(ContentSecurityPolicyToken) ) @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.REPORT_TO @classmethod def _parse(cls, parsable): parser = cls._parse_type(parsable) parser.parse_separator(' ') parser.parse_parsable('token', ContentSecurityPolicyToken) return cls(parser['token']), parser.parsed_length def compose(self): composer = self._compose_type() composer.compose_separator(' ') composer.compose_parsable(self.token) return composer.composed class ContentSecurityPolicyDirectiveNoValueBase(ContentSecurityPolicyDirectiveBase): @classmethod @abc.abstractmethod def get_type(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = cls._parse_type(parsable) return cls(), parser.parsed_length def compose(self): composer = self._compose_type() return composer.composed def _asdict(self): return collections.OrderedDict([ ('type', self.get_type()), ('value', None), ]) class ContentSecurityPolicyDirectiveBlockAllMixedContent(ContentSecurityPolicyDirectiveNoValueBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.BLOCK_ALL_MIXED_CONTENT class ContentSecurityPolicyDirectiveUpgradeInsecureRequests(ContentSecurityPolicyDirectiveNoValueBase): @classmethod def get_type(cls): return ContentSecurityPolicyDirectiveType.UPGRADE_INSECURE_REQUESTS @attr.s class HttpHeaderFieldValueContentSecurityPolicy(ParsableBase, Serializable): directives = attr.ib( validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(ContentSecurityPolicyDirectiveBase) ) ) @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_array( 'directives', separator=';', item_class=ContentSecurityPolicyDirectiveVariant, separator_spaces=' ', skip_empty=True, ) return cls(**parser), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string_array(self.directives, '; ') return composer.composed class HttpHeaderFieldValueExpectCTComponentEnforce(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'enforce' @attr.s class HttpHeaderFieldValueExpectCT(FieldsCommaSeparated): max_age = attr.ib( converter=HttpHeaderFieldValueComponentMaxAge.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueComponentMaxAge) ) enforce = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueExpectCTComponentEnforce.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueExpectCTComponentEnforce)), default=False ) report_uri = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueComponentReportURI.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueComponentReportURI)), default=None ) class HttpHeaderFieldValueContentTypeCharset(FieldValueComponentString): @classmethod def get_canonical_name(cls): return 'charset' class HttpHeaderFieldValueContentTypeBoundary(FieldValueComponentString): @classmethod def get_canonical_name(cls): return 'boundary' @attr.s class HttpHeaderFieldValueContentType(FieldsSemicolonSeparated): _MIME_TYPES_REQUIRE_BOUNDARY = (MimeTypeRegistry.MESSAGE, MimeTypeRegistry.MULTIPART) mime_type = attr.ib( converter=FieldValueMimeType.convert, validator=attr.validators.instance_of(FieldValueMimeType), metadata={'human_readable_name': 'MIME type'} ) charset = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueContentTypeCharset.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueContentTypeCharset)), default=None ) boundary = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueContentTypeBoundary.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueContentTypeBoundary)), default=None ) def __attrs_post_init__(self): if self.mime_type.registry in self._MIME_TYPES_REQUIRE_BOUNDARY and self.boundary is None: raise InvalidValue(None, type(self), 'boundary') if self.mime_type.registry not in self._MIME_TYPES_REQUIRE_BOUNDARY and self.boundary is not None: raise InvalidValue(self.boundary.value, type(self), 'boundary') class HttpHeaderXContentTypeOptions(StringEnumCaseInsensitiveParsable, enum.Enum): NOSNIFF = FieldValueStringEnumParams( code='nosniff' ) class HttpHeaderFieldValueXContentTypeOptions(FieldValueStringEnum): @classmethod def _get_value_type(cls): return HttpHeaderXContentTypeOptions class HttpHeaderPragma(StringEnumCaseInsensitiveParsable, enum.Enum): NO_CACHE = FieldValueStringEnumParams( code='no-cache' ) class HttpHeaderFieldValuePragma(FieldValueStringEnum): @classmethod def _get_value_type(cls): return HttpHeaderPragma class HttpHeaderFieldValuePublicKeyPinningPin(FieldValueComponentStringBase64): @classmethod def get_canonical_name(cls): return 'pin-sha256' @attr.s class HttpHeaderFieldValuePublicKeyPinning(FieldsSemicolonSeparated): pin_sha256 = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValuePublicKeyPinningPin.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValuePublicKeyPinningPin)), metadata={'human_readable_name': 'Pin (SHA-256)'} ) max_age = attr.ib( converter=HttpHeaderFieldValueComponentMaxAge.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueComponentMaxAge), default=None ) include_subdomains = attr.ib( converter=HttpHeaderFieldValueComponentIncludeSubDomains.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueComponentIncludeSubDomains), default=False ) report_uri = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueComponentReportURI.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueComponentReportURI)), default=None ) class HttpHeaderFieldValueServer(FieldValueString): pass class HttpHeaderFieldValueSetCookieParamDomain(FieldValueComponentString): @classmethod def get_canonical_name(cls): return 'Domain' class HttpHeaderFieldValueSetCookieParamPath(FieldValueComponentString): @classmethod def get_canonical_name(cls): return 'Path' class HttpHeaderFieldValueSetCookieParamSecure(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'Secure' class HttpHeaderFieldValueSetCookieParamHttpOnly(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'HttpOnly' class HttpHeaderSetCookieComponentSameSite(StringEnumCaseInsensitiveParsable, enum.Enum): STRICT = FieldValueStringEnumParams( code='STRICT' ) LAX = FieldValueStringEnumParams( code='Lax' ) NONE = FieldValueStringEnumParams( code='None' ) class HttpHeaderFieldValueSetCookieParamSameSite(FieldValueComponentStringEnum): @classmethod def get_canonical_name(cls): return 'SameSite' @classmethod def _get_value_type(cls): return HttpHeaderSetCookieComponentSameSite @attr.s class HttpHeaderFieldValueSetCookieParams(FieldsSemicolonSeparated): expires = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueComponentExpires.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueComponentExpires)), default=None ) max_age = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueComponentMaxAge.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueComponentMaxAge)), default=None ) domain = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueSetCookieParamDomain.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueSetCookieParamDomain)), default=None ) path = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueSetCookieParamPath.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueSetCookieParamPath)), default=None ) secure = attr.ib( converter=HttpHeaderFieldValueSetCookieParamSecure.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueSetCookieParamSecure), default=HttpHeaderFieldValueSetCookieParamSecure(False) ) http_only = attr.ib( converter=HttpHeaderFieldValueSetCookieParamHttpOnly.convert, validator=attr.validators.instance_of(HttpHeaderFieldValueSetCookieParamHttpOnly), default=HttpHeaderFieldValueSetCookieParamHttpOnly(False) ) same_site = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueSetCookieParamSameSite.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueSetCookieParamSameSite)), default=None ) @attr.s class HttpHeaderFieldValueSetCookie(FieldValueBase): # pylint: disable=too-many-instance-attributes name = attr.ib(validator=attr.validators.instance_of(str)) value = attr.ib(validator=attr.validators.instance_of(str)) expires = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueComponentExpires.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueComponentExpires)), default=None ) max_age = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueComponentMaxAge.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueComponentMaxAge)), default=None ) domain = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueSetCookieParamDomain.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueSetCookieParamDomain)), default=None ) path = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueSetCookieParamPath.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueSetCookieParamPath)), default=None ) secure = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueSetCookieParamSecure.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueSetCookieParamSecure)), default=HttpHeaderFieldValueSetCookieParamSecure(False) ) http_only = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueSetCookieParamHttpOnly.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueSetCookieParamHttpOnly)), default=HttpHeaderFieldValueSetCookieParamHttpOnly(False) ) same_site = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueSetCookieParamSameSite.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueSetCookieParamSameSite)), default=None ) @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_until_separator('name', '=') parser.parse_separator('=') parser.parse_string_until_separator_or_end('value', '; ') parser.parse_separator(' ', min_length=0) if parser.unparsed: parser.parse_separator(';') parser.parse_separator(' ', min_length=0) parser.parse_parsable('params', HttpHeaderFieldValueSetCookieParams) attributes = { 'name': parser['name'], 'value': parser['value'], } params = parser['params'] attributes.update({ name: getattr(params, name) for name in attr.fields_dict(type(params)) }) return cls(**attributes), len(parsable) def compose(self): composer = ComposerText() composer.compose_parsable(HttpHeaderFieldValueComponent(self.name, self.value)) params = {} for name, attribute in attr.fields_dict(type(self)).items(): value = getattr(self, name) if value != attribute.default and attribute.name not in ['name', 'value', ]: params[name] = getattr(self, name) if params: composer.compose_separator('; ') composer.compose_parsable(HttpHeaderFieldValueSetCookieParams(**params)) return composer.composed class HttpHeaderXFrameOptions(StringEnumCaseInsensitiveParsable, enum.Enum): DENY = FieldValueStringEnumParams( code='DENY' ) SAMEORIGIN = FieldValueStringEnumParams( code='SAMEORIGIN' ) class HttpHeaderFieldValueXFrameOptions(FieldValueStringEnum): @classmethod def _get_value_type(cls): return HttpHeaderXFrameOptions class HttpHeaderXXSSProtectionState(StringEnumParsable, enum.Enum): ENABLED = FieldValueStringEnumParams( code='1', human_readable_name='enabled' ) DISABLED = FieldValueStringEnumParams( code='0', human_readable_name='disabled' ) class HttpHeaderFieldValueXXSSProtectionState(FieldValueComponentStringEnumOption): @classmethod def _get_value_type(cls): return HttpHeaderXXSSProtectionState class HttpHeaderXXSSProtectionMode(StringEnumParsable, enum.Enum): BLOCK = FieldValueStringEnumParams( code='block' ) class HttpHeaderFieldValueXXSSProtectionMode(FieldValueComponentStringEnum): @classmethod def get_canonical_name(cls): return 'mode' @classmethod def _get_value_type(cls): return HttpHeaderXXSSProtectionMode @attr.s class HttpHeaderFieldValueXXSSProtection(FieldsSemicolonSeparated): state = attr.ib( converter=HttpHeaderFieldValueXXSSProtectionState, validator=attr.validators.instance_of(HttpHeaderFieldValueXXSSProtectionState), ) mode = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueXXSSProtectionMode.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueXXSSProtectionMode)), default=None ) report = attr.ib( converter=attr.converters.optional(HttpHeaderFieldValueComponentReport.convert), validator=attr.validators.optional(attr.validators.instance_of(HttpHeaderFieldValueComponentReport)), default=None ) class HttpHeaderReferrerPolicy(StringEnumCaseInsensitiveParsable, enum.Enum): NO_REFERRER = FieldValueStringEnumParams( code='no-referrer' ) NO_REFERRER_WHEN_DOWNGRADE = FieldValueStringEnumParams( code='no-referrer-when-downgrade' ) ORIGIN = FieldValueStringEnumParams( code='origin' ) ORIGIN_WHEN_CROSS_ORIGIN = FieldValueStringEnumParams( code='origin-when-cross-origin' ) SAME_ORIGIN = FieldValueStringEnumParams( code='same-origin' ) STRICT_ORIGIN = FieldValueStringEnumParams( code='strict-origin' ) STRICT_ORIGIN_WHEN_CROSS_ORIGIN = FieldValueStringEnumParams( code='strict-origin-when-cross-origin' ) UNSAFE_URL = FieldValueStringEnumParams( code='unsafe-url' ) class HttpHeaderFieldValueReferrerPolicy(FieldValueStringEnum): @classmethod def _get_value_type(cls): return HttpHeaderReferrerPolicy @attr.s(frozen=True) class HttpHeaderFieldNameParams(Serializable): code = attr.ib(validator=attr.validators.instance_of(str)) normalized_name = attr.ib(validator=attr.validators.instance_of(str)) def _as_markdown(self, level): return self._markdown_result(self.normalized_name, level) class HttpHeaderFieldName(StringEnumCaseInsensitiveParsable, enum.Enum): AGE = HttpHeaderFieldNameParams( code='age', normalized_name='Age' ) CACHE_CONTROL = HttpHeaderFieldNameParams( code='cache-control', normalized_name='Cache-Control' ) CONTENT_TYPE = HttpHeaderFieldNameParams( code='content-type', normalized_name='Content-Type' ) CONTENT_SECURITY_POLICY = HttpHeaderFieldNameParams( code='content-security-policy', normalized_name='Content-Security-Policy' ) CONTENT_SECURITY_POLICY_REPORT_ONLY = HttpHeaderFieldNameParams( code='content-security-policy-report-only', normalized_name='Content-Security-Policy-Report-Only' ) DATE = HttpHeaderFieldNameParams( code='date', normalized_name='Date' ) ETAG = HttpHeaderFieldNameParams( code='etag', normalized_name='ETag' ) EXPECT_CT = HttpHeaderFieldNameParams( code='expect-ct', normalized_name='Expect-CT' ) EXPECT_STAPLE = HttpHeaderFieldNameParams( code='expect-staple', normalized_name='Expect-Staple' ) EXPIRES = HttpHeaderFieldNameParams( code='expires', normalized_name='Expires' ) LAST_MODIFIED = HttpHeaderFieldNameParams( code='last-modified', normalized_name='Last-Modified' ) NETWORK_ERROR_LOGGING = HttpHeaderFieldNameParams( code='nel', normalized_name='NEL', ) PRAGMA = HttpHeaderFieldNameParams( code='pragma', normalized_name='Pragma' ) PUBLIC_KEY_PINNING = HttpHeaderFieldNameParams( code='public-key-pinning', normalized_name='Public-Key-Pinning' ) SERVER = HttpHeaderFieldNameParams( code='server', normalized_name='Server' ) SET_COOKIE = HttpHeaderFieldNameParams( code='set-cookie', normalized_name='Set-Cookie' ) REFERRER_POLICY = HttpHeaderFieldNameParams( code='referrer-policy', normalized_name='Referrer-Policy' ) STRICT_TRANSPORT_SECURITY = HttpHeaderFieldNameParams( code='strict-transport-security', normalized_name='Strict-Transport-Security' ) X_CONTENT_SECURITY_POLICY = HttpHeaderFieldNameParams( code='x-content-security-policy', normalized_name='X-Content-Security-Policy' ) X_CONTENT_TYPE_OPTIONS = HttpHeaderFieldNameParams( code='x-content-type-options', normalized_name='X-Content-Type-Options' ) X_FRAME_OPTIONS = HttpHeaderFieldNameParams( code='x-frame-options', normalized_name='X-Frame-Options' ) X_XSS_PROTECTION = HttpHeaderFieldNameParams( code='x-xss-protection', normalized_name='X-XSS-Protection' ) @classmethod def from_name(cls, name): found_items = [ item for item in cls if item.value.code == name.lower() ] if len(found_items) != 1: raise InvalidValue(name, cls, 'name') return found_items[0] class HttpHeaderFieldBase(NameValueVariantBase): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod def get_separator(cls): return ':' @classmethod def _compose_name_and_separator(cls, name): composer = cls._compose_name(name) composer.compose_separator(cls.get_separator()) return composer def _compose_name_and_value(self, name, value): composer = self._compose_name_and_separator(name) composer.compose_separator(' ') composer.compose_string(value) return composer.composed @attr.s class HttpHeaderFieldParsedBase(HttpHeaderFieldBase): @classmethod @abc.abstractmethod def get_header_field_name(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_value_class(cls): raise NotImplementedError() @classmethod def get_canonical_name(cls): return cls.get_header_field_name().value.code @classmethod def _parse(cls, parsable): parser = cls._parse_name_and_separator(parsable) parser.parse_separator(' ', min_length=0, max_length=None) parser.parse_string_until_separator('value', ['\r\n', ]) value = cls._get_value_class().parse_exact_size(parser['value'].encode('ascii')) return cls(value), parser.parsed_length def compose(self): return self._compose_name_and_value( self.get_header_field_name().value.normalized_name, bytes(self.value.compose()).decode('ascii') ) def _asdict(self): return collections.OrderedDict([ ('name', self.get_header_field_name().value.normalized_name), ('value', self.value) ]) class HttpHeaderFieldETag(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.ETAG @classmethod def _get_value_class(cls): return HttpHeaderFieldValueETag class HttpHeaderFieldAge(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.AGE @classmethod def _get_value_class(cls): return HttpHeaderFieldValueAge class HttpHeaderFieldContentType(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.CONTENT_TYPE @classmethod def _get_value_class(cls): return HttpHeaderFieldValueContentType class HttpHeaderFieldContentSecurityPolicy(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.CONTENT_SECURITY_POLICY @classmethod def _get_value_class(cls): return HttpHeaderFieldValueContentSecurityPolicy class HttpHeaderFieldXContentSecurityPolicy(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.X_CONTENT_SECURITY_POLICY @classmethod def _get_value_class(cls): return HttpHeaderFieldValueContentSecurityPolicy class HttpHeaderFieldContentSecurityPolicyReportOnly(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.CONTENT_SECURITY_POLICY_REPORT_ONLY @classmethod def _get_value_class(cls): return HttpHeaderFieldValueContentSecurityPolicy class HttpHeaderFieldCacheControlResponse(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.CACHE_CONTROL @classmethod def _get_value_class(cls): return HttpHeaderFieldValueCacheControlResponse class HttpHeaderFieldDate(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.DATE @classmethod def _get_value_class(cls): return HttpHeaderFieldValueDate class HttpHeaderFieldExpires(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.EXPIRES @classmethod def _get_value_class(cls): return HttpHeaderFieldValueExpires class HttpHeaderFieldLastModified(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.LAST_MODIFIED @classmethod def _get_value_class(cls): return HttpHeaderFieldValueLastModified class HttpHeaderFieldNetworkErrorLogging(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.NETWORK_ERROR_LOGGING @classmethod def _get_value_class(cls): return HttpHeaderFieldValueNetworkErrorLogging class HttpHeaderFieldSTS(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.STRICT_TRANSPORT_SECURITY @classmethod def _get_value_class(cls): return HttpHeaderFieldValueSTS class HttpHeaderFieldExpectCT(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.EXPECT_CT @classmethod def _get_value_class(cls): return HttpHeaderFieldValueExpectCT class HttpHeaderFieldExpectStaple(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.EXPECT_STAPLE @classmethod def _get_value_class(cls): return HttpHeaderFieldValueExpectStaple class HttpHeaderFieldPragma(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.PRAGMA @classmethod def _get_value_class(cls): return HttpHeaderFieldValuePragma class HttpHeaderFieldPublicKeyPinning(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.PUBLIC_KEY_PINNING @classmethod def _get_value_class(cls): return HttpHeaderFieldValuePublicKeyPinning class HttpHeaderFieldReferrerPolicy(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.REFERRER_POLICY @classmethod def _get_value_class(cls): return HttpHeaderFieldValueReferrerPolicy class HttpHeaderFieldSetCookie(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.SET_COOKIE @classmethod def _get_value_class(cls): return HttpHeaderFieldValueSetCookie class HttpHeaderFieldServer(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.SERVER @classmethod def _get_value_class(cls): return HttpHeaderFieldValueServer class HttpHeaderFieldXContentTypeOptions(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.X_CONTENT_TYPE_OPTIONS @classmethod def _get_value_class(cls): return HttpHeaderFieldValueXContentTypeOptions class HttpHeaderFieldXFrameOptions(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.X_FRAME_OPTIONS @classmethod def _get_value_class(cls): return HttpHeaderFieldValueXFrameOptions class HttpHeaderFieldXXSSProtection(HttpHeaderFieldParsedBase): @classmethod def get_header_field_name(cls): return HttpHeaderFieldName.X_XSS_PROTECTION @classmethod def _get_value_class(cls): return HttpHeaderFieldValueXXSSProtection class HttpHeaderFieldParsedVariant(VariantParsable): @classmethod def _get_variants(cls): return collections.OrderedDict([ (header_class.get_header_field_name(), [header_class, ]) for header_class in get_leaf_classes(HttpHeaderFieldParsedBase) ]) @attr.s class HttpHeaderFieldUnparsed(FieldParsableBase, Serializable): name = attr.ib(validator=attr.validators.instance_of(str)) value = attr.ib(validator=attr.validators.instance_of(str)) @classmethod def get_separator(cls): return ':' @classmethod def _parse(cls, parsable): parser = cls._parse_name(parsable) parser.parse_separator(cls.get_separator()) parser.parse_separator(' ', min_length=0, max_length=None) parser.parse_string_until_separator('value', '\r\n') return cls(parser['name'], parser['value']), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string_array([self.name, self.value], self.get_separator() + ' ') return composer.composed class HttpHeaderFields(ListParsable): @classmethod def get_param(cls): return ListParamParsable( item_class=HttpHeaderFieldParsedVariant, fallback_class=HttpHeaderFieldUnparsed, separator_class=ParserCRLF ) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/httpx/parse.py000066400000000000000000000022301524413560000274370ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 from cryptoparser.common.field import ( FieldValueComponentDateTime, FieldValueComponentQuotedString, FieldValueComponentString, FieldValueComponentTimeDelta, NameValuePair, ) class HttpHeaderFieldValueComponent(NameValuePair): pass class HttpHeaderFieldValueComponentExpires(FieldValueComponentDateTime): @classmethod def get_canonical_name(cls): return 'expires' class HttpHeaderFieldValueComponentMaxAge(FieldValueComponentTimeDelta): @classmethod def get_canonical_name(cls): return 'max-age' @classmethod def _check_name(cls, name): cls._check_name_insensitive(name) class HttpHeaderFieldValueComponentReport(FieldValueComponentString): @classmethod def get_canonical_name(cls): return 'report' @classmethod def _check_name(cls, name): cls._check_name_insensitive(name) class HttpHeaderFieldValueComponentReportURI(FieldValueComponentQuotedString): @classmethod def get_canonical_name(cls): return 'report-uri' @classmethod def _check_name(cls, name): cls._check_name_insensitive(name) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/httpx/version.py000066400000000000000000000012421524413560000300140ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import enum import attr from cryptoparser.common.base import Serializable @attr.s(frozen=True) class HttpVersionParams(Serializable): code = attr.ib(validator=attr.validators.instance_of(str)) name = attr.ib(validator=attr.validators.instance_of(str)) @property def identifier(self): return self.code def _asdict(self): return self.identifier def _as_markdown(self, level): return self._markdown_result(self.name, level) class HttpVersion(enum.Enum): HTTP1_0 = HttpVersionParams(code='http1_0', name='HTTP/1.0') HTTP1_1 = HttpVersionParams(code='http1_1', name='HTTP/1.1') cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ike/000077500000000000000000000000001524413560000253575ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ike/__init__.py000066400000000000000000000000431524413560000274650ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ike/common.py000066400000000000000000000340321524413560000272230ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import enum import ipaddress import typing import attr from cryptodatahub.common.exception import InvalidValue from cryptodatahub.ike.algorithm import ( Ikev1IdType, Ikev1PayloadType, Ikev2IdType, Ikev2PayloadType, ) from cryptoparser.common.exception import InvalidType from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary @attr.s class IkePayloadTypeUnknown(ParsableBase): """Wrapper for an unknown IKE payload-type code (RFC 2408 §3.10, RFC 7296 §3.2).""" code: int = attr.ib(validator=attr.validators.instance_of(int)) @classmethod def get_byte_num(cls): return 1 @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('code', cls.get_byte_num()) return cls(parser['code']), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.code, self.get_byte_num()) return composer.composed_bytes class _IkePayloadTypeFactoryBase(ParsableBase): """Yield payload-type enum member or :class:`IkePayloadTypeUnknown` for unregistered wire codes.""" @classmethod @abc.abstractmethod def _get_enum_class(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('code', 1) for enum_item in cls._get_enum_class(): if enum_item.value.code == parser['code']: return enum_item, 1 return IkePayloadTypeUnknown(parser['code']), 1 def compose(self): raise NotImplementedError() class Ikev1PayloadTypeFactory(_IkePayloadTypeFactoryBase): @classmethod def _get_enum_class(cls): return Ikev1PayloadType def compose(self): raise NotImplementedError() class Ikev2PayloadTypeFactory(_IkePayloadTypeFactoryBase): @classmethod def _get_enum_class(cls): return Ikev2PayloadType def compose(self): raise NotImplementedError() @attr.s class IkeIdentificationFqdnMixin: """Fully-qualified domain name (RFC 2407 §4.6.2.1, RFC 7296 §3.5).""" identifier: str = attr.ib(validator=attr.validators.instance_of(str)) @classmethod def get_id_type_ikev1(cls): return Ikev1IdType.FQDN @classmethod def get_id_type_ikev2(cls): return Ikev2IdType.FQDN @classmethod def _decode_identifier(cls, id_data): return id_data.decode('ascii', errors='replace') def _encode_identifier(self): return self.identifier.encode('ascii') @attr.s class IkeIdentificationUserFqdnMixin: """User fully-qualified domain name (RFC 2407 §4.6.2.1).""" identifier: str = attr.ib(validator=attr.validators.instance_of(str)) @classmethod def get_id_type_ikev1(cls): return Ikev1IdType.USER_FQDN @classmethod def get_id_type_ikev2(cls): return None @classmethod def _decode_identifier(cls, id_data): return id_data.decode('ascii', errors='replace') def _encode_identifier(self): return self.identifier.encode('ascii') @attr.s class IkeIdentificationRfc822AddrMixin: """RFC 822 email address (RFC 7296 §3.5).""" identifier: str = attr.ib(validator=attr.validators.instance_of(str)) @classmethod def get_id_type_ikev1(cls): return None @classmethod def get_id_type_ikev2(cls): return Ikev2IdType.RFC822_ADDR @classmethod def _decode_identifier(cls, id_data): return id_data.decode('ascii', errors='replace') def _encode_identifier(self): return self.identifier.encode('ascii') @attr.s class IkeIdentificationIpv4AddrMixin: """IPv4 address (RFC 2407 §4.6.2.1, RFC 7296 §3.5).""" identifier: ipaddress.IPv4Address = attr.ib( converter=ipaddress.IPv4Address, validator=attr.validators.instance_of(ipaddress.IPv4Address), ) @classmethod def get_id_type_ikev1(cls): return Ikev1IdType.IPV4_ADDR @classmethod def get_id_type_ikev2(cls): return Ikev2IdType.IPV4_ADDR @classmethod def _decode_identifier(cls, id_data): if len(id_data) != 4: raise InvalidType() return ipaddress.IPv4Address(id_data) def _encode_identifier(self): return self.identifier.packed @attr.s class IkeIdentificationIpv6AddrMixin: """IPv6 address (RFC 2407 §4.6.2.1, RFC 7296 §3.5).""" identifier: ipaddress.IPv6Address = attr.ib( converter=ipaddress.IPv6Address, validator=attr.validators.instance_of(ipaddress.IPv6Address), ) @classmethod def get_id_type_ikev1(cls): return Ikev1IdType.IPV6_ADDR @classmethod def get_id_type_ikev2(cls): return Ikev2IdType.IPV6_ADDR @classmethod def _decode_identifier(cls, id_data): if len(id_data) != 16: raise InvalidType() return ipaddress.IPv6Address(id_data) def _encode_identifier(self): return self.identifier.packed @attr.s class IkeIdentificationDerAsn1DnMixin: """DER-encoded ASN.1 X.500 Distinguished Name (RFC 2407 §4.6.2.1, RFC 7296 §3.5).""" identifier: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)), ) @classmethod def get_id_type_ikev1(cls): return Ikev1IdType.DER_ASN1_DN @classmethod def get_id_type_ikev2(cls): return Ikev2IdType.DER_ASN1_DN @classmethod def _decode_identifier(cls, id_data): return id_data def _encode_identifier(self): return bytes(self.identifier) @attr.s class IkeIdentificationKeyIdMixin: """Opaque vendor-specific key identifier (RFC 2407 §4.6.2.1, RFC 7296 §3.5).""" identifier: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)), ) @classmethod def get_id_type_ikev1(cls): return Ikev1IdType.KEY_ID @classmethod def get_id_type_ikev2(cls): return Ikev2IdType.KEY_ID @classmethod def _decode_identifier(cls, id_data): return id_data def _encode_identifier(self): return bytes(self.identifier) @attr.s class IkeIdentificationDerAsn1GnMixin: """DER-encoded ASN.1 X.500 General Name (RFC 2407 §4.6.2.1, RFC 7296 §3.5).""" identifier: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)), ) @classmethod def get_id_type_ikev1(cls): return Ikev1IdType.DER_ASN1_GN @classmethod def get_id_type_ikev2(cls): return Ikev2IdType.DER_ASN1_GN @classmethod def _decode_identifier(cls, id_data): return id_data def _encode_identifier(self): return bytes(self.identifier) @attr.s class IkeIdentificationFcNameMixin: """Fibre Channel name (RFC 4595, RFC 7296 §3.5).""" identifier: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)), ) @classmethod def get_id_type_ikev1(cls): return None @classmethod def get_id_type_ikev2(cls): return Ikev2IdType.FC_NAME @classmethod def _decode_identifier(cls, id_data): return id_data def _encode_identifier(self): return bytes(self.identifier) @attr.s class IkeIdentificationNullMixin: """NULL identification (RFC 7619).""" identifier: bytes = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)), ) @classmethod def get_id_type_ikev1(cls): return None @classmethod def get_id_type_ikev2(cls): return Ikev2IdType.NULL @classmethod def _decode_identifier(cls, id_data): return id_data def _encode_identifier(self): return bytes(self.identifier) @attr.s class DataAttributeBase(ParsableBase): """Data attribute base parser. .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ !A! Attribute Type ! AF=0 Attribute Length ! !F! ! AF=1 Attribute Value ! +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ . AF=0 Attribute Value . . AF=1 Not Transmitted . +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod @abc.abstractmethod def get_type(cls): raise NotImplementedError() @classmethod def _get_format(cls): raise NotImplementedError() @classmethod def _parse_header(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('format', 1, DataAttributeFormat) if parser['format'] != cls._get_format(): raise InvalidType() parser.parse_numeric_enum_coded('type', type(cls.get_type())) if parser['type'] != cls.get_type(): raise InvalidType() return parser def _compose_header(self): composer = ComposerBinary() composer.compose_numeric(self._get_format().value, 1) composer.compose_numeric_enum_coded(self.get_type()) return composer class DataAttributeFormat(enum.IntEnum): """Data attribute types.""" TYPE_LENGTH_VALUE = 0x00 TYPE_VALUE = 0x80 @attr.s class DataAttributeTypeValue(DataAttributeBase): """Data attribute type/value parser.""" @classmethod @abc.abstractmethod def get_type(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod def _get_format(cls): return DataAttributeFormat.TYPE_VALUE @attr.s class DataAttributeTypeValueEnumCoded(DataAttributeTypeValue): """Data attribute type/value parser where the type is an enum.""" value: typing.Any = attr.ib() @classmethod @abc.abstractmethod def get_type(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_enum_type(cls): raise NotImplementedError() @value.validator def _validate_value(self, _, value): enum_type = self._get_enum_type() if not isinstance(value, enum_type): raise InvalidValue(value, type(self), 'value') @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric_enum_coded('value', cls._get_enum_type()) return cls(value=parser['value']), parser.parsed_length def compose(self): composer = self._compose_header() composer.compose_numeric_enum_coded(self.value) return composer.composed_bytes @attr.s class DataAttributeKeyLength(DataAttributeTypeValue): """Key Length transform attribute (TV format). The Key Length attribute has the following format: 0 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ |A| Attribute Type | Key Length (in bits) | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar key_length: Key length in bits """ value: int = attr.ib(validator=attr.validators.and_( attr.validators.instance_of(int), attr.validators.ge(0), attr.validators.lt(2**64) )) @classmethod @abc.abstractmethod def get_type(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric('value', 2) return cls(value=parser['value']), parser.parsed_length def compose(self): composer = self._compose_header() composer.compose_numeric(self.value, 2) return composer.composed_bytes @attr.s class DataAttributeLength(DataAttributeBase): """Length transform attribute (TLV format). The Key Length attribute has the following format: 0 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ |A| Attribute Type | Key Length (in bits) | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar key_length: Key length in bits """ value: int = attr.ib(validator=attr.validators.and_( attr.validators.instance_of(int), attr.validators.ge(0), attr.validators.lt(2**64) )) @classmethod @abc.abstractmethod def get_type(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_size(cls): raise NotImplementedError() @classmethod def _get_format(cls): return DataAttributeFormat.TYPE_VALUE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) if parser['format'] == DataAttributeFormat.TYPE_LENGTH_VALUE: parser.parse_numeric('size', 2) parser.parse_numeric('value', parser['size']) else: parser.parse_numeric('value', cls._get_size()) return cls(value=parser['value']), parser.parsed_length def compose(self): composer = self._compose_header() if self._get_format() == DataAttributeFormat.TYPE_LENGTH_VALUE: composer.compose_numeric(self._get_size(), 2) composer.compose_numeric(self.value, self._get_size()) else: composer.compose_numeric(self.value, self._get_size()) return composer.composed_bytes cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ike/ikev1.py000066400000000000000000001331511524413560000267540ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 """IKEv1 payload parsers.""" # pylint: disable=too-many-lines import abc import collections import enum import typing import asn1crypto.x509 import attr from cryptodatahub.ike.algorithm import ( Ikev1PayloadType, Ikev1ProtocolId, Ikev1EncryptionAlgorithm, Ikev1HashAlgorithm, Ikev1TransformId, Ikev1AttributeType, Ikev1Doi, Ikev1DiffieHellmanGroup, Ikev1AuthenticationMethod, Ikev1LifeType, Ikev1NotifyType, Ikev1CertificateType, Ikev1IdType, ) from cryptodatahub.common.algorithm import IpProtocolNumber from cryptoparser.common.base import VariantParsable from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary from cryptoparser.common.exception import InvalidType, NotEnoughData from cryptoparser.ike.common import ( DataAttributeBase, DataAttributeKeyLength, DataAttributeLength, DataAttributeTypeValueEnumCoded, DataAttributeFormat, IkeIdentificationDerAsn1DnMixin, IkeIdentificationDerAsn1GnMixin, IkeIdentificationFqdnMixin, IkeIdentificationIpv4AddrMixin, IkeIdentificationIpv6AddrMixin, IkeIdentificationKeyIdMixin, IkeIdentificationUserFqdnMixin, IkePayloadTypeUnknown, Ikev1PayloadTypeFactory, ) @attr.s class Ikev1PayloadBase(ParsableBase): """Payload header parser, according to RFC2408. The generic payload header has the following structure: .. code-block:: text 0 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload | RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :cvar HEADER_SIZE: Size of the header in bytes :ivar next_payload: Type of the next payload (1 byte) """ HEADER_SIZE = 4 next_payload: typing.Optional[typing.Union[Ikev1PayloadType, IkePayloadTypeUnknown]] = attr.ib( init=False, default=Ikev1PayloadType.NONE, validator=attr.validators.optional( attr.validators.instance_of((Ikev1PayloadType, IkePayloadTypeUnknown)) ), ) @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod @abc.abstractmethod def get_payload_type(cls): raise NotImplementedError() @classmethod def _parse_header(cls, parsable): """Parse payload header from bytes. :param parsable: Bytes to parse :type parsable: bytes :return: Tuple of (parsed header, number of bytes parsed) :rtype: tuple(PayloadBase, int) :raises NotEnoughData: If there are not enough bytes to parse """ if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) # ``Ikev1PayloadTypeFactory`` yields an :class:`Ikev1PayloadType` # member when the wire code is registered, or an # :class:`IkePayloadTypeUnknown` wrapper for private-use / new # codes (IANA "IKE Payload Types" 128-255 per RFC 2408 §3.10). parser.parse_parsable('next_payload', Ikev1PayloadTypeFactory) parser.parse_numeric('reserved', 1) # Skip reserved byte parser.parse_numeric('payload_length', 2) if parser.unparsed_length < parser['payload_length'] - cls.HEADER_SIZE: raise NotEnoughData(parser['payload_length'] - cls.HEADER_SIZE - parser.unparsed_length) return parser def compose_header(self, payload_length): """Compose payload header to bytes. :return: Composed header bytes :rtype: bytes """ composer = ComposerBinary() if isinstance(self.next_payload, IkePayloadTypeUnknown): composer.compose_parsable(self.next_payload) else: composer.compose_numeric_enum_coded(self.next_payload) composer.compose_numeric(0, 1) # Reserved byte composer.compose_numeric(self.HEADER_SIZE + payload_length, 2) return composer @attr.s class Ikev1PayloadUnparsed(Ikev1PayloadBase): """Opaque IKEv1 payload wrapper for unknown payload type codes (RFC 2408 §3.10).""" payload_data: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)), ) payload_type: typing.Union[Ikev1PayloadType, IkePayloadTypeUnknown, None] = attr.ib( default=None, validator=attr.validators.optional( attr.validators.instance_of((Ikev1PayloadType, IkePayloadTypeUnknown)) ), ) # pylint: disable=invalid-overridden-method,arguments-differ def get_payload_type(self): return self.payload_type @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('payload_data', parser['payload_length'] - cls.HEADER_SIZE) payload = cls(payload_data=parser['payload_data']) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer = self.compose_header(len(self.payload_data)) composer.compose_raw(self.payload_data) return composer.composed_bytes @attr.s class Ikev1AttributeAuthenticationMethod(DataAttributeTypeValueEnumCoded): """Authentication Method transform attribute (TLV format).""" @classmethod def get_type(cls): return Ikev1AttributeType.AUTHENTICATION_METHOD @classmethod def _get_enum_type(cls): return Ikev1AuthenticationMethod @attr.s class Ikev1AttributeDiffieHellmanGroup(DataAttributeTypeValueEnumCoded): """Diffie-Hellman Group transform attribute (TLV format).""" @classmethod def get_type(cls): return Ikev1AttributeType.GROUP_DESCRIPTION @classmethod def _get_enum_type(cls): return Ikev1DiffieHellmanGroup @attr.s class Ikev1AttributeKeyLength(DataAttributeKeyLength): """Key Length transform attribute (TV format).""" @classmethod def get_type(cls): return Ikev1AttributeType.KEY_LENGTH @attr.s class Ikev1AttributeEncryptionAlgorithm(DataAttributeTypeValueEnumCoded): """Encryption Algorithm transform attribute (TLV format).""" @classmethod def get_type(cls): return Ikev1AttributeType.ENCRYPTION_ALGORITHM @classmethod def _get_enum_type(cls): return Ikev1EncryptionAlgorithm @attr.s class Ikev1AttributeHashAlgorithm(DataAttributeTypeValueEnumCoded): """Hash Algorithm transform attribute (TLV format).""" @classmethod def get_type(cls): return Ikev1AttributeType.HASH_ALGORITHM @classmethod def _get_enum_type(cls): return Ikev1HashAlgorithm @attr.s class Ikev1AttributeLifeType(DataAttributeTypeValueEnumCoded): """Life Type transform attribute (TLV format).""" @classmethod def get_type(cls): return Ikev1AttributeType.LIFE_TYPE @classmethod def _get_enum_type(cls): return Ikev1LifeType @attr.s class Ikev1AttributeLifeDuration(DataAttributeLength): """Lifetime transform attribute (TLV format).""" @classmethod def _get_format(cls): return DataAttributeFormat.TYPE_LENGTH_VALUE @classmethod def get_type(cls): return Ikev1AttributeType.LIFE_DURATION @classmethod def _get_size(cls): return 4 @attr.s class Ikev1PayloadTransform(Ikev1PayloadBase): """Transform Payload parser. The Transform Payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload | RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Transform # | Transform-Id | RESERVED2 | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ SA Attributes ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar transform_number: Transform number (1 byte) :ivar transform_id: Transform ID (1 byte) :ivar sa_attributes: Security Association attributes (variable length) """ transform_id: Ikev1TransformId = attr.ib(validator=attr.validators.instance_of(Ikev1TransformId)) attributes: list[DataAttributeBase] = attr.ib( validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(DataAttributeBase), ) ) transform_number: typing.Optional[int] = attr.ib( init=False, default=None, validator=attr.validators.optional(attr.validators.instance_of(int)) ) @classmethod def get_payload_type(cls): return Ikev1PayloadType.TRANSFORM def get_attribute_by_type(self, attribute_type: Ikev1AttributeType) -> DataAttributeBase: for attribute in self.attributes: if attribute.get_type() == attribute_type: return attribute raise KeyError(attribute_type) @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric('transform_number', 1) parser.parse_numeric_enum_coded('transform_id', Ikev1TransformId) parser.parse_numeric('reserved2', 2) # Skip reserved2 field attributes = [] parser_attributes = ParserBinary(parsable[parser.parsed_length:parser['payload_length']]) while parser_attributes.unparsed_length > 0: parser_attributes.parse_parsable('attribute', Ikev1AttributeVariantServer) attribute = parser_attributes['attribute'] attributes.append(attribute) payload = cls( transform_id=parser['transform_id'], attributes=attributes ) payload.next_payload = parser['next_payload'] payload.transform_number = parser['transform_number'] return payload, parser.parsed_length + parser_attributes.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_numeric(self.transform_number, 1) composer_payload.compose_numeric_enum_coded(self.transform_id) composer_payload.compose_numeric(0, 2) # Reserved2 field for attribute in self.attributes: composer_payload.compose_parsable(attribute) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev1PayloadProposal(Ikev1PayloadBase): """Proposal Payload parser. The Proposal Payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload | RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Proposal # | Protocol-Id | SPI Size |# of Transforms| +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | SPI (variable) | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar next_payload: Next payload type (1 byte) :ivar protocol_id: Protocol ID (1 byte) :ivar spi_size: Size of SPI in bytes (1 byte) :ivar transform_count: Number of transforms (1 byte) :ivar spi: Security Parameter Index (variable length) """ protocol_id: Ikev1ProtocolId = attr.ib(validator=attr.validators.instance_of(Ikev1ProtocolId)) transforms: list[Ikev1PayloadTransform] = attr.ib(validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(Ikev1PayloadTransform), )) spi: bytes = attr.ib(default=b'', converter=bytes, validator=attr.validators.instance_of(bytes)) next_payload: typing.Optional[Ikev1PayloadType] = attr.ib( init=False, default=None, validator=attr.validators.optional(attr.validators.instance_of(Ikev1PayloadType)) ) proposal_number: typing.Optional[int] = attr.ib( init=False, default=None, validator=attr.validators.optional(attr.validators.instance_of(int)) ) @classmethod def get_payload_type(cls): return Ikev1PayloadType.PROPOSAL @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric('proposal_number', 1) parser.parse_numeric_enum_coded('protocol_id', Ikev1ProtocolId) parser.parse_numeric('spi_size', 1) parser.parse_numeric('transform_count', 1) parser.parse_raw('spi', parser['spi_size']) transforms = [] for _ in range(parser['transform_count']): parser.parse_parsable('transform', Ikev1PayloadTransform) transforms.append(parser['transform']) payload = cls( protocol_id=parser['protocol_id'], spi=parser['spi'], transforms=transforms, ) payload.next_payload = parser['next_payload'] payload.proposal_number = parser['proposal_number'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_numeric(self.proposal_number, 1) composer_payload.compose_numeric_enum_coded(self.protocol_id) composer_payload.compose_numeric(len(self.spi), 1) composer_payload.compose_numeric(len(self.transforms), 1) composer_payload.compose_raw(self.spi) for transform_number, transform in enumerate(self.transforms): transform.next_payload = ( transform.get_payload_type() if transform_number < len(self.transforms) - 1 else Ikev1PayloadType.NONE ) transform.transform_number = transform_number + 1 composer_payload.compose_parsable(transform) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes class Ikev1Situation(enum.IntFlag): """IKEv1 situation.""" SIT_IDENTITY_ONLY = 1 << 0 SIT_SECRECY = 1 << 1 SIT_INTEGRITY = 1 << 2 @attr.s class Ikev1PayloadSecurityAssociation(Ikev1PayloadBase): """Security Association payload parser. The Security Association payload has the following format: .. code-block:: text 0 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload | RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Domain of Interpretation (DOI) | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Situation ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar doi: Domain of Interpretation (4 bytes) :ivar situation: Situation field (variable length) """ doi: Ikev1Doi = attr.ib(validator=attr.validators.instance_of(Ikev1Doi)) situation: list[Ikev1Situation] = attr.ib(validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(Ikev1Situation), )) proposals: list[Ikev1PayloadProposal] = attr.ib(validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(Ikev1PayloadProposal), )) @classmethod def get_payload_type(cls): return Ikev1PayloadType.SECURITY_ASSOCIATION @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric_enum_coded('doi', Ikev1Doi) parser.parse_numeric_flags('situation', 4, Ikev1Situation) proposals = [] parser_proposal = ParserBinary(parsable[parser.parsed_length:parser['payload_length']]) while parser_proposal.unparsed_length > 0: parser_proposal.parse_parsable('proposal', Ikev1PayloadProposal) proposal = parser_proposal['proposal'] proposals.append(proposal) payload = cls( doi=parser['doi'], situation=parser['situation'], proposals=proposals ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length + parser_proposal.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_numeric_enum_coded(self.doi) composer_payload.compose_numeric_flags(self.situation, 4) for proposal_number, proposal in enumerate(self.proposals): proposal.next_payload = ( proposal.get_payload_type() if proposal_number < len(self.proposals) - 1 else Ikev1PayloadType.NONE ) proposal.proposal_number = proposal_number + 1 composer_payload.compose_parsable(proposal) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev1PayloadKeyExchange(Ikev1PayloadBase): """Key Exchange payload parser. The Key Exchange payload has the following format: .. code-block:: text 0 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload | RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Key Exchange ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar next_payload: Next payload type (1 byte) :ivar key_exchange_data: Key exchange data (variable length) """ key_exchange_data: bytes = attr.ib(converter=bytes, validator=attr.validators.instance_of(bytes)) @classmethod def get_payload_type(cls): return Ikev1PayloadType.KEY_EXCHANGE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('key_exchange_data', parser['payload_length'] - cls.HEADER_SIZE) payload = cls(key_exchange_data=parser['key_exchange_data']) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_raw(self.key_exchange_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev1PayloadNonce(Ikev1PayloadBase): """Nonce payload parser. The Nonce payload has the following format: .. code-block:: text 0 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Nonce ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar nonce_data: Nonce data (variable length) """ nonce_data: bytes = attr.ib(converter=bytes, validator=attr.validators.instance_of(bytes)) @classmethod def get_payload_type(cls): return Ikev1PayloadType.NONCE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('nonce_data', parser['payload_length'] - cls.HEADER_SIZE) payload = cls(nonce_data=parser['nonce_data']) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_raw(self.nonce_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev1PayloadHash(Ikev1PayloadBase): """Hash payload parser. The Hash payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload | RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Hash Data ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar hash_data: Hash data (variable length) """ hash_data: bytes = attr.ib(converter=bytes, validator=attr.validators.instance_of(bytes)) @classmethod def get_payload_type(cls): return Ikev1PayloadType.HASH @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('hash_data', parser['payload_length'] - cls.HEADER_SIZE) payload = cls(hash_data=parser['hash_data']) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_raw(self.hash_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev1PayloadDoiProtocolSpiBase(Ikev1PayloadBase): """Base class for IKEv1 payloads with DOI, Protocol-Id, and SPI Size fields. Shared structure: .. code-block:: text +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Domain of Interpretation (DOI) | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Protocol-Id | SPI Size | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar doi: Domain of Interpretation (4 bytes) :ivar protocol_id: Protocol ID (1 byte) :ivar spi_size: Size of SPI in bytes (1 byte) """ DOI_PROTOCOL_SPI_SIZE = 6 doi: Ikev1Doi = attr.ib(validator=attr.validators.instance_of(Ikev1Doi)) protocol_id: Ikev1ProtocolId = attr.ib(validator=attr.validators.instance_of(Ikev1ProtocolId)) spi_size: int = attr.ib(validator=attr.validators.instance_of(int)) @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod def _parse_header(cls, parsable): parser = super()._parse_header(parsable) if parser.unparsed_length < cls.DOI_PROTOCOL_SPI_SIZE: raise NotEnoughData(cls.DOI_PROTOCOL_SPI_SIZE - parser.unparsed_length) parser.parse_numeric_enum_coded('doi', Ikev1Doi) parser.parse_numeric_enum_coded('protocol_id', Ikev1ProtocolId) parser.parse_numeric('spi_size', 1) return parser def _compose_doi_protocol_spi(self, composer): """Compose DOI, Protocol-Id, and SPI Size to composer.""" composer.compose_numeric_enum_coded(self.doi) composer.compose_numeric_enum_coded(self.protocol_id) composer.compose_numeric(self.spi_size, 1) @attr.s class Ikev1PayloadNotification(Ikev1PayloadDoiProtocolSpiBase): """Notification payload parser. The Notification payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ! Next Payload ! RESERVED ! Payload Length ! +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ! Domain of Interpretation (DOI) ! +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ! Protocol-ID ! SPI Size ! Notify Message Type ! +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ! ! ~ Security Parameter Index (SPI) ~ ! ! +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ! ! ~ Notification Data ~ ! ! +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar notify_message_type: Notify message type (1 byte) :ivar spi: Security Parameter Index (variable length) :ivar notification_data: Notification data (variable length) """ notify_type: Ikev1NotifyType = attr.ib(validator=attr.validators.instance_of(Ikev1NotifyType)) spi: bytes = attr.ib(converter=bytes, validator=attr.validators.instance_of(bytes)) notification_data: bytes = attr.ib(converter=bytes, validator=attr.validators.instance_of(bytes)) @classmethod def get_payload_type(cls): return Ikev1PayloadType.NOTIFICATION @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric_enum_coded('notify_type', Ikev1NotifyType) parser.parse_raw('spi', parser['spi_size']) parser.parse_raw('notification_data', parser['payload_length'] - parser.parsed_length) payload = cls( doi=parser['doi'], protocol_id=parser['protocol_id'], spi_size=parser['spi_size'], notify_type=parser['notify_type'], spi=parser['spi'], notification_data=parser['notification_data'] ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() self._compose_doi_protocol_spi(composer_payload) composer_payload.compose_numeric_enum_coded(self.notify_type) composer_payload.compose_raw(self.spi) composer_payload.compose_raw(self.notification_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev1PayloadDelete(Ikev1PayloadDoiProtocolSpiBase): """Delete payload parser. The Delete payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload | RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Domain of Interpretation (DOI) | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Protocol-Id | SPI Size | # of SPIs | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Security Parameter Index(es) (SPI) ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar spis: Security Parameter Index(es) (variable length, list of spi_size bytes each) """ spis: typing.Sequence[typing.Union[bytes, bytearray]] = attr.ib( validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of((bytes, bytearray)), ) ) @classmethod def get_payload_type(cls): return Ikev1PayloadType.DELETE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric('spi_count', 2) spis = [] for _ in range(parser['spi_count']): parser.parse_raw('spi', parser['spi_size']) spis.append(parser['spi']) payload = cls( doi=parser['doi'], protocol_id=parser['protocol_id'], spi_size=parser['spi_size'], spis=spis, ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() self._compose_doi_protocol_spi(composer_payload) composer_payload.compose_numeric(len(self.spis), 2) for spi in self.spis: composer_payload.compose_raw(spi) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev1PayloadVendorId(Ikev1PayloadBase): """Vendor ID payload parser. The Vendor ID payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ! Next Payload !C! RESERVED ! Payload Length ! +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ! ! ~ Vendor ID (VID) ~ ! ! +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar vendor_id: Vendor ID (variable length) """ vendor_id: bytes = attr.ib(converter=bytes, validator=attr.validators.instance_of(bytes)) @classmethod def get_payload_type(cls): return Ikev1PayloadType.VENDOR_ID @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('vendor_id', parser['payload_length'] - cls.HEADER_SIZE) payload = cls(vendor_id=parser['vendor_id']) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_raw(self.vendor_id) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev1PayloadCertificateBase(Ikev1PayloadBase): """Certificate payload base (RFC 2408 §3.9, §3.10).""" cert_encoding: Ikev1CertificateType = attr.ib(validator=attr.validators.instance_of(Ikev1CertificateType)) @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod def _parse_header(cls, parsable): parser = super()._parse_header(parsable) parser.parse_numeric_enum_coded('cert_encoding', Ikev1CertificateType) return parser def _compose_cert_encoding(self, composer): composer.compose_numeric_enum_coded(self.cert_encoding) @attr.s class Ikev1PayloadCertificateRequest(Ikev1PayloadCertificateBase): """Certificate Request payload parser (RFC 2408 §3.10). .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload | RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Cert Encoding | | +-+-+-+-+-+-+-+-+ + ~ Certification Authority ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ certification_authority: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)) ) @classmethod def get_payload_type(cls): return Ikev1PayloadType.CERTIFICATE_REQUEST @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('certification_authority', parser['payload_length'] - cls.HEADER_SIZE - 1) payload = cls( cert_encoding=parser['cert_encoding'], certification_authority=parser['certification_authority'] ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() self._compose_cert_encoding(composer_payload) composer_payload.compose_raw(self.certification_authority) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes def get_distinguished_name(self) -> typing.Optional[collections.OrderedDict]: """Parse the Certification Authority field as X.501 Distinguished Name (RFC 2408 §3.10).""" raw = bytes(self.certification_authority) if not raw: return None try: return asn1crypto.x509.Name.load(raw).native except (ValueError, TypeError): return None @attr.s class Ikev1PayloadCertificate(Ikev1PayloadCertificateBase): """Certificate payload parser (RFC 2408 §3.9). .. code-block:: text 0 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Cert Encoding | | +-+-+-+-+-+-+-+-+ + ~ Certificate Data ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ certificate_data: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)) ) @classmethod def get_payload_type(cls): return Ikev1PayloadType.CERTIFICATE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('certificate_data', parser['payload_length'] - cls.HEADER_SIZE - 1) payload = cls( cert_encoding=parser['cert_encoding'], certificate_data=parser['certificate_data'] ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() self._compose_cert_encoding(composer_payload) composer_payload.compose_raw(self.certificate_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev1PayloadIdentificationBase(Ikev1PayloadBase): """Identification payload parser (RFC 2408 §3.8, RFC 2407 §4.6.2). .. code-block:: text 0 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload | RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | ID Type | Protocol ID | Port | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Identification Data ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ protocol_id: IpProtocolNumber = attr.ib(validator=attr.validators.instance_of(IpProtocolNumber)) port: int = attr.ib(validator=attr.validators.instance_of(int)) @classmethod def get_payload_type(cls): return Ikev1PayloadType.IDENTIFICATION # The following abstract methods are concretely provided by the # ``IkeIdentification*Mixin`` mixins listed first in the concrete # subclass MRO; they appear here so static-analysis tools see them # as part of the contract. @classmethod @abc.abstractmethod def get_id_type_ikev1(cls) -> Ikev1IdType: raise NotImplementedError() @classmethod @abc.abstractmethod def _decode_identifier(cls, id_data: bytes): raise NotImplementedError() @abc.abstractmethod def _encode_identifier(self) -> bytes: raise NotImplementedError() @property def id_type(self) -> Ikev1IdType: """Wire-format ``ID Type`` field — fixed per concrete subclass via the :meth:`get_id_type_ikev1` mixin classmethod.""" return self.get_id_type_ikev1() @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric_enum_coded('id_type', Ikev1IdType) if parser['id_type'] != cls.get_id_type_ikev1(): raise InvalidType() parser.parse_numeric_enum_coded('protocol_id', IpProtocolNumber) parser.parse_numeric('port', 2) id_data_length = parser['payload_length'] - cls.HEADER_SIZE - 4 if id_data_length < 0: raise NotEnoughData(bytes_needed=-id_data_length) parser.parse_raw('id_data', id_data_length) payload = cls( protocol_id=parser['protocol_id'], port=parser['port'], identifier=cls._decode_identifier(bytes(parser['id_data'])), ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_numeric_enum_coded(self.id_type) composer_payload.compose_numeric_enum_coded(self.protocol_id) composer_payload.compose_numeric(self.port, 2) composer_payload.compose_raw(self._encode_identifier()) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev1PayloadIdentificationFqdn( IkeIdentificationFqdnMixin, Ikev1PayloadIdentificationBase, ): """IKEv1 IDENTIFICATION with fully-qualified domain name (RFC 2407 §4.6.2.1).""" @attr.s class Ikev1PayloadIdentificationUserFqdn( IkeIdentificationUserFqdnMixin, Ikev1PayloadIdentificationBase, ): """IKEv1 IDENTIFICATION with user fully-qualified domain name (RFC 2407 §4.6.2.1).""" @attr.s class Ikev1PayloadIdentificationIpv4Addr( IkeIdentificationIpv4AddrMixin, Ikev1PayloadIdentificationBase, ): """IKEv1 IDENTIFICATION with IPv4 address (RFC 2407 §4.6.2.1).""" @attr.s class Ikev1PayloadIdentificationIpv6Addr( IkeIdentificationIpv6AddrMixin, Ikev1PayloadIdentificationBase, ): """IKEv1 IDENTIFICATION with IPv6 address (RFC 2407 §4.6.2.1).""" @attr.s class Ikev1PayloadIdentificationDerAsn1Dn( IkeIdentificationDerAsn1DnMixin, Ikev1PayloadIdentificationBase, ): """IKEv1 IDENTIFICATION with DER-encoded ASN.1 X.500 Distinguished Name (RFC 2407 §4.6.2.1).""" @attr.s class Ikev1PayloadIdentificationKeyId( IkeIdentificationKeyIdMixin, Ikev1PayloadIdentificationBase, ): """IKEv1 IDENTIFICATION with opaque vendor-specific key identifier (RFC 2407 §4.6.2.1).""" @attr.s class Ikev1PayloadIdentificationDerAsn1Gn( IkeIdentificationDerAsn1GnMixin, Ikev1PayloadIdentificationBase, ): """IKEv1 IDENTIFICATION with DER-encoded ASN.1 X.500 General Name (RFC 2407 §4.6.2.1).""" class Ikev1PayloadIdentificationVariant(VariantParsable): """Variant dispatcher for IKEv1 Identification payload (RFC 2408 §3.8, RFC 2407 §4.6.2).""" @classmethod def _get_variants(cls): return collections.OrderedDict([ (subclass.get_id_type_ikev1(), [subclass]) for subclass in ( Ikev1PayloadIdentificationFqdn, Ikev1PayloadIdentificationUserFqdn, Ikev1PayloadIdentificationIpv4Addr, Ikev1PayloadIdentificationIpv6Addr, Ikev1PayloadIdentificationDerAsn1Dn, Ikev1PayloadIdentificationDerAsn1Gn, Ikev1PayloadIdentificationKeyId, ) ]) @attr.s class Ikev1PayloadSignature(Ikev1PayloadBase): """Signature payload parser (RFC 2408 §3.12). .. code-block:: text 0 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload | RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Signature Data ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ signature_data: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)) ) @classmethod def get_payload_type(cls): return Ikev1PayloadType.SIGNATURE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('signature_data', parser['payload_length'] - cls.HEADER_SIZE) payload = cls(signature_data=parser['signature_data']) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_raw(self.signature_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes IKEV1_PAYLOAD_CLASSES_BY_TYPE = { Ikev1PayloadType.SECURITY_ASSOCIATION: Ikev1PayloadSecurityAssociation, Ikev1PayloadType.KEY_EXCHANGE: Ikev1PayloadKeyExchange, Ikev1PayloadType.IDENTIFICATION: Ikev1PayloadIdentificationVariant, Ikev1PayloadType.HASH: Ikev1PayloadHash, Ikev1PayloadType.SIGNATURE: Ikev1PayloadSignature, Ikev1PayloadType.NONCE: Ikev1PayloadNonce, Ikev1PayloadType.NOTIFICATION: Ikev1PayloadNotification, Ikev1PayloadType.DELETE: Ikev1PayloadDelete, Ikev1PayloadType.VENDOR_ID: Ikev1PayloadVendorId, Ikev1PayloadType.CERTIFICATE: Ikev1PayloadCertificate, Ikev1PayloadType.CERTIFICATE_REQUEST: Ikev1PayloadCertificateRequest, } class Ikev1AttributeVariantBase(VariantParsable): @classmethod @abc.abstractmethod def get_parsed_extensions(cls): raise NotImplementedError() @classmethod def _get_variants(cls): variants = cls.get_parsed_extensions() # variants.update([ # (extension_type, (Ikev1AttributeUnparsed, )) # for extension_type in Ikev1AttributeType # if extension_type not in variants # ]) return variants class Ikev1AttributeVariantServer(Ikev1AttributeVariantBase): @classmethod def get_parsed_extensions(cls): return collections.OrderedDict([ (Ikev1AttributeType.ENCRYPTION_ALGORITHM, [Ikev1AttributeEncryptionAlgorithm, ]), (Ikev1AttributeType.HASH_ALGORITHM, [Ikev1AttributeHashAlgorithm, ]), (Ikev1AttributeType.LIFE_TYPE, [Ikev1AttributeLifeType, ]), (Ikev1AttributeType.KEY_LENGTH, [Ikev1AttributeKeyLength, ]), (Ikev1AttributeType.GROUP_DESCRIPTION, [Ikev1AttributeDiffieHellmanGroup, ]), (Ikev1AttributeType.AUTHENTICATION_METHOD, [Ikev1AttributeAuthenticationMethod, ]), (Ikev1AttributeType.LIFE_DURATION, [Ikev1AttributeLifeDuration, ]), ]) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ike/ikev2.py000066400000000000000000002074771524413560000267720ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 """IKEv2 message parsers.""" # pylint: disable=too-many-lines import abc import collections import enum import typing import asn1crypto.algos import attr from cryptodatahub.common.algorithm import Signature from cryptodatahub.common.exception import InvalidValue from cryptodatahub.ike.algorithm import ( Ikev2TransformAttributeType, Ikev2PayloadType, Ikev2AuthenticationMethod, Ikev2DiffieHellmanGroup, Ikev2EncryptionAlgorithm, Ikev2HashAlgorithm, Ikev2IdType, Ikev2IntegrityAlgorithm, Ikev2NotifyType, Ikev2ProtocolId, Ikev2PseudorandomFunction, Ikev2TransformType, Ikev2CertificateType, ) from cryptoparser.common.base import TwoByteEnumParsable, VariantParsable from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary from cryptoparser.common.exception import NotEnoughData, TooMuchData, InvalidType from cryptoparser.ike.common import ( DataAttributeLength, DataAttributeTypeValue, DataAttributeFormat, IkeIdentificationDerAsn1DnMixin, IkeIdentificationDerAsn1GnMixin, IkeIdentificationFcNameMixin, IkeIdentificationFqdnMixin, IkeIdentificationIpv4AddrMixin, IkeIdentificationIpv6AddrMixin, IkeIdentificationKeyIdMixin, IkeIdentificationNullMixin, IkeIdentificationRfc822AddrMixin, IkePayloadTypeUnknown, Ikev2PayloadTypeFactory, ) class Ikev2ProposalFlags(enum.IntFlag): """Proposal flags.""" LAST_SUBSTRUCT = 0x80 class Ikev2PayloadFlags(enum.IntFlag): """Payload flags.""" CRITICAL = 0x80 # Critical bit flag @attr.s class Ikev2PayloadBase(ParsableBase): """Payload header parser, according to RFC7296. The generic payload header has the following structure: .. code-block:: text 0 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :cvar HEADER_SIZE: Size of the header in bytes :ivar next_payload: Type of the next payload (1 byte) :ivar critical: Critical bit flag (1 bit) """ HEADER_SIZE = 4 flags: set[Ikev2PayloadFlags] = attr.ib( validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(Ikev2PayloadFlags), ) ) next_payload: typing.Optional[typing.Union[Ikev2PayloadType, IkePayloadTypeUnknown]] = attr.ib( init=False, default=None, validator=attr.validators.optional( attr.validators.instance_of((Ikev2PayloadType, IkePayloadTypeUnknown)) ), ) @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod @abc.abstractmethod def get_payload_type(cls): raise NotImplementedError() @classmethod def _parse_header(cls, parsable): """Parse payload header from bytes. :param parsable: Bytes to parse :type parsable: bytes :return: Tuple of (parsed header, number of bytes parsed) :rtype: tuple(PayloadBase, int) :raises NotEnoughData: If there are not enough bytes to parse """ if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) # ``Ikev2PayloadTypeFactory`` yields an :class:`Ikev2PayloadType` # member when the wire code is registered, or an # :class:`IkePayloadTypeUnknown` wrapper for private-use / new # codes (IANA "IKEv2 Payload Types" 128-255). parser.parse_parsable('next_payload', Ikev2PayloadTypeFactory) parser.parse_numeric_flags('flags', 1, Ikev2PayloadFlags) parser.parse_numeric('payload_length', 2) if parser.unparsed_length < parser['payload_length'] - cls.HEADER_SIZE: raise NotEnoughData(parser['payload_length'] - cls.HEADER_SIZE - parser.unparsed_length) return parser def compose_header(self, payload_length): """Compose payload header to bytes. :return: Composed header bytes :rtype: bytes """ assert self.next_payload is not None composer = ComposerBinary() if isinstance(self.next_payload, IkePayloadTypeUnknown): composer.compose_parsable(self.next_payload) else: composer.compose_numeric(self.next_payload.value.code, 1) composer.compose_numeric_flags(self.flags, 1) composer.compose_numeric(self.HEADER_SIZE + payload_length, 2) return composer @attr.s class Ikev2PayloadUnparsed(Ikev2PayloadBase): """Opaque IKEv2 payload wrapper for unknown payload type codes (RFC 7296 §3.2).""" payload_data: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)), ) payload_type: typing.Union[Ikev2PayloadType, IkePayloadTypeUnknown, None] = attr.ib( default=None, validator=attr.validators.optional( attr.validators.instance_of((Ikev2PayloadType, IkePayloadTypeUnknown)) ), ) # pylint: disable=invalid-overridden-method,arguments-differ def get_payload_type(self): return self.payload_type @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('payload_data', parser['payload_length'] - cls.HEADER_SIZE) payload = cls(flags=set(), payload_data=parser['payload_data']) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer = self.compose_header(len(self.payload_data)) composer.compose_raw(self.payload_data) return composer.composed_bytes class TransformNextPayload(enum.IntEnum): """Transform next payload.""" LAST = 0x00 MORE = 0x03 @attr.s class Transform(ParsableBase): """Transform payload parser. The transform payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next payload | RESERVED | Transform Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ |Transform Type | RESERVED | Transform ID | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Transform Attributes ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar transform_id: Transform ID :ivar next_payload: Next payload """ HEADER_SIZE = 8 transform_id: typing.Any = attr.ib() next_payload: typing.Optional[TransformNextPayload] = attr.ib( init=False, default=None, validator=attr.validators.optional(attr.validators.instance_of(TransformNextPayload)) ) @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod @abc.abstractmethod def get_transform_type(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_transform_id_class(cls): raise NotImplementedError() @transform_id.validator def _validate_transform_id(self, _, value): transform_id_class = self._get_transform_id_class() if not isinstance(value, transform_id_class): raise InvalidValue(value, type(self), 'value') @classmethod def _parse_header(cls, parsable): """Parse transform from bytes. :param parsable: Bytes to parse :type parsable: bytes :return: Tuple of (parsed transform, number of bytes parsed) :rtype: tuple(Transform, int) :raises NotEnoughData: If there are not enough bytes to parse """ if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric('next_payload', 1, TransformNextPayload) parser.parse_numeric('reserved', 1) parser.parse_numeric('transform_length', 2) parser.parse_numeric_enum_coded('transform_type', Ikev2TransformType) if parser['transform_type'] != cls.get_transform_type(): raise InvalidType() parser.parse_numeric('reserved2', 1) parser.parse_numeric_enum_coded('transform_id', cls._get_transform_id_class()) return parser def compose_header(self, transform_length): """Compose transform to bytes. :param last_substruc: Whether this is the last substructure :type last_substruc: bool :return: Composed transform bytes :rtype: bytes """ assert self.next_payload is not None composer = ComposerBinary() composer.compose_numeric(self.next_payload.value, 1) composer.compose_numeric(0, 1) # reserved composer.compose_numeric(transform_length + self.HEADER_SIZE, 2) composer.compose_numeric(self.get_transform_type().value.code, 1) composer.compose_numeric(0, 1) # reserved2 composer.compose_numeric(self.transform_id.value.code, 2) return composer class TransformAttributeKeyLength(DataAttributeLength): @classmethod def get_type(cls): return Ikev2TransformAttributeType.KEY_LENGTH @classmethod def _get_size(cls): return 2 @attr.s class TransformAttributeSignatureAlgorithm(DataAttributeTypeValue): """Signature Algorithm transform attribute (TLV format). The Signature Algorithm attribute has the following format: 0 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ |A| Attribute Type | Attribute Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Signature Algorithm Value ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar signature_algorithm: Signature algorithm value """ signature_algorithm: typing.Union[bytearray, bytes] = attr.ib( validator=attr.validators.instance_of((bytearray, bytes)) ) @classmethod def get_type(cls): return Ikev2TransformAttributeType.SIGNATURE_ALGORITHM @classmethod def _get_format(cls): return DataAttributeFormat.TYPE_LENGTH_VALUE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) if parser.unparsed_length < 2: raise NotEnoughData(2 - parser.unparsed_length) parser.parse_numeric('length', 2) parser.parse_raw('signature_algorithm', parser['length']) return cls(signature_algorithm=bytes(parser['signature_algorithm'])), parser.parsed_length def compose(self): composer = self._compose_header() composer.compose_numeric(len(self.signature_algorithm), 2) composer.compose_raw(self.signature_algorithm) return composer.composed_bytes @attr.s class TransformNoAttributes(Transform): """Transform payload parser for transforms with no attributes.""" @classmethod @abc.abstractmethod def get_transform_type(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_transform_id_class(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) transform = cls( transform_id=parser['transform_id'], ) transform.next_payload = parser['next_payload'] return transform, parser.parsed_length def compose(self): return self.compose_header(transform_length=0).composed_bytes @attr.s class TransformAttributes(Transform): """Transform payload parser for transforms with some attributes.""" @classmethod @abc.abstractmethod def get_transform_type(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_transform_id_class(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _parse_attributes(cls, parsable): raise NotImplementedError() @abc.abstractmethod def _get_attributes(self): raise NotImplementedError() @classmethod def _parse(cls, parsable): header_parser = cls._parse_header(parsable) attributes_size = header_parser['transform_length'] - cls.HEADER_SIZE attributes, attributes_length = cls._parse_attributes(header_parser.unparsed[:attributes_size]) transform = cls( transform_id=header_parser['transform_id'], **attributes ) transform.next_payload = header_parser['next_payload'] return transform, header_parser.parsed_length + attributes_length def compose(self): payload_composer = ComposerBinary() for attribute in self._get_attributes(): payload_composer.compose_parsable(attribute) header_composer = self.compose_header(transform_length=payload_composer.composed_length) return header_composer.composed_bytes + payload_composer.composed_bytes class Ikev2TransformIntegrity(TransformNoAttributes): """Transform payload parser for integrity algorithm.""" @classmethod def get_transform_type(cls): return Ikev2TransformType.INTEG @classmethod def _get_transform_id_class(cls): return Ikev2IntegrityAlgorithm class Ikev2TransformPrf(TransformNoAttributes): """Transform payload parser for pseudorandom function.""" @classmethod def get_transform_type(cls): return Ikev2TransformType.PRF @classmethod def _get_transform_id_class(cls): return Ikev2PseudorandomFunction class Ikev2TransformDhGroup(TransformNoAttributes): """Transform payload parser for Diffie-Hellman group.""" @classmethod def get_transform_type(cls): return Ikev2TransformType.DH @classmethod def _get_transform_id_class(cls): return Ikev2DiffieHellmanGroup @attr.s class Ikev2TransformEncryptionAlgorithm(TransformAttributes): """Transform payload parser for encryption algorithm.""" key_length: typing.Optional[int] = attr.ib( default=None, validator=attr.validators.optional(attr.validators.instance_of(int)), ) @classmethod def get_transform_type(cls): return Ikev2TransformType.ENCR @classmethod def _get_transform_id_class(cls): return Ikev2EncryptionAlgorithm @classmethod def _parse_attributes(cls, parsable): if not parsable: return {'key_length': None}, 0 parser = ParserBinary(parsable) parser.parse_parsable('key_length', TransformAttributeKeyLength) return {'key_length': parser['key_length'].value}, parser.parsed_length def _get_attributes(self): if self.key_length is None: return [] return [TransformAttributeKeyLength(value=self.key_length)] class Ikev2ProposalNextPayload(enum.IntEnum): """Proposal next payload.""" LAST = 0x00 MORE = 0x02 @attr.s class Ikev2Proposal(ParsableBase): """Proposal payload parser. The proposal payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Last Substruc | RESERVED | Proposal Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Proposal Num | Protocol ID | SPI Size |Num Transforms| +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ~ SPI (variable) ~ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar flags: Proposal flags (Ikev2ProposalFlags) :ivar protocol_id: Protocol ID :ivar spi: Security Parameter Index :ivar transforms: List of transforms """ HEADER_SIZE = 8 protocol_id: Ikev2ProtocolId = attr.ib(validator=attr.validators.instance_of(Ikev2ProtocolId)) transforms: list[Transform] = attr.ib(validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(Transform) )) spi: bytes = attr.ib(default=b'', converter=bytes, validator=attr.validators.instance_of(bytes)) last: typing.Optional[Ikev2ProposalNextPayload] = attr.ib( init=False, default=None, validator=attr.validators.optional(attr.validators.instance_of(Ikev2ProposalNextPayload)) ) proposal_number: typing.Optional[int] = attr.ib( init=False, default=None, validator=attr.validators.optional(attr.validators.instance_of(int)) ) @classmethod def _parse(cls, parsable): """Parse proposal from bytes. :param parsable: Bytes to parse :type parsable: bytes :return: Tuple of (parsed proposal, number of bytes parsed) :rtype: tuple(Ikev2ProposalPayload, int) :raises NotEnoughData: If there are not enough bytes to parse """ if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric_flags('last', 1, Ikev2ProposalNextPayload) parser.parse_numeric('reserved', 1) parser.parse_numeric('proposal_length', 2) parser.parse_numeric('proposal_number', 1) parser.parse_numeric_enum_coded('protocol_id', Ikev2ProtocolId) parser.parse_numeric('spi_size', 1) parser.parse_numeric('transform_count', 1) parser.parse_raw('spi', parser['spi_size']) transforms = [] for _ in range(parser['transform_count']): parser.parse_parsable('transform', Ikev2TransformVariantInitiator) transforms.append(parser['transform']) proposal = cls( protocol_id=parser['protocol_id'], spi=parser['spi'], transforms=transforms ) proposal.last = parser['last'] proposal.proposal_number = parser['proposal_number'] return proposal, parser.parsed_length def compose(self): """Compose proposal to bytes. :param last_substruc: Whether this is the last substructure :type last_substruc: bool :return: Composed proposal bytes :rtype: bytes """ assert self.last is not None header_composer = ComposerBinary() header_composer.compose_numeric(self.last.value, 1) header_composer.compose_numeric(0, 1) # reserved payload_composer = ComposerBinary() for i, transform in enumerate(self.transforms): transform.next_payload = ( TransformNextPayload.MORE if i < len(self.transforms) - 1 else TransformNextPayload.LAST ) payload_composer.compose_parsable(transform) header_composer.compose_numeric(self.HEADER_SIZE + len(payload_composer.composed_bytes), 2) header_composer.compose_numeric(self.proposal_number, 1) header_composer.compose_numeric(self.protocol_id.value.code, 1) header_composer.compose_numeric(len(self.spi), 1) header_composer.compose_numeric(len(self.transforms), 1) if self.spi: header_composer.compose_raw(self.spi) return header_composer.composed_bytes + payload_composer.composed_bytes @attr.s class Ikev2PayloadSecurityAssociation(Ikev2PayloadBase): """Security Association payload parser. The Security Association payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ proposals: list[Ikev2Proposal] = attr.ib(validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(Ikev2Proposal) )) @classmethod def get_payload_type(cls): return Ikev2PayloadType.SA def get_transform_by_type(self, transform_type: Ikev2TransformType) -> Transform: for proposal in self.proposals: for transform in proposal.transforms: if transform.get_transform_type() == transform_type: return transform raise KeyError(transform_type) @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) proposals = [] parser_proposal = ParserBinary(parser.unparsed[:parser['payload_length'] - cls.HEADER_SIZE:]) while parser_proposal.unparsed_length > 0: parser_proposal.parse_parsable('proposal', Ikev2Proposal) proposal = parser_proposal['proposal'] proposals.append(proposal) payload = cls( flags=parser['flags'], proposals=proposals ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length + parser_proposal.parsed_length def compose(self): composer_payload = ComposerBinary() for proposal_number, proposal in enumerate(self.proposals): proposal.last = ( Ikev2ProposalNextPayload.MORE if proposal_number < len(self.proposals) - 1 else Ikev2ProposalNextPayload.LAST ) proposal.proposal_number = proposal_number + 1 composer_payload.compose_parsable(proposal) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev2PayloadNonce(Ikev2PayloadBase): """Nonce payload parser. The nonce payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Nonce Data ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar nonce_data: Random data generated by the transmitting entity """ nonce_data: bytes = attr.ib( converter=bytes, validator=attr.validators.and_( attr.validators.instance_of(bytes), attr.validators.min_len(16), attr.validators.max_len(256) ) ) @classmethod def get_payload_type(cls): return Ikev2PayloadType.NONCE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) nonce_data_length = parser['payload_length'] - cls.HEADER_SIZE if nonce_data_length < 16: raise NotEnoughData(bytes_needed=16 - nonce_data_length) if nonce_data_length > 256: raise TooMuchData(bytes_needed=nonce_data_length - 256) parser.parse_raw('nonce_data', nonce_data_length) payload = cls( flags=parser['flags'], nonce_data=parser['nonce_data'] ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_raw(self.nonce_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev2PayloadKeyExchange(Ikev2PayloadBase): """Key exchange payload parser. The key exchange payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Diffie-Hellman Group Num | RESERVED | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Key Exchange Data ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar dh_group: Diffie-Hellman group number :ivar key_exchange_data: Diffie-Hellman public value """ dh_group: Ikev2DiffieHellmanGroup = attr.ib(validator=attr.validators.instance_of(Ikev2DiffieHellmanGroup)) key_exchange_data: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)) ) @classmethod def get_payload_type(cls): return Ikev2PayloadType.KE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric_enum_coded('dh_group', Ikev2DiffieHellmanGroup) parser.parse_numeric('reserved2', 2) key_exchange_length = parser['payload_length'] - 8 parser.parse_raw('key_exchange_data', key_exchange_length) payload = cls( flags=parser['flags'], dh_group=parser['dh_group'], key_exchange_data=parser['key_exchange_data'] ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_numeric(self.dh_group.value.code, 2) composer_payload.compose_numeric(0, 2) # reserved2 composer_payload.compose_raw(self.key_exchange_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev2PayloadDelete(Ikev2PayloadBase): """Delete payload parser. The Delete payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Protocol ID | SPI Size | Num of SPIs | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Security Parameter Index(es) (SPI) ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar protocol_id: Protocol ID :ivar spi_size: Length in octets of the SPI :ivar num_spis: Number of SPIs contained in the payload :ivar spis: List of Security Parameter Indexes """ protocol_id: Ikev2ProtocolId = attr.ib(validator=attr.validators.instance_of(Ikev2ProtocolId)) spis: list[int] = attr.ib(validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(int), )) @classmethod def get_payload_type(cls): return Ikev2PayloadType.DELETE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric_enum_coded('protocol_id', Ikev2ProtocolId) parser.parse_numeric('spi_size', 1) parser.parse_numeric('num_spis', 2) parser.parse_numeric_array('spis', parser['num_spis'], 8) payload = cls( flags=parser['flags'], protocol_id=parser['protocol_id'], spis=parser['spis'] ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_numeric_enum_coded(self.protocol_id) composer_payload.compose_numeric(len(self.spis) * 8, 1) composer_payload.compose_numeric(len(self.spis), 2) composer_payload.compose_numeric_array(self.spis, 8) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes class Ikev2NotifyTypeFactory(TwoByteEnumParsable): @classmethod def get_enum_class(cls): return Ikev2NotifyType @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s class Ikev2PayloadNotifyBase(Ikev2PayloadBase): """Notify payload parser. The notify payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Protocol ID | SPI Size | Notify Message Type | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Security Parameter Index (SPI) ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Notification Data ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar protocol_id: Protocol ID (1 byte) :ivar spi_size: Size of SPI in bytes (1 byte) :ivar notify_message_type: Type of notification message (2 bytes) :ivar spi: Security Parameter Index (variable length) :ivar data: Notification data (variable length) """ protocol_id: Ikev2ProtocolId = attr.ib(validator=attr.validators.instance_of(Ikev2ProtocolId)) type: Ikev2NotifyType = attr.ib(validator=attr.validators.instance_of(Ikev2NotifyType)) spi: bytes = attr.ib(converter=bytes, validator=attr.validators.instance_of(bytes)) @classmethod @abc.abstractmethod def _parse_type(cls, parser, name): raise NotImplementedError() @classmethod @abc.abstractmethod def _parse_data(cls, parser, notification_data_length): raise NotImplementedError() @abc.abstractmethod def _compose_data(self, composer): raise NotImplementedError() @classmethod def get_payload_type(cls): return Ikev2PayloadType.NOTIFY @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric_enum_coded('protocol_id', Ikev2ProtocolId) parser.parse_numeric('spi_size', 1) cls._parse_type(parser, 'type') if parser['spi_size'] > 0: parser.parse_raw('spi', parser['spi_size']) spi = parser['spi'] else: spi = b'' del parser['spi_size'] if 'spi' in parser: del parser['spi'] notification_data_length = parser['payload_length'] - (cls.HEADER_SIZE + 4) - len(spi) cls._parse_data(parser, notification_data_length) next_payload = parser['next_payload'] del parser['next_payload'] del parser['payload_length'] payload = cls( **parser, spi=spi, ) payload.next_payload = next_payload return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_numeric_enum_coded(self.protocol_id) composer_payload.compose_numeric(len(self.spi), 1) composer_payload.compose_numeric_enum_coded(self.type) if self.spi: composer_payload.compose_raw(self.spi) self._compose_data(composer_payload) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes class Ikev2PayloadNotifyNoData(Ikev2PayloadNotifyBase): @classmethod def _parse_data(cls, parser, notification_data_length): pass def _compose_data(self, composer): pass @classmethod @abc.abstractmethod def _get_message_type(cls): raise NotImplementedError() @classmethod def _parse_type(cls, parser, name): parser.parse_parsable(name, Ikev2NotifyTypeFactory) if parser[name] != cls._get_message_type(): raise InvalidType() class Ikev2PayloadNotifyAuthenticationFailed(Ikev2PayloadNotifyNoData): @classmethod def _get_message_type(cls): return Ikev2NotifyType.AUTHENTICATION_FAILED class Ikev2NotifyPayloadUseTransportMode(Ikev2PayloadNotifyNoData): @classmethod def _get_message_type(cls): return Ikev2NotifyType.USE_TRANSPORT_MODE class Ikev2NotifyPayloadHttpCertLookupSupported(Ikev2PayloadNotifyNoData): @classmethod def _get_message_type(cls): return Ikev2NotifyType.HTTP_CERT_LOOKUP_SUPPORTED class Ikev2NotifyPayloadIkev2FragmentationSupported(Ikev2PayloadNotifyNoData): @classmethod def _get_message_type(cls): return Ikev2NotifyType.IKEV2_FRAGMENTATION_SUPPORTED class Ikev2NotifyPayloadIntermediateExchangeSupported(Ikev2PayloadNotifyNoData): @classmethod def _get_message_type(cls): return Ikev2NotifyType.INTERMEDIATE_EXCHANGE_SUPPORTED class Ikev2NotifyPayloadUsePpk(Ikev2PayloadNotifyNoData): @classmethod def _get_message_type(cls): return Ikev2NotifyType.USE_PPK class Ikev2NotifyPayloadRedirectSupported(Ikev2PayloadNotifyNoData): @classmethod def _get_message_type(cls): return Ikev2NotifyType.REDIRECT_SUPPORTED class Ikev2NotifyPayloadChildlessIkev2Supported(Ikev2PayloadNotifyNoData): @classmethod def _get_message_type(cls): return Ikev2NotifyType.CHILDLESS_IKEV2_SUPPORTED @attr.s class Ikev2PayloadNotifyUnparsed(Ikev2PayloadNotifyBase): data: typing.Union[bytes, bytearray] = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def _parse_data(cls, parser, notification_data_length): parser.parse_raw('data', notification_data_length) def _compose_data(self, composer): composer.compose_raw(self.data) @classmethod def _parse_type(cls, parser, name): parser.parse_numeric_enum_coded(name, Ikev2NotifyType) @attr.s class Ikev2PayloadNotifyParsedBase(Ikev2PayloadNotifyBase): @classmethod @abc.abstractmethod def _parse_data(cls, parser, notification_data_length): raise NotImplementedError() @abc.abstractmethod def _compose_data(self, composer): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_message_type(cls): raise NotImplementedError() @classmethod def _parse_type(cls, parser, name): parser.parse_numeric_enum_coded(name, Ikev2NotifyType) if parser[name] != cls._get_message_type(): raise InvalidType() @attr.s class Ikev2NotifyPayloadInvalidKe(Ikev2PayloadNotifyParsedBase): """Invalid KE payload notification data parser.""" dh_group: Ikev2DiffieHellmanGroup = attr.ib(validator=attr.validators.instance_of(Ikev2DiffieHellmanGroup)) @classmethod def _get_message_type(cls): return Ikev2NotifyType.INVALID_KE_PAYLOAD @classmethod def _parse_data(cls, parser, notification_data_length): parser.parse_numeric_enum_coded('dh_group', Ikev2DiffieHellmanGroup) def _compose_data(self, composer): composer.compose_numeric_enum_coded(self.dh_group) @attr.s class Ikev2NotifyPayloadCookie(Ikev2PayloadNotifyParsedBase): """Invalid KE payload notification data parser.""" cookie: typing.Union[bytes, bytearray] = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def _get_message_type(cls): return Ikev2NotifyType.COOKIE @classmethod def _parse_data(cls, parser, notification_data_length): parser.parse_raw('cookie', notification_data_length) def _compose_data(self, composer): composer.compose_raw(self.cookie) @attr.s class Ikev2NotifyPayloadSetWindowSize(Ikev2PayloadNotifyParsedBase): """Set window size payload notification data parser.""" window_size: int = attr.ib(validator=[ attr.validators.instance_of(int), attr.validators.in_(range(0, 2 ** 32)), ]) @classmethod def _get_message_type(cls): return Ikev2NotifyType.SET_WINDOW_SIZE @classmethod def _parse_data(cls, parser, notification_data_length): if notification_data_length != 4: raise InvalidValue(notification_data_length, cls, 'notification_data_length') parser.parse_numeric('window_size', 4) def _compose_data(self, composer): composer.compose_numeric(self.window_size, 4) @attr.s class Ikev2NotifyPayloadNatDetectionBase(Ikev2PayloadNotifyParsedBase): hash_data: typing.Union[bytes, bytearray] = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod @abc.abstractmethod def _get_message_type(cls): raise NotImplementedError() @classmethod def _parse_data(cls, parser, notification_data_length): parser.parse_raw('hash_data', notification_data_length) def _compose_data(self, composer): composer.compose_raw(self.hash_data) @attr.s class Ikev2NotifyPayloadNatDetectionSourceIp(Ikev2NotifyPayloadNatDetectionBase): """NAT detection source IP payload notification data parser.""" @classmethod def _get_message_type(cls): return Ikev2NotifyType.NAT_DETECTION_SOURCE_IP @attr.s class Ikev2NotifyPayloadNatDetectionDestinationIp(Ikev2NotifyPayloadNatDetectionBase): """NAT detection destination IP payload notification data parser.""" @classmethod def _get_message_type(cls): return Ikev2NotifyType.NAT_DETECTION_DESTINATION_IP @attr.s class Ikev2NotifyPayloadSignatureHashAlgorithms(Ikev2PayloadNotifyParsedBase): """Signature hash algorithms notification (RFC 7427 §4).""" hash_algorithms: tuple[Ikev2HashAlgorithm, ...] = attr.ib( converter=tuple, validator=attr.validators.deep_iterable(attr.validators.instance_of(Ikev2HashAlgorithm)), ) @classmethod def _get_message_type(cls): return Ikev2NotifyType.SIGNATURE_HASH_ALGORITHMS @classmethod def _parse_data(cls, parser, notification_data_length): if notification_data_length % 2 != 0: raise InvalidValue(notification_data_length, cls, 'notification_data_length') parser.parse_numeric_array( 'hash_algorithms', notification_data_length // 2, 2, Ikev2HashAlgorithm.from_code, ) def _compose_data(self, composer): for hash_algorithm in self.hash_algorithms: composer.compose_numeric(hash_algorithm.value.code, 2) class Ikev2NotifyPayloadVariantBase(VariantParsable): @classmethod @abc.abstractmethod def get_parsed_notifies(cls): raise NotImplementedError() @classmethod def _get_variants(cls): variants = cls.get_parsed_notifies() variants.update([ (notify_type, (Ikev2PayloadNotifyUnparsed, )) for notify_type in Ikev2NotifyType if notify_type not in variants ]) return variants class Ikev2NotifyPayloadVariantResponder(Ikev2NotifyPayloadVariantBase): @classmethod def get_parsed_notifies(cls): return collections.OrderedDict([ (Ikev2NotifyType.COOKIE, [Ikev2NotifyPayloadCookie, ]), (Ikev2NotifyType.INVALID_KE_PAYLOAD, [Ikev2NotifyPayloadInvalidKe, ]), (Ikev2NotifyType.SET_WINDOW_SIZE, [Ikev2NotifyPayloadSetWindowSize, ]), (Ikev2NotifyType.NAT_DETECTION_SOURCE_IP, [Ikev2NotifyPayloadNatDetectionSourceIp, ]), (Ikev2NotifyType.NAT_DETECTION_DESTINATION_IP, [Ikev2NotifyPayloadNatDetectionDestinationIp, ]), (Ikev2NotifyType.USE_TRANSPORT_MODE, [Ikev2NotifyPayloadUseTransportMode, ]), (Ikev2NotifyType.HTTP_CERT_LOOKUP_SUPPORTED, [Ikev2NotifyPayloadHttpCertLookupSupported, ]), (Ikev2NotifyType.SIGNATURE_HASH_ALGORITHMS, [Ikev2NotifyPayloadSignatureHashAlgorithms, ]), (Ikev2NotifyType.IKEV2_FRAGMENTATION_SUPPORTED, [Ikev2NotifyPayloadIkev2FragmentationSupported, ]), (Ikev2NotifyType.INTERMEDIATE_EXCHANGE_SUPPORTED, [Ikev2NotifyPayloadIntermediateExchangeSupported, ]), (Ikev2NotifyType.USE_PPK, [Ikev2NotifyPayloadUsePpk, ]), (Ikev2NotifyType.REDIRECT_SUPPORTED, [Ikev2NotifyPayloadRedirectSupported, ]), (Ikev2NotifyType.CHILDLESS_IKEV2_SUPPORTED, [Ikev2NotifyPayloadChildlessIkev2Supported, ]), ]) @attr.s class Ikev2PayloadCertificateBase(Ikev2PayloadBase): """Certificate payload base (RFC 7296 §3.6, §3.7).""" cert_encoding: Ikev2CertificateType = attr.ib(validator=attr.validators.instance_of(Ikev2CertificateType)) @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod def _parse_header(cls, parsable): parser = super()._parse_header(parsable) parser.parse_numeric_enum_coded('cert_encoding', Ikev2CertificateType) return parser def _compose_cert_encoding(self, composer): composer.compose_numeric_enum_coded(self.cert_encoding) @attr.s class Ikev2PayloadCertificateRequest(Ikev2PayloadCertificateBase): """Certificate Request payload parser (RFC 7296 §3.7). .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Cert Encoding | | +-+-+-+-+-+-+-+-+ | ~ Certification Authority ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ certification_authority: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)) ) # RFC 7296 §3.7: SHA-1 SPKI hashes are 20 octets. _SHA1_HASH_LENGTH_BYTES: typing.ClassVar[int] = 20 @classmethod def get_payload_type(cls): return Ikev2PayloadType.CERTREQ @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('certification_authority', parser['payload_length'] - cls.HEADER_SIZE - 1) payload = cls( flags=parser['flags'], cert_encoding=parser['cert_encoding'], certification_authority=parser['certification_authority'] ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() self._compose_cert_encoding(composer_payload) composer_payload.compose_raw(self.certification_authority) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes def get_authority_hashes(self) -> list[bytes]: """Split the Certification Authority field into 20-octet SHA-1(SPKI) hashes (RFC 7296 §3.7).""" raw = bytes(self.certification_authority) aligned_length = len(raw) - (len(raw) % self._SHA1_HASH_LENGTH_BYTES) return [ raw[offset:offset + self._SHA1_HASH_LENGTH_BYTES] for offset in range(0, aligned_length, self._SHA1_HASH_LENGTH_BYTES) ] @attr.s class Ikev2PayloadCertificate(Ikev2PayloadCertificateBase): """Certificate payload parser (RFC 7296 §3.6). .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Cert Encoding | | +-+-+-+-+-+-+-+-+ | ~ Certificate Data ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ certificate_data: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)) ) @classmethod def get_payload_type(cls): return Ikev2PayloadType.CERT @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('certificate_data', parser['payload_length'] - cls.HEADER_SIZE - 1) payload = cls( flags=parser['flags'], cert_encoding=parser['cert_encoding'], certificate_data=parser['certificate_data'] ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() self._compose_cert_encoding(composer_payload) composer_payload.compose_raw(self.certificate_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev2PayloadIdentificationBase(Ikev2PayloadBase): """Identification payload parser (RFC 7296 §3.5). .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | ID Type | RESERVED | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Identification Data ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ @classmethod @abc.abstractmethod def get_payload_type(cls): raise NotImplementedError() # The following abstract methods are concretely provided by the # ``IkeIdentification*Mixin`` mixins listed first in the concrete # subclass MRO; they appear here so static-analysis tools see them # as part of the contract. @classmethod @abc.abstractmethod def get_id_type_ikev2(cls) -> Ikev2IdType: raise NotImplementedError() @classmethod @abc.abstractmethod def _decode_identifier(cls, id_data: bytes): raise NotImplementedError() @abc.abstractmethod def _encode_identifier(self) -> bytes: raise NotImplementedError() @property def id_type(self) -> Ikev2IdType: """Wire-format ``ID Type`` field — fixed per concrete subclass via the :meth:`get_id_type_ikev2` mixin classmethod.""" return self.get_id_type_ikev2() @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric_enum_coded('id_type', Ikev2IdType) if parser['id_type'] != cls.get_id_type_ikev2(): raise InvalidType() parser.parse_numeric('reserved', 3) id_data_length = parser['payload_length'] - cls.HEADER_SIZE - 4 if id_data_length < 0: raise NotEnoughData(bytes_needed=-id_data_length) parser.parse_raw('id_data', id_data_length) payload = cls( flags=parser['flags'], identifier=cls._decode_identifier(bytes(parser['id_data'])), ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_numeric_enum_coded(self.id_type) composer_payload.compose_numeric(0, 3) composer_payload.compose_raw(self._encode_identifier()) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes # --- IDi (initiator) concrete subclasses -------------------------------- @attr.s class Ikev2PayloadIdentificationInitiatorBase(Ikev2PayloadIdentificationBase): # pylint: disable=abstract-method """IDi payload base (RFC 7296 §3.5).""" @classmethod def get_payload_type(cls): return Ikev2PayloadType.IDI @attr.s class Ikev2PayloadIdentificationInitiatorFqdn( IkeIdentificationFqdnMixin, Ikev2PayloadIdentificationInitiatorBase, ): """IDi with fully-qualified domain name (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationInitiatorRfc822Addr( IkeIdentificationRfc822AddrMixin, Ikev2PayloadIdentificationInitiatorBase, ): """IDi with RFC 822 email address (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationInitiatorIpv4Addr( IkeIdentificationIpv4AddrMixin, Ikev2PayloadIdentificationInitiatorBase, ): """IDi with IPv4 address (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationInitiatorIpv6Addr( IkeIdentificationIpv6AddrMixin, Ikev2PayloadIdentificationInitiatorBase, ): """IDi with IPv6 address (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationInitiatorDerAsn1Dn( IkeIdentificationDerAsn1DnMixin, Ikev2PayloadIdentificationInitiatorBase, ): """IDi with DER-encoded ASN.1 X.500 Distinguished Name (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationInitiatorKeyId( IkeIdentificationKeyIdMixin, Ikev2PayloadIdentificationInitiatorBase, ): """IDi with opaque vendor-specific key identifier (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationInitiatorDerAsn1Gn( IkeIdentificationDerAsn1GnMixin, Ikev2PayloadIdentificationInitiatorBase, ): """IDi with DER-encoded ASN.1 X.500 General Name (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationInitiatorFcName( IkeIdentificationFcNameMixin, Ikev2PayloadIdentificationInitiatorBase, ): """IDi with Fibre Channel name (RFC 4595).""" @attr.s class Ikev2PayloadIdentificationInitiatorNull( IkeIdentificationNullMixin, Ikev2PayloadIdentificationInitiatorBase, ): """IDi with NULL identification (RFC 7619).""" class Ikev2PayloadIdentificationInitiator(VariantParsable): """Variant dispatcher for IKEv2 IDi payload (RFC 7296 §3.5).""" @classmethod def _get_variants(cls): return collections.OrderedDict([ (subclass.get_id_type_ikev2(), [subclass]) for subclass in ( Ikev2PayloadIdentificationInitiatorFqdn, Ikev2PayloadIdentificationInitiatorRfc822Addr, Ikev2PayloadIdentificationInitiatorIpv4Addr, Ikev2PayloadIdentificationInitiatorIpv6Addr, Ikev2PayloadIdentificationInitiatorDerAsn1Dn, Ikev2PayloadIdentificationInitiatorDerAsn1Gn, Ikev2PayloadIdentificationInitiatorKeyId, Ikev2PayloadIdentificationInitiatorFcName, Ikev2PayloadIdentificationInitiatorNull, ) ]) # --- IDr (responder) concrete subclasses -------------------------------- @attr.s class Ikev2PayloadIdentificationResponderBase(Ikev2PayloadIdentificationBase): # pylint: disable=abstract-method """IDr payload base (RFC 7296 §3.5).""" @classmethod def get_payload_type(cls): return Ikev2PayloadType.IDR @attr.s class Ikev2PayloadIdentificationResponderFqdn( IkeIdentificationFqdnMixin, Ikev2PayloadIdentificationResponderBase, ): """IDr with fully-qualified domain name (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationResponderRfc822Addr( IkeIdentificationRfc822AddrMixin, Ikev2PayloadIdentificationResponderBase, ): """IDr with RFC 822 email address (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationResponderIpv4Addr( IkeIdentificationIpv4AddrMixin, Ikev2PayloadIdentificationResponderBase, ): """IDr with IPv4 address (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationResponderIpv6Addr( IkeIdentificationIpv6AddrMixin, Ikev2PayloadIdentificationResponderBase, ): """IDr with IPv6 address (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationResponderDerAsn1Dn( IkeIdentificationDerAsn1DnMixin, Ikev2PayloadIdentificationResponderBase, ): """IDr with DER-encoded ASN.1 X.500 Distinguished Name (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationResponderKeyId( IkeIdentificationKeyIdMixin, Ikev2PayloadIdentificationResponderBase, ): """IDr with opaque vendor-specific key identifier (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationResponderDerAsn1Gn( IkeIdentificationDerAsn1GnMixin, Ikev2PayloadIdentificationResponderBase, ): """IDr with DER-encoded ASN.1 X.500 General Name (RFC 7296 §3.5).""" @attr.s class Ikev2PayloadIdentificationResponderFcName( IkeIdentificationFcNameMixin, Ikev2PayloadIdentificationResponderBase, ): """IDr with Fibre Channel name (RFC 4595).""" @attr.s class Ikev2PayloadIdentificationResponderNull( IkeIdentificationNullMixin, Ikev2PayloadIdentificationResponderBase, ): """IDr with NULL identification (RFC 7619).""" class Ikev2PayloadIdentificationResponder(VariantParsable): """Variant dispatcher for IKEv2 IDr payload (RFC 7296 §3.5).""" @classmethod def _get_variants(cls): return collections.OrderedDict([ (subclass.get_id_type_ikev2(), [subclass]) for subclass in ( Ikev2PayloadIdentificationResponderFqdn, Ikev2PayloadIdentificationResponderRfc822Addr, Ikev2PayloadIdentificationResponderIpv4Addr, Ikev2PayloadIdentificationResponderIpv6Addr, Ikev2PayloadIdentificationResponderDerAsn1Dn, Ikev2PayloadIdentificationResponderDerAsn1Gn, Ikev2PayloadIdentificationResponderKeyId, Ikev2PayloadIdentificationResponderFcName, Ikev2PayloadIdentificationResponderNull, ) ]) @attr.s class Ikev2AuthDigitalSignatureEnvelope(ParsableBase): """Digital-Signature Authentication Data envelope (RFC 7427 §3). .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | ASN.1 Length | AlgorithmIdentifier ASN.1 object | +-+-+-+-+-+-+-+-+ | ~ AlgorithmIdentifier ASN.1 object continuing ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ~ Signature Value ~ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ signature_algorithm: Signature = attr.ib( validator=attr.validators.instance_of(Signature) ) signature: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)), metadata={'human_friendly': False}, ) @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('algorithm_identifier_length', 1) parser.parse_raw('algorithm_identifier', parser['algorithm_identifier_length']) parser.parse_raw('signature', len(parsable) - parser.parsed_length) algorithm = asn1crypto.algos.SignedDigestAlgorithm.load( bytes(parser['algorithm_identifier']), ) signature_algorithm = Signature.from_oid(algorithm['algorithm'].dotted) return cls( signature_algorithm=signature_algorithm, signature=parser['signature'], ), parser.parsed_length @classmethod def from_bytes(cls, parsable): """Public entry point for callers outside this module.""" envelope, _ = cls._parse(parsable) return envelope def compose(self): algorithm_identifier = asn1crypto.algos.SignedDigestAlgorithm({ 'algorithm': self.signature_algorithm.value.oid, }).dump() composer = ComposerBinary() composer.compose_numeric(len(algorithm_identifier), 1) composer.compose_raw(algorithm_identifier) composer.compose_raw(self.signature) return composer.composed_bytes @attr.s class Ikev2PayloadAuthentication(Ikev2PayloadBase): """Authentication payload parser (RFC 7296 §3.8). .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Auth Method | RESERVED | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ~ Authentication Data ~ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ auth_method: Ikev2AuthenticationMethod = attr.ib( validator=attr.validators.instance_of(Ikev2AuthenticationMethod) ) auth_data: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)) ) @classmethod def get_payload_type(cls): return Ikev2PayloadType.AUTH @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric_enum_coded('auth_method', Ikev2AuthenticationMethod) parser.parse_numeric('reserved', 3) auth_data_length = parser['payload_length'] - cls.HEADER_SIZE - 4 if auth_data_length < 0: raise NotEnoughData(bytes_needed=-auth_data_length) parser.parse_raw('auth_data', auth_data_length) payload = cls( flags=parser['flags'], auth_method=parser['auth_method'], auth_data=parser['auth_data'], ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_numeric_enum_coded(self.auth_method) composer_payload.compose_numeric(0, 3) composer_payload.compose_raw(self.auth_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes def parse_digital_signature_envelope(self) -> Ikev2AuthDigitalSignatureEnvelope: """Parse ``auth_data`` as an RFC 7427 envelope; only valid when ``auth_method == DIGITAL_SIGNATURE``.""" if self.auth_method != Ikev2AuthenticationMethod.DIGITAL_SIGNATURE: raise InvalidType(self.auth_method) return Ikev2AuthDigitalSignatureEnvelope.from_bytes(bytes(self.auth_data)) @attr.s class Ikev2PayloadEncryptedAndAuthenticated(Ikev2PayloadBase): """Encrypted payload parser (RFC 7296 §3.14). .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Initialization Vector | | (length is block size for encryption algorithm) | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ~ Encrypted IKE Payloads ~ + +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | Padding (0-255 octets) | +-+-+-+-+-+-+-+-+ +-+-+-+-+-+-+-+-+ | | Pad Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ ~ Integrity Checksum Data ~ +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ """ encrypted_data: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)) ) @classmethod def get_payload_type(cls): return Ikev2PayloadType.SK @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('encrypted_data', parser['payload_length'] - cls.HEADER_SIZE) payload = cls( flags=parser['flags'], encrypted_data=parser['encrypted_data'], ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_raw(self.encrypted_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev2PayloadEap(Ikev2PayloadBase): """EAP payload parser (RFC 7296 §3.16).""" eap_data: typing.Union[bytes, bytearray] = attr.ib( validator=attr.validators.instance_of((bytes, bytearray)) ) @classmethod def get_payload_type(cls): return Ikev2PayloadType.EAP @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('eap_data', parser['payload_length'] - cls.HEADER_SIZE) payload = cls( flags=parser['flags'], eap_data=parser['eap_data'], ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_raw(self.eap_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes @attr.s class Ikev2PayloadVendorId(Ikev2PayloadBase): """Vendor ID payload parser. The Vendor ID payload has the following format: .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload |C| RESERVED | Payload Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | | ~ Vendor ID (VID) ~ | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar vendor_id: Vendor ID data """ vendor_id: typing.Union[bytes, bytearray] = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def get_payload_type(cls): return Ikev2PayloadType.VENDOR_ID @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('vendor_id', parser['payload_length'] - cls.HEADER_SIZE) payload = cls( flags=parser['flags'], vendor_id=parser['vendor_id'] ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_raw(self.vendor_id) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes IKEV2_PAYLOAD_CLASSES_BY_TYPE = { Ikev2PayloadType.SA: Ikev2PayloadSecurityAssociation, Ikev2PayloadType.NONCE: Ikev2PayloadNonce, Ikev2PayloadType.KE: Ikev2PayloadKeyExchange, Ikev2PayloadType.NOTIFY: Ikev2NotifyPayloadVariantResponder, Ikev2PayloadType.IDI: Ikev2PayloadIdentificationInitiator, Ikev2PayloadType.IDR: Ikev2PayloadIdentificationResponder, Ikev2PayloadType.CERTREQ: Ikev2PayloadCertificateRequest, Ikev2PayloadType.CERT: Ikev2PayloadCertificate, Ikev2PayloadType.AUTH: Ikev2PayloadAuthentication, Ikev2PayloadType.SK: Ikev2PayloadEncryptedAndAuthenticated, Ikev2PayloadType.EAP: Ikev2PayloadEap, Ikev2PayloadType.VENDOR_ID: Ikev2PayloadVendorId, } class Ikev2TransformVariantBase(VariantParsable): @classmethod @abc.abstractmethod def get_parsed_transforms(cls): raise NotImplementedError() @classmethod def _get_variants(cls): variants = cls.get_parsed_transforms() variants.update([ (Ikev2TransformType.ENCR, [Ikev2TransformEncryptionAlgorithm, ]), (Ikev2TransformType.PRF, [Ikev2TransformPrf, ]), (Ikev2TransformType.DH, [Ikev2TransformDhGroup, ]), (Ikev2TransformType.INTEG, [Ikev2TransformIntegrity, ]), ]) return variants class Ikev2TransformVariantInitiator(Ikev2TransformVariantBase): @classmethod def get_parsed_transforms(cls): return collections.OrderedDict() cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ike/isakmp.py000066400000000000000000000276741524413560000272350ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 """ISAKMP header parser.""" import enum import typing import attr from cryptodatahub.ike.algorithm import Ikev1PayloadType, Ikev2PayloadType, Ikev1ExchangeType, Ikev2ExchangeType from cryptodatahub.ike.version import IkeVersion from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary from cryptoparser.common.exception import NotEnoughData from cryptoparser.ike.version import IsakmpProtocolVersion from cryptoparser.ike.ikev1 import ( Ikev1PayloadBase, Ikev1PayloadUnparsed, IKEV1_PAYLOAD_CLASSES_BY_TYPE, ) from cryptoparser.ike.ikev2 import ( Ikev2PayloadBase, Ikev2PayloadUnparsed, IKEV2_PAYLOAD_CLASSES_BY_TYPE, ) class IsakmpFlags(enum.IntFlag): """ISAKMP flags.""" ENCRYPTION = 1 << 0 COMMIT = 1 << 1 AUTHENTICATION_ONLY = 1 << 2 INITIATOR = 1 << 3 VERSION = 1 << 4 RESPONSE = 1 << 5 @attr.s class IsakmpMessage(ParsableBase): """ISAKMP message parser. .. code-block:: text 1 2 3 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 2 3 4 5 6 7 8 9 0 1 +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | IKE SA Initiator's SPI | | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | IKE SA Responder's SPI | | | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Next Payload | MjVer | MnVer | Exchange Type | Flags | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Message ID | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Length | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ | Payload Data (variable) | +-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+-+ :ivar initiator_spi: Initiator SPI :ivar responder_spi: Responder SPI :ivar exchange_type: Exchange Type :ivar flags: Flags :ivar message_id: Message ID :ivar payloads: Payloads :ivar version: Version """ HEADER_SIZE = 28 version: IsakmpProtocolVersion = attr.ib(validator=attr.validators.instance_of(IsakmpProtocolVersion)) initiator_spi: int = attr.ib(validator=attr.validators.and_( attr.validators.instance_of(int), attr.validators.ge(0), attr.validators.lt(2**64) )) responder_spi: int = attr.ib(validator=attr.validators.and_( attr.validators.instance_of(int), attr.validators.ge(0), attr.validators.lt(2**64))) exchange_type: typing.Union[Ikev1ExchangeType, Ikev2ExchangeType] = attr.ib( validator=attr.validators.instance_of((Ikev1ExchangeType, Ikev2ExchangeType)) ) flags: list[IsakmpFlags] = attr.ib(validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(IsakmpFlags) )) message_id: int = attr.ib(validator=attr.validators.and_( attr.validators.instance_of(int), attr.validators.ge(0), attr.validators.lt(2**32) )) payloads: list[typing.Union[Ikev1PayloadBase, Ikev2PayloadBase]] = attr.ib( validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of((Ikev1PayloadBase, Ikev2PayloadBase)) ) ) def get_payloads_by_type( self, payload_type: typing.Union[Ikev1PayloadType, Ikev2PayloadType] ) -> list[typing.Union[Ikev1PayloadBase, Ikev2PayloadBase]]: return [ payload for payload in self.payloads if payload.get_payload_type() == payload_type ] def get_payload_by_type( self, payload_type: typing.Union[Ikev1PayloadType, Ikev2PayloadType] ) -> typing.Union[Ikev1PayloadBase, Ikev2PayloadBase]: payloads = self.get_payloads_by_type(payload_type) if not payloads: raise KeyError(payload_type) if len(payloads) > 1: raise IndexError(payload_type) return payloads[0] @classmethod def parse_header_and_body(cls, parsable): """Parse the 28-byte ISAKMP header and slice out the payload body. Returns ``(header, body_bytes, next_payload)``. ``header`` is a dict of the concrete header fields; ``body_bytes`` is ``parsable[HEADER_SIZE:header.length]``; ``next_payload`` is the first payload-type code (enum member or :class:`IkePayloadTypeUnknown`). Callers walk the plaintext payload chain with :meth:`parse_payload_chain`; encryption-aware callers (IKEv1 ``ENCRYPTION`` flag per RFC 2408 §3.1, wire-level state that cryptoparser does not track) wrap ``body_bytes`` opaquely instead. """ if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric('initiator_spi', 8) parser.parse_numeric('responder_spi', 8) parser.parse_raw('next_payload', 1) parser.parse_parsable('version', IsakmpProtocolVersion) if parser['version'].major == IkeVersion.V1: next_payload_type = Ikev1PayloadType parser.parse_numeric_enum_coded('exchange_type', Ikev1ExchangeType) elif parser['version'].major == IkeVersion.V2: next_payload_type = Ikev2PayloadType parser.parse_numeric_enum_coded('exchange_type', Ikev2ExchangeType) else: raise NotImplementedError(parser['version']) next_payload_parser = ParserBinary(parser['next_payload']) next_payload_parser.parse_numeric_enum_coded('value', next_payload_type) next_payload = next_payload_parser['value'] parser.parse_numeric_flags('flags', 1, IsakmpFlags) parser.parse_numeric('message_id', 4) parser.parse_numeric('length', 4) body_bytes = bytes(parsable[parser.parsed_length:parser['length']]) header = { 'initiator_spi': parser['initiator_spi'], 'responder_spi': parser['responder_spi'], 'version': parser['version'], 'exchange_type': parser['exchange_type'], 'flags': parser['flags'], 'message_id': parser['message_id'], } return header, body_bytes, next_payload @classmethod def parse_payload_chain(cls, version, body_bytes, next_payload): """Walk the payload chain from ``body_bytes`` starting at ``next_payload``. Returns ``(payloads, consumed_length)``. Do NOT call this for IKEv1 messages with the ``ENCRYPTION`` flag set: ``body_bytes`` is ciphertext under SKEYID_e (RFC 2408 §3.1) and walking it as plaintext misreads the first block's payload-length field. That case is a caller-side concern. """ if version.major == IkeVersion.V1: payload_classes = IKEV1_PAYLOAD_CLASSES_BY_TYPE payload_none = Ikev1PayloadType.NONE unparsed_class = Ikev1PayloadUnparsed elif version.major == IkeVersion.V2: payload_classes = IKEV2_PAYLOAD_CLASSES_BY_TYPE payload_none = Ikev2PayloadType.NONE unparsed_class = Ikev2PayloadUnparsed else: raise NotImplementedError(version) payloads = [] parser_payload = ParserBinary(body_bytes) while parser_payload.unparsed_length > 0 and next_payload != payload_none: # ``next_payload`` is either a known enum member or — for # private-use / unimplemented codes (IANA "IKE Payload # Types" 128-255 per RFC 2408 §3.10 / RFC 7296 §3.2) — an # :class:`IkePayloadTypeUnknown` wrapper produced by # ``Ikev*PayloadTypeFactory``. Both cases fall through the # payload-class map's ``.get`` fallback to # ``Ikev*PayloadUnparsed`` so the payload walk continues. payload_class = payload_classes.get(next_payload, unparsed_class) parser_payload.parse_parsable('payload', payload_class) payload = parser_payload['payload'] if isinstance(payload, (Ikev1PayloadUnparsed, Ikev2PayloadUnparsed)): # Preserve the on-wire payload-type code so # ``compose()`` can reproduce it in the previous # payload's ``next_payload`` field. payload.payload_type = next_payload payloads.append(payload) next_payload = payload.next_payload return payloads, parser_payload.parsed_length @classmethod def _parse(cls, parsable): header, body_bytes, next_payload = cls.parse_header_and_body(parsable) payloads, consumed_length = cls.parse_payload_chain(header['version'], body_bytes, next_payload) return cls(payloads=payloads, **header), cls.HEADER_SIZE + consumed_length @classmethod def compose_payload_chain(cls, version, payloads): """Compose the payload chain: link each payload's ``next_payload`` to its successor's type code and concatenate the encoded bytes. """ if version.major == IkeVersion.V1: payload_none = Ikev1PayloadType.NONE elif version.major == IkeVersion.V2: payload_none = Ikev2PayloadType.NONE else: raise NotImplementedError(version) payload_composer = ComposerBinary() for i, payload in enumerate(payloads): payload.next_payload = ( payload_none if i == len(payloads) - 1 else payloads[i + 1].get_payload_type() ) payload_composer.compose_parsable(payload) return payload_composer.composed_bytes @classmethod def compose_header( # pylint: disable=too-many-arguments,too-many-positional-arguments cls, initiator_spi, responder_spi, version, first_payload_type, exchange_type, flags, message_id, ike_total_length, ): """Compose the 28-octet IKE message header (RFC 7296 §3.1). Standalone so callers can build an outer header without touching the payload chain — needed by IKEv2 SK-envelope compose paths that supply the encrypted body directly and must not have ``payload.next_payload`` overwritten (see :meth:`compose_payload_chain`). """ header_composer = ComposerBinary() header_composer.compose_numeric(initiator_spi, 8) header_composer.compose_numeric(responder_spi, 8) header_composer.compose_numeric_enum_coded(first_payload_type) header_composer.compose_parsable(version) header_composer.compose_numeric_enum_coded(exchange_type) header_composer.compose_numeric_flags(flags, 1) header_composer.compose_numeric(message_id, 4) header_composer.compose_numeric(ike_total_length, 4) return header_composer.composed_bytes def compose(self): if self.version.major == IkeVersion.V1: payload_none = Ikev1PayloadType.NONE elif self.version.major == IkeVersion.V2: payload_none = Ikev2PayloadType.NONE else: raise NotImplementedError(self.version) payload_bytes = self.compose_payload_chain(self.version, self.payloads) first_payload_type = ( self.payloads[0].get_payload_type() if self.payloads else payload_none ) header_bytes = self.compose_header( self.initiator_spi, self.responder_spi, self.version, first_payload_type, self.exchange_type, self.flags, self.message_id, self.HEADER_SIZE + len(payload_bytes), ) return header_bytes + payload_bytes cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ike/version.py000066400000000000000000000044141524413560000274210ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 """ISAKMP version handling.""" import abc import attr from cryptodatahub.common.grade import GradeableSimple, Grade from cryptodatahub.ike.version import IkeVersion from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary from cryptoparser.common.exception import NotEnoughData, InvalidType from cryptoparser.common.base import OneByteEnumParsable class IsakmpVersionFactory(OneByteEnumParsable): """ISAKMP version.""" @classmethod def get_enum_class(cls): return IkeVersion @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s(frozen=True, order=False) class IsakmpProtocolVersion(ParsableBase, GradeableSimple): """ISAKMP protocol version parser.""" HEADER_SIZE = 1 major = attr.ib(validator=attr.validators.instance_of(IkeVersion)) minor: int = attr.ib(validator=attr.validators.instance_of(int)) def __lt__(self, other): if self.major.value.code != other.major.value.code: return self.major.value.code < other.major.value.code return self.minor < other.minor @property def grade(self): if self.major == IkeVersion.V1: return Grade.DEPRECATED return Grade.SECURE def __str__(self): return f"IKEv{self.major.value.code} ({self.minor})" @classmethod def _parse(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric('version', 1) major_version = (parser['version'] >> 4) & 0x0f minor_version = parser['version'] & 0x0f if major_version not in [v.value.code for v in IkeVersion]: raise InvalidType() return cls( major=next(v for v in IkeVersion if v.value.code == major_version), minor=minor_version, ), parser.parsed_length def compose(self): composer = ComposerBinary() version = (self.major.value.code << 4) | self.minor composer.compose_numeric(version, 1) return composer.composed_bytes @property def version(self): """Get version as string.""" return f"{self.major.value.code}.{self.minor}" cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ssh/000077500000000000000000000000001524413560000254045ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ssh/__init__.py000066400000000000000000000000431524413560000275120ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ssh/key.py000066400000000000000000001115501524413560000265510ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import abc import base64 import binascii import collections import datetime import enum import itertools import textwrap from collections import OrderedDict import ipaddress import attr from cryptodatahub.common.algorithm import Authentication, Hash from cryptodatahub.common.parameter import ECParamWellKnown from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.key import ( PublicKey, PublicKeyParamsDsa, PublicKeyParamsEcdsa, PublicKeyParamsEddsa, PublicKeyParamsRsa, PublicKeySize, ) from cryptodatahub.common.utils import hash_bytes from cryptodatahub.ssh.algorithm import SshHostKeyAlgorithm, SshHostKeyType, SshEllipticCurveIdentifier from cryptoparser.common.base import ( FourByteEnumComposer, FourByteEnumParsable, Serializable, StringEnumParsable, VariantParsable, VectorParamParsable, VectorParamString, VectorParsable, VectorParsableDerived, VectorString, ) from cryptoparser.common.exception import InvalidType, NotEnoughData from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary, ComposerText from cryptoparser.common.x509 import PublicKeyX509 @attr.s class SshPublicKeyBase: host_key_algorithm = attr.ib( converter=SshHostKeyAlgorithm, validator=attr.validators.instance_of(SshHostKeyAlgorithm) ) public_key = attr.ib( validator=attr.validators.instance_of(PublicKey) ) _HEADER_SIZE = 4 @classmethod def get_host_key_algorithms(cls): raise NotImplementedError() @property @abc.abstractmethod def key_bytes(self): raise NotImplementedError() @classmethod def _fingerprint(cls, hash_type, key_bytes, prefix): digest = hash_bytes(hash_type, key_bytes) if hash_type == Hash.MD5: fingerprint = ':'.join(textwrap.wrap(binascii.hexlify(digest).decode('ascii'), 2)) else: fingerprint = base64.b64encode(digest).decode('ascii') return ':'.join((prefix, fingerprint)) @property def key_size(self): return PublicKeySize(self.public_key.key_type, self.public_key.key_size) @property def fingerprints(self): key_bytes = self.key_bytes return OrderedDict([ (hash_type, self._fingerprint(hash_type, key_bytes, prefix)) for hash_type, prefix in [(Hash.SHA2_256, 'SHA256'), (Hash.SHA1, 'SHA1'), (Hash.MD5, 'MD5')] ]) def host_key_asdict(self): known_hosts = base64.b64encode(self.key_bytes).decode('ascii') public_key_dict = ( [('key_type', self.host_key_algorithm.value.key_type.value)] + list(self.public_key._asdict().items()) + [('known_hosts', known_hosts)] ) public_key_dict = OrderedDict(public_key_dict) public_key_dict['fingerprints'] = self.fingerprints return public_key_dict def _asdict(self): return self.host_key_asdict() @classmethod def _parse_host_key_algorithm(cls, parsable): if len(parsable) < cls._HEADER_SIZE: raise NotEnoughData(cls._HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_string('host_key_algorithm', 4, 'ascii', SshHostKeyAlgorithm.from_code) if parser['host_key_algorithm'] not in cls.get_host_key_algorithms(): raise InvalidType() return parser def _compose_host_key_algorithm(self): composer = ComposerBinary() host_key_algorithm_bytes = self.host_key_algorithm.value.code.encode('ascii') composer.compose_bytes(host_key_algorithm_bytes, 4) return composer class SshHostKeyBase(SshPublicKeyBase): @classmethod @abc.abstractmethod def get_host_key_algorithms(cls): raise NotImplementedError() class SshHostKeyParserBase(ParsableBase): @classmethod @abc.abstractmethod def _parse_host_key_algorithm(cls, parsable): raise NotImplementedError() @abc.abstractmethod def _compose_host_key_algorithm(self): raise NotImplementedError() @classmethod @abc.abstractmethod def _parse_host_key(cls, parser): raise NotImplementedError() @abc.abstractmethod def _compose_host_key_params(self, composer): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = cls._parse_host_key_algorithm(parsable) host_key = cls._parse_host_key(parser) return cls(parser['host_key_algorithm'], host_key), parser.parsed_length def compose(self): composer = self._compose_host_key_algorithm() self._compose_host_key_params(composer) return composer.composed @attr.s class SshHostKeyDSSBase(SshHostKeyBase): @property @abc.abstractmethod def key_bytes(self): raise NotImplementedError() @classmethod def get_host_key_algorithms(cls): return filter( lambda host_key_algorithm: ( host_key_algorithm.value.key_type == SshHostKeyType.HOST_KEY and host_key_algorithm.value.signature.value.key_type == Authentication.DSS ), SshHostKeyAlgorithm ) @classmethod def _parse_host_key(cls, parser): for param_name in ['p', 'q', 'g', 'y']: parser.parse_ssh_mpint(param_name) public_key = PublicKey.from_params(PublicKeyParamsDsa( prime=parser['p'], generator=parser['g'], order=parser['q'], public_key_value=parser['y'], )) for param_name in ['p', 'q', 'g', 'y']: del parser[param_name] return public_key def _compose_host_key_params(self, composer): params = self.public_key.params for param_name in ['prime', 'order', 'generator', 'public_key_value']: value = getattr(params, param_name) composer.compose_ssh_mpint(value) def host_key_asdict(self): key_dict = OrderedDict([]) key_dict.update(SshHostKeyBase.host_key_asdict(self)) key_dict.update(OrderedDict([ (param_name, getattr(self, param_name)) for param_name in attr.fields_dict(SshHostKeyDSSBase).keys() ])) return key_dict @attr.s class SshHostKeyDSS(SshHostKeyDSSBase, SshHostKeyParserBase): @property def key_bytes(self): return self.compose() @attr.s class SshHostKeyRSABase(SshHostKeyBase): @property @abc.abstractmethod def key_bytes(self): raise NotImplementedError() @classmethod def get_host_key_algorithms(cls): return filter( lambda host_key_algorithm: ( host_key_algorithm.value.key_type == SshHostKeyType.HOST_KEY and host_key_algorithm.value.signature.value.key_type == Authentication.RSA ), SshHostKeyAlgorithm ) @classmethod def _parse_host_key(cls, parser): parser.parse_ssh_mpint('e') parser.parse_ssh_mpint('n') public_key = PublicKey.from_params(PublicKeyParamsRsa( modulus=parser['n'], public_exponent=parser['e'], )) del parser['e'] del parser['n'] return public_key def _compose_host_key_params(self, composer): params = self.public_key.params composer.compose_ssh_mpint(params.public_exponent) composer.compose_ssh_mpint(params.modulus) @attr.s class SshHostKeyRSA(SshHostKeyRSABase, SshHostKeyParserBase): @property def key_bytes(self): return self.compose() @attr.s class SshHostKeyECDSABase(SshHostKeyBase): @property @abc.abstractmethod def key_bytes(self): raise NotImplementedError() @classmethod def get_host_key_algorithms(cls): return filter( lambda host_key_algorithm: ( host_key_algorithm.value.key_type == SshHostKeyType.HOST_KEY and host_key_algorithm.value.signature.value.key_type == Authentication.ECDSA ), SshHostKeyAlgorithm ) @classmethod def _parse_host_key(cls, parser): parser.parse_string('curve_identifier', 4, 'ascii', SshEllipticCurveIdentifier.from_code) parser.parse_bytes('curve_data', 4) public_key = PublicKey.from_params(PublicKeyParamsEcdsa.from_octet_bit_string( parser['curve_identifier'].value.key_parameter, parser['curve_data'], )) del parser['curve_identifier'] del parser['curve_data'] return public_key def _compose_host_key_params(self, composer): key_parameter = self.public_key.params.key_parameter for elliptic_curve_identifier in SshEllipticCurveIdentifier: if elliptic_curve_identifier.value.key_parameter == key_parameter: composer.compose_string(elliptic_curve_identifier.value.code, 'ascii', 4) break else: raise NotImplementedError(key_parameter) composer.compose_bytes(self.public_key.params.octet_bit_string, 4) @attr.s class SshHostKeyECDSA(SshHostKeyECDSABase, SshHostKeyParserBase): @property def key_bytes(self): return self.compose() @attr.s class SshHostKeyEDDSABase(SshHostKeyBase): @property @abc.abstractmethod def key_bytes(self): raise NotImplementedError() @classmethod def get_host_key_algorithms(cls): return filter( lambda host_key_algorithm: ( host_key_algorithm.value.key_type == SshHostKeyType.HOST_KEY and host_key_algorithm.value.signature.value.key_type == Authentication.EDDSA ), SshHostKeyAlgorithm ) @classmethod def _parse_host_key(cls, parser): parser.parse_bytes('key_data', 4) public_key = PublicKey.from_params(PublicKeyParamsEddsa( key_parameter=ECParamWellKnown.CURVE25519, key_data=parser['key_data'], )) del parser['key_data'] return public_key def _compose_host_key_params(self, composer): composer.compose_bytes(self.public_key.params.key_data, 4) @attr.s class SshHostKeyEDDSA(SshHostKeyEDDSABase, SshHostKeyParserBase): @property def key_bytes(self): return self.compose() @attr.s(frozen=True) class SshCertTypeParams(Serializable): code = attr.ib(validator=attr.validators.instance_of(int)) name = attr.ib(validator=attr.validators.instance_of(str)) def _as_markdown(self, level): return self._markdown_result(self.name, level) class SshCertType(FourByteEnumComposer, enum.Enum): SSH_CERT_TYPE_USER = SshCertTypeParams(code=1, name='User') SSH_CERT_TYPE_HOST = SshCertTypeParams(code=2, name='Host') class SshCertTypeFactory(FourByteEnumParsable): @classmethod def get_enum_class(cls): return SshCertType @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s class SshCertSignature(ParsableBase): signature_type = attr.ib(validator=attr.validators.instance_of(SshHostKeyAlgorithm)) signature_data = attr.ib(attr.validators.instance_of((bytes, bytearray))) @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_string('signature_type', 4, 'ascii', SshHostKeyAlgorithm.from_code) parser.parse_bytes('signature_data', 4) return cls(**parser), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_string(self.signature_type.value.code, 'ascii', 4) composer.compose_bytes(self.signature_data, 4) return composer.composed @attr.s class SshString(ParsableBase): value = attr.ib(validator=attr.validators.instance_of(str)) @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_string('value', 4, 'ascii') return cls(**parser), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_string(self.value, 'ascii', 4) return composer.composed class SshCertValidPrincipals(VectorParsable): @classmethod def get_param(cls): return VectorParamParsable( item_class=SshString, fallback_class=None, min_byte_num=0, max_byte_num=2 ** 32 - 4 ) @attr.s(frozen=True) class SshCertExtensionParam: code = attr.ib(validator=attr.validators.instance_of(str)) critical = attr.ib(validator=attr.validators.instance_of(bool)) class SshCertExtensionName(StringEnumParsable, enum.Enum): FORCE_COMMAND = SshCertExtensionParam( code='force-command', critical=True, ) SOURCE_ADDRESS = SshCertExtensionParam( code='source-address', critical=True, ) NO_PRESENCE_REQUIRED = SshCertExtensionParam( code='no-presence-required', critical=False, ) PERMIT_X11_FORWARDING = SshCertExtensionParam( code='permit-X11-forwarding', critical=False, ) PERMIT_AGENT_FORWARDING = SshCertExtensionParam( code='permit-agent-forwarding', critical=False, ) PERMIT_PORT_FORWARDING = SshCertExtensionParam( code='permit-port-forwarding', critical=False, ) PERMIT_PTY = SshCertExtensionParam( code='permit-pty', critical=False, ) PERMIT_USER_RC = SshCertExtensionParam( code='permit-user-rc', critical=False, ) class SshCertConstraintVector(VectorParsableDerived): @classmethod def get_param(cls): return VectorParamParsable( item_class=SshCertExtensionParsed, fallback_class=SshCertExtensionUnparsed, min_byte_num=0, max_byte_num=2 ** 32 - 1 ) @attr.s class SshCertExtensionBase(ParsableBase): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s class SshCertExtensionUnparsed(SshCertExtensionBase): extension_name = attr.ib(validator=attr.validators.instance_of(str)) extension_data = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_string('extension_name', 4, 'ascii') parser.parse_bytes('extension_data', 4) return SshCertExtensionUnparsed(**parser), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_string(self.extension_name, 'ascii', 4) payload_composer.compose_bytes(self.extension_data, 4) return payload_composer.composed_bytes @attr.s class SshCertExtensionParsed(SshCertExtensionBase): extension_name = attr.ib(init=False, validator=attr.validators.instance_of(SshCertExtensionName)) def __attrs_post_init__(self): self.extension_name = self.get_extension_name() attr.validate(self) @classmethod @abc.abstractmethod def get_extension_name(cls): raise NotImplementedError() @classmethod def _parse_header(cls, parsable): header_parser = ParserBinary(parsable) header_parser.parse_parsable('extension_name', SshCertExtensionName, 4) if header_parser['extension_name'] != cls.get_extension_name(): raise InvalidType() return header_parser def _compose_header(self): header_composer = ComposerBinary() header_composer.compose_string(self.extension_name.value.code, 'ascii', 4) return header_composer class SshCertExtensionNoData(SshCertExtensionParsed): @classmethod @abc.abstractmethod def get_extension_name(cls): raise NotImplementedError() @classmethod def _parse_header(cls, parsable): header_parser = super()._parse_header(parsable) header_parser.parse_numeric('extension_length', 4) return header_parser def _compose_header(self): header_composer = super()._compose_header() header_composer.compose_numeric(0, 4) return header_composer @classmethod def _parse(cls, parsable): header_parser = cls._parse_header(parsable) return cls(), header_parser.parsed_length def compose(self): header_composer = self._compose_header() return header_composer.composed_bytes class SshCertExtensionNoPrecenseRequired(SshCertExtensionNoData): @classmethod def get_extension_name(cls): return SshCertExtensionName.NO_PRESENCE_REQUIRED class SshCertExtensionPermitX11Forwarding(SshCertExtensionNoData): @classmethod def get_extension_name(cls): return SshCertExtensionName.PERMIT_X11_FORWARDING class SshCertExtensionPermitAgentForwarding(SshCertExtensionNoData): @classmethod def get_extension_name(cls): return SshCertExtensionName.PERMIT_AGENT_FORWARDING class SshCertExtensionPermitPortForwarding(SshCertExtensionNoData): @classmethod def get_extension_name(cls): return SshCertExtensionName.PERMIT_PORT_FORWARDING class SshCertExtensionPermitPTY(SshCertExtensionNoData): @classmethod def get_extension_name(cls): return SshCertExtensionName.PERMIT_PTY class SshCertExtensionPermitUserRC(SshCertExtensionNoData): @classmethod def get_extension_name(cls): return SshCertExtensionName.PERMIT_USER_RC @attr.s class SshCertExtensionForceCommand(SshCertExtensionParsed): command = attr.ib(validator=attr.validators.instance_of(str)) @classmethod def get_extension_name(cls): return SshCertExtensionName.FORCE_COMMAND @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_string('command', 4, 'ascii') return cls(parser['command']), parser.parsed_length def compose(self): body_composer = ComposerBinary() body_composer.compose_string(self.command, 'ascii', 4) header_composer = self._compose_header() return header_composer.composed_bytes + body_composer.composed_bytes class VectorParamNetorkAddress(VectorParamString): def get_item_size(self, item): return len(str(item)) class NetworkVector(VectorString): @classmethod def get_param(cls): return VectorParamNetorkAddress( min_byte_num=0, max_byte_num=2 ** 32 - 1, separator=',', item_class=ipaddress.ip_network, fallback_class=None, ) def compose(self): composer = ComposerBinary() address_composer = ComposerText() address_composer.compose_string_array(self._items) composer.compose_bytes(address_composer.composed, 4) return composer.composed @attr.s class SshCertExtensionSourceAddress(SshCertExtensionParsed): addresses = attr.ib( converter=NetworkVector, validator=attr.validators.instance_of(NetworkVector) ) @classmethod def get_extension_name(cls): return SshCertExtensionName.SOURCE_ADDRESS @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_parsable('addresses', NetworkVector) return cls(parser['addresses']), parser.parsed_length def compose(self): composer = self._compose_header() composer.compose_parsable(self.addresses) return composer.composed_bytes class SshCertConstraintVariant(VariantParsable): _VARIANTS = collections.OrderedDict([ (SshCertExtensionName.NO_PRESENCE_REQUIRED, (SshCertExtensionNoPrecenseRequired, )), (SshCertExtensionName.PERMIT_X11_FORWARDING, (SshCertExtensionPermitX11Forwarding, )), (SshCertExtensionName.PERMIT_AGENT_FORWARDING, (SshCertExtensionPermitAgentForwarding, )), (SshCertExtensionName.PERMIT_PORT_FORWARDING, (SshCertExtensionPermitPortForwarding, )), (SshCertExtensionName.PERMIT_PTY, (SshCertExtensionPermitPTY, )), (SshCertExtensionName.PERMIT_USER_RC, (SshCertExtensionPermitUserRC, )), (SshCertExtensionName.FORCE_COMMAND, (SshCertExtensionForceCommand, )), (SshCertExtensionName.SOURCE_ADDRESS, (SshCertExtensionSourceAddress, )), ]) @classmethod @abc.abstractmethod def _get_variants(cls): raise NotImplementedError() class SshCertificateBase: @classmethod @abc.abstractmethod def _parse_host_key_algorithm(cls, parsable): raise NotImplementedError() @abc.abstractmethod def _compose_host_key_algorithm(self): raise NotImplementedError() @classmethod @abc.abstractmethod def _parse_host_key(cls, parser): raise NotImplementedError() @abc.abstractmethod def _compose_host_key_params(self, composer): raise NotImplementedError() def host_key_asdict(self): return attr.asdict(self, recurse=False, dict_factory=OrderedDict) @attr.s class SshHostCertificateV00Base(ParsableBase, SshCertificateBase): # pylint: disable=too-many-instance-attributes certificate_type = attr.ib(validator=attr.validators.instance_of(SshCertType)) key_id = attr.ib( validator=attr.validators.instance_of(str), metadata={'human_readable_name': 'Key ID'}, ) valid_principals = attr.ib( converter=SshCertValidPrincipals, validator=attr.validators.instance_of(SshCertValidPrincipals) ) valid_after = attr.ib(validator=attr.validators.instance_of(datetime.datetime)) valid_before = attr.ib(validator=attr.validators.optional(attr.validators.instance_of(datetime.datetime))) constraints = attr.ib( converter=SshCertConstraintVector, validator=attr.validators.instance_of(SshCertConstraintVector) ) nonce = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) reserved = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) signature_key = attr.ib(validator=attr.validators.instance_of(SshPublicKeyBase)) signature = attr.ib(validator=attr.validators.instance_of(SshCertSignature)) @classmethod @abc.abstractmethod def _parse_host_key_algorithm(cls, parsable): raise NotImplementedError() @abc.abstractmethod def _compose_host_key_algorithm(self): raise NotImplementedError() @classmethod @abc.abstractmethod def _parse_host_key(cls, parser): raise NotImplementedError() @abc.abstractmethod def _compose_host_key_params(self, composer): raise NotImplementedError() @classmethod def _parse_host_cert_params(cls, parser): parser.parse_parsable('certificate_type', SshCertTypeFactory) parser.parse_string('key_id', 4, 'ascii') parser.parse_parsable('valid_principals', SshCertValidPrincipals) parser.parse_timestamp('valid_after') parser.parse_timestamp('valid_before') parser.parse_parsable('constraints', SshCertConstraintVector) parser.parse_bytes('nonce', 4) parser.parse_bytes('reserved', 4) parser.parse_parsable('signature_key', SshHostPublicKeyVariant, 4) parser.parse_parsable('signature', SshCertSignature, 4) def _compose_host_cert_params(self, composer): composer.compose_parsable(self.certificate_type) composer.compose_string(self.key_id, 'ascii', 4) composer.compose_parsable(self.valid_principals) composer.compose_timestamp(self.valid_after) composer.compose_timestamp(self.valid_before) composer.compose_parsable(self.constraints) composer.compose_bytes(self.nonce, 4) composer.compose_bytes(self.reserved, 4) composer.compose_parsable(self.signature_key, 4) composer.compose_parsable(self.signature, 4) @classmethod def _parse(cls, parsable): parser = cls._parse_host_key_algorithm(parsable) public_key = cls._parse_host_key(parser) cls._parse_host_cert_params(parser) return cls(public_key=public_key, **parser), parser.parsed_length def compose(self): composer = self._compose_host_key_algorithm() self._compose_host_key_params(composer) self._compose_host_cert_params(composer) return composer.composed class SshHostCertificateBase: def _asdict(self): key_dict = OrderedDict([]) key_dict.update(SshPublicKeyBase.host_key_asdict(self)) key_dict.update(SshCertificateBase.host_key_asdict(self)) return key_dict @attr.s class SshHostCertificateV00DSSBase(SshHostCertificateBase, SshHostKeyDSSBase, SshHostCertificateV00Base): @property def key_bytes(self): return self.compose() class SshHostCertificateV00DSS(SshHostCertificateV00DSSBase, SshHostCertificateV00Base): @classmethod def get_host_key_algorithms(cls): return [SshHostKeyAlgorithm.SSH_DSS_CERT_V00_OPENSSH_COM, ] @attr.s class SshHostCertificateV00RSABase(SshHostCertificateBase, SshHostKeyRSABase, SshHostCertificateV00Base): @property def key_bytes(self): return self.compose() @attr.s class SshHostCertificateV00RSA(SshHostCertificateV00RSABase, SshHostCertificateV00Base): @classmethod def get_host_key_algorithms(cls): return [SshHostKeyAlgorithm.SSH_RSA_CERT_V00_OPENSSH_COM, ] class SshCertExtensionVariant(SshCertConstraintVariant): @classmethod def _get_variants(cls): return collections.OrderedDict([ (constraint_name, constraint_classes) for constraint_name, constraint_classes in cls._VARIANTS.items() if constraint_name.value.critical is False ]) class SshCertExtensionVector(VectorParsableDerived): @classmethod def get_param(cls): return VectorParamParsable( item_class=SshCertExtensionVariant, fallback_class=SshCertExtensionUnparsed, min_byte_num=0, max_byte_num=2 ** 32 - 1 ) class SshCertCriticalOptionVariant(SshCertConstraintVariant): @classmethod def _get_variants(cls): return collections.OrderedDict([ (constraint_name, constraint_classes) for constraint_name, constraint_classes in cls._VARIANTS.items() if constraint_name.value.critical is True ]) class SshCertCriticalOptionVector(VectorParsableDerived): @classmethod def get_param(cls): return VectorParamParsable( item_class=SshCertCriticalOptionVariant, fallback_class=SshCertExtensionUnparsed, min_byte_num=0, max_byte_num=2 ** 32 - 1 ) @attr.s class SshHostCertificateV01Base(ParsableBase, SshCertificateBase): # pylint: disable=too-many-instance-attributes nonce = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) serial = attr.ib(validator=attr.validators.instance_of(int)) certificate_type = attr.ib(validator=attr.validators.instance_of(SshCertType)) key_id = attr.ib(validator=attr.validators.instance_of(str)) valid_principals = attr.ib( converter=SshCertValidPrincipals, validator=attr.validators.instance_of(SshCertValidPrincipals) ) valid_after = attr.ib(validator=attr.validators.instance_of(datetime.datetime)) valid_before = attr.ib(validator=attr.validators.optional(attr.validators.instance_of(datetime.datetime))) critical_options = attr.ib( converter=SshCertCriticalOptionVector, validator=attr.validators.instance_of(SshCertCriticalOptionVector) ) extensions = attr.ib( converter=SshCertExtensionVector, validator=attr.validators.instance_of(SshCertExtensionVector) ) reserved = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) signature_key = attr.ib(validator=attr.validators.instance_of(SshPublicKeyBase)) signature = attr.ib( validator=attr.validators.instance_of(SshCertSignature), metadata={'human_friendly': False}, ) @classmethod @abc.abstractmethod def _parse_host_key_algorithm(cls, parsable): raise NotImplementedError() @abc.abstractmethod def _compose_host_key_algorithm(self): raise NotImplementedError() @classmethod @abc.abstractmethod def _parse_host_key(cls, parser): raise NotImplementedError() @abc.abstractmethod def _compose_host_key_params(self, composer): raise NotImplementedError() @classmethod def _parse_host_cert_params(cls, parser): parser.parse_numeric('serial', 8) parser.parse_parsable('certificate_type', SshCertTypeFactory) parser.parse_string('key_id', 4, 'ascii') parser.parse_parsable('valid_principals', SshCertValidPrincipals) parser.parse_timestamp('valid_after') parser.parse_timestamp('valid_before') parser.parse_parsable('critical_options', SshCertCriticalOptionVector) parser.parse_parsable('extensions', SshCertExtensionVector) parser.parse_bytes('reserved', 4) parser.parse_parsable('signature_key', SshHostPublicKeyVariant, 4) parser.parse_parsable('signature', SshCertSignature, 4) def _compose_host_cert_params(self, composer): composer.compose_numeric(self.serial, 8) composer.compose_parsable(self.certificate_type) composer.compose_string(self.key_id, 'ascii', 4) composer.compose_parsable(self.valid_principals) composer.compose_timestamp(self.valid_after) composer.compose_timestamp(self.valid_before) composer.compose_parsable(self.critical_options) composer.compose_parsable(self.extensions) composer.compose_bytes(self.reserved, 4) composer.compose_parsable(self.signature_key, 4) composer.compose_parsable(self.signature, 4) @classmethod def _parse(cls, parsable): parser = cls._parse_host_key_algorithm(parsable) parser.parse_bytes('nonce', 4) public_key = cls._parse_host_key(parser) cls._parse_host_cert_params(parser) return cls(public_key=public_key, **parser), parser.parsed_length def compose(self): composer = self._compose_host_key_algorithm() composer.compose_bytes(self.nonce, 4) self._compose_host_key_params(composer) self._compose_host_cert_params(composer) return composer.composed @attr.s class SshHostCertificateV01DSSBase(SshHostCertificateBase, SshHostKeyDSSBase, SshHostCertificateV01Base): @property def key_bytes(self): return self.compose() class SshHostCertificateV01DSS(SshHostCertificateV01DSSBase, SshHostCertificateV01Base): @classmethod def get_host_key_algorithms(cls): return [SshHostKeyAlgorithm.SSH_DSS_CERT_V01_OPENSSH_COM, ] @attr.s class SshHostCertificateV01RSABase(SshHostCertificateBase, SshHostKeyRSABase, SshHostCertificateV01Base): @property def key_bytes(self): return self.compose() @attr.s class SshHostCertificateV01RSA(SshHostCertificateV01RSABase, SshHostCertificateV01Base): @classmethod def get_host_key_algorithms(cls): return [SshHostKeyAlgorithm.SSH_RSA_CERT_V01_OPENSSH_COM, ] @attr.s class SshHostCertificateV01ECDSABase(SshHostCertificateBase, SshHostKeyECDSABase, SshHostCertificateV01Base): @property def key_bytes(self): return self.compose() class SshHostCertificateV01ECDSA(SshHostCertificateV01ECDSABase, SshHostCertificateV01Base): @classmethod def get_host_key_algorithms(cls): return [ SshHostKeyAlgorithm.ECDSA_SHA2_NISTP256_CERT_V01_OPENSSH_COM, SshHostKeyAlgorithm.ECDSA_SHA2_NISTP384_CERT_V01_OPENSSH_COM, SshHostKeyAlgorithm.ECDSA_SHA2_NISTP521_CERT_V01_OPENSSH_COM, SshHostKeyAlgorithm.ECDSA_SHA2_SECP256K1_OID_CERT_V01_OPENSSH_COM, ] @attr.s class SshHostCertificateV01EDDSABase(SshHostCertificateBase, SshHostKeyEDDSABase, SshHostCertificateV01Base): @property def key_bytes(self): return self.compose() @attr.s class SshHostCertificateV01EDDSA(SshHostCertificateV01EDDSABase, SshHostCertificateV01Base): @classmethod def get_host_key_algorithms(cls): return [SshHostKeyAlgorithm.SSH_ED25519_CERT_V01_OPENSSH_COM, ] @attr.s class SshX509Certificate(ParsableBase, SshHostKeyBase): _NOT_DEFINED_HOST_KEY_ALGORITHMS_BY_PUBLIC_KEY_TYPE = { Authentication.RSA: SshHostKeyAlgorithm.X509V3_SIGN_RSA, Authentication.DSS: SshHostKeyAlgorithm.X509V3_SIGN_DSS, } @classmethod def get_host_key_algorithms(cls): return filter( lambda host_key_algorithm: host_key_algorithm.value.key_type == SshHostKeyType.X509_CERTIFICATE, SshHostKeyAlgorithm ) @property def key_bytes(self): return self.public_key.public_key.der @classmethod def _parse(cls, parsable): try: parser = cls._parse_host_key_algorithm(parsable) except (InvalidValue, NotEnoughData): parser = ParserBinary(parsable) host_key_algorithm = None else: host_key_algorithm = parser['host_key_algorithm'] if host_key_algorithm is None: public_key_length = parser.unparsed_length else: parser.parse_numeric('public_key_length', 4) public_key_length = parser['public_key_length'] try: parser.parse_raw('public_key', public_key_length, PublicKeyX509.from_der) except (NotEnoughData, InvalidValue) as e: raise InvalidValue(parser.unparsed, cls, 'public_key') from e public_key = parser['public_key'] if host_key_algorithm is None: host_key_algorithm = cls._NOT_DEFINED_HOST_KEY_ALGORITHMS_BY_PUBLIC_KEY_TYPE.get(public_key.key_type, None) if host_key_algorithm is None: raise InvalidType() return cls(host_key_algorithm, public_key), parser.parsed_length def compose(self): public_key_bytes = self.public_key.der if self.host_key_algorithm in [SshHostKeyAlgorithm.X509V3_SIGN_RSA, SshHostKeyAlgorithm.X509V3_SIGN_DSS]: composer = ComposerBinary() else: composer = self._compose_host_key_algorithm() composer.compose_numeric(len(public_key_bytes), 4) composer.compose_raw(public_key_bytes) return composer.composed @attr.s class SshX509CertificateChain(ParsableBase, SshHostKeyBase): issuer_certificates = attr.ib( validator=attr.validators.deep_iterable(member_validator=attr.validators.instance_of(PublicKey)) ) ocsp_responses = attr.ib( validator=attr.validators.deep_iterable(member_validator=attr.validators.instance_of((bytes, bytearray))) ) @classmethod def get_host_key_algorithms(cls): return filter( lambda host_key_algorithm: host_key_algorithm.value.key_type == SshHostKeyType.X509_CERTIFICATE_CHAIN, SshHostKeyAlgorithm ) @property def key_bytes(self): return self.public_key.key_bytes @classmethod def _parse(cls, parsable): parser = cls._parse_host_key_algorithm(parsable) parser.parse_numeric('certificate_count', 4) certificates = [] for _ in range(parser['certificate_count']): parser.parse_bytes('certificate', 4) certificates.append(PublicKeyX509.from_der(bytes(parser['certificate']))) parser.parse_numeric('ocsp_response_count', 4) ocsp_responses = [] for _ in range(parser['ocsp_response_count']): parser.parse_bytes('ocsp_response', 4) ocsp_responses.append(parser['ocsp_response']) return cls( parser['host_key_algorithm'], certificates[0], certificates[1:], ocsp_responses, ), parser.parsed_length def compose(self): composer = self._compose_host_key_algorithm() composer.compose_numeric(len(self.issuer_certificates) + 1, 4) for certificate in [self.public_key] + self.issuer_certificates: composer.compose_bytes(certificate.der, 4) composer.compose_numeric(len(self.ocsp_responses), 4) for ocsp_response in self.ocsp_responses: composer.compose_bytes(ocsp_response, 4) return composer.composed def _asdict(self): return collections.OrderedDict([ ('key_type', self.host_key_algorithm.value.key_type.value), ('certificate_chain', [self.public_key] + self.issuer_certificates), ]) class SshHostPublicKeyVariant(VariantParsable): _VARIANTS = collections.OrderedDict(itertools.chain.from_iterable([ [ (host_key_algorithm, (ssh_key_class, )) for host_key_algorithm in ssh_key_class.get_host_key_algorithms() ] for ssh_key_class in [ SshHostKeyDSS, SshHostKeyECDSA, SshHostKeyEDDSA, SshHostKeyRSA, SshHostCertificateV00DSS, SshHostCertificateV00RSA, SshHostCertificateV01DSS, SshHostCertificateV01ECDSA, SshHostCertificateV01EDDSA, SshHostCertificateV01RSA, SshX509Certificate, SshX509CertificateChain, ] ])) @classmethod def _get_variants(cls): return cls._VARIANTS cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ssh/record.py000066400000000000000000000042161524413560000272370ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import attr from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary from cryptoparser.common.exception import NotEnoughData from cryptoparser.ssh.subprotocol import ( SshMessageBase, SshMessageVariantInit, SshMessageVariantKexDH, SshMessageVariantKexDHGroup, ) @attr.s class SshRecordBase(ParsableBase): HEADER_SIZE = 6 packet = attr.ib(validator=attr.validators.instance_of(SshMessageBase)) @classmethod @abc.abstractmethod def _get_variant_class(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('packet_length', 4) if parser['packet_length'] > parser.unparsed_length: raise NotEnoughData(parser['packet_length'] - parser.unparsed_length) parser.parse_numeric('padding_length', 1) parser.parse_parsable('packet', cls._get_variant_class()) parser.parse_raw('padding', parser['padding_length']) return cls(packet=parser['packet']), parser.parsed_length def compose(self): body_composer = ComposerBinary() body_composer.compose_parsable(self.packet) payload_length = body_composer.composed_length padding_length = 8 - ((payload_length + 5) % 8) if padding_length < 4: padding_length += 8 packet_length = payload_length + padding_length + 1 for _ in range(padding_length): body_composer.compose_numeric(0, 1) header_composer = ComposerBinary() header_composer.compose_numeric(packet_length, 4) header_composer.compose_numeric(padding_length, 1) return header_composer.composed + body_composer.composed class SshRecordInit(SshRecordBase): @classmethod def _get_variant_class(cls): return SshMessageVariantInit class SshRecordKexDH(SshRecordBase): @classmethod def _get_variant_class(cls): return SshMessageVariantKexDH class SshRecordKexDHGroup(SshRecordBase): @classmethod def _get_variant_class(cls): return SshMessageVariantKexDHGroup cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ssh/subprotocol.py000066400000000000000000000471111524413560000303350ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import collections import enum import hashlib import random import attr from cryptodatahub.common.exception import InvalidValue from cryptodatahub.ssh.algorithm import ( SshKexAlgorithm, SshHostKeyAlgorithm, SshEncryptionAlgorithm, SshMacAlgorithm, SshCompressionAlgorithm, ) from cryptoparser.common.base import ( VariantParsable, VectorParamString, VectorString, ) from cryptoparser.common.classes import LanguageTag from cryptoparser.common.exception import InvalidType, TooMuchData from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary, ParserText, ComposerText from cryptoparser.common.utils import bytes_to_hex_string from cryptoparser.ssh.key import SshPublicKeyBase, SshHostPublicKeyVariant from cryptoparser.ssh.version import ( SshProtocolVersion, SshSoftwareVersionBase, SshSoftwareVersionParsedVariant, SshSoftwareVersionUnparsed, ) class SshMessageCode(enum.IntEnum): DISCONNECT = 0x1 IGNORE = 0x2 UNIMPLEMENTED = 0x3 DEBUG = 0x4 SERVICE_REQUEST = 0x5 SERVICE_ACCEPT = 0x6 KEXINIT = 0x14 NEWKEYS = 0x15 DH_KEX_INIT = 0x1e DH_KEX_REPLY = 0x1f DH_GEX_GROUP = 0x1f DH_GEX_INIT = 0x20 DH_GEX_REPLY = 0x21 DH_GEX_REQUEST = 0x22 class SshMessageBase(ParsableBase): @classmethod @abc.abstractmethod def get_message_code(cls): raise NotImplementedError() @classmethod def _parse_header(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('message_code', 1, SshMessageCode) if parser['message_code'] != cls.get_message_code(): raise InvalidType() return parser @classmethod def _compose_header(cls): composer = ComposerBinary() composer.compose_numeric(cls.get_message_code(), 1) return composer @attr.s class SshProtocolMessage(ParsableBase): protocol_version = attr.ib(validator=attr.validators.instance_of(SshProtocolVersion)) software_version = attr.ib(validator=attr.validators.instance_of(SshSoftwareVersionBase)) comment = attr.ib(validator=attr.validators.optional(attr.validators.instance_of(str)), default=None) @comment.validator def comment_validator(self, _, value): # pylint: disable=no-self-use if value is not None: if '\r' in value or '\n' in value: raise InvalidValue(value, SshProtocolMessage, 'comment') try: value.encode('ascii') except UnicodeEncodeError as e: raise InvalidValue(value, SshProtocolMessage, 'comment') from e @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_by_length('protocol', min_length=3, max_length=3) if parser['protocol'] != 'SSH': raise InvalidValue(parser['protocol'], SshProtocolMessage, 'protocol') parser.parse_string('separator', '-') parser.parse_parsable('protocol_version', SshProtocolVersion) parser.parse_string('separator', '-') parser.parse_string_until_separator('software_version_and_comment', '\n') software_version_and_comment = parser['software_version_and_comment'].split(' ') if software_version_and_comment[-1][-1] == '\r': software_version_and_comment[-1] = software_version_and_comment[-1][:-1] software_version_parser = ParserText(software_version_and_comment[0].encode('ascii')) try: software_version_parser.parse_parsable('value', SshSoftwareVersionParsedVariant) except InvalidValue: software_version_parser.parse_parsable('value', SshSoftwareVersionUnparsed) if len(software_version_and_comment) > 1: comment = ' '.join(software_version_and_comment[1:]) else: comment = None parser.parse_separator('\n') if parser.parsed_length > 255: raise TooMuchData(parser.parsed_length - 255) return SshProtocolMessage( parser['protocol_version'], software_version_parser['value'], comment ), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string('SSH') composer.compose_separator('-') composer.compose_parsable(self.protocol_version) composer.compose_separator('-') composer.compose_string(self.software_version) if self.comment is not None: composer.compose_separator(' ') composer.compose_string(self.comment) composer.compose_separator('\r\n') return composer.composed class SshAlgorithmVector(VectorString): @classmethod def get_param(cls): return VectorParamString( min_byte_num=0, max_byte_num=2 ** 32 - 1, separator=',', item_class=cls.get_item_class(), fallback_class=str ) @classmethod @abc.abstractmethod def get_item_class(cls): raise NotImplementedError() class SshKexAlgorithmVector(SshAlgorithmVector): @classmethod def get_item_class(cls): return SshKexAlgorithm class SshHostKeyAlgorithmVector(SshAlgorithmVector): @classmethod def get_item_class(cls): return SshHostKeyAlgorithm class SshEncryptionAlgorithmVector(SshAlgorithmVector): @classmethod def get_item_class(cls): return SshEncryptionAlgorithm class SshMacAlgorithmVector(SshAlgorithmVector): @classmethod def get_item_class(cls): return SshMacAlgorithm class SshCompressionAlgorithmVector(SshAlgorithmVector): @classmethod def get_item_class(cls): return SshCompressionAlgorithm class VectorParamSshLanguage(VectorParamString): def __init__(self): super().__init__( min_byte_num=0, max_byte_num=2 ** 32 - 1, separator=',', item_class=LanguageTag, fallback_class=None, ) def get_item_size(self, item): return len(item.compose()) class SshLanguageVector(VectorString): @classmethod def get_param(cls): return VectorParamSshLanguage() @attr.s class SshKeyExchangeInit(SshMessageBase): # pylint: disable=too-many-instance-attributes kex_algorithms = attr.ib( converter=SshKexAlgorithmVector, validator=attr.validators.instance_of(SshKexAlgorithmVector) ) host_key_algorithms = attr.ib( converter=SshHostKeyAlgorithmVector, validator=attr.validators.instance_of(SshHostKeyAlgorithmVector) ) encryption_algorithms_client_to_server = attr.ib( converter=SshEncryptionAlgorithmVector, validator=attr.validators.instance_of(SshEncryptionAlgorithmVector) ) encryption_algorithms_server_to_client = attr.ib( converter=SshEncryptionAlgorithmVector, validator=attr.validators.instance_of(SshEncryptionAlgorithmVector) ) mac_algorithms_client_to_server = attr.ib( converter=SshMacAlgorithmVector, validator=attr.validators.instance_of(SshMacAlgorithmVector) ) mac_algorithms_server_to_client = attr.ib( converter=SshMacAlgorithmVector, validator=attr.validators.instance_of(SshMacAlgorithmVector) ) compression_algorithms_client_to_server = attr.ib( converter=SshCompressionAlgorithmVector, validator=attr.validators.instance_of(SshCompressionAlgorithmVector) ) compression_algorithms_server_to_client = attr.ib( converter=SshCompressionAlgorithmVector, validator=attr.validators.instance_of(SshCompressionAlgorithmVector) ) languages_client_to_server = attr.ib( converter=SshLanguageVector, validator=attr.validators.instance_of(SshLanguageVector), default=() ) languages_server_to_client = attr.ib( converter=SshLanguageVector, validator=attr.validators.instance_of(SshLanguageVector), default=() ) first_kex_packet_follows = attr.ib(validator=attr.validators.instance_of(int), default=0) cookie = attr.ib( validator=attr.validators.instance_of((bytearray, bytes)), default=bytearray.fromhex(f'{random.getrandbits(128):16x}'.zfill(32)) ) reserved = attr.ib(validator=attr.validators.instance_of(int), default=0x00000000) @classmethod def _get_cipher_attributes(cls): for attribute in attr.fields(cls): if (attribute.validator and isinstance(attribute.validator.type, type) and issubclass(attribute.validator.type, ParsableBase)): yield attribute @classmethod def _parse(cls, parsable): header_parser = cls._parse_header(parsable) body_parser = ParserBinary(parsable[header_parser.parsed_length:]) body_parser.parse_raw('cookie', 16) for attribute in cls._get_cipher_attributes(): body_parser.parse_parsable(attribute.name, attribute.validator.type) body_parser.parse_numeric('first_kex_packet_follows', 1, bool) body_parser.parse_numeric('reserved', 4) return SshKeyExchangeInit(**dict(body_parser)), header_parser.parsed_length + body_parser.parsed_length def compose(self): composer = self._compose_header() composer.compose_raw(self.cookie) for attribute in self._get_cipher_attributes(): composer.compose_parsable(getattr(self, attribute.name)) composer.compose_numeric(1 if self.first_kex_packet_follows else 0, 1) composer.compose_numeric(self.reserved, 4) return composer.composed @classmethod def get_message_code(cls): return SshMessageCode.KEXINIT @staticmethod def _hassh(algorithm_vectors): hassh_text = ';'.join([ ','.join([ algorithm if isinstance(algorithm, str) else algorithm.value.code for algorithm in algorithms ]) for algorithms in algorithm_vectors ]) message = hashlib.md5() message.update(hassh_text.encode('ascii')) return bytes_to_hex_string(message.digest(), lowercase=True) @property def hassh(self): return self._hassh([ self.kex_algorithms, self.encryption_algorithms_client_to_server, self.mac_algorithms_client_to_server, self.compression_algorithms_client_to_server, ]) @property def hassh_server(self): return self._hassh([ self.kex_algorithms, self.encryption_algorithms_server_to_client, self.mac_algorithms_server_to_client, self.compression_algorithms_server_to_client, ]) class SshReasonCode(enum.IntEnum): HOST_NOT_ALLOWED_TO_CONNECT = 0x01 PROTOCOL_ERROR = 0x02 KEY_EXCHANGE_FAILED = 0x03 RESERVED = 0x04 MAC_ERROR = 0x05 COMPRESSION_ERROR = 0x06 SERVICE_NOT_AVAILABLE = 0x07 PROTOCOL_VERSION_NOT_SUPPORTED = 0x08 HOST_KEY_NOT_VERIFIABLE = 0x09 CONNECTION_LOST = 0x0a BY_APPLICATION = 0x0b TOO_MANY_CONNECTIONS = 0x0c AUTH_CANCELLED_BY_USER = 0x0d NO_MORE_AUTH_METHODS_AVAILABLE = 0x0e ILLEGAL_USER_NAME = 0x0f @attr.s class SshDisconnectMessage(SshMessageBase): reason = attr.ib(validator=attr.validators.instance_of(SshReasonCode)) description = attr.ib( converter=str, validator=attr.validators.instance_of(str) ) language = attr.ib( default='US', validator=attr.validators.instance_of(str) ) @classmethod def get_message_code(cls): return SshMessageCode.DISCONNECT @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric('reason', 4, SshReasonCode) parser.parse_string('description', 4, 'utf-8') parser.parse_string('language', 4, 'ascii') return SshDisconnectMessage(parser['reason'], parser['description'], parser['language']), parser.parsed_length def compose(self): composer = self._compose_header() composer.compose_numeric(self.reason.value, 4) composer.compose_string(self.description, 'utf-8', 4) composer.compose_string(self.language, 'ascii', 4) return composer.composed @attr.s class SshUnimplementedMessage(SshMessageBase): sequence_number = attr.ib(validator=attr.validators.instance_of(int)) @classmethod def get_message_code(cls): return SshMessageCode.UNIMPLEMENTED @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric('sequence_number', 4) return SshUnimplementedMessage(parser['sequence_number']), parser.parsed_length def compose(self): composer = self._compose_header() composer.compose_numeric(self.sequence_number, 4) return composer.composed @attr.s class SshDHKeyExchangeInitBase(SshMessageBase): ephemeral_public_key = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod @abc.abstractmethod def get_message_code(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): header_parser = cls._parse_header(parsable) body_parser = ParserBinary(parsable[header_parser.parsed_length:]) body_parser.parse_bytes('ephemeral_public_key', 4) return cls( body_parser['ephemeral_public_key'] ), header_parser.parsed_length + body_parser.parsed_length def compose(self): composer = self._compose_header() composer.compose_bytes(self.ephemeral_public_key, 4) return composer.composed class SshDHKeyExchangeInit(SshDHKeyExchangeInitBase): @classmethod def get_message_code(cls): return SshMessageCode.DH_KEX_INIT class SshDHGroupExchangeInit(SshDHKeyExchangeInitBase): @classmethod def get_message_code(cls): return SshMessageCode.DH_GEX_INIT @attr.s class SshDHKeyExchangeReplyBase(SshMessageBase): host_public_key = attr.ib(validator=attr.validators.instance_of(SshPublicKeyBase)) ephemeral_public_key = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) signature = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod @abc.abstractmethod def get_message_code(cls): raise NotImplementedError() @staticmethod def _parse_ssh_key(parser): parser.parse_parsable('host_public_key', SshHostPublicKeyVariant, 4) parser.parse_bytes('ephemeral_public_key', 4) parser.parse_bytes('signature', 4) return parser['host_public_key'], parser['ephemeral_public_key'], parser['signature'] @staticmethod def _compose_ssh_key(composer, host_public_key, ephemeral_public_key_bytes, signature_bytes): composer.compose_parsable(host_public_key, 4) composer.compose_bytes(ephemeral_public_key_bytes, 4) composer.compose_bytes(signature_bytes, 4) @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) host_public_key, ephemeral_public_key, signature = cls._parse_ssh_key(parser) return cls( host_public_key, ephemeral_public_key, signature, ), parser.parsed_length def compose(self): composer = self._compose_header() self._compose_ssh_key(composer, self.host_public_key, self.ephemeral_public_key, self.signature) return composer.composed class SshDHKeyExchangeReply(SshDHKeyExchangeReplyBase): @classmethod def get_message_code(cls): return SshMessageCode.DH_KEX_REPLY class SshDHGroupExchangeReply(SshDHKeyExchangeReplyBase): @classmethod def get_message_code(cls): return SshMessageCode.DH_GEX_REPLY @attr.s class SshDHGroupExchangeRequest(SshMessageBase): gex_min = attr.ib(validator=attr.validators.instance_of(int)) gex_number = attr.ib(validator=attr.validators.instance_of(int)) gex_max = attr.ib(validator=attr.validators.instance_of(int)) @classmethod def get_message_code(cls): return SshMessageCode.DH_GEX_REQUEST @classmethod def _parse(cls, parsable): header_parser = cls._parse_header(parsable) body_parser = ParserBinary(parsable[header_parser.parsed_length:]) body_parser.parse_numeric('gex_min', 4) body_parser.parse_numeric('gex_number', 4) body_parser.parse_numeric('gex_max', 4) return SshDHGroupExchangeRequest( body_parser['gex_min'], body_parser['gex_number'], body_parser['gex_max'], ), header_parser.parsed_length + body_parser.parsed_length def compose(self): composer = self._compose_header() composer.compose_numeric(self.gex_min, 4) composer.compose_numeric(self.gex_number, 4) composer.compose_numeric(self.gex_max, 4) return composer.composed @attr.s class SshDHGroupExchangeGroup(SshMessageBase): p = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) g = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def get_message_code(cls): return SshMessageCode.DH_GEX_GROUP @classmethod def _parse(cls, parsable): header_parser = cls._parse_header(parsable) body_parser = ParserBinary(parsable[header_parser.parsed_length:]) body_parser.parse_bytes('p', 4) body_parser.parse_bytes('g', 4) return SshDHGroupExchangeGroup( body_parser['p'], body_parser['g'], ), header_parser.parsed_length + body_parser.parsed_length def compose(self): composer = self._compose_header() composer.compose_bytes(self.p, 4) composer.compose_bytes(self.g, 4) return composer.composed @attr.s class SshNewKeys(SshMessageBase): @classmethod def get_message_code(cls): return SshMessageCode.NEWKEYS @classmethod def _parse(cls, parsable): header_parser = cls._parse_header(parsable) return SshNewKeys(), header_parser.parsed_length def compose(self): composer = self._compose_header() return composer.composed class SshMessageVariantInit(VariantParsable): _VARIANTS = collections.OrderedDict([ (SshMessageCode.DISCONNECT, (SshDisconnectMessage, )), (SshMessageCode.UNIMPLEMENTED, (SshUnimplementedMessage, )), (SshMessageCode.KEXINIT, (SshKeyExchangeInit, )), ]) @classmethod def _get_variants(cls): return cls._VARIANTS class SshMessageVariantKexDH(VariantParsable): _VARIANTS = collections.OrderedDict([ (SshMessageCode.DISCONNECT, (SshDisconnectMessage, )), (SshMessageCode.KEXINIT, (SshKeyExchangeInit, )), (SshMessageCode.UNIMPLEMENTED, (SshUnimplementedMessage, )), (SshMessageCode.DH_KEX_INIT, (SshDHKeyExchangeInit, )), (SshMessageCode.DH_KEX_REPLY, (SshDHKeyExchangeReply, )), (SshMessageCode.NEWKEYS, (SshNewKeys, )), ]) @classmethod def _get_variants(cls): return cls._VARIANTS class SshMessageVariantKexDHGroup(VariantParsable): _VARIANTS = collections.OrderedDict([ (SshMessageCode.DISCONNECT, (SshDisconnectMessage, )), (SshMessageCode.KEXINIT, (SshKeyExchangeInit, )), (SshMessageCode.UNIMPLEMENTED, (SshUnimplementedMessage, )), (SshMessageCode.DH_GEX_REQUEST, (SshDHGroupExchangeRequest, )), (SshMessageCode.DH_GEX_GROUP, (SshDHGroupExchangeGroup, )), (SshMessageCode.DH_GEX_INIT, (SshDHGroupExchangeInit, )), (SshMessageCode.DH_GEX_REPLY, (SshDHGroupExchangeReply, )), (SshMessageCode.NEWKEYS, (SshNewKeys, )), ]) @classmethod def _get_variants(cls): return cls._VARIANTS cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/ssh/version.py000066400000000000000000000143111524413560000274430ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import collections import enum import attr from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.grade import Grade, GradeableSimple from cryptoparser.common.base import ProtocolVersionBase, Serializable, VariantParsable from cryptoparser.common.exception import InvalidType from cryptoparser.common.parse import ParsableBase, ParserText, ComposerText class SshVersion(enum.IntEnum): SSH1 = 1 SSH2 = 2 @attr.s(hash=True) class SshProtocolVersion(ProtocolVersionBase, GradeableSimple): major = attr.ib(converter=SshVersion, validator=attr.validators.instance_of(SshVersion)) minor = attr.ib(validator=attr.validators.instance_of(int), default=0) @property def grade(self): if self.major == SshVersion.SSH1: return Grade.INSECURE return Grade.SECURE def __str__(self): return f'SSH {self.major}.{self.minor}' @classmethod def _parse(cls, parsable): parser = ParserText(parsable) try: parser.parse_numeric('major') parser.parse_separator('.') parser.parse_numeric('minor') except InvalidValue as e: raise InvalidValue(parsable, SshProtocolVersion) from e return SshProtocolVersion(parser['major'], parser['minor']), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_numeric(self.major.value) composer.compose_separator('.') composer.compose_numeric(self.minor) return composer.composed @property def identifier(self): return f'ssh{self.major}' @property def supported_versions(self): if self.major == SshVersion.SSH1 and self.minor == 99: return [SshVersion.SSH1, SshVersion.SSH2] return [self.major, ] class SshSoftwareVersionBase(ParsableBase): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s class SshSoftwareVersionUnparsed(SshSoftwareVersionBase, Serializable): raw = attr.ib(validator=attr.validators.instance_of(str)) @raw.validator def raw_validator(self, _, value): # pylint: disable=no-self-use if '\r' in value or '\n' in value or ' ' in value: raise InvalidValue(value, SshSoftwareVersionUnparsed, 'raw') try: value.encode('ascii') except UnicodeEncodeError as e: raise InvalidValue(value, SshSoftwareVersionUnparsed, 'raw') from e def _as_markdown(self, level): return self._markdown_result(self.raw, level) @classmethod def _parse(cls, parsable): parser = ParserText(parsable) parser.parse_string_by_length('raw', len(parsable)) return SshSoftwareVersionUnparsed(parser['raw']), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string(self.raw) return composer.composed @attr.s class SshSoftwareVersionParsedBase(SshSoftwareVersionBase): version = attr.ib(default=None, validator=attr.validators.optional(attr.validators.instance_of(str))) @classmethod @abc.abstractmethod def _get_vendor(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_version_separator(cls): raise NotImplementedError() @property def vendor(self): return self._get_vendor() def _asdict(self): return collections.OrderedDict([ ('vendor', self.vendor), ('version', self.version), ]) @classmethod def _parse(cls, parsable): parser = ParserText(parsable) version_separator = cls._get_version_separator() if version_separator is None: parser.parse_string_by_length('vendor') else: parser.parse_string_until_separator_or_end('vendor', version_separator) if parser['vendor'] != cls._get_vendor(): raise InvalidType() if parser.unparsed_length > 0 and version_separator is not None: parser.parse_separator(version_separator) parser.parse_string_by_length('version') version = parser['version'] else: version = None return cls(version), parser.parsed_length def compose(self): composer = ComposerText() composer.compose_string(self.vendor) if self.version is not None: composer.compose_separator(self._get_version_separator()) composer.compose_string(self.version) return composer.composed @attr.s class SshSoftwareVersionCryptlib(SshSoftwareVersionParsedBase): @classmethod def _get_vendor(cls): return 'cryptlib' @classmethod def _get_version_separator(cls): return None @attr.s class SshSoftwareVersionDropbear(SshSoftwareVersionParsedBase): @classmethod def _get_vendor(cls): return 'dropbear' @classmethod def _get_version_separator(cls): return '_' @attr.s class SshSoftwareVersionIPSSH(SshSoftwareVersionParsedBase): @classmethod def _get_vendor(cls): return 'IPSSH' @classmethod def _get_version_separator(cls): return '-' @attr.s class SshSoftwareVersionMonacaSSH(SshSoftwareVersionParsedBase): @classmethod def _get_vendor(cls): return 'Monaca' @classmethod def _get_version_separator(cls): return None @attr.s class SshSoftwareVersionOpenSSH(SshSoftwareVersionParsedBase): @classmethod def _get_vendor(cls): return 'OpenSSH' @classmethod def _get_version_separator(cls): return '_' class SshSoftwareVersionParsedVariant(VariantParsable): _VARIANTS = collections.OrderedDict([ (SshSoftwareVersionCryptlib, (SshSoftwareVersionCryptlib, )), (SshSoftwareVersionDropbear, (SshSoftwareVersionDropbear, )), (SshSoftwareVersionIPSSH, (SshSoftwareVersionIPSSH, )), (SshSoftwareVersionMonacaSSH, (SshSoftwareVersionMonacaSSH, )), (SshSoftwareVersionOpenSSH, (SshSoftwareVersionOpenSSH, )), ]) @classmethod def _get_variants(cls): return cls._VARIANTS cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/000077500000000000000000000000001524413560000254115ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/__init__.py000066400000000000000000000000431524413560000275170ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/algorithm.py000066400000000000000000000015771524413560000277630ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc from cryptodatahub.tls.algorithm import TlsECPointFormat, TlsNamedCurve, TlsSignatureAndHashAlgorithm from cryptoparser.common.base import OneByteEnumParsable, TwoByteEnumParsable class TlsNamedCurveFactory(TwoByteEnumParsable): @classmethod def get_enum_class(cls): return TlsNamedCurve @abc.abstractmethod def compose(self): raise NotImplementedError() class TlsSignatureAndHashAlgorithmFactory(TwoByteEnumParsable): @classmethod def get_enum_class(cls): return TlsSignatureAndHashAlgorithm @abc.abstractmethod def compose(self): raise NotImplementedError() class TlsECPointFormatFactory(OneByteEnumParsable): @classmethod def get_enum_class(cls): return TlsECPointFormat @abc.abstractmethod def compose(self): raise NotImplementedError() cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/ciphersuite.py000066400000000000000000000012201524413560000303020ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import abc from cryptodatahub.tls.algorithm import TlsCipherSuite, SslCipherKind from cryptoparser.common.base import TwoByteEnumParsable, ThreeByteEnumParsable class TlsCipherSuiteFactory(TwoByteEnumParsable): @classmethod def get_enum_class(cls): return TlsCipherSuite @abc.abstractmethod def compose(self): raise NotImplementedError() class SslCipherKindFactory(ThreeByteEnumParsable): @classmethod def get_enum_class(cls): return SslCipherKind @abc.abstractmethod def compose(self): raise NotImplementedError() cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/extension.py000066400000000000000000001270721524413560000300100ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import abc import collections import enum import attr from cryptodatahub.common.exception import InvalidValue from cryptodatahub.tls.algorithm import ( TlsCertificateCompressionAlgorithm, TlsExtensionType, TlsNamedCurve, TlsNextProtocolName, TlsProtocolName, TlsPskKeyExchangeMode, TlsTokenBindingParamater, ) from cryptoparser.common.base import ( OneByteEnumParsable, Opaque, OpaqueParam, ProtocolVersionMajorMinorBase, OpaqueEnumParsable, TwoByteEnumParsable, VariantParsable, Vector, VectorEnumCodeNumeric, VectorEnumCodeString, VectorParamEnumCodeNumeric, VectorParamEnumCodeString, VectorParamNumeric, VectorParamParsable, VectorParsable, VectorParsableDerived, ) from cryptoparser.common.exception import NotEnoughData, InvalidType from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary from cryptoparser.common.x509 import SignedCertificateTimestampList from cryptoparser.tls.algorithm import ( TlsECPointFormatFactory, TlsNamedCurveFactory, TlsSignatureAndHashAlgorithmFactory, ) from cryptoparser.tls.grease import TlsInvalidTypeOneByte, TlsInvalidTypeTwoByte from cryptoparser.tls.version import TlsProtocolVersion class TlsExtensionTypeFactory(TwoByteEnumParsable): @classmethod def get_enum_class(cls): return TlsExtensionType @abc.abstractmethod def compose(self): raise NotImplementedError() class TlsExtensionsBase(VectorParsable): @classmethod @abc.abstractmethod def get_param(cls): raise NotImplementedError() def get_item_by_type(self, extension_type): try: item = next( extension for extension in self if extension.extension_type == extension_type ) except StopIteration as e: raise KeyError from e return item class TlsExtensionsClient(TlsExtensionsBase): @classmethod def get_param(cls): return VectorParamParsable( item_class=TlsExtensionVariantClient, fallback_class=TlsExtensionUnparsed, min_byte_num=0, max_byte_num=2 ** 16 - 1 ) class TlsExtensionsServer(TlsExtensionsBase): @classmethod def get_param(cls): return VectorParamParsable( item_class=TlsExtensionVariantServer, fallback_class=TlsExtensionUnparsed, min_byte_num=0, max_byte_num=2 ** 16 - 1 ) @attr.s class TlsExtensionBase(ParsableBase): extension_type = attr.ib(init=False, validator=attr.validators.instance_of(TlsExtensionType)) @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod @abc.abstractmethod def _parse_type(cls, parser, name): raise NotImplementedError() @abc.abstractmethod def _compose_type(self, composer): raise NotImplementedError() @classmethod def _check_header(cls, parsable): parser = ParserBinary(parsable) cls._parse_type(parser, 'extension_type') parser.parse_numeric('extension_length', 2) if parser.unparsed_length < parser['extension_length']: raise NotEnoughData(parser['extension_length'] + parser.parsed_length) return parser def _compose_header(self, payload_length): header_composer = ComposerBinary() self._compose_type(header_composer) header_composer.compose_numeric(payload_length, 2) return header_composer.composed_bytes @attr.s class TlsExtensionUnparsed(TlsExtensionBase): extension_type = attr.ib(validator=attr.validators.instance_of((TlsExtensionType, TlsInvalidTypeTwoByte))) extension_data = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def _parse_type(cls, parser, name): parser.parse_parsable(name, TlsInvalidTypeTwoByte) @classmethod def _parse(cls, parsable): parser = cls._check_header(parsable) parser.parse_raw('extension_data', parser['extension_length']) return TlsExtensionUnparsed(parser['extension_type'], parser['extension_data']), parser.parsed_length def _compose_type(self, composer): composer.compose_parsable(self.extension_type) def compose(self): return self._compose_header(len(self.extension_data)) + self.extension_data @attr.s class TlsExtensionParsed(TlsExtensionBase): def __attrs_post_init__(self): self.extension_type = self.get_extension_type() attr.validate(self) @classmethod def _parse_type(cls, parser, name): parser.parse_parsable(name, TlsExtensionTypeFactory) def _compose_type(self, composer): composer.compose_numeric_enum_coded(self.extension_type) @classmethod @abc.abstractmethod def get_extension_type(cls): raise NotImplementedError() @classmethod def _parse_header(cls, parsable): parser = super()._check_header(parsable) if parser['extension_type'] != cls.get_extension_type(): raise InvalidType() return parser class TlsExtensionUnusedData(TlsExtensionParsed): @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('extension_data', parser['extension_length']) if parser['extension_data']: raise InvalidValue(parser['extension_data'], cls) return cls(), parser.parsed_length def compose(self): return self._compose_header(0) @classmethod @abc.abstractmethod def get_extension_type(cls): raise NotImplementedError() class TlsServerNameType(enum.IntEnum): HOST_NAME = 0x00 class TlsServerName(Vector): @classmethod def get_param(cls): return VectorParamNumeric( item_size=1, min_byte_num=1, max_byte_num=2 ** 16 - 1, ) @attr.s class TlsExtensionServerNameClient(TlsExtensionParsed): host_name = attr.ib(validator=attr.validators.instance_of(str)) name_type = attr.ib(validator=attr.validators.in_(TlsServerNameType), default=TlsServerNameType.HOST_NAME) @classmethod def get_extension_type(cls): return TlsExtensionType.SERVER_NAME @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_numeric('server_name_list_length', 2) parser.parse_numeric('server_name_type', 1, TlsServerNameType) parser.parse_parsable('server_name', TlsServerName) return cls(bytearray(parser['server_name']).decode('idna')), parser.parsed_length def compose(self): composer = ComposerBinary() idna_encoded_host_name = self.host_name.encode('idna') composer.compose_numeric(3 + len(idna_encoded_host_name), 2) composer.compose_numeric(self.name_type, 1) composer.compose_bytes(idna_encoded_host_name, 2) header_bytes = self._compose_header(composer.composed_length) return header_bytes + composer.composed_bytes @attr.s class TlsExtensionServerNameServer(TlsExtensionUnusedData): @classmethod def get_extension_type(cls): return TlsExtensionType.SERVER_NAME class TlsECPointFormatVector(VectorEnumCodeNumeric): @classmethod def get_param(cls): return VectorParamEnumCodeNumeric( item_class=TlsECPointFormatFactory, fallback_class=TlsInvalidTypeOneByte, min_byte_num=1, max_byte_num=2 ** 8 - 1, ) @attr.s class TlsExtensionECPointFormats(TlsExtensionParsed): point_formats = attr.ib( converter=TlsECPointFormatVector, validator=attr.validators.instance_of(TlsECPointFormatVector) ) @classmethod def get_extension_type(cls): return TlsExtensionType.EC_POINT_FORMATS @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('point_formats', TlsECPointFormatVector) return TlsExtensionECPointFormats(parser['point_formats']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.point_formats) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsEllipticCurveVector(VectorEnumCodeNumeric): @classmethod def get_param(cls): return VectorParamEnumCodeNumeric( item_class=TlsNamedCurveFactory, fallback_class=TlsInvalidTypeTwoByte, min_byte_num=1, max_byte_num=2 ** 16 - 1 ) @attr.s class TlsExtensionEllipticCurves(TlsExtensionParsed): elliptic_curves = attr.ib( converter=TlsEllipticCurveVector, validator=attr.validators.instance_of(TlsEllipticCurveVector) ) @classmethod def get_extension_type(cls): return TlsExtensionType.SUPPORTED_GROUPS @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('elliptic_curves', TlsEllipticCurveVector) return TlsExtensionEllipticCurves(parser['elliptic_curves']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.elliptic_curves) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsSupportedVersionVector(VectorParsableDerived): @classmethod def get_param(cls): return VectorParamParsable( item_class=TlsProtocolVersion, fallback_class=TlsInvalidTypeTwoByte, min_byte_num=2, max_byte_num=2 ** 8 - 2 ) @attr.s class TlsExtensionSupportedVersionsBase(TlsExtensionParsed): @classmethod def get_extension_type(cls): return TlsExtensionType.SUPPORTED_VERSIONS @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s class TlsExtensionSupportedVersionsClient(TlsExtensionSupportedVersionsBase): supported_versions = attr.ib( converter=TlsSupportedVersionVector, validator=attr.validators.instance_of(TlsSupportedVersionVector) ) @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('supported_versions', TlsSupportedVersionVector) return TlsExtensionSupportedVersionsClient(parser['supported_versions']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.supported_versions) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes @attr.s class TlsExtensionSupportedVersionsServer(TlsExtensionSupportedVersionsBase): selected_version = attr.ib(validator=attr.validators.instance_of(TlsProtocolVersion)) @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('selected_version', TlsProtocolVersion) return TlsExtensionSupportedVersionsServer(parser['selected_version']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.selected_version) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsSignatureAndHashAlgorithmVector(VectorEnumCodeNumeric): @classmethod def get_param(cls): return VectorParamEnumCodeNumeric( item_class=TlsSignatureAndHashAlgorithmFactory, fallback_class=TlsInvalidTypeTwoByte, min_byte_num=2, max_byte_num=2 ** 16 - 2 ) @attr.s class TlsExtensionSignatureAlgorithmsBase(TlsExtensionParsed): hash_and_signature_algorithms = attr.ib( converter=TlsSignatureAndHashAlgorithmVector, validator=attr.validators.instance_of(TlsSignatureAndHashAlgorithmVector) ) @classmethod @abc.abstractmethod def get_extension_type(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('hash_and_signature_algorithms', TlsSignatureAndHashAlgorithmVector) return cls(parser['hash_and_signature_algorithms']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.hash_and_signature_algorithms) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsExtensionSignatureAlgorithms(TlsExtensionSignatureAlgorithmsBase): @classmethod def get_extension_type(cls): return TlsExtensionType.SIGNATURE_ALGORITHMS class TlsExtensionSignatureAlgorithmsCert(TlsExtensionSignatureAlgorithmsBase): @classmethod def get_extension_type(cls): return TlsExtensionType.SIGNATURE_ALGORITHMS_CERT class TlsExtensionDelegatedCredentials(TlsExtensionSignatureAlgorithmsBase): @classmethod def get_extension_type(cls): return TlsExtensionType.DELEGATED_CREDENTIALS class TlsKeyExchangeVector(Vector): @classmethod def get_param(cls): return VectorParamNumeric(item_size=1, min_byte_num=1, max_byte_num=2 ** 16 - 1) @attr.s class TlsKeyShareEntry(ParsableBase): group = attr.ib(validator=attr.validators.instance_of(TlsNamedCurve)) key_exchange = attr.ib(converter=TlsKeyExchangeVector) @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_parsable('group', TlsNamedCurveFactory) parser.parse_parsable('key_exchange', TlsKeyExchangeVector) return TlsKeyShareEntry(parser['group'], parser['key_exchange']), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric_enum_coded(self.group) composer.compose_parsable(self.key_exchange) return composer.composed_bytes @attr.s class TlsKeyShareEntryInvalidType(ParsableBase): group = attr.ib(validator=attr.validators.instance_of(TlsInvalidTypeTwoByte)) data = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_parsable('group', TlsInvalidTypeTwoByte) parser.parse_bytes('data', 2) return TlsKeyShareEntryInvalidType(parser['group'], parser['data']), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_parsable(self.group) composer.compose_bytes(self.data, 2) return composer.composed_bytes class TlsKeyShareEntryVector(VectorParsable): @classmethod def get_param(cls): return VectorParamParsable( item_class=TlsKeyShareEntry, fallback_class=TlsKeyShareEntryInvalidType, min_byte_num=0, max_byte_num=2 ** 16 - 1 ) class TlsExtensionKeyShareBase(TlsExtensionParsed): @classmethod @abc.abstractmethod def get_extension_type(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s class TlsExtensionKeyShareServer(TlsExtensionKeyShareBase): key_share_entry = attr.ib(validator=attr.validators.instance_of(TlsKeyShareEntry)) @classmethod def get_extension_type(cls): return TlsExtensionType.KEY_SHARE @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('key_share_entry', TlsKeyShareEntry) return TlsExtensionKeyShareServer(parser['key_share_entry']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.key_share_entry) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes @attr.s class TlsExtensionKeyShareClientHelloRetry(TlsExtensionKeyShareBase): selected_group = attr.ib(converter=TlsNamedCurve) @classmethod def get_extension_type(cls): return TlsExtensionType.KEY_SHARE @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) if parser['extension_length'] != 2: raise InvalidType() parser.parse_parsable('selected_group', TlsNamedCurveFactory) return cls(parser['selected_group']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_numeric_enum_coded(self.selected_group) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes @attr.s class TlsExtensionKeyShareEntriesClientBase(TlsExtensionKeyShareBase): key_share_entries = attr.ib(converter=TlsKeyShareEntryVector) @classmethod @abc.abstractmethod def get_extension_type(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('key_share_entries', TlsKeyShareEntryVector) return cls(parser['key_share_entries']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.key_share_entries) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsExtensionKeyShareClient(TlsExtensionKeyShareEntriesClientBase): @classmethod def get_extension_type(cls): return TlsExtensionType.KEY_SHARE class TlsExtensionKeyShareReservedClient(TlsExtensionKeyShareEntriesClientBase): @classmethod def get_extension_type(cls): return TlsExtensionType.KEY_SHARE_RESERVED class TlsCertificateStatusType(enum.IntEnum): OCSP = 1 class TlsCertificateStatusRequestExtensions(Opaque): @classmethod def get_param(cls): return OpaqueParam( min_byte_num=0, max_byte_num=2 ** 16 - 1, ) class TlsCertificateStatusRequestResponderId(Opaque): @classmethod def get_param(cls): return OpaqueParam( min_byte_num=1, max_byte_num=2 ** 16 - 1, ) class TlsCertificateStatusRequestResponderIdList(VectorParsable): @classmethod def get_param(cls): return VectorParamParsable( item_class=TlsCertificateStatusRequestResponderId, fallback_class=None, min_byte_num=0, max_byte_num=2 ** 16 - 1 ) class TlsExtensionCertificateStatusRequestClient(TlsExtensionParsed): def __init__(self, responder_id_list=(), extensions=()): super().__init__() self.responder_id_list = TlsCertificateStatusRequestResponderIdList(responder_id_list) self.request_extensions = TlsCertificateStatusRequestExtensions(extensions) @classmethod def get_extension_type(cls): return TlsExtensionType.STATUS_REQUEST @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_numeric('status_type', 1, TlsCertificateStatusType) parser.parse_parsable('responder_id_list', TlsCertificateStatusRequestResponderIdList) parser.parse_parsable('extensions', TlsCertificateStatusRequestExtensions) return cls( parser['responder_id_list'], parser['extensions'], ), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_numeric(TlsCertificateStatusType.OCSP, 1) payload_composer.compose_parsable(self.responder_id_list) payload_composer.compose_parsable(self.request_extensions) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsExtensionCertificateStatusRequestServer(TlsExtensionUnusedData): @classmethod def get_extension_type(cls): return TlsExtensionType.STATUS_REQUEST class TlsRenegotiatedConnection(Opaque): @classmethod def get_param(cls): return OpaqueParam( min_byte_num=0, max_byte_num=2 ** 8 - 1, ) @attr.s class TlsExtensionRenegotiationInfo(TlsExtensionParsed): renegotiated_connection = attr.ib( default=TlsRenegotiatedConnection([]), validator=attr.validators.instance_of(TlsRenegotiatedConnection) ) @classmethod def get_extension_type(cls): return TlsExtensionType.RENEGOTIATION_INFO @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('renegotiated_connection', TlsRenegotiatedConnection) return TlsExtensionRenegotiationInfo(parser['renegotiated_connection']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.renegotiated_connection) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes @attr.s class TlsExtensionSessionTicket(TlsExtensionParsed): session_ticket = attr.ib( default=bytearray([]), validator=attr.validators.instance_of((bytes, bytearray)) ) @classmethod def get_extension_type(cls): return TlsExtensionType.SESSION_TICKET @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_raw('session_ticket', parser['extension_length']) return TlsExtensionSessionTicket(parser['session_ticket']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_raw(self.session_ticket) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsProtocolNameFactory(OpaqueEnumParsable): @classmethod def get_enum_class(cls): return TlsProtocolName @classmethod def get_param(cls): return OpaqueParam( min_byte_num=1, max_byte_num=2 ** 8 - 1 ) class TlsProtocolNameList(VectorEnumCodeString): @classmethod def get_param(cls): return VectorParamEnumCodeString( item_class=TlsProtocolNameFactory, min_byte_num=2, max_byte_num=2 ** 16 - 1 ) @attr.s class TlsExtensionApplicationLayerProtocolBase(TlsExtensionParsed): protocol_names = attr.ib( converter=TlsProtocolNameList, validator=attr.validators.instance_of(TlsProtocolNameList), ) @classmethod @abc.abstractmethod def get_extension_type(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('protocol_names', TlsProtocolNameList) return cls(parser['protocol_names']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.protocol_names) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsExtensionApplicationLayerProtocolNegotiation(TlsExtensionApplicationLayerProtocolBase): @classmethod def get_extension_type(cls): return TlsExtensionType.APPLICATION_LAYER_PROTOCOL_NEGOTIATION class TlsExtensionApplicationLayerProtocolSettings(TlsExtensionApplicationLayerProtocolBase): @classmethod def get_extension_type(cls): return TlsExtensionType.APPLICATION_LAYER_PROTOCOL_SETTINGS class TlsExtensionOldApplicationLayerProtocolSettings(TlsExtensionApplicationLayerProtocolBase): @classmethod def get_extension_type(cls): return TlsExtensionType.OLD_APPLICATION_LAYER_PROTOCOL_SETTINGS class TlsExtensionNextProtocolNegotiationClient(TlsExtensionUnusedData): @classmethod def get_extension_type(cls): return TlsExtensionType.NEXT_PROTOCOL_NEGOTIATION class TlsNextProtocolNameFactory(OpaqueEnumParsable): @classmethod def get_enum_class(cls): return TlsNextProtocolName @classmethod def get_param(cls): return OpaqueParam( min_byte_num=1, max_byte_num=2 ** 8 - 1 ) class TlsNextProtocolNameList(VectorEnumCodeString): @classmethod def get_param(cls): return VectorParamEnumCodeString( item_class=TlsNextProtocolNameFactory, min_byte_num=1, max_byte_num=2 ** 16 - 1 ) @attr.s class TlsExtensionNextProtocolNegotiationServer(TlsExtensionParsed): protocol_names = attr.ib( converter=TlsNextProtocolNameList, validator=attr.validators.instance_of(TlsNextProtocolNameList), ) @classmethod def get_extension_type(cls): return TlsExtensionType.NEXT_PROTOCOL_NEGOTIATION @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) cls._parse_type(parser, 'extension_type') if parser['extension_type'] != cls.get_extension_type(): raise InvalidType() parser.parse_parsable('protocol_names', TlsNextProtocolNameList) return cls(parser['protocol_names']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.protocol_names) header_composer = ComposerBinary() self._compose_type(header_composer) return header_composer.composed_bytes + payload_composer.composed_bytes class TlsExtensionChannelId(TlsExtensionUnusedData): @classmethod def get_extension_type(cls): return TlsExtensionType.CHANNEL_ID class TlsExtensionEncryptThenMAC(TlsExtensionUnusedData): @classmethod def get_extension_type(cls): return TlsExtensionType.ENCRYPT_THEN_MAC class TlsExtensionExtendedMasterSecret(TlsExtensionUnusedData): @classmethod def get_extension_type(cls): return TlsExtensionType.EXTENDED_MASTER_SECRET @attr.s class TlsExtensionShortRecordHeader(TlsExtensionUnusedData): @classmethod def get_extension_type(cls): return TlsExtensionType.SHORT_RECORD_HEADER class TlsTokenBindingProtocolVersion(ProtocolVersionMajorMinorBase): pass class TlsTokenBindingParamaterFactory(OneByteEnumParsable): @classmethod def get_enum_class(cls): return TlsTokenBindingParamater @abc.abstractmethod def compose(self): raise NotImplementedError() class TlsTokenBindingParamaterVector(VectorEnumCodeNumeric): @classmethod def get_param(cls): return VectorParamEnumCodeNumeric( item_class=TlsTokenBindingParamaterFactory, fallback_class=TlsInvalidTypeOneByte, min_byte_num=1, max_byte_num=2 ** 8 - 1, ) @attr.s class TlsExtensionTokenBinding(TlsExtensionParsed): protocol_version = attr.ib(validator=attr.validators.instance_of(TlsTokenBindingProtocolVersion)) parameters = attr.ib( converter=TlsTokenBindingParamaterVector, validator=attr.validators.instance_of(TlsTokenBindingParamaterVector) ) @classmethod def get_extension_type(cls): return TlsExtensionType.TOKEN_BINDING @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('protocol_version', TlsTokenBindingProtocolVersion) parser.parse_parsable('parameters', TlsTokenBindingParamaterVector) return cls(parser['protocol_version'], parser['parameters']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.protocol_version) payload_composer.compose_parsable(self.parameters) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsPskKeyExchangeModeFactory(OneByteEnumParsable): @classmethod def get_enum_class(cls): return TlsPskKeyExchangeMode @abc.abstractmethod def compose(self): raise NotImplementedError() class TlsPskKeyExchangeModeVector(VectorEnumCodeNumeric): @classmethod def get_param(cls): return VectorParamEnumCodeNumeric( item_class=TlsPskKeyExchangeModeFactory, fallback_class=TlsInvalidTypeOneByte, min_byte_num=1, max_byte_num=2 ** 8 - 1, ) @attr.s class TlsExtensionPskKeyExchangeModes(TlsExtensionParsed): key_exchange_modes = attr.ib( converter=TlsPskKeyExchangeModeVector, validator=attr.validators.instance_of(TlsPskKeyExchangeModeVector) ) @classmethod def get_extension_type(cls): return TlsExtensionType.PSK_KEY_EXCHANGE_MODES @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('key_exchange_modes', TlsPskKeyExchangeModeVector) return TlsExtensionPskKeyExchangeModes(parser['key_exchange_modes']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.key_exchange_modes) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes @attr.s class TlsExtensionRecordSizeLimit(TlsExtensionParsed): record_size_limit = attr.ib(validator=attr.validators.instance_of(int)) @classmethod def get_extension_type(cls): return TlsExtensionType.RECORD_SIZE_LIMIT @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric('record_size_limit', 2) return TlsExtensionRecordSizeLimit(parser['record_size_limit']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_numeric(self.record_size_limit, 2) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes @attr.s class TlsExtensionServerPadding(TlsExtensionParsed): padding_size = attr.ib(validator=attr.validators.instance_of(int)) @classmethod def get_extension_type(cls): return TlsExtensionType.SERVER_PADDING @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric('padding_size', 2) return TlsExtensionServerPadding(parser['padding_size']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_numeric(self.padding_size, 2) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes @attr.s class TlsExtensionSignedCertificateTimestampClient(TlsExtensionUnusedData): @classmethod def get_extension_type(cls): return TlsExtensionType.SIGNED_CERTIFICATE_TIMESTAMP @attr.s class TlsExtensionSignedCertificateTimestampServer(TlsExtensionParsed): scts = attr.ib( converter=SignedCertificateTimestampList, validator=attr.validators.optional(attr.validators.instance_of(SignedCertificateTimestampList)) ) @classmethod def get_extension_type(cls): return TlsExtensionType.SIGNED_CERTIFICATE_TIMESTAMP @classmethod def _parse(cls, parsable): parser = super(cls, cls)._parse_header(parsable) parser.parse_parsable('scts', SignedCertificateTimestampList) return cls(parser['scts']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.scts) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsCertificateCompressionAlgorithmFactory(TwoByteEnumParsable): @classmethod def get_enum_class(cls): return TlsCertificateCompressionAlgorithm @abc.abstractmethod def compose(self): raise NotImplementedError() class TlsCertificateCompressionAlgorithmVector(VectorEnumCodeNumeric): @classmethod def get_param(cls): return VectorParamEnumCodeNumeric( item_class=TlsCertificateCompressionAlgorithmFactory, fallback_class=TlsInvalidTypeTwoByte, min_byte_num=2, max_byte_num=2 ** 8 - 2 ) @attr.s class TlsExtensionCompressCertificate(TlsExtensionParsed): compression_algorithms = attr.ib( converter=TlsCertificateCompressionAlgorithmVector, validator=attr.validators.instance_of(TlsCertificateCompressionAlgorithmVector), ) @classmethod def get_extension_type(cls): return TlsExtensionType.COMPRESS_CERTIFICATE @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_parsable('compression_algorithms', TlsCertificateCompressionAlgorithmVector) return cls(parser['compression_algorithms']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.compression_algorithms) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes @attr.s class TlsExtensionPadding(TlsExtensionParsed): length = attr.ib(validator=attr.validators.instance_of(int)) @classmethod def get_extension_type(cls): return TlsExtensionType.PADDING @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) parser.parse_raw('padding', parser['extension_length']) try: non_zero_int = next(byte for byte in parser['padding'] if byte != 0) raise InvalidValue(bytes((non_zero_int,)), cls) except StopIteration: pass return cls(parser['extension_length']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_raw(self.length * b'\x00') header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsEncryptedClientHelloType(enum.IntEnum): OUTER = 0 INNER = 1 @attr.s class TlsExtensionEncryptedClientHelloBase(TlsExtensionParsed): @classmethod def get_extension_type(cls): return TlsExtensionType.ENCRYPTED_CLIENT_HELLO @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod @abc.abstractmethod def get_encrypted_client_hello_type(cls): raise NotImplementedError() @classmethod def _parse_client_hello_type(cls, parser): parser.parse_numeric('client_hello_type', 1, TlsEncryptedClientHelloType) if parser['client_hello_type'] != cls.get_encrypted_client_hello_type(): raise InvalidType() def compose_type(self): composer = ComposerBinary() composer.compose_numeric(self.get_encrypted_client_hello_type(), 1) return composer @attr.s class TlsExtensionEncryptedClientHelloInner(TlsExtensionEncryptedClientHelloBase): @classmethod def get_encrypted_client_hello_type(cls): return TlsEncryptedClientHelloType.INNER @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('payload', parser['extension_length']) body_parser = ParserBinary(parser['payload']) cls._parse_client_hello_type(body_parser) return cls(), parser.parsed_length def compose(self): payload_composer = self.compose_type() header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes @attr.s class TlsExtensionEncryptedClientHelloOuter(TlsExtensionEncryptedClientHelloBase): data = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def get_encrypted_client_hello_type(cls): return TlsEncryptedClientHelloType.OUTER @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_raw('payload', parser['extension_length']) body_parser = ParserBinary(parser['payload']) cls._parse_client_hello_type(body_parser) body_parser.parse_raw('data', body_parser.unparsed_length) return cls(body_parser['data']), parser.parsed_length def compose(self): payload_composer = self.compose_type() payload_composer.compose_raw(self.data) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsExtensionPostHandshakeAuthentication(TlsExtensionUnusedData): @classmethod def get_extension_type(cls): return TlsExtensionType.POST_HANDSHAKE_AUTH class TlsTrustAnchorIdentifier(Opaque): @classmethod def get_param(cls): return OpaqueParam( min_byte_num=1, max_byte_num=2 ** 8 - 1 ) class TlsTrustAnchorIdentifierList(VectorParsable): @classmethod def get_param(cls): return VectorParamParsable( item_class=TlsTrustAnchorIdentifier, fallback_class=None, min_byte_num=0, max_byte_num=2 ** 16 - 1 ) @attr.s class TlsExtensionTrustAnchors(TlsExtensionParsed): trust_anchor_identifiers = attr.ib( converter=TlsTrustAnchorIdentifierList, validator=attr.validators.instance_of(TlsTrustAnchorIdentifierList), ) @classmethod def get_extension_type(cls): return TlsExtensionType.TRUST_ANCHORS @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_parsable('trust_anchor_identifiers', TlsTrustAnchorIdentifierList) return TlsExtensionTrustAnchors(parser['trust_anchor_identifiers']), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.trust_anchor_identifiers) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsExtensionVariantBase(VariantParsable): @classmethod @abc.abstractmethod def get_parsed_extensions(cls): raise NotImplementedError() @classmethod def _get_variants(cls): variants = cls.get_parsed_extensions() variants.update([ (extension_type, (TlsExtensionUnparsed, )) for extension_type in TlsExtensionType if extension_type not in variants ]) return variants class TlsExtensionVariantClient(TlsExtensionVariantBase): @classmethod def get_parsed_extensions(cls): return collections.OrderedDict([ (TlsExtensionType.APPLICATION_LAYER_PROTOCOL_NEGOTIATION, [TlsExtensionApplicationLayerProtocolNegotiation, ]), (TlsExtensionType.APPLICATION_LAYER_PROTOCOL_SETTINGS, [TlsExtensionApplicationLayerProtocolSettings, ]), (TlsExtensionType.OLD_APPLICATION_LAYER_PROTOCOL_SETTINGS, [TlsExtensionOldApplicationLayerProtocolSettings, ]), (TlsExtensionType.CHANNEL_ID, [TlsExtensionChannelId, ]), (TlsExtensionType.COMPRESS_CERTIFICATE, [TlsExtensionCompressCertificate, ]), (TlsExtensionType.ENCRYPT_THEN_MAC, [TlsExtensionEncryptThenMAC, ]), (TlsExtensionType.EXTENDED_MASTER_SECRET, [TlsExtensionExtendedMasterSecret, ]), (TlsExtensionType.RENEGOTIATION_INFO, [TlsExtensionRenegotiationInfo, ]), (TlsExtensionType.NEXT_PROTOCOL_NEGOTIATION, [TlsExtensionNextProtocolNegotiationClient, ]), (TlsExtensionType.PADDING, [TlsExtensionPadding, ]), (TlsExtensionType.SERVER_NAME, [TlsExtensionServerNameClient, ]), (TlsExtensionType.SESSION_TICKET, [TlsExtensionSessionTicket, ]), (TlsExtensionType.STATUS_REQUEST, [TlsExtensionCertificateStatusRequestClient, ]), (TlsExtensionType.SUPPORTED_GROUPS, [TlsExtensionEllipticCurves, ]), (TlsExtensionType.DELEGATED_CREDENTIALS, [TlsExtensionDelegatedCredentials, ]), (TlsExtensionType.EC_POINT_FORMATS, [TlsExtensionECPointFormats, ]), (TlsExtensionType.KEY_SHARE, [TlsExtensionKeyShareClient, ]), (TlsExtensionType.KEY_SHARE_RESERVED, [TlsExtensionKeyShareReservedClient, ]), (TlsExtensionType.POST_HANDSHAKE_AUTH, [TlsExtensionPostHandshakeAuthentication, ]), (TlsExtensionType.PSK_KEY_EXCHANGE_MODES, [TlsExtensionPskKeyExchangeModes, ]), (TlsExtensionType.RECORD_SIZE_LIMIT, [TlsExtensionRecordSizeLimit, ]), (TlsExtensionType.SERVER_PADDING, [TlsExtensionServerPadding, ]), (TlsExtensionType.SHORT_RECORD_HEADER, [TlsExtensionShortRecordHeader, ]), (TlsExtensionType.SIGNATURE_ALGORITHMS, [TlsExtensionSignatureAlgorithms, ]), (TlsExtensionType.SIGNATURE_ALGORITHMS_CERT, [TlsExtensionSignatureAlgorithmsCert, ]), (TlsExtensionType.SIGNED_CERTIFICATE_TIMESTAMP, [TlsExtensionSignedCertificateTimestampClient, ]), (TlsExtensionType.SUPPORTED_VERSIONS, [TlsExtensionSupportedVersionsClient, ]), (TlsExtensionType.ENCRYPTED_CLIENT_HELLO, [TlsExtensionEncryptedClientHelloInner, TlsExtensionEncryptedClientHelloOuter]), (TlsExtensionType.TOKEN_BINDING, [TlsExtensionTokenBinding, ]), (TlsExtensionType.TRUST_ANCHORS, [TlsExtensionTrustAnchors, ]), ]) class TlsExtensionVariantServer(TlsExtensionVariantBase): @classmethod def get_parsed_extensions(cls): return collections.OrderedDict([ (TlsExtensionType.APPLICATION_LAYER_PROTOCOL_NEGOTIATION, [TlsExtensionApplicationLayerProtocolNegotiation, ]), (TlsExtensionType.CHANNEL_ID, [TlsExtensionChannelId, ]), (TlsExtensionType.EC_POINT_FORMATS, [TlsExtensionECPointFormats, ]), (TlsExtensionType.ENCRYPT_THEN_MAC, [TlsExtensionEncryptThenMAC, ]), (TlsExtensionType.EXTENDED_MASTER_SECRET, [TlsExtensionExtendedMasterSecret, ]), (TlsExtensionType.KEY_SHARE, [TlsExtensionKeyShareClientHelloRetry, TlsExtensionKeyShareServer]), (TlsExtensionType.NEXT_PROTOCOL_NEGOTIATION, [TlsExtensionNextProtocolNegotiationServer, ]), (TlsExtensionType.RECORD_SIZE_LIMIT, [TlsExtensionRecordSizeLimit, ]), (TlsExtensionType.RENEGOTIATION_INFO, [TlsExtensionRenegotiationInfo, ]), (TlsExtensionType.SERVER_NAME, [TlsExtensionServerNameServer, ]), (TlsExtensionType.SESSION_TICKET, [TlsExtensionSessionTicket, ]), (TlsExtensionType.SIGNED_CERTIFICATE_TIMESTAMP, [TlsExtensionSignedCertificateTimestampServer, ]), (TlsExtensionType.STATUS_REQUEST, [TlsExtensionCertificateStatusRequestServer, ]), (TlsExtensionType.SUPPORTED_VERSIONS, [TlsExtensionSupportedVersionsServer, ]), ]) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/grease.py000066400000000000000000000063231524413560000272350ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import enum import random import attr from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.types import CryptoDataEnumCodedBase from cryptodatahub.tls.algorithm import TlsGreaseOneByte, TlsGreaseTwoByte from cryptoparser.common.base import ParsableBase from cryptoparser.common.parse import ParserBinary, ComposerBinary class TlsInvalidType(enum.IntEnum): GREASE = 0 UNKNOWN = 1 @attr.s(frozen=True) class TlsInvalidTypeParamsBase: code = attr.ib(validator=attr.validators.instance_of(int)) value_type = attr.ib(validator=attr.validators.in_(TlsInvalidType)) @classmethod def get_code_size(cls): raise NotImplementedError() class TlsInvalidTypeParamsOneByte(TlsInvalidTypeParamsBase): @classmethod def get_code_size(cls): return 1 class TlsInvalidTypeParamsTwoByte(TlsInvalidTypeParamsBase): @classmethod def get_code_size(cls): return 2 @attr.s(frozen=True) class TlsInvalidTypeBase(ParsableBase): code = attr.ib(validator=attr.validators.instance_of((int, CryptoDataEnumCodedBase))) value = attr.ib(init=False, validator=attr.validators.instance_of(TlsInvalidTypeParamsBase)) def __attrs_post_init__(self): if isinstance(self.code, self.get_grease_enum()): value_type = TlsInvalidType.GREASE object.__setattr__(self, 'code', self.code.value.code) elif isinstance(self.code, CryptoDataEnumCodedBase): value_type = TlsInvalidType.UNKNOWN object.__setattr__(self, 'code', self.code.value.code) else: try: object.__setattr__(self, 'code', self.get_grease_enum().from_code(self.code).value.code) value_type = TlsInvalidType.GREASE except InvalidValue: value_type = TlsInvalidType.UNKNOWN object.__setattr__(self, 'value', self.get_param_class()(self.code, value_type)) @classmethod @abc.abstractmethod def get_param_class(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def get_grease_enum(cls): raise NotImplementedError() @classmethod def get_byte_num(cls): return cls.get_param_class().get_code_size() @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('code', cls.get_byte_num()) code = parser['code'] return cls(code), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.code, self.get_byte_num()) return composer.composed_bytes @classmethod def from_random(cls): grease_enum_type = cls.get_grease_enum() return cls(random.choice(list(grease_enum_type))) class TlsInvalidTypeOneByte(TlsInvalidTypeBase): @classmethod def get_param_class(cls): return TlsInvalidTypeParamsOneByte @classmethod def get_grease_enum(cls): return TlsGreaseOneByte class TlsInvalidTypeTwoByte(TlsInvalidTypeBase): @classmethod def get_param_class(cls): return TlsInvalidTypeParamsTwoByte @classmethod def get_grease_enum(cls): return TlsGreaseTwoByte cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/ldap.py000066400000000000000000000131461524413560000267100ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import enum import re import attr import asn1crypto.core from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import NotEnoughData from cryptoparser.common.parse import ParsableBase class LDAPClass(enum.IntEnum): UNIVERSAL = 0 APPLICATION = 1 CONTEXT = 2 class LDAPResultCode(enum.IntEnum): SUCCESS = 0 OPERATIONS_ERROR = 1 PROTOCOL_ERROR = 2 TIME_LIMIT_EXCEEDED = 3 SIZE_LIMIT_EXCEEDED = 4 COMPARE_FALSE = 5 COMPARE_TRUE = 6 AUTH_METHOD_NOT_SUPPORTED = 7 STRONGER_AUTH_REQUIRED = 8 REFERRAL = 10 ADMIN_LIMIT_EXCEEDED = 11 UNAVAILABLE_CRITICAL_EXTENSION = 12 CONFIDENTIALITY_REQUIRED = 13 SASL_BIND_IN_PROGRESS = 14 NO_SUCH_ATTRIBUTE = 16 UNDEFINED_ATTRIBUTE_TYPE = 17 INAPPROPRIATE_MATCHING = 18 CONSTRAINT_VIOLATION = 19 ATTRIBUTE_OR_VALUE_EXISTS = 20 INVALID_ATTRIBUTE_SYNTAX = 21 NO_SUCH_OBJECT = 32 ALIAS_PROBLEM = 33 INVALID_DN_SYNTAX = 34 ALIAS_DEREFERENCING_PROBLEM = 36 INAPPROPRIATE_AUTHENTICATION = 48 INVALID_CREDENTIALS = 49 INSUFFICIENT_ACCESS_RIGHTS = 50 BUSY = 51 UNAVAILABLE = 52 UNWILLING_TO_PERFORM = 53 LOOP_DETECT = 54 NAMING_VIOLATION = 64 OBJECT_CLASS_VIOLATION = 65 NOT_ALLOWED_ON_NON_LEAF = 66 NOT_ALLOWED_ON_RDN = 67 ENTRY_ALREADY_EXISTS = 68 OBJECT_CLASS_MODS_PROHIBITED = 69 AFFECTS_MULTIPLE_DSAS = 71 OTHER = 80 class LDAPResultCodeEnum(asn1crypto.core.Enumerated): _map = {result_code.value: result_code for result_code in list(LDAPResultCode)} class LDAPOID(asn1crypto.core.OctetString): pass class LDAPControl(asn1crypto.core.Sequence): _fields = [ ('controlType', LDAPOID), ('criticality', asn1crypto.core.Boolean, {'default': False}), ('controlValue', asn1crypto.core.OctetString, {'optional': True}), ] class LDAPControls(asn1crypto.core.SequenceOf): _child_spec = LDAPControl class LDAPExtendedRequest(asn1crypto.core.Sequence): _fields = [ ('requestName', LDAPOID, {'implicit': (LDAPClass.CONTEXT.value, 0)}), ('requestValue', asn1crypto.core.OctetString, {'implicit': (LDAPClass.CONTEXT.value, 1), 'optional': True}), ] class LDAPDN(asn1crypto.core.OctetString): pass class LDAPString(asn1crypto.core.OctetString): pass class LDAPURI(LDAPString): pass class LDAPReferral(asn1crypto.core.SequenceOf): _child_spec = LDAPURI class LDAPExtendedResponse(asn1crypto.core.Sequence): _fields = [ ('resultCode', LDAPResultCodeEnum), ('matchedDN', LDAPDN), ('diagnosticMessage', LDAPString), ('referral', LDAPReferral, {'implicit': (LDAPClass.CONTEXT.value, 3), 'optional': True}), ('responseName', LDAPOID, {'implicit': (LDAPClass.CONTEXT.value, 10), 'optional': True}), ('responseValue', asn1crypto.core.OctetString, {'implicit': (LDAPClass.CONTEXT.value, 11), 'optional': True}), ] class LDAPProtocolOp(asn1crypto.core.Choice): _alternatives = [ ('extendedReq', LDAPExtendedRequest, {'implicit': (LDAPClass.APPLICATION.value, 23)}), ('extendedResp', LDAPExtendedResponse, {'implicit': (LDAPClass.APPLICATION.value, 24)}), ] class LDAPMessage(asn1crypto.core.Sequence): _fields = [ ('messageID', asn1crypto.core.Integer), ('protocolOp', LDAPProtocolOp), ('controls', LDAPControls, {'implicit': (LDAPClass.CONTEXT.value, 0), 'optional': True}), ] class LDAPMessageParsableBase(ParsableBase): HEADER_SIZE = 6 _NOT_ENOUGH_DATA_REGEX = re.compile( r'Insufficient data - ([0-9]+) bytes requested but only ([0-9]+) available' ) @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @classmethod def _parse_asn1(cls, parsable): try: message = LDAPMessage.load(bytes(parsable)) # ensure recursive parsing _ = message.native except ValueError as e: match = cls._NOT_ENOUGH_DATA_REGEX.match(e.args[0]) if match: bytes_requested = int(match.group(1)) bytes_available = int(match.group(2)) raise NotEnoughData(bytes_requested - bytes_available) from e raise InvalidValue(parsable, cls) from e return message class LDAPExtendedRequestStartTLS(LDAPMessageParsableBase): @classmethod def _parse(cls, parsable): asn1_message = cls._parse_asn1(parsable) return LDAPExtendedRequestStartTLS(), len(asn1_message.dump()) def compose(self): return LDAPMessage({ 'messageID': 1, 'protocolOp': { 'extendedReq': { 'requestName': b'1.3.6.1.4.1.1466.20037' } } }).dump() @attr.s class LDAPExtendedResponseStartTLS(LDAPMessageParsableBase): result_code = attr.ib(validator=attr.validators.in_(LDAPResultCode)) @classmethod def _parse(cls, parsable): asn1_message = cls._parse_asn1(parsable) return LDAPExtendedResponseStartTLS( asn1_message['protocolOp'].chosen['resultCode'].native ), len(asn1_message.dump()) def compose(self): return LDAPMessage({ 'messageID': 1, 'protocolOp': { 'extendedResp': { 'resultCode': self.result_code.value, 'matchedDN': b'', 'diagnosticMessage': b'' } } }).dump() cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/mysql.py000066400000000000000000000372151524413560000271400ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import enum import attr from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.base import Serializable, OneByteEnumComposer, OneByteEnumParsable from cryptoparser.common.parse import ByteOrder, ComposerBinary, ParsableBase, ParserBinary from cryptoparser.common.exception import NotEnoughData class MySQLVersion(enum.IntEnum): MYSQL_9 = 0x09 MYSQL_10 = 0x0a class MySQLCapability(enum.IntEnum): CLIENT_LONG_PASSWORD = 0x00000001 CLIENT_FOUND_ROWS = 0x00000002 CLIENT_LONG_FLAG = 0x00000004 CLIENT_CONNECT_WITH_DB = 0x00000008 CLIENT_NO_SCHEMA = 0x00000010 CLIENT_COMPRESS = 0x00000020 CLIENT_ODBC = 0x00000040 CLIENT_LOCAL_FILES = 0x00000080 CLIENT_IGNORE_SPACE = 0x00000100 CLIENT_PROTOCOL_41 = 0x00000200 CLIENT_INTERACTIVE = 0x00000400 CLIENT_SSL = 0x00000800 CLIENT_IGNORE_SIGPIPE = 0x00001000 CLIENT_TRANSACTIONS = 0x00002000 CLIENT_RESERVED = 0x00004000 CLIENT_SECURE_CONNECTION = 0x00008000 CLIENT_MULTI_STATEMENTS = 0x00010000 CLIENT_MULTI_RESULTS = 0x00020000 CLIENT_PS_MULTI_RESULTS = 0x00040000 CLIENT_PLUGIN_AUTH = 0x00080000 CLIENT_CONNECT_ATTRS = 0x00100000 CLIENT_PLUGIN_AUTH_LENENC_CLIENT_DATA = 0x00200000 CLIENT_CAN_HANDLE_EXPIRED_PASSWORDS = 0x00400000 CLIENT_SESSION_TRACK = 0x00800000 CLIENT_DEPRECATE_EOF = 0x01000000 class MySQLStatusFlag(enum.IntEnum): SERVER_STATUS_IN_TRANS = 0x0001 SERVER_STATUS_AUTOCOMMIT = 0x0002 SERVER_MORE_RESULTS_EXISTS = 0x0008 SERVER_STATUS_NO_GOOD_INDEX_USED = 0x0010 SERVER_STATUS_NO_INDEX_USED = 0x0020 SERVER_STATUS_CURSOR_EXISTS = 0x0040 SERVER_STATUS_LAST_ROW_SENT = 0x0080 SERVER_STATUS_DB_DROPPED = 0x0100 SERVER_STATUS_NO_BACKSLASH_ESCAPES = 0x0200 SERVER_STATUS_METADATA_CHANGED = 0x0400 SERVER_QUERY_WAS_SLOW = 0x0800 SERVER_PS_OUT_PARAMS = 0x1000 SERVER_STATUS_IN_TRANS_READONLY = 0x2000 SERVER_SESSION_STATE_CHANGED = 0x4000 @attr.s(frozen=True) class MySQLCharacterSetParams(Serializable): code = attr.ib(validator=attr.validators.instance_of(int)) name = attr.ib(validator=attr.validators.instance_of(str)) collate_name = attr.ib(validator=attr.validators.instance_of(str)) class MySQLCharacterSetFactory(OneByteEnumParsable): @classmethod def get_enum_class(cls): return MySQLCharacterSet @abc.abstractmethod def compose(self): raise NotImplementedError() class MySQLCharacterSet(OneByteEnumComposer, enum.Enum): BIG5 = MySQLCharacterSetParams( code=1, name='big5', collate_name='big5_chinese_ci', ) DEC8 = MySQLCharacterSetParams( code=3, name='dec8', collate_name='dec8_swedish_ci', ) CP850 = MySQLCharacterSetParams( code=4, name='cp850', collate_name='cp850_general_ci', ) HP8 = MySQLCharacterSetParams( code=6, name='hp8', collate_name='hp8_english_ci', ) KOI8R = MySQLCharacterSetParams( code=7, name='koi8r', collate_name='koi8r_general_ci', ) LATIN1 = MySQLCharacterSetParams( code=8, name='latin1', collate_name='latin1_swedish_ci', ) LATIN2 = MySQLCharacterSetParams( code=9, name='latin2', collate_name='latin2_general_ci', ) SWE7 = MySQLCharacterSetParams( code=10, name='swe7', collate_name='swe7_swedish_ci', ) ASCII = MySQLCharacterSetParams( code=11, name='ascii', collate_name='ascii_general_ci', ) UJIS = MySQLCharacterSetParams( code=12, name='ujis', collate_name='ujis_japanese_ci', ) SJIS = MySQLCharacterSetParams( code=13, name='sjis', collate_name='sjis_japanese_ci', ) HEBREW = MySQLCharacterSetParams( code=16, name='hebrew', collate_name='hebrew_general_ci', ) TIS620 = MySQLCharacterSetParams( code=18, name='tis620', collate_name='tis620_thai_ci', ) EUCKR = MySQLCharacterSetParams( code=19, name='euckr', collate_name='euckr_korean_ci', ) KOI8U = MySQLCharacterSetParams( code=22, name='koi8u', collate_name='koi8u_general_ci', ) GB2312 = MySQLCharacterSetParams( code=24, name='gb2312', collate_name='gb2312_chinese_ci', ) GREEK = MySQLCharacterSetParams( code=25, name='greek', collate_name='greek_general_ci', ) CP1250 = MySQLCharacterSetParams( code=26, name='cp1250', collate_name='cp1250_general_ci', ) GBK = MySQLCharacterSetParams( code=28, name='gbk', collate_name='gbk_chinese_ci', ) LATIN5 = MySQLCharacterSetParams( code=30, name='latin5', collate_name='latin5_turkish_ci', ) ARMSCII8 = MySQLCharacterSetParams( code=32, name='armscii8', collate_name='armscii8_general_ci', ) UTF8 = MySQLCharacterSetParams( code=33, name='utf8', collate_name='utf8_general_ci', ) UCS2 = MySQLCharacterSetParams( code=35, name='ucs2', collate_name='ucs2_general_ci', ) CP866 = MySQLCharacterSetParams( code=36, name='cp866', collate_name='cp866_general_ci', ) KEYBCS2 = MySQLCharacterSetParams( code=37, name='keybcs2', collate_name='keybcs2_general_ci', ) MACCE = MySQLCharacterSetParams( code=38, name='macce', collate_name='macce_general_ci', ) MACROMAN = MySQLCharacterSetParams( code=39, name='macroman', collate_name='macroman_general_ci', ) CP852 = MySQLCharacterSetParams( code=40, name='cp852', collate_name='cp852_general_ci', ) LATIN7 = MySQLCharacterSetParams( code=41, name='latin7', collate_name='latin7_general_ci', ) CP1251 = MySQLCharacterSetParams( code=51, name='cp1251', collate_name='cp1251_general_ci', ) UTF16 = MySQLCharacterSetParams( code=54, name='utf16', collate_name='utf16_general_ci', ) UTF16LE = MySQLCharacterSetParams( code=56, name='utf16le', collate_name='utf16le_general_ci', ) CP1256 = MySQLCharacterSetParams( code=57, name='cp1256', collate_name='cp1256_general_ci', ) CP1257 = MySQLCharacterSetParams( code=59, name='cp1257', collate_name='cp1257_general_ci', ) UTF32 = MySQLCharacterSetParams( code=60, name='utf32', collate_name='utf32_general_ci', ) BINARY = MySQLCharacterSetParams( code=63, name='binary', collate_name='binary', ) GEOSTD8 = MySQLCharacterSetParams( code=92, name='geostd8', collate_name='geostd8_general_ci', ) CP932 = MySQLCharacterSetParams( code=95, name='cp932', collate_name='cp932_japanese_ci', ) EUCJPMS = MySQLCharacterSetParams( code=97, name='eucjpms', collate_name='eucjpms_japanese_ci', ) GB18030 = MySQLCharacterSetParams( code=248, name='gb18030', collate_name='gb18030_chinese_ci', ) UTF8MB4 = MySQLCharacterSetParams( code=255, name='utf8mb4', collate_name='utf8mb4_0900_ai_ci', ) class MySQLPacketBase(ParsableBase): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s class MySQLRecord(ParsableBase): HEADER_SIZE = 4 packet_number = attr.ib(validator=attr.validators.instance_of(int)) packet_bytes = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def _parse(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable, byte_order=ByteOrder.LITTLE_ENDIAN) parser.parse_numeric('packet_length', 3) parser.parse_numeric('packet_number', 1) parser.parse_raw('packet_bytes', parser['packet_length']) return MySQLRecord( packet_number=parser['packet_number'], packet_bytes=parser['packet_bytes'], ), parser.parsed_length def compose(self): composer = ComposerBinary(byte_order=ByteOrder.LITTLE_ENDIAN) composer.compose_numeric(len(self.packet_bytes), 3) composer.compose_numeric(self.packet_number, 1) composer.compose_raw(self.packet_bytes) return composer.composed_bytes @attr.s class MySQLHandshakeV10(MySQLPacketBase): # pylint: disable=too-many-instance-attributes protocol_version = attr.ib(validator=attr.validators.in_(MySQLVersion)) server_version = attr.ib(validator=attr.validators.instance_of(str)) connection_id = attr.ib(validator=attr.validators.instance_of(int)) auth_plugin_data = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) capabilities = attr.ib(validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(MySQLCapability), )) character_set = attr.ib( default=MySQLCharacterSet.UTF8, validator=attr.validators.optional(attr.validators.in_(MySQLCharacterSet)) ) states = attr.ib(default={}, validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(MySQLStatusFlag), )) auth_plugin_data_2 = attr.ib( default=None, validator=attr.validators.optional(attr.validators.instance_of((bytes, bytearray))), ) auth_plugin_name = attr.ib( default=None, validator=attr.validators.optional(attr.validators.instance_of(str)) ) MINIMUM_SIZE = 33 @classmethod def _parse(cls, parsable): if len(parsable) < cls.MINIMUM_SIZE: raise NotEnoughData(cls.MINIMUM_SIZE - len(parsable)) parser = ParserBinary(parsable, byte_order=ByteOrder.LITTLE_ENDIAN) parser.parse_numeric('protocol_version', 1, MySQLVersion) parser.parse_string_null_terminated('server_version', 'ascii') parser.parse_numeric('connection_id', 4) parser.parse_raw('auth_plugin_data', 8) parser.parse_raw('filler', 1) del parser['filler'] parser.parse_numeric_flags('capabilities', 2, MySQLCapability) parser.parse_parsable('character_set', MySQLCharacterSetFactory) parser.parse_numeric_flags('states', 2, MySQLStatusFlag) parser.parse_numeric_flags('capabilities_2', 2, MySQLCapability, shift_left=16) capabilities = set(parser['capabilities']) | set(parser['capabilities_2']) del parser['capabilities_2'] parser.parse_numeric('auth_plugin_data_len', 1) auth_plugin_data_len = parser['auth_plugin_data_len'] del parser['auth_plugin_data_len'] parser.parse_raw('reserved', 10) del parser['reserved'] if MySQLCapability.CLIENT_PLUGIN_AUTH in capabilities: if not auth_plugin_data_len: raise InvalidValue(auth_plugin_data_len, cls, 'auth_plugin_data_len') auth_plugin_data_2_len = auth_plugin_data_len - 8 parser.parse_raw('auth_plugin_data_2', auth_plugin_data_2_len) if MySQLCapability.CLIENT_PLUGIN_AUTH in capabilities: parser.parse_string_null_terminated('auth_plugin_name', 'ascii') params = dict(parser) params['capabilities'] = capabilities return cls(**params), parser.parsed_length def compose(self): composer = ComposerBinary(byte_order=ByteOrder.LITTLE_ENDIAN) composer.compose_numeric(self.protocol_version, 1) composer.compose_string_null_terminated(self.server_version, 'ascii') composer.compose_numeric(self.connection_id, 4) composer.compose_raw(self.auth_plugin_data) composer.compose_raw(b'\x00') # filler capabilities = [capability for capability in self.capabilities if capability.value < 2 ** 16] composer.compose_numeric_flags(capabilities, 2) capabilities_2 = [capability for capability in self.capabilities if capability.value >= 2 ** 16] composer.compose_parsable(self.character_set) composer.compose_numeric_flags(self.states, 2) composer.compose_numeric_flags(capabilities_2, 2, shift_right=16) if MySQLCapability.CLIENT_PLUGIN_AUTH in self.capabilities: auth_plugin_data_len = 8 if self.auth_plugin_data_2: auth_plugin_data_len += len(self.auth_plugin_data_2) composer.compose_numeric(auth_plugin_data_len, 1) else: composer.compose_numeric(0, 1) composer.compose_raw(10 * b'\x00') # reserved if self.auth_plugin_data_2: composer.compose_raw(self.auth_plugin_data_2) if MySQLCapability.CLIENT_PLUGIN_AUTH in self.capabilities: composer.compose_string_null_terminated(self.auth_plugin_name, 'ascii') return composer.composed_bytes @attr.s class MySQLHandshakeSslRequest(MySQLPacketBase): capabilities = attr.ib(validator=attr.validators.deep_iterable( member_validator=attr.validators.instance_of(MySQLCapability), )) max_packet_size = attr.ib(default=0xffff, validator=attr.validators.instance_of(int)) character_set = attr.ib(default=None, validator=attr.validators.optional(attr.validators.in_(MySQLCharacterSet))) MINIMUM_SIZE = 5 def __attrs_post_init__(self): if MySQLCapability.CLIENT_PROTOCOL_41 in self.capabilities: if self.character_set is None: self.character_set = MySQLCharacterSet.UTF8 else: if self.max_packet_size >= 2 ** 24: raise ValueError(self.max_packet_size) for capability in self.capabilities: if capability.value >= 2 ** 16: raise ValueError(self.max_packet_size) @classmethod def _parse(cls, parsable): if len(parsable) < cls.MINIMUM_SIZE: raise NotEnoughData(cls.MINIMUM_SIZE - len(parsable)) parser = ParserBinary(parsable, byte_order=ByteOrder.LITTLE_ENDIAN) parser.parse_numeric_flags('capabilities', 2, MySQLCapability) if MySQLCapability.CLIENT_PROTOCOL_41 in parser['capabilities']: parser.parse_numeric_flags('capabilities_2', 2, MySQLCapability, shift_left=16) parser.parse_numeric('max_packet_size', 4) parser.parse_parsable('character_set', MySQLCharacterSetFactory) parser.parse_raw('reserved', 23) del parser['reserved'] character_set = parser['character_set'] capabilities = parser['capabilities'] | parser['capabilities_2'] else: parser.parse_numeric('max_packet_size', 3) capabilities = parser['capabilities'] character_set = None return cls(capabilities, parser['max_packet_size'], character_set), parser.parsed_length def compose(self): composer = ComposerBinary(byte_order=ByteOrder.LITTLE_ENDIAN) if MySQLCapability.CLIENT_PROTOCOL_41 in self.capabilities: composer.compose_numeric_flags(self.capabilities, 4) composer.compose_numeric(self.max_packet_size, 4) composer.compose_parsable(self.character_set) composer.compose_raw(23 * b'\x00') else: composer.compose_numeric_flags(self.capabilities, 2) composer.compose_numeric(self.max_packet_size, 3) return composer.composed_bytes cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/openvpn.py000066400000000000000000000156031524413560000274550ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import collections import enum import attr from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.base import VariantParsable from cryptoparser.common.exception import InvalidType, NotEnoughData from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary class OpenVpnOpCode(enum.IntEnum): CONTROL_V1 = 0x04 ACK_V1 = 0x05 HARD_RESET_CLIENT_V2 = 0x07 HARD_RESET_SERVER_V2 = 0x08 class OpenVpnPacketWrapperTcp(ParsableBase): def __init__(self, payload): self.payload = payload @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_bytes('payload', 2) return OpenVpnPacketWrapperTcp(parser['payload']), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_bytes(self.payload, 2) return composer.composed_bytes @attr.s class OpenVpnPacketBase(ParsableBase): HEADER_SIZE = 10 session_id = attr.ib(validator=attr.validators.instance_of(int)) packet_id_array = attr.ib( validator=attr.validators.deep_iterable(member_validator=attr.validators.instance_of(int)) ) remote_session_id = attr.ib(validator=attr.validators.optional(attr.validators.instance_of(int))) @classmethod @abc.abstractmethod def get_op_code(cls): return NotImplementedError() # pragma: no cover @classmethod def parse_header(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric('packet_type', 1) if parser['packet_type'] >> 3 != cls.get_op_code(): raise InvalidType() parser.parse_numeric('session_id', 8) parser.parse_numeric('packet_id_array_length', 1) if parser['packet_id_array_length']: parser.parse_numeric_array('packet_id_array', parser['packet_id_array_length'], 4) parser.parse_numeric('remote_session_id', 8) packet_id_array = parser['packet_id_array'] remote_session_id = parser['remote_session_id'] else: packet_id_array = [] remote_session_id = None return parser['session_id'], packet_id_array, remote_session_id, parser.parsed_length def _compose_header(self): composer = ComposerBinary() composer.compose_numeric(self.get_op_code() << 3, 1) composer.compose_numeric(self.session_id, 8) composer.compose_numeric(len(self.packet_id_array), 1) if self.packet_id_array: composer.compose_numeric_array(self.packet_id_array, 4) composer.compose_numeric(self.remote_session_id, 8) return composer.composed_bytes @attr.s class OpenVpnPacketControlV1(OpenVpnPacketBase): packet_id = attr.ib(validator=attr.validators.instance_of(int)) payload = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def get_op_code(cls): return OpenVpnOpCode.CONTROL_V1 def compose(self): composer = ComposerBinary() composer.compose_numeric(self.packet_id, 4) composer.compose_raw(self.payload) return self._compose_header() + composer.composed_bytes @classmethod def _parse(cls, parsable): session_id, packet_id_array, remote_session_id, header_length = cls.parse_header(parsable) body_parser = ParserBinary(parsable[header_length:]) body_parser.parse_numeric('packet_id', 4) body_parser.parse_raw('payload', body_parser.unparsed_length) return OpenVpnPacketControlV1( session_id, packet_id_array, remote_session_id, body_parser['packet_id'], body_parser['payload'] ), header_length + body_parser.parsed_length @attr.s(init=False) class OpenVpnPacketAckV1(OpenVpnPacketBase): def __init__(self, session_id, remote_session_id, packet_id_array): super().__init__(session_id, packet_id_array, remote_session_id) @classmethod def get_op_code(cls): return OpenVpnOpCode.ACK_V1 def compose(self): return self._compose_header() @classmethod def _parse(cls, parsable): session_id, packet_id_array, remote_session_id, header_length = cls.parse_header(parsable) return OpenVpnPacketAckV1(session_id, remote_session_id, packet_id_array), header_length @attr.s(init=False) class OpenVpnPacketHardResetClientV2(OpenVpnPacketBase): def __init__(self, session_id, packet_id): super().__init__(session_id, packet_id_array=[], remote_session_id=None) self.packet_id = packet_id @classmethod def get_op_code(cls): return OpenVpnOpCode.HARD_RESET_CLIENT_V2 @classmethod def _parse(cls, parsable): session_id, packet_id_array, _, header_length = cls.parse_header(parsable) if packet_id_array: raise InvalidValue(packet_id_array, cls, 'packet_id_array') body_parser = ParserBinary(parsable[header_length:]) body_parser.parse_numeric('packet_id', 4) return OpenVpnPacketHardResetClientV2( session_id, body_parser['packet_id'] ), header_length + body_parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.packet_id, 4) return self._compose_header() + composer.composed_bytes @attr.s(init=False) class OpenVpnPacketHardResetServerV2(OpenVpnPacketBase): def __init__(self, session_id, remote_session_id, packet_id_array, packet_id): super().__init__(session_id, packet_id_array, remote_session_id) self.packet_id = packet_id @classmethod def get_op_code(cls): return OpenVpnOpCode.HARD_RESET_SERVER_V2 @classmethod def _parse(cls, parsable): session_id, packet_id_array, remote_session_id, header_length = cls.parse_header(parsable) body_parser = ParserBinary(parsable[header_length:]) body_parser.parse_numeric('packet_id', 4) return OpenVpnPacketHardResetServerV2( session_id, remote_session_id, packet_id_array, body_parser['packet_id'] ), header_length + body_parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.packet_id, 4) return self._compose_header() + composer.composed_bytes class OpenVpnPacketVariant(VariantParsable): _VARIANTS = collections.OrderedDict([ (OpenVpnOpCode.ACK_V1, [OpenVpnPacketAckV1, ]), (OpenVpnOpCode.CONTROL_V1, [OpenVpnPacketControlV1, ]), (OpenVpnOpCode.HARD_RESET_CLIENT_V2, [OpenVpnPacketHardResetClientV2, ]), (OpenVpnOpCode.HARD_RESET_SERVER_V2, [OpenVpnPacketHardResetServerV2, ]), ]) @classmethod def _get_variants(cls): return cls._VARIANTS cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/postgresql.py000066400000000000000000000031351524413560000301700ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import attr from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary from cryptoparser.common.exception import NotEnoughData @attr.s class Sync(ParsableBase): MESSAGE_SIZE = 1 COMMAND = b'S' @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_raw('command', cls.MESSAGE_SIZE) if parser['command'] != cls.COMMAND: raise InvalidValue(parser['command'], cls, 'command') return cls(), cls.MESSAGE_SIZE def compose(self): composer = ComposerBinary() composer.compose_raw(self.COMMAND) return composer.composed_bytes class SslRequest(ParsableBase): MESSAGE_SIZE = 8 REQUEST_CODE = 80877103 @classmethod def _parse(cls, parsable): if len(parsable) < cls.MESSAGE_SIZE: raise NotEnoughData(cls.MESSAGE_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric('length', 4) if parser['length'] != cls.MESSAGE_SIZE: raise InvalidValue(parser['length'], cls, 'length') parser.parse_numeric('request_code', 4) if parser['request_code'] != cls.REQUEST_CODE: raise InvalidValue(parser['request_code'], cls, 'request_code') return cls(), cls.MESSAGE_SIZE def compose(self): composer = ComposerBinary() composer.compose_numeric(self.MESSAGE_SIZE, 4) composer.compose_numeric(self.REQUEST_CODE, 4) return composer.composed_bytes cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/rdp.py000066400000000000000000000155541524413560000265620ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import enum import attr from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary, ByteOrder from cryptoparser.common.exception import NotEnoughData, InvalidType @attr.s class TPKT(ParsableBase): HEADER_SIZE = 4 version = attr.ib(validator=attr.validators.instance_of(int)) message = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def _parse(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric('version', 1) if parser['version'] != 3: raise InvalidValue(parser['version'], TPKT, 'version') parser.parse_numeric('reserved', 1) parser.parse_numeric('packet_length', 2) if len(parsable) < parser['packet_length']: raise NotEnoughData(parser['packet_length'] - len(parsable)) parser.parse_raw('message', parser['packet_length'] - 4) return TPKT(parser['version'], parser['message']), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.version, 1) composer.compose_numeric(0, 1) # reserved composer.compose_numeric(len(self.message) + 4, 2) composer.compose_raw(self.message) return composer.composed_bytes class COTPType(enum.IntEnum): CONNECTION_REQUEST = 0xe CONNECTION_CONFIRM = 0xd DISCONNECT_REQUEST = 0x8 DISCONNECT_CONFIRM = 0xc DATA = 0xf EXPEDITED_DATA = 0x1 DATA_ACKNOWLEDGEMENT = 0x6 EXPEDITED_DATA_ANOWLEDGEMENT = 0x2 REJECT = 0x5 @attr.s class COTPConnectionBase(ParsableBase): HEADER_SIZE = 7 src_ref = attr.ib(validator=attr.validators.instance_of(int)) user_data = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) dst_ref = attr.ib(default=0, validator=attr.validators.instance_of(int)) class_option = attr.ib(default=0) @classmethod @abc.abstractmethod def _get_type(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric('length_indicator', 1) if parser.unparsed_length < parser['length_indicator']: raise NotEnoughData(parser['length_indicator'] - parser.unparsed_length) parser.parse_numeric('pdu_type', 1) pdu_type = parser['pdu_type'] >> 4 if pdu_type != cls._get_type(): raise InvalidType() parser.parse_numeric('src_ref', 2) parser.parse_numeric('dst_ref', 2) parser.parse_numeric('class_option', 1) parser.parse_raw('user_data', parser['length_indicator'] - parser.parsed_length + 1) return COTPConnectionRequest( src_ref=parser['src_ref'], dst_ref=parser['dst_ref'], class_option=parser['class_option'], user_data=parser['user_data'], ), parser.parsed_length def compose(self): body_composer = ComposerBinary() body_composer.compose_numeric(self._get_type() << 4, 1) body_composer.compose_numeric(self.src_ref, 2) body_composer.compose_numeric(self.dst_ref, 2) body_composer.compose_numeric(self.class_option, 1) body_composer.compose_raw(self.user_data) body = body_composer.composed_bytes header_composer = ComposerBinary() header_composer.compose_numeric(len(body), 1) return header_composer.composed_bytes + body def __attrs_post_init__(self): if self.class_option != 0: raise InvalidValue(self.class_option, COTPConnectionRequest, 'class_option') @attr.s class COTPConnectionRequest(COTPConnectionBase): @classmethod def _get_type(cls): return COTPType.CONNECTION_REQUEST @attr.s class COTPConnectionConfirm(COTPConnectionBase): @classmethod def _get_type(cls): return COTPType.CONNECTION_CONFIRM class RDPProtocol(enum.IntEnum): RDP = 0x00000000 SSL = 0x00000001 HYBRID = 0x00000002 RDSTLS = 0x00000004 HYBRID_EX = 0x00000008 class RDPNegotiationRequestFlags(enum.IntEnum): RESTRICTED_ADMIN_MODE_REQUIRED = 0x01 REDIRECTED_AUTHENTICATION_MODE_REQUIRED = 0x02 CORRELATION_INFO_PRESENT = 0x08 class RDPNegotiationResponseFlags(enum.IntEnum): EXTENDED_CLIENT_DATA_SUPPORTED = 0x01 DYNVC_GFX_PROTOCOL_SUPPORTED = 0x02 NEGRSP_FLAG_RESERVED = 0x04 RESTRICTED_ADMIN_MODE_SUPPORTED = 0x08 REDIRECTED_AUTHENTICATION_MODE_SUPPORTED = 0x10 class RDPPacketType(enum.IntEnum): NEG_REQ = 1 NEG_RSP = 2 @attr.s class RDPNegotiationBase(ParsableBase): PACKET_LENGTH = 8 flags = attr.ib(validator=attr.validators.deep_iterable(member_validator=attr.validators.instance_of( (RDPNegotiationRequestFlags, RDPNegotiationResponseFlags) ))) protocol = attr.ib(validator=attr.validators.deep_iterable(member_validator=attr.validators.in_(RDPProtocol))) @classmethod @abc.abstractmethod def _get_type(cls): raise NotImplementedError() @classmethod @abc.abstractmethod def _get_flag_type(cls): raise NotImplementedError() @classmethod def _parse(cls, parsable): if len(parsable) < cls.PACKET_LENGTH: raise NotEnoughData(cls.PACKET_LENGTH - len(parsable)) parser = ParserBinary(parsable, ByteOrder.LITTLE_ENDIAN) parser.parse_numeric('type', 1, RDPPacketType) if parser['type'] != cls._get_type(): raise InvalidType() parser.parse_numeric_flags('flags', 1, cls._get_flag_type()) parser.parse_numeric('length', 2) if parser['length'] != cls.PACKET_LENGTH: raise InvalidValue(parser['length'], cls, 'packet length') parser.parse_numeric_flags('protocol', 4, RDPProtocol) return cls(parser['flags'], parser['protocol']), parser.parsed_length def compose(self): composer = ComposerBinary(ByteOrder.LITTLE_ENDIAN) composer.compose_numeric(self._get_type(), 1) composer.compose_numeric_flags(self.flags, 1) composer.compose_numeric(self.PACKET_LENGTH, 2) composer.compose_numeric_flags(self.protocol, 4) return composer.composed_bytes @attr.s class RDPNegotiationRequest(RDPNegotiationBase): @classmethod def _get_type(cls): return RDPPacketType.NEG_REQ @classmethod def _get_flag_type(cls): return RDPNegotiationRequestFlags @attr.s class RDPNegotiationResponse(RDPNegotiationBase): @classmethod def _get_type(cls): return RDPPacketType.NEG_RSP @classmethod def _get_flag_type(cls): return RDPNegotiationResponseFlags cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/record.py000066400000000000000000000075451524413560000272540ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import attr from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary from cryptoparser.common.exception import NotEnoughData from cryptoparser.tls.version import TlsVersion, TlsProtocolVersion from cryptoparser.tls.subprotocol import TlsContentType from cryptoparser.tls.subprotocol import SslMessageBase, SslMessageType, SslSubprotocolMessageParser @attr.s class TlsRecord(ParsableBase): HEADER_SIZE = 5 fragment = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) protocol_version = attr.ib( default=TlsProtocolVersion(TlsVersion.TLS1), validator=attr.validators.instance_of(TlsProtocolVersion), ) content_type = attr.ib( default=TlsContentType.HANDSHAKE, validator=attr.validators.instance_of(TlsContentType), ) @classmethod def parse_header(cls, parsable): if len(parsable) < cls.HEADER_SIZE: raise NotEnoughData(cls.HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) try: parser.parse_numeric('content_type', 1, TlsContentType) except InvalidValue as e: raise InvalidValue(e.value, TlsContentType) from e parser.parse_parsable('protocol_version', TlsProtocolVersion) parser.parse_numeric('fragment_length', 2) return parser @classmethod def _parse(cls, parsable): parser = cls.parse_header(parsable) parser.parse_raw('fragment', parser['fragment_length']) return TlsRecord( content_type=parser['content_type'], protocol_version=parser['protocol_version'], fragment=parser['fragment'], ), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.content_type, 1) composer.compose_parsable(self.protocol_version) composer.compose_bytes(self.fragment, 2) return composer.composed_bytes @attr.s class SslRecord(ParsableBase): message = attr.ib(validator=attr.validators.instance_of(SslMessageBase)) protocol_version = attr.ib(init=False, default=TlsVersion.SSL2, validator=attr.validators.in_(TlsVersion)) @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('record_length_0', 1) parser.parse_numeric('record_length_1', 1) if parser['record_length_0'] & 0x80: record_length = ((parser['record_length_0'] & 0x7f) * (2 ** 8)) + parser['record_length_1'] padding_length = 0 else: record_length = ((parser['record_length_0'] & 0x3f) * (2 ** 8)) + parser['record_length_1'] parser.parse_numeric('padding_length', 1) padding_length = parser['padding_length'] if record_length > parser.unparsed_length: raise NotEnoughData(record_length - parser.unparsed_length) try: parser.parse_numeric('message_type', 1, SslMessageType) except InvalidValue as e: raise InvalidValue(e.value, SslMessageType) from e parser.parse_variant('message', SslSubprotocolMessageParser(parser['message_type'])) parser.parse_raw('padding', padding_length) return SslRecord(message=parser['message']), parser.parsed_length def compose(self): body_composer = ComposerBinary() message_type = self.message.get_message_type() body_composer.compose_numeric(message_type, 1) body_composer.compose_parsable(self.message) header_composer = ComposerBinary() header_composer.compose_numeric(body_composer.composed_length | (2 ** 15), 2) return header_composer.composed_bytes + body_composer.composed_bytes @property def content_type(self): return self.message.get_message_type() cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/subprotocol.py000066400000000000000000001346631524413560000303530ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import abc import calendar import collections import datetime import enum import hashlib import random import attr from cryptodatahub.common.exception import InvalidValue from cryptodatahub.tls.algorithm import SslCipherKind, TlsCipherSuite, TlsCipherSuiteExtension, TlsCompressionMethod from cryptoparser.common.base import ( OneByteEnumParsable, Opaque, OpaqueParam, VariantParsable, Vector, VectorEnumCodeNumeric, VectorParamEnumCodeNumeric, VectorParamNumeric, VectorParamParsable, VectorParsable, ) from cryptoparser.common.exception import NotEnoughData, InvalidType from cryptoparser.common.parse import ParsableBase, ParserBinary, ComposerBinary, ComposerText from cryptoparser.tls.extension import ( TlsCertificateStatusType, TlsExtensionType, TlsExtensionsClient, TlsExtensionsServer, TlsSignatureAndHashAlgorithmVector, ) from cryptoparser.tls.grease import TlsInvalidType, TlsInvalidTypeOneByte, TlsInvalidTypeTwoByte from cryptoparser.tls.version import TlsProtocolVersion, TlsVersion from cryptoparser.tls.ciphersuite import SslCipherKindFactory, TlsCipherSuiteFactory class TlsContentType(enum.IntEnum): CHANGE_CIPHER_SPEC = 0x14 ALERT = 0x15 HANDSHAKE = 0x16 APPLICATION_DATA = 0x17 HEARTBEAT = 0x18 @attr.s class SubprotocolParser: _subprotocol_type = attr.ib(validator=attr.validators.instance_of(enum.IntEnum)) @classmethod @abc.abstractmethod def _get_subprotocol_parsers(cls): raise NotImplementedError() @classmethod def register_subprotocol_parser(cls, subprotocol_type, parsable_class): subprotocol_parsers = cls._get_subprotocol_parsers() subprotocol_parsers[subprotocol_type] = parsable_class def parse(self, parsable): subprotocol_parsers = self._get_subprotocol_parsers() if self._subprotocol_type in subprotocol_parsers: parsed_object, parsed_length = subprotocol_parsers[self._subprotocol_type].parse_immutable(parsable) return parsed_object, parsed_length raise InvalidValue(self._subprotocol_type, TlsSubprotocolMessageBase) class TlsSubprotocolMessageBase(ParsableBase): @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() class TlsAlertLevel(enum.IntEnum): WARNING = 0x01 FATAL = 0x02 class TlsAlertDescription(enum.IntEnum): CLOSE_NOTIFY = 0x00 UNEXPECTED_MESSAGE = 0x0a BAD_RECORD_MAC = 0x14 RECORD_OVERFLOW = 0x16 HANDSHAKE_FAILURE = 0x28 BAD_CERTIFICATE = 0x2a UNSUPPORTED_CERTIFICATE = 0x2b CERTIFICATE_REVOKED = 0x2c CERTIFICATE_EXPIRED = 0x2d CERTIFICATE_UNKNOWN = 0x2e ILLEGAL_PARAMETER = 0x2f UNKNOWN_CA = 0x30 ACCESS_DENIED = 0x30 DECODE_ERROR = 0x32 DECRYPT_ERROR = 0x33 PROTOCOL_VERSION = 0x46 INSUFFICIENT_SECURITY = 0x47 INTERNAL_ERROR = 0x50 INAPPROPRIATE_FALLBACK = 0x56 USER_CANCELED = 0x5a MISSING_EXTENSION = 0x6d UNSUPPORTED_EXTENSION = 0x6e CERTIFICATE_UNOBTAINABLE = 0x6f UNRECOGNIZED_NAME = 0x70 BAD_CERTIFICATE_STATUS_RESPONSE = 0x71 BAD_CERTIFICATE_HASH_VALUE = 0x72 UNKNOWN_PSK_IDENTITY = 0x73 CERTIFICATE_REQUIRED = 0x74 NO_APPLICATION_PROTOCOL = 0x78 @attr.s class TlsAlertMessage(TlsSubprotocolMessageBase): _SIZE = 2 level = attr.ib() description = attr.ib() @classmethod def _parse(cls, parsable): if len(parsable) < cls._SIZE: raise NotEnoughData(cls._SIZE - len(parsable)) parser = ParserBinary(parsable) parser.parse_numeric('level', 1) parser.parse_numeric('description', 1) return TlsAlertMessage(parser['level'], parser['description']), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.level, 1) composer.compose_numeric(self.description, 1) return composer.composed_bytes @level.validator def _validator_level(self, attribute, value): # pylint: disable=unused-argument try: self.level = TlsAlertLevel(value) except ValueError as e: raise InvalidValue(value, TlsAlertLevel, 'level') from e @description.validator def _validator_description(self, attribute, value): # pylint: disable=unused-argument try: self.description = TlsAlertDescription(value) except ValueError as e: raise InvalidValue(value, TlsAlertDescription) from e class TlsChangeCipherSpecType(enum.IntEnum): CHANGE_CIPHER_SPEC = 0x01 @attr.s class TlsChangeCipherSpecMessage(TlsSubprotocolMessageBase): _change_cipher_spec_type = attr.ib( default=TlsChangeCipherSpecType.CHANGE_CIPHER_SPEC, validator=attr.validators.in_(TlsChangeCipherSpecType) ) @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('change_cipher_spec_type', 1, TlsChangeCipherSpecType) return TlsChangeCipherSpecMessage(parser['change_cipher_spec_type']), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self._change_cipher_spec_type, 1) return composer.composed_bytes @attr.s class TlsApplicationDataMessage(TlsSubprotocolMessageBase): data = attr.ib(attr.validators.instance_of(bytearray)) @classmethod def _parse(cls, parsable): return TlsApplicationDataMessage(parsable), len(parsable) def compose(self): return self.data class TlsHandshakeType(enum.IntEnum): HELLO_REQUEST = 0x00 CLIENT_HELLO = 0x01 SERVER_HELLO = 0x02 HELLO_VERIFY_REQUEST = 0x03 NEW_SESSION_TICKET = 0x04 HELLO_RETRY_REQUEST = 0x06 ENCRYPTED_EXTENSIONS = 0x08 CERTIFICATE = 0x0b SERVER_KEY_EXCHANGE = 0x0c CERTIFICATE_REQUEST = 0x0d SERVER_HELLO_DONE = 0x0e CERTIFICATE_VERIFY = 0x0f CLIENT_KEY_EXCHANGE = 0x10 CLIENT_CERTIFICATE_REQUEST = 0x11 FINISHED = 0x14 CLIENT_CERTIFICATE_URL = 0x15 CERTIFICATE_STATUS = 0x16 SUPPLEMENTAL_DATA = 0x17 KEY_UPDATE = 0x18 COMPRESSED_CERTIFICATE = 0x19 EKT_KEY = 0x15 MESSAGE_HASH = 254 @attr.s class TlsHandshakeMessage(TlsSubprotocolMessageBase): """The payload of a handshake record. """ _HEADER_SIZE = 4 @classmethod @abc.abstractmethod def get_handshake_type(cls): raise NotImplementedError() @classmethod def _parse_handshake_header(cls, parsable): if len(parsable) < cls._HEADER_SIZE: raise NotEnoughData(cls._HEADER_SIZE - len(parsable)) parser = ParserBinary(parsable) try: parser.parse_numeric('handshake_type', 1, TlsHandshakeType) except InvalidValue as e: raise e if parser['handshake_type'] != cls.get_handshake_type(): raise InvalidType() try: parser.parse_bytes('payload', 3) except NotEnoughData as e: raise NotEnoughData(e.bytes_needed) from e return parser def _compose_header(self, payload_length): composer = ComposerBinary() composer.compose_numeric(self.get_handshake_type(), 1) composer.compose_numeric(payload_length, 3) return composer.composed_bytes class TlsHandshakeHelloRandomBytes(Vector): @classmethod def _parse(cls, parsable): composer = ComposerBinary() vector_param = cls.get_param() composer.compose_numeric(vector_param.min_byte_num, vector_param.item_num_size) vector, parsed_length = super()._parse(composer.composed_bytes + parsable) return cls(vector), parsed_length - vector_param.item_num_size def compose(self): return super().compose()[self.get_param().item_num_size:] @classmethod def get_param(cls): return OpaqueParam(min_byte_num=28, max_byte_num=28) @attr.s class TlsHandshakeHelloRandom(ParsableBase): time = attr.ib(validator=attr.validators.instance_of(datetime.datetime)) random = attr.ib(validator=attr.validators.instance_of(TlsHandshakeHelloRandomBytes)) @time.default def _default_time(self): # pylint: disable=no-self-use return datetime.datetime.now(datetime.timezone.utc) @random.default def _default_random(self): # pylint: disable=no-self-use return TlsHandshakeHelloRandomBytes( bytearray.fromhex(f'{random.getrandbits(224):28x}'.zfill(56)) ) @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_timestamp('time', item_size=4) parser.parse_parsable('random', TlsHandshakeHelloRandomBytes) return TlsHandshakeHelloRandom(parser['time'], parser['random']), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(int(calendar.timegm(self.time.utctimetuple())), 4) composer.compose_parsable(self.random) return composer.composed_bytes @attr.s class TlsHandshakeHello(TlsHandshakeMessage): @classmethod def _parse_hello_header(cls, parsable): parser = ParserBinary(parsable) parser.parse_parsable('protocol_version', TlsProtocolVersion) parser.parse_parsable('random', TlsHandshakeHelloRandom) parser.parse_parsable('session_id', TlsSessionIdVector) return parser def _compose_header(self, payload_length): composer = ComposerBinary() handshake_header_bytes = super()._compose_header( payload_length + composer.composed_length ) return handshake_header_bytes + composer.composed_bytes @classmethod def _parse_extensions(cls, handshake_header_parser, parser, extensions_class): if parser.parsed_length >= len(handshake_header_parser['payload']): return None parser.parse_parsable('extensions', extensions_class) return parser @staticmethod def _compose_extensions(extensions): extension_bytes = bytearray() for extension in extensions: extension_bytes += extension.compose() payload_composer = ComposerBinary() if extensions: payload_composer.compose_numeric(len(extension_bytes), 2) return payload_composer.composed_bytes + extension_bytes @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() @abc.abstractmethod def compose(self): raise NotImplementedError() class TlsCipherSuiteVector(VectorParsable): @classmethod def get_param(cls): return VectorParamEnumCodeNumeric( item_class=TlsCipherSuiteFactory, fallback_class=TlsInvalidTypeTwoByte, min_byte_num=2, max_byte_num=2 ** 16 - 2 ) class TlsCompressionMethodFactory(OneByteEnumParsable): @classmethod def get_enum_class(cls): return TlsCompressionMethod @abc.abstractmethod def compose(self): raise NotImplementedError() class TlsCompressionMethodVector(VectorEnumCodeNumeric): @classmethod def get_param(cls): return VectorParamEnumCodeNumeric( item_class=TlsCompressionMethodFactory, fallback_class=TlsInvalidTypeOneByte, min_byte_num=1, max_byte_num=2 ** 8 - 1, ) class TlsSessionIdVector(Vector): @classmethod def get_param(cls): return VectorParamNumeric(item_size=1, min_byte_num=0, max_byte_num=32) @attr.s(frozen=True) class TlsJA4Fingerprint: fingerprint = attr.ib(validator=attr.validators.instance_of(str)) fingerprint_original = attr.ib(validator=attr.validators.instance_of(str)) fingerprint_raw = attr.ib(validator=attr.validators.instance_of(str)) fingerprint_raw_original = attr.ib(validator=attr.validators.instance_of(str)) @attr.s class TlsHandshakeClientHello(TlsHandshakeHello): # pylint: disable=too-many-instance-attributes cipher_suites = attr.ib( converter=TlsCipherSuiteVector, validator=attr.validators.instance_of(TlsCipherSuiteVector) ) protocol_version = attr.ib( default=TlsProtocolVersion(TlsVersion.TLS1_2), validator=attr.validators.instance_of(TlsProtocolVersion), ) random = attr.ib( default=TlsHandshakeHelloRandom(), validator=attr.validators.instance_of(TlsHandshakeHelloRandom), ) session_id = attr.ib( default=TlsSessionIdVector(()), converter=TlsSessionIdVector, validator=attr.validators.instance_of(TlsSessionIdVector), ) compression_methods = attr.ib( default=TlsCompressionMethodVector([TlsCompressionMethod.NULL, ]), converter=TlsCompressionMethodVector, validator=attr.validators.instance_of(TlsCompressionMethodVector), ) extensions = attr.ib( default=TlsExtensionsClient(()), converter=TlsExtensionsClient, validator=attr.validators.instance_of(TlsExtensionsClient) ) fallback_scsv = attr.ib(default=False, validator=attr.validators.instance_of(bool)) empty_renegotiation_info_scsv = attr.ib(default=True, validator=attr.validators.instance_of(bool)) @classmethod def get_handshake_type(cls): return TlsHandshakeType.CLIENT_HELLO @classmethod def _parse(cls, parsable): handshake_header_parser = cls._parse_handshake_header(parsable) parser = cls._parse_hello_header(handshake_header_parser['payload']) parser.parse_parsable('cipher_suites', TlsCipherSuiteVector) parser.parse_parsable('compression_methods', TlsCompressionMethodVector) extension_parser = cls._parse_extensions(handshake_header_parser, parser, TlsExtensionsClient) cipher_suites = [] fallback_scsv = False empty_renegotiation_info_scsv = False for cipher_suite in parser['cipher_suites']: if cipher_suite.value.code == TlsCipherSuiteExtension.FALLBACK_SCSV.value.code: fallback_scsv = True elif cipher_suite.value.code == TlsCipherSuiteExtension.EMPTY_RENEGOTIATION_INFO_SCSV.value.code: empty_renegotiation_info_scsv = True else: cipher_suites.append(cipher_suite) return TlsHandshakeClientHello( cipher_suites, parser['protocol_version'], parser['random'], parser['session_id'], parser['compression_methods'], extensions=parser['extensions'] if extension_parser else TlsExtensionsClient([]), fallback_scsv=fallback_scsv, empty_renegotiation_info_scsv=empty_renegotiation_info_scsv, ), handshake_header_parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.protocol_version) payload_composer.compose_parsable(self.random) payload_composer.compose_parsable(self.session_id) if self.fallback_scsv: self.cipher_suites.append(TlsCipherSuiteExtension.FALLBACK_SCSV) if self.empty_renegotiation_info_scsv: self.cipher_suites.append(TlsCipherSuiteExtension.EMPTY_RENEGOTIATION_INFO_SCSV) payload_composer.compose_numeric(len(self.cipher_suites) * self.cipher_suites.get_param().item_num_size, 2) payload_composer.compose_numeric_array_enum_coded(self.cipher_suites) if self.fallback_scsv: del self.cipher_suites[-1] if self.empty_renegotiation_info_scsv: del self.cipher_suites[-1] payload_composer.compose_parsable(self.compression_methods) extension_bytes = self._compose_extensions(self.extensions) header_bytes = self._compose_header(payload_composer.composed_length + len(extension_bytes)) return header_bytes + payload_composer.composed_bytes + extension_bytes def ja3(self): parser = ParserBinary(self.protocol_version.compose()) parser.parse_numeric('tls_protocol_version', 2) cipher_suites = [cipher_suite.value.code for cipher_suite in self.cipher_suites] extension_types = [] named_curves = [] ec_point_formats = [] for extension in self.extensions: if (not isinstance(extension.extension_type, TlsInvalidTypeTwoByte) or extension.extension_type.value.value_type != TlsInvalidType.GREASE): extension_types.append(extension.extension_type.value.code) if extension.extension_type == TlsExtensionType.SUPPORTED_GROUPS: named_curves = [ named_curve.value.code for named_curve in extension.elliptic_curves if (not isinstance(named_curve, TlsInvalidTypeTwoByte) or named_curve.value.value_type != TlsInvalidType.GREASE) ] elif extension.extension_type == TlsExtensionType.EC_POINT_FORMATS: ec_point_formats = [ point_format.value.code for point_format in extension.point_formats if (not isinstance(point_format, TlsInvalidTypeOneByte) or point_format.value.value_type != TlsInvalidType.GREASE) ] composer = ComposerText() composer.compose_numeric(parser['tls_protocol_version']) for numeric_array in (cipher_suites, extension_types, named_curves, ec_point_formats): composer.compose_separator(',') composer.compose_numeric_array(numeric_array, '-') return composer.composed.decode('ascii') _JA4_VERSION_STRINGS = { TlsVersion.TLS1_3: '13', TlsVersion.TLS1_2: '12', TlsVersion.TLS1_1: '11', TlsVersion.TLS1: '10', TlsVersion.SSL3: 's3', TlsVersion.SSL2: 's2', } @staticmethod def _ja4_is_grease(item): return isinstance(item, TlsInvalidTypeTwoByte) and item.value.value_type == TlsInvalidType.GREASE @staticmethod def _ja4_sha256(data): return hashlib.sha256(data).hexdigest()[:12] @classmethod def _ja4_hashes(cls, cipher_hexes, extension_hexes, signature_algorithm_hexes): cipher_composer = ComposerText() cipher_composer.compose_string_array(cipher_hexes, ',') cipher_hash = cls._ja4_sha256(cipher_composer.composed) if cipher_hexes else '000000000000' extension_composer = ComposerText() extension_composer.compose_string_array(extension_hexes, ',') if signature_algorithm_hexes: extension_composer.compose_separator('_') extension_composer.compose_string_array(signature_algorithm_hexes, ',') extension_hash = cls._ja4_sha256(extension_composer.composed) if extension_hexes else '000000000000' return cipher_hash, extension_hash @staticmethod def _ja4_raw(header, cipher_hexes, extension_hexes, signature_algorithm_hexes): composer = ComposerText() composer.compose_string(header) for hexes in (cipher_hexes, extension_hexes, signature_algorithm_hexes): composer.compose_separator('_') composer.compose_string_array(hexes, ',') return composer.composed.decode('ascii') @staticmethod def _ja4_alpn_characters(value): return value[0] + value[-1] def _ja4_version_string(self): protocol_versions = [] for extension in self.extensions: if extension.extension_type == TlsExtensionType.SUPPORTED_VERSIONS: protocol_versions = [ version for version in extension.supported_versions if not self._ja4_is_grease(version) ] break if not protocol_versions: protocol_versions = [self.protocol_version] highest_protocol_version = max( protocol_versions, key=lambda protocol_version: protocol_version.version.value.code ) return self._JA4_VERSION_STRINGS.get(highest_protocol_version.version, '00') def _ja4_alpn_value(self): for extension in self.extensions: if extension.extension_type == TlsExtensionType.APPLICATION_LAYER_PROTOCOL_NEGOTIATION: return self._ja4_alpn_characters(list(extension.protocol_names)[0].value.code) return '00' def _ja4_signature_algorithm_hexes(self): for extension in self.extensions: if extension.extension_type == TlsExtensionType.SIGNATURE_ALGORITHMS: return [ f'{algorithm.value.code:04x}' for algorithm in extension.hash_and_signature_algorithms if not self._ja4_is_grease(algorithm) ] return [] def ja4(self): cipher_hexes = [ f'{cipher_suite.value.code:04x}' for cipher_suite in self.cipher_suites if not self._ja4_is_grease(cipher_suite) ] extension_hexes = [ f'{extension.extension_type.value.code:04x}' for extension in self.extensions if not self._ja4_is_grease(extension.extension_type) ] signature_algorithm_hexes = self._ja4_signature_algorithm_hexes() sni = 'd' if any( extension.extension_type == TlsExtensionType.SERVER_NAME for extension in self.extensions ) else 'i' header = ( f't{self._ja4_version_string()}{sni}' f'{min(len(cipher_hexes), 99):02d}{min(len(extension_hexes), 99):02d}' f'{self._ja4_alpn_value()}' ) excluded_extension_hexes = ( f'{TlsExtensionType.SERVER_NAME.value.code:04x}', f'{TlsExtensionType.APPLICATION_LAYER_PROTOCOL_NEGOTIATION.value.code:04x}', ) sorted_cipher_hexes = sorted(cipher_hexes) sorted_extension_hexes = sorted( extension_hex for extension_hex in extension_hexes if extension_hex not in excluded_extension_hexes ) cipher_hash, extension_hash = self._ja4_hashes( sorted_cipher_hexes, sorted_extension_hexes, signature_algorithm_hexes ) cipher_hash_original, extension_hash_original = self._ja4_hashes( cipher_hexes, extension_hexes, signature_algorithm_hexes ) return TlsJA4Fingerprint( fingerprint=f'{header}_{cipher_hash}_{extension_hash}', fingerprint_original=f'{header}_{cipher_hash_original}_{extension_hash_original}', fingerprint_raw=self._ja4_raw( header, sorted_cipher_hexes, sorted_extension_hexes, signature_algorithm_hexes ), fingerprint_raw_original=self._ja4_raw( header, cipher_hexes, extension_hexes, signature_algorithm_hexes ), ) @attr.s class TlsHandshakeServerHello(TlsHandshakeHello): protocol_version = attr.ib( default=TlsProtocolVersion(TlsVersion.TLS1_2), validator=attr.validators.instance_of(TlsProtocolVersion), ) random = attr.ib( default=TlsHandshakeHelloRandom(), validator=attr.validators.instance_of(TlsHandshakeHelloRandom), ) session_id = attr.ib( default=TlsSessionIdVector(random.randint(0, 255) for i in range(32)), converter=TlsSessionIdVector, validator=attr.validators.instance_of(TlsSessionIdVector), ) compression_method = attr.ib( default=TlsCompressionMethod.NULL, converter=TlsCompressionMethod, validator=attr.validators.in_(TlsCompressionMethod), ) cipher_suite = attr.ib(default=None, validator=attr.validators.in_(TlsCipherSuite)) extensions = attr.ib( default=TlsExtensionsServer([]), converter=TlsExtensionsServer, validator=attr.validators.instance_of(TlsExtensionsServer) ) @classmethod def get_handshake_type(cls): return TlsHandshakeType.SERVER_HELLO @classmethod def _parse(cls, parsable): handshake_header_parser = cls._parse_handshake_header(parsable) parser = cls._parse_hello_header(handshake_header_parser['payload']) parser.parse_parsable('cipher_suite', TlsCipherSuiteFactory) parser.parse_parsable('compression_method', TlsCompressionMethodFactory) extension_parser = cls._parse_extensions(handshake_header_parser, parser, TlsExtensionsServer) return TlsHandshakeServerHello( protocol_version=parser['protocol_version'], random=parser['random'], session_id=parser['session_id'], compression_method=parser['compression_method'], cipher_suite=parser['cipher_suite'], extensions=parser['extensions'] if extension_parser else TlsExtensionsServer([]), ), handshake_header_parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.protocol_version) payload_composer.compose_parsable(self.random) payload_composer.compose_parsable(self.session_id) payload_composer.compose_numeric_enum_coded(self.cipher_suite) payload_composer.compose_numeric_enum_coded(self.compression_method) extension_bytes = self._compose_extensions(self.extensions) header_bytes = self._compose_header(payload_composer.composed_length + len(extension_bytes)) return header_bytes + payload_composer.composed_bytes + extension_bytes class TlsCertificateType(enum.IntEnum): X509 = 0 RAW_PUBLIC_KEY = 2 @attr.s class TlsCertificate(ParsableBase): certificate = attr.ib(validator=attr.validators.instance_of(bytes)) @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_bytes('certificate', 3) return TlsCertificate(bytes(parser['certificate'])), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_bytes(self.certificate, 3) return composer.composed_bytes class TlsCertificates(VectorParsable): @classmethod def get_param(cls): return VectorParamParsable( item_class=TlsCertificate, fallback_class=None, min_byte_num=1, max_byte_num=2 ** 24 - 1 ) @attr.s class TlsHandshakeServerCertificate(TlsHandshakeMessage): certificate_chain = attr.ib(validator=attr.validators.instance_of(TlsCertificates)) @classmethod def get_handshake_type(cls): return TlsHandshakeType.CERTIFICATE @classmethod def _parse(cls, parsable): handshake_header_parser = cls._parse_handshake_header(parsable) parser = ParserBinary(handshake_header_parser['payload']) parser.parse_parsable('certificates', TlsCertificates) return TlsHandshakeServerCertificate( parser['certificates'] ), handshake_header_parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.certificate_chain) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes @attr.s class TlsCertificateEntry(ParsableBase): certificate = attr.ib(validator=attr.validators.instance_of(TlsCertificate)) extensions = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) def __attrs_post_init__(self): if len(self.extensions) > 2 ** 16 - 1: raise InvalidValue(len(self.extensions), self.__class__, 'extensions') @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_parsable('certificate', TlsCertificate) parser.parse_bytes('extensions', 2) return TlsCertificateEntry( parser['certificate'], bytes(parser['extensions']), ), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_parsable(self.certificate) composer.compose_bytes(self.extensions, 2) return composer.composed_bytes class TlsCertificateEntryVector(VectorParsable): @classmethod def get_param(cls): return VectorParamParsable( item_class=TlsCertificateEntry, fallback_class=None, min_byte_num=0, max_byte_num=2 ** 24 - 1, ) @attr.s class TlsHandshakeCertificate(TlsHandshakeMessage): certificate_request_context = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) certificate_entries = attr.ib(validator=attr.validators.instance_of(TlsCertificateEntryVector)) def __attrs_post_init__(self): if len(self.certificate_request_context) > 255: raise InvalidValue(len(self.certificate_request_context), self.__class__, 'certificate_request_context') @classmethod def get_handshake_type(cls): return TlsHandshakeType.CERTIFICATE @classmethod def _parse(cls, parsable): handshake_header_parser = cls._parse_handshake_header(parsable) parser = ParserBinary(handshake_header_parser['payload']) parser.parse_numeric('certificate_request_context_length', 1) parser.parse_raw( 'certificate_request_context', parser['certificate_request_context_length'], ) try: parser.parse_parsable('certificate_entries', TlsCertificateEntryVector) except (NotEnoughData, InvalidValue) as e: raise InvalidType() from e return TlsHandshakeCertificate( bytes(parser['certificate_request_context']), parser['certificate_entries'], ), handshake_header_parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_numeric(len(self.certificate_request_context), 1) payload_composer.compose_raw(self.certificate_request_context) payload_composer.compose_parsable(self.certificate_entries) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class TlsHandshakeCertificateStatus(TlsHandshakeMessage): def __init__(self, status_type, status): super().__init__() self.status_type = status_type self.status = status @classmethod def get_handshake_type(cls): return TlsHandshakeType.CERTIFICATE_STATUS @classmethod def _parse(cls, parsable): handshake_header_parser = cls._parse_handshake_header(parsable) parser = ParserBinary(handshake_header_parser['payload']) parser.parse_numeric('status_type', 1, TlsCertificateStatusType) parser.parse_bytes('status', 3) return TlsHandshakeCertificateStatus( parser['status_type'], parser['status'] ), handshake_header_parser.parsed_length def compose(self): body_composer = ComposerBinary() body_composer.compose_numeric(self.status_type, 1) body_composer.compose_bytes(self.status, 3) header_bytes = self._compose_header(body_composer.composed_length) return header_bytes + body_composer.composed_bytes class TlsHandshakeServerHelloDone(TlsHandshakeMessage): @classmethod def get_handshake_type(cls): return TlsHandshakeType.SERVER_HELLO_DONE @classmethod def _parse(cls, parsable): handshake_header_parser = cls._parse_handshake_header(parsable) if handshake_header_parser['payload']: raise InvalidValue(bytes(handshake_header_parser['payload']), TlsHandshakeServerHelloDone, 'payload') return TlsHandshakeServerHelloDone(), handshake_header_parser.parsed_length def compose(self): return self._compose_header(0) class TlsECCurveType(enum.IntEnum): EXPLICIT_PRIME = 1 EXPLICIT_CHAR2 = 2 NAMED_CURVE = 3 @attr.s class TlsHandshakeServerKeyExchange(TlsHandshakeMessage): param_bytes = attr.ib(validator=attr.validators.instance_of(bytes)) @classmethod def get_handshake_type(cls): return TlsHandshakeType.SERVER_KEY_EXCHANGE @classmethod def _parse(cls, parsable): handshake_header_parser = cls._parse_handshake_header(parsable) return TlsHandshakeServerKeyExchange( bytes(handshake_header_parser['payload']) ), handshake_header_parser.parsed_length def compose(self): return self._compose_header(len(self.param_bytes)) + bytes(self.param_bytes) class TlsClientCertificateType(enum.IntEnum): RSA_SIGN = 0x01 DSS_SIGN = 0x02 RSA_FIXED_DH = 0x03 DSS_FIXED_DH = 0x04 ECDSA_SIGN = 0x40 RSA_FIXED_ECDH = 0x41 ECDSA_FIXED_ECDH = 0x42 GOST_SIGN256 = 0x43 GOST_SIGN512 = 0x44 class TlsClientCertificateTypeVector(Vector): @classmethod def get_param(cls): return VectorParamNumeric( item_size=1, min_byte_num=1, max_byte_num=2 ** 8 - 1, numeric_class=TlsClientCertificateType ) class TlsDistinguishedName(Opaque): @classmethod def get_param(cls): return OpaqueParam(min_byte_num=1, max_byte_num=2 ** 16 - 1) class TlsDistinguishedNameVector(VectorParsable): @classmethod def get_param(cls): return VectorParamParsable( item_class=TlsDistinguishedName, fallback_class=None, min_byte_num=0, max_byte_num=2 ** 16 - 1 ) @attr.s class TlsHandshakeCertificateRequest(TlsHandshakeMessage): certificate_types = attr.ib(converter=TlsClientCertificateTypeVector) certificate_authorities = attr.ib(converter=TlsDistinguishedNameVector) supported_signature_algorithms = attr.ib(default=None) def __attrs_post_init__(self): if self.supported_signature_algorithms is not None: self.supported_signature_algorithms = TlsSignatureAndHashAlgorithmVector( self.supported_signature_algorithms ) attr.validate(self) @classmethod def get_handshake_type(cls): return TlsHandshakeType.CERTIFICATE_REQUEST @classmethod def _parse(cls, parsable): handshake_header_parser = cls._parse_handshake_header(parsable) parser = ParserBinary(handshake_header_parser['payload']) parser.parse_parsable('certificate_types', TlsClientCertificateTypeVector) parser_remaining = ParserBinary(parser.unparsed) parser_remaining.parse_numeric('vector_length', 2) if parser_remaining['vector_length'] + 2 == parser.unparsed_length: supported_signature_algorithms = None else: parser.parse_parsable('supported_signature_algorithms', TlsSignatureAndHashAlgorithmVector) supported_signature_algorithms = parser['supported_signature_algorithms'] parser.parse_parsable('certificate_authorities', TlsDistinguishedNameVector) msg = TlsHandshakeCertificateRequest( certificate_types=parser['certificate_types'], certificate_authorities=parser['certificate_authorities'], supported_signature_algorithms=supported_signature_algorithms, ) return msg, handshake_header_parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.certificate_types) if self.supported_signature_algorithms is not None: payload_composer.compose_parsable(self.supported_signature_algorithms) payload_composer.compose_parsable(self.certificate_authorities) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes TLS_HANDSHAKE_HELLO_RETRY_REQUEST_RANDOM_BYTES = ( b'\xcf\x21\xad\x74\xe5\x9a\x61\x11\xbe\x1d\x8c\x02\x1e\x65\xb8\x91' + b'\xc2\xa2\x11\x16\x7a\xbb\x8c\x5e\x07\x9e\x09\xe2\xc8\xa8\x33\x9c' + b'' ) TLS_HANDSHAKE_HELLO_RETRY_REQUEST_RANDOM = TlsHandshakeHelloRandom.parse_exact_size( TLS_HANDSHAKE_HELLO_RETRY_REQUEST_RANDOM_BYTES ) @attr.s class TlsHandshakeHelloRetryRequest(TlsHandshakeHello): cipher_suite = attr.ib(default=None, validator=attr.validators.in_(TlsCipherSuite)) protocol_version = attr.ib( default=TlsProtocolVersion(TlsVersion.TLS1_3), validator=attr.validators.instance_of(TlsProtocolVersion), ) random_bytes = attr.ib( default=TLS_HANDSHAKE_HELLO_RETRY_REQUEST_RANDOM, validator=attr.validators.instance_of(TlsHandshakeHelloRandom), ) session_id = attr.ib( default=TlsSessionIdVector(random.randint(0, 255) for i in range(32)), converter=TlsSessionIdVector, validator=attr.validators.instance_of(TlsSessionIdVector), ) compression_method = attr.ib( default=TlsCompressionMethod.NULL, converter=TlsCompressionMethod, validator=attr.validators.in_(TlsCompressionMethod), ) extensions = attr.ib( default=TlsExtensionsServer([]), converter=TlsExtensionsServer, validator=attr.validators.instance_of(TlsExtensionsServer) ) @classmethod def get_handshake_type(cls): return TlsHandshakeType.HELLO_RETRY_REQUEST @classmethod def _parse(cls, parsable): handshake_header_parser = cls._parse_handshake_header(parsable) parser = cls._parse_hello_header(handshake_header_parser['payload']) parser.parse_parsable('cipher_suite', TlsCipherSuiteFactory) parser.parse_parsable('compression_method', TlsCompressionMethodFactory) compression_method = parser['compression_method'] session_id = parser['session_id'] extension_parser = cls._parse_extensions(handshake_header_parser, parser, TlsExtensionsServer) return TlsHandshakeHelloRetryRequest( protocol_version=parser['protocol_version'], random_bytes=parser['random'], session_id=session_id, compression_method=compression_method, cipher_suite=parser['cipher_suite'], extensions=parser['extensions'] if extension_parser else TlsExtensionsClient([]), ), handshake_header_parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_parsable(self.protocol_version) payload_composer.compose_parsable(self.random_bytes) payload_composer.compose_parsable(self.session_id) payload_composer.compose_numeric_enum_coded(self.cipher_suite) payload_composer.compose_numeric_enum_coded(self.compression_method) extension_bytes = self._compose_extensions(self.extensions) header_bytes = self._compose_header(payload_composer.composed_length + len(extension_bytes)) return header_bytes + payload_composer.composed_bytes + extension_bytes @attr.s class TlsHandshakeEncryptedExtensions(TlsHandshakeMessage): extension_data = attr.ib(validator=attr.validators.instance_of((bytes, bytearray))) @classmethod def get_handshake_type(cls): return TlsHandshakeType.ENCRYPTED_EXTENSIONS @classmethod def _parse(cls, parsable): parser = cls._parse_handshake_header(parsable) return TlsHandshakeEncryptedExtensions( parser['payload'], ), parser.parsed_length def compose(self): payload_composer = ComposerBinary() payload_composer.compose_raw(self.extension_data) header_bytes = self._compose_header(payload_composer.composed_length) return header_bytes + payload_composer.composed_bytes class SslMessageBase(ParsableBase): @classmethod def get_message_type(cls): return NotImplementedError() # pragma: no cover # pylint: disable=duplicate-code @classmethod @abc.abstractmethod def _parse(cls, parsable): raise NotImplementedError() # pylint: disable=duplicate-code @abc.abstractmethod def compose(self): raise NotImplementedError() class SslMessageType(enum.IntEnum): ERROR = 0x00 CLIENT_HELLO = 0x01 CLIENT_MASTER_KEY = 0x02 CLIENT_FINISHED = 0x03 SERVER_HELLO = 0x04 SERVER_VERIFY = 0x05 SERVER_FINISHED = 0x06 REQUEST_CERTIFICATE = 0x07 CLIENT_CERTIFICATE = 0x08 class SslCertificateType(enum.IntEnum): X509_CERTIFICATE = 0x01 class SslAuthenticationType(enum.IntEnum): MD5_WITH_RSA_ENCRYPTION = 0x01 class SslErrorType(enum.IntEnum): NO_CIPHER_ERROR = 0x0001 NO_CERTIFICATE_ERROR = 0x0002 BAD_CERTIFICATE_ERROR = 0x0003 UNSUPPORTED_CERTIFICATE_TYPE_ERROR = 0x0004 @attr.s class SslErrorMessage(SslMessageBase): error_type = attr.ib(validator=attr.validators.in_(SslErrorType)) @classmethod def get_message_type(cls): return SslMessageType.ERROR @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('error_type', 2, SslErrorType) return SslErrorMessage(parser['error_type']), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.error_type, 2) return composer.composed_bytes @attr.s class SslHandshakeClientHello(SslMessageBase): cipher_kinds = attr.ib(validator=attr.validators.deep_iterable(member_validator=attr.validators.in_(SslCipherKind))) session_id = attr.ib(validator=attr.validators.instance_of(bytes)) challenge = attr.ib(validator=attr.validators.instance_of(bytes)) @session_id.default def _default_session_id(self): # pylint: disable=no-self-use return b'' @challenge.default def _default_challenge(self): # pylint: disable=no-self-use return bytes(bytearray.fromhex(f'{random.getrandbits(128):16x}'.zfill(32))) @classmethod def get_message_type(cls): return SslMessageType.CLIENT_HELLO @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_parsable('version', TlsProtocolVersion) parser.parse_numeric('cipher_kinds_length', 2) parser.parse_numeric('session_id_length', 2) parser.parse_numeric('challenge_length', 2) parser.parse_parsable_array('cipher_kinds', parser['cipher_kinds_length'], SslCipherKindFactory) parser.parse_raw('session_id', parser['session_id_length']) parser.parse_raw('challenge', parser['challenge_length']) return SslHandshakeClientHello( cipher_kinds=parser['cipher_kinds'], session_id=bytes(parser['session_id']), challenge=bytes(parser['challenge']), ), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_parsable(TlsProtocolVersion(TlsVersion.SSL2)) composer.compose_numeric(len(self.cipher_kinds) * 3, 2) composer.compose_numeric(len(self.session_id), 2) composer.compose_numeric(len(self.challenge), 2) composer.compose_numeric_array_enum_coded(self.cipher_kinds) composer.compose_raw(self.session_id) composer.compose_raw(self.challenge) return composer.composed_bytes @attr.s class SslHandshakeServerHello(SslMessageBase): certificate = attr.ib(validator=attr.validators.instance_of(bytes)) cipher_kinds = attr.ib(validator=attr.validators.deep_iterable(member_validator=attr.validators.in_(SslCipherKind))) connection_id = attr.ib(validator=attr.validators.instance_of(bytes)) session_id_hit = attr.ib(default=False, validator=attr.validators.instance_of(bool)) @connection_id.default def _default_connection_id(self): # pylint: disable=no-self-use return b'' @classmethod def get_message_type(cls): return SslMessageType.SERVER_HELLO @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('session_id_hit', 1) parser.parse_numeric('certificate_type', 1, SslCertificateType) parser.parse_parsable('version', TlsProtocolVersion) parser.parse_numeric('certificate_length', 2) parser.parse_numeric('cipher_kinds_length', 2) parser.parse_numeric('connection_id_length', 2) parser.parse_raw('certificate', parser['certificate_length']) parser.parse_parsable_array('cipher_kinds', parser['cipher_kinds_length'], SslCipherKindFactory) parser.parse_raw('connection_id', parser['connection_id_length']) return SslHandshakeServerHello( certificate=bytes(parser['certificate']), cipher_kinds=parser['cipher_kinds'], connection_id=bytes(parser['connection_id']), session_id_hit=bool(parser['session_id_hit']), ), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(1 if self.session_id_hit else 0, 1) composer.compose_numeric(SslCertificateType.X509_CERTIFICATE, 1) composer.compose_parsable(TlsProtocolVersion(TlsVersion.SSL2)) composer.compose_numeric(len(self.certificate), 2) composer.compose_numeric(len(self.cipher_kinds) * 3, 2) composer.compose_numeric(len(self.connection_id), 2) composer.compose_raw(self.certificate) composer.compose_numeric_array_enum_coded(self.cipher_kinds) composer.compose_raw(self.connection_id) return composer.composed_bytes class TlsHandshakeMessageVariant(VariantParsable): _VARIANTS = collections.OrderedDict([ (TlsHandshakeType.CLIENT_HELLO, [TlsHandshakeClientHello, ]), (TlsHandshakeType.SERVER_HELLO, [TlsHandshakeServerHello, ]), (TlsHandshakeType.ENCRYPTED_EXTENSIONS, [TlsHandshakeEncryptedExtensions, ]), (TlsHandshakeType.CERTIFICATE, [TlsHandshakeCertificate, TlsHandshakeServerCertificate, ]), (TlsHandshakeType.SERVER_KEY_EXCHANGE, [TlsHandshakeServerKeyExchange, ]), (TlsHandshakeType.CERTIFICATE_REQUEST, [TlsHandshakeCertificateRequest, ]), (TlsHandshakeType.CERTIFICATE_STATUS, [TlsHandshakeCertificateStatus, ]), (TlsHandshakeType.SERVER_HELLO_DONE, [TlsHandshakeServerHelloDone, ]), (TlsHandshakeType.HELLO_RETRY_REQUEST, [TlsHandshakeHelloRetryRequest, ]), ]) @classmethod def _get_variants(cls): return cls._VARIANTS class TlsSubprotocolMessageParser(SubprotocolParser): _SUBPROTOCOL_PARSERS = { TlsContentType.CHANGE_CIPHER_SPEC: TlsChangeCipherSpecMessage, TlsContentType.ALERT: TlsAlertMessage, TlsContentType.HANDSHAKE: TlsHandshakeMessageVariant, } @classmethod def _get_subprotocol_parsers(cls): return cls._SUBPROTOCOL_PARSERS class SslSubprotocolMessageParser(SubprotocolParser): _SUBPROTOCOL_PARSERS = { SslMessageType.ERROR: SslErrorMessage, SslMessageType.CLIENT_HELLO: SslHandshakeClientHello, SslMessageType.SERVER_HELLO: SslHandshakeServerHello, } @classmethod def _get_subprotocol_parsers(cls): return cls._SUBPROTOCOL_PARSERS cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/cryptoparser/tls/version.py000066400000000000000000000055701524413560000274570ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import functools import attr from cryptodatahub.common.grade import Grade, GradeableSimple from cryptodatahub.tls.version import TlsVersion from cryptoparser.common.base import TwoByteEnumParsable, ProtocolVersionBase from cryptoparser.common.parse import ParserBinary, ComposerBinary class TlsVersionFactory(TwoByteEnumParsable): @classmethod def get_enum_class(cls): return TlsVersion @abc.abstractmethod def compose(self): raise NotImplementedError() @attr.s(order=False, eq=False, hash=True) @functools.total_ordering class TlsProtocolVersion(ProtocolVersionBase, GradeableSimple): version = attr.ib(validator=attr.validators.instance_of(TlsVersion)) @property def grade(self): if self.version in (TlsVersion.TLS1_3, TlsVersion.TLS1_2): return Grade.SECURE if self.version in (TlsVersion.TLS1, TlsVersion.TLS1_1) or self.is_draft or self.is_google_experimental: return Grade.DEPRECATED if self.version in (TlsVersion.SSL2, TlsVersion.SSL3): return Grade.INSECURE raise NotImplementedError(self.version) @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_parsable('version', TlsVersionFactory) return cls(**parser), parser.parsed_length def compose(self): composer = ComposerBinary() composer.compose_numeric(self.major, 1) composer.compose_numeric(self.minor, 1) return composer.composed_bytes @property def major(self): return (self.version.value.code & 0xff00) >> 8 @property def minor(self): return self.version.value.code & 0x00ff @property def is_draft(self): return self.major == 0x7f @property def is_google_experimental(self): return self.major == 0x7e def __eq__(self, other): return self.version.value.code == other.version.value.code def __lt__(self, other): if self.major == other.major: return self.minor < other.minor if self.is_draft: return other.version == TlsVersion.TLS1_3 if other.is_draft: return self.version != TlsVersion.TLS1_3 return self.major < other.major @property def identifier(self): return self.version.name.lower() def __str__(self): if self.is_draft: return f'TLS 1.3 Draft {self.minor}' if self.is_google_experimental: return f'TLS 1.3 Google Experiment {self.minor}' if self.version == TlsVersion.SSL3: return 'SSL 3.0' if self.version == TlsVersion.SSL2: return 'SSL 2.0' return f'TLS 1.{self.minor - 1}' def _as_markdown(self, level): return self._markdown_result(str(self), level) def _asdict(self): return self.identifier cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/debian/000077500000000000000000000000001524413560000232745ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/debian/changelog000066400000000000000000000064301524413560000251510ustar00rootroot00000000000000python-cryptoparser (1.6.0) unstable; urgency=low * add trust anchors extension related messages (#104) * add server padding extension related messages (#104) * add application-layer protocol settings extension related messages with the old code point (#104) * bound the outer extension payload of the encrypted client hello by the extension length (#104) * preserve the transform number when parsing the IKEv1 transform payload (#105) * follow the replacement of the named group with the key parameter (#106) * make the parameter classes of the invalid extension types immutable (#104) -- Szilárd Pfeiffer Tue, 25 Aug 2026 00:00:00 +0200 python-cryptoparser (1.5.0) unstable; urgency=low * add IKEv1 and IKEv2 certificate payload parsing (#103) * add IKEv1 identification payload parsing (#103) * add IKEv1 signature payload parsing (#103) * add IKEv2 identification payload parsing (#103) * add IKEv2 authentication payload parsing (#103) * add IKEv2 extensible authentication protocol payload parsing (#103) * add IKEv2 encrypted and authenticated payload parsing (#103) * keep IKEv1 and IKEv2 payloads of unknown type unparsed instead of rejecting the message (#103) * parse the IKEv1 and IKEv2 identification payload data according to the identification type (#103) * parse the digital signature envelope of the IKEv2 authentication payload (#103) * split the ISAKMP message parsing and composing into header and payload chain steps (#103) * use typed values in the IKEv2 signature hash algorithms notify payload (#103) * add getter for the distinguished name of the IKEv1 certificate request payload (#103) * add getter for the certification authority hashes of the IKEv2 certificate request payload (#103) -- Szilárd Pfeiffer Fri, 31 Jul 2026 00:00:00 +0200 python-cryptoparser (1.4.0) unstable; urgency=low * add getter for multiple payloads with the same type (#93) * add IKEv2 NAT detection source IP and destination IP notify payload parsing (#93) * add IKEv2 set window size notify payload parsing (#93) * add IKEv2 use transport mode notify payload parsing (#93) * add IKEv2 HTTP certificate lookup supported notify payload parsing (#93) * add IKEv2 signature hash algorithms notify payload parsing (#93) * add IKEv2 intermediate exchange supported notify payload parsing (#93) * add IKEv2 use PPK notify payload parsing (#93) * add IKEv2 redirect supported notify payload parsing (#93) * add IKEv2 fragmentation supported notify payload parsing (#93) * add childless IKEv2 supported notify payload parsing (#93) -- Szilárd Pfeiffer Fri, 17 Jul 2026 15:37:42 +0200 python-cryptoparser (1.3.0) unstable; urgency=low * add Debian and RPM packaging (#102) * add JA4 tag generation for the client hello message (#100) * add JA4X tag generation for X.509 certificates (#101) * add certificate-related messages for protocol version 1.3 (#94) * make IKEv2 transform key length optional for fixed-key ciphers (#99) -- Szilárd Pfeiffer Mon, 15 Jun 2026 21:14:01 +0200 python-cryptoparser (1.2.1) unstable; urgency=low * Initial Debian packaging. -- Szilárd Pfeiffer Sun, 14 Jun 2026 00:00:00 +0200 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/debian/control000066400000000000000000000021571524413560000247040ustar00rootroot00000000000000Source: python-cryptoparser Section: python Priority: optional Maintainer: Szilárd Pfeiffer Build-Depends: debhelper-compat (= 12), dh-python, python3-all, python3-setuptools, python3-asn1crypto, python3-attr, python3-cryptodatahub (>= 1.2.1), python3-pyfakefs , python3-urllib3 Standards-Version: 4.7.2 Rules-Requires-Root: no Homepage: https://gitlab.com/coroner/cryptoparser Vcs-Browser: https://gitlab.com/coroner/cryptoparser Vcs-Git: https://gitlab.com/coroner/cryptoparser.git Package: python3-cryptoparser Architecture: all Depends: ${python3:Depends}, ${misc:Depends}, python3-asn1crypto, python3-attr, python3-cryptodatahub (>= 1.2.1), python3-urllib3 Description: Analysis-oriented security protocol parser and generator CryptoParser is a Python library for parsing and generating security protocol messages. It supports TLS, SSH, DNS, LDAP, RDP, and IKE protocols with a focus on cryptographic algorithm analysis. cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/debian/copyright000066400000000000000000000010511524413560000252240ustar00rootroot00000000000000Format: https://www.debian.org/doc/packaging-manuals/copyright-format/1.0/ Upstream-Name: CryptoParser Upstream-Contact: Szilárd Pfeiffer Source: https://gitlab.com/coroner/cryptoparser Files: * Copyright: 2018-2026 Szilárd Pfeiffer License: MPL-2.0 Files: debian/* Copyright: 2026 Szilárd Pfeiffer License: MPL-2.0 License: MPL-2.0 On Debian systems, the complete text of the Mozilla Public License 2.0 can be found in '/usr/share/common-licenses/MPL-2.0'. cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/debian/rules000077500000000000000000000002121524413560000243470ustar00rootroot00000000000000#!/usr/bin/make -f export SETUPTOOLS_SCM_PRETEND_VERSION := 1.6.0 %: dh $@ --with python3 --buildsystem=pybuild override_dh_auto_test: cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/debian/source/000077500000000000000000000000001524413560000245745ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/debian/source/format000066400000000000000000000000151524413560000260030ustar00rootroot000000000000003.0 (native) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/debian/watch000066400000000000000000000001641524413560000243260ustar00rootroot00000000000000version=4 opts=uversionmangle=s/(rc|a|b|c)/~$1/ \ https://pypi.debian.net/CryptoParser/CryptoParser-(.+)\.tar\.gz cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/docs/000077500000000000000000000000001524413560000230025ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/docs/.gitignore000066400000000000000000000000141524413560000247650ustar00rootroot00000000000000_build html cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/docs/changelog.rst000066400000000000000000000000361524413560000254620ustar00rootroot00000000000000.. include:: ../CHANGELOG.rst cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/docs/conf.py000066400000000000000000000037561524413560000243140ustar00rootroot00000000000000#!/usr/bin/env python # SPDX-License-Identifier: MPL-2.0 # -*- coding: utf-8 -*- # pylint: disable=invalid-name import datetime import os import pathlib import sys import urllib sys.path.insert(0, os.path.abspath('..')) from cryptoparser.__setup__ import ( # noqa: E402, pylint: disable=wrong-import-position __author__, __description__, __title__, __version__, ) extensions = [ 'myst_parser', 'sphinx_sitemap', ] templates_path = ['_templates'] source_suffix = '.rst' master_doc = 'index' project = __title__ copyright = f'{datetime.datetime.now().year}, {__author__}' # pylint: disable=redefined-builtin if 'READTHEDOCS' in os.environ: version = release = os.environ['READTHEDOCS_VERSION'] html_baseurl = os.environ['READTHEDOCS_CANONICAL_URL'] _baseurl_parsed = urllib.parse.urlparse(html_baseurl) _baseurl = urllib.parse.ParseResult(_baseurl_parsed.scheme, _baseurl_parsed.netloc, '/', '', '', '').geturl() sitemap_url_scheme = "{link}" _robots_txt_lines = [ 'User-agent: *', '', 'Disallow: # Allow everything', '', ] for lang in ('en',): for tag in ('latest', 'stable'): _robots_txt_lines.append(f'Sitemap: {_baseurl}{lang}/{tag}/sitemap.xml') _html_extra_dir_name = 'readthedocs' _html_extra_path = pathlib.Path(_html_extra_dir_name) _html_extra_path.mkdir(exist_ok=True) with open(_html_extra_path / 'robots.txt', 'w+', encoding='ascii') as _robots_txt_file: _robots_txt_file.write(os.linesep.join(_robots_txt_lines)) html_extra_path = [ _html_extra_dir_name ] else: version = release = __version__ exclude_patterns = ['_build'] html_title = __title__ + ' — ' + __description__ html_theme = 'alabaster' html_sidebars = { '**': [ 'about.html', 'navigation.html', 'relations.html', 'searchbox.html', 'donate.html', ] } html_theme_options = { 'description': __description__, 'fixed_sidebar': True, } cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/docs/development.rst000066400000000000000000000003171524413560000260570ustar00rootroot00000000000000----------- Development ----------- If you want to setup a development environment, you are in need of `uv `__. .. code:: shell $ cd cryptoparser $ uv sync --extra tests cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/docs/features.rst000066400000000000000000000147401524413560000253600ustar00rootroot00000000000000-------- Features -------- Supported Protocols =================== Internet Key Exchange (IKE) --------------------------- - `ISAKMP `__ - `IKEv1 `__ - `IKEv2 `__ Secure Shell (SSH) ------------------ - `SSH 2.0 `__ Secure Socket Layer (SSL) ------------------------- - `SSL 2.0 `__ - `SSL 3.0 `__ Transport Layer Security (TLS) ------------------------------ - `TLS 1.0 `__ - `TLS 1.1 `__ - `TLS 1.2 `__ - `TLS 1.3 `__ Domain Name System (DNS) ------------------------ - `DNSSEC `__ (Domain Name System Security Extensions) Protocol Specific Features ========================== Internet Key Exchange (IKE) --------------------------- - protocol versions - notify payloads of IKEv2 protocol extensions - certificate payloads - certificate request payloads - identification payloads (per-ID-type parsed data) - signature payloads - authentication payloads (digital signature envelope) - extensible authentication protocol (EAP) payloads - encrypted and authenticated (SK) payloads Hypertext Transfer Protocol (HTTP) ---------------------------------- 1. supports header wire format parsing 2. supports detailed parsing of generic headers (`Content-Type `__, `NEL `__ (Network Error Logging), `Server `__, `Set-Cookie `__) 3. supports detailed parsing of caching headers (`Age `__, `Cache-Control `__, `Date `__, `ETag `__, `Expires `__, `Last-Modified `__, `Pragma `__) 4. supports detailed parsing of security headers (`Content Security Policy `__ (CSP), `Content-Security-Policy-Report-Only `__, `Expect-CT `__, `Expect-Staple `__, `HTTP Public Key Pinning `__ (HPKP), `Referrer-Policy `__, `Strict-Transport-Security `__, `X-Content-Type-Options `__, `X-Frame-Options `__, `X-XSS-Protection `__) Transport Layer Security (TLS) ------------------------------ Only features that cannot be or difficultly implemented by some of the most popular SSL/TLS implementations (eg: `GnuTls `__, `LibreSSL `__, `OpenSSL `__, `wolfSSL `__, ...) are listed. - generic 1. supports `Generate Random Extensions And Sustain Extensibility `__ (GREASE) values for - protocol version - extension type - ciphers suite - signature algorithms - named group 2. supports easy `JA3 fingerprint `__ generation - protocol versions 1. support not only the final, but also draft versions - cipher suites 1. supports each cipher suites discussed on `ciphersuite.info `__ 2. supports `GOST `__ (national standards of the Russian Federation and CIS countries) cipher suites 3. supports `ShangMi (SM) `__ (national standards of China) cipher suites - application layer - supports TLS handshake-related `MySQL `__ messages - supports TLS handshake-related `OpenVPN `__ messages - supports TLS handshake-related `PostgreSQL `__ messages - supports TLS handshake-related `RDP `__ messages Secure Shell (SSH) ------------------ - cipher suites 1. identifies as much encryption algorithms as possible (more than 200, compared to 70+ currently supported by OpenSSH) 2. supports `HASSH fingerprint `__ calculation - public keys 1. supports host keys, certificates (both ``V00`` and ``V01``), X.509 certificates and chains Domain Name System (DNS) ------------------------ - e-mail authentication, reporting - `Domain-based Message Authentication, Reporting, and Conformance `__ (DMARC) - `Sender Policy Framework `__ (SPF) - `SMTP MTA Strict Transport Security `__ (MTA-STS) - `SMTP TLS Reporting `__ (TLSRPT) - DNSSEC (Domain Name System Security Extensions) - `DNSKEY `__ - `DS `__ - `RRSIG `__ - `SSHFP `__ (SSH host key fingerprints) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/docs/index.rst000066400000000000000000000044161524413560000246500ustar00rootroot00000000000000.. meta:: :google-site-verification: 2AAgZNptPaMHxDeXJegA8i8aW1jURVBpQseacnHQr8Q .. meta:: :description: An analysis oriented security protocol parser and generator .. meta:: :keywords: cryptoparser,cryptolyzer,cryptography,cryptographic algorithms,tls handshake,ssl handshake,ssh handshake, starttls,opportunistic tls,ssh host keys,ssh host certificates,http caching headers,http security header, dnssec records,email authentication .. meta:: :author: Szilárd Pfeiffer ======= Summary ======= .. include:: ../README.md :parser: myst_parser.sphinx_ :start-after: :end-before: .. include:: ../README.md :parser: myst_parser.sphinx_ :start-after: :end-before: Why CryptoParser? ================= .. include:: ../README.md :parser: myst_parser.sphinx_ :start-after: :end-before: ======= Details ======= The main purpose of creating this library is the fact, that cryptography protocol analysis differs in many aspects from establishing a connection using a cryptographic protocol. Analysis is mostly testing where we trigger special and corner cases of the protocol and we also trying to establish connection with hardly supported, experimental, obsoleted or even deprecated mechanisms or algorithms which are may or may not supported by the latest or any version of an implementation of the cryptographic protocol. On the one hand it is neither a comprehensive nor a secure implementation of any cryptographic protocol. On the one hand library implements only the absolutely necessary parts of the protocol. On the other it contains completely insecure algorithms and mechanisms. It is not designed and contraindicated to use this library establishing secure connections. If you are searching for cryptographic protocol implementation, there are several existing wrappers and native implementations for Python (eg: M2Crypto, pyOpenSSL, Paramiko, ...). .. toctree:: :maxdepth: 3 features development ======= History ======= .. toctree:: :maxdepth: 2 changelog cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/llms.txt000066400000000000000000000140411524413560000235620ustar00rootroot00000000000000# CryptoParser An analysis-oriented cryptographic protocol parser and generator. Parses and composes TLS, SSL, SSH, DNSSEC, IKE, and HTTP security headers from raw bytes to typed Python objects. Not a comprehensive or secure implementation — designed for testing and analysis purposes. ## Files ### Project Root - `pyproject.toml` — Build system, dependencies, metadata (setuptools + setuptools-scm) - `setup.py` — Delegates to pyproject.toml - `README.md` — Project overview, install instructions, license - `CHANGELOG.rst` — Versioned changelog - `CONTRIBUTING.rst` — Contribution guidelines - `LICENSE.txt` — MPL-2.0 ### Common Infrastructure (`cryptoparser/common/`) - `parse.py` — Core parsing/composing engine: `ParserBinary`, `ParserText`, `ComposerBinary`, `ComposerText`, `ParsableBase`, `ByteOrder` - `base.py` — Reusable parseable types: `Serializable`, `ProtocolVersionBase`, `VectorParsable`, `VariantParsable`, `Opaque`, `ListParsable`, enum parsers (1/2/3-byte) - `field.py` — Structured text field parsing: `FieldParsableBase`, `NameValuePair`, quoted strings, URLs, datetimes, percentages, key=value components - `exception.py` — Parser exceptions: `InvalidDataLength`, `NotEnoughData`, `TooMuchData`, `InvalidType` - `classes.py` — Generic domain classes: `LanguageTag` (RFC 5646) - `utils.py` — Utility: `get_leaf_classes()`, `bytes_to_hex_string()` - `x509.py` — X.509 Certificate Transparency: `SignedCertificateTimestamp`, `SignedCertificateTimestampList` ### TLS/SSL (`cryptoparser/tls/`) - `record.py` — TLS record layer (5-byte header: content type, version, fragment length), dispatches to subprotocol parsers - `subprotocol.py` — TLS/SSL handshake messages: `ClientHello`, `ServerHello`, `Certificate`, `ServerKeyExchange`, `Finished`, `Alert`, `ApplicationData`, `ChangeCipherSpec` - `extension.py` — TLS extensions: SNI, ALPN, supported groups, key share, signature algorithms, PSK, GREASE, and ~40 more - `version.py` — TLS/SSL version handling: `TlsProtocolVersion`, `TlsVersionFactory` - `ciphersuite.py` — Ciphersuite factories: `TlsCipherSuiteFactory`, `SslCipherKindFactory` - `algorithm.py` — Algorithm factories: named curves, signature algorithms, EC point formats - `grease.py` — GREASE (Generate Random Extensions And Sustain Extensibility) handling - `ldap.py` — LDAP StartTLS protocol parser - `mysql.py` — MySQL protocol parser (handshake, SSL request) - `openvpn.py` — OpenVPN control channel packet parser - `postgresql.py` — PostgreSQL StartTLS protocol parser - `rdp.py` — RDP protocol parser (TPKT/X.224 encapsulation to extract TLS) ### SSH (`cryptoparser/ssh/`) - `record.py` — SSH record layer (binary packet framing: length, padding, message code) - `subprotocol.py` — SSH transport messages: `SshMessageKexInit`, `SshMessageKexDHInit`, `SshMessageKexDHReply`, `SshMessageNewKeys` - `key.py` — SSH public/private key parsing: RSA, DSA, ECDSA, Ed25519, Ed448, X25519, X448; OpenSSH private key format, known_hosts, authorized_keys - `version.py` — SSH version banner parsing: `SshVersion`, `SshProtocolVersion`, `SshVersionBanner`, `SshComment` ### DNSSEC (`cryptoparser/dnsrec/`) - `record.py` — DNS resource records: DNSKEY, DS, RRSIG, NSEC, NSEC3, NSEC3PARAM, CDNSKEY, CDS, TLSA, SSHFP - `txt.py` — DNS TXT record content: SPF, DKIM, DMARC, MTA-STS, TLSRPT, SMIMEA, CAA ### HTTP Headers (`cryptoparser/httpx/`) - `header.py` — HTTP security headers: HSTS, CSP, HPKP, Expect-CT, X-Frame-Options, X-Content-Type-Options, Referrer-Policy, Feature-Policy, Permissions-Policy, Cross-Origin headers, Set-Cookie, Cache-Control, WWW-Authenticate - `parse.py` — HTTP-specific field value components - `version.py` — `HttpVersion` enum (HTTP/1.0, HTTP/1.1) ### IKE (`cryptoparser/ike/`) - `isakmp.py` — ISAKMP header parser (cookies, version, exchange type, flags) - `ikev1.py` — IKEv1 payloads: SA, Proposal, Transform, KE, ID, Cert, Auth, Nonce, Notify, Delete, VendorID - `ikev2.py` — IKEv2 payloads: SA, Proposal, Transform, KE, IDi/IDr, Cert, Auth, Nonce, Notify, TS initiator/responder, Encrypted - `common.py` — Common IKE data attributes (AF flag, transform types) - `version.py` — ISAKMP version handling: `IsakmpVersion`, `IsakmpProtocolVersion` ### Tests (`test/`) - `test/common/` — Common infrastructure tests: base classes, parsing engine, exceptions, field parsing, X.509 - `test/tls/` — TLS/SSL tests: records, handshakes, extensions, cipher suites, alerts, version, wrapped protocols - `test/ssh/` — SSH tests: key parsing, records, subprotocols, version banners, cipher suites - `test/dnsrec/` — DNS tests: DNSSEC binary records, TXT structured text - `test/httpx/` — HTTP header tests: security headers, HTTP version - `test/ike/` — IKE tests: IKEv1/IKEv2 payloads, SA/transform negotiation, ISAKMP header ## Architecture The project uses a layered parser/composer design: 1. **Core** (`common/parse.py`) — `ParsableBase` abstract interface with `_parse()` returning `(object, bytes_consumed)`. `ParserBinary`/`ParserText` for input, `ComposerBinary`/`ComposerText` for output. 2. **Reusable types** (`common/base.py`) — Enum parsers (1/2/3-byte), vectors, variants, opaque blobs, lists, strings, protocol versions built on the core engine. 3. **Structured text** (`common/field.py`) — Name=value field parsing with typed components (strings, URLs, datetimes, quoted strings). 4. **Protocol modules** — Each protocol builds on common types to parse protocol-specific messages from raw bytes to typed Python objects. 5. **Enums/constants** — Defined in the `cryptodatahub` external library (submodule at `submodules/cryptodatahub/`). ## Key Patterns - All parsable classes implement `_parse(cls, parseable, bytes_consumed=None)` classmethod - All composable classes implement `compose(self, composable)` method - Protocol version classes implement `__lt__`, `__le__`, etc. for gradeable security assessment - Factory classes wrap enum-based parsers (one-byte, two-byte) using `OneByteEnumParsable`/`TwoByteEnumParsable` - `@attr.s` (attrs library) is used heavily for data classes with validators cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/pyproject.toml000066400000000000000000000047031524413560000247720ustar00rootroot00000000000000[build-system] requires = ['setuptools', 'setuptools-scm'] build-backend = 'setuptools.build_meta' [project] name = 'CryptoParser' version = '1.6.0' description = 'An analysis oriented security protocol parser and generator' authors = [ {name = 'Szilárd Pfeiffer', email = 'coroner@pfeifferszilard.hu'} ] maintainers = [ {name = 'Szilárd Pfeiffer', email = 'coroner@pfeifferszilard.hu'} ] classifiers=[ 'Development Status :: 5 - Production/Stable', 'Environment :: Console', 'Intended Audience :: Information Technology', 'Intended Audience :: Science/Research', 'Intended Audience :: System Administrators', 'Natural Language :: English', 'Operating System :: MacOS', 'Operating System :: Microsoft :: Windows', 'Operating System :: POSIX', 'Programming Language :: Python :: 3.9', 'Programming Language :: Python :: 3.10', 'Programming Language :: Python :: 3.11', 'Programming Language :: Python :: 3.12', 'Programming Language :: Python :: 3.13', 'Programming Language :: Python :: 3.14', 'Programming Language :: Python :: Implementation :: CPython', 'Programming Language :: Python :: Implementation :: PyPy', 'Programming Language :: Python', 'Topic :: Internet', 'Topic :: Security', 'Topic :: Security :: Cryptography', 'Topic :: Software Development :: Libraries :: Python Modules', 'Topic :: Software Development :: Testing :: Traffic Generation', 'Topic :: Software Development :: Testing', ] keywords=['ssl', 'tls', 'gost', 'ja3', 'ldap', 'rdp', 'ssh', 'hsts', 'dns', 'ike'] readme = {file = 'README.md', content-type = 'text/markdown'} license = {text = 'MPL-2.0'} requires-python = '>=3.9' dependencies = [ 'asn1crypto', 'attrs', 'cryptodatahub==1.6.0', 'urllib3', ] [project.optional-dependencies] tests = [ 'pyfakefs', 'coverage', ] docs = [ 'myst-parser', 'sphinx', 'sphinx-sitemap', 'docutils', ] [project.urls] Homepage = 'https://gitlab.com/coroner/cryptoparser' Changelog = 'https://cryptoparser.readthedocs.io/en/latest/changelog' Documentation = 'https://cryptoparser.readthedocs.io/en/latest/' Issues = 'https://gitlab.com/coroner/cryptoparser/-/issues' Source = 'https://gitlab.com/coroner/cryptoparser' [tool.variables] technical_name = 'cryptoparser' [tool.setuptools.packages.find] exclude = ['submodules'] [tool.ruff] line-length = 120 target-version = 'py39' [tool.ruff.lint] select = ['E', 'F', 'W', 'UP', 'B'] cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/setup.py000077500000000000000000000001401524413560000235620ustar00rootroot00000000000000#!/usr/bin/env python # SPDX-License-Identifier: MPL-2.0 import setuptools setuptools.setup() cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/submodules/000077500000000000000000000000001524413560000242345ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/submodules/cryptodatahub/000077500000000000000000000000001524413560000271055ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/000077500000000000000000000000001524413560000230315ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/__init__.py000066400000000000000000000000431524413560000251370ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/000077500000000000000000000000001524413560000243215ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/__init__.py000066400000000000000000000000431524413560000264270ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/certs/000077500000000000000000000000001524413560000254415ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/certs/ecc256.badssl.com.pem000066400000000000000000000030461524413560000311620ustar00rootroot00000000000000-----BEGIN CERTIFICATE----- MIIEXjCCA0agAwIBAgISAxbpbfOZVQALObIfcle/CC1lMA0GCSqGSIb3DQEBCwUA MDIxCzAJBgNVBAYTAlVTMRYwFAYDVQQKEw1MZXQncyBFbmNyeXB0MQswCQYDVQQD EwJSMzAeFw0yMzA0MjMyMjU4NTRaFw0yMzA3MjIyMjU4NTNaMBcxFTATBgNVBAMM DCouYmFkc3NsLmNvbTBZMBMGByqGSM49AgEGCCqGSM49AwEHA0IABCahOdqN0HmZ gKIGPLFq5tFcQKRzLZ2Y8ZH5wBbOjQxxs5BlqRPli7OEWQxWNct+iLTCYRvqIB88 ptS+pntd45WjggJSMIICTjAOBgNVHQ8BAf8EBAMCB4AwHQYDVR0lBBYwFAYIKwYB BQUHAwEGCCsGAQUFBwMCMAwGA1UdEwEB/wQCMAAwHQYDVR0OBBYEFC4cTT4KpISH Qm1+Bl8S6C7mCMUdMB8GA1UdIwQYMBaAFBQusxe3WFbLrlAJQOYfr52LFMLGMFUG CCsGAQUFBwEBBEkwRzAhBggrBgEFBQcwAYYVaHR0cDovL3IzLm8ubGVuY3Iub3Jn MCIGCCsGAQUFBzAChhZodHRwOi8vcjMuaS5sZW5jci5vcmcvMCMGA1UdEQQcMBqC DCouYmFkc3NsLmNvbYIKYmFkc3NsLmNvbTBMBgNVHSAERTBDMAgGBmeBDAECATA3 BgsrBgEEAYLfEwEBATAoMCYGCCsGAQUFBwIBFhpodHRwOi8vY3BzLmxldHNlbmNy eXB0Lm9yZzCCAQMGCisGAQQB1nkCBAIEgfQEgfEA7wB2ALc++yTfnE26dfI5xbpY 9Gxd/ELPep81xJ4dCYEl7bSZAAABh7COYsMAAAQDAEcwRQIhAJ1reE5Oo7YRwLGg 0QI9cm8BX3igarNq2bssyh0Ci4yeAiAt2aCEvJkrgtbWLXSjh7hMW2WuRN1C+opV onLe1LdM0gB1AHoyjFTYty22IOo44FIe6YQWcDIThU070ivBOlejUutSAAABh7CO YtYAAAQDAEYwRAIgLyfHa/MB7YaD07krLoEzQAfGS1/C4bjv757A8SBRce8CIHVX Hfz0PGfAabtxdlbbpT3F4S4TU1HPD0yprlPhFKrIMA0GCSqGSIb3DQEBCwUAA4IB AQCuunVuLKZQ6O/ZQpiP4WqqDkMXaBcFLM3o5EMU6Otn2Qzn+oFIe27EMAVelvUt hy2xEyec1Gd5d41+ik9gytHP7SZlp5dMXdZdIfzgGd5zVmRJJ7tyz3PaxxgQjctj MJFkAAYhR1Aa2iimbu66nqE118YZyALV3sQ778b1IRSfvZ0QQ1sLxkDFURpUeo1f log+EnCupj8N0DYE7q4qMB40m1nFH4BWpnfLzkNy3XJiYcZ+JJ1chxUnqy/105Dw l7tL+Xjbiuf2oXTseRbUcRc9uOtzdaYR3VtF1paIrESgQwPI4rj3tmXi6WuMag6h 7UZLQbPrxY7OYEt+uXyzZbZ5 -----END CERTIFICATE----- rsa8192.badssl.com.pem000066400000000000000000000062271524413560000312310ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/certs-----BEGIN CERTIFICATE----- MIIJIDCCCAigAwIBAgIQDBboHFj/nFQ7X6MbvDiVMDANBgkqhkiG9w0BAQsFADBg MQswCQYDVQQGEwJVUzEVMBMGA1UEChMMRGlnaUNlcnQgSW5jMRkwFwYDVQQLExB3 d3cuZGlnaWNlcnQuY29tMR8wHQYDVQQDExZSYXBpZFNTTCBUTFMgUlNBIENBIEcx MB4XDTIzMDMyODAwMDAwMFoXDTI0MDMyNzIzNTk1OVowFzEVMBMGA1UEAwwMKi5i YWRzc2wuY29tMIIEIjANBgkqhkiG9w0BAQEFAAOCBA8AMIIECgKCBAEA2gk0w0XT FdWSvMYNgxJ1dyiSe5THwjxJfe35/Rqt3BSLVCTHBTgZKiBwXYMK11gsdrpBqWD3 Mp081yIjU5BC1uuUCV6su61TmWg5CoIQbM+imx8zXazBdu7xqvKXSK41wtb9HZhj 8cvQePK4WjnGePfMB/NoMww2ogny3Q9gCpFTYiIXXS/sjZ06fCRsfK32RsbOBRyT DfJWMu+vz+lwBbFkcKXTYUGnEyOQaS+Tu8RU564oYbYacMr7Ijvgz+Fsb71KmdEk Y809LnM7xZ4wfvlO7jXvMomhWlZBGL8xHPz7XgGlUXza77wZhsfepxeWzS/HbT3k 2qBtzpq4hTPnqd9yrZCpO/J+AtClt1tYw+orErVZ60+nWG4ZLpTqeY5wjagglPp6 yyA1jvhyz0U66ult9KjgZEGCC+QJXZPKwnKkqXnGWzxvqthUQWqRIoMnr6+LCP1e U3jvy0pmtPPLsY6o8xEfDPycY8fI1lKkuP0teo8jSVY4eVCl0X3hnm6nTrWP8+c8 I+Qc3j7V0RFJqWcPD7ouB7V6P5fSHLdmNWzamJVzrva494RCctmJ15rtfQ+/nwd0 pqNyL1GC3+OJ6Go2B1wjY1z/jEhrDcptLFJo9+nX6u331KA/3q6i+X52Od+kDmwl xeM2wkNgo4da3yu6TAnONnviopSIDjCp6/j5FdHTf1Xf1oYCVRJWje2Ur8r0l31u JV6q+uv/1+CitnPe9cB16KqEDQlm4igzwb6ikP/QJIU9XUu4p9uTY1fKTSLNku78 VEVNI72W9eY5uCz4kcw02CrbTKOVnEILNZKGpsKs8FeZmz7EXiKBiiXIB3i4JLvd SrU13RnXptHRg3imiKRObZ267gs3ho+un6TaoC7fE6RieBPZ+OaCbnKDMq1RnSPq iC/fizBS9b5KCnY8JcCIizfcSGEoGCVliNkXJhlU5qKaS4XCCNhVKDQ3M5vWCWCn dMEk563oIySanqIYWceYnD3b9hgEj6A88EKZnY50A8rn2jINEIU4ds+rQEAJjxA/ Zv1xaaTpiERGRj9Ln+9veppOHgap7KZpz8hELQ8zCLf3o3yWUCIBn3f0eLlkOwT+ rYVvVeg88F/5LmCPGegnn5D+IcjHuB2SGErvZKiCAE0CYF/xkONg3mo51Jli3cYB BV8wbU4C/sl4lTVeamUfnajY5ofPBEkRsDIhIqtGvxDk6Yr84rVvngDXTbP05eU3 9sSbActXbX2iHtrSzjGCdKV6UtAohuZtLTkut4e/8KJIZSuvqvf9dk0Z9jRnwmJv EmHQq5d63kwp/PtmFg6FpB7h1+Y9fXANoIuaC0u342VDOGoWCH5QUxKamhedy1Uz zh5PeHkZJ79jVQIDAQABo4IDHTCCAxkwHwYDVR0jBBgwFoAUDNtsgkkPSmcKuBTu esRIUojrVjgwHQYDVR0OBBYEFJzZKG+o1iyDszBUsyY2Ye6iJqhJMCMGA1UdEQQc MBqCDCouYmFkc3NsLmNvbYIKYmFkc3NsLmNvbTAOBgNVHQ8BAf8EBAMCBaAwHQYD VR0lBBYwFAYIKwYBBQUHAwEGCCsGAQUFBwMCMD8GA1UdHwQ4MDYwNKAyoDCGLmh0 dHA6Ly9jZHAucmFwaWRzc2wuY29tL1JhcGlkU1NMVExTUlNBQ0FHMS5jcmwwPgYD VR0gBDcwNTAzBgZngQwBAgEwKTAnBggrBgEFBQcCARYbaHR0cDovL3d3dy5kaWdp Y2VydC5jb20vQ1BTMHYGCCsGAQUFBwEBBGowaDAmBggrBgEFBQcwAYYaaHR0cDov L3N0YXR1cy5yYXBpZHNzbC5jb20wPgYIKwYBBQUHMAKGMmh0dHA6Ly9jYWNlcnRz LnJhcGlkc3NsLmNvbS9SYXBpZFNTTFRMU1JTQUNBRzEuY3J0MAkGA1UdEwQCMAAw ggF9BgorBgEEAdZ5AgQCBIIBbQSCAWkBZwB1AO7N0GTV2xrOxVy3nbTNE6Iyh0Z8 vOzew1FIWUZxH7WbAAABhyniBv8AAAQDAEYwRAIgdTxb3FalokIXg75PlNXUinxZ 085Pg6teiwa02m6SorcCIAjlUa6XdVWs8G/taSuBeHR/j1lRVkvLM2C69sIv4QaE AHYAc9meiRtMlnigIH1HneayxhzQUV5xGSqMa4AQesF3crUAAAGHKeIHPQAABAMA RzBFAiAMUSADwj0B9AS3uNB0aa7+6J01Odxay5kpGp2BwM8nbgIhANh5TLAjRRg+ zrU+MYRskNWcgEQJ8bSeLiLLXA8YA76aAHYASLDja9qmRzQP5WoC+p0w6xxSActW 3SyB2bu/qznYhHMAAAGHKeIHJAAABAMARzBFAiA73w18S9dUznaz8hDG4qAm3PYT 8otBiybOJXOgZgvENAIhAKf6LW24VTCn9iS46yPeuyGm8gj6Uc+x0CYQIT5XSJqE MA0GCSqGSIb3DQEBCwUAA4IBAQB/TzOg8MSmahNcoI7Nv9INBoUQt12jfxdmSrSd SvDWwrFLzxeN5KfIdXQ3h+Os09epfQSvEibgpr9z6135/T5vmCocRtbd6tM7i0Fs b7kDIqK3E+I0StOTW5qafsYO5CEuOQSr7DBUmN973MRMhmHpqhwy3S4lUGlcBwMc R18mtzlk88qTIZIoU0oOc31bxchtsYxeddAyGhvqaFeQ3PegWc/dz7xUPx3WSqBR 1k/gzFgUBbPKsRuISAXAYdAm7kPFceOA8WxZwWjFD8vf/4mrfzKoly0k1aRY6jTB auM/9ji/JXH3+BEawzDeVZWVzb0BSN6yIBJv/C/PeBeUPKqx -----END CERTIFICATE----- rsa8192.badssl.com_certificate.crt000066400000000000000000000062271524413560000336020ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/certs-----BEGIN CERTIFICATE----- MIIJIDCCCAigAwIBAgIQDBboHFj/nFQ7X6MbvDiVMDANBgkqhkiG9w0BAQsFADBg MQswCQYDVQQGEwJVUzEVMBMGA1UEChMMRGlnaUNlcnQgSW5jMRkwFwYDVQQLExB3 d3cuZGlnaWNlcnQuY29tMR8wHQYDVQQDExZSYXBpZFNTTCBUTFMgUlNBIENBIEcx MB4XDTIzMDMyODAwMDAwMFoXDTI0MDMyNzIzNTk1OVowFzEVMBMGA1UEAwwMKi5i YWRzc2wuY29tMIIEIjANBgkqhkiG9w0BAQEFAAOCBA8AMIIECgKCBAEA2gk0w0XT FdWSvMYNgxJ1dyiSe5THwjxJfe35/Rqt3BSLVCTHBTgZKiBwXYMK11gsdrpBqWD3 Mp081yIjU5BC1uuUCV6su61TmWg5CoIQbM+imx8zXazBdu7xqvKXSK41wtb9HZhj 8cvQePK4WjnGePfMB/NoMww2ogny3Q9gCpFTYiIXXS/sjZ06fCRsfK32RsbOBRyT DfJWMu+vz+lwBbFkcKXTYUGnEyOQaS+Tu8RU564oYbYacMr7Ijvgz+Fsb71KmdEk Y809LnM7xZ4wfvlO7jXvMomhWlZBGL8xHPz7XgGlUXza77wZhsfepxeWzS/HbT3k 2qBtzpq4hTPnqd9yrZCpO/J+AtClt1tYw+orErVZ60+nWG4ZLpTqeY5wjagglPp6 yyA1jvhyz0U66ult9KjgZEGCC+QJXZPKwnKkqXnGWzxvqthUQWqRIoMnr6+LCP1e U3jvy0pmtPPLsY6o8xEfDPycY8fI1lKkuP0teo8jSVY4eVCl0X3hnm6nTrWP8+c8 I+Qc3j7V0RFJqWcPD7ouB7V6P5fSHLdmNWzamJVzrva494RCctmJ15rtfQ+/nwd0 pqNyL1GC3+OJ6Go2B1wjY1z/jEhrDcptLFJo9+nX6u331KA/3q6i+X52Od+kDmwl xeM2wkNgo4da3yu6TAnONnviopSIDjCp6/j5FdHTf1Xf1oYCVRJWje2Ur8r0l31u JV6q+uv/1+CitnPe9cB16KqEDQlm4igzwb6ikP/QJIU9XUu4p9uTY1fKTSLNku78 VEVNI72W9eY5uCz4kcw02CrbTKOVnEILNZKGpsKs8FeZmz7EXiKBiiXIB3i4JLvd SrU13RnXptHRg3imiKRObZ267gs3ho+un6TaoC7fE6RieBPZ+OaCbnKDMq1RnSPq iC/fizBS9b5KCnY8JcCIizfcSGEoGCVliNkXJhlU5qKaS4XCCNhVKDQ3M5vWCWCn dMEk563oIySanqIYWceYnD3b9hgEj6A88EKZnY50A8rn2jINEIU4ds+rQEAJjxA/ Zv1xaaTpiERGRj9Ln+9veppOHgap7KZpz8hELQ8zCLf3o3yWUCIBn3f0eLlkOwT+ rYVvVeg88F/5LmCPGegnn5D+IcjHuB2SGErvZKiCAE0CYF/xkONg3mo51Jli3cYB BV8wbU4C/sl4lTVeamUfnajY5ofPBEkRsDIhIqtGvxDk6Yr84rVvngDXTbP05eU3 9sSbActXbX2iHtrSzjGCdKV6UtAohuZtLTkut4e/8KJIZSuvqvf9dk0Z9jRnwmJv EmHQq5d63kwp/PtmFg6FpB7h1+Y9fXANoIuaC0u342VDOGoWCH5QUxKamhedy1Uz zh5PeHkZJ79jVQIDAQABo4IDHTCCAxkwHwYDVR0jBBgwFoAUDNtsgkkPSmcKuBTu esRIUojrVjgwHQYDVR0OBBYEFJzZKG+o1iyDszBUsyY2Ye6iJqhJMCMGA1UdEQQc MBqCDCouYmFkc3NsLmNvbYIKYmFkc3NsLmNvbTAOBgNVHQ8BAf8EBAMCBaAwHQYD VR0lBBYwFAYIKwYBBQUHAwEGCCsGAQUFBwMCMD8GA1UdHwQ4MDYwNKAyoDCGLmh0 dHA6Ly9jZHAucmFwaWRzc2wuY29tL1JhcGlkU1NMVExTUlNBQ0FHMS5jcmwwPgYD VR0gBDcwNTAzBgZngQwBAgEwKTAnBggrBgEFBQcCARYbaHR0cDovL3d3dy5kaWdp Y2VydC5jb20vQ1BTMHYGCCsGAQUFBwEBBGowaDAmBggrBgEFBQcwAYYaaHR0cDov L3N0YXR1cy5yYXBpZHNzbC5jb20wPgYIKwYBBQUHMAKGMmh0dHA6Ly9jYWNlcnRz LnJhcGlkc3NsLmNvbS9SYXBpZFNTTFRMU1JTQUNBRzEuY3J0MAkGA1UdEwQCMAAw ggF9BgorBgEEAdZ5AgQCBIIBbQSCAWkBZwB1AO7N0GTV2xrOxVy3nbTNE6Iyh0Z8 vOzew1FIWUZxH7WbAAABhyniBv8AAAQDAEYwRAIgdTxb3FalokIXg75PlNXUinxZ 085Pg6teiwa02m6SorcCIAjlUa6XdVWs8G/taSuBeHR/j1lRVkvLM2C69sIv4QaE AHYAc9meiRtMlnigIH1HneayxhzQUV5xGSqMa4AQesF3crUAAAGHKeIHPQAABAMA RzBFAiAMUSADwj0B9AS3uNB0aa7+6J01Odxay5kpGp2BwM8nbgIhANh5TLAjRRg+ zrU+MYRskNWcgEQJ8bSeLiLLXA8YA76aAHYASLDja9qmRzQP5WoC+p0w6xxSActW 3SyB2bu/qznYhHMAAAGHKeIHJAAABAMARzBFAiA73w18S9dUznaz8hDG4qAm3PYT 8otBiybOJXOgZgvENAIhAKf6LW24VTCn9iS46yPeuyGm8gj6Uc+x0CYQIT5XSJqE MA0GCSqGSIb3DQEBCwUAA4IBAQB/TzOg8MSmahNcoI7Nv9INBoUQt12jfxdmSrSd SvDWwrFLzxeN5KfIdXQ3h+Os09epfQSvEibgpr9z6135/T5vmCocRtbd6tM7i0Fs b7kDIqK3E+I0StOTW5qafsYO5CEuOQSr7DBUmN973MRMhmHpqhwy3S4lUGlcBwMc R18mtzlk88qTIZIoU0oOc31bxchtsYxeddAyGhvqaFeQ3PegWc/dz7xUPx3WSqBR 1k/gzFgUBbPKsRuISAXAYdAm7kPFceOA8WxZwWjFD8vf/4mrfzKoly0k1aRY6jTB auM/9ji/JXH3+BEawzDeVZWVzb0BSN6yIBJv/C/PeBeUPKqx -----END CERTIFICATE----- rsa8192.badssl.com_root_ca.crt000066400000000000000000000024161524413560000327420ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/certs-----BEGIN CERTIFICATE----- MIIDjjCCAnagAwIBAgIQAzrx5qcRqaC7KGSxHQn65TANBgkqhkiG9w0BAQsFADBh MQswCQYDVQQGEwJVUzEVMBMGA1UEChMMRGlnaUNlcnQgSW5jMRkwFwYDVQQLExB3 d3cuZGlnaWNlcnQuY29tMSAwHgYDVQQDExdEaWdpQ2VydCBHbG9iYWwgUm9vdCBH MjAeFw0xMzA4MDExMjAwMDBaFw0zODAxMTUxMjAwMDBaMGExCzAJBgNVBAYTAlVT MRUwEwYDVQQKEwxEaWdpQ2VydCBJbmMxGTAXBgNVBAsTEHd3dy5kaWdpY2VydC5j b20xIDAeBgNVBAMTF0RpZ2lDZXJ0IEdsb2JhbCBSb290IEcyMIIBIjANBgkqhkiG 9w0BAQEFAAOCAQ8AMIIBCgKCAQEAuzfNNNx7a8myaJCtSnX/RrohCgiN9RlUyfuI 2/Ou8jqJkTx65qsGGmvPrC3oXgkkRLpimn7Wo6h+4FR1IAWsULecYxpsMNzaHxmx 1x7e/dfgy5SDN67sH0NO3Xss0r0upS/kqbitOtSZpLYl6ZtrAGCSYP9PIUkY92eQ q2EGnI/yuum06ZIya7XzV+hdG82MHauVBJVJ8zUtluNJbd134/tJS7SsVQepj5Wz tCO7TG1F8PapspUwtP1MVYwnSlcUfIKdzXOS0xZKBgyMUNGPHgm+F6HmIcr9g+UQ vIOlCsRnKPZzFBQ9RnbDhxSJITRNrw9FDKZJobq7nMWxM4MphQIDAQABo0IwQDAP BgNVHRMBAf8EBTADAQH/MA4GA1UdDwEB/wQEAwIBhjAdBgNVHQ4EFgQUTiJUIBiV 5uNu5g/6+rkS7QYXjzkwDQYJKoZIhvcNAQELBQADggEBAGBnKJRvDkhj6zHd6mcY 1Yl9PMWLSn/pvtsrF9+wX3N3KjITOYFnQoQj8kVnNeyIv/iPsGEMNKSuIEyExtv4 NeF22d+mQrvHRAiGfzZ0JFrabA0UWTW98kndth/Jsw1HKj2ZL7tcu7XUIOGZX1NG Fdtom/DzMNU+MeKNhJ7jitralj41E6Vf8PlwUHBHQRFXGU7Aj64GxJUTFy8bJZ91 8rGOmaFvE7FBcf6IKshPECBV1/MUReXgRPTqh5Uykw7+U0b6LJ3/iyK5S9kJRaTe pLiaWN0bfVKfjllDiIGknibVb63dDcY3fe0Dkhvld1927jyNxF1WW6LZZm6zNTfl MrY= -----END CERTIFICATE----- cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/certs/snakeoil_cert.pem000066400000000000000000000025231524413560000307700ustar00rootroot00000000000000-----BEGIN CERTIFICATE----- MIIDwDCCAqigAwIBAgIUa3I/sWdZ+gcSTE6JZt8WEHIumLYwDQYJKoZIhvcNAQEL BQAwXzELMAkGA1UEBhMCWFgxFTATBgNVBAcMDERlZmF1bHQgQ2l0eTEcMBoGA1UE CgwTRGVmYXVsdCBDb21wYW55IEx0ZDEbMBkGA1UEAwwSRGVmYXVsdCBDb21wYW55 IENBMB4XDTIzMDUyOTE3MzMyNloXDTI1MDgzMTE3MzMyNlowVjELMAkGA1UEBhMC WFgxFTATBgNVBAcMDERlZmF1bHQgQ2l0eTEcMBoGA1UECgwTRGVmYXVsdCBDb21w YW55IEx0ZDESMBAGA1UEAwwJbG9jYWxob3N0MIIBIjANBgkqhkiG9w0BAQEFAAOC AQ8AMIIBCgKCAQEAyb+f0tHAhl6Pw9NKGgdvwu2kDWB1K9lN8YKUhzWanqwZCIDf XwrqegnrmGxMTz5hnOc1Sgk0pwIzeuJfFFmdH6mkvCfPXnQPhm5aiIyJ8uaz3Uwq A0EGYGd6qQcsbEh4nxep79DVQMul8NyI6q8EvWrdHx8hq4/GKVUl3fad1x7pmQs8 D6B489vIKvrvkqejgf7WpOubuJNbzSTlpX4nRp8JVG58cEDR2QgP6PLNT7HCvB1n MpfYOTe4Np+JYiHAqjO9KqlxnurXkQRSMTBLkAQ0/Nj6nNXB+JxR3ZSenMpM0WhA LVqvIQAKnujzwGdQS7l4BLs9tMgpblkshI+glQIDAQABo30wezAfBgNVHSMEGDAW gBTHdyAReTARJdyUeGx0lZTrdW/yRTAJBgNVHRMEAjAAMAsGA1UdDwQEAwIE8DAh BgNVHREEGjAYhwR/AAABhxAAAAAAAAAAAAAAAAAAAAABMB0GA1UdDgQWBBRYFnhD PfaWgm8mhMleIchZSbGqtDANBgkqhkiG9w0BAQsFAAOCAQEAYJqlrt2+YrRIrP7l TQW8J3C1fAapQJBAhWY+Zj2MhLuwSR+Hj/WBR56WRryT/BntMy5R3uUFnfP6k/3I dLDRsxrpj4nxRe+0JZLHymbdTyjoUmt6hTxdNWEdYnpdeRIf31BG8cmbhmzn5USp xvuOklLAWjNSwvzks0hg7GWTY7Fs9jCFQYpoPIOmjUmwz01bTZ/93gO6HTRt5zGA pglIH8yPfeEnBpzwbD4zcysq+OIMBL9sOucRj5RWU15Y4Lm/zoB9PuIAY53ILZfA cjoWOvHiCuzVFKNDblsuo3vONcOFORsOtYQzL4wRaRZz/D10ShE67pvwkciJtUpB yLk5xg== -----END CERTIFICATE----- cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/classes.py000066400000000000000000000476451524413560000263500ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import abc import base64 import collections import enum import json import pathlib import attr import pyfakefs.fake_filesystem_unittest from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.grade import AttackNamed, AttackType, Grade, GradeableVulnerabilities, Vulnerability from cryptodatahub.common.types import CryptoDataEnumBase, CryptoDataParamsNamed from cryptoparser.common.base import ( ComposerBinary, FourByteEnumComposer, FourByteEnumParsable, ListParamParsable, ListParsable, NumericRangeParsableBase, OneByteEnumComposer, OneByteEnumParsable, OpaqueEnumComposer, OpaqueEnumParsable, OpaqueParam, ParsableBase, ParserBinary, Serializable, SerializableTextEncoder, StringEnumCaseInsensitiveParsable, StringEnumParsable, ThreeByteEnumComposer, ThreeByteEnumParsable, TwoByteEnumComposer, TwoByteEnumParsable, VariantParsable, VariantParsableExact, ) from cryptoparser.common.exception import TooMuchData, InvalidType from cryptoparser.common.field import ( FieldsJson, FieldsSemicolonSeparated, FieldValueStringEnum, FieldValueComponentBool, FieldValueComponentFloat, FieldValueComponentNumber, FieldValueComponentOption, FieldValueComponentPercent, FieldValueComponentQuotedString, FieldValueComponentString, FieldValueComponentStringBase64, FieldValueComponentStringEnum, FieldValueComponentStringEnumParams, FieldValueComponentTimeDelta, FieldValueComponentUrl, FieldValueStringEnumParams, FieldValueTimeDelta, NameValuePairListSemicolonSeparated, ) from cryptoparser.common.parse import ParserCRLF from cryptoparser.common.x509 import PublicKeyX509 class NByteParsable(ParsableBase): def __init__(self, value): if value < 0 or value >= 2 ** (8 * self.get_byte_size()): raise ValueError self.value = value def __int__(self): return self.value @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('value', cls.get_byte_size()) return cls(parser['value']), cls.get_byte_size() def compose(self): composer = ComposerBinary() composer.compose_numeric(self.value, self.get_byte_size()) return composer.composed_bytes def __repr__(self): return f'0x{self.value:>0{self.get_byte_size() * 2}x}' def __eq__(self, other): return self.get_byte_size() == other.get_byte_size() and self.value == other.value @classmethod def get_byte_size(cls): raise NotImplementedError() class OneByteParsable(NByteParsable): @classmethod def get_byte_size(cls): return 1 class TwoByteParsable(NByteParsable): @classmethod def get_byte_size(cls): return 2 class ConditionalParsable(NByteParsable): def __int__(self): return self.value @classmethod def _parse(cls, parsable): parser = ParserBinary(parsable) parser.parse_numeric('value', cls.get_byte_size()) cls.check_parsed(parser['value']) return cls(parser['value']), cls.get_byte_size() def compose(self): composer = ComposerBinary() composer.compose_numeric(self.value, self.get_byte_size()) return composer.composed_bytes @classmethod @abc.abstractmethod def check_parsed(cls, value): raise NotImplementedError() class OneByteOddParsable(ConditionalParsable): @classmethod def get_byte_size(cls): return 1 @classmethod def check_parsed(cls, value): if value % 2 == 0: raise InvalidValue(value, OneByteOddParsable) class TwoByteEvenParsable(ConditionalParsable): @classmethod def get_byte_size(cls): return 2 @classmethod def check_parsed(cls, value): if value % 2 != 0: raise InvalidValue(value, TwoByteEvenParsable) class AlwaysUnknowTypeParsable(ParsableBase): @classmethod def _parse(cls, parsable): raise InvalidValue(parsable, AlwaysUnknowTypeParsable) def compose(self): raise TooMuchData() class AlwaysInvalidTypeParsable(ParsableBase): @classmethod def _parse(cls, parsable): raise InvalidType() def compose(self): raise TooMuchData() class AlwaysInvalidTypeVariantParsable(VariantParsable): @classmethod def _get_variants(cls): return collections.OrderedDict([ (AlwaysInvalidTypeParsable, (AlwaysInvalidTypeParsable, )) ]) class SerializableEnumVariantParsable(VariantParsable): @classmethod def _get_variants(cls): return collections.OrderedDict([ (SerializableEnum, (SerializableEnum, )) ]) class AlwaysTestStringComposer(ParsableBase): def __eq__(self, other): return isinstance(other, AlwaysTestStringComposer) @classmethod def _parse(cls, parsable): if parsable[:4] != b'test': raise InvalidValue(parsable, cls) return AlwaysTestStringComposer(), 4 def compose(self): return b'test' @attr.s class SerializableEnumValue(Serializable): code = attr.ib(validator=attr.validators.instance_of(int)) def as_json(self): return json.dumps({'code': self.code}) def _as_markdown(self, level): return False, self.code @classmethod def get_code_size(cls): return 2 class OpaqueEnumFactory(OpaqueEnumParsable): @classmethod def get_enum_class(cls): return OpaqueEnum @classmethod def get_param(cls): return OpaqueParam( min_byte_num=1, max_byte_num=32 ) @attr.s class OpaqueEnumParams: code = attr.ib(validator=attr.validators.instance_of(str)) class OpaqueEnum(OpaqueEnumComposer): ALPHA = OpaqueEnumParams(code='άλφα') BETA = OpaqueEnumParams(code='βήτα') GAMMA = OpaqueEnumParams(code='γάμμα') class TestObject: pass class SerializableSimpleTypes(Serializable): # pylint: disable=too-many-instance-attributes def __init__(self): self.int_value = 1 self.float_value = 1.0 self.bool_value = False self.str_value = 'string' self.bytearray_value = bytearray(b'\x00\x01\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f') self.none_value = None class SerializableIterables(Serializable): def __init__(self): self.dict_key = collections.OrderedDict([ (1, 'int'), ('str', 'string'), (SerializableStringEnum.FIRST, 'enum'), ]) self.dict_value = collections.OrderedDict([ ('int', 1), ('string', 'str'), ('enum', SerializableStringEnum.FIRST), ]) self.list_value = list(['value', ]) self.tuple_value = tuple(['value', ]) class SerializableEnumFactory(TwoByteEnumParsable): @classmethod def get_enum_class(cls): return SerializableEnum @abc.abstractmethod def compose(self): raise NotImplementedError() class SerializableEnum(TwoByteEnumComposer): FIRST = SerializableEnumValue( code=0x0001, ) SECOND = SerializableEnumValue( code=0x0002, ) class SerializableStringEnum(enum.Enum): FIRST = '1' SECOND = '2' class SerializableEnums(Serializable): def __init__(self): self.param_enum = SerializableEnum.FIRST self.string_enum = SerializableStringEnum.SECOND class SerializableHidden(Serializable): def __init__(self): self._invisible_value = None self.visible_value = 'value' class SerializableSingle(Serializable): def _asdict(self): return 'single' class SerializableUnhandled(Serializable): def __init__(self): self.complex_number = 1 + 2j @attr.s class SerializableHumanReadable(Serializable): attr_2 = attr.ib(default='value 2', metadata={'human_readable_name': 'Human Readable Name 2'}) attr_1 = attr.ib(default='value 1', metadata={'human_readable_name': 'Human Readable Name 1'}) @attr.s class SerializableHumanFriendly(Serializable): human_friendly = attr.ib(default='human-friendly', metadata={'human_friendly': True}) human_friendly_by_default = attr.ib(default='human-friendly-by-default') non_human_friendly = attr.ib(default='non-human-friendly', metadata={'human_friendly': False}) @attr.s class SerializableAttributeOrder(Serializable): attr_b = attr.ib(default='b') attr_a = attr.ib(default='a') class Class: def __init__(self): self.attr_b = 'b' self.attr_a = 'a' @attr.s class ClassAttr: attr_b = attr.ib(default='b') attr_a = attr.ib(default='a') class ClassAsDict: def __init__(self): self.attr_a = 'a' self.attr_b = 'b' def _asdict(self): return collections.OrderedDict([ ('attr_b', self.attr_b), ('attr_a', self.attr_a), ]) @attr.s class ClassAttrAsDict: attr_a = attr.ib(default='a') attr_b = attr.ib(default='b') def _asdict(self): return collections.OrderedDict([ ('attr_b', self.attr_b), ('attr_a', self.attr_a), ]) class ClassCryptoDataEnum(CryptoDataEnumBase): ONE = CryptoDataParamsNamed('one', 'long one') class ClassGradeable(GradeableVulnerabilities): @classmethod def get_gradeable_name(cls): return 'gradeable' def __str__(self): return 'value' class SerializableRecursive(Serializable): # pylint: disable=too-many-instance-attributes def __init__(self): self.json_object = Class() self.json_attr_object = ClassAttr() self.json_attr_as_dict = ClassAttrAsDict() self.json_asdict_object = ClassAsDict() self.json_crypto_data_hub_enum = ClassCryptoDataEnum.ONE self.json_gradeable = ClassGradeable([ Vulnerability(AttackType.MITM, Grade.INSECURE, AttackNamed.NOFS) ]) self.json_serializable_hidden = SerializableHidden() self.json_serializable_single = 'single' self.json_serializable_in_list = list([SerializableHidden(), 'single']) self.json_serializable_in_tuple = tuple([SerializableHidden(), 'single']) self.json_serializable_in_dict = dict({'key1': SerializableHidden(), 'key2': 'single'}) class SerializableEmptyValues(Serializable): def __init__(self): self.value = None self.list = [] self.tuple = tuple() self.dict = {} class FlagEnum(enum.IntEnum): ONE = 1 TWO = 2 FOUR = 4 EIGHT = 8 @attr.s class StringEnumParams: code = attr.ib() def _check_code(self, code): if code != self.code: raise InvalidType() @classmethod def get_code_size(cls): return 2 class StringEnum(StringEnumParsable, enum.Enum): ONE = StringEnumParams( code='one', ) TWO = StringEnumParams( code='two', ) THREE = StringEnumParams( code='three', ) class StringEnumA(StringEnumParsable, enum.Enum): A = StringEnumParams(code='a') class StringEnumAA(StringEnumParsable, enum.Enum): AA = StringEnumParams(code='aa') class StringEnumAAA(StringEnumParsable, enum.Enum): AAA = StringEnumParams(code='aaa') class VariantParsableTest(VariantParsable): @classmethod def _get_variants(cls): return collections.OrderedDict([ (StringEnumA, [StringEnumA, ]), (StringEnumAA, [StringEnumAA, ]), (StringEnumAAA, [StringEnumAAA, ]), ]) class VariantParsableExactTest(VariantParsableExact): @classmethod def _get_variants(cls): return collections.OrderedDict([ (StringEnumA, [StringEnumA, ]), (StringEnumAA, [StringEnumAA, ]), (StringEnumAAA, [StringEnumAAA, ]), ]) class EnumStringValue(enum.Enum): ONE = 'one' TWO = 'two' THREE = 'three' class ListParamParsableTest(ListParamParsable): pass class ListParsableTest(ListParsable): @classmethod def get_param(cls): return ListParamParsableTest( item_class=AlwaysTestStringComposer, fallback_class=None, separator_class=ParserCRLF, ) @attr.s class NByteEnumParam: code = attr.ib(validator=attr.validators.instance_of(int)) class NByteEnumTest(enum.Enum): ONE = NByteEnumParam(code=1) TWO = NByteEnumParam(code=2) THREE = NByteEnumParam(code=3) FOUR = NByteEnumParam(code=4) class OneByteEnumParsableTest(OneByteEnumParsable): @classmethod def get_enum_class(cls): return NByteEnumTest @abc.abstractmethod def compose(self): raise NotImplementedError() class OneByteEnumComposerTest(OneByteEnumComposer, enum.Enum): ONE = NByteEnumParam(code=1) TWO = NByteEnumParam(code=2) THREE = NByteEnumParam(code=3) FOUR = NByteEnumParam(code=4) class TwoByteEnumParsableTest(TwoByteEnumParsable): @classmethod def get_enum_class(cls): return NByteEnumTest @abc.abstractmethod def compose(self): raise NotImplementedError() class TwoByteEnumComposerTest(TwoByteEnumComposer, enum.Enum): ONE = NByteEnumParam(code=1) TWO = NByteEnumParam(code=2) THREE = NByteEnumParam(code=3) FOUR = NByteEnumParam(code=4) class ThreeByteEnumParsableTest(ThreeByteEnumParsable): @classmethod def get_enum_class(cls): return NByteEnumTest @abc.abstractmethod def compose(self): raise NotImplementedError() class ThreeByteEnumComposerTest(ThreeByteEnumComposer, enum.Enum): ONE = NByteEnumParam(code=1) TWO = NByteEnumParam(code=2) THREE = NByteEnumParam(code=3) FOUR = NByteEnumParam(code=4) class FourByteEnumParsableTest(FourByteEnumParsable): @classmethod def get_enum_class(cls): return NByteEnumTest @abc.abstractmethod def compose(self): raise NotImplementedError() class FourByteEnumComposerTest(FourByteEnumComposer, enum.Enum): ONE = NByteEnumParam(code=1) TWO = NByteEnumParam(code=2) THREE = NByteEnumParam(code=3) FOUR = NByteEnumParam(code=4) class FieldValueTimeDeltaTest(FieldValueTimeDelta): @classmethod def get_name(cls): return 'testTimeDelta' class FieldValueEnumTest(StringEnumCaseInsensitiveParsable, enum.Enum): FIRST = FieldValueStringEnumParams(code='first', human_readable_name='FiRsT') SECOND = FieldValueStringEnumParams(code='second') class FieldValueStringEnumTest(FieldValueStringEnum): @classmethod def _get_value_type(cls): return FieldValueEnumTest class FieldValueComponentOptionTest(FieldValueComponentOption): @classmethod def get_canonical_name(cls): return 'testOption' class FieldValueComponentStringTest(FieldValueComponentString): @classmethod def get_canonical_name(cls): return 'testString' class FieldValueComponentUrlTest(FieldValueComponentUrl): @classmethod def get_canonical_name(cls): return 'testUrl' class ComponentStringEnumTest(StringEnumParsable, enum.Enum): ONE = FieldValueComponentStringEnumParams( code='one' ) TWO = FieldValueComponentStringEnumParams( code='two' ) THREE = FieldValueComponentStringEnumParams( code='three' ) class FieldValueComponentStringBase64Test(FieldValueComponentStringBase64): @classmethod def get_canonical_name(cls): return 'testStringBase64' class FieldValueComponentBoolTest(FieldValueComponentBool): @classmethod def get_canonical_name(cls): return 'testBool' class FieldValueComponentFloatTest(FieldValueComponentFloat): @classmethod def get_canonical_name(cls): return 'testFloat' class FieldValueComponentStringEnumTest(FieldValueComponentStringEnum): @classmethod def get_canonical_name(cls): return 'testStringEnum' @classmethod def _get_value_type(cls): return ComponentStringEnumTest class FieldValueComponentOptionalStringTest(FieldValueComponentQuotedString): @classmethod def get_canonical_name(cls): return 'testOptionalString' class FieldValueComponentQuotedStringTest(FieldValueComponentQuotedString): @classmethod def get_canonical_name(cls): return 'testQuotedString' class FieldValueComponentNumberTest(FieldValueComponentNumber): @classmethod def get_canonical_name(cls): return 'testNumber' class FieldValueComponentPercentTest(FieldValueComponentPercent): @classmethod def get_canonical_name(cls): return 'testPercent' class FieldValueComponentTimeDeltaTest(FieldValueComponentTimeDelta): @classmethod def get_canonical_name(cls): return 'testTimeDelta' @attr.s class FieldValueComplexTestBase: # pylint: disable=too-many-instance-attributes time_delta = attr.ib( converter=FieldValueComponentTimeDeltaTest.convert, validator=attr.validators.instance_of(FieldValueComponentTimeDeltaTest) ) string = attr.ib( converter=FieldValueComponentStringTest.convert, validator=attr.validators.instance_of(FieldValueComponentStringTest), default='default' ) url = attr.ib( converter=FieldValueComponentUrlTest.convert, validator=attr.validators.instance_of(FieldValueComponentUrlTest), default='https://example.com' ) base64_string = attr.ib( converter=FieldValueComponentStringBase64Test.convert, validator=attr.validators.instance_of(FieldValueComponentStringBase64Test), default=base64.b64encode('default'.encode('ascii')).decode('ascii') ) number = attr.ib( converter=FieldValueComponentNumberTest.convert, validator=attr.validators.instance_of(FieldValueComponentNumberTest), default=0 ) percent = attr.ib( converter=FieldValueComponentPercentTest.convert, validator=attr.validators.instance_of(FieldValueComponentPercentTest), default=100 ) @attr.s class FieldValueMultipleTest( # pylint: disable=too-many-instance-attributes FieldsSemicolonSeparated, FieldValueComplexTestBase): option = attr.ib( converter=FieldValueComponentOptionTest.convert, validator=attr.validators.instance_of(FieldValueComponentOptionTest), default=False ) optional_string = attr.ib( converter=attr.converters.optional(FieldValueComponentOptionalStringTest.convert), validator=attr.validators.optional( attr.validators.instance_of(FieldValueComponentOptionalStringTest) ), default=None ) class FieldValueJsonTest( # pylint: disable=too-many-instance-attributes FieldsJson, FieldValueComplexTestBase): pass @attr.s class FieldValueMultipleExtendableTest(FieldValueMultipleTest): extensions = attr.ib( default=None, validator=attr.validators.optional(attr.validators.instance_of(NameValuePairListSemicolonSeparated)), metadata={'extension': True}, ) class SerializableUpperCaseEncoder(SerializableTextEncoder): def __call__(self, obj, level): _, string_result = super().__call__(obj, level) return False, string_result.upper() class TestClasses: class TestKeyBase(pyfakefs.fake_filesystem_unittest.TestCase): def setUp(self): self.setUpPyfakefs() self.__certs_dir = pathlib.PurePath(__file__).parent.parent / 'common' / 'certs' self.fs.add_real_directory(str(self.__certs_dir)) def _get_pem_str(self, public_key_file_name): public_key_path = self.__certs_dir / public_key_file_name with open(str(public_key_path), encoding='ascii') as pem_file: return pem_file.read() def _get_public_key_x509(self, public_key_file_name): return PublicKeyX509.from_pem(self._get_pem_str(public_key_file_name)) class NumericRangeParsableTest(NumericRangeParsableBase): @classmethod def _get_value_min(cls): return 0x01 @classmethod def _get_value_max(cls): return 0xfe @classmethod def _get_value_length(cls): return 1 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/test_algorithm.py000066400000000000000000000017011524413560000277170ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptodatahub.common.algorithm import Authentication, MAC, Signature from cryptodatahub.common.exception import InvalidValue class TestAlgortihmOIDBase(unittest.TestCase): def test_error_not_found(self): with self.assertRaisesRegex(InvalidValue, '\'1.2.3.4.5.6.7.8\' is not a valid Authentication oid value'): Authentication.from_oid('1.2.3.4.5.6.7.8') def test_from_oid(self): self.assertEqual( Authentication.from_oid(Authentication.RSA.value.oid), Authentication.RSA ) class TestAlgortihmParam(unittest.TestCase): def test_str(self): self.assertEqual(str(Signature.RSA_WITH_SHA2_224.value), 'SHA-224 with RSA Encryption') class TestMAC(unittest.TestCase): def test_digest_size(self): self.assertEqual(MAC.SHA2_256.value.digest_size, 256) self.assertEqual(MAC.POLY1305.value.digest_size, 128) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/test_base.py000066400000000000000000000725001524413560000266500ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import json import unittest from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import NotEnoughData, TooMuchData from cryptoparser.common.base import ( Opaque, OpaqueEnumParsable, OpaqueParam, ProtocolVersionMajorMinorBase, Serializable, SerializableTextEncoder, Vector, VectorEnumCodeString, VectorParamEnumCodeString, VectorParamNumeric, VectorParamParsable, VectorParamString, VectorParsable, VectorParsableDerived, VectorString, ) from cryptoparser.common.parse import ComposerBinary from .classes import ( AlwaysTestStringComposer, ConditionalParsable, EnumStringValue, FourByteEnumComposerTest, FourByteEnumParsableTest, ListParsableTest, NByteEnumTest, NumericRangeParsableTest, OneByteEnumComposerTest, OneByteEnumParsableTest, OneByteOddParsable, OneByteParsable, OpaqueEnum, OpaqueEnumFactory, SerializableAttributeOrder, SerializableEmptyValues, SerializableEnums, SerializableHidden, SerializableHumanFriendly, SerializableHumanReadable, SerializableIterables, SerializableRecursive, SerializableSimpleTypes, SerializableSingle, SerializableUnhandled, SerializableUpperCaseEncoder, StringEnum, StringEnumA, StringEnumAA, StringEnumAAA, TestObject, ThreeByteEnumComposerTest, ThreeByteEnumParsableTest, TwoByteEnumComposerTest, TwoByteEnumParsableTest, TwoByteEvenParsable, TwoByteParsable, VariantParsableTest, VariantParsableExactTest, ) class TestProtocolVersionMajorMinorBase(unittest.TestCase): def setUp(self): self.protocol_version = ProtocolVersionMajorMinorBase(1, 2) self.protocol_version_bytes = b'\x01\x02' def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: ProtocolVersionMajorMinorBase.parse_exact_size(b'\x00') self.assertEqual(context_manager.exception.bytes_needed, 1) def test_parse(self): self.assertEqual( ProtocolVersionMajorMinorBase.parse_exact_size(self.protocol_version_bytes), self.protocol_version ) def test_compose(self): self.assertEqual(self.protocol_version.compose(), self.protocol_version_bytes) def test_identifier(self): self.assertEqual(self.protocol_version.identifier, '1_2') def test_markdown(self): self.assertEqual(self.protocol_version.as_markdown(), '1.2') def test_str(self): self.assertEqual(str(self.protocol_version), '1.2') class VectorNumericTestErrors(Vector): @classmethod def get_param(cls): return VectorParamNumeric(item_size=2, min_byte_num=4, max_byte_num=6) class VectorNumericTest(Vector): @classmethod def get_param(cls): return VectorParamNumeric(item_size=2, min_byte_num=0, max_byte_num=0xff) class VectorStringTest(VectorString): @classmethod def get_param(cls): return VectorParamString( min_byte_num=0, max_byte_num=16, separator=';', item_class=StringEnum, fallback_class=str ) class VectorOneByteParsableTest(VectorParsable): @classmethod def get_param(cls): return VectorParamParsable(item_class=OneByteParsable, min_byte_num=0, max_byte_num=0xff, fallback_class=None) class VectorTwoByteParsableTest(VectorParsable): @classmethod def get_param(cls): return VectorParamParsable(item_class=TwoByteParsable, min_byte_num=0, max_byte_num=0xffff, fallback_class=None) class VectorConsditionalParsableTest(VectorParsableDerived): @classmethod def get_param(cls): return VectorParamParsable( item_class=ConditionalParsable, min_byte_num=0, max_byte_num=0xff, fallback_class=None ) class VectorFallbackParsableTest(VectorParsableDerived): @classmethod def get_param(cls): return VectorParamParsable( item_class=OneByteOddParsable, min_byte_num=0, max_byte_num=0xff, fallback_class=TwoByteEvenParsable ) class OpaqueTest(Opaque): @classmethod def get_param(cls): return OpaqueParam(min_byte_num=3, max_byte_num=3) class TestVectorNumeric(unittest.TestCase): def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: VectorNumericTestErrors(items=[1, ]) self.assertEqual(context_manager.exception.bytes_needed, VectorNumericTestErrors.get_param().min_byte_num) with self.assertRaises(TooMuchData) as context_manager: VectorNumericTestErrors(items=[1, 2, 3, 4, ]) self.assertEqual(context_manager.exception.bytes_needed, VectorNumericTestErrors.get_param().max_byte_num) vector = VectorNumericTestErrors(items=[1, 2, ]) with self.assertRaises(NotEnoughData) as context_manager: del vector[0] self.assertEqual(context_manager.exception.bytes_needed, VectorNumericTestErrors.get_param().min_byte_num) vector = VectorNumericTestErrors(items=[1, 2, 3]) with self.assertRaises(TooMuchData) as context_manager: vector.append(0xff) self.assertEqual(context_manager.exception.bytes_needed, VectorNumericTestErrors.get_param().max_byte_num) def test_parse(self): self.assertEqual(len(VectorNumericTest.parse_exact_size(b'\x00')), 0) self.assertEqual( [1, 2, ], list(VectorNumericTest.parse_exact_size(b'\x04\x00\x01\x00\x02')) ) def test_compose(self): self.assertEqual( b'\x00', VectorNumericTest([]).compose(), ) self.assertEqual( b'\x04\x00\x01\x00\x02', VectorNumericTest([1, 2, ]).compose(), ) def test_container(self): vector = VectorNumericTest(items=[]) vector.append(1) self.assertEqual(vector[0], 1) self.assertEqual(len(vector), 1) self.assertEqual(str(vector), '[1]') vector.insert(0, 0) self.assertEqual(vector[0], 0) self.assertEqual(len(vector), 2) self.assertEqual(str(vector), '[0, 1]') del vector[0] self.assertEqual(vector[0], 1) self.assertEqual(len(vector), 1) self.assertEqual(str(vector), '[1]') vector[0] = 0 self.assertEqual(vector[0], 0) self.assertEqual(len(vector), 1) self.assertEqual(str(vector), '[0]') class TestVectorString(unittest.TestCase): def test_error(self): pass def test_parse(self): self.assertEqual(len(VectorStringTest.parse_exact_size(b'\x00')), 0) self.assertEqual( [StringEnum.ONE, StringEnum.TWO, StringEnum.THREE, ], list(VectorStringTest.parse_exact_size(b'\x0fone;two;three')) ) self.assertEqual( [StringEnum.ONE, StringEnum.TWO, StringEnum.THREE, 'four', ], list(VectorStringTest.parse_exact_size(b'\x14one;two;three;four')) ) def test_compose(self): self.assertEqual( b'\x00', VectorStringTest([]).compose(), ) self.assertEqual( b'\x07one;two', VectorStringTest([StringEnum.ONE, StringEnum.TWO, ]).compose(), ) def test_json(self): self.assertEqual('[]', VectorStringTest([]).as_json()) self.assertEqual( VectorStringTest([StringEnum.ONE, StringEnum.TWO, ]).as_json(), '[{"ONE": {"code": "one"}}, {"TWO": {"code": "two"}}]', ) def test_markdown(self): self.assertEqual('-', VectorStringTest([]).as_markdown()) self.assertEqual( VectorStringTest([StringEnum.ONE, StringEnum.TWO, ]).as_markdown(), '\n'.join([ '1. ONE', '2. TWO', '', ]) ) class StringEnumFactory(OpaqueEnumParsable): @classmethod def get_enum_class(cls): return StringEnum @classmethod def get_param(cls): return OpaqueParam( min_byte_num=1, max_byte_num=2 ** 8 - 1 ) class VectorEnumCodeStringTest(VectorEnumCodeString): @classmethod def get_param(cls): return VectorParamEnumCodeString( item_class=StringEnumFactory, min_byte_num=0, max_byte_num=2 ** 16 - 1 ) class TestVectorEnumCodeString(unittest.TestCase): def test_error(self): pass def test_parse(self): self.assertEqual( [StringEnum.ONE, StringEnum.TWO, StringEnum.THREE, ], list(VectorEnumCodeStringTest.parse_exact_size(b'\x00\x0e\x03one\x03two\x05three')) ) def test_compose(self): self.assertEqual( b'\x00\x00', VectorEnumCodeStringTest([]).compose(), ) self.assertEqual( b'\x00\x08\x03one\x03two', VectorEnumCodeStringTest([StringEnum.ONE, StringEnum.TWO, ]).compose(), ) def test_json(self): self.assertEqual('[]', VectorEnumCodeStringTest([]).as_json()) self.assertEqual( VectorEnumCodeStringTest([StringEnum.ONE, StringEnum.TWO, ]).as_json(), '[{"ONE": {"code": "one"}}, {"TWO": {"code": "two"}}]', ) def test_markdown(self): self.assertEqual('-', VectorEnumCodeStringTest([]).as_markdown()) self.assertEqual( VectorEnumCodeStringTest([StringEnum.ONE, StringEnum.TWO, ]).as_markdown(), '\n'.join([ '1. ONE', '2. TWO', '', ]) ) class TestVectorParsable(unittest.TestCase): def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: VectorOneByteParsableTest.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, 1) with self.assertRaises(NotEnoughData) as context_manager: VectorTwoByteParsableTest.parse_exact_size(b'\x00') self.assertEqual(context_manager.exception.bytes_needed, 1) with self.assertRaises(NotEnoughData) as context_manager: VectorOneByteParsableTest.parse_exact_size(b'\x03') self.assertEqual(context_manager.exception.bytes_needed, 3) with self.assertRaises(NotEnoughData) as context_manager: VectorTwoByteParsableTest.parse_exact_size(b'\x00\x03') self.assertEqual(context_manager.exception.bytes_needed, 3) def test_parse(self): self.assertEqual(len(VectorOneByteParsableTest.parse_exact_size(b'\x00')), 0) self.assertEqual(len(VectorTwoByteParsableTest.parse_exact_size(b'\x00\x00')), 0) self.assertEqual( [0, 1, 0, 2, ], list(map(int, VectorOneByteParsableTest.parse_exact_size(b'\x04\x00\x01\x00\x02'))) ) self.assertEqual( [1, 2, ], list(map(int, VectorTwoByteParsableTest.parse_exact_size(b'\x00\x04\x00\x01\x00\x02'))) ) self.assertEqual( [0x01, 0x0200], list(map(int, VectorFallbackParsableTest.parse_exact_size(b'\x03\x01\x02\x00'))) ) def test_compose(self): self.assertEqual(b'\x00', VectorOneByteParsableTest([]).compose()) self.assertEqual(b'\x00\x00', VectorTwoByteParsableTest([]).compose()) self.assertEqual( b'\x04\x00\x01\x00\x02', VectorOneByteParsableTest([ OneByteParsable(0), OneByteParsable(1), OneByteParsable(0), OneByteParsable(2) ]).compose(), ) self.assertEqual( b'\x00\x04\x00\x01\x00\x02', VectorTwoByteParsableTest([ TwoByteParsable(1), TwoByteParsable(2) ]).compose(), ) def test_container(self): vector = VectorOneByteParsableTest(items=[]) vector.append(OneByteParsable(1)) self.assertEqual(vector[0], OneByteParsable(1)) self.assertNotEqual(vector[0], TwoByteParsable(1)) self.assertEqual(len(vector), 1) self.assertEqual(str(vector), '[0x01]') vector.insert(0, TwoByteParsable(0)) self.assertEqual(vector[0], TwoByteParsable(0)) self.assertNotEqual(vector[0], OneByteParsable(0)) self.assertEqual(len(vector), 2) self.assertEqual(str(vector), '[0x0000, 0x01]') del vector[0] self.assertEqual(vector[0], OneByteParsable(1)) self.assertNotEqual(vector[0], TwoByteParsable(1)) self.assertEqual(len(vector), 1) self.assertEqual(str(vector), '[0x01]') vector[0] = TwoByteParsable(0) self.assertEqual(vector[0], TwoByteParsable(0)) self.assertNotEqual(vector[0], OneByteParsable(0)) self.assertEqual(len(vector), 1) self.assertEqual(str(vector), '[0x0000]') class TestVectorDerived(unittest.TestCase): def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: VectorConsditionalParsableTest.parse_exact_size(b'\x05\x01\x02\x00') self.assertEqual(context_manager.exception.bytes_needed, 2) def test_parse(self): self.assertEqual( [0x01, 0x0200], list(map(int, VectorConsditionalParsableTest.parse_exact_size(b'\x03\x01\x02\x00'))) ) self.assertEqual( [0x0200, 0x01], list(map(int, VectorConsditionalParsableTest.parse_exact_size(b'\x03\x02\x00\x01'))) ) def test_compose(self): self.assertEqual( b'\x03\x01\x00\x02', VectorConsditionalParsableTest([ OneByteParsable(1), TwoByteParsable(2), ]).compose() ) self.assertEqual( b'\x03\x00\x02\x01', VectorConsditionalParsableTest([ TwoByteParsable(2), OneByteParsable(1), ]).compose() ) class TestOpaque(unittest.TestCase): def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: OpaqueTest.parse_exact_size(b'\x03\x01\x02') self.assertEqual(context_manager.exception.bytes_needed, 1) def test_parse(self): self.assertEqual( [1, 2, 3], list(OpaqueTest.parse_exact_size(b'\x03\x01\x02\x03')) ) def test_compose(self): self.assertEqual( b'\x03\x01\x02\x03', OpaqueTest([1, 2, 3]).compose() ) self.assertEqual( b'\x03\x01\x02\x03', OpaqueTest(b'\x01\x02\x03').compose() ) self.assertEqual( b'\x03\x01\x02\x03', OpaqueTest(bytearray(b'\x01\x02\x03')).compose() ) class TestOpaqueEnum(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: OpaqueEnumFactory.parse_exact_size(b'\x0a' + 'δέλτα'.encode()) self.assertEqual(context_manager.exception.value, 'δέλτα') def test_parse(self): self.assertEqual( OpaqueEnum.ALPHA, OpaqueEnumFactory.parse_exact_size(b'\x08' + 'άλφα'.encode()) ) def test_compose(self): self.assertEqual( b'\x0a' + 'γάμμα'.encode(), OpaqueEnum.GAMMA.compose() # pylint: disable=no-member ) def test_repr(self): self.assertEqual( repr(OpaqueEnum.GAMMA), 'OpaqueEnum.GAMMA' ) class TestNByteEnumParsable(unittest.TestCase): def test_compose(self): composer = ComposerBinary() composer.compose_parsable(OneByteEnumComposerTest.ONE) self.assertEqual(composer.composed_bytes, b'\x01') composer = ComposerBinary() composer.compose_parsable(TwoByteEnumComposerTest.TWO) self.assertEqual(composer.composed_bytes, b'\x00\x02') composer = ComposerBinary() composer.compose_parsable(ThreeByteEnumComposerTest.THREE) self.assertEqual(composer.composed_bytes, b'\x00\x00\x03') composer = ComposerBinary() composer.compose_parsable(FourByteEnumComposerTest.FOUR) self.assertEqual(composer.composed_bytes, b'\x00\x00\x00\x04') def test_parse(self): self.assertEqual( OneByteEnumParsableTest.parse_exact_size(b'\x01'), NByteEnumTest.ONE ) self.assertEqual( TwoByteEnumParsableTest.parse_exact_size(b'\x00\x02'), NByteEnumTest.TWO ) self.assertEqual( ThreeByteEnumParsableTest.parse_exact_size(b'\x00\x00\x03'), NByteEnumTest.THREE ) self.assertEqual( FourByteEnumParsableTest.parse_exact_size(b'\x00\x00\x00\x04'), NByteEnumTest.FOUR ) class TestEnumString(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: StringEnum.parse_exact_size(b'four') self.assertEqual(context_manager.exception.value, b'four') with self.assertRaises(InvalidValue) as context_manager: StringEnum.parse_exact_size(b'\xffthree') self.assertEqual(context_manager.exception.value, b'\xffthree') def test_parse(self): self.assertEqual(StringEnum.parse_exact_size(b'one'), StringEnum.ONE) def test_compose(self): self.assertEqual(StringEnum.ONE.compose(), b'one') class TestSerializable(unittest.TestCase): def test_json(self): self.assertEqual( SerializableSimpleTypes().as_json(), '{' + '"bool_value": false, ' + '"bytearray_value": "00:01:01:02:03:04:05:06:07:08:09:0A:0B:0C:0D:0E:0F", ' + '"float_value": 1.0, ' + '"int_value": 1, ' + '"none_value": null, ' + '"str_value": "string"' + '}' ) self.assertEqual( SerializableIterables().as_json(), '{' + '"dict_key": ' + '{' + '"1": "int", ' + '"str": "string", ' + '"FIRST": "enum"' + '}, ' + '"dict_value": ' + '{' + '"int": 1, ' + '"string": "str", ' + '"enum": {"FIRST": "1"}' + '}, ' + '"list_value": ["value"], ' + '"tuple_value": ["value"]' + '}' ) self.assertEqual( SerializableEnums().as_json(), '{"param_enum": {"FIRST": {"code": 1}}, "string_enum": {"SECOND": "2"}}' ) self.assertEqual( SerializableSingle().as_json(), '"single"' ) self.assertEqual( SerializableHidden().as_json(), '{"visible_value": "value"}' ) self.assertEqual( SerializableUnhandled().as_json(), '{"complex_number": "(1+2j)"}' ) self.assertEqual( SerializableHumanFriendly().as_json(), '{' + '"human_friendly": "human-friendly", ' + '"human_friendly_by_default": "human-friendly-by-default", ' + '"non_human_friendly": "non-human-friendly"' + '}' ) self.assertEqual( SerializableRecursive().as_json(), '{' + '"json_asdict_object": {"attr_b": "b", "attr_a": "a"}, ' + '"json_attr_as_dict": {"attr_b": "b", "attr_a": "a"}, ' + '"json_attr_object": {"attr_b": "b", "attr_a": "a"}, ' + '"json_crypto_data_hub_enum": "ONE", ' + '"json_gradeable": {"vulnerabilities": [{"attack_type": "MITM", "grade": "INSECURE", "named": "NOFS"}]}, ' + '"json_object": {"attr_a": "a", "attr_b": "b"}, ' + '"json_serializable_hidden": {"visible_value": "value"}, ' + '"json_serializable_in_dict": {"key1": {"visible_value": "value"}, "key2": "single"}, ' + '"json_serializable_in_list": [{"visible_value": "value"}, "single"], ' + '"json_serializable_in_tuple": [{"visible_value": "value"}, "single"], ' + '"json_serializable_single": "single"' + '}' ) self.assertEqual(json.dumps(TestObject()), '{}') self.assertEqual( json.dumps(EnumStringValue.ONE), '{"ONE": "one"}' ) def test_markdown(self): self.assertEqual( SerializableSimpleTypes().as_markdown(), '\n'.join([ '* Bool Value: no', '* Bytearray Value: 00:01:01:02:03:04:05:06:07:08:09:0A:0B:0C:0D:0E:0F', '* Float Value: 1.0', '* Int Value: 1', '* None Value: n/a', '* Str Value: string', '' ]) ) self.assertEqual( SerializableIterables().as_markdown(), '\n'.join([ '* Dict Key:', ' * 1: int', ' * Str: string', ' * FIRST: enum', '* Dict Value:', ' * Int: 1', ' * String: str', ' * Enum: FIRST', '* List Value:', ' 1. value', '* Tuple Value:', ' 1. value', '' ]) ) self.assertEqual( SerializableEnums().as_markdown(), '\n'.join([ '* Param Enum: 1', '* String Enum: SECOND', '' ]) ) self.assertEqual( SerializableHidden().as_markdown(), '* Visible Value: value\n' ) self.assertEqual( SerializableUnhandled().as_markdown(), '* Complex Number: (1+2j)\n' ) self.assertEqual( SerializableHumanReadable().as_markdown(), '* Human Readable Name 2: value 2\n' '* Human Readable Name 1: value 1\n' ) self.assertEqual( SerializableHumanFriendly().as_markdown(), '* Human Friendly: human-friendly\n' '* Human Friendly By Default: human-friendly-by-default\n' ) self.assertEqual( SerializableAttributeOrder().as_markdown(), '\n'.join([ '* Attr B: b', '* Attr A: a' ]) + '\n' ) self.assertEqual( SerializableEmptyValues().as_markdown(), '\n'.join([ '* Dict: -', '* List: -', '* Tuple: -', '* Value: n/a', '', ]) ) self.assertEqual( SerializableRecursive().as_markdown(), '\n'.join([ '* Json Asdict Object:', ' * Attr B: b', ' * Attr A: a', '* Json Attr As Dict:', ' * Attr B: b', ' * Attr A: a', '* Json Attr Object:', ' * Attr B: b', ' * Attr A: a', '* Json Crypto Data Hub Enum: one', '* Json Gradeable: value', '* Json Object:', ' * Attr A: a', ' * Attr B: b', '* Json Serializable Hidden:', ' * Visible Value: value', '* Json Serializable In Dict:', ' * Key1:', ' * Visible Value: value', ' * Key2: single', '* Json Serializable In List:', ' 1.', ' * Visible Value: value', ' 2. single', '* Json Serializable In Tuple:', ' 1.', ' * Visible Value: value', ' 2. single', '* Json Serializable Single: single', '', ]) ) Serializable.post_text_encoder = SerializableUpperCaseEncoder() self.assertEqual( SerializableRecursive().as_markdown(), '\n'.join([ '* Json Asdict Object:', ' * Attr B: B', ' * Attr A: A', '* Json Attr As Dict:', ' * Attr B: B', ' * Attr A: A', '* Json Attr Object:', ' * Attr B: B', ' * Attr A: A', '* Json Crypto Data Hub Enum: ONE', '* Json Gradeable: VALUE', '* Json Object:', ' * Attr A: A', ' * Attr B: B', '* Json Serializable Hidden:', ' * Visible Value: VALUE', '* Json Serializable In Dict:', ' * Key1:', ' * Visible Value: VALUE', ' * Key2: SINGLE', '* Json Serializable In List:', ' 1.', ' * Visible Value: VALUE', ' 2. SINGLE', '* Json Serializable In Tuple:', ' 1.', ' * Visible Value: VALUE', ' 2. SINGLE', '* Json Serializable Single: SINGLE', '', ]) ) Serializable.post_text_encoder = SerializableTextEncoder() class TestListParsable(unittest.TestCase): def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: ListParsableTest.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, 2) with self.assertRaises(InvalidValue) as context_manager: ListParsableTest.parse_exact_size(b'test') self.assertEqual(context_manager.exception.value, b'test') with self.assertRaises(InvalidValue) as context_manager: ListParsableTest.parse_exact_size(b'nottest\r\n\r\n') self.assertEqual(context_manager.exception.value, b'nottest\r\n\r\n') with self.assertRaises(InvalidValue) as context_manager: ListParsableTest.parse_exact_size(b'test\r\n\r\ntest\r\n\r\n') self.assertEqual(context_manager.exception.value, b'test\r\n\r\ntest\r\n\r\n') def test_parse(self): list_parsable = ListParsableTest.parse_exact_size(b'\r\n') self.assertEqual(list_parsable, ListParsableTest([])) list_parsable = ListParsableTest.parse_exact_size(b'test\r\n\r\n') self.assertEqual(list_parsable, ListParsableTest([AlwaysTestStringComposer(), ])) list_parsable = ListParsableTest.parse_exact_size(b'test\r\ntest\r\n\r\n') self.assertEqual(list_parsable, ListParsableTest([AlwaysTestStringComposer(), AlwaysTestStringComposer()])) def test_compose(self): self.assertEqual( ListParsableTest([]).compose(), b'\r\n' ) self.assertEqual( ListParsableTest([AlwaysTestStringComposer(), ]).compose(), b'test\r\n\r\n' ) self.assertEqual( ListParsableTest([AlwaysTestStringComposer(), AlwaysTestStringComposer(), ]).compose(), b'test\r\ntest\r\n\r\n' ) class TestVariantParsable(unittest.TestCase): def test_error(self): with self.assertRaises(TooMuchData) as context_manager: VariantParsableTest.parse_exact_size(b'aaaa') self.assertEqual(context_manager.exception.bytes_needed, 1) def test_parse(self): parsable = bytearray(b'a') self.assertEqual(VariantParsableTest.parse_mutable(parsable), StringEnumA.A) self.assertEqual(parsable, b'') parsable = bytearray(b'aa') self.assertEqual(VariantParsableTest.parse_mutable(parsable), StringEnumA.A) self.assertEqual(parsable, b'a') parsable = bytearray(b'aaa') self.assertEqual(VariantParsableTest.parse_mutable(parsable), StringEnumA.A) self.assertEqual(parsable, b'aa') class TestVariantParsableExact(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: VariantParsableExactTest.parse_exact_size(b'aaaa') self.assertEqual(context_manager.exception.value, b'aaaa') def test_parse(self): parsable = bytearray(b'a') self.assertEqual(VariantParsableExactTest.parse_mutable(parsable), StringEnumA.A) self.assertEqual(parsable, b'') parsable = bytearray(b'aa') self.assertEqual(VariantParsableExactTest.parse_mutable(parsable), StringEnumAA.AA) self.assertEqual(parsable, b'') parsable = bytearray(b'aaa') self.assertEqual(VariantParsableExactTest.parse_mutable(parsable), StringEnumAAA.AAA) self.assertEqual(parsable, b'') class TestNumericRangeParsable(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: NumericRangeParsableTest.parse_exact_size(b'\x00') self.assertEqual(context_manager.exception.value, 0x00) with self.assertRaises(InvalidValue) as context_manager: NumericRangeParsableTest.parse_exact_size(b'\xff') self.assertEqual(context_manager.exception.value, 0xff) def test_parse(self): self.assertEqual(NumericRangeParsableTest.parse_exact_size(b'\x01'), NumericRangeParsableTest(1)) self.assertEqual(NumericRangeParsableTest(1).compose(), b'\x01') def test_str(self): self.assertEqual(str(NumericRangeParsableTest(1)), '1') def test_as_markdown(self): self.assertEqual(NumericRangeParsableTest(1).as_markdown(), '1') cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/test_classes.py000066400000000000000000000030031524413560000273630ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.classes import LanguageTag class TestLanguageTag(unittest.TestCase): def setUp(self): self.language_tag = LanguageTag('a') def test_error(self): with self.assertRaises(InvalidValue): self.language_tag.primary_subtag = 'a1' with self.assertRaises(InvalidValue): self.language_tag.primary_subtag = '' with self.assertRaises(InvalidValue): self.language_tag.primary_subtag = 9 * 'a' with self.assertRaises(InvalidValue): self.language_tag.subsequent_subtags = ['a1', 'a#'] with self.assertRaises(InvalidValue): self.language_tag.subsequent_subtags = ['a1', 9 * 'a'] with self.assertRaises(InvalidValue): LanguageTag.parse_exact_size(b'') def test_parse(self): language_tag = LanguageTag.parse_exact_size(b'a') self.assertEqual(language_tag.primary_subtag, 'a') self.assertEqual(language_tag.subsequent_subtags, []) language_tag = LanguageTag.parse_exact_size(b'a-b-c') self.assertEqual(language_tag.primary_subtag, 'a') self.assertEqual(language_tag.subsequent_subtags, ['b', 'c', ]) def test_compose(self): self.assertEqual(LanguageTag('a').compose(), b'a') self.assertEqual(LanguageTag('a', []).compose(), b'a') self.assertEqual(LanguageTag('a', ['b', 'c', ]).compose(), b'a-b-c') cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/test_exception.py000066400000000000000000000025551524413560000277370ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import InvalidType, NotEnoughData, TooMuchData class TestException(unittest.TestCase): def test_str(self): with self.assertRaisesRegex( NotEnoughData, 'not enough data received from target; missing_byte_count="10"' ) as context_manager: raise NotEnoughData(10) self.assertEqual(context_manager.exception.bytes_needed, 10) with self.assertRaisesRegex( TooMuchData, 'too much data received from target; rest_byte_count="10"' ) as context_manager: raise TooMuchData(10) self.assertEqual(context_manager.exception.bytes_needed, 10) with self.assertRaisesRegex( InvalidValue, '0xa is not a valid str member name value' ) as context_manager: raise InvalidValue(10, str, 'member name') self.assertEqual(context_manager.exception.value, 10) with self.assertRaisesRegex(InvalidValue, '0xa is not a valid str') as context_manager: raise InvalidValue(10, str) self.assertEqual(context_manager.exception.value, 10) with self.assertRaisesRegex(InvalidType, 'invalid type value received from target') as context_manager: raise InvalidType(10) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/test_field.py000066400000000000000000000641531524413560000270260ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest import datetime from collections import OrderedDict import attr import urllib3 from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import InvalidType from cryptoparser.common.field import ( FieldValueDateTime, FieldValueString, NameValuePair, NameValuePairListSemicolonSeparated, ) from .classes import ( ComponentStringEnumTest, FieldValueJsonTest, FieldValueMultipleTest, FieldValueMultipleExtendableTest, FieldValueEnumTest, FieldValueStringEnumTest, FieldValueComponentBoolTest, FieldValueComponentFloatTest, FieldValueComponentNumberTest, FieldValueComponentOptionTest, FieldValueComponentQuotedStringTest, FieldValueComponentPercentTest, FieldValueComponentStringEnumTest, FieldValueComponentStringTest, FieldValueComponentTimeDeltaTest, FieldValueComponentUrlTest, FieldValueTimeDeltaTest, ) class TestFieldValueString(unittest.TestCase): def test_parse(self): self.assertEqual(FieldValueString.parse_exact_size(b'value').value, 'value') def test_compose(self): self.assertEqual(FieldValueString('value').compose(), b'value') def test_convert(self): self.assertEqual(FieldValueString.convert(None), None) self.assertEqual(FieldValueString.convert(bytearray(b'non-string-value')), bytearray(b'non-string-value')) self.assertEqual(FieldValueString.convert('value'), FieldValueString('value')) def test_markdown(self): self.assertEqual(FieldValueString('value').as_markdown(), 'value') class TestFieldValueStringEnum(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: FieldValueStringEnumTest('value not in enum') self.assertEqual(context_manager.exception.value, 'value not in enum') with self.assertRaises(InvalidValue) as context_manager: FieldValueStringEnumTest.parse_exact_size(b'value not in enum') self.assertEqual(context_manager.exception.value, 'value not in enum') def test_parse(self): self.assertEqual( FieldValueStringEnumTest.parse_exact_size(b'first'), FieldValueStringEnumTest(FieldValueEnumTest.FIRST) ) self.assertEqual( FieldValueStringEnumTest.parse_exact_size(b'FIRST'), FieldValueStringEnumTest(FieldValueEnumTest.FIRST) ) def test_compose(self): self.assertEqual( FieldValueStringEnumTest(FieldValueEnumTest.SECOND).compose(), b'second' ) def test_markdown(self): self.assertEqual( FieldValueStringEnumTest(FieldValueEnumTest.FIRST).as_markdown(), 'FiRsT' ) self.assertEqual( FieldValueEnumTest.FIRST.value.as_markdown(), 'FiRsT' ) self.assertEqual( FieldValueStringEnumTest(FieldValueEnumTest.SECOND).as_markdown(), 'second' ) self.assertEqual( FieldValueEnumTest.SECOND.value.as_markdown(), 'second' ) class TestFieldValueComponentOption(unittest.TestCase): def test_parse(self): component, _ = FieldValueComponentOptionTest.parse_immutable(b'option') self.assertEqual(component.value, False) component, _ = FieldValueComponentOptionTest.parse_immutable(b'testOption') self.assertEqual(component.value, True) def test_compose(self): self.assertEqual(FieldValueComponentOptionTest(False).compose(), b'') self.assertEqual(FieldValueComponentOptionTest(True).compose(), b'testOption') def test_as_markdown(self): self.assertEqual(FieldValueComponentOptionTest(False).as_markdown(), 'no') self.assertEqual(FieldValueComponentOptionTest(True).as_markdown(), 'yes') class TestFieldValueComponentString(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidType): FieldValueComponentStringTest.parse_exact_size(b'shortNam') with self.assertRaises(InvalidType): FieldValueComponentStringTest.parse_exact_size(b'wrongName=value') def test_parse(self): component = FieldValueComponentStringTest.parse_exact_size(b'testString=value') self.assertEqual(component.value, 'value') def test_compose(self): self.assertEqual(FieldValueComponentStringTest('value').compose(), b'testString=value') def test_as_markdown(self): self.assertEqual(FieldValueComponentStringTest('value').as_markdown(), 'value') class TestFieldValueComponentUrl(unittest.TestCase): _component_url_https = FieldValueComponentUrlTest('https://example.com') _component_url_https_bytes = b'testUrl=https://example.com' _component_url_mailto = FieldValueComponentUrlTest('mailto:user@example.com') _component_url_mailto_bytes = b'testUrl=mailto:user@example.com' def test_error_invalid_value(self): with self.assertRaises(InvalidValue) as context_manager: FieldValueComponentUrlTest.parse_exact_size(b'testUrl=https://example.com:port') self.assertEqual(context_manager.exception.value, 'https://example.com:port') with self.assertRaises(InvalidValue) as context_manager: FieldValueComponentUrlTest(None) self.assertEqual(context_manager.exception.value, None) def test_parse(self): component = FieldValueComponentUrlTest.parse_exact_size(b'testUrl=https://example.com') self.assertEqual(component.value, urllib3.util.parse_url('https://example.com')) component = FieldValueComponentUrlTest.parse_exact_size(b'testUrl=mailto:user@example.com') self.assertEqual(component.value, urllib3.util.parse_url('mailto:user@example.com')) def test_compose(self): self.assertEqual( FieldValueComponentUrlTest('https://example.com').compose(), b'testUrl=https://example.com' ) self.assertEqual( FieldValueComponentUrlTest('mailto:user@example.com').compose(), b'testUrl=mailto:user@example.com' ) def test_as_markdown(self): self.assertEqual(self._component_url_https.as_markdown(), 'https://example.com') self.assertEqual(self._component_url_mailto.as_markdown(), 'mailto:user@example.com') class TestFieldValueComponentStringEnum(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: FieldValueComponentStringEnumTest.parse_exact_size( # pylint: disable=expression-not-assigned b'testStringEnum=non-existing-value' ) self.assertEqual(context_manager.exception.value, 'non-existing-value') with self.assertRaises(InvalidValue) as context_manager: FieldValueComponentStringEnumTest('non-existing-value') self.assertEqual(context_manager.exception.value, 'non-existing-value') def test_parse(self): self.assertEqual( FieldValueComponentStringEnumTest.parse_exact_size(b'testStringEnum=one'), FieldValueComponentStringEnumTest(ComponentStringEnumTest.ONE) ) def test_compose(self): self.assertEqual( FieldValueComponentStringEnumTest(ComponentStringEnumTest.TWO).compose(), b'testStringEnum=two' ) class TestFieldValueComponentQuotedString(unittest.TestCase): def test_error(self): component = FieldValueComponentQuotedStringTest.parse_exact_size(b'testQuotedString="value') self.assertEqual(component.value, 'value') component = FieldValueComponentQuotedStringTest.parse_exact_size(b'testQuotedString=value"') self.assertEqual(component.value, 'value') component = FieldValueComponentQuotedStringTest.parse_exact_size(b'testQuotedString=""value"') self.assertEqual(component.value, 'value') component = FieldValueComponentQuotedStringTest.parse_exact_size(b'testQuotedString="value""') self.assertEqual(component.value, 'value') def test_parse(self): component = FieldValueComponentQuotedStringTest.parse_exact_size(b'testQuotedString=value') self.assertEqual(component.value, 'value') component = FieldValueComponentQuotedStringTest.parse_exact_size(b'testQuotedString="value"') self.assertEqual(component.value, 'value') def test_compose(self): self.assertEqual(FieldValueComponentQuotedStringTest('value').compose(), b'testQuotedString="value"') def test_markdown(self): self.assertEqual(FieldValueComponentQuotedStringTest('value').as_markdown(), 'value') class TestFieldValueComponentBool(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: FieldValueComponentBoolTest.parse_exact_size(b'testBool=str') self.assertEqual(context_manager.exception.value, b'str') def test_parse(self): component = FieldValueComponentBoolTest.parse_exact_size(b'testBool=yes') self.assertTrue(component.value) component = FieldValueComponentBoolTest.parse_exact_size(b'testBool=no') self.assertFalse(component.value) def test_compose(self): self.assertEqual(FieldValueComponentBoolTest(True).compose(), b'testBool=yes') self.assertEqual(FieldValueComponentBoolTest(False).compose(), b'testBool=no') def test_as_markdown(self): self.assertEqual(FieldValueComponentBoolTest(True).as_markdown(), 'yes') self.assertEqual(FieldValueComponentBoolTest(False).as_markdown(), 'no') class TestFieldValueComponentFloat(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: FieldValueComponentFloatTest.parse_exact_size(b'testFloat=str') self.assertEqual(context_manager.exception.value, b'str') def test_parse(self): component = FieldValueComponentFloatTest.parse_exact_size(b'testFloat=1') self.assertEqual(component.value, 1.0) component = FieldValueComponentFloatTest.parse_exact_size(b'testFloat=1.0') self.assertEqual(component.value, 1.0) def test_compose(self): self.assertEqual(FieldValueComponentFloatTest(1).compose(), b'testFloat=1.0') self.assertEqual(FieldValueComponentFloatTest(1.0).compose(), b'testFloat=1.0') def test_as_markdown(self): self.assertEqual(FieldValueComponentFloatTest(1).as_markdown(), '1.0') self.assertEqual(FieldValueComponentFloatTest(1.0).as_markdown(), '1.0') class TestFieldValueComponentNumber(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: FieldValueComponentNumberTest.parse_exact_size(b'testNumber=notnumeric') self.assertEqual(context_manager.exception.value, b'notnumeric') def test_parse(self): component = FieldValueComponentNumberTest.parse_exact_size(b'testNumber=1234') self.assertEqual(component.value, 1234) def test_compose(self): self.assertEqual(FieldValueComponentNumberTest(1234).compose(), b'testNumber=1234') def test_as_markdown(self): self.assertEqual(FieldValueComponentNumberTest(1234).as_markdown(), '1234') class TestFieldValueComponentPercent(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: FieldValueComponentPercentTest.parse_exact_size(b'testPercent=101') self.assertEqual(context_manager.exception.value, 101) class TestFieldValueComponentTimeDelta(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: FieldValueComponentTimeDeltaTest.parse_exact_size(b'testTimeDelta=notnumeric') self.assertEqual(context_manager.exception.value, b'notnumeric') with self.assertRaises(InvalidValue) as context_manager: FieldValueComponentTimeDeltaTest.parse_exact_size( b'testTimeDelta=' + str(2 ** 48).encode('ascii') ) self.assertEqual(context_manager.exception.value, 2 ** 48) def test_parse(self): component = FieldValueComponentTimeDeltaTest.parse_exact_size(b'testTimeDelta=86401') self.assertEqual(component.value, datetime.timedelta(days=1, seconds=1)) def test_compose(self): self.assertEqual( FieldValueComponentTimeDeltaTest(datetime.timedelta(days=1, seconds=1)).compose(), b'testTimeDelta=86401' ) def test_as_markdown(self): self.assertEqual( FieldValueComponentTimeDeltaTest(datetime.timedelta(days=1, seconds=1)).as_markdown(), '1 day, 0:00:01' ) class TestFieldJson(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: FieldValueJsonTest.parse_exact_size(b'not-a-valid-json') self.assertEqual(context_manager.exception.value, 'not-a-valid-json') def test_parse(self): header_field = FieldValueJsonTest.parse_exact_size(b'{"testTimeDelta": 1}') self.assertEqual( header_field, FieldValueJsonTest(datetime.timedelta(seconds=1)) ) self.assertEqual( header_field.string.value, # pylint: disable=no-member attr.fields_dict(FieldValueJsonTest)['string'].default ) self.assertEqual( header_field.number.value, # pylint: disable=no-member attr.fields_dict(FieldValueJsonTest)['number'].default ) parsed_header_field = FieldValueJsonTest.parse_exact_size( b'{"testTimeDelta": 1, "testString": "string", "testNumber": 1}' ) header_field = FieldValueJsonTest( time_delta=datetime.timedelta(seconds=1), string='string', number=1 ) self.assertEqual(parsed_header_field, header_field) parsed_header_field = FieldValueJsonTest.parse_exact_size(b'{' + b', '.join([ b'"testTimeDelta": 1', b'"testString": "string"', b'"testStringBase64": "ZGVmYXVsdA=="', b'"optional_string": "optional_string"', b'"testNumber": 1', ]) + b'}') header_field = FieldValueJsonTest( time_delta=datetime.timedelta(seconds=1), string='string', number=1 ) self.assertEqual(parsed_header_field, header_field) def test_compose(self): header_field = FieldValueJsonTest(datetime.timedelta(seconds=1)) self.assertEqual( header_field.compose(), b'{' + b', '.join([ b'"testTimeDelta": 1', b'"testString": "default"', b'"testUrl": "https://example.com"', b'"testStringBase64": "ZGVmYXVsdA=="', b'"testNumber": 0', b'"testPercent": 100', ]) + b'}' ) class TestFieldValueMultiple(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: FieldValueMultipleTest.parse_exact_size(b'') self.assertEqual(context_manager.exception.value, None) def test_init_param_default(self): header_field = FieldValueMultipleTest(datetime.timedelta(1)) self.assertEqual(header_field.option.value, False) self.assertEqual(header_field.number.value, 0) self.assertEqual(header_field.string.value, 'default') def test_init_param_convert(self): header_field_from_value = FieldValueMultipleTest( time_delta=datetime.timedelta(1), ) header_field_from_component = FieldValueMultipleTest( time_delta=FieldValueComponentTimeDeltaTest(datetime.timedelta(1)) ) self.assertEqual(header_field_from_value, header_field_from_component) def test_parse(self): header_field = FieldValueMultipleTest.parse_exact_size(b'testTimeDelta=1') self.assertEqual( header_field, FieldValueMultipleTest(datetime.timedelta(seconds=1)) ) self.assertEqual( header_field.option.value, # pylint: disable=no-member attr.fields_dict(FieldValueMultipleTest)['option'].default ) self.assertEqual( header_field.string.value, # pylint: disable=no-member attr.fields_dict(FieldValueMultipleTest)['string'].default ) self.assertEqual( header_field.number.value, # pylint: disable=no-member attr.fields_dict(FieldValueMultipleTest)['number'].default ) parsed_header_field = FieldValueMultipleTest.parse_exact_size( b'testTimeDelta=1; testOption; testString=string; testNumber=1' ) header_field = FieldValueMultipleTest( time_delta=datetime.timedelta(seconds=1), option=True, string='string', number=1 ) self.assertEqual(parsed_header_field, header_field) parsed_header_field = FieldValueMultipleTest.parse_exact_size(b'; '.join([ b'testTimeDelta=1', b'testOption', b'testString=string', b'testStringBase64="ZGVmYXVsdA=="', b'testOptionalString=optional_string', b'testNumber=1', ])) header_field = FieldValueMultipleTest( time_delta=datetime.timedelta(seconds=1), option=True, string='string', optional_string='optional_string', number=1 ) self.assertEqual(parsed_header_field, header_field) def test_compose(self): header_field = FieldValueMultipleTest(datetime.timedelta(seconds=1)) self.assertEqual( header_field.compose(), b'; '.join([ b'testTimeDelta=1', b'testString=default', b'testUrl=https://example.com', b'testStringBase64="ZGVmYXVsdA=="', b'testNumber=0', b'testPercent=100', ]) ) header_field.option.value = True self.assertEqual( header_field.compose(), b'; '.join([ b'testTimeDelta=1', b'testString=default', b'testUrl=https://example.com', b'testStringBase64="ZGVmYXVsdA=="', b'testNumber=0', b'testPercent=100', b'testOption', ]) ) class TestFieldValueMultipleExtendable(unittest.TestCase): def test_parse(self): parsed_header_field = FieldValueMultipleExtendableTest.parse_exact_size( b'testTimeDelta=1; testExtension1=value1; testExtension2=value2' ) header_field = FieldValueMultipleExtendableTest( time_delta=datetime.timedelta(seconds=1), extensions=NameValuePairListSemicolonSeparated( OrderedDict([('testExtension1', 'value1'), ('testExtension2', 'value2')]) ) ) self.assertEqual(parsed_header_field, header_field) def test_compose(self): header_field = FieldValueMultipleExtendableTest( datetime.timedelta(seconds=1), extensions=NameValuePairListSemicolonSeparated( OrderedDict([('testExtension1', 'value1'), ('testExtension2', 'value2')]) ) ) self.assertEqual( header_field.compose(), b'; '.join([ b'testTimeDelta=1', b'testString=default', b'testUrl=https://example.com', b'testStringBase64="ZGVmYXVsdA=="', b'testNumber=0', b'testPercent=100', b'testExtension1=value1', b'testExtension2=value2', ]) ) class TestFieldValueComponentKeyValue(unittest.TestCase): def test_parse_error(self): component = NameValuePair.parse_exact_size(b'name="') self.assertEqual(component.name, 'name') self.assertEqual(component.value, '') component = NameValuePair.parse_exact_size(b'name="value') self.assertEqual(component.name, 'name') self.assertEqual(component.value, 'value') def test_parse(self): component = NameValuePair.parse_exact_size(b'name=value') self.assertEqual(component.name, 'name') self.assertEqual(component.value, 'value') component = NameValuePair.parse_exact_size(b'name=') self.assertEqual(component.name, 'name') self.assertEqual(component.value, '') def test_parse_quoted(self): component = NameValuePair.parse_exact_size(b'name=""') self.assertEqual(component.name, 'name') self.assertEqual(component.value, '') component = NameValuePair.parse_exact_size(b'name="value"') self.assertEqual(component.name, 'name') self.assertEqual(component.value, 'value') def test_compose(self): self.assertEqual(NameValuePair('name', 'value').compose(), b'name=value') def test_compose_quoted(self): self.assertEqual(NameValuePair('name', 'value', quoted=True).compose(), b'name="value"') class TestFieldValueDateTime(unittest.TestCase): def test_parse(self): http_header_field = FieldValueDateTime.parse_exact_size(b'Wed, 21 Oct 2015 07:28:00 GMT') self.assertEqual( http_header_field.value, datetime.datetime(2015, 10, 21, 7, 28, tzinfo=datetime.timezone.utc) ) http_header_field = FieldValueDateTime.parse_exact_size(b'Wed, 21 Oct 2015 07:28:00 +01:00') self.assertEqual( http_header_field.value, datetime.datetime(2015, 10, 21, 7, 28, tzinfo=datetime.timezone(datetime.timedelta(hours=1))) ) def test_compose(self): self.assertEqual( FieldValueDateTime( datetime.datetime(2015, 10, 21, 7, 28, tzinfo=datetime.timezone.utc) ).compose(), b'Wed, 21 Oct 2015 07:28:00 GMT' ) class TestFieldValueTimeDelta(unittest.TestCase): def test_parse(self): http_header_field = FieldValueTimeDeltaTest.parse_exact_size(b'86401') self.assertEqual(http_header_field.value, datetime.timedelta(days=1, seconds=1)) def test_compose(self): self.assertEqual( FieldValueTimeDeltaTest(datetime.timedelta(days=1, seconds=1)).compose(), b'86401' ) class TestNameValuePairList(unittest.TestCase): _EMPTY_BYTES = b'' _EMPTY = NameValuePairListSemicolonSeparated(OrderedDict([])) _OPTION_BYTES = b'option' _OPTION = NameValuePairListSemicolonSeparated(OrderedDict([('option', None)])) _KEY_VALUE_BYTES = b'key=value' _KEY_VALUE = NameValuePairListSemicolonSeparated(OrderedDict([('key', 'value')])) _KEY_VALUE_OPTION_BYTES = b'key=value; option' _KEY_VALUE_OPTION = NameValuePairListSemicolonSeparated(OrderedDict([('key', 'value'), ('option', None)])) _OPTION_KEY_VALUE_BYTES = b'option; key=value' _OPTION_KEY_VALUE = NameValuePairListSemicolonSeparated(OrderedDict([('option', None), ('key', 'value')])) def test_error(self): self.assertEqual( NameValuePairListSemicolonSeparated.parse_exact_size(b';;;option'), self._OPTION ) self.assertEqual( NameValuePairListSemicolonSeparated.parse_exact_size(b'option;;;'), self._OPTION ) def test_markdown(self): self.assertEqual(self._EMPTY.as_markdown(), '-') self.assertEqual(self._OPTION.as_markdown(), '* Option: n/a\n') self.assertEqual(self._KEY_VALUE.as_markdown(), '* Key: value\n') self.assertEqual(self._KEY_VALUE_OPTION.as_markdown(), '* Key: value\n* Option: n/a\n') self.assertEqual(self._OPTION_KEY_VALUE.as_markdown(), '* Option: n/a\n* Key: value\n') def test_parse_mixed_values(self): self.assertEqual( NameValuePairListSemicolonSeparated.parse_exact_size(self._EMPTY_BYTES), self._EMPTY ) self.assertEqual( NameValuePairListSemicolonSeparated.parse_exact_size(self._OPTION_BYTES), self._OPTION ) self.assertEqual( NameValuePairListSemicolonSeparated.parse_exact_size(self._KEY_VALUE_BYTES), self._KEY_VALUE ) self.assertEqual( NameValuePairListSemicolonSeparated.parse_exact_size(self._KEY_VALUE_OPTION_BYTES), self._KEY_VALUE_OPTION ) self.assertEqual( NameValuePairListSemicolonSeparated.parse_exact_size(self._OPTION_KEY_VALUE_BYTES), self._OPTION_KEY_VALUE ) def test_compose_mixed_values(self): self.assertEqual(self._EMPTY.compose(), self._EMPTY_BYTES) self.assertEqual(self._OPTION.compose(), self._OPTION_BYTES) self.assertEqual(self._KEY_VALUE.compose(), self._KEY_VALUE_BYTES) self.assertEqual(self._KEY_VALUE_OPTION.compose(), self._KEY_VALUE_OPTION_BYTES) self.assertEqual(self._OPTION_KEY_VALUE.compose(), self._OPTION_KEY_VALUE_BYTES) def test_parse_single_option(self): header_field_value = NameValuePairListSemicolonSeparated.parse_exact_size(b'option') self.assertEqual(header_field_value.value, OrderedDict([('option', None), ])) def test_parse_muliple_options(self): header_field_value = NameValuePairListSemicolonSeparated.parse_exact_size(b'option1;option2;option3') self.assertEqual( header_field_value.value, OrderedDict([('option1', None), ('option2', None), ('option3', None), ]) ) header_field_value = NameValuePairListSemicolonSeparated.parse_exact_size( b' option1; \toption2;\t \toption3' ) self.assertEqual( header_field_value.value, OrderedDict([('option1', None), ('option2', None), ('option3', None), ]) ) def test_compose_options(self): self.assertEqual( NameValuePairListSemicolonSeparated(OrderedDict([])).compose(), b'' ) self.assertEqual(self._OPTION.compose(), self._OPTION_BYTES) self.assertEqual( NameValuePairListSemicolonSeparated( OrderedDict([('option1', None), ('option2', None), ('option3', None), ]) ).compose(), b'option1; option2; option3' ) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/test_parse.py000066400000000000000000001403421524413560000270500ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import datetime import unittest from unittest import mock import asn1crypto.x509 from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import NotEnoughData, TooMuchData, InvalidType from cryptoparser.common.parse import ParserBinary, ParserText, ParsableBase, ComposerBinary, ComposerText, ByteOrder from cryptoparser.tls.ciphersuite import TlsCipherSuiteFactory from .classes import ( AlwaysInvalidTypeVariantParsable, AlwaysTestStringComposer, ConditionalParsable, FlagEnum, OneByteOddParsable, OneByteParsable, SerializableEnum, SerializableEnumFactory, SerializableEnumVariantParsable, StringEnum, TwoByteParsable, ) class TestParsableBase(unittest.TestCase): _ALPHA_BETA_GAMMA = 'αβγ' _ALPHA_BETA_GAMMA_BYTES = 'αβγ'.encode() _ALPHA_BETA_GAMMA_LEN_BYTES = bytes((len(_ALPHA_BETA_GAMMA_BYTES),)) _ALPHA_BETA_GAMMA_HASHMARK_BYTES = 'αβγ#'.encode() class TestParsable(TestParsableBase): def test_error(self): with self.assertRaises(TooMuchData) as context_manager: OneByteParsable.parse_exact_size(b'\x01\x02') self.assertEqual(context_manager.exception.bytes_needed, 1) with self.assertRaises(NotEnoughData) as context_manager: OneByteParsable.parse_immutable(b'') self.assertEqual(context_manager.exception.bytes_needed, 1) with self.assertRaises(TypeError): # pylint: disable=protected-access,abstract-class-instantiated ParsableBase()._parse(b'') with self.assertRaises(TypeError): # pylint: disable=abstract-class-instantiated ParsableBase().compose() parser = ParserBinary(b'\xff\xff') with self.assertRaises(InvalidValue): parser.parse_parsable('cipher_suite', TlsCipherSuiteFactory) with self.assertRaises(InvalidValue) as context_manager: AlwaysInvalidTypeVariantParsable.parse_exact_size(b'\x01\x02\x03\x04') self.assertEqual(context_manager.exception.value, b'\x01\x02\x03\x04') with self.assertRaises(InvalidValue) as context_manager: AlwaysInvalidTypeVariantParsable(0) self.assertEqual(context_manager.exception.value, 0) def test_parse(self): _, parsed_length = OneByteParsable.parse_immutable(b'\x01\x02') self.assertEqual(parsed_length, 1) parsable = bytearray([0x01, 0x02]) OneByteParsable.parse_mutable(parsable) self.assertEqual(parsable, b'\x02') parsed_value, parsed_length = SerializableEnumFactory.parse_immutable(b'\x00\x01') self.assertEqual(parsed_value, SerializableEnum.FIRST) self.assertEqual(parsed_length, 2) def test_repr(self): self.assertEqual(repr(SerializableEnum.FIRST), 'SerializableEnum.FIRST') AlwaysInvalidTypeVariantParsable.register_variant_parser(SerializableEnumFactory, SerializableEnumFactory) class TestParserBase(TestParsableBase): def test_mapping(self): parser = ParserBinary(b'\x01\x02') self.assertEqual(len(parser), 0) self.assertEqual(parser.parsed_length, 0) self.assertEqual(parser.unparsed_length, 2) self.assertEqual(dict(parser), {}) parser.parse_numeric('first_byte', 1) self.assertEqual(len(parser), 1) self.assertEqual(parser.parsed_length, 1) self.assertEqual(parser.unparsed_length, 1) self.assertEqual(dict(parser), {'first_byte': 1}) parser.parse_numeric('second_byte', 1) self.assertEqual(len(parser), 2) self.assertEqual(parser.parsed_length, 2) self.assertEqual(parser.unparsed_length, 0) self.assertEqual(dict(parser), {'first_byte': 1, 'second_byte': 2}) class TestParserBinary(TestParsableBase): def test_error(self): parser = ParserBinary(b'\x00') parser.parse_numeric('one_byte', 1) with self.assertRaises(NotEnoughData) as context_manager: parser.parse_numeric('one_byte', 1) self.assertEqual(context_manager.exception.bytes_needed, 1) parser = ParserBinary(b'\x00\x00\x00\x00') parser.parse_numeric('one_byte', 1) with self.assertRaises(NotEnoughData) as context_manager: parser.parse_numeric_array('four_byte_array', item_num=2, item_size=3) self.assertEqual(context_manager.exception.bytes_needed, 3) parser = ParserBinary(b'\x00\x00\x00\x00') parser.parse_numeric('one_byte', 1) with self.assertRaises(NotEnoughData) as context_manager: parser.parse_bytes('four_byte_array', 4) self.assertEqual(context_manager.exception.bytes_needed, 1) parser = ParserBinary(b'\x01\xff') with self.assertRaises(InvalidValue) as context_manager: parser.parse_bytes('one_byte_array', 1, converter=mock.Mock(name='mock', side_effect=ValueError)) self.assertEqual(context_manager.exception.value, 1) parser = ParserBinary(b'\x00\x00\x00\x00\x00') with self.assertRaises(NotImplementedError): parser.parse_numeric('five_byte_numeric', 5) parser = ParserBinary(b'\xff\xff') with self.assertRaises(InvalidValue): parser.parse_numeric('two_byte_numeric', 2, OneByteParsable) parser = ParserBinary(b'\x10') with self.assertRaises(InvalidValue): parser.parse_numeric('flags', 1, FlagEnum) with self.assertRaises(InvalidValue) as context_manager: AlwaysInvalidTypeVariantParsable.parse_immutable(b'\x00\x00') AlwaysInvalidTypeVariantParsable.register_variant_parser(SerializableEnumFactory, SerializableEnumFactory) with self.assertRaises(InvalidValue) as context_manager: AlwaysInvalidTypeVariantParsable.parse_exact_size(b'\x01\x02') self.assertEqual(context_manager.exception.value, 0x0102) def test_parse_numeric(self): parser = ParserBinary(b'\x01\x02') parser.parse_numeric('first_byte', 1) parser.parse_numeric('second_byte', 1) self.assertEqual(parser['first_byte'], 0x01) self.assertEqual(parser['second_byte'], 0x02) parser = ParserBinary(b'\x01\x02') parser.parse_numeric('first_two_bytes', 2) self.assertEqual(parser['first_two_bytes'], 0x0102) parser = ParserBinary(b'\x01\x02\x03') parser.parse_numeric('first_two_bytes', 3) self.assertEqual(parser['first_two_bytes'], 0x010203) parser = ParserBinary(b'\x01\x02\x03\x04') parser.parse_numeric('first_four_bytes', 4) self.assertEqual(parser['first_four_bytes'], 0x01020304) def test_parse_byte_order(self): parser = ParserBinary(b'\x01\x02\x03\x04', byte_order=ByteOrder.BIG_ENDIAN) parser.parse_numeric('number', 4) self.assertEqual(parser['number'], 0x01020304) parser = ParserBinary(b'\x01\x02\x03\x04', byte_order=ByteOrder.LITTLE_ENDIAN) parser.parse_numeric('number', 4) self.assertEqual(parser['number'], 0x04030201) parser = ParserBinary(b'\x01\x02\x03\x04', byte_order=ByteOrder.NETWORK) parser.parse_numeric('number', 4) self.assertEqual(parser['number'], 0x01020304) parser = ParserBinary(b'\x01\x02\x03', byte_order=ByteOrder.BIG_ENDIAN) parser.parse_numeric('number', 3) self.assertEqual(parser['number'], 0x010203) parser = ParserBinary(b'\x01\x02\x03', byte_order=ByteOrder.LITTLE_ENDIAN) parser.parse_numeric('number', 3) self.assertEqual(parser['number'], 0x030201) parser = ParserBinary(b'\x01\x02\x03', byte_order=ByteOrder.NETWORK) parser.parse_numeric('number', 3) self.assertEqual(parser['number'], 0x010203) def test_parse_numeric_flags(self): parser = ParserBinary(b'\x01') parser.parse_numeric_flags('flags', 1, FlagEnum) self.assertEqual(parser['flags'], {FlagEnum.ONE, }) parser = ParserBinary(b'\x03') parser.parse_numeric_flags('flags', 1, FlagEnum) self.assertEqual(parser['flags'], {FlagEnum.ONE, FlagEnum.TWO}) def test_parse_mpint(self): parser = ParserBinary(b'\x00\x00\x00\x00') parser.parse_mpint('mpint', 4) self.assertEqual(parser['mpint'], 0) parser = ParserBinary(b'\x09\xa3\x78\xf9\xb2\xe3\x32\xa7') parser.parse_mpint('mpint', 8) self.assertEqual(parser['mpint'], 0x9a378f9b2e332a7) parser = ParserBinary(b'\x00\x80') parser.parse_mpint('mpint', 2) self.assertEqual(parser['mpint'], 0x80) parser = ParserBinary(b'\xed\xcc') parser.parse_mpint('mpint', 2) self.assertEqual(parser['mpint'], 0xedcc) parser = ParserBinary(b'\xff\x21\x52\x41\x11') parser.parse_mpint('mpint', 5) self.assertEqual(parser['mpint'], 0xff21524111) def test_parse_ssh_mpint(self): parser = ParserBinary(b'\x00') with self.assertRaises(NotEnoughData) as context_manager: parser.parse_ssh_mpint('mpint') self.assertEqual(context_manager.exception.bytes_needed, 3) parser = ParserBinary(b'\x00\x00\x00\x00') parser.parse_ssh_mpint('mpint') self.assertEqual(parser['mpint'], 0) parser = ParserBinary(b'\x00\x00\x00\x08\x09\xa3\x78\xf9\xb2\xe3\x32\xa7') parser.parse_ssh_mpint('mpint') self.assertEqual(parser['mpint'], 0x9a378f9b2e332a7) parser = ParserBinary(b'\x00\x00\x00\x02\x00\x80') parser.parse_ssh_mpint('mpint') self.assertEqual(parser['mpint'], 0x80) parser = ParserBinary(b'\x00\x00\x00\x02\xed\xcc') parser.parse_ssh_mpint('mpint') self.assertEqual(parser['mpint'], -0x1234) parser = ParserBinary(b'\x00\x00\x00\x05\xff\x21\x52\x41\x11') parser.parse_ssh_mpint('mpint') self.assertEqual(parser['mpint'], -0xdeadbeef) def test_parse_numeric_array(self): parser = ParserBinary(b'\x01\x02') parser.parse_numeric_array('one_byte_array', item_num=2, item_size=1) self.assertEqual(parser['one_byte_array'], [1, 2]) parser = ParserBinary(b'\x00\x01\x00\x02') parser.parse_numeric_array('two_byte_array', item_num=2, item_size=2) self.assertEqual(parser['two_byte_array'], [1, 2]) parser = ParserBinary(b'\x00\x00\x01\x00\x00\x02') parser.parse_numeric_array('three_byte_array', item_num=2, item_size=3) self.assertEqual(parser['three_byte_array'], [1, 2]) parser = ParserBinary(b'\x00\x00\x00\x01\x00\x00\x00\x02') parser.parse_numeric_array('four_byte_array', item_num=2, item_size=4) self.assertEqual(parser['four_byte_array'], [1, 2]) def test_parse_byte_array(self): parser = ParserBinary(b'\x01\x02') parser.parse_raw('two_byte_array', size=2) self.assertEqual(parser['two_byte_array'], b'\x01\x02') parser = ParserBinary(b'not X.509 certificate') with self.assertRaises(InvalidValue) as context_manager: parser.parse_raw('name', 21, asn1crypto.x509.Certificate.load) self.assertEqual(context_manager.exception.value, b'not X.509 certificate') def test_parse_string(self): parser = ParserBinary(b'\x02\xff\xff') with self.assertRaises(InvalidValue): parser.parse_string('non-utf-8-string', 1, 'utf-8') parser = ParserBinary(self._ALPHA_BETA_GAMMA_LEN_BYTES + self._ALPHA_BETA_GAMMA_BYTES) parser.parse_string('utf-8-string', 1, 'utf-8') self.assertEqual(parser['utf-8-string'], self._ALPHA_BETA_GAMMA) def test_parse_string_null_terminated(self): parser = ParserBinary(b'non-null-terminated-string') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_null_terminated('non-utf-8-string', 1, 'utf-8') self.assertEqual(context_manager.exception.value, b'non-null-terminated-string') parser = ParserBinary(b'1\x00') parser.parse_string_null_terminated('one-byte-string', 'utf-8', int) self.assertEqual(parser['one-byte-string'], 1) parser = ParserBinary(self._ALPHA_BETA_GAMMA_BYTES + b'\x00remaining-data') parser.parse_string_null_terminated('utf-8-string', 'utf-8') self.assertEqual(parser['utf-8-string'], self._ALPHA_BETA_GAMMA) self.assertEqual(parser.unparsed, b'remaining-data') def test_parse_parsable(self): parser = ParserBinary(b'\x01\x02\x03\x04') parser.parse_parsable('first_byte', OneByteParsable) self.assertEqual( b'\x01', parser['first_byte'].compose() ) parser.parse_parsable('second_byte', OneByteParsable) self.assertEqual( b'\x02', parser['second_byte'].compose() ) parser = ParserBinary(b'\x01\x02') parser.parse_parsable('byte', OneByteParsable, 1) self.assertEqual( b'\x02', parser['byte'].compose() ) parser = ParserBinary(b'\x02\x01\x02') with self.assertRaises(TooMuchData) as context_manager: parser.parse_parsable('byte', OneByteParsable, 1) self.assertEqual(context_manager.exception.bytes_needed, 1) def test_parse_parsable_array(self): parser = ParserBinary(b'\x01\x02\x03\x04') parser.parse_parsable_array('array', items_size=4, item_class=OneByteParsable) self.assertEqual( [0x01, 0x02, 0x03, 0x04], list(map(int, parser['array'])) ) parser = ParserBinary(b'\x01\x02\x03\x04') parser.parse_parsable_array('array', items_size=4, item_class=TwoByteParsable) self.assertEqual( [0x0102, 0x0304], list(map(int, parser['array'])) ) parser = ParserBinary(b'\x01\x02') with self.assertRaises(InvalidValue): parser.parse_parsable_array('array', items_size=2, item_class=OneByteOddParsable) parser = ParserBinary(b'\x00') with self.assertRaises(NotEnoughData) as context_manager: parser.parse_parsable_array('array', items_size=3, item_class=OneByteOddParsable) self.assertEqual(context_manager.exception.bytes_needed, 2) def test_parse_parsable_derived_array(self): parser = ParserBinary(b'\x01\x02\x00') parser.parse_parsable_derived_array( 'array', items_size=3, item_base_class=ConditionalParsable, fallback_class=None ) self.assertEqual( [0x01, 0x0200], list(map(int, parser['array'])) ) self.assertEqual(parser.unparsed_length, 0) self.assertEqual(parser.unparsed, b'') parser = ParserBinary(b'\x00\x01') with self.assertRaises(InvalidValue): parser.parse_parsable_derived_array( 'array', items_size=2, item_base_class=ConditionalParsable, fallback_class=None ) parser = ParserBinary(b'\x00\x01') parser.parse_parsable_derived_array( 'array', items_size=2, item_base_class=ConditionalParsable, fallback_class=TwoByteParsable ) self.assertEqual( [0x01, ], list(map(int, parser['array'])) ) self.assertEqual(parser.unparsed_length, 0) self.assertEqual(parser.unparsed, b'') def test_parse_variant_parsable(self): AlwaysInvalidTypeVariantParsable.register_variant_parser(SerializableEnumFactory, SerializableEnumFactory) self.assertEqual( AlwaysInvalidTypeVariantParsable.parse_exact_size(b'\x00\x01').value, AlwaysInvalidTypeVariantParsable(SerializableEnum.FIRST).variant.value ) def test_parse_timestamp(self): parser = ParserBinary(b'\x00\x00\x00\x00\x00\x00\x00\x00') parser.parse_timestamp('timestamp') self.assertEqual(parser['timestamp'], datetime.datetime.fromtimestamp(0, datetime.timezone.utc)) parser = ParserBinary(b'\x00\x00\x00\x00') parser.parse_timestamp('timestamp', item_size=4) self.assertEqual(parser['timestamp'], datetime.datetime.fromtimestamp(0, datetime.timezone.utc)) parser = ParserBinary(b'\x00\x00\x00\x00\x00\x00\x00\xff') parser.parse_timestamp('timestamp', milliseconds=True) self.assertEqual( parser['timestamp'], datetime.datetime.fromtimestamp(0, datetime.timezone.utc) + datetime.timedelta(microseconds=255000) ) parser = ParserBinary(b'\x00\x00\x00\xff') parser.parse_timestamp('timestamp', milliseconds=True, item_size=4) self.assertEqual( parser['timestamp'], datetime.datetime.fromtimestamp(0, datetime.timezone.utc) + datetime.timedelta(microseconds=255000) ) parser = ParserBinary(b'\xff\xff\xff\xff\xff\xff\xff\xff') parser.parse_timestamp('timestamp') self.assertEqual(parser['timestamp'], None) parser = ParserBinary(b'\xff\xff\xff\xff') parser.parse_timestamp('timestamp', item_size=4) self.assertEqual(parser['timestamp'], None) parser = ParserBinary(b'\x00\x00\x00\x00\xff\xff\xff\xff') parser.parse_timestamp('timestamp') self.assertEqual(parser['timestamp'], datetime.datetime.fromtimestamp(0xffffffff, datetime.timezone.utc)) parser = ParserBinary(b'\x00\x00\xff\xff') parser.parse_timestamp('timestamp', item_size=4) self.assertEqual(parser['timestamp'], datetime.datetime.fromtimestamp(0x0000ffff, datetime.timezone.utc)) class TestParserText(TestParsableBase): def test_error(self): parser = ParserText(b'\xff') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_until_separator('string', '#') self.assertEqual(context_manager.exception.value, b'\xff') def test_separator(self): parser = ParserText(b';') self.assertEqual(parser.unparsed_length, 1) parser.parse_separator(';') self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b';;test;;') parser.parse_separator(';', max_length=None) self.assertEqual(parser.unparsed_length, 6) parser.parse_string_by_length('test', 4, 4) self.assertEqual(parser.unparsed_length, 2) parser.parse_separator(';', max_length=None) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b';;') with self.assertRaises(InvalidValue) as context_manager: parser.parse_separator(';', max_length=1) self.assertEqual(context_manager.exception.value, b';;') with self.assertRaises(InvalidValue) as context_manager: parser.parse_separator(';', min_length=3) self.assertEqual(context_manager.exception.value, b';;') def test_parse_numeric(self): parser = ParserText(b'1#') parser.parse_numeric('number') self.assertEqual(parser['number'], 1) parser = ParserText(b'NaN') with self.assertRaises(InvalidValue) as context_manager: parser.parse_numeric('number') self.assertEqual(context_manager.exception.value, b'NaN') parser = ParserText(b'1a') parser.parse_numeric('number') self.assertEqual(parser['number'], 1) self.assertEqual(parser.unparsed_length, 1) parser = ParserText(b'1.2#') parser.parse_numeric_array('number', 2, '.') self.assertEqual(parser['number'], [1, 2]) parser = ParserText(b'1#') with self.assertRaises(InvalidValue) as context_manager: parser.parse_numeric_array('number', 2, '.') self.assertEqual(context_manager.exception.value, b'1') def test_parse_float(self): parser = ParserText(b'1.2#') parser.parse_float('number') self.assertEqual(parser['number'], 1.2) parser = ParserText(b'1.#') parser.parse_float('number') self.assertEqual(parser['number'], 1.0) parser = ParserText(b'1#') parser.parse_float('number') self.assertEqual(parser['number'], 1.0) parser = ParserText(b'NaN') with self.assertRaises(InvalidValue) as context_manager: parser.parse_float('number') self.assertEqual(context_manager.exception.value, b'NaN') def test_parse_string_until_separator(self): parser = ParserText(b'a#') parser.parse_string_until_separator('string', '#') self.assertEqual(parser['string'], 'a') self.assertEqual(parser.unparsed_length, 1) parser = ParserText(b'12#') parser.parse_string_until_separator('number', '#', int) self.assertEqual(parser['number'], 12) self.assertEqual(parser.unparsed_length, 1) parser = ParserText(b'three#one') parser.parse_string_until_separator('number', '#', StringEnum) self.assertEqual(parser['number'], StringEnum.THREE) self.assertEqual(parser.unparsed_length, 4) parser = ParserText(self._ALPHA_BETA_GAMMA_HASHMARK_BYTES, 'utf-8') parser.parse_string_until_separator('alphabet', '#') self.assertEqual(parser['alphabet'], 'αβγ') self.assertEqual(parser.unparsed_length, 1) parser = ParserText(b'ab') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_until_separator('string', '#') self.assertEqual(context_manager.exception.value, b'ab') parser = ParserText(b'ab') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_until_separator('string', '#') self.assertEqual(context_manager.exception.value, b'ab') parser = ParserText(b'12a#') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_until_separator('string', '#', int) self.assertEqual(context_manager.exception.value, b'12a#') parser = ParserText(self._ALPHA_BETA_GAMMA_HASHMARK_BYTES, 'ascii') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_until_separator('alphabet', '#') self.assertEqual(context_manager.exception.value, self._ALPHA_BETA_GAMMA_HASHMARK_BYTES) def test_parse_string_until_separator_or_end(self): parser = ParserText(b'ab') parser.parse_string_until_separator_or_end('string', '#') self.assertEqual(parser['string'], 'ab') self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b'12') parser.parse_string_until_separator_or_end('number', '#', int) self.assertEqual(parser['number'], 12) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b'a#b') parser.parse_string_until_separator_or_end('string', '#') self.assertEqual(parser['string'], 'a') self.assertEqual(parser.unparsed_length, 2) def test_parse_string(self): parser = ParserText(b'abc') parser.parse_string('string', 'abc') self.assertEqual(parser['string'], 'abc') parser = ParserText(b'abcd') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string('string', 'bcd') self.assertEqual(context_manager.exception.value, b'abc') parser = ParserText(b'abc') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string('string', 'abcd') self.assertEqual(context_manager.exception.value, b'abc') def test_parse_string_by_length(self): parser = ParserText(b'abc') parser.parse_string_by_length('string', 1, None) self.assertEqual(parser['string'], 'abc') self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b'12') parser.parse_string_by_length('number', 1, None, int) self.assertEqual(parser['number'], 12) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b'12ab') parser.parse_string_by_length('string', 1, 2, int) self.assertEqual(parser['string'], 12) self.assertEqual(parser.unparsed_length, 2) parser = ParserText(b'12') with self.assertRaises(NotEnoughData) as context_manager: parser.parse_string_by_length('string', 3, 3) self.assertEqual(context_manager.exception.bytes_needed, 1) parser = ParserText(b'12ab') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_by_length('string', 3, 4, int) self.assertEqual(context_manager.exception.value, '12ab') parser = ParserText(self._ALPHA_BETA_GAMMA_HASHMARK_BYTES, 'ascii') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_by_length('alphabet') self.assertEqual(context_manager.exception.value, self._ALPHA_BETA_GAMMA_HASHMARK_BYTES) def test_parse_bool(self): parser = ParserText(b'yes') parser.parse_bool('bool') self.assertEqual(parser['bool'], True) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b'no') parser.parse_bool('bool') self.assertEqual(parser['bool'], False) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b'abcd') with self.assertRaises(InvalidValue) as context_manager: parser.parse_bool('bool') self.assertEqual(context_manager.exception.value, b'abcd') class TestParserTextStringArray(TestParsableBase): def test_empty(self): parser = ParserText(b'') parser.parse_string_array('array', ',', skip_empty=True) self.assertEqual(parser['array'], []) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b'') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_array('array', ',', skip_empty=False) self.assertEqual(context_manager.exception.value, b'') self.assertEqual(parser.unparsed_length, 0) def test_separator_only(self): parser = ParserText(b',,,') parser.parse_string_array('array', ',', skip_empty=True) self.assertEqual(parser['array'], []) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b',,,') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_array('array', ',', skip_empty=False) self.assertEqual(context_manager.exception.value, b',,,') self.assertEqual(parser.unparsed_length, 3) def test_space_only(self): parser = ParserText(b' ') parser.parse_string_array('array', ',', separator_spaces=' ', skip_empty=True) self.assertEqual(parser['array'], []) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b' ') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_array('array', ',', separator_spaces=' ') self.assertEqual(context_manager.exception.value, b'') self.assertEqual(parser.unparsed_length, 3) def test_separator_and_spaces(self): parser = ParserText(b' , ,, ,,,') parser.parse_string_array('array', ',', separator_spaces=' ', skip_empty=True) self.assertEqual(parser['array'], []) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b' , ,, ,,,') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_array('array', ',', separator_spaces=' ') self.assertEqual(context_manager.exception.value, b', ,, ,,,') self.assertEqual(parser.unparsed_length, 12) def test_one_character_separator(self): parser = ParserText(b'a,b') parser.parse_string_array('array', ',') self.assertEqual(parser['array'], ['a', 'b']) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b'a,b') parser.parse_string_array('array', ',', item_class=ord) self.assertEqual(parser['array'], [ord('a'), ord('b')]) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b'one,two') parser.parse_string_array('array', ',', item_class=StringEnum) self.assertEqual(parser['array'], [StringEnum.ONE, StringEnum.TWO]) self.assertEqual(parser.unparsed_length, 0) def test_separator_spaces(self): parser = ParserText(b' a; \tb\t;\t c') parser.parse_string_array('array', ';', separator_spaces=' \t') self.assertEqual(parser['array'], ['a', 'b', 'c']) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b' one; \ttwo\t;\t three') parser.parse_string_array('array', ';', item_class=StringEnum, separator_spaces=' \t') self.assertEqual(parser['array'], [StringEnum.ONE, StringEnum.TWO, StringEnum.THREE]) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b' a \t b ; \tc') parser.parse_string_array('array', ';', separator_spaces='\t') self.assertEqual(parser['array'], [' a \t b ', ' \tc']) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b' a ') parser.parse_string_array('array', ';', separator_spaces=' ') self.assertEqual(parser['array'], ['a', ]) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b' one ') parser.parse_string_array('array', ';', item_class=StringEnum, separator_spaces=' ') self.assertEqual(parser['array'], [StringEnum.ONE, ]) self.assertEqual(parser.unparsed_length, 0) def test_starts_with_separator(self): parser = ParserText(b'; a; b; c') parser.parse_string_array('array', ';', separator_spaces=' ', skip_empty=True) self.assertEqual(parser['array'], ['a', 'b', 'c']) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b'; a; b; c') with self.assertRaises(InvalidValue) as context_manager: parser.parse_string_array('array', ';', separator_spaces=' ') self.assertEqual(context_manager.exception.value, b'; a; b; c') self.assertEqual(parser.unparsed_length, 9) def test_ends_with_separator(self): parser = ParserText(b'a; b; c; ') parser.parse_string_array('array', ';', separator_spaces=' ', skip_empty=True) self.assertEqual(parser['array'], ['a', 'b', 'c']) self.assertEqual(parser.unparsed_length, 0) parser = ParserText(b'one; two; three; ') parser.parse_string_array('array', ';', item_class=StringEnum, separator_spaces=' ', skip_empty=True) self.assertEqual(parser['array'], [StringEnum.ONE, StringEnum.TWO, StringEnum.THREE]) self.assertEqual(parser.unparsed_length, 0) def test_ends_without_separator(self): parser = ParserText(b'a; b; c') parser.parse_string_array('array', ';', separator_spaces=' ', max_item_num=2) self.assertEqual(parser['array'], ['a', 'b']) self.assertEqual(parser.unparsed_length, 1) parser = ParserText(b'one; two; three') parser.parse_string_array('array', ';', item_class=StringEnum, separator_spaces=' ', max_item_num=2) self.assertEqual(parser['array'], [StringEnum.ONE, StringEnum.TWO]) self.assertEqual(parser.unparsed_length, 5) class TestParserTextDateTime(TestParsableBase): def test_parse_date_time(self): parser = ParserText(b'Wed, 21 Oct 2015 07:28:00 GMT') parser.parse_date_time('datetime') self.assertEqual( parser['datetime'], datetime.datetime(2015, 10, 21, 7, 28, tzinfo=datetime.timezone.utc) ) self.assertEqual(parser.unparsed_length, 0) datetime_value = b'not a date' parser = ParserText(datetime_value) with self.assertRaises(InvalidValue) as context_manager: parser.parse_date_time('datetime') self.assertEqual(context_manager.exception.value, datetime_value) self.assertEqual(parser.unparsed_length, len(datetime_value)) class TestParserTextTimeDelta(TestParsableBase): def test_parse_time_delta(self): parser = ParserText(b'86400') parser.parse_time_delta('timedelta') self.assertEqual(parser['timedelta'], datetime.timedelta(1)) self.assertEqual(parser.unparsed_length, 0) timedelta_value = str(int(datetime.timedelta.max.total_seconds())).encode('ascii') parser = ParserText(timedelta_value) with self.assertRaises(InvalidValue) as context_manager: parser.parse_time_delta('timedelta') self.assertEqual( context_manager.exception.value, int(datetime.timedelta.max.total_seconds()) ) self.assertEqual(parser.unparsed_length, len(timedelta_value)) class TestComposerBinary(TestParsableBase): def test_error(self): composer = ComposerBinary() for size in (1, 2, 4): min_value = 0 max_value = 2 ** (size * 8) with self.assertRaises(InvalidValue) as context_manager: composer.compose_numeric(max_value + 1, size) self.assertEqual(context_manager.exception.value, max_value + 1) with self.assertRaises(InvalidValue) as context_manager: composer.compose_numeric(min_value - 1, size) self.assertEqual(context_manager.exception.value, min_value - 1) def test_compose_numeric_to_right_size(self): composer = ComposerBinary() composer.compose_numeric(0x01, 1) self.assertEqual(composer.composed_bytes, b'\x01') composer = ComposerBinary() composer.compose_numeric(0x01, 2) self.assertEqual(composer.composed_bytes, b'\x00\x01') composer = ComposerBinary() composer.compose_numeric(0x01, 3) self.assertEqual(composer.composed_bytes, b'\x00\x00\x01') composer = ComposerBinary() composer.compose_numeric(0x01, 4) self.assertEqual(composer.composed_bytes, b'\x00\x00\x00\x01') def test_compose_numeric_to_rigth_order(self): composer = ComposerBinary() composer.compose_numeric(0x01, 1) self.assertEqual(composer.composed_bytes, b'\x01') composer = ComposerBinary() composer.compose_numeric(0x0102, 2) self.assertEqual(composer.composed_bytes, b'\x01\x02') composer = ComposerBinary() composer.compose_numeric(0x010203, 3) self.assertEqual(composer.composed_bytes, b'\x01\x02\x03') composer = ComposerBinary() composer.compose_numeric(0x01020304, 4) self.assertEqual(composer.composed_bytes, b'\x01\x02\x03\x04') def test_compose_byte_order(self): composer = ComposerBinary(byte_order=ByteOrder.BIG_ENDIAN) composer.compose_numeric(0x01020304, 4) self.assertEqual(composer.composed_bytes, b'\x01\x02\x03\x04') composer = ComposerBinary(byte_order=ByteOrder.LITTLE_ENDIAN) composer.compose_numeric(0x01020304, 4) self.assertEqual(composer.composed_bytes, b'\x04\x03\x02\x01') composer = ComposerBinary(byte_order=ByteOrder.NETWORK) composer.compose_numeric(0x01020304, 4) self.assertEqual(composer.composed_bytes, b'\x01\x02\x03\x04') composer = ComposerBinary(byte_order=ByteOrder.BIG_ENDIAN) composer.compose_numeric(0x010203, 3) self.assertEqual(composer.composed_bytes, b'\x01\x02\x03') composer = ComposerBinary(byte_order=ByteOrder.LITTLE_ENDIAN) composer.compose_numeric(0x010203, 3) self.assertEqual(composer.composed_bytes, b'\x03\x02\x01') composer = ComposerBinary(byte_order=ByteOrder.NETWORK) composer.compose_numeric(0x010203, 3) self.assertEqual(composer.composed_bytes, b'\x01\x02\x03') def test_compose_numeric_array(self): composer = ComposerBinary() composer.compose_numeric_array(values=[1, 2, 3, 4], item_size=1) self.assertEqual(composer.composed_bytes, b'\x01\x02\x03\x04') composer = ComposerBinary() composer.compose_numeric_array(values=[1, 2, 3, 4], item_size=2) self.assertEqual(composer.composed_bytes, b'\x00\x01\x00\x02\x00\x03\x00\x04') composer = ComposerBinary() composer.compose_numeric_array(values=[1, 2, 3, 4], item_size=3) self.assertEqual(composer.composed_bytes, b'\x00\x00\x01\x00\x00\x02\x00\x00\x03\x00\x00\x04') composer = ComposerBinary() composer.compose_numeric_array(values=[1, 2, 3, 4], item_size=4) self.assertEqual(composer.composed_bytes, b'\x00\x00\x00\x01\x00\x00\x00\x02\x00\x00\x00\x03\x00\x00\x00\x04') def test_compose_numeric_enum_coded(self): composer = ComposerBinary() composer.compose_numeric_enum_coded(SerializableEnum.FIRST) self.assertEqual(composer.composed_bytes, b'\x00\x01') def test_compose_numeric_array_enum_coded(self): composer = ComposerBinary() composer.compose_numeric_array_enum_coded(values=[]) self.assertEqual(composer.composed_bytes, b'') composer = ComposerBinary() composer.compose_numeric_array_enum_coded(values=[ SerializableEnum.FIRST, SerializableEnum.SECOND, ]) self.assertEqual(composer.composed_bytes, b'\x00\x01\x00\x02') def test_compose_numeric_flags(self): composer = ComposerBinary() composer.compose_numeric_flags([FlagEnum.ONE, ], 1) self.assertEqual(composer.composed_bytes, b'\x01') composer = ComposerBinary() composer.compose_numeric_flags([FlagEnum.ONE, FlagEnum.TWO, ], 1) self.assertEqual(composer.composed_bytes, b'\x03') def test_compose_mpint(self): composer = ComposerBinary() with self.assertRaises(InvalidValue) as context_manager: composer.compose_mpint(1024, 1) self.assertEqual(context_manager.exception.value, 1) composer = ComposerBinary() composer.compose_mpint(1024, 2) self.assertEqual(composer.composed_bytes, b'\x04\x00') composer = ComposerBinary() composer.compose_mpint(1024, 10) self.assertEqual(composer.composed_bytes, b'\x00\x00\x00\x00\x00\x00\x00\x00\x04\x00') composer = ComposerBinary() composer.compose_mpint(-1024, 10) self.assertEqual(composer.composed_bytes, b'\xff\xff\xff\xff\xff\xff\xff\xff\xfc\x00') composer = ComposerBinary(byte_order=ByteOrder.LITTLE_ENDIAN) composer.compose_mpint(1000, 2) self.assertEqual(composer.composed_bytes, b'\xe8\x03') composer = ComposerBinary(byte_order=ByteOrder.LITTLE_ENDIAN) composer.compose_mpint(1000, 10) self.assertEqual(composer.composed_bytes, b'\xe8\x03\x00\x00\x00\x00\x00\x00\x00\x00') composer = ComposerBinary(byte_order=ByteOrder.LITTLE_ENDIAN) composer.compose_mpint(-1000, 10) self.assertEqual(composer.composed_bytes, b'\x18\xfc\xff\xff\xff\xff\xff\xff\xff\xff') def test_compose_ssh_mpint(self): composer = ComposerBinary() composer.compose_ssh_mpint(0) self.assertEqual(composer.composed_bytes, b'\x00\x00\x00\x00') composer = ComposerBinary() composer.compose_ssh_mpint(0x9a378f9b2e332a7) self.assertEqual(composer.composed_bytes, b'\x00\x00\x00\x08\x09\xa3\x78\xf9\xb2\xe3\x32\xa7') composer = ComposerBinary() composer.compose_ssh_mpint(0x80) self.assertEqual(composer.composed_bytes, b'\x00\x00\x00\x02\x00\x80') composer = ComposerBinary() composer.compose_ssh_mpint(-0x1234) self.assertEqual(composer.composed_bytes, b'\x00\x00\x00\x02\xed\xcc') composer = ComposerBinary() composer.compose_ssh_mpint(-0xdeadbeef) self.assertEqual(composer.composed_bytes, b'\x00\x00\x00\x05\xff\x21\x52\x41\x11') def test_compose_raw(self): composer = ComposerBinary() composer.compose_raw(b'\x01\x02\x03\x04') self.assertEqual(composer.composed_bytes, b'\x01\x02\x03\x04') def test_compose_string(self): composer = ComposerBinary() with self.assertRaises(InvalidValue) as context_manager: composer.compose_string(self._ALPHA_BETA_GAMMA, 'ascii', 1) self.assertEqual(context_manager.exception.value, self._ALPHA_BETA_GAMMA) composer = ComposerBinary() composer.compose_string(self._ALPHA_BETA_GAMMA, 'utf-8', 1) self.assertEqual(composer.composed_bytes[1:], self._ALPHA_BETA_GAMMA_BYTES) def test_compose_string_null_terminated(self): composer = ComposerBinary() with self.assertRaises(InvalidValue) as context_manager: composer.compose_string_null_terminated(self._ALPHA_BETA_GAMMA, 'ascii') self.assertEqual(context_manager.exception.value, self._ALPHA_BETA_GAMMA) composer = ComposerBinary() composer.compose_string_null_terminated(self._ALPHA_BETA_GAMMA, 'utf-8') self.assertEqual(composer.composed_bytes, self._ALPHA_BETA_GAMMA_BYTES + b'\x00') def test_compose_string_enum_coded(self): composer = ComposerBinary() composer.compose_string_enum_coded(StringEnum.ONE, 2) self.assertEqual(composer.composed_bytes, b'\x00\x03one') def test_compose_multiple(self): composer = ComposerBinary() one_byte_parsable = OneByteParsable(0x01) composer.compose_parsable(one_byte_parsable) self.assertEqual(composer.composed_bytes, b'\x01') composer.compose_numeric(0x02, 1) self.assertEqual(composer.composed_bytes, b'\x01\x02') self.assertEqual(composer.composed_length, 2) composer.compose_numeric(0x0304, 2) self.assertEqual(composer.composed_bytes, b'\x01\x02\x03\x04') self.assertEqual(composer.composed_length, 4) composer.compose_numeric(0x050607, 3) self.assertEqual(composer.composed_bytes, b'\x01\x02\x03\x04\x05\x06\x07') self.assertEqual(composer.composed_length, 7) composer.compose_numeric(0x08090a0b, 4) self.assertEqual(composer.composed_bytes, b'\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b') self.assertEqual(composer.composed_length, 11) def test_compose_parsable(self): composer = ComposerText() composer.compose_parsable(StringEnum.THREE) composer.compose_parsable(StringEnum.ONE) self.assertEqual(composer.composed, b'threeone') composer = ComposerBinary() composer.compose_parsable(OneByteParsable(0x01)) composer.compose_parsable(TwoByteParsable(0x0203)) self.assertEqual( b'\x01\x02\x03', composer.composed_bytes ) def test_compose_parsable_array(self): composer = ComposerText() with self.assertRaises(InvalidType): composer.compose_parsable_array([ StringEnum.THREE, StringEnum.ONE, 'four' ]) composer = ComposerText() with self.assertRaises(InvalidType): composer.compose_parsable_array([ StringEnum.THREE, StringEnum.ONE, 'four' ], fallback_class=int) composer = ComposerText() composer.compose_parsable_array([ StringEnum.THREE, StringEnum.ONE, 'four' ], fallback_class=str) self.assertEqual(composer.composed, b'three,one,four') composer = ComposerBinary() parsable_array = [OneByteParsable(0x01), TwoByteParsable(0x0203), ] composer.compose_parsable_array(parsable_array) self.assertEqual( b'\x01\x02\x03', composer.composed_bytes ) def test_compose_enum(self): composer = ComposerBinary() composer.compose_parsable(SerializableEnum.SECOND) self.assertEqual(b'\x00\x02', composer.composed_bytes) composer = ComposerBinary() composer.compose_parsable(SerializableEnum.FIRST, item_size=1) self.assertEqual(b'\x02\x00\x01', composer.composed_bytes) def test_compose_variant_parsable(self): composer = ComposerBinary() composer.compose_parsable(SerializableEnumVariantParsable(SerializableEnum.FIRST)) self.assertEqual(b'\x00\x01', composer.composed_bytes) def test_compose_timestamp(self): composer = ComposerBinary() date_time = datetime.datetime.fromtimestamp(0, datetime.timezone.utc) composer.compose_timestamp(date_time) self.assertEqual(b'\x00\x00\x00\x00\x00\x00\x00\x00', composer.composed_bytes) composer = ComposerBinary() date_time = datetime.datetime.fromtimestamp(0, datetime.timezone.utc) + datetime.timedelta(microseconds=255000) composer.compose_timestamp(date_time, milliseconds=True) self.assertEqual(b'\x00\x00\x00\x00\x00\x00\x00\xff', composer.composed_bytes) composer = ComposerBinary() date_time = datetime.datetime.fromtimestamp(0xffffffff, datetime.timezone.utc) date_time.replace(tzinfo=None) composer.compose_timestamp(date_time) self.assertEqual(b'\x00\x00\x00\x00\xff\xff\xff\xff', composer.composed_bytes) composer = ComposerBinary() composer.compose_timestamp(None) self.assertEqual(b'\xff\xff\xff\xff\xff\xff\xff\xff', composer.composed_bytes) class TestComposerText(TestParsableBase): def test_compose_numeric(self): composer = ComposerText() composer.compose_numeric(1) self.assertEqual(composer.composed, b'1') composer.compose_numeric(2) self.assertEqual(composer.composed, b'12') self.assertEqual(composer.composed_length, 2) def test_compose_numeric_array(self): composer = ComposerText() composer.compose_numeric_array([1, 2], separator=',') self.assertEqual(composer.composed, b'1,2') self.assertEqual(composer.composed_length, 3) def test_compose_string(self): composer = ComposerText() composer.compose_string('abc') self.assertEqual(composer.composed, b'abc') self.assertEqual(composer.composed_length, 3) composer = ComposerText('utf-8') for index, char in enumerate(self._ALPHA_BETA_GAMMA): composer.compose_string(char) expected_composed = self._ALPHA_BETA_GAMMA[0:index + 1].encode('utf-8') self.assertEqual(composer.composed, expected_composed) self.assertEqual(composer.composed_length, (index + 1) * 2) composer = ComposerText() with self.assertRaises(InvalidValue) as context_manager: composer.compose_string(self._ALPHA_BETA_GAMMA) self.assertEqual(context_manager.exception.value, self._ALPHA_BETA_GAMMA) def test_compose_string_array(self): composer = ComposerText() composer.compose_string_array(['a', 'b', 'c'], '#') self.assertEqual(composer.composed, b'a#b#c') self.assertEqual(composer.composed_length, 5) composer = ComposerText('utf-8') composer.compose_string_array(list(self._ALPHA_BETA_GAMMA), '') self.assertEqual(composer.composed, self._ALPHA_BETA_GAMMA.encode('utf-8')) self.assertEqual(composer.composed_length, len(self._ALPHA_BETA_GAMMA) * 2) composer = ComposerText() composer.compose_string_array( [AlwaysTestStringComposer(), AlwaysTestStringComposer(), AlwaysTestStringComposer()], '#' ) self.assertEqual(composer.composed, b'test#test#test') self.assertEqual(composer.composed_length, len(AlwaysTestStringComposer().compose()) * 3 + 2) def test_compose_separator(self): composer = ComposerText() composer.compose_separator('#') self.assertEqual(composer.composed, b'#') composer.compose_separator('string') self.assertEqual(composer.composed, b'#string') composer.compose_separator('#') self.assertEqual(composer.composed, b'#string#') composer = ComposerText('utf-8') composer.compose_separator(self._ALPHA_BETA_GAMMA) self.assertEqual(composer.composed, self._ALPHA_BETA_GAMMA.encode('utf-8')) def test_compose_date_time(self): composer = ComposerText() composer.compose_date_time(datetime.datetime(2015, 10, 21, 7, 28), '%a, %d %b %Y %H:%M:%S GMT') self.assertEqual(composer.composed, b'Wed, 21 Oct 2015 07:28:00 GMT') def test_compose_time_delta(self): composer = ComposerText() composer.compose_time_delta(datetime.timedelta(1)) self.assertEqual(composer.composed, b'86400') def test_compose_bool(self): composer = ComposerText() composer.compose_bool(True) self.assertEqual(composer.composed, b'yes') composer = ComposerText() composer.compose_bool(False) self.assertEqual(composer.composed, b'no') cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/test_utils.py000066400000000000000000000026671524413560000271050ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptoparser.common.utils import bytes_from_hex_string, bytes_to_hex_string class TestBytesToHexString(unittest.TestCase): def test_separator(self): self.assertEqual(bytes_to_hex_string(b''), '') self.assertEqual(bytes_to_hex_string(b'\xde\xad\xbe\xef'), 'DEADBEEF') self.assertEqual(bytes_to_hex_string(b'\xde\xad\xbe\xef', separator=':'), 'DE:AD:BE:EF') def test_lowercase(self): self.assertEqual(bytes_to_hex_string(b''), '') self.assertEqual(bytes_to_hex_string(b'\xde\xad\xbe\xef'), 'DEADBEEF') self.assertEqual(bytes_to_hex_string(b'\xde\xad\xbe\xef', lowercase=True), 'deadbeef') class TestBytesFromHexString(unittest.TestCase): def test_error_odd_length_string(self): with self.assertRaises(ValueError) as context_manager: bytes_from_hex_string('0d:d') self.assertEqual(type(context_manager.exception), ValueError) def test_error_non_hex_string(self): with self.assertRaises(ValueError) as context_manager: bytes_from_hex_string('no:th:ex') self.assertEqual(type(context_manager.exception), ValueError) def test_separator(self): self.assertEqual(bytes_from_hex_string(''), b'') self.assertEqual(bytes_from_hex_string('DEADBEEF'), b'\xde\xad\xbe\xef') self.assertEqual(bytes_from_hex_string('DE:AD:BE:EF', separator=':'), b'\xde\xad\xbe\xef') cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/common/test_x509.py000066400000000000000000000037761524413560000264540ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import hashlib from test.common.classes import TestClasses from cryptodatahub.common.key import PublicKeyX509Base from cryptodatahub.common.entity import Entity from cryptoparser.common.x509 import SignedCertificateTimestampList class TestTlsPubKeys(TestClasses.TestKeyBase): def test_signed_certificate_timestamps(self): certificate = self._get_public_key_x509('rsa8192.badssl.com_root_ca.crt') self.assertEqual(certificate.signed_certificate_timestamps, SignedCertificateTimestampList([])) certificate = self._get_public_key_x509('rsa8192.badssl.com_certificate.crt') self.assertIn( Entity.GOOGLE, [sct.log.operator for sct in certificate.signed_certificate_timestamps] ) def test_ja4x(self): certificate = self._get_public_key_x509('rsa8192.badssl.com_certificate.crt') ja4x = certificate.ja4x issuer_oid_hexes, subject_oid_hexes, extension_oid_hexes = ja4x.fingerprint_raw.split('_') # the subject relative distinguished names contain the common name (OID 2.5.4.3 -> 550403) self.assertIn('550403', subject_oid_hexes.split(',')) self.assertIn('550403', issuer_oid_hexes.split(',')) self.assertTrue(extension_oid_hexes) # the fingerprint is the per-section truncated SHA-256 of the raw OID lists self.assertEqual(ja4x.fingerprint, '_'.join( hashlib.sha256(oid_hexes.encode('ascii')).hexdigest()[:12] for oid_hexes in ja4x.fingerprint_raw.split('_') )) self.assertEqual(certificate.ja4x, ja4x) def test_asdict(self): certificate = self._get_public_key_x509('rsa8192.badssl.com_certificate.crt') dict_result = certificate._asdict() self.assertIn( Entity.GOOGLE, [sct.log.operator for sct in dict_result.pop('signed_certificate_timestamps')], ) dict_result.pop('ja4x') self.assertEqual(PublicKeyX509Base._asdict(certificate), dict_result) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/dnsrec/000077500000000000000000000000001524413560000243075ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/dnsrec/__init__.py000066400000000000000000000000431524413560000264150ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/dnsrec/test_record.py000066400000000000000000000555641524413560000272150ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import base64 import collections import datetime import unittest from cryptodatahub.common.algorithm import Authentication, KeyExchange from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.parameter import ECParamWellKnown from cryptodatahub.common.key import ( PublicKey, PublicKeyParamsDsa, PublicKeyParamsEcdsa, PublicKeyParamsEddsa, PublicKeyParamsRsa, ) from cryptodatahub.dnsrec.algorithm import ( DnsSecAlgorithm, DnsSecDigestType, DnsRrType, SshFpAlgorithm, SshFpFingerprintType, ) from cryptoparser.common.exception import NotEnoughData from cryptoparser.dnsrec.record import ( DnsNameUncompressed, DnsRecordDnskey, DnsRecordDs, DnsRecordMx, DnsRecordRrsig, DnsRecordSshfp, DnsRecordTxt, DnsRrTypePrivate, DnsSecFlag, DnsSecProtocol, ) class TestDnsRecordDnskey(unittest.TestCase): def setUp(self): self.header_bytes = ( b'\x01\x00' + # flags: DNS_ZONE_KEY b'\x03' + # version b'' ) def test_error_inconsistent_algorithm(self): public_key_rsa = PublicKey.from_params(PublicKeyParamsRsa( public_exponent=2 ** 2048 - 1, modulus=2 ** 1024 - 1, )) with self.assertRaises(InvalidValue) as context_manager: DnsRecordDnskey( flags=[DnsSecFlag.DNS_ZONE_KEY], algorithm=DnsSecAlgorithm.DH, key=public_key_rsa, protocol=DnsSecProtocol.V3, ) self.assertEqual(context_manager.exception.value, KeyExchange.DH) with self.assertRaises(InvalidValue) as context_manager: DnsRecordDnskey( flags=[DnsSecFlag.DNS_ZONE_KEY], algorithm=DnsSecAlgorithm.ECCGOST, key=public_key_rsa, protocol=DnsSecProtocol.V3, ) self.assertEqual(context_manager.exception.value, Authentication.GOST_R3410_01) def test_error_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: DnsRecordDnskey.parse_exact_size(self.header_bytes) self.assertEqual( context_manager.exception.bytes_needed, DnsRecordDnskey.HEADER_SIZE - len(self.header_bytes) ) def test_key_tag(self): record_bytes = self.header_bytes + ( b'\x01' + # algorithm: RSAMD5 b'\x03' + # exponent_length: 3 b'\x01\x00\x01' + # exponent: 65537 124 * b'\x00' + # modulus b'\x11\x22\x44\x88' b'' ) dns_record = DnsRecordDnskey.parse_exact_size(record_bytes) self.assertEqual(dns_record.algorithm, DnsSecAlgorithm.RSAMD5) self.assertEqual(dns_record.key_tag, 0x2244) record_bytes = self.header_bytes + ( b'\x05' + # algorithm: RSASHA1 b'\x03' + # exponent_length: 3 b'\x01\x00\x01' + # exponent: 65537 128 * b'\xff' + # modulus b'' ) dns_record = DnsRecordDnskey.parse_exact_size(record_bytes) self.assertEqual(dns_record.algorithm, DnsSecAlgorithm.RSASHA1) self.assertEqual(dns_record.key_tag, 1799) def test_asdict(self): self.assertEqual( DnsRecordDnskey( flags=[DnsSecFlag.DNS_ZONE_KEY], algorithm=DnsSecAlgorithm.RSASHA1, key=PublicKey.from_params(PublicKeyParamsRsa( public_exponent=2 ** 2048 - 1, modulus=2 ** 1024 - 1, )), protocol=DnsSecProtocol.V3, )._asdict(), collections.OrderedDict([ ('key_tag', 1540), ('flags', [DnsSecFlag.DNS_ZONE_KEY]), ('algorithm', DnsSecAlgorithm.RSASHA1), ('key', PublicKey.from_params(PublicKeyParamsRsa( public_exponent=2 ** 2048 - 1, modulus=2 ** 1024 - 1, ))), ('protocol', DnsSecProtocol.V3), ]) ) def test_parse_rsa_key(self): record_bytes = self.header_bytes + ( b'\x05' + # algorithm: RSASHA1 b'\x03' + # exponent_length: 3 b'\x01\x00\x01' + # exponent: 65537 128 * b'\xff' + # modulus b'' ) dns_record = DnsRecordDnskey.parse_exact_size(record_bytes) self.assertEqual(dns_record.algorithm, DnsSecAlgorithm.RSASHA1) self.assertEqual(dns_record.key, PublicKey.from_params(PublicKeyParamsRsa( public_exponent=65537, modulus=2 ** 1024 - 1, ))) self.assertEqual(dns_record.compose(), record_bytes) record_bytes = self.header_bytes + ( b'\x05' + # algorithm: RSASHA1 b'\x00' + # exponent length extender mark b'\x01\x00' + # exponent_length: 256 256 * b'\xff' + # exponent 128 * b'\xff' + # modulus b'' ) dns_record = DnsRecordDnskey.parse_exact_size(record_bytes) self.assertEqual(dns_record.algorithm, DnsSecAlgorithm.RSASHA1) self.assertEqual(dns_record.key, PublicKey.from_params(PublicKeyParamsRsa( public_exponent=2 ** 2048 - 1, modulus=2 ** 1024 - 1, ))) self.assertEqual(dns_record.compose(), record_bytes) def test_parse_dsa_key(self): record_bytes = self.header_bytes + ( b'\x03' + # algorithm: DSA b'\x08' + # key size parameter 20 * b'\xff' + # q b'\x80' + (1024 // 8 - 1) * b'\x00' + # p b'\x40' + (1024 // 8 - 1) * b'\x00' + # g b'\x20' + (1024 // 8 - 1) * b'\x00' + # y b'' ) dns_record = DnsRecordDnskey.parse_exact_size(record_bytes) self.assertEqual(dns_record.algorithm, DnsSecAlgorithm.DSA) self.assertEqual(dns_record.key, PublicKey.from_params(PublicKeyParamsDsa( prime=2 ** 1023, generator=2 ** 1022, order=2 ** 160 - 1, public_key_value=2 ** 1021, ))) self.assertEqual(dns_record.compose(), record_bytes) def test_parse_ecdsa_key(self): record_bytes = self.header_bytes + ( b'\x0c' + # algorithm: ECCGOST b'\x80' + (256 // 8 - 1) * b'\x00' + # point_x b'\x40' + (256 // 8 - 1) * b'\x00' + # point_y b'' ) dns_record = DnsRecordDnskey.parse_exact_size(record_bytes) self.assertEqual(dns_record.algorithm, DnsSecAlgorithm.ECCGOST) self.assertEqual(dns_record.key, PublicKey.from_params(PublicKeyParamsEcdsa( key_parameter=ECParamWellKnown.GC256B, point_x=2 ** 255, point_y=2 ** 254, ))) self.assertEqual(dns_record.compose(), record_bytes) record_bytes = self.header_bytes + ( b'\x0d' + # algorithm: ECDSAP256SHA256 b'\x80' + (256 // 8 - 1) * b'\x00' + # point_x b'\x40' + (256 // 8 - 1) * b'\x00' + # point_y b'' ) dns_record = DnsRecordDnskey.parse_exact_size(record_bytes) self.assertEqual(dns_record.algorithm, DnsSecAlgorithm.ECDSAP256SHA256) self.assertEqual(dns_record.key, PublicKey.from_params(PublicKeyParamsEcdsa( key_parameter=ECParamWellKnown.SECP256K1, point_x=2 ** 255, point_y=2 ** 254, ))) self.assertEqual(dns_record.compose(), record_bytes) record_bytes = self.header_bytes + ( b'\x0e' + # algorithm: ECDSAP384SHA384 b'\x80' + (384 // 8 - 1) * b'\x00' + # point_x b'\x40' + (384 // 8 - 1) * b'\x00' + # point_y b'' ) dns_record = DnsRecordDnskey.parse_exact_size(record_bytes) self.assertEqual(dns_record.algorithm, DnsSecAlgorithm.ECDSAP384SHA384) self.assertEqual(dns_record.key, PublicKey.from_params(PublicKeyParamsEcdsa( key_parameter=ECParamWellKnown.SECP384R1, point_x=2 ** 383, point_y=2 ** 382, ))) self.assertEqual(dns_record.compose(), record_bytes) def test_parse_eddsa_key(self): record_bytes = self.header_bytes + ( b'\x0f' + # algorithm: ED25519 32 * b'\xff' + # key_data b'' ) dns_record = DnsRecordDnskey.parse_exact_size(record_bytes) self.assertEqual(dns_record.algorithm, DnsSecAlgorithm.ED25519) self.assertEqual(dns_record.key, PublicKey.from_params(PublicKeyParamsEddsa( key_parameter=ECParamWellKnown.CURVE25519, key_data=32 * b'\xff', ))) self.assertEqual(dns_record.compose(), record_bytes) record_bytes = self.header_bytes + ( b'\x10' + # algorithm: ED448 56 * b'\xff' + # key_data b'' ) dns_record = DnsRecordDnskey.parse_exact_size(record_bytes) self.assertEqual(dns_record.algorithm, DnsSecAlgorithm.ED448) self.assertEqual(dns_record.key, PublicKey.from_params(PublicKeyParamsEddsa( key_parameter=ECParamWellKnown.CURVE448, key_data=56 * b'\xff', ))) self.assertEqual(dns_record.compose(), record_bytes) def test_real(self): # RFC 4034 Section 5.4 public_key_bytes = base64.b64decode( 'AQOeiiR0GOMYkDshWoSKz9Xz' 'fwJr1AYtsmx3TGkJaNXVbfi/' '2pHm822aJ5iI9BMzNXxeYCmZ' 'DRD99WYwYqUSdjMmmAphXdvx' 'egXd/M5+X7OrzKBaMbCVdFLU' 'Uh6DhweJBjEVv5f2wwjM9Xzc' 'nOf+EPbtG9DMBmADjFDc2w/r' 'ljwvFw==' ) public_key = DnsRecordDnskey.parse_key(public_key_bytes, DnsSecAlgorithm.RSASHA1) self.assertEqual(DnsRecordDnskey.compose_key(public_key), public_key_bytes) dns_record = DnsRecordDnskey( flags=[DnsSecFlag.DNS_ZONE_KEY], algorithm=DnsSecAlgorithm.RSASHA1, key=public_key, protocol=DnsSecProtocol.V3, ) self.assertEqual(dns_record.key_tag, 60485) # RFC 4034 Section 2.3 public_key_bytes = base64.b64decode( 'AQPSKmynfzW4kyBv015MUG2DeIQ3' 'Cbl+BBZH4b/0PY1kxkmvHjcZc8no' 'kfzj31GajIQKY+5CptLr3buXA10h' 'WqTkF7H6RfoRqXQeogmMHfpftf6z' 'Mv1LyBUgia7za6ZEzOJBOztyvhjL' '742iU/TpPSEDhm2SNKLijfUppn1U' 'aNvv4w==' ) public_key = DnsRecordDnskey.parse_key(public_key_bytes, DnsSecAlgorithm.RSASHA1) self.assertEqual(DnsRecordDnskey.compose_key(public_key), public_key_bytes) dns_record = DnsRecordDnskey( flags=[DnsSecFlag.DNS_ZONE_KEY], algorithm=DnsSecAlgorithm.RSASHA1, key=public_key, protocol=DnsSecProtocol.V3, ) self.assertEqual(dns_record.key_tag, 2642) # RFC 5702 Section 6.1 public_key_bytes = base64.b64decode( 'AwEAAcFcGsaxxdgiuuGmCkVI' 'my4h99CqT7jwY3pexPGcnUFtR2Fh36BponcwtkZ4cAgtvd4Qs8P' 'kxUdp6p/DlUmObdk=' ) public_key = DnsRecordDnskey.parse_key(public_key_bytes, DnsSecAlgorithm.RSASHA256) self.assertEqual(DnsRecordDnskey.compose_key(public_key), public_key_bytes) dns_record = DnsRecordDnskey( flags=[DnsSecFlag.DNS_ZONE_KEY], algorithm=DnsSecAlgorithm.RSASHA256, key=public_key, protocol=DnsSecProtocol.V3, ) self.assertEqual(dns_record.key_tag, 9033) # RFC 5702 Section 6.2 public_key_bytes = base64.b64decode( 'AwEAAdHoNTOW+et86KuJOWRD' 'p1pndvwb6Y83nSVXXyLA3DLroROUkN6X0O6pnWnjJQujX/AyhqFD' 'xj13tOnD9u/1kTg7cV6rklMrZDtJCQ5PCl/D7QNPsgVsMu1J2Q8g' 'pMpztNFLpPBz1bWXjDtaR7ZQBlZ3PFY12ZTSncorffcGmhOL' ) public_key = DnsRecordDnskey.parse_key(public_key_bytes, DnsSecAlgorithm.RSASHA512) self.assertEqual(DnsRecordDnskey.compose_key(public_key), public_key_bytes) dns_record = DnsRecordDnskey( flags=[DnsSecFlag.DNS_ZONE_KEY], algorithm=DnsSecAlgorithm.RSASHA512, key=public_key, protocol=DnsSecProtocol.V3, ) self.assertEqual(dns_record.key_tag, 3740) # RFC 5933 Section 2.2 public_key_bytes = base64.b64decode( 'aRS/DcPWGQj2wVJydT8EcAVoC0kXn5pDVm2I' 'MvDDPXeD32dsSKcmq8KNVzigjL4OXZTV+t/6' 'w4X1gpNrZiC01g==' ) public_key = DnsRecordDnskey.parse_key(public_key_bytes, DnsSecAlgorithm.ECCGOST) self.assertEqual(DnsRecordDnskey.compose_key(public_key), public_key_bytes) dns_record = DnsRecordDnskey( flags=[DnsSecFlag.DNS_ZONE_KEY], algorithm=DnsSecAlgorithm.ECCGOST, key=public_key, protocol=DnsSecProtocol.V3, ) self.assertEqual(dns_record.key_tag, 59732) class TestDnsRecordDs(unittest.TestCase): def setUp(self): self.record_bytes = bytes( b'\x00\x01' + # key_tag: 1 b'\x01' + # algorithm: RSAMD5 b'\x02' + # digest_type: SHA_256 32 * b'\xff' + # digest b'' ) self.record = DnsRecordDs( key_tag=1, algorithm=DnsSecAlgorithm.RSAMD5, digest_type=DnsSecDigestType.SHA_256, digest=32 * b'\xff', ) def test_error_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: DnsRecordDs.parse_exact_size(b'\x00') self.assertEqual( context_manager.exception.bytes_needed, DnsRecordDs.HEADER_SIZE - 1 ) def test_parse(self): self.assertEqual(DnsRecordDs.parse_exact_size(self.record_bytes), self.record) def test_compose(self): self.assertEqual(self.record.compose(), self.record_bytes) class TestDnsNameUncompressed(unittest.TestCase): def setUp(self): self.label_empty_bytes = b'\x00' self.label_empty_name = DnsNameUncompressed([]) self.label_single_bytes = b'\x01a\x00' self.label_single_name = DnsNameUncompressed(['a']) self.label_multiple_bytes = b'\x01a\x02bb\x03ccc\x00' self.label_multiple_name = DnsNameUncompressed(['a', 'bb', 'ccc']) def test_error_convert_invalid_value(self): with self.assertRaises(InvalidValue) as context_manager: DnsNameUncompressed.convert(None) self.assertEqual(context_manager.exception.value, None) def test_parse(self): self.assertEqual(DnsNameUncompressed.parse_exact_size(self.label_empty_bytes), self.label_empty_name) self.assertEqual(DnsNameUncompressed.parse_exact_size(self.label_single_bytes), self.label_single_name) self.assertEqual(DnsNameUncompressed.parse_exact_size(self.label_multiple_bytes), self.label_multiple_name) def test_compose(self): self.assertEqual(self.label_empty_name.compose(), self.label_empty_bytes) self.assertEqual(self.label_single_name.compose(), self.label_single_bytes) self.assertEqual(self.label_multiple_name.compose(), self.label_multiple_bytes) def test_convert(self): self.assertEqual(DnsNameUncompressed.convert(self.label_empty_name), self.label_empty_name) self.assertEqual(DnsNameUncompressed.convert(self.label_single_name), self.label_single_name) self.assertEqual(DnsNameUncompressed.convert(self.label_multiple_name), self.label_multiple_name) self.assertEqual(DnsNameUncompressed.convert(''), self.label_empty_name) self.assertEqual(DnsNameUncompressed.convert('a'), self.label_single_name) self.assertEqual(DnsNameUncompressed.convert('a.bb.ccc'), self.label_multiple_name) def test_str(self): self.assertEqual(str(self.label_empty_name), '') self.assertEqual(str(self.label_single_name), 'a') self.assertEqual(str(self.label_multiple_name), 'a.bb.ccc') def test_as_markdown(self): self.assertEqual(self.label_empty_name.as_markdown(), '') self.assertEqual(self.label_single_name.as_markdown(), 'a') self.assertEqual(self.label_multiple_name.as_markdown(), 'a.bb.ccc') class TestDnsRecordRrsig(unittest.TestCase): def setUp(self): self.record_bytes = bytes( b'\x00\x01' + # type_covered: A b'\x01' + # algorithm: RSAMD5 b'\x03' + # labels b'\x00\x00\x0e\x10' + # original_ttl: 3600 b'\x00\x00\x00\x01' + # signature_expiration b'\x00\x00\x00\x02' + # signature_inception b'\xab\xcd' + # key_tag b'\x06signer\x00' + # signers_name 32 * b'\xff' + # signature b'' ) self.record = DnsRecordRrsig( type_covered=DnsRrType.A, algorithm=DnsSecAlgorithm.RSAMD5, labels=3, original_ttl=3600, signature_expiration=datetime.datetime(1970, 1, 1, 0, 0, 1, tzinfo=datetime.timezone.utc), signature_inception=datetime.datetime(1970, 1, 1, 0, 0, 2, tzinfo=datetime.timezone.utc), key_tag=0xabcd, signers_name='signer', signature=32 * b'\xff', ) self.record_private_type_bytes = bytes( b'\xff\xfe' + # type_covered: A b'\x01' + # algorithm: RSAMD5 b'\x03' + # labels b'\x00\x00\x0e\x10' + # original_ttl: 3600 b'\x00\x00\x00\x01' + # signature_expiration b'\x00\x00\x00\x02' + # signature_inception b'\xab\xcd' + # key_tag b'\x06signer\x00' + # signers_name 32 * b'\xff' + # signature b'' ) self.record_private_types = DnsRecordRrsig( type_covered=DnsRrTypePrivate(0xfffe), algorithm=DnsSecAlgorithm.RSAMD5, labels=3, original_ttl=3600, signature_expiration=datetime.datetime(1970, 1, 1, 0, 0, 1, tzinfo=datetime.timezone.utc), signature_inception=datetime.datetime(1970, 1, 1, 0, 0, 2, tzinfo=datetime.timezone.utc), key_tag=0xabcd, signers_name='signer', signature=32 * b'\xff', ) def test_error_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: DnsRecordRrsig.parse_exact_size(b'\x00') self.assertEqual( context_manager.exception.bytes_needed, DnsRecordRrsig.HEADER_SIZE - 1 ) def test_parse(self): self.assertEqual(DnsRecordRrsig.parse_exact_size(self.record_bytes), self.record) self.assertEqual(DnsRecordRrsig.parse_exact_size(self.record_private_type_bytes), self.record_private_types) def test_compose(self): self.assertEqual(self.record.compose(), self.record_bytes) self.assertEqual(self.record_private_types.compose(), self.record_private_type_bytes) class TestDnsRecordMx(unittest.TestCase): def setUp(self): self.record_bytes = bytes( b'\x00\x01' + # priority: 1 b'\x08exchange\x00' + # exchange b'' ) self.record = DnsRecordMx( priority=1, exchange='exchange', ) def test_error_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: DnsRecordMx.parse_exact_size(b'\x00') self.assertEqual( context_manager.exception.bytes_needed, DnsRecordMx.HEADER_SIZE - 1 ) def test_parse(self): self.assertEqual(DnsRecordMx.parse_exact_size(self.record_bytes), self.record) def test_compose(self): self.assertEqual(self.record.compose(), self.record_bytes) class TestDnsRecordTxt(unittest.TestCase): def setUp(self): self.record_bytes_single = bytes( b'\x05' + # length: 5 b'value' + b'' ) self.record_single = DnsRecordTxt(value='value') self.record_bytes_multiple = bytes( b'\x06' + # length: 6 b'value1' + b'\x06' + # length: 6 b'value2' + b'' ) self.record_multiple = DnsRecordTxt(value='value1value2') def test_error_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: DnsRecordTxt.parse_exact_size(b'') self.assertEqual( context_manager.exception.bytes_needed, DnsRecordTxt.HEADER_SIZE ) def test_parse(self): self.assertEqual(DnsRecordTxt.parse_exact_size(self.record_bytes_single), self.record_single) self.assertEqual(DnsRecordTxt.parse_exact_size(self.record_bytes_multiple), self.record_multiple) def test_compose(self): self.assertEqual(self.record_single.compose(), self.record_bytes_single) class TestDnsRecordSshfp(unittest.TestCase): def setUp(self): self.record_rsa_sha1_bytes = bytes( b'\x01' + # algorithm: RSA b'\x01' + # fingerprint type: SHA-1 b'\xde\xad\xbe\xef' * 5 # fingerprint: 20 bytes ) self.record_rsa_sha1 = DnsRecordSshfp( algorithm=SshFpAlgorithm.RSA, fingerprint_type=SshFpFingerprintType.SHA1, fingerprint=b'\xde\xad\xbe\xef' * 5, ) self.record_ecdsa_sha256_bytes = bytes( b'\x03' + # algorithm: ECDSA b'\x02' + # fingerprint type: SHA-256 b'\xab\xcd\xef\x01' * 8 # fingerprint: 32 bytes ) self.record_ecdsa_sha256 = DnsRecordSshfp( algorithm=SshFpAlgorithm.ECDSA, fingerprint_type=SshFpFingerprintType.SHA2_256, fingerprint=b'\xab\xcd\xef\x01' * 8, ) def test_error_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: DnsRecordSshfp.parse_exact_size(b'\x01') self.assertEqual( context_manager.exception.bytes_needed, DnsRecordSshfp.HEADER_SIZE - 1 ) def test_parse(self): self.assertEqual( DnsRecordSshfp.parse_exact_size(self.record_rsa_sha1_bytes), self.record_rsa_sha1 ) self.assertEqual( DnsRecordSshfp.parse_exact_size(self.record_ecdsa_sha256_bytes), self.record_ecdsa_sha256 ) def test_compose(self): self.assertEqual(self.record_rsa_sha1.compose(), self.record_rsa_sha1_bytes) self.assertEqual(self.record_ecdsa_sha256.compose(), self.record_ecdsa_sha256_bytes) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/dnsrec/test_txt.py000066400000000000000000000276441524413560000265540ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import collections import unittest import ipaddress from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import InvalidType from cryptoparser.common.field import NameValuePairListSemicolonSeparated from cryptoparser.dnsrec.txt import ( DmarcAlignment, DmarcFailureReportingFormat, DmarcFailureReportingOption, DmarcPolicyOption, DmarcPolicyVersion, DmarcReportingInterval, DnsRecordTxtValueDmarc, DnsRecordTxtValueMtaSts, DnsRecordTxtValueSpf, DnsRecordTxtValueSpfDirectiveA, DnsRecordTxtValueSpfDirectiveAll, DnsRecordTxtValueSpfDirectiveExists, DnsRecordTxtValueSpfDirectiveInclude, DnsRecordTxtValueSpfDirectiveIp4, DnsRecordTxtValueSpfDirectiveIp6, DnsRecordTxtValueSpfDirectiveMx, DnsRecordTxtValueSpfDirectivePtr, DnsRecordTxtValueSpfModifierExplanation, DnsRecordTxtValueSpfModifierRedirect, DnsRecordTxtValueSpfModifierUnknown, DnsRecordTxtValueTlsRpt, MtaStsPolicyVersion, SpfQualifier, SpfVersion, TlsRptVersion, ) class TestDnsRecordTxtValueDmarcReportingInterval(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: DmarcReportingInterval.parse_exact_size(b'ri=5000000000') self.assertEqual(context_manager.exception.value, 5000000000) class TestDnsRecordDmarc(unittest.TestCase): _record_minimal = DnsRecordTxtValueDmarc(DmarcPolicyVersion.DMARC1, DmarcPolicyOption.NONE) _record_minimal_bytes = b'v=DMARC1;p=none' _record_full = DnsRecordTxtValueDmarc( version=DmarcPolicyVersion.DMARC1, policy=DmarcPolicyOption.NONE, alignment_dkim=DmarcAlignment.STRICT, alignment_aspf=DmarcAlignment.STRICT, failure_option=DmarcFailureReportingOption.ANY_FAILURE, percent=99, reporting_url_aggregated='mailto:dmarc-report@example.com', reporting_url_failure='https://example.com/dmarc/report/failure', reporting_format=DmarcFailureReportingFormat.AUTHENTICATION_FAILURE_REPORTING_FORMAT, reporting_interval=3600, subdomain_policy=DmarcPolicyOption.NONE, ) _record_full_bytes = b'; '.join([ b'v=DMARC1', b'p=none', b'adkim=s', b'aspf=s', b'fo=1', b'pct=99', b'rua=mailto:dmarc-report@example.com', b'ruf=https://example.com/dmarc/report/failure', b'rf=afrf', b'ri=3600', b'sp=none', ]) def test_parse(self): self.assertEqual(DnsRecordTxtValueDmarc.parse_exact_size(self._record_minimal_bytes), self._record_minimal) self.assertEqual(DnsRecordTxtValueDmarc.parse_exact_size(self._record_full_bytes), self._record_full) def test_compose(self): self.assertEqual(self._record_full.compose(), self._record_full_bytes) class TestDnsRecordMtaSts(unittest.TestCase): _record_minimal = DnsRecordTxtValueMtaSts(MtaStsPolicyVersion.STSV1, '20160831085700Z') _record_minimal_bytes = b'v=STSv1; id=20160831085700Z' _record_full = DnsRecordTxtValueMtaSts( version=MtaStsPolicyVersion.STSV1, identifier='20160831085700Z', extensions=NameValuePairListSemicolonSeparated( collections.OrderedDict([('extension_name', 'extension_value')]) ), ) _record_full_bytes = b'v=STSv1; id=20160831085700Z; extension_name=extension_value' def test_parse(self): self.assertEqual(DnsRecordTxtValueMtaSts.parse_exact_size(self._record_minimal_bytes), self._record_minimal) self.assertEqual(DnsRecordTxtValueMtaSts.parse_exact_size(self._record_full_bytes), self._record_full) def test_compose(self): self.assertEqual(self._record_minimal.compose(), self._record_minimal_bytes) self.assertEqual(self._record_full.compose(), self._record_full_bytes) class TestDnsRecordTxtValueTlsRpt(unittest.TestCase): _record_minimal = DnsRecordTxtValueTlsRpt(TlsRptVersion.TLSRPTV1, 'https://example.com/tlsrpt/report/failure') _record_minimal_bytes = b'v=TLSRPTv1; rua=https://example.com/tlsrpt/report/failure' _record_full = DnsRecordTxtValueTlsRpt( version=TlsRptVersion.TLSRPTV1, reporting_url_aggregated='mailto:tls-report@example.com', extensions=NameValuePairListSemicolonSeparated( collections.OrderedDict([('extension_name', 'extension_value')]) ), ) _record_full_bytes = b'v=TLSRPTv1; rua=mailto:tls-report@example.com; extension_name=extension_value' def test_parse(self): self.assertEqual(DnsRecordTxtValueTlsRpt.parse_exact_size(self._record_minimal_bytes), self._record_minimal) self.assertEqual(DnsRecordTxtValueTlsRpt.parse_exact_size(self._record_full_bytes), self._record_full) def test_compose(self): self.assertEqual(self._record_minimal.compose(), self._record_minimal_bytes) self.assertEqual(self._record_full.compose(), self._record_full_bytes) class TestDnsRecordTxtValueSpfDirective(unittest.TestCase): def test_parse_domain_optional(self): directive = DnsRecordTxtValueSpfDirectivePtr.parse_exact_size(b'ptr') self.assertEqual(directive.domain, None) self.assertEqual(directive.compose(), b'ptr') directive = DnsRecordTxtValueSpfDirectivePtr.parse_exact_size(b'ptr:domain') self.assertEqual(directive.domain.value, 'domain') self.assertEqual(directive.compose(), b'ptr:domain') def test_parse_domain_required(self): directive = DnsRecordTxtValueSpfDirectiveInclude.parse_exact_size(b'include:domain') self.assertEqual(directive.domain.value, 'domain') self.assertEqual(directive.compose(), b'include:domain') def test_parse_single_cidr_length(self): directive = DnsRecordTxtValueSpfDirectiveIp4.parse_exact_size(b'ip4:1.1.1.1') self.assertEqual(directive.ipv4_network, ipaddress.IPv4Network('1.1.1.1/32')) self.assertEqual(directive.compose(), b'ip4:1.1.1.1') directive = DnsRecordTxtValueSpfDirectiveIp4.parse_exact_size(b'ip4:1.1.1.0/24') self.assertEqual(directive.ipv4_network, ipaddress.IPv4Network('1.1.1.0/24')) self.assertEqual(directive.compose(), b'ip4:1.1.1.0/24') directive = DnsRecordTxtValueSpfDirectiveIp6.parse_exact_size(b'ip6:::1') self.assertEqual(directive.ipv6_network, ipaddress.IPv6Network('::1/128')) self.assertEqual(directive.compose(), b'ip6:::1') directive = DnsRecordTxtValueSpfDirectiveIp6.parse_exact_size(b'ip6:::1:0/120') self.assertEqual(directive.ipv6_network, ipaddress.IPv6Network('::1:0/120')) self.assertEqual(directive.compose(), b'ip6:::1:0/120') with self.assertRaises(InvalidValue) as context_manager: DnsRecordTxtValueSpfDirectiveMx('example.com', ipv4_cidr_length=-1) self.assertEqual(context_manager.exception.value, -1) with self.assertRaises(InvalidValue) as context_manager: DnsRecordTxtValueSpfDirectiveMx('example.com', ipv4_cidr_length=33) self.assertEqual(context_manager.exception.value, 33) with self.assertRaises(InvalidValue) as context_manager: DnsRecordTxtValueSpfDirectiveMx('example.com', ipv6_cidr_length=-1) self.assertEqual(context_manager.exception.value, -1) with self.assertRaises(InvalidValue) as context_manager: DnsRecordTxtValueSpfDirectiveMx('example.com', ipv6_cidr_length=129) self.assertEqual(context_manager.exception.value, 129) def test_parse_dual_cidr_length(self): directive = DnsRecordTxtValueSpfDirectiveMx.parse_exact_size(b'mx:example.com') self.assertEqual(directive.domain.value, 'example.com') self.assertEqual(directive.compose(), b'mx:example.com') directive = DnsRecordTxtValueSpfDirectiveMx.parse_exact_size(b'mx:example.com/24') self.assertEqual(directive.domain.value, 'example.com') self.assertEqual(directive.ipv4_cidr_length, 24) self.assertEqual(directive.ipv6_cidr_length, None) self.assertEqual(directive.compose(), b'mx:example.com/24') directive = DnsRecordTxtValueSpfDirectiveMx.parse_exact_size(b'mx:example.com/24/64') self.assertEqual(directive.domain.value, 'example.com') self.assertEqual(directive.ipv4_cidr_length, 24) self.assertEqual(directive.ipv6_cidr_length, 64) self.assertEqual(directive.compose(), b'mx:example.com/24/64') class TestDnsRecordTxtValueSpf(unittest.TestCase): _record_minimal = DnsRecordTxtValueSpf(version=SpfVersion.SPF1, terms=[]) _record_minimal_bytes = b'v=spf1' _record_full = DnsRecordTxtValueSpf( version=SpfVersion.SPF1, terms=[ DnsRecordTxtValueSpfDirectiveExists(domain='%{ir}.%{l1r+-}._spf.%{d}'), DnsRecordTxtValueSpfDirectiveA('example.com', 32, 128), DnsRecordTxtValueSpfDirectiveMx('example.com', 32, 128), DnsRecordTxtValueSpfDirectivePtr('example.com'), DnsRecordTxtValueSpfDirectiveIp4('1.2.3.4'), DnsRecordTxtValueSpfDirectiveIp6('::1:2:3:4'), DnsRecordTxtValueSpfDirectiveInclude('_spf.example.com'), DnsRecordTxtValueSpfModifierRedirect('redirect.example.com'), DnsRecordTxtValueSpfModifierExplanation('exp.example.com'), DnsRecordTxtValueSpfModifierUnknown('modifier_key', 'modifier_value'), DnsRecordTxtValueSpfDirectiveAll(SpfQualifier.FAIL), ], ) _record_full_bytes = b' '.join([ b'v=spf1', b'exists:%{ir}.%{l1r+-}._spf.%{d}', b'a:example.com/32/128', b'mx:example.com/32/128', b'ptr:example.com', b'ip4:1.2.3.4', b'ip6:::1:2:3:4', b'include:_spf.example.com', b'redirect=redirect.example.com', b'exp=exp.example.com', b'modifier_key=modifier_value', b'-all', ]) def test_error_non_spf(self): with self.assertRaises(InvalidType): DnsRecordTxtValueSpf.parse_exact_size(b'v=STSv1') def test_parse(self): self.assertEqual(DnsRecordTxtValueSpf.parse_exact_size(self._record_minimal_bytes), self._record_minimal) self.assertEqual(DnsRecordTxtValueSpf.parse_exact_size(self._record_full_bytes), self._record_full) def test_compose(self): self.assertEqual(self._record_minimal.compose(), self._record_minimal_bytes) self.assertEqual(self._record_full.compose(), self._record_full_bytes) def test_as_markdown(self): self.assertEqual(self._record_minimal.as_markdown(), '\n'.join([ '* Version: SPF1', '* Terms: -', '', ])) self.assertEqual(self._record_full.as_markdown(), '\n'.join([ '* Version: SPF1', '* Terms:', ' * Exists:', ' * Domain: %{ir}.%{l1r+-}._spf.%{d}', ' * Qualifier: n/a', ' * A/AAAA records:', ' * Domain: example.com', ' * Ipv4 Cidr Length: 32', ' * Ipv6 Cidr Length: 128', ' * Qualifier: n/a', ' * MX records:', ' * Domain: example.com', ' * Ipv4 Cidr Length: 32', ' * Ipv6 Cidr Length: 128', ' * Qualifier: n/a', ' * PTR records:', ' * Domain: example.com', ' * Qualifier: n/a', ' * IPv4 records:', ' * Ipv4 Network: 1.2.3.4/32', ' * Qualifier: n/a', ' * IPv6 records:', ' * Ipv6 Network: ::1:2:3:4/128', ' * Qualifier: n/a', ' * Include:', ' * Domain: _spf.example.com', ' * Qualifier: n/a', ' * Redirect: redirect.example.com', ' * Explanation: exp.example.com', ' * Modifier Key: modifier_value', ' * All:', ' * Qualifier: Fail', '', ])) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/httpx/000077500000000000000000000000001524413560000242005ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/httpx/__init__.py000066400000000000000000000000431524413560000263060ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/httpx/classes.py000066400000000000000000000032431524413560000262110ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest class TestCasesBasesHttpHeader: class MinimalHeader(unittest.TestCase): _header_minimal = None _header_minimal_bytes = None _header_minimal_markdown = None def test_parse_minimal(self): parsed_header = self._header_minimal.parse_exact_size(self._header_minimal_bytes) self.assertEqual(parsed_header, self._header_minimal) def test_compose_minimal(self): self.assertEqual(self._header_minimal.compose(), self._header_minimal_bytes) def test_markdown(self): self.assertEqual(self._header_minimal.as_markdown(), self._header_minimal_markdown) class FullHeaderBase(unittest.TestCase): _header_full = None _header_full_bytes = None class FullHeader(FullHeaderBase): def test_parse_full(self): parsed_header = self._header_full.parse_exact_size(self._header_full_bytes) self.assertEqual(parsed_header, self._header_full) def test_compose_full(self): self.assertEqual(self._header_full.compose(), self._header_full_bytes) class CaseInsensitiveHeader(FullHeaderBase): _header_full_upper_case_bytes = None _header_full_lower_case_bytes = None def test_parse_upper_case(self): parsed_header = self._header_full.parse_exact_size(self._header_full_upper_case_bytes) self.assertEqual(parsed_header, self._header_full) def test_parse_lower_case(self): parsed_header = self._header_full.parse_exact_size(self._header_full_lower_case_bytes) self.assertEqual(parsed_header, self._header_full) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/httpx/test_header.py000066400000000000000000001206301524413560000270430ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines # -*- coding: utf-8 -*- import os import unittest import datetime from cryptodatahub.common.algorithm import Hash from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import InvalidType from cryptoparser.httpx.header import ( ContentSecurityPolicyDirectiveBaseUri, ContentSecurityPolicyDirectiveBlockAllMixedContent, ContentSecurityPolicyDirectiveChildSrc, ContentSecurityPolicyDirectiveConnectSrc, ContentSecurityPolicyDirectiveDefaultSrc, ContentSecurityPolicyDirectiveFontSrc, ContentSecurityPolicyDirectiveFormAction, ContentSecurityPolicyDirectiveFrameAncestors, ContentSecurityPolicyDirectiveFrameSrc, ContentSecurityPolicyDirectiveImgSrc, ContentSecurityPolicyDirectiveManifestSrc, ContentSecurityPolicyDirectiveMediaSrc, ContentSecurityPolicyDirectiveObjectSrc, ContentSecurityPolicyDirectivePluginTypes, ContentSecurityPolicyDirectivePrefetchSrc, ContentSecurityPolicyDirectiveReferrer, ContentSecurityPolicyDirectiveReportTo, ContentSecurityPolicyDirectiveReportUri, ContentSecurityPolicyDirectiveRequireTrustedTypesFor, ContentSecurityPolicyDirectiveSandbox, ContentSecurityPolicyDirectiveScriptSrc, ContentSecurityPolicyDirectiveScriptSrcAttr, ContentSecurityPolicyDirectiveScriptSrcElem, ContentSecurityPolicyDirectiveStyleSrc, ContentSecurityPolicyDirectiveStyleSrcAttr, ContentSecurityPolicyDirectiveStyleSrcElem, ContentSecurityPolicyDirectiveUpgradeInsecureRequests, ContentSecurityPolicyDirectiveWebrtc, ContentSecurityPolicyDirectiveWorkerSrc, ContentSecurityPolicyReferrerPolicy, ContentSecurityPolicySourceHash, ContentSecurityPolicySourceHost, ContentSecurityPolicySourceKeyword, ContentSecurityPolicySourceNonce, ContentSecurityPolicySourceScheme, ContentSecurityPolicyTrustedTypeSinkGroup, ContentSecurityPolicyWebRtcType, FieldValueMimeType, HttpHeaderFields, HttpHeaderFieldAge, HttpHeaderFieldCacheControlResponse, HttpHeaderFieldContentType, HttpHeaderFieldContentSecurityPolicy, HttpHeaderFieldContentSecurityPolicyReportOnly, HttpHeaderFieldDate, HttpHeaderFieldETag, HttpHeaderFieldExpectCT, HttpHeaderFieldExpectStaple, HttpHeaderFieldExpires, HttpHeaderFieldLastModified, HttpHeaderFieldName, HttpHeaderFieldNetworkErrorLogging, HttpHeaderFieldPragma, HttpHeaderFieldPublicKeyPinning, HttpHeaderFieldReferrerPolicy, HttpHeaderFieldServer, HttpHeaderFieldSetCookie, HttpHeaderFieldSTS, HttpHeaderFieldUnparsed, HttpHeaderFieldValueCacheControlResponse, HttpHeaderFieldValueContentSecurityPolicy, HttpHeaderFieldValueContentType, HttpHeaderFieldValueExpectCT, HttpHeaderFieldValueExpectStaple, HttpHeaderFieldValueNetworkErrorLogging, HttpHeaderFieldValuePragma, HttpHeaderFieldValuePublicKeyPinning, HttpHeaderFieldValueReferrerPolicy, HttpHeaderFieldValueSetCookie, HttpHeaderFieldValueSTS, HttpHeaderFieldValueXContentTypeOptions, HttpHeaderFieldValueXFrameOptions, HttpHeaderFieldValueXXSSProtection, HttpHeaderFieldXContentSecurityPolicy, HttpHeaderFieldXContentTypeOptions, HttpHeaderFieldXFrameOptions, HttpHeaderFieldXXSSProtection, HttpHeaderReferrerPolicy, HttpHeaderPragma, HttpHeaderSetCookieComponentSameSite, HttpHeaderXContentTypeOptions, HttpHeaderXFrameOptions, HttpHeaderXXSSProtectionMode, HttpHeaderXXSSProtectionState, MimeTypeRegistry, ) from .classes import TestCasesBasesHttpHeader class TestHttpHeaderFieldValueCacheControlResponse( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader, TestCasesBasesHttpHeader.CaseInsensitiveHeader): _header_minimal = HttpHeaderFieldValueCacheControlResponse() _header_minimal_bytes = b'' _header_minimal_markdown = '' _header_full = HttpHeaderFieldValueCacheControlResponse( max_age=datetime.timedelta(seconds=1), s_maxage=datetime.timedelta(seconds=2), must_revalidate=True, no_cache=True, no_store=True, public=True, private=True, no_transform=True, ) _header_full_bytes = b'max-age=1, s-maxage=2, must-revalidate, no-cache, no-store, public, private, no-transform' _header_full_upper_case_bytes = b', '.join([ b'MAX-AGE=1', b'S-MAXAGE=2', b'MUST-REVALIDATE', b'NO-CACHE', b'NO-STORE', b'PUBLIC', b'PRIVATE', b'NO-TRANSFORM', ]) _header_full_lower_case_bytes = b', '.join([ b'max-age=1', b's-maxage=2', b'must-revalidate', b'no-cache', b'no-store', b'public', b'private', b'no-transform', ]) _header_minimal_markdown = os.linesep.join([ '* Max Age: n/a', '* S Maxage: n/a', '* Must Revalidate: no', '* Proxy Revalidate: no', '* No Cache: no', '* No Store: no', '* Public: no', '* Private: no', '* No Transform: no', '', ]) class TestContentSecurityPolicySourceHash(unittest.TestCase): def test_error_wrong_prefix(self): with self.assertRaises(InvalidType): ContentSecurityPolicySourceHash.parse_exact_size(b'notavalidhashalgorithm') def test_parse(self): self.assertEqual( ContentSecurityPolicySourceHash.parse_exact_size(b'sha256-bGlnaHQgd29yay4='), ContentSecurityPolicySourceHash(Hash.SHA2_256, bytearray(b'light work.')) ) def test_compose(self): self.assertEqual( ContentSecurityPolicySourceHash(Hash.SHA2_256, bytearray(b'light work.')).compose(), b'sha256-bGlnaHQgd29yay4=' ) class TestContentSecurityPolicySourceHost(unittest.TestCase): def test_parse(self): self.assertEqual( ContentSecurityPolicySourceHost.parse_exact_size(b'http://example.com'), ContentSecurityPolicySourceHost('http://example.com') ) def test_compose(self): self.assertEqual( ContentSecurityPolicySourceHost('http://example.com').compose(), b'http://example.com' ) class TestContentSecurityPolicySourceNonce(unittest.TestCase): def test_error_wrong_prefix(self): with self.assertRaises(InvalidType): ContentSecurityPolicySourceNonce.parse_exact_size(b'bGlnaHQgd29yay4=') def test_parse(self): self.assertEqual( ContentSecurityPolicySourceNonce.parse_exact_size(b'nonce-bGlnaHQgd29yay4='), ContentSecurityPolicySourceNonce(bytearray(b'light work.')) ) def test_compose(self): self.assertEqual( ContentSecurityPolicySourceNonce(bytearray(b'light work.')).compose(), b'nonce-bGlnaHQgd29yay4=' ) class TestContentSecurityPolicySourceScheme(unittest.TestCase): def test_parse(self): self.assertEqual( ContentSecurityPolicySourceScheme.parse_exact_size(b'http:'), ContentSecurityPolicySourceScheme('http') ) def test_compose(self): self.assertEqual( ContentSecurityPolicySourceScheme('http').compose(), b'http:' ) class TestContentSecurityPolicyDirectivesFetch( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader): _header_minimal = ContentSecurityPolicyDirectiveDefaultSrc([ContentSecurityPolicySourceKeyword.SELF]) _header_minimal_bytes = ' '.join([ 'default-src', '\'self\'', ]).encode('ascii') _header_minimal_markdown = os.linesep.join([ '* Type: default-src', '* Value:', ' 1.', ' * Type: KEYWORD', ' * Value: \'self\'', '', ]) _header_full = ContentSecurityPolicyDirectiveDefaultSrc([ ContentSecurityPolicySourceKeyword.SELF, ContentSecurityPolicySourceScheme('http'), ContentSecurityPolicySourceHost('https://example.com'), ContentSecurityPolicySourceNonce(bytearray(b'light work.')), ContentSecurityPolicySourceHash(Hash.SHA2_256, bytearray(b'light work.')), ]) _header_full_bytes = ' '.join([ 'default-src', '\'self\'', 'http:', 'https://example.com', 'nonce-bGlnaHQgd29yay4=', 'sha256-bGlnaHQgd29yay4=', ]).encode('ascii') def test_min_source_length(self): with self.assertRaises(InvalidValue) as context_manager: ContentSecurityPolicyDirectiveDefaultSrc.parse_exact_size(b'default-src') self.assertEqual(context_manager.exception.value, b'') def test_error_invalid_value_type(self): with self.assertRaises(InvalidValue): ContentSecurityPolicyDirectiveDefaultSrc([None]) class TestContentSecurityPolicyDirectiveFrameAncestors( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader): _header_minimal = ContentSecurityPolicyDirectiveFrameAncestors([ContentSecurityPolicySourceKeyword.SELF]) _header_minimal_bytes = ' '.join([ 'frame-ancestors', '\'self\'', ]).encode('ascii') _header_minimal_markdown = os.linesep.join([ '* Type: frame-ancestors', '* Value:', ' 1.', ' * Type: KEYWORD', ' * Value: \'self\'', '', ]) _header_full = ContentSecurityPolicyDirectiveFrameAncestors([ ContentSecurityPolicySourceKeyword.SELF, ContentSecurityPolicySourceScheme('http'), ContentSecurityPolicySourceHost('https://example.com'), ]) _header_full_bytes = ' '.join([ 'frame-ancestors', '\'self\'', 'http:', 'https://example.com', ]).encode('ascii') def test_error_invalid_value_type(self): with self.assertRaises(InvalidValue): ContentSecurityPolicyDirectiveFrameAncestors([ ContentSecurityPolicySourceKeyword(ContentSecurityPolicySourceKeyword.REPORT_SAMPLE) ]) with self.assertRaises(InvalidValue): ContentSecurityPolicyDirectiveFrameAncestors([ ContentSecurityPolicySourceNonce(bytearray(b'light work.')) ]) with self.assertRaises(InvalidValue): ContentSecurityPolicyDirectiveFrameAncestors([ ContentSecurityPolicySourceHash(Hash.SHA2_256, bytearray(b'light work.')) ]) class TestContentSecurityPolicyDirectiveWebrtc( TestCasesBasesHttpHeader.MinimalHeader): _header_minimal = ContentSecurityPolicyDirectiveWebrtc(ContentSecurityPolicyWebRtcType.ALLOW) _header_minimal_bytes = ' '.join([ 'webrtc', '\'allow\'', ]).encode('ascii') _header_minimal_markdown = '\'allow\'' def test_error_invalid_value_type(self): with self.assertRaises(InvalidValue) as context_manager: ContentSecurityPolicyDirectiveWebrtc('not-a-webrtc-type') self.assertEqual(context_manager.exception.value, 'not-a-webrtc-type') class TestContentSecurityPolicyDirectiveRequireTrustedTypesFor( TestCasesBasesHttpHeader.MinimalHeader): _header_minimal = ContentSecurityPolicyDirectiveRequireTrustedTypesFor([ ContentSecurityPolicyTrustedTypeSinkGroup.SCRIPT ]) _header_minimal_bytes = ' '.join([ 'require-trusted-types-for', '\'script\'', ]).encode('ascii') _header_minimal_markdown = os.linesep.join([ '* Type: require-trusted-types-for', '* Sink Groups:', ' 1. \'script\'', '', ]) class TestContentSecurityPolicyDirectiveReportUri( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader): _header_minimal = ContentSecurityPolicyDirectiveReportUri(['http://example.com']) _header_minimal_bytes = ' '.join([ 'report-uri', 'http://example.com', ]).encode('ascii') _header_minimal_markdown = os.linesep.join([ '* Type: report-uri', '* URI references:', ' 1. http://example.com', '', ]) _header_full = ContentSecurityPolicyDirectiveReportUri(['http://example.com/1', 'http://example.com/2']) _header_full_bytes = ' '.join([ 'report-uri', 'http://example.com/1', 'http://example.com/2', ]).encode('ascii') def test_min_reference_length(self): with self.assertRaises(InvalidValue) as context_manager: ContentSecurityPolicyDirectiveReportUri.parse_exact_size(b'report-uri') self.assertEqual(context_manager.exception.value, b'') class TestContentSecurityPolicyDirectiveReportTo( TestCasesBasesHttpHeader.MinimalHeader): _header_minimal = ContentSecurityPolicyDirectiveReportTo('token') _header_minimal_bytes = ' '.join([ 'report-to', 'token', ]).encode('ascii') _header_minimal_markdown = os.linesep.join([ '* Token: token', '', ]) class TestContentSecurityPolicyDirectiveSandbox( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader): _header_minimal = ContentSecurityPolicyDirectiveSandbox(['token']) _header_minimal_bytes = ' '.join([ 'sandbox', 'token', ]).encode('ascii') _header_minimal_markdown = os.linesep.join([ '* Type: sandbox', '* Tokens:', ' 1. token', '', ]) _header_full = ContentSecurityPolicyDirectiveSandbox(['token1', 'token2']) _header_full_bytes = ' '.join([ 'sandbox', 'token1', 'token2', ]).encode('ascii') class TestContentSecurityPolicyDirectivePluginTypes( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader): _header_minimal = ContentSecurityPolicyDirectivePluginTypes([ FieldValueMimeType('html', MimeTypeRegistry.TEXT) ]) _header_minimal_bytes = ' '.join([ 'plugin-types', 'text/html', ]).encode('ascii') _header_minimal_markdown = os.linesep.join([ '* Type: plugin-types', '* MIME Types:', ' 1.', ' * Type: html', ' * Registry: TEXT', '', ]) _header_full = ContentSecurityPolicyDirectivePluginTypes([ FieldValueMimeType('html', MimeTypeRegistry.TEXT), FieldValueMimeType('csv', MimeTypeRegistry.TEXT), ]) _header_full_bytes = ' '.join([ 'plugin-types', 'text/html', 'text/csv', ]).encode('ascii') class TestContentSecurityPolicyDirectiveNoValue( TestCasesBasesHttpHeader.MinimalHeader): _header_minimal = ContentSecurityPolicyDirectiveBlockAllMixedContent() _header_minimal_bytes = b'block-all-mixed-content' _header_minimal_markdown = os.linesep.join([ '* Type: block-all-mixed-content', '* Value: n/a', '', ]) class TestHttpHeaderFieldValueContentSecurityPolicy( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader): _header_minimal = HttpHeaderFieldValueContentSecurityPolicy([ ContentSecurityPolicyDirectiveDefaultSrc([ ContentSecurityPolicySourceKeyword.SELF, ContentSecurityPolicySourceHash(Hash.SHA2_256, bytearray(b'light work.')), ContentSecurityPolicySourceNonce(bytearray(b'light work.')), ContentSecurityPolicySourceScheme('http'), ContentSecurityPolicySourceHost('http://example.com'), ]) ]) _header_minimal_bytes = ' '.join([ 'default-src', '\'self\'', 'sha256-bGlnaHQgd29yay4=', 'nonce-bGlnaHQgd29yay4=', 'http:', 'http://example.com', ]).encode('ascii') _header_minimal_markdown = os.linesep.join([ '* Directives:', ' 1.', ' * Type: default-src', ' * Value:', ' 1.', ' * Type: KEYWORD', ' * Value: \'self\'', ' 2.', ' * Type: HASH', ' * Value:', ' * Hash Algorithm: SHA-256', ' * Hash Value: bGlnaHQgd29yay4=', ' 3.', ' * Type: NONCE', ' * Value: bGlnaHQgd29yay4=', ' 4.', ' * Type: SCHEME', ' * Value: http', ' 5.', ' * Type: HOST', ' * Value: http://example.com', '', ]) _header_full = HttpHeaderFieldValueContentSecurityPolicy([ ContentSecurityPolicyDirectiveChildSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveConnectSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveDefaultSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveFontSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveFrameSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveImgSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveManifestSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveMediaSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveObjectSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectivePrefetchSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveScriptSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveScriptSrcElem([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveScriptSrcAttr([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveStyleSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveStyleSrcElem([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveStyleSrcAttr([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveWebrtc( ContentSecurityPolicyWebRtcType.ALLOW, ), ContentSecurityPolicyDirectiveWorkerSrc([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveBaseUri([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveSandbox([ 'token' ]), ContentSecurityPolicyDirectiveFormAction([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveFrameAncestors([ ContentSecurityPolicySourceKeyword.SELF, ]), ContentSecurityPolicyDirectiveReportUri([ 'http://example.com' ]), ContentSecurityPolicyDirectiveReportTo( 'token' ), ContentSecurityPolicyDirectiveBlockAllMixedContent(), ContentSecurityPolicyDirectiveUpgradeInsecureRequests(), ContentSecurityPolicyDirectiveReferrer( ContentSecurityPolicyReferrerPolicy.NO_REFERRER, ), ContentSecurityPolicyDirectivePluginTypes([ FieldValueMimeType('html', MimeTypeRegistry.TEXT) ]), ]) _header_full_bytes = '; '.join([ 'child-src \'self\'', 'connect-src \'self\'', 'default-src \'self\'', 'font-src \'self\'', 'frame-src \'self\'', 'img-src \'self\'', 'manifest-src \'self\'', 'media-src \'self\'', 'object-src \'self\'', 'prefetch-src \'self\'', 'script-src \'self\'', 'script-src-elem \'self\'', 'script-src-attr \'self\'', 'style-src \'self\'', 'style-src-elem \'self\'', 'style-src-attr \'self\'', 'webrtc \'allow\'', 'worker-src \'self\'', 'base-uri \'self\'', 'sandbox token', 'form-action \'self\'', 'frame-ancestors \'self\'', 'report-uri http://example.com', 'report-to token', 'block-all-mixed-content', 'upgrade-insecure-requests', 'referrer "no-referrer"', 'plugin-types text/html', ]).encode('ascii') class TestHttpHeaderFieldValueContentType( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader): _header_minimal = HttpHeaderFieldValueContentType( FieldValueMimeType('html', MimeTypeRegistry.TEXT) ) _header_minimal_bytes = b'text/html' _header_minimal_markdown = os.linesep.join([ '* MIME type:', ' * Type: html', ' * Registry: TEXT', '* Charset: n/a', '* Boundary: n/a', '', ]) _header_full = HttpHeaderFieldValueContentType( FieldValueMimeType('bhttp', MimeTypeRegistry.MESSAGE), charset='utf-8', boundary='boundary_pattern', ) _header_full_bytes = b'message/bhttp; charset=utf-8; boundary=boundary_pattern' def test_error_invalid_parameter(self): with self.assertRaises(InvalidValue) as context_manager: HttpHeaderFieldValueContentType( FieldValueMimeType('html', MimeTypeRegistry.TEXT), boundary='pattern', ) self.assertEqual(context_manager.exception.value, 'pattern') with self.assertRaises(InvalidValue) as context_manager: HttpHeaderFieldValueContentType( FieldValueMimeType('bhttp', MimeTypeRegistry.MESSAGE), ) self.assertEqual(context_manager.exception.value, None) class TestHttpHeaderFieldValueNetworkErrorLogging( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader): _header_minimal = HttpHeaderFieldValueNetworkErrorLogging( report_to="network-errors", max_age=1, ) _header_minimal_bytes = b'{"report_to": "network-errors", "max_age": 1}' _header_minimal_markdown = os.linesep.join([ '* Report To: network-errors', '* Max Age: 0:00:01', '* Include Subdomains: n/a', '* Success Fraction: n/a', '* Failure Fraction: n/a', '', ]) _header_full = HttpHeaderFieldValueNetworkErrorLogging( report_to="network-errors", max_age=datetime.timedelta(1), include_subdomains=True, success_fraction=0.1, failure_fraction=0.9, ) _header_full_bytes = b''.join([ b'{', b'"report_to": "network-errors", ', b'"max_age": 86400, ', b'"include_subdomains": true, ', b'"success_fraction": 0.1, ', b'"failure_fraction": 0.9', b'}', ]) class TestHttpHeaderFieldValuePragma( TestCasesBasesHttpHeader.FullHeader, TestCasesBasesHttpHeader.CaseInsensitiveHeader): _header_full = HttpHeaderFieldValuePragma(HttpHeaderPragma.NO_CACHE) _header_full_bytes = b'no-cache' _header_full_upper_case_bytes = b'NO-CACHE' _header_full_lower_case_bytes = b'no-cache' class TestHttpHeaderFieldValueSTS( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader, TestCasesBasesHttpHeader.CaseInsensitiveHeader): _header_minimal = HttpHeaderFieldValueSTS(max_age=datetime.timedelta(seconds=1)) _header_minimal_bytes = b'max-age=1' _header_minimal_markdown = os.linesep.join([ '* Max Age: 0:00:01', '* Include Subdomains: no', '* Preload: no', '', ]) _header_full = HttpHeaderFieldValueSTS(max_age=datetime.timedelta(seconds=1), include_subdomains=True, preload=True) _header_full_bytes = b'max-age=1; includeSubDomains; preload' _header_full_upper_case_bytes = b'MAX-AGE=1; INCLUDESUBDOMAINS; PRELOAD' _header_full_lower_case_bytes = b'max-age=1; includesubdomains; preload' class TestHttpHeaderFieldValueExpectStaple( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader, TestCasesBasesHttpHeader.CaseInsensitiveHeader): _header_minimal = HttpHeaderFieldValueExpectStaple(max_age=datetime.timedelta(seconds=1)) _header_minimal_bytes = b'max-age=1' _header_minimal_markdown = os.linesep.join([ '* Max Age: 0:00:01', '* Include Subdomains: no', '* Preload: no', '* Report Uri: n/a', '', ]) _header_full = HttpHeaderFieldValueExpectStaple( max_age=datetime.timedelta(seconds=1), include_subdomains=True, preload=True, report_uri="http://example.com" ) _header_full_bytes = b'max-age=1; includeSubDomains; preload; report-uri="http://example.com"' _header_full_upper_case_bytes = b'MAX-AGE=1; INCLUDESUBDOMAINS; PRELOAD; REPORT-URI="http://example.com"' _header_full_lower_case_bytes = b'max-age=1; includesubdomains; preload; report-uri="http://example.com"' class TestHttpHeaderFieldValueExpectCT( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader, TestCasesBasesHttpHeader.CaseInsensitiveHeader): _header_minimal = HttpHeaderFieldValueExpectCT(max_age=datetime.timedelta(seconds=1)) _header_minimal_bytes = b'max-age=1' _header_minimal_markdown = os.linesep.join([ '* Max Age: 0:00:01', '* Enforce: no', '* Report Uri: n/a', '', ]) _header_full = HttpHeaderFieldValueExpectCT( max_age=datetime.timedelta(seconds=1), enforce=True, report_uri='http://example.com' ) _header_full_bytes = b'max-age=1, enforce, report-uri="http://example.com"' _header_full_upper_case_bytes = b'MAX-AGE=1, ENFORCE, REPORT-URI="http://example.com"' _header_full_lower_case_bytes = b'max-age=1, enforce, report-uri="http://example.com"' class TestHttpHeaderFieldValuePublicKeyPinning( TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader, TestCasesBasesHttpHeader.CaseInsensitiveHeader): _header_minimal = HttpHeaderFieldValuePublicKeyPinning( pin_sha256='cGluLXNoYTI1Ng==', max_age=datetime.timedelta(seconds=1), ) _header_minimal_bytes = b'pin-sha256="cGluLXNoYTI1Ng=="; max-age=1' _header_minimal_markdown = os.linesep.join([ '* Pin (SHA-256): cGluLXNoYTI1Ng==', '* Max Age: 0:00:01', '* Include Subdomains: no', '* Report Uri: n/a', '', ]) _header_full = HttpHeaderFieldValuePublicKeyPinning( pin_sha256='cGluLXNoYTI1Ng==', max_age=datetime.timedelta(seconds=1), include_subdomains=True, report_uri='http://example.com' ) _header_full_bytes = b'; '.join([ b'pin-sha256="cGluLXNoYTI1Ng=="', b'max-age=1', b'includeSubDomains', b'report-uri="http://example.com"', ]) _header_full_upper_case_bytes = b'; '.join([ b'PIN-SHA256="cGluLXNoYTI1Ng=="', b'MAX-AGE=1', b'INCLUDESUBDOMAINS', b'REPORT-URI="http://example.com"', ]) _header_full_lower_case_bytes = b'; '.join([ b'pin-sha256="cGluLXNoYTI1Ng=="', b'max-age=1', b'includesubdomains', b'report-uri="http://example.com"', ]) class TestHttpHeaderFieldValueSetCookie(TestCasesBasesHttpHeader.MinimalHeader, TestCasesBasesHttpHeader.FullHeader): _header_minimal = HttpHeaderFieldValueSetCookie('name', 'value') _header_minimal_bytes = b'name=value' _header_minimal_markdown = os.linesep.join([ '* Name: name', '* Value: value', '* Expires: n/a', '* Max Age: n/a', '* Domain: n/a', '* Path: n/a', '* Secure: no', '* Http Only: no', '* Same Site: n/a', '' ]) _header_full = HttpHeaderFieldValueSetCookie( name='name', value='value', expires=datetime.datetime.fromtimestamp(0, datetime.timezone.utc), max_age=datetime.timedelta(seconds=1), domain='example.com', path='/', secure=True, http_only=True, same_site=HttpHeaderSetCookieComponentSameSite.LAX, ) _header_full_bytes = b'; '.join([ b'name=value', b'expires=Thu, 01 Jan 1970 00:00:00 GMT', b'max-age=1', b'Domain=example.com', b'Path=/', b'Secure', b'HttpOnly', b'SameSite=Lax', ]) class TestHttpHeaderFieldValueXContentTypeOptions(TestCasesBasesHttpHeader.FullHeader): _header_full = HttpHeaderFieldValueXContentTypeOptions(HttpHeaderXContentTypeOptions.NOSNIFF) _header_full_bytes = b'nosniff' class TestHttpHeaderFieldValueXFrameOptions(TestCasesBasesHttpHeader.FullHeader): _header_full = HttpHeaderFieldValueXFrameOptions(HttpHeaderXFrameOptions.SAMEORIGIN) _header_full_bytes = b'SAMEORIGIN' class TestHttpHeaderFieldValueXXSSProtection(TestCasesBasesHttpHeader.FullHeader): _header_full = HttpHeaderFieldValueXXSSProtection( HttpHeaderXXSSProtectionState.ENABLED, HttpHeaderXXSSProtectionMode.BLOCK, 'http://example.com' ) _header_full_bytes = b'1; mode=block; report=http://example.com' class TestHttpHeaderFieldValueReferrerPolicy(TestCasesBasesHttpHeader.FullHeader): _header_full = HttpHeaderFieldValueReferrerPolicy(HttpHeaderReferrerPolicy.SAME_ORIGIN) _header_full_bytes = b'same-origin' class TestHttpHeaderFieldName(unittest.TestCase): def test_markdown(self): self.assertEqual( HttpHeaderFieldName.STRICT_TRANSPORT_SECURITY.value.as_markdown(), 'Strict-Transport-Security', ) def test_from_name(self): with self.assertRaises(InvalidValue) as context_manager: HttpHeaderFieldName.from_name('non-existing-name') self.assertEqual(context_manager.exception.value, 'non-existing-name') self.assertEqual( HttpHeaderFieldName.from_name('strict-transport-security'), HttpHeaderFieldName.STRICT_TRANSPORT_SECURITY ) self.assertEqual( HttpHeaderFieldName.from_name('STRICT-TRANSPORT-SECURITY'), HttpHeaderFieldName.STRICT_TRANSPORT_SECURITY ) self.assertEqual( HttpHeaderFieldName.from_name('Strict-Transport-Security'), HttpHeaderFieldName.STRICT_TRANSPORT_SECURITY ) class TestHttpHeaderFields(unittest.TestCase): def setUp(self): self.headers_bytes = b'\r\n'.join([ b'Age: 1', b'Cache-Control: no-cache', b'Content-Type: text/html', b'Date: Thu, 01 Jan 1970 00:00:00 GMT', b'ETag: 12345678', b'Expect-CT: max-age=1', b'Expect-Staple: max-age=1', b'Expires: Thu, 01 Jan 1970 00:00:00 GMT', b'Last-Modified: Thu, 01 Jan 1970 00:00:00 GMT', b'NEL: {"report_to": "network-errors", "max_age": 1}', b'Pragma: no-cache', b'Public-Key-Pinning: pin-sha256="cGluLXNoYTI1Ng=="; max-age=1', b'Referrer-Policy: origin', b'Server: server', b'Set-Cookie: name=value', b'Strict-Transport-Security: max-age=1', b'X-Unparsed: Value', b'X-Content-Type-Options: nosniff', b'X-Frame-Options: SAMEORIGIN', b'X-XSS-Protection: 1', b'Content-Security-Policy: default-src \'self\'', b'Content-Security-Policy-Report-Only: default-src \'self\'', b'X-Content-Security-Policy: default-src \'self\'', b'', b'', ]) self.headers = HttpHeaderFields([ HttpHeaderFieldAge(datetime.timedelta(seconds=1)), HttpHeaderFieldCacheControlResponse(HttpHeaderFieldValueCacheControlResponse(no_cache=True)), HttpHeaderFieldContentType(FieldValueMimeType('html', MimeTypeRegistry.TEXT)), HttpHeaderFieldDate(datetime.datetime.fromtimestamp(0, datetime.timezone.utc)), HttpHeaderFieldETag('12345678'), HttpHeaderFieldExpectCT(HttpHeaderFieldValueExpectCT(datetime.timedelta(seconds=1))), HttpHeaderFieldExpectStaple(HttpHeaderFieldValueExpectStaple(datetime.timedelta(seconds=1))), HttpHeaderFieldExpires(datetime.datetime.fromtimestamp(0, datetime.timezone.utc)), HttpHeaderFieldLastModified(datetime.datetime.fromtimestamp(0, datetime.timezone.utc)), HttpHeaderFieldNetworkErrorLogging(HttpHeaderFieldValueNetworkErrorLogging( report_to="network-errors", max_age=1 )), HttpHeaderFieldPragma(HttpHeaderPragma.NO_CACHE), HttpHeaderFieldPublicKeyPinning(HttpHeaderFieldValuePublicKeyPinning( pin_sha256='cGluLXNoYTI1Ng==', max_age=datetime.timedelta(seconds=1), )), HttpHeaderFieldReferrerPolicy(HttpHeaderReferrerPolicy.ORIGIN), HttpHeaderFieldServer('server'), HttpHeaderFieldSetCookie(HttpHeaderFieldValueSetCookie('name', 'value')), HttpHeaderFieldSTS(HttpHeaderFieldValueSTS(datetime.timedelta(seconds=1))), HttpHeaderFieldUnparsed('X-Unparsed', 'Value'), HttpHeaderFieldXContentTypeOptions(HttpHeaderXContentTypeOptions.NOSNIFF), HttpHeaderFieldXFrameOptions(HttpHeaderXFrameOptions.SAMEORIGIN), HttpHeaderFieldXXSSProtection(HttpHeaderXXSSProtectionState.ENABLED), HttpHeaderFieldContentSecurityPolicy([ ContentSecurityPolicyDirectiveDefaultSrc([ContentSecurityPolicySourceKeyword.SELF]) ]), HttpHeaderFieldContentSecurityPolicyReportOnly([ ContentSecurityPolicyDirectiveDefaultSrc([ContentSecurityPolicySourceKeyword.SELF]) ]), HttpHeaderFieldXContentSecurityPolicy([ ContentSecurityPolicyDirectiveDefaultSrc([ContentSecurityPolicySourceKeyword.SELF]) ]), ]) def test_parse(self): self.assertEqual( HttpHeaderFields.parse_exact_size(self.headers_bytes), self.headers ) def test_compose(self): self.assertEqual(self.headers.compose(), self.headers_bytes) def test_markdown(self): self.assertEqual(self.headers.as_markdown(), os.linesep.join([ '1.', ' * Name: Age', ' * Value: 1', '2.', ' * Name: Cache-Control', ' * Value:', ' * Max Age: n/a', ' * S Maxage: n/a', ' * Must Revalidate: no', ' * Proxy Revalidate: no', ' * No Cache: yes', ' * No Store: no', ' * Public: no', ' * Private: no', ' * No Transform: no', '3.', ' * Name: Content-Type', ' * Value:', ' * MIME type:', ' * Type: html', ' * Registry: TEXT', ' * Charset: n/a', ' * Boundary: n/a', '4.', ' * Name: Date', ' * Value: 1970-01-01 00:00:00+00:00', '5.', ' * Name: ETag', ' * Value: 12345678', '6.', ' * Name: Expect-CT', ' * Value:', ' * Max Age: 0:00:01', ' * Enforce: no', ' * Report Uri: n/a', '7.', ' * Name: Expect-Staple', ' * Value:', ' * Max Age: 0:00:01', ' * Include Subdomains: no', ' * Preload: no', ' * Report Uri: n/a', '8.', ' * Name: Expires', ' * Value: 1970-01-01 00:00:00+00:00', '9.', ' * Name: Last-Modified', ' * Value: 1970-01-01 00:00:00+00:00', '10.', ' * Name: NEL', ' * Value:', ' * Report To: network-errors', ' * Max Age: 0:00:01', ' * Include Subdomains: n/a', ' * Success Fraction: n/a', ' * Failure Fraction: n/a', '11.', ' * Name: Pragma', ' * Value: no-cache', '12.', ' * Name: Public-Key-Pinning', ' * Value:', ' * Pin (SHA-256): cGluLXNoYTI1Ng==', ' * Max Age: 0:00:01', ' * Include Subdomains: no', ' * Report Uri: n/a', '13.', ' * Name: Referrer-Policy', ' * Value: origin', '14.', ' * Name: Server', ' * Value: server', '15.', ' * Name: Set-Cookie', ' * Value:', ' * Name: name', ' * Value: value', ' * Expires: n/a', ' * Max Age: n/a', ' * Domain: n/a', ' * Path: n/a', ' * Secure: no', ' * Http Only: no', ' * Same Site: n/a', '16.', ' * Name: Strict-Transport-Security', ' * Value:', ' * Max Age: 0:00:01', ' * Include Subdomains: no', ' * Preload: no', '17.', ' * Name: X-Unparsed', ' * Value: Value', '18.', ' * Name: X-Content-Type-Options', ' * Value: nosniff', '19.', ' * Name: X-Frame-Options', ' * Value: SAMEORIGIN', '20.', ' * Name: X-XSS-Protection', ' * Value:', ' * State: enabled', ' * Mode: n/a', ' * Report: n/a', '21.', ' * Name: Content-Security-Policy', ' * Value:', ' * Directives:', ' 1.', ' * Type: default-src', ' * Value:', ' 1.', ' * Type: KEYWORD', ' * Value: \'self\'', '22.', ' * Name: Content-Security-Policy-Report-Only', ' * Value:', ' * Directives:', ' 1.', ' * Type: default-src', ' * Value:', ' 1.', ' * Type: KEYWORD', ' * Value: \'self\'', '23.', ' * Name: X-Content-Security-Policy', ' * Value:', ' * Directives:', ' 1.', ' * Type: default-src', ' * Value:', ' 1.', ' * Type: KEYWORD', ' * Value: \'self\'', '', ])) class TestHttpHeaderFieldUnparsed(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: HttpHeaderFieldUnparsed.parse_immutable(b'name: value') self.assertEqual(context_manager.exception.value, b'value') with self.assertRaises(InvalidValue) as context_manager: HttpHeaderFieldUnparsed.parse_immutable(b'name value') self.assertEqual(context_manager.exception.value, b'') def test_parse(self): parsable = b'name: value\r\n' header, parsed_length = HttpHeaderFieldUnparsed.parse_immutable(parsable) self.assertEqual(header.name, 'name') self.assertEqual(header.value, 'value') self.assertEqual(parsable[parsed_length:], b'\r\n') parsable = b'name: value\r\n' header, parsed_length = HttpHeaderFieldUnparsed.parse_immutable(parsable) self.assertEqual(header.name, 'name') self.assertEqual(header.value, 'value') self.assertEqual(parsable[parsed_length:], b'\r\n') def test_compose(self): header = HttpHeaderFieldUnparsed('name', 'value') self.assertEqual(header.compose(), b'name: value') def test_markdown(self): header = HttpHeaderFieldUnparsed('name', 'value') self.assertEqual(header.as_markdown(), os.linesep.join([ '* Name: name', '* Value: value', '' ])) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/httpx/test_version.py000066400000000000000000000007441524413560000273030ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptoparser.httpx.version import HttpVersion class TestHttpVersion(unittest.TestCase): def test_markdown(self): self.assertEqual(HttpVersion.HTTP1_0.value.as_json(), '"http1_0"') self.assertEqual(HttpVersion.HTTP1_0.value.as_markdown(), 'HTTP/1.0') self.assertEqual(HttpVersion.HTTP1_1.value.as_json(), '"http1_1"') self.assertEqual(HttpVersion.HTTP1_1.value.as_markdown(), 'HTTP/1.1') cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ike/000077500000000000000000000000001524413560000236015ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ike/__init__.py000066400000000000000000000000431524413560000257070ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ike/classes.py000066400000000000000000000253461524413560000256220ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import collections import unittest from cryptodatahub.ike.algorithm import ( Ikev1PayloadType, Ikev2NotifyType, Ikev2PayloadType, Ikev2ProtocolId, Ikev2TransformType, Ikev2PseudorandomFunction, ) from cryptoparser.common.parse import ComposerBinary from cryptoparser.ike.ikev1 import Ikev1PayloadBase, Ikev1PayloadDoiProtocolSpiBase from cryptoparser.ike.ikev2 import ( Ikev2NotifyPayloadNatDetectionBase, Ikev2PayloadBase, Ikev2PayloadNotifyBase, Ikev2PayloadNotifyNoData, Transform, TransformNextPayload, ) class Ikev1PayloadBaseTest(Ikev1PayloadBase): """Concrete implementation of Ikev1PayloadBase for testing.""" def __init__(self, test_data): super().__init__() self.test_data = test_data self.next_payload = Ikev1PayloadType.NONE @classmethod def get_payload_type(cls): return Ikev1PayloadType.NONE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) test_data_length = parser['payload_length'] - cls.HEADER_SIZE parser.parse_raw('test_data', test_data_length) payload = cls( test_data=parser['test_data'] ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_raw(self.test_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes class Ikev1PayloadDoiProtocolSpiBaseTest(Ikev1PayloadDoiProtocolSpiBase): """Concrete implementation of Ikev1PayloadDoiProtocolSpiBase for testing.""" def __init__(self, doi, protocol_id, spi_size, extra_data=b''): super().__init__(doi=doi, protocol_id=protocol_id, spi_size=spi_size) self.extra_data = extra_data self.next_payload = Ikev1PayloadType.NONE @classmethod def get_payload_type(cls): return Ikev1PayloadType.NONE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) extra_data_length = parser['payload_length'] - parser.parsed_length parser.parse_raw('extra_data', extra_data_length) payload = cls( doi=parser['doi'], protocol_id=parser['protocol_id'], spi_size=parser['spi_size'], extra_data=bytes(parser['extra_data']) ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() self._compose_doi_protocol_spi(composer_payload) composer_payload.compose_raw(self.extra_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes class Ikev2PayloadBaseTest(Ikev2PayloadBase): """Concrete implementation of Ikev2PayloadBase for testing.""" def __init__(self, flags, test_data): super().__init__(flags=flags) self.test_data = test_data self.next_payload = Ikev2PayloadType.NONE @classmethod def get_payload_type(cls): return Ikev2PayloadType.NONE @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) test_data_length = parser['payload_length'] - cls.HEADER_SIZE parser.parse_raw('test_data', test_data_length) payload = cls( flags=parser['flags'], test_data=parser['test_data'] ) payload.next_payload = parser['next_payload'] return payload, parser.parsed_length def compose(self): composer_payload = ComposerBinary() composer_payload.compose_raw(self.test_data) composer_header = self.compose_header(composer_payload.composed_length) return composer_header.composed_bytes + composer_payload.composed_bytes class Ikev2PayloadNotifyNoDataTest(Ikev2PayloadNotifyNoData): """Concrete implementation of Ikev2PayloadNotifyNoData for testing.""" def __init__(self, flags, protocol_id, notify_type, spi): super().__init__(flags=flags, protocol_id=protocol_id, type=notify_type, spi=spi) self.next_payload = Ikev2PayloadType.NONE @classmethod def _get_message_type(cls): return Ikev2NotifyType.AUTHENTICATION_FAILED @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric_enum_coded('protocol_id', Ikev2ProtocolId) parser.parse_numeric('spi_size', 1) cls._parse_type(parser, 'type') if parser['spi_size'] > 0: parser.parse_raw('spi', parser['spi_size']) spi = parser['spi'] else: spi = b'' del parser['spi_size'] if 'spi' in parser: del parser['spi'] notification_data_length = parser['payload_length'] - (cls.HEADER_SIZE + 4) cls._parse_data(parser, notification_data_length) next_payload = parser['next_payload'] del parser['next_payload'] del parser['payload_length'] payload = cls( flags=parser['flags'], protocol_id=parser['protocol_id'], notify_type=parser['type'], spi=spi, ) payload.next_payload = next_payload return payload, parser.parsed_length class Ikev2PayloadNotifyBaseTest(Ikev2PayloadNotifyBase): """Concrete implementation of Ikev2PayloadNotifyBase for testing.""" def __init__( # pylint: disable=too-many-arguments,too-many-positional-arguments self, flags, protocol_id, notify_type, spi, test_data ): super().__init__(flags=flags, protocol_id=protocol_id, type=notify_type, spi=spi) self.test_data = test_data self.next_payload = Ikev2PayloadType.NONE @classmethod def _parse_type(cls, parser, name): parser.parse_numeric_enum_coded(name, Ikev2NotifyType) @classmethod def _parse_data(cls, parser, notification_data_length): if notification_data_length > 0: parser.parse_raw('test_data', notification_data_length) else: # For empty data, we need to handle it differently since parser is not a dict pass def _compose_data(self, composer): composer.compose_raw(self.test_data) @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) parser.parse_numeric_enum_coded('protocol_id', Ikev2ProtocolId) parser.parse_numeric('spi_size', 1) cls._parse_type(parser, 'type') if parser['spi_size'] > 0: parser.parse_raw('spi', parser['spi_size']) spi = parser['spi'] else: spi = b'' del parser['spi_size'] if 'spi' in parser: del parser['spi'] notification_data_length = parser['payload_length'] - (cls.HEADER_SIZE + 4) - len(spi) cls._parse_data(parser, notification_data_length) next_payload = parser['next_payload'] del parser['next_payload'] del parser['payload_length'] test_data = parser.get('test_data', b'') if 'test_data' in parser: del parser['test_data'] payload = cls( flags=parser['flags'], protocol_id=parser['protocol_id'], notify_type=parser['type'], spi=spi, test_data=test_data, ) payload.next_payload = next_payload return payload, parser.parsed_length class TransformTest(Transform): """Concrete implementation of Transform base class for testing.""" def __init__(self, transform_id): super().__init__(transform_id=transform_id) self.next_payload = TransformNextPayload.LAST @classmethod def get_transform_type(cls): return Ikev2TransformType.PRF @classmethod def _get_transform_id_class(cls): return Ikev2PseudorandomFunction @classmethod def _parse(cls, parsable): parser = cls._parse_header(parsable) transform = cls( transform_id=parser['transform_id'], ) transform.next_payload = parser['next_payload'] return transform, parser.parsed_length def compose(self): return self.compose_header(transform_length=0).composed_bytes class Ikev2NotifyPayloadNatDetectionBaseTest(unittest.TestCase): _HASH_DATA = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f\x10\x11\x12\x13' _NOTIFY_TYPE: Ikev2NotifyType _PAYLOAD_CLASS: type[Ikev2NotifyPayloadNatDetectionBase] _NOTIFY_TYPE_BYTES: bytes def setUp(self): payload_dict = collections.OrderedDict([ ('next_payload', b'\x00'), ('flags', b'\x00'), ('payload_length', b'\x00\x1c'), ('protocol_id', b'\x01'), ('spi_size', b'\x00'), ('notify_type', self._NOTIFY_TYPE_BYTES), ('hash_data', self._HASH_DATA), ]) self.payload_bytes = b''.join(payload_dict.values()) self.nat_payload = self._PAYLOAD_CLASS( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=self._NOTIFY_TYPE, spi=b'', hash_data=self._HASH_DATA ) self.nat_payload.next_payload = Ikev2PayloadType.NONE def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual(self._PAYLOAD_CLASS._get_message_type(), self._NOTIFY_TYPE) def test_parse(self): parsed_payload = self._PAYLOAD_CLASS.parse_exact_size(self.payload_bytes) self.assertEqual(parsed_payload.hash_data, self._HASH_DATA) # pylint: disable=no-member self.assertEqual(parsed_payload.type, self._NOTIFY_TYPE) def test_compose(self): self.assertEqual(self.nat_payload.compose(), self.payload_bytes) def test_hash_data_storage(self): self.assertEqual(self.nat_payload.hash_data, self._HASH_DATA) # pylint: disable=no-member different_hash = b'\xff\xfe\xfd\xfc\xfb\xfa' payload_2 = self._PAYLOAD_CLASS( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=self._NOTIFY_TYPE, spi=b'', hash_data=different_hash ) self.assertEqual(payload_2.hash_data, different_hash) # pylint: disable=no-member def test_round_trip_hash_preservation(self): composed_bytes = self.nat_payload.compose() parsed_payload = self._PAYLOAD_CLASS.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.hash_data, self.nat_payload.hash_data) # pylint: disable=no-member self.assertEqual(parsed_payload.type, self.nat_payload.type) self.assertEqual(parsed_payload.spi, self.nat_payload.spi) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ike/test_ikev1_payload.py000066400000000000000000002043441524413560000277510ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import collections import ipaddress import unittest from cryptodatahub.common.algorithm import IpProtocolNumber from cryptodatahub.common.exception import InvalidValue from cryptodatahub.ike.algorithm import ( Ikev1PayloadType, Ikev1AttributeType, Ikev1AuthenticationMethod, Ikev1DiffieHellmanGroup, Ikev1HashAlgorithm, Ikev1LifeType, Ikev1TransformId, Ikev1EncryptionAlgorithm, Ikev1Doi, Ikev1ProtocolId, Ikev1NotifyType, Ikev1CertificateType, Ikev1IdType ) from cryptoparser.common.exception import InvalidType, NotEnoughData from cryptoparser.ike.ikev1 import ( Ikev1AttributeAuthenticationMethod, Ikev1AttributeDiffieHellmanGroup, Ikev1AttributeEncryptionAlgorithm, Ikev1AttributeHashAlgorithm, Ikev1AttributeKeyLength, Ikev1AttributeLifeDuration, Ikev1AttributeLifeType, Ikev1PayloadCertificate, Ikev1PayloadCertificateRequest, Ikev1PayloadDelete, Ikev1PayloadHash, Ikev1PayloadIdentificationDerAsn1Dn, Ikev1PayloadIdentificationDerAsn1Gn, Ikev1PayloadIdentificationFqdn, Ikev1PayloadIdentificationIpv4Addr, Ikev1PayloadIdentificationIpv6Addr, Ikev1PayloadIdentificationKeyId, Ikev1PayloadIdentificationUserFqdn, Ikev1PayloadIdentificationVariant, Ikev1PayloadKeyExchange, Ikev1PayloadNonce, Ikev1PayloadNotification, Ikev1PayloadSignature, Ikev1PayloadTransform, Ikev1PayloadVendorId, ) from .classes import Ikev1PayloadBaseTest, Ikev1PayloadDoiProtocolSpiBaseTest class TestIkev1PayloadBaseTest(unittest.TestCase): """Test the Ikev1PayloadBaseTest helper class, focusing only on unique aspects.""" def setUp(self): self.test_data = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' self.test_payload_empty = Ikev1PayloadBaseTest(test_data=b'') self.test_payload_with_data = Ikev1PayloadBaseTest(test_data=self.test_data) self.test_dict_empty = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x04'), # 4 bytes total (header only) ]) self.test_bytes_empty = b''.join(self.test_dict_empty.values()) self.test_dict_with_data = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x14'), # 4 + 16 = 20 bytes total ('test_data', self.test_data), ]) self.test_bytes_with_data = b''.join(self.test_dict_with_data.values()) def test_constructor_with_test_data(self): payload = Ikev1PayloadBaseTest(test_data=self.test_data) self.assertEqual(payload.test_data, self.test_data) self.assertEqual(payload.next_payload, Ikev1PayloadType.NONE) def test_constructor_with_empty_data(self): payload = Ikev1PayloadBaseTest(test_data=b'') self.assertEqual(payload.test_data, b'') self.assertEqual(payload.next_payload, Ikev1PayloadType.NONE) def test_get_payload_type_returns_none(self): self.assertEqual(Ikev1PayloadBaseTest.get_payload_type(), Ikev1PayloadType.NONE) def test_test_data_storage(self): different_data = b'\xaa\xbb\xcc\xdd' payload = Ikev1PayloadBaseTest(test_data=different_data) self.assertEqual(payload.test_data, different_data) def test_parse_empty_test_data(self): parsed: Ikev1PayloadBaseTest = Ikev1PayloadBaseTest.parse_exact_size(self.test_bytes_empty) self.assertEqual(parsed.test_data, b'') self.assertEqual(parsed.next_payload, Ikev1PayloadType.NONE) def test_parse_with_test_data(self): parsed: Ikev1PayloadBaseTest = Ikev1PayloadBaseTest.parse_exact_size(self.test_bytes_with_data) self.assertEqual(parsed.test_data, self.test_data) self.assertEqual(parsed.next_payload, Ikev1PayloadType.NONE) def test_compose_empty_test_data(self): composed_bytes = self.test_payload_empty.compose() self.assertEqual(composed_bytes, self.test_bytes_empty) def test_compose_with_test_data(self): composed_bytes = self.test_payload_with_data.compose() self.assertEqual(composed_bytes, self.test_bytes_with_data) def test_round_trip_test_data_preservation(self): composed_bytes = self.test_payload_with_data.compose() parsed: Ikev1PayloadBaseTest = Ikev1PayloadBaseTest.parse_exact_size(composed_bytes) self.assertEqual(parsed.test_data, self.test_payload_with_data.test_data) class TestIkev1PayloadBase(unittest.TestCase): def setUp(self): self.test_data = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' self.test_payload_minimal = Ikev1PayloadBaseTest( test_data=b'', ) self.test_payload_minimal.next_payload = Ikev1PayloadType.NONE self.test_payload_with_data = Ikev1PayloadBaseTest( test_data=self.test_data ) self.test_payload_with_data.next_payload = Ikev1PayloadType.KEY_EXCHANGE self.test_dict_minimal = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x04'), ]) self.test_bytes_minimal = b''.join(self.test_dict_minimal.values()) self.test_dict_with_data = collections.OrderedDict([ ('next_payload', b'\x04'), # KEY_EXCHANGE = 0x04 ('reserved', b'\x00'), ('payload_length', b'\x00\x14'), # 4 + 16 = 20 bytes total ('test_data', self.test_data), ]) self.test_bytes_with_data = b''.join(self.test_dict_with_data.values()) def test_get_payload_type(self): self.assertEqual(Ikev1PayloadBaseTest.get_payload_type(), Ikev1PayloadType.NONE) def test_parse(self): parsed_minimal: Ikev1PayloadBaseTest = Ikev1PayloadBaseTest.parse_exact_size(self.test_bytes_minimal) self.assertEqual(parsed_minimal.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed_minimal.test_data, b'') parsed_with_data: Ikev1PayloadBaseTest = Ikev1PayloadBaseTest.parse_exact_size(self.test_bytes_with_data) self.assertEqual(parsed_with_data.next_payload, Ikev1PayloadType.KEY_EXCHANGE) self.assertEqual(parsed_with_data.test_data, self.test_data) def test_compose(self): composed_minimal = self.test_payload_minimal.compose() self.assertEqual(composed_minimal, self.test_bytes_minimal) composed_with_data = self.test_payload_with_data.compose() self.assertEqual(composed_with_data, self.test_bytes_with_data) def test_round_trip(self): composed_minimal = self.test_payload_minimal.compose() parsed_minimal: Ikev1PayloadBaseTest = Ikev1PayloadBaseTest.parse_exact_size(composed_minimal) self.assertEqual(parsed_minimal.test_data, self.test_payload_minimal.test_data) self.assertEqual(parsed_minimal.next_payload, self.test_payload_minimal.next_payload) composed_with_data = self.test_payload_with_data.compose() parsed_with_data: Ikev1PayloadBaseTest = Ikev1PayloadBaseTest.parse_exact_size(composed_with_data) self.assertEqual(parsed_with_data.test_data, self.test_payload_with_data.test_data) self.assertEqual(parsed_with_data.next_payload, self.test_payload_with_data.next_payload) def test_next_payload(self): self.test_payload_minimal.next_payload = Ikev1PayloadType.SECURITY_ASSOCIATION composed = self.test_payload_minimal.compose() parsed: Ikev1PayloadBaseTest = Ikev1PayloadBaseTest.parse_exact_size(composed) self.assertEqual(parsed.next_payload, Ikev1PayloadType.SECURITY_ASSOCIATION) def test_payload_length(self): test_data = b'\x00\x01\x02\x03\x04\x05' payload = Ikev1PayloadBaseTest(test_data) payload.next_payload = Ikev1PayloadType.NONCE composed = payload.compose() self.assertEqual(len(composed), 10) # 1 + 1 + 2 + 6 = 10 bytes total self.assertEqual(composed[2:4], b'\x00\x0a') # 10 = 0x000a def test_error_parse_not_enough_data(self): incomplete_data = self.test_bytes_minimal[:-1] with self.assertRaises(NotEnoughData) as context_manager: Ikev1PayloadBaseTest.parse_exact_size(incomplete_data) self.assertEqual(context_manager.exception.bytes_needed, 1) def test_error_parse_payload_length_mismatch(self): malformed_dict = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x14'), # 20 bytes total ]) malformed_data = b''.join(malformed_dict.values()) with self.assertRaises(NotEnoughData) as context_manager: Ikev1PayloadBaseTest.parse_exact_size(malformed_data) self.assertEqual(context_manager.exception.bytes_needed, 16) class TestIkev1PayloadSaIdentifyingBase(unittest.TestCase): """Test Ikev1PayloadDoiProtocolSpiBase (RFC 2408 Section 2.4 Identifying Security Associations).""" def setUp(self): self.doi = Ikev1Doi.IPSEC self.protocol_id = Ikev1ProtocolId.ISAKMP self.spi_size = 8 self.extra_data = b'\x00\x01\x02\x03' self.test_payload = Ikev1PayloadDoiProtocolSpiBaseTest( doi=self.doi, protocol_id=self.protocol_id, spi_size=self.spi_size, extra_data=self.extra_data ) self.test_payload.next_payload = Ikev1PayloadType.NONE def test_parse(self): composed = self.test_payload.compose() parsed: Ikev1PayloadDoiProtocolSpiBaseTest = Ikev1PayloadDoiProtocolSpiBaseTest.parse_exact_size(composed) self.assertEqual(parsed.doi, self.doi) self.assertEqual(parsed.protocol_id, self.protocol_id) self.assertEqual(parsed.spi_size, self.spi_size) self.assertEqual(parsed.extra_data, self.extra_data) def test_compose_round_trip(self): composed = self.test_payload.compose() parsed: Ikev1PayloadDoiProtocolSpiBaseTest = Ikev1PayloadDoiProtocolSpiBaseTest.parse_exact_size(composed) recomposed = parsed.compose() self.assertEqual(recomposed, composed) def test_different_doi_protocol_spi_values(self): for doi, protocol_id, spi_size in [ (Ikev1Doi.ISAKMP, Ikev1ProtocolId.IPSEC_ESP, 4), (Ikev1Doi.GDOI, Ikev1ProtocolId.IPSEC_AH, 0), ]: payload = Ikev1PayloadDoiProtocolSpiBaseTest( doi=doi, protocol_id=protocol_id, spi_size=spi_size, extra_data=b'' ) composed = payload.compose() parsed: Ikev1PayloadDoiProtocolSpiBaseTest = Ikev1PayloadDoiProtocolSpiBaseTest.parse_exact_size(composed) self.assertEqual(parsed.doi, doi) self.assertEqual(parsed.protocol_id, protocol_id) self.assertEqual(parsed.spi_size, spi_size) def test_error_not_enough_data(self): minimal = Ikev1PayloadDoiProtocolSpiBaseTest( doi=self.doi, protocol_id=self.protocol_id, spi_size=0, extra_data=b'' ) minimal.next_payload = Ikev1PayloadType.NONE minimal_bytes = minimal.compose() incomplete = minimal_bytes[:-1] with self.assertRaises(NotEnoughData) as context_manager: Ikev1PayloadDoiProtocolSpiBaseTest.parse_exact_size(incomplete) self.assertEqual(context_manager.exception.bytes_needed, 1) truncated_doi_protocol_spi = ( b'\x00\x00' # next_payload, reserved b'\x00\x09' # payload_length = 9 b'\x00\x00\x00\x00\x00' # 5 bytes (need 6 for DOI+Protocol+SPI) ) with self.assertRaises(NotEnoughData) as context_manager: Ikev1PayloadDoiProtocolSpiBaseTest.parse_exact_size(truncated_doi_protocol_spi) self.assertEqual(context_manager.exception.bytes_needed, 1) class TestIkev1AttributeKeyLength(unittest.TestCase): def setUp(self): self.key_length_value = 128 # 128-bit key length self.key_length_attribute = Ikev1AttributeKeyLength(value=self.key_length_value) def test_get_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeKeyLength.get_type(), Ikev1AttributeType.KEY_LENGTH) def test_key_length_value_support(self): different_key_lengths = [64, 128, 192, 256] # bits for key_length in different_key_lengths: attribute = Ikev1AttributeKeyLength(value=key_length) self.assertEqual(attribute.value, key_length) def test_round_trip(self): composed_bytes = self.key_length_attribute.compose() parsed_attribute: Ikev1AttributeKeyLength = Ikev1AttributeKeyLength.parse_exact_size(composed_bytes) self.assertEqual(parsed_attribute.value, self.key_length_attribute.value) class TestIkev1AttributeEncryptionAlgorithm(unittest.TestCase): def setUp(self): self.encryption_algorithm = Ikev1EncryptionAlgorithm.AES_CBC self.encryption_attribute = Ikev1AttributeEncryptionAlgorithm(value=self.encryption_algorithm) def test_get_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeEncryptionAlgorithm.get_type(), Ikev1AttributeType.ENCRYPTION_ALGORITHM) def test_get_enum_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeEncryptionAlgorithm._get_enum_type(), Ikev1EncryptionAlgorithm) def test_encryption_algorithm_support(self): different_encryption_algorithms = [ Ikev1EncryptionAlgorithm.DES_CBC, Ikev1EncryptionAlgorithm.DES3_CBC, Ikev1EncryptionAlgorithm.AES_CBC, ] for encryption_algorithm in different_encryption_algorithms: attribute = Ikev1AttributeEncryptionAlgorithm(value=encryption_algorithm) self.assertEqual(attribute.value, encryption_algorithm) def test_round_trip(self): composed_bytes = self.encryption_attribute.compose() parsed_attribute: Ikev1AttributeEncryptionAlgorithm = Ikev1AttributeEncryptionAlgorithm.parse_exact_size( composed_bytes ) self.assertEqual(parsed_attribute.value, self.encryption_attribute.value) class TestIkev1AttributeAuthenticationMethod(unittest.TestCase): def setUp(self): self.auth_method = Ikev1AuthenticationMethod.PRE_SHARED_KEY self.auth_attribute = Ikev1AttributeAuthenticationMethod(value=self.auth_method) def test_get_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeAuthenticationMethod.get_type(), Ikev1AttributeType.AUTHENTICATION_METHOD) def test_get_enum_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeAuthenticationMethod._get_enum_type(), Ikev1AuthenticationMethod) def test_authentication_method_support(self): different_auth_methods = [ Ikev1AuthenticationMethod.PRE_SHARED_KEY, Ikev1AuthenticationMethod.DSS_SIGNATURES, Ikev1AuthenticationMethod.RSA_SIGNATURES, ] for auth_method in different_auth_methods: attribute = Ikev1AttributeAuthenticationMethod(value=auth_method) self.assertEqual(attribute.value, auth_method) def test_round_trip(self): composed_bytes = self.auth_attribute.compose() parsed_attribute: Ikev1AttributeAuthenticationMethod = Ikev1AttributeAuthenticationMethod.parse_exact_size( composed_bytes ) self.assertEqual(parsed_attribute.value, self.auth_attribute.value) def test_error_invalid_value(self): with self.assertRaises(InvalidValue): Ikev1AttributeAuthenticationMethod(value="invalid") class TestIkev1AttributeDiffieHellmanGroup(unittest.TestCase): def setUp(self): self.dh_group = Ikev1DiffieHellmanGroup.MODP_1024_BIT self.dh_attribute = Ikev1AttributeDiffieHellmanGroup(value=self.dh_group) def test_get_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeDiffieHellmanGroup.get_type(), Ikev1AttributeType.GROUP_DESCRIPTION) def test_get_enum_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeDiffieHellmanGroup._get_enum_type(), Ikev1DiffieHellmanGroup) def test_diffie_hellman_group_support(self): different_dh_groups = [ Ikev1DiffieHellmanGroup.MODP_768_BIT, Ikev1DiffieHellmanGroup.MODP_1024_BIT, Ikev1DiffieHellmanGroup.MODP_1536_BIT, ] for dh_group in different_dh_groups: attribute = Ikev1AttributeDiffieHellmanGroup(value=dh_group) self.assertEqual(attribute.value, dh_group) def test_round_trip(self): composed_bytes = self.dh_attribute.compose() parsed_attribute: Ikev1AttributeDiffieHellmanGroup = Ikev1AttributeDiffieHellmanGroup.parse_exact_size( composed_bytes ) self.assertEqual(parsed_attribute.value, self.dh_attribute.value) class TestIkev1AttributeHashAlgorithm(unittest.TestCase): def setUp(self): self.hash_algorithm = Ikev1HashAlgorithm.MD5 self.hash_attribute = Ikev1AttributeHashAlgorithm(value=self.hash_algorithm) def test_get_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeHashAlgorithm.get_type(), Ikev1AttributeType.HASH_ALGORITHM) def test_get_enum_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeHashAlgorithm._get_enum_type(), Ikev1HashAlgorithm) def test_hash_algorithm_support(self): different_hash_algorithms = [ Ikev1HashAlgorithm.MD5, Ikev1HashAlgorithm.SHA, ] for hash_algorithm in different_hash_algorithms: attribute = Ikev1AttributeHashAlgorithm(value=hash_algorithm) self.assertEqual(attribute.value, hash_algorithm) def test_round_trip(self): composed_bytes = self.hash_attribute.compose() parsed_attribute: Ikev1AttributeHashAlgorithm = Ikev1AttributeHashAlgorithm.parse_exact_size(composed_bytes) self.assertEqual(parsed_attribute.value, self.hash_attribute.value) class TestIkev1AttributeLifeType(unittest.TestCase): def setUp(self): self.life_type = Ikev1LifeType.SECONDS self.life_type_attribute = Ikev1AttributeLifeType(value=self.life_type) def test_get_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeLifeType.get_type(), Ikev1AttributeType.LIFE_TYPE) def test_get_enum_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeLifeType._get_enum_type(), Ikev1LifeType) def test_life_type_support(self): different_life_types = [ Ikev1LifeType.SECONDS, Ikev1LifeType.KILOBYTES, ] for life_type in different_life_types: attribute = Ikev1AttributeLifeType(value=life_type) self.assertEqual(attribute.value, life_type) def test_round_trip(self): composed_bytes = self.life_type_attribute.compose() parsed_attribute: Ikev1AttributeLifeType = Ikev1AttributeLifeType.parse_exact_size(composed_bytes) self.assertEqual(parsed_attribute.value, self.life_type_attribute.value) class TestIkev1AttributeLifeDuration(unittest.TestCase): def setUp(self): self.life_duration_value = 3600 # 1 hour in seconds self.life_duration_attribute = Ikev1AttributeLifeDuration(value=self.life_duration_value) def test_get_type(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeLifeDuration.get_type(), Ikev1AttributeType.LIFE_DURATION) def test_get_size(self): # pylint: disable=protected-access self.assertEqual(Ikev1AttributeLifeDuration._get_size(), 4) def test_life_duration_value_support(self): different_durations = [3600, 86400, 604800] # seconds: 1 hour, 1 day, 1 week for duration in different_durations: attribute = Ikev1AttributeLifeDuration(value=duration) self.assertEqual(attribute.value, duration) def test_round_trip(self): composed_bytes = self.life_duration_attribute.compose() parsed_attribute: Ikev1AttributeLifeDuration = Ikev1AttributeLifeDuration.parse_exact_size(composed_bytes) self.assertEqual(parsed_attribute.value, self.life_duration_attribute.value) class TestIkev1PayloadTransform(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.transform_id = Ikev1TransformId.KEY_IKE self.attributes = [] self.transform = Ikev1PayloadTransform(transform_id=self.transform_id, attributes=self.attributes) # Test data for transform without attributes self.test_dict_no_attrs = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x08'), ('transform_number', b'\x01'), # Transform number 1 ('transform_id', b'\x01'), # KEY_IKE = 0x01 ('reserved2', b'\x00\x00'), ]) self.test_bytes_no_attrs = b''.join(self.test_dict_no_attrs.values()) # Test data for transform with authentication method attribute self.auth_attribute = Ikev1AttributeAuthenticationMethod(value=Ikev1AuthenticationMethod.PRE_SHARED_KEY) self.test_dict_auth_attr = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x0c'), # 12 bytes total ('transform_number', b'\x01'), # Transform number 1 ('transform_id', b'\x01'), # KEY_IKE = 0x01 ('reserved2', b'\x00\x00'), ('attr_format_type', b'\x80\x03'), # AF=1, type=3 (AUTH_METHOD) ('attr_value', b'\x00\x01'), # PSK = 0x0001 ]) self.test_bytes_auth_attr = b''.join(self.test_dict_auth_attr.values()) # Test data for transform with key length attribute self.key_length_attribute = Ikev1AttributeKeyLength(value=128) self.test_dict_key_length_attr = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x0c'), # 12 bytes total ('transform_number', b'\x01'), # Transform number 1 ('transform_id', b'\x01'), # KEY_IKE = 0x01 ('reserved2', b'\x00\x00'), ('attr_format_type', b'\x80\x0e'), # AF=1, type=14 (KEY_LENGTH) ('attr_value', b'\x00\x80'), # 128 = 0x0080 ]) self.test_bytes_key_length_attr = b''.join(self.test_dict_key_length_attr.values()) # Test data for transform with life duration attribute self.life_duration_attribute = Ikev1AttributeLifeDuration(value=3600) self.test_dict_life_duration_attr = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x10'), # 16 bytes total ('transform_number', b'\x01'), # Transform number 1 ('transform_id', b'\x01'), # KEY_IKE = 0x01 ('reserved2', b'\x00\x00'), ('attr_format_type', b'\x00\x0c'), # AF=0, type=12 (LIFE_DURATION) ('attr_length', b'\x00\x04'), # 4 bytes value length ('attr_value', b'\x00\x00\x0e\x10'), # 3600 = 0x00000e10 ]) self.test_bytes_life_duration_attr = b''.join(self.test_dict_life_duration_attr.values()) # Test data for transform with multiple attributes self.test_dict_multi_attrs = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x10'), # 16 bytes total ('transform_number', b'\x01'), # Transform number 1 ('transform_id', b'\x01'), # KEY_IKE = 0x01 ('reserved2', b'\x00\x00'), ('attr1_format_type', b'\x80\x03'), # AF=1, type=3 (AUTH_METHOD) ('attr1_value', b'\x00\x01'), # PSK = 0x0001 ('attr2_format_type', b'\x80\x0e'), # AF=1, type=14 (KEY_LENGTH) ('attr2_value', b'\x00\x80'), # 128 = 0x0080 ]) self.test_bytes_multi_attrs = b''.join(self.test_dict_multi_attrs.values()) # Transform objects for compose tests self.transform_auth_attr = Ikev1PayloadTransform( transform_id=Ikev1TransformId.KEY_IKE, attributes=[self.auth_attribute] ) self.transform_auth_attr.transform_number = 1 self.transform_auth_attr.next_payload = Ikev1PayloadType.NONE self.transform_key_length_attr = Ikev1PayloadTransform( transform_id=Ikev1TransformId.KEY_IKE, attributes=[self.key_length_attribute] ) self.transform_key_length_attr.transform_number = 1 self.transform_key_length_attr.next_payload = Ikev1PayloadType.NONE self.transform_life_duration_attr = Ikev1PayloadTransform( transform_id=Ikev1TransformId.KEY_IKE, attributes=[self.life_duration_attribute] ) self.transform_life_duration_attr.transform_number = 1 self.transform_life_duration_attr.next_payload = Ikev1PayloadType.NONE self.transform_multi_attrs = Ikev1PayloadTransform( transform_id=Ikev1TransformId.KEY_IKE, attributes=[self.auth_attribute, self.key_length_attribute] ) self.transform_multi_attrs.transform_number = 1 self.transform_multi_attrs.next_payload = Ikev1PayloadType.NONE def test_get_payload_type(self): self.assertEqual(Ikev1PayloadTransform.get_payload_type(), Ikev1PayloadType.TRANSFORM) def test_transform_id_storage(self): self.assertEqual(self.transform.transform_id, self.transform_id) def test_attributes_storage(self): self.assertEqual(self.transform.attributes, self.attributes) def test_get_attribute_by_type(self): transform = Ikev1PayloadTransform( transform_id=Ikev1TransformId.KEY_IKE, attributes=[self.auth_attribute], ) attribute = transform.get_attribute_by_type(Ikev1AttributeType.AUTHENTICATION_METHOD) self.assertIsInstance(attribute, Ikev1AttributeAuthenticationMethod) self.assertEqual(attribute.value, Ikev1AuthenticationMethod.PRE_SHARED_KEY) def test_error_get_attribute_by_type_not_found(self): transform = Ikev1PayloadTransform( transform_id=Ikev1TransformId.KEY_IKE, attributes=[], ) with self.assertRaises(KeyError): transform.get_attribute_by_type(Ikev1AttributeType.AUTHENTICATION_METHOD) def test_transform_number_initialization(self): self.assertIsNone(self.transform.transform_number) def test_transform_with_attributes(self): auth_attribute = Ikev1AttributeAuthenticationMethod(value=Ikev1AuthenticationMethod.PRE_SHARED_KEY) transform_with_attrs = Ikev1PayloadTransform( transform_id=Ikev1TransformId.KEY_IKE, attributes=[auth_attribute] ) self.assertEqual(len(transform_with_attrs.attributes), 1) self.assertEqual(transform_with_attrs.attributes[0], auth_attribute) def test_parse(self): parsed_transform: Ikev1PayloadTransform = Ikev1PayloadTransform.parse_exact_size(self.test_bytes_no_attrs) self.assertEqual(parsed_transform.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed_transform.transform_id, Ikev1TransformId.KEY_IKE) self.assertEqual(len(parsed_transform.attributes), 0) def test_compose(self): self.transform.transform_number = 1 self.transform.next_payload = Ikev1PayloadType.NONE composed_bytes = self.transform.compose() self.assertEqual(composed_bytes, self.test_bytes_no_attrs) def test_round_trip(self): self.transform.transform_number = 1 self.transform.next_payload = Ikev1PayloadType.NONE composed_bytes = self.transform.compose() parsed_transform: Ikev1PayloadTransform = Ikev1PayloadTransform.parse_exact_size(composed_bytes) self.assertEqual(parsed_transform.transform_id, self.transform.transform_id) self.assertEqual(parsed_transform.next_payload, self.transform.next_payload) self.assertEqual(parsed_transform.transform_number, self.transform.transform_number) self.assertEqual(len(parsed_transform.attributes), len(self.transform.attributes)) def test_parse_with_attributes(self): parsed_transform: Ikev1PayloadTransform = Ikev1PayloadTransform.parse_exact_size(self.test_bytes_auth_attr) self.assertEqual(parsed_transform.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed_transform.transform_id, Ikev1TransformId.KEY_IKE) self.assertEqual(len(parsed_transform.attributes), 1) attribute = parsed_transform.attributes[0] self.assertIsInstance(attribute, Ikev1AttributeAuthenticationMethod) self.assertEqual(attribute.value, Ikev1AuthenticationMethod.PRE_SHARED_KEY) def test_parse_with_key_length_attribute(self): parsed_transform: Ikev1PayloadTransform = Ikev1PayloadTransform.parse_exact_size( self.test_bytes_key_length_attr ) self.assertEqual(len(parsed_transform.attributes), 1) attribute = parsed_transform.attributes[0] self.assertIsInstance(attribute, Ikev1AttributeKeyLength) self.assertEqual(attribute.value, 128) def test_parse_with_life_duration_attribute(self): parsed_transform: Ikev1PayloadTransform = Ikev1PayloadTransform.parse_exact_size( self.test_bytes_life_duration_attr ) self.assertEqual(len(parsed_transform.attributes), 1) attribute = parsed_transform.attributes[0] self.assertIsInstance(attribute, Ikev1AttributeLifeDuration) self.assertEqual(attribute.value, 3600) def test_compose_with_attributes(self): composed_bytes = self.transform_auth_attr.compose() self.assertEqual(composed_bytes, self.test_bytes_auth_attr) def test_compose_with_key_length_attribute(self): composed_bytes = self.transform_key_length_attr.compose() self.assertEqual(composed_bytes, self.test_bytes_key_length_attr) def test_compose_with_life_duration_attribute(self): composed_bytes = self.transform_life_duration_attr.compose() self.assertEqual(composed_bytes, self.test_bytes_life_duration_attr) def test_compose_with_multiple_attributes(self): composed_bytes = self.transform_multi_attrs.compose() self.assertEqual(composed_bytes, self.test_bytes_multi_attrs) def test_round_trip_with_attributes(self): composed_bytes = self.transform_multi_attrs.compose() parsed_transform: Ikev1PayloadTransform = Ikev1PayloadTransform.parse_exact_size(composed_bytes) self.assertEqual(parsed_transform.transform_id, self.transform_multi_attrs.transform_id) self.assertEqual(parsed_transform.next_payload, self.transform_multi_attrs.next_payload) self.assertEqual(parsed_transform.transform_number, self.transform_multi_attrs.transform_number) self.assertEqual(len(parsed_transform.attributes), len(self.transform_multi_attrs.attributes)) # Verify individual attributes self.assertIsInstance(parsed_transform.attributes[0], Ikev1AttributeAuthenticationMethod) self.assertEqual(parsed_transform.attributes[0].value, Ikev1AuthenticationMethod.PRE_SHARED_KEY) self.assertIsInstance(parsed_transform.attributes[1], Ikev1AttributeKeyLength) self.assertEqual(parsed_transform.attributes[1].value, 128) class TestIkev1PayloadKeyExchange(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.key_exchange_data = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' self.large_key_data = bytes(range(256)) self.test_dict_small_key = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x14'), # 4 + 16 = 20 bytes total ('key_exchange_data', self.key_exchange_data), ]) self.test_bytes_small_key = b''.join(self.test_dict_small_key.values()) self.test_dict_empty_key = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x04'), ]) self.test_bytes_empty_key = b''.join(self.test_dict_empty_key.values()) self.test_dict_large_key = collections.OrderedDict([ ('next_payload', b'\x0a'), # NONCE = 0x0a (10) ('reserved', b'\x00'), ('payload_length', b'\x01\x04'), # 4 + 256 = 260 bytes total ('key_exchange_data', self.large_key_data), ]) self.test_bytes_large_key = b''.join(self.test_dict_large_key.values()) # KE objects for compose tests self.ke_small = Ikev1PayloadKeyExchange(key_exchange_data=self.key_exchange_data) self.ke_small.next_payload = Ikev1PayloadType.NONE self.ke_empty = Ikev1PayloadKeyExchange(key_exchange_data=b'') self.ke_empty.next_payload = Ikev1PayloadType.NONE self.ke_large = Ikev1PayloadKeyExchange(key_exchange_data=self.large_key_data) self.ke_large.next_payload = Ikev1PayloadType.NONCE def test_get_payload_type(self): self.assertEqual(Ikev1PayloadKeyExchange.get_payload_type(), Ikev1PayloadType.KEY_EXCHANGE) def test_constructor_with_key_data(self): ke = Ikev1PayloadKeyExchange(key_exchange_data=self.key_exchange_data) self.assertEqual(ke.key_exchange_data, self.key_exchange_data) self.assertEqual(ke.next_payload, Ikev1PayloadType.NONE) def test_constructor_with_empty_data(self): ke = Ikev1PayloadKeyExchange(key_exchange_data=b'') self.assertEqual(ke.key_exchange_data, b'') def test_constructor_with_large_data(self): ke = Ikev1PayloadKeyExchange(key_exchange_data=self.large_key_data) self.assertEqual(ke.key_exchange_data, self.large_key_data) self.assertEqual(len(ke.key_exchange_data), 256) def test_key_exchange_data_storage(self): different_data = b'\xaa\xbb\xcc\xdd\xee\xff' ke = Ikev1PayloadKeyExchange(key_exchange_data=different_data) self.assertEqual(ke.key_exchange_data, different_data) def test_parse_small_key_data(self): parsed: Ikev1PayloadKeyExchange = Ikev1PayloadKeyExchange.parse_exact_size(self.test_bytes_small_key) self.assertEqual(parsed.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed.key_exchange_data, self.key_exchange_data) def test_parse_empty_key_data(self): parsed: Ikev1PayloadKeyExchange = Ikev1PayloadKeyExchange.parse_exact_size(self.test_bytes_empty_key) self.assertEqual(parsed.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed.key_exchange_data, b'') def test_parse_large_key_data(self): parsed: Ikev1PayloadKeyExchange = Ikev1PayloadKeyExchange.parse_exact_size(self.test_bytes_large_key) self.assertEqual(parsed.next_payload, Ikev1PayloadType.NONCE) self.assertEqual(parsed.key_exchange_data, self.large_key_data) self.assertEqual(len(parsed.key_exchange_data), 256) def test_compose_small_key_data(self): composed_bytes = self.ke_small.compose() self.assertEqual(composed_bytes, self.test_bytes_small_key) def test_compose_empty_key_data(self): composed_bytes = self.ke_empty.compose() self.assertEqual(composed_bytes, self.test_bytes_empty_key) def test_compose_large_key_data(self): composed_bytes = self.ke_large.compose() self.assertEqual(composed_bytes, self.test_bytes_large_key) def test_round_trip_small_key(self): composed_bytes = self.ke_small.compose() parsed: Ikev1PayloadKeyExchange = Ikev1PayloadKeyExchange.parse_exact_size(composed_bytes) self.assertEqual(parsed.key_exchange_data, self.ke_small.key_exchange_data) self.assertEqual(parsed.next_payload, self.ke_small.next_payload) def test_round_trip_empty_key(self): composed_bytes = self.ke_empty.compose() parsed: Ikev1PayloadKeyExchange = Ikev1PayloadKeyExchange.parse_exact_size(composed_bytes) self.assertEqual(parsed.key_exchange_data, self.ke_empty.key_exchange_data) self.assertEqual(parsed.next_payload, self.ke_empty.next_payload) def test_round_trip_large_key(self): composed_bytes = self.ke_large.compose() parsed: Ikev1PayloadKeyExchange = Ikev1PayloadKeyExchange.parse_exact_size(composed_bytes) self.assertEqual(parsed.key_exchange_data, self.ke_large.key_exchange_data) self.assertEqual(parsed.next_payload, self.ke_large.next_payload) self.assertEqual(len(parsed.key_exchange_data), len(self.ke_large.key_exchange_data)) def test_payload_length_calculation(self): test_data = b'\x12\x34\x56\x78\x9a\xbc' ke = Ikev1PayloadKeyExchange(key_exchange_data=test_data) ke.next_payload = Ikev1PayloadType.VENDOR_ID composed = ke.compose() self.assertEqual(len(composed), 10) # 4 + 6 = 10 bytes total self.assertEqual(composed[2:4], b'\x00\x0a') # 10 = 0x000a def test_different_next_payload_types(self): ke = Ikev1PayloadKeyExchange(key_exchange_data=self.key_exchange_data) ke.next_payload = Ikev1PayloadType.SECURITY_ASSOCIATION composed = ke.compose() parsed: Ikev1PayloadKeyExchange = Ikev1PayloadKeyExchange.parse_exact_size(composed) self.assertEqual(parsed.next_payload, Ikev1PayloadType.SECURITY_ASSOCIATION) self.assertEqual(parsed.key_exchange_data, self.key_exchange_data) class TestIkev1PayloadHash(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.hash_data = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' self.large_hash_data = bytes(range(256)) self.test_dict_small_hash = collections.OrderedDict([ ('next_payload', b'\x0a'), # NONCE = 0x0a (10) ('reserved', b'\x00'), ('payload_length', b'\x00\x14'), # 4 + 16 = 20 bytes total ('hash_data', self.hash_data), ]) self.test_bytes_small_hash = b''.join(self.test_dict_small_hash.values()) self.test_dict_empty_hash = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x04'), ]) self.test_bytes_empty_hash = b''.join(self.test_dict_empty_hash.values()) self.test_dict_large_hash = collections.OrderedDict([ ('next_payload', b'\x04'), # KEY_EXCHANGE = 0x04 ('reserved', b'\x00'), ('payload_length', b'\x01\x04'), # 4 + 256 = 260 bytes total ('hash_data', self.large_hash_data), ]) self.test_bytes_large_hash = b''.join(self.test_dict_large_hash.values()) self.hash_small = Ikev1PayloadHash(hash_data=self.hash_data) self.hash_small.next_payload = Ikev1PayloadType.NONCE self.hash_empty = Ikev1PayloadHash(hash_data=b'') self.hash_empty.next_payload = Ikev1PayloadType.NONE self.hash_large = Ikev1PayloadHash(hash_data=self.large_hash_data) self.hash_large.next_payload = Ikev1PayloadType.KEY_EXCHANGE def test_get_payload_type(self): self.assertEqual(Ikev1PayloadHash.get_payload_type(), Ikev1PayloadType.HASH) def test_parse_small_hash_data(self): parsed: Ikev1PayloadHash = Ikev1PayloadHash.parse_exact_size(self.test_bytes_small_hash) self.assertEqual(parsed.next_payload, Ikev1PayloadType.NONCE) self.assertEqual(parsed.hash_data, self.hash_data) def test_parse_empty_hash_data(self): parsed: Ikev1PayloadHash = Ikev1PayloadHash.parse_exact_size(self.test_bytes_empty_hash) self.assertEqual(parsed.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed.hash_data, b'') def test_parse_large_hash_data(self): parsed: Ikev1PayloadHash = Ikev1PayloadHash.parse_exact_size(self.test_bytes_large_hash) self.assertEqual(parsed.next_payload, Ikev1PayloadType.KEY_EXCHANGE) self.assertEqual(parsed.hash_data, self.large_hash_data) self.assertEqual(len(parsed.hash_data), 256) def test_compose_small_hash_data(self): composed_bytes = self.hash_small.compose() self.assertEqual(composed_bytes, self.test_bytes_small_hash) def test_compose_empty_hash_data(self): composed_bytes = self.hash_empty.compose() self.assertEqual(composed_bytes, self.test_bytes_empty_hash) def test_compose_large_hash_data(self): composed_bytes = self.hash_large.compose() self.assertEqual(composed_bytes, self.test_bytes_large_hash) def test_round_trip_small_hash(self): composed_bytes = self.hash_small.compose() parsed: Ikev1PayloadHash = Ikev1PayloadHash.parse_exact_size(composed_bytes) self.assertEqual(parsed.hash_data, self.hash_small.hash_data) self.assertEqual(parsed.next_payload, self.hash_small.next_payload) def test_round_trip_empty_hash(self): composed_bytes = self.hash_empty.compose() parsed: Ikev1PayloadHash = Ikev1PayloadHash.parse_exact_size(composed_bytes) self.assertEqual(parsed.hash_data, self.hash_empty.hash_data) self.assertEqual(parsed.next_payload, self.hash_empty.next_payload) def test_round_trip_large_hash(self): composed_bytes = self.hash_large.compose() parsed: Ikev1PayloadHash = Ikev1PayloadHash.parse_exact_size(composed_bytes) self.assertEqual(parsed.hash_data, self.hash_large.hash_data) self.assertEqual(parsed.next_payload, self.hash_large.next_payload) self.assertEqual(len(parsed.hash_data), len(self.hash_large.hash_data)) def test_payload_length_calculation(self): test_data = b'\x00\x01\x02\x03\x04\x05' payload = Ikev1PayloadHash(hash_data=test_data) payload.next_payload = Ikev1PayloadType.VENDOR_ID composed = payload.compose() self.assertEqual(len(composed), 10) # 4 + 6 = 10 bytes total self.assertEqual(composed[2:4], b'\x00\x0a') # 10 = 0x000a def test_error_parse_not_enough_data(self): incomplete_data = b'\x00\x00\x00' with self.assertRaises(NotEnoughData) as context_manager: Ikev1PayloadHash.parse_exact_size(incomplete_data) self.assertEqual(context_manager.exception.bytes_needed, 1) def test_error_parse_payload_length_mismatch(self): malformed_dict = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x14'), # 20 bytes total, but missing data ]) malformed_data = b''.join(malformed_dict.values()) with self.assertRaises(NotEnoughData) as context_manager: Ikev1PayloadHash.parse_exact_size(malformed_data) self.assertEqual(context_manager.exception.bytes_needed, 16) class TestIkev1PayloadNonce(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.nonce_data = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' self.large_nonce_data = bytes(range(256)) self.test_dict_small_nonce = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x14'), ('nonce_data', self.nonce_data), ]) self.test_bytes_small_nonce = b''.join(self.test_dict_small_nonce.values()) self.test_dict_empty_nonce = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x04'), ]) self.test_bytes_empty_nonce = b''.join(self.test_dict_empty_nonce.values()) self.test_dict_large_nonce = collections.OrderedDict([ ('next_payload', b'\x04'), # KEY_EXCHANGE = 0x04 ('reserved', b'\x00'), ('payload_length', b'\x01\x04'), ('nonce_data', self.large_nonce_data), ]) self.test_bytes_large_nonce = b''.join(self.test_dict_large_nonce.values()) self.nonce_small = Ikev1PayloadNonce(nonce_data=self.nonce_data) self.nonce_small.next_payload = Ikev1PayloadType.NONE self.nonce_empty = Ikev1PayloadNonce(nonce_data=b'') self.nonce_empty.next_payload = Ikev1PayloadType.NONE self.nonce_large = Ikev1PayloadNonce(nonce_data=self.large_nonce_data) self.nonce_large.next_payload = Ikev1PayloadType.KEY_EXCHANGE def test_get_payload_type(self): self.assertEqual(Ikev1PayloadNonce.get_payload_type(), Ikev1PayloadType.NONCE) def test_payload_length_calculation(self): test_data = b'\x12\x34\x56\x78\x9a\xbc' nonce = Ikev1PayloadNonce(nonce_data=test_data) nonce.next_payload = Ikev1PayloadType.VENDOR_ID composed = nonce.compose() self.assertEqual(len(composed), 10) self.assertEqual(composed[2:4], b'\x00\x0a') # 10 = 0x000a def test_different_next_payload_types(self): nonce = Ikev1PayloadNonce(nonce_data=self.nonce_data) nonce.next_payload = Ikev1PayloadType.SECURITY_ASSOCIATION composed = nonce.compose() parsed: Ikev1PayloadNonce = Ikev1PayloadNonce.parse_exact_size(composed) self.assertEqual(parsed.next_payload, Ikev1PayloadType.SECURITY_ASSOCIATION) self.assertEqual(parsed.nonce_data, self.nonce_data) def test_nonce_with_special_characters(self): special_nonce = b'Nonce\x00\x01\x02\x03\xff\xfe\xfd' nonce = Ikev1PayloadNonce(nonce_data=special_nonce) nonce.next_payload = Ikev1PayloadType.NOTIFICATION composed = nonce.compose() parsed: Ikev1PayloadNonce = Ikev1PayloadNonce.parse_exact_size(composed) self.assertEqual(parsed.nonce_data, special_nonce) self.assertEqual(parsed.next_payload, Ikev1PayloadType.NOTIFICATION) self.assertEqual(len(parsed.nonce_data), len(special_nonce)) class TestIkev1PayloadNotification(unittest.TestCase): """Test Notification-specific payload fields: notify_type, spi, notification_data.""" def setUp(self): self.notify_type = Ikev1NotifyType.NO_PROPOSAL_CHOSEN self.spi = b'\x00\x01\x02\x03\x04\x05\x06\x07' self.notification_data = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' def test_get_payload_type(self): self.assertEqual(Ikev1PayloadNotification.get_payload_type(), Ikev1PayloadType.NOTIFICATION) def test_payload_length_calculation(self): test_spi = b'\x12\x34\x56\x78' test_notification_data = b'\x9a\xbc\xde\xf0' notification = Ikev1PayloadNotification( doi=Ikev1Doi.IPSEC, protocol_id=Ikev1ProtocolId.IPSEC_ESP, spi_size=4, notify_type=Ikev1NotifyType.INVALID_SPI, spi=test_spi, notification_data=test_notification_data ) notification.next_payload = Ikev1PayloadType.VENDOR_ID composed = notification.compose() self.assertEqual(len(composed), 20) self.assertEqual(composed[2:4], b'\x00\x14') # 20 = 0x0014 def test_notify_type_spi_notification_data_round_trip(self): notification = Ikev1PayloadNotification( doi=Ikev1Doi.IPSEC, protocol_id=Ikev1ProtocolId.ISAKMP, spi_size=8, notify_type=self.notify_type, spi=self.spi, notification_data=self.notification_data ) composed = notification.compose() parsed: Ikev1PayloadNotification = Ikev1PayloadNotification.parse_exact_size(composed) self.assertEqual(parsed.notify_type, self.notify_type) self.assertEqual(parsed.spi, self.spi) self.assertEqual(parsed.notification_data, self.notification_data) class TestIkev1PayloadDelete(unittest.TestCase): def setUp(self): self.doi = Ikev1Doi.IPSEC self.protocol_id = Ikev1ProtocolId.ISAKMP self.spi_size = 8 self.spis = [ b'\x00\x01\x02\x03\x04\x05\x06\x07', b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', ] def test_get_payload_type(self): self.assertEqual(Ikev1PayloadDelete.get_payload_type(), Ikev1PayloadType.DELETE) def test_payload_length_calculation(self): spis_4byte = [b'\x12\x34\x56\x78', b'\x9a\xbc\xde\xf0'] delete = Ikev1PayloadDelete( doi=Ikev1Doi.IPSEC, protocol_id=Ikev1ProtocolId.IPSEC_ESP, spi_size=4, spis=spis_4byte, ) delete.next_payload = Ikev1PayloadType.VENDOR_ID composed = delete.compose() self.assertEqual(len(composed), 20) # 4 + 4 + 1 + 1 + 2 + 8 = 20 bytes total self.assertEqual(composed[2:4], b'\x00\x14') # 20 = 0x0014 def test_spis_round_trip(self): delete = Ikev1PayloadDelete( doi=self.doi, protocol_id=self.protocol_id, spi_size=self.spi_size, spis=self.spis, ) composed = delete.compose() parsed: Ikev1PayloadDelete = Ikev1PayloadDelete.parse_exact_size(composed) self.assertEqual(parsed.spis, self.spis) class TestIkev1PayloadVendorId(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.vendor_id = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f\x10\x11\x12\x13' self.large_vendor_id = bytes(range(256)) self.test_dict_small_vendor_id = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x18'), # 4 + 20 = 24 bytes total ('vendor_id', self.vendor_id), ]) self.test_bytes_small_vendor_id = b''.join(self.test_dict_small_vendor_id.values()) self.test_dict_empty_vendor_id = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x04'), ]) self.test_bytes_empty_vendor_id = b''.join(self.test_dict_empty_vendor_id.values()) self.test_dict_large_vendor_id = collections.OrderedDict([ ('next_payload', b'\x04'), # KEY_EXCHANGE = 0x04 ('reserved', b'\x00'), ('payload_length', b'\x01\x04'), # 4 + 256 = 260 bytes total ('vendor_id', self.large_vendor_id), ]) self.test_bytes_large_vendor_id = b''.join(self.test_dict_large_vendor_id.values()) # Vendor ID objects for compose tests self.vendor_id_small = Ikev1PayloadVendorId(vendor_id=self.vendor_id) self.vendor_id_small.next_payload = Ikev1PayloadType.NONE self.vendor_id_empty = Ikev1PayloadVendorId(vendor_id=b'') self.vendor_id_empty.next_payload = Ikev1PayloadType.NONE self.vendor_id_large = Ikev1PayloadVendorId(vendor_id=self.large_vendor_id) self.vendor_id_large.next_payload = Ikev1PayloadType.KEY_EXCHANGE def test_get_payload_type(self): self.assertEqual(Ikev1PayloadVendorId.get_payload_type(), Ikev1PayloadType.VENDOR_ID) def test_vendor_id_with_special_characters(self): special_vendor_id = b'Vendor\x00\x01\x02\x03\xff\xfe\xfd' vendor_id = Ikev1PayloadVendorId(vendor_id=special_vendor_id) vendor_id.next_payload = Ikev1PayloadType.NONCE composed = vendor_id.compose() parsed: Ikev1PayloadVendorId = Ikev1PayloadVendorId.parse_exact_size(composed) self.assertEqual(parsed.vendor_id, special_vendor_id) self.assertEqual(parsed.next_payload, Ikev1PayloadType.NONCE) self.assertEqual(len(parsed.vendor_id), len(special_vendor_id)) class TestIkev1PayloadCertificateRequest(unittest.TestCase): _CERT_ENCODING = Ikev1CertificateType.X509_CERTIFICATE_SIGNATURE _CERTIFICATION_AUTHORITY = ( b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' b'\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f' ) _CERTREQ_BYTES = bytes.fromhex( '00' # next_payload = NONE '00' # reserved '0025' # payload_length = 37 (4 header + 1 cert_encoding + 32 certification authority) '04' # cert_encoding = X.509 certificate signature ) + _CERTIFICATION_AUTHORITY def test_get_payload_type(self): self.assertEqual(Ikev1PayloadCertificateRequest.get_payload_type(), Ikev1PayloadType.CERTIFICATE_REQUEST) def test_parse(self): parsed_certreq: Ikev1PayloadCertificateRequest = Ikev1PayloadCertificateRequest.parse_exact_size( self._CERTREQ_BYTES) self.assertEqual(parsed_certreq.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed_certreq.cert_encoding, self._CERT_ENCODING) self.assertEqual(parsed_certreq.certification_authority, self._CERTIFICATION_AUTHORITY) def test_compose(self): certreq_payload = Ikev1PayloadCertificateRequest( cert_encoding=self._CERT_ENCODING, certification_authority=self._CERTIFICATION_AUTHORITY, ) self.assertEqual(certreq_payload.compose(), self._CERTREQ_BYTES) def test_round_trip(self): certreq_payload = Ikev1PayloadCertificateRequest( cert_encoding=self._CERT_ENCODING, certification_authority=self._CERTIFICATION_AUTHORITY, ) composed_bytes = certreq_payload.compose() parsed_payload: Ikev1PayloadCertificateRequest = Ikev1PayloadCertificateRequest.parse_exact_size( composed_bytes) self.assertEqual(parsed_payload.cert_encoding, certreq_payload.cert_encoding) self.assertEqual(parsed_payload.certification_authority, certreq_payload.certification_authority) self.assertEqual(parsed_payload.next_payload, certreq_payload.next_payload) def test_error_parse_not_enough_data(self): incomplete_data = self._CERTREQ_BYTES[:-5] with self.assertRaises(NotEnoughData) as context_manager: Ikev1PayloadCertificateRequest.parse_exact_size(incomplete_data) self.assertEqual(context_manager.exception.bytes_needed, 5) def test_get_distinguished_name_returns_none_for_empty(self): payload = Ikev1PayloadCertificateRequest( cert_encoding=self._CERT_ENCODING, certification_authority=b'', ) self.assertIsNone(payload.get_distinguished_name()) def test_get_distinguished_name_returns_none_for_non_asn1(self): payload = Ikev1PayloadCertificateRequest( cert_encoding=self._CERT_ENCODING, certification_authority=b'\x00\x01\x02not-asn1', ) self.assertIsNone(payload.get_distinguished_name()) def test_get_distinguished_name_returns_ordered_dict_for_der_encoded_name(self): # DER-encoded ``CN=Example CA`` Distinguished Name. dn_der = bytes.fromhex('301531133011060355040313 0a4578616d706c65204341'.replace(' ', '')) payload = Ikev1PayloadCertificateRequest( cert_encoding=self._CERT_ENCODING, certification_authority=dn_der, ) parsed = payload.get_distinguished_name() self.assertIsNotNone(parsed) self.assertEqual(parsed.get('common_name'), 'Example CA') class TestIkev1PayloadCertificate(unittest.TestCase): _CERT_ENCODING = Ikev1CertificateType.X509_CERTIFICATE_SIGNATURE _CERTIFICATE_DATA = ( b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' b'\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f' ) _CERT_BYTES = bytes.fromhex( '00' # next_payload = NONE '00' # reserved '0025' # payload_length = 37 (4 header + 1 cert_encoding + 32 data) '04' # cert_encoding = X.509 certificate signature ) + _CERTIFICATE_DATA def test_get_payload_type(self): self.assertEqual(Ikev1PayloadCertificate.get_payload_type(), Ikev1PayloadType.CERTIFICATE) def test_parse(self): parsed_cert: Ikev1PayloadCertificate = Ikev1PayloadCertificate.parse_exact_size(self._CERT_BYTES) self.assertEqual(parsed_cert.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed_cert.cert_encoding, self._CERT_ENCODING) self.assertEqual(parsed_cert.certificate_data, self._CERTIFICATE_DATA) def test_compose(self): cert_payload = Ikev1PayloadCertificate( cert_encoding=self._CERT_ENCODING, certificate_data=self._CERTIFICATE_DATA, ) self.assertEqual(cert_payload.compose(), self._CERT_BYTES) def test_round_trip(self): cert_payload = Ikev1PayloadCertificate( cert_encoding=self._CERT_ENCODING, certificate_data=self._CERTIFICATE_DATA, ) composed_bytes = cert_payload.compose() parsed_payload: Ikev1PayloadCertificate = Ikev1PayloadCertificate.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.cert_encoding, cert_payload.cert_encoding) self.assertEqual(parsed_payload.certificate_data, cert_payload.certificate_data) self.assertEqual(parsed_payload.next_payload, cert_payload.next_payload) def test_error_parse_not_enough_data(self): incomplete_data = self._CERT_BYTES[:-5] with self.assertRaises(NotEnoughData) as context_manager: Ikev1PayloadCertificate.parse_exact_size(incomplete_data) self.assertEqual(context_manager.exception.bytes_needed, 5) class TestIkev1PayloadIdentification(unittest.TestCase): _PROTOCOL_ID = IpProtocolNumber.UDP _PORT = 500 # IKE _IDENTIFIER = ipaddress.IPv4Address('192.0.2.1') _ID_BYTES = bytes.fromhex( '00' # next_payload = NONE '00' # reserved '000c' # payload_length = 12 (4 header + 1 id_type + 1 protocol_id + 2 port + 4 id_data) '01' # id_type = IPV4_ADDR '11' # protocol_id = 17 (UDP) '01f4' # port = 500 ) + _IDENTIFIER.packed def test_get_payload_type(self): self.assertEqual( Ikev1PayloadIdentificationIpv4Addr.get_payload_type(), Ikev1PayloadType.IDENTIFICATION, ) def test_parse(self): parsed_id = Ikev1PayloadIdentificationIpv4Addr.parse_exact_size(self._ID_BYTES) self.assertEqual(parsed_id.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed_id.id_type, Ikev1IdType.IPV4_ADDR) self.assertEqual(parsed_id.protocol_id, self._PROTOCOL_ID) self.assertEqual(parsed_id.port, self._PORT) self.assertEqual(parsed_id.identifier, self._IDENTIFIER) # pylint: disable=no-member def test_compose(self): id_payload = Ikev1PayloadIdentificationIpv4Addr( protocol_id=self._PROTOCOL_ID, port=self._PORT, identifier=self._IDENTIFIER, ) self.assertEqual(id_payload.compose(), self._ID_BYTES) def test_round_trip(self): id_payload = Ikev1PayloadIdentificationIpv4Addr( protocol_id=self._PROTOCOL_ID, port=self._PORT, identifier=self._IDENTIFIER, ) composed_bytes = id_payload.compose() parsed_payload = Ikev1PayloadIdentificationIpv4Addr.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.id_type, id_payload.id_type) self.assertEqual(parsed_payload.protocol_id, id_payload.protocol_id) self.assertEqual(parsed_payload.port, id_payload.port) identifier_parsed = parsed_payload.identifier # pylint: disable=no-member identifier_expected = id_payload.identifier # pylint: disable=no-member self.assertEqual(identifier_parsed, identifier_expected) self.assertEqual(parsed_payload.next_payload, id_payload.next_payload) def test_round_trip_fqdn(self): fqdn = 'gateway.example.com' payload = Ikev1PayloadIdentificationFqdn( protocol_id=IpProtocolNumber.HOPOPT, port=0, identifier=fqdn, ) payload.next_payload = Ikev1PayloadType.NONE composed_bytes = payload.compose() parsed_payload = Ikev1PayloadIdentificationFqdn.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.id_type, Ikev1IdType.FQDN) self.assertEqual(parsed_payload.protocol_id, IpProtocolNumber.HOPOPT) self.assertEqual(parsed_payload.port, 0) self.assertEqual(parsed_payload.identifier, fqdn) # pylint: disable=no-member def test_round_trip_der_asn1_dn(self): dn_data = ( b'\x30\x21\x31\x1f\x30\x1d\x06\x03\x55\x04\x03\x0c\x16' b'gateway.example.com' ) payload = Ikev1PayloadIdentificationDerAsn1Dn( protocol_id=IpProtocolNumber.HOPOPT, port=0, identifier=dn_data, ) payload.next_payload = Ikev1PayloadType.NONE composed_bytes = payload.compose() parsed_payload = Ikev1PayloadIdentificationDerAsn1Dn.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.id_type, Ikev1IdType.DER_ASN1_DN) self.assertEqual(parsed_payload.identifier, dn_data) # pylint: disable=no-member def test_error_parse_not_enough_data(self): incomplete_data = self._ID_BYTES[:-2] with self.assertRaises(NotEnoughData) as context_manager: Ikev1PayloadIdentificationIpv4Addr.parse_exact_size(incomplete_data) self.assertEqual(context_manager.exception.bytes_needed, 2) def test_round_trip_der_asn1_gn(self): gn_data = b'\x30\x0c\xa0\x0a\x82\x08gateway' payload = Ikev1PayloadIdentificationDerAsn1Gn( protocol_id=IpProtocolNumber.HOPOPT, port=0, identifier=gn_data, ) payload.next_payload = Ikev1PayloadType.NONE parsed = Ikev1PayloadIdentificationDerAsn1Gn.parse_exact_size(payload.compose()) self.assertEqual(parsed.id_type, Ikev1IdType.DER_ASN1_GN) self.assertEqual(parsed.identifier, gn_data) # pylint: disable=no-member def test_error_parse_negative_id_data_length(self): # payload_length=6 → id_data_length = 6 - 4 - 4 = -2 → NotEnoughData bogus = b'\x00\x00\x00\x06\x02\x00\x00\x00\x00\x00' with self.assertRaises(NotEnoughData): Ikev1PayloadIdentificationFqdn.parse_exact_size(bogus) def test_error_parse_id_type_mismatch(self): # Wire id_type = 0x01 (IPV4_ADDR) but caller uses FQDN subclass. wire = b'\x00\x00\x00\x0d\x01\x00\x00\x00\xc0\x00\x02\x01\x00' with self.assertRaises(InvalidType): Ikev1PayloadIdentificationFqdn.parse_exact_size(wire[:13]) def test_round_trip_user_fqdn(self): payload = Ikev1PayloadIdentificationUserFqdn( protocol_id=IpProtocolNumber.HOPOPT, port=0, identifier='admin@example.com', ) payload.next_payload = Ikev1PayloadType.NONE parsed = Ikev1PayloadIdentificationUserFqdn.parse_exact_size(payload.compose()) self.assertEqual(parsed.id_type, Ikev1IdType.USER_FQDN) self.assertEqual(parsed.identifier, 'admin@example.com') # pylint: disable=no-member def test_round_trip_ipv6(self): addr = ipaddress.IPv6Address('2001:db8::42') payload = Ikev1PayloadIdentificationIpv6Addr( protocol_id=IpProtocolNumber.HOPOPT, port=0, identifier=addr, ) payload.next_payload = Ikev1PayloadType.NONE parsed = Ikev1PayloadIdentificationIpv6Addr.parse_exact_size(payload.compose()) self.assertEqual(parsed.id_type, Ikev1IdType.IPV6_ADDR) self.assertEqual(parsed.identifier, addr) # pylint: disable=no-member def test_round_trip_key_id(self): payload = Ikev1PayloadIdentificationKeyId( protocol_id=IpProtocolNumber.HOPOPT, port=0, identifier=b'\xde\xad\xbe\xef', ) payload.next_payload = Ikev1PayloadType.NONE parsed = Ikev1PayloadIdentificationKeyId.parse_exact_size(payload.compose()) self.assertEqual(parsed.id_type, Ikev1IdType.KEY_ID) self.assertEqual(parsed.identifier, b'\xde\xad\xbe\xef') # pylint: disable=no-member def test_variant_dispatch_picks_subclass_by_id_type(self): payload = Ikev1PayloadIdentificationFqdn( protocol_id=IpProtocolNumber.HOPOPT, port=0, identifier='www.example.com', ) payload.next_payload = Ikev1PayloadType.NONE parsed = Ikev1PayloadIdentificationVariant.parse_exact_size(payload.compose()) self.assertIsInstance(parsed, Ikev1PayloadIdentificationFqdn) def test_user_fqdn_mixin_returns_none_for_ikev2(self): self.assertIsNone(Ikev1PayloadIdentificationUserFqdn.get_id_type_ikev2()) class TestIkev1PayloadSignature(unittest.TestCase): _SIGNATURE_DATA = ( b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' b'\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f' ) _SIG_BYTES = bytes.fromhex( '00' # next_payload = NONE '00' # reserved '0024' # payload_length = 36 (4 header + 32 signature_data) ) + _SIGNATURE_DATA def test_get_payload_type(self): self.assertEqual(Ikev1PayloadSignature.get_payload_type(), Ikev1PayloadType.SIGNATURE) def test_parse(self): parsed_sig: Ikev1PayloadSignature = Ikev1PayloadSignature.parse_exact_size(self._SIG_BYTES) self.assertEqual(parsed_sig.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed_sig.signature_data, self._SIGNATURE_DATA) def test_compose(self): sig_payload = Ikev1PayloadSignature(signature_data=self._SIGNATURE_DATA) self.assertEqual(sig_payload.compose(), self._SIG_BYTES) def test_round_trip(self): sig_payload = Ikev1PayloadSignature(signature_data=self._SIGNATURE_DATA) composed_bytes = sig_payload.compose() parsed_payload: Ikev1PayloadSignature = Ikev1PayloadSignature.parse_exact_size( composed_bytes) self.assertEqual(parsed_payload.signature_data, sig_payload.signature_data) self.assertEqual(parsed_payload.next_payload, sig_payload.next_payload) def test_error_parse_not_enough_data(self): incomplete_data = self._SIG_BYTES[:-4] with self.assertRaises(NotEnoughData) as context_manager: Ikev1PayloadSignature.parse_exact_size(incomplete_data) self.assertEqual(context_manager.exception.bytes_needed, 4) if __name__ == '__main__': unittest.main() cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ike/test_ikev1_sa.py000066400000000000000000000450521524413560000267220ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import collections import unittest from cryptodatahub.ike.algorithm import ( Ikev1PayloadType, Ikev1AuthenticationMethod, Ikev1TransformId, Ikev1ProtocolId, Ikev1Doi ) from cryptoparser.ike.ikev1 import ( Ikev1AttributeAuthenticationMethod, Ikev1PayloadTransform, Ikev1PayloadProposal, Ikev1PayloadSecurityAssociation, Ikev1Situation ) class TestIkev1PayloadProposal(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.simple_transform = Ikev1PayloadTransform( transform_id=Ikev1TransformId.KEY_IKE, attributes=[] ) self.auth_attribute = Ikev1AttributeAuthenticationMethod(value=Ikev1AuthenticationMethod.PRE_SHARED_KEY) self.transform_with_attr = Ikev1PayloadTransform( transform_id=Ikev1TransformId.KEY_IKE, attributes=[self.auth_attribute] ) self.test_dict_single_transform = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x10'), # 16 bytes total ('proposal_number', b'\x01'), ('protocol_id', b'\x01'), # ISAKMP = 0x01 ('spi_size', b'\x00'), ('transform_count', b'\x01'), ('transform_data', b'\x00\x00\x00\x08\x01\x01\x00\x00'), # 8-byte transform ]) self.test_bytes_single_transform = b''.join(self.test_dict_single_transform.values()) self.spi_data = b'\x12\x34\x56\x78' self.test_dict_with_spi = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x14'), # 20 bytes total ('proposal_number', b'\x01'), ('protocol_id', b'\x03'), # IPSEC_ESP = 0x03 ('spi_size', b'\x04'), # 4-byte SPI ('transform_count', b'\x01'), ('spi', self.spi_data), ('transform_data', b'\x00\x00\x00\x08\x01\x01\x00\x00'), # 8-byte transform ]) self.test_bytes_with_spi = b''.join(self.test_dict_with_spi.values()) self.test_dict_multi_transforms = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x18'), # 24 bytes total ('proposal_number', b'\x01'), ('protocol_id', b'\x01'), # ISAKMP = 0x01 ('spi_size', b'\x00'), ('transform_count', b'\x02'), ('transform1_data', b'\x03\x00\x00\x08\x01\x01\x00\x00'), # next=TRANSFORM ('transform2_data', b'\x00\x00\x00\x08\x02\x01\x00\x00'), # next=NONE ]) self.test_bytes_multi_transforms = b''.join(self.test_dict_multi_transforms.values()) self.proposal_single = Ikev1PayloadProposal( protocol_id=Ikev1ProtocolId.ISAKMP, transforms=[self.simple_transform], spi=b'' ) self.proposal_single.proposal_number = 1 self.proposal_single.next_payload = Ikev1PayloadType.NONE self.proposal_with_spi = Ikev1PayloadProposal( protocol_id=Ikev1ProtocolId.IPSEC_ESP, transforms=[self.simple_transform], spi=self.spi_data ) self.proposal_with_spi.proposal_number = 1 self.proposal_with_spi.next_payload = Ikev1PayloadType.NONE self.second_transform = Ikev1PayloadTransform( transform_id=Ikev1TransformId.KEY_IKE, attributes=[] ) self.proposal_multi = Ikev1PayloadProposal( protocol_id=Ikev1ProtocolId.ISAKMP, transforms=[self.simple_transform, self.second_transform], spi=b'' ) self.proposal_multi.proposal_number = 1 self.proposal_multi.next_payload = Ikev1PayloadType.NONE def test_get_payload_type(self): self.assertEqual(Ikev1PayloadProposal.get_payload_type(), Ikev1PayloadType.PROPOSAL) def test_constructor_with_single_transform(self): proposal = Ikev1PayloadProposal( protocol_id=Ikev1ProtocolId.ISAKMP, transforms=[self.simple_transform], spi=b'' ) self.assertEqual(proposal.protocol_id, Ikev1ProtocolId.ISAKMP) self.assertEqual(len(proposal.transforms), 1) self.assertEqual(proposal.transforms[0], self.simple_transform) self.assertEqual(proposal.spi, b'') def test_constructor_with_spi(self): proposal = Ikev1PayloadProposal( protocol_id=Ikev1ProtocolId.IPSEC_ESP, transforms=[self.simple_transform], spi=self.spi_data ) self.assertEqual(proposal.protocol_id, Ikev1ProtocolId.IPSEC_ESP) self.assertEqual(proposal.spi, self.spi_data) def test_constructor_with_multiple_transforms(self): proposal = Ikev1PayloadProposal( protocol_id=Ikev1ProtocolId.ISAKMP, transforms=[self.simple_transform, self.second_transform], spi=b'' ) self.assertEqual(len(proposal.transforms), 2) self.assertEqual(proposal.transforms[0], self.simple_transform) self.assertEqual(proposal.transforms[1], self.second_transform) def test_parse_single_transform(self): parsed: Ikev1PayloadProposal = Ikev1PayloadProposal.parse_exact_size(self.test_bytes_single_transform) self.assertEqual(parsed.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed.proposal_number, 1) self.assertEqual(parsed.protocol_id, Ikev1ProtocolId.ISAKMP) self.assertEqual(parsed.spi, b'') self.assertEqual(len(parsed.transforms), 1) self.assertEqual(parsed.transforms[0].transform_id, Ikev1TransformId.KEY_IKE) def test_parse_with_spi(self): parsed: Ikev1PayloadProposal = Ikev1PayloadProposal.parse_exact_size(self.test_bytes_with_spi) self.assertEqual(parsed.protocol_id, Ikev1ProtocolId.IPSEC_ESP) self.assertEqual(parsed.spi, self.spi_data) self.assertEqual(len(parsed.transforms), 1) def test_parse_multiple_transforms(self): parsed: Ikev1PayloadProposal = Ikev1PayloadProposal.parse_exact_size(self.test_bytes_multi_transforms) self.assertEqual(parsed.proposal_number, 1) self.assertEqual(parsed.protocol_id, Ikev1ProtocolId.ISAKMP) self.assertEqual(len(parsed.transforms), 2) self.assertEqual(parsed.transforms[0].transform_id, Ikev1TransformId.KEY_IKE) self.assertEqual(parsed.transforms[1].transform_id, Ikev1TransformId.KEY_IKE) def test_compose_single_transform(self): composed_bytes = self.proposal_single.compose() self.assertEqual(composed_bytes, self.test_bytes_single_transform) def test_compose_with_spi(self): composed_bytes = self.proposal_with_spi.compose() self.assertEqual(composed_bytes, self.test_bytes_with_spi) def test_compose_multiple_transforms(self): composed_bytes = self.proposal_multi.compose() self.assertEqual(composed_bytes, self.test_bytes_multi_transforms) def test_round_trip_single_transform(self): composed_bytes = self.proposal_single.compose() parsed: Ikev1PayloadProposal = Ikev1PayloadProposal.parse_exact_size(composed_bytes) self.assertEqual(parsed.protocol_id, self.proposal_single.protocol_id) self.assertEqual(parsed.spi, self.proposal_single.spi) self.assertEqual(len(parsed.transforms), len(self.proposal_single.transforms)) self.assertEqual(parsed.proposal_number, self.proposal_single.proposal_number) self.assertEqual(parsed.next_payload, self.proposal_single.next_payload) def test_round_trip_with_spi(self): composed_bytes = self.proposal_with_spi.compose() parsed: Ikev1PayloadProposal = Ikev1PayloadProposal.parse_exact_size(composed_bytes) self.assertEqual(parsed.protocol_id, self.proposal_with_spi.protocol_id) self.assertEqual(parsed.spi, self.proposal_with_spi.spi) self.assertEqual(len(parsed.transforms), len(self.proposal_with_spi.transforms)) def test_round_trip_multiple_transforms(self): composed_bytes = self.proposal_multi.compose() parsed: Ikev1PayloadProposal = Ikev1PayloadProposal.parse_exact_size(composed_bytes) self.assertEqual(parsed.protocol_id, self.proposal_multi.protocol_id) self.assertEqual(len(parsed.transforms), len(self.proposal_multi.transforms)) self.assertEqual(parsed.transforms[0].next_payload, Ikev1PayloadType.TRANSFORM) self.assertEqual(parsed.transforms[1].next_payload, Ikev1PayloadType.NONE) def test_transform_numbering_in_compose(self): composed_bytes = self.proposal_multi.compose() self.assertEqual(composed_bytes[12], 1) # First transform number at byte 12 self.assertEqual(composed_bytes[20], 2) # Second transform number at byte 20 def test_spi_size_calculation(self): different_spi = b'\xaa\xbb\xcc\xdd\xee\xff' # 6 bytes proposal = Ikev1PayloadProposal( protocol_id=Ikev1ProtocolId.IPSEC_AH, transforms=[self.simple_transform], spi=different_spi ) proposal.proposal_number = 1 proposal.next_payload = Ikev1PayloadType.NONE composed_bytes = proposal.compose() self.assertEqual(composed_bytes[6], 6) # SPI size field at byte 6 class TestIkev1PayloadSecurityAssociation(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): # Simple transform for proposals self.simple_transform = Ikev1PayloadTransform( transform_id=Ikev1TransformId.KEY_IKE, attributes=[] ) # Simple proposal for SA self.simple_proposal = Ikev1PayloadProposal( protocol_id=Ikev1ProtocolId.ISAKMP, transforms=[self.simple_transform], spi=b'' ) # Second proposal for multi-proposal tests self.second_proposal = Ikev1PayloadProposal( protocol_id=Ikev1ProtocolId.IPSEC_ESP, transforms=[self.simple_transform], spi=b'\x12\x34\x56\x78' ) # Test data for SA with single proposal self.test_dict_single_proposal = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x1c'), # 28 bytes total ('doi', b'\x00\x00\x00\x01'), # IPSEC = 0x00000001 ('situation', b'\x00\x00\x00\x01'), # SIT_IDENTITY_ONLY = 0x00000001 # Proposal: next=NONE, length=16, prop_num=1, protocol=ISAKMP, spi_size=0, transform_count=1 ('proposal_data', b'\x00\x00\x00\x10\x01\x01\x00\x01'), # Transform: next=NONE, length=8, transform_num=1, transform_id=KEY_IKE ('transform_data', b'\x00\x00\x00\x08\x01\x01\x00\x00'), ]) self.test_bytes_single_proposal = b''.join(self.test_dict_single_proposal.values()) # Test data for SA with multiple proposals self.test_dict_multi_proposals = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x30'), # 48 bytes total ('doi', b'\x00\x00\x00\x01'), # IPSEC = 0x00000001 ('situation', b'\x00\x00\x00\x01'), # SIT_IDENTITY_ONLY = 0x00000001 # First proposal: next=PROPOSAL, length=16, prop_num=1 ('proposal1_data', b'\x02\x00\x00\x10\x01\x01\x00\x01'), ('transform1_data', b'\x00\x00\x00\x08\x01\x01\x00\x00'), # Second proposal: next=NONE, length=20, prop_num=2, protocol=ESP, 4-byte SPI ('proposal2_data', b'\x00\x00\x00\x14\x02\x03\x04\x01'), ('spi_data', b'\x12\x34\x56\x78'), ('transform2_data', b'\x00\x00\x00\x08\x01\x01\x00\x00'), ]) self.test_bytes_multi_proposals = b''.join(self.test_dict_multi_proposals.values()) # Test data for different situation flags self.test_dict_secrecy_integrity = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE = 0x00 ('reserved', b'\x00'), ('payload_length', b'\x00\x1c'), # 28 bytes total ('doi', b'\x00\x00\x00\x01'), # IPSEC = 0x00000001 ('situation', b'\x00\x00\x00\x06'), # SIT_SECRECY | SIT_INTEGRITY = 0x02 | 0x04 = 0x06 ('proposal_data', b'\x00\x00\x00\x10\x01\x01\x00\x01'), ('transform_data', b'\x00\x00\x00\x08\x01\x01\x00\x00'), ]) self.test_bytes_secrecy_integrity = b''.join(self.test_dict_secrecy_integrity.values()) # SA objects for compose tests self.sa_single = Ikev1PayloadSecurityAssociation( doi=Ikev1Doi.IPSEC, situation={Ikev1Situation.SIT_IDENTITY_ONLY}, proposals=[self.simple_proposal] ) self.sa_single.next_payload = Ikev1PayloadType.NONE self.sa_multi = Ikev1PayloadSecurityAssociation( doi=Ikev1Doi.IPSEC, situation={Ikev1Situation.SIT_IDENTITY_ONLY}, proposals=[self.simple_proposal, self.second_proposal] ) self.sa_multi.next_payload = Ikev1PayloadType.NONE self.sa_flags = Ikev1PayloadSecurityAssociation( doi=Ikev1Doi.IPSEC, situation={Ikev1Situation.SIT_SECRECY, Ikev1Situation.SIT_INTEGRITY}, proposals=[self.simple_proposal] ) self.sa_flags.next_payload = Ikev1PayloadType.NONE def test_get_payload_type(self): self.assertEqual(Ikev1PayloadSecurityAssociation.get_payload_type(), Ikev1PayloadType.SECURITY_ASSOCIATION) def test_constructor_with_single_proposal(self): sa = Ikev1PayloadSecurityAssociation( doi=Ikev1Doi.IPSEC, situation={Ikev1Situation.SIT_IDENTITY_ONLY}, proposals=[self.simple_proposal] ) self.assertEqual(sa.doi, Ikev1Doi.IPSEC) self.assertEqual(sa.situation, {Ikev1Situation.SIT_IDENTITY_ONLY}) self.assertEqual(len(sa.proposals), 1) self.assertEqual(sa.proposals[0], self.simple_proposal) def test_constructor_with_multiple_proposals(self): sa = Ikev1PayloadSecurityAssociation( doi=Ikev1Doi.IPSEC, situation={Ikev1Situation.SIT_IDENTITY_ONLY}, proposals=[self.simple_proposal, self.second_proposal] ) self.assertEqual(len(sa.proposals), 2) self.assertEqual(sa.proposals[0], self.simple_proposal) self.assertEqual(sa.proposals[1], self.second_proposal) def test_constructor_with_situation_flags(self): sa = Ikev1PayloadSecurityAssociation( doi=Ikev1Doi.IPSEC, situation={Ikev1Situation.SIT_SECRECY, Ikev1Situation.SIT_INTEGRITY}, proposals=[self.simple_proposal] ) self.assertEqual(sa.situation, {Ikev1Situation.SIT_SECRECY, Ikev1Situation.SIT_INTEGRITY}) def test_doi_storage(self): sa = Ikev1PayloadSecurityAssociation( doi=Ikev1Doi.GDOI, situation={Ikev1Situation.SIT_IDENTITY_ONLY}, proposals=[self.simple_proposal] ) self.assertEqual(sa.doi, Ikev1Doi.GDOI) def test_parse_single_proposal(self): parsed: Ikev1PayloadSecurityAssociation = Ikev1PayloadSecurityAssociation.parse_exact_size( self.test_bytes_single_proposal ) self.assertEqual(parsed.next_payload, Ikev1PayloadType.NONE) self.assertEqual(parsed.doi, Ikev1Doi.IPSEC) self.assertEqual(parsed.situation, {Ikev1Situation.SIT_IDENTITY_ONLY}) self.assertEqual(len(parsed.proposals), 1) self.assertEqual(parsed.proposals[0].protocol_id, Ikev1ProtocolId.ISAKMP) def test_parse_multiple_proposals(self): parsed: Ikev1PayloadSecurityAssociation = Ikev1PayloadSecurityAssociation.parse_exact_size( self.test_bytes_multi_proposals ) self.assertEqual(parsed.doi, Ikev1Doi.IPSEC) self.assertEqual(len(parsed.proposals), 2) self.assertEqual(parsed.proposals[0].protocol_id, Ikev1ProtocolId.ISAKMP) self.assertEqual(parsed.proposals[1].protocol_id, Ikev1ProtocolId.IPSEC_ESP) def test_parse_situation_flags(self): parsed: Ikev1PayloadSecurityAssociation = Ikev1PayloadSecurityAssociation.parse_exact_size( self.test_bytes_secrecy_integrity ) self.assertEqual(parsed.situation, {Ikev1Situation.SIT_SECRECY, Ikev1Situation.SIT_INTEGRITY}) def test_compose_single_proposal(self): composed_bytes = self.sa_single.compose() self.assertEqual(composed_bytes, self.test_bytes_single_proposal) def test_compose_multiple_proposals(self): composed_bytes = self.sa_multi.compose() self.assertEqual(composed_bytes, self.test_bytes_multi_proposals) def test_compose_situation_flags(self): composed_bytes = self.sa_flags.compose() self.assertEqual(composed_bytes, self.test_bytes_secrecy_integrity) def test_round_trip_single_proposal(self): composed_bytes = self.sa_single.compose() parsed: Ikev1PayloadSecurityAssociation = Ikev1PayloadSecurityAssociation.parse_exact_size(composed_bytes) self.assertEqual(parsed.doi, self.sa_single.doi) self.assertEqual(parsed.situation, self.sa_single.situation) self.assertEqual(len(parsed.proposals), len(self.sa_single.proposals)) self.assertEqual(parsed.next_payload, self.sa_single.next_payload) def test_round_trip_multiple_proposals(self): composed_bytes = self.sa_multi.compose() parsed: Ikev1PayloadSecurityAssociation = Ikev1PayloadSecurityAssociation.parse_exact_size(composed_bytes) self.assertEqual(parsed.doi, self.sa_multi.doi) self.assertEqual(len(parsed.proposals), len(self.sa_multi.proposals)) # Verify proposal chaining (first proposal points to next, last to NONE) self.assertEqual(parsed.proposals[0].next_payload, Ikev1PayloadType.PROPOSAL) self.assertEqual(parsed.proposals[1].next_payload, Ikev1PayloadType.NONE) def test_round_trip_situation_flags(self): composed_bytes = self.sa_flags.compose() parsed: Ikev1PayloadSecurityAssociation = Ikev1PayloadSecurityAssociation.parse_exact_size(composed_bytes) self.assertEqual(parsed.situation, self.sa_flags.situation) def test_proposal_numbering_in_compose(self): composed_bytes = self.sa_multi.compose() # Verify proposal numbers are written correctly (1-based indexing) self.assertEqual(composed_bytes[16], 1) # First proposal number at byte 16 self.assertEqual(composed_bytes[32], 2) # Second proposal number at byte 32 if __name__ == '__main__': unittest.main() cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ike/test_ikev2_notify.py000066400000000000000000000701421524413560000276260ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import collections import unittest from cryptodatahub.common.exception import InvalidValue from cryptodatahub.ike.algorithm import Ikev2HashAlgorithm, Ikev2NotifyType, Ikev2ProtocolId from cryptoparser.common.exception import NotEnoughData, InvalidType from cryptoparser.ike.ikev2 import ( Ikev2PayloadFlags, Ikev2PayloadType, Ikev2PayloadNotifyUnparsed, Ikev2NotifyPayloadChildlessIkev2Supported, Ikev2NotifyPayloadCookie, Ikev2NotifyPayloadHttpCertLookupSupported, Ikev2NotifyPayloadIkev2FragmentationSupported, Ikev2NotifyPayloadIntermediateExchangeSupported, Ikev2NotifyPayloadNatDetectionDestinationIp, Ikev2NotifyPayloadNatDetectionSourceIp, Ikev2NotifyPayloadRedirectSupported, Ikev2NotifyPayloadSetWindowSize, Ikev2NotifyPayloadSignatureHashAlgorithms, Ikev2NotifyPayloadUsePpk, Ikev2NotifyPayloadUseTransportMode, Ikev2NotifyPayloadVariantResponder, ) from . import classes as _ike_test_classes from .classes import Ikev2PayloadNotifyBaseTest, Ikev2PayloadNotifyNoDataTest class TestIkev2PayloadNotifyBase(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.protocol_id = Ikev2ProtocolId.IKE self.notify_type = Ikev2NotifyType.AUTHENTICATION_FAILED self.spi = b'\x00\x01\x02\x03' self.test_data = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' self.notify_payload_minimal = Ikev2PayloadNotifyBaseTest( flags=set(), protocol_id=self.protocol_id, notify_type=self.notify_type, spi=b'', test_data=b'' ) self.notify_payload_minimal.next_payload = Ikev2PayloadType.NONE self.notify_payload_with_data = Ikev2PayloadNotifyBaseTest( flags={Ikev2PayloadFlags.CRITICAL}, protocol_id=self.protocol_id, notify_type=self.notify_type, spi=self.spi, test_data=self.test_data ) self.notify_payload_with_data.next_payload = Ikev2PayloadType.KE self.notify_dict_minimal = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE ('flags', b'\x00'), # No flags ('payload_length', b'\x00\x08'), # 4 bytes header + 4 bytes notify header + 0 bytes data ('protocol_id', b'\x01'), # IKE ('spi_size', b'\x00'), # 0 bytes SPI ('notify_type', b'\x00\x18'), # AUTHENTICATION_FAILED (0x0018) ]) self.notify_bytes_minimal = b''.join(self.notify_dict_minimal.values()) self.notify_dict_with_data = collections.OrderedDict([ ('next_payload', b'\x22'), # KE ('flags', b'\x80'), # CRITICAL ('payload_length', b'\x00\x1c'), # 4 bytes header + 4 bytes notify header + 4 bytes SPI + 16 bytes data ('protocol_id', b'\x01'), # IKE ('spi_size', b'\x04'), # 4 bytes SPI ('notify_type', b'\x00\x18'), # AUTHENTICATION_FAILED (0x0018) ('spi', self.spi), # SPI data ('test_data', self.test_data), # Notification data ]) self.notify_bytes_with_data = b''.join(self.notify_dict_with_data.values()) def test_get_payload_type(self): self.assertEqual(Ikev2PayloadNotifyBaseTest.get_payload_type(), Ikev2PayloadType.NOTIFY) def test_parse(self): parsed_minimal: Ikev2PayloadNotifyBaseTest = Ikev2PayloadNotifyBaseTest.parse_exact_size( self.notify_bytes_minimal ) self.assertEqual(parsed_minimal.flags, set()) self.assertEqual(parsed_minimal.next_payload, Ikev2PayloadType.NONE) self.assertEqual(parsed_minimal.protocol_id, self.protocol_id) self.assertEqual(parsed_minimal.type, self.notify_type) self.assertEqual(parsed_minimal.spi, b'') self.assertEqual(parsed_minimal.test_data, b'') parsed_with_data: Ikev2PayloadNotifyBaseTest = Ikev2PayloadNotifyBaseTest.parse_exact_size( self.notify_bytes_with_data ) self.assertEqual(parsed_with_data.flags, {Ikev2PayloadFlags.CRITICAL}) self.assertEqual(parsed_with_data.next_payload, Ikev2PayloadType.KE) self.assertEqual(parsed_with_data.protocol_id, self.protocol_id) self.assertEqual(parsed_with_data.type, self.notify_type) self.assertEqual(parsed_with_data.spi, self.spi) self.assertEqual(parsed_with_data.test_data, self.test_data) def test_compose(self): composed_minimal = self.notify_payload_minimal.compose() self.assertEqual(composed_minimal, self.notify_bytes_minimal) composed_with_data = self.notify_payload_with_data.compose() self.assertEqual(composed_with_data, self.notify_bytes_with_data) def test_round_trip(self): composed_minimal = self.notify_payload_minimal.compose() parsed_minimal: Ikev2PayloadNotifyBaseTest = Ikev2PayloadNotifyBaseTest.parse_exact_size(composed_minimal) self.assertEqual(parsed_minimal.protocol_id, self.notify_payload_minimal.protocol_id) self.assertEqual(parsed_minimal.type, self.notify_payload_minimal.type) self.assertEqual(parsed_minimal.spi, self.notify_payload_minimal.spi) self.assertEqual(parsed_minimal.test_data, self.notify_payload_minimal.test_data) self.assertEqual(parsed_minimal.flags, self.notify_payload_minimal.flags) self.assertEqual(parsed_minimal.next_payload, self.notify_payload_minimal.next_payload) composed_with_data = self.notify_payload_with_data.compose() parsed_with_data: Ikev2PayloadNotifyBaseTest = Ikev2PayloadNotifyBaseTest.parse_exact_size(composed_with_data) self.assertEqual(parsed_with_data.protocol_id, self.notify_payload_with_data.protocol_id) self.assertEqual(parsed_with_data.type, self.notify_payload_with_data.type) self.assertEqual(parsed_with_data.spi, self.notify_payload_with_data.spi) self.assertEqual(parsed_with_data.test_data, self.notify_payload_with_data.test_data) self.assertEqual(parsed_with_data.flags, self.notify_payload_with_data.flags) self.assertEqual(parsed_with_data.next_payload, self.notify_payload_with_data.next_payload) def test_error_parse_not_enough_data(self): incomplete_data = self.notify_bytes_minimal[:-2] with self.assertRaises(NotEnoughData) as context_manager: Ikev2PayloadNotifyBaseTest.parse_exact_size(incomplete_data) self.assertEqual(context_manager.exception.bytes_needed, 2) def test_error_invalid_protocol_id(self): with self.assertRaises(TypeError): Ikev2PayloadNotifyBaseTest( flags=set(), protocol_id="invalid", notify_type=self.notify_type, spi=b'', test_data=b'' ) def test_error_invalid_notify_type(self): with self.assertRaises(TypeError): Ikev2PayloadNotifyBaseTest( flags=set(), protocol_id=self.protocol_id, notify_type="invalid", spi=b'', test_data=b'' ) def test_error_invalid_spi(self): with self.assertRaises(TypeError): Ikev2PayloadNotifyBaseTest( flags=set(), protocol_id=self.protocol_id, notify_type=self.notify_type, spi=None, test_data=b'' ) class TestIkev2PayloadNotifyNoData(unittest.TestCase): def setUp(self): self.wrong_notify_dict = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE ('flags', b'\x00'), # No flags ('payload_length', b'\x00\x08'), # 4 bytes header + 4 bytes notify header ('protocol_id', b'\x01'), # IKE ('spi_size', b'\x00'), # 0 bytes SPI ('notify_type', b'\x00\x01'), # UNSUPPORTED_CRITICAL_PAYLOAD (not AUTHENTICATION_FAILED) ]) self.wrong_notify_bytes = b''.join(self.wrong_notify_dict.values()) def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual(Ikev2PayloadNotifyNoDataTest._get_message_type(), Ikev2NotifyType.AUTHENTICATION_FAILED) def test_error_invalid_notify_type(self): with self.assertRaises(InvalidType): Ikev2PayloadNotifyNoDataTest.parse_exact_size(self.wrong_notify_bytes) class TestIkev2PayloadNotifyUnparsed(unittest.TestCase): def setUp(self): self.protocol_id = Ikev2ProtocolId.IKE self.notify_type = Ikev2NotifyType.INVALID_SELECTORS # Different from AUTHENTICATION_FAILED self.spi = b'\x00\x01\x02\x03' self.data = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' self.notify_payload_with_data = Ikev2PayloadNotifyUnparsed( flags={Ikev2PayloadFlags.CRITICAL}, protocol_id=self.protocol_id, type=self.notify_type, spi=self.spi, data=self.data ) self.notify_payload_with_data.next_payload = Ikev2PayloadType.KE def test_any_notify_type_support(self): different_notify_types = [ Ikev2NotifyType.AUTHENTICATION_FAILED, Ikev2NotifyType.INVALID_SELECTORS, Ikev2NotifyType.UNSUPPORTED_CRITICAL_PAYLOAD, ] for notify_type in different_notify_types: payload = Ikev2PayloadNotifyUnparsed( flags=set(), protocol_id=self.protocol_id, type=notify_type, spi=b'', data=b'\x00\x01\x02\x03' ) self.assertEqual(payload.type, notify_type) self.assertEqual(payload.data, b'\x00\x01\x02\x03') def test_raw_data_storage(self): payload = Ikev2PayloadNotifyUnparsed( flags=set(), protocol_id=self.protocol_id, type=self.notify_type, spi=b'', data=self.data ) self.assertEqual(payload.data, self.data) different_data = b'\xff\xfe\xfd\xfc' payload_2 = Ikev2PayloadNotifyUnparsed( flags=set(), protocol_id=self.protocol_id, type=self.notify_type, spi=b'', data=different_data ) self.assertEqual(payload_2.data, different_data) def test_round_trip_data_preservation(self): payload_no_spi = Ikev2PayloadNotifyUnparsed( flags=set(), protocol_id=self.protocol_id, type=self.notify_type, spi=b'', data=self.data ) payload_no_spi.next_payload = Ikev2PayloadType.NONE composed_bytes = payload_no_spi.compose() parsed_payload: Ikev2PayloadNotifyUnparsed = Ikev2PayloadNotifyUnparsed.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.data, payload_no_spi.data) # pylint: disable=no-member self.assertEqual(parsed_payload.type, payload_no_spi.type) self.assertEqual(parsed_payload.spi, payload_no_spi.spi) def test_parse_with_spi(self): data_with_spi = b'\x00\x01\x02\x03\x04\x05\x06\x07' payload_with_spi = Ikev2PayloadNotifyUnparsed( flags=set(), protocol_id=self.protocol_id, type=self.notify_type, spi=self.spi, data=data_with_spi ) payload_with_spi.next_payload = Ikev2PayloadType.NONE composed_bytes = payload_with_spi.compose() parsed_payload: Ikev2PayloadNotifyUnparsed = Ikev2PayloadNotifyUnparsed.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.protocol_id, self.protocol_id) self.assertEqual(parsed_payload.type, self.notify_type) self.assertEqual(parsed_payload.spi, self.spi) self.assertEqual(parsed_payload.data, data_with_spi) # pylint: disable=no-member self.assertEqual(parsed_payload.flags, set()) self.assertEqual(parsed_payload.next_payload, Ikev2PayloadType.NONE) class TestIkev2NotifyPayloadCookie(unittest.TestCase): _PROTOCOL_ID = Ikev2ProtocolId.IKE _COOKIE_DATA = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual(Ikev2NotifyPayloadCookie._get_message_type(), Ikev2NotifyType.COOKIE) def test_cookie_data_storage(self): payload = Ikev2NotifyPayloadCookie( flags=set(), protocol_id=self._PROTOCOL_ID, type=Ikev2NotifyType.COOKIE, spi=b'', cookie=self._COOKIE_DATA ) self.assertEqual(payload.cookie, self._COOKIE_DATA) different_cookie = b'\xff\xfe\xfd\xfc\xfb\xfa' payload_2 = Ikev2NotifyPayloadCookie( flags=set(), protocol_id=self._PROTOCOL_ID, type=Ikev2NotifyType.COOKIE, spi=b'', cookie=different_cookie ) self.assertEqual(payload_2.cookie, different_cookie) def test_round_trip_cookie_preservation(self): cookie_payload = Ikev2NotifyPayloadCookie( flags=set(), protocol_id=self._PROTOCOL_ID, type=Ikev2NotifyType.COOKIE, spi=b'', cookie=self._COOKIE_DATA ) cookie_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = cookie_payload.compose() parsed_payload: Ikev2NotifyPayloadCookie = Ikev2NotifyPayloadCookie.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.cookie, cookie_payload.cookie) # pylint: disable=no-member self.assertEqual(parsed_payload.type, cookie_payload.type) self.assertEqual(parsed_payload.spi, cookie_payload.spi) class TestIkev2NotifyPayloadSetWindowSize(unittest.TestCase): _PROTOCOL_ID = Ikev2ProtocolId.IKE _WINDOW_SIZE = 5 def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual(Ikev2NotifyPayloadSetWindowSize._get_message_type(), Ikev2NotifyType.SET_WINDOW_SIZE) def test_window_size_storage(self): payload = Ikev2NotifyPayloadSetWindowSize( flags=set(), protocol_id=self._PROTOCOL_ID, type=Ikev2NotifyType.SET_WINDOW_SIZE, spi=b'', window_size=self._WINDOW_SIZE ) self.assertEqual(payload.window_size, self._WINDOW_SIZE) different_window_size = 10 payload_2 = Ikev2NotifyPayloadSetWindowSize( flags=set(), protocol_id=self._PROTOCOL_ID, type=Ikev2NotifyType.SET_WINDOW_SIZE, spi=b'', window_size=different_window_size ) self.assertEqual(payload_2.window_size, different_window_size) def test_round_trip_window_size_preservation(self): window_size_payload = Ikev2NotifyPayloadSetWindowSize( flags=set(), protocol_id=self._PROTOCOL_ID, type=Ikev2NotifyType.SET_WINDOW_SIZE, spi=b'', window_size=self._WINDOW_SIZE ) window_size_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = window_size_payload.compose() parsed_payload: Ikev2NotifyPayloadSetWindowSize = \ Ikev2NotifyPayloadSetWindowSize.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.window_size, window_size_payload.window_size) # pylint: disable=no-member self.assertEqual(parsed_payload.type, window_size_payload.type) self.assertEqual(parsed_payload.spi, window_size_payload.spi) def test_error_invalid_notification_data_length(self): wrong_length_bytes = bytes.fromhex( '00' # next_payload = NONE '00' # flags = 0 '000b' # payload_length = 11 (8 header + 3 data bytes) '01' # protocol_id = IKE '00' # spi_size = 0 '4001' # notify_type = SET_WINDOW_SIZE 'aaaaaa' # 3 bytes data (must be exactly 4) ) with self.assertRaises(InvalidValue): Ikev2NotifyPayloadSetWindowSize.parse_exact_size(wrong_length_bytes) class TestIkev2NotifyPayloadNatDetectionSourceIp(_ike_test_classes.Ikev2NotifyPayloadNatDetectionBaseTest): _NOTIFY_TYPE = Ikev2NotifyType.NAT_DETECTION_SOURCE_IP _PAYLOAD_CLASS = Ikev2NotifyPayloadNatDetectionSourceIp _NOTIFY_TYPE_BYTES = b'\x40\x04' class TestIkev2NotifyPayloadNatDetectionDestinationIp(_ike_test_classes.Ikev2NotifyPayloadNatDetectionBaseTest): _NOTIFY_TYPE = Ikev2NotifyType.NAT_DETECTION_DESTINATION_IP _PAYLOAD_CLASS = Ikev2NotifyPayloadNatDetectionDestinationIp _NOTIFY_TYPE_BYTES = b'\x40\x05' class TestIkev2NotifyPayloadUseTransportMode(unittest.TestCase): def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual(Ikev2NotifyPayloadUseTransportMode._get_message_type(), Ikev2NotifyType.USE_TRANSPORT_MODE) def test_round_trip_preservation(self): transport_mode_payload = Ikev2NotifyPayloadUseTransportMode( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=Ikev2NotifyType.USE_TRANSPORT_MODE, spi=b'' ) transport_mode_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = transport_mode_payload.compose() parsed_payload: Ikev2NotifyPayloadUseTransportMode = \ Ikev2NotifyPayloadUseTransportMode.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.type, transport_mode_payload.type) self.assertEqual(parsed_payload.spi, transport_mode_payload.spi) self.assertEqual(parsed_payload.protocol_id, transport_mode_payload.protocol_id) self.assertEqual(parsed_payload.flags, transport_mode_payload.flags) self.assertEqual(parsed_payload.next_payload, transport_mode_payload.next_payload) class TestIkev2NotifyPayloadHttpCertLookupSupported(unittest.TestCase): def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual( Ikev2NotifyPayloadHttpCertLookupSupported._get_message_type(), Ikev2NotifyType.HTTP_CERT_LOOKUP_SUPPORTED ) def test_round_trip_preservation(self): http_cert_payload = Ikev2NotifyPayloadHttpCertLookupSupported( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=Ikev2NotifyType.HTTP_CERT_LOOKUP_SUPPORTED, spi=b'' ) http_cert_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = http_cert_payload.compose() parsed_payload: Ikev2NotifyPayloadHttpCertLookupSupported = \ Ikev2NotifyPayloadHttpCertLookupSupported.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.type, http_cert_payload.type) self.assertEqual(parsed_payload.spi, http_cert_payload.spi) self.assertEqual(parsed_payload.protocol_id, http_cert_payload.protocol_id) self.assertEqual(parsed_payload.flags, http_cert_payload.flags) self.assertEqual(parsed_payload.next_payload, http_cert_payload.next_payload) class TestIkev2NotifyPayloadIkev2FragmentationSupported(unittest.TestCase): def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual( Ikev2NotifyPayloadIkev2FragmentationSupported._get_message_type(), Ikev2NotifyType.IKEV2_FRAGMENTATION_SUPPORTED, ) def test_round_trip_preservation(self): fragmentation_payload = Ikev2NotifyPayloadIkev2FragmentationSupported( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=Ikev2NotifyType.IKEV2_FRAGMENTATION_SUPPORTED, spi=b'', ) fragmentation_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = fragmentation_payload.compose() parsed_payload = Ikev2NotifyPayloadIkev2FragmentationSupported.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.type, fragmentation_payload.type) class TestIkev2NotifyPayloadIntermediateExchangeSupported(unittest.TestCase): def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual( Ikev2NotifyPayloadIntermediateExchangeSupported._get_message_type(), Ikev2NotifyType.INTERMEDIATE_EXCHANGE_SUPPORTED, ) def test_round_trip_preservation(self): intermediate_exchange_payload = Ikev2NotifyPayloadIntermediateExchangeSupported( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=Ikev2NotifyType.INTERMEDIATE_EXCHANGE_SUPPORTED, spi=b'', ) intermediate_exchange_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = intermediate_exchange_payload.compose() parsed_payload = Ikev2NotifyPayloadIntermediateExchangeSupported.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.type, intermediate_exchange_payload.type) class TestIkev2NotifyPayloadUsePpk(unittest.TestCase): def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual(Ikev2NotifyPayloadUsePpk._get_message_type(), Ikev2NotifyType.USE_PPK) def test_round_trip_preservation(self): use_ppk_payload = Ikev2NotifyPayloadUsePpk( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=Ikev2NotifyType.USE_PPK, spi=b'', ) use_ppk_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = use_ppk_payload.compose() parsed_payload = Ikev2NotifyPayloadUsePpk.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.type, use_ppk_payload.type) class TestIkev2NotifyPayloadRedirectSupported(unittest.TestCase): def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual( Ikev2NotifyPayloadRedirectSupported._get_message_type(), Ikev2NotifyType.REDIRECT_SUPPORTED, ) def test_round_trip_preservation(self): redirect_supported_payload = Ikev2NotifyPayloadRedirectSupported( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=Ikev2NotifyType.REDIRECT_SUPPORTED, spi=b'', ) redirect_supported_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = redirect_supported_payload.compose() parsed_payload = Ikev2NotifyPayloadRedirectSupported.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.type, redirect_supported_payload.type) class TestIkev2NotifyPayloadChildlessIkev2Supported(unittest.TestCase): def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual( Ikev2NotifyPayloadChildlessIkev2Supported._get_message_type(), Ikev2NotifyType.CHILDLESS_IKEV2_SUPPORTED, ) def test_round_trip_preservation(self): childless_ikev2_payload = Ikev2NotifyPayloadChildlessIkev2Supported( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=Ikev2NotifyType.CHILDLESS_IKEV2_SUPPORTED, spi=b'', ) childless_ikev2_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = childless_ikev2_payload.compose() parsed_payload = Ikev2NotifyPayloadChildlessIkev2Supported.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.type, childless_ikev2_payload.type) class TestIkev2NotifyPayloadSignatureHashAlgorithms(unittest.TestCase): # Five 16-bit hash algorithm identifiers per RFC 7427 (IANA registry # values: 1=SHA1, 2=SHA2-256, 3=SHA2-384, 4=SHA2-512, 5=IDENTITY) _HASH_ALGORITHMS = ( Ikev2HashAlgorithm.SHA1, Ikev2HashAlgorithm.SHA2_256, Ikev2HashAlgorithm.SHA2_384, Ikev2HashAlgorithm.SHA2_512, Ikev2HashAlgorithm.IDENTITY, ) _PAYLOAD_BYTES = bytes.fromhex( '00' # next_payload = NONE '00' # flags = 0 '0012' # payload_length = 18 (8 header + 10 data) '01' # protocol_id = IKE '00' # spi_size = 0 '402f' # notify_type = SIGNATURE_HASH_ALGORITHMS (16431) '00010002000300040005' # five 16-bit hash algorithm IDs ) def setUp(self): self.payload = Ikev2NotifyPayloadSignatureHashAlgorithms( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=Ikev2NotifyType.SIGNATURE_HASH_ALGORITHMS, spi=b'', hash_algorithms=self._HASH_ALGORITHMS, ) self.payload.next_payload = Ikev2PayloadType.NONE def test_get_message_type(self): # pylint: disable=protected-access self.assertEqual( Ikev2NotifyPayloadSignatureHashAlgorithms._get_message_type(), Ikev2NotifyType.SIGNATURE_HASH_ALGORITHMS, ) def test_parse(self): parsed = Ikev2NotifyPayloadSignatureHashAlgorithms.parse_exact_size(self._PAYLOAD_BYTES) self.assertEqual(parsed.type, Ikev2NotifyType.SIGNATURE_HASH_ALGORITHMS) self.assertEqual(parsed.hash_algorithms, self._HASH_ALGORITHMS) # pylint: disable=no-member def test_compose(self): self.assertEqual(self.payload.compose(), self._PAYLOAD_BYTES) def test_round_trip(self): composed = self.payload.compose() parsed = Ikev2NotifyPayloadSignatureHashAlgorithms.parse_exact_size(composed) self.assertEqual(parsed.hash_algorithms, self.payload.hash_algorithms) # pylint: disable=no-member def test_error_invalid_notification_data_length(self): odd_length_bytes = bytes.fromhex( '00' # next_payload = NONE '00' # flags = 0 '0009' # payload_length = 9 (8 header + 1 data byte) '01' # protocol_id = IKE '00' # spi_size = 0 '402f' # notify_type = SIGNATURE_HASH_ALGORITHMS 'aa' # 1 byte data (must be even number of bytes) ) with self.assertRaises(InvalidValue): Ikev2NotifyPayloadSignatureHashAlgorithms.parse_exact_size(odd_length_bytes) class TestIkev2NotifyPayloadVariantResponder(unittest.TestCase): _PROTOCOL_ID = Ikev2ProtocolId.IKE def test_parse_other_notify_type(self): other_notify = Ikev2PayloadNotifyUnparsed( flags=set(), protocol_id=self._PROTOCOL_ID, type=Ikev2NotifyType.AUTHENTICATION_FAILED, spi=b'', data=b'\x00\x01\x02\x03' ) other_notify.next_payload = Ikev2PayloadType.NONE composed_bytes = other_notify.compose() parsed_payload = Ikev2NotifyPayloadVariantResponder.parse_exact_size(composed_bytes) self.assertIsInstance(parsed_payload, Ikev2PayloadNotifyUnparsed) self.assertEqual(parsed_payload.type, Ikev2NotifyType.AUTHENTICATION_FAILED) self.assertEqual(parsed_payload.data, b'\x00\x01\x02\x03') # pylint: disable=no-member def test_compose(self): cookie_data = b'\x00\x01\x02\x03' cookie_payload = Ikev2NotifyPayloadCookie( flags=set(), protocol_id=self._PROTOCOL_ID, type=Ikev2NotifyType.COOKIE, spi=b'', cookie=cookie_data ) cookie_payload.next_payload = Ikev2PayloadType.NONE variant_parsable = Ikev2NotifyPayloadVariantResponder(variant=cookie_payload) composed_bytes = variant_parsable.compose() self.assertEqual(composed_bytes, cookie_payload.compose()) def test_round_trip(self): cookie_data = b'\x00\x01\x02\x03\x04\x05' cookie_payload = Ikev2NotifyPayloadCookie( flags=set(), protocol_id=self._PROTOCOL_ID, type=Ikev2NotifyType.COOKIE, spi=b'', cookie=cookie_data ) cookie_payload.next_payload = Ikev2PayloadType.NONE variant_parsable = Ikev2NotifyPayloadVariantResponder(variant=cookie_payload) composed_bytes = variant_parsable.compose() parsed_payload = Ikev2NotifyPayloadVariantResponder.parse_exact_size(composed_bytes) self.assertIsInstance(parsed_payload, Ikev2NotifyPayloadCookie) self.assertEqual(parsed_payload.cookie, cookie_data) # pylint: disable=no-member self.assertEqual(parsed_payload.type, cookie_payload.type) self.assertEqual(parsed_payload.spi, cookie_payload.spi) if __name__ == '__main__': unittest.main() cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ike/test_ikev2_payload.py000066400000000000000000001507461524413560000277600ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import collections import ipaddress import unittest from cryptodatahub.common.algorithm import Signature from cryptodatahub.common.exception import InvalidValue from cryptodatahub.ike.algorithm import ( Ikev2AuthenticationMethod, Ikev2CertificateType, Ikev2DiffieHellmanGroup, Ikev2IdType, Ikev2NotifyType, Ikev2ProtocolId, Ikev2TransformAttributeType, ) from cryptoparser.common.exception import InvalidType, NotEnoughData, TooMuchData from cryptoparser.ike.ikev2 import ( Ikev2AuthDigitalSignatureEnvelope, Ikev2NotifyPayloadInvalidKe, Ikev2PayloadAuthentication, Ikev2PayloadCertificate, Ikev2PayloadCertificateRequest, Ikev2PayloadDelete, Ikev2PayloadEap, Ikev2PayloadEncryptedAndAuthenticated, Ikev2PayloadFlags, Ikev2PayloadIdentificationInitiator, Ikev2PayloadIdentificationInitiatorDerAsn1Gn, Ikev2PayloadIdentificationInitiatorFcName, Ikev2PayloadIdentificationInitiatorFqdn, Ikev2PayloadIdentificationInitiatorIpv6Addr, Ikev2PayloadIdentificationInitiatorKeyId, Ikev2PayloadIdentificationInitiatorNull, Ikev2PayloadIdentificationInitiatorRfc822Addr, Ikev2PayloadIdentificationResponder, Ikev2PayloadIdentificationResponderDerAsn1Dn, Ikev2PayloadIdentificationResponderFqdn, Ikev2PayloadIdentificationResponderIpv4Addr, Ikev2PayloadIdentificationResponderIpv6Addr, Ikev2PayloadKeyExchange, Ikev2PayloadNonce, Ikev2PayloadNotifyAuthenticationFailed, Ikev2PayloadNotifyUnparsed, Ikev2PayloadType, Ikev2PayloadVendorId, TransformAttributeKeyLength, ) from .classes import Ikev2PayloadBaseTest, Ikev2PayloadNotifyNoDataTest class TestIkev2PayloadBase(unittest.TestCase): def setUp(self): self.test_data = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' self.test_payload_minimal = Ikev2PayloadBaseTest( flags=set(), test_data=b'', ) self.test_payload_minimal.next_payload = Ikev2PayloadType.NONE self.test_dict_minimal = collections.OrderedDict([ ('next_payload', b'\x00'), ('flags', b'\x00'), ('payload_length', b'\x00\x04'), ('test_data', b''), ]) self.test_bytes_minimal = b''.join(self.test_dict_minimal.values()) self.test_payload_with_data = Ikev2PayloadBaseTest( flags={Ikev2PayloadFlags.CRITICAL}, test_data=self.test_data ) self.test_payload_with_data.next_payload = Ikev2PayloadType.NONE self.test_dict_with_data = collections.OrderedDict([ ('next_payload', b'\x00'), ('flags', b'\x80'), ('payload_length', b'\x00\x14'), ('test_data', self.test_data), ]) self.test_bytes_with_data = b''.join(self.test_dict_with_data.values()) def test_parse(self): parsed_payload = Ikev2PayloadBaseTest.parse_exact_size(self.test_bytes_minimal) self.assertEqual(parsed_payload.flags, self.test_payload_minimal.flags) self.assertEqual(parsed_payload.test_data, self.test_payload_minimal.test_data) self.assertEqual(parsed_payload.next_payload, self.test_payload_minimal.next_payload) parsed_payload = Ikev2PayloadBaseTest.parse_exact_size(self.test_bytes_with_data) self.assertEqual(parsed_payload.flags, self.test_payload_with_data.flags) self.assertEqual(parsed_payload.test_data, self.test_payload_with_data.test_data) self.assertEqual(parsed_payload.next_payload, self.test_payload_with_data.next_payload) def test_compose(self): self.assertEqual(self.test_payload_minimal.compose(), self.test_bytes_minimal) self.assertEqual(self.test_payload_with_data.compose(), self.test_bytes_with_data) def test_error_parse_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: Ikev2PayloadBaseTest.parse_exact_size(b'\x00') self.assertEqual(context_manager.exception.bytes_needed, 3) def test_error_payload_validation(self): with self.assertRaises(TypeError) as context_manager: Ikev2PayloadBaseTest( flags={'invalid_flag'}, test_data=self.test_data ) exception_str = str(context_manager.exception) self.assertIn("flags", exception_str) self.assertIn("must be", exception_str) self.assertIn("Ikev2PayloadFlags", exception_str) def test_next_payload(self): payload = Ikev2PayloadBaseTest( flags={Ikev2PayloadFlags.CRITICAL}, test_data=self.test_data ) payload.next_payload = Ikev2PayloadType.KE composed = payload.compose() self.assertEqual(composed[0], Ikev2PayloadType.KE.value.code) parsed, _ = Ikev2PayloadBaseTest._parse(composed) # pylint: disable=protected-access self.assertEqual(parsed.next_payload, Ikev2PayloadType.KE) class TestIkev2PayloadKeyExchange(unittest.TestCase): def setUp(self): self.dh_group = Ikev2DiffieHellmanGroup.MODP_GROUP_2048_BIT self.key_exchange_data = b'\x00\x01\x02\x03\x04\x05\x06\x07' self.key_exchange_payload = Ikev2PayloadKeyExchange( flags=set(), dh_group=self.dh_group, key_exchange_data=self.key_exchange_data ) self.key_exchange_payload.next_payload = Ikev2PayloadType.NONE self.key_exchange_dict = collections.OrderedDict([ ('next_payload', b'\x00'), ('flags', b'\x00'), ('payload_length', b'\x00\x10'), ('dh_group', b'\x00\x0e'), ('reserved2', b'\x00\x00'), ('key_exchange_data', self.key_exchange_data), ]) self.key_exchange_bytes = b''.join(self.key_exchange_dict.values()) def test_get_payload_type(self): self.assertEqual(Ikev2PayloadKeyExchange.get_payload_type(), Ikev2PayloadType.KE) def test_parse(self): parsed_ke = Ikev2PayloadKeyExchange.parse_exact_size(self.key_exchange_bytes) self.assertEqual(parsed_ke.dh_group, self.dh_group) self.assertEqual(parsed_ke.key_exchange_data, self.key_exchange_data) def test_compose(self): composed_bytes = self.key_exchange_payload.compose() self.assertGreater(len(composed_bytes), Ikev2PayloadKeyExchange.HEADER_SIZE) parsed_ke = Ikev2PayloadKeyExchange.parse_exact_size(composed_bytes) self.assertEqual(parsed_ke.key_exchange_data, self.key_exchange_payload.key_exchange_data) def test_round_trip(self): composed_bytes = self.key_exchange_payload.compose() parsed_payload: Ikev2PayloadKeyExchange = Ikev2PayloadKeyExchange.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.dh_group, self.key_exchange_payload.dh_group) self.assertEqual(parsed_payload.key_exchange_data, self.key_exchange_payload.key_exchange_data) self.assertEqual(parsed_payload.flags, self.key_exchange_payload.flags) self.assertEqual(parsed_payload.next_payload, self.key_exchange_payload.next_payload) class TestIkev2PayloadNonce(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.nonce_data = b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' self.nonce_payload = Ikev2PayloadNonce( flags=set(), nonce_data=self.nonce_data ) self.nonce_payload.next_payload = Ikev2PayloadType.NONE self.nonce_dict = collections.OrderedDict([ ('next_payload', b'\x00'), ('flags', b'\x00'), ('payload_length', b'\x00\x14'), ('nonce_data', self.nonce_data), ]) self.nonce_bytes = b''.join(self.nonce_dict.values()) self.nonce_dict_too_small = collections.OrderedDict([ ('next_payload', b'\x00'), ('flags', b'\x00'), ('payload_length', b'\x00\x13'), ('nonce_data', b'\x00' * 15), ]) self.nonce_bytes_too_small = b''.join(self.nonce_dict_too_small.values()) self.nonce_dict_too_large = collections.OrderedDict([ ('next_payload', b'\x00'), ('flags', b'\x00'), ('payload_length', b'\x01\x05'), ('nonce_data', b'\x00' * 257), ]) self.nonce_bytes_too_large = b''.join(self.nonce_dict_too_large.values()) def test_get_payload_type(self): self.assertEqual(Ikev2PayloadNonce.get_payload_type(), Ikev2PayloadType.NONCE) def test_parse(self): parsed_nonce = Ikev2PayloadNonce.parse_exact_size(self.nonce_bytes) self.assertEqual(parsed_nonce.nonce_data, self.nonce_data) self.assertEqual(parsed_nonce.next_payload, Ikev2PayloadType.NONE) self.assertEqual(parsed_nonce.flags, set()) def test_compose(self): composed_bytes = self.nonce_payload.compose() self.assertIsInstance(composed_bytes, (bytes, bytearray)) self.assertGreater(len(composed_bytes), Ikev2PayloadNonce.HEADER_SIZE) parsed_nonce = Ikev2PayloadNonce.parse_exact_size(composed_bytes) self.assertEqual(parsed_nonce.nonce_data, self.nonce_payload.nonce_data) self.assertEqual(parsed_nonce.next_payload, self.nonce_payload.next_payload) self.assertEqual(parsed_nonce.flags, self.nonce_payload.flags) def test_error_invalid_nonce_data(self): with self.assertRaises(Exception) as context_manager: Ikev2PayloadNonce( flags=set(), nonce_data="not_bytes" ) self.assertTrue(len(str(context_manager.exception)) > 0) def test_minimal_nonce_data(self): nonce_empty = Ikev2PayloadNonce( flags=set(), nonce_data=b'\x00' * 16 ) nonce_empty.next_payload = Ikev2PayloadType.NONE composed_bytes = nonce_empty.compose() parsed_nonce = Ikev2PayloadNonce.parse_exact_size(composed_bytes) self.assertEqual(parsed_nonce.nonce_data, b'\x00' * 16) def test_too_small_nonce_data(self): with self.assertRaises(NotEnoughData) as context_manager: Ikev2PayloadNonce.parse_exact_size(self.nonce_bytes_too_small) self.assertEqual(context_manager.exception.bytes_needed, 16 - 15) def test_too_large_nonce_data(self): with self.assertRaises(TooMuchData) as context_manager: Ikev2PayloadNonce.parse_exact_size(self.nonce_bytes_too_large) self.assertEqual(context_manager.exception.bytes_needed, 257 - 256) def test_error_nonce_data_validation_min_length(self): with self.assertRaises(ValueError): Ikev2PayloadNonce(flags=set(), nonce_data=b'\x01\x02\x03\x04\x05\x06\x07\x08') def test_error_nonce_data_validation_max_length(self): with self.assertRaises(ValueError): Ikev2PayloadNonce(flags=set(), nonce_data=b'\x01' * 257) def test_round_trip(self): composed_bytes = self.nonce_payload.compose() parsed_payload: Ikev2PayloadNonce = Ikev2PayloadNonce.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.nonce_data, self.nonce_payload.nonce_data) self.assertEqual(parsed_payload.flags, self.nonce_payload.flags) self.assertEqual(parsed_payload.next_payload, self.nonce_payload.next_payload) class TestIkev2PayloadDelete(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.protocol_id = Ikev2ProtocolId.IKE self.spis = [0x1234567890abcdef, 0xfedcba0987654321] self.delete_payload = Ikev2PayloadDelete( flags=set(), protocol_id=self.protocol_id, spis=self.spis ) self.delete_payload.next_payload = Ikev2PayloadType.NONE self.delete_dict = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE ('flags', b'\x00'), ('payload_length', b'\x00\x18'), # 24 bytes (header + protocol_id + spi_size + num_spis + spis) ('protocol_id', b'\x01'), # IKE protocol ('spi_size', b'\x10'), # 16 bytes (2 SPIs * 8 bytes each) ('num_spis', b'\x00\x02'), # 2 SPIs ('spis', b'\x12\x34\x56\x78\x90\xab\xcd\xef\xfe\xdc\xba\x09\x87\x65\x43\x21'), ]) self.delete_bytes = b''.join(self.delete_dict.values()) self.empty_delete_payload = Ikev2PayloadDelete( flags={Ikev2PayloadFlags.CRITICAL}, protocol_id=self.protocol_id, spis=[] ) self.empty_delete_payload.next_payload = Ikev2PayloadType.NONE self.empty_delete_dict = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE ('flags', b'\x80'), # CRITICAL (0x80) ('payload_length', b'\x00\x08'), # 8 bytes (header + protocol_id + spi_size + num_spis) ('protocol_id', b'\x01'), # IKE protocol ('spi_size', b'\x00'), ('num_spis', b'\x00\x00'), ]) self.empty_delete_bytes = b''.join(self.empty_delete_dict.values()) def test_get_payload_type(self): self.assertEqual(Ikev2PayloadDelete.get_payload_type(), Ikev2PayloadType.DELETE) def test_parse(self): parsed_delete = Ikev2PayloadDelete.parse_exact_size(self.delete_bytes) self.assertEqual(parsed_delete.protocol_id, self.protocol_id) self.assertEqual(parsed_delete.spis, self.spis) self.assertEqual(parsed_delete.next_payload, Ikev2PayloadType.NONE) self.assertEqual(parsed_delete.flags, set()) def test_compose(self): composed_bytes = self.delete_payload.compose() self.assertIsInstance(composed_bytes, (bytes, bytearray)) self.assertGreater(len(composed_bytes), Ikev2PayloadDelete.HEADER_SIZE) parsed_delete = Ikev2PayloadDelete.parse_exact_size(composed_bytes) self.assertEqual(parsed_delete.protocol_id, self.delete_payload.protocol_id) self.assertEqual(parsed_delete.spis, self.delete_payload.spis) self.assertEqual(parsed_delete.next_payload, self.delete_payload.next_payload) self.assertEqual(parsed_delete.flags, self.delete_payload.flags) def test_empty_spis(self): composed_bytes = self.empty_delete_payload.compose() parsed_delete = Ikev2PayloadDelete.parse_exact_size(composed_bytes) self.assertEqual(parsed_delete.spis, []) self.assertEqual(parsed_delete.protocol_id, self.protocol_id) self.assertEqual(parsed_delete.flags, {Ikev2PayloadFlags.CRITICAL}) def test_error_parse_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: Ikev2PayloadDelete.parse_exact_size(b'\x00\x01\x02') self.assertEqual(context_manager.exception.bytes_needed, 1) def test_error_invalid_protocol_id(self): with self.assertRaises(TypeError): Ikev2PayloadDelete( flags=set(), protocol_id="invalid_protocol", spis=self.spis ) def test_error_invalid_spis_type(self): with self.assertRaises(TypeError): Ikev2PayloadDelete( flags=set(), protocol_id=self.protocol_id, spis=["not_an_int"] ) def test_round_trip(self): composed_bytes = self.delete_payload.compose() parsed_payload: Ikev2PayloadDelete = Ikev2PayloadDelete.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.protocol_id, self.delete_payload.protocol_id) self.assertEqual(parsed_payload.spis, self.delete_payload.spis) self.assertEqual(parsed_payload.flags, self.delete_payload.flags) self.assertEqual(parsed_payload.next_payload, self.delete_payload.next_payload) class TestIkev2PayloadNotifyNoData(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.protocol_id = Ikev2ProtocolId.IKE self.notify_type = Ikev2NotifyType.AUTHENTICATION_FAILED self.spi = b'\x00\x01\x02\x03\x04\x05\x06\x07' self.notify_payload_minimal = Ikev2PayloadNotifyNoDataTest( flags=set(), protocol_id=self.protocol_id, notify_type=self.notify_type, spi=b'' ) self.notify_payload_minimal.next_payload = Ikev2PayloadType.NONE self.notify_dict_minimal = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE ('flags', b'\x00'), ('payload_length', b'\x00\x08'), # 8 bytes (header + protocol_id + spi_size + notify_type) ('protocol_id', b'\x01'), # IKE protocol ('spi_size', b'\x00'), ('notify_type', b'\x00\x18'), # AUTHENTICATION_FAILED (0x0018) ]) self.notify_bytes_minimal = b''.join(self.notify_dict_minimal.values()) self.notify_payload_with_spi = Ikev2PayloadNotifyNoDataTest( flags={Ikev2PayloadFlags.CRITICAL}, protocol_id=self.protocol_id, notify_type=self.notify_type, spi=self.spi ) self.notify_payload_with_spi.next_payload = Ikev2PayloadType.NONE self.notify_dict_with_spi = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE ('flags', b'\x80'), # CRITICAL (0x80) ('payload_length', b'\x00\x10'), # 16 bytes (header + protocol_id + spi_size + notify_type + spi) ('protocol_id', b'\x01'), # IKE protocol ('spi_size', b'\x08'), ('notify_type', b'\x00\x18'), # AUTHENTICATION_FAILED (0x0018) ('spi', self.spi), ]) self.notify_bytes_with_spi = b''.join(self.notify_dict_with_spi.values()) def test_parse(self): parsed_notify = Ikev2PayloadNotifyNoDataTest.parse_exact_size(self.notify_bytes_minimal) self.assertEqual(parsed_notify.protocol_id, self.protocol_id) self.assertEqual(parsed_notify.type, self.notify_type) self.assertEqual(parsed_notify.spi, b'') self.assertEqual(parsed_notify.next_payload, Ikev2PayloadType.NONE) self.assertEqual(parsed_notify.flags, set()) parsed_notify = Ikev2PayloadNotifyNoDataTest.parse_exact_size(self.notify_bytes_with_spi) self.assertEqual(parsed_notify.protocol_id, self.protocol_id) self.assertEqual(parsed_notify.type, self.notify_type) self.assertEqual(parsed_notify.spi, self.spi) self.assertEqual(parsed_notify.next_payload, Ikev2PayloadType.NONE) self.assertEqual(parsed_notify.flags, {Ikev2PayloadFlags.CRITICAL}) def test_compose(self): self.assertEqual(self.notify_payload_minimal.compose(), self.notify_bytes_minimal) self.assertEqual(self.notify_payload_with_spi.compose(), self.notify_bytes_with_spi) def test_error_parse_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: Ikev2PayloadNotifyNoDataTest.parse_exact_size(b'\x00\x01\x02') self.assertEqual(context_manager.exception.bytes_needed, 1) def test_error_payload_validation(self): with self.assertRaises(TypeError) as context_manager: Ikev2PayloadNotifyNoDataTest( flags={'invalid_flag'}, protocol_id=self.protocol_id, notify_type=self.notify_type, spi=self.spi ) exception_str = str(context_manager.exception) self.assertIn("flags", exception_str) self.assertIn("must be", exception_str) self.assertIn("Ikev2PayloadFlags", exception_str) def test_error_invalid_protocol_id(self): with self.assertRaises(TypeError): Ikev2PayloadNotifyNoDataTest( flags=set(), protocol_id="invalid_protocol", notify_type=self.notify_type, spi=self.spi ) def test_error_invalid_notify_type(self): with self.assertRaises(TypeError): Ikev2PayloadNotifyNoDataTest( flags=set(), protocol_id=self.protocol_id, notify_type="invalid_type", spi=self.spi ) def test_error_invalid_spi(self): # The spi attribute has converter=bytes, so test with something that can't be converted with self.assertRaises((TypeError, ValueError)): Ikev2PayloadNotifyNoDataTest( flags=set(), protocol_id=self.protocol_id, notify_type=self.notify_type, spi=None ) class TestIkev2PayloadNotifyAuthenticationFailed(unittest.TestCase): def setUp(self): self.protocol_id = Ikev2ProtocolId.IKE self.notify_dict_authentication_failed = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE ('flags', b'\x00'), ('payload_length', b'\x00\x08'), # 8 bytes (header + protocol_id + spi_size + notify_type) ('protocol_id', b'\x01'), # IKE protocol ('spi_size', b'\x00'), ('notify_type', b'\x00\x18'), # AUTHENTICATION_FAILED (0x0018) ]) self.notify_bytes_authentication_failed = b''.join(self.notify_dict_authentication_failed.values()) def test_get_message_type(self): self.assertEqual(Ikev2PayloadNotifyAuthenticationFailed._get_message_type(), # pylint: disable=protected-access Ikev2NotifyType.AUTHENTICATION_FAILED) def test_parse_authentication_failed(self): parsed_notify = Ikev2PayloadNotifyAuthenticationFailed.parse_exact_size( self.notify_bytes_authentication_failed) self.assertEqual(parsed_notify.type, Ikev2NotifyType.AUTHENTICATION_FAILED) self.assertEqual(parsed_notify.protocol_id, self.protocol_id) class TestIkev2PayloadNotifyUnparsed(unittest.TestCase): def setUp(self): self.notify_type = Ikev2NotifyType.INVALID_SYNTAX self.notify_data = ( b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' b'\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f' ) self.notify_dict_with_data = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE ('flags', b'\x00'), ('payload_length', b'\x00\x28'), # 40 bytes (8 header + 32 data) ('protocol_id', b'\x01'), # IKE protocol ('spi_size', b'\x00'), ('notify_type', b'\x00\x07'), # INVALID_SYNTAX (0x0007) ]) self.notify_bytes_with_data = b''.join(self.notify_dict_with_data.values()) + self.notify_data def test_parse(self): parsed_notify: Ikev2PayloadNotifyUnparsed = Ikev2PayloadNotifyUnparsed.parse_exact_size( self.notify_bytes_with_data) self.assertEqual(parsed_notify.data, self.notify_data) # pylint: disable=no-member def test_compose(self): notify_payload = Ikev2PayloadNotifyUnparsed( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=self.notify_type, spi=b'', data=self.notify_data ) notify_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = notify_payload.compose() self.assertEqual(composed_bytes, self.notify_bytes_with_data) def test_round_trip(self): original_payload = Ikev2PayloadNotifyUnparsed( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=self.notify_type, spi=b'', data=self.notify_data ) original_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = original_payload.compose() parsed_payload: Ikev2PayloadNotifyUnparsed = Ikev2PayloadNotifyUnparsed.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.data, original_payload.data) # pylint: disable=no-member class TestIkev2NotifyPayloadInvalidKe(unittest.TestCase): def setUp(self): self.dh_group = Ikev2DiffieHellmanGroup.MODP_GROUP_2048_BIT self.invalid_ke_dict = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE ('flags', b'\x00'), ('payload_length', b'\x00\x0a'), # 10 bytes (8 header + 2 dh_group) ('protocol_id', b'\x01'), # IKE protocol ('spi_size', b'\x00'), ('notify_type', b'\x00\x11'), # INVALID_KE_PAYLOAD (0x0011) ('dh_group', b'\x00\x0e'), # MODP_GROUP_2048_BIT (0x000e) ]) self.invalid_ke_bytes = b''.join(self.invalid_ke_dict.values()) self.invalid_ke_payload = Ikev2NotifyPayloadInvalidKe( flags=set(), protocol_id=Ikev2ProtocolId.IKE, type=Ikev2NotifyType.INVALID_KE_PAYLOAD, spi=b'', dh_group=self.dh_group ) self.invalid_ke_payload.next_payload = Ikev2PayloadType.NONE def test_get_message_type(self): self.assertEqual(Ikev2NotifyPayloadInvalidKe._get_message_type(), # pylint: disable=protected-access Ikev2NotifyType.INVALID_KE_PAYLOAD) def test_parse(self): parsed_notify: Ikev2NotifyPayloadInvalidKe = Ikev2NotifyPayloadInvalidKe.parse_exact_size(self.invalid_ke_bytes) self.assertEqual(parsed_notify.type, Ikev2NotifyType.INVALID_KE_PAYLOAD) self.assertEqual(parsed_notify.protocol_id, Ikev2ProtocolId.IKE) self.assertEqual(parsed_notify.dh_group, self.dh_group) # pylint: disable=no-member def test_compose(self): composed_bytes = self.invalid_ke_payload.compose() self.assertEqual(composed_bytes, self.invalid_ke_bytes) def test_round_trip(self): composed_bytes = self.invalid_ke_payload.compose() parsed_payload: Ikev2NotifyPayloadInvalidKe = Ikev2NotifyPayloadInvalidKe.parse_exact_size(composed_bytes) # Verify all attributes are preserved self.assertEqual(parsed_payload.dh_group, self.invalid_ke_payload.dh_group) # pylint: disable=no-member self.assertEqual(parsed_payload.flags, self.invalid_ke_payload.flags) self.assertEqual(parsed_payload.next_payload, self.invalid_ke_payload.next_payload) self.assertEqual(parsed_payload.spi, self.invalid_ke_payload.spi) class TestIkev2PayloadCertificateRequest(unittest.TestCase): _CERT_ENCODING = Ikev2CertificateType.X509_CERTIFICATE_SIGNATURE _CERTIFICATION_AUTHORITY = ( b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' b'\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f' ) _CERTREQ_BYTES = bytes.fromhex( '00' # next_payload = NONE '00' # flags '0025' # payload_length = 37 (4 header + 1 cert_encoding + 32 certification authority) '04' # cert_encoding = X.509 certificate signature ) + _CERTIFICATION_AUTHORITY def test_get_payload_type(self): self.assertEqual(Ikev2PayloadCertificateRequest.get_payload_type(), Ikev2PayloadType.CERTREQ) def test_parse(self): parsed_certreq: Ikev2PayloadCertificateRequest = Ikev2PayloadCertificateRequest.parse_exact_size( self._CERTREQ_BYTES) self.assertEqual(parsed_certreq.next_payload, Ikev2PayloadType.NONE) self.assertEqual(parsed_certreq.flags, set()) self.assertEqual(parsed_certreq.cert_encoding, self._CERT_ENCODING) self.assertEqual(parsed_certreq.certification_authority, self._CERTIFICATION_AUTHORITY) def test_compose(self): certreq_payload = Ikev2PayloadCertificateRequest( flags=set(), cert_encoding=self._CERT_ENCODING, certification_authority=self._CERTIFICATION_AUTHORITY, ) certreq_payload.next_payload = Ikev2PayloadType.NONE self.assertEqual(certreq_payload.compose(), self._CERTREQ_BYTES) def test_round_trip(self): certreq_payload = Ikev2PayloadCertificateRequest( flags=set(), cert_encoding=self._CERT_ENCODING, certification_authority=self._CERTIFICATION_AUTHORITY, ) certreq_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = certreq_payload.compose() parsed_payload: Ikev2PayloadCertificateRequest = Ikev2PayloadCertificateRequest.parse_exact_size( composed_bytes) self.assertEqual(parsed_payload.cert_encoding, certreq_payload.cert_encoding) self.assertEqual(parsed_payload.certification_authority, certreq_payload.certification_authority) self.assertEqual(parsed_payload.flags, certreq_payload.flags) self.assertEqual(parsed_payload.next_payload, certreq_payload.next_payload) def test_error_parse_not_enough_data(self): incomplete_data = self._CERTREQ_BYTES[:-5] with self.assertRaises(NotEnoughData) as context_manager: Ikev2PayloadCertificateRequest.parse_exact_size(incomplete_data) self.assertEqual(context_manager.exception.bytes_needed, 5) def test_get_authority_hashes_returns_empty_list_for_empty(self): payload = Ikev2PayloadCertificateRequest( flags=set(), cert_encoding=self._CERT_ENCODING, certification_authority=b'', ) self.assertEqual(payload.get_authority_hashes(), []) def test_get_authority_hashes_splits_single_20_byte_hash(self): hash_a = b'\x00' * 20 payload = Ikev2PayloadCertificateRequest( flags=set(), cert_encoding=self._CERT_ENCODING, certification_authority=hash_a, ) self.assertEqual(payload.get_authority_hashes(), [hash_a]) def test_get_authority_hashes_splits_multiple_concatenated_hashes(self): hash_a = b'\x00' * 20 hash_b = b'\xff' * 20 hash_c = b'\x42' * 20 payload = Ikev2PayloadCertificateRequest( flags=set(), cert_encoding=self._CERT_ENCODING, certification_authority=hash_a + hash_b + hash_c, ) self.assertEqual(payload.get_authority_hashes(), [hash_a, hash_b, hash_c]) def test_get_authority_hashes_drops_trailing_partial_chunk(self): hash_a = b'\x00' * 20 payload = Ikev2PayloadCertificateRequest( flags=set(), cert_encoding=self._CERT_ENCODING, certification_authority=hash_a + b'\xaa\xbb\xcc', # 3 stray octets ) self.assertEqual(payload.get_authority_hashes(), [hash_a]) def test_get_authority_hashes_accepts_bytearray_certification_authority(self): hash_a = b'\x00' * 20 payload = Ikev2PayloadCertificateRequest( flags=set(), cert_encoding=self._CERT_ENCODING, certification_authority=bytearray(hash_a), ) self.assertEqual(payload.get_authority_hashes(), [hash_a]) class TestIkev2PayloadCertificate(unittest.TestCase): _CERT_ENCODING = Ikev2CertificateType.X509_CERTIFICATE_SIGNATURE _CERTIFICATE_DATA = ( b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' b'\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f' ) _CERT_BYTES = bytes.fromhex( '00' # next_payload = NONE '00' # flags '0025' # payload_length = 37 (4 header + 1 cert_encoding + 32 data) '04' # cert_encoding = X.509 certificate signature ) + _CERTIFICATE_DATA def test_get_payload_type(self): self.assertEqual(Ikev2PayloadCertificate.get_payload_type(), Ikev2PayloadType.CERT) def test_parse(self): parsed_cert: Ikev2PayloadCertificate = Ikev2PayloadCertificate.parse_exact_size(self._CERT_BYTES) self.assertEqual(parsed_cert.next_payload, Ikev2PayloadType.NONE) self.assertEqual(parsed_cert.flags, set()) self.assertEqual(parsed_cert.cert_encoding, self._CERT_ENCODING) self.assertEqual(parsed_cert.certificate_data, self._CERTIFICATE_DATA) def test_compose(self): cert_payload = Ikev2PayloadCertificate( flags=set(), cert_encoding=self._CERT_ENCODING, certificate_data=self._CERTIFICATE_DATA, ) cert_payload.next_payload = Ikev2PayloadType.NONE self.assertEqual(cert_payload.compose(), self._CERT_BYTES) def test_round_trip(self): cert_payload = Ikev2PayloadCertificate( flags=set(), cert_encoding=self._CERT_ENCODING, certificate_data=self._CERTIFICATE_DATA, ) cert_payload.next_payload = Ikev2PayloadType.NONE composed_bytes = cert_payload.compose() parsed_payload: Ikev2PayloadCertificate = Ikev2PayloadCertificate.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.cert_encoding, cert_payload.cert_encoding) self.assertEqual(parsed_payload.certificate_data, cert_payload.certificate_data) self.assertEqual(parsed_payload.flags, cert_payload.flags) self.assertEqual(parsed_payload.next_payload, cert_payload.next_payload) def test_error_parse_not_enough_data(self): incomplete_data = self._CERT_BYTES[:-5] with self.assertRaises(NotEnoughData) as context_manager: Ikev2PayloadCertificate.parse_exact_size(incomplete_data) self.assertEqual(context_manager.exception.bytes_needed, 5) class TestIkev2PayloadIdentificationInitiator(unittest.TestCase): _IDENTIFIER = 'scanner@cryptolyzer' _ID_BYTES = bytes.fromhex( '00' # next_payload = NONE '00' # flags '001b' # payload_length = 27 (4 header + 1 id_type + 3 reserved + 19 id_data) '03' # id_type = RFC822_ADDR '000000' # reserved ) + _IDENTIFIER.encode('ascii') def test_get_payload_type(self): self.assertEqual( Ikev2PayloadIdentificationInitiatorRfc822Addr.get_payload_type(), Ikev2PayloadType.IDI, ) def test_parse(self): parsed = Ikev2PayloadIdentificationInitiatorRfc822Addr.parse_exact_size(self._ID_BYTES) self.assertEqual(parsed.next_payload, Ikev2PayloadType.NONE) self.assertEqual(parsed.flags, set()) self.assertEqual(parsed.id_type, Ikev2IdType.RFC822_ADDR) self.assertEqual(parsed.identifier, self._IDENTIFIER) # pylint: disable=no-member def test_compose(self): id_payload = Ikev2PayloadIdentificationInitiatorRfc822Addr( flags=set(), identifier=self._IDENTIFIER, ) id_payload.next_payload = Ikev2PayloadType.NONE self.assertEqual(id_payload.compose(), self._ID_BYTES) def test_round_trip(self): id_payload = Ikev2PayloadIdentificationInitiatorRfc822Addr( flags=set(), identifier=self._IDENTIFIER, ) id_payload.next_payload = Ikev2PayloadType.NONE composed = id_payload.compose() parsed = Ikev2PayloadIdentificationInitiatorRfc822Addr.parse_exact_size(composed) self.assertEqual(parsed.id_type, id_payload.id_type) identifier_parsed = parsed.identifier # pylint: disable=no-member identifier_expected = id_payload.identifier # pylint: disable=no-member self.assertEqual(identifier_parsed, identifier_expected) self.assertEqual(parsed.flags, id_payload.flags) self.assertEqual(parsed.next_payload, id_payload.next_payload) class TestIkev2PayloadIdentificationResponder(unittest.TestCase): _IDENTIFIER = ipaddress.IPv4Address('192.0.2.1') def test_get_payload_type(self): self.assertEqual( Ikev2PayloadIdentificationResponderIpv4Addr.get_payload_type(), Ikev2PayloadType.IDR, ) def test_round_trip(self): id_payload = Ikev2PayloadIdentificationResponderIpv4Addr( flags=set(), identifier=self._IDENTIFIER, ) id_payload.next_payload = Ikev2PayloadType.NONE composed = id_payload.compose() parsed = Ikev2PayloadIdentificationResponderIpv4Addr.parse_exact_size(composed) self.assertEqual(parsed.id_type, Ikev2IdType.IPV4_ADDR) self.assertEqual(parsed.identifier, self._IDENTIFIER) # pylint: disable=no-member def test_round_trip_der_asn1_dn(self): dn_bytes = ( b'\x30\x1a\x31\x18\x30\x16\x06\x03\x55\x04\x03\x0c\x0f' b'gateway.test.org' ) payload = Ikev2PayloadIdentificationResponderDerAsn1Dn( flags=set(), identifier=dn_bytes, ) payload.next_payload = Ikev2PayloadType.NONE composed = payload.compose() parsed = Ikev2PayloadIdentificationResponderDerAsn1Dn.parse_exact_size(composed) self.assertEqual(parsed.id_type, Ikev2IdType.DER_ASN1_DN) self.assertEqual(parsed.identifier, dn_bytes) # pylint: disable=no-member def test_ipv4_decode_rejects_wrong_length(self): with self.assertRaises(InvalidType): Ikev2PayloadIdentificationResponderIpv4Addr._decode_identifier( # pylint: disable=protected-access b'\xc0\x00\x02', ) def test_ipv6_decode_rejects_wrong_length(self): with self.assertRaises(InvalidType): Ikev2PayloadIdentificationResponderIpv6Addr._decode_identifier( # pylint: disable=protected-access b'\x00' * 15, ) class TestIkev2PayloadIdentificationVariants(unittest.TestCase): def test_initiator_fqdn_round_trip(self): payload = Ikev2PayloadIdentificationInitiatorFqdn( flags=set(), identifier='gateway.example.com', ) payload.next_payload = Ikev2PayloadType.NONE parsed = Ikev2PayloadIdentificationInitiatorFqdn.parse_exact_size(payload.compose()) self.assertEqual(parsed.id_type, Ikev2IdType.FQDN) self.assertEqual(parsed.identifier, 'gateway.example.com') # pylint: disable=no-member def test_initiator_ipv6_round_trip(self): addr = ipaddress.IPv6Address('2001:db8::1') payload = Ikev2PayloadIdentificationInitiatorIpv6Addr( flags=set(), identifier=addr, ) payload.next_payload = Ikev2PayloadType.NONE parsed = Ikev2PayloadIdentificationInitiatorIpv6Addr.parse_exact_size(payload.compose()) self.assertEqual(parsed.id_type, Ikev2IdType.IPV6_ADDR) self.assertEqual(parsed.identifier, addr) # pylint: disable=no-member def test_initiator_key_id_round_trip(self): payload = Ikev2PayloadIdentificationInitiatorKeyId( flags=set(), identifier=b'\x01\x02\x03', ) payload.next_payload = Ikev2PayloadType.NONE parsed = Ikev2PayloadIdentificationInitiatorKeyId.parse_exact_size(payload.compose()) self.assertEqual(parsed.id_type, Ikev2IdType.KEY_ID) self.assertEqual(parsed.identifier, b'\x01\x02\x03') # pylint: disable=no-member def test_initiator_der_asn1_gn_round_trip(self): payload = Ikev2PayloadIdentificationInitiatorDerAsn1Gn( flags=set(), identifier=b'\xab\xcd\xef', ) payload.next_payload = Ikev2PayloadType.NONE parsed = Ikev2PayloadIdentificationInitiatorDerAsn1Gn.parse_exact_size(payload.compose()) self.assertEqual(parsed.id_type, Ikev2IdType.DER_ASN1_GN) self.assertEqual(parsed.identifier, b'\xab\xcd\xef') # pylint: disable=no-member def test_initiator_fc_name_round_trip(self): payload = Ikev2PayloadIdentificationInitiatorFcName( flags=set(), identifier=b'\x10\x20\x30\x40\x50\x60\x70\x80', ) payload.next_payload = Ikev2PayloadType.NONE parsed = Ikev2PayloadIdentificationInitiatorFcName.parse_exact_size(payload.compose()) self.assertEqual(parsed.id_type, Ikev2IdType.FC_NAME) self.assertEqual(parsed.identifier, b'\x10\x20\x30\x40\x50\x60\x70\x80') # pylint: disable=no-member def test_initiator_null_round_trip(self): payload = Ikev2PayloadIdentificationInitiatorNull( flags=set(), identifier=b'', ) payload.next_payload = Ikev2PayloadType.NONE parsed = Ikev2PayloadIdentificationInitiatorNull.parse_exact_size(payload.compose()) self.assertEqual(parsed.id_type, Ikev2IdType.NULL) self.assertEqual(parsed.identifier, b'') # pylint: disable=no-member def test_id_payload_rejects_negative_id_data_length(self): # payload_length=6 → id_data_length = 6 - 4 - 4 = -2 → NotEnoughData bogus = b'\x00\x00\x00\x06\x02\x00\x00\x00\x00\x00' with self.assertRaises(NotEnoughData): Ikev2PayloadIdentificationInitiatorFqdn.parse_exact_size(bogus) def test_id_payload_rejects_id_type_mismatch(self): # Wire id_type = 0x01 (IPV4_ADDR); caller uses FQDN subclass. wire = b'\x00\x00\x00\x0c\x01\x00\x00\x00\xc0\x00\x02\x01' with self.assertRaises(InvalidType): Ikev2PayloadIdentificationInitiatorFqdn.parse_exact_size(wire) def test_initiator_variant_dispatch_picks_subclass(self): source = Ikev2PayloadIdentificationInitiatorFqdn( flags=set(), identifier='initiator.example', ) source.next_payload = Ikev2PayloadType.NONE parsed = Ikev2PayloadIdentificationInitiator.parse_exact_size(source.compose()) self.assertIsInstance(parsed, Ikev2PayloadIdentificationInitiatorFqdn) def test_responder_variant_dispatch_picks_subclass(self): source = Ikev2PayloadIdentificationResponderFqdn( flags=set(), identifier='responder.example', ) source.next_payload = Ikev2PayloadType.NONE parsed = Ikev2PayloadIdentificationResponder.parse_exact_size(source.compose()) self.assertIsInstance(parsed, Ikev2PayloadIdentificationResponderFqdn) def test_rfc822_mixin_returns_none_for_ikev1(self): self.assertIsNone(Ikev2PayloadIdentificationInitiatorRfc822Addr.get_id_type_ikev1()) def test_fc_name_mixin_returns_none_for_ikev1(self): self.assertIsNone(Ikev2PayloadIdentificationInitiatorFcName.get_id_type_ikev1()) def test_null_mixin_returns_none_for_ikev1(self): self.assertIsNone(Ikev2PayloadIdentificationInitiatorNull.get_id_type_ikev1()) class TestIkev2PayloadAuthentication(unittest.TestCase): _AUTH_METHOD = Ikev2AuthenticationMethod.RSA_DIGITAL_SIGNATURE _AUTH_DATA = b'\xab' * 256 # 2048-bit RSA signature _AUTH_BYTES = bytes.fromhex( '00' # next_payload = NONE '00' # flags '0108' # payload_length = 264 (4 header + 1 auth_method + 3 reserved + 256 auth_data) '01' # auth_method = RSA_DIGITAL_SIGNATURE '000000' # reserved ) + _AUTH_DATA def test_get_payload_type(self): self.assertEqual(Ikev2PayloadAuthentication.get_payload_type(), Ikev2PayloadType.AUTH) def test_parse(self): parsed: Ikev2PayloadAuthentication = Ikev2PayloadAuthentication.parse_exact_size(self._AUTH_BYTES) self.assertEqual(parsed.auth_method, self._AUTH_METHOD) self.assertEqual(parsed.auth_data, self._AUTH_DATA) def test_round_trip_rsa(self): auth_payload = Ikev2PayloadAuthentication( flags=set(), auth_method=self._AUTH_METHOD, auth_data=self._AUTH_DATA, ) auth_payload.next_payload = Ikev2PayloadType.NONE composed = auth_payload.compose() parsed: Ikev2PayloadAuthentication = Ikev2PayloadAuthentication.parse_exact_size(composed) self.assertEqual(parsed.auth_method, self._AUTH_METHOD) self.assertEqual(parsed.auth_data, self._AUTH_DATA) def test_round_trip_ecdsa_256(self): payload = Ikev2PayloadAuthentication( flags=set(), auth_method=Ikev2AuthenticationMethod.ECDSA_SHA_256_P_256, auth_data=b'\xcd' * 64, # r||s for P-256 ) payload.next_payload = Ikev2PayloadType.NONE composed = payload.compose() parsed: Ikev2PayloadAuthentication = Ikev2PayloadAuthentication.parse_exact_size(composed) self.assertEqual(parsed.auth_method, Ikev2AuthenticationMethod.ECDSA_SHA_256_P_256) self.assertEqual(len(parsed.auth_data), 64) def test_error_parse_negative_auth_data_length(self): # payload_length=6 → auth_data_length = 6 - 4 - 4 = -2 → NotEnoughData bogus = b'\x00\x00\x00\x06\x01\x00\x00\x00\x00\x00' with self.assertRaises(NotEnoughData): Ikev2PayloadAuthentication.parse_exact_size(bogus) def test_round_trip_null(self): payload = Ikev2PayloadAuthentication( flags=set(), auth_method=Ikev2AuthenticationMethod.NULL_AUTHENTICATION, auth_data=b'\xef' * 32, # HMAC(SK_pi, ...) output sized to the prf ) payload.next_payload = Ikev2PayloadType.NONE composed = payload.compose() parsed: Ikev2PayloadAuthentication = Ikev2PayloadAuthentication.parse_exact_size(composed) self.assertEqual(parsed.auth_method, Ikev2AuthenticationMethod.NULL_AUTHENTICATION) class TestIkev2AuthDigitalSignatureEnvelope(unittest.TestCase): # AlgorithmIdentifier for ecdsa-with-SHA256 (OID 1.2.840.10045.4.3.2), # encoded as a DER SEQUENCE so asn1crypto can decode the OID. _SIGNATURE_ALGORITHM = Signature.ECDSA_WITH_SHA2_256 _ALGORITHM_IDENTIFIER = bytes.fromhex('300a06082a8648ce3d040302') _SIGNATURE = b'\xab' * 64 # P-256 r||s _ENVELOPE_BYTES = bytes([len(_ALGORITHM_IDENTIFIER)]) + _ALGORITHM_IDENTIFIER + _SIGNATURE def test_round_trip(self): envelope = Ikev2AuthDigitalSignatureEnvelope( signature_algorithm=self._SIGNATURE_ALGORITHM, signature=self._SIGNATURE, ) composed = envelope.compose() self.assertEqual(composed, self._ENVELOPE_BYTES) parsed = Ikev2AuthDigitalSignatureEnvelope.from_bytes(composed) self.assertEqual(parsed.signature_algorithm, self._SIGNATURE_ALGORITHM) self.assertEqual(parsed.signature, self._SIGNATURE) def test_from_auth_payload(self): payload = Ikev2PayloadAuthentication( flags=set(), auth_method=Ikev2AuthenticationMethod.DIGITAL_SIGNATURE, auth_data=self._ENVELOPE_BYTES, ) payload.next_payload = Ikev2PayloadType.NONE envelope: Ikev2AuthDigitalSignatureEnvelope = payload.parse_digital_signature_envelope() self.assertEqual(envelope.signature_algorithm, self._SIGNATURE_ALGORITHM) self.assertEqual(envelope.signature, self._SIGNATURE) def test_wrong_auth_method_rejected(self): payload = Ikev2PayloadAuthentication( flags=set(), auth_method=Ikev2AuthenticationMethod.RSA_DIGITAL_SIGNATURE, auth_data=self._ENVELOPE_BYTES, ) payload.next_payload = Ikev2PayloadType.NONE with self.assertRaises(InvalidType): payload.parse_digital_signature_envelope() def test_parse_unknown_oid_raises_invalid_value(self): # AlgorithmIdentifier for the made-up OID 1.2.3.4.5 (not in # cryptodatahub's Signature registry). # DER: SEQUENCE(0x30) len=6 OID(0x06) len=4 first-two(0x2A) 03 04 05 algorithm_identifier = bytes.fromhex('300606042a030405') envelope_bytes = ( bytes([len(algorithm_identifier)]) + algorithm_identifier + b'\xcc' * 32 ) with self.assertRaises(InvalidValue): Ikev2AuthDigitalSignatureEnvelope.from_bytes(envelope_bytes) class TestIkev2PayloadEncryptedAndAuthenticated(unittest.TestCase): # 16-octet IV + 16 octets of "encrypted" body + 1 octet pad length + 16 octet ICV _ENCRYPTED_DATA = b'\x11' * 16 + b'\x22' * 16 + b'\x00' + b'\x33' * 16 _SK_BYTES = bytes.fromhex( '23' # next_payload = IDI (35 = inner first payload after decrypt) '00' # flags '0035' # payload_length = 53 (4 header + 49 encrypted_data) ) + _ENCRYPTED_DATA def test_get_payload_type(self): self.assertEqual(Ikev2PayloadEncryptedAndAuthenticated.get_payload_type(), Ikev2PayloadType.SK) def test_parse(self): parsed: Ikev2PayloadEncryptedAndAuthenticated parsed = Ikev2PayloadEncryptedAndAuthenticated.parse_exact_size(self._SK_BYTES) self.assertEqual(parsed.next_payload, Ikev2PayloadType.IDI) self.assertEqual(parsed.flags, set()) self.assertEqual(parsed.encrypted_data, self._ENCRYPTED_DATA) def test_round_trip(self): sk_payload = Ikev2PayloadEncryptedAndAuthenticated( flags=set(), encrypted_data=self._ENCRYPTED_DATA, ) sk_payload.next_payload = Ikev2PayloadType.IDI composed = sk_payload.compose() self.assertEqual(composed, self._SK_BYTES) parsed: Ikev2PayloadEncryptedAndAuthenticated parsed = Ikev2PayloadEncryptedAndAuthenticated.parse_exact_size(composed) self.assertEqual(parsed.encrypted_data, self._ENCRYPTED_DATA) self.assertEqual(parsed.next_payload, Ikev2PayloadType.IDI) class TestIkev2PayloadEap(unittest.TestCase): _EAP_DATA = b'\x01\x01\x00\x05\x01' # EAP-Request/Identity-ish _EAP_BYTES = bytes.fromhex( '00' # next_payload = NONE '00' # flags '0009' # payload_length = 9 (4 header + 5 eap_data) ) + _EAP_DATA def test_get_payload_type(self): self.assertEqual(Ikev2PayloadEap.get_payload_type(), Ikev2PayloadType.EAP) def test_round_trip(self): eap_payload = Ikev2PayloadEap( flags=set(), eap_data=self._EAP_DATA, ) eap_payload.next_payload = Ikev2PayloadType.NONE composed = eap_payload.compose() self.assertEqual(composed, self._EAP_BYTES) parsed: Ikev2PayloadEap = Ikev2PayloadEap.parse_exact_size(composed) self.assertEqual(parsed.eap_data, self._EAP_DATA) class TestIkev2PayloadVendorId(unittest.TestCase): def setUp(self): self.vendor_id = ( b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' b'\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f' ) self.vendor_dict = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE ('flags', b'\x80'), # CRITICAL (0x80) ('payload_length', b'\x00\x24'), # 36 bytes (4 header + 32 vendor_id) ('vendor_id', self.vendor_id), ]) self.vendor_bytes = b''.join(self.vendor_dict.values()) self.vendor_payload = Ikev2PayloadVendorId( flags={Ikev2PayloadFlags.CRITICAL}, vendor_id=self.vendor_id ) self.vendor_payload.next_payload = Ikev2PayloadType.NONE def test_get_payload_type(self): self.assertEqual(Ikev2PayloadVendorId.get_payload_type(), Ikev2PayloadType.VENDOR_ID) def test_parse(self): parsed_vendor: Ikev2PayloadVendorId = Ikev2PayloadVendorId.parse_exact_size(self.vendor_bytes) self.assertEqual(parsed_vendor.next_payload, Ikev2PayloadType.NONE) self.assertEqual(parsed_vendor.flags, {Ikev2PayloadFlags.CRITICAL}) self.assertEqual(parsed_vendor.vendor_id, self.vendor_id) def test_compose(self): composed_bytes = self.vendor_payload.compose() self.assertEqual(composed_bytes, self.vendor_bytes) def test_round_trip(self): composed_bytes = self.vendor_payload.compose() parsed_payload: Ikev2PayloadVendorId = Ikev2PayloadVendorId.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.vendor_id, self.vendor_payload.vendor_id) self.assertEqual(parsed_payload.flags, self.vendor_payload.flags) self.assertEqual(parsed_payload.next_payload, self.vendor_payload.next_payload) def test_error_parse_not_enough_data(self): incomplete_data = self.vendor_bytes[:-10] with self.assertRaises(NotEnoughData) as context_manager: Ikev2PayloadVendorId.parse_exact_size(incomplete_data) self.assertEqual(context_manager.exception.bytes_needed, 10) class TestTransformAttributeKeyLength(unittest.TestCase): def setUp(self): self.key_length_value = 128 self.key_length_attribute = TransformAttributeKeyLength(value=self.key_length_value) def test_parse(self): composed_bytes = self.key_length_attribute.compose() parsed_attribute = TransformAttributeKeyLength.parse_exact_size(composed_bytes) self.assertEqual(parsed_attribute.value, self.key_length_attribute.value) def test_compose(self): composed_bytes = self.key_length_attribute.compose() self.assertEqual(len(composed_bytes), 4) self.assertEqual(composed_bytes[0], 0x80) self.assertEqual(composed_bytes[1], Ikev2TransformAttributeType.KEY_LENGTH.value.code) self.assertEqual(composed_bytes[2:4], b'\x00\x80') def test_round_trip(self): composed_bytes = self.key_length_attribute.compose() parsed_attribute = TransformAttributeKeyLength.parse_exact_size(composed_bytes) self.assertEqual(parsed_attribute.value, self.key_length_attribute.value) def test_error_parse_not_enough_data(self): incomplete_data = b'\x80\x0e\x00' with self.assertRaises(NotEnoughData) as context_manager: TransformAttributeKeyLength.parse_exact_size(incomplete_data) self.assertEqual(context_manager.exception.bytes_needed, 1) if __name__ == '__main__': unittest.main() cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ike/test_ikev2_sa.py000066400000000000000000000612171524413560000267240ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import collections import unittest from cryptodatahub.common.exception import InvalidValue from cryptodatahub.ike.algorithm import ( Ikev2DiffieHellmanGroup, Ikev2EncryptionAlgorithm, Ikev2IntegrityAlgorithm, Ikev2ProtocolId, Ikev2PseudorandomFunction, Ikev2TransformType, ) from cryptoparser.common.exception import NotEnoughData from cryptoparser.ike.ikev2 import ( Ikev2Proposal, Ikev2ProposalNextPayload, Ikev2PayloadSecurityAssociation, Ikev2PayloadFlags, Ikev2PayloadType, Ikev2TransformPrf, Ikev2TransformDhGroup, Ikev2TransformEncryptionAlgorithm, Ikev2TransformIntegrity, TransformAttributeSignatureAlgorithm, TransformNextPayload, ) from .classes import TransformTest class TestTransformAttributeSignatureAlgorithm(unittest.TestCase): def setUp(self): self.signature_algorithm = ( b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' b'\x10\x11\x12\x13\x14\x15\x16\x17\x18\x19\x1a\x1b\x1c\x1d\x1e\x1f' ) self.signature_algorithm_payload = TransformAttributeSignatureAlgorithm( signature_algorithm=self.signature_algorithm ) self.signature_algorithm_dict = collections.OrderedDict([ ('format', b'\x00'), # TYPE_LENGTH_VALUE ('type', b'\x12'), # SIGNATURE_ALGORITHM = 18 ('length', b'\x00\x20'), # 32 bytes length ('signature_algorithm', self.signature_algorithm), ]) self.signature_algorithm_bytes = b''.join(self.signature_algorithm_dict.values()) def test_parse(self): parsed_signature_algorithm = TransformAttributeSignatureAlgorithm.parse_exact_size( self.signature_algorithm_bytes ) self.assertEqual(parsed_signature_algorithm.signature_algorithm, self.signature_algorithm) def test_compose(self): composed_bytes = self.signature_algorithm_payload.compose() parsed_signature_algorithm = TransformAttributeSignatureAlgorithm.parse_exact_size(composed_bytes) self.assertEqual( parsed_signature_algorithm.signature_algorithm, self.signature_algorithm_payload.signature_algorithm ) def test_round_trip(self): composed_bytes = self.signature_algorithm_payload.compose() parsed_payload: TransformAttributeSignatureAlgorithm = ( TransformAttributeSignatureAlgorithm.parse_exact_size(composed_bytes) ) self.assertEqual(parsed_payload.signature_algorithm, self.signature_algorithm_payload.signature_algorithm) def test_error_parse_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: TransformAttributeSignatureAlgorithm.parse_exact_size(b'\x00\x12') self.assertEqual(context_manager.exception.bytes_needed, 2) incomplete = self.signature_algorithm_bytes[:-1] with self.assertRaises(NotEnoughData) as context_manager: TransformAttributeSignatureAlgorithm.parse_exact_size(incomplete) self.assertEqual(context_manager.exception.bytes_needed, 1) def test_error_invalid_type(self): wrong_type_header = b'\x00\x11\x00\x20' # wrong type 0x11 instead of 0x12 wrong_type_bytes = wrong_type_header + self.signature_algorithm with self.assertRaises(InvalidValue): TransformAttributeSignatureAlgorithm.parse_exact_size(wrong_type_bytes) def test_error_invalid_signature_algorithm_type(self): with self.assertRaises(TypeError): TransformAttributeSignatureAlgorithm( signature_algorithm="not_bytes" ) class TestTransform(unittest.TestCase): def setUp(self): self.transform_id = Ikev2PseudorandomFunction.PRF_HMAC_SHA1 self.transform = TransformTest( transform_id=self.transform_id ) self.transform.next_payload = TransformNextPayload.LAST self.transform_dict = collections.OrderedDict([ ('next_payload', b'\x00'), # LAST ('reserved1', b'\x00'), ('transform_length', b'\x00\x08'), # 8 bytes header only ('transform_type', b'\x02'), # PRF ('reserved2', b'\x00'), ('transform_id', b'\x00\x02'), # PRF_HMAC_SHA1 ]) self.transform_bytes = b''.join(self.transform_dict.values()) def test_parse(self): parsed_transform = TransformTest.parse_exact_size(self.transform_bytes) self.assertEqual(parsed_transform.transform_id, self.transform_id) self.assertEqual(parsed_transform.next_payload, TransformNextPayload.LAST) def test_compose(self): composed_bytes = self.transform.compose() self.assertEqual(len(composed_bytes), 8) # Header only, no attributes parsed_transform = TransformTest.parse_exact_size(composed_bytes) self.assertEqual(parsed_transform.transform_id, self.transform.transform_id) def test_round_trip(self): composed_bytes = self.transform.compose() parsed_payload: TransformTest = TransformTest.parse_exact_size(composed_bytes) self.assertEqual(parsed_payload.transform_id, self.transform.transform_id) self.assertEqual(parsed_payload.next_payload, self.transform.next_payload) def test_get_transform_type(self): self.assertEqual(TransformTest.get_transform_type(), Ikev2TransformType.PRF) def test_get_transform_id_class(self): self.assertEqual( TransformTest._get_transform_id_class(), # pylint: disable=protected-access Ikev2PseudorandomFunction ) def test_error_parse_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: TransformTest.parse_exact_size(b'\x00\x01\x02') self.assertEqual(context_manager.exception.bytes_needed, 5) incomplete = self.transform_bytes[:-1] with self.assertRaises(NotEnoughData) as context_manager: TransformTest.parse_exact_size(incomplete) self.assertEqual(context_manager.exception.bytes_needed, 1) def test_error_invalid_transform_id(self): with self.assertRaises(InvalidValue): TransformTest( transform_id="invalid_id" ) class TestIkev2TransformPrf(unittest.TestCase): def test_get_transform_type(self): self.assertEqual(Ikev2TransformPrf.get_transform_type(), Ikev2TransformType.PRF) def test_get_transform_id_class(self): self.assertEqual( Ikev2TransformPrf._get_transform_id_class(), # pylint: disable=protected-access Ikev2PseudorandomFunction ) class TestIkev2TransformDhGroup(unittest.TestCase): def test_get_transform_type(self): self.assertEqual(Ikev2TransformDhGroup.get_transform_type(), Ikev2TransformType.DH) def test_get_transform_id_class(self): self.assertEqual( Ikev2TransformDhGroup._get_transform_id_class(), # pylint: disable=protected-access Ikev2DiffieHellmanGroup ) class TestIkev2TransformIntegrity(unittest.TestCase): def test_get_transform_type(self): self.assertEqual(Ikev2TransformIntegrity.get_transform_type(), Ikev2TransformType.INTEG) def test_get_transform_id_class(self): self.assertEqual( Ikev2TransformIntegrity._get_transform_id_class(), # pylint: disable=protected-access Ikev2IntegrityAlgorithm ) class TestIkev2TransformEncryptionAlgorithm(unittest.TestCase): def setUp(self): self.encryption_algorithm = Ikev2EncryptionAlgorithm.ENCR_AES_CBC self.key_length = 128 self.transform = Ikev2TransformEncryptionAlgorithm( transform_id=self.encryption_algorithm, key_length=self.key_length ) self.transform.next_payload = TransformNextPayload.LAST def test_get_transform_type(self): self.assertEqual(Ikev2TransformEncryptionAlgorithm.get_transform_type(), Ikev2TransformType.ENCR) def test_get_transform_id_class(self): self.assertEqual( Ikev2TransformEncryptionAlgorithm._get_transform_id_class(), # pylint: disable=protected-access Ikev2EncryptionAlgorithm ) def test_key_length_value_support(self): different_key_lengths = [128, 192, 256] for key_length in different_key_lengths: transform = Ikev2TransformEncryptionAlgorithm( transform_id=self.encryption_algorithm, key_length=key_length ) self.assertEqual(transform.key_length, key_length) def test_parse(self): composed_bytes = self.transform.compose() parsed_transform = Ikev2TransformEncryptionAlgorithm.parse_exact_size(composed_bytes) self.assertEqual(parsed_transform.transform_id, self.transform.transform_id) self.assertEqual(parsed_transform.key_length, self.transform.key_length) # pylint: disable=no-member def test_compose(self): composed_bytes = self.transform.compose() self.assertGreater(len(composed_bytes), 8) def test_round_trip(self): composed_bytes = self.transform.compose() parsed_transform = Ikev2TransformEncryptionAlgorithm.parse_exact_size(composed_bytes) self.assertEqual(parsed_transform.transform_id, self.transform.transform_id) self.assertEqual(parsed_transform.key_length, self.transform.key_length) # pylint: disable=no-member def test_error_parse_not_enough_data(self): incomplete_data = b'\x00\x00\x00\x08\x01\x00' with self.assertRaises(NotEnoughData) as context_manager: Ikev2TransformEncryptionAlgorithm.parse_exact_size(incomplete_data) self.assertGreater(context_manager.exception.bytes_needed, 0) def test_round_trip_without_key_length(self): transform_without_key = Ikev2TransformEncryptionAlgorithm( transform_id=Ikev2EncryptionAlgorithm.ENCR_DES, key_length=None, ) transform_without_key.next_payload = TransformNextPayload.LAST composed_bytes = transform_without_key.compose() parsed_transform = Ikev2TransformEncryptionAlgorithm.parse_exact_size(composed_bytes) self.assertEqual(parsed_transform.transform_id, Ikev2EncryptionAlgorithm.ENCR_DES) self.assertIsNone(parsed_transform.key_length) # pylint: disable=no-member class TestIkev2Proposal(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.transform_prf = Ikev2TransformPrf( transform_id=Ikev2PseudorandomFunction.PRF_HMAC_SHA1 ) self.transform_prf.next_payload = TransformNextPayload.LAST self.protocol_id = Ikev2ProtocolId.IKE self.spi = b'\x00\x01\x02\x03' self.proposal_minimal = Ikev2Proposal( protocol_id=self.protocol_id, transforms=[self.transform_prf], spi=b'' ) self.proposal_minimal.last = Ikev2ProposalNextPayload.LAST self.proposal_minimal.proposal_number = 1 self.proposal_with_spi = Ikev2Proposal( protocol_id=self.protocol_id, transforms=[self.transform_prf], spi=self.spi ) self.proposal_with_spi.last = Ikev2ProposalNextPayload.MORE self.proposal_with_spi.proposal_number = 2 self.proposal_dict_minimal = collections.OrderedDict([ ('last', b'\x00'), # LAST ('reserved', b'\x00'), ('proposal_length', b'\x00\x10'), # 16 bytes total (8 header + 8 transform) ('proposal_number', b'\x01'), # 1 ('protocol_id', b'\x01'), # IKE ('spi_size', b'\x00'), # 0 bytes ('transform_count', b'\x01'), # 1 transform # No SPI data # Transform data: next_payload(1) + reserved(1) + length(2) + type(1) + reserved(1) + id(2) ('transform_data', b'\x00\x00\x00\x08\x02\x00\x00\x02'), # PRF_HMAC_SHA1 ]) self.proposal_bytes_minimal = b''.join(self.proposal_dict_minimal.values()) self.proposal_minimal = Ikev2Proposal( protocol_id=self.protocol_id, transforms=[self.transform_prf], spi=b'' ) self.proposal_minimal.last = Ikev2ProposalNextPayload.LAST self.proposal_minimal.proposal_number = 1 self.proposal_dict_with_spi = collections.OrderedDict([ ('last', b'\x02'), # MORE enum value ('reserved', b'\x00'), ('proposal_length', b'\x00\x10'), # 16 bytes total (8 header + 8 transform) ('proposal_number', b'\x02'), # 2 ('protocol_id', b'\x01'), # IKE ('spi_size', b'\x04'), # 4 bytes ('transform_count', b'\x01'), # 1 transform ('spi', self.spi), # SPI data ('transform_data', b'\x00\x00\x00\x08\x02\x00\x00\x02'), # PRF_HMAC_SHA1 ]) self.proposal_bytes_with_spi = b''.join(self.proposal_dict_with_spi.values()) self.proposal_with_spi = Ikev2Proposal( protocol_id=self.protocol_id, transforms=[self.transform_prf], spi=self.spi ) self.proposal_with_spi.last = Ikev2ProposalNextPayload.MORE self.proposal_with_spi.proposal_number = 2 def test_parse(self): parsed_proposal = Ikev2Proposal.parse_exact_size(self.proposal_bytes_minimal) self.assertEqual(parsed_proposal.protocol_id, self.protocol_id) self.assertEqual(parsed_proposal.spi, b'') self.assertEqual(len(parsed_proposal.transforms), 1) self.assertEqual(parsed_proposal.last, set()) self.assertEqual(parsed_proposal.proposal_number, 1) parsed_proposal_spi = Ikev2Proposal.parse_exact_size(self.proposal_bytes_with_spi) self.assertEqual(parsed_proposal_spi.protocol_id, self.protocol_id) self.assertEqual(parsed_proposal_spi.spi, self.spi) self.assertEqual(len(parsed_proposal_spi.transforms), 1) self.assertEqual(parsed_proposal_spi.last, {Ikev2ProposalNextPayload.MORE}) self.assertEqual(parsed_proposal_spi.proposal_number, 2) def test_compose(self): composed_bytes = self.proposal_minimal.compose() self.assertEqual(composed_bytes, self.proposal_bytes_minimal) composed_bytes = self.proposal_with_spi.compose() self.assertEqual(composed_bytes, self.proposal_bytes_with_spi) def test_round_trip(self): composed_bytes = self.proposal_minimal.compose() parsed_proposal: Ikev2Proposal = Ikev2Proposal.parse_exact_size(composed_bytes) self.assertEqual(parsed_proposal.protocol_id, self.proposal_minimal.protocol_id) self.assertEqual(parsed_proposal.spi, self.proposal_minimal.spi) self.assertEqual(len(parsed_proposal.transforms), len(self.proposal_minimal.transforms)) self.assertEqual(parsed_proposal.last, set()) self.assertEqual(parsed_proposal.proposal_number, self.proposal_minimal.proposal_number) def test_error_parse_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: Ikev2Proposal.parse_exact_size(b'\x00\x01\x02') self.assertEqual(context_manager.exception.bytes_needed, Ikev2Proposal.HEADER_SIZE - 3) incomplete = self.proposal_bytes_minimal[:-1] with self.assertRaises(NotEnoughData) as context_manager: Ikev2Proposal.parse_exact_size(incomplete) self.assertEqual(context_manager.exception.bytes_needed, 1) def test_error_invalid_protocol_id(self): with self.assertRaises(TypeError): Ikev2Proposal( protocol_id="invalid_protocol", transforms=[self.transform_prf] ) def test_error_invalid_transforms(self): with self.assertRaises(TypeError): Ikev2Proposal( protocol_id=self.protocol_id, transforms=["invalid_transform"] ) def test_error_invalid_spi_type(self): with self.assertRaises(TypeError): Ikev2Proposal( protocol_id=self.protocol_id, transforms=[self.transform_prf], spi="invalid_spi" ) class TestIkev2PayloadSecurityAssociation(unittest.TestCase): # pylint: disable=too-many-instance-attributes def setUp(self): self.transform_prf = Ikev2TransformPrf( transform_id=Ikev2PseudorandomFunction.PRF_HMAC_SHA1 ) self.transform_prf.next_payload = TransformNextPayload.LAST self.proposal_single = Ikev2Proposal( protocol_id=Ikev2ProtocolId.IKE, transforms=[self.transform_prf], spi=b'' ) self.proposal_with_spi = Ikev2Proposal( protocol_id=Ikev2ProtocolId.IKE, transforms=[self.transform_prf], spi=b'\x00\x01\x02\x03' ) self.sa_single_proposal = Ikev2PayloadSecurityAssociation( flags=set(), proposals=[self.proposal_single] ) self.sa_single_proposal.next_payload = Ikev2PayloadType.NONE self.sa_multiple_proposals = Ikev2PayloadSecurityAssociation( flags={Ikev2PayloadFlags.CRITICAL}, proposals=[self.proposal_single, self.proposal_with_spi] ) self.sa_multiple_proposals.next_payload = Ikev2PayloadType.KE # Single proposal: header(4) + proposal(16) = 20 bytes total self.sa_dict_single = collections.OrderedDict([ ('next_payload', b'\x00'), # NONE ('flags', b'\x00'), # No flags ('payload_length', b'\x00\x14'), # 20 bytes total ('proposals_data', b'\x00\x00\x00\x10\x01\x01\x00\x01\x00\x00\x00\x08\x02\x00\x00\x02'), ]) self.sa_bytes_single = b''.join(self.sa_dict_single.values()) # Multiple proposals: header(4) + proposal1(16) + proposal2(20) = 40 bytes total self.sa_dict_multiple = collections.OrderedDict([ ('next_payload', b'\x22'), # KE ('flags', b'\x80'), # CRITICAL ('payload_length', b'\x00\x28'), # 40 bytes total ('proposals_data', ( b'\x02\x00\x00\x10\x01\x01\x00\x01\x00\x00\x00\x08\x02\x00\x00\x02' + # proposal 1 b'\x00\x00\x00\x10\x02\x01\x04\x01\x00\x01\x02\x03\x00\x00\x00\x08\x02\x00\x00\x02' # proposal 2 )), ]) self.sa_bytes_multiple = b''.join(self.sa_dict_multiple.values()) def test_parse(self): parsed_sa = Ikev2PayloadSecurityAssociation.parse_exact_size(self.sa_bytes_single) self.assertEqual(parsed_sa.next_payload, Ikev2PayloadType.NONE) self.assertEqual(parsed_sa.flags, set()) self.assertEqual(len(parsed_sa.proposals), 1) self.assertEqual(parsed_sa.proposals[0].protocol_id, Ikev2ProtocolId.IKE) parsed_sa_multiple = Ikev2PayloadSecurityAssociation.parse_exact_size(self.sa_bytes_multiple) self.assertEqual(parsed_sa_multiple.next_payload, Ikev2PayloadType.KE) self.assertEqual(parsed_sa_multiple.flags, {Ikev2PayloadFlags.CRITICAL}) self.assertEqual(len(parsed_sa_multiple.proposals), 2) def test_compose(self): composed_bytes = self.sa_single_proposal.compose() self.assertEqual(composed_bytes, self.sa_bytes_single) composed_bytes_multiple = self.sa_multiple_proposals.compose() self.assertEqual(composed_bytes_multiple, self.sa_bytes_multiple) def test_round_trip(self): composed_bytes = self.sa_single_proposal.compose() parsed_sa: Ikev2PayloadSecurityAssociation = Ikev2PayloadSecurityAssociation.parse_exact_size(composed_bytes) self.assertEqual(parsed_sa.next_payload, self.sa_single_proposal.next_payload) self.assertEqual(parsed_sa.flags, self.sa_single_proposal.flags) self.assertEqual(len(parsed_sa.proposals), len(self.sa_single_proposal.proposals)) def test_get_payload_type(self): self.assertEqual(Ikev2PayloadSecurityAssociation.get_payload_type(), Ikev2PayloadType.SA) def test_error_parse_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: Ikev2PayloadSecurityAssociation.parse_exact_size(b'\x00\x01\x02') self.assertEqual(context_manager.exception.bytes_needed, 1) incomplete = self.sa_bytes_single[:-1] with self.assertRaises(NotEnoughData) as context_manager: Ikev2PayloadSecurityAssociation.parse_exact_size(incomplete) self.assertEqual(context_manager.exception.bytes_needed, 1) def test_error_invalid_proposals(self): with self.assertRaises(TypeError): Ikev2PayloadSecurityAssociation( flags=set(), proposals=["invalid_proposal"] ) def test_get_transform_by_type(self): transform_prf = Ikev2TransformPrf( transform_id=Ikev2PseudorandomFunction.PRF_HMAC_SHA1 ) transform_prf.next_payload = TransformNextPayload.LAST transform_dh = Ikev2TransformDhGroup( transform_id=Ikev2DiffieHellmanGroup.MODP_GROUP_2048_BIT ) transform_dh.next_payload = TransformNextPayload.LAST transform_encr = Ikev2TransformEncryptionAlgorithm( transform_id=Ikev2EncryptionAlgorithm.ENCR_AES_CBC, key_length=128 ) transform_encr.next_payload = TransformNextPayload.LAST proposal = Ikev2Proposal( protocol_id=Ikev2ProtocolId.IKE, transforms=[transform_prf, transform_dh, transform_encr], spi=b'' ) sa_payload = Ikev2PayloadSecurityAssociation( flags=set(), proposals=[proposal] ) found_prf = sa_payload.get_transform_by_type(Ikev2TransformType.PRF) self.assertEqual(found_prf, transform_prf) self.assertEqual(found_prf.get_transform_type(), Ikev2TransformType.PRF) found_dh = sa_payload.get_transform_by_type(Ikev2TransformType.DH) self.assertEqual(found_dh, transform_dh) self.assertEqual(found_dh.get_transform_type(), Ikev2TransformType.DH) found_encr = sa_payload.get_transform_by_type(Ikev2TransformType.ENCR) self.assertEqual(found_encr, transform_encr) self.assertEqual(found_encr.get_transform_type(), Ikev2TransformType.ENCR) transform_integ = Ikev2TransformIntegrity( transform_id=Ikev2IntegrityAlgorithm.AUTH_HMAC_SHA1_96 ) transform_integ.next_payload = TransformNextPayload.LAST proposal_with_integ = Ikev2Proposal( protocol_id=Ikev2ProtocolId.IKE, transforms=[transform_integ], spi=b'' ) sa_payload_with_integ = Ikev2PayloadSecurityAssociation( flags=set(), proposals=[proposal_with_integ] ) found_integ = sa_payload_with_integ.get_transform_by_type(Ikev2TransformType.INTEG) self.assertEqual(found_integ, transform_integ) self.assertEqual(found_integ.get_transform_type(), Ikev2TransformType.INTEG) def test_get_transform_by_type_multiple_proposals(self): transform_prf1 = Ikev2TransformPrf( transform_id=Ikev2PseudorandomFunction.PRF_HMAC_SHA1 ) transform_prf1.next_payload = TransformNextPayload.LAST transform_prf2 = Ikev2TransformPrf( transform_id=Ikev2PseudorandomFunction.PRF_HMAC_SHA2_256 ) transform_prf2.next_payload = TransformNextPayload.LAST proposal1 = Ikev2Proposal( protocol_id=Ikev2ProtocolId.IKE, transforms=[transform_prf1], spi=b'' ) proposal2 = Ikev2Proposal( protocol_id=Ikev2ProtocolId.IKE, transforms=[transform_prf2], spi=b'' ) sa_payload = Ikev2PayloadSecurityAssociation( flags=set(), proposals=[proposal1, proposal2] ) found_prf = sa_payload.get_transform_by_type(Ikev2TransformType.PRF) self.assertEqual(found_prf.get_transform_type(), Ikev2TransformType.PRF) self.assertIn(found_prf, [transform_prf1, transform_prf2]) def test_get_transform_by_type_not_found(self): sa_payload = Ikev2PayloadSecurityAssociation( flags=set(), proposals=[self.proposal_single] ) with self.assertRaises(KeyError) as context_manager: sa_payload.get_transform_by_type(Ikev2TransformType.DH) self.assertEqual(context_manager.exception.args[0], Ikev2TransformType.DH) def test_get_transform_by_type_empty_proposals(self): sa_payload = Ikev2PayloadSecurityAssociation( flags=set(), proposals=[] ) with self.assertRaises(KeyError) as context_manager: sa_payload.get_transform_by_type(Ikev2TransformType.PRF) self.assertEqual(context_manager.exception.args[0], Ikev2TransformType.PRF) if __name__ == '__main__': unittest.main() cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ike/test_isakmp.py000066400000000000000000000452021524413560000265010ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 """ Test ISAKMP header parsing and composition. """ import collections import unittest from test.ike.classes import Ikev2PayloadBaseTest from cryptodatahub.common.exception import InvalidValue from cryptodatahub.ike.algorithm import ( Ikev1ExchangeType, Ikev1PayloadType, Ikev2DiffieHellmanGroup, Ikev2ExchangeType, Ikev2PayloadType, ) from cryptodatahub.ike.version import IkeVersion from cryptoparser.common.exception import InvalidType, NotEnoughData from cryptoparser.ike.common import IkePayloadTypeUnknown from cryptoparser.ike.ikev1 import Ikev1PayloadKeyExchange, Ikev1PayloadNonce, Ikev1PayloadUnparsed from cryptoparser.ike.ikev2 import Ikev2PayloadNonce, Ikev2PayloadKeyExchange, Ikev2PayloadUnparsed from cryptoparser.ike.isakmp import IsakmpFlags, IsakmpMessage from cryptoparser.ike.version import IsakmpProtocolVersion class TestISAKMPHeader(unittest.TestCase): def setUp(self): self.header_dict = collections.OrderedDict([ ('initiator_cookie', b'\x00' * 8), ('responder_cookie', b'\x00' * 8), ('next_payload', b'\x00'), # NONE ('protocol_version', b'\x11'), # ISAKMP v1.1 ('exchange_type', b'\x01'), # BASE ('flags', b'\x00'), ('message_id', b'\x00' * 4), ('length', b'\x00\x00\x00\x1c'), # 28 bytes ]) self.header_bytes = b''.join(self.header_dict.values()) self.header = IsakmpMessage( version=IsakmpProtocolVersion(IkeVersion.V1, 1), initiator_spi=0, responder_spi=0, exchange_type=Ikev1ExchangeType.BASE, flags=set(), message_id=0, payloads=[] ) def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: IsakmpMessage.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, IsakmpMessage.HEADER_SIZE) with self.assertRaises(InvalidType): IsakmpMessage.parse_exact_size(b'\x00' * (IsakmpMessage.HEADER_SIZE - 1) + b'\xff') def test_parse(self): header = IsakmpMessage.parse_exact_size(self.header_bytes) self.assertEqual(header.initiator_spi, 0) self.assertEqual(header.responder_spi, 0) self.assertEqual(header.version, IsakmpProtocolVersion(IkeVersion.V1, 1)) self.assertEqual(header.exchange_type, Ikev1ExchangeType.BASE) self.assertEqual(header.flags, set()) self.assertEqual(header.message_id, 0) def test_compose(self): self.assertEqual(self.header.compose(), self.header_bytes) def test_flags(self): header_with_flags = IsakmpMessage( version=IsakmpProtocolVersion(IkeVersion.V1, 1), initiator_spi=0, responder_spi=0, exchange_type=Ikev1ExchangeType.BASE, flags={IsakmpFlags.ENCRYPTION, IsakmpFlags.COMMIT}, message_id=0, payloads=[] ) header_bytes_with_flags = bytearray(self.header_bytes) header_bytes_with_flags[19] = 0x03 # Set ENCRYPTION and COMMIT flags self.assertEqual(header_with_flags.compose(), bytes(header_bytes_with_flags)) def test_ikev1_compose_payload_type(self): payload = Ikev1PayloadNonce(nonce_data=b'A' * 16) message = IsakmpMessage( version=IsakmpProtocolVersion(IkeVersion.V1, 0), initiator_spi=0, responder_spi=0, exchange_type=Ikev1ExchangeType.INFORMATIONAL, flags=set(), message_id=0, payloads=[payload], ) composed = message.compose() self.assertEqual(composed[16], Ikev1PayloadType.NONCE.value.code) def test_ikev2_parsing(self): header_dict_v2 = collections.OrderedDict([ ('initiator_cookie', b'\x00' * 8), ('responder_cookie', b'\x00' * 8), ('next_payload', b'\x00'), # NONE ('protocol_version', b'\x20'), # ISAKMP v2.0 ('exchange_type', b'\x22'), # IKE_SA_INIT ('flags', b'\x00'), ('message_id', b'\x00' * 4), ('length', b'\x00\x00\x00\x1c'), # 28 bytes ]) header_bytes_v2 = b''.join(header_dict_v2.values()) header = IsakmpMessage.parse_exact_size(header_bytes_v2) self.assertEqual(header.version, IsakmpProtocolVersion(IkeVersion.V2, 0)) self.assertEqual(header.exchange_type, Ikev2ExchangeType.IKE_SA_INIT) def test_invalid_payload_type_ikev1(self): header_dict = collections.OrderedDict([ ('initiator_cookie', b'\x00' * 8), ('responder_cookie', b'\x00' * 8), ('next_payload', b'\xff'), # Invalid payload type ('protocol_version', b'\x11'), # ISAKMP v1.1 ('exchange_type', b'\x01'), # BASE ('flags', b'\x00'), ('message_id', b'\x00' * 4), ('length', b'\x00\x00\x00\x1c'), # 28 bytes ]) header_bytes = b''.join(header_dict.values()) with self.assertRaises(InvalidValue): IsakmpMessage.parse_exact_size(header_bytes) def test_ikev2_compose_with_empty_payloads(self): header_v2 = IsakmpMessage( version=IsakmpProtocolVersion(IkeVersion.V2, 0), initiator_spi=0, responder_spi=0, exchange_type=Ikev2ExchangeType.IKE_SA_INIT, flags=set(), message_id=0, payloads=[] ) # Empty payloads compose to a header-only 28-octet IKE message # (RFC 7296 §3.1) with ``next_payload = NONE``. composed = header_v2.compose() self.assertEqual(len(composed), IsakmpMessage.HEADER_SIZE) # Byte offset 16 (after 8-octet initiator + 8-octet responder SPI) # carries ``next_payload``; NONE is 0 for both IKEv1 / IKEv2. self.assertEqual(composed[16], 0) def test_payload_parsing_with_invalid_next_payload(self): header_dict = collections.OrderedDict([ ('initiator_cookie', b'\x00' * 8), ('responder_cookie', b'\x00' * 8), ('next_payload', b'\x01'), # SECURITY_ASSOCIATION ('protocol_version', b'\x11'), # ISAKMP v1.1 ('exchange_type', b'\x01'), # BASE ('flags', b'\x00'), ('message_id', b'\x00' * 4), ('length', b'\x00\x00\x00\x20'), # 32 bytes (header + 4 bytes payload) ]) payload_data = b'\x00\x00\x00\x04' # Invalid/minimal payload header_bytes = b''.join(header_dict.values()) + payload_data with self.assertRaises((NotEnoughData, TypeError, AttributeError)): IsakmpMessage.parse_exact_size(header_bytes) def test_payload_composition_loop(self): payload = Ikev2PayloadBaseTest( flags=set(), test_data=b'\x00\x01\x02\x03' ) message = IsakmpMessage( version=IsakmpProtocolVersion(IkeVersion.V2, 0), initiator_spi=0, responder_spi=0, exchange_type=Ikev2ExchangeType.IKE_SA_INIT, flags=set(), message_id=0, payloads=[payload] ) message_bytes = message.compose() payload_bytes = payload.compose() self.assertEqual(message_bytes[-len(payload_bytes):], payload_bytes) def test_payload_parsing_loop(self): header_dict = collections.OrderedDict([ ('initiator_cookie', b'\x00' * 8), ('responder_cookie', b'\x00' * 8), ('next_payload', b'\x28'), # NONCE payload type (0x28 = 40 decimal) ('protocol_version', b'\x20'), # ISAKMP v2.0 ('exchange_type', b'\x22'), # IKE_SA_INIT ('flags', b'\x00'), ('message_id', b'\x00' * 4), ('length', b'\x00\x00\x00\x30'), # 48 bytes total (28 header + 20 payload) ]) nonce_payload = ( b'\x00' + # Next payload = NONE (0x00) b'\x00' + # Flags = 0 (no critical bit) b'\x00\x14' + # Payload length = 20 bytes (header 4 + data 16) b'A' * 16 # 16 bytes of nonce data (minimum required for NONCE) ) header_bytes = b''.join(header_dict.values()) + nonce_payload message = IsakmpMessage.parse_exact_size(header_bytes) self.assertEqual(len(message.payloads), 1) self.assertEqual(message.payloads[0].get_payload_type(), Ikev2PayloadType.NONCE) def test_invalid_payload_type_ikev2(self): header_dict = collections.OrderedDict([ ('initiator_cookie', b'\x00' * 8), ('responder_cookie', b'\x00' * 8), ('next_payload', b'\xff'), # Invalid payload type ('protocol_version', b'\x20'), # ISAKMP v2.0 ('exchange_type', b'\x22'), # IKE_SA_INIT ('flags', b'\x00'), ('message_id', b'\x00' * 4), ('length', b'\x00\x00\x00\x1c'), # 28 bytes ]) header_bytes = b''.join(header_dict.values()) with self.assertRaises(InvalidValue): IsakmpMessage.parse_exact_size(header_bytes) def test_unsupported_payload_type_ikev2(self): # ``next_payload`` is a valid enum (TSI) for which no parser # class is registered in :data:`IKEV2_PAYLOAD_CLASSES_BY_TYPE`. # Fall back to :class:`Ikev2PayloadUnparsed` and keep walking # the payload chain instead of raising ``InvalidType``. payload_body = b'\xab\xcd' payload_length = len(payload_body) + Ikev2PayloadUnparsed.HEADER_SIZE # 6 message_length = 28 + payload_length header_dict = collections.OrderedDict([ ('initiator_cookie', b'\x00' * 8), ('responder_cookie', b'\x00' * 8), ('next_payload', bytes([Ikev2PayloadType.TSI.value.code])), ('protocol_version', b'\x20'), # ISAKMP v2.0 ('exchange_type', b'\x22'), # IKE_SA_INIT ('flags', b'\x00'), ('message_id', b'\x00' * 4), ('length', message_length.to_bytes(4, 'big')), ]) payload_bytes = b''.join([ bytes([Ikev2PayloadType.NONE.value.code]), b'\x00', payload_length.to_bytes(2, 'big'), payload_body, ]) message = IsakmpMessage.parse_exact_size(b''.join(header_dict.values()) + payload_bytes) self.assertEqual(len(message.payloads), 1) self.assertIsInstance(message.payloads[0], Ikev2PayloadUnparsed) self.assertEqual(bytes(message.payloads[0].payload_data), payload_body) def test_get_payload_by_type_ikev2(self): nonce_payload = Ikev2PayloadNonce( flags=set(), nonce_data=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' ) ke_payload = Ikev2PayloadKeyExchange( flags=set(), dh_group=Ikev2DiffieHellmanGroup.MODP_GROUP_2048_BIT, key_exchange_data=b'\x00\x01\x02\x03\x04\x05\x06\x07' ) message = IsakmpMessage( version=IsakmpProtocolVersion(IkeVersion.V2, 0), initiator_spi=0, responder_spi=0, exchange_type=Ikev2ExchangeType.IKE_SA_INIT, flags=set(), message_id=0, payloads=[nonce_payload, ke_payload] ) found_nonce = message.get_payload_by_type(Ikev2PayloadType.NONCE) self.assertEqual(found_nonce, nonce_payload) self.assertEqual(found_nonce.get_payload_type(), Ikev2PayloadType.NONCE) found_ke = message.get_payload_by_type(Ikev2PayloadType.KE) self.assertEqual(found_ke, ke_payload) self.assertEqual(found_ke.get_payload_type(), Ikev2PayloadType.KE) def test_get_payload_by_type_raises_index_error_on_multiple(self): first_nonce = Ikev2PayloadNonce( flags=set(), nonce_data=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', ) second_nonce = Ikev2PayloadNonce( flags=set(), nonce_data=b'\xff\xfe\xfd\xfc\xfb\xfa\xf9\xf8\xf7\xf6\xf5\xf4\xf3\xf2\xf1\xf0', ) message = IsakmpMessage( version=IsakmpProtocolVersion(IkeVersion.V2, 0), initiator_spi=0, responder_spi=0, exchange_type=Ikev2ExchangeType.IKE_SA_INIT, flags=set(), message_id=0, payloads=[first_nonce, second_nonce], ) with self.assertRaises(IndexError): message.get_payload_by_type(Ikev2PayloadType.NONCE) self.assertEqual( message.get_payloads_by_type(Ikev2PayloadType.NONCE), [first_nonce, second_nonce], ) def test_get_payloads_by_type_returns_empty_list_when_missing(self): message = IsakmpMessage( version=IsakmpProtocolVersion(IkeVersion.V2, 0), initiator_spi=0, responder_spi=0, exchange_type=Ikev2ExchangeType.IKE_SA_INIT, flags=set(), message_id=0, payloads=[], ) self.assertEqual( message.get_payloads_by_type(Ikev2PayloadType.NONCE), [], ) def test_get_payload_by_type_ikev1(self): ke_payload = Ikev1PayloadKeyExchange(key_exchange_data=b'\x00\x01\x02\x03') nonce_payload = Ikev1PayloadNonce(nonce_data=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08') message = IsakmpMessage( version=IsakmpProtocolVersion(IkeVersion.V1, 1), initiator_spi=0, responder_spi=0, exchange_type=Ikev1ExchangeType.BASE, flags=set(), message_id=0, payloads=[ke_payload, nonce_payload] ) found_ke = message.get_payload_by_type(Ikev1PayloadType.KEY_EXCHANGE) self.assertEqual(found_ke, ke_payload) self.assertEqual(found_ke.get_payload_type(), Ikev1PayloadType.KEY_EXCHANGE) found_nonce = message.get_payload_by_type(Ikev1PayloadType.NONCE) self.assertEqual(found_nonce, nonce_payload) self.assertEqual(found_nonce.get_payload_type(), Ikev1PayloadType.NONCE) def test_get_payload_by_type_not_found(self): nonce_payload = Ikev2PayloadNonce( flags=set(), nonce_data=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' ) message = IsakmpMessage( version=IsakmpProtocolVersion(IkeVersion.V2, 0), initiator_spi=0, responder_spi=0, exchange_type=Ikev2ExchangeType.IKE_SA_INIT, flags=set(), message_id=0, payloads=[nonce_payload] ) with self.assertRaises(KeyError) as context_manager: message.get_payload_by_type(Ikev2PayloadType.KE) self.assertEqual(context_manager.exception.args[0], Ikev2PayloadType.KE) def test_get_payload_by_type_empty_payloads(self): message = IsakmpMessage( version=IsakmpProtocolVersion(IkeVersion.V2, 0), initiator_spi=0, responder_spi=0, exchange_type=Ikev2ExchangeType.IKE_SA_INIT, flags=set(), message_id=0, payloads=[] ) with self.assertRaises(KeyError) as context_manager: message.get_payload_by_type(Ikev2PayloadType.NONCE) self.assertEqual(context_manager.exception.args[0], Ikev2PayloadType.NONCE) class TestIkePayloadTypeUnknown(unittest.TestCase): def test_get_byte_num(self): self.assertEqual(IkePayloadTypeUnknown.get_byte_num(), 1) def test_round_trip(self): wrapper = IkePayloadTypeUnknown(code=0xB8) composed = wrapper.compose() self.assertEqual(composed, b'\xb8') parsed = IkePayloadTypeUnknown.parse_exact_size(composed) self.assertEqual(parsed.code, 0xB8) class TestIkev1PayloadUnparsed(unittest.TestCase): def test_round_trip(self): payload = Ikev1PayloadUnparsed(payload_data=b'\xaa\xbb\xcc\xdd') payload.payload_type = Ikev1PayloadType.HASH payload.next_payload = Ikev1PayloadType.NONE composed = payload.compose() parsed = Ikev1PayloadUnparsed.parse_exact_size(composed) self.assertEqual(parsed.payload_data, b'\xaa\xbb\xcc\xdd') self.assertEqual(parsed.next_payload, Ikev1PayloadType.NONE) def test_wire_payload_type_returned(self): payload = Ikev1PayloadUnparsed(payload_data=b'\x00') payload.payload_type = Ikev1PayloadType.HASH self.assertEqual(payload.get_payload_type(), Ikev1PayloadType.HASH) def test_unknown_next_payload_wrapped(self): # Header carries next_payload = 0xB8 (private-use in RFC 2408 §3.10). # Ikev1PayloadUnparsed._parse_header should route through # IkePayloadTypeUnknown instead of raising on the enum lookup. payload_body = b'\x11' header = b'\xb8\x00' + (4 + len(payload_body)).to_bytes(2, 'big') + payload_body parsed = Ikev1PayloadUnparsed.parse_exact_size(header) self.assertIsInstance(parsed.next_payload, IkePayloadTypeUnknown) self.assertEqual(parsed.next_payload.code, 0xB8) def test_compose_with_unknown_next_payload(self): payload = Ikev1PayloadUnparsed(payload_data=b'\x99') payload.payload_type = Ikev1PayloadType.HASH payload.next_payload = IkePayloadTypeUnknown(code=0xB8) composed = payload.compose() # First octet is the raw private-use code. self.assertEqual(composed[0], 0xB8) class TestIkev2PayloadUnparsedRoundTrip(unittest.TestCase): def test_round_trip(self): payload = Ikev2PayloadUnparsed(flags=set(), payload_data=b'\xaa\xbb\xcc\xdd') payload.payload_type = Ikev2PayloadType.SK payload.next_payload = Ikev2PayloadType.NONE composed = payload.compose() parsed = Ikev2PayloadUnparsed.parse_exact_size(composed) self.assertEqual(parsed.payload_data, b'\xaa\xbb\xcc\xdd') self.assertEqual(parsed.next_payload, Ikev2PayloadType.NONE) def test_wire_payload_type_returned(self): payload = Ikev2PayloadUnparsed(flags=set(), payload_data=b'\x00') payload.payload_type = Ikev2PayloadType.SK self.assertEqual(payload.get_payload_type(), Ikev2PayloadType.SK) def test_unknown_next_payload_wrapped(self): payload_body = b'\x11' header = b'\xb8\x00' + (4 + len(payload_body)).to_bytes(2, 'big') + payload_body parsed = Ikev2PayloadUnparsed.parse_exact_size(header) self.assertIsInstance(parsed.next_payload, IkePayloadTypeUnknown) self.assertEqual(parsed.next_payload.code, 0xB8) def test_compose_with_unknown_next_payload(self): payload = Ikev2PayloadUnparsed(flags=set(), payload_data=b'\x99') payload.payload_type = Ikev2PayloadType.SK payload.next_payload = IkePayloadTypeUnknown(code=0xB8) self.assertEqual(payload.compose()[0], 0xB8) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ike/test_version.py000066400000000000000000000061251524413560000267030ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptodatahub.common.grade import Grade from cryptodatahub.ike.version import IkeVersion from cryptoparser.common.exception import InvalidType, NotEnoughData from cryptoparser.ike.version import IsakmpProtocolVersion, IsakmpVersionFactory class TestIsakmpProtocolVersion(unittest.TestCase): def setUp(self): self.version_bytes = b'\x11' self.version = IsakmpProtocolVersion(IkeVersion.V1, 1) def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: IsakmpProtocolVersion.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, 1) with self.assertRaises(InvalidType): IsakmpProtocolVersion.parse_exact_size(b'\x00') def test_parse(self): version = IsakmpProtocolVersion.parse_exact_size(self.version_bytes) self.assertEqual(version.major, IkeVersion.V1) self.assertEqual(version.minor, 1) def test_compose(self): self.assertEqual(self.version.compose(), self.version_bytes) def test_versions(self): version_1_0 = IsakmpProtocolVersion(IkeVersion.V1, 0) self.assertEqual(version_1_0.compose(), b'\x10') self.assertEqual(IsakmpProtocolVersion.parse_exact_size(b'\x10'), version_1_0) version_1_1 = IsakmpProtocolVersion(IkeVersion.V1, 1) self.assertEqual(version_1_1.compose(), b'\x11') self.assertEqual(IsakmpProtocolVersion.parse_exact_size(b'\x11'), version_1_1) version_2_0 = IsakmpProtocolVersion(IkeVersion.V2, 0) self.assertEqual(version_2_0.compose(), b'\x20') self.assertEqual(IsakmpProtocolVersion.parse_exact_size(b'\x20'), version_2_0) def test_grade(self): version_v1 = IsakmpProtocolVersion(IkeVersion.V1, 0) self.assertEqual(version_v1.grade, Grade.DEPRECATED) version_v2 = IsakmpProtocolVersion(IkeVersion.V2, 0) self.assertEqual(version_v2.grade, Grade.SECURE) def test_str(self): version_1_1 = IsakmpProtocolVersion(IkeVersion.V1, 1) self.assertEqual(str(version_1_1), "IKEv1 (1)") version_2_0 = IsakmpProtocolVersion(IkeVersion.V2, 0) self.assertEqual(str(version_2_0), "IKEv2 (0)") def test_version(self): version_1_1 = IsakmpProtocolVersion(IkeVersion.V1, 1) self.assertEqual(version_1_1.version, "1.1") version_2_0 = IsakmpProtocolVersion(IkeVersion.V2, 0) self.assertEqual(version_2_0.version, "2.0") def test_ordering(self): version_1_0 = IsakmpProtocolVersion(IkeVersion.V1, 0) version_1_1 = IsakmpProtocolVersion(IkeVersion.V1, 1) version_2_0 = IsakmpProtocolVersion(IkeVersion.V2, 0) self.assertLess(version_1_0, version_1_1) self.assertLess(version_1_1, version_2_0) self.assertEqual(sorted([version_2_0, version_1_1, version_1_0]), [version_1_0, version_1_1, version_2_0]) class TestIsakmpVersionFactory(unittest.TestCase): def test_get_enum_class(self): self.assertEqual(IsakmpVersionFactory.get_enum_class(), IkeVersion) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ssh/000077500000000000000000000000001524413560000236265ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ssh/__init__.py000066400000000000000000000000431524413560000257340ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ssh/test_ciphersuites.py000066400000000000000000000017501524413560000277510ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptodatahub.ssh.algorithm import ( SshCompressionAlgorithm, SshEncryptionAlgorithm, SshHostKeyAlgorithm, SshKexAlgorithm, SshMacAlgorithm, ) class TestSshAlgorithm(unittest.TestCase): def test_str(self): self.assertEqual(str(SshEncryptionAlgorithm.ACSS_OPENSSH_ORG.value), 'acss@openssh.org') self.assertEqual(str(SshMacAlgorithm.CRYPTICORE_MAC_SSH_COM.value), 'crypticore-mac@ssh.com') self.assertEqual(str(SshKexAlgorithm.DIFFIE_HELLMAN_GROUP1_SHA1.value), 'diffie-hellman-group1-sha1') self.assertEqual(str(SshHostKeyAlgorithm.SSH_ED25519.value), 'ssh-ed25519') self.assertEqual(str(SshCompressionAlgorithm.ZLIB_OPENSSH_COM.value), 'zlib@openssh.com') class TestSshAlgorithmMac(unittest.TestCase): def test_size(self): self.assertEqual(SshMacAlgorithm.HMAC_SHA2_256.value.size, 256) self.assertEqual(SshMacAlgorithm.HMAC_SHA2_256_96.value.size, 96) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ssh/test_key.py000066400000000000000000001662411524413560000260410ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import collections import datetime import ipaddress import unittest from collections import OrderedDict from test.common.classes import TestClasses from cryptodatahub.common.algorithm import Authentication, Hash from cryptodatahub.common.parameter import ECParamWellKnown from cryptodatahub.common.key import ( PublicKey, PublicKeySize, PublicKeyParamsDsa, PublicKeyParamsEcdsa, PublicKeyParamsEddsa, PublicKeyParamsRsa, ) from cryptodatahub.common.exception import InvalidValue from cryptodatahub.ssh.algorithm import SshHostKeyAlgorithm from cryptoparser.common.exception import InvalidType, NotEnoughData from cryptoparser.ssh.key import ( SshCertType, SshCertExtensionVector, SshCertExtensionForceCommand, SshCertExtensionNoPrecenseRequired, SshCertExtensionPermitX11Forwarding, SshCertExtensionPermitAgentForwarding, SshCertExtensionPermitPortForwarding, SshCertExtensionPermitPTY, SshCertExtensionPermitUserRC, SshCertExtensionSourceAddress, SshCertExtensionUnparsed, SshCertConstraintVector, SshCertCriticalOptionVector, SshCertSignature, SshCertValidPrincipals, SshHostPublicKeyVariant, SshHostCertificateV00DSS, SshHostCertificateV00RSA, SshHostCertificateV01DSS, SshHostCertificateV01ECDSA, SshHostCertificateV01EDDSA, SshHostCertificateV01RSA, SshHostKeyDSS, SshHostKeyECDSA, SshHostKeyEDDSA, SshHostKeyRSA, SshString, SshX509Certificate, SshX509CertificateChain, ) class TestPublicKeyBase(unittest.TestCase): def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: SshHostKeyRSA.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, 4) with self.assertRaises(InvalidValue) as context_manager: SshHostKeyRSA.parse_exact_size(b'\x00\x00\x00\x16' + b'non-existing-type-name') self.assertEqual(context_manager.exception.value, 'non-existing-type-name') class TestString(unittest.TestCase): def setUp(self): self.string_bytes = bytes( b'\x00\x00\x00\x06' + b'string' + b'' ) self.string = SshString('string') def test_parse(self): string = SshString.parse_exact_size(self.string_bytes) self.assertEqual(string.value, 'string') def test_compose(self): self.assertEqual(self.string.compose(), self.string_bytes) class TestHostKeyDSS(TestPublicKeyBase): def setUp(self): self.host_key_bytes = bytes( b'\x00\x00\x00\x07' + # host_key_algorithm_length b'ssh-dss' + # host_key_algorithm b'\x00\x00\x00\x04' + # p_length b'\x01\x01\x02\x03' + # p b'\x00\x00\x00\x04' + # q_length b'\x04\x05\x06\x07' + # q b'\x00\x00\x00\x04' + # g_length b'\x08\x09\x0a\x0b' + # g b'\x00\x00\x00\x04' + # y_length b'\x0c\x0d\x0e\x0f' + # y b'' ) self.host_key = SshHostKeyDSS( host_key_algorithm=SshHostKeyAlgorithm.SSH_DSS, public_key=PublicKey.from_params(PublicKeyParamsDsa( prime=0x01010203, order=0x04050607, generator=0x08090a0b, public_key_value=0x0c0d0e0f, )), ) def test_parse(self): host_key = SshHostPublicKeyVariant.parse_exact_size(self.host_key_bytes) self.assertEqual(host_key.public_key.params.prime, 0x01010203) self.assertEqual(host_key.public_key.params.order, 0x04050607) self.assertEqual(host_key.public_key.params.generator, 0x08090a0b) self.assertEqual(host_key.public_key.params.public_key_value, 0x0c0d0e0f) self.assertEqual(host_key.public_key.key_size, 32) def test_compose(self): self.assertEqual(self.host_key.compose(), self.host_key_bytes) def test_asdict(self): self.assertEqual(self.host_key._asdict(), OrderedDict([ ('key_type', 'host key'), ('algorithm', Authentication.DSS), ('size', PublicKeySize(Authentication.DSS, 32)), ('fingerprints', OrderedDict([ (Hash.SHA2_256, 'SHA256:wdClb94C9Lyi38P1o/SEG38glOh3ea5CJl84bZVx2yM='), (Hash.SHA1, 'SHA1:fOmDMlRkSkplVc2vGTmkRY65j/c='), (Hash.MD5, 'MD5:f2:4f:70:62:fc:36:fa:20:25:62:5d:95:1c:6c:5e:63'), ])), ('known_hosts', 'AAAAB3NzaC1kc3MAAAAEAQECAwAAAAQEBQYHAAAABAgJCgsAAAAEDA0ODw=='), ('host_key_algorithm', SshHostKeyAlgorithm.SSH_DSS), ('public_key', PublicKey.from_params(PublicKeyParamsDsa( prime=0x01010203, generator=0x08090a0b, order=0x04050607, public_key_value=0x0c0d0e0f, ))), ])) class TestHostKeyRSA(TestPublicKeyBase): def setUp(self): self.host_key_bytes = bytes( b'\x00\x00\x00\x07' + # host_key_algorithm_length b'ssh-rsa' + # host_key_algorithm b'\x00\x00\x00\x04' + # e_length b'\x01\x01\x02\x03' + # e b'\x00\x00\x00\x04' + # n_length b'\x04\x05\x06\x07' + # n b'' ) self.host_key = SshHostKeyRSA( host_key_algorithm=SshHostKeyAlgorithm.SSH_RSA, public_key=PublicKey.from_params(PublicKeyParamsRsa( modulus=0x04050607, public_exponent=0x01010203, )), ) def test_parse(self): host_key = SshHostPublicKeyVariant.parse_exact_size(self.host_key_bytes) public_key_params = host_key.public_key.params self.assertEqual(public_key_params.public_exponent, 0x01010203) self.assertEqual(public_key_params.modulus, 0x04050607) self.assertEqual(host_key.public_key.key_size, 32) def test_compose(self): self.assertEqual(self.host_key.compose(), self.host_key_bytes) def test_key_size(self): self.assertEqual(self.host_key.key_size, PublicKeySize(Authentication.RSA, 32)) def test_asdict(self): self.assertEqual(self.host_key._asdict(), OrderedDict([ ('key_type', 'host key'), ('algorithm', Authentication.RSA), ('size', PublicKeySize(Authentication.RSA, 32)), ('fingerprints', OrderedDict([ (Hash.SHA2_256, 'SHA256:ZuSq5GtQTjPj8LAwY4UE4gGILhIAh5kDaDkkEYLaRU0='), (Hash.SHA1, 'SHA1:KAG3KmsLUs4OClEUj62npdXcJTg='), (Hash.MD5, 'MD5:0b:40:11:ce:71:86:01:02:2c:7c:9e:13:d9:37:3b:aa'), ])), ('known_hosts', 'AAAAB3NzaC1yc2EAAAAEAQECAwAAAAQEBQYH'), ])) class TestHostKeyECDSA(TestPublicKeyBase): def setUp(self): self.point_x_bytes = b'\x80' + (256 // 8 - 1) * b'\x00' self.point_y_bytes = b'\x40' + (256 // 8 - 1) * b'\x00' self.host_key_bytes = bytes( b'\x00\x00\x00\x13' + # host_key_algorithm_length b'ecdsa-sha2-nistp256' + # host_key_algorithm b'\x00\x00\x00\x08' + # curve_name_length b'nistp256' + # curve_name b'\x00\x00\x00\x41' + # curve_data_length b'\04' + # curve_data self.point_x_bytes + self.point_y_bytes + b'' ) self.host_key = SshHostKeyECDSA( host_key_algorithm=SshHostKeyAlgorithm.ECDSA_SHA2_NISTP256, public_key=PublicKey.from_params(PublicKeyParamsEcdsa( key_parameter=ECParamWellKnown.PRIME256V1, point_x=2 ** 255, point_y=2 ** 254, )), ) def test_parse(self): host_key = SshHostPublicKeyVariant.parse_exact_size(self.host_key_bytes) self.assertEqual(host_key.public_key.params.key_parameter, ECParamWellKnown.PRIME256V1) self.assertEqual(host_key.public_key.params.point_x, 2 ** 255) self.assertEqual(host_key.public_key.params.point_y, 2 ** 254) self.assertEqual(host_key.public_key.key_size, 256) def test_compose(self): self.assertEqual(self.host_key.compose(), self.host_key_bytes) def test_asdict(self): self.assertEqual(self.host_key._asdict(), OrderedDict([ ('key_type', 'host key'), ('algorithm', Authentication.ECDSA), ('size', PublicKeySize(Authentication.ECDSA, 256)), ('fingerprints', OrderedDict([ (Hash.SHA2_256, 'SHA256:+baTTAvJKIn0rfi1HVlDxDb/lIzi41H9UoCkFPyyO4I='), (Hash.SHA1, 'SHA1:WiBKhHvCyV8LpdXgJWrJr9WAhqw='), (Hash.MD5, 'MD5:86:c6:d5:ca:3e:5e:82:95:31:80:8a:30:b3:a3:6e:80'), ])), ('known_hosts', ( 'AAAAE2VjZHNhLXNoYTItbmlzdHAyNTYAAAAIbmlzdHAyNTYAAABBBIAAAAAAAAAA' 'AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAQAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA' 'AAAAAAAAAAA=' )), ])) class TestHostKeyEDDSA(TestPublicKeyBase): def setUp(self): self.host_key_bytes = bytes( b'\x00\x00\x00\x0b' + # host_key_algorithm_length b'ssh-ed25519' + # host_key_algorithm b'\x00\x00\x00\x20' + # key_data_length b'\x00\x01\x02\x03' * 8 + # key_data b'' ) self.host_key = SshHostKeyEDDSA( host_key_algorithm=SshHostKeyAlgorithm.SSH_ED25519, public_key=PublicKey.from_params(PublicKeyParamsEddsa( key_parameter=ECParamWellKnown.CURVE25519, key_data=b'\x00\x01\x02\x03' * 8, )), ) def test_parse(self): host_key = SshHostPublicKeyVariant.parse_exact_size(self.host_key_bytes) self.assertEqual(host_key.public_key.params.key_data, b'\x00\x01\x02\x03' * 8) self.assertEqual(host_key.public_key.key_size, 256) def test_compose(self): self.assertEqual(self.host_key.compose(), self.host_key_bytes) def test_asdict(self): self.assertEqual(self.host_key._asdict(), OrderedDict([ ('key_type', 'host key'), ('algorithm', Authentication.EDDSA), ('size', PublicKeySize(Authentication.EDDSA, 256)), ('fingerprints', OrderedDict([ (Hash.SHA2_256, 'SHA256:tE7ReEqO7s6dpo8PFRQXhAe4Vdy9HkT0UawUx+NfClk='), (Hash.SHA1, 'SHA1:0ql0OFSSY06SHsZhtEAJ1Wx2AUo='), (Hash.MD5, 'MD5:5c:83:5d:46:c1:7e:b0:47:70:6c:2c:30:a1:28:87:c0'), ])), ('known_hosts', 'AAAAC3NzaC1lZDI1NTE5AAAAIAABAgMAAQIDAAECAwABAgMAAQIDAAECAwABAgMAAQID'), ])) class TestCertType(unittest.TestCase): def test_as_markdown(self): self.assertEqual(SshCertType.SSH_CERT_TYPE_USER.value.as_markdown(), 'User') self.assertEqual(SshCertType.SSH_CERT_TYPE_HOST.value.as_markdown(), 'Host') class TestCertExtensionUnparsed(unittest.TestCase): def setUp(self): self.extension_dict = collections.OrderedDict([ ('extension_name_length', b'\x00\x00\x00\x10'), ('extension_name', b'extension name 1'), ('extension_data_length', b'\x00\x00\x00\x10'), ('extension_data', b'extension data 1'), ]) self.extension_bytes = b''.join(self.extension_dict.values()) self.extension = SshCertExtensionUnparsed( 'extension name 1', b'extension data 1', ) def test_parse(self): extension = SshCertExtensionUnparsed.parse_exact_size(self.extension_bytes) self.assertEqual(extension.extension_name, self.extension.extension_name) self.assertEqual(extension.extension_data, self.extension.extension_data) def test_compose(self): self.assertEqual(self.extension.compose(), self.extension_bytes) class TestHostCertExtensionsUnparsed(unittest.TestCase): def setUp(self): self.extensions_bytes = ( b'\x00\x00\x00\x50' + b'\x00\x00\x00\x10' + b'extension name 1' + b'\x00\x00\x00\x10' + b'extension data 1' + b'\x00\x00\x00\x10' + b'extension name 2' + b'\x00\x00\x00\x10' + b'extension data 2' + b'' ) self.extensions = SshCertExtensionVector([ SshCertExtensionUnparsed('extension name 1', b'extension data 1'), SshCertExtensionUnparsed('extension name 2', b'extension data 2'), ]) def test_parse(self): self.assertEqual( SshCertExtensionVector.parse_exact_size(self.extensions_bytes), self.extensions ) def test_compose(self): self.assertEqual(self.extensions.compose(), self.extensions_bytes) class TestHostCertExtensionsNoData(unittest.TestCase): def setUp(self): self.extensions_bytes = ( b'\x00\x00\x00\x9e' + b'\x00\x00\x00\x14' + b'no-presence-required' + b'\x00\x00\x00\x00' + b'\x00\x00\x00\x15' + b'permit-X11-forwarding' + b'\x00\x00\x00\x00' + b'\x00\x00\x00\x17' + b'permit-agent-forwarding' + b'\x00\x00\x00\x00' + b'\x00\x00\x00\x16' + b'permit-port-forwarding' + b'\x00\x00\x00\x00' + b'\x00\x00\x00\x0a' + b'permit-pty' + b'\x00\x00\x00\x00' + b'\x00\x00\x00\x0e' + b'permit-user-rc' + b'\x00\x00\x00\x00' + b'' ) self.extensions = SshCertExtensionVector([ SshCertExtensionNoPrecenseRequired(), SshCertExtensionPermitX11Forwarding(), SshCertExtensionPermitAgentForwarding(), SshCertExtensionPermitPortForwarding(), SshCertExtensionPermitPTY(), SshCertExtensionPermitUserRC(), ]) def test_parse(self): self.assertEqual( SshCertExtensionVector.parse_exact_size(self.extensions_bytes), self.extensions ) def test_compose(self): self.assertEqual(self.extensions.compose(), self.extensions_bytes) class TestHostCertExtensionsWithData(unittest.TestCase): def setUp(self): self.extensions_bytes = ( b'\x00\x00\x00\x4b' + b'\x00\x00\x00\x0d' + b'force-command' b'\x00\x00\x00\x07' + b'command' b'\x00\x00\x00\x0e' + b'source-address' + b'\x00\x00\x00\x19' + b'192.168.0.0/16,10.0.0.0/8' + b'' ) self.extensions = SshCertCriticalOptionVector([ SshCertExtensionForceCommand('command'), SshCertExtensionSourceAddress([ ipaddress.IPv4Network('192.168.0.0/16'), ipaddress.IPv4Network('10.0.0.0/8'), ]), ]) def test_parse(self): self.assertEqual( SshCertCriticalOptionVector.parse_exact_size(self.extensions_bytes), self.extensions ) def test_compose(self): self.assertEqual(self.extensions.compose(), self.extensions_bytes) class TestHostCertificateDSSBase(TestPublicKeyBase): def setUp(self): self.host_key_bytes = bytes( b'\x00\x00\x00\x07' + # certificate_type b'ssh-dss' + b'\x00\x00\x00\x04' + # p_length b'\x01\x01\x02\x03' + # p b'\x00\x00\x00\x04' + # q_length b'\x04\x05\x06\x07' + # q b'\x00\x00\x00\x04' + # g_length b'\x08\x09\x0a\x0b' + # g b'\x00\x00\x00\x04' + # y_length b'\x0c\x0d\x0e\x0f' + # y b'\x00\x00\x00\x13' + b'\x00\x00\x00\x07' + # signature_type b'ssh-dss' + b'\x00\x00\x00\x04' + # signature_data b'\x00\x01\x02\x03' + b'' ) class TestHostCertificateV00DSS(TestHostCertificateDSSBase): def setUp(self): super().setUp() self.host_cert_bytes = bytes( b'\x00\x00\x00\x1c' + b'ssh-dss-cert-v00@openssh.com' + b'\x00\x00\x00\x04' + # p_length b'\x01\x01\x02\x03' + # p b'\x00\x00\x00\x04' + # q_length b'\x04\x05\x06\x07' + # q b'\x00\x00\x00\x04' + # g_length b'\x08\x09\x0a\x0b' + # g b'\x00\x00\x00\x04' + # y_length b'\x0c\x0d\x0e\x0f' + # y b'\x00\x00\x00\x02' + # certificate_type (SshCertType.SSH_CERT_TYPE_HOST) b'\x00\x00\x00\x08' + # key_id b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x00\x00\x00' + # valid_principals b'\x00\x00\x00\x00\x00\x00\x00\x00' + # valid_after b'\xff\xff\xff\xff\xff\xff\xff\xff' + # valid_before b'\x00\x00\x00\x00' + # constraints b'\x00\x00\x00\x04' + # nonce b'\x00\x01\x02\x03' + b'\x00\x00\x00\x00' + # reserved b'\x00\x00\x00\x2b' + # signature_key self.host_key_bytes + b'' ) self.host_cert = SshHostCertificateV00DSS( host_key_algorithm=SshHostKeyAlgorithm.SSH_DSS_CERT_V00_OPENSSH_COM, nonce=b'\x00\x01\x02\x03', public_key=PublicKey.from_params(PublicKeyParamsDsa( prime=0x01010203, generator=0x08090a0b, order=0x04050607, public_key_value=0x0c0d0e0f, )), certificate_type=SshCertType.SSH_CERT_TYPE_HOST, key_id='\x00\x01\x02\x03\x04\x05\x06\x07', valid_principals=SshCertValidPrincipals([]), valid_after=datetime.datetime.fromtimestamp(0, datetime.timezone.utc), valid_before=None, constraints=SshCertConstraintVector([]), reserved=b'', signature_key=SshHostKeyDSS( SshHostKeyAlgorithm.SSH_DSS, PublicKey.from_params(PublicKeyParamsDsa( prime=0x01010203, generator=0x08090a0b, order=0x04050607, public_key_value=0x0c0d0e0f, )) ), signature=SshCertSignature( SshHostKeyAlgorithm.SSH_DSS, b'\x00\x01\x02\x03', ) ) def test_parse(self): host_cert = SshHostPublicKeyVariant.parse_exact_size(self.host_cert_bytes) self.assertEqual(host_cert.host_key_algorithm, self.host_cert.host_key_algorithm) self.assertEqual(host_cert.nonce, self.host_cert.nonce) self.assertEqual( host_cert.public_key.params.prime, self.host_cert.public_key.params.prime ) self.assertEqual( host_cert.public_key.params.order, self.host_cert.public_key.params.order ) self.assertEqual( host_cert.public_key.params.generator, self.host_cert.public_key.params.generator ) self.assertEqual( host_cert.public_key.params.public_key_value, self.host_cert.public_key.params.public_key_value ) self.assertEqual(host_cert.certificate_type, self.host_cert.certificate_type) self.assertEqual(host_cert.key_id, self.host_cert.key_id) self.assertEqual(host_cert.valid_principals, self.host_cert.valid_principals) self.assertEqual(host_cert.valid_after, self.host_cert.valid_after) self.assertEqual(host_cert.valid_before, self.host_cert.valid_before) self.assertEqual(host_cert.constraints, self.host_cert.constraints) self.assertEqual(host_cert.reserved, self.host_cert.reserved) self.assertEqual(host_cert.signature_key, self.host_cert.signature_key) self.assertEqual(host_cert.signature, self.host_cert.signature) def test_compose(self): self.assertEqual(self.host_cert.compose(), self.host_cert_bytes) self.assertEqual(self.host_cert.key_bytes, self.host_cert_bytes) def test_asdict(self): self.assertEqual(self.host_cert._asdict(), OrderedDict([ ('key_type', 'host certificate'), ('algorithm', Authentication.DSS), ('size', PublicKeySize(Authentication.DSS, 32)), ('fingerprints', OrderedDict([ (Hash.SHA2_256, 'SHA256:JLbl6U9Dd3zrR/pS86OLLJUsL6NU7NOnSd554vL1md0='), (Hash.SHA1, 'SHA1:eGd4yAwFTXy+Wdi5xsuEJUNsj0M='), (Hash.MD5, 'MD5:01:63:07:cb:45:12:bf:96:4c:fd:62:cd:40:6d:e0:99') ])), ('known_hosts', ( 'AAAAHHNzaC1kc3MtY2VydC12MDBAb3BlbnNzaC5jb20AAAAEAQECAwAAAAQEBQYH' 'AAAABAgJCgsAAAAEDA0ODwAAAAIAAAAIAAECAwQFBgcAAAAAAAAAAAAAAAD/////' '/////wAAAAAAAAAEAAECAwAAAAAAAAArAAAAB3NzaC1kc3MAAAAEAQECAwAAAAQE' 'BQYHAAAABAgJCgsAAAAEDA0ODwAAABMAAAAHc3NoLWRzcwAAAAQAAQID' )), ('host_key_algorithm', SshHostKeyAlgorithm.SSH_DSS_CERT_V00_OPENSSH_COM), ('public_key', PublicKey.from_params(PublicKeyParamsDsa( prime=0x01010203, generator=0x08090a0b, order=0x04050607, public_key_value=0x0c0d0e0f, ))), ('certificate_type', SshCertType.SSH_CERT_TYPE_HOST), ('key_id', '\x00\x01\x02\x03\x04\x05\x06\x07'), ('valid_principals', SshCertValidPrincipals([])), ('valid_after', datetime.datetime.fromtimestamp(0, datetime.timezone.utc)), ('valid_before', None), ('constraints', SshCertConstraintVector([])), ('nonce', b'\x00\x01\x02\x03'), ('reserved', b''), ('signature_key', SshHostKeyDSS( host_key_algorithm=SshHostKeyAlgorithm.SSH_DSS, public_key=PublicKey.from_params(PublicKeyParamsDsa( prime=0x01010203, generator=0x08090a0b, order=0x04050607, public_key_value=0x0c0d0e0f, )), )), ('signature', SshCertSignature( signature_type=SshHostKeyAlgorithm.SSH_DSS, signature_data=b'\x00\x01\x02\x03' )) ])) class TestHostCertificateV01DSS(TestHostCertificateDSSBase): def setUp(self): super().setUp() self.host_cert_bytes = bytes( b'\x00\x00\x00\x1c' + b'ssh-dss-cert-v01@openssh.com' + b'\x00\x00\x00\x04' + # nonce b'\x00\x01\x02\x03' + b'\x00\x00\x00\x04' + # p_length b'\x01\x01\x02\x03' + # p b'\x00\x00\x00\x04' + # q_length b'\x04\x05\x06\x07' + # q b'\x00\x00\x00\x04' + # g_length b'\x08\x09\x0a\x0b' + # g b'\x00\x00\x00\x04' + # y_length b'\x0c\x0d\x0e\x0f' + # y b'\x01\x02\x03\x04\x05\x06\x07\x08' + # serial b'\x00\x00\x00\x02' + # certificate_type (SshCertType.SSH_CERT_TYPE_HOST) b'\x00\x00\x00\x08' + # key_id b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x00\x00\x00' + # valid_principals b'\x00\x00\x00\x00\x00\x00\x00\x00' + # valid_after b'\xff\xff\xff\xff\xff\xff\xff\xff' + # valid_before b'\x00\x00\x00\x00' + # critical_options b'\x00\x00\x00\x00' + # extensions b'\x00\x00\x00\x00' + # reserved b'\x00\x00\x00\x2b' + # signature_key self.host_key_bytes + b'' ) self.host_cert = SshHostCertificateV01DSS( host_key_algorithm=SshHostKeyAlgorithm.SSH_DSS_CERT_V01_OPENSSH_COM, nonce=b'\x00\x01\x02\x03', public_key=PublicKey.from_params(PublicKeyParamsDsa( prime=0x01010203, generator=0x08090a0b, order=0x04050607, public_key_value=0x0c0d0e0f, )), serial=0x0102030405060708, certificate_type=SshCertType.SSH_CERT_TYPE_HOST, key_id='\x00\x01\x02\x03\x04\x05\x06\x07', valid_principals=SshCertValidPrincipals([]), valid_after=datetime.datetime.fromtimestamp(0, datetime.timezone.utc), valid_before=None, critical_options=SshCertCriticalOptionVector([]), extensions=SshCertExtensionVector([]), reserved=b'', signature_key=SshHostKeyDSS( SshHostKeyAlgorithm.SSH_DSS, PublicKey.from_params(PublicKeyParamsDsa( prime=0x01010203, generator=0x08090a0b, order=0x04050607, public_key_value=0x0c0d0e0f, )), ), signature=SshCertSignature( SshHostKeyAlgorithm.SSH_DSS, b'\x00\x01\x02\x03', ) ) def test_parse(self): host_cert = SshHostPublicKeyVariant.parse_exact_size(self.host_cert_bytes) self.assertEqual(host_cert.host_key_algorithm, self.host_cert.host_key_algorithm) self.assertEqual(host_cert.nonce, self.host_cert.nonce) self.assertEqual( host_cert.public_key.params.prime, self.host_cert.public_key.params.prime ) self.assertEqual( host_cert.public_key.params.order, self.host_cert.public_key.params.order ) self.assertEqual( host_cert.public_key.params.generator, self.host_cert.public_key.params.generator ) self.assertEqual( host_cert.public_key.params.public_key_value, self.host_cert.public_key.params.public_key_value ) self.assertEqual(host_cert.certificate_type, self.host_cert.certificate_type) self.assertEqual(host_cert.key_id, self.host_cert.key_id) self.assertEqual(host_cert.valid_principals, self.host_cert.valid_principals) self.assertEqual(host_cert.valid_after, self.host_cert.valid_after) self.assertEqual(host_cert.valid_before, self.host_cert.valid_before) self.assertEqual(host_cert.critical_options, self.host_cert.critical_options) self.assertEqual(host_cert.extensions, SshCertExtensionVector([])) self.assertEqual(host_cert.reserved, self.host_cert.reserved) self.assertEqual(host_cert.signature_key, self.host_cert.signature_key) self.assertEqual(host_cert.signature, self.host_cert.signature) def test_compose(self): self.assertEqual(self.host_cert.compose(), self.host_cert_bytes) self.assertEqual(self.host_cert.key_bytes, self.host_cert_bytes) def test_asdict(self): self.assertEqual(self.host_cert._asdict(), OrderedDict([ ('key_type', 'host certificate'), ('algorithm', Authentication.DSS), ('size', PublicKeySize(Authentication.DSS, 32)), ('fingerprints', OrderedDict([ (Hash.SHA2_256, 'SHA256:81ny5W3AQivvf/ffzLM0b13/eAG82/GsSFNaBnVOF4A='), (Hash.SHA1, 'SHA1:Kx6Nzq+MWzKiT4q0lxdQbZzecSo='), (Hash.MD5, 'MD5:94:9b:e9:f1:27:4c:7a:60:ba:5b:da:47:99:46:13:06') ])), ('known_hosts', ( 'AAAAHHNzaC1kc3MtY2VydC12MDFAb3BlbnNzaC5jb20AAAAEAAECAwAAAAQBAQID' 'AAAABAQFBgcAAAAECAkKCwAAAAQMDQ4PAQIDBAUGBwgAAAACAAAACAABAgMEBQYH' 'AAAAAAAAAAAAAAAA//////////8AAAAAAAAAAAAAAAAAAAArAAAAB3NzaC1kc3MA' 'AAAEAQECAwAAAAQEBQYHAAAABAgJCgsAAAAEDA0ODwAAABMAAAAHc3NoLWRzcwAA' 'AAQAAQID' )), ('host_key_algorithm', SshHostKeyAlgorithm.SSH_DSS_CERT_V01_OPENSSH_COM), ('public_key', PublicKey.from_params(PublicKeyParamsDsa( prime=0x01010203, generator=0x08090a0b, order=0x04050607, public_key_value=0x0c0d0e0f, ))), ('nonce', b'\x00\x01\x02\x03'), ('serial', 0x0102030405060708), ('certificate_type', SshCertType.SSH_CERT_TYPE_HOST), ('key_id', '\x00\x01\x02\x03\x04\x05\x06\x07'), ('valid_principals', SshCertValidPrincipals([])), ('valid_after', datetime.datetime.fromtimestamp(0, datetime.timezone.utc)), ('valid_before', None), ('critical_options', SshCertCriticalOptionVector([])), ('extensions', SshCertExtensionVector([])), ('reserved', b''), ('signature_key', SshHostKeyDSS( host_key_algorithm=SshHostKeyAlgorithm.SSH_DSS, public_key=PublicKey.from_params(PublicKeyParamsDsa( prime=0x01010203, generator=0x08090a0b, order=0x04050607, public_key_value=0x0c0d0e0f, )), )), ('signature', SshCertSignature( signature_type=SshHostKeyAlgorithm.SSH_DSS, signature_data=b'\x00\x01\x02\x03' )) ])) class TestHostCertificateRSABase(TestPublicKeyBase): def setUp(self): self.host_key_bytes = bytes( b'\x00\x00\x00\x07' + # certificate_type b'ssh-rsa' + b'\x00\x00\x00\x03' + # e b'\x01\x00\x01' + b'\x00\x00\x00\x04' + # n b'\x01\x01\x02\x03' + b'\x00\x00\x00\x13' + b'\x00\x00\x00\x07' + # signature_type b'ssh-rsa' + b'\x00\x00\x00\x04' + # signature_data b'\x00\x01\x02\x03' + b'' ) class TestHostCertificateV00RSA(TestHostCertificateRSABase): def setUp(self): super().setUp() self.host_cert_bytes = bytes( b'\x00\x00\x00\x1c' + b'ssh-rsa-cert-v00@openssh.com' + b'\x00\x00\x00\x01' + # e b'\x03' + b'\x00\x00\x00\x04' + # n b'\x01\x01\x02\x03' + b'\x00\x00\x00\x02' + # certificate_type (SshCertType.SSH_CERT_TYPE_HOST) b'\x00\x00\x00\x08' + # key_id b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x00\x00\x00' + # valid_principals b'\x00\x00\x00\x00\x00\x00\x00\x00' + # valid_after b'\xff\xff\xff\xff\xff\xff\xff\xff' + # valid_before b'\x00\x00\x00\x00' + # constraints b'\x00\x00\x00\x04' + # nonce b'\x00\x01\x02\x03' + b'\x00\x00\x00\x00' + # reserved b'\x00\x00\x00\x1a' + # signature_key self.host_key_bytes + b'' ) self.host_cert = SshHostCertificateV00RSA( host_key_algorithm=SshHostKeyAlgorithm.SSH_RSA_CERT_V00_OPENSSH_COM, nonce=b'\x00\x01\x02\x03', public_key=PublicKey.from_params(PublicKeyParamsRsa( public_exponent=0x03, modulus=0x01010203, )), certificate_type=SshCertType.SSH_CERT_TYPE_HOST, key_id='\x00\x01\x02\x03\x04\x05\x06\x07', valid_principals=SshCertValidPrincipals([]), valid_after=datetime.datetime.fromtimestamp(0, datetime.timezone.utc), valid_before=None, constraints=SshCertConstraintVector([]), reserved=b'', signature_key=SshHostKeyRSA( SshHostKeyAlgorithm.SSH_RSA, PublicKey.from_params(PublicKeyParamsRsa( public_exponent=0x010001, modulus=0x01010203, )), ), signature=SshCertSignature( SshHostKeyAlgorithm.SSH_RSA, b'\x00\x01\x02\x03', ) ) def test_parse(self): host_cert = SshHostPublicKeyVariant.parse_exact_size(self.host_cert_bytes) self.assertEqual(host_cert.host_key_algorithm, self.host_cert.host_key_algorithm) self.assertEqual(host_cert.nonce, self.host_cert.nonce) self.assertEqual(host_cert.public_key.params.public_exponent, self.host_cert.public_key.params.public_exponent) self.assertEqual(host_cert.public_key.params.modulus, self.host_cert.public_key.params.modulus) self.assertEqual(host_cert.certificate_type, self.host_cert.certificate_type) self.assertEqual(host_cert.key_id, self.host_cert.key_id) self.assertEqual(host_cert.valid_principals, self.host_cert.valid_principals) self.assertEqual(host_cert.valid_after, self.host_cert.valid_after) self.assertEqual(host_cert.valid_before, self.host_cert.valid_before) self.assertEqual(host_cert.constraints, self.host_cert.constraints) self.assertEqual(host_cert.reserved, self.host_cert.reserved) self.assertEqual(host_cert.signature_key, self.host_cert.signature_key) self.assertEqual(host_cert.signature, self.host_cert.signature) def test_compose(self): self.assertEqual(self.host_cert.compose(), self.host_cert_bytes) self.assertEqual(self.host_cert.key_bytes, self.host_cert_bytes) def test_asdict(self): self.assertEqual(self.host_cert._asdict(), OrderedDict([ ('key_type', 'host certificate'), ('algorithm', Authentication.RSA), ('size', PublicKeySize(Authentication.RSA, 32)), ('fingerprints', OrderedDict([ (Hash.SHA2_256, 'SHA256:YJVpL0zTCstfPryV5C1tD3boAzBrzRlAMjrAosxw4pA='), (Hash.SHA1, 'SHA1:zNRNUIcyRZvu8MuhAmtALdUNCMM='), (Hash.MD5, 'MD5:9c:ce:d6:f1:7f:96:d1:5f:9d:14:a3:32:74:a1:18:88') ])), ('known_hosts', ( 'AAAAHHNzaC1yc2EtY2VydC12MDBAb3BlbnNzaC5jb20AAAABAwAAAAQBAQIDAAAA' 'AgAAAAgAAQIDBAUGBwAAAAAAAAAAAAAAAP//////////AAAAAAAAAAQAAQIDAAAA' 'AAAAABoAAAAHc3NoLXJzYQAAAAMBAAEAAAAEAQECAwAAABMAAAAHc3NoLXJzYQAA' 'AAQAAQID' )), ('host_key_algorithm', SshHostKeyAlgorithm.SSH_RSA_CERT_V00_OPENSSH_COM), ('public_key', PublicKey.from_params(PublicKeyParamsRsa( public_exponent=0x03, modulus=0x01010203, ))), ('certificate_type', SshCertType.SSH_CERT_TYPE_HOST), ('key_id', '\x00\x01\x02\x03\x04\x05\x06\x07'), ('valid_principals', SshCertValidPrincipals([])), ('valid_after', datetime.datetime.fromtimestamp(0, datetime.timezone.utc)), ('valid_before', None), ('constraints', SshCertConstraintVector([])), ('nonce', b'\x00\x01\x02\x03'), ('reserved', b''), ('signature_key', SshHostKeyRSA( host_key_algorithm=SshHostKeyAlgorithm.SSH_RSA, public_key=PublicKey.from_params(PublicKeyParamsRsa( public_exponent=0x010001, modulus=0x01010203, )), )), ('signature', SshCertSignature( signature_type=SshHostKeyAlgorithm.SSH_RSA, signature_data=b'\x00\x01\x02\x03' )) ])) class TestHostCertificateV01RSA(TestHostCertificateRSABase): def setUp(self): super().setUp() self.host_cert_bytes = bytes( b'\x00\x00\x00\x1c' + b'ssh-rsa-cert-v01@openssh.com' + b'\x00\x00\x00\x04' + # nonce b'\x00\x01\x02\x03' + b'\x00\x00\x00\x01' + # e b'\x03' + b'\x00\x00\x00\x04' + # n b'\x01\x01\x02\x03' + b'\x01\x02\x03\x04\x05\x06\x07\x08' + # serial b'\x00\x00\x00\x02' + # certificate_type (SshCertType.SSH_CERT_TYPE_HOST) b'\x00\x00\x00\x08' + # key_id b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x00\x00\x00' + # valid_principals b'\x00\x00\x00\x00\x00\x00\x00\x00' + # valid_after b'\xff\xff\xff\xff\xff\xff\xff\xff' + # valid_before b'\x00\x00\x00\x00' + # critical_options b'\x00\x00\x00\x00' + # extensions b'\x00\x00\x00\x00' + # reserved b'\x00\x00\x00\x1a' + # signature_key self.host_key_bytes + b'' ) self.host_cert = SshHostCertificateV01RSA( host_key_algorithm=SshHostKeyAlgorithm.SSH_RSA_CERT_V01_OPENSSH_COM, nonce=b'\x00\x01\x02\x03', public_key=PublicKey.from_params(PublicKeyParamsRsa( public_exponent=0x03, modulus=0x01010203, )), serial=0x0102030405060708, certificate_type=SshCertType.SSH_CERT_TYPE_HOST, key_id='\x00\x01\x02\x03\x04\x05\x06\x07', valid_principals=SshCertValidPrincipals([]), valid_after=datetime.datetime.fromtimestamp(0, datetime.timezone.utc), valid_before=None, critical_options=SshCertCriticalOptionVector([]), extensions=SshCertExtensionVector([]), reserved=b'', signature_key=SshHostKeyRSA( SshHostKeyAlgorithm.SSH_RSA, PublicKey.from_params(PublicKeyParamsRsa( public_exponent=0x010001, modulus=0x01010203, )) ), signature=SshCertSignature( SshHostKeyAlgorithm.SSH_RSA, b'\x00\x01\x02\x03', ) ) def test_parse(self): host_cert = SshHostPublicKeyVariant.parse_exact_size(self.host_cert_bytes) self.assertEqual(host_cert.host_key_algorithm, self.host_cert.host_key_algorithm) self.assertEqual(host_cert.nonce, self.host_cert.nonce) self.assertEqual(host_cert.public_key.params.public_exponent, self.host_cert.public_key.params.public_exponent) self.assertEqual(host_cert.public_key.params.modulus, self.host_cert.public_key.params.modulus) self.assertEqual(host_cert.serial, self.host_cert.serial) self.assertEqual(host_cert.certificate_type, self.host_cert.certificate_type) self.assertEqual(host_cert.key_id, self.host_cert.key_id) self.assertEqual(host_cert.valid_principals, self.host_cert.valid_principals) self.assertEqual(host_cert.valid_after, self.host_cert.valid_after) self.assertEqual(host_cert.valid_before, self.host_cert.valid_before) self.assertEqual(host_cert.critical_options, self.host_cert.critical_options) self.assertEqual(host_cert.extensions, SshCertExtensionVector([])) self.assertEqual(host_cert.reserved, self.host_cert.reserved) self.assertEqual(host_cert.signature_key, self.host_cert.signature_key) self.assertEqual(host_cert.signature, self.host_cert.signature) def test_compose(self): self.assertEqual(self.host_cert.compose(), self.host_cert_bytes) self.assertEqual(self.host_cert.key_bytes, self.host_cert_bytes) def test_asdict(self): self.assertEqual(self.host_cert._asdict(), OrderedDict([ ('key_type', 'host certificate'), ('algorithm', Authentication.RSA), ('size', PublicKeySize(Authentication.RSA, 32)), ('fingerprints', OrderedDict([ (Hash.SHA2_256, 'SHA256:vqtvWzklbsElMSGXE0G7Gk7WDHuXzd83KOC0y0Rv9TY='), (Hash.SHA1, 'SHA1:9Iem7KW/rAvahTzVMTmDFg93MBk='), (Hash.MD5, 'MD5:54:14:e3:b7:d8:7e:fa:d0:3c:52:4b:6a:8f:c4:72:6d') ])), ('known_hosts', ( 'AAAAHHNzaC1yc2EtY2VydC12MDFAb3BlbnNzaC5jb20AAAAEAAECAwAAAAEDAAAA' 'BAEBAgMBAgMEBQYHCAAAAAIAAAAIAAECAwQFBgcAAAAAAAAAAAAAAAD/////////' '/wAAAAAAAAAAAAAAAAAAABoAAAAHc3NoLXJzYQAAAAMBAAEAAAAEAQECAwAAABMA' 'AAAHc3NoLXJzYQAAAAQAAQID' )), ('host_key_algorithm', SshHostKeyAlgorithm.SSH_RSA_CERT_V01_OPENSSH_COM), ('public_key', PublicKey.from_params(PublicKeyParamsRsa( public_exponent=0x03, modulus=0x01010203, ))), ('nonce', b'\x00\x01\x02\x03'), ('serial', 0x0102030405060708), ('certificate_type', SshCertType.SSH_CERT_TYPE_HOST), ('key_id', '\x00\x01\x02\x03\x04\x05\x06\x07'), ('valid_principals', SshCertValidPrincipals([])), ('valid_after', datetime.datetime.fromtimestamp(0, datetime.timezone.utc)), ('valid_before', None), ('critical_options', SshCertCriticalOptionVector([])), ('extensions', SshCertExtensionVector([])), ('reserved', b''), ('signature_key', SshHostKeyRSA( host_key_algorithm=SshHostKeyAlgorithm.SSH_RSA, public_key=PublicKey.from_params(PublicKeyParamsRsa( public_exponent=0x010001, modulus=0x01010203, )), )), ('signature', SshCertSignature( signature_type=SshHostKeyAlgorithm.SSH_RSA, signature_data=b'\x00\x01\x02\x03' )) ])) class TestHostCertificateECDSABase(TestPublicKeyBase): def setUp(self): self.point_x_bytes = b'\x80' + (256 // 8 - 1) * b'\x00' self.point_y_bytes = b'\x40' + (256 // 8 - 1) * b'\x00' self.host_key_bytes = bytes( b'\x00\x00\x00\x13' + # certificate_type b'ecdsa-sha2-nistp256' + b'\x00\x00\x00\x08' + # curve_name_length b'nistp256' + # curve_name b'\x00\x00\x00\x41' + # curve_data_length b'\x04' + # curve_data self.point_x_bytes + self.point_y_bytes + b'\x00\x00\x00\x1f' + b'\x00\x00\x00\x13' + # signature_type b'ecdsa-sha2-nistp256' + b'\x00\x00\x00\x04' + # signature_data b'\x00\x01\x02\x03' + b'' ) class TestHostCertificateV01ECDSA(TestHostCertificateECDSABase): def setUp(self): super().setUp() self.point_x_bytes = b'\x80' + (256 // 8 - 1) * b'\x00' self.point_y_bytes = b'\x40' + (256 // 8 - 1) * b'\x00' self.host_cert_bytes = bytes( b'\x00\x00\x00\x28' + b'ecdsa-sha2-nistp256-cert-v01@openssh.com' + b'\x00\x00\x00\x04' + # nonce b'\x00\x01\x02\x03' + b'\x00\x00\x00\x08' + # curve_name_length b'nistp256' + # curve_name b'\x00\x00\x00\x41' + # curve_data_length b'\x04' + # curve_data self.point_x_bytes + self.point_y_bytes + b'\x01\x02\x03\x04\x05\x06\x07\x08' + # serial b'\x00\x00\x00\x02' + # certificate_type (SshCertType.SSH_CERT_TYPE_HOST) b'\x00\x00\x00\x08' + # key_id b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x00\x00\x00' + # valid_principals b'\x00\x00\x00\x00\x00\x00\x00\x00' + # valid_after b'\xff\xff\xff\xff\xff\xff\xff\xff' + # valid_before b'\x00\x00\x00\x00' + # critical_options b'\x00\x00\x00\x00' + # extensions b'\x00\x00\x00\x00' + # reserved b'\x00\x00\x00\x68' + # signature_key self.host_key_bytes + b'' ) self.host_cert = SshHostCertificateV01ECDSA( host_key_algorithm=SshHostKeyAlgorithm.ECDSA_SHA2_NISTP256_CERT_V01_OPENSSH_COM, nonce=b'\x00\x01\x02\x03', public_key=PublicKey.from_params(PublicKeyParamsEcdsa( key_parameter=ECParamWellKnown.PRIME256V1, point_x=2 ** 255, point_y=2 ** 254, )), serial=0x0102030405060708, certificate_type=SshCertType.SSH_CERT_TYPE_HOST, key_id='\x00\x01\x02\x03\x04\x05\x06\x07', valid_principals=SshCertValidPrincipals([]), valid_after=datetime.datetime.fromtimestamp(0, datetime.timezone.utc), valid_before=None, critical_options=SshCertCriticalOptionVector([]), extensions=SshCertExtensionVector([]), reserved=b'', signature_key=SshHostKeyECDSA( host_key_algorithm=SshHostKeyAlgorithm.ECDSA_SHA2_NISTP256, public_key=PublicKey.from_params(PublicKeyParamsEcdsa( key_parameter=ECParamWellKnown.PRIME256V1, point_x=2 ** 255, point_y=2 ** 254, )), ), signature=SshCertSignature( SshHostKeyAlgorithm.ECDSA_SHA2_NISTP256, b'\x00\x01\x02\x03', ) ) def test_parse(self): host_cert = SshHostPublicKeyVariant.parse_exact_size(self.host_cert_bytes) self.assertEqual(host_cert.host_key_algorithm, self.host_cert.host_key_algorithm) self.assertEqual(host_cert.nonce, self.host_cert.nonce) self.assertEqual(host_cert.public_key.params.key_parameter, ECParamWellKnown.PRIME256V1) self.assertEqual(host_cert.public_key.params.point_x, 2 ** 255) self.assertEqual(host_cert.public_key.params.point_y, 2 ** 254) self.assertEqual(host_cert.serial, self.host_cert.serial) self.assertEqual(host_cert.certificate_type, self.host_cert.certificate_type) self.assertEqual(host_cert.key_id, self.host_cert.key_id) self.assertEqual(host_cert.valid_principals, self.host_cert.valid_principals) self.assertEqual(host_cert.valid_after, self.host_cert.valid_after) self.assertEqual(host_cert.valid_before, self.host_cert.valid_before) self.assertEqual(host_cert.critical_options, self.host_cert.critical_options) self.assertEqual(host_cert.extensions, SshCertExtensionVector([])) self.assertEqual(host_cert.reserved, self.host_cert.reserved) self.assertEqual(host_cert.signature_key, self.host_cert.signature_key) self.assertEqual(host_cert.signature, self.host_cert.signature) def test_compose(self): self.assertEqual(self.host_cert.compose(), self.host_cert_bytes) self.assertEqual(self.host_cert.key_bytes, self.host_cert_bytes) def test_asdict(self): self.assertEqual(self.host_cert._asdict(), OrderedDict([ ('key_type', 'host certificate'), ('algorithm', Authentication.ECDSA), ('size', PublicKeySize(Authentication.ECDSA, 256)), ('fingerprints', OrderedDict([ (Hash.SHA2_256, 'SHA256:aRBBmsAqpnH3wrOh35fa//LB/JGcuTUAPuRPEj9RWRQ='), (Hash.SHA1, 'SHA1:9E5qN+JPZBX4ZBfKEou/WNWtWR8='), (Hash.MD5, 'MD5:e9:dd:20:42:a4:36:6b:90:60:29:f8:69:5f:be:e7:13') ])), ('known_hosts', ( 'AAAAKGVjZHNhLXNoYTItbmlzdHAyNTYtY2VydC12MDFAb3BlbnNzaC5jb20AAAAE' 'AAECAwAAAAhuaXN0cDI1NgAAAEEEgAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA' 'AAAAAABAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAECAwQFBgcIAAAA' 'AgAAAAgAAQIDBAUGBwAAAAAAAAAAAAAAAP//////////AAAAAAAAAAAAAAAAAAAA' 'aAAAABNlY2RzYS1zaGEyLW5pc3RwMjU2AAAACG5pc3RwMjU2AAAAQQSAAAAAAAAA' 'AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAEAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA' 'AAAAAAAAAAAAAAAAHwAAABNlY2RzYS1zaGEyLW5pc3RwMjU2AAAABAABAgM=' )), ('host_key_algorithm', SshHostKeyAlgorithm.ECDSA_SHA2_NISTP256_CERT_V01_OPENSSH_COM), ('public_key', PublicKey.from_params(PublicKeyParamsEcdsa( key_parameter=ECParamWellKnown.PRIME256V1, point_x=2 ** 255, point_y=2 ** 254, ))), ('nonce', b'\x00\x01\x02\x03'), ('serial', 0x0102030405060708), ('certificate_type', SshCertType.SSH_CERT_TYPE_HOST), ('key_id', '\x00\x01\x02\x03\x04\x05\x06\x07'), ('valid_principals', SshCertValidPrincipals([])), ('valid_after', datetime.datetime.fromtimestamp(0, datetime.timezone.utc)), ('valid_before', None), ('critical_options', SshCertCriticalOptionVector([])), ('extensions', SshCertExtensionVector([])), ('reserved', b''), ('signature_key', SshHostKeyECDSA( host_key_algorithm=SshHostKeyAlgorithm.ECDSA_SHA2_NISTP256, public_key=PublicKey.from_params(PublicKeyParamsEcdsa( key_parameter=ECParamWellKnown.PRIME256V1, point_x=2 ** 255, point_y=2 ** 254, )), )), ('signature', SshCertSignature( signature_type=SshHostKeyAlgorithm.ECDSA_SHA2_NISTP256, signature_data=b'\x00\x01\x02\x03' )) ])) class TestHostCertificateEDDSABase(TestPublicKeyBase): def setUp(self): self.host_key_bytes = bytes( b'\x00\x00\x00\x17' + # certificate_type b'\x00\x00\x00\x0b' + # host_key_algorithm b'ssh-ed25519' + b'\x00\x00\x00\x04' + # key_data_length b'\x00\x01\x02\x03' + # key_data b'\x00\x00\x00\x17' + b'\x00\x00\x00\x0b' + # signature_type b'ssh-ed25519' + b'\x00\x00\x00\x04' + # signature_data b'\x00\x01\x02\x03' + b'' ) class TestHostCertificateV01EDDSA(TestHostCertificateEDDSABase): def setUp(self): super().setUp() self.host_cert_bytes = bytes( b'\x00\x00\x00\x20' + b'ssh-ed25519-cert-v01@openssh.com' + b'\x00\x00\x00\x04' + # nonce b'\x00\x01\x02\x03' + b'\x00\x00\x00\x04' + # key_data_length b'\x00\x01\x02\x03' + # key_data b'\x01\x02\x03\x04\x05\x06\x07\x08' + # serial b'\x00\x00\x00\x02' + # certificate_type (SshCertType.SSH_CERT_TYPE_HOST) b'\x00\x00\x00\x08' + # key_id b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x00\x00\x00' + # valid_principals b'\x00\x00\x00\x00\x00\x00\x00\x00' + # valid_after b'\xff\xff\xff\xff\xff\xff\xff\xff' + # valid_before b'\x00\x00\x00\x00' + # critical_options b'\x00\x00\x00\x00' + # extensions b'\x00\x00\x00\x00' + # reserved self.host_key_bytes + b'' ) self.host_cert = SshHostCertificateV01EDDSA( host_key_algorithm=SshHostKeyAlgorithm.SSH_ED25519_CERT_V01_OPENSSH_COM, nonce=b'\x00\x01\x02\x03', public_key=PublicKey.from_params(PublicKeyParamsEddsa( key_parameter=ECParamWellKnown.CURVE25519, key_data=b'\x00\x01\x02\x03', )), serial=0x0102030405060708, certificate_type=SshCertType.SSH_CERT_TYPE_HOST, key_id='\x00\x01\x02\x03\x04\x05\x06\x07', valid_principals=SshCertValidPrincipals([]), valid_after=datetime.datetime.fromtimestamp(0, datetime.timezone.utc), valid_before=None, critical_options=SshCertCriticalOptionVector([]), extensions=SshCertExtensionVector([]), reserved=b'', signature_key=SshHostKeyEDDSA( host_key_algorithm=SshHostKeyAlgorithm.SSH_ED25519, public_key=PublicKey.from_params(PublicKeyParamsEddsa( key_parameter=ECParamWellKnown.CURVE25519, key_data=b'\x00\x01\x02\x03', )), ), signature=SshCertSignature( SshHostKeyAlgorithm.SSH_ED25519, b'\x00\x01\x02\x03', ) ) def test_parse(self): host_cert = SshHostPublicKeyVariant.parse_exact_size(self.host_cert_bytes) self.assertEqual(host_cert.host_key_algorithm, self.host_cert.host_key_algorithm) self.assertEqual(host_cert.nonce, self.host_cert.nonce) self.assertEqual(host_cert.public_key.params.key_data, b'\x00\x01\x02\x03') self.assertEqual(host_cert.serial, self.host_cert.serial) self.assertEqual(host_cert.certificate_type, self.host_cert.certificate_type) self.assertEqual(host_cert.key_id, self.host_cert.key_id) self.assertEqual(host_cert.valid_principals, self.host_cert.valid_principals) self.assertEqual(host_cert.valid_after, self.host_cert.valid_after) self.assertEqual(host_cert.valid_before, self.host_cert.valid_before) self.assertEqual(host_cert.critical_options, self.host_cert.critical_options) self.assertEqual(host_cert.extensions, SshCertExtensionVector([])) self.assertEqual(host_cert.reserved, self.host_cert.reserved) self.assertEqual(host_cert.signature_key, self.host_cert.signature_key) self.assertEqual(host_cert.signature, self.host_cert.signature) def test_compose(self): self.assertEqual(self.host_cert.compose(), self.host_cert_bytes) self.assertEqual(self.host_cert.key_bytes, self.host_cert_bytes) def test_asdict(self): self.assertEqual(self.host_cert._asdict(), OrderedDict([ ('key_type', 'host certificate'), ('algorithm', Authentication.EDDSA), ('size', PublicKeySize(Authentication.EDDSA, 256)), ('fingerprints', OrderedDict([ (Hash.SHA2_256, 'SHA256:IDjjwI5W2lkjfR/gnU0pvSw6E340LushvP/N9A1HrWg='), (Hash.SHA1, 'SHA1:EXpev7tY2XCP2R0G6eqcqNgBpOc='), (Hash.MD5, 'MD5:74:06:06:3f:aa:1f:a8:34:1f:1e:6c:16:26:4c:fd:6e') ])), ('known_hosts', ( 'AAAAIHNzaC1lZDI1NTE5LWNlcnQtdjAxQG9wZW5zc2guY29tAAAABAABAgMAAAAE' 'AAECAwECAwQFBgcIAAAAAgAAAAgAAQIDBAUGBwAAAAAAAAAAAAAAAP//////////' 'AAAAAAAAAAAAAAAAAAAAFwAAAAtzc2gtZWQyNTUxOQAAAAQAAQIDAAAAFwAAAAtz' 'c2gtZWQyNTUxOQAAAAQAAQID' )), ('host_key_algorithm', SshHostKeyAlgorithm.SSH_ED25519_CERT_V01_OPENSSH_COM), ('public_key', PublicKey.from_params(PublicKeyParamsEddsa( key_parameter=ECParamWellKnown.CURVE25519, key_data=b'\x00\x01\x02\x03', ))), ('nonce', b'\x00\x01\x02\x03'), ('serial', 0x0102030405060708), ('certificate_type', SshCertType.SSH_CERT_TYPE_HOST), ('key_id', '\x00\x01\x02\x03\x04\x05\x06\x07'), ('valid_principals', SshCertValidPrincipals([])), ('valid_after', datetime.datetime.fromtimestamp(0, datetime.timezone.utc)), ('valid_before', None), ('critical_options', SshCertCriticalOptionVector([])), ('extensions', SshCertExtensionVector([])), ('reserved', b''), ('signature_key', SshHostKeyEDDSA( host_key_algorithm=SshHostKeyAlgorithm.SSH_ED25519, public_key=PublicKey.from_params(PublicKeyParamsEddsa( key_parameter=ECParamWellKnown.CURVE25519, key_data=b'\x00\x01\x02\x03', )), )), ('signature', SshCertSignature( signature_type=SshHostKeyAlgorithm.SSH_ED25519, signature_data=b'\x00\x01\x02\x03' )) ])) class TestX509CertificateChain(TestClasses.TestKeyBase): def setUp(self): super().setUp() x509_certificate = self._get_public_key_x509('snakeoil_cert.pem') self.x509_certificate_bytes = x509_certificate.der self.x509v3_ssh_rsa_certificate_bytes = bytes( b'\x00\x00\x00\x0e' + # host_key_algorithm b'x509v3-ssh-rsa' + b'\x00\x00\x00\x01' + # certificate_count b'\x00\x00\x03\xc4' + # certificate_length self.x509_certificate_bytes + b'\x00\x00\x00\x01' + # ocsp_response_count b'\x00\x00\x00\x04' + # ocsp_response_length b'\x00\x01\x02\x03' + b'' ) self.x509v3_ssh_rsa_certificate = SshX509CertificateChain( SshHostKeyAlgorithm.X509V3_SSH_RSA, x509_certificate, [], [b'\x00\x01\x02\x03'] ) def test_key_bytes(self): self.assertEqual( self.x509v3_ssh_rsa_certificate.key_bytes, self.x509v3_ssh_rsa_certificate.public_key.key_bytes ) def test_asdict(self): dict_result = self.x509v3_ssh_rsa_certificate._asdict() self.assertEqual(list(dict_result.keys())[-2:], ['key_type', 'certificate_chain']) self.assertEqual(dict_result.pop('key_type'), 'X.509 certificate chain') self.assertEqual(dict_result.pop('certificate_chain'), [self.x509v3_ssh_rsa_certificate.public_key]) def test_parse(self): x509_certificate = SshX509CertificateChain.parse_exact_size(self.x509v3_ssh_rsa_certificate_bytes) self.assertEqual(x509_certificate.host_key_algorithm, SshHostKeyAlgorithm.X509V3_SSH_RSA) self.assertEqual(x509_certificate.public_key, self.x509v3_ssh_rsa_certificate.public_key) self.assertEqual(x509_certificate.issuer_certificates, []) self.assertEqual(x509_certificate.ocsp_responses, [b'\x00\x01\x02\x03']) def test_compose(self): self.assertEqual(self.x509v3_ssh_rsa_certificate.compose(), self.x509v3_ssh_rsa_certificate_bytes) class TestX509Certificate(TestClasses.TestKeyBase): def setUp(self): super().setUp() self.x509v3_sign_rsa_header = bytes( b'\x00\x00\x00\x14' + # host_key_algorithm b'x509v3-sign-rsa-sha1' + b'\x00\x00\x03\xc4' + # public_key_length b'' ) self.x509_certificate = self._get_public_key_x509('snakeoil_cert.pem') self.x509_certificate_bytes = self.x509_certificate.der self.x509v3_sign_rsa_certificate = SshX509Certificate( SshHostKeyAlgorithm.X509V3_SIGN_RSA, self.x509_certificate ) self.x509v3_sign_rsa_sha1_certificate = SshX509Certificate( SshHostKeyAlgorithm.X509V3_SIGN_RSA_SHA1, self.x509_certificate ) def test_error_invalid_type(self): x509_certificate_bytes = self._get_public_key_x509('ecc256.badssl.com.pem').der with self.assertRaises(InvalidType): SshX509Certificate.parse_exact_size(x509_certificate_bytes) def test_error_invalid_certificate_value(self): with self.assertRaises(InvalidValue) as context_manager: SshX509Certificate.parse_exact_size(self.x509v3_sign_rsa_header) self.assertEqual(context_manager.exception.value, b'') def test_parse_with_host_key_type(self): x509_certificate = SshX509Certificate.parse_exact_size( self.x509v3_sign_rsa_header + self.x509_certificate_bytes ) self.assertEqual(x509_certificate.host_key_algorithm, SshHostKeyAlgorithm.X509V3_SIGN_RSA_SHA1) self.assertEqual(x509_certificate.public_key, self.x509v3_sign_rsa_certificate.public_key) def test_parse_without_host_key_type(self): x509_certificate = SshX509Certificate.parse_exact_size(self.x509_certificate_bytes) self.assertEqual(x509_certificate.host_key_algorithm, SshHostKeyAlgorithm.X509V3_SIGN_RSA) self.assertEqual(x509_certificate.public_key, self.x509v3_sign_rsa_certificate.public_key) def test_key_bytes(self): self.assertEqual(self.x509v3_sign_rsa_certificate.key_bytes, self.x509_certificate.public_key.key_bytes) def test_compose_with_host_key_type(self): self.assertEqual( self.x509v3_sign_rsa_sha1_certificate.compose(), self.x509v3_sign_rsa_header + self.x509_certificate_bytes ) def test_compose_without_host_key_type(self): self.assertEqual(self.x509v3_sign_rsa_certificate.compose(), self.x509_certificate_bytes) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ssh/test_record.py000066400000000000000000000055301524413560000265200ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptoparser.common.exception import NotEnoughData from cryptoparser.ssh.record import SshRecordInit, SshRecordKexDH, SshRecordKexDHGroup from cryptoparser.ssh.subprotocol import SshDisconnectMessage, SshReasonCode class TestRecord(unittest.TestCase): def setUp(self): self.test_packet = SshDisconnectMessage( SshReasonCode.PROTOCOL_ERROR, 'αβγ', 'en-US' ) self.test_record = SshRecordInit(self.test_packet) self.test_record_bytes = bytes( b'\x00\x00\x00\x24' + # length = 0x01020304 b'\x0b' + # padding length = 0x00 b'\x01' + # message code = DISCONNECT b'\x00\x00\x00\x02' + # reason = PROTOCOL_ERROR b'\x00\x00\x00\x06' + # description length = 6 'αβγ'.encode() + # description b'\x00\x00\x00\x05' + # language length = 5 b'en-US' + # language b'\x00\x00\x00\x00\x00\x00\x00\x00' + # padding b'\x00\x00\x00' + # padding b'' ) def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: SshRecordInit.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, 4) with self.assertRaises(NotEnoughData) as context_manager: SshRecordInit.parse_exact_size( b'\x01\x02\x03\x04' # length = 0x01020304 ) self.assertEqual(context_manager.exception.bytes_needed, 0x01020304) def test_parse(self): record = SshRecordInit.parse_exact_size(self.test_record_bytes) self.assertEqual(record.packet.reason, SshReasonCode.PROTOCOL_ERROR) self.assertEqual(record.packet.description, 'αβγ') self.assertEqual(record.packet.language, 'en-US') record = SshRecordKexDH.parse_exact_size(self.test_record_bytes) self.assertEqual(record.packet.reason, SshReasonCode.PROTOCOL_ERROR) self.assertEqual(record.packet.description, 'αβγ') self.assertEqual(record.packet.language, 'en-US') record = SshRecordKexDHGroup.parse_exact_size(self.test_record_bytes) self.assertEqual(record.packet.reason, SshReasonCode.PROTOCOL_ERROR) self.assertEqual(record.packet.description, 'αβγ') self.assertEqual(record.packet.language, 'en-US') def test_compose(self): self.assertEqual( self.test_record.compose(), self.test_record_bytes ) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ssh/test_subprotocol.py000066400000000000000000000447601524413560000276250ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import collections import unittest from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.key import PublicKey, PublicKeyParamsRsa from cryptodatahub.ssh.algorithm import ( SshCompressionAlgorithm, SshEncryptionAlgorithm, SshHostKeyAlgorithm, SshKexAlgorithm, SshMacAlgorithm, ) from cryptoparser.common.classes import LanguageTag from cryptoparser.common.exception import TooMuchData from cryptoparser.ssh.key import SshHostKeyRSA from cryptoparser.ssh.subprotocol import ( SshDHGroupExchangeInit, SshDHGroupExchangeGroup, SshDHGroupExchangeReply, SshDHGroupExchangeRequest, SshDHKeyExchangeInit, SshDHKeyExchangeReply, SshKeyExchangeInit, SshMessageVariantInit, SshMessageVariantKexDH, SshMessageVariantKexDHGroup, SshNewKeys, SshProtocolMessage, SshUnimplementedMessage, ) from cryptoparser.ssh.version import SshProtocolVersion, SshSoftwareVersionUnparsed, SshVersion class TestProtocolMessage(unittest.TestCase): def test_error(self): with self.assertRaises(InvalidValue) as context_manager: SshProtocolMessage.parse_exact_size(b'ABC') self.assertEqual(context_manager.exception.value, 'ABC') with self.assertRaises(InvalidValue) as context_manager: SshProtocolMessage.parse_exact_size(b'SSH-2.0\r\n') self.assertEqual(context_manager.exception.value, b'\r') with self.assertRaises(InvalidValue) as context_manager: SshProtocolMessage.parse_exact_size(b'SSH-2.0-software_version\r') self.assertEqual(context_manager.exception.value, b'software_version\r') with self.assertRaises(TooMuchData) as context_manager: SshProtocolMessage.parse_exact_size(b'SSH-2.0-software_version ' + b'X' * 255 + b'\r\n') self.assertEqual(context_manager.exception.bytes_needed, len(b'SSH-2.0-software_version ') + len(b'\r\n')) def test_parse(self): message = SshProtocolMessage.parse_exact_size(b'SSH-1.1-software_version\r\n') self.assertEqual(message.protocol_version, SshProtocolVersion(SshVersion.SSH1, 1)) self.assertEqual(message.software_version, SshSoftwareVersionUnparsed('software_version')) self.assertEqual(message.comment, None) message = SshProtocolMessage.parse_exact_size(b'SSH-2.2-software_version\r\n') self.assertEqual(message.protocol_version, SshProtocolVersion(SshVersion.SSH2, 2)) self.assertEqual(message.software_version, SshSoftwareVersionUnparsed('software_version')) self.assertEqual(message.comment, None) message = SshProtocolMessage.parse_exact_size(b'SSH-2.0-software_version comment\r\n') self.assertEqual(message.protocol_version, SshProtocolVersion(SshVersion.SSH2, 0)) self.assertEqual(message.software_version, SshSoftwareVersionUnparsed('software_version')) self.assertEqual(message.comment, 'comment') message = SshProtocolMessage.parse_exact_size(b'SSH-2.0-software_version comment with spaces\r\n') self.assertEqual(message.protocol_version, SshProtocolVersion(SshVersion.SSH2, 0)) self.assertEqual(message.software_version, SshSoftwareVersionUnparsed('software_version')) self.assertEqual(message.comment, 'comment with spaces') def test_comment(self): with self.assertRaises(InvalidValue): SshProtocolMessage( SshProtocolVersion(SshVersion.SSH2, 2), SshSoftwareVersionUnparsed('software_version'), comment='αβγ', ) with self.assertRaises(InvalidValue): SshProtocolMessage( SshProtocolVersion(SshVersion.SSH2, 2), SshSoftwareVersionUnparsed('software_version'), comment='comment\r', ) with self.assertRaises(InvalidValue): SshProtocolMessage( SshProtocolVersion(SshVersion.SSH2, 2), SshSoftwareVersionUnparsed('software_version'), comment='comment\n', ) def test_compose(self): self.assertEqual( SshProtocolMessage( SshProtocolVersion(SshVersion.SSH2, 2), SshSoftwareVersionUnparsed('software_version'), ).compose(), b'SSH-2.2-software_version\r\n' ) self.assertEqual( SshProtocolMessage( SshProtocolVersion(SshVersion.SSH2, 2), SshSoftwareVersionUnparsed('software_version'), 'comment' ).compose(), b'SSH-2.2-software_version comment\r\n' ) self.assertEqual( SshProtocolMessage( SshProtocolVersion(SshVersion.SSH2, 2), SshSoftwareVersionUnparsed('software_version'), 'comment with spaces' ).compose(), b'SSH-2.2-software_version comment with spaces\r\n' ) class TestKeyExchangeInitMessage(unittest.TestCase): def setUp(self): self.key_exchange_init_bytes = bytes( b'\x14' + # message_code = SshMessageCode.KEXINIT b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + # cookie b'\x00\x00\x00\x3d' + b'diffie-hellman-group1-sha1,ecdh-sha2-nistp256,unparsable-algo' + # kex_algorithms b'\x00\x00\x00\x2f' + b'ssh-ed25519,ecdsa-sha2-nistp256,unparsable-algo' + # host_key_algorithms b'\x00\x00\x00\x31' + b'aes128-cbc,aes256-gcm@openssh.com,unparsable-algo' + # encryption_algorithms_client_to_server b'\x00\x00\x00\x31' + b'aes256-gcm@openssh.com,aes128-cbc,unparsable-algo' + # encryption_algorithms_server_to_client b'\x00\x00\x00\x2e' + b'hmac-sha1,umac-128@openssh.com,unparsable-algo' + # mac_algorithms_client_to_server b'\x00\x00\x00\x2e' + b'umac-128@openssh.com,hmac-sha1,unparsable-algo' + # mac_algorithms_server_to_client b'\x00\x00\x00\x25' + b'none,zlib@openssh.com,unparsable-algo' + # compression_algorithms_client_to_server b'\x00\x00\x00\x25' + b'zlib@openssh.com,none,unparsable-algo' + # compression_algorithms_server_to_client b'\x00\x00\x00\x0b' + b'en-UK,en-US' + # languages_client_to_server b'\x00\x00\x00\x0b' + b'en-US,en-UK' + # languages_server_to_client b'\x00' + # first_kex_packet_follows b'\x00\x01\x02\x03' + # reserved b'' ) self.key_exchange_init = SshKeyExchangeInit( kex_algorithms=[ SshKexAlgorithm.DIFFIE_HELLMAN_GROUP1_SHA1, SshKexAlgorithm.ECDH_SHA2_NISTP256, 'unparsable-algo', ], host_key_algorithms=[ SshHostKeyAlgorithm.SSH_ED25519, SshHostKeyAlgorithm.ECDSA_SHA2_NISTP256, 'unparsable-algo', ], encryption_algorithms_client_to_server=[ SshEncryptionAlgorithm.AES128_CBC, SshEncryptionAlgorithm.AES256_GCM_OPENSSH_COM, 'unparsable-algo', ], encryption_algorithms_server_to_client=[ SshEncryptionAlgorithm.AES256_GCM_OPENSSH_COM, SshEncryptionAlgorithm.AES128_CBC, 'unparsable-algo', ], mac_algorithms_client_to_server=[ SshMacAlgorithm.HMAC_SHA1, SshMacAlgorithm.UMAC_128_OPENSSH_COM, 'unparsable-algo', ], mac_algorithms_server_to_client=[ SshMacAlgorithm.UMAC_128_OPENSSH_COM, SshMacAlgorithm.HMAC_SHA1, 'unparsable-algo', ], compression_algorithms_client_to_server=[ SshCompressionAlgorithm.NONE, SshCompressionAlgorithm.ZLIB_OPENSSH_COM, 'unparsable-algo', ], compression_algorithms_server_to_client=[ SshCompressionAlgorithm.ZLIB_OPENSSH_COM, SshCompressionAlgorithm.NONE, 'unparsable-algo', ], languages_client_to_server=[ LanguageTag('en', ['UK', ]), LanguageTag('en', ['US', ]), ], languages_server_to_client=[ LanguageTag('en', ['US', ]), LanguageTag('en', ['UK', ]), ], cookie=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', reserved=0x00010203, ) def test_parse(self): SshMessageVariantInit.parse_exact_size(self.key_exchange_init_bytes) def test_compose(self): self.assertEqual(self.key_exchange_init.compose(), self.key_exchange_init_bytes) def test_hassh(self): self.assertEqual(self.key_exchange_init.hassh, 'cc40dc455f685d8f57f8794262e30422') self.assertEqual(self.key_exchange_init.hassh_server, '55a954fe89f7f04218e3013996beee76') class TestUnimplementedMessage(unittest.TestCase): def setUp(self): self.unimplemented_bytes = bytes( b'\x03' + # message_code = SshMessageCode.UNIMPLEMENTED b'\x01\x02\x03\x04' + # sequence_number b'' ) self.unimplemented = SshUnimplementedMessage( sequence_number=0x01020304 ) def test_parse(self): message = SshMessageVariantInit.parse_exact_size(self.unimplemented_bytes) self.assertEqual(message.sequence_number, 0x01020304) def test_compose(self): self.assertEqual(self.unimplemented.compose(), self.unimplemented_bytes) class TestDHKeyExchangeInit(unittest.TestCase): def setUp(self): self.dh_key_exchange_init_bytes = bytes( b'\x1e' + # message_code = SshMessageCode.DH_KEX_INIT b'\x00\x00\x00\x10' + # ephemeral_public_key length b'\x00\x01\x02\x03\x04\x05\x06\x07' + # ephemeral_public_key length b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'' ) self.dh_key_exchange_init = SshDHKeyExchangeInit( ephemeral_public_key=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' ) def test_parse(self): message = SshMessageVariantKexDH.parse_exact_size(self.dh_key_exchange_init_bytes) self.assertEqual(message.ephemeral_public_key, self.dh_key_exchange_init.ephemeral_public_key) def test_compose(self): self.assertEqual(self.dh_key_exchange_init.compose(), self.dh_key_exchange_init_bytes) class TestDHGroupExchangeInit(unittest.TestCase): def setUp(self): self.dh_group_exchange_init_bytes = bytes( b'\x20' + # message_code = SshMessageCode.DH_KEX_INIT b'\x00\x00\x00\x10' + # ephemeral_public_key length b'\x00\x01\x02\x03\x04\x05\x06\x07' + # ephemeral_public_key length b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'' ) self.dh_group_exchange_init = SshDHGroupExchangeInit( ephemeral_public_key=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' ) def test_parse(self): message = SshMessageVariantKexDHGroup.parse_exact_size(self.dh_group_exchange_init_bytes) self.assertEqual(message, self.dh_group_exchange_init) def test_compose(self): self.assertEqual(self.dh_group_exchange_init.compose(), self.dh_group_exchange_init_bytes) class TestDHKeyExchangeReply(unittest.TestCase): def setUp(self): self.dh_key_exchange_reply_dict = collections.OrderedDict([ ('message_code', b'\x1f'), # DH_KEX_REPLY ('host_key_length', b'\x00\x00\x00\x23'), ('host_public_key', ( b'\x00\x00\x00\x07' + b'ssh-rsa' + b'\x00\x00\x00\x08' + b'\x01\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x00\x00\x08' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'' )), ('ephemeral_public_key_length', b'\x00\x00\x00\x10'), ('ephemeral_public_key', ( b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'' )), ('signature_length', b'\x00\x00\x00\x10'), ('signature_key', ( b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'' )), ]) self.dh_key_exchange_reply_bytes = b''.join(self.dh_key_exchange_reply_dict.values()) self.dh_key_exchange_reply = SshDHKeyExchangeReply( host_public_key=SshHostKeyRSA( host_key_algorithm=SshHostKeyAlgorithm.SSH_RSA, public_key=PublicKey.from_params(PublicKeyParamsRsa( modulus=0x08090a0b0c0d0e0f, public_exponent=0x0101020304050607, )), ), ephemeral_public_key=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', signature=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', ) def test_parse(self): message = SshMessageVariantKexDH.parse_exact_size(self.dh_key_exchange_reply_bytes) self.assertEqual(message, self.dh_key_exchange_reply) def test_compose(self): self.assertEqual(self.dh_key_exchange_reply.compose(), self.dh_key_exchange_reply_bytes) class TestDHGroupExchangeReply(unittest.TestCase): def setUp(self): self.dh_group_exchange_reply_dict = collections.OrderedDict([ ('message_code', b'\x21'), # DH_GEX_REPLY ('host_key_length', b'\x00\x00\x00\x23'), ('host_public_key', ( b'\x00\x00\x00\x07' + b'ssh-rsa' + b'\x00\x00\x00\x08' + b'\x01\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x00\x00\x08' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'' )), ('ephemeral_public_key_length', b'\x00\x00\x00\x10'), ('ephemeral_public_keylic_key', ( b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'' )), ('signature_length', b'\x00\x00\x00\x10'), ('signature_key', ( b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'' )), ]) self.dh_group_exchange_reply_bytes = b''.join(self.dh_group_exchange_reply_dict.values()) self.dh_group_exchange_reply = SshDHGroupExchangeReply( host_public_key=SshHostKeyRSA( host_key_algorithm=SshHostKeyAlgorithm.SSH_RSA, public_key=PublicKey.from_params(PublicKeyParamsRsa( modulus=0x08090a0b0c0d0e0f, public_exponent=0x0101020304050607, )), ), ephemeral_public_key=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', signature=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', ) def test_parse(self): message = SshMessageVariantKexDHGroup.parse_exact_size(self.dh_group_exchange_reply_bytes) self.assertEqual(message, self.dh_group_exchange_reply) def test_compose(self): self.assertEqual(self.dh_group_exchange_reply.compose(), self.dh_group_exchange_reply_bytes) class TestDHGroupExchangeRequest(unittest.TestCase): def setUp(self): self.dh_group_exchange_reply_dict = collections.OrderedDict([ ('message_code', b'\x22'), # DH_GEX_REQUEST ('gex_min', b'\x00\x00\x04\x00'), ('gex_number', b'\x00\x00\x08\x00'), ('gex_max', b'\x00\x00\x10\x00'), ]) self.dh_group_exchange_reply_bytes = b''.join(self.dh_group_exchange_reply_dict.values()) self.dh_group_exchange_reply = SshDHGroupExchangeRequest( gex_min=1024, gex_number=2048, gex_max=4096, ) def test_parse(self): message = SshMessageVariantKexDHGroup.parse_exact_size(self.dh_group_exchange_reply_bytes) self.assertEqual(message, self.dh_group_exchange_reply) def test_compose(self): self.assertEqual(self.dh_group_exchange_reply.compose(), self.dh_group_exchange_reply_bytes) class TestDHGroupExchangeGroup(unittest.TestCase): def setUp(self): self.dh_group_exchange_group_dict = collections.OrderedDict([ ('message_code', b'\x1f'), # DH_GEX_GROUP ('p_length', b'\x00\x00\x00\x10'), ('p', ( b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'' )), ('g_length', b'\x00\x00\x00\x10'), ('g', ( b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'' )), ]) self.dh_group_exchange_group_bytes = b''.join(self.dh_group_exchange_group_dict.values()) self.dh_group_exchange_group = SshDHGroupExchangeGroup( p=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', g=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', ) def test_parse(self): message = SshMessageVariantKexDHGroup.parse_exact_size(self.dh_group_exchange_group_bytes) self.assertEqual(message, self.dh_group_exchange_group) def test_compose(self): self.assertEqual(self.dh_group_exchange_group.compose(), self.dh_group_exchange_group_bytes) class TestNewKeys(unittest.TestCase): def setUp(self): self.new_keys_dict = collections.OrderedDict([ ('message_code', b'\x15'), # NEWKEYS ]) self.new_keys_bytes = b''.join(self.new_keys_dict.values()) self.new_keys = SshNewKeys() def test_parse(self): message = SshNewKeys.parse_exact_size(self.new_keys_bytes) self.assertEqual(message, self.new_keys) def test_compose(self): self.assertEqual(self.new_keys.compose(), self.new_keys_bytes) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/ssh/test_version.py000066400000000000000000000173321524413560000267320ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import collections import unittest from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.grade import Grade from cryptoparser.common.parse import ParserText from cryptoparser.ssh.version import ( SshVersion, SshProtocolVersion, SshSoftwareVersionCryptlib, SshSoftwareVersionDropbear, SshSoftwareVersionIPSSH, SshSoftwareVersionMonacaSSH, SshSoftwareVersionOpenSSH, SshSoftwareVersionUnparsed, SshSoftwareVersionParsedVariant ) class TestSshVersion(unittest.TestCase): def test_error(self): parsable = b'3.0' expected_error_message = '3 is not a valid SshVersion' with self.assertRaisesRegex(ValueError, expected_error_message): # pylint: disable=expression-not-assigned SshProtocolVersion.parse_exact_size(parsable) expected_error_message = '\'.0\' is not a valid SshProtocolVersion' with self.assertRaisesRegex(InvalidValue, expected_error_message): # pylint: disable=expression-not-assigned SshProtocolVersion.parse_exact_size(b'.0') expected_error_message = '\'2.\' is not a valid SshProtocolVersion' with self.assertRaisesRegex(InvalidValue, expected_error_message): # pylint: disable=expression-not-assigned SshProtocolVersion.parse_exact_size(b'2.') def test_parse(self): version = SshProtocolVersion.parse_exact_size(b'1.0') self.assertEqual(version, SshProtocolVersion(SshVersion.SSH1)) self.assertEqual(version.supported_versions, [SshVersion.SSH1, ]) version = SshProtocolVersion.parse_exact_size(b'1.99') self.assertEqual(version, SshProtocolVersion(SshVersion.SSH1, 99)) self.assertEqual(version.supported_versions, [SshVersion.SSH1, SshVersion.SSH2]) version = SshProtocolVersion.parse_exact_size(b'2.0') self.assertEqual(version, SshProtocolVersion(SshVersion.SSH2)) self.assertEqual(version.supported_versions, [SshVersion.SSH2, ]) def test_compose(self): self.assertEqual(b'2.0', SshProtocolVersion(SshVersion.SSH2, 0).compose()) self.assertEqual(b'1.1', SshProtocolVersion(SshVersion.SSH1, 1).compose()) def test_lt(self): self.assertLess( SshProtocolVersion(SshVersion.SSH1), SshProtocolVersion(SshVersion.SSH2) ) self.assertLess( SshProtocolVersion(SshVersion.SSH2, 0), SshProtocolVersion(SshVersion.SSH2, 1) ) self.assertLess( SshProtocolVersion(SshVersion.SSH1, 1), SshProtocolVersion(SshVersion.SSH2, 0) ) self.assertGreater( SshProtocolVersion(SshVersion.SSH2, 0), SshProtocolVersion(SshVersion.SSH1, 1) ) def test_eq(self): self.assertEqual( SshProtocolVersion(SshVersion.SSH1, 0), SshProtocolVersion(SshVersion.SSH1, 0) ) self.assertEqual( SshProtocolVersion(SshVersion.SSH2, 0), SshProtocolVersion(SshVersion.SSH2, 0) ) def test_as_json(self): self.assertEqual(SshProtocolVersion(SshVersion.SSH1, 0).as_json(), '\"ssh1\"') self.assertEqual(SshProtocolVersion(SshVersion.SSH1, 1).as_json(), '\"ssh1\"') self.assertEqual(SshProtocolVersion(SshVersion.SSH2, 0).as_json(), '\"ssh2\"') self.assertEqual(SshProtocolVersion(SshVersion.SSH2, 1).as_json(), '\"ssh2\"') def test_str(self): self.assertEqual(str(SshProtocolVersion(SshVersion.SSH1, 0)), 'SSH 1.0') self.assertEqual(str(SshProtocolVersion(SshVersion.SSH1, 1)), 'SSH 1.1') self.assertEqual(str(SshProtocolVersion(SshVersion.SSH2, 0)), 'SSH 2.0') self.assertEqual(str(SshProtocolVersion(SshVersion.SSH2, 1)), 'SSH 2.1') def test_grade(self): self.assertEqual(SshProtocolVersion(SshVersion.SSH2).grade, Grade.SECURE) self.assertEqual(SshProtocolVersion(SshVersion.SSH1).grade, Grade.INSECURE) self.assertEqual(SshProtocolVersion(SshVersion.SSH1, 99).grade, Grade.INSECURE) class TestSshSoftwareVersion(unittest.TestCase): @staticmethod def _get_software_version(raw): parser = ParserText(raw) parser.parse_parsable('software_version', SshSoftwareVersionParsedVariant) return parser['software_version'] def test_parse(self): software_version = self._get_software_version(b'cryptlib') self.assertEqual(software_version.vendor, 'cryptlib') self.assertEqual(software_version, SshSoftwareVersionCryptlib()) software_version = self._get_software_version(b'dropbear_2020.81') self.assertEqual(software_version.vendor, 'dropbear') self.assertEqual(software_version, SshSoftwareVersionDropbear('2020.81')) software_version = self._get_software_version(b'IPSSH-6.9.0') self.assertEqual(software_version.vendor, 'IPSSH') self.assertEqual(software_version, SshSoftwareVersionIPSSH('6.9.0')) software_version = self._get_software_version(b'Monaca') self.assertEqual(software_version.vendor, 'Monaca') self.assertEqual(software_version, SshSoftwareVersionMonacaSSH()) software_version = self._get_software_version(b'OpenSSH_8.6') self.assertEqual(software_version.vendor, 'OpenSSH') self.assertEqual(software_version, SshSoftwareVersionOpenSSH('8.6')) parser = ParserText(b'unknown.ssh.server-1.2.3') parser.parse_parsable('software_version', SshSoftwareVersionUnparsed) software_version = parser['software_version'] self.assertEqual(software_version.raw, 'unknown.ssh.server-1.2.3') def test_compose(self): software_version = SshSoftwareVersionCryptlib() self.assertEqual(software_version.compose(), b'cryptlib') self.assertEqual( software_version._asdict(), collections.OrderedDict([('vendor', 'cryptlib'), ('version', None)]) ) software_version = SshSoftwareVersionDropbear('2020.81') self.assertEqual(software_version.compose(), b'dropbear_2020.81') self.assertEqual( software_version._asdict(), collections.OrderedDict([('vendor', 'dropbear'), ('version', '2020.81')]) ) software_version = SshSoftwareVersionIPSSH('6.9.0') self.assertEqual(software_version.compose(), b'IPSSH-6.9.0') self.assertEqual( software_version._asdict(), collections.OrderedDict([('vendor', 'IPSSH'), ('version', '6.9.0')]) ) software_version = SshSoftwareVersionMonacaSSH() self.assertEqual(software_version.compose(), b'Monaca') self.assertEqual( software_version._asdict(), collections.OrderedDict([('vendor', 'Monaca'), ('version', None)]) ) software_version = SshSoftwareVersionOpenSSH('8.6') self.assertEqual(software_version.compose(), b'OpenSSH_8.6') self.assertEqual( software_version._asdict(), collections.OrderedDict([('vendor', 'OpenSSH'), ('version', '8.6')]) ) software_version = SshSoftwareVersionUnparsed('unknown.ssh.server-1.2.3') self.assertEqual(software_version.compose(), b'unknown.ssh.server-1.2.3') self.assertEqual(software_version.as_markdown(), 'unknown.ssh.server-1.2.3') def test_error_raw(self): with self.assertRaises(InvalidValue): SshSoftwareVersionUnparsed('αβγ') with self.assertRaises(InvalidValue): SshSoftwareVersionUnparsed('software_version ') with self.assertRaises(InvalidValue): SshSoftwareVersionUnparsed('software_version\r') with self.assertRaises(InvalidValue): SshSoftwareVersionUnparsed('software_version\n') cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/000077500000000000000000000000001524413560000236335ustar00rootroot00000000000000cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/__init__.py000066400000000000000000000000431524413560000257410ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/classes.py000066400000000000000000000014421524413560000256430ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 from cryptoparser.tls.extension import TlsExtensionUnusedData, TlsExtensionType from cryptoparser.tls.subprotocol import TlsSubprotocolMessageBase, TlsHandshakeMessage, TlsHandshakeType class TestMessage(TlsSubprotocolMessageBase): @classmethod def get_handshake_type(cls): raise NotImplementedError class TestVariantMessage(TlsHandshakeMessage): @classmethod def get_handshake_type(cls): return TlsHandshakeType.SERVER_HELLO_DONE @classmethod def _parse(cls, parsable): raise NotImplementedError def compose(self): raise NotImplementedError class TestUnusedDataExtension(TlsExtensionUnusedData): @classmethod def get_extension_type(cls): return TlsExtensionType.RENEGOTIATION_INFO cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_alert.py000066400000000000000000000025441524413560000263600ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import NotEnoughData from cryptoparser.tls.subprotocol import TlsAlertMessage, TlsAlertLevel, TlsAlertDescription class TestAlert(unittest.TestCase): def test_error(self): with self.assertRaisesRegex(InvalidValue, '0xff is not a valid TlsAlertLevel'): # pylint: disable=expression-not-assigned TlsAlertMessage.parse_exact_size(b'\xff\x00') with self.assertRaisesRegex(InvalidValue, '0xff is not a valid TlsAlertDescription'): # pylint: disable=expression-not-assigned TlsAlertMessage.parse_exact_size(b'\x01\xff') with self.assertRaises(NotEnoughData) as context_manager: # pylint: disable=expression-not-assigned TlsAlertMessage.parse_exact_size(b'\xff') self.assertGreaterEqual(context_manager.exception.bytes_needed, 1) def test_parse(self): self.assertEqual( TlsAlertMessage.parse_exact_size(b'\x02\x28'), TlsAlertMessage(TlsAlertLevel.FATAL, TlsAlertDescription.HANDSHAKE_FAILURE) ) def test_compose(self): self.assertEqual( b'\x02\x28', TlsAlertMessage(TlsAlertLevel.FATAL, TlsAlertDescription.HANDSHAKE_FAILURE).compose() ) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_application_data.py000066400000000000000000000012751524413560000305450ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptoparser.tls.subprotocol import TlsApplicationDataMessage class TestRecord(unittest.TestCase): _APPLICATION_DATA_MESSAGE_BYTES = b'\x01\x02\x03\x04' def test_error(self): pass def test_parse(self): self.assertEqual( TlsApplicationDataMessage.parse_exact_size(self._APPLICATION_DATA_MESSAGE_BYTES), TlsApplicationDataMessage(data=self._APPLICATION_DATA_MESSAGE_BYTES) ) def test_compose(self): self.assertEqual( self._APPLICATION_DATA_MESSAGE_BYTES, TlsApplicationDataMessage(data=self._APPLICATION_DATA_MESSAGE_BYTES).compose() ) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_change_cipher_spec.py000066400000000000000000000031341524413560000310360ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import NotEnoughData from cryptoparser.tls.record import TlsRecord from cryptoparser.tls.subprotocol import TlsChangeCipherSpecMessage, TlsChangeCipherSpecType, TlsContentType from cryptoparser.tls.version import TlsProtocolVersion, TlsVersion class TestRecord(unittest.TestCase): def test_error(self): with self.assertRaisesRegex(InvalidValue, '0xff is not a valid TlsChangeCipherSpecType'): # pylint: disable=expression-not-assigned TlsChangeCipherSpecMessage.parse_exact_size(b'\xff') with self.assertRaises(NotEnoughData) as context_manager: # pylint: disable=expression-not-assigned TlsChangeCipherSpecMessage.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, 1) def test_parse(self): self.assertEqual( TlsChangeCipherSpecMessage.parse_exact_size(b'\x01'), TlsChangeCipherSpecMessage(TlsChangeCipherSpecType.CHANGE_CIPHER_SPEC) ) def test_compose(self): self.assertEqual( b'\x01', TlsChangeCipherSpecMessage(TlsChangeCipherSpecType.CHANGE_CIPHER_SPEC).compose() ) def test_record(self): self.assertEqual( b'\x14\x03\x03\x00\x01\x01', TlsRecord( TlsChangeCipherSpecMessage().compose(), TlsProtocolVersion(TlsVersion.TLS1_2), TlsContentType.CHANGE_CIPHER_SPEC, ).compose() ) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_ciphersuite.py000066400000000000000000000027341524413560000275760ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptoparser.tls.ciphersuite import TlsCipherSuite from cryptoparser.tls.version import TlsProtocolVersion, TlsVersion class TestTlsCipherSuite(unittest.TestCase): def test_str(self): for cipher_suite in filter(lambda cipher_suite: cipher_suite.value.iana_name, TlsCipherSuite): self.assertIn(cipher_suite.value.iana_name, str(cipher_suite.value)) for cipher_suite in filter(lambda cipher_suite: cipher_suite.value.iana_name is None, TlsCipherSuite): self.assertIn(cipher_suite.name, str(cipher_suite.value)) for cipher_suite in filter(lambda cipher_suite: cipher_suite.value.openssl_name, TlsCipherSuite): self.assertIn(cipher_suite.value.openssl_name, str(cipher_suite.value)) def test_initial_version(self): self.assertEqual( TlsProtocolVersion(TlsCipherSuite.TLS_AES_128_GCM_SHA256.value.initial_version), TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_15) ) self.assertEqual( TlsCipherSuite.TLS_AES_128_GCM_SHA256.value.initial_version, TlsVersion.TLS1_3_DRAFT_15 ) self.assertEqual( TlsProtocolVersion(TlsCipherSuite.TLS_RSA_WITH_NULL_SHA256.value.initial_version), TlsProtocolVersion(TlsVersion.TLS1_2) ) self.assertEqual( TlsCipherSuite.TLS_RSA_WITH_NULL_SHA256.value.initial_version, TlsVersion.TLS1_2 ) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_extension.py000066400000000000000000001361311524413560000272650ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import collections import datetime import unittest from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.types import Base64Data from cryptodatahub.common.stores import CertificateTransparencyLog from cryptodatahub.tls.algorithm import ( TlsECPointFormat, TlsCertificateCompressionAlgorithm, TlsNamedCurve, TlsNextProtocolName, TlsProtocolName, TlsPskKeyExchangeMode, TlsSignatureAndHashAlgorithm, TlsTokenBindingParamater, ) from cryptoparser.common.exception import NotEnoughData, InvalidType from cryptoparser.common.x509 import CtExtensions, CtVersion, SignedCertificateTimestamp from cryptoparser.tls.extension import ( TlsCertificateStatusRequestExtensions, TlsCertificateStatusRequestResponderId, TlsCertificateStatusRequestResponderIdList, TlsExtensionApplicationLayerProtocolNegotiation, TlsExtensionApplicationLayerProtocolSettings, TlsExtensionCertificateStatusRequestClient, TlsExtensionCertificateStatusRequestServer, TlsExtensionChannelId, TlsExtensionCompressCertificate, TlsExtensionECPointFormats, TlsExtensionEllipticCurves, TlsExtensionEncryptedClientHelloInner, TlsExtensionEncryptedClientHelloOuter, TlsExtensionEncryptThenMAC, TlsExtensionExtendedMasterSecret, TlsExtensionKeyShareClient, TlsExtensionKeyShareClientHelloRetry, TlsExtensionKeyShareServer, TlsExtensionKeyShareReservedClient, TlsExtensionNextProtocolNegotiationClient, TlsExtensionNextProtocolNegotiationServer, TlsExtensionOldApplicationLayerProtocolSettings, TlsExtensionPadding, TlsExtensionPostHandshakeAuthentication, TlsExtensionPskKeyExchangeModes, TlsExtensionRecordSizeLimit, TlsExtensionRenegotiationInfo, TlsExtensionServerNameClient, TlsExtensionServerNameServer, TlsExtensionServerPadding, TlsExtensionSessionTicket, TlsExtensionShortRecordHeader, TlsExtensionSignatureAlgorithms, TlsExtensionSignatureAlgorithmsCert, TlsExtensionSignedCertificateTimestampClient, TlsExtensionSignedCertificateTimestampServer, TlsExtensionSupportedVersionsClient, TlsExtensionSupportedVersionsServer, TlsExtensionTokenBinding, TlsExtensionTrustAnchors, TlsExtensionUnparsed, TlsExtensionParsed, TlsExtensionType, TlsNextProtocolNameList, TlsProtocolNameList, TlsRenegotiatedConnection, TlsTokenBindingProtocolVersion, TlsTrustAnchorIdentifier, TlsTrustAnchorIdentifierList, ) from cryptoparser.tls.grease import TlsGreaseOneByte, TlsGreaseTwoByte, TlsInvalidTypeOneByte, TlsInvalidTypeTwoByte from cryptoparser.tls.version import TlsVersion, TlsProtocolVersion from .classes import TestUnusedDataExtension class TestExtensionUnparsed(unittest.TestCase): def test_error(self): extension_missing_data_dict = collections.OrderedDict([ ('extension_type', b'\xff\x01'), ('extension_length', b'\x00\x01'), ]) extension_missing_data_bytes = b''.join(extension_missing_data_dict.values()) with self.assertRaises(NotEnoughData) as context_manager: # pylint: disable=expression-not-assigned TlsExtensionUnparsed.parse_exact_size(extension_missing_data_bytes) self.assertEqual(context_manager.exception.bytes_needed, 5) def test_parse_and_compose(self): extension_minimal_dict = collections.OrderedDict([ ('extension_type', b'\xff\x01'), ('extension_length', b'\x00\x00'), ]) extension_minimal_bytes = b''.join(extension_minimal_dict.values()) extension_minimal = TlsExtensionUnparsed.parse_exact_size(extension_minimal_bytes) self.assertEqual(extension_minimal.compose(), extension_minimal_bytes) extension_with_data_dict = collections.OrderedDict([ ('extension_type', b'\xff\x01'), ('extension_length', b'\x00\x04'), ('extension_data', b'\xde\xad\xbe\xaf'), ]) extension_with_data_bytes = b''.join(extension_with_data_dict.values()) extension_with_data = TlsExtensionUnparsed.parse_exact_size(extension_with_data_bytes) self.assertEqual(extension_with_data.compose(), extension_with_data_bytes) class ExtensionInvalidType(TlsExtensionParsed): @classmethod def get_extension_type(cls): return 0xffff @classmethod def _parse(cls, parsable): parser = super()._parse_header(parsable) return ExtensionInvalidType(), parser.parsed_length def compose(self): raise NotImplementedError class TestExtensionParsed(unittest.TestCase): def test_error(self): extension_invalid_type_dict = collections.OrderedDict([ ('extension_type', b'\x00\x00'), ('extension_length', b'\x00\x00'), ]) extension_invalid_type_bytes = b''.join(extension_invalid_type_dict.values()) with self.assertRaises(InvalidType): # pylint: disable=expression-not-assigned ExtensionInvalidType.parse_exact_size(extension_invalid_type_bytes) class TestExtensionHostnameClient(unittest.TestCase): def test_parse(self): extension_hostname_dict = collections.OrderedDict([ ('extension_type', b'\x00\x00'), ('extension_length', b'\x00\x14'), ('name_list_length', b'\x00\x12'), ('name_type', b'\x00'), ('name_length', b'\x00\x0f'), ('name', b'\x77\x77\x77\x2e\x65\x78\x61\x6d\x70\x6c\x65\x2e\x63\x6f\x6d') ]) extension_hostname_bytes = b''.join(extension_hostname_dict.values()) extension_hostname = TlsExtensionServerNameClient.parse_exact_size(extension_hostname_bytes) self.assertEqual(extension_hostname.host_name, 'www.example.com') self.assertEqual(extension_hostname.compose(), extension_hostname_bytes) extension_hostname_internationalized_dict = collections.OrderedDict([ ('extension_type', b'\x00\x00'), ('extension_length', b'\x00\x1e'), ('name_list_length', b'\x00\x1c'), ('name_type', b'\x00'), ('name_length', b'\x00\x19'), ( 'name', b'\x78\x6e\x2d\x2d\x73\x6c\x61\x6e\x64\x2d\x79\x73\x61\x2e\x69\x63\x6f\x6d\x2e\x6d\x75\x73\x65\x75\x6d' ) ]) extension_hostname_internationalized_bytes = b''.join(extension_hostname_internationalized_dict.values()) extension_hostname_internationalized = TlsExtensionServerNameClient.parse_exact_size( extension_hostname_internationalized_bytes ) self.assertEqual(extension_hostname_internationalized.host_name, 'ísland.icom.museum') self.assertEqual(extension_hostname_internationalized.compose(), extension_hostname_internationalized_bytes) class TestExtensionECPointFormat(unittest.TestCase): def test_parse(self): extension_ec_point_formats_dict = collections.OrderedDict([ ('extension_type', b'\x00\x0b'), ('extension_length', b'\x00\x03'), ('ec_point_format_list_length', b'\x02'), ('ec_point_format_list', b'\x00\x0b'), ]) extension_ec_point_formats_bytes = b''.join(extension_ec_point_formats_dict.values()) extension_point_formats = TlsExtensionECPointFormats.parse_exact_size(extension_ec_point_formats_bytes) self.assertEqual( list(extension_point_formats.point_formats), [ TlsECPointFormat.UNCOMPRESSED, TlsInvalidTypeOneByte(TlsGreaseOneByte.GREASE_0B), ] ) self.assertEqual(extension_point_formats.compose(), extension_ec_point_formats_bytes) class TestExtensionHostnameServer(unittest.TestCase): def test_parse(self): extension_hostname_dict = collections.OrderedDict([ ('extension_type', b'\x00\x00'), ('extension_length', b'\x00\x00'), ]) extension_hostname_bytes = b''.join(extension_hostname_dict.values()) extension_hostname = TlsExtensionServerNameServer.parse_exact_size(extension_hostname_bytes) self.assertEqual(extension_hostname.compose(), extension_hostname_bytes) class TestExtensionEllipticCurves(unittest.TestCase): def test_parse(self): extension_elliptic_curves_dict = collections.OrderedDict([ ('extension_type', b'\x00\x0a'), ('extension_length', b'\x00\x0a'), ('elliptic_curve_list_length', b'\x00\x08'), ('elliptic_curve_list', b'\x00\x1d\x00\x17\x00\x18\x0a\x0a'), ]) extension_elliptic_curves_bytes = b''.join(extension_elliptic_curves_dict.values()) extension_elliptic_curves = TlsExtensionEllipticCurves.parse_exact_size(extension_elliptic_curves_bytes) self.assertEqual( list(extension_elliptic_curves.elliptic_curves), [ TlsNamedCurve.X25519, TlsNamedCurve.SECP256R1, TlsNamedCurve.SECP384R1, TlsInvalidTypeTwoByte(TlsGreaseTwoByte.GREASE_0A0A), ] ) self.assertEqual(extension_elliptic_curves.compose(), extension_elliptic_curves_bytes) class TestExtensionSupportedVersions(unittest.TestCase): def test_parse(self): extension_supported_versions_dict = collections.OrderedDict([ ('extension_type', b'\x00\x2b'), ('extension_length', b'\x00\x09'), ('supported_version_list_length', b'\x08'), ('supported_version_list', b'\x03\x02\x03\x03\x7f\x18\x0a\x0a'), ]) extension_supported_versions_bytes = b''.join(extension_supported_versions_dict.values()) extension_supported_versions = TlsExtensionSupportedVersionsClient.parse_exact_size( extension_supported_versions_bytes ) self.assertEqual( list(extension_supported_versions.supported_versions), [ TlsProtocolVersion(TlsVersion.TLS1_1), TlsProtocolVersion(TlsVersion.TLS1_2), TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_24), TlsInvalidTypeTwoByte(TlsGreaseTwoByte.GREASE_0A0A.value.code), ] ) self.assertEqual(extension_supported_versions.compose(), extension_supported_versions_bytes) extension_supported_versions_dict = collections.OrderedDict([ ('extension_type', b'\x00\x2b'), ('extension_length', b'\x00\x02'), ('selected_version', b'\x03\x03'), ]) extension_supported_versions_bytes = b''.join(extension_supported_versions_dict.values()) extension_supported_versions = TlsExtensionSupportedVersionsServer.parse_exact_size( extension_supported_versions_bytes ) self.assertEqual( extension_supported_versions.selected_version, TlsProtocolVersion(TlsVersion.TLS1_2) ) self.assertEqual(extension_supported_versions.compose(), extension_supported_versions_bytes) class TestExtensionTokenBinding(unittest.TestCase): def test_parse(self): extension_token_binding_dict = collections.OrderedDict([ ('extension_type', b'\x00\x18'), ('extension_length', b'\x00\x06'), ('protocol_version', b'\x01\x02'), ('supported_version_list', b'\x03\x02\x01\x00'), ]) extension_token_binding_bytes = b''.join(extension_token_binding_dict.values()) extension_token_binding = TlsExtensionTokenBinding.parse_exact_size( extension_token_binding_bytes ) self.assertEqual( list(extension_token_binding.parameters), [ TlsTokenBindingParamater.ECDSAP256, TlsTokenBindingParamater.RSA2048_PSS, TlsTokenBindingParamater.RSA2048_PKCS1_5, ] ) self.assertEqual( extension_token_binding.protocol_version, TlsTokenBindingProtocolVersion(1, 2), ) self.assertEqual(extension_token_binding.compose(), extension_token_binding_bytes) class TestExtensionCompressCertificateAlgorithms(unittest.TestCase): def test_parse(self): extension_compress_certificate_dict = collections.OrderedDict([ ('extension_type', b'\x00\x1b'), ('extension_length', b'\x00\x07'), ('compression_algorithm_list_length', b'\x06'), ('compression_algorithm_list', b'\x00\x03\x00\x02\x00\x01'), ]) extension_compress_certificate_bytes = b''.join(extension_compress_certificate_dict.values()) extension_compress_certificate = TlsExtensionCompressCertificate.parse_exact_size( extension_compress_certificate_bytes ) self.assertEqual(extension_compress_certificate.extension_type, TlsExtensionType.COMPRESS_CERTIFICATE) self.assertEqual( list(extension_compress_certificate.compression_algorithms), [ TlsCertificateCompressionAlgorithm.ZSTD, TlsCertificateCompressionAlgorithm.BROTLI, TlsCertificateCompressionAlgorithm.ZLIB, ] ) self.assertEqual(extension_compress_certificate.compose(), extension_compress_certificate_bytes) class TestExtensionSignatureAlgorithms(unittest.TestCase): def test_parse(self): extension_signature_algorithms_dict = collections.OrderedDict([ ('extension_type', b'\x00\x0d'), ('extension_length', b'\x00\x0c'), ('signature_algorithm_list_length', b'\x00\x0a'), ('signature_algorithm_list', b'\x01\x00\x02\x01\x03\x02\x04\x03\x0a\x0a'), ]) extension_signature_algorithms_bytes = b''.join(extension_signature_algorithms_dict.values()) extension_signature_algorithms = TlsExtensionSignatureAlgorithms.parse_exact_size( extension_signature_algorithms_bytes ) self.assertEqual(extension_signature_algorithms.extension_type, TlsExtensionType.SIGNATURE_ALGORITHMS) self.assertEqual( list(extension_signature_algorithms.hash_and_signature_algorithms), [ TlsSignatureAndHashAlgorithm.ANONYMOUS_MD5, TlsSignatureAndHashAlgorithm.RSA_SHA1, TlsSignatureAndHashAlgorithm.DSA_SHA224, TlsSignatureAndHashAlgorithm.ECDSA_SHA256, TlsInvalidTypeTwoByte(TlsGreaseTwoByte.GREASE_0A0A), ] ) self.assertEqual(extension_signature_algorithms.compose(), extension_signature_algorithms_bytes) class TestExtensionSignatureAlgorithmsCert(unittest.TestCase): def test_parse(self): extension_signature_algorithms_cert_dict = collections.OrderedDict([ ('extension_type', b'\x00\x32'), ('extension_length', b'\x00\x0a'), ('signature_algorithm_list_length', b'\x00\x08'), ('signature_algorithm_list', b'\x01\x00\x02\x01\x03\x02\x04\x03'), ]) extension_signature_algorithms_cert_bytes = b''.join(extension_signature_algorithms_cert_dict.values()) extension_signature_algorithms_cert = TlsExtensionSignatureAlgorithmsCert.parse_exact_size( extension_signature_algorithms_cert_bytes ) self.assertEqual(extension_signature_algorithms_cert.extension_type, TlsExtensionType.SIGNATURE_ALGORITHMS_CERT) self.assertEqual( list(extension_signature_algorithms_cert.hash_and_signature_algorithms), [ TlsSignatureAndHashAlgorithm.ANONYMOUS_MD5, TlsSignatureAndHashAlgorithm.RSA_SHA1, TlsSignatureAndHashAlgorithm.DSA_SHA224, TlsSignatureAndHashAlgorithm.ECDSA_SHA256, ] ) self.assertEqual(extension_signature_algorithms_cert.compose(), extension_signature_algorithms_cert_bytes) class TestExtensionSignedCertificateTimestampClient(unittest.TestCase): def test_parse_minimal(self): extension_sct_dict = collections.OrderedDict([ ('extension_type', b'\x00\x12'), ('extension_length', b'\x00\x00'), ]) extension_sct_bytes = b''.join(extension_sct_dict.values()) extension_sct = TlsExtensionSignedCertificateTimestampClient.parse_exact_size( extension_sct_bytes ) self.assertEqual( extension_sct.extension_type, TlsExtensionType.SIGNED_CERTIFICATE_TIMESTAMP ) self.assertEqual(extension_sct.compose(), extension_sct_bytes) class TestExtensionSignedCertificateTimestampServer(unittest.TestCase): def setUp(self): extension_sct_dict = collections.OrderedDict([ ('extension_type', b'\x00\x12'), ('extension_length', b'\x00\x84'), ('signed_certificate_timestamp_list_length', b'\x00\x82'), ('signed_certificate_timestamp_list', b''.join([ b'\x00\x3f', # length b'\x00', # version b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', # log b'\x00\x00\x00\x00\x00\x00\x00\x01', # timestamp b'\x00\x00', # extensions b'\x00\x00', # signature_algorithm b'\x00\x10\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', # signature b'\x00\x3f', # length b'\x00', # version b'\x96\x06\xc0\x2c\x69\x00\x33\xaa\x1d\x14\x5f\x59\xc6\xe2\x64\x8d', b'\x05\x49\xf0\xdf\x96\xaa\xb8\xdb\x91\x5a\x70\xd8\xec\xf3\x90\xa5', # log b'\x00\x00\x00\x00\x00\x00\x00\x01', # timestamp b'\x00\x00', # extensions b'\x00\x00', # signature_algorithm b'\x00\x10\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', # signature ])) ]) self.extension_sct_bytes = b''.join(extension_sct_dict.values()) def test_parse(self): extension_sct = TlsExtensionSignedCertificateTimestampServer.parse_exact_size( self.extension_sct_bytes ) self.assertEqual(extension_sct.extension_type, TlsExtensionType.SIGNED_CERTIFICATE_TIMESTAMP) sct = SignedCertificateTimestamp( CtVersion.V1, Base64Data(2 * b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), datetime.datetime(1970, 1, 1, 0, 0, 0, 1000, tzinfo=datetime.timezone.utc), CtExtensions([]), TlsSignatureAndHashAlgorithm.ANONYMOUS_NONE, b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', ) self.assertEqual(extension_sct.scts[0], sct) sct.log = CertificateTransparencyLog.AKAMAI_CT_LOG.value self.assertEqual(extension_sct.scts[1], sct) self.assertEqual(extension_sct.compose(), self.extension_sct_bytes) def test_markdown(self): scts = TlsExtensionSignedCertificateTimestampServer.parse_exact_size( self.extension_sct_bytes ).scts self.assertEqual(scts[0].as_markdown(), '\n'.join([ '* Version: V1', '* Log: AAECAwQFBgcICQoLDA0ODwABAgMEBQYHCAkKCwwNDg8=', '* Timestamp: 1970-01-01 00:00:00.001000+00:00', '* Extensions: -', '* Signature Algorithm: none with no encryption', '' ])) self.assertEqual(scts[1].as_markdown(), '\n'.join([ '* Version: V1', '* Log: Akamai CT Log (lgbALGkAM6odFF9ZxuJkjQVJ8N+WqrjbkVpw2OzzkKU=)', '* Timestamp: 1970-01-01 00:00:00.001000+00:00', '* Extensions: -', '* Signature Algorithm: none with no encryption', '' ])) class TestExtensionKeyShareClient(unittest.TestCase): def test_parse(self): extension_key_share_dict = collections.OrderedDict([ ('extension_type', b'\x00\x33'), ('extension_length', b'\x00\x2a'), ('key_share_length', b'\x00\x28'), ('group_grease', b'\x0a\x0a'), ('key_exchange_length_grease', b'\x00\x00'), ('group', b'\x00\x1d'), ('key_exchange_length', b'\x00\x20'), ('key_exchange', b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'') ]) extension_key_share_bytes = b''.join(extension_key_share_dict.values()) extension_key_share = TlsExtensionKeyShareClient.parse_exact_size( extension_key_share_bytes ) key_share_entries = extension_key_share.key_share_entries self.assertEqual(len(key_share_entries), 2) self.assertEqual(key_share_entries[0].group, TlsInvalidTypeTwoByte(TlsGreaseTwoByte.GREASE_0A0A)) self.assertEqual(key_share_entries[1].group, TlsNamedCurve.X25519) self.assertEqual( bytearray(key_share_entries[1].key_exchange), extension_key_share_dict['key_exchange'] ) self.assertEqual(extension_key_share.compose(), extension_key_share_bytes) class TestExtensionKeyShareReservedClient(unittest.TestCase): def test_parse(self): extension_key_share_dict = collections.OrderedDict([ ('extension_type', b'\x00\x28'), ('extension_length', b'\x00\x26'), ('key_share_length', b'\x00\x24'), ('group', b'\x00\x1d'), ('key_exchange_length', b'\x00\x20'), ('key_exchange', b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'') ]) extension_key_share_bytes = b''.join(extension_key_share_dict.values()) extension_key_share = TlsExtensionKeyShareReservedClient.parse_exact_size( extension_key_share_bytes ) key_share_entries = extension_key_share.key_share_entries # pylint: disable=no-member self.assertEqual(len(key_share_entries), 1) self.assertEqual(key_share_entries[0].group, TlsNamedCurve.X25519) self.assertEqual( bytearray(key_share_entries[0].key_exchange), extension_key_share_dict['key_exchange'] ) self.assertEqual(extension_key_share.compose(), extension_key_share_bytes) class TestExtensionKeyShareClientHelloRetry(unittest.TestCase): def test_error(self): extension_key_share_dict = collections.OrderedDict([ ('extension_type', b'\x00\x33'), ('extension_length', b'\x00\x04'), ('data', b'\x00\x00\x00\x00'), ]) extension_key_share_bytes = b''.join(extension_key_share_dict.values()) with self.assertRaises(InvalidType): # pylint: disable=expression-not-assigned TlsExtensionKeyShareClientHelloRetry.parse_exact_size( extension_key_share_bytes ) def test_parse(self): extension_key_share_dict = collections.OrderedDict([ ('extension_type', b'\x00\x33'), ('extension_length', b'\x00\x02'), ('group', b'\x00\x1d'), ]) extension_key_share_bytes = b''.join(extension_key_share_dict.values()) extension_key_share = TlsExtensionKeyShareClientHelloRetry.parse_exact_size( extension_key_share_bytes ) self.assertEqual(extension_key_share.selected_group, TlsNamedCurve.X25519) self.assertEqual(extension_key_share.compose(), extension_key_share_bytes) class TestExtensionKeyShareServer(unittest.TestCase): def test_parse(self): extension_key_share_dict = collections.OrderedDict([ ('extension_type', b'\x00\x33'), ('extension_length', b'\x00\x24'), ('group', b'\x00\x1d'), ('key_exchange_length', b'\x00\x20'), ('key_exchange', b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'') ]) extension_key_share_bytes = b''.join(extension_key_share_dict.values()) extension_key_share = TlsExtensionKeyShareServer.parse_exact_size( extension_key_share_bytes ) self.assertEqual(extension_key_share.key_share_entry.group, TlsNamedCurve.X25519) self.assertEqual( bytearray(extension_key_share.key_share_entry.key_exchange), extension_key_share_dict['key_exchange'] ) self.assertEqual(extension_key_share.compose(), extension_key_share_bytes) class TestExtensionCertificateStatusRequestClient(unittest.TestCase): def setUp(self): self.status_request_minimal_bytes = bytes( b'\x00\x05' + # handshake_type = STATUS_REQUEST b'\x00\x05' + # length = 0x05 b'\x01' + # status_type = OCSP b'\x00\x00' + # responder_id_list_length = 0x00 b'\x00\x00' + # request_extensions_length = 0x00 b'' ) self.status_request_minimal = TlsExtensionCertificateStatusRequestClient() self.request_extensions = b'\x00\x01\x02\x03\x04\x05\x06\x07' self.status_request_bytes = bytes( b'\x00\x05' + # handshake_type = STATUS_REQUEST b'\x00\x15' + # length = 0x05 b'\x01' + # status_type = OCSP b'\x00\x08' + # responder_id_list_length = 0x08 b'\x00\x01\x00\x00\x03\x01\x02\x03' + # responder_id_list b'\x00\x08' + # request_extensions_length = 0x08 self.request_extensions + # request_extensions b'' ) self.status_request = TlsExtensionCertificateStatusRequestClient( responder_id_list=TlsCertificateStatusRequestResponderIdList([ TlsCertificateStatusRequestResponderId(b'\x00'), TlsCertificateStatusRequestResponderId(b'\x01\x02\x03') ]), extensions=TlsCertificateStatusRequestExtensions(self.request_extensions) ) def test_parse(self): status_request_minimal = TlsExtensionCertificateStatusRequestClient.parse_exact_size( self.status_request_minimal_bytes ) self.assertEqual(status_request_minimal.responder_id_list, TlsCertificateStatusRequestResponderIdList([])) self.assertEqual(status_request_minimal.request_extensions, TlsCertificateStatusRequestExtensions([])) status_request = TlsExtensionCertificateStatusRequestClient.parse_exact_size(self.status_request_bytes) self.assertEqual( status_request.responder_id_list, TlsCertificateStatusRequestResponderIdList([ TlsCertificateStatusRequestResponderId(b'\x00'), TlsCertificateStatusRequestResponderId(b'\x01\x02\x03') ]) ) self.assertEqual( status_request.request_extensions, TlsCertificateStatusRequestExtensions(self.request_extensions) ) def test_compose(self): self.assertEqual(self.status_request_minimal.compose(), self.status_request_minimal_bytes) self.assertEqual(self.status_request.compose(), self.status_request_bytes) class TestExtensionCertificateStatusRequestServer(unittest.TestCase): def setUp(self): self.status_request_empty_bytes = bytes( b'\x00\x05' + # handshake_type = STATUS_REQUEST b'\x00\x00' + # length = 0x05 b'' ) self.status_request_empty = TlsExtensionCertificateStatusRequestServer() def test_parse(self): status_request_empty = TlsExtensionCertificateStatusRequestServer.parse_exact_size( self.status_request_empty_bytes ) self.assertEqual(status_request_empty, self.status_request_empty) def test_compose(self): self.assertEqual(self.status_request_empty.compose(), self.status_request_empty_bytes) class TestExtensionChannelId(unittest.TestCase): def test_parse(self): extension_channel_id_dict = collections.OrderedDict([ ('extension_type', b'\x75\x50'), ('extension_length', b'\x00\x00'), ]) extension_channel_id_bytes = b''.join(extension_channel_id_dict.values()) extension_channel_id = TlsExtensionChannelId.parse_exact_size( extension_channel_id_bytes ) self.assertEqual(extension_channel_id.compose(), extension_channel_id_bytes) class TestTlsExtensionPskKeyExchangeModes(unittest.TestCase): def test_parse(self): extension_pks_key_exchange_modes_dict = collections.OrderedDict([ ('extension_type', b'\x00\x2d'), ('extension_length', b'\x00\x03'), ('supported_version_list', b'\x02\x01\x00'), ]) extension_pks_key_exchange_modes_bytes = b''.join(extension_pks_key_exchange_modes_dict.values()) extension_pks_key_exchange_modes = TlsExtensionPskKeyExchangeModes.parse_exact_size( extension_pks_key_exchange_modes_bytes ) self.assertEqual( list(extension_pks_key_exchange_modes.key_exchange_modes), [ TlsPskKeyExchangeMode.PSK_DH_KE, TlsPskKeyExchangeMode.PSK_KE, ] ) self.assertEqual(extension_pks_key_exchange_modes.compose(), extension_pks_key_exchange_modes_bytes) class TestExtensionRenegotiationInfo(unittest.TestCase): def test_parse(self): extension_renegotiation_info_dict = collections.OrderedDict([ ('extension_type', b'\xff\x01'), ('extension_length', b'\x00\x09'), ('renegotiated_connection_length', b'\x08'), ('renegotiated_connection', b'\x00\x01\x02\x03\x04\x05\x06\x07'), ]) extension_renegotiation_info_bytes = b''.join(extension_renegotiation_info_dict.values()) extension_renegotiation_info = TlsExtensionRenegotiationInfo.parse_exact_size( extension_renegotiation_info_bytes ) self.assertEqual( extension_renegotiation_info.renegotiated_connection, TlsRenegotiatedConnection(b'\x00\x01\x02\x03\x04\x05\x06\x07') ) self.assertEqual(extension_renegotiation_info.compose(), extension_renegotiation_info_bytes) class TestExtensionSessionTicket(unittest.TestCase): def test_parse(self): extension_session_ticket_dict = collections.OrderedDict([ ('extension_type', b'\x00\x23'), ('extension_length', b'\x00\x08'), ('session_ticket', b'\x00\x01\x02\x03\x04\x05\x06\x07'), ]) extension_session_ticket_bytes = b''.join(extension_session_ticket_dict.values()) extension_session_ticket = TlsExtensionSessionTicket.parse_exact_size( extension_session_ticket_bytes ) self.assertEqual(extension_session_ticket.session_ticket, b'\x00\x01\x02\x03\x04\x05\x06\x07') self.assertEqual(extension_session_ticket.compose(), extension_session_ticket_bytes) class TestExtensionNextProtocolNegotiationClient(unittest.TestCase): def test_parse(self): extension_next_protocol_names_dict = collections.OrderedDict([ ('extension_type', b'\x33\x74'), ('extension_length', b'\x00\x00'), ]) extension_next_protocol_names_bytes = b''.join(extension_next_protocol_names_dict.values()) extension_next_protocol_names_mac = TlsExtensionNextProtocolNegotiationClient.parse_exact_size( extension_next_protocol_names_bytes ) self.assertEqual(extension_next_protocol_names_mac.compose(), extension_next_protocol_names_bytes) class TestExtensionNextProtocolNegotiationServer(unittest.TestCase): def test_parse(self): extension_next_protocol_names_dict = collections.OrderedDict([ ('extension_type', b'\x33\x74'), ('extension_length', b'\x00\x10'), ('protocol_name_h2_length', b'\x08'), ('protocol_name_h2', b'http/1.1'), ('protocol_name_h2c_length', b'\x06'), ('protocol_name_h2c', b'spdy/1'), ]) extension_next_protocol_names_bytes = b''.join(extension_next_protocol_names_dict.values()) extension_next_protocol_names = TlsExtensionNextProtocolNegotiationServer.parse_exact_size( extension_next_protocol_names_bytes ) self.assertEqual( extension_next_protocol_names.protocol_names, TlsNextProtocolNameList([TlsNextProtocolName.HTTP_1_1, TlsNextProtocolName.SPDY_1]) ) self.assertEqual(extension_next_protocol_names.compose(), extension_next_protocol_names_bytes) class TestExtensionApplicationLayerProtocolNegotiation(unittest.TestCase): def test_parse(self): extension_alpn_dict = collections.OrderedDict([ ('extension_type', b'\x00\x10'), ('extension_length', b'\x00\x09'), ('protocol_name_list_length', b'\x00\x07'), ('protocol_name_h2_length', b'\x02'), ('protocol_name_h2', b'h2'), ('protocol_name_h2c_length', b'\x03'), ('protocol_name_h2c', b'h2c'), ]) extension_alpn_bytes = b''.join(extension_alpn_dict.values()) extension_alpn = TlsExtensionApplicationLayerProtocolNegotiation.parse_exact_size( extension_alpn_bytes ) self.assertEqual(extension_alpn.protocol_names, TlsProtocolNameList([TlsProtocolName.H2, TlsProtocolName.H2C])) self.assertEqual(extension_alpn.compose(), extension_alpn_bytes) class TestExtensionApplicationLayerProtocolSettings(unittest.TestCase): def test_parse(self): extension_alpn_dict = collections.OrderedDict([ ('extension_type', b'\x44\xcd'), ('extension_length', b'\x00\x09'), ('protocol_name_list_length', b'\x00\x07'), ('protocol_name_h2_length', b'\x02'), ('protocol_name_h2', b'h2'), ('protocol_name_h2c_length', b'\x03'), ('protocol_name_h2c', b'h2c'), ]) extension_alpn_bytes = b''.join(extension_alpn_dict.values()) extension_alpn = TlsExtensionApplicationLayerProtocolSettings.parse_exact_size( extension_alpn_bytes ) self.assertEqual(extension_alpn.protocol_names, TlsProtocolNameList([TlsProtocolName.H2, TlsProtocolName.H2C])) self.assertEqual(extension_alpn.compose(), extension_alpn_bytes) class TestExtensionOldApplicationLayerProtocolSettings(unittest.TestCase): def test_parse(self): extension_alpn_dict = collections.OrderedDict([ ('extension_type', b'\x44\x69'), ('extension_length', b'\x00\x09'), ('protocol_name_list_length', b'\x00\x07'), ('protocol_name_h2_length', b'\x02'), ('protocol_name_h2', b'h2'), ('protocol_name_h2c_length', b'\x03'), ('protocol_name_h2c', b'h2c'), ]) extension_alpn_bytes = b''.join(extension_alpn_dict.values()) extension_alpn = TlsExtensionOldApplicationLayerProtocolSettings.parse_exact_size( extension_alpn_bytes ) self.assertEqual(extension_alpn.protocol_names, TlsProtocolNameList([TlsProtocolName.H2, TlsProtocolName.H2C])) self.assertEqual(extension_alpn.compose(), extension_alpn_bytes) class TestExtensionUnusedData(unittest.TestCase): def test_error(self): extension_unused_data_dict = collections.OrderedDict([ ('extension_type', b'\xff\x01'), ('extension_length', b'\x00\x01'), ('extension_data', b'\xff'), ]) extension_unused_data_bytes = b''.join(extension_unused_data_dict.values()) with self.assertRaises(InvalidValue) as context_manager: # pylint: disable=expression-not-assigned TestUnusedDataExtension.parse_exact_size(extension_unused_data_bytes) self.assertEqual(context_manager.exception.value, b'\xff') class TestExtensionEncryptThenMAC(unittest.TestCase): def test_parse(self): extension_encrypt_then_mac_dict = collections.OrderedDict([ ('extension_type', b'\x00\x16'), ('extension_length', b'\x00\x00'), ]) extension_encrypt_then_mac_bytes = b''.join(extension_encrypt_then_mac_dict.values()) extension_encrypt_then_mac = TlsExtensionEncryptThenMAC.parse_exact_size(extension_encrypt_then_mac_bytes) self.assertEqual(extension_encrypt_then_mac.compose(), extension_encrypt_then_mac_bytes) class TestExtensionExtendedMasterSecret(unittest.TestCase): def test_parse(self): extension_extended_master_secret_dict = collections.OrderedDict([ ('extension_type', b'\x00\x17'), ('extension_length', b'\x00\x00'), ]) extension_extended_master_secret_bytes = b''.join(extension_extended_master_secret_dict.values()) extended_master_secret = TlsExtensionExtendedMasterSecret.parse_exact_size( extension_extended_master_secret_bytes ) self.assertEqual(extended_master_secret.compose(), extension_extended_master_secret_bytes) class TestExtensionExtendedShortRecordHeader(unittest.TestCase): def test_parse(self): extension_short_record_header_dict = collections.OrderedDict([ ('extension_type', b'\xff\x03'), ('extension_length', b'\x00\x00'), ]) extension_short_record_header_bytes = b''.join(extension_short_record_header_dict.values()) short_record_header = TlsExtensionShortRecordHeader.parse_exact_size( extension_short_record_header_bytes ) self.assertEqual(short_record_header.compose(), extension_short_record_header_bytes) class TestExtensionRecordSizeLimit(unittest.TestCase): def test_parse(self): extension_record_size_limit_dict = collections.OrderedDict([ ('extension_type', b'\x00\x1c'), ('extension_length', b'\x00\x02'), ('record_size_limit', b'\x00\xff'), ]) extension_record_size_limit_bytes = b''.join(extension_record_size_limit_dict.values()) extension_record_size_limit = TlsExtensionRecordSizeLimit.parse_exact_size(extension_record_size_limit_bytes) self.assertEqual(extension_record_size_limit.record_size_limit, 0xff) self.assertEqual(extension_record_size_limit.compose(), extension_record_size_limit_bytes) class TestExtensionServerPadding(unittest.TestCase): def test_parse(self): extension_server_padding_dict = collections.OrderedDict([ ('extension_type', b'\x12\xe0'), ('extension_length', b'\x00\x02'), ('padding_size', b'\x00\x00'), ]) extension_server_padding_bytes = b''.join(extension_server_padding_dict.values()) extension_server_padding = TlsExtensionServerPadding.parse_exact_size(extension_server_padding_bytes) self.assertEqual(extension_server_padding.padding_size, 0) self.assertEqual(extension_server_padding.compose(), extension_server_padding_bytes) class TestExtensionTrustAnchors(unittest.TestCase): def test_parse(self): extension_trust_anchors_dict = collections.OrderedDict([ ('extension_type', b'\xca\x34'), ('extension_length', b'\x00\x02'), ('trust_anchor_identifiers_length', b'\x00\x00'), ]) extension_trust_anchors_bytes = b''.join(extension_trust_anchors_dict.values()) extension_trust_anchors = TlsExtensionTrustAnchors.parse_exact_size(extension_trust_anchors_bytes) self.assertEqual(extension_trust_anchors.trust_anchor_identifiers, TlsTrustAnchorIdentifierList([])) self.assertEqual(extension_trust_anchors.compose(), extension_trust_anchors_bytes) def test_parse_non_empty_list(self): extension_trust_anchors_dict = collections.OrderedDict([ ('extension_type', b'\xca\x34'), ('extension_length', b'\x00\x07'), ('trust_anchor_identifiers_length', b'\x00\x05'), ('trust_anchor_identifier_1_length', b'\x02'), ('trust_anchor_identifier_1', b'\x01\x02'), ('trust_anchor_identifier_2_length', b'\x01'), ('trust_anchor_identifier_2', b'\x03'), ]) extension_trust_anchors_bytes = b''.join(extension_trust_anchors_dict.values()) extension_trust_anchors = TlsExtensionTrustAnchors.parse_exact_size(extension_trust_anchors_bytes) self.assertEqual( extension_trust_anchors.trust_anchor_identifiers, TlsTrustAnchorIdentifierList([ TlsTrustAnchorIdentifier(b'\x01\x02'), TlsTrustAnchorIdentifier(b'\x03'), ]) ) self.assertEqual(extension_trust_anchors.compose(), extension_trust_anchors_bytes) class TestExtensionPadding(unittest.TestCase): def test_error_non_zero_padding(self): extension_padding_dict = collections.OrderedDict([ ('extension_type', b'\x00\x15'), ('extension_length', b'\x00\x04'), ('extension_data', b'\x00\x00\x00\x01'), ]) extension_padding_bytes = b''.join(extension_padding_dict.values()) with self.assertRaises(InvalidValue) as context_manager: TlsExtensionPadding.parse_exact_size(extension_padding_bytes) self.assertEqual(context_manager.exception.value, b'\x01') def test_parse(self): extension_padding_minimal_dict = collections.OrderedDict([ ('extension_type', b'\x00\x15'), ('extension_length', b'\x00\x00'), ]) extension_padding_minimal_bytes = b''.join(extension_padding_minimal_dict.values()) extension_padding_minimal = TlsExtensionPadding.parse_exact_size(extension_padding_minimal_bytes) self.assertEqual(extension_padding_minimal.compose(), extension_padding_minimal_bytes) extension_padding_with_data_dict = collections.OrderedDict([ ('extension_type', b'\x00\x15'), ('extension_length', b'\x00\x04'), ('extension_data', b'\x00\x00\x00\x00'), ]) extension_padding_with_data_bytes = b''.join(extension_padding_with_data_dict.values()) extension_padding_with_data = TlsExtensionPadding.parse_exact_size(extension_padding_with_data_bytes) self.assertEqual(extension_padding_with_data.compose(), extension_padding_with_data_bytes) class ExtensionEncryptedClientHelloBase(unittest.TestCase): def test_parse(self): extension_encrypted_client_hello_dict = collections.OrderedDict([ ('extension_type', b'\xfe\x0d'), ('extension_length', b'\x00\x01'), ('hello_type', b'\x00'), ]) extension_encrypted_client_hello_bytes = b''.join(extension_encrypted_client_hello_dict.values()) with self.assertRaises(InvalidType): # pylint: disable=expression-not-assigned TlsExtensionEncryptedClientHelloInner.parse_exact_size(extension_encrypted_client_hello_bytes) class ExtensionEncryptedClientHelloInner(unittest.TestCase): def test_parse(self): extension_encrypted_client_hello_dict = collections.OrderedDict([ ('extension_type', b'\xfe\x0d'), ('extension_length', b'\x00\x01'), ('hello_type', b'\x01'), ]) extension_encrypted_client_hello_bytes = b''.join(extension_encrypted_client_hello_dict.values()) extension_encrypted_client_hello = TlsExtensionEncryptedClientHelloInner.parse_exact_size( extension_encrypted_client_hello_bytes ) self.assertEqual( extension_encrypted_client_hello.compose(), extension_encrypted_client_hello_bytes ) class ExtensionEncryptedClientHelloOuter(unittest.TestCase): def test_parse(self): extension_encrypted_client_hello_dict = collections.OrderedDict([ ('extension_type', b'\xfe\x0d'), ('extension_length', b'\x00\x11'), ('hello_type', b'\x00'), ('hello_data', b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ]) extension_encrypted_client_hello_bytes = b''.join(extension_encrypted_client_hello_dict.values()) extension_encrypted_client_hello = TlsExtensionEncryptedClientHelloOuter.parse_exact_size( extension_encrypted_client_hello_bytes ) self.assertEqual( extension_encrypted_client_hello.compose(), extension_encrypted_client_hello_bytes ) def test_parse_followed_by_another_extension(self): extension_encrypted_client_hello_dict = collections.OrderedDict([ ('extension_type', b'\xfe\x0d'), ('extension_length', b'\x00\x11'), ('hello_type', b'\x00'), ('hello_data', b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ]) extension_encrypted_client_hello_bytes = b''.join(extension_encrypted_client_hello_dict.values()) extension_extended_master_secret_dict = collections.OrderedDict([ ('extension_type', b'\x00\x17'), ('extension_length', b'\x00\x00'), ]) extension_extended_master_secret_bytes = b''.join(extension_extended_master_secret_dict.values()) extension_encrypted_client_hello, parsed_length = TlsExtensionEncryptedClientHelloOuter.parse_immutable( extension_encrypted_client_hello_bytes + extension_extended_master_secret_bytes ) self.assertEqual(parsed_length, len(extension_encrypted_client_hello_bytes)) self.assertEqual( extension_encrypted_client_hello.data, extension_encrypted_client_hello_dict['hello_data'] ) self.assertEqual( extension_encrypted_client_hello.compose(), extension_encrypted_client_hello_bytes ) class TestExtensionPostHandshakeAuthentication(unittest.TestCase): def test_parse(self): extension_post_handshake_authentication_dict = collections.OrderedDict([ ('extension_type', b'\x00\x31'), ('extension_length', b'\x00\x00'), ]) extension_post_handshake_authentication_bytes = b''.join(extension_post_handshake_authentication_dict.values()) extension_post_handshake_authentication = TlsExtensionPostHandshakeAuthentication.parse_exact_size( extension_post_handshake_authentication_bytes ) self.assertEqual( extension_post_handshake_authentication.compose(), extension_post_handshake_authentication_bytes ) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_grease.py000066400000000000000000000032551524413560000265170ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from unittest import mock from cryptoparser.tls.grease import ( TlsInvalidType, TlsInvalidTypeOneByte, TlsInvalidTypeParamsOneByte, TlsInvalidTypeParamsTwoByte, TlsInvalidTypeTwoByte, ) class TestGrease(unittest.TestCase): def test_parse(self): self.assertEqual( TlsInvalidTypeOneByte.parse_exact_size(b'\x2a').value, TlsInvalidTypeParamsOneByte(0x2a, TlsInvalidType.GREASE) ) self.assertEqual( TlsInvalidTypeOneByte.parse_exact_size(b'\x2b').value, TlsInvalidTypeParamsOneByte(0x2b, TlsInvalidType.UNKNOWN) ) self.assertEqual( TlsInvalidTypeTwoByte.parse_exact_size(b'\x2a\x2a').value, TlsInvalidTypeParamsTwoByte(0x2a2a, TlsInvalidType.GREASE) ) self.assertEqual( TlsInvalidTypeTwoByte.parse_exact_size(b'\x2b\x2b').value, TlsInvalidTypeParamsTwoByte(0x2b2b, TlsInvalidType.UNKNOWN) ) def test_compose(self): self.assertEqual( b'\x2a', TlsInvalidTypeOneByte(0x2a).compose() ) self.assertEqual( b'\x2b', TlsInvalidTypeOneByte(0x2b).compose() ) self.assertEqual( b'\x2a\x2a', TlsInvalidTypeTwoByte(0x2a2a).compose() ) self.assertEqual( b'\x2b\x2b', TlsInvalidTypeTwoByte(0x2b2b).compose() ) @mock.patch('random.choice') def test_random(self, mock_choice): mock_choice.side_effect = [0xabcd, ] self.assertEqual(TlsInvalidTypeTwoByte.from_random(), TlsInvalidTypeTwoByte(0xabcd)) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_handshake.py000066400000000000000000001241361524413560000272010ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 # pylint: disable=too-many-lines import unittest import collections import copy import datetime import hashlib from cryptodatahub.common.exception import InvalidValue from cryptodatahub.tls.algorithm import ( TlsCipherSuiteExtension, TlsECPointFormat, TlsGreaseOneByte, TlsGreaseTwoByte, TlsNamedCurve, TlsProtocolName, TlsSignatureAndHashAlgorithm, ) from cryptoparser.common.exception import InvalidType, NotEnoughData from cryptoparser.tls.ciphersuite import TlsCipherSuite, SslCipherKind from cryptoparser.tls.extension import ( TlsExtensionApplicationLayerProtocolNegotiation, TlsExtensionSignatureAlgorithms, TlsExtensionSupportedVersionsClient, TlsExtensionSupportedVersionsServer, TlsExtensionUnparsed, TlsExtensionEllipticCurves, TlsExtensionECPointFormats, TlsECPointFormatVector, TlsEllipticCurveVector, TlsSupportedVersionVector, ) from cryptoparser.tls.grease import TlsInvalidTypeOneByte, TlsInvalidTypeTwoByte from cryptoparser.tls.subprotocol import ( TLS_HANDSHAKE_HELLO_RETRY_REQUEST_RANDOM, TLS_HANDSHAKE_HELLO_RETRY_REQUEST_RANDOM_BYTES, SslHandshakeClientHello, SslHandshakeServerHello, SslMessageType, TlsAlertMessage, TlsCertificate, TlsCertificateEntry, TlsCertificateEntryVector, TlsCertificateStatusType, TlsCertificates, TlsClientCertificateType, TlsCipherSuiteVector, TlsCompressionMethod, TlsCompressionMethodVector, TlsContentType, TlsDistinguishedName, TlsExtensionType, TlsExtensionsClient, TlsExtensionsServer, TlsHandshakeCertificate, TlsHandshakeServerCertificate, TlsHandshakeCertificateRequest, TlsHandshakeCertificateStatus, TlsHandshakeClientHello, TlsHandshakeEncryptedExtensions, TlsHandshakeHelloRandom, TlsHandshakeHelloRandomBytes, TlsHandshakeHelloRetryRequest, TlsHandshakeMessageVariant, TlsHandshakeServerHello, TlsHandshakeServerHelloDone, TlsHandshakeServerKeyExchange, TlsHandshakeType, TlsSessionIdVector, TlsSubprotocolMessageParser, ) from cryptoparser.tls.record import TlsRecord from cryptoparser.tls.version import TlsVersion, TlsProtocolVersion from .classes import TestMessage class TestSubprotocolParser(unittest.TestCase): def test_error(self): subprotocol_parser = TlsSubprotocolMessageParser(TlsContentType.HEARTBEAT) with self.assertRaises(InvalidValue) as context_manager: subprotocol_parser.parse( b'\x18' + # type = heartbeat b'\x03\x03' + # version = TLS 1.2 b'\x00\x01' + # length = 1 b'\x00' ) self.assertEqual(context_manager.exception.value, 0x18) def test_registered_parser(self): tls_message_dict = collections.OrderedDict([ ('level', b'\x02'), # FATAL ('description', b'\x28'), # HANDSHAKE_FAILURE ]) tls_message_bytes = b''.join(tls_message_dict.values()) tls_parser = TlsSubprotocolMessageParser(TlsContentType.ALERT) tls_parser.parse(tls_message_bytes) tls_parser.register_subprotocol_parser(TlsContentType.ALERT, TestMessage) with self.assertRaises(NotImplementedError): tls_parser.parse(tls_message_bytes) tls_parser.register_subprotocol_parser(TlsContentType.ALERT, TlsAlertMessage) parsed_object, _ = tls_parser.parse(tls_message_bytes) self.assertEqual(parsed_object.compose(), tls_message_bytes) class TestVariantParsable(unittest.TestCase): def setUp(self): self.server_hello_done_dict = collections.OrderedDict([ ('handshake_type', b'\x0e'), # SERVER_HELLO_DONE ('length', b'\x00\x00\x00'), # 0x00 ]) self.server_hello_done_bytes = b''.join(self.server_hello_done_dict.values()) self.server_hello_done = TlsHandshakeServerHelloDone() def test_error(self): invalid_tls_message_dict = collections.OrderedDict([ ('content_type', b'\x17'), ('data', b'\x00\x00\x00'), ]) invalid_tls_message_bytes = b''.join(invalid_tls_message_dict.values()) with self.assertRaisesRegex(InvalidValue, 'is not a valid TlsHandshakeMessageVariant'): TlsHandshakeMessageVariant.parse_exact_size(invalid_tls_message_bytes) def test_compose(self): self.assertEqual(TlsHandshakeMessageVariant(self.server_hello_done).compose(), self.server_hello_done_bytes) class TestTlsCipherSuiteVector(unittest.TestCase): def test_parse(self): cipher_suites = TlsCipherSuiteVector.parse_exact_size(b'\x00\x02\x00\x00') self.assertEqual(cipher_suites, TlsCipherSuiteVector([TlsCipherSuite.TLS_NULL_WITH_NULL_NULL])) cipher_suites = TlsCipherSuiteVector.parse_exact_size(b'\x00\x04\x56\x00\x00\xff') self.assertEqual( cipher_suites, TlsCipherSuiteVector([ TlsInvalidTypeTwoByte(TlsCipherSuiteExtension.FALLBACK_SCSV), TlsInvalidTypeTwoByte(TlsCipherSuiteExtension.EMPTY_RENEGOTIATION_INFO_SCSV), ]) ) cipher_suites = TlsCipherSuiteVector.parse_exact_size(b'\x00\x06\x56\x00\x00\x00\x00\xff') self.assertEqual( cipher_suites, TlsCipherSuiteVector([ TlsInvalidTypeTwoByte(TlsCipherSuiteExtension.FALLBACK_SCSV), TlsCipherSuite.TLS_NULL_WITH_NULL_NULL, TlsInvalidTypeTwoByte(TlsCipherSuiteExtension.EMPTY_RENEGOTIATION_INFO_SCSV), ]) ) class TestTlsHandshake(unittest.TestCase): def setUp(self): self.server_hello_done_dict = collections.OrderedDict([ ('handshake_type', b'\x0e'), # SERVER_HELLO_DONE ('length', b'\x00\x00\x00'), ]) self.server_hello_done_bytes = b''.join(self.server_hello_done_dict.values()) self.server_hello_done_record_dict = collections.OrderedDict([ ('content_type', b'\x16'), # HANDSHAKE ('protocol_version', b'\x03\x01'), # TLS1 ('length', b'\x00\x04'), ]) self.server_hello_done_record_bytes = b''.join(self.server_hello_done_record_dict.values()) def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: # pylint: disable=expression-not-assigned TlsHandshakeClientHello.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, 4) with self.assertRaises(InvalidType): # pylint: disable=expression-not-assigned TlsHandshakeClientHello.parse_exact_size(self.server_hello_done_bytes) with self.assertRaises(NotEnoughData) as context_manager: TlsHandshakeClientHello.parse_exact_size( b'\x01' # handshake_type: CLIENT_HELLO b'\x00\x00\x03' + # handshake_length = 3 b'\x03\x03' + # version = TLS 1.2 b'' ) self.assertEqual(context_manager.exception.bytes_needed, 1) with self.assertRaises(InvalidValue) as context_manager: TlsHandshakeClientHello.parse_exact_size( b'\xff' # handshake_type: INVALID b'\x00\x00\x02' + # handshake_length = 2 b'\x03\x03' + # version = TLS 1.2 b'' ) self.assertEqual(context_manager.exception.value, 0xff) def test_parse(self): record = TlsRecord.parse_exact_size( self.server_hello_done_record_bytes + self.server_hello_done_bytes ) self.assertEqual(record.protocol_version, TlsProtocolVersion(TlsVersion.TLS1)) record.protocol_version = TlsProtocolVersion(TlsVersion.TLS1_2) self.assertEqual(record.protocol_version, TlsProtocolVersion(TlsVersion.TLS1_2)) class TestTlsHandshakeClientHello(unittest.TestCase): def setUp(self): self.client_hello_minimal_dict = collections.OrderedDict([ ('handshake_type ', b'\x01'), # CLIENT_HELLO ('length ', b'\x00\x00\x37'), ('version ', b'\x03\x03'), ('random ', b'\x5b\x6c\xd5\x80\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b''), ('session_id_length', b'\x00'), ('cipher_suite_length', b'\x00\x10'), ('cipher_suites', b'\x0a\x0a\x00\x01\x00\x02\x00\x03' + b'\x00\x04\x00\x05\x56\x00\x00\xff' + b''), ('compression_method_length', b'\x01'), ('compression_methods', b'\x00'), ]) self.client_hello_minimal_bytes = b''.join(self.client_hello_minimal_dict.values()) self.client_hello_minimal_extensions_dict = collections.OrderedDict([ ('extensions_length', b'\x00\x0d'), ('extension_type', b'\x00\x2b'), # SUPPORTED_VERSIONS ('extension_length', b'\x00\x05'), ('supported_version_list_length', b'\x04'), ('supported_version_list', b'\x03\x02\x03\x03'), # TLS1_1, TLS1_2 ('extension_grease', b'\x0a\x0a'), ('extension_grease_length', b'\x00\x00'), ]) self.client_hello_minimal_extensions_bytes = b''.join(self.client_hello_minimal_extensions_dict.values()) self.client_hello_extension_bytes = bytearray( self.client_hello_minimal_bytes + self.client_hello_minimal_extensions_bytes + b'' ) self.client_hello_extension_bytes[3] += ( len(self.client_hello_extension_bytes) - len(self.client_hello_minimal_bytes) ) self.random_time = datetime.datetime(2018, 8, 10, tzinfo=datetime.timezone.utc) self.client_hello_minimal = TlsHandshakeClientHello( TlsCipherSuiteVector([ TlsInvalidTypeTwoByte(TlsGreaseTwoByte.GREASE_0A0A), TlsCipherSuite.TLS_RSA_WITH_NULL_MD5, TlsCipherSuite.TLS_RSA_WITH_NULL_SHA, TlsCipherSuite.TLS_RSA_EXPORT_WITH_RC4_40_MD5, TlsCipherSuite.TLS_RSA_WITH_RC4_128_MD5, TlsCipherSuite.TLS_RSA_WITH_RC4_128_SHA, ]), TlsProtocolVersion(TlsVersion.TLS1_2), TlsHandshakeHelloRandom( self.random_time, TlsHandshakeHelloRandomBytes(bytearray( b'\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'' )) ), TlsSessionIdVector(()), TlsCompressionMethodVector([TlsCompressionMethod.NULL, ]), TlsExtensionsClient(()), fallback_scsv=True, empty_renegotiation_info_scsv=True, ) def test_parse(self): client_hello_minimal = TlsHandshakeClientHello.parse_exact_size(self.client_hello_minimal_bytes) self.assertEqual(client_hello_minimal.get_handshake_type(), TlsHandshakeType.CLIENT_HELLO) self.assertEqual( client_hello_minimal.protocol_version, TlsProtocolVersion(TlsVersion.TLS1_2) ) self.assertEqual( client_hello_minimal.random, self.client_hello_minimal.random ) self.assertEqual( client_hello_minimal.random.time, self.random_time, ) self.assertEqual( client_hello_minimal.random.random, self.client_hello_minimal.random.random ) self.assertEqual( client_hello_minimal.cipher_suites, self.client_hello_minimal.cipher_suites ) self.assertEqual( client_hello_minimal.compression_methods, self.client_hello_minimal.compression_methods ) self.assertEqual( client_hello_minimal.extensions, self.client_hello_minimal.extensions ) self.assertTrue(client_hello_minimal.fallback_scsv) self.assertTrue(client_hello_minimal.empty_renegotiation_info_scsv) client_hello_extension = TlsHandshakeClientHello.parse_exact_size(self.client_hello_extension_bytes) self.assertEqual(len(client_hello_extension.extensions), 2) self.assertEqual( client_hello_extension.extensions.get_item_by_type(TlsExtensionType.SUPPORTED_VERSIONS), TlsExtensionSupportedVersionsClient(TlsSupportedVersionVector([ TlsProtocolVersion(TlsVersion.TLS1_1), TlsProtocolVersion(TlsVersion.TLS1_2), ])) ) self.assertEqual( client_hello_extension.extensions[1], TlsExtensionUnparsed(TlsInvalidTypeTwoByte(TlsGreaseTwoByte.GREASE_0A0A), b'') ) with self.assertRaises(KeyError): client_hello_extension.extensions.get_item_by_type(TlsGreaseTwoByte.GREASE_0A0A) def test_compose(self): self.assertEqual( self.client_hello_minimal.compose(), self.client_hello_minimal_bytes ) client_hello_extension = TlsHandshakeClientHello.parse_exact_size(self.client_hello_extension_bytes) self.assertEqual( client_hello_extension.compose(), self.client_hello_extension_bytes ) def test_ja3(self): client_hello_minimal = copy.copy(self.client_hello_minimal) self.assertEqual(client_hello_minimal.ja3(), '771,2570-1-2-3-4-5,,,') client_hello_minimal.extensions.append( TlsExtensionEllipticCurves(TlsEllipticCurveVector([TlsNamedCurve.SECT163K1])) ) self.assertEqual(client_hello_minimal.ja3(), '771,2570-1-2-3-4-5,10,1,') client_hello_minimal.extensions[0].elliptic_curves.append(TlsInvalidTypeTwoByte(TlsGreaseTwoByte.GREASE_0A0A)) self.assertEqual(client_hello_minimal.ja3(), '771,2570-1-2-3-4-5,10,1,') client_hello_minimal.extensions.append( TlsExtensionECPointFormats(TlsECPointFormatVector([TlsECPointFormat.UNCOMPRESSED])) ) self.assertEqual(client_hello_minimal.ja3(), '771,2570-1-2-3-4-5,10-11,1,0') client_hello_minimal.extensions[1].point_formats.append(TlsInvalidTypeOneByte(TlsGreaseOneByte.GREASE_0B)) self.assertEqual(client_hello_minimal.ja3(), '771,2570-1-2-3-4-5,10-11,1,0') @staticmethod def _ja4_hash(text): return hashlib.sha256(text.encode('ascii')).hexdigest()[:12] def test_ja4_sub_hashes(self): # canonical example from the JA4 technical specification ciphers = '002f,0035,009c,009d,1301,1302,1303,c013,c014,c02b,c02c,c02f,c030,cca8,cca9' extensions = '0005,000a,000b,000d,0012,0015,0017,001b,0023,002b,002d,0033,4469,ff01' signature_algorithms = '0403,0804,0401,0503,0805,0501,0806,0601' self.assertEqual(self._ja4_hash(ciphers), '8daaf6152771') self.assertEqual(self._ja4_hash(extensions + '_' + signature_algorithms), 'e5627efa2ab1') def test_ja4(self): client_hello_minimal = copy.copy(self.client_hello_minimal) # the GREASE cipher (2570 in the JA3 tag) is excluded; five ciphers remain cipher_hash = self._ja4_hash('0001,0002,0003,0004,0005') result = client_hello_minimal.ja4() self.assertEqual(result.fingerprint, f't12i050000_{cipher_hash}_000000000000') # ciphers are already in ascending order and there are no extensions, so the original-order # hashed form equals the sorted one self.assertEqual(result.fingerprint_original, result.fingerprint) self.assertEqual(result.fingerprint_raw, 't12i050000_0001,0002,0003,0004,0005__') self.assertEqual(result.fingerprint_raw_original, 't12i050000_0001,0002,0003,0004,0005__') def test_ja4_original_order(self): client_hello = TlsHandshakeClientHello( TlsCipherSuiteVector([ TlsCipherSuite.TLS_RSA_WITH_NULL_SHA, # 0x0002 TlsCipherSuite.TLS_RSA_WITH_NULL_MD5, # 0x0001 ]), TlsProtocolVersion(TlsVersion.TLS1_2), ) result = client_hello.ja4() # the sorted hash is over 0001,0002; the original-order hash is over 0002,0001, so they differ self.assertEqual(result.fingerprint, f't12i020000_{self._ja4_hash("0001,0002")}_000000000000') self.assertEqual(result.fingerprint_original, f't12i020000_{self._ja4_hash("0002,0001")}_000000000000') self.assertNotEqual(result.fingerprint, result.fingerprint_original) self.assertEqual(result.fingerprint_raw, 't12i020000_0001,0002__') self.assertEqual(result.fingerprint_raw_original, 't12i020000_0002,0001__') def test_ja4_supported_versions_and_extensions(self): client_hello_minimal = copy.copy(self.client_hello_minimal) client_hello_minimal.extensions.append( TlsExtensionSupportedVersionsClient(TlsSupportedVersionVector([ TlsProtocolVersion(TlsVersion.TLS1_3), ])) ) cipher_hash = self._ja4_hash('0001,0002,0003,0004,0005') extension_hash = self._ja4_hash('002b') result = client_hello_minimal.ja4() # version is the highest in supported_versions (TLS 1.3 -> "13"), not the legacy protocol version self.assertEqual(result.fingerprint, f't13i050100_{cipher_hash}_{extension_hash}') self.assertEqual(result.fingerprint_raw, 't13i050100_0001,0002,0003,0004,0005_002b_') def test_ja4_alpn_and_signature_algorithms(self): client_hello_minimal = copy.copy(self.client_hello_minimal) client_hello_minimal.extensions.append( TlsExtensionApplicationLayerProtocolNegotiation([TlsProtocolName.H2]) ) client_hello_minimal.extensions.append( TlsExtensionSignatureAlgorithms([TlsSignatureAndHashAlgorithm.RSA_PSS_RSAE_SHA256]) ) signature_algorithm_hex = f'{TlsSignatureAndHashAlgorithm.RSA_PSS_RSAE_SHA256.value.code:04x}' cipher_hash = self._ja4_hash('0001,0002,0003,0004,0005') extension_hash = self._ja4_hash(f'000d_{signature_algorithm_hex}') result = client_hello_minimal.ja4() # the ALPN value "h2" is in the prefix; SNI (0000) and ALPN (0010) are excluded from the # extension hash, leaving signature_algorithms (000d) with the algorithms appended self.assertEqual(result.fingerprint, f't12i0502h2_{cipher_hash}_{extension_hash}') self.assertEqual( result.fingerprint_raw, f't12i0502h2_0001,0002,0003,0004,0005_000d_{signature_algorithm_hex}' ) class TestTlsHandshakeServerHello(unittest.TestCase): def setUp(self): self.server_hello_minimal_dict = collections.OrderedDict([ ('handshake_type', b'\x02'), # SERVER_HELLO ('length', b'\x00\x00\x26'), ('version', b'\x03\x03'), # TLS1_2 ('random', b'\x5b\x6c\xd5\x80\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b''), ('session_id_length', b'\x00'), ('cipher_suite', b'\x00\x01'), ('compression_method', b'\x00'), ]) self.server_hello_minimal_bytes = b''.join(self.server_hello_minimal_dict.values()) self.server_hello_minimal = TlsHandshakeServerHello( TlsProtocolVersion(TlsVersion.TLS1_2), TlsHandshakeHelloRandom( datetime.datetime(2018, 8, 10, tzinfo=datetime.timezone.utc), TlsHandshakeHelloRandomBytes(bytearray( b'\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'' )) ), TlsSessionIdVector(()), TlsCompressionMethod.NULL, TlsCipherSuite.TLS_RSA_WITH_NULL_MD5, TlsExtensionsServer(()) ) self.server_hello_minimal_bytes = b''.join(self.server_hello_minimal_dict.values()) self.server_hello_minimal_extensions_dict = collections.OrderedDict([ ('extensions_length', b'\x00\x0a'), ('extension_type', b'\x00\x2b'), # SUPPORTED_VERSIONS ('extension_length', b'\x00\x05'), ('selected_version', b'\x03\x03'), # TLS1_2 ('extension_grease', b'\x0a\x0a'), ('extension_grease_length', b'\x00\x00'), ]) self.server_hello_minimal_extensions_bytes = b''.join(self.server_hello_minimal_extensions_dict.values()) self.server_hello_extension_bytes = bytearray( self.server_hello_minimal_bytes + self.server_hello_minimal_extensions_bytes + b'' ) self.server_hello_extension_bytes[3] += ( len(self.server_hello_extension_bytes) - len(self.server_hello_minimal_bytes) ) def test_parse(self): server_hello_minimal = TlsHandshakeServerHello.parse_exact_size(self.server_hello_minimal_bytes) self.assertEqual(server_hello_minimal.get_handshake_type(), TlsHandshakeType.SERVER_HELLO) self.assertEqual( server_hello_minimal.protocol_version, TlsProtocolVersion(TlsVersion.TLS1_2) ) self.assertEqual( server_hello_minimal.random, self.server_hello_minimal.random ) self.assertEqual( server_hello_minimal.cipher_suite, self.server_hello_minimal.cipher_suite ) self.assertEqual( server_hello_minimal.compression_method, self.server_hello_minimal.compression_method ) self.assertEqual( server_hello_minimal.extensions, self.server_hello_minimal.extensions ) server_hello_extension = TlsHandshakeServerHello.parse_exact_size(self.server_hello_extension_bytes) self.assertEqual(len(server_hello_extension.extensions), 2) self.assertEqual( server_hello_extension.extensions.get_item_by_type(TlsExtensionType.SUPPORTED_VERSIONS), TlsExtensionSupportedVersionsServer(TlsProtocolVersion(TlsVersion.TLS1_2)) ) self.assertEqual( server_hello_extension.extensions[1], TlsExtensionUnparsed(TlsInvalidTypeTwoByte(TlsGreaseTwoByte.GREASE_0A0A), b'') ) with self.assertRaises(KeyError): server_hello_extension.extensions.get_item_by_type(TlsGreaseTwoByte.GREASE_0A0A) def test_compose(self): self.assertEqual( self.server_hello_minimal.compose(), self.server_hello_minimal_bytes ) class TestTlsHandshakeHelloRetryRequest(unittest.TestCase): def setUp(self): self.hello_retry_request_minimal_dict = collections.OrderedDict([ ('handshake_type', b'\x06'), # HELLO_RETRY_REQUEST ('length', b'\x00\x00\x26'), ('version', b'\x03\x03'), # TLS1_2 ('random', b'\x5b\x6c\xd5\x80\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b''), ('session_id_length', b'\x00'), ('cipher_suite', b'\x00\x01'), ('compression_method', b'\x00'), ]) self.hello_retry_request_minimal_bytes = b''.join(self.hello_retry_request_minimal_dict.values()) self.hello_retry_request_minimal = TlsHandshakeHelloRetryRequest( TlsCipherSuite.TLS_RSA_WITH_NULL_MD5, TlsProtocolVersion(TlsVersion.TLS1_2), TlsHandshakeHelloRandom( datetime.datetime(2018, 8, 10, tzinfo=datetime.timezone.utc), TlsHandshakeHelloRandomBytes(bytearray( b'\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'' )) ), TlsSessionIdVector(()), TlsCompressionMethod.NULL, TlsExtensionsClient(()) ) def test_parse(self): hello_retry_request_minimal = TlsHandshakeHelloRetryRequest.parse_exact_size( self.hello_retry_request_minimal_bytes ) self.assertEqual(hello_retry_request_minimal.get_handshake_type(), TlsHandshakeType.HELLO_RETRY_REQUEST) self.assertEqual( hello_retry_request_minimal.protocol_version, TlsProtocolVersion(TlsVersion.TLS1_2) ) self.assertEqual( hello_retry_request_minimal.random_bytes, self.hello_retry_request_minimal.random_bytes ) self.assertEqual( hello_retry_request_minimal.cipher_suite, self.hello_retry_request_minimal.cipher_suite ) self.assertEqual( hello_retry_request_minimal.compression_method, self.hello_retry_request_minimal.compression_method ) self.assertEqual( hello_retry_request_minimal.extensions, self.hello_retry_request_minimal.extensions ) def test_compose(self): self.assertEqual( self.hello_retry_request_minimal.compose(), self.hello_retry_request_minimal_bytes ) def test_random(self): self.assertEqual( TLS_HANDSHAKE_HELLO_RETRY_REQUEST_RANDOM.compose(), TLS_HANDSHAKE_HELLO_RETRY_REQUEST_RANDOM_BYTES ) self.assertEqual( TlsHandshakeHelloRandom.parse_exact_size(TLS_HANDSHAKE_HELLO_RETRY_REQUEST_RANDOM_BYTES), TLS_HANDSHAKE_HELLO_RETRY_REQUEST_RANDOM ) class TestTlsHandshakeServerCertificate(unittest.TestCase): def setUp(self): self.certificate_minimal_dict = collections.OrderedDict([ ('handshake_type', b'\x0b'), # CERTIFICATE ('length', b'\x00\x00\x31'), ('cretificates', b'\x00\x00\x2e'), ('peer_cretificate_length', b'\x00\x00\x10'), ('peer_certificate_bytes', b'peer certificate'), ('intermrdiate_cretificate_length', b'\x00\x00\x18'), ('intermrdiate_certificate_bytes', b'intermediate certificate'), ]) self.certificate_minimal_bytes = b''.join(self.certificate_minimal_dict.values()) self.certificate_minimal = TlsHandshakeServerCertificate( TlsCertificates([ TlsCertificate(b'peer certificate'), TlsCertificate(b'intermediate certificate'), ]) ) def test_parse(self): certificate_minimal = TlsHandshakeServerCertificate.parse_exact_size(self.certificate_minimal_bytes) self.assertEqual( certificate_minimal.certificate_chain, self.certificate_minimal.certificate_chain ) def test_compose(self): self.assertEqual( self.certificate_minimal.compose(), self.certificate_minimal_bytes ) class TestTlsHandshakeCertificate(unittest.TestCase): def setUp(self): self.certificate_minimal_dict = collections.OrderedDict([ ('handshake_type', b'\x0b'), # CERTIFICATE ('length', b'\x00\x00\x1a'), ('certificate_request_context_length', b'\x01'), ('certificate_request_context', b'\xaa'), ('certificate_entries_length', b'\x00\x00\x15'), ('entry_certificate_length', b'\x00\x00\x10'), ('entry_certificate', b'first-cert-bytes'), ('entry_extensions_length', b'\x00\x00'), ]) self.certificate_minimal_bytes = b''.join(self.certificate_minimal_dict.values()) self.certificate_minimal = TlsHandshakeCertificate( certificate_request_context=b'\xaa', certificate_entries=TlsCertificateEntryVector([ TlsCertificateEntry( TlsCertificate(b'first-cert-bytes'), b'', ) ]), ) def test_parse(self): certificate = TlsHandshakeCertificate.parse_exact_size(self.certificate_minimal_bytes) self.assertEqual( certificate.certificate_request_context, self.certificate_minimal.certificate_request_context ) self.assertEqual( len(certificate.certificate_entries), len(self.certificate_minimal.certificate_entries) ) def test_compose(self): self.assertEqual( self.certificate_minimal.compose(), self.certificate_minimal_bytes ) def test_error_certificate_request_context_too_long(self): with self.assertRaises(InvalidValue) as context_manager: TlsHandshakeCertificate( certificate_request_context=b'\x00' * 256, certificate_entries=TlsCertificateEntryVector([]), ) self.assertEqual(context_manager.exception.value, 256) def test_error_entry_extensions_too_long(self): with self.assertRaises(InvalidValue) as context_manager: TlsCertificateEntry( TlsCertificate(b'cert'), b'\x00' * (2 ** 16), ) self.assertEqual(context_manager.exception.value, 2 ** 16) def test_error_parse_invalid_entries(self): bad_dict = collections.OrderedDict([ ('handshake_type', b'\x0b'), # CERTIFICATE ('length', b'\x00\x00\x05'), ('certificate_request_context_length', b'\x01'), ('certificate_request_context', b'\xaa'), ('certificate_entries_length', b'\x00\x00\x0f'), # claims more data than available ]) bad_bytes = b''.join(bad_dict.values()) with self.assertRaises(InvalidType): TlsHandshakeCertificate.parse_exact_size(bad_bytes) class TestTlsHandshakeCertificateRequestTls10(unittest.TestCase): def setUp(self): self.certificate_request_dict = collections.OrderedDict([ ('handshake_type', b'\x0d'), # CERTIFICATE_REQUEST ('length', b'\x00\x00\x19'), ('certificate_types_length', b'\x04'), ('certificate_types', b'\x01\02\x03\x04'), ('certificate_authorities_length', b'\x00\x12'), ('certificate_authority_length', b'\x00\x10'), ('certificate_authority', b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ]) self.certificate_request_bytes = b''.join(self.certificate_request_dict.values()) self.certificate_request = TlsHandshakeCertificateRequest( certificate_types=[ TlsClientCertificateType.RSA_SIGN, TlsClientCertificateType.DSS_SIGN, TlsClientCertificateType.RSA_FIXED_DH, TlsClientCertificateType.DSS_FIXED_DH, ], certificate_authorities=[ TlsDistinguishedName(b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ] ) def test_parse(self): certificate_request = TlsHandshakeCertificateRequest.parse_exact_size(self.certificate_request_bytes) self.assertEqual( list(certificate_request.certificate_types), [ TlsClientCertificateType.RSA_SIGN, TlsClientCertificateType.DSS_SIGN, TlsClientCertificateType.RSA_FIXED_DH, TlsClientCertificateType.DSS_FIXED_DH, ] ) self.assertEqual( list(certificate_request.certificate_authorities), [TlsDistinguishedName(b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ] ) self.assertEqual(certificate_request.supported_signature_algorithms, None) def test_compose(self): self.assertEqual( self.certificate_request.compose(), self.certificate_request_bytes ) class TestTlsHandshakeCertificateRequestTls12(unittest.TestCase): def setUp(self): self.certificate_request_dict = collections.OrderedDict([ ('handshake_type', b'\x0d'), # CERTIFICATE_REQUEST ('length', b'\x00\x00\x23'), ('certificate_types_length', b'\x04'), ('certificate_types', b'\x01\02\x03\x04'), ('signature_algorithm_list_length', b'\x00\x08'), ('signature_algorithm_list', b'\x01\x00\x02\x01\x03\x02\x04\x03'), ('certificate_authorities_length', b'\x00\x12'), ('certificate_authority_length', b'\x00\x10'), ('certificate_authority', b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ]) self.certificate_request_bytes = b''.join(self.certificate_request_dict.values()) self.certificate_request = TlsHandshakeCertificateRequest( certificate_types=[ TlsClientCertificateType.RSA_SIGN, TlsClientCertificateType.DSS_SIGN, TlsClientCertificateType.RSA_FIXED_DH, TlsClientCertificateType.DSS_FIXED_DH, ], supported_signature_algorithms=[ TlsSignatureAndHashAlgorithm.ANONYMOUS_MD5, TlsSignatureAndHashAlgorithm.RSA_SHA1, TlsSignatureAndHashAlgorithm.DSA_SHA224, TlsSignatureAndHashAlgorithm.ECDSA_SHA256, ], certificate_authorities=[ TlsDistinguishedName(b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ] ) def test_parse(self): certificate_request = TlsHandshakeCertificateRequest.parse_exact_size(self.certificate_request_bytes) self.assertEqual( list(certificate_request.certificate_types), [ TlsClientCertificateType.RSA_SIGN, TlsClientCertificateType.DSS_SIGN, TlsClientCertificateType.RSA_FIXED_DH, TlsClientCertificateType.DSS_FIXED_DH, ] ) self.assertEqual( list(certificate_request.supported_signature_algorithms), [ TlsSignatureAndHashAlgorithm.ANONYMOUS_MD5, TlsSignatureAndHashAlgorithm.RSA_SHA1, TlsSignatureAndHashAlgorithm.DSA_SHA224, TlsSignatureAndHashAlgorithm.ECDSA_SHA256, ] ) self.assertEqual( list(certificate_request.certificate_authorities), [TlsDistinguishedName(b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ] ) def test_compose(self): self.assertEqual( self.certificate_request.compose(), self.certificate_request_bytes ) class TestTlsHandshakeServerHelloDone(unittest.TestCase): def setUp(self): self.server_hello_done_dict = collections.OrderedDict([ ('handshake_type', b'\x0e'), # SERVER_HELLO_DONE ('length', b'\x00\x00\x00'), # 0x00 ]) self.server_hello_done_bytes = b''.join(self.server_hello_done_dict.values()) self.server_hello_done = TlsHandshakeServerHelloDone() def test_error(self): error_regex = 'b\'\\\\x00\' is not a valid TlsHandshakeServerHelloDone payload value' with self.assertRaisesRegex(InvalidValue, error_regex): # pylint: disable=expression-not-assigned TlsHandshakeServerHelloDone.parse_exact_size(b'\x0e\x00\x00\x01\x00') def test_parse(self): server_hello_done = TlsHandshakeServerHelloDone.parse_exact_size(self.server_hello_done_bytes) self.assertEqual(server_hello_done.get_handshake_type(), TlsHandshakeType.SERVER_HELLO_DONE) def test_compose(self): self.assertEqual(self.server_hello_done.compose(), self.server_hello_done_bytes) class TestTlsHandshakeServerKeyExcahnge(unittest.TestCase): def setUp(self): self.param_bytes = b'\x00\x01\x02\x03\x04\x05\x06\x07' self.server_key_exchange_dict = collections.OrderedDict([ ('handshake_type', b'\x0c'), # SERVER_KEY_EXCHANGE ('length', b'\x00\x00\x08'), ('param_bytes', self.param_bytes), ]) self.server_key_exchange_bytes = b''.join(self.server_key_exchange_dict.values()) self.server_key_exchange = TlsHandshakeServerKeyExchange(self.param_bytes) def test_parse(self): server_key_exchange = TlsHandshakeServerKeyExchange.parse_exact_size(self.server_key_exchange_bytes) self.assertEqual(server_key_exchange.get_handshake_type(), TlsHandshakeType.SERVER_KEY_EXCHANGE) self.assertEqual(server_key_exchange.param_bytes, self.param_bytes) def test_compose(self): self.assertEqual(self.server_key_exchange.compose(), self.server_key_exchange_bytes) class TestTlsHandshakeCertificateStatus(unittest.TestCase): def setUp(self): self.status_bytes = b'\x00\x01\x02\x03\x04\x05\x06\x07' self.certificate_status_bytes = bytes( b'\x16' + # handshake_type = CERTIFICATE_STATUS b'\x00\x00\x0c' + # length = 0x0c b'\x01' + # status_type = OCSP b'\x00\x00\x08' + # length = 0x08 self.status_bytes + # status_bytes b'' ) self.certificate_status = TlsHandshakeCertificateStatus(TlsCertificateStatusType.OCSP, self.status_bytes) def test_parse(self): certificate_status = TlsHandshakeCertificateStatus.parse_exact_size(self.certificate_status_bytes) self.assertEqual(certificate_status.get_handshake_type(), TlsHandshakeType.CERTIFICATE_STATUS) self.assertEqual(certificate_status.status, self.status_bytes) def test_compose(self): self.assertEqual(self.certificate_status.compose(), self.certificate_status_bytes) class TestTlsHandshakeEncryptedExtensions(unittest.TestCase): def setUp(self): self.extension_data = b'\x00\x10\x00\x0b\x00\x0c\x00\x00' self.encrypted_extensions_dict = collections.OrderedDict([ ('handshake_type', b'\x08'), # ENCRYPTED_EXTENSIONS ('length', b'\x00\x00\x08'), ('extension_data', self.extension_data), ]) self.encrypted_extensions_bytes = b''.join(self.encrypted_extensions_dict.values()) self.encrypted_extensions = TlsHandshakeEncryptedExtensions(self.extension_data) def test_parse(self): encrypted_extensions = TlsHandshakeEncryptedExtensions.parse_exact_size(self.encrypted_extensions_bytes) self.assertEqual(encrypted_extensions.get_handshake_type(), TlsHandshakeType.ENCRYPTED_EXTENSIONS) self.assertEqual(encrypted_extensions.extension_data, self.extension_data) def test_compose(self): self.assertEqual( self.encrypted_extensions.compose(), self.encrypted_extensions_bytes ) class TestSslHandshakeClientHello(unittest.TestCase): def setUp(self): self.client_hello_dict = collections.OrderedDict([ ('version', b'\x00\x02'), # SSL2 ('cipher_kinds_length', b'\x00\x06'), ('session_id_length', b'\x00\x08'), ('challenge_length', b'\x00\x10'), ('cipher_kinds', b'\x01\x00\x80\x07\x00\xc0'), ('session_id', b'\x00\x01\x02\x03\x04\x05\x06\x07'), ('challenge', b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'') ]) self.client_hello_bytes = b''.join(self.client_hello_dict.values()) self.client_hello = SslHandshakeClientHello( cipher_kinds=[ SslCipherKind.SSL_CK_RC4_128_WITH_MD5, SslCipherKind.SSL_CK_DES_192_EDE3_CBC_WITH_MD5 ], session_id=b'\x00\x01\x02\x03\x04\x05\x06\x07', challenge=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' ) def test_default(self): self.assertEqual(SslHandshakeClientHello([]).session_id, b'') self.assertNotEqual(SslHandshakeClientHello([]).challenge, SslHandshakeClientHello([]).challenge) self.assertEqual(SslHandshakeServerHello(b'', []).connection_id, b'') def test_parse(self): client_hello_minimal = SslHandshakeClientHello.parse_exact_size(self.client_hello_bytes) self.assertEqual(client_hello_minimal.get_message_type(), SslMessageType.CLIENT_HELLO) def test_compose(self): self.assertEqual(self.client_hello.compose(), self.client_hello_bytes) class TestSslHandshakeServerHello(unittest.TestCase): def setUp(self): self.server_hello_done_dict = collections.OrderedDict([ ('session_id_hit', b'\x00'), # False ('certificate_type', b'\x01'), # X509_CERTIFICATE ('version', b'\x00\x02'), # SSL2 ('certificate_length', b'\x00\x0b'), ('cipher_kinds_length', b'\x00\x06'), ('connection_id_length', b'\x00\x10'), ('certificate', b'certificate'), ('cipher_kinds', b'\x01\x00\x80\x07\x00\xc0'), ('connection_id', b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b''), ]) self.server_hello_bytes = b''.join(self.server_hello_done_dict.values()) self.server_hello = SslHandshakeServerHello( certificate=b'certificate', cipher_kinds=[ SslCipherKind.SSL_CK_RC4_128_WITH_MD5, SslCipherKind.SSL_CK_DES_192_EDE3_CBC_WITH_MD5 ], connection_id=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f', session_id_hit=False ) def test_parse(self): server_hello_minimal = SslHandshakeServerHello.parse_exact_size(self.server_hello_bytes) self.assertEqual(server_hello_minimal.get_message_type(), SslMessageType.SERVER_HELLO) def test_compose(self): self.assertEqual(self.server_hello.compose(), self.server_hello_bytes) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_ldap.py000066400000000000000000000105461524413560000261720ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import copy import unittest import collections from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import NotEnoughData from cryptoparser.tls.ldap import ( LDAPExtendedRequestStartTLS, LDAPExtendedResponseStartTLS, LDAPResultCode, ) class TestLDAPExtendedRequest(unittest.TestCase): def setUp(self): self.ldap_extended_request_dict = collections.OrderedDict([ ('message_sequence', b'\x30\x1d'), ('message_id', b'\x02\x01\x01'), ('protocol_op', b'\x77\x18'), ('extended_request', b'\x80'), ('request_name', ( b'\x16\x31\x2e\x33\x2e\x36\x2e\x31\x2e\x34\x2e\x31\x2e\x31\x34\x36' + b'\x36\x2e\x32\x30\x30\x33\x37' )), ]) self.ldap_extended_request_bytes = b''.join(self.ldap_extended_request_dict.values()) self.ldap_extended_request = LDAPExtendedRequestStartTLS() def test_parse(self): LDAPExtendedRequestStartTLS.parse_exact_size(self.ldap_extended_request_bytes) def test_compose(self): self.assertEqual(self.ldap_extended_request.compose(), self.ldap_extended_request_bytes) class TestLDAPExtendedResponseMinimal(unittest.TestCase): def setUp(self): self.ldap_extended_response_dict = collections.OrderedDict([ ('message_sequence', b'\x30\x0c'), ('message_id', b'\x02\x01\x01'), ('protocol_op', b'\x78\x07'), ('extended_response', b''), ('result_code', b'\x0a\x01\x07'), ('matched_dn', b'\x04\x00'), ('diagnostic_message', b'\x04\x00'), ('referral', b''), ]) self.ldap_extended_response_bytes = b''.join(self.ldap_extended_response_dict.values()) self.ldap_extended_response = LDAPExtendedResponseStartTLS(LDAPResultCode.AUTH_METHOD_NOT_SUPPORTED) def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: # pylint: disable=expression-not-assigned LDAPExtendedResponseStartTLS.parse_exact_size( self.ldap_extended_response_bytes[:LDAPExtendedResponseStartTLS.HEADER_SIZE] ) self.assertEqual( context_manager.exception.bytes_needed, len(self.ldap_extended_response_bytes) - LDAPExtendedResponseStartTLS.HEADER_SIZE ) ldap_extended_response_dict = copy.copy(self.ldap_extended_response_dict) ldap_extended_response_dict['protocol_op'] = b'\xff\xff' ldap_extended_response_bytes = b''.join(ldap_extended_response_dict.values()) with self.assertRaises(InvalidValue) as context_manager: # pylint: disable=expression-not-assigned LDAPExtendedResponseStartTLS.parse_exact_size(ldap_extended_response_bytes) def test_parse(self): ldap_extended_response = LDAPExtendedResponseStartTLS.parse_exact_size(self.ldap_extended_response_bytes) self.assertEqual(ldap_extended_response.result_code, self.ldap_extended_response.result_code) def test_compose(self): self.assertEqual(self.ldap_extended_response.compose(), self.ldap_extended_response_bytes) class TestLDAPExtendedResponseFull(unittest.TestCase): def setUp(self): self.ldap_extended_response_dict = collections.OrderedDict([ ('message_sequence', b'\x30\x34'), ('message_id', b'\x02\x01\x01'), ('protocol_op', b'\x78\x2f'), ('extended_response', b''), ('result_code', b'\x0a\x01\x07'), ('matched_dn', b'\x04\x08\x00\x01\x02\x03\x04\x05\x06\x07'), ('diagnostic_message', b'\x04\x08\x00\x01\x02\x03\x04\x05\x06\x07'), ('referral', b''), ('response_name', ( b'\x8a\x16\x31\x2e\x33\x2e\x36\x2e\x31\x2e\x34\x2e\x31\x2e\x31\x34' + b'\x36\x36\x2e\x32\x30\x30\x33\x37' )), ]) self.ldap_extended_response_bytes = b''.join(self.ldap_extended_response_dict.values()) self.ldap_extended_response = LDAPExtendedResponseStartTLS(LDAPResultCode.AUTH_METHOD_NOT_SUPPORTED) def test_parse(self): ldap_extended_response = LDAPExtendedResponseStartTLS.parse_exact_size(self.ldap_extended_response_bytes) self.assertEqual(ldap_extended_response.result_code, self.ldap_extended_response.result_code) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_mysql.py000066400000000000000000000200741524413560000264140ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import NotEnoughData from cryptoparser.tls.mysql import ( MySQLRecord, MySQLCapability, MySQLCharacterSet, MySQLStatusFlag, MySQLHandshakeSslRequest, MySQLHandshakeV10, MySQLVersion, ) class TestMySQLRecord(unittest.TestCase): def setUp(self): self.test_record = MySQLRecord( packet_number=1, packet_bytes=b'\x01\x02\x03\x04' ) self.test_record_bytes = bytes( b'\x04\x00\x00' + # packet_length b'\x01' + # packet_number b'\x01\x02\x03\x04' ) def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: MySQLRecord.parse_exact_size(b'\x00') self.assertEqual(context_manager.exception.bytes_needed, MySQLRecord.HEADER_SIZE - 1) def test_parse(self): MySQLRecord.parse_exact_size(self.test_record_bytes) def test_compose(self): self.assertEqual(self.test_record.compose(), self.test_record_bytes) class TestMySQLHandshake10(unittest.TestCase): def setUp(self): self.handshake_minimal = MySQLHandshakeV10( protocol_version=MySQLVersion.MYSQL_9, server_version='1.2.3.4', connection_id=0x01020304, auth_plugin_data=b'\x01\x02\x03\x04\x05\x06\x07\x08', capabilities={ MySQLCapability.CLIENT_SSL, MySQLCapability.CLIENT_MULTI_STATEMENTS, }, character_set=MySQLCharacterSet.UTF8, states={MySQLStatusFlag.SERVER_STATUS_IN_TRANS, }, ) self.handshake_full = MySQLHandshakeV10( protocol_version=MySQLVersion.MYSQL_9, server_version='1.2.3.4', connection_id=0x01020304, auth_plugin_data=b'\x01\x02\x03\x04\x05\x06\x07\x08', capabilities={ MySQLCapability.CLIENT_SSL, MySQLCapability.CLIENT_MULTI_STATEMENTS, MySQLCapability.CLIENT_PLUGIN_AUTH, }, character_set=MySQLCharacterSet.UTF8, states={MySQLStatusFlag.SERVER_STATUS_IN_TRANS, }, auth_plugin_name='auth_plugin_name', auth_plugin_data_2=b'\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d' ) self.handshake_bytes_base = bytes( b'\x09' + # protocol_version b'1.2.3.4\x00' + # server_version b'\x04\x03\x02\x01' + # connection_id b'\x01\x02\x03\x04\x05\x06\x07\x08' + # auth_plugin_data b'\x00' + # filler b'\x00\x08' + # capabilities b'\x21' + # character_set b'\x01\x00' + # states b'' ) self.handshake_bytes_minimal = self.handshake_bytes_base + bytes( b'\x01\x00' + # capabilities_2 b'\x00' + # auth_plugin_data_len 10 * b'\x00' + # reserved b'' ) self.handshake_bytes_full = self.handshake_bytes_base + bytes( b'\x09\x00' + # capabilities_2 b'\x15' + # auth_plugin_data_len 10 * b'\x00' + # reserved b'\x01\x02\x03\x04\x05\x06\x07\x08' + # auth_plugin_data_2 b'\x09\x0a\x0b\x0c\x0d' + b'auth_plugin_name\x00' + # auth_plugin_name b'' ) def test_error_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: MySQLHandshakeV10.parse_exact_size(b'\x00') self.assertEqual(context_manager.exception.bytes_needed, MySQLHandshakeV10.MINIMUM_SIZE - 1) def test_error_no_auth_plugin_data(self): handshake_bytes_no_auth_plugin_data = self.handshake_bytes_base + bytes( b'\x09\x00' + # capabilities_2 b'\x00' + # auth_plugin_data_len 10 * b'\x00' + # reserved b'' ) with self.assertRaises(InvalidValue) as context_manager: MySQLHandshakeV10.parse_exact_size(handshake_bytes_no_auth_plugin_data) self.assertEqual(context_manager.exception.value, 0) def test_parse(self): handshake_minimal = MySQLHandshakeV10.parse_exact_size(self.handshake_bytes_minimal) self.assertEqual(handshake_minimal, self.handshake_minimal) handshake_full = MySQLHandshakeV10.parse_exact_size(self.handshake_bytes_full) self.assertEqual(handshake_full, self.handshake_full) def test_compose(self): self.assertEqual(self.handshake_minimal.compose(), self.handshake_bytes_minimal) self.assertEqual(self.handshake_full.compose(), self.handshake_bytes_full) class TestMySQLHandshakeSslRequest(unittest.TestCase): def setUp(self): self.handshake_bytes_with_client_41_capability = bytes( b'\x00\x02\x00\x00' + # capabilities b'\x04\x03\x02\x01' + # max_packet_size b'\x21' + # character_set 23 * b'\x00' + # filler b'' ) self.handshake_with_client_41_capability = MySQLHandshakeSslRequest( capabilities={MySQLCapability.CLIENT_PROTOCOL_41}, max_packet_size=0x01020304, character_set=MySQLCharacterSet.UTF8 ) self.handshake_bytes_without_client_41_capability = bytes( b'\x00\x00' + # capabilities b'\x03\x02\x01' + # max_packet_size b'' ) self.handshake_without_client_41_capability = MySQLHandshakeSslRequest( capabilities=set(), max_packet_size=0x010203, ) def test_error_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: MySQLHandshakeSslRequest.parse_exact_size(b'\x00') self.assertEqual(context_manager.exception.bytes_needed, MySQLHandshakeSslRequest.MINIMUM_SIZE - 1) def test_with_client_41_character_set_is_none(self): ssl_request = MySQLHandshakeSslRequest( capabilities=set([MySQLCapability.CLIENT_PROTOCOL_41]), max_packet_size=1, character_set=None ) self.assertEqual(ssl_request.character_set, MySQLCharacterSet.UTF8) def test_error_without_client_41_capability_too_large(self): with self.assertRaises(ValueError) as context_manager: MySQLHandshakeSslRequest(capabilities=set([MySQLCapability.CLIENT_MULTI_STATEMENTS]), max_packet_size=1) self.assertEqual(context_manager.exception.args[0], 1) def test_error_without_client_41_max_packet_size_too_large(self): # pylint: disable=invalid-name with self.assertRaises(ValueError) as context_manager: MySQLHandshakeSslRequest(capabilities=set(), max_packet_size=2 ** 24) self.assertEqual(context_manager.exception.args[0], 2 ** 24) def test_with_client_41_capability(self): self.assertEqual( MySQLHandshakeSslRequest.parse_exact_size(self.handshake_bytes_with_client_41_capability), self.handshake_with_client_41_capability, ) self.assertEqual( self.handshake_with_client_41_capability.compose(), self.handshake_bytes_with_client_41_capability, ) def test_without_client_41_capability(self): self.assertEqual( MySQLHandshakeSslRequest.parse_exact_size(self.handshake_bytes_without_client_41_capability), self.handshake_without_client_41_capability, ) self.assertEqual( self.handshake_without_client_41_capability.compose(), self.handshake_bytes_without_client_41_capability, ) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_openvpn.py000066400000000000000000000204151524413560000267330ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest import collections from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import InvalidType, NotEnoughData from cryptoparser.tls.openvpn import ( OpenVpnPacketAckV1, OpenVpnPacketBase, OpenVpnPacketControlV1, OpenVpnPacketHardResetClientV2, OpenVpnPacketHardResetServerV2, OpenVpnPacketVariant, OpenVpnPacketWrapperTcp, ) class TestOpenVpnPacketWrapperTcp(unittest.TestCase): def test_error_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: OpenVpnPacketWrapperTcp.parse_exact_size(b'\x00\x02\x00') self.assertEqual(context_manager.exception.bytes_needed, 1) def test_parse(self): packet = OpenVpnPacketWrapperTcp.parse_exact_size(b'\x00\x08\x00\x01\x02\x03\x04\x05\x06\x07') self.assertEqual(packet.payload, b'\x00\x01\x02\x03\x04\x05\x06\x07') def test_compose(self): packet = OpenVpnPacketWrapperTcp(b'\x00\x01\x02\x03\x04\x05\x06\x07') self.assertEqual(packet.compose(), b'\x00\x08\x00\x01\x02\x03\x04\x05\x06\x07') class TestOpenVpnPacketBase(unittest.TestCase): def test_error_not_enough_data(self): with self.assertRaises(NotEnoughData) as context_manager: OpenVpnPacketBase.parse_header(b'\x00') self.assertEqual(context_manager.exception.bytes_needed, OpenVpnPacketBase.HEADER_SIZE - 1) def test_error_wrong_packet_type(self): with self.assertRaises(InvalidType): OpenVpnPacketAckV1.parse_header(b'\x00' * OpenVpnPacketBase.HEADER_SIZE) def test_packet_id_array(self): header_without_packet_array_dict = collections.OrderedDict([ ('opcode', b'\x20'), ('session_id', b'\x00\x01\x02\x03\x04\x05\x06\x07'), ]) header_without_packet_array_bytes = b''.join(header_without_packet_array_dict.values()) session_id, packet_id_array, remote_session_id, header_length = OpenVpnPacketControlV1.parse_header( header_without_packet_array_bytes + b'\x00' ) self.assertEqual(session_id, 0x0001020304050607) self.assertEqual(packet_id_array, []) self.assertEqual(remote_session_id, None) self.assertEqual(header_length, len(header_without_packet_array_bytes) + 1) packet_array_dict = collections.OrderedDict([ ('packet_id_array_length', b'\x01'), ('packet_id_array', b'\x04\x05\x06\x07'), ('remote_session_id', b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ]) packet_array_bytes = b''.join(packet_array_dict.values()) session_id, packet_id_array, remote_session_id, header_length = OpenVpnPacketControlV1.parse_header( header_without_packet_array_bytes + packet_array_bytes ) self.assertEqual(session_id, 0x0001020304050607) self.assertEqual(packet_id_array, [0x04050607]) self.assertEqual(remote_session_id, 0x08090a0b0c0d0e0f) self.assertEqual(header_length, len(header_without_packet_array_bytes) + len(packet_array_bytes)) class TestOpenVpnPacketControlV1(unittest.TestCase): def setUp(self): self.control_dict = collections.OrderedDict([ ('opcode', b'\x20'), ('session_id', b'\x00\x01\x02\x03\x04\x05\x06\x07'), ('packet_id_array_length', b'\x01'), ('packet_id_array', b'\x04\x05\x06\x07'), ('remote_session_id', b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ('packet_id', b'\x00\x01\x02\x03'), ('payload', b'\x00\x01\x02\x03'), ]) self.control_bytes = b''.join(self.control_dict.values()) self.control = OpenVpnPacketControlV1( session_id=0x0001020304050607, packet_id_array=[0x04050607], remote_session_id=0x08090a0b0c0d0e0f, packet_id=0x00010203, payload=b'\x00\x01\x02\x03', ) def test_parse(self): control = OpenVpnPacketVariant.parse_exact_size(self.control_bytes) self.assertEqual(control.session_id, 0x0001020304050607) self.assertEqual(control.packet_id_array, [0x04050607]) self.assertEqual(control.remote_session_id, 0x08090a0b0c0d0e0f) self.assertEqual(control.packet_id, 0x00010203) def test_compose(self): self.assertEqual(self.control.compose(), self.control_bytes) class TestOpenVpnPacketAckV1(unittest.TestCase): def setUp(self): self.ack_dict = collections.OrderedDict([ ('opcode', b'\x28'), ('session_id', b'\x00\x01\x02\x03\x04\x05\x06\x07'), ('packet_id_array_length', b'\x01'), ('packet_id_array', b'\x04\x05\x06\x07'), ('remote_session_id', b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ]) self.ack_bytes = b''.join(self.ack_dict.values()) self.ack = OpenVpnPacketAckV1( session_id=0x0001020304050607, packet_id_array=[0x04050607], remote_session_id=0x08090a0b0c0d0e0f, ) def test_parse(self): ack = OpenVpnPacketVariant.parse_exact_size(self.ack_bytes) self.assertEqual(ack.session_id, 0x0001020304050607) self.assertEqual(ack.packet_id_array, [0x04050607]) def test_compose(self): self.assertEqual(self.ack.compose(), self.ack_bytes) class TestOpenVpnPacketHardResetClientV2(unittest.TestCase): def setUp(self): self.hard_reset_client_dict = collections.OrderedDict([ ('opcode', b'\x38'), ('session_id', b'\x00\x01\x02\x03\x04\x05\x06\x07'), ('packet_id_array_length', b'\x00'), ('packet_id', b'\x00\x01\x02\x03'), ]) self.hard_reset_client_bytes = b''.join(self.hard_reset_client_dict.values()) self.hard_reset_client = OpenVpnPacketHardResetClientV2( session_id=0x0001020304050607, packet_id=0x00010203, ) def test_error_non_empty_packet_id_array(self): hard_reset_client_dict = collections.OrderedDict([ ('opcode', b'\x38'), ('session_id', b'\x00\x01\x02\x03\x04\x05\x06\x07'), ('packet_id_array_length', b'\x01'), ('packet_id_array', b'\x00\x01\x02\x03'), ('remote_session_id', b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ]) hard_reset_client_bytes = b''.join(hard_reset_client_dict.values()) with self.assertRaises(InvalidValue) as context_manager: OpenVpnPacketVariant.parse_exact_size(hard_reset_client_bytes) self.assertEqual(context_manager.exception.value, [0x00010203]) def test_parse(self): hard_reset_client = OpenVpnPacketVariant.parse_exact_size(self.hard_reset_client_bytes) self.assertEqual(hard_reset_client.session_id, 0x0001020304050607) self.assertEqual(hard_reset_client.packet_id_array, []) self.assertEqual(hard_reset_client.packet_id, 0x00010203) def test_compose(self): self.assertEqual(self.hard_reset_client.compose(), self.hard_reset_client_bytes) class TestOpenVpnPacketHardResetServerV2(unittest.TestCase): def setUp(self): self.hard_reset_server_dict = collections.OrderedDict([ ('opcode', b'\x40'), ('session_id', b'\x00\x01\x02\x03\x04\x05\x06\x07'), ('packet_id_array_length', b'\x01'), ('packet_id_array', b'\x04\x05\x06\x07'), ('remote_session_id', b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f'), ('packet_id', b'\x00\x01\x02\x03'), ]) self.hard_reset_server_bytes = b''.join(self.hard_reset_server_dict.values()) self.hard_reset_server = OpenVpnPacketHardResetServerV2( session_id=0x0001020304050607, packet_id_array=[0x04050607], remote_session_id=0x08090a0b0c0d0e0f, packet_id=0x00010203, ) def test_parse(self): hard_reset_server = OpenVpnPacketVariant.parse_exact_size(self.hard_reset_server_bytes) self.assertEqual(hard_reset_server.session_id, 0x0001020304050607) self.assertEqual(hard_reset_server.packet_id_array, [0x04050607]) self.assertEqual(hard_reset_server.remote_session_id, 0x08090a0b0c0d0e0f) self.assertEqual(hard_reset_server.packet_id, 0x00010203) def test_compose(self): self.assertEqual(self.hard_reset_server.compose(), self.hard_reset_server_bytes) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_postgresql.py000066400000000000000000000033231524413560000274500ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import NotEnoughData from cryptoparser.tls.postgresql import SslRequest, Sync class TestSslRequest(unittest.TestCase): def setUp(self): self.ssl_request = SslRequest() self.ssl_request_bytes = b'\x00\x00\x00\x08\x04\xd2\x16\x2f' def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: SslRequest.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, SslRequest.MESSAGE_SIZE) with self.assertRaises(InvalidValue) as context_manager: SslRequest.parse_exact_size(b'\x00\x00\x00\x04\x01\x02\x03\x04') self.assertEqual(context_manager.exception.value, 4) with self.assertRaises(InvalidValue) as context_manager: SslRequest.parse_exact_size(b'\x00\x00\x00\x08\x01\x02\x03\x04') self.assertEqual(context_manager.exception.value, 0x01020304) def test_parse(self): SslRequest.parse_exact_size(self.ssl_request_bytes) def test_compose(self): self.assertEqual(self.ssl_request.compose(), self.ssl_request_bytes) class TestSync(unittest.TestCase): def setUp(self): self.ssl_request = Sync() self.ssl_request_bytes = b'S' def test_error(self): with self.assertRaises(InvalidValue) as context_manager: Sync.parse_exact_size(b'X') self.assertEqual(context_manager.exception.value, b'X') def test_parse(self): Sync.parse_exact_size(self.ssl_request_bytes) def test_compose(self): self.assertEqual(self.ssl_request.compose(), self.ssl_request_bytes) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_rdp.py000066400000000000000000000206731524413560000260410ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest import collections from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import NotEnoughData, InvalidType from cryptoparser.tls.rdp import ( COTPConnectionConfirm, COTPConnectionRequest, RDPNegotiationRequest, RDPNegotiationRequestFlags, RDPNegotiationResponse, RDPNegotiationResponseFlags, RDPProtocol, TPKT, ) class TestTPKT(unittest.TestCase): def setUp(self): self.tpkt_dict = collections.OrderedDict([ ('version', b'\x03'), ('reserved', b'\x00'), ('packet_length', b'\x00\x14'), ('message', b'\x00\x01\x02\x03\x04\x05\x06\x07' + b'\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' + b'') ]) self.tpkt_bytes = b''.join(self.tpkt_dict.values()) self.tpkt = TPKT( version=3, message=b'\x00\x01\x02\x03\x04\x05\x06\x07\x08\x09\x0a\x0b\x0c\x0d\x0e\x0f' ) def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: TPKT.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, TPKT.HEADER_SIZE) with self.assertRaises(InvalidValue) as context_manager: TPKT.parse_exact_size(b'\x01\x00\x00\x00') self.assertEqual(context_manager.exception.value, 1) with self.assertRaises(NotEnoughData) as context_manager: TPKT.parse_exact_size(b'\x03\x00\x00\x05') self.assertEqual(context_manager.exception.bytes_needed, 1) def test_parse(self): tpkt = TPKT.parse_exact_size(self.tpkt_bytes) self.assertEqual(tpkt.version, 3) self.assertEqual(tpkt.message, self.tpkt_dict['message']) def test_compose(self): self.assertEqual(self.tpkt.compose(), self.tpkt_bytes) class TestCOTPConnectionRequest(unittest.TestCase): def setUp(self): self.cotp_connection_request_dict = collections.OrderedDict([ ('length', b'\x0e'), ('type', b'\xe0'), ('src_ref', b'\x01\x02'), ('dst_ref', b'\x03\x04'), ('class_option', b'\x00'), ('user_data', b'\x00\x01\x02\x03\x04\x05\x06\x07'), ]) self.cotp_connection_request_bytes = b''.join(self.cotp_connection_request_dict.values()) self.cotp_connection_request = COTPConnectionRequest( src_ref=0x0102, dst_ref=0x0304, class_option=0x00, user_data=b'\x00\x01\x02\x03\x04\x05\x06\x07' ) def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: COTPConnectionRequest.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, COTPConnectionRequest.HEADER_SIZE) with self.assertRaises(NotEnoughData) as context_manager: COTPConnectionRequest.parse_exact_size(b'\x07\xe0\x00\x00\x00\x00\x00') self.assertEqual(context_manager.exception.bytes_needed, 1) with self.assertRaises(InvalidType) as context_manager: COTPConnectionRequest.parse_exact_size(b'\x07\xf0\x00\x00\x00\x00\x00\x00') with self.assertRaises(InvalidValue) as context_manager: COTPConnectionRequest.parse_exact_size(b'\x07\xe0\x00\x00\x00\x00\x01\x00') self.assertEqual(context_manager.exception.value, 1) def test_parse(self): cotp_connection_request = COTPConnectionRequest.parse_exact_size(self.cotp_connection_request_bytes) self.assertEqual(cotp_connection_request.src_ref, self.cotp_connection_request.src_ref) self.assertEqual(cotp_connection_request.dst_ref, self.cotp_connection_request.dst_ref) self.assertEqual(cotp_connection_request.class_option, self.cotp_connection_request.class_option) self.assertEqual(cotp_connection_request.user_data, self.cotp_connection_request.user_data) def test_compose(self): self.assertEqual(self.cotp_connection_request.compose(), self.cotp_connection_request_bytes) class TestCOTPConnectionConfirm(unittest.TestCase): def setUp(self): self.cotp_connection_confirm_dict = collections.OrderedDict([ ('length', b'\x0e'), ('type', b'\xd0'), ('src_ref', b'\x01\x02'), ('dst_ref', b'\x03\x04'), ('class_option', b'\x00'), ('user_data', b'\x00\x01\x02\x03\x04\x05\x06\x07'), ]) self.cotp_connection_confirm_bytes = b''.join(self.cotp_connection_confirm_dict.values()) self.cotp_connection_confirm = COTPConnectionConfirm( src_ref=0x0102, dst_ref=0x0304, class_option=0x00, user_data=b'\x00\x01\x02\x03\x04\x05\x06\x07' ) def test_parse(self): cotp_connection_confirm = COTPConnectionConfirm.parse_exact_size(self.cotp_connection_confirm_bytes) self.assertEqual(cotp_connection_confirm.src_ref, self.cotp_connection_confirm.src_ref) self.assertEqual(cotp_connection_confirm.dst_ref, self.cotp_connection_confirm.dst_ref) self.assertEqual(cotp_connection_confirm.class_option, self.cotp_connection_confirm.class_option) self.assertEqual(cotp_connection_confirm.user_data, self.cotp_connection_confirm.user_data) def test_compose(self): self.assertEqual(self.cotp_connection_confirm.compose(), self.cotp_connection_confirm_bytes) class TestRDPNegotiationRequest(unittest.TestCase): def setUp(self): self.rdp_negotiation_request_dict = collections.OrderedDict([ ('type', b'\x01'), ('flags', b'\x03'), ('length', b'\x08\x00'), ('protocol', b'\x03\x00\x00\x00'), ]) self.rdp_negotiation_request_bytes = b''.join(self.rdp_negotiation_request_dict.values()) self.rdp_negotiation_request = RDPNegotiationRequest( flags={ RDPNegotiationRequestFlags.RESTRICTED_ADMIN_MODE_REQUIRED, RDPNegotiationRequestFlags.REDIRECTED_AUTHENTICATION_MODE_REQUIRED }, protocol={RDPProtocol.SSL, RDPProtocol.HYBRID} ) def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: RDPNegotiationRequest.parse_exact_size(b'\x01') self.assertEqual(context_manager.exception.bytes_needed, RDPNegotiationRequest.PACKET_LENGTH - 1) with self.assertRaises(InvalidType) as context_manager: RDPNegotiationRequest.parse_exact_size(b'\x02\x03\x08\x00\x00\x00\x00\x00') with self.assertRaises(NotEnoughData) as context_manager: RDPNegotiationRequest.parse_exact_size(b'\x01\x03\x08\x00\x00\x00') self.assertEqual(context_manager.exception.bytes_needed, 2) with self.assertRaises(InvalidValue) as context_manager: RDPNegotiationRequest.parse_exact_size(b'\x01\x03\x02\x00\x00\x00\x00\x00') self.assertEqual(context_manager.exception.value, 2) def test_parse(self): rdp_negotiation_request = RDPNegotiationRequest.parse_exact_size(self.rdp_negotiation_request_bytes) self.assertEqual(rdp_negotiation_request.flags, self.rdp_negotiation_request.flags) self.assertEqual(rdp_negotiation_request.protocol, self.rdp_negotiation_request.protocol) def test_compose(self): self.assertEqual(self.rdp_negotiation_request.compose(), self.rdp_negotiation_request_bytes) class TestRDPNegotiationResponse(unittest.TestCase): def setUp(self): self.rdp_negotiation_request_dict = collections.OrderedDict([ ('type', b'\x02'), ('flags', b'\x03'), ('length', b'\x08\x00'), ('protocol', b'\x03\x00\x00\x00'), ]) self.rdp_negotiation_request_bytes = b''.join(self.rdp_negotiation_request_dict.values()) self.rdp_negotiation_request = RDPNegotiationResponse( flags={ RDPNegotiationResponseFlags.EXTENDED_CLIENT_DATA_SUPPORTED, RDPNegotiationResponseFlags.DYNVC_GFX_PROTOCOL_SUPPORTED }, protocol={RDPProtocol.SSL, RDPProtocol.HYBRID} ) def test_parse(self): rdp_negotiation_request = RDPNegotiationResponse.parse_exact_size(self.rdp_negotiation_request_bytes) self.assertEqual(rdp_negotiation_request.flags, self.rdp_negotiation_request.flags) self.assertEqual(rdp_negotiation_request.protocol, self.rdp_negotiation_request.protocol) def test_compose(self): self.assertEqual(self.rdp_negotiation_request.compose(), self.rdp_negotiation_request_bytes) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_record.py000066400000000000000000000144101524413560000265220ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptodatahub.common.exception import InvalidValue from cryptoparser.common.exception import NotEnoughData from cryptoparser.tls.record import TlsRecord, SslRecord from cryptoparser.tls.subprotocol import TlsContentType, TlsSubprotocolMessageBase from cryptoparser.tls.version import TlsVersion, TlsProtocolVersion from cryptoparser.tls.subprotocol import TlsAlertMessage, TlsAlertLevel, TlsAlertDescription from cryptoparser.tls.subprotocol import SslErrorMessage, SslErrorType, SslMessageType class TestTlsSubprotocolMessageBase(unittest.TestCase): def test_error(self): error_message = 'Can\'t instantiate abstract class TlsSubprotocolMessageBase' with self.assertRaisesRegex(BaseException, error_message) as context_manager: TlsSubprotocolMessageBase() # pylint: disable=abstract-class-instantiated self.assertTrue(isinstance(context_manager.exception, (TypeError, AssertionError))) class TestTlsRecord(unittest.TestCase): def setUp(self): self.test_message = TlsAlertMessage( level=TlsAlertLevel.FATAL, description=TlsAlertDescription.HANDSHAKE_FAILURE ) self.test_record = TlsRecord( fragment=self.test_message.compose(), protocol_version=TlsProtocolVersion(TlsVersion.TLS1), content_type=TlsContentType.ALERT, ) self.test_record_bytes = bytes( b'\x15' + # type = ALERT b'\x03\x01' + # version = TLS1 b'\x00\x02' + # length = 2 b'\x02' + # level = FATAL b'\x28' + # description = HANDSHAKE_FAILURE b'' ) def test_error(self): with self.assertRaises(NotEnoughData) as context_manager: TlsRecord.parse_exact_size(b'') self.assertEqual(context_manager.exception.bytes_needed, TlsRecord.HEADER_SIZE) with self.assertRaisesRegex(InvalidValue, '0xff is not a valid TlsContentType'): TlsRecord.parse_exact_size( b'\xff' + # type = INVALID b'\x03\x03' + # version = TLS 1.2 b'\x00\x00' + # length = 0 b'' ) with self.assertRaises(NotEnoughData) as context_manager: TlsRecord.parse_exact_size( b'\x15' + # type = alert b'\x03\x03' + # version = TLS 1.2 b'\x00\x01' + # length = 1 (alert message is 2 bytes!) b'' ) self.assertEqual(context_manager.exception.bytes_needed, 1) TlsRecord.parse_exact_size( b'\x15' + # type = alert b'\x03\x03' + # version = TLS 1.2 b'\x00\x02' + # length = 2 b'\x02\x28' ) with self.assertRaises(NotEnoughData) as context_manager: TlsRecord.parse_exact_size( b'\x16' + # type = handshake b'\x03\x03' + # version = TLS 1.2 b'\x00\x02' + # length = 2 (handshake message is at least 4 bytes!) b'' ) self.assertEqual(context_manager.exception.bytes_needed, 2) def test_parse(self): record = TlsRecord.parse_exact_size(self.test_record_bytes) self.assertEqual(record.fragment, self.test_message.compose()) self.assertEqual(record.protocol_version, TlsProtocolVersion(TlsVersion.TLS1)) self.assertEqual(record.content_type, TlsContentType.ALERT) def test_compose(self): self.assertEqual( self.test_record.compose(), self.test_record_bytes ) class TestSslRecord(unittest.TestCase): def setUp(self): self.test_message = SslErrorMessage( error_type=SslErrorType.NO_CIPHER_ERROR ) self.test_record = SslRecord( message=self.test_message ) self.test_record_bytes = bytes( b'\x80\x03' + # length = 3 b'\x00' + # message_type = ERROR b'\x00\x01' + # error_type = NO_CIPHER_ERROR b'' ) def test_error(self): with self.assertRaisesRegex(InvalidValue, '0xff is not a valid SslMessageType'): SslRecord.parse_exact_size( b'\x80\x00' + # length = 0 b'\xff' + # type = INVALID b'' ) with self.assertRaises(NotEnoughData) as context_manager: SslRecord.parse_exact_size( b'\x80\x03' + # length = 3 (with length bytes) b'' ) self.assertEqual(context_manager.exception.bytes_needed, 3) with self.assertRaises(InvalidValue) as context_manager: SslRecord.parse_exact_size( b'\x80\x03' + # length = 3 b'\x00' + # message_type = ERROR b'\x00\xff' + # error_type = INVALID b'' ) with self.assertRaises(NotEnoughData) as context_manager: SslRecord.parse_exact_size( b'\x80\x03' + # length = 3 b'\x00' + # type = ERROR b'' ) self.assertEqual(context_manager.exception.bytes_needed, 2) with self.assertRaises(NotEnoughData) as context_manager: SslRecord.parse_exact_size( b'\x81\x03' + # length = 256 + 3 b'' ) self.assertEqual(context_manager.exception.bytes_needed, 256 + 3 - 0) with self.assertRaises(NotEnoughData) as context_manager: SslRecord.parse_exact_size( b'\x01\x03\x01' + # length = 256 + 3 b'' ) self.assertEqual(context_manager.exception.bytes_needed, 256 + 3 - 0) def test_setter(self): record = SslRecord(self.test_message) record.message = self.test_message record.protocol_version = TlsVersion.SSL2 def test_parse(self): record = SslRecord.parse_exact_size(self.test_record_bytes) self.assertEqual( record.message, self.test_message ) self.assertEqual(record.content_type, SslMessageType.ERROR) def test_compose(self): self.assertEqual( self.test_record.compose(), self.test_record_bytes ) cryptoparser-v1.6.0-e68cede0561acb18bbc8d10453335e03b3d63224/test/tls/test_version.py000066400000000000000000000153051524413560000267350ustar00rootroot00000000000000# SPDX-License-Identifier: MPL-2.0 import unittest from cryptodatahub.common.exception import InvalidValue from cryptodatahub.common.grade import Grade from cryptoparser.common.exception import NotEnoughData from cryptoparser.tls.version import TlsVersion, TlsProtocolVersion class TestTlsProtocolVersion(unittest.TestCase): def test_parse(self): parsable = b'\x03\xff' expected_error_message = ' is not a valid TlsVersionFactory' with self.assertRaisesRegex(InvalidValue, expected_error_message): # pylint: disable=expression-not-assigned TlsProtocolVersion.parse_exact_size(parsable) expected_error_message = ' is not a valid TlsVersionFactory' with self.assertRaisesRegex(InvalidValue, expected_error_message): # pylint: disable=expression-not-assigned TlsProtocolVersion.parse_exact_size(b'\x8f\x00') with self.assertRaises(NotEnoughData) as context_manager: TlsProtocolVersion.parse_exact_size(b'\xff') self.assertEqual(context_manager.exception.bytes_needed, 1) def test_compose(self): self.assertEqual( b'\x03\x03', TlsProtocolVersion(TlsVersion.TLS1_2).compose() ) self.assertEqual( b'\x7f\x12', TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_18).compose() ) def test_lt(self): self.assertLess( TlsProtocolVersion(TlsVersion.SSL2), TlsProtocolVersion(TlsVersion.SSL3) ) self.assertGreater( TlsProtocolVersion(TlsVersion.SSL3), TlsProtocolVersion(TlsVersion.SSL2) ) self.assertLess( TlsProtocolVersion(TlsVersion.SSL3), TlsProtocolVersion(TlsVersion.TLS1) ) self.assertGreater( TlsProtocolVersion(TlsVersion.TLS1), TlsProtocolVersion(TlsVersion.SSL3) ) self.assertLess( TlsProtocolVersion(TlsVersion.TLS1_1), TlsProtocolVersion(TlsVersion.TLS1_2) ) self.assertGreater( TlsProtocolVersion(TlsVersion.TLS1_2), TlsProtocolVersion(TlsVersion.TLS1_1) ) self.assertLess( TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_1), TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_2) ) self.assertGreater( TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_2), TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_1) ) self.assertLess( TlsProtocolVersion(TlsVersion.TLS1_2), TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_0) ) self.assertGreater( TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_0), TlsProtocolVersion(TlsVersion.TLS1_2) ) self.assertLess( TlsProtocolVersion(TlsVersion.TLS1_2), TlsProtocolVersion(TlsVersion.TLS1_3) ) self.assertGreater( TlsProtocolVersion(TlsVersion.TLS1_3), TlsProtocolVersion(TlsVersion.TLS1_2) ) self.assertLess( TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_28), TlsProtocolVersion(TlsVersion.TLS1_3) ) self.assertGreater( TlsProtocolVersion(TlsVersion.TLS1_3), TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_28) ) self.assertLess( TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_0), TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_28) ) self.assertGreater( TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_28), TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_0) ) def test_set(self): self.assertEqual( 2, len(set([ TlsProtocolVersion(TlsVersion.TLS1_1), TlsProtocolVersion(TlsVersion.TLS1_2) ])) ) self.assertEqual( 1, len(set([ TlsProtocolVersion(TlsVersion.TLS1_1), TlsProtocolVersion(TlsVersion.TLS1_1) ])) ) self.assertEqual( 2, len(set([ TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_1), TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_2) ])) ) self.assertEqual( 1, len(set([ TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_1), TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_1) ])) ) def test_as_json(self): self.assertEqual(TlsProtocolVersion(TlsVersion.SSL3).as_json(), '\"ssl3\"') self.assertEqual(TlsProtocolVersion(TlsVersion.TLS1).as_json(), '\"tls1\"') self.assertEqual(TlsProtocolVersion(TlsVersion.TLS1_2).as_json(), '\"tls1_2\"') self.assertEqual(TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_24).as_json(), '\"tls1_3_draft_24\"') self.assertEqual( TlsProtocolVersion(TlsVersion.TLS1_3_GOOGLE_EXPERIMENT_2).as_json(), '\"tls1_3_google_experiment_2\"' ) def test_as_markdown(self): self.assertEqual(TlsProtocolVersion(TlsVersion.SSL3).as_markdown(), 'SSL 3.0') self.assertEqual(TlsProtocolVersion(TlsVersion.TLS1).as_markdown(), 'TLS 1.0') self.assertEqual(TlsProtocolVersion(TlsVersion.TLS1_2).as_markdown(), 'TLS 1.2') self.assertEqual(TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_24).as_markdown(), 'TLS 1.3 Draft 24') self.assertEqual( TlsProtocolVersion(TlsVersion.TLS1_3_GOOGLE_EXPERIMENT_2).as_markdown(), 'TLS 1.3 Google Experiment 2' ) def test_str(self): self.assertEqual(str(TlsProtocolVersion(TlsVersion.SSL2)), 'SSL 2.0') self.assertEqual(str(TlsProtocolVersion(TlsVersion.SSL3)), 'SSL 3.0') self.assertEqual(str(TlsProtocolVersion(TlsVersion.TLS1)), 'TLS 1.0') self.assertEqual(str(TlsProtocolVersion(TlsVersion.TLS1_2)), 'TLS 1.2') self.assertEqual(str(TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_24)), 'TLS 1.3 Draft 24') self.assertEqual(str(TlsProtocolVersion(TlsVersion.TLS1_3_GOOGLE_EXPERIMENT_2)), 'TLS 1.3 Google Experiment 2') def test_grade(self): self.assertEqual(TlsProtocolVersion(TlsVersion.SSL2).grade, Grade.INSECURE) self.assertEqual(TlsProtocolVersion(TlsVersion.SSL3).grade, Grade.INSECURE) self.assertEqual(TlsProtocolVersion(TlsVersion.TLS1).grade, Grade.DEPRECATED) self.assertEqual(TlsProtocolVersion(TlsVersion.TLS1_1).grade, Grade.DEPRECATED) self.assertEqual(TlsProtocolVersion(TlsVersion.TLS1_3_DRAFT_28).grade, Grade.DEPRECATED) self.assertEqual(TlsProtocolVersion(TlsVersion.TLS1_2).grade, Grade.SECURE) self.assertEqual(TlsProtocolVersion(TlsVersion.TLS1_3).grade, Grade.SECURE)