diff --git a/tests/test_main.py b/tests/test_main.py index bea31a8..ee0d55e 100644 --- a/tests/test_main.py +++ b/tests/test_main.py @@ -3,6 +3,7 @@ from __future__ import annotations import asyncio +import ipaddress import sys from unittest.mock import MagicMock, patch @@ -64,6 +65,34 @@ def test_parse_stun_server__without_port_uses_stun_default(self): ) +class TestParseHostport: + def test_parse_hostport__bracketed_ipv6_without_port_uses_default(self): + """Return default port and IPv6Address when bracketed IPv6 address has no port.""" + from voip.__main__ import _parse_hostport + + assert _parse_hostport(None, None, "[::1]", default_port=5061) == ( + ipaddress.IPv6Address("::1"), + 5061, + ) + + def test_parse_hostport__bracketed_ipv6_with_port(self): + """Return explicit port and IPv6Address when bracketed IPv6 address includes a port.""" + from voip.__main__ import _parse_hostport + + assert _parse_hostport(None, None, "[::1]:5061") == ( + ipaddress.IPv6Address("::1"), + 5061, + ) + + def test_parse_hostport__unbracketed_ipv6_raises_bad_parameter(self): + """Raise BadParameter when an unbracketed IPv6 literal is given.""" + import click + from voip.__main__ import _parse_hostport + + with pytest.raises(click.BadParameter, match="enclosed in brackets"): + _parse_hostport(None, None, "::1") + + class TestVoIPCommand: def test_voip__verbose_flag(self): """Accept -v flag without error.""" diff --git a/voip/__main__.py b/voip/__main__.py index dfb5b92..b0f4510 100644 --- a/voip/__main__.py +++ b/voip/__main__.py @@ -3,6 +3,7 @@ import dataclasses import ipaddress import logging +import re import ssl import time @@ -29,13 +30,21 @@ SIP_TLS_PORT = 5061 +HOSTPORT_PATTERN: re.Pattern[str] = re.compile( + r"^(?:\[(?P[0-9a-fA-F:]+)\]|(?P[^:\[\]]+))" + r"(?::(?P\d+))?$" +) + + def _parse_hostport( ctx, param, value: str, default_port: int = 5061 -) -> tuple[str, int]: - """Parse `HOST[:PORT]` or `[IPv6HOST][:PORT]` into a `(host, port)` tuple. +) -> tuple[ipaddress.IPv4Address | ipaddress.IPv6Address | str, int]: + """Parse `HOST[:PORT]` or `[IPv6HOST][:PORT]` into a typed `(host, port)` tuple. IPv6 addresses must be enclosed in square brackets per RFC 2732, e.g. - ``[::1]:5061``. The returned host is the bare address without brackets. + ``[::1]:5061``. The returned host is an + [`IPv4Address`][ipaddress.IPv4Address] or [`IPv6Address`][ipaddress.IPv6Address] + when the value is a numeric IP address, otherwise a plain hostname string. Args: ctx: Click context. @@ -44,38 +53,24 @@ def _parse_hostport( default_port: Port to use when not specified. Returns: - Tuple of (host, port). + Tuple of (host, port) where host is an IP address object or hostname string. Raises: - click.BadParameter: When port is invalid. + click.BadParameter: When value is malformed (unbracketed IPv6 or invalid port). """ - if value.startswith("["): - bracket_end = value.find("]") - if bracket_end == -1: + if not (match := HOSTPORT_PATTERN.fullmatch(value)): + if value.count(":") > 1: raise click.BadParameter( - f"Unclosed bracket in IPv6 address: {value!r}.", param=param + f"IPv6 address must be enclosed in brackets, e.g. [{value}].", + param=param, ) - host = value[1:bracket_end] - remainder = value[bracket_end + 1 :] - if not remainder: - return host, default_port - if not remainder.startswith(":"): - raise click.BadParameter( - f"Expected ':port' after ']' in {value!r}.", param=param - ) - try: - return host, int(remainder[1:]) - except ValueError: - raise click.BadParameter( - f"Invalid port in {value!r}.", param=param - ) from None - host, _, port_str = value.rpartition(":") - if not host: - return value, default_port + raise click.BadParameter(f"Invalid host:port value: {value!r}.", param=param) + raw_host = match.group("ipv6") or match.group("host") + port = int(match.group("port")) if match.group("port") else default_port try: - return host, int(port_str) + return ipaddress.ip_address(raw_host), port except ValueError: - raise click.BadParameter(f"Invalid port in {value!r}.", param=param) from None + return raw_host, port def _parse_stun_server(ctx, param, value: str | None) -> tuple[str, int] | None: @@ -91,7 +86,8 @@ def _parse_stun_server(ctx, param, value: str | None) -> tuple[str, int] | None: """ if value is None or value.lower() == "none": return None - return _parse_hostport(ctx, param, value, default_port=3478) + host, port = _parse_hostport(ctx, param, value, default_port=3478) + return str(host), port class ConsoleMessageProtocol(SessionInitiationProtocol): @@ -220,8 +216,7 @@ def sip(ctx, aor, password, username, proxy, stun_server, no_tls, no_verify_tls) else: default_port = SIP_TCP_PORT if parsed_aor.scheme == "sip" else SIP_TLS_PORT port = parsed_aor.port if parsed_aor.port is not None else default_port - # asyncio.create_connection requires a plain str host, not an ipaddress object. - proxy_addr = (str(parsed_aor.host), port) + proxy_addr = (parsed_aor.host, port) use_tls = not no_tls and proxy_addr[1] != SIP_TCP_PORT # Build the canonical AOR; IPv6 hosts must be enclosed in brackets per RFC 2732. @@ -245,7 +240,7 @@ def sip(ctx, aor, password, username, proxy, stun_server, no_tls, no_verify_tls) async def _connect_sip( session_factory, - proxy_addr: tuple[str, int], + proxy_addr: tuple[ipaddress.IPv4Address | ipaddress.IPv6Address | str, int], use_tls: bool, no_verify_tls: bool, ) -> None: @@ -259,7 +254,7 @@ async def _connect_sip( ssl_context.verify_mode = ssl.CERT_NONE await loop.create_connection( session_factory, - host=proxy_addr[0], + host=str(proxy_addr[0]), port=proxy_addr[1], ssl=ssl_context, ) diff --git a/voip/sip/protocol.py b/voip/sip/protocol.py index 6101a58..e9fdb1b 100644 --- a/voip/sip/protocol.py +++ b/voip/sip/protocol.py @@ -163,7 +163,9 @@ def call_received(self, request: Request) -> None: #: When ``None`` the caller connects directly to the registrar server. #: The address may differ from the registrar domain derived from #: `aor` (e.g. ``proxy.carrier.com`` vs ``carrier.com``). - outbound_proxy: tuple[str, int] | None = None + outbound_proxy: ( + tuple[ipaddress.IPv4Address | ipaddress.IPv6Address | str, int] | None + ) = None aor: str username: str | None = None password: str | None = None