Compare commits

...

3 Commits

Author SHA1 Message Date
Kittycannon fe57c4390f reorganization 2026-08-22 09:42:23 -06:00
Kittycannon 7692c04057 some other 2026-08-18 00:38:30 -06:00
Kittycannon d678a38ad4 segment with ops and mul_long. Should be easier to work from then on. 2026-08-18 00:29:18 -06:00
21 changed files with 608 additions and 229 deletions
+102 -58
View File
@@ -1,4 +1,7 @@
#include <iostream>
#include <cassert>
#include <string>
#include <array>
#include <bigmath/BigMath.hpp>
#include <bigmath/chain/segment.hpp>
@@ -6,109 +9,150 @@
#include <ckitty/memory/buffers.hpp>
#include <cassert>
#include <iostream>
#include <string>
using namespace bigmath;
void test_segment() {
std::cout << "Testing segment...\n";
std::cout << "Testing segment..." << std::endl;
// 1. Zero Initialization
// 1. Zero Initialization & String Overload Construction
segment s1;
for (u8 i = 0; i < segment::digit_count; ++i) {
assert(s1[i] == 0 && "Default segment must be zero-initialized");
}
std::cout << " [Init] Default segment value: " << static_cast<std::string>(s1) << std::endl;
// 2. Set and Get Digits (Nibble Packing)
segment s2;
s2.set(0, 5); // Even index (low nibble)
s2.set(1, 9); // Odd index (high nibble)
s2.set(63, 7); // Boundary test
// Initialize using string overload
std::string init_str = "1234567890123456789012345678901234567890123456789012345678901234";
segment s_str(init_str);
std::cout << " [Init] Constructed from string: " << static_cast<std::string>(s_str) << std::endl;
assert(s_str[0] == 4 && "Lowest digit (LSB) must match string end");
assert(s_str[63] == 1 && "Highest digit (MSB) must match string start");
assert(s2[0] == 5 && "Failed to set/get low nibble");
assert(s2[1] == 9 && "Failed to set/get high nibble");
// 2. Set, Get & String Serialization Round-trip
segment s2("95");
std::cout << " [Set/Get] Initial s2: " << static_cast<std::string>(s2) << std::endl;
s2.set(63, 7); // Set boundary MSB digit
std::cout << " [Set/Get] After setting MSB (index 63 = 7): " << static_cast<std::string>(s2) << std::endl;
assert(s2[0] == 5 && "Failed to get low nibble from string");
assert(s2[1] == 9 && "Failed to get high nibble from string");
assert(s2[63] == 7 && "Failed boundary set/get");
// 3. Fast Byte Comparison
segment a, b;
a.set(10, 5);
b.set(10, 3);
segment a("500");
segment b("300");
std::cout << " [Compare] Comparing " << static_cast<std::string>(a) << " vs " << static_cast<std::string>(b) << ": " << int(a.compare(b)) << std::endl;
assert(a.compare(b) == 1 && "a should be greater than b");
assert(b.compare(a) == -1 && "b should be less than a");
b.set(10, 5);
assert(a.compare(b) == 0 && "a and b should be equal");
segment b_eq("500");
assert(a.compare(b_eq) == 0 && "a and b should be equal");
// High order takes precedence over low order in compare
a.set(0, 9); // Low order high value
b.set(1, 1); // High order low value
assert(a.compare(b) == -1 && "Higher-order digit must dominate comparison");
segment low_high_val("09");
segment high_low_val("10");
std::cout << " [Compare] Dominance test: " << static_cast<std::string>(low_high_val) << " vs " << static_cast<std::string>(high_low_val) << ": " << int(low_high_val.compare(high_low_val)) << std::endl;
assert(low_high_val.compare(high_low_val) == -1 && "Higher-order digit must dominate comparison");
// 4. Branchless Segment Addition & Subtraction
segment x, y;
x.set(0, 7);
y.set(0, 5);
// 4. Calculations: Addition
segment add_a("9999");
segment add_b("1");
segment n1000("10000");
std::cout << " [Add] Computing " << static_cast<std::string>(add_a) << " + " << static_cast<std::string>(add_b);
u8 carry = add_a.add(add_b, 0);
std::cout << " = " << static_cast<std::string>(add_a) << " (carry: " << static_cast<int>(carry) << ")" << std::endl;
assert(carry == 0 && "Addition within digit_count must not produce carry out");
assert(add_a.compare(n1000) == 0 && "9999 + 1 must equal 10000");
u8 carry = x.add(y, 0);
segment x("7");
segment y("5");
std::cout << " [Add] Computing " << static_cast<std::string>(x) << " + " << static_cast<std::string>(y);
carry = x.add(y, 0);
std::cout << " = " << static_cast<std::string>(x) << " (carry: " << static_cast<int>(carry) << ")" << std::endl;
assert(x[0] == 2 && "7 + 5 % 10 should be 2");
assert(carry == 1 && "7 + 5 should produce carry 1");
assert(carry == 0 && "No segment overflow for 12");
assert(x.compare(segment(12)) == 0 && "String conversion must yield '12'");
// 7 - 9 with dummy-10 borrow
segment sub1, sub2;
sub1.set(0, 7);
sub2.set(0, 9);
u8 borrow = sub1.sub(sub2, 0);
// 5. Calculations: Subtraction
segment sub_a("1000");
segment sub_b("1");
segment n999("999");
std::cout << " [Sub] Computing " << static_cast<std::string>(sub_a) << " - " << static_cast<std::string>(sub_b);
u8 borrow = sub_a.sub(sub_b, 0);
std::cout << " = " << static_cast<std::string>(sub_a) << " (borrow: " << static_cast<int>(borrow) << ")" << std::endl;
assert(borrow == 0 && "Valid subtraction must not generate final borrow");
assert(sub_a.compare(n999) == 0 && "1000 - 1 must equal 999");
segment sub1("7");
segment sub2("9");
std::cout << " [Sub] Computing " << static_cast<std::string>(sub1) << " - " << static_cast<std::string>(sub2);
borrow = sub1.sub(sub2, 0);
std::cout << " = " << static_cast<std::string>(sub1) << " (borrow: " << static_cast<int>(borrow) << ")" << std::endl;
assert(sub1[0] == 8 && "7 - 9 with borrow should yield 8");
assert(borrow == 1 && "7 - 9 should produce borrow 1");
assert(borrow == 1 && "7 - 9 must produce borrow 1");
std::cout << " [PASS] segment tests passed!\n";
// 6. Calculations: Multiplication (mul_long)
segment m1("123456789");
segment m2("987654321");
segment m_res("121932631112635269");
std::cout << " [Mul] Computing " << static_cast<std::string>(m1) << " * " << static_cast<std::string>(m2) << std::endl;
std::array<segment, 2> res = m1.mul_long(m2);
std::cout << " Low Segment : " << static_cast<std::string>(res[0]) << std::endl;
std::cout << " High Segment: " << static_cast<std::string>(res[1]) << std::endl;
assert(res[0].compare(m_res) == 0 && "Mul low segment mismatch");
assert(res[1].isZero() && "Mul high segment should be 0 for small values");
std::cout << " [PASS] segment tests passed!" << std::endl;
}
/* * /
void test_chain() {
std::cout << "Testing chain...\n";
// 1. Constructors & Invariants
chain c1; // Default constructor
chain c2(12345); // Value constructor
chain c1;
chain c2(12345);
assert(static_cast<std::string>(c2) == "12345" && "Chain value constructor failed");
// 2. Addition without Growth
chain num1(150);
chain num2(250);
// 2. Addition & Chain Growth
chain num1("150");
chain num2("250");
num1.add(num2);
// Assumes operator std::string() or describe() formatted check
// e.g., 150 + 250 = 400
assert(static_cast<std::string>(num1) == "400" && "150 + 250 must equal 400");
// 3. Subtraction & Borrow Ripple
chain a(1000);
chain b(1);
// 3. Subtraction & Borrow Ripple Across Links
chain a("10000000000000000000000000000000000000000000000000000000000000000");
chain b("1");
a.sub(b);
// Result should be 999
std::string expected_nines(64, '9');
assert(static_cast<std::string>(a) == expected_nines && "Borrow ripple across digits failed");
// 4. Shrink Verification (Trailing zeros)
chain zero_chain(0);
// Force growth then test shrinking back down
zero_chain.add(chain(0));
zero_chain.add(chain(0));
assert(zero_chain.shrink() == false && "Cannot shrink below 1 link invariant");
// 5. Copy Constructor & Copy Assignment
chain copy_src(9999);
// 5. Deep Copy & Independence Verification
chain copy_src("9999");
chain copy_dst = copy_src;
// Mutate source to verify deep copy
copy_src.add(chain(1));
// 6. Move Constructor & Move Assignment
chain move_src(8888);
copy_src.add(chain("1"));
assert(static_cast<std::string>(copy_src) == "10000" && "Source should be updated to 10000");
assert(static_cast<std::string>(copy_dst) == "9999" && "Destination must retain original copy value 9999");
// 6. Move Constructor Verification
chain move_src("8888");
chain move_dst = std::move(move_src);
// Verify move_src is left in a valid empty/default state without memory leaks
assert(static_cast<std::string>(move_dst) == "8888" && "Moved target must inherit value");
std::cout << " [PASS] chain tests passed!\n";
}
// */
int main() {
test_segment();
test_chain();
//test_chain();
std::cout << "\nAll unit tests passed successfully!\n";
return 0;
}
}
+220 -143
View File
@@ -8,182 +8,152 @@ namespace bigmath {
chain::link::link(const segment& _s) : s(_s) {}
chain::chain() noexcept : chain(0) {}
// ------------------- //
chain::chain(u64 n) noexcept {
start = new link(n);
middle = start;
end = start;
chain::chain_iterator::chain_iterator(chain::link* ptr) : current_(ptr) {}
// Dereference operators
chain::chain_iterator::reference chain::chain_iterator::operator*() const {
return current_->s;
}
chain::chain_iterator::pointer chain::chain_iterator::operator->() const {
return &(current_->s);
}
// Pre-increment: ++it
chain::chain_iterator& chain::chain_iterator::operator++() {
if (current_) current_ = current_->next;
return *this;
}
// Post-increment: it++
chain::chain_iterator chain::chain_iterator::operator++(int) {
chain_iterator temp = *this;
++(*this);
return temp;
}
// Comparison operators
bool chain::chain_iterator::operator==(const chain_iterator& other) const {
return current_ == other.current_;
}
bool chain::chain_iterator::operator!=(const chain_iterator& other) const {
return current_ != other.current_;
}
// Utility accessor to get the underlying Node pointer
chain::link* chain::chain_iterator::node() const { return current_; }
// --- Chain Constructors & Destructor ---
// Default Constructor: 1 element chain where all pointers reference the same link
chain::chain() noexcept {
_start = new link();
_mid_0 = _start;
_mid_1 = _start;
_end = _start;
}
// Value Constructor (u64): Initializes single link with value n
chain::chain(u64 n) noexcept {
_start = new link(n);
_mid_0 = _start;
_mid_1 = _start;
_end = _start;
}
chain::chain(std::string_view str) noexcept {
// TODO
}
// Destructor: Deallocates all doubly-linked nodes in sequence
chain::~chain() noexcept {
link* current = _start;
while (current != nullptr) {
link* next_node = current->next;
delete current;
current = next_node;
}
_start = _mid_0 = _mid_1 = _end = nullptr;
}
// Copy Constructor: Deep copy of the doubly-linked chain and recalculation of midpoints
chain::chain(const chain& c) noexcept {
// May be unspecified after a move
if (!c.start) return;
// 1. Copy the head node
start = new link(c.start->s);
link* src_curr = c.start->next;
link* dst_prev = start;
// Track the middle node offset in the source chain
link* src_middle = c.middle;
if (c.start == src_middle) {
middle = start;
if (!c._start) {
_start = _mid_0 = _mid_1 = _end = nullptr;
return;
}
// 2. Traverse and deep-copy remaining nodes
while (src_curr) {
link* new_node = new link(src_curr->s);
link* src_curr = c._start;
_start = new link(src_curr->s);
link* dst_prev = _start;
src_curr = src_curr->next;
while (src_curr != nullptr) {
link* new_node = new link(src_curr->s);
dst_prev->next = new_node;
new_node->prev = dst_prev;
if (src_curr == src_middle) {
middle = new_node;
}
dst_prev = new_node;
src_curr = src_curr->next;
}
_end = dst_prev;
end = dst_prev;
// Recalculate middle pointers using the helper function
auto [m0, m1] = findMidpoints(_start, _end);
_mid_0 = m0;
_mid_1 = m1;
}
// Move Constructor: Steals ownership of resources and invalidates source
chain::chain(chain&& c) noexcept
: start(c.start), middle(c.middle), end(c.end) {
// Reset source object
c.start = c.middle = c.end = nullptr;
: _start(c._start), _mid_0(c._mid_0), _mid_1(c._mid_1), _end(c._end) {
c._start = nullptr;
c._mid_0 = nullptr;
c._mid_1 = nullptr;
c._end = nullptr;
}
// Copy Assignment: Uses Copy-and-Swap idiom for exception safety and cleanliness
chain& chain::operator=(const chain& c) noexcept {
if (this != &c) {
chain temp(c);
std::swap(start, temp.start);
std::swap(middle, temp.middle);
std::swap(end, temp.end);
std::swap(_start, temp._start);
std::swap(_mid_0, temp._mid_0);
std::swap(_mid_1, temp._mid_1);
std::swap(_end, temp._end);
}
return *this;
}
// Move Assignment: Swaps pointers with move target to reuse/clean resources safely
chain& chain::operator=(chain&& c) noexcept {
if (this != &c) {
chain temp(std::move(*this));
start = c.start; middle = c.middle; end = c.end;
c.start = c.middle = c.end = nullptr;
std::swap(_start, c._start);
std::swap(_mid_0, c._mid_0);
std::swap(_mid_1, c._mid_1);
std::swap(_end, c._end);
}
return *this;
}
chain::~chain() noexcept {
link* ptr = start;
link* ptr_next;
while (ptr) {
ptr_next = ptr->next;
delete ptr;
ptr = ptr_next;
void chain::ensureValid() noexcept {
if (!_start) {
_start = new link();
_mid_0 = _start;
_mid_1 = _start;
_end = _start;
}
start = middle = end = nullptr;
}
chain::operator bool() const noexcept {
return start && middle && end;
}
std::array<chain::link*, 2> chain::split() const noexcept {
// find link* pointers by walking to the middle
if(!start) return { nullptr, nullptr };
if(start == middle) return { start, end };
// <- L [* * *] [{*} * * *] R <- //
chain::link* a_left = end;
chain::link* a_mid = end;
chain::link* a_right = middle->next;
// this must be odd
while(a_left != a_right) {
a_mid = a_left->prev;
a_left = a_left->prev;
a_right = a_right->next;
}
chain::link* b_left = middle;
chain::link* b_right = start;
// this must be even
while(b_left->next != b_right) {
b_left = b_left->prev;
b_right = b_right->next;
}
return { b_left, a_mid };
}
void chain::grow() noexcept {
// (*) -> (*) [{*} *]
// * (*) * -> * (*) {*} [* *]
// * * (*) * * -> * * (*) {*} * [* *]
// create the data
link* nl0 = new link();
link* nl1 = new link();
// setup pointers
end->next = nl0;
nl0->prev = end;
nl0->next = nl1;
nl1->prev = nl0;
// move over one step the middle
middle = middle->next;
// set new end pointer
end = nl1;
}
bool chain::shrink() noexcept {
// Cannot shrink if chain has fewer than 3 links,
// as removing 2 would leave fewer than 1 link or break invariants.
if (end == nullptr || end->prev == nullptr || end->prev->prev == nullptr) {
return false;
}
link* l1 = end; // Most significant link
link* l0 = end->prev; // Second most significant link
// Only shrink if BOTH trailing segments are completely zero
if (!l1->s.isZero() || !l0->s.isZero()) {
return false;
}
// New end will be the link right before l0
link* new_end = l0->prev;
new_end->next = nullptr;
// Shift middle pointer left by one step (reversing grow())
middle = middle->prev;
// Update end pointer
end = new_end;
// Free the two removed trailing links
delete l0;
delete l1;
return true;
}
// ------------------- //
u64 chain::count(u64 until) const noexcept {
u64 n = 1;
link* ptr = start;
u64 n = 0;
link* ptr = _start;
while (ptr->next && n < until) {
while (ptr && n < until) {
n++;
ptr = ptr->next;
}
@@ -191,13 +161,118 @@ namespace bigmath {
return n;
}
/**
* Finds m0 and m1 given the endpoints 'a' and 'b' of a power-of-two chain.
*
* @param a Start pointer of the chain segment
* @param b End pointer of the chain segment
*/
pair<chain::link*, chain::link*> chain::findMidpoints(chain::link* a, chain::link* b) {
// Base case: 1-element chain (2^0)
if (a == b) {
return { a, a };
}
link* slow = a;
link* fast = a;
// Fast moves 2 steps while slow moves 1 step.
// Stops when fast reaches 'b' or 'b's predecessor.
while (fast != b && fast->next != b) {
slow = slow->next;
fast = fast->next->next;
}
link* m0 = slow;
link* m1 = slow->next;
return { m0, m1 };
}
chain::chain_iterator chain::begin() const {
return chain::chain_iterator(_start);
}
chain::chain_iterator chain::end() const {
return chain::chain_iterator(nullptr);
}
void chain::grow() noexcept {
// Base case: handle empty/uninitialized chain
if (!_start) {
_start = new link();
_mid_0 = _start;
_mid_1 = _start;
_end = _start;
return;
}
// 1. Allocate the first link of the new upper half
link* new_head = new link();
link* new_tail = new_head;
// 2. Traverse the current chain [start->next ... end]
// For every existing link, construct one corresponding new link
link* curr = _start->next;
while (curr != nullptr) {
link* next_link = new link();
new_tail->next = next_link;
next_link->prev = new_tail;
new_tail = next_link;
curr = curr->next;
}
// 3. Connect old lower half to new upper half
_end->next = new_head;
new_head->prev = _end;
// 4. Update core pointers for double length [start ... new_tail]
_mid_0 = _end;
_mid_1 = new_head;
_end = new_tail;
}
bool chain::shrink() noexcept {
// Base case: Cannot shrink a 1-node chain (must remain at least 2^0)
if (!_start || _start == _end) {
return false;
}
// Step 1: Disconnect top half [_mid_1, _end] from bottom half [_start, _mid_0]
link* upper_curr = _mid_1;
if (upper_curr) {
upper_curr->prev = nullptr;
}
_mid_0->next = nullptr;
// Step 2: Delete top half nodes
while (upper_curr != nullptr) {
link* next_node = upper_curr->next;
delete upper_curr;
upper_curr = next_node;
}
// Step 3: Lower half becomes full active chain range
_end = _mid_0;
// Step 4: Recalculate midpoints for the halved range
auto [m0, m1] = findMidpoints(_start, _end);
_mid_0 = m0;
_mid_1 = m1;
return true;
}
// ============================================================================ //
// ADDITION / SUBTRACTION //
// ============================================================================ //
void chain::add(const chain& c) noexcept {
link* curr_a = start;
const link* curr_b = c.start;
ensureValid();
link* curr_a = _start;
const link* curr_b = c._start;
u8 carry = 0;
// Traverse both chains from least significant link (start) to most significant
@@ -229,8 +304,10 @@ namespace bigmath {
}
void chain::sub(const chain& c) noexcept {
link* curr_a = start;
const link* curr_b = c.start;
ensureValid();
link* curr_a = _start;
const link* curr_b = c._start;
u8 borrow = 0;
// Traverse both chains from least significant link (start) to most significant
+62 -8
View File
@@ -29,16 +29,60 @@ namespace bigmath {
};
public:
class chain_iterator {
public:
// Standard iterator traits required by C++ STL algorithms
using iterator_category = std::forward_iterator_tag;
using value_type = segment;
using difference_type = std::ptrdiff_t;
using pointer = segment*;
using reference = segment&;
private:
link* current_;
public:
chain_iterator(link* ptr = nullptr);
// Dereference operators
reference operator*() const;
pointer operator->() const;
// Pre-increment: ++it
chain_iterator& operator++();
// Post-increment: it++
chain_iterator operator++(int);
// Comparison operators
bool operator==(const chain_iterator& other) const;
bool operator!=(const chain_iterator& other) const;
// Utility accessor to get the underlying Node pointer
link* node() const;
};
private:
/// @brief The start of the number.
link* start = nullptr;
/// @brief The start (RIGHT) of the number.
link* _start = nullptr;
/// @brief The middle of the number.
link* middle = nullptr;
/// @brief The middle right of the number.
link* _mid_0 = nullptr;
/// @brief The end of the number.
link* end = nullptr;
/// @brief The middle left of the number.
link* _mid_1 = nullptr;
/// @brief The end (LEFT) of the number.
link* _end = nullptr;
public:
@@ -46,6 +90,8 @@ namespace bigmath {
chain(u64 n) noexcept;
chain(std::string_view str) noexcept;
chain(const chain& c) noexcept;
chain(chain&& c) noexcept;
@@ -56,6 +102,10 @@ namespace bigmath {
~chain() noexcept;
private:
void ensureValid() noexcept;
public:
void add(const chain& c) noexcept;
@@ -78,9 +128,9 @@ namespace bigmath {
public:
operator bool() const noexcept;
chain_iterator begin() const;
std::array<link*, 2> split() const noexcept;
chain_iterator end() const;
/**
* Grows the chain by two links.
@@ -93,6 +143,10 @@ namespace bigmath {
u64 count(u64 until = static_cast<u64>(-1)) const noexcept;
private:
static pair<link*, link*> findMidpoints(link* a, link* b);
};
}
View File
View File
+46 -20
View File
@@ -1,5 +1,7 @@
#include "segment.hpp"
#include <bigmath/chain/segment_ops.hpp>
#include <bigmath/util/describe.hpp>
namespace bigmath {
@@ -18,6 +20,36 @@ namespace bigmath {
}
}
segment::segment(std::string_view str) : segment() {
// Trim optional leading zeroes or empty strings
while (!str.empty() && str.front() == '0') {
str.remove_prefix(1);
}
if (str.empty()) {
return; // Segment remains 0
}
// Populate digits from right to left (LSD -> MSD)
u32 len = static_cast<u32>(str.length());
u32 max_digits = std::min(len, static_cast<u32>(digit_count));
for (u32 i = 0; i < max_digits; ++i) {
// Read character from rightmost available digit
char ch = str[len - 1 - i];
if (ch >= '0' && ch <= '9') {
u8 digit = static_cast<u8>(ch - '0');
// Store packed digit into segment (low nibble for even, high for odd)
u32 byte_idx = i >> 1;
u32 shift = (i & 1) * 4;
digits[byte_idx] |= (digit << shift);
}
}
}
segment::segment(const segment& s) {
std::copy(s.begin(), s.end(), digits);
}
@@ -34,32 +66,27 @@ namespace bigmath {
void segment::set(u8 i, u8 n) {
u8 j = i >> 1;
u8 k = (i & 1) * 4;
u8 b = digits[j];
u8 b = digits[j];
b &= 0xF0u >> k;
n &= 0xFu;
digits[j] = b | (n << k);
}
u8 segment::add(const segment& s, u8 carry) {
u8 m = carry;
u8 i = segment::digit_count;
while (i-- > 0) {
m = static_cast<u8>(operator[](i) + s[i] + m);
set(i, m % 10);
m = m / 10;
}
return m;
return bigmath::add(digits, s.digits, digits, carry, byte_count);
}
u8 segment::sub(const segment& s, u8 borrow) {
u8 m = borrow;
u8 i = segment::digit_count;
while (i-- > 0) {
m = static_cast<u8>(10 + operator[](i) - s[i] - m);
set(i, m % 10);
m = 1 - m / 10;
}
return m;
return bigmath::sub(digits, s.digits, digits, borrow, byte_count);
}
std::array<segment, 2> segment::mul_long(const segment& s) const {
segment low;
segment high;
bigmath::mul_long(digits, s.digits, high.digits, low.digits, byte_count);
return { low, high };
}
i8 segment::compare(const segment& s) const {
@@ -75,9 +102,8 @@ namespace bigmath {
bool segment::isZero() const noexcept {
// advances twice as fast
u8 i = segment::byte_count;
while (i-- > 0) {
if(digits[i]) return false;
for (u8 i = 0; i < segment::byte_count; i++) {
if (digits[i]) return false;
}
return true;
}
+20
View File
@@ -2,6 +2,8 @@
#include <bigmath/BigMath.hpp>
#include <bigmath/util/util.hpp>
namespace bigmath {
/**
@@ -23,6 +25,8 @@ namespace bigmath {
segment(u64 n);
segment(std::string_view str);
segment(const segment& s);
~segment();
@@ -33,10 +37,26 @@ namespace bigmath {
void set(u8 i, u8 n);
segment extract(u8 from, u8 to, u8 offset) const;
void shift_right(u8 places);
void shift_left(u8 places);
u8 add(const segment& s, u8 carry);
u8 sub(const segment& s, u8 borrow);
std::array<segment, 2> mul_long(const segment& s) const;
std::array<segment, 2> mul_karatsuba(const segment& s) const;
div_result<segment> div_long(const segment& s) const;
segment div_newtonraphson(const segment& s) const;
segment mod_long(const segment& s) const;
i8 compare(const segment& s) const;
bool isZero() const noexcept;
+137
View File
@@ -0,0 +1,137 @@
#include "segment.hpp"
namespace bigmath {
u8 add(u8* a, const u8* b, u8* c, u8 carry, u8 length) {
for (u8 i = 0; i < length; ++i) {
u8 val_a = a[i];
u8 val_b = b[i];
// 1. Raw binary addition including incoming carry
u16 raw_sum = static_cast<u16>(static_cast<u16>(val_a) + val_b + carry);
// 2. Extract nibbles of inputs
u8 low_a = val_a & 0x0F;
u8 low_b = val_b & 0x0F;
// 3. Compute lower nibble sum to detect half-carry (if lower sum > 15)
u8 low_sum = static_cast<u8>(low_a + low_b + carry);
// 4. Branchless condition checks (evaluates to 1 if true, 0 if false):
// Low correction needed if: low_sum > 9 OR half-carry occurred (low_sum > 15)
u8 fix_low_mask = static_cast<u8>((low_sum > 9) || (low_sum > 15));
// High correction needed if raw_sum (after optional low adjustment) > 0x99 or byte carried
u8 fix_high_mask = static_cast<u8>((raw_sum + (fix_low_mask * 0x06)) > 0x99);
// 5. Construct BCD correction value branchlessly (0x00, 0x06, 0x60, or 0x66)
u8 correction = static_cast<u8>((fix_low_mask * 0x06) + (fix_high_mask * 0x60));
// 6. Apply BCD adjustment
u16 adjusted_sum = raw_sum + correction;
// 7. Store packed BCD byte and calculate outgoing carry
c[i] = static_cast<u8>(adjusted_sum & 0xFF);
carry = static_cast<u8>(adjusted_sum >> 8);
}
return carry;
}
u8 sub(const u8* a, const u8* b, u8* c, u8 borrow, u8 length) {
for (u8 i = 0; i < length; ++i) {
u8 val_a = a[i];
u8 val_b = b[i];
// 1. Extract nibbles
u8 low_a = val_a & 0x0F;
u8 low_b = val_b & 0x0F;
// 2. Compute lower nibble subtraction
// Add 10 to ensure positive result; if low_a < (low_b + borrow), a borrow occurs
u16 low_diff = static_cast<u16>(10 + low_a - low_b - borrow);
u8 low_digit = static_cast<u8>(low_diff % 10);
u8 low_borrow = static_cast<u8>(1 - (low_diff / 10)); // 1 if borrowed, 0 otherwise
// 3. Extract high nibbles
u8 high_a = val_a >> 4;
u8 high_b = val_b >> 4;
// 4. Compute upper nibble subtraction including lower nibble borrow
u16 high_diff = static_cast<u16>(10 + high_a - high_b - low_borrow);
u8 high_digit = static_cast<u8>(high_diff % 10);
borrow = static_cast<u8>(1 - (high_diff / 10)); // Next byte's incoming borrow
// 5. Pack resulting BCD digits back into single byte
c[i] = low_digit | (high_digit << 4);
}
return borrow;
}
void mul_long(const u8* a, const u8* b, u8* c_high, u8* c_low, u8 length) {
const u32 blen = length;
const u32 blen2 = blen * 2;
const u32 blen4 = blen * 4;
// Step 1: Unpack raw BCD bytes into raw digit buffers (0-9)
u8 digits_a[blen2];
u8 digits_b[blen2];
for (u32 i = 0; i < blen; ++i) {
u8 byte_a = a[i];
u8 byte_b = b[i];
digits_a[i * 2] = byte_a & 0x0F;
digits_a[i * 2 + 1] = byte_a >> 4;
digits_b[i * 2] = byte_b & 0x0F;
digits_b[i * 2 + 1] = byte_b >> 4;
}
// Step 2: Accumulate digit products into temporary array of size blen4
u16 accum[blen4] = {};
for (u32 i = 0; i < blen2; ++i) {
u8 d1 = digits_a[i];
if (d1 == 0) continue; // Skip zeros
for (u32 j = 0; j < blen2; ++j) {
u8 d2 = digits_b[j];
accum[i + j] += static_cast<u16>(static_cast<u16>(d1) * static_cast<u16>(d2));
}
}
// Step 3: Base-10 Carry Propagation and Direct Packed BCD Storage
u32 carry = 0;
u8 out_digits[blen4];
for (u32 i = 0; i < blen4; ++i) {
u32 sum = accum[i] + carry;
out_digits[i] = static_cast<u8>(sum % 10);
carry = sum / 10;
}
// Step 4: Pack result digits directly into low and high segment bytes
for (u32 i = 0; i < blen; ++i) {
u8 low_digit0 = out_digits[i * 2];
u8 low_digit1 = out_digits[i * 2 + 1];
u8 high_digit0 = out_digits[i * 2 + blen2];
u8 high_digit1 = out_digits[i * 2 + blen2 + 1];
// No longer using the object version,
// so we set the bytes ourselves!
c_low[i] = low_digit0 | low_digit1 << 4;
//low.set(static_cast<u8>(i * 2), low_digit0);
//low.set(static_cast<u8>(i * 2 + 1), low_digit1);
c_high[i] = high_digit0 | high_digit1 << 4;
//high.set(static_cast<u8>(i * 2), high_digit0);
//high.set(static_cast<u8>(i * 2 + 1), high_digit1);
}
// end!
}
}
+19
View File
@@ -0,0 +1,19 @@
#pragma once
#include <bigmath/BigMath.hpp>
namespace bigmath {
u8 add(u8* a, const u8* b, u8* c, u8 carry, u8 length);
u8 sub(u8* a, const u8* b, u8* c, u8 borrow, u8 length);
void mul_long(const u8* a, const u8* b, u8* c_high, u8* c_low, u8 length);
void mul_karatsuba(const u8* a, const u8* b, u8* c_high, u8* c_low, u8 length);
void div_long(const u8* a, const u8* b, u8* quot, u8* rem, u8 length);
void div_newtonraphson(const u8* a, const u8* b, u8* quot, u8 length);
}
View File
View File
View File
View File
View File
View File
View File
View File
View File
+2
View File
@@ -12,4 +12,6 @@ namespace bigmath {
u64 ipow(u64 b, u64 p);
using std::pair;
}