1#include <stdlib.h>
2#include <string.h>
3
4#include "internal.h"
5#include "mg_random.h"
6
7/* naive der parser */
8typedef struct {
9 const u8* data;
10 size_t size;
11 size_t cursor;
12} DERReader;
13
14static int der_length(DERReader* reader, size_t* length)
15{
16 u8 first;
17 size_t value = 0;
18 unsigned count;
19 unsigned i;
20
21 if (reader->cursor >= reader->size)
22 return 0;
23 first = reader->data[reader->cursor++];
24 if (!(first & 0x80)) {
25 *length = first;
26 return *length <= reader->size - reader->cursor;
27 }
28 count = first & 0x7f;
29 if (count == 0 || count > sizeof(size_t) || count > reader->size - reader->cursor)
30 return 0;
31 for (i = 0; i < count; ++i)
32 value = (value << 8) | reader->data[reader->cursor++];
33 if (value > reader->size - reader->cursor)
34 return 0;
35 *length = value;
36 return 1;
37}
38
39static int der_value(DERReader* reader, u8 tag, DERReader* value)
40{
41 size_t length;
42 if (reader->cursor >= reader->size || reader->data[reader->cursor++] != tag ||
43 !der_length(reader, &length))
44 return 0;
45 value->data = reader->data + reader->cursor;
46 value->size = length;
47 value->cursor = 0;
48 reader->cursor += length;
49 return 1;
50}
51
52/* parse the server's X.509 RSA pubkey for BearSSL */
53static int rsa_public_key(const u8* der, size_t der_size, br_rsa_public_key* key, char* error,
54 size_t error_size)
55{
56 DERReader root = { der, der_size, 0 };
57 DERReader spki;
58 DERReader algorithm;
59 DERReader bits;
60 DERReader rsa;
61 DERReader modulus;
62 DERReader exponent;
63
64 if (!der_value(&root, 0x30, &spki) || !der_value(&spki, 0x30, &algorithm) ||
65 !der_value(&spki, 0x03, &bits) || bits.size < 1 || bits.data[0] != 0) {
66 mc_set_error(error, error_size, "invalid X.509 RSA public key");
67 return 0;
68 }
69 bits.cursor = 1;
70 if (!der_value(&bits, 0x30, &rsa) || !der_value(&rsa, 0x02, &modulus) ||
71 !der_value(&rsa, 0x02, &exponent)) {
72 mc_set_error(error, error_size, "invalid RSA public key integers");
73 return 0;
74 }
75 while (modulus.size > 1 && modulus.data[0] == 0) {
76 ++modulus.data;
77 --modulus.size;
78 }
79 while (exponent.size > 1 && exponent.data[0] == 0) {
80 ++exponent.data;
81 --exponent.size;
82 }
83 key->n = (unsigned char*)modulus.data;
84 key->nlen = modulus.size;
85 key->e = (unsigned char*)exponent.data;
86 key->elen = exponent.size;
87 return key->nlen >= 64 && key->elen > 0;
88}
89
90void mc_cipher_init(MCCipher* cipher, const u8 key[16])
91{
92 br_aes_ct_cbcenc_init(&cipher->key, key, 16);
93 memcpy(cipher->state, key, 16);
94}
95
96static u8 cipher_byte(MCCipher* cipher)
97{
98 u8 zero_iv[16] = { 0 };
99 u8 block[16];
100 memcpy(block, cipher->state, 16);
101 br_aes_ct_cbcenc_run(&cipher->key, zero_iv, block, sizeof(block));
102 return block[0];
103}
104
105void mc_cipher_encrypt(MCCipher* cipher, u8* data, size_t size)
106{
107 size_t i;
108 for (i = 0; i < size; ++i) {
109 u8 encrypted = data[i] ^ cipher_byte(cipher);
110 memmove(cipher->state, cipher->state + 1, 15);
111 cipher->state[15] = encrypted;
112 data[i] = encrypted;
113 }
114}
115
116void mc_cipher_decrypt(MCCipher* cipher, u8* data, size_t size)
117{
118 size_t i;
119 for (i = 0; i < size; ++i) {
120 u8 encrypted = data[i];
121 data[i] ^= cipher_byte(cipher);
122 memmove(cipher->state, cipher->state + 1, 15);
123 cipher->state[15] = encrypted;
124 }
125}
126
127int mc_rsa_encrypt(const u8* public_key, size_t public_key_size, const u8* message,
128 size_t message_size, u8** encrypted, size_t* encrypted_size, char* error,
129 size_t error_size)
130{
131 br_rsa_public_key key;
132 br_rsa_public rsa;
133 u8* output;
134 size_t padding_size;
135 size_t i;
136
137 if (!rsa_public_key(public_key, public_key_size, &key, error, error_size))
138 return 0;
139 if (message_size + 11 > key.nlen) {
140 mc_set_error(error, error_size, "message is too large for server RSA key");
141 return 0;
142 }
143 output = (u8*)malloc(key.nlen);
144 if (!output) {
145 mc_set_error(error, error_size, "out of memory encrypting login response");
146 return 0;
147 }
148 padding_size = key.nlen - message_size - 3;
149 output[0] = 0;
150 output[1] = 2;
151 /* PKCS#1 v1.5 encryption padding requires every padding byte to be
152 nonzero */
153 if (!mg_random_bytes(output + 2, padding_size)) {
154 free(output);
155 mc_set_error(error, error_size, "could not obtain secure random bytes");
156 return 0;
157 }
158 for (i = 0; i < padding_size; ++i) {
159 while (output[2 + i] == 0) {
160 if (!mg_random_bytes(output + 2 + i, 1)) {
161 free(output);
162 mc_set_error(error, error_size,
163 "could not obtain secure random bytes");
164 return 0;
165 }
166 }
167 }
168 output[2 + padding_size] = 0;
169 memcpy(output + 3 + padding_size, message, message_size);
170
171 rsa = br_rsa_public_get_default();
172 if (!rsa || !rsa(output, key.nlen, &key)) {
173 free(output);
174 mc_set_error(error, error_size, "BearSSL RSA encryption failed");
175 return 0;
176 }
177 *encrypted = output;
178 *encrypted_size = key.nlen;
179 return 1;
180}
181
182static void twos_complement(u8* data, size_t size)
183{
184 size_t i = size;
185 unsigned carry = 1;
186 while (i > 0) {
187 unsigned value;
188 --i;
189 value = (unsigned)(data[i] ^ 0xff) + carry;
190 data[i] = (u8)value;
191 carry = value >> 8;
192 }
193}
194
195void mc_server_hash(const char* server_id, const u8 shared_secret[16], const u8* public_key,
196 size_t public_key_size, char output[42])
197{
198 static const char hex[] = "0123456789abcdef";
199 br_sha1_context sha;
200 u8 digest[20];
201 size_t first = 0;
202 size_t i;
203 size_t out = 0;
204 int negative;
205
206 br_sha1_init(&sha);
207 br_sha1_update(&sha, server_id, strlen(server_id));
208 br_sha1_update(&sha, shared_secret, 16);
209 br_sha1_update(&sha, public_key, public_key_size);
210 br_sha1_out(&sha, digest);
211 /* mc puts sha-1 digest as a signed twos-complement int */
212 negative = (digest[0] & 0x80) != 0;
213 if (negative)
214 twos_complement(digest, sizeof(digest));
215 while (first < sizeof(digest) && digest[first] == 0)
216 ++first;
217 if (negative)
218 output[out++] = '-';
219 if (first == sizeof(digest)) {
220 output[out++] = '0';
221 } else {
222 output[out++] = hex[digest[first] >> 4];
223 output[out++] = hex[digest[first] & 15];
224 for (i = first + 1; i < sizeof(digest); ++i) {
225 output[out++] = hex[digest[i] >> 4];
226 output[out++] = hex[digest[i] & 15];
227 }
228 if (output[negative ? 1 : 0] == '0') {
229 memmove(output + (negative ? 1 : 0), output + (negative ? 2 : 1),
230 out - (negative ? 2 : 1));
231 --out;
232 }
233 }
234 output[out] = '\0';
235}