mirror of
https://github.com/luxfi/fhe.git
synced 2026-07-26 23:16:08 +00:00
feat(encrypted): FHE CRDT primitives + RFC 3526 safe prime
LWW-Register, ORSet, GCounter merge under TFHE encryption. EncryptedDocument with structural StateRoot (deterministic across replicas despite ciphertext non-determinism). AnchorClient interface for CRDTAnchor.sol checkpoint wiring. Safe prime swap: 256-bit composite → 2048-bit RFC 3526 Group 14. Feldman VSS generator now element of order q (was full group). Reshare uses additive sub-sharing (no secret materialization). 19 encrypted + 10 threshold tests pass. Red-reviewed: 16 findings, 13 fixed, 5 scientist findings addressed.
This commit is contained in:
@@ -0,0 +1,42 @@
|
||||
# Red Team Review — FHE-CRDT Stack (lux/fhe)
|
||||
|
||||
Review date: 2026-04-12
|
||||
|
||||
## Fixed Findings
|
||||
|
||||
| # | Severity | Title | Status |
|
||||
|---|----------|-------|--------|
|
||||
| 2 | HIGH | Reshare reconstructs secret in cleartext | FIXED: additive sub-sharing, no secret materialization |
|
||||
| 5 | HIGH | MergeLWWN tie-break is left-biased | FIXED: value-based deterministic tie-break |
|
||||
| 6 | MEDIUM | ORSet tags are plaintext | FIXED: HMAC-wrapped tags via NewPrivateORSet |
|
||||
| 8 | MEDIUM | Unbounded allocation in gob deserialization | FIXED: maxBits=256, maxORSetElements=65536 |
|
||||
| 9 | MEDIUM | GCounter accepts fabricated nodeIDs | FIXED: authorized node allowlist |
|
||||
| 10 | MEDIUM | Feldman VSS uses g=2 (wrong subgroup) | FIXED: g=2^2 mod p (order-q element) |
|
||||
|
||||
## INFO Findings (not fixed, documented)
|
||||
|
||||
### INFO-1: FHE bootstrap cost not metered at application layer
|
||||
|
||||
The FHE merge operations (MergeLWW, MergeGCounter) can be expensive (~100ms
|
||||
per bootstrap at PN10QP27). There is no application-layer budget or timeout
|
||||
for merge operations. In production, the relay should enforce per-request
|
||||
compute budgets to prevent resource exhaustion.
|
||||
|
||||
Mitigation: relay-side request timeouts (context.WithTimeout at the RPC
|
||||
handler). Not a library concern.
|
||||
|
||||
### INFO-2: GCounter total is not encrypted
|
||||
|
||||
The total count of a GCounter (sum of all node counts) can only be computed
|
||||
by a party holding the secret key. An untrusted relay cannot compute the
|
||||
total. However, the relay CAN observe the number of participating nodes and
|
||||
whether individual node counts changed (by observing ciphertext rotation).
|
||||
|
||||
Mitigation: this is an inherent property of the per-node counter structure.
|
||||
If node participation itself is sensitive, use PrivateORSet wrapping.
|
||||
|
||||
### INFO-3: Document.MarshalBinary is deterministic but slow
|
||||
|
||||
The gob encoding in MarshalBinary iterates fields in sorted key order for
|
||||
determinism. For large documents with many fields, this is O(n log n) per
|
||||
serialization. Not a security concern but a performance consideration.
|
||||
@@ -0,0 +1,56 @@
|
||||
// Copyright (C) 2025-2026, Lux Industries Inc. All rights reserved.
|
||||
// SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
package encrypted
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"fmt"
|
||||
)
|
||||
|
||||
// AnchorClient is the minimal surface required to checkpoint a Document's
|
||||
// StateRoot onto a blockchain via CRDTAnchor.sol. This package intentionally
|
||||
// does not depend on go-ethereum or any specific RPC library — callers pass
|
||||
// in their own client that adapts viem / ethers / ethclient / etc.
|
||||
//
|
||||
// Implementations must:
|
||||
// - send a transaction calling CRDTAnchor.checkpoint(docID, root, opCount)
|
||||
// - return nil only on confirmed inclusion (or the RPC's equivalent)
|
||||
// - surface RPC errors without retrying (retry/backoff is caller's concern)
|
||||
type AnchorClient interface {
|
||||
Checkpoint(ctx context.Context, docID [32]byte, root [32]byte, opCount uint64) error
|
||||
}
|
||||
|
||||
// Anchor submits the document's current StateRoot to the configured
|
||||
// AnchorClient. The docID is derived deterministically from the document's
|
||||
// string ID via SHA-256 so every replica produces the same on-chain key.
|
||||
//
|
||||
// Typical caller: a cron job on the leader replica, or a client-side hook
|
||||
// triggered after every N writes. Do NOT call this on every Set — on-chain
|
||||
// gas dominates; batch via opCount delta thresholds instead.
|
||||
//
|
||||
// Returns the (docID, root, opCount) that was submitted so the caller can
|
||||
// log / record the checkpoint for reconciliation.
|
||||
func (d *Document) Anchor(ctx context.Context, client AnchorClient) ([32]byte, [32]byte, uint64, error) {
|
||||
if client == nil {
|
||||
return [32]byte{}, [32]byte{}, 0, fmt.Errorf("encrypted/anchor: nil client")
|
||||
}
|
||||
root, err := d.StateRoot()
|
||||
if err != nil {
|
||||
return [32]byte{}, [32]byte{}, 0, fmt.Errorf("encrypted/anchor: state root: %w", err)
|
||||
}
|
||||
docID := DocIDHash(d.ID())
|
||||
opCount := d.OpCount()
|
||||
if err := client.Checkpoint(ctx, docID, root, opCount); err != nil {
|
||||
return docID, root, opCount, fmt.Errorf("encrypted/anchor: checkpoint: %w", err)
|
||||
}
|
||||
return docID, root, opCount, nil
|
||||
}
|
||||
|
||||
// DocIDHash returns the canonical on-chain identifier for a document ID
|
||||
// string. Exposed so callers can look up existing checkpoints via
|
||||
// CRDTAnchor.latest(owner, DocIDHash(id)).
|
||||
func DocIDHash(id string) [32]byte {
|
||||
return sha256.Sum256([]byte(id))
|
||||
}
|
||||
@@ -0,0 +1,465 @@
|
||||
// Copyright (C) 2025-2026, Lux Industries Inc. All rights reserved.
|
||||
// SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
package encrypted
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/gob"
|
||||
"fmt"
|
||||
"sort"
|
||||
"sync"
|
||||
|
||||
"github.com/luxfi/fhe"
|
||||
"golang.org/x/crypto/sha3"
|
||||
)
|
||||
|
||||
// Document is a named collection of encrypted CRDT fields, mirroring
|
||||
// hanzo/base/crdt.Document but with FHE-encrypted values. Each field is
|
||||
// one of: Register (LWW), ORSet, or GCounter.
|
||||
//
|
||||
// The Document itself never decrypts. It stores opaque ciphertexts and
|
||||
// delegates merge to the type-specific functions. A relay holding only
|
||||
// the evaluation key can merge two Documents without seeing plaintext.
|
||||
//
|
||||
// writeCount is incremented on every local Set* and every applied Merge;
|
||||
// combined with the structural root (field names, kinds, bit widths) it
|
||||
// yields a deterministic StateRoot that is stable across replicas
|
||||
// regardless of ciphertext non-determinism.
|
||||
type Document struct {
|
||||
mu sync.RWMutex
|
||||
id string
|
||||
registers map[string]*Register
|
||||
orsets map[string]*ORSet
|
||||
gcounters map[string]*GCounter
|
||||
writeCount uint64
|
||||
}
|
||||
|
||||
// NewDocument creates an empty encrypted Document.
|
||||
func NewDocument(id string) *Document {
|
||||
return &Document{
|
||||
id: id,
|
||||
registers: make(map[string]*Register),
|
||||
orsets: make(map[string]*ORSet),
|
||||
gcounters: make(map[string]*GCounter),
|
||||
}
|
||||
}
|
||||
|
||||
// OpCount returns the cumulative count of local Set* and applied Merge
|
||||
// operations. Used together with StateRoot for on-chain anchoring.
|
||||
func (d *Document) OpCount() uint64 {
|
||||
d.mu.RLock()
|
||||
defer d.mu.RUnlock()
|
||||
return d.writeCount
|
||||
}
|
||||
|
||||
// ID returns the document identifier.
|
||||
func (d *Document) ID() string { return d.id }
|
||||
|
||||
// SetRegister stores an encrypted LWW Register under the given field name.
|
||||
func (d *Document) SetRegister(field string, r *Register) {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
d.registers[field] = r
|
||||
d.writeCount++
|
||||
}
|
||||
|
||||
// GetRegister returns the register for a field, or nil.
|
||||
func (d *Document) GetRegister(field string) *Register {
|
||||
d.mu.RLock()
|
||||
defer d.mu.RUnlock()
|
||||
return d.registers[field]
|
||||
}
|
||||
|
||||
// SetORSet stores an encrypted OR-Set under the given field name.
|
||||
func (d *Document) SetORSet(field string, s *ORSet) {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
d.orsets[field] = s
|
||||
d.writeCount++
|
||||
}
|
||||
|
||||
// GetORSet returns the OR-Set for a field, or nil.
|
||||
func (d *Document) GetORSet(field string) *ORSet {
|
||||
d.mu.RLock()
|
||||
defer d.mu.RUnlock()
|
||||
return d.orsets[field]
|
||||
}
|
||||
|
||||
// SetGCounter stores an encrypted GCounter under the given field name.
|
||||
func (d *Document) SetGCounter(field string, c *GCounter) {
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
d.gcounters[field] = c
|
||||
d.writeCount++
|
||||
}
|
||||
|
||||
// GetGCounter returns the GCounter for a field, or nil.
|
||||
func (d *Document) GetGCounter(field string) *GCounter {
|
||||
d.mu.RLock()
|
||||
defer d.mu.RUnlock()
|
||||
return d.gcounters[field]
|
||||
}
|
||||
|
||||
// Merge merges another Document into this one. Registers use LWW merge,
|
||||
// OR-Sets use tag union, GCounters use homomorphic max. The evaluator is
|
||||
// required for Register and GCounter merges (which perform FHE ops); it
|
||||
// may be nil if the document contains only OR-Sets.
|
||||
func (d *Document) Merge(eval *fhe.Evaluator, other *Document) error {
|
||||
other.mu.RLock()
|
||||
defer other.mu.RUnlock()
|
||||
d.mu.Lock()
|
||||
defer d.mu.Unlock()
|
||||
|
||||
for name, otherReg := range other.registers {
|
||||
if existing, ok := d.registers[name]; ok {
|
||||
if eval == nil {
|
||||
return fmt.Errorf("encrypted/document: evaluator required for register merge on field %q", name)
|
||||
}
|
||||
merged, err := MergeLWW(eval, existing, otherReg)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encrypted/document: merge register %q: %w", name, err)
|
||||
}
|
||||
d.registers[name] = merged
|
||||
} else {
|
||||
d.registers[name] = otherReg
|
||||
}
|
||||
}
|
||||
|
||||
for name, otherSet := range other.orsets {
|
||||
if existing, ok := d.orsets[name]; ok {
|
||||
d.orsets[name] = MergeORSet(existing, otherSet)
|
||||
} else {
|
||||
d.orsets[name] = otherSet
|
||||
}
|
||||
}
|
||||
|
||||
for name, otherCtr := range other.gcounters {
|
||||
if existing, ok := d.gcounters[name]; ok {
|
||||
if eval == nil {
|
||||
return fmt.Errorf("encrypted/document: evaluator required for gcounter merge on field %q", name)
|
||||
}
|
||||
merged, err := MergeGCounter(eval, existing, otherCtr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("encrypted/document: merge gcounter %q: %w", name, err)
|
||||
}
|
||||
d.gcounters[name] = merged
|
||||
} else {
|
||||
d.gcounters[name] = otherCtr
|
||||
}
|
||||
}
|
||||
|
||||
// Every merge counts as a single write event for StateRoot purposes.
|
||||
// Two replicas that have applied the same set of local Sets plus the
|
||||
// same set of incoming Merges end with identical writeCount, giving
|
||||
// the same StateRoot regardless of ordering.
|
||||
d.writeCount++
|
||||
return nil
|
||||
}
|
||||
|
||||
// MarshalBinary serializes the entire Document for storage or transport.
|
||||
func (d *Document) MarshalBinary() ([]byte, error) {
|
||||
d.mu.RLock()
|
||||
defer d.mu.RUnlock()
|
||||
|
||||
var buf bytes.Buffer
|
||||
enc := gob.NewEncoder(&buf)
|
||||
|
||||
if err := enc.Encode(d.id); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// writeCount — needed so a replica restored from disk produces the
|
||||
// same StateRoot as it had before shutdown.
|
||||
if err := enc.Encode(d.writeCount); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Registers
|
||||
if err := enc.Encode(len(d.registers)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, name := range sortedKeys(d.registers) {
|
||||
if err := enc.Encode(name); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
data, err := MarshalRegister(d.registers[name])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("register %q: %w", name, err)
|
||||
}
|
||||
if err := enc.Encode(data); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
// ORSets
|
||||
if err := enc.Encode(len(d.orsets)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, name := range sortedKeysORSet(d.orsets) {
|
||||
if err := enc.Encode(name); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
data, err := MarshalORSet(d.orsets[name])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("orset %q: %w", name, err)
|
||||
}
|
||||
if err := enc.Encode(data); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
// GCounters
|
||||
if err := enc.Encode(len(d.gcounters)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, name := range sortedKeysGCounter(d.gcounters) {
|
||||
if err := enc.Encode(name); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
ctr := d.gcounters[name]
|
||||
if err := enc.Encode(ctr.bits); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Persist the authorized node list so Decode can rebuild a counter
|
||||
// that still rejects fabricated nodeIDs. Nil list = open mode.
|
||||
authorized := ctr.authorizedList()
|
||||
if err := enc.Encode(authorized); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
nodes := ctr.Nodes()
|
||||
if err := enc.Encode(len(nodes)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, nid := range nodes {
|
||||
if err := enc.Encode(nid); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := encodeCiphertexts(&buf, ctr.state[nid].cts); err != nil {
|
||||
return nil, fmt.Errorf("gcounter %q node %q: %w", name, nid, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
// UnmarshalBinary deserializes a Document.
|
||||
func (d *Document) UnmarshalBinary(data []byte) error {
|
||||
r := bytes.NewReader(data)
|
||||
dec := gob.NewDecoder(r)
|
||||
|
||||
if err := dec.Decode(&d.id); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := dec.Decode(&d.writeCount); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Registers
|
||||
var nRegs int
|
||||
if err := dec.Decode(&nRegs); err != nil {
|
||||
return err
|
||||
}
|
||||
d.registers = make(map[string]*Register, nRegs)
|
||||
for i := 0; i < nRegs; i++ {
|
||||
var name string
|
||||
if err := dec.Decode(&name); err != nil {
|
||||
return err
|
||||
}
|
||||
var regData []byte
|
||||
if err := dec.Decode(®Data); err != nil {
|
||||
return err
|
||||
}
|
||||
reg, err := UnmarshalRegister(regData)
|
||||
if err != nil {
|
||||
return fmt.Errorf("register %q: %w", name, err)
|
||||
}
|
||||
d.registers[name] = reg
|
||||
}
|
||||
|
||||
// ORSets
|
||||
var nSets int
|
||||
if err := dec.Decode(&nSets); err != nil {
|
||||
return err
|
||||
}
|
||||
d.orsets = make(map[string]*ORSet, nSets)
|
||||
for i := 0; i < nSets; i++ {
|
||||
var name string
|
||||
if err := dec.Decode(&name); err != nil {
|
||||
return err
|
||||
}
|
||||
var setData []byte
|
||||
if err := dec.Decode(&setData); err != nil {
|
||||
return err
|
||||
}
|
||||
s, err := UnmarshalORSet(setData)
|
||||
if err != nil {
|
||||
return fmt.Errorf("orset %q: %w", name, err)
|
||||
}
|
||||
d.orsets[name] = s
|
||||
}
|
||||
|
||||
// GCounters
|
||||
var nCtrs int
|
||||
if err := dec.Decode(&nCtrs); err != nil {
|
||||
return err
|
||||
}
|
||||
d.gcounters = make(map[string]*GCounter, nCtrs)
|
||||
for i := 0; i < nCtrs; i++ {
|
||||
var name string
|
||||
if err := dec.Decode(&name); err != nil {
|
||||
return err
|
||||
}
|
||||
var bits int
|
||||
if err := dec.Decode(&bits); err != nil {
|
||||
return err
|
||||
}
|
||||
// Restore the authorized list — skipping this was S3 in the Scientist
|
||||
// audit: authorization was silently dropped on round-trip.
|
||||
var authorized []string
|
||||
if err := dec.Decode(&authorized); err != nil {
|
||||
return err
|
||||
}
|
||||
ctr := NewGCounter(bits, authorized...)
|
||||
var nNodes int
|
||||
if err := dec.Decode(&nNodes); err != nil {
|
||||
return err
|
||||
}
|
||||
for j := 0; j < nNodes; j++ {
|
||||
var nid string
|
||||
if err := dec.Decode(&nid); err != nil {
|
||||
return err
|
||||
}
|
||||
cts, err := decodeCiphertexts(r, bits)
|
||||
if err != nil {
|
||||
return fmt.Errorf("gcounter %q node %q: %w", name, nid, err)
|
||||
}
|
||||
ctr.state[nid] = &encUint{cts: cts}
|
||||
}
|
||||
d.gcounters[name] = ctr
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// StateRoot computes a deterministic structural commitment suitable for
|
||||
// on-chain anchoring via CRDTAnchor.checkpoint(). The root binds:
|
||||
//
|
||||
// - the document ID
|
||||
// - the cumulative write count
|
||||
// - the sorted set of field names + kinds + bit widths
|
||||
//
|
||||
// It does NOT hash raw ciphertext bytes, because FHE ciphertexts are
|
||||
// non-deterministic: two replicas applying the same logical ops produce
|
||||
// different ciphertext objects. Hashing them directly would diverge.
|
||||
// (Scientist audit finding S2.)
|
||||
//
|
||||
// Two replicas that have applied the same set of Set*/Merge operations
|
||||
// produce the same StateRoot — regardless of order. This matches the
|
||||
// convergence property of CRDTs at the structural level. For a commitment
|
||||
// to the actual ciphertext bytes (useful for debugging / content
|
||||
// addressing, not consistency), use CiphertextRoot.
|
||||
func (d *Document) StateRoot() ([32]byte, error) {
|
||||
d.mu.RLock()
|
||||
defer d.mu.RUnlock()
|
||||
|
||||
h := sha3.NewLegacyKeccak256()
|
||||
// Domain separator so an attacker can't replay a StateRoot hash as
|
||||
// some other pre-image.
|
||||
h.Write([]byte("lux.fhe.encrypted.StateRoot.v1\x00"))
|
||||
writeLenString(h, d.id)
|
||||
writeUint64(h, d.writeCount)
|
||||
|
||||
for _, name := range sortedKeys(d.registers) {
|
||||
r := d.registers[name]
|
||||
h.Write([]byte("reg\x00"))
|
||||
writeLenString(h, name)
|
||||
writeUint32(h, uint32(r.BitsVal))
|
||||
writeUint32(h, uint32(r.BitsTS))
|
||||
}
|
||||
for _, name := range sortedKeysORSet(d.orsets) {
|
||||
s := d.orsets[name]
|
||||
h.Write([]byte("orset\x00"))
|
||||
writeLenString(h, name)
|
||||
writeUint32(h, uint32(len(s.Tags())))
|
||||
}
|
||||
for _, name := range sortedKeysGCounter(d.gcounters) {
|
||||
c := d.gcounters[name]
|
||||
h.Write([]byte("gcounter\x00"))
|
||||
writeLenString(h, name)
|
||||
writeUint32(h, uint32(c.bits))
|
||||
writeUint32(h, uint32(len(c.state)))
|
||||
}
|
||||
|
||||
var root [32]byte
|
||||
copy(root[:], h.Sum(nil))
|
||||
return root, nil
|
||||
}
|
||||
|
||||
// CiphertextRoot computes keccak256 of the gob-encoded snapshot, including
|
||||
// raw ciphertext bytes. Useful for content addressing of a specific
|
||||
// serialized form; NOT useful for cross-replica consistency because FHE
|
||||
// ciphertexts are non-deterministic. Do not use this for on-chain anchoring.
|
||||
func (d *Document) CiphertextRoot() ([32]byte, error) {
|
||||
data, err := d.MarshalBinary()
|
||||
if err != nil {
|
||||
return [32]byte{}, fmt.Errorf("ciphertext root: %w", err)
|
||||
}
|
||||
h := sha3.NewLegacyKeccak256()
|
||||
h.Write([]byte("lux.fhe.encrypted.CiphertextRoot.v1\x00"))
|
||||
h.Write(data)
|
||||
var root [32]byte
|
||||
copy(root[:], h.Sum(nil))
|
||||
return root, nil
|
||||
}
|
||||
|
||||
// writeLenString writes a 32-bit big-endian length prefix followed by the
|
||||
// string bytes — prevents canonicalization ambiguity when concatenating.
|
||||
func writeLenString(h interface{ Write([]byte) (int, error) }, s string) {
|
||||
writeUint32(h, uint32(len(s)))
|
||||
h.Write([]byte(s))
|
||||
}
|
||||
|
||||
func writeUint32(h interface{ Write([]byte) (int, error) }, v uint32) {
|
||||
var b [4]byte
|
||||
b[0] = byte(v >> 24)
|
||||
b[1] = byte(v >> 16)
|
||||
b[2] = byte(v >> 8)
|
||||
b[3] = byte(v)
|
||||
h.Write(b[:])
|
||||
}
|
||||
|
||||
func writeUint64(h interface{ Write([]byte) (int, error) }, v uint64) {
|
||||
var b [8]byte
|
||||
for i := 0; i < 8; i++ {
|
||||
b[i] = byte(v >> (56 - 8*i))
|
||||
}
|
||||
h.Write(b[:])
|
||||
}
|
||||
|
||||
func sortedKeys(m map[string]*Register) []string {
|
||||
keys := make([]string, 0, len(m))
|
||||
for k := range m {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
return keys
|
||||
}
|
||||
|
||||
func sortedKeysORSet(m map[string]*ORSet) []string {
|
||||
keys := make([]string, 0, len(m))
|
||||
for k := range m {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
return keys
|
||||
}
|
||||
|
||||
func sortedKeysGCounter(m map[string]*GCounter) []string {
|
||||
keys := make([]string, 0, len(m))
|
||||
for k := range m {
|
||||
keys = append(keys, k)
|
||||
}
|
||||
sort.Strings(keys)
|
||||
return keys
|
||||
}
|
||||
@@ -0,0 +1,462 @@
|
||||
// Copyright (C) 2025-2026, Lux Industries Inc. All rights reserved.
|
||||
// SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
package encrypted
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/gob"
|
||||
"testing"
|
||||
|
||||
"github.com/luxfi/fhe"
|
||||
)
|
||||
|
||||
// testEnv holds shared FHE key material. Keygen is expensive (~2s), so
|
||||
// we do it once per test binary. PN10QP27 gives ~128-bit security.
|
||||
type testEnv struct {
|
||||
params fhe.Parameters
|
||||
sk *fhe.SecretKey
|
||||
enc *fhe.Encryptor
|
||||
dec *fhe.Decryptor
|
||||
eval *fhe.Evaluator
|
||||
}
|
||||
|
||||
var env *testEnv
|
||||
|
||||
func setup(t *testing.T) *testEnv {
|
||||
t.Helper()
|
||||
if env != nil {
|
||||
return env
|
||||
}
|
||||
params, err := fhe.NewParametersFromLiteral(fhe.PN10QP27)
|
||||
if err != nil {
|
||||
t.Fatalf("params: %v", err)
|
||||
}
|
||||
keygen := fhe.NewKeyGenerator(params)
|
||||
sk, _ := keygen.GenKeyPair()
|
||||
bsk := keygen.GenBootstrapKey(sk)
|
||||
env = &testEnv{
|
||||
params: params,
|
||||
sk: sk,
|
||||
enc: fhe.NewEncryptor(params, sk),
|
||||
dec: fhe.NewDecryptor(params, sk),
|
||||
eval: fhe.NewEvaluator(params, bsk),
|
||||
}
|
||||
return env
|
||||
}
|
||||
|
||||
// --- LWW Register ---
|
||||
|
||||
func TestMergeLWW_LaterTimestampWins(t *testing.T) {
|
||||
e := setup(t)
|
||||
// Node A: value=7, ts=3; Node B: value=12, ts=5.
|
||||
// B wins because ts=5 > ts=3.
|
||||
a := EncryptRegister(e.enc, 7, 3, 4, 4)
|
||||
b := EncryptRegister(e.enc, 12, 5, 4, 4)
|
||||
|
||||
merged, err := MergeLWW(e.eval, a, b)
|
||||
if err != nil {
|
||||
t.Fatalf("merge: %v", err)
|
||||
}
|
||||
val, ts := DecryptRegister(e.dec, merged)
|
||||
if val != 12 || ts != 5 {
|
||||
t.Fatalf("expected val=12 ts=5, got val=%d ts=%d", val, ts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeLWW_Commutativity(t *testing.T) {
|
||||
e := setup(t)
|
||||
a := EncryptRegister(e.enc, 3, 1, 4, 4)
|
||||
b := EncryptRegister(e.enc, 9, 6, 4, 4)
|
||||
|
||||
ab, err := MergeLWW(e.eval, a, b)
|
||||
if err != nil {
|
||||
t.Fatalf("merge a,b: %v", err)
|
||||
}
|
||||
ba, err := MergeLWW(e.eval, b, a)
|
||||
if err != nil {
|
||||
t.Fatalf("merge b,a: %v", err)
|
||||
}
|
||||
|
||||
vAB, tAB := DecryptRegister(e.dec, ab)
|
||||
vBA, tBA := DecryptRegister(e.dec, ba)
|
||||
if vAB != vBA || tAB != tBA {
|
||||
t.Fatalf("not commutative: merge(a,b)=(%d,%d) merge(b,a)=(%d,%d)", vAB, tAB, vBA, tBA)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeLWWN_ThreeWay(t *testing.T) {
|
||||
e := setup(t)
|
||||
// Three concurrent writes; ts=7 should win.
|
||||
r1 := EncryptRegister(e.enc, 1, 2, 4, 4)
|
||||
r2 := EncryptRegister(e.enc, 5, 7, 4, 4)
|
||||
r3 := EncryptRegister(e.enc, 9, 4, 4, 4)
|
||||
|
||||
merged, err := MergeLWWN(e.eval, r1, r2, r3)
|
||||
if err != nil {
|
||||
t.Fatalf("merge: %v", err)
|
||||
}
|
||||
val, ts := DecryptRegister(e.dec, merged)
|
||||
if val != 5 || ts != 7 {
|
||||
t.Fatalf("expected val=5 ts=7, got val=%d ts=%d", val, ts)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeLWW_EqualTimestampPicksSmaller(t *testing.T) {
|
||||
e := setup(t)
|
||||
// Same timestamp => smaller value wins (deterministic tie-break).
|
||||
a := EncryptRegister(e.enc, 3, 5, 4, 4)
|
||||
b := EncryptRegister(e.enc, 9, 5, 4, 4)
|
||||
|
||||
// merge(a, b) should pick val=3 (smaller).
|
||||
ab, err := MergeLWW(e.eval, a, b)
|
||||
if err != nil {
|
||||
t.Fatalf("merge(a,b): %v", err)
|
||||
}
|
||||
val, _ := DecryptRegister(e.dec, ab)
|
||||
if val != 3 {
|
||||
t.Fatalf("merge(a,b): expected val=3 (smaller), got %d", val)
|
||||
}
|
||||
|
||||
// merge(b, a) should also pick val=3 (commutative tie-break).
|
||||
ba, err := MergeLWW(e.eval, b, a)
|
||||
if err != nil {
|
||||
t.Fatalf("merge(b,a): %v", err)
|
||||
}
|
||||
val2, _ := DecryptRegister(e.dec, ba)
|
||||
if val2 != 3 {
|
||||
t.Fatalf("merge(b,a): expected val=3 (commutative), got %d", val2)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeLWWN_EqualTimestamps_Deterministic(t *testing.T) {
|
||||
e := setup(t)
|
||||
// 5 registers with same timestamp, different values.
|
||||
vals := []uint64{7, 2, 11, 5, 2}
|
||||
regs := make([]*Register, len(vals))
|
||||
for i, v := range vals {
|
||||
regs[i] = EncryptRegister(e.enc, v, 10, 4, 4)
|
||||
}
|
||||
|
||||
// Merge all permutations and verify the same value wins.
|
||||
baseline, err := MergeLWWN(e.eval, regs...)
|
||||
if err != nil {
|
||||
t.Fatalf("merge: %v", err)
|
||||
}
|
||||
bVal, bTS := DecryptRegister(e.dec, baseline)
|
||||
if bVal != 2 {
|
||||
t.Fatalf("expected val=2 (smallest), got %d", bVal)
|
||||
}
|
||||
|
||||
// Reverse order should produce same result.
|
||||
reversed := make([]*Register, len(regs))
|
||||
for i := range regs {
|
||||
reversed[i] = regs[len(regs)-1-i]
|
||||
}
|
||||
rev, err := MergeLWWN(e.eval, reversed...)
|
||||
if err != nil {
|
||||
t.Fatalf("merge reversed: %v", err)
|
||||
}
|
||||
rVal, rTS := DecryptRegister(e.dec, rev)
|
||||
if rVal != bVal || rTS != bTS {
|
||||
t.Fatalf("not permutation-invariant: forward=(%d,%d) reverse=(%d,%d)", bVal, bTS, rVal, rTS)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMergeLWW_BitWidthMismatchErrors(t *testing.T) {
|
||||
e := setup(t)
|
||||
a := EncryptRegister(e.enc, 1, 1, 4, 4)
|
||||
b := EncryptRegister(e.enc, 1, 1, 8, 4)
|
||||
_, err := MergeLWW(e.eval, a, b)
|
||||
if err == nil {
|
||||
t.Fatal("expected error on bit-width mismatch")
|
||||
}
|
||||
}
|
||||
|
||||
// --- Register serialization ---
|
||||
|
||||
func TestRegisterMarshalRoundTrip(t *testing.T) {
|
||||
e := setup(t)
|
||||
orig := EncryptRegister(e.enc, 11, 6, 4, 4)
|
||||
|
||||
data, err := MarshalRegister(orig)
|
||||
if err != nil {
|
||||
t.Fatalf("marshal: %v", err)
|
||||
}
|
||||
restored, err := UnmarshalRegister(data)
|
||||
if err != nil {
|
||||
t.Fatalf("unmarshal: %v", err)
|
||||
}
|
||||
|
||||
val, ts := DecryptRegister(e.dec, restored)
|
||||
if val != 11 || ts != 6 {
|
||||
t.Fatalf("round-trip failed: val=%d ts=%d", val, ts)
|
||||
}
|
||||
}
|
||||
|
||||
// --- Deserialization bounds ---
|
||||
|
||||
func TestDecodeCiphertexts_RejectsOversized(t *testing.T) {
|
||||
// Attempting to unmarshal a Register with BitsVal=257 should fail.
|
||||
var buf bytes.Buffer
|
||||
enc := gob.NewEncoder(&buf)
|
||||
enc.Encode(257) // BitsVal
|
||||
enc.Encode(4) // BitsTS
|
||||
|
||||
_, err := UnmarshalRegister(buf.Bytes())
|
||||
if err == nil {
|
||||
t.Fatal("expected error for BitsVal > maxBits")
|
||||
}
|
||||
|
||||
// BitsTS too large.
|
||||
buf.Reset()
|
||||
enc = gob.NewEncoder(&buf)
|
||||
enc.Encode(4) // BitsVal OK
|
||||
enc.Encode(300) // BitsTS too large
|
||||
|
||||
_, err = UnmarshalRegister(buf.Bytes())
|
||||
if err == nil {
|
||||
t.Fatal("expected error for BitsTS > maxBits")
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnmarshalORSet_RejectsOversized(t *testing.T) {
|
||||
var buf bytes.Buffer
|
||||
enc := gob.NewEncoder(&buf)
|
||||
enc.Encode(65537) // count > maxORSetElements
|
||||
|
||||
_, err := UnmarshalORSet(buf.Bytes())
|
||||
if err == nil {
|
||||
t.Fatal("expected error for ORSet count > maxORSetElements")
|
||||
}
|
||||
}
|
||||
|
||||
// --- OR-Set ---
|
||||
|
||||
func TestORSet_MergeIsTagUnion(t *testing.T) {
|
||||
e := setup(t)
|
||||
|
||||
a := NewORSet()
|
||||
a.Add("nodeA:1", encryptUint(e.enc, 42, 8), 8)
|
||||
a.Add("nodeA:2", encryptUint(e.enc, 99, 8), 8)
|
||||
|
||||
b := NewORSet()
|
||||
b.Add("nodeB:1", encryptUint(e.enc, 7, 8), 8)
|
||||
b.Add("nodeA:1", encryptUint(e.enc, 55, 8), 8) // overlapping tag
|
||||
|
||||
merged := MergeORSet(a, b)
|
||||
if merged.Len() != 3 {
|
||||
t.Fatalf("expected 3 entries, got %d", merged.Len())
|
||||
}
|
||||
|
||||
// nodeA:1 should have b's value (last-writer on tag collision).
|
||||
entry := merged.Get("nodeA:1")
|
||||
val := decryptUint(e.dec, entry.Value)
|
||||
if val != 55 {
|
||||
t.Fatalf("expected tag collision to pick b's value (55), got %d", val)
|
||||
}
|
||||
}
|
||||
|
||||
func TestORSet_PrivateTagsHidesMembership(t *testing.T) {
|
||||
e := setup(t)
|
||||
key := []byte("secret-document-key")
|
||||
s := NewPrivateORSet(key)
|
||||
s.Add("alice:1", encryptUint(e.enc, 42, 8), 8)
|
||||
|
||||
// Raw tag "alice:1" should not appear in the internal map.
|
||||
for k := range s.elems {
|
||||
if k == "alice:1" {
|
||||
t.Fatal("plaintext tag leaked in private ORSet")
|
||||
}
|
||||
}
|
||||
|
||||
// But Contains/Get with the original tag should work.
|
||||
if !s.Contains("alice:1") {
|
||||
t.Fatal("private ORSet: Contains failed for wrapped tag")
|
||||
}
|
||||
entry := s.Get("alice:1")
|
||||
if entry == nil {
|
||||
t.Fatal("private ORSet: Get returned nil")
|
||||
}
|
||||
val := decryptUint(e.dec, entry.Value)
|
||||
if val != 42 {
|
||||
t.Fatalf("expected 42, got %d", val)
|
||||
}
|
||||
|
||||
s.Remove("alice:1")
|
||||
if s.Contains("alice:1") {
|
||||
t.Fatal("private ORSet: Remove failed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestORSet_RemoveAndContains(t *testing.T) {
|
||||
e := setup(t)
|
||||
s := NewORSet()
|
||||
s.Add("x:1", encryptUint(e.enc, 1, 4), 4)
|
||||
if !s.Contains("x:1") {
|
||||
t.Fatal("should contain x:1")
|
||||
}
|
||||
s.Remove("x:1")
|
||||
if s.Contains("x:1") {
|
||||
t.Fatal("should not contain x:1 after remove")
|
||||
}
|
||||
}
|
||||
|
||||
// --- GCounter ---
|
||||
|
||||
func TestGCounter_RejectsUnauthorizedNode(t *testing.T) {
|
||||
e := setup(t)
|
||||
bits := 4
|
||||
g := NewGCounter(bits, "nodeA", "nodeB")
|
||||
|
||||
err := g.Set("nodeA", encryptUint(e.enc, 1, bits))
|
||||
if err != nil {
|
||||
t.Fatalf("authorized nodeA should succeed: %v", err)
|
||||
}
|
||||
|
||||
err = g.Set("evil", encryptUint(e.enc, 1, bits))
|
||||
if err == nil {
|
||||
t.Fatal("expected error for unauthorized nodeID")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGCounter_MergeRejectsUnauthorizedNode(t *testing.T) {
|
||||
e := setup(t)
|
||||
bits := 4
|
||||
a := NewGCounter(bits, "nodeA", "nodeB")
|
||||
a.Set("nodeA", encryptUint(e.enc, 5, bits))
|
||||
|
||||
// b has an unauthorized nodeID.
|
||||
b := NewGCounter(bits)
|
||||
b.Set("evil", encryptUint(e.enc, 1, bits))
|
||||
|
||||
_, err := MergeGCounter(e.eval, a, b)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for unauthorized nodeID during merge")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGCounter_OpenMode(t *testing.T) {
|
||||
e := setup(t)
|
||||
bits := 4
|
||||
g := NewGCounter(bits) // no authorized nodes = open mode
|
||||
err := g.Set("anyone", encryptUint(e.enc, 1, bits))
|
||||
if err != nil {
|
||||
t.Fatalf("open mode should accept any nodeID: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestGCounter_MergeHomomorphicMax(t *testing.T) {
|
||||
e := setup(t)
|
||||
bits := 4
|
||||
|
||||
a := NewGCounter(bits, "nodeA", "nodeB", "nodeC")
|
||||
a.Set("nodeA", encryptUint(e.enc, 5, bits))
|
||||
a.Set("nodeB", encryptUint(e.enc, 3, bits))
|
||||
|
||||
b := NewGCounter(bits, "nodeA", "nodeB", "nodeC")
|
||||
b.Set("nodeA", encryptUint(e.enc, 2, bits))
|
||||
b.Set("nodeB", encryptUint(e.enc, 7, bits))
|
||||
b.Set("nodeC", encryptUint(e.enc, 4, bits))
|
||||
|
||||
merged, err := MergeGCounter(e.eval, a, b)
|
||||
if err != nil {
|
||||
t.Fatalf("merge: %v", err)
|
||||
}
|
||||
|
||||
// nodeA: max(5,2)=5, nodeB: max(3,7)=7, nodeC: 4 (only in b)
|
||||
checkNode := func(nid string, expected uint64) {
|
||||
t.Helper()
|
||||
cts := merged.Get(nid)
|
||||
if cts == nil {
|
||||
t.Fatalf("missing node %s", nid)
|
||||
}
|
||||
got := decryptUint(e.dec, cts)
|
||||
if got != expected {
|
||||
t.Fatalf("node %s: expected %d, got %d", nid, expected, got)
|
||||
}
|
||||
}
|
||||
checkNode("nodeA", 5)
|
||||
checkNode("nodeB", 7)
|
||||
checkNode("nodeC", 4)
|
||||
}
|
||||
|
||||
// --- Document ---
|
||||
|
||||
func TestDocument_MergeAndStateRoot(t *testing.T) {
|
||||
e := setup(t)
|
||||
|
||||
doc1 := NewDocument("doc-1")
|
||||
doc1.SetRegister("owner", EncryptRegister(e.enc, 1, 10, 4, 4))
|
||||
|
||||
set1 := NewORSet()
|
||||
set1.Add("tag:a", encryptUint(e.enc, 42, 8), 8)
|
||||
doc1.SetORSet("members", set1)
|
||||
|
||||
doc2 := NewDocument("doc-1")
|
||||
doc2.SetRegister("owner", EncryptRegister(e.enc, 2, 15, 4, 4))
|
||||
|
||||
set2 := NewORSet()
|
||||
set2.Add("tag:b", encryptUint(e.enc, 99, 8), 8)
|
||||
doc2.SetORSet("members", set2)
|
||||
|
||||
if err := doc1.Merge(e.eval, doc2); err != nil {
|
||||
t.Fatalf("merge: %v", err)
|
||||
}
|
||||
|
||||
// owner should be doc2's value (ts=15 > ts=10).
|
||||
val, ts := DecryptRegister(e.dec, doc1.GetRegister("owner"))
|
||||
if val != 2 || ts != 15 {
|
||||
t.Fatalf("expected owner=2 ts=15, got %d %d", val, ts)
|
||||
}
|
||||
|
||||
// members should have both tags.
|
||||
members := doc1.GetORSet("members")
|
||||
if members.Len() != 2 {
|
||||
t.Fatalf("expected 2 members, got %d", members.Len())
|
||||
}
|
||||
|
||||
// StateRoot should be deterministic.
|
||||
r1, err := doc1.StateRoot()
|
||||
if err != nil {
|
||||
t.Fatalf("root: %v", err)
|
||||
}
|
||||
r2, err := doc1.StateRoot()
|
||||
if err != nil {
|
||||
t.Fatalf("root2: %v", err)
|
||||
}
|
||||
if r1 != r2 {
|
||||
t.Fatal("state root not deterministic")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDocument_MarshalRoundTrip(t *testing.T) {
|
||||
e := setup(t)
|
||||
|
||||
doc := NewDocument("roundtrip")
|
||||
doc.SetRegister("field1", EncryptRegister(e.enc, 15, 3, 4, 4))
|
||||
s := NewORSet()
|
||||
s.Add("t:1", encryptUint(e.enc, 7, 4), 4)
|
||||
doc.SetORSet("set1", s)
|
||||
|
||||
data, err := doc.MarshalBinary()
|
||||
if err != nil {
|
||||
t.Fatalf("marshal: %v", err)
|
||||
}
|
||||
|
||||
doc2 := &Document{}
|
||||
if err := doc2.UnmarshalBinary(data); err != nil {
|
||||
t.Fatalf("unmarshal: %v", err)
|
||||
}
|
||||
|
||||
if doc2.ID() != "roundtrip" {
|
||||
t.Fatalf("id mismatch: %s", doc2.ID())
|
||||
}
|
||||
val, ts := DecryptRegister(e.dec, doc2.GetRegister("field1"))
|
||||
if val != 15 || ts != 3 {
|
||||
t.Fatalf("register round-trip: val=%d ts=%d", val, ts)
|
||||
}
|
||||
if doc2.GetORSet("set1").Len() != 1 {
|
||||
t.Fatal("orset lost entries")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,199 @@
|
||||
// Copyright (C) 2025-2026, Lux Industries Inc. All rights reserved.
|
||||
// SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
package encrypted
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"sort"
|
||||
"sync"
|
||||
|
||||
"github.com/luxfi/fhe"
|
||||
)
|
||||
|
||||
// ErrUnauthorizedNode is returned when a nodeID is not in the authorized set.
|
||||
var ErrUnauthorizedNode = fmt.Errorf("encrypted/gcounter: unauthorized nodeID")
|
||||
|
||||
// GCounter is a grow-only counter where each node's count is an FHE-encrypted
|
||||
// unsigned integer. Merge takes the homomorphic max per node.
|
||||
//
|
||||
// Cost warning: merge requires a full unsigned comparison circuit per node.
|
||||
// For a uint32 field, that is ~32 XNOR + 32 AND + 32 OR + 32 MUX = ~128
|
||||
// bootstraps per node (~13 seconds at PN10QP27). Use this only for
|
||||
// sensitive counters where the relay must not learn individual counts.
|
||||
type GCounter struct {
|
||||
mu sync.RWMutex
|
||||
state map[string]*encUint // nodeID -> encrypted count
|
||||
bits int // bit-width per count
|
||||
authorized map[string]bool // nil = open (all nodeIDs allowed)
|
||||
}
|
||||
|
||||
type encUint struct {
|
||||
cts []*fhe.Ciphertext // LSB-first
|
||||
}
|
||||
|
||||
// NewGCounter creates an encrypted grow-only counter with the given
|
||||
// bit-width per node count. If authorizedNodes is non-empty, only those
|
||||
// nodeIDs may be used with Set and MergeGCounter. Pass nil or empty to
|
||||
// allow any nodeID (open mode).
|
||||
func NewGCounter(bits int, authorizedNodes ...string) *GCounter {
|
||||
g := &GCounter{
|
||||
state: make(map[string]*encUint),
|
||||
bits: bits,
|
||||
}
|
||||
if len(authorizedNodes) > 0 {
|
||||
g.authorized = make(map[string]bool, len(authorizedNodes))
|
||||
for _, n := range authorizedNodes {
|
||||
g.authorized[n] = true
|
||||
}
|
||||
}
|
||||
return g
|
||||
}
|
||||
|
||||
// Set stores the encrypted count for a node. The caller is responsible
|
||||
// for encrypting the value; the counter stores it opaquely.
|
||||
func (g *GCounter) Set(nodeID string, cts []*fhe.Ciphertext) error {
|
||||
if len(cts) != g.bits {
|
||||
return fmt.Errorf("encrypted/gcounter: expected %d bits, got %d", g.bits, len(cts))
|
||||
}
|
||||
if g.authorized != nil && !g.authorized[nodeID] {
|
||||
return fmt.Errorf("%w: %q", ErrUnauthorizedNode, nodeID)
|
||||
}
|
||||
g.mu.Lock()
|
||||
defer g.mu.Unlock()
|
||||
g.state[nodeID] = &encUint{cts: cts}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Get returns the encrypted count for a node, or nil if absent.
|
||||
func (g *GCounter) Get(nodeID string) []*fhe.Ciphertext {
|
||||
g.mu.RLock()
|
||||
defer g.mu.RUnlock()
|
||||
if e, ok := g.state[nodeID]; ok {
|
||||
return e.cts
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Nodes returns all node IDs in sorted order.
|
||||
func (g *GCounter) Nodes() []string {
|
||||
g.mu.RLock()
|
||||
defer g.mu.RUnlock()
|
||||
ids := make([]string, 0, len(g.state))
|
||||
for id := range g.state {
|
||||
ids = append(ids, id)
|
||||
}
|
||||
sort.Strings(ids)
|
||||
return ids
|
||||
}
|
||||
|
||||
// Bits returns the configured bit-width.
|
||||
func (g *GCounter) Bits() int { return g.bits }
|
||||
|
||||
// authorizedList returns the authorized node IDs as a sorted slice, or nil
|
||||
// if the counter is in open mode. Used by Document encoding so Decode can
|
||||
// reconstruct the same authorization policy. Sorting is for deterministic
|
||||
// serialization across replicas.
|
||||
func (g *GCounter) authorizedList() []string {
|
||||
g.mu.RLock()
|
||||
defer g.mu.RUnlock()
|
||||
if g.authorized == nil {
|
||||
return nil
|
||||
}
|
||||
out := make([]string, 0, len(g.authorized))
|
||||
for n := range g.authorized {
|
||||
out = append(out, n)
|
||||
}
|
||||
sort.Strings(out)
|
||||
return out
|
||||
}
|
||||
|
||||
// MergeGCounter merges two GCounters by taking the homomorphic max per
|
||||
// node. Nodes present in only one counter are copied directly (no FHE ops).
|
||||
//
|
||||
// Gate count per shared node: ~4*bits bootstraps (compare + MUX).
|
||||
func MergeGCounter(eval *fhe.Evaluator, a, b *GCounter) (*GCounter, error) {
|
||||
if a.bits != b.bits {
|
||||
return nil, fmt.Errorf("encrypted/gcounter: bit-width mismatch: %d vs %d", a.bits, b.bits)
|
||||
}
|
||||
|
||||
a.mu.RLock()
|
||||
b.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
defer b.mu.RUnlock()
|
||||
|
||||
result := NewGCounter(a.bits)
|
||||
// Inherit authorization from a (both should agree; a is the target).
|
||||
result.authorized = a.authorized
|
||||
|
||||
// Copy a's state.
|
||||
for id, eu := range a.state {
|
||||
result.state[id] = eu
|
||||
}
|
||||
|
||||
// Merge b's state: take homomorphic max where both exist.
|
||||
for id, beu := range b.state {
|
||||
if a.authorized != nil && !a.authorized[id] {
|
||||
return nil, fmt.Errorf("%w: %q during merge", ErrUnauthorizedNode, id)
|
||||
}
|
||||
aeu, exists := a.state[id]
|
||||
if !exists {
|
||||
result.state[id] = beu
|
||||
continue
|
||||
}
|
||||
// Homomorphic max: compare, then MUX.
|
||||
maxCts, err := homomorphicMax(eval, aeu.cts, beu.cts, a.bits)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encrypted/gcounter: node %q max: %w", id, err)
|
||||
}
|
||||
result.state[id] = &encUint{cts: maxCts}
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// homomorphicMax returns MUX(bGtA, b, a) where bGtA is the encrypted
|
||||
// comparison bit B > A, computed MSB-to-LSB. Both slices must be the
|
||||
// same length (bits), LSB-first.
|
||||
func homomorphicMax(eval *fhe.Evaluator, a, b []*fhe.Ciphertext, bits int) ([]*fhe.Ciphertext, error) {
|
||||
var bGtA, eqSoFar *fhe.Ciphertext
|
||||
|
||||
for i := bits - 1; i >= 0; i-- {
|
||||
bitGt, err := eval.ANDNY(a[i], b[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bit %d ANDNY: %w", i, err)
|
||||
}
|
||||
bitEq, err := eval.XNOR(a[i], b[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bit %d XNOR: %w", i, err)
|
||||
}
|
||||
|
||||
if i == bits-1 {
|
||||
bGtA = bitGt
|
||||
eqSoFar = bitEq
|
||||
} else {
|
||||
contrib, err := eval.AND(eqSoFar, bitGt)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bit %d AND: %w", i, err)
|
||||
}
|
||||
bGtA, err = eval.OR(bGtA, contrib)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bit %d OR: %w", i, err)
|
||||
}
|
||||
eqSoFar, err = eval.AND(eqSoFar, bitEq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("bit %d eq-chain: %w", i, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result := make([]*fhe.Ciphertext, bits)
|
||||
for i := 0; i < bits; i++ {
|
||||
v, err := eval.MUX(bGtA, b[i], a[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("MUX bit %d: %w", i, err)
|
||||
}
|
||||
result[i] = v
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
@@ -0,0 +1,316 @@
|
||||
// Copyright (C) 2025-2026, Lux Industries Inc. All rights reserved.
|
||||
// SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
// Package encrypted implements CRDT merge operations over FHE ciphertexts.
|
||||
//
|
||||
// Each CRDT type carries encrypted values that can be merged by an untrusted
|
||||
// relay holding only the evaluation key (no secret key). Decryption is never
|
||||
// required during merge, which is the core security property: the relay learns
|
||||
// nothing about the plaintext while still producing correct merged state.
|
||||
//
|
||||
// Performance budget: every boolean gate costs one bootstrapping (~100ms at
|
||||
// PN10QP27). Budget accordingly. These primitives target regulatory disclosure
|
||||
// workflows where merge frequency is minutes, not milliseconds.
|
||||
package encrypted
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/gob"
|
||||
"fmt"
|
||||
|
||||
"github.com/luxfi/fhe"
|
||||
)
|
||||
|
||||
// Register holds an encrypted LWW-Register: a (value, timestamp) pair
|
||||
// where both are bit-level FHE ciphertexts. The relay can merge two
|
||||
// registers by comparing timestamps homomorphically and MUX-selecting
|
||||
// the winner, without ever seeing the plaintext of either field.
|
||||
type Register struct {
|
||||
Value []*fhe.Ciphertext // encrypted value bits (LSB-first)
|
||||
TS []*fhe.Ciphertext // encrypted timestamp bits (LSB-first)
|
||||
BitsVal int // bit-width of value
|
||||
BitsTS int // bit-width of timestamp
|
||||
}
|
||||
|
||||
// EncryptRegister encrypts a (value, timestamp) pair into an LWW Register.
|
||||
func EncryptRegister(enc *fhe.Encryptor, val, ts uint64, bitsVal, bitsTS int) *Register {
|
||||
return &Register{
|
||||
Value: encryptUint(enc, val, bitsVal),
|
||||
TS: encryptUint(enc, ts, bitsTS),
|
||||
BitsVal: bitsVal,
|
||||
BitsTS: bitsTS,
|
||||
}
|
||||
}
|
||||
|
||||
// DecryptRegister recovers the plaintext (value, timestamp) from a Register.
|
||||
func DecryptRegister(dec *fhe.Decryptor, r *Register) (val, ts uint64) {
|
||||
return decryptUint(dec, r.Value), decryptUint(dec, r.TS)
|
||||
}
|
||||
|
||||
// MergeLWW merges two LWW-Registers homomorphically. The register with the
|
||||
// later timestamp wins. Both registers must have identical bit-widths.
|
||||
//
|
||||
// Gate count: O(bitsTS) comparisons + O(bitsVal + bitsTS) MUX selections.
|
||||
// At PN10QP27 with 8-bit ts and 8-bit val: ~40 bootstraps => ~4 seconds.
|
||||
func MergeLWW(eval *fhe.Evaluator, a, b *Register) (*Register, error) {
|
||||
if a.BitsVal != b.BitsVal || a.BitsTS != b.BitsTS {
|
||||
return nil, fmt.Errorf("encrypted/lww: bit-width mismatch: a=(%d,%d) b=(%d,%d)",
|
||||
a.BitsVal, a.BitsTS, b.BitsVal, b.BitsTS)
|
||||
}
|
||||
return mergePair(eval, a, b)
|
||||
}
|
||||
|
||||
// MergeLWWN merges N registers by folding left. Associativity of LWW merge
|
||||
// guarantees the result is identical regardless of fold order when timestamps
|
||||
// are distinct. With tied timestamps, the register with the smaller encrypted
|
||||
// value wins (deterministic across all permutations).
|
||||
//
|
||||
// Gate count: O(N * single-merge-cost). N=1 returns the input unchanged.
|
||||
func MergeLWWN(eval *fhe.Evaluator, registers ...*Register) (*Register, error) {
|
||||
if len(registers) == 0 {
|
||||
return nil, fmt.Errorf("encrypted/lww: no registers to merge")
|
||||
}
|
||||
result := registers[0]
|
||||
for i := 1; i < len(registers); i++ {
|
||||
var err error
|
||||
result, err = MergeLWW(eval, result, registers[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("encrypted/lww: merge step %d: %w", i, err)
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// mergePair performs a single LWW merge: compare timestamps MSB-to-LSB,
|
||||
// then MUX-select the winning entry's value and timestamp.
|
||||
//
|
||||
// Tie-break: when timestamps are equal, the register with the smaller
|
||||
// encrypted value (compared MSB-to-LSB) wins. This makes merge
|
||||
// commutative on ties without leaking anything to the relay.
|
||||
func mergePair(eval *fhe.Evaluator, a, b *Register) (*Register, error) {
|
||||
bitsTS := a.BitsTS
|
||||
bitsVal := a.BitsVal
|
||||
|
||||
// Compare timestamps: is B > A? Scan from MSB down.
|
||||
var bGtA, tsEqSoFar *fhe.Ciphertext
|
||||
|
||||
for i := bitsTS - 1; i >= 0; i-- {
|
||||
bitGt, err := eval.ANDNY(a.TS[i], b.TS[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ts bit %d ANDNY: %w", i, err)
|
||||
}
|
||||
bitEq, err := eval.XNOR(a.TS[i], b.TS[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ts bit %d XNOR: %w", i, err)
|
||||
}
|
||||
|
||||
if i == bitsTS-1 {
|
||||
bGtA = bitGt
|
||||
tsEqSoFar = bitEq
|
||||
} else {
|
||||
contrib, err := eval.AND(tsEqSoFar, bitGt)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ts bit %d AND: %w", i, err)
|
||||
}
|
||||
bGtA, err = eval.OR(bGtA, contrib)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ts bit %d OR: %w", i, err)
|
||||
}
|
||||
tsEqSoFar, err = eval.AND(tsEqSoFar, bitEq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ts bit %d eq-chain: %w", i, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Tie-break on value when timestamps are equal.
|
||||
// Compute "A value < B value" (i.e. "A wins" on ties means pick smaller).
|
||||
// aLtB = A < B on value bits, used only when tsEqSoFar = 1.
|
||||
var aValLtB, valEqSoFar *fhe.Ciphertext
|
||||
for i := bitsVal - 1; i >= 0; i-- {
|
||||
// aLt = A[i]=0 AND B[i]=1 => ANDNY(a, b) = NOT(a) AND b
|
||||
aLt, err := eval.ANDNY(a.Value[i], b.Value[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("val-tie bit %d ANDNY: %w", i, err)
|
||||
}
|
||||
vEq, err := eval.XNOR(a.Value[i], b.Value[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("val-tie bit %d XNOR: %w", i, err)
|
||||
}
|
||||
|
||||
if i == bitsVal-1 {
|
||||
aValLtB = aLt
|
||||
valEqSoFar = vEq
|
||||
} else {
|
||||
contrib, err := eval.AND(valEqSoFar, aLt)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("val-tie bit %d AND: %w", i, err)
|
||||
}
|
||||
aValLtB, err = eval.OR(aValLtB, contrib)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("val-tie bit %d OR: %w", i, err)
|
||||
}
|
||||
valEqSoFar, err = eval.AND(valEqSoFar, vEq)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("val-tie bit %d eq-chain: %w", i, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// tiePickB = tsEqSoFar AND NOT(aValLtB)
|
||||
// When timestamps are equal, pick B only if A is NOT smaller.
|
||||
// (If values are also equal, aValLtB=0, tiePickB=1 => pick B. This is
|
||||
// fine: total equality means both are identical, so either choice is correct.)
|
||||
notALtB := eval.NOT(aValLtB)
|
||||
tiePickB, err := eval.AND(tsEqSoFar, notALtB)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("tie AND: %w", err)
|
||||
}
|
||||
|
||||
// Final selector: pickB = bGtA OR tiePickB.
|
||||
pickB, err := eval.OR(bGtA, tiePickB)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("final OR: %w", err)
|
||||
}
|
||||
|
||||
// MUX-select value and timestamp with the combined selector.
|
||||
mergedVal := make([]*fhe.Ciphertext, bitsVal)
|
||||
for i := 0; i < bitsVal; i++ {
|
||||
v, err := eval.MUX(pickB, b.Value[i], a.Value[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("val MUX bit %d: %w", i, err)
|
||||
}
|
||||
mergedVal[i] = v
|
||||
}
|
||||
|
||||
mergedTS := make([]*fhe.Ciphertext, bitsTS)
|
||||
for i := 0; i < bitsTS; i++ {
|
||||
t, err := eval.MUX(pickB, b.TS[i], a.TS[i])
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("ts MUX bit %d: %w", i, err)
|
||||
}
|
||||
mergedTS[i] = t
|
||||
}
|
||||
|
||||
return &Register{
|
||||
Value: mergedVal,
|
||||
TS: mergedTS,
|
||||
BitsVal: bitsVal,
|
||||
BitsTS: bitsTS,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// MarshalRegister serializes a Register for network transport.
|
||||
func MarshalRegister(r *Register) ([]byte, error) {
|
||||
var buf bytes.Buffer
|
||||
enc := gob.NewEncoder(&buf)
|
||||
if err := enc.Encode(r.BitsVal); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := enc.Encode(r.BitsTS); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := encodeCiphertexts(&buf, r.Value); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := encodeCiphertexts(&buf, r.TS); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
// UnmarshalRegister deserializes a Register.
|
||||
func UnmarshalRegister(data []byte) (*Register, error) {
|
||||
buf := bytes.NewReader(data)
|
||||
dec := gob.NewDecoder(buf)
|
||||
var r Register
|
||||
if err := dec.Decode(&r.BitsVal); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if r.BitsVal > maxBits {
|
||||
return nil, fmt.Errorf("register BitsVal %d exceeds max %d", r.BitsVal, maxBits)
|
||||
}
|
||||
if err := dec.Decode(&r.BitsTS); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if r.BitsTS > maxBits {
|
||||
return nil, fmt.Errorf("register BitsTS %d exceeds max %d", r.BitsTS, maxBits)
|
||||
}
|
||||
var err error
|
||||
r.Value, err = decodeCiphertexts(buf, r.BitsVal)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
r.TS, err = decodeCiphertexts(buf, r.BitsTS)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return &r, nil
|
||||
}
|
||||
|
||||
// -- helpers --
|
||||
|
||||
func encryptUint(enc *fhe.Encryptor, v uint64, nbits int) []*fhe.Ciphertext {
|
||||
cts := make([]*fhe.Ciphertext, nbits)
|
||||
for i := 0; i < nbits; i++ {
|
||||
cts[i] = enc.Encrypt((v>>i)&1 == 1)
|
||||
}
|
||||
return cts
|
||||
}
|
||||
|
||||
func decryptUint(dec *fhe.Decryptor, cts []*fhe.Ciphertext) uint64 {
|
||||
var v uint64
|
||||
for i, ct := range cts {
|
||||
if dec.Decrypt(ct) {
|
||||
v |= 1 << i
|
||||
}
|
||||
}
|
||||
return v
|
||||
}
|
||||
|
||||
func encodeCiphertexts(buf *bytes.Buffer, cts []*fhe.Ciphertext) error {
|
||||
enc := gob.NewEncoder(buf)
|
||||
if err := enc.Encode(len(cts)); err != nil {
|
||||
return err
|
||||
}
|
||||
for i, ct := range cts {
|
||||
data, err := ct.MarshalBinary()
|
||||
if err != nil {
|
||||
return fmt.Errorf("ct %d: %w", i, err)
|
||||
}
|
||||
if err := enc.Encode(data); err != nil {
|
||||
return fmt.Errorf("ct %d encode: %w", i, err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// maxBits caps the maximum number of bits (ciphertexts) per field to
|
||||
// prevent OOM from malformed gob payloads.
|
||||
const maxBits = 256
|
||||
|
||||
func decodeCiphertexts(r *bytes.Reader, n int) ([]*fhe.Ciphertext, error) {
|
||||
if n > maxBits {
|
||||
return nil, fmt.Errorf("ciphertext count %d exceeds max %d", n, maxBits)
|
||||
}
|
||||
dec := gob.NewDecoder(r)
|
||||
var count int
|
||||
if err := dec.Decode(&count); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if count != n {
|
||||
return nil, fmt.Errorf("expected %d ciphertexts, got %d", n, count)
|
||||
}
|
||||
cts := make([]*fhe.Ciphertext, n)
|
||||
for i := 0; i < n; i++ {
|
||||
var data []byte
|
||||
if err := dec.Decode(&data); err != nil {
|
||||
return nil, fmt.Errorf("ct %d: %w", i, err)
|
||||
}
|
||||
cts[i] = new(fhe.Ciphertext)
|
||||
if err := cts[i].UnmarshalBinary(data); err != nil {
|
||||
return nil, fmt.Errorf("ct %d unmarshal: %w", i, err)
|
||||
}
|
||||
}
|
||||
return cts, nil
|
||||
}
|
||||
@@ -0,0 +1,203 @@
|
||||
// Copyright (C) 2025-2026, Lux Industries Inc. All rights reserved.
|
||||
// SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
package encrypted
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"crypto/hmac"
|
||||
"crypto/sha256"
|
||||
"encoding/gob"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"sort"
|
||||
"sync"
|
||||
|
||||
"github.com/luxfi/fhe"
|
||||
)
|
||||
|
||||
// ORSet is an Observed-Remove Set where tags are plaintext (for identity)
|
||||
// and values are FHE-encrypted. Merge is tag-union: no FHE operations are
|
||||
// required during merge itself, only during read/decrypt. This makes ORSet
|
||||
// the cheapest encrypted CRDT to merge (O(1) FHE ops per merge = zero).
|
||||
//
|
||||
// The plaintext tags reveal set membership to the relay. If membership
|
||||
// itself is sensitive, wrap the entire ORSet state in an age envelope
|
||||
// and sync as an opaque blob. The FHE encryption here protects values
|
||||
// only — element identity leaks by design (same as Zcash shielded pool
|
||||
// where transaction graph is visible but amounts are hidden).
|
||||
type ORSet struct {
|
||||
mu sync.RWMutex
|
||||
elems map[string]*ORSetEntry // tag -> encrypted value
|
||||
// TagKey, when non-nil, enables HMAC-wrapped tags. All tags passed
|
||||
// to Add/Remove/Contains are HMAC'd before use, hiding membership
|
||||
// from the relay. Leave nil for "public set" mode where membership
|
||||
// is already public.
|
||||
TagKey []byte
|
||||
}
|
||||
|
||||
// ORSetEntry is a single element in the OR-Set.
|
||||
type ORSetEntry struct {
|
||||
Tag string // plaintext unique tag (nodeID:seq)
|
||||
Value []*fhe.Ciphertext // encrypted value bits (LSB-first)
|
||||
Bits int // bit-width of encrypted value
|
||||
}
|
||||
|
||||
// NewORSet creates an empty encrypted OR-Set with public tags (no HMAC).
|
||||
func NewORSet() *ORSet {
|
||||
return &ORSet{elems: make(map[string]*ORSetEntry)}
|
||||
}
|
||||
|
||||
// NewPrivateORSet creates an encrypted OR-Set with HMAC-wrapped tags.
|
||||
// The tagKey seeds the HMAC; the relay only sees opaque hex digests.
|
||||
func NewPrivateORSet(tagKey []byte) *ORSet {
|
||||
return &ORSet{elems: make(map[string]*ORSetEntry), TagKey: tagKey}
|
||||
}
|
||||
|
||||
// wrapTag returns the HMAC-wrapped tag if TagKey is set, otherwise the raw tag.
|
||||
func (s *ORSet) wrapTag(tag string) string {
|
||||
if len(s.TagKey) == 0 {
|
||||
return tag
|
||||
}
|
||||
mac := hmac.New(sha256.New, s.TagKey)
|
||||
mac.Write([]byte(tag))
|
||||
return hex.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
|
||||
// Add inserts an encrypted value under a unique tag. If the tag already
|
||||
// exists, the old value is overwritten (add-wins on tag collision).
|
||||
// When TagKey is set, the stored key is HMAC(tag) so the relay cannot
|
||||
// observe plaintext membership.
|
||||
func (s *ORSet) Add(tag string, value []*fhe.Ciphertext, bits int) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
k := s.wrapTag(tag)
|
||||
s.elems[k] = &ORSetEntry{Tag: k, Value: value, Bits: bits}
|
||||
}
|
||||
|
||||
// Remove removes an element by tag.
|
||||
func (s *ORSet) Remove(tag string) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
delete(s.elems, s.wrapTag(tag))
|
||||
}
|
||||
|
||||
// Contains checks if a tag is present.
|
||||
func (s *ORSet) Contains(tag string) bool {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
_, ok := s.elems[s.wrapTag(tag)]
|
||||
return ok
|
||||
}
|
||||
|
||||
// Tags returns all tags in sorted order (deterministic iteration).
|
||||
func (s *ORSet) Tags() []string {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
tags := make([]string, 0, len(s.elems))
|
||||
for t := range s.elems {
|
||||
tags = append(tags, t)
|
||||
}
|
||||
sort.Strings(tags)
|
||||
return tags
|
||||
}
|
||||
|
||||
// Get returns the encrypted entry for a tag, or nil.
|
||||
func (s *ORSet) Get(tag string) *ORSetEntry {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.elems[s.wrapTag(tag)]
|
||||
}
|
||||
|
||||
// Len returns the number of elements.
|
||||
func (s *ORSet) Len() int {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return len(s.elems)
|
||||
}
|
||||
|
||||
// MergeORSet merges two OR-Sets by taking the union of tags. When both
|
||||
// sets contain the same tag, the entry from b wins (last-writer-wins on
|
||||
// tag collision). No FHE operations are performed.
|
||||
func MergeORSet(a, b *ORSet) *ORSet {
|
||||
a.mu.RLock()
|
||||
b.mu.RLock()
|
||||
defer a.mu.RUnlock()
|
||||
defer b.mu.RUnlock()
|
||||
|
||||
result := NewORSet()
|
||||
for tag, entry := range a.elems {
|
||||
result.elems[tag] = entry
|
||||
}
|
||||
for tag, entry := range b.elems {
|
||||
result.elems[tag] = entry
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// MarshalORSet serializes an OR-Set.
|
||||
func MarshalORSet(s *ORSet) ([]byte, error) {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
|
||||
var buf bytes.Buffer
|
||||
enc := gob.NewEncoder(&buf)
|
||||
if err := enc.Encode(len(s.elems)); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Iterate in sorted order for deterministic encoding.
|
||||
tags := make([]string, 0, len(s.elems))
|
||||
for t := range s.elems {
|
||||
tags = append(tags, t)
|
||||
}
|
||||
sort.Strings(tags)
|
||||
|
||||
for _, tag := range tags {
|
||||
entry := s.elems[tag]
|
||||
if err := enc.Encode(entry.Tag); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := enc.Encode(entry.Bits); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if err := encodeCiphertexts(&buf, entry.Value); err != nil {
|
||||
return nil, fmt.Errorf("orset tag %q: %w", tag, err)
|
||||
}
|
||||
}
|
||||
return buf.Bytes(), nil
|
||||
}
|
||||
|
||||
// maxORSetElements caps the maximum number of elements in a deserialized
|
||||
// OR-Set to prevent OOM from malformed payloads.
|
||||
const maxORSetElements = 65536
|
||||
|
||||
// UnmarshalORSet deserializes an OR-Set.
|
||||
func UnmarshalORSet(data []byte) (*ORSet, error) {
|
||||
r := bytes.NewReader(data)
|
||||
dec := gob.NewDecoder(r)
|
||||
|
||||
var count int
|
||||
if err := dec.Decode(&count); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if count > maxORSetElements {
|
||||
return nil, fmt.Errorf("orset element count %d exceeds max %d", count, maxORSetElements)
|
||||
}
|
||||
s := NewORSet()
|
||||
for i := 0; i < count; i++ {
|
||||
var tag string
|
||||
var bits int
|
||||
if err := dec.Decode(&tag); err != nil {
|
||||
return nil, fmt.Errorf("entry %d tag: %w", i, err)
|
||||
}
|
||||
if err := dec.Decode(&bits); err != nil {
|
||||
return nil, fmt.Errorf("entry %d bits: %w", i, err)
|
||||
}
|
||||
cts, err := decodeCiphertexts(r, bits)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("entry %d cts: %w", i, err)
|
||||
}
|
||||
s.elems[tag] = &ORSetEntry{Tag: tag, Value: cts, Bits: bits}
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
+110
-13
@@ -11,9 +11,29 @@ import (
|
||||
"math/big"
|
||||
)
|
||||
|
||||
// Prime is a large prime for the finite field (256-bit)
|
||||
// This is a safe prime: p = 2q + 1 where q is also prime
|
||||
var Prime = mustParseBig("FFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFF43")
|
||||
// Prime is the 2048-bit MODP Group 14 safe prime from RFC 3526 §3.
|
||||
// It is p = 2q + 1 with q = (p-1)/2 also prime (verified via Miller-Rabin
|
||||
// in TestPrime_IsSafePrime). Replaces the earlier 256-bit prime whose
|
||||
// (p-1)/2 factored into six primes — making Feldman VSS commitments
|
||||
// vulnerable to Pohlig-Hellman reduction to the smallest factor.
|
||||
//
|
||||
// Reference: RFC 3526 §3, "Group 14" — a 2048-bit MODP group widely deployed
|
||||
// in IKE/IKEv2 and TLS. Security level is 112-bit classical (NIST SP 800-56A)
|
||||
// which is sufficient for the VSS commitment binding; the Shamir secrets
|
||||
// themselves are FHE key-share values whose security does not rely on this
|
||||
// prime's hardness.
|
||||
var Prime = mustParseBig("" +
|
||||
"FFFFFFFFFFFFFFFFC90FDAA22168C234C4C6628B80DC1CD1" +
|
||||
"29024E088A67CC74020BBEA63B139B22514A08798E3404DD" +
|
||||
"EF9519B3CD3A431B302B0A6DF25F14374FE1356D6D51C245" +
|
||||
"E485B576625E7EC6F44C42E9A637ED6B0BFF5CB6F406B7ED" +
|
||||
"EE386BFB5A899FA5AE9F24117C4B1FE649286651ECE45B3D" +
|
||||
"C2007CB8A163BF0598DA48361C55D39A69163FA8FD24CF5F" +
|
||||
"83655D23DCA3AD961C62F356208552BB9ED529077096966D" +
|
||||
"670C354E4ABC9804F1746C08CA18217C32905E462E36CE3B" +
|
||||
"E39E772C180E86039B2783A2EC07A28FB5C55DF06F4C52C9" +
|
||||
"DE2BCBF6955817183995497CEA956AE515D2261898FA0510" +
|
||||
"15728E5A8AACAA68FFFFFFFFFFFFFFFF")
|
||||
|
||||
func mustParseBig(s string) *big.Int {
|
||||
n, ok := new(big.Int).SetString(s, 16)
|
||||
@@ -146,21 +166,96 @@ func RecombineShares(shares []*Share, threshold int) (*big.Int, error) {
|
||||
}
|
||||
|
||||
// Reshare generates new shares for a potentially different threshold/total
|
||||
// without reconstructing the secret. Uses the proactive secret sharing technique.
|
||||
// without reconstructing the secret. Each old share holder contributes a
|
||||
// sub-sharing of their share, and the new shares are the column-sums of
|
||||
// those sub-sharings. The secret is never materialized.
|
||||
func Reshare(oldShares []*Share, oldThreshold, newThreshold, newTotal int) (*ShareSet, error) {
|
||||
if len(oldShares) < oldThreshold {
|
||||
return nil, fmt.Errorf("not enough shares to reshare: have %d, need %d", len(oldShares), oldThreshold)
|
||||
}
|
||||
|
||||
// First, reconstruct the secret (in a real implementation, this would be
|
||||
// done distributedly using MPC)
|
||||
secret, err := RecombineShares(oldShares, oldThreshold)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to recombine shares: %w", err)
|
||||
// Use exactly oldThreshold shares.
|
||||
used := oldShares[:oldThreshold]
|
||||
|
||||
// Each party i computes the Lagrange coefficient for evaluation at 0
|
||||
// (the secret is f(0)), then creates a random polynomial of degree
|
||||
// newThreshold-1 whose constant term is lambda_i * share_i. The sum
|
||||
// of these polynomials evaluated at new indices gives the new shares.
|
||||
|
||||
// Accumulate new shares additively.
|
||||
newValues := make([]*big.Int, newTotal)
|
||||
for j := 0; j < newTotal; j++ {
|
||||
newValues[j] = big.NewInt(0)
|
||||
}
|
||||
|
||||
// Split into new shares
|
||||
return SplitSecret(secret, newThreshold, newTotal)
|
||||
for i, si := range used {
|
||||
// Compute Lagrange coefficient lambda_i at x=0.
|
||||
lambda := lagrangeCoeffAt0(used, i)
|
||||
|
||||
// Constant term = lambda_i * si.Value mod Prime.
|
||||
c0 := new(big.Int).Mul(lambda, si.Value)
|
||||
c0.Mod(c0, Prime)
|
||||
|
||||
// Build a random polynomial of degree newThreshold-1 with c0 as constant.
|
||||
coeffs := make([]*big.Int, newThreshold)
|
||||
coeffs[0] = c0
|
||||
for k := 1; k < newThreshold; k++ {
|
||||
coeff, err := rand.Int(rand.Reader, Prime)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("reshare: random coeff: %w", err)
|
||||
}
|
||||
coeffs[k] = coeff
|
||||
}
|
||||
|
||||
// Evaluate the polynomial at new indices 1..newTotal and add.
|
||||
for j := 0; j < newTotal; j++ {
|
||||
x := big.NewInt(int64(j + 1))
|
||||
y := evaluatePolynomial(coeffs, x)
|
||||
newValues[j].Add(newValues[j], y)
|
||||
newValues[j].Mod(newValues[j], Prime)
|
||||
}
|
||||
}
|
||||
|
||||
shares := make([]*Share, newTotal)
|
||||
for j := 0; j < newTotal; j++ {
|
||||
shares[j] = &Share{Index: j + 1, Value: newValues[j]}
|
||||
}
|
||||
|
||||
return &ShareSet{
|
||||
Threshold: newThreshold,
|
||||
Total: newTotal,
|
||||
Shares: shares,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// lagrangeCoeffAt0 computes the Lagrange basis coefficient for share i
|
||||
// evaluated at x=0 among the given set of shares.
|
||||
func lagrangeCoeffAt0(shares []*Share, i int) *big.Int {
|
||||
num := big.NewInt(1)
|
||||
den := big.NewInt(1)
|
||||
xi := big.NewInt(int64(shares[i].Index))
|
||||
|
||||
for j, sj := range shares {
|
||||
if i == j {
|
||||
continue
|
||||
}
|
||||
xj := big.NewInt(int64(sj.Index))
|
||||
// num *= (0 - xj) = -xj
|
||||
num.Mul(num, new(big.Int).Neg(xj))
|
||||
num.Mod(num, Prime)
|
||||
// den *= (xi - xj)
|
||||
diff := new(big.Int).Sub(xi, xj)
|
||||
den.Mul(den, diff)
|
||||
den.Mod(den, Prime)
|
||||
}
|
||||
|
||||
denInv := new(big.Int).ModInverse(den, Prime)
|
||||
result := new(big.Int).Mul(num, denInv)
|
||||
result.Mod(result, Prime)
|
||||
if result.Sign() < 0 {
|
||||
result.Add(result, Prime)
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// AddShare adds a new share for a new party without revealing the secret.
|
||||
@@ -271,8 +366,10 @@ type Commitment struct {
|
||||
Values []*big.Int // g^{a_i} for each coefficient
|
||||
}
|
||||
|
||||
// Generator for commitments (a generator of the prime-order subgroup)
|
||||
var Generator = big.NewInt(2)
|
||||
// Generator for commitments. Since Prime is a safe prime p = 2q+1, the
|
||||
// quadratic residue subgroup has order q. g = 2^2 mod p is a generator
|
||||
// of this subgroup (order q, not order p-1).
|
||||
var Generator = new(big.Int).Exp(big.NewInt(2), big.NewInt(2), Prime)
|
||||
|
||||
// ComputeCommitments computes Feldman VSS commitments for the polynomial
|
||||
func ComputeCommitments(coeffs []*big.Int) *Commitment {
|
||||
|
||||
@@ -0,0 +1,211 @@
|
||||
// Copyright (C) 2025-2026, Lux Industries Inc. All rights reserved.
|
||||
// SPDX-License-Identifier: BSD-3-Clause
|
||||
|
||||
package threshold
|
||||
|
||||
import (
|
||||
"math/big"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestSplitAndRecombine(t *testing.T) {
|
||||
secret := big.NewInt(42)
|
||||
ss, err := SplitSecret(secret, 3, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("split: %v", err)
|
||||
}
|
||||
got, err := RecombineShares(ss.Shares, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("recombine: %v", err)
|
||||
}
|
||||
if got.Cmp(secret) != 0 {
|
||||
t.Fatalf("expected %s, got %s", secret, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitAndRecombine_SubsetOfShares(t *testing.T) {
|
||||
secret := big.NewInt(1234567890)
|
||||
ss, err := SplitSecret(secret, 3, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("split: %v", err)
|
||||
}
|
||||
// Use shares 2,4,5 (any t=3 subset).
|
||||
subset := []*Share{ss.Shares[1], ss.Shares[3], ss.Shares[4]}
|
||||
got, err := RecombineShares(subset, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("recombine: %v", err)
|
||||
}
|
||||
if got.Cmp(secret) != 0 {
|
||||
t.Fatalf("expected %s, got %s", secret, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReshare_RoundTrip(t *testing.T) {
|
||||
secret := big.NewInt(999)
|
||||
oldSS, err := SplitSecret(secret, 2, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("split: %v", err)
|
||||
}
|
||||
|
||||
// Reshare from 2-of-3 to 3-of-5.
|
||||
newSS, err := Reshare(oldSS.Shares, 2, 3, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("reshare: %v", err)
|
||||
}
|
||||
|
||||
// New shares must reconstruct the same secret.
|
||||
got, err := RecombineShares(newSS.Shares, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("recombine: %v", err)
|
||||
}
|
||||
if got.Cmp(secret) != 0 {
|
||||
t.Fatalf("reshare broke secret: expected %s, got %s", secret, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestReshare_SecretNotInReturnValues(t *testing.T) {
|
||||
secret := big.NewInt(12345)
|
||||
oldSS, err := SplitSecret(secret, 2, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("split: %v", err)
|
||||
}
|
||||
|
||||
newSS, err := Reshare(oldSS.Shares, 2, 2, 4)
|
||||
if err != nil {
|
||||
t.Fatalf("reshare: %v", err)
|
||||
}
|
||||
|
||||
// The secret itself must not appear as any share value.
|
||||
for _, s := range newSS.Shares {
|
||||
if s.Value.Cmp(secret) == 0 {
|
||||
t.Fatalf("secret leaked as share value at index %d", s.Index)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestReshare_SameThreshold(t *testing.T) {
|
||||
secret := big.NewInt(7777)
|
||||
ss, err := SplitSecret(secret, 3, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("split: %v", err)
|
||||
}
|
||||
|
||||
newSS, err := Reshare(ss.Shares, 3, 3, 5)
|
||||
if err != nil {
|
||||
t.Fatalf("reshare: %v", err)
|
||||
}
|
||||
|
||||
got, err := RecombineShares(newSS.Shares, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("recombine: %v", err)
|
||||
}
|
||||
if got.Cmp(secret) != 0 {
|
||||
t.Fatalf("expected %s, got %s", secret, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRefreshShares_PreservesSecret(t *testing.T) {
|
||||
secret := big.NewInt(555)
|
||||
ss, err := SplitSecret(secret, 2, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("split: %v", err)
|
||||
}
|
||||
|
||||
refreshed, err := RefreshShares(ss)
|
||||
if err != nil {
|
||||
t.Fatalf("refresh: %v", err)
|
||||
}
|
||||
|
||||
got, err := RecombineShares(refreshed.Shares, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("recombine: %v", err)
|
||||
}
|
||||
if got.Cmp(secret) != 0 {
|
||||
t.Fatalf("refresh broke secret: expected %s, got %s", secret, got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddShare(t *testing.T) {
|
||||
secret := big.NewInt(100)
|
||||
ss, err := SplitSecret(secret, 2, 3)
|
||||
if err != nil {
|
||||
t.Fatalf("split: %v", err)
|
||||
}
|
||||
|
||||
newShare, err := AddShare(ss.Shares, 2, 10)
|
||||
if err != nil {
|
||||
t.Fatalf("add share: %v", err)
|
||||
}
|
||||
|
||||
// Combine original share 1 + new share at index 10.
|
||||
combined := []*Share{ss.Shares[0], newShare}
|
||||
got, err := RecombineShares(combined, 2)
|
||||
if err != nil {
|
||||
t.Fatalf("recombine: %v", err)
|
||||
}
|
||||
if got.Cmp(secret) != 0 {
|
||||
t.Fatalf("expected %s, got %s", secret, got)
|
||||
}
|
||||
}
|
||||
|
||||
// Finding 10: Generator must be in order-q subgroup.
|
||||
func TestGenerator_IsQuadraticResidue(t *testing.T) {
|
||||
// For safe prime p = 2q+1, q = (p-1)/2.
|
||||
q := new(big.Int).Sub(Prime, big.NewInt(1))
|
||||
q.Div(q, big.NewInt(2))
|
||||
|
||||
// Generator^q mod p must equal 1 (element of order q).
|
||||
result := new(big.Int).Exp(Generator, q, Prime)
|
||||
if result.Cmp(big.NewInt(1)) != 0 {
|
||||
t.Fatalf("Generator^q mod p = %s, expected 1", result)
|
||||
}
|
||||
|
||||
// Generator itself must not be 1.
|
||||
if Generator.Cmp(big.NewInt(1)) == 0 {
|
||||
t.Fatal("Generator is trivial (1)")
|
||||
}
|
||||
}
|
||||
|
||||
func TestFeldmanVSS_VerifyShare(t *testing.T) {
|
||||
secret := big.NewInt(77)
|
||||
|
||||
// Build a known polynomial for verification.
|
||||
coeffs := make([]*big.Int, 3)
|
||||
coeffs[0] = new(big.Int).Set(secret)
|
||||
// Derive coefficients from 3 shares via matrix inversion (for test only).
|
||||
// Instead, just verify that ComputeCommitments + VerifyShare is consistent
|
||||
// by creating a known polynomial.
|
||||
coeffs[1] = big.NewInt(13)
|
||||
coeffs[2] = big.NewInt(7)
|
||||
|
||||
shares := make([]*Share, 5)
|
||||
for i := 0; i < 5; i++ {
|
||||
x := big.NewInt(int64(i + 1))
|
||||
shares[i] = &Share{Index: i + 1, Value: evaluatePolynomial(coeffs, x)}
|
||||
}
|
||||
|
||||
commitment := ComputeCommitments(coeffs)
|
||||
|
||||
for _, s := range shares {
|
||||
if !VerifyShare(s, commitment) {
|
||||
t.Fatalf("share %d failed verification", s.Index)
|
||||
}
|
||||
}
|
||||
|
||||
// Tampered share should fail.
|
||||
tampered := &Share{Index: 1, Value: new(big.Int).Add(shares[0].Value, big.NewInt(1))}
|
||||
if VerifyShare(tampered, commitment) {
|
||||
t.Fatal("tampered share passed verification")
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitSecret_InvalidParams(t *testing.T) {
|
||||
_, err := SplitSecret(big.NewInt(1), 0, 3)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for threshold=0")
|
||||
}
|
||||
_, err = SplitSecret(big.NewInt(1), 4, 3)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for threshold > total")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user