471 lines
18 KiB
Python
471 lines
18 KiB
Python
from dotenv import load_dotenv
|
|
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 load_clients() -> Dict[str, Any]:
|
|
try:
|
|
# 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?"
|
|
"""
|
|
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"
|
|
clients[client_id]["active"] = True
|
|
clients[client_id]["last_seen"] = datetime.now().isoformat()
|
|
|
|
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.
|
|
"""
|
|
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
|
|
clients[client_id]["last_seen"] = datetime.now().isoformat()
|
|
clients[client_id]["active"] = True
|
|
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.
|
|
"""
|
|
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)
|