Compare commits

...
133 Commits
Author SHA1 Message Date
公明andGitHub a496bd7aee Add files via upload 2026-07-14 15:00:05 +08:00
公明andGitHub e98da9d1b3 Update config.example.yaml 2026-07-14 14:59:52 +08:00
公明andGitHub a1a62657af Add files via upload 2026-07-14 14:58:20 +08:00
公明andGitHub 7a3e74a3af Add files via upload 2026-07-14 14:56:57 +08:00
公明andGitHub a36acddcf4 Add files via upload 2026-07-14 14:51:56 +08:00
公明andGitHub 4016eab95f Add files via upload 2026-07-14 14:49:37 +08:00
公明andGitHub 51efdd7b60 Add files via upload 2026-07-14 14:46:49 +08:00
公明andGitHub c24fbf602c Add files via upload 2026-07-14 14:45:09 +08:00
公明andGitHub 4b1f71b4a6 Add files via upload 2026-07-14 14:43:13 +08:00
公明andGitHub 29bf0ff4da Add files via upload 2026-07-14 14:41:30 +08:00
公明andGitHub ccfb85e6cb Add files via upload 2026-07-14 14:39:32 +08:00
公明andGitHub efa8b262b1 Add files via upload 2026-07-14 14:36:57 +08:00
公明andGitHub bbbe77e90c Add files via upload 2026-07-14 14:34:48 +08:00
公明andGitHub 6bafb8fe70 Add files via upload 2026-07-14 14:22:21 +08:00
公明andGitHub 7d1e16b97b Add files via upload 2026-07-14 11:52:27 +08:00
公明andGitHub 1be10cdf2c Add files via upload 2026-07-14 11:51:35 +08:00
公明andGitHub 95909ce999 Add files via upload 2026-07-14 11:49:40 +08:00
公明andGitHub a603adb467 Add files via upload 2026-07-14 11:48:09 +08:00
acfacfe1b3 feat(工作流):支持本地图编排策略包导入导出 (#195)
* docs: define local workflow package mvp

* docs: fix workflow package api contract

* feat: add local workflow package mvp

* feat(workflow): add package API client

* feat(workflow): add package import interface

* feat(workflow): connect package import flow

* fix(workflow): map package validation errors

* fix(workflow): handle import key generation errors

* feat(workflow): localize package import states

* fix(workflow): reset package import modal on open

---------

Co-authored-by: ruanmingchen <“ruanm@chenchen”>
2026-07-14 11:28:37 +08:00
公明andGitHub 2e5c1ff286 Add files via upload 2026-07-13 19:54:50 +08:00
公明andGitHub 4ccb330ded Add files via upload 2026-07-13 19:24:48 +08:00
公明andGitHub a28dad4827 Add files via upload 2026-07-13 19:24:11 +08:00
公明andGitHub a6b7c7e7be Add files via upload 2026-07-13 19:16:51 +08:00
公明andGitHub e5a4703e91 Add files via upload 2026-07-13 19:08:57 +08:00
公明andGitHub 8ea0ee4f1b Add files via upload 2026-07-13 19:07:19 +08:00
公明andGitHub 386ec7b835 Add files via upload 2026-07-13 19:06:02 +08:00
公明andGitHub 078f5cb222 Add files via upload 2026-07-13 19:02:42 +08:00
公明andGitHub 7351793683 Add files via upload 2026-07-13 18:59:47 +08:00
公明andGitHub 68e3ead4d7 Add files via upload 2026-07-13 18:56:52 +08:00
公明andGitHub 597712a7b9 Add files via upload 2026-07-13 18:55:06 +08:00
公明andGitHub f77cc09477 Add files via upload 2026-07-13 17:25:30 +08:00
公明andGitHub 5423bd5e1e Add files via upload 2026-07-13 17:22:31 +08:00
公明andGitHub a64c18df6d Add files via upload 2026-07-13 17:12:48 +08:00
公明andGitHub 33b0b56514 Add files via upload 2026-07-13 17:04:21 +08:00
公明andGitHub 23ec222b77 Add files via upload 2026-07-13 16:53:08 +08:00
公明andGitHub d36352e4dc Add files via upload 2026-07-13 16:19:35 +08:00
公明andGitHub ea22b1f3ba Add files via upload 2026-07-13 16:17:12 +08:00
公明andGitHub 8a3229aa5a Add files via upload 2026-07-13 16:10:07 +08:00
公明andGitHub a497c4bfcd Add files via upload 2026-07-13 16:03:40 +08:00
公明andGitHub bd9116e9c3 Add files via upload 2026-07-13 16:01:14 +08:00
公明andGitHub 6f39720669 Add files via upload 2026-07-13 15:57:17 +08:00
公明andGitHub f8481024ed Add files via upload 2026-07-13 15:52:25 +08:00
公明andGitHub 6b3a7d81d2 Add files via upload 2026-07-13 15:47:13 +08:00
公明andGitHub 25cf3c567b Add files via upload 2026-07-13 15:43:04 +08:00
公明andGitHub 41f683ce6c Update config.example.yaml 2026-07-12 13:27:31 +08:00
公明andGitHub ffae94fb2c Add files via upload 2026-07-12 13:21:20 +08:00
公明andGitHub fe3c845ff8 Add files via upload 2026-07-12 13:20:08 +08:00
公明andGitHub 0f1a6ad25a Add files via upload 2026-07-12 13:17:37 +08:00
公明andGitHub 3b3f73461b Add files via upload 2026-07-12 13:15:20 +08:00
公明andGitHub 3ce80fd00f Add files via upload 2026-07-12 13:14:23 +08:00
公明andGitHub 08cb8a68fb Add files via upload 2026-07-12 13:13:00 +08:00
公明andGitHub 2f38693891 Add files via upload 2026-07-12 13:10:56 +08:00
公明andGitHub 6fa17a3093 Add files via upload 2026-07-12 13:09:12 +08:00
公明andGitHub 46cea9459a Add files via upload 2026-07-12 13:07:20 +08:00
公明andGitHub 987bd0a03c Add files via upload 2026-07-11 12:10:57 +08:00
公明andGitHub 6529c13ffb Update config.example.yaml 2026-07-11 11:53:24 +08:00
公明andGitHub f2ba5093d7 Add files via upload 2026-07-11 11:47:26 +08:00
公明andGitHub 3c6cf633e1 Add files via upload 2026-07-11 11:44:25 +08:00
公明andGitHub 142977413e Add files via upload 2026-07-11 11:42:25 +08:00
公明andGitHub 62efc81993 Add files via upload 2026-07-11 11:41:00 +08:00
公明andGitHub 9ac5fd33ec Add files via upload 2026-07-11 11:40:09 +08:00
公明andGitHub d1d67b07d3 Add files via upload 2026-07-11 11:38:14 +08:00
公明andGitHub 211c36654a Add files via upload 2026-07-11 11:36:47 +08:00
公明andGitHub 3894ba6054 Add files via upload 2026-07-11 11:34:41 +08:00
公明andGitHub fa0dd6c721 Add files via upload 2026-07-11 11:32:44 +08:00
公明andGitHub 1cf10981ef Add files via upload 2026-07-11 11:30:50 +08:00
公明andGitHub 79d162ccb4 Add files via upload 2026-07-11 11:21:15 +08:00
公明andGitHub 722806797f Add files via upload 2026-07-11 10:47:31 +08:00
公明andGitHub 2fcda3a57d Add files via upload 2026-07-11 10:47:04 +08:00
公明andGitHub bec296ae3f Delete internal/multiagent directory 2026-07-11 10:45:21 +08:00
612015b6d7 修复工作流检查点序列化 (#193)
Co-authored-by: ruanmingchen <“ruanm@chenchen”>
2026-07-11 01:38:10 +08:00
公明andGitHub ca0bcc21d4 Update config.example.yaml 2026-07-10 21:37:32 +08:00
公明andGitHub c33e5f2026 Add files via upload 2026-07-10 21:35:49 +08:00
公明andGitHub 0251230654 Add files via upload 2026-07-10 21:31:23 +08:00
公明andGitHub b76e06ff92 Add files via upload 2026-07-10 21:29:13 +08:00
公明andGitHub aca97ffc94 Add files via upload 2026-07-10 21:27:03 +08:00
公明andGitHub 24052717bd Add files via upload 2026-07-10 21:25:33 +08:00
公明andGitHub 478b52b011 Add files via upload 2026-07-10 21:23:33 +08:00
公明andGitHub 6f8e324a75 Add files via upload 2026-07-10 21:21:56 +08:00
公明andGitHub 1abfd3d22a Add files via upload 2026-07-10 21:19:55 +08:00
公明andGitHub e87011b081 Add files via upload 2026-07-10 20:20:05 +08:00
公明andGitHub e020ffed49 Add files via upload 2026-07-10 19:47:53 +08:00
公明andGitHub 823fb47a81 Add files via upload 2026-07-10 19:46:14 +08:00
公明andGitHub 46a9b42fde Add files via upload 2026-07-10 18:58:17 +08:00
公明andGitHub 2995847e0e Add files via upload 2026-07-10 18:56:33 +08:00
公明andGitHub 1b64f5d8a0 Add files via upload 2026-07-10 18:54:04 +08:00
公明andGitHub a145687508 Add files via upload 2026-07-10 18:50:57 +08:00
公明andGitHub 49d2175872 Add files via upload 2026-07-10 18:49:42 +08:00
公明andGitHub cd87a0b965 Add files via upload 2026-07-10 18:47:30 +08:00
公明andGitHub 3bfa846db1 Add files via upload 2026-07-10 18:45:07 +08:00
公明andGitHub bc3246d157 Add files via upload 2026-07-10 16:51:06 +08:00
公明andGitHub 542d1d2411 Add files via upload 2026-07-10 16:48:49 +08:00
公明andGitHub 0ff8c58fbd Add files via upload 2026-07-10 16:46:26 +08:00
公明andGitHub 25a76a8c97 Add files via upload 2026-07-10 16:44:09 +08:00
公明andGitHub a2744d6936 Add files via upload 2026-07-10 16:42:35 +08:00
公明andGitHub 225d9ba0d9 Add files via upload 2026-07-10 16:40:21 +08:00
公明andGitHub 9dfb9d6c68 Add files via upload 2026-07-10 16:39:13 +08:00
公明andGitHub 04efc50161 Add files via upload 2026-07-10 16:37:14 +08:00
公明andGitHub 46f891a114 Add files via upload 2026-07-09 19:31:32 +08:00
公明andGitHub bd80010cda Add files via upload 2026-07-09 19:24:20 +08:00
公明andGitHub a1201240d4 Add files via upload 2026-07-09 19:23:15 +08:00
公明andGitHub 5266e2d95f Delete config.yaml 2026-07-09 19:21:21 +08:00
公明andGitHub 1af455b762 Add files via upload 2026-07-09 19:20:38 +08:00
公明andGitHub 8da02f94ac Rename gitignore to .gitignore 2026-07-09 19:19:10 +08:00
公明andGitHub 7f1bc1d229 Add files via upload 2026-07-09 19:18:21 +08:00
公明andGitHub 0e61c17db3 Add files via upload 2026-07-09 17:25:07 +08:00
公明andGitHub 2f1a95e8bf Add files via upload 2026-07-09 16:52:37 +08:00
公明andGitHub 878cf158b0 Add files via upload 2026-07-09 16:47:00 +08:00
公明andGitHub e092f6c590 Add files via upload 2026-07-09 16:10:46 +08:00
公明andGitHub deafabc276 Add files via upload 2026-07-09 16:08:11 +08:00
公明andGitHub 87c83c92fd Add files via upload 2026-07-09 16:04:06 +08:00
公明andGitHub 898f69a388 Add files via upload 2026-07-09 16:01:59 +08:00
公明andGitHub 0fa739d438 Add files via upload 2026-07-09 16:00:23 +08:00
公明andGitHub a7c3d28cbc Add files via upload 2026-07-09 15:57:54 +08:00
公明andGitHub 1c576db75d Add files via upload 2026-07-09 15:55:55 +08:00
公明andGitHub 56ee11a97a Add files via upload 2026-07-09 14:46:50 +08:00
公明andGitHub 95da42effb Add files via upload 2026-07-09 14:41:37 +08:00
公明andGitHub 64c3f28780 Add files via upload 2026-07-09 14:38:09 +08:00
公明andGitHub c8fece70ef Add files via upload 2026-07-09 14:35:58 +08:00
公明andGitHub 46ecb2f13b Add files via upload 2026-07-09 14:34:23 +08:00
公明andGitHub 6484950be3 Add files via upload 2026-07-09 14:32:36 +08:00
公明andGitHub f68d69d562 Add files via upload 2026-07-09 14:30:44 +08:00
公明andGitHub 798a76ec3c Add files via upload 2026-07-09 11:44:54 +08:00
公明andGitHub 37e553ba8a Add files via upload 2026-07-09 11:43:27 +08:00
公明andGitHub 16a854d0f5 Add files via upload 2026-07-09 11:41:24 +08:00
公明andGitHub c57b681a19 Add files via upload 2026-07-09 11:38:57 +08:00
公明andGitHub f405e02d70 Add files via upload 2026-07-09 11:37:16 +08:00
公明andGitHub 0531ee7292 Add files via upload 2026-07-09 11:33:15 +08:00
公明andGitHub 69e8d7020c Add files via upload 2026-07-09 11:26:31 +08:00
公明andGitHub 526e8626e9 Merge pull request #187 from chaojixinren/fix-knowledge-handler-method-name
fix: 修复 KnowledgeHandler 方法调用错误
2026-07-09 11:12:15 +08:00
公明andGitHub cb52ff37f5 Add files via upload 2026-07-08 16:39:12 +08:00
chaojixinren 461682acaf fix: 修复 KnowledgeHandler 方法调用错误
- 将 app.go 中的 RebuildIndex 调用改为 StartIndex
- 该方法在最近提交中已重命名,但调用点未同步更新
- 修复编译错误: app.knowledgeHandler.RebuildIndex undefined
2026-07-08 03:28:54 -04:00
公明andGitHub 7522bab98c Add files via upload 2026-07-08 15:28:38 +08:00
215 changed files with 31249 additions and 2875 deletions
+44
View File
@@ -0,0 +1,44 @@
# Runtime data
data/
*.db
*.db-shm
*.db-wal
*.sqlite
*.sqlite3
# Local configuration and secrets
config.yaml
config.local.yaml
.env
.env.*
*.pem
*.key
*.crt
# Build outputs
cyberstrike-ai
bin/
dist/
build/
coverage.out
coverage.html
# Logs and temporary files
*.log
tmp/
temp/
# Python
venv/
.venv/
__pycache__/
*.py[cod]
.pytest_cache/
# Go
vendor/
# macOS / editors
.DS_Store
.idea/
.vscode/
+103 -431
View File
@@ -7,31 +7,14 @@
[中文](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.
<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>
> [!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.
## 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.*
<details>
<summary><strong>More interface screenshots</strong></summary>
### Core Features Overview
<table>
@@ -115,32 +101,48 @@ If CyberStrikeAI helps you, you can support the project via **WeChat Pay** or **
</tr>
</table>
</details>
</div>
## Highlights
- 🤖 Agentic execution layer for translating natural-language intent into precise, governed, auditable security action
- 🧩 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
- 🧰 100+ curated security tool recipes, YAML-based extensions, and role-scoped tool control
- 📄 Large-result pagination, compression, and searchable archives
- 🔗 Attack-chain intelligence with graph views, risk scoring, project facts, and step-by-step replay
- 🧑‍⚖️ 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
- 📚 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
- 📁 Conversation grouping with pinning, rename, and batch management
- 📂 **Project management**: shared facts (blackboard) across sessions, `upsert_project_fact` + `links` to chain paths; attack-chain and project fact graph views
- 🛡️ Vulnerability management with CRUD operations, severity tracking, status workflow, and statistics
- 📋 Batch task management: create task queues, add multiple tasks, and execute them sequentially
- 🎭 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)
- 🧩 **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)
- 🎯 **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/`
- 📱 **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))
- 🧑‍⚖️ **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)
- 🐚 **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.**
### Agents and orchestration
- 🤖 **Agentic execution** translates natural-language intent into governed, auditable security actions.
- 🧩 **Eino orchestration** supports single-agent execution plus Deep, Plan-Execute, and Supervisor multi-agent modes.
- 🔀 **Graph workflows** combine Agents, tools, conditions, approvals, and outputs into reusable flows.
- 🎭 **Role-based testing** provides focused prompts and tool policies for common security scenarios.
### Tools and knowledge
- 🧰 **Security tools** include 100+ curated YAML recipes with custom extensions and role-scoped access.
- 🔌 **MCP integration** supports HTTP, stdio, SSE, external federation, and dynamic tool discovery.
- 🎯 **Agent Skills** follow the standard Skill layout and support progressive, on-demand loading.
- 📚 **Knowledge base** combines query rewriting, vector retrieval, reranking, and result post-processing.
- 🖼️ **Vision analysis** uses a separate vision model for screenshots, captchas, and UI while retaining text summaries only.
### Governance and audit
- 🧑‍⚖️ **Human in the loop** provides approval modes, tool allowlists, audit-agent review, and traceable decisions.
- 🔐 **Platform RBAC** supports multiple users, system and custom roles, scoped permissions, ownership, and explicit assignments.
- 🔒 **Security and audit** provide authenticated access, audit logs, SQLite persistence, and operational evidence retention.
- 📄 **Result governance** supports pagination, compression, archival, and search for large tool outputs.
### 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
@@ -149,11 +151,19 @@ CyberStrikeAI includes optional integrations under `plugins/`.
- **Burp Suite extension**: `plugins/burp-suite/cyberstrikeai-burp-extension/`
Build output: `plugins/burp-suite/cyberstrikeai-burp-extension/dist/cyberstrikeai-burp-extension.jar`
Docs: `plugins/burp-suite/cyberstrikeai-burp-extension/README.md`
- **Browser extension (Chrome / Edge)**: `plugins/browser-extension/cyberstrikeai-browser-extension/`
Capture Network traffic in DevTools and send it to CyberStrikeAI for AI-assisted security testing—aligned with the Burp plugin.
Install: `chrome://extensions/` → Load unpacked → F12 → **CyberStrikeAI** tab
Package output: `plugins/browser-extension/cyberstrikeai-browser-extension/dist/cyberstrikeai-browser-extension.zip`
Docs: `plugins/browser-extension/cyberstrikeai-browser-extension/README.md` / `README.zh-CN.md`
## Tool Overview
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
- **Web & App Scanners** sqlmap, nikto, dirb, gobuster, feroxbuster, ffuf, httpx
- **Vulnerability Scanners** nuclei, wpscan, wafw00f, dalfox, xsser
@@ -170,12 +180,16 @@ CyberStrikeAI ships with 100+ curated tools covering the whole kill chain:
- **CTF Utilities** stegsolve, zsteg, hash-identifier, fcrackzip, pdfcrack, cyberchef
- **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
### Quick Start (One-Command Deployment)
**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/))
**One-Command Deployment:**
@@ -193,6 +207,12 @@ The `run.sh` script will automatically:
- ✅ Build the project
- ✅ 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.
**First-Time Configuration:**
@@ -201,12 +221,12 @@ The `run.sh` script will automatically:
- Go to `Settings` → Fill in your API credentials:
```yaml
openai:
api_key: "sk-your-key"
api_key: "${OPENAI_API_KEY}"
base_url: "https://api.openai.com/v1" # or https://api.deepseek.com/v1
model: "gpt-4o" # or deepseek-chat, claude-3-opus, etc.
```
- 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:
**macOS (Homebrew):**
@@ -237,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.
### 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`
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.
@@ -254,401 +274,32 @@ Requirements / tips:
* `rsync` is recommended/required for the safe code sync.
* 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.
**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 Einos 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 projects `.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.
⚠️ **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.
### Automation Hooks
- **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
## 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
auth:
password: "change-me"
session_duration_hours: 12
server:
host: "0.0.0.0"
host: "127.0.0.1"
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:
api_key: "sk-xxx"
base_url: "https://api.deepseek.com/v1"
model: "deepseek-chat"
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
api_key: "${OPENAI_API_KEY}"
base_url: "https://api.openai.com/v1"
model: "your-model"
```
### Tool Definition Example (`tools/nmap.yaml`)
```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
```
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.
## Related documentation
- [Documentation index](docs/README.md): deployment, configuration, security model, API, knowledge base, C2, WebShell, MCP, development, testing, and troubleshooting.
- [Deployment guide](docs/en-US/deployment.md): source/binary startup, HTTPS, reverse proxy, systemd, backup, upgrade, and rollback.
- [Runbooks](docs/en-US/runbooks.md): production setup, external MCP, knowledge base, authorized Web testing, and C2 cleanup workflows.
- [Security hardening](docs/en-US/security-hardening.md): launch baseline, HITL allowlist, reverse proxy, file permissions, and periodic review.
- [API recipes](docs/en-US/api-recipes.md): examples for login, Agent, streaming, multi-agent, uploads, vulnerabilities, KB, and audit export.
- [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.
- [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): Setup, commands, and troubleshooting for WeChat, WeCom, DingTalk, Lark, Telegram, Slack, Discord, and QQ Bot.
- [HITL best practices](docs/en-US/hitl-best-practices.md): reviewer modes, allowlists, Audit Agent prompts, and separate small-model configuration.
- **New users:** [Deployment](docs/en-US/deployment.md) → [Configuration](docs/en-US/configuration.md) → [Troubleshooting](docs/en-US/troubleshooting.md)
- **Operators:** [Configuration profiles](docs/en-US/configuration-profiles.md) → [Security hardening](docs/en-US/security-hardening.md) → [Runbooks](docs/en-US/runbooks.md)
- **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)
- **Contributors:** [Developer guide](docs/en-US/developer-guide.md) → [Testing](docs/en-US/testing.md) → [Contributing](docs/en-US/contributing-guide.md)
- **All topics:** [English documentation](docs/en-US/README.md) · [Bilingual documentation index](docs/README.md)
## Project Layout
@@ -663,6 +314,7 @@ CyberStrikeAI/
├── agents/ # Multi-agent Markdown (orchestrator.md + sub-agent *.md)
├── docs/ # Topic docs (deployment, config, security, API, knowledge base, C2, WebShell, etc.)
├── images/ # Docs screenshots & diagrams
├── scripts/ # Repository maintenance checks, including documentation validation
├── config.yaml # Runtime configuration
├── run.sh # Convenience launcher
└── README*.md
@@ -704,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
CyberStrikeAI is licensed under the Apache License 2.0.
+102 -430
View File
@@ -6,31 +6,14 @@
[中文](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 框架。
<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>
> [!IMPORTANT]
> 仅可对自有系统或已获得明确授权的目标使用 CyberStrikeAI。在共享或生产环境启用高风险工具、WebShell 或 C2 前,请先阅读[安全模型](docs/zh-CN/security-model.md)和[安全加固指南](docs/zh-CN/security-hardening.md)。
## 界面与集成预览
@@ -53,6 +36,9 @@ CyberStrikeAI 基于 Go 构建,为 AI 原生安全运营提供完整底座:1
*仪表盘提供系统运行状态、安全漏洞、工具使用情况和知识库的全面概览,帮助用户快速了解平台核心功能和当前状态。*
<details>
<summary><strong>查看更多界面截图</strong></summary>
### 核心功能概览
<table>
@@ -114,32 +100,48 @@ CyberStrikeAI 基于 Go 构建,为 AI 原生安全运营提供完整底座:1
</tr>
</table>
</details>
</div>
## 特性速览
- 🤖 面向智能体时代的执行层,将自然语言意图转化为精准、受控、可审计的安全行动
- 🧩 基于 Eino 的单智能体与多智能体编排,支持 Deep、Plan-Execute、Supervisor 等模式
- 🔌 MCP 原生工具执行,支持 HTTP / stdio / SSE 传输、外部 MCP 联邦与动态工具发现
- 🧰 100+ 精选安全工具配方、YAML 扩展机制与按角色收敛的工具控制
- 📄 大结果分页、压缩与全文检索
- 🔗 攻击链智能分析,支持图谱视图、风险打分、项目事实沉淀与步骤回放
- 🧑‍⚖️ 人机协同治理,支持审批模式、免审批白名单、审计 Agent 复核与可追溯决策
- 🔒 Web 登录保护、审计日志、SQLite 持久化与行动证据留存
- 📚 知识库(RAG):**Eino MultiQuery** 查询改写 + 多路向量检索 + **HTTP 精排**DashScope `gte-rerank` / Cohere 兼容)+ 后处理(去重、预算);索引侧为 **Eino Compose** 流水线
- 📁 对话分组管理:支持分组创建、置顶、重命名、删除等操作
- 📂 **项目管理**:共享事实(黑板)跨会话沉淀认知,`upsert_project_fact` + `links` 串联攻击路径;聊天攻击链与项目事实图可视化
- 🛡️ 漏洞管理功能:完整的漏洞 CRUD 操作,支持严重程度分级、状态流转、按对话/严重程度/状态过滤,以及统计看板
- 📋 批量任务管理:创建任务队列,批量添加任务,依次顺序执行,支持任务编辑与状态跟踪
- 🎭 角色化测试:预设安全测试角色(渗透测试、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)
- 🎯 **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+ 领域示例仍可绑定角色
- 📱 **机器人**:个人微信、企业微信、钉钉、飞书、Telegram、Slack、Discord、QQ 机器人,在手机或 IM 中与 CyberStrikeAI 对话(详见 [机器人使用说明](docs/zh-CN/robot.md)
- 🧑‍⚖️ **人机协同(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 类规则(如命令拒绝正则)。**仅限授权测试。**
### 智能体与编排
- 🤖 **智能体执行层**:将自然语言意图转化为受控、可审计的安全行动。
- 🧩 **Eino 编排**:支持单智能体及 Deep、Plan-Execute、Supervisor 多智能体模式。
- 🔀 **图工作流**:通过 Agent、工具、条件、审批和输出节点构建可复用流程。
- 🎭 **角色化测试**:为常见安全场景提供聚焦的提示词和工具策略。
### 工具与知识扩展
- 🧰 **安全工具**:提供 100+ 精选 YAML 工具配方,支持自定义扩展和按角色控制。
- 🔌 **MCP 集成**:支持 HTTP、stdio、SSE、外部 MCP 联邦和动态工具发现。
- 🎯 **Agent Skills**:遵循标准 Skill 目录结构,支持渐进式按需加载。
- 📚 **知识库**:组合查询改写、向量检索、精排和结果后处理能力。
- 🖼️ **视觉分析**:使用独立视觉模型分析截图、验证码和 UI,对话中仅保留文字摘要。
### 安全治理与审计
- 🧑‍⚖️ **人机协同**:支持审批模式、工具白名单、审计 Agent 复核和决策追踪。
- 🔐 **平台 RBAC**:支持多用户、系统及自定义角色、权限 Scope、资源归属和显式授权。
- 🔒 **安全与审计**:提供登录保护、审计日志、SQLite 持久化和行动证据留存。
- 📄 **结果治理**:支持大结果分页、压缩、归档和检索
### 安全运营管理
- 📁 **对话管理**:支持分组、置顶、重命名和批量管理。
- 📂 **项目与攻击链**:关联跨会话事实、风险评分、图谱视图和步骤回放。
- 🛡️ **漏洞管理**:支持严重程度分级、状态流转、过滤和统计看板。
- 📋 **批量任务**:支持任务队列、编辑、状态跟踪和结果留存。
- 📱 **机器人接入**:支持个人微信、企业微信、钉钉、飞书、Telegram、Slack、Discord 和 QQ。
### 授权安全操作
- 🐚 **WebShell 管理**:提供连接管理、虚拟终端、文件操作和 AI 辅助工作流。
- 📡 **内置 C2**:提供监听器、加密 Beacon、会话、任务队列、Payload 辅助和实时事件。
> WebShell、C2 及其他高风险能力仅限自有系统或已获得明确授权的测试环境。使用前请阅读[安全模型](docs/zh-CN/security-model.md)和[安全加固指南](docs/zh-CN/security-hardening.md)。
## 插件(Plugins
@@ -148,11 +150,19 @@ CyberStrikeAI 基于 Go 构建,为 AI 原生安全运营提供完整底座:1
- **Burp Suite 插件**`plugins/burp-suite/cyberstrikeai-burp-extension/`
构建产物:`plugins/burp-suite/cyberstrikeai-burp-extension/dist/cyberstrikeai-burp-extension.jar`
说明文档:`plugins/burp-suite/cyberstrikeai-burp-extension/README.zh-CN.md`
- **浏览器扩展(Chrome / Edge**`plugins/browser-extension/cyberstrikeai-browser-extension/`
在 DevTools 中捕获 Network 流量并发送到 CyberStrikeAI 进行 AI 辅助安全测试,能力与 Burp 插件对齐。
安装:`chrome://extensions/` → 加载已解压 → F12 → **CyberStrikeAI** 标签页
打包产物:`plugins/browser-extension/cyberstrikeai-browser-extension/dist/cyberstrikeai-browser-extension.zip`
说明文档:`plugins/browser-extension/cyberstrikeai-browser-extension/README.zh-CN.md`
## 工具概览
系统预置 100+ 渗透/攻防工具,覆盖完整攻击链:
<details>
<summary><strong>查看完整工具分类</strong></summary>
- **网络扫描**nmap、masscan、rustscan、arp-scan、nbtscan
- **Web 应用扫描**sqlmap、nikto、dirb、gobuster、feroxbuster、ffuf、httpx
- **漏洞扫描**nuclei、wpscan、wafw00f、dalfox、xsser
@@ -169,12 +179,16 @@ CyberStrikeAI 基于 Go 构建,为 AI 原生安全运营提供完整底座:1
- **CTF 实用工具**stegsolve、zsteg、hash-identifier、fcrackzip、pdfcrack、cyberchef
- **系统辅助**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/))
**一条命令部署:**
@@ -192,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` 写错时程序会在终端提示正确写法。
**首次配置:**
@@ -200,12 +220,12 @@ chmod +x run.sh && ./run.sh
- 进入 `设置` → 填写 API 配置信息:
```yaml
openai:
api_key: "sk-your-key"
api_key: "${OPENAI_API_KEY}"
base_url: "https://api.openai.com/v1" # 或 https://api.deepseek.com/v1
model: "gpt-4o" # 或 deepseek-chat, claude-3-opus 等
```
- 或启动前直接编辑 `config.yaml` 文件
2. **登录系统** - 使用控制台显示自动生成密码(或在 `config.yaml` 中设置 `auth.password`
2. **登录系统** - 首次启动时控制台显示自动生成的 `admin` 初始密码;也可在「平台权限 → 用户管理」中创建账号
3. **安装安全工具(可选)** - 按需安装 `tools/` 目录中的工具;未安装的工具在执行时会自动跳过或改用替代方案。常用示例:
**macOSHomebrew):**
@@ -236,7 +256,7 @@ go build -o cyberstrike-ai cmd/server/main.go
**说明:** Python 虚拟环境(`venv/`)由 `run.sh` 自动创建和管理。需要 Python 的工具(如 `api-fuzzer`、`http-framework-test` 等)会自动使用该环境。
### CyberStrikeAI 版本更新(无兼容性问题)
### 版本升级与兼容性
1. (首次使用)启用脚本:`chmod +x upgrade.sh`
2. 一键升级:`./upgrade.sh`(可选参数:`--tag vX.Y.Z`、`--no-venv`、`--yes`)。本地的 `tools/`、`roles/`、`skills/` 会始终保留不被覆盖。
@@ -252,401 +272,32 @@ go build -o cyberstrike-ai cmd/server/main.go
* 建议/需要 `rsync` 用于安全同步代码。
* 如果遇到 GitHub API 限流,运行前设置 `export GITHUB_TOKEN="..."` 再执行 `./upgrade.sh`。
⚠️ **注意:** 仅适用于无兼容性变更的版本更新。若版本存在兼容性调整,此方法不适用
**举例:** 无兼容性变更如 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. 重启服务或重新加载配置,角色会出现在角色选择下拉菜单中。
### 多代理模式(EinoDeep / 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(与对话共用数据库),服务重启后仍可继续使用。
### 内置 C2Command & 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 联邦**:在设置中注册第三方 MCPHTTP/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 服务器。
⚠️ **升级前必读:** 请查看目标版本的 Release Notes,确认配置、数据库和 API 是否变化。即使只是补丁版本也应先备份,不能仅凭版本号判断兼容性
### 知识库功能
- **向量检索**: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 / 精排 / 预取候选数。
- **检索日志**:记录所有知识检索操作,便于审计与调试。
## 配置
**知识库配置步骤:**
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-rerankCohere→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。
## 配置参考
请以 [`config.example.yaml`](config.example.yaml) 作为权威配置模板,只复制当前环境需要的配置。最少需要配置服务监听地址和一个 OpenAI 兼容模型:
```yaml
auth:
password: "change-me"
session_duration_hours: 12
server:
host: "0.0.0.0"
host: "127.0.0.1"
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:
api_key: "sk-xxx"
base_url: "https://api.deepseek.com/v1"
model: "deepseek-chat"
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: "" # Deeporchestrator.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
api_key: "${OPENAI_API_KEY}"
base_url: "https://api.openai.com/v1"
model: "your-model"
```
### 工具模版示例(`tools/nmap.yaml`
```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
```
不要提交真实凭证。将服务暴露到 localhost 之外前,请阅读[配置参考](docs/zh-CN/configuration.md)、[推荐配置画像](docs/zh-CN/configuration-profiles.md)和[安全加固指南](docs/zh-CN/security-hardening.md)。
## 相关文档
- [文档导航](docs/README.md):部署、配置、安全模型、API、知识库、C2、WebShell、MCP、开发、测试、排错等完整专题入口。
- [部署指南](docs/zh-CN/deployment.md):源码/二进制运行、HTTPS、反向代理、systemd、备份、升级与回滚。
- [运维 Runbooks](docs/zh-CN/runbooks.md):生产部署、外部 MCP、知识库、授权 Web 测试、C2 清理等可执行流程。
- [安全加固指南](docs/zh-CN/security-hardening.md):上线前基线、HITL 白名单、反向代理、文件权限和周期巡检。
- [API Recipes](docs/zh-CN/api-recipes.md):登录、Agent、流式、多代理、上传、漏洞、知识库和审计导出调用示例。
- [配置参考](docs/zh-CN/configuration.md)`config.yaml` 各配置段、推荐值和修改建议。
- [安全模型](docs/zh-CN/security-model.md):认证、工具执行、HITL、审计、C2/WebShell 和数据安全边界。
- [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):个人微信、企业微信、钉钉、飞书、Telegram、Slack、Discord、QQ 机器人的配置、命令与排查。
- [人机协同最佳实践](docs/zh-CN/hitl-best-practices.md):审批方模式、白名单、审计 Agent 提示词策略与独立小模型配置。
- **新用户:** [部署指南](docs/zh-CN/deployment.md) → [配置参考](docs/zh-CN/configuration.md) → [排错指南](docs/zh-CN/troubleshooting.md)
- **运维人员:** [配置画像](docs/zh-CN/configuration-profiles.md) → [安全加固](docs/zh-CN/security-hardening.md) → [运维 Runbooks](docs/zh-CN/runbooks.md)
- **集成开发:** [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/developer-guide.md) → [测试指南](docs/zh-CN/testing.md) → [贡献规范](docs/zh-CN/contributing-guide.md)
- **全部专题:** [中文文档](docs/zh-CN/README.md) · [双语文档索引](docs/README.md)
## 项目结构
@@ -661,6 +312,7 @@ CyberStrikeAI/
├── agents/ # 多代理 Markdownorchestrator.md + 子代理 *.md
├── docs/ # 专题文档(部署、配置、安全、API、知识库、C2、WebShell 等)
├── images/ # 文档配图
├── scripts/ # 仓库维护检查,包括文档校验
├── config.yaml # 运行配置
├── run.sh # 启动脚本
└── README*.md
@@ -700,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** 开源许可。
+16 -9
View File
@@ -5,6 +5,7 @@ import (
"cyberstrike-ai/internal/app"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/logger"
"cyberstrike-ai/internal/termout"
"flag"
"fmt"
"os"
@@ -35,11 +36,20 @@ func main() {
fmt.Fprintf(os.Stderr, "无效的 -config 路径 %q。\n若同时需要 HTTPS,请写成: ./cyberstrike-ai --https -config config.yaml-config 后必须是 yaml 文件路径)。\n", cp)
os.Exit(2)
}
localConfig, err := config.EnsureLocalConfig(cp)
if err != nil {
fmt.Printf("加载配置失败: %v\n", err)
return
}
cfg, err := config.Load(cp)
if err != nil {
fmt.Printf("加载配置失败: %v\n", err)
return
}
if localConfig.Created {
termout.PrintConfigCreated()
}
if *httpsBootstrap {
config.ApplyDevHTTPSBootstrap(cfg)
@@ -53,15 +63,12 @@ func main() {
if config.MainWebUIUsesHTTPS(&cfg.Server) {
scheme = "https"
}
fmt.Println()
fmt.Printf("→ Web 界面: %s://127.0.0.1:%d/\n", scheme, port)
if scheme == "https" && cfg.Server.TLSAutoSelfSign {
fmt.Println(" (内存自签证书:浏览器首次需确认「继续访问」)")
}
if scheme == "https" && config.ServerHTTPRedirectEnabled(&cfg.Server) {
fmt.Printf(" http://127.0.0.1:%d/ 将自动跳转到 HTTPS\n", port)
}
fmt.Println()
termout.PrintStartupWebUI(termout.StartupWebUIOptions{
Scheme: scheme,
Port: port,
SelfSigned: scheme == "https" && cfg.Server.TLSAutoSelfSign,
HTTPRedirect: scheme == "https" && config.ServerHTTPRedirectEnabled(&cfg.Server),
})
// MCP 启用且 auth_header_value 为空时,自动生成随机密钥并写回配置
if err := config.EnsureMCPAuth(cp, cfg); err != nil {
+35 -7
View File
@@ -10,11 +10,14 @@
# ============================================
# 前端显示的版本号(可选,不填则显示默认版本)
version: "v1.6.51"
version: "v1.7.3"
# 服务器配置
server:
host: 0.0.0.0 # 监听地址,0.0.0.0 表示监听所有网络接口
port: 8080 # 服务端口;未启用 TLS 时为 http://localhost:8080
# 其他可信 Web 集成的精确 Origin。Chromium 浏览器插件会自动识别,无需配置;不要使用通配符。
# cors_allowed_origins:
# - https://trusted-integration.example
# --- 可选:HTTPS + HTTP/2(缓解浏览器对同源 HTTP/1.1 的并发连接数限制,多路 Deep 流式更稳)---
# 启用 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 关闭)。
@@ -28,7 +31,6 @@ server:
tls_auto_self_sign: true
# 认证配置
auth:
password: # Web 登录密码,请修改为强密码
session_duration_hours: 12 # 登录有效期(小时),超时后需重新登录
# 日志配置
log:
@@ -203,14 +205,20 @@ multi_agent:
tool_search_always_visible_tools: [read_file, glob, grep, analyze_image, write_file, edit_file, execute, task, transfer_to_agent, exit, write_todos, skill, tool_search, TaskCreate, TaskGet, TaskUpdate, TaskList, record_vulnerability, list_vulnerabilities, get_vulnerability, list_knowledge_risk_types, search_knowledge_base, webshell_exec, webshell_file_list, webshell_file_read, webshell_file_write, manage_webshell_list, manage_webshell_add, manage_webshell_update, manage_webshell_delete, manage_webshell_test, batch_task_list, batch_task_get, batch_task_start, batch_task_rerun, batch_task_pause, batch_task_update_metadata, batch_task_update_schedule, batch_task_schedule_enabled, batch_task_update_task, batch_task_remove_task, batch_task_delete, batch_task_create, batch_task_add_task, http-framework-test, exec] # 后端内置常驻工具白名单(优先于 always_visible 数量策略)
plantask_enable: true # P0:主代理挂载 TaskCreate/Get/Update/List 结构化任务板;需 eino_skills 可用且 skills_dir 存在
plantask_rel_dir: .eino/plantask # 任务文件相对 skills_dir,按会话分子目录:skills/.eino/plantask/<conversationId>/
reduction_enable: true # true:大工具输出截断/落盘以控上下文;依赖与 plantask 相同的 eino local 写盘后端,无后端时不挂载
reduction_enable: true # true:大工具输出截断/落盘以控上下文;后端会独立创建,不依赖 eino_skills 是否启用
reduction_max_length_for_trunc: 50000 # 单条工具结果超过该字符数(bytes)时截断并落盘(由 reduction 中间件处理)
reduction_max_tokens_for_clear: 160000 # 历史工具结果清理阈值(tokens),超阈值时在模型调用前清理旧结果
reduction_max_tokens_for_clear: 60000 # 历史工具结果清理阈值(tokens),应低于 max_total_tokens * summarization_trigger_ratio
reduction_root_dir: "" # 非空:截断/清理内容落盘根路径;空:使用系统临时目录下按会话隔离的默认路径
reduction_clear_exclude: [] # 不参与「清理阶段」的工具名额外列表(会与 task/transfer/exit 等内置排除项合并);需要时用 YAML 列表填写
reduction_sub_agents: true # true:子代理也挂 reductionfalse:仅编排主代理使用 reduction
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_user_intent_ledger_max_runes: 96000 # 压缩后注入模型上下文的「原始用户输入与约束账本」总字符上限;DB 原始消息不裁剪
summarization_user_intent_ledger_entry_max_runes: 16000 # 账本中单条用户消息的字符上限;超出仅裁剪模型可见账本,不影响 DB 原文
latest_user_message_max_runes: 48000 # 本轮最新 user 进入模型上下文的字符上限;超出时全文落盘,仅注入 head/tail 预览
latest_user_message_head_runes: 24000 # 超长本轮 user 的头部预览字符数
latest_user_message_tail_runes: 24000 # 超长本轮 user 的尾部预览字符数
plan_execute_user_input_budget_ratio: 0.35 # plan_execute 中 userInput 预算比例(planner/replanner/executor 共用)
plan_execute_executed_steps_budget_ratio: 0.2 # plan_execute 中 executed_steps 预算比例
plan_execute_max_step_result_runes: 4000 # plan_execute 每步结果最大字符数(超出截断)
@@ -267,10 +275,11 @@ security:
# MCP (Model Context Protocol) 用于工具注册和调用
mcp:
enabled: false # 是否启用 MCP 服务器(http模式)
host: 0.0.0.0 # MCP 服务器监听地址
host: 127.0.0.1 # MCP 服务器监听地址;需要远程访问时再显式修改并配置网络层访问控制
port: 8081 # MCP 服务器端口
auth_header: "X-MCP-Token" # 鉴权:请求需携带该 header 且值与 auth_header_value 一致方可调用。留空表示不鉴权
auth_header_value: "" # 鉴权密钥值(与 auth_header 配合使用,建议使用随机字符串)
auth_header: "X-MCP-Token" # 可选的全局服务凭证 Header;普通调用请使用用户 Authorization: Bearer Token
auth_header_value: "" # 全局服务凭证值,仅 allow_global_access=true 时生效
allow_global_access: false # 高风险兼容模式:静态密钥映射为全局服务身份;默认请使用用户 Bearer Token
# 外部 MCP 配置
external_mcp:
servers: {}
@@ -339,6 +348,11 @@ knowledge:
robots:
wechat: # 微信 iLink(个人微信 ClawBot,扫码绑定)
enabled: false
# 鉴权默认 user_binding;专用机器人可改 service_account,并必须限制真实发送者
auth:
mode: user_binding # user_binding | service_account
# service_user_id: "admin 或专用服务账号的 RBAC user ID"
# allowed_external_users: ["t:tenant|u:sender"]
bot_token: ""
ilink_bot_id: ""
ilink_user_id: ""
@@ -347,6 +361,8 @@ robots:
bot_agent: CyberStrikeAI/1.0
wecom: # 企业微信
enabled: false
auth:
mode: user_binding
token: ""
encoding_aes_key: ""
corp_id: ""
@@ -354,30 +370,42 @@ robots:
agent_id: 0
dingtalk: # 钉钉
enabled: false
auth:
mode: user_binding
client_id: ""
client_secret: ""
allow_conversation_id_fallback: false
lark: # 飞书
enabled: false
auth:
mode: user_binding
app_id: ""
app_secret: ""
verify_token: ""
allow_chat_id_fallback: false
telegram: # Telegram
enabled: false
auth:
mode: user_binding
bot_token: ""
bot_username: ""
allow_group_messages: false
slack: # Slack
enabled: false
auth:
mode: user_binding
bot_token: ""
app_token: ""
discord: # Discord
enabled: false
auth:
mode: user_binding
bot_token: ""
allow_guild_messages: false
qq: # QQ 机器人
enabled: false
auth:
mode: user_binding
app_id: ""
client_secret: ""
sandbox: true
+68 -49
View File
@@ -1,66 +1,85 @@
# CyberStrikeAI Documentation
Documentation is split by language:
[中文](#中文文档) | [English](#english-documentation)
- [中文文档](zh-CN/)
- [English docs](en-US/)
CyberStrikeAI documentation is organized by user journey. Start with deployment, then move to the topic that matches your task.
## 中文文档
- [部署指南](zh-CN/deployment.md)
- [运维 Runbooks](zh-CN/runbooks.md)
- [配置画像](zh-CN/configuration-profiles.md)
- [安全加固指南](zh-CN/security-hardening.md)
- [API Recipes](zh-CN/api-recipes.md)
- [贡献规范](zh-CN/contributing-guide.md)
- [配置参考](zh-CN/configuration.md)
- [安全模型](zh-CN/security-model.md)
### 按目标开始
- **快速体验**[部署指南](zh-CN/deployment.md) → [配置参考](zh-CN/configuration.md) → [排错指南](zh-CN/troubleshooting.md)
- **生产部署**[配置画像](zh-CN/configuration-profiles.md) → [安全加固](zh-CN/security-hardening.md) → [运维 Runbooks](zh-CN/runbooks.md) → [审计与监控](zh-CN/audit-and-monitoring.md)
- **接入与自动化**[API 参考](zh-CN/api-reference.md) → [API Recipes](zh-CN/api-recipes.md) → [MCP 联邦](zh-CN/mcp-federation.md)
- **参与开发**[开发者指南](zh-CN/developer-guide.md) → [测试指南](zh-CN/testing.md) → [贡献规范](zh-CN/contributing-guide.md)
### 核心概念与编排
- [架构说明](zh-CN/architecture.md)
- [API 参考](zh-CN/api-reference.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)
- [安全模型](zh-CN/security-model.md)
- [Agent 与角色](zh-CN/agent-and-role-guide.md)
- [Skills 指南](zh-CN/skills-guide.md)
- [插件开发](zh-CN/plugin-development.md)
- [发布流程](zh-CN/release-process.md)
- [测试指南](zh-CN/testing.md)
- [图编排使用说明](zh-CN/workflow-graph.md)
- [Eino 多代理](zh-CN/MULTI_AGENT_EINO.md)
- [图编排](zh-CN/workflow-graph.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/frontend-i18n.md)
- [Eino 多代理改造说明](zh-CN/MULTI_AGENT_EINO.md)
- [WebShell 管理](zh-CN/webshell.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)
- [Architecture](en-US/architecture.md)
- [API Reference](en-US/api-reference.md)
- [Troubleshooting](en-US/troubleshooting.md)
- [Audit and Monitoring](en-US/audit-and-monitoring.md)
- [Knowledge Base](en-US/knowledge-base.md)
- [C2 Guide](en-US/c2.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)
- [Security Model](en-US/security-model.md)
- [Agents and Roles](en-US/agent-and-role-guide.md)
- [Skills](en-US/skills-guide.md)
- [Eino Multi-Agent](en-US/MULTI_AGENT_EINO.md)
- [Graph Orchestration](en-US/workflow-graph.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)
- [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)
- [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 -28
View File
@@ -1,29 +1,32 @@
# English Docs
# English Documentation
- [Deployment Guide](deployment.md): deployment modes, HTTPS, reverse proxy, systemd, backup, upgrade, and acceptance checks.
- [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.
- [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.
- [Contributing Guide](contributing-guide.md): checklists for APIs, config, tools, frontend, DB, high-risk features, and docs.
- [Configuration Reference](configuration.md): `config.yaml` fields, hot-apply boundaries, recommended values, and source anchors.
- [Security Model](security-model.md): trust boundaries, HITL, tool execution, C2/WebShell, and data safety.
- [Architecture](architecture.md): request flow, module relationships, complexity hotspots, and design trade-offs.
- [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.
- [Audit and Monitoring](audit-and-monitoring.md): platform audit, tool monitoring, HITL logs, and retention.
- [Knowledge Base](knowledge-base.md): indexing pipeline, retrieval tuning, log analysis, and content writing.
- [C2 Guide](c2.md): lifecycle, task classification, event review, and safety guidance.
- [WebShell Management](webshell.md): operation tiers, naming, AI guardrails, and troubleshooting.
- [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.
- [Skills Guide](skills-guide.md): Skill structure, progressive disclosure, anti-patterns, and local-tool risk.
- [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.
- [Testing Guide](testing.md): test layers, regression focus, test data, and failure cases.
- [Graph Orchestration Guide](workflow-graph.md)
- [HITL Best Practices](hitl-best-practices.md)
- [Robot / Chatbot Guide](robot.md)
- [Vision Analysis](VISION.md)
- [Frontend i18n](frontend-i18n.md)
- [Eino Multi-Agent Notes](MULTI_AGENT_EINO.md)
[Documentation home](../README.md) | [中文](../zh-CN/README.md)
## Choose a path
- **Try locally**: [Deployment](deployment.md) → [Configuration](configuration.md) → [Troubleshooting](troubleshooting.md)
- **Run in production**: [Configuration Profiles](configuration-profiles.md) → [Security Hardening](security-hardening.md) → [Runbooks](runbooks.md) → [Audit and Monitoring](audit-and-monitoring.md)
- **Integrate and automate**: [API Reference](api-reference.md) → [API Recipes](api-recipes.md) → [MCP Federation](mcp-federation.md)
- **Contribute code**: [Developer Guide](developer-guide.md) → [Testing](testing.md) → [Contributing](contributing-guide.md)
## Concepts and orchestration
- [Architecture](architecture.md) · [Security Model](security-model.md) · [RBAC](rbac.md)
- [Agents and Roles](agent-and-role-guide.md) · [Skills](skills-guide.md) · [Eino Multi-Agent](MULTI_AGENT_EINO.md)
- [Graph Orchestration](workflow-graph.md) · [HITL Best Practices](hitl-best-practices.md)
## Feature guides
- [Knowledge Base](knowledge-base.md) · [Robot / Chatbot](robot.md) · [Vision](VISION.md)
- [WebShell](webshell.md) · [C2](c2.md) · [MCP Federation](mcp-federation.md)
## Operations and reference
- [Deployment](deployment.md) · [Configuration](configuration.md) · [Configuration Profiles](configuration-profiles.md)
- [Security Hardening](security-hardening.md) · [Audit and Monitoring](audit-and-monitoring.md) · [Runbooks](runbooks.md)
- [API Reference](api-reference.md) · [API Recipes](api-recipes.md) · [Troubleshooting](troubleshooting.md)
## Development and release
- [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)
+2 -3
View File
@@ -21,7 +21,7 @@ server:
tls_enabled: true
tls_auto_self_sign: true
auth:
password: "dev-only-change-me"
session_duration_hours: 12
audit:
enabled: true
retention_days: 7
@@ -45,7 +45,7 @@ server:
port: 8080
tls_enabled: false
auth:
password: "<long-random-password>"
session_duration_hours: 12
audit:
enabled: true
retention_days: 30
@@ -90,7 +90,6 @@ Goal: long-running production red-team or security platform.
```yaml
auth:
password: "<managed-secret>"
session_duration_hours: 8
audit:
enabled: true
+6 -2
View File
@@ -11,8 +11,10 @@ server:
host: 0.0.0.0
port: 8080
tls_enabled: true
# Optional: other trusted Web integrations; Chromium extensions need no entry.
# cors_allowed_origins:
# - https://trusted-integration.example
auth:
password: "change-me"
session_duration_hours: 12
openai:
provider: openai
@@ -24,7 +26,9 @@ agent:
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
+10
View File
@@ -103,6 +103,16 @@ Update:
- `docs/zh-CN/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
Prioritize:
+374
View File
@@ -0,0 +1,374 @@
# CyberStrikeAI RBAC Administration Guide
[中文](../zh-CN/rbac.md)
CyberStrikeAI can execute Agents, MCP tools, WebShell operations, C2 actions, and batch jobs. RBAC therefore applies beyond navigation visibility: it is enforced across HTTP APIs, resource queries, Agent contexts, built-in and external MCP tools, background jobs, and chatbot execution.
---
## 1. Two different kinds of roles
| Concept | Management location | Purpose |
|---------|---------------------|---------|
| **Platform role (RBAC Role)** | **Platform permissions** | Controls which features and resources a user may access |
| **AI testing role (Agent Role)** | **Roles** / `roles/*.yaml` | Controls Agent prompts, methodology, and candidate tools |
An AI testing role is not an authorization boundary. Selecting a penetration-testing role does not grant platform permissions, and granting RBAC permissions does not change the Agent prompt.
---
## 2. Authorization model
An operation is allowed only when all relevant checks pass:
```text
enabled account
+ required permission for the route/tool
+ scope attached to that permission
+ resource owner / explicit assignment / parent inheritance
+ additional rules for process-global operations
```
Request flow:
1. Login issues a Bearer token whose session contains user, roles, permissions, and per-permission scopes.
2. HTTP middleware maps the route to a permission, for example `GET /api/projects``project:read`.
3. Resource-ID requests also check ownership, explicit assignments, or supported parent inheritance.
4. Agent execution receives an immutable Principal through `context.Context`.
5. Built-in MCP tools authorize both the tool and resource IDs in tool arguments. External MCP has separate restrictions.
6. Denials are written to RBAC/audit logs.
Frontend button hiding is only a usability feature. The server is the security boundary.
---
## 3. Built-in platform roles
| Role | Scope | Default capability |
|------|-------|--------------------|
| **Administrator `admin`** | `all` | Every known permission, including RBAC, configuration, terminal, audit deletion, and global definition management |
| **Operator `operator`** | `assigned` | Normal read/write/execute work; excludes RBAC, core configuration, terminal, audit management, external MCP execution, and several global definition writes |
| **Auditor `auditor`** | `all` | Read permissions across modules plus `audit:read`; no writes |
| **Viewer `viewer`** | `assigned` | Read-only access within authorized resources |
System roles cannot be edited or deleted. Their grants are rebuilt from the current permission catalog during upgrade, preventing stale grants from older versions. Create custom roles for different job functions.
An account without a role can still authenticate but has almost no business capability; do not treat “no role” as a complete job profile.
---
## 4. Permission catalog
Permissions use `module:action`. Common actions are `read`, `write`, `delete`, and `execute`. The authoritative catalog for the running build is available in Platform permissions or `GET /api/rbac/metadata`.
| Module | Permissions |
|--------|-------------|
| Account | `auth:self` |
| Dashboard | `dashboard:read` |
| Chat | `chat:read`, `chat:write`, `chat:delete` |
| Agent | `agent:execute`, `agent:local-execute` |
| HITL | `hitl:read`, `hitl:write` |
| Tasks | `tasks:read`, `tasks:write`, `tasks:delete` |
| Projects | `project:read`, `project:write`, `project:delete` |
| Vulnerabilities | `vulnerability:read`, `vulnerability:write`, `vulnerability:delete` |
| WebShell | `webshell:read`, `webshell:write`, `webshell:delete` |
| C2 | `c2:read`, `c2:write`, `c2:delete` |
| MCP | `mcp:read`, `mcp:execute`, `mcp:write`, `mcp:external:execute` |
| Knowledge | `knowledge:read`, `knowledge:write`, `knowledge:delete` |
| Skills | `skills:read`, `skills:write`, `skills:delete` |
| Markdown Agents | `agents:read`, `agents:write`, `agents:delete` |
| AI testing roles | `roles:read`, `roles:write`, `roles:delete` |
| Workflows | `workflow:read`, `workflow:execute`, `workflow:write`, `workflow:delete` |
| Configuration | `config:read`, `config:write` |
| Terminal | `terminal:execute` |
| Audit | `audit:read`, `audit:delete` |
| RBAC | `rbac:read`, `rbac:write` |
| Notifications | `notification:read`, `notification:write` |
| Robots | `robot:read`, `robot:write` |
| Files | `files:read`, `files:write`, `files:delete` |
| Attack chain | `attackchain:read`, `attackchain:write` |
| FOFA | `fofa:execute` |
| OpenAPI | `openapi:read` |
| Chat groups | `group:read`, `group:write`, `group:delete` |
| Monitor | `monitor:read`, `monitor:write`, `monitor:delete` |
Important distinctions:
- `agent:execute` runs Agents but does not grant local filesystem, shell, or arbitrary configured command access.
- `agent:local-execute` is the local execution fallback and should be limited to trusted operators.
- `mcp:execute` protects the authenticated MCP HTTP entry point.
- `mcp:external:execute` allows Agent calls to external MCP tools and currently also requires `all` scope.
- `mcp:write` manages external MCP configuration; it is separate from external tool execution.
- `robot:write` manages robot configuration and the test endpoint. Chatbot conversations use the bound user or configured service account's business permissions.
---
## 5. Resource scopes
Each role has one scope:
| Scope | Meaning | Typical use |
|-------|---------|-------------|
| `all` | All resources covered by the permission | Administrator, global auditor |
| `assigned` | Explicitly assigned resources and supported parent-resource inheritance | Project member, assigned asset operator |
| `own` | Primarily resources created by/owned by the user; some resource types also support explicit assignment or parent inheritance | Personal workspace, isolated robot identity |
Users may have multiple roles. Permissions are unioned, while scopes are merged **for the same permission only**:
```text
all > assigned > own
```
Example:
```text
Global audit role: project:read + all
Personal editor: project:write + own
Effective:
project:read → all
project:write → own
```
A global read role does not widen an unrelated write permission. Authorization code must use `ScopeFor(permission)`, not the user's broadest display scope.
### Process-global restrictions
Some definitions have no owner. Their mutations require the corresponding permission with `all` scope even if the user has a `write` key:
- AI testing roles, Skills, and Markdown Agents.
- External MCP configuration.
- Robot configuration.
- Workflow definitions.
- Knowledge mutations other than search.
- Global HITL allowlist, reviewer, and audit policy.
- C2 Profile mutations.
- Some global monitor statistics.
---
## 6. Ownership, assignments, and inheritance
Use Platform permissions → Member details → Resource assignments. Directly assignable resource types include:
- `project`
- `conversation`
- `vulnerability`
- `webshell`
- `batch_task`
- `c2_listener`
A batch request accepts at most 100 resources. Duplicate grants are skipped.
Supported inheritance includes:
| Child resource | Parent access source |
|----------------|----------------------|
| Conversation | Project |
| Vulnerability | Project or related conversation |
| Message, process detail, attack chain | Conversation |
| C2 Session | Listener |
| C2 Task/file/event | Session, Task, or Listener chain |
Assigning a project therefore usually avoids assigning each conversation and vulnerability separately. The concrete route/tool server check remains authoritative.
---
## 7. Web administration workflow
### Create a user
1. Sign in as an administrator and open **Platform permissions**.
2. Create a user with username, display name, an eight-character-or-longer password, and enabled status.
3. Assign one or more platform roles.
4. For `assigned` roles, configure resource assignments.
5. Have the user sign in again and verify roles, permission count, and scope in the top-right user menu.
### Create a custom role
1. Give the role a job-oriented name and description.
2. Select `all`, `assigned`, or `own`.
3. Select only required permissions.
4. Test list, detail, mutation, deletion, Agent, and tool behavior with a test account.
5. Assign it to production users only after verification.
System roles are immutable; create a custom role instead of modifying them.
### When changes take effect
- Updating a user, password, enabled state, or role membership revokes that user's sessions; they must sign in again.
- Updating or deleting a custom role revokes all sessions; all users must sign in again.
- Robots resolve the bound user/service account on every message, so disablement and role changes affect the next message.
- Background batch jobs resolve a Principal from the task owner rather than trusting frontend state.
---
## 8. Suggested role templates
### Read-only project member
```text
Scope: assigned
dashboard:read
chat:read
project:read
vulnerability:read
files:read
attackchain:read
```
### Daily security operator
```text
Scope: assigned
agent:execute
chat:read / chat:write
project:read / project:write
vulnerability:read / vulnerability:write
tasks:read / tasks:write
files:read / files:write
hitl:read / hitl:write
```
Add `agent:local-execute` or `terminal:execute` only when local commands are required. Add individual `:delete` permissions only when deletion is part of the job.
### Robot service account
```text
Scope: own (isolated workspace) or assigned (specific projects)
agent:execute
chat:read / chat:write
optional project, vulnerability, and knowledge permissions
```
`admin` can be used as a robot service account, but exact sender allowlisting still applies. Every allowlisted sender receives full permissions and shares admin-owned data. See the [Robot guide](robot.md).
---
## 9. Agent, MCP, and robot boundaries
### Agent
The HTTP user becomes an immutable Principal propagated to single-agent, multi-agent, workflow, and tool contexts. A long-running task may survive an SSE disconnect while retaining that identity.
### Built-in MCP
Every built-in tool requires an explicit authorization policy. WebShell tools check both `webshell:read/write/delete` and the target `connection_id`; project, vulnerability, task, and C2 tools validate resource arguments as well. An unregistered built-in policy fails closed. Other local/configured tools require `agent:local-execute`.
### External MCP
Agent calls to external MCP require `mcp:external:execute` with `all` scope because an external service's resource model is not protected by local ownership and assignments.
### Robots
- `user_binding`: each platform sender binds their own RBAC user.
- `service_account`: exact allowlisted senders share one RBAC user.
- Platform signature verification authenticates message origin, not business authorization.
- Run `whoami` to inspect the effective Principal.
---
## 10. RBAC API
All requests use:
```http
Authorization: Bearer <token>
```
Management routes require `rbac:read` or `rbac:write`; the resource picker requires `rbac:write`.
| Method | Path | Purpose |
|--------|------|---------|
| GET | `/api/rbac/me` | Current user, roles, permissions, overall scope, per-permission scopes |
| GET | `/api/rbac/metadata` | Permission catalog, roles, grants, and scopes |
| GET/POST | `/api/rbac/users` | List/create users |
| PUT/DELETE | `/api/rbac/users/:id` | Update/delete a user |
| GET/POST | `/api/rbac/roles` | List/create roles |
| PUT/DELETE | `/api/rbac/roles/:id` | Update/delete a custom role |
| GET | `/api/rbac/resources?type=project&q=...` | Search assignable resources |
| GET/POST | `/api/rbac/resource-assignments` | List/create assignments |
| DELETE | `/api/rbac/resource-assignments/:id` | Revoke an assignment |
Create a user:
```bash
curl -X POST http://localhost:8080/api/rbac/users \
-H "Authorization: Bearer $TOKEN" \
-H "Content-Type: application/json" \
-d '{
"username": "operator01",
"display_name": "Security Operator 01",
"password": "change-me-123",
"enabled": true,
"roles": ["operator"]
}'
```
Create a custom role:
```bash
curl -X POST http://localhost:8080/api/rbac/roles \
-H "Authorization: Bearer $TOKEN" \
-H "Content-Type: application/json" \
-d '{
"name": "Project Auditor",
"description": "Read assigned projects",
"scope": "assigned",
"permissions": ["chat:read", "project:read", "vulnerability:read"]
}'
```
Assign projects:
```bash
curl -X POST http://localhost:8080/api/rbac/resource-assignments \
-H "Authorization: Bearer $TOKEN" \
-H "Content-Type: application/json" \
-d '{
"user_id": "USER_ID",
"resource_type": "project",
"resource_ids": ["PROJECT_ID_1", "PROJECT_ID_2"]
}'
```
---
## 11. Audit and operations recommendations
- Use individual administrator accounts instead of sharing one password.
- Name custom roles by job function and document purpose/owner.
- Review high-risk permissions separately: `terminal:execute`, `agent:local-execute`, `c2:write/delete`, `webshell:write/delete`, `rbac:write`, and `config:write`.
- Periodically review `all` roles, service accounts, robot allowlists, and dormant users.
- On offboarding, disable the account first, then revoke robot bindings, assignments, and sessions.
- Monitor RBAC denials, user/role changes, resource assignments, and robot service-account execution in audit logs.
- Pair RBAC with HITL for dangerous tools; permission to invoke does not bypass approval policy.
---
## 12. Troubleshooting
### A button is missing
The frontend hides actions based on `/api/rbac/me`. Verify the required permission. Direct API calls are still rejected server-side.
### Permission exists but the resource is denied
Inspect the scope for that specific permission, not only the overall display scope. Then check owner, explicit assignment, and parent assignment.
### Role changed but the user sees old access
Role changes revoke sessions. Sign in again. Robots resolve again on the next message.
### A global mutation is denied despite `write`
Process-global definitions require the corresponding permission with `all` scope. Create a dedicated global administration role instead of widening unrelated permissions.
### Agent chat works but commands fail
`agent:execute` and `agent:local-execute` are separate. Grant local execution only when necessary and combine it with HITL, tool allowlists, and audit.
### External MCP requires global scope
The user needs `mcp:external:execute`, and that permission's scope must be `all`.
+128 -22
View File
@@ -2,7 +2,7 @@
[中文](../zh-CN/robot.md)
This document explains how to chat with CyberStrikeAI from **personal WeChat**, **DingTalk**, **Lark (Feishu)**, and **WeCom (Enterprise WeChat)** using long-lived connections or HTTP callbacks—no need to open a browser on the server. Following the steps below helps avoid common mistakes.
This guide covers **Personal WeChat, WeCom, DingTalk, Lark, Telegram, Slack, Discord, and QQ Bot**, including platform connectivity, RBAC identity binding, service-account allowlists, commands, verification, and troubleshooting.
---
@@ -15,10 +15,18 @@ This document explains how to chat with CyberStrikeAI from **personal WeChat**,
- **Personal WeChat**: Open **WeChat / iLink****Generate QR code and bind**, then scan with WeChat (see [Section 3.4](#34-personal-wechat-wechat--ilink))
- **DingTalk**: Enable and fill in Client ID / Client Secret
- **Lark**: Enable and fill in App ID / App Secret
5. Click **Apply configuration** to save (WeChat binding saves and enables automatically on success—usually no extra click needed)
6. **Restart the CyberStrikeAI process** (DingTalk/Lark: saving alone does not establish the connection; WeChat auto-restarts the iLink poll after binding—usually no manual restart needed)
5. Click **Apply configuration** to save and automatically restart the corresponding bot connection. WeChat binding saves and enables automatically on success.
Settings are written to the `robots` section of `config.yaml`; you can also edit the file directly. **After changing DingTalk or Lark config, you must restart for the long-lived connection to take effect.** Personal WeChat binding automatically writes `robots.wechat` and restarts the iLink long poll.
Settings are written to the `robots` section of `config.yaml`; you can also edit the file directly. Web-based **Apply configuration** restarts the corresponding connection automatically. Restart the CyberStrikeAI process only when editing YAML directly. Personal WeChat binding automatically writes `robots.wechat` and restarts the iLink long poll.
### Shortest path to first use
After the platform connection works, configure the business identity before sending normal prompts:
- **Multiple users**: choose User binding → each user generates a code from the top-right Web user menu → sends the bind command to the bot → runs `whoami` to verify.
- **Only you**: run `whoami` first and copy the sender ID → choose Service account → set User ID to `admin` or another RBAC user → paste the exact sender allowlist → apply configuration → run `whoami` again.
Start normal AI chat only after the response shows an authorized status and the expected effective identity.
---
@@ -82,7 +90,7 @@ If you only have a **custom bot** Webhook URL (`oapi.dingtalk.com/robot/send?acc
- In CyberStrikeAI: System settings → Robot settings → DingTalk.
- Enable “Enable DingTalk robot”.
- Paste the Client ID and Client Secret from step 3.
- Click **Apply configuration**, then **restart CyberStrikeAI**.
- Click **Apply configuration**; CyberStrikeAI restarts the DingTalk connection automatically.
---
@@ -105,7 +113,7 @@ If you only have a **custom bot** Webhook URL (`oapi.dingtalk.com/robot/send?acc
| App Secret | From Lark open platform app credentials |
| Verify Token | Optional; for event subscription |
**Lark setup in short**: Log in to [Lark Open Platform](https://open.feishu.cn) → Create an enterprise app → In “Credentials and basic info” get **App ID** and **App Secret** → In “Application capabilities” enable **Robot** and the right permissions → Add **event subscription** and **permissions** below → Publish the app → Enter App ID and App Secret in CyberStrikeAI robot settings → Save and **restart** the app.
**Lark setup in short**: Log in to [Lark Open Platform](https://open.feishu.cn) → Create an enterprise app → In “Credentials and basic info” get **App ID** and **App Secret** → In “Application capabilities” enable **Robot** and the right permissions → Add **event subscription** and **permissions** below → Publish the app → Enter App ID and App Secret in CyberStrikeAI robot settings → **Apply configuration**.
**Event subscription**
The long-lived connection only receives message events if you subscribe to them. In the apps **Events and callbacks** (事件与回调) → **Event subscription** (事件订阅), add the event **Receive message** (**im.message.receive_v1**). Without it, the connection succeeds but no message events are delivered (no logs when users send messages).
@@ -272,12 +280,79 @@ In **Permission management** (权限管理), enable the following (names and ide
---
## 4. Bot commands
## 4. RBAC authorization and bot commands
Platform credentials and callback signatures authenticate the messaging platform. CyberStrikeAI RBAC determines what the sender can actually do. Each bot instance uses one authorization mode.
### 4.1 Choose an authorization mode
| Scenario | Recommended mode | Identity and data behavior |
|----------|------------------|----------------------------|
| Shared WeCom, Lark, DingTalk, or Slack bot | `user_binding` | Each sender binds their own Web user; permissions and resources remain isolated |
| Personal WeChat, single-user bot, fixed automation entry | `service_account` | Allowlisted senders share the configured RBAC user's permissions and owned resources |
Both modes resolve user status, roles, per-permission scope, and resource assignments before every message. Basic AI chat requires:
```text
agent:execute
chat:read
chat:write
```
Grant project, role, local execution, WebShell, C2, or MCP permissions only when those features are required. Conversation deletion also requires `chat:delete`.
### 4.2 User-binding mode (default)
Administrator:
1. Open System settings → Robot settings → select a platform.
2. Set Authorization policy to `user_binding` and apply the configuration.
Each user:
1. Sign in to the Web UI and open the top-right user menu → **Bind robot account**.
2. Generate a binding code; a five-minute countdown starts.
3. Send the full command to the target bot, for example `bind 7C6E-BD4C`.
4. Send `whoami` and confirm the effective RBAC identity is their own Web user.
Codes are stored only as hashes and are single-use. When the countdown ends, the UI marks the code expired, disables copying, and refreshes the binding list; the server also rejects it. Generating a new code immediately invalidates the previous unused code. Users can send `unbind` or revoke a binding from the Web dialog.
### 4.3 Service-account mode
1. Connect the bot to its messaging platform.
2. Have each intended sender run `whoami` and copy the exact sender ID. For Personal WeChat it usually resembles `xxxx@im.wechat`; never substitute `ilink_bot_id` or configured `ilink_user_id`.
3. In Robot settings, select `service_account`.
4. Enter the RBAC **User ID**, not its display name. `admin` is allowed; every allowlisted sender then receives full platform permissions and the UI shows a red warning.
5. Add one exact sender ID per line. Matching is case-sensitive and `*` wildcards are rejected.
6. Apply configuration and run `whoami` again to verify the effective user, roles, and scope.
Example:
```yaml
robots:
wechat:
auth:
mode: service_account
service_user_id: admin
allowed_external_users:
- "o9cq806s32Sm2_kyOmkyaV7Rn1lU@im.wechat"
```
Service-account mode rejects `bind` and `unbind`. All allowlisted senders share conversations, projects, and other resources owned by the service account. Use `user_binding` when that sharing is undesirable.
### 4.4 Inspect the effective identity
Send `whoami`. The response includes platform, exact sender ID, authorization mode and status, effective RBAC user and ID, roles, scope, and permission count. A non-allowlisted sender sees only the denial status and no service-account details.
### 4.5 Command list
Send these **text commands** to the bot on any connected platform (text only):
| Command | Description |
|---------|-------------|
| **绑定 \<code\>** or **bind \<code\>** | Bind the verified platform sender to the RBAC user that generated the code |
| **解绑** or **unbind** | Remove the current platform identity binding |
| **身份** or **whoami** | Show sender ID, authorization mode, binding status, and the effective RBAC user, roles, and scope |
| **帮助** (help) | Show command help |
| **列表** or **对话列表** (list) | List all conversation titles and IDs |
| **切换 \<conversationID\>** or **继续 \<conversationID\>** | Continue in the given conversation |
@@ -292,6 +367,8 @@ Send these **text commands** to the bot on any connected platform (text only):
Any other text is sent to the AI as a user message, same as in the web UI (e.g. penetration testing, security analysis).
Group messages are authorized as the actual sender, never as a group ID. In service-account mode, explicitly allowlisted senders intentionally share the configured account.
---
## 5. How to use (do I need to @ the bot?)
@@ -310,14 +387,17 @@ Summary: **Personal WeChat and direct chat**—just send; **DingTalk/Lark in a g
1. CyberStrikeAI web UI → System settings → Robot settings → **WeChat / iLink****Generate QR code and bind**.
2. Scan with WeChat and confirm (enter pairing code on the web page if prompted).
3. After binding, send “帮助” in the WeChat private chat to test.
3. Send `whoami` in the WeChat private chat and copy the sender ID.
4. Choose `user_binding`, or configure `service_account` with the RBAC user and exact sender allowlist.
5. Apply configuration, run `whoami` again, then send a normal message.
**DingTalk / Lark**
1. **In the open platform**: Complete app creation, copy credentials, enable the bot (DingTalk: **Stream mode**), set permissions, and publish (Section 3).
2. **In CyberStrikeAI**: System settings → Robot settings → Enable the platform, paste Client ID/App ID and Client Secret/App Secret → **Apply configuration**.
3. **Restart the CyberStrikeAI process** (otherwise the long-lived connection is not established).
4. **On your phone**: Open DingTalk or Lark, find the bot (direct chat or @ in a group), send “帮助” or any message to test.
3. **Choose authorization**: use `user_binding` for multiple users, or configure a service account and exact allowlist for a dedicated bot.
4. **Apply configuration**; the Web UI restarts the corresponding connection automatically.
5. **On your phone**: Open the bot, run `whoami` first, then send a normal message.
If the bot does not respond, see **Section 9 (troubleshooting)** and **Section 10 (common pitfalls)**.
@@ -331,6 +411,11 @@ Example `robots` section in `config.yaml`:
robots:
wechat: # Personal WeChat iLink (auto-filled after QR bind; usually no manual edit)
enabled: true
auth:
mode: service_account
service_user_id: admin
allowed_external_users:
- "exact sender ID copied from whoami"
bot_token: "your_bot_token@im.bot:..."
ilink_bot_id: "your_bot_id@im.bot"
ilink_user_id: "your_user_id@im.wechat"
@@ -339,10 +424,14 @@ robots:
bot_agent: "CyberStrikeAI/1.0"
dingtalk:
enabled: true
auth:
mode: user_binding
client_id: "your_dingtalk_app_key"
client_secret: "your_dingtalk_app_secret"
lark:
enabled: true
auth:
mode: user_binding
app_id: "your_lark_app_id"
app_secret: "your_lark_app_secret"
verify_token: ""
@@ -372,7 +461,7 @@ robots:
sandbox: true
```
After changing DingTalk/Lark/WeCom/Telegram/Slack/Discord/QQ settings, **Apply configuration** restarts the corresponding connections. Personal WeChat QR binding saves and restarts automatically.
Authorization is configured independently per platform; omitting `auth` defaults to `user_binding`. **Apply configuration** restarts the corresponding connections. Restart the process only after editing YAML directly. Personal WeChat QR binding saves and restarts automatically.
---
@@ -380,20 +469,24 @@ After changing DingTalk/Lark/WeCom/Telegram/Slack/Discord/QQ settings, **Apply c
You can verify bot logic with the **test API** (no DingTalk/Lark client needed):
1. Log in to the CyberStrikeAI web UI (so you have a session).
2. Call the test endpoint with curl (include your session Cookie):
1. Sign in with an account that has global `robot:write` permission and obtain a Bearer token.
2. Call the test endpoint with curl:
```bash
# Replace YOUR_COOKIE with the Cookie from your browser (F12 → Network → any request → Request headers → Cookie)
# Adjust the URL, username, and password for your deployment
TOKEN=$(curl -s -X POST "http://localhost:8080/api/auth/login" \
-H "Content-Type: application/json" \
-d '{"username":"admin","password":"YOUR_PASSWORD"}' | jq -r '.token')
curl -X POST "http://localhost:8080/api/robot/test" \
-H "Content-Type: application/json" \
-H "Cookie: YOUR_COOKIE" \
-H "Authorization: Bearer $TOKEN" \
-d '{"platform":"dingtalk","user_id":"test_user","text":"帮助"}'
```
If the JSON response contains `"reply":"【CyberStrikeAI 机器人命令】..."`, command handling works. You can also try `"text":"列表"` or `"text":"当前"`.
If the JSON response contains `"reply":"【CyberStrikeAI 机器人命令】..."`, command handling works. `help`, `version`, and `whoami` work before binding. `list`, `current`, and normal AI messages enforce RBAC: the test `platform + user_id` must already be bound or exactly match the service-account allowlist.
API: `POST /api/robot/test` (requires login). Body: `{"platform":"optional","user_id":"optional","text":"required"}`. Response: `{"reply":"..."}`.
API: `POST /api/robot/test` (requires global `robot:write`). Body: `{"platform":"optional","user_id":"optional","text":"required"}`. Response: `{"reply":"..."}`. This endpoint simulates bot business logic only; it does not validate a third-party callback signature or long-lived connection.
---
@@ -407,7 +500,7 @@ Check in this order:
Robot settings should show “Connected” or a bound Bot ID; `robots.wechat.bot_token` in `config.yaml` must not be empty.
2. **Enabled?**
Confirm “Enable WeChat robot” is checked; restart CyberStrikeAI if you just changed settings.
Confirm “Enable WeChat robot” is checked and click **Apply configuration** if you just changed settings.
3. **Application logs**
- On startup: `微信 iLink 长轮询已启动`;
@@ -430,8 +523,8 @@ Check in this order:
1. **Client ID / Client Secret match the open platform exactly**
Copy from “Credentials and basic info”; avoid typing. Watch **0** vs **o** and **1** vs **l** (e.g. `ding9gf9tiozuc504aer` has **504**, not 5o4).
2. **Did you restart after saving?**
The long-lived connection is created at **startup**. “Apply configuration” only updates the config file; you **must restart the CyberStrikeAI process** for the DingTalk connection to start.
2. **Did you apply the configuration?**
Web changes require **Apply configuration**, which restarts the corresponding connection automatically. Restart the process only after editing `config.yaml` directly.
3. **Application logs**
- On startup you should see: `钉钉 Stream 正在连接…`, `钉钉 Stream 已启动(无需公网),等待收消息`.
@@ -441,6 +534,15 @@ Check in this order:
4. **Open platform**
The app must be **published**. Under “Robot” you must enable **Stream** for receiving messages (HTTP callback only is not enough). Permission management must include robot receive/send message permissions.
### 9.3 Reply says unbound, sender denied, or permission missing
1. Run `whoami` and inspect the authorization mode and status.
2. In `user_binding`, generate a code from the top-right Web user menu and send the complete bind command from the same platform identity. Regenerate expired or already-used codes.
3. In `service_account`, copy the exact sender ID from `whoami` into that platform's allowlist. Preserve case, tenant prefixes, and suffixes such as `@im.wechat`.
4. If an effective user is shown but permissions are missing, grant at least `agent:execute`, `chat:read`, and `chat:write` for normal AI chat.
5. A missing or disabled service user is rejected when applying configuration.
6. If `admin` is denied, the usual cause is an allowlist mismatch—not insufficient admin permissions.
---
## 10. Common pitfalls
@@ -448,7 +550,11 @@ Check in this order:
- **Personal WeChat vs WeCom**: Personal WeChat uses `robots.wechat` + web QR bind; WeCom uses `robots.wecom` + admin callback URL—they are completely different.
- **WeChat QR expired**: QR codes last ~5 minutes; regenerate instead of reusing an old one.
- **Wrong bot type**: The “Custom” bot added in a DingTalk **group** (Webhook + sign secret) **cannot** be used for two-way chat. Only the **enterprise internal app** bot from the open platform is supported.
- **Saved but not restarted**: After changing DingTalk/Lark robot settings you **must restart** the app (WeChat QR bind restarts the connection automatically).
- **Configuration not applied**: Click **Apply configuration** after Web changes; connections restart automatically. A process restart is needed only for direct YAML edits.
- **Bot ID used as sender ID**: Copy the sender ID from `whoami`; do not use `ilink_bot_id`, configured `ilink_user_id`, a group ID, or a display name.
- **Reusing an expired code**: Codes last five minutes and are single-use; generating a new code immediately invalidates the old one.
- **Assuming service-account users are isolated**: All allowlisted senders share that account's conversations and owned resources. Use `user_binding` for isolation.
- **Assuming admin removes the allowlist**: It does not. The sender must still match exactly, but every matching sender gets full permissions.
- **Client ID typo**: If the platform shows `504`, use `504` (not `5o4`); prefer copy/paste.
- **DingTalk: only HTTP callback, no Stream**: This app receives messages via **Stream**. In the open platform, message reception must be **Stream mode**.
- **App not published**: After changing the bot or permissions in the open platform, **publish a new version** under “Version management and release”, or changes wont apply.
@@ -459,5 +565,5 @@ Check in this order:
- All platforms: **text messages only**; other types (e.g. image, voice) are not supported and may be ignored.
- Personal WeChat: **private chat only**—group @-bot is not supported.
- Conversations are shared with the web UI: conversations created from the bot appear in the web “Conversations” list and vice versa.
- Bot data is shared with the web UI: under `user_binding` it belongs to the bound user; under `service_account` it belongs to the service account and is shared by allowlisted senders.
- Bot execution uses the same **Eino single/multi-agent** path as the web UI (`ProcessMessageForRobot`, with progress callbacks and process details stored in the DB); only the final reply is sent back to personal WeChat/DingTalk/Lark/WeCom in one message (no SSE). Default: `robot_default_agent_mode: eino_single`.
+1 -1
View File
@@ -48,7 +48,7 @@ config.yaml
```yaml
auth:
password: "<long-random-password>"
session_duration_hours: 12
server:
host: 127.0.0.1
port: 8080
+1 -1
View File
@@ -6,7 +6,7 @@ This checklist covers pre-production and continuous hardening for CyberStrikeAI.
## 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.
- Restrict access by IP, VPN, or bastion.
- Enable `audit.enabled`.
+1 -1
View File
@@ -37,7 +37,7 @@ Page inaccessible:
Login fails:
- wrong `auth.password`;
- wrong RBAC user password;
- config not applied/restarted;
- stale cookie;
- 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` | 无资源 IDdetail 仅含错误码与包 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 -27
View File
@@ -1,29 +1,32 @@
# 中文文档
- [部署指南](deployment.md):部署形态、HTTPS、反向代理、systemd、备份、升级和验收。
- [运维 Runbooks](runbooks.md):生产部署、外部 MCP、知识库、Web 测试、C2 清理和工具排障的操作步骤。
- [配置画像](configuration-profiles.md):本地开发、内网团队、知识库、高审计生产、C2 演练等推荐配置。
- [安全加固指南](security-hardening.md):上线前基线、反向代理、HITL 白名单、文件权限和周期巡检。
- [API Recipes](api-recipes.md):登录、Agent、流式、多代理、上传、漏洞、知识库、MCP 和审计导出示例。
- [贡献规范](contributing-guide.md):新增 API、配置、工具、前端、数据库、高风险能力和文档的 checklist。
- [配置参考](configuration.md)`config.yaml` 字段、热应用边界、参数建议和源码锚点。
- [安全模型](security-model.md):信任边界、HITL、工具执行、C2/WebShell 与数据安全。
- [架构说明](architecture.md):请求路径、模块关系、复杂度热点和设计取舍。
- [API 参考](api-reference.md):认证、OpenAPI、SSE、稳定性分层和常用接口。
- [排错指南](troubleshooting.md):诊断顺序、最小命令、常见误判和故障模板。
- [审计与监控](audit-and-monitoring.md):平台审计、工具监控、HITL 日志和保留策略。
- [知识库](knowledge-base.md):索引链路、检索调参、日志分析和内容写法。
- [C2 使用说明](c2.md):生命周期、任务分级、事件复盘和安全建议。
- [WebShell 管理](webshell.md):操作分层、连接命名、AI 约束和排错。
- [MCP 联邦](mcp-federation.md):内置 MCP、外部 MCP、生命周期和工具命名。
- [Agent 与角色](agent-and-role-guide.md):角色、子代理、Skill、编排模式和工具可见性。
- [Skills 指南](skills-guide.md):Skill 结构、渐进式披露、反模式和本地工具风险。
- [插件开发](plugin-development.md):API 插件、MCP 插件、资源包插件和安全边界。
- [发布流程](release-process.md):发布风险、配置兼容、数据库迁移和验收。
- [测试指南](testing.md):测试分层、回归重点、测试数据和失败用例。
- [图编排使用说明](workflow-graph.md)
- [人机协同最佳实践](hitl-best-practices.md)
- [机器人使用说明](robot.md)
- [视觉分析](VISION.md)
- [前端国际化方案](frontend-i18n.md)
- [Eino 多代理改造说明](MULTI_AGENT_EINO.md)
[文档首页](../README.md) | [English](../en-US/README.md)
## 按目标开始
- **快速体验**[部署指南](deployment.md) → [配置参考](configuration.md) → [排错指南](troubleshooting.md)
- **生产部署**[配置画像](configuration-profiles.md) → [安全加固](security-hardening.md) → [运维 Runbooks](runbooks.md) → [审计与监控](audit-and-monitoring.md)
- **接入与自动化**[API 参考](api-reference.md) → [API Recipes](api-recipes.md) → [MCP 联邦](mcp-federation.md)
- **参与开发**[开发者指南](developer-guide.md) → [测试指南](testing.md) → [贡献规范](contributing-guide.md)
## 核心概念与编排
- [架构说明](architecture.md) · [安全模型](security-model.md) · [RBAC](rbac.md)
- [Agent 与角色](agent-and-role-guide.md) · [Skills](skills-guide.md) · [Eino 多代理](MULTI_AGENT_EINO.md)
- [图编排](workflow-graph.md) · [人机协同最佳实践](hitl-best-practices.md)
## 功能指南
- [知识库](knowledge-base.md) · [机器人接入](robot.md) · [视觉分析](VISION.md)
- [WebShell](webshell.md) · [C2](c2.md) · [MCP 联邦](mcp-federation.md)
## 运维与参考
- [部署指南](deployment.md) · [配置参考](configuration.md) · [配置画像](configuration-profiles.md)
- [安全加固](security-hardening.md) · [审计与监控](audit-and-monitoring.md) · [运维 Runbooks](runbooks.md)
- [API 参考](api-reference.md) · [API Recipes](api-recipes.md) · [排错指南](troubleshooting.md)
## 开发与发布
- [开发者指南](developer-guide.md) · [插件开发](plugin-development.md) · [前端国际化](frontend-i18n.md)
- [测试指南](testing.md) · [贡献规范](contributing-guide.md) · [发布流程](release-process.md)
+2 -3
View File
@@ -21,7 +21,7 @@ server:
tls_enabled: true
tls_auto_self_sign: true
auth:
password: "dev-only-change-me"
session_duration_hours: 12
audit:
enabled: true
retention_days: 7
@@ -54,7 +54,7 @@ server:
port: 8080
tls_enabled: false
auth:
password: "<long-random-password>"
session_duration_hours: 12
audit:
enabled: true
retention_days: 30
@@ -107,7 +107,6 @@ multi_agent:
```yaml
auth:
password: "<managed-secret>"
session_duration_hours: 8
audit:
enabled: true
+7 -4
View File
@@ -5,14 +5,16 @@ CyberStrikeAI 的主配置文件是 `config.yaml`。大多数配置也可以在
## 基础配置
```yaml
version: "v1.6.51"
version: "vX.Y.Z" # 占位符;请使用 config.example.yaml 中当前发布版本的值
server:
host: 0.0.0.0
port: 8080
tls_enabled: true
tls_auto_self_sign: true
# 可选:其他可信 Web 集成;Chromium 浏览器插件无需配置
# cors_allowed_origins:
# - https://trusted-integration.example
auth:
password: "change-me"
session_duration_hours: 12
log:
level: info
@@ -22,8 +24,9 @@ log:
- `version`:前端展示版本。
- `server.host/port`Web 服务监听地址和端口。
- `server.tls_*`HTTPS 配置。生产环境建议使用 `tls_cert_path``tls_key_path`
- `auth.password`Web 登录密码,必须改为强密码
- `auth.session_duration_hours`:登录会话有效期
- Chromium 浏览器插件的合法 `chrome-extension://<32位插件ID>` Origin 会被自动识别,无需配置。插件仍需按域授权,并使用密码登录与 Bearer Token 调用 API
- `server.cors_allowed_origins`:仅供其他可信 Web 集成使用的额外 Origin 精确白名单;不支持 `*`,修改后需重启服务
- `auth.session_duration_hours`:登录会话有效期(小时)。登录密码由 RBAC 用户管理,首次启动时在控制台输出 `admin` 初始密码。
- `log.output`:可以是 `stdout``stderr` 或文件路径。
## 模型配置
+10
View File
@@ -103,6 +103,16 @@ docs/en-US/
- `docs/zh-CN/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 关注点
代码评审优先看:
+386
View File
@@ -0,0 +1,386 @@
# CyberStrikeAI RBAC 使用与管理指南
[English](../en-US/rbac.md)
CyberStrikeAI 是可执行 Agent、MCP、WebShell、C2 和批量任务的安全自动化平台。RBAC 不仅控制页面是否可见,还会贯穿 HTTP API、资源查询、Agent 上下文、内置/外部 MCP 工具、后台任务和机器人执行链路。
---
## 一、先区分两种“角色”
| 概念 | 管理入口 | 作用 |
|------|----------|------|
| **平台角色(RBAC Role** | 左侧 **平台权限** | 决定用户能调用哪些功能、能访问哪些资源 |
| **AI 测试角色(Agent Role** | 左侧 **角色** / `roles/*.yaml` | 决定 Agent 的提示词、测试方法和可选工具集合 |
AI 测试角色不是安全授权边界。即使选择了“渗透测试”角色,用户仍必须拥有对应的平台权限;反过来,RBAC 有权限也不会自动改变 Agent 提示词。
---
## 二、授权模型
一次访问同时满足以下条件才会放行:
```text
有效账号
+ 路由/工具所需 permission
+ 该 permission 对应的 scope
+ 目标资源 owner / 显式授权 / 父资源继承
+ 全局操作的额外限制
```
处理链路:
1. 登录后签发 Bearer Token,会话包含用户、角色、权限和逐权限 Scope。
2. HTTP 中间件把路由映射为权限,例如 `GET /api/projects``project:read`
3. 对带资源 ID 的请求继续校验 owner、显式资源授权或父资源继承。
4. Agent 启动时把不可变 Principal 写入 `context.Context`
5. 内置 MCP 工具根据工具和参数再次检查权限与资源;外部 MCP 也有独立限制。
6. 拒绝事件写入 RBAC/审计日志。
前端隐藏按钮只用于改善体验,不构成安全边界;真正的拒绝发生在服务端。
---
## 三、内置平台角色
| 角色 | Scope | 默认能力 |
|------|-------|----------|
| **管理员 `admin`** | `all` | 所有已知权限,包括 RBAC、配置、终端、审计删除和全局定义管理 |
| **操作员 `operator`** | `assigned` | 日常读写与执行能力;不含 RBAC、核心配置、终端、审计管理、外部 MCP 执行和部分全局定义写权限 |
| **审计员 `auditor`** | `all` | 各模块只读权限和 `audit:read`,不执行写操作 |
| **只读用户 `viewer`** | `assigned` | 各模块只读权限,仅查看被授权范围 |
系统角色不可修改或删除,升级时会按当前版本的权限目录重新构建授权,避免旧版本残留权限。需要不同组合时创建自定义角色。
没有分配任何角色的账号仍可登录,但基本没有业务权限;不要依赖“无角色”作为完整岗位配置。
---
## 四、权限命名与目录
权限使用 `模块:动作` 命名。常见动作:
- `read`:查看、列表、查询、导出。
- `write`:创建、更新、执行或管理。
- `delete`:删除。
- `execute`:执行 Agent、终端、工作流或特定能力。
当前权限按模块分组如下;运行版本的权威目录以“平台权限”页面或 `GET /api/rbac/metadata` 为准。
| 模块 | 权限 |
|------|------|
| 账号 | `auth:self` |
| 仪表盘 | `dashboard:read` |
| 对话 | `chat:read``chat:write``chat:delete` |
| Agent | `agent:execute``agent:local-execute` |
| HITL | `hitl:read``hitl:write` |
| 任务 | `tasks:read``tasks:write``tasks:delete` |
| 项目 | `project:read``project:write``project:delete` |
| 漏洞 | `vulnerability:read``vulnerability:write``vulnerability:delete` |
| WebShell | `webshell:read``webshell:write``webshell:delete` |
| C2 | `c2:read``c2:write``c2:delete` |
| MCP | `mcp:read``mcp:execute``mcp:write``mcp:external:execute` |
| 知识库 | `knowledge:read``knowledge:write``knowledge:delete` |
| Skills | `skills:read``skills:write``skills:delete` |
| Markdown Agents | `agents:read``agents:write``agents:delete` |
| AI 测试角色 | `roles:read``roles:write``roles:delete` |
| 工作流 | `workflow:read``workflow:execute``workflow:write``workflow:delete` |
| 系统配置 | `config:read``config:write` |
| 终端 | `terminal:execute` |
| 审计 | `audit:read``audit:delete` |
| RBAC | `rbac:read``rbac:write` |
| 通知 | `notification:read``notification:write` |
| 机器人 | `robot:read``robot:write` |
| 文件 | `files:read``files:write``files:delete` |
| 攻击链 | `attackchain:read``attackchain:write` |
| FOFA | `fofa:execute` |
| OpenAPI | `openapi:read` |
| 对话分组 | `group:read``group:write``group:delete` |
| 执行监控 | `monitor:read``monitor:write``monitor:delete` |
特殊权限说明:
- `agent:execute` 允许运行 Agent,但不自动允许本地文件系统、Shell 或任意配置命令。
- `agent:local-execute` 是本地执行兜底权限,应仅授予可信操作员。
- `mcp:execute` 用于访问认证后的 MCP HTTP 入口。
- `mcp:external:execute` 用于 Agent 调用外部 MCP 工具,当前还要求该权限的 Scope 为 `all`
- 管理外部 MCP 配置使用 `mcp:write`,与执行外部工具是两项权限。
- `robot:write` 管理机器人配置和测试入口;机器人聊天本身使用绑定用户或服务账号的业务权限。
---
## 五、资源 Scope
每个角色包含一个 Scope
| Scope | 含义 | 适合场景 |
|-------|------|----------|
| `all` | 访问该权限覆盖的所有资源 | 管理员、全局审计员 |
| `assigned` | 访问管理员指定的资源及系统支持的父资源继承范围 | 项目成员、指定资产操作员 |
| `own` | 以本人创建/归属资源为主;部分资源仍可通过显式授权或父资源关系访问 | 个人工作区、机器人独立身份 |
权限和 Scope 是绑定在一起计算的。一个用户可拥有多个角色,权限取并集;**同一个权限**的 Scope 取最宽值:
```text
all > assigned > own
```
示例:
- “全局审计”角色:`project:read` + `all`
- “个人项目编辑”角色:`project:write` + `own`
最终结果是:
```text
project:read → all
project:write → own
```
全局读取不会把无关的写权限扩大为全局写入。服务端授权必须使用 `ScopeFor(permission)`,不能使用用户的最宽总 Scope。
### 全局对象限制
部分对象是进程级共享定义,没有 owner。即使用户拥有 `write`,若该权限 Scope 不是 `all`,服务端仍拒绝修改。例如:
- AI 测试角色、Skills、Markdown Agents。
- 外部 MCP 配置。
- 机器人配置。
- 工作流定义。
- 知识库写操作(搜索除外)。
- HITL 全局白名单、默认审核方和审计策略。
- C2 Profile 写操作。
- 部分全局监控统计。
---
## 六、资源归属、显式授权与继承
可在“平台权限 → 成员详情 → 资源授权”中给用户分配资源。当前可直接选择的主要类型:
- 项目 `project`
- 对话 `conversation`
- 漏洞 `vulnerability`
- WebShell `webshell`
- 批量任务队列 `batch_task`
- C2 Listener `c2_listener`
一次批量授权最多 100 个资源。重复授权会跳过,不会创建重复记录。
部分子资源会继承父资源访问能力:
| 子资源 | 可继承的父资源 |
|--------|----------------|
| 对话 | 所属项目 |
| 漏洞 | 所属项目或关联对话 |
| 消息、过程详情、攻击链 | 所属对话 |
| C2 Session | Listener |
| C2 Task / 文件 /事件 | Session、Task 或 Listener 链路 |
因此,给用户授权一个项目,通常不需要再逐个授权该项目中的每条对话和漏洞。仍应以具体页面/API 的服务端检查结果为准。
---
## 七、Web 管理流程
### 7.1 创建用户
1. 使用管理员进入左侧 **平台权限**
2. 创建平台用户,设置用户名、显示名称、至少 8 位密码和启用状态。
3. 分配一个或多个平台角色。
4. 若角色 Scope 为 `assigned`,继续配置资源授权。
5. 让用户重新登录并在右上角用户菜单确认角色、权限数量和 Scope。
### 7.2 创建自定义角色
1. 新建平台角色并填写清晰的岗位名称与说明。
2. 选择 `all``assigned``own`
3. 只勾选岗位实际需要的权限。
4. 先用测试账号验证列表、详情、写操作、删除和 Agent 工具调用。
5. 再批量分配给正式用户。
系统角色不可编辑;复制其思路创建自定义角色即可。
### 7.3 权限变更何时生效
- 更新用户、密码、启用状态或角色后,该用户现有会话会被撤销,需要重新登录。
- 更新或删除自定义角色后,平台会撤销全部现有会话,所有用户需重新登录。
- 机器人每条消息实时解析绑定用户/服务账号权限;用户禁用或角色调整会立即影响下一条消息。
- 后台批量任务会根据任务 owner 重新解析 Principal,不应依赖创建任务时的前端状态。
---
## 八、推荐角色模板
以下是起点,不是固定策略。
### 只读项目成员
```text
Scope: assigned
dashboard:read
chat:read
project:read
vulnerability:read
files:read
attackchain:read
```
### 日常安全操作员
```text
Scope: assigned
agent:execute
chat:read / chat:write
project:read / project:write
vulnerability:read / vulnerability:write
tasks:read / tasks:write
files:read / files:write
hitl:read / hitl:write
```
只有确实需要本机命令时才增加 `agent:local-execute``terminal:execute`;需要删除时再增加对应 `:delete`
### 机器人专用账号
```text
Scope: own(独立工作区)或 assigned(指定项目)
agent:execute
chat:read / chat:write
按需增加 project、vulnerability、knowledge 等权限
```
也可使用 `admin` 作为机器人服务账号,但发送者仍需精确白名单;白名单内每个人都会获得完整权限并共享 admin 数据。详见[机器人指南](robot.md)。
---
## 九、Agent、MCP 与机器人边界
### Agent
HTTP 登录用户会被转换为不可变 Principal,传入单 Agent、多 Agent、工作流和工具执行上下文。长任务脱离 SSE 连接后仍保留身份,但不会因为前端按钮可见而绕过服务端权限。
### 内置 MCP
每个内置工具必须有显式授权策略。例如 WebShell 工具会同时检查 `webshell:read/write/delete``connection_id` 的资源访问;漏洞、项目、任务与 C2 工具也会检查参数指向的资源。
未登记授权策略的内置工具默认拒绝。普通本地/配置工具需要 `agent:local-execute`
### 外部 MCP
Agent 调用外部 MCP 工具需要 `mcp:external:execute`,且当前要求 Scope 为 `all`。这是因为外部服务的资源模型通常不受本地 owner/assignment 约束。
### 机器人
- `user_binding`:平台发送者绑定自己的 RBAC 用户。
- `service_account`:精确白名单发送者统一使用一个 RBAC 用户。
- 平台验签只做来源认证,不代替业务授权。
- 发送 `身份` / `whoami` 可检查实际 Principal。
---
## 十、RBAC API
所有接口使用:
```http
Authorization: Bearer <token>
```
管理接口需要 `rbac:read``rbac:write`,资源选择器需要 `rbac:write`
| 方法 | 路径 | 说明 |
|------|------|------|
| GET | `/api/rbac/me` | 当前用户、角色、权限、总 Scope 与逐权限 Scope |
| GET | `/api/rbac/metadata` | 权限目录、角色、角色权限和 Scope 列表 |
| GET/POST | `/api/rbac/users` | 列出/创建用户 |
| PUT/DELETE | `/api/rbac/users/:id` | 更新/删除用户 |
| GET/POST | `/api/rbac/roles` | 列出/创建角色 |
| PUT/DELETE | `/api/rbac/roles/:id` | 更新/删除自定义角色 |
| GET | `/api/rbac/resources?type=project&q=...` | 分页搜索可授权资源 |
| GET/POST | `/api/rbac/resource-assignments` | 列出/创建资源授权 |
| DELETE | `/api/rbac/resource-assignments/:id` | 撤销资源授权 |
创建用户示例:
```bash
curl -X POST http://localhost:8080/api/rbac/users \
-H "Authorization: Bearer $TOKEN" \
-H "Content-Type: application/json" \
-d '{
"username": "operator01",
"display_name": "安全操作员 01",
"password": "change-me-123",
"enabled": true,
"roles": ["operator"]
}'
```
创建自定义角色示例:
```bash
curl -X POST http://localhost:8080/api/rbac/roles \
-H "Authorization: Bearer $TOKEN" \
-H "Content-Type: application/json" \
-d '{
"name": "项目审计员",
"description": "只读查看指定项目",
"scope": "assigned",
"permissions": ["chat:read", "project:read", "vulnerability:read"]
}'
```
批量授权项目示例:
```bash
curl -X POST http://localhost:8080/api/rbac/resource-assignments \
-H "Authorization: Bearer $TOKEN" \
-H "Content-Type: application/json" \
-d '{
"user_id": "USER_ID",
"resource_type": "project",
"resource_ids": ["PROJECT_ID_1", "PROJECT_ID_2"]
}'
```
---
## 十一、审计与运维建议
- 使用个人账号管理平台,避免多人共享管理员密码。
- 自定义角色按岗位命名,描述中写明用途和负责人。
- 高风险权限单独审批:`terminal:execute``agent:local-execute``c2:write/delete``webshell:write/delete``rbac:write``config:write`
- 定期检查 `all` Scope 角色、服务账号、机器人白名单和长期未使用用户。
- 用户离职时先禁用账号,再撤销机器人绑定、资源授权和会话。
- 在日志审计中关注 `rbac/access_denied`、角色/用户变更、资源授权、机器人服务账号执行。
- 配合 HITL 控制高风险工具;RBAC 允许调用不等于可以跳过审批。
---
## 十二、常见问题
### 页面按钮看不到
检查用户是否有对应权限;前端会根据 `/api/rbac/me` 隐藏无权操作。直接调用 API 仍会由服务端拒绝。
### 有权限但返回“无权访问该资源”
检查该权限的 Scope,而不是只看用户总 Scope;再检查资源 owner、显式授权和父资源授权。
### 角色改了但用户仍是旧权限
角色变更会撤销会话。让用户重新登录;机器人下一条消息会重新解析权限。
### `write` 权限存在但全局配置仍被拒绝
全局对象写操作要求对应权限的 Scope 为 `all`。创建一个 `all` Scope 的专用管理角色,而不是扩大无关权限。
### Agent 能对话但不能运行命令
`agent:execute``agent:local-execute` 分离。按需授予本地执行权限,并结合 HITL、工具白名单和审计。
### 外部 MCP 提示需要 global scope
`mcp:external:execute` 外,该权限的 Scope 还必须为 `all`。外部 MCP 的数据边界不由本地资源授权自动保护。
+134 -19
View File
@@ -2,7 +2,7 @@
[English](../en-US/robot.md)
本文档说明如何通过**个人微信**、**钉钉**、**飞书**与 **企业微信** 与 CyberStrikeAI 对话(长连接 / 回调模式),在手机端即可使用,无需在服务器上打开网页。按下面步骤操作可避免常见弯路
本文档说明如何通过**个人微信、企业微信、钉钉、飞书、Telegram、Slack、Discord 和 QQ 机器人**使用 CyberStrikeAI,包括平台接入、RBAC 身份绑定、服务账号白名单、命令、验证与故障排查
---
@@ -15,10 +15,18 @@
- **个人微信**:点击「微信 / iLink」→「生成二维码并绑定」,用微信扫码确认(见 [3.4 个人微信](#34-个人微信-wechat--ilink)
- **钉钉**:勾选并填写 Client ID / Client Secret
- **飞书**:勾选并填写 App ID / App Secret
5. 点击 **应用配置** 保存微信扫码绑定成功后会**自动保存并启用**,一般无需再点
6. **重启 CyberStrikeAI 应用**(钉钉/飞书:只保存不重启,长连接不会建立;微信绑定成功后会自动重启连接,通常无需手动重启)
5. 点击 **应用配置** 保存;程序会自动重启对应机器人连接。微信扫码绑定成功后会自动保存并启用,一般无需再点
配置会写入 `config.yaml``robots` 段,也可在配置文件中直接编辑。**修改钉钉/飞书配置后必须重启,长连接才会生效。** 个人微信绑定成功后程序会自动写入 `robots.wechat` 并重启 iLink 长轮询。
配置会写入 `config.yaml``robots` 段,也可在配置文件中直接编辑。通过 Web 点击“应用配置”会自动重启对应连接;直接手工修改 `config.yaml` 时,需要重启 CyberStrikeAI 进程。个人微信绑定成功后程序会自动写入 `robots.wechat` 并重启 iLink 长轮询。
### 最短使用路径
平台连接成功后,不要直接开始普通对话,先完成业务身份配置:
- **多人使用**:机器人设置选择“逐用户绑定” → 每位用户在 Web 右上角头像生成绑定码 → 在机器人中发送绑定命令 → 发送 `身份` 验证。
- **只有自己使用**:先在机器人中发送 `身份` 复制发送者 ID → 机器人设置选择“专用服务账号” → User ID 填 `admin` 或其他 RBAC 用户 → 粘贴发送者白名单 → 应用配置 → 再次发送 `身份` 验证。
看到“鉴权状态:已授权”且“实际身份”正确后,即可直接发送普通文本与 AI 对话。
---
@@ -300,12 +308,88 @@
---
## 四、机器人命令
## 四、RBAC 鉴权与机器人命令
平台 Token、签名或长连接凭证只负责证明“消息来自该平台”;真正能执行哪些操作,由 CyberStrikeAI 的 RBAC 决定。每个机器人实例都必须选择一种业务鉴权模式。
### 4.1 应该选择哪种模式
| 使用场景 | 推荐模式 | 身份与数据范围 |
|----------|----------|----------------|
| 企业微信、飞书、钉钉、Slack 等多人共享机器人 | `user_binding` | 每个发送者绑定自己的 Web 用户,权限和数据互相隔离 |
| 个人微信、单人专属机器人、固定自动化入口 | `service_account` | 白名单发送者统一使用配置的 RBAC 用户,并共享该账号的数据 |
两种模式都会在**每条消息**执行前重新读取用户状态、角色、逐权限 Scope 和资源授权。用户被禁用或权限被收回后,下一条消息立即失效。
机器人执行普通 AI 对话至少需要以下权限:
```text
agent:execute
chat:read
chat:write
```
使用项目、角色、本地命令、WebShell、C2 或外部 MCP 时,还需按功能增加对应权限。删除对话需要 `chat:delete`
### 4.2 逐用户绑定模式(默认)
管理员操作:
1. 系统设置 → 机器人设置 → 选择平台。
2. 在“业务鉴权策略”中选择“逐用户绑定(`user_binding`)”。
3. 点击“应用配置”。
每位使用者操作:
1. 登录 CyberStrikeAI Web,点击右上角头像 → **绑定机器人账号**
2. 点击 **生成绑定码**,页面开始 5 分钟倒计时。
3. 在目标机器人中发送页面给出的完整命令,例如 `绑定 7C6E-BD4C`
4. 发送 `身份``whoami`,确认“鉴权状态:已授权”且“实际身份”是自己的 Web 用户。
绑定码仅保存哈希、只能使用一次。倒计时结束后前端会标记失效、禁用复制并刷新绑定列表;服务端也会拒绝过期码。重新生成会让此前尚未使用的旧码立即失效。用户可发送 `解绑`,或在 Web 绑定窗口中撤销绑定。
### 4.3 专用服务账号模式
1. 先让机器人正常连接平台。
2. 目标使用者向机器人发送 `身份` / `whoami`,复制返回的完整“发送者 ID”。个人微信的 ID 通常形如 `xxxx@im.wechat`;必须以命令返回值为准,不能用 `ilink_bot_id` 或配置中的 `ilink_user_id` 代替。
3. 系统设置 → 机器人设置 → 选择平台 → 业务鉴权策略选择“专用服务账号(`service_account`)”。
4. 填写服务账号的 **RBAC User ID**,不是显示名称。可以填写 `admin`;此时白名单发送者拥有完整平台权限,界面会显示红色风险提示。
5. 在“允许的平台发送者 ID”中每行填写一个完整 ID。必须精确匹配、区分大小写,不允许 `*` 通配符。
6. 点击“应用配置”,再发送 `身份` 确认“实际身份”和角色正确。
示例:
```yaml
robots:
wechat:
auth:
mode: service_account
service_user_id: admin
allowed_external_users:
- "o9cq806s32Sm2_kyOmkyaV7Rn1lU@im.wechat"
```
服务账号模式不接受 `绑定` / `解绑` 命令。多个白名单发送者会共享服务账号创建的对话、项目和其他 `own` 范围资源;若不希望共享,请使用逐用户绑定。
### 4.4 如何检查当前身份
发送:
```text
身份
```
返回内容包含:平台、真实发送者 ID、鉴权模式、鉴权状态、实际 RBAC 用户、RBAC User ID、平台角色、资源范围和有效权限数量。不在服务账号白名单中的发送者只会看到拒绝状态,不会看到服务账号详情。
### 4.5 命令列表
在任一已接入平台(钉钉/飞书/微信/Telegram/Slack/Discord/QQ 等)向机器人发送以下**文本命令**(仅支持文本):
| 命令 | 说明 |
|------|------|
| **绑定 \<绑定码\>** | 将当前平台发送者绑定到生成绑定码的 RBAC 用户 |
| **解绑** | 解除当前平台账号绑定;也可在 Web 端的绑定列表中撤销 |
| **身份****whoami** | 显示平台发送者 ID、鉴权模式、绑定状态及当前实际 RBAC 用户、角色和资源范围 |
| **帮助** | 显示命令帮助与说明 |
| **列表****对话列表** | 列出所有对话的标题与对话 ID |
| **切换 \<对话ID\>****继续 \<对话ID\>** | 指定对话 ID,后续消息在该对话中继续 |
@@ -320,6 +404,8 @@
除以上命令外,**直接输入任意文字**会作为用户消息发给 AI,与 Web 端对话逻辑一致(渗透测试/安全分析等)。
群聊消息按实际发送者鉴权,不使用群 ID 作为业务身份。服务账号模式除外:白名单发送者会明确共享配置的服务账号权限和资源。
---
## 五、如何使用(要 @ 机器人吗?)
@@ -338,14 +424,17 @@
1. CyberStrikeAI Web 端 → 系统设置 → 机器人设置 → **微信 / iLink****生成二维码并绑定**
2. 手机微信扫码确认(如需配对数字则在 Web 页填写)。
3. 绑定成功后,在手机微信私聊中发「帮助」测试。
3. 在手机微信私聊中发`身份`,复制发送者 ID。
4. 回到机器人设置选择 `user_binding`,或选择 `service_account` 并填写服务账号与发送者白名单。
5. 点击应用配置,在微信中再次发送 `身份`,确认实际 RBAC 身份后再发送普通消息。
**钉钉 / 飞书**
1. **在开放平台**:按第三节完成应用创建、凭证复制、机器人开通(钉钉务必选 **Stream 模式**)、权限与发布。
2. **在 CyberStrikeAI**:系统设置 → 机器人设置 → 勾选对应平台,粘贴 Client ID/App ID、Client Secret/App Secret → 点击 **应用配置**
3. **重启 CyberStrikeAI 进程**(否则长连接不会建立)
4. **在手机钉钉/飞书**:找到该机器人(单聊直接发,群聊需 @机器人),发「帮助」或任意内容测试。
3. **选择鉴权模式**:多人使用建议 `user_binding`;专用机器人配置服务账号与发送者白名单
4. **应用配置**Web 会自动重启对应连接。
5. **在手机钉钉/飞书**:找到机器人(单聊直接发,群聊需 @),先发 `身份` 检查鉴权,再发普通内容测试。
若发消息没反应,先看 **第九节排查****第十节常见弯路**
@@ -359,6 +448,11 @@
robots:
wechat: # 个人微信 iLink(扫码绑定后自动写入,一般无需手填)
enabled: true
auth:
mode: service_account
service_user_id: admin
allowed_external_users:
- "从身份命令复制的完整发送者 ID"
bot_token: "your_bot_token@im.bot:..."
ilink_bot_id: "your_bot_id@im.bot"
ilink_user_id: "your_user_id@im.wechat"
@@ -367,10 +461,14 @@ robots:
bot_agent: "CyberStrikeAI/1.0"
dingtalk:
enabled: true
auth:
mode: user_binding
client_id: "your_dingtalk_app_key"
client_secret: "your_dingtalk_app_secret"
lark:
enabled: true
auth:
mode: user_binding
app_id: "your_lark_app_id"
app_secret: "your_lark_app_secret"
verify_token: ""
@@ -400,7 +498,7 @@ robots:
sandbox: true
```
修改钉钉/飞书/企业微信/Telegram/Slack/Discord/QQ 配置后,点击 **应用配置** 会自动重启对应连接。个人微信扫码绑定成功后会自动写入并重启 iLink 连接。
每个平台的 `auth` 独立配置;省略时默认为 `user_binding`。修改配置后,在 Web 点击 **应用配置** 会自动重启对应连接;手工编辑 YAML 则需重启进程。个人微信扫码绑定成功后会自动写入并重启 iLink 连接。
---
@@ -408,20 +506,24 @@ robots:
在未安装钉钉或飞书时,可用**测试接口**验证机器人逻辑是否正常:
1. 先登录 CyberStrikeAI Web 端(保证有登录态)
2. 使用 curl 调用测试接口(需携带登录后的 Cookie
1. 使用具有全局 `robot:write` 权限的账号登录并获取 Bearer Token
2. 使用 curl 调用测试接口:
```bash
# 将 YOUR_COOKIE 替换为登录后获得的 Cookie(浏览器 F12 → 网络 → 任意请求 → 请求头中的 Cookie)
# 先登录;请按实际地址、用户名和密码修改
TOKEN=$(curl -s -X POST "http://localhost:8080/api/auth/login" \
-H "Content-Type: application/json" \
-d '{"username":"admin","password":"YOUR_PASSWORD"}' | jq -r '.token')
curl -X POST "http://localhost:8080/api/robot/test" \
-H "Content-Type: application/json" \
-H "Cookie: YOUR_COOKIE" \
-H "Authorization: Bearer $TOKEN" \
-d '{"platform":"dingtalk","user_id":"test_user","text":"帮助"}'
```
若返回 JSON 中含有 `"reply":"【CyberStrikeAI 机器人命令】..."`,说明命令处理正常。可再试 `"text":"列表"``"text":"当前"`
若返回 JSON 中含有 `"reply":"【CyberStrikeAI 机器人命令】..."`,说明命令处理正常。`帮助``版本``身份` 可在未绑定时执行;`列表``当前` 和普通 AI 消息会走真实 RBAC,测试用 `platform + user_id` 必须已经绑定,或与服务账号模式的发送者白名单精确匹配
接口说明:`POST /api/robot/test`(需登录),请求体 `{"platform":"可选","user_id":"可选","text":"必填"}`,响应 `{"reply":"回复内容"}`
接口说明:`POST /api/robot/test`(需全局 `robot:write`),请求体 `{"platform":"可选","user_id":"可选","text":"必填"}`,响应 `{"reply":"回复内容"}`该接口仅模拟机器人业务逻辑,不验证第三方平台签名或长连接。
---
@@ -458,8 +560,8 @@ curl -X POST "http://localhost:8080/api/robot/test" \
1. **Client ID / Client Secret 是否与开放平台完全一致**
从「凭证与基础信息」里**复制粘贴**,不要手打。注意数字 **0** 与字母 **o**、数字 **1** 与字母 **l**(例如 `ding9gf9tiozuc504aer` 中间是 **504** 不是 5o4)。
2. **是否在保存配置后重启了应用**
机器人长连接在**应用启动时**建立。在 Web 端点击应用配置」只写入配置文件,**必须重启 CyberStrikeAI 进程**后钉钉连接才会生效
2. **配置是否已应用**
在 Web 端修改后必须点击应用配置”,程序会自动重启对应连接。若直接手工编辑 `config.yaml`,则需重启 CyberStrikeAI 进程。
3. **看程序日志**
- 启动后应看到:`钉钉 Stream 正在连接…``钉钉 Stream 已启动(无需公网),等待收消息`
@@ -469,6 +571,15 @@ curl -X POST "http://localhost:8080/api/robot/test" \
4. **开放平台侧**
应用需已**发布**;在「机器人」能力中需开启**流式接入(Stream)** 用于接收消息(仅 HTTP 回调不够);权限管理里需有机器人接收、发送消息等权限。
### 9.3 收到回复但提示未绑定、白名单拒绝或权限不足
1. 先发送 `身份`,查看“鉴权模式”和“鉴权状态”。
2. `user_binding` 显示未绑定:在 Web 右上角头像中生成绑定码,并在同一个平台账号中发送完整绑定命令。绑定码过期或已经使用时需重新生成。
3. `service_account` 显示白名单拒绝:把 `身份` 返回的完整发送者 ID 原样加入当前平台的白名单,注意大小写、租户前缀和 `@im.wechat` 等后缀。
4. 显示实际身份但提示缺少权限:在“平台权限”检查该 RBAC 用户的角色。普通 AI 对话至少需要 `agent:execute``chat:read``chat:write`
5. 服务账号不存在或被禁用:应用配置会拒绝保存;恢复用户或选择其他已启用 RBAC 用户。
6. 使用 `admin` 时仍被拒绝:通常是发送者不在精确白名单中,而不是 admin 权限不足。
---
## 十、常见弯路(避免踩坑)
@@ -476,7 +587,11 @@ curl -X POST "http://localhost:8080/api/robot/test" \
- **个人微信与企业微信混淆**:个人微信走 `robots.wechat` + Web 扫码绑定;企业微信走 `robots.wecom` + 管理后台回调 URL,二者完全不同。
- **个人微信二维码过期**:二维码约 5 分钟有效,过期需重新生成,不要一直扫旧码。
- **用错了机器人类型**:在钉钉**群里**添加的「自定义」机器人(Webhook + 加签)**不能**用来做对话,本程序只支持**开放平台「企业内部应用」**里的机器人。
- **只保存没重启**:钉钉/飞书改完配置后必须**重启应用**,否则长连接不会建立(个人微信扫码绑定会自动重启连接)
- **改完没有点应用配置**:Web 中修改机器人配置后要点击“应用配置”;程序会自动重启对应连接。只有手工编辑 YAML 时才需要重启进程
- **把 Bot ID 当成发送者 ID**:服务账号白名单必须填写 `身份` 命令返回的“发送者 ID”,不要填 `ilink_bot_id``ilink_user_id`、群 ID 或显示昵称。
- **绑定码过期后继续使用**:绑定码 5 分钟有效且只能使用一次;新生成的码会让旧码立即失效。
- **服务账号误以为数据隔离**:同一服务账号白名单中的发送者共享该账号的对话和 `own` 范围资源;需要隔离时应使用 `user_binding`
- **admin 配置后任意人都能用**:不会。即使服务账号是 `admin`,发送者仍必须与白名单精确匹配;但白名单中的人将拥有完整权限。
- **Client ID 抄错**:开放平台是 `504` 就填 `504`,不要填成 `5o4`;尽量用复制粘贴。
- **钉钉只开了 HTTP 回调没开 Stream**:本程序通过 **Stream 长连接**收消息,开放平台里机器人的消息接收方式必须选 **Stream 模式**
- **应用没发布**:开放平台里修改了机器人或权限后,要在「版本管理与发布」里**发布新版本**,否则不生效。
@@ -487,5 +602,5 @@ curl -X POST "http://localhost:8080/api/robot/test" \
- 各平台均**仅处理文本消息**;其他类型(如图片、语音)会提示暂不支持或忽略。
- 个人微信仅支持**私聊**,不支持群聊 @ 机器人。
- 会话与 Web 端共用同一套对话数据:在机器人里创建的对话会在 Web 端「对话」列表中看到,反之亦然
- 会话与 Web 端共用同一套数据:`user_binding` 下归属于绑定用户;`service_account` 下归属于服务账号,并由白名单发送者共享
- 机器人执行与 **Eino 单/多代理** 相同逻辑(`ProcessMessageForRobot`,含进度回调与过程详情入库),仅不向客户端推送 SSE,最后一次性回复个人微信/钉钉/飞书/企业微信。默认 `robot_default_agent_mode: eino_single`
+1 -1
View File
@@ -48,7 +48,7 @@ config.yaml
```yaml
auth:
password: "<long-random-password>"
session_duration_hours: 12
server:
host: 127.0.0.1
port: 8080
+1 -1
View File
@@ -6,7 +6,7 @@
## 上线前必做
- 修改 `auth.password` 为长随机密码
- 首次部署后立即修改 `admin` 初始密码(Web 界面或平台权限 → 用户管理)
- 使用 HTTPS,或放在可信反向代理之后。
- 限制来源 IP、VPN 或堡垒机访问。
- 开启 `audit.enabled`
+2 -2
View File
@@ -16,9 +16,9 @@ CyberStrikeAI 面向授权安全测试场景,内置命令执行、MCP 工具
## 认证与会话
`auth.password` 是 Web 登录密码。建议:
Web 登录凭据由 RBAC 用户管理(默认内置 `admin` 账号)。建议:
- 首次部署立即修改默认密码
- 首次部署立即修改 `admin` 初始密码(控制台首次启动会输出)
- 使用长随机密码,并限制分享范围。
- 将服务放在内网、VPN、堡垒机或反向代理认证后面。
- 生产环境开启 HTTPS,避免明文传输 Cookie。
+3 -3
View File
@@ -23,12 +23,12 @@ https://127.0.0.1:8080/
检查:
- `config.yaml` 中的 `auth.password`
- 是否修改后未重启或未应用配置
- RBAC 用户密码是否正确(默认 `admin`;首次启动密码见控制台输出)
- 是否修改密码后旧会话已失效,需重新登录
- 浏览器 Cookie 是否异常,可尝试无痕窗口。
- 审计日志中是否有登录失败节流。
生产环境忘记密码时,需在服务器上修改 `config.yaml` 并重启服务
生产环境忘记密码时,需在服务器上通过 RBAC 用户管理重置,或直接更新数据库中的用户密码哈希
## 模型无响应
Binary file not shown.

Before

Width:  |  Height:  |  Size: 86 KiB

After

Width:  |  Height:  |  Size: 88 KiB

+19 -16
View File
@@ -91,15 +91,15 @@ func NewAgent(cfg *config.OpenAIConfig, agentCfg *config.AgentConfig, mcpServer
llmClient := openai.NewClient(cfg, httpClient, logger)
return &Agent{
openAIClient: llmClient,
config: cfg,
agentConfig: agentCfg,
mcpServer: mcpServer,
externalMCPMgr: externalMCPMgr,
logger: logger,
maxIterations: maxIterations,
toolNameMapping: make(map[string]string), // 初始化工具名称映射
toolDescriptionMode: "short",
openAIClient: llmClient,
config: cfg,
agentConfig: agentCfg,
mcpServer: mcpServer,
externalMCPMgr: externalMCPMgr,
logger: logger,
maxIterations: maxIterations,
toolNameMapping: make(map[string]string), // 初始化工具名称映射
toolDescriptionMode: "short",
}
}
@@ -120,6 +120,9 @@ type ChatMessage struct {
ToolName string `json:"tool_name,omitempty"`
// ReasoningContent 对应 OpenAI/DeepSeek 的 reasoning_content;思考模式 + 工具调用后续跑须回传(见 DeepSeek 文档)。
ReasoningContent string `json:"reasoning_content,omitempty"`
// ModelFacingTrace is runtime-only metadata: true means Content was already the exact
// payload seen at the model boundary and must be restored byte-for-byte.
ModelFacingTrace bool `json:"-"`
}
// MarshalJSON 自定义JSON序列化,将tool_calls中的arguments转换为JSON字符串
@@ -658,7 +661,7 @@ func (a *Agent) UpdateToolDescriptionMode(mode string) {
mode = "short"
}
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报错
@@ -780,25 +783,25 @@ func (a *Agent) ExecuteMCPToolForConversation(ctx context.Context, conversationI
}
// BeginLocalToolExecution 在非 CallTool 路径工具开始时写入 running 状态,供 MCP 监控页展示「执行中」。
func (a *Agent) BeginLocalToolExecution(toolName string, args map[string]interface{}) string {
func (a *Agent) BeginLocalToolExecution(ctx context.Context, toolName string, args map[string]interface{}) string {
if a == nil || a.mcpServer == nil {
return ""
}
return a.mcpServer.BeginToolExecution(toolName, args)
return a.mcpServer.BeginToolExecution(ctx, toolName, args)
}
// FinishLocalToolExecution 完成 BeginLocalToolExecution 创建的记录;executionID 为空时一次性写入已完成记录。
func (a *Agent) FinishLocalToolExecution(executionID, toolName string, args map[string]interface{}, resultText string, invokeErr error) string {
func (a *Agent) FinishLocalToolExecution(ctx context.Context, executionID, toolName string, args map[string]interface{}, resultText string, invokeErr error) string {
if a == nil || a.mcpServer == nil {
return ""
}
return a.mcpServer.FinishToolExecution(executionID, toolName, args, resultText, invokeErr)
return a.mcpServer.FinishToolExecution(ctx, executionID, toolName, args, resultText, invokeErr)
}
// RecordLocalToolExecution 将非 CallTool 路径完成的工具调用写入 MCP 监控库(与 CallTool 落库一致),返回 executionId。
// 用于 Eino filesystem execute 等场景,使助手气泡「渗透测试详情」与常规 MCP 一致可点进监控。
func (a *Agent) RecordLocalToolExecution(toolName string, args map[string]interface{}, resultText string, invokeErr error) string {
return a.FinishLocalToolExecution("", toolName, args, resultText, invokeErr)
func (a *Agent) RecordLocalToolExecution(ctx context.Context, toolName string, args map[string]interface{}, resultText string, invokeErr error) string {
return a.FinishLocalToolExecution(ctx, "", toolName, args, resultText, invokeErr)
}
// UpdateMCPExecutionDisplayResult 将监控库中的工具结果更新为送入模型的展示正文(reduction 后)。
+27
View File
@@ -5,6 +5,31 @@ import (
"strings"
)
const ModelFacingTraceVersionKey = "cyberstrike_model_facing_trace_version"
// IsModelFacingTraceJSON reports whether a persisted trace was produced from the final
// model-boundary state. Legacy traces have no version marker and require one-time migration.
func IsModelFacingTraceJSON(traceInputJSON string) bool {
var raw []map[string]interface{}
if err := json.Unmarshal([]byte(strings.TrimSpace(traceInputJSON)), &raw); err != nil {
return false
}
for _, msg := range raw {
extra, _ := msg["extra"].(map[string]interface{})
v, ok := extra[ModelFacingTraceVersionKey]
if !ok {
continue
}
switch n := v.(type) {
case float64:
return n >= 1
case int:
return n >= 1
}
}
return false
}
// ParseTraceMessages 解析落库的 last_react_inputOpenAI 风格 messages JSON 数组)。
func ParseTraceMessages(traceInputJSON string) ([]ChatMessage, error) {
traceInputJSON = strings.TrimSpace(traceInputJSON)
@@ -15,6 +40,7 @@ func ParseTraceMessages(traceInputJSON string) ([]ChatMessage, error) {
if err := json.Unmarshal([]byte(traceInputJSON), &raw); err != nil {
return nil, err
}
modelFacing := IsModelFacingTraceJSON(traceInputJSON)
out := make([]ChatMessage, 0, len(raw))
for _, msgMap := range raw {
msg := ChatMessage{}
@@ -23,6 +49,7 @@ func ParseTraceMessages(traceInputJSON string) ([]ChatMessage, error) {
continue
}
msg.Role = role
msg.ModelFacingTrace = modelFacing
if content, ok := msgMap["content"].(string); ok {
msg.Content = content
}
+19
View File
@@ -55,3 +55,22 @@ func TestMergeAssistantTraceOutput(t *testing.T) {
t.Fatalf("expected merged output, got %q", out[len(out)-1].Content)
}
}
func TestParseTraceMessagesMarksVersionedModelFacingTrace(t *testing.T) {
raw := `[{"role":"system","content":"s","extra":{"cyberstrike_model_facing_trace_version":1}},{"role":"user","content":"u"},{"role":"tool","content":"exact","tool_call_id":"c1"}]`
if !IsModelFacingTraceJSON(raw) {
t.Fatal("versioned trace not detected")
}
msgs, err := ParseTraceMessages(raw)
if err != nil {
t.Fatal(err)
}
for i, msg := range msgs {
if !msg.ModelFacingTrace {
t.Fatalf("message %d missing model-facing marker", i)
}
}
if IsModelFacingTraceJSON(`[{"role":"user","content":"legacy"}]`) {
t.Fatal("legacy trace incorrectly marked model-facing")
}
}
+193 -39
View File
@@ -8,6 +8,7 @@ import (
"fmt"
"net"
"net/http"
"net/url"
"os"
"path/filepath"
"strings"
@@ -16,6 +17,7 @@ import (
"cyberstrike-ai/internal/agent"
"cyberstrike-ai/internal/audit"
"cyberstrike-ai/internal/authctx"
"cyberstrike-ai/internal/c2"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/database"
@@ -64,6 +66,7 @@ type App struct {
slackCancel context.CancelFunc // Slack Socket Mode 取消函数
discordCancel context.CancelFunc // Discord Gateway 取消函数
qqCancel context.CancelFunc // QQ WebSocket 取消函数
alertCancel context.CancelFunc // 漏洞提醒持久化投递 worker
c2Manager *c2.Manager // C2 管理器(未启用 C2 时为 nil)
c2Watchdog *c2.SessionWatchdog // C2 会话看门狗
c2WatchdogCancel context.CancelFunc // 看门狗取消函数
@@ -81,13 +84,7 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
router := gin.Default()
// CORS中间件
router.Use(corsMiddleware())
// 认证管理器
authManager, err := security.NewAuthManager(cfg.Auth.Password, cfg.Auth.SessionDurationHours)
if err != nil {
return nil, fmt.Errorf("初始化认证失败: %w", err)
}
router.Use(corsMiddleware(cfg.Server.CORSAllowedOrigins))
// 初始化数据库
dbPath := cfg.Database.Path
@@ -105,21 +102,51 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
return nil, fmt.Errorf("初始化数据库失败: %w", err)
}
// 认证管理器(数据库初始化后挂载 RBAC)
authManager := security.NewAuthManager(cfg.Auth.SessionDurationHours)
if generatedPassword, err := authManager.AttachRBACStore(db); err != nil {
return nil, fmt.Errorf("初始化RBAC失败: %w", err)
} else if generatedPassword != "" {
config.PrintBootstrapAdminPassword(generatedPassword)
}
for platform, userID := range cfg.Robots.ServiceAccountUserIDs() {
user, userErr := db.GetRBACUserByID(userID)
if userErr != nil || !user.Enabled {
return nil, fmt.Errorf("robots.%s.auth.service_user_id 必须指向已启用的 RBAC 用户", platform)
}
}
auditSvc := audit.NewService(db, cfg, log.Logger)
audit.RegisterConversationCreateHook(auditSvc)
auditSvc.PurgeExpired()
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.PurgeExpired()
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.PurgeExpired()
hitl.StartRetentionLoop(hitlRetention, log.Logger)
// 创建MCP服务器(带数据库持久化)
mcpServer := mcp.NewServerWithStorage(log.Logger, db)
mcpServer.SetToolAuthorizer(mcpToolAuthorizer(db))
mcpServer.ConfigureHTTPToolCallTimeoutFromAgentMinutes(cfg.Agent.ToolTimeoutMinutes)
// 创建安全工具执行器
@@ -134,15 +161,9 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
registerProjectFactTools(mcpServer, db, 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服务器相同的存储)
externalMCPMgr := mcp.NewExternalMCPManagerWithStorage(log.Logger, db)
externalMCPMgr.SetToolAuthorizer(externalMCPToolAuthorizer())
if cfg.ExternalMCP.Servers != nil {
externalMCPMgr.LoadConfigs(&cfg.ExternalMCP)
// 启动所有启用的外部MCP客户端
@@ -168,7 +189,7 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
var knowledgeHandler *handler.KnowledgeHandler
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 {
// 确定知识库数据库路径
knowledgeDBPath := cfg.Database.KnowledgeDBPath
@@ -311,7 +332,7 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
}
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)
plantaskRel := strings.TrimSpace(cfg.MultiAgent.EinoMiddleware.PlantaskRelDir)
if plantaskRel == "" {
@@ -337,7 +358,7 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
}
markdownAgentsHandler := handler.NewMarkdownAgentsHandler(agentsDir)
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)
@@ -360,17 +381,21 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
attackChainHandler := handler.NewAttackChainHandler(db, &cfg.OpenAI, log.Logger)
vulnerabilityHandler := handler.NewVulnerabilityHandler(db, log.Logger)
projectHandler := handler.NewProjectHandler(db, log.Logger)
rbacHandler := handler.NewRBACHandler(db, log.Logger)
rbacHandler.SetAudit(auditSvc)
rbacHandler.SetAuthManager(authManager)
workflowHandler := handler.NewWorkflowHandler(db, log.Logger)
workflowHandler.SetAudit(auditSvc)
workflowHandler.SetRuntime(agent, cfg)
vulnerabilityHandler.SetAudit(auditSvc)
webshellHandler := handler.NewWebShellHandler(log.Logger, db)
webshellHandler.SetAudit(auditSvc)
chatUploadsHandler := handler.NewChatUploadsHandler(log.Logger)
chatUploadsHandler := handler.NewChatUploadsHandler(log.Logger, db)
chatUploadsHandler.SetAudit(auditSvc)
registerWebshellTools(mcpServer, db, webshellHandler, log.Logger)
registerWebshellManagementTools(mcpServer, db, webshellHandler, log.Logger)
configHandler := handler.NewConfigHandler(configPath, cfg, mcpServer, executor, agent, attackChainHandler, externalMCPMgr, log.Logger)
configHandler.SetDB(db)
configHandler.SetAudit(auditSvc)
agentHandler.SetHitlToolWhitelistSaver(configHandler)
agentHandler.SetHitlAuditStrategySaver(configHandler)
@@ -403,6 +428,8 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
conversationHandler.SetTaskStopper(agentHandler)
auditHandler := handler.NewAuditHandler(db, auditSvc, log.Logger)
robotHandler := handler.NewRobotHandler(cfg, db, agentHandler, log.Logger)
robotHandler.SetAudit(auditSvc)
db.SetVulnerabilityCreatedHook(robotHandler.NotifyNewVulnerability)
openAPIHandler := handler.NewOpenAPIHandler(db, log.Logger, conversationHandler, agentHandler)
// 创建 App 实例(部分字段稍后填充)
@@ -431,6 +458,9 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
}
// 飞书/钉钉长连接(无需公网),启用时在后台启动;后续前端应用配置时会通过 RestartRobotConnections 重启
app.startRobotConnections()
alertCtx, alertCancel := context.WithCancel(context.Background())
app.alertCancel = alertCancel
go robotHandler.RunVulnerabilityAlertWorker(alertCtx)
// 设置漏洞工具注册器(内置工具,必须设置)
vulnerabilityRegistrar := func() error {
@@ -535,6 +565,8 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
terminalHandler,
app.c2Handler,
auditHandler,
auditSvc,
rbacHandler,
mcpServer,
authManager,
openAPIHandler,
@@ -547,17 +579,30 @@ func New(cfg *config.Config, log *logger.Logger, configPath string) (*App, error
// mcpHandlerWithAuth 在鉴权通过后转发到 MCP 处理;若配置了 auth_header 则校验请求头,否则直接放行
func (a *App) mcpHandlerWithAuth(w http.ResponseWriter, r *http.Request) {
cfg := a.config.MCP
if cfg.AuthHeader != "" {
actual := []byte(r.Header.Get(cfg.AuthHeader))
expected := []byte(cfg.AuthHeaderValue)
if subtle.ConstantTimeCompare(actual, expected) != 1 {
a.logger.Logger.Debug("MCP 鉴权失败:header 缺失或值不匹配", zap.String("header", cfg.AuthHeader))
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusUnauthorized)
w.Write([]byte(`{"error":"unauthorized"}`))
if authHeader := strings.TrimSpace(r.Header.Get("Authorization")); len(authHeader) > 7 && strings.EqualFold(authHeader[:7], "Bearer ") {
if session, ok := a.auth.ValidateToken(strings.TrimSpace(authHeader[7:])); ok && session.Permissions["mcp:execute"] {
principal := authctx.NewPrincipalWithScopes(session.UserID, session.Username, session.Scope, session.Permissions, session.PermissionScopes)
a.mcpServer.HandleHTTP(w, r.WithContext(authctx.WithPrincipal(r.Context(), principal)))
return
}
}
if !cfg.AllowGlobalAccess || strings.TrimSpace(cfg.AuthHeader) == "" || strings.TrimSpace(cfg.AuthHeaderValue) == "" {
http.Error(w, "use an authorized user bearer token; global MCP service access is disabled", http.StatusUnauthorized)
return
}
if subtle.ConstantTimeCompare([]byte(r.Header.Get(cfg.AuthHeader)), []byte(cfg.AuthHeaderValue)) != 1 {
a.logger.Logger.Debug("MCP 鉴权失败:header 缺失或值不匹配", zap.String("header", cfg.AuthHeader))
w.Header().Set("Content-Type", "application/json")
w.WriteHeader(http.StatusUnauthorized)
w.Write([]byte(`{"error":"unauthorized"}`))
return
}
permissions := make(map[string]bool, len(security.PermissionCatalog))
for permission := range security.PermissionCatalog {
permissions[permission] = true
}
principal := authctx.NewPrincipal("service:mcp", "mcp-service", database.RBACScopeAll, permissions)
r = r.WithContext(authctx.WithPrincipal(r.Context(), principal))
a.mcpServer.HandleHTTP(w, r)
}
@@ -602,20 +647,20 @@ func (a *App) RunWithContext(ctx context.Context) error {
}
switch tlsMode {
case mainTLSFromFiles:
a.logger.Info("启动 HTTPS 主服务(已启用 HTTP/2 协商)",
a.logger.Debug("启动 HTTPS 主服务(已启用 HTTP/2 协商)",
zap.String("address", addr),
zap.String("cert", certFile),
)
case mainTLSInMemorySelfSigned:
a.logger.Info("启动 HTTPS 主服务(内存自签证书,仅测试;已启用 HTTP/2 协商)",
a.logger.Debug("启动 HTTPS 主服务(内存自签证书,仅测试;已启用 HTTP/2 协商)",
zap.String("address", addr),
)
}
if httpRedirect {
a.logger.Info("已启用 HTTP→HTTPS 自动跳转(同端口嗅探分流)", zap.String("address", addr))
a.logger.Debug("已启用 HTTP→HTTPS 自动跳转(同端口嗅探分流)", zap.String("address", addr))
}
} else {
a.logger.Info("启动 HTTP 主服务", zap.String("address", addr))
a.logger.Debug("启动 HTTP 主服务", zap.String("address", addr))
}
// 监听 context 取消,优雅关闭 HTTP 服务器
@@ -677,6 +722,10 @@ func (a *App) Shutdown() {
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 5*time.Second)
_ = einoobserve.ShutdownOtel(shutdownCtx)
shutdownCancel()
if a.alertCancel != nil {
a.alertCancel()
a.alertCancel = nil
}
// 停止钉钉/飞书长连接
a.robotMu.Lock()
@@ -818,6 +867,8 @@ func setupRoutes(
terminalHandler *handler.TerminalHandler,
c2Handler *handler.C2Handler,
auditHandler *handler.AuditHandler,
auditSvc *audit.Service,
rbacHandler *handler.RBACHandler,
mcpServer *mcp.Server,
authManager *security.AuthManager,
openAPIHandler *handler.OpenAPIHandler,
@@ -827,11 +878,15 @@ func setupRoutes(
// 认证相关路由
authRoutes := api.Group("/auth")
loginRL := security.NewRateLimiter(10, 1*time.Minute)
{
authRoutes.POST("/login", authHandler.Login)
authRoutes.POST("/login", security.RateLimitMiddleware(loginRL), authHandler.Login)
authRoutes.POST("/logout", security.AuthMiddleware(authManager), authHandler.Logout)
authRoutes.POST("/change-password", security.AuthMiddleware(authManager), authHandler.ChangePassword)
authRoutes.POST("/change-password", security.AuthMiddleware(authManager), security.RequirePermission("auth:self"), authHandler.ChangePassword)
authRoutes.GET("/validate", security.AuthMiddleware(authManager), authHandler.Validate)
authRoutes.POST("/robot-binding-code", security.AuthMiddleware(authManager), security.RequirePermission("auth:self"), robotHandler.CreateRobotBindingCode)
authRoutes.GET("/robot-bindings", security.AuthMiddleware(authManager), security.RequirePermission("auth:self"), robotHandler.ListMyRobotBindings)
authRoutes.DELETE("/robot-bindings/:id", security.AuthMiddleware(authManager), security.RequirePermission("auth:self"), robotHandler.DeleteMyRobotBinding)
}
// 机器人回调(无需登录,供企业微信/钉钉/飞书服务器调用)
@@ -848,7 +903,31 @@ func setupRoutes(
protected := api.Group("")
protected.Use(security.AuthMiddleware(authManager))
protected.Use(security.RBACMiddlewareWithDenyHook(app.db, func(c *gin.Context, reason, permission string) {
if auditSvc != nil {
auditSvc.Record(c, audit.Entry{
Level: "warn", Category: "rbac", Action: "access_denied", Result: "failure",
Message: "RBAC 拒绝访问", ResourceType: "route", ResourceID: c.FullPath(),
Detail: map[string]interface{}{"reason": reason, "permission": permission, "method": c.Request.Method},
})
}
}))
{
protected.GET("/rbac/me", rbacHandler.Me)
protected.GET("/rbac/metadata", rbacHandler.Metadata)
protected.GET("/rbac/users", rbacHandler.ListUsers)
protected.POST("/rbac/users", rbacHandler.CreateUser)
protected.PUT("/rbac/users/:id", rbacHandler.UpdateUser)
protected.DELETE("/rbac/users/:id", rbacHandler.DeleteUser)
protected.GET("/rbac/roles", rbacHandler.ListRoles)
protected.POST("/rbac/roles", rbacHandler.CreateRole)
protected.PUT("/rbac/roles/:id", rbacHandler.UpdateRole)
protected.DELETE("/rbac/roles/:id", rbacHandler.DeleteRole)
protected.GET("/rbac/resource-assignments", rbacHandler.ListResourceAssignments)
protected.GET("/rbac/resources", rbacHandler.ListAssignableResources)
protected.POST("/rbac/resource-assignments", rbacHandler.AssignResource)
protected.DELETE("/rbac/resource-assignments/:id", rbacHandler.DeleteResourceAssignment)
// 机器人测试(需登录):POST /api/robot/testbody: {"platform":"dingtalk","user_id":"test","text":"帮助"},用于验证机器人逻辑
protected.POST("/robot/test", robotHandler.HandleRobotTest)
@@ -918,6 +997,7 @@ func setupRoutes(
protected.GET("/conversations", conversationHandler.ListConversations)
protected.GET("/conversations/:id", conversationHandler.GetConversation)
protected.GET("/messages/:id/process-details", conversationHandler.GetMessageProcessDetails)
protected.GET("/process-details/:id", conversationHandler.GetProcessDetail)
protected.PUT("/conversations/:id", conversationHandler.UpdateConversation)
protected.PUT("/conversations/:id/project", conversationHandler.SetConversationProject)
protected.DELETE("/conversations/:id", conversationHandler.DeleteConversation)
@@ -1135,6 +1215,8 @@ func setupRoutes(
protected.DELETE("/vulnerabilities/batch", vulnerabilityHandler.BatchDeleteVulnerabilities)
protected.GET("/vulnerabilities/filter-options", vulnerabilityHandler.GetVulnerabilityFilterOptions)
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.POST("/vulnerabilities", vulnerabilityHandler.CreateVulnerability)
protected.PUT("/vulnerabilities/:id", vulnerabilityHandler.UpdateVulnerability)
@@ -1244,6 +1326,11 @@ func setupRoutes(
protected.POST("/workflows/runs/:runId/resume", workflowHandler.ResumeRun)
protected.POST("/workflows/validate", workflowHandler.Validate)
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/:id", workflowHandler.Get)
protected.POST("/workflows", workflowHandler.Create)
@@ -1444,7 +1531,7 @@ func registerWebshellTools(mcpServer *mcp.Server, db *database.DB, webshellHandl
}
mcpServer.RegisterTool(writeTool, writeHandler)
logger.Info("WebShell 工具注册成功")
logger.Debug("WebShell 工具注册成功")
}
// registerWebshellManagementTools 注册 WebShell 连接管理 MCP 工具
@@ -1465,7 +1552,13 @@ func registerWebshellManagementTools(mcpServer *mcp.Server, db *database.DB, web
},
}
listHandler := func(ctx context.Context, args map[string]interface{}) (*mcp.ToolResult, error) {
connections, err := db.ListWebshellConnections()
connections := []database.WebShellConnection{}
var err error
if principal, ok := authctx.PrincipalFromContext(ctx); ok {
connections, err = db.ListWebshellConnectionsForAccess(principal.UserID, principal.ScopeFor("webshell:read"))
} else {
return &mcp.ToolResult{Content: []mcp.Content{{Type: "text", Text: "缺少认证身份"}}, IsError: true}, nil
}
if err != nil {
return &mcp.ToolResult{
Content: []mcp.Content{{Type: "text", Text: "获取连接列表失败: " + err.Error()}},
@@ -1580,6 +1673,10 @@ func registerWebshellManagementTools(mcpServer *mcp.Server, db *database.DB, web
IsError: true,
}, nil
}
if principal, ok := authctx.PrincipalFromContext(ctx); ok {
_ = db.SetResourceOwner("webshell", conn.ID, principal.UserID)
_ = db.AssignResourceToUser(principal.UserID, "webshell", conn.ID)
}
return &mcp.ToolResult{
Content: []mcp.Content{{
@@ -1805,7 +1902,7 @@ func registerWebshellManagementTools(mcpServer *mcp.Server, db *database.DB, web
}
mcpServer.RegisterTool(testTool, testHandler)
logger.Info("WebShell 管理工具注册成功")
logger.Debug("WebShell 管理工具注册成功")
}
// initializeKnowledge 初始化知识库组件(用于动态初始化)
@@ -1972,13 +2069,36 @@ func initializeKnowledge(
return knowledgeHandler, nil
}
// corsMiddleware CORS中间件
func corsMiddleware() gin.HandlerFunc {
// corsMiddleware allows same-origin requests, valid Chromium extension
// 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) {
c.Writer.Header().Set("Access-Control-Allow-Origin", "*")
c.Writer.Header().Set("Access-Control-Allow-Credentials", "true")
origin := strings.TrimSpace(c.GetHeader("Origin"))
if origin != "" {
c.Writer.Header().Add("Vary", "Origin")
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"})
return
}
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-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-Max-Age", "600")
if c.Request.Method == "OPTIONS" {
c.AbortWithStatus(204)
@@ -1988,3 +2108,37 @@ func corsMiddleware() gin.HandlerFunc {
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
}
+30 -13
View File
@@ -4,11 +4,13 @@ import (
"context"
"encoding/json"
"fmt"
"path/filepath"
"strconv"
"strings"
"time"
"cyberstrike-ai/internal/agent"
"cyberstrike-ai/internal/authctx"
"cyberstrike-ai/internal/c2"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/mcp"
@@ -29,7 +31,7 @@ func registerC2Tools(mcpServer *mcp.Server, c2Manager *c2.Manager, logger *zap.L
registerC2EventTool(mcpServer, c2Manager, logger)
registerC2ProfileTool(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) {
@@ -66,16 +68,16 @@ tcp_reverse 默认仅接受 CSB1 加密 BeaconAES-GCM + ImplantToken)才登
InputSchema: map[string]interface{}{
"type": "object",
"properties": map[string]interface{}{
"action": map[string]interface{}{"type": "string", "description": "操作: list/get/create/update/start/stop/delete", "enum": []string{"list", "get", "create", "update", "start", "stop", "delete"}},
"listener_id": map[string]interface{}{"type": "string", "description": "监听器 IDget/update/start/stop/delete 需要)"},
"name": map[string]interface{}{"type": "string", "description": "监听器名称(create/update"},
"type": map[string]interface{}{"type": "string", "description": "监听器类型(create", "enum": []string{"tcp_reverse", "http_beacon", "https_beacon", "websocket"}},
"action": map[string]interface{}{"type": "string", "description": "操作: list/get/create/update/start/stop/delete", "enum": []string{"list", "get", "create", "update", "start", "stop", "delete"}},
"listener_id": map[string]interface{}{"type": "string", "description": "监听器 IDget/update/start/stop/delete 需要)"},
"name": map[string]interface{}{"type": "string", "description": "监听器名称(create/update"},
"type": map[string]interface{}{"type": "string", "description": "监听器类型(create", "enum": []string{"tcp_reverse", "http_beacon", "https_beacon", "websocket"}},
"bind_host": map[string]interface{}{"type": "string", "description": "绑定地址,默认 127.0.0.1;外网监听常用 0.0.0.0"},
"callback_host": map[string]interface{}{"type": "string", "description": "可选:植入端/Payload 回连主机名(公网 IP 或域名)。写入 config_json;生成 oneliner/beacon 时优先于 bind_host。update 时传入空字符串可清除"},
"bind_port": map[string]interface{}{"type": "integer", "description": fmt.Sprintf("绑定端口(create 必填)。须 ≠ %d(当前本服务 Web/API 端口,配置 server.port", webListenPort), "minimum": 1, "maximum": 65535},
"profile_id": map[string]interface{}{"type": "string", "description": "Malleable Profile ID"},
"remark": map[string]interface{}{"type": "string", "description": "备注"},
"config": map[string]interface{}{"type": "object", "description": "高级配置(beacon 路径/TLS/OPSEC 等),create/update 可用。tcp_reverse 可选 allow_legacy_shell:true 允许未加密经典 shell(默认 false"},
"bind_port": map[string]interface{}{"type": "integer", "description": fmt.Sprintf("绑定端口(create 必填)。须 ≠ %d(当前本服务 Web/API 端口,配置 server.port", webListenPort), "minimum": 1, "maximum": 65535},
"profile_id": map[string]interface{}{"type": "string", "description": "Malleable Profile ID"},
"remark": map[string]interface{}{"type": "string", "description": "备注"},
"config": map[string]interface{}{"type": "object", "description": "高级配置(beacon 路径/TLS/OPSEC 等),create/update 可用。tcp_reverse 可选 allow_legacy_shell:true 允许未加密经典 shell(默认 false"},
},
"required": []string{"action"},
},
@@ -85,7 +87,7 @@ tcp_reverse 默认仅接受 CSB1 加密 BeaconAES-GCM + ImplantToken)才登
switch action {
case "list":
listeners, err := m.DB().ListC2Listeners()
listeners, err := m.DB().ListC2ListenersForAccess(c2ToolAccess(ctx))
if err != nil {
return makeC2Result(nil, err)
}
@@ -128,6 +130,10 @@ tcp_reverse 默认仅接受 CSB1 加密 BeaconAES-GCM + ImplantToken)才登
if err != nil {
return makeC2Result(nil, err)
}
if principal, ok := authctx.PrincipalFromContext(ctx); ok {
_ = m.DB().SetResourceOwner("c2_listener", listener.ID, principal.UserID)
_ = m.DB().AssignResourceToUser(principal.UserID, "c2_listener", listener.ID)
}
implantToken := listener.ImplantToken
listener.EncryptionKey = ""
listener.ImplantToken = ""
@@ -264,7 +270,7 @@ func registerC2SessionTool(s *mcp.Server, m *c2.Manager, l *zap.Logger) {
if v, ok := params["suspicious"].(bool); ok && v {
filter.Suspicious = true
}
sessions, err := m.DB().ListC2Sessions(filter)
sessions, err := m.DB().ListC2SessionsForAccess(filter, c2ToolAccess(ctx))
return makeC2Result(map[string]interface{}{"sessions": sessions, "count": len(sessions)}, err)
case "get":
@@ -494,7 +500,7 @@ func registerC2TaskManageTool(s *mcp.Server, m *c2.Manager, l *zap.Logger) {
if limit := int(getFloat64(params, "limit")); limit > 0 {
filter.Limit = limit
}
tasks, err := m.DB().ListC2Tasks(filter)
tasks, err := m.DB().ListC2TasksForAccess(filter, c2ToolAccess(ctx))
return makeC2Result(map[string]interface{}{"tasks": tasks, "count": len(tasks)}, err)
case "cancel":
@@ -602,6 +608,9 @@ func registerC2PayloadTool(s *mcp.Server, m *c2.Manager, l *zap.Logger, webListe
if err != nil {
return makeC2Result(nil, err)
}
if principal, ok := authctx.PrincipalFromContext(ctx); ok {
_ = m.DB().RecordC2PayloadArtifact(filepath.Base(result.OutputPath), result.PayloadID, result.ListenerID, principal.UserID)
}
return makeC2Result(map[string]interface{}{
"payload_id": result.PayloadID, "download_path": result.DownloadPath,
"os": result.OS, "arch": result.Arch, "size_bytes": result.SizeBytes,
@@ -648,11 +657,19 @@ func registerC2EventTool(s *mcp.Server, m *c2.Manager, l *zap.Logger) {
filter.Since = &t
}
}
events, err := m.DB().ListC2Events(filter)
events, err := m.DB().ListC2EventsForAccess(filter, c2ToolAccess(ctx))
return makeC2Result(map[string]interface{}{"events": events, "count": len(events)}, err)
})
}
func c2ToolAccess(ctx context.Context) database.RBACListAccess {
principal, ok := authctx.PrincipalFromContext(ctx)
if !ok {
return database.RBACListAccess{Scope: database.RBACScopeAssigned}
}
return database.RBACListAccess{UserID: principal.UserID, Scope: principal.ScopeFor("c2:read")}
}
// ============================================================================
// c2_profile — Malleable Profile 管理工具(新增)
// ============================================================================
+105
View File
@@ -0,0 +1,105 @@
package app
import (
"net/http"
"net/http/httptest"
"testing"
"github.com/gin-gonic/gin"
)
func TestCORSMiddlewareAllowsSameOriginAndRejectsForeignOrigin(t *testing.T) {
gin.SetMode(gin.TestMode)
router := gin.New()
router.Use(corsMiddleware(nil))
router.GET("/test", func(c *gin.Context) { c.Status(http.StatusNoContent) })
same := httptest.NewRequest(http.MethodGet, "http://app.example/test", nil)
same.Host = "app.example"
same.Header.Set("Origin", "http://app.example")
sameW := httptest.NewRecorder()
router.ServeHTTP(sameW, same)
if sameW.Code != http.StatusNoContent || sameW.Header().Get("Access-Control-Allow-Origin") != "http://app.example" {
t.Fatalf("same-origin response = %d, allow-origin=%q", sameW.Code, sameW.Header().Get("Access-Control-Allow-Origin"))
}
foreign := httptest.NewRequest(http.MethodGet, "http://app.example/test", nil)
foreign.Host = "app.example"
foreign.Header.Set("Origin", "https://evil.example")
foreignW := httptest.NewRecorder()
router.ServeHTTP(foreignW, foreign)
if 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)
}
})
}
}
+249
View File
@@ -0,0 +1,249 @@
package app
import (
"context"
"fmt"
"strings"
"cyberstrike-ai/internal/agent"
"cyberstrike-ai/internal/authctx"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/mcp"
"cyberstrike-ai/internal/mcp/builtin"
)
func mcpToolAuthorizer(db *database.DB) func(context.Context, string, map[string]interface{}) error {
return func(ctx context.Context, toolName string, args map[string]interface{}) error {
principal, ok := authctx.PrincipalFromContext(ctx)
if !ok {
return fmt.Errorf("missing authenticated principal")
}
require := func(permission string) error {
if !principal.HasPermission(permission) {
return fmt.Errorf("missing permission %s", permission)
}
return nil
}
resource := func(permission, resourceType, argument string) error {
if err := require(permission); err != nil {
return err
}
id := mcpAuthorizationString(args, argument)
if id == "" || db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), resourceType, id) {
return fmt.Errorf("no access to %s %s", resourceType, id)
}
return nil
}
switch toolName {
case builtin.ToolWebshellExec, builtin.ToolWebshellFileWrite:
return resource("webshell:write", "webshell", "connection_id")
case builtin.ToolWebshellFileList, builtin.ToolWebshellFileRead:
return resource("webshell:read", "webshell", "connection_id")
case builtin.ToolManageWebshellList:
return require("webshell:read")
case builtin.ToolManageWebshellAdd:
return require("webshell:write")
case builtin.ToolManageWebshellUpdate, builtin.ToolManageWebshellTest:
return resource("webshell:write", "webshell", "connection_id")
case builtin.ToolManageWebshellDelete:
return resource("webshell:delete", "webshell", "connection_id")
case builtin.ToolRecordVulnerability:
if err := require("vulnerability:write"); err != nil {
return err
}
conversationID := mcpAuthorizationString(args, "conversation_id")
if conversationID == "" {
conversationID = mcpAuthorizationConversationID(ctx)
}
if conversationID == "" || db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor("vulnerability:write"), "conversation", conversationID) {
return fmt.Errorf("no access to conversation %s", conversationID)
}
return nil
case builtin.ToolListVulnerabilities:
if err := require("vulnerability:read"); err != nil {
return err
}
conversationID := mcpAuthorizationConversationID(ctx)
if conversationID == "" || db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor("vulnerability:read"), "conversation", conversationID) {
return fmt.Errorf("no access to conversation %s", conversationID)
}
return nil
case builtin.ToolGetVulnerability:
return resource("vulnerability:read", "vulnerability", "id")
case builtin.ToolUpsertProjectFact, builtin.ToolDeprecateProjectFact, builtin.ToolRestoreProjectFact:
return authorizeProjectTool(ctx, principal, db, "project:write")
case builtin.ToolGetProjectFact, builtin.ToolListProjectFacts, builtin.ToolSearchProjectFacts:
return authorizeProjectTool(ctx, principal, db, "project:read")
case builtin.ToolListKnowledgeRiskTypes, builtin.ToolSearchKnowledgeBase:
return require("knowledge:read")
case builtin.ToolAnalyzeImage:
return require("agent:execute")
case builtin.ToolBatchTaskList:
return require("tasks:read")
case builtin.ToolBatchTaskGet:
return resource("tasks:read", "batch_task", "queue_id")
case builtin.ToolBatchTaskCreate:
if err := require("tasks:write"); err != nil {
return err
}
if projectID := mcpAuthorizationString(args, "project_id"); projectID != "" && (db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor("tasks:write"), "project", projectID)) {
return fmt.Errorf("no access to project %s", projectID)
}
return nil
case builtin.ToolBatchTaskDelete, builtin.ToolBatchTaskRemove:
return resource("tasks:delete", "batch_task", "queue_id")
case builtin.ToolBatchTaskStart, builtin.ToolBatchTaskRerun, builtin.ToolBatchTaskPause,
builtin.ToolBatchTaskUpdateMetadata, builtin.ToolBatchTaskUpdateSchedule,
builtin.ToolBatchTaskScheduleEnabled, builtin.ToolBatchTaskAdd, builtin.ToolBatchTaskUpdate:
return resource("tasks:write", "batch_task", "queue_id")
case builtin.ToolC2Listener:
return authorizeC2Action(principal, db, args, "c2_listener", "listener_id")
case builtin.ToolC2Session, builtin.ToolC2Task, builtin.ToolC2File:
if toolName == builtin.ToolC2File && mcpAuthorizationString(args, "action") == "get_result" {
return authorizeC2Action(principal, db, args, "c2_task", "task_id")
}
return authorizeC2Action(principal, db, args, "c2_session", "session_id")
case builtin.ToolC2TaskManage:
return authorizeC2Action(principal, db, args, "c2_task", "task_id")
case builtin.ToolC2Payload:
return resource("c2:write", "c2_listener", "listener_id")
case builtin.ToolC2Event:
if id := mcpAuthorizationString(args, "session_id"); id != "" {
return resource("c2:read", "c2_session", "session_id")
}
if principal.ScopeFor("c2:read") != database.RBACScopeAll {
return fmt.Errorf("unfiltered C2 event list requires global scope")
}
return require("c2:read")
case builtin.ToolC2Profile:
// Profiles are process-global and do not yet have an owner. Writes are
// therefore reserved for global scope; reads require c2:read.
if mcpAuthorizationString(args, "action") == "list" || mcpAuthorizationString(args, "action") == "get" {
return require("c2:read")
}
permission := "c2:write"
if mcpAuthorizationString(args, "action") == "delete" {
permission = "c2:delete"
}
if principal.ScopeFor(permission) != database.RBACScopeAll {
return fmt.Errorf("C2 profile mutation requires global scope")
}
if mcpAuthorizationString(args, "action") == "delete" {
return require("c2:delete")
}
return require("c2:write")
default:
if builtin.IsBuiltinTool(toolName) {
return fmt.Errorf("no authorization policy registered for builtin tool %s", toolName)
}
if principal.HasPermission("agent:local-execute") {
return nil
}
return fmt.Errorf("missing agent:local-execute")
}
}
}
func externalMCPToolAuthorizer() func(context.Context, string, map[string]interface{}) error {
return func(ctx context.Context, toolName string, _ map[string]interface{}) error {
principal, ok := authctx.PrincipalFromContext(ctx)
if !ok {
return fmt.Errorf("missing authenticated principal")
}
if !principal.HasPermission("mcp:external:execute") {
return fmt.Errorf("missing permission mcp:external:execute")
}
if principal.ScopeFor("mcp:external:execute") != database.RBACScopeAll {
return fmt.Errorf("external MCP invocation requires global scope")
}
if strings.TrimSpace(toolName) == "" {
return fmt.Errorf("missing external tool name")
}
return nil
}
}
func authorizeC2Action(principal authctx.Principal, db *database.DB, args map[string]interface{}, resourceType, argument string) error {
action := mcpAuthorizationString(args, "action")
permission := "c2:write"
if action == "list" || action == "get" || action == "get_result" || action == "wait" {
permission = "c2:read"
} else if action == "delete" || action == "delete_batch" {
permission = "c2:delete"
}
if !principal.HasPermission(permission) {
return fmt.Errorf("missing permission %s", permission)
}
id := mcpAuthorizationString(args, argument)
if action == "delete_batch" {
ids := mcpAuthorizationStrings(args, argument+"s")
if len(ids) == 0 {
return fmt.Errorf("missing resource identifiers %ss", argument)
}
for _, candidate := range ids {
if db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), resourceType, candidate) {
return fmt.Errorf("no access to %s %s", resourceType, candidate)
}
}
return nil
}
if id == "" {
if action == "create" || action == "list" {
return nil
}
return fmt.Errorf("missing resource identifier %s", argument)
}
if db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), resourceType, id) {
return fmt.Errorf("no access to %s %s", resourceType, id)
}
return nil
}
func mcpAuthorizationStrings(args map[string]interface{}, key string) []string {
values := []string{}
switch raw := args[key].(type) {
case []string:
for _, value := range raw {
if value = strings.TrimSpace(value); value != "" {
values = append(values, value)
}
}
case []interface{}:
for _, item := range raw {
if value, ok := item.(string); ok {
if value = strings.TrimSpace(value); value != "" {
values = append(values, value)
}
}
}
}
return values
}
func authorizeProjectTool(ctx context.Context, principal authctx.Principal, db *database.DB, permission string) error {
if !principal.HasPermission(permission) {
return fmt.Errorf("missing permission %s", permission)
}
conversationID := mcpAuthorizationConversationID(ctx)
if conversationID == "" || db == nil || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), "conversation", conversationID) {
return fmt.Errorf("no access to conversation %s", conversationID)
}
projectID, err := db.GetConversationProjectID(conversationID)
if err != nil || strings.TrimSpace(projectID) == "" || !db.UserCanAccessResource(principal.UserID, principal.ScopeFor(permission), "project", projectID) {
return fmt.Errorf("no access to project %s", projectID)
}
return nil
}
func mcpAuthorizationConversationID(ctx context.Context) string {
if id := strings.TrimSpace(agent.ConversationIDFromContext(ctx)); id != "" {
return id
}
return strings.TrimSpace(mcp.MCPConversationIDFromContext(ctx))
}
func mcpAuthorizationString(args map[string]interface{}, key string) string {
value, _ := args[key].(string)
return strings.TrimSpace(value)
}
+97
View File
@@ -0,0 +1,97 @@
package app
import (
"context"
"path/filepath"
"strings"
"testing"
"time"
"cyberstrike-ai/internal/authctx"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/mcp/builtin"
"cyberstrike-ai/internal/security"
"go.uber.org/zap"
)
func TestMCPToolAuthorizerEnforcesPermissionAndResource(t *testing.T) {
db, err := database.NewDB(filepath.Join(t.TempDir(), "mcp-authz.db"), zap.NewNop())
if err != nil {
t.Fatal(err)
}
defer db.Close()
user, err := db.CreateRBACUser("mcp-user", "MCP User", "hash", true, nil)
if err != nil {
t.Fatal(err)
}
for _, id := range []string{"ws_allowed", "ws_hidden"} {
if err := db.CreateWebshellConnection(&database.WebShellConnection{ID: id, URL: "http://127.0.0.1/" + id, Type: "php", Method: "post", CmdParam: "cmd", CreatedAt: time.Now()}); err != nil {
t.Fatal(err)
}
}
if err := db.AssignResourceToUser(user.ID, "webshell", "ws_allowed"); err != nil {
t.Fatal(err)
}
principal := authctx.NewPrincipal(user.ID, user.Username, database.RBACScopeAssigned, map[string]bool{"mcp:write": true, "webshell:write": true})
ctx := authctx.WithPrincipal(context.Background(), principal)
authorize := mcpToolAuthorizer(db)
if err := authorize(ctx, builtin.ToolWebshellExec, map[string]interface{}{"connection_id": "ws_allowed"}); err != nil {
t.Fatalf("allowed resource denied: %v", err)
}
if err := authorize(ctx, builtin.ToolWebshellExec, map[string]interface{}{"connection_id": "ws_hidden"}); err == nil {
t.Fatal("foreign webshell resource was allowed")
}
if err := authorize(ctx, builtin.ToolManageWebshellDelete, map[string]interface{}{"connection_id": "ws_allowed"}); err == nil {
t.Fatal("delete without webshell:delete was allowed")
}
}
func TestEveryBuiltinMCPToolHasExplicitAuthorizationPolicy(t *testing.T) {
db, err := database.NewDB(filepath.Join(t.TempDir(), "mcp-policy-inventory.db"), zap.NewNop())
if err != nil {
t.Fatal(err)
}
defer db.Close()
permissions := map[string]bool{}
for permission := range security.PermissionCatalog {
permissions[permission] = true
}
ctx := authctx.WithPrincipal(context.Background(), authctx.NewPrincipal("admin", "admin", database.RBACScopeAll, permissions))
authorize := mcpToolAuthorizer(db)
args := map[string]interface{}{
"action": "get", "connection_id": "x", "queue_id": "x", "listener_id": "x",
"session_id": "x", "task_id": "x", "id": "x", "conversation_id": "x",
}
for _, toolName := range builtin.GetAllBuiltinTools() {
err := authorize(ctx, toolName, args)
if err != nil && strings.Contains(err.Error(), "no authorization policy registered") {
t.Errorf("builtin tool %s has no explicit policy", toolName)
}
}
}
func TestExternalMCPRequiresDedicatedPermission(t *testing.T) {
authorize := externalMCPToolAuthorizer()
ctx := authctx.WithPrincipal(context.Background(), authctx.NewPrincipal("u1", "user", database.RBACScopeAssigned, map[string]bool{"agent:execute": true}))
if err := authorize(ctx, "server::tool", nil); err == nil {
t.Fatal("agent:execute alone authorized an external MCP tool")
}
ctx = authctx.WithPrincipal(context.Background(), authctx.NewPrincipal("u1", "user", database.RBACScopeAll, map[string]bool{"mcp:external:execute": true}))
if err := authorize(ctx, "server::tool", nil); err != nil {
t.Fatalf("dedicated external MCP permission rejected: %v", err)
}
}
func TestConfiguredCommandToolRequiresLocalExecutePermission(t *testing.T) {
authorize := mcpToolAuthorizer(nil)
agentOnly := authctx.WithPrincipal(context.Background(), authctx.NewPrincipal("u1", "user", database.RBACScopeAssigned, map[string]bool{"agent:execute": true}))
if err := authorize(agentOnly, "nmap_scan", nil); err == nil {
t.Fatal("agent:execute alone authorized a configured command tool")
}
local := authctx.WithPrincipal(context.Background(), authctx.NewPrincipal("u1", "user", database.RBACScopeAssigned, map[string]bool{"agent:local-execute": true}))
if err := authorize(local, "nmap_scan", nil); err != nil {
t.Fatalf("agent:local-execute rejected: %v", err)
}
}
+59
View File
@@ -0,0 +1,59 @@
package app
import (
"bytes"
"net/http"
"net/http/httptest"
"path/filepath"
"testing"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/mcp"
"cyberstrike-ai/internal/security"
"go.uber.org/zap"
)
func TestStandaloneMCPPrefersUserRBACAndDisablesGlobalTokenByDefault(t *testing.T) {
db, err := database.NewDB(filepath.Join(t.TempDir(), "mcp-http-auth.db"), zap.NewNop())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
auth := security.NewAuthManager(12)
if _, err := auth.AttachRBACStore(db); err != nil {
t.Fatal(err)
}
hash, err := security.HashPassword("admin-secret")
if err != nil {
t.Fatal(err)
}
if err := db.UpdateRBACAdminPassword(hash); err != nil {
t.Fatal(err)
}
token, _, err := auth.Authenticate("admin", "admin-secret")
if err != nil {
t.Fatal(err)
}
server := mcp.NewServer(zap.NewNop())
server.SetToolAuthorizer(mcpToolAuthorizer(db))
a := &App{config: &config.Config{MCP: config.MCPConfig{AuthHeader: "X-MCP-Token", AuthHeaderValue: "static-secret"}}, auth: auth, mcpServer: server}
body := []byte(`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`)
userReq := httptest.NewRequest(http.MethodPost, "/mcp", bytes.NewReader(body))
userReq.Header.Set("Authorization", "Bearer "+token)
userW := httptest.NewRecorder()
a.mcpHandlerWithAuth(userW, userReq)
if userW.Code != http.StatusOK {
t.Fatalf("user bearer status = %d: %s", userW.Code, userW.Body.String())
}
staticReq := httptest.NewRequest(http.MethodPost, "/mcp", bytes.NewReader(body))
staticReq.Header.Set("X-MCP-Token", "static-secret")
staticW := httptest.NewRecorder()
a.mcpHandlerWithAuth(staticW, staticReq)
if staticW.Code != http.StatusUnauthorized {
t.Fatalf("global static token status = %d, want 401", staticW.Code)
}
}
+2 -2
View File
@@ -90,7 +90,7 @@ func registerProjectFactTools(mcpServer *mcp.Server, db *database.DB, cfg *confi
"description": "可选:关联的漏洞记录 ID",
},
"links": map[string]interface{}{
"type": "array",
"type": "array",
"description": "可选:关系边(from → 当前 fact)。finding 至少 1 条 {from:target/*, type:discovered_on}finding 上记录 exploit 用 {from:exploit/*, type:exploits}。省略保留已有边;传 [] 清空全部关系边。",
"items": map[string]interface{}{
"type": "object",
@@ -357,7 +357,7 @@ func registerProjectFactTools(mcpServer *mcp.Server, db *database.DB, cfg *confi
})
if logger != nil {
logger.Info("项目黑板 MCP 工具注册成功")
logger.Debug("项目黑板 MCP 工具注册成功")
}
}
+11 -5
View File
@@ -6,6 +6,7 @@ import (
"strings"
"cyberstrike-ai/internal/agent"
"cyberstrike-ai/internal/authctx"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/mcp"
"cyberstrike-ai/internal/mcp/builtin"
@@ -189,7 +190,7 @@ func registerVulnerabilityTools(mcpServer *mcp.Server, db *database.DB, logger *
registerListVulnerabilitiesTool(mcpServer, db, logger)
registerGetVulnerabilityTool(mcpServer, db, logger)
if logger != nil {
logger.Info("漏洞 MCP 工具注册成功", zap.Strings("tools", []string{
logger.Debug("漏洞 MCP 工具注册成功", zap.Strings("tools", []string{
builtin.ToolRecordVulnerability,
builtin.ToolListVulnerabilities,
builtin.ToolGetVulnerability,
@@ -200,7 +201,7 @@ func registerVulnerabilityTools(mcpServer *mcp.Server, db *database.DB, logger *
func registerRecordVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, logger *zap.Logger) {
tool := mcp.Tool{
Name: builtin.ToolRecordVulnerability,
Description: "记录发现的漏洞详情到漏洞管理系统。必须按“仅看本记录即可复现”的标准填写:目标、触发点、前置条件、复现步骤、证据/POC、实际影响修复建议和复测方式。边渗透边记录:每验证出一条可复现漏洞后立即调用,勿等会话结束。记录前可先 list_vulnerabilities 避免重复。",
Description: "记录发现的漏洞详情到漏洞管理系统。必须按“仅看本记录即可复现”的标准填写:目标、漏洞类型、触发点、复现步骤、证据/POC、实际影响修复建议;前置条件与复测方式为推荐填写项。边渗透边记录:每验证出一条可复现漏洞后立即调用,勿等会话结束。记录前可先 list_vulnerabilities 避免重复。",
ShortDescription: "记录可复现的漏洞详情到漏洞管理系统",
InputSchema: map[string]interface{}{
"type": "object",
@@ -228,7 +229,7 @@ func registerRecordVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, log
},
"preconditions": map[string]interface{}{
"type": "string",
"description": "前置条件:登录状态、权限、账号、Header/Cookie、特定数据、网络位置、环境/版本等;无前置条件写“无”。",
"description": "前置条件(推荐填写):登录状态、权限、账号、Header/Cookie、特定数据、网络位置、环境/版本等;无前置条件写“无”。",
},
"reproduction_steps": map[string]interface{}{
"type": "string",
@@ -248,7 +249,7 @@ func registerRecordVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, log
},
"retest_notes": map[string]interface{}{
"type": "string",
"description": "复测方式:修复后如何验证漏洞已关闭,包括应返回的状态码、错误信息或访问控制结果。",
"description": "复测方式(推荐填写):修复后如何验证漏洞已关闭,包括应返回的状态码、错误信息或访问控制结果。",
},
},
"required": []string{"title", "description", "severity", "vulnerability_type", "target", "reproduction_steps", "evidence", "impact", "recommendation"},
@@ -281,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
}
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 := ""
@@ -313,6 +314,11 @@ func registerRecordVulnerabilityTool(mcpServer *mcp.Server, db *database.DB, log
}
return textResult(fmt.Sprintf("记录漏洞失败: %v", err), true), nil
}
if principal, ok := authctx.PrincipalFromContext(ctx); ok {
_ = db.SetResourceOwner("vulnerability", created.ID, principal.UserID)
_ = db.AssignResourceToUser(principal.UserID, "vulnerability", created.ID)
}
db.NotifyVulnerabilityCreated(created)
if logger != nil {
logger.Info("漏洞记录成功",
+6 -1
View File
@@ -63,7 +63,12 @@ func (s *Service) Record(c *gin.Context, e Entry) {
}
}
if strings.TrimSpace(e.Actor) == "" {
e.Actor = "admin"
if c != nil {
e.Actor = strings.TrimSpace(c.GetString(security.ContextUsernameKey))
}
if e.Actor == "" {
e.Actor = "admin"
}
}
maxDetail := s.cfg.Audit.MaxDetailBytesEffective()
detail := SanitizeDetail(e.Detail, maxDetail)
+72
View File
@@ -0,0 +1,72 @@
package authctx
import (
"context"
"strings"
)
// Principal is the immutable authorization identity propagated beyond the
// transport layer into Agent, MCP and background task contexts.
type Principal struct {
UserID string
Username string
Permissions map[string]bool
PermissionScopes map[string]string
Scope string
}
type principalContextKey struct{}
func NewPrincipal(userID, username, scope string, permissions map[string]bool) Principal {
return NewPrincipalWithScopes(userID, username, scope, permissions, nil)
}
func NewPrincipalWithScopes(userID, username, scope string, permissions map[string]bool, permissionScopes map[string]string) Principal {
permissionCopy := make(map[string]bool, len(permissions))
scopeCopy := make(map[string]string, len(permissionScopes))
for permission, allowed := range permissions {
if allowed {
permissionCopy[permission] = true
if permissionScope := strings.TrimSpace(permissionScopes[permission]); permissionScope != "" {
scopeCopy[permission] = permissionScope
}
}
}
return Principal{
UserID: strings.TrimSpace(userID), Username: strings.TrimSpace(username),
Scope: strings.TrimSpace(scope), Permissions: permissionCopy, PermissionScopes: scopeCopy,
}
}
func WithPrincipal(ctx context.Context, principal Principal) context.Context {
if ctx == nil {
ctx = context.Background()
}
if strings.TrimSpace(principal.UserID) == "" {
return ctx
}
return context.WithValue(ctx, principalContextKey{}, principal)
}
func PrincipalFromContext(ctx context.Context) (Principal, bool) {
if ctx == nil {
return Principal{}, false
}
principal, ok := ctx.Value(principalContextKey{}).(Principal)
return principal, ok && strings.TrimSpace(principal.UserID) != ""
}
func (p Principal) HasPermission(permission string) bool {
return p.Permissions[strings.TrimSpace(permission)]
}
// ScopeFor returns the scope attached to the permission that authorizes the
// current action. Falling back to Scope keeps explicit service principals and
// legacy callers compatible without reintroducing cross-role scope widening.
func (p Principal) ScopeFor(permission string) string {
permission = strings.TrimSpace(permission)
if scope := strings.TrimSpace(p.PermissionScopes[permission]); scope != "" {
return scope
}
return strings.TrimSpace(p.Scope)
}
+4 -4
View File
@@ -5,9 +5,9 @@ import (
"fmt"
"net"
"os"
"strconv"
"os/exec"
"path/filepath"
"strconv"
"strings"
"text/template"
@@ -173,15 +173,16 @@ func (b *PayloadBuilder) BuildBeacon(in PayloadBuilderInput) (*BuildResult, erro
}
// 交叉编译
payloadID := "p_" + strings.ReplaceAll(uuid.New().String(), "-", "")[:14]
binName := strings.TrimSpace(in.OutputName)
if binName == "" {
binName = fmt.Sprintf("beacon_%s_%s", goos, goarch)
binName = fmt.Sprintf("beacon_%s_%s_%s", goos, goarch, payloadID)
}
if goos == "windows" && !strings.HasSuffix(binName, ".exe") {
binName += ".exe"
}
binPath := filepath.Join(b.outputDir, binName)
if err := os.MkdirAll(b.outputDir, 0755); err != nil {
return nil, fmt.Errorf("mkdir output: %w", err)
}
@@ -214,7 +215,6 @@ func (b *PayloadBuilder) BuildBeacon(in PayloadBuilderInput) (*BuildResult, erro
return nil, fmt.Errorf("stat output: %w", err)
}
payloadID := "p_" + strings.ReplaceAll(uuid.New().String(), "-", "")[:14]
return &BuildResult{
PayloadID: payloadID,
ListenerID: listener.ID,
+317 -234
View File
@@ -2,15 +2,17 @@ package config
import (
"crypto/rand"
"encoding/base64"
"encoding/hex"
"encoding/json"
"fmt"
"io/fs"
"os"
"path/filepath"
"strconv"
"strings"
"cyberstrike-ai/internal/termout"
"gopkg.in/yaml.v3"
)
@@ -41,6 +43,20 @@ type Config struct {
Vision VisionConfig `yaml:"vision,omitempty" json:"vision,omitempty"`
}
type EnsureLocalConfigResult struct {
Created bool
ExamplePath string
}
const (
DefaultSummarizationUserIntentLedgerMaxRunes = 96000
DefaultSummarizationUserIntentLedgerEntryMaxRunes = 16000
DefaultLatestUserMessageMaxRunes = 48000
DefaultLatestUserMessageHeadRunes = 24000
DefaultLatestUserMessageTailRunes = 24000
DefaultSummarizationOutputReserveTokens = 8192
)
// ProjectConfig 项目黑板(跨对话共享事实)配置。
type ProjectConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"`
@@ -253,8 +269,20 @@ type MultiAgentEinoMiddlewareConfig struct {
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 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 *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 int `yaml:"summarization_user_intent_ledger_max_runes,omitempty" json:"summarization_user_intent_ledger_max_runes,omitempty"`
// SummarizationUserIntentLedgerEntryMaxRunes caps each user message entry inside the immutable user input ledger.
SummarizationUserIntentLedgerEntryMaxRunes int `yaml:"summarization_user_intent_ledger_entry_max_runes,omitempty" json:"summarization_user_intent_ledger_entry_max_runes,omitempty"`
// LatestUserMessageMaxRunes caps the current user turn inserted into model context; full text is persisted as an artifact when capped.
LatestUserMessageMaxRunes int `yaml:"latest_user_message_max_runes,omitempty" json:"latest_user_message_max_runes,omitempty"`
// LatestUserMessageHeadRunes keeps the head preview for an oversized current user turn.
LatestUserMessageHeadRunes int `yaml:"latest_user_message_head_runes,omitempty" json:"latest_user_message_head_runes,omitempty"`
// LatestUserMessageTailRunes keeps the tail preview for an oversized current user turn.
LatestUserMessageTailRunes int `yaml:"latest_user_message_tail_runes,omitempty" json:"latest_user_message_tail_runes,omitempty"`
// SummarizationRetryMaxAttempts 已废弃:summarization 与 run loop 共用 run_retry_max_attempts 及 isEinoTransientRunError。
SummarizationRetryMaxAttempts int `yaml:"summarization_retry_max_attempts,omitempty" json:"summarization_retry_max_attempts,omitempty"`
// PlanExecuteUserInputBudgetRatio caps planner/replanner/executor userInput prompt budget ratio (default 0.35).
@@ -295,6 +323,13 @@ func (c MultiAgentEinoMiddlewareConfig) SummarizationTriggerRatioEffective() flo
return v
}
func (c MultiAgentEinoMiddlewareConfig) SummarizationOutputReserveTokensEffective() int {
if c.SummarizationOutputReserveTokens > 0 {
return c.SummarizationOutputReserveTokens
}
return DefaultSummarizationOutputReserveTokens
}
func (c MultiAgentEinoMiddlewareConfig) SummarizationEmitInternalEventsEffective() bool {
if c.SummarizationEmitInternalEvents != nil {
return *c.SummarizationEmitInternalEvents
@@ -302,6 +337,41 @@ func (c MultiAgentEinoMiddlewareConfig) SummarizationEmitInternalEventsEffective
return true
}
func (c MultiAgentEinoMiddlewareConfig) SummarizationUserIntentLedgerMaxRunesEffective() int {
if c.SummarizationUserIntentLedgerMaxRunes > 0 {
return c.SummarizationUserIntentLedgerMaxRunes
}
return DefaultSummarizationUserIntentLedgerMaxRunes
}
func (c MultiAgentEinoMiddlewareConfig) SummarizationUserIntentLedgerEntryMaxRunesEffective() int {
if c.SummarizationUserIntentLedgerEntryMaxRunes > 0 {
return c.SummarizationUserIntentLedgerEntryMaxRunes
}
return DefaultSummarizationUserIntentLedgerEntryMaxRunes
}
func (c MultiAgentEinoMiddlewareConfig) LatestUserMessageMaxRunesEffective() int {
if c.LatestUserMessageMaxRunes > 0 {
return c.LatestUserMessageMaxRunes
}
return DefaultLatestUserMessageMaxRunes
}
func (c MultiAgentEinoMiddlewareConfig) LatestUserMessageHeadRunesEffective() int {
if c.LatestUserMessageHeadRunes > 0 {
return c.LatestUserMessageHeadRunes
}
return DefaultLatestUserMessageHeadRunes
}
func (c MultiAgentEinoMiddlewareConfig) LatestUserMessageTailRunesEffective() int {
if c.LatestUserMessageTailRunes > 0 {
return c.LatestUserMessageTailRunes
}
return DefaultLatestUserMessageTailRunes
}
func (c MultiAgentEinoMiddlewareConfig) PlanExecuteUserInputBudgetRatioEffective() float64 {
v := c.PlanExecuteUserInputBudgetRatio
if v <= 0 {
@@ -398,14 +468,19 @@ type MultiAgentSubConfig struct {
// MultiAgentPublic 返回给前端的精简信息(不含子代理指令全文)。
type MultiAgentPublic struct {
Enabled bool `json:"enabled"`
RobotDefaultAgentMode string `json:"robot_default_agent_mode,omitempty"`
BatchUseMultiAgent bool `json:"batch_use_multi_agent"`
SubAgentCount int `json:"sub_agent_count"`
Orchestration string `json:"orchestration,omitempty"`
PlanExecuteLoopMaxIterations int `json:"plan_execute_loop_max_iterations"`
ToolSearchAlwaysVisibleTools []string `json:"tool_search_always_visible_tools,omitempty"`
ToolSearchAlwaysVisibleEffectiveTools []string `json:"tool_search_always_visible_effective_tools,omitempty"`
Enabled bool `json:"enabled"`
RobotDefaultAgentMode string `json:"robot_default_agent_mode,omitempty"`
BatchUseMultiAgent bool `json:"batch_use_multi_agent"`
SubAgentCount int `json:"sub_agent_count"`
Orchestration string `json:"orchestration,omitempty"`
PlanExecuteLoopMaxIterations int `json:"plan_execute_loop_max_iterations"`
SummarizationUserIntentLedgerMaxRunes int `json:"summarization_user_intent_ledger_max_runes"`
SummarizationUserIntentLedgerEntryMaxRunes int `json:"summarization_user_intent_ledger_entry_max_runes"`
LatestUserMessageMaxRunes int `json:"latest_user_message_max_runes"`
LatestUserMessageHeadRunes int `json:"latest_user_message_head_runes"`
LatestUserMessageTailRunes int `json:"latest_user_message_tail_runes"`
ToolSearchAlwaysVisibleTools []string `json:"tool_search_always_visible_tools,omitempty"`
ToolSearchAlwaysVisibleEffectiveTools []string `json:"tool_search_always_visible_effective_tools,omitempty"`
}
// NormalizeAgentMode 解析代理模式(eino_single | deep | plan_execute | supervisor);空值默认 eino_single。
@@ -445,10 +520,15 @@ func NormalizeMultiAgentOrchestration(s string) string {
// MultiAgentAPIUpdate 设置页/API 仅更新多代理标量字段;写入 YAML 时不覆盖 sub_agents 等块。
type MultiAgentAPIUpdate struct {
Enabled bool `json:"enabled"`
RobotDefaultAgentMode string `json:"robot_default_agent_mode,omitempty"`
BatchUseMultiAgent bool `json:"batch_use_multi_agent"`
PlanExecuteLoopMaxIterations *int `json:"plan_execute_loop_max_iterations,omitempty"`
Enabled bool `json:"enabled"`
RobotDefaultAgentMode string `json:"robot_default_agent_mode,omitempty"`
BatchUseMultiAgent bool `json:"batch_use_multi_agent"`
PlanExecuteLoopMaxIterations *int `json:"plan_execute_loop_max_iterations,omitempty"`
SummarizationUserIntentLedgerMaxRunes *int `json:"summarization_user_intent_ledger_max_runes,omitempty"`
SummarizationUserIntentLedgerEntryMaxRunes *int `json:"summarization_user_intent_ledger_entry_max_runes,omitempty"`
LatestUserMessageMaxRunes *int `json:"latest_user_message_max_runes,omitempty"`
LatestUserMessageHeadRunes *int `json:"latest_user_message_head_runes,omitempty"`
LatestUserMessageTailRunes *int `json:"latest_user_message_tail_runes,omitempty"`
// 指针区分「JSON 未传该字段」与「传空数组要清空」;省略时不应覆盖 YAML 中的常驻工具白名单。
ToolSearchAlwaysVisibleTools *[]string `json:"tool_search_always_visible_tools,omitempty"`
}
@@ -468,14 +548,50 @@ type RobotsConfig struct {
// RobotWechatConfig 微信 iLink 机器人配置(个人微信 ClawBot / iLink 协议)
type RobotWechatConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"`
BotToken string `yaml:"bot_token,omitempty" json:"bot_token,omitempty"`
ILinkBotID string `yaml:"ilink_bot_id,omitempty" json:"ilink_bot_id,omitempty"`
ILinkUserID string `yaml:"ilink_user_id,omitempty" json:"ilink_user_id,omitempty"`
BaseURL string `yaml:"base_url,omitempty" json:"base_url,omitempty"` // 默认 https://ilinkai.weixin.qq.com
BotType string `yaml:"bot_type,omitempty" json:"bot_type,omitempty"` // get_bot_qrcode 参数,默认 3
BotAgent string `yaml:"bot_agent,omitempty" json:"bot_agent,omitempty"` // base_info.bot_agent
GetUpdatesBuf string `yaml:"get_updates_buf,omitempty" json:"get_updates_buf,omitempty"` // 长轮询游标(运行时)
Enabled bool `yaml:"enabled" json:"enabled"`
BotToken string `yaml:"bot_token,omitempty" json:"bot_token,omitempty"`
ILinkBotID string `yaml:"ilink_bot_id,omitempty" json:"ilink_bot_id,omitempty"`
ILinkUserID string `yaml:"ilink_user_id,omitempty" json:"ilink_user_id,omitempty"`
BaseURL string `yaml:"base_url,omitempty" json:"base_url,omitempty"` // 默认 https://ilinkai.weixin.qq.com
BotType string `yaml:"bot_type,omitempty" json:"bot_type,omitempty"` // get_bot_qrcode 参数,默认 3
BotAgent string `yaml:"bot_agent,omitempty" json:"bot_agent,omitempty"` // base_info.bot_agent
GetUpdatesBuf string `yaml:"get_updates_buf,omitempty" json:"get_updates_buf,omitempty"` // 长轮询游标(运行时)
Auth RobotAuthorizationConfig `yaml:"auth,omitempty" json:"auth,omitempty"`
}
const (
RobotAuthModeUserBinding = "user_binding"
RobotAuthModeServiceAccount = "service_account"
)
// RobotAuthorizationConfig controls how a verified platform sender becomes
// an RBAC principal. service_account is intentionally fail-closed unless an
// explicit non-admin service user and sender allowlist are both configured.
type RobotAuthorizationConfig struct {
Mode string `yaml:"mode,omitempty" json:"mode,omitempty"`
ServiceUserID string `yaml:"service_user_id,omitempty" json:"service_user_id,omitempty"`
AllowedExternalUsers []string `yaml:"allowed_external_users,omitempty" json:"allowed_external_users,omitempty"`
}
func (c RobotAuthorizationConfig) EffectiveMode() string {
mode := strings.ToLower(strings.TrimSpace(c.Mode))
if mode == "" {
return RobotAuthModeUserBinding
}
return mode
}
func (c RobotAuthorizationConfig) ExternalUserAllowed(externalUserID string) bool {
externalUserID = strings.TrimSpace(externalUserID)
if externalUserID == "" {
return false
}
for _, allowed := range c.AllowedExternalUsers {
if strings.TrimSpace(allowed) == externalUserID {
return true
}
}
return false
}
// RobotSessionConfig 机器人会话隔离策略
@@ -493,12 +609,13 @@ func (c RobotSessionConfig) StrictUserIdentityEnabled() bool {
// RobotWecomConfig 企业微信机器人配置
type RobotWecomConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"`
Token string `yaml:"token" json:"token"` // 回调 URL 校验 Token
EncodingAESKey string `yaml:"encoding_aes_key" json:"encoding_aes_key"` // EncodingAESKey
CorpID string `yaml:"corp_id" json:"corp_id"` // 企业 ID
Secret string `yaml:"secret" json:"secret"` // 应用 Secret
AgentID int64 `yaml:"agent_id" json:"agent_id"` // 应用 AgentId
Enabled bool `yaml:"enabled" json:"enabled"`
Token string `yaml:"token" json:"token"` // 回调 URL 校验 Token
EncodingAESKey string `yaml:"encoding_aes_key" json:"encoding_aes_key"` // EncodingAESKey
CorpID string `yaml:"corp_id" json:"corp_id"` // 企业 ID
Secret string `yaml:"secret" json:"secret"` // 应用 Secret
AgentID int64 `yaml:"agent_id" json:"agent_id"` // 应用 AgentId
Auth RobotAuthorizationConfig `yaml:"auth,omitempty" json:"auth,omitempty"`
}
// ValidateWecomConfig 校验企业微信机器人配置;启用时必须配置 token,否则回调无法防伪造。
@@ -514,55 +631,145 @@ func ValidateWecomConfig(w RobotWecomConfig) error {
// RobotDingtalkConfig 钉钉机器人配置
type RobotDingtalkConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"`
ClientID string `yaml:"client_id" json:"client_id"` // 应用 Key (AppKey)
ClientSecret string `yaml:"client_secret" json:"client_secret"` // 应用 Secret
AllowConversationIDFallback bool `yaml:"allow_conversation_id_fallback" json:"allow_conversation_id_fallback"` // sender_id 缺失时是否允许回退到会话 ID
Enabled bool `yaml:"enabled" json:"enabled"`
ClientID string `yaml:"client_id" json:"client_id"` // 应用 Key (AppKey)
ClientSecret string `yaml:"client_secret" json:"client_secret"` // 应用 Secret
AllowConversationIDFallback bool `yaml:"allow_conversation_id_fallback" json:"allow_conversation_id_fallback"` // sender_id 缺失时是否允许回退到会话 ID
Auth RobotAuthorizationConfig `yaml:"auth,omitempty" json:"auth,omitempty"`
}
// RobotLarkConfig 飞书机器人配置
type RobotLarkConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"`
AppID string `yaml:"app_id" json:"app_id"` // 应用 App ID
AppSecret string `yaml:"app_secret" json:"app_secret"` // 应用 App Secret
VerifyToken string `yaml:"verify_token" json:"verify_token"` // 事件订阅 Verification Token(可选)
AllowChatIDFallback bool `yaml:"allow_chat_id_fallback" json:"allow_chat_id_fallback"` // 用户 ID 缺失时是否允许回退到 chat_id
Enabled bool `yaml:"enabled" json:"enabled"`
AppID string `yaml:"app_id" json:"app_id"` // 应用 App ID
AppSecret string `yaml:"app_secret" json:"app_secret"` // 应用 App Secret
VerifyToken string `yaml:"verify_token" json:"verify_token"` // 事件订阅 Verification Token(可选)
AllowChatIDFallback bool `yaml:"allow_chat_id_fallback" json:"allow_chat_id_fallback"` // 用户 ID 缺失时是否允许回退到 chat_id
Auth RobotAuthorizationConfig `yaml:"auth,omitempty" json:"auth,omitempty"`
}
// RobotTelegramConfig Telegram 机器人配置(Bot API 长轮询)
type RobotTelegramConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"`
BotToken string `yaml:"bot_token" json:"bot_token"`
BotUsername string `yaml:"bot_username,omitempty" json:"bot_username,omitempty"` // 可选,用于群聊 @ 识别;留空则启动时 getMe
AllowGroupMessages bool `yaml:"allow_group_messages" json:"allow_group_messages"` // 群聊中仅响应 @ 机器人
UpdateOffset int64 `yaml:"update_offset,omitempty" json:"update_offset,omitempty"`
Enabled bool `yaml:"enabled" json:"enabled"`
BotToken string `yaml:"bot_token" json:"bot_token"`
BotUsername string `yaml:"bot_username,omitempty" json:"bot_username,omitempty"` // 可选,用于群聊 @ 识别;留空则启动时 getMe
AllowGroupMessages bool `yaml:"allow_group_messages" json:"allow_group_messages"` // 群聊中仅响应 @ 机器人
UpdateOffset int64 `yaml:"update_offset,omitempty" json:"update_offset,omitempty"`
Auth RobotAuthorizationConfig `yaml:"auth,omitempty" json:"auth,omitempty"`
}
// RobotSlackConfig Slack 机器人配置(Socket Mode,无需公网回调)
type RobotSlackConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"`
BotToken string `yaml:"bot_token" json:"bot_token"` // xoxb-
AppToken string `yaml:"app_token" json:"app_token"` // xapp-connections:write
Enabled bool `yaml:"enabled" json:"enabled"`
BotToken string `yaml:"bot_token" json:"bot_token"` // xoxb-
AppToken string `yaml:"app_token" json:"app_token"` // xapp-connections:write
Auth RobotAuthorizationConfig `yaml:"auth,omitempty" json:"auth,omitempty"`
}
// RobotDiscordConfig Discord 机器人配置(Gateway WebSocket
type RobotDiscordConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"`
BotToken string `yaml:"bot_token" json:"bot_token"`
AllowGuildMessages bool `yaml:"allow_guild_messages" json:"allow_guild_messages"` // 服务器频道中仅响应 @ 机器人
Enabled bool `yaml:"enabled" json:"enabled"`
BotToken string `yaml:"bot_token" json:"bot_token"`
AllowGuildMessages bool `yaml:"allow_guild_messages" json:"allow_guild_messages"` // 服务器频道中仅响应 @ 机器人
Auth RobotAuthorizationConfig `yaml:"auth,omitempty" json:"auth,omitempty"`
}
// RobotQQConfig QQ 机器人配置(QQ 开放平台 WebSocket
type RobotQQConfig struct {
Enabled bool `yaml:"enabled" json:"enabled"`
AppID string `yaml:"app_id" json:"app_id"`
ClientSecret string `yaml:"client_secret" json:"client_secret"`
Sandbox bool `yaml:"sandbox" json:"sandbox"` // 沙箱环境(上线前测试)
Enabled bool `yaml:"enabled" json:"enabled"`
AppID string `yaml:"app_id" json:"app_id"`
ClientSecret string `yaml:"client_secret" json:"client_secret"`
Sandbox bool `yaml:"sandbox" json:"sandbox"` // 沙箱环境(上线前测试)
Auth RobotAuthorizationConfig `yaml:"auth,omitempty" json:"auth,omitempty"`
}
func (c RobotsConfig) AuthorizationFor(platform string) RobotAuthorizationConfig {
switch strings.ToLower(strings.TrimSpace(platform)) {
case "wechat":
return c.Wechat.Auth
case "wecom":
return c.Wecom.Auth
case "dingtalk":
return c.Dingtalk.Auth
case "lark":
return c.Lark.Auth
case "telegram":
return c.Telegram.Auth
case "slack":
return c.Slack.Auth
case "discord":
return c.Discord.Auth
case "qq":
return c.QQ.Auth
default:
return RobotAuthorizationConfig{}
}
}
func ValidateRobotAuthorization(c RobotAuthorizationConfig, path string) error {
switch c.EffectiveMode() {
case RobotAuthModeUserBinding:
return nil
case RobotAuthModeServiceAccount:
serviceUserID := strings.TrimSpace(c.ServiceUserID)
if serviceUserID == "" {
return fmt.Errorf("%s.auth.service_user_id 不能为空", path)
}
if len(c.AllowedExternalUsers) == 0 {
return fmt.Errorf("%s.auth.allowed_external_users 至少配置一个真实发送者", path)
}
seen := map[string]bool{}
for _, userID := range c.AllowedExternalUsers {
userID = strings.TrimSpace(userID)
if userID == "" || userID == "*" {
return fmt.Errorf("%s.auth.allowed_external_users 不允许空值或通配符", path)
}
if seen[userID] {
return fmt.Errorf("%s.auth.allowed_external_users 包含重复用户", path)
}
seen[userID] = true
}
return nil
default:
return fmt.Errorf("%s.auth.mode 仅支持 user_binding 或 service_account", path)
}
}
func ValidateRobotsAuthorization(c RobotsConfig) error {
items := []struct {
path string
auth RobotAuthorizationConfig
}{
{"robots.wechat", c.Wechat.Auth}, {"robots.wecom", c.Wecom.Auth},
{"robots.dingtalk", c.Dingtalk.Auth}, {"robots.lark", c.Lark.Auth},
{"robots.telegram", c.Telegram.Auth}, {"robots.slack", c.Slack.Auth},
{"robots.discord", c.Discord.Auth}, {"robots.qq", c.QQ.Auth},
}
for _, item := range items {
if err := ValidateRobotAuthorization(item.auth, item.path); err != nil {
return err
}
}
return nil
}
func (c RobotsConfig) ServiceAccountUserIDs() map[string]string {
out := map[string]string{}
for _, platform := range []string{"wechat", "wecom", "dingtalk", "lark", "telegram", "slack", "discord", "qq"} {
auth := c.AuthorizationFor(platform)
if auth.EffectiveMode() == RobotAuthModeServiceAccount {
out[platform] = strings.TrimSpace(auth.ServiceUserID)
}
}
return out
}
type ServerConfig struct {
Host string `yaml:"host" json:"host"`
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 bool `yaml:"tls_enabled,omitempty" json:"tls_enabled,omitempty"`
// TLSCertPath / TLSKeyPath 非空时从 PEM 文件加载证书(生产环境推荐)。
@@ -580,11 +787,12 @@ type LogConfig struct {
}
type MCPConfig struct {
Enabled bool `yaml:"enabled"`
Host string `yaml:"host"`
Port int `yaml:"port"`
AuthHeader string `yaml:"auth_header,omitempty"` // 鉴权 header 名,留空表示不鉴权
AuthHeaderValue string `yaml:"auth_header_value,omitempty"` // 鉴权 header 值,需与请求中该 header 一致
Enabled bool `yaml:"enabled"`
Host string `yaml:"host"`
Port int `yaml:"port"`
AuthHeader string `yaml:"auth_header,omitempty"` // 可选的全局服务凭证 header;普通调用优先使用用户 Bearer Token
AuthHeaderValue string `yaml:"auth_header_value,omitempty"` // 全局服务凭证,仅 allow_global_access=true 时接受
AllowGlobalAccess bool `yaml:"allow_global_access,omitempty"` // 静态服务密钥是否映射为全局服务身份(默认关闭)
}
type OpenAIConfig struct {
@@ -800,11 +1008,7 @@ func normalizeHitlModeForPrompt(mode string) string {
}
type AuthConfig struct {
Password string `yaml:"password" json:"password"`
SessionDurationHours int `yaml:"session_duration_hours" json:"session_duration_hours"`
GeneratedPassword string `yaml:"-" json:"-"`
GeneratedPasswordPersisted bool `yaml:"-" json:"-"`
GeneratedPasswordPersistErr string `yaml:"-" json:"-"`
SessionDurationHours int `yaml:"session_duration_hours" json:"session_duration_hours"`
}
// MonitorConfig MCP 状态监控(tool_executions)保留策略。
@@ -964,23 +1168,6 @@ func Load(path string) (*Config, error) {
if cfg.Audit.MaxDetailBytes <= 0 {
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 != "" {
inlineTools := append([]ToolConfig(nil), cfg.Security.Tools...)
@@ -1036,170 +1223,64 @@ func Load(path string) (*Config, error) {
if err := ValidateWecomConfig(cfg.Robots.Wecom); err != nil {
return nil, err
}
if err := ValidateRobotsAuthorization(cfg.Robots); err != nil {
return nil, err
}
return &cfg, nil
}
func generateStrongPassword(length int) (string, error) {
if length <= 0 {
length = 24
func EnsureLocalConfig(path string) (EnsureLocalConfigResult, error) {
path = strings.TrimSpace(path)
if path == "" {
path = "config.yaml"
}
bytesLen := length
randomBytes := make([]byte, bytesLen)
if _, err := rand.Read(randomBytes); err != nil {
return "", err
if _, err := os.Stat(path); err == nil {
return EnsureLocalConfigResult{}, nil
} else if !os.IsNotExist(err) {
return EnsureLocalConfigResult{}, fmt.Errorf("检查配置文件失败: %w", 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 += " "
examplePath := filepath.Join(filepath.Dir(path), "config.example.yaml")
if _, err := os.Stat(examplePath); err != nil {
if os.IsNotExist(err) {
if alt := "config.example.yaml"; examplePath != alt {
if _, altErr := os.Stat(alt); altErr == nil {
examplePath = alt
} else {
return EnsureLocalConfigResult{}, fmt.Errorf("配置文件 %s 不存在,且未找到模板 %s", path, examplePath)
}
newLine += comment
} else {
return EnsureLocalConfigResult{}, fmt.Errorf("配置文件 %s 不存在,且未找到模板 %s", path, examplePath)
}
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] ⚠️ 无法自动写入配置文件中的密码。")
return EnsureLocalConfigResult{}, fmt.Errorf("检查配置模板失败: %w", err)
}
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("----------------------------------------------------------------")
data, err := os.ReadFile(examplePath)
if err != nil {
return EnsureLocalConfigResult{}, fmt.Errorf("读取配置模板失败: %w", err)
}
if dir := filepath.Dir(path); dir != "." && dir != "" {
if err := os.MkdirAll(dir, 0700); err != nil {
return EnsureLocalConfigResult{}, fmt.Errorf("创建配置目录失败: %w", err)
}
}
if err := os.WriteFile(path, data, fs.FileMode(0600)); err != nil {
return EnsureLocalConfigResult{}, fmt.Errorf("创建配置文件失败: %w", err)
}
return EnsureLocalConfigResult{
Created: true,
ExamplePath: examplePath,
}, nil
}
func PrintBootstrapAdminPassword(password string) {
termout.PrintBootstrapAdminCredentials(password)
}
// generateRandomToken 生成用于 MCP 鉴权的随机字符串(64 位十六进制)
@@ -1268,9 +1349,10 @@ func persistMCPAuth(path string, mcp *MCPConfig) error {
return os.WriteFile(path, []byte(strings.Join(lines, "\n")), 0644)
}
// EnsureMCPAuth 在 MCP 启用且 auth_header_value 为空时,自动生成随机密钥并写回配置
// EnsureMCPAuth only provisions the privileged static service credential when
// global service access was explicitly enabled.
func EnsureMCPAuth(path string, cfg *Config) error {
if !cfg.MCP.Enabled || strings.TrimSpace(cfg.MCP.AuthHeaderValue) != "" {
if !cfg.MCP.Enabled || !cfg.MCP.AllowGlobalAccess || strings.TrimSpace(cfg.MCP.AuthHeaderValue) != "" {
return nil
}
token, err := generateRandomToken()
@@ -1294,8 +1376,9 @@ func PrintMCPConfigJSON(mcp MCPConfig) {
hostForURL = "localhost"
}
url := fmt.Sprintf("http://%s:%d/mcp", hostForURL, mcp.Port)
headers := map[string]string{}
if mcp.AuthHeader != "" {
headers := map[string]string{"Authorization": "Bearer <USER_SESSION_TOKEN>"}
if mcp.AllowGlobalAccess && mcp.AuthHeader != "" {
delete(headers, "Authorization")
headers[mcp.AuthHeader] = mcp.AuthHeaderValue
}
serverEntry := map[string]interface{}{
@@ -1555,8 +1638,8 @@ func Default() *Config {
Output: "stdout",
},
MCP: MCPConfig{
Enabled: true,
Host: "0.0.0.0",
Enabled: false,
Host: "127.0.0.1",
Port: 8081,
},
OpenAI: OpenAIConfig{
+111 -48
View File
@@ -7,68 +7,71 @@ import (
"testing"
)
func TestPersistAuthPasswordQuotesYAMLSpecialCharacters(t *testing.T) {
func TestEnsureLocalConfigCreatesFromExample(t *testing.T) {
dir := t.TempDir()
examplePath := filepath.Join(dir, "config.example.yaml")
configPath := filepath.Join(dir, "config.yaml")
example := []byte(`auth:
session_duration_hours: 12
server:
host: 127.0.0.1
port: 8080
`)
if err := os.WriteFile(examplePath, example, 0644); err != nil {
t.Fatalf("write example: %v", err)
}
result, err := EnsureLocalConfig(configPath)
if err != nil {
t.Fatalf("EnsureLocalConfig: %v", err)
}
if !result.Created {
t.Fatal("Created = false, want true")
}
if result.ExamplePath != examplePath {
t.Fatalf("ExamplePath = %q, want %q", result.ExamplePath, examplePath)
}
cfg, err := Load(configPath)
if err != nil {
t.Fatalf("Load generated config: %v", err)
}
if cfg.Auth.SessionDurationHours != 12 {
t.Fatalf("SessionDurationHours = %d, want 12", cfg.Auth.SessionDurationHours)
}
second, err := EnsureLocalConfig(configPath)
if err != nil {
t.Fatalf("EnsureLocalConfig existing: %v", err)
}
if second.Created {
t.Fatal("Created = true for existing config, want false")
}
}
func TestLoadIgnoresLegacyAuthPasswordField(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 登录密码",
` password: "legacy-password"`,
" session_duration_hours: 12",
"log:",
" level: info",
"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)
}
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)
t.Fatalf("Load: %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)
if cfg.Auth.SessionDurationHours != 12 {
t.Fatalf("SessionDurationHours = %d, want 12", cfg.Auth.SessionDurationHours)
}
}
@@ -91,3 +94,63 @@ func TestHitlAuditModelEffectiveFallsBackToMainConfig(t *testing.T) {
t.Fatalf("expected audit model override, got %q", got.Model)
}
}
func TestSummarizationUserIntentLedgerRunesEffective(t *testing.T) {
var zero MultiAgentEinoMiddlewareConfig
if got := zero.SummarizationUserIntentLedgerMaxRunesEffective(); got != DefaultSummarizationUserIntentLedgerMaxRunes {
t.Fatalf("default ledger max runes = %d, want %d", got, DefaultSummarizationUserIntentLedgerMaxRunes)
}
if got := zero.SummarizationUserIntentLedgerEntryMaxRunesEffective(); got != DefaultSummarizationUserIntentLedgerEntryMaxRunes {
t.Fatalf("default ledger entry max runes = %d, want %d", got, DefaultSummarizationUserIntentLedgerEntryMaxRunes)
}
custom := MultiAgentEinoMiddlewareConfig{
SummarizationUserIntentLedgerMaxRunes: 12345,
SummarizationUserIntentLedgerEntryMaxRunes: 2345,
}
if got := custom.SummarizationUserIntentLedgerMaxRunesEffective(); got != 12345 {
t.Fatalf("custom ledger max runes = %d", got)
}
if got := custom.SummarizationUserIntentLedgerEntryMaxRunesEffective(); got != 2345 {
t.Fatalf("custom ledger entry max runes = %d", got)
}
}
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) {
var zero MultiAgentEinoMiddlewareConfig
if got := zero.LatestUserMessageMaxRunesEffective(); got != DefaultLatestUserMessageMaxRunes {
t.Fatalf("default latest user max runes = %d, want %d", got, DefaultLatestUserMessageMaxRunes)
}
if got := zero.LatestUserMessageHeadRunesEffective(); got != DefaultLatestUserMessageHeadRunes {
t.Fatalf("default latest user head runes = %d, want %d", got, DefaultLatestUserMessageHeadRunes)
}
if got := zero.LatestUserMessageTailRunesEffective(); got != DefaultLatestUserMessageTailRunes {
t.Fatalf("default latest user tail runes = %d, want %d", got, DefaultLatestUserMessageTailRunes)
}
custom := MultiAgentEinoMiddlewareConfig{
LatestUserMessageMaxRunes: 100,
LatestUserMessageHeadRunes: 40,
LatestUserMessageTailRunes: 60,
}
if got := custom.LatestUserMessageMaxRunesEffective(); got != 100 {
t.Fatalf("custom latest user max runes = %d", got)
}
if got := custom.LatestUserMessageHeadRunesEffective(); got != 40 {
t.Fatalf("custom latest user head runes = %d", got)
}
if got := custom.LatestUserMessageTailRunesEffective(); got != 60 {
t.Fatalf("custom latest user tail runes = %d", got)
}
}
+24
View File
@@ -43,3 +43,27 @@ func TestValidateWecomConfig(t *testing.T) {
})
}
}
func TestValidateRobotAuthorization(t *testing.T) {
tests := []struct {
name string
cfg RobotAuthorizationConfig
wantErr bool
}{
{name: "default user binding", cfg: RobotAuthorizationConfig{}, wantErr: false},
{name: "explicit user binding", cfg: RobotAuthorizationConfig{Mode: RobotAuthModeUserBinding}, wantErr: false},
{name: "service account", cfg: RobotAuthorizationConfig{Mode: RobotAuthModeServiceAccount, ServiceUserID: "svc-1", AllowedExternalUsers: []string{"t:x|u:y"}}, wantErr: false},
{name: "missing service user", cfg: RobotAuthorizationConfig{Mode: RobotAuthModeServiceAccount, AllowedExternalUsers: []string{"t:x|u:y"}}, wantErr: true},
{name: "admin allowed with exact sender", cfg: RobotAuthorizationConfig{Mode: RobotAuthModeServiceAccount, ServiceUserID: "admin", AllowedExternalUsers: []string{"t:x|u:y"}}, wantErr: false},
{name: "allowlist required", cfg: RobotAuthorizationConfig{Mode: RobotAuthModeServiceAccount, ServiceUserID: "svc-1"}, wantErr: true},
{name: "wildcard forbidden", cfg: RobotAuthorizationConfig{Mode: RobotAuthModeServiceAccount, ServiceUserID: "svc-1", AllowedExternalUsers: []string{"*"}}, wantErr: true},
{name: "unknown mode", cfg: RobotAuthorizationConfig{Mode: "open"}, wantErr: true},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
if err := ValidateRobotAuthorization(tt.cfg, "robots.lark"); (err != nil) != tt.wantErr {
t.Fatalf("ValidateRobotAuthorization() error=%v wantErr=%v", err, tt.wantErr)
}
})
}
}
+38 -28
View File
@@ -9,41 +9,47 @@ import (
// AuditLog platform operation audit record.
type AuditLog struct {
ID string `json:"id"`
CreatedAt time.Time `json:"createdAt"`
Level string `json:"level"`
Category string `json:"category"`
Action string `json:"action"`
Result string `json:"result"`
Actor string `json:"actor"`
SessionHint string `json:"sessionHint,omitempty"`
ClientIP string `json:"clientIp,omitempty"`
UserAgent string `json:"userAgent,omitempty"`
ResourceType string `json:"resourceType,omitempty"`
ResourceID string `json:"resourceId,omitempty"`
ResourceAvailable *bool `json:"resourceAvailable,omitempty"` // API-only: whether linked resource still exists
Message string `json:"message"`
Detail map[string]interface{} `json:"detail,omitempty"`
ID string `json:"id"`
CreatedAt time.Time `json:"createdAt"`
Level string `json:"level"`
Category string `json:"category"`
Action string `json:"action"`
Result string `json:"result"`
Actor string `json:"actor"`
SessionHint string `json:"sessionHint,omitempty"`
ClientIP string `json:"clientIp,omitempty"`
UserAgent string `json:"userAgent,omitempty"`
ResourceType string `json:"resourceType,omitempty"`
ResourceID string `json:"resourceId,omitempty"`
ResourceAvailable *bool `json:"resourceAvailable,omitempty"` // API-only: whether linked resource still exists
Message string `json:"message"`
Detail map[string]interface{} `json:"detail,omitempty"`
}
// ListAuditLogsFilter query parameters.
type ListAuditLogsFilter struct {
Level string
Category string
Action string
Result string
Query string
ResourceType string
ResourceID string
Since *time.Time
Until *time.Time
Limit int
Offset int
Actor string
Level string
Category string
Action string
Result string
Query string
ResourceType string
ResourceID string
RelatedUserID string
Since *time.Time
Until *time.Time
Limit int
Offset int
}
func buildAuditLogsWhere(filter ListAuditLogsFilter) (string, []interface{}) {
conditions := []string{"1=1"}
args := []interface{}{}
if filter.Actor != "" {
conditions = append(conditions, "actor = ?")
args = append(args, filter.Actor)
}
if filter.Level != "" {
conditions = append(conditions, "level = ?")
args = append(args, filter.Level)
@@ -68,6 +74,10 @@ func buildAuditLogsWhere(filter ListAuditLogsFilter) (string, []interface{}) {
conditions = append(conditions, "resource_id = ?")
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 {
conditions = append(conditions, sqliteEpochGE("created_at", ">="))
args = append(args, formatSQLiteUTC(*filter.Since))
@@ -78,8 +88,8 @@ func buildAuditLogsWhere(filter ListAuditLogsFilter) (string, []interface{}) {
}
if q := strings.TrimSpace(filter.Query); q != "" {
like := "%" + q + "%"
conditions = append(conditions, "(message LIKE ? OR resource_id LIKE ? OR action LIKE ? OR category LIKE ?)")
args = append(args, like, like, like, like)
conditions = append(conditions, "(message LIKE ? OR resource_id LIKE ? OR action LIKE ? OR category LIKE ? OR detail_json LIKE ?)")
args = append(args, like, like, like, like, like)
}
return strings.Join(conditions, " AND "), args
}
+13
View File
@@ -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) {
root, err := os.Getwd()
if err != nil {
+49 -1
View File
@@ -137,7 +137,7 @@ func (db *DB) GetBatchQueue(queueID string) (*BatchTaskQueueRow, error) {
// GetAllBatchQueues 获取所有批量任务队列
func (db *DB) GetAllBatchQueues() ([]*BatchTaskQueueRow, error) {
rows, err := db.Query(
"SELECT "+batchQueueSelectColumns+" FROM batch_task_queues ORDER BY created_at DESC",
"SELECT " + batchQueueSelectColumns + " FROM batch_task_queues ORDER BY created_at DESC",
)
if err != nil {
return nil, fmt.Errorf("查询批量任务队列列表失败: %w", err)
@@ -168,6 +168,10 @@ func (db *DB) GetAllBatchQueues() ([]*BatchTaskQueueRow, error) {
// ListBatchQueues 列出批量任务队列(支持筛选和分页)
func (db *DB) ListBatchQueues(limit, offset int, status, keyword string) ([]*BatchTaskQueueRow, error) {
return db.ListBatchQueuesForAccess(limit, offset, status, keyword, "", "")
}
func (db *DB) ListBatchQueuesForAccess(limit, offset int, status, keyword, userID, scope string) ([]*BatchTaskQueueRow, error) {
query := "SELECT " + batchQueueSelectColumns + " FROM batch_task_queues WHERE 1=1"
args := []interface{}{}
@@ -182,6 +186,26 @@ func (db *DB) ListBatchQueues(limit, offset int, status, keyword string) ([]*Bat
query += " AND (id LIKE ? OR title LIKE ?)"
args = append(args, "%"+keyword+"%", "%"+keyword+"%")
}
userID = strings.TrimSpace(userID)
if userID != "" && scope != RBACScopeAll {
query += ` AND (
owner_user_id = ?
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments ra
WHERE ra.user_id = ? AND ra.resource_type = 'batch_task' AND ra.resource_id = batch_task_queues.id
)
OR (
project_id IS NOT NULL AND project_id <> '' AND (
EXISTS (SELECT 1 FROM projects p WHERE p.id = batch_task_queues.project_id AND p.owner_user_id = ?)
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments pra
WHERE pra.user_id = ? AND pra.resource_type = 'project' AND pra.resource_id = batch_task_queues.project_id
)
)
)
)`
args = append(args, userID, userID, userID, userID)
}
query += " ORDER BY created_at DESC LIMIT ? OFFSET ?"
args = append(args, limit, offset)
@@ -216,6 +240,10 @@ func (db *DB) ListBatchQueues(limit, offset int, status, keyword string) ([]*Bat
// CountBatchQueues 统计批量任务队列总数(支持筛选条件)
func (db *DB) CountBatchQueues(status, keyword string) (int, error) {
return db.CountBatchQueuesForAccess(status, keyword, "", "")
}
func (db *DB) CountBatchQueuesForAccess(status, keyword, userID, scope string) (int, error) {
query := "SELECT COUNT(*) FROM batch_task_queues WHERE 1=1"
args := []interface{}{}
@@ -230,6 +258,26 @@ func (db *DB) CountBatchQueues(status, keyword string) (int, error) {
query += " AND (id LIKE ? OR title LIKE ?)"
args = append(args, "%"+keyword+"%", "%"+keyword+"%")
}
userID = strings.TrimSpace(userID)
if userID != "" && scope != RBACScopeAll {
query += ` AND (
owner_user_id = ?
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments ra
WHERE ra.user_id = ? AND ra.resource_type = 'batch_task' AND ra.resource_id = batch_task_queues.id
)
OR (
project_id IS NOT NULL AND project_id <> '' AND (
EXISTS (SELECT 1 FROM projects p WHERE p.id = batch_task_queues.project_id AND p.owner_user_id = ?)
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments pra
WHERE pra.user_id = ? AND pra.resource_type = 'project' AND pra.resource_id = batch_task_queues.project_id
)
)
)
)`
args = append(args, userID, userID, userID, userID)
}
var count int
err := db.QueryRow(query, args...).Scan(&count)
+507 -32
View File
@@ -45,20 +45,21 @@ func validC2TextIDForDelete(id string) bool {
// C2Listener 监听器实体
type C2Listener struct {
ID string `json:"id"`
Name string `json:"name"`
Type string `json:"type"` // tcp_reverse|http_beacon|https_beacon|websocket|dns
BindHost string `json:"bindHost"` // 默认 127.0.0.1
BindPort int `json:"bindPort"` // 1-65535
ProfileID string `json:"profileId"` // 可空:关联 c2_profiles.id
EncryptionKey string `json:"-"` // base64(AES-256),前端不返回
ImplantToken string `json:"-"` // beacon 携带的鉴权 token,前端不返回
Status string `json:"status"` // stopped|running|error
ConfigJSON string `json:"configJson"` // TLS 证书路径 / URI 模式 / 上限并发 等
Remark string `json:"remark"`
CreatedAt time.Time `json:"createdAt"`
ID string `json:"id"`
Name string `json:"name"`
Type string `json:"type"` // tcp_reverse|http_beacon|https_beacon|websocket|dns
BindHost string `json:"bindHost"` // 默认 127.0.0.1
BindPort int `json:"bindPort"` // 1-65535
ProfileID string `json:"profileId"` // 可空:关联 c2_profiles.id
EncryptionKey string `json:"-"` // base64(AES-256),前端不返回
ImplantToken string `json:"-"` // beacon 携带的鉴权 token,前端不返回
Status string `json:"status"` // stopped|running|error
ConfigJSON string `json:"configJson"` // TLS 证书路径 / URI 模式 / 上限并发 等
Remark string `json:"remark"`
OwnerUserID string `json:"ownerUserId,omitempty"`
CreatedAt time.Time `json:"createdAt"`
StartedAt *time.Time `json:"startedAt,omitempty"`
LastError string `json:"lastError,omitempty"`
LastError string `json:"lastError,omitempty"`
}
// C2Session 已上线会话
@@ -132,17 +133,17 @@ type C2Event struct {
// C2Profile Malleable Profile
type C2Profile struct {
ID string `json:"id"`
Name string `json:"name"`
UserAgent string `json:"userAgent"`
URIs []string `json:"uris"`
RequestHeaders map[string]string `json:"requestHeaders,omitempty"`
ResponseHeaders map[string]string `json:"responseHeaders,omitempty"`
BodyTemplate string `json:"bodyTemplate"`
JitterMinMS int `json:"jitterMinMs"`
JitterMaxMS int `json:"jitterMaxMs"`
Extra map[string]interface{} `json:"extra,omitempty"`
CreatedAt time.Time `json:"createdAt"`
ID string `json:"id"`
Name string `json:"name"`
UserAgent string `json:"userAgent"`
URIs []string `json:"uris"`
RequestHeaders map[string]string `json:"requestHeaders,omitempty"`
ResponseHeaders map[string]string `json:"responseHeaders,omitempty"`
BodyTemplate string `json:"bodyTemplate"`
JitterMinMS int `json:"jitterMinMs"`
JitterMaxMS int `json:"jitterMaxMs"`
Extra map[string]interface{} `json:"extra,omitempty"`
CreatedAt time.Time `json:"createdAt"`
}
// ----------------------------------------------------------------------------
@@ -165,12 +166,12 @@ func (db *DB) CreateC2Listener(l *C2Listener) error {
}
query := `
INSERT INTO c2_listeners (id, name, type, bind_host, bind_port, profile_id, encryption_key,
implant_token, status, config_json, remark, created_at, started_at, last_error)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
implant_token, status, config_json, remark, owner_user_id, created_at, started_at, last_error)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`
_, err := db.Exec(query,
l.ID, l.Name, l.Type, l.BindHost, l.BindPort, l.ProfileID, l.EncryptionKey,
l.ImplantToken, l.Status, l.ConfigJSON, l.Remark, l.CreatedAt, l.StartedAt, l.LastError,
l.ImplantToken, l.Status, l.ConfigJSON, l.Remark, l.OwnerUserID, l.CreatedAt, l.StartedAt, l.LastError,
)
if err != nil {
db.logger.Error("创建 C2 监听器失败", zap.Error(err), zap.String("id", l.ID))
@@ -190,12 +191,12 @@ func (db *DB) UpdateC2Listener(l *C2Listener) error {
query := `
UPDATE c2_listeners SET
name = ?, type = ?, bind_host = ?, bind_port = ?, profile_id = ?, encryption_key = ?,
implant_token = ?, status = ?, config_json = ?, remark = ?, started_at = ?, last_error = ?
implant_token = ?, status = ?, config_json = ?, remark = ?, owner_user_id = ?, started_at = ?, last_error = ?
WHERE id = ?
`
res, err := db.Exec(query,
l.Name, l.Type, l.BindHost, l.BindPort, l.ProfileID, l.EncryptionKey,
l.ImplantToken, l.Status, l.ConfigJSON, l.Remark, l.StartedAt, l.LastError, l.ID,
l.ImplantToken, l.Status, l.ConfigJSON, l.Remark, l.OwnerUserID, l.StartedAt, l.LastError, l.ID,
)
if err != nil {
db.logger.Error("更新 C2 监听器失败", zap.Error(err), zap.String("id", l.ID))
@@ -231,7 +232,7 @@ func (db *DB) GetC2Listener(id string) (*C2Listener, error) {
SELECT id, name, type, bind_host, bind_port, COALESCE(profile_id, ''),
COALESCE(encryption_key, ''), COALESCE(implant_token, ''), status,
COALESCE(config_json, '{}'), COALESCE(remark, ''),
created_at, started_at, COALESCE(last_error, '')
COALESCE(owner_user_id, ''), created_at, started_at, COALESCE(last_error, '')
FROM c2_listeners WHERE id = ?
`
var l C2Listener
@@ -240,7 +241,7 @@ func (db *DB) GetC2Listener(id string) (*C2Listener, error) {
&l.ID, &l.Name, &l.Type, &l.BindHost, &l.BindPort, &l.ProfileID,
&l.EncryptionKey, &l.ImplantToken, &l.Status,
&l.ConfigJSON, &l.Remark,
&l.CreatedAt, &startedAt, &l.LastError,
&l.OwnerUserID, &l.CreatedAt, &startedAt, &l.LastError,
)
if err == sql.ErrNoRows {
return nil, nil
@@ -261,7 +262,7 @@ func (db *DB) ListC2Listeners() ([]*C2Listener, error) {
SELECT id, name, type, bind_host, bind_port, COALESCE(profile_id, ''),
COALESCE(encryption_key, ''), COALESCE(implant_token, ''), status,
COALESCE(config_json, '{}'), COALESCE(remark, ''),
created_at, started_at, COALESCE(last_error, '')
COALESCE(owner_user_id, ''), created_at, started_at, COALESCE(last_error, '')
FROM c2_listeners ORDER BY created_at DESC
`
rows, err := db.Query(query)
@@ -277,6 +278,47 @@ func (db *DB) ListC2Listeners() ([]*C2Listener, error) {
&l.ID, &l.Name, &l.Type, &l.BindHost, &l.BindPort, &l.ProfileID,
&l.EncryptionKey, &l.ImplantToken, &l.Status,
&l.ConfigJSON, &l.Remark,
&l.OwnerUserID, &l.CreatedAt, &startedAt, &l.LastError,
); err != nil {
db.logger.Warn("扫描 c2_listeners 行失败", zap.Error(err))
continue
}
if startedAt.Valid {
t := startedAt.Time
l.StartedAt = &t
}
list = append(list, &l)
}
return list, rows.Err()
}
// ListC2ListenersForAccess lists listeners visible to the resolved RBAC scope.
func (db *DB) ListC2ListenersForAccess(access RBACListAccess) ([]*C2Listener, error) {
conditions := []string{"1=1"}
args := []interface{}{}
appendC2ListenerAccessFilter(&conditions, &args, access)
query := `
SELECT id, name, type, bind_host, bind_port, COALESCE(profile_id, ''),
COALESCE(encryption_key, ''), COALESCE(implant_token, ''), status,
COALESCE(config_json, '{}'), COALESCE(remark, ''),
COALESCE(owner_user_id, ''), created_at, started_at, COALESCE(last_error, '')
FROM c2_listeners
WHERE ` + strings.Join(conditions, " AND ") + `
ORDER BY created_at DESC
`
rows, err := db.Query(query, args...)
if err != nil {
return nil, err
}
defer rows.Close()
var list []*C2Listener
for rows.Next() {
var l C2Listener
var startedAt sql.NullTime
if err := rows.Scan(
&l.ID, &l.Name, &l.Type, &l.BindHost, &l.BindPort, &l.ProfileID,
&l.EncryptionKey, &l.ImplantToken, &l.Status,
&l.ConfigJSON, &l.Remark, &l.OwnerUserID,
&l.CreatedAt, &startedAt, &l.LastError,
); err != nil {
db.logger.Warn("扫描 c2_listeners 行失败", zap.Error(err))
@@ -291,6 +333,26 @@ func (db *DB) ListC2Listeners() ([]*C2Listener, error) {
return list, rows.Err()
}
func appendC2ListenerAccessFilter(conditions *[]string, args *[]interface{}, access RBACListAccess) {
if access.Scope == RBACScopeAll {
return
}
if access.UserID == "" {
*conditions = append(*conditions, "1=0")
return
}
clauses := []string{"owner_user_id = ?"}
*args = append(*args, access.UserID)
if access.Scope == RBACScopeAssigned {
clauses = append(clauses, `EXISTS (
SELECT 1 FROM rbac_resource_assignments ra
WHERE ra.user_id = ? AND ra.resource_type = 'c2_listener' AND ra.resource_id = c2_listeners.id
)`)
*args = append(*args, access.UserID)
}
*conditions = append(*conditions, "("+strings.Join(clauses, " OR ")+")")
}
// DeleteC2Listener 级联删除(会话/任务/文件/事件随之消失)
func (db *DB) DeleteC2Listener(id string) error {
res, err := db.Exec(`DELETE FROM c2_listeners WHERE id = ?`, id)
@@ -550,6 +612,109 @@ func (db *DB) ListC2Sessions(filter ListC2SessionsFilter) ([]*C2Session, error)
return list, rows.Err()
}
// ListC2SessionsForAccess lists sessions whose parent listener is visible.
func (db *DB) ListC2SessionsForAccess(filter ListC2SessionsFilter, access RBACListAccess) ([]*C2Session, error) {
conditions, args := buildC2SessionsWhere(filter)
appendC2SessionAccessFilter(&conditions, &args, access)
query := `
SELECT id, listener_id, implant_uuid, COALESCE(hostname,''), COALESCE(username,''),
COALESCE(os,''), COALESCE(arch,''), COALESCE(pid, 0), COALESCE(process_name,''),
COALESCE(is_admin, 0), COALESCE(internal_ip,''), COALESCE(external_ip,''),
COALESCE(user_agent,''), COALESCE(sleep_seconds, 5), COALESCE(jitter_percent, 0),
status, first_seen_at, last_check_in, COALESCE(metadata_json, '{}'),
COALESCE(note, '')
FROM c2_sessions
WHERE ` + strings.Join(conditions, " AND ") + `
ORDER BY last_check_in DESC
`
if filter.Limit > 0 {
query += fmt.Sprintf(" LIMIT %d", filter.Limit)
}
rows, err := db.Query(query, args...)
if err != nil {
return nil, err
}
defer rows.Close()
return db.scanC2SessionRows(rows)
}
func buildC2SessionsWhere(filter ListC2SessionsFilter) ([]string, []interface{}) {
conditions := []string{"1=1"}
args := []interface{}{}
if filter.ListenerID != "" {
conditions = append(conditions, "listener_id = ?")
args = append(args, filter.ListenerID)
}
if filter.Status != "" {
conditions = append(conditions, "status = ?")
args = append(args, filter.Status)
}
if filter.OS != "" {
conditions = append(conditions, "os = ?")
args = append(args, filter.OS)
}
if filter.Search != "" {
conditions = append(conditions, "(hostname LIKE ? OR username LIKE ? OR internal_ip LIKE ?)")
kw := "%" + filter.Search + "%"
args = append(args, kw, kw, kw)
}
if filter.Suspicious {
conditions = append(conditions, `status = 'dead' AND (
hostname LIKE 'tcp_%' OR LOWER(COALESCE(username,'')) = 'unknown' OR COALESCE(pid, 0) = 0
)`)
}
return conditions, args
}
func (db *DB) scanC2SessionRows(rows *sql.Rows) ([]*C2Session, error) {
var list []*C2Session
for rows.Next() {
var s C2Session
var isAdminInt int
var metadataJSON string
if err := rows.Scan(
&s.ID, &s.ListenerID, &s.ImplantUUID, &s.Hostname, &s.Username,
&s.OS, &s.Arch, &s.PID, &s.ProcessName,
&isAdminInt, &s.InternalIP, &s.ExternalIP,
&s.UserAgent, &s.SleepSeconds, &s.JitterPercent,
&s.Status, &s.FirstSeenAt, &s.LastCheckIn, &metadataJSON,
&s.Note,
); err != nil {
db.logger.Warn("扫描 c2_sessions 行失败", zap.Error(err))
continue
}
s.IsAdmin = isAdminInt != 0
if metadataJSON != "" && metadataJSON != "{}" {
_ = json.Unmarshal([]byte(metadataJSON), &s.Metadata)
}
list = append(list, &s)
}
return list, rows.Err()
}
func appendC2SessionAccessFilter(conditions *[]string, args *[]interface{}, access RBACListAccess) {
if access.Scope == RBACScopeAll {
return
}
if access.UserID == "" {
*conditions = append(*conditions, "1=0")
return
}
clauses := []string{`EXISTS (
SELECT 1 FROM c2_listeners
WHERE c2_listeners.id = c2_sessions.listener_id AND c2_listeners.owner_user_id = ?
)`}
*args = append(*args, access.UserID)
if access.Scope == RBACScopeAssigned {
clauses = append(clauses, `EXISTS (
SELECT 1 FROM rbac_resource_assignments ra
WHERE ra.user_id = ? AND ra.resource_type = 'c2_listener' AND ra.resource_id = c2_sessions.listener_id
)`)
*args = append(*args, access.UserID)
}
*conditions = append(*conditions, "("+strings.Join(clauses, " OR ")+")")
}
// DeleteC2Session 级联删除其 tasks/files
func (db *DB) DeleteC2Session(id string) error {
res, err := db.Exec(`DELETE FROM c2_sessions WHERE id = ?`, id)
@@ -601,6 +766,29 @@ func (db *DB) DeleteC2SessionsByIDs(ids []string) (int64, error) {
return res.RowsAffected()
}
func (db *DB) DeleteC2SessionsByIDsForAccess(ids []string, access RBACListAccess) (int64, error) {
if access.Scope == RBACScopeAll {
return db.DeleteC2SessionsByIDs(ids)
}
clean := cleanC2IDs(ids)
if len(clean) == 0 {
return 0, ErrNoValidC2SessionIDs
}
placeholders := strings.Repeat("?,", len(clean)-1) + "?"
args := make([]interface{}, 0, len(clean)+2)
for _, id := range clean {
args = append(args, id)
}
conditions := []string{"id IN (" + placeholders + ")"}
appendC2SessionAccessFilter(&conditions, &args, access)
query := `DELETE FROM c2_sessions WHERE ` + strings.Join(conditions, " AND ")
res, err := db.Exec(query, args...)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
// ----------------------------------------------------------------------------
// CRUDC2 任务
// ----------------------------------------------------------------------------
@@ -778,6 +966,39 @@ func buildC2TasksWhere(filter ListC2TasksFilter) (where string, args []interface
return strings.Join(conditions, " AND "), args
}
func appendC2TaskAccessFilter(conditions *[]string, args *[]interface{}, access RBACListAccess) {
if access.Scope == RBACScopeAll {
return
}
if access.UserID == "" {
*conditions = append(*conditions, "1=0")
return
}
clauses := []string{`EXISTS (
SELECT 1 FROM c2_sessions s
JOIN c2_listeners l ON l.id = s.listener_id
WHERE s.id = c2_tasks.session_id AND l.owner_user_id = ?
)`}
*args = append(*args, access.UserID)
if access.Scope == RBACScopeAssigned {
clauses = append(clauses, `EXISTS (
SELECT 1 FROM c2_sessions s
JOIN rbac_resource_assignments ra ON ra.resource_id = s.listener_id
WHERE s.id = c2_tasks.session_id
AND ra.user_id = ? AND ra.resource_type = 'c2_listener'
)`)
*args = append(*args, access.UserID)
}
*conditions = append(*conditions, "("+strings.Join(clauses, " OR ")+")")
}
func buildC2TasksWhereForAccess(filter ListC2TasksFilter, access RBACListAccess) (string, []interface{}) {
where, args := buildC2TasksWhere(filter)
conditions := []string{where}
appendC2TaskAccessFilter(&conditions, &args, access)
return strings.Join(conditions, " AND "), args
}
// CountC2Tasks 与 ListC2Tasks 相同过滤条件下的记录总数
func (db *DB) CountC2Tasks(filter ListC2TasksFilter) (int64, error) {
where, args := buildC2TasksWhere(filter)
@@ -787,6 +1008,14 @@ func (db *DB) CountC2Tasks(filter ListC2TasksFilter) (int64, error) {
return n, err
}
func (db *DB) CountC2TasksForAccess(filter ListC2TasksFilter, access RBACListAccess) (int64, error) {
where, args := buildC2TasksWhereForAccess(filter, access)
query := `SELECT COUNT(*) FROM c2_tasks WHERE ` + where
var n int64
err := db.QueryRow(query, args...).Scan(&n)
return n, err
}
// CountC2TasksQueuedOrPending 统计 queued/pending 状态任务数(仪表盘「待审任务」)
func (db *DB) CountC2TasksQueuedOrPending(sessionID string) (int64, error) {
conditions := []string{"status IN ('queued', 'pending')"}
@@ -801,6 +1030,15 @@ func (db *DB) CountC2TasksQueuedOrPending(sessionID string) (int64, error) {
return n, err
}
func (db *DB) CountC2TasksQueuedOrPendingForAccess(sessionID string, access RBACListAccess) (int64, error) {
filter := ListC2TasksFilter{SessionID: sessionID}
where, args := buildC2TasksWhereForAccess(filter, access)
query := `SELECT COUNT(*) FROM c2_tasks WHERE status IN ('queued', 'pending') AND ` + where
var n int64
err := db.QueryRow(query, args...).Scan(&n)
return n, err
}
// ListC2Tasks 任务列表,按创建时间倒序
func (db *DB) ListC2Tasks(filter ListC2TasksFilter) ([]*C2Task, error) {
where, args := buildC2TasksWhere(filter)
@@ -866,6 +1104,74 @@ func (db *DB) ListC2Tasks(filter ListC2TasksFilter) ([]*C2Task, error) {
return list, rows.Err()
}
func (db *DB) ListC2TasksForAccess(filter ListC2TasksFilter, access RBACListAccess) ([]*C2Task, error) {
where, args := buildC2TasksWhereForAccess(filter, access)
query := `
SELECT id, session_id, task_type, COALESCE(payload_json, '{}'),
status, COALESCE(result_text, ''), COALESCE(result_blob_path, ''),
COALESCE(error, ''), COALESCE(source, 'manual'),
COALESCE(conversation_id, ''), COALESCE(approval_status, ''),
created_at, sent_at, started_at, completed_at, COALESCE(duration_ms, 0)
FROM c2_tasks
WHERE ` + where + `
ORDER BY created_at DESC
`
limit := filter.Limit
offset := filter.Offset
if offset < 0 {
offset = 0
}
if limit > 0 {
if limit > 1000 {
limit = 1000
}
query += ` LIMIT ? OFFSET ?`
args = append(args, limit, offset)
}
rows, err := db.Query(query, args...)
if err != nil {
return nil, err
}
defer rows.Close()
return db.scanC2TaskRows(rows)
}
func (db *DB) scanC2TaskRows(rows *sql.Rows) ([]*C2Task, error) {
var list []*C2Task
for rows.Next() {
var t C2Task
var payloadJSON string
var sentAt, startedAt, completedAt sql.NullTime
if err := rows.Scan(
&t.ID, &t.SessionID, &t.TaskType, &payloadJSON,
&t.Status, &t.ResultText, &t.ResultBlobPath,
&t.Error, &t.Source,
&t.ConversationID, &t.ApprovalStatus,
&t.CreatedAt, &sentAt, &startedAt, &completedAt, &t.DurationMS,
); err != nil {
db.logger.Warn("扫描 c2_tasks 行失败", zap.Error(err))
continue
}
if payloadJSON != "" && payloadJSON != "{}" {
_ = json.Unmarshal([]byte(payloadJSON), &t.Payload)
}
if sentAt.Valid {
x := sentAt.Time
t.SentAt = &x
}
if startedAt.Valid {
x := startedAt.Time
t.StartedAt = &x
}
if completedAt.Valid {
x := completedAt.Time
t.CompletedAt = &x
}
list = append(list, &t)
}
return list, rows.Err()
}
// PopQueuedC2Tasks 取出某会话所有 queued/approved 任务(用于 beacon 拉取),原子置为 sent
func (db *DB) PopQueuedC2Tasks(sessionID string, limit int) ([]*C2Task, error) {
if limit <= 0 {
@@ -978,6 +1284,29 @@ func (db *DB) DeleteC2TasksByIDs(ids []string) (int64, error) {
return res.RowsAffected()
}
func (db *DB) DeleteC2TasksByIDsForAccess(ids []string, access RBACListAccess) (int64, error) {
if access.Scope == RBACScopeAll {
return db.DeleteC2TasksByIDs(ids)
}
clean := cleanC2IDs(ids)
if len(clean) == 0 {
return 0, ErrNoValidC2TaskIDs
}
placeholders := strings.Repeat("?,", len(clean)-1) + "?"
args := make([]interface{}, 0, len(clean)+2)
for _, id := range clean {
args = append(args, id)
}
conditions := []string{"id IN (" + placeholders + ")"}
appendC2TaskAccessFilter(&conditions, &args, access)
query := `DELETE FROM c2_tasks WHERE ` + strings.Join(conditions, " AND ")
res, err := db.Exec(query, args...)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
// ----------------------------------------------------------------------------
// CRUDC2 文件
// ----------------------------------------------------------------------------
@@ -1024,6 +1353,27 @@ func (db *DB) ListC2FilesBySession(sessionID string) ([]*C2File, error) {
return list, rows.Err()
}
func cleanC2IDs(ids []string) []string {
const maxBatch = 500
if len(ids) > maxBatch {
ids = ids[:maxBatch]
}
clean := make([]string, 0, len(ids))
seen := make(map[string]struct{}, len(ids))
for _, id := range ids {
id = strings.TrimSpace(id)
if !validC2TextIDForDelete(id) {
continue
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
clean = append(clean, id)
}
return clean
}
// ----------------------------------------------------------------------------
// CRUDC2 事件审计
// ----------------------------------------------------------------------------
@@ -1093,6 +1443,56 @@ func buildC2EventsWhere(filter ListC2EventsFilter) (where string, args []interfa
return strings.Join(conditions, " AND "), args
}
func appendC2EventAccessFilter(conditions *[]string, args *[]interface{}, access RBACListAccess) {
if access.Scope == RBACScopeAll {
return
}
if access.UserID == "" {
*conditions = append(*conditions, "1=0")
return
}
clauses := []string{`EXISTS (
SELECT 1 FROM c2_sessions s
JOIN c2_listeners l ON l.id = s.listener_id
WHERE s.id = c2_events.session_id AND l.owner_user_id = ?
)`}
*args = append(*args, access.UserID)
if access.Scope == RBACScopeAssigned {
clauses = append(clauses, `EXISTS (
SELECT 1 FROM c2_sessions s
JOIN rbac_resource_assignments ra ON ra.resource_id = s.listener_id
WHERE s.id = c2_events.session_id
AND ra.user_id = ? AND ra.resource_type = 'c2_listener'
)`)
*args = append(*args, access.UserID)
}
clauses = append(clauses, `EXISTS (
SELECT 1 FROM c2_tasks t
JOIN c2_sessions s ON s.id = t.session_id
JOIN c2_listeners l ON l.id = s.listener_id
WHERE t.id = c2_events.task_id AND l.owner_user_id = ?
)`)
*args = append(*args, access.UserID)
if access.Scope == RBACScopeAssigned {
clauses = append(clauses, `EXISTS (
SELECT 1 FROM c2_tasks t
JOIN c2_sessions s ON s.id = t.session_id
JOIN rbac_resource_assignments ra ON ra.resource_id = s.listener_id
WHERE t.id = c2_events.task_id
AND ra.user_id = ? AND ra.resource_type = 'c2_listener'
)`)
*args = append(*args, access.UserID)
}
*conditions = append(*conditions, "("+strings.Join(clauses, " OR ")+")")
}
func buildC2EventsWhereForAccess(filter ListC2EventsFilter, access RBACListAccess) (string, []interface{}) {
where, args := buildC2EventsWhere(filter)
conditions := []string{where}
appendC2EventAccessFilter(&conditions, &args, access)
return strings.Join(conditions, " AND "), args
}
// CountC2Events 与 ListC2Events 相同过滤条件下的记录总数
func (db *DB) CountC2Events(filter ListC2EventsFilter) (int64, error) {
where, args := buildC2EventsWhere(filter)
@@ -1102,6 +1502,14 @@ func (db *DB) CountC2Events(filter ListC2EventsFilter) (int64, error) {
return n, err
}
func (db *DB) CountC2EventsForAccess(filter ListC2EventsFilter, access RBACListAccess) (int64, error) {
where, args := buildC2EventsWhereForAccess(filter, access)
query := `SELECT COUNT(*) FROM c2_events WHERE ` + where
var n int64
err := db.QueryRow(query, args...).Scan(&n)
return n, err
}
// ListC2Events 事件查询,按创建时间倒序
func (db *DB) ListC2Events(filter ListC2EventsFilter) ([]*C2Event, error) {
where, args := buildC2EventsWhere(filter)
@@ -1143,6 +1551,50 @@ func (db *DB) ListC2Events(filter ListC2EventsFilter) ([]*C2Event, error) {
return list, rows.Err()
}
func (db *DB) ListC2EventsForAccess(filter ListC2EventsFilter, access RBACListAccess) ([]*C2Event, error) {
where, args := buildC2EventsWhereForAccess(filter, access)
limit := filter.Limit
if limit <= 0 || limit > 1000 {
limit = 200
}
offset := filter.Offset
if offset < 0 {
offset = 0
}
query := `
SELECT id, level, category, COALESCE(session_id, ''), COALESCE(task_id, ''),
message, COALESCE(data_json, ''), created_at
FROM c2_events
WHERE ` + where + `
ORDER BY created_at DESC
LIMIT ? OFFSET ?
`
args = append(args, limit, offset)
rows, err := db.Query(query, args...)
if err != nil {
return nil, err
}
defer rows.Close()
return scanC2EventRows(rows)
}
func scanC2EventRows(rows *sql.Rows) ([]*C2Event, error) {
var list []*C2Event
for rows.Next() {
var e C2Event
var dataJSON string
if err := rows.Scan(&e.ID, &e.Level, &e.Category, &e.SessionID, &e.TaskID,
&e.Message, &dataJSON, &e.CreatedAt); err != nil {
continue
}
if dataJSON != "" {
_ = json.Unmarshal([]byte(dataJSON), &e.Data)
}
list = append(list, &e)
}
return list, rows.Err()
}
// DeleteC2EventsByIDs 按主键批量删除事件,返回实际删除行数
func (db *DB) DeleteC2EventsByIDs(ids []string) (int64, error) {
if len(ids) == 0 {
@@ -1181,6 +1633,29 @@ func (db *DB) DeleteC2EventsByIDs(ids []string) (int64, error) {
return res.RowsAffected()
}
func (db *DB) DeleteC2EventsByIDsForAccess(ids []string, access RBACListAccess) (int64, error) {
if access.Scope == RBACScopeAll {
return db.DeleteC2EventsByIDs(ids)
}
clean := cleanC2IDs(ids)
if len(clean) == 0 {
return 0, ErrNoValidC2EventIDs
}
placeholders := strings.Repeat("?,", len(clean)-1) + "?"
args := make([]interface{}, 0, len(clean)+4)
for _, id := range clean {
args = append(args, id)
}
conditions := []string{"id IN (" + placeholders + ")"}
appendC2EventAccessFilter(&conditions, &args, access)
query := `DELETE FROM c2_events WHERE ` + strings.Join(conditions, " AND ")
res, err := db.Exec(query, args...)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
// ----------------------------------------------------------------------------
// CRUDC2 Malleable Profile
// ----------------------------------------------------------------------------
+30
View File
@@ -0,0 +1,30 @@
package database
import (
"strings"
"time"
)
func (db *DB) RecordC2PayloadArtifact(filename, payloadID, listenerID, ownerUserID string) error {
filename = strings.TrimSpace(filename)
if filename == "" || strings.TrimSpace(listenerID) == "" || strings.TrimSpace(ownerUserID) == "" {
return nil
}
_, err := db.Exec(`
INSERT INTO c2_payload_artifacts(filename, payload_id, listener_id, owner_user_id, created_at)
VALUES (?, ?, ?, ?, ?)
ON CONFLICT(filename) DO UPDATE SET payload_id=excluded.payload_id, listener_id=excluded.listener_id, owner_user_id=excluded.owner_user_id, created_at=excluded.created_at
`, filename, payloadID, listenerID, ownerUserID, time.Now())
return err
}
func (db *DB) UserCanAccessC2Payload(userID, scope, filename string) bool {
if scope == RBACScopeAll {
return true
}
var listenerID, ownerUserID string
if err := db.QueryRow(`SELECT listener_id, owner_user_id FROM c2_payload_artifacts WHERE filename = ?`, strings.TrimSpace(filename)).Scan(&listenerID, &ownerUserID); err != nil {
return false
}
return ownerUserID == strings.TrimSpace(userID) || db.UserCanAccessResource(userID, scope, "c2_listener", listenerID)
}
+56
View File
@@ -0,0 +1,56 @@
package database
import (
"strings"
"time"
)
func (db *DB) UpsertChatUploadArtifact(relativePath, conversationID, ownerUserID string) error {
relativePath = strings.TrimSpace(relativePath)
conversationID = strings.TrimSpace(conversationID)
ownerUserID = strings.TrimSpace(ownerUserID)
if relativePath == "" || conversationID == "" || ownerUserID == "" {
return nil
}
_, err := db.Exec(`
INSERT INTO chat_upload_artifacts(relative_path, conversation_id, owner_user_id, created_at)
VALUES (?, ?, ?, ?)
ON CONFLICT(relative_path) DO UPDATE SET conversation_id=excluded.conversation_id, owner_user_id=excluded.owner_user_id
`, relativePath, conversationID, ownerUserID, time.Now())
return err
}
func (db *DB) GetChatUploadArtifact(relativePath string) (conversationID, ownerUserID string, ok bool) {
err := db.QueryRow(`SELECT conversation_id, owner_user_id FROM chat_upload_artifacts WHERE relative_path = ?`, strings.TrimSpace(relativePath)).Scan(&conversationID, &ownerUserID)
return conversationID, ownerUserID, err == nil
}
func (db *DB) DeleteChatUploadArtifactPath(relativePath string) error {
path := strings.Trim(strings.TrimSpace(relativePath), "/")
if path == "" {
return nil
}
_, err := db.Exec(`DELETE FROM chat_upload_artifacts WHERE relative_path = ? OR relative_path LIKE ? ESCAPE '\'`, path, escapeLikePrefix(path)+"/%")
return err
}
func (db *DB) RenameChatUploadArtifactPath(oldPath, newPath string) error {
oldPath = strings.Trim(strings.TrimSpace(oldPath), "/")
newPath = strings.Trim(strings.TrimSpace(newPath), "/")
if oldPath == "" || newPath == "" {
return nil
}
_, err := db.Exec(`
UPDATE chat_upload_artifacts
SET relative_path = CASE
WHEN relative_path = ? THEN ?
ELSE ? || substr(relative_path, length(?) + 1)
END
WHERE relative_path = ? OR relative_path LIKE ? ESCAPE '\'
`, oldPath, newPath, newPath, oldPath, oldPath, escapeLikePrefix(oldPath)+"/%")
return err
}
func escapeLikePrefix(value string) string {
return strings.NewReplacer(`\`, `\\`, `%`, `\%`, `_`, `\_`).Replace(value)
}
+376 -5
View File
@@ -384,6 +384,31 @@ func appendConversationProjectFilter(where string, args []interface{}, projectID
return where + fmt.Sprintf(" AND %s = ?", col), append(args, pid)
}
func appendConversationAccessFilter(where string, args []interface{}, userID, scope, alias string) (string, []interface{}) {
userID = strings.TrimSpace(userID)
if userID == "" || scope == RBACScopeAll {
return where, args
}
prefix := ""
if alias != "" {
prefix = alias + "."
}
where += fmt.Sprintf(` AND (%sowner_user_id = ? OR EXISTS (
SELECT 1 FROM rbac_resource_assignments ra
WHERE ra.user_id = ? AND ra.resource_type = 'conversation' AND ra.resource_id = %sid
) OR EXISTS (
SELECT 1 FROM projects p
WHERE p.id = %sproject_id AND (
p.owner_user_id = ? OR EXISTS (
SELECT 1 FROM rbac_resource_assignments pra
WHERE pra.user_id = ? AND pra.resource_type = 'project' AND pra.resource_id = p.id
)
)
))`, prefix, prefix, prefix)
args = append(args, userID, userID, userID, userID)
return where, args
}
// CountConversations 统计对话数量。
func (db *DB) CountConversations(search, projectID string) (int, error) {
var count int
@@ -410,6 +435,33 @@ func (db *DB) CountConversations(search, projectID string) (int, error) {
return count, nil
}
func (db *DB) CountConversationsForAccess(search, projectID, userID, scope string) (int, error) {
var count int
var err error
if search != "" {
searchPattern := "%" + search + "%"
where := ` WHERE (c.title LIKE ?
OR EXISTS (SELECT 1 FROM messages m WHERE m.conversation_id = c.id AND m.content LIKE ?))`
args := []interface{}{searchPattern, searchPattern}
where, args = appendConversationProjectFilter(where, args, projectID, "c")
where, args = appendConversationAccessFilter(where, args, userID, scope, "c")
err = db.QueryRow(`SELECT COUNT(*) FROM conversations c`+where, args...).Scan(&count)
} else {
where := ""
args := []interface{}{}
where, args = appendConversationProjectFilter(where, args, projectID, "")
where, args = appendConversationAccessFilter(where, args, userID, scope, "")
if where != "" {
where = " WHERE" + strings.TrimPrefix(where, " AND")
}
err = db.QueryRow(`SELECT COUNT(*) FROM conversations`+where, args...).Scan(&count)
}
if err != nil {
return 0, fmt.Errorf("统计对话失败: %w", err)
}
return count, nil
}
func conversationOrderClause(sortBy, tableAlias string) string {
col := "updated_at"
if strings.TrimSpace(strings.ToLower(sortBy)) == "created_at" {
@@ -503,6 +555,81 @@ func (db *DB) ListConversations(limit, offset int, search, sortBy, projectID str
return conversations, nil
}
func (db *DB) ListConversationsForAccess(limit, offset int, search, sortBy, projectID, userID, scope string) ([]*Conversation, error) {
if scope == RBACScopeAll || strings.TrimSpace(userID) == "" {
return db.ListConversations(limit, offset, search, sortBy, projectID)
}
var rows *sql.Rows
var err error
if search != "" {
searchPattern := "%" + search + "%"
orderClause := conversationOrderClause(sortBy, "c")
where := ` WHERE (c.title LIKE ?
OR EXISTS (SELECT 1 FROM messages m WHERE m.conversation_id = c.id AND m.content LIKE ?))`
args := []interface{}{searchPattern, searchPattern}
where, args = appendConversationProjectFilter(where, args, projectID, "c")
where, args = appendConversationAccessFilter(where, args, userID, scope, "c")
args = append(args, limit, offset)
rows, err = db.Query(
`SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id
FROM conversations c`+where+`
`+orderClause+`
LIMIT ? OFFSET ?`, args...)
} else {
orderClause := conversationOrderClause(sortBy, "")
where := ""
args := []interface{}{}
where, args = appendConversationProjectFilter(where, args, projectID, "")
where, args = appendConversationAccessFilter(where, args, userID, scope, "")
if where != "" {
where = " WHERE" + strings.TrimPrefix(where, " AND")
}
args = append(args, limit, offset)
rows, err = db.Query(
"SELECT id, title, COALESCE(pinned, 0), created_at, updated_at, project_id FROM conversations"+where+" "+orderClause+" LIMIT ? OFFSET ?",
args...)
}
if err != nil {
return nil, fmt.Errorf("查询对话列表失败: %w", err)
}
defer rows.Close()
return scanConversationRows(rows)
}
func scanConversationRows(rows *sql.Rows) ([]*Conversation, error) {
var conversations []*Conversation
for rows.Next() {
var conv Conversation
var createdAt, updatedAt string
var pinned int
var projectID sql.NullString
if err := rows.Scan(&conv.ID, &conv.Title, &pinned, &createdAt, &updatedAt, &projectID); err != nil {
return nil, fmt.Errorf("扫描对话失败: %w", err)
}
if projectID.Valid {
conv.ProjectID = strings.TrimSpace(projectID.String)
}
var err1, err2 error
conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt)
if err1 != nil {
conv.CreatedAt, err1 = time.Parse("2006-01-02 15:04:05", createdAt)
}
if err1 != nil {
conv.CreatedAt, _ = time.Parse(time.RFC3339, createdAt)
}
conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05.999999999-07:00", updatedAt)
if err2 != nil {
conv.UpdatedAt, err2 = time.Parse("2006-01-02 15:04:05", updatedAt)
}
if err2 != nil {
conv.UpdatedAt, _ = time.Parse(time.RFC3339, updatedAt)
}
conv.Pinned = pinned != 0
conversations = append(conversations, &conv)
}
return conversations, rows.Err()
}
const ungroupedConversationsSQL = `
FROM conversations c
WHERE NOT EXISTS (
@@ -521,6 +648,18 @@ func (db *DB) CountUngroupedConversations(projectID string) (int, error) {
return count, nil
}
func (db *DB) CountUngroupedConversationsForAccess(projectID, userID, scope string) (int, error) {
where := ungroupedConversationsSQL
args := []interface{}{}
where, args = appendConversationProjectFilter(where, args, projectID, "c")
where, args = appendConversationAccessFilter(where, args, userID, scope, "c")
var count int
if err := db.QueryRow(`SELECT COUNT(*) `+where, args...).Scan(&count); err != nil {
return 0, fmt.Errorf("统计未分组对话失败: %w", err)
}
return count, nil
}
// ListUngroupedConversations 列出不在任何分组中的对话(最近对话侧栏)。
func (db *DB) ListUngroupedConversations(limit, offset int, sortBy, projectID string) ([]*Conversation, error) {
orderClause := conversationOrderClause(sortBy, "c")
@@ -578,6 +717,30 @@ func (db *DB) ListUngroupedConversations(limit, offset int, sortBy, projectID st
return conversations, rows.Err()
}
func (db *DB) ListUngroupedConversationsForAccess(limit, offset int, sortBy, projectID, userID, scope string) ([]*Conversation, error) {
if scope == RBACScopeAll || strings.TrimSpace(userID) == "" {
return db.ListUngroupedConversations(limit, offset, sortBy, projectID)
}
orderClause := conversationOrderClause(sortBy, "c")
where := ungroupedConversationsSQL
args := []interface{}{}
where, args = appendConversationProjectFilter(where, args, projectID, "c")
where, args = appendConversationAccessFilter(where, args, userID, scope, "c")
args = append(args, limit, offset)
rows, err := db.Query(
`SELECT c.id, c.title, COALESCE(c.pinned, 0), c.created_at, c.updated_at, c.project_id `+
where+`
`+orderClause+`
LIMIT ? OFFSET ?`,
args...,
)
if err != nil {
return nil, fmt.Errorf("查询未分组对话失败: %w", err)
}
defer rows.Close()
return scanConversationRows(rows)
}
// GetConversationTitle 获取对话标题(轻量查询,不加载消息)
func (db *DB) GetConversationTitle(id string) (string, error) {
var title string
@@ -1144,6 +1307,12 @@ ORDER BY created_at ASC, rowid ASC`, assistantMessageID)
// AddProcessDetail 添加过程详情事件
func (db *DB) AddProcessDetail(messageID, conversationID, eventType, message string, data interface{}) error {
_, err := db.AddProcessDetailWithID(messageID, conversationID, eventType, message, data)
return err
}
// AddProcessDetailWithID 添加过程详情事件并返回记录 ID。
func (db *DB) AddProcessDetailWithID(messageID, conversationID, eventType, message string, data interface{}) (string, error) {
id := uuid.New().String()
var dataJSON string
@@ -1161,10 +1330,10 @@ func (db *DB) AddProcessDetail(messageID, conversationID, eventType, message str
id, messageID, conversationID, eventType, message, dataJSON, time.Now(),
)
if err != nil {
return fmt.Errorf("添加过程详情失败: %w", err)
return "", fmt.Errorf("添加过程详情失败: %w", err)
}
return nil
return id, nil
}
// GetProcessDetails 获取消息的过程详情
@@ -1203,11 +1372,45 @@ func (db *DB) GetProcessDetails(messageID string) ([]ProcessDetail, error) {
return details, nil
}
// GetProcessDetailByID 获取单条过程详情。
func (db *DB) GetProcessDetailByID(id string) (*ProcessDetail, error) {
var detail ProcessDetail
var createdAt string
err := db.QueryRow(
"SELECT id, message_id, conversation_id, event_type, message, data, created_at FROM process_details WHERE id = ?",
id,
).Scan(&detail.ID, &detail.MessageID, &detail.ConversationID, &detail.EventType, &detail.Message, &detail.Data, &createdAt)
if err != nil {
return nil, fmt.Errorf("查询过程详情失败: %w", err)
}
var parseErr error
detail.CreatedAt, parseErr = time.Parse("2006-01-02 15:04:05.999999999-07:00", createdAt)
if parseErr != nil {
detail.CreatedAt, parseErr = time.Parse("2006-01-02 15:04:05", createdAt)
}
if parseErr != nil {
detail.CreatedAt, _ = time.Parse(time.RFC3339, createdAt)
}
return &detail, nil
}
// ProcessDetailsSummary 过程详情摘要(用于折叠态展示,避免全量加载)。
type ProcessDetailsSummary struct {
Total int `json:"total"`
IterationCount int `json:"iterationCount"`
MaxIteration int `json:"maxIteration"`
Total int `json:"total"`
IterationCount int `json:"iterationCount"`
MaxIteration int `json:"maxIteration"`
ToolCount int `json:"toolCount"`
ToolExecutions []ProcessDetailsToolExecution `json:"toolExecutions,omitempty"`
MCPExecutionIDs []string `json:"mcpExecutionIds,omitempty"`
}
type ProcessDetailsToolExecution struct {
ProcessDetailID string `json:"processDetailId,omitempty"`
ToolName string `json:"toolName,omitempty"`
ToolCallID string `json:"toolCallId,omitempty"`
ExecutionID string `json:"executionId,omitempty"`
Status string `json:"status,omitempty"`
}
// GetProcessDetailsSummary 统计消息的过程详情数量与迭代轮次。
@@ -1225,6 +1428,144 @@ func (db *DB) GetProcessDetailsSummary(messageID string) (*ProcessDetailsSummary
return summary, nil
}
if err := db.QueryRow(
"SELECT COUNT(*) FROM process_details WHERE message_id = ? AND event_type = 'tool_call'",
messageID,
).Scan(&summary.ToolCount); err != nil {
return nil, fmt.Errorf("统计工具调用详情失败: %w", err)
}
execRows, err := db.Query(
"SELECT id, event_type, data FROM process_details WHERE message_id = ? AND event_type IN ('tool_call', 'tool_result') ORDER BY created_at ASC, rowid ASC",
messageID,
)
if err != nil {
return nil, fmt.Errorf("查询工具执行摘要失败: %w", err)
}
seenExecIDs := make(map[string]bool)
// A provider may reuse a fallback toolCallId across streaming rounds. Keep a
// FIFO per ID instead of a single index so every persisted call gets at most
// one result. Results without an ID fall back to the oldest unmatched call.
toolIndexesByCallID := make(map[string][]int)
lastMatchedToolIndexByCallID := make(map[string]int)
matchedToolIndexes := make([]bool, 0)
nextUnmatchedToolIdx := 0
for execRows.Next() {
var detailID string
var eventType string
var dataJSON string
if err := execRows.Scan(&detailID, &eventType, &dataJSON); err != nil {
execRows.Close()
return nil, fmt.Errorf("扫描工具执行摘要失败: %w", err)
}
if dataJSON == "" {
continue
}
var payload map[string]interface{}
if err := json.Unmarshal([]byte(dataJSON), &payload); err != nil {
continue
}
toolName, _ := payload["toolName"].(string)
toolName = strings.TrimSpace(toolName)
toolCallID, _ := payload["toolCallId"].(string)
toolCallID = strings.TrimSpace(toolCallID)
execID, _ := payload["executionId"].(string)
execID = strings.TrimSpace(execID)
status := ""
if eventType == "tool_result" {
if success, ok := payload["success"].(bool); ok {
if success {
status = "completed"
} else {
status = "failed"
}
} else if isErr, ok := payload["isError"].(bool); ok && isErr {
status = "failed"
}
}
if eventType == "tool_call" {
summary.ToolExecutions = append(summary.ToolExecutions, ProcessDetailsToolExecution{
ProcessDetailID: strings.TrimSpace(detailID),
ToolName: toolName,
ToolCallID: toolCallID,
// This summary is reconstructed from persisted history, not live
// execution state. Until a matching result is found the honest state
// is "result_missing", never "running".
Status: "result_missing",
})
matchedToolIndexes = append(matchedToolIndexes, false)
if toolCallID != "" {
toolIndexesByCallID[toolCallID] = append(toolIndexesByCallID[toolCallID], len(summary.ToolExecutions)-1)
}
}
if eventType == "tool_result" {
idx := -1
if toolCallID != "" {
queue := toolIndexesByCallID[toolCallID]
for len(queue) > 0 {
candidate := queue[0]
queue = queue[1:]
if candidate >= 0 && candidate < len(matchedToolIndexes) && !matchedToolIndexes[candidate] {
idx = candidate
break
}
}
toolIndexesByCallID[toolCallID] = queue
if idx < 0 {
// Multiple persisted result events for one call (for example an
// agent-facing reduced result replacing an earlier preview) update
// that call instead of consuming an unrelated FIFO entry.
if previous, ok := lastMatchedToolIndexByCallID[toolCallID]; ok {
idx = previous
}
}
}
if idx < 0 {
for nextUnmatchedToolIdx < len(matchedToolIndexes) && matchedToolIndexes[nextUnmatchedToolIdx] {
nextUnmatchedToolIdx++
}
if nextUnmatchedToolIdx < len(matchedToolIndexes) {
idx = nextUnmatchedToolIdx
nextUnmatchedToolIdx++
}
}
if idx >= 0 && idx < len(summary.ToolExecutions) {
matchedToolIndexes[idx] = true
if toolCallID != "" {
lastMatchedToolIndexByCallID[toolCallID] = idx
}
if summary.ToolExecutions[idx].ToolName == "" {
summary.ToolExecutions[idx].ToolName = toolName
}
if summary.ToolExecutions[idx].ToolCallID == "" {
summary.ToolExecutions[idx].ToolCallID = toolCallID
}
summary.ToolExecutions[idx].ExecutionID = execID
if status != "" {
summary.ToolExecutions[idx].Status = status
}
} else {
summary.ToolExecutions = append(summary.ToolExecutions, ProcessDetailsToolExecution{
ProcessDetailID: strings.TrimSpace(detailID),
ToolName: toolName,
ToolCallID: toolCallID,
ExecutionID: execID,
Status: status,
})
matchedToolIndexes = append(matchedToolIndexes, true)
}
}
if execID != "" && !seenExecIDs[execID] {
seenExecIDs[execID] = true
summary.MCPExecutionIDs = append(summary.MCPExecutionIDs, execID)
}
}
if err := execRows.Err(); err != nil {
execRows.Close()
return nil, fmt.Errorf("遍历工具执行摘要失败: %w", err)
}
execRows.Close()
rows, err := db.Query(
"SELECT data FROM process_details WHERE message_id = ? AND event_type = 'iteration' ORDER BY created_at ASC, rowid ASC",
messageID,
@@ -1304,6 +1645,36 @@ func (db *DB) GetProcessDetailsPage(messageID string, limit, offset int) ([]Proc
return details, total, nil
}
// GetProcessDetailOffset 返回某条过程详情在所属消息详情流中的零基 offset。
func (db *DB) GetProcessDetailOffset(messageID, detailID string) (int, error) {
messageID = strings.TrimSpace(messageID)
detailID = strings.TrimSpace(detailID)
if messageID == "" || detailID == "" {
return 0, fmt.Errorf("messageID and detailID are required")
}
var createdAt string
var rowID int64
if err := db.QueryRow(
"SELECT created_at, rowid FROM process_details WHERE message_id = ? AND id = ?",
messageID, detailID,
).Scan(&createdAt, &rowID); err != nil {
if err == sql.ErrNoRows {
return 0, fmt.Errorf("过程详情不存在")
}
return 0, fmt.Errorf("查询过程详情锚点失败: %w", err)
}
var offset int
if err := db.QueryRow(
`SELECT COUNT(*) FROM process_details
WHERE message_id = ?
AND (created_at < ? OR (created_at = ? AND rowid < ?))`,
messageID, createdAt, createdAt, rowID,
).Scan(&offset); err != nil {
return 0, fmt.Errorf("计算过程详情锚点位置失败: %w", err)
}
return offset, nil
}
// GetProcessDetailsByConversation 获取对话的所有过程详情(按消息分组)
func (db *DB) GetProcessDetailsByConversation(conversationID string) (map[string][]ProcessDetail, error) {
rows, err := db.Query(
+94 -6
View File
@@ -58,6 +58,7 @@ type DB struct {
checkpointDone chan struct{}
closeOnce sync.Once
closeErr error
vulnerabilityCreatedHook func(*Vulnerability)
}
// startPassiveCheckpointLoop 启动后台 PASSIVE checkpoint 循环。
@@ -113,10 +114,10 @@ func (db *DB) runPassiveCheckpoint(trigger string) {
return
}
if busy > 0 {
db.logger.Info("SQLite PASSIVE checkpoint 完成(部分推进)", fields...)
db.logger.Debug("SQLite PASSIVE checkpoint 完成(部分推进)", fields...)
return
}
db.logger.Info("SQLite PASSIVE checkpoint 完成(成功)", fields...)
db.logger.Debug("SQLite PASSIVE checkpoint 完成(成功)", fields...)
}
// NewDB 创建数据库连接
@@ -225,6 +226,8 @@ func (db *DB) initTables() error {
start_time DATETIME NOT NULL,
end_time DATETIME,
duration_ms INTEGER,
owner_user_id TEXT,
conversation_id TEXT,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
);`
@@ -300,6 +303,7 @@ func (db *DB) initTables() error {
id TEXT PRIMARY KEY,
name TEXT NOT NULL,
icon TEXT,
owner_user_id TEXT,
created_at DATETIME NOT NULL,
updated_at DATETIME NOT NULL
);`
@@ -322,6 +326,7 @@ func (db *DB) initTables() error {
session_key TEXT PRIMARY KEY,
conversation_id TEXT NOT NULL,
role_name TEXT NOT NULL DEFAULT '默认',
agent_mode TEXT NOT NULL DEFAULT 'eino_single',
updated_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
FOREIGN KEY (conversation_id) REFERENCES conversations(id) ON DELETE CASCADE
);`
@@ -400,6 +405,33 @@ func (db *DB) initTables() error {
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 := `
CREATE TABLE IF NOT EXISTS batch_task_queues (
@@ -478,6 +510,7 @@ func (db *DB) initTables() error {
status TEXT NOT NULL DEFAULT 'stopped',
config_json TEXT NOT NULL DEFAULT '{}',
remark TEXT NOT NULL DEFAULT '',
owner_user_id TEXT,
created_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
started_at DATETIME,
last_error TEXT
@@ -634,6 +667,28 @@ func (db *DB) initTables() error {
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 := `
CREATE INDEX IF NOT EXISTS idx_messages_conversation_id ON messages(conversation_id);
@@ -698,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_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_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 {
@@ -746,6 +804,9 @@ func (db *DB) initTables() error {
if _, err := db.Exec(createRobotUserSessionsTable); err != nil {
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 {
return fmt.Errorf("创建projects表失败: %w", err)
@@ -783,10 +844,22 @@ func (db *DB) initTables() error {
return fmt.Errorf("创建audit_logs表失败: %w", err)
}
if err := db.initRBACTables(); err != nil {
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{
"workflow_definitions": createWorkflowDefinitionsTable,
"workflow_runs": createWorkflowRunsTable,
"workflow_node_runs": createWorkflowNodeRunsTable,
"workflow_definitions": createWorkflowDefinitionsTable,
"workflow_runs": createWorkflowRunsTable,
"workflow_node_runs": createWorkflowNodeRunsTable,
"workflow_package_inspections": createWorkflowPackageInspectionsTable,
"workflow_package_imports": createWorkflowPackageImportsTable,
} {
if _, err := db.Exec(ddl); err != nil {
return fmt.Errorf("创建%s表失败: %w", tableName, err)
@@ -853,12 +926,27 @@ func (db *DB) initTables() error {
if err := db.migrateWorkflowRunsTable(); err != nil {
db.logger.Warn("迁移workflow_runs表失败", zap.Error(err))
}
if err := db.migrateRBACOwnershipColumns(); err != nil {
db.logger.Warn("迁移RBAC资源归属字段失败", zap.Error(err))
}
if _, err := db.Exec(createIndexes); err != nil {
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
}
+60 -23
View File
@@ -10,20 +10,28 @@ import (
// ConversationGroup 对话分组
type ConversationGroup struct {
ID string `json:"id"`
Name string `json:"name"`
Icon string `json:"icon"`
Pinned bool `json:"pinned"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
ID string `json:"id"`
Name string `json:"name"`
Icon string `json:"icon"`
Pinned bool `json:"pinned"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
OwnerUserID string `json:"-"`
}
// GroupExistsByName 检查分组名称是否已存在
func (db *DB) GroupExistsByName(name string, excludeID string) (bool, error) {
return db.groupExistsByNameForOwner(name, excludeID, "")
}
func (db *DB) groupExistsByNameForOwner(name, excludeID, ownerUserID string) (bool, error) {
var count int
var err error
if excludeID != "" {
if ownerUserID != "" && excludeID != "" {
err = db.QueryRow("SELECT COUNT(*) FROM conversation_groups WHERE name = ? AND owner_user_id = ? AND id != ?", name, ownerUserID, excludeID).Scan(&count)
} else if ownerUserID != "" {
err = db.QueryRow("SELECT COUNT(*) FROM conversation_groups WHERE name = ? AND owner_user_id = ?", name, ownerUserID).Scan(&count)
} else if excludeID != "" {
err = db.QueryRow(
"SELECT COUNT(*) FROM conversation_groups WHERE name = ? AND id != ?",
name, excludeID,
@@ -43,9 +51,13 @@ func (db *DB) GroupExistsByName(name string, excludeID string) (bool, error) {
}
// CreateGroup 创建分组
func (db *DB) CreateGroup(name, icon string) (*ConversationGroup, error) {
func (db *DB) CreateGroup(name, icon string, owners ...string) (*ConversationGroup, error) {
ownerUserID := ""
if len(owners) > 0 {
ownerUserID = owners[0]
}
// 检查名称是否已存在
exists, err := db.GroupExistsByName(name, "")
exists, err := db.groupExistsByNameForOwner(name, "", ownerUserID)
if err != nil {
return nil, err
}
@@ -61,27 +73,39 @@ func (db *DB) CreateGroup(name, icon string) (*ConversationGroup, error) {
}
_, err = db.Exec(
"INSERT INTO conversation_groups (id, name, icon, pinned, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?)",
id, name, icon, 0, now, now,
"INSERT INTO conversation_groups (id, name, icon, pinned, owner_user_id, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?)",
id, name, icon, 0, ownerUserID, now, now,
)
if err != nil {
return nil, fmt.Errorf("创建分组失败: %w", err)
}
return &ConversationGroup{
ID: id,
Name: name,
Icon: icon,
Pinned: false,
CreatedAt: now,
UpdatedAt: now,
ID: id,
Name: name,
Icon: icon,
Pinned: false,
CreatedAt: now,
UpdatedAt: now,
OwnerUserID: ownerUserID,
}, nil
}
// ListGroups 列出所有分组
func (db *DB) ListGroups() ([]*ConversationGroup, error) {
return db.ListGroupsForAccess("", RBACScopeAll)
}
func (db *DB) ListGroupsForAccess(userID, scope string) ([]*ConversationGroup, error) {
query := "SELECT id, name, icon, COALESCE(pinned, 0), COALESCE(owner_user_id, ''), created_at, updated_at FROM conversation_groups"
args := []interface{}{}
if scope != RBACScopeAll {
query += " WHERE owner_user_id = ?"
args = append(args, userID)
}
query += " ORDER BY COALESCE(pinned, 0) DESC, created_at ASC"
rows, err := db.Query(
"SELECT id, name, icon, COALESCE(pinned, 0), created_at, updated_at FROM conversation_groups ORDER BY COALESCE(pinned, 0) DESC, created_at ASC",
query, args...,
)
if err != nil {
return nil, fmt.Errorf("查询分组列表失败: %w", err)
@@ -94,7 +118,7 @@ func (db *DB) ListGroups() ([]*ConversationGroup, error) {
var createdAt, updatedAt string
var pinned int
if err := rows.Scan(&group.ID, &group.Name, &group.Icon, &pinned, &createdAt, &updatedAt); err != nil {
if err := rows.Scan(&group.ID, &group.Name, &group.Icon, &pinned, &group.OwnerUserID, &createdAt, &updatedAt); err != nil {
return nil, fmt.Errorf("扫描分组失败: %w", err)
}
@@ -131,9 +155,9 @@ func (db *DB) GetGroup(id string) (*ConversationGroup, error) {
var pinned int
err := db.QueryRow(
"SELECT id, name, icon, COALESCE(pinned, 0), created_at, updated_at FROM conversation_groups WHERE id = ?",
"SELECT id, name, icon, COALESCE(pinned, 0), COALESCE(owner_user_id, ''), created_at, updated_at FROM conversation_groups WHERE id = ?",
id,
).Scan(&group.ID, &group.Name, &group.Icon, &pinned, &createdAt, &updatedAt)
).Scan(&group.ID, &group.Name, &group.Icon, &pinned, &group.OwnerUserID, &createdAt, &updatedAt)
if err != nil {
if err == sql.ErrNoRows {
return nil, fmt.Errorf("分组不存在")
@@ -164,10 +188,23 @@ func (db *DB) GetGroup(id string) (*ConversationGroup, error) {
return &group, nil
}
func (db *DB) UserCanAccessGroup(userID, scope, groupID string) bool {
if scope == RBACScopeAll {
return true
}
var count int
err := db.QueryRow(`SELECT COUNT(*) FROM conversation_groups WHERE id = ? AND owner_user_id = ?`, groupID, userID).Scan(&count)
return err == nil && count > 0
}
// UpdateGroup 更新分组
func (db *DB) UpdateGroup(id, name, icon string) error {
existing, err := db.GetGroup(id)
if err != nil {
return err
}
// 检查名称是否已存在(排除当前分组)
exists, err := db.GroupExistsByName(name, id)
exists, err := db.groupExistsByNameForOwner(name, id, existing.OwnerUserID)
if err != nil {
return err
}
+126 -7
View File
@@ -46,8 +46,8 @@ func (db *DB) SaveToolExecution(exec *mcp.ToolExecution) error {
query := `
INSERT OR REPLACE INTO tool_executions
(id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
(id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms, owner_user_id, conversation_id, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
`
_, err = db.Exec(query,
@@ -60,6 +60,8 @@ func (db *DB) SaveToolExecution(exec *mcp.ToolExecution) error {
exec.StartTime,
endTime,
durationMs,
strings.TrimSpace(exec.OwnerUserID),
strings.TrimSpace(exec.ConversationID),
time.Now(),
)
@@ -90,6 +92,10 @@ func (db *DB) UpdateToolExecutionResult(id string, result *mcp.ToolResult) error
// CountToolExecutions 统计工具执行记录总数
func (db *DB) CountToolExecutions(status, toolName string) (int, error) {
return db.CountToolExecutionsForAccess(status, toolName, RBACListAccess{Scope: RBACScopeAll})
}
func (db *DB) CountToolExecutionsForAccess(status, toolName string, access RBACListAccess) (int, error) {
query := `SELECT COUNT(*) FROM tool_executions`
args := []interface{}{}
conditions := []string{}
@@ -108,6 +114,7 @@ func (db *DB) CountToolExecutions(status, toolName string) (int, error) {
query += ` AND ` + conditions[i]
}
}
query, args = appendToolExecutionAccessSQL(query, args, access, len(conditions) > 0)
var count int
err := db.QueryRow(query, args...).Scan(&count)
if err != nil {
@@ -135,7 +142,7 @@ func (db *DB) LoadToolExecutionsWithPagination(offset, limit int, status, toolNa
}
query := `
SELECT id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms
SELECT id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms, COALESCE(owner_user_id, ''), COALESCE(conversation_id, '')
FROM tool_executions
`
args := []interface{}{}
@@ -183,6 +190,8 @@ func (db *DB) LoadToolExecutionsWithPagination(offset, limit int, status, toolNa
&exec.StartTime,
&endTime,
&durationMs,
&exec.OwnerUserID,
&exec.ConversationID,
)
if err != nil {
db.logger.Warn("加载执行记录失败", zap.Error(err))
@@ -335,8 +344,62 @@ func (db *DB) LoadToolStatsSummary(topN int) (*ToolStatsSummaryResult, error) {
return result, nil
}
func (db *DB) LoadToolStatsSummaryForAccess(topN int, access RBACListAccess) (*ToolStatsSummaryResult, error) {
if access.Scope == RBACScopeAll {
return db.LoadToolStatsSummary(topN)
}
if topN <= 0 {
topN = 6
}
if topN > 100 {
topN = 100
}
result := &ToolStatsSummaryResult{TopTools: make([]*mcp.ToolStats, 0, topN)}
fromSQL, args := appendToolExecutionAccessSQL(` FROM tool_executions`, nil, access, false)
var lastCall sql.NullString
err := db.QueryRow(`SELECT COUNT(*),
COALESCE(SUM(CASE WHEN status = 'completed' THEN 1 ELSE 0 END), 0),
COALESCE(SUM(CASE WHEN status IN ('failed', 'cancelled') THEN 1 ELSE 0 END), 0),
MAX(start_time), COUNT(DISTINCT tool_name)`+fromSQL, args...).Scan(
&result.Summary.TotalCalls, &result.Summary.SuccessCalls, &result.Summary.FailedCalls,
&lastCall, &result.Summary.ToolCount,
)
if err != nil {
return nil, err
}
if lastCall.Valid {
parsed := parseDBTime(lastCall.String)
result.Summary.LastCallTime = &parsed
}
rows, err := db.Query(`SELECT tool_name, COUNT(*),
SUM(CASE WHEN status = 'completed' THEN 1 ELSE 0 END),
SUM(CASE WHEN status IN ('failed', 'cancelled') THEN 1 ELSE 0 END), MAX(start_time)`+
fromSQL+` GROUP BY tool_name ORDER BY COUNT(*) DESC, tool_name ASC LIMIT ?`, append(args, topN)...)
if err != nil {
return nil, err
}
defer rows.Close()
for rows.Next() {
var stat mcp.ToolStats
var last sql.NullString
if err := rows.Scan(&stat.ToolName, &stat.TotalCalls, &stat.SuccessCalls, &stat.FailedCalls, &last); err != nil {
return nil, err
}
if last.Valid {
parsed := parseDBTime(last.String)
stat.LastCallTime = &parsed
}
result.TopTools = append(result.TopTools, &stat)
}
return result, rows.Err()
}
// LoadToolExecutionListPage 分页加载执行记录列表(不含 arguments/result,供监控列表使用)
func (db *DB) LoadToolExecutionListPage(offset, limit int, status, toolName string) ([]*mcp.ToolExecution, error) {
return db.LoadToolExecutionListPageForAccess(offset, limit, status, toolName, RBACListAccess{Scope: RBACScopeAll})
}
func (db *DB) LoadToolExecutionListPageForAccess(offset, limit int, status, toolName string, access RBACListAccess) ([]*mcp.ToolExecution, error) {
if limit <= 0 {
limit = 20
}
@@ -345,11 +408,13 @@ func (db *DB) LoadToolExecutionListPage(offset, limit int, status, toolName stri
}
query := `
SELECT id, tool_name, status, start_time, end_time, duration_ms
SELECT id, tool_name, status, start_time, end_time, duration_ms, COALESCE(owner_user_id, ''), COALESCE(conversation_id, '')
FROM tool_executions
`
whereSQL, args := toolExecutionsFilterSQL(status, toolName)
query += whereSQL + ` ORDER BY start_time DESC LIMIT ? OFFSET ?`
query += whereSQL
query, args = appendToolExecutionAccessSQL(query, args, access, whereSQL != "")
query += ` ORDER BY start_time DESC LIMIT ? OFFSET ?`
args = append(args, limit, offset)
rows, err := db.Query(query, args...)
@@ -371,6 +436,8 @@ func (db *DB) LoadToolExecutionListPage(offset, limit int, status, toolName stri
&exec.StartTime,
&endTime,
&durationMs,
&exec.OwnerUserID,
&exec.ConversationID,
); err != nil {
db.logger.Warn("加载执行记录列表失败", zap.Error(err))
continue
@@ -387,10 +454,35 @@ func (db *DB) LoadToolExecutionListPage(offset, limit int, status, toolName stri
return executions, nil
}
func appendToolExecutionAccessSQL(query string, args []interface{}, access RBACListAccess, hasWhere bool) (string, []interface{}) {
if access.Scope == RBACScopeAll {
return query, args
}
userID := strings.TrimSpace(access.UserID)
joiner := " WHERE "
if hasWhere {
joiner = " AND "
}
if userID == "" {
return query + joiner + "1=0", args
}
query += joiner + `(
owner_user_id = ?
OR (conversation_id IS NOT NULL AND conversation_id <> '' AND (
EXISTS (SELECT 1 FROM conversations c WHERE c.id = tool_executions.conversation_id AND c.owner_user_id = ?)
OR EXISTS (SELECT 1 FROM rbac_resource_assignments ra WHERE ra.user_id = ? AND ra.resource_type = 'conversation' AND ra.resource_id = tool_executions.conversation_id)
OR EXISTS (SELECT 1 FROM conversations c JOIN projects p ON p.id = c.project_id WHERE c.id = tool_executions.conversation_id AND p.owner_user_id = ?)
OR EXISTS (SELECT 1 FROM conversations c JOIN rbac_resource_assignments pra ON pra.resource_id = c.project_id WHERE c.id = tool_executions.conversation_id AND pra.user_id = ? AND pra.resource_type = 'project')
))
)`
args = append(args, userID, userID, userID, userID, userID)
return query, args
}
// GetToolExecution 根据ID获取单条工具执行记录
func (db *DB) GetToolExecution(id string) (*mcp.ToolExecution, error) {
query := `
SELECT id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms
SELECT id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms, COALESCE(owner_user_id, ''), COALESCE(conversation_id, '')
FROM tool_executions
WHERE id = ?
`
@@ -414,6 +506,8 @@ func (db *DB) GetToolExecution(id string) (*mcp.ToolExecution, error) {
&exec.StartTime,
&endTime,
&durationMs,
&exec.OwnerUserID,
&exec.ConversationID,
)
if err != nil {
return nil, err
@@ -448,6 +542,29 @@ func (db *DB) GetToolExecution(id string) (*mcp.ToolExecution, error) {
return &exec, nil
}
// UserCanAccessToolExecution enforces ownership for monitor detail and mutation
// endpoints. Legacy records without an owner or conversation fail closed for
// non-global users.
func (db *DB) UserCanAccessToolExecution(userID, scope, executionID string) bool {
userID = strings.TrimSpace(userID)
executionID = strings.TrimSpace(executionID)
if userID == "" || executionID == "" {
return false
}
if scope == RBACScopeAll {
return true
}
var ownerUserID, conversationID sql.NullString
if err := db.QueryRow(`SELECT owner_user_id, conversation_id FROM tool_executions WHERE id = ?`, executionID).Scan(&ownerUserID, &conversationID); err != nil {
return false
}
if strings.TrimSpace(ownerUserID.String) == userID {
return true
}
conversation := strings.TrimSpace(conversationID.String)
return conversation != "" && db.UserCanAccessResource(userID, scope, "conversation", conversation)
}
// CancelOrphanedRunningToolExecutions 将仍为 running 的记录批量标记为 cancelled(如进程重启后无对应执行协程)。
func (db *DB) CancelOrphanedRunningToolExecutions(endTime time.Time, errMsg string) (int64, error) {
errMsg = strings.TrimSpace(errMsg)
@@ -584,7 +701,7 @@ func (db *DB) GetToolExecutionsByIds(ids []string) ([]*mcp.ToolExecution, error)
}
query := `
SELECT id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms
SELECT id, tool_name, arguments, status, result, error, start_time, end_time, duration_ms, COALESCE(owner_user_id, ''), COALESCE(conversation_id, '')
FROM tool_executions
WHERE id IN (` + strings.Join(placeholders, ",") + `)
`
@@ -614,6 +731,8 @@ func (db *DB) GetToolExecutionsByIds(ids []string) ([]*mcp.ToolExecution, error)
&exec.StartTime,
&endTime,
&durationMs,
&exec.OwnerUserID,
&exec.ConversationID,
)
if err != nil {
db.logger.Warn("加载执行记录失败", zap.Error(err))
@@ -0,0 +1,108 @@
package database
import (
"path/filepath"
"testing"
"go.uber.org/zap"
)
func TestProcessDetailsSummaryPairsMixedIdentifiedAndIDLessResults(t *testing.T) {
db, conversationID, messageID := setupProcessDetailsSummaryTest(t)
for _, id := range []string{"call-1", "call-2", "call-3", "call-4"} {
if err := db.AddProcessDetail(messageID, conversationID, "tool_call", "call", map[string]interface{}{
"toolName": "http-framework-test", "toolCallId": id,
}); err != nil {
t.Fatalf("AddProcessDetail(tool_call): %v", err)
}
}
results := []map[string]interface{}{
{"toolName": "http-framework-test", "toolCallId": "call-1", "success": true},
{"toolName": "http-framework-test", "toolCallId": "call-2", "success": true},
{"toolName": "http-framework-test", "success": true},
{"toolName": "http-framework-test", "success": true},
}
for _, result := range results {
if err := db.AddProcessDetail(messageID, conversationID, "tool_result", "result", result); err != nil {
t.Fatalf("AddProcessDetail(tool_result): %v", err)
}
}
summary, err := db.GetProcessDetailsSummary(messageID)
if err != nil {
t.Fatalf("GetProcessDetailsSummary: %v", err)
}
if len(summary.ToolExecutions) != 4 {
t.Fatalf("tool executions = %d, want 4", len(summary.ToolExecutions))
}
for i, execution := range summary.ToolExecutions {
if execution.Status != "completed" {
t.Fatalf("execution %d status = %q, want completed", i, execution.Status)
}
}
}
func TestProcessDetailsSummaryPairsRepeatedToolCallIDsFIFO(t *testing.T) {
db, conversationID, messageID := setupProcessDetailsSummaryTest(t)
for i := 0; i < 2; i++ {
if err := db.AddProcessDetail(messageID, conversationID, "tool_call", "call", map[string]interface{}{
"toolName": "execute", "toolCallId": "legacy-reused-id",
}); err != nil {
t.Fatalf("AddProcessDetail(tool_call): %v", err)
}
}
for i := 0; i < 2; i++ {
if err := db.AddProcessDetail(messageID, conversationID, "tool_result", "result", map[string]interface{}{
"toolName": "execute", "toolCallId": "legacy-reused-id", "success": true,
}); err != nil {
t.Fatalf("AddProcessDetail(tool_result): %v", err)
}
}
summary, err := db.GetProcessDetailsSummary(messageID)
if err != nil {
t.Fatalf("GetProcessDetailsSummary: %v", err)
}
if len(summary.ToolExecutions) != 2 {
t.Fatalf("tool executions = %d, want 2", len(summary.ToolExecutions))
}
for i, execution := range summary.ToolExecutions {
if execution.Status != "completed" {
t.Fatalf("execution %d status = %q, want completed", i, execution.Status)
}
}
}
func TestProcessDetailsSummaryDoesNotReportPersistedOrphanAsRunning(t *testing.T) {
db, conversationID, messageID := setupProcessDetailsSummaryTest(t)
if err := db.AddProcessDetail(messageID, conversationID, "tool_call", "call", map[string]interface{}{
"toolName": "execute", "toolCallId": "orphan",
}); err != nil {
t.Fatalf("AddProcessDetail(tool_call): %v", err)
}
summary, err := db.GetProcessDetailsSummary(messageID)
if err != nil {
t.Fatalf("GetProcessDetailsSummary: %v", err)
}
if len(summary.ToolExecutions) != 1 || summary.ToolExecutions[0].Status != "result_missing" {
t.Fatalf("tool executions = %#v, want result_missing", summary.ToolExecutions)
}
}
func setupProcessDetailsSummaryTest(t *testing.T) (*DB, string, string) {
t.Helper()
db, err := NewDB(filepath.Join(t.TempDir(), "process-details.db"), zap.NewNop())
if err != nil {
t.Fatalf("NewDB: %v", err)
}
t.Cleanup(func() { _ = db.Close() })
conversation, err := db.CreateConversation("process details", ConversationCreateMeta{})
if err != nil {
t.Fatalf("CreateConversation: %v", err)
}
message, err := db.AddMessage(conversation.ID, "assistant", "done", nil)
if err != nil {
t.Fatalf("AddMessage: %v", err)
}
return db, conversation.ID, message.ID
}
+66 -5
View File
@@ -58,11 +58,11 @@ type ProjectFact struct {
// ProjectFactListFilter 事实列表筛选。
type ProjectFactListFilter struct {
Category string
Confidence string
Search string
RelatedVulnerabilityID string
ExcludeDeprecated bool // 为 true 时排除 confidence=deprecated
Category string
Confidence string
Search string
RelatedVulnerabilityID string
ExcludeDeprecated bool // 为 true 时排除 confidence=deprecated
}
// CreateProject 创建项目。
@@ -143,6 +143,19 @@ func appendProjectListFilters(query string, args []interface{}, status, search s
return query, args
}
func appendProjectAccessFilter(query string, args []interface{}, userID, scope string) (string, []interface{}) {
userID = strings.TrimSpace(userID)
if userID == "" || scope == RBACScopeAll {
return query, args
}
query += ` AND (owner_user_id = ? OR EXISTS (
SELECT 1 FROM rbac_resource_assignments ra
WHERE ra.user_id = ? AND ra.resource_type = 'project' AND ra.resource_id = projects.id
))`
args = append(args, userID, userID)
return query, args
}
// CountProjects 统计项目数量。
func (db *DB) CountProjects(status, search string) (int, error) {
query := `SELECT COUNT(*) FROM projects WHERE 1=1`
@@ -155,6 +168,18 @@ func (db *DB) CountProjects(status, search string) (int, error) {
return count, nil
}
func (db *DB) CountProjectsForAccess(status, search, userID, scope string) (int, error) {
query := `SELECT COUNT(*) FROM projects WHERE 1=1`
args := []interface{}{}
query, args = appendProjectListFilters(query, args, status, search)
query, args = appendProjectAccessFilter(query, args, userID, scope)
var count int
if err := db.QueryRow(query, args...).Scan(&count); err != nil {
return 0, fmt.Errorf("统计项目失败: %w", err)
}
return count, nil
}
// ListProjects 列出项目。
func (db *DB) ListProjects(status, search string, limit, offset int) ([]*Project, error) {
if limit <= 0 {
@@ -189,6 +214,42 @@ func (db *DB) ListProjects(status, search string, limit, offset int) ([]*Project
return out, rows.Err()
}
func (db *DB) ListProjectsForAccess(status, search string, limit, offset int, userID, scope string) ([]*Project, error) {
if scope == RBACScopeAll || strings.TrimSpace(userID) == "" {
return db.ListProjects(status, search, limit, offset)
}
if limit <= 0 {
limit = 50
}
query := `SELECT id, name, COALESCE(description,''), COALESCE(scope_json,''), status, pinned, created_at, updated_at
FROM projects WHERE 1=1`
args := []interface{}{}
query, args = appendProjectListFilters(query, args, status, search)
query, args = appendProjectAccessFilter(query, args, userID, scope)
query += " ORDER BY pinned DESC, updated_at DESC LIMIT ? OFFSET ?"
args = append(args, limit, offset)
rows, err := db.Query(query, args...)
if err != nil {
return nil, fmt.Errorf("列出项目失败: %w", err)
}
defer rows.Close()
var out []*Project
for rows.Next() {
var p Project
var pinned int
var createdAt, updatedAt string
if err := rows.Scan(&p.ID, &p.Name, &p.Description, &p.ScopeJSON, &p.Status, &pinned, &createdAt, &updatedAt); err != nil {
return nil, err
}
p.Pinned = pinned != 0
p.CreatedAt = parseDBTime(createdAt)
p.UpdatedAt = parseDBTime(updatedAt)
out = append(out, &p)
}
return out, rows.Err()
}
// UpdateProject 更新项目。
func (db *DB) UpdateProject(p *Project) error {
p.UpdatedAt = time.Now()
+25 -4
View File
@@ -33,6 +33,10 @@ type ProjectDashboardSummary struct {
// GetProjectDashboardSummary 聚合跨项目近期事实(仅活跃项目、排除 deprecated)。
func (db *DB) GetProjectDashboardSummary(factLimit int) (*ProjectDashboardSummary, error) {
return db.GetProjectDashboardSummaryForAccess(factLimit, "", "")
}
func (db *DB) GetProjectDashboardSummaryForAccess(factLimit int, userID, scope string) (*ProjectDashboardSummary, error) {
if factLimit <= 0 {
factLimit = 5
}
@@ -44,25 +48,42 @@ func (db *DB) GetProjectDashboardSummary(factLimit int) (*ProjectDashboardSummar
RecentFacts: []ProjectDashboardFact{},
}
if err := db.QueryRow(`SELECT COUNT(*) FROM projects WHERE status = 'active'`).Scan(&out.Totals.ActiveProjects); err != nil {
projectAccess := ""
args := []interface{}{}
userID = strings.TrimSpace(userID)
if userID != "" && scope != RBACScopeAll {
projectAccess = ` AND (
p.owner_user_id = ?
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments ra
WHERE ra.user_id = ? AND ra.resource_type = 'project' AND ra.resource_id = p.id
)
)`
args = append(args, userID, userID)
}
if err := db.QueryRow(`SELECT COUNT(*) FROM projects p WHERE p.status = 'active'`+projectAccess, args...).Scan(&out.Totals.ActiveProjects); err != nil {
return nil, fmt.Errorf("统计活跃项目失败: %w", err)
}
if err := db.QueryRow(
`SELECT COUNT(*) FROM project_facts f
INNER JOIN projects p ON p.id = f.project_id
WHERE f.confidence != 'deprecated' AND p.status = 'active'`,
WHERE f.confidence != 'deprecated' AND p.status = 'active'`+projectAccess,
args...,
).Scan(&out.Totals.TotalFacts); err != nil {
return nil, fmt.Errorf("统计事实失败: %w", err)
}
queryArgs := append([]interface{}{}, args...)
queryArgs = append(queryArgs, factLimit)
rows, err := db.Query(
`SELECT f.id, f.project_id, p.name, f.fact_key, f.category, f.summary, f.confidence, f.pinned, f.updated_at
FROM project_facts f
INNER JOIN projects p ON p.id = f.project_id
WHERE f.confidence != 'deprecated' AND p.status = 'active'
WHERE f.confidence != 'deprecated' AND p.status = 'active'`+projectAccess+`
ORDER BY f.pinned DESC, f.updated_at DESC
LIMIT ?`,
factLimit,
queryArgs...,
)
if err != nil {
return nil, fmt.Errorf("查询近期事实失败: %w", err)
File diff suppressed because it is too large Load Diff
+621
View File
@@ -0,0 +1,621 @@
package database
import (
"path/filepath"
"strings"
"testing"
"time"
"cyberstrike-ai/internal/mcp"
"go.uber.org/zap"
)
func newRBACTestDB(t *testing.T) *DB {
t.Helper()
db, err := NewDB(filepath.Join(t.TempDir(), "rbac.db"), zap.NewNop())
if err != nil {
t.Fatalf("NewDB: %v", err)
}
t.Cleanup(func() { _ = db.Close() })
return db
}
func TestRBACToolExecutionOwnershipAccess(t *testing.T) {
db := newRBACTestDB(t)
for _, exec := range []*mcp.ToolExecution{
{ID: "exec-u1", ToolName: "one", Status: "completed", StartTime: time.Now(), OwnerUserID: "u1"},
{ID: "exec-u2", ToolName: "two", Status: "completed", StartTime: time.Now(), OwnerUserID: "u2"},
{ID: "exec-legacy", ToolName: "legacy", Status: "completed", StartTime: time.Now()},
} {
if err := db.SaveToolExecution(exec); err != nil {
t.Fatal(err)
}
}
access := RBACListAccess{UserID: "u1", Scope: RBACScopeAssigned}
rows, err := db.LoadToolExecutionListPageForAccess(0, 20, "", "", access)
if err != nil {
t.Fatal(err)
}
if len(rows) != 1 || rows[0].ID != "exec-u1" {
t.Fatalf("rows = %#v, want only exec-u1", rows)
}
summary, err := db.LoadToolStatsSummaryForAccess(10, access)
if err != nil {
t.Fatal(err)
}
if summary.Summary.TotalCalls != 1 || summary.Summary.ToolCount != 1 || len(summary.TopTools) != 1 || summary.TopTools[0].ToolName != "one" {
t.Fatalf("scoped summary = %#v", summary)
}
if !db.UserCanAccessToolExecution("u1", RBACScopeAssigned, "exec-u1") {
t.Fatal("owner could not access execution")
}
if db.UserCanAccessToolExecution("u1", RBACScopeAssigned, "exec-u2") {
t.Fatal("foreign execution was accessible")
}
if db.UserCanAccessToolExecution("u1", RBACScopeAssigned, "exec-legacy") {
t.Fatal("ownerless legacy execution did not fail closed")
}
}
func TestRBACGroupAndUploadOwnership(t *testing.T) {
db := newRBACTestDB(t)
group1, err := db.CreateGroup("u1 group", "", "u1")
if err != nil {
t.Fatal(err)
}
group2, err := db.CreateGroup("u2 group", "", "u2")
if err != nil {
t.Fatal(err)
}
groups, err := db.ListGroupsForAccess("u1", RBACScopeAssigned)
if err != nil {
t.Fatal(err)
}
if len(groups) != 1 || groups[0].ID != group1.ID {
t.Fatalf("groups = %#v, want only %s (not %s)", groups, group1.ID, group2.ID)
}
if db.UserCanAccessGroup("u1", RBACScopeAssigned, group2.ID) {
t.Fatal("foreign group was accessible")
}
conversation, err := db.CreateConversation("upload", ConversationCreateMeta{})
if err != nil {
t.Fatal(err)
}
if err := db.UpsertChatUploadArtifact("2026-07-10/"+conversation.ID+"/a.txt", conversation.ID, "u1"); err != nil {
t.Fatal(err)
}
if conv, owner, ok := db.GetChatUploadArtifact("2026-07-10/" + conversation.ID + "/a.txt"); !ok || conv != conversation.ID || owner != "u1" {
t.Fatalf("artifact = conv=%q owner=%q ok=%v", conv, owner, ok)
}
if err := db.RenameChatUploadArtifactPath("2026-07-10/"+conversation.ID+"/a.txt", "2026-07-10/"+conversation.ID+"/b.txt"); err != nil {
t.Fatal(err)
}
if _, _, ok := db.GetChatUploadArtifact("2026-07-10/" + conversation.ID + "/b.txt"); !ok {
t.Fatal("renamed artifact metadata missing")
}
}
func TestSystemRoleBootstrapDoesNotLeakManagementReadPermissions(t *testing.T) {
db := newRBACTestDB(t)
catalog := map[string]string{
"auth:self": "self", "project:read": "projects", "project:write": "project writes",
"agent:local-execute": "local tools",
"rbac:read": "rbac", "config:read": "config", "audit:read": "audit", "terminal:execute": "terminal",
"mcp:execute": "invoke", "mcp:write": "manage", "mcp:external:execute": "external invoke",
"workflow:execute": "run", "workflow:write": "manage definitions", "knowledge:write": "manage knowledge",
}
if err := db.BootstrapRBAC("hash", catalog); err != nil {
t.Fatal(err)
}
viewer, err := db.CreateRBACUser("viewer-policy", "Viewer", "hash", true, []string{RBACSystemRoleViewer})
if err != nil {
t.Fatal(err)
}
viewerAccess, err := db.ResolveRBACAccess(viewer.ID)
if err != nil {
t.Fatal(err)
}
if !viewerAccess.Permissions["project:read"] || viewerAccess.Permissions["rbac:read"] || viewerAccess.Permissions["config:read"] || viewerAccess.Permissions["audit:read"] {
t.Fatalf("unexpected viewer permissions: %#v", viewerAccess.Permissions)
}
auditor, err := db.CreateRBACUser("auditor-policy", "Auditor", "hash", true, []string{RBACSystemRoleAuditor})
if err != nil {
t.Fatal(err)
}
auditorAccess, err := db.ResolveRBACAccess(auditor.ID)
if err != nil {
t.Fatal(err)
}
if !auditorAccess.Permissions["audit:read"] || auditorAccess.Permissions["config:read"] || auditorAccess.Permissions["rbac:read"] {
t.Fatalf("unexpected auditor permissions: %#v", auditorAccess.Permissions)
}
operator, err := db.CreateRBACUser("operator-policy", "Operator", "hash", true, []string{RBACSystemRoleOperator})
if err != nil {
t.Fatal(err)
}
operatorAccess, err := db.ResolveRBACAccess(operator.ID)
if err != nil {
t.Fatal(err)
}
if !operatorAccess.Permissions["mcp:execute"] || operatorAccess.Permissions["mcp:write"] || operatorAccess.Permissions["mcp:external:execute"] {
t.Fatalf("unexpected operator MCP permissions: %#v", operatorAccess.Permissions)
}
if !operatorAccess.Permissions["workflow:execute"] || operatorAccess.Permissions["workflow:write"] || operatorAccess.Permissions["knowledge:write"] {
t.Fatalf("operator received global definition mutation permissions: %#v", operatorAccess.Permissions)
}
if !operatorAccess.Permissions["agent:local-execute"] {
t.Fatalf("operator is missing explicit local tool permission: %#v", operatorAccess.Permissions)
}
}
func TestPermissionScopeDoesNotWidenAcrossUnrelatedRoles(t *testing.T) {
db := newRBACTestDB(t)
catalog := map[string]string{"auth:self": "self", "project:read": "read", "project:write": "write", "audit:read": "audit"}
if err := db.BootstrapRBAC("hash", catalog); err != nil {
t.Fatal(err)
}
ownWrite, err := db.UpsertRBACRole("", "own-writer", "", RBACScopeOwn, []string{"project:write"})
if err != nil {
t.Fatal(err)
}
user, err := db.CreateRBACUser("mixed-scope", "Mixed", "hash", true, []string{RBACSystemRoleAuditor, ownWrite.ID})
if err != nil {
t.Fatal(err)
}
access, err := db.ResolveRBACAccess(user.ID)
if err != nil {
t.Fatal(err)
}
if access.Scope != RBACScopeAll {
t.Fatalf("compatibility scope = %q, want all", access.Scope)
}
if got := access.PermissionScopes["project:read"]; got != RBACScopeAll {
t.Fatalf("project:read scope = %q, want all", got)
}
if got := access.PermissionScopes["project:write"]; got != RBACScopeOwn {
t.Fatalf("project:write scope widened to %q, want own", got)
}
}
func TestRoleRejectsUnknownPermission(t *testing.T) {
db := newRBACTestDB(t)
if err := db.BootstrapRBAC("hash", map[string]string{"auth:self": "self"}); err != nil {
t.Fatal(err)
}
if _, err := db.UpsertRBACRole("", "future-role", "", RBACScopeAssigned, []string{"future:permission"}); err == nil {
t.Fatal("unknown permission was persisted")
}
if _, err := db.Exec(`INSERT INTO rbac_permissions (key, description, created_at) VALUES ('stale:permission', '', ?)`, time.Now()); err != nil {
t.Fatal(err)
}
if err := db.BootstrapRBAC("hash", map[string]string{"auth:self": "self"}); err != nil {
t.Fatal(err)
}
var count int
if err := db.QueryRow(`SELECT COUNT(*) FROM rbac_permissions WHERE key = 'stale:permission'`).Scan(&count); err != nil || count != 0 {
t.Fatalf("stale permission survived bootstrap: count=%d err=%v", count, err)
}
}
func TestRBACProjectAndConversationListAccess(t *testing.T) {
db := newRBACTestDB(t)
p1, _ := db.CreateProject(&Project{Name: "visible"})
p2, _ := db.CreateProject(&Project{Name: "hidden"})
if err := db.SetResourceOwner("project", p1.ID, "u1"); err != nil {
t.Fatal(err)
}
c1, _ := db.CreateConversation("visible conv", ConversationCreateMeta{ProjectID: p1.ID})
c2, _ := db.CreateConversation("hidden conv", ConversationCreateMeta{ProjectID: p2.ID})
_ = db.SetResourceOwner("conversation", c1.ID, "u1")
_ = db.SetResourceOwner("conversation", c2.ID, "u2")
projects, err := db.ListProjectsForAccess("", "", 50, 0, "u1", RBACScopeOwn)
if err != nil {
t.Fatal(err)
}
if len(projects) != 1 || projects[0].ID != p1.ID {
t.Fatalf("projects = %#v, want only %s", projects, p1.ID)
}
convs, err := db.ListConversationsForAccess(50, 0, "", "", "", "u1", RBACScopeOwn)
if err != nil {
t.Fatal(err)
}
if len(convs) != 1 || convs[0].ID != c1.ID {
t.Fatalf("conversations = %#v, want only %s", convs, c1.ID)
}
}
func TestRBACVulnerabilityAccessInheritsProject(t *testing.T) {
db := newRBACTestDB(t)
user, err := db.CreateRBACUser("u1", "User 1", "hash", true, nil)
if err != nil {
t.Fatal(err)
}
p1, _ := db.CreateProject(&Project{Name: "visible"})
p2, _ := db.CreateProject(&Project{Name: "hidden"})
if err := db.AssignResourceToUser(user.ID, "project", p1.ID); err != nil {
t.Fatal(err)
}
v1, _ := db.CreateVulnerability(&Vulnerability{ProjectID: p1.ID, Title: "v1", Severity: "high"})
v2, _ := db.CreateVulnerability(&Vulnerability{ProjectID: p2.ID, Title: "v2", Severity: "high"})
items, err := db.ListVulnerabilitiesForAccess(50, 0, VulnerabilityListFilter{}, RBACListAccess{UserID: user.ID, Scope: RBACScopeAssigned})
if err != nil {
t.Fatal(err)
}
if len(items) != 1 || items[0].ID != v1.ID {
t.Fatalf("vulnerabilities = %#v, want only %s; hidden %s", items, v1.ID, v2.ID)
}
if !db.UserCanAccessResource(user.ID, RBACScopeAssigned, "vulnerability", v1.ID) {
t.Fatalf("expected project assignment to allow vulnerability detail")
}
if db.UserCanAccessResource(user.ID, RBACScopeAssigned, "vulnerability", v2.ID) {
t.Fatalf("unexpected access to hidden vulnerability")
}
}
func TestRBACConversationAccessInheritsProject(t *testing.T) {
db := newRBACTestDB(t)
user, err := db.CreateRBACUser("project-member", "Project Member", "hash", true, nil)
if err != nil {
t.Fatal(err)
}
project, err := db.CreateProject(&Project{Name: "assigned project"})
if err != nil {
t.Fatal(err)
}
conversation, err := db.CreateConversation("project conversation", ConversationCreateMeta{ProjectID: project.ID})
if err != nil {
t.Fatal(err)
}
if err := db.AssignResourceToUser(user.ID, "project", project.ID); err != nil {
t.Fatal(err)
}
rows, err := db.ListConversationsForAccess(50, 0, "", "", "", user.ID, RBACScopeAssigned)
if err != nil {
t.Fatal(err)
}
if len(rows) != 1 || rows[0].ID != conversation.ID {
t.Fatalf("conversations = %#v, want project conversation %s", rows, conversation.ID)
}
if !db.UserCanAccessResource(user.ID, RBACScopeAssigned, "conversation", conversation.ID) {
t.Fatal("expected project assignment to allow conversation detail")
}
}
func TestRBACBatchResourceAssignmentValidationAndAtomicity(t *testing.T) {
db := newRBACTestDB(t)
user, err := db.CreateRBACUser("batch-member", "Batch Member", "hash", true, nil)
if err != nil {
t.Fatal(err)
}
p1, err := db.CreateProject(&Project{Name: "p1"})
if err != nil {
t.Fatal(err)
}
p2, err := db.CreateProject(&Project{Name: "p2"})
if err != nil {
t.Fatal(err)
}
p3, err := db.CreateProject(&Project{Name: "p3"})
if err != nil {
t.Fatal(err)
}
options, err := db.ListAssignableRBACResources("project", "p1", 50)
if err != nil {
t.Fatal(err)
}
if len(options) != 1 || options[0].ID != p1.ID || options[0].Label != "p1" {
t.Fatalf("resource options = %#v, want p1", options)
}
firstPage, err := db.ListAssignableRBACResourcesPage("project", "", 2, 0)
if err != nil {
t.Fatal(err)
}
secondPage, err := db.ListAssignableRBACResourcesPage("project", "", 2, 2)
if err != nil {
t.Fatal(err)
}
if len(firstPage) != 2 || len(secondPage) != 1 {
t.Fatalf("paged resource options = %d + %d, want 2 + 1", len(firstPage), len(secondPage))
}
seen := map[string]bool{}
for _, option := range append(firstPage, secondPage...) {
seen[option.ID] = true
}
if !seen[p1.ID] || !seen[p2.ID] || !seen[p3.ID] {
t.Fatalf("paged resource options missed resources: %#v", seen)
}
if _, err := db.ListAssignableRBACResources("secret_table", "", 50); err == nil {
t.Fatal("expected unsupported picker resource type to fail")
}
if _, err := db.AssignResourcesToUser(user.ID, "unknown_type", []string{p1.ID}); err == nil {
t.Fatal("expected unsupported resource type to fail")
}
if _, err := db.AssignResourcesToUser(user.ID, "project", []string{p1.ID, "missing-project"}); err == nil {
t.Fatal("expected missing resource to fail the entire batch")
}
rows, err := db.ListRBACResourceAssignments(user.ID)
if err != nil {
t.Fatal(err)
}
if len(rows) != 0 {
t.Fatalf("partial grants persisted after failed batch: %#v", rows)
}
created, err := db.AssignResourcesToUser(user.ID, "project", []string{p1.ID, p1.ID, p2.ID})
if err != nil {
t.Fatal(err)
}
if created != 2 {
t.Fatalf("created = %d, want 2 unique grants", created)
}
created, err = db.AssignResourcesToUser(user.ID, "project", []string{p1.ID, p2.ID})
if err != nil {
t.Fatal(err)
}
if created != 0 {
t.Fatalf("idempotent retry created = %d, want 0", created)
}
rows, err = db.ListRBACResourceAssignments(user.ID)
if err != nil {
t.Fatal(err)
}
if len(rows) != 2 {
t.Fatalf("assignment count = %d, want 2", len(rows))
}
}
func TestRBACWebshellAndBatchListAccess(t *testing.T) {
db := newRBACTestDB(t)
ws1 := WebShellConnection{ID: "ws_visible", URL: "http://a", Type: "php", Method: "post", CreatedAt: time.Now()}
ws2 := WebShellConnection{ID: "ws_hidden", URL: "http://b", Type: "php", Method: "post", CreatedAt: time.Now()}
if err := db.CreateWebshellConnection(&ws1); err != nil {
t.Fatal(err)
}
if err := db.CreateWebshellConnection(&ws2); err != nil {
t.Fatal(err)
}
_ = db.SetResourceOwner("webshell", ws1.ID, "u1")
_ = db.SetResourceOwner("webshell", ws2.ID, "u2")
webshells, err := db.ListWebshellConnectionsForAccess("u1", RBACScopeOwn)
if err != nil {
t.Fatal(err)
}
if len(webshells) != 1 || webshells[0].ID != ws1.ID {
t.Fatalf("webshells = %#v, want only %s", webshells, ws1.ID)
}
if err := db.CreateBatchQueue("q_visible", "visible", "", "eino_single", "manual", "", nil, "", 1, []map[string]interface{}{{"id": "t1", "message": "a"}}); err != nil {
t.Fatal(err)
}
if err := db.CreateBatchQueue("q_hidden", "hidden", "", "eino_single", "manual", "", nil, "", 1, []map[string]interface{}{{"id": "t2", "message": "b"}}); err != nil {
t.Fatal(err)
}
_ = db.SetResourceOwner("batch_task", "q_visible", "u1")
_ = db.SetResourceOwner("batch_task", "q_hidden", "u2")
queues, err := db.ListBatchQueuesForAccess(50, 0, "all", "", "u1", RBACScopeOwn)
if err != nil {
t.Fatal(err)
}
if len(queues) != 1 || queues[0].ID != "q_visible" {
t.Fatalf("queues = %#v, want only q_visible", queues)
}
}
func TestRBACC2AccessInheritsListener(t *testing.T) {
db := newRBACTestDB(t)
now := time.Now()
l1 := &C2Listener{ID: "l_visible", Name: "visible", Type: "http_beacon", BindHost: "127.0.0.1", BindPort: 9001, OwnerUserID: "u1", CreatedAt: now}
l2 := &C2Listener{ID: "l_hidden", Name: "hidden", Type: "http_beacon", BindHost: "127.0.0.1", BindPort: 9002, OwnerUserID: "u2", CreatedAt: now}
if err := db.CreateC2Listener(l1); err != nil {
t.Fatal(err)
}
if err := db.CreateC2Listener(l2); err != nil {
t.Fatal(err)
}
if err := db.UpsertC2Session(&C2Session{ID: "s_visible", ListenerID: l1.ID, ImplantUUID: "implant-visible", Status: "active", FirstSeenAt: now, LastCheckIn: now}); err != nil {
t.Fatal(err)
}
if err := db.UpsertC2Session(&C2Session{ID: "s_hidden", ListenerID: l2.ID, ImplantUUID: "implant-hidden", Status: "active", FirstSeenAt: now, LastCheckIn: now}); err != nil {
t.Fatal(err)
}
if err := db.CreateC2Task(&C2Task{ID: "t_visible", SessionID: "s_visible", TaskType: "shell", Status: "queued", CreatedAt: now}); err != nil {
t.Fatal(err)
}
if err := db.CreateC2Task(&C2Task{ID: "t_hidden", SessionID: "s_hidden", TaskType: "shell", Status: "queued", CreatedAt: now}); err != nil {
t.Fatal(err)
}
if err := db.AppendC2Event(&C2Event{ID: "e_visible", Level: "info", Category: "task", SessionID: "s_visible", TaskID: "t_visible", Message: "visible", CreatedAt: now}); err != nil {
t.Fatal(err)
}
if err := db.AppendC2Event(&C2Event{ID: "e_hidden", Level: "info", Category: "task", SessionID: "s_hidden", TaskID: "t_hidden", Message: "hidden", CreatedAt: now}); err != nil {
t.Fatal(err)
}
access := RBACListAccess{UserID: "u1", Scope: RBACScopeOwn}
listeners, err := db.ListC2ListenersForAccess(access)
if err != nil {
t.Fatal(err)
}
if len(listeners) != 1 || listeners[0].ID != l1.ID {
t.Fatalf("listeners = %#v, want only %s", listeners, l1.ID)
}
sessions, err := db.ListC2SessionsForAccess(ListC2SessionsFilter{}, access)
if err != nil {
t.Fatal(err)
}
if len(sessions) != 1 || sessions[0].ID != "s_visible" {
t.Fatalf("sessions = %#v, want only s_visible", sessions)
}
tasks, err := db.ListC2TasksForAccess(ListC2TasksFilter{}, access)
if err != nil {
t.Fatal(err)
}
if len(tasks) != 1 || tasks[0].ID != "t_visible" {
t.Fatalf("tasks = %#v, want only t_visible", tasks)
}
events, err := db.ListC2EventsForAccess(ListC2EventsFilter{}, access)
if err != nil {
t.Fatal(err)
}
if len(events) != 1 || events[0].ID != "e_visible" {
t.Fatalf("events = %#v, want only e_visible", events)
}
if !db.UserCanAccessResource("u1", RBACScopeOwn, "c2_task", "t_visible") {
t.Fatalf("expected listener ownership to allow task detail")
}
if db.UserCanAccessResource("u1", RBACScopeOwn, "c2_task", "t_hidden") {
t.Fatalf("unexpected access to hidden task")
}
}
func TestRBACC2AssignedDeleteIsScoped(t *testing.T) {
db := newRBACTestDB(t)
user, err := db.CreateRBACUser("u1", "User 1", "hash", true, nil)
if err != nil {
t.Fatal(err)
}
now := time.Now()
if err := db.CreateC2Listener(&C2Listener{ID: "l_assigned", Name: "assigned", Type: "http_beacon", BindHost: "127.0.0.1", BindPort: 9001, CreatedAt: now}); err != nil {
t.Fatal(err)
}
if err := db.CreateC2Listener(&C2Listener{ID: "l_hidden", Name: "hidden", Type: "http_beacon", BindHost: "127.0.0.1", BindPort: 9002, CreatedAt: now}); err != nil {
t.Fatal(err)
}
if err := db.AssignResourceToUser(user.ID, "c2_listener", "l_assigned"); err != nil {
t.Fatal(err)
}
for _, row := range []struct {
sessionID string
listener string
taskID string
eventID string
}{
{"s_assigned", "l_assigned", "t_assigned", "e_assigned"},
{"s_hidden", "l_hidden", "t_hidden", "e_hidden"},
} {
if err := db.UpsertC2Session(&C2Session{ID: row.sessionID, ListenerID: row.listener, ImplantUUID: row.sessionID + "_uuid", Status: "active", FirstSeenAt: now, LastCheckIn: now}); err != nil {
t.Fatal(err)
}
if err := db.CreateC2Task(&C2Task{ID: row.taskID, SessionID: row.sessionID, TaskType: "shell", Status: "queued", CreatedAt: now}); err != nil {
t.Fatal(err)
}
if err := db.AppendC2Event(&C2Event{ID: row.eventID, Level: "info", Category: "task", SessionID: row.sessionID, TaskID: row.taskID, Message: row.eventID, CreatedAt: now}); err != nil {
t.Fatal(err)
}
}
access := RBACListAccess{UserID: user.ID, Scope: RBACScopeAssigned}
n, err := db.DeleteC2TasksByIDsForAccess([]string{"t_assigned", "t_hidden"}, access)
if err != nil {
t.Fatal(err)
}
if n != 1 {
t.Fatalf("deleted tasks = %d, want 1", n)
}
if task, _ := db.GetC2Task("t_hidden"); task == nil {
t.Fatalf("hidden task was deleted")
}
n, err = db.DeleteC2EventsByIDsForAccess([]string{"e_assigned", "e_hidden"}, access)
if err != nil {
t.Fatal(err)
}
if n != 1 {
t.Fatalf("deleted events = %d, want 1", n)
}
hiddenEvents, err := db.ListC2Events(ListC2EventsFilter{TaskID: "t_hidden"})
if err != nil {
t.Fatal(err)
}
if len(hiddenEvents) != 1 {
t.Fatalf("hidden event count = %d, want 1", len(hiddenEvents))
}
}
func TestRBACAssignmentLabelsAndWeakTitles(t *testing.T) {
db := newRBACTestDB(t)
user, err := db.CreateRBACUser("label-member", "Label Member", "hash", true, nil)
if err != nil {
t.Fatal(err)
}
project, err := db.CreateProject(&Project{Name: "Alpha Project"})
if err != nil {
t.Fatal(err)
}
conversation, err := db.CreateConversation("1", ConversationCreateMeta{})
if err != nil {
t.Fatal(err)
}
if _, err := db.AssignResourcesToUser(user.ID, "project", []string{project.ID}); err != nil {
t.Fatal(err)
}
options, err := db.ListAssignableRBACResources("conversation", "", 10)
if err != nil {
t.Fatal(err)
}
if len(options) == 0 {
t.Fatal("expected conversation options")
}
for _, option := range options {
if option.ID == conversation.ID && !strings.Contains(option.Label, "1 ·") {
t.Fatalf("weak conversation label = %q, want suffix with short id", option.Label)
}
}
rows, err := db.ListRBACResourceAssignments(user.ID)
if err != nil {
t.Fatal(err)
}
if len(rows) != 1 {
t.Fatalf("assignments = %#v, want 1", rows)
}
if rows[0].ResourceLabel != "Alpha Project" {
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")
}
}
+174
View File
@@ -0,0 +1,174 @@
package database
import (
"database/sql"
"fmt"
"strings"
"time"
"github.com/google/uuid"
)
// RobotUserBinding maps one tenant-scoped platform identity to one RBAC user.
// external_user_id must be derived from the verified platform event, never
// from user-controlled message content.
type RobotUserBinding struct {
ID string `json:"id"`
Platform string `json:"platform"`
ExternalUserID string `json:"externalUserId"`
RBACUserID string `json:"rbacUserId"`
Enabled bool `json:"enabled"`
CreatedAt time.Time `json:"createdAt"`
UpdatedAt time.Time `json:"updatedAt"`
}
func normalizeRobotIdentity(platform, externalUserID string) (string, string, error) {
platform = strings.ToLower(strings.TrimSpace(platform))
externalUserID = strings.TrimSpace(externalUserID)
if platform == "" || externalUserID == "" {
return "", "", fmt.Errorf("robot platform and external user identity are required")
}
return platform, externalUserID, nil
}
func (db *DB) CreateRobotBindingCode(userID, codeHash string, expiresAt time.Time) error {
userID = strings.TrimSpace(userID)
codeHash = strings.TrimSpace(codeHash)
if userID == "" || codeHash == "" || !expiresAt.After(time.Now()) {
return fmt.Errorf("invalid robot binding code")
}
now := time.Now()
tx, err := db.Begin()
if err != nil {
return err
}
defer func() { _ = tx.Rollback() }()
// Keep only the newest active code per user and remove expired/used secrets.
if _, err = tx.Exec(`DELETE FROM robot_binding_codes WHERE rbac_user_id = ? OR expires_at <= ? OR used_at IS NOT NULL`, userID, now); err != nil {
return err
}
if _, err = tx.Exec(`INSERT INTO robot_binding_codes (code_hash, rbac_user_id, expires_at, created_at) VALUES (?, ?, ?, ?)`, codeHash, userID, expiresAt, now); err != nil {
return err
}
return tx.Commit()
}
// ConsumeRobotBindingCode atomically consumes a single-use code and binds the
// verified platform identity. Existing bindings are deliberately replaced so
// users can recover from stale or incorrect associations with a fresh code.
func (db *DB) ConsumeRobotBindingCode(platform, externalUserID, codeHash string) (*RBACUser, error) {
platform, externalUserID, err := normalizeRobotIdentity(platform, externalUserID)
if err != nil {
return nil, err
}
codeHash = strings.TrimSpace(codeHash)
if codeHash == "" {
return nil, fmt.Errorf("binding code is required")
}
tx, err := db.Begin()
if err != nil {
return nil, err
}
defer func() { _ = tx.Rollback() }()
var userID string
now := time.Now()
if err = tx.QueryRow(`
SELECT c.rbac_user_id
FROM robot_binding_codes c
JOIN rbac_users u ON u.id = c.rbac_user_id
WHERE c.code_hash = ? AND c.used_at IS NULL AND c.expires_at > ? AND u.enabled = 1
`, codeHash, now).Scan(&userID); err != nil {
if err == sql.ErrNoRows {
return nil, fmt.Errorf("binding code is invalid or expired")
}
return nil, err
}
result, err := tx.Exec(`UPDATE robot_binding_codes SET used_at = ? WHERE code_hash = ? AND used_at IS NULL`, now, codeHash)
if err != nil {
return nil, err
}
if affected, _ := result.RowsAffected(); affected != 1 {
return nil, fmt.Errorf("binding code has already been used")
}
if _, err = tx.Exec(`
INSERT INTO robot_user_bindings (id, platform, external_user_id, rbac_user_id, enabled, created_at, updated_at)
VALUES (?, ?, ?, ?, 1, ?, ?)
ON CONFLICT(platform, external_user_id) DO UPDATE SET
rbac_user_id = excluded.rbac_user_id,
enabled = 1,
updated_at = excluded.updated_at
`, uuid.New().String(), platform, externalUserID, userID, now, now); err != nil {
return nil, err
}
if err = tx.Commit(); err != nil {
return nil, err
}
return db.GetRBACUserByID(userID)
}
func (db *DB) ResolveRobotRBACAccess(platform, externalUserID string) (*RBACAccess, error) {
platform, externalUserID, err := normalizeRobotIdentity(platform, externalUserID)
if err != nil {
return nil, err
}
var userID string
err = db.QueryRow(`
SELECT b.rbac_user_id
FROM robot_user_bindings b
JOIN rbac_users u ON u.id = b.rbac_user_id
WHERE b.platform = ? AND b.external_user_id = ? AND b.enabled = 1 AND u.enabled = 1
`, platform, externalUserID).Scan(&userID)
if err == sql.ErrNoRows {
return nil, fmt.Errorf("robot identity is not bound")
}
if err != nil {
return nil, err
}
return db.ResolveRBACAccess(userID)
}
func (db *DB) ListRobotUserBindings(userID string) ([]RobotUserBinding, error) {
rows, err := db.Query(`
SELECT id, platform, external_user_id, rbac_user_id, enabled, created_at, updated_at
FROM robot_user_bindings WHERE rbac_user_id = ? ORDER BY updated_at DESC
`, strings.TrimSpace(userID))
if err != nil {
return nil, err
}
defer rows.Close()
var out []RobotUserBinding
for rows.Next() {
var b RobotUserBinding
var enabled int
var createdAt, updatedAt string
if err := rows.Scan(&b.ID, &b.Platform, &b.ExternalUserID, &b.RBACUserID, &enabled, &createdAt, &updatedAt); err != nil {
return nil, err
}
b.Enabled = enabled != 0
b.CreatedAt = parseDBTime(createdAt)
b.UpdatedAt = parseDBTime(updatedAt)
out = append(out, b)
}
return out, rows.Err()
}
func (db *DB) DeleteRobotUserBindingForUser(bindingID, userID string) error {
result, err := db.Exec(`DELETE FROM robot_user_bindings WHERE id = ? AND rbac_user_id = ?`, strings.TrimSpace(bindingID), strings.TrimSpace(userID))
if err != nil {
return err
}
if affected, _ := result.RowsAffected(); affected != 1 {
return sql.ErrNoRows
}
return nil
}
func (db *DB) DeleteRobotIdentityBinding(platform, externalUserID string) error {
platform, externalUserID, err := normalizeRobotIdentity(platform, externalUserID)
if err != nil {
return err
}
_, err = db.Exec(`DELETE FROM robot_user_bindings WHERE platform = ? AND external_user_id = ?`, platform, externalUserID)
return err
}
+89
View File
@@ -0,0 +1,89 @@
package database_test
import (
"testing"
"time"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/security"
"go.uber.org/zap"
)
func TestRobotBindingCodeIsSingleUseAndPermissionsAreResolvedLive(t *testing.T) {
db, err := database.NewDB(t.TempDir()+"/robot-identity.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("bound-user", "Bound User", "hash", true, []string{database.RBACSystemRoleOperator})
if err != nil {
t.Fatal(err)
}
if err := db.CreateRobotBindingCode(user.ID, "code-hash", time.Now().Add(time.Minute)); err != nil {
t.Fatal(err)
}
bound, err := db.ConsumeRobotBindingCode("LARK", "t:tenant|u:user", "code-hash")
if err != nil || bound.ID != user.ID {
t.Fatalf("consume binding code: user=%v err=%v", bound, err)
}
if _, err := db.ConsumeRobotBindingCode("lark", "t:tenant|u:other", "code-hash"); err == nil {
t.Fatal("single-use binding code was accepted twice")
}
access, err := db.ResolveRobotRBACAccess("lark", "t:tenant|u:user")
if err != nil || !access.Permissions["agent:execute"] {
t.Fatalf("resolved access does not include live role permissions: %#v err=%v", access, err)
}
disabled := false
if err := db.UpdateRBACUser(user.ID, user.DisplayName, &disabled, nil); err != nil {
t.Fatal(err)
}
if _, err := db.ResolveRobotRBACAccess("lark", "t:tenant|u:user"); err == nil {
t.Fatal("disabled RBAC user retained robot access")
}
}
func TestRobotBindingCodeExpiryAndOwnerScopedRevocation(t *testing.T) {
db, err := database.NewDB(t.TempDir()+"/robot-revoke.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)
}
u1, _ := db.CreateRBACUser("binding-owner", "Owner", "hash", true, nil)
u2, _ := db.CreateRBACUser("binding-other", "Other", "hash", true, nil)
now := time.Now()
if _, err := db.Exec(`INSERT INTO robot_binding_codes (code_hash, rbac_user_id, expires_at, created_at) VALUES (?, ?, ?, ?)`, "expired-hash", u1.ID, now.Add(-time.Minute), now.Add(-2*time.Minute)); err != nil {
t.Fatal(err)
}
if _, err := db.ConsumeRobotBindingCode("wecom", "t:corp|u:expired", "expired-hash"); err == nil {
t.Fatal("expired binding code was accepted")
}
if err := db.CreateRobotBindingCode(u1.ID, "valid-hash", time.Now().Add(time.Minute)); err != nil {
t.Fatal(err)
}
if _, err := db.ConsumeRobotBindingCode("wecom", "t:corp|u:one", "valid-hash"); err != nil {
t.Fatal(err)
}
bindings, err := db.ListRobotUserBindings(u1.ID)
if err != nil || len(bindings) != 1 {
t.Fatalf("bindings=%v err=%v", bindings, err)
}
if err := db.DeleteRobotUserBindingForUser(bindings[0].ID, u2.ID); err == nil {
t.Fatal("another user revoked a binding they do not own")
}
if _, err := db.ResolveRobotRBACAccess("wecom", "t:corp|u:one"); err != nil {
t.Fatalf("unauthorized revocation changed binding: %v", err)
}
if err := db.DeleteRobotUserBindingForUser(bindings[0].ID, u1.ID); err != nil {
t.Fatal(err)
}
if _, err := db.ResolveRobotRBACAccess("wecom", "t:corp|u:one"); err == nil {
t.Fatal("revoked binding still resolves")
}
}
+15 -6
View File
@@ -12,6 +12,7 @@ type RobotSessionBinding struct {
SessionKey string
ConversationID string
RoleName string
AgentMode string
UpdatedAt time.Time
}
@@ -24,9 +25,9 @@ func (db *DB) GetRobotSessionBinding(sessionKey string) (*RobotSessionBinding, e
var b RobotSessionBinding
var updatedAt string
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,
).Scan(&b.SessionKey, &b.ConversationID, &b.RoleName, &updatedAt)
).Scan(&b.SessionKey, &b.ConversationID, &b.RoleName, &b.AgentMode, &updatedAt)
if err != nil {
if err == sql.ErrNoRows {
return nil, nil
@@ -43,28 +44,36 @@ func (db *DB) GetRobotSessionBinding(sessionKey string) (*RobotSessionBinding, e
if strings.TrimSpace(b.RoleName) == "" {
b.RoleName = "默认"
}
if strings.TrimSpace(b.AgentMode) == "" {
b.AgentMode = "eino_single"
}
return &b, nil
}
// UpsertRobotSessionBinding 写入或更新机器人会话绑定(包含角色)。
func (db *DB) UpsertRobotSessionBinding(sessionKey, conversationID, roleName string) error {
func (db *DB) UpsertRobotSessionBinding(sessionKey, conversationID, roleName, agentMode string) error {
sessionKey = strings.TrimSpace(sessionKey)
conversationID = strings.TrimSpace(conversationID)
roleName = strings.TrimSpace(roleName)
agentMode = strings.TrimSpace(agentMode)
if sessionKey == "" || conversationID == "" {
return nil
}
if roleName == "" {
roleName = "默认"
}
if agentMode == "" {
agentMode = "eino_single"
}
_, err := db.Exec(`
INSERT INTO robot_user_sessions (session_key, conversation_id, role_name, updated_at)
VALUES (?, ?, ?, ?)
INSERT INTO robot_user_sessions (session_key, conversation_id, role_name, agent_mode, updated_at)
VALUES (?, ?, ?, ?, ?)
ON CONFLICT(session_key) DO UPDATE SET
conversation_id = excluded.conversation_id,
role_name = excluded.role_name,
agent_mode = excluded.agent_mode,
updated_at = excluded.updated_at
`, sessionKey, conversationID, roleName, time.Now())
`, sessionKey, conversationID, roleName, agentMode, time.Now())
if err != nil {
return fmt.Errorf("写入机器人会话绑定失败: %w", err)
}
+74 -8
View File
@@ -23,6 +23,11 @@ type VulnerabilityListFilter struct {
TaskTag string
}
type RBACListAccess struct {
UserID string
Scope string
}
func escapeVulnerabilityLikePattern(s string) string {
s = strings.ReplaceAll(s, `\`, `\\`)
s = strings.ReplaceAll(s, `%`, `\%`)
@@ -89,6 +94,40 @@ func (f VulnerabilityListFilter) appendWhere(query string, args []interface{}) (
return query, args
}
func appendVulnerabilityAccessFilter(query string, args []interface{}, access RBACListAccess) (string, []interface{}) {
userID := strings.TrimSpace(access.UserID)
if userID == "" || access.Scope == RBACScopeAll {
return query, args
}
query += ` AND (
owner_user_id = ?
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments ra
WHERE ra.user_id = ? AND ra.resource_type = 'vulnerability' AND ra.resource_id = vulnerabilities.id
)
OR (
project_id IS NOT NULL AND project_id <> '' AND (
EXISTS (SELECT 1 FROM projects p WHERE p.id = vulnerabilities.project_id AND p.owner_user_id = ?)
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments pra
WHERE pra.user_id = ? AND pra.resource_type = 'project' AND pra.resource_id = vulnerabilities.project_id
)
)
)
OR (
conversation_id IS NOT NULL AND conversation_id <> '' AND (
EXISTS (SELECT 1 FROM conversations c WHERE c.id = vulnerabilities.conversation_id AND c.owner_user_id = ?)
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments cra
WHERE cra.user_id = ? AND cra.resource_type = 'conversation' AND cra.resource_id = vulnerabilities.conversation_id
)
)
)
)`
args = append(args, userID, userID, userID, userID, userID, userID)
return query, args
}
// Vulnerability 漏洞
type Vulnerability struct {
ID string `json:"id"`
@@ -152,7 +191,6 @@ func (db *DB) CreateVulnerability(vuln *Vulnerability) (*Vulnerability, error) {
if err != nil {
return nil, fmt.Errorf("创建漏洞失败: %w", err)
}
return vuln, nil
}
@@ -190,6 +228,10 @@ func (db *DB) GetVulnerability(id string) (*Vulnerability, error) {
// ListVulnerabilities 列出漏洞
func (db *DB) ListVulnerabilities(limit, offset int, filter VulnerabilityListFilter) ([]*Vulnerability, error) {
return db.ListVulnerabilitiesForAccess(limit, offset, filter, RBACListAccess{})
}
func (db *DB) ListVulnerabilitiesForAccess(limit, offset int, filter VulnerabilityListFilter, access RBACListAccess) ([]*Vulnerability, error) {
query := `
SELECT id, COALESCE(conversation_id,''), COALESCE(project_id,''), title, description, severity, status, conversation_tag, task_tag,
vulnerability_type, target,
@@ -203,6 +245,7 @@ func (db *DB) ListVulnerabilities(limit, offset int, filter VulnerabilityListFil
`
args := []interface{}{}
query, args = filter.appendWhere(query, args)
query, args = appendVulnerabilityAccessFilter(query, args, access)
query += " ORDER BY created_at DESC LIMIT ? OFFSET ?"
args = append(args, limit, offset)
@@ -235,9 +278,14 @@ func (db *DB) ListVulnerabilities(limit, offset int, filter VulnerabilityListFil
// CountVulnerabilities 统计漏洞总数(支持筛选条件)
func (db *DB) CountVulnerabilities(filter VulnerabilityListFilter) (int, error) {
return db.CountVulnerabilitiesForAccess(filter, RBACListAccess{})
}
func (db *DB) CountVulnerabilitiesForAccess(filter VulnerabilityListFilter, access RBACListAccess) (int, error) {
query := "SELECT COUNT(*) FROM vulnerabilities WHERE 1=1"
args := []interface{}{}
query, args = filter.appendWhere(query, args)
query, args = appendVulnerabilityAccessFilter(query, args, access)
var count int
err := db.QueryRow(query, args...).Scan(&count)
@@ -275,6 +323,10 @@ func (db *DB) UpdateVulnerability(id string, vuln *Vulnerability) error {
// DeleteVulnerabilitiesByFilter 按筛选条件批量删除漏洞,返回实际删除条数
func (db *DB) DeleteVulnerabilitiesByFilter(filter VulnerabilityListFilter) (int64, error) {
return db.DeleteVulnerabilitiesByFilterForAccess(filter, RBACListAccess{})
}
func (db *DB) DeleteVulnerabilitiesByFilterForAccess(filter VulnerabilityListFilter, access RBACListAccess) (int64, error) {
tx, err := db.Begin()
if err != nil {
return 0, fmt.Errorf("开启事务失败: %w", err)
@@ -284,6 +336,7 @@ func (db *DB) DeleteVulnerabilitiesByFilter(filter VulnerabilityListFilter) (int
where := "WHERE 1=1"
args := []interface{}{}
where, args = filter.appendWhere(where, args)
where, args = appendVulnerabilityAccessFilter(where, args, access)
clearQuery := `UPDATE project_facts SET related_vulnerability_id = NULL
WHERE related_vulnerability_id IN (SELECT id FROM vulnerabilities ` + where + `)`
@@ -329,11 +382,16 @@ func (db *DB) DeleteVulnerability(id string) error {
// GetVulnerabilityStats 获取漏洞统计(筛选条件与 ListVulnerabilities / CountVulnerabilities 一致)
func (db *DB) GetVulnerabilityStats(filter VulnerabilityListFilter) (map[string]interface{}, error) {
return db.GetVulnerabilityStatsForAccess(filter, RBACListAccess{})
}
func (db *DB) GetVulnerabilityStatsForAccess(filter VulnerabilityListFilter, access RBACListAccess) (map[string]interface{}, error) {
stats := make(map[string]interface{})
where := "WHERE 1=1"
args := []interface{}{}
where, args = filter.appendWhere(where, args)
where, args = appendVulnerabilityAccessFilter(where, args, access)
// 总漏洞数
var totalCount int
@@ -389,6 +447,10 @@ func (db *DB) GetVulnerabilityStats(filter VulnerabilityListFilter) (map[string]
// GetVulnerabilityFilterOptions 获取漏洞筛选建议项
func (db *DB) GetVulnerabilityFilterOptions() (map[string][]string, error) {
return db.GetVulnerabilityFilterOptionsForAccess(RBACListAccess{})
}
func (db *DB) GetVulnerabilityFilterOptionsForAccess(access RBACListAccess) (map[string][]string, error) {
collect := func(query string, args ...interface{}) ([]string, error) {
rows, err := db.Query(query, args...)
if err != nil {
@@ -409,31 +471,35 @@ func (db *DB) GetVulnerabilityFilterOptions() (map[string][]string, error) {
return items, nil
}
vulnIDs, err := collect(`SELECT DISTINCT id FROM vulnerabilities ORDER BY created_at DESC LIMIT 500`)
where := "WHERE 1=1"
accessArgs := []interface{}{}
where, accessArgs = appendVulnerabilityAccessFilter(where, accessArgs, access)
vulnIDs, err := collect(`SELECT DISTINCT id FROM vulnerabilities `+where+` ORDER BY created_at DESC LIMIT 500`, accessArgs...)
if err != nil {
return nil, fmt.Errorf("查询漏洞ID建议失败: %w", err)
}
conversationIDs, err := collect(`SELECT DISTINCT conversation_id FROM vulnerabilities WHERE conversation_id IS NOT NULL AND conversation_id <> '' ORDER BY created_at DESC LIMIT 500`)
conversationIDs, err := collect(`SELECT DISTINCT conversation_id FROM vulnerabilities `+where+` AND conversation_id IS NOT NULL AND conversation_id <> '' ORDER BY created_at DESC LIMIT 500`, accessArgs...)
if err != nil {
return nil, fmt.Errorf("查询会话ID建议失败: %w", err)
}
taskIDs, err := collect(`SELECT DISTINCT id FROM batch_tasks WHERE id <> '' ORDER BY rowid DESC LIMIT 500`)
taskIDs, err := collect(`SELECT DISTINCT bt.id FROM batch_tasks bt JOIN vulnerabilities ON bt.conversation_id = vulnerabilities.conversation_id `+where+` AND bt.id <> '' ORDER BY bt.rowid DESC LIMIT 500`, accessArgs...)
if err != nil {
return nil, fmt.Errorf("查询任务ID建议失败: %w", err)
}
queueIDs, err := collect(`SELECT DISTINCT queue_id FROM batch_tasks WHERE queue_id <> '' ORDER BY rowid DESC LIMIT 500`)
queueIDs, err := collect(`SELECT DISTINCT bt.queue_id FROM batch_tasks bt JOIN vulnerabilities ON bt.conversation_id = vulnerabilities.conversation_id `+where+` AND bt.queue_id <> '' ORDER BY bt.rowid DESC LIMIT 500`, accessArgs...)
if err != nil {
return nil, fmt.Errorf("查询队列ID建议失败: %w", err)
}
conversationTags, err := collect(`SELECT DISTINCT conversation_tag FROM vulnerabilities WHERE conversation_tag IS NOT NULL AND conversation_tag <> '' ORDER BY conversation_tag LIMIT 500`)
conversationTags, err := collect(`SELECT DISTINCT conversation_tag FROM vulnerabilities `+where+` AND conversation_tag IS NOT NULL AND conversation_tag <> '' ORDER BY conversation_tag LIMIT 500`, accessArgs...)
if err != nil {
return nil, fmt.Errorf("查询对话标签建议失败: %w", err)
}
taskTags, err := collect(`SELECT DISTINCT task_tag FROM vulnerabilities WHERE task_tag IS NOT NULL AND task_tag <> '' ORDER BY task_tag LIMIT 500`)
taskTags, err := collect(`SELECT DISTINCT task_tag FROM vulnerabilities `+where+` AND task_tag IS NOT NULL AND task_tag <> '' ORDER BY task_tag LIMIT 500`, accessArgs...)
if err != nil {
return nil, fmt.Errorf("查询任务标签建议失败: %w", err)
}
projectIDs, err := collect(`SELECT DISTINCT project_id FROM vulnerabilities WHERE project_id IS NOT NULL AND project_id <> '' ORDER BY created_at DESC LIMIT 200`)
projectIDs, err := collect(`SELECT DISTINCT project_id FROM vulnerabilities `+where+` AND project_id IS NOT NULL AND project_id <> '' ORDER BY created_at DESC LIMIT 200`, accessArgs...)
if err != nil {
return nil, fmt.Errorf("查询项目ID建议失败: %w", err)
}
+215
View File
@@ -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)
}
}
+20 -2
View File
@@ -2,6 +2,7 @@ package database
import (
"database/sql"
"strings"
"time"
"go.uber.org/zap"
@@ -59,13 +60,30 @@ func (db *DB) UpsertWebshellConnectionState(connectionID, stateJSON string) erro
// ListWebshellConnections 列出所有 WebShell 连接,按创建时间倒序
func (db *DB) ListWebshellConnections() ([]WebShellConnection, error) {
return db.ListWebshellConnectionsForAccess("", "")
}
func (db *DB) ListWebshellConnectionsForAccess(userID, scope string) ([]WebShellConnection, error) {
query := `
SELECT id, url, password, type, method, cmd_param, remark,
COALESCE(encoding, '') AS encoding, COALESCE(os, '') AS os, created_at
FROM webshell_connections
ORDER BY created_at DESC
WHERE 1=1
`
rows, err := db.Query(query)
args := []interface{}{}
userID = strings.TrimSpace(userID)
if userID != "" && scope != RBACScopeAll {
query += ` AND (
owner_user_id = ?
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments ra
WHERE ra.user_id = ? AND ra.resource_type = 'webshell' AND ra.resource_id = webshell_connections.id
)
)`
args = append(args, userID, userID)
}
query += ` ORDER BY created_at DESC`
rows, err := db.Query(query, args...)
if err != nil {
db.logger.Error("查询 WebShell 连接列表失败", zap.Error(err))
return nil, err
+286
View File
@@ -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)
}
}
+143 -33
View File
@@ -18,12 +18,14 @@ import (
"cyberstrike-ai/internal/agent"
"cyberstrike-ai/internal/audit"
"cyberstrike-ai/internal/authctx"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/reasoning"
"cyberstrike-ai/internal/mcp/builtin"
"cyberstrike-ai/internal/multiagent"
"cyberstrike-ai/internal/openai"
"cyberstrike-ai/internal/reasoning"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"github.com/robfig/cron/v3"
@@ -188,8 +190,8 @@ type AgentHandler struct {
hitlWhitelistSaver HitlToolWhitelistSaver
hitlStrategySaver HitlAuditStrategySaver
hitlDefaultReviewerSaver HitlDefaultReviewerSaver
auditLLM *openai.Client
audit *audit.Service
auditLLM *openai.Client
audit *audit.Service
}
// SetAudit wires platform audit logging.
@@ -332,13 +334,13 @@ type ChatReasoningRequest struct {
// ChatRequest 聊天请求
type ChatRequest struct {
Message string `json:"message" binding:"required"`
ConversationID string `json:"conversationId,omitempty"`
ProjectID string `json:"projectId,omitempty"` // 新对话绑定的项目(可选;未指定时可用 config.project.default_project_id
Role string `json:"role,omitempty"` // 角色名称
Attachments []ChatAttachment `json:"attachments,omitempty"`
WebShellConnectionID string `json:"webshellConnectionId,omitempty"` // WebShell 管理 - AI 助手:当前选中的连接 ID,仅使用 webshell_* 工具
Hitl *HITLRequest `json:"hitl,omitempty"`
Message string `json:"message" binding:"required"`
ConversationID string `json:"conversationId,omitempty"`
ProjectID string `json:"projectId,omitempty"` // 新对话绑定的项目(可选;未指定时可用 config.project.default_project_id
Role string `json:"role,omitempty"` // 角色名称
Attachments []ChatAttachment `json:"attachments,omitempty"`
WebShellConnectionID string `json:"webshellConnectionId,omitempty"` // WebShell 管理 - AI 助手:当前选中的连接 ID,仅使用 webshell_* 工具
Hitl *HITLRequest `json:"hitl,omitempty"`
Reasoning *ChatReasoningRequest `json:"reasoning,omitempty"`
// Orchestration 仅对 /api/multi-agent、/api/multi-agent/streamdeep | plan_execute | supervisor;空则等同 deep。机器人/批量等无请求体时由服务端默认 deep。/api/eino-agent* 不使用此字段。
Orchestration string `json:"orchestration,omitempty"`
@@ -727,7 +729,15 @@ func (h *AgentHandler) runRobotMultiAgentWithRetry(
}
// ProcessMessageForRobot 供机器人(企业微信/钉钉/飞书)调用:Eino 单/多代理执行路径(含 progressCallback、过程详情),仅不发送 SSE,最后返回完整回复
func (h *AgentHandler) ProcessMessageForRobot(ctx context.Context, platform, 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)
if ownerUserID == "" {
return "", "", fmt.Errorf("authenticated robot principal is required")
}
if !principal.HasPermission("agent:execute") || !principal.HasPermission("chat:read") || !principal.HasPermission("chat:write") {
return "", "", fmt.Errorf("机器人账号缺少 agent:execute、chat:read 或 chat:write 权限")
}
ctx = authctx.WithPrincipal(ctx, principal)
if conversationID == "" {
title := safeTruncateString(message, 50)
src := "robot"
@@ -736,13 +746,17 @@ func (h *AgentHandler) ProcessMessageForRobot(ctx context.Context, platform, con
}
meta := audit.ConversationCreateMeta(src)
meta.ProjectID = effectiveProjectID(h.config, "")
if meta.ProjectID != "" && (!principal.HasPermission("project:read") || !h.db.UserCanAccessResource(ownerUserID, principal.ScopeFor("project:read"), "project", meta.ProjectID)) {
meta.ProjectID = ""
}
conv, createErr := h.db.CreateConversation(title, meta)
if createErr != nil {
return "", "", fmt.Errorf("创建对话失败: %w", createErr)
}
conversationID = conv.ID
_ = h.db.SetResourceOwner("conversation", conversationID, ownerUserID)
} else {
if _, getErr := h.db.GetConversation(conversationID); getErr != nil {
if _, getErr := h.db.GetConversation(conversationID); getErr != nil || !h.db.UserCanAccessResource(ownerUserID, principal.ScopeFor("chat:write"), "conversation", conversationID) {
return "", "", fmt.Errorf("对话不存在")
}
}
@@ -800,18 +814,14 @@ func (h *AgentHandler) ProcessMessageForRobot(ctx context.Context, platform, con
}
progressCallback := h.createProgressCallback(taskCtx, cancelWithCause, conversationID, assistantMessageID, nil)
robotMode := "eino_single"
if h.config != nil {
robotMode = config.NormalizeRobotAgentMode(h.config.MultiAgent)
}
robotMode := config.NormalizeAgentMode(agentMode)
switch robotMode {
case "eino_single":
return h.runRobotEinoSingleWithRetry(taskCtx, conversationID, finalMessage, agentHistoryMessages, roleTools, progressCallback, assistantMessageID, &taskStatus)
case "deep", "plan_execute", "supervisor":
if h.config == nil || !h.config.MultiAgent.Enabled {
h.logger.Warn("机器人配置为多代理模式但未启用 multi_agent,回退 Eino 单代理",
zap.String("robot_mode", robotMode))
return h.runRobotEinoSingleWithRetry(taskCtx, conversationID, finalMessage, agentHistoryMessages, roleTools, progressCallback, assistantMessageID, &taskStatus)
taskStatus = "failed"
return "", conversationID, fmt.Errorf("机器人对话模式 %s 需要启用 Eino 多代理", robotMode)
}
return h.runRobotMultiAgentWithRetry(taskCtx, conversationID, finalMessage, robotMode, agentHistoryMessages, roleTools, progressCallback, assistantMessageID, &taskStatus)
}
@@ -844,6 +854,36 @@ func (h *AgentHandler) publishProgressToTaskEventBus(conversationID, eventType,
h.taskEventBus.Publish(conversationID, sseLine)
}
// enrichProgressEventData 为 SSE / taskEventBus 事件补齐 conversationId、messageId,便于前端懒加载过程详情。
func enrichProgressEventData(data interface{}, conversationID, assistantMessageID string) interface{} {
if strings.TrimSpace(conversationID) == "" && strings.TrimSpace(assistantMessageID) == "" {
return data
}
var m map[string]interface{}
switch v := data.(type) {
case map[string]interface{}:
m = make(map[string]interface{}, len(v)+2)
for k, val := range v {
m[k] = val
}
case nil:
m = make(map[string]interface{}, 2)
default:
m = map[string]interface{}{"payload": data}
}
if id := strings.TrimSpace(assistantMessageID); id != "" {
if existing, ok := m["messageId"]; !ok || strings.TrimSpace(fmt.Sprint(existing)) == "" {
m["messageId"] = id
}
}
if id := strings.TrimSpace(conversationID); id != "" {
if existing, ok := m["conversationId"]; !ok || strings.TrimSpace(fmt.Sprint(existing)) == "" {
m["conversationId"] = id
}
}
return m
}
// createProgressCallback 创建进度回调函数,用于保存processDetails
// sendEventFunc: 可选的流式事件发送函数,如果为nil则不发送流式事件
func (h *AgentHandler) createProgressCallback(runCtx context.Context, cancelRun context.CancelCauseFunc, conversationID, assistantMessageID string, sendEventFunc func(eventType, message string, data interface{})) agent.ProgressCallback {
@@ -988,11 +1028,21 @@ func (h *AgentHandler) createProgressCallback(runCtx context.Context, cancelRun
}
}
// 流式:写 HTTP SSE;非流式(机器人等):镜像到 taskEventBus 供 Web 订阅
if sendEventFunc != nil {
sendEventFunc(eventType, message, data)
} else {
h.publishProgressToTaskEventBus(conversationID, eventType, message, data)
// 工具输出片段不在详情区实时展示;完整结果由 tool_result 落库后按需拉取。
if eventType == "tool_result_delta" {
return
}
deferToolProgressSend := eventType == "tool_call" || eventType == "tool_result"
// 流式:写 HTTP SSE;非流式(机器人等):镜像到 taskEventBus 供 Web 订阅。
// 工具事件需先落库拿 processDetailId,再向前端发送摘要,避免大 payload 默认进入浏览器。
if !deferToolProgressSend {
clientData := enrichProgressEventData(data, conversationID, assistantMessageID)
if sendEventFunc != nil {
sendEventFunc(eventType, message, clientData)
} else {
h.publishProgressToTaskEventBus(conversationID, eventType, message, clientData)
}
}
// 保存tool_call事件中的参数
@@ -1360,9 +1410,28 @@ func (h *AgentHandler) createProgressCallback(runCtx context.Context, cancelRun
// 在关键过程事件落库前,先把「规划中」与聚合中的 thinking / reasoning_chain 流落库
flushResponsePlan()
flushThinkingStreams()
if err := h.db.AddProcessDetail(assistantMessageID, conversationID, eventType, message, data); err != nil {
processDetailID, err := h.db.AddProcessDetailWithID(assistantMessageID, conversationID, eventType, message, data)
if err != nil {
h.logger.Warn("保存过程详情失败", zap.Error(err), zap.String("eventType", eventType))
}
if deferToolProgressSend {
clientData := enrichProgressEventData(summarizeProcessDetailData(eventType, data), conversationID, assistantMessageID)
if m, ok := clientData.(map[string]interface{}); ok {
m["processDetailId"] = processDetailID
}
if sendEventFunc != nil {
sendEventFunc(eventType, message, clientData)
} else {
h.publishProgressToTaskEventBus(conversationID, eventType, message, clientData)
}
}
} else if deferToolProgressSend {
clientData := enrichProgressEventData(summarizeProcessDetailData(eventType, data), conversationID, assistantMessageID)
if sendEventFunc != nil {
sendEventFunc(eventType, message, clientData)
} else {
h.publishProgressToTaskEventBus(conversationID, eventType, message, clientData)
}
}
}
}
@@ -1429,6 +1498,10 @@ func (h *AgentHandler) CancelAgentLoop(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if !h.agentConversationAllowed(c, req.ConversationID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
if req.ContinueAfter {
if h.tasks.GetTask(req.ConversationID) == nil {
@@ -1495,10 +1568,10 @@ func (h *AgentHandler) CancelAgentLoop(c *gin.Context) {
}
c.JSON(http.StatusOK, gin.H{
"status": "cancelling",
"conversationId": req.ConversationID,
"message": msg,
"continueAfter": false,
"status": "cancelling",
"conversationId": req.ConversationID,
"message": msg,
"continueAfter": false,
"interruptWithNote": false,
})
}
@@ -1510,6 +1583,10 @@ func (h *AgentHandler) SubscribeAgentTaskEvents(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "conversationId is required"})
return
}
if !h.agentConversationAllowed(c, conversationID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
if h.tasks.GetTask(conversationID) == nil {
c.JSON(http.StatusNotFound, gin.H{"error": "no active task for this conversation"})
return
@@ -1530,9 +1607,8 @@ func (h *AgentHandler) SubscribeAgentTaskEvents(c *gin.Context) {
flusher, _ := c.Writer.(http.Flusher)
ctx := c.Request.Context()
var writeMu sync.Mutex
stopKeepalive := make(chan struct{})
go sseKeepalive(c, stopKeepalive, &writeMu)
defer close(stopKeepalive)
stopKeepalive := runSSEKeepalive(c, &writeMu)
defer stopKeepalive()
for {
select {
@@ -1588,6 +1664,9 @@ func (h *AgentHandler) enrichCompletedTasksWithConversationTitles(tasks []*Compl
// ListAgentTasks 列出所有运行中的任务
func (h *AgentHandler) ListAgentTasks(c *gin.Context) {
tasks := h.tasks.GetActiveTasks()
tasks = filterSlice(tasks, func(task *AgentTask) bool {
return task != nil && h.agentConversationAllowed(c, task.ConversationID)
})
h.enrichAgentTasksWithConversationTitles(tasks)
c.JSON(http.StatusOK, gin.H{
"tasks": tasks,
@@ -1597,12 +1676,30 @@ func (h *AgentHandler) ListAgentTasks(c *gin.Context) {
// ListCompletedTasks 列出最近完成的任务历史
func (h *AgentHandler) ListCompletedTasks(c *gin.Context) {
tasks := h.tasks.GetCompletedTasks()
tasks = filterSlice(tasks, func(task *CompletedTask) bool {
return task != nil && h.agentConversationAllowed(c, task.ConversationID)
})
h.enrichCompletedTasksWithConversationTitles(tasks)
c.JSON(http.StatusOK, gin.H{
"tasks": tasks,
})
}
func (h *AgentHandler) agentConversationAllowed(c *gin.Context, conversationID string) bool {
session, ok := security.CurrentSession(c)
return ok && h.db != nil && h.db.UserCanAccessResource(session.UserID, session.Scope, "conversation", strings.TrimSpace(conversationID))
}
func filterSlice[T any](items []T, keep func(T) bool) []T {
out := make([]T, 0, len(items))
for _, item := range items {
if keep(item) {
out = append(out, item)
}
}
return out
}
// BatchTaskRequest 批量任务请求
type BatchTaskRequest struct {
Title string `json:"title"` // 任务标题(可选)
@@ -1654,6 +1751,12 @@ func (h *AgentHandler) CreateBatchQueue(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "没有有效的任务"})
return
}
if session, ok := security.CurrentSession(c); ok && h.db != nil && session.Scope != database.RBACScopeAll && strings.TrimSpace(req.ProjectID) != "" {
if !h.db.UserCanAccessResource(session.UserID, session.Scope, "project", strings.TrimSpace(req.ProjectID)) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权在该项目下创建批量任务"})
return
}
}
agentMode := config.NormalizeAgentMode(req.AgentMode)
scheduleMode := normalizeBatchQueueScheduleMode(req.ScheduleMode)
@@ -1678,6 +1781,10 @@ func (h *AgentHandler) CreateBatchQueue(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": createErr.Error()})
return
}
if session, ok := security.CurrentSession(c); ok && h.db != nil {
_ = h.db.SetResourceOwner("batch_task", queue.ID, session.UserID)
_ = h.db.AssignResourceToUser(session.UserID, "batch_task", queue.ID)
}
started := false
if req.ExecuteNow {
ok, err := h.startBatchQueueExecution(queue.ID, false)
@@ -1765,7 +1872,8 @@ func (h *AgentHandler) ListBatchQueues(c *gin.Context) {
}
// 获取队列列表和总数
queues, total, err := h.batchTaskManager.ListQueues(limit, offset, status, keyword)
session, _ := security.CurrentSession(c)
queues, total, err := h.batchTaskManager.ListQueuesForAccess(limit, offset, status, keyword, session.UserID, session.Scope)
if err != nil {
h.logger.Error("获取批量任务队列列表失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -2233,6 +2341,7 @@ func (h *AgentHandler) loadHistoryFromAgentTrace(conversationID string) ([]agent
}
messageCount := len(messagesArray)
modelFacingTrace := agent.IsModelFacingTraceJSON(traceInputJSON)
h.logger.Info("使用保存的代理轨迹恢复历史上下文",
zap.String("conversationId", conversationID),
@@ -2247,6 +2356,7 @@ func (h *AgentHandler) loadHistoryFromAgentTrace(conversationID string) ([]agent
agentMessages := make([]agent.ChatMessage, 0, len(messagesArray))
for _, msgMap := range messagesArray {
msg := agent.ChatMessage{}
msg.ModelFacingTrace = modelFacingTrace
// 解析role
if role, ok := msgMap["role"].(string); ok {
@@ -97,3 +97,26 @@ func TestCreateProgressCallback_FlushesReasoningOnDone(t *testing.T) {
t.Fatalf("expected reasoning_chain persisted on done, got %+v", details)
}
}
func TestEnrichProgressEventData(t *testing.T) {
t.Run("fills ids", func(t *testing.T) {
out := enrichProgressEventData(map[string]interface{}{"source": "eino"}, "conv-1", "msg-1")
m, ok := out.(map[string]interface{})
if !ok {
t.Fatalf("expected map, got %T", out)
}
if m["conversationId"] != "conv-1" || m["messageId"] != "msg-1" {
t.Fatalf("unexpected enrichment: %+v", m)
}
})
t.Run("preserves existing ids", func(t *testing.T) {
out := enrichProgressEventData(map[string]interface{}{
"conversationId": "keep-conv",
"messageId": "keep-msg",
}, "conv-1", "msg-1")
m := out.(map[string]interface{})
if m["conversationId"] != "keep-conv" || m["messageId"] != "keep-msg" {
t.Fatalf("should not overwrite existing ids: %+v", m)
}
})
}
+20 -8
View File
@@ -6,6 +6,7 @@ import (
"cyberstrike-ai/internal/audit"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
@@ -32,8 +33,8 @@ func (h *AuditHandler) Meta(c *gin.Context) {
retentionDays = h.audit.RetentionDays()
}
c.JSON(http.StatusOK, gin.H{
"enabled": enabled,
"retention_days": retentionDays,
"enabled": enabled,
"retention_days": retentionDays,
"default_page_size": 20,
"max_page_size": 100,
"max_export": 5000,
@@ -46,7 +47,7 @@ func (h *AuditHandler) Summary(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": "database unavailable"})
return
}
base := auditFilterFromQuery(c)
base := auditFilterForAccess(c, auditFilterFromQuery(c))
total, err := h.db.CountAuditLogs(base)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -68,9 +69,9 @@ func (h *AuditHandler) Summary(c *gin.Context) {
return
}
c.JSON(http.StatusOK, gin.H{
"total": total,
"failures": failures,
"recent_7d": recent7d,
"total": total,
"failures": failures,
"recent_7d": recent7d,
"has_filters": c.Query("category") != "" || c.Query("action") != "" || c.Query("result") != "" ||
c.Query("q") != "" || c.Query("since") != "" || c.Query("until") != "",
})
@@ -82,7 +83,7 @@ func (h *AuditHandler) ListLogs(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": "database unavailable"})
return
}
filter := auditFilterFromQuery(c)
filter := auditFilterForAccess(c, auditFilterFromQuery(c))
page, pageSize := auditPaginationFromQuery(c)
filter.Limit = pageSize
filter.Offset = (page - 1) * pageSize
@@ -116,6 +117,10 @@ func (h *AuditHandler) GetLog(c *gin.Context) {
c.JSON(http.StatusNotFound, gin.H{"error": "审计记录不存在"})
return
}
if session, ok := security.CurrentSession(c); !ok || (session.Scope != database.RBACScopeAll && row.Actor != session.Username) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
audit.ApplyResourceAvailability(h.db, row)
c.JSON(http.StatusOK, gin.H{"log": row})
}
@@ -126,7 +131,7 @@ func (h *AuditHandler) ExportLogs(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": "database unavailable"})
return
}
filter := auditFilterFromQuery(c)
filter := auditFilterForAccess(c, auditFilterFromQuery(c))
filter.Limit = 5000
filter.Offset = 0
@@ -145,3 +150,10 @@ func (h *AuditHandler) ExportLogs(c *gin.Context) {
"logs": logs,
})
}
func auditFilterForAccess(c *gin.Context, filter database.ListAuditLogsFilter) database.ListAuditLogsFilter {
if session, ok := security.CurrentSession(c); ok && session.Scope != database.RBACScopeAll {
filter.Actor = session.Username
}
return filter
}
+9 -7
View File
@@ -10,13 +10,15 @@ import (
func auditFilterFromQuery(c *gin.Context) database.ListAuditLogsFilter {
filter := database.ListAuditLogsFilter{
Level: c.Query("level"),
Category: c.Query("category"),
Action: c.Query("action"),
Result: c.Query("result"),
Query: c.Query("q"),
ResourceType: c.Query("resource_type"),
ResourceID: c.Query("resource_id"),
Actor: c.Query("actor"),
Level: c.Query("level"),
Category: c.Query("category"),
Action: c.Query("action"),
Result: c.Query("result"),
Query: c.Query("q"),
ResourceType: c.Query("resource_type"),
ResourceID: c.Query("resource_id"),
RelatedUserID: c.Query("related_user_id"),
}
if since := c.Query("since"); since != "" {
if t, err := database.ParseRFC3339Time(since); err == nil {
+19
View File
@@ -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)
}
}
+44 -18
View File
@@ -38,6 +38,7 @@ func NewAuthHandler(manager *security.AuthManager, cfg *config.Config, configPat
}
type loginRequest struct {
Username string `json:"username"`
Password string `json:"password" binding:"required"`
}
@@ -54,7 +55,7 @@ func (h *AuthHandler) Login(c *gin.Context) {
return
}
token, expiresAt, err := h.manager.Authenticate(req.Password)
token, expiresAt, err := h.manager.Authenticate(req.Username, req.Password)
if err != nil {
if h.audit != nil {
h.audit.Record(c, audit.Entry{
@@ -63,11 +64,13 @@ func (h *AuthHandler) Login(c *gin.Context) {
Action: "login",
Result: "failure",
Message: "登录失败:密码错误",
Actor: strings.TrimSpace(req.Username),
})
}
c.JSON(http.StatusUnauthorized, gin.H{"error": "密码错误"})
return
}
session, _ := h.manager.ValidateToken(token)
if h.audit != nil {
h.audit.Record(c, audit.Entry{
@@ -76,6 +79,7 @@ func (h *AuthHandler) Login(c *gin.Context) {
Result: "success",
SessionHint: audit.HintFromToken(token),
Message: "登录成功",
Actor: session.Username,
Detail: map[string]interface{}{
"expires_at": expiresAt.UTC().Format(time.RFC3339),
},
@@ -86,6 +90,15 @@ func (h *AuthHandler) Login(c *gin.Context) {
"token": token,
"expires_at": expiresAt.UTC().Format(time.RFC3339),
"session_duration_hr": h.manager.SessionDurationHours(),
"user": gin.H{
"id": session.UserID,
"username": session.Username,
"display_name": session.DisplayName,
},
"roles": session.Roles,
"permissions": permissionKeys(session.Permissions),
"permission_scopes": session.PermissionScopes,
"scope": session.Scope,
})
}
@@ -139,7 +152,11 @@ func (h *AuthHandler) ChangePassword(c *gin.Context) {
return
}
if !h.manager.CheckPassword(oldPassword) {
session, _ := security.CurrentSession(c)
if session.Username == "" {
session.Username = "admin"
}
if !h.manager.CheckUserPassword(session.Username, oldPassword) {
if h.audit != nil {
h.audit.Record(c, audit.Entry{
Level: "warn",
@@ -153,27 +170,17 @@ func (h *AuthHandler) ChangePassword(c *gin.Context) {
return
}
if err := config.PersistAuthPassword(h.configPath, newPassword); err != nil {
if session.UserID == "" {
session.UserID = "admin"
}
if err := h.manager.UpdateUserPassword(session.UserID, newPassword); err != nil {
if h.logger != nil {
h.logger.Error("保存新密码失败", zap.Error(err))
h.logger.Error("更新用户密码失败", zap.Error(err))
}
c.JSON(http.StatusInternalServerError, gin.H{"error": "保存新密码失败,请重试"})
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 = ""
if h.logger != nil {
h.logger.Info("登录密码已更新,所有会话已失效")
}
@@ -207,5 +214,24 @@ func (h *AuthHandler) Validate(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{
"token": session.Token,
"expires_at": session.ExpiresAt.UTC().Format(time.RFC3339),
"user": gin.H{
"id": session.UserID,
"username": session.Username,
"display_name": session.DisplayName,
},
"roles": session.Roles,
"permissions": permissionKeys(session.Permissions),
"permission_scopes": session.PermissionScopes,
"scope": session.Scope,
})
}
func permissionKeys(perms map[string]bool) []string {
keys := make([]string, 0, len(perms))
for key, ok := range perms {
if ok {
keys = append(keys, key)
}
}
return keys
}
+12 -1
View File
@@ -11,6 +11,7 @@ import (
"cyberstrike-ai/internal/agent"
"cyberstrike-ai/internal/audit"
"cyberstrike-ai/internal/authctx"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/mcp"
"cyberstrike-ai/internal/multiagent"
@@ -109,6 +110,13 @@ func (h *AgentHandler) tryFinalizeBatchQueue(queueID string) {
// executeOneBatchSubTask 执行单条批量子任务(各自独立会话)。
func (h *AgentHandler) executeOneBatchSubTask(queueID string, queue *BatchTaskQueue, task *BatchTask) {
ownerUserID := h.db.GetResourceOwner("batch_task", queueID)
access, accessErr := h.db.ResolveRBACAccess(ownerUserID)
if accessErr != nil || access == nil || !access.User.Enabled {
h.batchTaskManager.UpdateTaskStatus(queueID, task.ID, BatchTaskStatusFailed, "", "队列所有者不存在或已禁用")
return
}
principal := authctx.NewPrincipalWithScopes(access.User.ID, access.User.Username, access.Scope, access.Permissions, access.PermissionScopes)
title := safeTruncateString(task.Message, 50)
batchMeta := audit.ConversationCreateMeta("batch_task")
batchMeta.ProjectID = effectiveProjectID(h.config, queue.ProjectID)
@@ -119,6 +127,8 @@ func (h *AgentHandler) executeOneBatchSubTask(queueID string, queue *BatchTaskQu
return
}
conversationID := conv.ID
_ = h.db.SetResourceOwner("conversation", conversationID, access.User.ID)
_ = h.db.AssignResourceToUser(access.User.ID, "conversation", conversationID)
h.batchTaskManager.UpdateTaskStatusWithConversationID(queueID, task.ID, BatchTaskStatusRunning, "", "", conversationID)
@@ -156,7 +166,8 @@ func (h *AgentHandler) executeOneBatchSubTask(queueID string, queue *BatchTaskQu
h.logger.Info("执行批量任务", zap.String("queueId", queueID), zap.String("taskId", task.ID), zap.String("message", task.Message), zap.String("role", queue.Role), zap.String("conversationId", conversationID))
baseCtx, cancelWithCause := context.WithCancelCause(context.Background())
principalCtx := authctx.WithPrincipal(context.Background(), principal)
baseCtx, cancelWithCause := context.WithCancelCause(principalCtx)
taskCtx, timeoutCancel := context.WithTimeout(baseCtx, 6*time.Hour)
registered := false
+8 -4
View File
@@ -98,8 +98,8 @@ type BatchTaskManager struct {
logger *zap.Logger
queues map[string]*BatchTaskQueue
taskCancels map[string]map[string]context.CancelFunc // queueID -> taskID -> 取消函数
singleRunTasks map[string]string // queueID -> taskID,单条执行完成后暂停队列
queueExecutors map[string]struct{} // executeBatchQueue 协程活跃标记(与队列 status 解耦)
singleRunTasks map[string]string // queueID -> taskID,单条执行完成后暂停队列
queueExecutors map[string]struct{} // executeBatchQueue 协程活跃标记(与队列 status 解耦)
mu sync.RWMutex
}
@@ -426,20 +426,24 @@ func (m *BatchTaskManager) GetAllQueues() []*BatchTaskQueue {
// ListQueues 列出队列(支持筛选和分页)
func (m *BatchTaskManager) ListQueues(limit, offset int, status, keyword string) ([]*BatchTaskQueue, int, error) {
return m.ListQueuesForAccess(limit, offset, status, keyword, "", "")
}
func (m *BatchTaskManager) ListQueuesForAccess(limit, offset int, status, keyword, userID, scope string) ([]*BatchTaskQueue, int, error) {
var queues []*BatchTaskQueue
var total int
// 如果数据库可用,从数据库查询
if m.db != nil {
// 获取总数
count, err := m.db.CountBatchQueues(status, keyword)
count, err := m.db.CountBatchQueuesForAccess(status, keyword, userID, scope)
if err != nil {
return nil, 0, fmt.Errorf("统计队列总数失败: %w", err)
}
total = count
// 获取队列列表(只获取ID
queueRows, err := m.db.ListBatchQueues(limit, offset, status, keyword)
queueRows, err := m.db.ListBatchQueuesForAccess(limit, offset, status, keyword, userID, scope)
if err != nil {
return nil, 0, fmt.Errorf("查询队列列表失败: %w", err)
}
+20 -2
View File
@@ -9,7 +9,9 @@ import (
"strings"
"time"
"cyberstrike-ai/internal/authctx"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/mcp"
"cyberstrike-ai/internal/mcp/builtin"
@@ -74,7 +76,14 @@ func RegisterBatchTaskMCPTools(mcpServer *mcp.Server, h *AgentHandler, logger *z
if offset > 100000 {
offset = 100000
}
queues, total, err := h.batchTaskManager.ListQueues(pageSize, offset, status, keyword)
queues := []*BatchTaskQueue{}
total := 0
var err error
if principal, ok := authctx.PrincipalFromContext(ctx); ok {
queues, total, err = h.batchTaskManager.ListQueuesForAccess(pageSize, offset, status, keyword, principal.UserID, principal.ScopeFor("tasks:read"))
} else {
return batchMCPTextResult("缺少认证身份", true), nil
}
if err != nil {
return batchMCPTextResult(fmt.Sprintf("列出队列失败: %v", err), true), nil
}
@@ -215,11 +224,20 @@ func RegisterBatchTaskMCPTools(mcpServer *mcp.Server, h *AgentHandler, logger *z
executeNow = false
}
projectID := strings.TrimSpace(mcpArgString(args, "project_id"))
if principal, ok := authctx.PrincipalFromContext(ctx); ok && projectID != "" && principal.ScopeFor("tasks:write") != database.RBACScopeAll {
if h.db == nil || !h.db.UserCanAccessResource(principal.UserID, principal.ScopeFor("tasks:write"), "project", projectID) {
return batchMCPTextResult("无权访问目标项目", true), nil
}
}
concurrency := int(mcpArgFloat(args, "concurrency"))
queue, createErr := h.batchTaskManager.CreateBatchQueue(title, role, agentMode, scheduleMode, cronExpr, projectID, nextRunAt, concurrency, tasks)
if createErr != nil {
return batchMCPTextResult("创建队列失败: "+createErr.Error(), true), nil
}
if principal, ok := authctx.PrincipalFromContext(ctx); ok && h.db != nil {
_ = h.db.SetResourceOwner("batch_task", queue.ID, principal.UserID)
_ = h.db.AssignResourceToUser(principal.UserID, "batch_task", queue.ID)
}
started := false
if executeNow {
ok, err := h.startBatchQueueExecution(queue.ID, false)
@@ -646,7 +664,7 @@ schedule_mode 为 cron 时必须提供有效 cron_expr;为 manual 时会清除
return batchMCPJSONResult(queue)
})
logger.Info("批量任务 MCP 工具已注册", zap.Int("count", 12))
logger.Debug("批量任务 MCP 工具已注册", zap.Int("count", 12))
}
// --- batch_task_list 精简结构(避免把每条子任务的 result 等大段文本塞进列表上下文) ---
+104 -14
View File
@@ -17,6 +17,7 @@ import (
"cyberstrike-ai/internal/audit"
"cyberstrike-ai/internal/c2"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"github.com/google/uuid"
@@ -59,7 +60,7 @@ func (h *C2Handler) SetManager(m *c2.Manager) {
// ListListeners 获取监听器列表
func (h *C2Handler) ListListeners(c *gin.Context) {
listeners, err := h.mgr().DB().ListC2Listeners()
listeners, err := h.mgr().DB().ListC2ListenersForAccess(c2AccessFromContext(c))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -109,6 +110,11 @@ func (h *C2Handler) CreateListener(c *gin.Context) {
c.JSON(code, gin.H{"error": err.Error()})
return
}
if session, ok := security.CurrentSession(c); ok {
listener.OwnerUserID = session.UserID
_ = h.mgr().DB().SetResourceOwner("c2_listener", listener.ID, session.UserID)
_ = h.mgr().DB().AssignResourceToUser(session.UserID, "c2_listener", listener.ID)
}
implantToken := listener.ImplantToken
listener.EncryptionKey = ""
listener.ImplantToken = ""
@@ -282,7 +288,7 @@ func (h *C2Handler) ListSessions(c *gin.Context) {
filter.Suspicious = true
}
sessions, err := h.mgr().DB().ListC2Sessions(filter)
sessions, err := h.mgr().DB().ListC2SessionsForAccess(filter, c2AccessFromContext(c))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -304,10 +310,10 @@ func (h *C2Handler) GetSession(c *gin.Context) {
}
// 获取最近任务
tasks, _ := h.mgr().DB().ListC2Tasks(database.ListC2TasksFilter{
tasks, _ := h.mgr().DB().ListC2TasksForAccess(database.ListC2TasksFilter{
SessionID: id,
Limit: 20,
})
}, c2AccessFromContext(c))
c.JSON(http.StatusOK, gin.H{
"session": session,
@@ -341,7 +347,7 @@ func (h *C2Handler) DeleteSessions(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "ids is required"})
return
}
n, err := h.mgr().DB().DeleteC2SessionsByIDs(req.IDs)
n, err := h.mgr().DB().DeleteC2SessionsByIDsForAccess(req.IDs, c2AccessFromContext(c))
if err != nil {
if errors.Is(err, database.ErrNoValidC2SessionIDs) {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -433,24 +439,25 @@ func (h *C2Handler) ListTasks(c *gin.Context) {
}
}
tasks, err := h.mgr().DB().ListC2Tasks(filter)
access := c2AccessFromContext(c)
tasks, err := h.mgr().DB().ListC2TasksForAccess(filter, access)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
// 仪表盘「待审任务」为全局 queued/pending 数量,与列表 session 过滤无关
pendingN, _ := h.mgr().DB().CountC2TasksQueuedOrPending("")
pendingN, _ := h.mgr().DB().CountC2TasksQueuedOrPendingForAccess("", access)
if !paginated {
c.JSON(http.StatusOK, gin.H{
"tasks": tasks,
"pending_queued_count": pendingN,
"tasks": tasks,
"pending_queued_count": pendingN,
})
return
}
total, err := h.mgr().DB().CountC2Tasks(filter)
total, err := h.mgr().DB().CountC2TasksForAccess(filter, access)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -477,7 +484,7 @@ func (h *C2Handler) DeleteTasks(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "ids is required"})
return
}
n, err := h.mgr().DB().DeleteC2TasksByIDs(req.IDs)
n, err := h.mgr().DB().DeleteC2TasksByIDsForAccess(req.IDs, c2AccessFromContext(c))
if err != nil {
if errors.Is(err, database.ErrNoValidC2TaskIDs) {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -522,6 +529,21 @@ func (h *C2Handler) CreateTask(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if strings.TrimSpace(req.SessionID) == "" {
req.SessionID = strings.TrimSpace(c.Param("id"))
}
if !h.c2ResourceAllowed(c, "c2_session", req.SessionID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
if conversationID := strings.TrimSpace(req.ConversationID); conversationID != "" {
session, ok := security.CurrentSession(c)
if !ok || !h.mgr().DB().UserCanAccessResource(session.UserID, session.Scope, "conversation", conversationID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权关联目标对话"})
return
}
req.ConversationID = conversationID
}
input := c2.EnqueueTaskInput{
SessionID: req.SessionID,
@@ -621,6 +643,10 @@ func (h *C2Handler) PayloadOneliner(c *gin.Context) {
c.JSON(http.StatusNotFound, gin.H{"error": "listener not found"})
return
}
if !h.c2ResourceAllowed(c, "c2_listener", req.ListenerID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
host := c2.ResolveBeaconDialHost(listener, strings.TrimSpace(req.Host), h.logger, listener.ID)
@@ -684,6 +710,10 @@ func (h *C2Handler) PayloadBuild(c *gin.Context) {
c.JSON(http.StatusNotFound, gin.H{"error": "listener not found"})
return
}
if !h.c2ResourceAllowed(c, "c2_listener", req.ListenerID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
builder := c2.NewPayloadBuilder(h.mgr(), h.logger, "", "")
input := c2.PayloadBuilderInput{
@@ -700,6 +730,9 @@ func (h *C2Handler) PayloadBuild(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if session, ok := security.CurrentSession(c); ok {
_ = h.mgr().DB().RecordC2PayloadArtifact(filepath.Base(result.OutputPath), result.PayloadID, result.ListenerID, session.UserID)
}
c.JSON(http.StatusOK, gin.H{
"payload": result,
@@ -718,6 +751,11 @@ func (h *C2Handler) PayloadDownload(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid payload id"})
return
}
session, ok := security.CurrentSession(c)
if !ok || !h.mgr().DB().UserCanAccessC2Payload(session.UserID, session.Scope, filename) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
builder := c2.NewPayloadBuilder(h.mgr(), h.logger, "", "")
storageDir := builder.GetPayloadStoragePath()
@@ -779,7 +817,8 @@ func (h *C2Handler) ListEvents(c *gin.Context) {
}
}
events, err := h.mgr().DB().ListC2Events(filter)
access := c2AccessFromContext(c)
events, err := h.mgr().DB().ListC2EventsForAccess(filter, access)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -788,7 +827,7 @@ func (h *C2Handler) ListEvents(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"events": events})
return
}
total, err := h.mgr().DB().CountC2Events(filter)
total, err := h.mgr().DB().CountC2EventsForAccess(filter, access)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -814,7 +853,7 @@ func (h *C2Handler) DeleteEvents(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "ids is required"})
return
}
n, err := h.mgr().DB().DeleteC2EventsByIDs(req.IDs)
n, err := h.mgr().DB().DeleteC2EventsByIDsForAccess(req.IDs, c2AccessFromContext(c))
if err != nil {
if errors.Is(err, database.ErrNoValidC2EventIDs) {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -851,6 +890,9 @@ func (h *C2Handler) EventStream(c *gin.Context) {
if !ok {
return false
}
if !h.c2EventAllowed(c, e) {
return true
}
data, _ := json.Marshal(e)
fmt.Fprintf(w, "data: %s\n\n", data)
return true
@@ -964,6 +1006,10 @@ func (h *C2Handler) UploadFileForImplant(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "session_id and remote_path required"})
return
}
if !h.c2ResourceAllowed(c, "c2_session", sessionID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
file, header, err := c.Request.FormFile("file")
if err != nil {
@@ -1018,6 +1064,10 @@ func (h *C2Handler) ListFiles(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "session_id required"})
return
}
if !h.c2ResourceAllowed(c, "c2_session", sessionID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
files, err := h.mgr().DB().ListC2FilesBySession(sessionID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -1038,6 +1088,10 @@ func (h *C2Handler) DownloadResultFile(c *gin.Context) {
c.JSON(http.StatusNotFound, gin.H{"error": "task not found"})
return
}
if !h.c2ResourceAllowed(c, "c2_task", taskID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
if task.ResultBlobPath == "" {
c.JSON(http.StatusNotFound, gin.H{"error": "no result file for this task"})
return
@@ -1053,6 +1107,42 @@ func osCreate(path string) (*os.File, error) {
return os.Create(path)
}
func c2AccessFromContext(c *gin.Context) database.RBACListAccess {
session, ok := security.CurrentSession(c)
if !ok {
return database.RBACListAccess{}
}
return database.RBACListAccess{UserID: session.UserID, Scope: session.Scope}
}
func (h *C2Handler) c2ResourceAllowed(c *gin.Context, resourceType, resourceID string) bool {
session, ok := security.CurrentSession(c)
if !ok {
return false
}
return h.mgr().DB().UserCanAccessResource(session.UserID, session.Scope, resourceType, resourceID)
}
func (h *C2Handler) c2EventAllowed(c *gin.Context, e *c2.Event) bool {
if e == nil {
return false
}
session, ok := security.CurrentSession(c)
if !ok {
return false
}
if session.Scope == database.RBACScopeAll {
return true
}
if strings.TrimSpace(e.SessionID) != "" {
return h.mgr().DB().UserCanAccessResource(session.UserID, session.Scope, "c2_session", e.SessionID)
}
if strings.TrimSpace(e.TaskID) != "" {
return h.mgr().DB().UserCanAccessResource(session.UserID, session.Scope, "c2_task", e.TaskID)
}
return false
}
// ============================================================================
// 辅助函数(firstNonEmpty 已在 vulnerability.go 中定义)
// ============================================================================
+90 -3
View File
@@ -13,6 +13,8 @@ import (
"unicode/utf8"
"cyberstrike-ai/internal/audit"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
@@ -27,6 +29,7 @@ const (
type ChatUploadsHandler struct {
logger *zap.Logger
audit *audit.Service
db *database.DB
}
// SetAudit wires platform audit logging.
@@ -35,8 +38,32 @@ func (h *ChatUploadsHandler) SetAudit(s *audit.Service) {
}
// NewChatUploadsHandler 创建处理器
func NewChatUploadsHandler(logger *zap.Logger) *ChatUploadsHandler {
return &ChatUploadsHandler{logger: logger}
func NewChatUploadsHandler(logger *zap.Logger, databases ...*database.DB) *ChatUploadsHandler {
h := &ChatUploadsHandler{logger: logger}
if len(databases) > 0 {
h.db = databases[0]
}
return h
}
func (h *ChatUploadsHandler) pathAllowed(c *gin.Context, relativePath string) bool {
session, ok := security.CurrentSession(c)
if !ok || h.db == nil {
return false
}
if session.Scope == database.RBACScopeAll {
return true
}
rel := filepath.ToSlash(filepath.Clean(filepath.FromSlash(strings.TrimSpace(relativePath))))
rel = strings.Trim(rel, "/")
if conversationID, ownerUserID, found := h.db.GetChatUploadArtifact(rel); found {
return strings.TrimSpace(ownerUserID) == session.UserID || h.db.UserCanAccessResource(session.UserID, session.Scope, "conversation", conversationID)
}
parts := strings.Split(strings.Trim(rel, "/"), "/")
if len(parts) < 2 || parts[1] == "" || parts[1] == "_manual" {
return false
}
return h.db.UserCanAccessResource(session.UserID, session.Scope, "conversation", parts[1])
}
func (h *ChatUploadsHandler) absRoot() (string, error) {
@@ -175,6 +202,21 @@ func (h *ChatUploadsHandler) List(c *gin.Context) {
}
folders = filteredFolders
}
files = filterSlice(files, func(file ChatUploadFileItem) bool {
return h.pathAllowed(c, file.RelativePath)
})
folders = filterSlice(folders, func(folder string) bool {
if h.pathAllowed(c, folder) {
return true
}
prefix := strings.TrimSuffix(folder, "/") + "/"
for _, file := range files {
if strings.HasPrefix(file.RelativePath, prefix) {
return true
}
}
return false
})
sort.Strings(folders)
sort.Slice(files, func(i, j int) bool {
return files[i].ModifiedUnix > files[j].ModifiedUnix
@@ -185,6 +227,10 @@ func (h *ChatUploadsHandler) List(c *gin.Context) {
// Download GET /api/chat-uploads/download?path=...
func (h *ChatUploadsHandler) Download(c *gin.Context) {
p := c.Query("path")
if !h.pathAllowed(c, p) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
abs, err := h.resolveUnderChatUploads(p)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -209,6 +255,10 @@ func (h *ChatUploadsHandler) Delete(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid body"})
return
}
if !h.pathAllowed(c, body.Path) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
abs, err := h.resolveUnderChatUploads(body.Path)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -238,6 +288,7 @@ func (h *ChatUploadsHandler) Delete(c *gin.Context) {
return
}
}
_ = h.db.DeleteChatUploadArtifactPath(filepath.ToSlash(filepath.Clean(filepath.FromSlash(body.Path))))
if h.audit != nil {
h.audit.RecordOK(c, "file", "delete", "删除对话附件", "chat_upload", body.Path, nil)
}
@@ -272,6 +323,10 @@ func (h *ChatUploadsHandler) Mkdir(c *gin.Context) {
if parent == "." {
parent = ""
}
if !h.pathAllowed(c, parent) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
root, err := h.absRoot()
if err != nil {
@@ -327,6 +382,10 @@ func (h *ChatUploadsHandler) Rename(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid body"})
return
}
if !h.pathAllowed(c, body.Path) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
newName := strings.TrimSpace(body.NewName)
if newName == "" || strings.ContainsAny(newName, `/\`) {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid newName"})
@@ -354,6 +413,8 @@ func (h *ChatUploadsHandler) Rename(c *gin.Context) {
return
}
newRel, _ := filepath.Rel(root, newAbs)
oldRel := filepath.ToSlash(filepath.Clean(filepath.FromSlash(body.Path)))
_ = h.db.RenameChatUploadArtifactPath(oldRel, filepath.ToSlash(newRel))
c.JSON(http.StatusOK, gin.H{"ok": true, "relativePath": filepath.ToSlash(newRel)})
}
@@ -365,6 +426,10 @@ type chatUploadContentBody struct {
// GetContent GET /api/chat-uploads/content?path=...
func (h *ChatUploadsHandler) GetContent(c *gin.Context) {
p := c.Query("path")
if !h.pathAllowed(c, p) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
abs, err := h.resolveUnderChatUploads(p)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -398,6 +463,10 @@ func (h *ChatUploadsHandler) PutContent(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid body"})
return
}
if !h.pathAllowed(c, body.Path) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
if !utf8.ValidString(body.Content) {
c.JSON(http.StatusBadRequest, gin.H{"error": "content must be valid UTF-8"})
return
@@ -444,6 +513,10 @@ func (h *ChatUploadsHandler) Upload(c *gin.Context) {
var targetDir string
targetRel := strings.TrimSpace(c.PostForm("relativeDir"))
if targetRel != "" {
if !h.pathAllowed(c, targetRel) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
absDir, err := h.resolveUnderChatUploads(targetRel)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
@@ -467,13 +540,17 @@ func (h *ChatUploadsHandler) Upload(c *gin.Context) {
targetDir = absDir
} else {
convID := strings.TrimSpace(c.PostForm("conversationId"))
dateStr := time.Now().Format("2006-01-02")
if !h.pathAllowed(c, filepath.ToSlash(filepath.Join(dateStr, convID))) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
convDir := convID
if convDir == "" {
convDir = "_manual"
} else {
convDir = strings.ReplaceAll(convDir, string(filepath.Separator), "_")
}
dateStr := time.Now().Format("2006-01-02")
targetDir = filepath.Join(root, dateStr, convDir)
if err := os.MkdirAll(targetDir, 0755); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -514,6 +591,16 @@ func (h *ChatUploadsHandler) Upload(c *gin.Context) {
}
rel, _ := filepath.Rel(root, fullPath)
absSaved, _ := filepath.Abs(fullPath)
if session, ok := security.CurrentSession(c); ok {
conversationID := strings.TrimSpace(c.PostForm("conversationId"))
if conversationID == "" {
parts := strings.Split(filepath.ToSlash(rel), "/")
if len(parts) >= 2 {
conversationID = parts[1]
}
}
_ = h.db.UpsertChatUploadArtifact(filepath.ToSlash(rel), conversationID, session.UserID)
}
if h.audit != nil {
h.audit.RecordOK(c, "file", "upload", "上传对话附件", "chat_upload", filepath.ToSlash(rel), map[string]interface{}{
"name": unique,
+90 -10
View File
@@ -16,6 +16,7 @@ import (
"cyberstrike-ai/internal/agents"
"cyberstrike-ai/internal/audit"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/knowledge"
"cyberstrike-ai/internal/mcp"
"cyberstrike-ai/internal/mcp/builtin"
@@ -89,11 +90,32 @@ type ConfigHandler struct {
appUpdater AppUpdater // App更新器(可选)
robotRestarter RobotRestarter // 机器人连接重启器(可选),ApplyConfig 时重启钉钉/飞书
audit *audit.Service
db *database.DB
logger *zap.Logger
mu sync.RWMutex
lastEmbeddingConfig *config.EmbeddingConfig // 上一次的嵌入模型配置(用于检测变更)
}
func (h *ConfigHandler) SetDB(db *database.DB) {
h.db = db
}
func (h *ConfigHandler) validateRobotServiceAccounts(robots config.RobotsConfig) error {
if h.db == nil {
return fmt.Errorf("RBAC 服务不可用,无法校验机器人服务账号")
}
for platform, userID := range robots.ServiceAccountUserIDs() {
user, err := h.db.GetRBACUserByID(userID)
if err != nil {
return fmt.Errorf("robots.%s.auth.service_user_id 对应用户不存在", platform)
}
if !user.Enabled {
return fmt.Errorf("robots.%s.auth.service_user_id 对应用户已禁用", platform)
}
}
return nil
}
// AttackChainUpdater 攻击链处理器更新接口
type AttackChainUpdater interface {
UpdateConfig(cfg *config.OpenAIConfig)
@@ -319,13 +341,18 @@ func (h *ConfigHandler) GetConfig(c *gin.Context) {
subAgentCount = len(agents.MergeYAMLAndMarkdown(h.config.MultiAgent.SubAgents, load.SubAgents))
}
multiPub := config.MultiAgentPublic{
Enabled: h.config.MultiAgent.Enabled,
RobotDefaultAgentMode: config.NormalizeRobotAgentMode(h.config.MultiAgent),
BatchUseMultiAgent: h.config.MultiAgent.BatchUseMultiAgent,
SubAgentCount: subAgentCount,
Orchestration: config.NormalizeMultiAgentOrchestration(h.config.MultiAgent.Orchestration),
PlanExecuteLoopMaxIterations: h.config.MultiAgent.PlanExecuteLoopMaxIterations,
ToolSearchAlwaysVisibleTools: append([]string(nil), h.config.MultiAgent.EinoMiddleware.ToolSearchAlwaysVisibleTools...),
Enabled: h.config.MultiAgent.Enabled,
RobotDefaultAgentMode: config.NormalizeRobotAgentMode(h.config.MultiAgent),
BatchUseMultiAgent: h.config.MultiAgent.BatchUseMultiAgent,
SubAgentCount: subAgentCount,
Orchestration: config.NormalizeMultiAgentOrchestration(h.config.MultiAgent.Orchestration),
PlanExecuteLoopMaxIterations: h.config.MultiAgent.PlanExecuteLoopMaxIterations,
SummarizationUserIntentLedgerMaxRunes: h.config.MultiAgent.EinoMiddleware.SummarizationUserIntentLedgerMaxRunesEffective(),
SummarizationUserIntentLedgerEntryMaxRunes: h.config.MultiAgent.EinoMiddleware.SummarizationUserIntentLedgerEntryMaxRunesEffective(),
LatestUserMessageMaxRunes: h.config.MultiAgent.EinoMiddleware.LatestUserMessageMaxRunesEffective(),
LatestUserMessageHeadRunes: h.config.MultiAgent.EinoMiddleware.LatestUserMessageHeadRunesEffective(),
LatestUserMessageTailRunes: h.config.MultiAgent.EinoMiddleware.LatestUserMessageTailRunesEffective(),
ToolSearchAlwaysVisibleTools: append([]string(nil), h.config.MultiAgent.EinoMiddleware.ToolSearchAlwaysVisibleTools...),
ToolSearchAlwaysVisibleEffectiveTools: mergeToolNameLists(
h.config.MultiAgent.EinoMiddleware.ToolSearchAlwaysVisibleTools,
builtin.GetAllBuiltinTools(),
@@ -822,6 +849,14 @@ func (h *ConfigHandler) UpdateConfig(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if err := config.ValidateRobotsAuthorization(*req.Robots); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if err := h.validateRobotServiceAccounts(*req.Robots); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
h.config.Robots = *req.Robots
h.logger.Info("更新机器人配置",
zap.Bool("wechat_enabled", h.config.Robots.Wechat.Enabled),
@@ -853,6 +888,41 @@ func (h *ConfigHandler) UpdateConfig(c *gin.Context) {
if req.MultiAgent.PlanExecuteLoopMaxIterations != nil {
h.config.MultiAgent.PlanExecuteLoopMaxIterations = *req.MultiAgent.PlanExecuteLoopMaxIterations
}
if req.MultiAgent.SummarizationUserIntentLedgerMaxRunes != nil {
v := *req.MultiAgent.SummarizationUserIntentLedgerMaxRunes
if v < 0 {
v = 0
}
h.config.MultiAgent.EinoMiddleware.SummarizationUserIntentLedgerMaxRunes = v
}
if req.MultiAgent.SummarizationUserIntentLedgerEntryMaxRunes != nil {
v := *req.MultiAgent.SummarizationUserIntentLedgerEntryMaxRunes
if v < 0 {
v = 0
}
h.config.MultiAgent.EinoMiddleware.SummarizationUserIntentLedgerEntryMaxRunes = v
}
if req.MultiAgent.LatestUserMessageMaxRunes != nil {
v := *req.MultiAgent.LatestUserMessageMaxRunes
if v < 0 {
v = 0
}
h.config.MultiAgent.EinoMiddleware.LatestUserMessageMaxRunes = v
}
if req.MultiAgent.LatestUserMessageHeadRunes != nil {
v := *req.MultiAgent.LatestUserMessageHeadRunes
if v < 0 {
v = 0
}
h.config.MultiAgent.EinoMiddleware.LatestUserMessageHeadRunes = v
}
if req.MultiAgent.LatestUserMessageTailRunes != nil {
v := *req.MultiAgent.LatestUserMessageTailRunes
if v < 0 {
v = 0
}
h.config.MultiAgent.EinoMiddleware.LatestUserMessageTailRunes = v
}
if req.MultiAgent.ToolSearchAlwaysVisibleTools != nil {
h.config.MultiAgent.EinoMiddleware.ToolSearchAlwaysVisibleTools = dedupeToolNameList(*req.MultiAgent.ToolSearchAlwaysVisibleTools)
}
@@ -861,6 +931,11 @@ func (h *ConfigHandler) UpdateConfig(c *gin.Context) {
zap.String("robot_default_agent_mode", config.NormalizeRobotAgentMode(h.config.MultiAgent)),
zap.Bool("batch_use_multi_agent", h.config.MultiAgent.BatchUseMultiAgent),
zap.Int("plan_execute_loop_max_iterations", h.config.MultiAgent.PlanExecuteLoopMaxIterations),
zap.Int("summarization_user_intent_ledger_max_runes", h.config.MultiAgent.EinoMiddleware.SummarizationUserIntentLedgerMaxRunesEffective()),
zap.Int("summarization_user_intent_ledger_entry_max_runes", h.config.MultiAgent.EinoMiddleware.SummarizationUserIntentLedgerEntryMaxRunesEffective()),
zap.Int("latest_user_message_max_runes", h.config.MultiAgent.EinoMiddleware.LatestUserMessageMaxRunesEffective()),
zap.Int("latest_user_message_head_runes", h.config.MultiAgent.EinoMiddleware.LatestUserMessageHeadRunesEffective()),
zap.Int("latest_user_message_tail_runes", h.config.MultiAgent.EinoMiddleware.LatestUserMessageTailRunesEffective()),
zap.Int("tool_search_always_visible_tools", len(h.config.MultiAgent.EinoMiddleware.ToolSearchAlwaysVisibleTools)),
)
}
@@ -1287,7 +1362,7 @@ func (h *ConfigHandler) ApplyConfig(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": "初始化知识库失败: " + err.Error()})
return
}
h.logger.Info("知识库动态初始化完成,工具已注册")
h.logger.Debug("知识库动态初始化完成,工具已注册")
}
// 检查嵌入模型配置是否变更(需要在锁外执行,避免阻塞)
@@ -1366,10 +1441,10 @@ func (h *ConfigHandler) ApplyConfig(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": "重新加载工具配置失败: " + err.Error()})
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服务器中的工具
h.mcpServer.ClearTools()
@@ -1932,6 +2007,11 @@ func updateMultiAgentConfig(doc *yaml.Node, cfg config.MultiAgentConfig) {
setBoolInMap(maNode, "batch_use_multi_agent", cfg.BatchUseMultiAgent)
setIntInMap(maNode, "plan_execute_loop_max_iterations", cfg.PlanExecuteLoopMaxIterations)
mwNode := ensureMap(maNode, "eino_middleware")
setIntInMap(mwNode, "summarization_user_intent_ledger_max_runes", cfg.EinoMiddleware.SummarizationUserIntentLedgerMaxRunesEffective())
setIntInMap(mwNode, "summarization_user_intent_ledger_entry_max_runes", cfg.EinoMiddleware.SummarizationUserIntentLedgerEntryMaxRunesEffective())
setIntInMap(mwNode, "latest_user_message_max_runes", cfg.EinoMiddleware.LatestUserMessageMaxRunesEffective())
setIntInMap(mwNode, "latest_user_message_head_runes", cfg.EinoMiddleware.LatestUserMessageHeadRunesEffective())
setIntInMap(mwNode, "latest_user_message_tail_runes", cfg.EinoMiddleware.LatestUserMessageTailRunesEffective())
setFlowStringSliceInMap(mwNode, "tool_search_always_visible_tools", dedupeToolNameList(cfg.EinoMiddleware.ToolSearchAlwaysVisibleTools))
}
+153 -35
View File
@@ -8,6 +8,7 @@ import (
"cyberstrike-ai/internal/audit"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
@@ -69,12 +70,23 @@ func (h *ConversationHandler) CreateConversation(c *gin.Context) {
meta := audit.ConversationCreateMetaFromGin(c, "api")
meta.ProjectID = strings.TrimSpace(req.ProjectID)
if !h.conversationProjectAllowed(c, meta.ProjectID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问目标项目"})
return
}
conv, err := h.db.CreateConversation(title, meta)
if err != nil {
h.logger.Error("创建对话失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if session, ok := security.CurrentSession(c); ok {
_ = h.db.SetResourceOwner("conversation", conv.ID, session.UserID)
_ = h.db.AssignResourceToUser(session.UserID, "conversation", conv.ID)
if conv.ProjectID != "" {
_ = h.db.AssignResourceToUser(session.UserID, "project", conv.ProjectID)
}
}
c.JSON(http.StatusOK, conv)
}
@@ -91,11 +103,28 @@ func (h *ConversationHandler) SetConversationProject(c *gin.Context) {
c.JSON(http.StatusNotFound, gin.H{"error": "对话不存在"})
return
}
if err := h.db.SetConversationProjectID(id, req.ProjectID); err != nil {
projectID := strings.TrimSpace(req.ProjectID)
if !h.conversationProjectAllowed(c, projectID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问目标项目"})
return
}
if err := h.db.SetConversationProjectID(id, projectID); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"success": true, "projectId": strings.TrimSpace(req.ProjectID)})
c.JSON(http.StatusOK, gin.H{"success": true, "projectId": projectID})
}
func (h *ConversationHandler) conversationProjectAllowed(c *gin.Context, projectID string) bool {
projectID = strings.TrimSpace(projectID)
if projectID == "" {
return true
}
session, ok := security.CurrentSession(c)
if !ok {
return false
}
return h.db.UserCanAccessResource(session.UserID, session.Scope, "project", projectID)
}
// ListConversations 列出对话
@@ -118,19 +147,20 @@ func (h *ConversationHandler) ListConversations(c *gin.Context) {
excludeGrouped := strings.TrimSpace(search) == "" && projectID == "" &&
(c.Query("exclude_grouped") == "true" || c.Query("exclude_grouped") == "1")
sortBy := strings.TrimSpace(c.Query("sort_by"))
session, _ := security.CurrentSession(c)
var conversations []*database.Conversation
var total int
var err error
if excludeGrouped {
conversations, err = h.db.ListUngroupedConversations(limit, offset, sortBy, projectID)
conversations, err = h.db.ListUngroupedConversationsForAccess(limit, offset, sortBy, projectID, session.UserID, session.Scope)
if err == nil {
total, err = h.db.CountUngroupedConversations(projectID)
total, err = h.db.CountUngroupedConversationsForAccess(projectID, session.UserID, session.Scope)
}
} else {
conversations, err = h.db.ListConversations(limit, offset, search, sortBy, projectID)
conversations, err = h.db.ListConversationsForAccess(limit, offset, search, sortBy, projectID, session.UserID, session.Scope)
if err == nil {
total, err = h.db.CountConversations(search, projectID)
total, err = h.db.CountConversationsForAccess(search, projectID, session.UserID, session.Scope)
}
}
if err != nil {
@@ -176,10 +206,17 @@ func (h *ConversationHandler) GetConversation(c *gin.Context) {
c.JSON(http.StatusOK, conv)
}
const (
defaultProcessDetailsPageLimit = 50
maxProcessDetailsPageLimit = 500
)
// GetMessageProcessDetails 获取指定消息的过程详情(按需加载)
// 查询参数:
// - summary=1:仅返回摘要(total / iterationCount / maxIteration
// - limit + offset:分页返回 processDetails(未指定 limit 时保持全量兼容
// - limit + offset:分页返回 processDetails(未指定 limit 时默认 50 条
// - anchorId:返回包含该过程详情锚点的一页,适合从工具按钮精准定位
// - full=1:显式返回全量 processDetails(用于导出/兼容旧集成,不建议 UI 展开时使用)
func (h *ConversationHandler) GetMessageProcessDetails(c *gin.Context) {
messageID := c.Param("id")
if messageID == "" {
@@ -199,52 +236,109 @@ func (h *ConversationHandler) GetMessageProcessDetails(c *gin.Context) {
return
}
limitStr := strings.TrimSpace(c.Query("limit"))
if limitStr != "" {
limit, err := strconv.Atoi(limitStr)
if err != nil || limit <= 0 {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid limit"})
return
}
if limit > 500 {
limit = 500
}
offset, _ := strconv.Atoi(strings.TrimSpace(c.Query("offset")))
if offset < 0 {
offset = 0
}
details, total, err := h.db.GetProcessDetailsPage(messageID, limit, offset)
fullStr := strings.TrimSpace(c.Query("full"))
if fullStr == "1" || strings.EqualFold(fullStr, "true") || strings.EqualFold(fullStr, "yes") {
details, err := h.db.GetProcessDetails(messageID)
if err != nil {
h.logger.Error("分页获取过程详情失败", zap.Error(err))
h.logger.Error("获取过程详情失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
details = database.DedupeConsecutiveProcessDetails(details)
out := processDetailsToJSON(h.logger, details)
out := processDetailsToJSON(h.logger, details, true)
c.JSON(http.StatusOK, gin.H{
"processDetails": out,
"total": total,
"offset": offset,
"limit": limit,
"hasMore": offset+len(out) < total,
"total": len(out),
"offset": 0,
"limit": len(out),
"hasMore": false,
})
return
}
details, err := h.db.GetProcessDetails(messageID)
limitStr := strings.TrimSpace(c.Query("limit"))
limit := defaultProcessDetailsPageLimit
if limitStr != "" {
parsedLimit, err := strconv.Atoi(limitStr)
if err != nil || parsedLimit <= 0 {
c.JSON(http.StatusBadRequest, gin.H{"error": "invalid limit"})
return
}
limit = parsedLimit
}
if limit > maxProcessDetailsPageLimit {
limit = maxProcessDetailsPageLimit
}
offset, _ := strconv.Atoi(strings.TrimSpace(c.Query("offset")))
if offset < 0 {
offset = 0
}
anchorID := strings.TrimSpace(c.Query("anchorId"))
if anchorID != "" {
anchorOffset, err := h.db.GetProcessDetailOffset(messageID, anchorID)
if err != nil {
h.logger.Warn("获取过程详情锚点位置失败", zap.Error(err), zap.String("messageID", messageID), zap.String("anchorID", anchorID))
c.JSON(http.StatusNotFound, gin.H{"error": err.Error()})
return
}
offset = anchorOffset - limit/3
if offset < 0 {
offset = 0
}
}
details, total, err := h.db.GetProcessDetailsPage(messageID, limit, offset)
if err != nil {
h.logger.Error("获取过程详情失败", zap.Error(err))
h.logger.Error("分页获取过程详情失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
details = database.DedupeConsecutiveProcessDetails(details)
out := processDetailsToJSON(h.logger, details)
c.JSON(http.StatusOK, gin.H{"processDetails": out, "total": len(out)})
out := processDetailsToJSON(h.logger, details, false)
// A page may end between tool_call and tool_result. Return the full-history
// execution summary so the UI can render terminal status without pretending
// that an unloaded result is still running.
summary, summaryErr := h.db.GetProcessDetailsSummary(messageID)
if summaryErr != nil {
h.logger.Warn("获取分页工具执行状态失败", zap.Error(summaryErr), zap.String("messageID", messageID))
}
var toolExecutions []database.ProcessDetailsToolExecution
if summary != nil {
toolExecutions = summary.ToolExecutions
}
c.JSON(http.StatusOK, gin.H{
"processDetails": out,
"toolExecutions": toolExecutions,
"total": total,
"offset": offset,
"limit": limit,
"hasMore": offset+len(out) < total,
})
}
func processDetailsToJSON(logger *zap.Logger, details []database.ProcessDetail) []map[string]interface{} {
// GetProcessDetail 获取单条完整过程详情。列表接口默认不给工具 payload,用户点开单条工具时再拉这里。
func (h *ConversationHandler) GetProcessDetail(c *gin.Context) {
id := strings.TrimSpace(c.Param("id"))
if id == "" {
c.JSON(http.StatusBadRequest, gin.H{"error": "process detail id required"})
return
}
detail, err := h.db.GetProcessDetailByID(id)
if err != nil {
h.logger.Error("获取过程详情失败", zap.Error(err))
c.JSON(http.StatusNotFound, gin.H{"error": "过程详情不存在"})
return
}
out := processDetailsToJSON(h.logger, []database.ProcessDetail{*detail}, true)
if len(out) == 0 {
c.JSON(http.StatusNotFound, gin.H{"error": "过程详情不存在"})
return
}
c.JSON(http.StatusOK, gin.H{"processDetail": out[0]})
}
func processDetailsToJSON(logger *zap.Logger, details []database.ProcessDetail, includeToolPayload bool) []map[string]interface{} {
out := make([]map[string]interface{}, 0, len(details))
for _, d := range details {
var data interface{}
@@ -253,6 +347,9 @@ func processDetailsToJSON(logger *zap.Logger, details []database.ProcessDetail)
logger.Warn("解析过程详情数据失败", zap.Error(err))
}
}
if !includeToolPayload {
data = summarizeProcessDetailData(d.EventType, data)
}
out = append(out, map[string]interface{}{
"id": d.ID,
"messageId": d.MessageID,
@@ -266,6 +363,27 @@ func processDetailsToJSON(logger *zap.Logger, details []database.ProcessDetail)
return out
}
func summarizeProcessDetailData(eventType string, data interface{}) interface{} {
m, ok := data.(map[string]interface{})
if !ok || (eventType != "tool_call" && eventType != "tool_result") {
return data
}
allow := map[string]bool{
"toolName": true, "toolCallId": true, "index": true, "total": true,
"success": true, "isError": true, "executionId": true,
"einoAgent": true, "einoRole": true, "einoScope": true, "orchestration": true,
"agentFacing": true,
}
out := make(map[string]interface{}, len(allow)+1)
for k, v := range m {
if allow[k] {
out[k] = v
}
}
out["_payloadDeferred"] = true
return out
}
// UpdateConversationRequest 更新对话请求
type UpdateConversationRequest struct {
Title string `json:"title"`
@@ -0,0 +1,75 @@
package handler
import (
"encoding/json"
"fmt"
"net/http/httptest"
"path/filepath"
"testing"
"cyberstrike-ai/internal/database"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
func TestProcessDetailsPageIncludesTerminalToolStatusAcrossPageBoundary(t *testing.T) {
gin.SetMode(gin.TestMode)
db, err := database.NewDB(filepath.Join(t.TempDir(), "process-details-page.db"), zap.NewNop())
if err != nil {
t.Fatalf("NewDB: %v", err)
}
t.Cleanup(func() { _ = db.Close() })
conversation, err := db.CreateConversation("page boundary", database.ConversationCreateMeta{})
if err != nil {
t.Fatalf("CreateConversation: %v", err)
}
message, err := db.AddMessage(conversation.ID, "assistant", "done", nil)
if err != nil {
t.Fatalf("AddMessage: %v", err)
}
for i := 1; i <= 4; i++ {
id := fmt.Sprintf("call-%d", i)
if err := db.AddProcessDetail(message.ID, conversation.ID, "tool_call", "call", map[string]interface{}{
"toolName": "http-framework-test", "toolCallId": id, "index": i, "total": 4,
}); err != nil {
t.Fatalf("AddProcessDetail(tool_call): %v", err)
}
}
for i := 1; i <= 4; i++ {
id := fmt.Sprintf("call-%d", i)
if err := db.AddProcessDetail(message.ID, conversation.ID, "tool_result", "result", map[string]interface{}{
"toolName": "http-framework-test", "toolCallId": id, "success": true,
}); err != nil {
t.Fatalf("AddProcessDetail(tool_result): %v", err)
}
}
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest("GET", "/api/messages/"+message.ID+"/process-details?limit=6&offset=0", nil)
c.Params = gin.Params{{Key: "id", Value: message.ID}}
NewConversationHandler(db, zap.NewNop()).GetMessageProcessDetails(c)
if w.Code != 200 {
t.Fatalf("status = %d: %s", w.Code, w.Body.String())
}
var response struct {
HasMore bool `json:"hasMore"`
ProcessDetails []map[string]interface{} `json:"processDetails"`
ToolExecutions []database.ProcessDetailsToolExecution `json:"toolExecutions"`
}
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatalf("decode response: %v", err)
}
if !response.HasMore || len(response.ProcessDetails) != 6 {
t.Fatalf("page hasMore=%v details=%d, want true/6", response.HasMore, len(response.ProcessDetails))
}
if len(response.ToolExecutions) != 4 {
t.Fatalf("tool executions = %d, want 4", len(response.ToolExecutions))
}
for i, execution := range response.ToolExecutions {
if execution.Status != "completed" {
t.Fatalf("execution %d status = %q, want completed", i, execution.Status)
}
}
}
+118
View File
@@ -0,0 +1,118 @@
package handler
import (
"bytes"
"encoding/json"
"net/http"
"net/http/httptest"
"path/filepath"
"testing"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
func TestCreateConversationRequiresProjectAccess(t *testing.T) {
gin.SetMode(gin.TestMode)
db, user := setupConversationRBACTest(t)
project, err := db.CreateProject(&database.Project{Name: "hidden"})
if err != nil {
t.Fatalf("CreateProject: %v", err)
}
handler := NewConversationHandler(db, zap.NewNop())
w := performConversationRequest(user, http.MethodPost, "/api/conversations", map[string]string{
"title": "blocked",
"projectId": project.ID,
}, handler.CreateConversation)
if w.Code != http.StatusForbidden {
t.Fatalf("status = %d, want %d: %s", w.Code, http.StatusForbidden, w.Body.String())
}
if err := db.AssignResourceToUser(user.ID, "project", project.ID); err != nil {
t.Fatalf("AssignResourceToUser: %v", err)
}
w = performConversationRequest(user, http.MethodPost, "/api/conversations", map[string]string{
"title": "allowed",
"projectId": project.ID,
}, handler.CreateConversation)
if w.Code != http.StatusOK {
t.Fatalf("status = %d, want %d: %s", w.Code, http.StatusOK, w.Body.String())
}
}
func TestSetConversationProjectRequiresProjectAccess(t *testing.T) {
gin.SetMode(gin.TestMode)
db, user := setupConversationRBACTest(t)
project, err := db.CreateProject(&database.Project{Name: "hidden"})
if err != nil {
t.Fatalf("CreateProject: %v", err)
}
conv, err := db.CreateConversation("owned", database.ConversationCreateMeta{})
if err != nil {
t.Fatalf("CreateConversation: %v", err)
}
if err := db.SetResourceOwner("conversation", conv.ID, user.ID); err != nil {
t.Fatalf("SetResourceOwner: %v", err)
}
if err := db.AssignResourceToUser(user.ID, "conversation", conv.ID); err != nil {
t.Fatalf("AssignResourceToUser conversation: %v", err)
}
handler := NewConversationHandler(db, zap.NewNop())
w := performConversationRequest(user, http.MethodPut, "/api/conversations/"+conv.ID+"/project", map[string]string{
"projectId": project.ID,
}, func(c *gin.Context) {
c.Params = gin.Params{{Key: "id", Value: conv.ID}}
handler.SetConversationProject(c)
})
if w.Code != http.StatusForbidden {
t.Fatalf("status = %d, want %d: %s", w.Code, http.StatusForbidden, w.Body.String())
}
if err := db.AssignResourceToUser(user.ID, "project", project.ID); err != nil {
t.Fatalf("AssignResourceToUser project: %v", err)
}
w = performConversationRequest(user, http.MethodPut, "/api/conversations/"+conv.ID+"/project", map[string]string{
"projectId": project.ID,
}, func(c *gin.Context) {
c.Params = gin.Params{{Key: "id", Value: conv.ID}}
handler.SetConversationProject(c)
})
if w.Code != http.StatusOK {
t.Fatalf("status = %d, want %d: %s", w.Code, http.StatusOK, w.Body.String())
}
}
func setupConversationRBACTest(t *testing.T) (*database.DB, *database.RBACUser) {
t.Helper()
db, err := database.NewDB(filepath.Join(t.TempDir(), "conversation-rbac.db"), zap.NewNop())
if err != nil {
t.Fatalf("NewDB: %v", err)
}
t.Cleanup(func() { _ = db.Close() })
user, err := db.CreateRBACUser("operator1", "Operator One", "hash", true, nil)
if err != nil {
t.Fatalf("CreateRBACUser: %v", err)
}
return db, user
}
func performConversationRequest(user *database.RBACUser, method, path string, body map[string]string, handler gin.HandlerFunc) *httptest.ResponseRecorder {
payload, _ := json.Marshal(body)
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(method, path, bytes.NewReader(payload))
c.Request.Header.Set("Content-Type", "application/json")
c.Set(security.ContextSessionKey, security.Session{
UserID: user.ID,
Username: user.Username,
Permissions: map[string]bool{"chat:write": true},
Scope: database.RBACScopeAssigned,
})
handler(c)
return w
}
@@ -13,11 +13,11 @@ import (
)
// rebindEinoRunningTask 中断并继续 / 空正文续跑:重建 cancel 链与超时 ctx,保持任务 running。
func (h *AgentHandler) rebindEinoRunningTask(conversationID string, timeoutCancel context.CancelFunc) (context.Context, context.CancelCauseFunc, context.Context, context.CancelFunc) {
func (h *AgentHandler) rebindEinoRunningTask(parent context.Context, conversationID string, timeoutCancel context.CancelFunc) (context.Context, context.CancelCauseFunc, context.Context, context.CancelFunc) {
if timeoutCancel != nil {
timeoutCancel()
}
baseCtx, cancelWithCause := context.WithCancelCause(context.Background())
baseCtx, cancelWithCause := context.WithCancelCause(detachedAgentContext(parent))
h.tasks.BindTaskCancel(conversationID, cancelWithCause)
taskCtx, newTimeoutCancel := context.WithTimeout(baseCtx, 600*time.Minute)
h.tasks.UpdateTaskStatus(conversationID, "running")
+8 -8
View File
@@ -116,7 +116,7 @@ func (h *AgentHandler) EinoSingleAgentLoopStream(c *gin.Context) {
"userMessageId": prep.UserMessageID,
})
}
if h.runRoleWorkflowStreamIfBound(&req, prep, sendEvent) {
if h.runRoleWorkflowStreamIfBound(c, &req, prep, sendEvent) {
return
}
@@ -139,9 +139,8 @@ func (h *AgentHandler) EinoSingleAgentLoopStream(c *gin.Context) {
"conversationId": conversationID,
})
stopKeepalive := make(chan struct{})
go sseKeepalive(c, stopKeepalive, &sseWriteMu)
defer close(stopKeepalive)
stopKeepalive := runSSEKeepalive(c, &sseWriteMu)
defer stopKeepalive()
if h.config == nil {
taskStatus = "failed"
@@ -154,7 +153,7 @@ func (h *AgentHandler) EinoSingleAgentLoopStream(c *gin.Context) {
var result *multiagent.RunResult
var runErr error
baseCtx, cancelWithCause = context.WithCancelCause(context.Background())
baseCtx, cancelWithCause = context.WithCancelCause(detachedAgentContext(c.Request.Context()))
taskCtx, timeoutCancel := context.WithTimeout(baseCtx, 600*time.Minute)
if _, err := h.tasks.StartTask(conversationID, req.Message, cancelWithCause); err != nil {
@@ -247,7 +246,7 @@ func (h *AgentHandler) EinoSingleAgentLoopStream(c *gin.Context) {
if h.tryContinueOnEinoEmptyResponse(taskCtx, mw, conversationID, result, &emptyResponseContinueAttempt, &curHistory, &curFinalMessage, progressCallback) {
mainIterationOffset += segmentMainIterationMax
timeoutCancel()
baseCtx, cancelWithCause, taskCtx, timeoutCancel = h.rebindEinoRunningTask(conversationID, timeoutCancel)
baseCtx, cancelWithCause, taskCtx, timeoutCancel = h.rebindEinoRunningTask(taskCtx, conversationID, timeoutCancel)
continue
}
timeoutCancel()
@@ -279,7 +278,7 @@ func (h *AgentHandler) EinoSingleAgentLoopStream(c *gin.Context) {
})
mainIterationOffset += segmentMainIterationMax
timeoutCancel()
baseCtx, cancelWithCause = context.WithCancelCause(context.Background())
baseCtx, cancelWithCause = context.WithCancelCause(detachedAgentContext(baseCtx))
h.tasks.BindTaskCancel(conversationID, cancelWithCause)
taskCtx, timeoutCancel = context.WithTimeout(baseCtx, 600*time.Minute)
h.tasks.UpdateTaskStatus(conversationID, "running")
@@ -381,7 +380,8 @@ func (h *AgentHandler) EinoSingleAgentLoop(c *gin.Context) {
prep, err := h.prepareMultiAgentSession(&req, c, "eino_agent")
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
status, msg := multiAgentHTTPErrorStatus(err)
c.JSON(status, gin.H{"error": msg})
return
}
h.activateHITLForConversation(prep.ConversationID, req.Hitl)
+23 -2
View File
@@ -9,6 +9,7 @@ import (
"cyberstrike-ai/internal/audit"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/mcp"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
@@ -67,7 +68,7 @@ func (h *ExternalMCPHandler) GetExternalMCPs(c *gin.Context) {
errorMsg := externalMCPStatusError(h.manager, name, status)
result[name] = ExternalMCPResponse{
Config: cfg,
Config: externalMCPConfigForResponse(c, cfg),
Status: status,
ToolCount: toolCount,
Error: errorMsg,
@@ -113,13 +114,33 @@ func (h *ExternalMCPHandler) GetExternalMCP(c *gin.Context) {
}
c.JSON(http.StatusOK, ExternalMCPResponse{
Config: cfg,
Config: externalMCPConfigForResponse(c, cfg),
Status: status,
ToolCount: toolCount,
Error: externalMCPStatusError(h.manager, name, status),
})
}
func externalMCPConfigForResponse(c *gin.Context, cfg config.ExternalMCPServerConfig) config.ExternalMCPServerConfig {
if security.SessionHasPermission(c, "mcp:write") {
return cfg
}
copyCfg := cfg
if len(cfg.Env) > 0 {
copyCfg.Env = make(map[string]string, len(cfg.Env))
for key := range cfg.Env {
copyCfg.Env[key] = "***"
}
}
if len(cfg.Headers) > 0 {
copyCfg.Headers = make(map[string]string, len(cfg.Headers))
for key := range cfg.Headers {
copyCfg.Headers[key] = "***"
}
}
return copyCfg
}
// externalMCPStatusError 在 error/disconnected 状态下返回最近错误(含断连原因)。
func externalMCPStatusError(manager *mcp.ExternalMCPManager, name, status string) string {
if status != "error" && status != "disconnected" {
+76 -3
View File
@@ -5,6 +5,7 @@ import (
"time"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
@@ -43,7 +44,8 @@ func (h *GroupHandler) CreateGroup(c *gin.Context) {
return
}
group, err := h.db.CreateGroup(req.Name, req.Icon)
session, _ := security.CurrentSession(c)
group, err := h.db.CreateGroup(req.Name, req.Icon, session.UserID)
if err != nil {
h.logger.Error("创建分组失败", zap.Error(err))
// 如果是名称重复错误,返回400状态码
@@ -60,7 +62,8 @@ func (h *GroupHandler) CreateGroup(c *gin.Context) {
// ListGroups 列出所有分组
func (h *GroupHandler) ListGroups(c *gin.Context) {
groups, err := h.db.ListGroups()
session, _ := security.CurrentSession(c)
groups, err := h.db.ListGroupsForAccess(session.UserID, session.Scope)
if err != nil {
h.logger.Error("获取分组列表失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -73,6 +76,10 @@ func (h *GroupHandler) ListGroups(c *gin.Context) {
// GetGroup 获取分组
func (h *GroupHandler) GetGroup(c *gin.Context) {
id := c.Param("id")
if !h.groupAllowed(c, id) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
group, err := h.db.GetGroup(id)
if err != nil {
@@ -93,6 +100,10 @@ type UpdateGroupRequest struct {
// UpdateGroup 更新分组
func (h *GroupHandler) UpdateGroup(c *gin.Context) {
id := c.Param("id")
if !h.groupAllowed(c, id) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
var req UpdateGroupRequest
if err := c.ShouldBindJSON(&req); err != nil {
@@ -129,6 +140,10 @@ func (h *GroupHandler) UpdateGroup(c *gin.Context) {
// DeleteGroup 删除分组
func (h *GroupHandler) DeleteGroup(c *gin.Context) {
id := c.Param("id")
if !h.groupAllowed(c, id) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
if err := h.db.DeleteGroup(id); err != nil {
h.logger.Error("删除分组失败", zap.Error(err))
@@ -152,6 +167,14 @@ func (h *GroupHandler) AddConversationToGroup(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if !h.groupConversationAllowed(c, req.ConversationID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
if !h.groupAllowed(c, req.GroupID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该分组"})
return
}
if err := h.db.AddConversationToGroup(req.ConversationID, req.GroupID); err != nil {
h.logger.Error("添加对话到分组失败", zap.Error(err))
@@ -166,6 +189,14 @@ func (h *GroupHandler) AddConversationToGroup(c *gin.Context) {
func (h *GroupHandler) RemoveConversationFromGroup(c *gin.Context) {
conversationID := c.Param("conversationId")
groupID := c.Param("id")
if !h.groupAllowed(c, groupID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该分组"})
return
}
if !h.groupConversationAllowed(c, conversationID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
if err := h.db.RemoveConversationFromGroup(conversationID, groupID); err != nil {
h.logger.Error("从分组中移除对话失败", zap.Error(err))
@@ -189,6 +220,10 @@ type GroupConversation struct {
// GetGroupConversations 获取分组中的所有对话
func (h *GroupHandler) GetGroupConversations(c *gin.Context) {
groupID := c.Param("id")
if !h.groupAllowed(c, groupID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该分组"})
return
}
searchQuery := c.Query("search") // 获取搜索参数
var conversations []*database.Conversation
@@ -210,6 +245,9 @@ func (h *GroupHandler) GetGroupConversations(c *gin.Context) {
// 获取每个对话在分组中的置顶状态
groupConvs := make([]GroupConversation, 0, len(conversations))
for _, conv := range conversations {
if conv == nil || !h.groupConversationAllowed(c, conv.ID) {
continue
}
// 查询分组内置顶状态
var groupPinned int
err := h.db.QueryRow(
@@ -242,8 +280,14 @@ func (h *GroupHandler) GetAllMappings(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
filtered := mappings[:0]
for _, mapping := range mappings {
if h.groupConversationAllowed(c, mapping.ConversationID) && h.groupAllowed(c, mapping.GroupID) {
filtered = append(filtered, mapping)
}
}
c.JSON(http.StatusOK, mappings)
c.JSON(http.StatusOK, filtered)
}
// UpdateConversationPinnedRequest 更新对话置顶状态请求
@@ -254,6 +298,10 @@ type UpdateConversationPinnedRequest struct {
// UpdateConversationPinned 更新对话置顶状态
func (h *GroupHandler) UpdateConversationPinned(c *gin.Context) {
conversationID := c.Param("id")
if !h.groupConversationAllowed(c, conversationID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
var req UpdateConversationPinnedRequest
if err := c.ShouldBindJSON(&req); err != nil {
@@ -278,6 +326,10 @@ type UpdateGroupPinnedRequest struct {
// UpdateGroupPinned 更新分组置顶状态
func (h *GroupHandler) UpdateGroupPinned(c *gin.Context) {
groupID := c.Param("id")
if !h.groupAllowed(c, groupID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该分组"})
return
}
var req UpdateGroupPinnedRequest
if err := c.ShouldBindJSON(&req); err != nil {
@@ -303,6 +355,14 @@ type UpdateConversationPinnedInGroupRequest struct {
func (h *GroupHandler) UpdateConversationPinnedInGroup(c *gin.Context) {
groupID := c.Param("id")
conversationID := c.Param("conversationId")
if !h.groupAllowed(c, groupID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该分组"})
return
}
if !h.groupConversationAllowed(c, conversationID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
var req UpdateConversationPinnedInGroupRequest
if err := c.ShouldBindJSON(&req); err != nil {
@@ -318,3 +378,16 @@ func (h *GroupHandler) UpdateConversationPinnedInGroup(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"message": "更新成功"})
}
func (h *GroupHandler) groupConversationAllowed(c *gin.Context, conversationID string) bool {
session, ok := security.CurrentSession(c)
if !ok {
return false
}
return h.db.UserCanAccessResource(session.UserID, session.Scope, "conversation", conversationID)
}
func (h *GroupHandler) groupAllowed(c *gin.Context, groupID string) bool {
session, ok := security.CurrentSession(c)
return ok && h.db.UserCanAccessGroup(session.UserID, session.Scope, groupID)
}
+19 -3
View File
@@ -611,6 +611,7 @@ func (h *AgentHandler) ListHITLPending(c *gin.Context) {
offset := (page - 1) * pageSize
q, args := h.buildHitlListQuery(false)
q, args = h.appendHitlListFilters(q, args, c)
q, args = appendConversationAccessSQL(q, args, "conversation_id", notificationAccessFromContext(c))
total, err := h.countHitlQuery(q, args)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -649,6 +650,10 @@ func (h *AgentHandler) DecideHITLInterrupt(c *gin.Context) {
c.JSON(500, gin.H{"error": "hitl manager unavailable"})
return
}
if !h.hitlInterruptAllowed(c, req.InterruptID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
if err := h.hitlManager.ResolveInterrupt(req.InterruptID, req.Decision, req.Comment, req.EditedArguments); err != nil {
c.JSON(http.StatusConflict, gin.H{"error": err.Error()})
return
@@ -673,6 +678,10 @@ func (h *AgentHandler) DismissHITLInterrupt(c *gin.Context) {
c.JSON(500, gin.H{"error": "hitl manager unavailable"})
return
}
if !h.hitlInterruptAllowed(c, req.InterruptID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
res, err := h.db.Exec(`UPDATE hitl_interrupts SET status='cancelled', decision='reject',
decision_comment='dismissed by user', decided_at=CURRENT_TIMESTAMP, decided_by='human'
WHERE id=? AND status='pending'`, req.InterruptID)
@@ -728,7 +737,6 @@ func (h *AgentHandler) interceptHITLForEinoTool(runCtx context.Context, cancelRu
return arguments, nil
}
type hitlConfigReq struct {
ConversationID string `json:"conversationId" binding:"required"`
HITLRequest
@@ -740,6 +748,10 @@ func (h *AgentHandler) GetHITLConversationConfig(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "conversationId is required"})
return
}
if !h.hitlConversationAllowed(c, conversationID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
cfg, err := h.loadHITLConversationConfig(conversationID)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -770,6 +782,10 @@ func (h *AgentHandler) UpsertHITLConversationConfig(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if !h.hitlConversationAllowed(c, req.ConversationID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
req.Mode = normalizeHitlMode(req.Mode)
req.Reviewer = normalizeHitlReviewer(req.Reviewer)
if strings.TrimSpace(req.Reviewer) == "" {
@@ -868,8 +884,8 @@ func (h *AgentHandler) SetHITLGlobalToolWhitelist(c *gin.Context) {
h.audit.RecordOK(c, "hitl", "tool_whitelist_update", "HITL 全局白名单更新", "hitl_config", "tool_whitelist", nil)
}
c.JSON(http.StatusOK, gin.H{
"ok": true,
"toolWhitelist": h.hitlConfigGlobalToolWhitelist(),
"ok": true,
"toolWhitelist": h.hitlConfigGlobalToolWhitelist(),
"hitlGlobalToolWhitelist": h.hitlConfigGlobalToolWhitelist(),
"hitlGlobalWhitelistMerged": false,
})
+69 -1
View File
@@ -10,6 +10,7 @@ import (
"time"
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
)
@@ -163,6 +164,7 @@ func (h *AgentHandler) ListHITLLogs(c *gin.Context) {
q, args := h.buildHitlListQuery(true)
q, args = h.appendHitlListFilters(q, args, c)
q, args = appendConversationAccessSQL(q, args, "conversation_id", notificationAccessFromContext(c))
total, err := h.countHitlQuery(q, args)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -207,6 +209,7 @@ func (h *AgentHandler) DeleteHITLLogs(c *gin.Context) {
if request.All {
where, args := h.buildHitlLogsWhere(true)
where, args = h.appendHitlListFilters(where, args, c)
where, args = appendConversationAccessSQL(where, args, "conversation_id", notificationAccessFromContext(c))
deleted, err = h.db.DeleteHitlInterruptLogsMatching(where, args)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -222,7 +225,12 @@ func (h *AgentHandler) DeleteHITLLogs(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "审计日志 ID 列表不能为空"})
return
}
deleted, err = h.db.DeleteHitlInterruptLogsByIDs(request.IDs)
ids, filterErr := h.filterAllowedHitlInterruptIDs(c, request.IDs)
if filterErr != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": filterErr.Error()})
return
}
deleted, err = h.db.DeleteHitlInterruptLogsByIDs(ids)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -259,5 +267,65 @@ func (h *AgentHandler) GetHITLLog(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if !h.hitlConversationAllowed(c, cid) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
c.JSON(http.StatusOK, hitlInterruptRowToMap(rowID, cid, mode, toolName, toolCallID, payload, rowStatus, decidedBy, messageID, decision, comment, createdAt, decidedAt))
}
func (h *AgentHandler) filterAllowedHitlInterruptIDs(c *gin.Context, ids []string) ([]string, error) {
clean := make([]string, 0, len(ids))
seen := map[string]struct{}{}
for _, id := range ids {
id = strings.TrimSpace(id)
if id == "" {
continue
}
if _, ok := seen[id]; ok {
continue
}
seen[id] = struct{}{}
clean = append(clean, id)
}
if len(clean) == 0 {
return clean, nil
}
query := `SELECT id, conversation_id FROM hitl_interrupts WHERE id IN (` + buildPlaceholders(len(clean)) + `)`
args := make([]interface{}, 0, len(clean))
for _, id := range clean {
args = append(args, id)
}
rows, err := h.db.Query(query, args...)
if err != nil {
return nil, err
}
defer rows.Close()
allowed := make([]string, 0, len(clean))
for rows.Next() {
var id, conversationID string
if err := rows.Scan(&id, &conversationID); err != nil {
continue
}
if h.hitlConversationAllowed(c, conversationID) {
allowed = append(allowed, id)
}
}
return allowed, rows.Err()
}
func (h *AgentHandler) hitlInterruptAllowed(c *gin.Context, interruptID string) bool {
var conversationID string
if err := h.db.QueryRow(`SELECT conversation_id FROM hitl_interrupts WHERE id = ?`, strings.TrimSpace(interruptID)).Scan(&conversationID); err != nil {
return false
}
return h.hitlConversationAllowed(c, conversationID)
}
func (h *AgentHandler) hitlConversationAllowed(c *gin.Context, conversationID string) bool {
session, ok := security.CurrentSession(c)
if !ok {
return false
}
return h.db.UserCanAccessResource(session.UserID, session.Scope, "conversation", conversationID)
}
+123 -15
View File
@@ -69,7 +69,7 @@ func (h *MonitorHandler) SetAgentHandler(ah *AgentHandler) {
h.agentHandler = ah
}
const monitorPageTopTools = 6
const monitorPageTopTools = 3
// MonitorStatsSummary 工具调用汇总
type MonitorStatsSummary struct {
@@ -120,9 +120,22 @@ func (h *MonitorHandler) Monitor(c *gin.Context) {
// 解析工具筛选参数(兼容 mcp__tool 与内部 mcp::tool
toolName := normalizeToolNameFilter(c.Query("tool"))
executions, total := h.loadExecutionListWithPagination(page, pageSize, status, toolName)
access := notificationAccessFromContext(c)
executions, total := h.loadExecutionListWithPagination(page, pageSize, status, toolName, access)
h.enrichExecutionsConversationID(executions)
summary, topTools := h.loadStatsSummary(monitorPageTopTools)
var summary *MonitorStatsSummary
var topTools []*mcp.ToolStats
if access.Scope == database.RBACScopeAll {
summary, topTools = h.loadStatsSummary(monitorPageTopTools)
} else if h.db != nil {
if scoped, err := h.db.LoadToolStatsSummaryForAccess(monitorPageTopTools, access); err == nil {
summary, topTools = dbStatsSummaryToMonitor(scoped), scoped.TopTools
} else {
summary, topTools = summarizeAccessibleExecutionPage(executions, monitorPageTopTools)
}
} else {
summary, topTools = summarizeAccessibleExecutionPage(executions, monitorPageTopTools)
}
totalPages := (total + pageSize - 1) / pageSize
if totalPages == 0 {
@@ -142,6 +155,31 @@ func (h *MonitorHandler) Monitor(c *gin.Context) {
})
}
func summarizeAccessibleExecutionPage(executions []*mcp.ToolExecution, topN int) (*MonitorStatsSummary, []*mcp.ToolStats) {
stats := map[string]*mcp.ToolStats{}
for _, exec := range executions {
if exec == nil {
continue
}
stat := stats[exec.ToolName]
if stat == nil {
stat = &mcp.ToolStats{ToolName: exec.ToolName}
stats[exec.ToolName] = stat
}
stat.TotalCalls++
if exec.Status == "failed" || exec.Status == "cancelled" {
stat.FailedCalls++
} else if exec.Status == "completed" {
stat.SuccessCalls++
}
started := exec.StartTime
if stat.LastCallTime == nil || started.After(*stat.LastCallTime) {
stat.LastCallTime = &started
}
}
return summarizeToolStats(stats, topN)
}
func (h *MonitorHandler) monitorRetentionDays() int {
if h.monitorRetention != nil {
return h.monitorRetention.RetentionDays()
@@ -154,9 +192,9 @@ func (h *MonitorHandler) loadExecutions() []*mcp.ToolExecution {
return executions
}
func (h *MonitorHandler) loadExecutionListWithPagination(page, pageSize int, status, toolName string) ([]*mcp.ToolExecution, int) {
func (h *MonitorHandler) loadExecutionListWithPagination(page, pageSize int, status, toolName string, access database.RBACListAccess) ([]*mcp.ToolExecution, int) {
if h.db == nil {
allExecutions := h.mcpServer.GetAllExecutions()
allExecutions := filterToolExecutionsForAccess(h.mcpServer.GetAllExecutions(), access, h.db)
if status != "" || toolName != "" {
filtered := make([]*mcp.ToolExecution, 0)
for _, exec := range allExecutions {
@@ -189,13 +227,13 @@ func (h *MonitorHandler) loadExecutionListWithPagination(page, pageSize int, sta
}
offset := (page - 1) * pageSize
executions, err := h.db.LoadToolExecutionListPage(offset, pageSize, status, toolName)
executions, err := h.db.LoadToolExecutionListPageForAccess(offset, pageSize, status, toolName, access)
if err != nil {
h.logger.Warn("从数据库加载执行记录列表失败,回退到内存数据", zap.Error(err))
return h.loadExecutionListWithPaginationFromMemory(page, pageSize, status, toolName)
return h.loadExecutionListWithPaginationFromMemory(page, pageSize, status, toolName, access)
}
total, err := h.db.CountToolExecutions(status, toolName)
total, err := h.db.CountToolExecutionsForAccess(status, toolName, access)
if err != nil {
h.logger.Warn("获取执行记录总数失败", zap.Error(err))
total = offset + len(executions)
@@ -207,8 +245,8 @@ func (h *MonitorHandler) loadExecutionListWithPagination(page, pageSize int, sta
return executions, total
}
func (h *MonitorHandler) loadExecutionListWithPaginationFromMemory(page, pageSize int, status, toolName string) ([]*mcp.ToolExecution, int) {
allExecutions := h.mcpServer.GetAllExecutions()
func (h *MonitorHandler) loadExecutionListWithPaginationFromMemory(page, pageSize int, status, toolName string, access database.RBACListAccess) ([]*mcp.ToolExecution, int) {
allExecutions := filterToolExecutionsForAccess(h.mcpServer.GetAllExecutions(), access, h.db)
if status != "" || toolName != "" {
filtered := make([]*mcp.ToolExecution, 0)
for _, exec := range allExecutions {
@@ -260,6 +298,50 @@ func slimToolExecution(exec *mcp.ToolExecution) *mcp.ToolExecution {
return slim
}
func filterToolExecutionsForAccess(executions []*mcp.ToolExecution, access database.RBACListAccess, db *database.DB) []*mcp.ToolExecution {
if access.Scope == database.RBACScopeAll {
return executions
}
out := make([]*mcp.ToolExecution, 0, len(executions))
for _, exec := range executions {
if toolExecutionVisible(exec, access, db) {
out = append(out, exec)
}
}
return out
}
func toolExecutionVisible(exec *mcp.ToolExecution, access database.RBACListAccess, db *database.DB) bool {
if exec == nil || strings.TrimSpace(access.UserID) == "" {
return false
}
if access.Scope == database.RBACScopeAll || strings.TrimSpace(exec.OwnerUserID) == strings.TrimSpace(access.UserID) {
return true
}
conversationID := strings.TrimSpace(exec.ConversationID)
return conversationID != "" && db != nil && db.UserCanAccessResource(access.UserID, access.Scope, "conversation", conversationID)
}
func (h *MonitorHandler) monitorExecutionAllowed(c *gin.Context, id string) bool {
access := notificationAccessFromContext(c)
if access.Scope == database.RBACScopeAll {
return true
}
id = strings.TrimSpace(id)
if id == "" {
return false
}
if exec, ok := h.mcpServer.GetExecution(id); ok {
return toolExecutionVisible(exec, access, h.db)
}
if h.externalMCPMgr != nil {
if exec, ok := h.externalMCPMgr.GetExecution(id); ok {
return toolExecutionVisible(exec, access, h.db)
}
}
return h.db != nil && h.db.UserCanAccessToolExecution(access.UserID, access.Scope, id)
}
func (h *MonitorHandler) loadExecutionsWithPagination(page, pageSize int, status, toolName string) ([]*mcp.ToolExecution, int) {
if h.db == nil {
allExecutions := h.mcpServer.GetAllExecutions()
@@ -453,6 +535,10 @@ func (h *MonitorHandler) loadStatsMap() map[string]*mcp.ToolStats {
// GetExecution 获取特定执行记录
func (h *MonitorHandler) GetExecution(c *gin.Context) {
id := c.Param("id")
if !h.monitorExecutionAllowed(c, id) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
// 先从内部MCP服务器查找
exec, exists := h.mcpServer.GetExecution(id)
@@ -493,6 +579,10 @@ func (h *MonitorHandler) CancelExecution(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "执行记录ID不能为空"})
return
}
if !h.monitorExecutionAllowed(c, id) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
note := ""
dec := json.NewDecoder(c.Request.Body)
var body struct {
@@ -575,7 +665,7 @@ func (h *MonitorHandler) lookupExecution(id string) *mcp.ToolExecution {
return nil
}
// BatchGetToolNames 批量获取工具执行的工具名称(消除前端 N+1 请求)
// BatchGetToolNames 批量获取工具执行摘要(消除前端 N+1 请求)
func (h *MonitorHandler) BatchGetToolNames(c *gin.Context) {
var req struct {
IDs []string `json:"ids"`
@@ -585,24 +675,32 @@ func (h *MonitorHandler) BatchGetToolNames(c *gin.Context) {
return
}
result := make(map[string]string, len(req.IDs))
type executionSummary struct {
ToolName string `json:"toolName"`
Status string `json:"status"`
}
result := make(map[string]executionSummary, len(req.IDs))
for _, id := range req.IDs {
if !h.monitorExecutionAllowed(c, id) {
continue
}
// 先从内部MCP服务器查找
if exec, exists := h.mcpServer.GetExecution(id); exists {
result[id] = exec.ToolName
result[id] = executionSummary{ToolName: exec.ToolName, Status: exec.Status}
continue
}
// 再从外部MCP管理器查找
if h.externalMCPMgr != nil {
if exec, exists := h.externalMCPMgr.GetExecution(id); exists {
result[id] = exec.ToolName
result[id] = executionSummary{ToolName: exec.ToolName, Status: exec.Status}
continue
}
}
// 最后从数据库查找
if h.db != nil {
if exec, err := h.db.GetToolExecution(id); err == nil && exec != nil {
result[id] = exec.ToolName
result[id] = executionSummary{ToolName: exec.ToolName, Status: exec.Status}
}
}
}
@@ -750,6 +848,10 @@ func (h *MonitorHandler) DeleteExecution(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "执行记录ID不能为空"})
return
}
if !h.monitorExecutionAllowed(c, id) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问该资源"})
return
}
// 如果使用数据库,先获取执行记录信息,然后删除并更新统计
if h.db != nil {
@@ -818,6 +920,12 @@ func (h *MonitorHandler) DeleteExecutions(c *gin.Context) {
c.JSON(http.StatusBadRequest, gin.H{"error": "执行记录ID列表不能为空"})
return
}
for _, id := range request.IDs {
if !h.monitorExecutionAllowed(c, id) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问一个或多个执行记录"})
return
}
}
// 如果使用数据库,先获取执行记录信息,然后删除并更新统计
if h.db != nil {
+8 -7
View File
@@ -133,7 +133,7 @@ func (h *AgentHandler) MultiAgentLoopStream(c *gin.Context) {
"userMessageId": prep.UserMessageID,
})
}
if h.runRoleWorkflowStreamIfBound(&req, prep, sendEvent) {
if h.runRoleWorkflowStreamIfBound(c, &req, prep, sendEvent) {
return
}
@@ -156,14 +156,13 @@ func (h *AgentHandler) MultiAgentLoopStream(c *gin.Context) {
"conversationId": conversationID,
})
stopKeepalive := make(chan struct{})
go sseKeepalive(c, stopKeepalive, &sseWriteMu)
defer close(stopKeepalive)
stopKeepalive := runSSEKeepalive(c, &sseWriteMu)
defer stopKeepalive()
var result *multiagent.RunResult
var runErr error
baseCtx, cancelWithCause = context.WithCancelCause(context.Background())
baseCtx, cancelWithCause = context.WithCancelCause(detachedAgentContext(c.Request.Context()))
taskCtx, timeoutCancel := context.WithTimeout(baseCtx, 600*time.Minute)
if _, err := h.tasks.StartTask(conversationID, req.Message, cancelWithCause); err != nil {
@@ -259,7 +258,7 @@ func (h *AgentHandler) MultiAgentLoopStream(c *gin.Context) {
if h.tryContinueOnEinoEmptyResponse(taskCtx, mw, conversationID, result, &emptyResponseContinueAttempt, &curHistory, &curFinalMessage, progressCallback) {
mainIterationOffset += segmentMainIterationMax
timeoutCancel()
baseCtx, cancelWithCause, taskCtx, timeoutCancel = h.rebindEinoRunningTask(conversationID, timeoutCancel)
baseCtx, cancelWithCause, taskCtx, timeoutCancel = h.rebindEinoRunningTask(taskCtx, conversationID, timeoutCancel)
continue
}
timeoutCancel()
@@ -291,7 +290,7 @@ func (h *AgentHandler) MultiAgentLoopStream(c *gin.Context) {
})
mainIterationOffset += segmentMainIterationMax
timeoutCancel()
baseCtx, cancelWithCause = context.WithCancelCause(context.Background())
baseCtx, cancelWithCause = context.WithCancelCause(detachedAgentContext(baseCtx))
h.tasks.BindTaskCancel(conversationID, cancelWithCause)
taskCtx, timeoutCancel = context.WithTimeout(baseCtx, 600*time.Minute)
h.tasks.UpdateTaskStatus(conversationID, "running")
@@ -541,6 +540,8 @@ func formatInterruptContinueUserMessage(note string) string {
func multiAgentHTTPErrorStatus(err error) (int, string) {
msg := err.Error()
switch {
case strings.Contains(msg, "无权访问"):
return http.StatusForbidden, msg
case strings.Contains(msg, "对话不存在"):
return http.StatusNotFound, msg
case strings.Contains(msg, "未找到该 WebShell"):
+31 -5
View File
@@ -8,6 +8,7 @@ import (
"cyberstrike-ai/internal/audit"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/mcp/builtin"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
@@ -30,16 +31,34 @@ func (h *AgentHandler) prepareMultiAgentSession(req *ChatRequest, c *gin.Context
}
conversationID := strings.TrimSpace(req.ConversationID)
projectID := strings.TrimSpace(effectiveProjectID(h.config, req.ProjectID))
webshellID := strings.TrimSpace(req.WebShellConnectionID)
session, hasSession := security.CurrentSession(c)
if !hasSession || !session.Permissions["chat:write"] {
return nil, fmt.Errorf("无权写入对话")
}
canAccess := func(resourceType, resourceID string) bool {
if !hasSession || h.db == nil || strings.TrimSpace(resourceID) == "" {
return false
}
return h.db.UserCanAccessResource(session.UserID, session.Scope, resourceType, resourceID)
}
if projectID != "" && (!session.Permissions["project:read"] || !canAccess("project", projectID)) {
return nil, fmt.Errorf("无权访问目标项目")
}
if webshellID != "" && (!session.Permissions["webshell:write"] || !canAccess("webshell", webshellID)) {
return nil, fmt.Errorf("无权访问该 WebShell 连接")
}
createdNew := false
if conversationID == "" {
title := safeTruncateString(req.Message, 50)
var conv *database.Conversation
var err error
meta := audit.ConversationCreateMetaFromGin(c, source)
meta.ProjectID = effectiveProjectID(h.config, req.ProjectID)
if strings.TrimSpace(req.WebShellConnectionID) != "" {
meta.ProjectID = projectID
if webshellID != "" {
meta.Source = source + "_webshell"
meta.WebShellConnectionID = strings.TrimSpace(req.WebShellConnectionID)
meta.WebShellConnectionID = webshellID
conv, err = h.db.CreateConversationWithWebshell(meta.WebShellConnectionID, title, meta)
} else {
conv, err = h.db.CreateConversation(title, meta)
@@ -49,10 +68,17 @@ func (h *AgentHandler) prepareMultiAgentSession(req *ChatRequest, c *gin.Context
}
conversationID = conv.ID
createdNew = true
if hasSession {
_ = h.db.SetResourceOwner("conversation", conversationID, session.UserID)
_ = h.db.AssignResourceToUser(session.UserID, "conversation", conversationID)
}
} else {
if _, err := h.db.GetConversation(conversationID); err != nil {
return nil, fmt.Errorf("对话不存在")
}
if !canAccess("conversation", conversationID) {
return nil, fmt.Errorf("无权访问该对话")
}
}
agentHistoryMessages, err := h.loadHistoryFromAgentTrace(conversationID)
@@ -67,8 +93,8 @@ func (h *AgentHandler) prepareMultiAgentSession(req *ChatRequest, c *gin.Context
finalMessage := req.Message
var roleTools []string
if req.WebShellConnectionID != "" {
conn, errConn := h.db.GetWebshellConnection(strings.TrimSpace(req.WebShellConnectionID))
if webshellID != "" {
conn, errConn := h.db.GetWebshellConnection(webshellID)
if errConn != nil || conn == nil {
h.logger.Warn("WebShell AI 助手:未找到连接", zap.String("id", req.WebShellConnectionID), zap.Error(errConn))
return nil, fmt.Errorf("未找到该 WebShell 连接")
+206 -68
View File
@@ -9,6 +9,8 @@ import (
"time"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
@@ -89,6 +91,10 @@ func normalizedSinceSec(sinceMs int64) int64 {
return 0
}
func ptrTime(t time.Time) *time.Time {
return &t
}
func normalizeSinceMs(raw int64) int64 {
if raw > 0 {
return raw
@@ -126,8 +132,82 @@ func i18nText(english bool, zh string, en string) string {
return zh
}
func (h *NotificationHandler) loadPendingHITLItems(limit int, english bool) ([]NotificationSummaryItem, error) {
rows, err := h.db.Query(`
func notificationAccessFromContext(c *gin.Context) database.RBACListAccess {
session, ok := security.CurrentSession(c)
if !ok {
return database.RBACListAccess{}
}
return database.RBACListAccess{UserID: session.UserID, Scope: session.Scope}
}
func appendConversationAccessSQL(query string, args []interface{}, column string, access database.RBACListAccess) (string, []interface{}) {
userID := strings.TrimSpace(access.UserID)
if access.Scope == database.RBACScopeAll {
return query, args
}
if userID == "" {
return query + ` AND 1=0`, args
}
query += ` AND ` + column + ` IS NOT NULL AND ` + column + ` <> '' AND (
EXISTS (SELECT 1 FROM conversations c WHERE c.id = ` + column + ` AND c.owner_user_id = ?)
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments ra
WHERE ra.user_id = ? AND ra.resource_type = 'conversation' AND ra.resource_id = ` + column + `
)
OR EXISTS (
SELECT 1 FROM conversations c
JOIN projects p ON p.id = c.project_id
WHERE c.id = ` + column + ` AND p.owner_user_id = ?
)
OR EXISTS (
SELECT 1 FROM conversations c
JOIN rbac_resource_assignments pra ON pra.resource_id = c.project_id
WHERE c.id = ` + column + ` AND pra.user_id = ? AND pra.resource_type = 'project'
)
)`
args = append(args, userID, userID, userID, userID)
return query, args
}
func appendVulnerabilityNotificationAccessSQL(query string, args []interface{}, access database.RBACListAccess) (string, []interface{}) {
userID := strings.TrimSpace(access.UserID)
if access.Scope == database.RBACScopeAll {
return query, args
}
if userID == "" {
return query + ` AND 1=0`, args
}
query += ` AND (
owner_user_id = ?
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments ra
WHERE ra.user_id = ? AND ra.resource_type = 'vulnerability' AND ra.resource_id = vulnerabilities.id
)
OR (
project_id IS NOT NULL AND project_id <> '' AND (
EXISTS (SELECT 1 FROM projects p WHERE p.id = vulnerabilities.project_id AND p.owner_user_id = ?)
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments pra
WHERE pra.user_id = ? AND pra.resource_type = 'project' AND pra.resource_id = vulnerabilities.project_id
)
)
)
OR (
conversation_id IS NOT NULL AND conversation_id <> '' AND (
EXISTS (SELECT 1 FROM conversations c WHERE c.id = vulnerabilities.conversation_id AND c.owner_user_id = ?)
OR EXISTS (
SELECT 1 FROM rbac_resource_assignments cra
WHERE cra.user_id = ? AND cra.resource_type = 'conversation' AND cra.resource_id = vulnerabilities.conversation_id
)
)
)
)`
args = append(args, userID, userID, userID, userID, userID, userID)
return query, args
}
func (h *NotificationHandler) loadPendingHITLItems(limit int, english bool, access database.RBACListAccess) ([]NotificationSummaryItem, error) {
query := `
SELECT
id,
conversation_id,
@@ -135,9 +215,14 @@ func (h *NotificationHandler) loadPendingHITLItems(limit int, english bool) ([]N
COALESCE(CAST(strftime('%s', created_at) AS INTEGER), 0)
FROM hitl_interrupts
WHERE status = 'pending'
ORDER BY created_at DESC
`
args := []interface{}{}
query, args = appendConversationAccessSQL(query, args, "conversation_id", access)
query += ` ORDER BY created_at DESC
LIMIT ?
`, limit)
`
args = append(args, limit)
rows, err := h.db.Query(query, args...)
if err != nil {
return nil, err
}
@@ -170,9 +255,9 @@ func (h *NotificationHandler) loadPendingHITLItems(limit int, english bool) ([]N
return items, nil
}
func (h *NotificationHandler) loadVulnerabilityItems(sinceMs int64, limit int, english bool) ([]NotificationSummaryItem, map[string]int, error) {
func (h *NotificationHandler) loadVulnerabilityItems(sinceMs int64, limit int, english bool, access database.RBACListAccess) ([]NotificationSummaryItem, map[string]int, error) {
sinceSec := normalizedSinceSec(sinceMs)
rows, err := h.db.Query(`
query := `
SELECT
id,
title,
@@ -181,9 +266,15 @@ func (h *NotificationHandler) loadVulnerabilityItems(sinceMs int64, limit int, e
COALESCE(CAST(strftime('%s', created_at) AS INTEGER), 0)
FROM vulnerabilities
WHERE CAST(strftime('%s', created_at) AS INTEGER) > ?
`
args := []interface{}{sinceSec}
query, args = appendVulnerabilityNotificationAccessSQL(query, args, access)
query += `
ORDER BY created_at DESC
LIMIT ?
`, sinceSec, limit)
`
args = append(args, limit)
rows, err := h.db.Query(query, args...)
if err != nil {
return nil, nil, err
}
@@ -241,29 +332,23 @@ func (h *NotificationHandler) loadVulnerabilityItems(sinceMs int64, limit int, e
}
// loadC2SessionOnlineEvents 新会话上线(c2_eventssession + critical,与 Manager.IngestCheckIn 一致)
func (h *NotificationHandler) loadC2SessionOnlineEvents(sinceMs int64, limit int, english bool) ([]NotificationSummaryItem, int, error) {
func (h *NotificationHandler) loadC2SessionOnlineEvents(sinceMs int64, limit int, english bool, access database.RBACListAccess) ([]NotificationSummaryItem, int, error) {
sinceSec := normalizedSinceSec(sinceMs)
rows, err := h.db.Query(`
SELECT id, message, COALESCE(session_id, ''),
COALESCE(CAST(strftime('%s', created_at) AS INTEGER), 0)
FROM c2_events
WHERE category = 'session' AND level = 'critical'
AND CAST(strftime('%s', created_at) AS INTEGER) > ?
ORDER BY created_at DESC
LIMIT ?
`, sinceSec, limit)
events, err := h.db.ListC2EventsForAccess(database.ListC2EventsFilter{
Category: "session",
Level: "critical",
Since: ptrTime(time.Unix(sinceSec, 0)),
Limit: limit,
}, access)
if err != nil {
return nil, 0, err
}
defer rows.Close()
items := make([]NotificationSummaryItem, 0, limit)
for rows.Next() {
var id, message, sessionID string
var createdSec int64
if err := rows.Scan(&id, &message, &sessionID, &createdSec); err != nil {
for _, e := range events {
if e == nil {
continue
}
desc := strings.TrimSpace(message)
desc := strings.TrimSpace(e.Message)
if len(desc) > 220 {
desc = desc[:200] + "…"
}
@@ -271,19 +356,19 @@ func (h *NotificationHandler) loadC2SessionOnlineEvents(sinceMs int64, limit int
desc = i18nText(english, "新会话已建立", "A new session was created")
}
items = append(items, NotificationSummaryItem{
ID: "c2evt:" + id,
ID: "c2evt:" + e.ID,
Level: "p0",
Type: "c2_session_online",
Title: i18nText(english, "C2 新会话上线", "C2 new session online"),
Desc: desc,
Ts: unixSecToRFC3339(createdSec),
Ts: e.CreatedAt.UTC().Format(time.RFC3339),
Count: 1,
Actionable: false,
Read: false,
SessionID: sessionID,
SessionID: e.SessionID,
})
}
return items, len(items), rows.Err()
return items, len(items), nil
}
func (h *NotificationHandler) loadFailedExecutionItems(sinceMs int64, limit int, english bool) ([]NotificationSummaryItem, int, error) {
@@ -331,7 +416,7 @@ func (h *NotificationHandler) loadFailedExecutionItems(sinceMs int64, limit int,
return items, count, nil
}
func (h *NotificationHandler) summarizeLongRunningTasks(threshold time.Duration, english bool) ([]NotificationSummaryItem, int) {
func (h *NotificationHandler) summarizeLongRunningTasks(threshold time.Duration, english bool, access database.RBACListAccess) ([]NotificationSummaryItem, int) {
if h.agentHandler == nil || h.agentHandler.tasks == nil {
return nil, 0
}
@@ -342,6 +427,9 @@ func (h *NotificationHandler) summarizeLongRunningTasks(threshold time.Duration,
if t == nil {
continue
}
if !h.notificationConversationAllowed(access, t.ConversationID) {
continue
}
if now.Sub(t.StartedAt) >= threshold {
items = append(items, NotificationSummaryItem{
ID: "task_long:" + t.ConversationID,
@@ -360,7 +448,7 @@ func (h *NotificationHandler) summarizeLongRunningTasks(threshold time.Duration,
return items, len(items)
}
func (h *NotificationHandler) summarizeCompletedTasksSince(sinceMs int64, limit int, english bool) ([]NotificationSummaryItem, int) {
func (h *NotificationHandler) summarizeCompletedTasksSince(sinceMs int64, limit int, english bool, access database.RBACListAccess) ([]NotificationSummaryItem, int) {
if h.agentHandler == nil || h.agentHandler.tasks == nil {
return nil, 0
}
@@ -371,6 +459,9 @@ func (h *NotificationHandler) summarizeCompletedTasksSince(sinceMs int64, limit
if t == nil {
continue
}
if !h.notificationConversationAllowed(access, t.ConversationID) {
continue
}
if t.CompletedAt.After(since) {
items = append(items, NotificationSummaryItem{
ID: "task_completed:" + t.ConversationID + ":" + strconv.FormatInt(t.CompletedAt.Unix(), 10),
@@ -403,14 +494,16 @@ func buildPlaceholders(n int) string {
return strings.Join(out, ",")
}
func (h *NotificationHandler) readStatesByIDs(ids []string) (map[string]bool, error) {
func (h *NotificationHandler) readStatesByIDs(userID string, ids []string) (map[string]bool, error) {
result := make(map[string]bool, len(ids))
if len(ids) == 0 {
userID = strings.TrimSpace(userID)
if len(ids) == 0 || userID == "" {
return result, nil
}
holders := buildPlaceholders(len(ids))
query := "SELECT event_id FROM notification_reads WHERE event_id IN (" + holders + ")"
args := make([]interface{}, 0, len(ids))
query := "SELECT event_id FROM notification_reads_by_user WHERE user_id = ? AND event_id IN (" + holders + ")"
args := make([]interface{}, 0, len(ids)+1)
args = append(args, userID)
for _, id := range ids {
args = append(args, id)
}
@@ -429,7 +522,7 @@ func (h *NotificationHandler) readStatesByIDs(ids []string) (map[string]bool, er
return result, nil
}
func (h *NotificationHandler) applyReadStates(items []NotificationSummaryItem) ([]NotificationSummaryItem, error) {
func (h *NotificationHandler) applyReadStates(userID string, items []NotificationSummaryItem) ([]NotificationSummaryItem, error) {
markableIDs := make([]string, 0, len(items))
for _, item := range items {
if item.Actionable {
@@ -437,7 +530,7 @@ func (h *NotificationHandler) applyReadStates(items []NotificationSummaryItem) (
}
markableIDs = append(markableIDs, item.ID)
}
readMap, err := h.readStatesByIDs(markableIDs)
readMap, err := h.readStatesByIDs(userID, markableIDs)
if err != nil {
return items, err
}
@@ -494,34 +587,38 @@ func createNotificationReadTableIfNeeded(db *database.DB) error {
return fmt.Errorf("db is nil")
}
_, err := db.Exec(`
CREATE TABLE IF NOT EXISTS notification_reads (
event_id TEXT PRIMARY KEY,
read_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP
CREATE TABLE IF NOT EXISTS notification_reads_by_user (
user_id TEXT NOT NULL,
event_id TEXT NOT NULL,
read_at DATETIME NOT NULL DEFAULT CURRENT_TIMESTAMP,
PRIMARY KEY(user_id, event_id)
);
`)
if err != nil {
return err
}
_, idxErr := db.Exec(`CREATE INDEX IF NOT EXISTS idx_notification_reads_read_at ON notification_reads(read_at DESC);`)
_, idxErr := db.Exec(`CREATE INDEX IF NOT EXISTS idx_notification_reads_user_read_at ON notification_reads_by_user(user_id, read_at DESC);`)
return idxErr
}
func pruneNotificationReads(db *database.DB, maxRows int) error {
func pruneNotificationReads(db *database.DB, userID string, maxRows int) error {
if db == nil {
return fmt.Errorf("db is nil")
}
if maxRows <= 0 {
userID = strings.TrimSpace(userID)
if maxRows <= 0 || userID == "" {
return nil
}
_, err := db.Exec(`
DELETE FROM notification_reads
WHERE event_id NOT IN (
DELETE FROM notification_reads_by_user
WHERE user_id = ? AND event_id NOT IN (
SELECT event_id
FROM notification_reads
FROM notification_reads_by_user
WHERE user_id = ?
ORDER BY read_at DESC, rowid DESC
LIMIT ?
)
`, maxRows)
`, userID, userID, maxRows)
return err
}
@@ -564,6 +661,11 @@ func (h *NotificationHandler) MarkRead(c *gin.Context) {
c.JSON(http.StatusOK, gin.H{"ok": true, "marked": 0})
return
}
session, ok := security.CurrentSession(c)
if !ok || strings.TrimSpace(session.UserID) == "" {
c.JSON(http.StatusUnauthorized, gin.H{"error": "missing authenticated user"})
return
}
tx, err := h.db.Begin()
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to begin transaction"})
@@ -573,9 +675,9 @@ func (h *NotificationHandler) MarkRead(c *gin.Context) {
_ = tx.Rollback()
}()
stmt, err := tx.Prepare(`
INSERT INTO notification_reads(event_id, read_at)
VALUES(?, CURRENT_TIMESTAMP)
ON CONFLICT(event_id) DO UPDATE SET read_at = CURRENT_TIMESTAMP
INSERT INTO notification_reads_by_user(user_id, event_id, read_at)
VALUES(?, ?, CURRENT_TIMESTAMP)
ON CONFLICT(user_id, event_id) DO UPDATE SET read_at = CURRENT_TIMESTAMP
`)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to prepare statement"})
@@ -588,7 +690,7 @@ func (h *NotificationHandler) MarkRead(c *gin.Context) {
if !ok {
continue
}
if _, err := stmt.Exec(id); err != nil {
if _, err := stmt.Exec(session.UserID, id); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to mark read"})
return
}
@@ -598,7 +700,7 @@ func (h *NotificationHandler) MarkRead(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to commit read marks"})
return
}
if err := pruneNotificationReads(h.db, notificationReadMaxRows); err != nil {
if err := pruneNotificationReads(h.db, session.UserID, notificationReadMaxRows); err != nil {
h.logger.Warn("裁剪通知已读记录失败", zap.Error(err))
}
c.JSON(http.StatusOK, gin.H{"ok": true, "marked": marked})
@@ -626,30 +728,57 @@ func (h *NotificationHandler) GetSummary(c *gin.Context) {
if limit > 200 {
limit = 200
}
access := notificationAccessFromContext(c)
hitlItems, err := h.loadPendingHITLItems(limit, english)
if err != nil {
h.logger.Warn("加载 HITL 通知失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to summarize hitl notifications"})
return
hitlItems := []NotificationSummaryItem{}
if security.SessionHasPermission(c, "hitl:read") {
var err error
hitlItems, err = h.loadPendingHITLItems(limit, english, access)
if err != nil {
h.logger.Warn("加载 HITL 通知失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to summarize hitl notifications"})
return
}
}
vulnItems, vulnCounts, err := h.loadVulnerabilityItems(sinceMs, limit, english)
if err != nil {
h.logger.Warn("加载漏洞通知失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to summarize vulnerabilities"})
return
vulnItems := []NotificationSummaryItem{}
vulnCounts := map[string]int{
"newCriticalVulns": 0,
"newHighVulns": 0,
"newMediumVulns": 0,
"newLowVulns": 0,
"newInfoVulns": 0,
}
if security.SessionHasPermission(c, "vulnerability:read") {
var err error
vulnItems, vulnCounts, err = h.loadVulnerabilityItems(sinceMs, limit, english, access)
if err != nil {
h.logger.Warn("加载漏洞通知失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to summarize vulnerabilities"})
return
}
}
c2OnlineItems, c2OnlineCount, err := h.loadC2SessionOnlineEvents(sinceMs, limit, english)
if err != nil {
h.logger.Warn("加载 C2 会话上线通知失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to summarize c2 session events"})
return
c2OnlineItems := []NotificationSummaryItem{}
c2OnlineCount := 0
if security.SessionHasPermission(c, "c2:read") {
var err error
c2OnlineItems, c2OnlineCount, err = h.loadC2SessionOnlineEvents(sinceMs, limit, english, access)
if err != nil {
h.logger.Warn("加载 C2 会话上线通知失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to summarize c2 session events"})
return
}
}
longRunningItems, longRunningCount := h.summarizeLongRunningTasks(15*time.Minute, english)
completedItems, completedCount := h.summarizeCompletedTasksSince(sinceMs, limit, english)
longRunningItems := []NotificationSummaryItem{}
completedItems := []NotificationSummaryItem{}
longRunningCount := 0
completedCount := 0
if security.SessionHasPermission(c, "tasks:read") || security.SessionHasPermission(c, "chat:read") {
longRunningItems, longRunningCount = h.summarizeLongRunningTasks(15*time.Minute, english, access)
completedItems, completedCount = h.summarizeCompletedTasksSince(sinceMs, limit, english, access)
}
items := make([]NotificationSummaryItem, 0, len(hitlItems)+len(vulnItems)+len(c2OnlineItems)+len(longRunningItems)+len(completedItems))
items = append(items, hitlItems...)
@@ -658,7 +787,8 @@ func (h *NotificationHandler) GetSummary(c *gin.Context) {
items = append(items, longRunningItems...)
items = append(items, completedItems...)
items, err = h.applyReadStates(items)
session, _ := security.CurrentSession(c)
items, err := h.applyReadStates(session.UserID, items)
if err != nil {
h.logger.Warn("加载通知已读状态失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": "failed to load notification read states"})
@@ -697,3 +827,11 @@ func (h *NotificationHandler) GetSummary(c *gin.Context) {
Items: items,
})
}
func (h *NotificationHandler) notificationConversationAllowed(access database.RBACListAccess, conversationID string) bool {
conversationID = strings.TrimSpace(conversationID)
if conversationID == "" {
return access.Scope == database.RBACScopeAll
}
return h.db.UserCanAccessResource(access.UserID, access.Scope, "conversation", conversationID)
}
@@ -0,0 +1,45 @@
package handler
import (
"bytes"
"net/http"
"net/http/httptest"
"path/filepath"
"testing"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
func TestNotificationReadStateIsPerUser(t *testing.T) {
db, err := database.NewDB(filepath.Join(t.TempDir(), "notification-rbac.db"), zap.NewNop())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
h := NewNotificationHandler(db, nil, zap.NewNop())
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set(security.ContextSessionKey, security.Session{UserID: "u1", Scope: database.RBACScopeAssigned})
c.Next()
})
router.POST("/notifications/read", h.MarkRead)
w := httptest.NewRecorder()
req := httptest.NewRequest(http.MethodPost, "/notifications/read", bytes.NewBufferString(`{"eventIds":["vuln:v1"]}`))
req.Header.Set("Content-Type", "application/json")
router.ServeHTTP(w, req)
if w.Code != http.StatusOK {
t.Fatalf("mark read status = %d: %s", w.Code, w.Body.String())
}
u1, err := h.readStatesByIDs("u1", []string{"vuln:v1"})
if err != nil || !u1["vuln:v1"] {
t.Fatalf("u1 read state = %#v, err=%v", u1, err)
}
u2, err := h.readStatesByIDs("u2", []string{"vuln:v1"})
if err != nil || u2["vuln:v1"] {
t.Fatalf("u2 inherited u1 read state = %#v, err=%v", u2, err)
}
}
+37 -1
View File
@@ -4735,7 +4735,7 @@ func (h *OpenAPIHandler) GetOpenAPISpec(c *gin.Context) {
"get": map[string]interface{}{
"tags": []string{"对话交互"},
"summary": "获取消息过程详情",
"description": "按需加载指定消息的执行过程详情,包括工具调用、思考过程等事件。",
"description": "按需分页加载指定消息的执行过程详情,包括工具调用、思考过程等事件。默认返回 50 条;导出或旧集成需要全量时可显式传 full=1。",
"operationId": "getMessageProcessDetails",
"parameters": []map[string]interface{}{
{
@@ -4745,6 +4745,34 @@ func (h *OpenAPIHandler) GetOpenAPISpec(c *gin.Context) {
"description": "消息ID",
"schema": map[string]interface{}{"type": "string"},
},
{
"name": "summary",
"in": "query",
"required": false,
"description": "仅返回过程详情摘要(total / iterationCount / maxIteration",
"schema": map[string]interface{}{"type": "boolean"},
},
{
"name": "limit",
"in": "query",
"required": false,
"description": "分页大小,默认 50,最大 500",
"schema": map[string]interface{}{"type": "integer", "default": 50, "maximum": 500},
},
{
"name": "offset",
"in": "query",
"required": false,
"description": "分页偏移量,默认 0",
"schema": map[string]interface{}{"type": "integer", "default": 0},
},
{
"name": "full",
"in": "query",
"required": false,
"description": "显式返回全量过程详情;仅建议导出/兼容旧集成使用",
"schema": map[string]interface{}{"type": "boolean"},
},
},
"responses": map[string]interface{}{
"200": map[string]interface{}{
@@ -4769,6 +4797,14 @@ func (h *OpenAPIHandler) GetOpenAPISpec(c *gin.Context) {
},
},
},
"total": map[string]interface{}{"type": "integer", "description": "过程详情总数"},
"offset": map[string]interface{}{"type": "integer", "description": "当前分页偏移量"},
"limit": map[string]interface{}{"type": "integer", "description": "当前分页大小"},
"hasMore": map[string]interface{}{"type": "boolean", "description": "是否还有更多过程详情"},
"summary": map[string]interface{}{
"type": "object",
"description": "summary=1 时返回的摘要",
},
},
},
},
+16 -3
View File
@@ -9,6 +9,7 @@ import (
"cyberstrike-ai/internal/attackchain"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/project"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
@@ -70,6 +71,10 @@ func (h *ProjectHandler) CreateProject(c *gin.Context) {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if session, ok := security.CurrentSession(c); ok {
_ = h.db.SetResourceOwner("project", created.ID, session.UserID)
_ = h.db.AssignResourceToUser(session.UserID, "project", created.ID)
}
c.JSON(http.StatusOK, created)
}
@@ -82,7 +87,8 @@ func (h *ProjectHandler) GetDashboardSummary(c *gin.Context) {
if limit > 50 {
limit = 50
}
summary, err := h.db.GetProjectDashboardSummary(limit)
session, _ := security.CurrentSession(c)
summary, err := h.db.GetProjectDashboardSummaryForAccess(limit, session.UserID, session.Scope)
if err != nil {
h.logger.Error("获取项目仪表盘摘要失败", zap.Error(err))
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
@@ -106,7 +112,8 @@ func (h *ProjectHandler) ListProjects(c *gin.Context) {
if limit > 500 {
limit = 500
}
list, err := h.db.ListProjects(status, search, limit, offset)
session, _ := security.CurrentSession(c)
list, err := h.db.ListProjectsForAccess(status, search, limit, offset, session.UserID, session.Scope)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -114,7 +121,7 @@ func (h *ProjectHandler) ListProjects(c *gin.Context) {
if list == nil {
list = []*database.Project{}
}
total, err := h.db.CountProjects(status, search)
total, err := h.db.CountProjectsForAccess(status, search, session.UserID, session.Scope)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
@@ -635,6 +642,12 @@ func (h *ProjectHandler) DeleteFactEdge(c *gin.Context) {
func (h *ProjectHandler) PromoteAttackChain(c *gin.Context) {
projectID := c.Param("id")
conversationID := c.Param("conversationId")
session, ok := security.CurrentSession(c)
if !ok || !h.db.UserCanAccessResource(session.UserID, session.Scope, "project", projectID) ||
!h.db.UserCanAccessResource(session.UserID, session.Scope, "conversation", conversationID) {
c.JSON(http.StatusForbidden, gin.H{"error": "无权访问目标项目或来源对话"})
return
}
result, err := attackchain.PromoteToProject(h.db, projectID, conversationID)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
+429
View File
@@ -0,0 +1,429 @@
package handler
import (
"fmt"
"net/http"
"strconv"
"strings"
"cyberstrike-ai/internal/audit"
"cyberstrike-ai/internal/authctx"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
type RBACHandler struct {
db *database.DB
logger *zap.Logger
audit *audit.Service
auth *security.AuthManager
}
func NewRBACHandler(db *database.DB, logger *zap.Logger) *RBACHandler {
return &RBACHandler{db: db, logger: logger}
}
func (h *RBACHandler) SetAudit(s *audit.Service) {
h.audit = s
}
func (h *RBACHandler) SetAuthManager(m *security.AuthManager) {
h.auth = m
}
func (h *RBACHandler) Me(c *gin.Context) {
session, _ := security.CurrentSession(c)
resolvedScope := session.Scope
permissionScopes := session.PermissionScopes
if principal, ok := authctx.PrincipalFromContext(c.Request.Context()); ok {
resolvedScope = principal.Scope
permissionScopes = principal.PermissionScopes
}
c.JSON(http.StatusOK, gin.H{
"user": gin.H{
"id": session.UserID,
"username": session.Username,
"display_name": session.DisplayName,
},
"roles": session.Roles,
"permissions": permissionKeys(session.Permissions),
"scope": resolvedScope,
"permission_scopes": permissionScopes,
})
}
func (h *RBACHandler) Metadata(c *gin.Context) {
roles, err := h.db.ListRBACRoles()
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
rolePermissions := map[string][]string{}
for _, role := range roles {
keys, _ := h.db.ListRBACRolePermissionKeys(role.ID)
rolePermissions[role.ID] = keys
}
c.JSON(http.StatusOK, gin.H{
"permissions": security.PermissionCatalog,
"roles": roles,
"role_permissions": rolePermissions,
"scopes": []string{database.RBACScopeAll, database.RBACScopeAssigned, database.RBACScopeOwn},
})
}
func (h *RBACHandler) ListRoles(c *gin.Context) {
roles, err := h.db.ListRBACRoles()
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
out := make([]gin.H, 0, len(roles))
for _, role := range roles {
keys, _ := h.db.ListRBACRolePermissionKeys(role.ID)
out = append(out, gin.H{
"id": role.ID,
"name": role.Name,
"description": role.Description,
"scope": role.Scope,
"is_system": role.IsSystem,
"permissions": keys,
"created_at": role.CreatedAt,
"updated_at": role.UpdatedAt,
})
}
c.JSON(http.StatusOK, gin.H{"roles": out})
}
type upsertRBACRoleRequest struct {
ID string `json:"id"`
Name string `json:"name" binding:"required"`
Description string `json:"description"`
Scope string `json:"scope"`
Permissions []string `json:"permissions"`
}
func validateRBACPermissionKeys(keys []string) error {
for _, key := range keys {
key = strings.TrimSpace(key)
if key == "" {
continue
}
if _, ok := security.PermissionCatalog[key]; !ok {
return fmt.Errorf("未知权限: %s", key)
}
}
return nil
}
func (h *RBACHandler) CreateRole(c *gin.Context) {
var req upsertRBACRoleRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if err := validateRBACPermissionKeys(req.Permissions); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
role, err := h.db.UpsertRBACRole("", req.Name, req.Description, req.Scope, req.Permissions)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if h.audit != nil {
h.audit.RecordOK(c, "rbac", "create_role", "创建平台角色", "role", role.ID, nil)
}
c.JSON(http.StatusOK, gin.H{"role": role})
}
func (h *RBACHandler) UpdateRole(c *gin.Context) {
id := strings.TrimSpace(c.Param("id"))
existing, err := h.db.GetRBACRoleByID(id)
if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "角色不存在"})
return
}
if existing.IsSystem {
c.JSON(http.StatusBadRequest, gin.H{"error": "系统内置角色不可修改,请创建自定义角色"})
return
}
var req upsertRBACRoleRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if err := validateRBACPermissionKeys(req.Permissions); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
role, err := h.db.UpsertRBACRole(id, req.Name, req.Description, req.Scope, req.Permissions)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if existing.IsSystem {
role.IsSystem = true
}
if h.audit != nil {
h.audit.RecordOK(c, "rbac", "update_role", "更新平台角色", "role", id, nil)
}
if h.auth != nil {
h.auth.RevokeAllSessions()
}
c.JSON(http.StatusOK, gin.H{"role": role})
}
func (h *RBACHandler) DeleteRole(c *gin.Context) {
id := strings.TrimSpace(c.Param("id"))
if err := h.db.DeleteRBACRole(id); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if h.audit != nil {
h.audit.RecordOK(c, "rbac", "delete_role", "删除平台角色", "role", id, nil)
}
if h.auth != nil {
h.auth.RevokeAllSessions()
}
c.JSON(http.StatusOK, gin.H{"success": true})
}
func (h *RBACHandler) ListUsers(c *gin.Context) {
users, err := h.db.ListRBACUsers()
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
out := make([]gin.H, 0, len(users))
for _, user := range users {
roleIDs, _ := h.db.ListRBACUserRoleIDs(user.ID)
out = append(out, gin.H{
"id": user.ID,
"username": user.Username,
"display_name": user.DisplayName,
"enabled": user.Enabled,
"is_builtin": user.IsBuiltin,
"roles": roleIDs,
"created_at": user.CreatedAt,
"updated_at": user.UpdatedAt,
})
}
c.JSON(http.StatusOK, gin.H{"users": out})
}
type createRBACUserRequest struct {
Username string `json:"username" binding:"required"`
DisplayName string `json:"display_name"`
Password string `json:"password" binding:"required"`
Enabled *bool `json:"enabled"`
Roles []string `json:"roles"`
}
func (h *RBACHandler) CreateUser(c *gin.Context) {
var req createRBACUserRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if len(strings.TrimSpace(req.Password)) < 8 {
c.JSON(http.StatusBadRequest, gin.H{"error": "密码长度至少需要 8 位"})
return
}
hash, err := security.HashPassword(req.Password)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
enabled := true
if req.Enabled != nil {
enabled = *req.Enabled
}
user, err := h.db.CreateRBACUser(req.Username, req.DisplayName, hash, enabled, req.Roles)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if h.audit != nil {
h.audit.RecordOK(c, "rbac", "create_user", "创建平台用户", "user", user.ID, map[string]interface{}{"username": user.Username})
}
c.JSON(http.StatusOK, gin.H{"user": user})
}
type updateRBACUserRequest struct {
DisplayName *string `json:"display_name"`
Password *string `json:"password"`
Enabled *bool `json:"enabled"`
Roles *[]string `json:"roles"`
}
func (h *RBACHandler) UpdateUser(c *gin.Context) {
id := strings.TrimSpace(c.Param("id"))
user, err := h.db.GetRBACUserByID(id)
if err != nil {
c.JSON(http.StatusNotFound, gin.H{"error": "用户不存在"})
return
}
var req updateRBACUserRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
displayName := user.DisplayName
if req.DisplayName != nil {
displayName = *req.DisplayName
}
if err := h.db.UpdateRBACUser(id, displayName, req.Enabled, req.Roles); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if req.Password != nil && strings.TrimSpace(*req.Password) != "" {
if len(strings.TrimSpace(*req.Password)) < 8 {
c.JSON(http.StatusBadRequest, gin.H{"error": "密码长度至少需要 8 位"})
return
}
hash, err := security.HashPassword(*req.Password)
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
if err := h.db.UpdateRBACUserPassword(id, hash); err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
}
if h.audit != nil {
h.audit.RecordOK(c, "rbac", "update_user", "更新平台用户", "user", id, nil)
}
if h.auth != nil {
h.auth.RevokeUserSessions(id)
}
updated, _ := h.db.GetRBACUserByID(id)
c.JSON(http.StatusOK, gin.H{"user": updated})
}
func (h *RBACHandler) DeleteUser(c *gin.Context) {
id := strings.TrimSpace(c.Param("id"))
if err := h.db.DeleteRBACUser(id); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if h.audit != nil {
h.audit.RecordOK(c, "rbac", "delete_user", "删除平台用户", "user", id, nil)
}
if h.auth != nil {
h.auth.RevokeUserSessions(id)
}
c.JSON(http.StatusOK, gin.H{"success": true})
}
type assignResourceRequest struct {
UserID string `json:"user_id" binding:"required"`
ResourceType string `json:"resource_type" binding:"required"`
ResourceID string `json:"resource_id"`
ResourceIDs []string `json:"resource_ids"`
AutoDetect bool `json:"auto_detect"`
}
func (h *RBACHandler) AssignResource(c *gin.Context) {
var req assignResourceRequest
if err := c.ShouldBindJSON(&req); err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
resourceIDs := append([]string(nil), req.ResourceIDs...)
if strings.TrimSpace(req.ResourceID) != "" {
resourceIDs = append(resourceIDs, req.ResourceID)
}
if len(resourceIDs) == 0 {
c.JSON(http.StatusBadRequest, gin.H{"error": "至少需要一个资源 ID"})
return
}
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 {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if h.audit != nil {
for _, resourceID := range resourceIDs {
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{
"success": true,
"requested": len(resourceIDs),
"created": created,
"skipped": int64(len(resourceIDs)) - created,
"detected_types": detectedTypes,
})
}
func (h *RBACHandler) ListResourceAssignments(c *gin.Context) {
rows, err := h.db.ListRBACResourceAssignments(c.Query("user_id"))
if err != nil {
c.JSON(http.StatusInternalServerError, gin.H{"error": err.Error()})
return
}
c.JSON(http.StatusOK, gin.H{"assignments": rows})
}
func (h *RBACHandler) ListAssignableResources(c *gin.Context) {
limit, _ := strconv.Atoi(c.DefaultQuery("limit", "50"))
if limit <= 0 || limit > 50 {
limit = 50
}
offset, _ := strconv.Atoi(c.DefaultQuery("offset", "0"))
if offset < 0 {
offset = 0
}
resources, err := h.db.ListAssignableRBACResourcesPage(c.Query("type"), c.Query("q"), limit+1, offset)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
hasMore := len(resources) > limit
if hasMore {
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{
"resources": resources,
"has_more": hasMore,
"limit": limit,
"offset": offset,
"total": total,
})
}
func (h *RBACHandler) DeleteResourceAssignment(c *gin.Context) {
id := strings.TrimSpace(c.Param("id"))
assignment, err := h.db.DeleteRBACResourceAssignmentWithDetails(id)
if err != nil {
c.JSON(http.StatusBadRequest, gin.H{"error": err.Error()})
return
}
if h.audit != 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})
}
+206
View File
@@ -0,0 +1,206 @@
package handler
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"path/filepath"
"testing"
"time"
"cyberstrike-ai/internal/authctx"
"cyberstrike-ai/internal/database"
"cyberstrike-ai/internal/mcp"
"cyberstrike-ai/internal/security"
"github.com/gin-gonic/gin"
"go.uber.org/zap"
)
func TestDetachedAgentContextRetainsPrincipalWithoutParentCancellation(t *testing.T) {
parent, cancel := context.WithCancel(context.Background())
parent = authctx.WithPrincipal(parent, authctx.NewPrincipal("u1", "user", database.RBACScopeAssigned, map[string]bool{"agent:execute": true}))
detached := detachedAgentContext(parent)
cancel()
if err := detached.Err(); err != nil {
t.Fatalf("detached context inherited cancellation: %v", err)
}
principal, ok := authctx.PrincipalFromContext(detached)
if !ok || principal.UserID != "u1" || !principal.HasPermission("agent:execute") {
t.Fatalf("detached context lost principal: %#v, ok=%v", principal, ok)
}
}
func TestPromoteAttackChainRequiresSourceConversationAccess(t *testing.T) {
db, err := database.NewDB(filepath.Join(t.TempDir(), "promote-rbac.db"), zap.NewNop())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
project, _ := db.CreateProject(&database.Project{Name: "owned"})
conversation, _ := db.CreateConversation("foreign", database.ConversationCreateMeta{})
_ = db.SetResourceOwner("project", project.ID, "u1")
_ = db.SetResourceOwner("conversation", conversation.ID, "u2")
h := NewProjectHandler(db, zap.NewNop())
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set(security.ContextSessionKey, security.Session{UserID: "u1", Scope: database.RBACScopeOwn})
c.Next()
})
router.POST("/api/projects/:id/promote-attack-chain/:conversationId", h.PromoteAttackChain)
w := httptest.NewRecorder()
router.ServeHTTP(w, httptest.NewRequest(http.MethodPost, "/api/projects/"+project.ID+"/promote-attack-chain/"+conversation.ID, nil))
if w.Code != http.StatusForbidden {
t.Fatalf("status = %d, want 403: %s", w.Code, w.Body.String())
}
}
func TestVulnerabilityCannotBeReparentedToForeignProject(t *testing.T) {
db, err := database.NewDB(filepath.Join(t.TempDir(), "vuln-reparent-rbac.db"), zap.NewNop())
if err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = db.Close() })
owned, _ := db.CreateProject(&database.Project{Name: "owned"})
foreign, _ := db.CreateProject(&database.Project{Name: "foreign"})
_ = db.SetResourceOwner("project", owned.ID, "u1")
_ = db.SetResourceOwner("project", foreign.ID, "u2")
vulnerability, err := db.CreateVulnerability(&database.Vulnerability{Title: "v", Severity: "high", ProjectID: owned.ID})
if err != nil {
t.Fatal(err)
}
_ = db.SetResourceOwner("vulnerability", vulnerability.ID, "u1")
h := NewVulnerabilityHandler(db, zap.NewNop())
router := gin.New()
router.Use(func(c *gin.Context) {
c.Set(security.ContextSessionKey, security.Session{UserID: "u1", Scope: database.RBACScopeOwn})
c.Next()
})
router.PUT("/api/vulnerabilities/:id", h.UpdateVulnerability)
body, _ := json.Marshal(map[string]interface{}{"project_id": foreign.ID})
w := httptest.NewRecorder()
router.ServeHTTP(w, httptest.NewRequest(http.MethodPut, "/api/vulnerabilities/"+vulnerability.ID, bytes.NewReader(body)))
if w.Code != http.StatusForbidden {
t.Fatalf("status = %d, want 403: %s", w.Code, w.Body.String())
}
}
func TestAgentTaskEndpointsFilterAndRejectForeignConversations(t *testing.T) {
gin.SetMode(gin.TestMode)
db, user := setupConversationRBACTest(t)
allowed, _ := db.CreateConversation("allowed", database.ConversationCreateMeta{})
hidden, _ := db.CreateConversation("hidden", database.ConversationCreateMeta{})
if err := db.AssignResourceToUser(user.ID, "conversation", allowed.ID); err != nil {
t.Fatal(err)
}
tasks := NewAgentTaskManager()
if _, err := tasks.StartTask(allowed.ID, "visible", func(error) {}); err != nil {
t.Fatal(err)
}
if _, err := tasks.StartTask(hidden.ID, "secret", func(error) {}); err != nil {
t.Fatal(err)
}
h := &AgentHandler{db: db, tasks: tasks, logger: zap.NewNop()}
w := performAssignedHandler(user, http.MethodGet, "/api/agent-loop/tasks", nil, h.ListAgentTasks)
if w.Code != http.StatusOK {
t.Fatalf("list status = %d: %s", w.Code, w.Body.String())
}
var response struct {
Tasks []*AgentTask `json:"tasks"`
}
if err := json.Unmarshal(w.Body.Bytes(), &response); err != nil {
t.Fatal(err)
}
if len(response.Tasks) != 1 || response.Tasks[0].ConversationID != allowed.ID {
t.Fatalf("tasks = %#v, want only %s", response.Tasks, allowed.ID)
}
w = performAssignedHandler(user, http.MethodPost, "/api/agent-loop/cancel", map[string]string{"conversationId": hidden.ID}, h.CancelAgentLoop)
if w.Code != http.StatusForbidden {
t.Fatalf("cancel status = %d, want %d: %s", w.Code, http.StatusForbidden, w.Body.String())
}
}
func TestChatUploadPathAuthorizationFollowsConversationAccess(t *testing.T) {
db, user := setupConversationRBACTest(t)
allowed, _ := db.CreateConversation("allowed", database.ConversationCreateMeta{})
hidden, _ := db.CreateConversation("hidden", database.ConversationCreateMeta{})
if err := db.AssignResourceToUser(user.ID, "conversation", allowed.ID); err != nil {
t.Fatal(err)
}
h := NewChatUploadsHandler(zap.NewNop(), db)
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Set(security.ContextSessionKey, security.Session{UserID: user.ID, Scope: database.RBACScopeAssigned, Permissions: map[string]bool{"chat:write": true}})
if !h.pathAllowed(c, filepath.ToSlash(filepath.Join("2026-07-10", allowed.ID, "a.txt"))) {
t.Fatal("assigned conversation attachment should be accessible")
}
if h.pathAllowed(c, filepath.ToSlash(filepath.Join("2026-07-10", hidden.ID, "secret.txt"))) {
t.Fatal("foreign conversation attachment should be denied")
}
if h.pathAllowed(c, "2026-07-10/_manual/secret.txt") {
t.Fatal("unowned manual attachment should fail closed")
}
}
func TestPrepareMultiAgentSessionRejectsForeignConversation(t *testing.T) {
db, user := setupConversationRBACTest(t)
hidden, _ := db.CreateConversation("hidden", database.ConversationCreateMeta{})
h := &AgentHandler{db: db, logger: zap.NewNop()}
c, _ := gin.CreateTestContext(httptest.NewRecorder())
c.Set(security.ContextSessionKey, security.Session{UserID: user.ID, Scope: database.RBACScopeAssigned, Permissions: map[string]bool{"chat:write": true}})
_, err := h.prepareMultiAgentSession(&ChatRequest{ConversationID: hidden.ID, Message: "write"}, c, "test")
if err == nil || err.Error() != "无权访问该对话" {
t.Fatalf("err = %v, want unauthorized conversation", err)
}
}
func TestMonitorExecutionDetailRejectsForeignOwner(t *testing.T) {
db, user := setupConversationRBACTest(t)
for _, exec := range []*mcp.ToolExecution{
{ID: "exec-allowed", ToolName: "allowed", Status: "completed", StartTime: time.Now(), OwnerUserID: user.ID},
{ID: "exec-hidden", ToolName: "hidden", Status: "completed", StartTime: time.Now(), OwnerUserID: "another-user"},
} {
if err := db.SaveToolExecution(exec); err != nil {
t.Fatal(err)
}
}
h := NewMonitorHandler(mcp.NewServerWithStorage(zap.NewNop(), db), nil, db, zap.NewNop())
request := func(id string) *httptest.ResponseRecorder {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
c.Request = httptest.NewRequest(http.MethodGet, "/api/monitor/execution/"+id, nil)
c.Params = gin.Params{{Key: "id", Value: id}}
c.Set(security.ContextSessionKey, security.Session{UserID: user.ID, Scope: database.RBACScopeAssigned})
h.GetExecution(c)
return w
}
if w := request("exec-hidden"); w.Code != http.StatusForbidden {
t.Fatalf("hidden status = %d, want %d: %s", w.Code, http.StatusForbidden, w.Body.String())
}
if w := request("exec-allowed"); w.Code != http.StatusOK {
t.Fatalf("allowed status = %d, want %d: %s", w.Code, http.StatusOK, w.Body.String())
}
}
func performAssignedHandler(user *database.RBACUser, method, path string, body interface{}, handler gin.HandlerFunc) *httptest.ResponseRecorder {
w := httptest.NewRecorder()
c, _ := gin.CreateTestContext(w)
var req *http.Request
if body == nil {
req = httptest.NewRequest(method, path, nil)
} else {
payload, _ := json.Marshal(body)
req = httptest.NewRequest(method, path, bytes.NewReader(payload))
req.Header.Set("Content-Type", "application/json")
}
c.Request = req
c.Set(security.ContextSessionKey, security.Session{UserID: user.ID, Scope: database.RBACScopeAssigned})
handler(c)
return w
}

Some files were not shown because too many files have changed in this diff Show More