306 lines
9.9 KiB
Python
306 lines
9.9 KiB
Python
"""
|
|
Pending queue: if an agent is unreachable, actions are stored in SQLite
|
|
(table: pending_actions) and can be retried later via /pending endpoints.
|
|
|
|
broker.py — Redis mTLS message broker for NetAdmin
|
|
|
|
All connections use mutual TLS (mTLS):
|
|
- Backend presents backend.crt signed by the internal CA
|
|
- Redis server presents redis-server.crt signed by the same CA
|
|
- Both sides verify each other's certificate
|
|
|
|
Channels:
|
|
dns.commands / dns.acks
|
|
mail.commands / mail.acks
|
|
"""
|
|
|
|
import datetime
|
|
import json
|
|
import logging
|
|
import ssl
|
|
import time
|
|
import uuid
|
|
|
|
import redis
|
|
import redis.connection
|
|
|
|
from config import cfg
|
|
|
|
log = logging.getLogger(__name__)
|
|
|
|
_redis_client: redis.Redis | None = None
|
|
|
|
|
|
# ──────────────────────────────────────────────
|
|
# Pending action queue (SQLite — shared netadmin.db)
|
|
# ──────────────────────────────────────────────
|
|
|
|
from db import get_db as _queue_conn
|
|
|
|
|
|
def _init_queue_table():
|
|
try:
|
|
with _queue_conn() as conn:
|
|
conn.execute("""
|
|
CREATE TABLE IF NOT EXISTS pending_actions (
|
|
id VARCHAR(36) PRIMARY KEY,
|
|
service TEXT NOT NULL,
|
|
action TEXT NOT NULL,
|
|
payload TEXT NOT NULL,
|
|
status TEXT NOT NULL DEFAULT 'pending',
|
|
error TEXT,
|
|
created_at DATETIME NOT NULL DEFAULT NOW(),
|
|
updated_at DATETIME NOT NULL DEFAULT NOW()
|
|
)
|
|
""")
|
|
except Exception as e:
|
|
log.warning(f"[broker] queue table init: {e}")
|
|
|
|
|
|
_init_queue_table()
|
|
|
|
|
|
|
|
def queue_action(service: str, action: str, payload: dict) -> str:
|
|
"""Store a failed action for later retry. Returns the queue entry id."""
|
|
entry_id = str(uuid.uuid4())
|
|
with _queue_conn() as conn:
|
|
conn.execute(
|
|
"INSERT INTO pending_actions(id,service,action,payload,error) VALUES(%s,%s,%s,%s,%s)",
|
|
(entry_id, service, action, json.dumps(payload),
|
|
"Agent unreachable at " + datetime.datetime.utcnow().isoformat())
|
|
)
|
|
log.warning(f"[queue] Action '{action}' ({service}) queued (id={entry_id})")
|
|
return entry_id
|
|
|
|
|
|
def list_pending() -> list:
|
|
"""Return all non-done pending actions ordered by creation date."""
|
|
with _queue_conn() as conn:
|
|
rows = conn.execute(
|
|
"SELECT * FROM pending_actions WHERE status != 'done' ORDER BY created_at"
|
|
).fetchall()
|
|
return [dict(r) for r in rows]
|
|
|
|
|
|
def mark_done(entry_id: str):
|
|
with _queue_conn() as conn:
|
|
conn.execute(
|
|
"UPDATE pending_actions SET status='done', updated_at=NOW() WHERE id=%s",
|
|
(entry_id,)
|
|
)
|
|
|
|
|
|
def mark_failed(entry_id: str, error: str):
|
|
with _queue_conn() as conn:
|
|
conn.execute(
|
|
"UPDATE pending_actions SET status='failed', error=%s, updated_at=NOW() WHERE id=%s",
|
|
(error, entry_id)
|
|
)
|
|
|
|
|
|
def get_redis() -> redis.Redis:
|
|
global _redis_client
|
|
if _redis_client is not None:
|
|
return _redis_client
|
|
|
|
kwargs: dict = {
|
|
"host": cfg.redis.host,
|
|
"port": cfg.redis.port,
|
|
"db": cfg.redis.db,
|
|
"decode_responses": True,
|
|
"socket_timeout": cfg.redis.ack_timeout + 2,
|
|
"socket_connect_timeout": 5,
|
|
}
|
|
|
|
if cfg.redis.password:
|
|
kwargs["password"] = cfg.redis.password
|
|
|
|
tls = cfg.redis.tls
|
|
if tls.enabled:
|
|
kwargs["ssl"] = True
|
|
kwargs["ssl_certfile"] = tls.client_cert
|
|
kwargs["ssl_keyfile"] = tls.client_key
|
|
kwargs["ssl_ca_certs"] = tls.ca_cert
|
|
kwargs["ssl_cert_reqs"] = "required"
|
|
log.info(
|
|
f"[broker] Redis mTLS enabled "
|
|
f"(cert={tls.client_cert}, ca={tls.ca_cert})"
|
|
)
|
|
else:
|
|
log.warning("[broker] Redis TLS disabled — traffic is unencrypted")
|
|
|
|
_redis_client = redis.Redis(**kwargs)
|
|
return _redis_client
|
|
|
|
|
|
def publish_and_wait(channel: str, ack_channel: str,
|
|
action: str, payload: dict,
|
|
timeout: int | None = None) -> dict:
|
|
"""Publish a command and block until the agent ACKs or timeout.
|
|
|
|
Both the publish connection and the subscribe connection use the same
|
|
mTLS configuration.
|
|
|
|
Returns ACK dict: {"id": ..., "success": True/False, "error": "..."}
|
|
Raises RuntimeError on timeout or connection error.
|
|
"""
|
|
r = get_redis()
|
|
msg_id = str(uuid.uuid4())
|
|
timeout = timeout or cfg.redis.ack_timeout
|
|
|
|
message = json.dumps({
|
|
"id": msg_id,
|
|
"action": action,
|
|
"payload": payload,
|
|
# No token — authentication is handled by the mTLS certificate CN
|
|
})
|
|
|
|
# Build a separate subscribe connection with mTLS
|
|
sub_kwargs: dict = {
|
|
"host": cfg.redis.host,
|
|
"port": cfg.redis.port,
|
|
"db": cfg.redis.db,
|
|
"decode_responses": True,
|
|
}
|
|
if cfg.redis.password:
|
|
sub_kwargs["password"] = cfg.redis.password
|
|
tls = cfg.redis.tls
|
|
if tls.enabled:
|
|
sub_kwargs["ssl"] = True
|
|
sub_kwargs["ssl_certfile"] = tls.client_cert
|
|
sub_kwargs["ssl_keyfile"] = tls.client_key
|
|
sub_kwargs["ssl_ca_certs"] = tls.ca_cert
|
|
sub_kwargs["ssl_cert_reqs"] = "required"
|
|
|
|
sub_r = redis.Redis(**sub_kwargs)
|
|
pubsub = sub_r.pubsub()
|
|
pubsub.subscribe(ack_channel)
|
|
|
|
try:
|
|
# Drain any stale subscribe confirmation message
|
|
pubsub.get_message(timeout=0.1)
|
|
|
|
r.publish(channel, message)
|
|
log.debug(f"[broker] published '{action}' → {channel} (id={msg_id})")
|
|
|
|
deadline = time.monotonic() + timeout
|
|
while time.monotonic() < deadline:
|
|
remaining = deadline - time.monotonic()
|
|
msg = pubsub.get_message(timeout=min(remaining, 0.5))
|
|
if msg is None or msg["type"] != "message":
|
|
continue
|
|
try:
|
|
ack = json.loads(msg["data"])
|
|
except (json.JSONDecodeError, TypeError):
|
|
continue
|
|
if ack.get("id") == msg_id:
|
|
log.debug(f"[broker] ACK for {msg_id}: success={ack.get('success')}")
|
|
return ack
|
|
|
|
finally:
|
|
try:
|
|
pubsub.unsubscribe()
|
|
pubsub.close()
|
|
sub_r.close()
|
|
except Exception:
|
|
pass
|
|
|
|
raise RuntimeError(
|
|
f"Agent timeout ({timeout}s) for action '{action}'. "
|
|
"Check that the agent is running and connected to Redis."
|
|
)
|
|
|
|
|
|
def publish_or_queue(service: str, action: str, payload: dict) -> dict:
|
|
"""Try to publish and wait for ACK. If agent is unreachable, queue the action.
|
|
|
|
Returns:
|
|
{"success": True, "queued": False} — agent responded OK
|
|
{"success": True, "queued": True, "id": "..."} — queued for later
|
|
Raises HTTPException on agent error (agent reachable but returned failure).
|
|
"""
|
|
from fastapi import HTTPException
|
|
channel = f"{service}.commands"
|
|
ack_channel = f"{service}.acks"
|
|
try:
|
|
ack = publish_and_wait(channel, ack_channel, action, payload)
|
|
if not ack.get("success"):
|
|
raise HTTPException(500, f"Agent error: {ack.get('error', '')}")
|
|
return {"success": True, "queued": False}
|
|
except RuntimeError:
|
|
# Agent unreachable — queue for later
|
|
entry_id = queue_action(service, action, payload)
|
|
return {"success": True, "queued": True, "id": entry_id}
|
|
|
|
|
|
def retry_pending(entry_id: str) -> dict:
|
|
"""Retry a queued action. Returns result dict."""
|
|
with _queue_conn() as conn:
|
|
row = conn.execute(
|
|
"SELECT * FROM pending_actions WHERE id=%s", (entry_id,)
|
|
).fetchone()
|
|
if not row:
|
|
return {"success": False, "error": "Not found"}
|
|
|
|
conn2_ctx = _queue_conn()
|
|
with conn2_ctx as conn:
|
|
conn.execute(
|
|
"UPDATE pending_actions SET status='retrying', updated_at=NOW() WHERE id=%s",
|
|
(entry_id,)
|
|
)
|
|
|
|
service = row["service"]
|
|
action = row["action"]
|
|
payload = json.loads(row["payload"])
|
|
channel = f"{service}.commands"
|
|
ack_channel = f"{service}.acks"
|
|
|
|
try:
|
|
ack = publish_and_wait(channel, ack_channel, action, payload)
|
|
if ack.get("success"):
|
|
mark_done(entry_id)
|
|
return {"success": True, "queued": False}
|
|
else:
|
|
err = ack.get("error", "Unknown error")
|
|
mark_failed(entry_id, err)
|
|
return {"success": False, "error": err}
|
|
except RuntimeError as e:
|
|
mark_failed(entry_id, str(e))
|
|
return {"success": False, "error": str(e)}
|
|
|
|
|
|
def flush_pending_for_service(service: str):
|
|
"""Retry all pending actions for a given service (called when agent reconnects)."""
|
|
with _queue_conn() as conn:
|
|
rows = conn.execute(
|
|
"SELECT * FROM pending_actions WHERE service=%s AND status != 'done' ORDER BY created_at",
|
|
(service,)
|
|
).fetchall()
|
|
if not rows:
|
|
log.info(f"[queue] No pending actions for service '{service}'")
|
|
return
|
|
log.info(f"[queue] Flushing {len(rows)} pending action(s) for '{service}'")
|
|
for row in rows:
|
|
result = retry_pending(row["id"])
|
|
if result.get("success"):
|
|
log.info(f"[queue] Flushed action {row['id']} ({row['action']})")
|
|
else:
|
|
log.warning(f"[queue] Failed to flush {row['id']}: {result.get('error')}")
|
|
|
|
|
|
def publish_dns(action: str, payload: dict) -> dict:
|
|
return publish_and_wait("dns.commands", "dns.acks", action, payload)
|
|
|
|
|
|
def publish_mail(action: str, payload: dict) -> dict:
|
|
return publish_and_wait("mail.commands", "mail.acks", action, payload)
|
|
|
|
|
|
def check_connection() -> bool:
|
|
try:
|
|
get_redis().ping()
|
|
return True
|
|
except Exception:
|
|
return False
|