file transfer
This commit is contained in:
@@ -11,11 +11,11 @@ else:
|
||||
# Fallback to load default .env if any variables are not yet defined
|
||||
load_dotenv()
|
||||
|
||||
from fastapi import FastAPI, HTTPException, Depends, Header
|
||||
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
|
||||
from typing import Dict, Any, List, Optional
|
||||
import uvicorn
|
||||
import base64
|
||||
import hmac
|
||||
@@ -104,6 +104,28 @@ 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
|
||||
|
||||
@@ -303,6 +325,9 @@ class TelemetryPayload(BaseModel):
|
||||
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
|
||||
@@ -314,15 +339,45 @@ 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):
|
||||
if payload.username == "root" and payload.password == "seekright159@":
|
||||
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):
|
||||
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.
|
||||
"""
|
||||
@@ -335,7 +390,7 @@ async def get_whitelisted_commands(platform: str):
|
||||
raise HTTPException(status_code=500, detail=f"Failed to load commands: {e}")
|
||||
|
||||
@app.get("/api/get-command")
|
||||
async def get_command(client_id: str):
|
||||
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?"
|
||||
@@ -389,7 +444,7 @@ async def schedule_command(client_id: str, command: str, current_user: str = Dep
|
||||
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):
|
||||
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.
|
||||
"""
|
||||
@@ -548,7 +603,7 @@ async def get_logs_history(current_user: str = Depends(get_current_user)):
|
||||
return {}
|
||||
|
||||
@app.post("/api/command-result")
|
||||
async def receive_command_result(client_id: str, payload: CommandResultPayload):
|
||||
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.
|
||||
"""
|
||||
@@ -636,3 +691,278 @@ 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
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user