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, Depends, Header, Request, BackgroundTasks from fastapi.responses import PlainTextResponse from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel from typing import Dict, Any, List, Optional import uvicorn import base64 import hmac import hashlib import json from datetime import datetime, timedelta 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=["*"], ) # JWT-like HMAC Secure Token Utilities JWT_SECRET_KEY = os.getenv("JWT_SECRET_KEY", "seekright_rmm_secret_key_2026_default") def generate_token(username: str) -> str: # Expire in 1 day (24 hours) expiry = (datetime.utcnow() + timedelta(days=1)).isoformat() payload = { "username": username, "expires": expiry } payload_json = json.dumps(payload) payload_b64 = base64.urlsafe_b64encode(payload_json.encode('utf-8')).decode('utf-8').rstrip('=') sig = hmac.new(JWT_SECRET_KEY.encode('utf-8'), payload_b64.encode('utf-8'), hashlib.sha256).hexdigest() return f"{payload_b64}.{sig}" def verify_token(token: str) -> bool: try: parts = token.split(".") if len(parts) != 2: return False payload_b64, sig = parts # Verify signature expected_sig = hmac.new(JWT_SECRET_KEY.encode('utf-8'), payload_b64.encode('utf-8'), hashlib.sha256).hexdigest() if not hmac.compare_digest(sig, expected_sig): return False # Decode and check expiry padding = 4 - (len(payload_b64) % 4) if padding < 4: payload_b64 += "=" * padding payload_json = base64.urlsafe_b64decode(payload_b64.encode('utf-8')).decode('utf-8') payload = json.loads(payload_json) expires_dt = datetime.fromisoformat(payload["expires"]) if datetime.utcnow() > expires_dt: return False # Expired return payload["username"] == "root" except Exception: return False # FastAPI dependency to secure UI endpoints async def get_current_user(authorization: str = Header(None)): if not authorization or not authorization.startswith("Bearer "): raise HTTPException(status_code=401, detail="Missing or invalid authentication credentials") token = authorization.split(" ")[1] if not verify_token(token): raise HTTPException(status_code=401, detail="Authentication token is invalid or has expired") return "root" 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"] config_collection = db["config"] # --- Agent authentication --- # Shared secret agents present in X-Agent-Token. Auth mode ("grace"|"strict") is # stored in the config collection so it can be flipped live from the dashboard: # grace = accept authenticated and legacy (tokenless) agents; used during rollout # strict = reject any request without a valid token AGENT_TOKEN = os.getenv("AGENT_TOKEN", "") def get_agent_auth_mode() -> str: doc = config_collection.find_one({"_id": "agent_auth"}) return (doc or {}).get("mode", "grace") def is_valid_agent_token(token: Optional[str]) -> bool: return bool(AGENT_TOKEN) and token is not None and hmac.compare_digest(token, AGENT_TOKEN) async def require_agent_token(x_agent_token: str = Header(None)): """Dependency guarding every agent-facing endpoint. Enforces only in strict mode; in grace mode it lets legacy tokenless agents through so a fleet can migrate.""" if get_agent_auth_mode() == "strict" and not is_valid_agent_token(x_agent_token): raise HTTPException(status_code=401, detail="Invalid or missing agent token") return is_valid_agent_token(x_agent_token) MAX_ENTRIES_PER_CLIENT = 100 # Keep the latest 100 historical logs per client node ROCKETCHAT_WEBHOOK_URL = os.getenv("ROCKETCHAT_WEBHOOK_URL", "") # Alert states are now stored directly in the MongoDB client documents # to support multi-worker environments safely. 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): try: now_str = datetime.now().isoformat() # Atomically check if transitioned from offline to online (active is not True) old_doc = clients_collection.find_one_and_update( {"_id": client_id, "active": {"$ne": True}}, {"$set": {"active": True, "last_seen": now_str, "last_down_alert_time": None}}, return_document=False ) if old_doc: # Client transitioned from inactive/None to active! # Send UP recovery alert only if they had a down alert time logged if old_doc.get("last_down_alert_time"): send_rocketchat_notification( text=f"✅ **System UP:** Client `{client_id}` has recovered and is back online.", color="#2ecc71" ) else: # Client is already active, just update their last_seen timestamp clients_collection.update_one( {"_id": client_id}, {"$set": {"last_seen": now_str}} ) except Exception as e: print(f"[!] Error marking client active: {e}") 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*$"}}) # Load all documents cursor = clients_collection.find() data = {} for doc in cursor: cid = doc["_id"] # Compute active status dynamically for reads, but do NOT write back to database or alert last_seen_str = doc.get("last_seen") is_active = doc.get("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 = False except Exception: pass info = {k: v for k, v in doc.items() if k != "_id"} info["active"] = is_active data[cid] = info return data except Exception as e: print(f"[!] Error loading clients from MongoDB: {e}") return {} import asyncio async def monitor_heartbeats_loop(): while True: try: now = datetime.now() timeout_time = (now - timedelta(seconds=30)).isoformat() # Find clients that are active in DB but haven't checked in for 30s cursor = clients_collection.find({ "active": True, "last_seen": {"$lt": timeout_time} }) for doc in cursor: cid = doc["_id"] now_str = now.isoformat() # Atomically update to active = False # Ensures only one request/worker triggers the state transition and sends the alert old_doc = clients_collection.find_one_and_update( {"_id": cid, "active": True, "last_seen": doc["last_seen"]}, {"$set": {"active": False, "last_down_alert_time": now_str}}, return_document=False ) if old_doc: send_rocketchat_notification( text=f"🚨 **System DOWN:** Client `{cid}` has missed heartbeats for over 30 seconds.", color="#e74c3c" ) # Hourly reminders for STILL DOWN clients reminder_time = (now - timedelta(hours=1)).isoformat() cursor_still_down = clients_collection.find({ "active": False, "last_down_alert_time": {"$lt": reminder_time} }) for doc in cursor_still_down: cid = doc["_id"] old_last_down = doc["last_down_alert_time"] now_str = now.isoformat() # Atomically update reminder timestamp old_doc = clients_collection.find_one_and_update( {"_id": cid, "active": False, "last_down_alert_time": old_last_down}, {"$set": {"last_down_alert_time": now_str}}, return_document=False ) if old_doc: send_rocketchat_notification( text=f"🚨 **System STILL DOWN:** Client `{cid}` remains offline (reminder sent every hour).", color="#e74c3c" ) except Exception as e: print(f"[!] Error in heartbeat monitoring loop: {e}") await asyncio.sleep(10) @app.on_event("startup") async def startup_event(): asyncio.create_task(monitor_heartbeats_loop()) print("[*] Heartbeat monitoring background task initialized.") 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]] agent_version: Optional[str] = None speed_upload_mbps: Optional[float] = None speed_download_mbps: Optional[float] = None class CommandResultPayload(BaseModel): command: str returncode: int stdout: str stderr: str class LoginPayload(BaseModel): username: str password: str DASHBOARD_USERNAME = os.getenv("DASHBOARD_USERNAME", "root") DASHBOARD_PASSWORD = os.getenv("DASHBOARD_PASSWORD", "seekright159@") @app.post("/api/login") async def login(payload: LoginPayload): user_ok = hmac.compare_digest(payload.username, DASHBOARD_USERNAME) pass_ok = hmac.compare_digest(payload.password, DASHBOARD_PASSWORD) if user_ok and pass_ok: token = generate_token("root") return {"token": token} raise HTTPException(status_code=401, detail="Invalid username or password") class AuthModePayload(BaseModel): mode: str @app.get("/api/agent-auth-mode") async def read_agent_auth_mode(current_user: str = Depends(get_current_user)): return {"mode": get_agent_auth_mode()} @app.post("/api/set-agent-auth-mode") async def set_agent_auth_mode(payload: AuthModePayload, current_user: str = Depends(get_current_user)): if payload.mode not in ("grace", "strict"): raise HTTPException(status_code=400, detail="mode must be 'grace' or 'strict'") config_collection.update_one({"_id": "agent_auth"}, {"$set": {"mode": payload.mode}}, upsert=True) return {"status": "success", "mode": payload.mode} @app.get("/api/agent-versions") async def agent_versions(current_user: str = Depends(get_current_user)): """Migration dashboard: which agent version each node reports and whether its token is valid.""" out = {} for doc in clients_collection.find({}, {"agent_version": 1, "agent_token_ok": 1}): out[doc["_id"]] = { "agent_version": doc.get("agent_version", "unknown"), "agent_token_ok": doc.get("agent_token_ok", False), } return out @app.get("/api/get-commands") async def get_whitelisted_commands(platform: str, token_valid: bool = Depends(require_agent_token)): """ 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, token_valid: bool = Depends(require_agent_token)): """ 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") now_str = datetime.now().isoformat() # Dynamic fleet auto-registration clients_collection.update_one( {"_id": client_id}, {"$setOnInsert": {"pending_command": "none", "active": True, "last_seen": now_str}}, upsert=True ) # Get and reset command atomically updated_doc = clients_collection.find_one_and_update( {"_id": client_id}, {"$set": {"pending_command": "none"}}, return_document=False # Returns the state before update ) cmd = "none" if updated_doc: cmd = updated_doc.get("pending_command", "none") if cmd != "none": append_to_logs("command_polled", client_id, {"command": cmd}) mark_client_active(client_id) return {"command": cmd} @app.post("/api/schedule-command") async def schedule_command(client_id: str, command: str, current_user: str = Depends(get_current_user)): """ Endpoint for your Dashboard/UI to schedule a new command. Supports a comma-separated list of client IDs for batch fleet updates. """ 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: clients_collection.update_one( {"_id": cid}, {"$set": {"pending_command": command}}, upsert=True ) append_to_logs("command_scheduled", cid, {"command": command}) 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, token_valid: bool = Depends(require_agent_token)): """ 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") now_str = datetime.now().isoformat() # Auto-register client if we see it through telemetry first clients_collection.update_one( {"_id": client_id}, {"$setOnInsert": {"pending_command": "none", "active": True, "last_seen": now_str}}, upsert=True ) mark_client_active(client_id) # Get current alert states from database to prevent duplicate alerts client_doc = clients_collection.find_one({"_id": client_id}) or {} storage_alerts = client_doc.get("active_storage_alerts", []) gpu_alert_sent = client_doc.get("gpu_alert_sent", False) storage_modified = False gpu_modified = False # Storage warning checker for disk in payload.disks: mount = disk.get("mount", "/") percent = disk.get("percent", 0.0) if percent >= 80.0: if mount not in storage_alerts: storage_alerts.append(mount) storage_modified = True send_rocketchat_notification( text=f"⚠️ **Storage Warning:** Client `{client_id}` disk `{mount}` is at **{percent}%** capacity.", color="#f39c12" ) else: if mount in storage_alerts: storage_alerts.remove(mount) storage_modified = True send_rocketchat_notification( text=f"✅ **Storage Recovered:** Client `{client_id}` disk `{mount}` has cleared warning state and is at **{percent}%**.", color="#2ecc71" ) # GPU Failure Monitor if len(payload.gpus) == 0: if not gpu_alert_sent: gpu_alert_sent = True gpu_modified = True send_rocketchat_notification( text=f"🚨 **GPU Failure:** Client `{client_id}` is not reporting any GPU data! (nvidia-smi has failed or is missing)", color="#e74c3c" ) else: if gpu_alert_sent: gpu_alert_sent = False gpu_modified = True send_rocketchat_notification( text=f"✅ **GPU Recovered:** Client `{client_id}` is reporting GPU data successfully again.", color="#2ecc71" ) update_fields = { "telemetry": payload.dict() } if storage_modified: update_fields["active_storage_alerts"] = storage_alerts if gpu_modified: update_fields["gpu_alert_sent"] = gpu_alert_sent clients_collection.update_one( {"_id": client_id}, {"$set": update_fields} ) # Log telemetry history event append_to_logs("telemetry", client_id, payload.dict()) data = payload.dict() client_telemetry[client_id] = data # Log received stats 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"} 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(current_user: str = Depends(get_current_user)): """ Endpoint to view the live dashboard data of all clients. """ clients = load_clients() return {cid: info.get("telemetry", {}) for cid, info in clients.items() if info.get("telemetry")} @app.get("/api/clients") async def get_clients_api(current_user: str = Depends(get_current_user)): """ 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(current_user: str = Depends(get_current_user)): """ 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, token_valid: bool = Depends(require_agent_token)): """ 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(current_user: str = Depends(get_current_user)): """ 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], current_user: str = Depends(get_current_user)): """ 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) # ============================================================ # Remote File Transfer (video fetch from agent TAKELEAP folders) # ============================================================ import os FILE_TRANSFER_DIR = os.path.join(os.path.dirname(os.path.abspath(__file__)), "file_transfers") os.makedirs(FILE_TRANSFER_DIR, exist_ok=True) class FileStatusPayload(BaseModel): filename: str status: str message: Optional[str] = "" progress_percent: Optional[float] = None def set_file_transfer_state(client_id: str, filename: str, status: str, message: str = "", extra: Optional[Dict[str, Any]] = None): state = { "filename": filename, "status": status, "message": message, "updated": datetime.now().isoformat() } if extra: state.update(extra) clients_collection.update_one( {"_id": client_id}, {"$set": {"file_transfer": state}}, upsert=True ) @app.post("/api/request-file") async def request_file(client_id: str, filename: str, current_user: str = Depends(get_current_user)): safe_name = os.path.basename(filename.strip()) if not client_id.strip() or not safe_name: raise HTTPException(status_code=400, detail="client_id and filename are required") clients_collection.update_one( {"_id": client_id}, {"$set": {"pending_file_request": safe_name}}, upsert=True ) set_file_transfer_state(client_id, safe_name, "requested", "Waiting for agent to poll...") append_to_logs("file_requested", client_id, {"filename": safe_name}) return {"status": "success", "message": f"File '{safe_name}' requested from {client_id}"} @app.get("/api/get-file-request") async def get_file_request(client_id: str, token_valid: bool = Depends(require_agent_token)): if not client_id or not client_id.strip(): raise HTTPException(status_code=400, detail="client_id cannot be empty") doc = clients_collection.find_one_and_update( {"_id": client_id}, {"$set": {"pending_file_request": "none"}}, return_document=False ) fname = doc.get("pending_file_request", "none") if doc else "none" return {"filename": fname} @app.post("/api/file-transfer-status") async def update_file_transfer_status(client_id: str, payload: FileStatusPayload, token_valid: bool = Depends(require_agent_token)): if not client_id or not client_id.strip(): raise HTTPException(status_code=400, detail="client_id cannot be empty") set_file_transfer_state(client_id, payload.filename, payload.status, payload.message or "", {"progress_percent": payload.progress_percent}) append_to_logs("file_transfer_status", client_id, payload.dict()) return {"status": "success"} @app.post("/api/cancel-upload") async def cancel_upload(client_id: str, filename: str): safe_name = os.path.basename(filename.strip()) set_file_transfer_state(client_id, safe_name, "cancelled", "Upload cancelled by user") client_dir = os.path.join(FILE_TRANSFER_DIR, os.path.basename(client_id.strip())) if os.path.exists(client_dir): for f in os.listdir(client_dir): if f.startswith(f"{safe_name}.part"): try: os.remove(os.path.join(client_dir, f)) except: pass return {"status": "success"} @app.get("/api/received-chunks") async def received_chunks(client_id: str, filename: str, chunk_size: int, file_size: int, token_valid: bool = Depends(require_agent_token)): """Resume support: report which chunk indexes are already stored for this file. Parts whose size doesn't match the current chunking scheme (stale from an earlier run with a different chunk size) are deleted so they can't corrupt the final stitch.""" safe_name = os.path.basename(filename.strip()) if not client_id.strip() or not safe_name or chunk_size <= 0 or file_size < 0: raise HTTPException(status_code=400, detail="invalid parameters") total_chunks = 1 if file_size == 0 else (file_size + chunk_size - 1) // chunk_size last_expected = file_size - (total_chunks - 1) * chunk_size client_dir = os.path.join(FILE_TRANSFER_DIR, os.path.basename(client_id.strip())) received = [] if os.path.isdir(client_dir): prefix = f"{safe_name}.part" for entry in os.listdir(client_dir): if not entry.startswith(prefix) or not entry[len(prefix):].isdigit(): continue idx = int(entry[len(prefix):]) path = os.path.join(client_dir, entry) expected = chunk_size if idx < total_chunks - 1 else last_expected try: if idx < total_chunks and os.path.getsize(path) == expected: received.append(idx) else: os.remove(path) except OSError: pass # part is mid-write or locked; report as not received return {"received": sorted(received)} @app.post("/api/upload-chunk") async def upload_chunk(client_id: str, filename: str, chunk_index: int, request: Request, token_valid: bool = Depends(require_agent_token)): safe_name = os.path.basename(filename.strip()) if not client_id.strip() or not safe_name: raise HTTPException(status_code=400, detail="client_id and filename are required") doc = clients_collection.find_one({"_id": client_id}) if doc and doc.get("file_transfer", {}).get("filename") == safe_name and doc.get("file_transfer", {}).get("status") == "cancelled": return {"status": "cancelled"} client_dir = os.path.join(FILE_TRANSFER_DIR, os.path.basename(client_id.strip())) os.makedirs(client_dir, exist_ok=True) dest_path = os.path.join(client_dir, f"{safe_name}.part{chunk_index}") try: with open(dest_path, "wb") as f: async for chunk in request.stream(): f.write(chunk) except Exception as e: raise HTTPException(status_code=500, detail=f"Failed to store chunk: {e}") return {"status": "success", "chunk_index": chunk_index} def enforce_storage_limit(max_bytes: int = 20 * 1024**3): """Scan FILE_TRANSFER_DIR recursively and delete oldest completed files if total size > max_bytes.""" all_files = [] total_size = 0 for root, dirs, files in os.walk(FILE_TRANSFER_DIR): for name in files: if ".part" in name: continue path = os.path.join(root, name) try: stat = os.stat(path) all_files.append((stat.st_mtime, path, stat.st_size)) total_size += stat.st_size except FileNotFoundError: pass if total_size > max_bytes: all_files.sort(key=lambda x: x[0]) for mtime, path, size in all_files: if total_size <= max_bytes: break try: os.remove(path) total_size -= size print(f"Deleted old file {path} to free space.") except Exception as e: print(f"Error deleting file {path}: {e}") @app.post("/api/upload-complete") async def upload_complete(client_id: str, filename: str, total_chunks: int, background_tasks: BackgroundTasks, token_valid: bool = Depends(require_agent_token)): safe_name = os.path.basename(filename.strip()) client_dir = os.path.join(FILE_TRANSFER_DIR, os.path.basename(client_id.strip())) final_path = os.path.join(client_dir, safe_name) try: with open(final_path, "wb") as outfile: for i in range(total_chunks): part_path = os.path.join(client_dir, f"{safe_name}.part{i}") if not os.path.exists(part_path): raise HTTPException(status_code=400, detail=f"Missing chunk {i}") with open(part_path, "rb") as infile: outfile.write(infile.read()) os.remove(part_path) except Exception as e: set_file_transfer_state(client_id, safe_name, "error", f"Upload finalize failed: {e}") raise HTTPException(status_code=500, detail=f"Failed to finalize file: {e}") size_mb = round(os.path.getsize(final_path) / (1024 * 1024), 2) set_file_transfer_state(client_id, safe_name, "ready", f"File received ({size_mb} MB)", {"size_mb": size_mb, "progress_percent": 100.0}) append_to_logs("file_received", client_id, {"filename": safe_name, "size_mb": size_mb}) send_rocketchat_notification( text=f"📥 **File Received:** `{safe_name}` ({size_mb} MB) uploaded from client `{client_id}`.", color="#2ecc71" ) background_tasks.add_task(enforce_storage_limit) return {"status": "success", "size_mb": size_mb} @app.get("/api/file-transfers") async def get_file_transfers(current_user: str = Depends(get_current_user)): try: cursor = clients_collection.find({}, {"file_transfer": 1}) return {doc["_id"]: doc.get("file_transfer") for doc in cursor if doc.get("file_transfer")} except Exception as e: print(f"[!] Error loading file transfers: {e}") return {} from fastapi.responses import FileResponse, Response @app.get("/deploy_agent.sh") async def get_agent_installer(x_agent_token: str = Header(None), token: Optional[str] = None): # In strict mode the installer carries the live token, so downloading it must # itself be authenticated (header for agent self-update, ?token= for humans). if get_agent_auth_mode() == "strict" and not (is_valid_agent_token(x_agent_token) or is_valid_agent_token(token)): raise HTTPException(status_code=401, detail="Invalid or missing agent token") path = os.path.join(os.path.dirname(os.path.abspath(__file__)), "deploy_agent.sh") if not os.path.isfile(path): raise HTTPException(status_code=404, detail="Installer script not found on server") # Normalize CRLF -> LF (edited on Windows, executed by bash on Linux) and inject # the live agent token so a freshly downloaded installer bakes it into the agent. with open(path, "rb") as f: content = f.read().replace(b"\r\n", b"\n") content = content.replace(b"__AGENT_TOKEN__", AGENT_TOKEN.encode("utf-8")) return Response( content=content, media_type="text/x-shellscript", headers={"Content-Disposition": 'attachment; filename="deploy_agent.sh"'}, ) @app.get("/api/download-file") async def download_file(client_id: str, filename: str, token: Optional[str] = None, inline: bool = False, authorization: str = Header(None)): authed = False if authorization and authorization.startswith("Bearer ") and verify_token(authorization.split(" ")[1]): authed = True if token and verify_token(token): authed = True if not authed: raise HTTPException(status_code=401, detail="Authentication token is invalid or has expired") safe_name = os.path.basename(filename.strip()) path = os.path.join(FILE_TRANSFER_DIR, os.path.basename(client_id.strip()), safe_name) if not os.path.isfile(path): raise HTTPException(status_code=404, detail="File not found on server") media_type = "video/mp4" if safe_name.lower().endswith(".mp4") else "application/octet-stream" if inline: return FileResponse(path, media_type=media_type) return FileResponse(path, media_type=media_type, filename=safe_name) @app.post("/api/heartbeat") async def receive_heartbeat(client_id: str, payload: TelemetryPayload, token_valid: bool = Depends(require_agent_token)): await receive_telemetry(client_id, payload) # Record migration status so the dashboard can confirm the fleet is upgraded # and every node is authenticated before auth is flipped to strict. clients_collection.update_one( {"_id": client_id}, {"$set": {"agent_version": payload.agent_version or "unknown", "speed_upload_mbps": payload.speed_upload_mbps, "speed_download_mbps": payload.speed_download_mbps, "agent_token_ok": token_valid}}, ) doc = clients_collection.find_one_and_update( {"_id": client_id}, {"$set": {"pending_command": "none", "pending_file_request": "none"}}, return_document=False ) cmd = "none" fname = "none" if doc: cmd = doc.get("pending_command", "none") fname = doc.get("pending_file_request", "none") if cmd != "none": append_to_logs("command_polled", client_id, {"command": cmd}) return { "status": "success", "command": cmd, "filename": fname }