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
73 changes: 44 additions & 29 deletions agent/agent.go
Original file line number Diff line number Diff line change
Expand Up @@ -13,35 +13,62 @@ import (
"github.com/lsongdev/miya-agents/tools"
)

type LLM interface {
CreateChatCompletionStream(context.Context, *openai.ChatCompletionRequest) (<-chan openai.ChatCompletionResponse, error)
}
// StreamFunc is the only model capability required by the agent loop.
type StreamFunc func(context.Context, *openai.ChatCompletionRequest) (<-chan openai.ChatCompletionResponse, error)

type Agent struct {
Name string
Config *config.ProfileConfig
LLM LLM
// tools
toolsMap map[string]openai.Tool
toolsDefs []openai.ToolDef
Stream StreamFunc
tools []openai.Tool
}

func New(name string, cfg *config.ProfileConfig, stream StreamFunc) *Agent {
return &Agent{Name: name, Config: cfg, Stream: stream}
}

// Use adds tools to the agent and returns it for chaining.
func (a *Agent) Use(tools ...openai.Tool) *Agent {
a.tools = append(a.tools, tools...)
return a
}

func (a *Agent) tool(name string) (openai.Tool, bool) {
for _, tool := range a.tools {
if tool.Def().Function.Name == name {
return tool, true
}
}
return nil, false
}

func (a *Agent) toolDefs() []openai.ToolDef {
defs := make([]openai.ToolDef, len(a.tools))
for i, tool := range a.tools {
defs[i] = tool.Def()
}
return defs
}

func (a *Agent) RunAgentLoop(ctx context.Context, sess *session.Session, sink EventSink) error {
if a.Stream == nil {
return fmt.Errorf("agent has no model stream")
}
for {
req := openai.ChatCompletionRequest{
Model: a.Config.ModelName,
Messages: sess.Messages,
Tools: a.toolsDefs,
Tools: a.toolDefs(),
Stream: true,
}
resp, err := a.LLM.CreateChatCompletionStream(ctx, &req)
resp, err := a.Stream(ctx, &req)
if err != nil {
return fmt.Errorf("failed to create chat completion stream: %w", err)
return fmt.Errorf("create model stream: %w", err)
}
builder := openai.NewMessageBuilder()
for chunk := range resp {
if chunk.Error != nil {
return fmt.Errorf("API error: %s", chunk.Error.Message)
return fmt.Errorf("model stream: %s", chunk.Error.Message)
}
m := chunk.GetMessage()
if m == nil {
Expand All @@ -64,10 +91,9 @@ func (a *Agent) RunAgentLoop(ctx context.Context, sess *session.Session, sink Ev
}
respMessage := builder.Build()
if respMessage.IsEmpty() {
return fmt.Errorf("chat completion stream closed without a response")
return fmt.Errorf("model stream closed without a response")
}
sess.AppendResponse(respMessage)
// finish
if !respMessage.HasToolCall() {
if err := sink.Usage(UsageEvent{}); err != nil {
return err
Expand All @@ -81,12 +107,12 @@ func (a *Agent) RunAgentLoop(ctx context.Context, sess *session.Session, sink Ev
}
return nil
}
// Execute tool calls

for _, tc := range respMessage.ToolCalls {
if tc.ID == "" {
return fmt.Errorf("tool call %q is missing an id", tc.Function.Name)
}
tool, ok := a.toolsMap[tc.Function.Name]
tool, ok := a.tool(tc.Function.Name)
if err := sink.ToolCallStart(ToolCallEvent{
ID: tc.ID,
Name: tc.Function.Name,
Expand Down Expand Up @@ -156,12 +182,6 @@ func emitAttachedFileResult(sink EventSink, result string) (string, bool, error)
return fmt.Sprintf("Attached %s (%s, %d bytes) as %s.", attachment.Name, attachment.MimeType, attachment.Size, attachment.URI), true, nil
}

func (a *Agent) AddTool(tool openai.Tool) {
d := tool.Def()
a.toolsMap[d.Function.Name] = tool
a.toolsDefs = append(a.toolsDefs, d)
}

func (a *Agent) NewSession() *session.Session {
s := session.New(a.Name)
prompt := a.readSystemPrompt()
Expand Down Expand Up @@ -196,7 +216,7 @@ func (a *Agent) BuildTools() {
if workspace != "" {
_ = os.MkdirAll(workspace, 0755)
}
var tools = []openai.Tool{
a.Use(
&tools.WebFetchTool{},
&tools.WebSearchTool{},
&tools.ReadFileTool{Workspace: workspace},
Expand All @@ -209,11 +229,6 @@ func (a *Agent) BuildTools() {
DefaultTimeout: tools.ExecDefaultTimeoutSeconds,
RestrictToWorkspace: true,
},
&tools.SkillsTool{
Workspace: filepath.Join(config.ConfigPath, "skills"),
},
}
for _, t := range tools {
a.AddTool(t)
}
&tools.SkillsTool{Workspace: filepath.Join(config.ConfigPath, "skills")},
)
}
52 changes: 19 additions & 33 deletions agent/agent_loop_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,17 +12,15 @@ import (
"github.com/lsongdev/miya-agents/session"
)

type fakeLLM struct {
chunks []openai.ChatCompletionResponse
}

func (m *fakeLLM) CreateChatCompletionStream(context.Context, *openai.ChatCompletionRequest) (<-chan openai.ChatCompletionResponse, error) {
ch := make(chan openai.ChatCompletionResponse, len(m.chunks))
for _, chunk := range m.chunks {
ch <- chunk
func fakeStream(chunks ...openai.ChatCompletionResponse) StreamFunc {
return func(context.Context, *openai.ChatCompletionRequest) (<-chan openai.ChatCompletionResponse, error) {
ch := make(chan openai.ChatCompletionResponse, len(chunks))
for _, chunk := range chunks {
ch <- chunk
}
close(ch)
return ch, nil
}
close(ch)
return ch, nil
}

type discardSink struct{}
Expand All @@ -37,10 +35,7 @@ func (discardSink) Usage(UsageEvent) error { return nil }
func (discardSink) Done() error { return nil }

func TestRunAgentLoopRejectsEmptyStream(t *testing.T) {
ag := &Agent{
Config: &config.ProfileConfig{ModelName: "test"},
LLM: &fakeLLM{},
}
ag := New("test", &config.ProfileConfig{ModelName: "test"}, fakeStream())

err := ag.RunAgentLoop(context.Background(), session.New("test"), discardSink{})
if err == nil || !strings.Contains(err.Error(), "closed without a response") {
Expand All @@ -50,13 +45,10 @@ func TestRunAgentLoopRejectsEmptyStream(t *testing.T) {

func TestRunAgentLoopRejectsInterruptedStream(t *testing.T) {
message := openai.ChatCompletionMessage{Role: openai.RoleAssistant, Content: "partial"}
ag := &Agent{
Config: &config.ProfileConfig{ModelName: "test"},
LLM: &fakeLLM{chunks: []openai.ChatCompletionResponse{
{Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}},
{Error: &openai.Error{Type: "stream_error", Message: "connection reset"}},
}},
}
ag := New("test", &config.ProfileConfig{ModelName: "test"}, fakeStream(
openai.ChatCompletionResponse{Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}},
openai.ChatCompletionResponse{Error: &openai.Error{Type: "stream_error", Message: "connection reset"}},
))
sess := session.New("test")

err := ag.RunAgentLoop(context.Background(), sess, discardSink{})
Expand All @@ -75,12 +67,9 @@ func TestRunAgentLoopRejectsToolCallWithoutID(t *testing.T) {
Function: openai.FunctionCall{Name: "read_file", Arguments: `{}`},
}},
}
ag := &Agent{
Config: &config.ProfileConfig{ModelName: "test"},
LLM: &fakeLLM{chunks: []openai.ChatCompletionResponse{{
Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}},
}}},
}
ag := New("test", &config.ProfileConfig{ModelName: "test"}, fakeStream(
openai.ChatCompletionResponse{Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}},
))

err := ag.RunAgentLoop(context.Background(), session.New("test"), discardSink{})
if err == nil || !strings.Contains(err.Error(), "missing an id") {
Expand All @@ -98,12 +87,9 @@ func TestRunAgentLoopReturnsSaveError(t *testing.T) {
t.Cleanup(func() { config.ConfigPath = oldConfigPath })

message := openai.ChatCompletionMessage{Role: openai.RoleAssistant, Content: "done"}
ag := &Agent{
Config: &config.ProfileConfig{ModelName: "test"},
LLM: &fakeLLM{chunks: []openai.ChatCompletionResponse{{
Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}},
}}},
}
ag := New("test", &config.ProfileConfig{ModelName: "test"}, fakeStream(
openai.ChatCompletionResponse{Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}},
))

err := ag.RunAgentLoop(context.Background(), session.New("test"), discardSink{})
if err == nil || !strings.Contains(err.Error(), "save session") {
Expand Down
48 changes: 25 additions & 23 deletions agent/manager.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import (
"time"

"github.com/lsongdev/miya-agents/acp"
"github.com/lsongdev/miya-agents/anthropic"
"github.com/lsongdev/miya-agents/config"
"github.com/lsongdev/miya-agents/openai"
"github.com/lsongdev/miya-agents/session"
Expand All @@ -30,38 +31,39 @@ func NewAgentManager(config *config.Config) *Manager {
}
}

func (m *Manager) UseAgent(name string) (a *Agent, err error) {
ac, ok := m.config.Profiles[name]
func (m *Manager) UseAgent(name string) (*Agent, error) {
profile, ok := m.config.Profiles[name]
if !ok {
err = fmt.Errorf("agent not found: %s", name)
return
return nil, fmt.Errorf("agent not found: %s", name)
}
pc, ok := m.config.Providers[ac.Provider]
provider, ok := m.config.Providers[profile.Provider]
if !ok {
err = fmt.Errorf("provider not found: %s", ac.Provider)
return
return nil, fmt.Errorf("provider not found: %s", profile.Provider)
}
llm, err := openai.NewClient(&openai.Configuration{
API: pc.APIBase,
APIKey: pc.APIKey,
})
if err != nil {
return
}
a = &Agent{
Name: name,
LLM: llm,
Config: ac,
toolsMap: make(map[string]openai.Tool),
toolsDefs: []openai.ToolDef{},

var stream StreamFunc
switch strings.ToLower(strings.TrimSpace(provider.Type)) {
case "", "openai":
client, err := openai.NewClient(&openai.Configuration{API: provider.APIBase, APIKey: provider.APIKey})
if err != nil {
return nil, err
}
stream = client.CreateChatCompletionStream
case "anthropic":
client := anthropic.NewClient(&anthropic.Configuration{API: provider.APIBase, APIKey: provider.APIKey})
stream = client.CreateChatCompletionStream
default:
return nil, fmt.Errorf("unsupported provider type %q", provider.Type)
}

a := New(name, profile, stream)
a.BuildTools()
mcpManager := tools.NewMcpManager(m.config.McpServers)
for _, tool := range mcpManager.Tools {
a.AddTool(tool)
a.Use(tool)
}
a.AddTool(tools.NewSubagentTool(m))
return
a.Use(tools.NewSubagentTool(m))
return a, nil
}

func (m *Manager) defaultAgentName() (string, error) {
Expand Down
58 changes: 56 additions & 2 deletions agent/manager_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,6 +11,7 @@ import (
"github.com/lsongdev/miya-agents/acp"
"github.com/lsongdev/miya-agents/config"
"github.com/lsongdev/miya-agents/mcp"
"github.com/lsongdev/miya-agents/openai"
"github.com/lsongdev/miya-agents/session"
)

Expand Down Expand Up @@ -164,8 +165,61 @@ func TestUseAgentIncludesConfiguredMCPTools(t *testing.T) {
if err != nil {
t.Fatalf("UseAgent: %v", err)
}
if _, ok := ag.toolsMap["mcp_coffee_queryShopList"]; !ok {
t.Fatalf("missing MCP tool; tools = %#v", ag.toolsMap)
if _, ok := ag.tool("mcp_coffee_queryShopList"); !ok {
t.Fatalf("missing MCP tool; tools = %#v", ag.tools)
}
}

func TestUseAgentUsesAnthropicProvider(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/v1/messages" {
t.Fatalf("path = %q", r.URL.Path)
}
if got := r.Header.Get("x-api-key"); got != "test-key" {
t.Fatalf("x-api-key = %q", got)
}
var req map[string]any
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
t.Fatal(err)
}
if req["model"] != "claude-test" {
t.Fatalf("model = %#v", req["model"])
}

w.Header().Set("Content-Type", "text/event-stream")
fmt.Fprint(w, "event: message_start\ndata: {\"type\":\"message_start\",\"message\":{\"id\":\"msg_1\",\"type\":\"message\",\"role\":\"assistant\",\"model\":\"claude-test\",\"usage\":{\"input_tokens\":1,\"output_tokens\":0}}}\n\n")
fmt.Fprint(w, "event: content_block_delta\ndata: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"hello\"}}\n\n")
fmt.Fprint(w, "event: message_delta\ndata: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"},\"usage\":{\"output_tokens\":1}}\n\n")
fmt.Fprint(w, "event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n")
}))
defer server.Close()

m := NewAgentManager(&config.Config{
Profiles: map[string]*config.ProfileConfig{
"default": {Provider: "claude", ModelName: "claude-test", Workspace: t.TempDir()},
},
Providers: map[string]*config.ProviderConfig{
"claude": {Type: "anthropic", APIBase: server.URL, APIKey: "test-key"},
},
})
ag, err := m.UseAgent("default")
if err != nil {
t.Fatal(err)
}
stream, err := ag.Stream(context.Background(), &openai.ChatCompletionRequest{
Model: "claude-test", Messages: []openai.ChatCompletionMessage{openai.UserMessage("hi")}, Stream: true,
})
if err != nil {
t.Fatal(err)
}
builder := openai.NewMessageBuilder()
for chunk := range stream {
if message := chunk.GetMessage(); message != nil {
builder.Update(*message)
}
}
if got := builder.Build().Content; got != "hello" {
t.Fatalf("content = %q", got)
}
}

Expand Down
Loading
Loading