diff --git a/src/tpm2_packet.c b/src/tpm2_packet.c index e4809fb3..9bee4897 100644 --- a/src/tpm2_packet.c +++ b/src/tpm2_packet.c @@ -1157,6 +1157,8 @@ void TPM2_Packet_AppendPublicArea(TPM2_Packet* packet, TPMT_PUBLIC* publicArea) TPM2_Packet_AppendU16(packet, publicArea->type); TPM2_Packet_AppendU16(packet, publicArea->nameAlg); TPM2_Packet_AppendU32(packet, publicArea->objectAttributes); + if (publicArea->authPolicy.size > sizeof(publicArea->authPolicy.buffer)) + publicArea->authPolicy.size = sizeof(publicArea->authPolicy.buffer); TPM2_Packet_AppendU16(packet, publicArea->authPolicy.size); TPM2_Packet_AppendBytes(packet, publicArea->authPolicy.buffer, publicArea->authPolicy.size); @@ -1166,16 +1168,28 @@ void TPM2_Packet_AppendPublicArea(TPM2_Packet* packet, TPMT_PUBLIC* publicArea) switch (publicArea->type) { case TPM_ALG_KEYEDHASH: + if (publicArea->unique.keyedHash.size > + sizeof(publicArea->unique.keyedHash.buffer)) + publicArea->unique.keyedHash.size = + sizeof(publicArea->unique.keyedHash.buffer); TPM2_Packet_AppendU16(packet, publicArea->unique.keyedHash.size); TPM2_Packet_AppendBytes(packet, publicArea->unique.keyedHash.buffer, publicArea->unique.keyedHash.size); break; case TPM_ALG_SYMCIPHER: + if (publicArea->unique.sym.size > + sizeof(publicArea->unique.sym.buffer)) + publicArea->unique.sym.size = + sizeof(publicArea->unique.sym.buffer); TPM2_Packet_AppendU16(packet, publicArea->unique.sym.size); TPM2_Packet_AppendBytes(packet, publicArea->unique.sym.buffer, publicArea->unique.sym.size); break; case TPM_ALG_RSA: + if (publicArea->unique.rsa.size > + sizeof(publicArea->unique.rsa.buffer)) + publicArea->unique.rsa.size = + sizeof(publicArea->unique.rsa.buffer); TPM2_Packet_AppendU16(packet, publicArea->unique.rsa.size); TPM2_Packet_AppendBytes(packet, publicArea->unique.rsa.buffer, publicArea->unique.rsa.size); @@ -1186,6 +1200,10 @@ void TPM2_Packet_AppendPublicArea(TPM2_Packet* packet, TPMT_PUBLIC* publicArea) #ifdef WOLFTPM_MLDSA case TPM_ALG_MLDSA: case TPM_ALG_HASH_MLDSA: + if (publicArea->unique.mldsa.size > + sizeof(publicArea->unique.mldsa.buffer)) + publicArea->unique.mldsa.size = + sizeof(publicArea->unique.mldsa.buffer); TPM2_Packet_AppendU16(packet, publicArea->unique.mldsa.size); TPM2_Packet_AppendBytes(packet, publicArea->unique.mldsa.buffer, publicArea->unique.mldsa.size); @@ -1193,6 +1211,10 @@ void TPM2_Packet_AppendPublicArea(TPM2_Packet* packet, TPMT_PUBLIC* publicArea) #endif /* WOLFTPM_MLDSA */ #ifdef WOLFTPM_MLKEM case TPM_ALG_MLKEM: + if (publicArea->unique.mlkem.size > + sizeof(publicArea->unique.mlkem.buffer)) + publicArea->unique.mlkem.size = + sizeof(publicArea->unique.mlkem.buffer); TPM2_Packet_AppendU16(packet, publicArea->unique.mlkem.size); TPM2_Packet_AppendBytes(packet, publicArea->unique.mlkem.buffer, publicArea->unique.mlkem.size); diff --git a/tests/unit_tests.c b/tests/unit_tests.c index 791bf409..58c14093 100644 --- a/tests/unit_tests.c +++ b/tests/unit_tests.c @@ -3873,6 +3873,31 @@ static void test_TPM2_AppendSensitive_Clamp(void) printf("Test TPM2: %-40s Passed\n", "AppendSensitive clamp:"); } +static void test_TPM2_AppendPublic_Clamp(void) +{ + TPM2_Packet packet; + byte buf[1024]; + TPM2B_PUBLIC pub; + word16 policyCap, rsaCap; + + policyCap = (word16)sizeof(pub.publicArea.authPolicy.buffer); + rsaCap = (word16)sizeof(pub.publicArea.unique.rsa.buffer); + + XMEMSET(&pub, 0, sizeof(pub)); + pub.publicArea.type = TPM_ALG_RSA; + pub.publicArea.nameAlg = TPM_ALG_SHA256; + pub.publicArea.authPolicy.size = policyCap + 100; + pub.publicArea.unique.rsa.size = rsaCap + 100; + XMEMSET(&packet, 0, sizeof(packet)); + packet.buf = buf; + packet.size = sizeof(buf); + TPM2_Packet_AppendPublic(&packet, &pub); + AssertIntEQ(pub.publicArea.authPolicy.size, policyCap); + AssertIntEQ(pub.publicArea.unique.rsa.size, rsaCap); + + printf("Test TPM2: %-40s Passed\n", "AppendPublic clamp:"); +} + /* Roundtrip a maximum-size inner payload (size == buffer capacity) so the * parse-side ParseU16Buf clamp branch is exercised with valid data. */ static void test_TPM2_Sensitive_MaxRoundtrip(void) @@ -6266,6 +6291,7 @@ int unit_tests(int argc, char *argv[]) test_TPM2_TIS_ValidateRspSz(); test_TPM2_ParsePublic_EmptyClears(); test_TPM2_AppendSensitive_Clamp(); + test_TPM2_AppendPublic_Clamp(); test_TPM2_Sensitive_MaxRoundtrip(); test_KeySealTemplate(); test_SealAndKeyedHash_Boundaries();