diff --git a/src/adn_server/application/routing/lc_ta.py b/src/adn_server/application/routing/lc_ta.py index fd2a9e3..3ad23d7 100644 --- a/src/adn_server/application/routing/lc_ta.py +++ b/src/adn_server/application/routing/lc_ta.py @@ -49,6 +49,7 @@ from typing import Any from bitarray import bitarray from ...domain import int_id +from ...domain.dmr import bptc from ...domain.dmr.const import LC_OPT_G, LC_OPT_U from ...domain.talker_alias import DMRA_BLOCK_COUNT from ..talker_alias_use_cases import passthrough_complete, talker_alias_settings @@ -57,9 +58,43 @@ from .helpers import EMB_LC_SLICE logger = logging.getLogger(__name__) +# Distinct destination LCs kept encoded; a call start needs one per (TG, source radio). +_LC_SET_CACHE_MAX = 512 + + class LcTaMixin: """Talker Alias DMRA relay and embedded LC overlay on forward legs.""" + def _encode_lc_set(self, dst_lc: bytes) -> tuple[bitarray, bitarray, dict[int, Any]]: + """Header LC, terminator LC and embedded LC for one destination LC. + + About 250us of BPTC/RS work, done once per LC instead of once per target + leg: every leg opened for the same TG and radio gets the same encoding. + The results are shared between legs, and every reader only slices them. + """ + cache: dict[bytes, tuple[bitarray, bitarray, dict[int, Any]]] | None = getattr( + self, "_lc_set_cache", None + ) + if cache is None: + cache = self._lc_set_cache = {} + if not isinstance(dst_lc, bytes): # unhashable (bytearray): encode, don't keep + return ( + bptc.encode_header_lc(dst_lc), + bptc.encode_terminator_lc(dst_lc), + self._encode_emblc(dst_lc), + ) + codes = cache.get(dst_lc) + if codes is None: + codes = ( + bptc.encode_header_lc(dst_lc), + bptc.encode_terminator_lc(dst_lc), + self._encode_emblc(dst_lc), + ) + if len(cache) >= _LC_SET_CACHE_MAX: + cache.clear() + cache[dst_lc] = codes + return codes + def _get_stream_dmra_blocks(self, source_system: str, stream_id: bytes) -> dict[int, bytes] | None: if not self._get_dmra_blocks: return None diff --git a/src/adn_server/application/routing/obp_forward.py b/src/adn_server/application/routing/obp_forward.py index 8b7c17d..4277136 100644 --- a/src/adn_server/application/routing/obp_forward.py +++ b/src/adn_server/application/routing/obp_forward.py @@ -52,6 +52,7 @@ from typing import Any from ...domain import HBPF_DATA_SYNC, HBPF_SLT_VHEAD, int_id from ...domain.dmr import decode from ...domain.dmr.const import LC_OPT +from ..subscription.obp_source_ops import ensure_obp_source_for_tg_store, obp_source_needs_ensure from ..subscription.subscription_queries import store_has_table from .helpers import group_voice_tg_ingress_collision, obp_is_canonical_ingress, unit_data_reportable @@ -86,11 +87,6 @@ class ObpForwardMixin: return if not (79 <= dst_int < 9990 or dst_int > 9999): return - from ..subscription.obp_source_ops import ( - ensure_obp_source_for_tg_store, - obp_source_needs_ensure, - ) - store = self._subscription_store if not obp_source_needs_ensure(store, system_name, relay_table_key, dst_int): return diff --git a/src/adn_server/application/routing_use_cases.py b/src/adn_server/application/routing_use_cases.py index f5a50c3..c9c2265 100644 --- a/src/adn_server/application/routing_use_cases.py +++ b/src/adn_server/application/routing_use_cases.py @@ -585,9 +585,11 @@ class RoutingUseCases( logger.exception("(to_target) caught exception") _target_status[stream_id]["LAST"] = pkt_time return - _target_status[stream_id]["H_LC"] = bptc.encode_header_lc(dst_lc) - _target_status[stream_id]["T_LC"] = bptc.encode_terminator_lc(dst_lc) - _target_status[stream_id]["EMB_LC"] = self._encode_emblc(dst_lc) + ( + _target_status[stream_id]["H_LC"], + _target_status[stream_id]["T_LC"], + _target_status[stream_id]["EMB_LC"], + ) = self._encode_lc_set(dst_lc) self._init_talker_alias_embed( _target_status[stream_id], system_name, @@ -798,9 +800,11 @@ class RoutingUseCases( _bridge_tx_leg["TX_STREAM_ID"] = stream_id _bridge_tx_leg["TX_RFS"] = rf_src _bridge_tx_leg["TX_PEER"] = peer_id - _bridge_tx_leg["TX_H_LC"] = bptc.encode_header_lc(dst_lc) - _bridge_tx_leg["TX_T_LC"] = bptc.encode_terminator_lc(dst_lc) - _bridge_tx_leg["TX_EMB_LC"] = self._encode_emblc(dst_lc) + ( + _bridge_tx_leg["TX_H_LC"], + _bridge_tx_leg["TX_T_LC"], + _bridge_tx_leg["TX_EMB_LC"], + ) = self._encode_lc_set(dst_lc) if obp_flat_bridge_tx_idle(_ts_st, pkt_time) or _ts_st.get("TX_STREAM_ID") == stream_id: obp_publish_flat_bridge_tx(_ts_st, _bridge_tx_leg) self._dispatch_talker_alias_on_bridge_open( @@ -819,9 +823,7 @@ class RoutingUseCases( _ts_st["TX_STREAM_ID"] = stream_id _ts_st["TX_RFS"] = rf_src _ts_st["TX_PEER"] = peer_id - _ts_st["TX_H_LC"] = bptc.encode_header_lc(dst_lc) - _ts_st["TX_T_LC"] = bptc.encode_terminator_lc(dst_lc) - _ts_st["TX_EMB_LC"] = self._encode_emblc(dst_lc) + _ts_st["TX_H_LC"], _ts_st["TX_T_LC"], _ts_st["TX_EMB_LC"] = self._encode_lc_set(dst_lc) self._dispatch_talker_alias_on_bridge_open( _ts_st, system_name, diff --git a/src/adn_server/application/subscription/in_band_signalling_ops.py b/src/adn_server/application/subscription/in_band_signalling_ops.py index 87b9c2f..ba8677b 100644 --- a/src/adn_server/application/subscription/in_band_signalling_ops.py +++ b/src/adn_server/application/subscription/in_band_signalling_ops.py @@ -30,7 +30,7 @@ from adn_server.application.routing.helpers import is_special_tg from adn_server.application.subscription.routing_table_export import _legacy_to_type from adn_server.application.subscription.trigger_bytes import dst_in_triggers from adn_server.domain import bytes_3, int_id -from adn_server.domain.subscription import ActivationPolicy, Subscription, SubscriptionPhase +from adn_server.domain.subscription import ActivationPolicy, Subscription, SubscriptionPhase, SystemId logger = logging.getLogger(__name__) @@ -52,10 +52,7 @@ def apply_in_band_signalling_store( dst_group = int_id(dst_id) dst_id_b = dst_id if isinstance(dst_id, bytes) and len(dst_id) >= 3 else bytes_3(dst_group) - for sub in store.snapshot(): - if sub.system.value != system_name: - continue - + for sub in store.list_by_system(SystemId(system_name)): relay_table_key = sub.table_key() if relay_table_key[:1] == "#" and dst_group != 9: continue diff --git a/src/adn_server/application/subscription/obp_source_ops.py b/src/adn_server/application/subscription/obp_source_ops.py index e029c2b..2371f33 100644 --- a/src/adn_server/application/subscription/obp_source_ops.py +++ b/src/adn_server/application/subscription/obp_source_ops.py @@ -58,7 +58,7 @@ def obp_source_needs_ensure( active_tables = set(store.relay_tables_with_active_source(system_name, 1, dst_int)) pending_keys: list[str] = [] for key in (relay_table_key, "#" + relay_table_key): - if store.legs_in_table(key): + if store.has_table(key): pending_keys.append(key) if not pending_keys: return False @@ -75,7 +75,7 @@ def ensure_obp_source_for_tg_store( ) -> None: """Ensure OBP has ACTIVE TS1 source row in main and #reflector tables.""" for key in (relay_table_key, "#" + relay_table_key): - if not any(sub.table_key() == key for sub in store.snapshot()): + if not store.has_table(key): continue channel_tgid = dst_int patched = False diff --git a/src/adn_server/application/subscription/subscription_table_ops.py b/src/adn_server/application/subscription/subscription_table_ops.py index 1446dfa..6f5cd05 100644 --- a/src/adn_server/application/subscription/subscription_table_ops.py +++ b/src/adn_server/application/subscription/subscription_table_ops.py @@ -54,13 +54,12 @@ def _effective_tmout_minutes(tgid_int: int, tmout: float) -> float: def _table_has_legs(store: SubscriptionStore, table_key: str) -> bool: - return any(sub.table_key() == table_key for sub in store.snapshot()) + return store.has_table(table_key) def _remove_table(store: SubscriptionStore, table_key: str) -> None: - for sub in list(store.snapshot()): - if sub.table_key() == table_key: - store.remove(sub.subscription_id) + for sub in store.legs_in_table(table_key): + store.remove(sub.subscription_id) def _find_leg( diff --git a/src/adn_server/infrastructure/subscription_store.py b/src/adn_server/infrastructure/subscription_store.py index 6a7a90c..23984b1 100644 --- a/src/adn_server/infrastructure/subscription_store.py +++ b/src/adn_server/infrastructure/subscription_store.py @@ -42,6 +42,8 @@ class InMemorySubscriptionStore(SubscriptionStore): def __init__(self) -> None: self._items: dict[SubscriptionId, Subscription] = {} + # Per system, in the same order as _items (an upsert keeps its place). + self._by_system: dict[str, dict[SubscriptionId, Subscription]] = {} self._by_table: dict[str, list[Subscription]] = defaultdict(list) self._source_tables: dict[_IndexKey, set[str]] = {} self._active_target_counts: dict[_IndexKey, int] = {} @@ -63,6 +65,7 @@ class InMemorySubscriptionStore(SubscriptionStore): if old is not None: self._unindex(old) self._items[subscription.subscription_id] = subscription + self._by_system.setdefault(subscription.system.value, {})[subscription.subscription_id] = subscription self._index(subscription) self._revision += 1 @@ -70,12 +73,18 @@ class InMemorySubscriptionStore(SubscriptionStore): old = self._items.pop(sub_id, None) if old is None: return False + of_system = self._by_system.get(old.system.value) + if of_system is not None: + of_system.pop(sub_id, None) + if not of_system: + del self._by_system[old.system.value] self._unindex(old) self._revision += 1 return True def clear(self) -> None: self._items.clear() + self._by_system.clear() self._by_table.clear() self._source_tables.clear() self._active_target_counts.clear() @@ -86,6 +95,7 @@ class InMemorySubscriptionStore(SubscriptionStore): self.clear() for sub in subscriptions: self._items[sub.subscription_id] = sub + self._by_system.setdefault(sub.system.value, {})[sub.subscription_id] = sub self._index(sub) self._revision += 1 @@ -96,7 +106,7 @@ class InMemorySubscriptionStore(SubscriptionStore): return tuple(sub for sub in self._items.values() if sub.channel == channel) def list_by_system(self, system: SystemId) -> tuple[Subscription, ...]: - return tuple(sub for sub in self._items.values() if sub.system == system) + return tuple(self._by_system.get(system.value, {}).values()) def list_active(self) -> tuple[Subscription, ...]: return tuple(sub for sub in self._items.values() if sub.is_active()) diff --git a/tests/routing/test_forward_plan_cache.py b/tests/routing/test_forward_plan_cache.py index 12199ca..2724284 100644 --- a/tests/routing/test_forward_plan_cache.py +++ b/tests/routing/test_forward_plan_cache.py @@ -33,7 +33,8 @@ from tests.harness.deterministic import ( ) from adn_server.application.subscription.routing_table_import import subscriptions_from_routing_table -from adn_server.domain.subscription import SubscriptionPhase +from adn_server.domain.dmr import bptc +from adn_server.domain.subscription import SubscriptionPhase, SystemId from adn_server.infrastructure.subscription_store import InMemorySubscriptionStore TG = 213 @@ -155,3 +156,55 @@ def test_removing_a_missing_subscription_keeps_the_revision() -> None: before = store.revision assert store.remove(subs[0].subscription_id) is False assert store.revision == before + + +def test_list_by_system_keeps_the_order_of_a_full_scan() -> None: + table = active_routing_table(TG, (("MASTER-A", 2), ("OBP-0", 1), ("MASTER-B", 2))) + table |= active_routing_table(91, (("MASTER-A", 1), ("OBP-0", 1))) + table |= active_routing_table(92, (("MASTER-A", 2),)) + subs = subscriptions_from_routing_table(table) + store = InMemorySubscriptionStore() + + def agrees() -> None: + for system in ("MASTER-A", "OBP-0", "MASTER-B", "NONE"): + sid = SystemId(system) + scan = tuple(s for s in store.snapshot() if s.system == sid) + assert store.list_by_system(sid) == scan, system + + store.replace_all(subs) + agrees() + moved = next(s for s in subs if s.system.value == "MASTER-A") + moved.state.phase = SubscriptionPhase.IDLE + store.upsert(moved) # an upsert keeps its place + agrees() + store.remove(moved.subscription_id) + agrees() + store.upsert(moved) # back in, now last + agrees() + for sub in store.list_by_system(SystemId("OBP-0")): + store.remove(sub.subscription_id) + agrees() + store.clear() + agrees() + + +def test_legs_of_one_call_share_one_lc_encoding() -> None: + sc = _scenario() + _call(sc, 0x2002, bursts=1) + rows = [sc.protocols[name].STATUS[(0x2002).to_bytes(4, "big")] for name in OBPS] + assert rows[0]["H_LC"] is rows[1]["H_LC"] + assert rows[0]["EMB_LC"] is rows[1]["EMB_LC"] + lc = b"\x00\x00\x20" + TG.to_bytes(3, "big") + PacketSpec(dst_id=TG).rf_src.to_bytes(3, "big") + assert rows[0]["H_LC"] == bptc.encode_header_lc(lc) + assert rows[0]["T_LC"] == bptc.encode_terminator_lc(lc) + assert rows[0]["EMB_LC"] == bptc.encode_emblc(lc) + + +def test_lc_set_of_a_bytearray_is_encoded_without_caching() -> None: + sc = _scenario() + lc = bytearray(b"\x00\x00\x20" + TG.to_bytes(3, "big") + (2130035).to_bytes(3, "big")) + header, terminator, emb = sc.routing._encode_lc_set(lc) + assert header == bptc.encode_header_lc(bytes(lc)) + assert terminator == bptc.encode_terminator_lc(bytes(lc)) + assert emb == bptc.encode_emblc(bytes(lc)) + assert not getattr(sc.routing, "_lc_set_cache", {})