fix: improve proxy liveness and graceful shutdown
This commit is contained in:
+171
-25
@@ -3,11 +3,14 @@
|
||||
|
||||
import asyncio
|
||||
import logging
|
||||
import signal
|
||||
import sys
|
||||
|
||||
|
||||
KEEPALIVE_INTERVAL = 300
|
||||
RECONNECT_DELAY = 10
|
||||
PAC_CONNECT_TIMEOUT = 10.0
|
||||
PAC_READ_TIMEOUT = 600.0
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.INFO,
|
||||
@@ -17,9 +20,53 @@ logging.basicConfig(
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def parse_allow_client_writes(value: str | None) -> bool:
|
||||
"""Retourne False tant que l'option n'est pas explicitement vraie."""
|
||||
return value is not None and value.strip().lower() == "true"
|
||||
class ConfigurationError(ValueError):
|
||||
"""Configuration de l'add-on invalide."""
|
||||
|
||||
|
||||
def parse_allow_client_writes(value: bool | str | None) -> bool:
|
||||
"""Accepte uniquement un booléen ou les chaînes explicites true et false."""
|
||||
if value is None:
|
||||
return False
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if isinstance(value, str):
|
||||
normalized = value.strip().lower()
|
||||
if normalized == "true":
|
||||
return True
|
||||
if normalized == "false":
|
||||
return False
|
||||
raise ConfigurationError("allow_client_writes doit être true ou false")
|
||||
|
||||
|
||||
def parse_port(value: int | str, option_name: str) -> int:
|
||||
"""Valide un port TCP sans accepter de valeur ambiguë."""
|
||||
if isinstance(value, bool):
|
||||
raise ConfigurationError(f"{option_name} doit être un entier entre 1 et 65535")
|
||||
if isinstance(value, str):
|
||||
if not value.strip().isdigit():
|
||||
raise ConfigurationError(f"{option_name} doit être un entier entre 1 et 65535")
|
||||
value = int(value.strip())
|
||||
if not isinstance(value, int) or not 1 <= value <= 65535:
|
||||
raise ConfigurationError(f"{option_name} doit être un entier entre 1 et 65535")
|
||||
return value
|
||||
|
||||
|
||||
def validate_configuration(
|
||||
pac_host: str,
|
||||
pac_port: int | str,
|
||||
proxy_port: int | str,
|
||||
allow_client_writes: bool | str | None,
|
||||
) -> tuple[str, int, int, bool]:
|
||||
"""Retourne une configuration validée avant toute ouverture réseau."""
|
||||
if not isinstance(pac_host, str) or not pac_host.strip():
|
||||
raise ConfigurationError("pac_host doit être une chaîne non vide")
|
||||
return (
|
||||
pac_host.strip(),
|
||||
parse_port(pac_port, "pac_port"),
|
||||
parse_port(proxy_port, "proxy_port"),
|
||||
parse_allow_client_writes(allow_client_writes),
|
||||
)
|
||||
|
||||
|
||||
class ArkteosProxy:
|
||||
@@ -35,9 +82,30 @@ class ArkteosProxy:
|
||||
self.proxy_port = proxy_port
|
||||
self.allow_client_writes = allow_client_writes
|
||||
self.clients: set[asyncio.StreamWriter] = set()
|
||||
self.client_tasks: set[asyncio.Task[None]] = set()
|
||||
self.pac_writer: asyncio.StreamWriter | None = None
|
||||
self.pac_write_lock = asyncio.Lock()
|
||||
self.stop_event = asyncio.Event()
|
||||
self.server: asyncio.AbstractServer | None = None
|
||||
self.pac_reader_task: asyncio.Task[None] | None = None
|
||||
self.keepalive_task: asyncio.Task[None] | None = None
|
||||
|
||||
def request_stop(self) -> None:
|
||||
"""Déclenche l'arrêt sans bloquer un gestionnaire de signal."""
|
||||
if self.stop_event.is_set():
|
||||
return
|
||||
logger.info("Arrêt du proxy demandé")
|
||||
self.stop_event.set()
|
||||
if self.server is not None:
|
||||
self.server.close()
|
||||
if self.keepalive_task is not None:
|
||||
self.keepalive_task.cancel()
|
||||
if self.pac_reader_task is not None:
|
||||
self.pac_reader_task.cancel()
|
||||
if self.pac_writer is not None:
|
||||
self.pac_writer.close()
|
||||
for writer in tuple(self.clients):
|
||||
writer.close()
|
||||
|
||||
async def write_to_pac(self, data: bytes) -> None:
|
||||
"""Écrit un bloc complet vers la PAC sans l'altérer."""
|
||||
@@ -55,13 +123,21 @@ class ArkteosProxy:
|
||||
logger.info("Démarrage lecture PAC")
|
||||
try:
|
||||
while not self.stop_event.is_set():
|
||||
data = await reader.read(4096)
|
||||
try:
|
||||
async with asyncio.timeout(PAC_READ_TIMEOUT):
|
||||
data = await reader.read(4096)
|
||||
except TimeoutError:
|
||||
logger.warning("PAC restée silencieuse pendant %.0f secondes", PAC_READ_TIMEOUT)
|
||||
return
|
||||
if not data:
|
||||
logger.info("PAC a fermé la connexion")
|
||||
return
|
||||
await self.broadcast_to_clients(data)
|
||||
except (ConnectionError, OSError) as error:
|
||||
logger.info("Erreur lecture PAC : %s", error)
|
||||
except asyncio.CancelledError:
|
||||
if not self.stop_event.is_set():
|
||||
raise
|
||||
finally:
|
||||
logger.info("Arrêt lecture PAC")
|
||||
|
||||
@@ -80,10 +156,17 @@ class ArkteosProxy:
|
||||
logger.info("Démarrage keepalive PAC")
|
||||
try:
|
||||
while not self.stop_event.is_set():
|
||||
await asyncio.sleep(KEEPALIVE_INTERVAL)
|
||||
await self.send_keepalive()
|
||||
try:
|
||||
async with asyncio.timeout(KEEPALIVE_INTERVAL):
|
||||
await self.stop_event.wait()
|
||||
return
|
||||
except TimeoutError:
|
||||
await self.send_keepalive()
|
||||
except (ConnectionError, OSError) as error:
|
||||
logger.info("Erreur keepalive PAC : %s", error)
|
||||
except asyncio.CancelledError:
|
||||
if not self.stop_event.is_set():
|
||||
raise
|
||||
finally:
|
||||
logger.info("Arrêt keepalive PAC")
|
||||
|
||||
@@ -92,6 +175,9 @@ class ArkteosProxy:
|
||||
reader: asyncio.StreamReader,
|
||||
writer: asyncio.StreamWriter,
|
||||
) -> None:
|
||||
task = asyncio.current_task()
|
||||
if task is not None:
|
||||
self.client_tasks.add(task)
|
||||
peername = writer.get_extra_info("peername")
|
||||
logger.info("Nouveau client : %s", peername)
|
||||
self.clients.add(writer)
|
||||
@@ -113,8 +199,13 @@ class ArkteosProxy:
|
||||
return
|
||||
except (ConnectionError, OSError) as error:
|
||||
logger.info("Erreur client %s : %s", peername, error)
|
||||
except asyncio.CancelledError:
|
||||
if not self.stop_event.is_set():
|
||||
raise
|
||||
finally:
|
||||
await self.close_client(writer)
|
||||
if task is not None:
|
||||
self.client_tasks.discard(task)
|
||||
|
||||
async def close_client(self, writer: asyncio.StreamWriter) -> None:
|
||||
if writer not in self.clients:
|
||||
@@ -131,17 +222,38 @@ class ArkteosProxy:
|
||||
for writer in tuple(self.clients):
|
||||
await self.close_client(writer)
|
||||
|
||||
async def stop_client_tasks(self) -> None:
|
||||
current_task = asyncio.current_task()
|
||||
tasks = [task for task in self.client_tasks if task is not current_task]
|
||||
for task in tasks:
|
||||
task.cancel()
|
||||
if tasks:
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
async def run_pac_connection(self) -> None:
|
||||
reader, writer = await asyncio.open_connection(self.pac_host, self.pac_port)
|
||||
try:
|
||||
async with asyncio.timeout(PAC_CONNECT_TIMEOUT):
|
||||
reader, writer = await asyncio.open_connection(self.pac_host, self.pac_port)
|
||||
except TimeoutError as error:
|
||||
logger.warning("Délai de connexion PAC dépassé après %.0f secondes", PAC_CONNECT_TIMEOUT)
|
||||
raise ConnectionError("Délai de connexion PAC dépassé") from error
|
||||
|
||||
self.pac_writer = writer
|
||||
logger.info("Connecté à la PAC %s:%s", self.pac_host, self.pac_port)
|
||||
reader_task = asyncio.create_task(self.pac_reader(reader))
|
||||
keepalive_task = asyncio.create_task(self.pac_keepalive())
|
||||
self.pac_reader_task = asyncio.create_task(self.pac_reader(reader))
|
||||
self.keepalive_task = asyncio.create_task(self.pac_keepalive())
|
||||
try:
|
||||
await reader_task
|
||||
try:
|
||||
await self.pac_reader_task
|
||||
except asyncio.CancelledError:
|
||||
if not self.stop_event.is_set():
|
||||
raise
|
||||
finally:
|
||||
keepalive_task.cancel()
|
||||
await asyncio.gather(keepalive_task, return_exceptions=True)
|
||||
if self.keepalive_task is not None:
|
||||
self.keepalive_task.cancel()
|
||||
await asyncio.gather(self.keepalive_task, return_exceptions=True)
|
||||
self.keepalive_task = None
|
||||
self.pac_reader_task = None
|
||||
if self.pac_writer is writer:
|
||||
self.pac_writer = None
|
||||
writer.close()
|
||||
@@ -150,6 +262,14 @@ class ArkteosProxy:
|
||||
except (ConnectionError, OSError):
|
||||
pass
|
||||
await self.close_clients()
|
||||
await self.stop_client_tasks()
|
||||
|
||||
async def wait_for_reconnect_delay(self) -> None:
|
||||
try:
|
||||
async with asyncio.timeout(RECONNECT_DELAY):
|
||||
await self.stop_event.wait()
|
||||
except TimeoutError:
|
||||
pass
|
||||
|
||||
async def serve(self) -> None:
|
||||
mode = (
|
||||
@@ -158,36 +278,62 @@ class ArkteosProxy:
|
||||
else "Mode lecture seule : écritures des clients bloquées"
|
||||
)
|
||||
logger.info(mode)
|
||||
server = await asyncio.start_server(self.handle_client, "0.0.0.0", self.proxy_port)
|
||||
self.server = await asyncio.start_server(self.handle_client, "0.0.0.0", self.proxy_port)
|
||||
logger.info("Proxy en écoute sur 0.0.0.0:%s", self.proxy_port)
|
||||
try:
|
||||
while not self.stop_event.is_set():
|
||||
try:
|
||||
await self.run_pac_connection()
|
||||
except (ConnectionError, OSError) as error:
|
||||
logger.info("Échec connexion PAC : %s", error)
|
||||
if not self.stop_event.is_set():
|
||||
logger.info("Échec connexion PAC : %s", error)
|
||||
if not self.stop_event.is_set():
|
||||
logger.info("Nouvelle tentative de connexion PAC dans 10s...")
|
||||
await asyncio.sleep(RECONNECT_DELAY)
|
||||
await self.wait_for_reconnect_delay()
|
||||
finally:
|
||||
self.stop_event.set()
|
||||
server.close()
|
||||
await server.wait_closed()
|
||||
self.request_stop()
|
||||
if self.server is not None:
|
||||
await self.server.wait_closed()
|
||||
self.server = None
|
||||
await self.close_clients()
|
||||
await self.stop_client_tasks()
|
||||
logger.info("Proxy arrêté")
|
||||
|
||||
|
||||
def main() -> None:
|
||||
pac_host = sys.argv[1] if len(sys.argv) > 1 else "192.168.X.X"
|
||||
pac_port = int(sys.argv[2]) if len(sys.argv) > 2 else 9641
|
||||
proxy_port = int(sys.argv[3]) if len(sys.argv) > 3 else 9641
|
||||
allow_client_writes = parse_allow_client_writes(sys.argv[4] if len(sys.argv) > 4 else None)
|
||||
async def run_proxy(proxy: ArkteosProxy) -> None:
|
||||
loop = asyncio.get_running_loop()
|
||||
installed_signals: list[signal.Signals] = []
|
||||
for signal_name in (signal.SIGTERM, signal.SIGINT):
|
||||
try:
|
||||
loop.add_signal_handler(signal_name, proxy.request_stop)
|
||||
installed_signals.append(signal_name)
|
||||
except (NotImplementedError, RuntimeError):
|
||||
pass
|
||||
try:
|
||||
await proxy.serve()
|
||||
finally:
|
||||
for signal_name in installed_signals:
|
||||
loop.remove_signal_handler(signal_name)
|
||||
|
||||
|
||||
def main() -> int:
|
||||
try:
|
||||
pac_host, pac_port, proxy_port, allow_client_writes = validate_configuration(
|
||||
sys.argv[1] if len(sys.argv) > 1 else "192.168.X.X",
|
||||
sys.argv[2] if len(sys.argv) > 2 else "9641",
|
||||
sys.argv[3] if len(sys.argv) > 3 else "9641",
|
||||
sys.argv[4] if len(sys.argv) > 4 else None,
|
||||
)
|
||||
except ConfigurationError as error:
|
||||
logger.error("Configuration invalide : %s", error)
|
||||
return 2
|
||||
proxy = ArkteosProxy(pac_host, pac_port, proxy_port, allow_client_writes)
|
||||
try:
|
||||
asyncio.run(proxy.serve())
|
||||
asyncio.run(run_proxy(proxy))
|
||||
except KeyboardInterrupt:
|
||||
logger.info("Arrêt demandé")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
raise SystemExit(main())
|
||||
|
||||
Reference in New Issue
Block a user