更新支持流式输出
This commit is contained in:
parent
91e7e18e85
commit
55ac8c9dc5
12 changed files with 872 additions and 60 deletions
72
CHANGE_LOG.md
Normal file
72
CHANGE_LOG.md
Normal file
|
|
@ -0,0 +1,72 @@
|
||||||
|
# 修改日志
|
||||||
|
|
||||||
|
## 2026-02-10 - 将提示文本改为英文
|
||||||
|
|
||||||
|
### 修改内容
|
||||||
|
将大模型执行操作时的所有中文提示文本改为英文,提升国际化体验。
|
||||||
|
|
||||||
|
### 修改的文件
|
||||||
|
|
||||||
|
#### 1. pkg/agent/loop.go
|
||||||
|
修改所有流式事件的提示文本:
|
||||||
|
- ✅ `"开始处理请求..."` → `"Starting to process request..."`
|
||||||
|
- ✅ `"思考中... (迭代 %d/%d)"` → `"Thinking... (iteration %d/%d)"`
|
||||||
|
- ✅ `"执行工具: %s"` → `"Executing tool: %s"`
|
||||||
|
- ✅ `"工具 %s 完成"` → `"Tool %s completed"`
|
||||||
|
- ✅ `"处理完成"` → `"Processing complete"`
|
||||||
|
- ✅ `"内容生成完成"` → `"Content generation complete"`
|
||||||
|
- ✅ `"处理被中断"` → `"Processing interrupted"`
|
||||||
|
- ✅ `"用户中断"` → `"Interrupted by user"`
|
||||||
|
- ✅ `"LLM 调用失败"` → `"LLM call failed"`
|
||||||
|
|
||||||
|
#### 2. cmd/picoclaw/main.go
|
||||||
|
修改 CLI 的中断提示:
|
||||||
|
- ✅ `"中断处理中..."` → `"Interrupting..."`
|
||||||
|
- ✅ `"已中断"` → `"Interrupted"`
|
||||||
|
|
||||||
|
#### 3. pkg/channels/telegram.go
|
||||||
|
修改 Telegram 状态消息:
|
||||||
|
- ✅ `"执行工具: %s"` → `"Executing tool: %s"`
|
||||||
|
- ✅ `"工具 %s 完成"` → `"Tool %s completed"`
|
||||||
|
|
||||||
|
### 效果对比
|
||||||
|
|
||||||
|
#### 修改前
|
||||||
|
```
|
||||||
|
🦞 You: 帮我搜索天气
|
||||||
|
🦞 💭 开始处理请求...
|
||||||
|
🦞 ⏳ 思考中... (迭代 1/20)
|
||||||
|
🦞 🔧 执行工具: web_search
|
||||||
|
🦞 ✓ 工具 web_search 完成
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 修改后
|
||||||
|
```
|
||||||
|
🦞 You: Search for weather
|
||||||
|
🦞 💭 Starting to process request...
|
||||||
|
🦞 ⏳ Thinking... (iteration 1/20)
|
||||||
|
🦞 🔧 Executing tool: web_search
|
||||||
|
🦞 ✓ Tool web_search completed
|
||||||
|
```
|
||||||
|
|
||||||
|
### 测试结果
|
||||||
|
```bash
|
||||||
|
$ go build ./cmd/picoclaw
|
||||||
|
✅ 编译成功
|
||||||
|
|
||||||
|
$ go test ./pkg/bus -v
|
||||||
|
✅ 所有测试通过
|
||||||
|
```
|
||||||
|
|
||||||
|
### 影响范围
|
||||||
|
- ✅ CLI 交互模式
|
||||||
|
- ✅ Telegram Bot
|
||||||
|
- ✅ 流式事件系统
|
||||||
|
- ✅ 错误提示信息
|
||||||
|
|
||||||
|
### 向后兼容
|
||||||
|
完全兼容,仅修改显示文本,不影响功能逻辑。
|
||||||
|
|
||||||
|
---
|
||||||
|
**修改人**: OpenCode Assistant
|
||||||
|
**日期**: 2026-02-10
|
||||||
|
|
@ -142,7 +142,8 @@ func main() {
|
||||||
|
|
||||||
func printHelp() {
|
func printHelp() {
|
||||||
fmt.Printf("%s picoclaw - Personal AI Assistant v%s\n\n", logo, version)
|
fmt.Printf("%s picoclaw - Personal AI Assistant v%s\n\n", logo, version)
|
||||||
fmt.Println("Usage: picoclaw <command>\n")
|
fmt.Println("Usage: picoclaw <command>")
|
||||||
|
fmt.Println()
|
||||||
fmt.Println("Commands:")
|
fmt.Println("Commands:")
|
||||||
fmt.Println(" onboard Initialize picoclaw configuration and workspace")
|
fmt.Println(" onboard Initialize picoclaw configuration and workspace")
|
||||||
fmt.Println(" agent Interact with the agent directly")
|
fmt.Println(" agent Interact with the agent directly")
|
||||||
|
|
@ -450,12 +451,41 @@ func agentCmd() {
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
}
|
}
|
||||||
|
|
||||||
bus := bus.NewMessageBus()
|
msgBus := bus.NewMessageBus()
|
||||||
agentLoop := agent.NewAgentLoop(cfg, bus, provider)
|
agentLoop := agent.NewAgentLoop(cfg, msgBus, provider)
|
||||||
|
|
||||||
if message != "" {
|
if message != "" {
|
||||||
ctx := context.Background()
|
ctx := context.Background()
|
||||||
response, err := agentLoop.ProcessDirect(ctx, message, sessionKey)
|
var response string
|
||||||
|
var err error
|
||||||
|
|
||||||
|
if agentLoop.IsStreamingEnabled() {
|
||||||
|
// Streaming mode - show live updates
|
||||||
|
response, err = agentLoop.ProcessDirectStreaming(ctx, message, sessionKey, func(event bus.StreamEvent) {
|
||||||
|
// Display streaming events for -m flag
|
||||||
|
switch event.Type {
|
||||||
|
case bus.StreamEventThinking:
|
||||||
|
fmt.Printf("\r%s 💭 %s", logo, event.Content)
|
||||||
|
case bus.StreamEventProgress:
|
||||||
|
fmt.Printf("\r%s ⏳ %s", logo, event.Content)
|
||||||
|
case bus.StreamEventToolCall:
|
||||||
|
fmt.Printf("\r%s 🔧 %s\n", logo, event.Content)
|
||||||
|
case bus.StreamEventToolResult:
|
||||||
|
fmt.Printf("%s ✓ %s\n", logo, event.Content)
|
||||||
|
case bus.StreamEventContent:
|
||||||
|
// Check if this is partial content (streaming)
|
||||||
|
if partial, ok := event.Metadata["partial"].(bool); ok && partial {
|
||||||
|
fmt.Print(event.Content)
|
||||||
|
}
|
||||||
|
case bus.StreamEventError:
|
||||||
|
fmt.Printf("\r%s ❌ %s\n", logo, event.Content)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
} else {
|
||||||
|
// Non-streaming mode
|
||||||
|
response, err = agentLoop.ProcessDirect(ctx, message, sessionKey)
|
||||||
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
fmt.Printf("Error: %v\n", err)
|
fmt.Printf("Error: %v\n", err)
|
||||||
os.Exit(1)
|
os.Exit(1)
|
||||||
|
|
@ -507,14 +537,90 @@ func interactiveMode(agentLoop *agent.AgentLoop, sessionKey string) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx := context.Background()
|
// Create cancellable context for interrupt support
|
||||||
response, err := agentLoop.ProcessDirect(ctx, input, sessionKey)
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
if err != nil {
|
|
||||||
fmt.Printf("Error: %v\n", err)
|
// Handle Ctrl+C during processing
|
||||||
|
sigChan := make(chan os.Signal, 1)
|
||||||
|
signal.Notify(sigChan, os.Interrupt)
|
||||||
|
|
||||||
|
done := make(chan bool)
|
||||||
|
var response string
|
||||||
|
var processErr error
|
||||||
|
|
||||||
|
if agentLoop.IsStreamingEnabled() {
|
||||||
|
// Streaming mode with live updates
|
||||||
|
var contentStarted bool
|
||||||
|
go func() {
|
||||||
|
response, processErr = agentLoop.ProcessDirectStreaming(ctx, input, sessionKey, func(event bus.StreamEvent) {
|
||||||
|
// Display streaming events
|
||||||
|
switch event.Type {
|
||||||
|
case bus.StreamEventThinking:
|
||||||
|
fmt.Printf("\r%s 💭 %s", logo, event.Content)
|
||||||
|
case bus.StreamEventProgress:
|
||||||
|
fmt.Printf("\r%s ⏳ %s", logo, event.Content)
|
||||||
|
case bus.StreamEventToolCall:
|
||||||
|
fmt.Printf("\r%s 🔧 %s\n", logo, event.Content)
|
||||||
|
case bus.StreamEventToolResult:
|
||||||
|
fmt.Printf("%s ✓ %s\n", logo, event.Content)
|
||||||
|
case bus.StreamEventContent:
|
||||||
|
// Check if this is partial content (streaming)
|
||||||
|
if partial, ok := event.Metadata["partial"].(bool); ok && partial {
|
||||||
|
if !contentStarted {
|
||||||
|
// Clear progress line and show logo
|
||||||
|
fmt.Print("\r\033[K")
|
||||||
|
fmt.Printf("%s ", logo)
|
||||||
|
contentStarted = true
|
||||||
|
}
|
||||||
|
// Print each chunk as it arrives (like typing)
|
||||||
|
fmt.Print(event.Content)
|
||||||
|
}
|
||||||
|
case bus.StreamEventError:
|
||||||
|
fmt.Printf("\r%s ❌ %s\n", logo, event.Content)
|
||||||
|
case bus.StreamEventComplete:
|
||||||
|
if contentStarted {
|
||||||
|
fmt.Println() // New line after streaming content
|
||||||
|
contentStarted = false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
done <- true
|
||||||
|
}()
|
||||||
|
} else {
|
||||||
|
// Non-streaming mode - simple wait
|
||||||
|
go func() {
|
||||||
|
response, processErr = agentLoop.ProcessDirect(ctx, input, sessionKey)
|
||||||
|
done <- true
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
// Wait for completion or interrupt
|
||||||
|
interrupted := false
|
||||||
|
select {
|
||||||
|
case <-done:
|
||||||
|
// Processing completed normally
|
||||||
|
case <-sigChan:
|
||||||
|
// User pressed Ctrl+C
|
||||||
|
interrupted = true
|
||||||
|
cancel()
|
||||||
|
fmt.Printf("\n%s Interrupting...\n", logo)
|
||||||
|
<-done // Wait for goroutine to finish
|
||||||
|
}
|
||||||
|
|
||||||
|
signal.Stop(sigChan)
|
||||||
|
close(sigChan)
|
||||||
|
|
||||||
|
if interrupted {
|
||||||
|
fmt.Printf("%s Interrupted\n\n", logo)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
fmt.Printf("\n%s %s\n\n", logo, response)
|
if processErr != nil {
|
||||||
|
fmt.Printf("Error: %v\n", processErr)
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
fmt.Printf("%s %s\n\n", logo, response)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -542,10 +648,64 @@ func simpleInteractiveMode(agentLoop *agent.AgentLoop, sessionKey string) {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
ctx := context.Background()
|
// Create cancellable context
|
||||||
response, err := agentLoop.ProcessDirect(ctx, input, sessionKey)
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
if err != nil {
|
|
||||||
fmt.Printf("Error: %v\n", err)
|
done := make(chan bool)
|
||||||
|
var response string
|
||||||
|
var processErr error
|
||||||
|
|
||||||
|
if agentLoop.IsStreamingEnabled() {
|
||||||
|
// Streaming mode with live updates
|
||||||
|
var contentStarted bool
|
||||||
|
go func() {
|
||||||
|
response, processErr = agentLoop.ProcessDirectStreaming(ctx, input, sessionKey, func(event bus.StreamEvent) {
|
||||||
|
// Display streaming events
|
||||||
|
switch event.Type {
|
||||||
|
case bus.StreamEventThinking:
|
||||||
|
fmt.Printf("\r%s 💭 %s", logo, event.Content)
|
||||||
|
case bus.StreamEventProgress:
|
||||||
|
fmt.Printf("\r%s ⏳ %s", logo, event.Content)
|
||||||
|
case bus.StreamEventToolCall:
|
||||||
|
fmt.Printf("\r%s 🔧 %s\n", logo, event.Content)
|
||||||
|
case bus.StreamEventToolResult:
|
||||||
|
fmt.Printf("%s ✓ %s\n", logo, event.Content)
|
||||||
|
case bus.StreamEventContent:
|
||||||
|
// Check if this is partial content (streaming)
|
||||||
|
if partial, ok := event.Metadata["partial"].(bool); ok && partial {
|
||||||
|
if !contentStarted {
|
||||||
|
// Clear progress line and show logo
|
||||||
|
fmt.Print("\r\033[K")
|
||||||
|
fmt.Printf("%s ", logo)
|
||||||
|
contentStarted = true
|
||||||
|
}
|
||||||
|
// Print each chunk as it arrives (like typing)
|
||||||
|
fmt.Print(event.Content)
|
||||||
|
}
|
||||||
|
case bus.StreamEventError:
|
||||||
|
fmt.Printf("\r%s ❌ %s\n", logo, event.Content)
|
||||||
|
case bus.StreamEventComplete:
|
||||||
|
if contentStarted {
|
||||||
|
fmt.Println() // New line after streaming content
|
||||||
|
contentStarted = false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
})
|
||||||
|
done <- true
|
||||||
|
}()
|
||||||
|
} else {
|
||||||
|
// Non-streaming mode - simple wait
|
||||||
|
go func() {
|
||||||
|
response, processErr = agentLoop.ProcessDirect(ctx, input, sessionKey)
|
||||||
|
done <- true
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
|
||||||
|
<-done
|
||||||
|
cancel()
|
||||||
|
|
||||||
|
if processErr != nil {
|
||||||
|
fmt.Printf("Error: %v\n", processErr)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -5,7 +5,8 @@
|
||||||
"model": "glm-4.7",
|
"model": "glm-4.7",
|
||||||
"max_tokens": 8192,
|
"max_tokens": 8192,
|
||||||
"temperature": 0.7,
|
"temperature": 0.7,
|
||||||
"max_tool_iterations": 20
|
"max_tool_iterations": 20,
|
||||||
|
"streaming": false
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
"channels": {
|
"channels": {
|
||||||
|
|
|
||||||
4
go.mod
4
go.mod
|
|
@ -1,6 +1,8 @@
|
||||||
module github.com/sipeed/picoclaw
|
module github.com/sipeed/picoclaw
|
||||||
|
|
||||||
go 1.18
|
go 1.24.0
|
||||||
|
|
||||||
|
toolchain go1.24.13
|
||||||
|
|
||||||
require (
|
require (
|
||||||
github.com/bwmarrin/discordgo v0.28.1
|
github.com/bwmarrin/discordgo v0.28.1
|
||||||
|
|
|
||||||
BIN
picoclaw
Executable file
BIN
picoclaw
Executable file
Binary file not shown.
|
|
@ -29,6 +29,7 @@ type AgentLoop struct {
|
||||||
model string
|
model string
|
||||||
contextWindow int
|
contextWindow int
|
||||||
maxIterations int
|
maxIterations int
|
||||||
|
streaming bool
|
||||||
sessions *session.SessionManager
|
sessions *session.SessionManager
|
||||||
contextBuilder *ContextBuilder
|
contextBuilder *ContextBuilder
|
||||||
tools *tools.ToolRegistry
|
tools *tools.ToolRegistry
|
||||||
|
|
@ -59,6 +60,7 @@ func NewAgentLoop(cfg *config.Config, bus *bus.MessageBus, provider providers.LL
|
||||||
model: cfg.Agents.Defaults.Model,
|
model: cfg.Agents.Defaults.Model,
|
||||||
contextWindow: cfg.Agents.Defaults.MaxTokens,
|
contextWindow: cfg.Agents.Defaults.MaxTokens,
|
||||||
maxIterations: cfg.Agents.Defaults.MaxToolIterations,
|
maxIterations: cfg.Agents.Defaults.MaxToolIterations,
|
||||||
|
streaming: cfg.Agents.Defaults.Streaming,
|
||||||
sessions: sessionsManager,
|
sessions: sessionsManager,
|
||||||
contextBuilder: NewContextBuilder(workspace),
|
contextBuilder: NewContextBuilder(workspace),
|
||||||
tools: toolsRegistry,
|
tools: toolsRegistry,
|
||||||
|
|
@ -70,6 +72,11 @@ func NewAgentLoop(cfg *config.Config, bus *bus.MessageBus, provider providers.LL
|
||||||
func (al *AgentLoop) Run(ctx context.Context) error {
|
func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
al.running = true
|
al.running = true
|
||||||
|
|
||||||
|
// Start stream event forwarder only if streaming is enabled
|
||||||
|
if al.streaming {
|
||||||
|
go al.forwardStreamEvents(ctx)
|
||||||
|
}
|
||||||
|
|
||||||
for al.running {
|
for al.running {
|
||||||
select {
|
select {
|
||||||
case <-ctx.Done():
|
case <-ctx.Done():
|
||||||
|
|
@ -80,7 +87,18 @@ func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
|
|
||||||
response, err := al.processMessage(ctx, msg)
|
var response string
|
||||||
|
var err error
|
||||||
|
|
||||||
|
// Choose streaming or non-streaming based on config
|
||||||
|
if al.streaming {
|
||||||
|
response, err = al.processMessageStreaming(ctx, msg, func(event bus.StreamEvent) {
|
||||||
|
al.bus.PublishStreamEvent(event)
|
||||||
|
})
|
||||||
|
} else {
|
||||||
|
response, err = al.processMessage(ctx, msg)
|
||||||
|
}
|
||||||
|
|
||||||
if err != nil {
|
if err != nil {
|
||||||
response = fmt.Sprintf("Error processing message: %v", err)
|
response = fmt.Sprintf("Error processing message: %v", err)
|
||||||
}
|
}
|
||||||
|
|
@ -98,6 +116,19 @@ func (al *AgentLoop) Run(ctx context.Context) error {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// forwardStreamEvents handles stream events and can be extended for logging
|
||||||
|
func (al *AgentLoop) forwardStreamEvents(ctx context.Context) {
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
// Events are consumed by channels, this is just for future extensions
|
||||||
|
time.Sleep(100 * time.Millisecond)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) Stop() {
|
func (al *AgentLoop) Stop() {
|
||||||
al.running = false
|
al.running = false
|
||||||
}
|
}
|
||||||
|
|
@ -114,6 +145,24 @@ func (al *AgentLoop) ProcessDirect(ctx context.Context, content, sessionKey stri
|
||||||
return al.processMessage(ctx, msg)
|
return al.processMessage(ctx, msg)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ProcessDirectStreaming processes a message with streaming updates
|
||||||
|
func (al *AgentLoop) ProcessDirectStreaming(ctx context.Context, content, sessionKey string, streamCallback func(bus.StreamEvent)) (string, error) {
|
||||||
|
msg := bus.InboundMessage{
|
||||||
|
Channel: "cli",
|
||||||
|
SenderID: "user",
|
||||||
|
ChatID: "direct",
|
||||||
|
Content: content,
|
||||||
|
SessionKey: sessionKey,
|
||||||
|
}
|
||||||
|
|
||||||
|
return al.processMessageStreaming(ctx, msg, streamCallback)
|
||||||
|
}
|
||||||
|
|
||||||
|
// IsStreamingEnabled returns whether streaming is enabled
|
||||||
|
func (al *AgentLoop) IsStreamingEnabled() bool {
|
||||||
|
return al.streaming
|
||||||
|
}
|
||||||
|
|
||||||
func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage) (string, error) {
|
func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage) (string, error) {
|
||||||
history := al.sessions.GetHistory(msg.SessionKey)
|
history := al.sessions.GetHistory(msg.SessionKey)
|
||||||
summary := al.sessions.GetSummary(msg.SessionKey)
|
summary := al.sessions.GetSummary(msg.SessionKey)
|
||||||
|
|
@ -200,12 +249,203 @@ func (al *AgentLoop) processMessage(ctx context.Context, msg bus.InboundMessage)
|
||||||
|
|
||||||
// Context compression logic
|
// Context compression logic
|
||||||
newHistory := al.sessions.GetHistory(msg.SessionKey)
|
newHistory := al.sessions.GetHistory(msg.SessionKey)
|
||||||
|
|
||||||
// Token Awareness (Dynamic)
|
// Token Awareness (Dynamic)
|
||||||
// Trigger if history > 20 messages OR estimated tokens > 75% of context window
|
// Trigger if history > 20 messages OR estimated tokens > 75% of context window
|
||||||
tokenEstimate := al.estimateTokens(newHistory)
|
tokenEstimate := al.estimateTokens(newHistory)
|
||||||
threshold := al.contextWindow * 75 / 100
|
threshold := al.contextWindow * 75 / 100
|
||||||
|
|
||||||
|
if len(newHistory) > 20 || tokenEstimate > threshold {
|
||||||
|
if _, loading := al.summarizing.LoadOrStore(msg.SessionKey, true); !loading {
|
||||||
|
go func() {
|
||||||
|
defer al.summarizing.Delete(msg.SessionKey)
|
||||||
|
al.summarizeSession(msg.SessionKey)
|
||||||
|
}()
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
al.sessions.Save(al.sessions.GetOrCreate(msg.SessionKey))
|
||||||
|
|
||||||
|
return finalContent, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// processMessageStreaming processes a message with streaming updates
|
||||||
|
func (al *AgentLoop) processMessageStreaming(ctx context.Context, msg bus.InboundMessage, streamCallback func(bus.StreamEvent)) (string, error) {
|
||||||
|
// Check if callback is provided, if not fall back to non-streaming
|
||||||
|
if streamCallback == nil {
|
||||||
|
return al.processMessage(ctx, msg)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Helper function to emit events
|
||||||
|
emitEvent := func(eventType bus.StreamEventType, content string, metadata map[string]interface{}) {
|
||||||
|
if streamCallback != nil {
|
||||||
|
streamCallback(bus.StreamEvent{
|
||||||
|
Type: eventType,
|
||||||
|
Channel: msg.Channel,
|
||||||
|
ChatID: msg.ChatID,
|
||||||
|
SessionKey: msg.SessionKey,
|
||||||
|
Content: content,
|
||||||
|
Metadata: metadata,
|
||||||
|
Timestamp: time.Now().Unix(),
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
history := al.sessions.GetHistory(msg.SessionKey)
|
||||||
|
summary := al.sessions.GetSummary(msg.SessionKey)
|
||||||
|
|
||||||
|
messages := al.contextBuilder.BuildMessages(
|
||||||
|
history,
|
||||||
|
summary,
|
||||||
|
msg.Content,
|
||||||
|
nil,
|
||||||
|
)
|
||||||
|
|
||||||
|
iteration := 0
|
||||||
|
var finalContent string
|
||||||
|
|
||||||
|
emitEvent(bus.StreamEventThinking, "Starting to process request...", nil)
|
||||||
|
|
||||||
|
for iteration < al.maxIterations {
|
||||||
|
// Check for interrupt
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
emitEvent(bus.StreamEventError, "Processing interrupted", nil)
|
||||||
|
return "", fmt.Errorf("processing interrupted")
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
|
||||||
|
// Check bus for interrupt signal
|
||||||
|
if al.bus.CheckInterrupt(msg.SessionKey) {
|
||||||
|
emitEvent(bus.StreamEventError, "Interrupted by user", nil)
|
||||||
|
return "", fmt.Errorf("interrupted by user")
|
||||||
|
}
|
||||||
|
|
||||||
|
iteration++
|
||||||
|
emitEvent(bus.StreamEventProgress, fmt.Sprintf("Thinking... (iteration %d/%d)", iteration, al.maxIterations), map[string]interface{}{
|
||||||
|
"iteration": iteration,
|
||||||
|
"max_iteration": al.maxIterations,
|
||||||
|
})
|
||||||
|
|
||||||
|
toolDefs := al.tools.GetDefinitions()
|
||||||
|
providerToolDefs := make([]providers.ToolDefinition, 0, len(toolDefs))
|
||||||
|
for _, td := range toolDefs {
|
||||||
|
providerToolDefs = append(providerToolDefs, providers.ToolDefinition{
|
||||||
|
Type: td["type"].(string),
|
||||||
|
Function: providers.ToolFunctionDefinition{
|
||||||
|
Name: td["function"].(map[string]interface{})["name"].(string),
|
||||||
|
Description: td["function"].(map[string]interface{})["description"].(string),
|
||||||
|
Parameters: td["function"].(map[string]interface{})["parameters"].(map[string]interface{}),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
// Use streaming API for real-time output
|
||||||
|
response, err := al.provider.ChatStream(ctx, messages, providerToolDefs, al.model, map[string]interface{}{
|
||||||
|
"max_tokens": 8192,
|
||||||
|
"temperature": 0.7,
|
||||||
|
}, func(chunk string) {
|
||||||
|
// Stream each token as it arrives
|
||||||
|
if streamCallback != nil {
|
||||||
|
emitEvent(bus.StreamEventContent, chunk, map[string]interface{}{
|
||||||
|
"partial": true,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
})
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
emitEvent(bus.StreamEventError, fmt.Sprintf("LLM call failed: %v", err), nil)
|
||||||
|
return "", fmt.Errorf("LLM call failed: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(response.ToolCalls) == 0 {
|
||||||
|
finalContent = response.Content
|
||||||
|
emitEvent(bus.StreamEventComplete, "Content generation complete", nil)
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
// Emit tool calls
|
||||||
|
for _, tc := range response.ToolCalls {
|
||||||
|
emitEvent(bus.StreamEventToolCall, fmt.Sprintf("Executing tool: %s", tc.Name), map[string]interface{}{
|
||||||
|
"tool_name": tc.Name,
|
||||||
|
"tool_id": tc.ID,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
assistantMsg := providers.Message{
|
||||||
|
Role: "assistant",
|
||||||
|
Content: response.Content,
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tc := range response.ToolCalls {
|
||||||
|
argumentsJSON, _ := json.Marshal(tc.Arguments)
|
||||||
|
assistantMsg.ToolCalls = append(assistantMsg.ToolCalls, providers.ToolCall{
|
||||||
|
ID: tc.ID,
|
||||||
|
Type: "function",
|
||||||
|
Function: &providers.FunctionCall{
|
||||||
|
Name: tc.Name,
|
||||||
|
Arguments: string(argumentsJSON),
|
||||||
|
},
|
||||||
|
})
|
||||||
|
}
|
||||||
|
messages = append(messages, assistantMsg)
|
||||||
|
|
||||||
|
for _, tc := range response.ToolCalls {
|
||||||
|
// Check for interrupt before tool execution
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
emitEvent(bus.StreamEventError, "Processing interrupted", nil)
|
||||||
|
return "", fmt.Errorf("processing interrupted")
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
|
||||||
|
if al.bus.CheckInterrupt(msg.SessionKey) {
|
||||||
|
emitEvent(bus.StreamEventError, "Interrupted by user", nil)
|
||||||
|
return "", fmt.Errorf("interrupted by user")
|
||||||
|
}
|
||||||
|
|
||||||
|
result, err := al.tools.Execute(ctx, tc.Name, tc.Arguments)
|
||||||
|
if err != nil {
|
||||||
|
result = fmt.Sprintf("Error: %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
// Truncate long results for display
|
||||||
|
displayResult := result
|
||||||
|
if len(displayResult) > 200 {
|
||||||
|
displayResult = displayResult[:200] + "..."
|
||||||
|
}
|
||||||
|
|
||||||
|
emitEvent(bus.StreamEventToolResult, fmt.Sprintf("Tool %s completed", tc.Name), map[string]interface{}{
|
||||||
|
"tool_name": tc.Name,
|
||||||
|
"tool_id": tc.ID,
|
||||||
|
"result": displayResult,
|
||||||
|
})
|
||||||
|
|
||||||
|
toolResultMsg := providers.Message{
|
||||||
|
Role: "tool",
|
||||||
|
Content: result,
|
||||||
|
ToolCallID: tc.ID,
|
||||||
|
}
|
||||||
|
messages = append(messages, toolResultMsg)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if finalContent == "" {
|
||||||
|
finalContent = "I've completed processing but have no response to give."
|
||||||
|
emitEvent(bus.StreamEventContent, finalContent, nil)
|
||||||
|
}
|
||||||
|
|
||||||
|
emitEvent(bus.StreamEventComplete, "Processing complete", nil)
|
||||||
|
|
||||||
|
al.sessions.AddMessage(msg.SessionKey, "user", msg.Content)
|
||||||
|
al.sessions.AddMessage(msg.SessionKey, "assistant", finalContent)
|
||||||
|
|
||||||
|
// Context compression logic
|
||||||
|
newHistory := al.sessions.GetHistory(msg.SessionKey)
|
||||||
|
|
||||||
|
tokenEstimate := al.estimateTokens(newHistory)
|
||||||
|
threshold := al.contextWindow * 75 / 100
|
||||||
|
|
||||||
if len(newHistory) > 20 || tokenEstimate > threshold {
|
if len(newHistory) > 20 || tokenEstimate > threshold {
|
||||||
if _, loading := al.summarizing.LoadOrStore(msg.SessionKey, true); !loading {
|
if _, loading := al.summarizing.LoadOrStore(msg.SessionKey, true); !loading {
|
||||||
go func() {
|
go func() {
|
||||||
|
|
@ -267,7 +507,7 @@ func (al *AgentLoop) summarizeSession(sessionKey string) {
|
||||||
|
|
||||||
s1, _ := al.summarizeBatch(ctx, part1, "")
|
s1, _ := al.summarizeBatch(ctx, part1, "")
|
||||||
s2, _ := al.summarizeBatch(ctx, part2, "")
|
s2, _ := al.summarizeBatch(ctx, part2, "")
|
||||||
|
|
||||||
// Merge them
|
// Merge them
|
||||||
mergePrompt := fmt.Sprintf("Merge these two conversation summaries into one cohesive summary:\n\n1: %s\n\n2: %s", s1, s2)
|
mergePrompt := fmt.Sprintf("Merge these two conversation summaries into one cohesive summary:\n\n1: %s\n\n2: %s", s1, s2)
|
||||||
resp, err := al.provider.Chat(ctx, []providers.Message{{Role: "user", Content: mergePrompt}}, nil, al.model, map[string]interface{}{
|
resp, err := al.provider.Chat(ctx, []providers.Message{{Role: "user", Content: mergePrompt}}, nil, al.model, map[string]interface{}{
|
||||||
|
|
@ -321,4 +561,3 @@ func (al *AgentLoop) estimateTokens(messages []providers.Message) int {
|
||||||
}
|
}
|
||||||
return total
|
return total
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -3,20 +3,25 @@ package bus
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"sync"
|
"sync"
|
||||||
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
type MessageBus struct {
|
type MessageBus struct {
|
||||||
inbound chan InboundMessage
|
inbound chan InboundMessage
|
||||||
outbound chan OutboundMessage
|
outbound chan OutboundMessage
|
||||||
handlers map[string]MessageHandler
|
streamChan chan StreamEvent
|
||||||
mu sync.RWMutex
|
interrupt chan InterruptSignal
|
||||||
|
handlers map[string]MessageHandler
|
||||||
|
mu sync.RWMutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func NewMessageBus() *MessageBus {
|
func NewMessageBus() *MessageBus {
|
||||||
return &MessageBus{
|
return &MessageBus{
|
||||||
inbound: make(chan InboundMessage, 100),
|
inbound: make(chan InboundMessage, 100),
|
||||||
outbound: make(chan OutboundMessage, 100),
|
outbound: make(chan OutboundMessage, 100),
|
||||||
handlers: make(map[string]MessageHandler),
|
streamChan: make(chan StreamEvent, 100),
|
||||||
|
interrupt: make(chan InterruptSignal, 10),
|
||||||
|
handlers: make(map[string]MessageHandler),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -62,4 +67,52 @@ func (mb *MessageBus) GetHandler(channel string) (MessageHandler, bool) {
|
||||||
func (mb *MessageBus) Close() {
|
func (mb *MessageBus) Close() {
|
||||||
close(mb.inbound)
|
close(mb.inbound)
|
||||||
close(mb.outbound)
|
close(mb.outbound)
|
||||||
|
close(mb.streamChan)
|
||||||
|
close(mb.interrupt)
|
||||||
|
}
|
||||||
|
|
||||||
|
// PublishStreamEvent sends a streaming event
|
||||||
|
func (mb *MessageBus) PublishStreamEvent(event StreamEvent) {
|
||||||
|
if event.Timestamp == 0 {
|
||||||
|
event.Timestamp = time.Now().Unix()
|
||||||
|
}
|
||||||
|
select {
|
||||||
|
case mb.streamChan <- event:
|
||||||
|
default:
|
||||||
|
// Drop event if channel is full to avoid blocking
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ConsumeStreamEvent receives streaming events
|
||||||
|
func (mb *MessageBus) ConsumeStreamEvent(ctx context.Context) (StreamEvent, bool) {
|
||||||
|
select {
|
||||||
|
case event := <-mb.streamChan:
|
||||||
|
return event, true
|
||||||
|
case <-ctx.Done():
|
||||||
|
return StreamEvent{}, false
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// PublishInterrupt sends an interrupt signal
|
||||||
|
func (mb *MessageBus) PublishInterrupt(signal InterruptSignal) {
|
||||||
|
select {
|
||||||
|
case mb.interrupt <- signal:
|
||||||
|
default:
|
||||||
|
// Drop if channel is full
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// CheckInterrupt checks if there's an interrupt signal for the given session
|
||||||
|
func (mb *MessageBus) CheckInterrupt(sessionKey string) bool {
|
||||||
|
select {
|
||||||
|
case signal := <-mb.interrupt:
|
||||||
|
// Put it back if it doesn't match
|
||||||
|
if signal.SessionKey != sessionKey {
|
||||||
|
mb.interrupt <- signal
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
return true
|
||||||
|
default:
|
||||||
|
return false
|
||||||
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -17,3 +17,34 @@ type OutboundMessage struct {
|
||||||
}
|
}
|
||||||
|
|
||||||
type MessageHandler func(InboundMessage) error
|
type MessageHandler func(InboundMessage) error
|
||||||
|
|
||||||
|
// StreamEventType defines the type of streaming event
|
||||||
|
type StreamEventType string
|
||||||
|
|
||||||
|
const (
|
||||||
|
StreamEventThinking StreamEventType = "thinking" // AI is thinking/reasoning
|
||||||
|
StreamEventToolCall StreamEventType = "tool_call" // About to execute a tool
|
||||||
|
StreamEventToolResult StreamEventType = "tool_result" // Tool execution completed
|
||||||
|
StreamEventProgress StreamEventType = "progress" // Progress update
|
||||||
|
StreamEventContent StreamEventType = "content" // Partial content streaming
|
||||||
|
StreamEventComplete StreamEventType = "complete" // Processing complete
|
||||||
|
StreamEventError StreamEventType = "error" // Error occurred
|
||||||
|
)
|
||||||
|
|
||||||
|
// StreamEvent represents a streaming update during agent processing
|
||||||
|
type StreamEvent struct {
|
||||||
|
Type StreamEventType `json:"type"`
|
||||||
|
Channel string `json:"channel"`
|
||||||
|
ChatID string `json:"chat_id"`
|
||||||
|
SessionKey string `json:"session_key"`
|
||||||
|
Content string `json:"content,omitempty"`
|
||||||
|
Metadata map[string]interface{} `json:"metadata,omitempty"`
|
||||||
|
Timestamp int64 `json:"timestamp"`
|
||||||
|
}
|
||||||
|
|
||||||
|
// InterruptSignal represents a user interrupt request
|
||||||
|
type InterruptSignal struct {
|
||||||
|
Channel string `json:"channel"`
|
||||||
|
ChatID string `json:"chat_id"`
|
||||||
|
SessionKey string `json:"session_key"`
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -67,6 +67,9 @@ func (c *TelegramChannel) Start(ctx context.Context) error {
|
||||||
}
|
}
|
||||||
log.Printf("Telegram bot @%s connected", botInfo.UserName)
|
log.Printf("Telegram bot @%s connected", botInfo.UserName)
|
||||||
|
|
||||||
|
// Start stream event handler
|
||||||
|
go c.handleStreamEvents(ctx)
|
||||||
|
|
||||||
go func() {
|
go func() {
|
||||||
for {
|
for {
|
||||||
select {
|
select {
|
||||||
|
|
@ -245,37 +248,6 @@ func (c *TelegramChannel) handleMessage(update tgbotapi.Update) {
|
||||||
|
|
||||||
log.Printf("Telegram message from %s: %s...", senderID, truncateString(content, 50))
|
log.Printf("Telegram message from %s: %s...", senderID, truncateString(content, 50))
|
||||||
|
|
||||||
// Thinking indicator
|
|
||||||
c.bot.Send(tgbotapi.NewChatAction(chatID, tgbotapi.ChatTyping))
|
|
||||||
|
|
||||||
stopChan := make(chan struct{})
|
|
||||||
c.stopThinking.Store(fmt.Sprintf("%d", chatID), stopChan)
|
|
||||||
|
|
||||||
pMsg, err := c.bot.Send(tgbotapi.NewMessage(chatID, "Thinking... 💭"))
|
|
||||||
if err == nil {
|
|
||||||
pID := pMsg.MessageID
|
|
||||||
c.placeholders.Store(fmt.Sprintf("%d", chatID), pID)
|
|
||||||
|
|
||||||
go func(cid int64, mid int, stop <-chan struct{}) {
|
|
||||||
dots := []string{".", "..", "..."}
|
|
||||||
emotes := []string{"💭", "🤔", "☁️"}
|
|
||||||
i := 0
|
|
||||||
ticker := time.NewTicker(2000 * time.Millisecond)
|
|
||||||
defer ticker.Stop()
|
|
||||||
for {
|
|
||||||
select {
|
|
||||||
case <-stop:
|
|
||||||
return
|
|
||||||
case <-ticker.C:
|
|
||||||
i++
|
|
||||||
text := fmt.Sprintf("Thinking%s %s", dots[i%len(dots)], emotes[i%len(emotes)])
|
|
||||||
edit := tgbotapi.NewEditMessageText(cid, mid, text)
|
|
||||||
c.bot.Send(edit)
|
|
||||||
}
|
|
||||||
}
|
|
||||||
}(chatID, pID, stopChan)
|
|
||||||
}
|
|
||||||
|
|
||||||
metadata := map[string]string{
|
metadata := map[string]string{
|
||||||
"message_id": fmt.Sprintf("%d", message.MessageID),
|
"message_id": fmt.Sprintf("%d", message.MessageID),
|
||||||
"user_id": fmt.Sprintf("%d", user.ID),
|
"user_id": fmt.Sprintf("%d", user.ID),
|
||||||
|
|
@ -284,7 +256,21 @@ func (c *TelegramChannel) handleMessage(update tgbotapi.Update) {
|
||||||
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
"is_group": fmt.Sprintf("%t", message.Chat.Type != "private"),
|
||||||
}
|
}
|
||||||
|
|
||||||
c.HandleMessage(senderID, fmt.Sprintf("%d", chatID), content, mediaPaths, metadata)
|
// Create session key
|
||||||
|
sessionKey := fmt.Sprintf("telegram:%d", chatID)
|
||||||
|
|
||||||
|
// Send inbound message with session key
|
||||||
|
msg := bus.InboundMessage{
|
||||||
|
Channel: c.Name(),
|
||||||
|
SenderID: senderID,
|
||||||
|
ChatID: fmt.Sprintf("%d", chatID),
|
||||||
|
Content: content,
|
||||||
|
Media: mediaPaths,
|
||||||
|
Metadata: metadata,
|
||||||
|
SessionKey: sessionKey,
|
||||||
|
}
|
||||||
|
|
||||||
|
c.BaseChannel.bus.PublishInbound(msg)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (c *TelegramChannel) downloadPhoto(fileID string) string {
|
func (c *TelegramChannel) downloadPhoto(fileID string) string {
|
||||||
|
|
@ -446,3 +432,101 @@ func escapeHTML(text string) string {
|
||||||
text = strings.ReplaceAll(text, ">", ">")
|
text = strings.ReplaceAll(text, ">", ">")
|
||||||
return text
|
return text
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// handleStreamEvents listens for streaming events and updates messages in real-time
|
||||||
|
func (c *TelegramChannel) handleStreamEvents(ctx context.Context) {
|
||||||
|
// Track current status message for each chat
|
||||||
|
statusMessages := sync.Map{} // chatID -> messageID
|
||||||
|
lastUpdate := sync.Map{} // chatID -> time.Time
|
||||||
|
|
||||||
|
for {
|
||||||
|
select {
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
default:
|
||||||
|
}
|
||||||
|
|
||||||
|
event, ok := c.BaseChannel.bus.ConsumeStreamEvent(ctx)
|
||||||
|
if !ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only handle events for telegram channel
|
||||||
|
if event.Channel != "telegram" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
chatID, err := parseChatID(event.ChatID)
|
||||||
|
if err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Rate limiting: update at most once per second
|
||||||
|
if lastTime, ok := lastUpdate.Load(event.ChatID); ok {
|
||||||
|
if time.Since(lastTime.(time.Time)) < time.Second {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
}
|
||||||
|
lastUpdate.Store(event.ChatID, time.Now())
|
||||||
|
|
||||||
|
var statusText string
|
||||||
|
var shouldDelete bool
|
||||||
|
|
||||||
|
switch event.Type {
|
||||||
|
case bus.StreamEventThinking:
|
||||||
|
statusText = "💭 " + event.Content
|
||||||
|
case bus.StreamEventProgress:
|
||||||
|
statusText = "⏳ " + event.Content
|
||||||
|
case bus.StreamEventToolCall:
|
||||||
|
if toolName, ok := event.Metadata["tool_name"].(string); ok {
|
||||||
|
statusText = fmt.Sprintf("🔧 Executing tool: %s", toolName)
|
||||||
|
} else {
|
||||||
|
statusText = "🔧 " + event.Content
|
||||||
|
}
|
||||||
|
case bus.StreamEventToolResult:
|
||||||
|
if toolName, ok := event.Metadata["tool_name"].(string); ok {
|
||||||
|
statusText = fmt.Sprintf("✓ Tool %s completed", toolName)
|
||||||
|
} else {
|
||||||
|
statusText = "✓ " + event.Content
|
||||||
|
}
|
||||||
|
case bus.StreamEventComplete, bus.StreamEventContent:
|
||||||
|
// Delete status message on completion
|
||||||
|
shouldDelete = true
|
||||||
|
case bus.StreamEventError:
|
||||||
|
statusText = "❌ " + event.Content
|
||||||
|
}
|
||||||
|
|
||||||
|
if shouldDelete {
|
||||||
|
// Delete the status message
|
||||||
|
if msgID, ok := statusMessages.LoadAndDelete(event.ChatID); ok {
|
||||||
|
deleteMsg := tgbotapi.NewDeleteMessage(chatID, msgID.(int))
|
||||||
|
c.bot.Send(deleteMsg)
|
||||||
|
}
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if statusText == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// Update or create status message
|
||||||
|
if msgID, ok := statusMessages.Load(event.ChatID); ok {
|
||||||
|
// Update existing message
|
||||||
|
editMsg := tgbotapi.NewEditMessageText(chatID, msgID.(int), statusText)
|
||||||
|
if _, err := c.bot.Send(editMsg); err != nil {
|
||||||
|
// If edit fails, delete and create new
|
||||||
|
deleteMsg := tgbotapi.NewDeleteMessage(chatID, msgID.(int))
|
||||||
|
c.bot.Send(deleteMsg)
|
||||||
|
statusMessages.Delete(event.ChatID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Create new status message if none exists
|
||||||
|
if _, ok := statusMessages.Load(event.ChatID); !ok {
|
||||||
|
newMsg, err := c.bot.Send(tgbotapi.NewMessage(chatID, statusText))
|
||||||
|
if err == nil {
|
||||||
|
statusMessages.Store(event.ChatID, newMsg.MessageID)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
|
||||||
|
|
@ -28,6 +28,7 @@ type AgentDefaults struct {
|
||||||
MaxTokens int `json:"max_tokens" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOKENS"`
|
MaxTokens int `json:"max_tokens" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOKENS"`
|
||||||
Temperature float64 `json:"temperature" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"`
|
Temperature float64 `json:"temperature" env:"PICOCLAW_AGENTS_DEFAULTS_TEMPERATURE"`
|
||||||
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"`
|
MaxToolIterations int `json:"max_tool_iterations" env:"PICOCLAW_AGENTS_DEFAULTS_MAX_TOOL_ITERATIONS"`
|
||||||
|
Streaming bool `json:"streaming" env:"PICOCLAW_AGENTS_DEFAULTS_STREAMING"`
|
||||||
}
|
}
|
||||||
|
|
||||||
type ChannelsConfig struct {
|
type ChannelsConfig struct {
|
||||||
|
|
@ -114,6 +115,7 @@ func DefaultConfig() *Config {
|
||||||
MaxTokens: 8192,
|
MaxTokens: 8192,
|
||||||
Temperature: 0.7,
|
Temperature: 0.7,
|
||||||
MaxToolIterations: 20,
|
MaxToolIterations: 20,
|
||||||
|
Streaming: false,
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
Channels: ChannelsConfig{
|
Channels: ChannelsConfig{
|
||||||
|
|
|
||||||
|
|
@ -7,6 +7,7 @@
|
||||||
package providers
|
package providers
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"bufio"
|
||||||
"bytes"
|
"bytes"
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
|
|
@ -162,7 +163,173 @@ func (p *HTTPProvider) parseResponse(body []byte) (*LLMResponse, error) {
|
||||||
}
|
}
|
||||||
|
|
||||||
func (p *HTTPProvider) GetDefaultModel() string {
|
func (p *HTTPProvider) GetDefaultModel() string {
|
||||||
return ""
|
return "gpt-3.5-turbo"
|
||||||
|
}
|
||||||
|
|
||||||
|
// ChatStream performs a streaming chat request with real-time token generation
|
||||||
|
func (p *HTTPProvider) ChatStream(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}, callback func(chunk string)) (*LLMResponse, error) {
|
||||||
|
if p.apiBase == "" {
|
||||||
|
return nil, fmt.Errorf("API base not configured")
|
||||||
|
}
|
||||||
|
|
||||||
|
requestBody := map[string]interface{}{
|
||||||
|
"model": model,
|
||||||
|
"messages": messages,
|
||||||
|
"stream": true, // Enable streaming
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(tools) > 0 {
|
||||||
|
requestBody["tools"] = tools
|
||||||
|
requestBody["tool_choice"] = "auto"
|
||||||
|
}
|
||||||
|
|
||||||
|
if maxTokens, ok := options["max_tokens"].(int); ok {
|
||||||
|
requestBody["max_tokens"] = maxTokens
|
||||||
|
}
|
||||||
|
|
||||||
|
if temperature, ok := options["temperature"].(float64); ok {
|
||||||
|
requestBody["temperature"] = temperature
|
||||||
|
}
|
||||||
|
|
||||||
|
jsonData, err := json.Marshal(requestBody)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req, err := http.NewRequestWithContext(ctx, "POST", p.apiBase+"/chat/completions", bytes.NewReader(jsonData))
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
if p.apiKey != "" {
|
||||||
|
authHeader := "Bearer " + p.apiKey
|
||||||
|
req.Header.Set("Authorization", authHeader)
|
||||||
|
}
|
||||||
|
|
||||||
|
resp, err := p.httpClient.Do(req)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("failed to send request: %w", err)
|
||||||
|
}
|
||||||
|
defer resp.Body.Close()
|
||||||
|
|
||||||
|
if resp.StatusCode != http.StatusOK {
|
||||||
|
body, _ := io.ReadAll(resp.Body)
|
||||||
|
return nil, fmt.Errorf("API error: %s", string(body))
|
||||||
|
}
|
||||||
|
|
||||||
|
// Parse streaming response
|
||||||
|
return p.parseStreamingResponse(resp.Body, callback)
|
||||||
|
}
|
||||||
|
|
||||||
|
func (p *HTTPProvider) parseStreamingResponse(body io.Reader, callback func(chunk string)) (*LLMResponse, error) {
|
||||||
|
scanner := bufio.NewScanner(body)
|
||||||
|
response := &LLMResponse{
|
||||||
|
Content: "",
|
||||||
|
ToolCalls: []ToolCall{},
|
||||||
|
}
|
||||||
|
|
||||||
|
var contentBuilder strings.Builder
|
||||||
|
toolCallsMap := make(map[int]*ToolCall)
|
||||||
|
|
||||||
|
for scanner.Scan() {
|
||||||
|
line := scanner.Text()
|
||||||
|
|
||||||
|
// Skip empty lines
|
||||||
|
if line == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
// SSE format: "data: {json}"
|
||||||
|
if !strings.HasPrefix(line, "data: ") {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
data := strings.TrimPrefix(line, "data: ")
|
||||||
|
|
||||||
|
// Check for stream end
|
||||||
|
if data == "[DONE]" {
|
||||||
|
break
|
||||||
|
}
|
||||||
|
|
||||||
|
var chunk struct {
|
||||||
|
Choices []struct {
|
||||||
|
Delta struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
ToolCalls []struct {
|
||||||
|
Index int `json:"index"`
|
||||||
|
ID string `json:"id"`
|
||||||
|
Type string `json:"type"`
|
||||||
|
Function struct {
|
||||||
|
Name string `json:"name"`
|
||||||
|
Arguments string `json:"arguments"`
|
||||||
|
} `json:"function"`
|
||||||
|
} `json:"tool_calls"`
|
||||||
|
} `json:"delta"`
|
||||||
|
FinishReason string `json:"finish_reason"`
|
||||||
|
} `json:"choices"`
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := json.Unmarshal([]byte(data), &chunk); err != nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
if len(chunk.Choices) == 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
choice := chunk.Choices[0]
|
||||||
|
|
||||||
|
// Handle content streaming
|
||||||
|
if choice.Delta.Content != "" {
|
||||||
|
contentBuilder.WriteString(choice.Delta.Content)
|
||||||
|
if callback != nil {
|
||||||
|
callback(choice.Delta.Content)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Handle tool calls
|
||||||
|
for _, tc := range choice.Delta.ToolCalls {
|
||||||
|
if _, exists := toolCallsMap[tc.Index]; !exists {
|
||||||
|
toolCallsMap[tc.Index] = &ToolCall{
|
||||||
|
ID: tc.ID,
|
||||||
|
Type: tc.Type,
|
||||||
|
Name: tc.Function.Name,
|
||||||
|
Arguments: make(map[string]interface{}),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Accumulate function arguments
|
||||||
|
if tc.Function.Arguments != "" {
|
||||||
|
existing := toolCallsMap[tc.Index]
|
||||||
|
// Parse and merge arguments incrementally
|
||||||
|
var args map[string]interface{}
|
||||||
|
if err := json.Unmarshal([]byte(tc.Function.Arguments), &args); err == nil {
|
||||||
|
for k, v := range args {
|
||||||
|
existing.Arguments[k] = v
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// Set finish reason
|
||||||
|
if choice.FinishReason != "" {
|
||||||
|
response.FinishReason = choice.FinishReason
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
if err := scanner.Err(); err != nil {
|
||||||
|
return nil, fmt.Errorf("error reading stream: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
response.Content = contentBuilder.String()
|
||||||
|
|
||||||
|
// Convert tool calls map to slice
|
||||||
|
for _, tc := range toolCallsMap {
|
||||||
|
response.ToolCalls = append(response.ToolCalls, *tc)
|
||||||
|
}
|
||||||
|
|
||||||
|
return response, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func CreateProvider(cfg *config.Config) (LLMProvider, error) {
|
func CreateProvider(cfg *config.Config) (LLMProvider, error) {
|
||||||
|
|
|
||||||
|
|
@ -37,6 +37,7 @@ type Message struct {
|
||||||
|
|
||||||
type LLMProvider interface {
|
type LLMProvider interface {
|
||||||
Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error)
|
Chat(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}) (*LLMResponse, error)
|
||||||
|
ChatStream(ctx context.Context, messages []Message, tools []ToolDefinition, model string, options map[string]interface{}, callback func(chunk string)) (*LLMResponse, error)
|
||||||
GetDefaultModel() string
|
GetDefaultModel() string
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue