Enhance Agent Context with JSAPI Initialization and Message Handling

- Added initialization for the Agent JSAPI factory to support ctx.agent.* methods, improving agent interaction capabilities.
- Introduced a new agent object in the JSAPI context for calling other agents, enhancing modularity.
- Implemented an OnMessage callback in the context options to handle messages sent via ctx.Send(), allowing for more flexible message processing.
This commit is contained in:
Max 2026-01-25 19:42:56 +08:00
parent ef63718941
commit 3a9f32af12
14 changed files with 2557 additions and 0 deletions

View file

@ -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)

View file

@ -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")
}

300
agent/caller/jsapi.go Normal file
View file

@ -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)
}
}

145
agent/caller/jsapi_test.go Normal file
View file

@ -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)
}

View file

@ -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 ""
}

View file

@ -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")
}

44
agent/caller/types.go Normal file
View file

@ -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,
}
}

View file

@ -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)
}

View file

@ -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

View file

@ -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)
}

View file

@ -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{}{}
}

View file

@ -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")
}

View file

@ -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

View file

@ -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 {