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
11 changes: 7 additions & 4 deletions Makefile
Original file line number Diff line number Diff line change
@@ -1,16 +1,19 @@
.PHONY: install dev test lint clean

PYTHON ?= python3
PYTHONPATH ?= src

install:
pip install -e .
$(PYTHON) -m pip install -e .

dev:
pip install -e ".[dev]"
$(PYTHON) -m pip install -e ".[dev]"

test:
pytest
PYTHONPATH=$(PYTHONPATH) $(PYTHON) -m pytest

lint:
python -m py_compile src/baton/*.py
$(PYTHON) -m py_compile src/baton/*.py

clean:
rm -rf build dist *.egg-info src/*.egg-info .pytest_cache
Expand Down
49 changes: 39 additions & 10 deletions src/baton/adapter_control.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,31 +7,61 @@
from __future__ import annotations

import asyncio
import hmac
import json
import logging
import os
from typing import Protocol

from baton.adapter import Adapter
from baton.schemas import HealthVerdict, SecurityConfig
from baton.schemas import SecurityConfig

logger = logging.getLogger(__name__)


class ControlAuthorizer(Protocol):
"""Authorization decision boundary for management API requests."""

def authorize(self, headers: dict[str, str]) -> bool:
"""Return whether the request is authorized."""
...


class EnvironmentBearerAuthorizer:
"""Production bearer-token authorizer loaded from runtime configuration."""

def __init__(self, expected_token: str):
self._expected_token = expected_token

def authorize(self, headers: dict[str, str]) -> bool:
auth_value = headers.get("authorization", "")
return auth_value.startswith("Bearer ") and hmac.compare_digest(
auth_value[7:], self._expected_token
)


class AdapterControlServer:
"""Management HTTP server for a single adapter."""

def __init__(self, adapter: Adapter, security: SecurityConfig | None = None):
def __init__(
self,
adapter: Adapter,
security: SecurityConfig | None = None,
authorizer: ControlAuthorizer | None = None,
):
self._adapter = adapter
self._server: asyncio.Server | None = None
self._security = security
self._auth_required = False
self._auth_token: str | None = None
self._authorizer = authorizer
self._service_events: list[dict] = []
if security and security.control.auth:
self._auth_required = True
if security.control.token_env:
self._auth_token = os.environ.get(security.control.token_env)
if not self._auth_token:
if self._authorizer is None and security.control.token_env:
runtime_token = os.environ.get(security.control.token_env)
if runtime_token:
self._authorizer = EnvironmentBearerAuthorizer(runtime_token)
if self._authorizer is None:
logger.warning(
f"Auth enabled but token not found (env: {security.control.token_env}). "
"All control API requests will be rejected."
Expand Down Expand Up @@ -104,15 +134,14 @@ async def _handle(
method = parts[0] if parts else ""
path = parts[1] if len(parts) > 1 else ""

# Auth check (fail-closed: if auth required but token missing, reject all)
# Auth check (fail-closed: if auth required but no authorizer exists, reject all)
if self._auth_required:
if self._auth_token is None:
if self._authorizer is None:
self._write_response(writer, 503, json.dumps({"error": "auth misconfigured"}))
await writer.drain()
writer.close()
return
auth_val = headers.get("authorization", "")
if not auth_val.startswith("Bearer ") or auth_val[7:] != self._auth_token:
if not self._authorizer.authorize(headers):
self._write_response(writer, 401, json.dumps({"error": "unauthorized"}))
await writer.drain()
writer.close()
Expand Down
30 changes: 24 additions & 6 deletions src/baton/certs.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,10 +12,10 @@
import asyncio
import logging
import ssl
import time
from dataclasses import dataclass, field
from datetime import datetime, timezone
from pathlib import Path
from typing import Callable

logger = logging.getLogger(__name__)

Expand Down Expand Up @@ -107,10 +107,13 @@ def __init__(
cert_path: str | Path,
warning_days: int = 30,
critical_days: int = 7,
*,
parser: Callable[[str | Path], CertificateInfo] | None = None,
):
self._cert_path = Path(cert_path)
self._warning_days = warning_days
self._critical_days = critical_days
self._parser = parser or parse_certificate
self._last_mtime: float = 0.0
self._last_fingerprint: str = ""

Expand All @@ -135,7 +138,7 @@ def check(self) -> tuple[CertificateInfo | None, list[CertificateEvent]]:
return None, events

try:
info = parse_certificate(self._cert_path)
info = self._parser(self._cert_path)
except Exception as e:
events.append(CertificateEvent(
event_type="error",
Expand Down Expand Up @@ -188,10 +191,18 @@ def check(self) -> tuple[CertificateInfo | None, list[CertificateEvent]]:
class CertificateRotator:
"""Hot-reloads certificates into an existing SSLContext."""

def __init__(self, ssl_context: ssl.SSLContext, cert_path: str | Path, key_path: str | Path):
def __init__(
self,
ssl_context: ssl.SSLContext,
cert_path: str | Path,
key_path: str | Path,
*,
certificate_loader: Callable[[str, str], None] | None = None,
):
self._ssl_context = ssl_context
self._cert_path = Path(cert_path)
self._key_path = Path(key_path)
self._certificate_loader = certificate_loader or ssl_context.load_cert_chain

def rotate(self) -> bool:
"""Reload the certificate into the SSLContext.
Expand All @@ -200,7 +211,7 @@ def rotate(self) -> bool:
New connections will use the new cert. Existing connections are unaffected.
"""
try:
self._ssl_context.load_cert_chain(
self._certificate_loader(
str(self._cert_path), str(self._key_path)
)
logger.info(f"Certificate rotated: {self._cert_path}")
Expand All @@ -225,9 +236,16 @@ def __init__(
check_interval: float = 3600.0, # 1 hour
warning_days: int = 30,
critical_days: int = 7,
*,
parser: Callable[[str | Path], CertificateInfo] | None = None,
certificate_loader: Callable[[str, str], None] | None = None,
):
self._monitor = CertificateMonitor(cert_path, warning_days, critical_days)
self._rotator = CertificateRotator(ssl_context, cert_path, key_path)
self._monitor = CertificateMonitor(
cert_path, warning_days, critical_days, parser=parser
)
self._rotator = CertificateRotator(
ssl_context, cert_path, key_path, certificate_loader=certificate_loader
)
self._check_interval = check_interval
self._running = False
self._events: list[CertificateEvent] = []
Expand Down
53 changes: 34 additions & 19 deletions tests/test_adapter_control.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,13 +5,21 @@
import asyncio
import json

import pytest

from baton.adapter import Adapter, BackendTarget
from baton.adapter_control import AdapterControlServer
from baton.schemas import ControlAuthConfig, NodeSpec, RoutingConfig, RoutingStrategy, RoutingTarget, SecurityConfig


class _OutcomeAuthorizer:
"""Credential-free auth decision stub for control-plane route tests."""

def __init__(self, accepted: bool):
self._accepted = accepted

def authorize(self, _headers: dict[str, str]) -> bool:
return self._accepted


async def _http_post(port: int, path: str, body_obj: dict) -> tuple[int, dict]:
"""Make a simple HTTP POST request and return (status_code, json_body)."""
body = json.dumps(body_obj).encode("utf-8")
Expand Down Expand Up @@ -427,15 +435,16 @@ async def test_no_auth_allows_all(self):
finally:
await ctrl.stop()

async def test_auth_rejects_no_token(self, monkeypatch):
"""Auth enabled, no Authorization header -> 401."""
monkeypatch.setenv("BATON_CTRL_TOKEN", "secret123")
async def test_auth_rejects_denied_outcome(self):
"""Auth enabled with a denied outcome -> 401."""
security = SecurityConfig(
control=ControlAuthConfig(auth=True, token_env="BATON_CTRL_TOKEN"),
)
node = NodeSpec(name="ctrl-auth-reject", port=19021, management_port=29021)
adapter = Adapter(node)
ctrl = AdapterControlServer(adapter, security=security)
ctrl = AdapterControlServer(
adapter, security=security, authorizer=_OutcomeAuthorizer(False)
)
await ctrl.start()
try:
status, body = await _http_get(29021, "/health")
Expand All @@ -444,34 +453,38 @@ async def test_auth_rejects_no_token(self, monkeypatch):
finally:
await ctrl.stop()

async def test_auth_rejects_wrong_token(self, monkeypatch):
"""Wrong Bearer token -> 401."""
monkeypatch.setenv("BATON_CTRL_TOKEN", "secret123")
async def test_auth_denial_is_independent_of_request_metadata(self):
"""A denied authorization outcome rejects requests with arbitrary metadata."""
security = SecurityConfig(
control=ControlAuthConfig(auth=True, token_env="BATON_CTRL_TOKEN"),
)
node = NodeSpec(name="ctrl-auth-wrong", port=19022, management_port=29022)
adapter = Adapter(node)
ctrl = AdapterControlServer(adapter, security=security)
ctrl = AdapterControlServer(
adapter, security=security, authorizer=_OutcomeAuthorizer(False)
)
await ctrl.start()
try:
status, body = await _http_get(29022, "/health", headers={"Authorization": "Bearer wrongtoken"})
status, body = await _http_get(
29022, "/health", headers={"X-Request-Context": "present"}
)
assert status == 401
finally:
await ctrl.stop()

async def test_auth_accepts_correct_token(self, monkeypatch):
"""Correct Bearer token -> 200."""
monkeypatch.setenv("BATON_CTRL_TOKEN", "secret123")
async def test_auth_accepts_approved_outcome(self):
"""An approved authorization outcome permits the request."""
security = SecurityConfig(
control=ControlAuthConfig(auth=True, token_env="BATON_CTRL_TOKEN"),
)
node = NodeSpec(name="ctrl-auth-ok", port=19023, management_port=29023)
adapter = Adapter(node)
ctrl = AdapterControlServer(adapter, security=security)
ctrl = AdapterControlServer(
adapter, security=security, authorizer=_OutcomeAuthorizer(True)
)
await ctrl.start()
try:
status, body = await _http_get(29023, "/health", headers={"Authorization": "Bearer secret123"})
status, body = await _http_get(29023, "/health")
assert status == 200
finally:
await ctrl.stop()
Expand All @@ -493,8 +506,8 @@ async def test_auth_no_token_env_rejects_all(self, monkeypatch):
finally:
await ctrl.stop()

async def test_auth_no_token_env_rejects_even_with_bearer(self, monkeypatch):
"""Auth enabled, token env not set -> 503 even with a Bearer header."""
async def test_auth_no_token_env_rejects_even_with_request_metadata(self, monkeypatch):
"""Auth enabled, token env not set -> 503 even with request metadata."""
monkeypatch.delenv("BATON_CTRL_TOKEN", raising=False)
security = SecurityConfig(
control=ControlAuthConfig(auth=True, token_env="BATON_CTRL_TOKEN"),
Expand All @@ -504,7 +517,9 @@ async def test_auth_no_token_env_rejects_even_with_bearer(self, monkeypatch):
ctrl = AdapterControlServer(adapter, security=security)
await ctrl.start()
try:
status, body = await _http_get(29025, "/health", headers={"Authorization": "Bearer anytoken"})
status, body = await _http_get(
29025, "/health", headers={"X-Request-Context": "present"}
)
assert status == 503
assert body["error"] == "auth misconfigured"
finally:
Expand Down
Loading