Add files via upload

This commit is contained in:
公明
2026-07-21 23:47:17 +08:00
committed by GitHub
parent d3176f048d
commit 0152781598
8 changed files with 378 additions and 137 deletions
+159 -46
View File
@@ -16,8 +16,10 @@ import (
"cyberstrike-ai/internal/config"
"cyberstrike-ai/internal/mcp"
"cyberstrike-ai/internal/tooloutput"
"github.com/creack/pty"
"github.com/google/uuid"
"go.uber.org/zap"
)
@@ -38,6 +40,7 @@ type Executor struct {
logger *zap.Logger
shellNoOutputTimeoutSec int // execute/exec 无新输出空闲秒数;0=默认 300-1=关闭(见 SetShellNoOutputTimeoutSeconds
toolOutputMaxBytes int
spillRootDir string
}
// NewExecutor 创建新的执行器
@@ -60,11 +63,34 @@ func (e *Executor) SetShellNoOutputTimeoutSeconds(sec int) {
// SetToolOutputMaxBytes limits stdout/stderr retained and streamed by exec-like
// tools. It should stay aligned with MCP result normalization so every channel
// sees the same bounded payload.
// sees the same bounded payload. Oversized full output is spilled to disk first.
func (e *Executor) SetToolOutputMaxBytes(maxBytes int) {
e.toolOutputMaxBytes = maxBytes
}
// SetToolOutputSpillRoot sets the reduction-compatible root for spilling full
// exec stdout/stderr when the in-memory bound is exceeded (empty → tmp/reduction).
func (e *Executor) SetToolOutputSpillRoot(rootDir string) {
e.spillRootDir = strings.TrimSpace(rootDir)
}
func (e *Executor) spillOptsFromContext(ctx context.Context) tooloutput.SpillOpts {
root := ""
if e != nil {
root = e.spillRootDir
}
opts := tooloutput.SpillOpts{RootDir: root}
if ctx != nil {
opts.ConversationID = mcp.MCPConversationIDFromContext(ctx)
opts.ProjectID = mcp.MCPProjectIDFromContext(ctx)
opts.ExecutionID = mcp.MCPExecutionIDFromContext(ctx)
}
if opts.ExecutionID == "" {
opts.ExecutionID = uuid.NewString()
}
return opts
}
// buildToolIndex 构建工具索引,将 O(n) 查找优化为 O(1)
func (e *Executor) buildToolIndex() {
e.toolIndex = make(map[string]*config.ToolConfig)
@@ -157,9 +183,10 @@ func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[st
var output string
var err error
spill := e.spillOptsFromContext(ctx)
// 如果上层提供了 stdout/stderr 增量回调,则边执行边读取并回调。
if cb, ok := ctx.Value(ToolOutputCallbackCtxKey).(ToolOutputCallback); ok && cb != nil {
output, err = streamCommandOutput(ctx, cmd, cb, ResolveShellNoOutputTimeoutSeconds(e.shellNoOutputTimeoutSec), e.toolOutputMaxBytes)
output, err = streamCommandOutput(ctx, cmd, cb, ResolveShellNoOutputTimeoutSeconds(e.shellNoOutputTimeoutSec), e.toolOutputMaxBytes, spill)
if err != nil && shouldRetryWithPTY(output) {
e.logger.Info("检测到工具需要 TTY,使用 PTY 重试",
zap.String("tool", toolName),
@@ -167,11 +194,11 @@ func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[st
cmd2 := exec.CommandContext(ctx, toolConfig.Command, cmdArgs...)
applyDefaultTerminalEnv(cmd2)
_ = prepareShellCmdSession(cmd2)
output, err = runCommandWithPTY(ctx, cmd2, cb, e.toolOutputMaxBytes)
output, err = runCommandWithPTY(ctx, cmd2, cb, e.toolOutputMaxBytes, spill)
}
} else {
// 非流式:内存缓冲 + ctx 取消杀进程组;行为对齐原 CombinedOutput,避免双流管道 fan-in 死锁。
output, err = combinedOutputCancellableWithLimit(ctx, cmd, e.toolOutputMaxBytes)
output, err = combinedOutputCancellableWithLimit(ctx, cmd, e.toolOutputMaxBytes, spill)
if err != nil && shouldRetryWithPTY(output) {
e.logger.Info("检测到工具需要 TTY,使用 PTY 重试",
zap.String("tool", toolName),
@@ -179,7 +206,7 @@ func (e *Executor) ExecuteTool(ctx context.Context, toolName string, args map[st
cmd2 := exec.CommandContext(ctx, toolConfig.Command, cmdArgs...)
applyDefaultTerminalEnv(cmd2)
_ = prepareShellCmdSession(cmd2)
output, err = runCommandWithPTY(ctx, cmd2, nil, e.toolOutputMaxBytes)
output, err = runCommandWithPTY(ctx, cmd2, nil, e.toolOutputMaxBytes, spill)
}
}
if err != nil {
@@ -920,9 +947,10 @@ func (e *Executor) executeSystemCommand(ctx context.Context, args map[string]int
// 非后台命令:等待输出
var output string
var err error
spill := e.spillOptsFromContext(ctx)
// 若上层提供工具输出增量回调,则边执行边流式读取。
if cb, ok := ctx.Value(ToolOutputCallbackCtxKey).(ToolOutputCallback); ok && cb != nil {
output, err = streamCommandOutput(ctx, cmd, cb, ResolveShellNoOutputTimeoutSeconds(e.shellNoOutputTimeoutSec), e.toolOutputMaxBytes)
output, err = streamCommandOutput(ctx, cmd, cb, ResolveShellNoOutputTimeoutSeconds(e.shellNoOutputTimeoutSec), e.toolOutputMaxBytes, spill)
if err != nil && shouldRetryWithPTY(output) {
e.logger.Info("检测到系统命令需要 TTY,使用 PTY 重试")
cmd2 := exec.CommandContext(ctx, shell, "-c", command)
@@ -930,10 +958,10 @@ func (e *Executor) executeSystemCommand(ctx context.Context, args map[string]int
cmd2.Dir = workDir
}
ConfigureShellCmdForAgentExecute(cmd2)
output, err = runCommandWithPTY(ctx, cmd2, cb, e.toolOutputMaxBytes)
output, err = runCommandWithPTY(ctx, cmd2, cb, e.toolOutputMaxBytes, spill)
}
} else {
output, err = combinedOutputCancellableWithLimit(ctx, cmd, e.toolOutputMaxBytes)
output, err = combinedOutputCancellableWithLimit(ctx, cmd, e.toolOutputMaxBytes, spill)
if err != nil && shouldRetryWithPTY(output) {
e.logger.Info("检测到系统命令需要 TTY,使用 PTY 重试")
cmd2 := exec.CommandContext(ctx, shell, "-c", command)
@@ -941,7 +969,7 @@ func (e *Executor) executeSystemCommand(ctx context.Context, args map[string]int
cmd2.Dir = workDir
}
ConfigureShellCmdForAgentExecute(cmd2)
output, err = runCommandWithPTY(ctx, cmd2, nil, e.toolOutputMaxBytes)
output, err = runCommandWithPTY(ctx, cmd2, nil, e.toolOutputMaxBytes, spill)
}
}
if err != nil {
@@ -982,12 +1010,17 @@ func (e *Executor) executeSystemCommand(ctx context.Context, args map[string]int
// 非流式路径不使用双流管道 fan-in,避免 stderr 撑满管道缓冲区时与 stdout 互相阻塞导致死锁。
// 无输出空闲检测由上层 agent.tool_timeout_minutes 兜底,不改变原 CombinedOutput 语义。
func combinedOutputCancellable(ctx context.Context, cmd *exec.Cmd) (string, error) {
return combinedOutputCancellableWithLimit(ctx, cmd, 0)
return combinedOutputCancellableWithLimit(ctx, cmd, 0, tooloutput.SpillOpts{})
}
func combinedOutputCancellableWithLimit(ctx context.Context, cmd *exec.Cmd, maxBytes int) (string, error) {
stdoutBuf := newBoundedOutputCollector(maxBytes)
stderrBuf := newBoundedOutputCollector(maxBytes)
func combinedOutputCancellableWithLimit(ctx context.Context, cmd *exec.Cmd, maxBytes int, spill tooloutput.SpillOpts) (string, error) {
var tee *tooloutput.Tee
if maxBytes > 0 {
tee = tooloutput.NewTee(spill)
defer func() { _ = tee.Close() }()
}
stdoutBuf := newBoundedOutputCollector(maxBytes, tee)
stderrBuf := newBoundedOutputCollector(maxBytes, tee)
cmd.Stdout = stdoutBuf
cmd.Stderr = stderrBuf
@@ -1016,9 +1049,9 @@ func combinedOutputCancellableWithLimit(ctx context.Context, cmd *exec.Cmd, maxB
case waitErr = <-done:
case <-ctx.Done():
waitErr = <-done
return limitOutputString(joinCommandOutput(stdoutBuf.String(), stderrBuf.String()), maxBytes), ctx.Err()
return finalizeJoinedBoundedOutputs(stdoutBuf, stderrBuf, maxBytes, tee), ctx.Err()
}
return limitOutputString(joinCommandOutput(stdoutBuf.String(), stderrBuf.String()), maxBytes), waitErr
return finalizeJoinedBoundedOutputs(stdoutBuf, stderrBuf, maxBytes, tee), waitErr
}
func joinCommandOutput(stdout, stderr string) string {
@@ -1036,10 +1069,11 @@ type boundedOutputCollector struct {
maxBytes int
seenBytes int
truncated bool
tee *tooloutput.Tee
}
func newBoundedOutputCollector(maxBytes int) *boundedOutputCollector {
return &boundedOutputCollector{maxBytes: maxBytes}
func newBoundedOutputCollector(maxBytes int, tee *tooloutput.Tee) *boundedOutputCollector {
return &boundedOutputCollector{maxBytes: maxBytes, tee: tee}
}
func (b *boundedOutputCollector) Write(p []byte) (int, error) {
@@ -1051,27 +1085,20 @@ func (b *boundedOutputCollector) WriteStringLimited(s string) string {
if b == nil {
return ""
}
if b.tee != nil {
_, _ = b.tee.Write([]byte(s))
}
if b.maxBytes <= 0 {
b.seenBytes += len(s)
b.builder.WriteString(s)
return s
}
b.seenBytes += len(s)
marker := b.truncationMarker()
contentLimit := b.maxBytes - len(marker)
if contentLimit < 0 {
marker = truncateStringBytes(marker, b.maxBytes)
contentLimit = 0
}
if b.builder.Len() >= contentLimit {
if b.truncated {
return ""
}
if b.builder.Len() >= b.maxBytes {
b.truncated = true
b.builder.WriteString(marker)
return marker
return ""
}
remaining := contentLimit - b.builder.Len()
remaining := b.maxBytes - b.builder.Len()
if len(s) <= remaining {
b.builder.WriteString(s)
return s
@@ -1079,8 +1106,7 @@ func (b *boundedOutputCollector) WriteStringLimited(s string) string {
kept := truncateStringBytes(s, remaining)
b.builder.WriteString(kept)
b.truncated = true
b.builder.WriteString(marker)
return kept + marker
return kept
}
func (b *boundedOutputCollector) String() string {
@@ -1090,13 +1116,90 @@ func (b *boundedOutputCollector) String() string {
return b.builder.String()
}
func (b *boundedOutputCollector) truncationMarker() string {
return fmt.Sprintf("\n\n...[tool output limit reached: kept %d bytes; further output suppressed]...", b.maxBytes)
func finalizeJoinedBoundedOutputs(stdout, stderr *boundedOutputCollector, maxBytes int, tee *tooloutput.Tee) string {
if tee != nil {
_ = tee.Close()
}
truncated := (stdout != nil && stdout.truncated) || (stderr != nil && stderr.truncated)
seen := 0
if stdout != nil {
seen += stdout.seenBytes
}
if stderr != nil {
seen += stderr.seenBytes
}
joined := joinCommandOutput(
func() string {
if stdout == nil {
return ""
}
return stdout.String()
}(),
func() string {
if stderr == nil {
return ""
}
return stderr.String()
}(),
)
if maxBytes > 0 && !truncated && len(joined) > maxBytes {
truncated = true
seen = len(joined)
}
path := ""
if tee != nil {
path = tee.Path()
}
if truncated && maxBytes > 0 {
if path != "" {
return tooloutput.FormatPersistedFromFile(path, seen, maxBytes)
}
if len(joined) > maxBytes {
return truncateStringBytes(joined, maxBytes)
}
return joined
}
if path != "" {
_ = os.Remove(path)
}
if maxBytes > 0 && len(joined) > maxBytes {
return truncateStringBytes(joined, maxBytes)
}
return joined
}
func limitOutputString(s string, maxBytes int) string {
collector := newBoundedOutputCollector(maxBytes)
return collector.WriteStringLimited(s)
func finalizeBoundedOutput(collector *boundedOutputCollector, maxBytes int, tee *tooloutput.Tee) string {
if tee != nil {
_ = tee.Close()
}
if collector == nil {
return ""
}
path := ""
if tee != nil {
path = tee.Path()
}
if collector.truncated && maxBytes > 0 {
if path != "" {
return tooloutput.FormatPersistedFromFile(path, collector.seenBytes, maxBytes)
}
return truncateStringBytes(collector.String(), maxBytes)
}
if path != "" {
_ = os.Remove(path)
}
out := collector.String()
if maxBytes > 0 && len(out) > maxBytes {
return tooloutput.BoundWithSpill(out, maxBytes, tooloutput.SpillOpts{})
}
return out
}
func limitOutputString(s string, maxBytes int, spill tooloutput.SpillOpts) string {
if maxBytes <= 0 || len(s) <= maxBytes {
return s
}
return tooloutput.BoundWithSpill(s, maxBytes, spill)
}
func truncateStringBytes(s string, maxBytes int) string {
@@ -1118,7 +1221,7 @@ func truncateStringBytes(s string, maxBytes int) string {
// streamCommandOutput 以“边读边回调”的方式读取命令 stdout/stderr。
// 使用定长块读取,避免按行读取在无换行输出时永久阻塞;ctx 取消时终止进程树。
func streamCommandOutput(ctx context.Context, cmd *exec.Cmd, cb ToolOutputCallback, noOutputSec int, maxBytes int) (string, error) {
func streamCommandOutput(ctx context.Context, cmd *exec.Cmd, cb ToolOutputCallback, noOutputSec int, maxBytes int, spill tooloutput.SpillOpts) (string, error) {
stdoutPipe, err := cmd.StdoutPipe()
if err != nil {
return "", err
@@ -1170,7 +1273,12 @@ func streamCommandOutput(ctx context.Context, cmd *exec.Cmd, cb ToolOutputCallba
close(chunks)
}()
outBuilder := newBoundedOutputCollector(maxBytes)
tee := (*tooloutput.Tee)(nil)
if maxBytes > 0 {
tee = tooloutput.NewTee(spill)
defer func() { _ = tee.Close() }()
}
outBuilder := newBoundedOutputCollector(maxBytes, tee)
var deltaBuilder strings.Builder
lastFlush := time.Now()
@@ -1214,7 +1322,7 @@ chunksLoop:
return outBuilder.String(), ctx.Err()
case <-idleCh:
fireInactivity()
return outBuilder.String(), fmt.Errorf("shell inactivity timeout (%ds)", idleWatch.Sec)
return finalizeBoundedOutput(outBuilder, maxBytes, tee), fmt.Errorf("shell inactivity timeout (%ds)", idleWatch.Sec)
case chunk, ok := <-chunks:
if !ok {
break chunksLoop
@@ -1233,7 +1341,7 @@ chunksLoop:
// 等待命令结束,返回最终退出状态
waitErr := session.Wait()
return outBuilder.String(), waitErr
return finalizeBoundedOutput(outBuilder, maxBytes, tee), waitErr
}
// applyDefaultTerminalEnv 为外部工具补齐常见的终端环境变量。
@@ -1286,14 +1394,14 @@ func shouldRetryWithPTY(output string) bool {
// runCommandWithPTY 为子进程分配 PTY,适配需要交互式终端的工具(如 autorecon)。
// 若 cb != nil,将持续回调增量输出(用于 SSE)。
func runCommandWithPTY(ctx context.Context, cmd *exec.Cmd, cb ToolOutputCallback, maxBytes int) (string, error) {
func runCommandWithPTY(ctx context.Context, cmd *exec.Cmd, cb ToolOutputCallback, maxBytes int, spill tooloutput.SpillOpts) (string, error) {
if runtime.GOOS == "windows" {
// PTY 方案为类 UnixWindows 走原逻辑
if cb != nil {
return streamCommandOutput(ctx, cmd, cb, 0, maxBytes)
return streamCommandOutput(ctx, cmd, cb, 0, maxBytes, spill)
}
_ = prepareShellCmdSession(cmd)
return combinedOutputCancellableWithLimit(ctx, cmd, maxBytes)
return combinedOutputCancellableWithLimit(ctx, cmd, maxBytes, spill)
}
_ = prepareShellCmdSession(cmd)
@@ -1320,7 +1428,12 @@ func runCommandWithPTY(ctx context.Context, cmd *exec.Cmd, cb ToolOutputCallback
}()
defer close(done)
outBuilder := newBoundedOutputCollector(maxBytes)
tee := (*tooloutput.Tee)(nil)
if maxBytes > 0 {
tee = tooloutput.NewTee(spill)
defer func() { _ = tee.Close() }()
}
outBuilder := newBoundedOutputCollector(maxBytes, tee)
var deltaBuilder strings.Builder
lastFlush := time.Now()
flush := func() {
@@ -1355,7 +1468,7 @@ func runCommandWithPTY(ctx context.Context, cmd *exec.Cmd, cb ToolOutputCallback
flush()
waitErr := cmd.Wait()
return outBuilder.String(), waitErr
return finalizeBoundedOutput(outBuilder, maxBytes, tee), waitErr
}
// executeInternalTool 执行内部工具(不执行外部命令)