mirror of
https://github.com/usestrix/strix.git
synced 2026-10-07 02:58:26 +00:00
feat(hub): add Strix Hub multi-tenant task orchestration & dual-channel model routing console
This commit is contained in:
parent
2cc8167814
commit
636fd247fb
9 changed files with 2854 additions and 0 deletions
|
|
@ -97,6 +97,9 @@ Examples:
|
|||
# Extra files placed in the sandbox workspace
|
||||
strix --target ./my-project --workspace-file ./wordlist.txt
|
||||
strix --target https://app.com --workspace-file ./openapi.yaml:specs/openapi.yaml
|
||||
|
||||
# Launch Strix Hub web management & dual-channel model routing console
|
||||
strix hub --port 8888
|
||||
""",
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -431,6 +431,13 @@ def main() -> None:
|
|||
|
||||
sys.exit(run_auth(sys.argv[2:]))
|
||||
|
||||
# `strix hub …` launches the multi-tenant task orchestration and dual-channel router console.
|
||||
if len(sys.argv) > 1 and sys.argv[1] == "hub":
|
||||
from strix_hub.main import main as run_hub
|
||||
|
||||
run_hub()
|
||||
return
|
||||
|
||||
args = parse_arguments()
|
||||
|
||||
start_background_check()
|
||||
|
|
|
|||
94
strix_hub/README.md
Normal file
94
strix_hub/README.md
Normal file
|
|
@ -0,0 +1,94 @@
|
|||
# 🦅 Strix Hub (Multi-Tenant & Dual-Channel Orchestration Console)
|
||||
|
||||
[English](./README.md) | [简体中文](#简体中文)
|
||||
|
||||
---
|
||||
|
||||
## English
|
||||
|
||||
**Strix Hub** is an open-source, multi-tenant task orchestration and dual-channel model routing platform designed for [Strix](https://github.com/usestrix/strix) autonomous AI penetration testing.
|
||||
|
||||
### ✨ Key Features
|
||||
|
||||
1. **Zero-Intrusion Facade Architecture**:
|
||||
- Operates completely outside the core `strix/` codebase as a management facade.
|
||||
- 100% compatible with current and future upstream Strix versions without code merge conflicts.
|
||||
2. **Independent Dual-Channel Model Routing**:
|
||||
- **Root Agent (Brain)**: High-reasoning cloud models (e.g. Gemini 3.1 Pro, Claude 3.7 Sonnet, GPT-4o) for strategic planning and vulnerability discovery.
|
||||
- **Sub-agents (Muscles)**: Private/local models (e.g. Qwen 3.8 / DeepSeek / Ollama / vLLM / SGLang) for high-concurrency port scanning, fuzzing, and payload execution with **zero token costs and zero rate limits**.
|
||||
- Built-in **Auto Tool-Call Recovery** to seamlessly convert text-based XML tool outputs into standard OpenAI function calls.
|
||||
3. **Live Hot-Reloading**:
|
||||
- Switch model providers, base URLs, and API keys on-the-fly without interrupting running penetration tasks.
|
||||
4. **OS-Level Process Control**:
|
||||
- True process group lifecycle management supporting **Start**, **Pause (`SIGSTOP` with 0 CPU & 0 Token consumption)**, **Resume (`SIGCONT`)**, and **Terminate**.
|
||||
5. **Multi-Tenancy & Role-Based Access Control (RBAC)**:
|
||||
- Built-in SQLite database with user isolation (Users manage their own scans; Admins monitor the entire fleet).
|
||||
6. **Zero External Dependencies**:
|
||||
- Powered purely by Python standard libraries (`http.server`, `sqlite3`, `subprocess`, `threading`) and a single-file modern Dark Mode React SPA.
|
||||
|
||||
### 🚀 Quick Start
|
||||
|
||||
```bash
|
||||
# 1. Install & build Strix
|
||||
git clone https://github.com/usestrix/strix.git
|
||||
cd strix
|
||||
make dev-install
|
||||
|
||||
# 2. Launch Strix Hub (default port 8888)
|
||||
python -m strix_hub.main --port 8888
|
||||
```
|
||||
|
||||
Open `http://localhost:8888` in your browser.
|
||||
- **Default Admin Account**: `admin` / `admin123`
|
||||
|
||||
### ⚙️ Environment Variables (Optional)
|
||||
|
||||
```bash
|
||||
export LOCAL_LLM_MODEL="openai/Qwen3.8-27B-abliterated"
|
||||
export LOCAL_LLM_URL="http://127.0.0.1:8000/v1"
|
||||
export LOCAL_LLM_KEY="your-api-key"
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
<a name="简体中文"></a>
|
||||
## 简体中文
|
||||
|
||||
**Strix Hub** 是专为 [Strix](https://github.com/usestrix/strix) 打造的现代化、零侵入、多租户渗透测试任务编排与双渠道模型网关管理平台。
|
||||
|
||||
### 🌟 核心特性
|
||||
|
||||
1. **零侵入门面架构 (Zero-Intrusion Facade)**:
|
||||
- 100% 独立于官方 `strix/` 核心代码运行,完全解耦。
|
||||
- 官方上游仓库执行 `git pull` 升级时**零代码合并冲突**,天然适配当前与未来所有 Strix 版本。
|
||||
2. **独立双渠道模型路由网关 (Dual-Channel Routing)**:
|
||||
- **主控大脑 (Root Agent)**:可自由配置云端顶尖推理模型(如 Gemini 3.1 Pro、Claude 3.7 Sonnet、GPT-4o),负责全局渗透决策与漏洞挖掘。
|
||||
- **探测打手 (Sub-agents)**:可无缝对接局域网私有化模型(如本地部署的 Qwen 3.8-27B 无审查特化版 / DeepSeek / SGLang / vLLM),负责高并发端口扫描与 Web 模糊测试,**零 Token 成本、零外网限流**!
|
||||
- 内置 **Tool-Call 自动解析自愈器**,自动将开源模型输出的 XML 格式工具调用转为标准 OpenAI 协议,确保子智能体触手稳定并发派生。
|
||||
3. **运行时配置热更新 (Live Hot-Reload)**:
|
||||
- 任务在运行中或暂停中,均可在 Web 控制台一键热更模型渠道、Base URL 与 API Key,底层网关实时生效。
|
||||
4. **操作系统进程级精准调度**:
|
||||
- 采用进程组信号控制,真正实现任务 **一键暂停 (`SIGSTOP` 瞬间零 CPU、零 Token 消耗挂起)**、**继续 (`SIGCONT` 唤醒)** 与 **强力终止**。
|
||||
5. **多租户与 RBAC 权限隔离**:
|
||||
- 内置轻量 SQLite 持久化,普通用户仅能查看与操作自己创建的任务,管理员全局统一纳管。
|
||||
6. **零额外依赖 & 极速部署**:
|
||||
- 后端基于 Python 标准库实现,前端内置现代化 Single-File React SPA 暗黑大屏,一键即可启动。
|
||||
|
||||
### 🚀 快速启动
|
||||
|
||||
```bash
|
||||
# 1. 克隆并安装 Strix
|
||||
git clone https://github.com/usestrix/strix.git
|
||||
cd strix
|
||||
make dev-install
|
||||
|
||||
# 2. 启动 Strix Hub 服务 (默认端口 8888)
|
||||
python -m strix_hub.main --port 8888
|
||||
```
|
||||
|
||||
浏览器访问 `http://localhost:8888` 即可开始使用。
|
||||
- **默认管理员账户**:`admin` / `admin123`(登录后可在用户管理中修改或新增成员)
|
||||
|
||||
### 📄 License
|
||||
|
||||
Apache License 2.0
|
||||
3
strix_hub/__init__.py
Normal file
3
strix_hub/__init__.py
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
"""Strix Hub — Zero-Intrusion Web Task & Multi-Model Management Platform."""
|
||||
|
||||
__version__ = "1.0.0"
|
||||
402
strix_hub/db.py
Normal file
402
strix_hub/db.py
Normal file
|
|
@ -0,0 +1,402 @@
|
|||
"""SQLite database persistence for Strix Hub (users, tasks, sessions)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import hmac
|
||||
import os
|
||||
import secrets
|
||||
import sqlite3
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
DB_PATH = Path(os.environ.get("STRIX_HUB_DB", "/opt/strix/strix_hub.db"))
|
||||
if not DB_PATH.parent.exists():
|
||||
try:
|
||||
DB_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
except Exception:
|
||||
DB_PATH = Path.home() / ".strix" / "strix_hub.db"
|
||||
DB_PATH.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
|
||||
def get_connection() -> sqlite3.Connection:
|
||||
conn = sqlite3.connect(str(DB_PATH), check_same_thread=False)
|
||||
conn.row_factory = sqlite3.Row
|
||||
return conn
|
||||
|
||||
|
||||
def init_db() -> None:
|
||||
"""Initialize database tables and create default admin account if not exists."""
|
||||
with get_connection() as conn:
|
||||
conn.executescript("""
|
||||
CREATE TABLE IF NOT EXISTS users (
|
||||
id TEXT PRIMARY KEY,
|
||||
username TEXT UNIQUE NOT NULL,
|
||||
password_hash TEXT NOT NULL,
|
||||
salt TEXT NOT NULL,
|
||||
role TEXT NOT NULL DEFAULT 'user',
|
||||
created_at INTEGER NOT NULL
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS sessions (
|
||||
token TEXT PRIMARY KEY,
|
||||
user_id TEXT NOT NULL,
|
||||
expires_at INTEGER NOT NULL,
|
||||
FOREIGN KEY(user_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
|
||||
CREATE TABLE IF NOT EXISTS tasks (
|
||||
id TEXT PRIMARY KEY,
|
||||
owner_id TEXT NOT NULL,
|
||||
owner_username TEXT NOT NULL,
|
||||
target TEXT NOT NULL,
|
||||
scan_mode TEXT NOT NULL DEFAULT 'deep',
|
||||
instruction TEXT DEFAULT '',
|
||||
root_model TEXT NOT NULL DEFAULT 'openai/gemini-3.1-pro-preview',
|
||||
root_api_base TEXT DEFAULT '',
|
||||
root_api_key_masked TEXT DEFAULT '',
|
||||
root_api_key_raw TEXT DEFAULT '',
|
||||
subagent_model TEXT NOT NULL DEFAULT 'openai/gemini-3.5-flash',
|
||||
subagent_api_base TEXT DEFAULT '',
|
||||
subagent_api_key_masked TEXT DEFAULT '',
|
||||
subagent_api_key_raw TEXT DEFAULT '',
|
||||
api_base TEXT DEFAULT '',
|
||||
api_key_masked TEXT DEFAULT '',
|
||||
api_key_raw TEXT DEFAULT '',
|
||||
status TEXT NOT NULL DEFAULT 'pending',
|
||||
pid INTEGER DEFAULT NULL,
|
||||
run_dir_name TEXT DEFAULT '',
|
||||
vulns_count INTEGER DEFAULT 0,
|
||||
duration_seconds INTEGER DEFAULT 0,
|
||||
log_preview TEXT DEFAULT '',
|
||||
created_at INTEGER NOT NULL,
|
||||
updated_at INTEGER NOT NULL,
|
||||
FOREIGN KEY(owner_id) REFERENCES users(id) ON DELETE CASCADE
|
||||
);
|
||||
""")
|
||||
conn.commit()
|
||||
|
||||
# Dynamic Schema Migration (add columns if older table exists)
|
||||
cursor = conn.cursor()
|
||||
existing_cols = {row["name"] for row in cursor.execute("PRAGMA table_info(tasks)").fetchall()}
|
||||
new_columns = [
|
||||
("root_api_base", "TEXT DEFAULT ''"),
|
||||
("root_api_key_masked", "TEXT DEFAULT ''"),
|
||||
("root_api_key_raw", "TEXT DEFAULT ''"),
|
||||
("subagent_api_base", "TEXT DEFAULT ''"),
|
||||
("subagent_api_key_masked", "TEXT DEFAULT ''"),
|
||||
("subagent_api_key_raw", "TEXT DEFAULT ''"),
|
||||
]
|
||||
for col_name, col_type in new_columns:
|
||||
if col_name not in existing_cols:
|
||||
try:
|
||||
cursor.execute(f"ALTER TABLE tasks ADD COLUMN {col_name} {col_type}")
|
||||
except Exception:
|
||||
pass
|
||||
conn.commit()
|
||||
|
||||
# Seed default admin user if table is empty
|
||||
ensure_admin_user()
|
||||
|
||||
|
||||
def hash_password(password: str, salt: str | None = None) -> tuple[str, str]:
|
||||
if salt is None:
|
||||
salt = secrets.token_hex(16)
|
||||
key = hashlib.pbkdf2_hmac("sha256", password.encode("utf-8"), salt.encode("utf-8"), 100_000)
|
||||
return key.hex(), salt
|
||||
|
||||
|
||||
def verify_password(password: str, password_hash: str, salt: str) -> bool:
|
||||
key, _ = hash_password(password, salt)
|
||||
return hmac.compare_digest(key, password_hash)
|
||||
|
||||
|
||||
def ensure_admin_user() -> None:
|
||||
with get_connection() as conn:
|
||||
cursor = conn.cursor()
|
||||
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")
|
||||
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()
|
||||
|
||||
|
||||
# --- User Operations ---
|
||||
|
||||
def create_user(username: str, password: str, role: str = "user") -> dict[str, Any] | None:
|
||||
user_id = f"user_{secrets.token_hex(6)}"
|
||||
p_hash, salt = hash_password(password)
|
||||
now = int(time.time())
|
||||
try:
|
||||
with get_connection() as conn:
|
||||
conn.execute(
|
||||
"INSERT INTO users (id, username, password_hash, salt, role, created_at) VALUES (?, ?, ?, ?, ?, ?)",
|
||||
(user_id, username.strip(), p_hash, salt, role, now),
|
||||
)
|
||||
conn.commit()
|
||||
return {"id": user_id, "username": username.strip(), "role": role, "created_at": now}
|
||||
except sqlite3.IntegrityError:
|
||||
return None
|
||||
|
||||
|
||||
def get_user_by_username(username: str) -> dict[str, Any] | None:
|
||||
with get_connection() as conn:
|
||||
row = conn.execute("SELECT * FROM users WHERE username = ?", (username.strip(),)).fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
|
||||
def get_user_by_id(user_id: str) -> dict[str, Any] | None:
|
||||
with get_connection() as conn:
|
||||
row = conn.execute("SELECT id, username, role, created_at FROM users WHERE id = ?", (user_id,)).fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
|
||||
def list_users() -> list[dict[str, Any]]:
|
||||
with get_connection() as conn:
|
||||
rows = conn.execute("SELECT id, username, role, created_at FROM users ORDER BY created_at ASC").fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
|
||||
def delete_user(user_id: str) -> bool:
|
||||
with get_connection() as conn:
|
||||
cursor = conn.execute("DELETE FROM users WHERE id = ? AND role != 'admin'", (user_id,))
|
||||
conn.commit()
|
||||
return cursor.rowcount > 0
|
||||
|
||||
|
||||
# --- Session Operations ---
|
||||
|
||||
def create_session(user_id: str, ttl_seconds: int = 86400 * 7) -> str:
|
||||
token = secrets.token_urlsafe(32)
|
||||
expires_at = int(time.time()) + ttl_seconds
|
||||
with get_connection() as conn:
|
||||
conn.execute(
|
||||
"INSERT INTO sessions (token, user_id, expires_at) VALUES (?, ?, ?)",
|
||||
(token, user_id, expires_at),
|
||||
)
|
||||
conn.commit()
|
||||
return token
|
||||
|
||||
|
||||
def validate_session(token: str) -> dict[str, Any] | None:
|
||||
now = int(time.time())
|
||||
with get_connection() as conn:
|
||||
row = conn.execute(
|
||||
"""
|
||||
SELECT u.id, u.username, u.role
|
||||
FROM sessions s
|
||||
JOIN users u ON s.user_id = u.id
|
||||
WHERE s.token = ? AND s.expires_at > ?
|
||||
""",
|
||||
(token, now),
|
||||
).fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
|
||||
def delete_session(token: str) -> None:
|
||||
with get_connection() as conn:
|
||||
conn.execute("DELETE FROM sessions WHERE token = ?", (token,))
|
||||
conn.commit()
|
||||
|
||||
|
||||
# --- Task Operations ---
|
||||
|
||||
def mask_api_key(key: str) -> str:
|
||||
if not key:
|
||||
return ""
|
||||
if len(key) <= 8:
|
||||
return "****"
|
||||
return f"{key[:4]}...{key[-4:]}"
|
||||
|
||||
|
||||
def create_task(
|
||||
owner_id: str,
|
||||
owner_username: str,
|
||||
target: str,
|
||||
scan_mode: str = "deep",
|
||||
instruction: str = "",
|
||||
root_model: str = "openai/gemini-3.1-pro-preview",
|
||||
root_api_base: str = "",
|
||||
root_api_key: str = "",
|
||||
subagent_model: str = "openai/gemini-3.5-flash",
|
||||
subagent_api_base: str = "",
|
||||
subagent_api_key: str = "",
|
||||
) -> dict[str, Any]:
|
||||
task_id = f"task_{secrets.token_hex(6)}"
|
||||
now = int(time.time())
|
||||
root_masked = mask_api_key(root_api_key)
|
||||
sub_masked = mask_api_key(subagent_api_key)
|
||||
|
||||
with get_connection() as conn:
|
||||
conn.execute(
|
||||
"""
|
||||
INSERT INTO tasks (
|
||||
id, owner_id, owner_username, target, scan_mode, instruction,
|
||||
root_model, root_api_base, root_api_key_masked, root_api_key_raw,
|
||||
subagent_model, subagent_api_base, subagent_api_key_masked, subagent_api_key_raw,
|
||||
status, created_at, updated_at
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, 'pending', ?, ?)
|
||||
""",
|
||||
(
|
||||
task_id,
|
||||
owner_id,
|
||||
owner_username,
|
||||
target.strip(),
|
||||
scan_mode,
|
||||
instruction.strip(),
|
||||
root_model.strip(),
|
||||
root_api_base.strip(),
|
||||
root_masked,
|
||||
root_api_key.strip(),
|
||||
subagent_model.strip(),
|
||||
subagent_api_base.strip(),
|
||||
sub_masked,
|
||||
subagent_api_key.strip(),
|
||||
now,
|
||||
now,
|
||||
),
|
||||
)
|
||||
conn.commit()
|
||||
return get_task_by_id(task_id) or {}
|
||||
|
||||
|
||||
def get_task_by_id(task_id: str) -> dict[str, Any] | None:
|
||||
with get_connection() as conn:
|
||||
row = conn.execute(
|
||||
"""
|
||||
SELECT id, owner_id, owner_username, target, scan_mode, instruction,
|
||||
root_model, root_api_base, root_api_key_masked,
|
||||
subagent_model, subagent_api_base, subagent_api_key_masked,
|
||||
status, pid, run_dir_name, vulns_count, duration_seconds,
|
||||
log_preview, created_at, updated_at
|
||||
FROM tasks WHERE id = ?
|
||||
""",
|
||||
(task_id,),
|
||||
).fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
|
||||
def get_task_full(task_id: str) -> dict[str, Any] | None:
|
||||
"""Internal use: returns raw api_keys for task runner."""
|
||||
with get_connection() as conn:
|
||||
row = conn.execute("SELECT * FROM tasks WHERE id = ?", (task_id,)).fetchone()
|
||||
return dict(row) if row else None
|
||||
|
||||
|
||||
def list_tasks(user_id: str | None = None, is_admin: bool = False) -> list[dict[str, Any]]:
|
||||
with get_connection() as conn:
|
||||
if is_admin or user_id is None:
|
||||
rows = conn.execute(
|
||||
"""
|
||||
SELECT id, owner_id, owner_username, target, scan_mode, instruction,
|
||||
root_model, root_api_base, root_api_key_masked,
|
||||
subagent_model, subagent_api_base, subagent_api_key_masked,
|
||||
status, pid, run_dir_name, vulns_count, duration_seconds,
|
||||
created_at, updated_at
|
||||
FROM tasks ORDER BY created_at DESC
|
||||
"""
|
||||
).fetchall()
|
||||
else:
|
||||
rows = conn.execute(
|
||||
"""
|
||||
SELECT id, owner_id, owner_username, target, scan_mode, instruction,
|
||||
root_model, root_api_base, root_api_key_masked,
|
||||
subagent_model, subagent_api_base, subagent_api_key_masked,
|
||||
status, pid, run_dir_name, vulns_count, duration_seconds,
|
||||
created_at, updated_at
|
||||
FROM tasks WHERE owner_id = ? ORDER BY created_at DESC
|
||||
""",
|
||||
(user_id,),
|
||||
).fetchall()
|
||||
return [dict(r) for r in rows]
|
||||
|
||||
|
||||
def update_task_status(
|
||||
task_id: str,
|
||||
status: str,
|
||||
pid: int | None = None,
|
||||
run_dir_name: str | None = None,
|
||||
vulns_count: int | None = None,
|
||||
duration_seconds: int | None = None,
|
||||
log_preview: str | None = None,
|
||||
) -> None:
|
||||
now = int(time.time())
|
||||
updates = ["status = ?", "updated_at = ?"]
|
||||
params: list[Any] = [status, now]
|
||||
|
||||
if pid is not None:
|
||||
updates.append("pid = ?")
|
||||
params.append(pid)
|
||||
if run_dir_name is not None:
|
||||
updates.append("run_dir_name = ?")
|
||||
params.append(run_dir_name)
|
||||
if vulns_count is not None:
|
||||
updates.append("vulns_count = ?")
|
||||
params.append(vulns_count)
|
||||
if duration_seconds is not None:
|
||||
updates.append("duration_seconds = ?")
|
||||
params.append(duration_seconds)
|
||||
if log_preview is not None:
|
||||
updates.append("log_preview = ?")
|
||||
params.append(log_preview)
|
||||
|
||||
params.append(task_id)
|
||||
with get_connection() as conn:
|
||||
conn.execute(f"UPDATE tasks SET {', '.join(updates)} WHERE id = ?", params)
|
||||
conn.commit()
|
||||
|
||||
|
||||
def update_task_model_config(
|
||||
task_id: str,
|
||||
root_model: str | None = None,
|
||||
root_api_base: str | None = None,
|
||||
root_api_key: str | None = None,
|
||||
subagent_model: str | None = None,
|
||||
subagent_api_base: str | None = None,
|
||||
subagent_api_key: str | None = None,
|
||||
) -> None:
|
||||
"""Hot update task model and channel configuration."""
|
||||
now = int(time.time())
|
||||
updates = ["updated_at = ?"]
|
||||
params: list[Any] = [now]
|
||||
|
||||
if root_model is not None:
|
||||
updates.append("root_model = ?")
|
||||
params.append(root_model.strip())
|
||||
if root_api_base is not None:
|
||||
updates.append("root_api_base = ?")
|
||||
params.append(root_api_base.strip())
|
||||
if root_api_key is not None:
|
||||
updates.append("root_api_key_masked = ?")
|
||||
updates.append("root_api_key_raw = ?")
|
||||
params.append(mask_api_key(root_api_key))
|
||||
params.append(root_api_key.strip())
|
||||
|
||||
if subagent_model is not None:
|
||||
updates.append("subagent_model = ?")
|
||||
params.append(subagent_model.strip())
|
||||
if subagent_api_base is not None:
|
||||
updates.append("subagent_api_base = ?")
|
||||
params.append(subagent_api_base.strip())
|
||||
if subagent_api_key is not None:
|
||||
updates.append("subagent_api_key_masked = ?")
|
||||
updates.append("subagent_api_key_raw = ?")
|
||||
params.append(mask_api_key(subagent_api_key))
|
||||
params.append(subagent_api_key.strip())
|
||||
|
||||
params.append(task_id)
|
||||
with get_connection() as conn:
|
||||
conn.execute(f"UPDATE tasks SET {', '.join(updates)} WHERE id = ?", params)
|
||||
conn.commit()
|
||||
|
||||
|
||||
def delete_task(task_id: str) -> bool:
|
||||
with get_connection() as conn:
|
||||
cursor = conn.execute("DELETE FROM tasks WHERE id = ?", (task_id,))
|
||||
conn.commit()
|
||||
return cursor.rowcount > 0
|
||||
29
strix_hub/main.py
Normal file
29
strix_hub/main.py
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
"""Main CLI Entry Point for Strix Hub."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import sys
|
||||
|
||||
from strix_hub.server import serve
|
||||
|
||||
|
||||
def main() -> 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("--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()
|
||||
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG if args.debug else logging.INFO,
|
||||
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
|
||||
)
|
||||
|
||||
serve(host=args.host, port=args.port)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
313
strix_hub/model_router.py
Normal file
313
strix_hub/model_router.py
Normal file
|
|
@ -0,0 +1,313 @@
|
|||
"""Smart Dual-Channel Model Router Gateway for Strix Hub.
|
||||
|
||||
Intercepts requests from Strix and dynamically routes across DIFFERENT providers/channels:
|
||||
- Root Agent (Commander/Orchestrator) -> Dispatches to (root_model, root_api_base, root_api_key)
|
||||
- Subagents (Recon/Fuzzing/Testers) -> Dispatches to (subagent_model, subagent_api_base, subagent_api_key)
|
||||
|
||||
Supports hot-reloading channel URLs, keys, and model identifiers dynamically.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import threading
|
||||
from http import HTTPStatus
|
||||
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
||||
from typing import Any
|
||||
from urllib.error import HTTPError, URLError
|
||||
from urllib.parse import urlparse
|
||||
from urllib.request import Request, urlopen
|
||||
|
||||
logger = logging.getLogger("strix_hub.model_router")
|
||||
|
||||
|
||||
class ModelRouterServer:
|
||||
def __init__(
|
||||
self,
|
||||
host: str = "127.0.0.1",
|
||||
port: int = 18880,
|
||||
root_model: str = "openai/gemini-3.1-pro-preview",
|
||||
root_api_base: str = "",
|
||||
root_api_key: str = "",
|
||||
subagent_model: str = "openai/gemini-3.5-flash",
|
||||
subagent_api_base: str = "",
|
||||
subagent_api_key: str = "",
|
||||
):
|
||||
self.host = host
|
||||
self.port = port
|
||||
self.config_lock = threading.Lock()
|
||||
self.config: dict[str, Any] = {
|
||||
"root_model": root_model,
|
||||
"root_api_base": root_api_base.rstrip("/"),
|
||||
"root_api_key": root_api_key,
|
||||
"subagent_model": subagent_model,
|
||||
"subagent_api_base": subagent_api_base.rstrip("/"),
|
||||
"subagent_api_key": subagent_api_key,
|
||||
}
|
||||
self.server: ThreadingHTTPServer | None = None
|
||||
self.thread: threading.Thread | None = None
|
||||
|
||||
def update_config(self, **kwargs: Any) -> None:
|
||||
"""Hot-reload routing channels and credentials on the fly."""
|
||||
with self.config_lock:
|
||||
for k, v in kwargs.items():
|
||||
if k in self.config and v is not None:
|
||||
if k.endswith("_api_base") and isinstance(v, str):
|
||||
self.config[k] = v.rstrip("/")
|
||||
else:
|
||||
self.config[k] = v
|
||||
logger.info("ModelRouter config hot-updated: %s", self.config)
|
||||
|
||||
def start(self) -> None:
|
||||
handler = _create_router_handler(self)
|
||||
self.server = ThreadingHTTPServer((self.host, self.port), handler)
|
||||
self.thread = threading.Thread(target=self.server.serve_forever, daemon=True)
|
||||
self.thread.start()
|
||||
logger.info("Dual-Channel Model Router Gateway listening on http://%s:%d/v1", self.host, self.port)
|
||||
|
||||
def stop(self) -> None:
|
||||
if self.server:
|
||||
self.server.shutdown()
|
||||
self.server.server_close()
|
||||
self.server = None
|
||||
|
||||
|
||||
def _clean_model_name(model_str: str) -> str:
|
||||
"""Strip litellm/openai prefixes like 'openai/Qwen3.6-35B-A3B' -> 'Qwen3.6-35B-A3B'."""
|
||||
if "/" in model_str:
|
||||
return model_str.split("/", 1)[1]
|
||||
return model_str
|
||||
|
||||
|
||||
def _is_root_agent_payload(body: dict[str, Any]) -> bool:
|
||||
"""Detect if the LLM request originated from Root Agent vs a Sub-agent."""
|
||||
messages = body.get("messages", [])
|
||||
if not isinstance(messages, list):
|
||||
return True
|
||||
|
||||
text_corpus = ""
|
||||
for msg in messages:
|
||||
if isinstance(msg, dict):
|
||||
content = msg.get("content", "")
|
||||
if isinstance(content, str):
|
||||
text_corpus += " " + content
|
||||
elif isinstance(content, list):
|
||||
for part in content:
|
||||
if isinstance(part, dict) and "text" in part:
|
||||
text_corpus += " " + str(part["text"])
|
||||
|
||||
text_lower = text_corpus.lower()
|
||||
if "root agent" in text_lower or "overall mission" in text_lower or "spawn_child_agent" in text_lower:
|
||||
return True
|
||||
if "subagent" in text_lower or "child agent" in text_lower or "agent_finish" in text_lower:
|
||||
return False
|
||||
return True
|
||||
|
||||
|
||||
def _create_router_handler(router_instance: ModelRouterServer) -> type[BaseHTTPRequestHandler]:
|
||||
class RouterHandler(BaseHTTPRequestHandler):
|
||||
def log_message(self, format: str, *args: Any) -> None:
|
||||
logger.debug("Router %s - %s", self.address_string(), format % args)
|
||||
|
||||
def do_OPTIONS(self) -> None:
|
||||
self.send_response(HTTPStatus.NO_CONTENT)
|
||||
self.send_header("Access-Control-Allow-Origin", "*")
|
||||
self.send_header("Access-Control-Allow-Methods", "GET, POST, OPTIONS")
|
||||
self.send_header("Access-Control-Allow-Headers", "*")
|
||||
self.end_headers()
|
||||
|
||||
def do_GET(self) -> None:
|
||||
with router_instance.config_lock:
|
||||
cfg = dict(router_instance.config)
|
||||
|
||||
clean_root = _clean_model_name(cfg["root_model"])
|
||||
clean_sub = _clean_model_name(cfg["subagent_model"])
|
||||
|
||||
if self.path.endswith("/models"):
|
||||
models_payload = {
|
||||
"object": "list",
|
||||
"data": [
|
||||
{"id": clean_root, "object": "model", "owned_by": "strix_hub"},
|
||||
{"id": clean_sub, "object": "model", "owned_by": "strix_hub"},
|
||||
{"id": cfg["root_model"], "object": "model", "owned_by": "strix_hub"},
|
||||
{"id": cfg["subagent_model"], "object": "model", "owned_by": "strix_hub"},
|
||||
],
|
||||
}
|
||||
body = json.dumps(models_payload).encode("utf-8")
|
||||
self.send_response(HTTPStatus.OK)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(body)))
|
||||
self.end_headers()
|
||||
self.wfile.write(body)
|
||||
else:
|
||||
self.send_response(HTTPStatus.OK)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.end_headers()
|
||||
self.wfile.write(b'{"status": "strix_hub_dual_channel_router_ready"}')
|
||||
|
||||
def do_POST(self) -> None:
|
||||
length = int(self.headers.get("Content-Length") or 0)
|
||||
raw_body = self.rfile.read(length) if length else b""
|
||||
|
||||
try:
|
||||
payload = json.loads(raw_body.decode("utf-8")) if raw_body else {}
|
||||
except Exception:
|
||||
payload = {}
|
||||
|
||||
with router_instance.config_lock:
|
||||
cfg = dict(router_instance.config)
|
||||
|
||||
# 1. Determine if this request is for Root Agent or Sub-agent
|
||||
is_root = _is_root_agent_payload(payload) if isinstance(payload, dict) else True
|
||||
|
||||
if is_root:
|
||||
target_model = _clean_model_name(cfg["root_model"])
|
||||
target_base = cfg["root_api_base"]
|
||||
target_key = cfg["root_api_key"]
|
||||
role_label = "Root Agent (主控大脑)"
|
||||
else:
|
||||
target_model = _clean_model_name(cfg["subagent_model"])
|
||||
target_base = cfg["subagent_api_base"]
|
||||
target_key = cfg["subagent_api_key"]
|
||||
role_label = "Subagent (执行打手)"
|
||||
|
||||
# 2. Rewrite model in payload
|
||||
if isinstance(payload, dict):
|
||||
payload["model"] = target_model
|
||||
forward_body = json.dumps(payload).encode("utf-8")
|
||||
else:
|
||||
forward_body = raw_body
|
||||
|
||||
logger.info("ModelRouter: Forwarding [%s] -> Model: %s @ Channel: %s", role_label, target_model, target_base or "Default")
|
||||
|
||||
# 3. Determine Upstream URL
|
||||
req_path = self.path
|
||||
if req_path.startswith("/v1/"):
|
||||
sub_path = req_path[len("/v1") :]
|
||||
else:
|
||||
sub_path = req_path
|
||||
|
||||
upstream_url = f"{target_base}{sub_path}" if target_base else f"https://api.openai.com/v1{sub_path}"
|
||||
|
||||
# 4. Determine Authorization Header
|
||||
auth_header = f"Bearer {target_key}" if target_key else self.headers.get("Authorization", "")
|
||||
|
||||
headers = {
|
||||
"Content-Type": "application/json",
|
||||
"User-Agent": "StrixHub-DualChannelRouter/1.0",
|
||||
}
|
||||
if auth_header:
|
||||
headers["Authorization"] = auth_header
|
||||
|
||||
is_stream = payload.get("stream", False) if isinstance(payload, dict) else False
|
||||
|
||||
try:
|
||||
req = Request(upstream_url, data=forward_body, headers=headers, method="POST")
|
||||
with urlopen(req, timeout=300) as response:
|
||||
self.send_response(response.status)
|
||||
for k, v in response.getheaders():
|
||||
if k.lower() not in ["content-length", "transfer-encoding", "content-encoding"]:
|
||||
self.send_header(k, v)
|
||||
|
||||
if is_stream:
|
||||
self.send_header("Transfer-Encoding", "chunked")
|
||||
self.end_headers()
|
||||
while True:
|
||||
chunk = response.read(4096)
|
||||
if not chunk:
|
||||
break
|
||||
self.wfile.write(f"{len(chunk):X}\r\n".encode("ascii") + chunk + b"\r\n")
|
||||
self.wfile.flush()
|
||||
self.wfile.write(b"0\r\n\r\n")
|
||||
self.wfile.flush()
|
||||
else:
|
||||
raw_data = response.read()
|
||||
try:
|
||||
json_obj = json.loads(raw_data.decode("utf-8"))
|
||||
if isinstance(json_obj, dict) and "choices" in json_obj and json_obj["choices"]:
|
||||
msg = json_obj["choices"][0].get("message", {})
|
||||
content = msg.get("content", "") or ""
|
||||
if not msg.get("tool_calls") and "<tool_call>" in content:
|
||||
extracted = _parse_qwen_tool_calls(content)
|
||||
if extracted:
|
||||
msg["tool_calls"] = extracted
|
||||
json_obj["choices"][0]["finish_reason"] = "tool_calls"
|
||||
logger.info("ModelRouter: Auto-extracted %d tool calls from Qwen text output!", len(extracted))
|
||||
raw_data = json.dumps(json_obj, ensure_ascii=False).encode("utf-8")
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
self.send_header("Content-Length", str(len(raw_data)))
|
||||
self.end_headers()
|
||||
self.wfile.write(raw_data)
|
||||
except HTTPError as e:
|
||||
err_content = e.read()
|
||||
self.send_response(e.code)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(err_content)))
|
||||
self.end_headers()
|
||||
self.wfile.write(err_content)
|
||||
except Exception as exc:
|
||||
logger.exception("Router forwarding error to %s", upstream_url)
|
||||
err_msg = json.dumps({"error": {"message": str(exc), "type": "dual_channel_router_error"}}).encode("utf-8")
|
||||
self.send_response(HTTPStatus.BAD_GATEWAY)
|
||||
self.send_header("Content-Type", "application/json")
|
||||
self.send_header("Content-Length", str(len(err_msg)))
|
||||
self.end_headers()
|
||||
self.wfile.write(err_msg)
|
||||
|
||||
return RouterHandler
|
||||
|
||||
|
||||
def _parse_qwen_tool_calls(content: str) -> list[dict[str, Any]]:
|
||||
"""Extract tool calls from plain text Qwen XML/JSON tags."""
|
||||
if not content or "<tool_call>" not in content:
|
||||
return []
|
||||
|
||||
import re
|
||||
import uuid
|
||||
|
||||
tool_calls: list[dict[str, Any]] = []
|
||||
matches = re.findall(r"<tool_call>(.*?)</tool_call>", content, re.DOTALL)
|
||||
for block in matches:
|
||||
block = block.strip()
|
||||
func_match = re.search(r"<function=([a-zA-Z0-9_-]+)>", block)
|
||||
if func_match:
|
||||
func_name = func_match.group(1)
|
||||
params: dict[str, Any] = {}
|
||||
param_matches = re.findall(r"<parameter=([a-zA-Z0-9_-]+)>\s*(.*?)\s*</parameter>", block, re.DOTALL)
|
||||
for k, v in param_matches:
|
||||
v = v.strip()
|
||||
try:
|
||||
params[k] = json.loads(v)
|
||||
except Exception:
|
||||
params[k] = v
|
||||
tool_calls.append({
|
||||
"id": f"call_{uuid.uuid4().hex[:8]}",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": func_name,
|
||||
"arguments": json.dumps(params, ensure_ascii=False),
|
||||
},
|
||||
})
|
||||
continue
|
||||
|
||||
clean_json = re.sub(r"^```(json)?|```$", "", block, flags=re.MULTILINE).strip()
|
||||
try:
|
||||
parsed = json.loads(clean_json)
|
||||
if isinstance(parsed, dict) and "name" in parsed:
|
||||
args = parsed.get("arguments", {})
|
||||
args_str = json.dumps(args, ensure_ascii=False) if isinstance(args, dict) else str(args)
|
||||
tool_calls.append({
|
||||
"id": f"call_{uuid.uuid4().hex[:8]}",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": parsed["name"],
|
||||
"arguments": args_str,
|
||||
},
|
||||
})
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return tool_calls
|
||||
1661
strix_hub/server.py
Normal file
1661
strix_hub/server.py
Normal file
File diff suppressed because it is too large
Load diff
342
strix_hub/task_manager.py
Normal file
342
strix_hub/task_manager.py
Normal file
|
|
@ -0,0 +1,342 @@
|
|||
"""Task Lifecycle & Process Supervisor for Strix Hub.
|
||||
|
||||
Controls Strix scans via subprocesses and OS signals (SIGSTOP, SIGCONT, SIGTERM).
|
||||
Supports Dual-Channel Model Routing and Hot-Reloading.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import signal
|
||||
import subprocess
|
||||
import threading
|
||||
import time
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from strix_hub import db
|
||||
from strix_hub.model_router import ModelRouterServer
|
||||
|
||||
logger = logging.getLogger("strix_hub.task_manager")
|
||||
|
||||
STRIX_RUNS_DIR = Path(os.environ.get("STRIX_RUNS_DIR", "/opt/strix/strix_runs"))
|
||||
if not STRIX_RUNS_DIR.exists():
|
||||
try:
|
||||
STRIX_RUNS_DIR.mkdir(parents=True, exist_ok=True)
|
||||
except Exception:
|
||||
STRIX_RUNS_DIR = Path.cwd() / "strix_runs"
|
||||
STRIX_RUNS_DIR.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
# Active runners in memory: {task_id: {"process": Popen, "router": ModelRouterServer, "start_time": float}}
|
||||
_ACTIVE_TASKS: dict[str, dict[str, Any]] = {}
|
||||
_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def get_strix_bin_path() -> str:
|
||||
"""Find the strix executable binary path."""
|
||||
venv_strix = Path("/opt/strix/.venv/bin/strix")
|
||||
if venv_strix.is_file():
|
||||
return str(venv_strix)
|
||||
uv_bin = Path("/usr/local/bin/uv")
|
||||
if uv_bin.is_file():
|
||||
return "/usr/local/bin/uv run strix"
|
||||
return "strix"
|
||||
|
||||
|
||||
def start_task(task_id: str) -> dict[str, Any]:
|
||||
"""Launch a pentest task as a background subprocess with dual-channel routing."""
|
||||
task = db.get_task_full(task_id)
|
||||
if not task:
|
||||
raise ValueError("Task not found")
|
||||
|
||||
with _LOCK:
|
||||
if task_id in _ACTIVE_TASKS:
|
||||
existing = _ACTIVE_TASKS[task_id].get("process")
|
||||
if existing and existing.poll() is None:
|
||||
return db.get_task_by_id(task_id) or {}
|
||||
|
||||
# 1. Start dedicated Dual-Channel ModelRouter on a free port for this task
|
||||
port = 18800 + (abs(hash(task_id)) % 1000)
|
||||
|
||||
# Fallback to server env if task channel is left blank
|
||||
default_base = os.environ.get("OPENAI_BASE_URL") or os.environ.get("LLM_API_BASE", "https://api.openai.com/v1")
|
||||
default_key = os.environ.get("OPENAI_API_KEY") or os.environ.get("LLM_API_KEY", "")
|
||||
|
||||
root_base = task.get("root_api_base") or task.get("api_base") or default_base
|
||||
root_key = task.get("root_api_key_raw") or task.get("api_key_raw") or default_key
|
||||
|
||||
sub_base = task.get("subagent_api_base") or task.get("api_base") or default_base
|
||||
sub_key = task.get("subagent_api_key_raw") or task.get("api_key_raw") or default_key
|
||||
|
||||
router = ModelRouterServer(
|
||||
host="127.0.0.1",
|
||||
port=port,
|
||||
root_model=task["root_model"],
|
||||
root_api_base=root_base,
|
||||
root_api_key=root_key,
|
||||
subagent_model=task["subagent_model"],
|
||||
subagent_api_base=sub_base,
|
||||
subagent_api_key=sub_key,
|
||||
)
|
||||
router.start()
|
||||
|
||||
# 2. Build subprocess environment (directs LiteLLM / OpenAI to local ModelRouter)
|
||||
env = dict(os.environ)
|
||||
env["OPENAI_BASE_URL"] = f"http://127.0.0.1:{port}/v1"
|
||||
env["LLM_API_BASE"] = f"http://127.0.0.1:{port}/v1"
|
||||
root_m = task["root_model"]
|
||||
env["STRIX_LLM"] = root_m if root_m.startswith("openai/") else f"openai/{root_m}"
|
||||
# A dummy key to satisfy litellm validation if empty
|
||||
env["OPENAI_API_KEY"] = root_key or "strix-hub-key"
|
||||
env["LLM_API_KEY"] = root_key or "strix-hub-key"
|
||||
# Disable streaming so ModelRouter can accurately parse & auto-recover Qwen text tool-calls
|
||||
env["LLM_DISABLE_STREAMING"] = "true"
|
||||
|
||||
# 3. Assemble command line arguments
|
||||
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}"
|
||||
if instruction:
|
||||
clean_inst = instruction.replace('"', '\\"')
|
||||
cmd += f' --instruction "{clean_inst}"'
|
||||
|
||||
# Work directory
|
||||
work_dir = "/opt/strix" if Path("/opt/strix").is_dir() else str(Path.cwd())
|
||||
log_file_path = STRIX_RUNS_DIR / f"{task_id}.log"
|
||||
log_f = open(log_file_path, "w", encoding="utf-8")
|
||||
|
||||
# Spawn in a new process group so SIGSTOP/SIGCONT pauses all children (Docker/tools)
|
||||
proc = subprocess.Popen(
|
||||
cmd,
|
||||
shell=True,
|
||||
cwd=work_dir,
|
||||
env=env,
|
||||
stdout=log_f,
|
||||
stderr=subprocess.STDOUT,
|
||||
preexec_fn=os.setsid if hasattr(os, "setsid") else None,
|
||||
)
|
||||
|
||||
start_time = time.time()
|
||||
_ACTIVE_TASKS[task_id] = {
|
||||
"process": proc,
|
||||
"router": router,
|
||||
"start_time": start_time,
|
||||
"log_file": log_f,
|
||||
}
|
||||
|
||||
db.update_task_status(task_id, status="running", pid=proc.pid)
|
||||
|
||||
# Spawn background monitor thread for this task
|
||||
threading.Thread(target=_monitor_task_process, args=(task_id,), daemon=True).start()
|
||||
return db.get_task_by_id(task_id) or {}
|
||||
|
||||
|
||||
def hot_update_task_config(
|
||||
task_id: str,
|
||||
root_model: str | None = None,
|
||||
root_api_base: str | None = None,
|
||||
root_api_key: str | None = None,
|
||||
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
|
||||
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,
|
||||
)
|
||||
|
||||
# 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)
|
||||
|
||||
return True
|
||||
|
||||
|
||||
def pause_task(task_id: str) -> bool:
|
||||
"""Pause task execution using SIGSTOP."""
|
||||
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.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:
|
||||
pass
|
||||
|
||||
if info.get("router"):
|
||||
info["router"].stop()
|
||||
if info.get("log_file") and not info["log_file"].closed:
|
||||
info["log_file"].close()
|
||||
|
||||
_ACTIVE_TASKS.pop(task_id, None)
|
||||
db.update_task_status(task_id, status="stopped")
|
||||
logger.info("Stopped task %s", task_id)
|
||||
return True
|
||||
|
||||
|
||||
def _monitor_task_process(task_id: str) -> None:
|
||||
"""Monitor task process completion, update status and parse 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"]
|
||||
|
||||
while proc.poll() is None:
|
||||
time.sleep(1.0)
|
||||
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"
|
||||
new_status = "paused" if current_status == "paused" else "running"
|
||||
db.update_task_status(
|
||||
task_id,
|
||||
status=new_status,
|
||||
duration_seconds=elapsed,
|
||||
run_dir_name=run_dir_name,
|
||||
vulns_count=vulns_cnt,
|
||||
)
|
||||
|
||||
exit_code = proc.returncode
|
||||
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"
|
||||
db.update_task_status(
|
||||
task_id,
|
||||
status=final_status,
|
||||
duration_seconds=elapsed,
|
||||
run_dir_name=run_dir_name,
|
||||
vulns_count=vulns_cnt,
|
||||
)
|
||||
|
||||
router.stop()
|
||||
if info.get("log_file") and not info["log_file"].closed:
|
||||
info["log_file"].close()
|
||||
|
||||
with _LOCK:
|
||||
_ACTIVE_TASKS.pop(task_id, None)
|
||||
logger.info("Task %s completed with status [%s] (code %d)", 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."""
|
||||
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,
|
||||
)
|
||||
|
||||
if not runs:
|
||||
return None, 0
|
||||
|
||||
latest = runs[0]
|
||||
vuln_count = 0
|
||||
vulns_file = latest / "vulnerabilities.json"
|
||||
if vulns_file.is_file():
|
||||
try:
|
||||
with open(vulns_file, "r", encoding="utf-8") as f:
|
||||
vulns = json.load(f)
|
||||
if isinstance(vulns, list):
|
||||
vuln_count = len(vulns)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
return latest.name, vuln_count
|
||||
|
||||
|
||||
def get_task_logs(task_id: str, max_lines: int = 100) -> str:
|
||||
"""Read the latest log output for a task."""
|
||||
log_path = STRIX_RUNS_DIR / f"{task_id}.log"
|
||||
if not log_path.is_file():
|
||||
return "No log output available yet."
|
||||
try:
|
||||
with open(log_path, "r", encoding="utf-8", errors="replace") as f:
|
||||
lines = f.readlines()
|
||||
return "".join(lines[-max_lines:])
|
||||
except Exception as e:
|
||||
return f"Error reading logs: {e}"
|
||||
Loading…
Add table
Reference in a new issue