Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
48 changes: 48 additions & 0 deletions backend/app/dao/chat_session_dao.py
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@

from app.dao.base import TenantScopedBaseDAO
from app.models.chat_session import ChatSession
from app.models.group import Group, GroupMember
from app.models.participant import Participant


class ChatSessionDAO(TenantScopedBaseDAO[ChatSession]):
Expand Down Expand Up @@ -47,6 +49,52 @@ async def get_active_for_agent(
)
return (await session_db.execute(stmt)).scalar_one_or_none()

async def get_active_for_sandbox_agent(
self,
*,
tenant_id: uuid.UUID,
agent_id: uuid.UUID,
session_id: uuid.UUID,
db: Any = None,
) -> ChatSession | None:
"""Authorize a Session for one Agent's local sandbox execution.

Direct and external-channel group Sessions retain exact Agent ownership.
Native group Sessions are shared, so they require an active Agent
participant membership in the active tenant-owned Group instead.
"""
async with self.session(db=db, readonly=True) as session_db:
session_stmt = select(ChatSession).where(
ChatSession.tenant_id == tenant_id,
ChatSession.id == session_id,
ChatSession.deleted_at.is_(None),
)
chat_session = (await session_db.execute(session_stmt)).scalar_one_or_none()
if chat_session is None:
return None

if chat_session.group_id is None:
return chat_session if chat_session.agent_id == agent_id else None

if chat_session.session_type != "group" or chat_session.agent_id is not None:
return None

membership_stmt = (
select(GroupMember.id)
.join(Group, Group.id == GroupMember.group_id)
.join(Participant, Participant.id == GroupMember.participant_id)
.where(
Group.id == chat_session.group_id,
Group.tenant_id == tenant_id,
Group.deleted_at.is_(None),
GroupMember.removed_at.is_(None),
Participant.type == "agent",
Participant.ref_id == agent_id,
)
)
membership_id = (await session_db.execute(membership_stmt)).scalar_one_or_none()
return chat_session if membership_id is not None else None

async def get_including_deleted(self, session_id: uuid.UUID, db: Any = None) -> ChatSession | None:
"""Fetch a session by ID including soft-deleted records."""
tenant_id = self._require_tenant_id()
Expand Down
2 changes: 1 addition & 1 deletion backend/app/services/agent_tools.py
Original file line number Diff line number Diff line change
Expand Up @@ -1975,7 +1975,7 @@ async def _resolve_sandbox_execution_scope(
raise ValueError("Session sandbox execution requires a tenant")
tenant_uuid = parse_canonical_uuid(tenant_id, label="tenant_id")
session_uuid = parse_canonical_uuid(session_id, label="session_id")
chat_session = await chat_session_dao.get_active_for_agent(
chat_session = await chat_session_dao.get_active_for_sandbox_agent(
tenant_id=tenant_uuid,
agent_id=agent_id,
session_id=session_uuid,
Expand Down
197 changes: 197 additions & 0 deletions backend/tests/test_chat_session_dao.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,197 @@
"""Sandbox authorization contracts for ChatSessionDAO."""

from collections import deque
from types import SimpleNamespace
import uuid

import pytest
from sqlalchemy.dialects import postgresql

from app.dao.chat_session_dao import chat_session_dao


class _Result:
def __init__(self, values=None) -> None:
self.values = list(values or [])

def scalar_one_or_none(self):
return self.values[0] if self.values else None


class _RecordingDB:
def __init__(self, *results: _Result) -> None:
self.results = deque(results)
self.statements = []

async def execute(self, statement):
self.statements.append(statement)
if not self.results:
raise AssertionError("unexpected database query")
return self.results.popleft()


def _sql(statement) -> str:
return str(
statement.compile(
dialect=postgresql.dialect(),
compile_kwargs={"literal_binds": True},
)
)


def _session(
*,
tenant_id: uuid.UUID,
agent_id: uuid.UUID | None,
session_type: str,
group_id: uuid.UUID | None = None,
):
return SimpleNamespace(
id=uuid.uuid4(),
tenant_id=tenant_id,
agent_id=agent_id,
session_type=session_type,
group_id=group_id,
deleted_at=None,
)


@pytest.mark.asyncio
@pytest.mark.parametrize("session_type", ["direct", "group"])
async def test_sandbox_scope_preserves_exact_agent_ownership(session_type: str) -> None:
tenant_id = uuid.uuid4()
agent_id = uuid.uuid4()
chat_session = _session(
tenant_id=tenant_id,
agent_id=agent_id,
session_type=session_type,
)
db = _RecordingDB(_Result([chat_session]))

result = await chat_session_dao.get_active_for_sandbox_agent(
tenant_id=tenant_id,
agent_id=agent_id,
session_id=chat_session.id,
db=db,
)

assert result is chat_session
assert len(db.statements) == 1


@pytest.mark.asyncio
async def test_sandbox_scope_rejects_session_owned_by_another_agent() -> None:
tenant_id = uuid.uuid4()
chat_session = _session(
tenant_id=tenant_id,
agent_id=uuid.uuid4(),
session_type="direct",
)
db = _RecordingDB(_Result([chat_session]))

result = await chat_session_dao.get_active_for_sandbox_agent(
tenant_id=tenant_id,
agent_id=uuid.uuid4(),
session_id=chat_session.id,
db=db,
)

assert result is None
assert len(db.statements) == 1


@pytest.mark.asyncio
async def test_sandbox_scope_allows_active_native_group_agent_member() -> None:
tenant_id = uuid.uuid4()
agent_id = uuid.uuid4()
chat_session = _session(
tenant_id=tenant_id,
agent_id=None,
session_type="group",
group_id=uuid.uuid4(),
)
db = _RecordingDB(_Result([chat_session]), _Result([uuid.uuid4()]))

result = await chat_session_dao.get_active_for_sandbox_agent(
tenant_id=tenant_id,
agent_id=agent_id,
session_id=chat_session.id,
db=db,
)

assert result is chat_session
assert len(db.statements) == 2
membership_sql = _sql(db.statements[1])
assert "JOIN groups ON groups.id = group_members.group_id" in membership_sql
assert "JOIN participants ON participants.id = group_members.participant_id" in membership_sql
assert f"groups.tenant_id = '{tenant_id}'" in membership_sql
assert "groups.deleted_at IS NULL" in membership_sql
assert "group_members.removed_at IS NULL" in membership_sql
assert "participants.type = 'agent'" in membership_sql
assert f"participants.ref_id = '{agent_id}'" in membership_sql


@pytest.mark.asyncio
@pytest.mark.parametrize("reason", ["removed member", "deleted group", "cross-tenant group"])
async def test_sandbox_scope_rejects_inactive_native_group_membership(reason: str) -> None:
tenant_id = uuid.uuid4()
chat_session = _session(
tenant_id=tenant_id,
agent_id=None,
session_type="group",
group_id=uuid.uuid4(),
)
db = _RecordingDB(_Result([chat_session]), _Result())

result = await chat_session_dao.get_active_for_sandbox_agent(
tenant_id=tenant_id,
agent_id=uuid.uuid4(),
session_id=chat_session.id,
db=db,
)

assert result is None, reason


@pytest.mark.asyncio
@pytest.mark.parametrize("reason", ["deleted session", "cross-tenant session"])
async def test_sandbox_scope_rejects_inaccessible_session(reason: str) -> None:
tenant_id = uuid.uuid4()
session_id = uuid.uuid4()
db = _RecordingDB(_Result())

result = await chat_session_dao.get_active_for_sandbox_agent(
tenant_id=tenant_id,
agent_id=uuid.uuid4(),
session_id=session_id,
db=db,
)

assert result is None, reason
session_sql = _sql(db.statements[0])
assert f"chat_sessions.tenant_id = '{tenant_id}'" in session_sql
assert f"chat_sessions.id = '{session_id}'" in session_sql
assert "chat_sessions.deleted_at IS NULL" in session_sql


@pytest.mark.asyncio
async def test_sandbox_scope_rejects_malformed_owned_native_group_session() -> None:
tenant_id = uuid.uuid4()
agent_id = uuid.uuid4()
chat_session = _session(
tenant_id=tenant_id,
agent_id=agent_id,
session_type="group",
group_id=uuid.uuid4(),
)
db = _RecordingDB(_Result([chat_session]))

result = await chat_session_dao.get_active_for_sandbox_agent(
tenant_id=tenant_id,
agent_id=agent_id,
session_id=chat_session.id,
db=db,
)

assert result is None
assert len(db.statements) == 1
Loading