Ruby 4.1.0dev (2026-09-09 revision a925c74bd239398b97112e24f952c377345c5bff)
integer.c
1#include "prism/internal/integer.h"
2
3#include "prism/internal/allocator.h"
4#include "prism/internal/buffer.h"
5
6#include <assert.h>
7#include <inttypes.h>
8#include <stdbool.h>
9#include <stddef.h>
10#include <stdlib.h>
11#include <string.h>
12
17static void
18pm_integer_free(pm_integer_t *integer) {
19 if (integer->values) {
20 xfree(integer->values);
21 }
22}
23
28#define INTEGER_EXTRACT(integer, length_variable, values_variable) \
29 if ((integer)->values == NULL) { \
30 length_variable = 1; \
31 values_variable = &(integer)->value; \
32 } else { \
33 length_variable = (integer)->length; \
34 values_variable = (integer)->values; \
35 }
36
41static void
42big_add(pm_integer_t *destination, pm_integer_t *left, pm_integer_t *right, uint64_t base) {
43 size_t left_length;
44 uint32_t *left_values;
45 INTEGER_EXTRACT(left, left_length, left_values)
46
47 size_t right_length;
48 uint32_t *right_values;
49 INTEGER_EXTRACT(right, right_length, right_values)
50
51 size_t length = left_length < right_length ? right_length : left_length;
52 uint32_t *values = (uint32_t *) xmalloc(sizeof(uint32_t) * (length + 1));
53 if (values == NULL) return;
54
55 uint64_t carry = 0;
56 for (size_t index = 0; index < length; index++) {
57 uint64_t sum = carry + (index < left_length ? left_values[index] : 0) + (index < right_length ? right_values[index] : 0);
58 values[index] = (uint32_t) (sum % base);
59 carry = sum / base;
60 }
61
62 if (carry > 0) {
63 values[length] = (uint32_t) carry;
64 length++;
65 }
66
67 *destination = (pm_integer_t) { length, values, 0, false };
68}
69
75static void
76big_sub2(pm_integer_t *destination, pm_integer_t *a, pm_integer_t *b, pm_integer_t *c, uint64_t base) {
77 size_t a_length;
78 uint32_t *a_values;
79 INTEGER_EXTRACT(a, a_length, a_values)
80
81 size_t b_length;
82 uint32_t *b_values;
83 INTEGER_EXTRACT(b, b_length, b_values)
84
85 size_t c_length;
86 uint32_t *c_values;
87 INTEGER_EXTRACT(c, c_length, c_values)
88
89 uint32_t *values = (uint32_t*) xmalloc(sizeof(uint32_t) * a_length);
90 int64_t carry = 0;
91
92 for (size_t index = 0; index < a_length; index++) {
93 int64_t sub = (
94 carry +
95 a_values[index] -
96 (index < b_length ? b_values[index] : 0) -
97 (index < c_length ? c_values[index] : 0)
98 );
99
100 if (sub >= 0) {
101 values[index] = (uint32_t) sub;
102 carry = 0;
103 } else {
104 sub += 2 * (int64_t) base;
105 values[index] = (uint32_t) ((uint64_t) sub % base);
106 carry = sub / (int64_t) base - 2;
107 }
108 }
109
110 while (a_length > 1 && values[a_length - 1] == 0) a_length--;
111 *destination = (pm_integer_t) { a_length, values, 0, false };
112}
113
118static void
119karatsuba_multiply(pm_integer_t *destination, pm_integer_t *left, pm_integer_t *right, uint64_t base) {
120 size_t left_length;
121 uint32_t *left_values;
122 INTEGER_EXTRACT(left, left_length, left_values)
123
124 size_t right_length;
125 uint32_t *right_values;
126 INTEGER_EXTRACT(right, right_length, right_values)
127
128 if (left_length > right_length) {
129 size_t temporary_length = left_length;
130 left_length = right_length;
131 right_length = temporary_length;
132
133 uint32_t *temporary_values = left_values;
134 left_values = right_values;
135 right_values = temporary_values;
136 }
137
138 if (left_length <= 10) {
139 size_t length = left_length + right_length;
140 uint32_t *values = (uint32_t *) xcalloc(length, sizeof(uint32_t));
141 if (values == NULL) return;
142
143 for (size_t left_index = 0; left_index < left_length; left_index++) {
144 uint32_t carry = 0;
145 for (size_t right_index = 0; right_index < right_length; right_index++) {
146 uint64_t product = (uint64_t) left_values[left_index] * right_values[right_index] + values[left_index + right_index] + carry;
147 values[left_index + right_index] = (uint32_t) (product % base);
148 carry = (uint32_t) (product / base);
149 }
150 values[left_index + right_length] = carry;
151 }
152
153 while (length > 1 && values[length - 1] == 0) length--;
154 *destination = (pm_integer_t) { length, values, 0, false };
155 return;
156 }
157
158 if (left_length * 2 <= right_length) {
159 uint32_t *values = (uint32_t *) xcalloc(left_length + right_length, sizeof(uint32_t));
160
161 for (size_t start_offset = 0; start_offset < right_length; start_offset += left_length) {
162 size_t end_offset = start_offset + left_length;
163 if (end_offset > right_length) end_offset = right_length;
164
165 pm_integer_t sliced_left = {
166 .length = left_length,
167 .values = left_values,
168 .value = 0,
169 .negative = false
170 };
171
172 pm_integer_t sliced_right = {
173 .length = end_offset - start_offset,
174 .values = right_values + start_offset,
175 .value = 0,
176 .negative = false
177 };
178
179 pm_integer_t product;
180 karatsuba_multiply(&product, &sliced_left, &sliced_right, base);
181
182 uint32_t carry = 0;
183 for (size_t index = 0; index < product.length; index++) {
184 uint64_t sum = (uint64_t) values[start_offset + index] + product.values[index] + carry;
185 values[start_offset + index] = (uint32_t) (sum % base);
186 carry = (uint32_t) (sum / base);
187 }
188
189 if (carry > 0) values[start_offset + product.length] += carry;
190 pm_integer_free(&product);
191 }
192
193 *destination = (pm_integer_t) { left_length + right_length, values, 0, false };
194 return;
195 }
196
197 size_t half = left_length / 2;
198 pm_integer_t x0 = { half, left_values, 0, false };
199 pm_integer_t x1 = { left_length - half, left_values + half, 0, false };
200 pm_integer_t y0 = { half, right_values, 0, false };
201 pm_integer_t y1 = { right_length - half, right_values + half, 0, false };
202
203 pm_integer_t z0 = { 0 };
204 karatsuba_multiply(&z0, &x0, &y0, base);
205
206 pm_integer_t z2 = { 0 };
207 karatsuba_multiply(&z2, &x1, &y1, base);
208
209 // For simplicity to avoid considering negative values,
210 // use `z1 = (x0 + x1) * (y0 + y1) - z0 - z2` instead of original karatsuba algorithm.
211 pm_integer_t x01 = { 0 };
212 big_add(&x01, &x0, &x1, base);
213
214 pm_integer_t y01 = { 0 };
215 big_add(&y01, &y0, &y1, base);
216
217 pm_integer_t xy = { 0 };
218 karatsuba_multiply(&xy, &x01, &y01, base);
219
220 pm_integer_t z1;
221 big_sub2(&z1, &xy, &z0, &z2, base);
222
223 size_t length = left_length + right_length;
224 uint32_t *values = (uint32_t*) xcalloc(length, sizeof(uint32_t));
225
226 assert(z0.values != NULL);
227 memcpy(values, z0.values, sizeof(uint32_t) * z0.length);
228
229 assert(z2.values != NULL);
230 memcpy(values + 2 * half, z2.values, sizeof(uint32_t) * z2.length);
231
232 uint32_t carry = 0;
233 for(size_t index = 0; index < z1.length; index++) {
234 uint64_t sum = (uint64_t) carry + values[index + half] + z1.values[index];
235 values[index + half] = (uint32_t) (sum % base);
236 carry = (uint32_t) (sum / base);
237 }
238
239 for(size_t index = half + z1.length; carry > 0; index++) {
240 uint64_t sum = (uint64_t) carry + values[index];
241 values[index] = (uint32_t) (sum % base);
242 carry = (uint32_t) (sum / base);
243 }
244
245 while (length > 1 && values[length - 1] == 0) length--;
246 pm_integer_free(&z0);
247 pm_integer_free(&z1);
248 pm_integer_free(&z2);
249 pm_integer_free(&x01);
250 pm_integer_free(&y01);
251 pm_integer_free(&xy);
252
253 *destination = (pm_integer_t) { length, values, 0, false };
254}
255
263static const int8_t pm_integer_parse_digit_values[256] = {
264// 0 1 2 3 4 5 6 7 8 9 A B C D E F
265 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // 0x
266 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // 1x
267 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // 2x
268 0, 1, 2, 3, 4, 5, 6, 7, 8, 9, -1, -1, -1, -1, -1, -1, // 3x
269 -1, 10, 11, 12, 13, 14, 15, -1, -1, -1, -1, -1, -1, -1, -1, -1, // 4x
270 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, 0, // 5x
271 -1, 10, 11, 12, 13, 14, 15, -1, -1, -1, -1, -1, -1, -1, -1, -1, // 6x
272 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // 7x
273 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // 8x
274 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // 9x
275 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // Ax
276 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // Bx
277 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // Cx
278 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // Dx
279 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // Ex
280 -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, -1, // Fx
281};
282
286static uint8_t
287pm_integer_parse_digit(const uint8_t character) {
288 int8_t value = pm_integer_parse_digit_values[character];
289 assert(value != -1 && "invalid digit");
290
291 return (uint8_t) value;
292}
293
298static void
299pm_integer_from_uint64(pm_integer_t *integer, uint64_t value, uint64_t base) {
300 if (value < base) {
301 integer->value = (uint32_t) value;
302 return;
303 }
304
305 size_t length = 0;
306 uint64_t length_value = value;
307 while (length_value > 0) {
308 length++;
309 length_value /= base;
310 }
311
312 uint32_t *values = (uint32_t *) xmalloc(sizeof(uint32_t) * length);
313 if (values == NULL) return;
314
315 for (size_t value_index = 0; value_index < length; value_index++) {
316 values[value_index] = (uint32_t) (value % base);
317 value /= base;
318 }
319
320 integer->length = length;
321 integer->values = values;
322}
323
329static void
330pm_integer_normalize(pm_integer_t *integer) {
331 if (integer->values == NULL) {
332 return;
333 }
334
335 while (integer->length > 1 && integer->values[integer->length - 1] == 0) {
336 integer->length--;
337 }
338
339 if (integer->length > 1) {
340 return;
341 }
342
343 uint32_t value = integer->values[0];
344 bool negative = integer->negative && value != 0;
345
346 pm_integer_free(integer);
347 *integer = (pm_integer_t) { .values = NULL, .value = value, .length = 0, .negative = negative };
348}
349
354static void
355pm_integer_convert_base(pm_integer_t *destination, const pm_integer_t *source, uint64_t base_from, uint64_t base_to) {
356 size_t source_length;
357 const uint32_t *source_values;
358 INTEGER_EXTRACT(source, source_length, source_values)
359
360 size_t bigints_length = (source_length + 1) / 2;
361 assert(bigints_length > 0);
362
363 pm_integer_t *bigints = (pm_integer_t *) xcalloc(bigints_length, sizeof(pm_integer_t));
364 if (bigints == NULL) return;
365
366 for (size_t index = 0; index < source_length; index += 2) {
367 uint64_t value = source_values[index] + base_from * (index + 1 < source_length ? source_values[index + 1] : 0);
368 pm_integer_from_uint64(&bigints[index / 2], value, base_to);
369 }
370
371 pm_integer_t base = { 0 };
372 pm_integer_from_uint64(&base, base_from, base_to);
373
374 while (bigints_length > 1) {
375 pm_integer_t next_base;
376 karatsuba_multiply(&next_base, &base, &base, base_to);
377
378 pm_integer_free(&base);
379 base = next_base;
380
381 size_t next_length = (bigints_length + 1) / 2;
382 pm_integer_t *next_bigints = (pm_integer_t *) xcalloc(next_length, sizeof(pm_integer_t));
383
384 for (size_t bigints_index = 0; bigints_index < bigints_length; bigints_index += 2) {
385 if (bigints_index + 1 == bigints_length) {
386 next_bigints[bigints_index / 2] = bigints[bigints_index];
387 } else {
388 pm_integer_t multiplied = { 0 };
389 karatsuba_multiply(&multiplied, &base, &bigints[bigints_index + 1], base_to);
390
391 big_add(&next_bigints[bigints_index / 2], &bigints[bigints_index], &multiplied, base_to);
392 pm_integer_free(&bigints[bigints_index]);
393 pm_integer_free(&bigints[bigints_index + 1]);
394 pm_integer_free(&multiplied);
395 }
396 }
397
398 xfree_sized(bigints, bigints_length * sizeof(pm_integer_t));
399 bigints = next_bigints;
400 bigints_length = next_length;
401 }
402
403 *destination = bigints[0];
404 destination->negative = source->negative;
405 pm_integer_normalize(destination);
406
407 xfree_sized(bigints, bigints_length * sizeof(pm_integer_t));
408 pm_integer_free(&base);
409}
410
411#undef INTEGER_EXTRACT
412
416static void
417pm_integer_parse_powof2(pm_integer_t *integer, uint32_t base, const uint8_t *digits, size_t digits_length) {
418 size_t bit = 1;
419 while (base > (uint32_t) (1 << bit)) bit++;
420
421 size_t length = (digits_length * bit + 31) / 32;
422 uint32_t *values = (uint32_t *) xcalloc(length, sizeof(uint32_t));
423
424 for (size_t digit_index = 0; digit_index < digits_length; digit_index++) {
425 size_t bit_position = bit * (digits_length - digit_index - 1);
426 uint32_t value = digits[digit_index];
427
428 size_t index = bit_position / 32;
429 size_t shift = bit_position % 32;
430
431 values[index] |= value << shift;
432 if (32 - shift < bit) values[index + 1] |= value >> (32 - shift);
433 }
434
435 while (length > 1 && values[length - 1] == 0) length--;
436 *integer = (pm_integer_t) { .length = length, .values = values, .value = 0, .negative = false };
437 pm_integer_normalize(integer);
438}
439
443static void
444pm_integer_parse_decimal(pm_integer_t *integer, const uint8_t *digits, size_t digits_length) {
445 const size_t batch = 9;
446 const size_t length = (digits_length + batch - 1) / batch;
447
448 uint32_t *values = (uint32_t *) xcalloc(length, sizeof(uint32_t));
449 uint32_t value = 0;
450
451 for (size_t digits_index = 0; digits_index < digits_length; digits_index++) {
452 value = value * 10 + digits[digits_index];
453
454 size_t reverse_index = digits_length - digits_index - 1;
455 if (reverse_index % batch == 0) {
456 values[reverse_index / batch] = value;
457 value = 0;
458 }
459 }
460
461 // Convert base from 10**9 to 1<<32.
462 pm_integer_convert_base(integer, &((pm_integer_t) { .length = length, .values = values, .value = 0, .negative = false }), 1000000000, ((uint64_t) 1 << 32));
463 xfree_sized(values, length * sizeof(uint32_t));
464}
465
469static void
470pm_integer_parse_big(pm_integer_t *integer, uint32_t multiplier, const uint8_t *start, const uint8_t *end) {
471 // Allocate an array to store digits.
472 const size_t digits_capa = sizeof(uint8_t) * (size_t) (end - start);
473 uint8_t *digits = xmalloc(digits_capa);
474 size_t digits_length = 0;
475
476 for (; start < end; start++) {
477 if (*start == '_') continue;
478 digits[digits_length++] = pm_integer_parse_digit(*start);
479 }
480
481 // Construct pm_integer_t from the digits.
482 if (multiplier == 10) {
483 pm_integer_parse_decimal(integer, digits, digits_length);
484 } else {
485 pm_integer_parse_powof2(integer, multiplier, digits, digits_length);
486 }
487
488 xfree_sized(digits, digits_capa);
489}
490
496void
497pm_integer_parse(pm_integer_t *integer, pm_integer_base_t base, const uint8_t *start, const uint8_t *end) {
498 // Ignore unary +. Unary - is parsed differently and will not end up here.
499 // Instead, it will modify the parsed integer later.
500 if (*start == '+') start++;
501
502 // Determine the multiplier from the base, and skip past any prefixes.
503 uint32_t multiplier = 10;
504 switch (base) {
505 case PM_INTEGER_BASE_DEFAULT:
506 while (start < end && *start == '0') start++; // 01 -> 1
507 break;
508 case PM_INTEGER_BASE_BINARY:
509 start += 2; // 0b
510 multiplier = 2;
511 break;
512 case PM_INTEGER_BASE_OCTAL:
513 start++; // 0
514 if (*start == '_' || *start == 'o' || *start == 'O') start++; // o
515 multiplier = 8;
516 break;
517 case PM_INTEGER_BASE_DECIMAL:
518 if (*start == '0' && (end - start) > 1) start += 2; // 0d
519 break;
520 case PM_INTEGER_BASE_HEXADECIMAL:
521 start += 2; // 0x
522 multiplier = 16;
523 break;
524 case PM_INTEGER_BASE_UNKNOWN:
525 if (*start == '0' && (end - start) > 1) {
526 switch (start[1]) {
527 case '_': start += 2; multiplier = 8; break;
528 case '0': case '1': case '2': case '3': case '4': case '5': case '6': case '7': start++; multiplier = 8; break;
529 case 'b': case 'B': start += 2; multiplier = 2; break;
530 case 'o': case 'O': start += 2; multiplier = 8; break;
531 case 'd': case 'D': start += 2; break;
532 case 'x': case 'X': start += 2; multiplier = 16; break;
533 default: assert(false && "unreachable"); break;
534 }
535 }
536 break;
537 }
538
539 // It's possible that we've consumed everything at this point if there is an
540 // invalid integer. If this is the case, we'll just return 0.
541 if (start >= end) return;
542
543 const uint8_t *cursor = start;
544 uint64_t value = (uint64_t) pm_integer_parse_digit(*cursor++);
545
546 for (; cursor < end; cursor++) {
547 if (*cursor == '_') continue;
548 value = value * multiplier + (uint64_t) pm_integer_parse_digit(*cursor);
549
550 if (value > UINT32_MAX) {
551 // If the integer is too large to fit into a single uint32_t, then
552 // we'll parse it as a big integer.
553 pm_integer_parse_big(integer, multiplier, start, end);
554 return;
555 }
556 }
557
558 integer->value = (uint32_t) value;
559}
560
566int
567pm_integer_compare(const pm_integer_t *left, const pm_integer_t *right) {
568 if (left->negative != right->negative) return left->negative ? -1 : 1;
569 int negative = left->negative ? -1 : 1;
570
571 if (left->values == NULL && right->values == NULL) {
572 if (left->value < right->value) return -1 * negative;
573 if (left->value > right->value) return 1 * negative;
574 return 0;
575 }
576
577 if (left->values == NULL || left->length < right->length) return -1 * negative;
578 if (right->values == NULL || left->length > right->length) return 1 * negative;
579
580 for (size_t index = 0; index < left->length; index++) {
581 size_t value_index = left->length - index - 1;
582 uint32_t left_value = left->values[value_index];
583 uint32_t right_value = right->values[value_index];
584
585 if (left_value < right_value) return -1 * negative;
586 if (left_value > right_value) return 1 * negative;
587 }
588
589 return 0;
590}
591
595void pm_integers_reduce(pm_integer_t *numerator, pm_integer_t *denominator) {
596 // If either the numerator or denominator do not fit into a 32-bit integer,
597 // then this function is a no-op. In the future, we may consider reducing
598 // even the larger numbers, but for now we're going to keep it simple.
599 if (
600 // If the numerator doesn't fit into a 32-bit integer, return early.
601 numerator->length != 0 ||
602 // If the denominator doesn't fit into a 32-bit integer, return early.
603 denominator->length != 0 ||
604 // If the numerator is 0, then return early.
605 numerator->value == 0 ||
606 // If the denominator is 1, then return early.
607 denominator->value == 1
608 ) return;
609
610 // Find the greatest common divisor of the numerator and denominator.
611 uint32_t divisor = numerator->value;
612 uint32_t remainder = denominator->value;
613
614 while (remainder != 0) {
615 uint32_t temporary = remainder;
616 remainder = divisor % remainder;
617 divisor = temporary;
618 }
619
620 // Divide the numerator and denominator by the greatest common divisor.
621 numerator->value /= divisor;
622 denominator->value /= divisor;
623}
624
628void
629pm_integer_string(pm_buffer_t *buffer, const pm_integer_t *integer) {
630 if (integer->negative) {
631 pm_buffer_append_byte(buffer, '-');
632 }
633
634 // If the integer fits into a single uint32_t, then we can just append the
635 // value directly to the buffer.
636 if (integer->values == NULL) {
637 pm_buffer_append_format(buffer, "%" PRIu32, integer->value);
638 return;
639 }
640
641 // If the integer is two uint32_t values, then we can | them together and
642 // append the result to the buffer.
643 if (integer->length == 2) {
644 const uint64_t value = ((uint64_t) integer->values[0]) | ((uint64_t) integer->values[1] << 32);
645 pm_buffer_append_format(buffer, "%" PRIu64, value);
646 return;
647 }
648
649 // Otherwise, first we'll convert the base from 1<<32 to 10**9.
650 pm_integer_t converted = { 0 };
651 pm_integer_convert_base(&converted, integer, (uint64_t) 1 << 32, 1000000000);
652
653 if (converted.values == NULL) {
654 pm_buffer_append_format(buffer, "%" PRIu32, converted.value);
655 pm_integer_free(&converted);
656 return;
657 }
658
659 // Allocate a buffer that we'll copy the decimal digits into.
660 const size_t digits_length = converted.length * 9;
661 char *digits = xcalloc(digits_length, sizeof(char));
662 if (digits == NULL) return;
663
664 // Pack bigdecimal to digits.
665 for (size_t value_index = 0; value_index < converted.length; value_index++) {
666 uint32_t value = converted.values[value_index];
667
668 for (size_t digit_index = 0; digit_index < 9; digit_index++) {
669 digits[digits_length - 9 * value_index - digit_index - 1] = (char) ('0' + value % 10);
670 value /= 10;
671 }
672 }
673
674 size_t start_offset = 0;
675 while (start_offset < digits_length - 1 && digits[start_offset] == '0') start_offset++;
676
677 // Finally, append the string to the buffer and free the digits.
678 pm_buffer_append_string(buffer, digits + start_offset, digits_length - start_offset);
679 xfree_sized(digits, sizeof(char) * digits_length);
680 pm_integer_free(&converted);
681}
#define xfree
Old name of ruby_xfree.
Definition xmalloc.h:58
#define xmalloc
Old name of ruby_xmalloc.
Definition xmalloc.h:53
#define xcalloc
Old name of ruby_xcalloc.
Definition xmalloc.h:55
C99 shim for <stdbool.h>
A structure represents an arbitrary-sized integer.
Definition integer.h:16
size_t length
The number of allocated values.
Definition integer.h:21
uint32_t value
Embedded value for small integer.
Definition integer.h:32
uint32_t * values
List of 32-bit integers.
Definition integer.h:26
bool negative
Whether or not the integer is negative.
Definition integer.h:38