mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-08-10 13:10:57 +00:00
feat(tenancy): implement workspace isolation
This commit is contained in:
@@ -17,6 +17,7 @@ from unittest.mock import AsyncMock, Mock
|
||||
# this, running a stage test in isolation triggers a circular-import error:
|
||||
# stage.py → core.app → pipelinemgr → stage.stage_class (not yet bound).
|
||||
import langbot.pkg.pipeline.pipelinemgr # noqa: F401
|
||||
from langbot.pkg.api.http.context import ExecutionContext
|
||||
|
||||
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
|
||||
import langbot_plugin.api.entities.builtin.platform.message as platform_message
|
||||
@@ -40,6 +41,14 @@ class MockApplication:
|
||||
self.query_pool = self._create_mock_query_pool()
|
||||
self.instance_config = self._create_mock_instance_config()
|
||||
self.task_mgr = self._create_mock_task_manager()
|
||||
self.workspace_service = AsyncMock()
|
||||
self.workspace_service.get_execution_binding = AsyncMock(
|
||||
return_value=Mock(
|
||||
instance_uuid='test-instance',
|
||||
workspace_uuid='test-workspace',
|
||||
placement_generation=1,
|
||||
)
|
||||
)
|
||||
# Skill manager is optional; PreProcessor only touches it for the
|
||||
# local-agent runner. None keeps the skill-binding branch inert.
|
||||
self.skill_mgr = None
|
||||
@@ -83,6 +92,7 @@ class MockApplication:
|
||||
query_pool.cached_queries = {}
|
||||
query_pool.queries = []
|
||||
query_pool.condition = AsyncMock()
|
||||
query_pool.remove_query = AsyncMock(return_value=True)
|
||||
return query_pool
|
||||
|
||||
def _create_mock_instance_config(self):
|
||||
@@ -191,6 +201,9 @@ def sample_query(sample_message_chain, sample_message_event, mock_adapter):
|
||||
|
||||
# Use model_construct to bypass Pydantic validation for test purposes
|
||||
query = pipeline_query.Query.model_construct(
|
||||
instance_uuid='test-instance',
|
||||
workspace_uuid='test-workspace',
|
||||
placement_generation=1,
|
||||
query_id='test-query-id',
|
||||
launcher_type=provider_session.LauncherTypes.PERSON,
|
||||
launcher_id=12345,
|
||||
@@ -219,6 +232,17 @@ def sample_query(sample_message_chain, sample_message_event, mock_adapter):
|
||||
resp_message_chain=None,
|
||||
current_stage_name=None,
|
||||
)
|
||||
object.__setattr__(
|
||||
query,
|
||||
'_execution_context',
|
||||
ExecutionContext(
|
||||
instance_uuid='test-instance',
|
||||
workspace_uuid='test-workspace',
|
||||
placement_generation=1,
|
||||
bot_uuid='test-bot-uuid',
|
||||
pipeline_uuid='test-pipeline-uuid',
|
||||
),
|
||||
)
|
||||
return query
|
||||
|
||||
|
||||
|
||||
@@ -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'}
|
||||
|
||||
@@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
import yaml
|
||||
from langbot_plugin.api.entities.builtin.provider import session as provider_session
|
||||
|
||||
|
||||
def _preproc_module():
|
||||
@@ -43,7 +44,12 @@ def _prompt_preprocessing_context(default_prompt=None, prompt=None):
|
||||
|
||||
|
||||
async def _run_preprocessor(mock_app, sample_query, conversation):
|
||||
session = SimpleNamespace(launcher_type=sample_query.launcher_type, launcher_id=sample_query.launcher_id)
|
||||
session = provider_session.Session(
|
||||
launcher_type=sample_query.launcher_type,
|
||||
launcher_id=sample_query.launcher_id,
|
||||
sender_id=sample_query.sender_id,
|
||||
bot_uuid=sample_query.bot_uuid,
|
||||
)
|
||||
mock_app.sess_mgr.get_session = AsyncMock(return_value=session)
|
||||
mock_app.sess_mgr.get_conversation = AsyncMock(return_value=conversation)
|
||||
mock_app.plugin_connector.emit_event = AsyncMock(return_value=_prompt_preprocessing_context())
|
||||
|
||||
@@ -0,0 +1,65 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, MagicMock, Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from langbot.pkg.pipeline.controller import Controller
|
||||
from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError
|
||||
|
||||
|
||||
def _prepare_scheduler(mock_app):
|
||||
query_pool = MagicMock()
|
||||
query_pool.remove_query = AsyncMock(return_value=True)
|
||||
query_pool.__aenter__ = AsyncMock(return_value=query_pool)
|
||||
query_pool.__aexit__ = AsyncMock(return_value=None)
|
||||
query_pool.condition = SimpleNamespace(notify_all=Mock())
|
||||
mock_app.query_pool = query_pool
|
||||
|
||||
session = SimpleNamespace(_semaphore=SimpleNamespace(release=Mock()))
|
||||
mock_app.sess_mgr.get_session = AsyncMock(return_value=session)
|
||||
mock_app.pipeline_mgr = SimpleNamespace(get_pipeline_by_uuid=AsyncMock())
|
||||
return query_pool, session
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_controller_drops_stale_query_before_pipeline_lookup(
|
||||
mock_app,
|
||||
sample_query,
|
||||
):
|
||||
query_pool, session = _prepare_scheduler(mock_app)
|
||||
mock_app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError('stale generation')
|
||||
controller = Controller(mock_app)
|
||||
|
||||
await controller._process_query(sample_query)
|
||||
|
||||
mock_app.workspace_service.get_execution_binding.assert_awaited_once_with(
|
||||
'test-workspace',
|
||||
expected_generation=1,
|
||||
)
|
||||
mock_app.pipeline_mgr.get_pipeline_by_uuid.assert_not_awaited()
|
||||
query_pool.remove_query.assert_awaited_once_with(sample_query)
|
||||
session._semaphore.release.assert_called_once_with()
|
||||
query_pool.condition.notify_all.assert_called_once_with()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_controller_revalidates_generation_before_running_pipeline(
|
||||
mock_app,
|
||||
sample_query,
|
||||
):
|
||||
query_pool, session = _prepare_scheduler(mock_app)
|
||||
runtime_pipeline = SimpleNamespace(run=AsyncMock())
|
||||
mock_app.pipeline_mgr.get_pipeline_by_uuid.return_value = runtime_pipeline
|
||||
controller = Controller(mock_app)
|
||||
|
||||
await controller._process_query(sample_query)
|
||||
|
||||
mock_app.workspace_service.get_execution_binding.assert_awaited_once_with(
|
||||
'test-workspace',
|
||||
expected_generation=1,
|
||||
)
|
||||
runtime_pipeline.run.assert_awaited_once_with(sample_query)
|
||||
query_pool.remove_query.assert_awaited_once_with(sample_query)
|
||||
session._semaphore.release.assert_called_once_with()
|
||||
@@ -5,6 +5,9 @@ import pytest
|
||||
from langbot.pkg.api.http.service.pipeline import PipelineService
|
||||
|
||||
|
||||
WORKSPACE_UUID = 'workspace-a'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_pipeline_filters_protected_fields_without_mutating_input(mock_app):
|
||||
service = PipelineService(mock_app)
|
||||
@@ -27,7 +30,7 @@ async def test_update_pipeline_filters_protected_fields_without_mutating_input(m
|
||||
}
|
||||
original_pipeline_data = pipeline_data.copy()
|
||||
|
||||
await service.update_pipeline('pipeline-uuid', pipeline_data)
|
||||
await service.update_pipeline(WORKSPACE_UUID, 'pipeline-uuid', pipeline_data)
|
||||
|
||||
assert pipeline_data == original_pipeline_data
|
||||
|
||||
@@ -36,8 +39,9 @@ async def test_update_pipeline_filters_protected_fields_without_mutating_input(m
|
||||
assert updated_fields == {'name'}
|
||||
|
||||
mock_app.bot_service.update_bot.assert_awaited_once_with(
|
||||
WORKSPACE_UUID,
|
||||
'bot-uuid',
|
||||
{'use_pipeline_name': 'Updated pipeline'},
|
||||
)
|
||||
mock_app.pipeline_mgr.remove_pipeline.assert_awaited_once_with('pipeline-uuid')
|
||||
mock_app.pipeline_mgr.load_pipeline.assert_awaited_once_with(loaded_pipeline)
|
||||
mock_app.pipeline_mgr.remove_pipeline.assert_awaited_once_with('workspace-a', 'pipeline-uuid')
|
||||
mock_app.pipeline_mgr.load_pipeline.assert_awaited_once_with('workspace-a', loaded_pipeline)
|
||||
|
||||
@@ -6,6 +6,18 @@ import pytest
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
from importlib import import_module
|
||||
|
||||
from langbot.pkg.api.http.context import ExecutionContext
|
||||
from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError
|
||||
|
||||
|
||||
def _context(pipeline_uuid: str = 'test-uuid') -> ExecutionContext:
|
||||
return ExecutionContext(
|
||||
instance_uuid='test-instance',
|
||||
workspace_uuid='test-workspace',
|
||||
placement_generation=1,
|
||||
pipeline_uuid=pipeline_uuid,
|
||||
)
|
||||
|
||||
|
||||
def get_pipelinemgr_module():
|
||||
return import_module('langbot.pkg.pipeline.pipelinemgr')
|
||||
@@ -51,11 +63,12 @@ async def test_load_pipeline(mock_app):
|
||||
# Create test pipeline entity
|
||||
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
|
||||
pipeline_entity.uuid = 'test-uuid'
|
||||
pipeline_entity.workspace_uuid = 'test-workspace'
|
||||
pipeline_entity.stages = []
|
||||
pipeline_entity.config = {'test': 'config'}
|
||||
pipeline_entity.extensions_preferences = {'plugins': []}
|
||||
|
||||
await manager.load_pipeline(pipeline_entity)
|
||||
await manager.load_pipeline(_context(), pipeline_entity)
|
||||
|
||||
assert len(manager.pipelines) == 1
|
||||
assert manager.pipelines[0].pipeline_entity.uuid == 'test-uuid'
|
||||
@@ -75,19 +88,20 @@ async def test_get_pipeline_by_uuid(mock_app):
|
||||
# Create and add test pipeline
|
||||
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
|
||||
pipeline_entity.uuid = 'test-uuid'
|
||||
pipeline_entity.workspace_uuid = 'test-workspace'
|
||||
pipeline_entity.stages = []
|
||||
pipeline_entity.config = {}
|
||||
pipeline_entity.extensions_preferences = {'plugins': []}
|
||||
|
||||
await manager.load_pipeline(pipeline_entity)
|
||||
await manager.load_pipeline(_context(), pipeline_entity)
|
||||
|
||||
# Test retrieval
|
||||
result = await manager.get_pipeline_by_uuid('test-uuid')
|
||||
result = await manager.get_pipeline_by_uuid(_context(), 'test-uuid')
|
||||
assert result is not None
|
||||
assert result.pipeline_entity.uuid == 'test-uuid'
|
||||
|
||||
# Test non-existent UUID
|
||||
result = await manager.get_pipeline_by_uuid('non-existent')
|
||||
result = await manager.get_pipeline_by_uuid(_context('non-existent'), 'non-existent')
|
||||
assert result is None
|
||||
|
||||
|
||||
@@ -105,15 +119,16 @@ async def test_remove_pipeline(mock_app):
|
||||
# Create and add test pipeline
|
||||
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
|
||||
pipeline_entity.uuid = 'test-uuid'
|
||||
pipeline_entity.workspace_uuid = 'test-workspace'
|
||||
pipeline_entity.stages = []
|
||||
pipeline_entity.config = {}
|
||||
pipeline_entity.extensions_preferences = {'plugins': []}
|
||||
|
||||
await manager.load_pipeline(pipeline_entity)
|
||||
await manager.load_pipeline(_context(), pipeline_entity)
|
||||
assert len(manager.pipelines) == 1
|
||||
|
||||
# Remove pipeline
|
||||
await manager.remove_pipeline('test-uuid')
|
||||
await manager.remove_pipeline(_context(), 'test-uuid')
|
||||
assert len(manager.pipelines) == 0
|
||||
|
||||
|
||||
@@ -143,25 +158,104 @@ async def test_runtime_pipeline_execute(mock_app, sample_query):
|
||||
|
||||
# Create pipeline entity
|
||||
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
|
||||
pipeline_entity.uuid = 'test-pipeline-uuid'
|
||||
pipeline_entity.workspace_uuid = 'test-workspace'
|
||||
pipeline_entity.config = sample_query.pipeline_config
|
||||
pipeline_entity.extensions_preferences = {'plugins': []}
|
||||
|
||||
# Create runtime pipeline
|
||||
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [stage_container])
|
||||
runtime_pipeline = pipelinemgr.RuntimePipeline(
|
||||
mock_app,
|
||||
pipeline_entity,
|
||||
[stage_container],
|
||||
_context('test-pipeline-uuid'),
|
||||
)
|
||||
|
||||
# Mock plugin connector
|
||||
event_ctx = Mock()
|
||||
event_ctx.is_prevented_default = Mock(return_value=False)
|
||||
mock_app.plugin_connector.emit_event = AsyncMock(return_value=event_ctx)
|
||||
|
||||
# Add query to cached_queries to prevent KeyError in finally block
|
||||
mock_app.query_pool.cached_queries[sample_query.query_id] = sample_query
|
||||
|
||||
# Execute pipeline
|
||||
await runtime_pipeline.run(sample_query)
|
||||
|
||||
# Verify stage was called
|
||||
mock_stage.process.assert_called_once()
|
||||
mock_app.query_pool.remove_query.assert_awaited_once_with(sample_query)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_pipeline_rejects_stale_generation_before_side_effects(
|
||||
mock_app,
|
||||
sample_query,
|
||||
):
|
||||
pipelinemgr = get_pipelinemgr_module()
|
||||
persistence_pipeline = get_persistence_pipeline_module()
|
||||
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
|
||||
pipeline_entity.uuid = 'test-pipeline-uuid'
|
||||
pipeline_entity.workspace_uuid = 'test-workspace'
|
||||
pipeline_entity.config = sample_query.pipeline_config
|
||||
pipeline_entity.extensions_preferences = {'plugins': []}
|
||||
runtime_pipeline = pipelinemgr.RuntimePipeline(
|
||||
mock_app,
|
||||
pipeline_entity,
|
||||
[],
|
||||
_context('test-pipeline-uuid'),
|
||||
)
|
||||
mock_app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError('stale generation')
|
||||
|
||||
with pytest.raises(WorkspaceGenerationMismatchError):
|
||||
await runtime_pipeline.run(sample_query)
|
||||
|
||||
mock_app.plugin_connector.emit_event.assert_not_awaited()
|
||||
sample_query.adapter.reply_message.assert_not_awaited()
|
||||
sample_query.adapter.reply_message_chunk.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_pipeline_revalidates_after_awaited_stage(
|
||||
mock_app,
|
||||
sample_query,
|
||||
):
|
||||
pipelinemgr = get_pipelinemgr_module()
|
||||
stage = get_stage_module()
|
||||
persistence_pipeline = get_persistence_pipeline_module()
|
||||
entities = get_entities_module()
|
||||
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
|
||||
pipeline_entity.uuid = 'test-pipeline-uuid'
|
||||
pipeline_entity.workspace_uuid = 'test-workspace'
|
||||
pipeline_entity.config = sample_query.pipeline_config
|
||||
pipeline_entity.extensions_preferences = {'plugins': []}
|
||||
|
||||
result = entities.StageProcessResult(
|
||||
result_type=entities.ResultType.CONTINUE,
|
||||
new_query=sample_query,
|
||||
user_notice='must not be sent',
|
||||
console_notice='',
|
||||
debug_notice='',
|
||||
error_notice='',
|
||||
)
|
||||
|
||||
async def stage_process(*_args):
|
||||
mock_app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError(
|
||||
'generation changed during stage'
|
||||
)
|
||||
return result
|
||||
|
||||
mock_stage = Mock(spec=stage.PipelineStage)
|
||||
mock_stage.process = Mock(side_effect=stage_process)
|
||||
runtime_pipeline = pipelinemgr.RuntimePipeline(
|
||||
mock_app,
|
||||
pipeline_entity,
|
||||
[pipelinemgr.StageInstContainer(inst_name='TestStage', inst=mock_stage)],
|
||||
_context('test-pipeline-uuid'),
|
||||
)
|
||||
|
||||
with pytest.raises(WorkspaceGenerationMismatchError):
|
||||
await runtime_pipeline._execute_from_stage(0, sample_query)
|
||||
|
||||
sample_query.adapter.reply_message.assert_not_awaited()
|
||||
sample_query.adapter.reply_message_chunk.assert_not_awaited()
|
||||
|
||||
|
||||
def test_runtime_pipeline_prefers_local_agent_mcp_resources(mock_app):
|
||||
@@ -170,6 +264,8 @@ def test_runtime_pipeline_prefers_local_agent_mcp_resources(mock_app):
|
||||
persistence_pipeline = get_persistence_pipeline_module()
|
||||
|
||||
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
|
||||
pipeline_entity.uuid = 'test-uuid'
|
||||
pipeline_entity.workspace_uuid = 'test-workspace'
|
||||
pipeline_entity.config = {
|
||||
'ai': {
|
||||
'local-agent': {
|
||||
@@ -183,7 +279,7 @@ def test_runtime_pipeline_prefers_local_agent_mcp_resources(mock_app):
|
||||
'mcp_resource_agent_read_enabled': True,
|
||||
}
|
||||
|
||||
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [])
|
||||
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [], _context())
|
||||
|
||||
assert runtime_pipeline.mcp_resource_attachments == [{'server_uuid': 'srv-new', 'uri': 'file:///new.md'}]
|
||||
assert runtime_pipeline.mcp_resource_agent_read_enabled is False
|
||||
@@ -195,13 +291,15 @@ def test_runtime_pipeline_falls_back_to_extension_mcp_resources(mock_app):
|
||||
persistence_pipeline = get_persistence_pipeline_module()
|
||||
|
||||
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
|
||||
pipeline_entity.uuid = 'test-uuid'
|
||||
pipeline_entity.workspace_uuid = 'test-workspace'
|
||||
pipeline_entity.config = {'ai': {'local-agent': {}}}
|
||||
pipeline_entity.extensions_preferences = {
|
||||
'mcp_resources': [{'server_uuid': 'srv-old', 'uri': 'file:///old.md'}],
|
||||
'mcp_resource_agent_read_enabled': False,
|
||||
}
|
||||
|
||||
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [])
|
||||
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [], _context())
|
||||
|
||||
assert runtime_pipeline.mcp_resource_attachments == [{'server_uuid': 'srv-old', 'uri': 'file:///old.md'}]
|
||||
assert runtime_pipeline.mcp_resource_agent_read_enabled is False
|
||||
|
||||
@@ -6,10 +6,51 @@ Tests query management, ID generation, and async context handling.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import uuid
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from unittest.mock import Mock, patch
|
||||
|
||||
from langbot.pkg.pipeline.pool import QueryPool
|
||||
from langbot.pkg.api.http.context import ExecutionContext
|
||||
from langbot.pkg.pipeline.pool import (
|
||||
ExecutionContextMismatchError,
|
||||
ExecutionContextRequiredError,
|
||||
QueryNotFoundError,
|
||||
QueryPool,
|
||||
get_query_execution_context,
|
||||
)
|
||||
|
||||
|
||||
TEST_CONTEXT = ExecutionContext(
|
||||
instance_uuid='instance-test',
|
||||
workspace_uuid='workspace-test',
|
||||
placement_generation=1,
|
||||
)
|
||||
|
||||
|
||||
def oss_pool():
|
||||
"""Build the explicit singleton resolver used by the OSS compatibility path."""
|
||||
return QueryPool(singleton_context_resolver=lambda: TEST_CONTEXT)
|
||||
|
||||
|
||||
async def add_scoped_mock_query(pool, context, *, bot_uuid='bot-a'):
|
||||
"""Create a Query through the real pool while keeping SDK details mocked."""
|
||||
query = Mock()
|
||||
query.bot_uuid = bot_uuid
|
||||
query.pipeline_uuid = None
|
||||
query.query_id = pool.query_id_counter
|
||||
with patch('langbot.pkg.pipeline.pool.pipeline_query.Query', return_value=query):
|
||||
return await pool.add_query(
|
||||
bot_uuid=bot_uuid,
|
||||
launcher_type=Mock(),
|
||||
launcher_id='launcher-1',
|
||||
sender_id='sender-1',
|
||||
message_event=Mock(),
|
||||
message_chain=Mock(),
|
||||
adapter=Mock(),
|
||||
execution_context=context,
|
||||
)
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
@@ -39,7 +80,7 @@ class TestQueryPoolAddQuery:
|
||||
|
||||
async def test_add_query_adds_query_with_id(self):
|
||||
"""add_query creates, stores, and caches a Query with the correct ID."""
|
||||
pool = QueryPool()
|
||||
pool = oss_pool()
|
||||
|
||||
# Mock Query creation
|
||||
mock_query = Mock()
|
||||
@@ -62,12 +103,12 @@ class TestQueryPoolAddQuery:
|
||||
|
||||
# Query is added to list and cache
|
||||
assert pool.queries[0] is mock_query
|
||||
assert pool.cached_queries[0] is mock_query
|
||||
assert pool.cached_queries[('workspace-test', mock_query.query_uuid)] is mock_query
|
||||
assert mock_query.query_id == 0
|
||||
|
||||
async def test_add_query_increments_counter(self):
|
||||
"""Each add_query increments the counter."""
|
||||
pool = QueryPool()
|
||||
pool = oss_pool()
|
||||
|
||||
mock_query1 = Mock()
|
||||
mock_query1.query_id = 0
|
||||
@@ -103,7 +144,7 @@ class TestQueryPoolAddQuery:
|
||||
|
||||
async def test_add_query_appends_to_list(self):
|
||||
"""Query is appended to queries list."""
|
||||
pool = QueryPool()
|
||||
pool = oss_pool()
|
||||
|
||||
mock_query = Mock()
|
||||
mock_query.query_id = 0
|
||||
@@ -126,7 +167,7 @@ class TestQueryPoolAddQuery:
|
||||
|
||||
async def test_add_query_caches_query(self):
|
||||
"""Query is cached by query_id."""
|
||||
pool = QueryPool()
|
||||
pool = oss_pool()
|
||||
|
||||
mock_query = Mock()
|
||||
mock_query.query_id = 0
|
||||
@@ -144,12 +185,13 @@ class TestQueryPoolAddQuery:
|
||||
adapter=Mock(),
|
||||
)
|
||||
|
||||
assert 0 in pool.cached_queries
|
||||
assert pool.cached_queries[0] is mock_query
|
||||
cache_key = ('workspace-test', mock_query.query_uuid)
|
||||
assert cache_key in pool.cached_queries
|
||||
assert pool.cached_queries[cache_key] is mock_query
|
||||
|
||||
async def test_add_query_with_pipeline_uuid(self):
|
||||
"""Query can have pipeline_uuid set."""
|
||||
pool = QueryPool()
|
||||
pool = oss_pool()
|
||||
|
||||
mock_query = Mock()
|
||||
mock_query.query_id = 0
|
||||
@@ -175,7 +217,7 @@ class TestQueryPoolAddQuery:
|
||||
|
||||
async def test_add_query_sets_routed_by_rule_variable(self):
|
||||
"""Query has _routed_by_rule variable."""
|
||||
pool = QueryPool()
|
||||
pool = oss_pool()
|
||||
|
||||
mock_query = Mock()
|
||||
mock_query.query_id = 0
|
||||
@@ -201,7 +243,7 @@ class TestQueryPoolAddQuery:
|
||||
|
||||
async def test_add_query_notifier_condition(self):
|
||||
"""add_query notifies waiting consumers."""
|
||||
pool = QueryPool()
|
||||
pool = oss_pool()
|
||||
|
||||
mock_query = Mock()
|
||||
mock_query.query_id = 0
|
||||
@@ -237,7 +279,7 @@ class TestQueryPoolContext:
|
||||
|
||||
async def test_aenter_acquires_lock(self):
|
||||
"""__aenter__ acquires the pool lock."""
|
||||
pool = QueryPool()
|
||||
pool = oss_pool()
|
||||
|
||||
async with pool as p:
|
||||
# Lock is acquired
|
||||
@@ -260,7 +302,7 @@ class TestQueryPoolEdgeCases:
|
||||
|
||||
async def test_multiple_queries_cached_correctly(self):
|
||||
"""Multiple queries are cached separately."""
|
||||
pool = QueryPool()
|
||||
pool = oss_pool()
|
||||
|
||||
mock_queries = []
|
||||
for i in range(5):
|
||||
@@ -287,4 +329,107 @@ class TestQueryPoolEdgeCases:
|
||||
|
||||
# Each query is cached by its ID
|
||||
for i in range(5):
|
||||
assert pool.cached_queries[i] is mock_queries[i]
|
||||
query = mock_queries[i]
|
||||
assert pool.cached_queries[('workspace-test', query.query_uuid)] is query
|
||||
|
||||
|
||||
class TestQueryPoolWorkspaceIsolation:
|
||||
"""Regression coverage for trusted scope and scoped cache indexes."""
|
||||
|
||||
async def test_add_query_requires_execution_context_by_default(self):
|
||||
with pytest.raises(ExecutionContextRequiredError):
|
||||
await QueryPool().add_query(
|
||||
bot_uuid='bot-a',
|
||||
launcher_type=Mock(),
|
||||
launcher_id='launcher-1',
|
||||
sender_id='sender-1',
|
||||
message_event=Mock(),
|
||||
message_chain=Mock(),
|
||||
adapter=Mock(),
|
||||
)
|
||||
|
||||
async def test_serialized_scope_fields_are_not_trusted_context(self):
|
||||
forged_query = SimpleNamespace(
|
||||
instance_uuid='instance-test',
|
||||
workspace_uuid='workspace-test',
|
||||
placement_generation=1,
|
||||
bot_uuid='bot-a',
|
||||
pipeline_uuid=None,
|
||||
query_uuid='forged-query',
|
||||
)
|
||||
|
||||
with pytest.raises(ExecutionContextRequiredError):
|
||||
get_query_execution_context(forged_query)
|
||||
|
||||
async def test_query_lookup_is_workspace_scoped(self):
|
||||
pool = QueryPool()
|
||||
query = await add_scoped_mock_query(pool, TEST_CONTEXT)
|
||||
|
||||
uuid.UUID(query.query_uuid)
|
||||
assert await pool.get_query('workspace-test', query.query_uuid) is query
|
||||
assert await pool.get_query('workspace-other', query.query_uuid) is None
|
||||
assert await pool.get_query_by_legacy_id('workspace-test', 0) is query
|
||||
assert await pool.get_query_by_legacy_id('workspace-other', 0) is None
|
||||
with pytest.raises(QueryNotFoundError):
|
||||
await pool.require_query('workspace-other', query.query_uuid)
|
||||
|
||||
async def test_cache_separates_same_opaque_id_between_workspaces(self, monkeypatch):
|
||||
fixed_uuid = uuid.UUID('11111111-1111-4111-8111-111111111111')
|
||||
monkeypatch.setattr('langbot.pkg.pipeline.pool.uuid.uuid4', lambda: fixed_uuid)
|
||||
pool = QueryPool()
|
||||
context_a = TEST_CONTEXT
|
||||
context_b = ExecutionContext(
|
||||
instance_uuid='instance-test',
|
||||
workspace_uuid='workspace-other',
|
||||
placement_generation=1,
|
||||
)
|
||||
|
||||
query_a = await add_scoped_mock_query(pool, context_a)
|
||||
query_b = await add_scoped_mock_query(pool, context_b)
|
||||
|
||||
assert query_a.query_uuid == query_b.query_uuid
|
||||
assert await pool.get_query('workspace-test', query_a.query_uuid) is query_a
|
||||
assert await pool.get_query('workspace-other', query_b.query_uuid) is query_b
|
||||
|
||||
async def test_remove_query_cleans_both_scoped_indexes(self):
|
||||
pool = QueryPool()
|
||||
query = await add_scoped_mock_query(pool, TEST_CONTEXT)
|
||||
|
||||
assert await pool.remove_query(query) is True
|
||||
assert await pool.get_query('workspace-test', query.query_uuid) is None
|
||||
assert await pool.get_query_by_legacy_id('workspace-test', query.query_id) is None
|
||||
assert await pool.remove_query(query) is False
|
||||
|
||||
async def test_context_cannot_substitute_bot_identity(self):
|
||||
context = ExecutionContext(
|
||||
instance_uuid='instance-test',
|
||||
workspace_uuid='workspace-test',
|
||||
placement_generation=1,
|
||||
bot_uuid='bot-b',
|
||||
)
|
||||
|
||||
with pytest.raises(ExecutionContextMismatchError):
|
||||
await add_scoped_mock_query(QueryPool(), context, bot_uuid='bot-a')
|
||||
|
||||
async def test_query_counter_is_scoped_by_workspace_and_generation(self):
|
||||
pool = QueryPool()
|
||||
workspace_a = TEST_CONTEXT
|
||||
workspace_b = ExecutionContext(
|
||||
instance_uuid='instance-test',
|
||||
workspace_uuid='workspace-other',
|
||||
placement_generation=1,
|
||||
)
|
||||
next_generation = ExecutionContext(
|
||||
instance_uuid='instance-test',
|
||||
workspace_uuid='workspace-test',
|
||||
placement_generation=2,
|
||||
)
|
||||
|
||||
await add_scoped_mock_query(pool, workspace_a)
|
||||
await add_scoped_mock_query(pool, workspace_a)
|
||||
await add_scoped_mock_query(pool, workspace_b)
|
||||
|
||||
assert pool.get_query_count(workspace_a) == 2
|
||||
assert pool.get_query_count(workspace_b) == 1
|
||||
assert pool.get_query_count(next_generation) == 0
|
||||
assert pool.query_id_counter == 3
|
||||
|
||||
@@ -16,6 +16,8 @@ from unittest.mock import AsyncMock, Mock
|
||||
from importlib import import_module
|
||||
from types import SimpleNamespace
|
||||
|
||||
from langbot_plugin.api.entities.builtin.provider import session as provider_session
|
||||
|
||||
from tests.factories import (
|
||||
FakeApp,
|
||||
text_query,
|
||||
@@ -35,6 +37,20 @@ def get_entities_module():
|
||||
return import_module('langbot.pkg.pipeline.entities')
|
||||
|
||||
|
||||
def make_session(
|
||||
launcher_type: provider_session.LauncherTypes = provider_session.LauncherTypes.PERSON,
|
||||
launcher_id: int = 12345,
|
||||
) -> provider_session.Session:
|
||||
"""Build a scope-aware Session that matches the shared Query factory."""
|
||||
|
||||
return provider_session.Session(
|
||||
launcher_type=launcher_type,
|
||||
launcher_id=launcher_id,
|
||||
sender_id=12345,
|
||||
bot_uuid='test-bot-uuid',
|
||||
)
|
||||
|
||||
|
||||
class TestPreProcessorNormalText:
|
||||
"""Tests for normal text message preprocessing."""
|
||||
|
||||
@@ -46,9 +62,7 @@ class TestPreProcessorNormalText:
|
||||
|
||||
app = FakeApp()
|
||||
# Mock session manager to return a session
|
||||
mock_session = Mock()
|
||||
mock_session.launcher_type = Mock(value='person')
|
||||
mock_session.launcher_id = 12345
|
||||
mock_session = make_session()
|
||||
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
|
||||
|
||||
# Mock conversation
|
||||
@@ -92,9 +106,7 @@ class TestPreProcessorNormalText:
|
||||
preproc = get_preproc_module()
|
||||
|
||||
app = FakeApp()
|
||||
mock_session = Mock()
|
||||
mock_session.launcher_type = Mock(value='person')
|
||||
mock_session.launcher_id = 12345
|
||||
mock_session = make_session()
|
||||
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
|
||||
|
||||
mock_conversation = Mock()
|
||||
@@ -132,9 +144,7 @@ class TestPreProcessorEmptyMessage:
|
||||
entities = get_entities_module()
|
||||
|
||||
app = FakeApp()
|
||||
mock_session = Mock()
|
||||
mock_session.launcher_type = Mock(value='person')
|
||||
mock_session.launcher_id = 12345
|
||||
mock_session = make_session()
|
||||
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
|
||||
|
||||
mock_conversation = Mock()
|
||||
@@ -171,9 +181,7 @@ class TestPreProcessorImageSegment:
|
||||
preproc = get_preproc_module()
|
||||
|
||||
app = FakeApp()
|
||||
mock_session = Mock()
|
||||
mock_session.launcher_type = Mock(value='person')
|
||||
mock_session.launcher_id = 12345
|
||||
mock_session = make_session()
|
||||
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
|
||||
|
||||
mock_conversation = Mock()
|
||||
@@ -219,9 +227,7 @@ class TestPreProcessorImageSegment:
|
||||
preproc = get_preproc_module()
|
||||
|
||||
app = FakeApp()
|
||||
mock_session = Mock()
|
||||
mock_session.launcher_type = Mock(value='person')
|
||||
mock_session.launcher_id = 12345
|
||||
mock_session = make_session()
|
||||
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
|
||||
|
||||
mock_conversation = Mock()
|
||||
@@ -258,9 +264,7 @@ class TestPreProcessorModelSelection:
|
||||
preproc = get_preproc_module()
|
||||
|
||||
app = FakeApp()
|
||||
mock_session = Mock()
|
||||
mock_session.launcher_type = Mock(value='person')
|
||||
mock_session.launcher_id = 12345
|
||||
mock_session = make_session()
|
||||
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
|
||||
|
||||
mock_conversation = Mock()
|
||||
@@ -305,9 +309,7 @@ class TestPreProcessorModelSelection:
|
||||
preproc = get_preproc_module()
|
||||
|
||||
app = FakeApp()
|
||||
mock_session = Mock()
|
||||
mock_session.launcher_type = Mock(value='person')
|
||||
mock_session.launcher_id = 12345
|
||||
mock_session = make_session()
|
||||
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
|
||||
|
||||
mock_conversation = Mock()
|
||||
@@ -324,7 +326,7 @@ class TestPreProcessorModelSelection:
|
||||
mock_fallback = Mock()
|
||||
mock_fallback.model_entity = Mock(uuid='fallback-uuid', abilities=['func_call'])
|
||||
|
||||
async def mock_get_model(uuid):
|
||||
async def mock_get_model(_context, uuid):
|
||||
if uuid == 'primary-uuid':
|
||||
return mock_primary
|
||||
elif uuid == 'fallback-uuid':
|
||||
@@ -368,9 +370,7 @@ class TestPreProcessorVariables:
|
||||
preproc = get_preproc_module()
|
||||
|
||||
app = FakeApp()
|
||||
mock_session = Mock()
|
||||
mock_session.launcher_type = Mock(value='person')
|
||||
mock_session.launcher_id = 12345
|
||||
mock_session = make_session()
|
||||
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
|
||||
|
||||
mock_conversation = Mock()
|
||||
@@ -405,9 +405,10 @@ class TestPreProcessorVariables:
|
||||
preproc = get_preproc_module()
|
||||
|
||||
app = FakeApp()
|
||||
mock_session = Mock()
|
||||
mock_session.launcher_type = Mock(value='group')
|
||||
mock_session.launcher_id = 99999
|
||||
mock_session = make_session(
|
||||
provider_session.LauncherTypes.GROUP,
|
||||
99999,
|
||||
)
|
||||
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
|
||||
|
||||
mock_conversation = Mock()
|
||||
@@ -443,9 +444,7 @@ class TestPreProcessorToolSelection:
|
||||
preproc = get_preproc_module()
|
||||
|
||||
app = FakeApp()
|
||||
mock_session = Mock()
|
||||
mock_session.launcher_type = Mock(value='person')
|
||||
mock_session.launcher_id = 12345
|
||||
mock_session = make_session()
|
||||
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
|
||||
|
||||
mock_conversation = Mock()
|
||||
|
||||
@@ -9,6 +9,7 @@ import langbot_plugin.api.definition.abstract.platform.adapter as abstract_platf
|
||||
import langbot_plugin.api.definition.abstract.platform.event_logger as abstract_platform_logger
|
||||
|
||||
from langbot.pkg.pipeline.pool import QueryPool
|
||||
from langbot.pkg.api.http.context import ExecutionContext
|
||||
|
||||
|
||||
class DummyEventLogger(abstract_platform_logger.AbstractEventLogger):
|
||||
@@ -64,12 +65,18 @@ async def test_add_query_returns_created_query_and_preserves_side_effects(
|
||||
adapter=adapter,
|
||||
pipeline_uuid='test-pipeline-uuid',
|
||||
routed_by_rule=True,
|
||||
execution_context=ExecutionContext(
|
||||
instance_uuid='test-instance-uuid',
|
||||
workspace_uuid='test-workspace-uuid',
|
||||
placement_generation=1,
|
||||
),
|
||||
)
|
||||
|
||||
assert query is query_pool.queries[0]
|
||||
assert query_pool.cached_queries[0] is query
|
||||
assert query_pool.cached_queries[('test-workspace-uuid', query.query_uuid)] is query
|
||||
assert query_pool.query_id_counter == 1
|
||||
assert query.query_id == 0
|
||||
assert query.bot_uuid == 'test-bot-uuid'
|
||||
assert query.pipeline_uuid == 'test-pipeline-uuid'
|
||||
assert query.workspace_uuid == 'test-workspace-uuid'
|
||||
assert query.variables == {'_routed_by_rule': True}
|
||||
|
||||
Reference in New Issue
Block a user