Provide a private Unix listener for owner-metered Messages
Assistant: codex Assistant-Model: gpt-6-astra Assistant-Session: 01a07ff8-19d0-7820-b4d0-1353833cb7fc
This commit is contained in:
parent
dc77f434a8
commit
718e6730e4
4 changed files with 143 additions and 4 deletions
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue