Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
227 changes: 227 additions & 0 deletions examples/openrouter/main.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,227 @@
package main

import (
"context"
"fmt"
"log/slog"
"os"
"strings"
"time"

"github.com/astercloud/aster/pkg/agent"
"github.com/astercloud/aster/pkg/provider"
"github.com/astercloud/aster/pkg/sandbox"
"github.com/astercloud/aster/pkg/store"
"github.com/astercloud/aster/pkg/tools"
"github.com/astercloud/aster/pkg/tools/builtin"
"github.com/astercloud/aster/pkg/types"
)

// TestSuite 测试套件
type TestSuite struct {
cases []testCase
startTime time.Time
}

type testCase struct {
name string
passed bool
err string
duration time.Duration
}

func newTestSuite() *TestSuite {
return &TestSuite{startTime: time.Now()}
}

func (ts *TestSuite) add(name string, err error, duration time.Duration) {
tc := testCase{name: name, duration: duration, passed: err == nil}
if err != nil {
tc.err = err.Error()
}
ts.cases = append(ts.cases, tc)
}

func (ts *TestSuite) summary() bool {
fmt.Println("\n" + strings.Repeat("=", 60))
fmt.Println("📊 测试结果总结")
fmt.Println(strings.Repeat("=", 60))

passed, failed := 0, 0
for _, tc := range ts.cases {
if tc.passed {
fmt.Printf(" ✅ PASS %-30s (%v)\n", tc.name, tc.duration.Round(time.Millisecond))
passed++
} else {
fmt.Printf(" ❌ FAIL %-30s (%v)\n", tc.name, tc.duration.Round(time.Millisecond))
fmt.Printf(" └─ %s\n", tc.err)
failed++
}
}

fmt.Println(strings.Repeat("-", 60))
fmt.Printf(" 总计: %d 通过, %d 失败, 耗时 %v\n", passed, failed, time.Since(ts.startTime).Round(time.Millisecond))
if failed == 0 {
fmt.Println(" 🎉 所有测试通过!")
}
fmt.Println(strings.Repeat("=", 60))
return failed == 0
}

// runChatTest 执行对话测试并验证
func runChatTest(ctx context.Context, ag *agent.Agent, suite *TestSuite, name, prompt string, verify func() error) {
slog.Info("--- " + name + " ---")
start := time.Now()

result, err := ag.Chat(ctx, prompt)
if err != nil {
suite.add(name, err, time.Since(start))
return
}
if result == nil || result.Status != "ok" {
suite.add(name, fmt.Errorf("状态异常: %v", result), time.Since(start))
return
}

// 执行额外验证
if verify != nil {
time.Sleep(300 * time.Millisecond) // 等待文件操作完成
if err := verify(); err != nil {
suite.add(name, err, time.Since(start))
return
}
}

suite.add(name, nil, time.Since(start))
slog.Info("响应", "status", result.Status)
}

func main() {
// 检查 API Key
apiKey := os.Getenv("OPENROUTER_API_KEY")
if apiKey == "" {
slog.Error("需要设置 OPENROUTER_API_KEY 环境变量")
os.Exit(1)
}

baseURL := os.Getenv("OPENROUTER_BASE_URL")
if baseURL == "" {
baseURL = "https://openrouter.ai/api/v1"
}

ctx := context.Background()

// 创建依赖
toolRegistry := tools.NewRegistry()
builtin.RegisterAll(toolRegistry)

jsonStore, err := store.NewJSONStore(".aster")
if err != nil {
slog.Error("创建 Store 失败", "error", err)
os.Exit(1)
}

templateRegistry := agent.NewTemplateRegistry()
templateRegistry.Register(&types.AgentTemplateDefinition{
ID: "simple-assistant",
SystemPrompt: "你是一个有用的助手,可以读取和写入文件。当用户要求你读取或写入文件时,请使用可用的工具。",
Tools: []interface{}{"Read", "Write", "Bash"},
})

deps := &agent.Dependencies{
Store: jsonStore,
SandboxFactory: sandbox.NewFactory(),
ToolRegistry: toolRegistry,
ProviderFactory: &provider.OpenRouterFactory{},
TemplateRegistry: templateRegistry,
}

config := &types.AgentConfig{
TemplateID: "simple-assistant",
ModelConfig: &types.ModelConfig{
Provider: "openrouter",
Model: "anthropic/claude-haiku-4.5",
APIKey: apiKey,
BaseURL: baseURL,
ExecutionMode: types.ExecutionModeNonStreaming,
},
Sandbox: &types.SandboxConfig{
Kind: types.SandboxKindLocal,
WorkDir: "./workspace",
},
}

ag, err := agent.Create(ctx, config, deps)
if err != nil {
slog.Error("创建 Agent 失败", "error", err)
os.Exit(1)
}
defer func() { _ = ag.Close() }()

slog.Info("Agent 创建成功", "id", ag.ID())

// 准备测试
suite := newTestSuite()
testFile := "./workspace/test.txt"
_ = os.Remove(testFile)

// 订阅事件
eventCh := ag.Subscribe([]types.AgentChannel{types.ChannelProgress, types.ChannelMonitor}, nil)
go func() {
for envelope := range eventCh {
handleEvent(envelope.Event)
}
}()

// 执行测试
runChatTest(ctx, ag, suite, "创建测试文件",
"请创建一个名为 test.txt 的文件,内容为 'Hello World'",
func() error {
data, err := os.ReadFile(testFile)
if err != nil {
return fmt.Errorf("文件未创建: %w", err)
}
if strings.TrimSpace(string(data)) != "Hello World" {
return fmt.Errorf("内容不匹配: %s", string(data))
}
return nil
})

runChatTest(ctx, ag, suite, "读取文件内容",
"请读取 test.txt 文件的内容", nil)

runChatTest(ctx, ag, suite, "执行 Bash 命令",
"请执行 'ls -la' 命令", nil)

// 验证 Agent 状态
status := ag.Status()
if status.State != types.AgentStateReady || status.StepCount == 0 {
suite.add("Agent 状态检查", fmt.Errorf("状态: %s, 步骤: %d", status.State, status.StepCount), 0)
} else {
suite.add("Agent 状态检查", nil, 0)
}

// 输出结果
fmt.Printf("\n最终状态: Agent=%s, 状态=%s, 步骤=%d\n", status.AgentID, status.State, status.StepCount)

if !suite.summary() {
os.Exit(1)
}
}

func handleEvent(event interface{}) {
switch e := event.(type) {
case *types.ProgressToolStartEvent:
fmt.Printf("\n[工具] %s 开始\n", e.Call.Name)
case *types.ProgressToolEndEvent:
fmt.Printf("[工具] %s 完成\n", e.Call.Name)
case *types.ProgressToolErrorEvent:
fmt.Printf("[错误] %s: %s\n", e.Call.Name, e.Error)
case *types.ProgressDoneEvent:
fmt.Printf("[完成] 步骤 %d\n", e.Step)
case *types.MonitorStateChangedEvent:
fmt.Printf("[状态] %s\n", e.State)
case *types.MonitorErrorEvent:
fmt.Printf("[错误] %s: %s\n", e.Phase, e.Message)
}
}
2 changes: 2 additions & 0 deletions examples/providers/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -132,6 +132,8 @@ func main() {
}
```

**完整示例**:查看 [`examples/openrouter/`](../openrouter/) 获取包含工具调用验证的完整集成测试示例。

### 5. 多模态输入(图片)

```go
Expand Down
87 changes: 81 additions & 6 deletions pkg/provider/openai_compatible.go
Original file line number Diff line number Diff line change
Expand Up @@ -268,20 +268,95 @@ func (p *OpenAICompatibleProvider) convertMessages(messages []types.Message) []m
continue
}

// 检查 ContentBlocks 中是否包含特殊块类型
var toolUseBlocks []*types.ToolUseBlock
var toolResultBlocks []*types.ToolResultBlock
var otherBlocks []types.ContentBlock

for _, block := range msg.ContentBlocks {
switch b := block.(type) {
case *types.ToolUseBlock:
toolUseBlocks = append(toolUseBlocks, b)
case *types.ToolResultBlock:
toolResultBlocks = append(toolResultBlocks, b)
default:
otherBlocks = append(otherBlocks, block)
}
}

// 处理包含 ToolResultBlock 的消息
// OpenAI API 要求每个工具结果作为独立的 role: "tool" 消息
if len(toolResultBlocks) > 0 {
for _, tr := range toolResultBlocks {
toolMsg := map[string]interface{}{
"role": "tool",
"tool_call_id": tr.ToolUseID,
"content": tr.Content,
}
result = append(result, toolMsg)
}
continue // 跳过常规消息处理
}

msgMap := map[string]interface{}{
"role": string(msg.Role),
}

// 处理内容
if len(msg.ContentBlocks) > 0 {
// 复杂格式:ContentBlocks
content := p.convertContentBlocks(msg.ContentBlocks)
// 处理内容(排除 ToolUseBlock,它们会被单独处理)
if len(otherBlocks) > 0 {
content := p.convertContentBlocks(otherBlocks)
msgMap["content"] = content
} else {
// 简单格式:纯文本
} else if msg.Content != "" {
msgMap["content"] = msg.Content
}

// 处理 assistant 消息的工具调用
// 优先从 ToolCalls 字段获取,如果没有则从 ContentBlocks 中的 ToolUseBlock 提取
if msg.Role == types.RoleAssistant {
var toolCalls []map[string]interface{}

// 从 msg.ToolCalls 获取
for _, tc := range msg.ToolCalls {
argsJSON, _ := json.Marshal(tc.Arguments)
toolCalls = append(toolCalls, map[string]interface{}{
"id": tc.ID,
"type": "function",
"function": map[string]interface{}{
"name": tc.Name,
"arguments": string(argsJSON),
},
})
}

// 从 ContentBlocks 中的 ToolUseBlock 提取
for _, tu := range toolUseBlocks {
argsJSON, _ := json.Marshal(tu.Input)
toolCalls = append(toolCalls, map[string]interface{}{
"id": tu.ID,
"type": "function",
"function": map[string]interface{}{
"name": tu.Name,
"arguments": string(argsJSON),
},
})
}

if len(toolCalls) > 0 {
msgMap["tool_calls"] = toolCalls

// 注意:OpenRouter 转发到 Anthropic 时,要求非最后一条 assistant 消息必须有 content
// 因此这里不删除 content 字段,即使为空也保留空字符串
if msgMap["content"] == nil {
msgMap["content"] = ""
}
}
}

// 处理 tool 消息的工具调用 ID
if msg.Role == types.RoleTool && msg.ToolCallID != "" {
msgMap["tool_call_id"] = msg.ToolCallID
}

result = append(result, msgMap)
}

Expand Down
Loading