diff --git a/go/adk/pkg/a2a/executor.go b/go/adk/pkg/a2a/executor.go index 7b8fb95e6..3ecd08962 100644 --- a/go/adk/pkg/a2a/executor.go +++ b/go/adk/pkg/a2a/executor.go @@ -119,6 +119,16 @@ func (u *userIDInterceptor) Before(ctx context.Context, callCtx *a2asrv.CallCont return ctx, nil, nil } +// CallerUserID returns the authenticated user name from the a2asrv +// CallContext attached to ctx, or "" if none is set. +func CallerUserID(ctx context.Context) string { + callCtx, ok := a2asrv.CallContextFrom(ctx) + if !ok || callCtx.User == nil { + return "" + } + return callCtx.User.Name +} + // Execute applies kagent-specific request setup and delegates event generation // to the upstream ADK executor, which streams output as artifact updates. func (e *KAgentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.ExecutorContext) iter.Seq2[a2atype.Event, error] { @@ -129,8 +139,8 @@ func (e *KAgentExecutor) Execute(ctx context.Context, reqCtx *a2asrv.ExecutorCon } userID := "A2A_USER_" + reqCtx.ContextID - if callCtx, ok := a2asrv.CallContextFrom(ctx); ok && callCtx.User != nil && callCtx.User.Name != "" { - userID = callCtx.User.Name + if id := CallerUserID(ctx); id != "" { + userID = id } sessionID := reqCtx.ContextID diff --git a/go/adk/pkg/taskstore/store.go b/go/adk/pkg/taskstore/store.go index 4f090a50a..013e45017 100644 --- a/go/adk/pkg/taskstore/store.go +++ b/go/adk/pkg/taskstore/store.go @@ -8,6 +8,7 @@ import ( a2atype "github.com/a2aproject/a2a-go/v2/a2a" "github.com/a2aproject/a2a-go/v2/a2apb/v1/pbconv" a2ataskstore "github.com/a2aproject/a2a-go/v2/a2asrv/taskstore" + "github.com/kagent-dev/kagent/go/adk/pkg/a2a" "github.com/kagent-dev/kagent/go/adk/pkg/controllerclient" apiv1alpha1 "github.com/kagent-dev/kagent/go/api/gen/kagent/api/v1alpha1" "google.golang.org/grpc/codes" @@ -30,6 +31,12 @@ func NewKAgentTaskStore(client *controllerclient.Client) *KAgentTaskStore { return &KAgentTaskStore{client: client} } +// userID resolves the caller for a TaskStore call; auth.WithUserID's value +// doesn't reach this goroutine. +func userID(ctx context.Context) string { + return a2a.CallerUserID(ctx) +} + func (s *KAgentTaskStore) saveTask(ctx context.Context, task *a2atype.Task) (a2ataskstore.TaskVersion, error) { if task == nil { return a2ataskstore.TaskVersionMissing, fmt.Errorf("task cannot be nil") @@ -52,7 +59,7 @@ func (s *KAgentTaskStore) saveTask(ctx context.Context, task *a2atype.Task) (a2a if err != nil { return a2ataskstore.TaskVersionMissing, fmt.Errorf("encode task: %w", err) } - callContext, cancel := s.client.CallContext(ctx, "") + callContext, cancel := s.client.CallContext(ctx, userID(ctx)) defer cancel() _, err = s.client.TaskService().UpsertTask(callContext, &apiv1alpha1.UpsertTaskRequest{Task: encoded}) if err != nil { @@ -77,7 +84,7 @@ func (s *KAgentTaskStore) Update(ctx context.Context, update *a2ataskstore.Updat // Get implements taskstore.Store. func (s *KAgentTaskStore) Get(ctx context.Context, taskID a2atype.TaskID) (*a2ataskstore.StoredTask, error) { - callContext, cancel := s.client.CallContext(ctx, "") + callContext, cancel := s.client.CallContext(ctx, userID(ctx)) defer cancel() response, err := s.client.TaskService().GetTask(callContext, &apiv1alpha1.GetTaskRequest{TaskId: string(taskID)}) if err != nil { @@ -118,7 +125,7 @@ func (s *KAgentTaskStore) List(ctx context.Context, req *a2atype.ListTasksReques return &a2atype.ListTasksResponse{Tasks: []*a2atype.Task{}, PageSize: pageSize}, nil } - callContext, cancel := s.client.CallContext(ctx, "") + callContext, cancel := s.client.CallContext(ctx, userID(ctx)) defer cancel() response, err := s.client.TaskService().ListTasks(callContext, &apiv1alpha1.ListTasksRequest{SessionId: req.ContextID}) if err != nil { diff --git a/go/adk/pkg/taskstore/store_test.go b/go/adk/pkg/taskstore/store_test.go index 77795221c..873306878 100644 --- a/go/adk/pkg/taskstore/store_test.go +++ b/go/adk/pkg/taskstore/store_test.go @@ -8,6 +8,7 @@ import ( a2a "github.com/a2aproject/a2a-go/v2/a2a" a2apb "github.com/a2aproject/a2a-go/v2/a2apb/v1" "github.com/a2aproject/a2a-go/v2/a2apb/v1/pbconv" + "github.com/a2aproject/a2a-go/v2/a2asrv" a2ataskstore "github.com/a2aproject/a2a-go/v2/a2asrv/taskstore" "github.com/kagent-dev/kagent/go/adk/pkg/auth" "github.com/kagent-dev/kagent/go/adk/pkg/controllerclient" @@ -124,6 +125,23 @@ func TestGetDecodesCanonicalTask(t *testing.T) { assert.Equal(t, "done", stored.Task.History[0].Parts[0].Text()) } +func TestGetPrefersCallContextUserOverContextValue(t *testing.T) { + encoded, err := pbconv.ToProtoTask(&a2a.Task{ID: a2a.TaskID("task-4"), ContextID: "session-4"}) + require.NoError(t, err) + service := newTaskStore(t, &taskTestServer{get: func(ctx context.Context, _ *apiv1alpha1.GetTaskRequest) (*apiv1alpha1.GetTaskResponse, error) { + values, _ := metadata.FromIncomingContext(ctx) + assert.Equal(t, []string{"call-context-user"}, values.Get("x-user-id")) + return &apiv1alpha1.GetTaskResponse{Task: encoded}, nil + }}) + + baseCtx, callCtx := a2asrv.NewCallContext(t.Context(), nil) + callCtx.User = a2asrv.NewAuthenticatedUser("call-context-user", nil) + ctx := auth.WithUserID(baseCtx, "context-value-user") + + _, err = service.Get(ctx, a2a.TaskID("task-4")) + require.NoError(t, err) +} + func TestGetMapsNotFound(t *testing.T) { service := newTaskStore(t, &taskTestServer{get: func(context.Context, *apiv1alpha1.GetTaskRequest) (*apiv1alpha1.GetTaskResponse, error) { return nil, status.Error(codes.NotFound, "missing")