mirror of
https://github.com/Ed1s0nZ/CyberStrikeAI.git
synced 2026-08-01 08:37:41 +02:00
Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
a496bd7aee | ||
|
|
e98da9d1b3 | ||
|
|
a1a62657af | ||
|
|
7a3e74a3af | ||
|
|
a36acddcf4 | ||
|
|
4016eab95f | ||
|
|
51efdd7b60 | ||
|
|
c24fbf602c | ||
|
|
4b1f71b4a6 | ||
|
|
29bf0ff4da | ||
|
|
ccfb85e6cb | ||
|
|
efa8b262b1 | ||
|
|
bbbe77e90c | ||
|
|
6bafb8fe70 | ||
|
|
7d1e16b97b | ||
|
|
1be10cdf2c | ||
|
|
95909ce999 | ||
|
|
a603adb467 | ||
|
|
acfacfe1b3 | ||
|
|
2e5c1ff286 | ||
|
|
4ccb330ded | ||
|
|
a28dad4827 | ||
|
|
a6b7c7e7be | ||
|
|
e5a4703e91 | ||
|
|
8ea0ee4f1b | ||
|
|
386ec7b835 | ||
|
|
078f5cb222 | ||
|
|
7351793683 | ||
|
|
68e3ead4d7 | ||
|
|
597712a7b9 | ||
|
|
f77cc09477 | ||
|
|
5423bd5e1e | ||
|
|
a64c18df6d | ||
|
|
33b0b56514 | ||
|
|
23ec222b77 | ||
|
|
d36352e4dc | ||
|
|
ea22b1f3ba | ||
|
|
8a3229aa5a | ||
|
|
a497c4bfcd | ||
|
|
bd9116e9c3 | ||
|
|
6f39720669 | ||
|
|
f8481024ed | ||
|
|
6b3a7d81d2 | ||
|
|
25cf3c567b | ||
|
|
41f683ce6c | ||
|
|
ffae94fb2c | ||
|
|
fe3c845ff8 | ||
|
|
0f1a6ad25a | ||
|
|
3b3f73461b | ||
|
|
3ce80fd00f | ||
|
|
08cb8a68fb | ||
|
|
2f38693891 | ||
|
|
6fa17a3093 | ||
|
|
46cea9459a | ||
|
|
987bd0a03c | ||
|
|
6529c13ffb | ||
|
|
f2ba5093d7 | ||
|
|
3c6cf633e1 | ||
|
|
142977413e | ||
|
|
62efc81993 | ||
|
|
9ac5fd33ec | ||
|
|
d1d67b07d3 | ||
|
|
211c36654a | ||
|
|
3894ba6054 | ||
|
|
fa0dd6c721 | ||
|
|
1cf10981ef | ||
|
|
79d162ccb4 | ||
|
|
722806797f | ||
|
|
2fcda3a57d | ||
|
|
bec296ae3f | ||
|
|
612015b6d7 |
@@ -7,31 +7,14 @@
|
|||||||
|
|
||||||
[中文](README_CN.md) | [English](README.md)
|
[中文](README_CN.md) | [English](README.md)
|
||||||
|
|
||||||
**Community**: [Join us on Discord](https://discord.gg/8PjVCMu8Zw)
|
**The system of action for AI-native cybersecurity—where intent becomes governed execution, evidence becomes operational memory, and every operation improves the next.**
|
||||||
|
|
||||||
**CyberStrikeAI is building the agentic execution layer for modern cyber security.**
|
CyberStrikeAI connects planning, execution, human oversight, evidence, and replay in one auditable workspace. Built in Go, it combines Eino-powered agents, MCP-native tools, RAG knowledge, graph workflows, and attack-chain modeling and analysis for authorized security operations.
|
||||||
|
|
||||||
It brings AI agents, security tools, MCP-native integrations, knowledge systems, human oversight, and attack-chain intelligence into a unified workspace for authorized cyber engagements. Instead of treating tools, prompts, evidence, approvals, and reports as separate fragments, CyberStrikeAI turns security intent into auditable multi-agent workflows that can plan, execute, review, replay, and continuously accumulate operational context.
|
**Start here:** [Quick start](#quick-start-one-command-deployment) · [Documentation](docs/en-US/README.md) · [Security hardening](docs/en-US/security-hardening.md)
|
||||||
|
|
||||||
Built in Go, CyberStrikeAI provides a full-stack foundation for AI-native security operations: 100+ curated tool recipes, role-based testing, Agent Skills, Eino-powered single-agent and multi-agent orchestration, RAG knowledge retrieval, graph workflows, vulnerability and task lifecycle management, WebShell operations, chatbot access, and a lightweight built-in C2 framework for authorized lab and engagement scenarios.
|
> [!IMPORTANT]
|
||||||
|
> Use CyberStrikeAI only on systems you own or are explicitly authorized to test. For shared or production environments, review the [security model](docs/en-US/security-model.md) and [hardening guide](docs/en-US/security-hardening.md) before enabling high-risk tools, WebShell, or C2 capabilities.
|
||||||
<details>
|
|
||||||
<summary><strong>WeChat group</strong> (click to reveal QR code)</summary>
|
|
||||||
|
|
||||||
<img src="./images/wechat-group-cyberstrikeai-qr.jpg" alt="CyberStrikeAI WeChat group QR code" width="280">
|
|
||||||
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><strong>Sponsorship</strong> (click to expand)</summary>
|
|
||||||
|
|
||||||
If CyberStrikeAI helps you, you can support the project via **WeChat Pay** or **Alipay**:
|
|
||||||
|
|
||||||
<div align="center">
|
|
||||||
<img src="./images/sponsor-wechat-alipay-qr.jpg" alt="WeChat Pay and Alipay sponsorship QR codes" width="480">
|
|
||||||
</div>
|
|
||||||
|
|
||||||
</details>
|
|
||||||
|
|
||||||
## Interface & Integration Preview
|
## Interface & Integration Preview
|
||||||
|
|
||||||
@@ -54,6 +37,9 @@ If CyberStrikeAI helps you, you can support the project via **WeChat Pay** or **
|
|||||||
|
|
||||||
*The dashboard provides a comprehensive overview of system runtime status, security vulnerabilities, tool usage, and knowledge base, helping users quickly understand the platform's core features and current state.*
|
*The dashboard provides a comprehensive overview of system runtime status, security vulnerabilities, tool usage, and knowledge base, helping users quickly understand the platform's core features and current state.*
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><strong>More interface screenshots</strong></summary>
|
||||||
|
|
||||||
### Core Features Overview
|
### Core Features Overview
|
||||||
|
|
||||||
<table>
|
<table>
|
||||||
@@ -115,33 +101,48 @@ If CyberStrikeAI helps you, you can support the project via **WeChat Pay** or **
|
|||||||
</tr>
|
</tr>
|
||||||
</table>
|
</table>
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
## Highlights
|
## Highlights
|
||||||
|
|
||||||
- 🤖 Agentic execution layer for translating natural-language intent into precise, governed, auditable security action
|
### Agents and orchestration
|
||||||
- 🧩 Eino-powered single-agent and multi-agent orchestration with Deep, Plan-Execute, and Supervisor modes
|
|
||||||
- 🔌 MCP-native tool execution with HTTP/stdio/SSE transports, external MCP federation, and dynamic tool discovery
|
- 🤖 **Agentic execution** translates natural-language intent into governed, auditable security actions.
|
||||||
- 🧰 100+ curated security tool recipes, YAML-based extensions, and role-scoped tool control
|
- 🧩 **Eino orchestration** supports single-agent execution plus Deep, Plan-Execute, and Supervisor multi-agent modes.
|
||||||
- 📄 Large-result pagination, compression, and searchable archives
|
- 🔀 **Graph workflows** combine Agents, tools, conditions, approvals, and outputs into reusable flows.
|
||||||
- 🔗 Attack-chain intelligence with graph views, risk scoring, project facts, and step-by-step replay
|
- 🎭 **Role-based testing** provides focused prompts and tool policies for common security scenarios.
|
||||||
- 🧑⚖️ Human-in-the-loop governance with approval modes, allowlists, audit-agent review, and traceable decisions
|
|
||||||
- 🔒 Password-protected web UI, audit logs, SQLite persistence, and operational evidence retention
|
### Tools and knowledge
|
||||||
- 🔐 **Platform RBAC** with multi-user accounts, system/custom roles, per-permission scopes (`all` / `assigned` / `own`), ownership, and explicit assignments enforced across APIs, Agents, MCP, background jobs, and chatbots; see the [RBAC administration guide](docs/en-US/rbac.md)
|
|
||||||
- 📚 Knowledge base (RAG): **Eino MultiQuery** query rewrite + multi-path vector retrieval + **HTTP rerank** (DashScope `gte-rerank` / Cohere-compatible) + post-processing (dedupe, budget); **Eino Compose** indexing pipeline
|
- 🧰 **Security tools** include 100+ curated YAML recipes with custom extensions and role-scoped access.
|
||||||
- 📁 Conversation grouping with pinning, rename, and batch management
|
- 🔌 **MCP integration** supports HTTP, stdio, SSE, external federation, and dynamic tool discovery.
|
||||||
- 📂 **Project management**: shared facts (blackboard) across sessions, `upsert_project_fact` + `links` to chain paths; attack-chain and project fact graph views
|
- 🎯 **Agent Skills** follow the standard Skill layout and support progressive, on-demand loading.
|
||||||
- 🛡️ Vulnerability management with CRUD operations, severity tracking, status workflow, and statistics
|
- 📚 **Knowledge base** combines query rewriting, vector retrieval, reranking, and result post-processing.
|
||||||
- 📋 Batch task management: create task queues, add multiple tasks, and execute them sequentially
|
- 🖼️ **Vision analysis** uses a separate vision model for screenshots, captchas, and UI while retaining text summaries only.
|
||||||
- 🎭 Role-based testing: predefined security testing roles (Penetration Testing, CTF, Web App Scanning, etc.) with custom prompts and tool restrictions
|
|
||||||
- 🔀 **Graph orchestration**: visual workflow editor (Start / Agent / Tool / Condition / HITL / Output) with `{{previous.output}}` and `{{outputs.variable_name}}` for inter-node data passing; bind a graph to a role for automatic execution on chat. See [Graph orchestration guide](docs/en-US/workflow-graph.md)
|
### Governance and audit
|
||||||
- 🧩 **Agent orchestration (CloudWeGo Eino)**: **single-agent** via **`/api/eino-agent/stream`** (Eino ADK `ChatModelAgent`); **multi-agent** via **`/api/multi-agent/stream`** with **`deep`** (coordinator + `task` sub-agents), **`plan_execute`**, or **`supervisor`** (`orchestration` in the request body). ADK **summarization** compresses long contexts; pre-compaction **transcripts** land at `data/conversation_artifacts/<conversation-id>/summarization/transcript.txt` (full user/assistant/tool turns; static system omitted). Markdown under `agents/`: `orchestrator.md`, `orchestrator-plan-execute.md`, `orchestrator-supervisor.md`, plus sub-agent `*.md` (see [Multi-agent doc](docs/en-US/MULTI_AGENT_EINO.md))
|
|
||||||
- 🖼️ **Vision analysis (`analyze_image`)**: separate VL model (e.g. `qwen-vl-max`) via MCP for local screenshots, captchas, and UI; image bytes stay out of agent history (text summaries only). Configure `vision` in `config.yaml`; see [docs/en-US/VISION.md](docs/en-US/VISION.md)
|
- 🧑⚖️ **Human in the loop** provides approval modes, tool allowlists, audit-agent review, and traceable decisions.
|
||||||
- 🎯 **Skills (refactored for Eino)**: packs under `skills_dir` follow **Agent Skills** layout (`SKILL.md` + optional files); **multi-agent** sessions use the official Eino ADK **`skill`** tool for **progressive disclosure** (load by name), with optional **host filesystem / shell** via `multi_agent.eino_skills`; optional **`eino_middleware`** adds patchtoolcalls, tool_search, **plantask** (`TaskCreate` / `TaskList` boards under `skills_dir/.eino/plantask/`), reduction, file **checkpoints** (`checkpoint_dir`), ChatModel **retries**, session **output key**, and Deep tuning—20+ sample domains (SQLi, XSS, API security, …) ship under `skills/`
|
- 🔐 **Platform RBAC** supports multiple users, system and custom roles, scoped permissions, ownership, and explicit assignments.
|
||||||
- 📱 **Chatbot**: Personal WeChat, WeCom, DingTalk, Lark, Telegram, Slack, Discord, and QQ Bot—chat from mobile or IM apps (see [Robot / Chatbot guide](docs/en-US/robot.md))
|
- 🔒 **Security and audit** provide authenticated access, audit logs, SQLite persistence, and operational evidence retention.
|
||||||
- 🧑⚖️ **Human-in-the-loop (HITL)**: Chat sidebar to set approval mode and tool allowlists (listed tools skip approval); global list in `config.yaml` under `hitl.tool_whitelist`; the Audit Agent can use a separate lightweight model via `hitl.audit_model`; **Apply** can merge new tools into the file and update the running server without restart; dedicated **HITL** page for pending approvals. See [HITL best practices](docs/en-US/hitl-best-practices.md)
|
- 📄 **Result governance** supports pagination, compression, archival, and search for large tool outputs.
|
||||||
- 🐚 **WebShell management**: Add and manage WebShell connections (e.g. IceSword/AntSword compatible), use a virtual terminal for command execution, a built-in file manager for file operations, and an AI assistant tab that orchestrates tests and keeps per-connection conversation history; supports PHP, ASP, ASPX, JSP and custom shell types with configurable request method and command parameter.
|
|
||||||
- 📡 **Built-in C2**: AI-oriented lightweight command-and-control—**listeners** (TCP reverse, HTTP/HTTPS beacon, WebSocket), **encrypted** beacon channel, **session** and **task** queues with persistence, **payload** helpers (one-liner / build / download), **SSE** live events, REST under `/api/c2/*`, plus unified MCP tools (`c2_listener`, `c2_session`, **`c2_task`**, `c2_task_manage`, `c2_payload`, `c2_event`, `c2_profile`, `c2_file`); optional **HITL** approval for sensitive operations and OPSEC-style controls (e.g. command deny rules). **Authorized testing only.**
|
### Security operations
|
||||||
|
|
||||||
|
- 📁 **Conversation management** provides grouping, pinning, renaming, and batch organization.
|
||||||
|
- 📂 **Projects and attack chains** connect cross-session facts, risk scoring, graph views, and step-by-step replay.
|
||||||
|
- 🛡️ **Vulnerability management** provides severity classification, lifecycle tracking, filtering, and statistics.
|
||||||
|
- 📋 **Batch tasks** provide queued execution, editing, status tracking, and retained results.
|
||||||
|
- 📱 **Chatbots** connect Personal WeChat, WeCom, DingTalk, Lark, Telegram, Slack, Discord, and QQ Bot.
|
||||||
|
|
||||||
|
### Authorized security operations
|
||||||
|
|
||||||
|
- 🐚 **WebShell management** provides connection management, a virtual terminal, file operations, and AI-assisted workflows.
|
||||||
|
- 📡 **Built-in C2** provides listeners, encrypted beacons, sessions, task queues, payload helpers, and live events.
|
||||||
|
|
||||||
|
> WebShell, C2, and other high-risk capabilities are for systems you own or are explicitly authorized to test. See the [security model](docs/en-US/security-model.md) and [hardening guide](docs/en-US/security-hardening.md).
|
||||||
|
|
||||||
## Plugins
|
## Plugins
|
||||||
|
|
||||||
@@ -160,6 +161,9 @@ CyberStrikeAI includes optional integrations under `plugins/`.
|
|||||||
|
|
||||||
CyberStrikeAI ships with 100+ curated tools covering the whole kill chain:
|
CyberStrikeAI ships with 100+ curated tools covering the whole kill chain:
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><strong>View the complete tool categories</strong></summary>
|
||||||
|
|
||||||
- **Network Scanners** – nmap, masscan, rustscan, arp-scan, nbtscan
|
- **Network Scanners** – nmap, masscan, rustscan, arp-scan, nbtscan
|
||||||
- **Web & App Scanners** – sqlmap, nikto, dirb, gobuster, feroxbuster, ffuf, httpx
|
- **Web & App Scanners** – sqlmap, nikto, dirb, gobuster, feroxbuster, ffuf, httpx
|
||||||
- **Vulnerability Scanners** – nuclei, wpscan, wafw00f, dalfox, xsser
|
- **Vulnerability Scanners** – nuclei, wpscan, wafw00f, dalfox, xsser
|
||||||
@@ -176,12 +180,16 @@ CyberStrikeAI ships with 100+ curated tools covering the whole kill chain:
|
|||||||
- **CTF Utilities** – stegsolve, zsteg, hash-identifier, fcrackzip, pdfcrack, cyberchef
|
- **CTF Utilities** – stegsolve, zsteg, hash-identifier, fcrackzip, pdfcrack, cyberchef
|
||||||
- **System Helpers** – exec, create-file, delete-file, list-files, modify-file
|
- **System Helpers** – exec, create-file, delete-file, list-files, modify-file
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
|
See [tools/README_EN.md](tools/README_EN.md) for tool definitions, customization, and usage notes.
|
||||||
|
|
||||||
## Basic Usage
|
## Basic Usage
|
||||||
|
|
||||||
### Quick Start (One-Command Deployment)
|
### Quick Start (One-Command Deployment)
|
||||||
|
|
||||||
**Prerequisites:**
|
**Prerequisites:**
|
||||||
- Go 1.21+ ([Install](https://go.dev/dl/))
|
- Go 1.25+ ([Install](https://go.dev/dl/); required by `go.mod`)
|
||||||
- Python 3.10+ ([Install](https://www.python.org/downloads/))
|
- Python 3.10+ ([Install](https://www.python.org/downloads/))
|
||||||
|
|
||||||
**One-Command Deployment:**
|
**One-Command Deployment:**
|
||||||
@@ -199,6 +207,12 @@ The `run.sh` script will automatically:
|
|||||||
- ✅ Build the project
|
- ✅ Build the project
|
||||||
- ✅ Start the server
|
- ✅ Start the server
|
||||||
|
|
||||||
|
**Verify the startup:**
|
||||||
|
|
||||||
|
1. Confirm the terminal displays `● ONLINE` followed by the actual Web UI URL.
|
||||||
|
2. Open that URL; the default HTTPS mode uses a local self-signed certificate, so accept the browser warning once.
|
||||||
|
3. On a new installation, store the one-time `admin` password shown under `ADMIN SETUP REQUIRED`, sign in, and change it immediately.
|
||||||
|
|
||||||
**Networking defaults:** `run.sh` starts the server with **`--https`** and the repo **`config.yaml`** (local self-signed TLS; better for many concurrent streams). Use **`./run.sh --http`** for plain HTTP. In production, set **`server.tls_cert_path`** / **`server.tls_key_path`** in **`config.yaml`** (see comments there). For manual runs, add **`--https`** or **`CYBERSTRIKE_HTTPS=1`**; if **`-config`** is wrong, the binary prints a short usage hint on stderr.
|
**Networking defaults:** `run.sh` starts the server with **`--https`** and the repo **`config.yaml`** (local self-signed TLS; better for many concurrent streams). Use **`./run.sh --http`** for plain HTTP. In production, set **`server.tls_cert_path`** / **`server.tls_key_path`** in **`config.yaml`** (see comments there). For manual runs, add **`--https`** or **`CYBERSTRIKE_HTTPS=1`**; if **`-config`** is wrong, the binary prints a short usage hint on stderr.
|
||||||
|
|
||||||
**First-Time Configuration:**
|
**First-Time Configuration:**
|
||||||
@@ -207,12 +221,12 @@ The `run.sh` script will automatically:
|
|||||||
- Go to `Settings` → Fill in your API credentials:
|
- Go to `Settings` → Fill in your API credentials:
|
||||||
```yaml
|
```yaml
|
||||||
openai:
|
openai:
|
||||||
api_key: "sk-your-key"
|
api_key: "${OPENAI_API_KEY}"
|
||||||
base_url: "https://api.openai.com/v1" # or https://api.deepseek.com/v1
|
base_url: "https://api.openai.com/v1" # or https://api.deepseek.com/v1
|
||||||
model: "gpt-4o" # or deepseek-chat, claude-3-opus, etc.
|
model: "gpt-4o" # or deepseek-chat, claude-3-opus, etc.
|
||||||
```
|
```
|
||||||
- Or edit `config.yaml` directly before launching
|
- Or edit `config.yaml` directly before launching
|
||||||
2. **Login** - Use the auto-generated password shown in the console (or set `auth.password` in `config.yaml`)
|
2. **Login** - On first startup the console prints an auto-generated initial `admin` password; create accounts from **Platform permissions → User management**
|
||||||
3. **Install security tools (optional)** - Install tools from `tools/` as needed; missing tools are skipped or substituted at runtime. Common examples:
|
3. **Install security tools (optional)** - Install tools from `tools/` as needed; missing tools are skipped or substituted at runtime. Common examples:
|
||||||
|
|
||||||
**macOS (Homebrew):**
|
**macOS (Homebrew):**
|
||||||
@@ -243,9 +257,9 @@ If server logs show `client sent an HTTP request to an HTTPS server`, a client i
|
|||||||
|
|
||||||
**Note:** The Python virtual environment (`venv/`) is automatically created and managed by `run.sh`. Tools that require Python (like `api-fuzzer`, `http-framework-test`, etc.) will automatically use this environment.
|
**Note:** The Python virtual environment (`venv/`) is automatically created and managed by `run.sh`. Tools that require Python (like `api-fuzzer`, `http-framework-test`, etc.) will automatically use this environment.
|
||||||
|
|
||||||
### Version Update (No Breaking Changes)
|
### Upgrade and Compatibility
|
||||||
|
|
||||||
**CyberStrikeAI one-click upgrade (recommended):**
|
**CyberStrikeAI one-click upgrade:**
|
||||||
1. (First time) enable the script: `chmod +x upgrade.sh`
|
1. (First time) enable the script: `chmod +x upgrade.sh`
|
||||||
2. Upgrade with: `./upgrade.sh` (optional flags: `--tag vX.Y.Z`, `--no-venv`, `--yes`). Local `tools/`, `roles/`, and `skills/` are always preserved.
|
2. Upgrade with: `./upgrade.sh` (optional flags: `--tag vX.Y.Z`, `--no-venv`, `--yes`). Local `tools/`, `roles/`, and `skills/` are always preserved.
|
||||||
3. The script will back up your `config.yaml` and `data/`, upgrade the code from GitHub Release, update `config.yaml`'s `version`, then restart the server.
|
3. The script will back up your `config.yaml` and `data/`, upgrade the code from GitHub Release, update `config.yaml`'s `version`, then restart the server.
|
||||||
@@ -260,402 +274,32 @@ Requirements / tips:
|
|||||||
* `rsync` is recommended/required for the safe code sync.
|
* `rsync` is recommended/required for the safe code sync.
|
||||||
* If GitHub API rate-limits you, set `export GITHUB_TOKEN="..."` before running `./upgrade.sh`.
|
* If GitHub API rate-limits you, set `export GITHUB_TOKEN="..."` before running `./upgrade.sh`.
|
||||||
|
|
||||||
⚠️ **Note:** This procedure only applies to version updates without compatibility or breaking changes. If a release includes compatibility changes, this method may not apply.
|
⚠️ **Before upgrading:** review the target release notes for configuration, database, and API changes. Backups are required even for patch upgrades; a version number alone is not a compatibility guarantee.
|
||||||
|
|
||||||
**Examples:** No breaking changes — e.g. v1.3.1 → v1.3.2; with breaking changes — e.g. v1.3.1 → v1.4.0. The project follows [Semantic Versioning](https://semver.org/) (SemVer): when only the patch version (third number) changes, this upgrade path is usually safe; when the minor or major version changes, config, data, or APIs may have changed — check the release notes before using this method.
|
|
||||||
|
|
||||||
### Core Workflows
|
|
||||||
- **Conversation testing** – Natural-language prompts trigger toolchains with streaming SSE output.
|
|
||||||
- **Single vs multi-agent** – Chat UI switches between **Eino single-agent** (`/api/eino-agent/stream`) and **multi-agent** (`/api/multi-agent/stream` with `orchestration`: `deep` | `plan_execute` | `supervisor`). Multi mode requires `multi_agent.enabled: true`. MCP tools are bridged the same way for both paths.
|
|
||||||
- **Role-based testing** – Select from predefined security testing roles (Penetration Testing, CTF, Web App Scanning, API Security Testing, etc.) to customize AI behavior and tool availability. Each role applies custom system prompts and can restrict available tools for focused testing scenarios.
|
|
||||||
- **Graph orchestration** – Design flows on the **Graph Orchestration** page (drag nodes, connect edges, save); bind `workflow_id` on a role to run the graph on chat (Agent, MCP tools, condition branches). Use `{{outputs.variable_name}}` to pass data across non-adjacent nodes. See [Graph orchestration guide](docs/en-US/workflow-graph.md).
|
|
||||||
- **Tool monitor** – Inspect running jobs, execution logs, and large-result attachments.
|
|
||||||
- **History & audit** – Every conversation and tool invocation is stored in SQLite with replay.
|
|
||||||
- **Conversation groups** – Organize conversations into groups, pin important groups, rename or delete groups via context menu.
|
|
||||||
- **Vulnerability management** – Create, update, and track vulnerabilities discovered during testing. Filter by severity (critical/high/medium/low/info), status (open/confirmed/fixed/false_positive), and conversation. View statistics and export findings.
|
|
||||||
- **Batch task management** – Create task queues with multiple tasks, add or edit tasks before execution, and run them sequentially. Each task executes as a separate conversation, with status tracking (pending/running/completed/failed/cancelled) and full execution history.
|
|
||||||
- **WebShell management** – Add and manage WebShell connections (PHP/ASP/ASPX/JSP or custom). Use the virtual terminal to run commands, the file manager to list, read, edit, upload, and delete files, and the AI assistant tab to drive scripted tests with per-connection conversation history. Connections are stored in SQLite; supports GET/POST and configurable command parameter (e.g. IceSword/AntSword style).
|
|
||||||
- **Built-in C2** – Create/start **listeners**, generate **payloads**, track **sessions**, enqueue **tasks**, and subscribe to **events** (SSE) from the Web UI or `/api/c2/*`. Agents and external clients use the C2 MCP tool family (including **`c2_task`**); when HITL is enabled, high-risk tasks can require human approval. Intended **only** for systems you are explicitly authorized to test.
|
|
||||||
- **Settings** – Tweak provider keys, MCP enablement, tool toggles, and agent iteration limits.
|
|
||||||
- **Human-in-the-loop (HITL)** – Sidebar sets mode and allowlisted tools (comma- or newline-separated); global list lives in `config.yaml` under `hitl.tool_whitelist`. The Audit Agent can use a separate low-cost model through `hitl.audit_model`, useful when human reviewers cannot keep up. **Apply** updates browser/server and can merge new tools into the file (**no restart**). **New chat** keeps sidebar choices; **HITL** nav shows pending approvals. Removing a tool in the sidebar does not remove it from the global list in `config.yaml`—edit the file if needed.
|
|
||||||
|
|
||||||
### Built-in Safeguards
|
|
||||||
- Required-field validation prevents accidental blank API credentials.
|
|
||||||
- Auto-generated strong passwords when `auth.password` is empty.
|
|
||||||
- Unified auth middleware for every web/API call (Bearer token flow).
|
|
||||||
- Timeout and sandbox guards per tool, plus structured logging for triage.
|
|
||||||
|
|
||||||
## Advanced Usage
|
|
||||||
|
|
||||||
### Role-Based Testing
|
|
||||||
- **Predefined roles** – System includes 12+ predefined security testing roles (Penetration Testing, CTF, Web App Scanning, API Security Testing, Binary Analysis, Cloud Security Audit, etc.) in the `roles/` directory.
|
|
||||||
- **Custom prompts** – Each role can define a `user_prompt` that prepends to user messages, guiding the AI to adopt specialized testing methodologies and focus areas.
|
|
||||||
- **Tool restrictions** – Roles can specify a `tools` list to limit available tools, ensuring focused testing workflows (e.g., CTF role restricts to CTF-specific utilities).
|
|
||||||
- **Skills** – Skill packs live under `skills_dir` and load via the Eino ADK **`skill`** tool (**progressive disclosure**) in both **single- and multi-agent** sessions when **`multi_agent.eino_skills`** is enabled. Optional host **read_file / glob / grep / write / edit / execute** and **`eino_middleware`** (tool_search, plantask, reduction, checkpoints, summarization transcripts, etc.) apply per mode—see docs.
|
|
||||||
- **Easy role creation** – Create custom roles by adding YAML files to the `roles/` directory. Each role defines `name`, `description`, `user_prompt`, `icon`, `tools`, and `enabled` fields.
|
|
||||||
- **Web UI integration** – Select roles from a dropdown in the chat interface. Role selection affects both AI behavior and available tool suggestions.
|
|
||||||
|
|
||||||
**Creating a custom role (example):**
|
|
||||||
1. Create a YAML file in `roles/` (e.g., `roles/custom-role.yaml`):
|
|
||||||
```yaml
|
|
||||||
name: Custom Role
|
|
||||||
description: Specialized testing scenario
|
|
||||||
user_prompt: You are a specialized security tester focusing on API security...
|
|
||||||
icon: "\U0001F4E1"
|
|
||||||
tools:
|
|
||||||
- api-fuzzer
|
|
||||||
- arjun
|
|
||||||
- graphql-scanner
|
|
||||||
enabled: true
|
|
||||||
```
|
|
||||||
2. Restart the server or reload configuration; the role appears in the role selector dropdown.
|
|
||||||
|
|
||||||
### Multi-Agent Mode (Eino: Deep, Plan-Execute, Supervisor)
|
|
||||||
- **What it is** – Multi-agent orchestration on CloudWeGo **Eino** `adk/prebuilt` (alongside **Eino single-agent** on `/api/eino-agent*`): **`deep`** — coordinator + **`task`** sub-agents for complex security testing and delegated synthesis; **`plan_execute`** — planner / executor / replanner for structured loops; **`supervisor`** — expert-routing mode with **`transfer`** / **`exit`** for multiple specialist sub-agents. Client sends **`orchestration`**: `deep` | `plan_execute` | `supervisor` (default `deep`).
|
|
||||||
- **Markdown agents** – Under `agents_dir` (default `agents/`):
|
|
||||||
- **Deep orchestrator**: `orchestrator.md` *or* one `.md` with `kind: orchestrator`. Body or `multi_agent.orchestrator_instruction`, then Eino defaults.
|
|
||||||
- **Plan-Execute orchestrator**: fixed name **`orchestrator-plan-execute.md`** (plus optional `orchestrator_instruction_plan_execute` in YAML).
|
|
||||||
- **Supervisor orchestrator**: fixed name **`orchestrator-supervisor.md`** (plus optional `orchestrator_instruction_supervisor`); requires at least one sub-agent, and one-sub-agent runs emit a hint that expert routing has limited value.
|
|
||||||
- **Sub-agents** (for **deep** / **supervisor**): other `*.md` files (YAML front matter + body). Not used as **`task`** targets if marked orchestrator-only.
|
|
||||||
- **Management** – Web UI: **Agents → Agent management**; API `/api/multi-agent/markdown-agents`.
|
|
||||||
- **Config** – `multi_agent` in `config.yaml`: `enabled`, `robot_default_agent_mode`, `batch_use_multi_agent`, `max_iteration`, `plan_execute_loop_max_iterations`, per-mode orchestrator instruction fields, optional YAML `sub_agents` merged with disk (`id` clash → Markdown wins), **`eino_skills`**, **`eino_middleware`** (optional ADK middleware and Deep/Supervisor tuning).
|
|
||||||
- **Resilience & long runs** – `checkpoint_dir` enables ADK **resume** after process crashes (distinct from trace-based “interrupt & continue”). `deep_model_retry_max_retries` retries transient LLM API failures within a single call. **Summarization** writes a filtered **transcript** when compression fires; the summary message includes the path so the model can `read_file` for scan output and other pre-compaction details.
|
|
||||||
- **Details** – **[docs/en-US/MULTI_AGENT_EINO.md](docs/en-US/MULTI_AGENT_EINO.md)** (streaming, robots, batch, middleware caveats).
|
|
||||||
|
|
||||||
### Skills System (Agent Skills + Eino)
|
|
||||||
- **Layout** – Each skill is a directory with **required** `SKILL.md` only ([Agent Skills](https://platform.claude.com/docs/en/agents-and-tools/agent-skills/overview)): YAML front matter **only** `name` and `description`, plus Markdown body. Optional sibling files (`FORMS.md`, `REFERENCE.md`, `scripts/*`, …). **No** `SKILL.yaml` (not part of Claude or Eino specs); sections/scripts/progressive behavior are **derived at runtime** from Markdown and the filesystem.
|
|
||||||
- **Runtime refactor** – **`skills_dir`** is the single root for packs. **Multi-agent** loads them through Eino’s official **`skill`** middleware (**progressive disclosure**: model calls `skill` with a pack **name** instead of receiving full SKILL text up front). Configure via **`multi_agent.eino_skills`**: `disable`, `filesystem_tools` (host read/glob/grep/write/edit/execute), `skill_tool_name`.
|
|
||||||
- **Eino / RAG** – Packages are also split into `schema.Document` chunks for `FilesystemSkillsRetriever` (`skills.AsEinoRetriever()`) in **compose** graphs (e.g. knowledge/indexing pipelines).
|
|
||||||
- **HTTP API** – `/api/skills` listing and `depth` (`summary` | `full`), `section`, and `resource_path` remain for the web UI and ops; **model-side** skill loading in multi-agent uses the **`skill`** tool, not MCP.
|
|
||||||
- **Optional `eino_middleware`** – e.g. `tool_search` (dynamic MCP tool list), `patch_tool_calls`, **`plantask`** (Eino `TaskCreate` / `TaskGet` / `TaskUpdate` / `TaskList`; JSON under `skills_dir/.eino/plantask/<conversation-id>/`; Eino clears task files when **all** tasks are marked completed), `reduction`, **`checkpoint_dir`** (`data/eino-checkpoints/`), **`deep_model_retry_max_retries`**, **`deep_output_key`**, task-tool description prefix—see `config.yaml` and `internal/config/config.go`.
|
|
||||||
- **Shipped demo** – `skills/cyberstrike-eino-demo/`; see `skills/README.md`.
|
|
||||||
|
|
||||||
**Creating a skill:**
|
|
||||||
1. `mkdir skills/<skill-id>` and add standard `SKILL.md` (+ any optional files), or drop in an open-source skill folder as-is.
|
|
||||||
2. Use **multi-agent** with **`multi_agent.eino_skills`** enabled so the model can call the **`skill`** tool with that pack **name**.
|
|
||||||
|
|
||||||
### Tool Orchestration & Extensions
|
|
||||||
- **YAML recipes** in `tools/*.yaml` describe commands, arguments, prompts, and metadata.
|
|
||||||
- **Directory hot-reload** – pointing `security.tools_dir` to a folder is usually enough; inline definitions in `config.yaml` remain supported for quick experiments.
|
|
||||||
- **Large tool outputs** – outputs beyond `reduction_max_length_for_trunc` are summarized via Eino reduction with full content persisted under `tmp/reduction/`; use `read_file` on the path in `<persisted-output>`.
|
|
||||||
- **Result compression** – multi-megabyte logs can be summarized or losslessly compressed before persisting to keep SQLite lean.
|
|
||||||
|
|
||||||
**Creating a custom tool (typical flow)**
|
|
||||||
1. Copy an existing YAML file from `tools/` (for example `tools/nmap.yaml` or `tools/ffuf.yaml`).
|
|
||||||
2. Update `name`, `command`, `args`, and `short_description`.
|
|
||||||
3. Describe positional or flag parameters in `parameters[]` so the agent knows how to build CLI arguments.
|
|
||||||
4. Provide a longer `description`/`notes` block if the agent needs extra context or post-processing tips.
|
|
||||||
5. Restart the server or reload configuration; the new tool becomes available immediately and can be enabled/disabled from the Settings panel.
|
|
||||||
|
|
||||||
### Attack-Chain Intelligence
|
|
||||||
- AI parses each conversation to assemble targets, tools, vulnerabilities, and relationships.
|
|
||||||
- The web UI renders the chain as an interactive graph with severity scoring and step replay.
|
|
||||||
- Export the chain or raw findings to external reporting pipelines.
|
|
||||||
|
|
||||||
### WebShell Management
|
|
||||||
- **Connections** – From the Web UI, go to **WebShell Management** to add, edit, or delete WebShell connections. Each connection stores: Shell URL, password/key, shell type (PHP, ASP, ASPX, JSP, Custom), request method (GET/POST), command parameter name (default `cmd`), and an optional remark; all records persist in SQLite and are compatible with common clients such as IceSword and AntSword.
|
|
||||||
- **Virtual terminal** – After selecting a connection, use the **Virtual terminal** tab to run arbitrary commands with history and quick commands (whoami/id/ls/pwd etc.). Output is streamed in the browser, and Ctrl+L clears the screen.
|
|
||||||
- **File manager** – Use the **File manager** tab to list directories, read or edit files, delete files, create folders/files, upload files (including chunked uploads for large files), rename paths, and download selected files. Path navigation supports breadcrumbs, parent directory jumps, and name filtering.
|
|
||||||
- **AI assistant** – Use the **AI assistant** tab to chat with an agent that understands the current WebShell connection, automatically runs tools and shell commands, and maintains per-connection conversation history with a sidebar of previous sessions.
|
|
||||||
- **Connectivity test** – Use **Test connectivity** to verify that the shell URL, password, and command parameter are correct before running commands (sends a lightweight `echo 1` check).
|
|
||||||
- **Persistence** – All WebShell connections and AI conversations are stored in SQLite (same database as conversations), so they persist across restarts.
|
|
||||||
|
|
||||||
### Built-in C2 (Command & Control)
|
|
||||||
- **What it is** – A first-party, **AI-native** C2 stack: listeners accept implants (beacons), the server stores **sessions** and **tasks** in SQLite, pushes updates over an **event bus** (including **SSE**), and exposes everything through authenticated **REST** plus MCP.
|
|
||||||
- **Listeners & transports** – `tcp_reverse`, `http_beacon`, `https_beacon`, and `websocket`; per-listener crypto keys; running listeners can be **restored after restart** when marked running in the database.
|
|
||||||
- **Agent integration** – MCP exposes a small **C2 tool family** (listeners, sessions, **`c2_task`**, task management, payloads, events, profiles, files) so the same agent loop can orchestrate C2 alongside other tools; dangerous task types can go through the existing **HITL** bridge when your session policy requires it.
|
|
||||||
- **Safety** – Use **only** in lab or **fully authorized** engagements; combine network isolation, strong auth, and HITL/allowlists as your policy demands.
|
|
||||||
|
|
||||||
### MCP Everywhere
|
|
||||||
- **Web mode** – ships with HTTP MCP server automatically consumed by the UI.
|
|
||||||
- **MCP stdio mode** – `go run cmd/mcp-stdio/main.go` exposes the agent to Cursor/CLI.
|
|
||||||
- **External MCP federation** – register third-party MCP servers (HTTP, stdio, or SSE) from the UI, toggle them per engagement, and monitor their health and call volume in real time.
|
|
||||||
- **Optional MCP servers** – the [`mcp-servers/`](mcp-servers/README.md) directory provides standalone MCPs (e.g. reverse shell). They speak standard MCP over stdio and work with CyberStrikeAI (Settings → External MCP), Cursor, VS Code, and other MCP clients.
|
|
||||||
|
|
||||||
#### MCP stdio quick start
|
|
||||||
1. **Build the binary** (run from the project root):
|
|
||||||
```bash
|
|
||||||
go build -o cyberstrike-ai-mcp cmd/mcp-stdio/main.go
|
|
||||||
```
|
|
||||||
2. **Wire it up in Cursor**
|
|
||||||
Open `Settings → Tools & MCP → Add Custom MCP`, pick **Command**, then point to the compiled binary and your config:
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"mcpServers": {
|
|
||||||
"cyberstrike-ai": {
|
|
||||||
"command": "/absolute/path/to/cyberstrike-ai-mcp",
|
|
||||||
"args": [
|
|
||||||
"--config",
|
|
||||||
"/absolute/path/to/config.yaml"
|
|
||||||
]
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
Replace the paths with your local locations; Cursor will launch the stdio server automatically.
|
|
||||||
|
|
||||||
#### MCP HTTP quick start (Cursor / Claude Code)
|
|
||||||
The HTTP MCP server runs on a separate port (default `8081`) and supports **header-based authentication** so only clients that send the correct header can call tools.
|
|
||||||
|
|
||||||
1. **Enable MCP in config** – In `config.yaml` set `mcp.enabled: true` and optionally `mcp.host` / `mcp.port`. For auth (recommended if the port is reachable from the network), set:
|
|
||||||
- `mcp.auth_header` – header name (e.g. `X-MCP-Token`);
|
|
||||||
- `mcp.auth_header_value` – secret value. **Leave it empty** if you want the server to **auto-generate** a random token on first start and write it back to the config.
|
|
||||||
2. **Start the service** – Run `./run.sh` or `go run cmd/server/main.go`. The MCP endpoint is `http://<host>:<port>/mcp` (e.g. `http://localhost:8081/mcp`).
|
|
||||||
3. **Copy the JSON from the terminal** – When MCP is enabled, the server prints a **ready-to-paste** JSON block. If `auth_header_value` was empty, it will have been generated and saved; the printed JSON includes the URL and headers.
|
|
||||||
4. **Use in Cursor or Claude Code**:
|
|
||||||
- **Cursor**: Paste the block into `~/.cursor/mcp.json` (or your project’s `.cursor/mcp.json`) under `mcpServers`, or merge it into your existing `mcpServers`.
|
|
||||||
- **Claude Code**: Paste into `.mcp.json` or `~/.claude.json` under `mcpServers`.
|
|
||||||
|
|
||||||
Example of what the terminal prints (with auth enabled):
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"mcpServers": {
|
|
||||||
"cyberstrike-ai": {
|
|
||||||
"url": "http://localhost:8081/mcp",
|
|
||||||
"headers": {
|
|
||||||
"X-MCP-Token": "<auto-generated-or-your-value>"
|
|
||||||
},
|
|
||||||
"type": "http"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
If you do not set `auth_header` / `auth_header_value`, the endpoint accepts requests without authentication (suitable only for localhost or trusted networks).
|
|
||||||
|
|
||||||
#### External MCP federation (HTTP/stdio/SSE)
|
|
||||||
CyberStrikeAI supports connecting to external MCP servers via three transport modes:
|
|
||||||
- **HTTP mode** – traditional request/response over HTTP POST
|
|
||||||
- **stdio mode** – process-based communication via standard input/output
|
|
||||||
- **SSE mode** – Server-Sent Events for real-time streaming communication
|
|
||||||
|
|
||||||
To add an external MCP server:
|
|
||||||
1. Open the Web UI and navigate to **Settings → External MCP**.
|
|
||||||
2. Click **Add External MCP** and provide the configuration in JSON format:
|
|
||||||
|
|
||||||
**HTTP mode example:**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"my-http-mcp": {
|
|
||||||
"transport": "http",
|
|
||||||
"url": "http://127.0.0.1:8081/mcp",
|
|
||||||
"description": "HTTP MCP server",
|
|
||||||
"timeout": 30
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**stdio mode example:**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"my-stdio-mcp": {
|
|
||||||
"command": "python3",
|
|
||||||
"args": ["/path/to/mcp-server.py"],
|
|
||||||
"description": "stdio MCP server",
|
|
||||||
"timeout": 30
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**SSE mode example:**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"my-sse-mcp": {
|
|
||||||
"transport": "sse",
|
|
||||||
"url": "http://127.0.0.1:8082/sse",
|
|
||||||
"description": "SSE MCP server",
|
|
||||||
"timeout": 30
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
3. Click **Save** and then **Start** to connect to the server.
|
|
||||||
4. Monitor the connection status, tool count, and health in real time.
|
|
||||||
|
|
||||||
**SSE mode benefits:**
|
|
||||||
- Real-time bidirectional communication via Server-Sent Events
|
|
||||||
- Suitable for scenarios requiring continuous data streaming
|
|
||||||
- Lower latency for push-based notifications
|
|
||||||
|
|
||||||
A test SSE MCP server is available at `cmd/test-sse-mcp-server/` for validation purposes.
|
|
||||||
|
|
||||||
### Knowledge Base
|
|
||||||
- **Vector search** – AI agent can automatically search the knowledge base for relevant security knowledge during conversations using the `search_knowledge_base` tool.
|
|
||||||
- **RAG pipeline (always on)** – **MultiQuery** (LLM query rewrite) → vector prefetch & fusion → **HTTP rerank** (DashScope `gte-rerank` or Cohere-compatible `/v1/rerank`) → post-processing (normalized dedupe, char/token budget, final top_k). Rerank failures degrade to fusion order without breaking search.
|
|
||||||
- **Vector retrieval** – cosine similarity over stored embeddings with configurable threshold, aligned with Eino `retriever.Retriever` usage.
|
|
||||||
- **Auto-indexing** – scans the `knowledge_base/` directory for Markdown files and automatically indexes them with embeddings (Markdown header split + recursive chunking via Eino).
|
|
||||||
- **Web management** – create, update, delete knowledge items through the web UI, with category-based organization; settings page exposes MultiQuery / rerank / prefetch options.
|
|
||||||
- **Retrieval logs** – tracks all knowledge retrieval operations for audit and debugging.
|
|
||||||
|
|
||||||
**Setting up the knowledge base:**
|
|
||||||
1. **Enable in config** – set `knowledge.enabled: true` in `config.yaml`:
|
|
||||||
```yaml
|
|
||||||
knowledge:
|
|
||||||
enabled: true
|
|
||||||
base_path: knowledge_base
|
|
||||||
embedding:
|
|
||||||
provider: openai
|
|
||||||
model: text-embedding-v4
|
|
||||||
base_url: "https://api.openai.com/v1" # or your embedding API
|
|
||||||
api_key: "sk-xxx"
|
|
||||||
retrieval:
|
|
||||||
top_k: 5
|
|
||||||
similarity_threshold: 0.7
|
|
||||||
multi_query:
|
|
||||||
max_queries: 4 # LLM rewrite variants (always on)
|
|
||||||
rerank: # always on; empty fields inherit openai/embedding credentials
|
|
||||||
provider: "" # auto: dashscope | cohere from base_url
|
|
||||||
model: "" # empty: gte-rerank (DashScope) or rerank-multilingual-v3.0 (Cohere)
|
|
||||||
base_url: ""
|
|
||||||
api_key: ""
|
|
||||||
post_retrieve:
|
|
||||||
prefetch_top_k: 20 # vector candidates per MultiQuery variant; 0 = max(top_k×4, 20)
|
|
||||||
max_context_chars: 0
|
|
||||||
max_context_tokens: 0
|
|
||||||
```
|
|
||||||
2. **Add knowledge files** – place Markdown files in `knowledge_base/` directory, organized by category (e.g., `knowledge_base/SQL Injection/README.md`).
|
|
||||||
3. **Scan and index** – use the web UI to scan the knowledge base directory, which will automatically import files and build vector embeddings.
|
|
||||||
4. **Use in conversations** – the AI agent will automatically use `search_knowledge_base` when it needs security knowledge. You can also explicitly ask: "Search the knowledge base for SQL injection techniques".
|
|
||||||
|
|
||||||
**Knowledge base structure:**
|
|
||||||
- Files are organized by category (directory name becomes the category).
|
|
||||||
- Each Markdown file becomes a knowledge item with automatic chunking for vector search.
|
|
||||||
- The system supports incremental updates – modified files are re-indexed automatically.
|
|
||||||
|
|
||||||
|
|
||||||
### Automation Hooks
|
## Configuration
|
||||||
- **REST APIs** – everything the UI uses (auth, conversations, tool runs, monitor, vulnerabilities, roles) is available over JSON.
|
|
||||||
- **Multi-agent APIs** – `POST /api/multi-agent/stream` (SSE, when enabled), `POST /api/multi-agent` (non-streaming), Markdown agents under `/api/multi-agent/markdown-agents` (list/get/create/update/delete).
|
|
||||||
- **Role APIs** – manage security testing roles via `/api/roles` endpoints: `GET /api/roles` (list all roles), `GET /api/roles/:name` (get role), `POST /api/roles` (create role), `PUT /api/roles/:name` (update role), `DELETE /api/roles/:name` (delete role). Roles are stored as YAML files in the `roles/` directory and support hot-reload.
|
|
||||||
- **Vulnerability APIs** – manage vulnerabilities via `/api/vulnerabilities` endpoints: `GET /api/vulnerabilities` (list with filters), `POST /api/vulnerabilities` (create), `GET /api/vulnerabilities/:id` (get), `PUT /api/vulnerabilities/:id` (update), `DELETE /api/vulnerabilities/:id` (delete), `GET /api/vulnerabilities/stats` (statistics).
|
|
||||||
- **Batch Task APIs** – manage batch task queues via `/api/batch-tasks` endpoints: `POST /api/batch-tasks` (create queue), `GET /api/batch-tasks` (list queues), `GET /api/batch-tasks/:queueId` (get queue), `POST /api/batch-tasks/:queueId/start` (start execution), `POST /api/batch-tasks/:queueId/cancel` (cancel), `DELETE /api/batch-tasks/:queueId` (delete), `POST /api/batch-tasks/:queueId/tasks` (add task), `PUT /api/batch-tasks/:queueId/tasks/:taskId` (update task), `DELETE /api/batch-tasks/:queueId/tasks/:taskId` (delete task). Tasks execute sequentially, each creating a separate conversation with full status tracking.
|
|
||||||
- **WebShell APIs** – manage WebShell connections and execute commands via `/api/webshell/connections` (GET list, POST create, PUT update, DELETE delete) and `/api/webshell/exec` (command execution), `/api/webshell/fileop` (list/read/write/delete files).
|
|
||||||
- **C2 APIs** – manage listeners, sessions, tasks, payloads, files, and events under `/api/c2/*` (e.g. listeners CRUD/start/stop, session sleep, task create/cancel/wait, payload build/download, event stream).
|
|
||||||
- **Task control** – pause/resume/stop long scans, re-run steps with new params, or stream transcripts.
|
|
||||||
- **Audit & security** – rotate passwords via `/api/auth/change-password`, enforce short-lived sessions, and restrict MCP ports at the network layer when exposing the service.
|
|
||||||
|
|
||||||
## Configuration Reference
|
Use [`config.example.yaml`](config.example.yaml) as the authoritative configuration template and copy only the values required for your environment. At minimum, configure the server and an OpenAI-compatible model provider:
|
||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
auth:
|
|
||||||
password: "change-me"
|
|
||||||
session_duration_hours: 12
|
|
||||||
server:
|
server:
|
||||||
host: "0.0.0.0"
|
host: "127.0.0.1"
|
||||||
port: 8080
|
port: 8080
|
||||||
log:
|
|
||||||
level: "info"
|
|
||||||
output: "stdout"
|
|
||||||
mcp:
|
|
||||||
enabled: true
|
|
||||||
host: "0.0.0.0"
|
|
||||||
port: 8081
|
|
||||||
auth_header: "X-MCP-Token" # optional; leave empty for no auth
|
|
||||||
auth_header_value: "" # optional; leave empty to auto-generate on first start
|
|
||||||
openai:
|
openai:
|
||||||
api_key: "sk-xxx"
|
api_key: "${OPENAI_API_KEY}"
|
||||||
base_url: "https://api.deepseek.com/v1"
|
base_url: "https://api.openai.com/v1"
|
||||||
model: "deepseek-chat"
|
model: "your-model"
|
||||||
database:
|
|
||||||
path: "data/conversations.db"
|
|
||||||
knowledge_db_path: "data/knowledge.db" # Optional: separate DB for knowledge base
|
|
||||||
security:
|
|
||||||
tools_dir: "tools"
|
|
||||||
knowledge:
|
|
||||||
enabled: false # Enable knowledge base feature
|
|
||||||
base_path: "knowledge_base" # Path to knowledge base directory
|
|
||||||
embedding:
|
|
||||||
provider: "openai" # Embedding provider (currently only "openai")
|
|
||||||
model: "text-embedding-v4" # Embedding model name
|
|
||||||
base_url: "" # Leave empty to use OpenAI base_url
|
|
||||||
api_key: "" # Leave empty to use OpenAI api_key
|
|
||||||
retrieval:
|
|
||||||
top_k: 5 # Number of top results to return
|
|
||||||
similarity_threshold: 0.7 # Minimum cosine similarity (0-1)
|
|
||||||
multi_query:
|
|
||||||
max_queries: 4 # MultiQuery rewrite variants (always on)
|
|
||||||
rerank: # HTTP rerank (always on); empty fields inherit openai/embedding credentials
|
|
||||||
provider: ""
|
|
||||||
model: ""
|
|
||||||
base_url: ""
|
|
||||||
api_key: ""
|
|
||||||
post_retrieve:
|
|
||||||
prefetch_top_k: 20 # per MultiQuery variant; 0 = max(top_k×4, 20)
|
|
||||||
max_context_chars: 0
|
|
||||||
max_context_tokens: 0
|
|
||||||
roles_dir: "roles" # Role configuration directory (relative to config file)
|
|
||||||
skills_dir: "skills" # Skills directory (relative to config file)
|
|
||||||
agents_dir: "agents" # Multi-agent Markdown definitions (orchestrator + sub-agents)
|
|
||||||
multi_agent:
|
|
||||||
enabled: false
|
|
||||||
default_mode: "eino_single" # eino_single | multi (UI default when multi-agent is enabled)
|
|
||||||
robot_default_agent_mode: eino_single
|
|
||||||
batch_use_multi_agent: false
|
|
||||||
orchestrator_instruction: "" # Deep; used when orchestrator.md body is empty
|
|
||||||
# orchestrator_instruction_plan_execute / orchestrator_instruction_supervisor optional
|
|
||||||
# eino_skills: { disable: false, filesystem_tools: true, skill_tool_name: skill }
|
|
||||||
# eino_middleware: plantask_enable, checkpoint_dir, deep_model_retry_max_retries, deep_output_key, ...
|
|
||||||
project:
|
|
||||||
enabled: true # Enable project blackboard & fact MCP tools
|
|
||||||
fact_index_max_runes: 65000
|
|
||||||
fact_summary_max_runes: 24000
|
|
||||||
default_inject_deprecated: false
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### Tool Definition Example (`tools/nmap.yaml`)
|
Do not commit real credentials. Review the [configuration reference](docs/en-US/configuration.md), [recommended profiles](docs/en-US/configuration-profiles.md), and [security hardening guide](docs/en-US/security-hardening.md) before exposing the service beyond localhost.
|
||||||
|
|
||||||
```yaml
|
|
||||||
name: "nmap"
|
|
||||||
command: "nmap"
|
|
||||||
args: ["-sT", "-sV", "-sC"]
|
|
||||||
enabled: true
|
|
||||||
short_description: "Network mapping & service fingerprinting"
|
|
||||||
parameters:
|
|
||||||
- name: "target"
|
|
||||||
type: "string"
|
|
||||||
description: "IP or domain"
|
|
||||||
required: true
|
|
||||||
position: 0
|
|
||||||
- name: "ports"
|
|
||||||
type: "string"
|
|
||||||
flag: "-p"
|
|
||||||
description: "Range, e.g. 1-1000"
|
|
||||||
```
|
|
||||||
|
|
||||||
### Role Definition Example (`roles/penetration-testing.yaml`)
|
|
||||||
|
|
||||||
```yaml
|
|
||||||
name: Penetration Testing
|
|
||||||
description: Professional penetration testing expert for comprehensive security testing
|
|
||||||
user_prompt: You are a professional cybersecurity penetration testing expert. Please use professional penetration testing methods and tools to conduct comprehensive security testing on targets, including but not limited to SQL injection, XSS, CSRF, file inclusion, command execution and other common vulnerabilities.
|
|
||||||
icon: "\U0001F3AF"
|
|
||||||
tools:
|
|
||||||
- nmap
|
|
||||||
- sqlmap
|
|
||||||
- nuclei
|
|
||||||
- burpsuite
|
|
||||||
- metasploit
|
|
||||||
- httpx
|
|
||||||
- record_vulnerability
|
|
||||||
- list_knowledge_risk_types
|
|
||||||
- search_knowledge_base
|
|
||||||
enabled: true
|
|
||||||
```
|
|
||||||
|
|
||||||
## Related documentation
|
## Related documentation
|
||||||
|
|
||||||
- [Documentation index](docs/README.md): deployment, configuration, security model, API, knowledge base, C2, WebShell, MCP, development, testing, and troubleshooting.
|
- **New users:** [Deployment](docs/en-US/deployment.md) → [Configuration](docs/en-US/configuration.md) → [Troubleshooting](docs/en-US/troubleshooting.md)
|
||||||
- [Deployment guide](docs/en-US/deployment.md): source/binary startup, HTTPS, reverse proxy, systemd, backup, upgrade, and rollback.
|
- **Operators:** [Configuration profiles](docs/en-US/configuration-profiles.md) → [Security hardening](docs/en-US/security-hardening.md) → [Runbooks](docs/en-US/runbooks.md)
|
||||||
- [Runbooks](docs/en-US/runbooks.md): production setup, external MCP, knowledge base, authorized Web testing, and C2 cleanup workflows.
|
- **Integrators:** [API reference](docs/en-US/api-reference.md) → [API recipes](docs/en-US/api-recipes.md) → [MCP federation](docs/en-US/mcp-federation.md)
|
||||||
- [Security hardening](docs/en-US/security-hardening.md): launch baseline, HITL allowlist, reverse proxy, file permissions, and periodic review.
|
- **Contributors:** [Developer guide](docs/en-US/developer-guide.md) → [Testing](docs/en-US/testing.md) → [Contributing](docs/en-US/contributing-guide.md)
|
||||||
- [API recipes](docs/en-US/api-recipes.md): examples for login, Agent, streaming, multi-agent, uploads, vulnerabilities, KB, and audit export.
|
- **All topics:** [English documentation](docs/en-US/README.md) · [Bilingual documentation index](docs/README.md)
|
||||||
- [Configuration reference](docs/en-US/configuration.md): main `config.yaml` sections, recommended values, and update guidance.
|
|
||||||
- [Security model](docs/en-US/security-model.md): authentication, tool execution, HITL, audit, C2/WebShell, and data safety boundaries.
|
|
||||||
- [RBAC administration](docs/en-US/rbac.md): platform users, system/custom roles, permission catalog, per-permission scopes, resource assignments, Agent/MCP/robot boundaries, and API examples.
|
|
||||||
- [API reference](docs/en-US/api-reference.md): OpenAPI, authentication, Agent, projects, knowledge base, C2, WebShell, and other API entry points.
|
|
||||||
- [Multi-agent mode (Eino)](docs/en-US/MULTI_AGENT_EINO.md): **Deep**, **Plan-Execute**, **Supervisor**, `agents/*.md`, `eino_skills` / `eino_middleware`, APIs, and chat/stream behavior.
|
|
||||||
- [Graph orchestration guide](docs/en-US/workflow-graph.md): visual workflow design, node configuration, `previous` / `outputs` variable passing, and role binding.
|
|
||||||
- [Robot / Chatbot guide](docs/en-US/robot.md): Platform setup, RBAC user-binding/service-account modes, sender allowlists, commands, verification, and troubleshooting.
|
|
||||||
- [HITL best practices](docs/en-US/hitl-best-practices.md): reviewer modes, allowlists, Audit Agent prompts, and separate small-model configuration.
|
|
||||||
|
|
||||||
## Project Layout
|
## Project Layout
|
||||||
|
|
||||||
@@ -670,6 +314,7 @@ CyberStrikeAI/
|
|||||||
├── agents/ # Multi-agent Markdown (orchestrator.md + sub-agent *.md)
|
├── agents/ # Multi-agent Markdown (orchestrator.md + sub-agent *.md)
|
||||||
├── docs/ # Topic docs (deployment, config, security, API, knowledge base, C2, WebShell, etc.)
|
├── docs/ # Topic docs (deployment, config, security, API, knowledge base, C2, WebShell, etc.)
|
||||||
├── images/ # Docs screenshots & diagrams
|
├── images/ # Docs screenshots & diagrams
|
||||||
|
├── scripts/ # Repository maintenance checks, including documentation validation
|
||||||
├── config.yaml # Runtime configuration
|
├── config.yaml # Runtime configuration
|
||||||
├── run.sh # Convenience launcher
|
├── run.sh # Convenience launcher
|
||||||
└── README*.md
|
└── README*.md
|
||||||
@@ -711,6 +356,26 @@ CyberStrikeAI has joined [404Starlink](https://github.com/knownsec/404StarLink)
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
## Community and Support
|
||||||
|
|
||||||
|
- Join the community on [Discord](https://discord.gg/8PjVCMu8Zw).
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><strong>WeChat group</strong></summary>
|
||||||
|
|
||||||
|
<img src="./images/wechat-group-cyberstrikeai-qr.jpg" alt="CyberStrikeAI WeChat group QR code" width="280">
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><strong>Sponsorship via WeChat Pay or Alipay</strong></summary>
|
||||||
|
|
||||||
|
<div align="center">
|
||||||
|
<img src="./images/sponsor-wechat-alipay-qr.jpg" alt="WeChat Pay and Alipay sponsorship QR codes" width="480">
|
||||||
|
</div>
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
## License
|
## License
|
||||||
|
|
||||||
CyberStrikeAI is licensed under the Apache License 2.0.
|
CyberStrikeAI is licensed under the Apache License 2.0.
|
||||||
|
|||||||
+97
-432
@@ -6,31 +6,14 @@
|
|||||||
|
|
||||||
[中文](README_CN.md) | [English](README.md)
|
[中文](README_CN.md) | [English](README.md)
|
||||||
|
|
||||||
**社区**:[加入 Discord](https://discord.gg/8PjVCMu8Zw)
|
**CyberStrikeAI 是 AI 原生网络安全的智能执行中枢——让意图转化为受治理的行动,让证据沉淀为运营记忆,并让每次行动优化下一次行动。**
|
||||||
|
|
||||||
**CyberStrikeAI 正在构建现代网络安全的智能体执行层。**
|
CyberStrikeAI 将规划、执行、人工监督、证据与复盘连接在同一个可审计工作空间中。项目基于 Go 构建,融合 Eino 智能体、MCP 原生工具、RAG 知识、图工作流以及攻击链建模与分析能力,面向已获得明确授权的安全任务。
|
||||||
|
|
||||||
它将 AI 智能体、安全工具、MCP 原生集成、知识系统、人工监督与攻击链智能汇聚到一个面向授权安全任务的统一工作空间中。CyberStrikeAI 不再把工具、提示词、证据、审批和报告视为割裂环节,而是将安全意图转化为可规划、可执行、可审查、可复盘、可持续沉淀上下文的多智能体工作流。
|
**从这里开始:** [快速上手](#快速上手一条命令部署) · [中文文档](docs/zh-CN/README.md) · [安全加固](docs/zh-CN/security-hardening.md)
|
||||||
|
|
||||||
CyberStrikeAI 基于 Go 构建,为 AI 原生安全运营提供完整底座:100+ 精选工具配方、角色化测试、Agent Skills、基于 Eino 的单智能体与多智能体编排、RAG 知识检索、图工作流、漏洞与任务生命周期管理、WebShell 运营、机器人接入,以及面向授权实验室和安全任务场景的内置轻量 C2 框架。
|
> [!IMPORTANT]
|
||||||
|
> 仅可对自有系统或已获得明确授权的目标使用 CyberStrikeAI。在共享或生产环境启用高风险工具、WebShell 或 C2 前,请先阅读[安全模型](docs/zh-CN/security-model.md)和[安全加固指南](docs/zh-CN/security-hardening.md)。
|
||||||
<details>
|
|
||||||
<summary><strong>微信群</strong>(点击展开二维码)</summary>
|
|
||||||
|
|
||||||
<img src="./images/wechat-group-cyberstrikeai-qr.jpg" alt="CyberStrikeAI 微信群二维码" width="280">
|
|
||||||
|
|
||||||
</details>
|
|
||||||
|
|
||||||
<details>
|
|
||||||
<summary><strong>赞助</strong>(点击展开)</summary>
|
|
||||||
|
|
||||||
若 CyberStrikeAI 对您有帮助,可通过 **微信支付** 或 **支付宝** 赞助项目:
|
|
||||||
|
|
||||||
<div align="center">
|
|
||||||
<img src="./images/sponsor-wechat-alipay-qr.jpg" alt="微信与支付宝赞助二维码" width="480">
|
|
||||||
</div>
|
|
||||||
|
|
||||||
</details>
|
|
||||||
|
|
||||||
## 界面与集成预览
|
## 界面与集成预览
|
||||||
|
|
||||||
@@ -53,6 +36,9 @@ CyberStrikeAI 基于 Go 构建,为 AI 原生安全运营提供完整底座:1
|
|||||||
|
|
||||||
*仪表盘提供系统运行状态、安全漏洞、工具使用情况和知识库的全面概览,帮助用户快速了解平台核心功能和当前状态。*
|
*仪表盘提供系统运行状态、安全漏洞、工具使用情况和知识库的全面概览,帮助用户快速了解平台核心功能和当前状态。*
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><strong>查看更多界面截图</strong></summary>
|
||||||
|
|
||||||
### 核心功能概览
|
### 核心功能概览
|
||||||
|
|
||||||
<table>
|
<table>
|
||||||
@@ -114,33 +100,48 @@ CyberStrikeAI 基于 Go 构建,为 AI 原生安全运营提供完整底座:1
|
|||||||
</tr>
|
</tr>
|
||||||
</table>
|
</table>
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
</div>
|
</div>
|
||||||
|
|
||||||
## 特性速览
|
## 特性速览
|
||||||
|
|
||||||
- 🤖 面向智能体时代的执行层,将自然语言意图转化为精准、受控、可审计的安全行动
|
### 智能体与编排
|
||||||
- 🧩 基于 Eino 的单智能体与多智能体编排,支持 Deep、Plan-Execute、Supervisor 等模式
|
|
||||||
- 🔌 MCP 原生工具执行,支持 HTTP / stdio / SSE 传输、外部 MCP 联邦与动态工具发现
|
- 🤖 **智能体执行层**:将自然语言意图转化为受控、可审计的安全行动。
|
||||||
- 🧰 100+ 精选安全工具配方、YAML 扩展机制与按角色收敛的工具控制
|
- 🧩 **Eino 编排**:支持单智能体及 Deep、Plan-Execute、Supervisor 多智能体模式。
|
||||||
- 📄 大结果分页、压缩与全文检索
|
- 🔀 **图工作流**:通过 Agent、工具、条件、审批和输出节点构建可复用流程。
|
||||||
- 🔗 攻击链智能分析,支持图谱视图、风险打分、项目事实沉淀与步骤回放
|
- 🎭 **角色化测试**:为常见安全场景提供聚焦的提示词和工具策略。
|
||||||
- 🧑⚖️ 人机协同治理,支持审批模式、免审批白名单、审计 Agent 复核与可追溯决策
|
|
||||||
- 🔒 Web 登录保护、审计日志、SQLite 持久化与行动证据留存
|
### 工具与知识扩展
|
||||||
- 🔐 **平台 RBAC**:支持多用户、系统/自定义角色、逐权限 Scope(`all` / `assigned` / `own`)、资源归属与显式授权,并统一约束 API、Agent、MCP、后台任务和机器人;详见 [RBAC 权限管理](docs/zh-CN/rbac.md)
|
|
||||||
- 📚 知识库(RAG):**Eino MultiQuery** 查询改写 + 多路向量检索 + **HTTP 精排**(DashScope `gte-rerank` / Cohere 兼容)+ 后处理(去重、预算);索引侧为 **Eino Compose** 流水线
|
- 🧰 **安全工具**:提供 100+ 精选 YAML 工具配方,支持自定义扩展和按角色控制。
|
||||||
- 📁 对话分组管理:支持分组创建、置顶、重命名、删除等操作
|
- 🔌 **MCP 集成**:支持 HTTP、stdio、SSE、外部 MCP 联邦和动态工具发现。
|
||||||
- 📂 **项目管理**:共享事实(黑板)跨会话沉淀认知,`upsert_project_fact` + `links` 串联攻击路径;聊天攻击链与项目事实图可视化
|
- 🎯 **Agent Skills**:遵循标准 Skill 目录结构,支持渐进式按需加载。
|
||||||
- 🛡️ 漏洞管理功能:完整的漏洞 CRUD 操作,支持严重程度分级、状态流转、按对话/严重程度/状态过滤,以及统计看板
|
- 📚 **知识库**:组合查询改写、向量检索、精排和结果后处理能力。
|
||||||
- 📋 批量任务管理:创建任务队列,批量添加任务,依次顺序执行,支持任务编辑与状态跟踪
|
- 🖼️ **视觉分析**:使用独立视觉模型分析截图、验证码和 UI,对话中仅保留文字摘要。
|
||||||
- 🎭 角色化测试:预设安全测试角色(渗透测试、CTF、Web 应用扫描等),支持自定义提示词和工具限制
|
|
||||||
- 🔀 **图编排**:可视化流程编排(开始 / Agent / 工具 / 条件 / 审批 / 输出),节点间用 `{{previous.output}}` 或 `{{outputs.变量名}}` 传参;绑定角色后对话自动按图执行。详见 [图编排使用说明](docs/zh-CN/workflow-graph.md)
|
### 安全治理与审计
|
||||||
- 🧩 **Agent 编排(CloudWeGo Eino)**:**单代理** `POST /api/eino-agent/stream`(Eino ADK);**多代理** `POST /api/multi-agent/stream`,`orchestration` 选 **`deep`** / **`plan_execute`** / **`supervisor`**。ADK **Summarization** 在上下文过长时压缩历史;压缩前将可恢复 **转录** 写入 `data/conversation_artifacts/<会话ID>/summarization/transcript.txt`(保留完整 user/assistant/tool 轮次,省略静态 system)。`agents/` 下主代理与子代理 Markdown 见 [多代理说明](docs/zh-CN/MULTI_AGENT_EINO.md)
|
|
||||||
- 🖼️ **视觉分析(`analyze_image`)**:独立 Vision 模型(如 `qwen-vl-max`),MCP 工具分析本地截图/验证码/UI;图片仅在单次 VL 调用中出现,对话上下文只保留文字摘要。配置见 `config.yaml` → `vision` 与 [视觉分析说明](docs/zh-CN/VISION.md)
|
- 🧑⚖️ **人机协同**:支持审批模式、工具白名单、审计 Agent 复核和决策追踪。
|
||||||
- 🎯 **Skills(面向 Eino 重构)**:技能包放在 **`skills_dir`**,遵循 **Agent Skills** 目录规范(`SKILL.md` + 可选文件);**多代理** 下通过 Eino 官方 **`skill`** 工具 **渐进式披露**(按 name 加载)。**`multi_agent.eino_skills`** 控制是否启用、本机文件/Shell 工具、工具名覆盖;**`eino_middleware`** 可选 patch、tool_search、**plantask**(`TaskCreate` / `TaskList` 任务板,落在 `skills_dir/.eino/plantask/`)、reduction、文件型 **checkpoint**(`checkpoint_dir`)、ChatModel **重试**、会话 **输出键** 及 Deep 调参。20+ 领域示例仍可绑定角色
|
- 🔐 **平台 RBAC**:支持多用户、系统及自定义角色、权限 Scope、资源归属和显式授权。
|
||||||
- 📱 **机器人**:个人微信、企业微信、钉钉、飞书、Telegram、Slack、Discord、QQ 机器人,在手机或 IM 中与 CyberStrikeAI 对话(详见 [机器人使用说明](docs/zh-CN/robot.md))
|
- 🔒 **安全与审计**:提供登录保护、审计日志、SQLite 持久化和行动证据留存。
|
||||||
- 🧑⚖️ **人机协同(HITL)**:对话页侧栏配置协同模式与免审批工具白名单;全局列表在 `config.yaml` 的 `hitl.tool_whitelist`;审计 Agent 可通过 `hitl.audit_model` 使用独立小模型;点「应用」可将新增工具合并写入配置文件且**无需重启**即可生效;导航 **人机协同** 页处理待审批工具调用。详见 [人机协同最佳实践](docs/zh-CN/hitl-best-practices.md)
|
- 📄 **结果治理**:支持大结果分页、压缩、归档和检索。
|
||||||
- 🐚 **WebShell 管理**:添加与管理 WebShell 连接(兼容冰蝎/蚁剑等),通过虚拟终端执行命令、内置文件管理进行文件操作,并提供按连接维度保存历史的 AI 助手标签页;支持 PHP/ASP/ASPX/JSP 及自定义类型,可配置请求方法与命令参数。
|
|
||||||
- 📡 **内置 C2**:面向 AI 协同的轻量 **C2**——**多种监听器**(TCP 反向、HTTP/HTTPS Beacon、WebSocket)、**加密** Beacon 信道、**会话与任务**队列及持久化、**Payload** 辅助(一键命令 / 构建 / 下载)、**SSE** 实时事件、REST(`/api/c2/*`)及智能体侧 **一组 C2 MCP 工具**(如 `c2_listener`、`c2_session`、**`c2_task`**、`c2_task_manage`、`c2_payload`、`c2_event`、`c2_profile`、`c2_file`);敏感操作可对接 **人机协同(HITL)**,并支持 OPSEC 类规则(如命令拒绝正则)。**仅限授权测试。**
|
### 安全运营管理
|
||||||
|
|
||||||
|
- 📁 **对话管理**:支持分组、置顶、重命名和批量管理。
|
||||||
|
- 📂 **项目与攻击链**:关联跨会话事实、风险评分、图谱视图和步骤回放。
|
||||||
|
- 🛡️ **漏洞管理**:支持严重程度分级、状态流转、过滤和统计看板。
|
||||||
|
- 📋 **批量任务**:支持任务队列、编辑、状态跟踪和结果留存。
|
||||||
|
- 📱 **机器人接入**:支持个人微信、企业微信、钉钉、飞书、Telegram、Slack、Discord 和 QQ。
|
||||||
|
|
||||||
|
### 授权安全操作
|
||||||
|
|
||||||
|
- 🐚 **WebShell 管理**:提供连接管理、虚拟终端、文件操作和 AI 辅助工作流。
|
||||||
|
- 📡 **内置 C2**:提供监听器、加密 Beacon、会话、任务队列、Payload 辅助和实时事件。
|
||||||
|
|
||||||
|
> WebShell、C2 及其他高风险能力仅限自有系统或已获得明确授权的测试环境。使用前请阅读[安全模型](docs/zh-CN/security-model.md)和[安全加固指南](docs/zh-CN/security-hardening.md)。
|
||||||
|
|
||||||
## 插件(Plugins)
|
## 插件(Plugins)
|
||||||
|
|
||||||
@@ -159,6 +160,9 @@ CyberStrikeAI 基于 Go 构建,为 AI 原生安全运营提供完整底座:1
|
|||||||
|
|
||||||
系统预置 100+ 渗透/攻防工具,覆盖完整攻击链:
|
系统预置 100+ 渗透/攻防工具,覆盖完整攻击链:
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><strong>查看完整工具分类</strong></summary>
|
||||||
|
|
||||||
- **网络扫描**:nmap、masscan、rustscan、arp-scan、nbtscan
|
- **网络扫描**:nmap、masscan、rustscan、arp-scan、nbtscan
|
||||||
- **Web 应用扫描**:sqlmap、nikto、dirb、gobuster、feroxbuster、ffuf、httpx
|
- **Web 应用扫描**:sqlmap、nikto、dirb、gobuster、feroxbuster、ffuf、httpx
|
||||||
- **漏洞扫描**:nuclei、wpscan、wafw00f、dalfox、xsser
|
- **漏洞扫描**:nuclei、wpscan、wafw00f、dalfox、xsser
|
||||||
@@ -175,12 +179,16 @@ CyberStrikeAI 基于 Go 构建,为 AI 原生安全运营提供完整底座:1
|
|||||||
- **CTF 实用工具**:stegsolve、zsteg、hash-identifier、fcrackzip、pdfcrack、cyberchef
|
- **CTF 实用工具**:stegsolve、zsteg、hash-identifier、fcrackzip、pdfcrack、cyberchef
|
||||||
- **系统辅助**:exec、create-file、delete-file、list-files、modify-file
|
- **系统辅助**:exec、create-file、delete-file、list-files、modify-file
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
|
工具定义、自定义方式与使用说明见 [tools/README.md](tools/README.md)。
|
||||||
|
|
||||||
## 基础使用
|
## 基础使用
|
||||||
|
|
||||||
### 快速上手(一条命令部署)
|
### 快速上手(一条命令部署)
|
||||||
|
|
||||||
**环境要求:**
|
**环境要求:**
|
||||||
- Go 1.21+ ([下载安装](https://go.dev/dl/))
|
- Go 1.25+([下载安装](https://go.dev/dl/),以 `go.mod` 为准)
|
||||||
- Python 3.10+ ([下载安装](https://www.python.org/downloads/))
|
- Python 3.10+ ([下载安装](https://www.python.org/downloads/))
|
||||||
|
|
||||||
**一条命令部署:**
|
**一条命令部署:**
|
||||||
@@ -198,6 +206,12 @@ chmod +x run.sh && ./run.sh
|
|||||||
- ✅ 编译构建项目
|
- ✅ 编译构建项目
|
||||||
- ✅ 启动服务器
|
- ✅ 启动服务器
|
||||||
|
|
||||||
|
**验证是否启动成功:**
|
||||||
|
|
||||||
|
1. 确认终端显示 `● ONLINE`,并在其后给出实际 Web UI 地址。
|
||||||
|
2. 打开该地址;默认 HTTPS 使用本地自签证书,首次访问需接受一次浏览器证书提示。
|
||||||
|
3. 全新安装时,妥善保存 `ADMIN SETUP REQUIRED` 下仅展示一次的 `admin` 密码,登录后立即修改。
|
||||||
|
|
||||||
**网络默认:** `run.sh` 会以 **`--https`** 并传入项目根 **`config.yaml`** 启动(本机自签证书,多路流式场景更稳)。只要明文 HTTP 用 **`./run.sh --http`**。生产环境在 **`config.yaml`** 的 **`server.tls_cert_path` / `server.tls_key_path`** 配正式证书(见文件内注释)。手动启动可加 **`--https`** 或环境变量 **`CYBERSTRIKE_HTTPS=1`**;`-config` 写错时程序会在终端提示正确写法。
|
**网络默认:** `run.sh` 会以 **`--https`** 并传入项目根 **`config.yaml`** 启动(本机自签证书,多路流式场景更稳)。只要明文 HTTP 用 **`./run.sh --http`**。生产环境在 **`config.yaml`** 的 **`server.tls_cert_path` / `server.tls_key_path`** 配正式证书(见文件内注释)。手动启动可加 **`--https`** 或环境变量 **`CYBERSTRIKE_HTTPS=1`**;`-config` 写错时程序会在终端提示正确写法。
|
||||||
|
|
||||||
**首次配置:**
|
**首次配置:**
|
||||||
@@ -206,12 +220,12 @@ chmod +x run.sh && ./run.sh
|
|||||||
- 进入 `设置` → 填写 API 配置信息:
|
- 进入 `设置` → 填写 API 配置信息:
|
||||||
```yaml
|
```yaml
|
||||||
openai:
|
openai:
|
||||||
api_key: "sk-your-key"
|
api_key: "${OPENAI_API_KEY}"
|
||||||
base_url: "https://api.openai.com/v1" # 或 https://api.deepseek.com/v1
|
base_url: "https://api.openai.com/v1" # 或 https://api.deepseek.com/v1
|
||||||
model: "gpt-4o" # 或 deepseek-chat, claude-3-opus 等
|
model: "gpt-4o" # 或 deepseek-chat, claude-3-opus 等
|
||||||
```
|
```
|
||||||
- 或启动前直接编辑 `config.yaml` 文件
|
- 或启动前直接编辑 `config.yaml` 文件
|
||||||
2. **登录系统** - 使用控制台显示的自动生成密码(或在 `config.yaml` 中设置 `auth.password`)
|
2. **登录系统** - 首次启动时控制台会显示自动生成的 `admin` 初始密码;也可在「平台权限 → 用户管理」中创建账号
|
||||||
3. **安装安全工具(可选)** - 按需安装 `tools/` 目录中的工具;未安装的工具在执行时会自动跳过或改用替代方案。常用示例:
|
3. **安装安全工具(可选)** - 按需安装 `tools/` 目录中的工具;未安装的工具在执行时会自动跳过或改用替代方案。常用示例:
|
||||||
|
|
||||||
**macOS(Homebrew):**
|
**macOS(Homebrew):**
|
||||||
@@ -242,7 +256,7 @@ go build -o cyberstrike-ai cmd/server/main.go
|
|||||||
|
|
||||||
**说明:** Python 虚拟环境(`venv/`)由 `run.sh` 自动创建和管理。需要 Python 的工具(如 `api-fuzzer`、`http-framework-test` 等)会自动使用该环境。
|
**说明:** Python 虚拟环境(`venv/`)由 `run.sh` 自动创建和管理。需要 Python 的工具(如 `api-fuzzer`、`http-framework-test` 等)会自动使用该环境。
|
||||||
|
|
||||||
### CyberStrikeAI 版本更新(无兼容性问题)
|
### 版本升级与兼容性
|
||||||
|
|
||||||
1. (首次使用)启用脚本:`chmod +x upgrade.sh`
|
1. (首次使用)启用脚本:`chmod +x upgrade.sh`
|
||||||
2. 一键升级:`./upgrade.sh`(可选参数:`--tag vX.Y.Z`、`--no-venv`、`--yes`)。本地的 `tools/`、`roles/`、`skills/` 会始终保留不被覆盖。
|
2. 一键升级:`./upgrade.sh`(可选参数:`--tag vX.Y.Z`、`--no-venv`、`--yes`)。本地的 `tools/`、`roles/`、`skills/` 会始终保留不被覆盖。
|
||||||
@@ -258,402 +272,32 @@ go build -o cyberstrike-ai cmd/server/main.go
|
|||||||
* 建议/需要 `rsync` 用于安全同步代码。
|
* 建议/需要 `rsync` 用于安全同步代码。
|
||||||
* 如果遇到 GitHub API 限流,运行前设置 `export GITHUB_TOKEN="..."` 再执行 `./upgrade.sh`。
|
* 如果遇到 GitHub API 限流,运行前设置 `export GITHUB_TOKEN="..."` 再执行 `./upgrade.sh`。
|
||||||
|
|
||||||
⚠️ **注意:** 仅适用于无兼容性变更的版本更新。若版本存在兼容性调整,此方法不适用。
|
⚠️ **升级前必读:** 请查看目标版本的 Release Notes,确认配置、数据库和 API 是否变化。即使只是补丁版本也应先备份,不能仅凭版本号判断兼容性。
|
||||||
|
|
||||||
**举例:** 无兼容性变更如 v1.3.1 → v1.3.2;有兼容性变更如 v1.3.1 → v1.4.0。项目采用语义化版本(SemVer):仅第三位(补丁号)变更时通常可安全按上述步骤升级;次版本号或主版本号变更时可能涉及配置、数据或接口调整,需查阅 release notes 再决定是否适用本方法。
|
|
||||||
|
|
||||||
### 常用流程
|
|
||||||
- **对话测试**:自然语言触发多步工具编排,SSE 实时输出。
|
|
||||||
- **单代理 / 多代理**:聊天可选 **Eino 单代理**(`/api/eino-agent/stream`)与 **多代理**(`/api/multi-agent/stream` + `orchestration`)。多代理需 `multi_agent.enabled: true`。MCP 工具桥接一致。
|
|
||||||
- **角色化测试**:从预设的安全测试角色(渗透测试、CTF、Web 应用扫描、API 安全测试等)中选择,自定义 AI 行为和可用工具。每个角色可应用自定义系统提示词,并可限制可用工具列表,实现聚焦的测试场景。
|
|
||||||
- **图编排**:在 **图编排** 页拖拽节点、连线并保存流程;在角色中绑定 `workflow_id` 后,该角色对话将按图执行(Agent、MCP 工具、条件分支等)。跨节点传参优先用 `{{outputs.变量名}}`。详见 [图编排使用说明](docs/zh-CN/workflow-graph.md)。
|
|
||||||
- **工具监控**:查看任务队列、执行日志、大文件附件。
|
|
||||||
- **会话历史**:所有对话与工具调用保存在 SQLite,可随时重放。
|
|
||||||
- **对话分组**:将对话按项目或主题组织到不同分组,支持置顶、重命名、删除等操作,所有数据持久化存储。
|
|
||||||
- **漏洞管理**:在测试过程中创建、更新和跟踪发现的漏洞。支持按严重程度(严重/高/中/低/信息)、状态(待确认/已确认/已修复/误报)和对话进行过滤,查看统计信息并导出发现。
|
|
||||||
- **批量任务管理**:创建任务队列,批量添加多个任务,执行前可编辑或删除任务,然后依次顺序执行。每个任务会作为独立对话执行,支持完整的状态跟踪(待执行/执行中/已完成/失败/已取消)和执行历史。
|
|
||||||
- **WebShell 管理**:添加并管理 WebShell 连接(PHP/ASP/ASPX/JSP 或自定义类型)。使用虚拟终端执行命令(带命令历史与快捷命令),使用文件管理浏览、读取、编辑、上传与删除目标文件,并支持按路径导航和名称过滤。连接信息持久化存储于 SQLite,支持 GET/POST 及可配置命令参数(兼容冰蝎/蚁剑等)。
|
|
||||||
- **内置 C2**:在 Web 界面或 `/api/c2/*` 创建/启动 **监听器**、生成 **Payload**、查看 **会话**、下发 **任务** 并订阅 **事件(SSE)**。智能体与外部客户端通过 **C2 MCP 工具族**(含 **`c2_task`** 等)编排;开启人机协同时,高风险任务可走审批。**仅用于已获明确授权的目标。**
|
|
||||||
- **可视化配置**:在界面中切换模型、启停工具、设置迭代次数等。
|
|
||||||
- **人机协同(HITL)**:侧栏设置协同模式与免审批工具(逗号或换行);全局白名单见 `config.yaml` 的 `hitl.tool_whitelist`。审计 Agent 可通过 `hitl.audit_model` 单独配置低成本模型,适合人工审计压力较大时接管常规审批。点「**应用**」可写浏览器/服务端并合并新增工具进配置(**无需重启**)。**新对话**保留侧栏选择;导航 **人机协同** 处理待审批。从侧栏删掉工具不会自动从配置文件移除全局项,需手改 `config.yaml`。
|
|
||||||
|
|
||||||
### 默认安全措施
|
|
||||||
- 设置面板内置必填校验,防止漏配 API Key/Base URL/模型。
|
|
||||||
- `auth.password` 为空时自动生成 24 位强口令并写回 `config.yaml`。
|
|
||||||
- 所有 API(除登录外)都需携带 Bearer Token,统一鉴权中间件拦截。
|
|
||||||
- 每个工具执行都带有超时、日志和错误隔离。
|
|
||||||
|
|
||||||
## 进阶使用
|
|
||||||
|
|
||||||
### 角色化测试
|
|
||||||
- **预设角色**:系统内置 12+ 个预设的安全测试角色(渗透测试、CTF、Web 应用扫描、API 安全测试、二进制分析、云安全审计等),位于 `roles/` 目录。
|
|
||||||
- **自定义提示词**:每个角色可定义 `user_prompt`,会在用户消息前自动添加,引导 AI 采用特定的测试方法和关注重点。
|
|
||||||
- **工具限制**:角色可指定 `tools` 列表,限制可用工具,实现聚焦的测试流程(如 CTF 角色限制为 CTF 专用工具)。
|
|
||||||
- **Skills**:技能包位于 `skills_dir`;启用 **`multi_agent.eino_skills`** 后,**单代理与多代理**均可通过 Eino **`skill`** 工具按需加载。可选 **`eino_middleware`**(tool_search、plantask、reduction、checkpoint、Summarization 转录等)与本机 read_file/glob/grep 等见文档。
|
|
||||||
- **轻松创建角色**:通过在 `roles/` 目录添加 YAML 文件即可创建自定义角色。每个角色定义 `name`、`description`、`user_prompt`、`icon`、`tools`、`enabled` 字段。
|
|
||||||
- **Web 界面集成**:在聊天界面通过下拉菜单选择角色。角色选择会影响 AI 行为和可用工具建议。
|
|
||||||
|
|
||||||
**创建自定义角色示例:**
|
|
||||||
1. 在 `roles/` 目录创建 YAML 文件(如 `roles/custom-role.yaml`):
|
|
||||||
```yaml
|
|
||||||
name: 自定义角色
|
|
||||||
description: 专用测试场景
|
|
||||||
user_prompt: 你是一个专注于 API 安全的专业安全测试人员...
|
|
||||||
icon: "\U0001F4E1"
|
|
||||||
tools:
|
|
||||||
- api-fuzzer
|
|
||||||
- arjun
|
|
||||||
- graphql-scanner
|
|
||||||
enabled: true
|
|
||||||
```
|
|
||||||
2. 重启服务或重新加载配置,角色会出现在角色选择下拉菜单中。
|
|
||||||
|
|
||||||
### 多代理模式(Eino:Deep / Plan-Execute / Supervisor)
|
|
||||||
- **能力说明**:在 **Eino 单代理**(`/api/eino-agent*`)之外,多代理基于 CloudWeGo **Eino** `adk/prebuilt`:**`deep`**、**`plan_execute`**、**`supervisor`**;客户端 **`orchestration`** 选择(缺省 `deep`)。模式定位按 Eino ADK 最佳实践区分:**Deep** 适合复杂安全测试与 task 子代理协作;**Plan-Execute** 适合目标明确的规划 → 执行 → 重规划闭环;**Supervisor** 适合多个专业子代理动态分派的专家路由场景。
|
|
||||||
- **Markdown 定义**(`agents_dir`,默认 `agents/`):
|
|
||||||
- **Deep 主代理**:`orchestrator.md` 或唯一 `kind: orchestrator` 的 `.md`;正文或 `multi_agent.orchestrator_instruction`,再回退 Eino 默认。
|
|
||||||
- **Plan-Execute 主代理**:固定 **`orchestrator-plan-execute.md`**(另可配 `orchestrator_instruction_plan_execute`)。
|
|
||||||
- **Supervisor 主代理**:固定 **`orchestrator-supervisor.md`**(另可配 `orchestrator_instruction_supervisor`);至少需一名子代理,只有一名子代理时会提示专家路由价值有限。
|
|
||||||
- **子代理**(**deep** / **supervisor**):其余 `*.md`;标成 orchestrator 的不会进入 `task` 列表。
|
|
||||||
- **界面管理**:**Agents → Agent 管理**;API `/api/multi-agent/markdown-agents`。
|
|
||||||
- **配置项**:`multi_agent`:`enabled`、`robot_default_agent_mode`、`batch_use_multi_agent`、`max_iteration`、`plan_execute_loop_max_iterations`、各模式 orchestrator 指令字段、可选 YAML `sub_agents` 与目录合并(同 `id` → Markdown 优先)、**`eino_skills`**、**`eino_middleware`**。
|
|
||||||
- **长任务与恢复**:`checkpoint_dir` 支持进程崩溃后 ADK **断点续跑**(与基于 trace 的「中断继续」不同)。`deep_model_retry_max_retries` 在同一次 LLM 调用内重试瞬时 API 失败。**Summarization** 触发压缩时会写入过滤后的 **transcript**,摘要消息中带路径,模型可用 `read_file` 找回扫描输出等压缩前细节。
|
|
||||||
- **更多细节**:[docs/zh-CN/MULTI_AGENT_EINO.md](docs/zh-CN/MULTI_AGENT_EINO.md)(流式、机器人、批量、中间件差异)。
|
|
||||||
|
|
||||||
### Skills 技能系统(Agent Skills + Eino)
|
|
||||||
- **目录规范**:与 [Agent Skills](https://platform.claude.com/docs/en/agents-and-tools/agent-skills/overview) 一致,**仅**需目录下的 **`SKILL.md`**:YAML 头只用官方的 **`name` 与 `description`**,正文为 Markdown。可选同目录其他文件(`FORMS.md`、`REFERENCE.md`、`scripts/*` 等)。**不使用 `SKILL.yaml`**(Claude / Eino 官方均无此文件);章节、`scripts/` 列表、渐进式行为由运行时从正文与磁盘 **自动推导**。
|
|
||||||
- **运行侧重构**:**`skills_dir`** 为技能包唯一根目录;**多代理** 通过 Eino 官方 **`skill`** 中间件做 **渐进式披露**(模型按 **name** 调用 `skill`,而非一次性注入全文)。由 **`multi_agent.eino_skills`** 控制:`disable`、`filesystem_tools`(本机读写与 Shell)、`skill_tool_name`。
|
|
||||||
- **Eino / 知识流水线**:技能包可切分为 `schema.Document`,供 `FilesystemSkillsRetriever`(`skills.AsEinoRetriever()`)在 **compose** 图(如索引/编排)中使用。
|
|
||||||
- **HTTP 管理**:`/api/skills` 列表与 `depth=summary|full`、`section`、`resource_path` 等仍用于 Web 与运维;**模型侧** 多代理走 **`skill`** 工具,而非 MCP。
|
|
||||||
- **可选 `eino_middleware`**:如 `tool_search`(动态工具列表)、`patch_tool_calls`、**`plantask`**(Eino `TaskCreate` / `TaskGet` / `TaskUpdate` / `TaskList`;JSON 存于 `skills_dir/.eino/plantask/<会话ID>/`;**全部**任务标为 completed 后 Eino 会清理任务文件)、`reduction`、**`checkpoint_dir`**(如 `data/eino-checkpoints/`)、**`deep_model_retry_max_retries`**、**`deep_output_key`**、task 描述前缀等,见 `config.yaml` 与 `internal/config/config.go`。
|
|
||||||
- **自带示例**:`skills/cyberstrike-eino-demo/`;说明见 `skills/README.md`。
|
|
||||||
|
|
||||||
**新建技能:**
|
|
||||||
1. 在 `skills/` 下创建 `<skill-id>/`,放入标准 `SKILL.md`(及任意可选文件),或直接解压开源技能包到该目录。
|
|
||||||
2. 启用 **`multi_agent.eino_skills`** 并使用 **多代理** 会话,由模型通过 **`skill`** 工具按包 **name** 加载。
|
|
||||||
|
|
||||||
### 工具编排与扩展
|
|
||||||
- `tools/*.yaml` 定义命令、参数、提示词与元数据,可热加载。
|
|
||||||
- `security.tools_dir` 指向目录即可批量启用;仍支持在主配置里内联定义。
|
|
||||||
- **大工具输出**:超过 `reduction_max_length_for_trunc` 时由 Eino reduction 摘要,完整内容落盘至 `tmp/reduction/`;按 `<persisted-output>` 中的路径用 `read_file` 读取。
|
|
||||||
- **结果压缩/摘要**:多兆字节日志可先压缩或生成摘要再写入 SQLite,减小档案体积。
|
|
||||||
|
|
||||||
**自定义工具的一般步骤**
|
|
||||||
1. 复制 `tools/` 下现有示例(如 `tools/nmap.yaml` 或 `tools/ffuf.yaml`)。
|
|
||||||
2. 修改 `name`、`command`、`args`、`short_description` 等基础信息。
|
|
||||||
3. 在 `parameters[]` 中声明位置参数或带 flag 的参数,方便智能体自动拼装命令。
|
|
||||||
4. 视需要补充 `description` 或 `notes`,给 AI 额外上下文或结果解读提示。
|
|
||||||
5. 重启服务或在界面中重新加载配置,新工具即可在 Settings 面板中启用/禁用。
|
|
||||||
|
|
||||||
### 攻击链分析
|
|
||||||
- 智能体解析每次对话,抽取目标、工具、漏洞与因果关系。
|
|
||||||
- Web 端可交互式查看链路节点、风险级别及时间轴,支持导出报告。
|
|
||||||
|
|
||||||
### WebShell 管理
|
|
||||||
- **连接管理**:在 Web 界面进入 **WebShell 管理**,可添加、编辑或删除 WebShell 连接。每条连接包含:Shell 地址、密码/密钥、Shell 类型(PHP/ASP/ASPX/JSP/自定义)、请求方式(GET/POST)、命令参数名(默认 `cmd`)、备注等信息,并持久化存储在 SQLite,兼容冰蝎、蚁剑等常见客户端。
|
|
||||||
- **虚拟终端**:选择连接后,在 **虚拟终端** 标签页中执行任意命令,支持命令历史与常用快捷命令(whoami/id/ls/pwd 等),输出在浏览器中实时显示,支持 Ctrl+L 清屏。
|
|
||||||
- **文件管理**:在 **文件管理** 标签页中可列出目录、读取/编辑文件、删除文件、新建文件/目录、上传文件(大文件分片上传)、重命名路径以及下载勾选文件,并支持面包屑导航与名称过滤。
|
|
||||||
- **AI 助手**:在 **AI 助手** 标签页中与智能体对话,由系统自动结合当前 WebShell 连接执行工具与命令,侧边栏展示该连接下的所有历史会话,支持多轮追踪与查看。
|
|
||||||
- **连通性测试**:使用 **测试连通性** 可在执行命令前通过一次 `echo 1` 调用校验 Shell 地址、密码与命令参数是否正确。
|
|
||||||
- **持久化**:所有 WebShell 连接与相关 AI 会话均保存在 SQLite(与对话共用数据库),服务重启后仍可继续使用。
|
|
||||||
|
|
||||||
### 内置 C2(Command & Control)
|
|
||||||
- **定位**:平台内置的 **AI 原生** C2 能力栈——监听器接入植入体(Beacon),服务端以 SQLite 持久化 **会话** 与 **任务**,通过 **事件总线** 推送变更(含 **SSE**),并由鉴权后的 **REST** 与 MCP 统一对外。
|
|
||||||
- **监听器与传输**:支持 `tcp_reverse`、`http_beacon`、`https_beacon`、`websocket`;按监听器独立密钥;数据库中标记为运行中的监听器可在 **服务重启后尝试恢复**。
|
|
||||||
- **与智能体联动**:通过 **`c2_task` 等 C2 MCP 工具** 与现有对话/多代理工具链协同;在会话策略需要时,危险任务类型可走既有 **人机协同(HITL)** 审批流。
|
|
||||||
- **安全提示**:**仅**在实验环境或 **已获完整书面授权** 的对抗演练中使用;结合网络隔离、强鉴权及 HITL/白名单等策略管控风险。
|
|
||||||
|
|
||||||
### MCP 全场景
|
|
||||||
- **Web 模式**:自带 HTTP MCP 服务供前端调用。
|
|
||||||
- **MCP stdio 模式**:`go run cmd/mcp-stdio/main.go` 可接入 Cursor/命令行。
|
|
||||||
- **外部 MCP 联邦**:在设置中注册第三方 MCP(HTTP/stdio/SSE),按需启停并实时查看调用统计与健康度。
|
|
||||||
- **可选 MCP 服务**:项目中的 [`mcp-servers/`](mcp-servers/README_CN.md) 目录提供独立 MCP(如反向 Shell),采用标准 MCP stdio,可在 CyberStrikeAI(设置 → 外部 MCP)、Cursor、VS Code 等任意支持 MCP 的客户端中使用。
|
|
||||||
|
|
||||||
#### MCP stdio 快速集成
|
|
||||||
1. **编译可执行文件**(在项目根目录执行):
|
|
||||||
```bash
|
|
||||||
go build -o cyberstrike-ai-mcp cmd/mcp-stdio/main.go
|
|
||||||
```
|
|
||||||
2. **在 Cursor 中配置**
|
|
||||||
打开 `Settings → Tools & MCP → Add Custom MCP`,选择 **Command**,指定编译后的程序与配置文件:
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"mcpServers": {
|
|
||||||
"cyberstrike-ai": {
|
|
||||||
"command": "/absolute/path/to/cyberstrike-ai-mcp",
|
|
||||||
"args": [
|
|
||||||
"--config",
|
|
||||||
"/absolute/path/to/config.yaml"
|
|
||||||
]
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
将路径替换成你本地的实际地址,Cursor 会自动启动 stdio 版本的 MCP。
|
|
||||||
|
|
||||||
#### MCP HTTP 快速集成(Cursor / Claude Code)
|
|
||||||
HTTP MCP 服务在独立端口(默认 `8081`)运行,支持 **Header 鉴权**:仅携带正确 header 的客户端可调用工具。
|
|
||||||
|
|
||||||
1. **在配置中启用 MCP** – 在 `config.yaml` 中设置 `mcp.enabled: true`,并按需设置 `mcp.host` / `mcp.port`。若需鉴权(端口对外暴露时建议开启),可设置:
|
|
||||||
- `mcp.auth_header`:鉴权用的 header 名(如 `X-MCP-Token`);
|
|
||||||
- `mcp.auth_header_value`:鉴权密钥。**留空**时,首次启动会自动生成随机密钥并写回配置文件。
|
|
||||||
2. **启动服务** – 执行 `./run.sh` 或 `go run cmd/server/main.go`。MCP 端点为 `http://<host>:<port>/mcp`(例如 `http://localhost:8081/mcp`)。
|
|
||||||
3. **从终端复制 JSON** – 启用 MCP 后,启动时会在终端打印一段 **可直接复制的 JSON**。若 `auth_header_value` 留空,会自动生成并写入配置,打印内容中会包含 URL 与 headers。
|
|
||||||
4. **在 Cursor 或 Claude Code 中使用**:
|
|
||||||
- **Cursor**:将整段 JSON 粘贴到 `~/.cursor/mcp.json` 或项目下的 `.cursor/mcp.json` 的 `mcpServers` 中(或合并进现有 `mcpServers`)。
|
|
||||||
- **Claude Code**:粘贴到 `.mcp.json` 或 `~/.claude.json` 的 `mcpServers` 中。
|
|
||||||
|
|
||||||
终端打印示例(开启鉴权时):
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"mcpServers": {
|
|
||||||
"cyberstrike-ai": {
|
|
||||||
"url": "http://localhost:8081/mcp",
|
|
||||||
"headers": {
|
|
||||||
"X-MCP-Token": "<自动生成或你配置的值>"
|
|
||||||
},
|
|
||||||
"type": "http"
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
若不配置 `auth_header` / `auth_header_value`,则端点不鉴权(仅适合本机或可信网络)。
|
|
||||||
|
|
||||||
#### 外部 MCP 联邦(HTTP/stdio/SSE)
|
|
||||||
CyberStrikeAI 支持通过三种传输模式连接外部 MCP 服务器:
|
|
||||||
- **HTTP 模式** – 通过 HTTP POST 进行传统的请求/响应通信
|
|
||||||
- **stdio 模式** – 通过标准输入/输出进行进程间通信
|
|
||||||
- **SSE 模式** – 通过 Server-Sent Events 实现实时流式通信
|
|
||||||
|
|
||||||
添加外部 MCP 服务器:
|
|
||||||
1. 打开 Web 界面,进入 **设置 → 外部MCP**。
|
|
||||||
2. 点击 **添加外部MCP**,以 JSON 格式提供配置:
|
|
||||||
|
|
||||||
**HTTP 模式示例:**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"my-http-mcp": {
|
|
||||||
"transport": "http",
|
|
||||||
"url": "http://127.0.0.1:8081/mcp",
|
|
||||||
"description": "HTTP MCP 服务器",
|
|
||||||
"timeout": 30
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**stdio 模式示例:**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"my-stdio-mcp": {
|
|
||||||
"command": "python3",
|
|
||||||
"args": ["/path/to/mcp-server.py"],
|
|
||||||
"description": "stdio MCP 服务器",
|
|
||||||
"timeout": 30
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
**SSE 模式示例:**
|
|
||||||
```json
|
|
||||||
{
|
|
||||||
"my-sse-mcp": {
|
|
||||||
"transport": "sse",
|
|
||||||
"url": "http://127.0.0.1:8082/sse",
|
|
||||||
"description": "SSE MCP 服务器",
|
|
||||||
"timeout": 30
|
|
||||||
}
|
|
||||||
}
|
|
||||||
```
|
|
||||||
|
|
||||||
3. 点击 **保存**,然后点击 **启动** 连接服务器。
|
|
||||||
4. 实时监控连接状态、工具数量和健康度。
|
|
||||||
|
|
||||||
**SSE 模式优势:**
|
|
||||||
- 通过 Server-Sent Events 实现实时双向通信
|
|
||||||
- 适用于需要持续数据流的场景
|
|
||||||
- 对于基于推送的通知,延迟更低
|
|
||||||
|
|
||||||
可在 `cmd/test-sse-mcp-server/` 目录找到用于验证的测试 SSE MCP 服务器。
|
|
||||||
|
|
||||||
|
|
||||||
### 知识库功能
|
## 配置
|
||||||
- **向量检索**:AI 智能体在对话过程中可自动调用 `search_knowledge_base` 工具搜索知识库中的安全知识。
|
|
||||||
- **RAG 管线(始终启用)**:**MultiQuery**(LLM 查询改写)→ 向量预取与融合 → **HTTP 精排**(DashScope `gte-rerank` 或 Cohere 兼容 `/v1/rerank`)→ 后处理(规范化去重、字符/token 预算、最终 top_k)。精排失败时自动降级为融合排序,检索仍可用。
|
|
||||||
- **向量相似度**:基于嵌入余弦相似度与相似度阈值过滤(与 Eino `retriever.Retriever` 语义一致)。
|
|
||||||
- **自动索引**:扫描 `knowledge_base/` 目录下的 Markdown 文件,自动构建向量嵌入索引(Eino Markdown 标题切分 + 递归分块)。
|
|
||||||
- **Web 管理**:通过 Web 界面创建、更新、删除知识项,支持分类管理;设置页可配置 MultiQuery / 精排 / 预取候选数。
|
|
||||||
- **检索日志**:记录所有知识检索操作,便于审计与调试。
|
|
||||||
|
|
||||||
**知识库配置步骤:**
|
请以 [`config.example.yaml`](config.example.yaml) 作为权威配置模板,只复制当前环境需要的配置。最少需要配置服务监听地址和一个 OpenAI 兼容模型:
|
||||||
1. **启用功能**:在 `config.yaml` 中设置 `knowledge.enabled: true`:
|
|
||||||
```yaml
|
|
||||||
knowledge:
|
|
||||||
enabled: true
|
|
||||||
base_path: knowledge_base
|
|
||||||
embedding:
|
|
||||||
provider: openai
|
|
||||||
model: text-embedding-v4
|
|
||||||
base_url: "https://api.openai.com/v1" # 或你的嵌入模型 API
|
|
||||||
api_key: "sk-xxx"
|
|
||||||
retrieval:
|
|
||||||
top_k: 5
|
|
||||||
similarity_threshold: 0.7
|
|
||||||
multi_query:
|
|
||||||
max_queries: 4 # LLM 改写变体上限(始终启用)
|
|
||||||
rerank: # 精排始终启用;留空则继承 openai/embedding 凭据
|
|
||||||
provider: "" # 空=按 base_url 推断 dashscope | cohere
|
|
||||||
model: "" # 空=DashScope→gte-rerank,Cohere→rerank-multilingual-v3.0
|
|
||||||
base_url: ""
|
|
||||||
api_key: ""
|
|
||||||
post_retrieve:
|
|
||||||
prefetch_top_k: 20 # 每条 MultiQuery 变体的向量候选数;0=max(top_k×4, 20)
|
|
||||||
max_context_chars: 0
|
|
||||||
max_context_tokens: 0
|
|
||||||
```
|
|
||||||
2. **添加知识文件**:将 Markdown 文件放入 `knowledge_base/` 目录,按分类组织(如 `knowledge_base/SQL注入/README.md`)。
|
|
||||||
3. **扫描索引**:在 Web 界面中点击"扫描知识库",系统会自动导入文件并构建向量索引。
|
|
||||||
4. **对话中使用**:AI 智能体在需要安全知识时会自动调用知识检索工具。你也可以显式要求:"搜索知识库中关于 SQL 注入的技术"。
|
|
||||||
|
|
||||||
**知识库结构说明:**
|
|
||||||
- 文件按分类组织(目录名作为分类)。
|
|
||||||
- 每个 Markdown 文件自动切块并生成向量嵌入。
|
|
||||||
- 支持增量更新,修改后的文件会自动重新索引。
|
|
||||||
|
|
||||||
|
|
||||||
### 自动化与安全
|
|
||||||
- **REST API**:认证、会话、任务、监控、漏洞管理、角色管理等接口全部开放,可与 CI/CD 集成。
|
|
||||||
- **多代理 API**:`POST /api/multi-agent/stream`(SSE,需启用多代理)、`POST /api/multi-agent`(非流式);Markdown 子代理/主代理管理见 `/api/multi-agent/markdown-agents`(列表/读写/增删)。
|
|
||||||
- **角色管理 API**:通过 `/api/roles` 端点管理安全测试角色:`GET /api/roles`(列表)、`GET /api/roles/:name`(获取角色)、`POST /api/roles`(创建角色)、`PUT /api/roles/:name`(更新角色)、`DELETE /api/roles/:name`(删除角色)。角色以 YAML 文件形式存储在 `roles/` 目录,支持热加载。
|
|
||||||
- **漏洞管理 API**:通过 `/api/vulnerabilities` 端点管理漏洞:`GET /api/vulnerabilities`(列表,支持过滤)、`POST /api/vulnerabilities`(创建)、`GET /api/vulnerabilities/:id`(获取)、`PUT /api/vulnerabilities/:id`(更新)、`DELETE /api/vulnerabilities/:id`(删除)、`GET /api/vulnerabilities/stats`(统计)。
|
|
||||||
- **批量任务 API**:通过 `/api/batch-tasks` 端点管理批量任务队列:`POST /api/batch-tasks`(创建队列)、`GET /api/batch-tasks`(列表)、`GET /api/batch-tasks/:queueId`(获取队列)、`POST /api/batch-tasks/:queueId/start`(开始执行)、`POST /api/batch-tasks/:queueId/cancel`(取消)、`DELETE /api/batch-tasks/:queueId`(删除队列)、`POST /api/batch-tasks/:queueId/tasks`(添加任务)、`PUT /api/batch-tasks/:queueId/tasks/:taskId`(更新任务)、`DELETE /api/batch-tasks/:queueId/tasks/:taskId`(删除任务)。任务依次顺序执行,每个任务创建独立对话,支持完整状态跟踪。
|
|
||||||
- **WebShell API**:通过 `/api/webshell/connections`(GET 列表、POST 创建、PUT 更新、DELETE 删除)及 `/api/webshell/exec`(执行命令)、`/api/webshell/fileop`(列出/读取/写入/删除文件)管理 WebShell 连接与执行操作。
|
|
||||||
- **C2 API**:在 `/api/c2/*` 管理监听器、会话、任务、Payload、文件与事件(如监听器增删改查/启停、会话休眠、任务创建/取消/等待、Payload 构建/下载、事件流等)。
|
|
||||||
- **任务控制**:支持暂停/终止长任务、修改参数后重跑、流式获取日志。
|
|
||||||
- **安全管理**:`/api/auth/change-password` 可即时轮换口令;建议在暴露 MCP 端口时配合网络层 ACL。
|
|
||||||
|
|
||||||
## 配置参考
|
|
||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
auth:
|
|
||||||
password: "change-me"
|
|
||||||
session_duration_hours: 12
|
|
||||||
server:
|
server:
|
||||||
host: "0.0.0.0"
|
host: "127.0.0.1"
|
||||||
port: 8080
|
port: 8080
|
||||||
log:
|
|
||||||
level: "info"
|
|
||||||
output: "stdout"
|
|
||||||
mcp:
|
|
||||||
enabled: true
|
|
||||||
host: "0.0.0.0"
|
|
||||||
port: 8081
|
|
||||||
auth_header: "X-MCP-Token" # 可选;留空则不鉴权
|
|
||||||
auth_header_value: "" # 可选;留空则首次启动自动生成并写回
|
|
||||||
openai:
|
openai:
|
||||||
api_key: "sk-xxx"
|
api_key: "${OPENAI_API_KEY}"
|
||||||
base_url: "https://api.deepseek.com/v1"
|
base_url: "https://api.openai.com/v1"
|
||||||
model: "deepseek-chat"
|
model: "your-model"
|
||||||
database:
|
|
||||||
path: "data/conversations.db"
|
|
||||||
knowledge_db_path: "data/knowledge.db" # 可选:知识库独立数据库
|
|
||||||
security:
|
|
||||||
tools_dir: "tools"
|
|
||||||
knowledge:
|
|
||||||
enabled: false # 是否启用知识库功能
|
|
||||||
base_path: "knowledge_base" # 知识库目录路径
|
|
||||||
embedding:
|
|
||||||
provider: "openai" # 嵌入模型提供商(目前仅支持 openai)
|
|
||||||
model: "text-embedding-v4" # 嵌入模型名称
|
|
||||||
base_url: "" # 留空则使用 OpenAI 配置的 base_url
|
|
||||||
api_key: "" # 留空则使用 OpenAI 配置的 api_key
|
|
||||||
retrieval:
|
|
||||||
top_k: 5 # 检索返回的 Top-K 结果数量
|
|
||||||
similarity_threshold: 0.7 # 余弦相似度阈值(0-1),低于此值的结果将被过滤
|
|
||||||
multi_query:
|
|
||||||
max_queries: 4 # MultiQuery 改写变体上限(始终启用)
|
|
||||||
rerank: # HTTP 精排(始终启用);留空则继承 openai/embedding 凭据
|
|
||||||
provider: ""
|
|
||||||
model: ""
|
|
||||||
base_url: ""
|
|
||||||
api_key: ""
|
|
||||||
post_retrieve:
|
|
||||||
prefetch_top_k: 20 # 每条 MultiQuery 变体;0=max(top_k×4, 20)
|
|
||||||
max_context_chars: 0
|
|
||||||
max_context_tokens: 0
|
|
||||||
roles_dir: "roles" # 角色配置文件目录(相对于配置文件所在目录)
|
|
||||||
skills_dir: "skills" # Skills 目录(相对于配置文件所在目录)
|
|
||||||
agents_dir: "agents" # 多代理 Markdown(主代理 orchestrator.md + 子代理 *.md)
|
|
||||||
multi_agent:
|
|
||||||
enabled: false
|
|
||||||
default_mode: "eino_single" # eino_single | multi(开启多代理时的界面默认模式)
|
|
||||||
robot_default_agent_mode: eino_single
|
|
||||||
batch_use_multi_agent: false
|
|
||||||
orchestrator_instruction: "" # Deep;orchestrator.md 正文为空时使用
|
|
||||||
# orchestrator_instruction_plan_execute / orchestrator_instruction_supervisor 可选
|
|
||||||
# eino_skills: { disable: false, filesystem_tools: true, skill_tool_name: skill }
|
|
||||||
# eino_middleware: plantask_enable、checkpoint_dir、deep_model_retry_max_retries、deep_output_key 等
|
|
||||||
project:
|
|
||||||
enabled: true # 启用项目黑板与事实 MCP 工具
|
|
||||||
fact_index_max_runes: 65000
|
|
||||||
fact_summary_max_runes: 24000
|
|
||||||
default_inject_deprecated: false
|
|
||||||
```
|
```
|
||||||
|
|
||||||
### 工具模版示例(`tools/nmap.yaml`)
|
不要提交真实凭证。将服务暴露到 localhost 之外前,请阅读[配置参考](docs/zh-CN/configuration.md)、[推荐配置画像](docs/zh-CN/configuration-profiles.md)和[安全加固指南](docs/zh-CN/security-hardening.md)。
|
||||||
|
|
||||||
```yaml
|
|
||||||
name: "nmap"
|
|
||||||
command: "nmap"
|
|
||||||
args: ["-sT", "-sV", "-sC"]
|
|
||||||
enabled: true
|
|
||||||
short_description: "网络资产扫描与服务指纹识别"
|
|
||||||
parameters:
|
|
||||||
- name: "target"
|
|
||||||
type: "string"
|
|
||||||
description: "IP 或域名"
|
|
||||||
required: true
|
|
||||||
position: 0
|
|
||||||
- name: "ports"
|
|
||||||
type: "string"
|
|
||||||
flag: "-p"
|
|
||||||
description: "端口范围,如 1-1000"
|
|
||||||
```
|
|
||||||
|
|
||||||
### 角色配置示例(`roles/渗透测试.yaml`)
|
|
||||||
|
|
||||||
```yaml
|
|
||||||
name: 渗透测试
|
|
||||||
description: 专业渗透测试专家,全面深入的漏洞检测
|
|
||||||
user_prompt: 你是一个专业的网络安全渗透测试专家。请使用专业的渗透测试方法和工具,对目标进行全面的安全测试,包括但不限于SQL注入、XSS、CSRF、文件包含、命令执行等常见漏洞。
|
|
||||||
icon: "\U0001F3AF"
|
|
||||||
tools:
|
|
||||||
- nmap
|
|
||||||
- sqlmap
|
|
||||||
- nuclei
|
|
||||||
- burpsuite
|
|
||||||
- metasploit
|
|
||||||
- httpx
|
|
||||||
- record_vulnerability
|
|
||||||
- list_knowledge_risk_types
|
|
||||||
- search_knowledge_base
|
|
||||||
enabled: true
|
|
||||||
```
|
|
||||||
|
|
||||||
## 相关文档
|
## 相关文档
|
||||||
|
|
||||||
- [文档导航](docs/README.md):部署、配置、安全模型、API、知识库、C2、WebShell、MCP、开发、测试、排错等完整专题入口。
|
- **新用户:** [部署指南](docs/zh-CN/deployment.md) → [配置参考](docs/zh-CN/configuration.md) → [排错指南](docs/zh-CN/troubleshooting.md)
|
||||||
- [部署指南](docs/zh-CN/deployment.md):源码/二进制运行、HTTPS、反向代理、systemd、备份、升级与回滚。
|
- **运维人员:** [配置画像](docs/zh-CN/configuration-profiles.md) → [安全加固](docs/zh-CN/security-hardening.md) → [运维 Runbooks](docs/zh-CN/runbooks.md)
|
||||||
- [运维 Runbooks](docs/zh-CN/runbooks.md):生产部署、外部 MCP、知识库、授权 Web 测试、C2 清理等可执行流程。
|
- **集成开发:** [API 参考](docs/zh-CN/api-reference.md) → [API Recipes](docs/zh-CN/api-recipes.md) → [MCP 联邦](docs/zh-CN/mcp-federation.md)
|
||||||
- [安全加固指南](docs/zh-CN/security-hardening.md):上线前基线、HITL 白名单、反向代理、文件权限和周期巡检。
|
- **项目贡献:** [开发者指南](docs/zh-CN/developer-guide.md) → [测试指南](docs/zh-CN/testing.md) → [贡献规范](docs/zh-CN/contributing-guide.md)
|
||||||
- [API Recipes](docs/zh-CN/api-recipes.md):登录、Agent、流式、多代理、上传、漏洞、知识库和审计导出调用示例。
|
- **全部专题:** [中文文档](docs/zh-CN/README.md) · [双语文档索引](docs/README.md)
|
||||||
- [配置参考](docs/zh-CN/configuration.md):`config.yaml` 各配置段、推荐值和修改建议。
|
|
||||||
- [安全模型](docs/zh-CN/security-model.md):认证、工具执行、HITL、审计、C2/WebShell 和数据安全边界。
|
|
||||||
- [RBAC 权限管理](docs/zh-CN/rbac.md):平台用户、系统/自定义角色、权限目录、逐权限 Scope、资源授权、Agent/MCP/机器人边界与 API 示例。
|
|
||||||
- [API 参考](docs/zh-CN/api-reference.md):OpenAPI、认证、Agent、项目、知识库、C2、WebShell 等接口入口。
|
|
||||||
- [多代理模式(Eino)](docs/zh-CN/MULTI_AGENT_EINO.md):**Deep**、**Plan-Execute**、**Supervisor**、`agents/*.md`、`eino_skills` / `eino_middleware`、接口与流式说明。
|
|
||||||
- [图编排使用说明](docs/zh-CN/workflow-graph.md):可视化流程搭建、节点配置、`previous` / `outputs` 变量传参与角色绑定。
|
|
||||||
- [机器人使用说明](docs/zh-CN/robot.md):各平台接入、RBAC 逐用户绑定/服务账号模式、发送者白名单、命令、验证与排查。
|
|
||||||
- [人机协同最佳实践](docs/zh-CN/hitl-best-practices.md):审批方模式、白名单、审计 Agent 提示词策略与独立小模型配置。
|
|
||||||
|
|
||||||
## 项目结构
|
## 项目结构
|
||||||
|
|
||||||
@@ -668,6 +312,7 @@ CyberStrikeAI/
|
|||||||
├── agents/ # 多代理 Markdown(orchestrator.md + 子代理 *.md)
|
├── agents/ # 多代理 Markdown(orchestrator.md + 子代理 *.md)
|
||||||
├── docs/ # 专题文档(部署、配置、安全、API、知识库、C2、WebShell 等)
|
├── docs/ # 专题文档(部署、配置、安全、API、知识库、C2、WebShell 等)
|
||||||
├── images/ # 文档配图
|
├── images/ # 文档配图
|
||||||
|
├── scripts/ # 仓库维护检查,包括文档校验
|
||||||
├── config.yaml # 运行配置
|
├── config.yaml # 运行配置
|
||||||
├── run.sh # 启动脚本
|
├── run.sh # 启动脚本
|
||||||
└── README*.md
|
└── README*.md
|
||||||
@@ -707,6 +352,26 @@ CyberStrikeAI 现已加入 [404星链计划](https://github.com/knownsec/404Star
|
|||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
## 社区与支持
|
||||||
|
|
||||||
|
- 在 [Discord](https://discord.gg/8PjVCMu8Zw) 加入社区。
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><strong>微信群</strong></summary>
|
||||||
|
|
||||||
|
<img src="./images/wechat-group-cyberstrikeai-qr.jpg" alt="CyberStrikeAI 微信群二维码" width="280">
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><strong>通过微信支付或支付宝赞助</strong></summary>
|
||||||
|
|
||||||
|
<div align="center">
|
||||||
|
<img src="./images/sponsor-wechat-alipay-qr.jpg" alt="微信与支付宝赞助二维码" width="480">
|
||||||
|
</div>
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|
||||||
## 许可证
|
## 许可证
|
||||||
|
|
||||||
CyberStrikeAI 采用 **Apache License 2.0** 开源许可。
|
CyberStrikeAI 采用 **Apache License 2.0** 开源许可。
|
||||||
|
|||||||
+8
-11
@@ -5,6 +5,7 @@ import (
|
|||||||
"cyberstrike-ai/internal/app"
|
"cyberstrike-ai/internal/app"
|
||||||
"cyberstrike-ai/internal/config"
|
"cyberstrike-ai/internal/config"
|
||||||
"cyberstrike-ai/internal/logger"
|
"cyberstrike-ai/internal/logger"
|
||||||
|
"cyberstrike-ai/internal/termout"
|
||||||
"flag"
|
"flag"
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
@@ -47,8 +48,7 @@ func main() {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
if localConfig.Created {
|
if localConfig.Created {
|
||||||
cfg.Auth.GeneratedPassword = localConfig.GeneratedPassword
|
termout.PrintConfigCreated()
|
||||||
cfg.Auth.GeneratedPasswordPersisted = true
|
|
||||||
}
|
}
|
||||||
|
|
||||||
if *httpsBootstrap {
|
if *httpsBootstrap {
|
||||||
@@ -63,15 +63,12 @@ func main() {
|
|||||||
if config.MainWebUIUsesHTTPS(&cfg.Server) {
|
if config.MainWebUIUsesHTTPS(&cfg.Server) {
|
||||||
scheme = "https"
|
scheme = "https"
|
||||||
}
|
}
|
||||||
fmt.Println()
|
termout.PrintStartupWebUI(termout.StartupWebUIOptions{
|
||||||
fmt.Printf("→ Web 界面: %s://127.0.0.1:%d/\n", scheme, port)
|
Scheme: scheme,
|
||||||
if scheme == "https" && cfg.Server.TLSAutoSelfSign {
|
Port: port,
|
||||||
fmt.Println(" (内存自签证书:浏览器首次需确认「继续访问」)")
|
SelfSigned: scheme == "https" && cfg.Server.TLSAutoSelfSign,
|
||||||
}
|
HTTPRedirect: scheme == "https" && config.ServerHTTPRedirectEnabled(&cfg.Server),
|
||||||
if scheme == "https" && config.ServerHTTPRedirectEnabled(&cfg.Server) {
|
})
|
||||||
fmt.Printf(" (http://127.0.0.1:%d/ 将自动跳转到 HTTPS)\n", port)
|
|
||||||
}
|
|
||||||
fmt.Println()
|
|
||||||
|
|
||||||
// MCP 启用且 auth_header_value 为空时,自动生成随机密钥并写回配置
|
// MCP 启用且 auth_header_value 为空时,自动生成随机密钥并写回配置
|
||||||
if err := config.EnsureMCPAuth(cp, cfg); err != nil {
|
if err := config.EnsureMCPAuth(cp, cfg); err != nil {
|
||||||
|
|||||||
+5
-2
@@ -10,11 +10,14 @@
|
|||||||
# ============================================
|
# ============================================
|
||||||
|
|
||||||
# 前端显示的版本号(可选,不填则显示默认版本)
|
# 前端显示的版本号(可选,不填则显示默认版本)
|
||||||
version: "v1.7.0"
|
version: "v1.7.3"
|
||||||
# 服务器配置
|
# 服务器配置
|
||||||
server:
|
server:
|
||||||
host: 0.0.0.0 # 监听地址,0.0.0.0 表示监听所有网络接口
|
host: 0.0.0.0 # 监听地址,0.0.0.0 表示监听所有网络接口
|
||||||
port: 8080 # 服务端口;未启用 TLS 时为 http://localhost:8080
|
port: 8080 # 服务端口;未启用 TLS 时为 http://localhost:8080
|
||||||
|
# 其他可信 Web 集成的精确 Origin。Chromium 浏览器插件会自动识别,无需配置;不要使用通配符。
|
||||||
|
# cors_allowed_origins:
|
||||||
|
# - https://trusted-integration.example
|
||||||
# --- 可选:HTTPS + HTTP/2(缓解浏览器对同源 HTTP/1.1 的并发连接数限制,多路 Deep 流式更稳)---
|
# --- 可选:HTTPS + HTTP/2(缓解浏览器对同源 HTTP/1.1 的并发连接数限制,多路 Deep 流式更稳)---
|
||||||
# 启用 TLS 的条件(满足其一即可):tls_enabled: true,或 tls_auto_self_sign: true,或同时配置了 tls_cert_path + tls_key_path。
|
# 启用 TLS 的条件(满足其一即可):tls_enabled: true,或 tls_auto_self_sign: true,或同时配置了 tls_cert_path + tls_key_path。
|
||||||
# 启用后请用 https://127.0.0.1:<本端口>/ 访问;若仍用 http:// 访问同端口,将自动 308 跳转到 HTTPS(可用 tls_http_redirect: false 关闭)。
|
# 启用后请用 https://127.0.0.1:<本端口>/ 访问;若仍用 http:// 访问同端口,将自动 308 跳转到 HTTPS(可用 tls_http_redirect: false 关闭)。
|
||||||
@@ -28,7 +31,6 @@ server:
|
|||||||
tls_auto_self_sign: true
|
tls_auto_self_sign: true
|
||||||
# 认证配置
|
# 认证配置
|
||||||
auth:
|
auth:
|
||||||
password: # Web 登录密码,请修改为强密码
|
|
||||||
session_duration_hours: 12 # 登录有效期(小时),超时后需重新登录
|
session_duration_hours: 12 # 登录有效期(小时),超时后需重新登录
|
||||||
# 日志配置
|
# 日志配置
|
||||||
log:
|
log:
|
||||||
@@ -210,6 +212,7 @@ multi_agent:
|
|||||||
reduction_clear_exclude: [] # 不参与「清理阶段」的工具名额外列表(会与 task/transfer/exit 等内置排除项合并);需要时用 YAML 列表填写
|
reduction_clear_exclude: [] # 不参与「清理阶段」的工具名额外列表(会与 task/transfer/exit 等内置排除项合并);需要时用 YAML 列表填写
|
||||||
reduction_sub_agents: true # true:子代理也挂 reduction;false:仅编排主代理使用 reduction
|
reduction_sub_agents: true # true:子代理也挂 reduction;false:仅编排主代理使用 reduction
|
||||||
summarization_trigger_ratio: 0.8 # summarization 触发比例(max_total_tokens * ratio),建议 0.75~0.85
|
summarization_trigger_ratio: 0.8 # summarization 触发比例(max_total_tokens * ratio),建议 0.75~0.85
|
||||||
|
summarization_output_reserve_tokens: 8192 # 摘要模型输出预留 token;摘要输入预算 = 触发阈值 - 该值
|
||||||
summarization_emit_internal_events: true # true:发出 summarization 内部事件(便于诊断)
|
summarization_emit_internal_events: true # true:发出 summarization 内部事件(便于诊断)
|
||||||
summarization_user_intent_ledger_max_runes: 96000 # 压缩后注入模型上下文的「原始用户输入与约束账本」总字符上限;DB 原始消息不裁剪
|
summarization_user_intent_ledger_max_runes: 96000 # 压缩后注入模型上下文的「原始用户输入与约束账本」总字符上限;DB 原始消息不裁剪
|
||||||
summarization_user_intent_ledger_entry_max_runes: 16000 # 账本中单条用户消息的字符上限;超出仅裁剪模型可见账本,不影响 DB 原文
|
summarization_user_intent_ledger_entry_max_runes: 16000 # 账本中单条用户消息的字符上限;超出仅裁剪模型可见账本,不影响 DB 原文
|
||||||
|
|||||||
+68
-51
@@ -1,68 +1,85 @@
|
|||||||
# CyberStrikeAI Documentation
|
# CyberStrikeAI Documentation
|
||||||
|
|
||||||
Documentation is split by language:
|
[中文](#中文文档) | [English](#english-documentation)
|
||||||
|
|
||||||
- [中文文档](zh-CN/)
|
CyberStrikeAI documentation is organized by user journey. Start with deployment, then move to the topic that matches your task.
|
||||||
- [English docs](en-US/)
|
|
||||||
|
|
||||||
## 中文文档
|
## 中文文档
|
||||||
|
|
||||||
- [部署指南](zh-CN/deployment.md)
|
### 按目标开始
|
||||||
- [运维 Runbooks](zh-CN/runbooks.md)
|
|
||||||
- [配置画像](zh-CN/configuration-profiles.md)
|
- **快速体验**:[部署指南](zh-CN/deployment.md) → [配置参考](zh-CN/configuration.md) → [排错指南](zh-CN/troubleshooting.md)
|
||||||
- [安全加固指南](zh-CN/security-hardening.md)
|
- **生产部署**:[配置画像](zh-CN/configuration-profiles.md) → [安全加固](zh-CN/security-hardening.md) → [运维 Runbooks](zh-CN/runbooks.md) → [审计与监控](zh-CN/audit-and-monitoring.md)
|
||||||
- [API Recipes](zh-CN/api-recipes.md)
|
- **接入与自动化**:[API 参考](zh-CN/api-reference.md) → [API Recipes](zh-CN/api-recipes.md) → [MCP 联邦](zh-CN/mcp-federation.md)
|
||||||
- [贡献规范](zh-CN/contributing-guide.md)
|
- **参与开发**:[开发者指南](zh-CN/developer-guide.md) → [测试指南](zh-CN/testing.md) → [贡献规范](zh-CN/contributing-guide.md)
|
||||||
- [配置参考](zh-CN/configuration.md)
|
|
||||||
- [安全模型](zh-CN/security-model.md)
|
### 核心概念与编排
|
||||||
- [RBAC 权限管理](zh-CN/rbac.md)
|
|
||||||
- [架构说明](zh-CN/architecture.md)
|
- [架构说明](zh-CN/architecture.md)
|
||||||
- [API 参考](zh-CN/api-reference.md)
|
- [安全模型](zh-CN/security-model.md)
|
||||||
- [排错指南](zh-CN/troubleshooting.md)
|
|
||||||
- [审计与监控](zh-CN/audit-and-monitoring.md)
|
|
||||||
- [知识库](zh-CN/knowledge-base.md)
|
|
||||||
- [C2 使用说明](zh-CN/c2.md)
|
|
||||||
- [WebShell 管理](zh-CN/webshell.md)
|
|
||||||
- [MCP 联邦](zh-CN/mcp-federation.md)
|
|
||||||
- [Agent 与角色](zh-CN/agent-and-role-guide.md)
|
- [Agent 与角色](zh-CN/agent-and-role-guide.md)
|
||||||
- [Skills 指南](zh-CN/skills-guide.md)
|
- [Skills 指南](zh-CN/skills-guide.md)
|
||||||
- [插件开发](zh-CN/plugin-development.md)
|
- [Eino 多代理](zh-CN/MULTI_AGENT_EINO.md)
|
||||||
- [发布流程](zh-CN/release-process.md)
|
- [图编排](zh-CN/workflow-graph.md)
|
||||||
- [测试指南](zh-CN/testing.md)
|
|
||||||
- [图编排使用说明](zh-CN/workflow-graph.md)
|
|
||||||
- [人机协同最佳实践](zh-CN/hitl-best-practices.md)
|
- [人机协同最佳实践](zh-CN/hitl-best-practices.md)
|
||||||
- [机器人使用说明](zh-CN/robot.md)
|
|
||||||
|
### 功能指南
|
||||||
|
|
||||||
|
- [知识库](zh-CN/knowledge-base.md)
|
||||||
|
- [RBAC 权限管理](zh-CN/rbac.md)
|
||||||
|
- [机器人接入](zh-CN/robot.md)
|
||||||
- [视觉分析](zh-CN/VISION.md)
|
- [视觉分析](zh-CN/VISION.md)
|
||||||
- [前端国际化方案](zh-CN/frontend-i18n.md)
|
- [WebShell 管理](zh-CN/webshell.md)
|
||||||
- [Eino 多代理改造说明](zh-CN/MULTI_AGENT_EINO.md)
|
- [C2 使用说明](zh-CN/c2.md)
|
||||||
|
|
||||||
## English Docs
|
### 开发与发布
|
||||||
|
|
||||||
|
- [开发者指南](zh-CN/developer-guide.md)
|
||||||
|
- [插件开发](zh-CN/plugin-development.md)
|
||||||
|
- [前端国际化](zh-CN/frontend-i18n.md)
|
||||||
|
- [测试指南](zh-CN/testing.md)
|
||||||
|
- [贡献规范](zh-CN/contributing-guide.md)
|
||||||
|
- [发布流程](zh-CN/release-process.md)
|
||||||
|
|
||||||
|
## English Documentation
|
||||||
|
|
||||||
|
### Choose a path
|
||||||
|
|
||||||
|
- **Try locally**: [Deployment](en-US/deployment.md) → [Configuration](en-US/configuration.md) → [Troubleshooting](en-US/troubleshooting.md)
|
||||||
|
- **Run in production**: [Configuration Profiles](en-US/configuration-profiles.md) → [Security Hardening](en-US/security-hardening.md) → [Runbooks](en-US/runbooks.md) → [Audit and Monitoring](en-US/audit-and-monitoring.md)
|
||||||
|
- **Integrate and automate**: [API Reference](en-US/api-reference.md) → [API Recipes](en-US/api-recipes.md) → [MCP Federation](en-US/mcp-federation.md)
|
||||||
|
- **Contribute code**: [Developer Guide](en-US/developer-guide.md) → [Testing](en-US/testing.md) → [Contributing](en-US/contributing-guide.md)
|
||||||
|
|
||||||
|
### Concepts and orchestration
|
||||||
|
|
||||||
- [Deployment Guide](en-US/deployment.md)
|
|
||||||
- [Runbooks](en-US/runbooks.md)
|
|
||||||
- [Configuration Profiles](en-US/configuration-profiles.md)
|
|
||||||
- [Security Hardening](en-US/security-hardening.md)
|
|
||||||
- [API Recipes](en-US/api-recipes.md)
|
|
||||||
- [Contributing Guide](en-US/contributing-guide.md)
|
|
||||||
- [Configuration Reference](en-US/configuration.md)
|
|
||||||
- [Security Model](en-US/security-model.md)
|
|
||||||
- [RBAC Administration](en-US/rbac.md)
|
|
||||||
- [Architecture](en-US/architecture.md)
|
- [Architecture](en-US/architecture.md)
|
||||||
- [API Reference](en-US/api-reference.md)
|
- [Security Model](en-US/security-model.md)
|
||||||
- [Troubleshooting](en-US/troubleshooting.md)
|
- [Agents and Roles](en-US/agent-and-role-guide.md)
|
||||||
- [Audit and Monitoring](en-US/audit-and-monitoring.md)
|
- [Skills](en-US/skills-guide.md)
|
||||||
- [Knowledge Base](en-US/knowledge-base.md)
|
- [Eino Multi-Agent](en-US/MULTI_AGENT_EINO.md)
|
||||||
- [C2 Guide](en-US/c2.md)
|
- [Graph Orchestration](en-US/workflow-graph.md)
|
||||||
- [WebShell Management](en-US/webshell.md)
|
|
||||||
- [MCP Federation](en-US/mcp-federation.md)
|
|
||||||
- [Agent and Role Guide](en-US/agent-and-role-guide.md)
|
|
||||||
- [Skills Guide](en-US/skills-guide.md)
|
|
||||||
- [Plugin Development](en-US/plugin-development.md)
|
|
||||||
- [Release Process](en-US/release-process.md)
|
|
||||||
- [Testing Guide](en-US/testing.md)
|
|
||||||
- [Graph Orchestration Guide](en-US/workflow-graph.md)
|
|
||||||
- [HITL Best Practices](en-US/hitl-best-practices.md)
|
- [HITL Best Practices](en-US/hitl-best-practices.md)
|
||||||
- [Robot / Chatbot Guide](en-US/robot.md)
|
|
||||||
|
### Feature guides
|
||||||
|
|
||||||
|
- [Knowledge Base](en-US/knowledge-base.md)
|
||||||
|
- [RBAC Administration](en-US/rbac.md)
|
||||||
|
- [Robot / Chatbot](en-US/robot.md)
|
||||||
- [Vision Analysis](en-US/VISION.md)
|
- [Vision Analysis](en-US/VISION.md)
|
||||||
|
- [WebShell Management](en-US/webshell.md)
|
||||||
|
- [C2 Guide](en-US/c2.md)
|
||||||
|
|
||||||
|
### Development and release
|
||||||
|
|
||||||
|
- [Developer Guide](en-US/developer-guide.md)
|
||||||
|
- [Plugin Development](en-US/plugin-development.md)
|
||||||
- [Frontend i18n](en-US/frontend-i18n.md)
|
- [Frontend i18n](en-US/frontend-i18n.md)
|
||||||
- [Eino Multi-Agent Notes](en-US/MULTI_AGENT_EINO.md)
|
- [Testing](en-US/testing.md)
|
||||||
|
- [Contributing](en-US/contributing-guide.md)
|
||||||
|
- [Release Process](en-US/release-process.md)
|
||||||
|
|
||||||
|
## Documentation conventions
|
||||||
|
|
||||||
|
- Commands assume the repository root unless stated otherwise.
|
||||||
|
- Examples use placeholders; never commit real credentials or target systems without explicit authorization.
|
||||||
|
- Runtime behavior and configuration defaults are authoritative in `config.example.yaml` and the source code. If a document differs, report it as documentation drift.
|
||||||
|
|||||||
+31
-29
@@ -1,30 +1,32 @@
|
|||||||
# English Docs
|
# English Documentation
|
||||||
|
|
||||||
- [Deployment Guide](deployment.md): deployment modes, HTTPS, reverse proxy, systemd, backup, upgrade, and acceptance checks.
|
[Documentation home](../README.md) | [中文](../zh-CN/README.md)
|
||||||
- [Runbooks](runbooks.md): operational steps for production setup, external MCP, KB, Web testing, C2 cleanup, and tool debugging.
|
|
||||||
- [Configuration Profiles](configuration-profiles.md): recommended profiles for dev, internal team, knowledge-only, production, C2, and MCP automation.
|
## Choose a path
|
||||||
- [Security Hardening](security-hardening.md): pre-launch baseline, reverse proxy, HITL allowlist, file permissions, and periodic review.
|
|
||||||
- [API Recipes](api-recipes.md): examples for login, Agent, streaming, multi-agent, uploads, vulnerabilities, KB, MCP, and audit export.
|
- **Try locally**: [Deployment](deployment.md) → [Configuration](configuration.md) → [Troubleshooting](troubleshooting.md)
|
||||||
- [Contributing Guide](contributing-guide.md): checklists for APIs, config, tools, frontend, DB, high-risk features, and docs.
|
- **Run in production**: [Configuration Profiles](configuration-profiles.md) → [Security Hardening](security-hardening.md) → [Runbooks](runbooks.md) → [Audit and Monitoring](audit-and-monitoring.md)
|
||||||
- [Configuration Reference](configuration.md): `config.yaml` fields, hot-apply boundaries, recommended values, and source anchors.
|
- **Integrate and automate**: [API Reference](api-reference.md) → [API Recipes](api-recipes.md) → [MCP Federation](mcp-federation.md)
|
||||||
- [Security Model](security-model.md): trust boundaries, HITL, tool execution, C2/WebShell, and data safety.
|
- **Contribute code**: [Developer Guide](developer-guide.md) → [Testing](testing.md) → [Contributing](contributing-guide.md)
|
||||||
- [RBAC Administration](rbac.md): platform users, system/custom roles, permission catalog, per-permission scopes, resource assignments, Agent/MCP/robot boundaries, and API examples.
|
|
||||||
- [Architecture](architecture.md): request flow, module relationships, complexity hotspots, and design trade-offs.
|
## Concepts and orchestration
|
||||||
- [API Reference](api-reference.md): authentication, OpenAPI, SSE, stability tiers, and common endpoints.
|
|
||||||
- [Troubleshooting](troubleshooting.md): diagnostic order, minimal commands, common misdiagnoses, and issue template.
|
- [Architecture](architecture.md) · [Security Model](security-model.md) · [RBAC](rbac.md)
|
||||||
- [Audit and Monitoring](audit-and-monitoring.md): platform audit, tool monitoring, HITL logs, and retention.
|
- [Agents and Roles](agent-and-role-guide.md) · [Skills](skills-guide.md) · [Eino Multi-Agent](MULTI_AGENT_EINO.md)
|
||||||
- [Knowledge Base](knowledge-base.md): indexing pipeline, retrieval tuning, log analysis, and content writing.
|
- [Graph Orchestration](workflow-graph.md) · [HITL Best Practices](hitl-best-practices.md)
|
||||||
- [C2 Guide](c2.md): lifecycle, task classification, event review, and safety guidance.
|
|
||||||
- [WebShell Management](webshell.md): operation tiers, naming, AI guardrails, and troubleshooting.
|
## Feature guides
|
||||||
- [MCP Federation](mcp-federation.md): built-in MCP, external MCP, lifecycle, and tool naming.
|
|
||||||
- [Agent and Role Guide](agent-and-role-guide.md): roles, sub-agents, Skills, orchestration modes, and tool visibility.
|
- [Knowledge Base](knowledge-base.md) · [Robot / Chatbot](robot.md) · [Vision](VISION.md)
|
||||||
- [Skills Guide](skills-guide.md): Skill structure, progressive disclosure, anti-patterns, and local-tool risk.
|
- [WebShell](webshell.md) · [C2](c2.md) · [MCP Federation](mcp-federation.md)
|
||||||
- [Plugin Development](plugin-development.md): API plugins, MCP plugins, resource-pack plugins, and security boundaries.
|
|
||||||
- [Release Process](release-process.md): release risk, config compatibility, DB migrations, and acceptance checks.
|
## Operations and reference
|
||||||
- [Testing Guide](testing.md): test layers, regression focus, test data, and failure cases.
|
|
||||||
- [Graph Orchestration Guide](workflow-graph.md)
|
- [Deployment](deployment.md) · [Configuration](configuration.md) · [Configuration Profiles](configuration-profiles.md)
|
||||||
- [HITL Best Practices](hitl-best-practices.md)
|
- [Security Hardening](security-hardening.md) · [Audit and Monitoring](audit-and-monitoring.md) · [Runbooks](runbooks.md)
|
||||||
- [Robot / Chatbot Guide](robot.md)
|
- [API Reference](api-reference.md) · [API Recipes](api-recipes.md) · [Troubleshooting](troubleshooting.md)
|
||||||
- [Vision Analysis](VISION.md)
|
|
||||||
- [Frontend i18n](frontend-i18n.md)
|
## Development and release
|
||||||
- [Eino Multi-Agent Notes](MULTI_AGENT_EINO.md)
|
|
||||||
|
- [Developer Guide](developer-guide.md) · [Plugin Development](plugin-development.md) · [Frontend i18n](frontend-i18n.md)
|
||||||
|
- [Testing](testing.md) · [Contributing](contributing-guide.md) · [Release Process](release-process.md)
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ server:
|
|||||||
tls_enabled: true
|
tls_enabled: true
|
||||||
tls_auto_self_sign: true
|
tls_auto_self_sign: true
|
||||||
auth:
|
auth:
|
||||||
password: "dev-only-change-me"
|
session_duration_hours: 12
|
||||||
audit:
|
audit:
|
||||||
enabled: true
|
enabled: true
|
||||||
retention_days: 7
|
retention_days: 7
|
||||||
@@ -45,7 +45,7 @@ server:
|
|||||||
port: 8080
|
port: 8080
|
||||||
tls_enabled: false
|
tls_enabled: false
|
||||||
auth:
|
auth:
|
||||||
password: "<long-random-password>"
|
session_duration_hours: 12
|
||||||
audit:
|
audit:
|
||||||
enabled: true
|
enabled: true
|
||||||
retention_days: 30
|
retention_days: 30
|
||||||
@@ -90,7 +90,6 @@ Goal: long-running production red-team or security platform.
|
|||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
auth:
|
auth:
|
||||||
password: "<managed-secret>"
|
|
||||||
session_duration_hours: 8
|
session_duration_hours: 8
|
||||||
audit:
|
audit:
|
||||||
enabled: true
|
enabled: true
|
||||||
|
|||||||
@@ -11,8 +11,10 @@ server:
|
|||||||
host: 0.0.0.0
|
host: 0.0.0.0
|
||||||
port: 8080
|
port: 8080
|
||||||
tls_enabled: true
|
tls_enabled: true
|
||||||
|
# Optional: other trusted Web integrations; Chromium extensions need no entry.
|
||||||
|
# cors_allowed_origins:
|
||||||
|
# - https://trusted-integration.example
|
||||||
auth:
|
auth:
|
||||||
password: "change-me"
|
|
||||||
session_duration_hours: 12
|
session_duration_hours: 12
|
||||||
openai:
|
openai:
|
||||||
provider: openai
|
provider: openai
|
||||||
@@ -24,7 +26,9 @@ agent:
|
|||||||
tool_timeout_minutes: 60
|
tool_timeout_minutes: 60
|
||||||
```
|
```
|
||||||
|
|
||||||
Change the default password immediately. Use HTTPS or a trusted reverse proxy in any shared environment.
|
Change the initial `admin` password from the Web UI after first login. Use HTTPS or a trusted reverse proxy in any shared environment.
|
||||||
|
|
||||||
|
Valid Chromium `chrome-extension://<32-character-extension-id>` origins are recognized automatically. The extension must still obtain host permission and authenticate with a password and Bearer token. `server.cors_allowed_origins` remains available as an exact allowlist for other trusted Web integrations; wildcards are not accepted, and changing it requires a restart.
|
||||||
|
|
||||||
## Hot-Apply Boundaries
|
## Hot-Apply Boundaries
|
||||||
|
|
||||||
|
|||||||
@@ -103,6 +103,16 @@ Update:
|
|||||||
- `docs/zh-CN/README.md`
|
- `docs/zh-CN/README.md`
|
||||||
- `docs/en-US/README.md`
|
- `docs/en-US/README.md`
|
||||||
|
|
||||||
|
Before submitting documentation changes, run:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python3 scripts/check-docs.py
|
||||||
|
```
|
||||||
|
|
||||||
|
The check verifies local links, fenced code blocks, bilingual filename parity, locale index coverage, and the Go version documented by the root READMEs. Keep versioned examples derived from authoritative files such as `go.mod` and `config.example.yaml` whenever possible.
|
||||||
|
|
||||||
|
The same command runs automatically in the `Documentation` GitHub Actions workflow when documentation, `go.mod`, or the checker itself changes.
|
||||||
|
|
||||||
## Review Focus
|
## Review Focus
|
||||||
|
|
||||||
Prioritize:
|
Prioritize:
|
||||||
|
|||||||
@@ -48,7 +48,7 @@ config.yaml
|
|||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
auth:
|
auth:
|
||||||
password: "<long-random-password>"
|
session_duration_hours: 12
|
||||||
server:
|
server:
|
||||||
host: 127.0.0.1
|
host: 127.0.0.1
|
||||||
port: 8080
|
port: 8080
|
||||||
|
|||||||
@@ -6,7 +6,7 @@ This checklist covers pre-production and continuous hardening for CyberStrikeAI.
|
|||||||
|
|
||||||
## Before Going Live
|
## Before Going Live
|
||||||
|
|
||||||
- Change `auth.password` to a long random secret.
|
- Change the initial `admin` password from the Web UI after first login.
|
||||||
- Use HTTPS or a trusted reverse proxy.
|
- Use HTTPS or a trusted reverse proxy.
|
||||||
- Restrict access by IP, VPN, or bastion.
|
- Restrict access by IP, VPN, or bastion.
|
||||||
- Enable `audit.enabled`.
|
- Enable `audit.enabled`.
|
||||||
|
|||||||
@@ -37,7 +37,7 @@ Page inaccessible:
|
|||||||
|
|
||||||
Login fails:
|
Login fails:
|
||||||
|
|
||||||
- wrong `auth.password`;
|
- wrong RBAC user password;
|
||||||
- config not applied/restarted;
|
- config not applied/restarted;
|
||||||
- stale cookie;
|
- stale cookie;
|
||||||
- audit throttling repeated failures.
|
- audit throttling repeated failures.
|
||||||
|
|||||||
@@ -0,0 +1,99 @@
|
|||||||
|
# Local Workflow Package MVP Implementation Plan
|
||||||
|
|
||||||
|
> **For agentic workers:** REQUIRED SUB-SKILL: Use superpowers:subagent-driven-development (recommended) or superpowers:executing-plans to implement this plan task-by-task. Steps use checkbox (`- [ ]`) syntax for tracking.
|
||||||
|
|
||||||
|
**Goal:** Add secure, deterministic single-workflow ZIP export plus two-step, idempotent local package import without changing existing workflow APIs.
|
||||||
|
|
||||||
|
**Architecture:** `internal/workflow/package` owns package format, deterministic ZIP construction, archive inspection and import orchestration. `internal/database` owns SQLite schema, lifecycle state and the one transaction that rechecks the inspection snapshot, changes `workflow_definitions`, persists the import result and consumes the inspection. `internal/handler` maps the approved REST contract to these services and app routing/RBAC remains the enforcement boundary.
|
||||||
|
|
||||||
|
**Tech Stack:** Go, Gin, SQLite via `github.com/mattn/go-sqlite3`, archive/zip, SHA-256, existing Eino `ValidateGraphJSON`.
|
||||||
|
|
||||||
|
## Global Constraints
|
||||||
|
|
||||||
|
- Backend only: do not modify `web/templates`, `web/static`, i18n, or any other frontend file.
|
||||||
|
- Support exactly one `workflows/*.json` item; never execute package contents.
|
||||||
|
- Request ZIP maximum is 10 MiB and extracted total maximum is 20 MiB.
|
||||||
|
- Keep `workflow_definitions.version` local: create/rename is 1 and overwrite is the existing local version plus one.
|
||||||
|
- Use `workflow:read` only for export, `workflow:write` for every inspection/import route, and require existing RBAC `all` scope for package mutations.
|
||||||
|
- Preserve existing CRUD, validate, dry-run and run API response formats.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
### Task 1: Package format, canonical hashes and deterministic export
|
||||||
|
|
||||||
|
**Files:**
|
||||||
|
- Create: `internal/workflow/package/manifest.go`
|
||||||
|
- Create: `internal/workflow/package/exporter.go`
|
||||||
|
- Test: `internal/workflow/package/exporter_test.go`
|
||||||
|
|
||||||
|
**Interfaces:**
|
||||||
|
- Produces: `Export(database.WorkflowDefinition) ([]byte, ExportMetadata, error)`, `InspectArchive(context.Context, []byte) (*InspectionResult, error)`, and typed package errors exposing `Code`, safe `Message`, and safe `Details`.
|
||||||
|
- Consumes: `database.WorkflowDefinition` and the existing graph JSON fields only.
|
||||||
|
|
||||||
|
- [ ] **Step 1: Write failing package tests.** Cover two identical exports producing byte-identical ZIPs, lower-case `sha256:` hashes, `manifest.json`/`checksums.sha256`/one workflow entry, and source revision equal to the source workflow's local version.
|
||||||
|
- [ ] **Step 2: Run the package test.** Run `go test ./internal/workflow/package -run 'TestExport'`; expected failure is missing package export symbols.
|
||||||
|
- [ ] **Step 3: Implement canonical JSON and exporter.** Canonicalize JSON with `Decoder.UseNumber`, hash canonical graph JSON and canonical item JSON, derive a stable package id and fixed ZIP metadata, then write entries in lexical order.
|
||||||
|
- [ ] **Step 4: Run the package test.** Run `go test ./internal/workflow/package -run 'TestExport'`; expected result is PASS.
|
||||||
|
|
||||||
|
### Task 2: Safe package inspection and validation
|
||||||
|
|
||||||
|
**Files:**
|
||||||
|
- Create: `internal/workflow/package/inspector.go`
|
||||||
|
- Test: `internal/workflow/package/inspector_test.go`
|
||||||
|
|
||||||
|
**Interfaces:**
|
||||||
|
- Consumes: ZIP bytes and `workflow.ValidateGraphJSON(context.Context, string)`.
|
||||||
|
- Produces: validated manifest/workflow payload, package/content/graph hashes, node/edge counts, and contract error codes without archive paths.
|
||||||
|
|
||||||
|
- [ ] **Step 1: Write failing inspector tests.** Use an exported valid package and assert accepted parsing; add independent cases for path traversal, duplicate names, symlink entries, undeclared files, checksum mismatch, two workflow entries, extracted-size overflow, invalid manifest and invalid graph.
|
||||||
|
- [ ] **Step 2: Run the inspector test.** Run `go test ./internal/workflow/package -run 'TestInspect'`; expected failure is missing inspection implementation.
|
||||||
|
- [ ] **Step 3: Implement archive checks before parsing.** Reject non-exact paths, duplicate names, links, unexpected entries and declared/actual oversized extraction; validate checksums and Manifest 1.0; require exactly one declared workflow entry; then reuse `ValidateGraphJSON`.
|
||||||
|
- [ ] **Step 4: Run the inspector test.** Run `go test ./internal/workflow/package -run 'TestInspect'`; expected result is PASS.
|
||||||
|
|
||||||
|
### Task 3: SQLite package state, lifecycle, and transactional application
|
||||||
|
|
||||||
|
**Files:**
|
||||||
|
- Modify: `internal/database/database.go`
|
||||||
|
- Create: `internal/database/workflow_package.go`
|
||||||
|
- Test: `internal/database/workflow_package_test.go`
|
||||||
|
|
||||||
|
**Interfaces:**
|
||||||
|
- Produces: inspection create/read/expiry methods, `ApplyWorkflowPackageImport` and `PurgeWorkflowPackageLifecycle(time.Time)`.
|
||||||
|
- Consumes: primitive database request structs carrying inspected payload and immutable conflict snapshot; no browser-supplied workflow JSON.
|
||||||
|
|
||||||
|
- [ ] **Step 1: Write failing DB tests.** Assert migration tables/indexes exist; inspection expiry transitions to `expired`; first successful application consumes inspection; same actor/key/same hash returns stored import; same key/different hash rejects; changed target snapshot rejects; overwrite increments local version; create and rename start at version 1; rollback leaves workflow/import/inspection unchanged on failure.
|
||||||
|
- [ ] **Step 2: Run the DB test.** Run `go test ./internal/database -run 'TestWorkflowPackage'`; expected failure is missing migration and methods.
|
||||||
|
- [ ] **Step 3: Add exact DDL and transactional repository method.** Add the two contract tables and indexes to `initTables`; in one `BEGIN` transaction recheck owner/status/expiry/idempotency/snapshot, apply the allowed action, insert import row, mark inspection consumed, and commit. Add 24-hour expired-inspection and 90-day import cleanup.
|
||||||
|
- [ ] **Step 4: Run the DB test.** Run `go test ./internal/database -run 'TestWorkflowPackage'`; expected result is PASS.
|
||||||
|
|
||||||
|
### Task 4: Import orchestration, HTTP handlers, audit and routes
|
||||||
|
|
||||||
|
**Files:**
|
||||||
|
- Create: `internal/workflow/package/importer.go`
|
||||||
|
- Create: `internal/handler/workflow_package.go`
|
||||||
|
- Modify: `internal/handler/workflow.go`
|
||||||
|
- Modify: `internal/app/app.go`
|
||||||
|
- Modify: `internal/security/rbac_middleware.go`
|
||||||
|
- Test: `internal/handler/workflow_package_test.go`
|
||||||
|
|
||||||
|
**Interfaces:**
|
||||||
|
- Consumes: authenticated `security.Session`, `Idempotency-Key`, multipart `file`, database package state, and typed package errors.
|
||||||
|
- Produces: contract response envelopes, `application/zip` export headers, cache invalidation after committed writes, and audit events in category `workflow_package`.
|
||||||
|
|
||||||
|
- [ ] **Step 1: Write failing handler/RBAC tests.** Cover 403 mapping for read/write permissions, 10 MiB file limit, export headers/404, creator-only inspection/import reads, 201 first apply/200 idempotent replay, contract error body/status, and the existing validate/dry-run/runs routes still resolving.
|
||||||
|
- [ ] **Step 2: Run the handler test.** Run `go test ./internal/handler -run 'TestWorkflowPackage'`; expected failure is missing routes/handlers.
|
||||||
|
- [ ] **Step 3: Implement service and handlers.** Limit upload bytes before multipart parsing; persist only validated payload; perform request-hash and action validation in importer; map package errors to approved statuses; invalidate the compiled cache only after commit; record export/inspect/import success and failure audits.
|
||||||
|
- [ ] **Step 4: Register and authorize routes.** Register exact paths `GET /workflows/:id/package`, `POST|GET /workflow-package-inspections`, and `POST|GET /workflow-package-imports`; make the route mapper explicit and treat inspection/import POSTs as process-global workflow mutations.
|
||||||
|
- [ ] **Step 5: Run focused handler tests.** Run `go test ./internal/handler -run 'TestWorkflowPackage'`; expected result is PASS.
|
||||||
|
|
||||||
|
### Task 5: Lifecycle wiring and final verification
|
||||||
|
|
||||||
|
**Files:**
|
||||||
|
- Modify: `internal/app/app.go`
|
||||||
|
- Test: the tests from Tasks 1-4
|
||||||
|
|
||||||
|
- [ ] **Step 1: Write the failing lifecycle wiring test or startup-level assertion.** Assert startup invokes package lifecycle cleanup and that the retention loop has no workflow-definition side effect.
|
||||||
|
- [ ] **Step 2: Implement startup cleanup/loop.** Invoke `PurgeWorkflowPackageLifecycle(time.Now().UTC())` at startup and start an hourly package lifecycle loop after the database is ready.
|
||||||
|
- [ ] **Step 3: Run format and focused verification.** Run `gofmt -w` only on changed Go files, `go test ./internal/workflow/package`, `go test ./internal/database -run 'TestWorkflowPackage'`, `go test ./internal/handler -run 'TestWorkflowPackage'`, and `git diff --check`.
|
||||||
|
- [ ] **Step 4: Run compatible regression verification.** Run `go test ./internal/database ./internal/handler ./internal/workflow` in an environment with the required C compiler, then inspect `git diff --check` and `git status --short` before committing.
|
||||||
|
- [ ] **Step 5: Commit verified files.** Run `git add internal/workflow/package internal/database/database.go internal/database/workflow_package.go internal/handler/workflow.go internal/handler/workflow_package.go internal/security/rbac_middleware.go internal/app/app.go docs/superpowers/plans/2026-07-13-local-workflow-package-mvp.md` followed by `git commit -m "feat: add local workflow package mvp"`.
|
||||||
@@ -0,0 +1,334 @@
|
|||||||
|
# 本地图编排策略包 MVP:API 与数据模型契约 v1
|
||||||
|
|
||||||
|
> 本文是 [本地图编排策略包 MVP 设计](2026-07-13-local-workflow-package-mvp-design.md) 的实现前契约。前端与后端以本文的路径、字段、枚举、状态码和错误码为准;未经版本升级不得改变既有字段语义。
|
||||||
|
|
||||||
|
## 1. 范围与不变式
|
||||||
|
|
||||||
|
- 仅支持一个工作流的本地 `.csapkg.zip` 包。
|
||||||
|
- 仅处理 `workflow_definitions`;不导入 Role、Skill、MCP 配置、运行记录或任何可执行文件。
|
||||||
|
- 导入固定为“上传预检”和“确认应用”两步。预检不修改 `workflow_definitions`。
|
||||||
|
- 现有工作流 CRUD、`/validate`、`/dry-run`、运行 API 和 `workflow_definitions` 表结构保持兼容。
|
||||||
|
- 已有 `workflow_definitions.version` 始终是目标实例本地修订号:新建导入从 `1` 开始;覆盖导入由现有本地版本递增;包内 `source_revision` 只用于展示和审计。
|
||||||
|
|
||||||
|
## 2. 统一约定
|
||||||
|
|
||||||
|
### 2.1 认证与权限
|
||||||
|
|
||||||
|
所有 API 均位于现有受保护的 `/api` 路由组。
|
||||||
|
|
||||||
|
| 接口 | 所需权限 |
|
||||||
|
|---|---|
|
||||||
|
| 导出包 | `workflow:read` |
|
||||||
|
| 创建或读取预检 | `workflow:write` |
|
||||||
|
| 应用或读取导入结果 | `workflow:write` |
|
||||||
|
|
||||||
|
预检会保存短期、已验证的工作流载荷,故不把它降级为只读权限。`created_by` / `actor_user_id` 取当前已认证会话的 `UserID`。
|
||||||
|
|
||||||
|
RBAC 路由映射必须显式新增:`GET /workflows/:id/package` 映射 `workflow:read`;`/workflow-package-inspections` 与 `/workflow-package-imports` 的所有 MVP 路由映射 `workflow:write`。它们与既有工作流定义同属全局资产,写操作仅允许现有 RBAC 的 `all` 资源范围。
|
||||||
|
|
||||||
|
### 2.2 错误响应
|
||||||
|
|
||||||
|
新接口统一使用如下错误响应;不改变旧工作流 API 的 `{"error":"..."}` 兼容格式。
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"error": {
|
||||||
|
"code": "WFPKG_ID_CONFLICT",
|
||||||
|
"message": "目标实例已存在同 ID 工作流,请选择处理策略",
|
||||||
|
"details": {
|
||||||
|
"workflow_id": "web-src-hunting"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
`details` 仅包含可安全展示的结构化信息,不返回 Zip 路径、服务端文件路径、Token 或内部堆栈。
|
||||||
|
|
||||||
|
### 2.3 时间与哈希
|
||||||
|
|
||||||
|
- 所有时间字段使用 RFC 3339 UTC 字符串。
|
||||||
|
- 哈希固定为小写十六进制 `sha256:<64-hex>`。
|
||||||
|
- `graph_hash` 是对规范化 `graph_json` 的 SHA-256;`content_hash` 是工作流包项的 SHA-256。
|
||||||
|
|
||||||
|
## 3. REST API
|
||||||
|
|
||||||
|
### 3.1 导出单工作流包
|
||||||
|
|
||||||
|
```http
|
||||||
|
GET /api/workflows/{id}/package
|
||||||
|
Accept: application/zip
|
||||||
|
```
|
||||||
|
|
||||||
|
语义:从当前 `workflow_definitions` 行生成 `.csapkg.zip`。只读、幂等、无确认、无数据库写入。
|
||||||
|
|
||||||
|
成功响应:
|
||||||
|
|
||||||
|
```http
|
||||||
|
200 OK
|
||||||
|
Content-Type: application/zip
|
||||||
|
Content-Disposition: attachment; filename="web-src-hunting.csapkg.zip"
|
||||||
|
ETag: "sha256:3f..."
|
||||||
|
X-Workflow-Package-SHA256: sha256:3f...
|
||||||
|
```
|
||||||
|
|
||||||
|
失败:`404 WFPKG_WORKFLOW_NOT_FOUND`、`403`、`500 WFPKG_EXPORT_FAILED`。
|
||||||
|
|
||||||
|
### 3.2 创建预检
|
||||||
|
|
||||||
|
```http
|
||||||
|
POST /api/workflow-package-inspections
|
||||||
|
Content-Type: multipart/form-data
|
||||||
|
|
||||||
|
file=@web-src-hunting.csapkg.zip;type=application/zip
|
||||||
|
```
|
||||||
|
|
||||||
|
限制:请求体最大 10 MiB;Zip 解压总量最大 20 MiB;只允许 `manifest.json`、`checksums.sha256` 和一个 `workflows/*.json`。同一包不得有重复条目、软链接、路径穿越或未声明文件。
|
||||||
|
|
||||||
|
成功响应:`201 Created`。
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"inspection": {
|
||||||
|
"id": "wpi_01JQ2K6G7K8W2C1E3R4T5Y6U7I",
|
||||||
|
"status": "ready",
|
||||||
|
"expires_at": "2026-07-13T09:30:00Z",
|
||||||
|
"package": {
|
||||||
|
"package_format": "cyberstrikeai.workflow-package",
|
||||||
|
"format_version": "1.0",
|
||||||
|
"package_id": "pkg_01JWEBHUNT",
|
||||||
|
"package_hash": "sha256:af..."
|
||||||
|
},
|
||||||
|
"workflow": {
|
||||||
|
"source_id": "web-src-hunting",
|
||||||
|
"name": "Web SRC 猎洞",
|
||||||
|
"description": "面向 SRC Web 资产的侦察与漏洞候选流程",
|
||||||
|
"source_revision": 18,
|
||||||
|
"enabled": true,
|
||||||
|
"content_hash": "sha256:51...",
|
||||||
|
"graph_hash": "sha256:a9...",
|
||||||
|
"node_count": 10,
|
||||||
|
"edge_count": 11
|
||||||
|
},
|
||||||
|
"conflict": {
|
||||||
|
"state": "id_conflict",
|
||||||
|
"local_workflow": {
|
||||||
|
"id": "web-src-hunting",
|
||||||
|
"version": 12,
|
||||||
|
"content_hash": "sha256:42...",
|
||||||
|
"graph_hash": "sha256:17..."
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"warnings": []
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
`conflict.state` 固定枚举:
|
||||||
|
|
||||||
|
| 值 | 含义 |
|
||||||
|
|---|---|
|
||||||
|
| `none` | 目标不存在,可 `create`。 |
|
||||||
|
| `identical` | 目标同 ID 且 `content_hash` 相同。 |
|
||||||
|
| `id_conflict` | 目标同 ID,但内容不同。 |
|
||||||
|
|
||||||
|
无效包不创建 inspection,直接返回 `422`。常用错误码:`WFPKG_FILE_REQUIRED`、`WFPKG_FILE_TOO_LARGE`、`WFPKG_INVALID_ARCHIVE`、`WFPKG_UNSUPPORTED_FORMAT`、`WFPKG_INVALID_MANIFEST`、`WFPKG_CHECKSUM_MISMATCH`、`WFPKG_MULTIPLE_WORKFLOWS`、`WFPKG_WORKFLOW_INVALID`。
|
||||||
|
|
||||||
|
### 3.3 读取预检
|
||||||
|
|
||||||
|
```http
|
||||||
|
GET /api/workflow-package-inspections/{inspectionId}
|
||||||
|
```
|
||||||
|
|
||||||
|
用于前端刷新页面后恢复预检状态。仅 inspection 创建者可读取;不存在返回 `404 WFPKG_INSPECTION_NOT_FOUND`,已过期返回 `409 WFPKG_INSPECTION_EXPIRED`。
|
||||||
|
|
||||||
|
### 3.4 应用导入
|
||||||
|
|
||||||
|
```http
|
||||||
|
POST /api/workflow-package-imports
|
||||||
|
Content-Type: application/json
|
||||||
|
Idempotency-Key: 4b75a1eb-7ed1-4eb1-a074-389dba3d4d7b
|
||||||
|
```
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"inspection_id": "wpi_01JQ2K6G7K8W2C1E3R4T5Y6U7I",
|
||||||
|
"resolution": {
|
||||||
|
"action": "overwrite",
|
||||||
|
"new_workflow_id": ""
|
||||||
|
},
|
||||||
|
"confirm_overwrite": true
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
字段规则:
|
||||||
|
|
||||||
|
| 字段 | 规则 |
|
||||||
|
|---|---|
|
||||||
|
| `inspection_id` | 必填;必须是当前用户创建、状态为 `ready` 且未过期的 inspection。 |
|
||||||
|
| `resolution.action` | `create`、`keep_existing`、`overwrite`、`rename` 之一。 |
|
||||||
|
| `resolution.new_workflow_id` | 仅 `rename` 必填;去除首尾空格后 1–128 字符,不得包含控制字符。其他 action 必须传空字符串。 |
|
||||||
|
| `confirm_overwrite` | 仅 `overwrite` 时必须为 `true`。 |
|
||||||
|
| `Idempotency-Key` | 必填 UUID;同一用户、同一 key、同一请求返回原结果;同 key 不同请求返回冲突。 |
|
||||||
|
|
||||||
|
`request_hash` 固定为下列字段按键名升序、无空白序列化后的 SHA-256:`inspection_id`、`resolution.action`、`resolution.new_workflow_id`、`confirm_overwrite`。后端不得将 `Idempotency-Key` 自身计入该 hash。
|
||||||
|
|
||||||
|
动作与 inspection 状态的合法组合:
|
||||||
|
|
||||||
|
| `conflict.state` | 合法 action | 结果 |
|
||||||
|
|---|---|---|
|
||||||
|
| `none` | `create` | 新建目标工作流。 |
|
||||||
|
| `identical` | `keep_existing` | 返回 `skipped_identical`,不修改工作流。 |
|
||||||
|
| `id_conflict` | `keep_existing` | 返回 `kept_existing`,不修改工作流。 |
|
||||||
|
| `id_conflict` | `overwrite` | 完整替换 name、description、graph_json、enabled;本地 version 递增。 |
|
||||||
|
| `id_conflict` | `rename` | 以 `new_workflow_id` 新建副本,version 为 1。 |
|
||||||
|
|
||||||
|
后端在应用事务开始前必须重新读取目标工作流并比较 inspection 中记录的冲突快照;若本地内容在预检后变化,返回 `409 WFPKG_CONFLICT_CHANGED`,前端必须重新预检。
|
||||||
|
|
||||||
|
首次应用成功返回 `201 Created`:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"import": {
|
||||||
|
"id": "wpii_01JQ2M93S2PH0WY8X7B4F8R9QG",
|
||||||
|
"inspection_id": "wpi_01JQ2K6G7K8W2C1E3R4T5Y6U7I",
|
||||||
|
"status": "succeeded",
|
||||||
|
"result": "overwritten",
|
||||||
|
"action": "overwrite",
|
||||||
|
"source_workflow_id": "web-src-hunting",
|
||||||
|
"target_workflow_id": "web-src-hunting",
|
||||||
|
"workflow": {
|
||||||
|
"id": "web-src-hunting",
|
||||||
|
"version": 13,
|
||||||
|
"content_hash": "sha256:51...",
|
||||||
|
"graph_hash": "sha256:a9..."
|
||||||
|
},
|
||||||
|
"applied_at": "2026-07-13T09:05:00Z"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
同一幂等键重试返回 `200 OK` 和完全相同的 `import` 对象。成功写入后必须调用 `InvalidateCompiledCache(workflowID)`。
|
||||||
|
|
||||||
|
错误码:
|
||||||
|
|
||||||
|
| HTTP | code | 触发条件 |
|
||||||
|
|---:|---|---|
|
||||||
|
| 400 | `WFPKG_IDEMPOTENCY_KEY_REQUIRED` | 缺少或非 UUID 幂等键。 |
|
||||||
|
| 404 | `WFPKG_INSPECTION_NOT_FOUND` | inspection 不存在或不属于当前用户。 |
|
||||||
|
| 409 | `WFPKG_INSPECTION_EXPIRED` | inspection 已过期。 |
|
||||||
|
| 409 | `WFPKG_INSPECTION_CONSUMED` | inspection 已被其他幂等键成功应用。 |
|
||||||
|
| 409 | `WFPKG_IDEMPOTENCY_KEY_REUSED` | 同 key 的请求体不同。 |
|
||||||
|
| 409 | `WFPKG_ID_CONFLICT` | action 与预检冲突状态不匹配。 |
|
||||||
|
| 409 | `WFPKG_OVERWRITE_CONFIRMATION_REQUIRED` | overwrite 未确认。 |
|
||||||
|
| 409 | `WFPKG_CONFLICT_CHANGED` | 预检后本地工作流已改变。 |
|
||||||
|
| 422 | `WFPKG_INVALID_ACTION` | action 不在枚举中,或 action 与字段组合不合法。 |
|
||||||
|
| 422 | `WFPKG_INVALID_RENAME_ID` | rename ID 为空、含控制字符或已存在。 |
|
||||||
|
| 500 | `WFPKG_IMPORT_FAILED` | 事务失败;不修改目标工作流。 |
|
||||||
|
|
||||||
|
### 3.5 查询导入结果
|
||||||
|
|
||||||
|
```http
|
||||||
|
GET /api/workflow-package-imports/{importId}
|
||||||
|
```
|
||||||
|
|
||||||
|
仅导入创建者可读取。响应为 3.4 中的 `import` 对象。该接口不提供列表;MVP 的历史审计通过现有审计日志页面查看。
|
||||||
|
|
||||||
|
## 4. SQLite 数据模型
|
||||||
|
|
||||||
|
`workflow_definitions` 不增列、不改语义。新增两张表;DDL 即为迁移目标。
|
||||||
|
|
||||||
|
```sql
|
||||||
|
CREATE TABLE IF NOT EXISTS workflow_package_inspections (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
package_hash TEXT NOT NULL,
|
||||||
|
manifest_json TEXT NOT NULL,
|
||||||
|
workflow_payload_json TEXT NOT NULL,
|
||||||
|
inspection_json TEXT NOT NULL,
|
||||||
|
source_workflow_id TEXT NOT NULL,
|
||||||
|
source_revision INTEGER NOT NULL,
|
||||||
|
source_content_hash TEXT NOT NULL,
|
||||||
|
source_graph_hash TEXT NOT NULL,
|
||||||
|
local_conflict_state TEXT NOT NULL
|
||||||
|
CHECK (local_conflict_state IN ('none', 'identical', 'id_conflict')),
|
||||||
|
local_workflow_id TEXT,
|
||||||
|
local_content_hash TEXT,
|
||||||
|
local_graph_hash TEXT,
|
||||||
|
created_by TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL DEFAULT 'ready'
|
||||||
|
CHECK (status IN ('ready', 'consumed', 'expired')),
|
||||||
|
created_at DATETIME NOT NULL,
|
||||||
|
expires_at DATETIME NOT NULL,
|
||||||
|
consumed_at DATETIME
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_workflow_package_inspections_creator_expiry
|
||||||
|
ON workflow_package_inspections(created_by, expires_at);
|
||||||
|
|
||||||
|
CREATE TABLE IF NOT EXISTS workflow_package_imports (
|
||||||
|
id TEXT PRIMARY KEY,
|
||||||
|
inspection_id TEXT NOT NULL,
|
||||||
|
request_hash TEXT NOT NULL,
|
||||||
|
idempotency_key TEXT NOT NULL,
|
||||||
|
actor_user_id TEXT NOT NULL,
|
||||||
|
action TEXT NOT NULL
|
||||||
|
CHECK (action IN ('create', 'keep_existing', 'overwrite', 'rename')),
|
||||||
|
source_workflow_id TEXT NOT NULL,
|
||||||
|
target_workflow_id TEXT NOT NULL,
|
||||||
|
resulting_workflow_id TEXT,
|
||||||
|
result TEXT NOT NULL
|
||||||
|
CHECK (result IN ('created', 'overwritten', 'renamed', 'kept_existing', 'skipped_identical', 'failed')),
|
||||||
|
error_code TEXT,
|
||||||
|
error_message TEXT,
|
||||||
|
created_at DATETIME NOT NULL,
|
||||||
|
applied_at DATETIME,
|
||||||
|
FOREIGN KEY (inspection_id) REFERENCES workflow_package_inspections(id)
|
||||||
|
);
|
||||||
|
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_workflow_package_imports_actor_key
|
||||||
|
ON workflow_package_imports(actor_user_id, idempotency_key);
|
||||||
|
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_workflow_package_imports_inspection_success
|
||||||
|
ON workflow_package_imports(inspection_id)
|
||||||
|
WHERE result IN ('created', 'overwritten', 'renamed', 'kept_existing', 'skipped_identical');
|
||||||
|
```
|
||||||
|
|
||||||
|
### 4.1 表职责与生命周期
|
||||||
|
|
||||||
|
| 表 | 职责 | 保留规则 |
|
||||||
|
|---|---|---|
|
||||||
|
| `workflow_package_inspections` | 保存已验证的 Manifest、单工作流载荷、冲突快照和前端恢复所需摘要;不保存原始 Zip。 | `expires_at` 为创建后 30 分钟;到期改为 `expired`;清理任务可在 24 小时后删除。 |
|
||||||
|
| `workflow_package_imports` | 导入结果、幂等键和应用记录。 | 保留 90 天;删除不影响既有审计日志。 |
|
||||||
|
|
||||||
|
inspection 创建时写入 `workflow_payload_json`,应用时只读取该已验证载荷,不信任浏览器重新提交的工作流内容。`inspection_json` 是 3.2 成功响应中的安全摘要快照。
|
||||||
|
|
||||||
|
### 4.2 导入事务
|
||||||
|
|
||||||
|
导入应用必须在一个 SQLite 事务中完成以下操作:
|
||||||
|
|
||||||
|
1. 校验 inspection 所属用户、状态、有效期与幂等键。
|
||||||
|
2. 再次读取目标 `workflow_definitions`,验证冲突快照未变化。
|
||||||
|
3. 按 action 新建、覆盖、重命名或保持现有工作流。
|
||||||
|
4. 新增 `workflow_package_imports` 成功行,并把 inspection 改为 `consumed`。
|
||||||
|
5. 提交事务;提交后失效工作流编译缓存并写审计日志。
|
||||||
|
|
||||||
|
任一步失败必须回滚工作流、导入行和 inspection 状态。`workflow_package_imports.result='failed'` 仅在能安全独立记录失败时写入,绝不替代事务回滚。
|
||||||
|
|
||||||
|
## 5. 审计契约
|
||||||
|
|
||||||
|
复用现有 `audit_logs`,不在包表中复制审计全文:
|
||||||
|
|
||||||
|
| 事件 | category | action | resource |
|
||||||
|
|---|---|---|---|
|
||||||
|
| 导出成功 | `workflow_package` | `export` | `workflow/{id}` |
|
||||||
|
| 预检成功 | `workflow_package` | `inspect` | `inspection/{id}` |
|
||||||
|
| 预检失败 | `workflow_package` | `inspect` | 无资源 ID;detail 仅含错误码与包 hash |
|
||||||
|
| 应用成功 | `workflow_package` | `import` | `workflow/{resulting_workflow_id}` |
|
||||||
|
| 应用失败 | `workflow_package` | `import` | `inspection/{id}` |
|
||||||
|
|
||||||
|
## 6. 前后端并行边界
|
||||||
|
|
||||||
|
前端可依据本文直接完成:下载按钮、文件上传、预检结果页、冲突动作选择、`confirm_overwrite` 二次确认、导入结果页和错误码国际化。
|
||||||
|
|
||||||
|
后端可依据本文直接完成:路由、Handler、包解析服务、SQLite 迁移、事务、RBAC 映射、审计和单元/集成测试。
|
||||||
|
|
||||||
|
前端不得自行解析 Zip、计算最终冲突结论或直接提交工作流 JSON;后端是 Manifest、哈希、图校验、冲突复查和导入结果的唯一权威。
|
||||||
@@ -0,0 +1,134 @@
|
|||||||
|
# 本地图编排策略包 MVP 设计
|
||||||
|
|
||||||
|
## 决策摘要
|
||||||
|
|
||||||
|
本期只交付本地图编排策略管理,不接入公共市场、远程仓库、发布上传、账号、评分或订阅能力。目标是先建立稳定的工作流包格式和安全导入闭环;未来市场仅复用该包格式和本地安装器。
|
||||||
|
|
||||||
|
## 目标与非目标
|
||||||
|
|
||||||
|
目标:用户可将单个工作流导出为可离线传输、可审查的包,并在另一实例中完成预检后显式导入。
|
||||||
|
|
||||||
|
本期非目标:
|
||||||
|
|
||||||
|
- 批量导出、批量导入、按标签或角色筛选。
|
||||||
|
- 角色、Skill、工具元数据的实际导出或安装。
|
||||||
|
- 远程仓库配置、策略市场、下载、上传和发布者身份。
|
||||||
|
- 自动合并、三方 diff、跨实例 SemVer 升级、降级和回滚。
|
||||||
|
- 工作流运行记录、会话、项目数据、MCP 密钥或任何可执行载荷。
|
||||||
|
|
||||||
|
## 当前基础
|
||||||
|
|
||||||
|
- 工作流保存在 SQLite 的 `workflow_definitions`,包含 `id`、`name`、`description`、整型 `version`、`graph_json`、`enabled`。
|
||||||
|
- 保存前已有严格的 `ValidateGraphJSON` 校验;MVP 导入必须复用它。
|
||||||
|
- 当前工作流 CRUD、`/validate`、`/dry-run` 和运行 API 不改变。
|
||||||
|
- 角色和 Skill 分别存放于 `roles/`、`skills/`,本期不写入这两个目录。
|
||||||
|
|
||||||
|
## 用户流程
|
||||||
|
|
||||||
|
### 导出
|
||||||
|
|
||||||
|
1. 用户在图编排详情页选择“导出”。
|
||||||
|
2. 系统读取单个工作流定义,生成 `.csapkg.zip`。
|
||||||
|
3. 用户下载包并可解压审查 JSON 与 Manifest。
|
||||||
|
|
||||||
|
### 导入
|
||||||
|
|
||||||
|
1. 用户在图编排列表页选择“导入本地包”。
|
||||||
|
2. 系统上传并解析 Zip,但不写入数据库。
|
||||||
|
3. 系统检查包结构、文件哈希、工作流 JSON,并调用 `ValidateGraphJSON`。
|
||||||
|
4. 用户查看工作流名称、ID、节点/边数量、`graph_json` hash 和冲突结果。
|
||||||
|
5. 用户确认“创建”或在冲突时选择“保留本地 / 覆盖 / 另存为新 ID”。
|
||||||
|
6. 系统写入工作流、失效编译缓存、写入审计日志并返回结果。
|
||||||
|
|
||||||
|
导入始终为两步:预检不会写入;仅确认后的应用步骤会改变本地工作流。缺少本期未处理的工具依赖时可显示提示,但不得阻止仅保存定义的导入。
|
||||||
|
|
||||||
|
## 包格式
|
||||||
|
|
||||||
|
文件扩展名为 `.csapkg.zip`,解压后保持人类可读:
|
||||||
|
|
||||||
|
```text
|
||||||
|
web-src-hunting-1.0.0.csapkg.zip
|
||||||
|
├─ manifest.json
|
||||||
|
├─ checksums.sha256
|
||||||
|
└─ workflows/
|
||||||
|
└─ web-src-hunting.json
|
||||||
|
```
|
||||||
|
|
||||||
|
`manifest.json` 示例:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"package_format": "cyberstrikeai.workflow-package",
|
||||||
|
"format_version": "1.0",
|
||||||
|
"package_id": "pkg_01JWEBHUNT",
|
||||||
|
"created_at": "2026-07-13T10:00:00Z",
|
||||||
|
"items": [
|
||||||
|
{
|
||||||
|
"type": "workflow",
|
||||||
|
"path": "workflows/web-src-hunting.json",
|
||||||
|
"source_id": "web-src-hunting",
|
||||||
|
"source_revision": 18,
|
||||||
|
"content_hash": "sha256:...",
|
||||||
|
"graph_hash": "sha256:..."
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
`workflows/*.json` 保留现有工作流字段。现有整型 `version` 继续作为本地修订号;MVP 不引入 SemVer,也不将版本解释为跨实例升级语义。
|
||||||
|
|
||||||
|
## 冲突规则
|
||||||
|
|
||||||
|
| 目标状态 | 默认行为 | 可选动作 |
|
||||||
|
|---|---|---|
|
||||||
|
| 本地不存在同 ID | 创建 | 无 |
|
||||||
|
| 本地存在同 ID | 保留本地并报告冲突 | 覆盖、另存为新 ID、取消 |
|
||||||
|
| 包内容 hash 与本地一致 | 跳过,视为幂等成功 | 无 |
|
||||||
|
|
||||||
|
覆盖是破坏性操作,必须二次确认。另存为新 ID 时仅修改导入副本 ID,不修改任何角色绑定。
|
||||||
|
|
||||||
|
## 后端边界
|
||||||
|
|
||||||
|
新增 `internal/workflow/package` 包:
|
||||||
|
|
||||||
|
- `manifest.go`:包格式、JSON 解析和版本兼容。
|
||||||
|
- `exporter.go`:从 `workflow_definitions` 生成 Zip。
|
||||||
|
- `inspector.go`:安全解压、哈希校验、结构校验与 `ValidateGraphJSON` 复用。
|
||||||
|
- `importer.go`:冲突策略、数据库写入和编译缓存失效。
|
||||||
|
|
||||||
|
建议新增 API:
|
||||||
|
|
||||||
|
| API | 语义 | 写入 / 确认 |
|
||||||
|
|---|---|---|
|
||||||
|
| `POST /api/workflow-packages/exports` | 生成并下载单工作流包 | 无写入;无需确认 |
|
||||||
|
| `POST /api/workflow-packages/inspections` | 上传并预检包 | 无写入;无需确认 |
|
||||||
|
| `POST /api/workflow-package-imports` | 创建并应用导入计划 | 有写入;覆盖时需确认 |
|
||||||
|
|
||||||
|
现有 API 不变。权限建议沿用 `workflow:read` 用于导出,`workflow:write` 用于导入;若后续需要细粒度授权,再拆出 `workflow:export`、`workflow:import`。
|
||||||
|
|
||||||
|
## 安全与审计
|
||||||
|
|
||||||
|
- 仅允许 Manifest 与声明的 JSON 文本;拒绝 Zip 路径穿越、重复条目、软链接、超大解压和未知文件。
|
||||||
|
- 校验每个包项 SHA-256。
|
||||||
|
- 不导出或导入密钥、Token、MCP 连接配置、运行记录和可执行文件。
|
||||||
|
- 预检与实际导入分别写审计日志;应用记录包 hash、工作流 ID、策略和结果。
|
||||||
|
|
||||||
|
## 前端范围
|
||||||
|
|
||||||
|
- 图编排列表页:增加“导入本地包”。
|
||||||
|
- 图编排详情页:增加“导出”。
|
||||||
|
- 导入 Modal:上传、预检结果、冲突策略和确认应用。
|
||||||
|
- `workflows.js` 负责调用新 API;不增加策略市场、远程仓库或发布 UI。
|
||||||
|
|
||||||
|
## 验收与测试
|
||||||
|
|
||||||
|
- 导出 `web-src-hunting` 后可解压并审查 Manifest 与完整工作流 JSON。
|
||||||
|
- 导入包在确认前不改变数据库。
|
||||||
|
- 合法图可创建;非法 DAG 或节点参数由 `ValidateGraphJSON` 拒绝。
|
||||||
|
- 同 ID 默认不覆盖;覆盖须确认;另存生成新 ID。
|
||||||
|
- 包 hash 不匹配、路径穿越、未知文件、超限文件均被拒绝。
|
||||||
|
- 成功导入后工作流可由现有 GET API 读取,且编译缓存已失效。
|
||||||
|
|
||||||
|
## 后续扩展边界
|
||||||
|
|
||||||
|
v1 可增加批量导入导出、Role/Skill 可选项、工具依赖展示与导入历史。v2 可基于同一 `.csapkg.zip` 增加语义标识、SemVer、升级 diff、三方合并与回滚。公共策略市场、远程仓库和发布上传属于 v3 之后的独立子项目,不进入本期实现。
|
||||||
+30
-28
@@ -1,30 +1,32 @@
|
|||||||
# 中文文档
|
# 中文文档
|
||||||
|
|
||||||
- [部署指南](deployment.md):部署形态、HTTPS、反向代理、systemd、备份、升级和验收。
|
[文档首页](../README.md) | [English](../en-US/README.md)
|
||||||
- [运维 Runbooks](runbooks.md):生产部署、外部 MCP、知识库、Web 测试、C2 清理和工具排障的操作步骤。
|
|
||||||
- [配置画像](configuration-profiles.md):本地开发、内网团队、知识库、高审计生产、C2 演练等推荐配置。
|
## 按目标开始
|
||||||
- [安全加固指南](security-hardening.md):上线前基线、反向代理、HITL 白名单、文件权限和周期巡检。
|
|
||||||
- [API Recipes](api-recipes.md):登录、Agent、流式、多代理、上传、漏洞、知识库、MCP 和审计导出示例。
|
- **快速体验**:[部署指南](deployment.md) → [配置参考](configuration.md) → [排错指南](troubleshooting.md)
|
||||||
- [贡献规范](contributing-guide.md):新增 API、配置、工具、前端、数据库、高风险能力和文档的 checklist。
|
- **生产部署**:[配置画像](configuration-profiles.md) → [安全加固](security-hardening.md) → [运维 Runbooks](runbooks.md) → [审计与监控](audit-and-monitoring.md)
|
||||||
- [配置参考](configuration.md):`config.yaml` 字段、热应用边界、参数建议和源码锚点。
|
- **接入与自动化**:[API 参考](api-reference.md) → [API Recipes](api-recipes.md) → [MCP 联邦](mcp-federation.md)
|
||||||
- [安全模型](security-model.md):信任边界、HITL、工具执行、C2/WebShell 与数据安全。
|
- **参与开发**:[开发者指南](developer-guide.md) → [测试指南](testing.md) → [贡献规范](contributing-guide.md)
|
||||||
- [RBAC 权限管理](rbac.md):平台用户、系统/自定义角色、权限目录、逐权限 Scope、资源授权、Agent/MCP/机器人边界与 API 示例。
|
|
||||||
- [架构说明](architecture.md):请求路径、模块关系、复杂度热点和设计取舍。
|
## 核心概念与编排
|
||||||
- [API 参考](api-reference.md):认证、OpenAPI、SSE、稳定性分层和常用接口。
|
|
||||||
- [排错指南](troubleshooting.md):诊断顺序、最小命令、常见误判和故障模板。
|
- [架构说明](architecture.md) · [安全模型](security-model.md) · [RBAC](rbac.md)
|
||||||
- [审计与监控](audit-and-monitoring.md):平台审计、工具监控、HITL 日志和保留策略。
|
- [Agent 与角色](agent-and-role-guide.md) · [Skills](skills-guide.md) · [Eino 多代理](MULTI_AGENT_EINO.md)
|
||||||
- [知识库](knowledge-base.md):索引链路、检索调参、日志分析和内容写法。
|
- [图编排](workflow-graph.md) · [人机协同最佳实践](hitl-best-practices.md)
|
||||||
- [C2 使用说明](c2.md):生命周期、任务分级、事件复盘和安全建议。
|
|
||||||
- [WebShell 管理](webshell.md):操作分层、连接命名、AI 约束和排错。
|
## 功能指南
|
||||||
- [MCP 联邦](mcp-federation.md):内置 MCP、外部 MCP、生命周期和工具命名。
|
|
||||||
- [Agent 与角色](agent-and-role-guide.md):角色、子代理、Skill、编排模式和工具可见性。
|
- [知识库](knowledge-base.md) · [机器人接入](robot.md) · [视觉分析](VISION.md)
|
||||||
- [Skills 指南](skills-guide.md):Skill 结构、渐进式披露、反模式和本地工具风险。
|
- [WebShell](webshell.md) · [C2](c2.md) · [MCP 联邦](mcp-federation.md)
|
||||||
- [插件开发](plugin-development.md):API 插件、MCP 插件、资源包插件和安全边界。
|
|
||||||
- [发布流程](release-process.md):发布风险、配置兼容、数据库迁移和验收。
|
## 运维与参考
|
||||||
- [测试指南](testing.md):测试分层、回归重点、测试数据和失败用例。
|
|
||||||
- [图编排使用说明](workflow-graph.md)
|
- [部署指南](deployment.md) · [配置参考](configuration.md) · [配置画像](configuration-profiles.md)
|
||||||
- [人机协同最佳实践](hitl-best-practices.md)
|
- [安全加固](security-hardening.md) · [审计与监控](audit-and-monitoring.md) · [运维 Runbooks](runbooks.md)
|
||||||
- [机器人使用说明](robot.md)
|
- [API 参考](api-reference.md) · [API Recipes](api-recipes.md) · [排错指南](troubleshooting.md)
|
||||||
- [视觉分析](VISION.md)
|
|
||||||
- [前端国际化方案](frontend-i18n.md)
|
## 开发与发布
|
||||||
- [Eino 多代理改造说明](MULTI_AGENT_EINO.md)
|
|
||||||
|
- [开发者指南](developer-guide.md) · [插件开发](plugin-development.md) · [前端国际化](frontend-i18n.md)
|
||||||
|
- [测试指南](testing.md) · [贡献规范](contributing-guide.md) · [发布流程](release-process.md)
|
||||||
|
|||||||
@@ -21,7 +21,7 @@ server:
|
|||||||
tls_enabled: true
|
tls_enabled: true
|
||||||
tls_auto_self_sign: true
|
tls_auto_self_sign: true
|
||||||
auth:
|
auth:
|
||||||
password: "dev-only-change-me"
|
session_duration_hours: 12
|
||||||
audit:
|
audit:
|
||||||
enabled: true
|
enabled: true
|
||||||
retention_days: 7
|
retention_days: 7
|
||||||
@@ -54,7 +54,7 @@ server:
|
|||||||
port: 8080
|
port: 8080
|
||||||
tls_enabled: false
|
tls_enabled: false
|
||||||
auth:
|
auth:
|
||||||
password: "<long-random-password>"
|
session_duration_hours: 12
|
||||||
audit:
|
audit:
|
||||||
enabled: true
|
enabled: true
|
||||||
retention_days: 30
|
retention_days: 30
|
||||||
@@ -107,7 +107,6 @@ multi_agent:
|
|||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
auth:
|
auth:
|
||||||
password: "<managed-secret>"
|
|
||||||
session_duration_hours: 8
|
session_duration_hours: 8
|
||||||
audit:
|
audit:
|
||||||
enabled: true
|
enabled: true
|
||||||
|
|||||||
@@ -5,14 +5,16 @@ CyberStrikeAI 的主配置文件是 `config.yaml`。大多数配置也可以在
|
|||||||
## 基础配置
|
## 基础配置
|
||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
version: "v1.6.51"
|
version: "vX.Y.Z" # 占位符;请使用 config.example.yaml 中当前发布版本的值
|
||||||
server:
|
server:
|
||||||
host: 0.0.0.0
|
host: 0.0.0.0
|
||||||
port: 8080
|
port: 8080
|
||||||
tls_enabled: true
|
tls_enabled: true
|
||||||
tls_auto_self_sign: true
|
tls_auto_self_sign: true
|
||||||
|
# 可选:其他可信 Web 集成;Chromium 浏览器插件无需配置
|
||||||
|
# cors_allowed_origins:
|
||||||
|
# - https://trusted-integration.example
|
||||||
auth:
|
auth:
|
||||||
password: "change-me"
|
|
||||||
session_duration_hours: 12
|
session_duration_hours: 12
|
||||||
log:
|
log:
|
||||||
level: info
|
level: info
|
||||||
@@ -22,8 +24,9 @@ log:
|
|||||||
- `version`:前端展示版本。
|
- `version`:前端展示版本。
|
||||||
- `server.host/port`:Web 服务监听地址和端口。
|
- `server.host/port`:Web 服务监听地址和端口。
|
||||||
- `server.tls_*`:HTTPS 配置。生产环境建议使用 `tls_cert_path` 和 `tls_key_path`。
|
- `server.tls_*`:HTTPS 配置。生产环境建议使用 `tls_cert_path` 和 `tls_key_path`。
|
||||||
- `auth.password`:Web 登录密码,必须改为强密码。
|
- Chromium 浏览器插件的合法 `chrome-extension://<32位插件ID>` Origin 会被自动识别,无需配置。插件仍需按域授权,并使用密码登录与 Bearer Token 调用 API。
|
||||||
- `auth.session_duration_hours`:登录会话有效期。
|
- `server.cors_allowed_origins`:仅供其他可信 Web 集成使用的额外 Origin 精确白名单;不支持 `*`,修改后需重启服务。
|
||||||
|
- `auth.session_duration_hours`:登录会话有效期(小时)。登录密码由 RBAC 用户管理,首次启动时在控制台输出 `admin` 初始密码。
|
||||||
- `log.output`:可以是 `stdout`、`stderr` 或文件路径。
|
- `log.output`:可以是 `stdout`、`stderr` 或文件路径。
|
||||||
|
|
||||||
## 模型配置
|
## 模型配置
|
||||||
|
|||||||
@@ -103,6 +103,16 @@ docs/en-US/
|
|||||||
- `docs/zh-CN/README.md`
|
- `docs/zh-CN/README.md`
|
||||||
- `docs/en-US/README.md`
|
- `docs/en-US/README.md`
|
||||||
|
|
||||||
|
提交文档变更前运行:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
python3 scripts/check-docs.py
|
||||||
|
```
|
||||||
|
|
||||||
|
该检查会验证本地链接、代码块闭合、中英文文件名对齐、语言导航覆盖率,以及根 README 中的 Go 版本是否与 `go.mod` 一致。版本化示例应尽量从 `go.mod`、`config.example.yaml` 等权威文件派生,避免手工同步。
|
||||||
|
|
||||||
|
当文档、`go.mod` 或检查脚本变更时,GitHub Actions 中的 `Documentation` 工作流会自动运行同一条命令。
|
||||||
|
|
||||||
## Review 关注点
|
## Review 关注点
|
||||||
|
|
||||||
代码评审优先看:
|
代码评审优先看:
|
||||||
|
|||||||
@@ -48,7 +48,7 @@ config.yaml
|
|||||||
|
|
||||||
```yaml
|
```yaml
|
||||||
auth:
|
auth:
|
||||||
password: "<long-random-password>"
|
session_duration_hours: 12
|
||||||
server:
|
server:
|
||||||
host: 127.0.0.1
|
host: 127.0.0.1
|
||||||
port: 8080
|
port: 8080
|
||||||
|
|||||||
@@ -6,7 +6,7 @@
|
|||||||
|
|
||||||
## 上线前必做
|
## 上线前必做
|
||||||
|
|
||||||
- 修改 `auth.password` 为长随机密码。
|
- 首次部署后立即修改 `admin` 初始密码(Web 界面或平台权限 → 用户管理)。
|
||||||
- 使用 HTTPS,或放在可信反向代理之后。
|
- 使用 HTTPS,或放在可信反向代理之后。
|
||||||
- 限制来源 IP、VPN 或堡垒机访问。
|
- 限制来源 IP、VPN 或堡垒机访问。
|
||||||
- 开启 `audit.enabled`。
|
- 开启 `audit.enabled`。
|
||||||
|
|||||||
@@ -16,9 +16,9 @@ CyberStrikeAI 面向授权安全测试场景,内置命令执行、MCP 工具
|
|||||||
|
|
||||||
## 认证与会话
|
## 认证与会话
|
||||||
|
|
||||||
`auth.password` 是 Web 登录密码。建议:
|
Web 登录凭据由 RBAC 用户管理(默认内置 `admin` 账号)。建议:
|
||||||
|
|
||||||
- 首次部署立即修改默认密码。
|
- 首次部署后立即修改 `admin` 初始密码(控制台首次启动会输出)。
|
||||||
- 使用长随机密码,并限制分享范围。
|
- 使用长随机密码,并限制分享范围。
|
||||||
- 将服务放在内网、VPN、堡垒机或反向代理认证后面。
|
- 将服务放在内网、VPN、堡垒机或反向代理认证后面。
|
||||||
- 生产环境开启 HTTPS,避免明文传输 Cookie。
|
- 生产环境开启 HTTPS,避免明文传输 Cookie。
|
||||||
|
|||||||
@@ -23,12 +23,12 @@ https://127.0.0.1:8080/
|
|||||||
|
|
||||||
检查:
|
检查:
|
||||||
|
|
||||||
- `config.yaml` 中的 `auth.password`。
|
- RBAC 用户密码是否正确(默认 `admin`;首次启动密码见控制台输出)。
|
||||||
- 是否修改后未重启或未应用配置。
|
- 是否修改密码后旧会话已失效,需重新登录。
|
||||||
- 浏览器 Cookie 是否异常,可尝试无痕窗口。
|
- 浏览器 Cookie 是否异常,可尝试无痕窗口。
|
||||||
- 审计日志中是否有登录失败节流。
|
- 审计日志中是否有登录失败节流。
|
||||||
|
|
||||||
生产环境忘记密码时,需要在服务器上修改 `config.yaml` 并重启服务。
|
生产环境忘记密码时,需在服务器上通过 RBAC 用户管理重置,或直接更新数据库中的用户密码哈希。
|
||||||
|
|
||||||
## 模型无响应
|
## 模型无响应
|
||||||
|
|
||||||
|
|||||||
Binary file not shown.
|
Before Width: | Height: | Size: 86 KiB After Width: | Height: | Size: 88 KiB |
@@ -661,7 +661,7 @@ func (a *Agent) UpdateToolDescriptionMode(mode string) {
|
|||||||
mode = "short"
|
mode = "short"
|
||||||
}
|
}
|
||||||
a.toolDescriptionMode = mode
|
a.toolDescriptionMode = mode
|
||||||
a.logger.Info("Agent工具描述模式已更新", zap.String("tool_description_mode", mode))
|
a.logger.Debug("Agent工具描述模式已更新", zap.String("tool_description_mode", mode))
|
||||||
}
|
}
|
||||||
|
|
||||||
// RepairOrphanToolMessages 清理失去配对的tool消息和未完成的tool_calls,避免OpenAI报错
|
// RepairOrphanToolMessages 清理失去配对的tool消息和未完成的tool_calls,避免OpenAI报错
|
||||||
|
|||||||
+99
-28
@@ -66,6 +66,7 @@ type App struct {
|
|||||||
slackCancel context.CancelFunc // Slack Socket Mode 取消函数
|
slackCancel context.CancelFunc // Slack Socket Mode 取消函数
|
||||||
discordCancel context.CancelFunc // Discord Gateway 取消函数
|
discordCancel context.CancelFunc // Discord Gateway 取消函数
|
||||||
qqCancel context.CancelFunc // QQ WebSocket 取消函数
|
qqCancel context.CancelFunc // QQ WebSocket 取消函数
|
||||||
|
alertCancel context.CancelFunc // 漏洞提醒持久化投递 worker
|
||||||
c2Manager *c2.Manager // C2 管理器(未启用 C2 时为 nil)
|
c2Manager *c2.Manager // C2 管理器(未启用 C2 时为 nil)
|
||||||
c2Watchdog *c2.SessionWatchdog // C2 会话看门狗
|
c2Watchdog *c2.SessionWatchdog // C2 会话看门狗
|
||||||
c2WatchdogCancel context.CancelFunc // 看门狗取消函数
|
c2WatchdogCancel context.CancelFunc // 看门狗取消函数
|
||||||
@@ -83,7 +84,7 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
|
|||||||
router := gin.Default()
|
router := gin.Default()
|
||||||
|
|
||||||
// CORS中间件
|
// CORS中间件
|
||||||
router.Use(corsMiddleware())
|
router.Use(corsMiddleware(cfg.Server.CORSAllowedOrigins))
|
||||||
|
|
||||||
// 初始化数据库
|
// 初始化数据库
|
||||||
dbPath := cfg.Database.Path
|
dbPath := cfg.Database.Path
|
||||||
@@ -101,13 +102,12 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
|
|||||||
return nil, fmt.Errorf("初始化数据库失败: %w", err)
|
return nil, fmt.Errorf("初始化数据库失败: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
// 认证管理器(数据库初始化后挂载 RBAC,以兼容旧的单密码配置)
|
// 认证管理器(数据库初始化后挂载 RBAC)
|
||||||
authManager, err := security.NewAuthManager(cfg.Auth.Password, cfg.Auth.SessionDurationHours)
|
authManager := security.NewAuthManager(cfg.Auth.SessionDurationHours)
|
||||||
if err != nil {
|
if generatedPassword, err := authManager.AttachRBACStore(db); err != nil {
|
||||||
return nil, fmt.Errorf("初始化认证失败: %w", err)
|
|
||||||
}
|
|
||||||
if err := authManager.AttachRBACStore(db); err != nil {
|
|
||||||
return nil, fmt.Errorf("初始化RBAC失败: %w", err)
|
return nil, fmt.Errorf("初始化RBAC失败: %w", err)
|
||||||
|
} else if generatedPassword != "" {
|
||||||
|
config.PrintBootstrapAdminPassword(generatedPassword)
|
||||||
}
|
}
|
||||||
for platform, userID := range cfg.Robots.ServiceAccountUserIDs() {
|
for platform, userID := range cfg.Robots.ServiceAccountUserIDs() {
|
||||||
user, userErr := db.GetRBACUserByID(userID)
|
user, userErr := db.GetRBACUserByID(userID)
|
||||||
@@ -120,11 +120,26 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
|
|||||||
audit.RegisterConversationCreateHook(auditSvc)
|
audit.RegisterConversationCreateHook(auditSvc)
|
||||||
auditSvc.PurgeExpired()
|
auditSvc.PurgeExpired()
|
||||||
audit.StartRetentionLoop(auditSvc, log.Logger)
|
audit.StartRetentionLoop(auditSvc, log.Logger)
|
||||||
|
if err := db.PurgeWorkflowPackageLifecycle(time.Now().UTC()); err != nil {
|
||||||
|
log.Logger.Warn("清理过期工作流包记录失败", zap.Error(err))
|
||||||
|
}
|
||||||
|
go func() {
|
||||||
|
ticker := time.NewTicker(time.Hour)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for range ticker.C {
|
||||||
|
if err := db.PurgeWorkflowPackageLifecycle(time.Now().UTC()); err != nil {
|
||||||
|
log.Logger.Warn("清理过期工作流包记录失败", zap.Error(err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}()
|
||||||
|
|
||||||
monitorRetention := monitor.NewService(db, cfg, log.Logger)
|
monitorRetention := monitor.NewService(db, cfg, log.Logger)
|
||||||
monitorRetention.PurgeExpired()
|
monitorRetention.PurgeExpired()
|
||||||
monitor.StartRetentionLoop(monitorRetention, log.Logger)
|
monitor.StartRetentionLoop(monitorRetention, log.Logger)
|
||||||
|
|
||||||
|
if err := handler.NewHITLManager(db, log.Logger).EnsureSchema(); err != nil {
|
||||||
|
log.Logger.Warn("初始化 HITL 表失败", zap.Error(err))
|
||||||
|
}
|
||||||
hitlRetention := hitl.NewService(db, cfg, log.Logger)
|
hitlRetention := hitl.NewService(db, cfg, log.Logger)
|
||||||
hitlRetention.PurgeExpired()
|
hitlRetention.PurgeExpired()
|
||||||
hitl.StartRetentionLoop(hitlRetention, log.Logger)
|
hitl.StartRetentionLoop(hitlRetention, log.Logger)
|
||||||
@@ -146,13 +161,6 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
|
|||||||
registerProjectFactTools(mcpServer, db, cfg, log.Logger)
|
registerProjectFactTools(mcpServer, db, cfg, log.Logger)
|
||||||
registerVisionTools(mcpServer, cfg, log.Logger)
|
registerVisionTools(mcpServer, cfg, log.Logger)
|
||||||
|
|
||||||
if cfg.Auth.GeneratedPassword != "" {
|
|
||||||
config.PrintGeneratedPasswordWarning(cfg.Auth.GeneratedPassword, cfg.Auth.GeneratedPasswordPersisted, cfg.Auth.GeneratedPasswordPersistErr)
|
|
||||||
cfg.Auth.GeneratedPassword = ""
|
|
||||||
cfg.Auth.GeneratedPasswordPersisted = false
|
|
||||||
cfg.Auth.GeneratedPasswordPersistErr = ""
|
|
||||||
}
|
|
||||||
|
|
||||||
// 创建外部MCP管理器(使用与内部MCP服务器相同的存储)
|
// 创建外部MCP管理器(使用与内部MCP服务器相同的存储)
|
||||||
externalMCPMgr := mcp.NewExternalMCPManagerWithStorage(log.Logger, db)
|
externalMCPMgr := mcp.NewExternalMCPManagerWithStorage(log.Logger, db)
|
||||||
externalMCPMgr.SetToolAuthorizer(externalMCPToolAuthorizer())
|
externalMCPMgr.SetToolAuthorizer(externalMCPToolAuthorizer())
|
||||||
@@ -181,7 +189,7 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
|
|||||||
var knowledgeHandler *handler.KnowledgeHandler
|
var knowledgeHandler *handler.KnowledgeHandler
|
||||||
|
|
||||||
var knowledgeDBConn *database.DB
|
var knowledgeDBConn *database.DB
|
||||||
log.Logger.Info("检查知识库配置", zap.Bool("enabled", cfg.Knowledge.Enabled))
|
log.Logger.Debug("检查知识库配置", zap.Bool("enabled", cfg.Knowledge.Enabled))
|
||||||
if cfg.Knowledge.Enabled {
|
if cfg.Knowledge.Enabled {
|
||||||
// 确定知识库数据库路径
|
// 确定知识库数据库路径
|
||||||
knowledgeDBPath := cfg.Database.KnowledgeDBPath
|
knowledgeDBPath := cfg.Database.KnowledgeDBPath
|
||||||
@@ -324,7 +332,7 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
|
|||||||
}
|
}
|
||||||
|
|
||||||
skillsDir := skillpackage.SkillsRootFromConfig(cfg.SkillsDir, configPath)
|
skillsDir := skillpackage.SkillsRootFromConfig(cfg.SkillsDir, configPath)
|
||||||
log.Logger.Info("Skills 目录(Eino ADK skill 中间件 + Web 管理 API)", zap.String("skillsDir", skillsDir))
|
log.Logger.Debug("Skills 目录(Eino ADK skill 中间件 + Web 管理 API)", zap.String("skillsDir", skillsDir))
|
||||||
configDir := filepath.Dir(configPath)
|
configDir := filepath.Dir(configPath)
|
||||||
plantaskRel := strings.TrimSpace(cfg.MultiAgent.EinoMiddleware.PlantaskRelDir)
|
plantaskRel := strings.TrimSpace(cfg.MultiAgent.EinoMiddleware.PlantaskRelDir)
|
||||||
if plantaskRel == "" {
|
if plantaskRel == "" {
|
||||||
@@ -350,7 +358,7 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
|
|||||||
}
|
}
|
||||||
markdownAgentsHandler := handler.NewMarkdownAgentsHandler(agentsDir)
|
markdownAgentsHandler := handler.NewMarkdownAgentsHandler(agentsDir)
|
||||||
markdownAgentsHandler.SetAudit(auditSvc)
|
markdownAgentsHandler.SetAudit(auditSvc)
|
||||||
log.Logger.Info("多代理 Markdown 子 Agent 目录", zap.String("agentsDir", agentsDir))
|
log.Logger.Debug("多代理 Markdown 子 Agent 目录", zap.String("agentsDir", agentsDir))
|
||||||
|
|
||||||
// 创建处理器
|
// 创建处理器
|
||||||
agentHandler := handler.NewAgentHandler(agent, db, cfg, log.Logger)
|
agentHandler := handler.NewAgentHandler(agent, db, cfg, log.Logger)
|
||||||
@@ -421,6 +429,7 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
|
|||||||
auditHandler := handler.NewAuditHandler(db, auditSvc, log.Logger)
|
auditHandler := handler.NewAuditHandler(db, auditSvc, log.Logger)
|
||||||
robotHandler := handler.NewRobotHandler(cfg, db, agentHandler, log.Logger)
|
robotHandler := handler.NewRobotHandler(cfg, db, agentHandler, log.Logger)
|
||||||
robotHandler.SetAudit(auditSvc)
|
robotHandler.SetAudit(auditSvc)
|
||||||
|
db.SetVulnerabilityCreatedHook(robotHandler.NotifyNewVulnerability)
|
||||||
openAPIHandler := handler.NewOpenAPIHandler(db, log.Logger, conversationHandler, agentHandler)
|
openAPIHandler := handler.NewOpenAPIHandler(db, log.Logger, conversationHandler, agentHandler)
|
||||||
|
|
||||||
// 创建 App 实例(部分字段稍后填充)
|
// 创建 App 实例(部分字段稍后填充)
|
||||||
@@ -449,6 +458,9 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
|
|||||||
}
|
}
|
||||||
// 飞书/钉钉长连接(无需公网),启用时在后台启动;后续前端应用配置时会通过 RestartRobotConnections 重启
|
// 飞书/钉钉长连接(无需公网),启用时在后台启动;后续前端应用配置时会通过 RestartRobotConnections 重启
|
||||||
app.startRobotConnections()
|
app.startRobotConnections()
|
||||||
|
alertCtx, alertCancel := context.WithCancel(context.Background())
|
||||||
|
app.alertCancel = alertCancel
|
||||||
|
go robotHandler.RunVulnerabilityAlertWorker(alertCtx)
|
||||||
|
|
||||||
// 设置漏洞工具注册器(内置工具,必须设置)
|
// 设置漏洞工具注册器(内置工具,必须设置)
|
||||||
vulnerabilityRegistrar := func() error {
|
vulnerabilityRegistrar := func() error {
|
||||||
@@ -635,20 +647,20 @@ func (a *App) RunWithContext(ctx context.Context) error {
|
|||||||
}
|
}
|
||||||
switch tlsMode {
|
switch tlsMode {
|
||||||
case mainTLSFromFiles:
|
case mainTLSFromFiles:
|
||||||
a.logger.Info("启动 HTTPS 主服务(已启用 HTTP/2 协商)",
|
a.logger.Debug("启动 HTTPS 主服务(已启用 HTTP/2 协商)",
|
||||||
zap.String("address", addr),
|
zap.String("address", addr),
|
||||||
zap.String("cert", certFile),
|
zap.String("cert", certFile),
|
||||||
)
|
)
|
||||||
case mainTLSInMemorySelfSigned:
|
case mainTLSInMemorySelfSigned:
|
||||||
a.logger.Info("启动 HTTPS 主服务(内存自签证书,仅测试;已启用 HTTP/2 协商)",
|
a.logger.Debug("启动 HTTPS 主服务(内存自签证书,仅测试;已启用 HTTP/2 协商)",
|
||||||
zap.String("address", addr),
|
zap.String("address", addr),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
if httpRedirect {
|
if httpRedirect {
|
||||||
a.logger.Info("已启用 HTTP→HTTPS 自动跳转(同端口嗅探分流)", zap.String("address", addr))
|
a.logger.Debug("已启用 HTTP→HTTPS 自动跳转(同端口嗅探分流)", zap.String("address", addr))
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
a.logger.Info("启动 HTTP 主服务", zap.String("address", addr))
|
a.logger.Debug("启动 HTTP 主服务", zap.String("address", addr))
|
||||||
}
|
}
|
||||||
|
|
||||||
// 监听 context 取消,优雅关闭 HTTP 服务器
|
// 监听 context 取消,优雅关闭 HTTP 服务器
|
||||||
@@ -710,6 +722,10 @@ func (a *App) Shutdown() {
|
|||||||
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||||
_ = einoobserve.ShutdownOtel(shutdownCtx)
|
_ = einoobserve.ShutdownOtel(shutdownCtx)
|
||||||
shutdownCancel()
|
shutdownCancel()
|
||||||
|
if a.alertCancel != nil {
|
||||||
|
a.alertCancel()
|
||||||
|
a.alertCancel = nil
|
||||||
|
}
|
||||||
|
|
||||||
// 停止钉钉/飞书长连接
|
// 停止钉钉/飞书长连接
|
||||||
a.robotMu.Lock()
|
a.robotMu.Lock()
|
||||||
@@ -1199,6 +1215,8 @@ func setupRoutes(
|
|||||||
protected.DELETE("/vulnerabilities/batch", vulnerabilityHandler.BatchDeleteVulnerabilities)
|
protected.DELETE("/vulnerabilities/batch", vulnerabilityHandler.BatchDeleteVulnerabilities)
|
||||||
protected.GET("/vulnerabilities/filter-options", vulnerabilityHandler.GetVulnerabilityFilterOptions)
|
protected.GET("/vulnerabilities/filter-options", vulnerabilityHandler.GetVulnerabilityFilterOptions)
|
||||||
protected.GET("/vulnerabilities/stats", vulnerabilityHandler.GetVulnerabilityStats)
|
protected.GET("/vulnerabilities/stats", vulnerabilityHandler.GetVulnerabilityStats)
|
||||||
|
protected.GET("/vulnerability-alerts/subscription", vulnerabilityHandler.GetMyAlertSubscription)
|
||||||
|
protected.PUT("/vulnerability-alerts/subscription", vulnerabilityHandler.UpdateMyAlertSubscription)
|
||||||
protected.GET("/vulnerabilities/:id", vulnerabilityHandler.GetVulnerability)
|
protected.GET("/vulnerabilities/:id", vulnerabilityHandler.GetVulnerability)
|
||||||
protected.POST("/vulnerabilities", vulnerabilityHandler.CreateVulnerability)
|
protected.POST("/vulnerabilities", vulnerabilityHandler.CreateVulnerability)
|
||||||
protected.PUT("/vulnerabilities/:id", vulnerabilityHandler.UpdateVulnerability)
|
protected.PUT("/vulnerabilities/:id", vulnerabilityHandler.UpdateVulnerability)
|
||||||
@@ -1308,6 +1326,11 @@ func setupRoutes(
|
|||||||
protected.POST("/workflows/runs/:runId/resume", workflowHandler.ResumeRun)
|
protected.POST("/workflows/runs/:runId/resume", workflowHandler.ResumeRun)
|
||||||
protected.POST("/workflows/validate", workflowHandler.Validate)
|
protected.POST("/workflows/validate", workflowHandler.Validate)
|
||||||
protected.POST("/workflows/dry-run", workflowHandler.DryRun)
|
protected.POST("/workflows/dry-run", workflowHandler.DryRun)
|
||||||
|
protected.GET("/workflows/:id/package", workflowHandler.ExportPackage)
|
||||||
|
protected.POST("/workflow-package-inspections", workflowHandler.CreatePackageInspection)
|
||||||
|
protected.GET("/workflow-package-inspections/:inspectionId", workflowHandler.GetPackageInspection)
|
||||||
|
protected.POST("/workflow-package-imports", workflowHandler.ApplyPackageImport)
|
||||||
|
protected.GET("/workflow-package-imports/:importId", workflowHandler.GetPackageImport)
|
||||||
protected.GET("/workflows", workflowHandler.List)
|
protected.GET("/workflows", workflowHandler.List)
|
||||||
protected.GET("/workflows/:id", workflowHandler.Get)
|
protected.GET("/workflows/:id", workflowHandler.Get)
|
||||||
protected.POST("/workflows", workflowHandler.Create)
|
protected.POST("/workflows", workflowHandler.Create)
|
||||||
@@ -1508,7 +1531,7 @@ func registerWebshellTools(mcpServer *mcp.Server, db *database.DB, webshellHandl
|
|||||||
}
|
}
|
||||||
mcpServer.RegisterTool(writeTool, writeHandler)
|
mcpServer.RegisterTool(writeTool, writeHandler)
|
||||||
|
|
||||||
logger.Info("WebShell 工具注册成功")
|
logger.Debug("WebShell 工具注册成功")
|
||||||
}
|
}
|
||||||
|
|
||||||
// registerWebshellManagementTools 注册 WebShell 连接管理 MCP 工具
|
// registerWebshellManagementTools 注册 WebShell 连接管理 MCP 工具
|
||||||
@@ -1879,7 +1902,7 @@ func registerWebshellManagementTools(mcpServer *mcp.Server, db *database.DB, web
|
|||||||
}
|
}
|
||||||
mcpServer.RegisterTool(testTool, testHandler)
|
mcpServer.RegisterTool(testTool, testHandler)
|
||||||
|
|
||||||
logger.Info("WebShell 管理工具注册成功")
|
logger.Debug("WebShell 管理工具注册成功")
|
||||||
}
|
}
|
||||||
|
|
||||||
// initializeKnowledge 初始化知识库组件(用于动态初始化)
|
// initializeKnowledge 初始化知识库组件(用于动态初始化)
|
||||||
@@ -2046,22 +2069,36 @@ func initializeKnowledge(
|
|||||||
return knowledgeHandler, nil
|
return knowledgeHandler, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// corsMiddleware CORS中间件
|
// corsMiddleware allows same-origin requests, valid Chromium extension
|
||||||
func corsMiddleware() gin.HandlerFunc {
|
// origins, and exact origins explicitly configured by the operator. CORS is
|
||||||
|
// not an authentication boundary; API access still requires a valid session.
|
||||||
|
func corsMiddleware(configuredOrigins []string) gin.HandlerFunc {
|
||||||
|
allowedOrigins := make(map[string]struct{}, len(configuredOrigins))
|
||||||
|
for _, origin := range configuredOrigins {
|
||||||
|
if normalized, ok := normalizeCORSOrigin(origin); ok {
|
||||||
|
allowedOrigins[normalized] = struct{}{}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
return func(c *gin.Context) {
|
return func(c *gin.Context) {
|
||||||
origin := strings.TrimSpace(c.GetHeader("Origin"))
|
origin := strings.TrimSpace(c.GetHeader("Origin"))
|
||||||
if origin != "" {
|
if origin != "" {
|
||||||
parsed, err := url.Parse(origin)
|
c.Writer.Header().Add("Vary", "Origin")
|
||||||
if err != nil || parsed.Host == "" || !strings.EqualFold(parsed.Host, c.Request.Host) {
|
normalized, valid := normalizeCORSOrigin(origin)
|
||||||
|
_, explicitlyAllowed := allowedOrigins[normalized]
|
||||||
|
parsed, _ := url.Parse(origin)
|
||||||
|
sameHost := valid && strings.EqualFold(parsed.Host, c.Request.Host)
|
||||||
|
browserExtension := valid && isChromiumExtensionOrigin(parsed)
|
||||||
|
if !sameHost && !browserExtension && !explicitlyAllowed {
|
||||||
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error": "cross-origin request denied"})
|
c.AbortWithStatusJSON(http.StatusForbidden, gin.H{"error": "cross-origin request denied"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
c.Writer.Header().Set("Access-Control-Allow-Origin", origin)
|
c.Writer.Header().Set("Access-Control-Allow-Origin", origin)
|
||||||
c.Writer.Header().Set("Access-Control-Allow-Credentials", "true")
|
c.Writer.Header().Set("Access-Control-Allow-Credentials", "true")
|
||||||
c.Writer.Header().Add("Vary", "Origin")
|
|
||||||
}
|
}
|
||||||
c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Content-Length, Accept-Encoding, X-CSRF-Token, Authorization, accept, origin, Cache-Control, X-Requested-With")
|
c.Writer.Header().Set("Access-Control-Allow-Headers", "Content-Type, Content-Length, Accept-Encoding, X-CSRF-Token, Authorization, accept, origin, Cache-Control, X-Requested-With")
|
||||||
c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET, PUT, DELETE")
|
c.Writer.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS, GET, PUT, DELETE")
|
||||||
|
c.Writer.Header().Set("Access-Control-Max-Age", "600")
|
||||||
|
|
||||||
if c.Request.Method == "OPTIONS" {
|
if c.Request.Method == "OPTIONS" {
|
||||||
c.AbortWithStatus(204)
|
c.AbortWithStatus(204)
|
||||||
@@ -2071,3 +2108,37 @@ func corsMiddleware() gin.HandlerFunc {
|
|||||||
c.Next()
|
c.Next()
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// isChromiumExtensionOrigin accepts only Chrome's canonical 32-character
|
||||||
|
// extension IDs (letters a-p). It does not allow arbitrary custom schemes or
|
||||||
|
// web origins, and the extension must separately obtain host permission.
|
||||||
|
func isChromiumExtensionOrigin(origin *url.URL) bool {
|
||||||
|
if origin == nil || !strings.EqualFold(origin.Scheme, "chrome-extension") || origin.Port() != "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
id := strings.ToLower(origin.Hostname())
|
||||||
|
if len(id) != 32 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, ch := range id {
|
||||||
|
if ch < 'a' || ch > 'p' {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
// normalizeCORSOrigin validates and canonicalizes a serialized origin. CORS
|
||||||
|
// origins never contain credentials, paths, query strings, or fragments.
|
||||||
|
func normalizeCORSOrigin(raw string) (string, bool) {
|
||||||
|
raw = strings.TrimSpace(raw)
|
||||||
|
if raw == "" || raw == "*" || strings.EqualFold(raw, "null") {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
parsed, err := url.Parse(raw)
|
||||||
|
if err != nil || parsed.Scheme == "" || parsed.Host == "" || parsed.User != nil ||
|
||||||
|
(parsed.Path != "" && parsed.Path != "/") || parsed.RawQuery != "" || parsed.Fragment != "" {
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
return strings.ToLower(parsed.Scheme) + "://" + strings.ToLower(parsed.Host), true
|
||||||
|
}
|
||||||
|
|||||||
@@ -31,7 +31,7 @@ func registerC2Tools(mcpServer *mcp.Server, c2Manager *c2.Manager, logger *zap.L
|
|||||||
registerC2EventTool(mcpServer, c2Manager, logger)
|
registerC2EventTool(mcpServer, c2Manager, logger)
|
||||||
registerC2ProfileTool(mcpServer, c2Manager, logger)
|
registerC2ProfileTool(mcpServer, c2Manager, logger)
|
||||||
registerC2FileTool(mcpServer, c2Manager, logger)
|
registerC2FileTool(mcpServer, c2Manager, logger)
|
||||||
logger.Info("C2 MCP tools registered (8 unified tools)")
|
logger.Debug("C2 MCP tools registered (8 unified tools)")
|
||||||
}
|
}
|
||||||
|
|
||||||
func makeC2Result(data interface{}, err error) (*mcp.ToolResult, error) {
|
func makeC2Result(data interface{}, err error) (*mcp.ToolResult, error) {
|
||||||
|
|||||||
@@ -11,7 +11,7 @@ import (
|
|||||||
func TestCORSMiddlewareAllowsSameOriginAndRejectsForeignOrigin(t *testing.T) {
|
func TestCORSMiddlewareAllowsSameOriginAndRejectsForeignOrigin(t *testing.T) {
|
||||||
gin.SetMode(gin.TestMode)
|
gin.SetMode(gin.TestMode)
|
||||||
router := gin.New()
|
router := gin.New()
|
||||||
router.Use(corsMiddleware())
|
router.Use(corsMiddleware(nil))
|
||||||
router.GET("/test", func(c *gin.Context) { c.Status(http.StatusNoContent) })
|
router.GET("/test", func(c *gin.Context) { c.Status(http.StatusNoContent) })
|
||||||
|
|
||||||
same := httptest.NewRequest(http.MethodGet, "http://app.example/test", nil)
|
same := httptest.NewRequest(http.MethodGet, "http://app.example/test", nil)
|
||||||
@@ -32,3 +32,74 @@ func TestCORSMiddlewareAllowsSameOriginAndRejectsForeignOrigin(t *testing.T) {
|
|||||||
t.Fatalf("foreign-origin response = %d, want %d", foreignW.Code, http.StatusForbidden)
|
t.Fatalf("foreign-origin response = %d, want %d", foreignW.Code, http.StatusForbidden)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestCORSMiddlewareAllowsBrowserExtensionWithoutConfiguration(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(corsMiddleware(nil))
|
||||||
|
router.POST("/api/auth/login", func(c *gin.Context) { c.Status(http.StatusNoContent) })
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodOptions, "https://server.example/api/auth/login", nil)
|
||||||
|
req.Host = "server.example"
|
||||||
|
req.Header.Set("Origin", "chrome-extension://abcdefghijklmnopabcdefghijklmnop")
|
||||||
|
req.Header.Set("Access-Control-Request-Method", http.MethodPost)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(w, req)
|
||||||
|
|
||||||
|
if w.Code != http.StatusNoContent {
|
||||||
|
t.Fatalf("preflight response = %d, want %d", w.Code, http.StatusNoContent)
|
||||||
|
}
|
||||||
|
if got := w.Header().Get("Access-Control-Allow-Origin"); got != "chrome-extension://abcdefghijklmnopabcdefghijklmnop" {
|
||||||
|
t.Fatalf("allow-origin = %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCORSMiddlewareRejectsInvalidExtensionOrigins(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
for _, origin := range []string{
|
||||||
|
"chrome-extension://too-short",
|
||||||
|
"chrome-extension://qrstuvwxyzabcdefqrstuvwxyzabcdef",
|
||||||
|
"chrome-extension://abcdefghijklmnopabcdefghijklmnop:8443",
|
||||||
|
"moz-extension://abcdefghijklmnopabcdefghijklmnop",
|
||||||
|
} {
|
||||||
|
t.Run(origin, func(t *testing.T) {
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(corsMiddleware(nil))
|
||||||
|
router.GET("/test", func(c *gin.Context) { c.Status(http.StatusNoContent) })
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "https://server.example/test", nil)
|
||||||
|
req.Host = "server.example"
|
||||||
|
req.Header.Set("Origin", origin)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(w, req)
|
||||||
|
if w.Code != http.StatusForbidden {
|
||||||
|
t.Fatalf("response = %d, want %d", w.Code, http.StatusForbidden)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCORSMiddlewareRejectsUnsafeConfiguredEntries(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
for _, configured := range []string{
|
||||||
|
"*",
|
||||||
|
"null",
|
||||||
|
"https://trusted.example/extra",
|
||||||
|
"https://trusted.example?trusted=true",
|
||||||
|
} {
|
||||||
|
t.Run(configured, func(t *testing.T) {
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(corsMiddleware([]string{configured}))
|
||||||
|
router.GET("/test", func(c *gin.Context) { c.Status(http.StatusNoContent) })
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "https://server.example/test", nil)
|
||||||
|
req.Host = "server.example"
|
||||||
|
req.Header.Set("Origin", "https://trusted.example")
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(w, req)
|
||||||
|
if w.Code != http.StatusForbidden {
|
||||||
|
t.Fatalf("response = %d, want %d", w.Code, http.StatusForbidden)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -21,11 +21,15 @@ func TestStandaloneMCPPrefersUserRBACAndDisablesGlobalTokenByDefault(t *testing.
|
|||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
t.Cleanup(func() { _ = db.Close() })
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
auth, err := security.NewAuthManager("admin-secret", 12)
|
auth := security.NewAuthManager(12)
|
||||||
|
if _, err := auth.AttachRBACStore(db); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
hash, err := security.HashPassword("admin-secret")
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if err := auth.AttachRBACStore(db); err != nil {
|
if err := db.UpdateRBACAdminPassword(hash); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
token, _, err := auth.Authenticate("admin", "admin-secret")
|
token, _, err := auth.Authenticate("admin", "admin-secret")
|
||||||
|
|||||||
@@ -90,7 +90,7 @@ func registerProjectFactTools(mcpServer *mcp.Server, db *database.DB, cfg *confi
|
|||||||
"description": "可选:关联的漏洞记录 ID",
|
"description": "可选:关联的漏洞记录 ID",
|
||||||
},
|
},
|
||||||
"links": map[string]interface{}{
|
"links": map[string]interface{}{
|
||||||
"type": "array",
|
"type": "array",
|
||||||
"description": "可选:关系边(from → 当前 fact)。finding 至少 1 条 {from:target/*, type:discovered_on};finding 上记录 exploit 用 {from:exploit/*, type:exploits}。省略保留已有边;传 [] 清空全部关系边。",
|
"description": "可选:关系边(from → 当前 fact)。finding 至少 1 条 {from:target/*, type:discovered_on};finding 上记录 exploit 用 {from:exploit/*, type:exploits}。省略保留已有边;传 [] 清空全部关系边。",
|
||||||
"items": map[string]interface{}{
|
"items": map[string]interface{}{
|
||||||
"type": "object",
|
"type": "object",
|
||||||
@@ -357,7 +357,7 @@ func registerProjectFactTools(mcpServer *mcp.Server, db *database.DB, cfg *confi
|
|||||||
})
|
})
|
||||||
|
|
||||||
if logger != nil {
|
if logger != nil {
|
||||||
logger.Info("项目黑板 MCP 工具注册成功")
|
logger.Debug("项目黑板 MCP 工具注册成功")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -190,7 +190,7 @@ func registerVulnerabilityTools(mcpServer *mcp.Server, db *database.DB, logger *
|
|||||||
registerListVulnerabilitiesTool(mcpServer, db, logger)
|
registerListVulnerabilitiesTool(mcpServer, db, logger)
|
||||||
registerGetVulnerabilityTool(mcpServer, db, logger)
|
registerGetVulnerabilityTool(mcpServer, db, logger)
|
||||||
if logger != nil {
|
if logger != nil {
|
||||||
logger.Info("漏洞 MCP 工具注册成功", zap.Strings("tools", []string{
|
logger.Debug("漏洞 MCP 工具注册成功", zap.Strings("tools", []string{
|
||||||
builtin.ToolRecordVulnerability,
|
builtin.ToolRecordVulnerability,
|
||||||
builtin.ToolListVulnerabilities,
|
builtin.ToolListVulnerabilities,
|
||||||
builtin.ToolGetVulnerability,
|
builtin.ToolGetVulnerability,
|
||||||
@@ -201,7 +201,7 @@ func registerVulnerabilityTools(mcpServer *mcp.Server, db *database.DB, logger *
|
|||||||
func registerRecordVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, logger *zap.Logger) {
|
func registerRecordVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, logger *zap.Logger) {
|
||||||
tool := mcp.Tool{
|
tool := mcp.Tool{
|
||||||
Name: builtin.ToolRecordVulnerability,
|
Name: builtin.ToolRecordVulnerability,
|
||||||
Description: "记录发现的漏洞详情到漏洞管理系统。必须按“仅看本记录即可复现”的标准填写:目标、触发点、前置条件、复现步骤、证据/POC、实际影响、修复建议和复测方式。边渗透边记录:每验证出一条可复现漏洞后立即调用,勿等会话结束。记录前可先 list_vulnerabilities 避免重复。",
|
Description: "记录发现的漏洞详情到漏洞管理系统。必须按“仅看本记录即可复现”的标准填写:目标、漏洞类型、触发点、复现步骤、证据/POC、实际影响和修复建议;前置条件与复测方式为推荐填写项。边渗透边记录:每验证出一条可复现漏洞后立即调用,勿等会话结束。记录前可先 list_vulnerabilities 避免重复。",
|
||||||
ShortDescription: "记录可复现的漏洞详情到漏洞管理系统",
|
ShortDescription: "记录可复现的漏洞详情到漏洞管理系统",
|
||||||
InputSchema: map[string]interface{}{
|
InputSchema: map[string]interface{}{
|
||||||
"type": "object",
|
"type": "object",
|
||||||
@@ -229,7 +229,7 @@ func registerRecordVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, log
|
|||||||
},
|
},
|
||||||
"preconditions": map[string]interface{}{
|
"preconditions": map[string]interface{}{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "前置条件:登录状态、权限、账号、Header/Cookie、特定数据、网络位置、环境/版本等;无前置条件写“无”。",
|
"description": "前置条件(推荐填写):登录状态、权限、账号、Header/Cookie、特定数据、网络位置、环境/版本等;无前置条件可写“无”。",
|
||||||
},
|
},
|
||||||
"reproduction_steps": map[string]interface{}{
|
"reproduction_steps": map[string]interface{}{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
@@ -249,7 +249,7 @@ func registerRecordVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, log
|
|||||||
},
|
},
|
||||||
"retest_notes": map[string]interface{}{
|
"retest_notes": map[string]interface{}{
|
||||||
"type": "string",
|
"type": "string",
|
||||||
"description": "复测方式:修复后如何验证漏洞已关闭,包括应返回的状态码、错误信息或访问控制结果。",
|
"description": "复测方式(推荐填写):修复后如何验证漏洞已关闭,包括应返回的状态码、错误信息或访问控制结果。",
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
"required": []string{"title", "description", "severity", "vulnerability_type", "target", "reproduction_steps", "evidence", "impact", "recommendation"},
|
"required": []string{"title", "description", "severity", "vulnerability_type", "target", "reproduction_steps", "evidence", "impact", "recommendation"},
|
||||||
@@ -282,7 +282,7 @@ func registerRecordVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, log
|
|||||||
return textResult(fmt.Sprintf("错误: severity 必须是 critical、high、medium、low 或 info 之一,当前值: %s", severity), true), nil
|
return textResult(fmt.Sprintf("错误: severity 必须是 critical、high、medium、low 或 info 之一,当前值: %s", severity), true), nil
|
||||||
}
|
}
|
||||||
if missing := missingVulnerabilityReproFields(args); len(missing) > 0 {
|
if missing := missingVulnerabilityReproFields(args); len(missing) > 0 {
|
||||||
return textResult("错误: 漏洞记录缺少复现所需信息,请补充后再记录:\n- "+strings.Join(missing, "\n- ")+"\n\n最佳实践:漏洞管理中的单条记录应独立包含目标、前置条件、复现步骤、证据/POC、影响和修复/复测方式。", true), nil
|
return textResult("错误: 漏洞记录缺少必填信息,请补充后再记录:\n- "+strings.Join(missing, "\n- ")+"\n\n必填项用于确保单条记录可独立复现;前置条件和复测方式为推荐填写项。", true), nil
|
||||||
}
|
}
|
||||||
|
|
||||||
projectID := ""
|
projectID := ""
|
||||||
@@ -318,6 +318,7 @@ func registerRecordVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, log
|
|||||||
_ = db.SetResourceOwner("vulnerability", created.ID, principal.UserID)
|
_ = db.SetResourceOwner("vulnerability", created.ID, principal.UserID)
|
||||||
_ = db.AssignResourceToUser(principal.UserID, "vulnerability", created.ID)
|
_ = db.AssignResourceToUser(principal.UserID, "vulnerability", created.ID)
|
||||||
}
|
}
|
||||||
|
db.NotifyVulnerabilityCreated(created)
|
||||||
|
|
||||||
if logger != nil {
|
if logger != nil {
|
||||||
logger.Info("漏洞记录成功",
|
logger.Info("漏洞记录成功",
|
||||||
|
|||||||
+23
-198
@@ -2,7 +2,6 @@ package config
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"crypto/rand"
|
"crypto/rand"
|
||||||
"encoding/base64"
|
|
||||||
"encoding/hex"
|
"encoding/hex"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
@@ -12,6 +11,8 @@ import (
|
|||||||
"strconv"
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
||||||
|
"cyberstrike-ai/internal/termout"
|
||||||
|
|
||||||
"gopkg.in/yaml.v3"
|
"gopkg.in/yaml.v3"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -43,9 +44,8 @@ type Config struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type EnsureLocalConfigResult struct {
|
type EnsureLocalConfigResult struct {
|
||||||
Created bool
|
Created bool
|
||||||
GeneratedPassword string
|
ExamplePath string
|
||||||
ExamplePath string
|
|
||||||
}
|
}
|
||||||
|
|
||||||
const (
|
const (
|
||||||
@@ -54,6 +54,7 @@ const (
|
|||||||
DefaultLatestUserMessageMaxRunes = 48000
|
DefaultLatestUserMessageMaxRunes = 48000
|
||||||
DefaultLatestUserMessageHeadRunes = 24000
|
DefaultLatestUserMessageHeadRunes = 24000
|
||||||
DefaultLatestUserMessageTailRunes = 24000
|
DefaultLatestUserMessageTailRunes = 24000
|
||||||
|
DefaultSummarizationOutputReserveTokens = 8192
|
||||||
)
|
)
|
||||||
|
|
||||||
// ProjectConfig 项目黑板(跨对话共享事实)配置。
|
// ProjectConfig 项目黑板(跨对话共享事实)配置。
|
||||||
@@ -268,6 +269,8 @@ type MultiAgentEinoMiddlewareConfig struct {
|
|||||||
ReductionSubAgents bool `yaml:"reduction_sub_agents,omitempty" json:"reduction_sub_agents,omitempty"` // also attach to sub-agents
|
ReductionSubAgents bool `yaml:"reduction_sub_agents,omitempty" json:"reduction_sub_agents,omitempty"` // also attach to sub-agents
|
||||||
// SummarizationTriggerRatio controls summarization trigger threshold as max_total_tokens * ratio (default 0.8).
|
// SummarizationTriggerRatio controls summarization trigger threshold as max_total_tokens * ratio (default 0.8).
|
||||||
SummarizationTriggerRatio float64 `yaml:"summarization_trigger_ratio,omitempty" json:"summarization_trigger_ratio,omitempty"`
|
SummarizationTriggerRatio float64 `yaml:"summarization_trigger_ratio,omitempty" json:"summarization_trigger_ratio,omitempty"`
|
||||||
|
// SummarizationOutputReserveTokens reserves completion headroom for the summarization model call (default 8192).
|
||||||
|
SummarizationOutputReserveTokens int `yaml:"summarization_output_reserve_tokens,omitempty" json:"summarization_output_reserve_tokens,omitempty"`
|
||||||
// SummarizationEmitInternalEvents controls middleware internal event emission (default true).
|
// SummarizationEmitInternalEvents controls middleware internal event emission (default true).
|
||||||
SummarizationEmitInternalEvents *bool `yaml:"summarization_emit_internal_events,omitempty" json:"summarization_emit_internal_events,omitempty"`
|
SummarizationEmitInternalEvents *bool `yaml:"summarization_emit_internal_events,omitempty" json:"summarization_emit_internal_events,omitempty"`
|
||||||
// SummarizationUserIntentLedgerMaxRunes caps the DB-backed immutable user input ledger injected into model context.
|
// SummarizationUserIntentLedgerMaxRunes caps the DB-backed immutable user input ledger injected into model context.
|
||||||
@@ -320,6 +323,13 @@ func (c MultiAgentEinoMiddlewareConfig) SummarizationTriggerRatioEffective() flo
|
|||||||
return v
|
return v
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (c MultiAgentEinoMiddlewareConfig) SummarizationOutputReserveTokensEffective() int {
|
||||||
|
if c.SummarizationOutputReserveTokens > 0 {
|
||||||
|
return c.SummarizationOutputReserveTokens
|
||||||
|
}
|
||||||
|
return DefaultSummarizationOutputReserveTokens
|
||||||
|
}
|
||||||
|
|
||||||
func (c MultiAgentEinoMiddlewareConfig) SummarizationEmitInternalEventsEffective() bool {
|
func (c MultiAgentEinoMiddlewareConfig) SummarizationEmitInternalEventsEffective() bool {
|
||||||
if c.SummarizationEmitInternalEvents != nil {
|
if c.SummarizationEmitInternalEvents != nil {
|
||||||
return *c.SummarizationEmitInternalEvents
|
return *c.SummarizationEmitInternalEvents
|
||||||
@@ -757,6 +767,9 @@ func (c RobotsConfig) ServiceAccountUserIDs() map[string]string {
|
|||||||
type ServerConfig struct {
|
type ServerConfig struct {
|
||||||
Host string `yaml:"host" json:"host"`
|
Host string `yaml:"host" json:"host"`
|
||||||
Port int `yaml:"port" json:"port"`
|
Port int `yaml:"port" json:"port"`
|
||||||
|
// CORSAllowedOrigins contains additional, exact origins that may call the API.
|
||||||
|
// Same-origin browser requests are always allowed. Wildcards are intentionally unsupported.
|
||||||
|
CORSAllowedOrigins []string `yaml:"cors_allowed_origins,omitempty" json:"cors_allowed_origins,omitempty"`
|
||||||
// TLSEnabled 为 true 时主 Web UI 使用 HTTPS;现代浏览器在同源下会协商 HTTP/2,缓解 HTTP/1.1 每源并发连接数限制。
|
// TLSEnabled 为 true 时主 Web UI 使用 HTTPS;现代浏览器在同源下会协商 HTTP/2,缓解 HTTP/1.1 每源并发连接数限制。
|
||||||
TLSEnabled bool `yaml:"tls_enabled,omitempty" json:"tls_enabled,omitempty"`
|
TLSEnabled bool `yaml:"tls_enabled,omitempty" json:"tls_enabled,omitempty"`
|
||||||
// TLSCertPath / TLSKeyPath 非空时从 PEM 文件加载证书(生产环境推荐)。
|
// TLSCertPath / TLSKeyPath 非空时从 PEM 文件加载证书(生产环境推荐)。
|
||||||
@@ -995,11 +1008,7 @@ func normalizeHitlModeForPrompt(mode string) string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
type AuthConfig struct {
|
type AuthConfig struct {
|
||||||
Password string `yaml:"password" json:"password"`
|
SessionDurationHours int `yaml:"session_duration_hours" json:"session_duration_hours"`
|
||||||
SessionDurationHours int `yaml:"session_duration_hours" json:"session_duration_hours"`
|
|
||||||
GeneratedPassword string `yaml:"-" json:"-"`
|
|
||||||
GeneratedPasswordPersisted bool `yaml:"-" json:"-"`
|
|
||||||
GeneratedPasswordPersistErr string `yaml:"-" json:"-"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// MonitorConfig MCP 状态监控(tool_executions)保留策略。
|
// MonitorConfig MCP 状态监控(tool_executions)保留策略。
|
||||||
@@ -1159,23 +1168,6 @@ func Load(path string) (*Config, error) {
|
|||||||
if cfg.Audit.MaxDetailBytes <= 0 {
|
if cfg.Audit.MaxDetailBytes <= 0 {
|
||||||
cfg.Audit.MaxDetailBytes = 8192
|
cfg.Audit.MaxDetailBytes = 8192
|
||||||
}
|
}
|
||||||
if strings.TrimSpace(cfg.Auth.Password) == "" {
|
|
||||||
password, err := generateStrongPassword(24)
|
|
||||||
if err != nil {
|
|
||||||
return nil, fmt.Errorf("生成默认密码失败: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
cfg.Auth.Password = password
|
|
||||||
cfg.Auth.GeneratedPassword = password
|
|
||||||
|
|
||||||
if err := PersistAuthPassword(path, password); err != nil {
|
|
||||||
cfg.Auth.GeneratedPasswordPersisted = false
|
|
||||||
cfg.Auth.GeneratedPasswordPersistErr = err.Error()
|
|
||||||
} else {
|
|
||||||
cfg.Auth.GeneratedPasswordPersisted = true
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
// 如果配置了工具目录,从目录加载工具配置
|
// 如果配置了工具目录,从目录加载工具配置
|
||||||
if cfg.Security.ToolsDir != "" {
|
if cfg.Security.ToolsDir != "" {
|
||||||
inlineTools := append([]ToolConfig(nil), cfg.Security.Tools...)
|
inlineTools := append([]ToolConfig(nil), cfg.Security.Tools...)
|
||||||
@@ -1246,7 +1238,7 @@ func EnsureLocalConfig(path string) (EnsureLocalConfigResult, error) {
|
|||||||
|
|
||||||
if _, err := os.Stat(path); err == nil {
|
if _, err := os.Stat(path); err == nil {
|
||||||
return EnsureLocalConfigResult{}, nil
|
return EnsureLocalConfigResult{}, nil
|
||||||
} else if err != nil && !os.IsNotExist(err) {
|
} else if !os.IsNotExist(err) {
|
||||||
return EnsureLocalConfigResult{}, fmt.Errorf("检查配置文件失败: %w", err)
|
return EnsureLocalConfigResult{}, fmt.Errorf("检查配置文件失败: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -1281,181 +1273,14 @@ func EnsureLocalConfig(path string) (EnsureLocalConfigResult, error) {
|
|||||||
return EnsureLocalConfigResult{}, fmt.Errorf("创建配置文件失败: %w", err)
|
return EnsureLocalConfigResult{}, fmt.Errorf("创建配置文件失败: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
password, err := generateStrongPassword(24)
|
|
||||||
if err != nil {
|
|
||||||
return EnsureLocalConfigResult{}, fmt.Errorf("生成默认密码失败: %w", err)
|
|
||||||
}
|
|
||||||
if err := PersistAuthPassword(path, password); err != nil {
|
|
||||||
return EnsureLocalConfigResult{}, fmt.Errorf("写入默认密码失败: %w", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
return EnsureLocalConfigResult{
|
return EnsureLocalConfigResult{
|
||||||
Created: true,
|
Created: true,
|
||||||
GeneratedPassword: password,
|
ExamplePath: examplePath,
|
||||||
ExamplePath: examplePath,
|
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func generateStrongPassword(length int) (string, error) {
|
func PrintBootstrapAdminPassword(password string) {
|
||||||
if length <= 0 {
|
termout.PrintBootstrapAdminCredentials(password)
|
||||||
length = 24
|
|
||||||
}
|
|
||||||
|
|
||||||
bytesLen := length
|
|
||||||
randomBytes := make([]byte, bytesLen)
|
|
||||||
if _, err := rand.Read(randomBytes); err != nil {
|
|
||||||
return "", err
|
|
||||||
}
|
|
||||||
|
|
||||||
password := base64.RawURLEncoding.EncodeToString(randomBytes)
|
|
||||||
if len(password) > length {
|
|
||||||
password = password[:length]
|
|
||||||
}
|
|
||||||
return password, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func PersistAuthPassword(path, password string) error {
|
|
||||||
data, err := os.ReadFile(path)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
|
|
||||||
lines := strings.Split(string(data), "\n")
|
|
||||||
inAuthBlock := false
|
|
||||||
authIndent := -1
|
|
||||||
|
|
||||||
for i, line := range lines {
|
|
||||||
trimmed := strings.TrimSpace(line)
|
|
||||||
if !inAuthBlock {
|
|
||||||
if strings.HasPrefix(trimmed, "auth:") {
|
|
||||||
inAuthBlock = true
|
|
||||||
authIndent = len(line) - len(strings.TrimLeft(line, " "))
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if trimmed == "" || strings.HasPrefix(trimmed, "#") {
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
leadingSpaces := len(line) - len(strings.TrimLeft(line, " "))
|
|
||||||
if leadingSpaces <= authIndent {
|
|
||||||
// 离开 auth 块
|
|
||||||
inAuthBlock = false
|
|
||||||
authIndent = -1
|
|
||||||
// 继续寻找其它 auth 块(理论上没有)
|
|
||||||
if strings.HasPrefix(trimmed, "auth:") {
|
|
||||||
inAuthBlock = true
|
|
||||||
authIndent = leadingSpaces
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
if strings.HasPrefix(strings.TrimSpace(line), "password:") {
|
|
||||||
prefix := line[:len(line)-len(strings.TrimLeft(line, " "))]
|
|
||||||
comment := ""
|
|
||||||
if idx := yamlLineCommentIndex(line); idx >= 0 {
|
|
||||||
comment = strings.TrimRight(line[idx:], " ")
|
|
||||||
}
|
|
||||||
|
|
||||||
newLine := fmt.Sprintf("%spassword: %s", prefix, quoteYAMLString(password))
|
|
||||||
if comment != "" {
|
|
||||||
if !strings.HasPrefix(comment, " ") {
|
|
||||||
newLine += " "
|
|
||||||
}
|
|
||||||
newLine += comment
|
|
||||||
}
|
|
||||||
lines[i] = newLine
|
|
||||||
break
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return os.WriteFile(path, []byte(strings.Join(lines, "\n")), 0644)
|
|
||||||
}
|
|
||||||
|
|
||||||
func quoteYAMLString(value string) string {
|
|
||||||
node := yaml.Node{
|
|
||||||
Kind: yaml.ScalarNode,
|
|
||||||
Tag: "!!str",
|
|
||||||
Style: yaml.DoubleQuotedStyle,
|
|
||||||
Value: value,
|
|
||||||
}
|
|
||||||
data, err := yaml.Marshal(&node)
|
|
||||||
if err != nil {
|
|
||||||
return strconv.Quote(value)
|
|
||||||
}
|
|
||||||
return strings.TrimSuffix(string(data), "\n")
|
|
||||||
}
|
|
||||||
|
|
||||||
func yamlLineCommentIndex(line string) int {
|
|
||||||
inSingleQuote := false
|
|
||||||
inDoubleQuote := false
|
|
||||||
escaped := false
|
|
||||||
|
|
||||||
for i, r := range line {
|
|
||||||
if inDoubleQuote {
|
|
||||||
if escaped {
|
|
||||||
escaped = false
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if r == '\\' {
|
|
||||||
escaped = true
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if r == '"' {
|
|
||||||
inDoubleQuote = false
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if inSingleQuote {
|
|
||||||
if r == '\'' {
|
|
||||||
inSingleQuote = false
|
|
||||||
}
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
|
|
||||||
switch r {
|
|
||||||
case '"':
|
|
||||||
inDoubleQuote = true
|
|
||||||
case '\'':
|
|
||||||
inSingleQuote = true
|
|
||||||
case '#':
|
|
||||||
if i == 0 || isYAMLWhitespace(line[i-1]) {
|
|
||||||
return i
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return -1
|
|
||||||
}
|
|
||||||
|
|
||||||
func isYAMLWhitespace(b byte) bool {
|
|
||||||
return b == ' ' || b == '\t'
|
|
||||||
}
|
|
||||||
|
|
||||||
func PrintGeneratedPasswordWarning(password string, persisted bool, persistErr string) {
|
|
||||||
if strings.TrimSpace(password) == "" {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if persisted {
|
|
||||||
fmt.Println("[CyberStrikeAI] ✅ 已为您自动生成并写入 Web 登录密码。")
|
|
||||||
} else {
|
|
||||||
if persistErr != "" {
|
|
||||||
fmt.Printf("[CyberStrikeAI] ⚠️ 无法自动写入配置文件中的密码: %s\n", persistErr)
|
|
||||||
} else {
|
|
||||||
fmt.Println("[CyberStrikeAI] ⚠️ 无法自动写入配置文件中的密码。")
|
|
||||||
}
|
|
||||||
fmt.Println("请手动将以下随机密码写入 config.yaml 的 auth.password:")
|
|
||||||
}
|
|
||||||
|
|
||||||
fmt.Println("----------------------------------------------------------------")
|
|
||||||
fmt.Println("CyberStrikeAI Auto-Generated Web Password")
|
|
||||||
fmt.Printf("Password: %s\n", password)
|
|
||||||
fmt.Println("WARNING: Anyone with this password can fully control CyberStrikeAI.")
|
|
||||||
fmt.Println("Please store it securely and change it in config.yaml as soon as possible.")
|
|
||||||
fmt.Println("警告:持有此密码的人将拥有对 CyberStrikeAI 的完全控制权限。")
|
|
||||||
fmt.Println("请妥善保管,并尽快在 config.yaml 中修改 auth.password!")
|
|
||||||
fmt.Println("----------------------------------------------------------------")
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// generateRandomToken 生成用于 MCP 鉴权的随机字符串(64 位十六进制)
|
// generateRandomToken 生成用于 MCP 鉴权的随机字符串(64 位十六进制)
|
||||||
|
|||||||
@@ -7,78 +7,12 @@ import (
|
|||||||
"testing"
|
"testing"
|
||||||
)
|
)
|
||||||
|
|
||||||
func TestPersistAuthPasswordQuotesYAMLSpecialCharacters(t *testing.T) {
|
func TestEnsureLocalConfigCreatesFromExample(t *testing.T) {
|
||||||
dir := t.TempDir()
|
|
||||||
path := filepath.Join(dir, "config.yaml")
|
|
||||||
initial := strings.Join([]string{
|
|
||||||
"server:",
|
|
||||||
" host: 0.0.0.0",
|
|
||||||
"auth:",
|
|
||||||
" password: old-password # Web 登录密码",
|
|
||||||
" session_duration_hours: 12",
|
|
||||||
"log:",
|
|
||||||
" level: info",
|
|
||||||
"",
|
|
||||||
}, "\n")
|
|
||||||
if err := os.WriteFile(path, []byte(initial), 0644); err != nil {
|
|
||||||
t.Fatalf("write config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
want := `@abc:def # still password`
|
|
||||||
if err := PersistAuthPassword(path, want); err != nil {
|
|
||||||
t.Fatalf("PersistAuthPassword: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
data, err := os.ReadFile(path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read config: %v", err)
|
|
||||||
}
|
|
||||||
if !strings.Contains(string(data), `password: "@abc:def # still password" # Web 登录密码`) {
|
|
||||||
t.Fatalf("password was not safely quoted or comment was not preserved:\n%s", data)
|
|
||||||
}
|
|
||||||
|
|
||||||
cfg, err := Load(path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("Load after PersistAuthPassword: %v", err)
|
|
||||||
}
|
|
||||||
if cfg.Auth.Password != want {
|
|
||||||
t.Fatalf("Auth.Password = %q, want %q", cfg.Auth.Password, want)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestPersistAuthPasswordDoesNotTreatQuotedHashAsComment(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
|
||||||
path := filepath.Join(dir, "config.yaml")
|
|
||||||
initial := strings.Join([]string{
|
|
||||||
"auth:",
|
|
||||||
` password: "old#password"`,
|
|
||||||
" session_duration_hours: 12",
|
|
||||||
"",
|
|
||||||
}, "\n")
|
|
||||||
if err := os.WriteFile(path, []byte(initial), 0644); err != nil {
|
|
||||||
t.Fatalf("write config: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := PersistAuthPassword(path, "new-password"); err != nil {
|
|
||||||
t.Fatalf("PersistAuthPassword: %v", err)
|
|
||||||
}
|
|
||||||
|
|
||||||
data, err := os.ReadFile(path)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatalf("read config: %v", err)
|
|
||||||
}
|
|
||||||
if strings.Contains(string(data), "#password") {
|
|
||||||
t.Fatalf("old quoted password fragment was incorrectly preserved as a comment:\n%s", data)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestEnsureLocalConfigCreatesFromExampleWithGeneratedPassword(t *testing.T) {
|
|
||||||
dir := t.TempDir()
|
dir := t.TempDir()
|
||||||
examplePath := filepath.Join(dir, "config.example.yaml")
|
examplePath := filepath.Join(dir, "config.example.yaml")
|
||||||
configPath := filepath.Join(dir, "config.yaml")
|
configPath := filepath.Join(dir, "config.yaml")
|
||||||
|
|
||||||
example := []byte(`auth:
|
example := []byte(`auth:
|
||||||
password: "change-me-use-a-long-random-password"
|
|
||||||
session_duration_hours: 12
|
session_duration_hours: 12
|
||||||
server:
|
server:
|
||||||
host: 127.0.0.1
|
host: 127.0.0.1
|
||||||
@@ -95,9 +29,6 @@ server:
|
|||||||
if !result.Created {
|
if !result.Created {
|
||||||
t.Fatal("Created = false, want true")
|
t.Fatal("Created = false, want true")
|
||||||
}
|
}
|
||||||
if result.GeneratedPassword == "" {
|
|
||||||
t.Fatal("GeneratedPassword is empty")
|
|
||||||
}
|
|
||||||
if result.ExamplePath != examplePath {
|
if result.ExamplePath != examplePath {
|
||||||
t.Fatalf("ExamplePath = %q, want %q", result.ExamplePath, examplePath)
|
t.Fatalf("ExamplePath = %q, want %q", result.ExamplePath, examplePath)
|
||||||
}
|
}
|
||||||
@@ -106,11 +37,8 @@ server:
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("Load generated config: %v", err)
|
t.Fatalf("Load generated config: %v", err)
|
||||||
}
|
}
|
||||||
if cfg.Auth.Password == "change-me-use-a-long-random-password" {
|
if cfg.Auth.SessionDurationHours != 12 {
|
||||||
t.Fatal("auth.password still contains the template placeholder")
|
t.Fatalf("SessionDurationHours = %d, want 12", cfg.Auth.SessionDurationHours)
|
||||||
}
|
|
||||||
if cfg.Auth.Password != result.GeneratedPassword {
|
|
||||||
t.Fatalf("Auth.Password = %q, want generated password %q", cfg.Auth.Password, result.GeneratedPassword)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
second, err := EnsureLocalConfig(configPath)
|
second, err := EnsureLocalConfig(configPath)
|
||||||
@@ -122,6 +50,31 @@ server:
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestLoadIgnoresLegacyAuthPasswordField(t *testing.T) {
|
||||||
|
dir := t.TempDir()
|
||||||
|
path := filepath.Join(dir, "config.yaml")
|
||||||
|
initial := strings.Join([]string{
|
||||||
|
"auth:",
|
||||||
|
` password: "legacy-password"`,
|
||||||
|
" session_duration_hours: 12",
|
||||||
|
"server:",
|
||||||
|
" host: 127.0.0.1",
|
||||||
|
" port: 8080",
|
||||||
|
"",
|
||||||
|
}, "\n")
|
||||||
|
if err := os.WriteFile(path, []byte(initial), 0644); err != nil {
|
||||||
|
t.Fatalf("write config: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
cfg, err := Load(path)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("Load: %v", err)
|
||||||
|
}
|
||||||
|
if cfg.Auth.SessionDurationHours != 12 {
|
||||||
|
t.Fatalf("SessionDurationHours = %d, want 12", cfg.Auth.SessionDurationHours)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestHitlAuditModelEffectiveFallsBackToMainConfig(t *testing.T) {
|
func TestHitlAuditModelEffectiveFallsBackToMainConfig(t *testing.T) {
|
||||||
main := OpenAIConfig{
|
main := OpenAIConfig{
|
||||||
Provider: "openai",
|
Provider: "openai",
|
||||||
@@ -163,6 +116,17 @@ func TestSummarizationUserIntentLedgerRunesEffective(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestSummarizationOutputReserveTokensEffective(t *testing.T) {
|
||||||
|
var zero MultiAgentEinoMiddlewareConfig
|
||||||
|
if got := zero.SummarizationOutputReserveTokensEffective(); got != DefaultSummarizationOutputReserveTokens {
|
||||||
|
t.Fatalf("default output reserve = %d, want %d", got, DefaultSummarizationOutputReserveTokens)
|
||||||
|
}
|
||||||
|
custom := MultiAgentEinoMiddlewareConfig{SummarizationOutputReserveTokens: 4096}
|
||||||
|
if got := custom.SummarizationOutputReserveTokensEffective(); got != 4096 {
|
||||||
|
t.Fatalf("custom output reserve = %d", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestLatestUserMessageRunesEffective(t *testing.T) {
|
func TestLatestUserMessageRunesEffective(t *testing.T) {
|
||||||
var zero MultiAgentEinoMiddlewareConfig
|
var zero MultiAgentEinoMiddlewareConfig
|
||||||
if got := zero.LatestUserMessageMaxRunesEffective(); got != DefaultLatestUserMessageMaxRunes {
|
if got := zero.LatestUserMessageMaxRunesEffective(); got != DefaultLatestUserMessageMaxRunes {
|
||||||
|
|||||||
+17
-12
@@ -28,18 +28,19 @@ type AuditLog struct {
|
|||||||
|
|
||||||
// ListAuditLogsFilter query parameters.
|
// ListAuditLogsFilter query parameters.
|
||||||
type ListAuditLogsFilter struct {
|
type ListAuditLogsFilter struct {
|
||||||
Actor string
|
Actor string
|
||||||
Level string
|
Level string
|
||||||
Category string
|
Category string
|
||||||
Action string
|
Action string
|
||||||
Result string
|
Result string
|
||||||
Query string
|
Query string
|
||||||
ResourceType string
|
ResourceType string
|
||||||
ResourceID string
|
ResourceID string
|
||||||
Since *time.Time
|
RelatedUserID string
|
||||||
Until *time.Time
|
Since *time.Time
|
||||||
Limit int
|
Until *time.Time
|
||||||
Offset int
|
Limit int
|
||||||
|
Offset int
|
||||||
}
|
}
|
||||||
|
|
||||||
func buildAuditLogsWhere(filter ListAuditLogsFilter) (string, []interface{}) {
|
func buildAuditLogsWhere(filter ListAuditLogsFilter) (string, []interface{}) {
|
||||||
@@ -73,6 +74,10 @@ func buildAuditLogsWhere(filter ListAuditLogsFilter) (string, []interface{}) {
|
|||||||
conditions = append(conditions, "resource_id = ?")
|
conditions = append(conditions, "resource_id = ?")
|
||||||
args = append(args, filter.ResourceID)
|
args = append(args, filter.ResourceID)
|
||||||
}
|
}
|
||||||
|
if relatedUserID := strings.TrimSpace(filter.RelatedUserID); relatedUserID != "" {
|
||||||
|
conditions = append(conditions, `(resource_id = ? OR detail_json LIKE ? OR detail_json LIKE ?)`)
|
||||||
|
args = append(args, relatedUserID, `%"user_id":"`+relatedUserID+`"%`, `%"userId":"`+relatedUserID+`"%`)
|
||||||
|
}
|
||||||
if filter.Since != nil {
|
if filter.Since != nil {
|
||||||
conditions = append(conditions, sqliteEpochGE("created_at", ">="))
|
conditions = append(conditions, sqliteEpochGE("created_at", ">="))
|
||||||
args = append(args, formatSQLiteUTC(*filter.Since))
|
args = append(args, formatSQLiteUTC(*filter.Since))
|
||||||
|
|||||||
@@ -31,6 +31,19 @@ func TestBuildAuditLogsWhere_timeFilterSQL(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestBuildAuditLogsWhere_relatedUserID(t *testing.T) {
|
||||||
|
where, args := buildAuditLogsWhere(ListAuditLogsFilter{Category: "rbac", RelatedUserID: "user-123"})
|
||||||
|
if !strings.Contains(where, "resource_id = ?") || !strings.Contains(where, "detail_json LIKE ?") {
|
||||||
|
t.Fatalf("expected related-user predicates, got %q", where)
|
||||||
|
}
|
||||||
|
if len(args) != 4 {
|
||||||
|
t.Fatalf("expected category plus 3 related-user args, got %#v", args)
|
||||||
|
}
|
||||||
|
if args[1] != "user-123" || args[2] != `%"user_id":"user-123"%` || args[3] != `%"userId":"user-123"%` {
|
||||||
|
t.Fatalf("unexpected related-user args: %#v", args)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestListAuditLogs_timeFilterMixedStorageFormats(t *testing.T) {
|
func TestListAuditLogs_timeFilterMixedStorageFormats(t *testing.T) {
|
||||||
root, err := os.Getwd()
|
root, err := os.Getwd()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -58,6 +58,7 @@ type DB struct {
|
|||||||
checkpointDone chan struct{}
|
checkpointDone chan struct{}
|
||||||
closeOnce sync.Once
|
closeOnce sync.Once
|
||||||
closeErr error
|
closeErr error
|
||||||
|
vulnerabilityCreatedHook func(*Vulnerability)
|
||||||
}
|
}
|
||||||
|
|
||||||
// startPassiveCheckpointLoop 启动后台 PASSIVE checkpoint 循环。
|
// startPassiveCheckpointLoop 启动后台 PASSIVE checkpoint 循环。
|
||||||
@@ -113,10 +114,10 @@ func (db *DB) runPassiveCheckpoint(trigger string) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
if busy > 0 {
|
if busy > 0 {
|
||||||
db.logger.Info("SQLite PASSIVE checkpoint 完成(部分推进)", fields...)
|
db.logger.Debug("SQLite PASSIVE checkpoint 完成(部分推进)", fields...)
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
db.logger.Info("SQLite PASSIVE checkpoint 完成(成功)", fields...)
|
db.logger.Debug("SQLite PASSIVE checkpoint 完成(成功)", fields...)
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewDB 创建数据库连接
|
// NewDB 创建数据库连接
|
||||||
@@ -325,6 +326,7 @@ func (db *DB) initTables() error {
|
|||||||
session_key TEXT PRIMARY KEY,
|
session_key TEXT PRIMARY KEY,
|
||||||
conversation_id TEXT NOT NULL,
|
conversation_id TEXT NOT NULL,
|
||||||
role_name TEXT NOT NULL DEFAULT '默认',
|
role_name TEXT NOT NULL DEFAULT '默认',
|
||||||
|
agent_mode TEXT NOT NULL DEFAULT 'eino_single',
|
||||||
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||||
FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE
|
FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE
|
||||||
);`
|
);`
|
||||||
@@ -403,6 +405,33 @@ func (db *DB) initTables() error {
|
|||||||
FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE SET NULL
|
FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE SET NULL
|
||||||
);`
|
);`
|
||||||
|
|
||||||
|
createVulnerabilityAlertSubscriptionsTable := `
|
||||||
|
CREATE TABLE IF NOT EXISTS vulnerability_alert_subscriptions (
|
||||||
|
user_id TEXT PRIMARY KEY,
|
||||||
|
enabled INTEGER NOT NULL DEFAULT 0,
|
||||||
|
min_severity TEXT NOT NULL DEFAULT 'high',
|
||||||
|
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
FOREIGN KEY (user_id) REFERENCES rbac_users(id) ON DELETE CASCADE
|
||||||
|
);`
|
||||||
|
createVulnerabilityAlertDeliveriesTable := `
|
||||||
|
CREATE TABLE IF NOT EXISTS vulnerability_alert_deliveries (
|
||||||
|
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||||
|
vulnerability_id TEXT NOT NULL,
|
||||||
|
user_id TEXT NOT NULL,
|
||||||
|
platform TEXT NOT NULL,
|
||||||
|
external_user_id TEXT NOT NULL,
|
||||||
|
status TEXT NOT NULL DEFAULT 'pending',
|
||||||
|
attempts INTEGER NOT NULL DEFAULT 0,
|
||||||
|
next_attempt_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
last_error TEXT NOT NULL DEFAULT '',
|
||||||
|
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||||
|
UNIQUE(vulnerability_id, platform, external_user_id),
|
||||||
|
FOREIGN KEY (vulnerability_id) REFERENCES vulnerabilities(id) ON DELETE CASCADE,
|
||||||
|
FOREIGN KEY (user_id) REFERENCES rbac_users(id) ON DELETE CASCADE
|
||||||
|
);`
|
||||||
|
|
||||||
// 创建批量任务队列表
|
// 创建批量任务队列表
|
||||||
createBatchTaskQueuesTable := `
|
createBatchTaskQueuesTable := `
|
||||||
CREATE TABLE IF NOT EXISTS batch_task_queues (
|
CREATE TABLE IF NOT EXISTS batch_task_queues (
|
||||||
@@ -638,6 +667,28 @@ func (db *DB) initTables() error {
|
|||||||
FOREIGN KEY (run_id) REFERENCES workflow_runs(id) ON DELETE CASCADE
|
FOREIGN KEY (run_id) REFERENCES workflow_runs(id) ON DELETE CASCADE
|
||||||
);`
|
);`
|
||||||
|
|
||||||
|
createWorkflowPackageInspectionsTable := `
|
||||||
|
CREATE TABLE IF NOT EXISTS workflow_package_inspections (
|
||||||
|
id TEXT PRIMARY KEY, package_hash TEXT NOT NULL, manifest_json TEXT NOT NULL,
|
||||||
|
workflow_payload_json TEXT NOT NULL, inspection_json TEXT NOT NULL,
|
||||||
|
source_workflow_id TEXT NOT NULL, source_revision INTEGER NOT NULL,
|
||||||
|
source_content_hash TEXT NOT NULL, source_graph_hash TEXT NOT NULL,
|
||||||
|
local_conflict_state TEXT NOT NULL CHECK (local_conflict_state IN ('none','identical','id_conflict')),
|
||||||
|
local_workflow_id TEXT, local_content_hash TEXT, local_graph_hash TEXT,
|
||||||
|
created_by TEXT NOT NULL, status TEXT NOT NULL DEFAULT 'ready' CHECK (status IN ('ready','consumed','expired')),
|
||||||
|
created_at DATETIME NOT NULL, expires_at DATETIME NOT NULL, consumed_at DATETIME
|
||||||
|
);`
|
||||||
|
createWorkflowPackageImportsTable := `
|
||||||
|
CREATE TABLE IF NOT EXISTS workflow_package_imports (
|
||||||
|
id TEXT PRIMARY KEY, inspection_id TEXT NOT NULL, request_hash TEXT NOT NULL,
|
||||||
|
idempotency_key TEXT NOT NULL, actor_user_id TEXT NOT NULL,
|
||||||
|
action TEXT NOT NULL CHECK (action IN ('create','keep_existing','overwrite','rename')),
|
||||||
|
source_workflow_id TEXT NOT NULL, target_workflow_id TEXT NOT NULL, resulting_workflow_id TEXT,
|
||||||
|
result TEXT NOT NULL CHECK (result IN ('created','overwritten','renamed','kept_existing','skipped_identical','failed')),
|
||||||
|
error_code TEXT, error_message TEXT, created_at DATETIME NOT NULL, applied_at DATETIME,
|
||||||
|
FOREIGN KEY (inspection_id) REFERENCES workflow_package_inspections(id)
|
||||||
|
);`
|
||||||
|
|
||||||
// 创建索引
|
// 创建索引
|
||||||
createIndexes := `
|
createIndexes := `
|
||||||
CREATE INDEX IF NOT EXISTS idx_messages_conversation_id ON messages(conversation_id);
|
CREATE INDEX IF NOT EXISTS idx_messages_conversation_id ON messages(conversation_id);
|
||||||
@@ -702,6 +753,9 @@ func (db *DB) initTables() error {
|
|||||||
CREATE INDEX IF NOT EXISTS idx_workflow_runs_conversation ON workflow_runs(conversation_id);
|
CREATE INDEX IF NOT EXISTS idx_workflow_runs_conversation ON workflow_runs(conversation_id);
|
||||||
CREATE INDEX IF NOT EXISTS idx_workflow_runs_status ON workflow_runs(status);
|
CREATE INDEX IF NOT EXISTS idx_workflow_runs_status ON workflow_runs(status);
|
||||||
CREATE INDEX IF NOT EXISTS idx_workflow_node_runs_run ON workflow_node_runs(run_id);
|
CREATE INDEX IF NOT EXISTS idx_workflow_node_runs_run ON workflow_node_runs(run_id);
|
||||||
|
CREATE INDEX IF NOT EXISTS idx_workflow_package_inspections_creator_expiry ON workflow_package_inspections(created_by, expires_at);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_workflow_package_imports_actor_key ON workflow_package_imports(actor_user_id, idempotency_key);
|
||||||
|
CREATE UNIQUE INDEX IF NOT EXISTS uq_workflow_package_imports_inspection_success ON workflow_package_imports(inspection_id) WHERE result IN ('created','overwritten','renamed','kept_existing','skipped_identical');
|
||||||
`
|
`
|
||||||
|
|
||||||
if _, err := db.Exec(createConversationsTable); err != nil {
|
if _, err := db.Exec(createConversationsTable); err != nil {
|
||||||
@@ -750,6 +804,9 @@ func (db *DB) initTables() error {
|
|||||||
if _, err := db.Exec(createRobotUserSessionsTable); err != nil {
|
if _, err := db.Exec(createRobotUserSessionsTable); err != nil {
|
||||||
return fmt.Errorf("创建robot_user_sessions表失败: %w", err)
|
return fmt.Errorf("创建robot_user_sessions表失败: %w", err)
|
||||||
}
|
}
|
||||||
|
if err := db.migrateRobotUserSessionsTable(); err != nil {
|
||||||
|
return fmt.Errorf("迁移robot_user_sessions表失败: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
if _, err := db.Exec(createProjectsTable); err != nil {
|
if _, err := db.Exec(createProjectsTable); err != nil {
|
||||||
return fmt.Errorf("创建projects表失败: %w", err)
|
return fmt.Errorf("创建projects表失败: %w", err)
|
||||||
@@ -790,11 +847,19 @@ func (db *DB) initTables() error {
|
|||||||
if err := db.initRBACTables(); err != nil {
|
if err := db.initRBACTables(); err != nil {
|
||||||
return fmt.Errorf("创建RBAC表失败: %w", err)
|
return fmt.Errorf("创建RBAC表失败: %w", err)
|
||||||
}
|
}
|
||||||
|
if _, err := db.Exec(createVulnerabilityAlertSubscriptionsTable); err != nil {
|
||||||
|
return fmt.Errorf("创建漏洞提醒订阅表失败: %w", err)
|
||||||
|
}
|
||||||
|
if _, err := db.Exec(createVulnerabilityAlertDeliveriesTable); err != nil {
|
||||||
|
return fmt.Errorf("创建漏洞提醒投递表失败: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
for tableName, ddl := range map[string]string{
|
for tableName, ddl := range map[string]string{
|
||||||
"workflow_definitions": createWorkflowDefinitionsTable,
|
"workflow_definitions": createWorkflowDefinitionsTable,
|
||||||
"workflow_runs": createWorkflowRunsTable,
|
"workflow_runs": createWorkflowRunsTable,
|
||||||
"workflow_node_runs": createWorkflowNodeRunsTable,
|
"workflow_node_runs": createWorkflowNodeRunsTable,
|
||||||
|
"workflow_package_inspections": createWorkflowPackageInspectionsTable,
|
||||||
|
"workflow_package_imports": createWorkflowPackageImportsTable,
|
||||||
} {
|
} {
|
||||||
if _, err := db.Exec(ddl); err != nil {
|
if _, err := db.Exec(ddl); err != nil {
|
||||||
return fmt.Errorf("创建%s表失败: %w", tableName, err)
|
return fmt.Errorf("创建%s表失败: %w", tableName, err)
|
||||||
@@ -869,7 +934,19 @@ func (db *DB) initTables() error {
|
|||||||
return fmt.Errorf("创建索引失败: %w", err)
|
return fmt.Errorf("创建索引失败: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
db.logger.Info("数据库表初始化完成")
|
db.logger.Debug("数据库表初始化完成")
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) migrateRobotUserSessionsTable() error {
|
||||||
|
var count int
|
||||||
|
if err := db.QueryRow("SELECT COUNT(*) FROM pragma_table_info('robot_user_sessions') WHERE name='agent_mode'").Scan(&count); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if count == 0 {
|
||||||
|
_, err := db.Exec("ALTER TABLE robot_user_sessions ADD COLUMN agent_mode TEXT NOT NULL DEFAULT 'eino_single'")
|
||||||
|
return err
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
+179
-10
@@ -227,11 +227,32 @@ func (db *DB) addColumnIfMissing(table, name, stmt string) error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// RBACNeedsAdminPassword reports whether the built-in admin account still needs an initial password.
|
||||||
|
func (db *DB) RBACNeedsAdminPassword() (bool, error) {
|
||||||
|
var userCount int
|
||||||
|
if err := db.QueryRow(`SELECT COUNT(*) FROM rbac_users`).Scan(&userCount); err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
if userCount == 0 {
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
|
var hash sql.NullString
|
||||||
|
err := db.QueryRow(`
|
||||||
|
SELECT password_hash FROM rbac_users
|
||||||
|
WHERE username = 'admin' AND is_builtin = 1
|
||||||
|
LIMIT 1
|
||||||
|
`).Scan(&hash)
|
||||||
|
if err == sql.ErrNoRows {
|
||||||
|
return false, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return false, err
|
||||||
|
}
|
||||||
|
return !hash.Valid || strings.TrimSpace(hash.String) == "", nil
|
||||||
|
}
|
||||||
|
|
||||||
// BootstrapRBAC seeds the local admin account and system roles.
|
// BootstrapRBAC seeds the local admin account and system roles.
|
||||||
func (db *DB) BootstrapRBAC(adminPasswordHash string, permissions map[string]string) error {
|
func (db *DB) BootstrapRBAC(adminPasswordHash string, permissions map[string]string) error {
|
||||||
if strings.TrimSpace(adminPasswordHash) == "" {
|
|
||||||
return errors.New("admin password hash is required")
|
|
||||||
}
|
|
||||||
now := time.Now()
|
now := time.Now()
|
||||||
tx, err := db.Begin()
|
tx, err := db.Begin()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -302,6 +323,9 @@ func (db *DB) BootstrapRBAC(adminPasswordHash string, permissions map[string]str
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if userCount == 0 {
|
if userCount == 0 {
|
||||||
|
if strings.TrimSpace(adminPasswordHash) == "" {
|
||||||
|
return errors.New("admin password hash is required for initial bootstrap")
|
||||||
|
}
|
||||||
if _, err := tx.Exec(`
|
if _, err := tx.Exec(`
|
||||||
INSERT INTO rbac_users (id, username, display_name, password_hash, enabled, is_builtin, created_at, updated_at)
|
INSERT INTO rbac_users (id, username, display_name, password_hash, enabled, is_builtin, created_at, updated_at)
|
||||||
VALUES (?, 'admin', '管理员', ?, 1, 1, ?, ?)
|
VALUES (?, 'admin', '管理员', ?, 1, 1, ?, ?)
|
||||||
@@ -311,7 +335,7 @@ func (db *DB) BootstrapRBAC(adminPasswordHash string, permissions map[string]str
|
|||||||
if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_user_roles (user_id, role_id, created_at) VALUES ('admin', ?, ?)`, RBACSystemRoleAdmin, now); err != nil {
|
if _, err := tx.Exec(`INSERT OR IGNORE INTO rbac_user_roles (user_id, role_id, created_at) VALUES ('admin', ?, ?)`, RBACSystemRoleAdmin, now); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
} else {
|
} else if strings.TrimSpace(adminPasswordHash) != "" {
|
||||||
if _, err := tx.Exec(`UPDATE rbac_users SET password_hash = ?, updated_at = ? WHERE username = 'admin' AND is_builtin = 1 AND (password_hash = '' OR password_hash IS NULL)`, adminPasswordHash, now); err != nil {
|
if _, err := tx.Exec(`UPDATE rbac_users SET password_hash = ?, updated_at = ? WHERE username = 'admin' AND is_builtin = 1 AND (password_hash = '' OR password_hash IS NULL)`, adminPasswordHash, now); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -740,6 +764,37 @@ func (db *DB) ListAssignableRBACResourcesPage(resourceType, search string, limit
|
|||||||
return options, rows.Err()
|
return options, rows.Err()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// CountAssignableRBACResources returns the total rows matching the resource picker filter.
|
||||||
|
func (db *DB) CountAssignableRBACResources(resourceType, search string) (int, error) {
|
||||||
|
resourceType = strings.TrimSpace(resourceType)
|
||||||
|
if _, ok := rbacAssignableResourceTables[resourceType]; !ok {
|
||||||
|
return 0, fmt.Errorf("不支持的资源类型: %s", resourceType)
|
||||||
|
}
|
||||||
|
pattern := "%" + strings.ToLower(strings.NewReplacer(
|
||||||
|
`\`, `\\`, `%`, `\%`, `_`, `\_`,
|
||||||
|
).Replace(strings.TrimSpace(search))) + "%"
|
||||||
|
var query string
|
||||||
|
switch resourceType {
|
||||||
|
case "project":
|
||||||
|
query = `SELECT COUNT(*) FROM projects WHERE LOWER(name) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'`
|
||||||
|
case "conversation":
|
||||||
|
query = `SELECT COUNT(*) FROM conversations WHERE LOWER(COALESCE(NULLIF(TRIM(title), ''), id)) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'`
|
||||||
|
case "vulnerability":
|
||||||
|
query = `SELECT COUNT(*) FROM vulnerabilities WHERE LOWER(title) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'`
|
||||||
|
case "webshell":
|
||||||
|
query = `SELECT COUNT(*) FROM webshell_connections WHERE LOWER(COALESCE(NULLIF(remark, ''), url)) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'`
|
||||||
|
case "batch_task":
|
||||||
|
query = `SELECT COUNT(*) FROM batch_task_queues WHERE LOWER(COALESCE(NULLIF(title, ''), id)) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'`
|
||||||
|
case "c2_listener":
|
||||||
|
query = `SELECT COUNT(*) FROM c2_listeners WHERE LOWER(name) LIKE ? ESCAPE '\' OR LOWER(id) LIKE ? ESCAPE '\'`
|
||||||
|
}
|
||||||
|
var total int
|
||||||
|
if err := db.QueryRow(query, pattern, pattern).Scan(&total); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
return total, nil
|
||||||
|
}
|
||||||
|
|
||||||
func normalizeRBACResourceLabel(label, id string) string {
|
func normalizeRBACResourceLabel(label, id string) string {
|
||||||
label = strings.TrimSpace(label)
|
label = strings.TrimSpace(label)
|
||||||
if label == "" {
|
if label == "" {
|
||||||
@@ -944,6 +999,86 @@ func (db *DB) AssignResourcesToUser(userID, resourceType string, resourceIDs []s
|
|||||||
return created, nil
|
return created, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// AssignResourcesToUserAuto detects each resource's actual type before writing.
|
||||||
|
// The whole batch is validated first and committed atomically.
|
||||||
|
func (db *DB) AssignResourcesToUserAuto(userID string, resourceIDs []string) (int64, map[string]string, error) {
|
||||||
|
userID = strings.TrimSpace(userID)
|
||||||
|
if userID == "" || len(resourceIDs) == 0 {
|
||||||
|
return 0, nil, errors.New("user_id and resource_ids are required")
|
||||||
|
}
|
||||||
|
if len(resourceIDs) > RBACMaxBatchResourceAssignments {
|
||||||
|
return 0, nil, fmt.Errorf("一次最多授权 %d 个资源", RBACMaxBatchResourceAssignments)
|
||||||
|
}
|
||||||
|
uniqueIDs := make([]string, 0, len(resourceIDs))
|
||||||
|
seen := make(map[string]struct{}, len(resourceIDs))
|
||||||
|
for _, rawID := range resourceIDs {
|
||||||
|
id := strings.TrimSpace(rawID)
|
||||||
|
if id == "" {
|
||||||
|
return 0, nil, errors.New("资源 ID 不能为空")
|
||||||
|
}
|
||||||
|
if _, exists := seen[id]; exists {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[id] = struct{}{}
|
||||||
|
uniqueIDs = append(uniqueIDs, id)
|
||||||
|
}
|
||||||
|
|
||||||
|
tx, err := db.Begin()
|
||||||
|
if err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
defer func() { _ = tx.Rollback() }()
|
||||||
|
var userExists int
|
||||||
|
if err := tx.QueryRow(`SELECT COUNT(*) FROM rbac_users WHERE id = ?`, userID).Scan(&userExists); err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
if userExists == 0 {
|
||||||
|
return 0, nil, errors.New("用户不存在")
|
||||||
|
}
|
||||||
|
|
||||||
|
typeTablePairs := []struct{ resourceType, table string }{
|
||||||
|
{"project", "projects"}, {"conversation", "conversations"},
|
||||||
|
{"vulnerability", "vulnerabilities"}, {"webshell", "webshell_connections"},
|
||||||
|
{"batch_task", "batch_task_queues"}, {"c2_listener", "c2_listeners"},
|
||||||
|
}
|
||||||
|
detected := make(map[string]string, len(uniqueIDs))
|
||||||
|
for _, resourceID := range uniqueIDs {
|
||||||
|
for _, pair := range typeTablePairs {
|
||||||
|
var exists int
|
||||||
|
if err := tx.QueryRow(`SELECT COUNT(*) FROM `+pair.table+` WHERE id = ?`, resourceID).Scan(&exists); err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
if exists > 0 {
|
||||||
|
if previous := detected[resourceID]; previous != "" {
|
||||||
|
return 0, nil, fmt.Errorf("资源 ID 同时匹配多个类型: %s (%s, %s)", resourceID, previous, pair.resourceType)
|
||||||
|
}
|
||||||
|
detected[resourceID] = pair.resourceType
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if detected[resourceID] == "" {
|
||||||
|
return 0, nil, fmt.Errorf("资源不存在: %s", resourceID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
var created int64
|
||||||
|
for _, resourceID := range uniqueIDs {
|
||||||
|
result, err := tx.Exec(`
|
||||||
|
INSERT OR IGNORE INTO rbac_resource_assignments (id, user_id, resource_type, resource_id, created_at)
|
||||||
|
VALUES (?, ?, ?, ?, ?)
|
||||||
|
`, uuid.NewString(), userID, detected[resourceID], resourceID, time.Now())
|
||||||
|
if err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
if n, err := result.RowsAffected(); err == nil {
|
||||||
|
created += n
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := tx.Commit(); err != nil {
|
||||||
|
return 0, nil, err
|
||||||
|
}
|
||||||
|
return created, detected, nil
|
||||||
|
}
|
||||||
|
|
||||||
func (db *DB) ListRBACUsers() ([]RBACUser, error) {
|
func (db *DB) ListRBACUsers() ([]RBACUser, error) {
|
||||||
rows, err := db.Query(`SELECT id, username, display_name, password_hash, enabled, is_builtin, created_at, updated_at FROM rbac_users ORDER BY username ASC`)
|
rows, err := db.Query(`SELECT id, username, display_name, password_hash, enabled, is_builtin, created_at, updated_at FROM rbac_users ORDER BY username ASC`)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -1233,16 +1368,50 @@ func (db *DB) ListRBACResourceAssignments(userID string) ([]RBACResourceAssignme
|
|||||||
}
|
}
|
||||||
|
|
||||||
func (db *DB) DeleteRBACResourceAssignment(id string) error {
|
func (db *DB) DeleteRBACResourceAssignment(id string) error {
|
||||||
|
_, err := db.DeleteRBACResourceAssignmentWithDetails(id)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
// DeleteRBACResourceAssignmentWithDetails atomically removes an assignment and
|
||||||
|
// returns the deleted row so callers can write a complete, attributable audit
|
||||||
|
// event without racing a separate lookup against another delete.
|
||||||
|
func (db *DB) DeleteRBACResourceAssignmentWithDetails(id string) (*RBACResourceAssignment, error) {
|
||||||
id = strings.TrimSpace(id)
|
id = strings.TrimSpace(id)
|
||||||
if id == "" {
|
if id == "" {
|
||||||
return errors.New("assignment id is required")
|
return nil, errors.New("assignment id is required")
|
||||||
}
|
}
|
||||||
result, err := db.Exec(`DELETE FROM rbac_resource_assignments WHERE id = ?`, id)
|
tx, err := db.Begin()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return nil, err
|
||||||
}
|
}
|
||||||
if affected, err := result.RowsAffected(); err == nil && affected == 0 {
|
defer tx.Rollback()
|
||||||
return errors.New("资源授权不存在或已撤销")
|
|
||||||
|
var row RBACResourceAssignment
|
||||||
|
var createdAt string
|
||||||
|
err = tx.QueryRow(`
|
||||||
|
SELECT id, user_id, resource_type, resource_id, created_at
|
||||||
|
FROM rbac_resource_assignments
|
||||||
|
WHERE id = ?
|
||||||
|
`, id).Scan(&row.ID, &row.UserID, &row.ResourceType, &row.ResourceID, &createdAt)
|
||||||
|
if errors.Is(err, sql.ErrNoRows) {
|
||||||
|
return nil, errors.New("资源授权不存在或已撤销")
|
||||||
}
|
}
|
||||||
return nil
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
row.CreatedAt = parseDBTime(createdAt)
|
||||||
|
|
||||||
|
result, err := tx.Exec(`DELETE FROM rbac_resource_assignments WHERE id = ?`, id)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if affected, rowsErr := result.RowsAffected(); rowsErr != nil {
|
||||||
|
return nil, rowsErr
|
||||||
|
} else if affected != 1 {
|
||||||
|
return nil, errors.New("资源授权不存在或已撤销")
|
||||||
|
}
|
||||||
|
if err := tx.Commit(); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return &row, nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -579,3 +579,43 @@ func TestRBACAssignmentLabelsAndWeakTitles(t *testing.T) {
|
|||||||
t.Fatalf("assignment label = %q, want Alpha Project", rows[0].ResourceLabel)
|
t.Fatalf("assignment label = %q, want Alpha Project", rows[0].ResourceLabel)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestDeleteRBACResourceAssignmentWithDetails(t *testing.T) {
|
||||||
|
db := newRBACTestDB(t)
|
||||||
|
user, err := db.CreateRBACUser("revoke-member", "Revoke Member", "hash", true, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
project, err := db.CreateProject(&Project{Name: "Revoked Project"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := db.AssignResourcesToUser(user.ID, "project", []string{project.ID}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
rows, err := db.ListRBACResourceAssignments(user.ID)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(rows) != 1 {
|
||||||
|
t.Fatalf("assignments = %#v, want 1", rows)
|
||||||
|
}
|
||||||
|
|
||||||
|
deleted, err := db.DeleteRBACResourceAssignmentWithDetails(rows[0].ID)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if deleted.ID != rows[0].ID || deleted.UserID != user.ID || deleted.ResourceType != "project" || deleted.ResourceID != project.ID {
|
||||||
|
t.Fatalf("deleted assignment = %#v", deleted)
|
||||||
|
}
|
||||||
|
remaining, err := db.ListRBACResourceAssignments(user.ID)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(remaining) != 0 {
|
||||||
|
t.Fatalf("remaining assignments = %#v, want none", remaining)
|
||||||
|
}
|
||||||
|
if _, err := db.DeleteRBACResourceAssignmentWithDetails(rows[0].ID); err == nil {
|
||||||
|
t.Fatal("second delete unexpectedly succeeded")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ type RobotSessionBinding struct {
|
|||||||
SessionKey string
|
SessionKey string
|
||||||
ConversationID string
|
ConversationID string
|
||||||
RoleName string
|
RoleName string
|
||||||
|
AgentMode string
|
||||||
UpdatedAt time.Time
|
UpdatedAt time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -24,9 +25,9 @@ func (db *DB) GetRobotSessionBinding(sessionKey string) (*RobotSessionBinding, e
|
|||||||
var b RobotSessionBinding
|
var b RobotSessionBinding
|
||||||
var updatedAt string
|
var updatedAt string
|
||||||
err := db.QueryRow(
|
err := db.QueryRow(
|
||||||
"SELECT session_key, conversation_id, role_name, updated_at FROM robot_user_sessions WHERE session_key = ?",
|
"SELECT session_key, conversation_id, role_name, agent_mode, updated_at FROM robot_user_sessions WHERE session_key = ?",
|
||||||
sessionKey,
|
sessionKey,
|
||||||
).Scan(&b.SessionKey, &b.ConversationID, &b.RoleName, &updatedAt)
|
).Scan(&b.SessionKey, &b.ConversationID, &b.RoleName, &b.AgentMode, &updatedAt)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
if err == sql.ErrNoRows {
|
if err == sql.ErrNoRows {
|
||||||
return nil, nil
|
return nil, nil
|
||||||
@@ -43,28 +44,36 @@ func (db *DB) GetRobotSessionBinding(sessionKey string) (*RobotSessionBinding, e
|
|||||||
if strings.TrimSpace(b.RoleName) == "" {
|
if strings.TrimSpace(b.RoleName) == "" {
|
||||||
b.RoleName = "默认"
|
b.RoleName = "默认"
|
||||||
}
|
}
|
||||||
|
if strings.TrimSpace(b.AgentMode) == "" {
|
||||||
|
b.AgentMode = "eino_single"
|
||||||
|
}
|
||||||
return &b, nil
|
return &b, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpsertRobotSessionBinding 写入或更新机器人会话绑定(包含角色)。
|
// UpsertRobotSessionBinding 写入或更新机器人会话绑定(包含角色)。
|
||||||
func (db *DB) UpsertRobotSessionBinding(sessionKey, conversationID, roleName string) error {
|
func (db *DB) UpsertRobotSessionBinding(sessionKey, conversationID, roleName, agentMode string) error {
|
||||||
sessionKey = strings.TrimSpace(sessionKey)
|
sessionKey = strings.TrimSpace(sessionKey)
|
||||||
conversationID = strings.TrimSpace(conversationID)
|
conversationID = strings.TrimSpace(conversationID)
|
||||||
roleName = strings.TrimSpace(roleName)
|
roleName = strings.TrimSpace(roleName)
|
||||||
|
agentMode = strings.TrimSpace(agentMode)
|
||||||
if sessionKey == "" || conversationID == "" {
|
if sessionKey == "" || conversationID == "" {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
if roleName == "" {
|
if roleName == "" {
|
||||||
roleName = "默认"
|
roleName = "默认"
|
||||||
}
|
}
|
||||||
|
if agentMode == "" {
|
||||||
|
agentMode = "eino_single"
|
||||||
|
}
|
||||||
_, err := db.Exec(`
|
_, err := db.Exec(`
|
||||||
INSERT INTO robot_user_sessions (session_key, conversation_id, role_name, updated_at)
|
INSERT INTO robot_user_sessions (session_key, conversation_id, role_name, agent_mode, updated_at)
|
||||||
VALUES (?, ?, ?, ?)
|
VALUES (?, ?, ?, ?, ?)
|
||||||
ON CONFLICT(session_key) DO UPDATE SET
|
ON CONFLICT(session_key) DO UPDATE SET
|
||||||
conversation_id = excluded.conversation_id,
|
conversation_id = excluded.conversation_id,
|
||||||
role_name = excluded.role_name,
|
role_name = excluded.role_name,
|
||||||
|
agent_mode = excluded.agent_mode,
|
||||||
updated_at = excluded.updated_at
|
updated_at = excluded.updated_at
|
||||||
`, sessionKey, conversationID, roleName, time.Now())
|
`, sessionKey, conversationID, roleName, agentMode, time.Now())
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return fmt.Errorf("写入机器人会话绑定失败: %w", err)
|
return fmt.Errorf("写入机器人会话绑定失败: %w", err)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -191,7 +191,6 @@ func (db *DB) CreateVulnerability(vuln *Vulnerability) (*Vulnerability, error) {
|
|||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("创建漏洞失败: %w", err)
|
return nil, fmt.Errorf("创建漏洞失败: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
return vuln, nil
|
return vuln, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -0,0 +1,215 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"database/sql"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// VulnerabilityAlertSubscription is the single source of truth shared by Web
|
||||||
|
// settings and robot commands. Alerts are opt-in and user scoped.
|
||||||
|
type VulnerabilityAlertSubscription struct {
|
||||||
|
UserID string `json:"user_id"`
|
||||||
|
Enabled bool `json:"enabled"`
|
||||||
|
MinSeverity string `json:"min_severity"`
|
||||||
|
CreatedAt time.Time `json:"created_at"`
|
||||||
|
UpdatedAt time.Time `json:"updated_at"`
|
||||||
|
}
|
||||||
|
|
||||||
|
type VulnerabilityAlertRecipient struct {
|
||||||
|
UserID string
|
||||||
|
Platform string
|
||||||
|
ExternalUserID string
|
||||||
|
}
|
||||||
|
|
||||||
|
type VulnerabilityAlertDelivery struct {
|
||||||
|
ID int64
|
||||||
|
Vulnerability *Vulnerability
|
||||||
|
UserID string
|
||||||
|
Platform string
|
||||||
|
ExternalUserID string
|
||||||
|
Attempts int
|
||||||
|
}
|
||||||
|
|
||||||
|
var vulnerabilitySeverityRank = map[string]int{
|
||||||
|
"info": 0, "low": 1, "medium": 2, "high": 3, "critical": 4,
|
||||||
|
}
|
||||||
|
|
||||||
|
func NormalizeVulnerabilityAlertSeverity(value string) (string, error) {
|
||||||
|
value = strings.ToLower(strings.TrimSpace(value))
|
||||||
|
if value == "" {
|
||||||
|
value = "high"
|
||||||
|
}
|
||||||
|
if _, ok := vulnerabilitySeverityRank[value]; !ok {
|
||||||
|
return "", fmt.Errorf("invalid minimum severity %q", value)
|
||||||
|
}
|
||||||
|
return value, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) GetVulnerabilityAlertSubscription(userID string) (*VulnerabilityAlertSubscription, error) {
|
||||||
|
userID = strings.TrimSpace(userID)
|
||||||
|
var sub VulnerabilityAlertSubscription
|
||||||
|
var enabled int
|
||||||
|
var createdAt, updatedAt string
|
||||||
|
err := db.QueryRow(`SELECT user_id, enabled, min_severity, created_at, updated_at
|
||||||
|
FROM vulnerability_alert_subscriptions WHERE user_id = ?`, userID).
|
||||||
|
Scan(&sub.UserID, &enabled, &sub.MinSeverity, &createdAt, &updatedAt)
|
||||||
|
if err == sql.ErrNoRows {
|
||||||
|
now := time.Now()
|
||||||
|
return &VulnerabilityAlertSubscription{UserID: userID, MinSeverity: "high", CreatedAt: now, UpdatedAt: now}, nil
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
sub.Enabled = enabled != 0
|
||||||
|
sub.CreatedAt = parseDBTime(createdAt)
|
||||||
|
sub.UpdatedAt = parseDBTime(updatedAt)
|
||||||
|
return &sub, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) UpsertVulnerabilityAlertSubscription(userID string, enabled bool, minSeverity string) (*VulnerabilityAlertSubscription, error) {
|
||||||
|
userID = strings.TrimSpace(userID)
|
||||||
|
if userID == "" {
|
||||||
|
return nil, fmt.Errorf("user id is required")
|
||||||
|
}
|
||||||
|
severity, err := NormalizeVulnerabilityAlertSeverity(minSeverity)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
now := time.Now()
|
||||||
|
_, err = db.Exec(`INSERT INTO vulnerability_alert_subscriptions
|
||||||
|
(user_id, enabled, min_severity, created_at, updated_at) VALUES (?, ?, ?, ?, ?)
|
||||||
|
ON CONFLICT(user_id) DO UPDATE SET enabled = excluded.enabled,
|
||||||
|
min_severity = excluded.min_severity, updated_at = excluded.updated_at`,
|
||||||
|
userID, boolToInt(enabled), severity, now, now)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
return db.GetVulnerabilityAlertSubscription(userID)
|
||||||
|
}
|
||||||
|
|
||||||
|
// ListVulnerabilityAlertRecipients applies the same RBAC ownership/assignment
|
||||||
|
// boundaries as the vulnerability list, then expands only enabled robot bindings.
|
||||||
|
func (db *DB) ListVulnerabilityAlertRecipients(vuln *Vulnerability) ([]VulnerabilityAlertRecipient, error) {
|
||||||
|
if vuln == nil {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
rank, ok := vulnerabilitySeverityRank[strings.ToLower(strings.TrimSpace(vuln.Severity))]
|
||||||
|
if !ok {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
rows, err := db.Query(`
|
||||||
|
SELECT DISTINCT s.user_id, b.platform, b.external_user_id, s.min_severity
|
||||||
|
FROM vulnerability_alert_subscriptions s
|
||||||
|
JOIN rbac_users u ON u.id = s.user_id AND u.enabled = 1
|
||||||
|
JOIN robot_user_bindings b ON b.rbac_user_id = s.user_id AND b.enabled = 1
|
||||||
|
WHERE s.enabled = 1 AND (
|
||||||
|
EXISTS (SELECT 1 FROM vulnerabilities v WHERE v.id = ? AND v.owner_user_id = s.user_id)
|
||||||
|
OR EXISTS (SELECT 1 FROM rbac_resource_assignments ra WHERE ra.user_id = s.user_id AND ra.resource_type = 'vulnerability' AND ra.resource_id = ?)
|
||||||
|
OR (? <> '' AND (EXISTS (SELECT 1 FROM projects p WHERE p.id = ? AND p.owner_user_id = s.user_id)
|
||||||
|
OR EXISTS (SELECT 1 FROM rbac_resource_assignments pra WHERE pra.user_id = s.user_id AND pra.resource_type = 'project' AND pra.resource_id = ?)))
|
||||||
|
OR (? <> '' AND (EXISTS (SELECT 1 FROM conversations c WHERE c.id = ? AND c.owner_user_id = s.user_id)
|
||||||
|
OR EXISTS (SELECT 1 FROM rbac_resource_assignments cra WHERE cra.user_id = s.user_id AND cra.resource_type = 'conversation' AND cra.resource_id = ?)))
|
||||||
|
)`, vuln.ID, vuln.ID, vuln.ProjectID, vuln.ProjectID, vuln.ProjectID,
|
||||||
|
vuln.ConversationID, vuln.ConversationID, vuln.ConversationID)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
out := make([]VulnerabilityAlertRecipient, 0)
|
||||||
|
for rows.Next() {
|
||||||
|
var recipient VulnerabilityAlertRecipient
|
||||||
|
var minimum string
|
||||||
|
if err := rows.Scan(&recipient.UserID, &recipient.Platform, &recipient.ExternalUserID, &minimum); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if rank >= vulnerabilitySeverityRank[minimum] {
|
||||||
|
out = append(out, recipient)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) SetVulnerabilityCreatedHook(hook func(*Vulnerability)) {
|
||||||
|
db.vulnerabilityCreatedHook = hook
|
||||||
|
}
|
||||||
|
|
||||||
|
// NotifyVulnerabilityCreated must be called after resource ownership has been
|
||||||
|
// committed. Delivery runs asynchronously and never delays the write path.
|
||||||
|
func (db *DB) NotifyVulnerabilityCreated(vulnerability *Vulnerability) {
|
||||||
|
if db == nil || vulnerability == nil || db.vulnerabilityCreatedHook == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
created := *vulnerability
|
||||||
|
go db.vulnerabilityCreatedHook(&created)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) EnqueueVulnerabilityAlertDeliveries(vulnerabilityID string, recipients []VulnerabilityAlertRecipient) error {
|
||||||
|
now := time.Now()
|
||||||
|
tx, err := db.Begin()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer func() { _ = tx.Rollback() }()
|
||||||
|
for _, r := range recipients {
|
||||||
|
if _, err := tx.Exec(`INSERT INTO vulnerability_alert_deliveries
|
||||||
|
(vulnerability_id, user_id, platform, external_user_id, status, attempts, next_attempt_at, created_at, updated_at)
|
||||||
|
VALUES (?, ?, ?, ?, 'pending', 0, ?, ?, ?)
|
||||||
|
ON CONFLICT(vulnerability_id, platform, external_user_id) DO NOTHING`,
|
||||||
|
vulnerabilityID, r.UserID, r.Platform, r.ExternalUserID, now, now, now); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return tx.Commit()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) ListDueVulnerabilityAlertDeliveries(limit int) ([]VulnerabilityAlertDelivery, error) {
|
||||||
|
if limit <= 0 || limit > 100 {
|
||||||
|
limit = 50
|
||||||
|
}
|
||||||
|
rows, err := db.Query(`SELECT d.id, d.user_id, d.platform, d.external_user_id, d.attempts,
|
||||||
|
v.id, COALESCE(v.conversation_id,''), COALESCE(v.project_id,''), v.title, COALESCE(v.description,''),
|
||||||
|
v.severity, v.status, COALESCE(v.vulnerability_type,''), COALESCE(v.target,''),
|
||||||
|
COALESCE(v.impact,''), COALESCE(v.recommendation,''), v.created_at, v.updated_at
|
||||||
|
FROM vulnerability_alert_deliveries d JOIN vulnerabilities v ON v.id = d.vulnerability_id
|
||||||
|
WHERE d.status IN ('pending','retry') AND d.next_attempt_at <= ?
|
||||||
|
ORDER BY d.next_attempt_at, d.id LIMIT ?`, time.Now(), limit)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
defer rows.Close()
|
||||||
|
var out []VulnerabilityAlertDelivery
|
||||||
|
for rows.Next() {
|
||||||
|
var d VulnerabilityAlertDelivery
|
||||||
|
v := &Vulnerability{}
|
||||||
|
if err := rows.Scan(&d.ID, &d.UserID, &d.Platform, &d.ExternalUserID, &d.Attempts,
|
||||||
|
&v.ID, &v.ConversationID, &v.ProjectID, &v.Title, &v.Description, &v.Severity, &v.Status,
|
||||||
|
&v.Type, &v.Target, &v.Impact, &v.Recommendation, &v.CreatedAt, &v.UpdatedAt); err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
d.Vulnerability = v
|
||||||
|
out = append(out, d)
|
||||||
|
}
|
||||||
|
return out, rows.Err()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) MarkVulnerabilityAlertDeliverySent(id int64) error {
|
||||||
|
_, err := db.Exec(`UPDATE vulnerability_alert_deliveries SET status='sent', attempts=attempts+1, last_error='', updated_at=? WHERE id=?`, time.Now(), id)
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) MarkVulnerabilityAlertDeliveryFailed(id int64, attempts int, sendErr error) error {
|
||||||
|
status := "retry"
|
||||||
|
if attempts >= 5 {
|
||||||
|
status = "failed"
|
||||||
|
}
|
||||||
|
delay := time.Minute * time.Duration(1<<min(attempts, 6))
|
||||||
|
message := ""
|
||||||
|
if sendErr != nil {
|
||||||
|
message = sendErr.Error()
|
||||||
|
}
|
||||||
|
_, err := db.Exec(`UPDATE vulnerability_alert_deliveries SET status=?, attempts=?, next_attempt_at=?, last_error=?, updated_at=? WHERE id=?`,
|
||||||
|
status, attempts, time.Now().Add(delay), message, time.Now(), id)
|
||||||
|
return err
|
||||||
|
}
|
||||||
@@ -0,0 +1,88 @@
|
|||||||
|
package database_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"cyberstrike-ai/internal/database"
|
||||||
|
"cyberstrike-ai/internal/security"
|
||||||
|
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestVulnerabilityAlertSubscriptionIsOptInAndRBACScoped(t *testing.T) {
|
||||||
|
db, err := database.NewDB(t.TempDir()+"/alerts.db", zap.NewNop())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
if err := db.BootstrapRBAC("hash", security.PermissionCatalog); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
owner, _ := db.CreateRBACUser("alert-owner", "Owner", "hash", true, []string{database.RBACSystemRoleOperator})
|
||||||
|
other, _ := db.CreateRBACUser("alert-other", "Other", "hash", true, []string{database.RBACSystemRoleOperator})
|
||||||
|
|
||||||
|
for i, pair := range []struct {
|
||||||
|
user *database.RBACUser
|
||||||
|
external string
|
||||||
|
}{{owner, "owner"}, {other, "other"}} {
|
||||||
|
code := "code-" + string(rune('a'+i))
|
||||||
|
if err := db.CreateRobotBindingCode(pair.user.ID, code, time.Now().Add(time.Minute)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := db.ConsumeRobotBindingCode("wecom", pair.external, code); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
defaultSub, err := db.GetVulnerabilityAlertSubscription(owner.ID)
|
||||||
|
if err != nil || defaultSub.Enabled || defaultSub.MinSeverity != "high" {
|
||||||
|
t.Fatalf("unsafe default: %#v %v", defaultSub, err)
|
||||||
|
}
|
||||||
|
if _, err := db.UpsertVulnerabilityAlertSubscription(owner.ID, true, "high"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := db.UpsertVulnerabilityAlertSubscription(other.ID, true, "low"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
|
||||||
|
vuln, err := db.CreateVulnerability(&database.Vulnerability{Title: "owned", Severity: "high"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := db.SetResourceOwner("vulnerability", vuln.ID, owner.ID); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
recipients, err := db.ListVulnerabilityAlertRecipients(vuln)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(recipients) != 1 || recipients[0].UserID != owner.ID || recipients[0].ExternalUserID != "owner" {
|
||||||
|
t.Fatalf("alert escaped RBAC boundary: %#v", recipients)
|
||||||
|
}
|
||||||
|
if err := db.EnqueueVulnerabilityAlertDeliveries(vuln.ID, recipients); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := db.EnqueueVulnerabilityAlertDeliveries(vuln.ID, recipients); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
deliveries, err := db.ListDueVulnerabilityAlertDeliveries(10)
|
||||||
|
if err != nil || len(deliveries) != 1 {
|
||||||
|
t.Fatalf("outbox is not durable/deduplicated: %#v %v", deliveries, err)
|
||||||
|
}
|
||||||
|
if err := db.MarkVulnerabilityAlertDeliverySent(deliveries[0].ID); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
deliveries, _ = db.ListDueVulnerabilityAlertDeliveries(10)
|
||||||
|
if len(deliveries) != 0 {
|
||||||
|
t.Fatalf("sent delivery remained due: %#v", deliveries)
|
||||||
|
}
|
||||||
|
|
||||||
|
vuln.Severity = "medium"
|
||||||
|
recipients, err = db.ListVulnerabilityAlertRecipients(vuln)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(recipients) != 0 {
|
||||||
|
t.Fatalf("minimum severity was ignored: %#v", recipients)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,286 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
|
"database/sql"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
"unicode"
|
||||||
|
|
||||||
|
"github.com/google/uuid"
|
||||||
|
)
|
||||||
|
|
||||||
|
type WorkflowPackageInspection struct {
|
||||||
|
ID, PackageHash, ManifestJSON, WorkflowPayloadJSON, InspectionJSON string
|
||||||
|
SourceWorkflowID, SourceContentHash, SourceGraphHash string
|
||||||
|
SourceRevision int
|
||||||
|
LocalConflictState, LocalWorkflowID, LocalContentHash, LocalGraphHash string
|
||||||
|
CreatedBy, Status string
|
||||||
|
CreatedAt, ExpiresAt time.Time
|
||||||
|
ConsumedAt *time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
type WorkflowPackageImport struct {
|
||||||
|
ID, InspectionID, RequestHash, IdempotencyKey, ActorUserID string
|
||||||
|
Action, SourceWorkflowID, TargetWorkflowID, ResultingWorkflowID string
|
||||||
|
Result, ErrorCode, ErrorMessage string
|
||||||
|
CreatedAt time.Time
|
||||||
|
AppliedAt *time.Time
|
||||||
|
}
|
||||||
|
|
||||||
|
type WorkflowPackageApplyRequest struct {
|
||||||
|
InspectionID, RequestHash, IdempotencyKey, ActorUserID, Action, NewWorkflowID string
|
||||||
|
ConfirmOverwrite bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type WorkflowPackageStoreError struct{ Code, Message string }
|
||||||
|
|
||||||
|
func (e *WorkflowPackageStoreError) Error() string { return e.Code + ": " + e.Message }
|
||||||
|
func workflowPackageStoreError(code, message string) error {
|
||||||
|
return &WorkflowPackageStoreError{code, message}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) CreateWorkflowPackageInspection(v *WorkflowPackageInspection) error {
|
||||||
|
if v == nil || strings.TrimSpace(v.ID) == "" || strings.TrimSpace(v.CreatedBy) == "" {
|
||||||
|
return fmt.Errorf("workflow package inspection is incomplete")
|
||||||
|
}
|
||||||
|
if v.CreatedAt.IsZero() {
|
||||||
|
v.CreatedAt = time.Now().UTC()
|
||||||
|
}
|
||||||
|
if v.ExpiresAt.IsZero() {
|
||||||
|
v.ExpiresAt = v.CreatedAt.Add(30 * time.Minute)
|
||||||
|
}
|
||||||
|
_, err := db.Exec(`INSERT INTO workflow_package_inspections (id,package_hash,manifest_json,workflow_payload_json,inspection_json,source_workflow_id,source_revision,source_content_hash,source_graph_hash,local_conflict_state,local_workflow_id,local_content_hash,local_graph_hash,created_by,status,created_at,expires_at) VALUES (?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?,?)`, v.ID, v.PackageHash, v.ManifestJSON, v.WorkflowPayloadJSON, v.InspectionJSON, v.SourceWorkflowID, v.SourceRevision, v.SourceContentHash, v.SourceGraphHash, v.LocalConflictState, nullString(v.LocalWorkflowID), nullString(v.LocalContentHash), nullString(v.LocalGraphHash), v.CreatedBy, "ready", v.CreatedAt.UTC(), v.ExpiresAt.UTC())
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) GetWorkflowPackageInspection(id, actor string) (*WorkflowPackageInspection, error) {
|
||||||
|
now := time.Now().UTC()
|
||||||
|
_, _ = db.Exec(`UPDATE workflow_package_inspections SET status='expired' WHERE status='ready' AND expires_at <= ?`, now)
|
||||||
|
row, err := scanWorkflowPackageInspection(db.QueryRow(`SELECT id,package_hash,manifest_json,workflow_payload_json,inspection_json,source_workflow_id,source_revision,source_content_hash,source_graph_hash,local_conflict_state,COALESCE(local_workflow_id,''),COALESCE(local_content_hash,''),COALESCE(local_graph_hash,''),created_by,status,created_at,expires_at,consumed_at FROM workflow_package_inspections WHERE id=? AND created_by=?`, strings.TrimSpace(id), strings.TrimSpace(actor)))
|
||||||
|
if err == sql.ErrNoRows {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return row, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func scanWorkflowPackageInspection(s interface{ Scan(...any) error }) (*WorkflowPackageInspection, error) {
|
||||||
|
var v WorkflowPackageInspection
|
||||||
|
var consumed sql.NullTime
|
||||||
|
err := s.Scan(&v.ID, &v.PackageHash, &v.ManifestJSON, &v.WorkflowPayloadJSON, &v.InspectionJSON, &v.SourceWorkflowID, &v.SourceRevision, &v.SourceContentHash, &v.SourceGraphHash, &v.LocalConflictState, &v.LocalWorkflowID, &v.LocalContentHash, &v.LocalGraphHash, &v.CreatedBy, &v.Status, &v.CreatedAt, &v.ExpiresAt, &consumed)
|
||||||
|
if consumed.Valid {
|
||||||
|
t := consumed.Time
|
||||||
|
v.ConsumedAt = &t
|
||||||
|
}
|
||||||
|
return &v, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) GetWorkflowPackageImport(id, actor string) (*WorkflowPackageImport, error) {
|
||||||
|
v, err := scanWorkflowPackageImport(db.QueryRow(`SELECT id,inspection_id,request_hash,idempotency_key,actor_user_id,action,source_workflow_id,target_workflow_id,COALESCE(resulting_workflow_id,''),result,COALESCE(error_code,''),COALESCE(error_message,''),created_at,applied_at FROM workflow_package_imports WHERE id=? AND actor_user_id=?`, strings.TrimSpace(id), strings.TrimSpace(actor)))
|
||||||
|
if err == sql.ErrNoRows {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
return v, err
|
||||||
|
}
|
||||||
|
func scanWorkflowPackageImport(s interface{ Scan(...any) error }) (*WorkflowPackageImport, error) {
|
||||||
|
var v WorkflowPackageImport
|
||||||
|
var applied sql.NullTime
|
||||||
|
err := s.Scan(&v.ID, &v.InspectionID, &v.RequestHash, &v.IdempotencyKey, &v.ActorUserID, &v.Action, &v.SourceWorkflowID, &v.TargetWorkflowID, &v.ResultingWorkflowID, &v.Result, &v.ErrorCode, &v.ErrorMessage, &v.CreatedAt, &applied)
|
||||||
|
if applied.Valid {
|
||||||
|
t := applied.Time
|
||||||
|
v.AppliedAt = &t
|
||||||
|
}
|
||||||
|
return &v, err
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) ApplyWorkflowPackageImport(ctx context.Context, req WorkflowPackageApplyRequest) (*WorkflowPackageImport, bool, error) {
|
||||||
|
tx, err := db.BeginTx(ctx, nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, false, err
|
||||||
|
}
|
||||||
|
defer tx.Rollback()
|
||||||
|
var existingHash string
|
||||||
|
previous, prevErr := scanWorkflowPackageImport(tx.QueryRowContext(ctx, `SELECT id,inspection_id,request_hash,idempotency_key,actor_user_id,action,source_workflow_id,target_workflow_id,COALESCE(resulting_workflow_id,''),result,COALESCE(error_code,''),COALESCE(error_message,''),created_at,applied_at FROM workflow_package_imports WHERE actor_user_id=? AND idempotency_key=?`, req.ActorUserID, req.IdempotencyKey))
|
||||||
|
if prevErr == nil {
|
||||||
|
existingHash = previous.RequestHash
|
||||||
|
if existingHash != req.RequestHash {
|
||||||
|
return nil, false, workflowPackageStoreError("WFPKG_IDEMPOTENCY_KEY_REUSED", "幂等键已用于其他请求")
|
||||||
|
}
|
||||||
|
return previous, true, nil
|
||||||
|
}
|
||||||
|
if prevErr != sql.ErrNoRows {
|
||||||
|
return nil, false, prevErr
|
||||||
|
}
|
||||||
|
inspection, err := scanWorkflowPackageInspection(tx.QueryRowContext(ctx, `SELECT id,package_hash,manifest_json,workflow_payload_json,inspection_json,source_workflow_id,source_revision,source_content_hash,source_graph_hash,local_conflict_state,COALESCE(local_workflow_id,''),COALESCE(local_content_hash,''),COALESCE(local_graph_hash,''),created_by,status,created_at,expires_at,consumed_at FROM workflow_package_inspections WHERE id=? AND created_by=?`, req.InspectionID, req.ActorUserID))
|
||||||
|
if err == sql.ErrNoRows {
|
||||||
|
return nil, false, workflowPackageStoreError("WFPKG_INSPECTION_NOT_FOUND", "预检不存在")
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, false, err
|
||||||
|
}
|
||||||
|
now := time.Now().UTC()
|
||||||
|
if !inspection.ExpiresAt.After(now) || inspection.Status == "expired" {
|
||||||
|
_, _ = tx.ExecContext(ctx, `UPDATE workflow_package_inspections SET status='expired' WHERE id=?`, inspection.ID)
|
||||||
|
return nil, false, workflowPackageStoreError("WFPKG_INSPECTION_EXPIRED", "预检已过期")
|
||||||
|
}
|
||||||
|
if inspection.Status != "ready" {
|
||||||
|
return nil, false, workflowPackageStoreError("WFPKG_INSPECTION_CONSUMED", "预检已被使用")
|
||||||
|
}
|
||||||
|
var payload struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
GraphJSON string `json:"graph_json"`
|
||||||
|
Enabled bool `json:"enabled"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal([]byte(inspection.WorkflowPayloadJSON), &payload); err != nil {
|
||||||
|
return nil, false, fmt.Errorf("decode inspection payload: %w", err)
|
||||||
|
}
|
||||||
|
targetID := inspection.SourceWorkflowID
|
||||||
|
if req.Action == "rename" {
|
||||||
|
targetID = strings.TrimSpace(req.NewWorkflowID)
|
||||||
|
if !validWorkflowPackageID(targetID) {
|
||||||
|
return nil, false, workflowPackageStoreError("WFPKG_INVALID_RENAME_ID", "新工作流 ID 无效")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
sourceCurrent, err := scanWorkflowDefinition(tx.QueryRowContext(ctx, "SELECT "+workflowDefinitionColumns+" FROM workflow_definitions WHERE id=?", inspection.SourceWorkflowID))
|
||||||
|
if err == sql.ErrNoRows {
|
||||||
|
sourceCurrent = nil
|
||||||
|
} else if err != nil {
|
||||||
|
return nil, false, err
|
||||||
|
}
|
||||||
|
if err := checkWorkflowPackageSnapshot(inspection, sourceCurrent, inspection.SourceWorkflowID); err != nil {
|
||||||
|
return nil, false, err
|
||||||
|
}
|
||||||
|
current := sourceCurrent
|
||||||
|
if targetID != inspection.SourceWorkflowID {
|
||||||
|
current, err = scanWorkflowDefinition(tx.QueryRowContext(ctx, "SELECT "+workflowDefinitionColumns+" FROM workflow_definitions WHERE id=?", targetID))
|
||||||
|
if err == sql.ErrNoRows {
|
||||||
|
current = nil
|
||||||
|
} else if err != nil {
|
||||||
|
return nil, false, err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
result := ""
|
||||||
|
resultingID := ""
|
||||||
|
switch req.Action {
|
||||||
|
case "create":
|
||||||
|
if inspection.LocalConflictState != "none" || current != nil {
|
||||||
|
return nil, false, workflowPackageStoreError("WFPKG_ID_CONFLICT", "目标工作流已存在")
|
||||||
|
}
|
||||||
|
result = "created"
|
||||||
|
resultingID = targetID
|
||||||
|
case "keep_existing":
|
||||||
|
if inspection.LocalConflictState == "none" {
|
||||||
|
return nil, false, workflowPackageStoreError("WFPKG_ID_CONFLICT", "当前预检不允许保留本地")
|
||||||
|
}
|
||||||
|
if inspection.LocalConflictState == "identical" {
|
||||||
|
result = "skipped_identical"
|
||||||
|
} else {
|
||||||
|
result = "kept_existing"
|
||||||
|
}
|
||||||
|
resultingID = targetID
|
||||||
|
case "overwrite":
|
||||||
|
if inspection.LocalConflictState != "id_conflict" {
|
||||||
|
return nil, false, workflowPackageStoreError("WFPKG_ID_CONFLICT", "当前预检不允许覆盖")
|
||||||
|
}
|
||||||
|
if !req.ConfirmOverwrite {
|
||||||
|
return nil, false, workflowPackageStoreError("WFPKG_OVERWRITE_CONFIRMATION_REQUIRED", "覆盖需要确认")
|
||||||
|
}
|
||||||
|
result = "overwritten"
|
||||||
|
resultingID = targetID
|
||||||
|
case "rename":
|
||||||
|
if inspection.LocalConflictState != "id_conflict" || current != nil {
|
||||||
|
return nil, false, workflowPackageStoreError("WFPKG_ID_CONFLICT", "当前预检不允许另存")
|
||||||
|
}
|
||||||
|
result = "renamed"
|
||||||
|
resultingID = targetID
|
||||||
|
default:
|
||||||
|
return nil, false, workflowPackageStoreError("WFPKG_INVALID_ACTION", "导入动作无效")
|
||||||
|
}
|
||||||
|
if result == "created" || result == "renamed" {
|
||||||
|
_, err = tx.ExecContext(ctx, `INSERT INTO workflow_definitions (id,name,description,version,graph_json,enabled,created_at,updated_at) VALUES (?,?,?,?,?,?,?,?)`, resultingID, payload.Name, payload.Description, 1, payload.GraphJSON, boolToInt(payload.Enabled), now, now)
|
||||||
|
} else if result == "overwritten" {
|
||||||
|
_, err = tx.ExecContext(ctx, `UPDATE workflow_definitions SET name=?,description=?,version=version+1,graph_json=?,enabled=?,updated_at=? WHERE id=?`, payload.Name, payload.Description, payload.GraphJSON, boolToInt(payload.Enabled), now, resultingID)
|
||||||
|
}
|
||||||
|
if err != nil {
|
||||||
|
return nil, false, err
|
||||||
|
}
|
||||||
|
imp := &WorkflowPackageImport{ID: "wpii_" + strings.ReplaceAll(uuid.NewString(), "-", ""), InspectionID: inspection.ID, RequestHash: req.RequestHash, IdempotencyKey: req.IdempotencyKey, ActorUserID: req.ActorUserID, Action: req.Action, SourceWorkflowID: inspection.SourceWorkflowID, TargetWorkflowID: targetID, ResultingWorkflowID: resultingID, Result: result, CreatedAt: now, AppliedAt: &now}
|
||||||
|
_, err = tx.ExecContext(ctx, `INSERT INTO workflow_package_imports (id,inspection_id,request_hash,idempotency_key,actor_user_id,action,source_workflow_id,target_workflow_id,resulting_workflow_id,result,created_at,applied_at) VALUES (?,?,?,?,?,?,?,?,?,?,?,?)`, imp.ID, imp.InspectionID, imp.RequestHash, imp.IdempotencyKey, imp.ActorUserID, imp.Action, imp.SourceWorkflowID, imp.TargetWorkflowID, nullString(imp.ResultingWorkflowID), imp.Result, now, now)
|
||||||
|
if err != nil {
|
||||||
|
return nil, false, err
|
||||||
|
}
|
||||||
|
if _, err = tx.ExecContext(ctx, `UPDATE workflow_package_inspections SET status='consumed',consumed_at=? WHERE id=? AND status='ready'`, now, inspection.ID); err != nil {
|
||||||
|
return nil, false, err
|
||||||
|
}
|
||||||
|
if err = tx.Commit(); err != nil {
|
||||||
|
return nil, false, err
|
||||||
|
}
|
||||||
|
return imp, false, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func checkWorkflowPackageSnapshot(i *WorkflowPackageInspection, current *WorkflowDefinition, targetID string) error {
|
||||||
|
if i.LocalConflictState == "none" {
|
||||||
|
if current != nil {
|
||||||
|
return workflowPackageStoreError("WFPKG_CONFLICT_CHANGED", "本地工作流已变化")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if current == nil || current.ID != i.LocalWorkflowID || current.ID != targetID {
|
||||||
|
return workflowPackageStoreError("WFPKG_CONFLICT_CHANGED", "本地工作流已变化")
|
||||||
|
}
|
||||||
|
content, graph := workflowDefinitionPackageHashes(current)
|
||||||
|
if content != i.LocalContentHash || graph != i.LocalGraphHash {
|
||||||
|
return workflowPackageStoreError("WFPKG_CONFLICT_CHANGED", "本地工作流已变化")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
func workflowDefinitionPackageHashes(w *WorkflowDefinition) (string, string) {
|
||||||
|
var g any
|
||||||
|
dec := json.NewDecoder(strings.NewReader(w.GraphJSON))
|
||||||
|
dec.UseNumber()
|
||||||
|
_ = dec.Decode(&g)
|
||||||
|
graph, _ := json.Marshal(g)
|
||||||
|
payload := struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Description string `json:"description,omitempty"`
|
||||||
|
Version int `json:"version"`
|
||||||
|
GraphJSON string `json:"graph_json"`
|
||||||
|
Enabled bool `json:"enabled"`
|
||||||
|
}{w.ID, w.Name, w.Description, w.Version, string(graph), w.Enabled}
|
||||||
|
b, _ := json.Marshal(payload)
|
||||||
|
return workflowPackageHash(b), workflowPackageHash(graph)
|
||||||
|
}
|
||||||
|
func workflowPackageHash(b []byte) string {
|
||||||
|
s := sha256.Sum256(b)
|
||||||
|
return "sha256:" + hex.EncodeToString(s[:])
|
||||||
|
}
|
||||||
|
func validWorkflowPackageID(id string) bool {
|
||||||
|
if len(id) < 1 || len(id) > 128 {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, r := range id {
|
||||||
|
if unicode.IsControl(r) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func (db *DB) PurgeWorkflowPackageLifecycle(now time.Time) error {
|
||||||
|
now = now.UTC()
|
||||||
|
if _, err := db.Exec(`UPDATE workflow_package_inspections SET status='expired' WHERE status='ready' AND expires_at<=?`, now); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if _, err := db.Exec(`DELETE FROM workflow_package_inspections WHERE status='expired' AND expires_at<? AND NOT EXISTS (SELECT 1 FROM workflow_package_imports i WHERE i.inspection_id=workflow_package_inspections.id)`, now.Add(-24*time.Hour)); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
_, err := db.Exec(`DELETE FROM workflow_package_imports WHERE created_at<?`, now.AddDate(0, 0, -90))
|
||||||
|
return err
|
||||||
|
}
|
||||||
@@ -0,0 +1,74 @@
|
|||||||
|
package database
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestWorkflowPackageApplyOverwriteIsTransactionalAndIdempotent(t *testing.T) {
|
||||||
|
db, err := NewDB(filepath.Join(t.TempDir(), "workflow-package.db"), zap.NewNop())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
current := &WorkflowDefinition{ID: "wf-1", Name: "Local", Version: 12, GraphJSON: `{"nodes":[]}`, Enabled: true}
|
||||||
|
if err := db.UpsertWorkflowDefinition(current); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
current, _ = db.GetWorkflowDefinition("wf-1")
|
||||||
|
content, graph := workflowDefinitionPackageHashes(current)
|
||||||
|
payload, _ := json.Marshal(map[string]any{"id": "wf-1", "name": "Imported", "description": "new", "version": 18, "graph_json": `{"nodes":[]}`, "enabled": false})
|
||||||
|
now := time.Now().UTC()
|
||||||
|
inspection := &WorkflowPackageInspection{ID: "wpi_test", PackageHash: "sha256:pkg", ManifestJSON: "{}", WorkflowPayloadJSON: string(payload), InspectionJSON: "{}", SourceWorkflowID: "wf-1", SourceRevision: 18, SourceContentHash: "sha256:src", SourceGraphHash: "sha256:graph", LocalConflictState: "id_conflict", LocalWorkflowID: "wf-1", LocalContentHash: content, LocalGraphHash: graph, CreatedBy: "user-1", CreatedAt: now, ExpiresAt: now.Add(time.Minute)}
|
||||||
|
if err := db.CreateWorkflowPackageInspection(inspection); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
req := WorkflowPackageApplyRequest{InspectionID: inspection.ID, RequestHash: "sha256:req", IdempotencyKey: "key-1", ActorUserID: "user-1", Action: "overwrite", ConfirmOverwrite: true}
|
||||||
|
imp, replayed, err := db.ApplyWorkflowPackageImport(context.Background(), req)
|
||||||
|
if err != nil || replayed || imp.Result != "overwritten" {
|
||||||
|
t.Fatalf("apply = %#v replay=%v err=%v", imp, replayed, err)
|
||||||
|
}
|
||||||
|
updated, _ := db.GetWorkflowDefinition("wf-1")
|
||||||
|
if updated.Version != 13 || updated.Name != "Imported" || updated.Enabled {
|
||||||
|
t.Fatalf("updated workflow = %#v", updated)
|
||||||
|
}
|
||||||
|
replay, replayed, err := db.ApplyWorkflowPackageImport(context.Background(), req)
|
||||||
|
if err != nil || !replayed || replay.ID != imp.ID {
|
||||||
|
t.Fatalf("replay = %#v replay=%v err=%v", replay, replayed, err)
|
||||||
|
}
|
||||||
|
gotInspection, err := db.GetWorkflowPackageInspection(inspection.ID, "user-1")
|
||||||
|
if err != nil || gotInspection.Status != "consumed" {
|
||||||
|
t.Fatalf("inspection=%#v err=%v", gotInspection, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestWorkflowPackageApplyRejectsChangedConflictSnapshot(t *testing.T) {
|
||||||
|
db, err := NewDB(filepath.Join(t.TempDir(), "workflow-package-conflict.db"), zap.NewNop())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
if err := db.UpsertWorkflowDefinition(&WorkflowDefinition{ID: "wf-2", Name: "Local", GraphJSON: `{"nodes":[]}`, Enabled: true}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
local, _ := db.GetWorkflowDefinition("wf-2")
|
||||||
|
content, graph := workflowDefinitionPackageHashes(local)
|
||||||
|
payload, _ := json.Marshal(map[string]any{"id": "wf-2", "name": "Imported", "version": 2, "graph_json": `{"nodes":[]}`, "enabled": true})
|
||||||
|
now := time.Now().UTC()
|
||||||
|
inspection := &WorkflowPackageInspection{ID: "wpi_changed", PackageHash: "sha256:pkg", ManifestJSON: "{}", WorkflowPayloadJSON: string(payload), InspectionJSON: "{}", SourceWorkflowID: "wf-2", SourceRevision: 2, SourceContentHash: "sha256:src", SourceGraphHash: "sha256:graph", LocalConflictState: "id_conflict", LocalWorkflowID: "wf-2", LocalContentHash: content, LocalGraphHash: graph, CreatedBy: "user-1", CreatedAt: now, ExpiresAt: now.Add(time.Minute)}
|
||||||
|
if err := db.CreateWorkflowPackageInspection(inspection); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := db.UpsertWorkflowDefinition(&WorkflowDefinition{ID: "wf-2", Name: "Changed", GraphJSON: `{"nodes":[]}`, Enabled: true}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
_, _, err = db.ApplyWorkflowPackageImport(context.Background(), WorkflowPackageApplyRequest{InspectionID: inspection.ID, RequestHash: "sha256:req", IdempotencyKey: "key-2", ActorUserID: "user-1", Action: "overwrite", ConfirmOverwrite: true})
|
||||||
|
if e, ok := err.(*WorkflowPackageStoreError); !ok || e.Code != "WFPKG_CONFLICT_CHANGED" {
|
||||||
|
t.Fatalf("err=%v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -729,7 +729,7 @@ func (h *AgentHandler) runRobotMultiAgentWithRetry(
|
|||||||
}
|
}
|
||||||
|
|
||||||
// ProcessMessageForRobot 供机器人(企业微信/钉钉/飞书)调用:Eino 单/多代理执行路径(含 progressCallback、过程详情),仅不发送 SSE,最后返回完整回复
|
// ProcessMessageForRobot 供机器人(企业微信/钉钉/飞书)调用:Eino 单/多代理执行路径(含 progressCallback、过程详情),仅不发送 SSE,最后返回完整回复
|
||||||
func (h *AgentHandler) ProcessMessageForRobot(ctx context.Context, platform string, principal authctx.Principal, conversationID, message, role string) (response string, convID string, err error) {
|
func (h *AgentHandler) ProcessMessageForRobot(ctx context.Context, platform string, principal authctx.Principal, conversationID, message, role, agentMode string) (response string, convID string, err error) {
|
||||||
ownerUserID := strings.TrimSpace(principal.UserID)
|
ownerUserID := strings.TrimSpace(principal.UserID)
|
||||||
if ownerUserID == "" {
|
if ownerUserID == "" {
|
||||||
return "", "", fmt.Errorf("authenticated robot principal is required")
|
return "", "", fmt.Errorf("authenticated robot principal is required")
|
||||||
@@ -814,18 +814,14 @@ func (h *AgentHandler) ProcessMessageForRobot(ctx context.Context, platform stri
|
|||||||
}
|
}
|
||||||
progressCallback := h.createProgressCallback(taskCtx, cancelWithCause, conversationID, assistantMessageID, nil)
|
progressCallback := h.createProgressCallback(taskCtx, cancelWithCause, conversationID, assistantMessageID, nil)
|
||||||
|
|
||||||
robotMode := "eino_single"
|
robotMode := config.NormalizeAgentMode(agentMode)
|
||||||
if h.config != nil {
|
|
||||||
robotMode = config.NormalizeRobotAgentMode(h.config.MultiAgent)
|
|
||||||
}
|
|
||||||
switch robotMode {
|
switch robotMode {
|
||||||
case "eino_single":
|
case "eino_single":
|
||||||
return h.runRobotEinoSingleWithRetry(taskCtx, conversationID, finalMessage, agentHistoryMessages, roleTools, progressCallback, assistantMessageID, &taskStatus)
|
return h.runRobotEinoSingleWithRetry(taskCtx, conversationID, finalMessage, agentHistoryMessages, roleTools, progressCallback, assistantMessageID, &taskStatus)
|
||||||
case "deep", "plan_execute", "supervisor":
|
case "deep", "plan_execute", "supervisor":
|
||||||
if h.config == nil || !h.config.MultiAgent.Enabled {
|
if h.config == nil || !h.config.MultiAgent.Enabled {
|
||||||
h.logger.Warn("机器人配置为多代理模式但未启用 multi_agent,回退 Eino 单代理",
|
taskStatus = "failed"
|
||||||
zap.String("robot_mode", robotMode))
|
return "", conversationID, fmt.Errorf("机器人对话模式 %s 需要启用 Eino 多代理", robotMode)
|
||||||
return h.runRobotEinoSingleWithRetry(taskCtx, conversationID, finalMessage, agentHistoryMessages, roleTools, progressCallback, assistantMessageID, &taskStatus)
|
|
||||||
}
|
}
|
||||||
return h.runRobotMultiAgentWithRetry(taskCtx, conversationID, finalMessage, robotMode, agentHistoryMessages, roleTools, progressCallback, assistantMessageID, &taskStatus)
|
return h.runRobotMultiAgentWithRetry(taskCtx, conversationID, finalMessage, robotMode, agentHistoryMessages, roleTools, progressCallback, assistantMessageID, &taskStatus)
|
||||||
}
|
}
|
||||||
@@ -1611,9 +1607,8 @@ func (h *AgentHandler) SubscribeAgentTaskEvents(c *gin.Context) {
|
|||||||
flusher, _ := c.Writer.(http.Flusher)
|
flusher, _ := c.Writer.(http.Flusher)
|
||||||
ctx := c.Request.Context()
|
ctx := c.Request.Context()
|
||||||
var writeMu sync.Mutex
|
var writeMu sync.Mutex
|
||||||
stopKeepalive := make(chan struct{})
|
stopKeepalive := runSSEKeepalive(c, &writeMu)
|
||||||
go sseKeepalive(c, stopKeepalive, &writeMu)
|
defer stopKeepalive()
|
||||||
defer close(stopKeepalive)
|
|
||||||
|
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
|
|||||||
@@ -10,13 +10,15 @@ import (
|
|||||||
|
|
||||||
func auditFilterFromQuery(c *gin.Context) database.ListAuditLogsFilter {
|
func auditFilterFromQuery(c *gin.Context) database.ListAuditLogsFilter {
|
||||||
filter := database.ListAuditLogsFilter{
|
filter := database.ListAuditLogsFilter{
|
||||||
Level: c.Query("level"),
|
Actor: c.Query("actor"),
|
||||||
Category: c.Query("category"),
|
Level: c.Query("level"),
|
||||||
Action: c.Query("action"),
|
Category: c.Query("category"),
|
||||||
Result: c.Query("result"),
|
Action: c.Query("action"),
|
||||||
Query: c.Query("q"),
|
Result: c.Query("result"),
|
||||||
ResourceType: c.Query("resource_type"),
|
Query: c.Query("q"),
|
||||||
ResourceID: c.Query("resource_id"),
|
ResourceType: c.Query("resource_type"),
|
||||||
|
ResourceID: c.Query("resource_id"),
|
||||||
|
RelatedUserID: c.Query("related_user_id"),
|
||||||
}
|
}
|
||||||
if since := c.Query("since"); since != "" {
|
if since := c.Query("since"); since != "" {
|
||||||
if t, err := database.ParseRFC3339Time(since); err == nil {
|
if t, err := database.ParseRFC3339Time(since); err == nil {
|
||||||
|
|||||||
@@ -0,0 +1,19 @@
|
|||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http/httptest"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAuditFilterFromQueryIncludesActorAndMemberFilters(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
context, _ := gin.CreateTestContext(httptest.NewRecorder())
|
||||||
|
context.Request = httptest.NewRequest("GET", "/api/audit/logs?actor=operator-user&action=assign_resource&resource_type=conversation&related_user_id=user-1", nil)
|
||||||
|
|
||||||
|
filter := auditFilterFromQuery(context)
|
||||||
|
if filter.Actor != "operator-user" || filter.Action != "assign_resource" || filter.ResourceType != "conversation" || filter.RelatedUserID != "user-1" {
|
||||||
|
t.Fatalf("filter = %#v", filter)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -170,28 +170,10 @@ func (h *AuthHandler) ChangePassword(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
if session.UserID == "" || session.UserID == "admin" {
|
if session.UserID == "" {
|
||||||
if err := config.PersistAuthPassword(h.configPath, newPassword); err != nil {
|
session.UserID = "admin"
|
||||||
if h.logger != nil {
|
}
|
||||||
h.logger.Error("保存新密码失败", zap.Error(err))
|
if err := h.manager.UpdateUserPassword(session.UserID, newPassword); err != nil {
|
||||||
}
|
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "保存新密码失败,请重试"})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
if err := h.manager.UpdateConfig(newPassword, h.config.Auth.SessionDurationHours); err != nil {
|
|
||||||
if h.logger != nil {
|
|
||||||
h.logger.Error("更新认证配置失败", zap.Error(err))
|
|
||||||
}
|
|
||||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "更新认证配置失败"})
|
|
||||||
return
|
|
||||||
}
|
|
||||||
|
|
||||||
h.config.Auth.Password = newPassword
|
|
||||||
h.config.Auth.GeneratedPassword = ""
|
|
||||||
h.config.Auth.GeneratedPasswordPersisted = false
|
|
||||||
h.config.Auth.GeneratedPasswordPersistErr = ""
|
|
||||||
} else if err := h.manager.UpdateUserPassword(session.UserID, newPassword); err != nil {
|
|
||||||
if h.logger != nil {
|
if h.logger != nil {
|
||||||
h.logger.Error("更新用户密码失败", zap.Error(err))
|
h.logger.Error("更新用户密码失败", zap.Error(err))
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -664,7 +664,7 @@ schedule_mode 为 cron 时必须提供有效 cron_expr;为 manual 时会清除
|
|||||||
return batchMCPJSONResult(queue)
|
return batchMCPJSONResult(queue)
|
||||||
})
|
})
|
||||||
|
|
||||||
logger.Info("批量任务 MCP 工具已注册", zap.Int("count", 12))
|
logger.Debug("批量任务 MCP 工具已注册", zap.Int("count", 12))
|
||||||
}
|
}
|
||||||
|
|
||||||
// --- batch_task_list 精简结构(避免把每条子任务的 result 等大段文本塞进列表上下文) ---
|
// --- batch_task_list 精简结构(避免把每条子任务的 result 等大段文本塞进列表上下文) ---
|
||||||
|
|||||||
@@ -1362,7 +1362,7 @@ func (h *ConfigHandler) ApplyConfig(c *gin.Context) {
|
|||||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "初始化知识库失败: " + err.Error()})
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "初始化知识库失败: " + err.Error()})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
h.logger.Info("知识库动态初始化完成,工具已注册")
|
h.logger.Debug("知识库动态初始化完成,工具已注册")
|
||||||
}
|
}
|
||||||
|
|
||||||
// 检查嵌入模型配置是否变更(需要在锁外执行,避免阻塞)
|
// 检查嵌入模型配置是否变更(需要在锁外执行,避免阻塞)
|
||||||
@@ -1441,10 +1441,10 @@ func (h *ConfigHandler) ApplyConfig(c *gin.Context) {
|
|||||||
c.JSON(http.StatusInternalServerError, gin.H{"error": "重新加载工具配置失败: " + err.Error()})
|
c.JSON(http.StatusInternalServerError, gin.H{"error": "重新加载工具配置失败: " + err.Error()})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
h.logger.Info("已从 tools 目录重新加载工具配置", zap.Int("tools_count", len(h.config.Security.Tools)))
|
h.logger.Debug("已从 tools 目录重新加载工具配置", zap.Int("tools_count", len(h.config.Security.Tools)))
|
||||||
|
|
||||||
// 重新注册工具(根据新的启用状态)
|
// 重新注册工具(根据新的启用状态)
|
||||||
h.logger.Info("重新注册工具")
|
h.logger.Debug("重新注册工具")
|
||||||
|
|
||||||
// 清空MCP服务器中的工具
|
// 清空MCP服务器中的工具
|
||||||
h.mcpServer.ClearTools()
|
h.mcpServer.ClearTools()
|
||||||
|
|||||||
@@ -139,9 +139,8 @@ func (h *AgentHandler) EinoSingleAgentLoopStream(c *gin.Context) {
|
|||||||
"conversationId": conversationID,
|
"conversationId": conversationID,
|
||||||
})
|
})
|
||||||
|
|
||||||
stopKeepalive := make(chan struct{})
|
stopKeepalive := runSSEKeepalive(c, &sseWriteMu)
|
||||||
go sseKeepalive(c, stopKeepalive, &sseWriteMu)
|
defer stopKeepalive()
|
||||||
defer close(stopKeepalive)
|
|
||||||
|
|
||||||
if h.config == nil {
|
if h.config == nil {
|
||||||
taskStatus = "failed"
|
taskStatus = "failed"
|
||||||
|
|||||||
@@ -69,7 +69,7 @@ func (h *MonitorHandler) SetAgentHandler(ah *AgentHandler) {
|
|||||||
h.agentHandler = ah
|
h.agentHandler = ah
|
||||||
}
|
}
|
||||||
|
|
||||||
const monitorPageTopTools = 6
|
const monitorPageTopTools = 3
|
||||||
|
|
||||||
// MonitorStatsSummary 工具调用汇总
|
// MonitorStatsSummary 工具调用汇总
|
||||||
type MonitorStatsSummary struct {
|
type MonitorStatsSummary struct {
|
||||||
|
|||||||
@@ -156,9 +156,8 @@ func (h *AgentHandler) MultiAgentLoopStream(c *gin.Context) {
|
|||||||
"conversationId": conversationID,
|
"conversationId": conversationID,
|
||||||
})
|
})
|
||||||
|
|
||||||
stopKeepalive := make(chan struct{})
|
stopKeepalive := runSSEKeepalive(c, &sseWriteMu)
|
||||||
go sseKeepalive(c, stopKeepalive, &sseWriteMu)
|
defer stopKeepalive()
|
||||||
defer close(stopKeepalive)
|
|
||||||
|
|
||||||
var result *multiagent.RunResult
|
var result *multiagent.RunResult
|
||||||
var runErr error
|
var runErr error
|
||||||
|
|||||||
@@ -324,6 +324,7 @@ type assignResourceRequest struct {
|
|||||||
ResourceType string `json:"resource_type" binding:"required"`
|
ResourceType string `json:"resource_type" binding:"required"`
|
||||||
ResourceID string `json:"resource_id"`
|
ResourceID string `json:"resource_id"`
|
||||||
ResourceIDs []string `json:"resource_ids"`
|
ResourceIDs []string `json:"resource_ids"`
|
||||||
|
AutoDetect bool `json:"auto_detect"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *RBACHandler) AssignResource(c *gin.Context) {
|
func (h *RBACHandler) AssignResource(c *gin.Context) {
|
||||||
@@ -340,21 +341,33 @@ func (h *RBACHandler) AssignResource(c *gin.Context) {
|
|||||||
c.JSON(http.StatusBadRequest, gin.H{"error": "至少需要一个资源 ID"})
|
c.JSON(http.StatusBadRequest, gin.H{"error": "至少需要一个资源 ID"})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
created, err := h.db.AssignResourcesToUser(req.UserID, req.ResourceType, resourceIDs)
|
var created int64
|
||||||
|
var detectedTypes map[string]string
|
||||||
|
var err error
|
||||||
|
if req.AutoDetect {
|
||||||
|
created, detectedTypes, err = h.db.AssignResourcesToUserAuto(req.UserID, resourceIDs)
|
||||||
|
} else {
|
||||||
|
created, err = h.db.AssignResourcesToUser(req.UserID, req.ResourceType, resourceIDs)
|
||||||
|
}
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if h.audit != nil {
|
if h.audit != nil {
|
||||||
for _, resourceID := range resourceIDs {
|
for _, resourceID := range resourceIDs {
|
||||||
h.audit.RecordOK(c, "rbac", "assign_resource", "授权资源访问", req.ResourceType, strings.TrimSpace(resourceID), map[string]interface{}{"user_id": req.UserID})
|
resourceType := req.ResourceType
|
||||||
|
if detectedTypes != nil {
|
||||||
|
resourceType = detectedTypes[strings.TrimSpace(resourceID)]
|
||||||
|
}
|
||||||
|
h.audit.RecordOK(c, "rbac", "assign_resource", "授权资源访问", resourceType, strings.TrimSpace(resourceID), map[string]interface{}{"user_id": req.UserID})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
c.JSON(http.StatusOK, gin.H{
|
c.JSON(http.StatusOK, gin.H{
|
||||||
"success": true,
|
"success": true,
|
||||||
"requested": len(resourceIDs),
|
"requested": len(resourceIDs),
|
||||||
"created": created,
|
"created": created,
|
||||||
"skipped": int64(len(resourceIDs)) - created,
|
"skipped": int64(len(resourceIDs)) - created,
|
||||||
|
"detected_types": detectedTypes,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -385,22 +398,32 @@ func (h *RBACHandler) ListAssignableResources(c *gin.Context) {
|
|||||||
if hasMore {
|
if hasMore {
|
||||||
resources = resources[:limit]
|
resources = resources[:limit]
|
||||||
}
|
}
|
||||||
|
total, err := h.db.CountAssignableRBACResources(c.Query("type"), c.Query("q"))
|
||||||
|
if err != nil {
|
||||||
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||||
|
return
|
||||||
|
}
|
||||||
c.JSON(http.StatusOK, gin.H{
|
c.JSON(http.StatusOK, gin.H{
|
||||||
"resources": resources,
|
"resources": resources,
|
||||||
"has_more": hasMore,
|
"has_more": hasMore,
|
||||||
"limit": limit,
|
"limit": limit,
|
||||||
"offset": offset,
|
"offset": offset,
|
||||||
|
"total": total,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *RBACHandler) DeleteResourceAssignment(c *gin.Context) {
|
func (h *RBACHandler) DeleteResourceAssignment(c *gin.Context) {
|
||||||
id := strings.TrimSpace(c.Param("id"))
|
id := strings.TrimSpace(c.Param("id"))
|
||||||
if err := h.db.DeleteRBACResourceAssignment(id); err != nil {
|
assignment, err := h.db.DeleteRBACResourceAssignmentWithDetails(id)
|
||||||
|
if err != nil {
|
||||||
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if h.audit != nil {
|
if h.audit != nil {
|
||||||
h.audit.RecordOK(c, "rbac", "delete_resource_assignment", "撤销资源授权", "resource_assignment", id, nil)
|
h.audit.RecordOK(c, "rbac", "delete_resource_assignment", "撤销资源授权", assignment.ResourceType, assignment.ResourceID, map[string]interface{}{
|
||||||
|
"user_id": assignment.UserID,
|
||||||
|
"assignment_id": assignment.ID,
|
||||||
|
})
|
||||||
}
|
}
|
||||||
c.JSON(http.StatusOK, gin.H{"success": true})
|
c.JSON(http.StatusOK, gin.H{"success": true})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -2,7 +2,10 @@ package handler
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"bytes"
|
"bytes"
|
||||||
|
"cyberstrike-ai/internal/audit"
|
||||||
|
"cyberstrike-ai/internal/config"
|
||||||
"cyberstrike-ai/internal/database"
|
"cyberstrike-ai/internal/database"
|
||||||
|
"cyberstrike-ai/internal/security"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"net/http"
|
"net/http"
|
||||||
"net/http/httptest"
|
"net/http/httptest"
|
||||||
@@ -69,6 +72,40 @@ func TestRBACAssignResourceBatchIsAtomicAndLegacyCompatible(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRBACAssignResourceAutoDetectsActualType(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
db, err := database.NewDB(filepath.Join(t.TempDir(), "rbac-auto-detect.db"), zap.NewNop())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
user, err := db.CreateRBACUser("auto-member", "Auto Member", "hash", true, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
project, err := db.CreateProject(&database.Project{Name: "auto-project"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
h := NewRBACHandler(db, zap.NewNop())
|
||||||
|
router := gin.New()
|
||||||
|
router.POST("/api/rbac/resource-assignments", h.AssignResource)
|
||||||
|
|
||||||
|
response := performRBACJSONRequest(t, router, map[string]interface{}{
|
||||||
|
"user_id": user.ID, "resource_type": "conversation", "resource_ids": []string{project.ID}, "auto_detect": true,
|
||||||
|
})
|
||||||
|
if response.Code != http.StatusOK {
|
||||||
|
t.Fatalf("auto-detect status = %d, body = %s", response.Code, response.Body.String())
|
||||||
|
}
|
||||||
|
rows, err := db.ListRBACResourceAssignments(user.ID)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(rows) != 1 || rows[0].ResourceType != "project" || rows[0].ResourceID != project.ID {
|
||||||
|
t.Fatalf("auto-detected assignments = %#v, want project/%s", rows, project.ID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestRBACAssignableResourcesArePaged(t *testing.T) {
|
func TestRBACAssignableResourcesArePaged(t *testing.T) {
|
||||||
gin.SetMode(gin.TestMode)
|
gin.SetMode(gin.TestMode)
|
||||||
db, err := database.NewDB(filepath.Join(t.TempDir(), "rbac-picker.db"), zap.NewNop())
|
db, err := database.NewDB(filepath.Join(t.TempDir(), "rbac-picker.db"), zap.NewNop())
|
||||||
@@ -104,6 +141,60 @@ func TestRBACAssignableResourcesArePaged(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestRBACDeleteResourceAssignmentAuditsTargetResource(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
db, err := database.NewDB(filepath.Join(t.TempDir(), "rbac-revoke-audit.db"), zap.NewNop())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
defer db.Close()
|
||||||
|
user, err := db.CreateRBACUser("audit-member", "Audit Member", "hash", true, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
project, err := db.CreateProject(&database.Project{Name: "audit-project"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := db.AssignResourcesToUser(user.ID, "project", []string{project.ID}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
assignments, err := db.ListRBACResourceAssignments(user.ID)
|
||||||
|
if err != nil || len(assignments) != 1 {
|
||||||
|
t.Fatalf("assignments = %#v, err = %v", assignments, err)
|
||||||
|
}
|
||||||
|
|
||||||
|
h := NewRBACHandler(db, zap.NewNop())
|
||||||
|
h.SetAudit(audit.NewService(db, &config.Config{}, zap.NewNop()))
|
||||||
|
router := gin.New()
|
||||||
|
router.Use(func(c *gin.Context) {
|
||||||
|
c.Set(security.ContextUsernameKey, "operator-user")
|
||||||
|
c.Next()
|
||||||
|
})
|
||||||
|
router.DELETE("/api/rbac/resource-assignments/:id", h.DeleteResourceAssignment)
|
||||||
|
request := httptest.NewRequest(http.MethodDelete, "/api/rbac/resource-assignments/"+assignments[0].ID, nil)
|
||||||
|
recorder := httptest.NewRecorder()
|
||||||
|
router.ServeHTTP(recorder, request)
|
||||||
|
if recorder.Code != http.StatusOK {
|
||||||
|
t.Fatalf("status = %d, body = %s", recorder.Code, recorder.Body.String())
|
||||||
|
}
|
||||||
|
|
||||||
|
logs, err := db.ListAuditLogs(database.ListAuditLogsFilter{Category: "rbac", RelatedUserID: user.ID, Limit: 10})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(logs) != 1 {
|
||||||
|
t.Fatalf("audit logs = %#v, want one member-related revoke", logs)
|
||||||
|
}
|
||||||
|
log := logs[0]
|
||||||
|
if log.Action != "delete_resource_assignment" || log.Actor != "operator-user" || log.ResourceType != "project" || log.ResourceID != project.ID {
|
||||||
|
t.Fatalf("audit log = %#v", log)
|
||||||
|
}
|
||||||
|
if log.Detail["user_id"] != user.ID || log.Detail["assignment_id"] != assignments[0].ID {
|
||||||
|
t.Fatalf("audit detail = %#v", log.Detail)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func performRBACJSONRequest(t *testing.T, router http.Handler, payload map[string]interface{}) *httptest.ResponseRecorder {
|
func performRBACJSONRequest(t *testing.T, router http.Handler, payload map[string]interface{}) *httptest.ResponseRecorder {
|
||||||
t.Helper()
|
t.Helper()
|
||||||
body, err := json.Marshal(payload)
|
body, err := json.Marshal(payload)
|
||||||
|
|||||||
+417
-63
@@ -41,11 +41,14 @@ const (
|
|||||||
robotCmdContinue = "继续"
|
robotCmdContinue = "继续"
|
||||||
robotCmdNew = "新对话"
|
robotCmdNew = "新对话"
|
||||||
robotCmdClear = "清空"
|
robotCmdClear = "清空"
|
||||||
robotCmdCurrent = "当前"
|
robotCmdStatus = "状态"
|
||||||
robotCmdStop = "停止"
|
robotCmdStop = "停止"
|
||||||
robotCmdRoles = "角色"
|
robotCmdRoles = "角色"
|
||||||
robotCmdRolesList = "角色列表"
|
robotCmdRolesList = "角色列表"
|
||||||
robotCmdSwitchRole = "切换角色"
|
robotCmdSwitchRole = "切换角色"
|
||||||
|
robotCmdModes = "模式"
|
||||||
|
robotCmdModesList = "模式列表"
|
||||||
|
robotCmdSwitchMode = "切换模式"
|
||||||
robotCmdDelete = "删除"
|
robotCmdDelete = "删除"
|
||||||
robotCmdVersion = "版本"
|
robotCmdVersion = "版本"
|
||||||
robotCmdProjects = "项目"
|
robotCmdProjects = "项目"
|
||||||
@@ -56,35 +59,54 @@ const (
|
|||||||
robotCmdBindUser = "绑定"
|
robotCmdBindUser = "绑定"
|
||||||
robotCmdUnbindUser = "解绑"
|
robotCmdUnbindUser = "解绑"
|
||||||
robotCmdIdentity = "身份"
|
robotCmdIdentity = "身份"
|
||||||
|
robotCmdTask = "任务"
|
||||||
|
robotCmdRename = "重命名"
|
||||||
|
robotCmdPermissions = "权限"
|
||||||
|
robotCmdDoctor = "诊断"
|
||||||
|
robotCmdConfirm = "确认"
|
||||||
|
robotCmdCancel = "取消"
|
||||||
|
robotCmdVulnAlerts = "漏洞提醒"
|
||||||
robotBindingCodeTTL = 5 * time.Minute
|
robotBindingCodeTTL = 5 * time.Minute
|
||||||
)
|
)
|
||||||
|
|
||||||
|
type robotPendingConfirmation struct {
|
||||||
|
Action string
|
||||||
|
Target string
|
||||||
|
ExpiresAt time.Time
|
||||||
|
}
|
||||||
|
|
||||||
// RobotHandler 企业微信/钉钉/飞书等机器人回调处理
|
// RobotHandler 企业微信/钉钉/飞书等机器人回调处理
|
||||||
type RobotHandler struct {
|
type RobotHandler struct {
|
||||||
config *config.Config
|
config *config.Config
|
||||||
db *database.DB
|
db *database.DB
|
||||||
agentHandler *AgentHandler
|
agentHandler *AgentHandler
|
||||||
logger *zap.Logger
|
logger *zap.Logger
|
||||||
mu sync.RWMutex
|
mu sync.RWMutex
|
||||||
sessions map[string]string // key: "platform_userID", value: conversationID
|
sessions map[string]string // key: "platform_userID", value: conversationID
|
||||||
sessionRoles map[string]string // key: "platform_userID", value: roleName(默认"默认")
|
sessionRoles map[string]string // key: "platform_userID", value: roleName(默认"默认")
|
||||||
cancelMu sync.Mutex // 保护 runningCancels
|
sessionModes map[string]string // key: "platform_userID", value: agent mode
|
||||||
runningCancels map[string]context.CancelFunc // key: "platform_userID", 用于停止命令中断任务
|
cancelMu sync.Mutex // 保护 runningCancels
|
||||||
wecomReplay map[string]time.Time
|
runningCancels map[string]context.CancelFunc // key: "platform_userID", 用于停止命令中断任务
|
||||||
audit *audit.Service
|
wecomReplay map[string]time.Time
|
||||||
|
pendingConfirmations map[string]robotPendingConfirmation
|
||||||
|
alertWake chan struct{}
|
||||||
|
audit *audit.Service
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewRobotHandler 创建机器人处理器
|
// NewRobotHandler 创建机器人处理器
|
||||||
func NewRobotHandler(cfg *config.Config, db *database.DB, agentHandler *AgentHandler, logger *zap.Logger) *RobotHandler {
|
func NewRobotHandler(cfg *config.Config, db *database.DB, agentHandler *AgentHandler, logger *zap.Logger) *RobotHandler {
|
||||||
return &RobotHandler{
|
return &RobotHandler{
|
||||||
config: cfg,
|
config: cfg,
|
||||||
db: db,
|
db: db,
|
||||||
agentHandler: agentHandler,
|
agentHandler: agentHandler,
|
||||||
logger: logger,
|
logger: logger,
|
||||||
sessions: make(map[string]string),
|
sessions: make(map[string]string),
|
||||||
sessionRoles: make(map[string]string),
|
sessionRoles: make(map[string]string),
|
||||||
runningCancels: make(map[string]context.CancelFunc),
|
sessionModes: make(map[string]string),
|
||||||
wecomReplay: make(map[string]time.Time),
|
runningCancels: make(map[string]context.CancelFunc),
|
||||||
|
wecomReplay: make(map[string]time.Time),
|
||||||
|
pendingConfirmations: make(map[string]robotPendingConfirmation),
|
||||||
|
alertWake: make(chan struct{}, 1),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -175,26 +197,26 @@ func (h *RobotHandler) robotAccessDeniedMessage(platform string) string {
|
|||||||
return "当前平台账号尚未绑定 CyberStrikeAI 用户。请先在网页端生成绑定码,然后发送:绑定 XXXX-XXXX"
|
return "当前平台账号尚未绑定 CyberStrikeAI 用户。请先在网页端生成绑定码,然后发送:绑定 XXXX-XXXX"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *RobotHandler) loadSessionBinding(sk string) (convID, role string) {
|
func (h *RobotHandler) loadSessionBinding(sk string) (convID, role, agentMode string) {
|
||||||
if h.db == nil || strings.TrimSpace(sk) == "" {
|
if h.db == nil || strings.TrimSpace(sk) == "" {
|
||||||
return "", ""
|
return "", "", ""
|
||||||
}
|
}
|
||||||
binding, err := h.db.GetRobotSessionBinding(sk)
|
binding, err := h.db.GetRobotSessionBinding(sk)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
h.logger.Warn("读取机器人会话绑定失败", zap.String("session_key", sk), zap.Error(err))
|
h.logger.Warn("读取机器人会话绑定失败", zap.String("session_key", sk), zap.Error(err))
|
||||||
return "", ""
|
return "", "", ""
|
||||||
}
|
}
|
||||||
if binding == nil {
|
if binding == nil {
|
||||||
return "", ""
|
return "", "", ""
|
||||||
}
|
}
|
||||||
return binding.ConversationID, binding.RoleName
|
return binding.ConversationID, binding.RoleName, binding.AgentMode
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *RobotHandler) persistSessionBinding(sk, convID, role string) {
|
func (h *RobotHandler) persistSessionBinding(sk, convID, role, agentMode string) {
|
||||||
if h.db == nil || strings.TrimSpace(sk) == "" || strings.TrimSpace(convID) == "" {
|
if h.db == nil || strings.TrimSpace(sk) == "" || strings.TrimSpace(convID) == "" {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
if err := h.db.UpsertRobotSessionBinding(sk, convID, role); err != nil {
|
if err := h.db.UpsertRobotSessionBinding(sk, convID, role, agentMode); err != nil {
|
||||||
h.logger.Warn("写入机器人会话绑定失败", zap.String("session_key", sk), zap.Error(err))
|
h.logger.Warn("写入机器人会话绑定失败", zap.String("session_key", sk), zap.Error(err))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -219,7 +241,7 @@ func (h *RobotHandler) getOrCreateConversation(platform, userID, title string, a
|
|||||||
if convID != "" && access.Permissions["chat:read"] && h.db.UserCanAccessResource(ownerID, readScope, "conversation", convID) {
|
if convID != "" && access.Permissions["chat:read"] && h.db.UserCanAccessResource(ownerID, readScope, "conversation", convID) {
|
||||||
return convID, false
|
return convID, false
|
||||||
}
|
}
|
||||||
if persistedConvID, persistedRole := h.loadSessionBinding(sk); strings.TrimSpace(persistedConvID) != "" {
|
if persistedConvID, persistedRole, persistedMode := h.loadSessionBinding(sk); strings.TrimSpace(persistedConvID) != "" {
|
||||||
if !access.Permissions["chat:read"] || !h.db.UserCanAccessResource(ownerID, readScope, "conversation", persistedConvID) {
|
if !access.Permissions["chat:read"] || !h.db.UserCanAccessResource(ownerID, readScope, "conversation", persistedConvID) {
|
||||||
h.deleteSessionBinding(sk)
|
h.deleteSessionBinding(sk)
|
||||||
} else {
|
} else {
|
||||||
@@ -229,6 +251,9 @@ func (h *RobotHandler) getOrCreateConversation(platform, userID, title string, a
|
|||||||
if strings.TrimSpace(persistedRole) != "" {
|
if strings.TrimSpace(persistedRole) != "" {
|
||||||
h.sessionRoles[sk] = persistedRole
|
h.sessionRoles[sk] = persistedRole
|
||||||
}
|
}
|
||||||
|
if strings.TrimSpace(persistedMode) != "" {
|
||||||
|
h.sessionModes[sk] = config.NormalizeAgentMode(persistedMode)
|
||||||
|
}
|
||||||
h.mu.Unlock()
|
h.mu.Unlock()
|
||||||
return persistedConvID, false
|
return persistedConvID, false
|
||||||
}
|
}
|
||||||
@@ -256,9 +281,13 @@ func (h *RobotHandler) getOrCreateConversation(platform, userID, title string, a
|
|||||||
_ = h.db.SetResourceOwner("conversation", convID, ownerID)
|
_ = h.db.SetResourceOwner("conversation", convID, ownerID)
|
||||||
h.mu.Lock()
|
h.mu.Lock()
|
||||||
role := h.sessionRoles[sk]
|
role := h.sessionRoles[sk]
|
||||||
|
agentMode := h.sessionModes[sk]
|
||||||
h.sessions[sk] = convID
|
h.sessions[sk] = convID
|
||||||
h.mu.Unlock()
|
h.mu.Unlock()
|
||||||
h.persistSessionBinding(sk, convID, role)
|
if agentMode == "" {
|
||||||
|
agentMode = config.NormalizeRobotAgentMode(h.config.MultiAgent)
|
||||||
|
}
|
||||||
|
h.persistSessionBinding(sk, convID, role, agentMode)
|
||||||
return convID, true
|
return convID, true
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -267,9 +296,10 @@ func (h *RobotHandler) setConversation(platform, userID, convID string) {
|
|||||||
sk := h.sessionKey(platform, userID)
|
sk := h.sessionKey(platform, userID)
|
||||||
h.mu.Lock()
|
h.mu.Lock()
|
||||||
role := h.sessionRoles[sk]
|
role := h.sessionRoles[sk]
|
||||||
|
agentMode := h.sessionModes[sk]
|
||||||
h.sessions[sk] = convID
|
h.sessions[sk] = convID
|
||||||
h.mu.Unlock()
|
h.mu.Unlock()
|
||||||
h.persistSessionBinding(sk, convID, role)
|
h.persistSessionBinding(sk, convID, role, agentMode)
|
||||||
}
|
}
|
||||||
|
|
||||||
// getRole 获取当前用户使用的角色,未设置时返回"默认"
|
// getRole 获取当前用户使用的角色,未设置时返回"默认"
|
||||||
@@ -281,7 +311,7 @@ func (h *RobotHandler) getRole(platform, userID string) string {
|
|||||||
if strings.TrimSpace(role) != "" {
|
if strings.TrimSpace(role) != "" {
|
||||||
return role
|
return role
|
||||||
}
|
}
|
||||||
if _, persistedRole := h.loadSessionBinding(sk); strings.TrimSpace(persistedRole) != "" {
|
if _, persistedRole, _ := h.loadSessionBinding(sk); strings.TrimSpace(persistedRole) != "" {
|
||||||
h.mu.Lock()
|
h.mu.Lock()
|
||||||
h.sessionRoles[sk] = persistedRole
|
h.sessionRoles[sk] = persistedRole
|
||||||
h.mu.Unlock()
|
h.mu.Unlock()
|
||||||
@@ -296,8 +326,38 @@ func (h *RobotHandler) setRole(platform, userID, roleName string) {
|
|||||||
h.mu.Lock()
|
h.mu.Lock()
|
||||||
h.sessionRoles[sk] = roleName
|
h.sessionRoles[sk] = roleName
|
||||||
convID := h.sessions[sk]
|
convID := h.sessions[sk]
|
||||||
|
agentMode := h.sessionModes[sk]
|
||||||
h.mu.Unlock()
|
h.mu.Unlock()
|
||||||
h.persistSessionBinding(sk, convID, roleName)
|
h.persistSessionBinding(sk, convID, roleName, agentMode)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) getAgentMode(platform, userID string) string {
|
||||||
|
sk := h.sessionKey(platform, userID)
|
||||||
|
h.mu.RLock()
|
||||||
|
mode := h.sessionModes[sk]
|
||||||
|
h.mu.RUnlock()
|
||||||
|
if mode != "" {
|
||||||
|
return config.NormalizeAgentMode(mode)
|
||||||
|
}
|
||||||
|
if _, _, persistedMode := h.loadSessionBinding(sk); persistedMode != "" {
|
||||||
|
mode = config.NormalizeAgentMode(persistedMode)
|
||||||
|
h.mu.Lock()
|
||||||
|
h.sessionModes[sk] = mode
|
||||||
|
h.mu.Unlock()
|
||||||
|
return mode
|
||||||
|
}
|
||||||
|
return config.NormalizeRobotAgentMode(h.config.MultiAgent)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) setAgentMode(platform, userID, mode string) {
|
||||||
|
sk := h.sessionKey(platform, userID)
|
||||||
|
mode = config.NormalizeAgentMode(mode)
|
||||||
|
h.mu.Lock()
|
||||||
|
h.sessionModes[sk] = mode
|
||||||
|
convID := h.sessions[sk]
|
||||||
|
role := h.sessionRoles[sk]
|
||||||
|
h.mu.Unlock()
|
||||||
|
h.persistSessionBinding(sk, convID, role, mode)
|
||||||
}
|
}
|
||||||
|
|
||||||
// clearConversation 清空当前会话(切换到新对话)
|
// clearConversation 清空当前会话(切换到新对话)
|
||||||
@@ -379,7 +439,8 @@ func (h *RobotHandler) HandleMessage(platform, userID, text string) (reply strin
|
|||||||
h.cancelMu.Unlock()
|
h.cancelMu.Unlock()
|
||||||
}()
|
}()
|
||||||
role := h.getRole(platform, userID)
|
role := h.getRole(platform, userID)
|
||||||
resp, newConvID, err := h.agentHandler.ProcessMessageForRobot(ctx, platform, robotPrincipal(access), convID, text, role)
|
agentMode := h.getAgentMode(platform, userID)
|
||||||
|
resp, newConvID, err := h.agentHandler.ProcessMessageForRobot(ctx, platform, robotPrincipal(access), convID, text, role, agentMode)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
h.logger.Warn("机器人 Agent 执行失败", zap.String("platform", platform), zap.String("userID", userID), zap.Error(err))
|
h.logger.Warn("机器人 Agent 执行失败", zap.String("platform", platform), zap.String("userID", userID), zap.Error(err))
|
||||||
if errors.Is(err, context.Canceled) {
|
if errors.Is(err, context.Canceled) {
|
||||||
@@ -401,32 +462,54 @@ func (h *RobotHandler) robotMessageTimeout() time.Duration {
|
|||||||
return 10 * time.Hour
|
return 10 * time.Hour
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *RobotHandler) cmdHelp() string {
|
func (h *RobotHandler) cmdHelp(platform, userID string) string {
|
||||||
|
access, _ := h.resolveRobotAccess(platform, userID)
|
||||||
|
can := func(permission string) bool {
|
||||||
|
return access != nil && access.Permissions[permission]
|
||||||
|
}
|
||||||
var b strings.Builder
|
var b strings.Builder
|
||||||
b.WriteString("【CyberStrikeAI 机器人命令】\n\n")
|
b.WriteString("【CyberStrikeAI 机器人命令】\n\n")
|
||||||
b.WriteString("【通用 General】\n")
|
b.WriteString("【通用 General】\n")
|
||||||
b.WriteString("· 帮助 / help — 显示本帮助\n")
|
b.WriteString("· 帮助 / help — 显示本帮助\n")
|
||||||
b.WriteString("· 版本 / version — 显示当前版本号\n")
|
b.WriteString("· 版本 / version — 显示当前版本号\n")
|
||||||
b.WriteString("· 绑定 <绑定码> / bind <code> — 绑定网页端 RBAC 用户\n")
|
b.WriteString("· 绑定 <绑定码> / bind <code> — 绑定网页端 RBAC 用户\n")
|
||||||
b.WriteString("· 解绑 / unbind — 解除当前平台账号绑定\n")
|
b.WriteString("· 解绑 / unbind — 请求解除账号绑定(需确认)\n")
|
||||||
b.WriteString("· 身份 / whoami — 显示平台发送者、鉴权模式及当前实际 RBAC 身份\n")
|
b.WriteString("· 身份 / whoami — 显示平台发送者、鉴权模式及当前实际 RBAC 身份\n")
|
||||||
b.WriteString("\n【对话 Conversation】\n")
|
if can("chat:read") || can("chat:write") || can("chat:delete") {
|
||||||
b.WriteString("· 列表 / list — 列出所有对话标题与 ID\n")
|
b.WriteString("\n【对话 Conversation】\n")
|
||||||
b.WriteString("· 切换 <ID> / switch <ID> — 指定对话继续\n")
|
if can("chat:read") {
|
||||||
b.WriteString("· 新对话 / new — 开启新对话\n")
|
b.WriteString("· 列表 / list — 列出所有对话标题与 ID\n· 切换 <ID> / switch <ID> — 指定对话继续\n· 状态 / status — 汇总当前选择\n· 任务 / task — 查看当前任务状态\n")
|
||||||
b.WriteString("· 清空 / clear — 清空当前上下文\n")
|
}
|
||||||
b.WriteString("· 当前 / current — 显示当前对话、角色与项目\n")
|
if can("chat:write") {
|
||||||
b.WriteString("· 停止 / stop — 中断当前任务\n")
|
b.WriteString("· 新对话 / new;清空 / clear — 开启新对话\n· 重命名 <名称> / rename <name> — 修改当前对话标题\n")
|
||||||
b.WriteString("· 删除 <ID> / delete <ID> — 删除指定对话\n")
|
}
|
||||||
b.WriteString("\n【角色 Role】\n")
|
if can("chat:delete") {
|
||||||
b.WriteString("· 角色 / roles — 列出所有可用角色\n")
|
b.WriteString("· 删除 <ID> / delete <ID> — 删除指定对话(需确认)\n")
|
||||||
b.WriteString("· 角色 <名> / role <name> — 切换当前角色\n")
|
}
|
||||||
if h.projectsEnabled() {
|
}
|
||||||
|
if can("roles:read") {
|
||||||
|
b.WriteString("\n【角色 Role】\n· 角色 / roles — 列出所有可用角色\n· 角色 <名> / role <name> — 切换当前角色\n")
|
||||||
|
}
|
||||||
|
if can("agent:execute") {
|
||||||
|
b.WriteString("\n【模式 Mode】\n· 模式 / modes — 列出对话模式与当前选择\n· 模式 <名称> / mode <name> — 切换对话模式\n· 停止 / stop — 中断当前任务\n")
|
||||||
|
}
|
||||||
|
if can("vulnerability:read") {
|
||||||
|
b.WriteString("\n【漏洞提醒 Vulnerability alerts】\n· 漏洞提醒 — 查看订阅状态\n· 漏洞提醒 开启 / vuln alerts on — 开启提醒\n· 漏洞提醒 仅严重|高危以上|中危以上 / vuln alerts critical|high|medium — 设置最低级别\n· 漏洞提醒 关闭 / vuln alerts off — 关闭提醒\n")
|
||||||
|
}
|
||||||
|
b.WriteString("\n【诊断 Diagnostics】\n")
|
||||||
|
b.WriteString("· 权限 / permissions — 查看当前业务权限\n")
|
||||||
|
if can("config:read") {
|
||||||
|
b.WriteString("· 诊断 / doctor — 检查机器人关键配置状态\n")
|
||||||
|
}
|
||||||
|
b.WriteString("· 确认 / confirm;取消 / cancel — 处理高风险操作确认\n")
|
||||||
|
if h.projectsEnabled() && (can("project:read") || can("project:write")) {
|
||||||
b.WriteString("\n【项目 Project】\n")
|
b.WriteString("\n【项目 Project】\n")
|
||||||
b.WriteString("· 项目 / projects — 列出所有项目\n")
|
if can("project:read") {
|
||||||
b.WriteString("· 新建项目 <名称> / new project <name> — 创建并绑定当前对话\n")
|
b.WriteString("· 项目 / projects — 列出所有项目\n")
|
||||||
b.WriteString("· 绑定项目 <ID或名称> / bind project <ID|name> — 绑定到已有项目\n")
|
}
|
||||||
b.WriteString("· 解除项目 / unbind project — 解除项目绑定\n")
|
if can("project:write") {
|
||||||
|
b.WriteString("· 新建项目 <名称> / new project <name> — 创建并绑定当前对话\n· 绑定项目 <ID或名称> / bind project <ID|name> — 绑定已有项目\n· 解除项目 / unbind project — 解除项目绑定\n")
|
||||||
|
}
|
||||||
}
|
}
|
||||||
b.WriteString("\n──────────────\n")
|
b.WriteString("\n──────────────\n")
|
||||||
b.WriteString("除以上命令外,直接输入内容将发送给 AI 进行渗透测试/安全分析。")
|
b.WriteString("除以上命令外,直接输入内容将发送给 AI 进行渗透测试/安全分析。")
|
||||||
@@ -575,7 +658,7 @@ func (h *RobotHandler) cmdUnbindProject(platform, userID string) string {
|
|||||||
convID := h.sessions[sk]
|
convID := h.sessions[sk]
|
||||||
h.mu.RUnlock()
|
h.mu.RUnlock()
|
||||||
if convID == "" {
|
if convID == "" {
|
||||||
if persistedConvID, _ := h.loadSessionBinding(sk); persistedConvID != "" {
|
if persistedConvID, _, _ := h.loadSessionBinding(sk); persistedConvID != "" {
|
||||||
convID = persistedConvID
|
convID = persistedConvID
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -673,12 +756,10 @@ func (h *RobotHandler) cmdStop(platform, userID string) string {
|
|||||||
return "已停止当前任务。"
|
return "已停止当前任务。"
|
||||||
}
|
}
|
||||||
|
|
||||||
func (h *RobotHandler) cmdCurrent(platform, userID string) string {
|
func (h *RobotHandler) cmdStatus(platform, userID string) string {
|
||||||
h.mu.RLock()
|
convID := h.currentConversationID(platform, userID)
|
||||||
convID := h.sessions[h.sessionKey(platform, userID)]
|
|
||||||
h.mu.RUnlock()
|
|
||||||
if convID == "" {
|
if convID == "" {
|
||||||
return "当前没有进行中的对话。发送任意内容将创建新对话。"
|
return fmt.Sprintf("【当前状态】\n当前对话: 无\n当前角色: %s\n当前模式: %s\n当前项目: 无\n\n发送任意内容将创建新对话。", h.getRole(platform, userID), robotAgentModeLabel(h.getAgentMode(platform, userID)))
|
||||||
}
|
}
|
||||||
access, err := h.resolveRobotAccess(platform, userID)
|
access, err := h.resolveRobotAccess(platform, userID)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -692,14 +773,122 @@ func (h *RobotHandler) cmdCurrent(platform, userID string) string {
|
|||||||
return "当前对话 ID: " + convID + "(获取标题失败)"
|
return "当前对话 ID: " + convID + "(获取标题失败)"
|
||||||
}
|
}
|
||||||
role := h.getRole(platform, userID)
|
role := h.getRole(platform, userID)
|
||||||
reply := fmt.Sprintf("当前对话:「%s」\nID: %s\n当前角色: %s", conv.Title, conv.ID, role)
|
reply := fmt.Sprintf("【当前状态】\n当前对话: %s\n对话 ID: %s\n当前模式: %s\n当前角色: %s", conv.Title, conv.ID, robotAgentModeLabel(h.getAgentMode(platform, userID)), role)
|
||||||
if h.projectsEnabled() {
|
if h.projectsEnabled() {
|
||||||
projectID, _ := h.db.GetConversationProjectID(conv.ID)
|
projectID, _ := h.db.GetConversationProjectID(conv.ID)
|
||||||
reply += "\n当前项目: " + h.formatProjectLabel(projectID)
|
reply += "\n当前项目: " + h.formatProjectLabel(projectID)
|
||||||
|
} else {
|
||||||
|
reply += "\n当前项目: 未启用"
|
||||||
}
|
}
|
||||||
return reply
|
return reply
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) currentConversationID(platform, userID string) string {
|
||||||
|
sk := h.sessionKey(platform, userID)
|
||||||
|
h.mu.RLock()
|
||||||
|
convID := h.sessions[sk]
|
||||||
|
h.mu.RUnlock()
|
||||||
|
if convID != "" {
|
||||||
|
return convID
|
||||||
|
}
|
||||||
|
persistedConvID, persistedRole, persistedMode := h.loadSessionBinding(sk)
|
||||||
|
if persistedConvID == "" {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
h.mu.Lock()
|
||||||
|
h.sessions[sk] = persistedConvID
|
||||||
|
h.sessionRoles[sk] = persistedRole
|
||||||
|
h.sessionModes[sk] = config.NormalizeAgentMode(persistedMode)
|
||||||
|
h.mu.Unlock()
|
||||||
|
return persistedConvID
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) cmdTask(platform, userID string) string {
|
||||||
|
convID := h.currentConversationID(platform, userID)
|
||||||
|
if convID == "" {
|
||||||
|
return "【任务状态】\n当前没有对话,也没有正在执行的任务。"
|
||||||
|
}
|
||||||
|
if h.agentHandler == nil || h.agentHandler.tasks == nil {
|
||||||
|
return "任务状态服务不可用。"
|
||||||
|
}
|
||||||
|
task := h.agentHandler.tasks.GetTaskSnapshot(convID)
|
||||||
|
if task == nil {
|
||||||
|
return "【任务状态】\n状态: 空闲\n当前没有正在执行的任务。"
|
||||||
|
}
|
||||||
|
elapsed := time.Since(task.StartedAt).Round(time.Second)
|
||||||
|
return fmt.Sprintf("【任务状态】\n状态: %s\n已运行: %s\n对话 ID: %s\n模式: %s\n可用操作: 停止 / stop", task.Status, elapsed, convID, robotAgentModeLabel(h.getAgentMode(platform, userID)))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) cmdRename(platform, userID, title string) string {
|
||||||
|
title = strings.TrimSpace(title)
|
||||||
|
if title == "" {
|
||||||
|
return "请指定新标题,例如:重命名 外网资产排查"
|
||||||
|
}
|
||||||
|
title = safeTruncateString(title, 100)
|
||||||
|
convID := h.currentConversationID(platform, userID)
|
||||||
|
if convID == "" {
|
||||||
|
return "当前没有对话,无法重命名。"
|
||||||
|
}
|
||||||
|
access, err := h.resolveRobotAccess(platform, userID)
|
||||||
|
if err != nil || !h.db.UserCanAccessResource(access.User.ID, robotPrincipal(access).ScopeFor("chat:write"), "conversation", convID) {
|
||||||
|
return "当前对话不存在或无权修改。"
|
||||||
|
}
|
||||||
|
if err := h.db.UpdateConversationTitle(convID, title); err != nil {
|
||||||
|
return "重命名失败: " + err.Error()
|
||||||
|
}
|
||||||
|
h.recordRobotCommandAudit(access, platform, "conversation_rename", "conversation", convID, "机器人重命名当前对话")
|
||||||
|
return fmt.Sprintf("已将当前对话重命名为:「%s」", title)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) cmdPermissions(platform, userID string) string {
|
||||||
|
access, err := h.resolveRobotAccess(platform, userID)
|
||||||
|
if err != nil {
|
||||||
|
return h.robotAccessDeniedMessage(platform)
|
||||||
|
}
|
||||||
|
allowed := func(permission string) string {
|
||||||
|
if access.Permissions[permission] {
|
||||||
|
return "允许"
|
||||||
|
}
|
||||||
|
return "不允许"
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("【当前权限】\n执行 Agent: %s\n读取对话: %s\n编辑对话: %s\n删除对话: %s\n读取角色: %s\n读取项目: %s\n编辑项目: %s\n资源范围: %s", allowed("agent:execute"), allowed("chat:read"), allowed("chat:write"), allowed("chat:delete"), allowed("roles:read"), allowed("project:read"), allowed("project:write"), access.Scope)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) cmdDoctor() string {
|
||||||
|
configured := func(ok bool) string {
|
||||||
|
if ok {
|
||||||
|
return "正常"
|
||||||
|
}
|
||||||
|
return "未配置"
|
||||||
|
}
|
||||||
|
enabled := func(ok bool) string {
|
||||||
|
if ok {
|
||||||
|
return "已启用"
|
||||||
|
}
|
||||||
|
return "已关闭"
|
||||||
|
}
|
||||||
|
enabledInternalTools := 0
|
||||||
|
for _, tool := range h.config.Security.Tools {
|
||||||
|
if tool.Enabled {
|
||||||
|
enabledInternalTools++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
enabledExternal := 0
|
||||||
|
for _, server := range h.config.ExternalMCP.Servers {
|
||||||
|
if server.ExternalMCPEnable && !server.Disabled {
|
||||||
|
enabledExternal++
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("【配置诊断】\n主模型: %s\nEino 多代理: %s\n内置 MCP 工具: %d/%d 个已启用\nHTTP MCP 服务: %s\n外部 MCP: %d 个已启用\n知识库: %s\n项目功能: %s\n说明: 内置工具不依赖 HTTP MCP 服务;此命令只检查配置,不主动探测外部服务。", configured(strings.TrimSpace(h.config.OpenAI.Model) != "" && strings.TrimSpace(h.config.OpenAI.BaseURL) != ""), enabled(h.config.MultiAgent.Enabled), enabledInternalTools, len(h.config.Security.Tools), enabled(h.config.MCP.Enabled), enabledExternal, enabled(h.config.Knowledge.Enabled), enabled(h.config.Project.Enabled))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) recordRobotCommandAudit(access *database.RBACAccess, platform, action, resourceType, resourceID, message string) {
|
||||||
|
if h.audit == nil || access == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
h.audit.RecordSystem(audit.Entry{Category: "robot", Action: action, Result: "success", Actor: access.User.Username, ResourceType: resourceType, ResourceID: resourceID, Message: message + "(" + platform + ")"})
|
||||||
|
}
|
||||||
|
|
||||||
func (h *RobotHandler) cmdRoles() string {
|
func (h *RobotHandler) cmdRoles() string {
|
||||||
if h.config.Roles == nil || len(h.config.Roles) == 0 {
|
if h.config.Roles == nil || len(h.config.Roles) == 0 {
|
||||||
return "暂无可用角色。"
|
return "暂无可用角色。"
|
||||||
@@ -753,6 +942,55 @@ func (h *RobotHandler) cmdSwitchRole(platform, userID, roleName string) string {
|
|||||||
return fmt.Sprintf("已切换到角色:「%s」\n%s", roleName, role.Description)
|
return fmt.Sprintf("已切换到角色:「%s」\n%s", roleName, role.Description)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func robotAgentModeLabel(mode string) string {
|
||||||
|
switch config.NormalizeAgentMode(mode) {
|
||||||
|
case "deep":
|
||||||
|
return "Deep"
|
||||||
|
case "plan_execute":
|
||||||
|
return "Plan-Execute"
|
||||||
|
case "supervisor":
|
||||||
|
return "Supervisor"
|
||||||
|
default:
|
||||||
|
return "Eino 单代理"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseRobotAgentMode(input string) (string, bool) {
|
||||||
|
switch strings.ToLower(strings.TrimSpace(input)) {
|
||||||
|
case "eino_single", "eino-single", "single", "单代理", "eino单代理", "eino 单代理":
|
||||||
|
return "eino_single", true
|
||||||
|
case "deep":
|
||||||
|
return "deep", true
|
||||||
|
case "plan_execute", "plan-execute", "planexecute", "pe":
|
||||||
|
return "plan_execute", true
|
||||||
|
case "supervisor", "super", "sv":
|
||||||
|
return "supervisor", true
|
||||||
|
default:
|
||||||
|
return "", false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) cmdModes(platform, userID string) string {
|
||||||
|
current := h.getAgentMode(platform, userID)
|
||||||
|
multiStatus := "可用"
|
||||||
|
if h.config == nil || !h.config.MultiAgent.Enabled {
|
||||||
|
multiStatus = "不可用(需在系统设置中启用 Eino 多代理)"
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("【对话模式】\n· Eino 单代理 — 可用\n· Deep — %s\n· Plan-Execute — %s\n· Supervisor — %s\n\n当前模式: %s\n切换示例:模式 deep", multiStatus, multiStatus, multiStatus, robotAgentModeLabel(current))
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) cmdSwitchMode(platform, userID, input string) string {
|
||||||
|
mode, ok := parseRobotAgentMode(input)
|
||||||
|
if !ok {
|
||||||
|
return fmt.Sprintf("不支持的对话模式「%s」。发送「模式」查看可用模式。", strings.TrimSpace(input))
|
||||||
|
}
|
||||||
|
if mode != "eino_single" && (h.config == nil || !h.config.MultiAgent.Enabled) {
|
||||||
|
return fmt.Sprintf("无法切换到 %s:请先在系统设置中启用 Eino 多代理。", robotAgentModeLabel(mode))
|
||||||
|
}
|
||||||
|
h.setAgentMode(platform, userID, mode)
|
||||||
|
return fmt.Sprintf("已切换对话模式:%s\n后续消息和新对话将使用该模式。", robotAgentModeLabel(mode))
|
||||||
|
}
|
||||||
|
|
||||||
func (h *RobotHandler) cmdDelete(platform, userID, convID string) string {
|
func (h *RobotHandler) cmdDelete(platform, userID, convID string) string {
|
||||||
if convID == "" {
|
if convID == "" {
|
||||||
return "请指定对话 ID,例如:删除 xxx-xxx-xxx"
|
return "请指定对话 ID,例如:删除 xxx-xxx-xxx"
|
||||||
@@ -764,6 +1002,15 @@ func (h *RobotHandler) cmdDelete(platform, userID, convID string) string {
|
|||||||
if !h.db.UserCanAccessResource(access.User.ID, robotPrincipal(access).ScopeFor("chat:delete"), "conversation", convID) {
|
if !h.db.UserCanAccessResource(access.User.ID, robotPrincipal(access).ScopeFor("chat:delete"), "conversation", convID) {
|
||||||
return "对话不存在或无权访问。"
|
return "对话不存在或无权访问。"
|
||||||
}
|
}
|
||||||
|
h.setPendingConfirmation(platform, userID, "delete_conversation", convID)
|
||||||
|
return fmt.Sprintf("⚠️ 即将删除对话 ID: %s\n此操作不可撤销。请在 2 分钟内发送「确认」继续,或发送「取消」。", convID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) executeDelete(platform, userID, convID string) string {
|
||||||
|
access, err := h.resolveRobotAccess(platform, userID)
|
||||||
|
if err != nil || !h.db.UserCanAccessResource(access.User.ID, robotPrincipal(access).ScopeFor("chat:delete"), "conversation", convID) {
|
||||||
|
return "对话不存在或无权删除。"
|
||||||
|
}
|
||||||
sk := h.sessionKey(platform, userID)
|
sk := h.sessionKey(platform, userID)
|
||||||
h.mu.RLock()
|
h.mu.RLock()
|
||||||
currentConvID := h.sessions[sk]
|
currentConvID := h.sessions[sk]
|
||||||
@@ -773,6 +1020,7 @@ func (h *RobotHandler) cmdDelete(platform, userID, convID string) string {
|
|||||||
h.mu.Lock()
|
h.mu.Lock()
|
||||||
delete(h.sessions, sk)
|
delete(h.sessions, sk)
|
||||||
delete(h.sessionRoles, sk)
|
delete(h.sessionRoles, sk)
|
||||||
|
delete(h.sessionModes, sk)
|
||||||
h.mu.Unlock()
|
h.mu.Unlock()
|
||||||
h.deleteSessionBinding(sk)
|
h.deleteSessionBinding(sk)
|
||||||
}
|
}
|
||||||
@@ -782,6 +1030,7 @@ func (h *RobotHandler) cmdDelete(platform, userID, convID string) string {
|
|||||||
if err := h.db.DeleteConversation(convID); err != nil {
|
if err := h.db.DeleteConversation(convID); err != nil {
|
||||||
return "删除失败: " + err.Error()
|
return "删除失败: " + err.Error()
|
||||||
}
|
}
|
||||||
|
h.recordRobotCommandAudit(access, platform, "conversation_delete", "conversation", convID, "机器人删除对话")
|
||||||
return fmt.Sprintf("已删除对话 ID: %s", convID)
|
return fmt.Sprintf("已删除对话 ID: %s", convID)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -844,9 +1093,10 @@ func robotCommandPermission(text string) (string, bool) {
|
|||||||
case text == robotCmdList || text == robotCmdListAlt || text == "list",
|
case text == robotCmdList || text == robotCmdListAlt || text == "list",
|
||||||
strings.HasPrefix(text, robotCmdSwitch+" "), strings.HasPrefix(text, robotCmdContinue+" "),
|
strings.HasPrefix(text, robotCmdSwitch+" "), strings.HasPrefix(text, robotCmdContinue+" "),
|
||||||
strings.HasPrefix(text, "switch "), strings.HasPrefix(text, "continue "),
|
strings.HasPrefix(text, "switch "), strings.HasPrefix(text, "continue "),
|
||||||
text == robotCmdCurrent || text == "current":
|
text == robotCmdStatus || text == "status", text == robotCmdTask || text == "task":
|
||||||
return "chat:read", true
|
return "chat:read", true
|
||||||
case text == robotCmdNew || text == "new", text == robotCmdClear || text == "clear":
|
case text == robotCmdNew || text == "new", text == robotCmdClear || text == "clear",
|
||||||
|
strings.HasPrefix(text, robotCmdRename+" "), strings.HasPrefix(text, "rename "):
|
||||||
return "chat:write", true
|
return "chat:write", true
|
||||||
case strings.HasPrefix(text, robotCmdDelete+" "), strings.HasPrefix(text, "delete "):
|
case strings.HasPrefix(text, robotCmdDelete+" "), strings.HasPrefix(text, "delete "):
|
||||||
return "chat:delete", true
|
return "chat:delete", true
|
||||||
@@ -855,8 +1105,20 @@ func robotCommandPermission(text string) (string, bool) {
|
|||||||
case text == robotCmdRoles || text == robotCmdRolesList || text == "roles",
|
case text == robotCmdRoles || text == robotCmdRolesList || text == "roles",
|
||||||
strings.HasPrefix(text, robotCmdRoles+" "), strings.HasPrefix(text, robotCmdSwitchRole+" "), strings.HasPrefix(text, "role "):
|
strings.HasPrefix(text, robotCmdRoles+" "), strings.HasPrefix(text, robotCmdSwitchRole+" "), strings.HasPrefix(text, "role "):
|
||||||
return "roles:read", true
|
return "roles:read", true
|
||||||
|
case text == robotCmdModes || text == robotCmdModesList || text == "modes",
|
||||||
|
strings.HasPrefix(text, robotCmdModes+" "), strings.HasPrefix(text, robotCmdSwitchMode+" "), strings.HasPrefix(text, "mode "):
|
||||||
|
return "agent:execute", true
|
||||||
|
case text == robotCmdPermissions || text == "permissions":
|
||||||
|
return "", true
|
||||||
|
case text == robotCmdConfirm || text == "confirm", text == robotCmdCancel || text == "cancel":
|
||||||
|
return "", true
|
||||||
|
case text == robotCmdDoctor || text == "doctor":
|
||||||
|
return "config:read", true
|
||||||
case text == robotCmdProjects || text == robotCmdProjectsList || text == "projects":
|
case text == robotCmdProjects || text == robotCmdProjectsList || text == "projects":
|
||||||
return "project:read", true
|
return "project:read", true
|
||||||
|
case text == robotCmdVulnAlerts || strings.HasPrefix(text, robotCmdVulnAlerts+" "),
|
||||||
|
text == "vuln alerts" || strings.HasPrefix(text, "vuln alerts "):
|
||||||
|
return "vulnerability:read", true
|
||||||
case text == robotCmdUnbindProject || text == "unbind project",
|
case text == robotCmdUnbindProject || text == "unbind project",
|
||||||
strings.HasPrefix(text, robotCmdNewProject+" "), strings.HasPrefix(text, "new project "),
|
strings.HasPrefix(text, robotCmdNewProject+" "), strings.HasPrefix(text, "new project "),
|
||||||
strings.HasPrefix(text, robotCmdBindProject+" "), strings.HasPrefix(text, "bind project "):
|
strings.HasPrefix(text, robotCmdBindProject+" "), strings.HasPrefix(text, "bind project "):
|
||||||
@@ -883,6 +1145,7 @@ func (h *RobotHandler) cmdBindUser(platform, userID, code string) string {
|
|||||||
h.mu.Lock()
|
h.mu.Lock()
|
||||||
delete(h.sessions, sk)
|
delete(h.sessions, sk)
|
||||||
delete(h.sessionRoles, sk)
|
delete(h.sessionRoles, sk)
|
||||||
|
delete(h.sessionModes, sk)
|
||||||
h.mu.Unlock()
|
h.mu.Unlock()
|
||||||
h.deleteSessionBinding(sk)
|
h.deleteSessionBinding(sk)
|
||||||
name := strings.TrimSpace(user.DisplayName)
|
name := strings.TrimSpace(user.DisplayName)
|
||||||
@@ -903,6 +1166,15 @@ func (h *RobotHandler) cmdUnbindUser(platform, userID string) string {
|
|||||||
if h.config.Robots.AuthorizationFor(platform).EffectiveMode() != config.RobotAuthModeUserBinding {
|
if h.config.Robots.AuthorizationFor(platform).EffectiveMode() != config.RobotAuthModeUserBinding {
|
||||||
return "该机器人使用受控服务账号模式,无需用户解绑。"
|
return "该机器人使用受控服务账号模式,无需用户解绑。"
|
||||||
}
|
}
|
||||||
|
_, accessErr := h.resolveRobotAccess(platform, userID)
|
||||||
|
if accessErr != nil {
|
||||||
|
return "当前平台账号尚未绑定。"
|
||||||
|
}
|
||||||
|
h.setPendingConfirmation(platform, userID, "unbind_user", "")
|
||||||
|
return "⚠️ 即将解除当前平台账号绑定。请在 2 分钟内发送「确认」继续,或发送「取消」。"
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) executeUnbindUser(platform, userID string) string {
|
||||||
access, accessErr := h.resolveRobotAccess(platform, userID)
|
access, accessErr := h.resolveRobotAccess(platform, userID)
|
||||||
if accessErr != nil {
|
if accessErr != nil {
|
||||||
return "当前平台账号尚未绑定。"
|
return "当前平台账号尚未绑定。"
|
||||||
@@ -914,6 +1186,7 @@ func (h *RobotHandler) cmdUnbindUser(platform, userID string) string {
|
|||||||
h.mu.Lock()
|
h.mu.Lock()
|
||||||
delete(h.sessions, sk)
|
delete(h.sessions, sk)
|
||||||
delete(h.sessionRoles, sk)
|
delete(h.sessionRoles, sk)
|
||||||
|
delete(h.sessionModes, sk)
|
||||||
h.mu.Unlock()
|
h.mu.Unlock()
|
||||||
h.deleteSessionBinding(sk)
|
h.deleteSessionBinding(sk)
|
||||||
if h.audit != nil {
|
if h.audit != nil {
|
||||||
@@ -926,6 +1199,50 @@ func (h *RobotHandler) cmdUnbindUser(platform, userID string) string {
|
|||||||
return "已解除当前平台账号与 CyberStrikeAI 用户的绑定。"
|
return "已解除当前平台账号与 CyberStrikeAI 用户的绑定。"
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) setPendingConfirmation(platform, userID, action, target string) {
|
||||||
|
sk := h.sessionKey(platform, userID)
|
||||||
|
now := time.Now()
|
||||||
|
h.mu.Lock()
|
||||||
|
for key, pending := range h.pendingConfirmations {
|
||||||
|
if now.After(pending.ExpiresAt) {
|
||||||
|
delete(h.pendingConfirmations, key)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
h.pendingConfirmations[sk] = robotPendingConfirmation{Action: action, Target: target, ExpiresAt: now.Add(2 * time.Minute)}
|
||||||
|
h.mu.Unlock()
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) cmdConfirm(platform, userID string) string {
|
||||||
|
sk := h.sessionKey(platform, userID)
|
||||||
|
h.mu.Lock()
|
||||||
|
pending, ok := h.pendingConfirmations[sk]
|
||||||
|
delete(h.pendingConfirmations, sk)
|
||||||
|
h.mu.Unlock()
|
||||||
|
if !ok || time.Now().After(pending.ExpiresAt) {
|
||||||
|
return "当前没有待确认操作,或确认已超时。"
|
||||||
|
}
|
||||||
|
switch pending.Action {
|
||||||
|
case "delete_conversation":
|
||||||
|
return h.executeDelete(platform, userID, pending.Target)
|
||||||
|
case "unbind_user":
|
||||||
|
return h.executeUnbindUser(platform, userID)
|
||||||
|
default:
|
||||||
|
return "待确认操作无效,已取消。"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) cmdCancelConfirmation(platform, userID string) string {
|
||||||
|
sk := h.sessionKey(platform, userID)
|
||||||
|
h.mu.Lock()
|
||||||
|
_, ok := h.pendingConfirmations[sk]
|
||||||
|
delete(h.pendingConfirmations, sk)
|
||||||
|
h.mu.Unlock()
|
||||||
|
if !ok {
|
||||||
|
return "当前没有待确认操作。"
|
||||||
|
}
|
||||||
|
return "已取消待确认操作。"
|
||||||
|
}
|
||||||
|
|
||||||
// handleRobotCommand 处理机器人内置命令;若匹配到命令返回 (回复内容, true),否则返回 ("", false)
|
// handleRobotCommand 处理机器人内置命令;若匹配到命令返回 (回复内容, true),否则返回 ("", false)
|
||||||
func (h *RobotHandler) handleRobotCommand(platform, userID, text string) (string, bool) {
|
func (h *RobotHandler) handleRobotCommand(platform, userID, text string) (string, bool) {
|
||||||
if (strings.HasPrefix(text, robotCmdBindUser+" ") || strings.HasPrefix(text, "bind ")) && !strings.HasPrefix(text, "bind project ") {
|
if (strings.HasPrefix(text, robotCmdBindUser+" ") || strings.HasPrefix(text, "bind ")) && !strings.HasPrefix(text, "bind project ") {
|
||||||
@@ -945,10 +1262,20 @@ func (h *RobotHandler) handleRobotCommand(platform, userID, text string) (string
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
switch {
|
switch {
|
||||||
|
case text == robotCmdVulnAlerts || text == "vuln alerts":
|
||||||
|
return h.cmdVulnerabilityAlerts(platform, userID, ""), true
|
||||||
|
case strings.HasPrefix(text, robotCmdVulnAlerts+" "):
|
||||||
|
return h.cmdVulnerabilityAlerts(platform, userID, strings.TrimSpace(text[len(robotCmdVulnAlerts)+1:])), true
|
||||||
|
case strings.HasPrefix(text, "vuln alerts "):
|
||||||
|
return h.cmdVulnerabilityAlerts(platform, userID, strings.TrimSpace(text[len("vuln alerts "):])), true
|
||||||
case text == robotCmdHelp || text == "help" || text == "?" || text == "?":
|
case text == robotCmdHelp || text == "help" || text == "?" || text == "?":
|
||||||
return h.cmdHelp(), true
|
return h.cmdHelp(platform, userID), true
|
||||||
case text == robotCmdIdentity || text == "whoami":
|
case text == robotCmdIdentity || text == "whoami":
|
||||||
return h.cmdIdentity(platform, userID), true
|
return h.cmdIdentity(platform, userID), true
|
||||||
|
case text == robotCmdConfirm || text == "confirm":
|
||||||
|
return h.cmdConfirm(platform, userID), true
|
||||||
|
case text == robotCmdCancel || text == "cancel":
|
||||||
|
return h.cmdCancelConfirmation(platform, userID), true
|
||||||
case text == robotCmdList || text == robotCmdListAlt || text == "list":
|
case text == robotCmdList || text == robotCmdListAlt || text == "list":
|
||||||
return h.cmdList(platform, userID), true
|
return h.cmdList(platform, userID), true
|
||||||
case strings.HasPrefix(text, robotCmdSwitch+" ") || strings.HasPrefix(text, robotCmdContinue+" ") || strings.HasPrefix(text, "switch ") || strings.HasPrefix(text, "continue "):
|
case strings.HasPrefix(text, robotCmdSwitch+" ") || strings.HasPrefix(text, robotCmdContinue+" ") || strings.HasPrefix(text, "switch ") || strings.HasPrefix(text, "continue "):
|
||||||
@@ -968,8 +1295,18 @@ func (h *RobotHandler) handleRobotCommand(platform, userID, text string) (string
|
|||||||
return h.cmdNew(platform, userID), true
|
return h.cmdNew(platform, userID), true
|
||||||
case text == robotCmdClear || text == "clear":
|
case text == robotCmdClear || text == "clear":
|
||||||
return h.cmdClear(platform, userID), true
|
return h.cmdClear(platform, userID), true
|
||||||
case text == robotCmdCurrent || text == "current":
|
case text == robotCmdStatus || text == "status":
|
||||||
return h.cmdCurrent(platform, userID), true
|
return h.cmdStatus(platform, userID), true
|
||||||
|
case text == robotCmdTask || text == "task":
|
||||||
|
return h.cmdTask(platform, userID), true
|
||||||
|
case strings.HasPrefix(text, robotCmdRename+" ") || strings.HasPrefix(text, "rename "):
|
||||||
|
var title string
|
||||||
|
if strings.HasPrefix(text, robotCmdRename+" ") {
|
||||||
|
title = strings.TrimSpace(text[len(robotCmdRename)+1:])
|
||||||
|
} else {
|
||||||
|
title = strings.TrimSpace(text[len("rename "):])
|
||||||
|
}
|
||||||
|
return h.cmdRename(platform, userID, title), true
|
||||||
case text == robotCmdStop || text == "stop":
|
case text == robotCmdStop || text == "stop":
|
||||||
return h.cmdStop(platform, userID), true
|
return h.cmdStop(platform, userID), true
|
||||||
case text == robotCmdRoles || text == robotCmdRolesList || text == "roles":
|
case text == robotCmdRoles || text == robotCmdRolesList || text == "roles":
|
||||||
@@ -985,6 +1322,23 @@ func (h *RobotHandler) handleRobotCommand(platform, userID, text string) (string
|
|||||||
roleName = strings.TrimSpace(text[5:])
|
roleName = strings.TrimSpace(text[5:])
|
||||||
}
|
}
|
||||||
return h.cmdSwitchRole(platform, userID, roleName), true
|
return h.cmdSwitchRole(platform, userID, roleName), true
|
||||||
|
case text == robotCmdModes || text == robotCmdModesList || text == "modes":
|
||||||
|
return h.cmdModes(platform, userID), true
|
||||||
|
case strings.HasPrefix(text, robotCmdModes+" ") || strings.HasPrefix(text, robotCmdSwitchMode+" ") || strings.HasPrefix(text, "mode "):
|
||||||
|
var mode string
|
||||||
|
switch {
|
||||||
|
case strings.HasPrefix(text, robotCmdModes+" "):
|
||||||
|
mode = strings.TrimSpace(text[len(robotCmdModes)+1:])
|
||||||
|
case strings.HasPrefix(text, robotCmdSwitchMode+" "):
|
||||||
|
mode = strings.TrimSpace(text[len(robotCmdSwitchMode)+1:])
|
||||||
|
default:
|
||||||
|
mode = strings.TrimSpace(text[5:])
|
||||||
|
}
|
||||||
|
return h.cmdSwitchMode(platform, userID, mode), true
|
||||||
|
case text == robotCmdPermissions || text == "permissions":
|
||||||
|
return h.cmdPermissions(platform, userID), true
|
||||||
|
case text == robotCmdDoctor || text == "doctor":
|
||||||
|
return h.cmdDoctor(), true
|
||||||
case strings.HasPrefix(text, robotCmdDelete+" ") || strings.HasPrefix(text, "delete "):
|
case strings.HasPrefix(text, robotCmdDelete+" ") || strings.HasPrefix(text, "delete "):
|
||||||
var convID string
|
var convID string
|
||||||
if strings.HasPrefix(text, robotCmdDelete+" ") {
|
if strings.HasPrefix(text, robotCmdDelete+" ") {
|
||||||
|
|||||||
@@ -0,0 +1,101 @@
|
|||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"cyberstrike-ai/internal/config"
|
||||||
|
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRobotModeSwitch(t *testing.T) {
|
||||||
|
h := NewRobotHandler(&config.Config{MultiAgent: config.MultiAgentConfig{Enabled: true}}, nil, nil, zap.NewNop())
|
||||||
|
|
||||||
|
if got := h.cmdSwitchMode("lark", "user-1", "plan-execute"); !strings.Contains(got, "Plan-Execute") {
|
||||||
|
t.Fatalf("unexpected switch response: %s", got)
|
||||||
|
}
|
||||||
|
if got := h.getAgentMode("lark", "user-1"); got != "plan_execute" {
|
||||||
|
t.Fatalf("mode = %q, want plan_execute", got)
|
||||||
|
}
|
||||||
|
if got := h.cmdModes("lark", "user-1"); !strings.Contains(got, "当前模式: Plan-Execute") {
|
||||||
|
t.Fatalf("unexpected modes response: %s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRobotModeRejectsUnavailableMultiAgent(t *testing.T) {
|
||||||
|
h := NewRobotHandler(&config.Config{}, nil, nil, zap.NewNop())
|
||||||
|
|
||||||
|
if got := h.cmdSwitchMode("lark", "user-1", "deep"); !strings.Contains(got, "启用 Eino 多代理") {
|
||||||
|
t.Fatalf("unexpected rejection: %s", got)
|
||||||
|
}
|
||||||
|
if got := h.getAgentMode("lark", "user-1"); got != "eino_single" {
|
||||||
|
t.Fatalf("mode changed after rejection: %q", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestParseRobotAgentModeRejectsUnknownMode(t *testing.T) {
|
||||||
|
if mode, ok := parseRobotAgentMode("unknown"); ok || mode != "" {
|
||||||
|
t.Fatalf("parseRobotAgentMode returned (%q, %v), want empty,false", mode, ok)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRobotStatusCommandPermission(t *testing.T) {
|
||||||
|
for _, command := range []string{"状态", "status"} {
|
||||||
|
permission, recognized := robotCommandPermission(command)
|
||||||
|
if !recognized || permission != "chat:read" {
|
||||||
|
t.Fatalf("command %q returned permission=%q recognized=%v", command, permission, recognized)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, removed := range []string{"当前", "current"} {
|
||||||
|
if _, recognized := robotCommandPermission(removed); recognized {
|
||||||
|
t.Fatalf("removed command %q is still recognized", removed)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRobotBestPracticeCommandPermissions(t *testing.T) {
|
||||||
|
cases := map[string]string{
|
||||||
|
"任务": "chat:read",
|
||||||
|
"task": "chat:read",
|
||||||
|
"重命名 新标题": "chat:write",
|
||||||
|
"rename x": "chat:write",
|
||||||
|
"诊断": "config:read",
|
||||||
|
"doctor": "config:read",
|
||||||
|
}
|
||||||
|
for command, want := range cases {
|
||||||
|
permission, recognized := robotCommandPermission(command)
|
||||||
|
if !recognized || permission != want {
|
||||||
|
t.Fatalf("command %q returned permission=%q recognized=%v, want %q,true", command, permission, recognized, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRobotConfirmationCanBeCancelled(t *testing.T) {
|
||||||
|
h := NewRobotHandler(&config.Config{}, nil, nil, zap.NewNop())
|
||||||
|
h.setPendingConfirmation("lark", "user-1", "delete_conversation", "conv-1")
|
||||||
|
if got := h.cmdCancelConfirmation("lark", "user-1"); got != "已取消待确认操作。" {
|
||||||
|
t.Fatalf("unexpected cancel response: %s", got)
|
||||||
|
}
|
||||||
|
if got := h.cmdConfirm("lark", "user-1"); !strings.Contains(got, "没有待确认操作") {
|
||||||
|
t.Fatalf("confirmation survived cancellation: %s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRobotDoctorSeparatesInternalToolsFromHTTPMCP(t *testing.T) {
|
||||||
|
h := NewRobotHandler(&config.Config{
|
||||||
|
Security: config.SecurityConfig{Tools: []config.ToolConfig{
|
||||||
|
{Name: "enabled-tool", Enabled: true},
|
||||||
|
{Name: "disabled-tool", Enabled: false},
|
||||||
|
}},
|
||||||
|
MCP: config.MCPConfig{Enabled: false},
|
||||||
|
}, nil, nil, zap.NewNop())
|
||||||
|
|
||||||
|
got := h.cmdDoctor()
|
||||||
|
if !strings.Contains(got, "内置 MCP 工具: 1/2 个已启用") {
|
||||||
|
t.Fatalf("internal tool status missing: %s", got)
|
||||||
|
}
|
||||||
|
if !strings.Contains(got, "HTTP MCP 服务: 已关闭") {
|
||||||
|
t.Fatalf("HTTP MCP status missing: %s", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -1,6 +1,7 @@
|
|||||||
package handler
|
package handler
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"sync"
|
"sync"
|
||||||
@@ -13,33 +14,51 @@ import (
|
|||||||
// some proxies that treat connections as idle; 10s is a reasonable balance with traffic.
|
// some proxies that treat connections as idle; 10s is a reasonable balance with traffic.
|
||||||
const sseKeepaliveInterval = 10 * time.Second
|
const sseKeepaliveInterval = 10 * time.Second
|
||||||
|
|
||||||
// sseKeepalive sends periodic SSE traffic so proxies (e.g. nginx proxy_read_timeout), NATs,
|
// runSSEKeepalive starts periodic SSE heartbeats in a background goroutine.
|
||||||
|
// The returned stop function must be deferred (or called) before the handler returns so the
|
||||||
|
// goroutine exits before Gin finalizes the ResponseWriter (avoids "Write called after Handler finished").
|
||||||
|
//
|
||||||
|
// writeMu must be the same mutex used by the handler's event writes for this request: concurrent
|
||||||
|
// writes to http.ResponseWriter break chunked transfer encoding (browser: net::ERR_INVALID_CHUNKED_ENCODING).
|
||||||
|
func runSSEKeepalive(c *gin.Context, writeMu *sync.Mutex) func() {
|
||||||
|
if writeMu == nil {
|
||||||
|
return func() {}
|
||||||
|
}
|
||||||
|
stop := make(chan struct{})
|
||||||
|
var wg sync.WaitGroup
|
||||||
|
wg.Add(1)
|
||||||
|
go func() {
|
||||||
|
defer wg.Done()
|
||||||
|
sseKeepaliveLoop(c, stop, writeMu)
|
||||||
|
}()
|
||||||
|
var once sync.Once
|
||||||
|
return func() {
|
||||||
|
once.Do(func() {
|
||||||
|
close(stop)
|
||||||
|
wg.Wait()
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// sseKeepaliveLoop sends periodic SSE traffic so proxies (e.g. nginx proxy_read_timeout), NATs,
|
||||||
// and load balancers do not close long-running streams. Some intermediaries ignore comment-only
|
// and load balancers do not close long-running streams. Some intermediaries ignore comment-only
|
||||||
// lines, so we send both a comment and a minimal data frame (type heartbeat) per tick.
|
// lines, so we send both a comment and a minimal data frame (type heartbeat) per tick.
|
||||||
//
|
func sseKeepaliveLoop(c *gin.Context, stop <-chan struct{}, writeMu *sync.Mutex) {
|
||||||
// writeMu must be the same mutex used by sendEvent for this request: concurrent writes to
|
|
||||||
// http.ResponseWriter break chunked transfer encoding (browser: net::ERR_INVALID_CHUNKED_ENCODING).
|
|
||||||
func sseKeepalive(c *gin.Context, stop <-chan struct{}, writeMu *sync.Mutex) {
|
|
||||||
if writeMu == nil {
|
|
||||||
return
|
|
||||||
}
|
|
||||||
ticker := time.NewTicker(sseKeepaliveInterval)
|
ticker := time.NewTicker(sseKeepaliveInterval)
|
||||||
defer ticker.Stop()
|
defer ticker.Stop()
|
||||||
|
ctx := c.Request.Context()
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
case <-stop:
|
case <-stop:
|
||||||
return
|
return
|
||||||
case <-c.Request.Context().Done():
|
case <-ctx.Done():
|
||||||
return
|
return
|
||||||
case <-ticker.C:
|
case <-ticker.C:
|
||||||
select {
|
|
||||||
case <-stop:
|
|
||||||
return
|
|
||||||
case <-c.Request.Context().Done():
|
|
||||||
return
|
|
||||||
default:
|
|
||||||
}
|
|
||||||
writeMu.Lock()
|
writeMu.Lock()
|
||||||
|
if sseShuttingDown(stop, ctx) {
|
||||||
|
writeMu.Unlock()
|
||||||
|
return
|
||||||
|
}
|
||||||
if _, err := fmt.Fprintf(c.Writer, ": keepalive\n\n"); err != nil {
|
if _, err := fmt.Fprintf(c.Writer, ": keepalive\n\n"); err != nil {
|
||||||
writeMu.Unlock()
|
writeMu.Unlock()
|
||||||
return
|
return
|
||||||
@@ -56,3 +75,14 @@ func sseKeepalive(c *gin.Context, stop <-chan struct{}, writeMu *sync.Mutex) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func sseShuttingDown(stop <-chan struct{}, ctx context.Context) bool {
|
||||||
|
select {
|
||||||
|
case <-stop:
|
||||||
|
return true
|
||||||
|
case <-ctx.Done():
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -0,0 +1,61 @@
|
|||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"net/http/httptest"
|
||||||
|
"sync"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRunSSEKeepaliveStopsBeforeHandlerReturns(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/events", nil)
|
||||||
|
|
||||||
|
var writeMu sync.Mutex
|
||||||
|
stop := runSSEKeepalive(c, &writeMu)
|
||||||
|
stop()
|
||||||
|
|
||||||
|
// A second stop must be safe (channel already closed, goroutine already exited).
|
||||||
|
stop()
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunSSEKeepaliveExitsOnClientDisconnect(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
c.Request = httptest.NewRequest("GET", "/events", nil).WithContext(ctx)
|
||||||
|
|
||||||
|
var writeMu sync.Mutex
|
||||||
|
stop := runSSEKeepalive(c, &writeMu)
|
||||||
|
cancel()
|
||||||
|
|
||||||
|
done := make(chan struct{})
|
||||||
|
go func() {
|
||||||
|
stop()
|
||||||
|
close(done)
|
||||||
|
}()
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
case <-time.After(2 * time.Second):
|
||||||
|
t.Fatal("keepalive stop did not complete after client disconnect")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestRunSSEKeepaliveNilMutexIsNoop(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest("GET", "/events", nil)
|
||||||
|
|
||||||
|
stop := runSSEKeepalive(c, nil)
|
||||||
|
stop()
|
||||||
|
}
|
||||||
@@ -298,6 +298,18 @@ func (m *AgentTaskManager) GetTask(conversationID string) *AgentTask {
|
|||||||
return m.tasks[conversationID]
|
return m.tasks[conversationID]
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// GetTaskSnapshot 返回运行任务的只读副本,供状态展示使用,避免锁外读取可变任务字段。
|
||||||
|
func (m *AgentTaskManager) GetTaskSnapshot(conversationID string) *AgentTask {
|
||||||
|
m.mu.RLock()
|
||||||
|
defer m.mu.RUnlock()
|
||||||
|
task := m.tasks[conversationID]
|
||||||
|
if task == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
snapshot := *task
|
||||||
|
return &snapshot
|
||||||
|
}
|
||||||
|
|
||||||
// runStuckCancellingCleanup 定期将长时间处于「取消中」的任务强制结束,避免卡住无法发新消息
|
// runStuckCancellingCleanup 定期将长时间处于「取消中」的任务强制结束,避免卡住无法发新消息
|
||||||
func (m *AgentTaskManager) runStuckCancellingCleanup() {
|
func (m *AgentTaskManager) runStuckCancellingCleanup() {
|
||||||
ticker := time.NewTicker(cleanupInterval)
|
ticker := time.NewTicker(cleanupInterval)
|
||||||
|
|||||||
@@ -101,6 +101,7 @@ func (h *VulnerabilityHandler) CreateVulnerability(c *gin.Context) {
|
|||||||
_ = h.db.SetResourceOwner("vulnerability", created.ID, session.UserID)
|
_ = h.db.SetResourceOwner("vulnerability", created.ID, session.UserID)
|
||||||
_ = h.db.AssignResourceToUser(session.UserID, "vulnerability", created.ID)
|
_ = h.db.AssignResourceToUser(session.UserID, "vulnerability", created.ID)
|
||||||
}
|
}
|
||||||
|
h.db.NotifyVulnerabilityCreated(created)
|
||||||
|
|
||||||
if h.audit != nil {
|
if h.audit != nil {
|
||||||
h.audit.RecordOK(c, "vulnerability", "create", "创建漏洞记录", "vulnerability", created.ID, map[string]interface{}{
|
h.audit.RecordOK(c, "vulnerability", "create", "创建漏洞记录", "vulnerability", created.ID, map[string]interface{}{
|
||||||
|
|||||||
@@ -0,0 +1,227 @@
|
|||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"cyberstrike-ai/internal/database"
|
||||||
|
"cyberstrike-ai/internal/robot"
|
||||||
|
"cyberstrike-ai/internal/security"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
type updateVulnerabilityAlertRequest struct {
|
||||||
|
Enabled bool `json:"enabled"`
|
||||||
|
MinSeverity string `json:"min_severity"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *VulnerabilityHandler) GetMyAlertSubscription(c *gin.Context) {
|
||||||
|
session, ok := security.CurrentSession(c)
|
||||||
|
if !ok {
|
||||||
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "unauthorized"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sub, err := h.db.GetVulnerabilityAlertSubscription(session.UserID)
|
||||||
|
if err != nil {
|
||||||
|
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
bindings, _ := h.db.ListRobotUserBindings(session.UserID)
|
||||||
|
deliveryReady := false
|
||||||
|
for _, binding := range bindings {
|
||||||
|
if binding.Enabled && robot.SupportsProactive(binding.Platform) {
|
||||||
|
deliveryReady = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
c.JSON(http.StatusOK, gin.H{"subscription": sub, "bindings": bindings, "delivery_ready": deliveryReady})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *VulnerabilityHandler) UpdateMyAlertSubscription(c *gin.Context) {
|
||||||
|
session, ok := security.CurrentSession(c)
|
||||||
|
if !ok {
|
||||||
|
c.JSON(http.StatusUnauthorized, gin.H{"error": "unauthorized"})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
var req updateVulnerabilityAlertRequest
|
||||||
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
sub, err := h.db.UpsertVulnerabilityAlertSubscription(session.UserID, req.Enabled, req.MinSeverity)
|
||||||
|
if err != nil {
|
||||||
|
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if h.audit != nil {
|
||||||
|
h.audit.RecordOK(c, "vulnerability", "alert_subscription_update", "更新漏洞提醒订阅", "user", session.UserID, map[string]interface{}{
|
||||||
|
"enabled": sub.Enabled, "min_severity": sub.MinSeverity,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
c.JSON(http.StatusOK, sub)
|
||||||
|
}
|
||||||
|
|
||||||
|
func vulnerabilityAlertSeverityLabel(value string) string {
|
||||||
|
switch value {
|
||||||
|
case "critical":
|
||||||
|
return "严重"
|
||||||
|
case "high":
|
||||||
|
return "高危"
|
||||||
|
case "medium":
|
||||||
|
return "中危"
|
||||||
|
case "low":
|
||||||
|
return "低危"
|
||||||
|
default:
|
||||||
|
return "信息"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func formatVulnerabilityRobotAlert(v *database.Vulnerability) string {
|
||||||
|
return fmt.Sprintf("🚨 新漏洞提醒\n\n标题:%s\n严重程度:%s(%s)\n目标:%s\n影响:%s\n修复建议:%s\n漏洞 ID:%s\n\n请核实目标版本与实际暴露情况。",
|
||||||
|
strings.TrimSpace(v.Title), vulnerabilityAlertSeverityLabel(v.Severity), v.Severity,
|
||||||
|
fallbackAlertText(v.Target), fallbackAlertText(v.Impact), fallbackAlertText(v.Recommendation), v.ID)
|
||||||
|
}
|
||||||
|
|
||||||
|
func fallbackAlertText(value string) string {
|
||||||
|
value = strings.TrimSpace(value)
|
||||||
|
if value == "" {
|
||||||
|
return "未填写"
|
||||||
|
}
|
||||||
|
if len([]rune(value)) > 300 {
|
||||||
|
return string([]rune(value)[:300]) + "…"
|
||||||
|
}
|
||||||
|
return value
|
||||||
|
}
|
||||||
|
|
||||||
|
// NotifyNewVulnerability pushes an alert to every eligible bound robot identity.
|
||||||
|
// A failure on one platform never blocks vulnerability creation or other recipients.
|
||||||
|
func (h *RobotHandler) NotifyNewVulnerability(v *database.Vulnerability) {
|
||||||
|
if h == nil || h.db == nil || v == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
recipients, err := h.db.ListVulnerabilityAlertRecipients(v)
|
||||||
|
if err != nil {
|
||||||
|
h.logger.Warn("查询漏洞提醒接收人失败", zap.String("vulnerability_id", v.ID), zap.Error(err))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
eligible := recipients[:0]
|
||||||
|
for _, recipient := range recipients {
|
||||||
|
if robot.SupportsProactive(recipient.Platform) {
|
||||||
|
eligible = append(eligible, recipient)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := h.db.EnqueueVulnerabilityAlertDeliveries(v.ID, eligible); err != nil {
|
||||||
|
h.logger.Warn("写入漏洞提醒投递队列失败", zap.String("vulnerability_id", v.ID), zap.Error(err))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case h.alertWake <- struct{}{}:
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// RunVulnerabilityAlertWorker drains the durable outbox. Failed sends use
|
||||||
|
// exponential backoff and remain retryable across application restarts.
|
||||||
|
func (h *RobotHandler) RunVulnerabilityAlertWorker(ctx context.Context) {
|
||||||
|
if h == nil || h.db == nil {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
ticker := time.NewTicker(30 * time.Second)
|
||||||
|
defer ticker.Stop()
|
||||||
|
for {
|
||||||
|
h.processVulnerabilityAlertDeliveries(ctx)
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
case <-ticker.C:
|
||||||
|
case <-h.alertWake:
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) processVulnerabilityAlertDeliveries(ctx context.Context) {
|
||||||
|
deliveries, err := h.db.ListDueVulnerabilityAlertDeliveries(50)
|
||||||
|
if err != nil {
|
||||||
|
h.logger.Warn("读取漏洞提醒投递队列失败", zap.Error(err))
|
||||||
|
return
|
||||||
|
}
|
||||||
|
for _, delivery := range deliveries {
|
||||||
|
sendCtx, cancel := context.WithTimeout(ctx, 15*time.Second)
|
||||||
|
err := robot.SendProactive(sendCtx, h.config.Robots, delivery.Platform, delivery.ExternalUserID, formatVulnerabilityRobotAlert(delivery.Vulnerability))
|
||||||
|
cancel()
|
||||||
|
if err != nil {
|
||||||
|
attempts := delivery.Attempts + 1
|
||||||
|
_ = h.db.MarkVulnerabilityAlertDeliveryFailed(delivery.ID, attempts, err)
|
||||||
|
h.logger.Warn("发送漏洞机器人提醒失败", zap.String("platform", delivery.Platform), zap.String("user_id", delivery.UserID), zap.String("vulnerability_id", delivery.Vulnerability.ID), zap.Int("attempt", attempts), zap.Error(err))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
_ = h.db.MarkVulnerabilityAlertDeliverySent(delivery.ID)
|
||||||
|
h.logger.Info("漏洞机器人提醒已发送", zap.String("platform", delivery.Platform), zap.String("user_id", delivery.UserID), zap.String("vulnerability_id", delivery.Vulnerability.ID))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *RobotHandler) cmdVulnerabilityAlerts(platform, externalUserID, arg string) string {
|
||||||
|
access, err := h.resolveRobotAccess(platform, externalUserID)
|
||||||
|
if err != nil {
|
||||||
|
return h.robotAccessDeniedMessage(platform)
|
||||||
|
}
|
||||||
|
arg = strings.ToLower(strings.TrimSpace(arg))
|
||||||
|
if arg == "" || arg == "status" || arg == "状态" {
|
||||||
|
sub, err := h.db.GetVulnerabilityAlertSubscription(access.User.ID)
|
||||||
|
if err != nil {
|
||||||
|
return "读取漏洞提醒设置失败,请稍后重试。"
|
||||||
|
}
|
||||||
|
state := "已关闭"
|
||||||
|
if sub.Enabled {
|
||||||
|
state = "已开启"
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("漏洞提醒%s;最低提醒级别:%s(%s)。\n可在 Web「漏洞管理」页面同步修改。", state, vulnerabilityAlertSeverityLabel(sub.MinSeverity), sub.MinSeverity)
|
||||||
|
}
|
||||||
|
enabled, severity := true, ""
|
||||||
|
switch arg {
|
||||||
|
case "开启", "开", "on", "enable":
|
||||||
|
severity = "high"
|
||||||
|
case "关闭", "关", "off", "disable":
|
||||||
|
enabled, severity = false, "high"
|
||||||
|
case "仅严重", "严重", "critical":
|
||||||
|
severity = "critical"
|
||||||
|
case "高危以上", "高危", "high":
|
||||||
|
severity = "high"
|
||||||
|
case "中危以上", "中危", "medium":
|
||||||
|
severity = "medium"
|
||||||
|
case "低危以上", "低危", "low":
|
||||||
|
severity = "low"
|
||||||
|
case "全部", "all", "info":
|
||||||
|
severity = "info"
|
||||||
|
default:
|
||||||
|
return "用法:漏洞提醒 开启|关闭|仅严重|高危以上|中危以上"
|
||||||
|
}
|
||||||
|
current, _ := h.db.GetVulnerabilityAlertSubscription(access.User.ID)
|
||||||
|
if (arg == "开启" || arg == "开" || arg == "on" || arg == "enable" || !enabled) && current != nil {
|
||||||
|
severity = current.MinSeverity
|
||||||
|
}
|
||||||
|
sub, err := h.db.UpsertVulnerabilityAlertSubscription(access.User.ID, enabled, severity)
|
||||||
|
if err != nil {
|
||||||
|
return "更新漏洞提醒失败,请稍后重试。"
|
||||||
|
}
|
||||||
|
if !sub.Enabled {
|
||||||
|
return "已关闭漏洞提醒。Web 端设置已同步。"
|
||||||
|
}
|
||||||
|
bindings, _ := h.db.ListRobotUserBindings(access.User.ID)
|
||||||
|
deliveryReady := false
|
||||||
|
for _, binding := range bindings {
|
||||||
|
if binding.Enabled && robot.SupportsProactive(binding.Platform) {
|
||||||
|
deliveryReady = true
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !deliveryReady {
|
||||||
|
return fmt.Sprintf("已保存漏洞提醒:%s及以上。当前没有支持主动推送的已绑定账号;请绑定企业微信、飞书、Telegram、Slack 或 Discord。", vulnerabilityAlertSeverityLabel(sub.MinSeverity))
|
||||||
|
}
|
||||||
|
return fmt.Sprintf("已开启漏洞提醒:%s及以上漏洞将通过已绑定机器人推送。Web 端设置已同步。", vulnerabilityAlertSeverityLabel(sub.MinSeverity))
|
||||||
|
}
|
||||||
@@ -0,0 +1,51 @@
|
|||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"path/filepath"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"cyberstrike-ai/internal/config"
|
||||||
|
"cyberstrike-ai/internal/database"
|
||||||
|
"cyberstrike-ai/internal/security"
|
||||||
|
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestRobotVulnerabilityAlertCommandSharesSubscription(t *testing.T) {
|
||||||
|
db, err := database.NewDB(filepath.Join(t.TempDir(), "alert-command.db"), zap.NewNop())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
if err := db.BootstrapRBAC("hash", security.PermissionCatalog); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
user, err := db.CreateRBACUser("alert-user", "Alert User", "hash", true, []string{database.RBACSystemRoleOperator})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := db.CreateRobotBindingCode(user.ID, "alert-code", time.Now().Add(time.Minute)); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := db.ConsumeRobotBindingCode("wecom", "external-user", "alert-code"); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
h := NewRobotHandler(&config.Config{}, db, nil, zap.NewNop())
|
||||||
|
|
||||||
|
if got := h.HandleMessage("wecom", "external-user", "漏洞提醒 高危以上"); !strings.Contains(got, "已开启") {
|
||||||
|
t.Fatalf("enable reply: %s", got)
|
||||||
|
}
|
||||||
|
sub, err := db.GetVulnerabilityAlertSubscription(user.ID)
|
||||||
|
if err != nil || !sub.Enabled || sub.MinSeverity != "high" {
|
||||||
|
t.Fatalf("web subscription not updated: %#v %v", sub, err)
|
||||||
|
}
|
||||||
|
if got := h.HandleMessage("wecom", "external-user", "vuln alerts off"); !strings.Contains(got, "已关闭") {
|
||||||
|
t.Fatalf("disable reply: %s", got)
|
||||||
|
}
|
||||||
|
sub, _ = db.GetVulnerabilityAlertSubscription(user.ID)
|
||||||
|
if sub.Enabled {
|
||||||
|
t.Fatalf("subscription remained enabled: %#v", sub)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,319 @@
|
|||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/hex"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"io"
|
||||||
|
"mime"
|
||||||
|
"net/http"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"cyberstrike-ai/internal/database"
|
||||||
|
"cyberstrike-ai/internal/security"
|
||||||
|
workflowrunner "cyberstrike-ai/internal/workflow"
|
||||||
|
workflowpkg "cyberstrike-ai/internal/workflow/package"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/google/uuid"
|
||||||
|
)
|
||||||
|
|
||||||
|
type workflowPackageResolution struct {
|
||||||
|
Action string `json:"action"`
|
||||||
|
NewWorkflowID string `json:"new_workflow_id"`
|
||||||
|
}
|
||||||
|
type workflowPackageImportRequest struct {
|
||||||
|
InspectionID string `json:"inspection_id"`
|
||||||
|
Resolution workflowPackageResolution `json:"resolution"`
|
||||||
|
ConfirmOverwrite bool `json:"confirm_overwrite"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *WorkflowHandler) ExportPackage(c *gin.Context) {
|
||||||
|
wf, err := h.db.GetWorkflowDefinition(c.Param("id"))
|
||||||
|
if err != nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusInternalServerError, "WFPKG_EXPORT_FAILED", "导出工作流包失败", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if wf == nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusNotFound, "WFPKG_WORKFLOW_NOT_FOUND", "工作流不存在", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
pkg, meta, err := workflowpkg.Export(workflowPackageDocument(wf))
|
||||||
|
if err != nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusInternalServerError, "WFPKG_EXPORT_FAILED", "导出工作流包失败", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.Header("Content-Type", "application/zip")
|
||||||
|
c.Header("Content-Disposition", mime.FormatMediaType("attachment", map[string]string{"filename": meta.FileName}))
|
||||||
|
c.Header("ETag", `"`+meta.PackageHash+`"`)
|
||||||
|
c.Header("X-Workflow-Package-SHA256", meta.PackageHash)
|
||||||
|
if h.audit != nil {
|
||||||
|
h.audit.RecordOK(c, "workflow_package", "export", "导出工作流包", "workflow", wf.ID, map[string]interface{}{"package_hash": meta.PackageHash})
|
||||||
|
}
|
||||||
|
c.Data(http.StatusOK, "application/zip", pkg)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *WorkflowHandler) CreatePackageInspection(c *gin.Context) {
|
||||||
|
session, ok := security.CurrentSession(c)
|
||||||
|
if !ok || strings.TrimSpace(session.UserID) == "" {
|
||||||
|
writeWorkflowPackageError(c, http.StatusUnauthorized, "WFPKG_INSPECTION_NOT_FOUND", "未授权访问", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
c.Request.Body = http.MaxBytesReader(c.Writer, c.Request.Body, workflowpkg.MaxArchiveBytes+1)
|
||||||
|
file, _, err := c.Request.FormFile("file")
|
||||||
|
if err != nil {
|
||||||
|
var maxErr *http.MaxBytesError
|
||||||
|
if errors.As(err, &maxErr) {
|
||||||
|
writeWorkflowPackageError(c, http.StatusUnprocessableEntity, "WFPKG_FILE_TOO_LARGE", "工作流包文件超过大小限制", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
writeWorkflowPackageError(c, http.StatusUnprocessableEntity, "WFPKG_FILE_REQUIRED", "必须上传工作流包文件", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
defer file.Close()
|
||||||
|
archive, err := io.ReadAll(io.LimitReader(file, workflowpkg.MaxArchiveBytes+1))
|
||||||
|
if err != nil || len(archive) > workflowpkg.MaxArchiveBytes {
|
||||||
|
writeWorkflowPackageError(c, http.StatusUnprocessableEntity, "WFPKG_FILE_TOO_LARGE", "工作流包文件超过大小限制", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
inspected, err := workflowpkg.InspectArchive(c.Request.Context(), archive, workflowrunner.ValidateGraphJSON)
|
||||||
|
if err != nil {
|
||||||
|
code := workflowpkg.ErrorCode(err)
|
||||||
|
if code == "" {
|
||||||
|
code = "WFPKG_INVALID_ARCHIVE"
|
||||||
|
}
|
||||||
|
if h.audit != nil {
|
||||||
|
h.audit.RecordFail(c, "workflow_package", "inspect", "工作流包预检失败", map[string]interface{}{"code": code, "package_hash": workflowPackageHash(archive)})
|
||||||
|
}
|
||||||
|
writeWorkflowPackageError(c, http.StatusUnprocessableEntity, code, "工作流包预检失败", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
state, local, err := h.workflowPackageConflict(inspected.Document.ID, inspected.ContentHash)
|
||||||
|
if err != nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusInternalServerError, "WFPKG_IMPORT_FAILED", "读取本地工作流失败", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
summary := workflowPackageInspectionSummary{ID: "wpi_" + strings.ReplaceAll(uuid.NewString(), "-", ""), Status: "ready", ExpiresAt: time.Now().UTC().Add(30 * time.Minute), Package: workflowPackagePackageSummary{PackageFormat: inspected.Manifest.PackageFormat, FormatVersion: inspected.Manifest.FormatVersion, PackageID: inspected.Manifest.PackageID, PackageHash: inspected.PackageHash}, Workflow: workflowPackageWorkflowSummary{SourceID: inspected.Document.ID, Name: inspected.Document.Name, Description: inspected.Document.Description, SourceRevision: inspected.Document.Version, Enabled: inspected.Document.Enabled, ContentHash: inspected.ContentHash, GraphHash: inspected.GraphHash, NodeCount: inspected.NodeCount, EdgeCount: inspected.EdgeCount}, Conflict: workflowPackageConflictSummary{State: state}, Warnings: []string{}}
|
||||||
|
if local != nil {
|
||||||
|
content, graph, _, _ := workflowpkg.DocumentHashes(workflowPackageDocument(local))
|
||||||
|
summary.Conflict.LocalWorkflow = &workflowPackageLocalWorkflow{ID: local.ID, Version: local.Version, ContentHash: content, GraphHash: graph}
|
||||||
|
}
|
||||||
|
manifestJSON, _ := json.Marshal(inspected.Manifest)
|
||||||
|
payloadJSON, err := json.Marshal(inspected.Document)
|
||||||
|
if err != nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusInternalServerError, "WFPKG_IMPORT_FAILED", "保存预检失败", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
inspectionJSON, _ := json.Marshal(summary)
|
||||||
|
record := &database.WorkflowPackageInspection{ID: summary.ID, PackageHash: inspected.PackageHash, ManifestJSON: string(manifestJSON), WorkflowPayloadJSON: string(payloadJSON), InspectionJSON: string(inspectionJSON), SourceWorkflowID: inspected.Document.ID, SourceRevision: inspected.Document.Version, SourceContentHash: inspected.ContentHash, SourceGraphHash: inspected.GraphHash, LocalConflictState: state, CreatedBy: session.UserID, CreatedAt: time.Now().UTC(), ExpiresAt: summary.ExpiresAt}
|
||||||
|
if local != nil {
|
||||||
|
content, graph, _, _ := workflowpkg.DocumentHashes(workflowPackageDocument(local))
|
||||||
|
record.LocalWorkflowID = local.ID
|
||||||
|
record.LocalContentHash = content
|
||||||
|
record.LocalGraphHash = graph
|
||||||
|
}
|
||||||
|
if err := h.db.CreateWorkflowPackageInspection(record); err != nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusInternalServerError, "WFPKG_IMPORT_FAILED", "保存预检失败", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if h.audit != nil {
|
||||||
|
h.audit.RecordOK(c, "workflow_package", "inspect", "工作流包预检成功", "inspection", record.ID, map[string]interface{}{"package_hash": record.PackageHash, "workflow_id": record.SourceWorkflowID})
|
||||||
|
}
|
||||||
|
c.JSON(http.StatusCreated, gin.H{"inspection": summary})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *WorkflowHandler) GetPackageInspection(c *gin.Context) {
|
||||||
|
session, ok := security.CurrentSession(c)
|
||||||
|
if !ok {
|
||||||
|
writeWorkflowPackageError(c, http.StatusUnauthorized, "WFPKG_INSPECTION_NOT_FOUND", "未授权访问", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
v, err := h.db.GetWorkflowPackageInspection(c.Param("inspectionId"), session.UserID)
|
||||||
|
if err != nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusInternalServerError, "WFPKG_IMPORT_FAILED", "读取预检失败", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if v == nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusNotFound, "WFPKG_INSPECTION_NOT_FOUND", "预检不存在", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if v.Status == "expired" {
|
||||||
|
writeWorkflowPackageError(c, http.StatusConflict, "WFPKG_INSPECTION_EXPIRED", "预检已过期", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
var summary any
|
||||||
|
_ = json.Unmarshal([]byte(v.InspectionJSON), &summary)
|
||||||
|
c.JSON(http.StatusOK, gin.H{"inspection": summary})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *WorkflowHandler) ApplyPackageImport(c *gin.Context) {
|
||||||
|
session, ok := security.CurrentSession(c)
|
||||||
|
if !ok {
|
||||||
|
writeWorkflowPackageError(c, http.StatusUnauthorized, "WFPKG_INSPECTION_NOT_FOUND", "未授权访问", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
key := strings.TrimSpace(c.GetHeader("Idempotency-Key"))
|
||||||
|
if _, err := uuid.Parse(key); err != nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusBadRequest, "WFPKG_IDEMPOTENCY_KEY_REQUIRED", "必须提供 UUID 幂等键", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
var req workflowPackageImportRequest
|
||||||
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusUnprocessableEntity, "WFPKG_INVALID_ACTION", "导入请求无效", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
req.InspectionID = strings.TrimSpace(req.InspectionID)
|
||||||
|
req.Resolution.Action = strings.TrimSpace(req.Resolution.Action)
|
||||||
|
req.Resolution.NewWorkflowID = strings.TrimSpace(req.Resolution.NewWorkflowID)
|
||||||
|
if req.Resolution.Action != "rename" && req.Resolution.NewWorkflowID != "" {
|
||||||
|
writeWorkflowPackageError(c, http.StatusUnprocessableEntity, "WFPKG_INVALID_ACTION", "当前导入动作不接受新工作流 ID", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
requestHash := workflowPackageRequestHash(req)
|
||||||
|
imp, replayed, err := h.db.ApplyWorkflowPackageImport(c.Request.Context(), database.WorkflowPackageApplyRequest{InspectionID: req.InspectionID, RequestHash: requestHash, IdempotencyKey: key, ActorUserID: session.UserID, Action: req.Resolution.Action, NewWorkflowID: req.Resolution.NewWorkflowID, ConfirmOverwrite: req.ConfirmOverwrite})
|
||||||
|
if err != nil {
|
||||||
|
h.writeWorkflowPackageImportError(c, req.InspectionID, err)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
wf, _ := h.db.GetWorkflowDefinition(imp.ResultingWorkflowID)
|
||||||
|
if !replayed && (imp.Result == "created" || imp.Result == "overwritten" || imp.Result == "renamed") {
|
||||||
|
workflowrunner.InvalidateCompiledCache(imp.ResultingWorkflowID)
|
||||||
|
}
|
||||||
|
response := h.workflowPackageImportResponse(imp, wf)
|
||||||
|
if !replayed && h.audit != nil {
|
||||||
|
h.audit.RecordOK(c, "workflow_package", "import", "工作流包导入成功", "workflow", imp.ResultingWorkflowID, map[string]interface{}{"inspection_id": imp.InspectionID, "action": imp.Action, "result": imp.Result})
|
||||||
|
}
|
||||||
|
status := http.StatusCreated
|
||||||
|
if replayed {
|
||||||
|
status = http.StatusOK
|
||||||
|
}
|
||||||
|
c.JSON(status, gin.H{"import": response})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *WorkflowHandler) GetPackageImport(c *gin.Context) {
|
||||||
|
session, ok := security.CurrentSession(c)
|
||||||
|
if !ok {
|
||||||
|
writeWorkflowPackageError(c, http.StatusUnauthorized, "WFPKG_INSPECTION_NOT_FOUND", "未授权访问", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
imp, err := h.db.GetWorkflowPackageImport(c.Param("importId"), session.UserID)
|
||||||
|
if err != nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusInternalServerError, "WFPKG_IMPORT_FAILED", "读取导入结果失败", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if imp == nil {
|
||||||
|
writeWorkflowPackageError(c, http.StatusNotFound, "WFPKG_INSPECTION_NOT_FOUND", "导入结果不存在", nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
wf, _ := h.db.GetWorkflowDefinition(imp.ResultingWorkflowID)
|
||||||
|
c.JSON(http.StatusOK, gin.H{"import": h.workflowPackageImportResponse(imp, wf)})
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *WorkflowHandler) workflowPackageConflict(id, sourceHash string) (string, *database.WorkflowDefinition, error) {
|
||||||
|
local, err := h.db.GetWorkflowDefinition(id)
|
||||||
|
if err != nil {
|
||||||
|
return "", nil, err
|
||||||
|
}
|
||||||
|
if local == nil {
|
||||||
|
return "none", nil, nil
|
||||||
|
}
|
||||||
|
localHash, _, _, err := workflowpkg.DocumentHashes(workflowPackageDocument(local))
|
||||||
|
if err != nil {
|
||||||
|
return "", nil, err
|
||||||
|
}
|
||||||
|
if localHash == sourceHash {
|
||||||
|
return "identical", local, nil
|
||||||
|
}
|
||||||
|
return "id_conflict", local, nil
|
||||||
|
}
|
||||||
|
func workflowPackageDocument(w *database.WorkflowDefinition) workflowpkg.Document {
|
||||||
|
return workflowpkg.Document{ID: w.ID, Name: w.Name, Description: w.Description, Version: w.Version, GraphJSON: w.GraphJSON, Enabled: w.Enabled, UpdatedAt: w.UpdatedAt}
|
||||||
|
}
|
||||||
|
func workflowPackageRequestHash(req workflowPackageImportRequest) string {
|
||||||
|
value := struct {
|
||||||
|
ConfirmOverwrite bool `json:"confirm_overwrite"`
|
||||||
|
InspectionID string `json:"inspection_id"`
|
||||||
|
Resolution workflowPackageResolution `json:"resolution"`
|
||||||
|
}{req.ConfirmOverwrite, req.InspectionID, req.Resolution}
|
||||||
|
b, _ := json.Marshal(value)
|
||||||
|
sum := sha256.Sum256(b)
|
||||||
|
return "sha256:" + hex.EncodeToString(sum[:])
|
||||||
|
}
|
||||||
|
func workflowPackageHash(b []byte) string {
|
||||||
|
sum := sha256.Sum256(b)
|
||||||
|
return "sha256:" + hex.EncodeToString(sum[:])
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *WorkflowHandler) writeWorkflowPackageImportError(c *gin.Context, inspectionID string, err error) {
|
||||||
|
var e *database.WorkflowPackageStoreError
|
||||||
|
if errors.As(err, &e) {
|
||||||
|
status := http.StatusConflict
|
||||||
|
if e.Code == "WFPKG_INVALID_ACTION" || e.Code == "WFPKG_INVALID_RENAME_ID" {
|
||||||
|
status = http.StatusUnprocessableEntity
|
||||||
|
}
|
||||||
|
if h.audit != nil {
|
||||||
|
h.audit.RecordFail(c, "workflow_package", "import", "工作流包导入失败", map[string]interface{}{"code": e.Code, "inspection_id": inspectionID})
|
||||||
|
}
|
||||||
|
writeWorkflowPackageError(c, status, e.Code, e.Message, nil)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if h.audit != nil {
|
||||||
|
h.audit.RecordFail(c, "workflow_package", "import", "工作流包导入失败", map[string]interface{}{"code": "WFPKG_IMPORT_FAILED", "inspection_id": inspectionID})
|
||||||
|
}
|
||||||
|
writeWorkflowPackageError(c, http.StatusInternalServerError, "WFPKG_IMPORT_FAILED", "导入工作流包失败", nil)
|
||||||
|
}
|
||||||
|
func writeWorkflowPackageError(c *gin.Context, status int, code, message string, details map[string]any) {
|
||||||
|
body := gin.H{"code": code, "message": message}
|
||||||
|
if len(details) > 0 {
|
||||||
|
body["details"] = details
|
||||||
|
}
|
||||||
|
c.JSON(status, gin.H{"error": body})
|
||||||
|
}
|
||||||
|
|
||||||
|
type workflowPackageInspectionSummary struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Status string `json:"status"`
|
||||||
|
ExpiresAt time.Time `json:"expires_at"`
|
||||||
|
Package workflowPackagePackageSummary `json:"package"`
|
||||||
|
Workflow workflowPackageWorkflowSummary `json:"workflow"`
|
||||||
|
Conflict workflowPackageConflictSummary `json:"conflict"`
|
||||||
|
Warnings []string `json:"warnings"`
|
||||||
|
}
|
||||||
|
type workflowPackagePackageSummary struct {
|
||||||
|
PackageFormat string `json:"package_format"`
|
||||||
|
FormatVersion string `json:"format_version"`
|
||||||
|
PackageID string `json:"package_id"`
|
||||||
|
PackageHash string `json:"package_hash"`
|
||||||
|
}
|
||||||
|
type workflowPackageWorkflowSummary struct {
|
||||||
|
SourceID string `json:"source_id"`
|
||||||
|
Name string `json:"name"`
|
||||||
|
Description string `json:"description"`
|
||||||
|
SourceRevision int `json:"source_revision"`
|
||||||
|
Enabled bool `json:"enabled"`
|
||||||
|
ContentHash string `json:"content_hash"`
|
||||||
|
GraphHash string `json:"graph_hash"`
|
||||||
|
NodeCount int `json:"node_count"`
|
||||||
|
EdgeCount int `json:"edge_count"`
|
||||||
|
}
|
||||||
|
type workflowPackageConflictSummary struct {
|
||||||
|
State string `json:"state"`
|
||||||
|
LocalWorkflow *workflowPackageLocalWorkflow `json:"local_workflow,omitempty"`
|
||||||
|
}
|
||||||
|
type workflowPackageLocalWorkflow struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
Version int `json:"version"`
|
||||||
|
ContentHash string `json:"content_hash"`
|
||||||
|
GraphHash string `json:"graph_hash"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func (h *WorkflowHandler) workflowPackageImportResponse(imp *database.WorkflowPackageImport, wf *database.WorkflowDefinition) gin.H {
|
||||||
|
out := gin.H{"id": imp.ID, "inspection_id": imp.InspectionID, "status": "succeeded", "result": imp.Result, "action": imp.Action, "source_workflow_id": imp.SourceWorkflowID, "target_workflow_id": imp.TargetWorkflowID, "applied_at": imp.AppliedAt}
|
||||||
|
if wf != nil {
|
||||||
|
content, graph, _, _ := workflowpkg.DocumentHashes(workflowPackageDocument(wf))
|
||||||
|
out["workflow"] = gin.H{"id": wf.ID, "version": wf.Version, "content_hash": content, "graph_hash": graph}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
@@ -0,0 +1,78 @@
|
|||||||
|
package handler
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"mime/multipart"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"cyberstrike-ai/internal/database"
|
||||||
|
"cyberstrike-ai/internal/security"
|
||||||
|
workflowpkg "cyberstrike-ai/internal/workflow/package"
|
||||||
|
|
||||||
|
"github.com/gin-gonic/gin"
|
||||||
|
"github.com/google/uuid"
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestWorkflowPackageHandlerInspectionAndCreateImport(t *testing.T) {
|
||||||
|
gin.SetMode(gin.TestMode)
|
||||||
|
db, err := database.NewDB(filepath.Join(t.TempDir(), "workflow-package-handler.db"), zap.NewNop())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
h := NewWorkflowHandler(db, zap.NewNop())
|
||||||
|
pkg, _, err := workflowpkg.Export(workflowpkg.Document{ID: "wf-api", Name: "API workflow", Version: 4, Enabled: true, UpdatedAt: time.Now().UTC(), GraphJSON: `{"nodes":[{"id":"start-1","type":"start","label":"开始","position":{"x":0,"y":0},"config":{}},{"id":"out-1","type":"output","label":"输出","position":{"x":0,"y":120},"config":{"output_key":"result","source_binding":{"from":"inputs","field":"message"}}}],"edges":[{"id":"e1","source":"start-1","target":"out-1"}],"config":{"schema_version":1}}`})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var body bytes.Buffer
|
||||||
|
writer := multipart.NewWriter(&body)
|
||||||
|
part, err := writer.CreateFormFile("file", "wf-api.csapkg.zip")
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := part.Write(pkg); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if err := writer.Close(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
w := httptest.NewRecorder()
|
||||||
|
c, _ := gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/api/workflow-package-inspections", &body)
|
||||||
|
c.Request.Header.Set("Content-Type", writer.FormDataContentType())
|
||||||
|
c.Set(security.ContextSessionKey, security.Session{UserID: "user-1"})
|
||||||
|
h.CreatePackageInspection(c)
|
||||||
|
if w.Code != http.StatusCreated {
|
||||||
|
t.Fatalf("inspection status=%d body=%s", w.Code, w.Body.String())
|
||||||
|
}
|
||||||
|
var inspected struct {
|
||||||
|
Inspection struct {
|
||||||
|
ID string `json:"id"`
|
||||||
|
} `json:"inspection"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(w.Body.Bytes(), &inspected); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
applyBody := bytes.NewBufferString(`{"inspection_id":"` + inspected.Inspection.ID + `","resolution":{"action":"create","new_workflow_id":""},"confirm_overwrite":false}`)
|
||||||
|
w = httptest.NewRecorder()
|
||||||
|
c, _ = gin.CreateTestContext(w)
|
||||||
|
c.Request = httptest.NewRequest(http.MethodPost, "/api/workflow-package-imports", applyBody)
|
||||||
|
c.Request.Header.Set("Content-Type", "application/json")
|
||||||
|
c.Request.Header.Set("Idempotency-Key", uuid.NewString())
|
||||||
|
c.Set(security.ContextSessionKey, security.Session{UserID: "user-1"})
|
||||||
|
h.ApplyPackageImport(c)
|
||||||
|
if w.Code != http.StatusCreated {
|
||||||
|
t.Fatalf("import status=%d body=%s", w.Code, w.Body.String())
|
||||||
|
}
|
||||||
|
saved, _ := db.GetWorkflowDefinition("wf-api")
|
||||||
|
if saved == nil || saved.Version != 1 {
|
||||||
|
t.Fatalf("saved=%#v", saved)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -76,7 +76,7 @@ func RegisterKnowledgeTool(
|
|||||||
}
|
}
|
||||||
|
|
||||||
mcpServer.RegisterTool(listRiskTypesTool, listRiskTypesHandler)
|
mcpServer.RegisterTool(listRiskTypesTool, listRiskTypesHandler)
|
||||||
logger.Info("风险类型列表工具已注册", zap.String("toolName", listRiskTypesTool.Name))
|
logger.Debug("风险类型列表工具已注册", zap.String("toolName", listRiskTypesTool.Name))
|
||||||
|
|
||||||
// 注册第二个工具:搜索知识库(保持原有功能)
|
// 注册第二个工具:搜索知识库(保持原有功能)
|
||||||
searchTool := mcp.Tool{
|
searchTool := mcp.Tool{
|
||||||
@@ -271,7 +271,7 @@ func RegisterKnowledgeTool(
|
|||||||
}
|
}
|
||||||
|
|
||||||
mcpServer.RegisterTool(searchTool, searchHandler)
|
mcpServer.RegisterTool(searchTool, searchHandler)
|
||||||
logger.Info("知识检索工具已注册", zap.String("toolName", searchTool.Name))
|
logger.Debug("知识检索工具已注册", zap.String("toolName", searchTool.Name))
|
||||||
}
|
}
|
||||||
|
|
||||||
// contains 检查切片是否包含元素
|
// contains 检查切片是否包含元素
|
||||||
|
|||||||
@@ -0,0 +1,328 @@
|
|||||||
|
package multiagent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"fmt"
|
||||||
|
"strings"
|
||||||
|
"unicode/utf8"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/adk"
|
||||||
|
"github.com/cloudwego/eino/adk/middlewares/summarization"
|
||||||
|
"github.com/cloudwego/eino/schema"
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
toolOutputTruncationMarker = "\n\n...[tool output truncated; full text persisted in reduction cache or summarization transcript]...\n\n"
|
||||||
|
aggressiveToolTruncDivisor = 4
|
||||||
|
)
|
||||||
|
|
||||||
|
// isEinoContextOverflowError reports API-side context window rejections.
|
||||||
|
func isEinoContextOverflowError(err error) bool {
|
||||||
|
if err == nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
msg := strings.ToLower(strings.TrimSpace(err.Error()))
|
||||||
|
if msg == "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
markers := []string{
|
||||||
|
"context length",
|
||||||
|
"context_length",
|
||||||
|
"maximum context",
|
||||||
|
"max context",
|
||||||
|
"context window",
|
||||||
|
"context overflow",
|
||||||
|
"too many tokens",
|
||||||
|
"token limit",
|
||||||
|
"tokens exceed",
|
||||||
|
"exceeds the context",
|
||||||
|
"input is too long",
|
||||||
|
"prompt is too long",
|
||||||
|
"request too large",
|
||||||
|
}
|
||||||
|
for _, m := range markers {
|
||||||
|
if strings.Contains(msg, m) {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
|
||||||
|
func truncateBytesWithMarker(content string, maxBytes int, marker string) string {
|
||||||
|
if maxBytes <= 0 || len(content) <= maxBytes {
|
||||||
|
return content
|
||||||
|
}
|
||||||
|
if marker == "" {
|
||||||
|
marker = toolOutputTruncationMarker
|
||||||
|
}
|
||||||
|
budget := maxBytes - len(marker)
|
||||||
|
if budget <= 0 {
|
||||||
|
if len(marker) > maxBytes {
|
||||||
|
return marker[:maxBytes]
|
||||||
|
}
|
||||||
|
return marker
|
||||||
|
}
|
||||||
|
head := budget / 2
|
||||||
|
tail := budget - head
|
||||||
|
for head > 0 && !utf8.RuneStart(content[head]) {
|
||||||
|
head--
|
||||||
|
}
|
||||||
|
tailStart := len(content) - tail
|
||||||
|
for tailStart < len(content) && !utf8.RuneStart(content[tailStart]) {
|
||||||
|
tailStart++
|
||||||
|
}
|
||||||
|
return content[:head] + marker + content[tailStart:]
|
||||||
|
}
|
||||||
|
|
||||||
|
func cloneMessage(msg adk.Message) adk.Message {
|
||||||
|
if msg == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
cloned := *msg
|
||||||
|
return &cloned
|
||||||
|
}
|
||||||
|
|
||||||
|
func truncateMessageToolContent(msg adk.Message, maxBytes int, spillRef string) adk.Message {
|
||||||
|
if msg == nil || maxBytes <= 0 {
|
||||||
|
return msg
|
||||||
|
}
|
||||||
|
out := cloneMessage(msg)
|
||||||
|
marker := toolOutputTruncationMarker
|
||||||
|
if spillRef != "" {
|
||||||
|
marker = fmt.Sprintf("\n\n...[tool output truncated; retrieve full text via: %s]...\n\n", spillRef)
|
||||||
|
}
|
||||||
|
switch out.Role {
|
||||||
|
case schema.Tool:
|
||||||
|
out.Content = truncateBytesWithMarker(out.Content, maxBytes, marker)
|
||||||
|
case schema.Assistant:
|
||||||
|
if out.ReasoningContent != "" {
|
||||||
|
out.ReasoningContent = truncateBytesWithMarker(out.ReasoningContent, maxBytes, marker)
|
||||||
|
}
|
||||||
|
if out.Content != "" {
|
||||||
|
out.Content = truncateBytesWithMarker(out.Content, maxBytes, marker)
|
||||||
|
}
|
||||||
|
case schema.User:
|
||||||
|
if out.Content != "" {
|
||||||
|
out.Content = truncateBytesWithMarker(out.Content, maxBytes, marker)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return out
|
||||||
|
}
|
||||||
|
|
||||||
|
func countMessagesTokens(
|
||||||
|
ctx context.Context,
|
||||||
|
msgs []adk.Message,
|
||||||
|
counter summarization.TokenCounterFunc,
|
||||||
|
tools []*schema.ToolInfo,
|
||||||
|
) (int, error) {
|
||||||
|
if counter == nil {
|
||||||
|
return 0, nil
|
||||||
|
}
|
||||||
|
n, err := counter(ctx, &summarization.TokenCounterInput{Messages: msgs, Tools: tools})
|
||||||
|
if err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
return n, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func truncateRoundMessagesToTokenBudget(
|
||||||
|
ctx context.Context,
|
||||||
|
round messageRound,
|
||||||
|
tokenBudget int,
|
||||||
|
counter summarization.TokenCounterFunc,
|
||||||
|
toolMaxBytes int,
|
||||||
|
spillRef string,
|
||||||
|
) ([]adk.Message, error) {
|
||||||
|
if tokenBudget <= 0 || len(round.messages) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
msgs := append([]adk.Message(nil), round.messages...)
|
||||||
|
if n, err := countMessagesTokens(ctx, msgs, counter, nil); err != nil {
|
||||||
|
return nil, err
|
||||||
|
} else if n <= tokenBudget {
|
||||||
|
return msgs, nil
|
||||||
|
}
|
||||||
|
if toolMaxBytes <= 0 {
|
||||||
|
toolMaxBytes = 12000
|
||||||
|
}
|
||||||
|
for pass := 0; pass < 8 && toolMaxBytes >= 32; pass++ {
|
||||||
|
out := make([]adk.Message, 0, len(msgs))
|
||||||
|
for _, msg := range msgs {
|
||||||
|
switch {
|
||||||
|
case msg != nil && msg.Role == schema.Tool:
|
||||||
|
out = append(out, truncateMessageToolContent(msg, toolMaxBytes, spillRef))
|
||||||
|
case msg != nil && msg.Role == schema.Assistant:
|
||||||
|
out = append(out, truncateMessageToolContent(msg, toolMaxBytes, spillRef))
|
||||||
|
default:
|
||||||
|
out = append(out, msg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
n, err := countMessagesTokens(ctx, out, counter, nil)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if n <= tokenBudget {
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
msgs = out
|
||||||
|
toolMaxBytes /= 2
|
||||||
|
}
|
||||||
|
return msgs, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
type compactMessagesOpts struct {
|
||||||
|
maxTokens int
|
||||||
|
counter summarization.TokenCounterFunc
|
||||||
|
toolMaxBytes int
|
||||||
|
spillRef string
|
||||||
|
aggressive bool
|
||||||
|
logger *zap.Logger
|
||||||
|
phase string
|
||||||
|
}
|
||||||
|
|
||||||
|
func compactMessagesByDroppingRounds(
|
||||||
|
ctx context.Context,
|
||||||
|
messages []adk.Message,
|
||||||
|
opts compactMessagesOpts,
|
||||||
|
) ([]adk.Message, bool) {
|
||||||
|
if opts.maxTokens <= 0 || len(messages) == 0 || opts.counter == nil {
|
||||||
|
return messages, false
|
||||||
|
}
|
||||||
|
before, err := countMessagesTokens(ctx, messages, opts.counter, nil)
|
||||||
|
if err != nil || before <= opts.maxTokens {
|
||||||
|
return messages, false
|
||||||
|
}
|
||||||
|
|
||||||
|
systems := make([]adk.Message, 0, 1)
|
||||||
|
contextMsgs := make([]adk.Message, 0, len(messages))
|
||||||
|
for _, msg := range messages {
|
||||||
|
if msg != nil && msg.Role == schema.System && len(contextMsgs) == 0 {
|
||||||
|
systems = append(systems, msg)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if msg != nil {
|
||||||
|
contextMsgs = append(contextMsgs, msg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
rounds := splitMessagesIntoRounds(contextMsgs)
|
||||||
|
if len(rounds) == 0 {
|
||||||
|
return messages, false
|
||||||
|
}
|
||||||
|
|
||||||
|
startIdx := 0
|
||||||
|
if opts.aggressive {
|
||||||
|
startIdx = len(rounds) - 1
|
||||||
|
if startIdx < 0 {
|
||||||
|
startIdx = 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
dropped := 0
|
||||||
|
for len(rounds) > 1 || (opts.aggressive && len(rounds) == 1) {
|
||||||
|
if !opts.aggressive && len(rounds) <= 1 {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if opts.aggressive && len(rounds) == 1 {
|
||||||
|
// Fall through to latest-round truncation below.
|
||||||
|
break
|
||||||
|
}
|
||||||
|
rounds = rounds[1:]
|
||||||
|
dropped++
|
||||||
|
candidate := append([]adk.Message(nil), systems...)
|
||||||
|
for _, round := range rounds {
|
||||||
|
candidate = append(candidate, round.messages...)
|
||||||
|
}
|
||||||
|
after, countErr := countMessagesTokens(ctx, candidate, opts.counter, nil)
|
||||||
|
if countErr != nil {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
if after <= opts.maxTokens {
|
||||||
|
if opts.logger != nil {
|
||||||
|
opts.logger.Warn("eino context compacted by dropping older rounds",
|
||||||
|
zap.String("phase", opts.phase),
|
||||||
|
zap.Int("tokens_before", before),
|
||||||
|
zap.Int("tokens_after", after),
|
||||||
|
zap.Int("max_tokens", opts.maxTokens),
|
||||||
|
zap.Int("dropped_rounds", dropped),
|
||||||
|
zap.Bool("aggressive", opts.aggressive),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return candidate, true
|
||||||
|
}
|
||||||
|
if opts.aggressive {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(rounds) == 0 {
|
||||||
|
return messages, false
|
||||||
|
}
|
||||||
|
latest := rounds[len(rounds)-1]
|
||||||
|
truncated, truncErr := truncateRoundMessagesToTokenBudget(
|
||||||
|
ctx, latest, opts.maxTokens, opts.counter, opts.toolMaxBytes, opts.spillRef,
|
||||||
|
)
|
||||||
|
if truncErr != nil || len(truncated) == 0 {
|
||||||
|
if opts.logger != nil {
|
||||||
|
opts.logger.Warn("eino context still above budget after round compaction; passing through without local error",
|
||||||
|
zap.String("phase", opts.phase),
|
||||||
|
zap.Int("tokens_before", before),
|
||||||
|
zap.Int("max_tokens", opts.maxTokens),
|
||||||
|
zap.Bool("aggressive", opts.aggressive),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return messages, false
|
||||||
|
}
|
||||||
|
candidate := append([]adk.Message(nil), systems...)
|
||||||
|
if dropped > 0 || startIdx > 0 {
|
||||||
|
for _, round := range rounds[:len(rounds)-1] {
|
||||||
|
candidate = append(candidate, round.messages...)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
candidate = append(candidate, truncated...)
|
||||||
|
after, countErr := countMessagesTokens(ctx, candidate, opts.counter, nil)
|
||||||
|
if countErr != nil {
|
||||||
|
return messages, false
|
||||||
|
}
|
||||||
|
if opts.logger != nil {
|
||||||
|
opts.logger.Warn("eino context compacted by truncating latest round tool output",
|
||||||
|
zap.String("phase", opts.phase),
|
||||||
|
zap.Int("tokens_before", before),
|
||||||
|
zap.Int("tokens_after", after),
|
||||||
|
zap.Int("max_tokens", opts.maxTokens),
|
||||||
|
zap.Int("dropped_rounds", dropped),
|
||||||
|
zap.Bool("aggressive", opts.aggressive),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return candidate, true
|
||||||
|
}
|
||||||
|
|
||||||
|
func aggressiveCompactMessagesForOverflow(
|
||||||
|
ctx context.Context,
|
||||||
|
messages []adk.Message,
|
||||||
|
maxTotalTokens int,
|
||||||
|
modelName string,
|
||||||
|
toolMaxBytes int,
|
||||||
|
phase string,
|
||||||
|
logger *zap.Logger,
|
||||||
|
) []adk.Message {
|
||||||
|
if len(messages) == 0 || maxTotalTokens <= 0 {
|
||||||
|
return messages
|
||||||
|
}
|
||||||
|
budget := maxTotalTokens * 70 / 100
|
||||||
|
if budget < 4096 {
|
||||||
|
budget = 4096
|
||||||
|
}
|
||||||
|
aggressiveToolMax := toolMaxBytes / aggressiveToolTruncDivisor
|
||||||
|
if aggressiveToolMax < 2048 {
|
||||||
|
aggressiveToolMax = 2048
|
||||||
|
}
|
||||||
|
out, _ := compactMessagesByDroppingRounds(ctx, messages, compactMessagesOpts{
|
||||||
|
maxTokens: budget,
|
||||||
|
counter: einoSummarizationTokenCounter(modelName),
|
||||||
|
toolMaxBytes: aggressiveToolMax,
|
||||||
|
aggressive: true,
|
||||||
|
logger: logger,
|
||||||
|
phase: phase,
|
||||||
|
})
|
||||||
|
return out
|
||||||
|
}
|
||||||
@@ -0,0 +1,104 @@
|
|||||||
|
package multiagent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/adk"
|
||||||
|
"github.com/cloudwego/eino/schema"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestIsEinoContextOverflowError(t *testing.T) {
|
||||||
|
t.Parallel()
|
||||||
|
cases := []struct {
|
||||||
|
err error
|
||||||
|
want bool
|
||||||
|
}{
|
||||||
|
{nil, false},
|
||||||
|
{errors.New("context length exceeded"), true},
|
||||||
|
{errors.New("maximum context length"), true},
|
||||||
|
{errors.New("input is too long for model"), true},
|
||||||
|
{errors.New("HTTP 429 Too Many Requests"), false},
|
||||||
|
{errors.New("invalid api key"), false},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
if got := isEinoContextOverflowError(tc.err); got != tc.want {
|
||||||
|
t.Fatalf("isEinoContextOverflowError(%v) = %v, want %v", tc.err, got, tc.want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestTruncateRoundMessagesToTokenBudget(t *testing.T) {
|
||||||
|
huge := strings.Repeat("x", 8000)
|
||||||
|
round := messageRound{messages: []adk.Message{
|
||||||
|
assistantToolCallsMsg("", "c1"),
|
||||||
|
schema.ToolMessage(huge, "c1"),
|
||||||
|
}}
|
||||||
|
out, err := truncateRoundMessagesToTokenBudget(
|
||||||
|
context.Background(), round, 256, einoSummarizationTokenCounter("gpt-4o"), 512, "",
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
for _, msg := range out {
|
||||||
|
if msg != nil && msg.Role == schema.Tool && len(msg.Content) >= len(huge) {
|
||||||
|
t.Fatalf("expected truncated tool output, got len=%d", len(msg.Content))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBuildBudgetedSummarizationModelInputTruncatesOversizedLatestRound(t *testing.T) {
|
||||||
|
huge := strings.Repeat("x", 8000)
|
||||||
|
msgs := []adk.Message{
|
||||||
|
assistantToolCallsMsg("", "call-latest"),
|
||||||
|
schema.ToolMessage(huge, "call-latest"),
|
||||||
|
}
|
||||||
|
counter := einoSummarizationTokenCounter("gpt-4o")
|
||||||
|
input, dropped, err := buildBudgetedSummarizationModelInput(
|
||||||
|
context.Background(),
|
||||||
|
schema.SystemMessage("sys"),
|
||||||
|
schema.UserMessage("instr"),
|
||||||
|
msgs,
|
||||||
|
counter,
|
||||||
|
512,
|
||||||
|
summarizationInputBudgetOpts{toolMaxBytes: 256},
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if dropped != 0 {
|
||||||
|
t.Fatalf("expected no dropped rounds, got %d", dropped)
|
||||||
|
}
|
||||||
|
toolContent := ""
|
||||||
|
for _, msg := range input {
|
||||||
|
if msg != nil && msg.Role == schema.Tool {
|
||||||
|
toolContent = msg.Content
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if len(toolContent) >= len(huge) {
|
||||||
|
t.Fatalf("expected oversized tool output to be compacted, got len=%d", len(toolContent))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestModelInputSoftBudgetNeverErrors(t *testing.T) {
|
||||||
|
mw := &modelInputSoftBudgetMiddleware{
|
||||||
|
maxTokens: 4,
|
||||||
|
toolMaxBytes: 16,
|
||||||
|
counter: fixedTokenCounter(4),
|
||||||
|
phase: "test",
|
||||||
|
}
|
||||||
|
state := &adk.ChatModelAgentState{Messages: []adk.Message{
|
||||||
|
schema.UserMessage("u"),
|
||||||
|
assistantToolCallsMsg("", "c1"),
|
||||||
|
schema.ToolMessage(strings.Repeat("t", 200), "c1"),
|
||||||
|
}}
|
||||||
|
_, out, err := mw.BeforeModelRewriteState(context.Background(), state, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("soft budget must not error: %v", err)
|
||||||
|
}
|
||||||
|
if out == nil || len(out.Messages) == 0 {
|
||||||
|
t.Fatal("expected compacted messages")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -105,6 +105,11 @@ type einoADKRunLoopArgs struct {
|
|||||||
|
|
||||||
// EinoCallbacks 可选:为 ADK Runner 注入 eino [callbacks] 全链路观测(见 internal/einoobserve)。
|
// EinoCallbacks 可选:为 ADK Runner 注入 eino [callbacks] 全链路观测(见 internal/einoobserve)。
|
||||||
EinoCallbacks *config.MultiAgentEinoCallbacksConfig
|
EinoCallbacks *config.MultiAgentEinoCallbacksConfig
|
||||||
|
|
||||||
|
// MaxTotalTokens / ToolMaxBytes / ModelName 用于 context overflow 时的激进压缩续跑。
|
||||||
|
MaxTotalTokens int
|
||||||
|
ToolMaxBytes int
|
||||||
|
ModelName string
|
||||||
}
|
}
|
||||||
|
|
||||||
func runEinoADKAgentLoop(ctx context.Context, args *einoADKRunLoopArgs, baseMsgs []adk.Message) (*RunResult, error) {
|
func runEinoADKAgentLoop(ctx context.Context, args *einoADKRunLoopArgs, baseMsgs []adk.Message) (*RunResult, error) {
|
||||||
@@ -439,6 +444,7 @@ func runEinoADKAgentLoop(ctx context.Context, args *einoADKRunLoopArgs, baseMsgs
|
|||||||
iter = startRunnerIter(msgs)
|
iter = startRunnerIter(msgs)
|
||||||
}
|
}
|
||||||
transientRetrier := newEinoTransientRunRetrier(einoTransientRunRetryPolicyFromArgs(args))
|
transientRetrier := newEinoTransientRunRetrier(einoTransientRunRetryPolicyFromArgs(args))
|
||||||
|
var contextOverflowRetried bool
|
||||||
handleRunErr := func(runErr error) error {
|
handleRunErr := func(runErr error) error {
|
||||||
if runErr == nil {
|
if runErr == nil {
|
||||||
return nil
|
return nil
|
||||||
@@ -495,6 +501,31 @@ func runEinoADKAgentLoop(ctx context.Context, args *einoADKRunLoopArgs, baseMsgs
|
|||||||
if runErr == nil {
|
if runErr == nil {
|
||||||
return false, nil
|
return false, nil
|
||||||
}
|
}
|
||||||
|
if isEinoContextOverflowError(runErr) && !contextOverflowRetried {
|
||||||
|
contextOverflowRetried = true
|
||||||
|
restartMsgs, ctxSource := einoMessagesForRunRestart(args, baseMsgs, runAccumulatedMsgs, baseAccumulatedCount)
|
||||||
|
restartMsgs = aggressiveCompactMessagesForOverflow(
|
||||||
|
ctx, restartMsgs, args.MaxTotalTokens, args.ModelName, args.ToolMaxBytes, orchMode, logger,
|
||||||
|
)
|
||||||
|
if logger != nil {
|
||||||
|
logger.Warn("eino context overflow, retrying with aggressive compaction",
|
||||||
|
zap.Error(runErr),
|
||||||
|
zap.String("orchestration", orchMode),
|
||||||
|
zap.String("contextSource", string(ctxSource)),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if progress != nil {
|
||||||
|
progress("eino_context_overflow_retry", "上下文超限,正在激进压缩后重试…", map[string]interface{}{
|
||||||
|
"conversationId": conversationID,
|
||||||
|
"source": "eino",
|
||||||
|
"orchestration": orchMode,
|
||||||
|
"contextSource": string(ctxSource),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
msgs = restartMsgs
|
||||||
|
iter = startRunnerIter(msgs)
|
||||||
|
return true, nil
|
||||||
|
}
|
||||||
if !isEinoTransientRunError(runErr) {
|
if !isEinoTransientRunError(runErr) {
|
||||||
return false, handleRunErr(runErr)
|
return false, handleRunErr(runErr)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
package multiagent
|
package multiagent
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"cyberstrike-ai/internal/config"
|
||||||
|
|
||||||
"github.com/cloudwego/eino/adk"
|
"github.com/cloudwego/eino/adk"
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
)
|
)
|
||||||
@@ -13,17 +15,19 @@ import (
|
|||||||
// 2. continuation user dedup — drop stale session-resume injections
|
// 2. continuation user dedup — drop stale session-resume injections
|
||||||
// 3. pre-summarization tool-call/result reconciliation
|
// 3. pre-summarization tool-call/result reconciliation
|
||||||
// 4. summarization
|
// 4. summarization
|
||||||
// 5. total model-input hard budget
|
// 5. soft model-input budget (warn/compact only, never fail locally)
|
||||||
// 6. final tool-call/result reconciliation
|
// 6. final tool-call/result reconciliation
|
||||||
// 7. orphan tool prune (defense in depth)
|
// 7. orphan tool prune (defense in depth)
|
||||||
// 8. telemetry
|
// 8. malformed tool_search history repair
|
||||||
// 9. model-facing trace snapshot
|
// 9. telemetry
|
||||||
|
// 10. model-facing trace snapshot
|
||||||
type einoChatModelTailConfig struct {
|
type einoChatModelTailConfig struct {
|
||||||
logger *zap.Logger
|
logger *zap.Logger
|
||||||
phase string
|
phase string
|
||||||
summarization adk.ChatModelAgentMiddleware
|
summarization adk.ChatModelAgentMiddleware
|
||||||
modelName string
|
modelName string
|
||||||
maxTotalTokens int
|
maxTotalTokens int
|
||||||
|
toolMaxBytes int
|
||||||
conversationID string
|
conversationID string
|
||||||
trace *modelFacingTraceHolder
|
trace *modelFacingTraceHolder
|
||||||
skipOrphanPruner bool
|
skipOrphanPruner bool
|
||||||
@@ -40,11 +44,12 @@ func appendEinoChatModelTailMiddlewares(handlers []adk.ChatModelAgentMiddleware,
|
|||||||
handlers = append(handlers, newToolPairReconcilerMiddleware(cfg.logger, cfg.phase+"_pre_summarization"))
|
handlers = append(handlers, newToolPairReconcilerMiddleware(cfg.logger, cfg.phase+"_pre_summarization"))
|
||||||
handlers = append(handlers, cfg.summarization)
|
handlers = append(handlers, cfg.summarization)
|
||||||
}
|
}
|
||||||
handlers = append(handlers, newModelInputBudgetMiddleware(cfg.maxTotalTokens, cfg.modelName, cfg.logger, cfg.phase))
|
handlers = append(handlers, newModelInputSoftBudgetMiddleware(cfg.maxTotalTokens, cfg.toolMaxBytes, cfg.modelName, cfg.logger, cfg.phase))
|
||||||
handlers = append(handlers, newToolPairReconcilerMiddleware(cfg.logger, cfg.phase))
|
handlers = append(handlers, newToolPairReconcilerMiddleware(cfg.logger, cfg.phase))
|
||||||
if !cfg.skipOrphanPruner {
|
if !cfg.skipOrphanPruner {
|
||||||
handlers = append(handlers, newOrphanToolPrunerMiddleware(cfg.logger, cfg.phase))
|
handlers = append(handlers, newOrphanToolPrunerMiddleware(cfg.logger, cfg.phase))
|
||||||
}
|
}
|
||||||
|
handlers = append(handlers, newToolSearchResultSanitizerMiddleware(cfg.logger, cfg.phase))
|
||||||
if !cfg.skipTelemetry {
|
if !cfg.skipTelemetry {
|
||||||
if teleMw := newEinoModelInputTelemetryMiddleware(cfg.logger, cfg.modelName, cfg.conversationID, cfg.phase); teleMw != nil {
|
if teleMw := newEinoModelInputTelemetryMiddleware(cfg.logger, cfg.modelName, cfg.conversationID, cfg.phase); teleMw != nil {
|
||||||
handlers = append(handlers, teleMw)
|
handlers = append(handlers, teleMw)
|
||||||
@@ -57,3 +62,10 @@ func appendEinoChatModelTailMiddlewares(handlers []adk.ChatModelAgentMiddleware,
|
|||||||
}
|
}
|
||||||
return handlers
|
return handlers
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func toolMaxBytesFromMW(mwCfg *config.MultiAgentEinoMiddlewareConfig) int {
|
||||||
|
if mwCfg != nil {
|
||||||
|
return mwCfg.ReductionMaxLengthForTruncEffective()
|
||||||
|
}
|
||||||
|
return config.MultiAgentEinoMiddlewareConfig{}.ReductionMaxLengthForTruncEffective()
|
||||||
|
}
|
||||||
|
|||||||
@@ -132,6 +132,7 @@ func buildPlanExecuteExecutorHandlers(ctx context.Context, a *PlanExecuteRootArg
|
|||||||
summarization: sumMw,
|
summarization: sumMw,
|
||||||
modelName: a.ModelName,
|
modelName: a.ModelName,
|
||||||
maxTotalTokens: a.AppCfg.OpenAI.MaxTotalTokens,
|
maxTotalTokens: a.AppCfg.OpenAI.MaxTotalTokens,
|
||||||
|
toolMaxBytes: toolMaxBytesFromMW(a.MwCfg),
|
||||||
conversationID: a.ConversationID,
|
conversationID: a.ConversationID,
|
||||||
trace: a.ModelFacingTrace,
|
trace: a.ModelFacingTrace,
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -151,6 +151,7 @@ func RunEinoSingleChatModelAgent(
|
|||||||
summarization: mainSumMw,
|
summarization: mainSumMw,
|
||||||
modelName: appCfg.OpenAI.Model,
|
modelName: appCfg.OpenAI.Model,
|
||||||
maxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
maxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
||||||
|
toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware),
|
||||||
conversationID: conversationID,
|
conversationID: conversationID,
|
||||||
trace: modelFacingTrace,
|
trace: modelFacingTrace,
|
||||||
})
|
})
|
||||||
@@ -235,6 +236,9 @@ func RunEinoSingleChatModelAgent(
|
|||||||
DA: chatAgent,
|
DA: chatAgent,
|
||||||
ModelFacingTrace: modelFacingTrace,
|
ModelFacingTrace: modelFacingTrace,
|
||||||
EinoCallbacks: &ma.EinoCallbacks,
|
EinoCallbacks: &ma.EinoCallbacks,
|
||||||
|
MaxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
||||||
|
ToolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware),
|
||||||
|
ModelName: appCfg.OpenAI.Model,
|
||||||
EmptyResponseMessage: "(Eino ADK single-agent session completed but no assistant text was captured. Check process details or logs.) " +
|
EmptyResponseMessage: "(Eino ADK single-agent session completed but no assistant text was captured. Check process details or logs.) " +
|
||||||
"(Eino ADK 单代理会话已完成,但未捕获到助手文本输出。请查看过程详情或日志。)",
|
"(Eino ADK 单代理会话已完成,但未捕获到助手文本输出。请查看过程详情或日志。)",
|
||||||
}, baseMsgs)
|
}, baseMsgs)
|
||||||
|
|||||||
@@ -98,17 +98,21 @@ func newEinoSummarizationMiddleware(
|
|||||||
}
|
}
|
||||||
triggerRatio := 0.8
|
triggerRatio := 0.8
|
||||||
emitInternalEvents := true
|
emitInternalEvents := true
|
||||||
|
outputReserve := config.DefaultSummarizationOutputReserveTokens
|
||||||
userLedgerMaxRunes := config.DefaultSummarizationUserIntentLedgerMaxRunes
|
userLedgerMaxRunes := config.DefaultSummarizationUserIntentLedgerMaxRunes
|
||||||
userLedgerEntryMaxRunes := config.DefaultSummarizationUserIntentLedgerEntryMaxRunes
|
userLedgerEntryMaxRunes := config.DefaultSummarizationUserIntentLedgerEntryMaxRunes
|
||||||
|
toolMaxBytes := config.MultiAgentEinoMiddlewareConfig{}.ReductionMaxLengthForTruncEffective()
|
||||||
if mwCfg != nil {
|
if mwCfg != nil {
|
||||||
triggerRatio = mwCfg.SummarizationTriggerRatioEffective()
|
triggerRatio = mwCfg.SummarizationTriggerRatioEffective()
|
||||||
emitInternalEvents = mwCfg.SummarizationEmitInternalEventsEffective()
|
emitInternalEvents = mwCfg.SummarizationEmitInternalEventsEffective()
|
||||||
|
outputReserve = mwCfg.SummarizationOutputReserveTokensEffective()
|
||||||
userLedgerMaxRunes = mwCfg.SummarizationUserIntentLedgerMaxRunesEffective()
|
userLedgerMaxRunes = mwCfg.SummarizationUserIntentLedgerMaxRunesEffective()
|
||||||
userLedgerEntryMaxRunes = mwCfg.SummarizationUserIntentLedgerEntryMaxRunesEffective()
|
userLedgerEntryMaxRunes = mwCfg.SummarizationUserIntentLedgerEntryMaxRunesEffective()
|
||||||
|
toolMaxBytes = mwCfg.ReductionMaxLengthForTruncEffective()
|
||||||
}
|
}
|
||||||
// The ledger is merged into the leading system message and cannot be removed as
|
// The ledger is merged into the leading system message and cannot be removed as
|
||||||
// an ordinary conversation round. Bound it relative to the configured window so
|
// an ordinary conversation round. Bound it relative to the configured window so
|
||||||
// it cannot crowd out the summary/latest turn or trip the final fail-closed guard.
|
// it cannot crowd out the summary/latest turn.
|
||||||
ledgerWindowCap := modelFacingRuneBudget(maxTotal, 0.20)
|
ledgerWindowCap := modelFacingRuneBudget(maxTotal, 0.20)
|
||||||
userLedgerMaxRunes = minPositiveInt(userLedgerMaxRunes, ledgerWindowCap)
|
userLedgerMaxRunes = minPositiveInt(userLedgerMaxRunes, ledgerWindowCap)
|
||||||
userLedgerEntryMaxRunes = minPositiveInt(userLedgerEntryMaxRunes, userLedgerMaxRunes)
|
userLedgerEntryMaxRunes = minPositiveInt(userLedgerEntryMaxRunes, userLedgerMaxRunes)
|
||||||
@@ -137,11 +141,11 @@ func newEinoSummarizationMiddleware(
|
|||||||
if recentTrailMax > trigger/2 {
|
if recentTrailMax > trigger/2 {
|
||||||
recentTrailMax = trigger / 2
|
recentTrailMax = trigger / 2
|
||||||
}
|
}
|
||||||
// The summarization request itself needs output headroom. A trigger is not a hard
|
// Summarization input aligns with the trigger threshold, minus explicit output reserve.
|
||||||
// request limit: one large turn can jump far beyond it. Bound the actual summary
|
summaryInputMax := trigger - outputReserve
|
||||||
// model input to 60% of the configured context window and keep complete recent
|
if summaryInputMax < 4096 {
|
||||||
// rounds so tool_call/tool_result pairs are never split.
|
summaryInputMax = trigger * 80 / 100
|
||||||
summaryInputMax := int(float64(maxTotal) * 0.6)
|
}
|
||||||
if summaryInputMax < 4096 {
|
if summaryInputMax < 4096 {
|
||||||
summaryInputMax = 4096
|
summaryInputMax = 4096
|
||||||
}
|
}
|
||||||
@@ -153,6 +157,9 @@ func newEinoSummarizationMiddleware(
|
|||||||
baseRoot = filepath.Join(filepath.Dir(dbPath), "conversation_artifacts", sanitizeEinoPathSegment(conv), "summarization")
|
baseRoot = filepath.Join(filepath.Dir(dbPath), "conversation_artifacts", sanitizeEinoPathSegment(conv), "summarization")
|
||||||
}
|
}
|
||||||
base := baseRoot
|
base := baseRoot
|
||||||
|
if abs, err := filepath.Abs(base); err == nil {
|
||||||
|
base = abs
|
||||||
|
}
|
||||||
if mkErr := os.MkdirAll(base, 0o755); mkErr == nil {
|
if mkErr := os.MkdirAll(base, 0o755); mkErr == nil {
|
||||||
transcriptPath = filepath.Join(base, "transcript.txt")
|
transcriptPath = filepath.Join(base, "transcript.txt")
|
||||||
}
|
}
|
||||||
@@ -160,6 +167,7 @@ func newEinoSummarizationMiddleware(
|
|||||||
|
|
||||||
retryPolicy := einoTransientRunRetryPolicyFromMW(mwCfg)
|
retryPolicy := einoTransientRunRetryPolicyFromMW(mwCfg)
|
||||||
retryMax := retryPolicy.maxAttempts
|
retryMax := retryPolicy.maxAttempts
|
||||||
|
var summaryOverflowRetries int
|
||||||
|
|
||||||
// ModelOptions apply only to summarization Generate (same ChatModel instance as the agent).
|
// ModelOptions apply only to summarization Generate (same ChatModel instance as the agent).
|
||||||
// Strip thinking/reasoning on this call path; mark requests for empty-choices diagnostics.
|
// Strip thinking/reasoning on this call path; mark requests for empty-choices diagnostics.
|
||||||
@@ -189,13 +197,29 @@ func newEinoSummarizationMiddleware(
|
|||||||
zap.String("path", transcriptPath), zap.Error(werr))
|
zap.String("path", transcriptPath), zap.Error(werr))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
budget := summaryInputMax
|
||||||
|
aggressive := summaryOverflowRetries > 0
|
||||||
|
if aggressive {
|
||||||
|
budget = summaryInputMax * 70 / 100
|
||||||
|
if budget < 4096 {
|
||||||
|
budget = 4096
|
||||||
|
}
|
||||||
|
}
|
||||||
input, dropped, berr := buildBudgetedSummarizationModelInput(
|
input, dropped, berr := buildBudgetedSummarizationModelInput(
|
||||||
ctx, sysInstruction, userInstruction, originalMsgs, tokenCounter, summaryInputMax,
|
ctx, sysInstruction, userInstruction, originalMsgs, tokenCounter, budget,
|
||||||
|
summarizationInputBudgetOpts{
|
||||||
|
toolMaxBytes: toolMaxBytes,
|
||||||
|
spillRef: transcriptPath,
|
||||||
|
aggressive: aggressive,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
if logger != nil && (berr != nil || dropped > 0) {
|
if logger != nil && (berr != nil || dropped > 0 || aggressive) {
|
||||||
fields := []zap.Field{
|
fields := []zap.Field{
|
||||||
zap.Int("max_input_tokens", summaryInputMax),
|
zap.Int("max_input_tokens", budget),
|
||||||
|
zap.Int("trigger_context_tokens", trigger),
|
||||||
|
zap.Int("output_reserve_tokens", outputReserve),
|
||||||
zap.Int("dropped_rounds", dropped),
|
zap.Int("dropped_rounds", dropped),
|
||||||
|
zap.Bool("aggressive", aggressive),
|
||||||
}
|
}
|
||||||
if berr != nil {
|
if berr != nil {
|
||||||
fields = append(fields, zap.Error(berr))
|
fields = append(fields, zap.Error(berr))
|
||||||
@@ -220,6 +244,15 @@ func newEinoSummarizationMiddleware(
|
|||||||
Retry: &summarization.RetryConfig{
|
Retry: &summarization.RetryConfig{
|
||||||
MaxRetries: &retryMax,
|
MaxRetries: &retryMax,
|
||||||
ShouldRetry: func(_ context.Context, _ adk.Message, err error) bool {
|
ShouldRetry: func(_ context.Context, _ adk.Message, err error) bool {
|
||||||
|
if isEinoContextOverflowError(err) && summaryOverflowRetries < 1 {
|
||||||
|
summaryOverflowRetries++
|
||||||
|
if logger != nil {
|
||||||
|
logger.Warn("eino summarization context overflow, retrying with aggressive compaction",
|
||||||
|
zap.Error(err),
|
||||||
|
)
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
retry := isEinoTransientRunError(err)
|
retry := isEinoTransientRunError(err)
|
||||||
if retry && logger != nil {
|
if retry && logger != nil {
|
||||||
logger.Warn("eino summarization generate transient error, will retry if attempts remain",
|
logger.Warn("eino summarization generate transient error, will retry if attempts remain",
|
||||||
@@ -275,6 +308,13 @@ func newEinoSummarizationMiddleware(
|
|||||||
return mw, nil
|
return mw, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// summarizationInputBudgetOpts controls spill/truncation behavior when a round alone exceeds budget.
|
||||||
|
type summarizationInputBudgetOpts struct {
|
||||||
|
toolMaxBytes int
|
||||||
|
spillRef string
|
||||||
|
aggressive bool
|
||||||
|
}
|
||||||
|
|
||||||
// buildBudgetedSummarizationModelInput builds the exact payload sent to the summary model.
|
// buildBudgetedSummarizationModelInput builds the exact payload sent to the summary model.
|
||||||
// It retains the newest complete conversation rounds within budget and emits an explicit
|
// It retains the newest complete conversation rounds within budget and emits an explicit
|
||||||
// marker when older rounds are omitted. The full pre-compaction transcript is persisted
|
// marker when older rounds are omitted. The full pre-compaction transcript is persisted
|
||||||
@@ -286,6 +326,7 @@ func buildBudgetedSummarizationModelInput(
|
|||||||
originalMsgs []adk.Message,
|
originalMsgs []adk.Message,
|
||||||
tokenCounter summarization.TokenCounterFunc,
|
tokenCounter summarization.TokenCounterFunc,
|
||||||
maxTokens int,
|
maxTokens int,
|
||||||
|
opts summarizationInputBudgetOpts,
|
||||||
) ([]adk.Message, int, error) {
|
) ([]adk.Message, int, error) {
|
||||||
base := []adk.Message{sysInstruction, userInstruction}
|
base := []adk.Message{sysInstruction, userInstruction}
|
||||||
baseTokens, err := tokenCounter(ctx, &summarization.TokenCounterInput{Messages: base})
|
baseTokens, err := tokenCounter(ctx, &summarization.TokenCounterInput{Messages: base})
|
||||||
@@ -312,12 +353,36 @@ func buildBudgetedSummarizationModelInput(
|
|||||||
rounds := splitMessagesIntoRounds(contextMsgs)
|
rounds := splitMessagesIntoRounds(contextMsgs)
|
||||||
selectedReverse := make([]messageRound, 0, len(rounds))
|
selectedReverse := make([]messageRound, 0, len(rounds))
|
||||||
used := 0
|
used := 0
|
||||||
|
toolMaxBytes := opts.toolMaxBytes
|
||||||
|
if toolMaxBytes <= 0 {
|
||||||
|
toolMaxBytes = 12000
|
||||||
|
}
|
||||||
|
if opts.aggressive {
|
||||||
|
toolMaxBytes /= aggressiveToolTruncDivisor
|
||||||
|
if toolMaxBytes < 2048 {
|
||||||
|
toolMaxBytes = 2048
|
||||||
|
}
|
||||||
|
}
|
||||||
for i := len(rounds) - 1; i >= 0; i-- {
|
for i := len(rounds) - 1; i >= 0; i-- {
|
||||||
n, countErr := tokenCounter(ctx, &summarization.TokenCounterInput{Messages: rounds[i].messages})
|
n, countErr := tokenCounter(ctx, &summarization.TokenCounterInput{Messages: rounds[i].messages})
|
||||||
if countErr != nil {
|
if countErr != nil {
|
||||||
return nil, 0, countErr
|
return nil, 0, countErr
|
||||||
}
|
}
|
||||||
if used+n > remaining {
|
if used+n > remaining {
|
||||||
|
if len(selectedReverse) == 0 {
|
||||||
|
slot := remaining - used
|
||||||
|
if slot > 0 {
|
||||||
|
truncated, truncErr := truncateRoundMessagesToTokenBudget(
|
||||||
|
ctx, rounds[i], slot, tokenCounter, toolMaxBytes, opts.spillRef,
|
||||||
|
)
|
||||||
|
if truncErr != nil {
|
||||||
|
return nil, 0, truncErr
|
||||||
|
}
|
||||||
|
if len(truncated) > 0 {
|
||||||
|
selectedReverse = append(selectedReverse, messageRound{messages: truncated})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
used += n
|
used += n
|
||||||
|
|||||||
@@ -67,7 +67,7 @@ func TestBuildBudgetedSummarizationModelInputKeepsRecentCompleteRounds(t *testin
|
|||||||
}
|
}
|
||||||
input, dropped, err := buildBudgetedSummarizationModelInput(
|
input, dropped, err := buildBudgetedSummarizationModelInput(
|
||||||
context.Background(), schema.SystemMessage("summary-system"), schema.UserMessage("summary-instruction"),
|
context.Background(), schema.SystemMessage("summary-system"), schema.UserMessage("summary-instruction"),
|
||||||
msgs, fixedTokenCounter(2), 7,
|
msgs, fixedTokenCounter(2), 7, summarizationInputBudgetOpts{},
|
||||||
)
|
)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
|
|||||||
@@ -35,9 +35,6 @@ func isEinoTransientRunError(err error) bool {
|
|||||||
if msg == "" {
|
if msg == "" {
|
||||||
return false
|
return false
|
||||||
}
|
}
|
||||||
if strings.Contains(msg, "model input exceeds configured hard budget") {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
transientMarkers := []string{
|
transientMarkers := []string{
|
||||||
"406",
|
"406",
|
||||||
"429",
|
"429",
|
||||||
|
|||||||
@@ -33,7 +33,6 @@ func TestIsEinoTransientRunError(t *testing.T) {
|
|||||||
{"iteration limit", errors.New("max iteration reached"), false},
|
{"iteration limit", errors.New("max iteration reached"), false},
|
||||||
{"canceled", context.Canceled, false},
|
{"canceled", context.Canceled, false},
|
||||||
{"deadline", context.DeadlineExceeded, false},
|
{"deadline", context.DeadlineExceeded, false},
|
||||||
{"model hard budget containing 500", errors.New("model input exceeds configured hard budget after preserving the latest round: tokens=20500 max=19500 phase=test"), false},
|
|
||||||
{"auth", errors.New("invalid api key"), false},
|
{"auth", errors.New("invalid api key"), false},
|
||||||
}
|
}
|
||||||
for _, tc := range cases {
|
for _, tc := range cases {
|
||||||
|
|||||||
@@ -1,106 +0,0 @@
|
|||||||
package multiagent
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"fmt"
|
|
||||||
|
|
||||||
"github.com/cloudwego/eino/adk"
|
|
||||||
"github.com/cloudwego/eino/adk/middlewares/summarization"
|
|
||||||
"github.com/cloudwego/eino/schema"
|
|
||||||
"go.uber.org/zap"
|
|
||||||
)
|
|
||||||
|
|
||||||
// modelInputBudgetMiddleware is the final deterministic guard before a normal model call.
|
|
||||||
// Summarization remains the primary compaction strategy; this middleware only removes the
|
|
||||||
// oldest complete rounds if the finalized state still exceeds 65% of the configured
|
|
||||||
// window, leaving serialization/tokenizer headroom for the outbound HTTP guard.
|
|
||||||
type modelInputBudgetMiddleware struct {
|
|
||||||
adk.BaseChatModelAgentMiddleware
|
|
||||||
maxTokens int
|
|
||||||
counter summarization.TokenCounterFunc
|
|
||||||
logger *zap.Logger
|
|
||||||
phase string
|
|
||||||
}
|
|
||||||
|
|
||||||
func newModelInputBudgetMiddleware(maxTotalTokens int, modelName string, logger *zap.Logger, phase string) adk.ChatModelAgentMiddleware {
|
|
||||||
if maxTotalTokens <= 0 {
|
|
||||||
maxTotalTokens = 120000
|
|
||||||
}
|
|
||||||
limit := int(float64(maxTotalTokens) * 0.65)
|
|
||||||
if limit < 4096 {
|
|
||||||
limit = 4096
|
|
||||||
}
|
|
||||||
return &modelInputBudgetMiddleware{
|
|
||||||
maxTokens: limit,
|
|
||||||
counter: einoSummarizationTokenCounter(modelName),
|
|
||||||
logger: logger,
|
|
||||||
phase: phase,
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func (m *modelInputBudgetMiddleware) BeforeModelRewriteState(
|
|
||||||
ctx context.Context,
|
|
||||||
state *adk.ChatModelAgentState,
|
|
||||||
mc *adk.ModelContext,
|
|
||||||
) (context.Context, *adk.ChatModelAgentState, error) {
|
|
||||||
if m == nil || state == nil || len(state.Messages) == 0 {
|
|
||||||
return ctx, state, nil
|
|
||||||
}
|
|
||||||
count := func(msgs []adk.Message) (int, error) {
|
|
||||||
input := &summarization.TokenCounterInput{Messages: msgs}
|
|
||||||
if mc != nil {
|
|
||||||
input.Tools = mc.Tools
|
|
||||||
}
|
|
||||||
return m.counter(ctx, input)
|
|
||||||
}
|
|
||||||
before, err := count(state.Messages)
|
|
||||||
if err != nil {
|
|
||||||
return ctx, state, err
|
|
||||||
}
|
|
||||||
if before <= m.maxTokens {
|
|
||||||
return ctx, state, nil
|
|
||||||
}
|
|
||||||
|
|
||||||
systems := make([]adk.Message, 0, 1)
|
|
||||||
contextMsgs := make([]adk.Message, 0, len(state.Messages))
|
|
||||||
for _, msg := range state.Messages {
|
|
||||||
if msg != nil && msg.Role == schema.System && len(contextMsgs) == 0 {
|
|
||||||
systems = append(systems, msg)
|
|
||||||
continue
|
|
||||||
}
|
|
||||||
if msg != nil {
|
|
||||||
contextMsgs = append(contextMsgs, msg)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
rounds := splitMessagesIntoRounds(contextMsgs)
|
|
||||||
dropped := 0
|
|
||||||
var candidate []adk.Message
|
|
||||||
for len(rounds) > 1 {
|
|
||||||
rounds = rounds[1:]
|
|
||||||
dropped++
|
|
||||||
candidate = append(candidate[:0], systems...)
|
|
||||||
for _, round := range rounds {
|
|
||||||
candidate = append(candidate, round.messages...)
|
|
||||||
}
|
|
||||||
after, countErr := count(candidate)
|
|
||||||
if countErr != nil {
|
|
||||||
return ctx, state, countErr
|
|
||||||
}
|
|
||||||
if after <= m.maxTokens {
|
|
||||||
out := *state
|
|
||||||
out.Messages = append([]adk.Message(nil), candidate...)
|
|
||||||
if m.logger != nil {
|
|
||||||
m.logger.Warn("eino model input hard budget applied",
|
|
||||||
zap.String("phase", m.phase), zap.Int("tokens_before", before),
|
|
||||||
zap.Int("tokens_after", after), zap.Int("max_tokens", m.maxTokens),
|
|
||||||
zap.Int("dropped_rounds", dropped))
|
|
||||||
}
|
|
||||||
return ctx, &out, nil
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
return ctx, state, fmt.Errorf(
|
|
||||||
"model input exceeds configured hard budget after preserving the latest round: tokens=%d max=%d phase=%s",
|
|
||||||
before, m.maxTokens, m.phase,
|
|
||||||
)
|
|
||||||
}
|
|
||||||
@@ -1,50 +0,0 @@
|
|||||||
package multiagent
|
|
||||||
|
|
||||||
import (
|
|
||||||
"context"
|
|
||||||
"strings"
|
|
||||||
"testing"
|
|
||||||
|
|
||||||
"github.com/cloudwego/eino/adk"
|
|
||||||
"github.com/cloudwego/eino/schema"
|
|
||||||
)
|
|
||||||
|
|
||||||
func TestModelInputBudgetDropsOldestCompleteRoundsAndPersistsLatest(t *testing.T) {
|
|
||||||
mw := &modelInputBudgetMiddleware{
|
|
||||||
maxTokens: 7,
|
|
||||||
counter: fixedTokenCounter(4),
|
|
||||||
phase: "test",
|
|
||||||
}
|
|
||||||
state := &adk.ChatModelAgentState{Messages: []adk.Message{
|
|
||||||
schema.SystemMessage("system"),
|
|
||||||
schema.UserMessage("old-user"),
|
|
||||||
schema.AssistantMessage("old-answer", nil),
|
|
||||||
schema.UserMessage("latest-user"),
|
|
||||||
assistantToolCallsMsg("", "latest-call"),
|
|
||||||
schema.ToolMessage("latest-result", "latest-call"),
|
|
||||||
}}
|
|
||||||
_, out, err := mw.BeforeModelRewriteState(context.Background(), state, nil)
|
|
||||||
if err != nil {
|
|
||||||
t.Fatal(err)
|
|
||||||
}
|
|
||||||
joined := formatSummarizationTranscript(out.Messages)
|
|
||||||
if strings.Contains(joined, "old-user") || strings.Contains(joined, "old-answer") {
|
|
||||||
t.Fatalf("old rounds retained after hard budget: %s", joined)
|
|
||||||
}
|
|
||||||
if !strings.Contains(joined, "latest-user") || !strings.Contains(joined, "latest-result") {
|
|
||||||
t.Fatalf("latest rounds lost after hard budget: %s", joined)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
func TestModelInputBudgetFailsLocallyWhenLatestRoundAloneCannotFit(t *testing.T) {
|
|
||||||
mw := &modelInputBudgetMiddleware{maxTokens: 2, counter: fixedTokenCounter(4), phase: "test"}
|
|
||||||
state := &adk.ChatModelAgentState{Messages: []adk.Message{
|
|
||||||
schema.SystemMessage("system"),
|
|
||||||
assistantToolCallsMsg("", "latest-call"),
|
|
||||||
schema.ToolMessage("latest-result", "latest-call"),
|
|
||||||
}}
|
|
||||||
_, _, err := mw.BeforeModelRewriteState(context.Background(), state, nil)
|
|
||||||
if err == nil || !strings.Contains(err.Error(), "hard budget") {
|
|
||||||
t.Fatalf("expected local hard-budget error, got %v", err)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
@@ -0,0 +1,66 @@
|
|||||||
|
package multiagent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/adk"
|
||||||
|
"github.com/cloudwego/eino/adk/middlewares/summarization"
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
// modelInputSoftBudgetMiddleware is the final guard before a normal model call.
|
||||||
|
// It drops oldest complete rounds and truncates oversized tool output in the latest
|
||||||
|
// round, but never fails locally — API context limits are handled by overflow retry.
|
||||||
|
type modelInputSoftBudgetMiddleware struct {
|
||||||
|
adk.BaseChatModelAgentMiddleware
|
||||||
|
maxTokens int
|
||||||
|
toolMaxBytes int
|
||||||
|
counter summarization.TokenCounterFunc
|
||||||
|
logger *zap.Logger
|
||||||
|
phase string
|
||||||
|
}
|
||||||
|
|
||||||
|
func newModelInputSoftBudgetMiddleware(
|
||||||
|
maxTotalTokens int,
|
||||||
|
toolMaxBytes int,
|
||||||
|
modelName string,
|
||||||
|
logger *zap.Logger,
|
||||||
|
phase string,
|
||||||
|
) adk.ChatModelAgentMiddleware {
|
||||||
|
if maxTotalTokens <= 0 {
|
||||||
|
maxTotalTokens = 120000
|
||||||
|
}
|
||||||
|
if toolMaxBytes <= 0 {
|
||||||
|
toolMaxBytes = 12000
|
||||||
|
}
|
||||||
|
return &modelInputSoftBudgetMiddleware{
|
||||||
|
maxTokens: maxTotalTokens,
|
||||||
|
toolMaxBytes: toolMaxBytes,
|
||||||
|
counter: einoSummarizationTokenCounter(modelName),
|
||||||
|
logger: logger,
|
||||||
|
phase: phase,
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *modelInputSoftBudgetMiddleware) BeforeModelRewriteState(
|
||||||
|
ctx context.Context,
|
||||||
|
state *adk.ChatModelAgentState,
|
||||||
|
mc *adk.ModelContext,
|
||||||
|
) (context.Context, *adk.ChatModelAgentState, error) {
|
||||||
|
if m == nil || state == nil || len(state.Messages) == 0 {
|
||||||
|
return ctx, state, nil
|
||||||
|
}
|
||||||
|
compacted, changed := compactMessagesByDroppingRounds(ctx, state.Messages, compactMessagesOpts{
|
||||||
|
maxTokens: m.maxTokens,
|
||||||
|
counter: m.counter,
|
||||||
|
toolMaxBytes: m.toolMaxBytes,
|
||||||
|
phase: m.phase,
|
||||||
|
logger: m.logger,
|
||||||
|
})
|
||||||
|
if !changed {
|
||||||
|
return ctx, state, nil
|
||||||
|
}
|
||||||
|
out := *state
|
||||||
|
out.Messages = compacted
|
||||||
|
return ctx, &out, nil
|
||||||
|
}
|
||||||
@@ -252,6 +252,7 @@ func RunDeepAgent(
|
|||||||
summarization: subSumMw,
|
summarization: subSumMw,
|
||||||
modelName: appCfg.OpenAI.Model,
|
modelName: appCfg.OpenAI.Model,
|
||||||
maxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
maxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
||||||
|
toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware),
|
||||||
conversationID: conversationID,
|
conversationID: conversationID,
|
||||||
})
|
})
|
||||||
|
|
||||||
@@ -413,6 +414,7 @@ func RunDeepAgent(
|
|||||||
summarization: mainSumMw,
|
summarization: mainSumMw,
|
||||||
modelName: appCfg.OpenAI.Model,
|
modelName: appCfg.OpenAI.Model,
|
||||||
maxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
maxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
||||||
|
toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware),
|
||||||
conversationID: conversationID,
|
conversationID: conversationID,
|
||||||
trace: modelFacingTrace,
|
trace: modelFacingTrace,
|
||||||
})
|
})
|
||||||
@@ -430,6 +432,7 @@ func RunDeepAgent(
|
|||||||
summarization: mainSumMw,
|
summarization: mainSumMw,
|
||||||
modelName: appCfg.OpenAI.Model,
|
modelName: appCfg.OpenAI.Model,
|
||||||
maxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
maxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
||||||
|
toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware),
|
||||||
conversationID: conversationID,
|
conversationID: conversationID,
|
||||||
trace: modelFacingTrace,
|
trace: modelFacingTrace,
|
||||||
})
|
})
|
||||||
@@ -505,6 +508,7 @@ func RunDeepAgent(
|
|||||||
summarization: mainSumMw,
|
summarization: mainSumMw,
|
||||||
modelName: appCfg.OpenAI.Model,
|
modelName: appCfg.OpenAI.Model,
|
||||||
maxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
maxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
||||||
|
toolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware),
|
||||||
conversationID: conversationID,
|
conversationID: conversationID,
|
||||||
skipTrace: true,
|
skipTrace: true,
|
||||||
}),
|
}),
|
||||||
@@ -608,6 +612,9 @@ func RunDeepAgent(
|
|||||||
DA: da,
|
DA: da,
|
||||||
ModelFacingTrace: modelFacingTrace,
|
ModelFacingTrace: modelFacingTrace,
|
||||||
EinoCallbacks: &ma.EinoCallbacks,
|
EinoCallbacks: &ma.EinoCallbacks,
|
||||||
|
MaxTotalTokens: appCfg.OpenAI.MaxTotalTokens,
|
||||||
|
ToolMaxBytes: toolMaxBytesFromMW(&ma.EinoMiddleware),
|
||||||
|
ModelName: appCfg.OpenAI.Model,
|
||||||
EmptyResponseMessage: "(Eino multi-agent orchestration completed but no assistant text was captured. Check process details or logs.) " +
|
EmptyResponseMessage: "(Eino multi-agent orchestration completed but no assistant text was captured. Check process details or logs.) " +
|
||||||
"(Eino 多代理编排已完成,但未捕获到助手文本输出。请查看过程详情或日志。)",
|
"(Eino 多代理编排已完成,但未捕获到助手文本输出。请查看过程详情或日志。)",
|
||||||
}, baseMsgs)
|
}, baseMsgs)
|
||||||
|
|||||||
@@ -0,0 +1,76 @@
|
|||||||
|
package multiagent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/adk"
|
||||||
|
"github.com/cloudwego/eino/schema"
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
// toolSearchResultSanitizerMiddleware prevents malformed historical tool_search
|
||||||
|
// results (for example an HTML gateway error page) from crashing Eino's dynamic
|
||||||
|
// tool loader on every retry. Eino expects every tool_search result to be a JSON
|
||||||
|
// object containing selectedTools.
|
||||||
|
type toolSearchResultSanitizerMiddleware struct {
|
||||||
|
adk.BaseChatModelAgentMiddleware
|
||||||
|
logger *zap.Logger
|
||||||
|
phase string
|
||||||
|
}
|
||||||
|
|
||||||
|
func newToolSearchResultSanitizerMiddleware(logger *zap.Logger, phase string) adk.ChatModelAgentMiddleware {
|
||||||
|
return &toolSearchResultSanitizerMiddleware{logger: logger, phase: phase}
|
||||||
|
}
|
||||||
|
|
||||||
|
type toolSearchResultEnvelope struct {
|
||||||
|
SelectedTools []string `json:"selectedTools"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func validToolSearchResult(content string) bool {
|
||||||
|
var result toolSearchResultEnvelope
|
||||||
|
if err := json.Unmarshal([]byte(content), &result); err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
// Reject JSON values such as null. They unmarshal without an error but do not
|
||||||
|
// satisfy the object-shaped contract used by the toolsearch middleware.
|
||||||
|
return strings.HasPrefix(strings.TrimSpace(content), "{")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (m *toolSearchResultSanitizerMiddleware) BeforeModelRewriteState(
|
||||||
|
ctx context.Context,
|
||||||
|
state *adk.ChatModelAgentState,
|
||||||
|
_ *adk.ModelContext,
|
||||||
|
) (context.Context, *adk.ChatModelAgentState, error) {
|
||||||
|
if m == nil || state == nil || len(state.Messages) == 0 {
|
||||||
|
return ctx, state, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
var rewritten []adk.Message
|
||||||
|
repaired := 0
|
||||||
|
for i, msg := range state.Messages {
|
||||||
|
if msg == nil || msg.Role != schema.Tool || !IsToolSearchTool(msg.ToolName) || validToolSearchResult(msg.Content) {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if rewritten == nil {
|
||||||
|
rewritten = append([]adk.Message(nil), state.Messages...)
|
||||||
|
}
|
||||||
|
clone := *msg
|
||||||
|
clone.Content = `{"selectedTools":[],"_recovered":true,"reason":"invalid historical tool_search result"}`
|
||||||
|
rewritten[i] = &clone
|
||||||
|
repaired++
|
||||||
|
}
|
||||||
|
|
||||||
|
if repaired == 0 {
|
||||||
|
return ctx, state, nil
|
||||||
|
}
|
||||||
|
if m.logger != nil {
|
||||||
|
m.logger.Warn("invalid historical tool_search results repaired before model call",
|
||||||
|
zap.String("phase", m.phase),
|
||||||
|
zap.Int("repaired_count", repaired))
|
||||||
|
}
|
||||||
|
ns := *state
|
||||||
|
ns.Messages = rewritten
|
||||||
|
return ctx, &ns, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,56 @@
|
|||||||
|
package multiagent
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/adk"
|
||||||
|
"github.com/cloudwego/eino/schema"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestToolSearchResultSanitizerRepairsMalformedHistory(t *testing.T) {
|
||||||
|
good := &schema.Message{Role: schema.Tool, ToolName: "tool_search", Content: `{"selectedTools":["grep"]}`}
|
||||||
|
bad := &schema.Message{Role: schema.Tool, ToolName: "tool_search", Content: "<html>502 Bad Gateway</html>"}
|
||||||
|
other := &schema.Message{Role: schema.Tool, ToolName: "grep", Content: "plain text is valid for other tools"}
|
||||||
|
state := &adk.ChatModelAgentState{Messages: []adk.Message{good, bad, other}}
|
||||||
|
|
||||||
|
mw := newToolSearchResultSanitizerMiddleware(nil, "test")
|
||||||
|
_, got, err := mw.BeforeModelRewriteState(context.Background(), state, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("BeforeModelRewriteState: %v", err)
|
||||||
|
}
|
||||||
|
if got.Messages[0] != good || got.Messages[0].Content != good.Content {
|
||||||
|
t.Fatal("valid tool_search result was unexpectedly changed")
|
||||||
|
}
|
||||||
|
if got.Messages[1] == bad || !validToolSearchResult(got.Messages[1].Content) {
|
||||||
|
t.Fatalf("malformed result was not safely replaced: %q", got.Messages[1].Content)
|
||||||
|
}
|
||||||
|
if got.Messages[2] != other {
|
||||||
|
t.Fatal("non-tool_search result was unexpectedly changed")
|
||||||
|
}
|
||||||
|
if bad.Content != "<html>502 Bad Gateway</html>" {
|
||||||
|
t.Fatal("middleware mutated the original message")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestToolSearchResultSanitizerFastPath(t *testing.T) {
|
||||||
|
msg := &schema.Message{Role: schema.Tool, ToolName: "tool_search", Content: `{"selectedTools":[]}`}
|
||||||
|
state := &adk.ChatModelAgentState{Messages: []adk.Message{msg}}
|
||||||
|
mw := newToolSearchResultSanitizerMiddleware(nil, "test")
|
||||||
|
|
||||||
|
_, got, err := mw.BeforeModelRewriteState(context.Background(), state, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("BeforeModelRewriteState: %v", err)
|
||||||
|
}
|
||||||
|
if got != state {
|
||||||
|
t.Fatal("valid history should use the allocation-free fast path")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestValidToolSearchResultRejectsNonObjectJSON(t *testing.T) {
|
||||||
|
for _, content := range []string{"null", `[]`, `"text"`, `{"selectedTools":"grep"}`} {
|
||||||
|
if validToolSearchResult(content) {
|
||||||
|
t.Fatalf("expected invalid tool_search result: %s", content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,192 @@
|
|||||||
|
package robot
|
||||||
|
|
||||||
|
import (
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"net/http"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"cyberstrike-ai/internal/config"
|
||||||
|
|
||||||
|
"github.com/bwmarrin/discordgo"
|
||||||
|
lark "github.com/larksuite/oapi-sdk-go/v3"
|
||||||
|
larkim "github.com/larksuite/oapi-sdk-go/v3/service/im/v1"
|
||||||
|
"github.com/slack-go/slack"
|
||||||
|
)
|
||||||
|
|
||||||
|
// SendProactive sends a message without an inbound event. Platforms whose
|
||||||
|
// reply credentials are event-scoped deliberately return an error instead of
|
||||||
|
// pretending delivery succeeded.
|
||||||
|
func SendProactive(ctx context.Context, cfg config.RobotsConfig, platform, externalUserID, message string) error {
|
||||||
|
platform = strings.ToLower(strings.TrimSpace(platform))
|
||||||
|
userID := robotIdentityUserPart(externalUserID)
|
||||||
|
if userID == "" {
|
||||||
|
return fmt.Errorf("invalid robot recipient")
|
||||||
|
}
|
||||||
|
switch platform {
|
||||||
|
case "telegram":
|
||||||
|
if !cfg.Telegram.Enabled || strings.TrimSpace(cfg.Telegram.BotToken) == "" {
|
||||||
|
return fmt.Errorf("telegram is not configured")
|
||||||
|
}
|
||||||
|
id, err := strconv.ParseInt(userID, 10, 64)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("invalid telegram user id: %w", err)
|
||||||
|
}
|
||||||
|
return telegramSendReply(ctx, nilSafeHTTPClient(), strings.TrimSpace(cfg.Telegram.BotToken), id, message)
|
||||||
|
case "slack":
|
||||||
|
if !cfg.Slack.Enabled || strings.TrimSpace(cfg.Slack.BotToken) == "" {
|
||||||
|
return fmt.Errorf("slack is not configured")
|
||||||
|
}
|
||||||
|
api := slack.New(strings.TrimSpace(cfg.Slack.BotToken))
|
||||||
|
channel, _, _, err := api.OpenConversationContext(ctx, &slack.OpenConversationParameters{Users: []string{userID}})
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for _, chunk := range splitTextChunks(message, slackMaxMessageRunes) {
|
||||||
|
if _, _, err = api.PostMessageContext(ctx, channel.ID, slack.MsgOptionText(chunk, false)); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
case "discord":
|
||||||
|
if !cfg.Discord.Enabled || strings.TrimSpace(cfg.Discord.BotToken) == "" {
|
||||||
|
return fmt.Errorf("discord is not configured")
|
||||||
|
}
|
||||||
|
token := strings.TrimSpace(cfg.Discord.BotToken)
|
||||||
|
if !strings.HasPrefix(token, "Bot ") {
|
||||||
|
token = "Bot " + token
|
||||||
|
}
|
||||||
|
session, err := discordgo.New(token)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
channel, err := session.UserChannelCreate(userID)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
for _, chunk := range splitTextChunks(message, discordMaxMessageRunes) {
|
||||||
|
if _, err = session.ChannelMessageSend(channel.ID, chunk); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
case "wecom":
|
||||||
|
return sendWecomProactive(ctx, cfg.Wecom, userID, message)
|
||||||
|
case "lark":
|
||||||
|
return sendLarkProactive(ctx, cfg.Lark, externalUserID, message)
|
||||||
|
default:
|
||||||
|
return fmt.Errorf("platform %s does not support proactive alerts yet", platform)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func SupportsProactive(platform string) bool {
|
||||||
|
switch strings.ToLower(strings.TrimSpace(platform)) {
|
||||||
|
case "telegram", "slack", "discord", "wecom", "lark":
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func nilSafeHTTPClient() *http.Client { return &http.Client{Timeout: 15 * time.Second} }
|
||||||
|
|
||||||
|
func sendWecomProactive(ctx context.Context, cfg config.RobotWecomConfig, userID, message string) error {
|
||||||
|
if !cfg.Enabled || strings.TrimSpace(cfg.CorpID) == "" || strings.TrimSpace(cfg.Secret) == "" || cfg.AgentID == 0 {
|
||||||
|
return fmt.Errorf("wecom proactive API is not configured")
|
||||||
|
}
|
||||||
|
tokenURL := "https://qyapi.weixin.qq.com/cgi-bin/gettoken?corpid=" + cfg.CorpID + "&corpsecret=" + cfg.Secret
|
||||||
|
req, err := http.NewRequestWithContext(ctx, http.MethodGet, tokenURL, nil)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
resp, err := nilSafeHTTPClient().Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
body, err := io.ReadAll(resp.Body)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
var tokenResp struct {
|
||||||
|
ErrCode int `json:"errcode"`
|
||||||
|
ErrMsg string `json:"errmsg"`
|
||||||
|
AccessToken string `json:"access_token"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(body, &tokenResp); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if tokenResp.ErrCode != 0 || tokenResp.AccessToken == "" {
|
||||||
|
return fmt.Errorf("wecom token: %s", tokenResp.ErrMsg)
|
||||||
|
}
|
||||||
|
payload, _ := json.Marshal(map[string]interface{}{
|
||||||
|
"touser": userID, "msgtype": "text", "agentid": cfg.AgentID,
|
||||||
|
"text": map[string]string{"content": message}, "safe": 0,
|
||||||
|
})
|
||||||
|
sendReq, err := http.NewRequestWithContext(ctx, http.MethodPost, "https://qyapi.weixin.qq.com/cgi-bin/message/send?access_token="+tokenResp.AccessToken, bytes.NewReader(payload))
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
sendReq.Header.Set("Content-Type", "application/json")
|
||||||
|
sendResp, err := nilSafeHTTPClient().Do(sendReq)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
defer sendResp.Body.Close()
|
||||||
|
result, _ := io.ReadAll(sendResp.Body)
|
||||||
|
var parsed struct {
|
||||||
|
ErrCode int `json:"errcode"`
|
||||||
|
ErrMsg string `json:"errmsg"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(result, &parsed); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if parsed.ErrCode != 0 {
|
||||||
|
return fmt.Errorf("wecom send: %s", parsed.ErrMsg)
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func sendLarkProactive(ctx context.Context, cfg config.RobotLarkConfig, identity, message string) error {
|
||||||
|
if !cfg.Enabled || strings.TrimSpace(cfg.AppID) == "" || strings.TrimSpace(cfg.AppSecret) == "" {
|
||||||
|
return fmt.Errorf("lark is not configured")
|
||||||
|
}
|
||||||
|
receiveIDType, receiveID := "user_id", robotIdentityUserPart(identity)
|
||||||
|
if idx := strings.LastIndex(identity, "|o:"); idx >= 0 {
|
||||||
|
receiveIDType, receiveID = "open_id", strings.TrimSpace(identity[idx+3:])
|
||||||
|
}
|
||||||
|
if idx := strings.LastIndex(identity, "|n:"); idx >= 0 {
|
||||||
|
receiveIDType, receiveID = "union_id", strings.TrimSpace(identity[idx+3:])
|
||||||
|
}
|
||||||
|
if receiveID == "" {
|
||||||
|
return fmt.Errorf("invalid lark recipient")
|
||||||
|
}
|
||||||
|
content, _ := json.Marshal(larkTextContent{Text: message})
|
||||||
|
client := lark.NewClient(cfg.AppID, cfg.AppSecret)
|
||||||
|
resp, err := client.Im.Message.Create(ctx, larkim.NewCreateMessageReqBuilder().
|
||||||
|
ReceiveIdType(receiveIDType).
|
||||||
|
Body(larkim.NewCreateMessageReqBodyBuilder().ReceiveId(receiveID).MsgType(larkim.MsgTypeText).Content(string(content)).Build()).Build())
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
if resp == nil || !resp.Success() {
|
||||||
|
return fmt.Errorf("lark send failed")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func robotIdentityUserPart(identity string) string {
|
||||||
|
identity = strings.TrimSpace(identity)
|
||||||
|
if i := strings.LastIndex(identity, "|u:"); i >= 0 {
|
||||||
|
return strings.TrimSpace(identity[i+3:])
|
||||||
|
}
|
||||||
|
if strings.HasPrefix(identity, "u:") {
|
||||||
|
return strings.TrimSpace(identity[2:])
|
||||||
|
}
|
||||||
|
return identity
|
||||||
|
}
|
||||||
@@ -32,7 +32,6 @@ type Session struct {
|
|||||||
|
|
||||||
// AuthManager manages password-based authentication and session lifecycle.
|
// AuthManager manages password-based authentication and session lifecycle.
|
||||||
type AuthManager struct {
|
type AuthManager struct {
|
||||||
password string
|
|
||||||
sessionDuration time.Duration
|
sessionDuration time.Duration
|
||||||
db *database.DB
|
db *database.DB
|
||||||
|
|
||||||
@@ -41,39 +40,49 @@ type AuthManager struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// NewAuthManager creates a new AuthManager instance.
|
// NewAuthManager creates a new AuthManager instance.
|
||||||
func NewAuthManager(password string, sessionDurationHours int) (*AuthManager, error) {
|
func NewAuthManager(sessionDurationHours int) *AuthManager {
|
||||||
if strings.TrimSpace(password) == "" {
|
|
||||||
return nil, errors.New("auth password must be configured")
|
|
||||||
}
|
|
||||||
|
|
||||||
if sessionDurationHours <= 0 {
|
if sessionDurationHours <= 0 {
|
||||||
sessionDurationHours = 12
|
sessionDurationHours = 12
|
||||||
}
|
}
|
||||||
|
|
||||||
return &AuthManager{
|
return &AuthManager{
|
||||||
password: password,
|
|
||||||
sessionDuration: time.Duration(sessionDurationHours) * time.Hour,
|
sessionDuration: time.Duration(sessionDurationHours) * time.Hour,
|
||||||
sessions: make(map[string]Session),
|
sessions: make(map[string]Session),
|
||||||
}, nil
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// AttachRBACStore enables multi-user RBAC authentication and bootstraps the
|
// AttachRBACStore enables multi-user RBAC authentication. When no users exist yet,
|
||||||
// built-in admin account from the legacy auth.password value.
|
// it bootstraps the built-in admin account and returns the generated initial password.
|
||||||
func (a *AuthManager) AttachRBACStore(db *database.DB) error {
|
func (a *AuthManager) AttachRBACStore(db *database.DB) (generatedAdminPassword string, err error) {
|
||||||
if db == nil {
|
if db == nil {
|
||||||
return nil
|
return "", errors.New("database is required for authentication")
|
||||||
}
|
}
|
||||||
hash, err := HashPassword(a.password)
|
|
||||||
|
needsAdminPassword, err := db.RBACNeedsAdminPassword()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return "", err
|
||||||
}
|
}
|
||||||
if err := db.BootstrapRBAC(hash, PermissionCatalog); err != nil {
|
|
||||||
return err
|
adminPasswordHash := ""
|
||||||
|
if needsAdminPassword {
|
||||||
|
generatedAdminPassword, err = GenerateStrongPassword(24)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
adminPasswordHash, err = HashPassword(generatedAdminPassword)
|
||||||
|
if err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
if err := db.BootstrapRBAC(adminPasswordHash, PermissionCatalog); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
a.mu.Lock()
|
a.mu.Lock()
|
||||||
a.db = db
|
a.db = db
|
||||||
a.mu.Unlock()
|
a.mu.Unlock()
|
||||||
return nil
|
return generatedAdminPassword, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// Authenticate validates the password and creates a new session.
|
// Authenticate validates the password and creates a new session.
|
||||||
@@ -94,23 +103,9 @@ func (a *AuthManager) authenticateSession(username, password string) (Session, e
|
|||||||
|
|
||||||
a.mu.RLock()
|
a.mu.RLock()
|
||||||
db := a.db
|
db := a.db
|
||||||
legacyPassword := a.password
|
|
||||||
a.mu.RUnlock()
|
a.mu.RUnlock()
|
||||||
|
|
||||||
if db == nil {
|
if db == nil {
|
||||||
if password != legacyPassword {
|
return Session{}, errors.New("authentication store is not configured")
|
||||||
return Session{}, ErrInvalidPassword
|
|
||||||
}
|
|
||||||
return Session{
|
|
||||||
Token: token,
|
|
||||||
ExpiresAt: expiresAt,
|
|
||||||
UserID: "admin",
|
|
||||||
Username: "admin",
|
|
||||||
DisplayName: "管理员",
|
|
||||||
Roles: []string{database.RBACSystemRoleAdmin},
|
|
||||||
Permissions: allPermissions(),
|
|
||||||
Scope: database.RBACScopeAll,
|
|
||||||
}, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
username = strings.TrimSpace(strings.ToLower(username))
|
username = strings.TrimSpace(strings.ToLower(username))
|
||||||
@@ -187,10 +182,9 @@ func (a *AuthManager) CheckPassword(password string) bool {
|
|||||||
func (a *AuthManager) CheckUserPassword(username, password string) bool {
|
func (a *AuthManager) CheckUserPassword(username, password string) bool {
|
||||||
a.mu.RLock()
|
a.mu.RLock()
|
||||||
db := a.db
|
db := a.db
|
||||||
legacyPassword := a.password
|
|
||||||
a.mu.RUnlock()
|
a.mu.RUnlock()
|
||||||
if db == nil {
|
if db == nil {
|
||||||
return password == legacyPassword
|
return false
|
||||||
}
|
}
|
||||||
user, err := db.GetRBACUserByUsername(username)
|
user, err := db.GetRBACUserByUsername(username)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -212,7 +206,7 @@ func (a *AuthManager) UpdateUserPassword(userID, password string) error {
|
|||||||
db := a.db
|
db := a.db
|
||||||
a.mu.RUnlock()
|
a.mu.RUnlock()
|
||||||
if db == nil {
|
if db == nil {
|
||||||
return a.UpdateConfig(password, a.SessionDurationHours())
|
return errors.New("authentication store is not configured")
|
||||||
}
|
}
|
||||||
if err := db.UpdateRBACUserPassword(userID, hash); err != nil {
|
if err := db.UpdateRBACUserPassword(userID, hash); err != nil {
|
||||||
return err
|
return err
|
||||||
@@ -263,41 +257,6 @@ func (a *AuthManager) SessionDurationHours() int {
|
|||||||
return int(a.sessionDuration / time.Hour)
|
return int(a.sessionDuration / time.Hour)
|
||||||
}
|
}
|
||||||
|
|
||||||
// UpdateConfig updates the password and session duration, revoking existing sessions.
|
|
||||||
func (a *AuthManager) UpdateConfig(password string, sessionDurationHours int) error {
|
|
||||||
password = strings.TrimSpace(password)
|
|
||||||
if password == "" {
|
|
||||||
return errors.New("auth password must be configured")
|
|
||||||
}
|
|
||||||
|
|
||||||
if sessionDurationHours <= 0 {
|
|
||||||
sessionDurationHours = 12
|
|
||||||
}
|
|
||||||
|
|
||||||
hash := ""
|
|
||||||
if a.db != nil {
|
|
||||||
var err error
|
|
||||||
hash, err = HashPassword(password)
|
|
||||||
if err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
a.mu.Lock()
|
|
||||||
a.password = password
|
|
||||||
a.sessionDuration = time.Duration(sessionDurationHours) * time.Hour
|
|
||||||
a.sessions = make(map[string]Session)
|
|
||||||
db := a.db
|
|
||||||
a.mu.Unlock()
|
|
||||||
|
|
||||||
if db != nil {
|
|
||||||
if err := db.UpdateRBACAdminPassword(hash); err != nil {
|
|
||||||
return err
|
|
||||||
}
|
|
||||||
}
|
|
||||||
return nil
|
|
||||||
}
|
|
||||||
|
|
||||||
func allPermissions() map[string]bool {
|
func allPermissions() map[string]bool {
|
||||||
out := make(map[string]bool, len(PermissionCatalog))
|
out := make(map[string]bool, len(PermissionCatalog))
|
||||||
for key := range PermissionCatalog {
|
for key := range PermissionCatalog {
|
||||||
|
|||||||
@@ -0,0 +1,38 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"path/filepath"
|
||||||
|
"testing"
|
||||||
|
|
||||||
|
"cyberstrike-ai/internal/database"
|
||||||
|
|
||||||
|
"go.uber.org/zap"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestAttachRBACStoreBootstrapsAdminPassword(t *testing.T) {
|
||||||
|
db, err := database.NewDB(filepath.Join(t.TempDir(), "auth-bootstrap.db"), zap.NewNop())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewDB: %v", err)
|
||||||
|
}
|
||||||
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
|
||||||
|
manager := NewAuthManager(12)
|
||||||
|
generated, err := manager.AttachRBACStore(db)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("AttachRBACStore: %v", err)
|
||||||
|
}
|
||||||
|
if generated == "" {
|
||||||
|
t.Fatal("expected generated admin password on first bootstrap")
|
||||||
|
}
|
||||||
|
if !manager.CheckUserPassword("admin", generated) {
|
||||||
|
t.Fatal("generated password should authenticate admin")
|
||||||
|
}
|
||||||
|
|
||||||
|
second, err := manager.AttachRBACStore(db)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("AttachRBACStore second call: %v", err)
|
||||||
|
}
|
||||||
|
if second != "" {
|
||||||
|
t.Fatalf("expected no password on second bootstrap, got %q", second)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -20,11 +20,8 @@ func TestAuthManagerAuthenticatesCreatedRBACUser(t *testing.T) {
|
|||||||
}
|
}
|
||||||
t.Cleanup(func() { _ = db.Close() })
|
t.Cleanup(func() { _ = db.Close() })
|
||||||
|
|
||||||
manager, err := NewAuthManager("admin-secret", 12)
|
manager := NewAuthManager(12)
|
||||||
if err != nil {
|
if _, err := manager.AttachRBACStore(db); err != nil {
|
||||||
t.Fatalf("NewAuthManager: %v", err)
|
|
||||||
}
|
|
||||||
if err := manager.AttachRBACStore(db); err != nil {
|
|
||||||
t.Fatalf("AttachRBACStore: %v", err)
|
t.Fatalf("AttachRBACStore: %v", err)
|
||||||
}
|
}
|
||||||
hash, err := HashPassword("operator-secret")
|
hash, err := HashPassword("operator-secret")
|
||||||
|
|||||||
@@ -65,7 +65,7 @@ func (e *Executor) buildToolIndex() {
|
|||||||
e.toolIndex[e.config.Tools[i].Name] = &e.config.Tools[i]
|
e.toolIndex[e.config.Tools[i].Name] = &e.config.Tools[i]
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
e.logger.Info("工具索引构建完成",
|
e.logger.Debug("工具索引构建完成",
|
||||||
zap.Int("totalTools", len(e.config.Tools)),
|
zap.Int("totalTools", len(e.config.Tools)),
|
||||||
zap.Int("enabledTools", len(e.toolIndex)),
|
zap.Int("enabledTools", len(e.toolIndex)),
|
||||||
)
|
)
|
||||||
@@ -73,14 +73,14 @@ func (e *Executor) buildToolIndex() {
|
|||||||
|
|
||||||
// ExecuteTool 执行安全工具
|
// ExecuteTool 执行安全工具
|
||||||
func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[string]interface{}) (*mcp.ToolResult, error) {
|
func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[string]interface{}) (*mcp.ToolResult, error) {
|
||||||
e.logger.Info("ExecuteTool被调用",
|
e.logger.Debug("ExecuteTool被调用",
|
||||||
zap.String("toolName", toolName),
|
zap.String("toolName", toolName),
|
||||||
zap.Any("args", args),
|
zap.Any("args", args),
|
||||||
)
|
)
|
||||||
|
|
||||||
// 特殊处理:exec工具直接执行系统命令
|
// 特殊处理:exec工具直接执行系统命令
|
||||||
if toolName == "exec" {
|
if toolName == "exec" {
|
||||||
e.logger.Info("执行exec工具")
|
e.logger.Debug("执行exec工具")
|
||||||
return e.executeSystemCommand(ctx, args)
|
return e.executeSystemCommand(ctx, args)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -95,7 +95,7 @@ func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[st
|
|||||||
return nil, fmt.Errorf("工具 %s 未找到或未启用", toolName)
|
return nil, fmt.Errorf("工具 %s 未找到或未启用", toolName)
|
||||||
}
|
}
|
||||||
|
|
||||||
e.logger.Info("找到工具配置",
|
e.logger.Debug("找到工具配置",
|
||||||
zap.String("toolName", toolName),
|
zap.String("toolName", toolName),
|
||||||
zap.String("command", toolConfig.Command),
|
zap.String("command", toolConfig.Command),
|
||||||
zap.Strings("args", toolConfig.Args),
|
zap.Strings("args", toolConfig.Args),
|
||||||
@@ -103,7 +103,7 @@ func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[st
|
|||||||
|
|
||||||
// 特殊处理:内部工具(command 以 "internal:" 开头)
|
// 特殊处理:内部工具(command 以 "internal:" 开头)
|
||||||
if strings.HasPrefix(toolConfig.Command, "internal:") {
|
if strings.HasPrefix(toolConfig.Command, "internal:") {
|
||||||
e.logger.Info("执行内部工具",
|
e.logger.Debug("执行内部工具",
|
||||||
zap.String("toolName", toolName),
|
zap.String("toolName", toolName),
|
||||||
zap.String("command", toolConfig.Command),
|
zap.String("command", toolConfig.Command),
|
||||||
)
|
)
|
||||||
@@ -113,7 +113,7 @@ func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[st
|
|||||||
// 构建命令 - 根据工具类型使用不同的参数格式
|
// 构建命令 - 根据工具类型使用不同的参数格式
|
||||||
cmdArgs := e.buildCommandArgs(toolName, toolConfig, args)
|
cmdArgs := e.buildCommandArgs(toolName, toolConfig, args)
|
||||||
|
|
||||||
e.logger.Info("构建命令参数完成",
|
e.logger.Debug("构建命令参数完成",
|
||||||
zap.String("toolName", toolName),
|
zap.String("toolName", toolName),
|
||||||
zap.Strings("cmdArgs", cmdArgs),
|
zap.Strings("cmdArgs", cmdArgs),
|
||||||
zap.Int("argsCount", len(cmdArgs)),
|
zap.Int("argsCount", len(cmdArgs)),
|
||||||
@@ -142,7 +142,7 @@ func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[st
|
|||||||
attachNonInteractiveStdin(cmd)
|
attachNonInteractiveStdin(cmd)
|
||||||
_ = prepareShellCmdSession(cmd)
|
_ = prepareShellCmdSession(cmd)
|
||||||
|
|
||||||
e.logger.Info("执行安全工具",
|
e.logger.Debug("执行安全工具",
|
||||||
zap.String("tool", toolName),
|
zap.String("tool", toolName),
|
||||||
zap.Strings("args", cmdArgs),
|
zap.Strings("args", cmdArgs),
|
||||||
)
|
)
|
||||||
@@ -180,7 +180,7 @@ func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[st
|
|||||||
if exitCode != nil && toolConfig.AllowedExitCodes != nil {
|
if exitCode != nil && toolConfig.AllowedExitCodes != nil {
|
||||||
for _, allowedCode := range toolConfig.AllowedExitCodes {
|
for _, allowedCode := range toolConfig.AllowedExitCodes {
|
||||||
if *exitCode == allowedCode {
|
if *exitCode == allowedCode {
|
||||||
e.logger.Info("工具执行完成(退出码在允许列表中)",
|
e.logger.Debug("工具执行完成(退出码在允许列表中)",
|
||||||
zap.String("tool", toolName),
|
zap.String("tool", toolName),
|
||||||
zap.Int("exitCode", *exitCode),
|
zap.Int("exitCode", *exitCode),
|
||||||
zap.String("output", string(output)),
|
zap.String("output", string(output)),
|
||||||
@@ -215,7 +215,7 @@ func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[st
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
e.logger.Info("工具执行成功",
|
e.logger.Debug("工具执行成功",
|
||||||
zap.String("tool", toolName),
|
zap.String("tool", toolName),
|
||||||
zap.String("output", string(output)),
|
zap.String("output", string(output)),
|
||||||
)
|
)
|
||||||
@@ -233,7 +233,7 @@ func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[st
|
|||||||
|
|
||||||
// RegisterTools 注册工具到MCP服务器
|
// RegisterTools 注册工具到MCP服务器
|
||||||
func (e *Executor) RegisterTools(mcpServer *mcp.Server) {
|
func (e *Executor) RegisterTools(mcpServer *mcp.Server) {
|
||||||
e.logger.Info("开始注册工具",
|
e.logger.Debug("开始注册工具",
|
||||||
zap.Int("totalTools", len(e.config.Tools)),
|
zap.Int("totalTools", len(e.config.Tools)),
|
||||||
zap.Int("enabledTools", len(e.toolIndex)),
|
zap.Int("enabledTools", len(e.toolIndex)),
|
||||||
)
|
)
|
||||||
@@ -281,7 +281,7 @@ func (e *Executor) RegisterTools(mcpServer *mcp.Server) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
handler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) {
|
handler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) {
|
||||||
e.logger.Info("工具handler被调用",
|
e.logger.Debug("工具handler被调用",
|
||||||
zap.String("toolName", toolName),
|
zap.String("toolName", toolName),
|
||||||
zap.Any("args", args),
|
zap.Any("args", args),
|
||||||
)
|
)
|
||||||
@@ -289,14 +289,14 @@ func (e *Executor) RegisterTools(mcpServer *mcp.Server) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
mcpServer.RegisterTool(tool, handler)
|
mcpServer.RegisterTool(tool, handler)
|
||||||
e.logger.Info("注册安全工具成功",
|
e.logger.Debug("注册安全工具成功",
|
||||||
zap.String("tool", toolConfigCopy.Name),
|
zap.String("tool", toolConfigCopy.Name),
|
||||||
zap.String("command", toolConfigCopy.Command),
|
zap.String("command", toolConfigCopy.Command),
|
||||||
zap.Int("index", i),
|
zap.Int("index", i),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
e.logger.Info("工具注册完成",
|
e.logger.Debug("工具注册完成",
|
||||||
zap.Int("registeredCount", len(e.config.Tools)),
|
zap.Int("registeredCount", len(e.config.Tools)),
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,24 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/rand"
|
||||||
|
"encoding/base64"
|
||||||
|
)
|
||||||
|
|
||||||
|
// GenerateStrongPassword returns a URL-safe random password of the given length.
|
||||||
|
func GenerateStrongPassword(length int) (string, error) {
|
||||||
|
if length <= 0 {
|
||||||
|
length = 24
|
||||||
|
}
|
||||||
|
|
||||||
|
randomBytes := make([]byte, length)
|
||||||
|
if _, err := rand.Read(randomBytes); err != nil {
|
||||||
|
return "", err
|
||||||
|
}
|
||||||
|
|
||||||
|
password := base64.RawURLEncoding.EncodeToString(randomBytes)
|
||||||
|
if len(password) > length {
|
||||||
|
password = password[:length]
|
||||||
|
}
|
||||||
|
return password, nil
|
||||||
|
}
|
||||||
@@ -151,6 +151,9 @@ func permissionForRequest(method, fullPath string) string {
|
|||||||
return crudPermission(method, "knowledge")
|
return crudPermission(method, "knowledge")
|
||||||
case strings.HasPrefix(path, "/vulnerabilities"):
|
case strings.HasPrefix(path, "/vulnerabilities"):
|
||||||
return crudPermission(method, "vulnerability")
|
return crudPermission(method, "vulnerability")
|
||||||
|
case strings.HasPrefix(path, "/vulnerability-alerts"):
|
||||||
|
// This endpoint only changes the authenticated user's own preference.
|
||||||
|
return "vulnerability:read"
|
||||||
case strings.HasPrefix(path, "/projects"):
|
case strings.HasPrefix(path, "/projects"):
|
||||||
return crudPermission(method, "project")
|
return crudPermission(method, "project")
|
||||||
case strings.HasPrefix(path, "/webshell"):
|
case strings.HasPrefix(path, "/webshell"):
|
||||||
@@ -161,6 +164,10 @@ func permissionForRequest(method, fullPath string) string {
|
|||||||
return crudPermission(method, "files")
|
return crudPermission(method, "files")
|
||||||
case strings.HasPrefix(path, "/roles"):
|
case strings.HasPrefix(path, "/roles"):
|
||||||
return crudPermission(method, "roles")
|
return crudPermission(method, "roles")
|
||||||
|
case path == "/workflows/:id/package":
|
||||||
|
return "workflow:read"
|
||||||
|
case strings.HasPrefix(path, "/workflow-package-inspections"), strings.HasPrefix(path, "/workflow-package-imports"):
|
||||||
|
return "workflow:write"
|
||||||
case strings.HasPrefix(path, "/workflows"):
|
case strings.HasPrefix(path, "/workflows"):
|
||||||
if path == "/workflows/validate" || path == "/workflows/dry-run" || strings.HasSuffix(path, "/resume") {
|
if path == "/workflows/validate" || path == "/workflows/dry-run" || strings.HasSuffix(path, "/resume") {
|
||||||
return "workflow:execute"
|
return "workflow:execute"
|
||||||
@@ -254,6 +261,9 @@ func isProcessGlobalMutationPath(path string) bool {
|
|||||||
// Workflow runs inherit conversation access; definitions are global.
|
// Workflow runs inherit conversation access; definitions are global.
|
||||||
return !strings.HasPrefix(path, "/workflows/runs/") && path != "/workflows/validate" && path != "/workflows/dry-run"
|
return !strings.HasPrefix(path, "/workflows/runs/") && path != "/workflows/validate" && path != "/workflows/dry-run"
|
||||||
}
|
}
|
||||||
|
if strings.HasPrefix(path, "/workflow-package-inspections") || strings.HasPrefix(path, "/workflow-package-imports") {
|
||||||
|
return true
|
||||||
|
}
|
||||||
if strings.HasPrefix(path, "/knowledge") {
|
if strings.HasPrefix(path, "/knowledge") {
|
||||||
return path != "/knowledge/search"
|
return path != "/knowledge/search"
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,20 @@
|
|||||||
|
package security
|
||||||
|
|
||||||
|
import (
|
||||||
|
"net/http"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestWorkflowPackageRoutesHaveExplicitWorkflowPermissions(t *testing.T) {
|
||||||
|
if got := permissionForRequest(http.MethodGet, "/api/workflows/:id/package"); got != "workflow:read" {
|
||||||
|
t.Fatalf("export permission=%q", got)
|
||||||
|
}
|
||||||
|
for _, path := range []string{"/api/workflow-package-inspections", "/api/workflow-package-inspections/:inspectionId", "/api/workflow-package-imports", "/api/workflow-package-imports/:importId"} {
|
||||||
|
if got := permissionForRequest(http.MethodGet, path); got != "workflow:write" {
|
||||||
|
t.Fatalf("%s permission=%q", path, got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !isProcessGlobalMutationPath("/workflow-package-imports") || !isProcessGlobalMutationPath("/workflow-package-inspections") {
|
||||||
|
t.Fatal("package mutations must require all-resource scope")
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,67 @@
|
|||||||
|
package termout
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
// StartupWebUIOptions configures the startup Web UI banner.
|
||||||
|
type StartupWebUIOptions struct {
|
||||||
|
Scheme string
|
||||||
|
Port int
|
||||||
|
SelfSigned bool
|
||||||
|
HTTPRedirect bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// PrintConfigCreated prints a short notice when config.yaml is bootstrapped.
|
||||||
|
func PrintConfigCreated() {
|
||||||
|
s := New(os.Stdout)
|
||||||
|
s.Println("")
|
||||||
|
s.Println(s.Green("✔ ") + s.Bold("已创建 config.yaml") + s.Dim("(来自 config.example.yaml)"))
|
||||||
|
s.BlankLine()
|
||||||
|
}
|
||||||
|
|
||||||
|
// PrintStartupWebUI prints a colored startup banner for the Web UI.
|
||||||
|
func PrintStartupWebUI(opts StartupWebUIOptions) {
|
||||||
|
s := New(os.Stdout)
|
||||||
|
scheme := opts.Scheme
|
||||||
|
if scheme == "" {
|
||||||
|
scheme = "http"
|
||||||
|
}
|
||||||
|
port := opts.Port
|
||||||
|
if port <= 0 {
|
||||||
|
port = 8080
|
||||||
|
}
|
||||||
|
url := fmt.Sprintf("%s://127.0.0.1:%d/", scheme, port)
|
||||||
|
|
||||||
|
s.BlankLine()
|
||||||
|
s.Println(s.Bold(s.Cyan("CYBERSTRIKE AI")) + s.Dim(" / secure workspace"))
|
||||||
|
s.Println(s.Dim(strings.Repeat("─", 60)))
|
||||||
|
s.Println(s.Green("● ONLINE") + " " + s.Bold(s.White(url)))
|
||||||
|
if opts.SelfSigned {
|
||||||
|
s.Println(s.Dim(" TLS ") + s.Yellow("self-signed") + s.Dim(" · accept the browser warning once"))
|
||||||
|
}
|
||||||
|
if opts.HTTPRedirect {
|
||||||
|
s.Println(s.Dim(" Redirect ") + fmt.Sprintf("http://127.0.0.1:%d/ → HTTPS", port))
|
||||||
|
}
|
||||||
|
s.BlankLine()
|
||||||
|
}
|
||||||
|
|
||||||
|
// PrintBootstrapAdminCredentials prints the initial admin password banner.
|
||||||
|
func PrintBootstrapAdminCredentials(password string) {
|
||||||
|
password = strings.TrimSpace(password)
|
||||||
|
if password == "" {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
s := New(os.Stdout)
|
||||||
|
s.Println(s.Bold(s.Yellow("ADMIN SETUP REQUIRED")))
|
||||||
|
s.Println(s.Dim(strings.Repeat("─", 60)))
|
||||||
|
s.Println(s.Dim(" Username ") + s.Bold(s.White("admin")))
|
||||||
|
s.Println(s.Dim(" Password ") + s.Bold(s.Yellow(password)))
|
||||||
|
s.BlankLine()
|
||||||
|
s.Println(s.Yellow(" ! ") + s.White("Store this password securely. It is shown only once."))
|
||||||
|
s.Println(s.Dim(" Change it in Settings immediately after signing in."))
|
||||||
|
s.BlankLine()
|
||||||
|
}
|
||||||
@@ -0,0 +1,76 @@
|
|||||||
|
package termout
|
||||||
|
|
||||||
|
import (
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestDisplayWidthEmoji(t *testing.T) {
|
||||||
|
if got := displayWidth("🚀"); got != 2 {
|
||||||
|
t.Fatalf("displayWidth(emoji) = %d, want 2", got)
|
||||||
|
}
|
||||||
|
if got := displayWidth("ab"); got != 2 {
|
||||||
|
t.Fatalf("displayWidth(ab) = %d, want 2", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestDisplayWidthIgnoresANSI(t *testing.T) {
|
||||||
|
s := New(nil)
|
||||||
|
colored := s.Bold("admin")
|
||||||
|
if got := displayWidth(colored); got != 5 {
|
||||||
|
t.Fatalf("displayWidth colored = %d, want 5", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPadRightDisplay(t *testing.T) {
|
||||||
|
got := padRightDisplay("pwd", 10)
|
||||||
|
if displayWidth(got) != 10 {
|
||||||
|
t.Fatalf("padded width = %d, want 10", displayWidth(got))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestColorDisabledWithoutTTY(t *testing.T) {
|
||||||
|
s := New(nil)
|
||||||
|
if s.enabled {
|
||||||
|
t.Fatal("expected colors disabled for nil writer")
|
||||||
|
}
|
||||||
|
if got := s.Cyan("x"); got != "x" {
|
||||||
|
t.Fatalf("Cyan without TTY = %q, want plain text", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPrintBootstrapAdminCredentialsEmpty(t *testing.T) {
|
||||||
|
PrintBootstrapAdminCredentials(" ")
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestPrintStartupWebUIOptions(t *testing.T) {
|
||||||
|
PrintStartupWebUI(StartupWebUIOptions{
|
||||||
|
Scheme: "https",
|
||||||
|
Port: 8080,
|
||||||
|
SelfSigned: true,
|
||||||
|
HTTPRedirect: true,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestBoxRowAlignedWidth(t *testing.T) {
|
||||||
|
s := New(nil)
|
||||||
|
rows := []string{
|
||||||
|
s.Bold("CyberStrikeAI") + s.White(" is ready"),
|
||||||
|
s.Dim("Web UI ") + s.Bold("https://127.0.0.1:8080/"),
|
||||||
|
}
|
||||||
|
inner := maxDisplayWidth(rows...)
|
||||||
|
for _, row := range rows {
|
||||||
|
line := s.boxRow(inner, row)
|
||||||
|
if !strings.Contains(line, "│") {
|
||||||
|
t.Fatalf("box row missing border: %q", line)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMaxDisplayWidth(t *testing.T) {
|
||||||
|
short := "abc"
|
||||||
|
long := "https://127.0.0.1:8080/"
|
||||||
|
if got := maxDisplayWidth(short, long); got != displayWidth(long) {
|
||||||
|
t.Fatalf("maxDisplayWidth = %d, want %d", got, displayWidth(long))
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,108 @@
|
|||||||
|
package termout
|
||||||
|
|
||||||
|
import (
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
codeReset = "\033[0m"
|
||||||
|
codeBold = "\033[1m"
|
||||||
|
codeDim = "\033[2m"
|
||||||
|
codeRed = "\033[31m"
|
||||||
|
codeGreen = "\033[32m"
|
||||||
|
codeYellow = "\033[33m"
|
||||||
|
codeBlue = "\033[34m"
|
||||||
|
codeCyan = "\033[36m"
|
||||||
|
codeWhite = "\033[97m"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Style wraps ANSI styling with TTY / NO_COLOR awareness.
|
||||||
|
type Style struct {
|
||||||
|
out io.Writer
|
||||||
|
enabled bool
|
||||||
|
}
|
||||||
|
|
||||||
|
// New creates a Style writing to out (typically os.Stdout).
|
||||||
|
func New(out io.Writer) *Style {
|
||||||
|
return &Style{out: out, enabled: colorEnabled(out)}
|
||||||
|
}
|
||||||
|
|
||||||
|
func colorEnabled(w io.Writer) bool {
|
||||||
|
if strings.TrimSpace(os.Getenv("NO_COLOR")) != "" {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
force := strings.TrimSpace(os.Getenv("FORCE_COLOR"))
|
||||||
|
if force == "1" || strings.EqualFold(force, "true") || strings.EqualFold(force, "yes") {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
f, ok := w.(*os.File)
|
||||||
|
if !ok {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
stat, err := f.Stat()
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return stat.Mode()&os.ModeCharDevice != 0
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Style) paint(code, text string) string {
|
||||||
|
if !s.enabled || text == "" {
|
||||||
|
return text
|
||||||
|
}
|
||||||
|
return code + text + codeReset
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Style) Bold(text string) string { return s.paint(codeBold, text) }
|
||||||
|
func (s *Style) Dim(text string) string { return s.paint(codeDim, text) }
|
||||||
|
func (s *Style) Red(text string) string { return s.paint(codeRed, text) }
|
||||||
|
func (s *Style) Green(text string) string { return s.paint(codeGreen, text) }
|
||||||
|
func (s *Style) Yellow(text string) string { return s.paint(codeYellow, text) }
|
||||||
|
func (s *Style) Blue(text string) string { return s.paint(codeBlue, text) }
|
||||||
|
func (s *Style) Cyan(text string) string { return s.paint(codeCyan, text) }
|
||||||
|
func (s *Style) White(text string) string { return s.paint(codeWhite, text) }
|
||||||
|
|
||||||
|
func (s *Style) Println(text string) {
|
||||||
|
_, _ = fmt.Fprintln(s.out, text)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Style) Printf(format string, args ...interface{}) {
|
||||||
|
_, _ = fmt.Fprintf(s.out, format, args...)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Style) BlankLine() {
|
||||||
|
s.Println("")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Style) boxTop(innerWidth int) string {
|
||||||
|
return s.Cyan("╭" + strings.Repeat("─", innerWidth+2) + "╮")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Style) boxBottom(innerWidth int) string {
|
||||||
|
return s.Cyan("╰" + strings.Repeat("─", innerWidth+2) + "╯")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Style) boxRow(innerWidth int, content string) string {
|
||||||
|
return s.Cyan("│ ") + padRightDisplay(content, innerWidth) + s.Cyan(" │")
|
||||||
|
}
|
||||||
|
|
||||||
|
func (s *Style) printBox(rows []string, minInner, maxInner int) {
|
||||||
|
inner := maxDisplayWidth(rows...)
|
||||||
|
if inner < minInner {
|
||||||
|
inner = minInner
|
||||||
|
}
|
||||||
|
if maxInner > 0 && inner > maxInner {
|
||||||
|
inner = maxInner
|
||||||
|
}
|
||||||
|
|
||||||
|
s.BlankLine()
|
||||||
|
s.Println(s.boxTop(inner))
|
||||||
|
for _, row := range rows {
|
||||||
|
s.Println(s.boxRow(inner, row))
|
||||||
|
}
|
||||||
|
s.Println(s.boxBottom(inner))
|
||||||
|
s.BlankLine()
|
||||||
|
}
|
||||||
@@ -0,0 +1,73 @@
|
|||||||
|
package termout
|
||||||
|
|
||||||
|
import (
|
||||||
|
"regexp"
|
||||||
|
"strings"
|
||||||
|
"unicode/utf8"
|
||||||
|
|
||||||
|
"golang.org/x/text/width"
|
||||||
|
)
|
||||||
|
|
||||||
|
var ansiEscapeRe = regexp.MustCompile(`\x1b\[[0-9;]*m`)
|
||||||
|
|
||||||
|
// displayWidth returns the terminal display width of text, ignoring ANSI codes.
|
||||||
|
func displayWidth(text string) int {
|
||||||
|
plain := ansiEscapeRe.ReplaceAllString(text, "")
|
||||||
|
w := 0
|
||||||
|
for _, r := range plain {
|
||||||
|
w += runeDisplayWidth(r)
|
||||||
|
}
|
||||||
|
return w
|
||||||
|
}
|
||||||
|
|
||||||
|
func runeDisplayWidth(r rune) int {
|
||||||
|
if r == utf8.RuneError {
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
// Most emoji / symbols render as double-width in modern terminals.
|
||||||
|
if isEmojiLikeRune(r) {
|
||||||
|
return 2
|
||||||
|
}
|
||||||
|
switch width.LookupRune(r).Kind() {
|
||||||
|
case width.EastAsianWide, width.EastAsianFullwidth:
|
||||||
|
return 2
|
||||||
|
default:
|
||||||
|
return 1
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func isEmojiLikeRune(r rune) bool {
|
||||||
|
switch {
|
||||||
|
case r >= 0x1F300 && r <= 0x1FAFF: // pictographs / emoji
|
||||||
|
return true
|
||||||
|
case r >= 0x2600 && r <= 0x27BF: // misc symbols
|
||||||
|
return true
|
||||||
|
case r >= 0x2300 && r <= 0x23FF: // misc technical (⌚ etc.)
|
||||||
|
return true
|
||||||
|
case r >= 0x2B50 && r <= 0x2B55:
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func padRightDisplay(text string, target int) string {
|
||||||
|
if target <= 0 {
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
gap := target - displayWidth(text)
|
||||||
|
if gap <= 0 {
|
||||||
|
return text
|
||||||
|
}
|
||||||
|
return text + strings.Repeat(" ", gap)
|
||||||
|
}
|
||||||
|
|
||||||
|
func maxDisplayWidth(rows ...string) int {
|
||||||
|
max := 0
|
||||||
|
for _, row := range rows {
|
||||||
|
if w := displayWidth(row); w > max {
|
||||||
|
max = w
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return max
|
||||||
|
}
|
||||||
@@ -89,7 +89,7 @@ func RegisterAnalyzeImageTool(mcpServer *mcp.Server, cfg *config.Config, logger
|
|||||||
|
|
||||||
mcpServer.RegisterTool(tool, handler)
|
mcpServer.RegisterTool(tool, handler)
|
||||||
if logger != nil {
|
if logger != nil {
|
||||||
logger.Info("vision: analyze_image 工具已注册", zap.String("model", cfg.Vision.Model))
|
logger.Debug("vision: analyze_image 工具已注册", zap.String("model", cfg.Vision.Model))
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -2,6 +2,7 @@ package workflow
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"fmt"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
@@ -9,6 +10,7 @@ import (
|
|||||||
"cyberstrike-ai/internal/config"
|
"cyberstrike-ai/internal/config"
|
||||||
"cyberstrike-ai/internal/database"
|
"cyberstrike-ai/internal/database"
|
||||||
|
|
||||||
|
"github.com/cloudwego/eino/compose"
|
||||||
"go.uber.org/zap"
|
"go.uber.org/zap"
|
||||||
)
|
)
|
||||||
|
|
||||||
@@ -235,6 +237,89 @@ func TestExecuteEinoGraph_linearStartOutput(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestExecuteEinoGraph_checkpointRestoresStartOutput(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
checkpointStore, err := newFileCheckPointStore(t.TempDir())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("new checkpoint store: %v", err)
|
||||||
|
}
|
||||||
|
state := newWorkflowLocalState(map[string]interface{}{"message": "ping"}, "run-checkpoint")
|
||||||
|
node := graphNode{ID: "start-1", Type: "start"}
|
||||||
|
wf := compose.NewWorkflow[WorkflowInput, WorkflowOutput](
|
||||||
|
compose.WithGenLocalState(func(context.Context) *WorkflowLocalState { return state }),
|
||||||
|
)
|
||||||
|
start := wf.AddLambdaNode("start-1", compose.InvokableLambda(func(_ context.Context, input WorkflowInput) (WorkflowNodeOutput, error) {
|
||||||
|
result := startOutputMap(node, input.Message, input.ConversationID, input.ProjectID)
|
||||||
|
state.NodeOutputs[node.ID] = result
|
||||||
|
state.NodeOutputs["condition-1"] = conditionOutputMap(graphNode{ID: "condition-1", Type: "condition"}, "{{inputs.message}} == ping", true)
|
||||||
|
state.NodeOutputs["tool-1"] = toolOutputMap(graphNode{ID: "tool-1", Type: "tool"}, "tool result", "lookup", map[string]any{"id": "1"}, "exec-1", false)
|
||||||
|
state.NodeOutputs["agent-1"] = agentOutputMap(graphNode{ID: "agent-1", Type: "agent"}, "agent result", "chat", []string{"exec-1"})
|
||||||
|
state.NodeOutputs["hitl-1"] = hitlOutputMap(graphNode{ID: "hitl-1", Type: "hitl"}, "completed", "approved", "continue?", "reviewer", true)
|
||||||
|
state.NodeOutputs["output-1"] = outputNodeOutputMap(graphNode{ID: "output-1", Type: "output"}, "result", "ping")
|
||||||
|
state.NodeOutputs["end-1"] = endOutputMap(graphNode{ID: "end-1", Type: "end"}, "done")
|
||||||
|
state.LastOutput = result
|
||||||
|
state.Outputs["seed"] = "preserved"
|
||||||
|
return result, nil
|
||||||
|
}))
|
||||||
|
outputNode := wf.AddLambdaNode("out-1", compose.InvokableLambda(func(_ context.Context, input WorkflowNodeOutput) (WorkflowNodeOutput, error) {
|
||||||
|
return input, nil
|
||||||
|
}))
|
||||||
|
start.AddInput(compose.START)
|
||||||
|
outputNode.AddInput("start-1")
|
||||||
|
wf.End().AddInput("out-1", compose.ToField("out-1"))
|
||||||
|
runnable, err := wf.Compile(ctx,
|
||||||
|
compose.WithCheckPointStore(checkpointStore),
|
||||||
|
compose.WithInterruptAfterNodes([]string{"start-1"}),
|
||||||
|
)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("compile: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err = runnable.Invoke(ctx, workflowInputFromMap(state.Inputs), compose.WithCheckPointID("run-checkpoint"))
|
||||||
|
info, ok := compose.ExtractInterruptInfo(err)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("invoke error = %v, want checkpoint interrupt", err)
|
||||||
|
}
|
||||||
|
restored, ok := info.State.(*WorkflowLocalState)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("checkpoint state = %T, want *WorkflowLocalState", info.State)
|
||||||
|
}
|
||||||
|
for nodeID, wantType := range map[string]string{
|
||||||
|
"start-1": "StartOutput",
|
||||||
|
"condition-1": "ConditionOutput",
|
||||||
|
"tool-1": "ToolOutput",
|
||||||
|
"agent-1": "AgentOutput",
|
||||||
|
"hitl-1": "HITLOutput",
|
||||||
|
"output-1": "OutputNodeOutput",
|
||||||
|
"end-1": "NodeOutputEnvelope",
|
||||||
|
} {
|
||||||
|
if got := fmt.Sprintf("%T", restored.NodeOutputs[nodeID]["typed"]); got != "workflow."+wantType {
|
||||||
|
t.Fatalf("restored %s typed output = %s, want workflow.%s", nodeID, got, wantType)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if got := valueFromPath("previous.message", restored); got != "ping" {
|
||||||
|
t.Fatalf("restored previous.message = %v, want ping", got)
|
||||||
|
}
|
||||||
|
if got := valueFromPath("inputs.message", restored); got != "ping" {
|
||||||
|
t.Fatalf("restored inputs.message = %v, want ping", got)
|
||||||
|
}
|
||||||
|
if got := valueFromPath("outputs.seed", restored); got != "preserved" {
|
||||||
|
t.Fatalf("restored outputs.seed = %v, want preserved", got)
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := runnable.Invoke(ctx, WorkflowInput{}, compose.WithCheckPointID("run-checkpoint"))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("resume checkpoint: %v", err)
|
||||||
|
}
|
||||||
|
output, ok := result["out-1"].(map[string]any)
|
||||||
|
if !ok {
|
||||||
|
t.Fatalf("resumed output type = %T, want map[string]any", result["out-1"])
|
||||||
|
}
|
||||||
|
if got := output["output"]; got != "ping" {
|
||||||
|
t.Fatalf("resumed output = %v, want ping", got)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestExecuteEinoGraph_conditionBranch(t *testing.T) {
|
func TestExecuteEinoGraph_conditionBranch(t *testing.T) {
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
SetCheckpointDir(t.TempDir())
|
SetCheckpointDir(t.TempDir())
|
||||||
|
|||||||
@@ -0,0 +1,98 @@
|
|||||||
|
package workflowpackage
|
||||||
|
|
||||||
|
import (
|
||||||
|
"archive/zip"
|
||||||
|
"bytes"
|
||||||
|
"encoding/json"
|
||||||
|
"fmt"
|
||||||
|
"path"
|
||||||
|
"strings"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
// Export builds a deterministic, human-readable single-workflow package.
|
||||||
|
func Export(source Document) ([]byte, ExportMetadata, error) {
|
||||||
|
source.ID = strings.TrimSpace(source.ID)
|
||||||
|
source.Name = strings.TrimSpace(source.Name)
|
||||||
|
if source.ID == "" || source.Name == "" || source.Version <= 0 || !safePackageWorkflowID(source.ID) {
|
||||||
|
return nil, ExportMetadata{}, fmt.Errorf("workflow id, name and version are required")
|
||||||
|
}
|
||||||
|
if strings.TrimSpace(source.GraphJSON) == "" {
|
||||||
|
return nil, ExportMetadata{}, fmt.Errorf("workflow graph_json is required")
|
||||||
|
}
|
||||||
|
contentHash, graphHash, payload, err := DocumentHashes(source)
|
||||||
|
if err != nil {
|
||||||
|
return nil, ExportMetadata{}, err
|
||||||
|
}
|
||||||
|
workflowPath := path.Join("workflows", source.ID+".json")
|
||||||
|
createdAt := source.UpdatedAt.UTC()
|
||||||
|
if createdAt.IsZero() {
|
||||||
|
createdAt = time.Unix(0, 0).UTC()
|
||||||
|
}
|
||||||
|
manifest := Manifest{
|
||||||
|
PackageFormat: PackageFormat,
|
||||||
|
FormatVersion: FormatVersion,
|
||||||
|
PackageID: "pkg_" + strings.TrimPrefix(contentHash, "sha256:")[:16],
|
||||||
|
CreatedAt: createdAt.Format(time.RFC3339),
|
||||||
|
Items: []ManifestItem{{
|
||||||
|
Type: "workflow",
|
||||||
|
Path: workflowPath,
|
||||||
|
SourceID: source.ID,
|
||||||
|
SourceRevision: source.Version,
|
||||||
|
ContentHash: contentHash,
|
||||||
|
GraphHash: graphHash,
|
||||||
|
}},
|
||||||
|
}
|
||||||
|
manifestBytes, err := json.Marshal(manifest)
|
||||||
|
if err != nil {
|
||||||
|
return nil, ExportMetadata{}, fmt.Errorf("marshal manifest: %w", err)
|
||||||
|
}
|
||||||
|
checksums := fmt.Sprintf("%s manifest.json\n%s %s\n", strings.TrimPrefix(sha256Prefixed(manifestBytes), "sha256:"), strings.TrimPrefix(contentHash, "sha256:"), workflowPath)
|
||||||
|
|
||||||
|
var out bytes.Buffer
|
||||||
|
zw := zip.NewWriter(&out)
|
||||||
|
for _, entry := range []struct {
|
||||||
|
name string
|
||||||
|
data []byte
|
||||||
|
}{
|
||||||
|
{name: "checksums.sha256", data: []byte(checksums)},
|
||||||
|
{name: "manifest.json", data: manifestBytes},
|
||||||
|
{name: workflowPath, data: payload},
|
||||||
|
} {
|
||||||
|
header := &zip.FileHeader{Name: entry.name, Method: zip.Store}
|
||||||
|
header.SetModTime(time.Unix(0, 0).UTC())
|
||||||
|
writer, err := zw.CreateHeader(header)
|
||||||
|
if err != nil {
|
||||||
|
return nil, ExportMetadata{}, fmt.Errorf("write %s: %w", entry.name, err)
|
||||||
|
}
|
||||||
|
if _, err := writer.Write(entry.data); err != nil {
|
||||||
|
return nil, ExportMetadata{}, fmt.Errorf("write %s: %w", entry.name, err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := zw.Close(); err != nil {
|
||||||
|
return nil, ExportMetadata{}, fmt.Errorf("close package: %w", err)
|
||||||
|
}
|
||||||
|
pkg := out.Bytes()
|
||||||
|
return pkg, ExportMetadata{
|
||||||
|
PackageHash: sha256Prefixed(pkg),
|
||||||
|
ContentHash: contentHash,
|
||||||
|
GraphHash: graphHash,
|
||||||
|
SourceRevision: source.Version,
|
||||||
|
FileName: source.ID + ".csapkg.zip",
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// DocumentHashes returns the canonical package item and graph hashes together
|
||||||
|
// with the canonical item bytes used by export and inspection persistence.
|
||||||
|
func DocumentHashes(source Document) (string, string, []byte, error) {
|
||||||
|
graph, err := canonicalJSON([]byte(source.GraphJSON))
|
||||||
|
if err != nil {
|
||||||
|
return "", "", nil, fmt.Errorf("canonicalize graph_json: %w", err)
|
||||||
|
}
|
||||||
|
source.GraphJSON = string(graph)
|
||||||
|
payload, err := json.Marshal(source)
|
||||||
|
if err != nil {
|
||||||
|
return "", "", nil, fmt.Errorf("marshal workflow payload: %w", err)
|
||||||
|
}
|
||||||
|
return sha256Prefixed(payload), sha256Prefixed(graph), payload, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,58 @@
|
|||||||
|
package workflowpackage
|
||||||
|
|
||||||
|
import (
|
||||||
|
"archive/zip"
|
||||||
|
"bytes"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func testDocument() Document {
|
||||||
|
return Document{
|
||||||
|
ID: "web-src-hunting",
|
||||||
|
Name: "Web SRC 猎洞",
|
||||||
|
Description: "面向 SRC Web 资产的侦察与漏洞候选流程",
|
||||||
|
Version: 18,
|
||||||
|
Enabled: true,
|
||||||
|
GraphJSON: `{"nodes":[{"id":"start-1","type":"start","label":"开始","position":{"x":0,"y":0},"config":{}},{"id":"out-1","type":"output","label":"输出","position":{"x":0,"y":120},"config":{"output_key":"result","source_binding":{"from":"inputs","field":"message"}}}],"edges":[{"id":"e1","source":"start-1","target":"out-1"}],"config":{"schema_version":1}}`,
|
||||||
|
UpdatedAt: time.Date(2026, 7, 13, 10, 0, 0, 0, time.UTC),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestExportIsDeterministicAndSelfDescribing(t *testing.T) {
|
||||||
|
first, firstMeta, err := Export(testDocument())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("first export: %v", err)
|
||||||
|
}
|
||||||
|
second, secondMeta, err := Export(testDocument())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("second export: %v", err)
|
||||||
|
}
|
||||||
|
if !bytes.Equal(first, second) {
|
||||||
|
t.Fatal("identical document must produce byte-identical package")
|
||||||
|
}
|
||||||
|
if firstMeta.PackageHash != secondMeta.PackageHash || !strings.HasPrefix(firstMeta.PackageHash, "sha256:") {
|
||||||
|
t.Fatalf("unexpected deterministic package hash: %#v / %#v", firstMeta, secondMeta)
|
||||||
|
}
|
||||||
|
|
||||||
|
zr, err := zip.NewReader(bytes.NewReader(first), int64(len(first)))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("open package: %v", err)
|
||||||
|
}
|
||||||
|
if len(zr.File) != 3 {
|
||||||
|
t.Fatalf("zip entry count = %d, want 3", len(zr.File))
|
||||||
|
}
|
||||||
|
wantNames := []string{"checksums.sha256", "manifest.json", "workflows/web-src-hunting.json"}
|
||||||
|
for i, f := range zr.File {
|
||||||
|
if f.Name != wantNames[i] {
|
||||||
|
t.Fatalf("entry %d = %q, want %q", i, f.Name, wantNames[i])
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if firstMeta.SourceRevision != 18 {
|
||||||
|
t.Fatalf("source revision = %d, want 18", firstMeta.SourceRevision)
|
||||||
|
}
|
||||||
|
if !strings.HasPrefix(firstMeta.ContentHash, "sha256:") || !strings.HasPrefix(firstMeta.GraphHash, "sha256:") {
|
||||||
|
t.Fatalf("content/graph hashes must be sha256: %#v", firstMeta)
|
||||||
|
}
|
||||||
|
}
|
||||||
@@ -0,0 +1,225 @@
|
|||||||
|
package workflowpackage
|
||||||
|
|
||||||
|
import (
|
||||||
|
"archive/zip"
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"encoding/json"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"path"
|
||||||
|
"strings"
|
||||||
|
"unicode"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
MaxArchiveBytes = 10 << 20
|
||||||
|
MaxExtractedBytes = 20 << 20
|
||||||
|
)
|
||||||
|
|
||||||
|
// PackageError contains only a contract error code and safe, client-facing fields.
|
||||||
|
type PackageError struct {
|
||||||
|
Code string
|
||||||
|
Message string
|
||||||
|
Details map[string]any
|
||||||
|
}
|
||||||
|
|
||||||
|
func (e *PackageError) Error() string { return e.Code + ": " + e.Message }
|
||||||
|
|
||||||
|
func packageError(code, message string) error {
|
||||||
|
return &PackageError{Code: code, Message: message}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ErrorCode returns a package contract code without exposing internal errors.
|
||||||
|
func ErrorCode(err error) string {
|
||||||
|
var target *PackageError
|
||||||
|
if errors.As(err, &target) {
|
||||||
|
return target.Code
|
||||||
|
}
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
|
type InspectionResult struct {
|
||||||
|
PackageHash string
|
||||||
|
Manifest Manifest
|
||||||
|
Document Document
|
||||||
|
ContentHash string
|
||||||
|
GraphHash string
|
||||||
|
NodeCount int
|
||||||
|
EdgeCount int
|
||||||
|
}
|
||||||
|
|
||||||
|
// InspectArchive verifies an archive without executing any package content.
|
||||||
|
// validateGraph is injected by the application so this format package has no
|
||||||
|
// dependency on the workflow runtime or database driver.
|
||||||
|
func InspectArchive(ctx context.Context, archive []byte, validateGraph func(context.Context, string) error) (*InspectionResult, error) {
|
||||||
|
if len(archive) == 0 {
|
||||||
|
return nil, packageError("WFPKG_FILE_REQUIRED", "必须上传工作流包文件")
|
||||||
|
}
|
||||||
|
if len(archive) > MaxArchiveBytes {
|
||||||
|
return nil, packageError("WFPKG_FILE_TOO_LARGE", "工作流包文件超过大小限制")
|
||||||
|
}
|
||||||
|
zr, err := zip.NewReader(bytes.NewReader(archive), int64(len(archive)))
|
||||||
|
if err != nil {
|
||||||
|
return nil, packageError("WFPKG_INVALID_ARCHIVE", "工作流包不是有效 ZIP 文件")
|
||||||
|
}
|
||||||
|
entries := make(map[string][]byte, len(zr.File))
|
||||||
|
var extracted int64
|
||||||
|
for _, file := range zr.File {
|
||||||
|
if !safeArchivePath(file.Name) || file.FileInfo().IsDir() || file.FileInfo().Mode()&os.ModeSymlink != 0 {
|
||||||
|
return nil, packageError("WFPKG_INVALID_ARCHIVE", "工作流包包含不安全文件路径")
|
||||||
|
}
|
||||||
|
if _, exists := entries[file.Name]; exists {
|
||||||
|
return nil, packageError("WFPKG_INVALID_ARCHIVE", "工作流包包含重复文件")
|
||||||
|
}
|
||||||
|
if file.UncompressedSize64 > MaxExtractedBytes || extracted+int64(file.UncompressedSize64) > MaxExtractedBytes {
|
||||||
|
return nil, packageError("WFPKG_INVALID_ARCHIVE", "工作流包解压后超过大小限制")
|
||||||
|
}
|
||||||
|
reader, err := file.Open()
|
||||||
|
if err != nil {
|
||||||
|
return nil, packageError("WFPKG_INVALID_ARCHIVE", "无法读取工作流包文件")
|
||||||
|
}
|
||||||
|
data, readErr := io.ReadAll(io.LimitReader(reader, int64(MaxExtractedBytes)-extracted+1))
|
||||||
|
closeErr := reader.Close()
|
||||||
|
if readErr != nil || closeErr != nil || len(data) > MaxExtractedBytes-int(extracted) {
|
||||||
|
return nil, packageError("WFPKG_INVALID_ARCHIVE", "工作流包解压后超过大小限制")
|
||||||
|
}
|
||||||
|
extracted += int64(len(data))
|
||||||
|
entries[file.Name] = data
|
||||||
|
}
|
||||||
|
|
||||||
|
manifestRaw, hasManifest := entries["manifest.json"]
|
||||||
|
checksumsRaw, hasChecksums := entries["checksums.sha256"]
|
||||||
|
if !hasManifest || !hasChecksums {
|
||||||
|
return nil, packageError("WFPKG_UNSUPPORTED_FORMAT", "工作流包缺少必需文件")
|
||||||
|
}
|
||||||
|
manifest, err := parseManifest(manifestRaw)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(manifest.Items) != 1 || manifest.Items[0].Type != "workflow" {
|
||||||
|
return nil, packageError("WFPKG_MULTIPLE_WORKFLOWS", "工作流包必须且只能包含一个工作流")
|
||||||
|
}
|
||||||
|
item := manifest.Items[0]
|
||||||
|
workflowRaw, exists := entries[item.Path]
|
||||||
|
if !exists || !safeWorkflowPath(item.Path) || len(entries) != 3 {
|
||||||
|
return nil, packageError("WFPKG_INVALID_ARCHIVE", "工作流包包含未声明文件")
|
||||||
|
}
|
||||||
|
checksums, err := parseChecksums(checksumsRaw)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if len(checksums) != 2 || checksums["manifest.json"] != sha256Prefixed(manifestRaw) || checksums[item.Path] != sha256Prefixed(workflowRaw) {
|
||||||
|
return nil, packageError("WFPKG_CHECKSUM_MISMATCH", "工作流包校验和不匹配")
|
||||||
|
}
|
||||||
|
if item.ContentHash != sha256Prefixed(workflowRaw) || !validHash(item.ContentHash) || !validHash(item.GraphHash) {
|
||||||
|
return nil, packageError("WFPKG_CHECKSUM_MISMATCH", "工作流包内容校验和不匹配")
|
||||||
|
}
|
||||||
|
doc, err := parseDocument(workflowRaw)
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
if !safePackageWorkflowID(doc.ID) || doc.ID != item.SourceID || doc.Version != item.SourceRevision {
|
||||||
|
return nil, packageError("WFPKG_INVALID_MANIFEST", "工作流包清单与工作流内容不一致")
|
||||||
|
}
|
||||||
|
graph, err := canonicalJSON([]byte(doc.GraphJSON))
|
||||||
|
if err != nil || item.GraphHash != sha256Prefixed(graph) {
|
||||||
|
return nil, packageError("WFPKG_CHECKSUM_MISMATCH", "工作流图校验和不匹配")
|
||||||
|
}
|
||||||
|
if validateGraph == nil || validateGraph(ctx, string(graph)) != nil {
|
||||||
|
return nil, packageError("WFPKG_WORKFLOW_INVALID", "工作流图校验失败")
|
||||||
|
}
|
||||||
|
var graphShape struct {
|
||||||
|
Nodes []json.RawMessage `json:"nodes"`
|
||||||
|
Edges []json.RawMessage `json:"edges"`
|
||||||
|
}
|
||||||
|
if err := json.Unmarshal(graph, &graphShape); err != nil {
|
||||||
|
return nil, packageError("WFPKG_WORKFLOW_INVALID", "工作流图不是有效 JSON")
|
||||||
|
}
|
||||||
|
return &InspectionResult{
|
||||||
|
PackageHash: sha256Prefixed(archive),
|
||||||
|
Manifest: manifest,
|
||||||
|
Document: doc,
|
||||||
|
ContentHash: item.ContentHash,
|
||||||
|
GraphHash: item.GraphHash,
|
||||||
|
NodeCount: len(graphShape.Nodes),
|
||||||
|
EdgeCount: len(graphShape.Edges),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func safeArchivePath(name string) bool {
|
||||||
|
return name != "" && !strings.Contains(name, `\`) && !strings.HasPrefix(name, "/") && path.Clean(name) == name && !strings.HasPrefix(name, "../") && name != ".."
|
||||||
|
}
|
||||||
|
|
||||||
|
func safeWorkflowPath(name string) bool {
|
||||||
|
rest := strings.TrimPrefix(name, "workflows/")
|
||||||
|
return safeArchivePath(name) && strings.HasPrefix(name, "workflows/") && rest != "" && !strings.Contains(rest, "/") && strings.HasSuffix(rest, ".json")
|
||||||
|
}
|
||||||
|
|
||||||
|
func safePackageWorkflowID(id string) bool {
|
||||||
|
if id == "" || strings.ContainsAny(id, `/\`) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
for _, r := range id {
|
||||||
|
if unicode.IsControl(r) {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseManifest(raw []byte) (Manifest, error) {
|
||||||
|
var manifest Manifest
|
||||||
|
dec := json.NewDecoder(bytes.NewReader(raw))
|
||||||
|
dec.DisallowUnknownFields()
|
||||||
|
if err := dec.Decode(&manifest); err != nil {
|
||||||
|
return Manifest{}, packageError("WFPKG_INVALID_MANIFEST", "工作流包清单格式无效")
|
||||||
|
}
|
||||||
|
if err := consumeJSONEnd(dec); err != nil || manifest.PackageFormat != PackageFormat || manifest.FormatVersion != FormatVersion || strings.TrimSpace(manifest.PackageID) == "" || len(manifest.Items) == 0 {
|
||||||
|
return Manifest{}, packageError("WFPKG_INVALID_MANIFEST", "工作流包清单格式不受支持")
|
||||||
|
}
|
||||||
|
return manifest, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseDocument(raw []byte) (Document, error) {
|
||||||
|
var doc Document
|
||||||
|
dec := json.NewDecoder(bytes.NewReader(raw))
|
||||||
|
dec.DisallowUnknownFields()
|
||||||
|
if err := dec.Decode(&doc); err != nil || consumeJSONEnd(dec) != nil {
|
||||||
|
return Document{}, packageError("WFPKG_WORKFLOW_INVALID", "工作流定义格式无效")
|
||||||
|
}
|
||||||
|
doc.ID = strings.TrimSpace(doc.ID)
|
||||||
|
doc.Name = strings.TrimSpace(doc.Name)
|
||||||
|
if doc.ID == "" || doc.Name == "" || doc.Version <= 0 || strings.TrimSpace(doc.GraphJSON) == "" {
|
||||||
|
return Document{}, packageError("WFPKG_WORKFLOW_INVALID", "工作流定义缺少必需字段")
|
||||||
|
}
|
||||||
|
return doc, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func consumeJSONEnd(dec *json.Decoder) error {
|
||||||
|
var extra any
|
||||||
|
if err := dec.Decode(&extra); err != io.EOF {
|
||||||
|
if err == nil {
|
||||||
|
return fmt.Errorf("multiple JSON values")
|
||||||
|
}
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func parseChecksums(raw []byte) (map[string]string, error) {
|
||||||
|
entries := make(map[string]string)
|
||||||
|
for _, line := range strings.Split(strings.TrimSpace(string(raw)), "\n") {
|
||||||
|
parts := strings.SplitN(strings.TrimSpace(line), " ", 2)
|
||||||
|
if len(parts) != 2 || !validHash("sha256:"+parts[0]) || !safeArchivePath(parts[1]) {
|
||||||
|
return nil, packageError("WFPKG_CHECKSUM_MISMATCH", "工作流包校验和格式无效")
|
||||||
|
}
|
||||||
|
if _, exists := entries[parts[1]]; exists {
|
||||||
|
return nil, packageError("WFPKG_CHECKSUM_MISMATCH", "工作流包校验和重复")
|
||||||
|
}
|
||||||
|
entries[parts[1]] = "sha256:" + parts[0]
|
||||||
|
}
|
||||||
|
return entries, nil
|
||||||
|
}
|
||||||
@@ -0,0 +1,154 @@
|
|||||||
|
package workflowpackage
|
||||||
|
|
||||||
|
import (
|
||||||
|
"archive/zip"
|
||||||
|
"bytes"
|
||||||
|
"context"
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"io"
|
||||||
|
"os"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestInspectArchiveAcceptsSingleVerifiedWorkflow(t *testing.T) {
|
||||||
|
pkg, meta, err := Export(testDocument())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
result, err := InspectArchive(context.Background(), pkg, func(context.Context, string) error { return nil })
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("InspectArchive: %v", err)
|
||||||
|
}
|
||||||
|
if result.PackageHash != meta.PackageHash || result.Document.ID != "web-src-hunting" {
|
||||||
|
t.Fatalf("unexpected inspection: %#v", result)
|
||||||
|
}
|
||||||
|
if result.NodeCount != 2 || result.EdgeCount != 1 {
|
||||||
|
t.Fatalf("counts = %d/%d, want 2/1", result.NodeCount, result.EdgeCount)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInspectArchiveRejectsUnsafeArchiveShapes(t *testing.T) {
|
||||||
|
valid, _, err := Export(testDocument())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
cases := []struct {
|
||||||
|
name string
|
||||||
|
archive []byte
|
||||||
|
}{
|
||||||
|
{name: "duplicate entry", archive: appendZipEntry(t, valid, "manifest.json", []byte(`{}`), 0)},
|
||||||
|
{name: "path traversal", archive: appendZipEntry(t, valid, "../payload.json", []byte(`{}`), 0)},
|
||||||
|
{name: "symlink", archive: appendZipEntry(t, valid, "workflows/link.json", []byte("target"), 0o120777)},
|
||||||
|
{name: "undeclared file", archive: appendZipEntry(t, valid, "notes.txt", []byte("not allowed"), 0)},
|
||||||
|
}
|
||||||
|
for _, tc := range cases {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
_, err := InspectArchive(context.Background(), tc.archive, func(context.Context, string) error { return nil })
|
||||||
|
if ErrorCode(err) != "WFPKG_INVALID_ARCHIVE" {
|
||||||
|
t.Fatalf("code = %q, err = %v", ErrorCode(err), err)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestInspectArchiveRejectsChecksumMismatchAndInvalidWorkflow(t *testing.T) {
|
||||||
|
pkg, _, err := Export(testDocument())
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
badChecksum := replaceZipEntry(t, pkg, "checksums.sha256", []byte("00 manifest.json\n"), 0)
|
||||||
|
if _, err := InspectArchive(context.Background(), badChecksum, func(context.Context, string) error { return nil }); ErrorCode(err) != "WFPKG_CHECKSUM_MISMATCH" {
|
||||||
|
t.Fatalf("checksum code = %q, err = %v", ErrorCode(err), err)
|
||||||
|
}
|
||||||
|
if _, err := InspectArchive(context.Background(), pkg, func(context.Context, string) error { return errors.New("invalid graph") }); ErrorCode(err) != "WFPKG_WORKFLOW_INVALID" {
|
||||||
|
t.Fatalf("graph code = %q, err = %v", ErrorCode(err), err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func appendZipEntry(t *testing.T, archive []byte, name string, data []byte, mode os.FileMode) []byte {
|
||||||
|
t.Helper()
|
||||||
|
return rewriteZip(t, archive, func(zw *zip.Writer) error {
|
||||||
|
h := &zip.FileHeader{Name: name, Method: zip.Store}
|
||||||
|
h.SetMode(mode)
|
||||||
|
w, err := zw.CreateHeader(h)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
_, err = w.Write(data)
|
||||||
|
return err
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
func replaceZipEntry(t *testing.T, archive []byte, name string, data []byte, mode os.FileMode) []byte {
|
||||||
|
t.Helper()
|
||||||
|
zr, err := zip.NewReader(bytes.NewReader(archive), int64(len(archive)))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var out bytes.Buffer
|
||||||
|
zw := zip.NewWriter(&out)
|
||||||
|
for _, f := range zr.File {
|
||||||
|
if f.Name == name {
|
||||||
|
h := &zip.FileHeader{Name: name, Method: zip.Store}
|
||||||
|
h.SetMode(mode)
|
||||||
|
w, err := zw.CreateHeader(h)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if _, err := w.Write(data); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
r, err := f.Open()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
h := &zip.FileHeader{Name: f.Name, Method: zip.Store}
|
||||||
|
w, err := zw.CreateHeader(h)
|
||||||
|
if err == nil {
|
||||||
|
_, err = io.Copy(w, r)
|
||||||
|
}
|
||||||
|
_ = r.Close()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := zw.Close(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return out.Bytes()
|
||||||
|
}
|
||||||
|
|
||||||
|
func rewriteZip(t *testing.T, archive []byte, appendEntry func(*zip.Writer) error) []byte {
|
||||||
|
t.Helper()
|
||||||
|
zr, err := zip.NewReader(bytes.NewReader(archive), int64(len(archive)))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
var out bytes.Buffer
|
||||||
|
zw := zip.NewWriter(&out)
|
||||||
|
for _, f := range zr.File {
|
||||||
|
r, err := f.Open()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
h := &zip.FileHeader{Name: f.Name, Method: zip.Store}
|
||||||
|
w, err := zw.CreateHeader(h)
|
||||||
|
if err == nil {
|
||||||
|
_, err = io.Copy(w, r)
|
||||||
|
}
|
||||||
|
_ = r.Close()
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if err := appendEntry(zw); err != nil {
|
||||||
|
t.Fatal(fmt.Errorf("append entry: %w", err))
|
||||||
|
}
|
||||||
|
if err := zw.Close(); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return out.Bytes()
|
||||||
|
}
|
||||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user