Files
rmm-backend/central_api_prototype.py

505 lines
20 KiB
Python

import os
from dotenv import load_dotenv
# Dynamically select configuration file based on APP_ENV
app_env = os.getenv("APP_ENV", "development")
if app_env == "production":
load_dotenv(".env.production")
else:
load_dotenv(".env.development")
# Fallback to load default .env if any variables are not yet defined
load_dotenv()
from fastapi import FastAPI, HTTPException
from fastapi.responses import PlainTextResponse
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from typing import Dict, Any, List
import uvicorn
app = FastAPI(title="RMM Central API")
# Enable CORS for frontend dashboard queries (React dev server runs on a separate port)
app.add_middleware(
CORSMiddleware,
allow_origins=["*"],
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
import os
import json
from datetime import datetime
# In-memory storage for client hardware telemetry
client_telemetry = {}
from pymongo import MongoClient
# Database Connection: Configured via environment variable with local fallback
MONGODB_URI = os.getenv("MONGODB_URI", "mongodb://localhost:27017")
DB_NAME = os.getenv("MONGODB_DB", "rmm_db")
mongo_client = MongoClient(MONGODB_URI)
db = mongo_client[DB_NAME]
clients_collection = db["clients"]
logs_collection = db["logs"]
MAX_ENTRIES_PER_CLIENT = 100 # Keep the latest 100 historical logs per client node
ROCKETCHAT_WEBHOOK_URL = os.getenv("ROCKETCHAT_WEBHOOK_URL", "")
# In-memory storage to prevent duplicate/spam storage notifications
active_storage_alerts = set()
# Tracks the timestamp when a System DOWN notification was last sent for a client ID
# Key: client_id (str), Value: datetime object of the last notification time
last_down_notification_times = {}
def send_rocketchat_notification(text: str, color: str = "#808080"):
"""
Dispatches a formatted notification payload to the Rocket.Chat incoming webhook.
"""
if not ROCKETCHAT_WEBHOOK_URL:
try:
print(f"[Rocket.Chat Simulation] {text}")
except UnicodeEncodeError:
# Fallback for Windows consoles that do not support printing unicode emojis
sanitized_text = text.encode('ascii', errors='backslashreplace').decode('ascii')
print(f"[Rocket.Chat Simulation] {sanitized_text}")
return
payload = {
"text": text,
"attachments": [
{
"color": color,
"ts": datetime.now().isoformat()
}
]
}
try:
import urllib.request
import json
req = urllib.request.Request(
ROCKETCHAT_WEBHOOK_URL,
data=json.dumps(payload).encode('utf-8'),
headers={'Content-Type': 'application/json'}
)
with urllib.request.urlopen(req, timeout=3.0) as response:
pass
except Exception as e:
print(f"[!] Error sending Rocket.Chat notification: {e}")
def append_to_logs(log_type: str, client_id: str, detail: Any):
"""
Appends a new event log grouped by client_id in MongoDB logs collection.
Ensures fair-share log limits per system so noisy clients never overwrite others.
"""
try:
entry = {
"client_id": client_id,
"timestamp": datetime.now().isoformat(),
"type": log_type,
"detail": detail
}
logs_collection.insert_one(entry)
# Enforce fair-share logging limits per system (FIFO circular buffer)
count = logs_collection.count_documents({"client_id": client_id})
if count > MAX_ENTRIES_PER_CLIENT:
oldest_docs = logs_collection.find(
{"client_id": client_id},
{"_id": 1}
).sort("timestamp", 1).limit(count - MAX_ENTRIES_PER_CLIENT)
ids_to_delete = [doc["_id"] for doc in oldest_docs]
if ids_to_delete:
logs_collection.delete_many({"_id": {"$in": ids_to_delete}})
except Exception as e:
print(f"[!] Error writing log to MongoDB: {e}")
def mark_client_active(client_id: str, clients: Dict[str, Any]):
if client_id in clients:
was_active = clients[client_id].get("active", False)
if not was_active:
if client_id in last_down_notification_times:
send_rocketchat_notification(
text=f"✅ **System UP:** Client `{client_id}` has recovered and is back online.",
color="#2ecc71"
)
last_down_notification_times.pop(client_id, None)
clients[client_id]["active"] = True
clients[client_id]["last_seen"] = datetime.now().isoformat()
def load_clients() -> Dict[str, Any]:
try:
# Clean up any ghost clients with empty or whitespace-only keys
clients_collection.delete_many({"_id": {"$in": ["", None]}})
clients_collection.delete_many({"_id": {"$regex": "^\\s*$"}})
# If collection is empty, seed it with default registry
if clients_collection.count_documents({}) == 0:
default_data = {
"client_1": {"pending_command": "none", "last_seen": None, "active": False},
"client_2": {"pending_command": "none", "last_seen": None, "active": False},
"client_3": {"pending_command": "none", "last_seen": None, "active": False}
}
save_clients(default_data)
return default_data
# Load all documents
cursor = clients_collection.find()
data = {}
for doc in cursor:
cid = doc["_id"]
data[cid] = {k: v for k, v in doc.items() if k != "_id"}
# Compute active status dynamically based on 30-second heartbeat check-ins
modified = False
for cid, info in data.items():
last_seen_str = info.get("last_seen")
was_active = info.get("active", False)
is_active = False
if last_seen_str:
try:
last_seen_dt = datetime.fromisoformat(last_seen_str)
if (datetime.now() - last_seen_dt).total_seconds() < 30:
is_active = True
except Exception:
pass
if was_active != is_active:
info["active"] = is_active
modified = True
# Active status transitions alerts
if is_active:
# Transitions from offline (DOWN) to online (UP)
if cid in last_down_notification_times:
send_rocketchat_notification(
text=f"✅ **System UP:** Client `{cid}` has recovered and is back online.",
color="#2ecc71"
)
last_down_notification_times.pop(cid, None)
else:
# Transitions from online (UP) to offline (DOWN) for the first time
now = datetime.now()
last_down_notification_times[cid] = now
send_rocketchat_notification(
text=f"🚨 **System DOWN:** Client `{cid}` has missed heartbeats for over 30 seconds.",
color="#e74c3c"
)
else:
# No transition (state remains the same)
if not is_active:
# Client remains offline. Check if we should send a repeating reminder.
# We ONLY send a reminder if it is already tracked in last_down_notification_times.
# This prevents sending alerts for historically offline systems at server boot!
if cid in last_down_notification_times:
now = datetime.now()
elapsed_seconds = (now - last_down_notification_times[cid]).total_seconds()
if elapsed_seconds >= 3600: # 1 hour (3600 seconds)
last_down_notification_times[cid] = now
send_rocketchat_notification(
text=f"🚨 **System STILL DOWN:** Client `{cid}` remains offline (reminder sent every hour).",
color="#e74c3c"
)
if modified:
save_clients(data)
return data
except Exception as e:
print(f"[!] Error loading clients from MongoDB: {e}")
return {}
def save_clients(data: Dict[str, Any]):
try:
for cid, info in data.items():
# Use upsert to update existing client or insert new client automatically
clients_collection.update_one(
{"_id": cid},
{"$set": info},
upsert=True
)
except Exception as e:
print(f"[!] Error saving clients to MongoDB: {e}")
class TelemetryPayload(BaseModel):
cpu_percent: float
memory_percent: float
memory_total_gb: float
memory_free_gb: float
disks: List[Dict[str, Any]]
gpus: List[Dict[str, Any]]
class CommandResultPayload(BaseModel):
command: str
returncode: int
stdout: str
stderr: str
@app.get("/api/get-commands")
async def get_whitelisted_commands(platform: str):
"""
Endpoint for remote agents to fetch their OS-specific whitelisted commands.
"""
commands_file = os.path.join(os.path.dirname(os.path.abspath(__file__)), "commands.json")
try:
with open(commands_file, "r") as f:
data = json.load(f)
return data.get(platform, {})
except Exception as e:
raise HTTPException(status_code=500, detail=f"Failed to load commands: {e}")
@app.get("/api/get-command")
async def get_command(client_id: str):
"""
Endpoint for clients to poll for pending commands.
Client Agent hits this endpoint to ask: "Do I have any work to do?"
"""
if not client_id or not client_id.strip():
raise HTTPException(status_code=400, detail="client_id cannot be empty")
clients = load_clients()
# 1. Auto-register new client if we haven't seen it yet
if client_id not in clients:
print(f"[+] Dynamic fleet auto-registration: Registered new agent '{client_id}'")
clients[client_id] = {"pending_command": "none", "last_seen": None, "active": True}
# 2. Fetch the command
cmd = clients[client_id].get("pending_command", "none")
# Log if a valid pending command was polled
if cmd != "none":
append_to_logs("command_polled", client_id, {"command": cmd})
# 3. Reset the command to "none", mark active, and update last_seen timestamp
clients[client_id]["pending_command"] = "none"
mark_client_active(client_id, clients)
save_clients(clients)
return {"command": cmd}
@app.post("/api/schedule-command")
async def schedule_command(client_id: str, command: str):
"""
Endpoint for your Dashboard/UI to schedule a new command.
Supports a comma-separated list of client IDs for batch fleet updates.
"""
clients = load_clients()
target_ids = [cid.strip() for cid in client_id.split(",") if cid.strip()]
if not target_ids:
raise HTTPException(status_code=400, detail="No target client IDs specified")
for cid in target_ids:
if cid not in clients:
clients[cid] = {"pending_command": "none", "last_seen": None, "active": False}
clients[cid]["pending_command"] = command
# Log command scheduling per system
append_to_logs("command_scheduled", cid, {"command": command})
save_clients(clients)
return {"status": "success", "message": f"Command '{command}' scheduled for {', '.join(target_ids)}"}
@app.post("/api/telemetry")
async def receive_telemetry(client_id: str, payload: TelemetryPayload):
"""
Endpoint for agents to push their live hardware telemetry.
"""
if not client_id or not client_id.strip():
raise HTTPException(status_code=400, detail="client_id cannot be empty")
clients = load_clients()
# Auto-register new client if we see it through telemetry first
if client_id not in clients:
print(f"[+] Dynamic fleet auto-registration: Registered new agent '{client_id}' through telemetry")
clients[client_id] = {"pending_command": "none", "last_seen": None, "active": True}
# Update last_seen timestamp, active status, and store the current telemetry details
mark_client_active(client_id, clients)
clients[client_id]["telemetry"] = payload.dict()
save_clients(clients)
# Log telemetry history event
append_to_logs("telemetry", client_id, payload.dict())
# Storage limits warning checker
for disk in payload.disks:
mount = disk.get("mount", "/")
percent = disk.get("percent", 0.0)
alert_key = f"{client_id}:{mount}"
if percent >= 80.0:
if alert_key not in active_storage_alerts:
active_storage_alerts.add(alert_key)
send_rocketchat_notification(
text=f"⚠️ **Storage Warning:** Client `{client_id}` disk `{mount}` is at **{percent}%** capacity.",
color="#f39c12"
)
else:
if alert_key in active_storage_alerts:
active_storage_alerts.remove(alert_key)
send_rocketchat_notification(
text=f"✅ **Storage Recovered:** Client `{client_id}` disk `{mount}` has cleared warning state and is at **{percent}%**.",
color="#2ecc71"
)
data = payload.dict()
client_telemetry[client_id] = data
# Temporarily print the data to the console for the user to see!
print(f"\n[TELEMETRY RECEIVED] from {client_id}:")
print(f" CPU: {data['cpu_percent']}%")
print(f" RAM: {data['memory_percent']}% ({data['memory_free_gb']} GB Free / {data['memory_total_gb']} GB Total)")
for disk in data['disks']:
print(f" Drive {disk['mount']}: {disk['percent']}% Used ({disk['free_gb']} GB Free / {disk['total_gb']} GB Total)")
for gpu in data['gpus']:
power_str = f"Power: {gpu.get('power_draw', 'N/A')} / {gpu.get('power_limit', 'N/A')}"
fan_str = f"Fan: {gpu.get('fan_speed', 'N/A')}"
print(f" GPU [{gpu['name']}]: Core: {gpu['utilization']}, Temp: {gpu['temp']}, VRAM: {gpu['memory_used']}/{gpu['memory_total']}, {power_str}, {fan_str}")
return {"status": "success"}
@app.get("/api/telemetry")
async def get_all_telemetry():
"""
Endpoint to view the live dashboard data of all clients.
"""
return client_telemetry
@app.get("/api/clients")
async def get_clients_api():
"""
Endpoint for the React UI to fetch the live active client registry with computed statuses.
"""
return load_clients()
@app.get("/api/logs")
async def get_logs_history():
"""
Endpoint to retrieve logs history grouped by client_id.
"""
try:
cursor = logs_collection.find().sort("timestamp", 1)
grouped_logs = {}
for doc in cursor:
cid = doc.get("client_id")
if not cid:
continue
if cid not in grouped_logs:
grouped_logs[cid] = []
entry = {
"timestamp": doc.get("timestamp"),
"type": doc.get("type"),
"detail": doc.get("detail")
}
grouped_logs[cid].append(entry)
return grouped_logs
except Exception as e:
print(f"[!] Error loading logs from MongoDB: {e}")
return {}
@app.post("/api/command-result")
async def receive_command_result(client_id: str, payload: CommandResultPayload):
"""
Endpoint for remote agents to push execution results back to the central server.
"""
if not client_id or not client_id.strip():
raise HTTPException(status_code=400, detail="client_id cannot be empty")
append_to_logs("command_result", client_id, payload.dict())
return {"status": "success"}
@app.get("/api/get-raw-commands")
async def get_raw_commands():
"""
Endpoint for the UI to load the complete commands.json whitelist file.
"""
commands_file = os.path.join(os.path.dirname(os.path.abspath(__file__)), "commands.json")
try:
with open(commands_file, "r") as f:
return json.load(f)
except Exception as e:
raise HTTPException(status_code=500, detail=f"Failed to read commands: {e}")
@app.post("/api/save-commands")
async def save_commands(commands: Dict[str, Any]):
"""
Endpoint for the UI to save changes back to commands.json.
"""
commands_file = os.path.join(os.path.dirname(os.path.abspath(__file__)), "commands.json")
try:
with open(commands_file, "w") as f:
json.dump(commands, f, indent=2)
return {"status": "success", "message": "commands.json saved successfully"}
except Exception as e:
raise HTTPException(status_code=500, detail=f"Failed to save commands: {e}")
@app.get("/metrics", response_class=PlainTextResponse)
async def get_prometheus_metrics():
"""
Exposes the latest client telemetry in standard Prometheus text format.
"""
clients = load_clients()
metrics_lines = []
for client_id, data in clients.items():
# Active status (1.0 for active/online, 0.0 for inactive/offline)
active_val = 1.0 if data.get("active", False) else 0.0
metrics_lines.append(f'rmm_client_active{{client_id="{client_id}"}} {active_val}')
# Telemetry stats
telemetry = data.get("telemetry")
if telemetry:
cpu = telemetry.get("cpu_percent", 0.0)
ram = telemetry.get("memory_percent", 0.0)
ram_free = telemetry.get("memory_free_gb", 0.0)
ram_total = telemetry.get("memory_total_gb", 0.0)
metrics_lines.append(f'rmm_cpu_utilization{{client_id="{client_id}"}} {cpu}')
metrics_lines.append(f'rmm_memory_utilization{{client_id="{client_id}"}} {ram}')
metrics_lines.append(f'rmm_memory_free_bytes{{client_id="{client_id}"}} {ram_free * (1024**3)}')
metrics_lines.append(f'rmm_memory_total_bytes{{client_id="{client_id}"}} {ram_total * (1024**3)}')
# Disk mounts
for disk in telemetry.get("disks", []):
mount = disk.get("mount", "/")
disk_percent = disk.get("percent", 0.0)
metrics_lines.append(f'rmm_disk_utilization{{client_id="{client_id}",mount="{mount}"}} {disk_percent}')
# GPU utilization and temps
for gpu in telemetry.get("gpus", []):
gpu_name = gpu.get("name", "GPU")
# Utilization percent parsing
gpu_util_str = str(gpu.get("utilization", "0"))
gpu_util = float(gpu_util_str.replace("%", "").strip())
# Temperature parsing
gpu_temp_str = str(gpu.get("temp", "0"))
gpu_temp = float(gpu_temp_str.replace("C", "").strip())
metrics_lines.append(f'rmm_gpu_utilization{{client_id="{client_id}",gpu_name="{gpu_name}"}} {gpu_util}')
metrics_lines.append(f'rmm_gpu_temperature{{client_id="{client_id}",gpu_name="{gpu_name}"}} {gpu_temp}')
return "\n".join(metrics_lines) + "\n"
if __name__ == "__main__":
print("[*] Starting FastAPI Central Server...")
# uvicorn runs the FastAPI app on port 8000
uvicorn.run("central_api_prototype:app", host="0.0.0.0", port=8000, reload=True)