diff --git a/go/core/internal/a2a/a2a_handler_mux.go b/go/core/internal/a2a/a2a_handler_mux.go index 26dffe949..311681df7 100644 --- a/go/core/internal/a2a/a2a_handler_mux.go +++ b/go/core/internal/a2a/a2a_handler_mux.go @@ -58,13 +58,12 @@ func NewA2AHttpMux(agentPathPrefix, sandboxPathPrefix string, authenticator auth // newTaskQueryHandlers builds the request handler and the legacy (v0) JSON-RPC // handler for one agent. kagent persists tasks and is their source of truth, -// so with a store ListTasks is served from it instead of proxying to the agent -// runtime (whose legacy 0.3 transport returns ErrUnsupportedOperation for it); -// every other method, GetTask included, still delegates to the passthrough -// proxy. v0 has no native tasks/list, so the legacy handler is wrapped to -// serve that method from the store too (lowercase TaskState). Without a store -// both wires keep their native behavior, including v0's method-not-found for -// tasks/list. +// so with a store GetTask and ListTasks are served from it instead of proxying +// to the agent runtime; GetTask never falls through on a miss. v0 has no native +// tasks/list, so the legacy handler is +// wrapped to serve that method from the store too (lowercase TaskState). +// Without a store both wires keep their native behavior, including v0's +// method-not-found for tasks/list. func newTaskQueryHandlers(requestHandler a2asrv.RequestHandler, store TaskStore) (a2asrv.RequestHandler, http.Handler) { if store == nil { return requestHandler, a2av0.NewJSONRPCHandler(requestHandler) diff --git a/go/core/internal/a2a/task_query_store.go b/go/core/internal/a2a/task_query_store.go index a121af135..c2ed2c867 100644 --- a/go/core/internal/a2a/task_query_store.go +++ b/go/core/internal/a2a/task_query_store.go @@ -4,6 +4,7 @@ import ( "cmp" "context" "encoding/base64" + "errors" "fmt" "slices" "strconv" @@ -19,19 +20,19 @@ const ( maxTaskPageSize = 100 ) -// TaskStore is the subset of the persistent store ListTasks reads from. +// TaskStore is the subset of the persistent store task queries read from. // *database.Client satisfies it. GetSession errors (including a missing or // other-user session) surface to the caller. type TaskStore interface { + GetTask(ctx context.Context, taskID, userID string) (*a2atype.Task, error) GetSession(ctx context.Context, sessionID, userID string) (*dbpkg.Session, error) ListSessions(ctx context.Context, userID string) ([]dbpkg.Session, error) ListTasksForSession(ctx context.Context, sessionID, userID string) ([]*a2atype.Task, error) } -// storeTaskQueryHandler answers ListTasks from kagent's task store, which is -// the source of truth for persisted tasks. Every other method (including -// GetTask, which already resolves to the same store via the passthrough) is -// delegated to the embedded handler unchanged. +// storeTaskQueryHandler answers GetTask and ListTasks from kagent's task +// store, which is the source of truth for persisted tasks. GetTask never falls +// through to the embedded handler; every other method is delegated unchanged. type storeTaskQueryHandler struct { a2asrv.RequestHandler store TaskStore @@ -41,6 +42,22 @@ func newStoreTaskQueryHandler(delegate a2asrv.RequestHandler, store TaskStore) * return &storeTaskQueryHandler{RequestHandler: delegate, store: store} } +func (h *storeTaskQueryHandler) GetTask(ctx context.Context, req *a2atype.GetTaskRequest) (*a2atype.Task, error) { + userID := callerUserID(ctx) + if userID == "" { + return nil, a2atype.ErrTaskNotFound + } + + task, err := h.store.GetTask(ctx, string(req.ID), userID) + if err != nil { + if errors.Is(err, dbpkg.ErrNotFound) { + return nil, a2atype.ErrTaskNotFound + } + return nil, err + } + return shapeTask(task, req.HistoryLength, true), nil +} + // callerUserID returns the authenticated principal's user id, or "" when the // request carries no user identity. Task queries are scoped to this id so a // caller can never read another user's tasks. diff --git a/go/core/internal/a2a/task_query_store_test.go b/go/core/internal/a2a/task_query_store_test.go index d2c650590..d6526dcbb 100644 --- a/go/core/internal/a2a/task_query_store_test.go +++ b/go/core/internal/a2a/task_query_store_test.go @@ -21,8 +21,10 @@ import ( // fakeTaskStore is an in-memory TaskStore. Sessions are keyed by (id, userID) // so cross-user isolation is exercised the same way the real store enforces it. type fakeTaskStore struct { - sessions map[string]dbpkg.Session // key: sessionID -> session (carries UserID) - tasks map[string][]*a2atype.Task + sessions map[string]dbpkg.Session // key: sessionID -> session (carries UserID) + tasks map[string][]*a2atype.Task + getTaskIDs []string + getTaskUsers []string } func newFakeStore() *fakeTaskStore { @@ -40,6 +42,23 @@ func (f *fakeTaskStore) addTask(sessionID string, task *a2atype.Task) { f.tasks[sessionID] = append(f.tasks[sessionID], task) } +func (f *fakeTaskStore) GetTask(_ context.Context, taskID, userID string) (*a2atype.Task, error) { + f.getTaskIDs = append(f.getTaskIDs, taskID) + f.getTaskUsers = append(f.getTaskUsers, userID) + for sessionID, tasks := range f.tasks { + session, ok := f.sessions[sessionID] + if !ok || session.UserID != userID { + continue + } + for _, task := range tasks { + if string(task.ID) == taskID { + return task, nil + } + } + } + return nil, fmt.Errorf("task %s for user %s: %w", taskID, userID, dbpkg.ErrNotFound) +} + func (f *fakeTaskStore) GetSession(_ context.Context, sessionID, userID string) (*dbpkg.Session, error) { s, ok := f.sessions[sessionID] if !ok || s.UserID != userID { @@ -102,6 +121,106 @@ func storeWith(t *testing.T, user, session string, tasks ...*a2atype.Task) *stor return newStoreTaskQueryHandler(&PassthroughRequestHandler{}, store) } +type recordingGetTaskDelegate struct { + a2asrv.RequestHandler + calls int + task *a2atype.Task + err error +} + +func (d *recordingGetTaskDelegate) GetTask(context.Context, *a2atype.GetTaskRequest) (*a2atype.Task, error) { + d.calls++ + return d.task, d.err +} + +func TestGetTask_PersistentHitShapesHistoryAndIncludesArtifacts(t *testing.T) { + stored := newTask("t1", "s1", a2atype.TaskStateCompleted, 4, 2) + wantStored := newTask("t1", "s1", a2atype.TaskStateCompleted, 4, 2) + store := newFakeStore() + store.addSession("s1", "alice") + store.addTask("s1", stored) + delegate := &recordingGetTaskDelegate{} + h := newStoreTaskQueryHandler(delegate, store) + historyLength := 2 + + got, err := h.GetTask(userCtx("alice"), &a2atype.GetTaskRequest{ID: "t1", HistoryLength: &historyLength}) + require.NoError(t, err) + require.Equal(t, a2atype.TaskID("t1"), got.ID) + require.Len(t, got.History, 2) + require.Equal(t, []string{"t1-msg-2", "t1-msg-3"}, []string{got.History[0].ID, got.History[1].ID}) + require.Len(t, got.Artifacts, 2) + require.Equal(t, []a2atype.ArtifactID{"t1-art-0", "t1-art-1"}, []a2atype.ArtifactID{got.Artifacts[0].ID, got.Artifacts[1].ID}) + require.Equal(t, []string{"t1"}, store.getTaskIDs) + require.Equal(t, []string{"alice"}, store.getTaskUsers) + require.Zero(t, delegate.calls) + require.Equal(t, wantStored, stored) +} + +func TestGetTask_OtherOwnerReturnsNotFoundWithoutDelegation(t *testing.T) { + store := newFakeStore() + store.addSession("s1", "alice") + store.addTask("s1", newTask("persisted", "s1", a2atype.TaskStateCompleted, 1, 1)) + delegate := &recordingGetTaskDelegate{task: newTask("runtime", "s1", a2atype.TaskStateWorking, 0, 0)} + h := newStoreTaskQueryHandler(delegate, store) + req := &a2atype.GetTaskRequest{ID: "persisted"} + + got, err := h.GetTask(userCtx("mallory"), req) + require.Nil(t, got) + require.ErrorIs(t, err, a2atype.ErrTaskNotFound) + require.Equal(t, []string{"mallory"}, store.getTaskUsers) + require.Zero(t, delegate.calls) +} + +func TestGetTask_ShareContextStillUsesAuthenticatedCallerWithoutDelegation(t *testing.T) { + store := newFakeStore() + store.addSession("s1", "alice") + store.addTask("s1", newTask("persisted", "s1", a2atype.TaskStateCompleted, 1, 1)) + delegate := &recordingGetTaskDelegate{task: newTask("runtime", "s1", a2atype.TaskStateWorking, 0, 0)} + h := newStoreTaskQueryHandler(delegate, store) + ctx := auth.ShareContextTo(userCtx("bob"), &auth.ShareContext{SessionID: "s1", UserID: "alice"}) + + got, err := h.GetTask(ctx, &a2atype.GetTaskRequest{ID: "persisted"}) + require.Nil(t, got) + require.ErrorIs(t, err, a2atype.ErrTaskNotFound) + require.Equal(t, []string{"bob"}, store.getTaskUsers) + require.Zero(t, delegate.calls) +} + +func TestGetTask_PersistentMissReturnsNotFoundWithoutDelegation(t *testing.T) { + delegate := &recordingGetTaskDelegate{task: newTask("runtime", "s1", a2atype.TaskStateWorking, 0, 0)} + h := newStoreTaskQueryHandler(delegate, failingTaskStore{err: fmt.Errorf("lookup: %w", dbpkg.ErrNotFound)}) + + got, err := h.GetTask(userCtx("alice"), &a2atype.GetTaskRequest{ID: "missing"}) + require.Nil(t, got) + require.ErrorIs(t, err, a2atype.ErrTaskNotFound) + require.Zero(t, delegate.calls) +} + +func TestGetTask_AbsentIdentityReturnsNotFoundWithoutDelegation(t *testing.T) { + store := newFakeStore() + delegate := &recordingGetTaskDelegate{task: newTask("runtime", "s1", a2atype.TaskStateWorking, 0, 0)} + h := newStoreTaskQueryHandler(delegate, store) + req := &a2atype.GetTaskRequest{ID: "runtime"} + + got, err := h.GetTask(context.Background(), req) + require.Nil(t, got) + require.ErrorIs(t, err, a2atype.ErrTaskNotFound) + require.Empty(t, store.getTaskIDs) + require.Empty(t, store.getTaskUsers) + require.Zero(t, delegate.calls) +} + +func TestGetTask_BackendFailurePropagatesWithoutDelegation(t *testing.T) { + backendErr := fmt.Errorf("database connection refused") + delegate := &recordingGetTaskDelegate{} + h := newStoreTaskQueryHandler(delegate, failingTaskStore{err: backendErr}) + + got, err := h.GetTask(userCtx("alice"), &a2atype.GetTaskRequest{ID: "t1"}) + require.Nil(t, got) + require.ErrorIs(t, err, backendErr) + require.Zero(t, delegate.calls) +} + func TestListTasks_Pagination(t *testing.T) { tasks := []*a2atype.Task{ newTask("t1", "s1", a2atype.TaskStateWorking, 0, 0), @@ -266,6 +385,10 @@ func TestListTasks_AcrossAllUserSessions(t *testing.T) { // failingTaskStore fails every read with a backend error. type failingTaskStore struct{ err error } +func (f failingTaskStore) GetTask(context.Context, string, string) (*a2atype.Task, error) { + return nil, f.err +} + func (f failingTaskStore) GetSession(context.Context, string, string) (*dbpkg.Session, error) { return nil, f.err } @@ -350,6 +473,41 @@ func TestWire_ListTasksStateCasing(t *testing.T) { require.Equal(t, "input-required", v0state) } +func TestWire_GetTaskStoreErrorSemantics(t *testing.T) { + tests := []struct { + name string + storeErr error + wantCode float64 + }{ + {name: "not found", storeErr: fmt.Errorf("lookup: %w", dbpkg.ErrNotFound), wantCode: -32001}, + {name: "backend failure", storeErr: fmt.Errorf("database connection refused"), wantCode: -32603}, + } + wires := []struct { + name string + body string + new func(a2asrv.RequestHandler) http.Handler + }{ + {name: "v1", body: `{"jsonrpc":"2.0","id":1,"method":"GetTask","params":{"id":"missing"}}`, new: func(h a2asrv.RequestHandler) http.Handler { return a2asrv.NewJSONRPCHandler(h) }}, + {name: "v0", body: `{"jsonrpc":"2.0","id":1,"method":"tasks/get","params":{"id":"missing"}}`, new: func(h a2asrv.RequestHandler) http.Handler { return a2av0.NewJSONRPCHandler(h) }}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + for _, wire := range wires { + t.Run(wire.name, func(t *testing.T) { + delegate := &recordingGetTaskDelegate{task: newTask("runtime", "s1", a2atype.TaskStateWorking, 0, 0)} + h := newStoreTaskQueryHandler(delegate, failingTaskStore{err: tt.storeErr}) + + resp := rpcCall(t, withUser(wire.new(h), "alice"), wire.body) + errObj := resp["error"].(map[string]any) + require.Equal(t, tt.wantCode, errObj["code"]) + require.Zero(t, delegate.calls) + }) + } + }) + } +} + func TestWire_V0UnknownMethodDelegates(t *testing.T) { _, v0 := wireHandlers("alice") // A method the interceptor does not own must fall through to the v0 handler,