diff --git a/.github/workflows/wheels.yml b/.github/workflows/wheels.yml index 25b590e..c1e11a5 100644 --- a/.github/workflows/wheels.yml +++ b/.github/workflows/wheels.yml @@ -106,7 +106,7 @@ jobs: CIBW_MANYLINUX_X86_64_IMAGE: ${{ matrix.manylinux_image }} CIBW_TEST_COMMAND: python -m unittest discover {project}/tests - - uses: actions/upload-artifact@ea165f8d65b6e75b540449e92b4886f43607fa02 # v4.6.2 + - uses: actions/upload-artifact@330a01c490aca151604b8cf639adc76d48f6c5d4 # v5.0.0 with: name: cibw-wheels-${{ matrix.os }}-${{ matrix.python }}-${{ strategy.job-index }} path: ./wheelhouse/*.whl diff --git a/.gitignore b/.gitignore index 3e9d927..5ca1eb0 100644 --- a/.gitignore +++ b/.gitignore @@ -58,4 +58,7 @@ venv __pycache__ zzz/ .cache/ -compile_commands.json \ No newline at end of file +compile_commands.json + +# 运行示例/测试生成的密钥产物 +*.pem \ No newline at end of file diff --git a/GmSSL b/GmSSL index d655c06..7c9f029 160000 --- a/GmSSL +++ b/GmSSL @@ -1 +1 @@ -Subproject commit d655c06b3a6b0fe8cff900f293bf0e5aac6eb0a2 +Subproject commit 7c9f02904ef33e59c87b4f16621cc8fd434e7579 diff --git a/README.md b/README.md index de069ab..694cd4c 100644 --- a/README.md +++ b/README.md @@ -7,7 +7,7 @@ python wrapper (C extension) of [GmSSL](https://github.com/guanzhi/GmSSL) -使用的版本是 [GmSSL-3.1.0](https://github.com/guanzhi/GmSSL/releases/tag/v3.1.0) +使用的版本是 [GmSSL-3.2.0](https://github.com/guanzhi/GmSSL/releases/tag/v3.2.0) ## 安装 @@ -293,3 +293,11 @@ print('plaintext', plaintext) [SM9](docs/sm9.md) 如果要查看所有可用的 API ,可以看 [gmsslext.pyi](gmssl_pyx/gmsslext.pyi) 文件。 + +## 注意 + +GmSSL v3.2.0 修改了 SM9 PEM 文件头(即 BEGIN 与 END 行中的标识字符串),变更如下: + +"ENCRYPTED SM9 ENC MASTER KEY" => "ENCRYPTED PRIVATE KEY" + +"ENCRYPTED SM9 ENC PRIVATE KEY" => "ENCRYPTED PRIVATE KEY" diff --git a/gmssl_pyx/gmsslext.pyi b/gmssl_pyx/gmsslext.pyi index 184763c..17069a4 100644 --- a/gmssl_pyx/gmsslext.pyi +++ b/gmssl_pyx/gmsslext.pyi @@ -58,7 +58,7 @@ def sm2_verify_sm3_digest(public_key: bytes, digest: bytes, signature: bytes) -> """使用 SM2 验证 SM3 摘要和签名数据 Args: - public_key: 64 字节的私钥 + public_key: 64 字节的公钥 digest: SM3 摘要数据,长度为 32 字节 signature: 签名数据,编码格式为 ASN.1 DER ,模式为 rs @@ -223,9 +223,11 @@ def sm4_gcm_decrypt( iv: 初始化向量,也被叫做 nonce ,1 <= 长度 <= 64 aad: 附加数据,也被叫做 associated_data ciphertext: 密文数据 - tag: 标签 + tag: 标签,12 <= 长度 <= 16 Returns: 明文数据 + + Raises: InvalidValueError: tag 长度非法或 tag 不匹配(认证失败) """ ... @@ -259,7 +261,7 @@ class SM9PrivateKey: password: 密码 data: 加密的 ASN.1 DER 编码数据 - Return: 私钥 + Returns: 私钥 """ ... @@ -347,7 +349,7 @@ class SM9MasterKey: password: 密码 data: 加密的 ASN.1 DER 编码数据 - Return: 主密钥 + Returns: 主密钥 """ ... @@ -363,7 +365,7 @@ class SM9MasterKey: password: 密码 filepath: 文件路径 - Return: 主密钥 + Returns: 主密钥 """ def encrypt_to_pem(self, password: str, filepath: str) -> None: diff --git a/gmssl_pyx/gmsslext_sm9.c b/gmssl_pyx/gmsslext_sm9.c index 53158ec..47fb8d6 100644 --- a/gmssl_pyx/gmsslext_sm9.c +++ b/gmssl_pyx/gmsslext_sm9.c @@ -8,12 +8,33 @@ #include #include +#include #include "gmssl/sm9.h" +#include "gmssl/pem.h" +#include "gmssl/asn1.h" +#include "gmssl/mem.h" #include "gmsslext.h" #include "gmsslext_sm9.h" +static bool pem_starts_with(FILE *fp, const char *prefix) { + if (!fp || !prefix) return false; + + size_t len = strlen(prefix); + if (len == 0) return true; + + char buffer[128]; // 足够大 + size_t read_len = fread(buffer, 1, len, fp); + + // 重置文件指针到开头(重要!) + rewind(fp); + + if (read_len != len) return false; + + return memcmp(buffer, prefix, len) == 0; +} + /* * SM9 wrapper */ @@ -24,6 +45,7 @@ typedef struct { } SM9PrivateKeyObject; static void SM9PrivateKey_dealloc(SM9PrivateKeyObject *self) { + gmssl_secure_clear(&self->key, sizeof(self->key)); Py_TYPE(self)->tp_free((PyObject *)self); } @@ -181,6 +203,20 @@ static PyObject *SM9PrivateKey_encrypt_to_der(SM9PrivateKeyObject *self, return Py_BuildValue("y#", (char *)buf, (Py_ssize_t)len); } +static int sm9_enc_key_info_decrypt_from_pem_v3_1_1(SM9_ENC_KEY *key, const char *pass, FILE *fp) +{ + uint8_t buf[SM9_MAX_ENCED_PRIVATE_KEY_INFO_SIZE]; + const uint8_t *cp = buf; + size_t len; + + if (pem_read(fp, PEM_SM9_ENC_PRIVATE_KEY_V3_1_1, buf, &len, sizeof(buf)) != 1 + || sm9_enc_key_info_decrypt_from_der(key, pass, &cp, &len) != 1 + || asn1_length_is_zero(len) != 1) { + return -1; + } + return 1; +} + static PyObject *SM9PrivateKey_decrypt_from_pem(PyTypeObject *type, PyObject *args, PyObject *keywds) { @@ -215,7 +251,13 @@ static PyObject *SM9PrivateKey_decrypt_from_pem(PyTypeObject *type, fclose(fp); return NULL; } - ret = sm9_enc_key_info_decrypt_from_pem(&self->key, password, fp); + char begin_line[80]; + snprintf(begin_line, sizeof(begin_line), "-----BEGIN %s-----", PEM_SM9_ENC_PRIVATE_KEY_V3_1_1); + if (pem_starts_with(fp, begin_line)) { + ret = sm9_enc_key_info_decrypt_from_pem_v3_1_1(&self->key, password, fp); + } else { + ret = sm9_enc_key_info_decrypt_from_pem(&self->key, password, fp); + } if (ret != GMSSL_INNER_OK) { Py_DECREF(self); fclose(fp); @@ -365,6 +407,7 @@ typedef struct { } SM9MasterPublicKeyObject; static void SM9MasterPublicKey_dealloc(SM9MasterPublicKeyObject *self) { + gmssl_secure_clear(&self->master_public, sizeof(self->master_public)); Py_TYPE(self)->tp_free((PyObject *)self); } @@ -608,6 +651,7 @@ typedef struct { } SM9MasterKeyObject; static void SM9MasterKey_dealloc(SM9MasterKeyObject *self) { + gmssl_secure_clear(&self->master, sizeof(self->master)); Py_TYPE(self)->tp_free((PyObject *)self); } @@ -786,6 +830,20 @@ static PyObject *SM9MasterKey_encrypt_to_der(SM9MasterKeyObject *self, return Py_BuildValue("y#", buf, (Py_ssize_t)len); } +static int sm9_enc_master_key_info_decrypt_from_pem_v3_1_1(SM9_ENC_MASTER_KEY *msk, const char *pass, FILE *fp) +{ + uint8_t buf[SM9_MAX_ENCED_PRIVATE_KEY_INFO_SIZE]; + const uint8_t *cp = buf; + size_t len; + + if (pem_read(fp, PEM_SM9_ENC_MASTER_KEY_V3_1_1, buf, &len, sizeof(buf)) != 1 + || sm9_enc_master_key_info_decrypt_from_der(msk, pass, &cp, &len) != 1 + || asn1_length_is_zero(len) != 1) { + return -1; + } + return 1; +} + static PyObject *SM9MasterKey_decrypt_from_pem(PyTypeObject *type, PyObject *args, PyObject *keywds) { @@ -818,8 +876,14 @@ static PyObject *SM9MasterKey_decrypt_from_pem(PyTypeObject *type, fclose(fp); return NULL; } - int ret = - sm9_enc_master_key_info_decrypt_from_pem(&self->master, password, fp); + int ret; + char begin_line[80]; + snprintf(begin_line, sizeof(begin_line), "-----BEGIN %s-----", PEM_SM9_ENC_MASTER_KEY_V3_1_1); + if (pem_starts_with(fp, begin_line)) { + ret = sm9_enc_master_key_info_decrypt_from_pem_v3_1_1(&self->master, password, fp); + } else { + ret = sm9_enc_master_key_info_decrypt_from_pem(&self->master, password, fp); + } if (ret != GMSSL_INNER_OK) { Py_DECREF(self); fclose(fp); diff --git a/gmssl_pyx/gmsslext_sm9.h b/gmssl_pyx/gmsslext_sm9.h index 0a53e05..f15d01d 100644 --- a/gmssl_pyx/gmsslext_sm9.h +++ b/gmssl_pyx/gmsslext_sm9.h @@ -7,4 +7,7 @@ extern PyTypeObject GmsslextSM9PrivateKeyType; extern PyTypeObject GmsslextSM9MasterPublicKeyType; extern PyTypeObject GmsslextSM9MasterKeyType; +#define PEM_SM9_ENC_MASTER_KEY_V3_1_1 "ENCRYPTED SM9 ENC MASTER KEY" +#define PEM_SM9_ENC_PRIVATE_KEY_V3_1_1 "ENCRYPTED SM9 ENC PRIVATE KEY" + #endif // GMSSL_PYX_GMSSLEXT_SM9_H diff --git a/gmssl_pyx/gmsslmodule.c b/gmssl_pyx/gmsslmodule.c index b22461d..76386f3 100644 --- a/gmssl_pyx/gmsslmodule.c +++ b/gmssl_pyx/gmsslmodule.c @@ -4,9 +4,12 @@ #define PY_SSIZE_T_CLEAN #include +#include +#include #include "gmssl/rand.h" #include "gmssl/sm2.h" +#include "gmssl/sm2_z256.h" #include "gmssl/sm3.h" #include "gmssl/sm4.h" @@ -31,9 +34,16 @@ static PyObject *gmsslext_sm2_key_generate(PyObject *self, } // 整数字面量不是 Py_ssize_t 类型,需要强制转换,不然 Windows 会报错 // MemoryError - return Py_BuildValue("y#y#", (const char *)&sm2_key.public_key, - (Py_ssize_t)64, (const char *)sm2_key.private_key, - (Py_ssize_t)32); + uint8_t public_key[64]; + uint8_t private_key[32]; + if (sm2_z256_point_to_bytes(&sm2_key.public_key, public_key) == -1) { + PyErr_SetString(GmsslInnerError, + "libgmssl inner error in sm2_z256_point_to_bytes"); + return NULL; + } + sm2_z256_to_bytes(sm2_key.private_key, private_key); + return Py_BuildValue("y#y#", (const char *)public_key, (Py_ssize_t)64, + (const char *)private_key, (Py_ssize_t)32); } static PyObject *gmsslext_sm2_encrypt(PyObject *self, PyObject *args, @@ -62,7 +72,13 @@ static PyObject *gmsslext_sm2_encrypt(PyObject *self, PyObject *args, PyErr_SetString(InvalidValueError, "plaintext length not support"); return NULL; } - ret = sm2_key_set_public_key(&sm2_key, (SM2_POINT *)public_key); + SM2_Z256_POINT pub; + if (sm2_z256_point_from_bytes(&pub, (const uint8_t *)public_key) != + GMSSL_INNER_OK) { + PyErr_SetString(InvalidValueError, "invalid public key"); + return NULL; + } + ret = sm2_key_set_public_key(&sm2_key, &pub); if (ret != GMSSL_INNER_OK) { PyErr_SetString(InvalidValueError, "invalid public key"); return NULL; @@ -102,7 +118,9 @@ static PyObject *gmsslext_sm2_decrypt(PyObject *self, PyObject *args, PyErr_SetString(InvalidValueError, "ciphertext length not support"); return NULL; } - ret = sm2_key_set_private_key(&sm2_key, (uint8_t *)private_key); + sm2_z256_t priv; + sm2_z256_from_bytes(priv, (const uint8_t *)private_key); + ret = sm2_key_set_private_key(&sm2_key, priv); if (ret != GMSSL_INNER_OK) { PyErr_SetString(InvalidValueError, "invalid private key"); return NULL; @@ -141,7 +159,9 @@ static PyObject *gmsslext_sm2_sign_sm3_digest(PyObject *self, PyObject *args, PyErr_SetString(InvalidValueError, "expected 32bytes sm3 digest"); return NULL; } - ret = sm2_key_set_private_key(&sm2_key, (uint8_t *)private_key); + sm2_z256_t priv; + sm2_z256_from_bytes(priv, (const uint8_t *)private_key); + ret = sm2_key_set_private_key(&sm2_key, priv); if (ret != GMSSL_INNER_OK) { PyErr_SetString(InvalidValueError, "invalid private key"); return NULL; @@ -181,7 +201,13 @@ static PyObject *gmsslext_sm2_verify_sm3_digest(PyObject *self, PyObject *args, PyErr_SetString(InvalidValueError, "invalid sm3 digest"); return NULL; } - ret = sm2_key_set_public_key(&sm2_key, (SM2_POINT *)public_key); + SM2_Z256_POINT pub; + if (sm2_z256_point_from_bytes(&pub, (const uint8_t *)public_key) != + GMSSL_INNER_OK) { + PyErr_SetString(InvalidValueError, "invalid public key"); + return NULL; + } + ret = sm2_key_set_public_key(&sm2_key, &pub); if (ret != GMSSL_INNER_OK) { PyErr_SetString(InvalidValueError, "invalid public key"); return NULL; @@ -247,19 +273,27 @@ static PyObject *gmsslext_sm2_sign(PyObject *self, PyObject *args, "invalid public_key or private_key length"); return NULL; } - if (message_length <= 0) { - PyErr_SetString(InvalidValueError, "empty message"); + if (message_length < 0) { + PyErr_SetString(InvalidValueError, "invalid message"); return NULL; } SM2_KEY sm2_key; SM2_SIGN_CTX sign_ctx; - ret = sm2_key_set_public_key(&sm2_key, (SM2_POINT *)public_key); + SM2_Z256_POINT pub; + if (sm2_z256_point_from_bytes(&pub, (const uint8_t *)public_key) != + GMSSL_INNER_OK) { + PyErr_SetString(InvalidValueError, "invalid public_key"); + return NULL; + } + sm2_z256_t priv; + sm2_z256_from_bytes(priv, (const uint8_t *)private_key); + ret = sm2_key_set_public_key(&sm2_key, &pub); if (ret != GMSSL_INNER_OK) { PyErr_SetString(InvalidValueError, "invalid public_key"); return NULL; } - ret = sm2_key_set_private_key(&sm2_key, (uint8_t *)private_key); + ret = sm2_key_set_private_key(&sm2_key, priv); if (ret != GMSSL_INNER_OK) { PyErr_SetString(InvalidValueError, "invalid private_key"); return NULL; @@ -341,14 +375,20 @@ static PyObject *gmsslext_sm2_verify(PyObject *self, PyObject *args, PyErr_SetString(InvalidValueError, "empty signature"); return NULL; } - if (message_length <= 0) { - PyErr_SetString(InvalidValueError, "empty message"); + if (message_length < 0) { + PyErr_SetString(InvalidValueError, "invalid message"); return NULL; } SM2_KEY sm2_key; - SM2_SIGN_CTX sign_ctx; - ret = sm2_key_set_public_key(&sm2_key, (SM2_POINT *)public_key); + SM2_VERIFY_CTX sign_ctx; + SM2_Z256_POINT pub; + if (sm2_z256_point_from_bytes(&pub, (const uint8_t *)public_key) != + GMSSL_INNER_OK) { + PyErr_SetString(InvalidValueError, "invalid public_key"); + return NULL; + } + ret = sm2_key_set_public_key(&sm2_key, &pub); if (ret != GMSSL_INNER_OK) { PyErr_SetString(InvalidValueError, "invalid public_key"); return NULL; @@ -385,8 +425,8 @@ static PyObject *gmsslext_sm3_hash(PyObject *self, PyObject *args, if (!ok) { return NULL; } - if (message_length <= 0) { - PyErr_SetString(InvalidValueError, "empty message"); + if (message_length < 0) { + PyErr_SetString(InvalidValueError, "invalid message"); return NULL; } SM3_CTX sm3_ctx; @@ -416,8 +456,8 @@ static PyObject *gmsslext_sm3_hmac(PyObject *self, PyObject *args, PyErr_SetString(InvalidValueError, "empty key"); return NULL; } - if (message_length <= 0) { - PyErr_SetString(InvalidValueError, "empty message"); + if (message_length < 0) { + PyErr_SetString(InvalidValueError, "invalid message"); return NULL; } SM3_HMAC_CTX hmac_ctx; @@ -501,8 +541,14 @@ static PyObject *gmsslext_sm4_cbc_padding_encrypt(PyObject *self, return PyErr_NoMemory(); } sm4_set_encrypt_key(&sm4_key, (uint8_t *)key); - sm4_cbc_padding_encrypt(&sm4_key, (uint8_t *)iv, (uint8_t *)plaintext, + int ret = sm4_cbc_padding_encrypt(&sm4_key, (uint8_t *)iv, (uint8_t *)plaintext, plaintext_length, (uint8_t *)out, (size_t *)&outlen); + if (ret != GMSSL_INNER_OK) { + PyMem_RawFree(out); + PyErr_SetString(GmsslInnerError, + "libgmssl inner error in sm4_cbc_padding_encrypt"); + return NULL; + } PyObject *ciphertext_obj = Py_BuildValue("y#", out, outlen); PyMem_RawFree(out); return ciphertext_obj; @@ -546,8 +592,15 @@ static PyObject *gmsslext_sm4_cbc_padding_decrypt(PyObject *self, return PyErr_NoMemory(); } sm4_set_decrypt_key(&sm4_key, (uint8_t *)key); - sm4_cbc_padding_decrypt(&sm4_key, (uint8_t *)iv, (uint8_t *)ciphertext, - ciphertext_length, (uint8_t *)out, (size_t *)&outlen); + int ret = sm4_cbc_padding_decrypt(&sm4_key, (uint8_t *)iv, + (uint8_t *)ciphertext, ciphertext_length, + (uint8_t *)out, (size_t *)&outlen); + if (ret != GMSSL_INNER_OK) { + PyMem_RawFree(out); + PyErr_SetString(GmsslInnerError, + "libgmssl inner error in sm4_cbc_padding_decrypt"); + return NULL; + } PyObject *plaintext_obj = Py_BuildValue("y#", out, outlen); PyMem_RawFree(out); return plaintext_obj; @@ -639,11 +692,11 @@ static PyObject *gmsslext_sm4_ctr_decrypt(PyObject *self, PyObject *args, } sm4_set_encrypt_key(&sm4_key, (uint8_t *)key); - // sm4_ctr_decrypt 会修改 ctr ,会导致 Python 端调用者的 ctr 也发生改变,copy + // sm4_ctr_encrypt 会修改 ctr ,会导致 Python 端调用者的 ctr 也发生改变,copy // 一份来用 unsigned char temp_ctr[SM4_BLOCK_SIZE]; memcpy(temp_ctr, ctr, SM4_BLOCK_SIZE); - sm4_ctr_decrypt(&sm4_key, temp_ctr, (uint8_t *)ciphertext, ciphertext_length, + sm4_ctr_encrypt(&sm4_key, temp_ctr, (uint8_t *)ciphertext, ciphertext_length, (uint8_t *)out); PyObject *plaintext_obj = Py_BuildValue("y#", out, ciphertext_length); PyMem_RawFree(out); @@ -739,6 +792,10 @@ static PyObject *gmsslext_sm4_gcm_decrypt(PyObject *self, PyObject *args, PyErr_SetString(InvalidValueError, "empty ciphertext"); return NULL; } + if (tag_length < SM4_GCM_MIN_TAG_SIZE || tag_length > SM4_GCM_MAX_TAG_SIZE) { + PyErr_SetString(InvalidValueError, "invalid sm4 tag length"); + return NULL; + } SM4_KEY sm4_key; // 密文长度与明文一致 @@ -753,7 +810,8 @@ static PyObject *gmsslext_sm4_gcm_decrypt(PyObject *self, PyObject *args, (uint8_t *)tag, tag_length, (uint8_t *)out); if (ret != GMSSL_INNER_OK) { PyMem_RawFree(out); - PyErr_SetString(GmsslInnerError, "libgmssl inner error in sm4_gcm_decrypt"); + // key/iv/tag 长度都已在上面校验,这里失败只可能是 tag 不匹配(认证失败) + PyErr_SetString(InvalidValueError, "authentication failed"); return NULL; } PyObject *obj = Py_BuildValue("y#", out, ciphertext_length); @@ -939,7 +997,6 @@ PyMODINIT_FUNC PyInit_gmsslext(void) { if (PyModule_AddObject(m, "SM9MasterPublicKey", (PyObject *)&GmsslextSM9MasterPublicKeyType) < 0) { Py_DECREF(&GmsslextSM9MasterPublicKeyType); - Py_DECREF(&GmsslextSM9PrivateKeyType); Py_DECREF(m); return NULL; } @@ -947,8 +1004,6 @@ PyMODINIT_FUNC PyInit_gmsslext(void) { if (PyModule_AddObject(m, "SM9MasterKey", (PyObject *)&GmsslextSM9MasterKeyType) < 0) { Py_DECREF(&GmsslextSM9MasterKeyType); - Py_DECREF(&GmsslextSM9MasterPublicKeyType); - Py_DECREF(&GmsslextSM9PrivateKeyType); Py_DECREF(m); return NULL; } diff --git a/gmssl_pyx/sm2_utils.py b/gmssl_pyx/sm2_utils.py index a584456..924ccae 100644 --- a/gmssl_pyx/sm2_utils.py +++ b/gmssl_pyx/sm2_utils.py @@ -19,12 +19,22 @@ HexStr = str +def _is_on_curve(x: int, y: int, a: int, b: int, p: int) -> bool: + """点 (x, y) 是否在曲线 y^2 = x^3 + ax + b (mod p) 上,且 x、y 是合法域元素""" + return x < p and y < p and (y * y - (x**3 + a * x + b)) % p == 0 + + def decompress_sm2_public_key(key: bytes, a: int, b: int, p: int) -> bytes: prefix = key[0] + if prefix not in (2, 3): + raise InvalidValueError("invalid public key") x = int.from_bytes(key[1:], "big") # y^2 = (x^3 + ax + b) % p y_sq = (x**3 + a * x + b) % p y = pow(y_sq, (p + 1) // 4, p) + # 校验 y 是 x 对应的合法平方根,即点 (x, y) 必须在曲线上 + if not _is_on_curve(x, y, a, b, p): + raise InvalidValueError("invalid public key") # y 是偶数,前缀为 '\x02' ;奇数则是 '\x03' if (prefix - 2) != (y % 2): # y 的奇偶与前缀表示不同 @@ -38,19 +48,31 @@ def normalize_sm2_public_key(public_key: t.Union[HexStr, bytes]) -> bytes: Args: public_key: 16 进制字符串或者字节串 - Returns: 64 字节的字节串 + Returns: 64 字节的字节串,且是曲线上的点 + + Raises: InvalidValueError """ pk: bytes = public_key if not isinstance(public_key, bytes): - pk = binascii.unhexlify(public_key) + try: + pk = binascii.unhexlify(public_key) + except binascii.Error as e: + raise InvalidValueError("invalid public key") from e if len(pk) == 65: if pk[0] != 4: raise InvalidValueError("invalid public key") - return pk[1:] + raw = pk[1:] elif len(pk) == 64: - return pk + raw = pk elif len(pk) == 33: return decompress_sm2_public_key(pk, ga, gb, gp) else: raise InvalidValueError("invalid public key") + + # 校验 64 字节公钥 (x, y) 是曲线上的点 + x = int.from_bytes(raw[:32], "big") + y = int.from_bytes(raw[32:], "big") + if not _is_on_curve(x, y, ga, gb, gp): + raise InvalidValueError("invalid public key") + return raw diff --git a/tests/test_sm2.py b/tests/test_sm2.py index 72ca892..a3f7038 100644 --- a/tests/test_sm2.py +++ b/tests/test_sm2.py @@ -125,6 +125,24 @@ def test_error_inherit(self): self.assertTrue(issubclass(GmsslInnerError, Exception)) self.assertTrue(issubclass(InvalidValueError, GmsslInnerError)) + def test_all_zero_public_key(self): + # 全零公钥是无穷远点,不是合法公钥,必须拒绝 + # 回归测试:sm2_z256_point_from_bytes 对无穷远点返回 0 而非 -1, + # 之前包装层只判断 == -1 导致放行,sm2_encrypt 会产出无法解密的密文 + public_key = b"\x00" * 64 + private_key = secrets.token_bytes(32) + message = b"hello world" + with self.assertRaises(InvalidValueError): + sm2_encrypt(public_key, message) + with self.assertRaises(InvalidValueError): + sm2_verify(public_key, message, b"\x30\x06\x02\x01\x01\x02\x01\x01") + with self.assertRaises(InvalidValueError): + sm2_verify_sm3_digest( + public_key, secrets.token_bytes(32), b"\x30\x06\x02\x01\x01\x02\x01\x01" + ) + with self.assertRaises(InvalidValueError): + sm2_sign(private_key, public_key, message) + def test_sm2_sign_and_verify(self): public_key, private_key = sm2_key_generate() message_length = random.randint(1, 1024) @@ -163,6 +181,10 @@ def test_sm2_sign_and_verify(self): ) self.assertFalse(verify) + signature = sm2_sign(private_key, public_key, b"") + verify = sm2_verify(public_key, b"", signature) + self.assertTrue(verify) + def test_sm2_sign_and_verify_error(self): public_key, private_key = sm2_key_generate() message = b"hello world" @@ -172,9 +194,6 @@ def test_sm2_sign_and_verify_error(self): with self.assertRaises(InvalidValueError) as cm: sm2_sign(private_key, public_key[:63], message) self.assertEqual(str(cm.exception), "invalid public_key or private_key length") - with self.assertRaises(InvalidValueError) as cm: - sm2_sign(private_key, public_key, b"") - self.assertEqual(str(cm.exception), "empty message") with self.assertRaises(InvalidValueError) as cm: sm2_sign(private_key, public_key, message, signer_id=b"") self.assertEqual(str(cm.exception), "invalid signer_id length") @@ -182,9 +201,6 @@ def test_sm2_sign_and_verify_error(self): with self.assertRaises(InvalidValueError) as cm: sm2_verify(public_key[:63], message, b"signature") self.assertEqual(str(cm.exception), "invalid public_key") - with self.assertRaises(InvalidValueError) as cm: - sm2_verify(public_key, b"", b"signature") - self.assertEqual(str(cm.exception), "empty message") with self.assertRaises(InvalidValueError) as cm: sm2_verify(public_key, message, b"") self.assertEqual(str(cm.exception), "empty signature") @@ -209,6 +225,41 @@ def test_normalize_sm2_public_key(self): k1 = normalize_sm2_public_key(compressed_public_key) self.assertEqual(k1, raw_public_key) + def test_normalize_sm2_public_key_error(self): + # x 不在曲线上(x = 2^256 - 1 不是合法 x 坐标) + with self.assertRaises(InvalidValueError): + normalize_sm2_public_key(b"\x02" + b"\xff" * 32) + with self.assertRaises(InvalidValueError): + normalize_sm2_public_key(b"\x03" + b"\xff" * 32) + # x 超出素数 p + p = int( + "FFFFFFFEFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFF00000000FFFFFFFFFFFFFFFF", + 16, + ) + with self.assertRaises(InvalidValueError): + normalize_sm2_public_key(b"\x02" + p.to_bytes(32, "big")) + # 非法压缩前缀 + with self.assertRaises(InvalidValueError): + normalize_sm2_public_key(b"\x00" + b"\x01" * 32) + with self.assertRaises(InvalidValueError): + normalize_sm2_public_key(b"\x04" + b"\x01" * 32) + # 64 字节原始公钥不在曲线上(含全零 = 无穷远点) + with self.assertRaises(InvalidValueError): + normalize_sm2_public_key(b"\xff" * 64) + with self.assertRaises(InvalidValueError): + normalize_sm2_public_key(b"\x00" * 64) + # 65 字节非压缩公钥(0x04 前缀)不在曲线上 + with self.assertRaises(InvalidValueError): + normalize_sm2_public_key(b"\x04" + b"\xff" * 64) + # 非法 hex 字符串 + with self.assertRaises(InvalidValueError): + normalize_sm2_public_key("zz") + with self.assertRaises(InvalidValueError): + normalize_sm2_public_key("abc") + # 长度非法 + with self.assertRaises(InvalidValueError): + normalize_sm2_public_key(b"\x02" + b"\x01" * 31) + def test_randbytes(self): for i in range(1, 257): bs = rand_bytes(i) diff --git a/tests/test_sm3.py b/tests/test_sm3.py index 446f153..c704d99 100644 --- a/tests/test_sm3.py +++ b/tests/test_sm3.py @@ -26,10 +26,12 @@ def test_hash(self): got_hash = sm3_hash(message=message) self.assertEqual(got_hash, expected_hash) - def test_hash_error(self): - with self.assertRaises(InvalidValueError) as cm: - sm3_hash(b"") - self.assertEqual(str(cm.exception), "empty message") + expected_hash = binascii.unhexlify( + "1ab21d8355cfa17f8e61194831e81a8f22bec8c728fefb747ed035eb5082aa2b", + ) + message = b"" + got_hash = sm3_hash(message=message) + self.assertEqual(got_hash, expected_hash) def test_hmac(self): n = random.randint(1, 4096) @@ -47,11 +49,15 @@ def test_hmac(self): ) self.assertEqual(hmac_data.hex(), expected_hex) + message = b"" + hmac_data = sm3_hmac(key, message=message) + expected_hex = ( + "639486f482c0ec52cbd4900b9277b7c2132ff049e6a818b7dcdc13ebe1668dfd" + ) + self.assertEqual(hmac_data.hex(), expected_hex) + def test_hmac_error(self): key = secrets.token_bytes(32) - with self.assertRaises(InvalidValueError) as cm: - sm3_hmac(key, b"") - self.assertEqual(str(cm.exception), "empty message") with self.assertRaises(InvalidValueError) as cm: sm3_hmac(b"", b"hello") self.assertEqual(str(cm.exception), "empty key") diff --git a/tests/test_sm4.py b/tests/test_sm4.py index a133bf6..9a9d423 100644 --- a/tests/test_sm4.py +++ b/tests/test_sm4.py @@ -7,6 +7,7 @@ from gmssl_pyx import ( SM4_BLOCK_SIZE, SM4_KEY_SIZE, + GmsslInnerError, InvalidValueError, sm4_cbc_padding_decrypt, sm4_cbc_padding_encrypt, @@ -62,6 +63,18 @@ def test_cbc_encrypt_and_decrypt_error(self): sm4_cbc_padding_decrypt(key, iv, b"") self.assertEqual(str(cm.exception), "empty ciphertext") + # 密文长度不是 16 的倍数,必须抛异常而不是返回数据 + with self.assertRaises(GmsslInnerError): + sm4_cbc_padding_decrypt(key, iv, b"1" * 17) + # PKCS#7 padding 非法,必须抛异常而不是返回数据 + ciphertext = sm4_cbc_padding_encrypt(key, iv, b"hello world") + bad = bytearray(ciphertext) + for i in range(len(bad)): + bad[i] = i % 256 + with self.assertRaises(GmsslInnerError): + unexpected = sm4_cbc_padding_decrypt(key, iv, bytes(bad)) + print(f"unexpected: <{unexpected}>") + def test_ctr_encrypt_and_decrypt(self): for i in range(3): n = random.randint(1, 4096) @@ -151,6 +164,31 @@ def test_gcm_encrypt_and_decrypt_error(self): sm4_gcm_decrypt(key, iv, aad, b"", tag=secrets.token_bytes(16)) self.assertEqual(str(cm.exception), "empty ciphertext") + def test_gcm_decrypt_tag_error(self): + key = secrets.token_bytes(SM4_KEY_SIZE) + iv = secrets.token_bytes(SM4_BLOCK_SIZE) + aad = secrets.token_bytes(16) + plaintext = b"hello world" + ciphertext, tag = sm4_gcm_encrypt(key, iv, aad, plaintext=plaintext) + + # tag 长度非法(合法范围 12 ~ 16) + with self.assertRaises(InvalidValueError) as cm: + sm4_gcm_decrypt( + key, iv=iv, aad=aad, ciphertext=ciphertext, tag=b"1" * 8 + ) + self.assertEqual(str(cm.exception), "invalid sm4 tag length") + # tag 不匹配 → 认证失败,而不是库内部错误 + with self.assertRaises(InvalidValueError) as cm: + sm4_gcm_decrypt( + key, iv=iv, aad=aad, ciphertext=ciphertext, tag=secrets.token_bytes(16) + ) + self.assertEqual(str(cm.exception), "authentication failed") + # 正确 tag 仍可解密 + got_plaintext = sm4_gcm_decrypt( + key, iv=iv, aad=aad, ciphertext=ciphertext, tag=tag + ) + self.assertEqual(got_plaintext, plaintext) + def test_cbc_generated_data(self): key_path = script_dir / "data" / "sm4_generated_key.json" d = json.loads(key_path.read_text(encoding="utf-8")) diff --git a/tests/test_sm9.py b/tests/test_sm9.py index 264efc4..59a9272 100644 --- a/tests/test_sm9.py +++ b/tests/test_sm9.py @@ -99,7 +99,6 @@ def test_sm9_encrypt_and_decrypt_error(self): str(cm.exception), "invalid sm9 identity or ciphertext length" ) - @unittest.skip("Skip because of known issue") def test_sm9_master_key_der(self): identity = secrets.token_bytes(6) @@ -230,7 +229,6 @@ def _check(n: int): for f in fs: f.result() - @unittest.skip("Skip because of known issue") def test_sm9_master_key_encrypt_der(self): identity = secrets.token_bytes(6) @@ -295,9 +293,10 @@ def _check(n: int): for f in fs: f.result() - @unittest.skip("Skip because of known issue") def test_issue_1884(self): - master_der = bytes.fromhex("306602200084509d9f11799ba847a142b4c1ed860dd66943ecf79f544d4327882beb229d03420004654a84a614e4e3f670152a4253ef8fe5127ad7a5b0d85a6a009b3a95dadfb25d407d04cf4d90c3addd3b9829ed92a37d1be32af3ae32e5cb3b6fd5a1ddb8fa2c") + master_der = bytes.fromhex( + "306602200084509d9f11799ba847a142b4c1ed860dd66943ecf79f544d4327882beb229d03420004654a84a614e4e3f670152a4253ef8fe5127ad7a5b0d85a6a009b3a95dadfb25d407d04cf4d90c3addd3b9829ed92a37d1be32af3ae32e5cb3b6fd5a1ddb8fa2c" + ) identity = secrets.token_bytes(6) master = SM9MasterKey.from_der(master_der) key = master.extract_key(identity)