mirror of
https://gitcode.com/JianFeeeee/HomeAgent.git
synced 2026-09-21 09:28:14 +00:00
- agent: 工具轮请求尾部补 user 占位(zen 网关强制),tool 消息正确配对 - agent: 工具提醒/中断以 system 角色注入并带 [中断消息] 前缀,不进用户履历;系统提示词说明中断消息格式 - agentcli: 基于 ConPTY 的交互式终端(ptywin fork),terminal_create/read/write/resize/close/watch - webui: server 输出通道适配器(保留 reasoning_content/disable_thinking) - GUI: 沉浸式标题栏、icon 圆角重制、mascot 等打磨
392 lines
12 KiB
Go
392 lines
12 KiB
Go
package core
|
||
|
||
import (
|
||
"context"
|
||
"errors"
|
||
"fmt"
|
||
"log"
|
||
"strings"
|
||
|
||
agentAPI "gitcode.com/JianFeeeee/HomeAgent/internal/agent/api"
|
||
"gitcode.com/JianFeeeee/HomeAgent/internal/events"
|
||
sdk "gitcode.com/JianFeeeee/HomeAgent/internal/sdk"
|
||
)
|
||
|
||
func (a *Agent) process(input string, stageCtx *sdk.StageContext) (response string, toolsUsed []string, toolResults []ToolResultItem, err error) {
|
||
a.mu.Lock()
|
||
defer a.mu.Unlock()
|
||
|
||
if a.provider == nil {
|
||
return "", nil, nil, fmt.Errorf("agent: no LLM provider configured")
|
||
}
|
||
|
||
budget := ComputeTokenBudget(a.provider, a.systemPrompt)
|
||
|
||
memContext := a.buildMemoryContext(input, budget.MemoryTokens)
|
||
sysPrompt := a.buildSystemPrompt(memContext, input)
|
||
tools := a.buildToolDefs()
|
||
|
||
msgs := a.buildMessages(sysPrompt, input, budget.ContextTokens)
|
||
// 工具提醒(interrupt):以 system 角色注入,不让模型误认为用户发言
|
||
if a.interruptInput {
|
||
last := msgs[len(msgs)-1]
|
||
last.Role = "system"
|
||
last.Content = "[中断消息] " + last.Content
|
||
msgs[len(msgs)-1] = last
|
||
a.interruptInput = false
|
||
}
|
||
if blocks, ok := stageCtx.Extra["media_blocks"].([]agentAPI.ContentBlock); ok && len(blocks) > 0 {
|
||
if len(msgs) > 0 {
|
||
msgs[len(msgs)-1].Blocks = blocks
|
||
}
|
||
}
|
||
|
||
log.Printf("[agent] tool call loop start, max_ctx=%d target=%d fixed=%d mem=%d ctx=%d %d tools, %d events, personality=%t, docs=%d",
|
||
budget.MaxContext, budget.TargetUsage, budget.FixedTokens, budget.MemoryTokens, budget.ContextTokens,
|
||
len(tools), a.context.Len(),
|
||
a.personality != nil && a.personality.Content != "",
|
||
a.docStoreSize())
|
||
|
||
if a.runStage(sdk.StagePreAction, stageCtx) {
|
||
return *stageCtx.Response, toolsUsed, toolResults, nil
|
||
}
|
||
if len(stageCtx.ContextMsgs) > 0 {
|
||
for _, m := range stageCtx.ContextMsgs {
|
||
role, _ := m["role"].(string)
|
||
content, _ := m["content"].(string)
|
||
if role != "" {
|
||
msgs = append(msgs, agentAPI.Message{Role: role, Content: content})
|
||
}
|
||
}
|
||
}
|
||
|
||
for turn := 0; ; turn++ {
|
||
for _, interrupt := range a.drainInterrupts() {
|
||
msgs = append(msgs, agentAPI.Message{
|
||
Role: "system",
|
||
Content: "[中断消息] " + interrupt,
|
||
})
|
||
}
|
||
|
||
// zen 兼容网关要求请求的最后一条消息必须是 user(thinking 续写模式校验),
|
||
// 工具轮产出的 tool/assistant 消息作结尾会被 400 拒绝,故补一条 user 占位。
|
||
if last := msgs[len(msgs)-1]; last.Role != "user" {
|
||
msgs = append(msgs, agentAPI.Message{
|
||
Role: "user",
|
||
Content: "请根据以上工具结果继续。",
|
||
})
|
||
}
|
||
|
||
req := &agentAPI.CompletionRequest{
|
||
Messages: msgs,
|
||
MaxTokens: 4096,
|
||
Tools: tools,
|
||
ToolChoice: "auto",
|
||
DisableThinking: !a.thinkingEnabled,
|
||
}
|
||
|
||
var providers []agentAPI.Provider
|
||
if a.providerManager != nil {
|
||
// 精确模型名走 byModel 路由;AUTO/空走优先级链
|
||
var allProviders []agentAPI.Provider
|
||
if req.Model != "" && !strings.EqualFold(req.Model, "AUTO") {
|
||
allProviders = a.providerManager.ResolveForModel(req.Model)
|
||
} else {
|
||
allProviders = a.providerManager.OrderedProviders()
|
||
}
|
||
providers = make([]agentAPI.Provider, 0, len(allProviders))
|
||
for _, p := range allProviders {
|
||
if a.providerManager.IsAvailable(p.Name()) {
|
||
providers = append(providers, p)
|
||
}
|
||
}
|
||
}
|
||
if len(providers) == 0 {
|
||
providers = []agentAPI.Provider{a.provider}
|
||
}
|
||
var resp *agentAPI.CompletionResponse
|
||
var llmErr error
|
||
|
||
for pi, fbProvider := range providers {
|
||
if pi > 0 {
|
||
log.Printf("[agent] LLM fallback: trying provider %q (fallback #%d/%d)",
|
||
fbProvider.Name(), pi, len(providers)-1)
|
||
}
|
||
|
||
fCtx, fCancel := context.WithCancel(a.ctx)
|
||
a.llmMu.Lock()
|
||
a.cancelLLM = fCancel
|
||
a.llmMu.Unlock()
|
||
|
||
resp, llmErr = fbProvider.Chat(fCtx, req)
|
||
|
||
a.llmMu.Lock()
|
||
a.cancelLLM = nil
|
||
a.llmMu.Unlock()
|
||
fCancel()
|
||
|
||
if llmErr == nil {
|
||
a.providerManager.ResetAvailability(fbProvider.Name())
|
||
if fbProvider != a.provider {
|
||
a.provider = fbProvider
|
||
log.Printf("[agent] switched active provider to %q after fallback",
|
||
fbProvider.Name())
|
||
}
|
||
break
|
||
}
|
||
|
||
if errors.Is(llmErr, context.Canceled) {
|
||
break
|
||
}
|
||
var pe *agentAPI.ProviderError
|
||
if errors.As(llmErr, &pe) && (pe.StatusCode == 401 || pe.StatusCode == 403) {
|
||
a.providerManager.ReportStatus(fbProvider.Name(), pe.StatusCode)
|
||
log.Printf("[agent] provider %q marked unavailable (HTTP %d)", fbProvider.Name(), pe.StatusCode)
|
||
} else {
|
||
a.providerManager.MarkUnavailable(fbProvider.Name())
|
||
}
|
||
log.Printf("[agent] provider %q failed: %v", fbProvider.Name(), llmErr)
|
||
}
|
||
|
||
if llmErr != nil {
|
||
if errors.Is(llmErr, context.Canceled) && a.ctx.Err() == nil {
|
||
if a.currentOutputChannel == "_consolidation_" {
|
||
return "", toolsUsed, toolResults, fmt.Errorf("interrupted by user input")
|
||
}
|
||
continue
|
||
}
|
||
return "", toolsUsed, toolResults, fmt.Errorf("all %d providers failed, last error: %w",
|
||
len(providers), llmErr)
|
||
}
|
||
|
||
stageCtx.LLMText = resp.Content
|
||
stageCtx.ReasoningContent = resp.ReasoningContent
|
||
stageCtx.TokenUsage = map[string]int{
|
||
"prompt_tokens": resp.TokenUsage.Prompt,
|
||
"completion_tokens": resp.TokenUsage.Completion,
|
||
"total_tokens": resp.TokenUsage.Total,
|
||
}
|
||
stageCtx.ToolCalls = convertToolCalls(resp.ToolCalls)
|
||
for i := range stageCtx.ToolCalls {
|
||
if stageCtx.ToolCalls[i].Plugin == "" {
|
||
stageCtx.ToolCalls[i].Plugin = a.resolveToolPlugin(stageCtx.ToolCalls[i].Name)
|
||
}
|
||
}
|
||
if a.runStage(sdk.StagePostAction, stageCtx) {
|
||
return *stageCtx.Response, toolsUsed, toolResults, nil
|
||
}
|
||
resp.Content = stageCtx.LLMText
|
||
resp.ToolCalls = convertBackToolCalls(stageCtx.ToolCalls)
|
||
|
||
chainPayload := map[string]interface{}{
|
||
"content": resp.Content,
|
||
"reasoning": resp.ReasoningContent,
|
||
"tool_calls": resp.ToolCalls,
|
||
"phase": "intermediate",
|
||
"turn": turn,
|
||
}
|
||
if resp.TokenUsage.Total > 0 {
|
||
chainPayload["usage"] = map[string]int{
|
||
"prompt": resp.TokenUsage.Prompt,
|
||
"completion": resp.TokenUsage.Completion,
|
||
"total": resp.TokenUsage.Total,
|
||
}
|
||
}
|
||
a.publishEvent(events.EventAgentLLMChain, chainPayload)
|
||
|
||
if resp.ReasoningContent != "" {
|
||
a.publishEvent(events.EventReasoning, map[string]interface{}{
|
||
"content": resp.ReasoningContent,
|
||
"channel": a.currentOutputChannel,
|
||
})
|
||
}
|
||
|
||
if len(resp.ToolCalls) == 0 {
|
||
return resp.Content, toolsUsed, toolResults, nil
|
||
}
|
||
|
||
contentOnce := true
|
||
for _, tc := range resp.ToolCalls {
|
||
if len(a.interceptCh) > 0 {
|
||
for _, interrupt := range a.drainInterrupts() {
|
||
msgs = append(msgs, agentAPI.Message{Role: "system", Content: "[中断消息] " + interrupt})
|
||
}
|
||
a.publishEvent(events.EventToolCall, map[string]interface{}{
|
||
"tool": tc.Name,
|
||
"plugin": a.resolveToolPlugin(tc.Name),
|
||
"args": tc.Arguments,
|
||
"status": "interrupted",
|
||
"reason": "user interrupt before execution",
|
||
"channel": a.currentOutputChannel,
|
||
})
|
||
break
|
||
}
|
||
|
||
toolsUsed = append(toolsUsed, tc.Name)
|
||
pluginName := a.resolveToolPlugin(tc.Name)
|
||
log.Printf("[agent] executing tool: %s (plugin=%s, id=%s)", tc.Name, pluginName, tc.ID)
|
||
|
||
sdkTC := sdk.ToolCall{ID: tc.ID, Name: tc.Name, Plugin: pluginName, Arguments: tc.Arguments}
|
||
stageCtx.ToolCalls = []sdk.ToolCall{sdkTC}
|
||
stageCtx.ToolResults = nil
|
||
if a.runStage(sdk.StageBeforeToolcall, stageCtx) {
|
||
result := fmt.Sprintf("工具 %s 已被插件拒绝", tc.Name)
|
||
msgs = append(msgs, agentAPI.Message{Role: "assistant", ToolCalls: []agentAPI.ToolCall{tc}})
|
||
msgs = append(msgs, agentAPI.Message{Role: "tool", ToolCallID: tc.ID, Content: result})
|
||
a.publishEvent(events.EventToolCall, map[string]interface{}{
|
||
"tool": tc.Name,
|
||
"plugin": pluginName,
|
||
"args": tc.Arguments,
|
||
"result": result,
|
||
"status": "denied",
|
||
"channel": a.currentOutputChannel,
|
||
})
|
||
continue
|
||
}
|
||
tc.Arguments = stageCtx.ToolCalls[0].Arguments
|
||
|
||
if pluginName != "" && !a.pluginHealth.isHealthy(pluginName) {
|
||
result := fmt.Sprintf("插件 %s 处于崩溃状态,已跳过执行,等待自动恢复重载", pluginName)
|
||
log.Printf("[agent] skip tool %s: plugin %s unhealthy", tc.Name, pluginName)
|
||
msgs = append(msgs, agentAPI.Message{Role: "assistant", ToolCalls: []agentAPI.ToolCall{tc}})
|
||
msgs = append(msgs, agentAPI.Message{Role: "tool", ToolCallID: tc.ID, Content: result})
|
||
continue
|
||
}
|
||
|
||
result := a.executeToolCall(tc)
|
||
toolResults = append(toolResults, ToolResultItem{Name: tc.Name, Output: result})
|
||
log.Printf("[agent] tool %s result: %s", tc.Name, truncateStr(result, 100))
|
||
|
||
stageCtx.ToolResults = []sdk.ToolResult{{CallID: tc.ID, Name: tc.Name, Plugin: pluginName, Success: true, Result: result}}
|
||
a.runStage(sdk.StageAfterToolcall, stageCtx)
|
||
if len(stageCtx.ToolResults) > 0 {
|
||
if r, ok := stageCtx.ToolResults[0].Result.(string); ok {
|
||
result = r
|
||
}
|
||
}
|
||
|
||
msgContent := ""
|
||
if contentOnce {
|
||
msgContent = resp.Content
|
||
contentOnce = false
|
||
}
|
||
msgs = append(msgs, agentAPI.Message{Role: "assistant", Content: msgContent, ReasoningContent: resp.ReasoningContent, ToolCalls: []agentAPI.ToolCall{tc}})
|
||
msgs = append(msgs, agentAPI.Message{Role: "tool", ToolCallID: tc.ID, Content: result})
|
||
|
||
a.publishEvent(events.EventToolCall, map[string]interface{}{
|
||
"tool": tc.Name,
|
||
"plugin": pluginName,
|
||
"args": tc.Arguments,
|
||
"result": result,
|
||
"status": "ok",
|
||
"channel": a.currentOutputChannel,
|
||
})
|
||
|
||
if len(a.interceptCh) > 0 {
|
||
for _, interrupt := range a.drainInterrupts() {
|
||
msgs = append(msgs, agentAPI.Message{Role: "system", Content: "[中断消息] " + interrupt})
|
||
}
|
||
break
|
||
}
|
||
}
|
||
}
|
||
}
|
||
|
||
func convertToolCalls(tcs []agentAPI.ToolCall) []sdk.ToolCall {
|
||
if tcs == nil {
|
||
return nil
|
||
}
|
||
result := make([]sdk.ToolCall, len(tcs))
|
||
for i, tc := range tcs {
|
||
result[i] = sdk.ToolCall{ID: tc.ID, Name: tc.Name, Arguments: tc.Arguments}
|
||
}
|
||
return result
|
||
}
|
||
|
||
func convertBackToolCalls(tcs []sdk.ToolCall) []agentAPI.ToolCall {
|
||
if tcs == nil {
|
||
return nil
|
||
}
|
||
result := make([]agentAPI.ToolCall, len(tcs))
|
||
for i, tc := range tcs {
|
||
result[i] = agentAPI.ToolCall{ID: tc.ID, Name: tc.Name, Arguments: tc.Arguments}
|
||
}
|
||
return result
|
||
}
|
||
|
||
func (a *Agent) docStoreSize() int {
|
||
if a.docStore == nil {
|
||
return 0
|
||
}
|
||
s := a.docStore.Stats()
|
||
if n, ok := s["doc_count"]; ok {
|
||
if ni, ok := n.(int); ok {
|
||
return ni
|
||
}
|
||
}
|
||
return 0
|
||
}
|
||
|
||
func (a *Agent) formatMergedTimeline(maxTokens int) string {
|
||
a.context.mu.Lock()
|
||
events := make([]*ContextEvent, len(a.context.events))
|
||
copy(events, a.context.events)
|
||
a.context.mu.Unlock()
|
||
|
||
if len(events) == 0 {
|
||
return ""
|
||
}
|
||
|
||
// 第一轮:从最新到最旧,计算在预算内能放多少条
|
||
headerTokens := EstimateTokens("【对话时序】\n")
|
||
remaining := maxTokens - headerTokens
|
||
include := 0
|
||
for i := len(events) - 1; i >= 0; i-- {
|
||
e := events[i]
|
||
est := len(e.Source) + len(e.Input) + 40
|
||
if e.Response != "" {
|
||
est += 120
|
||
}
|
||
estTokens := est * 2
|
||
if remaining-estTokens < 0 && include > 0 {
|
||
break
|
||
}
|
||
remaining -= estTokens
|
||
include++
|
||
}
|
||
if include == 0 && len(events) > 0 {
|
||
include = 1
|
||
}
|
||
|
||
// 第二轮:按时间正序渲染
|
||
start := len(events) - include
|
||
if start < 0 {
|
||
start = 0
|
||
}
|
||
var sb strings.Builder
|
||
sb.WriteString("【对话时序】\n")
|
||
for _, e := range events[start:] {
|
||
sb.WriteString(fmt.Sprintf("[%s] %s: %s",
|
||
e.Timestamp.Format("15:04:05"), e.Source, e.Input))
|
||
if len(e.ToolsUsed) > 0 {
|
||
sb.WriteString(fmt.Sprintf(" → 调用工具: %s", strings.Join(e.ToolsUsed, ", ")))
|
||
}
|
||
if e.Response != "" {
|
||
sb.WriteString(fmt.Sprintf(" → %s", truncateStr(e.Response, 120)))
|
||
}
|
||
sb.WriteString("\n")
|
||
}
|
||
return sb.String()
|
||
}
|
||
|
||
func (a *Agent) buildMessages(sysPrompt, input string, ctxTokens int) []agentAPI.Message {
|
||
msgs := []agentAPI.Message{{Role: "system", Content: sysPrompt}}
|
||
|
||
if ctxTok := a.formatMergedTimeline(ctxTokens); ctxTok != "" {
|
||
msgs = append(msgs, agentAPI.Message{Role: "system", Content: ctxTok})
|
||
}
|
||
|
||
msgs = append(msgs, agentAPI.Message{Role: "user", Content: input})
|
||
return msgs
|
||
}
|