Files

340 lines
12 KiB
Python

#!/usr/bin/env python3
"""Proxy TCP pour PAC Arkteos REG3."""
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,
format="%(asctime)s %(levelname)s %(message)s",
stream=sys.stdout,
)
logger = logging.getLogger(__name__)
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:
def __init__(
self,
pac_host: str,
pac_port: int,
proxy_port: int,
allow_client_writes: bool = False,
) -> None:
self.pac_host = pac_host
self.pac_port = pac_port
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."""
async with self.pac_write_lock:
if self.pac_writer is None:
raise ConnectionError("PAC non connectée")
self.pac_writer.write(data)
await self.pac_writer.drain()
async def send_keepalive(self) -> None:
await self.write_to_pac(b"\x00")
logger.info("Keepalive envoyé à la PAC")
async def pac_reader(self, reader: asyncio.StreamReader) -> None:
logger.info("Démarrage lecture PAC")
try:
while not self.stop_event.is_set():
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")
async def broadcast_to_clients(self, data: bytes) -> None:
dead_clients: list[asyncio.StreamWriter] = []
for writer in tuple(self.clients):
try:
writer.write(data)
await writer.drain()
except (ConnectionError, OSError):
dead_clients.append(writer)
for writer in dead_clients:
await self.close_client(writer)
async def pac_keepalive(self) -> None:
logger.info("Démarrage keepalive PAC")
try:
while not self.stop_event.is_set():
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")
async def handle_client(
self,
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)
blocked_write_logged = False
try:
while not self.stop_event.is_set():
data = await reader.read(1024)
if not data:
return
if not self.allow_client_writes:
if not blocked_write_logged:
logger.info("Tentative d’écriture client ignorée")
blocked_write_logged = True
continue
try:
await self.write_to_pac(data)
except (ConnectionError, OSError) as error:
logger.info("Erreur envoi PAC depuis %s : %s", peername, error)
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:
return
self.clients.discard(writer)
writer.close()
try:
await writer.wait_closed()
except (ConnectionError, OSError):
pass
logger.info("Client déconnecté nettoyé")
async def close_clients(self) -> None:
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:
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)
self.pac_reader_task = asyncio.create_task(self.pac_reader(reader))
self.keepalive_task = asyncio.create_task(self.pac_keepalive())
try:
try:
await self.pac_reader_task
except asyncio.CancelledError:
if not self.stop_event.is_set():
raise
finally:
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()
try:
await writer.wait_closed()
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 = (
"Mode bidirectionnel : écritures des clients autorisées"
if self.allow_client_writes
else "Mode lecture seule : écritures des clients bloquées"
)
logger.info(mode)
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:
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 self.wait_for_reconnect_delay()
finally:
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é")
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(run_proxy(proxy))
except KeyboardInterrupt:
logger.info("Arrêt demandé")
return 0
if __name__ == "__main__":
raise SystemExit(main())