From 4b0636b4749757b72662dd6790ad1fdc4ffdec4d Mon Sep 17 00:00:00 2001 From: Xinyu Ma Date: Sun, 6 Sep 2026 17:58:50 -0700 Subject: [PATCH 1/8] encoding: Add dataclass TLV model Co-authored-by: Cursor --- src/ndn/encoding/__init__.py | 2 + src/ndn/encoding/tlv_model_v2.py | 980 +++++++++++++++++++++ tests/encoding/tlv_model_v2_test.py | 1235 +++++++++++++++++++++++++++ 3 files changed, 2217 insertions(+) create mode 100644 tests/encoding/tlv_model_v2_test.py diff --git a/src/ndn/encoding/__init__.py b/src/ndn/encoding/__init__.py index a3646e9..d5b788d 100644 --- a/src/ndn/encoding/__init__.py +++ b/src/ndn/encoding/__init__.py @@ -3,6 +3,7 @@ from .name import * from .signer import * from .tlv_model import * +from .tlv_model_v2 import tlv_encode, tlv_parse, NDNName, tlv_get_arg, tlv_set_arg from .ndn_format_0_3 import * from .ndnlp_v2 import * @@ -13,6 +14,7 @@ __all__.extend(name.__all__) __all__.extend(signer.__all__) __all__.extend(tlv_model.__all__) +__all__ += ['tlv_encode', 'tlv_parse', 'NDNName', 'tlv_get_arg', 'tlv_set_arg'] __all__.extend(ndn_format_0_3.__all__) __all__.extend(ndnlp_v2.__all__) diff --git a/src/ndn/encoding/tlv_model_v2.py b/src/ndn/encoding/tlv_model_v2.py index e69de29..f17dea4 100644 --- a/src/ndn/encoding/tlv_model_v2.py +++ b/src/ndn/encoding/tlv_model_v2.py @@ -0,0 +1,980 @@ +# ----------------------------------------------------------------------------- +# Copyright (C) 2019-2020 The python-ndn authors +# +# This file is part of python-ndn. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ----------------------------------------------------------------------------- +""" +Dataclass-based TLV encoding/decoding (v2 API). + +Usage:: + + from dataclasses import dataclass, field + from typing import List, Optional + from ndn.encoding import tlv_encode, tlv_parse, NDNName + + @dataclass + class Inner: + value: int = field(default=None, metadata={'tlv_type': 0x01}) + + @dataclass + class Outer: + name: NDNName = field(default=None, metadata={'tlv_type': 0x07}) + count: int = field(default=None, metadata={'tlv_type': 0x0a}) + payload: bytes = field(default=None, metadata={'tlv_type': 0x15}) + sub: Inner = field(default=None, metadata={'tlv_type': 0x16}) + tags: List[bytes] = field(default_factory=list, + metadata={'tlv_type': 0x17}) + + wire = tlv_encode(obj) + obj = tlv_parse(Outer, wire) + +Field-kind inference from Python annotation +------------------------------------------- ++--------------------------------------------+----------+------------------+ +| Annotation | Kind | Old equivalent | ++============================================+==========+==================+ +| int / Enum / Flag subclass | uint | UintField | ++--------------------------------------------+----------+------------------+ +| bool | bool | BoolField | ++--------------------------------------------+----------+------------------+ +| bytes / bytearray / memoryview | bytes | BytesField | ++--------------------------------------------+----------+------------------+ +| str | str | BytesField | +| | | (is_string=True) | ++--------------------------------------------+----------+------------------+ +| NDNName (sentinel) | name | NameField | ++--------------------------------------------+----------+------------------+ +| Any @dataclass type | model | ModelField | ++--------------------------------------------+----------+------------------+ +| List[T] | repeated | RepeatedField | ++--------------------------------------------+----------+------------------+ +| Dict[K, V] | map | MapField | ++--------------------------------------------+----------+------------------+ +| None + field_type='offset_marker' | (zero) | OffsetMarker | ++--------------------------------------------+----------+------------------+ +| bytes + field_type='sig_value' | (special)| SignatureValue | ++--------------------------------------------+----------+------------------+ +| NDNName + field_type='interest_name' | (special)| InterestNameField| ++--------------------------------------------+----------+------------------+ + +Supported metadata keys +----------------------- +``'tlv_type'`` int TLV type number (required except for offset_marker) +``'fixed_len'`` int Force uint value width: 1, 2, 4, or 8 bytes +``'ignore_critical' bool Suppress DecodeError for nested model parsing +``'field_type'`` str Explicit kind override when inference is insufficient + +For **map** fields (``Dict[K, V]``): +``'val_tlv_type'`` int TLV type for map values (required) + +For **sig_value** fields: +``'cover_start'`` str Name of the offset_marker field where sig coverage begins +``'digest_cover_start' str Same or different offset_marker; where digest coverage begins +``'digest_cover_end'`` str Offset_marker after sig_value; where digest coverage ends + +Signature machinery markers (set by caller before tlv_encode / tlv_parse): +``markers['##signer']`` Signer instance; absent means unsigned +``markers['##need_digest']`` True ⟹ insert/compute ParametersSha256DigestComponent + +Signature machinery markers (set by tlv_encode / tlv_parse internally): +``markers['##sig_covered_part']`` list[memoryview | bytes]: regions covered by sig +``markers['##sig_value_buf']`` writable memoryview into the placeholder bytes +``markers['##shrink_len']`` int: bytes trimmed from end after sig finalization +``markers['##digest_buf']`` writable memoryview into the digest component value +``markers[fname]`` int: recorded byte offset for each offset_marker field +""" +import dataclasses +import struct +import typing +from enum import Enum, Flag +from hashlib import sha256 + +from .tlv_type import VarBinaryStr, is_binary_str +from .tlv_var import write_tl_num, parse_tl_num, get_tl_num_size +from .name import Name, Component +from .tlv_model import DecodeError + + +__all__ = [ + 'tlv_encode', 'tlv_parse', 'NDNName', 'DecodeError', + 'tlv_get_arg', 'tlv_set_arg', +] + +# Kinds that occupy zero wire bytes and may not have a 'tlv_type' metadata key. +_ZERO_WIRE_KINDS = frozenset({'offset_marker'}) + + +# --------------------------------------------------------------------------- +# NDNName sentinel — used as a type annotation for NDN Name fields +# --------------------------------------------------------------------------- + +class NDNName: + """ + Sentinel annotation type that marks a field as an NDN Name. + + Use it wherever you would have used :class:`~ndn.encoding.NameField` in + the old metaclass API:: + + name: NDNName = field(default=None, metadata={'tlv_type': 0x07}) + # repeated Names: + names: List[NDNName] = field(default_factory=list, + metadata={'tlv_type': 0x07}) + + The actual runtime value is :any:`FormalName` (a list of encoded + component bytes), exactly as returned by the old NameField. + """ + + +# --------------------------------------------------------------------------- +# Annotation helpers +# --------------------------------------------------------------------------- + +def _unwrap_optional(annotation): + """Return T for Optional[T] = Union[T, None]; otherwise return unchanged.""" + if typing.get_origin(annotation) is typing.Union: + args = [a for a in typing.get_args(annotation) if a is not type(None)] + if len(args) == 1: + return args[0] + return annotation + + +def _infer_kind(annotation, metadata: dict) -> str: + """ + Determine the TLV field kind from a Python type annotation plus metadata. + + Returns one of: ``'uint'``, ``'bool'``, ``'bytes'``, ``'str'``, + ``'name'``, ``'model'``, ``'repeated'``. + + The ``'field_type'`` metadata key overrides automatic inference. + """ + if 'field_type' in metadata: + return metadata['field_type'] + + annotation = _unwrap_optional(annotation) + origin = typing.get_origin(annotation) + + if origin is list: + return 'repeated' + if origin is dict: + return 'map' + if annotation is NDNName: + return 'name' + # bool must be checked before int since bool is a subclass of int + if annotation is bool: + return 'bool' + if annotation is int or ( + isinstance(annotation, type) + and issubclass(annotation, (int, Enum, Flag)) + and annotation is not bool): + return 'uint' + if annotation in (bytes, bytearray, memoryview): + return 'bytes' + if annotation is str: + return 'str' + if dataclasses.is_dataclass(annotation): + return 'model' + + raise TypeError( + f'Cannot infer TLV field kind from annotation {annotation!r}. ' + f"Use metadata key 'field_type' to override." + ) + + +def _element_annotation(annotation): + """Extract T from List[T]; falls back to bytes.""" + annotation = _unwrap_optional(annotation) + args = typing.get_args(annotation) + return args[0] if args else bytes + + +def _map_annotations(annotation): + """Extract (K, V) from Dict[K, V]; falls back to (str, bytes).""" + annotation = _unwrap_optional(annotation) + args = typing.get_args(annotation) + if len(args) == 2: + return args[0], args[1] + return str, bytes + + +def _map_key_meta(metadata: dict) -> dict: + """Build a synthetic metadata dict for a map key sub-field.""" + return {'tlv_type': metadata['tlv_type']} + + +def _map_val_meta(metadata: dict) -> dict: + """Build a synthetic metadata dict for a map value sub-field.""" + m = {'tlv_type': metadata['val_tlv_type']} + if 'ignore_critical' in metadata: + m['ignore_critical'] = metadata['ignore_critical'] + return m + + +# --------------------------------------------------------------------------- +# Interest-name helpers (used by both pass-1 and pass-2) +# --------------------------------------------------------------------------- + +def _encoded_length_interest_name(fname: str, val, metadata: dict, + markers: dict) -> int: + """ + Size pass for an Interest Name field. + + Mirrors ``InterestNameField.encoded_length``. If ``markers['##need_digest']`` + is truthy and the name does not already contain a + ``ParametersSha256DigestComponent``, 34 extra bytes are reserved for one. + """ + if val is None: + return 0 + type_num = metadata['tlv_type'] + need_digest = markers.get('##need_digest', False) + + # Normalize to a list of component bytes. + if isinstance(val, str): + name = Name.from_str(val) + elif is_binary_str(val): + name = Name.decode(val)[0] + else: + name = list(val) + for i, comp in enumerate(name): + if isinstance(comp, str): + name[i] = Component.from_str(Component.escape_str(comp)) + elif not is_binary_str(comp): + raise TypeError(f'{fname}: invalid name component {comp!r}') + + # Locate an existing ParametersSha256DigestComponent (at most one allowed). + digest_pos = None + for i, comp in enumerate(name): + if Component.get_type(comp) == Component.TYPE_PARAMETERS_SHA256: + if len(Component.get_value(comp)) != 32: + raise ValueError( + f'{fname}: ParametersSha256DigestComponent must be 32 bytes') + if need_digest: + if digest_pos is None: + digest_pos = i + else: + raise ValueError( + f'{fname}: multiple ParametersSha256DigestComponent in name') + + markers[f'{fname}##digest_pos'] = digest_pos + markers[f'{fname}##preprocessed_name'] = name + + comp_total = sum(len(c) for c in name) + if need_digest and digest_pos is None: + # Reserve space for a new digest component: T(1B) + L(1B) + V(32B). + comp_total += (get_tl_num_size(Component.TYPE_PARAMETERS_SHA256) + + get_tl_num_size(32) + 32) + + markers[f'{fname}##name_value_len'] = comp_total + return get_tl_num_size(type_num) + get_tl_num_size(comp_total) + comp_total + + +def _encode_into_interest_name(fname: str, val, metadata: dict, markers: dict, + wire: VarBinaryStr, offset: int) -> int: + """ + Write pass for an Interest Name field. + + Mirrors ``InterestNameField.encode_into``. Appends non-digest name + components to ``markers['##sig_covered_part']`` (wire slices) and stores + the writable digest-value buffer in ``markers['##digest_buf']``. + """ + if val is None: + return 0 + type_num = metadata['tlv_type'] + name = markers[f'{fname}##preprocessed_name'] + comp_total = markers[f'{fname}##name_value_len'] + digest_pos = markers[f'{fname}##digest_pos'] + need_digest = markers.get('##need_digest', False) + sig_covered_part = markers.setdefault('##sig_covered_part', []) + + origin = offset + t_sz = write_tl_num(type_num, wire, offset); offset += t_sz + l_sz = write_tl_num(comp_total, wire, offset); offset += l_sz + cover_start = offset + + for i, comp in enumerate(name): + comp_len = len(comp) + wire[offset:offset + comp_len] = comp + if i == digest_pos: + if offset > cover_start: + sig_covered_part.append(wire[cover_start:offset]) + # Value of the digest component sits after T + L (each 1 byte for + # TYPE_PARAMETERS_SHA256=2 < 253 and length=32 < 253). + c_t_sz = get_tl_num_size(Component.TYPE_PARAMETERS_SHA256) + c_l_sz = get_tl_num_size(32) + markers['##digest_buf'] = wire[offset + c_t_sz + c_l_sz:offset + comp_len] + cover_start = offset + comp_len + offset += comp_len + + if offset > cover_start: + sig_covered_part.append(wire[cover_start:offset]) + + if need_digest and digest_pos is None: + # Append a new ParametersSha256DigestComponent at the end of the name. + c_t_sz = write_tl_num(Component.TYPE_PARAMETERS_SHA256, wire, offset) + offset += c_t_sz + c_l_sz = write_tl_num(32, wire, offset) + offset += c_l_sz + markers['##digest_buf'] = wire[offset:offset + 32] + # Keep the preprocessed name up-to-date for get_final_name use. + name.append(bytes(wire[offset - c_t_sz - c_l_sz:offset + 32])) + offset += 32 + + return offset - origin + + +# --------------------------------------------------------------------------- +# Post-encoding finalization (signature + SHA-256 digest) +# --------------------------------------------------------------------------- + +def _finalize_encode(markers: dict, mv: memoryview, model_end: int) -> int: + """ + Called by :func:`tlv_encode` after all bytes have been written. + + 1. Asks the signer to fill in the signature-value placeholder, updates the + inline length byte if the actual signature is shorter (ECDSA), and + records ``markers['##shrink_len']``. + 2. If ``markers['##need_digest']`` is set, computes ``SHA-256`` over the + digest-covered range and writes it into the name's digest-component + placeholder (``markers['##digest_buf']``). + + Returns *shrink_size* (0 for fixed-length signature schemes like HMAC/EdDSA). + All offsets in *markers* are absolute positions within *mv*. + """ + signer = markers.get('##signer') + shrink_size = 0 + + if signer is not None and '##sig_value_buf' in markers: + sig_value_buf = markers['##sig_value_buf'] + alloc_size = len(sig_value_buf) + real_size = signer.write_signature_value( + sig_value_buf, markers.get('##sig_covered_part', [])) + shrink_size = alloc_size - real_size + markers['##shrink_len'] = shrink_size + if shrink_size > 0: + if alloc_size >= 253: + raise ValueError( + f'Signature with variable length ≥ 253 bytes is not supported ' + f'(allocated {alloc_size})') + markers['##sig_wire_l_field'][0] = real_size + + if markers.get('##need_digest') and '##digest_buf' in markers: + d_start_field = markers.get('##_digest_cover_start_field') + d_end_field = markers.get('##_digest_cover_end_field') + d_start = markers[d_start_field] if (d_start_field and d_start_field in markers) else 0 + d_end = markers[d_end_field] if (d_end_field and d_end_field in markers) else model_end + d_end -= shrink_size + markers['##digest_buf'][:] = sha256(bytes(mv[d_start:d_end])).digest() + + return shrink_size + + +# --------------------------------------------------------------------------- +# Encoding — pass 1: size computation +# --------------------------------------------------------------------------- + +def _uint_value_len(val: int, fname: str, fixed_len) -> int: + if fixed_len is not None: + n = fixed_len + elif val <= 0xFF: + n = 1 + elif val <= 0xFFFF: + n = 2 + elif val <= 0xFFFFFFFF: + n = 4 + else: + n = 8 + if val >= 0x100 ** n: + raise ValueError(f'{fname}={val!r} cannot be encoded into {n} bytes') + return n + + +def _encoded_length_field(fname: str, val, kind: str, annotation, metadata: dict, + markers: dict) -> int: + """ + Compute the encoded byte count of one TLV field (T + L + V). + + Intermediate values are cached in *markers* under ``fname##...`` keys, + exactly mirroring the convention used by the v1 :class:`~ndn.encoding.Field` + subclasses. Returns 0 when the field is absent (*val* is ``None``/falsy + for bool). + """ + # Zero-wire kinds: handled before looking up tlv_type. + if kind == 'offset_marker': + return 0 + + if kind == 'sig_value': + signer = markers.get('##signer') + if signer is None: + return 0 + type_num = metadata['tlv_type'] + sig_size = signer.get_signature_value_size() + markers[f'{fname}##sig_size'] = sig_size + markers.setdefault('##sig_covered_part', []) + return get_tl_num_size(type_num) + get_tl_num_size(sig_size) + sig_size + + if kind == 'interest_name': + return _encoded_length_interest_name(fname, val, metadata, markers) + + type_num = metadata['tlv_type'] + + # BoolField: present if truthy, absent otherwise + if kind == 'bool': + return (get_tl_num_size(type_num) + 1) if val else 0 + + if val is None: + return 0 + + if kind == 'uint': + if isinstance(val, (Enum, Flag)): + val = val.value + if not isinstance(val, int) or val < 0: + raise TypeError(f'{fname}={val!r} is not a non-negative integer') + fixed_len = metadata.get('fixed_len') + vlen = _uint_value_len(val, fname, fixed_len) + markers[f'{fname}##encoded_length'] = vlen + # L for uint is always 1 byte because vlen ∈ {1,2,4,8} < 253 + return get_tl_num_size(type_num) + 1 + vlen + + if kind in ('bytes', 'str'): + if isinstance(val, str): + raw = val.encode('utf-8') + markers[f'{fname}##encoded_str'] = raw + else: + raw = val + n = len(raw) + return get_tl_num_size(type_num) + get_tl_num_size(n) + n + + if kind == 'name': + # Normalise to list-of-components or a pre-encoded binary blob + name_val = val + if isinstance(name_val, str): + name_val = Name.from_str(name_val) + elif not is_binary_str(name_val): + if hasattr(name_val, '__iter__'): + name_val = list(name_val) + for i, comp in enumerate(name_val): + if isinstance(comp, str): + name_val[i] = Component.from_str(Component.escape_str(comp)) + elif not is_binary_str(comp): + raise TypeError(f'{fname}: invalid name component type') + else: + raise TypeError(f'{fname}: invalid name type') + if isinstance(name_val, list): + total_with_tl = Name.encoded_length(name_val) + else: + total_with_tl = len(name_val) + markers[f'{fname}##preprocessed_name'] = name_val + markers[f'{fname}##encoded_length_with_tl'] = total_with_tl + return total_with_tl + + if kind == 'model': + inner_markers: dict = {} + length = _encoded_length_model(val, inner_markers) + markers[f'{fname}##inner_markers'] = inner_markers + markers[f'{fname}##encoded_length'] = length + return get_tl_num_size(type_num) + get_tl_num_size(length) + length + + if kind == 'repeated': + if not val: + return 0 + elem_ann = _element_annotation(annotation) + elem_kind = _infer_kind(elem_ann, metadata) + total = 0 + for i, ele in enumerate(val): + total += _encoded_length_field( + f'{fname}[{i}]', ele, elem_kind, elem_ann, metadata, markers) + return total + + if kind == 'map': + if not val: + return 0 + key_ann, val_ann = _map_annotations(annotation) + key_meta = _map_key_meta(metadata) + vl_meta = _map_val_meta(metadata) + key_kind = _infer_kind(key_ann, key_meta) + vl_kind = _infer_kind(val_ann, vl_meta) + total = 0 + for i, (k, v) in enumerate(val.items()): + total += _encoded_length_field( + f'{fname}[{i}#k]', k, key_kind, key_ann, key_meta, markers) + total += _encoded_length_field( + f'{fname}[{i}#v]', v, vl_kind, val_ann, vl_meta, markers) + return total + + raise TypeError(f'Unknown field kind {kind!r} for {fname!r}') + + +def _encoded_length_model(obj, markers: dict) -> int: + """Compute the total encoded length for all TLV fields of a dataclass object.""" + cls = type(obj) + hints = typing.get_type_hints(cls) + total = 0 + for f in dataclasses.fields(cls): + ann = hints[f.name] + kind = _infer_kind(ann, f.metadata) + if kind not in _ZERO_WIRE_KINDS and 'tlv_type' not in f.metadata: + continue + total += _encoded_length_field( + f.name, getattr(obj, f.name), kind, ann, f.metadata, markers) + markers['##encoded_length'] = total + return total + + +# --------------------------------------------------------------------------- +# Encoding — pass 2: write bytes +# --------------------------------------------------------------------------- + +def _encode_into_field(fname: str, val, kind: str, annotation, metadata: dict, + markers: dict, wire: VarBinaryStr, offset: int) -> int: + """ + Write one TLV field into *wire* at *offset*. + + *wire* must be a writable :class:`memoryview` (or :class:`bytearray`). + Returns the number of bytes written. Must be called after the matching + :func:`_encoded_length_field` call so that ``markers`` is populated. + """ + # Zero-wire kinds: handled before looking up tlv_type. + if kind == 'offset_marker': + markers[fname] = offset + return 0 + + if kind == 'sig_value': + signer = markers.get('##signer') + if signer is None: + return 0 + type_num = metadata['tlv_type'] + sig_size = markers[f'{fname}##sig_size'] + # Collect the covered region: from cover_start up to current offset. + cover_start_field = metadata.get('cover_start') + cover_start = markers.get(cover_start_field, 0) if cover_start_field else 0 + markers.setdefault('##sig_covered_part', []).append(wire[cover_start:offset]) + # Store digest-coverage field names for _finalize_encode. + for mkey in ('digest_cover_start', 'digest_cover_end'): + if mkey in metadata: + markers[f'##_{mkey}_field'] = metadata[mkey] + # Write T + L (stored for in-place shrink) + placeholder V. + t_sz = write_tl_num(type_num, wire, offset) + l_off = offset + t_sz + l_sz = write_tl_num(sig_size, wire, l_off) + markers['##sig_wire_l_field'] = wire[l_off:l_off + l_sz] + v_start = l_off + l_sz + markers['##sig_value_buf'] = wire[v_start:v_start + sig_size] + return t_sz + l_sz + sig_size + + if kind == 'interest_name': + return _encode_into_interest_name(fname, val, metadata, markers, wire, offset) + + type_num = metadata['tlv_type'] + + if kind == 'bool': + if val: + t_size = write_tl_num(type_num, wire, offset) + wire[offset + t_size] = 0 # L = 0 + return t_size + 1 + return 0 + + if val is None: + return 0 + + if kind == 'uint': + if isinstance(val, (Enum, Flag)): + val = val.value + vlen = markers[f'{fname}##encoded_length'] + t_size = write_tl_num(type_num, wire, offset) + if vlen == 1: + struct.pack_into('!BB', wire, offset + t_size, 1, val) + elif vlen == 2: + struct.pack_into('!BH', wire, offset + t_size, 2, val) + elif vlen == 4: + struct.pack_into('!BI', wire, offset + t_size, 4, val) + else: + struct.pack_into('!BQ', wire, offset + t_size, 8, val) + return t_size + 1 + vlen # T + L(1 byte) + V + + if kind in ('bytes', 'str'): + raw = markers.get(f'{fname}##encoded_str') + if raw is None: + raw = val.encode('utf-8') if isinstance(val, str) else val + n = len(raw) + t_size = write_tl_num(type_num, wire, offset) + l_size = write_tl_num(n, wire, offset + t_size) + v_start = offset + t_size + l_size + wire[v_start:v_start + n] = raw # zero-copy slice assignment + return t_size + l_size + n + + if kind == 'name': + name_val = markers[f'{fname}##preprocessed_name'] + name_len = markers[f'{fname}##encoded_length_with_tl'] + if isinstance(name_val, list): + Name.encode(name_val, wire, offset) + else: + wire[offset:offset + name_len] = name_val + return name_len + + if kind == 'model': + inner_markers = markers[f'{fname}##inner_markers'] + length = markers[f'{fname}##encoded_length'] + t_size = write_tl_num(type_num, wire, offset) + l_size = write_tl_num(length, wire, offset + t_size) + _encode_into_model(val, inner_markers, wire, offset + t_size + l_size) + return t_size + l_size + length + + if kind == 'repeated': + if not val: + return 0 + elem_ann = _element_annotation(annotation) + elem_kind = _infer_kind(elem_ann, metadata) + total = 0 + for i, ele in enumerate(val): + total += _encode_into_field( + f'{fname}[{i}]', ele, elem_kind, elem_ann, metadata, markers, + wire, offset + total) + return total + + if kind == 'map': + if not val: + return 0 + key_ann, val_ann = _map_annotations(annotation) + key_meta = _map_key_meta(metadata) + vl_meta = _map_val_meta(metadata) + key_kind = _infer_kind(key_ann, key_meta) + vl_kind = _infer_kind(val_ann, vl_meta) + total = 0 + for i, (k, v) in enumerate(val.items()): + total += _encode_into_field( + f'{fname}[{i}#k]', k, key_kind, key_ann, key_meta, markers, + wire, offset + total) + total += _encode_into_field( + f'{fname}[{i}#v]', v, vl_kind, val_ann, vl_meta, markers, + wire, offset + total) + return total + + raise TypeError(f'Unknown field kind {kind!r} for {fname!r}') + + +def _encode_into_model(obj, markers: dict, wire: VarBinaryStr, offset: int) -> None: + """Write all TLV fields of a dataclass object into *wire* starting at *offset*.""" + cls = type(obj) + hints = typing.get_type_hints(cls) + for f in dataclasses.fields(cls): + ann = hints[f.name] + kind = _infer_kind(ann, f.metadata) + if kind not in _ZERO_WIRE_KINDS and 'tlv_type' not in f.metadata: + continue + offset += _encode_into_field( + f.name, getattr(obj, f.name), kind, ann, f.metadata, markers, wire, offset) + + +# --------------------------------------------------------------------------- +# Public encode entry point +# --------------------------------------------------------------------------- + +def tlv_encode(obj, wire=None, offset: int = 0, markers: dict = None): + """ + Encode a dataclass TLV object. + + **Allocating form** — ``tlv_encode(obj)`` + Allocates a new :class:`bytearray`, fills it, and returns it. + + **In-place form** — ``tlv_encode(obj, wire, offset=0)`` + Encodes into an existing *wire* (:class:`bytearray` or writable + :class:`memoryview`) starting at *offset*. Returns a zero-copy + :class:`memoryview` slice of the written region. + + :param obj: dataclass instance to encode. + :param wire: optional writable buffer. + :param offset: starting byte offset within *wire*. + :param markers: optional shared markers dict (for multi-model coordination). + :return: :class:`bytearray` (allocating) or :class:`memoryview` (in-place). + """ + if markers is None: + markers = {} + total = _encoded_length_model(obj, markers) + if wire is None: + buf = bytearray(total) + mv = memoryview(buf) + _encode_into_model(obj, markers, mv, 0) + shrink = _finalize_encode(markers, mv, total) + if shrink: + # Can't resize bytearray while memoryview exports are live (the sig/digest + # slices in markers still reference mv). Return a trimmed copy instead. + return bytearray(mv[:total - shrink]) + return buf + mv = memoryview(wire) + _encode_into_model(obj, markers, mv, offset) + shrink = _finalize_encode(markers, mv, offset + total) + return mv[offset:offset + total - shrink] + + +# --------------------------------------------------------------------------- +# Parsing +# --------------------------------------------------------------------------- + +def _make_default_instance(cls): + """ + Create a dataclass instance with all fields set to their defaults. + + Uses ``object.__new__`` to bypass ``__init__``, then sets each field: + - ``field(default=X)`` → X + - ``field(default_factory=F)`` → F() + - no default → None (same behaviour as old TlvModel.parse) + """ + obj = object.__new__(cls) + for f in dataclasses.fields(cls): + if f.default is not dataclasses.MISSING: + object.__setattr__(obj, f.name, f.default) + elif f.default_factory is not dataclasses.MISSING: + object.__setattr__(obj, f.name, f.default_factory()) + else: + object.__setattr__(obj, f.name, None) + return obj + + +def _parse_value(fname: str, kind: str, annotation, metadata: dict, + wire, offset: int, length: int, offset_btl: int, + ignore_critical: bool): + """ + Parse a single TLV *value* (V only, not T or L) from *wire*. + + :param fname: field name (for error messages). + :param kind: field kind string. + :param annotation: resolved Python type annotation. + :param metadata: dataclass field metadata dict. + :param wire: memoryview of the full wire buffer. + :param offset: byte offset of V within *wire*. + :param length: byte length of V. + :param offset_btl: byte offset of the TLV's T field within *wire* + (used by NameField to pass to ``Name.decode``). + :param ignore_critical: forwarded to nested ``tlv_parse`` calls. + :return: the parsed Python value. + """ + if kind == 'bool': + return True + + if kind == 'uint': + if length == 1: + raw = struct.unpack_from('!B', wire, offset)[0] + elif length == 2: + raw = struct.unpack_from('!H', wire, offset)[0] + elif length == 4: + raw = struct.unpack_from('!I', wire, offset)[0] + elif length == 8: + raw = struct.unpack_from('!Q', wire, offset)[0] + else: + raise ValueError( + f'{fname}: uint value length must be 1, 2, 4, or 8; got {length}') + # Auto-convert to the annotated Enum/Flag type if applicable + inner = _unwrap_optional(annotation) + if (isinstance(inner, type) + and issubclass(inner, (Enum, Flag)) + and inner is not int): + try: + return inner(raw) + except ValueError: + pass + return raw + + if kind == 'bytes': + return wire[offset:offset + length] # zero-copy memoryview slice + + if kind == 'str': + return bytes(wire[offset:offset + length]).decode('utf-8') + + if kind == 'name': + return Name.decode(wire, offset_btl)[0] + + if kind == 'model': + inner_cls = _unwrap_optional(annotation) + ignore = metadata.get('ignore_critical', ignore_critical) + return tlv_parse(inner_cls, wire[offset:offset + length], ignore) + + raise TypeError(f'Unknown kind {kind!r} for {fname!r}') + + +def tlv_parse(cls, wire, ignore_critical: bool = False, markers: dict = None): + """ + Parse a TLV-encoded buffer into a fresh dataclass instance. + + Matching follows NDN ordering rules — fields are matched in their + declaration order within *cls* (parent class fields come first, as per + standard Python dataclass inheritance). + + Unknown critical TLV types (odd type numbers) raise + :exc:`~ndn.encoding.DecodeError` unless *ignore_critical* is ``True``. + + Bytes-typed fields (``bytes``, ``bytearray``, ``memoryview`` annotations) + are returned as zero-copy :class:`memoryview` slices into *wire*. + + :param cls: dataclass class to parse into. + :param wire: TLV-encoded buffer + (:class:`bytes`, :class:`bytearray`, or :class:`memoryview`). + :param ignore_critical: suppress :exc:`DecodeError` for unknown critical + TLV types. + :param markers: optional dict for out-of-band state (offset_marker positions, + sig/digest buffers). A fresh ``{}`` is used when ``None``. + :return: populated dataclass instance. + :raises DecodeError: unknown critical TLV type encountered. + """ + if markers is None: + markers = {} + + # Wrap in memoryview for zero-copy slicing throughout the parse + if isinstance(wire, memoryview): + mv = wire + else: + mv = memoryview(wire if isinstance(wire, (bytes, bytearray)) else bytes(wire)) + + hints = typing.get_type_hints(cls) + ordered = [] + for f in dataclasses.fields(cls): + ann = hints[f.name] + kind = _infer_kind(ann, f.metadata) + if kind not in _ZERO_WIRE_KINDS and 'tlv_type' not in f.metadata: + continue + ordered.append((f.name, f.metadata, kind, ann)) + + obj = _make_default_instance(cls) + offset = 0 + field_pos = 0 # lowest index still eligible for matching + + while offset < len(mv): + offset_btl = offset + typ, sz_t = parse_tl_num(mv, offset) + offset += sz_t + length, sz_l = parse_tl_num(mv, offset) + offset += sz_l + + found = False + for i in range(field_pos, len(ordered)): + fname, meta, kind, ann = ordered[i] + if kind == 'offset_marker': + continue # never matches a wire TLV type + + if meta['tlv_type'] != typ: + continue + + # Advance any offset_markers between field_pos and i. + for j in range(field_pos, i): + jname, _, jkind, _ = ordered[j] + if jkind == 'offset_marker': + markers[jname] = offset_btl + + if kind == 'repeated': + elem_ann = _element_annotation(ann) + elem_kind = _infer_kind(elem_ann, meta) + val = _parse_value(fname, elem_kind, elem_ann, meta, + mv, offset, length, offset_btl, ignore_critical) + lst = getattr(obj, fname) + if lst is None: + lst = [] + object.__setattr__(obj, fname, lst) + lst.append(val) + field_pos = i # stay at i to accept more elements + + elif kind == 'map': + # Two-phase parse: consume key, then immediately read value TLV. + key_ann, val_ann = _map_annotations(ann) + key_meta = _map_key_meta(meta) + vl_meta = _map_val_meta(meta) + key_kind = _infer_kind(key_ann, key_meta) + vl_kind = _infer_kind(val_ann, vl_meta) + + dct = getattr(obj, fname) + if dct is None: + dct = {} + object.__setattr__(obj, fname, dct) + idx = len(dct) + + key = _parse_value(f'{fname}[{idx}#k]', key_kind, key_ann, key_meta, + mv, offset, length, offset_btl, ignore_critical) + + # advance past key value → now at the value TLV + offset += length + offset_btl = offset + _val_typ, _sz_t2 = parse_tl_num(mv, offset) + offset += _sz_t2 + length, _sz_l2 = parse_tl_num(mv, offset) + offset += _sz_l2 + + val = _parse_value(f'{fname}[{idx}#v]', vl_kind, val_ann, vl_meta, + mv, offset, length, offset_btl, ignore_critical) + dct[key] = val + field_pos = i # stay at i to accept more pairs + + elif kind == 'sig_value': + # Extract sig buffer; append covered region to ##sig_covered_part. + sig_buf = mv[offset:offset + length] + markers['##sig_value_buf'] = sig_buf + cover_start_field = meta.get('cover_start') + if cover_start_field is not None: + cover_start = markers.get(cover_start_field) + if cover_start is not None: + markers.setdefault('##sig_covered_part', []).append( + mv[cover_start:offset_btl]) + object.__setattr__(obj, fname, sig_buf) + field_pos = i + 1 + + elif kind == 'interest_name': + # Decode name; split into sig-covered components and digest buf. + name = Name.decode(mv, offset_btl)[0] + sig_cp = markers.setdefault('##sig_covered_part', []) + for comp in name: + if Component.get_type(comp) == Component.TYPE_PARAMETERS_SHA256: + markers['##digest_buf'] = Component.get_value(comp) + else: + sig_cp.append(comp) + object.__setattr__(obj, fname, name) + field_pos = i + 1 + + else: + val = _parse_value(fname, kind, ann, meta, + mv, offset, length, offset_btl, ignore_critical) + object.__setattr__(obj, fname, val) + field_pos = i + 1 + + found = True + break + + if not found and (typ & 1) and not ignore_critical: + raise DecodeError( + f'unknown critical TLV type {typ:#x} is unrecognized, ' + f'redundant, or out-of-order') + + offset += length + + return obj + + +# --------------------------------------------------------------------------- +# Marker helpers (convenience wrappers for the markers dict) +# --------------------------------------------------------------------------- + +def tlv_get_arg(markers: dict, key: str, default=None): + """ + Read a value from the *markers* dict used by :func:`tlv_encode` / + :func:`tlv_parse`. + + Equivalent to ``markers.get(key, default)``. + """ + return markers.get(key, default) + + +def tlv_set_arg(markers: dict, key: str, val) -> None: + """ + Write a value into the *markers* dict used by :func:`tlv_encode` / + :func:`tlv_parse`. + + Equivalent to ``markers[key] = val``. + """ + markers[key] = val diff --git a/tests/encoding/tlv_model_v2_test.py b/tests/encoding/tlv_model_v2_test.py new file mode 100644 index 0000000..d58ac88 --- /dev/null +++ b/tests/encoding/tlv_model_v2_test.py @@ -0,0 +1,1235 @@ +# ----------------------------------------------------------------------------- +# Copyright (C) 2019-2020 The python-ndn authors +# +# This file is part of python-ndn. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ----------------------------------------------------------------------------- +"""Tests for the dataclass-based TLV v2 API (tlv_encode / tlv_parse).""" +from dataclasses import dataclass, field +from enum import IntEnum, IntFlag +from hashlib import sha256 + +import pytest + +from ndn.encoding import ( + tlv_encode, tlv_parse, NDNName, DecodeError, tlv_get_arg, tlv_set_arg, + # v1 equivalents used for binary-compatibility checks + TlvModel, UintField, BoolField, BytesField, NameField, ModelField, + RepeatedField, Name, Signer, +) + + +# --------------------------------------------------------------------------- +# Shared dataclass fixtures (defined at module scope for get_type_hints) +# --------------------------------------------------------------------------- + +@dataclass +class _Inner: + val: int = field(default=None, metadata={'tlv_type': 0x01}) + + +@dataclass +class _Outer: + inner: _Inner = field(default=None, metadata={'tlv_type': 0x02}) + + +@dataclass +class _RepeatedUint: + words: list[int] = field(default_factory=list, + metadata={'tlv_type': 0x01, 'fixed_len': 2}) + + +@dataclass +class _RepeatedModel: + items: list[_Inner] = field(default_factory=list, metadata={'tlv_type': 0x10}) + + +# --------------------------------------------------------------------------- +# TestUintField +# --------------------------------------------------------------------------- + +class TestUintField: + """UintField: variable-width and fixed-width non-negative integers.""" + + def test_min_width_1_byte(self): + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x03}) + + obj = M(x=0) + wire = tlv_encode(obj) + assert wire == b'\x03\x01\x00' + assert tlv_parse(M, wire).x == 0 + + def test_min_width_2_bytes(self): + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x03}) + + obj = M(x=0x0100) + wire = tlv_encode(obj) + assert wire == b'\x03\x02\x01\x00' + assert tlv_parse(M, wire).x == 0x0100 + + def test_min_width_4_bytes(self): + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x03}) + + obj = M(x=0x00010000) + wire = tlv_encode(obj) + assert wire == b'\x03\x04\x00\x01\x00\x00' + assert tlv_parse(M, wire).x == 0x00010000 + + def test_min_width_8_bytes(self): + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x03}) + + obj = M(x=0x0000000100000000) + wire = tlv_encode(obj) + assert wire == b'\x03\x08\x00\x00\x00\x01\x00\x00\x00\x00' + assert tlv_parse(M, wire).x == 0x0000000100000000 + + def test_fixed_len_1(self): + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x1b, 'fixed_len': 1}) + + obj = M(x=3) + wire = tlv_encode(obj) + # T=0x1b L=0x01 V=0x03 + assert wire == b'\x1b\x01\x03' + assert tlv_parse(M, wire).x == 3 + + def test_fixed_len_2(self): + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x03, 'fixed_len': 2}) + + obj = M(x=5) + wire = tlv_encode(obj) + assert wire == b'\x03\x02\x00\x05' + assert tlv_parse(M, wire).x == 5 + + def test_none_omitted(self): + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x03}) + y: int = field(default=None, metadata={'tlv_type': 0x05}) + + wire = tlv_encode(M(x=None, y=7)) + assert wire == b'\x05\x01\x07' + p = tlv_parse(M, wire) + assert p.x is None + assert p.y == 7 + + def test_optional_annotation(self): + @dataclass + class M: + x: int | None = field(default=None, metadata={'tlv_type': 0x03}) + + obj = M(x=42) + wire = tlv_encode(obj) + assert tlv_parse(M, wire).x == 42 + + +class TestUintFieldEnum: + """IntEnum and IntFlag auto-conversion on parse.""" + + def test_intenum_roundtrip(self): + class FaceType(IntEnum): + PERSISTENT = 1 + ON_DEMAND = 2 + + @dataclass + class M: + face_type: FaceType = field(default=None, metadata={'tlv_type': 0x84}) + + obj = M(face_type=FaceType.PERSISTENT) + wire = tlv_encode(obj) + p = tlv_parse(M, wire) + assert p.face_type == FaceType.PERSISTENT + assert isinstance(p.face_type, FaceType) + + def test_intflag_roundtrip(self): + class Flags(IntFlag): + A = 1 + B = 2 + + @dataclass + class M: + f: Flags = field(default=None, metadata={'tlv_type': 0x06}) + + obj = M(f=Flags.A | Flags.B) + wire = tlv_encode(obj) + p = tlv_parse(M, wire) + assert p.f == Flags.A | Flags.B + assert isinstance(p.f, Flags) + + def test_unknown_enum_value_returns_int(self): + class Color(IntEnum): + RED = 1 + + @dataclass + class M: + c: Color = field(default=None, metadata={'tlv_type': 0x01}) + + # Wire with value 99, which is not a valid Color + wire = b'\x01\x01\x63' + p = tlv_parse(M, wire) + assert p.c == 99 + assert type(p.c) is int + + +# --------------------------------------------------------------------------- +# TestBoolField +# --------------------------------------------------------------------------- + +class TestBoolField: + """BoolField: 0-length TLV present when truthy, absent otherwise.""" + + def test_present(self): + @dataclass + class M: + flag: bool = field(default=None, metadata={'tlv_type': 0x12}) + + wire = tlv_encode(M(flag=True)) + # T=0x12 L=0x00 + assert wire == b'\x12\x00' + assert tlv_parse(M, wire).flag is True + + def test_absent_when_false(self): + @dataclass + class M: + flag: bool = field(default=None, metadata={'tlv_type': 0x12}) + + assert tlv_encode(M(flag=False)) == b'' + assert tlv_encode(M(flag=None)) == b'' + + def test_absent_field_returns_none(self): + @dataclass + class M: + flag: bool = field(default=None, metadata={'tlv_type': 0x12}) + x: int = field(default=None, metadata={'tlv_type': 0x14}) + + wire = tlv_encode(M(flag=None, x=1)) + p = tlv_parse(M, wire) + assert p.flag is None + assert p.x == 1 + + def test_optional_annotation(self): + @dataclass + class M: + flag: bool | None = field(default=None, metadata={'tlv_type': 0x12}) + + wire = tlv_encode(M(flag=True)) + assert tlv_parse(M, wire).flag is True + + +# --------------------------------------------------------------------------- +# TestBytesField +# --------------------------------------------------------------------------- + +class TestBytesField: + """BytesField: raw bytes and UTF-8 strings.""" + + def test_bytes_roundtrip(self): + @dataclass + class M: + data: bytes = field(default=None, metadata={'tlv_type': 0x15}) + + obj = M(data=b'\x01\x02\x03') + wire = tlv_encode(obj) + assert wire == b'\x15\x03\x01\x02\x03' + p = tlv_parse(M, wire) + assert bytes(p.data) == b'\x01\x02\x03' + + def test_bytes_parse_is_memoryview(self): + @dataclass + class M: + data: bytes = field(default=None, metadata={'tlv_type': 0x15}) + + wire = b'\x15\x02\xde\xad' + p = tlv_parse(M, wire) + assert isinstance(p.data, memoryview) + + def test_str_field_roundtrip(self): + @dataclass + class M: + label: str = field(default=None, metadata={'tlv_type': 0x16}) + + obj = M(label='hello') + wire = tlv_encode(obj) + assert wire == b'\x16\x05hello' + assert tlv_parse(M, wire).label == 'hello' + + def test_str_field_unicode(self): + @dataclass + class M: + s: str = field(default=None, metadata={'tlv_type': 0x16}) + + obj = M(s='日本語') + wire = tlv_encode(obj) + assert tlv_parse(M, wire).s == '日本語' + + def test_bytearray_annotation(self): + @dataclass + class M: + data: bytearray = field(default=None, metadata={'tlv_type': 0x15}) + + wire = tlv_encode(M(data=bytearray(b'abc'))) + p = tlv_parse(M, wire) + assert bytes(p.data) == b'abc' + + def test_memoryview_annotation(self): + @dataclass + class M: + data: memoryview = field(default=None, metadata={'tlv_type': 0x15}) + + src = bytearray(b'\xca\xfe') + wire = tlv_encode(M(data=memoryview(src))) + p = tlv_parse(M, wire) + assert isinstance(p.data, memoryview) + assert bytes(p.data) == b'\xca\xfe' + + def test_large_value_multibyte_length(self): + """Length field uses multi-byte varint when value ≥ 253 bytes.""" + @dataclass + class M: + data: bytes = field(default=None, metadata={'tlv_type': 0x15}) + + payload = bytes(range(256)) + wire = tlv_encode(M(data=payload)) + # Length 256 → encoded as 0xFD 0x01 0x00 (3 bytes) + assert wire[1:4] == b'\xfd\x01\x00' + p = tlv_parse(M, wire) + assert bytes(p.data) == payload + + +# --------------------------------------------------------------------------- +# TestNameField +# --------------------------------------------------------------------------- + +class TestNameField: + """NDNName: NDN Name TLV via string, FormalName list, or binary.""" + + def test_from_string(self): + @dataclass + class M: + name: NDNName = field(default=None, metadata={'tlv_type': 0x07}) + + obj = M(name='/foo/bar') + wire = tlv_encode(obj) + # 0x07 0x0a [0x08 0x03 foo] [0x08 0x03 bar] + assert wire == b'\x07\x0a\x08\x03foo\x08\x03bar' + p = tlv_parse(M, wire) + assert Name.to_str(p.name) == '/foo/bar' + + def test_from_formal_name(self): + @dataclass + class M: + name: NDNName = field(default=None, metadata={'tlv_type': 0x07}) + + formal = Name.from_str('/a/b') + wire = tlv_encode(M(name=formal)) + p = tlv_parse(M, wire) + assert Name.to_str(p.name) == '/a/b' + + def test_empty_name(self): + @dataclass + class M: + name: NDNName = field(default=None, metadata={'tlv_type': 0x07}) + + wire = tlv_encode(M(name='/')) + p = tlv_parse(M, wire) + assert p.name == [] + + def test_none_omitted(self): + @dataclass + class M: + name: NDNName = field(default=None, metadata={'tlv_type': 0x07}) + + assert tlv_encode(M(name=None)) == b'' + + def test_repeated_names(self): + """List[NDNName] — multiple Name TLVs with the same type number.""" + @dataclass + class M: + names: list[NDNName] = field(default_factory=list, + metadata={'tlv_type': 0x07}) + + obj = M(names=['/foo', '/bar']) + wire = tlv_encode(obj) + p = tlv_parse(M, wire) + assert len(p.names) == 2 + assert Name.to_str(p.names[0]) == '/foo' + assert Name.to_str(p.names[1]) == '/bar' + + +# --------------------------------------------------------------------------- +# TestModelField +# --------------------------------------------------------------------------- + +class TestModelField: + """ModelField: nested dataclass, recursively encoded.""" + + def test_basic_nested(self): + wire = tlv_encode(_Outer(inner=_Inner(val=255))) + # 0x02 (outer T) 0x03 (outer L) 0x01 0x01 0xFF + assert wire == b'\x02\x03\x01\x01\xff' + p = tlv_parse(_Outer, wire) + assert p.inner.val == 255 + + def test_absent_nested(self): + wire = tlv_encode(_Outer(inner=None)) + assert wire == b'' + p = tlv_parse(_Outer, wire) + assert p.inner is None + + def test_deeply_nested(self): + @dataclass + class Level3: + x: int = field(default=None, metadata={'tlv_type': 0x01}) + + @dataclass + class Level2: + sub: Level3 = field(default=None, metadata={'tlv_type': 0x10}) + + @dataclass + class Level1: + sub: Level2 = field(default=None, metadata={'tlv_type': 0x20}) + + obj = Level1(sub=Level2(sub=Level3(x=7))) + wire = tlv_encode(obj) + p = tlv_parse(Level1, wire) + assert p.sub.sub.x == 7 + + def test_ignore_critical_propagated(self): + """ignore_critical in metadata is forwarded to nested parse.""" + @dataclass + class Inner: + x: int = field(default=None, metadata={'tlv_type': 0x01}) + + @dataclass + class Outer: + sub: Inner = field(default=None, + metadata={'tlv_type': 0x10, + 'ignore_critical': True}) + + # Inject an unknown critical TLV (type 0x03, odd) inside sub + inner_wire = b'\x03\x01\x00' + sub_wire = bytes([0x10, len(inner_wire)]) + inner_wire + # With ignore_critical via metadata, this must not raise + p = tlv_parse(Outer, sub_wire) + assert p.sub.x is None + + +# --------------------------------------------------------------------------- +# TestRepeatedField +# --------------------------------------------------------------------------- + +class TestRepeatedField: + """RepeatedField: multiple TLVs of the same type, no outer wrapper.""" + + def test_uint_elements(self): + wire = tlv_encode(_RepeatedUint(words=[0, 1, 2])) + # Each word: T=0x01 L=0x02 V=2-byte big-endian + assert wire == b'\x01\x02\x00\x00\x01\x02\x00\x01\x01\x02\x00\x02' + p = tlv_parse(_RepeatedUint, wire) + assert p.words == [0, 1, 2] + + def test_bytes_elements(self): + @dataclass + class M: + tags: list[bytes] = field(default_factory=list, + metadata={'tlv_type': 0x17}) + + obj = M(tags=[b'a', b'bb', b'ccc']) + wire = tlv_encode(obj) + assert wire == b'\x17\x01a\x17\x02bb\x17\x03ccc' + p = tlv_parse(M, wire) + assert [bytes(t) for t in p.tags] == [b'a', b'bb', b'ccc'] + + def test_str_elements(self): + @dataclass + class M: + labels: list[str] = field(default_factory=list, + metadata={'tlv_type': 0x16}) + + obj = M(labels=['hello', 'world']) + wire = tlv_encode(obj) + p = tlv_parse(M, wire) + assert p.labels == ['hello', 'world'] + + def test_model_elements(self): + wire = tlv_encode(_RepeatedModel(items=[_Inner(val=10), _Inner(val=20)])) + p = tlv_parse(_RepeatedModel, wire) + assert [i.val for i in p.items] == [10, 20] + + def test_empty_list_produces_no_bytes(self): + wire = tlv_encode(_RepeatedUint(words=[])) + assert wire == b'' + + def test_default_factory_list_initialised_on_parse(self): + """A repeated field with no default_factory should still get a list on parse.""" + @dataclass + class M: + items: list[int] = field(metadata={'tlv_type': 0x05}) + + wire = b'\x05\x01\x01\x05\x01\x02' + p = tlv_parse(M, wire) + assert p.items == [1, 2] + + +# --------------------------------------------------------------------------- +# TestOrdering +# --------------------------------------------------------------------------- + +class TestOrdering: + """NDN TLV ordering rules: forward-only matching, critical type handling.""" + + def test_unknown_even_type_skipped(self): + """Unknown even TLV types are silently ignored (non-critical).""" + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x04}) + + # Wire: unknown even type 0x02, then known type 0x04 + wire = b'\x02\x01\xff\x04\x01\x07' + p = tlv_parse(M, wire) + assert p.x == 7 + + def test_unknown_odd_type_raises(self): + """Unknown odd TLV types are critical — must raise DecodeError.""" + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x04}) + + wire = b'\x03\x01\x00\x04\x01\x07' + with pytest.raises(DecodeError): + tlv_parse(M, wire) + + def test_ignore_critical_suppresses_error(self): + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x04}) + + wire = b'\x03\x01\x00\x04\x01\x07' + p = tlv_parse(M, wire, ignore_critical=True) + assert p.x == 7 + + def test_out_of_order_critical_field_raises(self): + """ + A critical (odd type) field appearing out-of-order is unrecognised and + must raise DecodeError. Even types are non-critical and silently dropped. + """ + @dataclass + class M: + a: int = field(default=None, metadata={'tlv_type': 0x03}) # odd = critical + b: int = field(default=None, metadata={'tlv_type': 0x05}) # odd = critical + + # b (0x05) comes first; after matching it, field_pos advances past a (0x03). + # The parser then sees 0x03 as unknown critical → DecodeError. + wire = b'\x05\x01\x02\x03\x01\x01' + with pytest.raises(DecodeError): + tlv_parse(M, wire) + + def test_out_of_order_even_field_silently_dropped(self): + """Even (non-critical) fields seen out-of-order are silently skipped.""" + @dataclass + class M: + a: int = field(default=None, metadata={'tlv_type': 0x02}) # even + b: int = field(default=None, metadata={'tlv_type': 0x04}) # even + + # b before a — a is dropped silently (non-critical) + wire = b'\x04\x01\x02\x02\x01\x01' + p = tlv_parse(M, wire) + assert p.b == 2 + assert p.a is None + + def test_out_of_order_non_critical_silently_skipped(self): + @dataclass + class M: + a: int = field(default=None, metadata={'tlv_type': 0x04}) # even, non-critical when out of order + b: int = field(default=None, metadata={'tlv_type': 0x06}) + + # b before a — a is even so silently skipped + wire = b'\x06\x01\x02\x04\x01\x01' + p = tlv_parse(M, wire) + assert p.b == 2 + assert p.a is None + + +# --------------------------------------------------------------------------- +# TestInheritance +# --------------------------------------------------------------------------- + +class TestInheritance: + """Dataclass inheritance: parent fields come first (no IncludeBase needed).""" + + def test_parent_fields_encoded_first(self): + @dataclass + class Base: + x: int = field(default=None, metadata={'tlv_type': 0x01}) + + @dataclass + class Child(Base): + y: int = field(default=None, metadata={'tlv_type': 0x03}) + + wire = tlv_encode(Child(x=1, y=2)) + # x (0x01) must appear before y (0x03) + assert wire == b'\x01\x01\x01\x03\x01\x02' + p = tlv_parse(Child, wire) + assert p.x == 1 and p.y == 2 + + def test_child_only_fields(self): + @dataclass + class Base: + x: int = field(default=None, metadata={'tlv_type': 0x01}) + + @dataclass + class Child(Base): + y: int = field(default=None, metadata={'tlv_type': 0x03}) + + wire = tlv_encode(Child(x=None, y=5)) + assert wire == b'\x03\x01\x05' + p = tlv_parse(Child, wire) + assert p.x is None + assert p.y == 5 + + def test_non_tlv_fields_ignored(self): + """Fields without 'tlv_type' in metadata are silently skipped.""" + @dataclass + class M: + internal: str = field(default='ignored') # no metadata + x: int = field(default=None, metadata={'tlv_type': 0x01}) + + wire = tlv_encode(M(internal='should_not_appear', x=42)) + assert wire == b'\x01\x01\x2a' + p = tlv_parse(M, wire) + assert p.x == 42 + + +# --------------------------------------------------------------------------- +# TestZeroCopyAndInPlace +# --------------------------------------------------------------------------- + +class TestZeroCopyAndInPlace: + """Zero-copy memoryview slices and in-place buffer encoding.""" + + def test_bytes_parse_shares_buffer(self): + """Parsed bytes field is a memoryview slice — no copy.""" + @dataclass + class M: + data: bytes = field(default=None, metadata={'tlv_type': 0x15}) + + wire = bytearray(b'\x15\x04\xde\xad\xbe\xef') + p = tlv_parse(M, wire) + assert isinstance(p.data, memoryview) + assert bytes(p.data) == b'\xde\xad\xbe\xef' + + def test_inplace_encode_returns_memoryview(self): + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x03}) + + buf = bytearray(20) + mv = tlv_encode(M(x=7), buf, offset=5) + assert isinstance(mv, memoryview) + assert bytes(mv) == b'\x03\x01\x07' + # Bytes written at correct position + assert buf[5:8] == b'\x03\x01\x07' + # Surrounding bytes untouched + assert buf[:5] == b'\x00' * 5 + assert buf[8:] == b'\x00' * 12 + + def test_inplace_matches_standalone(self): + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x03}) + y: bytes = field(default=None, metadata={'tlv_type': 0x15}) + + obj = M(x=300, y=b'hello') + standalone = tlv_encode(obj) + buf = bytearray(len(standalone) + 10) + mv = tlv_encode(obj, buf, offset=3) + assert bytes(mv) == bytes(standalone) + + def test_inplace_memoryview_buffer(self): + """In-place encoding also works with a memoryview target.""" + @dataclass + class M: + x: int = field(default=None, metadata={'tlv_type': 0x05}) + + buf = bytearray(10) + mv_buf = memoryview(buf) + result = tlv_encode(M(x=1), mv_buf, offset=2) + assert bytes(result) == b'\x05\x01\x01' + + +# --------------------------------------------------------------------------- +# TestDefaultHandling +# --------------------------------------------------------------------------- + +class TestDefaultHandling: + """Default values and field initialisation during parse.""" + + def test_field_with_explicit_default_preserved_if_absent(self): + @dataclass + class M: + x: int = field(default=42, metadata={'tlv_type': 0x01}) + + # Wire that does not contain field x + wire = b'' + p = tlv_parse(M, wire) + assert p.x == 42 + + def test_field_without_default_is_none_if_absent(self): + @dataclass + class M: + x: int = field(metadata={'tlv_type': 0x01}) + + wire = b'' + p = tlv_parse(M, wire) + assert p.x is None + + def test_default_factory_list_preserved_if_absent(self): + wire = tlv_encode(_RepeatedUint(words=[])) + p = tlv_parse(_RepeatedUint, wire) + assert p.words == [] + + +# --------------------------------------------------------------------------- +# TestBinaryCompatibility +# --------------------------------------------------------------------------- + +class TestBinaryCompatibility: + """Byte-for-byte compatibility with the v1 TlvModel metaclass API.""" + + def test_uint_compat(self): + class V1(TlvModel): + sig_type = UintField(0x1b, fixed_len=1) + nonce = UintField(0x26) + + @dataclass + class V2: + sig_type: int = field(default=None, + metadata={'tlv_type': 0x1b, 'fixed_len': 1}) + nonce: int = field(default=None, metadata={'tlv_type': 0x26}) + + v1 = V1(); v1.sig_type = 3; v1.nonce = 42 + assert bytes(v1.encode()) == bytes(tlv_encode(V2(sig_type=3, nonce=42))) + + def test_bool_compat(self): + class V1(TlvModel): + flag = BoolField(0x12) + count = UintField(0x0a) + + @dataclass + class V2: + flag: bool = field(default=None, metadata={'tlv_type': 0x12}) + count: int = field(default=None, metadata={'tlv_type': 0x0a}) + + for flag_val in (True, False, None): + v1 = V1(); v1.flag = flag_val; v1.count = 5 + v2 = V2(flag=flag_val, count=5) + assert bytes(v1.encode()) == bytes(tlv_encode(v2)) + + def test_bytes_compat(self): + class V1(TlvModel): + raw = BytesField(0x15) + label = BytesField(0x16, is_string=True) + + @dataclass + class V2: + raw: bytes = field(default=None, metadata={'tlv_type': 0x15}) + label: str = field(default=None, metadata={'tlv_type': 0x16}) + + v1 = V1(); v1.raw = b'\x01\x02\x03'; v1.label = 'hi' + v2 = V2(raw=b'\x01\x02\x03', label='hi') + assert bytes(v1.encode()) == bytes(tlv_encode(v2)) + + def test_name_compat(self): + class V1(TlvModel): + name = NameField() + + @dataclass + class V2: + name: NDNName = field(default=None, metadata={'tlv_type': 0x07}) + + v1 = V1(); v1.name = '/foo/bar' + v2 = V2(name='/foo/bar') + assert bytes(v1.encode()) == bytes(tlv_encode(v2)) + + def test_model_compat(self): + class V1Inner(TlvModel): + val = UintField(0x01) + + class V1Outer(TlvModel): + inner = ModelField(0x10, V1Inner) + + @dataclass + class V2Inner: + val: int = field(default=None, metadata={'tlv_type': 0x01}) + + @dataclass + class V2Outer: + inner: V2Inner = field(default=None, metadata={'tlv_type': 0x10}) + + v1 = V1Outer(); v1.inner = V1Inner(); v1.inner.val = 99 + v2 = V2Outer(inner=V2Inner(val=99)) + assert bytes(v1.encode()) == bytes(tlv_encode(v2)) + + def test_repeated_uint_compat(self): + class V1(TlvModel): + words = RepeatedField(UintField(0x01, fixed_len=2)) + + v1 = V1(); v1.words = [0, 1, 2] + v2 = _RepeatedUint(words=[0, 1, 2]) + assert bytes(v1.encode()) == bytes(tlv_encode(v2)) + + def test_repeated_model_compat(self): + class V1Inner(TlvModel): + val = UintField(0x01) + + class V1Rep(TlvModel): + items = RepeatedField(ModelField(0x10, V1Inner)) + + v1 = V1Rep() + r1 = V1Inner(); r1.val = 10 + r2 = V1Inner(); r2.val = 20 + v1.items = [r1, r2] + + v2 = _RepeatedModel(items=[_Inner(val=10), _Inner(val=20)]) + assert bytes(v1.encode()) == bytes(tlv_encode(v2)) + + def test_parse_interop(self): + """Wire produced by v1 can be parsed by v2 and vice-versa.""" + class V1(TlvModel): + name = NameField() + count = UintField(0x0a) + + @dataclass + class V2: + name: NDNName = field(default=None, metadata={'tlv_type': 0x07}) + count: int = field(default=None, metadata={'tlv_type': 0x0a}) + + v1 = V1(); v1.name = '/test'; v1.count = 7 + wire_from_v1 = bytes(v1.encode()) + + p = tlv_parse(V2, wire_from_v1) + assert Name.to_str(p.name) == '/test' + assert p.count == 7 + + v2 = V2(name='/test', count=7) + wire_from_v2 = bytes(tlv_encode(v2)) + + p2 = V1.parse(wire_from_v2) + assert Name.to_str(p2.name) == '/test' + assert p2.count == 7 + + +# --------------------------------------------------------------------------- +# MapField tests +# --------------------------------------------------------------------------- + +@dataclass +class _StrBytesMap: + entries: dict[str, bytes] = field(default_factory=dict, metadata={ + 'tlv_type': 0x21, + 'val_tlv_type': 0x23, + }) + + +@dataclass +class _Inner2: + value: int = field(default=None, metadata={'tlv_type': 0x01}) + + +@dataclass +class _StrModelMap: + entries: dict[str, _Inner2] = field(default_factory=dict, metadata={ + 'tlv_type': 0x21, + 'val_tlv_type': 0x22, + }) + + +class TestMapField: + def test_str_bytes_roundtrip(self): + obj = _StrBytesMap(entries={'alpha': b'\x01\x02', 'beta': b'\x03'}) + wire = tlv_encode(obj) + p = tlv_parse(_StrBytesMap, wire) + assert list(p.entries.keys()) == ['alpha', 'beta'] + assert bytes(p.entries['alpha']) == b'\x01\x02' + assert bytes(p.entries['beta']) == b'\x03' + + def test_insertion_order_preserved(self): + """Dict round-trip must preserve the original key insertion order.""" + obj = _StrBytesMap(entries={'z': b'\x00', 'a': b'\x01', 'm': b'\x02'}) + p = tlv_parse(_StrBytesMap, tlv_encode(obj)) + assert list(p.entries.keys()) == ['z', 'a', 'm'] + + def test_empty_map_produces_no_bytes(self): + obj = _StrBytesMap(entries={}) + assert tlv_encode(obj) == b'' + + def test_none_map_produces_no_bytes(self): + obj = _StrBytesMap(entries=None) + assert tlv_encode(obj) == b'' + + def test_none_map_defaults_to_empty_on_parse(self): + """Parsing wire with no map TLVs leaves entries as the default_factory value.""" + p = tlv_parse(_StrBytesMap, b'') + assert p.entries == {} + + def test_str_model_map_roundtrip(self): + obj = _StrModelMap(entries={'x': _Inner2(value=7), 'y': _Inner2(value=99)}) + wire = tlv_encode(obj) + p = tlv_parse(_StrModelMap, wire) + assert list(p.entries.keys()) == ['x', 'y'] + assert p.entries['x'].value == 7 + assert p.entries['y'].value == 99 + + def test_v1_compat_wire(self): + """v2 map encoding must be byte-for-byte identical to v1 MapField.""" + from ndn.encoding import MapField, BytesField + + class V1Map(TlvModel): + entries = MapField(BytesField(0x21, is_string=True), BytesField(0x23)) + + v1 = V1Map() + v1.entries['alpha'] = b'\x01\x02' + v1.entries['beta'] = b'\x03' + v1_wire = bytes(v1.encode()) + + v2 = _StrBytesMap(entries={'alpha': b'\x01\x02', 'beta': b'\x03'}) + v2_wire = bytes(tlv_encode(v2)) + + assert v1_wire == v2_wire + + def test_v1_produced_wire_parsed_by_v2(self): + from ndn.encoding import MapField, BytesField + + class V1Map(TlvModel): + entries = MapField(BytesField(0x21, is_string=True), BytesField(0x23)) + + v1 = V1Map() + v1.entries['hello'] = b'\xde\xad' + wire = bytes(v1.encode()) + + p = tlv_parse(_StrBytesMap, wire) + assert bytes(p.entries['hello']) == b'\xde\xad' + + def test_bytes_values_are_memoryview_zero_copy(self): + obj = _StrBytesMap(entries={'k': b'\xca\xfe'}) + wire = tlv_encode(obj) + p = tlv_parse(_StrBytesMap, wire) + assert isinstance(p.entries['k'], memoryview) + + +# --------------------------------------------------------------------------- +# Signature machinery tests +# --------------------------------------------------------------------------- + +# ── Simple fixed-length mock signer (HMAC-like) ────────────────────────────── + +class _HmacSigner(Signer): + """Deterministic 32-byte 'HMAC' signer using SHA-256(key || content).""" + SIG_SIZE = 32 + + def __init__(self, key: bytes = b'secret'): + self._key = key + + def write_signature_info(self, sig_info): + sig_info.signature_type = 4 # HMAC_WITH_SHA256 + + def get_signature_value_size(self) -> int: + return self.SIG_SIZE + + def write_signature_value(self, wire, contents) -> int: + h = sha256(self._key) + for blk in contents: + h.update(bytes(blk)) + sig = h.digest() + wire[:] = sig + return len(sig) + + def verify(self, sig: bytes, contents) -> bool: + buf = bytearray(self.SIG_SIZE) + mv = memoryview(buf) + self.write_signature_value(mv, contents) + return bytes(buf) == bytes(sig) + + +# ── Variable-length mock signer (ECDSA-like, sometimes shorter) ─────────────── + +class _EcdsaSigner(Signer): + """Always signs with 71 bytes, but reports max 72 (tests shrink path).""" + MAX_SIZE = 72 + REAL_SIZE = 71 + + def write_signature_info(self, sig_info): + sig_info.signature_type = 3 # SHA256_WITH_ECDSA + + def get_signature_value_size(self): + return self.MAX_SIZE + + def write_signature_value(self, wire, contents): + for i in range(self.REAL_SIZE): + wire[i] = i & 0xFF + return self.REAL_SIZE + + +# ── Data-like model (no digest, no interest-name) ──────────────────────────── + +@dataclass +class _SigInfo: + signature_type: int = field(default=None, metadata={'tlv_type': 0x1b, 'fixed_len': 1}) + + +@dataclass +class _DataValue: + _sig_cover_start: None = field(default=None, metadata={'field_type': 'offset_marker'}) + name: NDNName = field(default=None, metadata={'tlv_type': 0x07}) + content: bytes | None = field(default=None, metadata={'tlv_type': 0x15}) + signature_info: _SigInfo | None = field(default=None, metadata={'tlv_type': 0x16}) + signature_value: bytes | None = field(default=None, metadata={ + 'tlv_type': 0x17, + 'field_type': 'sig_value', + 'cover_start': '_sig_cover_start', + }) + + +# ── Interest-like model (interest_name + digest) ────────────────────────────── + +@dataclass +class _InterestValue: + name: NDNName = field(default=None, metadata={ + 'tlv_type': 0x07, 'field_type': 'interest_name'}) + nonce: int | None = field(default=None, metadata={ + 'tlv_type': 0x0a, 'fixed_len': 4}) + _sig_cover_start: None = field(default=None, metadata={'field_type': 'offset_marker'}) + application_parameters: bytes | None = field(default=None, metadata={'tlv_type': 0x24}) + signature_info: _SigInfo | None = field(default=None, metadata={'tlv_type': 0x2c}) + signature_value: bytes | None = field(default=None, metadata={ + 'tlv_type': 0x2e, + 'field_type': 'sig_value', + 'cover_start': '_sig_cover_start', + 'digest_cover_start': '_sig_cover_start', + 'digest_cover_end': '_digest_cover_end', + }) + _digest_cover_end: None = field(default=None, metadata={'field_type': 'offset_marker'}) + + +class TestSignatureMachinery: + # ── offset_marker ────────────────────────────────────────────────────────── + + def test_offset_marker_produces_no_bytes(self): + @dataclass + class M: + _mark: None = field(default=None, metadata={'field_type': 'offset_marker'}) + v: int = field(default=None, metadata={'tlv_type': 0x01}) + + wire = tlv_encode(M(v=7)) + assert wire == bytes(tlv_encode(_Inner(val=7))) # only the uint TLV, no extra bytes + + def test_offset_marker_records_position_during_encode(self): + @dataclass + class M: + a: int = field(default=None, metadata={'tlv_type': 0x01}) + _mark: None = field(default=None, metadata={'field_type': 'offset_marker'}) + b: int = field(default=None, metadata={'tlv_type': 0x03}) + + markers = {} + tlv_encode(M(a=1, b=2), markers=markers) + # 'a' occupies 3 bytes (T=1, L=1, V=1), so _mark records offset 3. + assert markers['_mark'] == 3 + + def test_offset_marker_records_position_during_parse(self): + @dataclass + class M: + a: int = field(default=None, metadata={'tlv_type': 0x01}) + _mark: None = field(default=None, metadata={'field_type': 'offset_marker'}) + b: int = field(default=None, metadata={'tlv_type': 0x03}) + + wire = tlv_encode(M(a=1, b=2)) + markers = {} + tlv_parse(M, wire, markers=markers) + # offset_btl of 'b' = 3 (after 'a'), so _mark records 3. + assert markers.get('_mark') == 3 + + # ── sig_value / Data-like encoding ──────────────────────────────────────── + + def test_data_encode_produces_signature(self): + signer = _HmacSigner() + obj = _DataValue(name='/test', content=b'hello') + obj.signature_info = _SigInfo() + signer.write_signature_info(obj.signature_info) + + markers = {'##signer': signer} + wire = tlv_encode(obj, markers=markers) + + # Signature value TLV must be present. + assert 0x17 in bytes(wire) + # Wire must end with 32 sig bytes (preceded by TL 17 20). + assert wire[-34:-32] == b'\x17\x20' + + def test_data_encode_decode_roundtrip(self): + signer = _HmacSigner() + obj = _DataValue(name='/test/data', content=b'payload') + obj.signature_info = _SigInfo() + signer.write_signature_info(obj.signature_info) + + markers = {'##signer': signer} + wire = tlv_encode(obj, markers=markers) + + parse_markers = {} + p = tlv_parse(_DataValue, wire, markers=parse_markers) + assert Name.to_str(p.name) == '/test/data' + assert bytes(p.content) == b'payload' + assert p.signature_info.signature_type == 4 # HMAC_WITH_SHA256 + + def test_data_signature_verifies(self): + signer = _HmacSigner() + obj = _DataValue(name='/verify/me', content=b'data') + obj.signature_info = _SigInfo() + signer.write_signature_info(obj.signature_info) + + enc_markers = {'##signer': signer} + wire = tlv_encode(obj, markers=enc_markers) + enc_covered = enc_markers['##sig_covered_part'] + + parse_markers = {} + p = tlv_parse(_DataValue, wire, markers=parse_markers) + parse_covered = parse_markers.get('##sig_covered_part', []) + sig_buf = parse_markers['##sig_value_buf'] + + assert signer.verify(bytes(sig_buf), parse_covered) + + def test_data_signature_is_deterministic(self): + """Same object encoded twice with the same signer → identical wires.""" + signer = _HmacSigner() + obj1 = _DataValue(name='/det/test', content=b'hello') + obj1.signature_info = _SigInfo() + signer.write_signature_info(obj1.signature_info) + + obj2 = _DataValue(name='/det/test', content=b'hello') + obj2.signature_info = _SigInfo() + signer.write_signature_info(obj2.signature_info) + + w1 = tlv_encode(obj1, markers={'##signer': signer}) + w2 = tlv_encode(obj2, markers={'##signer': signer}) + assert bytes(w1) == bytes(w2) + + def test_ecdsa_signer_shrinks_wire(self): + signer = _EcdsaSigner() + obj = _DataValue(name='/shrink', content=b'x') + obj.signature_info = _SigInfo() + signer.write_signature_info(obj.signature_info) + + markers = {'##signer': signer} + wire = tlv_encode(obj, markers=markers) + + # Allocated 72 bytes, actual 71 → last byte trimmed. + assert markers['##shrink_len'] == 1 + # The sig_value TLV's L byte should now read 71 (0x47). + sig_tlv_idx = bytes(wire).index(0x17) # find sig_value type byte + assert wire[sig_tlv_idx + 1] == 71 + + def test_unsigned_data_produces_no_sig_tlv(self): + obj = _DataValue(name='/unsigned', content=b'ok') + wire = tlv_encode(obj) + assert b'\x17' not in bytes(wire) + + # ── interest_name + digest ──────────────────────────────────────────────── + + def test_interest_name_without_digest(self): + obj = _InterestValue(name='/plain/interest', nonce=42) + wire = tlv_encode(obj) + p = tlv_parse(_InterestValue, wire) + assert Name.to_str(p.name) == '/plain/interest' + assert p.nonce == 42 + + def test_interest_with_digest_appended(self): + """When ##need_digest is True and no digest component exists, one is appended.""" + app_param = b'\x01\x02\x03' + obj = _InterestValue(name='/digest/test', application_parameters=app_param) + obj.signature_info = _SigInfo() + signer = _HmacSigner() + signer.write_signature_info(obj.signature_info) + + markers = {'##signer': signer, '##need_digest': True} + wire = tlv_encode(obj, markers=markers) + + # Parse back and check digest component exists in name. + p = tlv_parse(_InterestValue, wire) + name_str = Name.to_str(p.name) + assert 'params-sha256=' in name_str + + def test_interest_digest_value_is_sha256(self): + """The ParametersSha256DigestComponent must equal SHA-256 of the digest-covered part.""" + from ndn.encoding.name import Component as C + app_param = b'\xde\xad\xbe\xef' + obj = _InterestValue(name='/verify/digest', application_parameters=app_param) + obj.signature_info = _SigInfo() + signer = _HmacSigner() + signer.write_signature_info(obj.signature_info) + + markers = {'##signer': signer, '##need_digest': True} + wire = tlv_encode(obj, markers=markers) + + # Locate the ParametersSha256DigestComponent in the encoded name. + p = tlv_parse(_InterestValue, wire) + digest_comp = None + for comp in p.name: + if C.get_type(comp) == C.TYPE_PARAMETERS_SHA256: + digest_comp = comp + break + assert digest_comp is not None + + digest_val = bytes(C.get_value(digest_comp)) + # Determine what the digest should cover: find where _sig_cover_start landed. + raw = bytes(wire) + sig_cover_start = markers.get('_sig_cover_start', 0) + d_end_field = '_digest_cover_end' + sig_cover_end = markers.get(d_end_field, len(raw)) + expected = sha256(raw[sig_cover_start:sig_cover_end]).digest() + assert digest_val == expected + + def test_interest_sig_covered_part_set_on_parse(self): + """After parsing an Interest, ##sig_covered_part is populated.""" + obj = _InterestValue(name='/parse/sig', application_parameters=b'\x00') + obj.signature_info = _SigInfo() + signer = _HmacSigner() + signer.write_signature_info(obj.signature_info) + + markers = {'##signer': signer, '##need_digest': True} + wire = tlv_encode(obj, markers=markers) + + parse_markers = {} + tlv_parse(_InterestValue, wire, markers=parse_markers) + assert '##sig_covered_part' in parse_markers + assert len(parse_markers['##sig_covered_part']) > 0 + + # ── tlv_get_arg / tlv_set_arg ───────────────────────────────────────────── + + def test_tlv_get_arg_missing_returns_default(self): + m = {} + assert tlv_get_arg(m, 'x', 42) == 42 + + def test_tlv_set_arg_stores_value(self): + m = {} + tlv_set_arg(m, 'key', 'value') + assert tlv_get_arg(m, 'key') == 'value' From 85a53cf2b8311e53ab0a64b3adb96ce3a4d81201 Mon Sep 17 00:00:00 2001 From: Xinyu Ma Date: Sun, 6 Sep 2026 17:59:01 -0700 Subject: [PATCH 2/8] encoding: Add dataclass packet format Co-authored-by: Cursor --- src/ndn/encoding/ndn_format_0_3_2.py | 329 ++++++++++++++++++++++++ tests/encoding/ndn_format_0_3_2_test.py | 73 ++++++ 2 files changed, 402 insertions(+) create mode 100644 src/ndn/encoding/ndn_format_0_3_2.py create mode 100644 tests/encoding/ndn_format_0_3_2_test.py diff --git a/src/ndn/encoding/ndn_format_0_3_2.py b/src/ndn/encoding/ndn_format_0_3_2.py new file mode 100644 index 0000000..181fc4c --- /dev/null +++ b/src/ndn/encoding/ndn_format_0_3_2.py @@ -0,0 +1,329 @@ +# ----------------------------------------------------------------------------- +# Copyright (C) 2019-2020 The python-ndn authors +# Licensed under the Apache License, Version 2.0 (the "License"); +# ----------------------------------------------------------------------------- +"""NDN Packet Format v0.3 models using the dataclass TLV API.""" +import dataclasses as dc +from typing import Optional + +from .name import Name, Component +from .signer import Signer +from .tlv_model_v2 import NDNName, tlv_encode, tlv_parse +from .tlv_type import BinaryStr, VarBinaryStr, NonStrictName, FormalName +from .tlv_var import get_tl_num_size, parse_and_check_tl, write_tl_num + +__all__ = [ + 'TypeNumber', 'ContentType', 'SignatureType', 'KeyLocator', + 'SignatureInfo', + 'Links', 'MetaInfo', 'InterestParam', 'SignaturePtrs', 'make_interest', + 'make_data', 'parse_interest', 'parse_data', 'Interest', 'Data', +] + + +class TypeNumber: + INTEREST = 0x05 + DATA = 0x06 + NAME = Name.TYPE_NAME + GENERIC_NAME_COMPONENT = Component.TYPE_GENERIC + IMPLICIT_SHA256_DIGEST_COMPONENT = Component.TYPE_IMPLICIT_SHA256 + PARAMETERS_SHA256_DIGEST_COMPONENT = Component.TYPE_PARAMETERS_SHA256 + CAN_BE_PREFIX = 0x21 + MUST_BE_FRESH = 0x12 + FORWARDING_HINT = 0x1e + NONCE = 0x0a + INTEREST_LIFETIME = 0x0c + HOP_LIMIT = 0x22 + APPLICATION_PARAMETERS = 0x24 + INTEREST_SIGNATURE_INFO = 0x2c + INTEREST_SIGNATURE_VALUE = 0x2e + META_INFO = 0x14 + CONTENT = 0x15 + SIGNATURE_INFO = 0x16 + SIGNATURE_VALUE = 0x17 + CONTENT_TYPE = 0x18 + FRESHNESS_PERIOD = 0x19 + FINAL_BLOCK_ID = 0x1a + SIGNATURE_TYPE = 0x1b + KEY_LOCATOR = 0x1c + KEY_DIGEST = 0x1d + SIGNATURE_NONCE = 0x26 + SIGNATURE_TIME = 0x28 + SIGNATURE_SEQ_NUM = 0x2a + DELEGATION = 0x1f + PREFERENCE = 0x1e + + +class ContentType: + BLOB = 0 + LINK = 1 + KEY = 2 + NACK = 3 + + +class SignatureType: + NOT_SIGNED = None + DIGEST_SHA256 = 0 + SHA256_WITH_RSA = 1 + SHA256_WITH_ECDSA = 3 + HMAC_WITH_SHA256 = 4 + ED25519 = 5 + NULL = 200 + + +@dc.dataclass +class KeyLocator: + name: NDNName = dc.field( + default=None, metadata={'tlv_type': TypeNumber.NAME}) + key_digest: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.KEY_DIGEST}) + + +@dc.dataclass +class SignatureInfo: + signature_type: Optional[int] = dc.field( + default=None, metadata={ + 'tlv_type': TypeNumber.SIGNATURE_TYPE, 'fixed_len': 1}) + key_locator: Optional[KeyLocator] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.KEY_LOCATOR}) + signature_nonce: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.SIGNATURE_NONCE}) + signature_time: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.SIGNATURE_TIME}) + signature_seq_num: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.SIGNATURE_SEQ_NUM}) + + +@dc.dataclass +class Links: + names: list[NDNName] = dc.field( + default_factory=list, metadata={'tlv_type': TypeNumber.NAME}) + + +@dc.dataclass +class InterestPacketValue: + name: NDNName = dc.field(default='/', metadata={ + 'tlv_type': TypeNumber.NAME, 'field_type': 'interest_name'}) + can_be_prefix: bool = dc.field( + default=False, metadata={'tlv_type': TypeNumber.CAN_BE_PREFIX}) + must_be_fresh: bool = dc.field( + default=False, metadata={'tlv_type': TypeNumber.MUST_BE_FRESH}) + forwarding_hint: Optional[Links] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.FORWARDING_HINT}) + nonce: Optional[int] = dc.field(default=None, metadata={ + 'tlv_type': TypeNumber.NONCE, 'fixed_len': 4}) + lifetime: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.INTEREST_LIFETIME}) + hop_limit: Optional[int] = dc.field(default=None, metadata={ + 'tlv_type': TypeNumber.HOP_LIMIT, 'fixed_len': 1}) + _sig_cover_start: None = dc.field( + default=None, metadata={'field_type': 'offset_marker'}) + _digest_cover_start: None = dc.field( + default=None, metadata={'field_type': 'offset_marker'}) + application_parameters: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.APPLICATION_PARAMETERS}) + signature_info: Optional[SignatureInfo] = dc.field( + default=None, metadata={ + 'tlv_type': TypeNumber.INTEREST_SIGNATURE_INFO}) + signature_value: Optional[bytes] = dc.field(default=None, metadata={ + 'tlv_type': TypeNumber.INTEREST_SIGNATURE_VALUE, + 'field_type': 'sig_value', + 'cover_start': '_sig_cover_start', + 'digest_cover_start': '_digest_cover_start', + 'digest_cover_end': '_digest_cover_end', + }) + _digest_cover_end: None = dc.field( + default=None, metadata={'field_type': 'offset_marker'}) + + +@dc.dataclass +class InterestPacket: + interest: Optional[InterestPacketValue] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.INTEREST}) + + +@dc.dataclass(init=False) +class MetaInfo: + content_type: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.CONTENT_TYPE}) + freshness_period: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.FRESHNESS_PERIOD}) + final_block_id: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.FINAL_BLOCK_ID}) + + def __init__(self, + content_type: Optional[int] = ContentType.BLOB, + freshness_period: Optional[int] = None, + final_block_id: Optional[BinaryStr] = None): + self.content_type = content_type + self.freshness_period = freshness_period + self.final_block_id = final_block_id + + @staticmethod + def from_dict(kwargs): + return MetaInfo(**{ + f.name: kwargs[f.name] + for f in dc.fields(MetaInfo) + if f.name in kwargs + }) + + +@dc.dataclass +class DataPacketValue: + _sig_cover_start: None = dc.field( + default=None, metadata={'field_type': 'offset_marker'}) + name: NDNName = dc.field( + default='/', metadata={'tlv_type': TypeNumber.NAME}) + meta_info: Optional[MetaInfo] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.META_INFO}) + content: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.CONTENT}) + signature_info: Optional[SignatureInfo] = dc.field(default=None, metadata={ + 'tlv_type': TypeNumber.SIGNATURE_INFO, 'ignore_critical': True}) + signature_value: Optional[bytes] = dc.field(default=None, metadata={ + 'tlv_type': TypeNumber.SIGNATURE_VALUE, + 'field_type': 'sig_value', + 'cover_start': '_sig_cover_start', + }) + + +@dc.dataclass +class DataPacket: + data: Optional[DataPacketValue] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.DATA}) + + +@dc.dataclass +class InterestParam: + can_be_prefix: bool = False + must_be_fresh: bool = False + nonce: Optional[int] = None + lifetime: Optional[int] = 4000 + hop_limit: Optional[int] = None + forwarding_hint: list[NonStrictName] = dc.field(default_factory=list) + + @staticmethod + def from_dict(kwargs): + return InterestParam(**{ + f.name: kwargs[f.name] + for f in dc.fields(InterestParam) + if f.name in kwargs + }) + + +@dc.dataclass +class SignaturePtrs: + signature_info: Optional[SignatureInfo] = None + signature_covered_part: list[BinaryStr] = dc.field(default_factory=list) + signature_value_buf: Optional[BinaryStr] = None + digest_covered_part: list[BinaryStr] = dc.field(default_factory=list) + digest_value_buf: Optional[BinaryStr] = None + + +Interest = tuple[FormalName, InterestParam, Optional[BinaryStr], SignaturePtrs] +Data = tuple[FormalName, MetaInfo, Optional[BinaryStr], SignaturePtrs] + + +def _wrap_tlv(type_num: int, value: BinaryStr) -> VarBinaryStr: + total = ( + get_tl_num_size(type_num) + + get_tl_num_size(len(value)) + + len(value) + ) + wire = bytearray(total) + offset = write_tl_num(type_num, wire, 0) + offset += write_tl_num(len(value), wire, offset) + wire[offset:] = value + return wire + + +def make_interest(name: NonStrictName, + interest_param: InterestParam, + app_param: Optional[BinaryStr] = None, + signer: Optional[Signer] = None, + need_final_name: bool = False): + value = InterestPacketValue( + name=name, + can_be_prefix=interest_param.can_be_prefix, + must_be_fresh=interest_param.must_be_fresh, + nonce=interest_param.nonce, + lifetime=interest_param.lifetime, + hop_limit=interest_param.hop_limit, + application_parameters=app_param, + ) + if interest_param.forwarding_hint: + value.forwarding_hint = Links( + names=list(interest_param.forwarding_hint)) + if signer is not None: + value.signature_info = SignatureInfo() + signer.write_signature_info(value.signature_info) + if value.application_parameters is None: + value.application_parameters = b'' + + markers = { + '##signer': signer, + '##need_digest': value.application_parameters is not None, + '##_digest_cover_start_field': '_digest_cover_start', + '##_digest_cover_end_field': '_digest_cover_end', + } + encoded_value = tlv_encode(value, markers=markers) + wire = _wrap_tlv(TypeNumber.INTEREST, encoded_value) + if need_final_name: + final_value = tlv_parse(InterestPacketValue, encoded_value) + return wire, final_value.name + return wire + + +def make_data(name: NonStrictName, + meta_info: MetaInfo, + content: Optional[BinaryStr] = None, + signer: Optional[Signer] = None) -> VarBinaryStr: + value = DataPacketValue(name=name, meta_info=meta_info, content=content) + if signer is not None: + value.signature_info = SignatureInfo() + signer.write_signature_info(value.signature_info) + encoded_value = tlv_encode(value, markers={'##signer': signer}) + return _wrap_tlv(TypeNumber.DATA, encoded_value) + + +def parse_interest(wire: BinaryStr, with_tl: bool = True) -> Interest: + value_wire = ( + parse_and_check_tl(wire, TypeNumber.INTEREST) + if with_tl else wire + ) + markers = {} + ret = tlv_parse(InterestPacketValue, value_wire, markers=markers) + params = InterestParam( + can_be_prefix=ret.can_be_prefix, + must_be_fresh=ret.must_be_fresh, + nonce=ret.nonce, + lifetime=ret.lifetime, + hop_limit=ret.hop_limit, + ) + if ret.forwarding_hint: + params.forwarding_hint.extend(ret.forwarding_hint.names) + + digest_parts = [] + digest_start = markers.get('_digest_cover_start') + if digest_start is not None: + digest_parts.append(memoryview(value_wire)[digest_start:]) + sig_ptrs = SignaturePtrs( + signature_info=ret.signature_info, + signature_covered_part=markers.get('##sig_covered_part', []), + signature_value_buf=ret.signature_value, + digest_covered_part=digest_parts, + digest_value_buf=markers.get('##digest_buf'), + ) + return ret.name, params, ret.application_parameters, sig_ptrs + + +def parse_data(wire: BinaryStr, with_tl: bool = True) -> Data: + value_wire = parse_and_check_tl(wire, TypeNumber.DATA) if with_tl else wire + markers = {} + ret = tlv_parse(DataPacketValue, value_wire, markers=markers) + meta_info = ret.meta_info if ret.meta_info is not None else MetaInfo() + sig_ptrs = SignaturePtrs( + signature_info=ret.signature_info, + signature_covered_part=markers.get('##sig_covered_part', []), + signature_value_buf=ret.signature_value, + ) + return ret.name, meta_info, ret.content, sig_ptrs diff --git a/tests/encoding/ndn_format_0_3_2_test.py b/tests/encoding/ndn_format_0_3_2_test.py new file mode 100644 index 0000000..87a186b --- /dev/null +++ b/tests/encoding/ndn_format_0_3_2_test.py @@ -0,0 +1,73 @@ +import hashlib + +from ndn.encoding import Name +from ndn.encoding.ndn_format_0_3_2 import ( + ContentType, + InterestParam, + MetaInfo, + SignatureType, + make_data, + make_interest, + parse_data, + parse_interest, +) +from ndn.security import DigestSha256Signer + + +def test_default_interest_wire_format(): + wire = make_interest('/local/ndn/prefix', InterestParam()) + assert wire == ( + b'\x05\x1a\x07\x14\x08\x05local\x08\x03ndn\x08\x06prefix' + b'\x0c\x02\x0f\xa0' + ) + + name, params, app_params, sig = parse_interest(wire) + assert name == Name.from_str('/local/ndn/prefix') + assert params.lifetime == 4000 + assert app_params is None + assert sig.signature_info is None + + +def test_signed_interest_wire_format_and_coverage(): + wire = make_interest( + '/local/ndn/prefix', + InterestParam(nonce=0x6c211166), + b'\x01\x02\x03\x04', + DigestSha256Signer(), + ) + assert wire == ( + b'\x05\x6f\x07\x36\x08\x05local\x08\x03ndn\x08\x06prefix' + b'\x02 \x8e\x6e\x36\xd7\xea\xbc\xde\x43\x75\x61\x40\xc9' + b'\x0b\xda\x09\xd5' + b'\x00\xd2\xa5\x77\xf2\xf5\x33\xb5\x69\xf0\x44\x1d\xf0\xa7\xf9\xe2' + b'\x0a\x04\x6c\x21\x11\x66\x0c\x02\x0f\xa0' + b'\x24\x04\x01\x02\x03\x04\x2c\x03\x1b\x01\x00' + b'\x2e \xea\xa8\xf0\x99\x08\x63\x78\x95\x1d\xe0\x5f\xf1' + b'\xde\xbb\xc1\x18' + b'\xb5\x21\x8b\x2f\xca\xa0\xb5\x1d\x18\xfa\xbc\x29\xf5\x4d\x58\xff' + ) + + _, _, _, sig = parse_interest(wire) + assert sig.signature_info.signature_type == SignatureType.DIGEST_SHA256 + signature = hashlib.sha256(b''.join(sig.signature_covered_part)).digest() + digest = hashlib.sha256(b''.join(sig.digest_covered_part)).digest() + assert signature == sig.signature_value_buf + assert digest == sig.digest_value_buf + + +def test_data_wire_format_and_coverage(): + wire = make_data( + '/local/ndn/prefix', MetaInfo(), signer=DigestSha256Signer()) + assert wire == ( + b"\x06\x42\x07\x14\x08\x05local\x08\x03ndn\x08\x06prefix" + b"\x14\x03\x18\x01\x00\x16\x03\x1b\x01\x00" + b"\x17 \x7f1\xe4\t\xc5z/\x1d\r\xdaVh8\xfd\xd9\x94" + b"\xd8\'S\x13[\xd7\x15\xa5\x9d%^\x80\xf2\xab\xf0\xb5" + ) + + name, meta_info, content, sig = parse_data(wire) + assert name == Name.from_str('/local/ndn/prefix') + assert meta_info.content_type == ContentType.BLOB + assert content is None + signature = hashlib.sha256(b''.join(sig.signature_covered_part)).digest() + assert signature == sig.signature_value_buf From 9cdee7e8512414a79420eb543670fffe47bc4304 Mon Sep 17 00:00:00 2001 From: Xinyu Ma Date: Sun, 6 Sep 2026 17:59:01 -0700 Subject: [PATCH 3/8] encoding: Add dataclass NDNLPv2 model Co-authored-by: Cursor --- src/ndn/encoding/ndnlp_v2_2.py | 133 ++++++++++++++++++++++++++++++ tests/encoding/ndnlp_v2_2_test.py | 58 +++++++++++++ 2 files changed, 191 insertions(+) create mode 100644 src/ndn/encoding/ndnlp_v2_2.py create mode 100644 tests/encoding/ndnlp_v2_2_test.py diff --git a/src/ndn/encoding/ndnlp_v2_2.py b/src/ndn/encoding/ndnlp_v2_2.py new file mode 100644 index 0000000..9939c00 --- /dev/null +++ b/src/ndn/encoding/ndnlp_v2_2.py @@ -0,0 +1,133 @@ +# ----------------------------------------------------------------------------- +# Copyright (C) 2019-2020 The python-ndn authors +# Licensed under the Apache License, Version 2.0 (the "License"); +# ----------------------------------------------------------------------------- +"""NDNLPv2 models using the dataclass TLV API.""" +import dataclasses as dc +from typing import Optional + +from .tlv_model import DecodeError +from .tlv_model_v2 import tlv_encode, tlv_parse +from .tlv_type import BinaryStr, VarBinaryStr +from .tlv_var import parse_and_check_tl + +__all__ = [ + 'LpTypeNumber', 'NackReason', 'NetworkNack', 'CachePolicy', + 'LpPacketValue', 'LpPacket', 'parse_network_nack', 'make_network_nack', + 'parse_lp_packet', 'parse_lp_packet_v2', +] + + +class LpTypeNumber: + FRAGMENT = 0x50 + SEQUENCE = 0x51 + FRAG_INDEX = 0x52 + FRAG_COUNT = 0x53 + HOP_COUNT = 0x54 + PIT_TOKEN = 0x62 + LP_PACKET = 0x64 + NACK = 0x0320 + NACK_REASON = 0x0321 + INCOMING_FACE_ID = 0x032C + NEXT_HOP_FACE_ID = 0x0330 + CACHE_POLICY = 0x0334 + CACHE_POLICY_TYPE = 0x0335 + CONGESTION_MARK = 0x0340 + ACK = 0x0344 + TX_SEQUENCE = 0x0348 + NON_DISCOVERY = 0x034C + PREFIX_ANNOUNCEMENT = 0x0350 + + +class NackReason: + NONE = 0 + CONGESTION = 50 + DUPLICATE = 100 + NO_ROUTE = 150 + + +@dc.dataclass +class NetworkNack: + nack_reason: Optional[int] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.NACK_REASON}) + + +@dc.dataclass +class CachePolicy: + cache_policy_type: Optional[int] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.CACHE_POLICY_TYPE}) + + +@dc.dataclass +class LpPacketValue: + frag_index: Optional[int] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.FRAG_INDEX}) + frag_count: Optional[int] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.FRAG_COUNT}) + pit_token: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.PIT_TOKEN}) + nack: Optional[NetworkNack] = dc.field( + default=None, metadata={ + 'tlv_type': LpTypeNumber.NACK, 'ignore_critical': False}) + incoming_face_id: Optional[int] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.INCOMING_FACE_ID}) + next_hop_face_id: Optional[int] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.NEXT_HOP_FACE_ID}) + cache_policy: Optional[CachePolicy] = dc.field( + default=None, metadata={ + 'tlv_type': LpTypeNumber.CACHE_POLICY, 'ignore_critical': False}) + congestion_mark: Optional[int] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.CONGESTION_MARK}) + tx_sequence: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.TX_SEQUENCE}) + ack: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.ACK}) + non_discovery: bool = dc.field( + default=False, metadata={'tlv_type': LpTypeNumber.NON_DISCOVERY}) + prefix_announcement: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.PREFIX_ANNOUNCEMENT}) + fragment: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.FRAGMENT}) + + +@dc.dataclass +class LpPacket: + lp_packet: Optional[LpPacketValue] = dc.field( + default=None, metadata={'tlv_type': LpTypeNumber.LP_PACKET}) + + +def parse_lp_packet(wire: BinaryStr, + with_tl: bool = True + ) -> tuple[Optional[int], Optional[BinaryStr]]: + ret = parse_lp_packet_v2(wire, with_tl) + reason = ret.nack.nack_reason if ret.nack is not None else None + return reason, ret.fragment + + +def parse_lp_packet_v2(wire: BinaryStr, with_tl: bool = True) -> LpPacketValue: + if with_tl: + wire = parse_and_check_tl(wire, LpTypeNumber.LP_PACKET) + ret = tlv_parse(LpPacketValue, wire, ignore_critical=True) + if ret.frag_index is not None or ret.frag_count is not None: + raise DecodeError('NDNLP fragmentation is not implemented yet.') + return ret + + +def parse_network_nack( + wire: BinaryStr, + with_tl: bool = True) -> tuple[Optional[int], Optional[BinaryStr]]: + if with_tl: + wire = parse_and_check_tl(wire, LpTypeNumber.LP_PACKET) + ret = tlv_parse(LpPacketValue, wire, ignore_critical=True) + if ret.nack is not None: + return ret.nack.nack_reason, ret.fragment + return None, None + + +def make_network_nack(encoded_interest: BinaryStr, + nack_reason: int) -> VarBinaryStr: + value = LpPacketValue( + nack=NetworkNack(nack_reason=nack_reason), + fragment=encoded_interest, + ) + return tlv_encode(LpPacket(lp_packet=value)) diff --git a/tests/encoding/ndnlp_v2_2_test.py b/tests/encoding/ndnlp_v2_2_test.py new file mode 100644 index 0000000..8c6c188 --- /dev/null +++ b/tests/encoding/ndnlp_v2_2_test.py @@ -0,0 +1,58 @@ +from ndn.encoding.ndnlp_v2_2 import ( + LpPacketValue, + LpTypeNumber, + NackReason, + NetworkNack, + make_network_nack, + parse_network_nack, + parse_lp_packet_v2, +) +from ndn.encoding.ndn_format_0_3_2 import ( + InterestParam, + make_interest, + parse_interest, +) +from ndn.encoding import DecodeError, Name, tlv_encode, write_tl_num +import pytest + + +def test_network_nack_wire_format(): + interest = make_interest( + '/localhost/nfd/faces/events', + InterestParam(must_be_fresh=True, can_be_prefix=True), + ) + lp_packet = make_network_nack(interest, NackReason.NO_ROUTE) + + assert lp_packet == ( + b"\x64\x36\xfd\x03\x20\x05\xfd\x03\x21\x01\x96" + b"\x50\x2b\x05\x29\x07\x1f\x08\tlocalhost\x08\x03nfd" + b"\x08\x05faces\x08\x06events\x21\x00\x12\x00\x0c\x02\x0f\xa0" + ) + + reason, encoded_interest = parse_network_nack(lp_packet) + name, params, _, _ = parse_interest(encoded_interest) + assert reason == NackReason.NO_ROUTE + assert name == Name.from_str('/localhost/nfd/faces/events') + assert params.can_be_prefix + assert params.must_be_fresh + + +def test_network_nack_parser_accepts_fragment_metadata(): + value = tlv_encode(LpPacketValue( + frag_index=0, + frag_count=1, + nack=NetworkNack(nack_reason=NackReason.NO_ROUTE), + fragment=b'\x05\x00', + )) + wire = bytearray(2 + len(value)) + offset = write_tl_num(LpTypeNumber.LP_PACKET, wire, 0) + offset += write_tl_num(len(value), wire, offset) + wire[offset:] = value + + assert parse_network_nack(wire) == (NackReason.NO_ROUTE, b'\x05\x00') + + +def test_nested_unknown_critical_field_is_rejected(): + wire = b'\x64\x06\xfd\x03\x20\x02\x01\x00' + with pytest.raises(DecodeError): + parse_lp_packet_v2(wire) From 4cc2e89fe1b820a96f4e46435e753dfdb108a516 Mon Sep 17 00:00:00 2001 From: Xinyu Ma Date: Sun, 6 Sep 2026 17:59:01 -0700 Subject: [PATCH 4/8] encoding: Add dataclass security model Co-authored-by: Cursor --- src/ndn/app_support/security_v2_2.py | 187 +++++++++++++++++++++++++++ tests/misc/security_v2_2_test.py | 83 ++++++++++++ 2 files changed, 270 insertions(+) create mode 100644 src/ndn/app_support/security_v2_2.py create mode 100644 tests/misc/security_v2_2_test.py diff --git a/src/ndn/app_support/security_v2_2.py b/src/ndn/app_support/security_v2_2.py new file mode 100644 index 0000000..4b585b5 --- /dev/null +++ b/src/ndn/app_support/security_v2_2.py @@ -0,0 +1,187 @@ +# ----------------------------------------------------------------------------- +# Copyright (C) 2019-2020 The python-ndn authors +# +# This file is part of python-ndn. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ----------------------------------------------------------------------------- +import dataclasses as dc +from datetime import datetime, timedelta, UTC +from typing import Optional + +from ..utils import timestamp +from ..encoding import ( + Component, + FormalName, + Name, + VarBinaryStr, + parse_and_check_tl, +) +from ..encoding.tlv_model_v2 import tlv_encode, tlv_parse +from ..encoding.ndn_format_0_3_2 import ( + ContentType, + DataPacketValue, + KeyLocator, + MetaInfo, + SignatureInfo, + TypeNumber, +) + + +KEY_COMPONENT = Component.from_str('KEY') +SELF_COMPONENT = Component.from_str('self') +SIGN_REQ_COMPONENT = Component.from_str('cert-request') + + +class SecurityV2TypeNumber: + VALIDITY_PERIOD = 0xFD + NOT_BEFORE = 0xFE + NOT_AFTER = 0xFF + ADDITIONAL_DESCRIPTION = 0x0102 + DESCRIPTION_ENTRY = 0x0200 + DESCRIPTION_KEY = 0x0201 + DESCRIPTION_VALUE = 0x0202 + + SAFE_BAG = 0x80 + ENCRYPTED_KEY_BAG = 0x81 + + +@dc.dataclass +class DescriptionEntry: + description_key: Optional[bytes] = dc.field( + default=None, metadata={ + 'tlv_type': SecurityV2TypeNumber.DESCRIPTION_KEY}) + description_value: Optional[bytes] = dc.field( + default=None, metadata={ + 'tlv_type': SecurityV2TypeNumber.DESCRIPTION_VALUE}) + + +@dc.dataclass +class AdditionalDescription: + description_entry: list[DescriptionEntry] = dc.field( + default_factory=list, metadata={ + 'tlv_type': SecurityV2TypeNumber.DESCRIPTION_ENTRY}) + + +@dc.dataclass +class CertificateV2Extension: + additional_description: Optional[AdditionalDescription] = dc.field( + default=None, metadata={ + 'tlv_type': SecurityV2TypeNumber.ADDITIONAL_DESCRIPTION}) + + +@dc.dataclass +class ValidityPeriod: + not_before: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': SecurityV2TypeNumber.NOT_BEFORE}) + not_after: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': SecurityV2TypeNumber.NOT_AFTER}) + + +@dc.dataclass +class CertificateV2SignatureInfo(SignatureInfo): + validity_period: Optional[ValidityPeriod] = dc.field( + default=None, metadata={ + 'tlv_type': SecurityV2TypeNumber.VALIDITY_PERIOD}) + additional_description: Optional[AdditionalDescription] = dc.field( + default=None, metadata={ + 'tlv_type': SecurityV2TypeNumber.ADDITIONAL_DESCRIPTION}) + + +@dc.dataclass +class CertificateV2Value(DataPacketValue): + signature_info: Optional[CertificateV2SignatureInfo] = dc.field( + default=None, metadata={ + 'tlv_type': TypeNumber.SIGNATURE_INFO, + 'ignore_critical': True, + }) + + +@dc.dataclass +class SafeBag: + certificate_v2: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.DATA}) + # We do not use ModelField due to 2 reasons: + # 1. The encoded length of CertificateV2 is unknown. + # 2. Generally we already have an encoded certificate when exporting a + # SafeBag. + encrypted_key_bag: Optional[bytes] = dc.field( + default=None, metadata={ + 'tlv_type': SecurityV2TypeNumber.ENCRYPTED_KEY_BAG}) + + +@dc.dataclass +class _CertificateEnvelope: + value: bytes = dc.field(metadata={'tlv_type': TypeNumber.DATA}) + + +def parse_certificate(wire) -> CertificateV2Value: + wire = parse_and_check_tl(wire, TypeNumber.DATA) + return tlv_parse(CertificateV2Value, wire) + + +def new_cert(key_name, issuer_id_component, pub_key, signer, + start_time, end_time) -> tuple[FormalName, VarBinaryStr]: + cert_name = Name.normalize(key_name) + [ + issuer_id_component, + Component.from_version(timestamp()), + ] + not_before = start_time.strftime('%Y%m%dT%H%M%S').encode() + not_after = end_time.strftime('%Y%m%dT%H%M%S').encode() + signature_info = CertificateV2SignatureInfo( + validity_period=ValidityPeriod( + not_before=not_before, + not_after=not_after, + ), + ) + signer.write_signature_info(signature_info) + if (signature_info.key_locator is not None + and not isinstance(signature_info.key_locator, KeyLocator)): + old_key_locator = signature_info.key_locator + signature_info.key_locator = KeyLocator( + name=old_key_locator.name, + key_digest=old_key_locator.key_digest, + ) + cert_val = CertificateV2Value( + name=cert_name, + content=pub_key, + meta_info=MetaInfo( + content_type=ContentType.KEY, + freshness_period=3600000, + ), + signature_info=signature_info, + ) + value = tlv_encode(cert_val, markers={'##signer': signer}) + return cert_name, tlv_encode(_CertificateEnvelope(value=value)) + + +def self_sign(key_name, pub_key, signer) -> tuple[FormalName, VarBinaryStr]: + end_time = datetime.now(UTC) + end_time = end_time.replace(year=end_time.year + 20) + return new_cert(key_name, SELF_COMPONENT, pub_key, signer, + datetime.fromisoformat('1970-01-01T00:00:00'), end_time) + + +def sign_req(key_name, pub_key, signer) -> tuple[FormalName, VarBinaryStr]: + start_time = datetime.now(UTC) + end_time = start_time + timedelta(days=10) + return new_cert(key_name, SIGN_REQ_COMPONENT, pub_key, signer, + datetime.now(UTC), end_time) + + +def derive_cert(key_name, issuer_id, pub_key, signer, + start_time, expire_sec) -> tuple[FormalName, VarBinaryStr]: + end_time = start_time + timedelta(seconds=expire_sec) + if isinstance(issuer_id, str): + issuer_id = Component.from_str(issuer_id) + return new_cert(key_name, issuer_id, pub_key, signer, start_time, end_time) diff --git a/tests/misc/security_v2_2_test.py b/tests/misc/security_v2_2_test.py new file mode 100644 index 0000000..bc22220 --- /dev/null +++ b/tests/misc/security_v2_2_test.py @@ -0,0 +1,83 @@ +import dataclasses as dc +import hashlib +from datetime import UTC, datetime + +from ndn.app_support.security_v2_2 import ( + CertificateV2SignatureInfo, + CertificateV2Value, + ContentType, + SafeBag, + SecurityV2TypeNumber, + new_cert, + parse_certificate, +) +from ndn.encoding import ( + Component, + Name, + SignatureType, + tlv_encode, + tlv_parse, +) +from ndn.encoding.ndn_format_0_3_2 import parse_data +from ndn.security import DigestSha256Signer, HmacSha256Signer + + +def test_certificate_models_use_dataclass_tlv_format(): + assert dc.is_dataclass(CertificateV2SignatureInfo) + assert dc.is_dataclass(CertificateV2Value) + assert dc.is_dataclass(SafeBag) + + +def test_new_cert_round_trip(): + start = datetime(2025, 1, 2, 3, 4, 5, tzinfo=UTC) + end = datetime(2026, 2, 3, 4, 5, 6, tzinfo=UTC) + cert_name, wire = new_cert( + '/test/KEY/key-id', + Component.from_str('issuer'), + b'public-key', + DigestSha256Signer(), + start, + end, + ) + + cert = parse_certificate(wire) + assert cert.name == cert_name + assert cert.content == b'public-key' + assert cert.meta_info.content_type == ContentType.KEY + assert cert.meta_info.freshness_period == 3600000 + assert cert.signature_info.signature_type == SignatureType.DIGEST_SHA256 + assert cert.signature_info.validity_period.not_before == b'20250102T030405' + assert cert.signature_info.validity_period.not_after == b'20260203T040506' + assert Name.is_prefix(Name.from_str('/test/KEY/key-id'), cert_name) + + _, _, _, sig = parse_data(wire) + covered = b''.join(sig.signature_covered_part) + assert hashlib.sha256(covered).digest() == sig.signature_value_buf + + +def test_new_cert_converts_legacy_signer_key_locator(): + _, wire = new_cert( + '/test/KEY/key-id', + Component.from_str('issuer'), + b'public-key', + HmacSha256Signer('/signer/key', b'secret'), + datetime(2025, 1, 1, tzinfo=UTC), + datetime(2026, 1, 1, tzinfo=UTC), + ) + + cert = parse_certificate(wire) + assert cert.signature_info.key_locator.name == Name.from_str('/signer/key') + + +def test_safe_bag_round_trip(): + safe_bag = SafeBag(certificate_v2=b'\x06\x00', encrypted_key_bag=b'key') + wire = tlv_encode(safe_bag) + + assert wire == ( + bytes([0x06, 0x02, 0x06, 0x00]) + + bytes([SecurityV2TypeNumber.ENCRYPTED_KEY_BAG, 0x03]) + + b'key' + ) + parsed = tlv_parse(SafeBag, wire) + assert parsed.certificate_v2 == b'\x06\x00' + assert parsed.encrypted_key_bag == b'key' From 96d14979c989ef75d712cf69fe0e0749a7d60727 Mon Sep 17 00:00:00 2001 From: Xinyu Ma Date: Sun, 6 Sep 2026 18:31:57 -0700 Subject: [PATCH 5/8] fix: Fix compilation error in pypy3 --- src/ndn/encoding/tlv_model_v2.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/src/ndn/encoding/tlv_model_v2.py b/src/ndn/encoding/tlv_model_v2.py index f17dea4..155fd20 100644 --- a/src/ndn/encoding/tlv_model_v2.py +++ b/src/ndn/encoding/tlv_model_v2.py @@ -100,6 +100,7 @@ class Outer: import typing from enum import Enum, Flag from hashlib import sha256 +from types import UnionType from .tlv_type import VarBinaryStr, is_binary_str from .tlv_var import write_tl_num, parse_tl_num, get_tl_num_size @@ -143,7 +144,7 @@ class NDNName: def _unwrap_optional(annotation): """Return T for Optional[T] = Union[T, None]; otherwise return unchanged.""" - if typing.get_origin(annotation) is typing.Union: + if typing.get_origin(annotation) in (typing.Union, UnionType): args = [a for a in typing.get_args(annotation) if a is not type(None)] if len(args) == 1: return args[0] From 4df603d0e61fa77d4d1ed42ee6af097410583ff2 Mon Sep 17 00:00:00 2001 From: Xinyu Ma Date: Sat, 26 Sep 2026 23:00:56 -0700 Subject: [PATCH 6/8] feat: Convert app_support to new model and improve performance --- src/ndn/app_support/light_versec/binary.py | 146 ++-- src/ndn/app_support/light_versec/checker.py | 2 +- src/ndn/app_support/nfd_mgmt_2.py | 303 ++++++++ src/ndn/app_support/security_v2_2.py | 11 +- src/ndn/appv2_2.py | 754 ++++++++++++++++++++ src/ndn/encoding/ndn_format_0_3_2.py | 20 +- src/ndn/encoding/tlv_model_v2.py | 198 ++--- src/ndn/transport/nfd_registerer.py | 13 +- tests/encoding/ndn_format_0_3_2_test.py | 15 +- tests/encoding/tlv_model_v2_test.py | 71 ++ tests/integration/app_v2_2_test.py | 295 ++++++++ tests/misc/nfd_mgmt_2_test.py | 83 +++ 12 files changed, 1751 insertions(+), 160 deletions(-) create mode 100644 src/ndn/app_support/nfd_mgmt_2.py create mode 100644 src/ndn/appv2_2.py create mode 100644 tests/integration/app_v2_2_test.py create mode 100644 tests/misc/nfd_mgmt_2_test.py diff --git a/src/ndn/app_support/light_versec/binary.py b/src/ndn/app_support/light_versec/binary.py index 8f88df8..103948d 100644 --- a/src/ndn/app_support/light_versec/binary.py +++ b/src/ndn/app_support/light_versec/binary.py @@ -20,7 +20,11 @@ # See the License for the specific language governing permissions and # limitations under the License. # ----------------------------------------------------------------------------- -import ndn.encoding as enc +import dataclasses as dc +from typing import Optional + +from ...encoding import BinaryStr +from ...encoding.tlv_model_v2 import tlv_encode, tlv_parse __all__ = [ @@ -62,63 +66,101 @@ class TypeNumber: NAMED_PATTERN_NUM = 0x69 -class UserFnArg(enc.TlvModel): +@dc.dataclass +class UserFnArg: # A given component - value = enc.BytesField(TypeNumber.COMPONENT_VALUE) + value: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.COMPONENT_VALUE}) # Referring to a previous matched pattern - tag = enc.UintField(TypeNumber.PATTERN_TAG) + tag: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.PATTERN_TAG}) -class UserFnCall(enc.TlvModel): - fn_id = enc.BytesField(TypeNumber.USER_FN_ID, is_string=True) - args = enc.RepeatedField(enc.ModelField(TypeNumber.FN_ARGS, UserFnArg)) +@dc.dataclass +class UserFnCall: + fn_id: Optional[str] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.USER_FN_ID}) + args: list[UserFnArg] = dc.field( + default_factory=list, metadata={'tlv_type': TypeNumber.FN_ARGS}) -class ConstraintOption(enc.TlvModel): +@dc.dataclass +class ConstraintOption: # Equal to a given NameComponent value - value = enc.BytesField(TypeNumber.COMPONENT_VALUE) + value: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.COMPONENT_VALUE}) # Equal to another pattern - tag = enc.UintField(TypeNumber.PATTERN_TAG) + tag: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.PATTERN_TAG}) # Decide by a user function call - fn = enc.ModelField(TypeNumber.USER_FN_CALL, UserFnCall) - - -class PatternConstraint(enc.TlvModel): - options = enc.RepeatedField( - enc.ModelField(TypeNumber.CONS_OPTION, ConstraintOption) - ) - - -class PatternEdge(enc.TlvModel): - dest = enc.UintField(TypeNumber.NODE_ID) - tag = enc.UintField(TypeNumber.PATTERN_TAG) - cons_sets = enc.RepeatedField( - enc.ModelField(TypeNumber.CONSTRAINT, PatternConstraint) - ) - - -class ValueEdge(enc.TlvModel): - dest = enc.UintField(TypeNumber.NODE_ID) - value = enc.BytesField(TypeNumber.COMPONENT_VALUE) - - -class Node(enc.TlvModel): - id = enc.UintField(TypeNumber.NODE_ID) - parent = enc.UintField(TypeNumber.PARENT_ID) - rule_name = enc.RepeatedField(enc.BytesField(TypeNumber.IDENTIFIER, is_string=True)) - v_edges = enc.RepeatedField(enc.ModelField(TypeNumber.VALUE_EDGE, ValueEdge)) - p_edges = enc.RepeatedField(enc.ModelField(TypeNumber.PATTERN_EDGE, PatternEdge)) - sign_cons = enc.RepeatedField(enc.UintField(TypeNumber.KEY_NODE_ID)) - - -class TagSymbol(enc.TlvModel): - tag = enc.UintField(TypeNumber.PATTERN_TAG) - ident = enc.BytesField(TypeNumber.IDENTIFIER, is_string=True) - - -class LvsModel(enc.TlvModel): - version = enc.UintField(TypeNumber.VERSION) - start_id = enc.UintField(TypeNumber.NODE_ID) - named_pattern_cnt = enc.UintField(TypeNumber.NAMED_PATTERN_NUM) - nodes = enc.RepeatedField(enc.ModelField(TypeNumber.NODE, Node)) - symbols = enc.RepeatedField(enc.ModelField(TypeNumber.TAG_SYMBOL, TagSymbol)) + fn: Optional[UserFnCall] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.USER_FN_CALL}) + + +@dc.dataclass +class PatternConstraint: + options: list[ConstraintOption] = dc.field( + default_factory=list, metadata={'tlv_type': TypeNumber.CONS_OPTION}) + + +@dc.dataclass +class PatternEdge: + dest: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.NODE_ID}) + tag: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.PATTERN_TAG}) + cons_sets: list[PatternConstraint] = dc.field( + default_factory=list, metadata={'tlv_type': TypeNumber.CONSTRAINT}) + + +@dc.dataclass +class ValueEdge: + dest: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.NODE_ID}) + value: Optional[bytes] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.COMPONENT_VALUE}) + + +@dc.dataclass +class Node: + id: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.NODE_ID}) + parent: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.PARENT_ID}) + rule_name: list[str] = dc.field( + default_factory=list, metadata={'tlv_type': TypeNumber.IDENTIFIER}) + v_edges: list[ValueEdge] = dc.field( + default_factory=list, metadata={'tlv_type': TypeNumber.VALUE_EDGE}) + p_edges: list[PatternEdge] = dc.field( + default_factory=list, metadata={'tlv_type': TypeNumber.PATTERN_EDGE}) + sign_cons: list[int] = dc.field( + default_factory=list, metadata={'tlv_type': TypeNumber.KEY_NODE_ID}) + + +@dc.dataclass +class TagSymbol: + tag: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.PATTERN_TAG}) + ident: Optional[str] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.IDENTIFIER}) + + +@dc.dataclass +class LvsModel: + version: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.VERSION}) + start_id: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.NODE_ID}) + named_pattern_cnt: Optional[int] = dc.field( + default=None, metadata={'tlv_type': TypeNumber.NAMED_PATTERN_NUM}) + nodes: list[Node] = dc.field( + default_factory=list, metadata={'tlv_type': TypeNumber.NODE}) + symbols: list[TagSymbol] = dc.field( + default_factory=list, metadata={'tlv_type': TypeNumber.TAG_SYMBOL}) + + def encode(self) -> bytearray: + return tlv_encode(self) + + @classmethod + def parse(cls, wire: BinaryStr) -> 'LvsModel': + return tlv_parse(cls, wire) diff --git a/src/ndn/app_support/light_versec/checker.py b/src/ndn/app_support/light_versec/checker.py index 02c32d7..190b6c6 100644 --- a/src/ndn/app_support/light_versec/checker.py +++ b/src/ndn/app_support/light_versec/checker.py @@ -26,7 +26,7 @@ from ...encoding import BinaryStr, Component, FormalName, Name, NonStrictName from ...security import Keychain -from ..security_v2 import parse_certificate +from ..security_v2_2 import parse_certificate from . import binary as bny from .compiler import top_order diff --git a/src/ndn/app_support/nfd_mgmt_2.py b/src/ndn/app_support/nfd_mgmt_2.py new file mode 100644 index 0000000..8d0383c --- /dev/null +++ b/src/ndn/app_support/nfd_mgmt_2.py @@ -0,0 +1,303 @@ +# ----------------------------------------------------------------------------- +# Copyright (C) 2019-2020 The python-ndn authors +# +# This file is part of python-ndn. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ----------------------------------------------------------------------------- +"""NFD management protocol models using the dataclass TLV API.""" +import dataclasses as dc +import struct +from typing import Optional + +from ..transport.face import Face +from ..utils import timestamp, gen_nonce_64 +from ..encoding import Component, Name, get_tl_num_size, write_tl_num, parse_and_check_tl +from ..encoding.tlv_model_v2 import NDNName, tlv_encode, tlv_parse +from ..encoding.ndn_format_0_3_2 import SignatureInfo, TypeNumber, write_signature_info +from ..security import DigestSha256Signer +from .nfd_mgmt import ( + FaceScope, FacePersistency, FaceLinkType, FaceFlags, RouteFlags, FaceEventKind, +) + +__all__ = [ + 'FaceScope', 'FacePersistency', 'FaceLinkType', 'FaceFlags', 'RouteFlags', 'FaceEventKind', + 'Strategy', 'ControlParametersValue', 'ControlParameters', 'ControlResponse', + 'FaceEventNotificationValue', 'FaceEventNotification', 'GeneralStatus', 'FaceStatus', + 'FaceStatusMsg', 'FaceQueryFilterValue', 'FaceQueryFilter', 'Route', 'RibEntry', 'RibStatus', + 'NextHopRecord', 'FibEntry', 'FibStatus', 'StrategyChoice', 'StrategyChoiceMsg', 'CsInfo', + 'make_command', 'make_command_v2', 'parse_response', +] + + +def _tlv(type_num: int): + return dc.field(default=None, metadata={'tlv_type': type_num}) + + +def _name(): + return _tlv(TypeNumber.NAME) + + +def _repeated(type_num: int): + return dc.field(default_factory=list, metadata={'tlv_type': type_num}) + + +@dc.dataclass +class Strategy: + name: NDNName = _name() + + +@dc.dataclass +class ControlParametersValue: + name: NDNName = _name() + face_id: Optional[int] = _tlv(0x69) + uri: Optional[str] = _tlv(0x72) + local_uri: Optional[str] = _tlv(0x81) + origin: Optional[int] = _tlv(0x6f) + cost: Optional[int] = _tlv(0x6a) + capacity: Optional[int] = _tlv(0x83) + count: Optional[int] = _tlv(0x84) + base_congestion_mark_interval: Optional[int] = _tlv(0x87) + default_congestion_threshold: Optional[int] = _tlv(0x88) + mtu: Optional[int] = _tlv(0x89) + flags: Optional[int] = _tlv(0x6c) + mask: Optional[int] = _tlv(0x70) + strategy: Optional[Strategy] = _tlv(0x6b) + expiration_period: Optional[int] = _tlv(0x6d) + face_persistency: Optional[FacePersistency] = _tlv(0x85) + + +@dc.dataclass +class ControlParameters: + cp: Optional[ControlParametersValue] = _tlv(0x68) + + +@dc.dataclass +class ControlResponse: + status_code: Optional[int] = _tlv(0x66) + status_text: Optional[str] = _tlv(0x67) + body: Optional[ControlParametersValue] = _tlv(0x68) + + +@dc.dataclass +class FaceEventNotificationValue: + face_event_kind: Optional[FaceEventKind] = _tlv(0xc1) + face_id: Optional[int] = _tlv(0x69) + uri: Optional[str] = _tlv(0x72) + local_uri: Optional[str] = _tlv(0x81) + face_scope: Optional[FaceScope] = _tlv(0x84) + face_persistency: Optional[FacePersistency] = _tlv(0x85) + link_type: Optional[FaceLinkType] = _tlv(0x86) + flags: Optional[FaceFlags] = _tlv(0x6c) + + +@dc.dataclass +class FaceEventNotification: + event: Optional[FaceEventNotificationValue] = _tlv(0xc0) + + +@dc.dataclass +class GeneralStatus: + nfd_version: Optional[str] = _tlv(0x80) + start_timestamp: Optional[int] = _tlv(0x81) + current_timestamp: Optional[int] = _tlv(0x82) + n_name_tree_entries: Optional[int] = _tlv(0x83) + n_fib_entries: Optional[int] = _tlv(0x84) + n_pit_entries: Optional[int] = _tlv(0x85) + n_measurement_entries: Optional[int] = _tlv(0x86) + n_cs_entries: Optional[int] = _tlv(0x87) + n_in_interests: Optional[int] = _tlv(0x90) + n_in_data: Optional[int] = _tlv(0x91) + n_in_nacks: Optional[int] = _tlv(0x97) + n_out_interests: Optional[int] = _tlv(0x92) + n_out_data: Optional[int] = _tlv(0x93) + n_out_nacks: Optional[int] = _tlv(0x98) + n_satisfied_interests: Optional[int] = _tlv(0x99) + n_unsatisfied_interests: Optional[int] = _tlv(0x9a) + # The following comes from DNMP's extension to NFD mgmt protocol: + # https://github.com/pollere/DNMP-v2/blob/c4359ae1af03824ec1ee8cd27a7d52c9151fa813/formats/forwarder-status.proto + # It does not show up in the standard protocol: + # https://redmine.named-data.net/projects/nfd/wiki/ForwarderStatus + n_fragmentation_errors: Optional[int] = _tlv(0xc8) + n_out_over_mtu: Optional[int] = _tlv(0xc9) + n_in_lp_invalid: Optional[int] = _tlv(0xca) + n_reassembly_timeouts: Optional[int] = _tlv(0xcb) + n_in_net_invalid: Optional[int] = _tlv(0xcc) + n_acknowledged: Optional[int] = _tlv(0xcd) + n_retransmitted: Optional[int] = _tlv(0xce) + n_retx_exhausted: Optional[int] = _tlv(0xcf) + n_congestion_marked: Optional[int] = _tlv(0xd0) + + +@dc.dataclass +class FaceStatus: + face_id: Optional[int] = _tlv(0x69) + uri: Optional[str] = _tlv(0x72) + local_uri: Optional[str] = _tlv(0x81) + expiration_period: Optional[int] = _tlv(0x6d) + face_scope: Optional[FaceScope] = _tlv(0x84) + face_persistency: Optional[FacePersistency] = _tlv(0x85) + link_type: Optional[FaceLinkType] = _tlv(0x86) + base_congestion_mark_interval: Optional[int] = _tlv(0x87) + default_congestion_threshold: Optional[int] = _tlv(0x88) + mtu: Optional[int] = _tlv(0x89) + n_in_interests: Optional[int] = _tlv(0x90) + n_in_data: Optional[int] = _tlv(0x91) + n_in_nacks: Optional[int] = _tlv(0x97) + n_out_interests: Optional[int] = _tlv(0x92) + n_out_data: Optional[int] = _tlv(0x93) + n_out_nacks: Optional[int] = _tlv(0x98) + n_in_bytes: Optional[int] = _tlv(0x94) + n_out_bytes: Optional[int] = _tlv(0x95) + flags: Optional[FaceFlags] = _tlv(0x6c) + + +@dc.dataclass +class FaceStatusMsg: + face_status: list[FaceStatus] = _repeated(0x80) + + +@dc.dataclass +class FaceQueryFilterValue: + face_id: Optional[int] = _tlv(0x69) + uri_scheme: Optional[str] = _tlv(0x83) + uri: Optional[str] = _tlv(0x72) + local_uri: Optional[str] = _tlv(0x81) + face_scope: Optional[FaceScope] = _tlv(0x84) + face_persistency: Optional[FacePersistency] = _tlv(0x85) + link_type: Optional[FaceLinkType] = _tlv(0x86) + + +@dc.dataclass +class FaceQueryFilter: + face_query_filter: Optional[FaceQueryFilterValue] = _tlv(0x96) + + +@dc.dataclass +class Route: + face_id: Optional[int] = _tlv(0x69) + origin: Optional[int] = _tlv(0x6f) + cost: Optional[int] = _tlv(0x6a) + flags: Optional[RouteFlags] = _tlv(0x6c) + expiration_period: Optional[int] = _tlv(0x6d) + + +@dc.dataclass +class RibEntry: + name: NDNName = _name() + routes: list[Route] = _repeated(0x81) + + +@dc.dataclass +class RibStatus: + entries: list[RibEntry] = _repeated(0x80) + + +@dc.dataclass +class NextHopRecord: + face_id: Optional[int] = _tlv(0x69) + cost: Optional[int] = _tlv(0x6a) + + +@dc.dataclass +class FibEntry: + name: NDNName = _name() + next_hop_records: list[NextHopRecord] = _repeated(0x81) + + +@dc.dataclass +class FibStatus: + entries: list[FibEntry] = _repeated(0x80) + + +@dc.dataclass +class StrategyChoice: + name: NDNName = _name() + strategy: Optional[Strategy] = _tlv(0x6b) + + +@dc.dataclass +class StrategyChoiceMsg: + strategy_choices: list[StrategyChoice] = _repeated(0x80) + + +@dc.dataclass +class CsInfo: + capacity: Optional[int] = _tlv(0x83) + flags: Optional[int] = _tlv(0x6c) + n_cs_entries: Optional[int] = _tlv(0x87) + n_hits: Optional[int] = _tlv(0x81) + n_misses: Optional[int] = _tlv(0x82) + + +def make_command(module, command, face: Face | None = None, **kwargs): + ret = make_command_v2(module, command, face, **kwargs) + + # Timestamp and nonce + ret.append(Component.from_bytes(struct.pack('!Q', timestamp()))) + ret.append(Component.from_bytes(struct.pack('!Q', gen_nonce_64()))) + + # SignatureInfo + signer = DigestSha256Signer() + sig_info = SignatureInfo() + write_signature_info(signer, sig_info) + buf = tlv_encode(sig_info) + ret.append(Component.from_bytes(bytes([TypeNumber.SIGNATURE_INFO, len(buf)]) + buf)) + + # SignatureValue + sig_size = signer.get_signature_value_size() + tlv_length = 1 + get_tl_num_size(sig_size) + sig_size + buf = bytearray(tlv_length) + buf[0] = TypeNumber.SIGNATURE_VALUE + offset = 1 + write_tl_num(sig_size, buf, 1) + signer.write_signature_value(memoryview(buf)[offset:], ret) + ret.append(Component.from_bytes(buf)) + + return ret + + +def make_command_v2(module, command, face: Face | None = None, **kwargs): + # V2 returns the Command Interest name for the NDNv3 signed Interest + # Note: this behavior is supported by NFD and YaNFD but has not been documented yet (on 06/26/2022): + # https://redmine.named-data.net/projects/nfd/wiki/ControlCommand + # Add ``app_param=b'', signer=sec.DigestSha256Signer(for_interest=True)`` to app.express when using this. + local = face.isLocalFace() if face else True + + if local: + ret = Name.from_str(f"/localhost/nfd/{module}/{command}") + else: + ret = Name.from_str(f"/localhop/nfd/{module}/{command}") + # Command parameters + cp = ControlParameters(cp=ControlParametersValue()) + for k, v in kwargs.items(): + if k == 'strategy': + cp.cp.strategy = Strategy(name=v) + else: + setattr(cp.cp, k, v) + ret.append(Component.from_bytes(tlv_encode(cp))) + return ret + + +def parse_response(buf): + buf = parse_and_check_tl(memoryview(buf), 0x65) + cr = tlv_parse(ControlResponse, buf) + ret = {} + ret['status_code'] = cr.status_code + ret['status_text'] = cr.status_text + params = cr.body + for f in dc.fields(ControlParametersValue): + val = getattr(params, f.name) if params is not None else None + if isinstance(val, memoryview): + val = bytes(val) + ret[f.name] = val + return ret diff --git a/src/ndn/app_support/security_v2_2.py b/src/ndn/app_support/security_v2_2.py index 4b585b5..0bb1822 100644 --- a/src/ndn/app_support/security_v2_2.py +++ b/src/ndn/app_support/security_v2_2.py @@ -31,10 +31,10 @@ from ..encoding.ndn_format_0_3_2 import ( ContentType, DataPacketValue, - KeyLocator, MetaInfo, SignatureInfo, TypeNumber, + write_signature_info, ) @@ -144,14 +144,7 @@ def new_cert(key_name, issuer_id_component, pub_key, signer, not_after=not_after, ), ) - signer.write_signature_info(signature_info) - if (signature_info.key_locator is not None - and not isinstance(signature_info.key_locator, KeyLocator)): - old_key_locator = signature_info.key_locator - signature_info.key_locator = KeyLocator( - name=old_key_locator.name, - key_digest=old_key_locator.key_digest, - ) + write_signature_info(signer, signature_info) cert_val = CertificateV2Value( name=cert_name, content=pub_key, diff --git a/src/ndn/appv2_2.py b/src/ndn/appv2_2.py new file mode 100644 index 0000000..b988043 --- /dev/null +++ b/src/ndn/appv2_2.py @@ -0,0 +1,754 @@ +# ----------------------------------------------------------------------------- +# Copyright (C) 2019-2022 The python-ndn authors +# +# This file is part of python-ndn. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ----------------------------------------------------------------------------- +import asyncio as aio +import typing +import struct +import logging +from hashlib import sha256 +from dataclasses import dataclass +from .transport.face import Face +from .transport.prefix_registerer import PrefixRegisterer +from . import security as sec +from . import encoding as enc +from . import name_tree +from . import types +from . import utils +from .encoding import ndnlp_v2_2 as ndnlp +from .encoding import ndn_format_0_3_2 as fmt +from .encoding.tlv_model_v2 import tlv_encode +from .client_conf import read_client_conf, default_face, default_keychain +from .transport.nfd_registerer import NfdRegister2 + + +DEFAULT_LIFETIME = 4000 + +ValidResult = types.ValidResult + +PktContext = dict[str, any] +r"""The context for NDN Interest or Data handling.""" + +ReplyFunc = typing.Callable[[enc.BinaryStr], bool] +r""" +Continuation function for :any:`IntHandler` to respond to an Interest. + +.. function:: (data: BinaryStr) -> bool + + :param data: an encoded Data packet. + :type data: :any:`BinaryStr` + :return: True for success, False upon error. +""" + +IntHandler = typing.Callable[[enc.FormalName, enc.BinaryStr | None, ReplyFunc, PktContext], None] +r""" +Interest handler function associated with a name prefix. + +The function should use the provided ``reply`` callback to reply with Data, which can handle PIT +token properly. + +.. function:: (name: FormalName, app_param: Optional[BinaryStr], reply: ReplyFunc, context: PktContext) -> None + + :param name: Interest name. + :type name: :any:`FormalName` + :param app_param: Interest ApplicationParameters value, or None if absent. + :type app_param: Optional[:any:`BinaryStr`] + :param reply: continuation function to respond with Data. + :type reply: :any:`ReplyFunc` + :param context: packet handler context. + :type context: :any:`PktContext` + +.. note:: + Interest handler function must be a normal function instead of an ``async`` one. + This is on purpose, because an Interest is supposed to be replied ASAP, + even it cannot finish the request in time. + To provide some feedback, a better practice is replying with an Application NACK + (or some equivalent Data packet saying the operation cannot be finished in time). + If you want to use ``await`` in the handler, please use ``asyncio.create_task`` to create a new coroutine. +""" + +Validator = typing.Callable[[enc.FormalName, fmt.SignaturePtrs, PktContext], + typing.Coroutine[any, None, ValidResult]] +r""" +Validator function that validates Interest or Data signature against trust policy. + +.. function:: (name: FormalName, sig: SignaturePtrs, context: PktContext) -> Coroutine[ValidResult] + + :param name: Interest or Data name. + :type name: :any:`FormalName` + :param sig: packet signature pointers. + :type sig: :any:`SignaturePtrs` + :param context: packet handler context. + :type context: :any:`PktContext` +""" + + +async def pass_all(_name, _sig, _context): + return types.ValidResult.PASS + + +@dataclass +class PrefixTreeNode: + callback: IntHandler = None + validator: Validator | None = None + + +@dataclass +class PendingIntEntry: + future: aio.Future + deadline: int + can_be_prefix: bool + must_be_fresh: bool + validator: Validator + implicit_sha256: enc.BinaryStr = b'' + task: aio.Task | None = None + + async def satisfy(self, data: types.DataTuple): + name, meta_info, content, sig, raw_packet = data + pkt_context = { + 'meta_info': meta_info, + 'sig_ptrs': sig, + 'raw_packet': raw_packet, + 'deadline': self.deadline, + } + if self.validator is not None: + try: + valid = await self.validator(name, sig, pkt_context) + except (TimeoutError, aio.CancelledError): + valid = ValidResult.TIMEOUT + else: + valid = ValidResult.FAIL + if self.future.cancelled() or self.future.done(): + # Don't know why but there was a race condition with timeout() + # The sequence was: Interest sent -> Data arrived -> timeout() -> satisfy() + # Cannot reproduce the scenario. Especially, delay in validator() does not trigger the race condition + # But anyway, let me add a guard check here. + return + if valid == ValidResult.PASS or valid == ValidResult.ALLOW_BYPASS: + self.future.set_result((name, content, pkt_context)) + else: + self.future.set_exception(types.ValidationFailure(name, meta_info, content, sig, valid)) + + +class InterestTreeNode: + pending_list: list[PendingIntEntry] + + def __init__(self): + self.pending_list = [] + + def append_interest(self, future: aio.Future, deadline: int, param: fmt.InterestParam, + validator: Validator, implicit_sha256: enc.BinaryStr): + self.pending_list.append( + PendingIntEntry(future, deadline, param.can_be_prefix, param.must_be_fresh, validator, implicit_sha256)) + + def nack_interest(self, nack_reason: int) -> bool: + for entry in self.pending_list: + entry.future.set_exception(types.InterestNack(nack_reason)) + return True + + def satisfy(self, data: types.DataTuple, is_prefix: bool) -> bool: + unsatisfied_entries = [] + raw_packet = data[4] + for entry in self.pending_list: + if entry.can_be_prefix or not is_prefix: + if len(entry.implicit_sha256) > 0: + data_sha256 = sha256(raw_packet).digest() + passed = data_sha256 == entry.implicit_sha256 + else: + passed = True + else: + passed = False + if passed: + # Try to validate the packet + aio.create_task(entry.satisfy(data)) + else: + unsatisfied_entries.append(entry) + if unsatisfied_entries: + self.pending_list = unsatisfied_entries + return False + else: + return True + + def timeout(self, future: aio.Future): + # Exception is raised by outside code. + for ele in self.pending_list: + if ele.future is future and ele.task is not None: + ele.task.cancel() + self.pending_list = [ele for ele in self.pending_list if ele.future is not future] + return not self.pending_list + + def cancel(self): + for entry in self.pending_list: + entry.future.cancel() + if entry.task is not None: + entry.task.cancel() + + +class NDNApp: + """ + An NDN application. + """ + # PIT and FIB here are not real PIT/FIB, but a data structure that handles expressed Interests (for PIT) + # and registered handlers & routes (for FIB). Since they share the functionality with real PIT and FIB, + # I borrow the word to have a shorter variable name. + _pit: name_tree.NameTrie = None + _fib: name_tree.NameTrie = None + face: Face = None + registerer: PrefixRegisterer = None + _autoreg_routes: list[enc.FormalName] + logger: logging.Logger + + def __init__(self, face=None, client_conf=None, registerer=None): + self.logger = logging.getLogger(__name__) + config = client_conf if client_conf else {} + if not face: + if 'transport' not in config: + config = read_client_conf() | config + if face is not None: + self.face = face + else: + self.face = default_face(config['transport']) + if registerer is not None: + self.registerer = registerer + else: + self.registerer = NfdRegister2() + self.registerer.set_app(app=self) + self.face.callback = self._receive + self._pit = name_tree.NameTrie() + self._fib = name_tree.NameTrie() + self._autoreg_routes = [] + + @staticmethod + def default_keychain(client_conf=None) -> sec.Keychain: + if not client_conf: + config = read_client_conf() + else: + config = read_client_conf() | client_conf + return default_keychain(config['pib'], config['tpm']) + + async def _receive(self, typ: int, data: enc.BinaryStr): + """ + Pipeline when a packet is received. + + :param typ: the Type. + :param data: the Value of the packet with TL. + """ + # if self.logger.isEnabledFor(logging.DEBUG): + # self.logger.debug('Packet received %s, %s' % (typ, bytes(data))) + if typ == ndnlp.LpTypeNumber.LP_PACKET: + try: + lp_pkt = ndnlp.parse_lp_packet_v2(data, with_tl=True) + except (enc.DecodeError, TypeError, ValueError, struct.error): + self.logger.warning('Unable to decode received packet') + return + if lp_pkt.nack is not None: + nack_reason = lp_pkt.nack.nack_reason + else: + nack_reason = None + pit_token = lp_pkt.pit_token + data = lp_pkt.fragment + typ, _ = enc.parse_tl_num(data) + else: + nack_reason = None + pit_token = None + + if nack_reason is not None: + try: + name, _, _, _ = fmt.parse_interest(data, with_tl=True) + except (enc.DecodeError, TypeError, ValueError, struct.error): + self.logger.warning('Unable to decode the fragment of LpPacket') + return + if self.logger.isEnabledFor(logging.DEBUG): + self.logger.debug('NetworkNack received %s, reason=%s', enc.Name.to_str(name), nack_reason) + self._on_nack(name, nack_reason) + else: + if typ == fmt.TypeNumber.INTEREST: + try: + name, param, app_param, sig = fmt.parse_interest(data, with_tl=True) + except (enc.DecodeError, TypeError, ValueError, struct.error): + self.logger.warning('Unable to decode received packet') + return + if self.logger.isEnabledFor(logging.DEBUG): + if pit_token: + self.logger.debug('Interest received %s w/ token=%s', + enc.Name.to_str(name), bytes(pit_token).hex()) + else: + self.logger.debug('Interest received %s', enc.Name.to_str(name)) + await self._on_interest(name, pit_token, param, app_param, sig, raw_packet=data) + elif typ == fmt.TypeNumber.DATA: + try: + name, meta_info, content, sig = fmt.parse_data(data, with_tl=True) + except (enc.DecodeError, TypeError, ValueError, struct.error): + self.logger.warning('Unable to decode received packet') + return + if self.logger.isEnabledFor(logging.DEBUG): + self.logger.debug('Data received %s', enc.Name.to_str(name)) + await self._on_data(name, meta_info, content, sig, raw_packet=data) + else: + self.logger.warning('Unable to decode received packet') + + @staticmethod + def make_data(name: enc.NonStrictName, content: enc.BinaryStr | None, + signer: enc.Signer | None, **kwargs): + r""" + Encode a data packet without requiring an NDNApp instance. + This is simply a wrapper of encoding.make_data. + I write this because most people seem not aware of the ``make_data`` function in the encoding package. + The corresponding ``make_interest`` is less useful (one should not reuse nonce) and thus not wrapped. + Sync protocol should use encoding.make_interest if necessary. + Also, since having a default signer encourages bad habit, + prepare_data is removed except for command Interests sent to NFD. + Please call ``keychain.get_signer({})`` to use the default certificate. + + :param name: the Name. + :type name: :any:`NonStrictName` + :param content: the Content. + :type content: Optional[:any:`BinaryStr`] + :param signer: the Signer used to sign the packet. + :type signer: Optional[:any:`Signer`] + :param kwargs: arguments for :any:`MetaInfo`. + :return: TLV encoded Data packet. + """ + if 'meta_info' in kwargs: + meta_info = kwargs['meta_info'] + else: + meta_info = fmt.MetaInfo.from_dict(kwargs) + return fmt.make_data(name, meta_info, content, signer=signer) + + async def _on_interest(self, name: enc.FormalName, pit_token: enc.BinaryStr | None, + param: fmt.InterestParam, app_param: enc.BinaryStr | None, sig: fmt.SignaturePtrs, + raw_packet: enc.BinaryStr): + trie_step = self._fib.longest_prefix(name) + if not trie_step: + self.logger.warning('No route: %s', name) + return + node: PrefixTreeNode = trie_step.value + if node.callback is None: + self.logger.warning('No callback: %s', name) + return + sig_required = app_param is not None or sig.signature_info is not None + if sig_required: + if not await sec.params_sha256_checker(name, sig): + self.logger.warning('Drop malformed Interest: %s', name) + return + + # Use context to handle misc parameters + if param.lifetime is not None: + deadline = utils.timestamp() + param.lifetime + else: + deadline = utils.timestamp() + DEFAULT_LIFETIME + context = { + 'int_param': param, + 'pit_token': pit_token, + 'sig_ptrs': sig, + 'raw_packet': raw_packet, + 'deadline': deadline, + } + + def reply(data: enc.BinaryStr) -> bool: + now = utils.timestamp() + if now > deadline: + self.logger.warning('Deadline passed, unable to reply to %s', enc.Name.to_str(name)) + return False + if pit_token is None: + self._put_raw_packet(data) + else: + self._put_raw_packet_with_pit_token(data, pit_token) + + # In case the validator blocks the pipeline, create a task + async def submit_interest(): + if sig_required: + # In v2, to enforce security, validator is required. Also, all interests with app_param are checked. + # The validator needs to manually pass it if the application wants to handle unsigned Interests with + # app_param. + if node.validator is not None: + valid = await node.validator(name, sig, context) + else: + valid = ValidResult.FAIL + else: + valid = ValidResult.PASS + if valid == ValidResult.PASS or valid == ValidResult.ALLOW_BYPASS: + node.callback(name, app_param, reply, context) + else: + self.logger.warning('Drop unvalidated Interest: %s', name) + return + aio.create_task(submit_interest()) + + def _put_raw_packet(self, data: enc.BinaryStr): + r""" + Send a raw Data packet. + + :param data: TLV encoded Data packet. + :type data: :any:`BinaryStr` + :raises NetworkError: the face to NFD is down. + """ + if not self.face.running: + raise types.NetworkError('cannot send packet before connected') + self.face.send(data) + + def _put_raw_packet_with_pit_token(self, data: enc.BinaryStr, pit_token: enc.BinaryStr): + r""" + Wrap a raw Data packet with PIT Token and send. + Used to reply an Interest with PIT Token provided. + + :param data: TLV encoded Data packet. + :type data: :any:`BinaryStr` + :param pit_token: The PIT Token provided. + :type pit_token: :any:`BinaryStr` + :raises NetworkError: the face to NFD is down. + """ + if not self.face.running: + raise types.NetworkError('cannot send packet before connected') + pkt = ndnlp.LpPacket(lp_packet=ndnlp.LpPacketValue(pit_token=pit_token, fragment=data)) + wire = tlv_encode(pkt) + self.face.send(wire) + + def _put_raw_packet_with_pit_token_nocopy(self, data: enc.BinaryStr, pit_token: enc.BinaryStr): + r""" + Wrap a raw Data packet with PIT Token and send. + Used to reply an Interest with PIT Token provided. + + This function is reserved as a backup because it assumes the face to be stream face. + + :param data: TLV encoded Data packet. + :type data: :any:`BinaryStr` + :param pit_token: The PIT Token provided. + :type pit_token: :any:`BinaryStr` + :raises NetworkError: the face to NFD is down. + """ + # To avoid extra copy, we manually encode the header and send it separately from Data body + # The format is: LP-T LP-L (PIT-TOKEN-TLV) FRAG-T FRAG-L + if not self.face.running: + raise types.NetworkError('cannot send packet before connected') + pt_wire = tlv_encode(ndnlp.LpPacketValue(pit_token=pit_token)) + frag_l = len(data) + lp_l = len(pt_wire) + enc.get_tl_num_size(ndnlp.LpTypeNumber.FRAGMENT) + enc.get_tl_num_size(frag_l) + wire_l = enc.get_tl_num_size(ndnlp.LpTypeNumber.LP_PACKET) + enc.get_tl_num_size(lp_l) + lp_l + wire = bytearray(wire_l) + pos = 0 + pos += enc.write_tl_num(ndnlp.LpTypeNumber.LP_PACKET, wire, pos) + pos += enc.write_tl_num(lp_l, wire, pos) + wire[pos:pos+len(pt_wire)] = pt_wire + pos += len(pt_wire) + pos += enc.write_tl_num(ndnlp.LpTypeNumber.FRAGMENT, wire, pos) + pos += enc.write_tl_num(frag_l, wire, pos) + self.face.send(wire) + self.face.send(data) + + def attach_handler(self, name: enc.NonStrictName, handler: IntHandler, + validator: Validator | None = None): + """ + Attach an Interest handler at a name prefix. + Incoming Interests under the specified name prefix will be dispatched to the handler. + + This only sets the handler within NDNApp, but does not send prefix registration commands + to the forwarder. + To register the prefix in the forwarder, use :any:`register`. + The handler association is retained even if the forwarder is disconnected. + + :param name: name prefix. + :type name: :any:`NonStrictName` + :param handler: Interest handler function. + :type handler: :any:`IntHandler` + :param validator: validator for signed Interests. + Non signed Interests, i.e. those without ApplicationParameters and SignatureInfo, are + passed to the handler directly without calling the validator. + Interests with malformed ParametersSha256DigestComponent are dropped silently. + If a validator is not provided (set to ``None``), signed Interests will be dropped. + Otherwise, signed Interests are passed to the validator. + Those failing the validation are dropped silently. + Those passing the validation are passed to the handler function. + :type validator: Optional[:any:`Validator`] + """ + name = enc.Name.normalize(name) + node = self._fib.setdefault(name, PrefixTreeNode()) + if node.callback: + raise ValueError(f'Duplicated handler attachment: {enc.Name.to_str(name)}') + node.callback = handler + node.validator = validator + + def detach_handler(self, name: enc.NonStrictName): + """ + Detach an Interest handler at a name prefix. + + This only deletes the handler within NDNApp, but does not unregister the prefix in the + forwarder. + To unregister the prefix in the forwarder, use :any:`unregister`. + + :param name: name prefix. This must exactly match the name passed to :any:`attach_handler`. + If there are Interest handlers attached to longer prefixes, each handler must + be removed explicitly. + :type name: :any:`NonStrictName` + """ + del self._fib[enc.Name.normalize(name)] + + async def register(self, name: enc.NonStrictName) -> bool: + """ + Register a prefix in the forwarder. + + This only sends the prefix registration command to the forwarder. + In order to receive incoming Interests, you also need to use :any:`attach_handler` to + attach an Interest handler function. + + :param name: name prefix. + :type name: :any:`NonStrictName` + + :raises ValueError: the prefix is already registered. + :raises NetworkError: the face to NFD is down now. + """ + name = enc.Name.normalize(name) + return await self.registerer.register(name) + + async def unregister(self, name: enc.NonStrictName) -> bool: + """ + Unregister a prefix in the forwarder. + + :param name: name prefix. + :type name: :any:`NonStrictName` + """ + name = enc.Name.normalize(name) + return await self.registerer.unregister(name) + + def express_raw_interest(self, + final_name: enc.NonStrictName, + interest_param: fmt.InterestParam, + raw_interest: enc.BinaryStr, + validator: Validator, + no_response: bool = False + ) -> typing.Coroutine[any, None, + tuple[enc.FormalName, enc.BinaryStr | None, PktContext]]: + if no_response: + self.face.send(raw_interest) + return None + if validator is None: + raise ValueError('Data Validator must not be None when expressing an Interest.') + final_name = enc.Name.normalize(final_name) + future = aio.get_running_loop().create_future() + # Handle implicit SHA256 + if enc.Component.get_type(final_name[-1]) == enc.Component.TYPE_IMPLICIT_SHA256: + node_name = final_name[:-1] + implicit_sha256 = enc.Component.get_value(final_name[-1]) + else: + node_name = final_name + implicit_sha256 = b'' + node: InterestTreeNode = self._pit.setdefault(node_name, InterestTreeNode()) + deadline = utils.timestamp() + if interest_param.lifetime is not None: + deadline += interest_param.lifetime + else: + deadline += DEFAULT_LIFETIME + node.append_interest(future, deadline, interest_param, validator, implicit_sha256) + self.face.send(raw_interest) + return self._wait_for_data(future, deadline, node_name, node) + + async def _wait_for_data(self, future: aio.Future, deadline: int, node_name: enc.FormalName, + node: InterestTreeNode): + lifetime = deadline - utils.timestamp() + if lifetime <= 0: + # This happens if the application sends an Interest, does some calculation, and then fetches the result. + # The Interest should be satisfied now. Thus, it should not be considered as an error. + lifetime = 100 + try: + data_name, content, pkt_context = await aio.wait_for(future, timeout=lifetime/1000.0) + except TimeoutError: + if node.timeout(future): + del self._pit[node_name] + raise types.InterestTimeout() + except aio.CancelledError: + raise types.InterestCanceled() + # ValidationError, InterestNack are passed to the parent caller + return data_name, content, pkt_context + + async def _on_data(self, name: enc.FormalName, meta_info: fmt.MetaInfo, + content: enc.BinaryStr | None, sig: fmt.SignaturePtrs, + raw_packet: enc.BinaryStr): + clean_list = [] + for prefix, node in self._pit.prefixes(name): + if node.satisfy((name, meta_info, content, sig, raw_packet), prefix != name): + clean_list.append(prefix) + for prefix in clean_list: + del self._pit[prefix] + + def _on_nack(self, name: enc.FormalName, nack_reason: int): + try: + node = self._pit[name] + except KeyError: + node = None + if node: + if node.nack_interest(nack_reason): + del self._pit[name] + + def express(self, name: enc.NonStrictName, validator: Validator, + app_param: enc.BinaryStr | None = None, + signer: enc.Signer | None = None, + **kwargs) -> typing.Coroutine[any, None, + tuple[enc.FormalName, enc.BinaryStr | None, PktContext]]: + r""" + Express an Interest. + + The Interest packet is sent immediately and a coroutine used to get the result is returned. + Awaiting on the returned coroutine will block until the Data is received. + It then returns the Data name, Data Content value, and :any:`PktContext`. + An exception is raised if NDNApp is unable to retrieve the Data. + + :param name: Interest name. + :type name: :any:`NonStrictName` + :param validator: validator for the retrieved Data packet. + :type validator: :any:`Validator` + :param app_param: Interest ApplicationParameters value. If this is not None, a signed + Interest is sent. NDNApp does not support sending parameterized + Interests that are not signed. + :type app_param: Optional[:any:`BinaryStr`] + :param signer: Signer for Interest signing. This is required if `app_param` is specified. + :type signer: Optional[:any:`Signer`] + :param kwargs: arguments for :any:`InterestParam`. + :return: A tuple of (Name, Content, PacketContext) after ``await``. + :rtype: Coroutine[Any, None, Tuple[:any:`FormalName`, Optional[:any:`BinaryStr`], :any:`PktContext`]] + + The following exceptions may be raised by ``express``: + + :raises NetworkError: the face to NFD is down before sending this Interest. + :raises ValueError: when the signer is missing but app_param presents. + + The following exceptions may be raised by the returned coroutine: + + :raises InterestNack: an NetworkNack is received. + :raises InterestTimeout: time out. + :raises ValidationFailure: unable to validate the Data packet. + :raises InterestCanceled: the face to NFD is shut down after sending this Interest. + """ + if not self.face.running: + raise types.NetworkError('cannot send packet before connected') + if app_param is not None and signer is None: + raise ValueError('An Interest with AppParam is required to be signed.') + if 'interest_param' in kwargs: + interest_param = kwargs['interest_param'] + else: + if 'nonce' not in kwargs: + kwargs['nonce'] = utils.gen_nonce() + interest_param = fmt.InterestParam.from_dict(kwargs) + interest, final_name = fmt.make_interest(name, interest_param, app_param, signer=signer, need_final_name=True) + no_response = kwargs.get('no_response', False) + return self.express_raw_interest(final_name, interest_param, interest, validator, no_response) + + def route(self, name: enc.NonStrictName, validator: Validator | None = None): + r""" + A decorator used to register a permanent route for a specific prefix. + The decorated function should be an :any:`IntHandler`. + + This function is non-blocking and can be called at any time. + It can be called before connecting to the forwarder. + Every time a forwarder connection is established, NDNApp will automatically send + prefix registration commands. + Errors in prefix registration are ignored. + + :param name: name prefix. + :type name: :any:`NonStrictName` + :param validator: validator for signed Interests. See :any:`attach_handler` for details. + :type validator: Optional[:any:`Validator`] + + :examples: + .. code-block:: python3 + + app = NDNApp() + + @app.route('/example/rpc') + def on_interest(name, app_param, reply, context): + pass + + """ + name = enc.Name.normalize(name) + + def decorator(func: IntHandler): + self._autoreg_routes.append(name) + self.attach_handler(name, func, validator) + if self.face.running: + aio.create_task(self.register(name)) + return func + return decorator + + def _clean_up(self): + for node in self._pit.itervalues(): + node.cancel() + # FIB is not cleared now + self._pit.clear() + + def shutdown(self): + """ + Manually shutdown the face to NFD. + """ + self.logger.info('Manually shutdown') + self.face.shutdown() + + async def main_loop(self, after_start: typing.Awaitable = None) -> bool: + """ + The main loop of NDNApp. + + :param after_start: the coroutine to start after connection to NFD is established. + :return: ``True`` if the connection is shutdown not by ``Ctrl+C``. + For example, manually or by the other side. + """ + async def starting_task(): + for name in self._autoreg_routes: + await self.register(name) + if after_start: + try: + await after_start + except Exception: + self.face.shutdown() + raise + + try: + await self.face.open() + except (FileNotFoundError, ConnectionError, OSError, PermissionError): + if after_start: + if isinstance(after_start, typing.Coroutine): + after_start.close() + elif isinstance(after_start, (aio.Task, aio.Future)): + after_start.cancel() + raise + task = aio.create_task(starting_task()) + self.logger.debug('Connected to NFD node, start running...') + try: + await self.face.run() + ret = True + except aio.CancelledError: + self.logger.info('Shutting down') + ret = False + finally: + self.face.shutdown() + self._clean_up() + await task + return ret + + def run_forever(self, after_start: typing.Awaitable = None): + """ + A non-async wrapper of :meth:`main_loop`. + + :param after_start: the coroutine to start after connection to NFD is established. + + :examples: + .. code-block:: python3 + + app = NDNApp() + + if __name__ == '__main__': + app.run_forever(after_start=main()) + """ + try: + aio.run(self.main_loop(after_start)) + except KeyboardInterrupt: + self.logger.info('Receiving Ctrl+C, exit') diff --git a/src/ndn/encoding/ndn_format_0_3_2.py b/src/ndn/encoding/ndn_format_0_3_2.py index 181fc4c..b81915e 100644 --- a/src/ndn/encoding/ndn_format_0_3_2.py +++ b/src/ndn/encoding/ndn_format_0_3_2.py @@ -14,7 +14,7 @@ __all__ = [ 'TypeNumber', 'ContentType', 'SignatureType', 'KeyLocator', - 'SignatureInfo', + 'SignatureInfo', 'write_signature_info', 'Links', 'MetaInfo', 'InterestParam', 'SignaturePtrs', 'make_interest', 'make_data', 'parse_interest', 'parse_data', 'Interest', 'Data', ] @@ -93,6 +93,20 @@ class SignatureInfo: default=None, metadata={'tlv_type': TypeNumber.SIGNATURE_SEQ_NUM}) +def write_signature_info(signer: Signer, signature_info: SignatureInfo) -> None: + """ + Let *signer* fill *signature_info*. + + Signers still assign the v1 ``KeyLocator`` model, which the dataclass + encoder cannot serialize, so it is converted to :class:`KeyLocator`. + """ + signer.write_signature_info(signature_info) + key_locator = signature_info.key_locator + if key_locator is not None and not isinstance(key_locator, KeyLocator): + signature_info.key_locator = KeyLocator( + name=key_locator.name, key_digest=key_locator.key_digest) + + @dc.dataclass class Links: names: list[NDNName] = dc.field( @@ -255,7 +269,7 @@ def make_interest(name: NonStrictName, names=list(interest_param.forwarding_hint)) if signer is not None: value.signature_info = SignatureInfo() - signer.write_signature_info(value.signature_info) + write_signature_info(signer, value.signature_info) if value.application_parameters is None: value.application_parameters = b'' @@ -280,7 +294,7 @@ def make_data(name: NonStrictName, value = DataPacketValue(name=name, meta_info=meta_info, content=content) if signer is not None: value.signature_info = SignatureInfo() - signer.write_signature_info(value.signature_info) + write_signature_info(signer, value.signature_info) encoded_value = tlv_encode(value, markers={'##signer': signer}) return _wrap_tlv(TypeNumber.DATA, encoded_value) diff --git a/src/ndn/encoding/tlv_model_v2.py b/src/ndn/encoding/tlv_model_v2.py index 155fd20..da696cd 100644 --- a/src/ndn/encoding/tlv_model_v2.py +++ b/src/ndn/encoding/tlv_model_v2.py @@ -98,6 +98,7 @@ class Outer: import dataclasses import struct import typing +import weakref from enum import Enum, Flag from hashlib import sha256 from types import UnionType @@ -222,6 +223,69 @@ def _map_val_meta(metadata: dict) -> dict: return m +# --------------------------------------------------------------------------- +# Per-class schema cache +# --------------------------------------------------------------------------- + +@dataclasses.dataclass(frozen=True, slots=True) +class _FieldSpec: + """Everything the encoder/parser needs about one field, resolved once.""" + name: str + kind: str + metadata: typing.Mapping + annotation: typing.Any + tlv_type: typing.Optional[int] + enum_cls: typing.Optional[type] = None + elem: typing.Optional['_FieldSpec'] = None + key: typing.Optional['_FieldSpec'] = None + val: typing.Optional['_FieldSpec'] = None + + +def _make_spec(name: str, annotation, metadata) -> _FieldSpec: + kind = _infer_kind(annotation, metadata) + annotation = _unwrap_optional(annotation) + enum_cls = elem = key = val = None + if kind == 'uint' and isinstance(annotation, type) and issubclass(annotation, (Enum, Flag)): + enum_cls = annotation + elif kind == 'repeated': + elem = _make_spec(name, _element_annotation(annotation), metadata) + elif kind == 'map': + key_ann, val_ann = _map_annotations(annotation) + key = _make_spec(name, key_ann, _map_key_meta(metadata)) + val = _make_spec(name, val_ann, _map_val_meta(metadata)) + return _FieldSpec(name, kind, metadata, annotation, metadata.get('tlv_type'), + enum_cls, elem, key, val) + + +_SCHEMA_CACHE: 'weakref.WeakKeyDictionary[type, tuple[_FieldSpec, ...]]' = weakref.WeakKeyDictionary() + + +def _get_schema(cls) -> tuple[_FieldSpec, ...]: + """ + Return the TLV field specs of dataclass *cls* in declaration order. + + Built on first use rather than at class definition so that forward + references to classes defined later in the same module can be resolved. + Fields with neither ``tlv_type`` nor ``field_type`` metadata are skipped. + """ + try: + return _SCHEMA_CACHE[cls] + except KeyError: + pass + hints = typing.get_type_hints(cls) + specs = [] + for f in dataclasses.fields(cls): + if 'tlv_type' not in f.metadata and 'field_type' not in f.metadata: + continue + spec = _make_spec(f.name, hints[f.name], f.metadata) + if spec.kind not in _ZERO_WIRE_KINDS and spec.tlv_type is None: + continue + specs.append(spec) + schema = tuple(specs) + _SCHEMA_CACHE[cls] = schema + return schema + + # --------------------------------------------------------------------------- # Interest-name helpers (used by both pass-1 and pass-2) # --------------------------------------------------------------------------- @@ -400,8 +464,7 @@ def _uint_value_len(val: int, fname: str, fixed_len) -> int: return n -def _encoded_length_field(fname: str, val, kind: str, annotation, metadata: dict, - markers: dict) -> int: +def _encoded_length_field(fname: str, val, spec: _FieldSpec, markers: dict) -> int: """ Compute the encoded byte count of one TLV field (T + L + V). @@ -410,6 +473,7 @@ def _encoded_length_field(fname: str, val, kind: str, annotation, metadata: dict subclasses. Returns 0 when the field is absent (*val* is ``None``/falsy for bool). """ + kind = spec.kind # Zero-wire kinds: handled before looking up tlv_type. if kind == 'offset_marker': return 0 @@ -418,16 +482,16 @@ def _encoded_length_field(fname: str, val, kind: str, annotation, metadata: dict signer = markers.get('##signer') if signer is None: return 0 - type_num = metadata['tlv_type'] + type_num = spec.tlv_type sig_size = signer.get_signature_value_size() markers[f'{fname}##sig_size'] = sig_size markers.setdefault('##sig_covered_part', []) return get_tl_num_size(type_num) + get_tl_num_size(sig_size) + sig_size if kind == 'interest_name': - return _encoded_length_interest_name(fname, val, metadata, markers) + return _encoded_length_interest_name(fname, val, spec.metadata, markers) - type_num = metadata['tlv_type'] + type_num = spec.tlv_type # BoolField: present if truthy, absent otherwise if kind == 'bool': @@ -441,7 +505,7 @@ def _encoded_length_field(fname: str, val, kind: str, annotation, metadata: dict val = val.value if not isinstance(val, int) or val < 0: raise TypeError(f'{fname}={val!r} is not a non-negative integer') - fixed_len = metadata.get('fixed_len') + fixed_len = spec.metadata.get('fixed_len') vlen = _uint_value_len(val, fname, fixed_len) markers[f'{fname}##encoded_length'] = vlen # L for uint is always 1 byte because vlen ∈ {1,2,4,8} < 253 @@ -489,28 +553,20 @@ def _encoded_length_field(fname: str, val, kind: str, annotation, metadata: dict if kind == 'repeated': if not val: return 0 - elem_ann = _element_annotation(annotation) - elem_kind = _infer_kind(elem_ann, metadata) + elem = spec.elem total = 0 for i, ele in enumerate(val): - total += _encoded_length_field( - f'{fname}[{i}]', ele, elem_kind, elem_ann, metadata, markers) + total += _encoded_length_field(f'{fname}[{i}]', ele, elem, markers) return total if kind == 'map': if not val: return 0 - key_ann, val_ann = _map_annotations(annotation) - key_meta = _map_key_meta(metadata) - vl_meta = _map_val_meta(metadata) - key_kind = _infer_kind(key_ann, key_meta) - vl_kind = _infer_kind(val_ann, vl_meta) + key_spec, val_spec = spec.key, spec.val total = 0 for i, (k, v) in enumerate(val.items()): - total += _encoded_length_field( - f'{fname}[{i}#k]', k, key_kind, key_ann, key_meta, markers) - total += _encoded_length_field( - f'{fname}[{i}#v]', v, vl_kind, val_ann, vl_meta, markers) + total += _encoded_length_field(f'{fname}[{i}#k]', k, key_spec, markers) + total += _encoded_length_field(f'{fname}[{i}#v]', v, val_spec, markers) return total raise TypeError(f'Unknown field kind {kind!r} for {fname!r}') @@ -518,16 +574,9 @@ def _encoded_length_field(fname: str, val, kind: str, annotation, metadata: dict def _encoded_length_model(obj, markers: dict) -> int: """Compute the total encoded length for all TLV fields of a dataclass object.""" - cls = type(obj) - hints = typing.get_type_hints(cls) total = 0 - for f in dataclasses.fields(cls): - ann = hints[f.name] - kind = _infer_kind(ann, f.metadata) - if kind not in _ZERO_WIRE_KINDS and 'tlv_type' not in f.metadata: - continue - total += _encoded_length_field( - f.name, getattr(obj, f.name), kind, ann, f.metadata, markers) + for spec in _get_schema(type(obj)): + total += _encoded_length_field(spec.name, getattr(obj, spec.name), spec, markers) markers['##encoded_length'] = total return total @@ -536,7 +585,7 @@ def _encoded_length_model(obj, markers: dict) -> int: # Encoding — pass 2: write bytes # --------------------------------------------------------------------------- -def _encode_into_field(fname: str, val, kind: str, annotation, metadata: dict, +def _encode_into_field(fname: str, val, spec: _FieldSpec, markers: dict, wire: VarBinaryStr, offset: int) -> int: """ Write one TLV field into *wire* at *offset*. @@ -545,6 +594,8 @@ def _encode_into_field(fname: str, val, kind: str, annotation, metadata: dict, Returns the number of bytes written. Must be called after the matching :func:`_encoded_length_field` call so that ``markers`` is populated. """ + kind = spec.kind + metadata = spec.metadata # Zero-wire kinds: handled before looking up tlv_type. if kind == 'offset_marker': markers[fname] = offset @@ -554,7 +605,7 @@ def _encode_into_field(fname: str, val, kind: str, annotation, metadata: dict, signer = markers.get('##signer') if signer is None: return 0 - type_num = metadata['tlv_type'] + type_num = spec.tlv_type sig_size = markers[f'{fname}##sig_size'] # Collect the covered region: from cover_start up to current offset. cover_start_field = metadata.get('cover_start') @@ -576,7 +627,7 @@ def _encode_into_field(fname: str, val, kind: str, annotation, metadata: dict, if kind == 'interest_name': return _encode_into_interest_name(fname, val, metadata, markers, wire, offset) - type_num = metadata['tlv_type'] + type_num = spec.tlv_type if kind == 'bool': if val: @@ -634,31 +685,23 @@ def _encode_into_field(fname: str, val, kind: str, annotation, metadata: dict, if kind == 'repeated': if not val: return 0 - elem_ann = _element_annotation(annotation) - elem_kind = _infer_kind(elem_ann, metadata) + elem = spec.elem total = 0 for i, ele in enumerate(val): total += _encode_into_field( - f'{fname}[{i}]', ele, elem_kind, elem_ann, metadata, markers, - wire, offset + total) + f'{fname}[{i}]', ele, elem, markers, wire, offset + total) return total if kind == 'map': if not val: return 0 - key_ann, val_ann = _map_annotations(annotation) - key_meta = _map_key_meta(metadata) - vl_meta = _map_val_meta(metadata) - key_kind = _infer_kind(key_ann, key_meta) - vl_kind = _infer_kind(val_ann, vl_meta) + key_spec, val_spec = spec.key, spec.val total = 0 for i, (k, v) in enumerate(val.items()): total += _encode_into_field( - f'{fname}[{i}#k]', k, key_kind, key_ann, key_meta, markers, - wire, offset + total) + f'{fname}[{i}#k]', k, key_spec, markers, wire, offset + total) total += _encode_into_field( - f'{fname}[{i}#v]', v, vl_kind, val_ann, vl_meta, markers, - wire, offset + total) + f'{fname}[{i}#v]', v, val_spec, markers, wire, offset + total) return total raise TypeError(f'Unknown field kind {kind!r} for {fname!r}') @@ -666,15 +709,9 @@ def _encode_into_field(fname: str, val, kind: str, annotation, metadata: dict, def _encode_into_model(obj, markers: dict, wire: VarBinaryStr, offset: int) -> None: """Write all TLV fields of a dataclass object into *wire* starting at *offset*.""" - cls = type(obj) - hints = typing.get_type_hints(cls) - for f in dataclasses.fields(cls): - ann = hints[f.name] - kind = _infer_kind(ann, f.metadata) - if kind not in _ZERO_WIRE_KINDS and 'tlv_type' not in f.metadata: - continue + for spec in _get_schema(type(obj)): offset += _encode_into_field( - f.name, getattr(obj, f.name), kind, ann, f.metadata, markers, wire, offset) + spec.name, getattr(obj, spec.name), spec, markers, wire, offset) # --------------------------------------------------------------------------- @@ -742,16 +779,14 @@ def _make_default_instance(cls): return obj -def _parse_value(fname: str, kind: str, annotation, metadata: dict, +def _parse_value(fname: str, spec: _FieldSpec, wire, offset: int, length: int, offset_btl: int, ignore_critical: bool): """ Parse a single TLV *value* (V only, not T or L) from *wire*. :param fname: field name (for error messages). - :param kind: field kind string. - :param annotation: resolved Python type annotation. - :param metadata: dataclass field metadata dict. + :param spec: resolved field spec. :param wire: memoryview of the full wire buffer. :param offset: byte offset of V within *wire*. :param length: byte length of V. @@ -760,6 +795,7 @@ def _parse_value(fname: str, kind: str, annotation, metadata: dict, :param ignore_critical: forwarded to nested ``tlv_parse`` calls. :return: the parsed Python value. """ + kind = spec.kind if kind == 'bool': return True @@ -776,12 +812,9 @@ def _parse_value(fname: str, kind: str, annotation, metadata: dict, raise ValueError( f'{fname}: uint value length must be 1, 2, 4, or 8; got {length}') # Auto-convert to the annotated Enum/Flag type if applicable - inner = _unwrap_optional(annotation) - if (isinstance(inner, type) - and issubclass(inner, (Enum, Flag)) - and inner is not int): + if spec.enum_cls is not None: try: - return inner(raw) + return spec.enum_cls(raw) except ValueError: pass return raw @@ -796,9 +829,8 @@ def _parse_value(fname: str, kind: str, annotation, metadata: dict, return Name.decode(wire, offset_btl)[0] if kind == 'model': - inner_cls = _unwrap_optional(annotation) - ignore = metadata.get('ignore_critical', ignore_critical) - return tlv_parse(inner_cls, wire[offset:offset + length], ignore) + ignore = spec.metadata.get('ignore_critical', ignore_critical) + return tlv_parse(spec.annotation, wire[offset:offset + length], ignore) raise TypeError(f'Unknown kind {kind!r} for {fname!r}') @@ -836,14 +868,7 @@ def tlv_parse(cls, wire, ignore_critical: bool = False, markers: dict = None): else: mv = memoryview(wire if isinstance(wire, (bytes, bytearray)) else bytes(wire)) - hints = typing.get_type_hints(cls) - ordered = [] - for f in dataclasses.fields(cls): - ann = hints[f.name] - kind = _infer_kind(ann, f.metadata) - if kind not in _ZERO_WIRE_KINDS and 'tlv_type' not in f.metadata: - continue - ordered.append((f.name, f.metadata, kind, ann)) + ordered = _get_schema(cls) obj = _make_default_instance(cls) offset = 0 @@ -858,23 +883,22 @@ def tlv_parse(cls, wire, ignore_critical: bool = False, markers: dict = None): found = False for i in range(field_pos, len(ordered)): - fname, meta, kind, ann = ordered[i] + spec = ordered[i] + kind = spec.kind if kind == 'offset_marker': continue # never matches a wire TLV type - if meta['tlv_type'] != typ: + if spec.tlv_type != typ: continue + fname = spec.name # Advance any offset_markers between field_pos and i. for j in range(field_pos, i): - jname, _, jkind, _ = ordered[j] - if jkind == 'offset_marker': - markers[jname] = offset_btl + if ordered[j].kind == 'offset_marker': + markers[ordered[j].name] = offset_btl if kind == 'repeated': - elem_ann = _element_annotation(ann) - elem_kind = _infer_kind(elem_ann, meta) - val = _parse_value(fname, elem_kind, elem_ann, meta, + val = _parse_value(fname, spec.elem, mv, offset, length, offset_btl, ignore_critical) lst = getattr(obj, fname) if lst is None: @@ -885,19 +909,13 @@ def tlv_parse(cls, wire, ignore_critical: bool = False, markers: dict = None): elif kind == 'map': # Two-phase parse: consume key, then immediately read value TLV. - key_ann, val_ann = _map_annotations(ann) - key_meta = _map_key_meta(meta) - vl_meta = _map_val_meta(meta) - key_kind = _infer_kind(key_ann, key_meta) - vl_kind = _infer_kind(val_ann, vl_meta) - dct = getattr(obj, fname) if dct is None: dct = {} object.__setattr__(obj, fname, dct) idx = len(dct) - key = _parse_value(f'{fname}[{idx}#k]', key_kind, key_ann, key_meta, + key = _parse_value(f'{fname}[{idx}#k]', spec.key, mv, offset, length, offset_btl, ignore_critical) # advance past key value → now at the value TLV @@ -908,7 +926,7 @@ def tlv_parse(cls, wire, ignore_critical: bool = False, markers: dict = None): length, _sz_l2 = parse_tl_num(mv, offset) offset += _sz_l2 - val = _parse_value(f'{fname}[{idx}#v]', vl_kind, val_ann, vl_meta, + val = _parse_value(f'{fname}[{idx}#v]', spec.val, mv, offset, length, offset_btl, ignore_critical) dct[key] = val field_pos = i # stay at i to accept more pairs @@ -917,7 +935,7 @@ def tlv_parse(cls, wire, ignore_critical: bool = False, markers: dict = None): # Extract sig buffer; append covered region to ##sig_covered_part. sig_buf = mv[offset:offset + length] markers['##sig_value_buf'] = sig_buf - cover_start_field = meta.get('cover_start') + cover_start_field = spec.metadata.get('cover_start') if cover_start_field is not None: cover_start = markers.get(cover_start_field) if cover_start is not None: @@ -939,7 +957,7 @@ def tlv_parse(cls, wire, ignore_critical: bool = False, markers: dict = None): field_pos = i + 1 else: - val = _parse_value(fname, kind, ann, meta, + val = _parse_value(fname, spec, mv, offset, length, offset_btl, ignore_critical) object.__setattr__(obj, fname, val) field_pos = i + 1 diff --git a/src/ndn/transport/nfd_registerer.py b/src/ndn/transport/nfd_registerer.py index 5dd9aea..8c2c3ed 100644 --- a/src/ndn/transport/nfd_registerer.py +++ b/src/ndn/transport/nfd_registerer.py @@ -21,7 +21,7 @@ from .. import security as sec from .. import types from .. import utils -from ..app_support import nfd_mgmt +from ..app_support import nfd_mgmt, nfd_mgmt_2 from .prefix_registerer import PrefixRegisterer @@ -32,6 +32,7 @@ async def pass_all(_name, _sig, _context): class NfdRegister(PrefixRegisterer): _prefix_register_semaphore: aio.Semaphore = None _last_command_timestamp: int = 0 + mgmt = nfd_mgmt def __init__(self): super().__init__() @@ -48,11 +49,11 @@ async def register(self, name: enc.NonStrictName) -> bool: await aio.sleep(0.001) try: _, reply, _ = await self.app.express( - name=nfd_mgmt.make_command_v2('rib', 'register', self.app.face, name=name), + name=self.mgmt.make_command_v2('rib', 'register', self.app.face, name=name), app_param=b'', signer=sec.DigestSha256Signer(for_interest=True), validator=pass_all, lifetime=1000) - ret = nfd_mgmt.parse_response(reply) + ret = self.mgmt.parse_response(reply) if ret['status_code'] != 200: logging.getLogger(__name__).error('Registration for %s failed: %s %s', enc.Name.to_str(name), ret["status_code"], ret["status_text"]) @@ -77,9 +78,13 @@ async def unregister(self, name: enc.NonStrictName) -> bool: await aio.sleep(0.001) try: await self.app.express( - nfd_mgmt.make_command_v2('rib', 'unregister', self.app.face, name=name), + self.mgmt.make_command_v2('rib', 'unregister', self.app.face, name=name), app_param=b'', signer=sec.DigestSha256Signer(for_interest=True), validator=pass_all, lifetime=1000) return True except (types.InterestNack, types.InterestTimeout, types.InterestCanceled, types.ValidationFailure): return False + + +class NfdRegister2(NfdRegister): + mgmt = nfd_mgmt_2 diff --git a/tests/encoding/ndn_format_0_3_2_test.py b/tests/encoding/ndn_format_0_3_2_test.py index 87a186b..0b2ea55 100644 --- a/tests/encoding/ndn_format_0_3_2_test.py +++ b/tests/encoding/ndn_format_0_3_2_test.py @@ -1,6 +1,7 @@ import hashlib from ndn.encoding import Name +from ndn.encoding import ndn_format_0_3 as v1 from ndn.encoding.ndn_format_0_3_2 import ( ContentType, InterestParam, @@ -11,7 +12,7 @@ parse_data, parse_interest, ) -from ndn.security import DigestSha256Signer +from ndn.security import DigestSha256Signer, HmacSha256Signer def test_default_interest_wire_format(): @@ -71,3 +72,15 @@ def test_data_wire_format_and_coverage(): assert content is None signature = hashlib.sha256(b''.join(sig.signature_covered_part)).digest() assert signature == sig.signature_value_buf + + +def test_key_locator_signer_matches_v1(): + signer = HmacSha256Signer('/local/KEY/1', b'secret') + data = make_data('/local/data', MetaInfo(), b'content', signer=signer) + assert data == v1.make_data('/local/data', v1.MetaInfo(), b'content', signer=signer) + _, _, _, sig = parse_data(data) + assert sig.signature_info.key_locator.name == Name.from_str('/local/KEY/1') + + interest = make_interest('/local/int', InterestParam(nonce=1), b'\x01', signer) + assert interest == v1.make_interest( + '/local/int', v1.InterestParam(nonce=1), b'\x01', signer) diff --git a/tests/encoding/tlv_model_v2_test.py b/tests/encoding/tlv_model_v2_test.py index d58ac88..2946267 100644 --- a/tests/encoding/tlv_model_v2_test.py +++ b/tests/encoding/tlv_model_v2_test.py @@ -1233,3 +1233,74 @@ def test_tlv_set_arg_stores_value(self): m = {} tlv_set_arg(m, 'key', 'value') assert tlv_get_arg(m, 'key') == 'value' + + +# --------------------------------------------------------------------------- +# Schema cache +# --------------------------------------------------------------------------- + +@dataclass +class _ForwardOuter: + inner: '_ForwardInner' = field(default=None, metadata={'tlv_type': 0x10}) + + +@dataclass +class _ForwardInner: + val: int = field(default=None, metadata={'tlv_type': 0x01}) + + +class TestSchemaCache: + def test_type_hints_resolved_once_per_class(self, monkeypatch): + import typing + from ndn.encoding import tlv_model_v2 + + @dataclass + class Inner: + v: int = field(default=None, metadata={'tlv_type': 0x01}) + + @dataclass + class Outer: + items: list[Inner] = field(default_factory=list, metadata={'tlv_type': 0x10}) + m: dict[str, bytes] = field(default_factory=dict, metadata={'tlv_type': 0x21, 'val_tlv_type': 0x23}) + + calls = [] + real = typing.get_type_hints + monkeypatch.setattr(tlv_model_v2.typing, 'get_type_hints', + lambda cls, *a, **k: calls.append(cls) or real(cls, *a, **k)) + obj = Outer(items=[Inner(v=1), Inner(v=2)], m={'k': b'v'}) + for _ in range(3): + assert tlv_parse(Outer, tlv_encode(obj)).items[1].v == 2 + assert sorted(c.__name__ for c in calls) == ['Inner', 'Outer'] + + def test_forward_reference_resolved_lazily(self): + wire = tlv_encode(_ForwardOuter(inner=_ForwardInner(val=5))) + assert wire == b'\x10\x03\x01\x01\x05' + assert tlv_parse(_ForwardOuter, wire).inner.val == 5 + + def test_local_class_not_kept_alive(self): + import gc + import weakref + + def make(): + @dataclass + class Local: + x: int = field(default=None, metadata={'tlv_type': 0x01}) + tlv_parse(Local, tlv_encode(Local(x=1))) + return weakref.ref(Local) + + ref = make() + gc.collect() + assert ref() is None + + def test_non_tlv_field_with_unsupported_annotation_ignored(self): + class Opaque: + pass + + @dataclass + class M: + cache: Opaque = None + x: int = field(default=None, metadata={'tlv_type': 0x01}) + + wire = tlv_encode(M(cache=Opaque(), x=3)) + assert wire == b'\x01\x01\x03' + assert tlv_parse(M, wire).x == 3 diff --git a/tests/integration/app_v2_2_test.py b/tests/integration/app_v2_2_test.py new file mode 100644 index 0000000..7242b73 --- /dev/null +++ b/tests/integration/app_v2_2_test.py @@ -0,0 +1,295 @@ +# ----------------------------------------------------------------------------- +# Copyright (C) 2019-2022 The python-ndn authors +# +# This file is part of python-ndn. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# ----------------------------------------------------------------------------- +import abc +import asyncio as aio +import pytest +from ndn import appv2_2 as app +from ndn import security as sec +from ndn import encoding as enc +from ndn import types +from ndn.app_support import nfd_mgmt_2 +from ndn.encoding import ndn_format_0_3_2 as fmt +from ndn.encoding.tlv_model_v2 import tlv_encode, tlv_parse +from ndn.transport.dummy_face import DummyFace +from ndn.transport.nfd_registerer import NfdRegister2 + + +class NDNAppTestSuite: + app = None + signer = None + + def test_main(self): + aio.run(self.comain()) + + async def comain(self): + face = DummyFace(self.face_proc) + self.signer = sec.DigestSha256Signer() + self.app = app.NDNApp(face) + face.app = self.app + await self.app.main_loop(self.app_main()) + + @abc.abstractmethod + async def face_proc(self, face: DummyFace): + pass + + @abc.abstractmethod + async def app_main(self): + pass + + +class TestConsumerBasic(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + await face.consume_output(b'\x050\x07(\x08\x07example\x08\x07testApp\x08\nrandomData' + b'\x38\x08\x00\x00\x01m\xa4\xf3\xffm\x12\x00\x0c\x02\x17p') + await face.input_packet(b'\x06B\x07(\x08\x07example\x08\x07testApp\x08\nrandomData' + b'\x38\x08\x00\x00\x01m\xa4\xf3\xffm\x14\x07\x18\x01\x00\x19\x02\x03\xe8' + b'\x15\rHello, world!') + + async def app_main(self): + name = '/example/testApp/randomData/t=1570430517101' + data_name, content, pkt_context = await self.app.express( + name, app.pass_all, must_be_fresh=True, can_be_prefix=False, lifetime=6000, nonce=None) + assert data_name == enc.Name.from_str(name) + assert pkt_context['meta_info'].freshness_period == 1000 + assert content == b'Hello, world!' + + +class TestInterestCancel(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + await face.consume_output(b'\x05\x15\x07\x0f\x08\rnot important\x0c\x02\x0f\xa0') + + async def app_main(self): + with pytest.raises(types.InterestCanceled): + await self.app.express('not important', app.pass_all, nonce=None) + + +class TestInterestNack(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + await face.consume_output(b'\x05)\x07\x1f\x08\tlocalhost\x08\x03nfd\x08\x05faces\x08\x06events' + b'\x21\x00\x12\x00\x0c\x02\x03\xe8') + await face.input_packet(b'\x64\x36\xfd\x03\x20\x05\xfd\x03\x21\x01\x96' + b'\x50\x2b\x05)\x07\x1f\x08\tlocalhost\x08\x03nfd\x08\x05faces\x08\x06events' + b'\x21\x00\x12\x00\x0c\x02\x03\xe8') + + async def app_main(self): + with pytest.raises(types.InterestNack) as nack: + await self.app.express('/localhost/nfd/faces/events', app.pass_all, nonce=None, lifetime=1000, + must_be_fresh=True, can_be_prefix=True) + assert nack.value.reason == 150 + + +class TestInterestTimeout(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + await face.consume_output(b'\x05\x14\x07\x0f\x08\rnot important\x0c\x01\x0a') + await aio.sleep(0.05) + + async def app_main(self): + with pytest.raises(types.InterestTimeout): + await self.app.express('not important', app.pass_all, nonce=None, lifetime=10) + + +class TestDataValidationFalure(NDNAppTestSuite): + @staticmethod + async def validator(_name, _sig, _context) -> types.ValidResult: + await aio.sleep(0.003) + return types.ValidResult.FAIL + + async def face_proc(self, face: DummyFace): + await face.consume_output(b'\x05\x1b\x07\x10\x08\x03not\x08\timportant\n\x04\x00\x00\x00\x00\x0c\x01\xfa') + await face.input_packet(b'\x06\x1d\x07\x10\x08\x03not\x08\timportant\x14\x03\x18\x01\x00\x15\x04test') + + async def app_main(self): + with pytest.raises(types.ValidationFailure) as e: + await self.app.express('/not/important', validator=self.validator, nonce=0, lifetime=250) + assert e.value.name == enc.Name.from_str('/not/important') + assert e.value.content == b'test' + assert e.value.result == types.ValidResult.FAIL + + +class TestInterestCanBePrefix(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + await face.consume_output(b'\x05\x0a\x07\x05\x08\x03not\x0c\x01\xfa' + b'\x05\x0c\x07\x05\x08\x03not\x21\x00\x0c\x01\xfa' + b'\x05\x15\x07\x10\x08\x03not\x08\timportant\x0c\x01\xfa') + await face.input_packet(b'\x06\x1d\x07\x10\x08\x03not\x08\timportant\x14\x03\x18\x01\x00\x15\x04test') + await aio.sleep(0.4) + + async def app_main(self): + future1 = self.app.express('/not', app.pass_all, nonce=None, lifetime=250, can_be_prefix=False) + future2 = self.app.express('/not', app.pass_all, nonce=None, lifetime=250, can_be_prefix=True) + future3 = self.app.express('/not/important', app.pass_all, nonce=None, lifetime=250, can_be_prefix=False) + name2, content2, _ = await future3 + name1, content1, _ = await future2 + with pytest.raises(types.InterestTimeout): + await future1 + assert name1 == enc.Name.from_str('/not/important') + assert content1 == b'test' + assert name2 == enc.Name.from_str('/not/important') + assert content2 == b'test' + + +class TestRoute(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + await face.ignore_output(0) + await face.input_packet(b'\x05\x15\x07\x10\x08\x03not\x08\timportant\x0c\x01\xfa') + await face.consume_output(b'\x06\x24\x07\x10\x08\x03not\x08\timportant\x14\x03\x18\x01\x00\x15\x04test' + b'\x16\x03\x1b\x01\xc8\x17\x00') + + async def app_main(self): + @self.app.route('/not') + def on_interest(name, _app_param, reply: app.ReplyFunc, _context): + data = self.app.make_data(name, b'test', signer=sec.NullSigner()) + assert reply(data) + + +class TestInvalidInterest(NDNAppTestSuite): + @staticmethod + async def validator(_name, _sig, _context) -> types.ValidResult: + await aio.sleep(0.003) + return types.ValidResult.FAIL + + async def face_proc(self, face: DummyFace): + await face.ignore_output(0) + await face.input_packet(b'\x05`\x072\x08\x03not\x08\timportant' + b'\x02 E\x8a\xeaxI}[\xb1\xcd\xf0\x01\xbe' + b'\xdb\xe9\x03\x085\xb1g+K\xa8jK,\xd0\xad' + b')\x07\x83\x96\xbb\x0c\x01\xfa$\x00,\x03' + b'\x1b\x01\x00. !\x93!zG[%\xcfs\xe89\\\x8f' + b'^\xd3\xa4\xb9\x13\xaa\x7f\xa6?\xd7\x13aVyS\xdc\x1dW\xea') + await aio.sleep(0.005) + + async def app_main(self): + @self.app.route('/not', validator=self.validator) + def on_interest(_name, _app_param, _reply: app.ReplyFunc, _context): + raise ValueError('This test fails') + + +class TestRoute2(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + await face.ignore_output(0) + await face.input_packet(b'\x05\x15\x07\x10\x08\x03not\x08\timportant\x0c\x01\xfa') + await face.consume_output(b'\x06\x24\x07\x10\x08\x03not\x08\timportant\x14\x03\x18\x01\x00\x15\x04test' + b'\x16\x03\x1b\x01\xc8\x17\x00') + + async def app_main(self): + @self.app.route('/not') + def on_interest(name, _app_param, reply: app.ReplyFunc, context): + assert context['raw_packet'] == b'\x05\x15\x07\x10\x08\x03not\x08\timportant\x0c\x01\xfa' + assert not context['sig_ptrs'].signature_info + reply(self.app.make_data(name, b'test', signer=sec.NullSigner())) + + +class TestConsumerRawPacket(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + await face.consume_output(b'\x050\x07(\x08\x07example\x08\x07testApp\x08\nrandomData' + b'\x38\x08\x00\x00\x01m\xa4\xf3\xffm\x12\x00\x0c\x02\x17p') + await face.input_packet(b'\x06B\x07(\x08\x07example\x08\x07testApp\x08\nrandomData' + b'\x38\x08\x00\x00\x01m\xa4\xf3\xffm\x14\x07\x18\x01\x00\x19\x02\x03\xe8' + b'\x15\rHello, world!') + + async def app_main(self): + name = '/example/testApp/randomData/t=1570430517101' + _, _, pkt_context = await self.app.express( + name, validator=app.pass_all, + must_be_fresh=True, can_be_prefix=False, lifetime=6000, nonce=None, need_raw_packet=True) + raw = pkt_context['raw_packet'] + assert (raw == b'\x06\x42\x07(\x08\x07example\x08\x07testApp\x08\nrandomData' + b'\x38\x08\x00\x00\x01m\xa4\xf3\xffm\x14\x07\x18\x01\x00\x19\x02\x03\xe8' + b'\x15\rHello, world!') + + +class TestCongestionMark(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + await face.ignore_output(0) + await face.input_packet(b'\x64\x1e\xfd\x03\x40\x01\x01\x50\x17' + b'\x05\x15\x07\x10\x08\x03not\x08\timportant\x0c\x01\xfa') + await face.consume_output(b'\x06\x24\x07\x10\x08\x03not\x08\timportant\x14\x03\x18\x01\x00\x15\x04test' + b'\x16\x03\x1b\x01\xc8\x17\x00', timeout=0.5) + + async def app_main(self): + @self.app.route('/not') + def on_interest(name, _app_param, reply: app.ReplyFunc, context): + assert context['raw_packet'] == b'\x05\x15\x07\x10\x08\x03not\x08\timportant\x0c\x01\xfa' + assert not context['sig_ptrs'].signature_info + reply(self.app.make_data(name, b'test', signer=sec.NullSigner())) + + +class TestImplicitSha256(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + await face.consume_output(b'\x05\x2d\x07\x28\x08\x04test\x01\x20' + b'\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff' + b'\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff\xff' + b'\x0c\x01\xfa' + b'\x05\x2d\x07\x28\x08\x04test\x01\x20' + b'\x54\x88\xf2\xc1\x1b\x56\x6d\x49\xe9\x90\x4f\xb5\x2a\xa6\xf6\xf9' + b'\xe6\x6a\x95\x41\x68\x10\x9c\xe1\x56\xee\xa2\xc9\x2c\x57\xe4\xc2' + b'\x0c\x01\xfa') + await face.input_packet(b'\x06\x13\x07\x06\x08\x04test\x14\x03\x18\x01\x00\x15\x04test') + await aio.sleep(0.4) + + async def app_main(self): + fut1 = self.app.express( + '/test/sha256digest=FFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFFF', + validator=app.pass_all, nonce=None, lifetime=250) + fut2 = self.app.express( + '/test/sha256digest=5488f2c11b566d49e9904fb52aa6f6f9e66a954168109ce156eea2c92c57e4c2', + validator=app.pass_all, nonce=None, lifetime=250) + name2, content2, _ = await fut2 + with pytest.raises(types.InterestTimeout): + await fut1 + assert name2 == enc.Name.from_str('/test') + assert content2 == b'test' + + +class TestPitToken(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + await face.ignore_output(0) + await face.input_packet(b'\x64\x1f\x62\x04\x01\x02\x03\x04\x50\x17' + b'\x05\x15\x07\x10\x08\x03not\x08\timportant\x0c\x01\xfa') + await face.consume_output(b'\x64\x2e\x62\x04\x01\x02\x03\x04\x50\x26' + b'\x06\x24\x07\x10\x08\x03not\x08\timportant\x14\x03\x18\x01\x00\x15\x04test' + b'\x16\x03\x1b\x01\xc8\x17\x00') + + async def app_main(self): + @self.app.route('/not') + def on_interest(name, _app_param, reply: app.ReplyFunc, _context): + data = self.app.make_data(name, b'test', signer=sec.NullSigner()) + assert reply(data) + + +class TestRegister(NDNAppTestSuite): + async def face_proc(self, face: DummyFace): + async with aio.timeout(1): + while not face.output_buf: + await aio.sleep(0.001) + interest, face.output_buf = face.output_buf, b'' + name, _, app_param, sig = fmt.parse_interest(interest) + assert name[:4] == enc.Name.from_str('/localhost/nfd/rib/register') + cp = tlv_parse(nfd_mgmt_2.ControlParameters, enc.Component.get_value(name[4])) + assert cp.cp.name == enc.Name.from_str('/test/prefix') + assert app_param == b'' + assert sig.signature_info.signature_type == fmt.SignatureType.DIGEST_SHA256 + + response = tlv_encode(nfd_mgmt_2.ControlResponse( + status_code=200, status_text='OK', body=nfd_mgmt_2.ControlParametersValue(name='/test/prefix'))) + content = bytes([0x65, len(response)]) + response + await face.input_packet(fmt.make_data(name, fmt.MetaInfo(), content, signer=sec.NullSigner())) + + async def app_main(self): + assert isinstance(self.app.registerer, NfdRegister2) + assert await self.app.register('/test/prefix') diff --git a/tests/misc/nfd_mgmt_2_test.py b/tests/misc/nfd_mgmt_2_test.py new file mode 100644 index 0000000..5d0858d --- /dev/null +++ b/tests/misc/nfd_mgmt_2_test.py @@ -0,0 +1,83 @@ +from ndn.app_support import nfd_mgmt as v1 +from ndn.app_support import nfd_mgmt_2 as v2 +from ndn.encoding import Name +from ndn.encoding.tlv_model_v2 import tlv_encode, tlv_parse + + +def test_make_command_v2_matches_v1(): + kwargs = dict(name='/example/prefix', face_id=300, origin=65, cost=10, + flags=1, expiration_period=3600000, + face_persistency=v1.FacePersistency.PERMANENT) + assert v2.make_command_v2('rib', 'register', **kwargs) == \ + v1.make_command_v2('rib', 'register', **kwargs) + assert v2.make_command_v2('strategy-choice', 'set', name='/a', strategy='/localhost/nfd/strategy/multicast') == \ + v1.make_command_v2('strategy-choice', 'set', name='/a', strategy='/localhost/nfd/strategy/multicast') + + +def test_make_command_matches_v1(monkeypatch): + for mod in (v1, v2): + monkeypatch.setattr(mod, 'timestamp', lambda: 1234567) + monkeypatch.setattr(mod, 'gen_nonce_64', lambda: 0xdeadbeef) + assert v2.make_command('faces', 'create', uri='udp4://127.0.0.1:6363') == \ + v1.make_command('faces', 'create', uri='udp4://127.0.0.1:6363') + + +def test_parse_response_matches_v1(): + cr = v1.ControlResponse() + cr.status_code = 200 + cr.status_text = 'OK' + cr.body = v1.ControlParametersValue() + cr.body.name = '/example' + cr.body.face_id = 5 + cr.body.uri = 'udp4://1.2.3.4:6363' + cr.body.face_persistency = v1.FacePersistency.ON_DEMAND + body = bytes(cr.encode()) + wire = bytes([0x65, len(body)]) + body + + ret = v2.parse_response(wire) + assert ret == v1.parse_response(wire) + assert ret['face_persistency'] is v2.FacePersistency.ON_DEMAND + assert Name.to_str(ret['name']) == '/example' + + +def test_parse_response_without_body(): + body = tlv_encode(v2.ControlResponse(status_code=404, status_text='Not found')) + ret = v2.parse_response(bytes([0x65, len(body)]) + body) + assert ret['status_code'] == 404 + assert ret['status_text'] == 'Not found' + assert ret['face_id'] is None + + +def test_face_status_interop(): + status = v1.FaceStatus() + status.face_id = 1 + status.uri = 'internal://' + status.face_scope = v1.FaceScope.LOCAL + status.link_type = v1.FaceLinkType.POINT_TO_POINT + status.flags = v1.FaceFlags.LOCAL_FIELDS_ENABLED | v1.FaceFlags.LP_RELIABILITY_ENABLED + status.n_in_bytes = 2 ** 40 + msg = v1.FaceStatusMsg() + msg.face_status = [status] + wire = bytes(msg.encode()) + + parsed = tlv_parse(v2.FaceStatusMsg, wire) + assert len(parsed.face_status) == 1 + fs = parsed.face_status[0] + assert fs.uri == 'internal://' + assert fs.face_scope is v2.FaceScope.LOCAL + assert fs.flags == v2.FaceFlags.LOCAL_FIELDS_ENABLED | v2.FaceFlags.LP_RELIABILITY_ENABLED + assert fs.n_in_bytes == 2 ** 40 + assert bytes(tlv_encode(parsed)) == wire + + +def test_rib_status_interop(): + rib = v2.RibStatus(entries=[ + v2.RibEntry(name='/a', routes=[v2.Route(face_id=1, origin=0, cost=0, + flags=v2.RouteFlags.CHILD_INHERIT)]), + v2.RibEntry(name='/b', routes=[]), + ]) + wire = bytes(tlv_encode(rib)) + parsed = v1.RibStatus.parse(wire) + assert [Name.to_str(e.name) for e in parsed.entries] == ['/a', '/b'] + assert parsed.entries[0].routes[0].flags == v1.RouteFlags.CHILD_INHERIT + assert bytes(parsed.encode()) == wire From 82d67ba36c0e3eaa8af7ceeef6539400d7233189 Mon Sep 17 00:00:00 2001 From: Xinyu Ma Date: Sun, 27 Sep 2026 12:57:06 -0700 Subject: [PATCH 7/8] Migrate examples and CLIs into new model --- examples/appv2/basic_packets/consumer.py | 6 +++--- examples/appv2/basic_packets/producer.py | 6 +++--- examples/appv2/forwarding_hint/consumer.py | 8 ++++---- examples/appv2/forwarding_hint/producer.py | 8 ++++---- examples/appv2/keychain_cert/fetch_certificate.py | 8 ++++---- examples/appv2/keychain_cert/keychain_register.py | 4 ++-- examples/dpdk_experimental/udp_consumer.py | 6 +++--- examples/dpdk_experimental/udp_producer.py | 6 +++--- src/ndn/bin/nfdc/cmd_get_face.py | 15 ++++++++------- src/ndn/bin/nfdc/cmd_get_route.py | 9 +++++---- src/ndn/bin/nfdc/cmd_get_status.py | 7 ++++--- src/ndn/bin/nfdc/cmd_get_strategy.py | 7 ++++--- src/ndn/bin/nfdc/cmd_new_face.py | 6 +++--- src/ndn/bin/nfdc/cmd_new_route.py | 4 ++-- src/ndn/bin/nfdc/cmd_remove_face.py | 11 ++++++----- src/ndn/bin/nfdc/cmd_remove_route.py | 4 ++-- src/ndn/bin/nfdc/cmd_remove_strategy.py | 4 ++-- src/ndn/bin/nfdc/cmd_set_strategy.py | 4 ++-- src/ndn/bin/nfdc/utils.py | 2 +- src/ndn/bin/sec/cmd_get_childitem.py | 2 +- src/ndn/bin/sec/cmd_get_signreq.py | 2 +- src/ndn/bin/sec/cmd_import_cert.py | 2 +- src/ndn/bin/sec/cmd_sign_cert.py | 2 +- src/ndn/bin/sec/utils.py | 2 +- src/ndn/security/tpm/tpm_osx_keychain.py | 4 ++-- 25 files changed, 72 insertions(+), 67 deletions(-) diff --git a/examples/appv2/basic_packets/consumer.py b/examples/appv2/basic_packets/consumer.py index cd3a0cd..4696988 100644 --- a/examples/appv2/basic_packets/consumer.py +++ b/examples/appv2/basic_packets/consumer.py @@ -16,7 +16,7 @@ # limitations under the License. # ----------------------------------------------------------------------------- import logging -from ndn import utils, appv2, types +from ndn import utils, appv2_2, types from ndn import encoding as enc @@ -26,7 +26,7 @@ style='{') -app = appv2.NDNApp() +app = appv2_2.NDNApp() async def main(): @@ -36,7 +36,7 @@ async def main(): print(f'Sending Interest {enc.Name.to_str(name)}, {enc.InterestParam(must_be_fresh=True, lifetime=6000)}') # TODO: Write a better validator data_name, content, pkt_context = await app.express( - name, validator=appv2.pass_all, + name, validator=appv2_2.pass_all, must_be_fresh=True, can_be_prefix=False, lifetime=6000) print(f'Received Data Name: {enc.Name.to_str(data_name)}') diff --git a/examples/appv2/basic_packets/producer.py b/examples/appv2/basic_packets/producer.py index ae4851a..e462a85 100644 --- a/examples/appv2/basic_packets/producer.py +++ b/examples/appv2/basic_packets/producer.py @@ -16,7 +16,7 @@ # limitations under the License. # ----------------------------------------------------------------------------- import logging -from ndn import appv2 +from ndn import appv2_2 from ndn import encoding as enc @@ -26,13 +26,13 @@ style='{') -app = appv2.NDNApp() +app = appv2_2.NDNApp() keychain = app.default_keychain() @app.route('/example/testApp') def on_interest(name: enc.FormalName, _app_param: enc.BinaryStr | None, - reply: appv2.ReplyFunc, context: appv2.PktContext): + reply: appv2_2.ReplyFunc, context: appv2_2.PktContext): print(f'>> I: {enc.Name.to_str(name)}, {context["int_param"]}') content = b"Hello, world!" reply(app.make_data(name, content=content, signer=keychain.get_signer({}), diff --git a/examples/appv2/forwarding_hint/consumer.py b/examples/appv2/forwarding_hint/consumer.py index 9fdbb78..490a437 100644 --- a/examples/appv2/forwarding_hint/consumer.py +++ b/examples/appv2/forwarding_hint/consumer.py @@ -16,7 +16,7 @@ # limitations under the License. # ----------------------------------------------------------------------------- import logging -from ndn import utils, appv2, types +from ndn import utils, appv2_2, types from ndn import encoding as enc @@ -26,7 +26,7 @@ style='{') -app = appv2.NDNApp() +app = appv2_2.NDNApp() async def express_int(name, fw_hint): @@ -34,13 +34,13 @@ async def express_int(name, fw_hint): if fw_hint is None: print(f'Sending Interest {enc.Name.to_str(name)}, {enc.InterestParam(must_be_fresh=True, lifetime=6000)}') data_name, content, pkt_context = await app.express( - name, validator=appv2.pass_all, + name, validator=appv2_2.pass_all, must_be_fresh=True, can_be_prefix=False, lifetime=6000) else: print(f'Sending Interest {enc.Name.to_str(name)}, ' f'{enc.InterestParam(must_be_fresh=True, lifetime=6000, forwarding_hint=[fw_hint])}') data_name, content, pkt_context = await app.express( - name, validator=appv2.pass_all, + name, validator=appv2_2.pass_all, must_be_fresh=True, can_be_prefix=False, lifetime=6000, forwarding_hint=[fw_hint]) print(f'Received Data Name: {enc.Name.to_str(data_name)}') diff --git a/examples/appv2/forwarding_hint/producer.py b/examples/appv2/forwarding_hint/producer.py index dcd0b66..823e7c6 100644 --- a/examples/appv2/forwarding_hint/producer.py +++ b/examples/appv2/forwarding_hint/producer.py @@ -1,5 +1,5 @@ import logging -from ndn import appv2 +from ndn import appv2_2 from ndn import encoding as enc @@ -9,13 +9,13 @@ style='{') -app = appv2.NDNApp() +app = appv2_2.NDNApp() keychain = app.default_keychain() @app.route('/repo/command') def on_cmd(name: enc.FormalName, _app_param: enc.BinaryStr | None, - reply: appv2.ReplyFunc, context: appv2.PktContext): + reply: appv2_2.ReplyFunc, context: appv2_2.PktContext): print(f'>> I: {enc.Name.to_str(name)}, {context["int_param"]}') content = b"Hello, world!" reply(app.make_data(name, content=content, signer=keychain.get_signer({}), @@ -30,7 +30,7 @@ def on_cmd(name: enc.FormalName, _app_param: enc.BinaryStr | None, # So we can dispatch by forwarding hints. @app.route('/') def on_fwd_hint(name: enc.FormalName, app_param: enc.BinaryStr | None, - reply: appv2.ReplyFunc, context: appv2.PktContext): + reply: appv2_2.ReplyFunc, context: appv2_2.PktContext): fwd_hints = context["int_param"].forwarding_hint if fwd_hints: fh_name = fwd_hints[0] diff --git a/examples/appv2/keychain_cert/fetch_certificate.py b/examples/appv2/keychain_cert/fetch_certificate.py index 3d17f15..9eb51af 100644 --- a/examples/appv2/keychain_cert/fetch_certificate.py +++ b/examples/appv2/keychain_cert/fetch_certificate.py @@ -17,9 +17,9 @@ # ----------------------------------------------------------------------------- import sys import logging -from ndn import appv2, types +from ndn import appv2_2, types from ndn import encoding as enc -from ndn.app_support import security_v2 as secv2 +from ndn.app_support import security_v2_2 as secv2 logging.basicConfig(format='[{asctime}]{levelname}:{message}', @@ -32,7 +32,7 @@ logging.fatal('Please input a KEY or CERT name') exit(0) -app = appv2.NDNApp() +app = appv2_2.NDNApp() async def main(): @@ -43,7 +43,7 @@ async def main(): f'{enc.InterestParam(must_be_fresh=True, can_be_prefix=can_be_prefix, lifetime=6000)}') # TODO: Write a better validator data_name, content, pkt_context = await app.express( - name, validator=appv2.pass_all, + name, validator=appv2_2.pass_all, must_be_fresh=True, can_be_prefix=can_be_prefix, lifetime=6000) print(f'Received Data Name: {enc.Name.to_str(data_name)}') diff --git a/examples/appv2/keychain_cert/keychain_register.py b/examples/appv2/keychain_cert/keychain_register.py index ce0b376..300b68c 100644 --- a/examples/appv2/keychain_cert/keychain_register.py +++ b/examples/appv2/keychain_cert/keychain_register.py @@ -16,7 +16,7 @@ # limitations under the License. # ----------------------------------------------------------------------------- import logging -from ndn import appv2 +from ndn import appv2_2 from ndn.app_support.keychain_register import attach_keychain_register @@ -26,7 +26,7 @@ style='{') -app = appv2.NDNApp() +app = appv2_2.NDNApp() keychain = app.default_keychain() attach_keychain_register(keychain, app) diff --git a/examples/dpdk_experimental/udp_consumer.py b/examples/dpdk_experimental/udp_consumer.py index cba3ae3..c21a6b3 100644 --- a/examples/dpdk_experimental/udp_consumer.py +++ b/examples/dpdk_experimental/udp_consumer.py @@ -17,7 +17,7 @@ # ----------------------------------------------------------------------------- import logging import sys -from ndn import utils, appv2, types +from ndn import utils, appv2_2, types from ndn import encoding as enc from ndn.transport.ndn_dpdk import NdnDpdkUdpFace, DpdkRegisterer @@ -42,7 +42,7 @@ face = NdnDpdkUdpFace(gql_url, self_addr, self_port, dpdk_addr, dpdk_port) registerer = DpdkRegisterer(face) -app = appv2.NDNApp(face=face, registerer=registerer) +app = appv2_2.NDNApp(face=face, registerer=registerer) keychain = app.default_keychain() @@ -53,7 +53,7 @@ async def main(): print(f'Sending Interest {enc.Name.to_str(name)}, {enc.InterestParam(must_be_fresh=True, lifetime=6000)}') # TODO: Write a better validator data_name, content, pkt_context = await app.express( - name, validator=appv2.pass_all, + name, validator=appv2_2.pass_all, must_be_fresh=True, can_be_prefix=False, lifetime=6000) print(f'Received Data Name: {enc.Name.to_str(data_name)}') diff --git a/examples/dpdk_experimental/udp_producer.py b/examples/dpdk_experimental/udp_producer.py index 4a0ef8d..c3b5e77 100644 --- a/examples/dpdk_experimental/udp_producer.py +++ b/examples/dpdk_experimental/udp_producer.py @@ -17,7 +17,7 @@ # ----------------------------------------------------------------------------- import logging import sys -from ndn import appv2 +from ndn import appv2_2 from ndn import encoding as enc from ndn.transport.ndn_dpdk import NdnDpdkUdpFace, DpdkRegisterer @@ -42,13 +42,13 @@ face = NdnDpdkUdpFace(gql_url, self_addr, self_port, dpdk_addr, dpdk_port) registerer = DpdkRegisterer(face) -app = appv2.NDNApp(face=face, registerer=registerer) +app = appv2_2.NDNApp(face=face, registerer=registerer) keychain = app.default_keychain() @app.route('/example/testApp') def on_interest(name: enc.FormalName, _app_param: enc.BinaryStr | None, - reply: appv2.ReplyFunc, context: appv2.PktContext): + reply: appv2_2.ReplyFunc, context: appv2_2.PktContext): print(f'>> I: {enc.Name.to_str(name)}, {context["int_param"]}') content = b"Hello, world!" reply(app.make_data(name, content=content, signer=keychain.get_signer({}), diff --git a/src/ndn/bin/nfdc/cmd_get_face.py b/src/ndn/bin/nfdc/cmd_get_face.py index 2da16c3..643f3ec 100644 --- a/src/ndn/bin/nfdc/cmd_get_face.py +++ b/src/ndn/bin/nfdc/cmd_get_face.py @@ -16,9 +16,10 @@ # limitations under the License. # ----------------------------------------------------------------------------- import argparse -from ...appv2 import NDNApp +from ...appv2_2 import NDNApp from ...encoding import Name, Component -from ...app_support.nfd_mgmt import FaceStatusMsg, FaceQueryFilter, FaceQueryFilterValue, parse_response +from ...encoding.tlv_model_v2 import tlv_encode, tlv_parse +from ...app_support.nfd_mgmt_2 import FaceStatusMsg, FaceQueryFilter, FaceQueryFilterValue, parse_response from .utils import express_interest @@ -35,7 +36,7 @@ def execute(args: argparse.Namespace): async def list_face(): try: data = await express_interest(app, "/localhost/nfd/faces/list") - msg = FaceStatusMsg.parse(data) + msg = tlv_parse(FaceStatusMsg, data) # TODO: Should calculate the length instead of using a fixed number print(f'{"FaceID":7}{"RemoteURI":<30}\t{"LocalURI":<30}') print(f'{"------":7}{"---------":<30}\t{"--------":<30}') @@ -53,7 +54,7 @@ async def exec_query(): msg = parse_response(data) print('Query failed with response', msg['status_code'], msg['status_text']) else: - msg = FaceStatusMsg.parse(data) + msg = tlv_parse(FaceStatusMsg, data) for f in msg.face_status: print() print(f'{"Face ID":>12}\t{f.face_id}') @@ -78,16 +79,16 @@ async def exec_query(): filt.face_query_filter = FaceQueryFilterValue() if face_id is not None: filt.face_query_filter.face_id = face_id - data_name = Name.from_str(name) + [Component.from_bytes(filt.encode())] + data_name = Name.from_str(name) + [Component.from_bytes(tlv_encode(filt))] if not await exec_query(): print('No face is found') else: filt.face_query_filter.uri = face_uri - data_name = Name.from_str(name) + [Component.from_bytes(filt.encode())] + data_name = Name.from_str(name) + [Component.from_bytes(tlv_encode(filt))] if not await exec_query(): filt.face_query_filter.uri = None filt.face_query_filter.local_uri = face_uri - data_name = Name.from_str(name) + [Component.from_bytes(filt.encode())] + data_name = Name.from_str(name) + [Component.from_bytes(tlv_encode(filt))] if not await exec_query(): print('No face is found') app.shutdown() diff --git a/src/ndn/bin/nfdc/cmd_get_route.py b/src/ndn/bin/nfdc/cmd_get_route.py index 91acebe..3fe6b3d 100644 --- a/src/ndn/bin/nfdc/cmd_get_route.py +++ b/src/ndn/bin/nfdc/cmd_get_route.py @@ -16,9 +16,10 @@ # limitations under the License. # ----------------------------------------------------------------------------- import argparse -from ...appv2 import NDNApp +from ...appv2_2 import NDNApp from ...encoding import Name -from ...app_support.nfd_mgmt import FibStatus, RibStatus +from ...encoding.tlv_model_v2 import tlv_parse +from ...app_support.nfd_mgmt_2 import FibStatus, RibStatus from .utils import express_interest @@ -36,9 +37,9 @@ def execute(args: argparse.Namespace): async def list_route(): try: fib_data = await express_interest(app, "/localhost/nfd/fib/list") - fib_msg = FibStatus.parse(fib_data) + fib_msg = tlv_parse(FibStatus, fib_data) rib_data = await express_interest(app, "/localhost/nfd/rib/list") - rib_msg = RibStatus.parse(rib_data) + rib_msg = tlv_parse(RibStatus, rib_data) # TODO: Should calculate the length instead of using a fixed number print('Forwarding Table (FIB)') for ent in fib_msg.entries: diff --git a/src/ndn/bin/nfdc/cmd_get_status.py b/src/ndn/bin/nfdc/cmd_get_status.py index e1273d1..9a57fc7 100644 --- a/src/ndn/bin/nfdc/cmd_get_status.py +++ b/src/ndn/bin/nfdc/cmd_get_status.py @@ -17,8 +17,9 @@ # ----------------------------------------------------------------------------- import argparse import datetime -from ...appv2 import NDNApp -from ...app_support.nfd_mgmt import GeneralStatus +from ...appv2_2 import NDNApp +from ...encoding.tlv_model_v2 import tlv_parse +from ...app_support.nfd_mgmt_2 import GeneralStatus from .utils import express_interest @@ -34,7 +35,7 @@ async def after_start(): try: data = await express_interest(app, "/localhost/nfd/status/general") - msg = GeneralStatus.parse(data) + msg = tlv_parse(GeneralStatus, data) print('General status:') print(f'{"version":>25}\t{msg.nfd_version}') diff --git a/src/ndn/bin/nfdc/cmd_get_strategy.py b/src/ndn/bin/nfdc/cmd_get_strategy.py index 03c0f57..c106665 100644 --- a/src/ndn/bin/nfdc/cmd_get_strategy.py +++ b/src/ndn/bin/nfdc/cmd_get_strategy.py @@ -16,9 +16,10 @@ # limitations under the License. # ----------------------------------------------------------------------------- import argparse -from ...appv2 import NDNApp +from ...appv2_2 import NDNApp from ...encoding import Name -from ...app_support.nfd_mgmt import StrategyChoiceMsg +from ...encoding.tlv_model_v2 import tlv_parse +from ...app_support.nfd_mgmt_2 import StrategyChoiceMsg from .utils import express_interest @@ -36,7 +37,7 @@ def execute(args: argparse.Namespace): async def list_strategy(): try: data = await express_interest(app, "/localhost/nfd/strategy-choice/list") - msg = StrategyChoiceMsg.parse(data) + msg = tlv_parse(StrategyChoiceMsg, data) for s in msg.strategy_choices: s_prefix = Name.to_str(s.name) if prefix and s_prefix != prefix: diff --git a/src/ndn/bin/nfdc/cmd_new_face.py b/src/ndn/bin/nfdc/cmd_new_face.py index f622365..de1b3ed 100644 --- a/src/ndn/bin/nfdc/cmd_new_face.py +++ b/src/ndn/bin/nfdc/cmd_new_face.py @@ -16,8 +16,8 @@ # limitations under the License. # ----------------------------------------------------------------------------- import argparse -from ...appv2 import NDNApp -from ...app_support.nfd_mgmt import parse_response, make_command_v2 +from ...appv2_2 import NDNApp +from ...app_support.nfd_mgmt_2 import parse_response, make_command_v2 from .utils import express_interest @@ -38,7 +38,7 @@ def execute(args: argparse.Namespace): uri = uri + ":6363" async def create_face(): - cmd = make_command_v2('faces', 'create', uri=uri.encode()) + cmd = make_command_v2('faces', 'create', uri=uri) res = await express_interest(app, cmd) msg = parse_response(res) print(f'{msg["status_code"]} {msg["status_text"]}') diff --git a/src/ndn/bin/nfdc/cmd_new_route.py b/src/ndn/bin/nfdc/cmd_new_route.py index b929fe2..8cb4c37 100644 --- a/src/ndn/bin/nfdc/cmd_new_route.py +++ b/src/ndn/bin/nfdc/cmd_new_route.py @@ -16,8 +16,8 @@ # limitations under the License. # ----------------------------------------------------------------------------- import argparse -from ...appv2 import NDNApp -from ...app_support.nfd_mgmt import make_command_v2, parse_response +from ...appv2_2 import NDNApp +from ...app_support.nfd_mgmt_2 import make_command_v2, parse_response from .utils import express_interest diff --git a/src/ndn/bin/nfdc/cmd_remove_face.py b/src/ndn/bin/nfdc/cmd_remove_face.py index b615071..a7b66bf 100644 --- a/src/ndn/bin/nfdc/cmd_remove_face.py +++ b/src/ndn/bin/nfdc/cmd_remove_face.py @@ -16,9 +16,10 @@ # limitations under the License. # ----------------------------------------------------------------------------- import argparse -from ...appv2 import NDNApp +from ...appv2_2 import NDNApp from ...encoding import Name, Component -from ...app_support.nfd_mgmt import FaceStatusMsg, FaceQueryFilter, FaceQueryFilterValue, parse_response, \ +from ...encoding.tlv_model_v2 import tlv_encode, tlv_parse +from ...app_support.nfd_mgmt_2 import FaceStatusMsg, FaceQueryFilter, FaceQueryFilterValue, parse_response, \ make_command_v2 from .utils import express_interest @@ -56,7 +57,7 @@ async def try_remove(): msg = parse_response(data) print('Query failed with response', msg['status_code'], msg['status_text']) else: - msg = FaceStatusMsg.parse(data) + msg = tlv_parse(FaceStatusMsg, data) for f in msg.face_status: await remove_face(f.face_id) return True @@ -66,11 +67,11 @@ async def try_remove(): filt = FaceQueryFilter() filt.face_query_filter = FaceQueryFilterValue() filt.face_query_filter.uri = uri - data_name = Name.from_str(name) + [Component.from_bytes(filt.encode())] + data_name = Name.from_str(name) + [Component.from_bytes(tlv_encode(filt))] if not await try_remove(): filt.face_query_filter.uri = None filt.face_query_filter.local_uri = uri - data_name = Name.from_str(name) + [Component.from_bytes(filt.encode())] + data_name = Name.from_str(name) + [Component.from_bytes(tlv_encode(filt))] if not await try_remove(): print('No face is found') finally: diff --git a/src/ndn/bin/nfdc/cmd_remove_route.py b/src/ndn/bin/nfdc/cmd_remove_route.py index 6a44375..a7535df 100644 --- a/src/ndn/bin/nfdc/cmd_remove_route.py +++ b/src/ndn/bin/nfdc/cmd_remove_route.py @@ -16,8 +16,8 @@ # limitations under the License. # ----------------------------------------------------------------------------- import argparse -from ...appv2 import NDNApp -from ...app_support.nfd_mgmt import parse_response, make_command_v2 +from ...appv2_2 import NDNApp +from ...app_support.nfd_mgmt_2 import parse_response, make_command_v2 from .utils import express_interest diff --git a/src/ndn/bin/nfdc/cmd_remove_strategy.py b/src/ndn/bin/nfdc/cmd_remove_strategy.py index 3fafeca..6fc7eda 100644 --- a/src/ndn/bin/nfdc/cmd_remove_strategy.py +++ b/src/ndn/bin/nfdc/cmd_remove_strategy.py @@ -16,8 +16,8 @@ # limitations under the License. # ----------------------------------------------------------------------------- import argparse -from ...appv2 import NDNApp -from ...app_support.nfd_mgmt import parse_response, make_command_v2 +from ...appv2_2 import NDNApp +from ...app_support.nfd_mgmt_2 import parse_response, make_command_v2 from .utils import express_interest diff --git a/src/ndn/bin/nfdc/cmd_set_strategy.py b/src/ndn/bin/nfdc/cmd_set_strategy.py index d617690..dfbea92 100644 --- a/src/ndn/bin/nfdc/cmd_set_strategy.py +++ b/src/ndn/bin/nfdc/cmd_set_strategy.py @@ -16,8 +16,8 @@ # limitations under the License. # ----------------------------------------------------------------------------- import argparse -from ...appv2 import NDNApp -from ...app_support.nfd_mgmt import parse_response, make_command_v2 +from ...appv2_2 import NDNApp +from ...app_support.nfd_mgmt_2 import parse_response, make_command_v2 from .utils import express_interest diff --git a/src/ndn/bin/nfdc/utils.py b/src/ndn/bin/nfdc/utils.py index ce086de..9615d5e 100644 --- a/src/ndn/bin/nfdc/utils.py +++ b/src/ndn/bin/nfdc/utils.py @@ -15,7 +15,7 @@ # See the License for the specific language governing permissions and # limitations under the License. # ----------------------------------------------------------------------------- -from ...appv2 import NDNApp, pass_all +from ...appv2_2 import NDNApp, pass_all from ...security import DigestSha256Signer from ...types import InterestNack, InterestTimeout, InterestCanceled, ValidationFailure diff --git a/src/ndn/bin/sec/cmd_get_childitem.py b/src/ndn/bin/sec/cmd_get_childitem.py index 82c35fd..cf78752 100644 --- a/src/ndn/bin/sec/cmd_get_childitem.py +++ b/src/ndn/bin/sec/cmd_get_childitem.py @@ -18,7 +18,7 @@ import argparse from base64 import standard_b64encode from ...encoding import Name, SignatureType -from ...app_support.security_v2 import parse_certificate +from ...app_support.security_v2_2 import parse_certificate from .utils import resolve_keychain diff --git a/src/ndn/bin/sec/cmd_get_signreq.py b/src/ndn/bin/sec/cmd_get_signreq.py index 6b3c333..ac5ce3e 100644 --- a/src/ndn/bin/sec/cmd_get_signreq.py +++ b/src/ndn/bin/sec/cmd_get_signreq.py @@ -18,7 +18,7 @@ import argparse import base64 from ...encoding import Name -from ...app_support.security_v2 import sign_req +from ...app_support.security_v2_2 import sign_req from .utils import resolve_keychain, infer_obj_name diff --git a/src/ndn/bin/sec/cmd_import_cert.py b/src/ndn/bin/sec/cmd_import_cert.py index 50e9356..8e9b2dc 100644 --- a/src/ndn/bin/sec/cmd_import_cert.py +++ b/src/ndn/bin/sec/cmd_import_cert.py @@ -20,7 +20,7 @@ import sys import argparse from ...encoding import Name -from ...app_support.security_v2 import parse_certificate +from ...app_support.security_v2_2 import parse_certificate from .utils import resolve_keychain diff --git a/src/ndn/bin/sec/cmd_sign_cert.py b/src/ndn/bin/sec/cmd_sign_cert.py index b693d28..aef9a44 100644 --- a/src/ndn/bin/sec/cmd_sign_cert.py +++ b/src/ndn/bin/sec/cmd_sign_cert.py @@ -21,7 +21,7 @@ import argparse from datetime import datetime, timedelta, UTC from ...encoding import Name -from ...app_support.security_v2 import parse_certificate, new_cert +from ...app_support.security_v2_2 import parse_certificate, new_cert from .utils import resolve_keychain, infer_obj_name diff --git a/src/ndn/bin/sec/utils.py b/src/ndn/bin/sec/utils.py index 0f99409..232123f 100644 --- a/src/ndn/bin/sec/utils.py +++ b/src/ndn/bin/sec/utils.py @@ -21,7 +21,7 @@ from ...platform import Platform from ...security import KeychainSqlite3 from ...client_conf import default_keychain -from ...app_support.security_v2 import KEY_COMPONENT +from ...app_support.security_v2_2 import KEY_COMPONENT from ...encoding import Name, FormalName diff --git a/src/ndn/security/tpm/tpm_osx_keychain.py b/src/ndn/security/tpm/tpm_osx_keychain.py index 97170fd..f35eee3 100644 --- a/src/ndn/security/tpm/tpm_osx_keychain.py +++ b/src/ndn/security/tpm/tpm_osx_keychain.py @@ -94,9 +94,9 @@ def _get_key(key_name: NonStrictName): g.dic = c_void_p() ret = sec.security.SecItemCopyMatching(g.query, pointer(g.dic)) if ret == sec.errSecItemNotFound: - raise KeyError(f"Unable to find key {key_name}") + raise KeyError(f"Unable to find key {Name.to_str(key_name)}") elif ret != sec.errSecSuccess: - raise RuntimeError(f"Error happened when searching specific key {key_name}") + raise RuntimeError(f"Error happened when searching specific key {Name.to_str(key_name)}") key_type = cfstring_to_string(cf.CFDictionaryGetValue(g.dic, sec.kSecAttrKeyType)) key_bits = cfnumber_to_number(cf.CFDictionaryGetValue(g.dic, sec.kSecAttrKeySizeInBits)) From 6d1314c97076194984b32c5b0af5dfce9c1530af Mon Sep 17 00:00:00 2001 From: Xinyu Ma Date: Mon, 28 Sep 2026 23:37:41 -0700 Subject: [PATCH 8/8] Fix review comments --- src/ndn/appv2.py | 7 +++++-- src/ndn/appv2_2.py | 11 +++++++++-- src/ndn/encoding/tlv_model_v2.py | 11 +++++++++++ 3 files changed, 25 insertions(+), 4 deletions(-) diff --git a/src/ndn/appv2.py b/src/ndn/appv2.py index e320459..b1c0da7 100644 --- a/src/ndn/appv2.py +++ b/src/ndn/appv2.py @@ -364,6 +364,7 @@ def reply(data: enc.BinaryStr) -> bool: self._put_raw_packet(data) else: self._put_raw_packet_with_pit_token(data, pit_token) + return True # In case the validator blocks the pipeline, create a task async def submit_interest(): @@ -437,8 +438,10 @@ def _put_raw_packet_with_pit_token_nocopy(self, data: enc.BinaryStr, pit_token: pt.pit_token = pit_token pt_wire = pt.encode() frag_l = len(data) - lp_l = len(pt_wire) + enc.get_tl_num_size(ndnlp.LpTypeNumber.FRAGMENT) + enc.get_tl_num_size(frag_l) - wire_l = enc.get_tl_num_size(ndnlp.LpTypeNumber.LP_PACKET) + enc.get_tl_num_size(lp_l) + lp_l + frag_header_l = (enc.get_tl_num_size(ndnlp.LpTypeNumber.FRAGMENT) + enc.get_tl_num_size(frag_l)) + lp_l = len(pt_wire) + frag_header_l + frag_l + wire_l = (enc.get_tl_num_size(ndnlp.LpTypeNumber.LP_PACKET) + + enc.get_tl_num_size(lp_l) + len(pt_wire) + frag_header_l) wire = bytearray(wire_l) pos = 0 pos += enc.write_tl_num(ndnlp.LpTypeNumber.LP_PACKET, wire, pos) diff --git a/src/ndn/appv2_2.py b/src/ndn/appv2_2.py index b988043..a2ae215 100644 --- a/src/ndn/appv2_2.py +++ b/src/ndn/appv2_2.py @@ -260,6 +260,10 @@ async def _receive(self, typ: int, data: enc.BinaryStr): nack_reason = None pit_token = lp_pkt.pit_token data = lp_pkt.fragment + if data is None: + # Only Nack and Data can reach this function. + self.logger.fatal('LP packet without a fragment reaching _receive branch. Unexpected behavior.') + return typ, _ = enc.parse_tl_num(data) else: nack_reason = None @@ -367,6 +371,7 @@ def reply(data: enc.BinaryStr) -> bool: self._put_raw_packet(data) else: self._put_raw_packet_with_pit_token(data, pit_token) + return True # In case the validator blocks the pipeline, create a task async def submit_interest(): @@ -435,8 +440,10 @@ def _put_raw_packet_with_pit_token_nocopy(self, data: enc.BinaryStr, pit_token: raise types.NetworkError('cannot send packet before connected') pt_wire = tlv_encode(ndnlp.LpPacketValue(pit_token=pit_token)) frag_l = len(data) - lp_l = len(pt_wire) + enc.get_tl_num_size(ndnlp.LpTypeNumber.FRAGMENT) + enc.get_tl_num_size(frag_l) - wire_l = enc.get_tl_num_size(ndnlp.LpTypeNumber.LP_PACKET) + enc.get_tl_num_size(lp_l) + lp_l + frag_header_l = (enc.get_tl_num_size(ndnlp.LpTypeNumber.FRAGMENT) + enc.get_tl_num_size(frag_l)) + lp_l = len(pt_wire) + frag_header_l + frag_l + wire_l = (enc.get_tl_num_size(ndnlp.LpTypeNumber.LP_PACKET) + + enc.get_tl_num_size(lp_l) + len(pt_wire) + frag_header_l) wire = bytearray(wire_l) pos = 0 pos += enc.write_tl_num(ndnlp.LpTypeNumber.LP_PACKET, wire, pos) diff --git a/src/ndn/encoding/tlv_model_v2.py b/src/ndn/encoding/tlv_model_v2.py index da696cd..f9232bf 100644 --- a/src/ndn/encoding/tlv_model_v2.py +++ b/src/ndn/encoding/tlv_model_v2.py @@ -450,6 +450,8 @@ def _finalize_encode(markers: dict, mv: memoryview, model_end: int) -> int: def _uint_value_len(val: int, fname: str, fixed_len) -> int: if fixed_len is not None: + if fixed_len not in (1, 2, 4, 8): + raise ValueError("uint fixed_len must be 1, 2, 4, or 8") n = fixed_len elif val <= 0xFF: n = 1 @@ -544,6 +546,8 @@ def _encoded_length_field(fname: str, val, spec: _FieldSpec, markers: dict) -> i return total_with_tl if kind == 'model': + if not isinstance(val, spec.annotation): + raise TypeError(f'{fname}={val!r} is not of type {spec.annotation!r}') inner_markers: dict = {} length = _encoded_length_model(val, inner_markers) markers[f'{fname}##inner_markers'] = inner_markers @@ -880,6 +884,8 @@ def tlv_parse(cls, wire, ignore_critical: bool = False, markers: dict = None): offset += sz_t length, sz_l = parse_tl_num(mv, offset) offset += sz_l + if length > len(mv) - offset: + raise IndexError('TLV length exceeds the input buffer') found = False for i in range(field_pos, len(ordered)): @@ -925,6 +931,11 @@ def tlv_parse(cls, wire, ignore_critical: bool = False, markers: dict = None): offset += _sz_t2 length, _sz_l2 = parse_tl_num(mv, offset) offset += _sz_l2 + if _val_typ != spec.val.tlv_type: + raise DecodeError( + f'{fname}: expected map value type {spec.val.tlv_type:#x}, got {_val_typ:#x}') + if length > len(mv) - offset: + raise IndexError('map value length exceeds the input buffer') val = _parse_value(f'{fname}[{idx}#v]', spec.val, mv, offset, length, offset_btl, ignore_critical)