Added Unix socket support for the Prometheus metrics server.
CI / Formatting (push) Successful in 5s
CI / Linting (push) Failing after 5s
CI / Tests (push) Successful in 8s
CI / Type Checking (push) Successful in 10s
CI / Spelling (push) Successful in 5s

This commit is contained in:
2026-06-11 15:07:44 -04:00
parent 7b395e1834
commit 446c0cd9ca
5 changed files with 255 additions and 45 deletions
+22 -11
View File
@@ -31,6 +31,7 @@ import uvloop
from crabstero import __version__ as crabstero_version from crabstero import __version__ as crabstero_version
from crabstero.bot import Crabstero from crabstero.bot import Crabstero
from crabstero.metrics import TcpMetricsAddress, UnixMetricsAddress
logger = logging.getLogger("crabstero") logger = logging.getLogger("crabstero")
@@ -145,8 +146,9 @@ def _parse_args(argv: list[str] | None = None) -> argparse.Namespace:
"--listen-metrics", "--listen-metrics",
default=os.environ.get("LISTEN_METRICS"), default=os.environ.get("LISTEN_METRICS"),
help=( help=(
"Enable Prometheus metrics endpoint on HOST:PORT" "Enable Prometheus metrics endpoint on HOST:PORT or unix:/path.sock"
" (e.g. 127.0.0.1:9090). Disabled by default." " (e.g. 127.0.0.1:9090 or unix:/run/crabstero.sock)."
" Disabled by default."
" (default: LISTEN_METRICS environment variable)." " (default: LISTEN_METRICS environment variable)."
), ),
) )
@@ -154,15 +156,24 @@ def _parse_args(argv: list[str] | None = None) -> argparse.Namespace:
args = parser.parse_args(argv) args = parser.parse_args(argv)
if args.listen_metrics is not None: if args.listen_metrics is not None:
host, sep, port_str = args.listen_metrics.rpartition(":") if args.listen_metrics.startswith("unix:"):
if not sep or not host: path = args.listen_metrics.removeprefix("unix:")
parser.error( if not path:
"--listen-metrics must be in HOST:PORT format (e.g. 127.0.0.1:9090)", parser.error("--listen-metrics Unix socket path must not be empty")
) args.listen_metrics = UnixMetricsAddress(path)
try: else:
args.listen_metrics = (host, int(port_str)) host, sep, port_str = args.listen_metrics.rpartition(":")
except ValueError: if not sep or not host:
parser.error(f"--listen-metrics port must be an integer, got '{port_str}'") parser.error(
"--listen-metrics must be in HOST:PORT or unix:/path.sock format "
"(e.g. 127.0.0.1:9090)",
)
try:
args.listen_metrics = TcpMetricsAddress(host, int(port_str))
except ValueError:
parser.error(
f"--listen-metrics port must be an integer, got '{port_str}'",
)
if args.token is None: if args.token is None:
parser.error( parser.error(
+4 -3
View File
@@ -36,6 +36,7 @@ from crabstero.metrics import (
ERRORS, ERRORS,
GUILD_COUNT, GUILD_COUNT,
INGESTION_ACTIVE, INGESTION_ACTIVE,
MetricsAddress,
MetricsServer, MetricsServer,
) )
from crabstero.tasks.ingestion import ingest_channel from crabstero.tasks.ingestion import ingest_channel
@@ -60,14 +61,14 @@ class Crabstero(commands.Bot):
database_path: str, database_path: str,
*, *,
ingest_only: bool = False, ingest_only: bool = False,
metrics_address: tuple[str, int] | None = None, metrics_address: MetricsAddress | None = None,
) -> None: ) -> None:
"""Configure intents, store configuration, and prepare ingestion state. """Configure intents, store configuration, and prepare ingestion state.
:param token: The Discord bot token. :param token: The Discord bot token.
:param database_path: The file path to the SQLite database. :param database_path: The file path to the SQLite database.
:param ingest_only: When True, the bot only ingests data and never responds. :param ingest_only: When True, the bot only ingests data and never responds.
:param metrics_address: Optional (host, port) for the Prometheus metrics server. :param metrics_address: Optional address for the Prometheus metrics server.
""" """
intents = discord.Intents.default() intents = discord.Intents.default()
intents.guilds = True intents.guilds = True
@@ -121,7 +122,7 @@ class Crabstero(commands.Bot):
self.ingest_cache.start() self.ingest_cache.start()
if self._metrics_address is not None: if self._metrics_address is not None:
server = MetricsServer(*self._metrics_address) server = MetricsServer(self._metrics_address)
await server.start() await server.start()
self._metrics_server = server self._metrics_server = server
+92 -13
View File
@@ -19,7 +19,13 @@ registry. The MetricsServer class wraps an aiohttp application that serves
the ``/metrics`` scrape endpoint. the ``/metrics`` scrape endpoint.
""" """
import contextlib
import errno
import logging import logging
import socket
import stat
from dataclasses import dataclass
from pathlib import Path
from aiohttp import web from aiohttp import web
from prometheus_client import Counter, Gauge, Histogram, Info from prometheus_client import Counter, Gauge, Histogram, Info
@@ -33,6 +39,25 @@ _FAST_BUCKETS = (0.001, 0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0,
_SLOW_BUCKETS = (0.1, 0.5, 1.0, 5.0, 10.0, 30.0, 60.0, 120.0, 300.0, 600.0) _SLOW_BUCKETS = (0.1, 0.5, 1.0, 5.0, 10.0, 30.0, 60.0, 120.0, 300.0, 600.0)
_MESSAGES_BUCKETS = (100, 500, 1000, 5000, 10000, 25000, 50000) _MESSAGES_BUCKETS = (100, 500, 1000, 5000, 10000, 25000, 50000)
@dataclass(frozen=True, slots=True)
class TcpMetricsAddress:
"""TCP listen address for the metrics server."""
host: str
port: int
@dataclass(frozen=True, slots=True)
class UnixMetricsAddress:
"""Unix socket listen address for the metrics server."""
path: str
type MetricsAddress = TcpMetricsAddress | UnixMetricsAddress
BUILD_INFO = Info("crabstero_build", "Build information") BUILD_INFO = Info("crabstero_build", "Build information")
BUILD_INFO.info({"version": __version__}) BUILD_INFO.info({"version": __version__})
@@ -119,21 +144,21 @@ CHANNEL_INGESTION_MESSAGES = Histogram(
class MetricsServer: class MetricsServer:
"""HTTP server that exposes a Prometheus ``/metrics`` scrape endpoint. """HTTP server that exposes a Prometheus ``/metrics`` scrape endpoint.
Uses ``aiohttp.web.AppRunner`` and ``TCPSite`` for async-native serving. Uses ``aiohttp.web.AppRunner`` with TCP or Unix socket sites for
async-native serving.
The handler is provided by ``prometheus_client.aiohttp.make_aiohttp_handler``, The handler is provided by ``prometheus_client.aiohttp.make_aiohttp_handler``,
which handles compression and content negotiation automatically. which handles compression and content negotiation automatically.
""" """
__slots__ = ("_host", "_port", "_runner") __slots__ = ("_address", "_port", "_runner")
def __init__(self, host: str, port: int) -> None: def __init__(self, address: MetricsAddress) -> None:
"""Store the listen address. """Store the listen address.
:param host: The hostname or IP to bind to. :param address: TCP or Unix socket address to bind to.
:param port: The TCP port to bind to. Use 0 for OS-assigned.
""" """
self._host = host self._address = address
self._port = port self._port = address.port if isinstance(address, TcpMetricsAddress) else None
self._runner: web.AppRunner | None = None self._runner: web.AppRunner | None = None
@property @property
@@ -143,6 +168,9 @@ class MetricsServer:
After ``start()``, this reflects the actual port (useful when After ``start()``, this reflects the actual port (useful when
the constructor received port 0 for OS assignment). the constructor received port 0 for OS assignment).
""" """
if self._port is None:
msg = "Unix socket metrics server does not have a TCP port"
raise RuntimeError(msg)
return self._port return self._port
async def start(self) -> None: async def start(self) -> None:
@@ -154,14 +182,65 @@ class MetricsServer:
app.router.add_get("/metrics", make_aiohttp_handler()) app.router.add_get("/metrics", make_aiohttp_handler())
self._runner = web.AppRunner(app) self._runner = web.AppRunner(app)
await self._runner.setup() await self._runner.setup()
site = web.TCPSite(self._runner, self._host, self._port) try:
await site.start() if isinstance(self._address, TcpMetricsAddress):
site: web.BaseSite = web.TCPSite(
self._runner,
self._address.host,
self._address.port,
)
await site.start()
# Resolve the actual bound port when the OS assigned one. # Resolve the actual bound port when the OS assigned one.
if self._port == 0: if self._address.port == 0:
self._port = self._runner.addresses[0][1] self._port = self._runner.addresses[0][1]
logger.info("Metrics server listening on %s:%d.", self._host, self._port) logger.info(
"Metrics server listening on %s:%d.",
self._address.host,
self.port,
)
else:
socket_path = Path(self._address.path)
try:
mode = socket_path.stat(follow_symlinks=False).st_mode
except FileNotFoundError:
pass
else:
if not stat.S_ISSOCK(mode):
msg = (
"Metrics Unix socket path already exists and is not a "
f"socket: {self._address.path}"
)
raise FileExistsError(msg)
with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as probe:
probe.settimeout(0.1)
try:
probe.connect(self._address.path)
except OSError as exc:
if exc.errno not in {errno.ECONNREFUSED, errno.ENOENT}:
raise
else:
msg = (
"Metrics Unix socket path is already in use: "
f"{self._address.path}"
)
raise OSError(errno.EADDRINUSE, msg)
with contextlib.suppress(FileNotFoundError):
socket_path.unlink()
site = web.UnixSite(self._runner, self._address.path)
await site.start()
logger.info(
"Metrics server listening on Unix socket %s.",
self._address.path,
)
except Exception:
await self._runner.cleanup()
self._runner = None
raise
async def stop(self) -> None: async def stop(self) -> None:
"""Shut down the HTTP server and release resources.""" """Shut down the HTTP server and release resources."""
+25 -3
View File
@@ -33,6 +33,7 @@ from crabstero.__main__ import (
_watchdog_interval, _watchdog_interval,
_watchdog_loop, _watchdog_loop,
) )
from crabstero.metrics import TcpMetricsAddress, UnixMetricsAddress
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import Callable, Coroutine from collections.abc import Callable, Coroutine
@@ -141,7 +142,7 @@ class TestParseArgs:
def test_listen_metrics_parses_address(self) -> None: def test_listen_metrics_parses_address(self) -> None:
"""--listen-metrics HOST:PORT sets listen_metrics to (host, port).""" """--listen-metrics HOST:PORT sets listen_metrics to (host, port)."""
args = _parse_args(["--token", "test", "--listen-metrics", "127.0.0.1:9090"]) args = _parse_args(["--token", "test", "--listen-metrics", "127.0.0.1:9090"])
assert args.listen_metrics == ("127.0.0.1", 9090) assert args.listen_metrics == TcpMetricsAddress("127.0.0.1", 9090)
def test_listen_metrics_default_none(self, monkeypatch: pytest.MonkeyPatch) -> None: def test_listen_metrics_default_none(self, monkeypatch: pytest.MonkeyPatch) -> None:
"""Omitting --listen-metrics defaults to None.""" """Omitting --listen-metrics defaults to None."""
@@ -152,13 +153,29 @@ class TestParseArgs:
def test_listen_metrics_ipv6(self) -> None: def test_listen_metrics_ipv6(self) -> None:
"""--listen-metrics [::1]:PORT parses IPv6 address correctly.""" """--listen-metrics [::1]:PORT parses IPv6 address correctly."""
args = _parse_args(["--token", "test", "--listen-metrics", "[::1]:9090"]) args = _parse_args(["--token", "test", "--listen-metrics", "[::1]:9090"])
assert args.listen_metrics == ("[::1]", 9090) assert args.listen_metrics == TcpMetricsAddress("[::1]", 9090)
def test_listen_metrics_from_env(self, monkeypatch: pytest.MonkeyPatch) -> None: def test_listen_metrics_from_env(self, monkeypatch: pytest.MonkeyPatch) -> None:
"""LISTEN_METRICS env var is used when --listen-metrics is omitted.""" """LISTEN_METRICS env var is used when --listen-metrics is omitted."""
monkeypatch.setenv("LISTEN_METRICS", "127.0.0.1:8080") monkeypatch.setenv("LISTEN_METRICS", "127.0.0.1:8080")
args = _parse_args(["--token", "test"]) args = _parse_args(["--token", "test"])
assert args.listen_metrics == ("127.0.0.1", 8080) assert args.listen_metrics == TcpMetricsAddress("127.0.0.1", 8080)
def test_listen_metrics_unix_socket(self) -> None:
"""--listen-metrics unix:/path sets listen_metrics to a Unix socket."""
args = _parse_args(
["--token", "test", "--listen-metrics", "unix:/run/crabstero.sock"],
)
assert args.listen_metrics == UnixMetricsAddress("/run/crabstero.sock")
def test_listen_metrics_unix_socket_from_env(
self,
monkeypatch: pytest.MonkeyPatch,
) -> None:
"""LISTEN_METRICS can enable metrics on a Unix socket."""
monkeypatch.setenv("LISTEN_METRICS", "unix:/run/crabstero.sock")
args = _parse_args(["--token", "test"])
assert args.listen_metrics == UnixMetricsAddress("/run/crabstero.sock")
def test_listen_metrics_invalid_format(self) -> None: def test_listen_metrics_invalid_format(self) -> None:
"""--listen-metrics with no colon raises SystemExit.""" """--listen-metrics with no colon raises SystemExit."""
@@ -170,6 +187,11 @@ class TestParseArgs:
with pytest.raises(SystemExit): with pytest.raises(SystemExit):
_parse_args(["--token", "test", "--listen-metrics", "127.0.0.1:abc"]) _parse_args(["--token", "test", "--listen-metrics", "127.0.0.1:abc"])
def test_listen_metrics_empty_unix_socket_path(self) -> None:
"""--listen-metrics unix: with no path raises SystemExit."""
with pytest.raises(SystemExit):
_parse_args(["--token", "test", "--listen-metrics", "unix:"])
def test_missing_token_exits(self, monkeypatch: pytest.MonkeyPatch) -> None: def test_missing_token_exits(self, monkeypatch: pytest.MonkeyPatch) -> None:
"""Missing token causes SystemExit.""" """Missing token causes SystemExit."""
monkeypatch.delenv("TOKEN", raising=False) monkeypatch.delenv("TOKEN", raising=False)
+112 -15
View File
@@ -17,17 +17,41 @@
Tests cover metric object registration and the MetricsServer HTTP endpoint. Tests cover metric object registration and the MetricsServer HTTP endpoint.
""" """
import errno
import socket
from dataclasses import dataclass
from http import HTTPStatus from http import HTTPStatus
from typing import TYPE_CHECKING from typing import TYPE_CHECKING, Literal, cast
import aiohttp import aiohttp
import pytest import pytest
from prometheus_client import generate_latest from prometheus_client import generate_latest
from crabstero.metrics import MetricsServer from crabstero.metrics import MetricsServer, TcpMetricsAddress, UnixMetricsAddress
if TYPE_CHECKING: if TYPE_CHECKING:
from collections.abc import AsyncGenerator from collections.abc import AsyncGenerator
from pathlib import Path
type MetricsTransport = Literal["tcp", "unix-socket"]
@dataclass(frozen=True, slots=True)
class MetricsEndpoint:
"""Client details for one running metrics transport."""
base_url: str
unix_socket_path: str | None = None
def client_session(self) -> aiohttp.ClientSession:
"""Create an aiohttp client session for this metrics transport."""
connector = (
aiohttp.UnixConnector(path=self.unix_socket_path)
if self.unix_socket_path is not None
else None
)
return aiohttp.ClientSession(connector=connector)
class TestMetricObjects: class TestMetricObjects:
@@ -75,23 +99,49 @@ class TestMetricObjects:
assert name in output, f"{name} not found in Prometheus output" assert name in output, f"{name} not found in Prometheus output"
@pytest.fixture @pytest.fixture(
async def metrics_server() -> AsyncGenerator[MetricsServer]: params=[
"""Start a MetricsServer on an OS-assigned port and stop it after the test.""" pytest.param("tcp", id="tcp"),
server = MetricsServer("127.0.0.1", 0) pytest.param("unix-socket", id="unix-socket"),
await server.start() ],
yield server )
await server.stop() async def metrics_endpoint(
request: pytest.FixtureRequest,
tmp_path: Path,
) -> AsyncGenerator[MetricsEndpoint]:
"""Start a MetricsServer on each supported transport."""
transport = cast("MetricsTransport", request.param)
match transport:
case "tcp":
server = MetricsServer(TcpMetricsAddress("127.0.0.1", 0))
await server.start()
endpoint = MetricsEndpoint(f"http://127.0.0.1:{server.port}")
case "unix-socket":
socket_path = tmp_path / "metrics.sock"
server = MetricsServer(UnixMetricsAddress(str(socket_path)))
await server.start()
endpoint = MetricsEndpoint(
"http://crabstero",
unix_socket_path=str(socket_path),
)
try:
yield endpoint
finally:
await server.stop()
class TestMetricsServer: class TestMetricsServer:
"""HTTP server serves Prometheus metrics on /metrics.""" """HTTP server serves Prometheus metrics on /metrics."""
async def test_serves_metrics_endpoint(self, metrics_server: MetricsServer) -> None: async def test_serves_metrics_endpoint(
self,
metrics_endpoint: MetricsEndpoint,
) -> None:
"""GET /metrics returns 200 with metric output containing our metrics.""" """GET /metrics returns 200 with metric output containing our metrics."""
async with ( async with (
aiohttp.ClientSession() as session, metrics_endpoint.client_session() as session,
session.get(f"http://127.0.0.1:{metrics_server.port}/metrics") as resp, session.get(f"{metrics_endpoint.base_url}/metrics") as resp,
): ):
assert resp.status == HTTPStatus.OK assert resp.status == HTTPStatus.OK
body = await resp.text() body = await resp.text()
@@ -99,11 +149,58 @@ class TestMetricsServer:
async def test_non_metrics_path_returns_404( async def test_non_metrics_path_returns_404(
self, self,
metrics_server: MetricsServer, metrics_endpoint: MetricsEndpoint,
) -> None: ) -> None:
"""GET on an unknown path returns 404.""" """GET on an unknown path returns 404."""
async with ( async with (
aiohttp.ClientSession() as session, metrics_endpoint.client_session() as session,
session.get(f"http://127.0.0.1:{metrics_server.port}/notfound") as resp, session.get(f"{metrics_endpoint.base_url}/notfound") as resp,
): ):
assert resp.status == HTTPStatus.NOT_FOUND assert resp.status == HTTPStatus.NOT_FOUND
async def test_stale_unix_socket_path_is_recovered(
self,
tmp_path: Path,
) -> None:
"""Startup removes an abandoned Unix socket file left by a crash."""
socket_path = tmp_path / "metrics.sock"
with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as stale_socket:
stale_socket.bind(str(socket_path))
server = MetricsServer(UnixMetricsAddress(str(socket_path)))
await server.start()
try:
connector = aiohttp.UnixConnector(path=str(socket_path))
async with (
aiohttp.ClientSession(connector=connector) as session,
session.get("http://crabstero/metrics") as resp,
):
assert resp.status == HTTPStatus.OK
body = await resp.text()
assert "crabstero_build_info" in body
finally:
await server.stop()
async def test_active_unix_socket_path_fails(
self,
tmp_path: Path,
) -> None:
"""Startup refuses to replace a Unix socket path that is still in use."""
socket_path = tmp_path / "metrics.sock"
with socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) as active_socket:
active_socket.bind(str(socket_path))
active_socket.listen(1)
server = MetricsServer(UnixMetricsAddress(str(socket_path)))
with pytest.raises(OSError, match="already in use") as exc_info:
await server.start()
assert exc_info.value.errno == errno.EADDRINUSE
async def test_non_socket_unix_path_fails(self, tmp_path: Path) -> None:
"""Startup refuses to replace a non-socket path."""
socket_path = tmp_path / "metrics.sock"
socket_path.write_text("")
server = MetricsServer(UnixMetricsAddress(str(socket_path)))
with pytest.raises(FileExistsError):
await server.start()