fix(hub): resolve command injection, run isolation, credentials logging, and CLI argv

This commit is contained in:
banxian1987 2026-08-26 15:07:20 +08:00
parent 636fd247fb
commit 91c750804a
6 changed files with 131 additions and 119 deletions

View file

@ -435,7 +435,7 @@ def main() -> None:
if len(sys.argv) > 1 and sys.argv[1] == "hub":
from strix_hub.main import main as run_hub
run_hub()
run_hub(sys.argv[2:])
return
args = parse_arguments()

View file

@ -118,12 +118,18 @@ def ensure_admin_user() -> None:
cursor.execute("SELECT id FROM users WHERE role = 'admin' LIMIT 1")
if cursor.fetchone() is None:
admin_id = f"user_{secrets.token_hex(6)}"
p_hash, salt = hash_password("admin123")
initial_password = os.environ.get("STRIX_HUB_ADMIN_PASSWORD", "admin123")
p_hash, salt = hash_password(initial_password)
cursor.execute(
"INSERT INTO users (id, username, password_hash, salt, role, created_at) VALUES (?, ?, ?, ?, ?, ?)",
(admin_id, "admin", p_hash, salt, "admin", int(time.time())),
)
conn.commit()
if initial_password == "admin123":
logger.warning(
"Default administrator account initialized (username: admin, password: admin123). "
"Please change your password immediately in user settings or set STRIX_HUB_ADMIN_PASSWORD."
)
# --- User Operations ---

View file

@ -9,13 +9,13 @@ import sys
from strix_hub.server import serve
def main() -> None:
def main(argv: list[str] | None = None) -> None:
parser = argparse.ArgumentParser(description="Strix Hub — Multi-Tenant Web Task & Model Control Platform")
parser.add_argument("--host", default="0.0.0.0", help="Host to bind to (default: 0.0.0.0)")
parser.add_argument("--host", default="127.0.0.1", help="Host to bind to (default: 127.0.0.1)")
parser.add_argument("--port", "-p", type=int, default=8888, help="Port to listen on (default: 8888)")
parser.add_argument("--debug", action="store_true", help="Enable debug logging")
args = parser.parse_args()
args = parser.parse_args(argv)
logging.basicConfig(
level=logging.DEBUG if args.debug else logging.INFO,

View file

@ -57,7 +57,11 @@ class ModelRouterServer:
self.config[k] = v.rstrip("/")
else:
self.config[k] = v
logger.info("ModelRouter config hot-updated: %s", self.config)
logger.info(
"ModelRouter config hot-updated (root_model=%s, subagent_model=%s)",
self.config.get("root_model"),
self.config.get("subagent_model"),
)
def start(self) -> None:
handler = _create_router_handler(self)

View file

@ -181,9 +181,9 @@ def make_hub_handler() -> type[BaseHTTPRequestHandler]:
self._send_json(HTTPStatus.OK, {
"presets": MODEL_PRESETS,
"local_defaults": {
"model": LOCAL_QWEN38_MODEL,
"url": LOCAL_QWEN38_URL,
"key": LOCAL_QWEN38_KEY,
"model": LOCAL_LLM_MODEL,
"url": LOCAL_LLM_URL,
"key": LOCAL_LLM_KEY,
}
})
return

View file

@ -94,16 +94,15 @@ def start_task(task_id: str) -> dict[str, Any]:
# Disable streaming so ModelRouter can accurately parse & auto-recover Qwen text tool-calls
env["LLM_DISABLE_STREAMING"] = "true"
# 3. Assemble command line arguments
# 3. Assemble command line arguments (list of args without shell=True to prevent injection)
target = task["target"]
scan_mode = task["scan_mode"]
instruction = task["instruction"]
strix_bin = get_strix_bin_path()
cmd = f"{strix_bin} -n --target {target} --scan-mode {scan_mode}"
cmd: list[str] = [strix_bin, "-n", "--target", target, "--scan-mode", scan_mode, "--run-name", task_id]
if instruction:
clean_inst = instruction.replace('"', '\\"')
cmd += f' --instruction "{clean_inst}"'
cmd.extend(["--instruction", instruction])
# Work directory
work_dir = "/opt/strix" if Path("/opt/strix").is_dir() else str(Path.cwd())
@ -113,7 +112,7 @@ def start_task(task_id: str) -> dict[str, Any]:
# Spawn in a new process group so SIGSTOP/SIGCONT pauses all children (Docker/tools)
proc = subprocess.Popen(
cmd,
shell=True,
shell=False,
cwd=work_dir,
env=env,
stdout=log_f,
@ -144,131 +143,120 @@ def hot_update_task_config(
subagent_model: str | None = None,
subagent_api_base: str | None = None,
subagent_api_key: str | None = None,
) -> bool:
"""Hot-reload model configuration in DB and live ModelRouter on the fly."""
# 1. Update in SQLite DB
) -> dict[str, Any]:
"""Hot-reload routing configuration for a running or stopped task."""
task = db.get_task_by_id(task_id)
if not task:
raise ValueError(f"Task {task_id} not found")
new_root_model = root_model or task.get("root_model", "")
new_root_api_base = root_api_base if root_api_base is not None else task.get("root_api_base", "")
new_root_key_raw = root_api_key if root_api_key is not None else task.get("root_api_key_raw", "")
new_subagent_model = subagent_model or task.get("subagent_model", "")
new_subagent_api_base = subagent_api_base if subagent_api_base is not None else task.get("subagent_api_base", "")
new_subagent_key_raw = subagent_api_key if subagent_api_key is not None else task.get("subagent_api_key_raw", "")
db.update_task_model_config(
task_id=task_id,
root_model=root_model,
root_api_base=root_api_base,
root_api_key=root_api_key,
subagent_model=subagent_model,
subagent_api_base=subagent_api_base,
subagent_api_key=subagent_api_key,
root_model=new_root_model,
root_api_base=new_root_api_base,
root_api_key=new_root_key_raw,
subagent_model=new_subagent_model,
subagent_api_base=new_subagent_api_base,
subagent_api_key=new_subagent_key_raw,
)
# 2. Hot-reload active ModelRouter instance if task is currently running / paused
with _LOCK:
info = _ACTIVE_TASKS.get(task_id)
if info and info.get("router"):
router: ModelRouterServer = info["router"]
update_payload: dict[str, Any] = {}
if root_model is not None:
update_payload["root_model"] = root_model
if root_api_base is not None:
update_payload["root_api_base"] = root_api_base
if root_api_key is not None:
update_payload["root_api_key"] = root_api_key
if subagent_model is not None:
update_payload["subagent_model"] = subagent_model
if subagent_api_base is not None:
update_payload["subagent_api_base"] = subagent_api_base
if subagent_api_key is not None:
update_payload["subagent_api_key"] = subagent_api_key
router.update_config(**update_payload)
logger.info("Hot-updated running ModelRouter for task %s", task_id)
router = info["router"]
router.update_config(
root_model=new_root_model,
root_api_base=new_root_api_base,
root_api_key=new_root_key_raw,
subagent_model=new_subagent_model,
subagent_api_base=new_subagent_api_base,
subagent_api_key=new_subagent_key_raw,
)
return True
return db.get_task_by_id(task_id) or {}
def pause_task(task_id: str) -> bool:
"""Pause task execution using SIGSTOP."""
def pause_task(task_id: str) -> dict[str, Any]:
"""Pause task process group using SIGSTOP."""
with _LOCK:
info = _ACTIVE_TASKS.get(task_id)
if not info or not info.get("process"):
return False
if not info:
raise ValueError(f"Task {task_id} is not currently active")
proc = info["process"]
if proc.poll() is not None:
return False
try:
pgid = os.getpgid(proc.pid)
os.killpg(pgid, signal.SIGSTOP)
db.update_task_status(task_id, status="paused")
logger.info("Paused task %s (PID %d, PGID %d)", task_id, proc.pid, pgid)
return True
except Exception:
logger.exception("Failed to pause task %s", task_id)
return False
def resume_task(task_id: str) -> bool:
"""Resume task execution using SIGCONT."""
with _LOCK:
info = _ACTIVE_TASKS.get(task_id)
if not info or not info.get("process"):
return False
proc = info["process"]
if proc.poll() is not None:
return False
try:
pgid = os.getpgid(proc.pid)
os.killpg(pgid, signal.SIGCONT)
db.update_task_status(task_id, status="running")
logger.info("Resumed task %s (PID %d, PGID %d)", task_id, proc.pid, pgid)
return True
except Exception:
logger.exception("Failed to resume task %s", task_id)
return False
def stop_task(task_id: str) -> bool:
"""Stop/Kill task execution using SIGTERM / SIGKILL."""
with _LOCK:
info = _ACTIVE_TASKS.get(task_id)
if not info or not info.get("process"):
db.update_task_status(task_id, status="stopped")
return True
proc = info["process"]
try:
pgid = os.getpgid(proc.pid)
os.killpg(pgid, signal.SIGTERM)
time.sleep(0.5)
if proc.poll() is None:
os.killpg(pgid, signal.SIGKILL)
except Exception:
os.killpg(os.getpgid(proc.pid), signal.SIGSTOP)
except ProcessLookupError:
pass
if info.get("router"):
info["router"].stop()
if info.get("log_file") and not info["log_file"].closed:
info["log_file"].close()
db.update_task_status(task_id, status="paused")
return db.get_task_by_id(task_id) or {}
def resume_task(task_id: str) -> dict[str, Any]:
"""Resume task process group using SIGCONT."""
with _LOCK:
info = _ACTIVE_TASKS.get(task_id)
if not info:
raise ValueError(f"Task {task_id} is not currently active")
proc = info["process"]
try:
os.killpg(os.getpgid(proc.pid), signal.SIGCONT)
except ProcessLookupError:
pass
db.update_task_status(task_id, status="running")
return db.get_task_by_id(task_id) or {}
def stop_task(task_id: str) -> dict[str, Any]:
"""Terminate task process group."""
with _LOCK:
info = _ACTIVE_TASKS.get(task_id)
if not info:
db.update_task_status(task_id, status="stopped")
return db.get_task_by_id(task_id) or {}
proc = info["process"]
try:
os.killpg(os.getpgid(proc.pid), signal.SIGTERM)
time.sleep(0.5)
if proc.poll() is None:
os.killpg(os.getpgid(proc.pid), signal.SIGKILL)
except ProcessLookupError:
pass
_ACTIVE_TASKS.pop(task_id, None)
db.update_task_status(task_id, status="stopped")
logger.info("Stopped task %s", task_id)
return True
return db.get_task_by_id(task_id) or {}
def _monitor_task_process(task_id: str) -> None:
"""Monitor task process completion, update status and parse findings."""
"""Monitor background execution and update status, durations, and findings."""
with _LOCK:
info = _ACTIVE_TASKS.get(task_id)
if not info:
return
proc: subprocess.Popen = info["process"]
router: ModelRouterServer = info["router"]
start_time: float = info["start_time"]
proc = info["process"]
router = info["router"]
start_time = info["start_time"]
while proc.poll() is None:
time.sleep(1.0)
time.sleep(2)
elapsed = int(time.time() - start_time)
run_dir_name, vulns_cnt = _inspect_strix_runs_dir(task_id)
current_task = db.get_task_by_id(task_id)
current_status = current_task.get("status") if current_task else "running"
if current_status == "stopped":
break
new_status = "paused" if current_status == "paused" else "running"
db.update_task_status(
task_id,
@ -282,7 +270,13 @@ def _monitor_task_process(task_id: str) -> None:
elapsed = int(time.time() - start_time)
run_dir_name, vulns_cnt = _inspect_strix_runs_dir(task_id)
final_status = "completed" if exit_code in [0, 2] else "failed"
# Preserve intentionally stopped status rather than overriding as failed
current_task = db.get_task_by_id(task_id)
if current_task and current_task.get("status") == "stopped":
final_status = "stopped"
else:
final_status = "completed" if exit_code in [0, 2] else "failed"
db.update_task_status(
task_id,
status=final_status,
@ -297,26 +291,34 @@ def _monitor_task_process(task_id: str) -> None:
with _LOCK:
_ACTIVE_TASKS.pop(task_id, None)
logger.info("Task %s completed with status [%s] (code %d)", task_id, final_status, exit_code)
logger.info("Task %s ended with status [%s] (code %s)", task_id, final_status, exit_code)
def _inspect_strix_runs_dir(task_id: str) -> tuple[str | None, int]:
"""Inspect newest run artifacts to find linked run_dir and vulnerability counts."""
"""Inspect run directory corresponding to task_id without cross-task pollution."""
if not STRIX_RUNS_DIR.is_dir():
return None, 0
runs = sorted(
[d for d in STRIX_RUNS_DIR.iterdir() if d.is_dir() and not d.name.startswith(".")],
key=lambda p: p.stat().st_mtime,
reverse=True,
)
# 1. Exact match by dedicated run name
target_dir = STRIX_RUNS_DIR / task_id
if not target_dir.is_dir():
# Fallback to searching matching run prefix
task = db.get_task_by_id(task_id)
if task and task.get("target"):
import re
slug = re.sub(r"[^a-zA-Z0-9_-]", "-", task["target"]).strip("-").lower()
matching = [
d for d in STRIX_RUNS_DIR.iterdir()
if d.is_dir() and d.name.startswith(slug[:20]) and not d.name.startswith(".")
]
if matching:
target_dir = sorted(matching, key=lambda p: p.stat().st_mtime, reverse=True)[0]
if not runs:
if not target_dir.is_dir():
return None, 0
latest = runs[0]
vuln_count = 0
vulns_file = latest / "vulnerabilities.json"
vulns_file = target_dir / "vulnerabilities.json"
if vulns_file.is_file():
try:
with open(vulns_file, "r", encoding="utf-8") as f:
@ -326,7 +328,7 @@ def _inspect_strix_runs_dir(task_id: str) -> tuple[str | None, int]:
except Exception:
pass
return latest.name, vuln_count
return target_dir.name, vuln_count
def get_task_logs(task_id: str, max_lines: int = 100) -> str: