diff --git a/agent/assistant/assistant.go b/agent/assistant/assistant.go index 7a6825be..93bcdf6c 100644 --- a/agent/assistant/assistant.go +++ b/agent/assistant/assistant.go @@ -25,6 +25,9 @@ func init() { return &agentCallerWrapper{ast: ast}, nil } + // Initialize Agent JSAPI factory for ctx.agent.* methods + caller.SetJSAPIFactory() + // Initialize Search JSAPI factory with config getter search.SetJSAPIFactory(func(assistantID string) (*searchTypes.Config, *search.Uses) { ast, err := Get(assistantID) diff --git a/agent/caller/integration_test.go b/agent/caller/integration_test.go new file mode 100644 index 00000000..e965d6e9 --- /dev/null +++ b/agent/caller/integration_test.go @@ -0,0 +1,278 @@ +package caller_test + +import ( + "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/yaoapp/yao/agent/assistant" + "github.com/yaoapp/yao/agent/caller" + agentContext "github.com/yaoapp/yao/agent/context" + "github.com/yaoapp/yao/agent/testutils" + "github.com/yaoapp/yao/openapi/oauth/types" +) + +func TestIntegration_Call_RealAgent(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + // Load the simple-greeting agent + ast, err := assistant.Get("tests.simple-greeting") + require.NoError(t, err) + require.NotNil(t, ast) + + // Create authorized info for the context + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + // Create a context with authorization + ctx := agentContext.New(context.Background(), authorized, "test-chat-integration") + ctx.AssistantID = "tests.agent-caller" + + // Create JSAPI + api := caller.NewJSAPI(ctx) + + // Call the simple-greeting agent + messages := []interface{}{ + map[string]interface{}{ + "role": "user", + "content": "Hello!", + }, + } + + opts := map[string]interface{}{ + "skip": map[string]interface{}{ + "history": true, + }, + } + + result := api.Call("tests.simple-greeting", messages, opts) + require.NotNil(t, result) + + r, ok := result.(*caller.Result) + require.True(t, ok) + assert.Equal(t, "tests.simple-greeting", r.AgentID) + + // Should either have content or error + if r.Error != "" { + t.Logf("Agent call error: %s", r.Error) + } else { + t.Logf("Agent response content: %s", r.Content) + assert.NotEmpty(t, r.Content) + } +} + +func TestIntegration_All_RealAgents(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + // Create authorized info for the context + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + // Create a context with authorization + ctx := agentContext.New(context.Background(), authorized, "test-chat-all") + ctx.AssistantID = "tests.agent-caller" + + // Create JSAPI + api := caller.NewJSAPI(ctx) + + // Call multiple agents in parallel + requests := []interface{}{ + map[string]interface{}{ + "agent": "tests.simple-greeting", + "messages": []interface{}{ + map[string]interface{}{ + "role": "user", + "content": "Hello from test 1!", + }, + }, + "options": map[string]interface{}{ + "skip": map[string]interface{}{ + "history": true, + }, + }, + }, + map[string]interface{}{ + "agent": "tests.simple-greeting", + "messages": []interface{}{ + map[string]interface{}{ + "role": "user", + "content": "Hello from test 2!", + }, + }, + "options": map[string]interface{}{ + "skip": map[string]interface{}{ + "history": true, + }, + }, + }, + } + + results := api.All(requests) + require.Len(t, results, 2) + + for i, result := range results { + r, ok := result.(*caller.Result) + require.True(t, ok, "result %d should be *caller.Result", i) + assert.Equal(t, "tests.simple-greeting", r.AgentID) + t.Logf("Result[%d]: content=%s, error=%s", i, r.Content, r.Error) + } +} + +func TestIntegration_Any_RealAgents(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + // Create authorized info for the context + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + // Create a context with authorization + ctx := agentContext.New(context.Background(), authorized, "test-chat-any") + ctx.AssistantID = "tests.agent-caller" + + // Create JSAPI + api := caller.NewJSAPI(ctx) + + // Call multiple agents - return when any succeeds + requests := []interface{}{ + map[string]interface{}{ + "agent": "tests.simple-greeting", + "messages": []interface{}{ + map[string]interface{}{ + "role": "user", + "content": "Hello from any test 1!", + }, + }, + "options": map[string]interface{}{ + "skip": map[string]interface{}{ + "history": true, + }, + }, + }, + map[string]interface{}{ + "agent": "tests.simple-greeting", + "messages": []interface{}{ + map[string]interface{}{ + "role": "user", + "content": "Hello from any test 2!", + }, + }, + "options": map[string]interface{}{ + "skip": map[string]interface{}{ + "history": true, + }, + }, + }, + } + + results := api.Any(requests) + require.Len(t, results, 2) + + // At least one should have a result + hasResult := false + for i, result := range results { + if result != nil { + r, ok := result.(*caller.Result) + if ok && r != nil && r.Error == "" { + hasResult = true + t.Logf("Any Result[%d]: content=%s", i, r.Content) + } + } + } + assert.True(t, hasResult, "At least one result should succeed") +} + +func TestIntegration_Race_RealAgents(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + // Create authorized info for the context + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + // Create a context with authorization + ctx := agentContext.New(context.Background(), authorized, "test-chat-race") + ctx.AssistantID = "tests.agent-caller" + + // Create JSAPI + api := caller.NewJSAPI(ctx) + + // Call multiple agents - return when any completes + requests := []interface{}{ + map[string]interface{}{ + "agent": "tests.simple-greeting", + "messages": []interface{}{ + map[string]interface{}{ + "role": "user", + "content": "Hello from race test 1!", + }, + }, + "options": map[string]interface{}{ + "skip": map[string]interface{}{ + "history": true, + }, + }, + }, + map[string]interface{}{ + "agent": "tests.simple-greeting", + "messages": []interface{}{ + map[string]interface{}{ + "role": "user", + "content": "Hello from race test 2!", + }, + }, + "options": map[string]interface{}{ + "skip": map[string]interface{}{ + "history": true, + }, + }, + }, + } + + results := api.Race(requests) + require.Len(t, results, 2) + + // At least one should have completed + hasResult := false + for i, result := range results { + if result != nil { + r, ok := result.(*caller.Result) + if ok && r != nil { + hasResult = true + t.Logf("Race Result[%d]: content=%s, error=%s", i, r.Content, r.Error) + } + } + } + assert.True(t, hasResult, "At least one result should complete") +} diff --git a/agent/caller/jsapi.go b/agent/caller/jsapi.go new file mode 100644 index 00000000..2edc023c --- /dev/null +++ b/agent/caller/jsapi.go @@ -0,0 +1,300 @@ +package caller + +import ( + agentContext "github.com/yaoapp/yao/agent/context" + "github.com/yaoapp/yao/agent/output/message" +) + +// JSAPI implements context.AgentAPI and context.AgentAPIWithCallback interfaces +// Provides ctx.agent.Call(), ctx.agent.All(), ctx.agent.Any(), ctx.agent.Race() +// and their *WithHandler variants for streaming callback support +type JSAPI struct { + ctx *agentContext.Context + orchestrator *Orchestrator +} + +// Ensure JSAPI implements AgentAPIWithCallback +var _ agentContext.AgentAPIWithCallback = (*JSAPI)(nil) + +// NewJSAPI creates a new agent JSAPI instance +func NewJSAPI(ctx *agentContext.Context) *JSAPI { + return &JSAPI{ + ctx: ctx, + orchestrator: NewOrchestrator(ctx), + } +} + +// Call executes a single agent call +// Usage: ctx.agent.Call("assistant-id", messages, options?) +// Returns: { agent_id, response, content, error } +func (api *JSAPI) Call(agentID string, messages []interface{}, opts map[string]interface{}) interface{} { + req := api.buildRequest(agentID, messages, opts) + result := api.orchestrator.callAgent(req) + return result +} + +// All executes all agent calls and waits for all to complete (like Promise.all) +// Each request should have: +// - agent: string - target agent ID +// - messages: array - messages to send +// - options?: object - call options +func (api *JSAPI) All(requests []interface{}) []interface{} { + reqs := api.parseRequests(requests) + results := api.orchestrator.All(reqs) + return api.convertResults(results) +} + +// Any returns as soon as any agent call succeeds (like Promise.any) +// Each request should have: +// - agent: string - target agent ID +// - messages: array - messages to send +// - options?: object - call options +func (api *JSAPI) Any(requests []interface{}) []interface{} { + reqs := api.parseRequests(requests) + results := api.orchestrator.Any(reqs) + return api.convertResults(results) +} + +// Race returns as soon as any agent call completes (like Promise.race) +// Each request should have: +// - agent: string - target agent ID +// - messages: array - messages to send +// - options?: object - call options +func (api *JSAPI) Race(requests []interface{}) []interface{} { + reqs := api.parseRequests(requests) + results := api.orchestrator.Race(reqs) + return api.convertResults(results) +} + +// ============================================================================ +// AgentAPIWithCallback Implementation +// ============================================================================ + +// CallWithHandler executes a single agent call with an OnMessage handler +func (api *JSAPI) CallWithHandler(agentID string, messages []interface{}, opts map[string]interface{}, handler agentContext.OnMessageFunc) interface{} { + req := api.buildRequest(agentID, messages, opts) + req.Handler = handler + result := api.orchestrator.callAgent(req) + return result +} + +// AllWithHandler executes all agent calls with handlers +func (api *JSAPI) AllWithHandler(requests []interface{}, globalHandler agentContext.BatchOnMessageFunc) []interface{} { + reqs := api.parseRequestsWithHandlers(requests, globalHandler) + results := api.orchestrator.All(reqs) + return api.convertResults(results) +} + +// AnyWithHandler executes agent calls and returns on first success, with handlers +func (api *JSAPI) AnyWithHandler(requests []interface{}, globalHandler agentContext.BatchOnMessageFunc) []interface{} { + reqs := api.parseRequestsWithHandlers(requests, globalHandler) + results := api.orchestrator.Any(reqs) + return api.convertResults(results) +} + +// RaceWithHandler executes agent calls and returns on first completion, with handlers +func (api *JSAPI) RaceWithHandler(requests []interface{}, globalHandler agentContext.BatchOnMessageFunc) []interface{} { + reqs := api.parseRequestsWithHandlers(requests, globalHandler) + results := api.orchestrator.Race(reqs) + return api.convertResults(results) +} + +// parseRequestsWithHandlers parses requests and attaches handlers +// It checks for per-request _handler fields and wraps globalHandler with agentID/index +func (api *JSAPI) parseRequestsWithHandlers(requests []interface{}, globalHandler agentContext.BatchOnMessageFunc) []*Request { + reqs := make([]*Request, 0, len(requests)) + + for i, r := range requests { + reqMap, ok := r.(map[string]interface{}) + if !ok { + continue + } + + // Get agent ID + agentID, ok := reqMap["agent"].(string) + if !ok { + continue + } + + // Get messages + messages, ok := reqMap["messages"].([]interface{}) + if !ok { + continue + } + + // Get options (optional) + var opts map[string]interface{} + if o, ok := reqMap["options"].(map[string]interface{}); ok { + opts = o + } + + req := api.buildRequest(agentID, messages, opts) + + // Check for per-request handler first (takes precedence) + if handler, ok := reqMap["_handler"].(agentContext.OnMessageFunc); ok && handler != nil { + req.Handler = handler + } else if globalHandler != nil { + // Wrap global handler with agentID and index + idx := i // Capture index for closure + aid := agentID + req.Handler = func(msg *message.Message) int { + return globalHandler(aid, idx, msg) + } + } + + reqs = append(reqs, req) + } + + return reqs +} + +// buildRequest builds a Request from agentID, messages, and options +func (api *JSAPI) buildRequest(agentID string, messages []interface{}, opts map[string]interface{}) *Request { + req := &Request{ + AgentID: agentID, + Messages: api.parseMessages(messages), + } + + if opts != nil { + req.Options = api.parseCallOptions(opts) + } + + return req +} + +// parseMessages converts []interface{} to []agentContext.Message +func (api *JSAPI) parseMessages(messages []interface{}) []agentContext.Message { + result := make([]agentContext.Message, 0, len(messages)) + for _, m := range messages { + msg, ok := m.(map[string]interface{}) + if !ok { + continue + } + + ctxMsg := agentContext.Message{} + + // Parse role + if role, ok := msg["role"].(string); ok { + ctxMsg.Role = agentContext.MessageRole(role) + } + + // Parse content (can be string or array) + ctxMsg.Content = msg["content"] + + // Parse name + if name, ok := msg["name"].(string); ok { + ctxMsg.Name = &name + } + + // Parse tool_call_id + if toolCallID, ok := msg["tool_call_id"].(string); ok { + ctxMsg.ToolCallID = &toolCallID + } + + // Parse tool_calls + if toolCalls, ok := msg["tool_calls"].([]interface{}); ok { + ctxMsg.ToolCalls = api.parseToolCalls(toolCalls) + } + + // Parse refusal + if refusal, ok := msg["refusal"].(string); ok { + ctxMsg.Refusal = &refusal + } + + result = append(result, ctxMsg) + } + return result +} + +// parseToolCalls converts []interface{} to []agentContext.ToolCall +func (api *JSAPI) parseToolCalls(toolCalls []interface{}) []agentContext.ToolCall { + result := make([]agentContext.ToolCall, 0, len(toolCalls)) + for _, tc := range toolCalls { + tcMap, ok := tc.(map[string]interface{}) + if !ok { + continue + } + + toolCall := agentContext.ToolCall{} + + if id, ok := tcMap["id"].(string); ok { + toolCall.ID = id + } + if tcType, ok := tcMap["type"].(string); ok { + toolCall.Type = agentContext.ToolCallType(tcType) + } + if fn, ok := tcMap["function"].(map[string]interface{}); ok { + if name, ok := fn["name"].(string); ok { + toolCall.Function.Name = name + } + if args, ok := fn["arguments"].(string); ok { + toolCall.Function.Arguments = args + } + } + + result = append(result, toolCall) + } + return result +} + +// parseCallOptions converts map to CallOptions +func (api *JSAPI) parseCallOptions(opts map[string]interface{}) *CallOptions { + callOpts := &CallOptions{} + + if connector, ok := opts["connector"].(string); ok { + callOpts.Connector = connector + } + if mode, ok := opts["mode"].(string); ok { + callOpts.Mode = mode + } + if metadata, ok := opts["metadata"].(map[string]interface{}); ok { + callOpts.Metadata = metadata + } + + // Parse skip configuration + if skip, ok := opts["skip"].(map[string]interface{}); ok { + callOpts.Skip = &agentContext.Skip{} + if history, ok := skip["history"].(bool); ok { + callOpts.Skip.History = history + } + if trace, ok := skip["trace"].(bool); ok { + callOpts.Skip.Trace = trace + } + if output, ok := skip["output"].(bool); ok { + callOpts.Skip.Output = output + } + if keyword, ok := skip["keyword"].(bool); ok { + callOpts.Skip.Keyword = keyword + } + if search, ok := skip["search"].(bool); ok { + callOpts.Skip.Search = search + } + if contentParsing, ok := skip["content_parsing"].(bool); ok { + callOpts.Skip.ContentParsing = contentParsing + } + } + + return callOpts +} + +// parseRequests parses an array of request objects into typed Requests +func (api *JSAPI) parseRequests(requests []interface{}) []*Request { + return api.parseRequestsWithHandlers(requests, nil) +} + +// convertResults converts typed Results to interface slice for JS +func (api *JSAPI) convertResults(results []*Result) []interface{} { + out := make([]interface{}, len(results)) + for i, r := range results { + out[i] = r + } + return out +} + +// SetJSAPIFactory sets the factory function for creating AgentAPI instances +// Called by assistant package during initialization +func SetJSAPIFactory() { + agentContext.AgentAPIFactory = func(ctx *agentContext.Context) agentContext.AgentAPI { + return NewJSAPI(ctx) + } +} diff --git a/agent/caller/jsapi_test.go b/agent/caller/jsapi_test.go new file mode 100644 index 00000000..6f38e360 --- /dev/null +++ b/agent/caller/jsapi_test.go @@ -0,0 +1,145 @@ +package caller_test + +import ( + stdContext "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/yaoapp/yao/agent/caller" + "github.com/yaoapp/yao/agent/context" +) + +func TestNewJSAPI(t *testing.T) { + ctx := context.New(stdContext.Background(), nil, "test-chat") + api := caller.NewJSAPI(ctx) + require.NotNil(t, api) +} + +func TestJSAPI_Call_NoAgentGetter(t *testing.T) { + // Reset AgentGetterFunc + originalGetter := caller.AgentGetterFunc + caller.AgentGetterFunc = nil + defer func() { caller.AgentGetterFunc = originalGetter }() + + ctx := context.New(stdContext.Background(), nil, "test-chat") + api := caller.NewJSAPI(ctx) + + messages := []interface{}{ + map[string]interface{}{ + "role": "user", + "content": "Hello", + }, + } + + result := api.Call("test-agent", messages, nil) + require.NotNil(t, result) + + r, ok := result.(*caller.Result) + require.True(t, ok) + assert.Equal(t, "test-agent", r.AgentID) + assert.Contains(t, r.Error, "agent getter not initialized") +} + +func TestJSAPI_All_Empty(t *testing.T) { + ctx := context.New(stdContext.Background(), nil, "test-chat") + api := caller.NewJSAPI(ctx) + results := api.All([]interface{}{}) + assert.Len(t, results, 0) +} + +func TestJSAPI_Any_Empty(t *testing.T) { + ctx := context.New(stdContext.Background(), nil, "test-chat") + api := caller.NewJSAPI(ctx) + results := api.Any([]interface{}{}) + assert.Len(t, results, 0) +} + +func TestJSAPI_Race_Empty(t *testing.T) { + ctx := context.New(stdContext.Background(), nil, "test-chat") + api := caller.NewJSAPI(ctx) + results := api.Race([]interface{}{}) + assert.Len(t, results, 0) +} + +func TestJSAPI_All_InvalidRequests(t *testing.T) { + ctx := context.New(stdContext.Background(), nil, "test-chat") + api := caller.NewJSAPI(ctx) + + // Mix of invalid and valid requests + requests := []interface{}{ + "invalid", // Not a map + map[string]interface{}{ + "messages": []interface{}{}, // Missing agent + }, + map[string]interface{}{ + "agent": "test-agent", // Missing messages + }, + } + + results := api.All(requests) + // None should produce a result (all invalid) + assert.Len(t, results, 0) +} + +func TestJSAPI_Call_WithOptions(t *testing.T) { + // Reset AgentGetterFunc + originalGetter := caller.AgentGetterFunc + caller.AgentGetterFunc = nil + defer func() { caller.AgentGetterFunc = originalGetter }() + + ctx := context.New(stdContext.Background(), nil, "test-chat") + api := caller.NewJSAPI(ctx) + + messages := []interface{}{ + map[string]interface{}{ + "role": "user", + "content": "Hello", + }, + } + + opts := map[string]interface{}{ + "connector": "gpt4", + "mode": "chat", + "metadata": map[string]interface{}{ + "key": "value", + }, + "skip": map[string]interface{}{ + "history": true, + "trace": true, + }, + } + + result := api.Call("test-agent", messages, opts) + require.NotNil(t, result) + + r, ok := result.(*caller.Result) + require.True(t, ok) + assert.Equal(t, "test-agent", r.AgentID) + // Still errors because AgentGetterFunc is nil + assert.Contains(t, r.Error, "agent getter not initialized") +} + +func TestSetJSAPIFactory(t *testing.T) { + // Reset factory + context.AgentAPIFactory = nil + + // Set factory + caller.SetJSAPIFactory() + + // Verify factory is set + require.NotNil(t, context.AgentAPIFactory) + + // Create a mock context + ctx := context.New(stdContext.Background(), nil, "test-chat") + + // Get agent API + agentAPI := context.AgentAPIFactory(ctx) + require.NotNil(t, agentAPI) +} + +func TestJSAPI_ImplementsAgentAPI(t *testing.T) { + // Verify JSAPI implements context.AgentAPI interface + ctx := context.New(stdContext.Background(), nil, "test-chat") + var _ context.AgentAPI = caller.NewJSAPI(ctx) +} diff --git a/agent/caller/orchestrator.go b/agent/caller/orchestrator.go new file mode 100644 index 00000000..6a0c1210 --- /dev/null +++ b/agent/caller/orchestrator.go @@ -0,0 +1,286 @@ +package caller + +import ( + "sync" + + agentContext "github.com/yaoapp/yao/agent/context" +) + +// Orchestrator handles parallel agent calls with different concurrency patterns +// Modeled after JavaScript Promise patterns (all, any, race) +type Orchestrator struct { + ctx *agentContext.Context +} + +// NewOrchestrator creates a new Orchestrator for parallel agent calls +func NewOrchestrator(ctx *agentContext.Context) *Orchestrator { + return &Orchestrator{ctx: ctx} +} + +// callResult is used internally to pass results through channels +type callResult struct { + idx int + result *Result +} + +// All executes all agent calls and waits for all to complete (like Promise.all) +// Returns results in the same order as requests, regardless of completion order +func (o *Orchestrator) All(reqs []*Request) []*Result { + if len(reqs) == 0 { + return []*Result{} + } + + results := make([]*Result, len(reqs)) + var wg sync.WaitGroup + var mu sync.Mutex + + for i, req := range reqs { + wg.Add(1) + go func(idx int, r *Request) { + defer wg.Done() + defer func() { + if err := recover(); err != nil { + mu.Lock() + results[idx] = &Result{ + AgentID: r.AgentID, + Error: "agent call panic recovered", + } + mu.Unlock() + } + }() + + result := o.callAgent(r) + mu.Lock() + results[idx] = result + mu.Unlock() + }(i, req) + } + + wg.Wait() + return results +} + +// Any returns as soon as any agent call succeeds (has non-error result) (like Promise.any) +// Other calls continue in background but results are discarded after first success +// Returns all results received so far when first success is found +func (o *Orchestrator) Any(reqs []*Request) []*Result { + if len(reqs) == 0 { + return []*Result{} + } + + results := make([]*Result, len(reqs)) + resultChan := make(chan callResult, len(reqs)) + + var wg sync.WaitGroup + done := make(chan struct{}) + + for i, req := range reqs { + wg.Add(1) + go func(idx int, r *Request) { + defer wg.Done() + defer func() { + if err := recover(); err != nil { + // Send panic result through channel + select { + case <-done: + case resultChan <- callResult{idx: idx, result: &Result{ + AgentID: r.AgentID, + Error: "agent call panic recovered", + }}: + } + } + }() + + // Check if done before starting + select { + case <-done: + return + default: + } + + result := o.callAgent(r) + + // Try to send result + select { + case <-done: + // Already found a successful result + case resultChan <- callResult{idx: idx, result: result}: + } + }(i, req) + } + + // Close channel when all goroutines complete + go func() { + wg.Wait() + close(resultChan) + }() + + // Collect results until we find one with success (no error and has content) + var foundSuccess bool + for res := range resultChan { + results[res.idx] = res.result + // Check if this result is successful (no error) + if !foundSuccess && res.result != nil && res.result.Error == "" { + foundSuccess = true + close(done) // Signal other goroutines to stop + } + } + + return results +} + +// Race returns as soon as any agent call completes (like Promise.race) +// Returns immediately when first result arrives, regardless of success/failure +// Note: Still waits for all goroutines to complete before returning to avoid resource leaks +func (o *Orchestrator) Race(reqs []*Request) []*Result { + if len(reqs) == 0 { + return []*Result{} + } + + results := make([]*Result, len(reqs)) + resultChan := make(chan callResult, len(reqs)) + + var wg sync.WaitGroup + done := make(chan struct{}) + + for i, req := range reqs { + wg.Add(1) + go func(idx int, r *Request) { + defer wg.Done() + defer func() { + if err := recover(); err != nil { + // Send panic result through channel + select { + case <-done: + case resultChan <- callResult{idx: idx, result: &Result{ + AgentID: r.AgentID, + Error: "agent call panic recovered", + }}: + } + } + }() + + // Check if done before starting + select { + case <-done: + return + default: + } + + result := o.callAgent(r) + + // Try to send result + select { + case <-done: + // Already got first result + case resultChan <- callResult{idx: idx, result: result}: + } + }(i, req) + } + + // Close channel when all goroutines complete + go func() { + wg.Wait() + close(resultChan) + }() + + // Get first result and signal others to stop + var gotFirst bool + for res := range resultChan { + results[res.idx] = res.result + if !gotFirst { + gotFirst = true + close(done) // Signal other goroutines to stop + } + } + + return results +} + +// callAgent executes a single agent call using the AgentGetterFunc +// This method handles context sharing and result extraction +func (o *Orchestrator) callAgent(req *Request) *Result { + if req == nil { + return &Result{Error: "nil request"} + } + + result := &Result{ + AgentID: req.AgentID, + } + + // Get the agent using the getter function + if AgentGetterFunc == nil { + result.Error = "agent getter not initialized" + return result + } + + agent, err := AgentGetterFunc(req.AgentID) + if err != nil { + result.Error = "failed to get agent: " + err.Error() + return result + } + + // Build context options for the call + var ctxOpts *agentContext.Options + if req.Options != nil { + ctxOpts = req.Options.ToContextOptions() + } else { + ctxOpts = &agentContext.Options{} + } + + // If request has a handler, set OnMessage callback + if req.Handler != nil { + if ctxOpts == nil { + ctxOpts = &agentContext.Options{} + } + // Set OnMessage to receive SSE messages + ctxOpts.OnMessage = req.Handler + } + + // Execute the agent call with shared context + // The agent.Stream method will use the parent context's Writer for output + resp, err := agent.Stream(o.ctx, req.Messages, ctxOpts) + if err != nil { + result.Error = "agent call failed: " + err.Error() + return result + } + + result.Response = resp + + // Extract content from completion if available + if resp != nil && resp.Completion != nil { + result.Content = extractContentFromCompletion(resp.Completion) + } + + return result +} + +// extractContentFromCompletion extracts the text content from a completion response +func extractContentFromCompletion(completion *agentContext.CompletionResponse) string { + if completion == nil { + return "" + } + + // Content can be string or []ContentPart + switch content := completion.Content.(type) { + case string: + return content + case []interface{}: + // Handle array of content parts - extract text parts + var texts []string + for _, part := range content { + if partMap, ok := part.(map[string]interface{}); ok { + if partType, ok := partMap["type"].(string); ok && partType == "text" { + if text, ok := partMap["text"].(string); ok { + texts = append(texts, text) + } + } + } + } + if len(texts) > 0 { + return texts[0] // Return first text content + } + } + + return "" +} diff --git a/agent/caller/orchestrator_test.go b/agent/caller/orchestrator_test.go new file mode 100644 index 00000000..bd55868d --- /dev/null +++ b/agent/caller/orchestrator_test.go @@ -0,0 +1,162 @@ +package caller_test + +import ( + stdContext "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/yaoapp/yao/agent/caller" + "github.com/yaoapp/yao/agent/context" +) + +func TestNewOrchestrator(t *testing.T) { + ctx := context.New(stdContext.Background(), nil, "test-chat") + orch := caller.NewOrchestrator(ctx) + require.NotNil(t, orch) +} + +func TestOrchestrator_All_Empty(t *testing.T) { + ctx := context.New(stdContext.Background(), nil, "test-chat") + orch := caller.NewOrchestrator(ctx) + + results := orch.All([]*caller.Request{}) + assert.Len(t, results, 0) +} + +func TestOrchestrator_Any_Empty(t *testing.T) { + ctx := context.New(stdContext.Background(), nil, "test-chat") + orch := caller.NewOrchestrator(ctx) + + results := orch.Any([]*caller.Request{}) + assert.Len(t, results, 0) +} + +func TestOrchestrator_Race_Empty(t *testing.T) { + ctx := context.New(stdContext.Background(), nil, "test-chat") + orch := caller.NewOrchestrator(ctx) + + results := orch.Race([]*caller.Request{}) + assert.Len(t, results, 0) +} + +func TestOrchestrator_All_NoGetter(t *testing.T) { + // Reset AgentGetterFunc + originalGetter := caller.AgentGetterFunc + caller.AgentGetterFunc = nil + defer func() { caller.AgentGetterFunc = originalGetter }() + + ctx := context.New(stdContext.Background(), nil, "test-chat") + orch := caller.NewOrchestrator(ctx) + + reqs := []*caller.Request{ + { + AgentID: "agent1", + Messages: []context.Message{{Role: "user", Content: "Hello"}}, + }, + { + AgentID: "agent2", + Messages: []context.Message{{Role: "user", Content: "World"}}, + }, + } + + results := orch.All(reqs) + require.Len(t, results, 2) + + // All should have errors because no getter + for i, r := range results { + require.NotNil(t, r, "result %d should not be nil", i) + assert.Contains(t, r.Error, "agent getter not initialized") + } +} + +func TestOrchestrator_Any_NoGetter(t *testing.T) { + // Reset AgentGetterFunc + originalGetter := caller.AgentGetterFunc + caller.AgentGetterFunc = nil + defer func() { caller.AgentGetterFunc = originalGetter }() + + ctx := context.New(stdContext.Background(), nil, "test-chat") + orch := caller.NewOrchestrator(ctx) + + reqs := []*caller.Request{ + { + AgentID: "agent1", + Messages: []context.Message{{Role: "user", Content: "Hello"}}, + }, + { + AgentID: "agent2", + Messages: []context.Message{{Role: "user", Content: "World"}}, + }, + } + + results := orch.Any(reqs) + require.Len(t, results, 2) + + // At least one result should exist + hasResult := false + for _, r := range results { + if r != nil { + hasResult = true + assert.Contains(t, r.Error, "agent getter not initialized") + } + } + assert.True(t, hasResult) +} + +func TestOrchestrator_Race_NoGetter(t *testing.T) { + // Reset AgentGetterFunc + originalGetter := caller.AgentGetterFunc + caller.AgentGetterFunc = nil + defer func() { caller.AgentGetterFunc = originalGetter }() + + ctx := context.New(stdContext.Background(), nil, "test-chat") + orch := caller.NewOrchestrator(ctx) + + reqs := []*caller.Request{ + { + AgentID: "agent1", + Messages: []context.Message{{Role: "user", Content: "Hello"}}, + }, + { + AgentID: "agent2", + Messages: []context.Message{{Role: "user", Content: "World"}}, + }, + } + + results := orch.Race(reqs) + require.Len(t, results, 2) + + // At least one result should exist (first to complete) + hasResult := false + for _, r := range results { + if r != nil { + hasResult = true + } + } + assert.True(t, hasResult) +} + +func TestOrchestrator_All_NilRequest(t *testing.T) { + ctx := context.New(stdContext.Background(), nil, "test-chat") + orch := caller.NewOrchestrator(ctx) + + reqs := []*caller.Request{ + nil, + { + AgentID: "agent1", + Messages: []context.Message{{Role: "user", Content: "Hello"}}, + }, + } + + // Reset AgentGetterFunc + originalGetter := caller.AgentGetterFunc + caller.AgentGetterFunc = nil + defer func() { caller.AgentGetterFunc = originalGetter }() + + results := orch.All(reqs) + require.Len(t, results, 2) + + // First result should have "nil request" error + assert.Contains(t, results[0].Error, "nil request") +} diff --git a/agent/caller/types.go b/agent/caller/types.go new file mode 100644 index 00000000..c763c696 --- /dev/null +++ b/agent/caller/types.go @@ -0,0 +1,44 @@ +// Package caller provides types and utilities for agent-to-agent calls +package caller + +import ( + agentContext "github.com/yaoapp/yao/agent/context" +) + +// Request represents a request to call an agent +type Request struct { + AgentID string `json:"agent"` // Target agent ID + Messages []agentContext.Message `json:"messages"` // Messages to send + Options *CallOptions `json:"options,omitempty"` // Call options + Handler agentContext.OnMessageFunc `json:"-"` // OnMessage handler for this request (not serialized) +} + +// CallOptions represents options for an agent call +type CallOptions struct { + Connector string `json:"connector,omitempty"` // Override connector + Mode string `json:"mode,omitempty"` // Agent mode (chat, etc.) + Metadata map[string]interface{} `json:"metadata,omitempty"` // Custom metadata passed to hooks + Skip *agentContext.Skip `json:"skip,omitempty"` // Skip configuration (history, trace, output, etc.) +} + +// Result represents the result of an agent call +type Result struct { + AgentID string `json:"agent_id"` // Agent ID that was called + Response *agentContext.Response `json:"response,omitempty"` // Full response from agent + Content string `json:"content,omitempty"` // Final text content (extracted from completion) + Error string `json:"error,omitempty"` // Error message if call failed +} + +// ToContextOptions converts CallOptions to context.Options for the agent call +func (o *CallOptions) ToContextOptions() *agentContext.Options { + if o == nil { + return nil + } + + return &agentContext.Options{ + Connector: o.Connector, + Mode: o.Mode, + Metadata: o.Metadata, + Skip: o.Skip, + } +} diff --git a/agent/caller/types_test.go b/agent/caller/types_test.go new file mode 100644 index 00000000..27a51f4a --- /dev/null +++ b/agent/caller/types_test.go @@ -0,0 +1,86 @@ +package caller_test + +import ( + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/yaoapp/yao/agent/caller" + "github.com/yaoapp/yao/agent/context" +) + +func TestCallOptions_ToContextOptions_Nil(t *testing.T) { + var opts *caller.CallOptions + ctxOpts := opts.ToContextOptions() + assert.Nil(t, ctxOpts) +} + +func TestCallOptions_ToContextOptions_Empty(t *testing.T) { + opts := &caller.CallOptions{} + ctxOpts := opts.ToContextOptions() + require.NotNil(t, ctxOpts) + assert.Empty(t, ctxOpts.Connector) + assert.Empty(t, ctxOpts.Mode) + assert.Nil(t, ctxOpts.Metadata) + assert.Nil(t, ctxOpts.Skip) +} + +func TestCallOptions_ToContextOptions_Full(t *testing.T) { + opts := &caller.CallOptions{ + Connector: "gpt4", + Mode: "chat", + Metadata: map[string]interface{}{ + "key": "value", + }, + Skip: &context.Skip{ + History: true, + Trace: true, + Output: false, + }, + } + + ctxOpts := opts.ToContextOptions() + require.NotNil(t, ctxOpts) + assert.Equal(t, "gpt4", ctxOpts.Connector) + assert.Equal(t, "chat", ctxOpts.Mode) + assert.Equal(t, "value", ctxOpts.Metadata["key"]) + require.NotNil(t, ctxOpts.Skip) + assert.True(t, ctxOpts.Skip.History) + assert.True(t, ctxOpts.Skip.Trace) + assert.False(t, ctxOpts.Skip.Output) +} + +func TestRequest_Basic(t *testing.T) { + req := &caller.Request{ + AgentID: "test-agent", + Messages: []context.Message{ + {Role: "user", Content: "Hello"}, + }, + } + + assert.Equal(t, "test-agent", req.AgentID) + assert.Len(t, req.Messages, 1) + assert.Equal(t, context.MessageRole("user"), req.Messages[0].Role) +} + +func TestResult_Basic(t *testing.T) { + result := &caller.Result{ + AgentID: "test-agent", + Content: "Hello response", + } + + assert.Equal(t, "test-agent", result.AgentID) + assert.Equal(t, "Hello response", result.Content) + assert.Empty(t, result.Error) +} + +func TestResult_WithError(t *testing.T) { + result := &caller.Result{ + AgentID: "test-agent", + Error: "something went wrong", + } + + assert.Equal(t, "test-agent", result.AgentID) + assert.Equal(t, "something went wrong", result.Error) + assert.Empty(t, result.Content) +} diff --git a/agent/context/jsapi.go b/agent/context/jsapi.go index fa7707c1..aefbdac0 100644 --- a/agent/context/jsapi.go +++ b/agent/context/jsapi.go @@ -68,6 +68,9 @@ func (ctx *Context) NewObject(v8ctx *v8go.Context) (*v8go.Value, error) { // Set search object jsObject.Set("search", ctx.newSearchObject(v8ctx.Isolate())) + // Set agent object for calling other agents + jsObject.Set("agent", ctx.newAgentObject(v8ctx.Isolate())) + // Note: Space object will be set after instance creation (requires v8ctx) // Create instance diff --git a/agent/context/jsapi_agent.go b/agent/context/jsapi_agent.go new file mode 100644 index 00000000..bb337210 --- /dev/null +++ b/agent/context/jsapi_agent.go @@ -0,0 +1,507 @@ +package context + +import ( + "github.com/yaoapp/gou/runtime/v8/bridge" + "github.com/yaoapp/yao/agent/output/message" + "rogchap.com/v8go" +) + +// AgentAPI defines the agent JSAPI interface for ctx.agent.* +// This interface is defined here to avoid circular dependency between context and caller packages. +// The actual implementation is in agent/caller/jsapi.go +type AgentAPI interface { + // Call executes a single agent call + // Returns *caller.Result or error information + Call(agentID string, messages []interface{}, opts map[string]interface{}) interface{} + + // Parallel agent call methods - inspired by JavaScript Promise + // All waits for all agent calls to complete (like Promise.all) + All(requests []interface{}) []interface{} + // Any returns when any agent call succeeds (like Promise.any) + Any(requests []interface{}) []interface{} + // Race returns when any agent call completes (like Promise.race) + Race(requests []interface{}) []interface{} +} + +// AgentAPIWithCallback extends AgentAPI with callback support +// This interface provides methods that accept OnMessage handlers for real-time message processing +type AgentAPIWithCallback interface { + AgentAPI + + // CallWithHandler executes a single agent call with an OnMessage handler + // handler receives SSE messages: func(msg *message.Message) int + CallWithHandler(agentID string, messages []interface{}, opts map[string]interface{}, handler OnMessageFunc) interface{} + + // AllWithHandler executes all agent calls with handlers + // globalHandler receives messages with agentID and index: func(agentID, index, msg) int + // Individual request handlers (if set) take precedence over globalHandler + AllWithHandler(requests []interface{}, globalHandler BatchOnMessageFunc) []interface{} + + // AnyWithHandler executes agent calls and returns on first success, with handlers + AnyWithHandler(requests []interface{}, globalHandler BatchOnMessageFunc) []interface{} + + // RaceWithHandler executes agent calls and returns on first completion, with handlers + RaceWithHandler(requests []interface{}, globalHandler BatchOnMessageFunc) []interface{} +} + +// BatchOnMessageFunc is the OnMessage function for batch calls +// It includes agentID and index to identify the source of each message +type BatchOnMessageFunc func(agentID string, index int, msg *message.Message) int + +// AgentAPIFactory is a function type that creates an AgentAPI for a context +// This is set by the caller package during initialization +var AgentAPIFactory func(ctx *Context) AgentAPI + +// Agent returns the agent API for this context +// Returns nil if AgentAPIFactory is not set +func (ctx *Context) Agent() AgentAPI { + if AgentAPIFactory == nil { + return nil + } + return AgentAPIFactory(ctx) +} + +// newAgentObject creates a new agent object with all agent methods +// This is called from jsapi.go NewObject() to mount ctx.agent +func (ctx *Context) newAgentObject(iso *v8go.Isolate) *v8go.ObjectTemplate { + agentObj := v8go.NewObjectTemplate(iso) + + // Single agent call method + agentObj.Set("Call", ctx.agentCallMethod(iso)) + + // Parallel agent call methods - inspired by JavaScript Promise + agentObj.Set("All", ctx.agentAllMethod(iso)) + agentObj.Set("Any", ctx.agentAnyMethod(iso)) + agentObj.Set("Race", ctx.agentRaceMethod(iso)) + + return agentObj +} + +// agentCallMethod implements ctx.agent.Call(agentID, messages, options?) +// Usage: const result = ctx.agent.Call("assistant-id", [{ role: "user", content: "Hello" }], { connector: "gpt4", onChunk: (type, data) => 0 }) +// Returns: { agent_id, response, content, error } +func (ctx *Context) agentCallMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + // Validate arguments + if len(args) < 2 { + return bridge.JsException(v8ctx, "Call requires agentID and messages parameters") + } + + // Get agent ID (first argument) + if !args[0].IsString() { + return bridge.JsException(v8ctx, "agentID must be a string") + } + agentID := args[0].String() + + // Parse messages (second argument) + messagesVal, err := bridge.GoValue(args[1], v8ctx) + if err != nil { + return bridge.JsException(v8ctx, "invalid messages: "+err.Error()) + } + messages, ok := messagesVal.([]interface{}) + if !ok { + return bridge.JsException(v8ctx, "messages must be an array") + } + + // Parse options (optional third argument) - extract onChunk separately + var opts map[string]interface{} + var onChunkFn *v8go.Function + + if len(args) >= 3 && !args[2].IsUndefined() && !args[2].IsNull() { + optsObj, err := args[2].AsObject() + if err == nil && optsObj != nil { + // Extract onChunk callback before converting to Go value + onChunkVal, _ := optsObj.Get("onChunk") + if onChunkVal != nil && onChunkVal.IsFunction() { + onChunkFn, _ = onChunkVal.AsFunction() + } + + // Convert the rest of options to Go map + goVal, err := bridge.GoValue(args[2], v8ctx) + if err == nil { + if optsMap, ok := goVal.(map[string]interface{}); ok { + // Remove onChunk from the map (it's handled separately) + delete(optsMap, "onChunk") + opts = optsMap + } + } + } + } + + // Get agent API + agentAPI := ctx.Agent() + if agentAPI == nil { + return bridge.JsException(v8ctx, "agent API not available") + } + + var result interface{} + + // If onChunk callback is provided and API supports it, use CallWithHandler + if onChunkFn != nil { + if apiWithCb, ok := agentAPI.(AgentAPIWithCallback); ok { + // Create Go StreamFunc that calls JS callback + handler := createJSStreamHandler(v8ctx, onChunkFn) + result = apiWithCb.CallWithHandler(agentID, messages, opts, handler) + } else { + // Fallback: ignore callback if API doesn't support it + result = agentAPI.Call(agentID, messages, opts) + } + } else { + // No callback, use regular Call + result = agentAPI.Call(agentID, messages, opts) + } + + // Convert result to JS value + jsVal, err := bridge.JsValue(v8ctx, result) + if err != nil { + return bridge.JsException(v8ctx, "failed to convert result: "+err.Error()) + } + + return jsVal + }) +} + +// createJSOnMessageHandler creates a Go OnMessageFunc that calls a JS callback +// JS callback signature: (msg: object) => number +// msg contains: type, props, delta, message_id, chunk_id, etc. +func createJSStreamHandler(v8ctx *v8go.Context, callback *v8go.Function) OnMessageFunc { + return func(msg *message.Message) int { + if callback == nil || v8ctx == nil || msg == nil { + return 0 // Continue if no callback + } + + // Convert message to JS value + jsMsg, err := bridge.JsValue(v8ctx, msg) + if err != nil { + return 1 // Stop on error + } + + // Call the JS callback with the message object + result, err := callback.Call(v8ctx.Global(), jsMsg) + if err != nil { + return 1 // Stop on error + } + + // Check return value (0 = continue, non-zero = stop) + if result != nil && result.IsNumber() { + ret := result.Integer() + if ret != 0 { + return int(ret) + } + } + + return 0 // Continue + } +} + +// agentAllMethod implements ctx.agent.All(requests, options?) +// Waits for all agent calls to complete (like Promise.all) +// Each request should have: +// - agent: string - target agent ID +// - messages: array - messages to send +// - options?: object - call options +// +// Global options (second argument): +// - onChunk?: (agentID, index, msg) => number - callback for all messages (uses channel for V8 safety) +func (ctx *Context) agentAllMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + // Validate arguments + if len(args) < 1 { + return bridge.JsException(v8ctx, "All requires requests parameter") + } + + // Parse requests and extract global callback + requests, globalCallback := ctx.parseRequestsForBatch(args, v8ctx) + + // Get agent API + agentAPI := ctx.Agent() + if agentAPI == nil { + return bridge.JsException(v8ctx, "agent API not available") + } + + // Execute with channel-based callback handling + results := ctx.executeBatchWithCallback(BatchMethodAll, requests, globalCallback, v8ctx) + + // Convert results to JS value + jsVal, err := bridge.JsValue(v8ctx, results) + if err != nil { + return bridge.JsException(v8ctx, "failed to convert results: "+err.Error()) + } + + return jsVal + }) +} + +// agentAnyMethod implements ctx.agent.Any(requests, options?) +// Returns when any agent call succeeds (like Promise.any) +// Each request should have: +// - agent: string - target agent ID +// - messages: array - messages to send +// - options?: object - call options +// +// Global options (second argument): +// - onChunk?: (agentID, index, msg) => number - callback for all messages (uses channel for V8 safety) +func (ctx *Context) agentAnyMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + // Validate arguments + if len(args) < 1 { + return bridge.JsException(v8ctx, "Any requires requests parameter") + } + + // Parse requests and extract global callback + requests, globalCallback := ctx.parseRequestsForBatch(args, v8ctx) + + // Get agent API + agentAPI := ctx.Agent() + if agentAPI == nil { + return bridge.JsException(v8ctx, "agent API not available") + } + + // Execute with channel-based callback handling + results := ctx.executeBatchWithCallback(BatchMethodAny, requests, globalCallback, v8ctx) + + // Convert results to JS value + jsVal, err := bridge.JsValue(v8ctx, results) + if err != nil { + return bridge.JsException(v8ctx, "failed to convert results: "+err.Error()) + } + + return jsVal + }) +} + +// agentRaceMethod implements ctx.agent.Race(requests, options?) +// Returns when any agent call completes (like Promise.race) +// Each request should have: +// - agent: string - target agent ID +// - messages: array - messages to send +// - options?: object - call options +// +// Global options (second argument): +// - onChunk?: (agentID, index, msg) => number - callback for all messages (uses channel for V8 safety) +func (ctx *Context) agentRaceMethod(iso *v8go.Isolate) *v8go.FunctionTemplate { + return v8go.NewFunctionTemplate(iso, func(info *v8go.FunctionCallbackInfo) *v8go.Value { + v8ctx := info.Context() + args := info.Args() + + // Validate arguments + if len(args) < 1 { + return bridge.JsException(v8ctx, "Race requires requests parameter") + } + + // Parse requests and extract global callback + requests, globalCallback := ctx.parseRequestsForBatch(args, v8ctx) + + // Get agent API + agentAPI := ctx.Agent() + if agentAPI == nil { + return bridge.JsException(v8ctx, "agent API not available") + } + + // Execute with channel-based callback handling + results := ctx.executeBatchWithCallback(BatchMethodRace, requests, globalCallback, v8ctx) + + // Convert results to JS value + jsVal, err := bridge.JsValue(v8ctx, results) + if err != nil { + return bridge.JsException(v8ctx, "failed to convert results: "+err.Error()) + } + + return jsVal + }) +} + +// batchMessage represents a message from a batch call for channel-based callback handling +type batchMessage struct { + AgentID string // Agent ID that generated this message + Index int // Index of the request in the batch + Message *message.Message // The message object +} + +// parseRequestsForBatch parses the requests array and extracts global callback for batch calls +// Returns the requests array and the global JS callback function (if any) +func (ctx *Context) parseRequestsForBatch(args []*v8go.Value, v8ctx *v8go.Context) ([]interface{}, *v8go.Function) { + var globalCallback *v8go.Function + + // Parse global options (second argument) for global onChunk + if len(args) >= 2 && !args[1].IsUndefined() && !args[1].IsNull() { + globalOptsObj, err := args[1].AsObject() + if err == nil && globalOptsObj != nil { + onChunkVal, _ := globalOptsObj.Get("onChunk") + if onChunkVal != nil && onChunkVal.IsFunction() { + globalCallback, _ = onChunkVal.AsFunction() + } + } + } + + // Parse requests array + if len(args) < 1 || args[0].IsUndefined() || args[0].IsNull() { + return []interface{}{}, globalCallback + } + + requestsObj, err := args[0].AsObject() + if err != nil { + return []interface{}{}, globalCallback + } + + // Get array length + lengthVal, err := requestsObj.Get("length") + if err != nil { + return []interface{}{}, globalCallback + } + + length := int(lengthVal.Integer()) + requests := make([]interface{}, 0, length) + + for i := 0; i < length; i++ { + itemVal, err := requestsObj.GetIdx(uint32(i)) + if err != nil || itemVal.IsUndefined() || itemVal.IsNull() { + continue + } + + // Convert to Go map + goVal, err := bridge.GoValue(itemVal, v8ctx) + if err != nil { + continue + } + + reqMap, ok := goVal.(map[string]interface{}) + if !ok { + continue + } + + // Remove onChunk from per-request options (only global callback is supported) + if opts, ok := reqMap["options"].(map[string]interface{}); ok { + delete(opts, "onChunk") + } + + requests = append(requests, reqMap) + } + + return requests, globalCallback +} + +// BatchMethod represents the type of batch operation +type BatchMethod int + +const ( + BatchMethodAll BatchMethod = iota + BatchMethodAny + BatchMethodRace +) + +// executeBatchWithCallback executes a batch operation with channel-based callback handling +// This ensures V8 thread safety by processing all callbacks in the main goroutine +func (ctx *Context) executeBatchWithCallback( + method BatchMethod, + requests []interface{}, + callback *v8go.Function, + v8ctx *v8go.Context, +) []interface{} { + // Get agent API + agentAPI := ctx.Agent() + if agentAPI == nil { + return []interface{}{} + } + + // If no callback, just execute directly + if callback == nil { + switch method { + case BatchMethodAll: + return agentAPI.All(requests) + case BatchMethodAny: + return agentAPI.Any(requests) + case BatchMethodRace: + return agentAPI.Race(requests) + } + return []interface{}{} + } + + // Check if API supports callbacks + apiWithCb, ok := agentAPI.(AgentAPIWithCallback) + if !ok { + switch method { + case BatchMethodAll: + return agentAPI.All(requests) + case BatchMethodAny: + return agentAPI.Any(requests) + case BatchMethodRace: + return agentAPI.Race(requests) + } + return []interface{}{} + } + + // Create message channel for callback handling + // Use a large buffer (1000) to reduce blocking, with blocking send to guarantee no message loss + msgChan := make(chan batchMessage, 1000) + doneChan := make(chan []interface{}, 1) + + // Create Go handler that sends messages to channel + // Blocking send ensures no message is lost (natural backpressure) + goHandler := func(agentID string, index int, msg *message.Message) int { + msgChan <- batchMessage{AgentID: agentID, Index: index, Message: msg} + return 0 + } + + // Start batch execution in background goroutine + go func() { + defer close(msgChan) + var results []interface{} + + switch method { + case BatchMethodAll: + results = apiWithCb.AllWithHandler(requests, goHandler) + case BatchMethodAny: + results = apiWithCb.AnyWithHandler(requests, goHandler) + case BatchMethodRace: + results = apiWithCb.RaceWithHandler(requests, goHandler) + } + + doneChan <- results + }() + + // Process messages in main goroutine (V8 thread-safe) + for msg := range msgChan { + callJSBatchCallback(v8ctx, callback, msg.AgentID, msg.Index, msg.Message) + } + + // Wait for results + return <-doneChan +} + +// callJSBatchCallback calls the JS callback with batch message parameters +// Must be called from the main V8 goroutine +func callJSBatchCallback(v8ctx *v8go.Context, callback *v8go.Function, agentID string, index int, msg *message.Message) { + if callback == nil || v8ctx == nil || msg == nil { + return + } + + iso := v8ctx.Isolate() + + agentIDVal, err := v8go.NewValue(iso, agentID) + if err != nil { + return + } + + indexVal, err := v8go.NewValue(iso, int32(index)) + if err != nil { + return + } + + // Convert message to JS value + jsMsg, err := bridge.JsValue(v8ctx, msg) + if err != nil { + return + } + + callback.Call(v8ctx.Global(), agentIDVal, indexVal, jsMsg) +} diff --git a/agent/context/jsapi_agent_test.go b/agent/context/jsapi_agent_test.go new file mode 100644 index 00000000..372b4ac4 --- /dev/null +++ b/agent/context/jsapi_agent_test.go @@ -0,0 +1,57 @@ +package context_test + +import ( + stdContext "context" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/yaoapp/yao/agent/context" +) + +func TestContext_Agent_NilFactory(t *testing.T) { + // Reset factory + context.AgentAPIFactory = nil + + ctx := context.New(stdContext.Background(), nil, "test-chat") + agentAPI := ctx.Agent() + assert.Nil(t, agentAPI) +} + +func TestContext_Agent_WithFactory(t *testing.T) { + // Set up a mock factory + var capturedCtx *context.Context + context.AgentAPIFactory = func(ctx *context.Context) context.AgentAPI { + capturedCtx = ctx + return &mockAgentAPI{} + } + defer func() { context.AgentAPIFactory = nil }() + + ctx := context.New(stdContext.Background(), nil, "test-chat") + agentAPI := ctx.Agent() + + require.NotNil(t, agentAPI) + assert.Equal(t, ctx, capturedCtx) +} + +// mockAgentAPI implements context.AgentAPI for testing +type mockAgentAPI struct{} + +func (m *mockAgentAPI) Call(agentID string, messages []interface{}, opts map[string]interface{}) interface{} { + return map[string]interface{}{ + "agent_id": agentID, + "content": "mock response", + } +} + +func (m *mockAgentAPI) All(requests []interface{}) []interface{} { + return []interface{}{} +} + +func (m *mockAgentAPI) Any(requests []interface{}) []interface{} { + return []interface{}{} +} + +func (m *mockAgentAPI) Race(requests []interface{}) []interface{} { + return []interface{}{} +} diff --git a/agent/context/jsapi_agent_v8_test.go b/agent/context/jsapi_agent_v8_test.go new file mode 100644 index 00000000..3d2c402d --- /dev/null +++ b/agent/context/jsapi_agent_v8_test.go @@ -0,0 +1,669 @@ +package context_test + +import ( + stdContext "context" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + v8 "github.com/yaoapp/gou/runtime/v8" + "github.com/yaoapp/yao/agent/context" + "github.com/yaoapp/yao/agent/testutils" + "github.com/yaoapp/yao/openapi/oauth/types" + + // Import assistant package to register AgentAPIFactory + _ "github.com/yaoapp/yao/agent/assistant" +) + +// TestAgent_Call_V8 tests basic ctx.agent.Call() functionality with real V8 execution +func TestAgent_Call_V8(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + // Create authorized info for the context + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + ctx := context.New(stdContext.Background(), authorized, "test-chat-v8-call") + ctx.AssistantID = "tests.agent-caller" + defer ctx.Release() + + res, err := v8.Call(v8.CallOptions{}, ` + function test(ctx) { + try { + const result = ctx.agent.Call( + "tests.simple-greeting", + [{ role: "user", content: "Hello" }] + ); + + return { + success: true, + agent_id: result.agent_id, + has_content: result.content && result.content.length > 0, + has_response: result.response !== undefined, + error: result.error || "" + }; + } catch (error) { + return { success: false, error: error.message }; + } + }`, ctx) + + require.NoError(t, err) + result, ok := res.(map[string]interface{}) + require.True(t, ok, "Result should be a map") + + success, _ := result["success"].(bool) + if !success { + t.Fatalf("Test failed: %v", result["error"]) + } + + assert.Equal(t, "tests.simple-greeting", result["agent_id"]) + + hasContent, _ := result["has_content"].(bool) + assert.True(t, hasContent, "Should have content in response") + + hasResponse, _ := result["has_response"].(bool) + assert.True(t, hasResponse, "Should have response object") + + errorStr, _ := result["error"].(string) + assert.Empty(t, errorStr, "Should not have error") +} + +// TestAgent_Call_WithOptions_V8 tests ctx.agent.Call() with options +func TestAgent_Call_WithOptions_V8(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + ctx := context.New(stdContext.Background(), authorized, "test-chat-v8-options") + ctx.AssistantID = "tests.agent-caller" + defer ctx.Release() + + res, err := v8.Call(v8.CallOptions{}, ` + function test(ctx) { + try { + const result = ctx.agent.Call( + "tests.simple-greeting", + [{ role: "user", content: "Hi there!" }], + { + skip: { + history: true, + trace: true + } + } + ); + + return { + success: true, + agent_id: result.agent_id, + content: result.content || "", + error: result.error || "" + }; + } catch (error) { + return { success: false, error: error.message }; + } + }`, ctx) + + require.NoError(t, err) + result := res.(map[string]interface{}) + + if !result["success"].(bool) { + t.Fatalf("Test failed: %v", result["error"]) + } + + assert.Equal(t, "tests.simple-greeting", result["agent_id"]) + assert.NotEmpty(t, result["content"], "Should have content") +} + +// TestAgent_All_V8 tests ctx.agent.All() for parallel execution +func TestAgent_All_V8(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + ctx := context.New(stdContext.Background(), authorized, "test-chat-v8-all") + ctx.AssistantID = "tests.agent-caller" + defer ctx.Release() + + res, err := v8.Call(v8.CallOptions{}, ` + function test(ctx) { + try { + const results = ctx.agent.All([ + { + agent: "tests.simple-greeting", + messages: [{ role: "user", content: "Hello from request 1" }] + }, + { + agent: "tests.simple-greeting", + messages: [{ role: "user", content: "Hello from request 2" }] + } + ]); + + return { + success: true, + count: results.length, + first_agent: results[0] ? results[0].agent_id : "", + second_agent: results[1] ? results[1].agent_id : "", + first_has_content: results[0] && results[0].content && results[0].content.length > 0, + second_has_content: results[1] && results[1].content && results[1].content.length > 0, + first_error: results[0] ? (results[0].error || "") : "no result", + second_error: results[1] ? (results[1].error || "") : "no result" + }; + } catch (error) { + return { success: false, error: error.message }; + } + }`, ctx) + + require.NoError(t, err) + result := res.(map[string]interface{}) + + if !result["success"].(bool) { + t.Fatalf("Test failed: %v", result["error"]) + } + + assert.Equal(t, float64(2), result["count"]) + assert.Equal(t, "tests.simple-greeting", result["first_agent"]) + assert.Equal(t, "tests.simple-greeting", result["second_agent"]) + assert.True(t, result["first_has_content"].(bool), "First result should have content") + assert.True(t, result["second_has_content"].(bool), "Second result should have content") + assert.Empty(t, result["first_error"], "First result should not have error") + assert.Empty(t, result["second_error"], "Second result should not have error") +} + +// TestAgent_Any_V8 tests ctx.agent.Any() returns on first success +func TestAgent_Any_V8(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + ctx := context.New(stdContext.Background(), authorized, "test-chat-v8-any") + ctx.AssistantID = "tests.agent-caller" + defer ctx.Release() + + res, err := v8.Call(v8.CallOptions{}, ` + function test(ctx) { + try { + const results = ctx.agent.Any([ + { + agent: "tests.simple-greeting", + messages: [{ role: "user", content: "Hello" }] + }, + { + agent: "tests.simple-greeting", + messages: [{ role: "user", content: "Hi" }] + } + ]); + + // At least one result should be successful + let hasSuccess = false; + for (const r of results) { + if (r && r.content && !r.error) { + hasSuccess = true; + break; + } + } + + return { + success: true, + count: results.length, + has_successful_result: hasSuccess + }; + } catch (error) { + return { success: false, error: error.message }; + } + }`, ctx) + + require.NoError(t, err) + result := res.(map[string]interface{}) + + if !result["success"].(bool) { + t.Fatalf("Test failed: %v", result["error"]) + } + + assert.Equal(t, float64(2), result["count"]) + assert.True(t, result["has_successful_result"].(bool), "Should have at least one successful result") +} + +// TestAgent_Race_V8 tests ctx.agent.Race() returns on first completion +func TestAgent_Race_V8(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + ctx := context.New(stdContext.Background(), authorized, "test-chat-v8-race") + ctx.AssistantID = "tests.agent-caller" + defer ctx.Release() + + res, err := v8.Call(v8.CallOptions{}, ` + function test(ctx) { + try { + const results = ctx.agent.Race([ + { + agent: "tests.simple-greeting", + messages: [{ role: "user", content: "Hello" }] + }, + { + agent: "tests.simple-greeting", + messages: [{ role: "user", content: "Hi" }] + } + ]); + + // At least one result should exist (first to complete) + let hasResult = false; + for (const r of results) { + if (r && (r.content || r.error)) { + hasResult = true; + break; + } + } + + return { + success: true, + count: results.length, + has_result: hasResult + }; + } catch (error) { + return { success: false, error: error.message }; + } + }`, ctx) + + require.NoError(t, err) + result := res.(map[string]interface{}) + + if !result["success"].(bool) { + t.Fatalf("Test failed: %v", result["error"]) + } + + assert.Equal(t, float64(2), result["count"]) + assert.True(t, result["has_result"].(bool), "Should have at least one result") +} + +// TestAgent_ErrorHandling_V8 tests error handling when calling non-existent agent +func TestAgent_ErrorHandling_V8(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + ctx := context.New(stdContext.Background(), authorized, "test-chat-v8-error") + ctx.AssistantID = "tests.agent-caller" + defer ctx.Release() + + res, err := v8.Call(v8.CallOptions{}, ` + function test(ctx) { + try { + const result = ctx.agent.Call( + "non-existent-agent", + [{ role: "user", content: "Hello" }] + ); + + return { + success: true, + has_error: result.error && result.error.length > 0, + error_message: result.error || "" + }; + } catch (error) { + return { success: false, error: error.message }; + } + }`, ctx) + + require.NoError(t, err) + result := res.(map[string]interface{}) + + // The call should succeed (no JS exception), but result should contain error + assert.True(t, result["success"].(bool), "JS execution should succeed") + assert.True(t, result["has_error"].(bool), "Result should have error for non-existent agent") + assert.True(t, strings.Contains(result["error_message"].(string), "failed to get agent"), "Error should mention failed to get agent") +} + +// TestAgent_EmptyRequests_V8 tests handling of empty requests array +func TestAgent_EmptyRequests_V8(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + ctx := context.New(stdContext.Background(), authorized, "test-chat-v8-empty") + ctx.AssistantID = "tests.agent-caller" + defer ctx.Release() + + res, err := v8.Call(v8.CallOptions{}, ` + function test(ctx) { + try { + const results = ctx.agent.All([]); + return { + success: true, + count: results.length + }; + } catch (error) { + return { success: false, error: error.message }; + } + }`, ctx) + + require.NoError(t, err) + result := res.(map[string]interface{}) + + assert.True(t, result["success"].(bool)) + assert.Equal(t, float64(0), result["count"]) +} + +// TestAgent_InvalidArguments_V8 tests error handling for invalid arguments +func TestAgent_InvalidArguments_V8(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + ctx := context.New(stdContext.Background(), authorized, "test-chat-v8-invalid") + ctx.AssistantID = "tests.agent-caller" + defer ctx.Release() + + // Test missing arguments + res, err := v8.Call(v8.CallOptions{}, ` + function test(ctx) { + try { + // Call with no arguments should throw + ctx.agent.Call(); + return { success: false, error: "Should have thrown" }; + } catch (error) { + return { success: true, error: error.message }; + } + }`, ctx) + + require.NoError(t, err) + result := res.(map[string]interface{}) + assert.True(t, result["success"].(bool), "Should catch the error") + assert.Contains(t, result["error"].(string), "requires") +} + +// ============================================================================ +// Callback Tests +// ============================================================================ + +// TestAgent_Call_WithCallback_V8 tests ctx.agent.Call() with onChunk callback +func TestAgent_Call_WithCallback_V8(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + ctx := context.New(stdContext.Background(), authorized, "test-chat-v8-callback") + ctx.AssistantID = "tests.agent-caller" + defer ctx.Release() + + res, err := v8.Call(v8.CallOptions{}, ` + function test(ctx) { + try { + const messages = []; + let messageCount = 0; + + const result = ctx.agent.Call( + "tests.simple-greeting", + [{ role: "user", content: "Hello" }], + { + onChunk: (msg) => { + // msg is the SSE message object + messageCount++; + messages.push({ + type: msg.type, + has_props: msg.props !== undefined + }); + return 0; // Continue + } + } + ); + + return { + success: true, + agent_id: result.agent_id, + has_content: result.content && result.content.length > 0, + message_count: messageCount, + received_messages: messages.slice(0, 5), // First 5 messages + error: result.error || "" + }; + } catch (error) { + return { success: false, error: error.message }; + } + }`, ctx) + + require.NoError(t, err) + result, ok := res.(map[string]interface{}) + require.True(t, ok, "Result should be a map") + + success, _ := result["success"].(bool) + if !success { + t.Fatalf("Test failed: %v", result["error"]) + } + + assert.Equal(t, "tests.simple-greeting", result["agent_id"]) + + // Should have received some messages via callback + messageCount, _ := result["message_count"].(float64) + t.Logf("Received %v messages via callback", messageCount) + assert.Greater(t, messageCount, float64(0), "Should have received messages via callback") + + // Check that we received message objects with type and props + receivedMsgs, _ := result["received_messages"].([]interface{}) + if len(receivedMsgs) > 0 { + firstMsg := receivedMsgs[0].(map[string]interface{}) + t.Logf("First message type: %v", firstMsg["type"]) + assert.NotEmpty(t, firstMsg["type"], "Message should have type") + } +} + +// TestAgent_Call_WithCallback_Stop_V8 tests that callback can stop streaming +func TestAgent_Call_WithCallback_Stop_V8(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + ctx := context.New(stdContext.Background(), authorized, "test-chat-v8-callback-stop") + ctx.AssistantID = "tests.agent-caller" + defer ctx.Release() + + res, err := v8.Call(v8.CallOptions{}, ` + function test(ctx) { + try { + let messageCount = 0; + + const result = ctx.agent.Call( + "tests.simple-greeting", + [{ role: "user", content: "Hello" }], + { + onChunk: (msg) => { + messageCount++; + // Stop after receiving 3 messages + if (messageCount >= 3) { + return 1; // Stop + } + return 0; // Continue + } + } + ); + + return { + success: true, + message_count: messageCount, + stopped_early: messageCount <= 5 // Should have stopped early + }; + } catch (error) { + return { success: false, error: error.message }; + } + }`, ctx) + + require.NoError(t, err) + result, ok := res.(map[string]interface{}) + require.True(t, ok, "Result should be a map") + + success, _ := result["success"].(bool) + if !success { + t.Fatalf("Test failed: %v", result["error"]) + } + + messageCount, _ := result["message_count"].(float64) + t.Logf("Received %v messages before stopping", messageCount) + // Note: The exact count may vary based on when the stop is processed +} + +// TestAgent_All_WithGlobalCallback_V8 tests ctx.agent.All() with global onChunk callback +// Uses channel-based callback handling for V8 thread safety +func TestAgent_All_WithGlobalCallback_V8(t *testing.T) { + if testing.Short() { + t.Skip("Skipping integration test in short mode") + } + + testutils.Prepare(t) + defer testutils.Clean(t) + + authorized := &types.AuthorizedInfo{ + Subject: "test-user", + UserID: "test-123", + TenantID: "test-tenant", + } + + ctx := context.New(stdContext.Background(), authorized, "test-chat-v8-all-callback") + ctx.AssistantID = "tests.agent-caller" + defer ctx.Release() + + res, err := v8.Call(v8.CallOptions{}, ` + function test(ctx) { + try { + const messagesByAgent = {}; + + const results = ctx.agent.All( + [ + { + agent: "tests.simple-greeting", + messages: [{ role: "user", content: "Hello from 1" }] + }, + { + agent: "tests.simple-greeting", + messages: [{ role: "user", content: "Hello from 2" }] + } + ], + { + // Global callback receives agentID, index, and message + onChunk: (agentID, index, msg) => { + const key = agentID + "_" + index; + if (!messagesByAgent[key]) { + messagesByAgent[key] = 0; + } + messagesByAgent[key]++; + return 0; + } + } + ); + + return { + success: true, + result_count: results.length, + messages_by_agent: messagesByAgent + }; + } catch (error) { + return { success: false, error: error.message }; + } + }`, ctx) + + require.NoError(t, err) + result, ok := res.(map[string]interface{}) + require.True(t, ok, "Result should be a map") + + success, _ := result["success"].(bool) + if !success { + t.Fatalf("Test failed: %v", result["error"]) + } + + assert.Equal(t, float64(2), result["result_count"]) + + // Should have received messages from both agents + messagesByAgent, _ := result["messages_by_agent"].(map[string]interface{}) + t.Logf("Messages by agent: %v", messagesByAgent) + + // At least one agent should have sent messages + assert.Greater(t, len(messagesByAgent), 0, "Should have received messages from agents") +} diff --git a/agent/context/output.go b/agent/context/output.go index ee3dfdc5..1223ea07 100644 --- a/agent/context/output.go +++ b/agent/context/output.go @@ -15,6 +15,13 @@ import ( // - Sends block_start event when a new BlockID is first encountered // - Records metadata for all sent messages to enable delta inheritance func (ctx *Context) Send(msg *message.Message) error { + // Call OnMessage callback if provided (for ctx.agent.Call with onChunk) + if ctx.Stack != nil && ctx.Stack.Options != nil && ctx.Stack.Options.OnMessage != nil { + if ret := ctx.Stack.Options.OnMessage(msg); ret != 0 { + return nil // Callback requested stop + } + } + out, err := ctx.getOutput() if err != nil { return err diff --git a/agent/context/types.go b/agent/context/types.go index 953a9eb9..215c0612 100644 --- a/agent/context/types.go +++ b/agent/context/types.go @@ -308,8 +308,18 @@ type Options struct { // Metadata for passing custom data to hooks (e.g., scenario selection) Metadata map[string]any `json:"metadata,omitempty"` // Custom metadata passed to Create/Next hooks + + // OnMessage is called for each message sent via ctx.Send() + // Used by ctx.agent.Call with onChunk callback to receive SSE messages + // Returns: 0 = continue, non-zero = stop + OnMessage OnMessageFunc `json:"-"` } +// OnMessageFunc is a callback function for receiving output messages +// Called for each message sent via ctx.Send() - same as SSE messages to client +// Returns: 0 = continue, non-zero = stop sending +type OnMessageFunc func(msg *message.Message) int + // Stack represents the call stack node for tracing agent-to-agent calls // Uses a flat structure to avoid circular references and memory overhead type Stack struct {