Skip to content
23 changes: 23 additions & 0 deletions python/packages/kagent-core/src/kagent/core/a2a/_context.py
Original file line number Diff line number Diff line change
@@ -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)


Expand All @@ -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)
Original file line number Diff line number Diff line change
Expand Up @@ -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__)
Expand Down Expand Up @@ -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)
Expand Down
24 changes: 15 additions & 9 deletions python/packages/kagent-core/src/kagent/core/a2a/_task_store.py
Original file line number Diff line number Diff line change
@@ -1,13 +1,16 @@
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
from typing_extensions import override

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.
Expand Down Expand Up @@ -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.
Expand All @@ -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
Expand All @@ -65,28 +68,30 @@ 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)
if task.id in self._save_events:
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

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()
Expand All @@ -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:
Expand Down
24 changes: 24 additions & 0 deletions python/packages/kagent-core/tests/test_request_context.py
Original file line number Diff line number Diff line change
@@ -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)
70 changes: 70 additions & 0 deletions python/packages/kagent-core/tests/test_task_store.py
Original file line number Diff line number Diff line change
@@ -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)
Loading