Provide a private Unix listener for owner-metered Messages
All checks were successful
CI Smoke / host-smoke (push) Successful in 0s
CI Smoke / container-smoke (push) Successful in 1s

Assistant: codex
Assistant-Model: gpt-6-astra
Assistant-Session: 01a07ff8-19d0-7820-b4d0-1353833cb7fc
This commit is contained in:
tegwick 2026-09-09 22:20:55 +02:00
parent dc77f434a8
commit 718e6730e4
4 changed files with 143 additions and 4 deletions

View file

@ -10,11 +10,16 @@ from __future__ import annotations
import hashlib
import http.client
import json
import os
import re
import socket
import stat
import threading
import time
from dataclasses import asdict, dataclass
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from socketserver import ThreadingUnixStreamServer
from typing import Any, NoReturn, Protocol, TypeGuard, cast
from urllib.parse import urlsplit
@ -345,6 +350,34 @@ class _OwnerHTTPServer(ThreadingHTTPServer):
owner: MessagesServer
class _OwnerUnixServer(ThreadingUnixStreamServer):
owner: MessagesServer
daemon_threads = True
def server_bind(self) -> None:
self._slots = threading.BoundedSemaphore(16)
super().server_bind()
def process_request(self, request: socket.socket | tuple[bytes, socket.socket], client_address: Any) -> None:
if not isinstance(request, socket.socket):
raise TypeError("Unix stream socket required")
if not self._slots.acquire(blocking=False):
request.close()
return
request.settimeout(10)
try:
super().process_request(request, client_address)
except BaseException:
self._slots.release()
raise
def process_request_thread(self, request: socket.socket | tuple[bytes, socket.socket], client_address: Any) -> None:
try:
super().process_request_thread(request, client_address)
finally:
self._slots.release()
class _Handler(BaseHTTPRequestHandler):
def log_message(self, format: str, *args: Any) -> None:
pass
@ -360,7 +393,7 @@ class _Handler(BaseHTTPRequestHandler):
self.wfile.write(raw)
def do_POST(self) -> None:
owner = cast(_OwnerHTTPServer, self.server).owner
owner = cast(_OwnerHTTPServer | _OwnerUnixServer, self.server).owner
upstream = None
sent = False
try:
@ -496,6 +529,7 @@ class MessagesServer:
host: str = "127.0.0.1",
port: int = 0,
allow_test_http: bool = False,
unix_path: Path | None = None,
) -> None:
endpoint = urlsplit(upstream_url)
if endpoint.scheme != "https" and not (
@ -521,13 +555,36 @@ class MessagesServer:
raise RequestRefused("explicit provider credential required")
self.policy, self.meter, self.endpoint = policy, meter, endpoint
self._provider_key = provider_key
self._httpd = _OwnerHTTPServer((host, port), _Handler)
self._unix_path = unix_path
self._httpd: _OwnerHTTPServer | _OwnerUnixServer
if unix_path is not None:
if host != "127.0.0.1" or port != 0:
raise RequestRefused("Unix route cannot also select a TCP listener")
parent = unix_path.parent
metadata = parent.lstat()
if (
not unix_path.is_absolute()
or parent.resolve() != parent
or not stat.S_ISDIR(metadata.st_mode)
or metadata.st_uid != os.getuid()
or stat.S_IMODE(metadata.st_mode) != 0o700
):
raise RequestRefused("Unix route requires a private owner directory")
# bind refuses existing paths, including dangling symlinks. Never unlink
# someone else's socket to recover a crashed or duplicated owner.
self._httpd = _OwnerUnixServer(str(unix_path), _Handler)
unix_path.chmod(0o600)
self._unix_inode = unix_path.stat().st_ino
else:
self._httpd = _OwnerHTTPServer((host, port), _Handler)
self._httpd.owner = self
self._thread: threading.Thread | None = None
@property
def port(self) -> int:
return int(self._httpd.server_address[1])
if self._unix_path is not None:
raise RequestRefused("Unix route has no TCP port")
return int(cast(_OwnerHTTPServer, self._httpd).server_address[1])
def start(self) -> None:
self._thread = threading.Thread(target=self._httpd.serve_forever, daemon=True)
@ -538,3 +595,9 @@ class MessagesServer:
self._httpd.shutdown()
self._thread.join()
self._httpd.server_close()
if self._unix_path is not None:
try:
if self._unix_path.lstat().st_ino == self._unix_inode:
self._unix_path.unlink()
except FileNotFoundError:
pass