#![allow(clippy::needless_range_loop)]
use crate::algebra::finite_field::Field;
use crate::algebra::finite_field::GF256;
pub struct Vector;
impl Vector {
pub fn add_inplace(a: &mut [u8], b: &[u8]) {
assert_eq!(a.len(), b.len());
for i in 0..a.len() {
a[i] ^= b[i];
}
}
pub fn multiply_alpha_inplace(field: &GF256, result: &mut [u8]) {
for elem in result.iter_mut() {
*elem = field.mul_alpha(*elem);
}
}
pub fn scalar_vector_multiply_inplace<F: Field>(field: &F, scalar: u8, result: &mut [u8]) {
for elem in result.iter_mut() {
*elem = field.mul(scalar, *elem);
}
}
}
pub struct Matrix;
impl Matrix {
pub fn multiply<F: Field>(field: &F, a: &[Vec<u8>], b: &[Vec<u8>]) -> Vec<Vec<u8>> {
let m = a.len();
let n = if m > 0 { a[0].len() } else { 0 };
if n != b.len() {
panic!("The number of columns in A must be equal to the number of rows in B");
}
let p = if n > 0 { b[0].len() } else { 0 };
let mut result = vec![vec![0u8; p]; m];
for (i, row) in result.iter_mut().enumerate() {
for (j, &a_ij) in a[i].iter().take(n).enumerate() {
for (k, &val) in b[j].iter().enumerate() {
row[k] ^= field.mul(a_ij, val);
}
}
}
result
}
pub fn permute_rows_inplace(a: &mut [Vec<u8>], p: &[usize]) {
let n = a.len();
let mut visited = vec![false; n];
for i in 0..n {
if visited[i] || p[i] == i {
continue;
}
let mut j = i;
while !visited[j] {
visited[j] = true;
let k = p[j];
if k != i {
a.swap(j, k);
}
j = k;
}
}
}
pub fn lu_decomp<F: Field>(field: &F, a: &mut [Vec<u8>]) -> (Vec<usize>, usize) {
let m = a.len();
let n = if m > 0 { a[0].len() } else { 0 };
let mut p: Vec<usize> = (0..m).collect();
let mut i = 0;
for j in 0..n {
let mut pivot_found = false;
for k in i..a.len() {
if a[k][j] != 0 {
p.swap(i, k);
a.swap(i, k);
pivot_found = true;
break;
}
}
if pivot_found {
for k in i + 1..m {
let l = field.divide(a[k][j], a[i][j]);
a[k][j] = 0;
a[k][i] = l;
for col in (j + 1)..n {
a[k][col] = field.add(a[k][col], field.mul(l, a[i][col]));
}
}
i += 1;
if i == m {
break;
}
}
}
(p, i)
}
pub fn lu_decomp_incr<F: Field>(
field: &F,
a: &mut [Vec<u8>],
q: &mut [usize],
r: usize,
) -> (Vec<usize>, usize) {
let m = a.len();
let n = if m > 0 { a[0].len() } else { 0 };
let mut p = (0..m).collect::<Vec<_>>();
for i in 0..r {
for k in r..m {
let l = field.divide(a[k][q[i]], a[i][q[i]]);
a[k][q[i]] = l;
for col in i + 1..n {
a[k][q[col]] = field.add(a[k][q[col]], field.mul(l, a[i][q[col]]));
}
}
}
let mut i = r;
for j in r..n {
let mut pivot_found = false;
for k in i..a.len() {
if a[k][q[j]] != 0 {
p.swap(i, k);
a.swap(i, k);
pivot_found = true;
break;
}
}
if pivot_found {
q.swap(i, j);
for k in i + 1..m {
let l = field.divide(a[k][q[i]], a[i][q[i]]);
a[k][q[i]] = l;
for col in i + 1..n {
a[k][q[col]] = field.add(a[k][q[col]], field.mul(l, a[i][q[col]]));
}
}
i += 1;
if i == m {
break;
}
}
}
(p, i)
}
pub fn lu_decomp_incr_binary(
a: &mut [Vec<u8>],
q: &mut [usize],
r: usize,
) -> (Vec<usize>, usize) {
let m = a.len();
let n = if m > 0 { a[0].len() } else { 0 };
let mut p = (0..m).collect::<Vec<_>>();
for i in 0..r {
assert_eq!(a[i][q[i]], 1);
for k in r..m {
if a[k][q[i]] == 1 {
for col in i + 1..n {
a[k][q[col]] ^= a[i][q[col]];
}
}
}
}
let mut i = r;
for j in r..n {
let mut pivot_found = false;
for k in i..a.len() {
if a[k][q[j]] != 0 {
p.swap(i, k);
a.swap(i, k);
pivot_found = true;
break;
}
}
if pivot_found {
q.swap(i, j);
assert_eq!(a[i][q[i]], 1);
for k in i + 1..m {
if a[k][q[i]] == 1 {
for col in i + 1..n {
a[k][q[col]] ^= a[i][q[col]];
}
}
}
i += 1;
if i == m {
break;
}
}
}
(p, i)
}
pub fn lu_decomp_incr_mixed(
field: &GF256,
a: &mut [Vec<u8>],
q: &mut [usize],
r: usize,
binary_rows: usize,
) -> (Vec<usize>, usize) {
let m = a.len();
let n = if m > 0 { a[0].len() } else { 0 };
let mut p = (0..m).collect::<Vec<_>>();
for i in 0..r {
let binary_pivot = i < binary_rows && a[i][q[i]] == 1;
for k in r..m {
let akqi = a[k][q[i]];
if akqi == 0 {
continue;
}
let pivot = a[i][q[i]];
if pivot == 0 {
continue;
}
let l = if binary_pivot {
akqi
} else {
field.divide(akqi, pivot)
};
a[k][q[i]] = l;
if l == 0 {
continue;
}
if binary_pivot && l == 1 {
for col in i + 1..n {
let u = a[i][q[col]];
if u != 0 {
a[k][q[col]] ^= u;
}
}
} else {
for col in i + 1..n {
let u = a[i][q[col]];
if u != 0 {
a[k][q[col]] = field.add(a[k][q[col]], field.mul(l, u));
}
}
}
}
}
let mut i = r;
for j in r..n {
let mut pivot_found = false;
for k in i..a.len() {
if a[k][q[j]] != 0 {
p.swap(i, k);
a.swap(i, k);
pivot_found = true;
break;
}
}
if pivot_found {
q.swap(i, j);
for k in i + 1..m {
let l = field.divide(a[k][q[i]], a[i][q[i]]);
a[k][q[i]] = l;
if l == 0 {
continue;
}
if l == 1 {
for col in i + 1..n {
let u = a[i][q[col]];
if u != 0 {
a[k][q[col]] ^= u;
}
}
} else {
for col in i + 1..n {
let u = a[i][q[col]];
if u != 0 {
a[k][q[col]] = field.add(a[k][q[col]], field.mul(l, u));
}
}
}
}
i += 1;
if i == m {
break;
}
}
}
(p, i)
}
}
pub struct LinearSys;
impl LinearSys {
pub fn lin_solve<F: Field>(
field: &F,
a: &mut [Vec<u8>],
b: &mut [Vec<u8>],
) -> Result<(), String> {
let m = a.len();
let n = if m > 0 { a[0].len() } else { 0 };
let (p, r) = Matrix::lu_decomp(field, a);
if r == n {
let a_n = &mut a[..n];
Matrix::permute_rows_inplace(b, &p);
let b_n = &mut b[..n];
Self::lu_solve(field, a_n, b_n)?;
Ok(())
} else {
Err(format!("The matrix is not invertible, rank = {}", r))
}
}
pub fn lu_solve<F: Field>(field: &F, a: &[Vec<u8>], b: &mut [Vec<u8>]) -> Result<(), String> {
let n = a.len();
if n != b.len() {
return Err("The number of rows in A and b must be the same".to_string());
}
for j in 0..n.saturating_sub(1) {
for i in j + 1..n {
Self::combined_vector_operation_inplace(field, b, i, a[i][j], j);
}
}
for j in (0..n).rev() {
Vector::scalar_vector_multiply_inplace(field, field.inverse(a[j][j]), &mut b[j]);
for i in (0..j).rev() {
Self::combined_vector_operation_inplace(field, b, i, a[i][j], j);
}
}
Ok(())
}
pub fn lu_solve_incr<F: Field>(
field: &F,
a: &[Vec<u8>],
b: &mut [Vec<u8>],
q: &[usize],
) -> Result<(), String> {
let n = a.len();
if n != b.len() {
return Err("The number of rows in A and b must be the same".to_string());
}
for j in 0..n.saturating_sub(1) {
for i in j + 1..n {
let l = a[i][q[j]];
if l != 0 {
Self::combined_vector_operation_inplace(field, b, i, l, j);
}
}
}
for j in (0..n).rev() {
let diag = a[j][q[j]];
if diag == 0 {
return Err(format!(
"Singular matrix: diagonal element at position {j} is zero"
));
}
Vector::scalar_vector_multiply_inplace(field, field.inverse(diag), &mut b[j]);
for i in (0..j).rev() {
let l = a[i][q[j]];
if l != 0 {
Self::combined_vector_operation_inplace(field, b, i, l, j);
}
}
}
Ok(())
}
fn combined_vector_operation_inplace<F: Field>(
field: &F,
b: &mut [Vec<u8>],
i: usize,
scalar: u8,
j: usize,
) {
for k in 0..b[i].len() {
b[i][k] = field.add(b[i][k], field.mul(scalar, b[j][k]));
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_simple_linear_system() {
let field = GF256::default();
let mut a = vec![vec![1, 1], vec![2, 1]];
let x = vec![vec![3, 7, 19], vec![5, 6, 20]];
let mut b = Matrix::multiply(&field, &a, &x);
let result = LinearSys::lin_solve(&field, &mut a, &mut b);
assert!(result.is_ok());
assert_eq!(b, x);
}
#[test]
fn test_singular_matrix() {
let field = GF256::default();
let mut a = vec![
vec![1, 1],
vec![1, 1], ];
let mut b = vec![vec![3], vec![5]];
let result = LinearSys::lin_solve(&field, &mut a, &mut b);
assert!(result.is_err());
}
#[test]
fn lu_decomp_incr_mixed_matches_gf256_on_binary_prefix() {
let field = GF256::default();
let n = 8;
let binary_rows = 5;
let mut q_ref = (0..n).collect::<Vec<_>>();
let mut q_mix = q_ref.clone();
let mut a_ref = vec![
vec![1, 0, 1, 0, 0, 0, 0, 0],
vec![0, 1, 1, 0, 0, 0, 0, 0],
vec![1, 1, 0, 1, 0, 0, 0, 0],
vec![0, 0, 1, 1, 1, 0, 0, 0],
vec![1, 0, 0, 1, 0, 1, 0, 0],
vec![2, 3, 1, 0, 4, 0, 1, 0],
vec![1, 5, 0, 2, 3, 1, 1, 0],
];
let mut a_mix = a_ref.clone();
let r0 = 0;
let (p_ref, r_ref) = Matrix::lu_decomp_incr(&field, &mut a_ref, &mut q_ref, r0);
let (p_mix, r_mix) =
Matrix::lu_decomp_incr_mixed(&field, &mut a_mix, &mut q_mix, r0, binary_rows);
assert_eq!(r_ref, r_mix);
assert_eq!(p_ref, p_mix);
assert_eq!(q_ref, q_mix);
assert_eq!(a_ref, a_mix);
}
}