feat(tenancy): implement workspace isolation

This commit is contained in:
Junyan Qin
2026-07-19 09:58:59 +08:00
parent 9eb292992d
commit c6f826fe2d
271 changed files with 31162 additions and 6106 deletions
+337 -24
View File
@@ -25,6 +25,49 @@ from tests.factories import (
import langbot_plugin.api.entities.builtin.provider.session as provider_session
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.pipeline.pool import (
ExecutionContextMismatchError,
ExecutionContextRequiredError,
bind_execution_context,
)
from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError
def execution_context(
workspace_uuid='workspace-test',
*,
bot_uuid='test-bot',
pipeline_uuid=None,
placement_generation=1,
):
return ExecutionContext(
instance_uuid='instance-test',
workspace_uuid=workspace_uuid,
placement_generation=placement_generation,
bot_uuid=bot_uuid,
pipeline_uuid=pipeline_uuid,
)
def aggregation_key(
context,
*,
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
bot_uuid='test-bot',
pipeline_uuid=None,
):
return (
context.instance_uuid,
context.workspace_uuid,
context.placement_generation,
bot_uuid,
pipeline_uuid,
launcher_type.value,
launcher_id,
)
def get_aggregator_module():
"""Lazy import to avoid circular import issues."""
@@ -36,12 +79,66 @@ def make_aggregator_app():
app = FakeApp()
# Ensure query_pool has add_query method
app.query_pool.add_query = AsyncMock()
async def resolve_context(
context,
*,
bot_uuid,
pipeline_uuid,
query_uuid=None,
):
if context is None:
raise ExecutionContextRequiredError('ExecutionContext required in test')
return bind_execution_context(
context,
bot_uuid=bot_uuid,
pipeline_uuid=pipeline_uuid,
query_uuid=query_uuid,
)
app.query_pool.resolve_execution_context = AsyncMock(side_effect=resolve_context)
# Add pipeline_mgr mock
app.pipeline_mgr = AsyncMock()
app.pipeline_mgr.get_pipeline_by_uuid = AsyncMock(return_value=None)
app.workspace_service = Mock()
app.workspace_service.get_execution_binding = AsyncMock(
return_value=Mock(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
)
)
return app
def enable_aggregation(app, *, delay=10.0):
pipeline = Mock()
pipeline.pipeline_entity.config = {
'trigger': {
'message-aggregation': {
'enabled': True,
'delay': delay,
}
}
}
app.pipeline_mgr.get_pipeline_by_uuid = AsyncMock(return_value=pipeline)
def scoped_message_kwargs(context, *, launcher_id=12345, text='hello'):
chain = text_chain(text)
return {
'execution_context': context,
'bot_uuid': context.bot_uuid,
'launcher_type': provider_session.LauncherTypes.PERSON,
'launcher_id': launcher_id,
'sender_id': launcher_id,
'message_event': friend_message_event(chain),
'message_chain': chain,
'adapter': mock_adapter(),
'pipeline_uuid': context.pipeline_uuid,
}
class TestPendingMessage:
"""Tests for PendingMessage dataclass."""
@@ -54,6 +151,7 @@ class TestPendingMessage:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
execution_context=execution_context(pipeline_uuid='test-pipeline'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -77,9 +175,14 @@ class TestSessionBuffer:
"""SessionBuffer should be created with correct fields."""
aggregator = get_aggregator_module()
buffer = aggregator.SessionBuffer(session_id='test-session')
context = execution_context()
key = aggregation_key(context)
buffer = aggregator.SessionBuffer(
aggregation_key=key,
execution_context=context,
)
assert buffer.session_id == 'test-session'
assert buffer.aggregation_key == key
assert buffer.messages == []
assert buffer.timer_task is None
assert buffer.last_message_time is not None
@@ -93,6 +196,7 @@ class TestSessionBuffer:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -103,8 +207,10 @@ class TestSessionBuffer:
pipeline_uuid=None,
)
context = execution_context()
buffer = aggregator.SessionBuffer(
session_id='test-session',
aggregation_key=aggregation_key(context),
execution_context=context,
messages=[pending],
)
@@ -127,7 +233,7 @@ class TestMessageAggregatorInit:
class TestMessageAggregatorSessionId:
"""Tests for session ID generation."""
"""Tests for scoped aggregation key generation."""
def test_session_id_format(self):
"""Session ID should be correctly formatted."""
@@ -136,13 +242,24 @@ class TestMessageAggregatorSessionId:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
session_id = agg._get_session_id(
context = execution_context()
session_id = agg._get_aggregation_key(
context,
bot_uuid='bot-123',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=45678,
pipeline_uuid=None,
)
assert session_id == 'bot-123:person:45678'
assert session_id == (
'instance-test',
'workspace-test',
1,
'bot-123',
None,
'person',
45678,
)
def test_session_id_different_launchers(self):
"""Different launcher types should produce different IDs."""
@@ -151,16 +268,21 @@ class TestMessageAggregatorSessionId:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
person_id = agg._get_session_id(
context = execution_context()
person_id = agg._get_aggregation_key(
context,
bot_uuid='bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=123,
pipeline_uuid=None,
)
group_id = agg._get_session_id(
group_id = agg._get_aggregation_key(
context,
bot_uuid='bot',
launcher_type=provider_session.LauncherTypes.GROUP,
launcher_id=123,
pipeline_uuid=None,
)
assert person_id != group_id
@@ -177,7 +299,7 @@ class TestMessageAggregatorConfig:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config(None)
enabled, delay = await agg._get_aggregation_config(execution_context(), None)
assert enabled == False
assert delay == 1.5
@@ -191,7 +313,10 @@ class TestMessageAggregatorConfig:
app.pipeline_mgr.get_pipeline_by_uuid = AsyncMock(return_value=None)
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config('unknown-pipeline')
enabled, delay = await agg._get_aggregation_config(
execution_context(pipeline_uuid='unknown-pipeline'),
'unknown-pipeline',
)
assert enabled == False
assert delay == 1.5
@@ -217,7 +342,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config('test-pipeline')
enabled, delay = await agg._get_aggregation_config(
execution_context(pipeline_uuid='test-pipeline'),
'test-pipeline',
)
assert enabled == True
assert delay == 2.0
@@ -243,7 +371,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config('test-pipeline')
enabled, delay = await agg._get_aggregation_config(
execution_context(pipeline_uuid='test-pipeline'),
'test-pipeline',
)
assert delay == 1.0 # Clamped to minimum
@@ -268,7 +399,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config('test-pipeline')
enabled, delay = await agg._get_aggregation_config(
execution_context(pipeline_uuid='test-pipeline'),
'test-pipeline',
)
assert delay == 10.0 # Clamped to maximum
@@ -293,7 +427,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config('test-pipeline')
enabled, delay = await agg._get_aggregation_config(
execution_context(pipeline_uuid='test-pipeline'),
'test-pipeline',
)
assert delay == 1.5 # Default
@@ -314,6 +451,7 @@ class TestMessageAggregatorAddMessage:
adapter = mock_adapter()
await agg.add_message(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -353,6 +491,7 @@ class TestMessageAggregatorAddMessage:
adapter = mock_adapter()
await agg.add_message(
execution_context=execution_context(pipeline_uuid='test-pipeline'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -394,6 +533,7 @@ class TestMessageAggregatorAddMessage:
# Add messages up to MAX_BUFFER_MESSAGES
for i in range(aggregator.MAX_BUFFER_MESSAGES):
await agg.add_message(
execution_context=execution_context(pipeline_uuid='test-pipeline'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -405,7 +545,14 @@ class TestMessageAggregatorAddMessage:
)
# Buffer should be flushed (empty or no buffer)
session_id = agg._get_session_id('test-bot', provider_session.LauncherTypes.PERSON, 12345)
context = execution_context(pipeline_uuid='test-pipeline')
session_id = agg._get_aggregation_key(
context,
'test-bot',
provider_session.LauncherTypes.PERSON,
12345,
'test-pipeline',
)
assert session_id not in agg.buffers or len(agg.buffers[session_id].messages) == 0
@@ -424,6 +571,7 @@ class TestMessageAggregatorMerge:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -451,6 +599,7 @@ class TestMessageAggregatorMerge:
adapter = mock_adapter()
pending1 = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -462,6 +611,7 @@ class TestMessageAggregatorMerge:
)
pending2 = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -492,6 +642,7 @@ class TestMessageAggregatorMerge:
adapter = mock_adapter()
pending1 = aggregator.PendingMessage(
execution_context=execution_context(pipeline_uuid='test-pipeline-uuid'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -504,6 +655,7 @@ class TestMessageAggregatorMerge:
)
pending2 = aggregator.PendingMessage(
execution_context=execution_context(pipeline_uuid='test-pipeline-uuid'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -532,7 +684,8 @@ class TestMessageAggregatorFlush:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
await agg._flush_buffer('nonexistent-session')
context = execution_context()
await agg._flush_buffer(aggregation_key(context), context)
# Should not call query_pool
assert not app.query_pool.add_query.called
@@ -550,6 +703,7 @@ class TestMessageAggregatorFlush:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -560,17 +714,57 @@ class TestMessageAggregatorFlush:
pipeline_uuid=None,
)
context = execution_context()
key = aggregation_key(context)
buffer = aggregator.SessionBuffer(
session_id='test-session',
aggregation_key=key,
execution_context=context,
messages=[pending],
)
agg.buffers['test-session'] = buffer
agg.buffers[key] = buffer
await agg._flush_buffer('test-session')
await agg._flush_buffer(key, context)
assert app.query_pool.add_query.called
assert 'test-session' not in agg.buffers
assert key not in agg.buffers
@pytest.mark.asyncio
async def test_flush_drops_buffer_when_placement_generation_is_stale(self):
"""A debounce timer cannot enqueue work after its placement is fenced."""
aggregator = get_aggregator_module()
app = make_aggregator_app()
app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError('stale generation')
agg = aggregator.MessageAggregator(app)
context = execution_context(placement_generation=3)
pending = aggregator.PendingMessage(
execution_context=context,
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
sender_id=12345,
message_event=friend_message_event(text_chain('stale')),
message_chain=text_chain('stale'),
adapter=mock_adapter(),
pipeline_uuid=None,
)
key = aggregation_key(context)
agg.buffers[key] = aggregator.SessionBuffer(
aggregation_key=key,
execution_context=context,
messages=[pending],
)
with pytest.raises(WorkspaceGenerationMismatchError):
await agg._flush_buffer(key, context)
app.workspace_service.get_execution_binding.assert_awaited_once_with(
'workspace-test',
expected_generation=3,
)
app.query_pool.add_query.assert_not_awaited()
assert key not in agg.buffers
class TestMessageAggregatorFlushAll:
@@ -603,6 +797,7 @@ class TestMessageAggregatorFlushAll:
# Create two buffers
pending1 = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -614,6 +809,7 @@ class TestMessageAggregatorFlushAll:
)
pending2 = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=67890,
@@ -624,14 +820,131 @@ class TestMessageAggregatorFlushAll:
pipeline_uuid=None,
)
buffer1 = aggregator.SessionBuffer(session_id='session-1', messages=[pending1])
buffer2 = aggregator.SessionBuffer(session_id='session-2', messages=[pending2])
context = execution_context()
key1 = aggregation_key(context, launcher_id=12345)
key2 = aggregation_key(context, launcher_id=67890)
buffer1 = aggregator.SessionBuffer(
aggregation_key=key1,
execution_context=context,
messages=[pending1],
)
buffer2 = aggregator.SessionBuffer(
aggregation_key=key2,
execution_context=context,
messages=[pending2],
)
agg.buffers['session-1'] = buffer1
agg.buffers['session-2'] = buffer2
agg.buffers[key1] = buffer1
agg.buffers[key2] = buffer2
await agg.flush_all()
# Both buffers should be flushed
assert len(agg.buffers) == 0
assert app.query_pool.add_query.call_count == 2
class TestMessageAggregatorWorkspaceIsolation:
"""Regression coverage for fail-closed and cross-workspace behavior."""
@pytest.mark.asyncio
async def test_missing_execution_context_fails_closed(self):
app = make_aggregator_app()
agg = get_aggregator_module().MessageAggregator(app)
kwargs = scoped_message_kwargs(execution_context())
kwargs['execution_context'] = None
with pytest.raises(ExecutionContextRequiredError):
await agg.add_message(**kwargs)
app.query_pool.add_query.assert_not_awaited()
@pytest.mark.asyncio
async def test_same_launcher_in_two_workspaces_uses_separate_buffers(self):
app = make_aggregator_app()
enable_aggregation(app)
agg = get_aggregator_module().MessageAggregator(app)
await agg.add_message(**scoped_message_kwargs(execution_context('workspace-a', pipeline_uuid='test-pipeline')))
await agg.add_message(**scoped_message_kwargs(execution_context('workspace-b', pipeline_uuid='test-pipeline')))
assert len(agg.buffers) == 2
assert {key[1] for key in agg.buffers} == {'workspace-a', 'workspace-b'}
await agg.flush_all()
@pytest.mark.asyncio
async def test_same_launcher_in_two_bots_uses_separate_buffers(self):
app = make_aggregator_app()
enable_aggregation(app)
agg = get_aggregator_module().MessageAggregator(app)
await agg.add_message(
**scoped_message_kwargs(execution_context(bot_uuid='bot-a', pipeline_uuid='test-pipeline'))
)
await agg.add_message(
**scoped_message_kwargs(execution_context(bot_uuid='bot-b', pipeline_uuid='test-pipeline'))
)
assert len(agg.buffers) == 2
assert {key[3] for key in agg.buffers} == {'bot-a', 'bot-b'}
await agg.flush_all()
@pytest.mark.asyncio
async def test_timer_receives_exact_captured_execution_context(self, monkeypatch):
app = make_aggregator_app()
enable_aggregation(app)
agg = get_aggregator_module().MessageAggregator(app)
delayed_flush = AsyncMock()
monkeypatch.setattr(agg, '_delayed_flush', delayed_flush)
context = execution_context(pipeline_uuid='test-pipeline')
await agg.add_message(**scoped_message_kwargs(context))
await asyncio.sleep(0)
delayed_flush.assert_awaited_once()
assert delayed_flush.await_args.args[2] is context
await agg.flush_all()
@pytest.mark.asyncio
async def test_flush_rejects_context_from_another_workspace(self):
app = make_aggregator_app()
enable_aggregation(app)
agg = get_aggregator_module().MessageAggregator(app)
context_a = execution_context('workspace-a', pipeline_uuid='test-pipeline')
context_b = execution_context('workspace-b', pipeline_uuid='test-pipeline')
await agg.add_message(**scoped_message_kwargs(context_a))
key = next(iter(agg.buffers))
with pytest.raises(ExecutionContextMismatchError):
await agg._flush_buffer(key, context_b)
assert key in agg.buffers
await agg.flush_all()
def test_merge_rejects_messages_from_different_workspaces(self):
app = make_aggregator_app()
agg = get_aggregator_module().MessageAggregator(app)
aggregator = get_aggregator_module()
with pytest.raises(ExecutionContextMismatchError):
agg._merge_messages(
[
aggregator.PendingMessage(**scoped_message_kwargs(execution_context('workspace-a'))),
aggregator.PendingMessage(**scoped_message_kwargs(execution_context('workspace-b'))),
]
)
@pytest.mark.asyncio
async def test_flush_all_preserves_each_workspace_context(self):
app = make_aggregator_app()
enable_aggregation(app)
agg = get_aggregator_module().MessageAggregator(app)
await agg.add_message(**scoped_message_kwargs(execution_context('workspace-a', pipeline_uuid='test-pipeline')))
await agg.add_message(**scoped_message_kwargs(execution_context('workspace-b', pipeline_uuid='test-pipeline')))
await agg.flush_all()
forwarded_workspaces = {
call.kwargs['execution_context'].workspace_uuid for call in app.query_pool.add_query.await_args_list
}
assert forwarded_workspaces == {'workspace-a', 'workspace-b'}