mirror of
https://github.com/jkingsman/Remote-Terminal-for-MeshCore.git
synced 2026-08-08 09:43:03 +02:00
Pass 2
This commit is contained in:
+62
-38
@@ -1,22 +1,25 @@
|
||||
"""Web Push dispatch manager.
|
||||
|
||||
Handles filtering subscriptions by their preferences and sending push
|
||||
notifications concurrently when a new message arrives.
|
||||
Checks the global push-enabled conversation list (stored in app_settings)
|
||||
and sends push notifications to ALL registered devices when a matching
|
||||
incoming message arrives.
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
import json
|
||||
import logging
|
||||
from dataclasses import dataclass
|
||||
|
||||
from pywebpush import WebPushException
|
||||
|
||||
from app.push.send import send_push
|
||||
from app.push.vapid import get_vapid_private_key
|
||||
from app.repository.push_subscriptions import PushSubscriptionRepository
|
||||
from app.repository.settings import AppSettingsRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SEND_TIMEOUT = 10 # seconds per push send
|
||||
_SEND_TIMEOUT = 15 # seconds per push send
|
||||
_VAPID_CLAIMS = {"sub": "mailto:noreply@meshcore.local"}
|
||||
|
||||
|
||||
@@ -29,19 +32,6 @@ def _state_key_for_message(data: dict) -> str:
|
||||
return f"channel-{conversation_key}"
|
||||
|
||||
|
||||
def _matches_filter(sub: dict, data: dict) -> bool:
|
||||
"""Check whether a message event matches a subscription's filter."""
|
||||
mode = sub.get("filter_mode", "all_messages")
|
||||
if mode == "all_messages":
|
||||
return True
|
||||
if mode == "all_dms":
|
||||
return data.get("type") == "PRIV"
|
||||
if mode == "selected":
|
||||
key = _state_key_for_message(data)
|
||||
return key in (sub.get("filter_conversations") or [])
|
||||
return False
|
||||
|
||||
|
||||
def _build_payload(data: dict) -> str:
|
||||
"""Build the push notification JSON payload from a message event."""
|
||||
msg_type = data.get("type", "")
|
||||
@@ -53,11 +43,11 @@ def _build_payload(data: dict) -> str:
|
||||
title = f"Message from {sender_name}" if sender_name else "New direct message"
|
||||
body = text
|
||||
else:
|
||||
# Channel messages include "SenderName: text" in the text field
|
||||
title = f"#{channel_name}" if channel_name else "Channel message"
|
||||
title = channel_name if channel_name else "Channel message"
|
||||
body = text
|
||||
|
||||
conversation_key = data.get("conversation_key", "")
|
||||
state_key = _state_key_for_message(data)
|
||||
if msg_type == "PRIV":
|
||||
url_hash = f"#contact/{conversation_key}"
|
||||
else:
|
||||
@@ -67,7 +57,10 @@ def _build_payload(data: dict) -> str:
|
||||
{
|
||||
"title": title,
|
||||
"body": body,
|
||||
"tag": f"meshcore-{data.get('id', '')}",
|
||||
# Tag per conversation so different conversations coexist in the
|
||||
# notification tray, while repeated messages in the same
|
||||
# conversation replace each other.
|
||||
"tag": f"meshcore-{state_key}",
|
||||
"url_hash": url_hash,
|
||||
}
|
||||
)
|
||||
@@ -84,13 +77,31 @@ def _subscription_info(sub: dict) -> dict:
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class _SendResult:
|
||||
sub_id: str
|
||||
success: bool = False
|
||||
expired: bool = False
|
||||
|
||||
|
||||
class PushManager:
|
||||
async def dispatch_message(self, data: dict) -> None:
|
||||
"""Send push notifications for a message event to matching subscriptions."""
|
||||
"""Send push notifications for a message event to all devices."""
|
||||
# Don't notify for messages the operator just sent themselves
|
||||
if data.get("outgoing"):
|
||||
return
|
||||
|
||||
# Check the global conversation list
|
||||
state_key = _state_key_for_message(data)
|
||||
try:
|
||||
push_conversations = await AppSettingsRepository.get_push_conversations()
|
||||
except Exception:
|
||||
logger.debug("Push dispatch: failed to load push_conversations", exc_info=True)
|
||||
return
|
||||
|
||||
if state_key not in push_conversations:
|
||||
return
|
||||
|
||||
try:
|
||||
subs = await PushSubscriptionRepository.get_all()
|
||||
except Exception:
|
||||
@@ -100,21 +111,40 @@ class PushManager:
|
||||
if not subs:
|
||||
return
|
||||
|
||||
matching = [s for s in subs if _matches_filter(s, data)]
|
||||
if not matching:
|
||||
return
|
||||
|
||||
payload = _build_payload(data)
|
||||
vapid_key = get_vapid_private_key()
|
||||
if not vapid_key:
|
||||
logger.debug("Push dispatch: no VAPID key configured, skipping")
|
||||
return
|
||||
|
||||
tasks = [self._send_one(sub, payload, vapid_key) for sub in matching]
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
results = await asyncio.gather(
|
||||
*(self._send_one(sub, payload, vapid_key) for sub in subs),
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
async def _send_one(self, sub: dict, payload: str, vapid_key: str) -> None:
|
||||
# Batch-update all delivery outcomes in one transaction.
|
||||
success_ids: list[str] = []
|
||||
failure_ids: list[str] = []
|
||||
remove_ids: list[str] = []
|
||||
for r in results:
|
||||
if isinstance(r, _SendResult):
|
||||
if r.expired:
|
||||
remove_ids.append(r.sub_id)
|
||||
elif r.success:
|
||||
success_ids.append(r.sub_id)
|
||||
else:
|
||||
failure_ids.append(r.sub_id)
|
||||
if success_ids or failure_ids or remove_ids:
|
||||
try:
|
||||
await PushSubscriptionRepository.batch_record_outcomes(
|
||||
success_ids, failure_ids, remove_ids
|
||||
)
|
||||
except Exception:
|
||||
logger.debug("Push dispatch: failed to record outcomes", exc_info=True)
|
||||
|
||||
async def _send_one(self, sub: dict, payload: str, vapid_key: str) -> _SendResult:
|
||||
sub_id = sub["id"]
|
||||
result = _SendResult(sub_id=sub_id)
|
||||
try:
|
||||
async with asyncio.timeout(_SEND_TIMEOUT):
|
||||
await send_push(
|
||||
@@ -123,26 +153,20 @@ class PushManager:
|
||||
vapid_private_key=vapid_key,
|
||||
vapid_claims=_VAPID_CLAIMS,
|
||||
)
|
||||
await PushSubscriptionRepository.record_success(sub_id)
|
||||
result.success = True
|
||||
except WebPushException as e:
|
||||
status = getattr(e, "response", None)
|
||||
status_code = getattr(status, "status_code", 0) if status else 0
|
||||
if status_code in (404, 410):
|
||||
logger.info(
|
||||
"Push subscription expired (HTTP %d), removing %s",
|
||||
status_code,
|
||||
sub_id,
|
||||
)
|
||||
await PushSubscriptionRepository.delete(sub_id)
|
||||
if status_code in (403, 404, 410):
|
||||
logger.info("Push subscription expired (HTTP %d), removing %s", status_code, sub_id)
|
||||
result.expired = True
|
||||
else:
|
||||
logger.warning("Push send failed for %s: %s", sub_id, e)
|
||||
await PushSubscriptionRepository.record_failure(sub_id)
|
||||
except TimeoutError:
|
||||
logger.warning("Push send timed out for %s", sub_id)
|
||||
await PushSubscriptionRepository.record_failure(sub_id)
|
||||
except Exception:
|
||||
logger.debug("Push send error for %s", sub_id, exc_info=True)
|
||||
await PushSubscriptionRepository.record_failure(sub_id)
|
||||
return result
|
||||
|
||||
|
||||
push_manager = PushManager()
|
||||
|
||||
+194
-8
@@ -6,11 +6,201 @@ a thread executor to avoid blocking the event loop.
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import socket
|
||||
from typing import Any, cast
|
||||
|
||||
import requests
|
||||
import urllib3.connection
|
||||
import urllib3.connectionpool
|
||||
from pywebpush import webpush
|
||||
from requests.adapters import HTTPAdapter
|
||||
from requests.exceptions import ConnectionError as RequestsConnectionError
|
||||
from requests.exceptions import ConnectTimeout as RequestsConnectTimeout
|
||||
from urllib3.exceptions import ConnectTimeoutError, NameResolutionError, NewConnectionError
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
DEFAULT_TIMEOUT = object()
|
||||
DEFAULT_PUSH_CONNECT_TIMEOUT_SECONDS = 3
|
||||
IPV4_FALLBACK_CONNECT_TIMEOUT_SECONDS = 10
|
||||
DEFAULT_PUSH_READ_TIMEOUT_SECONDS = 10
|
||||
|
||||
|
||||
def _create_ipv4_connection(
|
||||
address: tuple[str, int],
|
||||
timeout: float | None | object = DEFAULT_TIMEOUT,
|
||||
source_address: tuple[str, int] | None = None,
|
||||
socket_options=None,
|
||||
) -> socket.socket:
|
||||
"""Create a socket connection using IPv4 only."""
|
||||
host, port = address
|
||||
if host.startswith("["):
|
||||
host = host.strip("[]")
|
||||
|
||||
err: OSError | None = None
|
||||
for res in socket.getaddrinfo(host, port, socket.AF_INET, socket.SOCK_STREAM):
|
||||
af, socktype, proto, _, sa = res
|
||||
sock = None
|
||||
try:
|
||||
sock = socket.socket(af, socktype, proto)
|
||||
if socket_options:
|
||||
for opt in socket_options:
|
||||
sock.setsockopt(*opt)
|
||||
if timeout is not DEFAULT_TIMEOUT:
|
||||
sock.settimeout(cast(float | None, timeout))
|
||||
if source_address:
|
||||
sock.bind(source_address)
|
||||
sock.connect(sa)
|
||||
return sock
|
||||
except OSError as exc:
|
||||
err = exc
|
||||
if sock is not None:
|
||||
sock.close()
|
||||
|
||||
if err is not None:
|
||||
raise err
|
||||
raise OSError("getaddrinfo returns an empty list")
|
||||
|
||||
|
||||
class IPv4HTTPConnection(urllib3.connection.HTTPConnection):
|
||||
"""urllib3 HTTP connection that resolves and connects via IPv4 only."""
|
||||
|
||||
def _new_conn(self) -> socket.socket:
|
||||
try:
|
||||
return _create_ipv4_connection(
|
||||
(self._dns_host, self.port),
|
||||
self.timeout,
|
||||
source_address=self.source_address,
|
||||
socket_options=self.socket_options,
|
||||
)
|
||||
except socket.gaierror as exc:
|
||||
raise NameResolutionError(self.host, self, exc) from exc
|
||||
except TimeoutError as exc:
|
||||
raise ConnectTimeoutError(
|
||||
self,
|
||||
f"Connection to {self.host} timed out. (connect timeout={self.timeout})",
|
||||
) from exc
|
||||
except OSError as exc:
|
||||
raise NewConnectionError(self, f"Failed to establish a new connection: {exc}") from exc
|
||||
|
||||
|
||||
class IPv4HTTPSConnection(urllib3.connection.HTTPSConnection):
|
||||
"""urllib3 HTTPS connection that resolves and connects via IPv4 only."""
|
||||
|
||||
def _new_conn(self) -> socket.socket:
|
||||
try:
|
||||
return _create_ipv4_connection(
|
||||
(self._dns_host, self.port),
|
||||
self.timeout,
|
||||
source_address=self.source_address,
|
||||
socket_options=self.socket_options,
|
||||
)
|
||||
except socket.gaierror as exc:
|
||||
raise NameResolutionError(self.host, self, exc) from exc
|
||||
except TimeoutError as exc:
|
||||
raise ConnectTimeoutError(
|
||||
self,
|
||||
f"Connection to {self.host} timed out. (connect timeout={self.timeout})",
|
||||
) from exc
|
||||
except OSError as exc:
|
||||
raise NewConnectionError(self, f"Failed to establish a new connection: {exc}") from exc
|
||||
|
||||
|
||||
class IPv4HTTPConnectionPool(urllib3.connectionpool.HTTPConnectionPool):
|
||||
ConnectionCls = cast(Any, IPv4HTTPConnection)
|
||||
|
||||
|
||||
class IPv4HTTPSConnectionPool(urllib3.connectionpool.HTTPSConnectionPool):
|
||||
ConnectionCls = cast(Any, IPv4HTTPSConnection)
|
||||
|
||||
|
||||
def _configure_pool_manager_for_ipv4(manager: Any) -> None:
|
||||
manager.pool_classes_by_scheme = manager.pool_classes_by_scheme.copy()
|
||||
manager.pool_classes_by_scheme["http"] = IPv4HTTPConnectionPool
|
||||
manager.pool_classes_by_scheme["https"] = IPv4HTTPSConnectionPool
|
||||
|
||||
|
||||
class IPv4HTTPAdapter(HTTPAdapter):
|
||||
"""requests adapter that uses IPv4-only urllib3 connection pools."""
|
||||
|
||||
def init_poolmanager(self, connections, maxsize, block=False, **pool_kwargs):
|
||||
super().init_poolmanager(connections, maxsize, block=block, **pool_kwargs)
|
||||
_configure_pool_manager_for_ipv4(self.poolmanager)
|
||||
|
||||
def proxy_manager_for(self, *args, **kwargs):
|
||||
manager = super().proxy_manager_for(*args, **kwargs)
|
||||
_configure_pool_manager_for_ipv4(manager)
|
||||
return manager
|
||||
|
||||
|
||||
def _build_default_requests_session() -> requests.Session:
|
||||
return requests.Session()
|
||||
|
||||
|
||||
def _build_ipv4_requests_session() -> requests.Session:
|
||||
session = requests.Session()
|
||||
adapter = IPv4HTTPAdapter()
|
||||
session.mount("http://", adapter)
|
||||
session.mount("https://", adapter)
|
||||
return session
|
||||
|
||||
|
||||
def _send_push_with_session(
|
||||
*,
|
||||
subscription_info: dict,
|
||||
payload: str,
|
||||
vapid_private_key: str,
|
||||
vapid_claims: dict,
|
||||
session: requests.Session,
|
||||
connect_timeout_seconds: int,
|
||||
) -> int:
|
||||
response = webpush(
|
||||
subscription_info=subscription_info,
|
||||
data=payload,
|
||||
vapid_private_key=vapid_private_key,
|
||||
vapid_claims=vapid_claims,
|
||||
content_encoding="aes128gcm",
|
||||
timeout=cast(Any, (connect_timeout_seconds, DEFAULT_PUSH_READ_TIMEOUT_SECONDS)),
|
||||
requests_session=session,
|
||||
)
|
||||
return response.status_code # type: ignore[union-attr]
|
||||
|
||||
|
||||
def _send_push_with_fallback(
|
||||
subscription_info: dict,
|
||||
payload: str,
|
||||
vapid_private_key: str,
|
||||
vapid_claims: dict,
|
||||
) -> int:
|
||||
"""Send using normal dual-stack resolution, then retry with IPv4-only on connect failures."""
|
||||
session = _build_default_requests_session()
|
||||
try:
|
||||
return _send_push_with_session(
|
||||
subscription_info=subscription_info,
|
||||
payload=payload,
|
||||
vapid_private_key=vapid_private_key,
|
||||
vapid_claims=vapid_claims,
|
||||
session=session,
|
||||
connect_timeout_seconds=DEFAULT_PUSH_CONNECT_TIMEOUT_SECONDS,
|
||||
)
|
||||
except (RequestsConnectTimeout, RequestsConnectionError) as exc:
|
||||
logger.info("Push delivery retrying via IPv4 after initial network failure: %s", exc)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
session = _build_ipv4_requests_session()
|
||||
try:
|
||||
return _send_push_with_session(
|
||||
subscription_info=subscription_info,
|
||||
payload=payload,
|
||||
vapid_private_key=vapid_private_key,
|
||||
vapid_claims=vapid_claims,
|
||||
session=session,
|
||||
connect_timeout_seconds=IPV4_FALLBACK_CONNECT_TIMEOUT_SECONDS,
|
||||
)
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
async def send_push(
|
||||
subscription_info: dict,
|
||||
@@ -23,7 +213,7 @@ async def send_push(
|
||||
Args:
|
||||
subscription_info: {"endpoint": ..., "keys": {"p256dh": ..., "auth": ...}}
|
||||
payload: JSON string to encrypt and send
|
||||
vapid_private_key: PEM-encoded VAPID private key
|
||||
vapid_private_key: base64url-encoded raw EC private key scalar
|
||||
vapid_claims: {"sub": "mailto:..."} or {"sub": "https://..."}
|
||||
|
||||
Returns:
|
||||
@@ -33,13 +223,9 @@ async def send_push(
|
||||
WebPushException: on push service error (caller handles 404/410 cleanup).
|
||||
"""
|
||||
loop = asyncio.get_running_loop()
|
||||
response = await loop.run_in_executor(
|
||||
return await loop.run_in_executor(
|
||||
None,
|
||||
lambda: webpush(
|
||||
subscription_info=subscription_info,
|
||||
data=payload,
|
||||
vapid_private_key=vapid_private_key,
|
||||
vapid_claims=vapid_claims,
|
||||
lambda: _send_push_with_fallback(
|
||||
subscription_info, payload, vapid_private_key, vapid_claims
|
||||
),
|
||||
)
|
||||
return response.status_code # type: ignore[union-attr]
|
||||
|
||||
+14
-19
@@ -1,7 +1,8 @@
|
||||
"""VAPID key management for Web Push.
|
||||
|
||||
Generates a P-256 key pair on first use and caches it in app_settings.
|
||||
The public key is served to browsers for PushManager.subscribe().
|
||||
Generates a P-256 key pair on first use and caches it in app_settings
|
||||
via ``AppSettingsRepository``. The public key is served to browsers
|
||||
for ``PushManager.subscribe()``.
|
||||
"""
|
||||
|
||||
import base64
|
||||
@@ -10,7 +11,7 @@ import logging
|
||||
from cryptography.hazmat.primitives.serialization import Encoding, PublicFormat
|
||||
from py_vapid import Vapid
|
||||
|
||||
from app.database import db
|
||||
from app.repository.settings import AppSettingsRepository
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -22,14 +23,10 @@ async def ensure_vapid_keys() -> tuple[str, str]:
|
||||
"""Read or generate VAPID keys. Call once at startup after DB connect."""
|
||||
global _cached_private_key, _cached_public_key
|
||||
|
||||
cursor = await db.conn.execute(
|
||||
"SELECT vapid_private_key, vapid_public_key FROM app_settings WHERE id = 1"
|
||||
)
|
||||
row = await cursor.fetchone()
|
||||
|
||||
if row and row["vapid_private_key"] and row["vapid_public_key"]:
|
||||
_cached_private_key = row["vapid_private_key"]
|
||||
_cached_public_key = row["vapid_public_key"]
|
||||
private, public = await AppSettingsRepository.get_vapid_keys()
|
||||
if private and public:
|
||||
_cached_private_key = private
|
||||
_cached_public_key = public
|
||||
logger.info("VAPID keys loaded from database")
|
||||
return _cached_private_key, _cached_public_key
|
||||
|
||||
@@ -37,19 +34,17 @@ async def ensure_vapid_keys() -> tuple[str, str]:
|
||||
vapid = Vapid()
|
||||
vapid.generate_keys()
|
||||
|
||||
# Private key as PEM for pywebpush
|
||||
_cached_private_key = vapid.private_pem().decode("utf-8")
|
||||
# Private key as base64url-encoded raw 32-byte EC scalar — the format
|
||||
# that pywebpush passes to ``Vapid.from_string()``.
|
||||
raw_priv = vapid.private_key.private_numbers().private_value.to_bytes(32, "big") # type: ignore[union-attr]
|
||||
_cached_private_key = base64.urlsafe_b64encode(raw_priv).rstrip(b"=").decode("ascii")
|
||||
|
||||
# Public key as uncompressed P-256 point, base64url-encoded (no padding)
|
||||
# for the browser Push API's applicationServerKey
|
||||
raw_pub = vapid.public_key.public_bytes(Encoding.X962, PublicFormat.UncompressedPoint) # type: ignore[union-attr]
|
||||
_cached_public_key = base64.urlsafe_b64encode(raw_pub).rstrip(b"=").decode("ascii")
|
||||
|
||||
await db.conn.execute(
|
||||
"UPDATE app_settings SET vapid_private_key = ?, vapid_public_key = ? WHERE id = 1",
|
||||
(_cached_private_key, _cached_public_key),
|
||||
)
|
||||
await db.conn.commit()
|
||||
await AppSettingsRepository.set_vapid_keys(_cached_private_key, _cached_public_key)
|
||||
logger.info("Generated and stored new VAPID key pair")
|
||||
|
||||
return _cached_private_key, _cached_public_key
|
||||
@@ -61,5 +56,5 @@ def get_vapid_public_key() -> str:
|
||||
|
||||
|
||||
def get_vapid_private_key() -> str:
|
||||
"""Return the cached VAPID private key (PEM). Must call ensure_vapid_keys() first."""
|
||||
"""Return the cached VAPID private key (base64url). Must call ensure_vapid_keys() first."""
|
||||
return _cached_private_key
|
||||
|
||||
Reference in New Issue
Block a user