Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
29 changes: 29 additions & 0 deletions tests/test_main.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
from __future__ import annotations

import asyncio
import ipaddress
import sys
from unittest.mock import MagicMock, patch

Expand Down Expand Up @@ -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."""
Expand Down
61 changes: 28 additions & 33 deletions voip/__main__.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
import dataclasses
import ipaddress
import logging
import re
import ssl
import time

Expand All @@ -29,13 +30,21 @@
SIP_TLS_PORT = 5061


HOSTPORT_PATTERN: re.Pattern[str] = re.compile(
r"^(?:\[(?P<ipv6>[0-9a-fA-F:]+)\]|(?P<host>[^:\[\]]+))"
r"(?::(?P<port>\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.
Expand All @@ -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:
Expand All @@ -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):
Expand Down Expand Up @@ -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.
Expand All @@ -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:
Expand All @@ -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,
)
Expand Down
4 changes: 3 additions & 1 deletion voip/sip/protocol.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down