Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
20 commits
Select commit Hold shift + click to select a range
d95fdb8
feat: implemented multiplication of two d x d matrices in encoded in …
Tom-Finke Sep 21, 2026
d0c0e40
fix: typo
Tom-Finke Sep 21, 2026
41c8723
fix: indentation
Tom-Finke Sep 24, 2026
f69e2ce
Merge branch 'main' into linalg_jiang-et-al
Tom-Finke Sep 24, 2026
d68e2ec
refactor: use broadcasted multiply ".*" for element-wise SecureArray …
Tom-Finke Sep 25, 2026
4e44446
Change multiplication to element-wise operation
Tom-Finke Sep 25, 2026
0b33072
fix: julia codeblock instead of doctest
Tom-Finke Sep 25, 2026
bb6f77e
refactor: use broadcasted multiply ".*" for element-wise SecureArray …
Tom-Finke Sep 25, 2026
cfbf52f
feat: extend main operations to include PlainArray
Tom-Finke Sep 25, 2026
f9ce3f9
refactor: matrix multiplication with * operator for SecureMatrix
Tom-Finke Sep 25, 2026
05fd07d
fix[docs]: fix references
Tom-Finke Sep 25, 2026
fc3b28f
fix: broadcasting for scalar multiplication
Tom-Finke Sep 25, 2026
da38bbc
Merge branch 'refactor_broadcast-element-wise-multiply' into linalg_j…
Tom-Finke Sep 29, 2026
3abb913
include lin alg test in runtests
Tom-Finke Oct 2, 2026
43afc71
removed dangling transposition comment
Tom-Finke Oct 2, 2026
0a6f9d9
Get num rows for mat mul from shape instead of square root
Tom-Finke Oct 2, 2026
93b405d
fix: element wise multiply instead of mat mul
Tom-Finke Oct 2, 2026
05e454c
use i instead of l as counter variable
Tom-Finke Oct 2, 2026
9e39b0e
fix: merge error, remove duplicated array style
Tom-Finke Oct 2, 2026
0f516a6
feat: resize PlainVector and SecureVector
Tom-Finke Oct 2, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion examples/simple_array_operations.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
2 changes: 1 addition & 1 deletion examples/simple_matrix_operations.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
2 changes: 1 addition & 1 deletion examples/simple_real_numbers.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
3 changes: 2 additions & 1 deletion src/SecureArithmetic.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -32,5 +32,6 @@ include("auxiliary.jl")
include("openfhe.jl")
include("unencrypted.jl")
include("arithmetic.jl")
include("linear_algebra.jl")

end # module SecureArithmetic
100 changes: 91 additions & 9 deletions src/arithmetic.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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)

Expand All @@ -36,15 +102,31 @@ 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
end

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)
58 changes: 58 additions & 0 deletions src/linear_algebra.jl
Original file line number Diff line number Diff line change
@@ -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
102 changes: 100 additions & 2 deletions src/openfhe.jl
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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))
Expand All @@ -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))
Expand Down Expand Up @@ -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))
Expand All @@ -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))
Expand All @@ -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))
Expand All @@ -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))
Expand Down Expand Up @@ -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)
Expand Down
Loading
Loading