diff --git a/cmd/synapbus/main.go b/cmd/synapbus/main.go index 8bda60c..f300955 100644 --- a/cmd/synapbus/main.go +++ b/cmd/synapbus/main.go @@ -23,6 +23,7 @@ import ( "github.com/prometheus/client_golang/prometheus/promhttp" "github.com/spf13/cobra" + "github.com/synapbus/synapbus/internal/a2a" "github.com/synapbus/synapbus/internal/actions" "github.com/synapbus/synapbus/internal/admin" "github.com/synapbus/synapbus/internal/agents" @@ -538,6 +539,14 @@ func runServe(cmd *cobra.Command, args []string) error { r.Mount("/mcp", mcpSrv.Handler()) }) + // A2A Gateway (requires auth: API key, managed key, or OAuth bearer) + a2aTaskStore := a2a.NewA2ATaskStore(db.DB) + a2aGateway := a2a.NewGateway(a2aTaskStore, msgService, agentService) + r.Group(func(r chi.Router) { + r.Use(agents.RequiredAuthMiddlewareWithOAuth(agentService, apiKeyService, oauthProvider)) + r.Post("/a2a", a2aGateway.HandleJSONRPC) + }) + // Create SSE hub and broadcaster for real-time events sseHub := api.NewSSEHub() sseBroadcaster := api.NewSSEBroadcaster(sseHub, agentService, channelService) diff --git a/internal/a2a/gateway.go b/internal/a2a/gateway.go new file mode 100644 index 0000000..8aaa7f0 --- /dev/null +++ b/internal/a2a/gateway.go @@ -0,0 +1,306 @@ +package a2a + +import ( + "context" + "encoding/json" + "fmt" + "io" + "log/slog" + "net/http" + + "github.com/google/uuid" + + "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/messaging" +) + +// MessagingService defines the messaging operations needed by the A2A gateway. +type MessagingService interface { + SendMessage(ctx context.Context, from, to, body string, opts messaging.SendOptions) (*messaging.Message, error) + GetConversation(ctx context.Context, id int64) (*messaging.Conversation, []*messaging.Message, error) +} + +// AgentService defines the agent operations needed by the A2A gateway. +type AgentService interface { + GetAgent(ctx context.Context, name string) (*agents.Agent, error) +} + +// Gateway handles inbound A2A JSON-RPC requests. +type Gateway struct { + taskStore *A2ATaskStore + msgService MessagingService + agentService AgentService + logger *slog.Logger +} + +// NewGateway creates a new A2A gateway. +func NewGateway(taskStore *A2ATaskStore, msgService MessagingService, agentService AgentService) *Gateway { + return &Gateway{ + taskStore: taskStore, + msgService: msgService, + agentService: agentService, + logger: slog.Default().With("component", "a2a-gateway"), + } +} + +// JSON-RPC 2.0 request/response types. + +type jsonRPCRequest struct { + JSONRPC string `json:"jsonrpc"` + ID json.RawMessage `json:"id"` + Method string `json:"method"` + Params json.RawMessage `json:"params,omitempty"` +} + +type jsonRPCResponse struct { + JSONRPC string `json:"jsonrpc"` + ID json.RawMessage `json:"id"` + Result any `json:"result,omitempty"` + Error *jsonRPCError `json:"error,omitempty"` +} + +type jsonRPCError struct { + Code int `json:"code"` + Message string `json:"message"` +} + +// Standard JSON-RPC 2.0 error codes. +const ( + errCodeParse = -32700 + errCodeInvalidReq = -32600 + errCodeNoMethod = -32601 + errCodeInvalidParams = -32602 + errCodeInternal = -32603 +) + +// HandleJSONRPC dispatches incoming JSON-RPC 2.0 requests to the appropriate handler. +func (g *Gateway) HandleJSONRPC(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + http.Error(w, `{"error":"method not allowed"}`, http.StatusMethodNotAllowed) + return + } + + body, err := io.ReadAll(io.LimitReader(r.Body, 1<<20)) // 1 MiB limit + if err != nil { + writeJSONRPC(w, nil, nil, &jsonRPCError{Code: errCodeParse, Message: "failed to read request body"}) + return + } + + var req jsonRPCRequest + if err := json.Unmarshal(body, &req); err != nil { + writeJSONRPC(w, nil, nil, &jsonRPCError{Code: errCodeParse, Message: "invalid JSON"}) + return + } + + if req.JSONRPC != "2.0" { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidReq, Message: "jsonrpc must be \"2.0\""}) + return + } + + // Extract the calling agent from auth context. + callerAgent, ok := agents.AgentFromContext(r.Context()) + callerName := "" + if ok && callerAgent != nil { + callerName = callerAgent.Name + } + + switch req.Method { + case "message.send": + g.handleMessageSend(w, r.Context(), req, callerName) + case "tasks.get": + g.handleTasksGet(w, r.Context(), req) + case "tasks.cancel": + g.handleTasksCancel(w, r.Context(), req) + default: + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeNoMethod, Message: fmt.Sprintf("unknown method: %s", req.Method)}) + } +} + +// message.send params + +type messageSendParams struct { + Message struct { + Body string `json:"body"` + Metadata struct { + TargetAgent string `json:"target_agent"` + } `json:"metadata"` + } `json:"message"` +} + +func (g *Gateway) handleMessageSend(w http.ResponseWriter, ctx context.Context, req jsonRPCRequest, callerName string) { + var params messageSendParams + if err := json.Unmarshal(req.Params, ¶ms); err != nil { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "invalid params: " + err.Error()}) + return + } + + targetAgent := params.Message.Metadata.TargetAgent + if targetAgent == "" { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "params.message.metadata.target_agent is required"}) + return + } + + messageBody := params.Message.Body + if messageBody == "" { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "params.message.body is required"}) + return + } + + // Validate target agent exists. + _, err := g.agentService.GetAgent(ctx, targetAgent) + if err != nil { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: fmt.Sprintf("target agent not found: %s", targetAgent)}) + return + } + + // Create task. + taskID := uuid.New().String() + contextID := uuid.New().String() + + // Determine the sender name for the DM. Use the caller's identity if + // authenticated, otherwise fall back to "a2a-gateway" so SendMessage + // has a non-empty from field. + senderName := callerName + if senderName == "" { + senderName = "a2a-gateway" + } + + // Build metadata containing the a2a_task_id. + metaJSON, _ := json.Marshal(map[string]string{"a2a_task_id": taskID}) + + // Send DM to target agent. + msg, err := g.msgService.SendMessage(ctx, senderName, targetAgent, messageBody, messaging.SendOptions{ + Subject: fmt.Sprintf("A2A Task %s", taskID), + Metadata: string(metaJSON), + }) + if err != nil { + g.logger.Error("failed to send DM for A2A task", "task_id", taskID, "error", err) + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInternal, Message: "failed to deliver message: " + err.Error()}) + return + } + + convID := msg.ConversationID + task := &A2ATask{ + ID: taskID, + ContextID: contextID, + TargetAgent: targetAgent, + SourceAgent: callerName, + ConversationID: &convID, + State: StateSubmitted, + } + + if err := g.taskStore.CreateTask(ctx, task); err != nil { + g.logger.Error("failed to create A2A task", "task_id", taskID, "error", err) + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInternal, Message: "failed to create task"}) + return + } + + g.logger.Info("A2A task created", + "task_id", taskID, + "target_agent", targetAgent, + "source_agent", callerName, + "message_id", msg.ID, + ) + + writeJSONRPC(w, req.ID, task, nil) +} + +// tasks.get params + +type tasksGetParams struct { + ID string `json:"id"` +} + +func (g *Gateway) handleTasksGet(w http.ResponseWriter, ctx context.Context, req jsonRPCRequest) { + var params tasksGetParams + if err := json.Unmarshal(req.Params, ¶ms); err != nil { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "invalid params: " + err.Error()}) + return + } + + if params.ID == "" { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "params.id is required"}) + return + } + + task, err := g.taskStore.GetTask(ctx, params.ID) + if err != nil { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: err.Error()}) + return + } + + // If the task has a conversation and is still SUBMITTED, check whether the + // target agent has replied, which indicates completion. + if task.State == StateSubmitted && task.ConversationID != nil { + _, msgs, err := g.msgService.GetConversation(ctx, *task.ConversationID) + if err == nil && len(msgs) > 1 { + // Check if the target agent sent a reply (any message from target after the first). + for _, m := range msgs[1:] { + if m.FromAgent == task.TargetAgent { + task.State = StateCompleted + _ = g.taskStore.UpdateTaskState(ctx, task.ID, StateCompleted) + break + } + } + } + } + + writeJSONRPC(w, req.ID, task, nil) +} + +// tasks.cancel params + +type tasksCancelParams struct { + ID string `json:"id"` +} + +func (g *Gateway) handleTasksCancel(w http.ResponseWriter, ctx context.Context, req jsonRPCRequest) { + var params tasksCancelParams + if err := json.Unmarshal(req.Params, ¶ms); err != nil { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "invalid params: " + err.Error()}) + return + } + + if params.ID == "" { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: "params.id is required"}) + return + } + + task, err := g.taskStore.GetTask(ctx, params.ID) + if err != nil { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: err.Error()}) + return + } + + // Cannot cancel a terminal task. + if task.State == StateCompleted || task.State == StateCanceled { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInvalidParams, Message: fmt.Sprintf("task is already in terminal state: %s", task.State)}) + return + } + + if err := g.taskStore.UpdateTaskState(ctx, task.ID, StateCanceled); err != nil { + writeJSONRPC(w, req.ID, nil, &jsonRPCError{Code: errCodeInternal, Message: "failed to cancel task"}) + return + } + + task.State = StateCanceled + g.logger.Info("A2A task canceled", "task_id", task.ID) + + writeJSONRPC(w, req.ID, task, nil) +} + +// writeJSONRPC writes a JSON-RPC 2.0 response. +func writeJSONRPC(w http.ResponseWriter, id json.RawMessage, result any, rpcErr *jsonRPCError) { + resp := jsonRPCResponse{ + JSONRPC: "2.0", + ID: id, + Result: result, + Error: rpcErr, + } + w.Header().Set("Content-Type", "application/json") + if rpcErr != nil { + // Use 200 for JSON-RPC errors (per spec), but set result to nil. + resp.Result = nil + } + json.NewEncoder(w).Encode(resp) +} diff --git a/internal/a2a/gateway_test.go b/internal/a2a/gateway_test.go new file mode 100644 index 0000000..858b1b7 --- /dev/null +++ b/internal/a2a/gateway_test.go @@ -0,0 +1,483 @@ +package a2a + +import ( + "bytes" + "context" + "database/sql" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "testing" + + _ "modernc.org/sqlite" + + "github.com/synapbus/synapbus/internal/agents" + "github.com/synapbus/synapbus/internal/messaging" + "github.com/synapbus/synapbus/internal/storage" +) + +// --- test helpers --- + +func newTestDB(t *testing.T) *sql.DB { + t.Helper() + dsn := fmt.Sprintf("file:%s?mode=memory&cache=shared", t.Name()) + db, err := sql.Open("sqlite", dsn) + if err != nil { + t.Fatalf("open database: %v", err) + } + t.Cleanup(func() { db.Close() }) + + if _, err := db.Exec("PRAGMA foreign_keys=ON"); err != nil { + t.Fatalf("enable foreign keys: %v", err) + } + + ctx := context.Background() + if err := storage.RunMigrations(ctx, db); err != nil { + t.Fatalf("run migrations: %v", err) + } + + // Seed a test user for owner_id FK + db.Exec(`INSERT OR IGNORE INTO users (id, username, password_hash, display_name) VALUES (1, 'testowner', 'hash', 'Test Owner')`) + + return db +} + +func seedAgent(t *testing.T, db *sql.DB, name string) { + t.Helper() + _, err := db.Exec( + `INSERT OR IGNORE INTO agents (name, display_name, type, capabilities, owner_id, api_key_hash, status) VALUES (?, ?, 'ai', '{}', 1, 'testhash', 'active')`, + name, name, + ) + if err != nil { + t.Fatalf("seed agent %s: %v", name, err) + } +} + +// mockMsgService implements MessagingService for testing. +type mockMsgService struct { + lastFrom string + lastTo string + lastBody string + lastOpts messaging.SendOptions + sendErr error + returnMsg *messaging.Message + convMsgs []*messaging.Message + getConvErr error +} + +func (m *mockMsgService) SendMessage(_ context.Context, from, to, body string, opts messaging.SendOptions) (*messaging.Message, error) { + m.lastFrom = from + m.lastTo = to + m.lastBody = body + m.lastOpts = opts + if m.sendErr != nil { + return nil, m.sendErr + } + if m.returnMsg != nil { + return m.returnMsg, nil + } + return &messaging.Message{ + ID: 1, + ConversationID: 100, + FromAgent: from, + ToAgent: to, + Body: body, + }, nil +} + +func (m *mockMsgService) GetConversation(_ context.Context, id int64) (*messaging.Conversation, []*messaging.Message, error) { + if m.getConvErr != nil { + return nil, nil, m.getConvErr + } + conv := &messaging.Conversation{ID: id} + return conv, m.convMsgs, nil +} + +// mockAgentService implements AgentService for testing. +type mockAgentService struct { + agents map[string]*agents.Agent +} + +func (m *mockAgentService) GetAgent(_ context.Context, name string) (*agents.Agent, error) { + if a, ok := m.agents[name]; ok { + return a, nil + } + return nil, fmt.Errorf("agent not found: %s", name) +} + +func newTestGateway(t *testing.T) (*Gateway, *mockMsgService, *mockAgentService, *A2ATaskStore) { + t.Helper() + db := newTestDB(t) + seedAgent(t, db, "target-bot") + seedAgent(t, db, "sender-bot") + + taskStore := NewA2ATaskStore(db) + msgSvc := &mockMsgService{} + agentSvc := &mockAgentService{ + agents: map[string]*agents.Agent{ + "target-bot": {ID: 1, Name: "target-bot", Status: "active"}, + "sender-bot": {ID: 2, Name: "sender-bot", Status: "active"}, + }, + } + + gw := NewGateway(taskStore, msgSvc, agentSvc) + return gw, msgSvc, agentSvc, taskStore +} + +func jsonRPCCall(method string, params any) []byte { + p, _ := json.Marshal(params) + req := map[string]any{ + "jsonrpc": "2.0", + "id": 1, + "method": method, + "params": json.RawMessage(p), + } + b, _ := json.Marshal(req) + return b +} + +func doRequest(t *testing.T, gw *Gateway, body []byte, agentCtx *agents.Agent) *httptest.ResponseRecorder { + t.Helper() + req := httptest.NewRequest(http.MethodPost, "/a2a", bytes.NewReader(body)) + req.Header.Set("Content-Type", "application/json") + if agentCtx != nil { + req = req.WithContext(agents.ContextWithAgent(req.Context(), agentCtx)) + } + w := httptest.NewRecorder() + gw.HandleJSONRPC(w, req) + return w +} + +func parseResponse(t *testing.T, w *httptest.ResponseRecorder) jsonRPCResponse { + t.Helper() + var resp jsonRPCResponse + if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { + t.Fatalf("unmarshal response: %v\nbody: %s", err, w.Body.String()) + } + return resp +} + +// --- tests --- + +func TestMessageSend_CreatesTaskAndDM(t *testing.T) { + gw, msgSvc, _, taskStore := newTestGateway(t) + + body := jsonRPCCall("message.send", map[string]any{ + "message": map[string]any{ + "body": "Hello target bot", + "metadata": map[string]string{ + "target_agent": "target-bot", + }, + }, + }) + + callerAgent := &agents.Agent{Name: "sender-bot"} + w := doRequest(t, gw, body, callerAgent) + resp := parseResponse(t, w) + + if resp.Error != nil { + t.Fatalf("unexpected error: code=%d message=%s", resp.Error.Code, resp.Error.Message) + } + + // Verify a task was returned. + resultBytes, _ := json.Marshal(resp.Result) + var task A2ATask + if err := json.Unmarshal(resultBytes, &task); err != nil { + t.Fatalf("unmarshal task result: %v", err) + } + + if task.ID == "" { + t.Error("task ID should not be empty") + } + if task.State != StateSubmitted { + t.Errorf("task state = %q, want %q", task.State, StateSubmitted) + } + if task.TargetAgent != "target-bot" { + t.Errorf("target_agent = %q, want %q", task.TargetAgent, "target-bot") + } + if task.SourceAgent != "sender-bot" { + t.Errorf("source_agent = %q, want %q", task.SourceAgent, "sender-bot") + } + + // Verify the DM was sent. + if msgSvc.lastTo != "target-bot" { + t.Errorf("DM to = %q, want %q", msgSvc.lastTo, "target-bot") + } + if msgSvc.lastFrom != "sender-bot" { + t.Errorf("DM from = %q, want %q", msgSvc.lastFrom, "sender-bot") + } + if msgSvc.lastBody != "Hello target bot" { + t.Errorf("DM body = %q, want %q", msgSvc.lastBody, "Hello target bot") + } + + // Verify metadata contains a2a_task_id. + var meta map[string]string + if err := json.Unmarshal([]byte(msgSvc.lastOpts.Metadata), &meta); err != nil { + t.Fatalf("unmarshal metadata: %v", err) + } + if meta["a2a_task_id"] != task.ID { + t.Errorf("metadata a2a_task_id = %q, want %q", meta["a2a_task_id"], task.ID) + } + + // Verify task is persisted. + stored, err := taskStore.GetTask(context.Background(), task.ID) + if err != nil { + t.Fatalf("GetTask: %v", err) + } + if stored.State != StateSubmitted { + t.Errorf("stored task state = %q, want %q", stored.State, StateSubmitted) + } +} + +func TestTasksGet_ReturnsTask(t *testing.T) { + gw, _, _, taskStore := newTestGateway(t) + + // Create a task directly. + convID := int64(100) + task := &A2ATask{ + ID: "test-task-123", + ContextID: "ctx-123", + TargetAgent: "target-bot", + SourceAgent: "sender-bot", + ConversationID: &convID, + State: StateSubmitted, + } + if err := taskStore.CreateTask(context.Background(), task); err != nil { + t.Fatalf("CreateTask: %v", err) + } + + body := jsonRPCCall("tasks.get", map[string]string{"id": "test-task-123"}) + w := doRequest(t, gw, body, nil) + resp := parseResponse(t, w) + + if resp.Error != nil { + t.Fatalf("unexpected error: code=%d message=%s", resp.Error.Code, resp.Error.Message) + } + + resultBytes, _ := json.Marshal(resp.Result) + var got A2ATask + if err := json.Unmarshal(resultBytes, &got); err != nil { + t.Fatalf("unmarshal task result: %v", err) + } + + if got.ID != "test-task-123" { + t.Errorf("task ID = %q, want %q", got.ID, "test-task-123") + } + if got.State != StateSubmitted { + t.Errorf("task state = %q, want %q", got.State, StateSubmitted) + } +} + +func TestTasksGet_CompletesOnReply(t *testing.T) { + gw, msgSvc, _, taskStore := newTestGateway(t) + + // Create a task. + convID := int64(100) + task := &A2ATask{ + ID: "task-reply-test", + ContextID: "ctx-456", + TargetAgent: "target-bot", + SourceAgent: "sender-bot", + ConversationID: &convID, + State: StateSubmitted, + } + if err := taskStore.CreateTask(context.Background(), task); err != nil { + t.Fatalf("CreateTask: %v", err) + } + + // Simulate the target agent having replied. + msgSvc.convMsgs = []*messaging.Message{ + {ID: 1, FromAgent: "sender-bot", Body: "Hello"}, + {ID: 2, FromAgent: "target-bot", Body: "Reply from target"}, + } + + body := jsonRPCCall("tasks.get", map[string]string{"id": "task-reply-test"}) + w := doRequest(t, gw, body, nil) + resp := parseResponse(t, w) + + if resp.Error != nil { + t.Fatalf("unexpected error: code=%d message=%s", resp.Error.Code, resp.Error.Message) + } + + resultBytes, _ := json.Marshal(resp.Result) + var got A2ATask + if err := json.Unmarshal(resultBytes, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.State != StateCompleted { + t.Errorf("task state = %q, want %q (target agent replied)", got.State, StateCompleted) + } +} + +func TestTasksCancel_TransitionsToCanceled(t *testing.T) { + gw, _, _, taskStore := newTestGateway(t) + + task := &A2ATask{ + ID: "task-cancel-test", + ContextID: "ctx-789", + TargetAgent: "target-bot", + SourceAgent: "sender-bot", + State: StateSubmitted, + } + if err := taskStore.CreateTask(context.Background(), task); err != nil { + t.Fatalf("CreateTask: %v", err) + } + + body := jsonRPCCall("tasks.cancel", map[string]string{"id": "task-cancel-test"}) + w := doRequest(t, gw, body, nil) + resp := parseResponse(t, w) + + if resp.Error != nil { + t.Fatalf("unexpected error: code=%d message=%s", resp.Error.Code, resp.Error.Message) + } + + resultBytes, _ := json.Marshal(resp.Result) + var got A2ATask + if err := json.Unmarshal(resultBytes, &got); err != nil { + t.Fatalf("unmarshal: %v", err) + } + + if got.State != StateCanceled { + t.Errorf("task state = %q, want %q", got.State, StateCanceled) + } + + // Verify persisted state. + stored, err := taskStore.GetTask(context.Background(), "task-cancel-test") + if err != nil { + t.Fatalf("GetTask: %v", err) + } + if stored.State != StateCanceled { + t.Errorf("stored state = %q, want %q", stored.State, StateCanceled) + } +} + +func TestTasksCancel_TerminalStateError(t *testing.T) { + gw, _, _, taskStore := newTestGateway(t) + + task := &A2ATask{ + ID: "task-already-done", + ContextID: "ctx-done", + TargetAgent: "target-bot", + State: StateCompleted, + } + if err := taskStore.CreateTask(context.Background(), task); err != nil { + t.Fatalf("CreateTask: %v", err) + } + + body := jsonRPCCall("tasks.cancel", map[string]string{"id": "task-already-done"}) + w := doRequest(t, gw, body, nil) + resp := parseResponse(t, w) + + if resp.Error == nil { + t.Fatal("expected error for canceling terminal task") + } + if resp.Error.Code != errCodeInvalidParams { + t.Errorf("error code = %d, want %d", resp.Error.Code, errCodeInvalidParams) + } +} + +func TestMessageSend_NonExistentAgent(t *testing.T) { + gw, _, _, _ := newTestGateway(t) + + body := jsonRPCCall("message.send", map[string]any{ + "message": map[string]any{ + "body": "Hello ghost", + "metadata": map[string]string{ + "target_agent": "does-not-exist", + }, + }, + }) + + w := doRequest(t, gw, body, nil) + resp := parseResponse(t, w) + + if resp.Error == nil { + t.Fatal("expected error for non-existent target agent") + } + if resp.Error.Code != errCodeInvalidParams { + t.Errorf("error code = %d, want %d", resp.Error.Code, errCodeInvalidParams) + } +} + +func TestInvalidJSONRPC(t *testing.T) { + gw, _, _, _ := newTestGateway(t) + + tests := []struct { + name string + body string + }{ + { + name: "not JSON", + body: "this is not json", + }, + { + name: "wrong jsonrpc version", + body: `{"jsonrpc":"1.0","id":1,"method":"message.send","params":{}}`, + }, + { + name: "unknown method", + body: `{"jsonrpc":"2.0","id":1,"method":"unknown.method","params":{}}`, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/a2a", bytes.NewBufferString(tt.body)) + req.Header.Set("Content-Type", "application/json") + w := httptest.NewRecorder() + gw.HandleJSONRPC(w, req) + + var resp jsonRPCResponse + if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil { + t.Fatalf("unmarshal response: %v\nbody: %s", err, w.Body.String()) + } + if resp.Error == nil { + t.Error("expected error response") + } + }) + } +} + +func TestUnauthenticatedRequest_NoAgentContext(t *testing.T) { + // This tests that message.send works even without an authenticated agent + // in context (caller is anonymous), using "a2a-gateway" as the sender. + gw, msgSvc, _, _ := newTestGateway(t) + + body := jsonRPCCall("message.send", map[string]any{ + "message": map[string]any{ + "body": "Hello from anonymous", + "metadata": map[string]string{ + "target_agent": "target-bot", + }, + }, + }) + + // No agent context — simulates unauthenticated-at-gateway-level + // (in practice, the auth middleware would block this; this tests the + // gateway's fallback behavior). + w := doRequest(t, gw, body, nil) + resp := parseResponse(t, w) + + if resp.Error != nil { + t.Fatalf("unexpected error: code=%d message=%s", resp.Error.Code, resp.Error.Message) + } + + // Should use "a2a-gateway" as sender when no caller agent. + if msgSvc.lastFrom != "a2a-gateway" { + t.Errorf("DM from = %q, want %q", msgSvc.lastFrom, "a2a-gateway") + } +} + +func TestHTTPMethodNotAllowed(t *testing.T) { + gw, _, _, _ := newTestGateway(t) + + req := httptest.NewRequest(http.MethodGet, "/a2a", nil) + w := httptest.NewRecorder() + gw.HandleJSONRPC(w, req) + + if w.Code != http.StatusMethodNotAllowed { + t.Errorf("status = %d, want %d", w.Code, http.StatusMethodNotAllowed) + } +} diff --git a/internal/a2a/taskstore.go b/internal/a2a/taskstore.go new file mode 100644 index 0000000..91187db --- /dev/null +++ b/internal/a2a/taskstore.go @@ -0,0 +1,132 @@ +// Package a2a provides the A2A (Agent-to-Agent) inbound gateway for SynapBus. +// External A2A-compliant agents can send tasks to SynapBus agents via JSON-RPC. +package a2a + +import ( + "context" + "database/sql" + "fmt" + "time" +) + +// Task states following the A2A protocol. +const ( + StateSubmitted = "SUBMITTED" + StateCompleted = "COMPLETED" + StateCanceled = "CANCELED" +) + +// A2ATask represents an inbound A2A task tracked by the gateway. +type A2ATask struct { + ID string `json:"id"` + ContextID string `json:"context_id"` + TargetAgent string `json:"target_agent"` + SourceAgent string `json:"source_agent"` + ConversationID *int64 `json:"conversation_id,omitempty"` + State string `json:"state"` + CreatedAt time.Time `json:"created_at"` + UpdatedAt time.Time `json:"updated_at"` +} + +// A2ATaskStore provides CRUD operations for A2A tasks backed by SQLite. +type A2ATaskStore struct { + db *sql.DB +} + +// NewA2ATaskStore creates a new task store. +func NewA2ATaskStore(db *sql.DB) *A2ATaskStore { + return &A2ATaskStore{db: db} +} + +// CreateTask inserts a new A2A task into the database. +func (s *A2ATaskStore) CreateTask(ctx context.Context, task *A2ATask) error { + _, err := s.db.ExecContext(ctx, + `INSERT INTO a2a_tasks (id, context_id, target_agent, source_agent, conversation_id, state, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, CURRENT_TIMESTAMP, CURRENT_TIMESTAMP)`, + task.ID, task.ContextID, task.TargetAgent, task.SourceAgent, task.ConversationID, task.State, + ) + if err != nil { + return fmt.Errorf("insert a2a task: %w", err) + } + return nil +} + +// GetTask returns an A2A task by its ID. +func (s *A2ATaskStore) GetTask(ctx context.Context, id string) (*A2ATask, error) { + var task A2ATask + var conversationID sql.NullInt64 + err := s.db.QueryRowContext(ctx, + `SELECT id, context_id, target_agent, source_agent, conversation_id, state, created_at, updated_at + FROM a2a_tasks WHERE id = ?`, id, + ).Scan(&task.ID, &task.ContextID, &task.TargetAgent, &task.SourceAgent, + &conversationID, &task.State, &task.CreatedAt, &task.UpdatedAt) + if err != nil { + if err == sql.ErrNoRows { + return nil, fmt.Errorf("a2a task not found: %s", id) + } + return nil, fmt.Errorf("get a2a task: %w", err) + } + if conversationID.Valid { + task.ConversationID = &conversationID.Int64 + } + return &task, nil +} + +// UpdateTaskState transitions a task to a new state. +func (s *A2ATaskStore) UpdateTaskState(ctx context.Context, id, state string) error { + result, err := s.db.ExecContext(ctx, + `UPDATE a2a_tasks SET state = ?, updated_at = CURRENT_TIMESTAMP WHERE id = ?`, + state, id, + ) + if err != nil { + return fmt.Errorf("update a2a task state: %w", err) + } + rows, err := result.RowsAffected() + if err != nil { + return fmt.Errorf("get rows affected: %w", err) + } + if rows == 0 { + return fmt.Errorf("a2a task not found: %s", id) + } + return nil +} + +// ListTasks returns tasks for a target agent, optionally filtered by state. +func (s *A2ATaskStore) ListTasks(ctx context.Context, targetAgent, state string) ([]*A2ATask, error) { + var query string + var args []any + + if state != "" { + query = `SELECT id, context_id, target_agent, source_agent, conversation_id, state, created_at, updated_at + FROM a2a_tasks WHERE target_agent = ? AND state = ? ORDER BY created_at DESC` + args = []any{targetAgent, state} + } else { + query = `SELECT id, context_id, target_agent, source_agent, conversation_id, state, created_at, updated_at + FROM a2a_tasks WHERE target_agent = ? ORDER BY created_at DESC` + args = []any{targetAgent} + } + + rows, err := s.db.QueryContext(ctx, query, args...) + if err != nil { + return nil, fmt.Errorf("list a2a tasks: %w", err) + } + defer rows.Close() + + var tasks []*A2ATask + for rows.Next() { + var task A2ATask + var conversationID sql.NullInt64 + if err := rows.Scan(&task.ID, &task.ContextID, &task.TargetAgent, &task.SourceAgent, + &conversationID, &task.State, &task.CreatedAt, &task.UpdatedAt); err != nil { + return nil, fmt.Errorf("scan a2a task: %w", err) + } + if conversationID.Valid { + task.ConversationID = &conversationID.Int64 + } + tasks = append(tasks, &task) + } + if tasks == nil { + tasks = []*A2ATask{} + } + return tasks, rows.Err() +} diff --git a/internal/storage/schema/010_a2a_tasks.sql b/internal/storage/schema/010_a2a_tasks.sql new file mode 100644 index 0000000..7459d9b --- /dev/null +++ b/internal/storage/schema/010_a2a_tasks.sql @@ -0,0 +1,14 @@ +-- A2A inbound gateway: task tracking for external A2A agents sending tasks to SynapBus agents. +CREATE TABLE IF NOT EXISTS a2a_tasks ( + id TEXT PRIMARY KEY, + context_id TEXT NOT NULL, + target_agent TEXT NOT NULL, + source_agent TEXT DEFAULT '', + conversation_id INTEGER, + state TEXT NOT NULL DEFAULT 'SUBMITTED', + created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE INDEX IF NOT EXISTS idx_a2a_tasks_target ON a2a_tasks(target_agent); +CREATE INDEX IF NOT EXISTS idx_a2a_tasks_state ON a2a_tasks(state); diff --git a/schema/010_a2a_tasks.sql b/schema/010_a2a_tasks.sql new file mode 100644 index 0000000..7459d9b --- /dev/null +++ b/schema/010_a2a_tasks.sql @@ -0,0 +1,14 @@ +-- A2A inbound gateway: task tracking for external A2A agents sending tasks to SynapBus agents. +CREATE TABLE IF NOT EXISTS a2a_tasks ( + id TEXT PRIMARY KEY, + context_id TEXT NOT NULL, + target_agent TEXT NOT NULL, + source_agent TEXT DEFAULT '', + conversation_id INTEGER, + state TEXT NOT NULL DEFAULT 'SUBMITTED', + created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP, + updated_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP +); + +CREATE INDEX IF NOT EXISTS idx_a2a_tasks_target ON a2a_tasks(target_agent); +CREATE INDEX IF NOT EXISTS idx_a2a_tasks_state ON a2a_tasks(state);