From 5f2b8053bccf9dddd3c577bb31a9de84e9f35be2 Mon Sep 17 00:00:00 2001 From: Evan Rauner Date: Mon, 17 Aug 2026 13:56:26 -0500 Subject: [PATCH 1/4] fix(adk): preserve A2A user identity before task resolution Signed-off-by: Evan Rauner --- go/adk/pkg/a2a/executor.go | 2 +- go/adk/pkg/a2a/executor_test.go | 26 ++++++++++++++++++++++++++ 2 files changed, 27 insertions(+), 1 deletion(-) diff --git a/go/adk/pkg/a2a/executor.go b/go/adk/pkg/a2a/executor.go index 2c091cbc6..855baf351 100644 --- a/go/adk/pkg/a2a/executor.go +++ b/go/adk/pkg/a2a/executor.go @@ -94,7 +94,7 @@ func (u *userIDInterceptor) Before(ctx context.Context, callCtx *a2asrv.CallCont } // Set the authenticated user so downstream code picks up the real identity. callCtx.User = &a2asrv.AuthenticatedUser{UserName: vals[0]} - return ctx, nil + return auth.WithUserID(ctx, vals[0]), nil } // newAgentMessage builds an agent message stamped with the request's context diff --git a/go/adk/pkg/a2a/executor_test.go b/go/adk/pkg/a2a/executor_test.go index 91af7fe16..78fe97bd2 100644 --- a/go/adk/pkg/a2a/executor_test.go +++ b/go/adk/pkg/a2a/executor_test.go @@ -1,12 +1,38 @@ package a2a import ( + "context" + "net/http" "testing" a2atype "github.com/a2aproject/a2a-go/a2a" "github.com/a2aproject/a2a-go/a2asrv" + "github.com/kagent-dev/kagent/go/adk/pkg/auth" ) +func TestUserIDCallInterceptor(t *testing.T) { + ctx, callCtx := a2asrv.WithCallContext(context.Background(), a2asrv.NewRequestMeta(map[string][]string{ + "x-user-id": {"initiating-user"}, + })) + + gotCtx, err := UserIDCallInterceptor().Before(ctx, callCtx, &a2asrv.Request{}) + if err != nil { + t.Fatalf("Before() error = %v", err) + } + if callCtx.User == nil || callCtx.User.Name() != "initiating-user" { + t.Fatalf("CallContext.User = %#v, want name %q", callCtx.User, "initiating-user") + } + + req, err := http.NewRequestWithContext(gotCtx, http.MethodGet, "http://example.com", nil) + if err != nil { + t.Fatalf("NewRequestWithContext() error = %v", err) + } + auth.NewKAgentTokenService("test-agent").AddHeaders(req) + if got := req.Header.Get("X-User-Id"); got != "initiating-user" { + t.Fatalf("X-User-Id = %q, want %q", got, "initiating-user") + } +} + // TestNewAgentMessage_StampsContextAndTaskID verifies agent messages carry the // request's context and task ids. A2A allows omitting them (the task is the // canonical carrier), but stamping them lets consumers that flatten task.history From 919e1bef5796a355d3bf687db3d71d135a8f7e20 Mon Sep 17 00:00:00 2001 From: Evan Rauner Date: Mon, 17 Aug 2026 19:34:45 -0500 Subject: [PATCH 2/4] fix(adk): scope A2A user for Python task callbacks Signed-off-by: Evan Rauner --- .../src/kagent/core/a2a/_context.py | 26 +++++++ .../src/kagent/core/a2a/_requests.py | 6 +- .../src/kagent/core/a2a/_task_store.py | 24 ++++--- .../kagent-core/tests/test_request_context.py | 24 +++++++ .../kagent-core/tests/test_task_store.py | 70 +++++++++++++++++++ 5 files changed, 138 insertions(+), 12 deletions(-) create mode 100644 python/packages/kagent-core/tests/test_request_context.py create mode 100644 python/packages/kagent-core/tests/test_task_store.py diff --git a/python/packages/kagent-core/src/kagent/core/a2a/_context.py b/python/packages/kagent-core/src/kagent/core/a2a/_context.py index 2081c9b8f..b7d29a8a2 100644 --- a/python/packages/kagent-core/src/kagent/core/a2a/_context.py +++ b/python/packages/kagent-core/src/kagent/core/a2a/_context.py @@ -1,5 +1,9 @@ +from collections.abc import Iterator +from contextlib import contextmanager from contextvars import ContextVar +from a2a.server.context import ServerCallContext + _current_user_id: ContextVar[str | None] = ContextVar("kagent_user_id", default=None) @@ -15,3 +19,25 @@ def set_request_user_id(user_id: str | None) -> None: def get_request_user_id() -> str | None: """Return the caller's user ID for the current async context.""" return _current_user_id.get() + + +def get_call_context_user_id(context: ServerCallContext | None) -> str | None: + """Return the effective user forwarded in an A2A server call context.""" + if context is None: + return None + headers = context.state.get("headers", {}) + user_id = headers.get("x-user-id") + return user_id if isinstance(user_id, str) and user_id else None + + +@contextmanager +def scoped_request_user_id(user_id: str | None) -> Iterator[None]: + """Temporarily expose a scoped user to controller HTTP request hooks.""" + if not user_id: + yield + return + token = _current_user_id.set(user_id) + try: + yield + finally: + _current_user_id.reset(token) diff --git a/python/packages/kagent-core/src/kagent/core/a2a/_requests.py b/python/packages/kagent-core/src/kagent/core/a2a/_requests.py index 13b36ffa9..3ca4af3b3 100644 --- a/python/packages/kagent-core/src/kagent/core/a2a/_requests.py +++ b/python/packages/kagent-core/src/kagent/core/a2a/_requests.py @@ -6,7 +6,7 @@ from a2a.server.tasks import TaskStore from a2a.types import MessageSendParams, Task -from ._context import set_request_user_id +from ._context import get_call_context_user_id, set_request_user_id # --- Configure Logging --- logger = logging.getLogger(__name__) @@ -43,13 +43,13 @@ async def build( task: Task | None = None, context: ServerCallContext | None = None, ) -> RequestContext: + user_id = get_call_context_user_id(context) + set_request_user_id(user_id) if context: headers = context.state.get("headers", {}) # Extract the authenticated user ID forwarded by the parent agent - user_id = headers.get("x-user-id", None) if user_id: context.user = KAgentUser(user_id=user_id) - set_request_user_id(user_id) # Propagate x-kagent-source so downstream code (e.g. session # creation) can tag this session as agent-originated. source = headers.get("x-kagent-source", None) diff --git a/python/packages/kagent-core/src/kagent/core/a2a/_task_store.py b/python/packages/kagent-core/src/kagent/core/a2a/_task_store.py index 134fa25e9..ab86a1594 100644 --- a/python/packages/kagent-core/src/kagent/core/a2a/_task_store.py +++ b/python/packages/kagent-core/src/kagent/core/a2a/_task_store.py @@ -1,6 +1,7 @@ import asyncio import httpx +from a2a.server.context import ServerCallContext from a2a.server.tasks import TaskStore from a2a.types import Message, Task from pydantic import BaseModel @@ -8,6 +9,8 @@ from kagent.core.a2a import read_metadata_value +from ._context import get_call_context_user_id, scoped_request_user_id + class KAgentTaskResponse(BaseModel): """Wrapper for KAgent controller API responses. @@ -47,7 +50,7 @@ def _clean_partial_events(self, history: list[Message]) -> list[Message]: return [item for item in history if not self._is_partial_event(item)] @override - async def save(self, task: Task, context=None) -> None: + async def save(self, task: Task, context: ServerCallContext | None = None) -> None: """Save a task to KAgent. Skips saving if the current event is a partial streaming chunk. @@ -56,7 +59,7 @@ async def save(self, task: Task, context=None) -> None: Args: task: The task to save - context: Server call context (unused, for a2a-sdk 0.3+ compatibility) + context: Server call context supplying the effective user, when available Raises: httpx.HTTPStatusError: If the API request fails @@ -65,7 +68,8 @@ async def save(self, task: Task, context=None) -> None: history = task.history or [] task.history = self._clean_partial_events(history) - response = await self.client.post("/api/tasks", json=task.model_dump(mode="json")) + with scoped_request_user_id(get_call_context_user_id(context)): + response = await self.client.post("/api/tasks", json=task.model_dump(mode="json")) response.raise_for_status() # Signal that save completed (event-based sync) @@ -73,12 +77,12 @@ async def save(self, task: Task, context=None) -> None: self._save_events[task.id].set() @override - async def get(self, task_id: str, context=None) -> Task | None: + async def get(self, task_id: str, context: ServerCallContext | None = None) -> Task | None: """Retrieve a task from KAgent. Args: task_id: The ID of the task to retrieve - context: Server call context (unused, for a2a-sdk 0.3+ compatibility) + context: Server call context supplying the effective user, when available Returns: The task if found, None otherwise @@ -86,7 +90,8 @@ async def get(self, task_id: str, context=None) -> Task | None: Raises: httpx.HTTPStatusError: If the API request fails (except 404) """ - response = await self.client.get(f"/api/tasks/{task_id}") + with scoped_request_user_id(get_call_context_user_id(context)): + response = await self.client.get(f"/api/tasks/{task_id}") if response.status_code == 404: return None response.raise_for_status() @@ -96,17 +101,18 @@ async def get(self, task_id: str, context=None) -> Task | None: return wrapped.data @override - async def delete(self, task_id: str, context=None) -> None: + async def delete(self, task_id: str, context: ServerCallContext | None = None) -> None: """Delete a task from KAgent. Args: task_id: The ID of the task to delete - context: Server call context (unused, for a2a-sdk 0.3+ compatibility) + context: Server call context supplying the effective user, when available Raises: httpx.HTTPStatusError: If the API request fails """ - response = await self.client.delete(f"/api/tasks/{task_id}") + with scoped_request_user_id(get_call_context_user_id(context)): + response = await self.client.delete(f"/api/tasks/{task_id}") response.raise_for_status() async def wait_for_save(self, task_id: str, timeout: float = 5.0) -> None: diff --git a/python/packages/kagent-core/tests/test_request_context.py b/python/packages/kagent-core/tests/test_request_context.py new file mode 100644 index 000000000..7ce3463cf --- /dev/null +++ b/python/packages/kagent-core/tests/test_request_context.py @@ -0,0 +1,24 @@ +from unittest.mock import AsyncMock, MagicMock + +import pytest +from a2a.server.agent_execution import SimpleRequestContextBuilder +from a2a.server.context import ServerCallContext + +from kagent.core.a2a import KAgentRequestContextBuilder +from kagent.core.a2a._context import get_request_user_id, set_request_user_id + + +@pytest.mark.asyncio +async def test_headerless_request_clears_previous_user(monkeypatch): + monkeypatch.setattr(SimpleRequestContextBuilder, "build", AsyncMock(return_value=MagicMock())) + builder = KAgentRequestContextBuilder(task_store=MagicMock()) + + set_request_user_id(None) + try: + await builder.build(context=ServerCallContext(state={"headers": {"x-user-id": "user-1"}})) + assert get_request_user_id() == "user-1" + + await builder.build(context=ServerCallContext(state={"headers": {}})) + assert get_request_user_id() is None + finally: + set_request_user_id(None) diff --git a/python/packages/kagent-core/tests/test_task_store.py b/python/packages/kagent-core/tests/test_task_store.py new file mode 100644 index 000000000..4b0a2d0ab --- /dev/null +++ b/python/packages/kagent-core/tests/test_task_store.py @@ -0,0 +1,70 @@ +from unittest.mock import AsyncMock, MagicMock + +import pytest +from a2a.server.context import ServerCallContext + +from kagent.core.a2a import KAgentTaskStore +from kagent.core.a2a._context import get_request_user_id, set_request_user_id + + +@pytest.mark.asyncio +async def test_get_scopes_user_from_server_call_context(): + observed_users: list[str | None] = [] + response = MagicMock(status_code=404) + + async def get_task(_: str): + observed_users.append(get_request_user_id()) + return response + + client = MagicMock() + client.get = AsyncMock(side_effect=get_task) + store = KAgentTaskStore(client) + context = ServerCallContext(state={"headers": {"x-user-id": "user-1"}}) + set_request_user_id("ambient-user") + try: + result = await store.get("task-1", context=context) + + assert result is None + assert observed_users == ["user-1"] + assert get_request_user_id() == "ambient-user" + finally: + set_request_user_id(None) + + +@pytest.mark.asyncio +async def test_get_restores_user_when_request_fails(): + client = MagicMock() + client.get = AsyncMock(side_effect=RuntimeError("controller unavailable")) + store = KAgentTaskStore(client) + context = ServerCallContext(state={"headers": {"x-user-id": "user-1"}}) + set_request_user_id("ambient-user") + try: + with pytest.raises(RuntimeError, match="controller unavailable"): + await store.get("task-1", context=context) + + assert get_request_user_id() == "ambient-user" + finally: + set_request_user_id(None) + + +@pytest.mark.asyncio +async def test_get_without_scoped_user_preserves_ambient_user(): + observed_users: list[str | None] = [] + + async def get_task(_: str): + observed_users.append(get_request_user_id()) + return MagicMock(status_code=404) + + client = MagicMock() + client.get = AsyncMock(side_effect=get_task) + store = KAgentTaskStore(client) + context = ServerCallContext(state={"headers": {}}) + set_request_user_id("ambient-user") + try: + result = await store.get("task-1", context=context) + + assert result is None + assert observed_users == ["ambient-user"] + assert get_request_user_id() == "ambient-user" + finally: + set_request_user_id(None) From d75a67f6e106abfb32742195376d5df462b48029 Mon Sep 17 00:00:00 2001 From: Evan Rauner Date: Thu, 20 Aug 2026 11:38:53 -0500 Subject: [PATCH 3/4] fix(adk): clear absent task store user context Signed-off-by: Evan Rauner --- python/packages/kagent-core/src/kagent/core/a2a/_context.py | 3 --- python/packages/kagent-core/tests/test_task_store.py | 4 ++-- 2 files changed, 2 insertions(+), 5 deletions(-) diff --git a/python/packages/kagent-core/src/kagent/core/a2a/_context.py b/python/packages/kagent-core/src/kagent/core/a2a/_context.py index b7d29a8a2..a78d2687b 100644 --- a/python/packages/kagent-core/src/kagent/core/a2a/_context.py +++ b/python/packages/kagent-core/src/kagent/core/a2a/_context.py @@ -33,9 +33,6 @@ def get_call_context_user_id(context: ServerCallContext | None) -> str | None: @contextmanager def scoped_request_user_id(user_id: str | None) -> Iterator[None]: """Temporarily expose a scoped user to controller HTTP request hooks.""" - if not user_id: - yield - return token = _current_user_id.set(user_id) try: yield diff --git a/python/packages/kagent-core/tests/test_task_store.py b/python/packages/kagent-core/tests/test_task_store.py index 4b0a2d0ab..4dc101df8 100644 --- a/python/packages/kagent-core/tests/test_task_store.py +++ b/python/packages/kagent-core/tests/test_task_store.py @@ -48,7 +48,7 @@ async def test_get_restores_user_when_request_fails(): @pytest.mark.asyncio -async def test_get_without_scoped_user_preserves_ambient_user(): +async def test_get_without_scoped_user_clears_then_restores_ambient_user(): observed_users: list[str | None] = [] async def get_task(_: str): @@ -64,7 +64,7 @@ async def get_task(_: str): result = await store.get("task-1", context=context) assert result is None - assert observed_users == ["ambient-user"] + assert observed_users == [None] assert get_request_user_id() == "ambient-user" finally: set_request_user_id(None) From 4c7721aa7b19df63e72451ae59215193f63064ea Mon Sep 17 00:00:00 2001 From: erauner Date: Fri, 21 Aug 2026 19:03:13 -0500 Subject: [PATCH 4/4] chore(adk): remove superseded Go identity changes Signed-off-by: erauner --- go/adk/pkg/a2a/executor.go | 2 +- go/adk/pkg/a2a/executor_test.go | 26 -------------------------- 2 files changed, 1 insertion(+), 27 deletions(-) diff --git a/go/adk/pkg/a2a/executor.go b/go/adk/pkg/a2a/executor.go index 855baf351..2c091cbc6 100644 --- a/go/adk/pkg/a2a/executor.go +++ b/go/adk/pkg/a2a/executor.go @@ -94,7 +94,7 @@ func (u *userIDInterceptor) Before(ctx context.Context, callCtx *a2asrv.CallCont } // Set the authenticated user so downstream code picks up the real identity. callCtx.User = &a2asrv.AuthenticatedUser{UserName: vals[0]} - return auth.WithUserID(ctx, vals[0]), nil + return ctx, nil } // newAgentMessage builds an agent message stamped with the request's context diff --git a/go/adk/pkg/a2a/executor_test.go b/go/adk/pkg/a2a/executor_test.go index 78fe97bd2..91af7fe16 100644 --- a/go/adk/pkg/a2a/executor_test.go +++ b/go/adk/pkg/a2a/executor_test.go @@ -1,38 +1,12 @@ package a2a import ( - "context" - "net/http" "testing" a2atype "github.com/a2aproject/a2a-go/a2a" "github.com/a2aproject/a2a-go/a2asrv" - "github.com/kagent-dev/kagent/go/adk/pkg/auth" ) -func TestUserIDCallInterceptor(t *testing.T) { - ctx, callCtx := a2asrv.WithCallContext(context.Background(), a2asrv.NewRequestMeta(map[string][]string{ - "x-user-id": {"initiating-user"}, - })) - - gotCtx, err := UserIDCallInterceptor().Before(ctx, callCtx, &a2asrv.Request{}) - if err != nil { - t.Fatalf("Before() error = %v", err) - } - if callCtx.User == nil || callCtx.User.Name() != "initiating-user" { - t.Fatalf("CallContext.User = %#v, want name %q", callCtx.User, "initiating-user") - } - - req, err := http.NewRequestWithContext(gotCtx, http.MethodGet, "http://example.com", nil) - if err != nil { - t.Fatalf("NewRequestWithContext() error = %v", err) - } - auth.NewKAgentTokenService("test-agent").AddHeaders(req) - if got := req.Header.Get("X-User-Id"); got != "initiating-user" { - t.Fatalf("X-User-Id = %q, want %q", got, "initiating-user") - } -} - // TestNewAgentMessage_StampsContextAndTaskID verifies agent messages carry the // request's context and task ids. A2A allows omitting them (the task is the // canonical carrier), but stamping them lets consumers that flatten task.history