diff --git a/PasarGuardNodeBridge/__init__.py b/PasarGuardNodeBridge/__init__.py index 5f385e6..d21a2a7 100644 --- a/PasarGuardNodeBridge/__init__.py +++ b/PasarGuardNodeBridge/__init__.py @@ -12,10 +12,10 @@ - Extensible with custom metadata via the `extra` argument Author: PasarGuard -Version: 0.9.0 +Version: 0.10.0 """ -__version__ = "0.9.0" +__version__ = "0.10.0" __author__ = "PasarGuard" @@ -31,12 +31,20 @@ InMemoryNodeRegistry, InMemoryUserSyncStore, LifecycleLease, + LifecycleLeaseLostError, LifecycleOperation, LifecycleStatus, NodeConfig, NodeLifecycleCoordinatorProtocol, NodeLifecycleState, NodeRegistryProtocol, + RevocationAwareUserSyncStoreProtocol, + StartupUserSyncLease, + UserRevocationConflictError, + UserRevocationResult, + UserSyncLease, + UserSyncLeaseLostError, + UserSyncStoreFullError, UserSyncStoreProtocol, ) from PasarGuardNodeBridge.utils import create_proxy, create_user @@ -53,6 +61,8 @@ def create_node( port: int, server_ca: str, api_key: str, + api_port: int | None = None, + max_message_size: int | None = None, **kwargs, ) -> PasarGuardNode: """ @@ -67,6 +77,10 @@ def create_node( port (int): Port number used to connect to the node. server_ca (str): The server's SSL certificate as a string (PEM format). api_key (str): API key used for authentication with the node. + api_port (int | None): Port for the maintenance JSON API. Defaults to + ``port`` for backwards compatibility with shared-port deployments. + max_message_size (int | None): Maximum gRPC message size. Ignored for + REST nodes. **kwargs: Additional optional arguments: - name (str): Node instance name for logging. Defaults to "default". - extra (dict): Optional dictionary to pass custom metadata or configuration. Defaults to {}. @@ -112,12 +126,16 @@ def create_node( HTTP CONNECT and SOCKS proxy schemes. """ + resolved_api_port = port if api_port is None else api_port + if connection is NodeType.grpc: return GrpcNode( address=address, port=port, + api_port=resolved_api_port, server_ca=server_ca, api_key=api_key, + max_message_size=max_message_size, **kwargs, ) @@ -125,6 +143,7 @@ def create_node( return RestNode( address=address, port=port, + api_port=resolved_api_port, server_ca=server_ca, api_key=api_key, **kwargs, @@ -167,14 +186,22 @@ async def create_node_from_registry( "create_node_from_config", "InMemoryUserSyncStore", "InMemoryNodeRegistry", + "UserSyncStoreFullError", "UserSyncStoreProtocol", "NodeRegistryProtocol", "NodeConfig", "ClaimedUser", + "UserSyncLease", + "UserSyncLeaseLostError", + "UserRevocationConflictError", + "UserRevocationResult", + "RevocationAwareUserSyncStoreProtocol", + "StartupUserSyncLease", "NodeLifecycleState", "NodeLifecycleCoordinatorProtocol", "LifecycleStatus", "LifecycleOperation", "LifecycleLease", + "LifecycleLeaseLostError", "InMemoryNodeLifecycleCoordinator", ] diff --git a/PasarGuardNodeBridge/abstract_node.py b/PasarGuardNodeBridge/abstract_node.py index fcab062..82bc660 100644 --- a/PasarGuardNodeBridge/abstract_node.py +++ b/PasarGuardNodeBridge/abstract_node.py @@ -17,6 +17,7 @@ async def start( keep_alive: int = 0, exclude_inbounds: list[str] = [], timeout: int | None = None, + reconcile_user_sync: bool = False, ) -> service.BaseInfoResponse | None: raise NotImplementedError @@ -58,13 +59,30 @@ async def get_user_online_ip_list( @abstractmethod async def sync_users( - self, users: list[service.User], flush_pending: bool = False, timeout: int | None = None + self, + users: list[service.User], + flush_pending: bool = False, + timeout: int | None = None, + revocation_id: str | None = None, ) -> service.Empty | None: raise NotImplementedError + async def reconcile_users( + self, + users: list[service.User], + flush_pending: bool = False, + timeout: int | None = None, + ) -> service.Empty | None: + raise NodeAPIError(501, "This node transport does not support authoritative user reconciliation") + @abstractmethod async def sync_users_chunked( - self, users: list[service.User], chunk_size: int = 100, flush_pending: bool = False, timeout: int | None = None + self, + users: list[service.User], + chunk_size: int = 100, + flush_pending: bool = False, + timeout: int | None = None, + revocation_id: str | None = None, ) -> list[service.User]: raise NotImplementedError @@ -114,7 +132,7 @@ async def _check_node_health(self): raise NotImplementedError @abstractmethod - async def _sync_batch_users(self, users: list[service.User]) -> list[service.User]: + async def _sync_batch_users(self, users: list[service.User], user_sync_epoch: int = 0) -> list[service.User]: """Sync a batch of users individually. Returns list of failed users to requeue.""" raise NotImplementedError diff --git a/PasarGuardNodeBridge/aiohttp_compat.py b/PasarGuardNodeBridge/aiohttp_compat.py index fa39713..9c74cf8 100644 --- a/PasarGuardNodeBridge/aiohttp_compat.py +++ b/PasarGuardNodeBridge/aiohttp_compat.py @@ -29,7 +29,7 @@ def json(self) -> Any: return json.loads(self.text) def raise_for_status(self) -> None: - if 400 <= self.status_code: + if 300 <= self.status_code: raise BufferedStatusError(self) @@ -90,10 +90,15 @@ async def _get_session(self) -> aiohttp.ClientSession: return self._session def request(self, *args, **kwargs): + # aiohttp follows redirects by default and preserves custom headers such + # as x-api-key across origins. Node API calls must stay pinned to the + # configured origin. + kwargs["allow_redirects"] = False return _LazyRequestContext(self, args, kwargs) async def get(self, *args, **kwargs) -> aiohttp.ClientResponse: session = await self._get_session() + kwargs["allow_redirects"] = False return await session.get(*args, **kwargs) async def close(self) -> None: diff --git a/PasarGuardNodeBridge/common/service.proto b/PasarGuardNodeBridge/common/service.proto index 738e7b9..a147444 100644 --- a/PasarGuardNodeBridge/common/service.proto +++ b/PasarGuardNodeBridge/common/service.proto @@ -11,6 +11,8 @@ message BaseInfoResponse { bool started = 1; string core_version = 2; string node_version = 3; + bool user_sync_epoch_supported = 4; + uint64 user_sync_epoch = 5; } enum BackendType { @@ -24,6 +26,7 @@ message Backend { repeated User users = 3; uint64 keep_alive = 4; repeated string exclude_inbounds = 5; + uint64 user_sync_epoch = 6; } // log @@ -150,16 +153,19 @@ message User { string email = 1; Proxy proxies = 2; repeated string inbounds = 3; + uint64 user_sync_epoch = 4; } message Users { repeated User users = 1; + uint64 user_sync_epoch = 2; } message UsersChunk { repeated User users = 1; uint64 index = 2; bool last = 3; + uint64 user_sync_epoch = 4; } // Routing (mirrors xray app/router/command, node-friendly shapes) diff --git a/PasarGuardNodeBridge/common/service_pb2.py b/PasarGuardNodeBridge/common/service_pb2.py index f35327c..208e212 100644 --- a/PasarGuardNodeBridge/common/service_pb2.py +++ b/PasarGuardNodeBridge/common/service_pb2.py @@ -24,7 +24,7 @@ -DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n)PasarGuardNodeBridge/common/service.proto\x12\x07service\"\x07\n\x05\x45mpty\"O\n\x10\x42\x61seInfoResponse\x12\x0f\n\x07started\x18\x01 \x01(\x08\x12\x14\n\x0c\x63ore_version\x18\x02 \x01(\t\x12\x14\n\x0cnode_version\x18\x03 \x01(\t\"\x89\x01\n\x07\x42\x61\x63kend\x12\"\n\x04type\x18\x01 \x01(\x0e\x32\x14.service.BackendType\x12\x0e\n\x06\x63onfig\x18\x02 \x01(\t\x12\x1c\n\x05users\x18\x03 \x03(\x0b\x32\r.service.User\x12\x12\n\nkeep_alive\x18\x04 \x01(\x04\x12\x18\n\x10\x65xclude_inbounds\x18\x05 \x03(\t\"\x15\n\x03Log\x12\x0e\n\x06\x64\x65tail\x18\x01 \x01(\t\"?\n\x04Stat\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x0c\n\x04type\x18\x02 \x01(\t\x12\x0c\n\x04link\x18\x03 \x01(\t\x12\r\n\x05value\x18\x04 \x01(\x03\",\n\x0cStatResponse\x12\x1c\n\x05stats\x18\x01 \x03(\x0b\x32\r.service.Stat\"K\n\x0bStatRequest\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\r\n\x05reset\x18\x02 \x01(\x08\x12\x1f\n\x04type\x18\x03 \x01(\x0e\x32\x11.service.StatType\"1\n\x12OnlineStatResponse\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x03\"\x8f\x01\n\x19StatsOnlineIpListResponse\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x38\n\x03ips\x18\x02 \x03(\x0b\x32+.service.StatsOnlineIpListResponse.IpsEntry\x1a*\n\x08IpsEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x03:\x02\x38\x01\"\x82\x01\n\x07Latency\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\r\n\x05\x61live\x18\x02 \x01(\x08\x12\r\n\x05\x64\x65lay\x18\x03 \x01(\x03\x12\x0c\n\x04link\x18\x04 \x01(\t\x12\x16\n\x0elast_seen_time\x18\x05 \x01(\x03\x12\x15\n\rlast_try_time\x18\x06 \x01(\x03\x12\x0e\n\x06source\x18\x07 \x01(\t\"\x1e\n\x0eLatencyRequest\x12\x0c\n\x04name\x18\x01 \x01(\t\"6\n\x0fLatencyResponse\x12#\n\tlatencies\x18\x01 \x03(\x0b\x32\x10.service.Latency\"\xcc\x01\n\x14\x42\x61\x63kendStatsResponse\x12\x15\n\rnum_goroutine\x18\x01 \x01(\r\x12\x0e\n\x06num_gc\x18\x02 \x01(\r\x12\r\n\x05\x61lloc\x18\x03 \x01(\x04\x12\x13\n\x0btotal_alloc\x18\x04 \x01(\x04\x12\x0b\n\x03sys\x18\x05 \x01(\x04\x12\x0f\n\x07mallocs\x18\x06 \x01(\x04\x12\r\n\x05\x66rees\x18\x07 \x01(\x04\x12\x14\n\x0clive_objects\x18\x08 \x01(\x04\x12\x16\n\x0epause_total_ns\x18\t \x01(\x04\x12\x0e\n\x06uptime\x18\n \x01(\r\"\xb4\x01\n\x13SystemStatsResponse\x12\x11\n\tmem_total\x18\x01 \x01(\x04\x12\x10\n\x08mem_used\x18\x02 \x01(\x04\x12\x11\n\tcpu_cores\x18\x03 \x01(\x04\x12\x11\n\tcpu_usage\x18\x04 \x01(\x01\x12 \n\x18incoming_bandwidth_speed\x18\x05 \x01(\x04\x12 \n\x18outgoing_bandwidth_speed\x18\x06 \x01(\x04\x12\x0e\n\x06uptime\x18\x07 \x01(\x04\"\x13\n\x05Vmess\x12\n\n\x02id\x18\x01 \x01(\t\"!\n\x05Vless\x12\n\n\x02id\x18\x01 \x01(\t\x12\x0c\n\x04\x66low\x18\x02 \x01(\t\"\x1a\n\x06Trojan\x12\x10\n\x08password\x18\x01 \x01(\t\"/\n\x0bShadowsocks\x12\x10\n\x08password\x18\x01 \x01(\t\x12\x0e\n\x06method\x18\x02 \x01(\t\"1\n\tWireguard\x12\x12\n\npublic_key\x18\x01 \x01(\t\x12\x10\n\x08peer_ips\x18\x02 \x03(\t\"\x18\n\x08Hysteria\x12\x0c\n\x04\x61uth\x18\x01 \x01(\t\"\xdd\x01\n\x05Proxy\x12\x1d\n\x05vmess\x18\x01 \x01(\x0b\x32\x0e.service.Vmess\x12\x1d\n\x05vless\x18\x02 \x01(\x0b\x32\x0e.service.Vless\x12\x1f\n\x06trojan\x18\x03 \x01(\x0b\x32\x0f.service.Trojan\x12)\n\x0bshadowsocks\x18\x04 \x01(\x0b\x32\x14.service.Shadowsocks\x12%\n\twireguard\x18\x05 \x01(\x0b\x32\x12.service.Wireguard\x12#\n\x08hysteria\x18\x06 \x01(\x0b\x32\x11.service.Hysteria\"H\n\x04User\x12\r\n\x05\x65mail\x18\x01 \x01(\t\x12\x1f\n\x07proxies\x18\x02 \x01(\x0b\x32\x0e.service.Proxy\x12\x10\n\x08inbounds\x18\x03 \x03(\t\"%\n\x05Users\x12\x1c\n\x05users\x18\x01 \x03(\x0b\x32\r.service.User\"G\n\nUsersChunk\x12\x1c\n\x05users\x18\x01 \x03(\x0b\x32\r.service.User\x12\r\n\x05index\x18\x02 \x01(\x04\x12\x0c\n\x04last\x18\x03 \x01(\x08\"5\n\x0bRoutingRule\x12\x14\n\x0coutbound_tag\x18\x01 \x01(\t\x12\x10\n\x08rule_tag\x18\x02 \x01(\t\";\n\x14RoutingRulesResponse\x12#\n\x05rules\x18\x01 \x03(\x0b\x32\x14.service.RoutingRule\"\"\n\x13\x42\x61lancerInfoRequest\x12\x0b\n\x03tag\x18\x01 \x01(\t\"I\n\x14\x42\x61lancerInfoResponse\x12\x17\n\x0foverride_target\x18\x01 \x01(\t\x12\x18\n\x10principle_target\x18\x02 \x03(\t\"\xba\x02\n\x10TestRouteRequest\x12\x13\n\x0binbound_tag\x18\x01 \x01(\t\x12\x0f\n\x07network\x18\x02 \x01(\t\x12\x11\n\ttarget_ip\x18\x03 \x01(\t\x12\x15\n\rtarget_domain\x18\x04 \x01(\t\x12\x13\n\x0btarget_port\x18\x05 \x01(\r\x12\x10\n\x08protocol\x18\x06 \x01(\t\x12\x0c\n\x04user\x18\x07 \x01(\t\x12=\n\nattributes\x18\x08 \x03(\x0b\x32).service.TestRouteRequest.AttributesEntry\x12\x17\n\x0f\x66ield_selectors\x18\t \x03(\t\x12\x16\n\x0epublish_result\x18\n \x01(\x08\x1a\x31\n\x0f\x41ttributesEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\t:\x02\x38\x01\"}\n\x0bRouteResult\x12\x14\n\x0coutbound_tag\x18\x01 \x01(\t\x12\x1b\n\x13outbound_group_tags\x18\x02 \x03(\t\x12\x13\n\x0binbound_tag\x18\x03 \x01(\t\x12\x0f\n\x07network\x18\x04 \x01(\t\x12\x15\n\rtarget_domain\x18\x05 \x01(\t\";\n\x15\x41\x64\x64RoutingRuleRequest\x12\x0c\n\x04rule\x18\x01 \x01(\t\x12\x14\n\x0cshould_reset\x18\x02 \x01(\x08\",\n\x18RemoveRoutingRuleRequest\x12\x10\n\x08rule_tag\x18\x01 \x01(\t\"E\n\x1dOverrideBalancerTargetRequest\x12\x14\n\x0c\x62\x61lancer_tag\x18\x01 \x01(\t\x12\x0e\n\x06target\x18\x02 \x01(\t*&\n\x0b\x42\x61\x63kendType\x12\x08\n\x04XRAY\x10\x00\x12\r\n\tWIREGUARD\x10\x01*_\n\x08StatType\x12\r\n\tOutbounds\x10\x00\x12\x0c\n\x08Outbound\x10\x01\x12\x0c\n\x08Inbounds\x10\x02\x12\x0b\n\x07Inbound\x10\x03\x12\r\n\tUsersStat\x10\x04\x12\x0c\n\x08UserStat\x10\x05\x32\xdc\t\n\x0bNodeService\x12\x36\n\x05Start\x12\x10.service.Backend\x1a\x19.service.BaseInfoResponse\"\x00\x12(\n\x04Stop\x12\x0e.service.Empty\x1a\x0e.service.Empty\"\x00\x12:\n\x0bGetBaseInfo\x12\x0e.service.Empty\x1a\x19.service.BaseInfoResponse\"\x00\x12+\n\x07GetLogs\x12\x0e.service.Empty\x1a\x0c.service.Log\"\x00\x30\x01\x12@\n\x0eGetSystemStats\x12\x0e.service.Empty\x1a\x1c.service.SystemStatsResponse\"\x00\x12\x42\n\x0fGetBackendStats\x12\x0e.service.Empty\x1a\x1d.service.BackendStatsResponse\"\x00\x12\x39\n\x08GetStats\x12\x14.service.StatRequest\x1a\x15.service.StatResponse\"\x00\x12J\n\x13GetOutboundsLatency\x12\x17.service.LatencyRequest\x1a\x18.service.LatencyResponse\"\x00\x12I\n\x12GetUserOnlineStats\x12\x14.service.StatRequest\x1a\x1b.service.OnlineStatResponse\"\x00\x12V\n\x18GetUserOnlineIpListStats\x12\x14.service.StatRequest\x1a\".service.StatsOnlineIpListResponse\"\x00\x12-\n\x08SyncUser\x12\r.service.User\x1a\x0e.service.Empty\"\x00(\x01\x12-\n\tSyncUsers\x12\x0e.service.Users\x1a\x0e.service.Empty\"\x00\x12;\n\x10SyncUsersChunked\x12\x13.service.UsersChunk\x1a\x0e.service.Empty\"\x00(\x01\x12\x43\n\x10ListRoutingRules\x12\x0e.service.Empty\x1a\x1d.service.RoutingRulesResponse\"\x00\x12P\n\x0fGetBalancerInfo\x12\x1c.service.BalancerInfoRequest\x1a\x1d.service.BalancerInfoResponse\"\x00\x12>\n\tTestRoute\x12\x19.service.TestRouteRequest\x1a\x14.service.RouteResult\"\x00\x12\x42\n\x0e\x41\x64\x64RoutingRule\x12\x1e.service.AddRoutingRuleRequest\x1a\x0e.service.Empty\"\x00\x12H\n\x11RemoveRoutingRule\x12!.service.RemoveRoutingRuleRequest\x1a\x0e.service.Empty\"\x00\x12R\n\x16OverrideBalancerTarget\x12&.service.OverrideBalancerTargetRequest\x1a\x0e.service.Empty\"\x00\x42#Z!github.com/pasarguard/node/commonb\x06proto3') +DESCRIPTOR = _descriptor_pool.Default().AddSerializedFile(b'\n)PasarGuardNodeBridge/common/service.proto\x12\x07service\"\x07\n\x05\x45mpty\"\x8b\x01\n\x10\x42\x61seInfoResponse\x12\x0f\n\x07started\x18\x01 \x01(\x08\x12\x14\n\x0c\x63ore_version\x18\x02 \x01(\t\x12\x14\n\x0cnode_version\x18\x03 \x01(\t\x12!\n\x19user_sync_epoch_supported\x18\x04 \x01(\x08\x12\x17\n\x0fuser_sync_epoch\x18\x05 \x01(\x04\"\xa2\x01\n\x07\x42\x61\x63kend\x12\"\n\x04type\x18\x01 \x01(\x0e\x32\x14.service.BackendType\x12\x0e\n\x06\x63onfig\x18\x02 \x01(\t\x12\x1c\n\x05users\x18\x03 \x03(\x0b\x32\r.service.User\x12\x12\n\nkeep_alive\x18\x04 \x01(\x04\x12\x18\n\x10\x65xclude_inbounds\x18\x05 \x03(\t\x12\x17\n\x0fuser_sync_epoch\x18\x06 \x01(\x04\"\x15\n\x03Log\x12\x0e\n\x06\x64\x65tail\x18\x01 \x01(\t\"?\n\x04Stat\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x0c\n\x04type\x18\x02 \x01(\t\x12\x0c\n\x04link\x18\x03 \x01(\t\x12\r\n\x05value\x18\x04 \x01(\x03\",\n\x0cStatResponse\x12\x1c\n\x05stats\x18\x01 \x03(\x0b\x32\r.service.Stat\"K\n\x0bStatRequest\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\r\n\x05reset\x18\x02 \x01(\x08\x12\x1f\n\x04type\x18\x03 \x01(\x0e\x32\x11.service.StatType\"1\n\x12OnlineStatResponse\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x03\"\x8f\x01\n\x19StatsOnlineIpListResponse\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\x38\n\x03ips\x18\x02 \x03(\x0b\x32+.service.StatsOnlineIpListResponse.IpsEntry\x1a*\n\x08IpsEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\x03:\x02\x38\x01\"\x82\x01\n\x07Latency\x12\x0c\n\x04name\x18\x01 \x01(\t\x12\r\n\x05\x61live\x18\x02 \x01(\x08\x12\r\n\x05\x64\x65lay\x18\x03 \x01(\x03\x12\x0c\n\x04link\x18\x04 \x01(\t\x12\x16\n\x0elast_seen_time\x18\x05 \x01(\x03\x12\x15\n\rlast_try_time\x18\x06 \x01(\x03\x12\x0e\n\x06source\x18\x07 \x01(\t\"\x1e\n\x0eLatencyRequest\x12\x0c\n\x04name\x18\x01 \x01(\t\"6\n\x0fLatencyResponse\x12#\n\tlatencies\x18\x01 \x03(\x0b\x32\x10.service.Latency\"\xcc\x01\n\x14\x42\x61\x63kendStatsResponse\x12\x15\n\rnum_goroutine\x18\x01 \x01(\r\x12\x0e\n\x06num_gc\x18\x02 \x01(\r\x12\r\n\x05\x61lloc\x18\x03 \x01(\x04\x12\x13\n\x0btotal_alloc\x18\x04 \x01(\x04\x12\x0b\n\x03sys\x18\x05 \x01(\x04\x12\x0f\n\x07mallocs\x18\x06 \x01(\x04\x12\r\n\x05\x66rees\x18\x07 \x01(\x04\x12\x14\n\x0clive_objects\x18\x08 \x01(\x04\x12\x16\n\x0epause_total_ns\x18\t \x01(\x04\x12\x0e\n\x06uptime\x18\n \x01(\r\"\xb4\x01\n\x13SystemStatsResponse\x12\x11\n\tmem_total\x18\x01 \x01(\x04\x12\x10\n\x08mem_used\x18\x02 \x01(\x04\x12\x11\n\tcpu_cores\x18\x03 \x01(\x04\x12\x11\n\tcpu_usage\x18\x04 \x01(\x01\x12 \n\x18incoming_bandwidth_speed\x18\x05 \x01(\x04\x12 \n\x18outgoing_bandwidth_speed\x18\x06 \x01(\x04\x12\x0e\n\x06uptime\x18\x07 \x01(\x04\"\x13\n\x05Vmess\x12\n\n\x02id\x18\x01 \x01(\t\"!\n\x05Vless\x12\n\n\x02id\x18\x01 \x01(\t\x12\x0c\n\x04\x66low\x18\x02 \x01(\t\"\x1a\n\x06Trojan\x12\x10\n\x08password\x18\x01 \x01(\t\"/\n\x0bShadowsocks\x12\x10\n\x08password\x18\x01 \x01(\t\x12\x0e\n\x06method\x18\x02 \x01(\t\"1\n\tWireguard\x12\x12\n\npublic_key\x18\x01 \x01(\t\x12\x10\n\x08peer_ips\x18\x02 \x03(\t\"\x18\n\x08Hysteria\x12\x0c\n\x04\x61uth\x18\x01 \x01(\t\"\xdd\x01\n\x05Proxy\x12\x1d\n\x05vmess\x18\x01 \x01(\x0b\x32\x0e.service.Vmess\x12\x1d\n\x05vless\x18\x02 \x01(\x0b\x32\x0e.service.Vless\x12\x1f\n\x06trojan\x18\x03 \x01(\x0b\x32\x0f.service.Trojan\x12)\n\x0bshadowsocks\x18\x04 \x01(\x0b\x32\x14.service.Shadowsocks\x12%\n\twireguard\x18\x05 \x01(\x0b\x32\x12.service.Wireguard\x12#\n\x08hysteria\x18\x06 \x01(\x0b\x32\x11.service.Hysteria\"a\n\x04User\x12\r\n\x05\x65mail\x18\x01 \x01(\t\x12\x1f\n\x07proxies\x18\x02 \x01(\x0b\x32\x0e.service.Proxy\x12\x10\n\x08inbounds\x18\x03 \x03(\t\x12\x17\n\x0fuser_sync_epoch\x18\x04 \x01(\x04\">\n\x05Users\x12\x1c\n\x05users\x18\x01 \x03(\x0b\x32\r.service.User\x12\x17\n\x0fuser_sync_epoch\x18\x02 \x01(\x04\"`\n\nUsersChunk\x12\x1c\n\x05users\x18\x01 \x03(\x0b\x32\r.service.User\x12\r\n\x05index\x18\x02 \x01(\x04\x12\x0c\n\x04last\x18\x03 \x01(\x08\x12\x17\n\x0fuser_sync_epoch\x18\x04 \x01(\x04\"5\n\x0bRoutingRule\x12\x14\n\x0coutbound_tag\x18\x01 \x01(\t\x12\x10\n\x08rule_tag\x18\x02 \x01(\t\";\n\x14RoutingRulesResponse\x12#\n\x05rules\x18\x01 \x03(\x0b\x32\x14.service.RoutingRule\"\"\n\x13\x42\x61lancerInfoRequest\x12\x0b\n\x03tag\x18\x01 \x01(\t\"I\n\x14\x42\x61lancerInfoResponse\x12\x17\n\x0foverride_target\x18\x01 \x01(\t\x12\x18\n\x10principle_target\x18\x02 \x03(\t\"\xba\x02\n\x10TestRouteRequest\x12\x13\n\x0binbound_tag\x18\x01 \x01(\t\x12\x0f\n\x07network\x18\x02 \x01(\t\x12\x11\n\ttarget_ip\x18\x03 \x01(\t\x12\x15\n\rtarget_domain\x18\x04 \x01(\t\x12\x13\n\x0btarget_port\x18\x05 \x01(\r\x12\x10\n\x08protocol\x18\x06 \x01(\t\x12\x0c\n\x04user\x18\x07 \x01(\t\x12=\n\nattributes\x18\x08 \x03(\x0b\x32).service.TestRouteRequest.AttributesEntry\x12\x17\n\x0f\x66ield_selectors\x18\t \x03(\t\x12\x16\n\x0epublish_result\x18\n \x01(\x08\x1a\x31\n\x0f\x41ttributesEntry\x12\x0b\n\x03key\x18\x01 \x01(\t\x12\r\n\x05value\x18\x02 \x01(\t:\x02\x38\x01\"}\n\x0bRouteResult\x12\x14\n\x0coutbound_tag\x18\x01 \x01(\t\x12\x1b\n\x13outbound_group_tags\x18\x02 \x03(\t\x12\x13\n\x0binbound_tag\x18\x03 \x01(\t\x12\x0f\n\x07network\x18\x04 \x01(\t\x12\x15\n\rtarget_domain\x18\x05 \x01(\t\";\n\x15\x41\x64\x64RoutingRuleRequest\x12\x0c\n\x04rule\x18\x01 \x01(\t\x12\x14\n\x0cshould_reset\x18\x02 \x01(\x08\",\n\x18RemoveRoutingRuleRequest\x12\x10\n\x08rule_tag\x18\x01 \x01(\t\"E\n\x1dOverrideBalancerTargetRequest\x12\x14\n\x0c\x62\x61lancer_tag\x18\x01 \x01(\t\x12\x0e\n\x06target\x18\x02 \x01(\t*&\n\x0b\x42\x61\x63kendType\x12\x08\n\x04XRAY\x10\x00\x12\r\n\tWIREGUARD\x10\x01*_\n\x08StatType\x12\r\n\tOutbounds\x10\x00\x12\x0c\n\x08Outbound\x10\x01\x12\x0c\n\x08Inbounds\x10\x02\x12\x0b\n\x07Inbound\x10\x03\x12\r\n\tUsersStat\x10\x04\x12\x0c\n\x08UserStat\x10\x05\x32\xdc\t\n\x0bNodeService\x12\x36\n\x05Start\x12\x10.service.Backend\x1a\x19.service.BaseInfoResponse\"\x00\x12(\n\x04Stop\x12\x0e.service.Empty\x1a\x0e.service.Empty\"\x00\x12:\n\x0bGetBaseInfo\x12\x0e.service.Empty\x1a\x19.service.BaseInfoResponse\"\x00\x12+\n\x07GetLogs\x12\x0e.service.Empty\x1a\x0c.service.Log\"\x00\x30\x01\x12@\n\x0eGetSystemStats\x12\x0e.service.Empty\x1a\x1c.service.SystemStatsResponse\"\x00\x12\x42\n\x0fGetBackendStats\x12\x0e.service.Empty\x1a\x1d.service.BackendStatsResponse\"\x00\x12\x39\n\x08GetStats\x12\x14.service.StatRequest\x1a\x15.service.StatResponse\"\x00\x12J\n\x13GetOutboundsLatency\x12\x17.service.LatencyRequest\x1a\x18.service.LatencyResponse\"\x00\x12I\n\x12GetUserOnlineStats\x12\x14.service.StatRequest\x1a\x1b.service.OnlineStatResponse\"\x00\x12V\n\x18GetUserOnlineIpListStats\x12\x14.service.StatRequest\x1a\".service.StatsOnlineIpListResponse\"\x00\x12-\n\x08SyncUser\x12\r.service.User\x1a\x0e.service.Empty\"\x00(\x01\x12-\n\tSyncUsers\x12\x0e.service.Users\x1a\x0e.service.Empty\"\x00\x12;\n\x10SyncUsersChunked\x12\x13.service.UsersChunk\x1a\x0e.service.Empty\"\x00(\x01\x12\x43\n\x10ListRoutingRules\x12\x0e.service.Empty\x1a\x1d.service.RoutingRulesResponse\"\x00\x12P\n\x0fGetBalancerInfo\x12\x1c.service.BalancerInfoRequest\x1a\x1d.service.BalancerInfoResponse\"\x00\x12>\n\tTestRoute\x12\x19.service.TestRouteRequest\x1a\x14.service.RouteResult\"\x00\x12\x42\n\x0e\x41\x64\x64RoutingRule\x12\x1e.service.AddRoutingRuleRequest\x1a\x0e.service.Empty\"\x00\x12H\n\x11RemoveRoutingRule\x12!.service.RemoveRoutingRuleRequest\x1a\x0e.service.Empty\"\x00\x12R\n\x16OverrideBalancerTarget\x12&.service.OverrideBalancerTargetRequest\x1a\x0e.service.Empty\"\x00\x42#Z!github.com/pasarguard/node/commonb\x06proto3') _globals = globals() _builder.BuildMessageAndEnumDescriptors(DESCRIPTOR, _globals) @@ -36,80 +36,80 @@ _globals['_STATSONLINEIPLISTRESPONSE_IPSENTRY']._serialized_options = b'8\001' _globals['_TESTROUTEREQUEST_ATTRIBUTESENTRY']._loaded_options = None _globals['_TESTROUTEREQUEST_ATTRIBUTESENTRY']._serialized_options = b'8\001' - _globals['_BACKENDTYPE']._serialized_start=2772 - _globals['_BACKENDTYPE']._serialized_end=2810 - _globals['_STATTYPE']._serialized_start=2812 - _globals['_STATTYPE']._serialized_end=2907 + _globals['_BACKENDTYPE']._serialized_start=2933 + _globals['_BACKENDTYPE']._serialized_end=2971 + _globals['_STATTYPE']._serialized_start=2973 + _globals['_STATTYPE']._serialized_end=3068 _globals['_EMPTY']._serialized_start=54 _globals['_EMPTY']._serialized_end=61 - _globals['_BASEINFORESPONSE']._serialized_start=63 - _globals['_BASEINFORESPONSE']._serialized_end=142 - _globals['_BACKEND']._serialized_start=145 - _globals['_BACKEND']._serialized_end=282 - _globals['_LOG']._serialized_start=284 - _globals['_LOG']._serialized_end=305 - _globals['_STAT']._serialized_start=307 - _globals['_STAT']._serialized_end=370 - _globals['_STATRESPONSE']._serialized_start=372 - _globals['_STATRESPONSE']._serialized_end=416 - _globals['_STATREQUEST']._serialized_start=418 - _globals['_STATREQUEST']._serialized_end=493 - _globals['_ONLINESTATRESPONSE']._serialized_start=495 - _globals['_ONLINESTATRESPONSE']._serialized_end=544 - _globals['_STATSONLINEIPLISTRESPONSE']._serialized_start=547 - _globals['_STATSONLINEIPLISTRESPONSE']._serialized_end=690 - _globals['_STATSONLINEIPLISTRESPONSE_IPSENTRY']._serialized_start=648 - _globals['_STATSONLINEIPLISTRESPONSE_IPSENTRY']._serialized_end=690 - _globals['_LATENCY']._serialized_start=693 - _globals['_LATENCY']._serialized_end=823 - _globals['_LATENCYREQUEST']._serialized_start=825 - _globals['_LATENCYREQUEST']._serialized_end=855 - _globals['_LATENCYRESPONSE']._serialized_start=857 - _globals['_LATENCYRESPONSE']._serialized_end=911 - _globals['_BACKENDSTATSRESPONSE']._serialized_start=914 - _globals['_BACKENDSTATSRESPONSE']._serialized_end=1118 - _globals['_SYSTEMSTATSRESPONSE']._serialized_start=1121 - _globals['_SYSTEMSTATSRESPONSE']._serialized_end=1301 - _globals['_VMESS']._serialized_start=1303 - _globals['_VMESS']._serialized_end=1322 - _globals['_VLESS']._serialized_start=1324 - _globals['_VLESS']._serialized_end=1357 - _globals['_TROJAN']._serialized_start=1359 - _globals['_TROJAN']._serialized_end=1385 - _globals['_SHADOWSOCKS']._serialized_start=1387 - _globals['_SHADOWSOCKS']._serialized_end=1434 - _globals['_WIREGUARD']._serialized_start=1436 - _globals['_WIREGUARD']._serialized_end=1485 - _globals['_HYSTERIA']._serialized_start=1487 - _globals['_HYSTERIA']._serialized_end=1511 - _globals['_PROXY']._serialized_start=1514 - _globals['_PROXY']._serialized_end=1735 - _globals['_USER']._serialized_start=1737 - _globals['_USER']._serialized_end=1809 - _globals['_USERS']._serialized_start=1811 - _globals['_USERS']._serialized_end=1848 - _globals['_USERSCHUNK']._serialized_start=1850 - _globals['_USERSCHUNK']._serialized_end=1921 - _globals['_ROUTINGRULE']._serialized_start=1923 - _globals['_ROUTINGRULE']._serialized_end=1976 - _globals['_ROUTINGRULESRESPONSE']._serialized_start=1978 - _globals['_ROUTINGRULESRESPONSE']._serialized_end=2037 - _globals['_BALANCERINFOREQUEST']._serialized_start=2039 - _globals['_BALANCERINFOREQUEST']._serialized_end=2073 - _globals['_BALANCERINFORESPONSE']._serialized_start=2075 - _globals['_BALANCERINFORESPONSE']._serialized_end=2148 - _globals['_TESTROUTEREQUEST']._serialized_start=2151 - _globals['_TESTROUTEREQUEST']._serialized_end=2465 - _globals['_TESTROUTEREQUEST_ATTRIBUTESENTRY']._serialized_start=2416 - _globals['_TESTROUTEREQUEST_ATTRIBUTESENTRY']._serialized_end=2465 - _globals['_ROUTERESULT']._serialized_start=2467 - _globals['_ROUTERESULT']._serialized_end=2592 - _globals['_ADDROUTINGRULEREQUEST']._serialized_start=2594 - _globals['_ADDROUTINGRULEREQUEST']._serialized_end=2653 - _globals['_REMOVEROUTINGRULEREQUEST']._serialized_start=2655 - _globals['_REMOVEROUTINGRULEREQUEST']._serialized_end=2699 - _globals['_OVERRIDEBALANCERTARGETREQUEST']._serialized_start=2701 - _globals['_OVERRIDEBALANCERTARGETREQUEST']._serialized_end=2770 - _globals['_NODESERVICE']._serialized_start=2910 - _globals['_NODESERVICE']._serialized_end=4154 + _globals['_BASEINFORESPONSE']._serialized_start=64 + _globals['_BASEINFORESPONSE']._serialized_end=203 + _globals['_BACKEND']._serialized_start=206 + _globals['_BACKEND']._serialized_end=368 + _globals['_LOG']._serialized_start=370 + _globals['_LOG']._serialized_end=391 + _globals['_STAT']._serialized_start=393 + _globals['_STAT']._serialized_end=456 + _globals['_STATRESPONSE']._serialized_start=458 + _globals['_STATRESPONSE']._serialized_end=502 + _globals['_STATREQUEST']._serialized_start=504 + _globals['_STATREQUEST']._serialized_end=579 + _globals['_ONLINESTATRESPONSE']._serialized_start=581 + _globals['_ONLINESTATRESPONSE']._serialized_end=630 + _globals['_STATSONLINEIPLISTRESPONSE']._serialized_start=633 + _globals['_STATSONLINEIPLISTRESPONSE']._serialized_end=776 + _globals['_STATSONLINEIPLISTRESPONSE_IPSENTRY']._serialized_start=734 + _globals['_STATSONLINEIPLISTRESPONSE_IPSENTRY']._serialized_end=776 + _globals['_LATENCY']._serialized_start=779 + _globals['_LATENCY']._serialized_end=909 + _globals['_LATENCYREQUEST']._serialized_start=911 + _globals['_LATENCYREQUEST']._serialized_end=941 + _globals['_LATENCYRESPONSE']._serialized_start=943 + _globals['_LATENCYRESPONSE']._serialized_end=997 + _globals['_BACKENDSTATSRESPONSE']._serialized_start=1000 + _globals['_BACKENDSTATSRESPONSE']._serialized_end=1204 + _globals['_SYSTEMSTATSRESPONSE']._serialized_start=1207 + _globals['_SYSTEMSTATSRESPONSE']._serialized_end=1387 + _globals['_VMESS']._serialized_start=1389 + _globals['_VMESS']._serialized_end=1408 + _globals['_VLESS']._serialized_start=1410 + _globals['_VLESS']._serialized_end=1443 + _globals['_TROJAN']._serialized_start=1445 + _globals['_TROJAN']._serialized_end=1471 + _globals['_SHADOWSOCKS']._serialized_start=1473 + _globals['_SHADOWSOCKS']._serialized_end=1520 + _globals['_WIREGUARD']._serialized_start=1522 + _globals['_WIREGUARD']._serialized_end=1571 + _globals['_HYSTERIA']._serialized_start=1573 + _globals['_HYSTERIA']._serialized_end=1597 + _globals['_PROXY']._serialized_start=1600 + _globals['_PROXY']._serialized_end=1821 + _globals['_USER']._serialized_start=1823 + _globals['_USER']._serialized_end=1920 + _globals['_USERS']._serialized_start=1922 + _globals['_USERS']._serialized_end=1984 + _globals['_USERSCHUNK']._serialized_start=1986 + _globals['_USERSCHUNK']._serialized_end=2082 + _globals['_ROUTINGRULE']._serialized_start=2084 + _globals['_ROUTINGRULE']._serialized_end=2137 + _globals['_ROUTINGRULESRESPONSE']._serialized_start=2139 + _globals['_ROUTINGRULESRESPONSE']._serialized_end=2198 + _globals['_BALANCERINFOREQUEST']._serialized_start=2200 + _globals['_BALANCERINFOREQUEST']._serialized_end=2234 + _globals['_BALANCERINFORESPONSE']._serialized_start=2236 + _globals['_BALANCERINFORESPONSE']._serialized_end=2309 + _globals['_TESTROUTEREQUEST']._serialized_start=2312 + _globals['_TESTROUTEREQUEST']._serialized_end=2626 + _globals['_TESTROUTEREQUEST_ATTRIBUTESENTRY']._serialized_start=2577 + _globals['_TESTROUTEREQUEST_ATTRIBUTESENTRY']._serialized_end=2626 + _globals['_ROUTERESULT']._serialized_start=2628 + _globals['_ROUTERESULT']._serialized_end=2753 + _globals['_ADDROUTINGRULEREQUEST']._serialized_start=2755 + _globals['_ADDROUTINGRULEREQUEST']._serialized_end=2814 + _globals['_REMOVEROUTINGRULEREQUEST']._serialized_start=2816 + _globals['_REMOVEROUTINGRULEREQUEST']._serialized_end=2860 + _globals['_OVERRIDEBALANCERTARGETREQUEST']._serialized_start=2862 + _globals['_OVERRIDEBALANCERTARGETREQUEST']._serialized_end=2931 + _globals['_NODESERVICE']._serialized_start=3071 + _globals['_NODESERVICE']._serialized_end=4315 # @@protoc_insertion_point(module_scope) diff --git a/PasarGuardNodeBridge/common/service_pb2.pyi b/PasarGuardNodeBridge/common/service_pb2.pyi index e488a5b..d0b1d8f 100644 --- a/PasarGuardNodeBridge/common/service_pb2.pyi +++ b/PasarGuardNodeBridge/common/service_pb2.pyi @@ -34,28 +34,34 @@ class Empty(_message.Message): def __init__(self) -> None: ... class BaseInfoResponse(_message.Message): - __slots__ = ("started", "core_version", "node_version") + __slots__ = ("started", "core_version", "node_version", "user_sync_epoch_supported", "user_sync_epoch") STARTED_FIELD_NUMBER: _ClassVar[int] CORE_VERSION_FIELD_NUMBER: _ClassVar[int] NODE_VERSION_FIELD_NUMBER: _ClassVar[int] + USER_SYNC_EPOCH_SUPPORTED_FIELD_NUMBER: _ClassVar[int] + USER_SYNC_EPOCH_FIELD_NUMBER: _ClassVar[int] started: bool core_version: str node_version: str - def __init__(self, started: bool = ..., core_version: _Optional[str] = ..., node_version: _Optional[str] = ...) -> None: ... + user_sync_epoch_supported: bool + user_sync_epoch: int + def __init__(self, started: bool = ..., core_version: _Optional[str] = ..., node_version: _Optional[str] = ..., user_sync_epoch_supported: bool = ..., user_sync_epoch: _Optional[int] = ...) -> None: ... class Backend(_message.Message): - __slots__ = ("type", "config", "users", "keep_alive", "exclude_inbounds") + __slots__ = ("type", "config", "users", "keep_alive", "exclude_inbounds", "user_sync_epoch") TYPE_FIELD_NUMBER: _ClassVar[int] CONFIG_FIELD_NUMBER: _ClassVar[int] USERS_FIELD_NUMBER: _ClassVar[int] KEEP_ALIVE_FIELD_NUMBER: _ClassVar[int] EXCLUDE_INBOUNDS_FIELD_NUMBER: _ClassVar[int] + USER_SYNC_EPOCH_FIELD_NUMBER: _ClassVar[int] type: BackendType config: str users: _containers.RepeatedCompositeFieldContainer[User] keep_alive: int exclude_inbounds: _containers.RepeatedScalarFieldContainer[str] - def __init__(self, type: _Optional[_Union[BackendType, str]] = ..., config: _Optional[str] = ..., users: _Optional[_Iterable[_Union[User, _Mapping]]] = ..., keep_alive: _Optional[int] = ..., exclude_inbounds: _Optional[_Iterable[str]] = ...) -> None: ... + user_sync_epoch: int + def __init__(self, type: _Optional[_Union[BackendType, str]] = ..., config: _Optional[str] = ..., users: _Optional[_Iterable[_Union[User, _Mapping]]] = ..., keep_alive: _Optional[int] = ..., exclude_inbounds: _Optional[_Iterable[str]] = ..., user_sync_epoch: _Optional[int] = ...) -> None: ... class Log(_message.Message): __slots__ = ("detail",) @@ -245,30 +251,36 @@ class Proxy(_message.Message): def __init__(self, vmess: _Optional[_Union[Vmess, _Mapping]] = ..., vless: _Optional[_Union[Vless, _Mapping]] = ..., trojan: _Optional[_Union[Trojan, _Mapping]] = ..., shadowsocks: _Optional[_Union[Shadowsocks, _Mapping]] = ..., wireguard: _Optional[_Union[Wireguard, _Mapping]] = ..., hysteria: _Optional[_Union[Hysteria, _Mapping]] = ...) -> None: ... class User(_message.Message): - __slots__ = ("email", "proxies", "inbounds") + __slots__ = ("email", "proxies", "inbounds", "user_sync_epoch") EMAIL_FIELD_NUMBER: _ClassVar[int] PROXIES_FIELD_NUMBER: _ClassVar[int] INBOUNDS_FIELD_NUMBER: _ClassVar[int] + USER_SYNC_EPOCH_FIELD_NUMBER: _ClassVar[int] email: str proxies: Proxy inbounds: _containers.RepeatedScalarFieldContainer[str] - def __init__(self, email: _Optional[str] = ..., proxies: _Optional[_Union[Proxy, _Mapping]] = ..., inbounds: _Optional[_Iterable[str]] = ...) -> None: ... + user_sync_epoch: int + def __init__(self, email: _Optional[str] = ..., proxies: _Optional[_Union[Proxy, _Mapping]] = ..., inbounds: _Optional[_Iterable[str]] = ..., user_sync_epoch: _Optional[int] = ...) -> None: ... class Users(_message.Message): - __slots__ = ("users",) + __slots__ = ("users", "user_sync_epoch") USERS_FIELD_NUMBER: _ClassVar[int] + USER_SYNC_EPOCH_FIELD_NUMBER: _ClassVar[int] users: _containers.RepeatedCompositeFieldContainer[User] - def __init__(self, users: _Optional[_Iterable[_Union[User, _Mapping]]] = ...) -> None: ... + user_sync_epoch: int + def __init__(self, users: _Optional[_Iterable[_Union[User, _Mapping]]] = ..., user_sync_epoch: _Optional[int] = ...) -> None: ... class UsersChunk(_message.Message): - __slots__ = ("users", "index", "last") + __slots__ = ("users", "index", "last", "user_sync_epoch") USERS_FIELD_NUMBER: _ClassVar[int] INDEX_FIELD_NUMBER: _ClassVar[int] LAST_FIELD_NUMBER: _ClassVar[int] + USER_SYNC_EPOCH_FIELD_NUMBER: _ClassVar[int] users: _containers.RepeatedCompositeFieldContainer[User] index: int last: bool - def __init__(self, users: _Optional[_Iterable[_Union[User, _Mapping]]] = ..., index: _Optional[int] = ..., last: bool = ...) -> None: ... + user_sync_epoch: int + def __init__(self, users: _Optional[_Iterable[_Union[User, _Mapping]]] = ..., index: _Optional[int] = ..., last: bool = ..., user_sync_epoch: _Optional[int] = ...) -> None: ... class RoutingRule(_message.Message): __slots__ = ("outbound_tag", "rule_tag") diff --git a/PasarGuardNodeBridge/controller.py b/PasarGuardNodeBridge/controller.py index cd79d73..27e36fe 100644 --- a/PasarGuardNodeBridge/controller.py +++ b/PasarGuardNodeBridge/controller.py @@ -1,10 +1,13 @@ import asyncio +import inspect import logging import math import ssl +import sys +import traceback from enum import IntEnum from json import JSONDecodeError -from typing import Optional +from typing import Optional, cast from uuid import UUID import aiohttp @@ -22,10 +25,16 @@ from PasarGuardNodeBridge.storage import ( ClaimedUser, LifecycleLease, + LifecycleLeaseLostError, LifecycleOperation, LifecycleStatus, NodeLifecycleCoordinatorProtocol, NodeLifecycleState, + RevocationAwareUserSyncStoreProtocol, + StartupUserSyncLease, + UserRevocationResult, + UserSyncLease, + UserSyncLeaseLostError, UserSyncStoreProtocol, get_default_lifecycle_coordinator, get_default_user_sync_store, @@ -34,6 +43,64 @@ # Default timeout configuration (module-level constants) DEFAULT_API_TIMEOUT = 10 # Default timeout for public API methods DEFAULT_INTERNAL_TIMEOUT = 15 # Default timeout for internal gRPC/HTTP operations +CLAIM_RECOVERY_TIMEOUT = 1.0 +SYNC_WORKER_CLEANUP_TIMEOUT = CLAIM_RECOVERY_TIMEOUT + 1.0 +MIN_CLAIM_RECHECK_DELAY = 0.01 +INITIAL_CLAIM_RETRY_DELAY = 1.0 +MAX_CLAIM_RETRY_DELAY = 30.0 +STALE_USER_SYNC_RETRY_LIMIT = 1 + + +def _sanitize_log_text(value: object, limit: int = 2048) -> str: + text = str(value) + sanitized_parts = [] + for character in text: + codepoint = ord(character) + if codepoint in (0x2028, 0x2029): + sanitized_parts.append(f"\\u{codepoint:04x}") + elif codepoint < 32 or codepoint in range(127, 160): + sanitized_parts.append(f"\\x{codepoint:02x}") + else: + sanitized_parts.append(character) + sanitized = "".join(sanitized_parts) + if len(sanitized) > limit: + return f"{sanitized[:limit]}...[truncated]" + return sanitized + + +class _SanitizingLoggerAdapter(logging.LoggerAdapter): + def log(self, level, msg, *args, **kwargs): + if not self.isEnabledFor(level): + return + + if args: + try: + record = logging.LogRecord("", level, "", 0, msg, args, None) + msg = record.getMessage() + args = () + except Exception: # noqa: BLE001,S110 - preserve logging's formatting-error behavior + # Preserve logging's normal formatting-error behavior while + # still sanitizing the format string itself. + pass + + msg, kwargs = self.process(msg, kwargs) + self.logger.log(level, msg, *args, **kwargs) + + def process(self, msg, kwargs): + sanitized_message = _sanitize_log_text(msg) + exc_info = kwargs.get("exc_info") + if exc_info: + try: + if isinstance(exc_info, BaseException): + exc_info = (type(exc_info), exc_info, exc_info.__traceback__) + elif exc_info is True or not isinstance(exc_info, tuple): + exc_info = sys.exc_info() + formatted_exception = "".join(traceback.format_exception(*exc_info)) + sanitized_message = f"{sanitized_message} | Traceback: {_sanitize_log_text(formatted_exception)}" + kwargs["exc_info"] = None + except Exception: # noqa: BLE001 - logging must not replace the operational exception + kwargs["exc_info"] = None + return sanitized_message, kwargs class NodeAPIError(Exception): @@ -53,6 +120,17 @@ class Health(IntEnum): class Controller: + _REVOCATION_STORE_METHODS = ( + "begin_user_revocation", + "abort_user_revocation", + "finalize_user_revocation", + "acquire_user_sync_lease", + "acquire_startup_user_sync_lease", + "retain_user_sync_lease_keys", + "heartbeat_user_sync_lease", + "release_user_sync_lease", + ) + def __init__( self, server_ca: str, @@ -78,12 +156,10 @@ def __init__( if extra is None: extra = {} if logger is None: - logger = logging.getLogger(self.name) - logger.setLevel(logging.INFO) - handler = logging.StreamHandler() - handler.setFormatter(logging.Formatter("%(asctime)s - %(name)s - %(levelname)s - %(message)s")) - logger.addHandler(handler) - self.logger = logger + # Libraries should not install output handlers implicitly. Applications + # that want bridge logs can configure this package logger or pass one. + logger = logging.getLogger("PasarGuardNodeBridge") + self.logger = _SanitizingLoggerAdapter(logger, {}) # Timeout configuration self._default_timeout = default_timeout @@ -119,6 +195,10 @@ def __init__( self._tasks: list[asyncio.Task] = [] self._node_version = "" self._core_version = "" + self._user_sync_epoch_supported = False + self._user_sync_epoch_capability_probed = False + self._user_sync_epoch_handshake_lock = asyncio.Lock() + self._user_sync_connection_generation = 0 self._extra = extra # Lazy worker sync mechanism @@ -184,7 +264,8 @@ async def _increment_user_sync_failure(self): if not self._hard_reset_event.is_set(): self._hard_reset_event.set() self.logger.critical( - f"[{self.name}] HARD RESET REQUIRED: User sync failed {self._user_sync_failure_count} times in a row" + f"[{self.name}] HARD RESET REQUIRED: User sync failed " + f"{self._user_sync_failure_count} times in a row" ) async def _reset_user_sync_failure_count(self): @@ -200,9 +281,339 @@ async def _reset_user_sync_failure_count(self): f"[{self.name}] User sync recovered after {old_count} failures, cleared hard reset event" ) + def _revocation_store(self, *, required: bool = False) -> RevocationAwareUserSyncStoreProtocol | None: + store = getattr(self, "_user_sync_store", None) + if store is not None and all( + callable(getattr(store, method, None)) for method in self._REVOCATION_STORE_METHODS + ): + return cast(RevocationAwareUserSyncStoreProtocol, store) + if required: + raise NodeAPIError(501, "Configured user sync store does not support coordinated user revocation") + return None + + @staticmethod + def _normalize_user_keys(user_keys: list[str]) -> list[str]: + unique_keys = list(dict.fromkeys(user_keys)) + if any(not isinstance(user_key, str) or not user_key for user_key in unique_keys): + raise ValueError("user keys must be non-empty strings") + return unique_keys + + async def _acquire_user_sync_lease( + self, + user_keys: list[str], + expected_generations: dict[str, int] | None = None, + revocation_id: str | None = None, + ) -> UserSyncLease | None: + store = self._revocation_store(required=revocation_id is not None) + if store is None: + return None + return await store.acquire_user_sync_lease( + self.node_id, + self.worker_id, + self._normalize_user_keys(user_keys), + max(self._sync_lease_seconds, 0.1), + expected_generations, + revocation_id, + ) + + async def _heartbeat_user_sync_lease(self, lease: UserSyncLease) -> None: + store = self._revocation_store(required=True) + interval = max(lease.lease_seconds / 3, 0.01) + try: + while True: + await asyncio.sleep(interval) + if not await store.heartbeat_user_sync_lease(lease): + raise RuntimeError("User sync execution lease was lost") + except asyncio.CancelledError: + pass + + @staticmethod + async def _await_cleanup_despite_cancellation(awaitable) -> asyncio.CancelledError | None: + """Finish distributed cleanup before propagating caller cancellation. + + The cleanup runs in its own task and ``asyncio.wait`` observes it + without propagating caller cancellation into that task. Any cleanup + failure remains authoritative and is raised by ``task.result``. + """ + cleanup = asyncio.ensure_future(awaitable) + caller_cancellation: asyncio.CancelledError | None = None + while not cleanup.done(): + try: + await asyncio.wait((cleanup,)) + except asyncio.CancelledError as exc: + current = asyncio.current_task() + if current is None or not current.cancelling(): + if cleanup.done(): + break + raise + caller_cancellation = exc + cleanup.result() + return caller_cancellation + + async def _release_user_sync_lease( + self, + lease: UserSyncLease | None, + heartbeat: asyncio.Task | None = None, + ) -> None: + heartbeat_error: Exception | None = None + caller_cancellation: asyncio.CancelledError | None = None + if heartbeat is not None: + heartbeat.cancel() + try: + await heartbeat + except asyncio.CancelledError as exc: + current = asyncio.current_task() + if current is not None and current.cancelling(): + caller_cancellation = exc + except Exception as exc: # noqa: BLE001 - report a failed distributed heartbeat after cleanup + heartbeat_error = exc + if heartbeat_error is not None: + self.logger.error( + f"[{self.name}] User sync execution lease heartbeat failed | " + f"Error: {type(heartbeat_error).__name__} - {heartbeat_error!s}" + ) + raise UserSyncLeaseLostError("user-sync lease ownership was lost before completion") from heartbeat_error + if lease is not None: + store = self._revocation_store(required=True) + cleanup_cancellation = await self._await_cleanup_despite_cancellation( + store.release_user_sync_lease(lease) + ) + if caller_cancellation is None: + caller_cancellation = cleanup_cancellation + if caller_cancellation is not None: + raise caller_cancellation + + async def _abandon_user_sync_lease( + self, + lease: UserSyncLease | None, + heartbeat: asyncio.Task | None = None, + ) -> None: + """Stop renewing a lease whose remote write outcome is unknown. + + The store record is deliberately retained. Once it expires, coordinated + revocation must fail closed until an operator explicitly reconciles and + releases that exact lease. + """ + caller_cancellation: asyncio.CancelledError | None = None + if heartbeat is not None: + heartbeat.cancel() + try: + await heartbeat + except asyncio.CancelledError as exc: + current = asyncio.current_task() + if current is not None and current.cancelling(): + caller_cancellation = exc + except Exception as exc: # noqa: BLE001 - the retained lease is already fail-closed + self.logger.error( + f"[{self.name}] User sync execution lease heartbeat failed | Error: {type(exc).__name__} - {exc!s}" + ) + if lease is not None: + self.logger.error( + f"[{self.name}] User sync outcome is unknown; retaining execution lease {lease.token!r} " + "for explicit reconciliation" + ) + if caller_cancellation is not None: + raise caller_cancellation + + async def _retain_unknown_user_sync_lease_keys( + self, + lease: UserSyncLease | None, + heartbeat: asyncio.Task | None, + unknown_user_keys: list[str], + ) -> None: + """Release known outcomes while retaining a poison lease for unknown keys.""" + if lease is None: + return + unknown_keys = self._normalize_user_keys(unknown_user_keys) + if set(unknown_keys) == set(lease.user_keys): + await self._abandon_user_sync_lease(lease, heartbeat) + return + + # Stop the task that holds the original lease value before replacing + # that value in the store. Otherwise it can race the narrowing update, + # observe a different lease object, and report a false ownership loss. + await self._abandon_user_sync_lease(None, heartbeat) + try: + store = self._revocation_store(required=True) + narrowed = await store.retain_user_sync_lease_keys(lease, unknown_keys) + except BaseException: + await self._abandon_user_sync_lease(lease) + raise + await self._abandon_user_sync_lease(narrowed) + + async def _acquire_direct_user_sync_lease( + self, + users: list[User], + revocation_id: str | None = None, + ) -> tuple[UserSyncLease | None, asyncio.Task | None]: + await self._probe_user_sync_epoch_capability() + if revocation_id is not None: + await self._ensure_user_sync_epoch_support() + user_keys = self._normalize_user_keys([user.email for user in users]) + lease = await self._acquire_user_sync_lease(user_keys, revocation_id=revocation_id) + if lease is None: + return None, None + denied_keys = set(user_keys).difference(lease.user_keys) + if denied_keys: + await self._release_user_sync_lease(lease) + raise NodeAPIError(409, f"User sync is fenced for {len(denied_keys)} user(s)") + heartbeat = asyncio.create_task(self._heartbeat_user_sync_lease(lease)) if lease.token else None + return lease, heartbeat + + async def _acquire_snapshot_user_sync_lease( + self, + users: list[User], + ) -> tuple[list[User], UserSyncLease | None, asyncio.Task | None]: + """Acquire a node-wide replacement permit and omit permanently fenced users.""" + await self._probe_user_sync_epoch_capability() + store = self._revocation_store() + if store is None: + return users, None, None + startup: StartupUserSyncLease = await store.acquire_startup_user_sync_lease( + self.node_id, + self.worker_id, + self._normalize_user_keys([user.email for user in users]), + max(self._sync_lease_seconds, 0.1), + ) + allowed_keys = set(startup.included_user_keys) + filtered_users = [user for user in users if user.email in allowed_keys] + lease = startup.lease + heartbeat = asyncio.create_task(self._heartbeat_user_sync_lease(lease)) if lease.token else None + return filtered_users, lease, heartbeat + + async def _acquire_reconciliation_user_sync_lease( + self, + users: list[User], + ) -> tuple[list[User], UserSyncLease, asyncio.Task]: + """Acquire the explicit full-snapshot recovery permit.""" + await self._ensure_user_sync_epoch_support() + store = self._revocation_store(required=True) + acquire_reconciliation = getattr(store, "acquire_user_sync_reconciliation_lease", None) + if not callable(acquire_reconciliation): + raise NodeAPIError(501, "Configured user sync store does not support authoritative reconciliation") + recovery: StartupUserSyncLease = await acquire_reconciliation( + self.node_id, + self.worker_id, + self._normalize_user_keys([user.email for user in users]), + max(self._sync_lease_seconds, 0.1), + ) + allowed_keys = set(recovery.included_user_keys) + filtered_users = [user for user in users if user.email in allowed_keys] + lease = recovery.lease + heartbeat = asyncio.create_task(self._heartbeat_user_sync_lease(lease)) + return filtered_users, lease, heartbeat + + async def _assert_user_sync_lease_owned(self, lease: UserSyncLease | None) -> None: + """Revalidate a permit immediately before a remote side effect.""" + if lease is None or not lease.token: + return + store = self._revocation_store(required=True) + if not await store.heartbeat_user_sync_lease(lease): + raise UserSyncLeaseLostError("user-sync execution lease was lost before transport") + + async def _observe_user_sync_epoch_capability( + self, + info: object, + expected_generation: int | None = None, + ) -> None: + lock = getattr(self, "_user_sync_epoch_handshake_lock", None) + if lock is None: + lock = asyncio.Lock() + self._user_sync_epoch_handshake_lock = lock + async with lock: + if expected_generation is not None and expected_generation != getattr( + self, "_user_sync_connection_generation", 0 + ): + return + supported = bool(getattr(info, "user_sync_epoch_supported", False)) + if supported: + store = self._revocation_store() + if store is not None: + advance_epoch = getattr(store, "advance_user_sync_epoch", None) + if not callable(advance_epoch): + self._user_sync_epoch_supported = False + self._user_sync_epoch_capability_probed = False + raise NodeAPIError(501, "Configured user sync store does not support epoch handshakes") + await advance_epoch(self.node_id, int(getattr(info, "user_sync_epoch", 0))) + self._user_sync_epoch_supported = supported + self._user_sync_epoch_capability_probed = True + + def _require_user_sync_epoch_support(self) -> None: + if not getattr(self, "_user_sync_epoch_supported", False): + raise NodeAPIError(426, "Node does not advertise monotonic user-sync epoch fencing") + + async def _ensure_user_sync_epoch_support(self) -> None: + await self._probe_user_sync_epoch_capability() + self._require_user_sync_epoch_support() + + async def _probe_user_sync_epoch_capability(self) -> None: + if getattr(self, "_user_sync_epoch_capability_probed", False): + return + info_method = getattr(self, "info", None) + if not callable(info_method): + self._user_sync_epoch_capability_probed = True + return + capability_generation = getattr(self, "_user_sync_connection_generation", 0) + info = await info_method() + if info is not None and not getattr(self, "_user_sync_epoch_capability_probed", False): + await self._observe_user_sync_epoch_capability(info, capability_generation) + + def _user_sync_epoch_for_transport(self, lease: UserSyncLease | None) -> int: + if lease is None or not getattr(self, "_user_sync_epoch_supported", False): + return 0 + return lease.epoch + + @staticmethod + def _is_stale_user_sync_rejection(error: BaseException) -> bool: + """Return whether the node rejected an epoch before applying anything.""" + if isinstance(error, NodeAPIError): + return error.code == 412 + status = getattr(error, "status", None) + return getattr(status, "name", None) == "FAILED_PRECONDITION" + + @staticmethod + def _accepts_user_sync_epoch(callback: object) -> bool: + """Preserve legacy transport hooks without masking callback errors.""" + try: + parameters = inspect.signature(callback).parameters + except (TypeError, ValueError): + return True + return "user_sync_epoch" in parameters or any( + parameter.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD) + for parameter in parameters.values() + ) + + async def begin_user_revocation(self, user_keys: list[str], revocation_id: str) -> UserRevocationResult: + """Fence keys and wait until older queued/in-flight updates cannot run.""" + if not revocation_id: + raise ValueError("revocation_id must not be empty") + await self._ensure_user_sync_epoch_support() + store = self._revocation_store(required=True) + return await store.begin_user_revocation(self.node_id, self._normalize_user_keys(user_keys), revocation_id) + + async def abort_user_revocation(self, user_keys: list[str], revocation_id: str) -> None: + """Release only this provisional revocation after authoritative restore.""" + if not revocation_id: + raise ValueError("revocation_id must not be empty") + store = self._revocation_store(required=True) + await store.abort_user_revocation(self.node_id, self._normalize_user_keys(user_keys), revocation_id) + + async def finalize_user_revocation(self, user_keys: list[str], revocation_id: str) -> None: + """Convert this operation's provisional fences to permanent tombstones.""" + if not revocation_id: + raise ValueError("revocation_id must not be empty") + store = self._revocation_store(required=True) + await store.finalize_user_revocation(self.node_id, self._normalize_user_keys(user_keys), revocation_id) + async def update_user(self, user: User): """Queue a user for sync. Automatically deduplicates by email.""" - await self._user_sync_store.enqueue_users(self.node_id, [user]) + lease = await self._acquire_user_sync_lease([user.email]) + try: + if lease is not None and user.email not in lease.user_keys: + return + await self._user_sync_store.enqueue_users(self.node_id, [user]) + finally: + await self._release_user_sync_lease(lease) self._work_available.set() # Ensure worker is running to process the update @@ -213,7 +624,16 @@ async def update_users(self, users: list[User]): if not users: return - await self._user_sync_store.enqueue_users(self.node_id, users) + lease = await self._acquire_user_sync_lease([user.email for user in users]) + try: + if lease is not None: + allowed_keys = set(lease.user_keys) + users = [user for user in users if user.email in allowed_keys] + if not users: + return + await self._user_sync_store.enqueue_users(self.node_id, users) + finally: + await self._release_user_sync_lease(lease) self._work_available.set() # Ensure worker is running to process the updates @@ -283,6 +703,20 @@ async def get_extra(self) -> dict: async with self._version_lock: return self._extra + @property + def extra(self) -> dict: + """Backward-compatible access to node metadata. + + New asynchronous code should prefer :meth:`get_extra` when it needs + metadata coordinated with controller state updates. + """ + + return self._extra + + @extra.setter + def extra(self, value: dict) -> None: + self._extra = value + @staticmethod def _parse_version(version: str) -> Version | None: """Parse semver-like strings into a packaging Version for comparison.""" @@ -318,7 +752,8 @@ async def _heartbeat_lifecycle_lease(self, lease: LifecycleLease) -> None: try: while True: await asyncio.sleep(interval) - await self._lifecycle_coordinator.heartbeat(lease) + if not await self._lifecycle_coordinator.heartbeat(lease): + raise LifecycleLeaseLostError(f"Lifecycle lease ownership was lost for node {self.node_id}") except asyncio.CancelledError: pass @@ -339,13 +774,31 @@ async def _release_lifecycle_lease( node_version: str = "", core_version: str = "", ) -> None: - heartbeat = self._lifecycle_heartbeat_tasks.pop(lease.token, None) + heartbeat = getattr(self, "_lifecycle_heartbeat_tasks", {}).pop(lease.token, None) + heartbeat_error: BaseException | None = None + caller_cancellation: asyncio.CancelledError | None = None if heartbeat is not None: heartbeat.cancel() - await heartbeat + try: + await heartbeat + except asyncio.CancelledError as exc: + current = asyncio.current_task() + if current is not None and current.cancelling(): + caller_cancellation = exc + except BaseException as exc: # preserve the unknown remote outcome + heartbeat_error = exc + + if heartbeat_error is not None: + raise heartbeat_error if observed is None: - await self._lifecycle_coordinator.release(lease) + cleanup_cancellation = await self._await_cleanup_despite_cancellation( + self._lifecycle_coordinator.release(lease) + ) + if caller_cancellation is None: + caller_cancellation = cleanup_cancellation + if caller_cancellation is not None: + raise caller_cancellation return state = NodeLifecycleState( @@ -357,11 +810,44 @@ async def _release_lifecycle_lease( node_version=node_version, core_version=core_version, ) - await self._lifecycle_coordinator.release(lease, state) + cleanup_cancellation = await self._await_cleanup_despite_cancellation( + self._lifecycle_coordinator.release(lease, state) + ) + if caller_cancellation is None: + caller_cancellation = cleanup_cancellation + if caller_cancellation is not None: + raise caller_cancellation + + async def _stop_lifecycle_heartbeat(self, lease: LifecycleLease | None) -> None: + """Stop renewal after a failed operation without masking its error.""" + if lease is None: + return + heartbeat = getattr(self, "_lifecycle_heartbeat_tasks", {}).pop(lease.token, None) + if heartbeat is None: + return + heartbeat.cancel() + try: + await heartbeat + except asyncio.CancelledError: + current = asyncio.current_task() + if current is not None and current.cancelling(): + raise + except Exception: + self.logger.exception("[%s] Lifecycle heartbeat failed during cleanup", self.name) async def get_lifecycle_state(self) -> NodeLifecycleState | None: return await self._lifecycle_coordinator.get_state(self.node_id) + async def reconcile_lifecycle(self, observed: LifecycleStatus) -> None: + """Acknowledge an inspected remote state after an expired operation. + + Reconciliation is intentionally explicit: an expired lifecycle request + is an unknown remote effect and must not be replaced by a new operation + until a caller has probed the node and supplied the observed state. + """ + if not await self._lifecycle_coordinator.reconcile(self.node_id, observed): + raise NodeAPIError(409, f"Node lifecycle operation is still active for {self.node_id}") + async def update_observed_lifecycle(self, observed: LifecycleStatus, expected_epoch: int | None = None) -> None: await self._lifecycle_coordinator.update_observed(self.node_id, observed, expected_epoch) @@ -400,15 +886,31 @@ async def connect(self, node_version: str, core_version: str, tasks: list | None task = asyncio.create_task(t()) self._tasks.append(task) + # Updates can be queued while disconnected, including by another + # controller sharing the store. Wake a worker after every successful + # reconnect; it will exit normally after the idle timeout if no work exists. + self._work_available.set() + await self._ensure_sync_worker_running() + async def disconnect(self): # Set shutdown event (no lock needed) self._shutdown_event.set() + handshake_lock = getattr(self, "_user_sync_epoch_handshake_lock", None) + if handshake_lock is None: + handshake_lock = asyncio.Lock() + self._user_sync_epoch_handshake_lock = handshake_lock + async with handshake_lock: + self._user_sync_connection_generation = getattr(self, "_user_sync_connection_generation", 0) + 1 + self._user_sync_epoch_supported = False + self._user_sync_epoch_capability_probed = False + # Cleanup tasks async with self._task_lock: await self._cleanup_tasks() - # Cleanup sync worker and pending users + # Stop this controller's worker without deleting work from the shared + # store. Pending updates are resumed by this or another controller. await self._cleanup_sync_worker() # Clear versions and set health atomically to prevent race condition @@ -443,19 +945,23 @@ async def _cleanup_tasks(self): self._tasks.clear() async def _cleanup_sync_worker(self): - """Clean up sync worker and pending users.""" + """Stop this controller's sync worker while preserving shared pending users.""" # Cancel sync worker if running async with self._sync_worker_lock: - if self._sync_worker_task and not self._sync_worker_task.done(): - self._sync_worker_task.cancel() + task = self._sync_worker_task + if task is not None and not task.done(): + task.cancel() try: - await asyncio.wait_for(self._sync_worker_task, timeout=2.0) + await asyncio.wait_for(task, timeout=SYNC_WORKER_CLEANUP_TIMEOUT) except (asyncio.CancelledError, asyncio.TimeoutError): pass + except Exception as e: # noqa: BLE001 - cleanup must complete even if worker recovery fails + self.logger.warning( + f"[{self.name}] Sync worker cleanup observed an error | Error: {type(e).__name__} - {e!s}" + ) + if self._sync_worker_task is task: self._sync_worker_task = None - # Clear pending users - await self._user_sync_store.clear(self.node_id) self._work_available.clear() def is_shutting_down(self) -> bool: @@ -466,17 +972,68 @@ async def _ensure_sync_worker_running(self): """Spawn sync worker if not already running.""" async with self._sync_worker_lock: if self._sync_worker_task is None or self._sync_worker_task.done(): - self._sync_worker_task = asyncio.create_task(self._sync_worker()) + task = asyncio.create_task(self._sync_worker()) + self._sync_worker_task = task + task.add_done_callback(self._clear_finished_sync_worker) + + def _clear_finished_sync_worker(self, task: asyncio.Task) -> None: + """Drop a completed worker reference without racing a replacement worker.""" + if self._sync_worker_task is task: + self._sync_worker_task = None + + async def _retire_sync_worker_if_idle(self) -> bool: + """Atomically retire the current worker unless a concurrent enqueue woke it.""" + current_task = asyncio.current_task() + async with self._sync_worker_lock: + # update_user(s) sets the event before taking this lock in + # _ensure_sync_worker_running(). If that enqueue won the race, the + # current worker must keep running and consume the queued update. + if self._work_available.is_set(): + return False + + # Publish retirement before leaving the worker. An enqueue that + # happens after this point will observe None and start a replacement. + # Keep the identity check so an older worker can never clear a newer + # worker that was installed while it was finishing. + if self._sync_worker_task is current_task: + self._sync_worker_task = None + return True async def _claim_pending_users(self, limit: int = 2000) -> list[ClaimedUser]: """Claim pending users from the configured sync store.""" + # Clear before crossing the storage await boundary. An enqueue that + # races with claim_users() will set the event afterwards and must not + # be erased when an eventually-empty claim returns. + self._work_available.clear() claimed = await self._user_sync_store.claim_users( self.node_id, self.worker_id, limit=limit, lease_seconds=self._sync_lease_seconds ) - if not claimed: - self._work_available.clear() + if claimed: + # Keep draining. The store may still contain more than one claim + # batch, and one harmless empty claim restores the idle state. + self._work_available.set() return claimed + async def _next_claim_delay(self) -> tuple[bool, float | None]: + """Return whether the store can report the next claimable-work deadline.""" + next_claim_delay = getattr(self._user_sync_store, "next_claim_delay", None) + if next_claim_delay is None: + return False, None + return True, await next_claim_delay(self.node_id) + + async def _wait_for_claim_recheck(self, delay: float) -> None: + """Sleep until a store lease may expire, while remaining locally wakeable.""" + # A distributed/custom store may transiently report a due deadline + # while another worker wins the claim. Always yield for a real, + # positive interval even when polling is explicitly disabled. + wait_delay = max(delay, self._sync_poll_interval, MIN_CLAIM_RECHECK_DELAY) + try: + await asyncio.wait_for(self._work_available.wait(), timeout=wait_delay) + except asyncio.TimeoutError: + # Re-arm the normal worker loop after the bounded wait. A local + # enqueue can also set the event and wake this wait early. + self._work_available.set() + async def _ack_claimed_users(self, claimed_users: list[ClaimedUser]): await self._user_sync_store.ack_users(self.node_id, [item.token for item in claimed_users]) @@ -485,26 +1042,63 @@ async def _requeue_claimed_users(self, claimed_users: list[ClaimedUser]): if claimed_users: self._work_available.set() + async def _recover_claimed_users(self, claimed_users: list[ClaimedUser], context: str) -> bool: + """Attempt bounded claim recovery; an unexpired store lease remains the fallback.""" + if not claimed_users: + return True + recovery_task = asyncio.create_task(self._requeue_claimed_users(claimed_users)) + + def abandon_recovery(error: BaseException) -> bool: + recovery_task.cancel() + recovery_task.add_done_callback(lambda task: None if task.cancelled() else task.exception()) + self.logger.error( + f"[{self.name}] Failed to recover {len(claimed_users)} claimed user(s) after {context} | " + f"Error: {type(error).__name__}" + ) + return False + + try: + done, _ = await asyncio.wait((recovery_task,), timeout=CLAIM_RECOVERY_TIMEOUT) + if not done: + return abandon_recovery(TimeoutError()) + await recovery_task + return True + except asyncio.CancelledError as requeue_error: + return abandon_recovery(requeue_error) + except Exception as requeue_error: # noqa: BLE001 - the lease is the fallback for any storage failure + return abandon_recovery(requeue_error) + async def _sync_worker(self): """Lazy worker that processes pending users and exits when idle.""" self.logger.debug(f"[{self.name}] Sync worker started") retry_delay = 1.0 max_retry_delay = 30.0 + claim_retry_delay = INITIAL_CLAIM_RETRY_DELAY supports_chunked, node_version = await self._supports_chunked_sync() + legacy_lease_recheck_deadline: float | None = None if not supports_chunked: self.logger.debug( f"[{self.name}] Chunked sync disabled for node version '{node_version or 'unknown'}' (< v0.2.0)" ) + claimed_users: list[ClaimedUser] = [] + user_sync_lease: UserSyncLease | None = None + user_sync_heartbeat: asyncio.Task | None = None + user_sync_remote_started = False + user_sync_remote_completed = False try: while not self.is_shutting_down(): # Wait for work or timeout try: await asyncio.wait_for(self._work_available.wait(), timeout=self._worker_idle_timeout) except asyncio.TimeoutError: - # No work for idle_timeout seconds, exit worker - self.logger.debug(f"[{self.name}] Sync worker idle, exiting") - break + if await self._retire_sync_worker_if_idle(): + self.logger.debug(f"[{self.name}] Sync worker idle, exiting") + break + # An enqueue set the wake event before acquiring the worker + # lock. Keep this worker rather than letting ensure() observe + # a task that is about to exit. + continue # Check health - don't sync if not connected or invalid health = await self.get_health() @@ -523,32 +1117,127 @@ async def _sync_worker(self): continue # Claim pending users atomically (only when healthy) - claimed_users = await self._claim_pending_users() + try: + claimed_users = await self._claim_pending_users() + claim_retry_delay = INITIAL_CLAIM_RETRY_DELAY + except asyncio.CancelledError: + raise + except Exception as e: # noqa: BLE001 - storage failures must not strand queued work + # _claim_pending_users clears the event before awaiting the + # store. Re-arm it so pending work remains discoverable even + # when an enqueue raced with this failed claim. Retrying in + # the existing worker also closes the completion/ensure race + # without ever creating a second worker. + self._work_available.set() + self.logger.warning( + f"[{self.name}] Failed to claim pending users, retrying in {claim_retry_delay}s | " + f"Error: {type(e).__name__} - {e!s}" + ) + await asyncio.sleep(claim_retry_delay) + claim_retry_delay = min(claim_retry_delay * 2, MAX_CLAIM_RETRY_DELAY) + continue if not claimed_users: - await asyncio.sleep(self._sync_poll_interval) + lease_aware, next_claim_delay = await self._next_claim_delay() + if lease_aware: + legacy_lease_recheck_deadline = None + if next_claim_delay is not None: + await self._wait_for_claim_recheck(next_claim_delay) + # A lease-aware store reporting no tracked work is + # genuinely idle. Loop directly back to the original + # event/idle-timeout wait without adding poll latency. + else: + # Backward compatibility for custom stores created + # before next_claim_delay existed. Wait for at most one + # configured lease horizon, then restore normal idle exit. + loop = asyncio.get_running_loop() + if legacy_lease_recheck_deadline is None: + legacy_lease_recheck_deadline = loop.time() + max(self._sync_lease_seconds, 0.0) + remaining = legacy_lease_recheck_deadline - loop.time() + if remaining > 0: + await self._wait_for_claim_recheck(remaining) + else: + legacy_lease_recheck_deadline = None + await asyncio.sleep(self._sync_poll_interval) + continue + legacy_lease_recheck_deadline = None + expected_generations = {item.user.email: item.generation for item in claimed_users} + user_sync_lease = await self._acquire_user_sync_lease(list(expected_generations), expected_generations) + if user_sync_lease is not None: + allowed_keys = set(user_sync_lease.user_keys) + suppressed_claims = [item for item in claimed_users if item.user.email not in allowed_keys] + if suppressed_claims: + await self._ack_claimed_users(suppressed_claims) + claimed_users = [item for item in claimed_users if item.user.email in allowed_keys] + if user_sync_lease.token: + user_sync_heartbeat = asyncio.create_task(self._heartbeat_user_sync_lease(user_sync_lease)) + if not claimed_users: + await self._release_user_sync_lease(user_sync_lease, user_sync_heartbeat) + user_sync_lease = None + user_sync_heartbeat = None continue users = [item.user for item in claimed_users] + user_sync_remote_started = False + user_sync_remote_completed = False # Prefer chunked sync for large batches to reduce per-request overhead - use_chunked = supports_chunked and len(users) >= 1000 + chunked_transport = getattr(self, "_sync_users_chunked_transport", None) + use_chunked = supports_chunked and len(users) >= 1000 and callable(chunked_transport) + if supports_chunked and len(users) >= 1000 and not use_chunked: + self.logger.debug( + f"[{self.name}] Protected chunked transport is unavailable; using compatible batch sync" + ) if use_chunked: # Aim for ~10 chunks, cap size to 2000 to stay under server limits chunk_size = min(2000, max(1, math.ceil(len(users) / 10))) - failed_users = await self.sync_users_chunked( - users=users, chunk_size=chunk_size, flush_pending=False, timeout=self._internal_timeout - ) + try: + async with self._node_lock: + await self._assert_user_sync_lease_owned(user_sync_lease) + user_sync_remote_started = True + if self._accepts_user_sync_epoch(chunked_transport): + await chunked_transport( + users, + chunk_size, + self._internal_timeout, + self._user_sync_epoch_for_transport(user_sync_lease), + ) + else: + await chunked_transport(users, chunk_size, self._internal_timeout) + failed_users = [] + user_sync_remote_completed = True + except Exception as e: # noqa: BLE001 - preserve the worker's retry contract + if self._is_stale_user_sync_rejection(e): + await self._release_user_sync_lease(user_sync_lease, user_sync_heartbeat) + user_sync_lease = None + user_sync_heartbeat = None + user_sync_remote_completed = True + await self._requeue_claimed_users(claimed_users) + claimed_users = [] + await asyncio.sleep(retry_delay) + retry_delay = min(retry_delay * 2, max_retry_delay) + continue + error_type = type(e).__name__ + self.logger.warning( + f"[{self.name}] Chunked sync failed for {len(users)} user(s) | Error: {error_type} - {e!s}" + ) + failed_users = users if failed_users: self.logger.warning( f"[{self.name}] {len(failed_users)}/{len(users)} users failed to chunk-sync " f"(chunk_size={chunk_size})" ) failed_emails = {user.email for user in failed_users} + failed_claims = [item for item in claimed_users if item.user.email in failed_emails] + await self._retain_unknown_user_sync_lease_keys( + user_sync_lease, user_sync_heartbeat, list(failed_emails) + ) + user_sync_lease = None + user_sync_heartbeat = None await self._ack_claimed_users( [item for item in claimed_users if item.user.email not in failed_emails] ) - await self._requeue_claimed_users( - [item for item in claimed_users if item.user.email in failed_emails] - ) + claimed_users = failed_claims + await self._requeue_claimed_users(claimed_users) + claimed_users = [] await self._increment_user_sync_failure() await asyncio.sleep(retry_delay) retry_delay = min(retry_delay * 2, max_retry_delay) @@ -557,21 +1246,37 @@ async def _sync_worker(self): f"[{self.name}] Chunk-synced {len(users)} user(s) with chunk_size={chunk_size}" ) await self._ack_claimed_users(claimed_users) + claimed_users = [] await self._reset_user_sync_failure_count() retry_delay = 1.0 else: # Batch sync users individually try: - failed_users = await self._sync_batch_users(users) + await self._assert_user_sync_lease_owned(user_sync_lease) + user_sync_remote_started = True + if self._accepts_user_sync_epoch(self._sync_batch_users): + failed_users = await self._sync_batch_users( + users, + self._user_sync_epoch_for_transport(user_sync_lease), + ) + else: + failed_users = await self._sync_batch_users(users) + user_sync_remote_completed = not failed_users if failed_users: self.logger.warning(f"[{self.name}] {len(failed_users)}/{len(users)} users failed to sync") failed_emails = {user.email for user in failed_users} + failed_claims = [item for item in claimed_users if item.user.email in failed_emails] + await self._retain_unknown_user_sync_lease_keys( + user_sync_lease, user_sync_heartbeat, list(failed_emails) + ) + user_sync_lease = None + user_sync_heartbeat = None await self._ack_claimed_users( [item for item in claimed_users if item.user.email not in failed_emails] ) - await self._requeue_claimed_users( - [item for item in claimed_users if item.user.email in failed_emails] - ) + claimed_users = failed_claims + await self._requeue_claimed_users(claimed_users) + claimed_users = [] await self._increment_user_sync_failure() # Exponential backoff on partial failure await asyncio.sleep(retry_delay) @@ -579,29 +1284,72 @@ async def _sync_worker(self): else: self.logger.debug(f"[{self.name}] Synced {len(users)} user(s)") await self._ack_claimed_users(claimed_users) + claimed_users = [] await self._reset_user_sync_failure_count() retry_delay = 1.0 # Reset retry delay on success except Exception as e: + stale_epoch = self._is_stale_user_sync_rejection(e) error_type = type(e).__name__ self.logger.warning( f"[{self.name}] Batch sync failed for {len(users)} user(s), requeuing | " f"Error: {error_type} - {str(e)}" ) await self._increment_user_sync_failure() + if stale_epoch: + # Release the known-unapplied write before any + # fallible queue cleanup can replace this error. + await self._release_user_sync_lease(user_sync_lease, user_sync_heartbeat) + user_sync_lease = None + user_sync_heartbeat = None + user_sync_remote_completed = True await self._requeue_claimed_users(claimed_users) + claimed_users = [] + if not stale_epoch: + if user_sync_remote_started and not user_sync_remote_completed: + await self._abandon_user_sync_lease(user_sync_lease, user_sync_heartbeat) + else: + await self._release_user_sync_lease(user_sync_lease, user_sync_heartbeat) + user_sync_lease = None + user_sync_heartbeat = None # Exponential backoff on failure await asyncio.sleep(retry_delay) retry_delay = min(retry_delay * 2, max_retry_delay) + if user_sync_lease is not None: + await self._release_user_sync_lease(user_sync_lease, user_sync_heartbeat) + user_sync_lease = None + user_sync_heartbeat = None + except asyncio.CancelledError: + await self._recover_claimed_users(claimed_users, "worker cancellation") + claimed_users = [] + if user_sync_remote_started and not user_sync_remote_completed: + await self._abandon_user_sync_lease(user_sync_lease, user_sync_heartbeat) + else: + await self._release_user_sync_lease(user_sync_lease, user_sync_heartbeat) + user_sync_lease = None + user_sync_heartbeat = None self.logger.debug(f"[{self.name}] Sync worker cancelled") except Exception as e: error_type = type(e).__name__ self.logger.error( f"[{self.name}] Unexpected error in sync worker | Error: {error_type} - {str(e)}", exc_info=True ) + await self._recover_claimed_users(claimed_users, "worker error") + claimed_users = [] + if user_sync_remote_started and not user_sync_remote_completed: + await self._abandon_user_sync_lease(user_sync_lease, user_sync_heartbeat) + else: + await self._release_user_sync_lease(user_sync_lease, user_sync_heartbeat) + user_sync_lease = None + user_sync_heartbeat = None finally: + if user_sync_lease is not None: + if user_sync_remote_started and not user_sync_remote_completed: + await self._abandon_user_sync_lease(user_sync_lease, user_sync_heartbeat) + else: + await self._release_user_sync_lease(user_sync_lease, user_sync_heartbeat) self.logger.debug(f"[{self.name}] Sync worker finished") async def _make_json_request( @@ -659,9 +1407,15 @@ async def _run_coordinated_update( lease = await self._acquire_lifecycle_lease(operation) try: - return await self._make_json_request(method="POST", endpoint=endpoint, json=json) - finally: - await self._release_lifecycle_lease(lease) + response = await self._make_json_request(method="POST", endpoint=endpoint, json=json) + except BaseException: + # A timeout/cancellation after dispatch has an unknown remote + # outcome. Keep the shared lease as poison until an explicit + # lifecycle reconciliation observes the actual node state. + await self._stop_lifecycle_heartbeat(lease) + raise + await self._release_lifecycle_lease(lease) + return response async def update_node(self) -> BufferedResponse: """Trigger a node update via the REST API.""" diff --git a/PasarGuardNodeBridge/grpclib.py b/PasarGuardNodeBridge/grpclib.py index a9eb7d0..a107187 100644 --- a/PasarGuardNodeBridge/grpclib.py +++ b/PasarGuardNodeBridge/grpclib.py @@ -1,6 +1,6 @@ import asyncio import logging -from contextlib import asynccontextmanager +from contextlib import AsyncExitStack, asynccontextmanager from typing import AsyncGenerator from grpclib.client import Channel, Stream @@ -11,8 +11,8 @@ from PasarGuardNodeBridge.abstract_node import PasarGuardNode from PasarGuardNodeBridge.common import service_grpc from PasarGuardNodeBridge.common import service_pb2 as service -from PasarGuardNodeBridge.controller import Health, NodeAPIError -from PasarGuardNodeBridge.storage import LifecycleOperation, LifecycleStatus +from PasarGuardNodeBridge.controller import STALE_USER_SYNC_RETRY_LIMIT, Health, NodeAPIError +from PasarGuardNodeBridge.storage import LifecycleLeaseLostError, LifecycleOperation, LifecycleStatus from PasarGuardNodeBridge.utils import format_host_for_url, grpc_to_http_status @@ -139,6 +139,18 @@ async def _handle_grpc_request(self, method, request, timeout: int | None = None except Exception as e: self._handle_error(e) + @asynccontextmanager + async def _open_grpc_stream(self, method, timeout: float): + """Enter a grpclib stream context with a bounded establishment time.""" + stack = AsyncExitStack() + try: + stream = await asyncio.wait_for( + stack.enter_async_context(method.open(metadata=self._metadata)), timeout=timeout + ) + yield stream + finally: + await asyncio.wait_for(stack.aclose(), timeout=timeout) + async def start( self, config: str, @@ -147,6 +159,7 @@ async def start( keep_alive: int = 0, exclude_inbounds: list[str] = [], timeout: int | None = None, + reconcile_user_sync: bool = False, ) -> service.BaseInfoResponse | None: """Start the node with proper task management""" timeout = timeout or self._default_timeout @@ -154,40 +167,86 @@ async def start( if health is Health.INVALID: raise NodeAPIError(code=-4, detail="Invalid node") - req = service.Backend( - type=backend_type, config=config, users=users, keep_alive=keep_alive, exclude_inbounds=exclude_inbounds - ) - lease = await self._acquire_lifecycle_lease(LifecycleOperation.START) + user_sync_lease = None + user_sync_heartbeat = None + remote_started = False + remote_completed = False try: - async with self._node_lock: - info: service.BaseInfoResponse = await self._handle_grpc_request( - method=self._client.Start, - request=req, - timeout=timeout, - ) + requested_users = users + for attempt in range(STALE_USER_SYNC_RETRY_LIMIT + 1): + remote_started = False + remote_completed = False + if reconcile_user_sync: + filtered_users, user_sync_lease, user_sync_heartbeat = ( + await self._acquire_reconciliation_user_sync_lease(requested_users) + ) + else: + filtered_users, user_sync_lease, user_sync_heartbeat = ( + await self._acquire_snapshot_user_sync_lease(requested_users) + ) + try: + req = service.Backend( + type=backend_type, + config=config, + users=filtered_users, + keep_alive=keep_alive, + exclude_inbounds=exclude_inbounds, + user_sync_epoch=self._user_sync_epoch_for_transport(user_sync_lease), + ) + async with self._node_lock: + await self._assert_user_sync_lease_owned(user_sync_lease) + capability_generation = getattr(self, "_user_sync_connection_generation", 0) + remote_started = True + try: + info: service.BaseInfoResponse = await self._handle_grpc_request( + method=self._client.Start, + request=req, + timeout=timeout, + ) + except Exception as exc: + if self._is_stale_user_sync_rejection(exc): + remote_completed = True + if attempt < STALE_USER_SYNC_RETRY_LIMIT: + continue + raise + remote_completed = True - if not info.started: - raise NodeAPIError(500, "Failed to start the node") + if not info.started: + raise NodeAPIError(500, "Failed to start the node") - try: - await self.connect(info.node_version, info.core_version) - except Exception as e: - await self.disconnect() - self._handle_error(e) + await self._observe_user_sync_epoch_capability(info, capability_generation) - await self._release_lifecycle_lease( - lease, - LifecycleStatus.HEALTHY, - desired=LifecycleStatus.HEALTHY, - node_version=info.node_version, - core_version=info.core_version, - ) - return info - except BaseException: - await self._release_lifecycle_lease(lease, LifecycleStatus.BROKEN, desired=LifecycleStatus.HEALTHY) + try: + await self.connect(info.node_version, info.core_version) + except Exception as e: + await self.disconnect() + self._handle_error(e) + + await self._release_lifecycle_lease( + lease, + LifecycleStatus.HEALTHY, + desired=LifecycleStatus.HEALTHY, + node_version=info.node_version, + core_version=info.core_version, + ) + return info + finally: + if remote_started and not remote_completed: + await self._abandon_user_sync_lease(user_sync_lease, user_sync_heartbeat) + else: + await self._release_user_sync_lease(user_sync_lease, user_sync_heartbeat) + user_sync_lease = None + user_sync_heartbeat = None + except BaseException as exc: + if remote_started and (not remote_completed or isinstance(exc, LifecycleLeaseLostError)): + await self._stop_lifecycle_heartbeat(lease) + else: + await self._release_lifecycle_lease(lease, LifecycleStatus.BROKEN, desired=LifecycleStatus.HEALTHY) raise + raise AssertionError("unreachable") + async def stop(self, timeout: int | None = None) -> None: """Stop the node with proper cleanup""" timeout = timeout or self._default_timeout @@ -198,32 +257,30 @@ async def stop(self, timeout: int | None = None) -> None: lease = await self._acquire_lifecycle_lease(LifecycleOperation.STOP) try: async with self._node_lock: - await self.disconnect() - - try: - await self._handle_grpc_request( - method=self._client.Stop, - request=service.Empty(), - timeout=timeout, - ) - except Exception: - pass - await self._release_lifecycle_lease( - lease, LifecycleStatus.STOPPED, desired=LifecycleStatus.STOPPED + await self._handle_grpc_request( + method=self._client.Stop, + request=service.Empty(), + timeout=timeout, ) + await self.disconnect() + await self._release_lifecycle_lease(lease, LifecycleStatus.STOPPED, desired=LifecycleStatus.STOPPED) except BaseException: - await self._release_lifecycle_lease(lease, LifecycleStatus.BROKEN, desired=LifecycleStatus.STOPPED) + await self._stop_lifecycle_heartbeat(lease) raise finally: await self._json_client.close() async def info(self, timeout: int | None = None) -> service.BaseInfoResponse | None: timeout = timeout or self._default_timeout - return await self._handle_grpc_request( + capability_generation = getattr(self, "_user_sync_connection_generation", 0) + response = await self._handle_grpc_request( method=self._client.GetBaseInfo, request=service.Empty(), timeout=timeout, ) + if response is not None: + await self._observe_user_sync_epoch_capability(response, capability_generation) + return response async def get_system_stats(self, timeout: int | None = None) -> service.SystemStatsResponse | None: timeout = timeout or self._default_timeout @@ -278,18 +335,95 @@ async def get_user_online_ip_list( ) async def sync_users( - self, users: list[service.User], flush_pending: bool = False, timeout: int | None = None + self, + users: list[service.User], + flush_pending: bool = False, + timeout: int | None = None, + revocation_id: str | None = None, ) -> service.Empty | None: + if revocation_id is not None: + raise NodeAPIError(400, "sync_users is a full replacement; use sync_users_chunked for revocation writes") timeout = timeout or self._default_timeout if flush_pending: await self.flush_pending_users() - async with self._node_lock: - return await self._handle_grpc_request( - method=self._client.SyncUsers, - request=service.Users(users=users), - timeout=timeout, - ) + requested_users = users + for attempt in range(STALE_USER_SYNC_RETRY_LIMIT + 1): + filtered_users, lease, heartbeat = await self._acquire_snapshot_user_sync_lease(requested_users) + remote_started = False + remote_completed = False + try: + async with self._node_lock: + await self._assert_user_sync_lease_owned(lease) + remote_started = True + try: + response = await self._handle_grpc_request( + method=self._client.SyncUsers, + request=service.Users( + users=filtered_users, + user_sync_epoch=self._user_sync_epoch_for_transport(lease), + ), + timeout=timeout, + ) + except Exception as exc: + if self._is_stale_user_sync_rejection(exc): + remote_completed = True + if attempt < STALE_USER_SYNC_RETRY_LIMIT: + continue + raise + remote_completed = True + return response + finally: + if remote_started and not remote_completed: + await self._abandon_user_sync_lease(lease, heartbeat) + else: + await self._release_user_sync_lease(lease, heartbeat) + + raise AssertionError("unreachable") + + async def reconcile_users( + self, + users: list[service.User], + flush_pending: bool = False, + timeout: int | None = None, + ) -> service.Empty | None: + """Recover expired/unknown user writes with an authoritative snapshot.""" + timeout = timeout or self._default_timeout + if flush_pending: + await self.flush_pending_users() + requested_users = users + for attempt in range(STALE_USER_SYNC_RETRY_LIMIT + 1): + filtered_users, lease, heartbeat = await self._acquire_reconciliation_user_sync_lease(requested_users) + remote_started = False + remote_completed = False + try: + async with self._node_lock: + await self._assert_user_sync_lease_owned(lease) + remote_started = True + try: + response = await self._handle_grpc_request( + method=self._client.SyncUsers, + request=service.Users( + users=filtered_users, + user_sync_epoch=self._user_sync_epoch_for_transport(lease), + ), + timeout=timeout, + ) + except Exception as exc: + if self._is_stale_user_sync_rejection(exc): + remote_completed = True + if attempt < STALE_USER_SYNC_RETRY_LIMIT: + continue + raise + remote_completed = True + return response + finally: + if remote_started and not remote_completed: + await self._abandon_user_sync_lease(lease, heartbeat) + else: + await self._release_user_sync_lease(lease, heartbeat) + + raise AssertionError("unreachable") async def sync_users_chunked( self, @@ -297,6 +431,7 @@ async def sync_users_chunked( chunk_size: int = 100, flush_pending: bool = False, timeout: int | None = None, + revocation_id: str | None = None, ) -> list[service.User]: """Send users via the client-streaming SyncUsersChunked RPC. Returns failed users.""" if chunk_size <= 0: @@ -306,33 +441,74 @@ async def sync_users_chunked( if flush_pending: await self.flush_pending_users() - async with self._node_lock: + for attempt in range(STALE_USER_SYNC_RETRY_LIMIT + 1): + lease, heartbeat = await self._acquire_direct_user_sync_lease(users, revocation_id) + remote_started = False + remote_completed = False try: - async with self._client.SyncUsersChunked.open(metadata=self._metadata) as stream: - if not users: - await asyncio.wait_for( - stream.send_message(service.UsersChunk(index=0, last=True)), - timeout=self._internal_timeout, - ) - else: - total_users = len(users) - for index, start in enumerate(range(0, total_users, chunk_size)): - chunk_users = users[start : start + chunk_size] - is_last = start + chunk_size >= total_users - await asyncio.wait_for( - stream.send_message(service.UsersChunk(users=chunk_users, index=index, last=is_last)), - timeout=self._internal_timeout, - ) - - await stream.end() - await asyncio.wait_for(stream.recv_message(), timeout=timeout) - return [] - except Exception as e: + async with self._node_lock: + await self._assert_user_sync_lease_owned(lease) + remote_started = True + await self._sync_users_chunked_transport( + users, + chunk_size, + timeout, + self._user_sync_epoch_for_transport(lease), + ) + remote_completed = True + return [] + except Exception as e: # noqa: BLE001 - direct API reports the failed batch + stale_epoch = self._is_stale_user_sync_rejection(e) + if stale_epoch: + remote_completed = True + if attempt < STALE_USER_SYNC_RETRY_LIMIT: + continue error_type = type(e).__name__ self.logger.warning( - f"[{self.name}] Chunked gRPC sync failed for {len(users)} user(s) | Error: {error_type} - {str(e)}" + f"[{self.name}] Chunked gRPC sync failed for {len(users)} user(s) | " + f"Error: {error_type} - {e!s}" ) return users + finally: + if remote_started and not remote_completed: + await self._abandon_user_sync_lease(lease, heartbeat) + else: + await self._release_user_sync_lease(lease, heartbeat) + + raise AssertionError("unreachable") + + async def _sync_users_chunked_transport( + self, + users: list[service.User], + chunk_size: int, + timeout: int, + user_sync_epoch: int = 0, + ) -> None: + async with self._open_grpc_stream(self._client.SyncUsersChunked, timeout) as stream: + if not users: + await asyncio.wait_for( + stream.send_message(service.UsersChunk(index=0, last=True, user_sync_epoch=user_sync_epoch)), + timeout=self._internal_timeout, + ) + else: + total_users = len(users) + for index, start in enumerate(range(0, total_users, chunk_size)): + chunk_users = users[start : start + chunk_size] + is_last = start + chunk_size >= total_users + await asyncio.wait_for( + stream.send_message( + service.UsersChunk( + users=chunk_users, + index=index, + last=is_last, + user_sync_epoch=user_sync_epoch, + ) + ), + timeout=self._internal_timeout, + ) + + await asyncio.wait_for(stream.end(), timeout=timeout) + await asyncio.wait_for(stream.recv_message(), timeout=timeout) async def list_routing_rules(self, timeout: int | None = None) -> service.RoutingRulesResponse | None: timeout = timeout or self._default_timeout @@ -416,25 +592,42 @@ async def override_balancer_target( timeout=timeout, ) - async def _sync_batch_users(self, users: list[service.User]) -> list[service.User]: + async def _sync_batch_users(self, users: list[service.User], user_sync_epoch: int = 0) -> list[service.User]: """Sync users via gRPC SyncUser stream. Returns failed users.""" failed = [] try: - async with self._client.SyncUser.open(metadata=self._metadata) as stream: - for user in users: + async with self._open_grpc_stream(self._client.SyncUser, self._internal_timeout) as stream: + for index, user in enumerate(users): try: - await asyncio.wait_for(stream.send_message(user), timeout=self._internal_timeout) + await asyncio.wait_for( + stream.send_message( + service.User( + email=user.email, + proxies=user.proxies, + inbounds=user.inbounds, + user_sync_epoch=user_sync_epoch, + ) + ), + timeout=self._internal_timeout, + ) except Exception as e: + if self._is_stale_user_sync_rejection(e): + raise error_type = type(e).__name__ self.logger.warning( - f"[{self.name}] Failed to sync user {user.email} | Error: {error_type} - {str(e)}" + f"[{self.name}] Failed to sync user at batch index {index} | Error: {error_type} - {e!s}" ) - failed.append(user) - await stream.end() + # A send failure leaves the stream unusable. Account for + # the remaining users without retrying the broken stream. + failed.extend(users[index:]) + return failed + await asyncio.wait_for(stream.end(), timeout=self._internal_timeout) except Exception as e: + if self._is_stale_user_sync_rejection(e): + raise # Stream-level failure - all users failed error_type = type(e).__name__ - self.logger.error(f"[{self.name}] Stream failed | Error: {error_type} - {str(e)}") + self.logger.error(f"[{self.name}] Stream failed | Error: {error_type} - {e!s}") return users return failed @@ -467,7 +660,8 @@ async def _check_node_health(self): if retries >= max_retries: if last_health != Health.BROKEN: self.logger.error( - f"[{self.name}] Health check failed after {max_retries} retries, setting health to BROKEN | " + f"[{self.name}] Health check failed after {max_retries} retries, " + "setting health to BROKEN | " f"Error: {error_type} - {str(e)}" ) await self.set_health(Health.BROKEN) diff --git a/PasarGuardNodeBridge/rest.py b/PasarGuardNodeBridge/rest.py index 9011aa6..28c6563 100644 --- a/PasarGuardNodeBridge/rest.py +++ b/PasarGuardNodeBridge/rest.py @@ -9,8 +9,8 @@ from PasarGuardNodeBridge.abstract_node import PasarGuardNode from PasarGuardNodeBridge.aiohttp_compat import BufferedStatusError, LazyClientSession, buffer_response, make_timeout from PasarGuardNodeBridge.common import service_pb2 as service -from PasarGuardNodeBridge.controller import Health, NodeAPIError -from PasarGuardNodeBridge.storage import LifecycleOperation, LifecycleStatus +from PasarGuardNodeBridge.controller import STALE_USER_SYNC_RETRY_LIMIT, Health, NodeAPIError +from PasarGuardNodeBridge.storage import LifecycleLeaseLostError, LifecycleOperation, LifecycleStatus from PasarGuardNodeBridge.utils import format_host_for_url ProtoMessageT = TypeVar("ProtoMessageT", bound=Message) @@ -167,6 +167,7 @@ async def start( keep_alive: int = 0, exclude_inbounds: list[str] = [], timeout: int | None = None, + reconcile_user_sync: bool = False, ) -> service.BaseInfoResponse | None: """Start the node with proper task management""" timeout = timeout or self._default_timeout @@ -175,43 +176,90 @@ async def start( raise NodeAPIError(code=-4, detail="Invalid node") lease = await self._acquire_lifecycle_lease(LifecycleOperation.START) + user_sync_lease = None + user_sync_heartbeat = None + remote_started = False + remote_completed = False try: - async with self._node_lock: - response = await self._make_request( - method="POST", - endpoint="start", - timeout=timeout, - proto_message=service.Backend( + requested_users = users + for attempt in range(STALE_USER_SYNC_RETRY_LIMIT + 1): + remote_started = False + remote_completed = False + if reconcile_user_sync: + filtered_users, user_sync_lease, user_sync_heartbeat = ( + await self._acquire_reconciliation_user_sync_lease(requested_users) + ) + else: + filtered_users, user_sync_lease, user_sync_heartbeat = ( + await self._acquire_snapshot_user_sync_lease(requested_users) + ) + try: + request = service.Backend( type=backend_type, config=config, - users=users, + users=filtered_users, keep_alive=keep_alive, exclude_inbounds=exclude_inbounds, - ), - proto_response_class=service.BaseInfoResponse, - ) - - if not response.started: - raise NodeAPIError(500, "Failed to start the node") - - try: - await self.connect(response.node_version, response.core_version) - except BaseException as e: - await self.disconnect() - self._handle_error(e) + user_sync_epoch=self._user_sync_epoch_for_transport(user_sync_lease), + ) + async with self._node_lock: + await self._assert_user_sync_lease_owned(user_sync_lease) + capability_generation = getattr(self, "_user_sync_connection_generation", 0) + remote_started = True + try: + response = await self._make_request( + method="POST", + endpoint="start", + timeout=timeout, + proto_message=request, + proto_response_class=service.BaseInfoResponse, + ) + except Exception as exc: + if self._is_stale_user_sync_rejection(exc): + remote_completed = True + if attempt < STALE_USER_SYNC_RETRY_LIMIT: + continue + raise + remote_completed = True + + if not response.started: + raise NodeAPIError(500, "Failed to start the node") + + await self._observe_user_sync_epoch_capability(response, capability_generation) + + try: + await self.connect(response.node_version, response.core_version) + except asyncio.CancelledError: + await self.disconnect() + raise + except Exception as e: + await self.disconnect() + self._handle_error(e) - await self._release_lifecycle_lease( - lease, - LifecycleStatus.HEALTHY, - desired=LifecycleStatus.HEALTHY, - node_version=response.node_version, - core_version=response.core_version, - ) - return response - except BaseException: - await self._release_lifecycle_lease(lease, LifecycleStatus.BROKEN, desired=LifecycleStatus.HEALTHY) + await self._release_lifecycle_lease( + lease, + LifecycleStatus.HEALTHY, + desired=LifecycleStatus.HEALTHY, + node_version=response.node_version, + core_version=response.core_version, + ) + return response + finally: + if remote_started and not remote_completed: + await self._abandon_user_sync_lease(user_sync_lease, user_sync_heartbeat) + else: + await self._release_user_sync_lease(user_sync_lease, user_sync_heartbeat) + user_sync_lease = None + user_sync_heartbeat = None + except BaseException as exc: + if remote_started and (not remote_completed or isinstance(exc, LifecycleLeaseLostError)): + await self._stop_lifecycle_heartbeat(lease) + else: + await self._release_lifecycle_lease(lease, LifecycleStatus.BROKEN, desired=LifecycleStatus.HEALTHY) raise + raise AssertionError("unreachable") + async def stop(self, timeout: int | None = None) -> None: """Stop the node with proper cleanup""" timeout = timeout or self._default_timeout @@ -222,17 +270,11 @@ async def stop(self, timeout: int | None = None) -> None: lease = await self._acquire_lifecycle_lease(LifecycleOperation.STOP) try: async with self._node_lock: + await self._make_request(method="PUT", endpoint="stop", timeout=timeout) await self.disconnect() - - try: - await self._make_request(method="PUT", endpoint="stop", timeout=timeout) - except Exception: - pass - await self._release_lifecycle_lease( - lease, LifecycleStatus.STOPPED, desired=LifecycleStatus.STOPPED - ) + await self._release_lifecycle_lease(lease, LifecycleStatus.STOPPED, desired=LifecycleStatus.STOPPED) except BaseException: - await self._release_lifecycle_lease(lease, LifecycleStatus.BROKEN, desired=LifecycleStatus.STOPPED) + await self._stop_lifecycle_heartbeat(lease) raise finally: await self._client.close() @@ -240,9 +282,13 @@ async def stop(self, timeout: int | None = None) -> None: async def info(self, timeout: int | None = None) -> service.BaseInfoResponse | None: timeout = timeout or self._default_timeout - return await self._make_request( + capability_generation = getattr(self, "_user_sync_connection_generation", 0) + response = await self._make_request( method="GET", endpoint="info", timeout=timeout, proto_response_class=service.BaseInfoResponse ) + if response is not None: + await self._observe_user_sync_epoch_capability(response, capability_generation) + return response async def get_system_stats(self, timeout: int | None = None) -> service.SystemStatsResponse | None: timeout = timeout or self._default_timeout @@ -301,20 +347,99 @@ async def get_user_online_ip_list( ) async def sync_users( - self, users: list[service.User], flush_pending: bool = False, timeout: int | None = None + self, + users: list[service.User], + flush_pending: bool = False, + timeout: int | None = None, + revocation_id: str | None = None, ) -> service.Empty | None: + if revocation_id is not None: + raise NodeAPIError(400, "sync_users is a full replacement; use sync_users_chunked for revocation writes") timeout = timeout or self._default_timeout if flush_pending: await self.flush_pending_users() - async with self._node_lock: - return await self._make_request( - method="PUT", - endpoint="users/sync", - timeout=timeout, - proto_message=service.Users(users=users), - proto_response_class=service.Empty, - ) + requested_users = users + for attempt in range(STALE_USER_SYNC_RETRY_LIMIT + 1): + filtered_users, lease, heartbeat = await self._acquire_snapshot_user_sync_lease(requested_users) + remote_started = False + remote_completed = False + try: + async with self._node_lock: + await self._assert_user_sync_lease_owned(lease) + remote_started = True + try: + response = await self._make_request( + method="PUT", + endpoint="users/sync", + timeout=timeout, + proto_message=service.Users( + users=filtered_users, + user_sync_epoch=self._user_sync_epoch_for_transport(lease), + ), + proto_response_class=service.Empty, + ) + except Exception as exc: + if self._is_stale_user_sync_rejection(exc): + remote_completed = True + if attempt < STALE_USER_SYNC_RETRY_LIMIT: + continue + raise + remote_completed = True + return response + finally: + if remote_started and not remote_completed: + await self._abandon_user_sync_lease(lease, heartbeat) + else: + await self._release_user_sync_lease(lease, heartbeat) + + raise AssertionError("unreachable") + + async def reconcile_users( + self, + users: list[service.User], + flush_pending: bool = False, + timeout: int | None = None, + ) -> service.Empty | None: + """Recover expired/unknown user writes with an authoritative snapshot.""" + timeout = timeout or self._default_timeout + if flush_pending: + await self.flush_pending_users() + requested_users = users + for attempt in range(STALE_USER_SYNC_RETRY_LIMIT + 1): + filtered_users, lease, heartbeat = await self._acquire_reconciliation_user_sync_lease(requested_users) + remote_started = False + remote_completed = False + try: + async with self._node_lock: + await self._assert_user_sync_lease_owned(lease) + remote_started = True + try: + response = await self._make_request( + method="PUT", + endpoint="users/sync", + timeout=timeout, + proto_message=service.Users( + users=filtered_users, + user_sync_epoch=self._user_sync_epoch_for_transport(lease), + ), + proto_response_class=service.Empty, + ) + except Exception as exc: + if self._is_stale_user_sync_rejection(exc): + remote_completed = True + if attempt < STALE_USER_SYNC_RETRY_LIMIT: + continue + raise + remote_completed = True + return response + finally: + if remote_started and not remote_completed: + await self._abandon_user_sync_lease(lease, heartbeat) + else: + await self._release_user_sync_lease(lease, heartbeat) + + raise AssertionError("unreachable") async def sync_users_chunked( self, @@ -322,6 +447,7 @@ async def sync_users_chunked( chunk_size: int = 100, flush_pending: bool = False, timeout: int | None = None, + revocation_id: str | None = None, ) -> list[service.User]: """Stream UsersChunk messages over HTTP/2 for large sync operations. Returns failed users.""" if chunk_size <= 0: @@ -331,6 +457,49 @@ async def sync_users_chunked( if flush_pending: await self.flush_pending_users() + for attempt in range(STALE_USER_SYNC_RETRY_LIMIT + 1): + lease, heartbeat = await self._acquire_direct_user_sync_lease(users, revocation_id) + remote_started = False + remote_completed = False + try: + async with self._node_lock: + await self._assert_user_sync_lease_owned(lease) + remote_started = True + await self._sync_users_chunked_transport( + users, + chunk_size, + timeout, + self._user_sync_epoch_for_transport(lease), + ) + remote_completed = True + return [] + except Exception as e: # noqa: BLE001 - direct API reports the failed batch + stale_epoch = self._is_stale_user_sync_rejection(e) + if stale_epoch: + remote_completed = True + if attempt < STALE_USER_SYNC_RETRY_LIMIT: + continue + error_type = type(e).__name__ + self.logger.warning( + f"[{self.name}] Chunked REST sync failed for {len(users)} user(s) | " + f"Error: {error_type} - {e!s}" + ) + return users + finally: + if remote_started and not remote_completed: + await self._abandon_user_sync_lease(lease, heartbeat) + else: + await self._release_user_sync_lease(lease, heartbeat) + + raise AssertionError("unreachable") + + async def _sync_users_chunked_transport( + self, + users: list[service.User], + chunk_size: int, + timeout: int, + user_sync_epoch: int = 0, + ) -> None: def _encode_varint(value: int) -> bytes: encoded = bytearray() while True: @@ -344,9 +513,10 @@ def _encode_varint(value: int) -> bytes: return bytes(encoded) async def _iter_chunks(): - # Send a terminating empty chunk when no users are provided if not users: - chunk_bytes = self._serialize_protobuf(service.UsersChunk(index=0, last=True)) + chunk_bytes = self._serialize_protobuf( + service.UsersChunk(index=0, last=True, user_sync_epoch=user_sync_epoch) + ) yield _encode_varint(len(chunk_bytes)) + chunk_bytes return @@ -354,30 +524,24 @@ async def _iter_chunks(): for index, start in enumerate(range(0, total_users, chunk_size)): chunk_users = users[start : start + chunk_size] chunk_bytes = self._serialize_protobuf( - service.UsersChunk(users=chunk_users, index=index, last=start + chunk_size >= total_users) + service.UsersChunk( + users=chunk_users, + index=index, + last=start + chunk_size >= total_users, + user_sync_epoch=user_sync_epoch, + ) ) - # Length-prefix each protobuf chunk to preserve framing server-side yield _encode_varint(len(chunk_bytes)) + chunk_bytes - async with self._node_lock: - try: - async with self._client.request( - method="PUT", - url="users/sync/chunked", - data=_iter_chunks(), - timeout=make_timeout(timeout), - ) as raw_response: - response = await buffer_response(raw_response) - response.raise_for_status() - data = response.content - self._deserialize_protobuf(service.Empty, data) - return [] - except Exception as e: - error_type = type(e).__name__ - self.logger.warning( - f"[{self.name}] Chunked REST sync failed for {len(users)} user(s) | Error: {error_type} - {str(e)}" - ) - return users + async with self._client.request( + method="PUT", + url="users/sync/chunked", + data=_iter_chunks(), + timeout=make_timeout(timeout), + ) as raw_response: + response = await buffer_response(raw_response) + response.raise_for_status() + self._deserialize_protobuf(service.Empty, response.content) async def list_routing_rules(self, timeout: int | None = None) -> service.RoutingRulesResponse | None: timeout = timeout or self._default_timeout @@ -474,21 +638,30 @@ async def override_balancer_target( proto_response_class=service.Empty, ) - async def _sync_batch_users(self, users: list[service.User]) -> list[service.User]: + async def _sync_batch_users(self, users: list[service.User], user_sync_epoch: int = 0) -> list[service.User]: """Sync users individually via PUT user/sync. Returns failed users.""" failed = [] - for user in users: + for index, user in enumerate(users): try: await self._make_request( method="PUT", endpoint="user/sync", timeout=self._internal_timeout, - proto_message=user, + proto_message=service.User( + email=user.email, + proxies=user.proxies, + inbounds=user.inbounds, + user_sync_epoch=user_sync_epoch, + ), proto_response_class=service.Empty, ) except Exception as e: + if self._is_stale_user_sync_rejection(e): + raise error_type = type(e).__name__ - self.logger.warning(f"[{self.name}] Failed to sync user {user.email} | Error: {error_type} - {str(e)}") + self.logger.warning( + f"[{self.name}] Failed to sync user at batch index {index} | Error: {error_type} - {e!s}" + ) failed.append(user) return failed @@ -521,7 +694,8 @@ async def _check_node_health(self): if retries >= max_retries: if last_health != Health.BROKEN: self.logger.error( - f"[{self.name}] Health check failed after {max_retries} retries, setting health to BROKEN | " + f"[{self.name}] Health check failed after {max_retries} retries, " + "setting health to BROKEN | " f"Error: {error_type} - {str(e)}" ) await self.set_health(Health.BROKEN) @@ -635,11 +809,12 @@ async def _receive_logs(response: aiohttp.ClientResponse) -> None: pass response = None + stream_task = None try: self.logger.debug(f"[{self.name}] Opening on-demand log stream") response = await self._client.get("/logs", timeout=make_timeout(None)) - if response.status >= 400: - buffered_response = await buffer_response(response) + if response.status >= 300: + buffered_response = await asyncio.wait_for(buffer_response(response), timeout=self._internal_timeout) buffered_response.raise_for_status() self.logger.debug(f"[{self.name}] On-demand log stream opened successfully") @@ -680,4 +855,9 @@ async def _receive_logs(response: aiohttp.ClientResponse) -> None: # Convert to NodeAPIError self._handle_error(e) finally: + if response is not None: + try: + response.close() + except Exception: + pass self.logger.debug(f"[{self.name}] On-demand log stream closed") diff --git a/PasarGuardNodeBridge/storage.py b/PasarGuardNodeBridge/storage.py index 374244f..086d838 100644 --- a/PasarGuardNodeBridge/storage.py +++ b/PasarGuardNodeBridge/storage.py @@ -1,6 +1,6 @@ import asyncio import time -from dataclasses import asdict, dataclass, field +from dataclasses import asdict, dataclass, field, replace from enum import Enum from typing import Any, Protocol from uuid import uuid4 @@ -35,6 +35,58 @@ def from_dict(cls, data: dict[str, Any]) -> "NodeConfig": class ClaimedUser: token: str user: User + generation: int = 0 + + +@dataclass(slots=True) +class UserSyncLease: + """A revocation-aware permit for applying user updates to a node.""" + + node_id: str + worker_id: str + token: str + user_keys: tuple[str, ...] + generations: dict[str, int] + lease_seconds: float = 30.0 + revocation_id: str | None = None + covers_all_users: bool = False + epoch: int = 0 + + +@dataclass(slots=True) +class StartupUserSyncLease: + """A node-wide startup snapshot permit and the keys safe to include.""" + + lease: UserSyncLease + included_user_keys: tuple[str, ...] + + +@dataclass(slots=True) +class UserRevocationResult: + """Keys acquired by this operation and keys already permanently revoked.""" + + active_user_keys: tuple[str, ...] + finalized_user_keys: tuple[str, ...] + + +class UserSyncStoreFullError(RuntimeError): + """Raised when accepting more distinct pending users would exceed the store bound.""" + + +class UserSyncLeaseLostError(RuntimeError): + """Raised when an expired execution lease leaves a user-sync outcome unknown.""" + + +class LifecycleLeaseLostError(RuntimeError): + """Raised when lifecycle ownership is lost before the remote effect is known.""" + + +class UserRevocationConflictError(RuntimeError): + """Raised when another operation owns one or more requested user fences.""" + + def __init__(self, conflicting_user_keys: tuple[str, ...]): + self.conflicting_user_keys = conflicting_user_keys + super().__init__(f"another revocation owns {len(conflicting_user_keys)} user fence(s)") class NodeRegistryProtocol(Protocol): @@ -51,11 +103,81 @@ async def claim_users( self, node_id: str, worker_id: str, limit: int, lease_seconds: float ) -> list[ClaimedUser]: ... + async def next_claim_delay(self, node_id: str) -> float | None: ... + async def ack_users(self, node_id: str, tokens: list[str]) -> None: ... async def requeue_users(self, node_id: str, claimed_users: list[ClaimedUser]) -> None: ... async def clear(self, node_id: str) -> None: ... +class RevocationAwareUserSyncStoreProtocol(UserSyncStoreProtocol, Protocol): + """Optional store capability required for coordinated permanent revocation. + + ``begin_user_revocation`` must atomically fence the supplied keys, discard + their pending and claimed payloads, and wait for intersecting unexpired + execution leases to drain. Implementations must also reject enqueue, claim, + and requeue attempts for fenced keys or obsolete generations. Expired + execution leases have an unknown remote outcome and must fail revocation + closed until explicitly reconciled; they cannot be silently discarded. + + Distinct operation IDs must be serialized for intersecting keys. A store + may fail fast with ``UserRevocationConflictError``; it must not leave a + partially acquired multi-key fence. Abort and finalize close authorized + admission before draining all intersecting execution leases. + + Passing ``revocation_id`` to ``acquire_user_sync_lease`` authorizes only + keys provisionally owned by that operation. This is required for the direct + removal/restore writes performed while the ordinary update fence is active. + """ + + async def begin_user_revocation( + self, node_id: str, user_keys: list[str], revocation_id: str + ) -> UserRevocationResult: ... + + async def abort_user_revocation(self, node_id: str, user_keys: list[str], revocation_id: str) -> None: ... + + async def finalize_user_revocation(self, node_id: str, user_keys: list[str], revocation_id: str) -> None: ... + + async def acquire_user_sync_lease( + self, + node_id: str, + worker_id: str, + user_keys: list[str], + lease_seconds: float, + expected_generations: dict[str, int] | None = None, + revocation_id: str | None = None, + ) -> UserSyncLease: ... + + async def acquire_startup_user_sync_lease( + self, + node_id: str, + worker_id: str, + user_keys: list[str], + lease_seconds: float, + ) -> StartupUserSyncLease: ... + + async def acquire_user_sync_reconciliation_lease( + self, + node_id: str, + worker_id: str, + user_keys: list[str], + lease_seconds: float, + ) -> StartupUserSyncLease: + """Replace expired unknown writes with one authoritative snapshot.""" + ... + + async def advance_user_sync_epoch(self, node_id: str, minimum_epoch: int) -> None: + """Atomically advance the next-epoch floor from a Node handshake.""" + ... + + async def retain_user_sync_lease_keys( + self, lease: UserSyncLease, retained_user_keys: list[str] + ) -> UserSyncLease: ... + + async def heartbeat_user_sync_lease(self, lease: UserSyncLease) -> bool: ... + async def release_user_sync_lease(self, lease: UserSyncLease) -> None: ... + + class LifecycleOperation(str, Enum): START = "start" STOP = "stop" @@ -103,9 +225,13 @@ async def try_acquire( ) -> LifecycleLease | None: ... async def release(self, lease: LifecycleLease, state_update: NodeLifecycleState | None = None) -> None: ... - async def heartbeat(self, lease: LifecycleLease) -> None: ... + async def heartbeat(self, lease: LifecycleLease) -> bool: ... async def get_state(self, node_id: str) -> NodeLifecycleState | None: ... + async def reconcile(self, node_id: str, observed: LifecycleStatus) -> bool: + """Clear an expired unknown lease after the remote state was inspected.""" + ... + async def update_observed( self, node_id: str, observed: LifecycleStatus, expected_epoch: int | None = None ) -> None: ... @@ -133,19 +259,154 @@ async def list_nodes(self) -> list[str]: return list(self._nodes) +@dataclass(slots=True) +class _UserRevocationState: + generation: int = 0 + active_owner: str | None = None + closing: bool = False + finalized: bool = False + + class InMemoryUserSyncStore: - def __init__(self): - self._pending: dict[str, dict[str, User]] = {} - self._claimed: dict[str, dict[str, tuple[User, float]]] = {} + def __init__(self, max_pending_users_per_node: int = 10_000): + if max_pending_users_per_node <= 0: + raise ValueError("max_pending_users_per_node must be positive") + self._max_pending_users_per_node = max_pending_users_per_node + self._pending: dict[str, dict[str, tuple[User, int]]] = {} + self._claimed: dict[str, dict[str, tuple[User, int, float]]] = {} self._lock = asyncio.Lock() + self._lease_changed = asyncio.Condition(self._lock) + self._revocations: dict[str, dict[str, _UserRevocationState]] = {} + self._user_sync_leases: dict[str, tuple[UserSyncLease, float]] = {} + self._user_sync_epochs: dict[str, int] = {} + self._startup_pending: set[str] = set() + + @staticmethod + def _unique_user_keys(user_keys: list[str]) -> tuple[str, ...]: + return tuple(dict.fromkeys(user_keys)) + + def _revocation_state(self, node_id: str, user_key: str) -> _UserRevocationState: + states = self._revocations.setdefault(node_id, {}) + return states.setdefault(user_key, _UserRevocationState()) + + def _is_user_fenced(self, node_id: str, user_key: str) -> bool: + state = self._revocations.get(node_id, {}).get(user_key) + return state is not None and (state.finalized or state.active_owner is not None) + + def _is_user_finalized(self, node_id: str, user_key: str) -> bool: + """Read a tombstone without allocating state for an unseen user.""" + state = self._revocations.get(node_id, {}).get(user_key) + return state is not None and state.finalized + + def _user_generation(self, node_id: str, user_key: str) -> int: + state = self._revocations.get(node_id, {}).get(user_key) + return state.generation if state is not None else 0 + + def _next_user_sync_epoch(self, node_id: str) -> int: + epoch = self._user_sync_epochs.get(node_id, 0) + 1 + self._user_sync_epochs[node_id] = epoch + return epoch + + def _next_intersecting_lease_delay( + self, + node_id: str, + user_keys: set[str], + now: float, + revocation_id: str | None = None, + ) -> float | None: + delays = [ + max(0.0, expires_at - now) + for lease, expires_at in self._user_sync_leases.values() + if lease.node_id == node_id + and (lease.covers_all_users or user_keys.intersection(lease.user_keys)) + and (revocation_id is None or lease.revocation_id == revocation_id) + ] + return min(delays) if delays else None + + def _has_lost_user_sync_lease( + self, + node_id: str, + user_keys: set[str], + now: float, + revocation_id: str | None = None, + ) -> bool: + return any( + lease.node_id == node_id + and expires_at <= now + and (lease.covers_all_users or user_keys.intersection(lease.user_keys)) + and (revocation_id is None or lease.revocation_id == revocation_id) + for lease, expires_at in self._user_sync_leases.values() + ) + + async def _wait_for_user_sync_leases( + self, + node_id: str, + user_keys: set[str], + revocation_id: str | None = None, + covers_all_users: bool = False, + ) -> None: + while True: + now = time.monotonic() + if covers_all_users: + matching_leases = [ + (lease, expires_at) + for lease, expires_at in self._user_sync_leases.values() + if lease.node_id == node_id and (revocation_id is None or lease.revocation_id == revocation_id) + ] + if any(expires_at <= now for _, expires_at in matching_leases): + raise UserSyncLeaseLostError("an expired user-sync execution lease has an unknown remote outcome") + delay = min((max(0.0, expires_at - now) for _, expires_at in matching_leases), default=None) + elif self._has_lost_user_sync_lease(node_id, user_keys, now, revocation_id): + raise UserSyncLeaseLostError("an expired user-sync execution lease has an unknown remote outcome") + else: + delay = self._next_intersecting_lease_delay(node_id, user_keys, now, revocation_id) + if delay is None: + return + try: + await asyncio.wait_for(self._lease_changed.wait(), timeout=max(delay, 0.001)) + except TimeoutError: + pass + + async def _wait_for_startup_completion(self, node_id: str) -> None: + """Wait until no startup is acquiring or holding a node-wide permit.""" + while True: + if node_id in self._startup_pending: + await self._lease_changed.wait() + continue + now = time.monotonic() + startup_leases = [ + (lease, expires_at) + for lease, expires_at in self._user_sync_leases.values() + if lease.node_id == node_id and lease.covers_all_users + ] + if any(expires_at <= now for _, expires_at in startup_leases): + raise UserSyncLeaseLostError("an expired startup execution lease has an unknown remote outcome") + if not startup_leases: + return + delay = min(max(0.0, expires_at - now) for _, expires_at in startup_leases) + try: + await asyncio.wait_for(self._lease_changed.wait(), timeout=max(delay, 0.001)) + except TimeoutError: + pass async def enqueue_users(self, node_id: str, users: list[User]) -> None: if not users: return async with self._lock: pending = self._pending.setdefault(node_id, {}) - for user in users: - pending[user.email] = user + claimed = self._claimed.setdefault(node_id, {}) + latest_users = {user.email: user for user in users if not self._is_user_fenced(node_id, user.email)} + if not latest_users: + return + tracked_emails = set(pending) + tracked_emails.update(user.email for user, _, _ in claimed.values()) + new_emails = set(latest_users).difference(tracked_emails) + if len(tracked_emails) + len(new_emails) > self._max_pending_users_per_node: + raise UserSyncStoreFullError( + f"pending user sync limit reached for node ({self._max_pending_users_per_node})" + ) + for user in latest_users.values(): + pending[user.email] = (user, self._user_generation(node_id, user.email)) async def claim_users(self, node_id: str, worker_id: str, limit: int, lease_seconds: float) -> list[ClaimedUser]: if limit <= 0: @@ -155,21 +416,41 @@ async def claim_users(self, node_id: str, worker_id: str, limit: int, lease_seco pending = self._pending.setdefault(node_id, {}) claimed = self._claimed.setdefault(node_id, {}) - for token, (user, expires_at) in list(claimed.items()): + for token, (user, generation, expires_at) in list(claimed.items()): if expires_at <= now: - pending.setdefault(user.email, user) + if not self._is_user_fenced(node_id, user.email) and generation == self._user_generation( + node_id, user.email + ): + pending.setdefault(user.email, (user, generation)) del claimed[token] result: list[ClaimedUser] = [] - for email, user in list(pending.items()): + for email, (user, generation) in list(pending.items()): + if self._is_user_fenced(node_id, email) or generation != self._user_generation(node_id, email): + del pending[email] + continue token = f"{worker_id}:{uuid4()}" - claimed[token] = (user, now + lease_seconds) - result.append(ClaimedUser(token=token, user=user)) + claimed[token] = (user, generation, now + lease_seconds) + result.append(ClaimedUser(token=token, user=user, generation=generation)) del pending[email] if len(result) >= limit: break return result + async def next_claim_delay(self, node_id: str) -> float | None: + """Return when tracked work can next be claimed, or ``None`` when none exists.""" + now = time.monotonic() + async with self._lock: + pending = self._pending.get(node_id) + if pending: + return 0.0 + + claimed = self._claimed.get(node_id) + if not claimed: + return None + + return max(0.0, min(expires_at for _, _, expires_at in claimed.values()) - now) + async def ack_users(self, node_id: str, tokens: list[str]) -> None: if not tokens: return @@ -185,14 +466,333 @@ async def requeue_users(self, node_id: str, claimed_users: list[ClaimedUser]) -> pending = self._pending.setdefault(node_id, {}) claimed = self._claimed.setdefault(node_id, {}) for item in claimed_users: - claimed.pop(item.token, None) - pending.setdefault(item.user.email, item.user) + owned_claim = claimed.pop(item.token, None) + if owned_claim is not None: + user, generation, _ = owned_claim + if ( + generation == item.generation + and generation == self._user_generation(node_id, user.email) + and not self._is_user_fenced(node_id, user.email) + ): + pending.setdefault(user.email, (user, generation)) async def clear(self, node_id: str) -> None: async with self._lock: self._pending.pop(node_id, None) self._claimed.pop(node_id, None) + async def begin_user_revocation( + self, node_id: str, user_keys: list[str], revocation_id: str + ) -> UserRevocationResult: + if not revocation_id: + raise ValueError("revocation_id must not be empty") + unique_keys = self._unique_user_keys(user_keys) + if not unique_keys: + return UserRevocationResult((), ()) + + async with self._lease_changed: + await self._wait_for_startup_completion(node_id) + active_keys = tuple( + user_key for user_key in unique_keys if not self._revocation_state(node_id, user_key).finalized + ) + states = [self._revocation_state(node_id, user_key) for user_key in active_keys] + conflicting_keys = tuple( + user_key + for user_key, state in zip(active_keys, states) + if state.closing or state.active_owner not in (None, revocation_id) + ) + if conflicting_keys: + raise UserRevocationConflictError(conflicting_keys) + + for state in states: + if state.active_owner is None: + state.generation += 1 + state.active_owner = revocation_id + + pending = self._pending.setdefault(node_id, {}) + for user_key in active_keys: + pending.pop(user_key, None) + + claimed = self._claimed.setdefault(node_id, {}) + for token, (user, _, _) in list(claimed.items()): + if user.email in active_keys: + claimed.pop(token, None) + + await self._wait_for_user_sync_leases(node_id, set(active_keys)) + finalized_keys = tuple(user_key for user_key in unique_keys if user_key not in active_keys) + return UserRevocationResult(active_keys, finalized_keys) + + async def abort_user_revocation(self, node_id: str, user_keys: list[str], revocation_id: str) -> None: + if not revocation_id: + raise ValueError("revocation_id must not be empty") + unique_keys = self._unique_user_keys(user_keys) + async with self._lease_changed: + affected_keys = { + user_key + for user_key in unique_keys + if (state := self._revocations.get(node_id, {}).get(user_key)) is not None + and not state.finalized + and state.active_owner == revocation_id + } + if not affected_keys: + return + for user_key in affected_keys: + self._revocation_state(node_id, user_key).closing = True + try: + await self._wait_for_user_sync_leases(node_id, affected_keys) + except BaseException: + for user_key in affected_keys: + state = self._revocation_state(node_id, user_key) + if state.active_owner == revocation_id and not state.finalized: + state.closing = False + if affected_keys: + self._lease_changed.notify_all() + raise + for user_key in affected_keys: + state = self._revocation_state(node_id, user_key) + state.active_owner = None + state.closing = False + if affected_keys: + self._lease_changed.notify_all() + + async def finalize_user_revocation(self, node_id: str, user_keys: list[str], revocation_id: str) -> None: + if not revocation_id: + raise ValueError("revocation_id must not be empty") + unique_keys = self._unique_user_keys(user_keys) + async with self._lease_changed: + affected_keys = { + user_key + for user_key in unique_keys + if (state := self._revocations.get(node_id, {}).get(user_key)) is not None + and not state.finalized + and state.active_owner == revocation_id + } + if not affected_keys: + return + for user_key in affected_keys: + state = self._revocation_state(node_id, user_key) + state.closing = True + try: + await self._wait_for_user_sync_leases(node_id, affected_keys) + except BaseException: + for user_key in affected_keys: + state = self._revocation_state(node_id, user_key) + if state.active_owner == revocation_id and not state.finalized: + state.closing = False + if affected_keys: + self._lease_changed.notify_all() + raise + for user_key in affected_keys: + state = self._revocation_state(node_id, user_key) + state.finalized = True + state.active_owner = None + state.closing = False + if affected_keys: + self._lease_changed.notify_all() + + async def acquire_user_sync_lease( + self, + node_id: str, + worker_id: str, + user_keys: list[str], + lease_seconds: float, + expected_generations: dict[str, int] | None = None, + revocation_id: str | None = None, + ) -> UserSyncLease: + unique_keys = self._unique_user_keys(user_keys) + async with self._lease_changed: + await self._wait_for_startup_completion(node_id) + now = time.monotonic() + allowed_keys = tuple( + key + for key in unique_keys + if ( + ( + (state := self._revocation_state(node_id, key)).active_owner == revocation_id + and not state.closing + and not state.finalized + ) + if revocation_id is not None + else not self._is_user_fenced(node_id, key) + ) + and ( + expected_generations is None or expected_generations.get(key) == self._user_generation(node_id, key) + ) + ) + generations = {key: self._user_generation(node_id, key) for key in allowed_keys} + token = f"{worker_id}:{uuid4()}" if allowed_keys else "" + lease = UserSyncLease( + node_id=node_id, + worker_id=worker_id, + token=token, + user_keys=allowed_keys, + generations=generations, + epoch=self._next_user_sync_epoch(node_id) if token else 0, + lease_seconds=lease_seconds, + revocation_id=revocation_id, + ) + if token: + self._user_sync_leases[token] = (lease, now + lease_seconds) + return lease + + async def advance_user_sync_epoch(self, node_id: str, minimum_epoch: int) -> None: + if minimum_epoch < 0: + raise ValueError("minimum_epoch must be non-negative") + async with self._lock: + self._user_sync_epochs[node_id] = max( + self._user_sync_epochs.get(node_id, 0), + minimum_epoch, + ) + + async def acquire_startup_user_sync_lease( + self, + node_id: str, + worker_id: str, + user_keys: list[str], + lease_seconds: float, + ) -> StartupUserSyncLease: + """Serialize a full replacement snapshot with every per-user write. + + A startup snapshot is node-wide: a key omitted from the payload is a + deletion just as surely as a key present in it is an update. Wait for + provisional revocations to reach abort/finalize, then atomically take a + wildcard execution lease before deciding which finalized keys to omit. + """ + unique_keys = self._unique_user_keys(user_keys) + async with self._lease_changed: + while node_id in self._startup_pending: + await self._lease_changed.wait() + while any(state.active_owner is not None for state in self._revocations.get(node_id, {}).values()): + await self._lease_changed.wait() + while node_id in self._startup_pending: + await self._lease_changed.wait() + self._startup_pending.add(node_id) + try: + await self._wait_for_user_sync_leases(node_id, set(), covers_all_users=True) + + now = time.monotonic() + included_keys = tuple(key for key in unique_keys if not self._is_user_finalized(node_id, key)) + generations = {key: self._user_generation(node_id, key) for key in unique_keys} + token = f"{worker_id}:{uuid4()}" + lease = UserSyncLease( + node_id=node_id, + worker_id=worker_id, + token=token, + user_keys=unique_keys, + generations=generations, + epoch=self._next_user_sync_epoch(node_id), + lease_seconds=lease_seconds, + covers_all_users=True, + ) + self._user_sync_leases[token] = (lease, now + lease_seconds) + return StartupUserSyncLease(lease=lease, included_user_keys=included_keys) + finally: + self._startup_pending.discard(node_id) + self._lease_changed.notify_all() + + async def acquire_user_sync_reconciliation_lease( + self, + node_id: str, + worker_id: str, + user_keys: list[str], + lease_seconds: float, + ) -> StartupUserSyncLease: + """Atomically supersede expired unknown writes with a full snapshot. + + Ordinary startup remains fail-closed on an expired execution lease. + This explicit recovery path waits for all still-live writes, removes + only expired records for this node, and then holds a node-wide permit + while the caller applies an authoritative replacement snapshot. If + that replacement is ambiguous, its own wildcard lease is abandoned and + recovery remains fail-closed. + """ + unique_keys = self._unique_user_keys(user_keys) + async with self._lease_changed: + while node_id in self._startup_pending: + await self._lease_changed.wait() + while any(state.active_owner is not None for state in self._revocations.get(node_id, {}).values()): + await self._lease_changed.wait() + while node_id in self._startup_pending: + await self._lease_changed.wait() + self._startup_pending.add(node_id) + try: + while True: + now = time.monotonic() + node_leases = [ + (token, lease, expires_at) + for token, (lease, expires_at) in self._user_sync_leases.items() + if lease.node_id == node_id + ] + live = [(lease, expires_at) for _, lease, expires_at in node_leases if expires_at > now] + if not live: + for token, _, _ in node_leases: + self._user_sync_leases.pop(token, None) + break + delay = min(max(0.001, expires_at - now) for _, expires_at in live) + try: + await asyncio.wait_for(self._lease_changed.wait(), timeout=delay) + except TimeoutError: + pass + + now = time.monotonic() + included_keys = tuple(key for key in unique_keys if not self._is_user_finalized(node_id, key)) + generations = {key: self._user_generation(node_id, key) for key in unique_keys} + token = f"{worker_id}:{uuid4()}" + lease = UserSyncLease( + node_id=node_id, + worker_id=worker_id, + token=token, + user_keys=unique_keys, + generations=generations, + epoch=self._next_user_sync_epoch(node_id), + lease_seconds=lease_seconds, + covers_all_users=True, + ) + self._user_sync_leases[token] = (lease, now + lease_seconds) + return StartupUserSyncLease(lease=lease, included_user_keys=included_keys) + finally: + self._startup_pending.discard(node_id) + self._lease_changed.notify_all() + + async def retain_user_sync_lease_keys(self, lease: UserSyncLease, retained_user_keys: list[str]) -> UserSyncLease: + """Atomically narrow a lease after a partial remote outcome.""" + retained_keys = self._unique_user_keys(retained_user_keys) + if lease.covers_all_users: + raise ValueError("a node-wide startup lease cannot be narrowed") + if not retained_keys or not set(retained_keys).issubset(lease.user_keys): + raise ValueError("retained_user_keys must be a non-empty subset of the lease") + async with self._lease_changed: + current = self._user_sync_leases.get(lease.token) + if current is None or current[0] != lease: + raise UserSyncLeaseLostError("user-sync execution lease is no longer owned") + narrowed = replace( + lease, + user_keys=retained_keys, + generations={key: lease.generations[key] for key in retained_keys}, + ) + self._user_sync_leases[lease.token] = (narrowed, current[1]) + self._lease_changed.notify_all() + return narrowed + + async def heartbeat_user_sync_lease(self, lease: UserSyncLease) -> bool: + if not lease.token: + return False + async with self._lock: + current = self._user_sync_leases.get(lease.token) + if current is not None and current[0] == lease and current[1] > time.monotonic(): + self._user_sync_leases[lease.token] = (lease, time.monotonic() + lease.lease_seconds) + return True + return False + + async def release_user_sync_lease(self, lease: UserSyncLease) -> None: + if not lease.token: + return + async with self._lease_changed: + current = self._user_sync_leases.get(lease.token) + if current is not None and current[0] == lease: + self._user_sync_leases.pop(lease.token, None) + self._lease_changed.notify_all() + class InMemoryNodeLifecycleCoordinator: def __init__(self): @@ -207,10 +807,10 @@ async def try_acquire( async with self._lock: current = self._leases.get(node_id) if current is not None: - _, expires_at = current - if expires_at > now: - return None - self._leases.pop(node_id, None) + # Never steal an expired lease. Its remote request may still + # complete after the local timeout, so allowing a second + # lifecycle effect would violate operation ordering. + return None state = self._states.get(node_id) or NodeLifecycleState(updated_at=now) epoch = state.epoch + 1 @@ -251,12 +851,37 @@ async def release(self, lease: LifecycleLease, state_update: NodeLifecycleState state.updated_at = now self._states[lease.node_id] = state - async def heartbeat(self, lease: LifecycleLease) -> None: + async def heartbeat(self, lease: LifecycleLease) -> bool: now = time.monotonic() async with self._lock: current = self._leases.get(lease.node_id) - if current is not None and current[0].token == lease.token: + if current is not None and current[0].token == lease.token and current[1] > now: self._leases[lease.node_id] = (lease, now + lease.lease_seconds) + return True + return False + + async def reconcile(self, node_id: str, observed: LifecycleStatus) -> bool: + """Clear only an expired lease after an authoritative state probe.""" + now = time.monotonic() + async with self._lock: + current = self._leases.get(node_id) + if current is not None and current[1] > now: + return False + if current is not None: + lease = current[0] + epoch = lease.epoch + 1 + self._leases.pop(node_id, None) + else: + epoch = (self._states.get(node_id) or NodeLifecycleState()).epoch + 1 + state = self._states.get(node_id) or NodeLifecycleState() + state.epoch = epoch + state.desired = observed + state.observed = observed + state.operation = None + state.owner = None + state.updated_at = now + self._states[node_id] = state + return True async def get_state(self, node_id: str) -> NodeLifecycleState | None: async with self._lock: diff --git a/README.md b/README.md index 89c6000..357162d 100644 --- a/README.md +++ b/README.md @@ -21,7 +21,8 @@ pip install pasarguard-node-bridge - Python `>=3.12` - A reachable PasarGuard node - Node service port (`port`) for gRPC or protobuf-REST -- Node JSON API port (`api_port`) for maintenance endpoints +- Node JSON API port (`api_port`) for maintenance endpoints. When omitted, the + public factory uses the service `port` for backwards compatibility. - Server CA certificate content (PEM string) - API key (UUID string) @@ -55,7 +56,7 @@ node = Bridge.create_node( - `connection`: `Bridge.NodeType.grpc` or `Bridge.NodeType.rest` - `address`: node host/IP - `port`: node service port -- `api_port`: node REST JSON API port +- `api_port`: optional node REST JSON API port; defaults to `port` - `server_ca`: PEM certificate content as string - `api_key`: UUID string - `name`: optional logger name @@ -126,6 +127,8 @@ await node.stop() ### 1. Queue-Based User Updates (recommended for frequent updates) `update_user` and `update_users` enqueue users and a background worker handles retries and batching. +If the configured per-node queue bound is reached, both methods raise +`Bridge.UserSyncStoreFullError`; callers may retry after queued work is processed. ```python await node.update_user(user) @@ -158,11 +161,86 @@ A `UserSyncStoreProtocol` implementation must provide these async methods: - `enqueue_users(node_id, users)` stores latest user payloads by email. - `claim_users(node_id, worker_id, limit, lease_seconds)` atomically leases work and returns `ClaimedUser` items. +- `next_claim_delay(node_id)` returns seconds until tracked work can next be claimed, `0.0` for immediately + claimable work, or `None` when no work is tracked. This lets an idle worker wake after another worker's lease expires. - `ack_users(node_id, tokens)` removes successfully synced claims. - `requeue_users(node_id, claimed_users)` makes failed claims available again. - `clear(node_id)` clears pending and claimed updates for a node. Delivery is at-least-once. A crashed worker may cause the same latest user payload to be synced again after its lease expires, so external adapters should use atomic claim/lease operations such as Redis Lua/transactions or NATS KV revision compare-and-set. +For compatibility, stores without `next_claim_delay` are rechecked after at most one configured lease interval, but new +implementations should provide it so workers can preserve the normal fast idle exit when the store is truly empty. + +#### Coordinated Permanent User Revocation + +Permanent deletion requires stronger ordering than the ordinary at-least-once queue. A +`RevocationAwareUserSyncStoreProtocol` implementation provides per-user generations, execution leases, and +operation-owned fences. Stores without this optional capability continue to support normal updates, while calls to the +revocation API fail closed with `NodeAPIError` code `501`. + +```python +revocation_id = "delete-request-42" +user_keys = [user.email] + +barrier = await node.begin_user_revocation(user_keys, revocation_id) +users_to_remove = [removed_user] if removed_user.email in barrier.active_user_keys else [] +try: + if users_to_remove: + failed = await node.sync_users_chunked(users_to_remove, revocation_id=revocation_id) + if failed: + raise RuntimeError("revocation update did not complete") + await commit_database_delete() +except BaseException: + users_to_restore = ( + [authoritative_restored_user] + if authoritative_restored_user.email in barrier.active_user_keys + else [] + ) + if users_to_restore: + failed = await node.sync_users_chunked(users_to_restore, revocation_id=revocation_id) + if failed: + raise RuntimeError("authoritative restore did not complete") + await node.abort_user_revocation(user_keys, revocation_id) + raise +else: + await node.finalize_user_revocation(user_keys, revocation_id) +``` + +`begin_user_revocation` returns `UserRevocationResult(active_user_keys, finalized_user_keys)`. Direct writes must contain +only `active_user_keys`; already-finalized keys are idempotently skipped. It discards older pending/claimed payloads and +waits for older in-flight user-sync leases to drain. A different operation that already owns any requested key causes +`UserRevocationConflictError`; retry the whole operation later. This fail-fast serialization prevents one operation's +rollback restore from racing another operation's permanent finalize. + +An abort must happen only after the authoritative user has been restored on the node. Abort and finalize close admission +before draining authorized writes. A successful finalize leaves a permanent tombstone, and later aborts or stale queue +generations cannot reopen it. Shared stores must keep these fences, generations, claims, and execution leases in the same +atomic consistency domain. A timeout, cancellation, partial transport failure, or expired heartbeat leaves the remote +outcome unknown; its lease is deliberately retained and revocation fails closed. Recover by calling +`reconcile_users(authoritative_users)`: the revocation-aware store waits for live writes, replaces expired poison with a +node-wide permit, and clears it only after the authoritative full snapshot is acknowledged. A failed reconciliation +retains a new node-wide poison lease, so a crash cannot silently reopen revocation. + +Node startup and `sync_users` are full replacement snapshots, so their execution leases cover every user on the node, +including users omitted from the request. A snapshot already in flight drains before a new fence opens. A snapshot which +encounters a provisional fence waits for its authoritative abort/finalize outcome; permanently finalized users are then +omitted from the request. An ambiguous snapshot timeout or cancellation retains the node-wide lease and fails every later +write or permanent revocation on that node closed until reconciliation. `sync_users` rejects `revocation_id`; operation- +owned removal and restore writes must use the partial `sync_users_chunked` transport and treat any returned users as a +failure. + +Fences are scoped to `node_id`. A controller which adds a new node concurrently with deletion must register that node in +the revocation topology before reading its authoritative startup snapshot; a fence on another node cannot protect it. + +Version `0.10.0` adds node-enforced monotonic user-sync epochs plus the required startup-snapshot, epoch-floor handshake, +and atomic lease-narrowing methods to `RevocationAwareUserSyncStoreProtocol`. Roll out the Node binary first, but do not +activate positive epochs yet. Then stop or drain **all** legacy Panel/Bridge workers, upgrade every shared-store adapter +and Bridge process together, and only then resume user mutations and permanent revocation. After a Node accepts its first +positive epoch it rejects legacy epoch-zero clients with HTTP `412` / gRPC `FailedPrecondition`; mixed old/new workers and +rollback to Bridge `0.9` are intentionally unsupported until the Node service is fully restarted and its backend is +recreated from an authoritative snapshot. A Bridge capability probe reads the Node's current epoch and atomically advances +the shared allocator before granting a write lease, so a restarted worker cannot reuse a stale lower epoch. Do not run +user sync or permanent revocation during this cutover. Lifecycle operations are coordinated through the same model. The default process-local coordinator prevents concurrent `start()`, `stop()`, `update_node()`, `update_core()`, and `update_geofiles()` calls from controllers for the same node in one process. Pass a shared `lifecycle_coordinator` in multi-process or multi-host deployments so only one worker can perform a lifecycle operation at a time. Read-only status cron jobs can call stats/info normally; if they write shared observed status, use the current lifecycle epoch so stale cron results cannot overwrite a newer reconnect result. @@ -191,7 +269,12 @@ if state is not None: ) ``` -A lifecycle adapter must atomically acquire/release leases and fence writes with the returned epoch. This prevents a cron status job or another worker from overwriting the result of a newer `start()`, `stop()`, or reconnect flow. +A lifecycle adapter must atomically acquire/release leases, return `False` when heartbeat ownership is lost, and fence +writes with the returned epoch. An expired lease is an unknown remote effect and must not be stolen. Calling +`reconcile_lifecycle(observed_status)` is safe only after an operator has independently established that the old worker and +request can no longer complete; a status probe alone is not such a guarantee. Reconciliation refuses to clear a still-live +lease. Applications must not automatically reconcile a timeout and launch a competing lifecycle operation, because the +Node management API does not yet carry a server-enforced lifecycle fencing token. Node connection configs can also be stored through a registry protocol: @@ -215,9 +298,10 @@ node = await Bridge.create_node_from_registry( ) ``` -### 2. Direct User Sync +### 2. Full User Snapshot Sync -Use direct sync when you want explicit control in your flow. +`sync_users` replaces the node's complete user set. Always pass the authoritative full snapshot; an empty list clears all +users. Use `sync_users_chunked` for partial updates. ```python await node.sync_users([user1, user2], timeout=15) @@ -260,6 +344,10 @@ node_ver = await node.node_version() core_ver = await node.core_version() node_ver2, core_ver2 = await node.get_versions() meta = await node.get_extra() + +# Backwards-compatible synchronous metadata attribute. Prefer get_extra() in +# new asynchronous code. +legacy_meta = node.extra ``` ### 6. On-Demand Log Streaming @@ -342,8 +430,13 @@ await node.override_balancer_target("balancer-tag", "outbound-tag") - `update_user(user)` (queued/background) - `update_users(users)` (queued/background) -- `sync_users(users, flush_pending=False, timeout=None)` (direct) -- `sync_users_chunked(users, chunk_size=100, flush_pending=False, timeout=None)` (direct streaming) +- `begin_user_revocation(user_keys, revocation_id)` → `UserRevocationResult` +- `abort_user_revocation(user_keys, revocation_id)` (release this provisional fence after restore) +- `finalize_user_revocation(user_keys, revocation_id)` (commit permanent tombstones) +- `sync_users(users, flush_pending=False, timeout=None, revocation_id=None)` (full replacement; `revocation_id` rejected) +- `reconcile_users(users, flush_pending=False, timeout=None)` (authoritative full replacement that recovers expired/unknown sync leases) +- `start(..., reconcile_user_sync=True)` (authoritative startup snapshot recovery after a crashed/expired writer) +- `sync_users_chunked(users, chunk_size=100, flush_pending=False, timeout=None, revocation_id=None)` (partial streaming) ### Routing @@ -377,6 +470,9 @@ except Bridge.NodeAPIError as e: print(e.code, e.detail) ``` +Local queue-capacity errors from `update_user` and `update_users` are surfaced +separately as `Bridge.UserSyncStoreFullError` so callers can apply backpressure. + ## Protobuf Access For direct protobuf usage: diff --git a/pyproject.toml b/pyproject.toml index 1916cb1..9042737 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "pasarguard-node-bridge" -version = "0.9.0" +version = "0.10.0" description = "python package to connect your project with PasarGuard node go" url = "https://github.com/PasarGuard/node_bridge_py" keywords = [ diff --git a/tests/test_constructor_compatibility.py b/tests/test_constructor_compatibility.py new file mode 100644 index 0000000..10c44bf --- /dev/null +++ b/tests/test_constructor_compatibility.py @@ -0,0 +1,121 @@ +import unittest +from unittest.mock import MagicMock, patch +from uuid import uuid4 + +from PasarGuardNodeBridge import NodeType, create_node, create_node_from_config +from PasarGuardNodeBridge.abstract_node import PasarGuardNode +from PasarGuardNodeBridge.controller import NodeAPIError +from PasarGuardNodeBridge.rest import Node as RestNode +from PasarGuardNodeBridge.storage import NodeConfig + + +class ConstructorCompatibilityTestCase(unittest.IsolatedAsyncioTestCase): + def setUp(self) -> None: + self.api_key = str(uuid4()) + self.ssl_context = MagicMock() + self.ssl_patch = patch( + "PasarGuardNodeBridge.controller.ssl.create_default_context", + return_value=self.ssl_context, + ) + self.ssl_patch.start() + self.addCleanup(self.ssl_patch.stop) + + def test_rest_node_config_does_not_forward_grpc_only_message_size(self) -> None: + config = NodeConfig( + connection="rest", + address="127.0.0.1", + port=2096, + api_port=2097, + server_ca="test-ca.pem", + api_key=self.api_key, + max_message_size=1234, + ) + + node = create_node_from_config(config) + + self.assertIsInstance(node, RestNode) + self.assertEqual(str(node._client._base_url), "https://127.0.0.1:2096/") + self.assertEqual(str(node._json_client._base_url), "https://127.0.0.1:2097/") + + def test_factory_defaults_api_port_to_service_port(self) -> None: + node = create_node( + connection=NodeType.rest, + address="127.0.0.1", + port=2096, + server_ca="test-ca.pem", + api_key=self.api_key, + ) + + self.assertEqual(str(node._json_client._base_url), "https://127.0.0.1:2096/") + + def test_factory_preserves_explicit_api_port(self) -> None: + node = create_node( + connection=NodeType.rest, + address="127.0.0.1", + port=2096, + api_port=2097, + server_ca="test-ca.pem", + api_key=self.api_key, + ) + + self.assertEqual(str(node._json_client._base_url), "https://127.0.0.1:2097/") + + async def test_extra_attribute_remains_available(self) -> None: + node = create_node( + connection=NodeType.rest, + address="127.0.0.1", + port=2096, + server_ca="test-ca.pem", + api_key=self.api_key, + extra={"region": "eu"}, + ) + + self.assertEqual(node.extra, {"region": "eu"}) + node.extra = {"region": "us"} + self.assertEqual(await node.get_extra(), {"region": "us"}) + + async def test_legacy_subclass_without_reconcile_method_remains_instantiable(self) -> None: + async def legacy_method(*_args, **_kwargs): + return None + + legacy_methods = { + name: legacy_method + for name in PasarGuardNode.__abstractmethods__ + if name != "reconcile_users" + } + legacy_type = type("LegacyNode", (PasarGuardNode,), legacy_methods) + node = legacy_type.__new__(legacy_type) + + with self.assertRaises(NodeAPIError) as error: + await node.reconcile_users([]) + + self.assertEqual(error.exception.code, 501) + + +class FactoryRoutingTestCase(unittest.TestCase): + @patch("PasarGuardNodeBridge.GrpcNode") + def test_grpc_factory_receives_message_size(self, grpc_node: MagicMock) -> None: + api_key = str(uuid4()) + + create_node( + connection=NodeType.grpc, + address="127.0.0.1", + port=2096, + api_port=2097, + server_ca="test-ca.pem", + api_key=api_key, + max_message_size=4096, + ) + + grpc_node.assert_called_once_with( + address="127.0.0.1", + port=2096, + api_port=2097, + server_ca="test-ca.pem", + api_key=api_key, + max_message_size=4096, + ) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_epoch_fencing.py b/tests/test_epoch_fencing.py new file mode 100644 index 0000000..ce2f1bc --- /dev/null +++ b/tests/test_epoch_fencing.py @@ -0,0 +1,363 @@ +import asyncio +import logging +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock + +from grpclib.const import Status +from grpclib.exceptions import GRPCError + +from PasarGuardNodeBridge.common import service_pb2 as service +from PasarGuardNodeBridge.controller import Controller, Health, NodeAPIError +from PasarGuardNodeBridge.grpclib import Node as GrpcNode +from PasarGuardNodeBridge.rest import Node as RestNode +from PasarGuardNodeBridge.storage import InMemoryUserSyncStore + + +def _configured_node(node_type, store: InMemoryUserSyncStore): + node = node_type.__new__(node_type) + node.name = "node-1" + node.node_id = "node-1" + node.worker_id = "worker-1" + node.logger = logging.getLogger("test-epoch-fencing") + node._user_sync_store = store + node._sync_lease_seconds = 30 + node._default_timeout = 1 + node._internal_timeout = 1 + node._node_lock = asyncio.Lock() + node._user_sync_epoch_supported = False + node._user_sync_epoch_capability_probed = False + node._user_sync_epoch_handshake_lock = asyncio.Lock() + node._user_sync_connection_generation = 0 + node.get_health = AsyncMock(return_value=Health.HEALTHY) + node._acquire_lifecycle_lease = AsyncMock(return_value=None) + node._release_lifecycle_lease = AsyncMock() + node.connect = AsyncMock() + node.disconnect = AsyncMock() + return node + + +class EpochTransportTests(unittest.IsolatedAsyncioTestCase): + async def test_start_stale_epoch_retries_once_without_poison(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + node = _configured_node(node_type, store) + node._user_sync_epoch_supported = True + node._user_sync_epoch_capability_probed = True + observed = [] + + async def transport(request): + observed.append(request.user_sync_epoch) + if len(observed) == 1: + if node_type is GrpcNode: + raise GRPCError(Status.FAILED_PRECONDITION, "stale user sync epoch") + raise NodeAPIError(412, "stale user sync epoch") + return service.BaseInfoResponse( + started=True, + node_version="0.4.0", + core_version="1.0.0", + user_sync_epoch_supported=True, + user_sync_epoch=request.user_sync_epoch, + ) + + if node_type is RestNode: + + async def make_request(**kwargs): + return await transport(kwargs["proto_message"]) + + node._make_request = make_request + else: + node._client = SimpleNamespace(Start=object()) + + async def grpc_request(**kwargs): + return await transport(kwargs["request"]) + + node._handle_grpc_request = grpc_request + + result = await node.start("{}", service.BackendType.XRAY, [service.User(email="42")]) + + self.assertTrue(result.started) + self.assertEqual(observed, [1, 2]) + self.assertEqual(store._user_sync_leases, {}) + + async def test_snapshot_stale_epoch_retry_is_bounded_and_does_not_poison(self): + for node_type in (RestNode, GrpcNode): + for method_name in ("sync_users", "reconcile_users"): + with self.subTest(node_type=node_type.__module__, method=method_name): + store = InMemoryUserSyncStore() + node = _configured_node(node_type, store) + node._user_sync_epoch_supported = True + node._user_sync_epoch_capability_probed = True + observed = [] + + async def always_stale(request): + observed.append(request.user_sync_epoch) + if node_type is GrpcNode: + raise GRPCError(Status.FAILED_PRECONDITION, "stale user sync epoch") + raise NodeAPIError(412, "stale user sync epoch") + + if node_type is RestNode: + + async def make_request(**kwargs): + return await always_stale(kwargs["proto_message"]) + + node._make_request = make_request + else: + node._client = SimpleNamespace(SyncUsers=object()) + + async def grpc_request(**kwargs): + return await always_stale(kwargs["request"]) + + node._handle_grpc_request = grpc_request + + with self.assertRaises((NodeAPIError, GRPCError)): + await getattr(node, method_name)([service.User(email="42")]) + + self.assertEqual(observed, [1, 2]) + self.assertEqual(store._user_sync_leases, {}) + + async def test_reversed_direct_delivery_retries_once_with_newer_epoch(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + first = _configured_node(node_type, store) + second = _configured_node(node_type, store) + first.worker_id = "worker-1" + second.worker_id = "worker-2" + for node in (first, second): + node._user_sync_epoch_supported = True + node._user_sync_epoch_capability_probed = True + + first_entered = asyncio.Event() + release_first = asyncio.Event() + applied_epoch = 0 + observed = {"worker-1": [], "worker-2": []} + + def transport_for(node): + async def transport(_users, _chunk_size, _timeout, user_sync_epoch): + nonlocal applied_epoch + observed[node.worker_id].append(user_sync_epoch) + if node.worker_id == "worker-1" and len(observed[node.worker_id]) == 1: + first_entered.set() + await release_first.wait() + if user_sync_epoch < applied_epoch: + if node_type is GrpcNode: + raise GRPCError(Status.FAILED_PRECONDITION, "stale user sync epoch") + raise NodeAPIError(412, "stale user sync epoch") + applied_epoch = user_sync_epoch + + return transport + + first._sync_users_chunked_transport = transport_for(first) + second._sync_users_chunked_transport = transport_for(second) + users = [service.User(email="42")] + + first_sync = asyncio.create_task(first.sync_users_chunked(users)) + await asyncio.wait_for(first_entered.wait(), timeout=1) + self.assertEqual(await second.sync_users_chunked(users), []) + release_first.set() + self.assertEqual(await asyncio.wait_for(first_sync, timeout=1), []) + + self.assertEqual(observed["worker-1"], [1, 3]) + self.assertEqual(observed["worker-2"], [2]) + self.assertEqual(applied_epoch, 3) + self.assertEqual(store._user_sync_leases, {}) + + async def test_stale_direct_delivery_retry_is_bounded_and_does_not_poison(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + node = _configured_node(node_type, store) + node._user_sync_epoch_supported = True + node._user_sync_epoch_capability_probed = True + observed = [] + + async def always_stale(_users, _chunk_size, _timeout, user_sync_epoch): + observed.append(user_sync_epoch) + if node_type is GrpcNode: + raise GRPCError(Status.FAILED_PRECONDITION, "stale user sync epoch") + raise NodeAPIError(412, "stale user sync epoch") + + node._sync_users_chunked_transport = always_stale + users = [service.User(email="42")] + + self.assertEqual(await node.sync_users_chunked(users), users) + self.assertEqual(observed, [1, 2]) + self.assertEqual(store._user_sync_leases, {}) + + async def test_rest_start_cancellation_during_connect_is_not_converted(self): + store = InMemoryUserSyncStore() + node = _configured_node(RestNode, store) + node._user_sync_epoch_supported = True + node._user_sync_epoch_capability_probed = True + connect_entered = asyncio.Event() + + async def transport(**_kwargs): + return service.BaseInfoResponse( + started=True, + node_version="0.4.0", + core_version="1.0.0", + user_sync_epoch_supported=True, + user_sync_epoch=1, + ) + + async def blocking_connect(_node_version, _core_version): + connect_entered.set() + await asyncio.Event().wait() + + node._make_request = transport + node.connect = blocking_connect + start = asyncio.create_task(node.start("{}", service.BackendType.XRAY, [service.User(email="42")])) + await asyncio.wait_for(connect_entered.wait(), timeout=1) + start.cancel() + + with self.assertRaises(asyncio.CancelledError): + await start + node.disconnect.assert_awaited_once() + node._release_lifecycle_lease.assert_awaited_once() + self.assertEqual(store._user_sync_leases, {}) + + async def test_fresh_reconcile_start_probes_capability_and_advances_floor(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + node = _configured_node(node_type, store) + captured: list[int] = [] + + async def response_for(request): + if request is None: + return service.BaseInfoResponse( + started=False, + user_sync_epoch_supported=True, + user_sync_epoch=40, + ) + captured.append(request.user_sync_epoch) + return service.BaseInfoResponse( + started=True, + node_version="0.4.0", + core_version="1.0.0", + user_sync_epoch_supported=True, + user_sync_epoch=request.user_sync_epoch, + ) + + if node_type is RestNode: + + async def make_request(**kwargs): + return await response_for(kwargs.get("proto_message")) + + node._make_request = make_request + else: + info_method = object() + start_method = object() + node._client = SimpleNamespace(GetBaseInfo=info_method, Start=start_method) + + async def grpc_request(**kwargs): + request = None if kwargs["method"] is info_method else kwargs["request"] + return await response_for(request) + + node._handle_grpc_request = grpc_request + + result = await node.start( + config="{}", + backend_type=service.BackendType.XRAY, + users=[service.User(email="42")], + reconcile_user_sync=True, + ) + + self.assertTrue(result.started) + self.assertEqual(captured, [41]) + + async def test_disconnect_invalidates_delayed_capability_response(self): + class BlockingAdvanceStore(InMemoryUserSyncStore): + def __init__(self): + super().__init__() + self.entered = asyncio.Event() + self.release = asyncio.Event() + + async def advance_user_sync_epoch(self, node_id: str, minimum_epoch: int) -> None: + self.entered.set() + await self.release.wait() + await super().advance_user_sync_epoch(node_id, minimum_epoch) + + store = BlockingAdvanceStore() + controller = object.__new__(Controller) + controller.node_id = "node-1" + controller._user_sync_store = store + controller._user_sync_epoch_supported = False + controller._user_sync_epoch_capability_probed = False + controller._user_sync_epoch_handshake_lock = asyncio.Lock() + controller._user_sync_connection_generation = 0 + controller._shutdown_event = asyncio.Event() + controller._task_lock = asyncio.Lock() + controller._cleanup_tasks = AsyncMock() + controller._cleanup_sync_worker = AsyncMock() + controller._health_lock = asyncio.Lock() + controller._version_lock = asyncio.Lock() + controller._node_version = "0.4.0" + controller._core_version = "1.0.0" + controller._health = Health.HEALTHY + + observation = asyncio.create_task( + controller._observe_user_sync_epoch_capability( + service.BaseInfoResponse( + user_sync_epoch_supported=True, + user_sync_epoch=10, + ), + expected_generation=0, + ) + ) + await store.entered.wait() + disconnect = asyncio.create_task(controller.disconnect()) + await asyncio.sleep(0) + store.release.set() + await observation + await disconnect + + self.assertFalse(controller._user_sync_epoch_supported) + self.assertFalse(controller._user_sync_epoch_capability_probed) + self.assertEqual(controller._user_sync_connection_generation, 1) + + async def test_probe_cannot_reapply_response_from_disconnected_generation(self): + controller = object.__new__(Controller) + controller.node_id = "node-1" + controller._user_sync_store = InMemoryUserSyncStore() + controller._user_sync_epoch_supported = False + controller._user_sync_epoch_capability_probed = False + controller._user_sync_epoch_handshake_lock = asyncio.Lock() + controller._user_sync_connection_generation = 0 + controller._shutdown_event = asyncio.Event() + controller._task_lock = asyncio.Lock() + controller._cleanup_tasks = AsyncMock() + controller._cleanup_sync_worker = AsyncMock() + controller._health_lock = asyncio.Lock() + controller._version_lock = asyncio.Lock() + controller._node_version = "0.4.0" + controller._core_version = "1.0.0" + controller._health = Health.HEALTHY + entered = asyncio.Event() + release = asyncio.Event() + + async def delayed_info(): + entered.set() + await release.wait() + return service.BaseInfoResponse( + user_sync_epoch_supported=True, + user_sync_epoch=10, + ) + + controller.info = delayed_info + probe = asyncio.create_task(controller._ensure_user_sync_epoch_support()) + await entered.wait() + await controller.disconnect() + release.set() + with self.assertRaises(NodeAPIError) as error: + await probe + + self.assertEqual(error.exception.code, 426) + self.assertFalse(controller._user_sync_epoch_supported) + self.assertFalse(controller._user_sync_epoch_capability_probed) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_security_hardening.py b/tests/test_security_hardening.py new file mode 100644 index 0000000..8cb1e7c --- /dev/null +++ b/tests/test_security_hardening.py @@ -0,0 +1,962 @@ +import asyncio +import io +import logging +import unittest +from types import MethodType, SimpleNamespace +from unittest.mock import AsyncMock, MagicMock, patch + +from PasarGuardNodeBridge.aiohttp_compat import LazyClientSession +from PasarGuardNodeBridge.common.service_pb2 import User +from PasarGuardNodeBridge.controller import Controller, Health, NodeAPIError, _SanitizingLoggerAdapter +from PasarGuardNodeBridge.grpclib import Node as GrpcNode +from PasarGuardNodeBridge.rest import Node as RestNode +from PasarGuardNodeBridge.storage import ClaimedUser, InMemoryUserSyncStore + + +class _ResponseContext: + def __init__(self, response=None): + self.response = response if response is not None else object() + + async def __aenter__(self): + return self.response + + async def __aexit__(self, exc_type, exc, traceback): + return False + + +class _HangingStreamContext: + async def __aenter__(self): + await asyncio.Event().wait() + + async def __aexit__(self, exc_type, exc, traceback): + return False + + +class _StreamMethod: + def open(self, **kwargs): + return _HangingStreamContext() + + +class _LifecycleStream: + def __init__(self, hanging_phase): + self.hanging_phase = hanging_phase + + async def send_message(self, message): + if self.hanging_phase == "send": + await asyncio.Event().wait() + + async def end(self): + if self.hanging_phase == "end": + await asyncio.Event().wait() + + +class _LifecycleStreamContext: + def __init__(self, hanging_phase): + self.hanging_phase = hanging_phase + self.stream = _LifecycleStream(hanging_phase) + + async def __aenter__(self): + return self.stream + + async def __aexit__(self, exc_type, exc, traceback): + if self.hanging_phase == "exit": + await asyncio.Event().wait() + return False + + +class _LifecycleStreamMethod: + def __init__(self, hanging_phase): + self.hanging_phase = hanging_phase + + def open(self, **kwargs): + return _LifecycleStreamContext(self.hanging_phase) + + +class _BrokenSendStream: + def __init__(self): + self.send_calls = 0 + + async def send_message(self, message): + self.send_calls += 1 + raise RuntimeError("stream closed") + + async def end(self): + raise AssertionError("end must not be called after a send failure") + + +class _BrokenSendStreamContext: + def __init__(self): + self.stream = _BrokenSendStream() + + async def __aenter__(self): + return self.stream + + async def __aexit__(self, exc_type, exc, traceback): + return False + + +class _BrokenSendStreamMethod: + def __init__(self): + self.context = _BrokenSendStreamContext() + + def open(self, **kwargs): + return self.context + + +class _HttpResponse: + def __init__(self, status): + self.status = status + self.headers = {} + self.url = "https://node.example/redirect" + self.reason = "Redirect" + self.charset = "utf-8" + self.closed = False + + async def read(self): + return b"" + + def close(self): + self.closed = True + + +class _StalledHttpResponse(_HttpResponse): + async def read(self): + await asyncio.Event().wait() + + +class _StaticRequestClient: + def __init__(self, response): + self.response = response + + def request(self, *args, **kwargs): + return _ResponseContext(self.response) + + async def get(self, *args, **kwargs): + return self.response + + +class RedirectPolicyTests(unittest.IsolatedAsyncioTestCase): + async def test_request_forces_redirects_off_even_if_caller_enables_them(self): + session = MagicMock() + session.request.return_value = _ResponseContext() + client = LazyClientSession.__new__(LazyClientSession) + client._get_session = AsyncMock(return_value=session) + + async with client.request("GET", "/info", allow_redirects=True): + pass + + self.assertFalse(session.request.call_args.kwargs["allow_redirects"]) + + async def test_get_forces_redirects_off(self): + session = MagicMock() + session.get = AsyncMock(return_value=object()) + client = LazyClientSession.__new__(LazyClientSession) + client._get_session = AsyncMock(return_value=session) + + await client.get("/logs") + + self.assertFalse(session.get.call_args.kwargs["allow_redirects"]) + + async def test_rest_redirect_responses_are_explicit_node_api_errors(self): + for status in (302, 307, 308): + with self.subTest(status=status): + node = RestNode.__new__(RestNode) + node._client = _StaticRequestClient(_HttpResponse(status)) + + with self.assertRaises(NodeAPIError) as error: + await node._make_request(method="GET", endpoint="info", timeout=1) + + self.assertEqual(error.exception.code, status) + + async def test_log_stream_redirect_is_an_explicit_node_api_error(self): + node = RestNode.__new__(RestNode) + node._client = _StaticRequestClient(_HttpResponse(302)) + node._internal_timeout = 0.1 + node.name = "node" + node.logger = _SanitizingLoggerAdapter(logging.getLogger("test.log-redirect"), {}) + + with self.assertRaises(NodeAPIError) as error: + async with node.stream_logs(): + self.fail("redirected log stream must not open") + + self.assertEqual(error.exception.code, 302) + + async def test_log_stream_redirect_body_read_is_bounded_and_response_is_closed(self): + response = _StalledHttpResponse(302) + node = RestNode.__new__(RestNode) + node._client = _StaticRequestClient(response) + node._internal_timeout = 0.01 + node.name = "node" + node.logger = _SanitizingLoggerAdapter(logging.getLogger("test.stalled-log-redirect"), {}) + + with self.assertRaises(NodeAPIError): + async with asyncio.timeout(0.2): + async with node.stream_logs(): + self.fail("redirected log stream must not open") + + self.assertTrue(response.closed) + + +class GrpcStreamTimeoutTests(unittest.IsolatedAsyncioTestCase): + async def test_user_sync_stream_open_timeout_returns_users_for_retry_accounting(self): + node = GrpcNode.__new__(GrpcNode) + node._client = SimpleNamespace(SyncUser=_StreamMethod()) + node._metadata = {} + node._internal_timeout = 0.01 + node.name = "node" + node.logger = _SanitizingLoggerAdapter(logging.getLogger("test.grpc-timeout"), {}) + users = [User(email="private@example.com")] + + failed = await asyncio.wait_for(node._sync_batch_users(users), timeout=0.2) + + self.assertEqual(failed, users) + + async def test_user_sync_send_end_and_exit_are_bounded(self): + users = [User(email="private@example.com")] + for hanging_phase in ("send", "end", "exit"): + with self.subTest(hanging_phase=hanging_phase): + node = GrpcNode.__new__(GrpcNode) + node._client = SimpleNamespace(SyncUser=_LifecycleStreamMethod(hanging_phase)) + node._metadata = {} + node._internal_timeout = 0.01 + node.name = "node" + node.logger = _SanitizingLoggerAdapter(logging.getLogger("test.grpc-lifecycle"), {}) + + failed = await asyncio.wait_for(node._sync_batch_users(users), timeout=0.2) + + self.assertEqual(failed, users) + + async def test_user_sync_stops_after_first_broken_stream_send(self): + method = _BrokenSendStreamMethod() + node = GrpcNode.__new__(GrpcNode) + node._client = SimpleNamespace(SyncUser=method) + node._metadata = {} + node._internal_timeout = 0.1 + node.name = "node" + node.logger = _SanitizingLoggerAdapter(logging.getLogger("test.grpc-broken-stream"), {}) + users = [User(email=f"user-{index}@example.com") for index in range(3)] + + failed = await node._sync_batch_users(users) + + self.assertEqual(failed, users) + self.assertEqual(method.context.stream.send_calls, 1) + + async def test_stream_open_timeout_increments_worker_failure_and_requeues(self): + node = GrpcNode.__new__(GrpcNode) + node._client = SimpleNamespace(SyncUser=_StreamMethod()) + node._metadata = {} + node._internal_timeout = 0.01 + node.name = "node" + node.logger = _SanitizingLoggerAdapter(logging.getLogger("test.grpc-accounting"), {}) + node._shutdown_event = asyncio.Event() + node._work_available = asyncio.Event() + node._work_available.set() + node._worker_idle_timeout = 0.1 + node._sync_poll_interval = 0 + node._sync_lease_seconds = 30 + node._health = Health.HEALTHY + node._health_lock = asyncio.Lock() + node._failure_count_lock = asyncio.Lock() + node._user_sync_failure_count = 0 + node._hard_reset_threshold = 5 + node._hard_reset_event = asyncio.Event() + node._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + user = User(email="private@example.com") + claimed = [ClaimedUser(token="claim", user=user)] + node._claim_pending_users = AsyncMock(return_value=claimed) + node._requeue_claimed_users = AsyncMock() + node._ack_claimed_users = AsyncMock() + node._sync_batch_users = MethodType(GrpcNode._sync_batch_users, node) + + async def stop_after_backoff(_delay): + node._shutdown_event.set() + + with patch("PasarGuardNodeBridge.controller.asyncio.sleep", side_effect=stop_after_backoff): + await node._sync_worker() + + self.assertEqual(node._user_sync_failure_count, 1) + node._requeue_claimed_users.assert_awaited_once_with(claimed) + + +class LoggingSafetyTests(unittest.IsolatedAsyncioTestCase): + async def test_default_logger_does_not_install_output_handler(self): + package_logger = logging.getLogger("PasarGuardNodeBridge") + handlers_before = list(package_logger.handlers) + ssl_context = MagicMock() + + with ( + patch("PasarGuardNodeBridge.controller.ssl.create_default_context", return_value=ssl_context), + patch("PasarGuardNodeBridge.controller.LazyClientSession"), + ): + Controller( + server_ca="certificate", + api_key="00000000-0000-0000-0000-000000000000", + service_url="https://node.example/", + ) + + self.assertEqual(package_logger.handlers, handlers_before) + + async def test_user_sync_log_omits_identifier_and_escapes_control_characters(self): + records = [] + + class _Handler(logging.Handler): + def emit(self, record): + records.append(record.getMessage()) + + logger = logging.getLogger("test.log-safety") + logger.handlers = [_Handler()] + logger.propagate = False + logger.setLevel(logging.WARNING) + + node = RestNode.__new__(RestNode) + node.name = "node\nforged" + node._internal_timeout = 1 + node._make_request = AsyncMock(side_effect=[object(), RuntimeError("remote\r\ninjected")]) + node.logger = _SanitizingLoggerAdapter(logger, {}) + + failed = await node._sync_batch_users([User(email="first@example.com"), User(email="private@example.com")]) + + self.assertEqual(len(failed), 1) + self.assertEqual(len(records), 1) + self.assertNotIn("private@example.com", records[0]) + self.assertNotIn("\n", records[0]) + self.assertNotIn("\r", records[0]) + self.assertIn("\\x0a", records[0]) + self.assertIn("\\x0d", records[0]) + self.assertIn("batch index 1", records[0]) + + async def test_exception_traceback_is_sanitized_before_final_formatting(self): + output = io.StringIO() + handler = logging.StreamHandler(output) + handler.setFormatter(logging.Formatter("%(levelname)s:%(message)s")) + logger = logging.getLogger("test.traceback-safety") + logger.handlers = [handler] + logger.propagate = False + logger.setLevel(logging.ERROR) + adapter = _SanitizingLoggerAdapter(logger, {}) + + try: + raise RuntimeError("remote\r\nforged") + except RuntimeError: + adapter.error("sync failed", exc_info=True) + + formatted = output.getvalue() + self.assertEqual(len(formatted.splitlines()), 1) + self.assertIn("\\x0d\\x0a", formatted) + + async def test_exc_info_accepts_true_tuple_and_exception_instance(self): + output = io.StringIO() + handler = logging.StreamHandler(output) + logger = logging.getLogger("test.exc-info-forms") + logger.handlers = [handler] + logger.propagate = False + logger.setLevel(logging.ERROR) + adapter = _SanitizingLoggerAdapter(logger, {}) + + try: + raise RuntimeError("operational failure") + except RuntimeError as error: + exc_tuple = (type(error), error, error.__traceback__) + adapter.error("true", exc_info=True) + adapter.error("tuple", exc_info=exc_tuple) + adapter.error("instance", exc_info=error) + + formatted = output.getvalue() + self.assertEqual(formatted.count("RuntimeError: operational failure"), 3) + + async def test_positional_log_arguments_and_unicode_separators_are_sanitized(self): + records = [] + + class _Handler(logging.Handler): + def emit(self, record): + records.append(record.getMessage()) + + logger = logging.getLogger("test.positional-log-safety") + logger.handlers = [_Handler()] + logger.propagate = False + logger.setLevel(logging.WARNING) + adapter = _SanitizingLoggerAdapter(logger, {}) + + adapter.warning("remote=%s", "first\r\nsecond\u2028third\u2029fourth") + + self.assertEqual(records, ["remote=first\\x0d\\x0asecond\\u2028third\\u2029fourth"]) + + async def test_connect_restarts_worker_to_discover_stored_pending_work(self): + controller = Controller.__new__(Controller) + controller._shutdown_event = asyncio.Event() + controller._shutdown_event.set() + controller._hard_reset_event = asyncio.Event() + controller._failure_count_lock = asyncio.Lock() + controller._user_sync_failure_count = 3 + controller._task_lock = asyncio.Lock() + controller._tasks = [] + controller._health_lock = asyncio.Lock() + controller._version_lock = asyncio.Lock() + controller._health = 0 + controller._node_version = "" + controller._core_version = "" + controller._work_available = asyncio.Event() + controller._ensure_sync_worker_running = AsyncMock() + + await controller.connect("0.2.0", "1.0.0") + + self.assertTrue(controller._work_available.is_set()) + controller._ensure_sync_worker_running.assert_awaited_once() + + +class SharedStoreDisconnectTests(unittest.IsolatedAsyncioTestCase): + @staticmethod + def _controller(store, worker_id): + controller = Controller.__new__(Controller) + controller.node_id = "node-1" + controller.worker_id = worker_id + controller._user_sync_store = store + controller._shutdown_event = asyncio.Event() + controller._task_lock = asyncio.Lock() + controller._tasks = [] + controller._sync_worker_lock = asyncio.Lock() + controller._sync_worker_task = None + controller._work_available = asyncio.Event() + controller._health_lock = asyncio.Lock() + controller._version_lock = asyncio.Lock() + controller._health = Health.HEALTHY + controller._node_version = "0.2.0" + controller._core_version = "1.0.0" + controller._sync_lease_seconds = 30 + controller._worker_idle_timeout = 1 + controller._sync_poll_interval = 0 + controller._internal_timeout = 0.1 + controller._failure_count_lock = asyncio.Lock() + controller._user_sync_failure_count = 0 + controller._hard_reset_threshold = 5 + controller._hard_reset_event = asyncio.Event() + controller.name = worker_id + controller.logger = _SanitizingLoggerAdapter(logging.getLogger("test.shared-store"), {}) + return controller + + @staticmethod + async def _wait_until(predicate, timeout=0.2): + async def poll(): + while not predicate(): + await asyncio.sleep(0.001) + + await asyncio.wait_for(poll(), timeout=timeout) + + async def test_disconnect_preserves_pending_work_for_second_controller(self): + store = InMemoryUserSyncStore() + first = self._controller(store, "worker-1") + second = self._controller(store, "worker-2") + await store.enqueue_users("node-1", [User(email="pending@example.com")]) + + await asyncio.wait_for(first.disconnect(), timeout=1.0) + claimed = await second._claim_pending_users() + + self.assertEqual([item.user.email for item in claimed], ["pending@example.com"]) + + async def test_enqueue_during_empty_claim_does_not_lose_wakeup(self): + controller = self._controller(InMemoryUserSyncStore(), "worker-1") + claim_observed_empty = asyncio.Event() + release_claim = asyncio.Event() + + async def claim_users(*_args, **_kwargs): + claim_observed_empty.set() + await release_claim.wait() + return [] + + controller._user_sync_store.claim_users = claim_users + claim_task = asyncio.create_task(controller._claim_pending_users()) + await asyncio.wait_for(claim_observed_empty.wait(), timeout=0.1) + + # Model a local enqueue after the store observed an empty queue but + # before claim_users() returns to the controller. + controller._work_available.set() + release_claim.set() + + self.assertEqual(await asyncio.wait_for(claim_task, timeout=0.1), []) + self.assertTrue(controller._work_available.is_set()) + + async def test_enqueue_racing_failed_claim_is_retried_by_existing_worker(self): + store = InMemoryUserSyncStore() + controller = self._controller(store, "worker-1") + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + first_claim_started = asyncio.Event() + release_first_claim = asyncio.Event() + processed = asyncio.Event() + original_claim = store.claim_users + claim_count = 0 + + async def fail_first_claim(*args, **kwargs): + nonlocal claim_count + claim_count += 1 + if claim_count == 1: + first_claim_started.set() + await release_first_claim.wait() + raise RuntimeError("transient store failure") + return await original_claim(*args, **kwargs) + + async def successful_sync(users): + processed.set() + return [] + + store.claim_users = fail_first_claim + controller._sync_batch_users = successful_sync + controller._work_available.set() + + with patch("PasarGuardNodeBridge.controller.INITIAL_CLAIM_RETRY_DELAY", 0.01): + await controller._ensure_sync_worker_running() + worker = controller._sync_worker_task + await asyncio.wait_for(first_claim_started.wait(), timeout=0.2) + + # ensure_sync_worker_running observes the still-running first worker + # while this update sets the wake event and enters the shared store. + await controller.update_user(User(email="pending@example.com")) + self.assertIs(controller._sync_worker_task, worker) + release_first_claim.set() + + await asyncio.wait_for(processed.wait(), timeout=0.2) + + self.assertGreaterEqual(claim_count, 2) + self.assertIs(controller._sync_worker_task, worker) + await asyncio.wait_for(controller.disconnect(), timeout=0.2) + self.assertIsNone(controller._sync_worker_task) + self.assertTrue(worker.done()) + + async def test_enqueue_before_idle_retirement_lock_keeps_current_worker(self): + store = InMemoryUserSyncStore() + controller = self._controller(store, "worker-1") + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + controller._worker_idle_timeout = 0.01 + processed = asyncio.Event() + + async def successful_sync(users): + processed.set() + return [] + + controller._sync_batch_users = successful_sync + await controller._ensure_sync_worker_running() + worker = controller._sync_worker_task + await controller._sync_worker_lock.acquire() + update_task = None + try: + # Hold retirement at its lock boundary until enqueue has published + # both the stored user and the wake event. + await asyncio.sleep(0.02) + update_task = asyncio.create_task(controller.update_user(User(email="pending@example.com"))) + await self._wait_until(controller._work_available.is_set) + finally: + controller._sync_worker_lock.release() + + try: + await asyncio.wait_for(update_task, timeout=0.2) + await asyncio.wait_for(processed.wait(), timeout=0.2) + self.assertIs(controller._sync_worker_task, worker) + finally: + await asyncio.wait_for(controller.disconnect(), timeout=0.2) + + self.assertTrue(worker.done()) + self.assertIsNone(controller._sync_worker_task) + + async def test_enqueue_after_idle_retirement_clear_spawns_replacement(self): + enqueue_started = asyncio.Event() + release_enqueue = asyncio.Event() + + class BlockingEnqueueStore(InMemoryUserSyncStore): + async def enqueue_users(self, node_id, users): + enqueue_started.set() + await release_enqueue.wait() + await super().enqueue_users(node_id, users) + + store = BlockingEnqueueStore() + controller = self._controller(store, "worker-1") + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + controller._worker_idle_timeout = 0.01 + processed = asyncio.Event() + + async def successful_sync(users): + processed.set() + return [] + + controller._sync_batch_users = successful_sync + await controller._ensure_sync_worker_running() + retiring_worker = controller._sync_worker_task + update_task = asyncio.create_task(controller.update_user(User(email="pending@example.com"))) + try: + await asyncio.wait_for(enqueue_started.wait(), timeout=0.2) + # The old worker has already captured the short timeout. Give the + # replacement a generous timeout before releasing the enqueue. + controller._worker_idle_timeout = 1.0 + await self._wait_until(lambda: controller._sync_worker_task is None) + release_enqueue.set() + await asyncio.wait_for(update_task, timeout=0.2) + replacement = controller._sync_worker_task + self.assertIsNotNone(replacement) + self.assertIsNot(replacement, retiring_worker) + await asyncio.wait_for(processed.wait(), timeout=0.2) + await asyncio.wait_for(retiring_worker, timeout=0.2) + self.assertIs(controller._sync_worker_task, replacement) + finally: + release_enqueue.set() + if not update_task.done(): + update_task.cancel() + await asyncio.gather(update_task, return_exceptions=True) + await asyncio.wait_for(controller.disconnect(), timeout=0.2) + + self.assertTrue(retiring_worker.done()) + self.assertTrue(replacement.done()) + self.assertIsNone(controller._sync_worker_task) + + async def test_idle_retirement_boundary_100x_never_strands_enqueued_work(self): + workers = [] + for iteration in range(100): + store = InMemoryUserSyncStore() + controller = self._controller(store, f"worker-{iteration}") + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + controller._worker_idle_timeout = 0.005 + processed = asyncio.Event() + + async def successful_sync(users, processed_event=processed): + processed_event.set() + return [] + + controller._sync_batch_users = successful_sync + await controller._ensure_sync_worker_running() + original_worker = controller._sync_worker_task + workers.append(original_worker) + try: + if iteration % 2 == 0: + # Enqueue wins: the current worker must survive retirement. + await controller._sync_worker_lock.acquire() + try: + await asyncio.sleep(0.01) + update_task = asyncio.create_task( + controller.update_user(User(email=f"pending-{iteration}@example.com")) + ) + await self._wait_until(controller._work_available.is_set) + finally: + controller._sync_worker_lock.release() + await asyncio.wait_for(update_task, timeout=0.5) + self.assertIs(controller._sync_worker_task, original_worker) + else: + # Retirement wins: clearing the reference must let enqueue + # publish a replacement instead of observing a dying task. + await self._wait_until(lambda current=controller: current._sync_worker_task is None) + controller._worker_idle_timeout = 1.0 + await asyncio.wait_for( + controller.update_user(User(email=f"pending-{iteration}@example.com")), timeout=0.5 + ) + self.assertIsNot(controller._sync_worker_task, original_worker) + + try: + await asyncio.wait_for(processed.wait(), timeout=0.5) + except TimeoutError: + self.fail(f"queued work was stranded at retirement iteration {iteration}") + finally: + if controller._sync_worker_lock.locked(): + controller._sync_worker_lock.release() + await asyncio.wait_for(controller.disconnect(), timeout=0.5) + + self.assertIsNone(controller._sync_worker_task) + + self.assertTrue(all(worker.done() for worker in workers)) + + async def test_persistent_claim_failure_backs_off_and_disconnect_cleans_worker(self): + class FailingClaimStore(InMemoryUserSyncStore): + def __init__(self): + super().__init__() + self.claim_times = [] + self.third_claim = asyncio.Event() + + async def claim_users(self, *args, **kwargs): + self.claim_times.append(asyncio.get_running_loop().time()) + if len(self.claim_times) == 3: + self.third_claim.set() + raise RuntimeError("store unavailable") + + store = FailingClaimStore() + controller = self._controller(store, "worker-1") + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + controller._work_available.set() + + with ( + patch("PasarGuardNodeBridge.controller.INITIAL_CLAIM_RETRY_DELAY", 0.01), + patch("PasarGuardNodeBridge.controller.MAX_CLAIM_RETRY_DELAY", 0.02), + ): + await controller._ensure_sync_worker_running() + worker = controller._sync_worker_task + await asyncio.wait_for(store.third_claim.wait(), timeout=0.2) + + self.assertEqual(len(store.claim_times), 3) + self.assertTrue(all(b > a for a, b in zip(store.claim_times, store.claim_times[1:]))) + self.assertGreaterEqual(store.claim_times[-1] - store.claim_times[0], 0.015) + self.assertIs(controller._sync_worker_task, worker) + + await asyncio.wait_for(controller.disconnect(), timeout=0.2) + + self.assertIsNone(controller._sync_worker_task) + self.assertTrue(worker.done()) + claim_count = len(store.claim_times) + await asyncio.sleep(0.025) + self.assertEqual(len(store.claim_times), claim_count) + + async def test_invalid_health_stops_claim_retry_without_restart(self): + store = InMemoryUserSyncStore() + controller = self._controller(store, "worker-1") + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + claim_failed = asyncio.Event() + claim_count = 0 + + async def fail_claim(*_args, **_kwargs): + nonlocal claim_count + claim_count += 1 + claim_failed.set() + raise RuntimeError("store unavailable") + + store.claim_users = fail_claim + controller._work_available.set() + + with patch("PasarGuardNodeBridge.controller.INITIAL_CLAIM_RETRY_DELAY", 0.01): + await controller._ensure_sync_worker_running() + worker = controller._sync_worker_task + await asyncio.wait_for(claim_failed.wait(), timeout=0.2) + await controller.set_health(Health.INVALID) + await asyncio.wait_for(worker, timeout=0.2) + await asyncio.sleep(0) + + self.assertEqual(claim_count, 1) + self.assertIsNone(controller._sync_worker_task) + + async def test_zero_claim_deadline_uses_positive_backoff_when_polling_disabled(self): + controller = self._controller(InMemoryUserSyncStore(), "worker-1") + controller._sync_poll_interval = 0 + observed_timeouts = [] + + async def capture_wait(awaitable, *, timeout): + awaitable.close() + observed_timeouts.append(timeout) + raise asyncio.TimeoutError + + with patch("PasarGuardNodeBridge.controller.asyncio.wait_for", side_effect=capture_wait): + await controller._wait_for_claim_recheck(0.0) + + self.assertEqual(len(observed_timeouts), 1) + self.assertGreater(observed_timeouts[0], 0.0) + self.assertTrue(controller._work_available.is_set()) + + async def test_tiny_positive_claim_deadline_uses_minimum_backoff(self): + controller = self._controller(InMemoryUserSyncStore(), "worker-1") + controller._sync_poll_interval = 0 + observed_timeouts = [] + + async def capture_wait(awaitable, *, timeout): + awaitable.close() + observed_timeouts.append(timeout) + raise asyncio.TimeoutError + + with patch("PasarGuardNodeBridge.controller.asyncio.wait_for", side_effect=capture_wait): + await controller._wait_for_claim_recheck(1e-9) + + self.assertEqual(observed_timeouts, [0.01]) + + async def test_lease_aware_empty_store_exits_without_poll_interval_delay(self): + controller = self._controller(InMemoryUserSyncStore(), "worker-1") + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + controller._worker_idle_timeout = 0.02 + controller._sync_poll_interval = 1.0 + controller._work_available.set() + loop = asyncio.get_running_loop() + started = loop.time() + + await asyncio.wait_for(controller._sync_worker(), timeout=0.2) + + self.assertLess(loop.time() - started, 0.15) + + async def test_zero_deadline_worker_cancels_without_hot_loop_or_task_leak(self): + class DueElsewhereStore(InMemoryUserSyncStore): + def __init__(self): + super().__init__() + self.claim_count = 0 + + async def claim_users(self, *args, **kwargs): + self.claim_count += 1 + return [] + + async def next_claim_delay(self, node_id): + return 0.0 + + store = DueElsewhereStore() + controller = self._controller(store, "worker-1") + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + controller._worker_idle_timeout = 1.0 + controller._sync_poll_interval = 0 + controller._work_available.set() + controller._sync_worker_task = asyncio.create_task(controller._sync_worker()) + worker = controller._sync_worker_task + + await asyncio.sleep(0.05) + + self.assertLessEqual(store.claim_count, 7) + await asyncio.wait_for(controller.disconnect(), timeout=0.5) + self.assertIsNone(controller._sync_worker_task) + self.assertTrue(worker.done()) + + async def test_cancel_after_claim_requeues_immediately_for_second_controller(self): + store = InMemoryUserSyncStore() + first = self._controller(store, "worker-1") + second = self._controller(store, "worker-2") + first._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + sync_started = asyncio.Event() + + async def blocking_sync(users): + sync_started.set() + await asyncio.Event().wait() + + first._sync_batch_users = blocking_sync + await store.enqueue_users("node-1", [User(email="pending@example.com")]) + first._work_available.set() + first._sync_worker_task = asyncio.create_task(first._sync_worker()) + await asyncio.wait_for(sync_started.wait(), timeout=0.2) + + await asyncio.wait_for(first.disconnect(), timeout=1.0) + claimed = await second._claim_pending_users() + + self.assertEqual([item.user.email for item in claimed], ["pending@example.com"]) + + async def test_flush_pending_users_remains_explicit_destructive_clear(self): + store = InMemoryUserSyncStore() + controller = self._controller(store, "worker-1") + await store.enqueue_users("node-1", [User(email="pending@example.com")]) + + await controller.flush_pending_users() + + self.assertEqual(await controller._claim_pending_users(), []) + + async def test_outer_worker_failure_requeues_claimed_users(self): + store = InMemoryUserSyncStore() + controller = self._controller(store, "worker-1") + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + controller._sync_batch_users = AsyncMock(return_value=[]) + controller._ack_claimed_users = AsyncMock(side_effect=RuntimeError("ack failed")) + await store.enqueue_users("node-1", [User(email="pending@example.com")]) + controller._work_available.set() + + async def stop_after_recovery(_delay): + controller._shutdown_event.set() + + with patch("PasarGuardNodeBridge.controller.asyncio.sleep", side_effect=stop_after_recovery): + await controller._sync_worker() + + recovered = await store.claim_users("node-1", "worker-2", limit=10, lease_seconds=30) + self.assertEqual([item.user.email for item in recovered], ["pending@example.com"]) + + async def test_partial_ack_then_failed_requeue_recovers_only_failed_claim(self): + store = InMemoryUserSyncStore() + controller = self._controller(store, "worker-1") + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + succeeded = User(email="ok@example.com") + failed = User(email="failed@example.com") + controller._sync_batch_users = AsyncMock(return_value=[failed]) + original_requeue = controller._requeue_claimed_users + requeue_calls = [] + + async def fail_once_then_requeue(claimed_users): + requeue_calls.append([item.user.email for item in claimed_users]) + if len(requeue_calls) <= 2: + raise RuntimeError("store unavailable") + await original_requeue(claimed_users) + + controller._requeue_claimed_users = fail_once_then_requeue + await store.enqueue_users("node-1", [succeeded, failed]) + controller._work_available.set() + + await controller._sync_worker() + + recovered = await store.claim_users("node-1", "worker-2", limit=10, lease_seconds=30) + self.assertEqual( + requeue_calls, + [["failed@example.com"], ["failed@example.com"], ["failed@example.com"]], + ) + self.assertEqual([item.user.email for item in recovered], ["failed@example.com"]) + + async def test_second_worker_wakes_after_failed_requeue_lease_expires(self): + store = InMemoryUserSyncStore() + first = self._controller(store, "worker-1") + second = self._controller(store, "worker-2") + first._sync_lease_seconds = 0.08 + second._worker_idle_timeout = 0.01 + second._sync_poll_interval = 0 + first._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + second._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + sync_started = asyncio.Event() + second_processed = asyncio.Event() + + async def blocking_sync(users): + sync_started.set() + await asyncio.Event().wait() + + async def successful_sync(users): + second_processed.set() + return [] + + first._sync_batch_users = blocking_sync + second._sync_batch_users = successful_sync + first._requeue_claimed_users = AsyncMock(side_effect=RuntimeError("store unavailable")) + await store.enqueue_users("node-1", [User(email="pending@example.com")]) + first._work_available.set() + first._sync_worker_task = asyncio.create_task(first._sync_worker()) + await asyncio.wait_for(sync_started.wait(), timeout=0.2) + + await asyncio.wait_for(first.disconnect(), timeout=0.2) + + self.assertIsNone(first._sync_worker_task) + second._work_available.set() + second._sync_worker_task = asyncio.create_task(second._sync_worker()) + second_worker = second._sync_worker_task + + await asyncio.sleep(0.02) + self.assertFalse(second_processed.is_set()) + self.assertFalse(second_worker.done()) + + await asyncio.wait_for(second_processed.wait(), timeout=0.5) + for _ in range(20): + if await store.next_claim_delay("node-1") is None: + break + await asyncio.sleep(0.005) + self.assertIsNone(await store.next_claim_delay("node-1")) + + await asyncio.wait_for(second.disconnect(), timeout=0.2) + self.assertIsNone(second._sync_worker_task) + self.assertTrue(second_worker.done()) + + async def test_outer_worker_failure_retries_failed_requeue(self): + store = InMemoryUserSyncStore() + controller = self._controller(store, "worker-1") + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + user = User(email="pending@example.com") + controller._sync_batch_users = AsyncMock(return_value=[user]) + original_requeue = controller._requeue_claimed_users + requeue_calls = 0 + + async def fail_once_then_requeue(claimed_users): + nonlocal requeue_calls + requeue_calls += 1 + if requeue_calls == 1: + raise RuntimeError("store unavailable") + await original_requeue(claimed_users) + + controller._requeue_claimed_users = fail_once_then_requeue + await store.enqueue_users("node-1", [user]) + controller._work_available.set() + + async def stop_after_recovery(_delay): + controller._shutdown_event.set() + + with patch("PasarGuardNodeBridge.controller.asyncio.sleep", side_effect=stop_after_recovery): + await controller._sync_worker() + + recovered = await store.claim_users("node-1", "worker-2", limit=10, lease_seconds=30) + self.assertEqual(requeue_calls, 2) + self.assertEqual([item.user.email for item in recovered], ["pending@example.com"]) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_stop_lifecycle.py b/tests/test_stop_lifecycle.py new file mode 100644 index 0000000..3f07f24 --- /dev/null +++ b/tests/test_stop_lifecycle.py @@ -0,0 +1,241 @@ +import asyncio +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock, Mock + +from PasarGuardNodeBridge.controller import Health, NodeAPIError +from PasarGuardNodeBridge.grpclib import Node as GrpcNode +from PasarGuardNodeBridge.rest import Node as RestNode +from PasarGuardNodeBridge.storage import ( + InMemoryNodeLifecycleCoordinator, + InMemoryUserSyncStore, + LifecycleLeaseLostError, + LifecycleOperation, + LifecycleStatus, +) + + +class StopLifecycleTests(unittest.IsolatedAsyncioTestCase): + async def test_cancellation_during_coordinator_release_waits_for_cleanup(self): + class SlowReleaseCoordinator(InMemoryNodeLifecycleCoordinator): + def __init__(self): + super().__init__() + self.release_entered = asyncio.Event() + self.finish_release = asyncio.Event() + + async def release(self, lease, state_update=None): + self.release_entered.set() + await self.finish_release.wait() + await super().release(lease, state_update) + + coordinator = SlowReleaseCoordinator() + node = RestNode.__new__(RestNode) + self._configure_node(node, coordinator) + lease = await coordinator.try_acquire(node.node_id, node.worker_id, LifecycleOperation.STOP, 30) + self.assertIsNotNone(lease) + release = asyncio.create_task( + node._release_lifecycle_lease( + lease, + observed=LifecycleStatus.STOPPED, + desired=LifecycleStatus.STOPPED, + ) + ) + await asyncio.wait_for(coordinator.release_entered.wait(), timeout=1) + release.cancel() + await asyncio.sleep(0) + self.assertFalse(release.done()) + + coordinator.finish_release.set() + with self.assertRaises(asyncio.CancelledError): + await release + state = await coordinator.get_state(node.node_id) + self.assertEqual(state.observed, LifecycleStatus.STOPPED) + self.assertIsNone(state.operation) + + def _configure_node(self, node, coordinator: InMemoryNodeLifecycleCoordinator) -> None: + node.node_id = "node-1" + node.name = "node-1" + node.worker_id = "worker-1" + node._default_timeout = 10 + node._node_lock = asyncio.Lock() + node._lifecycle_coordinator = coordinator + node._lifecycle_lease_seconds = 30 + node._lifecycle_heartbeat_tasks = {} + node.logger = Mock() + node.get_health = AsyncMock(return_value=Health.HEALTHY) + node.disconnect = AsyncMock() + node._json_client = SimpleNamespace(close=AsyncMock()) + + async def _assert_failed_stop_keeps_lease(self, node) -> None: + with self.assertRaises(NodeAPIError) as error: + await node.stop() + self.assertEqual(error.exception.code, 503) + state = await node._lifecycle_coordinator.get_state(node.node_id) + self.assertEqual(state.observed, LifecycleStatus.STOPPING) + competing = await node._lifecycle_coordinator.try_acquire( + node.node_id, "worker-2", LifecycleOperation.START, 30 + ) + self.assertIsNone(competing) + self.assertEqual(node._lifecycle_heartbeat_tasks, {}) + node.disconnect.assert_not_awaited() + + async def test_rest_stop_propagates_error_without_releasing_lease(self): + coordinator = InMemoryNodeLifecycleCoordinator() + node = RestNode.__new__(RestNode) + self._configure_node(node, coordinator) + node._client = SimpleNamespace(close=AsyncMock()) + node._make_request = AsyncMock(side_effect=NodeAPIError(503, "REST stop failed")) + await self._assert_failed_stop_keeps_lease(node) + node._make_request.assert_awaited_once_with(method="PUT", endpoint="stop", timeout=10) + + async def test_grpc_stop_propagates_error_without_releasing_lease(self): + coordinator = InMemoryNodeLifecycleCoordinator() + node = GrpcNode.__new__(GrpcNode) + self._configure_node(node, coordinator) + node._client = SimpleNamespace(Stop=AsyncMock()) + node._handle_grpc_request = AsyncMock(side_effect=NodeAPIError(503, "gRPC stop failed")) + await self._assert_failed_stop_keeps_lease(node) + node._handle_grpc_request.assert_awaited_once() + + async def test_failed_heartbeat_does_not_mask_stop_error(self): + coordinator = InMemoryNodeLifecycleCoordinator() + node = RestNode.__new__(RestNode) + self._configure_node(node, coordinator) + node._client = SimpleNamespace(close=AsyncMock()) + node._make_request = AsyncMock(side_effect=NodeAPIError(503, "REST stop failed")) + + async def heartbeat_that_fails_during_cleanup(): + try: + await asyncio.Event().wait() + except asyncio.CancelledError as exc: + raise RuntimeError("heartbeat failed") from exc + + lease = await coordinator.try_acquire(node.node_id, node.worker_id, LifecycleOperation.STOP, 30) + self.assertIsNotNone(lease) + node._acquire_lifecycle_lease = AsyncMock(return_value=lease) + node._lifecycle_heartbeat_tasks[lease.token] = asyncio.create_task(heartbeat_that_fails_during_cleanup()) + await asyncio.sleep(0) + with self.assertRaises(NodeAPIError) as error: + await node.stop() + self.assertEqual(error.exception.detail, "REST stop failed") + node.logger.exception.assert_called_once() + + async def test_release_propagates_caller_cancellation_after_recording_state(self): + coordinator = InMemoryNodeLifecycleCoordinator() + node = RestNode.__new__(RestNode) + self._configure_node(node, coordinator) + lease = await coordinator.try_acquire(node.node_id, node.worker_id, LifecycleOperation.STOP, 30) + self.assertIsNotNone(lease) + heartbeat_cancelling = asyncio.Event() + + async def slow_heartbeat_shutdown(): + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + heartbeat_cancelling.set() + await asyncio.Event().wait() + + node._lifecycle_heartbeat_tasks[lease.token] = asyncio.create_task(slow_heartbeat_shutdown()) + release = asyncio.create_task( + node._release_lifecycle_lease( + lease, + observed=LifecycleStatus.STOPPED, + desired=LifecycleStatus.STOPPED, + ) + ) + await asyncio.wait_for(heartbeat_cancelling.wait(), timeout=1) + release.cancel() + + with self.assertRaises(asyncio.CancelledError): + await release + state = await coordinator.get_state(node.node_id) + self.assertEqual(state.observed, LifecycleStatus.STOPPED) + self.assertIsNone(state.operation) + + async def test_ambiguous_start_retains_lifecycle_lease(self): + coordinator = InMemoryNodeLifecycleCoordinator() + node = RestNode.__new__(RestNode) + self._configure_node(node, coordinator) + node._user_sync_store = InMemoryUserSyncStore() + node._sync_lease_seconds = 30 + node._internal_timeout = 1 + node._user_sync_epoch_supported = True + node._user_sync_epoch_capability_probed = True + node._user_sync_epoch_handshake_lock = asyncio.Lock() + node._user_sync_connection_generation = 0 + node._make_request = AsyncMock(side_effect=TimeoutError) + + with self.assertRaises(TimeoutError): + await node.start( + config="{}", + backend_type=0, + users=[], + reconcile_user_sync=True, + ) + + competing = await coordinator.try_acquire("node-1", "worker-2", LifecycleOperation.STOP, 30) + self.assertIsNone(competing) + state = await coordinator.get_state("node-1") + self.assertEqual(state.operation, LifecycleOperation.START) + + async def test_management_timeout_retains_lifecycle_lease(self): + coordinator = InMemoryNodeLifecycleCoordinator() + node = RestNode.__new__(RestNode) + self._configure_node(node, coordinator) + node.check_connectivity = AsyncMock(return_value=True) + node._make_json_request = AsyncMock(side_effect=TimeoutError) + + with self.assertRaises(TimeoutError): + await node._run_coordinated_update( + LifecycleOperation.UPDATE_CORE, + "/node/core_update", + {"version": "latest"}, + ) + + competing = await coordinator.try_acquire("node-1", "worker-2", LifecycleOperation.STOP, 30) + self.assertIsNone(competing) + + async def test_successful_start_with_lost_heartbeat_keeps_poison(self): + class LostHeartbeatCoordinator(InMemoryNodeLifecycleCoordinator): + async def heartbeat(self, _lease): + return False + + coordinator = LostHeartbeatCoordinator() + node = RestNode.__new__(RestNode) + self._configure_node(node, coordinator) + node._lifecycle_lease_seconds = 0.001 + node._user_sync_store = InMemoryUserSyncStore() + node._sync_lease_seconds = 30 + node._internal_timeout = 1 + node._user_sync_epoch_supported = True + node._user_sync_epoch_capability_probed = True + node._user_sync_epoch_handshake_lock = asyncio.Lock() + node._user_sync_connection_generation = 0 + + async def successful_start(**_kwargs): + await asyncio.sleep(0.02) + return SimpleNamespace( + started=True, + node_version="0.4.0", + core_version="1.0.0", + user_sync_epoch_supported=True, + user_sync_epoch=1, + ) + + node._make_request = successful_start + node.connect = AsyncMock() + + with self.assertRaises(LifecycleLeaseLostError): + await node.start( + config="{}", + backend_type=0, + users=[], + reconcile_user_sync=True, + ) + + competing = await coordinator.try_acquire("node-1", "worker-2", LifecycleOperation.STOP, 30) + self.assertIsNone(competing) + + +if __name__ == "__main__": + unittest.main() diff --git a/tests/test_storage.py b/tests/test_storage.py index a8d9a9a..e65d2f2 100644 --- a/tests/test_storage.py +++ b/tests/test_storage.py @@ -1,22 +1,58 @@ import asyncio import unittest from typing import Any, cast -from unittest.mock import AsyncMock +from unittest.mock import AsyncMock, patch from PasarGuardNodeBridge.common.service_pb2 import User from PasarGuardNodeBridge.controller import Controller, NodeAPIError from PasarGuardNodeBridge.storage import ( + ClaimedUser, InMemoryNodeLifecycleCoordinator, InMemoryNodeRegistry, InMemoryUserSyncStore, + LifecycleLease, + LifecycleLeaseLostError, LifecycleOperation, LifecycleStatus, NodeConfig, NodeLifecycleState, + UserSyncStoreFullError, ) class InMemoryUserSyncStoreTests(unittest.IsolatedAsyncioTestCase): + async def test_execution_epochs_are_monotonic_and_survive_narrowing(self): + store = InMemoryUserSyncStore() + first = await store.acquire_user_sync_lease("node-1", "worker-1", ["a@example.com", "b@example.com"], 30) + narrowed = await store.retain_user_sync_lease_keys(first, ["a@example.com"]) + self.assertEqual(narrowed.epoch, first.epoch) + await store.release_user_sync_lease(narrowed) + + second = (await store.acquire_startup_user_sync_lease("node-1", "worker-2", ["a@example.com"], 30)).lease + self.assertGreater(second.epoch, first.epoch) + await store.release_user_sync_lease(second) + + other_node = await store.acquire_user_sync_lease("node-2", "worker-3", ["a@example.com"], 30) + self.assertEqual(other_node.epoch, 1) + await store.release_user_sync_lease(other_node) + + async def test_snapshot_reads_do_not_allocate_revocation_state_for_unseen_users(self): + store = InMemoryUserSyncStore() + + startup = await store.acquire_startup_user_sync_lease( + "node-1", "worker-1", [f"user-{index}@example.com" for index in range(100)], 30 + ) + + self.assertNotIn("node-1", store._revocations) + await store.release_user_sync_lease(startup.lease) + + recovery = await store.acquire_user_sync_reconciliation_lease( + "node-1", "worker-2", [f"other-{index}@example.com" for index in range(100)], 30 + ) + + self.assertNotIn("node-1", store._revocations) + await store.release_user_sync_lease(recovery.lease) + async def test_enqueue_claim_ack_removes_user(self): store = InMemoryUserSyncStore() await store.enqueue_users("node-1", [User(email="a@example.com")]) @@ -60,6 +96,19 @@ async def test_requeue_makes_failed_claim_available_again(self): self.assertEqual([item.user.email for item in claimed_again], ["a@example.com"]) + async def test_requeue_ignores_acknowledged_or_unknown_tokens(self): + store = InMemoryUserSyncStore() + await store.enqueue_users("node-1", [User(email="acked@example.com")]) + claimed = await store.claim_users("node-1", "worker-1", limit=10, lease_seconds=30) + await store.ack_users("node-1", [claimed[0].token]) + + await store.requeue_users( + "node-1", + [claimed[0], ClaimedUser(token="unknown", user=User(email="injected@example.com"))], + ) + + self.assertEqual(await store.claim_users("node-1", "worker-2", limit=10, lease_seconds=30), []) + async def test_expired_lease_becomes_claimable(self): store = InMemoryUserSyncStore() await store.enqueue_users("node-1", [User(email="a@example.com")]) @@ -69,6 +118,42 @@ async def test_expired_lease_becomes_claimable(self): self.assertEqual([item.user.email for item in claimed_again], ["a@example.com"]) + async def test_next_claim_delay_distinguishes_empty_pending_and_leased_work(self): + store = InMemoryUserSyncStore() + + self.assertIsNone(await store.next_claim_delay("node-1")) + + await store.enqueue_users("node-1", [User(email="a@example.com")]) + self.assertEqual(await store.next_claim_delay("node-1"), 0.0) + + with patch("PasarGuardNodeBridge.storage.time.monotonic", side_effect=[100.0, 100.025]): + await store.claim_users("node-1", "worker-1", limit=10, lease_seconds=0.1) + delay = await store.next_claim_delay("node-1") + self.assertIsNotNone(delay) + self.assertAlmostEqual(delay, 0.075) + + async def test_enqueue_rejects_work_above_per_node_bound_without_partial_write(self): + store = InMemoryUserSyncStore(max_pending_users_per_node=1) + await store.enqueue_users("node-1", [User(email="a@example.com")]) + + with self.assertRaises(UserSyncStoreFullError): + await store.enqueue_users("node-1", [User(email="a@example.com"), User(email="b@example.com")]) + + claimed = await store.claim_users("node-1", "worker-1", limit=10, lease_seconds=30) + self.assertEqual([item.user.email for item in claimed], ["a@example.com"]) + + async def test_claimed_users_count_toward_per_node_bound(self): + store = InMemoryUserSyncStore(max_pending_users_per_node=1) + await store.enqueue_users("node-1", [User(email="a@example.com")]) + await store.claim_users("node-1", "worker-1", limit=10, lease_seconds=30) + + with self.assertRaises(UserSyncStoreFullError): + await store.enqueue_users("node-1", [User(email="b@example.com")]) + + def test_non_positive_per_node_bound_is_rejected(self): + with self.assertRaises(ValueError): + InMemoryUserSyncStore(max_pending_users_per_node=0) + class InMemoryNodeRegistryTests(unittest.IsolatedAsyncioTestCase): async def test_registry_roundtrip(self): @@ -122,12 +207,16 @@ async def test_release_records_final_state(self): self.assertEqual(state.owner, None) self.assertEqual(state.node_version, "0.2.0") - async def test_stale_observed_update_is_ignored(self): + async def test_expired_lease_requires_reconciliation_before_new_operation(self): coordinator = InMemoryNodeLifecycleCoordinator() first = await coordinator.try_acquire("node-1", "worker-1", LifecycleOperation.START, 0.001) await asyncio.sleep(0.01) second = await coordinator.try_acquire("node-1", "worker-2", LifecycleOperation.STOP, 30) self.assertIsNotNone(first) + self.assertIsNone(second) + + self.assertTrue(await coordinator.reconcile("node-1", LifecycleStatus.HEALTHY)) + second = await coordinator.try_acquire("node-1", "worker-2", LifecycleOperation.STOP, 30) self.assertIsNotNone(second) await coordinator.update_observed("node-1", LifecycleStatus.BROKEN, expected_epoch=first.epoch) @@ -136,15 +225,32 @@ async def test_stale_observed_update_is_ignored(self): self.assertEqual(state.epoch, second.epoch) self.assertNotEqual(state.observed, LifecycleStatus.BROKEN) - async def test_expired_lifecycle_lease_can_be_reacquired(self): + async def test_active_lifecycle_lease_cannot_be_reconciled(self): coordinator = InMemoryNodeLifecycleCoordinator() - await coordinator.try_acquire("node-1", "worker-1", LifecycleOperation.RECONNECT, 0) - await asyncio.sleep(0.01) - - lease = await coordinator.try_acquire("node-1", "worker-2", LifecycleOperation.RECONNECT, 30) + await coordinator.try_acquire("node-1", "worker-1", LifecycleOperation.RECONNECT, 30) + + self.assertFalse(await coordinator.reconcile("node-1", LifecycleStatus.HEALTHY)) + self.assertIsNone(await coordinator.try_acquire("node-1", "worker-2", LifecycleOperation.RECONNECT, 30)) + + async def test_controller_detects_lifecycle_heartbeat_ownership_loss(self): + controller = Controller.__new__(Controller) + controller.node_id = "node-1" + controller._lifecycle_lease_seconds = 0.001 + controller._lifecycle_coordinator = cast( + Any, + type("LostCoordinator", (), {"heartbeat": AsyncMock(return_value=False)})(), + ) + lease = LifecycleLease( + node_id="node-1", + worker_id="worker-1", + operation=LifecycleOperation.START, + token="lost-token", + epoch=1, + lease_seconds=0.001, + ) - self.assertIsNotNone(lease) - self.assertEqual(lease.worker_id, "worker-2") + with self.assertRaises(LifecycleLeaseLostError): + await controller._heartbeat_lifecycle_lease(lease) async def test_node_update_is_exclusive_across_controllers(self): coordinator = InMemoryNodeLifecycleCoordinator() diff --git a/tests/test_user_revocation.py b/tests/test_user_revocation.py new file mode 100644 index 0000000..c9c7c75 --- /dev/null +++ b/tests/test_user_revocation.py @@ -0,0 +1,971 @@ +import asyncio +import logging +import unittest +from types import SimpleNamespace +from unittest.mock import AsyncMock, patch + +from PasarGuardNodeBridge.common import service_pb2 as service +from PasarGuardNodeBridge.common.service_pb2 import User +from PasarGuardNodeBridge.controller import Controller, Health, NodeAPIError +from PasarGuardNodeBridge.grpclib import Node as GrpcNode +from PasarGuardNodeBridge.rest import Node as RestNode +from PasarGuardNodeBridge.storage import ( + InMemoryUserSyncStore, + UserRevocationConflictError, + UserSyncLeaseLostError, +) + + +def _user(key: str, inbound: str = "active") -> User: + return User(email=key, inbounds=[inbound] if inbound else []) + + +def _controller(store, worker_id: str = "worker-1") -> Controller: + controller = object.__new__(Controller) + controller.name = "test" + controller.node_id = "node-1" + controller.worker_id = worker_id + controller.logger = logging.getLogger("test-user-revocation") + controller._user_sync_store = store + controller._sync_lease_seconds = 1.0 + controller._internal_timeout = 1 + controller._work_available = asyncio.Event() + controller._shutdown_event = asyncio.Event() + controller._worker_idle_timeout = 30.0 + controller._sync_poll_interval = 0.0 + controller._health = Health.HEALTHY + controller._user_sync_epoch_supported = True + controller._health_lock = asyncio.Lock() + controller._node_lock = asyncio.Lock() + controller._sync_worker_lock = asyncio.Lock() + controller._sync_worker_task = None + controller._ensure_sync_worker_running = AsyncMock() + controller._hard_reset_event = asyncio.Event() + controller._user_sync_failure_count = 0 + controller._hard_reset_threshold = 5 + controller._failure_count_lock = asyncio.Lock() + controller._supports_chunked_sync = AsyncMock(return_value=(False, "0.1.0")) + return controller + + +async def _wait_until(predicate, timeout: float = 1.0) -> None: + deadline = asyncio.get_running_loop().time() + timeout + while not predicate(): + if asyncio.get_running_loop().time() >= deadline: + raise TimeoutError("condition was not met") + await asyncio.sleep(0.001) + + +class UserRevocationStoreTests(unittest.IsolatedAsyncioTestCase): + async def test_empty_abort_and_finalize_do_not_wait_for_active_startup(self): + store = InMemoryUserSyncStore() + startup = await store.acquire_startup_user_sync_lease("node-1", "starter", ["42"], 30) + + await asyncio.wait_for(store.abort_user_revocation("node-1", ["42"], "unknown"), timeout=0.1) + await asyncio.wait_for(store.finalize_user_revocation("node-1", ["42"], "unknown"), timeout=0.1) + + await store.release_user_sync_lease(startup.lease) + + async def test_begin_discards_pending_and_blocks_new_updates(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + await store.enqueue_users("node-1", [_user("42")]) + + await controller.begin_user_revocation(["42"], "revoke-a") + await controller.update_user(_user("42", "new")) + + self.assertEqual(await store.claim_users("node-1", "reader", 10, 30), []) + + async def test_stale_claim_is_rejected_after_abort_and_authoritative_restore_wins(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + await store.enqueue_users("node-1", [_user("42", "stale")]) + stale_claim = (await store.claim_users("node-1", "old-worker", 10, 30))[0] + + await controller.begin_user_revocation(["42"], "revoke-a") + await controller.abort_user_revocation(["42"], "revoke-a") + + stale_lease = await store.acquire_user_sync_lease( + "node-1", + "old-worker", + ["42"], + 30, + {"42": stale_claim.generation}, + ) + self.assertEqual(stale_lease.user_keys, ()) + + await store.enqueue_users("node-1", [_user("42", "restored")]) + restored = await store.claim_users("node-1", "new-worker", 10, 30) + self.assertEqual([list(item.user.inbounds) for item in restored], [["restored"]]) + + async def test_overlapping_revocations_fail_without_releasing_owner(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + + await controller.begin_user_revocation(["42"], "revoke-a") + with self.assertRaises(UserRevocationConflictError) as error: + await controller.begin_user_revocation(["43", "42"], "revoke-b") + self.assertEqual(error.exception.conflicting_user_keys, ("42",)) + + await controller.update_users([_user("42"), _user("43")]) + unfenced = await store.claim_users("node-1", "reader", 10, 30) + self.assertEqual([item.user.email for item in unfenced], ["43"]) + + await controller.abort_user_revocation(["42"], "revoke-a") + result = await controller.begin_user_revocation(["42"], "revoke-b") + self.assertEqual(result.active_user_keys, ("42",)) + await controller.abort_user_revocation(["42"], "revoke-b") + await controller.update_user(_user("42", "restored")) + claimed = await store.claim_users("node-1", "reader", 10, 30) + self.assertEqual([list(item.user.inbounds) for item in claimed], [["restored"]]) + + async def test_finalize_is_permanent_and_later_abort_cannot_reopen(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + + await controller.begin_user_revocation(["42"], "revoke-a") + await controller.finalize_user_revocation(["42"], "revoke-a") + result = await controller.begin_user_revocation(["42", "43"], "revoke-b") + self.assertEqual(result.active_user_keys, ("43",)) + self.assertEqual(result.finalized_user_keys, ("42",)) + await controller.update_user(_user("42")) + + self.assertEqual(await store.claim_users("node-1", "reader", 10, 30), []) + + async def test_abort_closes_admission_and_waits_authorized_write(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + await controller.begin_user_revocation(["42"], "revoke-a") + lease = await store.acquire_user_sync_lease("node-1", "writer", ["42"], 30, revocation_id="revoke-a") + + abort = asyncio.create_task(controller.abort_user_revocation(["42"], "revoke-a")) + await asyncio.sleep(0) + self.assertFalse(abort.done()) + denied = await store.acquire_user_sync_lease("node-1", "late-writer", ["42"], 30, revocation_id="revoke-a") + self.assertEqual(denied.user_keys, ()) + + await store.release_user_sync_lease(lease) + await asyncio.wait_for(abort, timeout=1) + allowed = await store.acquire_user_sync_lease("node-1", "ordinary", ["42"], 30) + self.assertEqual(allowed.user_keys, ("42",)) + await store.release_user_sync_lease(allowed) + + async def test_finalize_closes_admission_and_waits_authorized_write(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + await controller.begin_user_revocation(["42"], "revoke-a") + lease = await store.acquire_user_sync_lease("node-1", "writer", ["42"], 30, revocation_id="revoke-a") + + finalize = asyncio.create_task(controller.finalize_user_revocation(["42"], "revoke-a")) + await asyncio.sleep(0) + self.assertFalse(finalize.done()) + denied = await store.acquire_user_sync_lease("node-1", "late-writer", ["42"], 30, revocation_id="revoke-a") + self.assertEqual(denied.user_keys, ()) + + await store.release_user_sync_lease(lease) + await asyncio.wait_for(finalize, timeout=1) + result = await controller.begin_user_revocation(["42"], "revoke-b") + self.assertEqual(result.active_user_keys, ()) + self.assertEqual(result.finalized_user_keys, ("42",)) + + async def test_abort_lost_lease_reopens_only_owner_admission(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + await controller.begin_user_revocation(["42"], "revoke-a") + lost = await store.acquire_user_sync_lease("node-1", "writer", ["42"], 0.01, revocation_id="revoke-a") + await asyncio.sleep(0.02) + + with self.assertRaises(UserSyncLeaseLostError): + await controller.abort_user_revocation(["42"], "revoke-a") + + owner = await store.acquire_user_sync_lease("node-1", "reconcile", ["42"], 30, revocation_id="revoke-a") + ordinary = await store.acquire_user_sync_lease("node-1", "ordinary", ["42"], 30) + self.assertEqual(owner.user_keys, ("42",)) + self.assertEqual(ordinary.user_keys, ()) + await store.release_user_sync_lease(owner) + await store.release_user_sync_lease(lost) + await controller.abort_user_revocation(["42"], "revoke-a") + admitted = await store.acquire_user_sync_lease("node-1", "ordinary-after-abort", ["42"], 30) + self.assertEqual(admitted.user_keys, ("42",)) + await store.release_user_sync_lease(admitted) + + async def test_finalize_lost_lease_reopens_only_owner_admission(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + await controller.begin_user_revocation(["42"], "revoke-a") + lost = await store.acquire_user_sync_lease("node-1", "writer", ["42"], 0.01, revocation_id="revoke-a") + await asyncio.sleep(0.02) + + with self.assertRaises(UserSyncLeaseLostError): + await controller.finalize_user_revocation(["42"], "revoke-a") + + owner = await store.acquire_user_sync_lease("node-1", "reconcile", ["42"], 30, revocation_id="revoke-a") + ordinary = await store.acquire_user_sync_lease("node-1", "ordinary", ["42"], 30) + self.assertEqual(owner.user_keys, ("42",)) + self.assertEqual(ordinary.user_keys, ()) + await store.release_user_sync_lease(owner) + await store.release_user_sync_lease(lost) + await controller.finalize_user_revocation(["42"], "revoke-a") + result = await controller.begin_user_revocation(["42"], "revoke-b") + self.assertEqual(result.finalized_user_keys, ("42",)) + + async def test_abort_cancellation_reopens_only_owner_admission(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + await controller.begin_user_revocation(["42"], "revoke-a") + active = await store.acquire_user_sync_lease("node-1", "writer", ["42"], 30, revocation_id="revoke-a") + abort = asyncio.create_task(controller.abort_user_revocation(["42"], "revoke-a")) + await asyncio.sleep(0) + abort.cancel() + with self.assertRaises(asyncio.CancelledError): + await abort + + owner = await store.acquire_user_sync_lease("node-1", "reconcile", ["42"], 30, revocation_id="revoke-a") + ordinary = await store.acquire_user_sync_lease("node-1", "ordinary", ["42"], 30) + self.assertEqual(owner.user_keys, ("42",)) + self.assertEqual(ordinary.user_keys, ()) + await store.release_user_sync_lease(owner) + await store.release_user_sync_lease(active) + await controller.abort_user_revocation(["42"], "revoke-a") + + async def test_finalize_cancellation_reopens_only_owner_admission(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + await controller.begin_user_revocation(["42"], "revoke-a") + active = await store.acquire_user_sync_lease("node-1", "writer", ["42"], 30, revocation_id="revoke-a") + finalize = asyncio.create_task(controller.finalize_user_revocation(["42"], "revoke-a")) + await asyncio.sleep(0) + finalize.cancel() + with self.assertRaises(asyncio.CancelledError): + await finalize + + owner = await store.acquire_user_sync_lease("node-1", "reconcile", ["42"], 30, revocation_id="revoke-a") + ordinary = await store.acquire_user_sync_lease("node-1", "ordinary", ["42"], 30) + self.assertEqual(owner.user_keys, ("42",)) + self.assertEqual(ordinary.user_keys, ()) + await store.release_user_sync_lease(owner) + await store.release_user_sync_lease(active) + await controller.finalize_user_revocation(["42"], "revoke-a") + + async def test_cancelled_begin_stays_fail_closed_until_explicit_abort(self): + store = InMemoryUserSyncStore() + lease = await store.acquire_user_sync_lease("node-1", "worker", ["42"], 30) + begin = asyncio.create_task(store.begin_user_revocation("node-1", ["42"], "revoke-a")) + await asyncio.sleep(0) + self.assertFalse(begin.done()) + + begin.cancel() + with self.assertRaises(asyncio.CancelledError): + await begin + denied = await store.acquire_user_sync_lease("node-1", "reader", ["42"], 30) + self.assertEqual(denied.user_keys, ()) + + await store.release_user_sync_lease(lease) + await store.abort_user_revocation("node-1", ["42"], "revoke-a") + allowed = await store.acquire_user_sync_lease("node-1", "reader", ["42"], 30) + self.assertEqual(allowed.user_keys, ("42",)) + + async def test_expired_execution_lease_fails_closed_until_explicit_release(self): + store = InMemoryUserSyncStore() + lease = await store.acquire_user_sync_lease("node-1", "worker", ["42"], 0.01) + await asyncio.sleep(0.02) + + with self.assertRaises(UserSyncLeaseLostError): + await store.begin_user_revocation("node-1", ["42"], "revoke-a") + denied = await store.acquire_user_sync_lease("node-1", "reader", ["42"], 30) + self.assertEqual(denied.user_keys, ()) # begin already installed the fail-closed fence + + await store.release_user_sync_lease(lease) + await store.begin_user_revocation("node-1", ["42"], "revoke-a") + + async def test_expired_execution_lease_is_recovered_by_authoritative_snapshot(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + controller = _controller(store) + controller._default_timeout = 1 + lost = await store.acquire_user_sync_lease("node-1", "dead-worker", ["42"], 0.01) + await asyncio.sleep(0.02) + captured: list[str] = [] + + async def transport(_captured=captured, **kwargs): + request = kwargs["proto_message"] if "proto_message" in kwargs else kwargs["request"] + _captured.extend(user.email for user in request.users) + return service.Empty() + + if node_type is RestNode: + controller._make_request = transport + else: + controller._client = SimpleNamespace(SyncUsers=object()) + controller._handle_grpc_request = transport + + await node_type.reconcile_users(controller, [_user("42"), _user("43")]) + + self.assertEqual(captured, ["42", "43"]) + self.assertNotIn(lost.token, store._user_sync_leases) + result = await controller.begin_user_revocation(["42"], "revoke-a") + self.assertEqual(result.active_user_keys, ("42",)) + + async def test_failed_reconciliation_retains_new_node_wide_poison(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + controller._default_timeout = 1 + controller._sync_lease_seconds = 0.01 + await store.acquire_user_sync_lease("node-1", "dead-worker", ["42"], 0.01) + await asyncio.sleep(0.02) + controller._make_request = AsyncMock(side_effect=asyncio.TimeoutError) + + with self.assertRaises(asyncio.TimeoutError): + await RestNode.reconcile_users(controller, [_user("42")]) + await asyncio.sleep(0.02) + + with self.assertRaises(UserSyncLeaseLostError): + await controller.begin_user_revocation(["unrelated"], "revoke-b") + + async def test_execution_lease_on_other_node_does_not_block_begin(self): + store = InMemoryUserSyncStore() + other_node_lease = await store.acquire_user_sync_lease("node-2", "worker", ["42"], 30) + + await asyncio.wait_for( + store.begin_user_revocation("node-1", ["42"], "revoke-a"), + timeout=0.1, + ) + await store.release_user_sync_lease(other_node_lease) + + async def test_direct_sync_requires_matching_revocation_owner(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + await controller.begin_user_revocation(["42"], "revoke-a") + + with self.assertRaises(NodeAPIError) as error: + await controller._acquire_direct_user_sync_lease([_user("42")]) + self.assertEqual(error.exception.code, 409) + + lease, heartbeat = await controller._acquire_direct_user_sync_lease([_user("42")], revocation_id="revoke-a") + self.assertEqual(lease.user_keys, ("42",)) + await controller._release_user_sync_lease(lease, heartbeat) + + async def test_ambiguous_full_snapshot_timeout_retains_node_wide_poison(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + controller._sync_lease_seconds = 0.01 + controller._default_timeout = 1 + controller._node_lock = asyncio.Lock() + controller._make_request = AsyncMock(side_effect=asyncio.TimeoutError) + + with self.assertRaises(asyncio.TimeoutError): + await RestNode.sync_users(controller, [_user("42")]) + await asyncio.sleep(0.02) + + with self.assertRaises(UserSyncLeaseLostError): + await store.acquire_user_sync_lease("node-1", "retry", ["other-user"], 30) + with self.assertRaises(UserSyncLeaseLostError): + await controller.begin_user_revocation(["42"], "revoke-a") + + async def test_full_snapshot_rejects_revocation_id(self): + controller = _controller(InMemoryUserSyncStore()) + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + with self.assertRaises(NodeAPIError) as error: + await node_type.sync_users(controller, [_user("42")], revocation_id="revoke-a") + self.assertEqual(error.exception.code, 400) + + +class UserRevocationWorkerTests(unittest.IsolatedAsyncioTestCase): + async def test_cancellation_during_store_release_waits_for_cleanup_then_propagates(self): + class SlowReleaseStore(InMemoryUserSyncStore): + def __init__(self): + super().__init__() + self.release_entered = asyncio.Event() + self.finish_release = asyncio.Event() + + async def release_user_sync_lease(self, lease): + self.release_entered.set() + await self.finish_release.wait() + await super().release_user_sync_lease(lease) + + store = SlowReleaseStore() + controller = _controller(store) + lease = await store.acquire_user_sync_lease("node-1", "worker-1", ["42"], 30) + release = asyncio.create_task(controller._release_user_sync_lease(lease)) + await asyncio.wait_for(store.release_entered.wait(), timeout=1) + release.cancel() + await asyncio.sleep(0) + self.assertFalse(release.done()) + + store.finish_release.set() + with self.assertRaises(asyncio.CancelledError): + await release + self.assertNotIn(lease.token, store._user_sync_leases) + + async def test_cleanup_error_is_not_masked_by_caller_cancellation(self): + class FailingReleaseStore(InMemoryUserSyncStore): + def __init__(self): + super().__init__() + self.release_entered = asyncio.Event() + self.finish_release = asyncio.Event() + + async def release_user_sync_lease(self, _lease): + self.release_entered.set() + await self.finish_release.wait() + raise RuntimeError("release failed") + + store = FailingReleaseStore() + controller = _controller(store) + lease = await store.acquire_user_sync_lease("node-1", "worker-1", ["42"], 30) + release = asyncio.create_task(controller._release_user_sync_lease(lease)) + await asyncio.wait_for(store.release_entered.wait(), timeout=1) + release.cancel() + store.finish_release.set() + + with self.assertRaisesRegex(RuntimeError, "release failed"): + await release + + async def test_worker_retries_stale_epoch_with_new_lease_without_poison(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + observed_epochs = [] + + async def sync_batch(_users, user_sync_epoch): + observed_epochs.append(user_sync_epoch) + if len(observed_epochs) == 1: + raise NodeAPIError(412, "stale user sync epoch") + controller._shutdown_event.set() + return [] + + controller._sync_batch_users = sync_batch + await store.enqueue_users("node-1", [_user("42")]) + controller._work_available.set() + + async def no_delay(_seconds): + return None + + with patch("PasarGuardNodeBridge.controller.asyncio.sleep", side_effect=no_delay): + await controller._sync_worker() + + self.assertEqual(observed_epochs, [1, 2]) + self.assertEqual(store._user_sync_leases, {}) + self.assertIsNone(await store.next_claim_delay("node-1")) + + async def test_release_lease_propagates_caller_cancellation_after_cleanup(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + lease = await store.acquire_user_sync_lease("node-1", "worker-1", ["42"], 30) + heartbeat_cancelling = asyncio.Event() + + async def slow_heartbeat_shutdown(): + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + heartbeat_cancelling.set() + await asyncio.Event().wait() + + heartbeat = asyncio.create_task(slow_heartbeat_shutdown()) + release = asyncio.create_task(controller._release_user_sync_lease(lease, heartbeat)) + await asyncio.wait_for(heartbeat_cancelling.wait(), timeout=1) + release.cancel() + + with self.assertRaises(asyncio.CancelledError): + await release + self.assertNotIn(lease.token, store._user_sync_leases) + + async def test_abandon_lease_propagates_caller_cancellation_and_keeps_poison(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + lease = await store.acquire_user_sync_lease("node-1", "worker-1", ["42"], 30) + heartbeat_cancelling = asyncio.Event() + + async def slow_heartbeat_shutdown(): + try: + await asyncio.Event().wait() + except asyncio.CancelledError: + heartbeat_cancelling.set() + await asyncio.Event().wait() + + heartbeat = asyncio.create_task(slow_heartbeat_shutdown()) + abandon = asyncio.create_task(controller._abandon_user_sync_lease(lease, heartbeat)) + await asyncio.wait_for(heartbeat_cancelling.wait(), timeout=1) + abandon.cancel() + + with self.assertRaises(asyncio.CancelledError): + await abandon + self.assertIn(lease.token, store._user_sync_leases) + + async def test_narrowing_stops_original_heartbeat_before_replacing_lease(self): + class YieldAfterNarrowStore(InMemoryUserSyncStore): + async def retain_user_sync_lease_keys(self, lease, retained_user_keys): + narrowed = await super().retain_user_sync_lease_keys(lease, retained_user_keys) + await asyncio.sleep(0.02) + return narrowed + + store = YieldAfterNarrowStore() + controller = _controller(store) + controller._sync_lease_seconds = 0.03 + lease, heartbeat = await controller._acquire_direct_user_sync_lease([_user("ok"), _user("failed")]) + + await controller._retain_unknown_user_sync_lease_keys(lease, heartbeat, ["failed"]) + + self.assertTrue(heartbeat.done()) + if not heartbeat.cancelled(): + self.assertIsNone(heartbeat.exception()) + retained, _ = store._user_sync_leases[lease.token] + self.assertEqual(retained.user_keys, ("failed",)) + + async def test_partial_failure_retains_poison_only_for_unknown_key(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + controller._sync_lease_seconds = 0.01 + lease = await controller._acquire_user_sync_lease(["ok", "failed"]) + + await controller._retain_unknown_user_sync_lease_keys(lease, None, ["failed"]) + + ok_result = await controller.begin_user_revocation(["ok"], "revoke-ok") + self.assertEqual(ok_result.active_user_keys, ("ok",)) + await controller.abort_user_revocation(["ok"], "revoke-ok") + await asyncio.sleep(0.11) + with self.assertRaises(UserSyncLeaseLostError): + await controller.begin_user_revocation(["failed"], "revoke-failed") + + async def test_partial_failure_split_is_atomic_with_concurrent_begin(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + controller._sync_lease_seconds = 0.01 + lease = await controller._acquire_user_sync_lease(["ok", "failed"]) + + failed_begin = asyncio.create_task(controller.begin_user_revocation(["failed"], "revoke-failed")) + await _wait_until( + lambda: ( + (state := store._revocations.get("node-1", {}).get("failed")) is not None + and state.active_owner == "revoke-failed" + ) + ) + + await controller._retain_unknown_user_sync_lease_keys(lease, None, ["failed"]) + + ok_result = await controller.begin_user_revocation(["ok"], "revoke-ok") + self.assertEqual(ok_result.active_user_keys, ("ok",)) + await controller.abort_user_revocation(["ok"], "revoke-ok") + with self.assertRaises(UserSyncLeaseLostError): + await asyncio.wait_for(failed_begin, timeout=1) + + async def test_begin_waits_for_claimed_inflight_sync(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + revoker = _controller(store, "worker-2") + entered = asyncio.Event() + release = asyncio.Event() + applied: list[list[str]] = [] + + async def sync_batch(users): + entered.set() + await release.wait() + applied.extend(list(user.inbounds) for user in users) + return [] + + controller._sync_batch_users = sync_batch + await store.enqueue_users("node-1", [_user("42", "stale")]) + controller._work_available.set() + worker = asyncio.create_task(controller._sync_worker()) + controller._sync_worker_task = worker + await asyncio.wait_for(entered.wait(), timeout=1) + + begin = asyncio.create_task(revoker.begin_user_revocation(["42"], "revoke-a")) + await asyncio.sleep(0) + self.assertFalse(begin.done()) + release.set() + await asyncio.wait_for(begin, timeout=1) + self.assertEqual(applied, [["stale"]]) + + await revoker.finalize_user_revocation(["42"], "revoke-a") + await controller.update_user(_user("42", "late")) + worker.cancel() + await worker + self.assertEqual(await store.claim_users("node-1", "reader", 10, 30), []) + + async def test_chunked_worker_uses_outer_lease_without_nested_admission_race(self): + store = InMemoryUserSyncStore(max_pending_users_per_node=2000) + controller = _controller(store) + revoker = _controller(store, "worker-2") + controller._supports_chunked_sync = AsyncMock(return_value=(True, "0.2.0")) + entered = asyncio.Event() + release = asyncio.Event() + + async def sync_chunked(_users, _chunk_size, _timeout): + entered.set() + await release.wait() + + controller._sync_users_chunked_transport = sync_chunked + await store.enqueue_users("node-1", [_user(str(index)) for index in range(1000)]) + controller._work_available.set() + worker = asyncio.create_task(controller._sync_worker()) + controller._sync_worker_task = worker + await asyncio.wait_for(entered.wait(), timeout=1) + + begin = asyncio.create_task(revoker.begin_user_revocation(["0"], "revoke-a")) + await asyncio.sleep(0) + self.assertFalse(begin.done()) + release.set() + result = await asyncio.wait_for(begin, timeout=1) + self.assertEqual(result.active_user_keys, ("0",)) + self.assertFalse(worker.done()) + + worker.cancel() + await worker + self.assertEqual(await store.claim_users("node-1", "reader", 2000, 30), []) + + async def test_legacy_subclass_without_transport_hook_falls_back_to_batch(self): + store = InMemoryUserSyncStore(max_pending_users_per_node=2000) + controller = _controller(store) + revoker = _controller(store, "worker-2") + controller._supports_chunked_sync = AsyncMock(return_value=(True, "0.2.0")) + controller.sync_users_chunked = AsyncMock() + entered = asyncio.Event() + release = asyncio.Event() + + async def sync_batch(_users): + entered.set() + await release.wait() + return [] + + controller._sync_batch_users = sync_batch + await store.enqueue_users("node-1", [_user(str(index)) for index in range(1000)]) + controller._work_available.set() + worker = asyncio.create_task(controller._sync_worker()) + controller._sync_worker_task = worker + await asyncio.wait_for(entered.wait(), timeout=1) + + begin = asyncio.create_task(revoker.begin_user_revocation(["0"], "revoke-a")) + await asyncio.sleep(0) + self.assertFalse(begin.done()) + release.set() + await asyncio.wait_for(begin, timeout=1) + controller.sync_users_chunked.assert_not_awaited() + + worker.cancel() + await worker + + async def test_worker_cancellation_cannot_requeue_stale_claim_after_begin(self): + store = InMemoryUserSyncStore() + controller = _controller(store) + controller._sync_lease_seconds = 0.01 + entered = asyncio.Event() + + async def sync_batch(_users): + entered.set() + await asyncio.Event().wait() + + controller._sync_batch_users = sync_batch + await store.enqueue_users("node-1", [_user("42", "stale")]) + controller._work_available.set() + worker = asyncio.create_task(controller._sync_worker()) + controller._sync_worker_task = worker + await asyncio.wait_for(entered.wait(), timeout=1) + + begin = asyncio.create_task(controller.begin_user_revocation(["42"], "revoke-a")) + await asyncio.sleep(0) + worker.cancel() + await worker + with self.assertRaises(UserSyncLeaseLostError): + await asyncio.wait_for(begin, timeout=1) + + self.assertEqual(await store.claim_users("node-1", "reader", 10, 30), []) + + +class StartupRevocationTests(unittest.IsolatedAsyncioTestCase): + @staticmethod + def _startup_controller(store): + controller = _controller(store) + controller._default_timeout = 1 + controller.get_health = AsyncMock(return_value=Health.HEALTHY) + controller._acquire_lifecycle_lease = AsyncMock(return_value=None) + controller._release_lifecycle_lease = AsyncMock() + controller.connect = AsyncMock() + controller.disconnect = AsyncMock() + return controller + + async def _run_start(self, node_type, controller, transport, users=None, **kwargs): + if node_type is RestNode: + controller._make_request = transport + else: + controller._client = SimpleNamespace(Start=object()) + controller._handle_grpc_request = transport + return await node_type.start( + controller, + config="{}", + backend_type=service.BackendType.XRAY, + users=users if users is not None else [_user("stale"), _user("keep")], + **kwargs, + ) + + async def _run_full_sync(self, node_type, controller, transport, users=None): + if node_type is RestNode: + controller._make_request = transport + else: + controller._client = SimpleNamespace(SyncUsers=object()) + controller._handle_grpc_request = transport + return await node_type.sync_users( + controller, + users=users if users is not None else [_user("stale"), _user("keep")], + ) + + async def test_begin_waits_for_startup_snapshot_already_inflight(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + controller = self._startup_controller(store) + revoker = _controller(store, "revoker") + entered = asyncio.Event() + release = asyncio.Event() + captured = [] + + async def transport( + _captured=captured, + _entered=entered, + _release=release, + **kwargs, + ): + request = kwargs["proto_message"] if "proto_message" in kwargs else kwargs["request"] + _captured.extend(user.email for user in request.users) + _entered.set() + await _release.wait() + return service.BaseInfoResponse(started=True, node_version="0.2.0", core_version="1.0.0") + + start = asyncio.create_task(self._run_start(node_type, controller, transport)) + await asyncio.wait_for(entered.wait(), timeout=1) + begin = asyncio.create_task(revoker.begin_user_revocation(["stale"], "revoke-a")) + await asyncio.sleep(0) + self.assertFalse(begin.done()) + + release.set() + await asyncio.wait_for(start, timeout=1) + result = await asyncio.wait_for(begin, timeout=1) + self.assertEqual(result.active_user_keys, ("stale",)) + self.assertEqual(captured, ["stale", "keep"]) + await revoker.abort_user_revocation(["stale"], "revoke-a") + + async def test_authoritative_start_recovers_expired_node_wide_poison(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + controller = self._startup_controller(store) + lost = (await store.acquire_startup_user_sync_lease("node-1", "dead-worker", ["stale"], 0.01)).lease + await asyncio.sleep(0.02) + captured: list[str] = [] + + async def transport(_captured=captured, **kwargs): + request = kwargs["proto_message"] if "proto_message" in kwargs else kwargs["request"] + _captured.extend(user.email for user in request.users) + return service.BaseInfoResponse(started=True, node_version="0.2.0", core_version="1.0.0") + + await self._run_start(node_type, controller, transport, reconcile_user_sync=True) + + self.assertEqual(captured, ["stale", "keep"]) + self.assertNotIn(lost.token, store._user_sync_leases) + + async def test_failed_reconcile_leaves_recoverable_wildcard_poison(self): + store = InMemoryUserSyncStore() + controller = self._startup_controller(store) + controller._sync_lease_seconds = 0.01 + + async def failed_transport(**_kwargs): + raise TimeoutError + + with self.assertRaises(TimeoutError): + await self._run_start( + RestNode, + controller, + failed_transport, + reconcile_user_sync=True, + ) + await asyncio.sleep(0.02) + + captured: list[int] = [] + + async def successful_transport(**kwargs): + request = kwargs["proto_message"] + captured.append(request.user_sync_epoch) + return service.BaseInfoResponse( + started=True, + node_version="0.2.0", + core_version="1.0.0", + user_sync_epoch_supported=True, + ) + + await self._run_start( + RestNode, + controller, + successful_transport, + reconcile_user_sync=True, + ) + self.assertEqual(len(captured), 1) + self.assertGreater(captured[0], 0) + self.assertEqual(store._user_sync_leases, {}) + + async def test_startup_waits_for_provisional_restore_and_abort(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + controller = self._startup_controller(store) + await controller.begin_user_revocation(["stale"], "revoke-a") + restore_lease = await store.acquire_user_sync_lease( + "node-1", "restore", ["stale"], 30, revocation_id="revoke-a" + ) + entered = asyncio.Event() + captured = [] + + async def transport(_captured=captured, _entered=entered, **kwargs): + request = kwargs["proto_message"] if "proto_message" in kwargs else kwargs["request"] + _captured.extend(user.email for user in request.users) + _entered.set() + return service.BaseInfoResponse(started=True, node_version="0.2.0", core_version="1.0.0") + + start = asyncio.create_task(self._run_start(node_type, controller, transport)) + await asyncio.sleep(0) + self.assertFalse(entered.is_set()) + + abort = asyncio.create_task(controller.abort_user_revocation(["stale"], "revoke-a")) + await asyncio.sleep(0) + self.assertFalse(abort.done()) + await store.release_user_sync_lease(restore_lease) + await asyncio.wait_for(abort, timeout=1) + await asyncio.wait_for(start, timeout=1) + self.assertEqual(captured, ["stale", "keep"]) + + async def test_node_wide_start_blocks_begin_for_omitted_key(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + controller = self._startup_controller(store) + revoker = _controller(store, "revoker") + entered = asyncio.Event() + release = asyncio.Event() + + async def transport(_entered=entered, _release=release, **_kwargs): + _entered.set() + await _release.wait() + return service.BaseInfoResponse(started=True, node_version="0.2.0", core_version="1.0.0") + + start = asyncio.create_task(self._run_start(node_type, controller, transport, users=[_user("keep")])) + await asyncio.wait_for(entered.wait(), timeout=1) + begin = asyncio.create_task(revoker.begin_user_revocation(["omitted"], "revoke-a")) + await asyncio.sleep(0) + self.assertFalse(begin.done()) + + release.set() + await asyncio.wait_for(start, timeout=1) + await asyncio.wait_for(begin, timeout=1) + await revoker.abort_user_revocation(["omitted"], "revoke-a") + + async def test_startup_and_incremental_leases_are_mutually_exclusive(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + controller = self._startup_controller(store) + old_update = await store.acquire_user_sync_lease("node-1", "old-update", ["omitted"], 30) + entered = asyncio.Event() + release = asyncio.Event() + + async def transport(_entered=entered, _release=release, **_kwargs): + _entered.set() + await _release.wait() + return service.BaseInfoResponse(started=True, node_version="0.2.0", core_version="1.0.0") + + start = asyncio.create_task(self._run_start(node_type, controller, transport, users=[_user("keep")])) + await asyncio.sleep(0) + self.assertFalse(entered.is_set()) + + new_update = asyncio.create_task(store.acquire_user_sync_lease("node-1", "new-update", ["omitted"], 30)) + await asyncio.sleep(0) + self.assertFalse(new_update.done()) + + await store.release_user_sync_lease(old_update) + await asyncio.wait_for(entered.wait(), timeout=1) + self.assertFalse(new_update.done()) + + release.set() + await asyncio.wait_for(start, timeout=1) + admitted = await asyncio.wait_for(new_update, timeout=1) + self.assertEqual(admitted.user_keys, ("omitted",)) + await store.release_user_sync_lease(admitted) + + async def test_startup_snapshot_filters_finalized_user(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + controller = self._startup_controller(store) + await controller.begin_user_revocation(["stale"], "revoke-a") + await controller.finalize_user_revocation(["stale"], "revoke-a") + captured = [] + + async def transport(_captured=captured, **kwargs): + request = kwargs["proto_message"] if "proto_message" in kwargs else kwargs["request"] + _captured.extend(user.email for user in request.users) + return service.BaseInfoResponse(started=True, node_version="0.2.0", core_version="1.0.0") + + await self._run_start(node_type, controller, transport) + self.assertEqual(captured, ["keep"]) + + async def test_full_sync_snapshot_is_node_wide_and_filters_finalized_user(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + controller = self._startup_controller(store) + await controller.begin_user_revocation(["stale"], "revoke-a") + await controller.finalize_user_revocation(["stale"], "revoke-a") + captured = [] + + async def transport(_captured=captured, **kwargs): + request = kwargs["proto_message"] if "proto_message" in kwargs else kwargs["request"] + _captured.extend(user.email for user in request.users) + return service.Empty() + + await self._run_full_sync(node_type, controller, transport) + self.assertEqual(captured, ["keep"]) + + empty_users, lease, heartbeat = await controller._acquire_snapshot_user_sync_lease([]) + self.assertEqual(empty_users, []) + self.assertTrue(lease.covers_all_users) + self.assertTrue(lease.token) + await controller._release_user_sync_lease(lease, heartbeat) + + async def test_ambiguous_start_retains_user_sync_poison(self): + for node_type in (RestNode, GrpcNode): + with self.subTest(node_type=node_type.__module__): + store = InMemoryUserSyncStore() + controller = self._startup_controller(store) + controller._sync_lease_seconds = 0.01 + + async def transport(**_kwargs): + raise TimeoutError + + with self.assertRaises(TimeoutError): + await self._run_start(node_type, controller, transport) + await asyncio.sleep(0.11) + with self.assertRaises(UserSyncLeaseLostError): + await store.acquire_user_sync_lease("node-1", "late-update", ["omitted"], 30) + with self.assertRaises(UserSyncLeaseLostError): + await controller.begin_user_revocation(["omitted"], "revoke-a") + + +class LegacyStoreCompatibilityTests(unittest.IsolatedAsyncioTestCase): + async def test_normal_updates_work_but_revocation_fails_closed(self): + class LegacyStore: + def __init__(self): + self.users = [] + + async def enqueue_users(self, _node_id, users): + self.users.extend(users) + + store = LegacyStore() + controller = _controller(store) + controller._ensure_sync_worker_running = AsyncMock() + + await controller.update_user(_user("42")) + self.assertEqual([user.email for user in store.users], ["42"]) + with self.assertRaises(NodeAPIError) as error: + await controller.begin_user_revocation(["42"], "revoke-a") + self.assertEqual(error.exception.code, 501) + + +if __name__ == "__main__": + unittest.main() diff --git a/uv.lock b/uv.lock index e3fb60d..6b49f3e 100644 --- a/uv.lock +++ b/uv.lock @@ -467,7 +467,7 @@ wheels = [ [[package]] name = "pasarguard-node-bridge" -version = "0.9.0" +version = "0.10.0" source = { editable = "." } dependencies = [ { name = "aiohttp" },