feat: add NetBox plugin store
CI / php-store (push) Canceled after 0s
CI / python-components (push) Canceled after 0s

This commit is contained in:
2026-08-24 20:51:25 +02:00
commit f36d6be511
135 changed files with 15160 additions and 0 deletions
+6
View File
@@ -0,0 +1,6 @@
__pycache__/
*.py[cod]
*.egg-info/
.buildcheck/
build/
dist/
+176
View File
@@ -0,0 +1,176 @@
# NetBox Store host agent (MVP)
This package is the deliberately small privileged boundary between the NetBox Store plugin and a
NetBox 4.6.5–4.6.8 Linux host. It does not import the Store implementation or assume Django: every
decision is revalidated against the configured JSON API immediately before a lifecycle action.
The daemon is **dry-run by default**. It serializes operations, persists idempotency and state in
SQLite, accepts one bounded JSON line per Unix-stream connection, verifies Linux peer credentials,
downloads only an approved immutable wheel, and invokes subprocesses only as fixed argument arrays
with `shell=False`.
## Install and one-time operator setup
Build/install into a dedicated administrative environment, copy
[`examples/agent.toml`](examples/agent.toml) to `/etc/netbox-store-agent/agent.toml`, replace all
site-specific paths, hosts, UIDs/GIDs and versions, then make the file root-owned and mode `0600`.
The optional bearer-token file must also be root-owned `0600`. Store credentials are sent only to
the exact scheme/host/port origin of `store.base_url`, never to a separately allow-listed artifact
host.
`allow_private_addresses = false` is intentionally fail-closed. If the Store or an approved
artifact host consciously runs on RFC1918/ULA infrastructure, set it to `true` only after pinning
every expected hostname in `allowed_hosts` and deploying correctly verified TLS (including the
private CA via `ca_file` when needed). The host allow-list and exact-origin credential rule remain
active; this switch only permits private/reserved DNS results.
The agent owns only these two configured files:
- `paths.include_path`, which contains only `STORE_PLUGINS = [...]`;
- `paths.requirements_path`, which contains the locked direct wheel references.
In the operator-owned NetBox `configuration.py`, add once, after the normal `PLUGINS` declaration:
```python
from store_plugins import STORE_PLUGINS
PLUGINS += STORE_PLUGINS
```
Place `store_plugins.py` where that import resolves in your deployment. Optionally add a one-time
`-r /opt/netbox/local_requirements_store.txt` line to the operator-owned
`/opt/netbox/local_requirements.txt` for reproducible maintenance installs. The agent never edits
either operator-owned file.
Install the example systemd units, review their `ReadWritePaths`, `SocketGroup`, and `ExecStart`, then:
```text
systemctl daemon-reload
systemctl enable --now netbox-store-agent.socket
netbox-store-agent capabilities
```
The canonical socket path used by the unit, example config, CLI default, and NetBox Store plugin is
`/run/netbox-store-agent/agent.sock`.
With peer checks enabled (the production default), Linux `SO_PEERCRED` must be available and either
the peer's effective UID must appear in `allowed_peer_uids` or its effective GID in
`allowed_peer_gids`; Unix-socket filesystem permissions still apply. Prefer allow-listing the exact
NetBox service UID, especially when socket access is granted through a supplementary group.
Keep `dry_run = true` through catalog and lifecycle acceptance tests. Enabling real mutation is an
explicit operator configuration change.
## Store JSON contract
Both endpoint templates may be configured with or without a trailing slash. A 404 is retried once
using the alternate form. Path parameters are URL-quoted. The agent accepts exactly these fields.
`GET /api/v1/plugins/{slug}/`:
```json
{
"api_version": "v1",
"slug": "netbox-example",
"name": "Example",
"summary": "Example plugin",
"description": "Description",
"repository_url": "https://git.example/repo",
"latest_version": "1.2.3",
"package_name": "netbox-example",
"import_name": "netbox_example",
"min_netbox_version": "4.6.5",
"max_netbox_version": "4.6.8",
"approved": true,
"status": "approved",
"releases": []
}
```
`GET /api/v1/plugins/{slug}/releases/{version}/` returns a release object directly:
```json
{
"version": "1.2.3",
"download_url": "https://store.example/artifacts/netbox_example-1.2.3-py3-none-any.whl",
"sha256": "0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef",
"artifact_size": 12345,
"commit_sha": "",
"min_netbox_version": "4.6.5",
"max_netbox_version": "4.6.8",
"published_at": "2026-08-24T12:00:00Z",
"approved": true,
"status": "approved",
"immutable": true,
"approved_payload_sha256": "abcdefabcdefabcdefabcdefabcdefabcdefabcdefabcdefabcdefabcdefabcd"
}
```
If release detail returns 404, exactly one matching object in plugin `releases[]` is accepted. The
historic field `approved_payload_sha256` is protocol-v1 naming: clients treat its lowercase 64-hex
value as an opaque Store marker, and the agent compares it exactly with the freshly fetched value.
Artifact `sha256`, in contrast, is always a lowercase 64-character SHA-256.
The plugin and release must be approved, the release immutable and compatible with the configured
NetBox version. The artifact must be a valid wheel whose filename distribution and version match
the catalog, and its byte count, digest, host, scheme and DNS addresses are checked while streaming.
Redirects, source distributions, private/reserved DNS targets (unless explicitly enabled for a test
environment), and catalog additions outside the v1 schema fail closed.
## Client protocol
Transport is Unix `SOCK_STREAM`, UTF-8 JSON-lines, exactly one request and one response per
connection, maximum 64 KiB. Response shape is always
`{"protocol_version":1,"status":...,"body":{...}}`.
```json
{"protocol_version":1,"method":"GET","path":"/v1/capabilities"}
```
```json
{"protocol_version":1,"method":"POST","path":"/v1/operations","idempotency_key":"20f4274f-d4e5-42bf-9164-967b1a774481","body":{"request_id":"eea17d87-8944-4ee2-a076-363338ab746d","action":"install","plugin_slug":"netbox-example","version":"1.2.3","approved_payload_sha256":"abcdefabcdefabcdefabcdefabcdefabcdefabcdefabcdefabcdefabcdefabcd","requested_by":"netbox:alice"}}
```
```json
{"protocol_version":1,"method":"GET","path":"/v1/operations/20f4274f-d4e5-42bf-9164-967b1a774481"}
```
`POST` returns 202. Repeating the same idempotency UUID with the identical body returns the existing
operation (`created:false`); a different body returns 409. States are `queued`, `running`,
`dry_run`, `succeeded`, `failed`, or `manual_recovery`.
CLI examples:
```text
netbox-store-agent submit install netbox-example --version 1.2.3 \
--approved-payload-sha256 abcdefabcdefabcdefabcdefabcdefabcdefabcdefabcdefabcdefabcdefabcd
netbox-store-agent status 20f4274f-d4e5-42bf-9164-967b1a774481
```
## Lifecycle and fail-closed boundaries
- `install` and `update` re-fetch plugin and release, exactly match the approval token, download and
validate the wheel, then use `pip --no-index --no-deps --only-binary=:all: --require-hashes`.
- New installs remain disabled. Updating an enabled plugin runs NetBox `migrate`, `collectstatic`,
and restarts every configured service (`netbox` and `netbox-rq` in the example).
- `enable` revalidates its installed release, writes the include, migrates, collects static files,
and restarts both services. `disable` writes the include and restarts both. `uninstall` is allowed
only after disable and revalidates current approved plugin identity first.
- Self-management slugs are denied for every action. Package/import names and all executable paths,
service units, managed paths, Store origins and NetBox compatibility are configuration/catalog
policy—not caller-controlled command fragments.
- A root-owned global file lock prevents simultaneous host mutations. Interrupted `running`
operations become `manual_recovery` at startup. Managed files are backed up per operation and
atomically replaced; best-effort restoration does not claim package/service rollback.
This MVP intentionally has no TUF metadata, dependency-wheel set, transactional virtualenv switch,
or reliable rollback for package installation/database migrations/service restarts. `--no-deps`
means an approved plugin's dependencies must already be provisioned by the operator/base image.
Any failure after host mutation is marked `manual_recovery`; inspect the journal, backup directory,
installed distributions, migrations, both services and both managed files before retrying. If the
Store is unavailable or an entry is no longer approved, lifecycle calls fail closed; emergency
manual recovery remains an operator procedure outside this API.
The security boundary still depends on root ownership and permissions of the daemon executable,
configuration, token, journal/state directories, socket and NetBox paths; TLS/CA integrity; Store
approval operations; artifact build provenance; and a correctly restricted NetBox service account.
+52
View File
@@ -0,0 +1,52 @@
[agent]
socket_path = "/run/netbox-store-agent/agent.sock"
journal_path = "/var/lib/netbox-store-agent/journal.sqlite3"
lock_path = "/run/lock/netbox-store-agent.lifecycle.lock"
backup_dir = "/var/lib/netbox-store-agent/backups"
dry_run = true
require_root = true
require_peer_credentials = true
# Replace/add the numeric UID or effective GID used by the NetBox service.
allowed_peer_uids = [0]
allowed_peer_gids = [0]
socket_mode = 0o660
socket_uid = 0
socket_gid = 0
max_request_bytes = 65536
connection_timeout_seconds = 5
worker_threads = 1
[store]
base_url = "https://store.example.invalid"
plugin_endpoint_template = "/api/v1/plugins/{plugin_slug}/"
release_endpoint_template = "/api/v1/plugins/{plugin_slug}/releases/{version}/"
timeout_seconds = 10
max_catalog_bytes = 1048576
max_artifact_bytes = 268435456
allow_private_addresses = false
allow_http_for_testing = false
# Add an artifact origin only when releases intentionally use that origin.
allowed_hosts = ["store.example.invalid"]
# bearer_token_file = "/etc/netbox-store-agent/store.token"
# ca_file = "/etc/ssl/certs/internal-store-ca.pem"
[paths]
allowed_root = "/opt/netbox"
include_path = "/opt/netbox/netbox/netbox/store_plugins.py"
requirements_path = "/opt/netbox/local_requirements_store.txt"
temp_dir = "/var/lib/netbox-store-agent/tmp"
[commands]
python_path = "/opt/netbox/venv/bin/python"
manage_path = "/opt/netbox/netbox/manage.py"
systemctl_path = "/usr/bin/systemctl"
services = ["netbox", "netbox-rq"]
command_timeout_seconds = 900
[policy]
netbox_version = "4.6.8"
min_supported_netbox = "4.6.5"
max_supported_netbox = "4.6.8"
self_plugin_slugs = ["netbox-store", "netbox-plugin-store", "netbox_plugin_store"]
allow_prereleases = false
require_release_for_enable = true
+42
View File
@@ -0,0 +1,42 @@
[build-system]
requires = ["setuptools>=75"]
build-backend = "setuptools.build_meta"
[project]
name = "mrblake-netbox-store-agent"
version = "0.1.0"
description = "Fail-closed host agent for curated NetBox plugin lifecycle operations"
readme = "README.md"
requires-python = ">=3.11"
license = {text = "MIT"}
authors = [{name = "MrBlake"}]
dependencies = ["packaging>=24,<27"]
classifiers = [
"Development Status :: 3 - Alpha",
"Environment :: No Input/Output (Daemon)",
"Operating System :: POSIX :: Linux",
"Programming Language :: Python :: 3 :: Only",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Topic :: System :: Systems Administration",
]
[project.optional-dependencies]
dev = ["pytest>=8,<9", "ruff>=0.9,<1"]
[project.scripts]
netbox-store-agent = "netbox_store_agent.cli:main"
[tool.setuptools]
package-dir = {"" = "src"}
[tool.setuptools.packages.find]
where = ["src"]
[tool.ruff]
line-length = 100
target-version = "py311"
[tool.ruff.lint]
select = ["E", "F", "I", "UP", "B", "S"]
ignore = ["S101"]
@@ -0,0 +1,5 @@
"""NetBox Store privileged host agent."""
from .constants import AGENT_VERSION, PROTOCOL_VERSION
__all__ = ["AGENT_VERSION", "PROTOCOL_VERSION"]
@@ -0,0 +1,3 @@
from .cli import main
raise SystemExit(main())
@@ -0,0 +1,497 @@
from __future__ import annotations
import hashlib
import http.client
import ipaddress
import json
import os
import socket
import ssl
import stat
import zipfile
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Protocol
from urllib.parse import quote, urlsplit
from packaging.utils import canonicalize_name, parse_wheel_filename
from packaging.version import InvalidVersion, Version
from .config import Config, StoreSettings
from .errors import CatalogError
from .util import DIST_NAME_RE, IMPORT_NAME_RE, SHA256_RE, SLUG_RE
class HttpStatusError(CatalogError):
def __init__(self, status: int, message: str):
super().__init__(message)
self.http_status = status
class Transport(Protocol):
def get_json(self, url: str, max_bytes: int) -> Any: ...
def download(
self,
url: str,
destination: Path,
*,
max_bytes: int,
expected_size: int,
expected_sha256: str,
) -> None: ...
class _PinnedHTTPConnection(http.client.HTTPConnection):
def __init__(self, host: str, port: int, address: str, timeout: float):
super().__init__(host, port=port, timeout=timeout)
self._address = address
def connect(self) -> None:
self.sock = socket.create_connection((self._address, self.port), self.timeout)
class _PinnedHTTPSConnection(http.client.HTTPSConnection):
def __init__(
self, host: str, port: int, address: str, timeout: float, context: ssl.SSLContext
):
super().__init__(host, port=port, timeout=timeout, context=context)
self._address = address
def connect(self) -> None:
raw_socket = socket.create_connection((self._address, self.port), self.timeout)
try:
self.sock = self._context.wrap_socket(raw_socket, server_hostname=self.host)
except Exception:
raw_socket.close()
raise
class SecureHTTPTransport:
"""Small no-redirect HTTP transport with a DNS-pinned connection."""
def __init__(self, settings: StoreSettings):
self.settings = settings
self._ssl_context = ssl.create_default_context(cafile=str(settings.ca_file) if settings.ca_file else None)
def _token(self) -> str | None:
path = self.settings.bearer_token_file
if path is None:
return None
try:
metadata = path.lstat()
except OSError as exc:
raise CatalogError("bearer token file cannot be read") from exc
if stat.S_ISLNK(metadata.st_mode) or not stat.S_ISREG(metadata.st_mode):
raise CatalogError("bearer token file must be a regular file, not a symlink")
if metadata.st_mode & (stat.S_IRWXG | stat.S_IRWXO):
raise CatalogError("bearer token file must not be accessible by group or other")
if hasattr(metadata, "st_uid") and metadata.st_uid != 0:
raise CatalogError("bearer token file must be owned by root")
try:
token = path.read_text(encoding="utf-8").strip()
except (OSError, UnicodeDecodeError) as exc:
raise CatalogError("bearer token file cannot be read as UTF-8") from exc
if not 1 <= len(token) <= 4096 or any(ord(char) < 33 or ord(char) > 126 for char in token):
raise CatalogError("bearer token file contains an invalid token")
return token
@staticmethod
def _is_public(address: str) -> bool:
ip = ipaddress.ip_address(address)
return not (
ip.is_private
or ip.is_loopback
or ip.is_link_local
or ip.is_multicast
or ip.is_reserved
or ip.is_unspecified
)
def _connection(self, url: str) -> tuple[http.client.HTTPConnection, str]:
parsed = urlsplit(url)
schemes = {"https", "http"} if self.settings.allow_http_for_testing else {"https"}
if parsed.scheme not in schemes or not parsed.hostname:
raise CatalogError("Store returned a URL with a forbidden scheme or missing host")
if parsed.username or parsed.password or parsed.fragment:
raise CatalogError("Store URL may not contain credentials or a fragment")
host = parsed.hostname.lower()
if host not in self.settings.allowed_hosts:
raise CatalogError("Store URL hostname is not allow-listed")
port = parsed.port or (443 if parsed.scheme == "https" else 80)
try:
records = socket.getaddrinfo(host, port, type=socket.SOCK_STREAM)
except OSError as exc:
raise CatalogError("Store hostname could not be resolved") from exc
addresses = list(dict.fromkeys(record[4][0] for record in records))
if not addresses:
raise CatalogError("Store hostname did not resolve to an address")
if not self.settings.allow_private_addresses and any(
not self._is_public(address) for address in addresses
):
raise CatalogError("Store hostname resolves to a non-public address")
address = addresses[0]
if parsed.scheme == "https":
connection: http.client.HTTPConnection = _PinnedHTTPSConnection(
host, port, address, self.settings.timeout_seconds, self._ssl_context
)
else:
connection = _PinnedHTTPConnection(host, port, address, self.settings.timeout_seconds)
target = parsed.path or "/"
if parsed.query:
target += "?" + parsed.query
return connection, target
def _headers(self, url: str) -> dict[str, str]:
headers = {"Accept": "application/json", "User-Agent": "netbox-store-agent/0.1"}
requested = urlsplit(url)
configured = urlsplit(self.settings.base_url)
requested_origin = (
requested.scheme,
requested.hostname.lower() if requested.hostname else "",
requested.port or (443 if requested.scheme == "https" else 80),
)
configured_origin = (
configured.scheme,
configured.hostname.lower() if configured.hostname else "",
configured.port or (443 if configured.scheme == "https" else 80),
)
# Artifact hosts may be separately allow-listed, but Store credentials
# are never delegated to them.
if requested_origin == configured_origin:
token = self._token()
if token:
headers["Authorization"] = f"Bearer {token}"
return headers
def _response(self, url: str) -> tuple[http.client.HTTPConnection, http.client.HTTPResponse]:
connection, target = self._connection(url)
headers = self._headers(url)
try:
connection.request("GET", target, headers=headers)
response = connection.getresponse()
except (OSError, http.client.HTTPException) as exc:
connection.close()
raise CatalogError("Store request failed") from exc
if response.status != 200:
response.read(4096)
connection.close()
raise HttpStatusError(response.status, f"Store returned HTTP {response.status}")
return connection, response
@staticmethod
def _declared_length(response: http.client.HTTPResponse, max_bytes: int) -> int | None:
raw = response.getheader("Content-Length")
if raw is None:
return None
try:
length = int(raw)
except ValueError as exc:
raise CatalogError("Store returned an invalid Content-Length") from exc
if length < 0 or length > max_bytes:
raise CatalogError("Store response exceeds the configured size limit")
return length
def get_json(self, url: str, max_bytes: int) -> Any:
connection, response = self._response(url)
try:
declared = self._declared_length(response, max_bytes)
body = response.read(max_bytes + 1)
if len(body) > max_bytes or (declared is not None and declared != len(body)):
raise CatalogError("Store response has an invalid or excessive size")
finally:
connection.close()
try:
return json.loads(body.decode("utf-8"))
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
raise CatalogError("Store returned invalid UTF-8 JSON") from exc
def download(
self,
url: str,
destination: Path,
*,
max_bytes: int,
expected_size: int,
expected_sha256: str,
) -> None:
connection, response = self._response(url)
descriptor: int | None = None
try:
declared = self._declared_length(response, max_bytes)
if declared is not None and declared != expected_size:
raise CatalogError("artifact Content-Length does not match the approved size")
flags = os.O_CREAT | os.O_EXCL | os.O_WRONLY | getattr(os, "O_NOFOLLOW", 0)
descriptor = os.open(destination, flags, 0o600)
digest = hashlib.sha256()
total = 0
while True:
chunk = response.read(1024 * 1024)
if not chunk:
break
total += len(chunk)
if total > max_bytes or total > expected_size:
raise CatalogError("artifact exceeds its approved size")
digest.update(chunk)
os.write(descriptor, chunk)
os.fsync(descriptor)
if total != expected_size:
raise CatalogError("artifact size does not match the approved size")
if digest.hexdigest() != expected_sha256:
raise CatalogError("artifact SHA-256 does not match the approved digest")
finally:
if descriptor is not None:
os.close(descriptor)
connection.close()
@dataclass(frozen=True)
class PluginMetadata:
slug: str
package_name: str
import_name: str
min_netbox_version: str
max_netbox_version: str
releases: tuple[dict[str, Any], ...]
@dataclass(frozen=True)
class ReleasePlan:
plugin: PluginMetadata
version: str
download_url: str
filename: str
artifact_sha256: str
artifact_size: int
approved_payload_sha256: str
def requirement(self) -> dict[str, Any]:
return {
"package_name": self.plugin.package_name,
"version": self.version,
"download_url": self.download_url,
"filename": self.filename,
"sha256": self.artifact_sha256,
"size": self.artifact_size,
}
def _object(value: Any, context: str) -> dict[str, Any]:
if not isinstance(value, dict) or not all(isinstance(key, str) for key in value):
raise CatalogError(f"{context} must be a JSON object")
return value
def _strict_fields(data: dict[str, Any], required: set[str], context: str) -> None:
missing = sorted(required - set(data))
unknown = sorted(set(data) - required)
if missing:
raise CatalogError(f"{context} is missing fields: {', '.join(missing)}")
if unknown:
raise CatalogError(f"{context} contains unknown fields: {', '.join(unknown)}")
def _bounded_string(value: Any, field: str, maximum: int = 4096, *, empty: bool = False) -> str:
if not isinstance(value, str) or len(value) > maximum or (not empty and not value):
raise CatalogError(f"{field} must be a bounded string")
if any(ord(char) < 32 for char in value):
raise CatalogError(f"{field} contains control characters")
return value
class StoreClient:
PLUGIN_FIELDS = {
"api_version",
"slug",
"name",
"summary",
"description",
"repository_url",
"latest_version",
"package_name",
"import_name",
"min_netbox_version",
"max_netbox_version",
"approved",
"status",
"releases",
}
RELEASE_FIELDS = {
"version",
"download_url",
"sha256",
"artifact_size",
"commit_sha",
"min_netbox_version",
"max_netbox_version",
"published_at",
"approved",
"status",
"immutable",
"approved_payload_sha256",
}
def __init__(self, config: Config, transport: Transport | None = None):
self.config = config
self.transport = transport or SecureHTTPTransport(config.store)
def _url(self, template: str, **values: str) -> str:
quoted = {key: quote(value, safe="") for key, value in values.items()}
return self.config.store.base_url + template.format(**quoted)
def _get_with_slash_fallback(self, url: str) -> Any:
candidates = (url, url[:-1] if url.endswith("/") else url + "/")
for index, candidate in enumerate(dict.fromkeys(candidates)):
try:
return self.transport.get_json(candidate, self.config.store.max_catalog_bytes)
except HttpStatusError as exc:
if exc.http_status == 404 and index == 0:
continue
raise
raise CatalogError("Store resource was not found")
@staticmethod
def _version(value: Any, field: str) -> Version:
text = _bounded_string(value, field, 100)
try:
return Version(text)
except InvalidVersion as exc:
raise CatalogError(f"{field} is not a valid version") from exc
def _check_compatibility(self, minimum: Any, maximum: Any, context: str) -> tuple[str, str]:
minimum_text = _bounded_string(minimum, f"{context}.min_netbox_version", 100)
maximum_text = _bounded_string(maximum, f"{context}.max_netbox_version", 100)
minimum_version = self._version(minimum_text, f"{context}.min_netbox_version")
maximum_version = self._version(maximum_text, f"{context}.max_netbox_version")
current = Version(self.config.policy.netbox_version)
if minimum_version > maximum_version or not minimum_version <= current <= maximum_version:
raise CatalogError(f"{context} is incompatible with configured NetBox")
return minimum_text, maximum_text
def get_plugin(self, slug: str) -> PluginMetadata:
if not SLUG_RE.fullmatch(slug):
raise CatalogError("plugin slug is invalid")
raw = self._get_with_slash_fallback(
self._url(self.config.store.plugin_endpoint_template, plugin_slug=slug)
)
data = _object(raw, "plugin")
_strict_fields(data, self.PLUGIN_FIELDS, "plugin")
if data["api_version"] != "v1" or data["slug"] != slug:
raise CatalogError("plugin identity or API version does not match the request")
if data["approved"] is not True or data["status"] != "approved":
raise CatalogError("plugin is not approved")
package_name = _bounded_string(data["package_name"], "plugin.package_name", 200)
import_name = _bounded_string(data["import_name"], "plugin.import_name", 200)
if not DIST_NAME_RE.fullmatch(package_name) or not IMPORT_NAME_RE.fullmatch(import_name):
raise CatalogError("plugin package_name or import_name is invalid")
minimum, maximum = self._check_compatibility(
data["min_netbox_version"], data["max_netbox_version"], "plugin"
)
for field in ("name", "summary", "description", "repository_url"):
_bounded_string(data[field], f"plugin.{field}", 65535, empty=field in {"summary", "description"})
if data["latest_version"] is not None:
self._version(data["latest_version"], "plugin.latest_version")
releases = data["releases"]
if not isinstance(releases, list) or len(releases) > 1000:
raise CatalogError("plugin.releases must be a bounded array")
release_objects = tuple(_object(item, "plugin.releases[]") for item in releases)
for release in release_objects:
_strict_fields(release, self.RELEASE_FIELDS, "plugin.releases[]")
return PluginMetadata(slug, package_name, import_name, minimum, maximum, release_objects)
def _parse_release(
self, plugin: PluginMetadata, raw: Any, requested_version: str
) -> ReleasePlan:
data = _object(raw, "release")
_strict_fields(data, self.RELEASE_FIELDS, "release")
version_text = _bounded_string(data["version"], "release.version", 100)
version = self._version(version_text, "release.version")
if version_text != requested_version:
raise CatalogError("release version does not match the request")
if version.is_prerelease and not self.config.policy.allow_prereleases:
raise CatalogError("prerelease versions are forbidden by policy")
if data["approved"] is not True or data["status"] != "approved" or data["immutable"] is not True:
raise CatalogError("release is not approved and immutable")
self._check_compatibility(
data["min_netbox_version"], data["max_netbox_version"], "release"
)
download_url = _bounded_string(data["download_url"], "release.download_url", 8192)
sha256 = _bounded_string(data["sha256"], "release.sha256", 64)
if not SHA256_RE.fullmatch(sha256):
raise CatalogError("release.sha256 must be lowercase SHA-256")
size = data["artifact_size"]
if isinstance(size, bool) or not isinstance(size, int) or not 1 <= size <= self.config.store.max_artifact_bytes:
raise CatalogError("release.artifact_size is invalid")
token = _bounded_string(
data["approved_payload_sha256"], "release.approved_payload_sha256", 64
)
if not SHA256_RE.fullmatch(token):
raise CatalogError("release.approved_payload_sha256 must be lowercase SHA-256")
commit = _bounded_string(data["commit_sha"], "release.commit_sha", 64, empty=True)
if commit and (len(commit) != 40 or any(char not in "0123456789abcdef" for char in commit)):
raise CatalogError("release.commit_sha is invalid")
if data["published_at"] is not None:
_bounded_string(data["published_at"], "release.published_at", 100)
parsed_url = urlsplit(download_url)
filename = Path(parsed_url.path).name
if not filename or not filename.endswith(".whl") or len(filename) > 255:
raise CatalogError("approved artifact must be a wheel with a safe filename")
# URL host/scheme/DNS are revalidated by the transport at download time.
return ReleasePlan(plugin, version_text, download_url, filename, sha256, size, token)
def get_release(self, slug: str, version: str) -> ReleasePlan:
plugin = self.get_plugin(slug)
url = self._url(
self.config.store.release_endpoint_template, plugin_slug=slug, version=version
)
try:
raw = self._get_with_slash_fallback(url)
except HttpStatusError as exc:
if exc.http_status != 404:
raise
matches = [item for item in plugin.releases if item.get("version") == version]
if len(matches) != 1:
raise CatalogError("approved release was not found") from exc
raw = matches[0]
return self._parse_release(plugin, raw, version)
def download_release(self, plan: ReleasePlan, directory: Path) -> Path:
destination = directory / plan.filename
self.transport.download(
plan.download_url,
destination,
max_bytes=self.config.store.max_artifact_bytes,
expected_size=plan.artifact_size,
expected_sha256=plan.artifact_sha256,
)
try:
distribution, wheel_version, _build, _tags = parse_wheel_filename(plan.filename)
except (InvalidVersion, ValueError) as exc:
raise CatalogError("artifact filename is not a valid wheel filename") from exc
if canonicalize_name(distribution) != canonicalize_name(plan.plugin.package_name):
raise CatalogError("wheel distribution does not match approved package_name")
if wheel_version != Version(plan.version):
raise CatalogError("wheel version does not match approved release version")
try:
with zipfile.ZipFile(destination) as wheel:
members = wheel.infolist()
if len(members) > 10_000:
raise CatalogError("wheel contains too many archive members")
expanded = 0
for member in members:
parts = Path(member.filename).parts
if (
not member.filename
or member.filename.startswith(("/", "\\"))
or ".." in parts
or "\x00" in member.filename
):
raise CatalogError("wheel contains an unsafe archive member")
expanded += member.file_size
if expanded > min(self.config.store.max_artifact_bytes * 20, 2 * 1024**3):
raise CatalogError("wheel expands beyond the configured safety limit")
if wheel.testzip() is not None:
raise CatalogError("wheel archive failed its integrity check")
except (OSError, zipfile.BadZipFile) as exc:
raise CatalogError("artifact is not a valid wheel archive") from exc
return destination
+100
View File
@@ -0,0 +1,100 @@
from __future__ import annotations
import argparse
import getpass
import json
import signal
import sys
import threading
import uuid
from pathlib import Path
from typing import Any
from .client import exchange
from .config import load_config
from .constants import ACTIONS, DEFAULT_CONFIG_PATH, PROTOCOL_VERSION
from .daemon import serve
from .errors import AgentError
def _print_response(value: dict[str, Any]) -> int:
print(json.dumps(value, indent=2, ensure_ascii=False, sort_keys=True))
return 0 if int(value.get("status", 500)) < 400 else 1
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(prog="netbox-store-agent")
parser.add_argument("--config", default=DEFAULT_CONFIG_PATH)
parser.add_argument("--socket", default="/run/netbox-store-agent/agent.sock")
subcommands = parser.add_subparsers(dest="command", required=True)
subcommands.add_parser("daemon", help="run the Unix-socket daemon")
subcommands.add_parser("validate-config", help="validate the root-owned TOML configuration")
subcommands.add_parser("capabilities", help="query daemon capabilities")
status = subcommands.add_parser("status", help="query one operation")
status.add_argument("operation_id")
submit = subcommands.add_parser("submit", help="submit one lifecycle operation")
submit.add_argument("action", choices=sorted(ACTIONS))
submit.add_argument("plugin_slug")
submit.add_argument("--version")
submit.add_argument("--approved-payload-sha256")
submit.add_argument("--idempotency-key", default=None)
submit.add_argument("--request-id", default=None)
submit.add_argument("--requested-by", default=f"cli:{getpass.getuser()}")
return parser
def main(argv: list[str] | None = None) -> int:
arguments = build_parser().parse_args(argv)
try:
if arguments.command == "validate-config":
load_config(arguments.config)
print("configuration is valid")
return 0
if arguments.command == "daemon":
config = load_config(arguments.config)
stop = threading.Event()
def request_stop(_signum: int, _frame: object) -> None:
stop.set()
signal.signal(signal.SIGTERM, request_stop)
signal.signal(signal.SIGINT, request_stop)
serve(config, stop)
return 0
if arguments.command == "capabilities":
request = {
"protocol_version": PROTOCOL_VERSION,
"method": "GET",
"path": "/v1/capabilities",
}
elif arguments.command == "status":
request = {
"protocol_version": PROTOCOL_VERSION,
"method": "GET",
"path": f"/v1/operations/{arguments.operation_id}",
}
else:
request = {
"protocol_version": PROTOCOL_VERSION,
"method": "POST",
"path": "/v1/operations",
"idempotency_key": arguments.idempotency_key or str(uuid.uuid4()),
"body": {
"request_id": arguments.request_id or str(uuid.uuid4()),
"action": arguments.action,
"plugin_slug": arguments.plugin_slug,
"version": arguments.version,
"approved_payload_sha256": arguments.approved_payload_sha256,
"requested_by": arguments.requested_by,
},
}
return _print_response(exchange(Path(arguments.socket), request))
except (AgentError, OSError) as exc:
print(f"error: {exc}", file=sys.stderr)
return 2
if __name__ == "__main__": # pragma: no cover
raise SystemExit(main())
@@ -0,0 +1,47 @@
from __future__ import annotations
import json
import socket
from pathlib import Path
from typing import Any
from .constants import DEFAULT_MAX_REQUEST_BYTES, PROTOCOL_VERSION
from .errors import ValidationError
def exchange(socket_path: str | Path, request: dict[str, Any], timeout: float = 10) -> dict[str, Any]:
payload = (json.dumps(request, separators=(",", ":"), ensure_ascii=False) + "\n").encode("utf-8")
if len(payload) > DEFAULT_MAX_REQUEST_BYTES:
raise ValidationError("request exceeds the protocol size limit")
connection = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
connection.settimeout(timeout)
try:
connection.connect(str(socket_path))
connection.sendall(payload)
buffer = bytearray()
while b"\n" not in buffer:
chunk = connection.recv(4096)
if not chunk:
raise ValidationError("agent closed without a complete response")
buffer.extend(chunk)
if len(buffer) > DEFAULT_MAX_REQUEST_BYTES:
raise ValidationError("agent response exceeds the protocol size limit")
line, separator, trailing = bytes(buffer).partition(b"\n")
if not separator or trailing:
raise ValidationError("agent returned more than one JSON line")
finally:
connection.close()
try:
response = json.loads(line.decode("utf-8"))
except (UnicodeDecodeError, json.JSONDecodeError) as exc:
raise ValidationError("agent returned invalid JSON") from exc
if (
not isinstance(response, dict)
or set(response) != {"protocol_version", "status", "body"}
or response.get("protocol_version") != PROTOCOL_VERSION
or isinstance(response.get("status"), bool)
or not isinstance(response.get("status"), int)
or not isinstance(response.get("body"), dict)
):
raise ValidationError("agent returned an invalid protocol response")
return response
+383
View File
@@ -0,0 +1,383 @@
from __future__ import annotations
import os
import stat
import string
import tomllib
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from urllib.parse import urlsplit
from packaging.version import InvalidVersion, Version
from .constants import DEFAULT_MAX_REQUEST_BYTES
from .errors import PolicyError, ValidationError
from .util import (
SERVICE_NAME_RE,
is_relative_to,
reject_unknown,
require_absolute_path,
require_mapping,
)
@dataclass(frozen=True)
class AgentSettings:
socket_path: Path
journal_path: Path
lock_path: Path
backup_dir: Path
dry_run: bool
require_root: bool
require_peer_credentials: bool
allowed_peer_uids: tuple[int, ...]
allowed_peer_gids: tuple[int, ...]
socket_mode: int
socket_uid: int
socket_gid: int
max_request_bytes: int
connection_timeout_seconds: float
worker_threads: int
@dataclass(frozen=True)
class StoreSettings:
base_url: str
plugin_endpoint_template: str
release_endpoint_template: str
timeout_seconds: float
max_catalog_bytes: int
max_artifact_bytes: int
allow_private_addresses: bool
allow_http_for_testing: bool
allowed_hosts: tuple[str, ...]
bearer_token_file: Path | None
ca_file: Path | None
@dataclass(frozen=True)
class PathSettings:
allowed_root: Path
include_path: Path
requirements_path: Path
temp_dir: Path
@dataclass(frozen=True)
class CommandSettings:
python_path: Path
manage_path: Path
systemctl_path: Path
services: tuple[str, ...]
command_timeout_seconds: int
@dataclass(frozen=True)
class PolicySettings:
netbox_version: str
min_supported_netbox: str
max_supported_netbox: str
self_plugin_slugs: tuple[str, ...]
allow_prereleases: bool
require_release_for_enable: bool
@dataclass(frozen=True)
class Config:
agent: AgentSettings
store: StoreSettings
paths: PathSettings
commands: CommandSettings
policy: PolicySettings
def _table(root: dict[str, Any], name: str, allowed: set[str]) -> dict[str, Any]:
value = require_mapping(root.get(name), name)
reject_unknown(value, allowed, name)
return value
def _bool(data: dict[str, Any], key: str, default: bool) -> bool:
value = data.get(key, default)
if not isinstance(value, bool):
raise ValidationError(f"{key} must be a boolean")
return value
def _int(data: dict[str, Any], key: str, default: int, minimum: int, maximum: int) -> int:
value = data.get(key, default)
if isinstance(value, bool) or not isinstance(value, int) or not minimum <= value <= maximum:
raise ValidationError(f"{key} must be an integer between {minimum} and {maximum}")
return value
def _float(data: dict[str, Any], key: str, default: float, minimum: float, maximum: float) -> float:
value = data.get(key, default)
if isinstance(value, bool) or not isinstance(value, (int, float)):
raise ValidationError(f"{key} must be numeric")
result = float(value)
if not minimum <= result <= maximum:
raise ValidationError(f"{key} must be between {minimum} and {maximum}")
return result
def _string(data: dict[str, Any], key: str, default: str | None = None) -> str:
value = data.get(key, default)
if not isinstance(value, str) or not value:
raise ValidationError(f"{key} must be a non-empty string")
if "\x00" in value:
raise ValidationError(f"{key} contains a NUL byte")
return value
def _int_tuple(data: dict[str, Any], key: str, default: tuple[int, ...]) -> tuple[int, ...]:
value = data.get(key, list(default))
if not isinstance(value, list) or not value:
raise ValidationError(f"{key} must be a non-empty array")
result: list[int] = []
for item in value:
if isinstance(item, bool) or not isinstance(item, int) or item < 0:
raise ValidationError(f"{key} must contain non-negative integers")
result.append(item)
return tuple(sorted(set(result)))
def _string_tuple(data: dict[str, Any], key: str, default: tuple[str, ...]) -> tuple[str, ...]:
value = data.get(key, list(default))
if not isinstance(value, list) or not value or not all(isinstance(item, str) and item for item in value):
raise ValidationError(f"{key} must be a non-empty string array")
return tuple(dict.fromkeys(value))
def _optional_path(data: dict[str, Any], key: str) -> Path | None:
value = data.get(key)
if value in (None, ""):
return None
return require_absolute_path(value, key)
def _validate_endpoint_template(value: str, field: str, expected: set[str]) -> None:
if not value.startswith("/") or value.startswith("//") or "?" in value or "#" in value:
raise ValidationError(f"{field} must be an absolute URL path without query or fragment")
fields = {name for _, name, _, _ in string.Formatter().parse(value) if name is not None}
if fields != expected:
raise ValidationError(f"{field} placeholders must be exactly: {', '.join(sorted(expected))}")
if ".." in value.split("/"):
raise ValidationError(f"{field} may not contain '..'")
def _validate_config_file(path: Path, require_root_owner: bool) -> None:
try:
metadata = path.lstat()
except FileNotFoundError as exc:
raise ValidationError(f"configuration file does not exist: {path}") from exc
if stat.S_ISLNK(metadata.st_mode) or not stat.S_ISREG(metadata.st_mode):
raise ValidationError("configuration file must be a regular file, not a symlink")
if os.name == "posix" and metadata.st_mode & (stat.S_IWGRP | stat.S_IWOTH):
raise ValidationError("configuration file must not be group/world writable")
if (
os.name == "posix"
and require_root_owner
and hasattr(metadata, "st_uid")
and metadata.st_uid != 0
):
raise ValidationError("configuration file must be owned by root")
def load_config(path: str | Path, *, allow_insecure_owner: bool = False) -> Config:
config_path = Path(path)
_validate_config_file(config_path, require_root_owner=not allow_insecure_owner)
try:
with config_path.open("rb") as handle:
root = tomllib.load(handle)
except (OSError, tomllib.TOMLDecodeError) as exc:
raise ValidationError(f"cannot read configuration: {exc}") from exc
reject_unknown(root, {"agent", "store", "paths", "commands", "policy"}, "configuration")
missing_tables = sorted({"agent", "store", "paths", "commands", "policy"} - set(root))
if missing_tables:
raise ValidationError(f"configuration is missing tables: {', '.join(missing_tables)}")
agent_data = _table(
root,
"agent",
{
"socket_path",
"journal_path",
"lock_path",
"backup_dir",
"dry_run",
"require_root",
"require_peer_credentials",
"allowed_peer_uids",
"allowed_peer_gids",
"socket_mode",
"socket_uid",
"socket_gid",
"max_request_bytes",
"connection_timeout_seconds",
"worker_threads",
},
)
store_data = _table(
root,
"store",
{
"base_url",
"plugin_endpoint_template",
"release_endpoint_template",
"timeout_seconds",
"max_catalog_bytes",
"max_artifact_bytes",
"allow_private_addresses",
"allow_http_for_testing",
"allowed_hosts",
"bearer_token_file",
"ca_file",
},
)
paths_data = _table(root, "paths", {"allowed_root", "include_path", "requirements_path", "temp_dir"})
commands_data = _table(
root,
"commands",
{"python_path", "manage_path", "systemctl_path", "services", "command_timeout_seconds"},
)
policy_data = _table(
root,
"policy",
{
"netbox_version",
"min_supported_netbox",
"max_supported_netbox",
"self_plugin_slugs",
"allow_prereleases",
"require_release_for_enable",
},
)
agent = AgentSettings(
socket_path=require_absolute_path(_string(agent_data, "socket_path"), "socket_path"),
journal_path=require_absolute_path(_string(agent_data, "journal_path"), "journal_path"),
lock_path=require_absolute_path(_string(agent_data, "lock_path"), "lock_path"),
backup_dir=require_absolute_path(_string(agent_data, "backup_dir"), "backup_dir"),
dry_run=_bool(agent_data, "dry_run", True),
require_root=_bool(agent_data, "require_root", True),
require_peer_credentials=_bool(agent_data, "require_peer_credentials", True),
allowed_peer_uids=_int_tuple(agent_data, "allowed_peer_uids", (0,)),
allowed_peer_gids=_int_tuple(agent_data, "allowed_peer_gids", (0,)),
socket_mode=_int(agent_data, "socket_mode", 0o660, 0, 0o777),
socket_uid=_int(agent_data, "socket_uid", 0, 0, 2**31 - 1),
socket_gid=_int(agent_data, "socket_gid", 0, 0, 2**31 - 1),
max_request_bytes=_int(
agent_data, "max_request_bytes", DEFAULT_MAX_REQUEST_BYTES, 1024, DEFAULT_MAX_REQUEST_BYTES
),
connection_timeout_seconds=_float(agent_data, "connection_timeout_seconds", 5, 0.1, 60),
worker_threads=_int(agent_data, "worker_threads", 1, 1, 4),
)
allow_http = _bool(store_data, "allow_http_for_testing", False)
base_url = _string(store_data, "base_url").rstrip("/")
parsed_url = urlsplit(base_url)
allowed_schemes = {"https", "http"} if allow_http else {"https"}
if parsed_url.scheme not in allowed_schemes or not parsed_url.hostname:
raise ValidationError("store.base_url must use HTTPS and contain a hostname")
if parsed_url.username or parsed_url.password or parsed_url.query or parsed_url.fragment:
raise ValidationError("store.base_url may not contain credentials, query, or fragment")
allowed_hosts = tuple(host.lower() for host in _string_tuple(store_data, "allowed_hosts", (parsed_url.hostname,)))
if parsed_url.hostname.lower() not in allowed_hosts:
raise ValidationError("store.base_url hostname must be present in store.allowed_hosts")
plugin_template = _string(
store_data, "plugin_endpoint_template", "/api/v1/plugins/{plugin_slug}"
)
release_template = _string(
store_data,
"release_endpoint_template",
"/api/v1/plugins/{plugin_slug}/releases/{version}",
)
_validate_endpoint_template(plugin_template, "plugin_endpoint_template", {"plugin_slug"})
_validate_endpoint_template(
release_template, "release_endpoint_template", {"plugin_slug", "version"}
)
store = StoreSettings(
base_url=base_url,
plugin_endpoint_template=plugin_template,
release_endpoint_template=release_template,
timeout_seconds=_float(store_data, "timeout_seconds", 10, 0.1, 120),
max_catalog_bytes=_int(store_data, "max_catalog_bytes", 1024 * 1024, 1024, 8 * 1024 * 1024),
max_artifact_bytes=_int(
store_data, "max_artifact_bytes", 256 * 1024 * 1024, 1024, 2 * 1024 * 1024 * 1024
),
allow_private_addresses=_bool(store_data, "allow_private_addresses", False),
allow_http_for_testing=allow_http,
allowed_hosts=allowed_hosts,
bearer_token_file=_optional_path(store_data, "bearer_token_file"),
ca_file=_optional_path(store_data, "ca_file"),
)
allowed_root = require_absolute_path(_string(paths_data, "allowed_root"), "allowed_root").resolve()
include_path = require_absolute_path(_string(paths_data, "include_path"), "include_path")
requirements_path = require_absolute_path(
_string(paths_data, "requirements_path"), "requirements_path"
)
for field, path_value in (("include_path", include_path), ("requirements_path", requirements_path)):
if not is_relative_to(path_value.resolve(strict=False), allowed_root):
raise ValidationError(f"paths.{field} must be below paths.allowed_root")
if include_path.resolve(strict=False) == requirements_path.resolve(strict=False):
raise ValidationError("paths.include_path and paths.requirements_path must be distinct")
if include_path.name == requirements_path.name:
raise ValidationError("managed files must have distinct basenames for unambiguous backups")
paths = PathSettings(
allowed_root=allowed_root,
include_path=include_path,
requirements_path=requirements_path,
temp_dir=require_absolute_path(_string(paths_data, "temp_dir"), "temp_dir"),
)
services = _string_tuple(commands_data, "services", ("netbox", "netbox-rq"))
if any(not SERVICE_NAME_RE.fullmatch(service) for service in services):
raise ValidationError("commands.services contains an invalid systemd unit name")
commands = CommandSettings(
python_path=require_absolute_path(_string(commands_data, "python_path"), "python_path"),
manage_path=require_absolute_path(_string(commands_data, "manage_path"), "manage_path"),
systemctl_path=require_absolute_path(_string(commands_data, "systemctl_path"), "systemctl_path"),
services=services,
command_timeout_seconds=_int(commands_data, "command_timeout_seconds", 900, 1, 7200),
)
versions: dict[str, str] = {}
for key, default in (
("netbox_version", "4.6.8"),
("min_supported_netbox", "4.6.5"),
("max_supported_netbox", "4.6.8"),
):
value = _string(policy_data, key, default)
try:
Version(value)
except InvalidVersion as exc:
raise ValidationError(f"policy.{key} is not a valid version") from exc
versions[key] = value
if not Version(versions["min_supported_netbox"]) <= Version(versions["netbox_version"]) <= Version(
versions["max_supported_netbox"]
):
raise PolicyError("configured NetBox version is outside the agent support range")
policy = PolicySettings(
netbox_version=versions["netbox_version"],
min_supported_netbox=versions["min_supported_netbox"],
max_supported_netbox=versions["max_supported_netbox"],
self_plugin_slugs=tuple(
slug.lower()
for slug in _string_tuple(
policy_data,
"self_plugin_slugs",
("netbox-store", "netbox-plugin-store", "netbox_plugin_store"),
)
),
allow_prereleases=_bool(policy_data, "allow_prereleases", False),
require_release_for_enable=_bool(policy_data, "require_release_for_enable", True),
)
return Config(agent=agent, store=store, paths=paths, commands=commands, policy=policy)
def enforce_runtime_identity(config: Config) -> None:
if config.agent.require_root and hasattr(os, "geteuid") and os.geteuid() != 0:
raise PolicyError("daemon must run as root")
@@ -0,0 +1,9 @@
from __future__ import annotations
AGENT_VERSION = "0.1.0"
PROTOCOL_VERSION = 1
DEFAULT_CONFIG_PATH = "/etc/netbox-store-agent/agent.toml"
DEFAULT_MAX_REQUEST_BYTES = 64 * 1024
ACTIONS = frozenset({"install", "update", "enable", "disable", "uninstall"})
TERMINAL_STATES = frozenset({"succeeded", "failed", "manual_recovery", "dry_run"})
+137
View File
@@ -0,0 +1,137 @@
from __future__ import annotations
import os
import socket
import stat
import struct
import threading
from pathlib import Path
from .config import Config, enforce_runtime_identity
from .errors import AgentError, AuthenticationError, ValidationError
from .journal import Journal
from .protocol import decode_request, encode_response, error_response
from .service import AgentService
def _authorize_peer(connection: socket.socket, config: Config) -> None:
if not config.agent.require_peer_credentials:
return
option = getattr(socket, "SO_PEERCRED", None)
if option is None:
raise AuthenticationError("platform does not expose Unix peer credentials")
try:
raw = connection.getsockopt(socket.SOL_SOCKET, option, struct.calcsize("3i"))
_pid, uid, gid = struct.unpack("3i", raw)
except (OSError, struct.error) as exc:
raise AuthenticationError("Unix peer credentials could not be verified") from exc
if uid not in config.agent.allowed_peer_uids and gid not in config.agent.allowed_peer_gids:
raise AuthenticationError("Unix peer is not allow-listed")
def _read_frame(connection: socket.socket, maximum: int) -> bytes:
buffer = bytearray()
while True:
chunk = connection.recv(min(4096, maximum + 1 - len(buffer)))
if not chunk:
raise ValidationError("connection closed before the JSON-line terminator")
buffer.extend(chunk)
if len(buffer) > maximum:
raise ValidationError("request exceeds the configured size limit")
newline = buffer.find(b"\n")
if newline >= 0:
if newline >= maximum or newline != len(buffer) - 1:
raise ValidationError("connection must contain exactly one JSON line")
return bytes(buffer[:newline])
def handle_connection(connection: socket.socket, config: Config, service: AgentService) -> None:
try:
connection.settimeout(config.agent.connection_timeout_seconds)
_authorize_peer(connection, config)
request = decode_request(_read_frame(connection, config.agent.max_request_bytes), config.agent.max_request_bytes)
result = service.handle(request)
except AgentError as exc:
result = error_response(exc)
except (OSError, TimeoutError):
result = error_response(ValidationError("connection timed out or failed"))
except Exception:
result = error_response(
AgentError("unexpected internal error", code="unexpected_error", status=500)
)
try:
try:
connection.shutdown(socket.SHUT_RD)
except OSError:
pass
connection.sendall(encode_response(result))
except OSError:
pass
finally:
connection.close()
def _systemd_socket() -> socket.socket | None:
try:
listen_pid = int(os.environ.get("LISTEN_PID", "0"))
listen_fds = int(os.environ.get("LISTEN_FDS", "0"))
except ValueError:
return None
if listen_pid != os.getpid() or listen_fds != 1:
return None
inherited = socket.socket(fileno=3)
if inherited.family != socket.AF_UNIX or inherited.type & socket.SOCK_STREAM != socket.SOCK_STREAM:
inherited.close()
raise ValidationError("systemd passed an unexpected socket type")
os.environ.pop("LISTEN_PID", None)
os.environ.pop("LISTEN_FDS", None)
os.environ.pop("LISTEN_FDNAMES", None)
return inherited
def _bind_socket(config: Config) -> socket.socket:
path = config.agent.socket_path
path.parent.mkdir(parents=True, exist_ok=True, mode=0o750)
if path.exists() or path.is_symlink():
metadata = path.lstat()
if stat.S_ISLNK(metadata.st_mode) or not stat.S_ISSOCK(metadata.st_mode):
raise ValidationError("configured socket path exists and is not a Unix socket")
path.unlink()
listener = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM)
try:
listener.bind(str(path))
os.chmod(path, config.agent.socket_mode)
os.chown(path, config.agent.socket_uid, config.agent.socket_gid)
listener.listen(32)
except Exception:
listener.close()
raise
return listener
def serve(config: Config, stop_event: threading.Event | None = None) -> None:
enforce_runtime_identity(config)
stop = stop_event or threading.Event()
listener = _systemd_socket()
owns_socket = listener is None
if listener is None:
listener = _bind_socket(config)
journal = Journal(
config.agent.journal_path, require_root_owner=config.agent.require_root
)
service = AgentService(config, journal)
listener.settimeout(1.0)
try:
while not stop.is_set():
try:
connection, _address = listener.accept()
except socket.timeout:
continue
handle_connection(connection, config, service)
finally:
listener.close()
service.close()
if owns_socket:
path = Path(config.agent.socket_path)
if path.exists() and stat.S_ISSOCK(path.lstat().st_mode):
path.unlink()
@@ -0,0 +1,55 @@
from __future__ import annotations
class AgentError(Exception):
"""Base error carrying a stable machine-readable code and HTTP-like status."""
code = "agent_error"
status = 500
def __init__(self, message: str, *, code: str | None = None, status: int | None = None):
super().__init__(message)
self.message = message
if code is not None:
self.code = code
if status is not None:
self.status = status
class ValidationError(AgentError):
code = "invalid_request"
status = 400
class AuthenticationError(AgentError):
code = "peer_not_allowed"
status = 403
class NotFoundError(AgentError):
code = "not_found"
status = 404
class ConflictError(AgentError):
code = "conflict"
status = 409
class CatalogError(AgentError):
code = "catalog_rejected"
status = 422
class PolicyError(AgentError):
code = "policy_rejected"
status = 422
class ExecutionError(AgentError):
code = "execution_failed"
status = 500
class ManualRecoveryRequired(ExecutionError):
code = "manual_recovery_required"
@@ -0,0 +1,353 @@
from __future__ import annotations
import secrets
import tempfile
from pathlib import Path
from typing import Callable
from .catalog import PluginMetadata, ReleasePlan, StoreClient
from .config import Config
from .errors import (
AgentError,
ConflictError,
ManualRecoveryRequired,
NotFoundError,
PolicyError,
)
from .journal import Journal, ManagedPlugin
from .locking import GlobalFileLock
from .managed_files import FileSnapshot, ManagedFiles
from .protocol import OperationRequest
from .runner import Runner, SubprocessRunner
class OperationProcessor:
def __init__(
self,
config: Config,
journal: Journal,
*,
store: StoreClient | None = None,
files: ManagedFiles | None = None,
runner: Runner | None = None,
lock_factory: Callable[[Path], GlobalFileLock] = GlobalFileLock,
):
self.config = config
self.journal = journal
self.store = store or StoreClient(config)
self.files = files or ManagedFiles(config)
self.runner = runner or SubprocessRunner(config)
self.lock_factory = lock_factory
self._mutated_operations: set[str] = set()
self._snapshots: dict[str, FileSnapshot] = {}
def _event(self, operation_id: str, step: str, message: str) -> None:
self.journal.add_event(operation_id, "info", step, message)
self.journal.transition(operation_id, "running", step=step)
def _assert_not_self(self, request: OperationRequest) -> None:
if request.plugin_slug.lower() in self.config.policy.self_plugin_slugs:
raise PolicyError("self-management is forbidden")
@staticmethod
def _same_identity(existing: ManagedPlugin, plugin: PluginMetadata) -> None:
if existing.package_name != plugin.package_name or existing.import_name != plugin.import_name:
raise PolicyError("approved package/import identity changed for an existing plugin")
@staticmethod
def _same_release(existing: ManagedPlugin, plan: ReleasePlan) -> None:
if existing.version != plan.version or existing.requirements != (plan.requirement(),):
raise PolicyError(
"managed artifact lock differs from the current immutable Store release"
)
def _plugins_with(
self, replacement: ManagedPlugin | None = None, *, delete_slug: str | None = None
) -> list[ManagedPlugin]:
plugins = []
replaced = False
for item in self.journal.list_managed_plugins():
if item.slug == delete_slug:
continue
if replacement is not None and item.slug == replacement.slug:
plugins.append(replacement)
replaced = True
else:
plugins.append(item)
if replacement is not None and not replaced:
plugins.append(replacement)
return plugins
def _run(self, operation_id: str, step: str, argv: list[str]) -> None:
self._event(operation_id, step, f"Executing fixed {step} command")
self.runner.run(argv)
def _pip_install(
self, operation_id: str, directory: Path, plugin: ManagedPlugin, wheel: Path
) -> None:
requirements = self.files.make_local_requirements(directory, plugin, wheel)
self._run(
operation_id,
"pip_install",
[
str(self.config.commands.python_path),
"-m",
"pip",
"install",
"--no-input",
"--disable-pip-version-check",
"--no-index",
"--no-deps",
"--only-binary=:all:",
"--require-hashes",
"-r",
str(requirements),
],
)
def _pip_uninstall(self, operation_id: str, package_name: str) -> None:
self._run(
operation_id,
"pip_uninstall",
[
str(self.config.commands.python_path),
"-m",
"pip",
"uninstall",
"--yes",
package_name,
],
)
def _netbox_prepare(self, operation_id: str) -> None:
python = str(self.config.commands.python_path)
manage = str(self.config.commands.manage_path)
self._run(operation_id, "migrate", [python, manage, "migrate", "--no-input"])
self._run(
operation_id,
"collectstatic",
[python, manage, "collectstatic", "--no-input"],
)
def _restart(self, operation_id: str) -> None:
self._run(
operation_id,
"restart",
[
str(self.config.commands.systemctl_path),
"restart",
*self.config.commands.services,
],
)
def _write(self, operation_id: str, plugins: list[ManagedPlugin]) -> FileSnapshot:
self._event(operation_id, "managed_files", "Writing managed include and requirements")
snapshot = self.files.write(operation_id, plugins)
self._snapshots[operation_id] = snapshot
return snapshot
def _mark_mutated(self, operation_id: str) -> None:
self._mutated_operations.add(operation_id)
def _release_for_request(self, request: OperationRequest) -> ReleasePlan:
assert request.version is not None
plan = self.store.get_release(request.plugin_slug, request.version)
if request.approved_payload_sha256 is None or not secrets.compare_digest(
request.approved_payload_sha256, plan.approved_payload_sha256
):
raise PolicyError("approval payload token no longer matches the Store")
return plan
def _dry_result(self, request: OperationRequest, **extra: object) -> dict[str, object]:
return {
"dry_run": True,
"action": request.action,
"plugin_slug": request.plugin_slug,
**extra,
}
def _execute(
self, operation_id: str, request: OperationRequest
) -> tuple[dict[str, object], bool]:
self._assert_not_self(request)
existing = self.journal.get_managed_plugin(request.plugin_slug)
action = request.action
host_mutated = False
if action in {"install", "update"}:
if action == "install" and existing is not None:
raise ConflictError("plugin is already managed; use update")
if action == "update" and existing is None:
raise NotFoundError("plugin is not managed; use install")
self._event(operation_id, "catalog", "Fetching approved release from Store")
plan = self._release_for_request(request)
if existing is not None:
self._same_identity(existing, plan.plugin)
enabled = existing.enabled if existing else False
managed = ManagedPlugin(
slug=request.plugin_slug,
package_name=plan.plugin.package_name,
import_name=plan.plugin.import_name,
version=plan.version,
enabled=enabled,
requirements=(plan.requirement(),),
)
temp_root = self.config.paths.temp_dir
if temp_root.is_symlink():
raise PolicyError("temporary directory may not be a symlink")
temp_root.mkdir(parents=True, exist_ok=True, mode=0o700)
with tempfile.TemporaryDirectory(prefix="operation-", dir=temp_root) as temporary:
self._event(operation_id, "artifact", "Downloading and verifying approved wheel")
wheel = self.store.download_release(plan, Path(temporary))
if self.config.agent.dry_run:
return self._dry_result(request, version=plan.version, artifact_verified=True), False
self._mark_mutated(operation_id)
host_mutated = True
self._pip_install(operation_id, Path(temporary), managed, wheel)
self._write(operation_id, self._plugins_with(managed))
if enabled:
self._netbox_prepare(operation_id)
self._restart(operation_id)
self.journal.upsert_managed_plugin(managed)
return {
"dry_run": False,
"action": action,
"plugin_slug": request.plugin_slug,
"version": plan.version,
"enabled": enabled,
}, host_mutated
if existing is None:
raise NotFoundError("plugin is not managed")
self._event(operation_id, "catalog", "Re-fetching approved plugin metadata")
if action == "enable" and self.config.policy.require_release_for_enable:
plan = self.store.get_release(existing.slug, existing.version)
plugin = plan.plugin
self._same_release(existing, plan)
else:
plugin = self.store.get_plugin(existing.slug)
self._same_identity(existing, plugin)
if action == "enable":
if existing.enabled:
raise ConflictError("plugin is already enabled")
replacement = ManagedPlugin(
existing.slug,
existing.package_name,
existing.import_name,
existing.version,
True,
existing.requirements,
)
if self.config.agent.dry_run:
return self._dry_result(request, version=existing.version, enabled=True), False
self._write(operation_id, self._plugins_with(replacement))
host_mutated = True
self._mark_mutated(operation_id)
self._netbox_prepare(operation_id)
self._restart(operation_id)
self.journal.upsert_managed_plugin(replacement)
return self._dry_result(request, dry_run=False, version=existing.version, enabled=True), True
if action == "disable":
if not existing.enabled:
raise ConflictError("plugin is already disabled")
replacement = ManagedPlugin(
existing.slug,
existing.package_name,
existing.import_name,
existing.version,
False,
existing.requirements,
)
if self.config.agent.dry_run:
return self._dry_result(request, version=existing.version, enabled=False), False
self._write(operation_id, self._plugins_with(replacement))
host_mutated = True
self._mark_mutated(operation_id)
self._restart(operation_id)
self.journal.upsert_managed_plugin(replacement)
return self._dry_result(request, dry_run=False, version=existing.version, enabled=False), True
if action == "uninstall":
if existing.enabled:
raise ConflictError("disable the plugin before uninstalling it")
if self.config.agent.dry_run:
return self._dry_result(request, version=existing.version, removed=True), False
self._mark_mutated(operation_id)
host_mutated = True
self._pip_uninstall(operation_id, existing.package_name)
self._write(operation_id, self._plugins_with(delete_slug=existing.slug))
self.journal.delete_managed_plugin(existing.slug)
return self._dry_result(request, dry_run=False, version=existing.version, removed=True), True
raise PolicyError("unsupported action")
def process(self, operation_id: str) -> None:
if not self.journal.claim(operation_id):
return
self.journal.add_event(operation_id, "info", "starting", "Operation worker started")
host_mutated = False
try:
request = self.journal.get_request(operation_id)
with self.lock_factory(self.config.agent.lock_path):
result, host_mutated = self._execute(operation_id, request)
state = "dry_run" if self.config.agent.dry_run else "succeeded"
self.journal.transition(operation_id, state, step="complete", result=result, finished=True)
self.journal.add_event(operation_id, "info", "complete", f"Operation {state}")
except AgentError as exc:
mutated = host_mutated or operation_id in self._mutated_operations
snapshot = self._snapshots.get(operation_id)
if mutated and snapshot is not None:
try:
self.files.restore(snapshot)
self.journal.add_event(
operation_id,
"warning",
"rollback",
"Managed files restored; package/service state still requires inspection.",
)
except AgentError:
exc = ManualRecoveryRequired("managed-file rollback failed")
state = (
"manual_recovery"
if mutated or isinstance(exc, ManualRecoveryRequired)
else "failed"
)
message = exc.message[:2000]
self.journal.transition(
operation_id,
state,
step="failed",
error_code=exc.code,
error_message=message,
finished=True,
)
self.journal.add_event(operation_id, "error", "failed", message)
except Exception:
mutated = host_mutated or operation_id in self._mutated_operations
snapshot = self._snapshots.get(operation_id)
if mutated and snapshot is not None:
try:
self.files.restore(snapshot)
except AgentError:
pass
state = "manual_recovery" if mutated else "failed"
self.journal.transition(
operation_id,
state,
step="failed",
error_code="unexpected_error",
error_message="Unexpected internal error; inspect root-owned service logs.",
finished=True,
)
self.journal.add_event(
operation_id,
"error",
"failed",
"Unexpected internal error; details were withheld from the client.",
)
finally:
self._mutated_operations.discard(operation_id)
self._snapshots.pop(operation_id, None)
@@ -0,0 +1,352 @@
from __future__ import annotations
import json
import os
import sqlite3
import stat
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from .errors import ConflictError, NotFoundError, ValidationError
from .protocol import OperationRequest
from .util import canonical_json, payload_sha256, utc_now
class _ClosingConnection(sqlite3.Connection):
"""sqlite context manager which also releases the OS handle on exit."""
def __exit__(self, exc_type: object, exc_value: object, traceback: object) -> bool:
try:
return super().__exit__(exc_type, exc_value, traceback)
finally:
self.close()
@dataclass(frozen=True)
class ManagedPlugin:
slug: str
package_name: str
import_name: str
version: str
enabled: bool
requirements: tuple[dict[str, Any], ...]
class Journal:
def __init__(self, path: str | Path, *, require_root_owner: bool = False):
self.path = Path(path)
self.require_root_owner = require_root_owner
self._prepare_path()
self._initialize()
def _prepare_path(self) -> None:
self.path.parent.mkdir(parents=True, exist_ok=True, mode=0o700)
if self.path.exists() or self.path.is_symlink():
metadata = self.path.lstat()
if stat.S_ISLNK(metadata.st_mode) or not stat.S_ISREG(metadata.st_mode):
raise ValidationError("journal path must be a regular file, not a symlink")
if os.name == "posix" and metadata.st_mode & (stat.S_IRWXG | stat.S_IRWXO):
raise ValidationError("journal file must not be accessible by group or other")
if os.name == "posix" and self.require_root_owner and metadata.st_uid != 0:
raise ValidationError("journal file must be owned by root")
return
flags = os.O_CREAT | os.O_EXCL | os.O_WRONLY
flags |= getattr(os, "O_NOFOLLOW", 0)
descriptor = os.open(self.path, flags, 0o600)
os.close(descriptor)
try:
os.chmod(self.path, 0o600)
except OSError:
pass
def _connect(self) -> sqlite3.Connection:
connection = sqlite3.connect(
self.path, timeout=30, isolation_level=None, factory=_ClosingConnection
)
connection.row_factory = sqlite3.Row
connection.execute("PRAGMA foreign_keys=ON")
connection.execute("PRAGMA synchronous=FULL")
return connection
def _initialize(self) -> None:
with self._connect() as connection:
connection.execute("PRAGMA journal_mode=WAL")
connection.executescript(
"""
CREATE TABLE IF NOT EXISTS operations (
operation_id TEXT PRIMARY KEY,
request_hash TEXT NOT NULL,
request_json TEXT NOT NULL,
request_id TEXT NOT NULL,
action TEXT NOT NULL,
plugin_slug TEXT NOT NULL,
version TEXT,
approved_payload_sha256 TEXT,
requested_by TEXT NOT NULL,
state TEXT NOT NULL,
current_step TEXT NOT NULL DEFAULT '',
submitted_at TEXT NOT NULL,
updated_at TEXT NOT NULL,
finished_at TEXT,
error_code TEXT,
error_message TEXT,
result_json TEXT
);
CREATE INDEX IF NOT EXISTS operations_state_idx ON operations(state, submitted_at);
CREATE TABLE IF NOT EXISTS operation_events (
id INTEGER PRIMARY KEY AUTOINCREMENT,
operation_id TEXT NOT NULL REFERENCES operations(operation_id) ON DELETE CASCADE,
created_at TEXT NOT NULL,
level TEXT NOT NULL,
step TEXT NOT NULL,
message TEXT NOT NULL
);
CREATE INDEX IF NOT EXISTS operation_events_operation_idx
ON operation_events(operation_id, id);
CREATE TABLE IF NOT EXISTS managed_plugins (
slug TEXT PRIMARY KEY,
package_name TEXT NOT NULL,
import_name TEXT NOT NULL,
version TEXT NOT NULL,
enabled INTEGER NOT NULL CHECK(enabled IN (0, 1)),
requirements_json TEXT NOT NULL,
updated_at TEXT NOT NULL
);
"""
)
@staticmethod
def _operation_dict(row: sqlite3.Row, events: list[dict[str, Any]] | None = None) -> dict[str, Any]:
result = json.loads(row["result_json"]) if row["result_json"] else None
operation = {
"operation_id": row["operation_id"],
"request_id": row["request_id"],
"action": row["action"],
"plugin_slug": row["plugin_slug"],
"version": row["version"],
"approved_payload_sha256": row["approved_payload_sha256"],
"requested_by": row["requested_by"],
"state": row["state"],
"current_step": row["current_step"],
"submitted_at": row["submitted_at"],
"updated_at": row["updated_at"],
"finished_at": row["finished_at"],
"error": (
{"code": row["error_code"], "message": row["error_message"]}
if row["error_code"]
else None
),
"result": result,
}
if events is not None:
operation["events"] = events
return operation
def submit(self, operation_id: str, request: OperationRequest) -> tuple[dict[str, Any], bool]:
request_data = request.as_dict()
request_json = canonical_json(request_data)
request_hash = payload_sha256(request_data)
now = utc_now()
with self._connect() as connection:
connection.execute("BEGIN IMMEDIATE")
existing = connection.execute(
"SELECT * FROM operations WHERE operation_id = ?", (operation_id,)
).fetchone()
if existing is not None:
connection.execute("COMMIT")
if existing["request_hash"] != request_hash:
raise ConflictError("idempotency_key is already bound to a different request")
return self._operation_dict(existing), False
connection.execute(
"""
INSERT INTO operations (
operation_id, request_hash, request_json, request_id, action, plugin_slug,
version, approved_payload_sha256, requested_by, state, submitted_at, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, 'queued', ?, ?)
""",
(
operation_id,
request_hash,
request_json,
request.request_id,
request.action,
request.plugin_slug,
request.version,
request.approved_payload_sha256,
request.requested_by,
now,
now,
),
)
connection.execute("COMMIT")
self.add_event(operation_id, "info", "queued", "Operation accepted")
return self.get_operation(operation_id), True
def get_operation(self, operation_id: str, *, include_events: bool = True) -> dict[str, Any]:
with self._connect() as connection:
row = connection.execute(
"SELECT * FROM operations WHERE operation_id = ?", (operation_id,)
).fetchone()
if row is None:
raise NotFoundError("operation not found")
events = None
if include_events:
event_rows = connection.execute(
"""
SELECT created_at, level, step, message FROM operation_events
WHERE operation_id = ? ORDER BY id ASC LIMIT 40
""",
(operation_id,),
).fetchall()
events = [dict(event) for event in event_rows]
return self._operation_dict(row, events)
def get_request(self, operation_id: str) -> OperationRequest:
with self._connect() as connection:
row = connection.execute(
"SELECT request_json FROM operations WHERE operation_id = ?", (operation_id,)
).fetchone()
if row is None:
raise NotFoundError("operation not found")
data = json.loads(row["request_json"])
return OperationRequest(**data)
def claim(self, operation_id: str) -> bool:
"""Atomically move a queued operation to running exactly once."""
now = utc_now()
with self._connect() as connection:
cursor = connection.execute(
"""
UPDATE operations SET state = 'running', current_step = 'starting', updated_at = ?
WHERE operation_id = ? AND state = 'queued'
""",
(now, operation_id),
)
return cursor.rowcount == 1
def transition(
self,
operation_id: str,
state: str,
*,
step: str = "",
error_code: str | None = None,
error_message: str | None = None,
result: dict[str, Any] | None = None,
finished: bool = False,
) -> None:
now = utc_now()
result_json = canonical_json(result) if result is not None else None
with self._connect() as connection:
cursor = connection.execute(
"""
UPDATE operations
SET state = ?, current_step = ?, updated_at = ?,
finished_at = CASE WHEN ? THEN ? ELSE finished_at END,
error_code = ?, error_message = ?, result_json = ?
WHERE operation_id = ?
""",
(
state,
step,
now,
1 if finished else 0,
now,
error_code,
error_message,
result_json,
operation_id,
),
)
if cursor.rowcount != 1:
raise NotFoundError("operation not found")
def add_event(self, operation_id: str, level: str, step: str, message: str) -> None:
safe_level = level if level in {"debug", "info", "warning", "error"} else "info"
safe_step = step[:80]
safe_message = "".join(char for char in message if char == "\t" or ord(char) >= 32)[:1000]
with self._connect() as connection:
connection.execute(
"""
INSERT INTO operation_events(operation_id, created_at, level, step, message)
VALUES (?, ?, ?, ?, ?)
""",
(operation_id, utc_now(), safe_level, safe_step, safe_message),
)
def queued_operations(self) -> list[str]:
with self._connect() as connection:
rows = connection.execute(
"SELECT operation_id FROM operations WHERE state = 'queued' ORDER BY submitted_at"
).fetchall()
return [row["operation_id"] for row in rows]
def recover_interrupted(self) -> int:
now = utc_now()
with self._connect() as connection:
cursor = connection.execute(
"""
UPDATE operations SET state = 'manual_recovery', current_step = 'agent_restart',
updated_at = ?, finished_at = ?, error_code = 'agent_restarted',
error_message = 'Agent restarted while the operation was running; inspect host state.'
WHERE state = 'running'
""",
(now, now),
)
return cursor.rowcount
@staticmethod
def _managed_from_row(row: sqlite3.Row) -> ManagedPlugin:
requirements = json.loads(row["requirements_json"])
return ManagedPlugin(
slug=row["slug"],
package_name=row["package_name"],
import_name=row["import_name"],
version=row["version"],
enabled=bool(row["enabled"]),
requirements=tuple(requirements),
)
def get_managed_plugin(self, slug: str) -> ManagedPlugin | None:
with self._connect() as connection:
row = connection.execute(
"SELECT * FROM managed_plugins WHERE slug = ?", (slug,)
).fetchone()
return self._managed_from_row(row) if row else None
def list_managed_plugins(self) -> list[ManagedPlugin]:
with self._connect() as connection:
rows = connection.execute("SELECT * FROM managed_plugins ORDER BY slug").fetchall()
return [self._managed_from_row(row) for row in rows]
def upsert_managed_plugin(self, plugin: ManagedPlugin) -> None:
requirements_json = canonical_json(list(plugin.requirements))
with self._connect() as connection:
connection.execute(
"""
INSERT INTO managed_plugins(
slug, package_name, import_name, version, enabled, requirements_json, updated_at
) VALUES (?, ?, ?, ?, ?, ?, ?)
ON CONFLICT(slug) DO UPDATE SET
package_name=excluded.package_name,
import_name=excluded.import_name,
version=excluded.version,
enabled=excluded.enabled,
requirements_json=excluded.requirements_json,
updated_at=excluded.updated_at
""",
(
plugin.slug,
plugin.package_name,
plugin.import_name,
plugin.version,
int(plugin.enabled),
requirements_json,
utc_now(),
),
)
def delete_managed_plugin(self, slug: str) -> None:
with self._connect() as connection:
connection.execute("DELETE FROM managed_plugins WHERE slug = ?", (slug,))
@@ -0,0 +1,80 @@
from __future__ import annotations
import os
import stat
from pathlib import Path
from types import TracebackType
from .errors import ConflictError, ValidationError
try:
import fcntl
except ImportError: # pragma: no cover - Linux is the production target.
fcntl = None
try:
import msvcrt
except ImportError: # pragma: no cover - Windows-only test fallback.
msvcrt = None
class GlobalFileLock:
def __init__(self, path: str | Path):
self.path = Path(path)
self._descriptor: int | None = None
def __enter__(self) -> "GlobalFileLock":
self.path.parent.mkdir(parents=True, exist_ok=True, mode=0o700)
if self.path.is_symlink():
raise ValidationError("lock path may not be a symlink")
flags = os.O_CREAT | os.O_RDWR | getattr(os, "O_NOFOLLOW", 0)
descriptor = os.open(self.path, flags, 0o600)
metadata = os.fstat(descriptor)
if not stat.S_ISREG(metadata.st_mode):
os.close(descriptor)
raise ValidationError("lock path must be a regular file")
if os.name == "posix" and metadata.st_mode & (stat.S_IWGRP | stat.S_IWOTH):
os.close(descriptor)
raise ValidationError("lock file must not be group/world writable")
if (
os.name == "posix"
and hasattr(os, "geteuid")
and os.geteuid() == 0
and metadata.st_uid != 0
):
os.close(descriptor)
raise ValidationError("lock file must be owned by root")
try:
if fcntl is not None:
fcntl.flock(descriptor, fcntl.LOCK_EX | fcntl.LOCK_NB)
elif msvcrt is not None: # pragma: no cover - exercised on Windows CI only.
os.lseek(descriptor, 0, os.SEEK_SET)
if os.fstat(descriptor).st_size == 0:
os.write(descriptor, b"0")
os.lseek(descriptor, 0, os.SEEK_SET)
msvcrt.locking(descriptor, msvcrt.LK_NBLCK, 1)
else: # pragma: no cover
raise ValidationError("platform has no supported file locking primitive")
except (BlockingIOError, OSError) as exc:
os.close(descriptor)
raise ConflictError("another lifecycle operation holds the global lock") from exc
self._descriptor = descriptor
return self
def __exit__(
self,
exc_type: type[BaseException] | None,
exc_value: BaseException | None,
traceback: TracebackType | None,
) -> None:
if self._descriptor is None:
return
try:
if fcntl is not None:
fcntl.flock(self._descriptor, fcntl.LOCK_UN)
elif msvcrt is not None: # pragma: no cover
os.lseek(self._descriptor, 0, os.SEEK_SET)
msvcrt.locking(self._descriptor, msvcrt.LK_UNLCK, 1)
finally:
os.close(self._descriptor)
self._descriptor = None
@@ -0,0 +1,224 @@
from __future__ import annotations
import json
import os
import stat
import tempfile
from dataclasses import dataclass
from pathlib import Path
from typing import Iterable
from urllib.parse import urlsplit
from packaging.utils import canonicalize_name
from .config import Config
from .errors import ManualRecoveryRequired, PolicyError, ValidationError
from .journal import ManagedPlugin
from .util import DIST_NAME_RE, IMPORT_NAME_RE, SHA256_RE, is_relative_to, require_uuid
@dataclass(frozen=True)
class PreviousFile:
path: Path
content: bytes | None
mode: int
@dataclass(frozen=True)
class FileSnapshot:
files: tuple[PreviousFile, ...]
class ManagedFiles:
MAX_MANAGED_FILE_BYTES = 2 * 1024 * 1024
def __init__(self, config: Config):
self.config = config
def _validate_target(self, path: Path) -> None:
root = self.config.paths.allowed_root
if path.is_symlink():
raise PolicyError(f"managed path may not be a symlink: {path}")
resolved = path.resolve(strict=False)
if not is_relative_to(resolved, root):
raise PolicyError(f"managed path escaped the configured root: {path}")
if path.exists() and not path.is_file():
raise PolicyError(f"managed path must be a regular file: {path}")
@staticmethod
def _validate_backup_root(path: Path) -> None:
current = path
while not current.exists() and current != current.parent:
current = current.parent
if current.is_symlink():
raise PolicyError("backup directory ancestry may not be a symlink")
def _read_previous(self, path: Path) -> PreviousFile:
self._validate_target(path)
if not path.exists():
return PreviousFile(path, None, 0o640)
metadata = path.stat()
if metadata.st_size > self.MAX_MANAGED_FILE_BYTES:
raise PolicyError("managed file exceeds the safe backup limit")
return PreviousFile(path, path.read_bytes(), stat.S_IMODE(metadata.st_mode))
def _atomic_write(self, path: Path, content: bytes, mode: int = 0o640) -> None:
self._validate_target(path)
path.parent.mkdir(parents=True, exist_ok=True, mode=0o750)
# Re-check after creating the parent to narrow symlink races.
self._validate_target(path)
descriptor, temporary_name = tempfile.mkstemp(prefix=f".{path.name}.", dir=path.parent)
temporary = Path(temporary_name)
try:
os.fchmod(descriptor, mode)
with os.fdopen(descriptor, "wb", closefd=True) as handle:
handle.write(content)
handle.flush()
os.fsync(handle.fileno())
descriptor = -1
os.replace(temporary, path)
if os.name == "posix":
directory_descriptor = os.open(path.parent, os.O_RDONLY)
try:
os.fsync(directory_descriptor)
finally:
os.close(directory_descriptor)
finally:
if descriptor >= 0:
os.close(descriptor)
if temporary.exists():
temporary.unlink()
def backup(self, operation_id: str) -> FileSnapshot:
require_uuid(operation_id, "operation_id")
root = self.config.agent.backup_dir
self._validate_backup_root(root)
root.mkdir(parents=True, exist_ok=True, mode=0o700)
operation_dir = root / operation_id
operation_dir.mkdir(mode=0o700)
previous = tuple(
self._read_previous(path)
for path in (self.config.paths.include_path, self.config.paths.requirements_path)
)
for item in previous:
if item.content is None:
(operation_dir / f"{item.path.name}.absent").touch(mode=0o600, exist_ok=False)
else:
backup_path = operation_dir / item.path.name
descriptor = os.open(
backup_path,
os.O_CREAT | os.O_EXCL | os.O_WRONLY | getattr(os, "O_NOFOLLOW", 0),
0o600,
)
try:
os.write(descriptor, item.content)
os.fsync(descriptor)
finally:
os.close(descriptor)
return FileSnapshot(previous)
def restore(self, snapshot: FileSnapshot) -> None:
try:
for previous in snapshot.files:
if previous.content is None:
self._validate_target(previous.path)
if previous.path.exists():
previous.path.unlink()
else:
self._atomic_write(previous.path, previous.content, previous.mode)
except OSError as exc:
raise ManualRecoveryRequired("managed-file rollback failed") from exc
@staticmethod
def render_include(plugins: Iterable[ManagedPlugin]) -> bytes:
imports = sorted({plugin.import_name for plugin in plugins if plugin.enabled})
if any(not IMPORT_NAME_RE.fullmatch(name) for name in imports):
raise ValidationError("managed plugin state contains an invalid import name")
values = json.dumps(imports, ensure_ascii=True, indent=2)
text = (
"# Generated by netbox-store-agent. Do not edit.\n"
f"STORE_PLUGINS = {values}\n"
)
return text.encode("utf-8")
def render_requirements(self, plugins: Iterable[ManagedPlugin]) -> bytes:
requirements: dict[str, dict[str, object]] = {}
schemes = {"https", "http"} if self.config.store.allow_http_for_testing else {"https"}
for plugin in plugins:
for raw in plugin.requirements:
if not isinstance(raw, dict):
raise ValidationError("managed requirement must be an object")
required = {
"package_name",
"version",
"download_url",
"filename",
"sha256",
"size",
}
if set(raw) != required:
raise ValidationError("managed requirement has an invalid schema")
package = raw["package_name"]
url = raw["download_url"]
digest = raw["sha256"]
if not isinstance(package, str) or not DIST_NAME_RE.fullmatch(package):
raise ValidationError("managed requirement package is invalid")
if not isinstance(url, str) or any(char in url for char in "\r\n\t "):
raise ValidationError("managed requirement URL is invalid")
parsed = urlsplit(url)
if (
parsed.scheme not in schemes
or not parsed.hostname
or parsed.hostname.lower() not in self.config.store.allowed_hosts
or parsed.username
or parsed.password
or parsed.fragment
):
raise ValidationError("managed requirement URL violates Store policy")
if not isinstance(digest, str) or not SHA256_RE.fullmatch(digest):
raise ValidationError("managed requirement digest is invalid")
key = canonicalize_name(package)
existing = requirements.get(key)
if existing is not None and existing != raw:
raise PolicyError(f"conflicting locked requirements for {package}")
requirements[key] = raw
lines = ["# Generated by netbox-store-agent. Do not edit."]
for key in sorted(requirements):
item = requirements[key]
lines.append(
f"{item['package_name']} @ {item['download_url']} "
f"--hash=sha256:{item['sha256']}"
)
return ("\n".join(lines) + "\n").encode("utf-8")
def write(self, operation_id: str, plugins: Iterable[ManagedPlugin]) -> FileSnapshot:
plugin_list = tuple(plugins)
include = self.render_include(plugin_list)
requirements = self.render_requirements(plugin_list)
snapshot = self.backup(operation_id)
try:
self._atomic_write(self.config.paths.include_path, include)
self._atomic_write(self.config.paths.requirements_path, requirements)
except Exception:
self.restore(snapshot)
raise
return snapshot
def make_local_requirements(self, directory: Path, plugin: ManagedPlugin, wheel: Path) -> Path:
requirement = plugin.requirements[0]
path = directory / "install-requirements.txt"
content = (
f"{plugin.package_name} @ {wheel.as_uri()} "
f"--hash=sha256:{requirement['sha256']}\n"
).encode("utf-8")
descriptor = os.open(
path,
os.O_CREAT | os.O_EXCL | os.O_WRONLY | getattr(os, "O_NOFOLLOW", 0),
0o600,
)
try:
os.write(descriptor, content)
os.fsync(descriptor)
finally:
os.close(descriptor)
return path
@@ -0,0 +1,171 @@
from __future__ import annotations
import json
import re
from dataclasses import dataclass
from typing import Any
from .constants import ACTIONS, DEFAULT_MAX_REQUEST_BYTES, PROTOCOL_VERSION
from .errors import AgentError, ValidationError
from .util import SHA256_RE, SLUG_RE, reject_unknown, require_mapping, require_uuid
STATUS_PATH_RE = re.compile(r"^/v1/operations/([0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12})$")
@dataclass(frozen=True)
class OperationRequest:
request_id: str
action: str
plugin_slug: str
version: str | None
approved_payload_sha256: str | None
requested_by: str
def as_dict(self) -> dict[str, Any]:
return {
"request_id": self.request_id,
"action": self.action,
"plugin_slug": self.plugin_slug,
"version": self.version,
"approved_payload_sha256": self.approved_payload_sha256,
"requested_by": self.requested_by,
}
@dataclass(frozen=True)
class Request:
protocol_version: int
method: str
path: str
idempotency_key: str | None = None
operation: OperationRequest | None = None
def _parse_operation_body(value: Any) -> OperationRequest:
body = require_mapping(value, "body")
allowed = {
"request_id",
"action",
"plugin_slug",
"version",
"approved_payload_sha256",
"requested_by",
}
reject_unknown(body, allowed, "body")
missing = sorted(allowed - set(body))
if missing:
raise ValidationError(f"body is missing fields: {', '.join(missing)}")
request_id = require_uuid(body["request_id"], "body.request_id")
action = body["action"]
if not isinstance(action, str) or action not in ACTIONS:
raise ValidationError(f"body.action must be one of: {', '.join(sorted(ACTIONS))}")
slug = body["plugin_slug"]
if not isinstance(slug, str) or not SLUG_RE.fullmatch(slug):
raise ValidationError("body.plugin_slug is invalid")
version = body["version"]
if version is not None and (
not isinstance(version, str)
or not 1 <= len(version) <= 100
or any(char.isspace() for char in version)
or "\x00" in version
):
raise ValidationError("body.version is invalid")
digest = body["approved_payload_sha256"]
# Semantically opaque: the agent never recomputes this Store approval
# marker. Protocol v1 nevertheless fixes its wire encoding to lowercase
# SHA-256 so malformed values are rejected before persistence.
if digest is not None and (not isinstance(digest, str) or not SHA256_RE.fullmatch(digest)):
raise ValidationError("body.approved_payload_sha256 must be lowercase SHA-256 or null")
requested_by = body["requested_by"]
if not isinstance(requested_by, str) or not 1 <= len(requested_by) <= 200:
raise ValidationError("body.requested_by must be a non-empty string up to 200 characters")
if not requested_by.isprintable():
raise ValidationError("body.requested_by contains control characters")
if action in {"install", "update"}:
if version is None or digest is None:
raise ValidationError(f"{action} requires version and approved_payload_sha256")
elif version is not None or digest is not None:
raise ValidationError(f"{action} requires version and approved_payload_sha256 to be null")
return OperationRequest(
request_id=request_id,
action=action,
plugin_slug=slug,
version=version,
approved_payload_sha256=digest,
requested_by=requested_by,
)
def parse_request(value: Any) -> Request:
data = require_mapping(value, "request")
base_allowed = {"protocol_version", "method", "path"}
if data.get("method") == "POST" and data.get("path") == "/v1/operations":
allowed = base_allowed | {"idempotency_key", "body"}
else:
allowed = base_allowed
reject_unknown(data, allowed, "request")
missing = sorted(base_allowed - set(data))
if missing:
raise ValidationError(f"request is missing fields: {', '.join(missing)}")
if data["protocol_version"] != PROTOCOL_VERSION:
raise ValidationError(f"unsupported protocol_version; expected {PROTOCOL_VERSION}")
method = data["method"]
path = data["path"]
if method not in {"GET", "POST"} or not isinstance(path, str):
raise ValidationError("invalid method or path")
if method == "GET" and path == "/v1/capabilities":
return Request(PROTOCOL_VERSION, method, path)
if method == "POST" and path == "/v1/operations":
if "idempotency_key" not in data or "body" not in data:
raise ValidationError("submit requires idempotency_key and body")
key = require_uuid(data["idempotency_key"], "idempotency_key")
return Request(
PROTOCOL_VERSION,
method,
path,
idempotency_key=key,
operation=_parse_operation_body(data["body"]),
)
match = STATUS_PATH_RE.fullmatch(path) if method == "GET" else None
if match:
operation_id = require_uuid(match.group(1), "operation id")
return Request(PROTOCOL_VERSION, method, path, idempotency_key=operation_id)
raise ValidationError("unknown method/path")
def decode_request(raw: bytes, max_bytes: int) -> Request:
if not raw or len(raw) > max_bytes:
raise ValidationError(f"request must contain 1 to {max_bytes} bytes")
try:
text = raw.decode("utf-8")
except UnicodeDecodeError as exc:
raise ValidationError("request must be UTF-8") from exc
if "\n" in text or "\r" in text:
raise ValidationError("request frame must contain exactly one JSON line")
try:
value = json.loads(text)
except json.JSONDecodeError as exc:
raise ValidationError("request contains invalid JSON") from exc
return parse_request(value)
def response(status: int, body: dict[str, Any]) -> dict[str, Any]:
return {"protocol_version": PROTOCOL_VERSION, "status": status, "body": body}
def error_response(error: AgentError) -> dict[str, Any]:
return response(error.status, {"error": {"code": error.code, "message": error.message}})
def encode_response(value: dict[str, Any]) -> bytes:
encoded = (json.dumps(value, separators=(",", ":"), ensure_ascii=False) + "\n").encode("utf-8")
if len(encoded) <= DEFAULT_MAX_REQUEST_BYTES:
return encoded
fallback = error_response(
AgentError("response exceeded protocol limit", code="response_too_large", status=500)
)
return (json.dumps(fallback, separators=(",", ":")) + "\n").encode("utf-8")
@@ -0,0 +1,70 @@
from __future__ import annotations
import os
import subprocess
from dataclasses import dataclass
from pathlib import Path
from typing import Protocol, Sequence
from .config import Config
from .errors import ExecutionError, ValidationError
@dataclass(frozen=True)
class CommandResult:
argv: tuple[str, ...]
stdout: str
stderr: str
class Runner(Protocol):
def run(self, argv: Sequence[str]) -> CommandResult: ...
class SubprocessRunner:
MAX_CAPTURE_CHARS = 16_000
def __init__(self, config: Config):
self.config = config
def run(self, argv: Sequence[str]) -> CommandResult:
command = tuple(argv)
if not command or not all(isinstance(item, str) and item for item in command):
raise ValidationError("command argv must contain non-empty strings")
if any("\x00" in item or "\r" in item or "\n" in item for item in command):
raise ValidationError("command argv contains a forbidden control character")
if not Path(command[0]).is_absolute():
raise ValidationError("command executable must be an absolute path")
environment = os.environ.copy()
for key in tuple(environment):
if key.startswith("PIP_") or key in {"PYTHONPATH", "PYTHONHOME"}:
environment.pop(key, None)
environment.update(
{
"PIP_DISABLE_PIP_VERSION_CHECK": "1",
"PIP_NO_INPUT": "1",
"PYTHONNOUSERSITE": "1",
}
)
try:
completed = subprocess.run(
command,
shell=False,
stdin=subprocess.DEVNULL,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
errors="replace",
timeout=self.config.commands.command_timeout_seconds,
check=False,
env=environment,
)
except (OSError, subprocess.TimeoutExpired) as exc:
raise ExecutionError(f"command could not complete: {command[0]}") from exc
stdout = completed.stdout[-self.MAX_CAPTURE_CHARS :]
stderr = completed.stderr[-self.MAX_CAPTURE_CHARS :]
if completed.returncode != 0:
raise ExecutionError(
f"command failed with exit code {completed.returncode}: {command[0]}"
)
return CommandResult(command, stdout, stderr)
@@ -0,0 +1,56 @@
from __future__ import annotations
from concurrent.futures import ThreadPoolExecutor
from typing import Any
from .config import Config
from .constants import ACTIONS, AGENT_VERSION, PROTOCOL_VERSION
from .executor import OperationProcessor
from .journal import Journal
from .protocol import Request, response
class AgentService:
def __init__(
self,
config: Config,
journal: Journal,
processor: OperationProcessor | None = None,
):
self.config = config
self.journal = journal
self.processor = processor or OperationProcessor(config, journal)
# Lifecycle mutations are deliberately serialized in-process; the
# root-owned file lock also protects against a second process.
self.pool = ThreadPoolExecutor(max_workers=1, thread_name_prefix="plugin-operation")
self.journal.recover_interrupted()
for operation_id in self.journal.queued_operations():
self.pool.submit(self.processor.process, operation_id)
def capabilities(self) -> dict[str, Any]:
return {
"agent_version": AGENT_VERSION,
"protocol_version": PROTOCOL_VERSION,
"supported_actions": sorted(ACTIONS),
"dry_run": self.config.agent.dry_run,
"netbox_version": self.config.policy.netbox_version,
"min_supported_netbox": self.config.policy.min_supported_netbox,
"max_supported_netbox": self.config.policy.max_supported_netbox,
"max_request_bytes": self.config.agent.max_request_bytes,
"execution": "serialized-host-lifecycle",
}
def handle(self, request: Request) -> dict[str, Any]:
if request.method == "GET" and request.path == "/v1/capabilities":
return response(200, self.capabilities())
if request.method == "POST" and request.path == "/v1/operations":
assert request.idempotency_key is not None and request.operation is not None
operation, created = self.journal.submit(request.idempotency_key, request.operation)
if created:
self.pool.submit(self.processor.process, request.idempotency_key)
return response(202, {"created": created, "operation": operation})
assert request.idempotency_key is not None
return response(200, {"operation": self.journal.get_operation(request.idempotency_key)})
def close(self) -> None:
self.pool.shutdown(wait=True, cancel_futures=False)
+74
View File
@@ -0,0 +1,74 @@
from __future__ import annotations
import hashlib
import json
import re
import uuid
from datetime import UTC, datetime
from pathlib import Path
from typing import Any
from .errors import ValidationError
SLUG_RE = re.compile(r"^[a-z0-9](?:[a-z0-9._-]{0,198}[a-z0-9])?$")
SHA256_RE = re.compile(r"^[a-f0-9]{64}$")
IMPORT_NAME_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*(?:\.[A-Za-z_][A-Za-z0-9_]*)*$")
DIST_NAME_RE = re.compile(r"^[A-Za-z0-9](?:[A-Za-z0-9._-]{0,198}[A-Za-z0-9])?$")
SERVICE_NAME_RE = re.compile(r"^[A-Za-z0-9_.@-]{1,128}$")
def utc_now() -> str:
return datetime.now(UTC).isoformat().replace("+00:00", "Z")
def canonical_json(value: Any) -> str:
return json.dumps(value, sort_keys=True, separators=(",", ":"), ensure_ascii=False)
def payload_sha256(value: Any) -> str:
return hashlib.sha256(canonical_json(value).encode("utf-8")).hexdigest()
def require_uuid(value: Any, field: str) -> str:
if not isinstance(value, str):
raise ValidationError(f"{field} must be a UUID string")
try:
parsed = uuid.UUID(value)
except (ValueError, AttributeError) as exc:
raise ValidationError(f"{field} must be a valid UUID") from exc
if str(parsed) != value:
raise ValidationError(f"{field} must use canonical UUID notation")
return str(parsed)
def reject_unknown(data: dict[str, Any], allowed: set[str], context: str) -> None:
unknown = sorted(set(data) - allowed)
if unknown:
raise ValidationError(f"{context} contains unknown fields: {', '.join(unknown)}")
def require_mapping(value: Any, context: str) -> dict[str, Any]:
if not isinstance(value, dict):
raise ValidationError(f"{context} must be a JSON object")
if not all(isinstance(k, str) for k in value):
raise ValidationError(f"{context} keys must be strings")
return value
def require_absolute_path(value: Any, field: str) -> Path:
if not isinstance(value, str) or not value:
raise ValidationError(f"{field} must be a non-empty absolute path")
path = Path(value)
if not path.is_absolute():
raise ValidationError(f"{field} must be absolute")
if "\x00" in value:
raise ValidationError(f"{field} contains a NUL byte")
return path
def is_relative_to(path: Path, parent: Path) -> bool:
try:
path.relative_to(parent)
except ValueError:
return False
return True
@@ -0,0 +1,29 @@
[Unit]
Description=Privileged NetBox Store host agent
Requires=netbox-store-agent.socket
After=network-online.target
Wants=network-online.target
[Service]
Type=simple
User=root
Group=root
ExecStart=/usr/local/bin/netbox-store-agent --config /etc/netbox-store-agent/agent.toml daemon
StateDirectory=netbox-store-agent
StateDirectoryMode=0700
RuntimeDirectory=netbox-store-agent
RuntimeDirectoryMode=0750
UMask=0077
NoNewPrivileges=true
PrivateTmp=true
ProtectHome=true
ProtectSystem=strict
ReadWritePaths=/opt/netbox /var/lib/netbox-store-agent /run/netbox-store-agent /run/lock
RestrictAddressFamilies=AF_UNIX AF_INET AF_INET6
LockPersonality=true
SystemCallArchitectures=native
Restart=on-failure
RestartSec=5s
[Install]
WantedBy=multi-user.target
@@ -0,0 +1,14 @@
[Unit]
Description=NetBox Store host-agent socket
[Socket]
ListenStream=/run/netbox-store-agent/agent.sock
SocketMode=0660
SocketUser=root
# Replace with the group of the NetBox service.
SocketGroup=netbox
DirectoryMode=0750
RemoveOnStop=true
[Install]
WantedBy=sockets.target
+215
View File
@@ -0,0 +1,215 @@
from __future__ import annotations
import hashlib
import json
import sys
import zipfile
from pathlib import Path
from typing import Any
sys.path.insert(0, str(Path(__file__).resolve().parents[1] / "src"))
from netbox_store_agent.catalog import HttpStatusError, PluginMetadata, ReleasePlan
from netbox_store_agent.config import (
AgentSettings,
CommandSettings,
Config,
PathSettings,
PolicySettings,
StoreSettings,
)
from netbox_store_agent.errors import ExecutionError
from netbox_store_agent.runner import CommandResult
def make_config(root: Path, *, dry_run: bool = True, require_peers: bool = False) -> Config:
root = root.resolve()
managed_root = root / "netbox"
managed_root.mkdir(parents=True, exist_ok=True)
state = root / "state"
return Config(
agent=AgentSettings(
socket_path=root / "agent.sock",
journal_path=state / "journal.sqlite3",
lock_path=state / "lifecycle.lock",
backup_dir=state / "backups",
dry_run=dry_run,
require_root=False,
require_peer_credentials=require_peers,
allowed_peer_uids=(0,),
allowed_peer_gids=(0,),
socket_mode=0o660,
socket_uid=0,
socket_gid=0,
max_request_bytes=65536,
connection_timeout_seconds=1,
worker_threads=1,
),
store=StoreSettings(
base_url="http://store.test",
plugin_endpoint_template="/api/v1/plugins/{plugin_slug}",
release_endpoint_template="/api/v1/plugins/{plugin_slug}/releases/{version}",
timeout_seconds=1,
max_catalog_bytes=1024 * 1024,
max_artifact_bytes=1024 * 1024,
allow_private_addresses=True,
allow_http_for_testing=True,
allowed_hosts=("store.test", "artifacts.test"),
bearer_token_file=None,
ca_file=None,
),
paths=PathSettings(
allowed_root=managed_root,
include_path=managed_root / "store_plugins.py",
requirements_path=managed_root / "store_requirements.txt",
temp_dir=state / "tmp",
),
commands=CommandSettings(
python_path=root / "bin" / "python",
manage_path=managed_root / "manage.py",
systemctl_path=root / "bin" / "systemctl",
services=("netbox", "netbox-rq"),
command_timeout_seconds=10,
),
policy=PolicySettings(
netbox_version="4.6.8",
min_supported_netbox="4.6.5",
max_supported_netbox="4.6.8",
self_plugin_slugs=("netbox-store", "netbox-plugin-store", "netbox_plugin_store"),
allow_prereleases=False,
require_release_for_enable=True,
),
)
def release_json(version: str = "1.2.3", **overrides: Any) -> dict[str, Any]:
result: dict[str, Any] = {
"version": version,
"download_url": f"http://artifacts.test/demo_plugin-{version}-py3-none-any.whl",
"sha256": "a" * 64,
"artifact_size": 100,
"commit_sha": "",
"min_netbox_version": "4.6.5",
"max_netbox_version": "4.6.8",
"published_at": None,
"approved": True,
"status": "approved",
"immutable": True,
"approved_payload_sha256": "c" * 64,
}
result.update(overrides)
return result
def plugin_json(releases: list[dict[str, Any]] | None = None, **overrides: Any) -> dict[str, Any]:
result: dict[str, Any] = {
"api_version": "v1",
"slug": "demo-plugin",
"name": "Demo",
"summary": "Summary",
"description": "Description",
"repository_url": "https://git.test/demo",
"latest_version": "1.2.3",
"package_name": "demo-plugin",
"import_name": "demo_plugin",
"min_netbox_version": "4.6.5",
"max_netbox_version": "4.6.8",
"approved": True,
"status": "approved",
"releases": releases or [],
}
result.update(overrides)
return result
class FakeTransport:
def __init__(self, responses: dict[str, Any], artifact: bytes | None = None):
self.responses = responses
self.artifact = artifact or b""
self.json_urls: list[str] = []
self.download_urls: list[str] = []
def get_json(self, url: str, max_bytes: int) -> Any:
self.json_urls.append(url)
if url not in self.responses:
raise HttpStatusError(404, "missing")
value = self.responses[url]
if isinstance(value, Exception):
raise value
assert len(json.dumps(value)) <= max_bytes
return value
def download(
self,
url: str,
destination: Path,
*,
max_bytes: int,
expected_size: int,
expected_sha256: str,
) -> None:
self.download_urls.append(url)
if len(self.artifact) != expected_size:
raise AssertionError("fixture size mismatch")
if hashlib.sha256(self.artifact).hexdigest() != expected_sha256:
raise AssertionError("fixture digest mismatch")
destination.write_bytes(self.artifact)
def wheel_bytes(path: Path, version: str = "1.2.3") -> bytes:
wheel = path / f"demo_plugin-{version}-py3-none-any.whl"
with zipfile.ZipFile(wheel, "w") as archive:
archive.writestr("demo_plugin/__init__.py", "")
archive.writestr(
f"demo_plugin-{version}.dist-info/WHEEL",
"Wheel-Version: 1.0\nGenerator: tests\nRoot-Is-Purelib: true\nTag: py3-none-any\n",
)
return wheel.read_bytes()
def plan(version: str = "1.2.3") -> ReleasePlan:
plugin = PluginMetadata(
"demo-plugin", "demo-plugin", "demo_plugin", "4.6.5", "4.6.8", ()
)
return ReleasePlan(
plugin,
version,
f"http://artifacts.test/demo_plugin-{version}-py3-none-any.whl",
f"demo_plugin-{version}-py3-none-any.whl",
"b" * 64,
1,
"c" * 64,
)
class FakeStore:
def __init__(self, release: ReleasePlan | None = None):
self.release = release or plan()
self.calls: list[tuple[str, ...]] = []
def get_plugin(self, slug: str) -> PluginMetadata:
self.calls.append(("plugin", slug))
return self.release.plugin
def get_release(self, slug: str, version: str) -> ReleasePlan:
self.calls.append(("release", slug, version))
return self.release
def download_release(self, release: ReleasePlan, directory: Path) -> Path:
self.calls.append(("download", release.version))
target = directory / release.filename
target.write_bytes(b"x")
return target
class FakeRunner:
def __init__(self, fail_step: str | None = None):
self.commands: list[tuple[str, ...]] = []
self.fail_step = fail_step
def run(self, argv: list[str]) -> CommandResult:
command = tuple(argv)
self.commands.append(command)
if self.fail_step and self.fail_step in command:
raise ExecutionError("injected command failure")
return CommandResult(command, "", "")
+124
View File
@@ -0,0 +1,124 @@
from __future__ import annotations
import hashlib
import tempfile
import unittest
from pathlib import Path
from support import FakeTransport, make_config, plugin_json, release_json, wheel_bytes
from netbox_store_agent.catalog import CatalogError, SecureHTTPTransport, StoreClient
from netbox_store_agent.config import StoreSettings
class CatalogTests(unittest.TestCase):
def setUp(self) -> None:
self.temporary = tempfile.TemporaryDirectory()
self.root = Path(self.temporary.name)
self.config = make_config(self.root)
def tearDown(self) -> None:
self.temporary.cleanup()
def test_trailing_slash_fallback_and_valid_wheel(self) -> None:
artifact = wheel_bytes(self.root)
release = release_json(
sha256=hashlib.sha256(artifact).hexdigest(), artifact_size=len(artifact)
)
responses = {
"http://store.test/api/v1/plugins/demo-plugin/": plugin_json(),
"http://store.test/api/v1/plugins/demo-plugin/releases/1.2.3/": release,
}
transport = FakeTransport(responses, artifact)
client = StoreClient(self.config, transport)
plan = client.get_release("demo-plugin", "1.2.3")
# Destination directory is caller-owned; use an existing operation directory.
operation_dir = self.root / "operation"
operation_dir.mkdir()
downloaded = client.download_release(plan, operation_dir)
self.assertTrue(downloaded.is_file())
self.assertEqual(transport.json_urls[0][-1], "n")
self.assertEqual(transport.json_urls[1][-1], "/")
def test_release_falls_back_to_plugin_releases_array(self) -> None:
release = release_json()
transport = FakeTransport(
{"http://store.test/api/v1/plugins/demo-plugin": plugin_json([release])}
)
plan = StoreClient(self.config, transport).get_release("demo-plugin", "1.2.3")
self.assertEqual(plan.version, "1.2.3")
def test_unapproved_or_mutable_release_rejected(self) -> None:
for change in ({"approved": False}, {"immutable": False}, {"status": "pending"}):
with self.subTest(change=change):
transport = FakeTransport(
{
"http://store.test/api/v1/plugins/demo-plugin": plugin_json(),
"http://store.test/api/v1/plugins/demo-plugin/releases/1.2.3": release_json(
**change
),
}
)
with self.assertRaises(CatalogError):
StoreClient(self.config, transport).get_release("demo-plugin", "1.2.3")
def test_unknown_catalog_field_fails_closed(self) -> None:
payload = plugin_json()
payload["internal_id"] = 7
transport = FakeTransport({"http://store.test/api/v1/plugins/demo-plugin": payload})
with self.assertRaises(CatalogError):
StoreClient(self.config, transport).get_plugin("demo-plugin")
def test_netbox_incompatibility_rejected(self) -> None:
transport = FakeTransport(
{
"http://store.test/api/v1/plugins/demo-plugin": plugin_json(
min_netbox_version="4.7.0"
)
}
)
with self.assertRaises(CatalogError):
StoreClient(self.config, transport).get_plugin("demo-plugin")
def test_wheel_distribution_must_match(self) -> None:
artifact = wheel_bytes(self.root)
release = release_json(
sha256=hashlib.sha256(artifact).hexdigest(), artifact_size=len(artifact)
)
payload = plugin_json(package_name="another-package")
transport = FakeTransport(
{
"http://store.test/api/v1/plugins/demo-plugin": payload,
"http://store.test/api/v1/plugins/demo-plugin/releases/1.2.3": release,
},
artifact,
)
client = StoreClient(self.config, transport)
plan = client.get_release("demo-plugin", "1.2.3")
directory = self.root / "mismatch"
directory.mkdir()
with self.assertRaises(CatalogError):
client.download_release(plan, directory)
def test_bearer_token_is_not_sent_to_artifact_origin(self) -> None:
settings = StoreSettings(
**{
**self.config.store.__dict__,
"base_url": "https://store.test:8443",
"allow_http_for_testing": False,
}
)
transport = SecureHTTPTransport(settings)
transport._token = lambda: "secret" # type: ignore[method-assign]
self.assertEqual(
transport._headers("https://store.test:8443/api")["Authorization"],
"Bearer secret",
)
self.assertNotIn(
"Authorization", transport._headers("https://artifacts.test/release.whl")
)
self.assertNotIn("Authorization", transport._headers("https://store.test/api"))
if __name__ == "__main__":
unittest.main()
+132
View File
@@ -0,0 +1,132 @@
from __future__ import annotations
import json
import tempfile
import unittest
from pathlib import Path
from support import * # noqa: F403
from netbox_store_agent.config import load_config
from netbox_store_agent.errors import ValidationError
class ConfigTests(unittest.TestCase):
def setUp(self) -> None:
self.temporary = tempfile.TemporaryDirectory()
self.root = Path(self.temporary.name).resolve()
def tearDown(self) -> None:
self.temporary.cleanup()
@staticmethod
def q(path: Path) -> str:
return json.dumps(path.as_posix())
def config_text(self, **changes: str) -> str:
root = self.root
values = {
"agent_extra": "",
"store_extra": "",
"paths_extra": "",
"commands_extra": "",
"policy_extra": "",
"base_url": '"https://store.test"',
"include": self.q(root / "netbox" / "store_plugins.py"),
}
values.update(changes)
return f"""
[agent]
socket_path = {self.q(root / 'agent.sock')}
journal_path = {self.q(root / 'state' / 'journal.sqlite3')}
lock_path = {self.q(root / 'state' / 'lock')}
backup_dir = {self.q(root / 'state' / 'backups')}
require_root = false
require_peer_credentials = false
{values['agent_extra']}
[store]
base_url = {values['base_url']}
allowed_hosts = ["store.test"]
{values['store_extra']}
[paths]
allowed_root = {self.q(root / 'netbox')}
include_path = {values['include']}
requirements_path = {self.q(root / 'netbox' / 'requirements.txt')}
temp_dir = {self.q(root / 'state' / 'tmp')}
{values['paths_extra']}
[commands]
python_path = {self.q(root / 'bin' / 'python')}
manage_path = {self.q(root / 'netbox' / 'manage.py')}
systemctl_path = {self.q(root / 'bin' / 'systemctl')}
{values['commands_extra']}
[policy]
{values['policy_extra']}
"""
def write(self, text: str) -> Path:
path = self.root / "agent.toml"
path.write_text(text, encoding="utf-8")
path.chmod(0o600)
return path
def test_defaults_are_dry_run_and_block_all_self_slugs(self) -> None:
config = load_config(self.write(self.config_text()), allow_insecure_owner=True)
self.assertTrue(config.agent.dry_run)
self.assertIn("netbox-plugin-store", config.policy.self_plugin_slugs)
self.assertIn("netbox_plugin_store", config.policy.self_plugin_slugs)
def test_unknown_setting_rejected(self) -> None:
with self.assertRaises(ValidationError):
load_config(
self.write(self.config_text(agent_extra="surprise = true")),
allow_insecure_owner=True,
)
def test_http_requires_explicit_testing_switch(self) -> None:
with self.assertRaises(ValidationError):
load_config(
self.write(self.config_text(base_url='"http://store.test"')),
allow_insecure_owner=True,
)
config = load_config(
self.write(
self.config_text(
base_url='"http://store.test"', store_extra="allow_http_for_testing = true"
)
),
allow_insecure_owner=True,
)
self.assertTrue(config.store.allow_http_for_testing)
def test_managed_path_escape_rejected(self) -> None:
with self.assertRaises(ValidationError):
load_config(
self.write(self.config_text(include=self.q(self.root / "outside.py"))),
allow_insecure_owner=True,
)
def test_protocol_size_cannot_exceed_64k(self) -> None:
with self.assertRaises(ValidationError):
load_config(
self.write(self.config_text(agent_extra="max_request_bytes = 65537")),
allow_insecure_owner=True,
)
def test_endpoint_placeholders_are_exact(self) -> None:
with self.assertRaises(ValidationError):
load_config(
self.write(
self.config_text(
store_extra='release_endpoint_template = "/api/{plugin_slug}/{other}"'
)
),
allow_insecure_owner=True,
)
if __name__ == "__main__":
unittest.main()
+80
View File
@@ -0,0 +1,80 @@
from __future__ import annotations
import json
import socket
import tempfile
import unittest
from pathlib import Path
from support import make_config
from netbox_store_agent.daemon import handle_connection
from netbox_store_agent.protocol import response
class FakeService:
def __init__(self) -> None:
self.requests = []
def handle(self, request: object) -> dict[str, object]:
self.requests.append(request)
return response(200, {"ok": True})
class DaemonFramingTests(unittest.TestCase):
def setUp(self) -> None:
self.temporary = tempfile.TemporaryDirectory()
self.config = make_config(Path(self.temporary.name), require_peers=False)
def tearDown(self) -> None:
self.temporary.cleanup()
def exchange(self, payload: bytes) -> dict[str, object]:
server, client = socket.socketpair()
try:
client.sendall(payload)
service = FakeService()
handle_connection(server, self.config, service) # type: ignore[arg-type]
data = bytearray()
while b"\n" not in data:
chunk = client.recv(4096)
if not chunk:
break
data.extend(chunk)
self.assertEqual(data.count(b"\n"), 1)
return json.loads(bytes(data).decode())
finally:
client.close()
def test_one_request_one_response(self) -> None:
result = self.exchange(
b'{"protocol_version":1,"method":"GET","path":"/v1/capabilities"}\n'
)
self.assertEqual(result["status"], 200)
self.assertEqual(set(result), {"protocol_version", "status", "body"})
def test_second_json_line_rejected(self) -> None:
result = self.exchange(
b'{"protocol_version":1,"method":"GET","path":"/v1/capabilities"}\n{}\n'
)
self.assertEqual(result["status"], 400)
def test_missing_newline_rejected_when_client_half_closes(self) -> None:
server, client = socket.socketpair()
try:
client.sendall(b"{}")
client.shutdown(socket.SHUT_WR)
service = FakeService()
handle_connection(server, self.config, service) # type: ignore[arg-type]
result = json.loads(client.recv(4096).decode())
self.assertEqual(result["status"], 400)
finally:
client.close()
def test_oversized_frame_rejected(self) -> None:
result = self.exchange(b"x" * 65536 + b"\n")
self.assertEqual(result["status"], 400)
if __name__ == "__main__":
unittest.main()
+177
View File
@@ -0,0 +1,177 @@
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from support import FakeRunner, FakeStore, make_config, plan
from netbox_store_agent.executor import OperationProcessor
from netbox_store_agent.journal import Journal, ManagedPlugin
from netbox_store_agent.protocol import OperationRequest
class ExecutorTests(unittest.TestCase):
def setUp(self) -> None:
self.temporary = tempfile.TemporaryDirectory()
self.root = Path(self.temporary.name)
self.key = "20f4274f-d4e5-42bf-9164-967b1a774481"
def tearDown(self) -> None:
self.temporary.cleanup()
def request(self, action: str, *, version: str | None = None, token: str | None = None) -> OperationRequest:
return OperationRequest(
"eea17d87-8944-4ee2-a076-363338ab746d",
action,
"demo-plugin",
version,
token,
"alice",
)
def seed(self, journal: Journal, *, enabled: bool) -> ManagedPlugin:
release = plan()
plugin = ManagedPlugin(
"demo-plugin",
"demo-plugin",
"demo_plugin",
"1.2.3",
enabled,
(release.requirement(),),
)
journal.upsert_managed_plugin(plugin)
return plugin
def test_dry_run_verifies_artifact_without_runner_or_state(self) -> None:
config = make_config(self.root, dry_run=True)
journal = Journal(config.agent.journal_path)
store = FakeStore()
runner = FakeRunner()
journal.submit(self.key, self.request("install", version="1.2.3", token="c" * 64))
OperationProcessor(config, journal, store=store, runner=runner).process(self.key)
operation = journal.get_operation(self.key)
self.assertEqual(operation["state"], "dry_run")
self.assertIn(("download", "1.2.3"), store.calls)
self.assertEqual(runner.commands, [])
self.assertIsNone(journal.get_managed_plugin("demo-plugin"))
self.assertFalse(config.paths.include_path.exists())
def test_install_is_disabled_and_uses_fixed_pip_argv(self) -> None:
config = make_config(self.root, dry_run=False)
journal = Journal(config.agent.journal_path)
runner = FakeRunner()
journal.submit(self.key, self.request("install", version="1.2.3", token="c" * 64))
OperationProcessor(config, journal, store=FakeStore(), runner=runner).process(self.key)
operation = journal.get_operation(self.key)
self.assertEqual(operation["state"], "succeeded")
managed = journal.get_managed_plugin("demo-plugin")
self.assertIsNotNone(managed)
self.assertFalse(managed.enabled)
self.assertEqual(len(runner.commands), 1)
command = runner.commands[0]
self.assertIn("--no-index", command)
self.assertIn("--no-deps", command)
self.assertIn("--require-hashes", command)
def test_approval_token_mismatch_fails_before_download(self) -> None:
config = make_config(self.root, dry_run=False)
journal = Journal(config.agent.journal_path)
store = FakeStore()
runner = FakeRunner()
journal.submit(self.key, self.request("install", version="1.2.3", token="d" * 64))
OperationProcessor(config, journal, store=store, runner=runner).process(self.key)
self.assertEqual(journal.get_operation(self.key)["state"], "failed")
self.assertFalse(any(call[0] == "download" for call in store.calls))
self.assertEqual(runner.commands, [])
def test_enable_runs_migrate_collectstatic_and_both_service_restart(self) -> None:
config = make_config(self.root, dry_run=False)
journal = Journal(config.agent.journal_path)
self.seed(journal, enabled=False)
runner = FakeRunner()
journal.submit(self.key, self.request("enable"))
OperationProcessor(config, journal, store=FakeStore(), runner=runner).process(self.key)
self.assertEqual(journal.get_operation(self.key)["state"], "succeeded")
self.assertTrue(journal.get_managed_plugin("demo-plugin").enabled)
flattened = [item for command in runner.commands for item in command]
self.assertIn("migrate", flattened)
self.assertIn("collectstatic", flattened)
self.assertIn("netbox", flattened)
self.assertIn("netbox-rq", flattened)
def test_uninstall_requires_disabled(self) -> None:
config = make_config(self.root, dry_run=False)
journal = Journal(config.agent.journal_path)
self.seed(journal, enabled=True)
runner = FakeRunner()
journal.submit(self.key, self.request("uninstall"))
OperationProcessor(config, journal, store=FakeStore(), runner=runner).process(self.key)
self.assertEqual(journal.get_operation(self.key)["state"], "failed")
self.assertEqual(runner.commands, [])
def test_disable_restarts_services_and_persists_disabled_state(self) -> None:
config = make_config(self.root, dry_run=False)
journal = Journal(config.agent.journal_path)
self.seed(journal, enabled=True)
runner = FakeRunner()
journal.submit(self.key, self.request("disable"))
OperationProcessor(config, journal, store=FakeStore(), runner=runner).process(self.key)
self.assertEqual(journal.get_operation(self.key)["state"], "succeeded")
self.assertFalse(journal.get_managed_plugin("demo-plugin").enabled)
self.assertEqual(len(runner.commands), 1)
self.assertIn("restart", runner.commands[0])
def test_uninstall_disabled_plugin_uses_fixed_package_and_deletes_state(self) -> None:
config = make_config(self.root, dry_run=False)
journal = Journal(config.agent.journal_path)
self.seed(journal, enabled=False)
runner = FakeRunner()
journal.submit(self.key, self.request("uninstall"))
OperationProcessor(config, journal, store=FakeStore(), runner=runner).process(self.key)
self.assertEqual(journal.get_operation(self.key)["state"], "succeeded")
self.assertIsNone(journal.get_managed_plugin("demo-plugin"))
self.assertEqual(len(runner.commands), 1)
self.assertEqual(runner.commands[0][-1], "demo-plugin")
def test_failed_pip_attempt_is_conservatively_manual_recovery(self) -> None:
config = make_config(self.root, dry_run=False)
journal = Journal(config.agent.journal_path)
runner = FakeRunner(fail_step="install")
journal.submit(self.key, self.request("install", version="1.2.3", token="c" * 64))
OperationProcessor(config, journal, store=FakeStore(), runner=runner).process(self.key)
self.assertEqual(journal.get_operation(self.key)["state"], "manual_recovery")
def test_failure_after_pip_requires_manual_recovery_and_restores_files(self) -> None:
config = make_config(self.root, dry_run=False)
config.paths.include_path.write_text("old include", encoding="utf-8")
config.paths.requirements_path.write_text("old requirements", encoding="utf-8")
journal = Journal(config.agent.journal_path)
self.seed(journal, enabled=True)
runner = FakeRunner(fail_step="restart")
journal.submit(self.key, self.request("update", version="1.2.3", token="c" * 64))
OperationProcessor(config, journal, store=FakeStore(), runner=runner).process(self.key)
self.assertEqual(journal.get_operation(self.key)["state"], "manual_recovery")
self.assertEqual(config.paths.include_path.read_text(), "old include")
self.assertEqual(config.paths.requirements_path.read_text(), "old requirements")
def test_self_management_fails_before_store_access(self) -> None:
config = make_config(self.root, dry_run=True)
journal = Journal(config.agent.journal_path)
store = FakeStore()
request = OperationRequest(
"eea17d87-8944-4ee2-a076-363338ab746d",
"install",
"netbox-plugin-store",
"1.2.3",
"c" * 64,
"alice",
)
journal.submit(self.key, request)
OperationProcessor(config, journal, store=store, runner=FakeRunner()).process(self.key)
self.assertEqual(journal.get_operation(self.key)["state"], "failed")
self.assertEqual(store.calls, [])
if __name__ == "__main__":
unittest.main()
+72
View File
@@ -0,0 +1,72 @@
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from support import * # noqa: F403
from netbox_store_agent.errors import ConflictError
from netbox_store_agent.journal import Journal, ManagedPlugin
from netbox_store_agent.protocol import OperationRequest
class JournalTests(unittest.TestCase):
def setUp(self) -> None:
self.temporary = tempfile.TemporaryDirectory()
self.journal = Journal(Path(self.temporary.name) / "journal.sqlite3")
self.key = "20f4274f-d4e5-42bf-9164-967b1a774481"
self.request = OperationRequest(
"eea17d87-8944-4ee2-a076-363338ab746d",
"install",
"demo-plugin",
"1.2.3",
"opaque",
"alice",
)
def tearDown(self) -> None:
self.temporary.cleanup()
def test_idempotency_same_payload_returns_existing(self) -> None:
first, created = self.journal.submit(self.key, self.request)
second, created_again = self.journal.submit(self.key, self.request)
self.assertTrue(created)
self.assertFalse(created_again)
self.assertEqual(first["operation_id"], second["operation_id"])
def test_idempotency_conflict(self) -> None:
self.journal.submit(self.key, self.request)
changed = OperationRequest(**{**self.request.as_dict(), "requested_by": "mallory"})
with self.assertRaises(ConflictError):
self.journal.submit(self.key, changed)
def test_claim_is_exactly_once(self) -> None:
self.journal.submit(self.key, self.request)
self.assertTrue(self.journal.claim(self.key))
self.assertFalse(self.journal.claim(self.key))
def test_running_recovery_is_manual(self) -> None:
self.journal.submit(self.key, self.request)
self.journal.claim(self.key)
self.assertEqual(self.journal.recover_interrupted(), 1)
operation = self.journal.get_operation(self.key)
self.assertEqual(operation["state"], "manual_recovery")
def test_managed_plugin_round_trip(self) -> None:
plugin = ManagedPlugin(
"demo-plugin",
"demo-plugin",
"demo_plugin",
"1.2.3",
False,
({"sha256": "a" * 64},),
)
self.journal.upsert_managed_plugin(plugin)
self.assertEqual(self.journal.get_managed_plugin("demo-plugin"), plugin)
self.journal.delete_managed_plugin("demo-plugin")
self.assertIsNone(self.journal.get_managed_plugin("demo-plugin"))
if __name__ == "__main__":
unittest.main()
+87
View File
@@ -0,0 +1,87 @@
from __future__ import annotations
import tempfile
import unittest
from pathlib import Path
from support import make_config, plan
from netbox_store_agent.errors import PolicyError
from netbox_store_agent.journal import ManagedPlugin
from netbox_store_agent.managed_files import ManagedFiles
class ManagedFilesTests(unittest.TestCase):
def setUp(self) -> None:
self.temporary = tempfile.TemporaryDirectory()
self.root = Path(self.temporary.name)
self.config = make_config(self.root, dry_run=False)
self.files = ManagedFiles(self.config)
release = plan()
self.plugin = ManagedPlugin(
"demo-plugin",
"demo-plugin",
"demo_plugin",
"1.2.3",
True,
(release.requirement(),),
)
def tearDown(self) -> None:
self.temporary.cleanup()
def test_include_exports_constant_without_referencing_plugins(self) -> None:
text = self.files.render_include([self.plugin]).decode()
self.assertIn('STORE_PLUGINS = [\n "demo_plugin"\n]', text)
self.assertNotIn("PLUGINS = list", text)
def test_atomic_write_backup_and_restore(self) -> None:
self.config.paths.include_path.write_text("old include", encoding="utf-8")
self.config.paths.requirements_path.write_text("old requirements", encoding="utf-8")
snapshot = self.files.write(
"20f4274f-d4e5-42bf-9164-967b1a774481", [self.plugin]
)
self.assertIn("STORE_PLUGINS", self.config.paths.include_path.read_text())
self.assertIn("--hash=sha256:", self.config.paths.requirements_path.read_text())
self.files.restore(snapshot)
self.assertEqual(self.config.paths.include_path.read_text(), "old include")
self.assertEqual(self.config.paths.requirements_path.read_text(), "old requirements")
def test_conflicting_locked_distribution_rejected(self) -> None:
other_requirement = {**self.plugin.requirements[0], "version": "2.0.0"}
other = ManagedPlugin(
"other", "demo-plugin", "other_plugin", "2.0.0", False, (other_requirement,)
)
with self.assertRaises(PolicyError):
self.files.render_requirements([self.plugin, other])
def test_requirement_host_must_be_allowlisted(self) -> None:
requirement = {**self.plugin.requirements[0], "download_url": "http://evil.test/x.whl"}
plugin = ManagedPlugin(
"demo-plugin", "demo-plugin", "demo_plugin", "1.2.3", False, (requirement,)
)
with self.assertRaises(Exception):
self.files.render_requirements([plugin])
def test_path_escape_rejected(self) -> None:
escaped = self.root / "outside.py"
changed = self.config.__class__(
self.config.agent,
self.config.store,
self.config.paths.__class__(
self.config.paths.allowed_root,
escaped,
self.config.paths.requirements_path,
self.config.paths.temp_dir,
),
self.config.commands,
self.config.policy,
)
with self.assertRaises(PolicyError):
ManagedFiles(changed).write(
"20f4274f-d4e5-42bf-9164-967b1a774481", [self.plugin]
)
if __name__ == "__main__":
unittest.main()
+83
View File
@@ -0,0 +1,83 @@
from __future__ import annotations
import json
import unittest
from support import * # noqa: F403
from netbox_store_agent.errors import ValidationError
from netbox_store_agent.protocol import decode_request, encode_response, parse_request, response
class ProtocolTests(unittest.TestCase):
def submit(self, **body_overrides: object) -> dict[str, object]:
body: dict[str, object] = {
"request_id": "eea17d87-8944-4ee2-a076-363338ab746d",
"action": "install",
"plugin_slug": "demo-plugin",
"version": "1.2.3",
"approved_payload_sha256": "c" * 64,
"requested_by": "netbox:alice",
}
body.update(body_overrides)
return {
"protocol_version": 1,
"method": "POST",
"path": "/v1/operations",
"idempotency_key": "20f4274f-d4e5-42bf-9164-967b1a774481",
"body": body,
}
def test_exact_submit_contract_and_opaque_token(self) -> None:
request = parse_request(self.submit())
self.assertEqual(request.operation.approved_payload_sha256, "c" * 64)
def test_unknown_top_level_field_rejected(self) -> None:
value = self.submit(extra=True)
value["unexpected"] = True
with self.assertRaises(ValidationError):
parse_request(value)
def test_unknown_body_field_rejected(self) -> None:
with self.assertRaises(ValidationError):
parse_request(self.submit(extra=True))
def test_non_lifecycle_fields_must_be_null(self) -> None:
with self.assertRaises(ValidationError):
parse_request(self.submit(action="disable"))
request = parse_request(
self.submit(action="disable", version=None, approved_payload_sha256=None)
)
self.assertEqual(request.operation.action, "disable")
def test_noncanonical_uuid_rejected(self) -> None:
with self.assertRaises(ValidationError):
parse_request(self.submit(request_id="EEA17D87-8944-4EE2-A076-363338AB746D"))
def test_multiple_lines_and_oversize_rejected(self) -> None:
raw = json.dumps(self.submit()).encode()
with self.assertRaises(ValidationError):
decode_request(raw + b"\n{}", 65536)
with self.assertRaises(ValidationError):
decode_request(b"x" * 10, 9)
def test_capabilities_and_status_paths(self) -> None:
cap = parse_request({"protocol_version": 1, "method": "GET", "path": "/v1/capabilities"})
self.assertEqual(cap.path, "/v1/capabilities")
status = parse_request(
{
"protocol_version": 1,
"method": "GET",
"path": "/v1/operations/20f4274f-d4e5-42bf-9164-967b1a774481",
}
)
self.assertEqual(status.idempotency_key, "20f4274f-d4e5-42bf-9164-967b1a774481")
def test_response_shape_and_single_line(self) -> None:
encoded = encode_response(response(200, {"ok": True}))
self.assertEqual(encoded.count(b"\n"), 1)
self.assertEqual(set(json.loads(encoded)), {"protocol_version", "status", "body"})
if __name__ == "__main__":
unittest.main()