Skip to content
5 changes: 4 additions & 1 deletion .gitignore
Original file line number Diff line number Diff line change
Expand Up @@ -3,4 +3,7 @@ __pycache__
*.pyc
*.egg-info
.vscode
.coverage
.coverage
data.txt
test_input.txt
*.tmp
170 changes: 170 additions & 0 deletions src/ur/file_fountain_decoder.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,170 @@
#
# Copyright © 2020 Foundation Devices, Inc. and Contributors
# Licensed under the "BSD-2-Clause Plus Patent License"
#

import os
import binascii
from .fountain_decoder import FountainDecoder, InvalidChecksum
from .crc32 import crc32


class FileFountainDecoder(FountainDecoder):
"""
Fountain decoder that uses a files for fragments to minimize RAM usage.
"""

# optimizations
MAX_MIXED_XOR_DEGREE = 3
MAX_MIXED_INDEXES = 6
LEN_SIMPLE_AVOID_MIX = 60

def __init__(self, workdir):
super().__init__()
self.mixed_part_indexes = set()
self.workdir = workdir
self._clear_files()

# optimization for MAX_MIXED_INDEXES
def _limit_mixed_set(self, mix_set):
if len(mix_set) > FileFountainDecoder.MAX_MIXED_INDEXES:
best = sorted(mix_set, key=len)[: FileFountainDecoder.MAX_MIXED_INDEXES]
mix_set.clear()
mix_set.update(best)

def _fragment_path(self, index):
return self.workdir + "/" + "%s.tmp" % index

def _encode_indexes(self, indexes):
b = bytes(indexes)
s = binascii.b2a_base64(b)
return s.replace(b"+", b"-").replace(b"/", b"_").replace(b"\n", b"").decode()

def _fragment_path_mixed(self, indexes):
return self._fragment_path(self._encode_indexes(indexes))

def _clear_files(self):
for name in os.listdir(self.workdir):
if name.endswith(".tmp"):
os.remove(self.workdir + "/" + name)

# overrides ------

def _clear_caches(self):
super()._clear_caches()
self.mixed_part_indexes.clear()
self._clear_files()

def join_fragments_from_files(self, out_path):
remaining = self.expected_message_len
checksum = 0

buf = bytearray(self.expected_fragment_len)
mv = memoryview(buf)

with open(out_path, "wb") as out:
for i in range(self.expected_part_count()):
if remaining <= 0:
break

with open(self._fragment_path(i), "rb") as frag:
while remaining > 0:
n = frag.readinto(buf)
if n == 0:
break

# pylint: disable=consider-using-min-builtin
if n > remaining:
n = remaining

chunk = mv[:n]
checksum = crc32(chunk, checksum)
out.write(chunk)
remaining -= n

return checksum

def _store_fragment(self, index, part):
filename = (
self._fragment_path(index)
if isinstance(index, int)
else self._fragment_path_mixed(index)
)
with open(filename, "wb") as f:
f.write(part.data)

def _finalize_message(self):
out_path = self.workdir + "/" + "data.txt"

checksum = self.join_fragments_from_files(out_path)

if checksum == self.expected_checksum:
self.result = out_path # or open file handle if you prefer
else:
self.result = InvalidChecksum()

def _retrieve_part_data(self, p):
if isinstance(p.data, str):
with open(p.data, "rb") as frag:
p.data = frag.read()

def reduce_mixed_by(self, p):
# avoid looking for mixed if already got many simple parts
if len(self.received_part_indexes) > FileFountainDecoder.LEN_SIMPLE_AVOID_MIX:
return

new_mixed_indexes = set()
for indexes in self.mixed_part_indexes:
# avoid XOR with this mixed parts
if len(indexes) > FileFountainDecoder.MAX_MIXED_XOR_DEGREE:
new_mixed_indexes.add(indexes)
continue

value_bytearray = bytearray(self.expected_fragment_len)
with open(self._fragment_path_mixed(indexes), "rb") as mixed_frag:
mixed_frag.readinto(value_bytearray)
reduced = self.reduced_part_by_part(
FountainDecoder.Part(indexes, value_bytearray), p
)
if reduced.is_simple():
self.queued_parts.append(reduced)
else:
self._store_fragment(reduced.indexes, reduced)
new_mixed_indexes.add(reduced.indexes)

self._limit_mixed_set(new_mixed_indexes)
self.mixed_part_indexes.clear()
self.mixed_part_indexes.update(new_mixed_indexes)

def process_mixed_part(self, p):
# Don't process duplicate parts
if p.indexes in self.mixed_part_indexes:
return

# Reduce this part by all the others
reduced = p

for index in self.received_part_indexes:
r = FountainDecoder.Part(frozenset([index]), self._fragment_path(index))
reduced = self.reduced_part_by_part(reduced, r)
if reduced.is_simple():
break

if not reduced.is_simple():
for indexes in self.mixed_part_indexes:
r = FountainDecoder.Part(indexes, self._fragment_path_mixed(indexes))
reduced = self.reduced_part_by_part(reduced, r)
if reduced.is_simple():
break

# If the part is now simple
if reduced.is_simple():
# Add it to the queue
self.queued_parts.append(reduced)
else:
# Reduce all the mixed parts by this one
self.reduce_mixed_by(reduced)
# Record this new mixed part
self._store_fragment(reduced.indexes, reduced)
self.mixed_part_indexes.add(reduced.indexes)
self._limit_mixed_set(self.mixed_part_indexes)
80 changes: 80 additions & 0 deletions src/ur/file_fountain_encoder.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
#
# Copyright © 2020 Foundation Devices, Inc. and Contributors
# Licensed under the "BSD-2-Clause Plus Patent License"
#

import os
from .fountain_encoder import FountainEncoder
from .utils import xor_into
from .constants import MAX_UINT32
from .crc32 import crc32


class FileFountainEncoder(FountainEncoder):
# pylint: disable=super-init-not-called
def __init__(
self,
file_path,
max_fragment_len,
first_seq_num=0,
min_fragment_len=10,
):
self.file_path = file_path

self.message_len = os.stat(file_path)[6]
assert self.message_len <= MAX_UINT32

self.checksum = self._compute_checksum()

self.fragment_len = self.find_nominal_fragment_length(
self.message_len,
min_fragment_len,
max_fragment_len,
)

self.seq_num = first_seq_num

def _read_range(self, offset, length):
with open(self.file_path, "rb") as f:
f.seek(offset)
return f.read(length)

def _compute_checksum(self):
checksum = 0
remaining = self.message_len

buf = bytearray(256)
mv = memoryview(buf)

with open(self.file_path, "rb") as f:
while remaining > 0:
n = f.readinto(buf)
if n == 0:
break

# pylint: disable=consider-using-min-builtin
if n > remaining:
n = remaining

checksum = crc32(mv[:n], checksum)
remaining -= n

return checksum

# overrides ------

# XOR selected fragments
def mix(self, indexes):
result = bytearray(self.fragment_len)
frag_len = self.fragment_len
msg_len = self.message_len

for index in indexes:
start = index * frag_len
if start >= msg_len:
continue

size = min(frag_len, msg_len - start)
xor_into(result, self._read_range(start, size))

return result
17 changes: 17 additions & 0 deletions src/ur/file_ur_decoder.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
#
# Copyright © 2020 Foundation Devices, Inc. and Contributors
# Licensed under the "BSD-2-Clause Plus Patent License"
#

from .ur_decoder import URDecoder
from .file_fountain_decoder import FileFountainDecoder


class FileURDecoder(URDecoder):
"""
UR decoder that uses a FileFountainDecoder to minimize RAM usage.
"""

def __init__(self, workdir):
super().__init__()
self.fountain_decoder = FileFountainDecoder(workdir)
29 changes: 29 additions & 0 deletions src/ur/file_ur_encoder.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,29 @@
#
# Copyright © 2020 Foundation Devices, Inc. and Contributors
# Licensed under the "BSD-2-Clause Plus Patent License"
#

from .ur_encoder import UREncoder
from .file_fountain_encoder import FileFountainEncoder


class FileUREncoder(UREncoder):
"""
UR decoder that uses a FileFountainEncoder to minimize RAM usage.
"""

# pylint: disable=super-init-not-called
def __init__(self, ur, max_fragment_len, first_seq_num=0, min_fragment_len=10):
self.ur = ur
self.fountain_encoder = FileFountainEncoder(
ur.cbor, max_fragment_len, first_seq_num, min_fragment_len
)

def next_part(self):
part = self.fountain_encoder.next_part()
if self.is_single_part():
with open(self.fountain_encoder.file_path, "rb") as f:
self.ur.cbor = f.read()
return UREncoder.encode(self.ur)

return UREncoder.encode_part(self.ur.type, part)
3 changes: 2 additions & 1 deletion src/ur/fountain_decoder.py
Original file line number Diff line number Diff line change
Expand Up @@ -164,7 +164,8 @@ def reduced_part_by_part(self, a, b):
# `a` is not reducable by `b`, so return a
return a

def _store_fragment(self, _index, part):
# pylint: disable=unused-argument
def _store_fragment(self, index, part):
# store whole part
self.simple_parts[part.indexes] = part

Expand Down
Loading