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
38 changes: 9 additions & 29 deletions agent/agent.go
Original file line number Diff line number Diff line change
Expand Up @@ -50,7 +50,10 @@ func (a *Agent) toolDefs() []openai.ToolDef {
return defs
}

func (a *Agent) RunAgentLoop(ctx context.Context, sess *session.Session, sink EventSink) error {
// Run advances sess until the agent finishes its turn. Persistence is owned by
// the caller; the same loop can therefore run persistent user sessions or
// ephemeral delegated tasks.
func (a *Agent) Run(ctx context.Context, sess *session.Session, sink EventSink) error {
if a.Stream == nil {
return fmt.Errorf("agent has no model stream")
}
Expand Down Expand Up @@ -99,9 +102,6 @@ func (a *Agent) RunAgentLoop(ctx context.Context, sess *session.Session, sink Ev
return err
}
a.AppendContextMaintenanceNotice(sess)
if err := sess.Save(); err != nil {
return fmt.Errorf("save session: %w", err)
}
if err := sink.Done(); err != nil {
return err
}
Expand Down Expand Up @@ -150,12 +150,14 @@ func (a *Agent) RunAgentLoop(ctx context.Context, sess *session.Session, sink Ev
}
sess.Messages = append(sess.Messages, openai.ToolResultMessage(tc.ID, tc.Function.Name, result))
}
if err := sess.Save(); err != nil {
return fmt.Errorf("save session: %w", err)
}
}
}

// RunAgentLoop is kept for callers using the previous name.
func (a *Agent) RunAgentLoop(ctx context.Context, sess *session.Session, sink EventSink) error {
return a.Run(ctx, sess, sink)
}

func emitAttachedFileResult(sink EventSink, result string) (string, bool, error) {
var attachment tools.AttachFileResult
if err := json.Unmarshal([]byte(result), &attachment); err != nil {
Expand Down Expand Up @@ -210,25 +212,3 @@ func (a *Agent) readSystemPrompt() string {
}
return string(data)
}

func (a *Agent) BuildTools() {
workspace := a.Config.GetWorkspace()
if workspace != "" {
_ = os.MkdirAll(workspace, 0755)
}
a.Use(
&tools.WebFetchTool{},
&tools.WebSearchTool{},
&tools.ReadFileTool{Workspace: workspace},
&tools.WriteFileTool{Workspace: workspace},
&tools.AppendFileTool{Workspace: workspace},
&tools.EditFileTool{Workspace: workspace},
&tools.AttachFileTool{Workspace: workspace},
&tools.ExecTool{
Workspace: workspace,
DefaultTimeout: tools.ExecDefaultTimeoutSeconds,
RestrictToWorkspace: true,
},
&tools.SkillsTool{Workspace: filepath.Join(config.ConfigPath, "skills")},
)
}
40 changes: 20 additions & 20 deletions agent/agent_loop_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,6 @@ package agent

import (
"context"
"os"
"path/filepath"
"strings"
"testing"

Expand Down Expand Up @@ -34,33 +32,33 @@ func (discardSink) SessionInfo(SessionInfoEvent) error { return nil }
func (discardSink) Usage(UsageEvent) error { return nil }
func (discardSink) Done() error { return nil }

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

err := ag.RunAgentLoop(context.Background(), session.New("test"), discardSink{})
err := ag.Run(context.Background(), session.New("test"), discardSink{})
if err == nil || !strings.Contains(err.Error(), "closed without a response") {
t.Fatalf("RunAgentLoop error = %v", err)
t.Fatalf("Run error = %v", err)
}
}

func TestRunAgentLoopRejectsInterruptedStream(t *testing.T) {
func TestRunRejectsInterruptedStream(t *testing.T) {
message := openai.ChatCompletionMessage{Role: openai.RoleAssistant, Content: "partial"}
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{})
err := ag.Run(context.Background(), sess, discardSink{})
if err == nil || !strings.Contains(err.Error(), "connection reset") {
t.Fatalf("RunAgentLoop error = %v", err)
t.Fatalf("Run error = %v", err)
}
if len(sess.Messages) != 0 {
t.Fatalf("interrupted response was saved: %#v", sess.Messages)
}
}

func TestRunAgentLoopRejectsToolCallWithoutID(t *testing.T) {
func TestRunRejectsToolCallWithoutID(t *testing.T) {
message := openai.ChatCompletionMessage{
Role: openai.RoleAssistant,
ToolCalls: []openai.ToolCall{{
Expand All @@ -71,28 +69,30 @@ func TestRunAgentLoopRejectsToolCallWithoutID(t *testing.T) {
openai.ChatCompletionResponse{Choices: []openai.ChatCompletionChoice{{Index: 0, Delta: &message}}},
))

err := ag.RunAgentLoop(context.Background(), session.New("test"), discardSink{})
err := ag.Run(context.Background(), session.New("test"), discardSink{})
if err == nil || !strings.Contains(err.Error(), "missing an id") {
t.Fatalf("RunAgentLoop error = %v", err)
t.Fatalf("Run error = %v", err)
}
}

func TestRunAgentLoopReturnsSaveError(t *testing.T) {
func TestRunDoesNotPersistSession(t *testing.T) {
oldConfigPath := config.ConfigPath
configPath := filepath.Join(t.TempDir(), "config-file")
if err := os.WriteFile(configPath, []byte("not a directory"), 0600); err != nil {
t.Fatal(err)
}
config.ConfigPath = configPath
config.ConfigPath = t.TempDir()
t.Cleanup(func() { config.ConfigPath = oldConfigPath })

message := openai.ChatCompletionMessage{Role: openai.RoleAssistant, Content: "done"}
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") {
t.Fatalf("RunAgentLoop error = %v", err)
if err := ag.Run(context.Background(), session.New("test"), discardSink{}); err != nil {
t.Fatalf("Run: %v", err)
}
sessions, err := session.List()
if err != nil {
t.Fatalf("List: %v", err)
}
if len(sessions) != 0 {
t.Fatalf("Run persisted %d sessions, want 0", len(sessions))
}
}
142 changes: 142 additions & 0 deletions agent/delegate_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,142 @@
package agent

import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"strings"
"testing"

"github.com/lsongdev/miya-agents/config"
"github.com/lsongdev/miya-agents/session"
)

func TestUseAgentSelectsProfileTools(t *testing.T) {
m := NewAgentManager(&config.Config{
Profiles: map[string]*config.ProfileConfig{
"default": {
Provider: "openai",
ModelName: "test",
Description: "general coordinator",
Workspace: t.TempDir(),
Tools: []string{"web_search", "delegate"},
},
"researcher": {
Provider: "openai",
ModelName: "test",
Description: "web research",
},
},
Providers: map[string]*config.ProviderConfig{
"openai": {APIBase: "http://example.invalid"},
},
})

ag, err := m.UseAgent("default")
if err != nil {
t.Fatalf("UseAgent: %v", err)
}
if _, ok := ag.tool("web_search"); !ok {
t.Fatal("web_search not configured")
}
delegate, ok := ag.tool("delegate")
if !ok {
t.Fatal("delegate not configured")
}
if _, ok := ag.tool("exec"); ok {
t.Fatal("exec should not be configured")
}

params := delegate.Def().Function.Parameters
properties := params["properties"].(map[string]any)
agent := properties["agent"].(map[string]any)
names := agent["enum"].([]string)
if strings.Join(names, ",") != "default,researcher" {
t.Fatalf("delegate agents = %v", names)
}
if !strings.Contains(delegate.Def().Function.Description, "researcher: web research") {
t.Fatalf("delegate description = %q", delegate.Def().Function.Description)
}
}

func TestUseAgentRejectsUnknownTool(t *testing.T) {
m := NewAgentManager(&config.Config{
Profiles: map[string]*config.ProfileConfig{
"default": {Provider: "openai", ModelName: "test", Tools: []string{"missing"}},
},
Providers: map[string]*config.ProviderConfig{
"openai": {APIBase: "http://example.invalid"},
},
})

_, err := m.UseAgent("default")
if err == nil || !strings.Contains(err.Error(), `unknown tool "missing"`) {
t.Fatalf("UseAgent error = %v", err)
}
}

func TestDelegateUsesEphemeralSession(t *testing.T) {
oldConfigPath := config.ConfigPath
config.ConfigPath = t.TempDir()
t.Cleanup(func() { config.ConfigPath = oldConfigPath })

server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/chat/completions" {
t.Fatalf("path = %q", r.URL.Path)
}
var req map[string]any
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
t.Fatal(err)
}
messages, _ := req["messages"].([]any)
if len(messages) == 0 {
t.Fatal("missing delegated task")
}

w.Header().Set("Content-Type", "text/event-stream")
fmt.Fprint(w, "data: {\"id\":\"chat_1\",\"object\":\"chat.completion.chunk\",\"model\":\"test\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\",\"content\":\"delegated result\"}}]}\n\n")
fmt.Fprint(w, "data: {\"id\":\"chat_1\",\"object\":\"chat.completion.chunk\",\"model\":\"test\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\n")
fmt.Fprint(w, "data: [DONE]\n\n")
}))
defer server.Close()

m := NewAgentManager(&config.Config{
Profiles: map[string]*config.ProfileConfig{
"researcher": {
Provider: "openai",
ModelName: "test",
Workspace: t.TempDir(),
Tools: []string{"web_search"},
},
},
Providers: map[string]*config.ProviderConfig{
"openai": {APIBase: server.URL},
},
})

result, err := m.Delegate(context.Background(), "researcher", "investigate this")
if err != nil {
t.Fatalf("Delegate: %v", err)
}
if result != "delegated result" {
t.Fatalf("result = %q", result)
}
sessions, err := session.List()
if err != nil {
t.Fatalf("List: %v", err)
}
if len(sessions) != 0 {
t.Fatalf("Delegate persisted %d sessions, want 0", len(sessions))
}
}

func TestDelegateLimitsRecursion(t *testing.T) {
m := NewAgentManager(&config.Config{})
ctx := context.WithValue(context.Background(), delegateDepthKey{}, maxDelegateDepth)
_, err := m.Delegate(ctx, "anything", "task")
if err == nil || !strings.Contains(err.Error(), "maximum delegation depth") {
t.Fatalf("Delegate error = %v", err)
}
}
Loading
Loading