diff --git a/examples/simple_array_operations.jl b/examples/simple_array_operations.jl index c70a61d..3bd0c25 100644 --- a/examples/simple_array_operations.jl +++ b/examples/simple_array_operations.jl @@ -26,7 +26,7 @@ function simple_array_operations(context) sa_scalar = sa1 * 4.0 - sa_mult = sa1 * sa2 + sa_mult = sa1 .* sa2 sa_shift1 = circshift(sa1, (0, 1, 0)) sa_shift2 = circshift(sa1, (1, -1, 1)) diff --git a/examples/simple_matrix_operations.jl b/examples/simple_matrix_operations.jl index 9508183..3d67b7a 100644 --- a/examples/simple_matrix_operations.jl +++ b/examples/simple_matrix_operations.jl @@ -31,7 +31,7 @@ function simple_matrix_operations(context) sm_scalar = sm1 * 4.0 - sm_mult = sm1 * sm2 + sm_mult = sm1 .* sm2 sm_shift1 = circshift(sm1, (0, 1)) sm_shift2 = circshift(sm1, (1, -1)) diff --git a/examples/simple_real_numbers.jl b/examples/simple_real_numbers.jl index c15200d..6e5e0a2 100644 --- a/examples/simple_real_numbers.jl +++ b/examples/simple_real_numbers.jl @@ -26,7 +26,7 @@ function simple_real_numbers(context) sv_scalar = sv1 * 4.0 - sv_mult = sv1 * sv2 + sv_mult = sv1 .* sv2 sv_shift1 = circshift(sv1, -1) sv_shift2 = circshift(sv1, 2) diff --git a/src/SecureArithmetic.jl b/src/SecureArithmetic.jl index dead71f..1f13dc9 100644 --- a/src/SecureArithmetic.jl +++ b/src/SecureArithmetic.jl @@ -18,7 +18,7 @@ export generate_keys, init_multiplication!, init_rotation!, init_bootstrapping! export encrypt, decrypt, decrypt!, bootstrap! # Query crypto objects -export level, capacity +export level, capacity, resize # Memory management export release_context_memory @@ -32,5 +32,6 @@ include("auxiliary.jl") include("openfhe.jl") include("unencrypted.jl") include("arithmetic.jl") +include("linear_algebra.jl") end # module SecureArithmetic diff --git a/src/arithmetic.jl b/src/arithmetic.jl index ba3a882..81f815a 100644 --- a/src/arithmetic.jl +++ b/src/arithmetic.jl @@ -2,26 +2,92 @@ Base.:+(sa1::SecureArray{B, N}, sa2::SecureArray{B, N}) where {B, N} = add(sa1, sa2) Base.:+(sa::SecureArray{B, N}, pa::PlainArray{B, N}) where {B, N} = add(sa, pa) Base.:+(pa::PlainArray{B, N}, sa::SecureArray{B, N}) where {B, N} = add(sa, pa) +Base.:+(pa1::PlainArray{B, N}, pa2::PlainArray{B, N}) where {B, N} = add(pa1, pa2) Base.:+(sa::SecureArray, scalar::Real) = add(sa, scalar) Base.:+(scalar::Real, sa::SecureArray) = add(sa, scalar) +Base.:+(pa::PlainArray, scalar::Real) = add(pa, scalar) +Base.:+(scalar::Real, pa::PlainArray) = add(pa, scalar) # Subtract Base.:-(sa1::SecureArray{B, N}, sa2::SecureArray{B, N}) where {B, N} = subtract(sa1, sa2) Base.:-(sa::SecureArray{B, N}, pa::PlainArray{B, N}) where {B, N} = subtract(sa, pa) Base.:-(pa::PlainArray{B, N}, sa::SecureArray{B, N}) where {B, N} = subtract(pa, sa) +Base.:-(pa1::PlainArray{B, N}, pa2::PlainArray{B, N}) where {B, N} = subtract(pa1, pa2) Base.:-(sa::SecureArray, scalar::Real) = subtract(sa, scalar) Base.:-(scalar::Real, sa::SecureArray) = subtract(scalar, sa) +Base.:-(pa::PlainArray, scalar::Real) = subtract(pa, scalar) +Base.:-(scalar::Real, pa::PlainArray) = subtract(scalar, pa) # Negate Base.:-(sa::SecureArray) = negate(sa) +Base.:-(pa::PlainArray) = negate(pa) -# Multiply -Base.:*(sa1::SecureArray{B, N}, sa2::SecureArray{B, N}) where {B, N} = multiply(sa1, sa2) -Base.:*(sa::SecureArray{B, N}, pa::PlainArray{B, N}) where {B, N} = multiply(sa, pa) -Base.:*(pa::PlainArray{B, N}, sa::SecureArray{B, N}) where {B, N} = multiply(sa, pa) +# Multiply (scalar) Base.:*(sa::SecureArray, scalar::Real) = multiply(sa, scalar) Base.:*(scalar::Real, sa::SecureArray) = multiply(sa, scalar) +Base.:*(pa::PlainArray, scalar::Real) = multiply(pa, scalar) +Base.:*(scalar::Real, pa::PlainArray) = multiply(pa, scalar) + +""" + SecureArrayStyle <: Base.Broadcast.BroadcastStyle + +Custom broadcast style for [`SecureArray`](@ref) and [`PlainArray`](@ref). + +Since `SecureArray` and `PlainArray` are not `AbstractArray` subtypes and their elements +(ciphertexts) cannot be iterated individually, the standard broadcast machinery — which builds +a lazy `Broadcasted` expression tree and materializes it element-by-element — cannot be used. + +Instead, we eagerly evaluate broadcast expressions by overriding +[`Base.Broadcast.broadcasted`](https://docs.julialang.org/en/v1/base/arrays/#Base.Broadcast.broadcasted) +for specific operations, returning the computed result +directly. This is the same approach Julia Base uses for `AbstractRange` operations in +[`base/broadcast.jl`](https://github.com/JuliaLang/julia/blob/d1c37793dd2ab0de6bca636e1d7f2ceb43150a9c/base/broadcast.jl#L1176), e.g., +`broadcasted(::DefaultArrayStyle{1}, ::typeof(*), x::Number, r::LinRange)`. + +We also override +[`Base.Broadcast.broadcastable`](https://docs.julialang.org/en/v1/base/arrays/#Base.Broadcast.broadcastable) +to return the objects as-is, since the default fallback (`collect(x)`) would attempt to call +`iterate` on them. + +## Supported broadcast operations + +- `.*` (element-wise multiply): `sa .* sa`, `sa .* pa`, `pa .* sa`, `pa .* pa` + +## Example + +```julia +sa1 .* sa2 # element-wise multiply (calls `multiply`) +sa1 * sa2 # matrix multiply for SecureMatrix (calls `row_mat_times_mat`) +``` + +See also: [`SecureArray`](@ref), [`PlainArray`](@ref), `multiply` +""" +struct SecureArrayStyle <: Base.Broadcast.BroadcastStyle end +Base.Broadcast.BroadcastStyle(::Type{<:SecureArray}) = SecureArrayStyle() +Base.Broadcast.BroadcastStyle(::Type{<:PlainArray}) = SecureArrayStyle() +# Win over scalars so e.g. `sa .* 2` stays in SecureArrayStyle (compare [SparseArrays.jl/src/higherorderfns.jl](https://github.com/JuliaSparse/SparseArrays.jl/blob/84b5114372a15d05b9a9a160d36f99b9a3d3cea6/src/higherorderfns.jl#L76)) +Base.Broadcast.BroadcastStyle(s::SecureArrayStyle, ::Base.Broadcast.DefaultArrayStyle{0}) = s +# Prevent the default `broadcastable(x) = collect(x)` from calling `iterate` on ciphertexts. +Base.Broadcast.broadcastable(sa::SecureArray) = sa +Base.Broadcast.broadcastable(pa::PlainArray) = pa + +# Element-wise multiply (a .* b) +@inline Base.Broadcast.broadcasted(::SecureArrayStyle, ::typeof(*), a::SecureArray{B, N}, b::SecureArray{B, N}) where {B, N} = multiply(a, b) +@inline Base.Broadcast.broadcasted(::SecureArrayStyle, ::typeof(*), a::SecureArray{B, N}, b::PlainArray{B, N}) where {B, N} = multiply(a, b) +@inline Base.Broadcast.broadcasted(::SecureArrayStyle, ::typeof(*), a::PlainArray{B, N}, b::SecureArray{B, N}) where {B, N} = multiply(b, a) +@inline Base.Broadcast.broadcasted(::SecureArrayStyle, ::typeof(*), a::PlainArray{B, N}, b::PlainArray{B, N}) where {B, N} = multiply(a, b) +@inline Base.Broadcast.broadcasted(::SecureArrayStyle, ::typeof(*), a::SecureArray, b::Real) = multiply(a, b) +@inline Base.Broadcast.broadcasted(::SecureArrayStyle, ::typeof(*), a::Real, b::SecureArray) = multiply(b, a) # Circular shift + +function check_shifts(arr::Union{SecureArray, PlainArray}, shifts) + if length(shifts) > ndims(arr) + throw(ArgumentError("Got rotation index with length $(length(shifts)), expected $(ndims(arr))")) + elseif length(shifts) < ndims(arr) + shifts = vcat(collect(shifts), zeros(Integer, ndims(arr) - length(shifts))) + end + return shifts +end """ circshift(sa::SecureArray, shifts) @@ -36,11 +102,7 @@ Note: To precompute all required rotation indexes, use `init_rotation!`. See also: [`SecureArray`](@ref), [`init_rotation!`](@ref) """ function Base.circshift(sa::SecureArray, shifts) - if length(shifts) > ndims(sa) - throw(ArgumentError("Got rotation index with length $(length(shifts)), expected $(ndims(sa))")) - elseif length(shifts) < ndims(sa) - shifts = vcat(collect(shifts), zeros(Integer, ndims(sa) - length(shifts))) - end + shifts = check_shifts(sa, shifts) if all(shifts .% size(sa) .== 0) return sa @@ -48,3 +110,23 @@ function Base.circshift(sa::SecureArray, shifts) rotate(sa, shifts) end + +function Base.circshift(pa::PlainArray, shifts) + shifts = check_shifts(pa, shifts) + + if all(shifts .% size(pa) .== 0) + return pa + end + + data = collect(pa) + shifted = circshift(data, shifts) + PlainArray(shifted, pa.context) +end + + +# Matrix Multiplication +# For matrices in column-major order, we have to swap the arguments order +Base.:*(sm1::SecureMatrix{B, N}, sm2::SecureMatrix{B, N}) where {B, N} = SecureArithmetic.row_mat_times_mat(sm2, sm1) +Base.:*(sm1::PlainMatrix{B}, sm2::SecureMatrix{B}) where {B} = SecureArithmetic.row_mat_times_mat(sm2, sm1) +Base.:*(sm1::SecureMatrix{B}, sm2::PlainMatrix{B}) where {B} = SecureArithmetic.row_mat_times_mat(sm2, sm1) +Base.:*(sm1::PlainMatrix{B}, sm2::PlainMatrix{B}) where {B} = SecureArithmetic.row_mat_times_mat(sm2, sm1) \ No newline at end of file diff --git a/src/linear_algebra.jl b/src/linear_algebra.jl new file mode 100644 index 0000000..2e79cb1 --- /dev/null +++ b/src/linear_algebra.jl @@ -0,0 +1,58 @@ +""" +Multiply two d x d matrices in row order + +!!! Column order matrices + If `sm1` and `sm1` encode matrices ``A`` and ``B`` in column order, then `sm1` and `sm2` also encode ``A^T`` and ``B^T`` in row order. + If `sm3` encodes ``(A \\times B)^T`` in row order, then `sm3` also encodes ``A \\times B`` in column order. + Since ``(A \\times B)^T = B^T \\times A^T``, we can simply swap the roles of ``A`` and ``B`` to use the algorithm on column-order matrices and get the correct result back, in column order. + +See https://eprint.iacr.org/2018/1041.pdf +""" +function row_mat_times_mat(sm1, sm2) + n = length(sm1) + d = sm1.shape[1] + shape = sm1.shape + ctx = sm1.context + + # Flatten the matrices so that rotations act on flat vectors instead of being 2d rotations + sm1 = reshape(sm1, (n)) + sm2 = reshape(sm2, (n)) + + # Step 1-1 + # Compute linear Transformation sigma on A + + + sm1_sigma = PlainArray(zeros(n), ctx) + + for k in -d+1:d-1 + if k >= 0 + u_k_sigma = PlainArray([0 <= i-d*k && i-d*k < d-k ? 1 : 0 for i in 0:n-1], ctx) + else + u_k_sigma = PlainArray([-k <= i-(d+k)*d && i-(d+k)*d < d ? 1 : 0 for i in 0:n-1], ctx) + end + # Note that Rot(ct; i) in the paper is a leftshift, i.e. circshift(ct, -i) + sm1_sigma += circshift(sm1, -(k)) .* u_k_sigma + end + + # Step 1-2 + # Compute linear Transformation tau on B + + + sm2_tau = PlainArray(zeros(n), ctx) + + for k in 0:d-1 + u_dk_tau = PlainArray([(i - k) / d in 0:d-1 ? 1 : 0 for i in 0:n-1], ctx) + sm2_tau += circshift(sm2, -(d*k)) .* u_dk_tau + end + + # Step 2 and 3 + sm3 = sm1_sigma .* sm2_tau + for k in 1:d-1 + v_k = PlainArray([0 <= (i % d) && (i % d) < (d-k) ? 1 : 0 for i in 0:n-1], ctx) + v_k_d = PlainArray([(d-k) <= (i % d) && (i % d) < d ? 1 : 0 for i in 0:n-1], ctx) + sm1_k = circshift(sm1_sigma, -(k)) .* v_k + circshift(sm1_sigma, -(k-d)) .* v_k_d + sm2_k = circshift(sm2_tau, -(d*k)) + sm3 += sm1_k .* sm2_k + end + return reshape(sm3, shape) +end diff --git a/src/openfhe.jl b/src/openfhe.jl index b555b3d..65f01be 100644 --- a/src/openfhe.jl +++ b/src/openfhe.jl @@ -257,6 +257,69 @@ function PlainArray(data::Vector{Float64}, context::SecureContext{<:OpenFHEBacke plain_array end +""" + resize(a::PlainVector{<:OpenFHEBackend}, n::Integer) + +Return a `PlainVector` containing `n` elements. +If `n` is smaller than the current length, the first `n` elements are retained. +If `n` is larger, the new elements are not guaranteed to be initialized. +When `n` exceeds the current capacity, new plaintexts are allocated. +When `n` frees a complete batch, excess plaintexts are dropped. + +See also: [`PlainVector`](@ref), [`capacity`](@ref) +""" +function resize(a::PlainVector{<:OpenFHEBackend}, n::Integer) + cc = get_crypto_context(a.context) + batch_size = OpenFHE.GetBatchSize(OpenFHE.GetEncodingParams(cc)) + n_plaintexts_needed = ceil(Int, n / batch_size) + n_plaintexts_current = length(a.data) + + if n_plaintexts_needed < n_plaintexts_current + new_data = a.data[1:n_plaintexts_needed] + elseif n_plaintexts_needed > n_plaintexts_current + new_data = copy(a.data) + for _ in (n_plaintexts_current+1):(n_plaintexts_needed - 1) + push!(new_data, OpenFHE.MakeCKKSPackedPlaintext(cc, zeros(batch_size))) + end + push!(new_data, OpenFHE.MakeCKKSPackedPlaintext(cc, zeros(batch_size))) + else + return PlainArray(a.data, (n,), a.capacity, a.context) + end + + PlainArray(new_data, (n,), n_plaintexts_needed * batch_size, a.context) +end + +""" + resize(a::SecureVector{<:OpenFHEBackend}, n::Integer) + +Return a `SecureVector` containing `n` elements. +If `n` is smaller than the current length, the first `n` elements are retained. +If `n` is larger, the new elements are not guaranteed to be initialized. +When `n` exceeds the current capacity, new ciphertexts are allocated via `Clone`. +When `n` frees a complete batch, excess ciphertexts are dropped. + +See also: [`SecureVector`](@ref), [`capacity`](@ref) +""" +function resize(a::SecureVector{<:OpenFHEBackend}, n::Integer) + cc = get_crypto_context(a.context) + batch_size = OpenFHE.GetBatchSize(OpenFHE.GetEncodingParams(cc)) + n_ciphertexts_needed = ceil(Int, n / batch_size) + n_ciphertexts_current = length(a.data) + + if n_ciphertexts_needed < n_ciphertexts_current + new_data = a.data[1:n_ciphertexts_needed] + elseif n_ciphertexts_needed > n_ciphertexts_current + new_data = copy(a.data) + for _ in (n_ciphertexts_current+1):n_ciphertexts_needed + push!(new_data, OpenFHE.Clone(a.data[1])) + end + else + return SecureArray(a.data, (n,), a.capacity, a.context) + end + + SecureArray(new_data, (n,), n_ciphertexts_needed * batch_size, a.context) +end + function Base.show(io::IO, pa::PlainArray{<:OpenFHEBackend}) print(io, collect(pa)) end @@ -389,6 +452,11 @@ function add(sa::SecureArray{<:OpenFHEBackend}, pa::PlainArray{<:OpenFHEBackend} secure_array end +function add(pa1::PlainArray{<:OpenFHEBackend}, pa2::PlainArray{<:OpenFHEBackend}) + data = vec(collect(pa1)) .+ vec(collect(pa2)) + PlainArray(Vector{Float64}(data), pa1.context, size(pa1)) +end + function add(sa::SecureArray{<:OpenFHEBackend}, scalar::Real) cc = get_crypto_context(sa) ciphertexts = Vector{OpenFHE.Ciphertext}(undef, length(sa.data)) @@ -400,6 +468,11 @@ function add(sa::SecureArray{<:OpenFHEBackend}, scalar::Real) secure_array end +function add(pa1::PlainArray{<:OpenFHEBackend}, scalar::Real) + data = vec(collect(pa1)) .+ scalar + PlainArray(Vector{Float64}(data), pa1.context, size(pa1)) +end + function subtract(sa1::SecureArray{<:OpenFHEBackend}, sa2::SecureArray{<:OpenFHEBackend}) cc = get_crypto_context(sa1) ciphertexts = Vector{OpenFHE.Ciphertext}(undef, length(sa1.data)) @@ -433,6 +506,11 @@ function subtract(pa::PlainArray{<:OpenFHEBackend}, sa::SecureArray{<:OpenFHEBac secure_array end +function subtract(pa1::PlainArray{<:OpenFHEBackend}, pa2::PlainArray{<:OpenFHEBackend}) + data = vec(collect(pa1)) .- vec(collect(pa2)) + PlainArray(Vector{Float64}(data), pa1.context, size(pa1)) +end + function subtract(sa::SecureArray{<:OpenFHEBackend}, scalar::Real) cc = get_crypto_context(sa) ciphertexts = Vector{OpenFHE.Ciphertext}(undef, length(sa.data)) @@ -455,6 +533,16 @@ function subtract(scalar::Real, sa::SecureArray{<:OpenFHEBackend}) secure_array end +function subtract(pa::PlainArray{<:OpenFHEBackend}, scalar::Real) + data = vec(collect(pa)) .- scalar + PlainArray(Vector{Float64}(data), pa.context, size(pa)) +end + +function subtract(scalar::Real, pa::PlainArray{<:OpenFHEBackend}) + data = scalar .- vec(collect(pa)) + PlainArray(Vector{Float64}(data), pa.context, size(pa)) +end + function negate(sa::SecureArray{<:OpenFHEBackend}) cc = get_crypto_context(sa) ciphertexts = Vector{OpenFHE.Ciphertext}(undef, length(sa.data)) @@ -466,6 +554,11 @@ function negate(sa::SecureArray{<:OpenFHEBackend}) secure_array end +function negate(pa::PlainArray{<:OpenFHEBackend}) + data = -vec(collect(pa)) + PlainArray(Vector{Float64}(data), pa.context, size(pa)) +end + function multiply(sa1::SecureArray{<:OpenFHEBackend}, sa2::SecureArray{<:OpenFHEBackend}) cc = get_crypto_context(sa1) ciphertexts = Vector{OpenFHE.Ciphertext}(undef, length(sa1.data)) @@ -488,6 +581,11 @@ function multiply(sa::SecureArray{<:OpenFHEBackend}, pa::PlainArray{<:OpenFHEBac secure_array end +function multiply(pa1::PlainArray{<:OpenFHEBackend}, pa2::PlainArray{<:OpenFHEBackend}) + data = vec(collect(pa1)) .* vec(collect(pa2)) + PlainArray(Vector{Float64}(data), pa1.context, size(pa1)) +end + function multiply(sa::SecureArray{<:OpenFHEBackend}, scalar::Real) cc = get_crypto_context(sa) ciphertexts = Vector{OpenFHE.Ciphertext}(undef, length(sa.data)) @@ -685,10 +783,10 @@ function rotate(sa::SecureArray{<:OpenFHEBackend, N}, shift) where N # operate with N-dimensional array in form of 1D sv = SecureArray(sa.data, (length(sa),), capacity(sa), sa.context) # apply main shift - sv_new = circshift(sv * main_mask, main_1d_shift) + sv_new = circshift(sv .* main_mask, main_1d_shift) # correct positions of elements in each dimension combination for i in eachindex(masks) - sv_new += circshift(sv * masks[i], masked_1d_shift[i]) + sv_new += circshift(sv .* masks[i], masked_1d_shift[i]) end SecureArray(sv_new.data, size(sa), capacity(sa), sa.context) diff --git a/src/types.jl b/src/types.jl index 5991ff8..d2e2624 100644 --- a/src/types.jl +++ b/src/types.jl @@ -183,6 +183,40 @@ See also: [`length`](@ref), [`SecureArray`](@ref), [`PlainArray`](@ref) """ capacity(a::Union{PlainArray, SecureArray}) = a.capacity +""" + reshape(a::SecureArray, shape) + +Return a `SecureArray` with the same data but a different shape. +The new shape must have the same total number of elements as the original. + +See also: [`SecureArray`](@ref) +""" +function Base.reshape(a::SecureArray, shape::NTuple{M, Int}) where M + if prod(shape) != length(a) + throw(DimensionMismatch("new shape $(shape) is incompatible with array of length $(length(a))")) + end + SecureArray(a.data, shape, a.capacity, a.context) +end + +Base.reshape(a::SecureArray, dims::Int...) = reshape(a, dims) + +""" + reshape(a::PlainArray, shape) + +Return a `PlainArray` with the same data but a different shape. +The new shape must have the same total number of elements as the original. + +See also: [`PlainArray`](@ref) +""" +function Base.reshape(a::PlainArray, shape::NTuple{M, Int}) where M + if prod(shape) != length(a) + throw(DimensionMismatch("new shape $(shape) is incompatible with array of length $(length(a))")) + end + PlainArray(a.data, shape, a.capacity, a.context) +end + +Base.reshape(a::PlainArray, dims::Int...) = reshape(a, dims) + # Get wrapper name of a potentially parametric type # Copied from: https://github.com/ClapeyronThermo/Clapeyron.jl/blob/f40c282e2236ff68d91f37c39b5c1e4230ae9ef0/src/utils/core_utils.jl#L17 # Original source: https://github.com/JuliaArrays/ArrayInterface.jl/blob/40d9a87be07ba323cca00f9e59e5285c13f7ee72/src/ArrayInterface.jl#L20 diff --git a/src/unencrypted.jl b/src/unencrypted.jl index d5f9b70..d7b80ff 100644 --- a/src/unencrypted.jl +++ b/src/unencrypted.jl @@ -97,6 +97,38 @@ function PlainArray(data::Array{<:Real}, context::SecureContext{<:Unencrypted}) PlainArray(data, size(data), length(data), context) end +""" + resize(a::PlainVector{<:Unencrypted}, n::Integer) + +Return a `PlainVector` containing `n` elements. +If `n` is smaller than the current length, the first `n` elements are retained. +If `n` is larger, the new elements are not guaranteed to be initialized. + +See also: [`PlainVector`](@ref), [`capacity`](@ref) +""" +function resize(a::PlainVector{<:Unencrypted}, n::Integer) + data = similar(a.data, n) + copy_len = min(n, length(a)) + data[1:copy_len] = a.data[1:copy_len] + PlainArray(data, (n,), n, a.context) +end + +""" + resize(a::SecureVector{<:Unencrypted}, n::Integer) + +Return a `SecureVector` containing `n` elements. +If `n` is smaller than the current length, the first `n` elements are retained. +If `n` is larger, the new elements are not guaranteed to be initialized. + +See also: [`SecureVector`](@ref), [`capacity`](@ref) +""" +function resize(a::SecureVector{<:Unencrypted}, n::Integer) + data = similar(a.data, n) + copy_len = min(n, length(a)) + data[1:copy_len] = a.data[1:copy_len] + SecureArray(data, (n,), n, a.context) +end + function Base.show(io::IO, pa::PlainArray{<:Unencrypted}) print(io, pa.data) end @@ -182,6 +214,15 @@ function add(sa::SecureArray{<:Unencrypted}, scalar::Real) SecureArray(sa.data .+ scalar, size(sa), capacity(sa), sa.context) end +function add(pa1::PlainArray{<:Unencrypted}, pa2::PlainArray{<:Unencrypted}) + PlainArray(pa1.data .+ pa2.data, size(pa1), capacity(pa1), pa1.context) +end + +function add(pa::PlainArray{<:Unencrypted}, scalar::Real) + PlainArray(pa.data .+ scalar, size(pa), capacity(pa), pa.context) +end + + function subtract(sa1::SecureArray{<:Unencrypted}, sa2::SecureArray{<:Unencrypted}) SecureArray(sa1.data .- sa2.data, size(sa1), capacity(sa1), sa1.context) end @@ -194,6 +235,10 @@ function subtract(pa::PlainArray{<:Unencrypted}, sa::SecureArray{<:Unencrypted}) SecureArray(pa.data .- sa.data, size(sa), capacity(sa), sa.context) end +function subtract(pa1::PlainArray{<:Unencrypted}, pa2::PlainArray{<:Unencrypted}) + PlainArray(pa1.data .- pa2.data, size(pa1), capacity(pa1), pa1.context) +end + function subtract(sa::SecureArray{<:Unencrypted}, scalar::Real) SecureArray(sa.data .- scalar, size(sa), capacity(sa), sa.context) end @@ -202,10 +247,22 @@ function subtract(scalar::Real, sa::SecureArray{<:Unencrypted}) SecureArray(scalar .- sa.data, size(sa), capacity(sa), sa.context) end +function subtract(pa::PlainArray{<:Unencrypted}, scalar::Real) + PlainArray(pa.data .- scalar, size(pa), capacity(pa), pa.context) +end + +function subtract(scalar::Real, pa::PlainArray{<:Unencrypted}) + PlainArray(scalar .- pa.data, size(pa), capacity(pa), pa.context) +end + function negate(sa::SecureArray{<:Unencrypted}) SecureArray(-sa.data, size(sa), capacity(sa), sa.context) end +function negate(pa::PlainArray{<:Unencrypted}) + PlainArray(-pa.data, size(pa), capacity(pa), pa.context) +end + function multiply(sa1::SecureArray{<:Unencrypted}, sa2::SecureArray{<:Unencrypted}) SecureArray(sa1.data .* sa2.data, size(sa1), capacity(sa1), sa1.context) end @@ -214,10 +271,18 @@ function multiply(sa::SecureArray{<:Unencrypted}, pa::PlainArray{<:Unencrypted}) SecureArray(sa.data .* pa.data, size(sa), capacity(sa), sa.context) end +function multiply(pa1::PlainArray{<:Unencrypted}, pa2::PlainArray{<:Unencrypted}) + PlainArray(pa1.data .* pa2.data, size(pa1), capacity(pa1), pa1.context) +end + function multiply(sa::SecureArray{<:Unencrypted}, scalar::Real) SecureArray(sa.data .* scalar, size(sa), capacity(sa), sa.context) end +function multiply(pa::PlainArray{<:Unencrypted}, scalar::Real) + PlainArray(pa.data .* scalar, size(pa), capacity(pa), sa.context) +end + function rotate(sa::SecureArray{<:Unencrypted, N}, shift) where N SecureArray(circshift(sa.data, shift), size(sa), capacity(sa), sa.context) end diff --git a/test/runtests.jl b/test/runtests.jl index 503c34a..0ad0909 100644 --- a/test/runtests.jl +++ b/test/runtests.jl @@ -5,5 +5,6 @@ using Test include("test_serialization.jl") include("test_examples.jl") include("test_benchmarks.jl") + include("test_linear_algebra.jl") end diff --git a/test/test_linear_algebra.jl b/test/test_linear_algebra.jl new file mode 100644 index 0000000..5404e42 --- /dev/null +++ b/test/test_linear_algebra.jl @@ -0,0 +1,77 @@ +module TestLinearAlgebra + +using Test +using SecureArithmetic +using OpenFHE + +function make_context() + level_budget = [4, 4] + levels_after_bootstrap = 10 + depth = levels_after_bootstrap + GetBootstrapDepth(level_budget, UNIFORM_TERNARY) + + parameters = CCParams{CryptoContextCKKSRNS}() + SetSecretKeyDist(parameters, UNIFORM_TERNARY) + SetSecurityLevel(parameters, HEStd_NotSet) + SetRingDim(parameters, 1 << 12) + SetScalingModSize(parameters, 59) + SetScalingTechnique(parameters, FLEXIBLEAUTO) + SetFirstModSize(parameters, 60) + SetMultiplicativeDepth(parameters, depth) + + cc = GenCryptoContext(parameters) + Enable(cc, PKE) + Enable(cc, KEYSWITCH) + Enable(cc, LEVELEDSHE) + Enable(cc, ADVANCEDSHE) + Enable(cc, FHE) + + EvalBootstrapSetup(cc; level_budget) + + context = SecureContext(OpenFHEBackend(cc)) + public_key, private_key = generate_keys(context) + + + return (context, public_key, private_key) +end + +@testset verbose=true showtiming=true "test_linear_algebra.jl" begin + +@testset verbose=true showtiming=true "square_col_mat_mat_mul" begin + (context, public_key, private_key) = make_context() + + m1 = collect([0.25 0.5 0.75; + 1.0 2.0 3.0; + 4.0 5.0 6.0]) + + m2 = collect([6.0 5.0 4.0; + 3.0 2.0 1.0; + 0.75 0.5 0.25]) + + pm1 = PlainMatrix(m1, context) + pm2 = PlainMatrix(m2, context) + + + sm1 = encrypt(pm1, public_key) + sm2 = encrypt(pm2, public_key) + + init_rotation!(context, private_key, (length(m1)), 1:length(m1)...) + init_multiplication!(context, private_key) + + + + expected = [ 3.5625 2.625 1.6875; + 14.25 10.5 6.75; + 43.5 33.0 22.5] + + @test collect(decrypt(sm1 * sm2, private_key)) ≈ expected + @test collect(decrypt(sm1 * pm2, private_key)) ≈ expected + @test collect(decrypt(pm1 * sm2, private_key)) ≈ expected + @test collect(pm1 * pm2) ≈ expected +end + +release_context_memory() +GC.gc() + +end # @testset "test_serialization.jl" + +end # module diff --git a/test/test_serialization.jl b/test/test_serialization.jl index f9b466a..877f035 100644 --- a/test/test_serialization.jl +++ b/test/test_serialization.jl @@ -153,7 +153,7 @@ end sv_restored = deserialize(io) # Multiplication requires eval mult keys — this proves they survived the roundtrip - sv_mult = sv_restored * sv_restored + sv_mult = sv_restored .* sv_restored result = collect(decrypt(sv_mult, sk_restored)) @test result ≈ x1 .^ 2 end diff --git a/test/test_unit.jl b/test/test_unit.jl index 3f45858..8cf2487 100644 --- a/test/test_unit.jl +++ b/test/test_unit.jl @@ -114,31 +114,51 @@ for backend in ((; name = "OpenFHE", BackendT = OpenFHEBackend, context = contex @test sv1 + sv2 isa SecureVector @test sv1 + pv1 isa SecureVector @test pv1 + sv1 isa SecureVector + @test pv1 + pv2 isa PlainVector @test sv1 + 3 isa SecureVector @test 4 + sv1 isa SecureVector + @test pv1 + 3 isa PlainVector + @test 4 + pv1 isa PlainVector + @test sm1 + sm2 isa SecureMatrix @test sm1 + pm1 isa SecureMatrix @test pm1 + sm1 isa SecureMatrix + @test pm1 + pm2 isa PlainMatrix @test sm1 + 3 isa SecureMatrix @test 4 + sm1 isa SecureMatrix + @test pm1 + 3 isa PlainMatrix + @test 4 + pm1 isa PlainMatrix + @test sa1 + sa2 isa SecureArray @test sa1 + pa1 isa SecureArray @test pa1 + sa1 isa SecureArray + @test pa1 + pa2 isa PlainArray @test sa1 + 3 isa SecureArray @test 4 + sa1 isa SecureArray + @test pa1 + 3 isa PlainArray + @test 4 + pa1 isa PlainArray end @testset verbose=true showtiming=true "subtract" begin @test sv1 - sv2 isa SecureVector @test sv1 - pv1 isa SecureVector @test pv1 - sv1 isa SecureVector + @test pv1 - pv2 isa PlainVector @test sv1 - 3 isa SecureVector @test 4 - sv1 isa SecureVector + @test pv1 - 3 isa PlainVector + @test 4 - pv1 isa PlainVector + @test sm1 - sm2 isa SecureMatrix @test sm1 - pm1 isa SecureMatrix @test pm1 - sm1 isa SecureMatrix + @test pm1 - pm2 isa PlainMatrix @test sm1 - 3 isa SecureMatrix @test 4 - sm1 isa SecureMatrix + @test pm1 - 3 isa PlainMatrix + @test 4 - pm1 isa PlainMatrix + + @test sa1 - sa2 isa SecureArray @test sa1 - pa1 isa SecureArray @test pa1 - sa1 isa SecureArray @@ -147,27 +167,38 @@ for backend in ((; name = "OpenFHE", BackendT = OpenFHEBackend, context = contex end @testset verbose=true showtiming=true "multiply" begin - @test sv1 * sv2 isa SecureVector - @test sv1 * pv1 isa SecureVector - @test pv1 * sv1 isa SecureVector + @test sv1 .* sv2 isa SecureVector + @test sv1 .* pv1 isa SecureVector + @test pv1 .* sv1 isa SecureVector + @test pv1 .* pv2 isa PlainVector @test sv1 * 3 isa SecureVector + @test sv1 .* 3 isa SecureVector @test 4 * sv1 isa SecureVector - @test sm1 * sm2 isa SecureMatrix - @test sm1 * pm1 isa SecureMatrix - @test pm1 * sm1 isa SecureMatrix + @test 4 .* sv1 isa SecureVector + @test sm1 .* sm2 isa SecureMatrix + @test sm1 .* pm1 isa SecureMatrix + @test pm1 .* sm1 isa SecureMatrix @test sm1 * 3 isa SecureMatrix + @test sm1 .* 3 isa SecureMatrix @test 4 * sm1 isa SecureMatrix - @test sa1 * sa2 isa SecureArray - @test sa1 * pa1 isa SecureArray - @test pa1 * sa1 isa SecureArray + @test 4 .* sm1 isa SecureMatrix + @test sa1 .* sa2 isa SecureArray + @test sa1 .* pa1 isa SecureArray + @test pa1 .* sa1 isa SecureArray @test sa1 * 3 isa SecureArray + @test sa1 .* 3 isa SecureArray @test 4 * sa1 isa SecureArray + @test 4 .* sa1 isa SecureArray end @testset verbose=true showtiming=true "negate" begin @test -sv2 isa SecureVector @test -sm2 isa SecureMatrix @test -sa2 isa SecureArray + + @test -pv2 isa PlainVector + @test -pm2 isa PlainMatrix + @test -pa2 isa PlainArray end sv_short = encrypt([1.0, 2.0, 3.0], public_key, context) @@ -201,12 +232,35 @@ for backend in ((; name = "OpenFHE", BackendT = OpenFHEBackend, context = contex @test collect(decrypt(circshift(sa3, 8), private_key)) ≈ circshift(a3, 8) @test collect(decrypt(circshift(sa4, (0, 3, 1, -3)), private_key)) ≈ circshift(a4, (0, 3, 1, -3)) @test collect(decrypt(circshift(sa4, [-1, -2, -1, 2]), private_key)) ≈ circshift(a4, [-1, -2, -1, 2]) + + + @test collect(circshift(pv1, (1,))) ≈ + [5.0, 0.25, 0.5, 0.75, 1.0, 2.0, 3.0, 4.0] + @test collect(circshift(pv1, -2)) ≈ + [0.75, 1.0, 2.0, 3.0, 4.0, 5.0, 0.25, 0.5] + @test collect(circshift(pm1, (1, -1))) ≈ circshift(m1, (1, -1)) + @test collect(circshift(pm1, (1, 1))) ≈ circshift(m1, (1, 1)) + @test collect(circshift(pm1, -1)) ≈ circshift(m1, -1) + @test collect(circshift(pm1, [0, 1])) ≈ circshift(m1, [0, 1]) + @test collect(circshift(pm1, (1, 0))) ≈ circshift(m1, (1, 0)) + @test collect(circshift(pm1, 2)) ≈ circshift(m1, 2) + @test collect(circshift(pm1, (0, 0))) ≈ m1 + @test collect(circshift(pa1, 1)) ≈ circshift(a1, 1) + @test collect(circshift(pa1, 10)) ≈ circshift(a1, 10) + @test collect(circshift(pa1, -14)) ≈ circshift(a1, -14) + @test collect(circshift(pa1, 7)) ≈ circshift(a1, 7) + @test collect(circshift(pa1, (3,))) ≈ circshift(a1, (3,)) + @test collect(circshift(pa1, 0)) ≈ circshift(a1, 0) + @test collect(circshift(pa3, 2)) ≈ circshift(a3, 2) + @test collect(circshift(pa3, 8)) ≈ circshift(a3, 8) + @test collect(circshift(pa4, (0, 3, 1, -3))) ≈ circshift(a4, (0, 3, 1, -3)) + @test collect(circshift(pa4, [-1, -2, -1, 2])) ≈ circshift(a4, [-1, -2, -1, 2]) end @testset verbose=true showtiming=true "multithreading" begin @test enable_multithreading() - @test collect(decrypt(sa1 + sa2, private_key)) ≈ a1 .+ a2 - @test collect(decrypt(sa1 * sa2, private_key)) ≈ a1 .* a2 + @test collect(decrypt(sa1 + sa2, private_key)) ≈ a1 + a2 + @test collect(decrypt(sa1 .* sa2, private_key)) ≈ a1 .* a2 @test collect(decrypt(circshift(sa1, -14), private_key)) ≈ circshift(a1, -14) @test !disable_multithreading() end @@ -224,6 +278,55 @@ for backend in ((; name = "OpenFHE", BackendT = OpenFHEBackend, context = contex @test size(sa1, 1) == size(pa1, 1) end + @testset verbose=true showtiming=true "resize" begin + @testset verbose=true showtiming=true "PlainVector" begin + pv_resized = pv1 + for resize_params in ( + (; n = 8, c = 8, what="no op"), + (; n = 9, c = 16, what="increase size, changing capacity"), + (; n = 16, c = 16, what="increase size without changing capacity"), + (; n = 9, c = 16, what="reduce size wihtout changing capacity"), + (; n = 7, c = 8, what="reduce size, changing capacity"), + ) + (; n, c, what) = resize_params + @testset verbose=true showtiming=true "$what" begin + pv_resized = resize(pv_resized, n) + @test pv_resized isa PlainVector + @test length(pv_resized) == n + @test collect(pv_resized)[1:min(8, n)] == x1[1:min(8, n)] + @test capacity(pv_resized) == (BackendT == Unencrypted ? n : c) + pv_resized + 1 + # Modifying the resized vector must not modify the original vector + @test collect(pv1) == x1 + end + end + end + + + @testset verbose=true showtiming=true "SecureVector" begin + sv_resized = sv1 + for resize_params in ( + (; n = 8, c = 8, what="no op"), + (; n = 9, c = 16, what="increase size, changing capacity"), + (; n = 16, c = 16, what="increase size without changing capacity"), + (; n = 9, c = 16, what="reduce size wihtout changing capacity"), + (; n = 7, c = 8, what="reduce size, changing capacity"), + ) + (; n, c, what) = resize_params + @testset verbose=true showtiming=true "$what" begin + sv_resized = resize(sv_resized, n) + @test sv_resized isa SecureVector + @test length(sv_resized) == n + @test collect(decrypt(sv_resized, private_key))[1:min(8, n)] ≈ x1[1:min(8, n)] + @test capacity(sv_resized) == (BackendT == Unencrypted ? n : c) + sv_resized + 1 + # Modifying the resized vector must not modify the original vector + @test collect(decrypt(sv1, private_key)) ≈ x1 + end + end + end + end + @testset verbose=true showtiming=true "capacity" begin @test capacity(pv1) == 8 @test capacity(sv1) == 8