mirror of
https://github.com/rennf93/roboco.git
synced 2026-08-03 07:23:24 +02:00
fix(messaging): thread say() by task_id into one per-channel group and enforce channel access first
This commit is contained in:
@@ -1748,6 +1748,42 @@ class MessagingService(BaseService):
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _task_group_name(task_id: UUID) -> str:
|
||||
"""Deterministic group name that threads a task's channel chatter.
|
||||
|
||||
Encoding the task id in the group name lets us locate the same
|
||||
per-(channel, task) group on every post without a new column —
|
||||
`GroupTable.name` is reused as the lookup key.
|
||||
"""
|
||||
return f"task:{task_id}"
|
||||
|
||||
async def _task_group_for_channel(
|
||||
self,
|
||||
channel: ChannelTable,
|
||||
task_id: UUID,
|
||||
) -> GroupTable:
|
||||
"""Return the per-(channel, task) group, creating it on first post.
|
||||
|
||||
A task's discussion on a given channel threads into ONE group so it
|
||||
is not scattered across the channel's standing groups. The group is
|
||||
keyed by `_task_group_name(task_id)` within the channel.
|
||||
"""
|
||||
wanted = self._task_group_name(task_id)
|
||||
groups = await self.list_groups_in_channel(cast("UUID", channel.id))
|
||||
for group in groups:
|
||||
if group.name == wanted:
|
||||
return group
|
||||
return await self.create_group(
|
||||
GroupCreateRequest(
|
||||
name=wanted,
|
||||
channel_id=cast("UUID", channel.id),
|
||||
allowed_roles=[],
|
||||
hierarchy_level=4,
|
||||
members=[],
|
||||
)
|
||||
)
|
||||
|
||||
async def post_to_channel(
|
||||
self,
|
||||
*,
|
||||
@@ -1760,23 +1796,24 @@ class MessagingService(BaseService):
|
||||
|
||||
The gateway `say` verb addresses channels by slug (`backend-cell`,
|
||||
`all-hands`, ...) and doesn't carry session/group IDs. This adapter
|
||||
resolves the channel by slug, picks the channel's default group,
|
||||
gets or creates the active session for that group, then sends a
|
||||
message via `send_message`.
|
||||
resolves the channel by slug, picks the target group, gets or
|
||||
creates the active session for that group, then sends a message via
|
||||
`send_message`. When `task_id` is supplied the message threads into
|
||||
the per-(channel, task) group so a task's discussion lives in one
|
||||
place per channel rather than scattering across standing groups;
|
||||
otherwise it falls back to the channel's default group.
|
||||
|
||||
Channel access is enforced inside `send_message` (channel writers
|
||||
list + role rules) when `agent_slug` is supplied — this adapter
|
||||
looks up the slug from `agent_id` and forwards it so the check
|
||||
always runs. If the slug lookup returns None (unknown / removed
|
||||
agent), we fail closed by raising ChannelAccessDeniedError rather
|
||||
than letting send_message silently skip validate_channel_access.
|
||||
Channel write-access is validated UP FRONT (before any group or
|
||||
session is created) so a denied write leaves no side effects. The
|
||||
agent slug is resolved from `agent_id`; if the lookup returns None
|
||||
(unknown / removed agent) we fail closed by raising
|
||||
ChannelAccessDeniedError. `send_message` re-validates access as
|
||||
defense in depth.
|
||||
"""
|
||||
from roboco.enforcement.channel_access import ChannelAccessDeniedError
|
||||
from roboco.services.repositories import get_agent_slug
|
||||
|
||||
channel = await self.get_channel_by_slug_or_raise(channel_slug)
|
||||
group = await self._default_group_for_channel(channel)
|
||||
session = await self.get_or_create_active_session(cast("UUID", group.id))
|
||||
agent_slug = await get_agent_slug(self.session, agent_id)
|
||||
if agent_slug is None:
|
||||
raise ChannelAccessDeniedError(
|
||||
@@ -1785,6 +1822,13 @@ class MessagingService(BaseService):
|
||||
action="write",
|
||||
message=(f"agent {agent_id} not found; cannot validate channel access"),
|
||||
)
|
||||
validate_channel_access(agent_slug, channel_slug, "write")
|
||||
|
||||
if task_id is not None:
|
||||
group = await self._task_group_for_channel(channel, task_id)
|
||||
else:
|
||||
group = await self._default_group_for_channel(channel)
|
||||
session = await self.get_or_create_active_session(cast("UUID", group.id))
|
||||
return await self.send_message(
|
||||
MessageCreateRequest(
|
||||
agent_id=agent_id,
|
||||
|
||||
@@ -43,6 +43,7 @@ from sqlalchemy import select
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from collections.abc import AsyncIterator
|
||||
from uuid import UUID
|
||||
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
@@ -2464,3 +2465,214 @@ async def test_get_messaging_service_factory(
|
||||
) -> None:
|
||||
svc = get_messaging_service(db_session)
|
||||
assert isinstance(svc, MessagingService)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# post_to_channel — task threading + per-role channel access
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
async def _seed_backend_dev(db_session: AsyncSession, slug: str) -> AgentTable:
|
||||
"""Create a backend DEVELOPER agent with a known slug for ACL checks."""
|
||||
agent = AgentTable(
|
||||
id=uuid4(),
|
||||
name="Dev",
|
||||
slug=slug,
|
||||
role=AgentRole.DEVELOPER,
|
||||
team=Team.BACKEND,
|
||||
status=AgentStatus.ACTIVE,
|
||||
model_config={},
|
||||
system_prompt="dev",
|
||||
capabilities=[],
|
||||
permissions={},
|
||||
metrics={},
|
||||
)
|
||||
db_session.add(agent)
|
||||
await db_session.flush()
|
||||
return agent
|
||||
|
||||
|
||||
async def _seed_task(
|
||||
db_session: AsyncSession, created_by: UUID
|
||||
) -> tuple[ProjectTable, TaskTable]:
|
||||
project = ProjectTable(
|
||||
id=uuid4(),
|
||||
name="Thread-Proj",
|
||||
slug=f"thread-proj-{uuid4().hex[:8]}",
|
||||
git_url="https://example.com/r.git",
|
||||
assigned_cell=Team.BACKEND,
|
||||
created_by=created_by,
|
||||
)
|
||||
db_session.add(project)
|
||||
await db_session.flush()
|
||||
task = TaskTable(
|
||||
id=uuid4(),
|
||||
title="t",
|
||||
description="d",
|
||||
acceptance_criteria=["ac"],
|
||||
status=TaskStatus.PENDING,
|
||||
priority=2,
|
||||
task_type=TaskType.CODE,
|
||||
nature=TaskNature.TECHNICAL,
|
||||
project_id=project.id,
|
||||
created_by=created_by,
|
||||
team=Team.BACKEND,
|
||||
)
|
||||
db_session.add(task)
|
||||
await db_session.flush()
|
||||
return project, task
|
||||
|
||||
|
||||
def _real_channel_req(slug: str) -> ChannelCreateRequest:
|
||||
"""Channel whose slug matches a real CHANNEL_ACCESS entry so the
|
||||
static ACL check in validate_channel_access runs for real (no patch)."""
|
||||
return ChannelCreateRequest(
|
||||
name=f"Channel {slug}",
|
||||
slug=slug,
|
||||
channel_type=ChannelType.CELL,
|
||||
description="desc",
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_to_channel_threads_messages_under_one_task_group(
|
||||
db_session: AsyncSession,
|
||||
) -> None:
|
||||
"""With task_id, both posts land in the same per-(channel, task) group."""
|
||||
svc = MessagingService(db_session)
|
||||
agent = await _seed_backend_dev(db_session, "be-dev-1")
|
||||
_project, task = await _seed_task(db_session, agent.id)
|
||||
await svc.create_channel(_real_channel_req("backend-cell"))
|
||||
|
||||
first = await svc.post_to_channel(
|
||||
agent_id=agent.id,
|
||||
channel_slug="backend-cell",
|
||||
content="first update",
|
||||
task_id=task.id,
|
||||
)
|
||||
second = await svc.post_to_channel(
|
||||
agent_id=agent.id,
|
||||
channel_slug="backend-cell",
|
||||
content="second update",
|
||||
task_id=task.id,
|
||||
)
|
||||
|
||||
assert first.group_id == second.group_id
|
||||
assert first.task_id == task.id
|
||||
assert second.task_id == task.id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_to_channel_task_group_is_distinct_from_default(
|
||||
db_session: AsyncSession,
|
||||
) -> None:
|
||||
"""A task post threads into a task-specific group, NOT the channel's
|
||||
default standing group used by untasked posts."""
|
||||
svc = MessagingService(db_session)
|
||||
agent = await _seed_backend_dev(db_session, "be-dev-1")
|
||||
_project, task = await _seed_task(db_session, agent.id)
|
||||
await svc.create_channel(_real_channel_req("backend-cell"))
|
||||
|
||||
untasked = await svc.post_to_channel(
|
||||
agent_id=agent.id,
|
||||
channel_slug="backend-cell",
|
||||
content="no task here",
|
||||
)
|
||||
tasked = await svc.post_to_channel(
|
||||
agent_id=agent.id,
|
||||
channel_slug="backend-cell",
|
||||
content="task scoped",
|
||||
task_id=task.id,
|
||||
)
|
||||
|
||||
assert tasked.group_id != untasked.group_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_to_channel_separate_tasks_get_separate_groups(
|
||||
db_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Two different tasks thread into two different groups on one channel."""
|
||||
svc = MessagingService(db_session)
|
||||
agent = await _seed_backend_dev(db_session, "be-dev-1")
|
||||
_project_a, task_a = await _seed_task(db_session, agent.id)
|
||||
_project_b, task_b = await _seed_task(db_session, agent.id)
|
||||
await svc.create_channel(_real_channel_req("backend-cell"))
|
||||
|
||||
msg_a = await svc.post_to_channel(
|
||||
agent_id=agent.id,
|
||||
channel_slug="backend-cell",
|
||||
content="task a",
|
||||
task_id=task_a.id,
|
||||
)
|
||||
msg_b = await svc.post_to_channel(
|
||||
agent_id=agent.id,
|
||||
channel_slug="backend-cell",
|
||||
content="task b",
|
||||
task_id=task_b.id,
|
||||
)
|
||||
|
||||
assert msg_a.group_id != msg_b.group_id
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_to_channel_rejects_agent_without_channel_access(
|
||||
db_session: AsyncSession,
|
||||
) -> None:
|
||||
"""A backend dev cannot write to #announcements (board/PM/CEO only) —
|
||||
even with a task_id, post_to_channel must reject before threading."""
|
||||
svc = MessagingService(db_session)
|
||||
agent = await _seed_backend_dev(db_session, "be-dev-1")
|
||||
_project, task = await _seed_task(db_session, agent.id)
|
||||
await svc.create_channel(_real_channel_req("announcements"))
|
||||
|
||||
with pytest.raises(ChannelAccessDeniedError):
|
||||
await svc.post_to_channel(
|
||||
agent_id=agent.id,
|
||||
channel_slug="announcements",
|
||||
content="should be blocked",
|
||||
task_id=task.id,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_to_channel_access_denied_creates_no_task_group(
|
||||
db_session: AsyncSession,
|
||||
) -> None:
|
||||
"""Rejection happens before any per-task group is created (no side effect
|
||||
leaks for a denied write)."""
|
||||
svc = MessagingService(db_session)
|
||||
agent = await _seed_backend_dev(db_session, "be-dev-1")
|
||||
_project, task = await _seed_task(db_session, agent.id)
|
||||
ch = await svc.create_channel(_real_channel_req("announcements"))
|
||||
|
||||
with pytest.raises(ChannelAccessDeniedError):
|
||||
await svc.post_to_channel(
|
||||
agent_id=agent.id,
|
||||
channel_slug="announcements",
|
||||
content="should be blocked",
|
||||
task_id=task.id,
|
||||
)
|
||||
|
||||
groups = await svc.list_groups_in_channel(ch.id)
|
||||
assert groups == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_post_to_channel_permitted_agent_succeeds(
|
||||
db_session: AsyncSession,
|
||||
) -> None:
|
||||
"""An agent on the channel's write list posts successfully (real ACL)."""
|
||||
svc = MessagingService(db_session)
|
||||
agent = await _seed_backend_dev(db_session, "be-dev-1")
|
||||
_project, task = await _seed_task(db_session, agent.id)
|
||||
await svc.create_channel(_real_channel_req("backend-cell"))
|
||||
|
||||
msg = await svc.post_to_channel(
|
||||
agent_id=agent.id,
|
||||
channel_slug="backend-cell",
|
||||
content="hello cell",
|
||||
task_id=task.id,
|
||||
)
|
||||
assert msg.content == "hello cell"
|
||||
assert msg.task_id == task.id
|
||||
|
||||
Reference in New Issue
Block a user