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..a78d2687b 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,22 @@ 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.""" + 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..4dc101df8 --- /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_clears_then_restores_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 == [None] + assert get_request_user_id() == "ambient-user" + finally: + set_request_user_id(None)