feat: add NetBox plugin store
This commit is contained in:
@@ -0,0 +1,6 @@
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.egg-info/
|
||||
.buildcheck/
|
||||
build/
|
||||
dist/
|
||||
@@ -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.
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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"})
|
||||
@@ -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)
|
||||
@@ -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
|
||||
@@ -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, "", "")
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
@@ -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()
|
||||
Reference in New Issue
Block a user