chore(merge): sync master into dev/4.11.x

This commit is contained in:
huanghuoguoguo
2026-07-31 19:29:38 +08:00
502 changed files with 77975 additions and 12729 deletions
+43
View File
@@ -20,6 +20,18 @@ from langbot.pkg.provider.modelmgr import token
from langbot.pkg.provider.modelmgr.modelmgr import ModelManager
from langbot.pkg.entity.persistence import model as persistence_model
from langbot.pkg.discover import engine as discover_engine
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.workspace.entities import WorkspaceExecutionBinding
TEST_INSTANCE_UUID = 'test-instance'
TEST_WORKSPACE_UUID = 'test-workspace'
TEST_GENERATION = 1
TEST_EXECUTION_CONTEXT = ExecutionContext(
instance_uuid=TEST_INSTANCE_UUID,
workspace_uuid=TEST_WORKSPACE_UUID,
placement_generation=TEST_GENERATION,
)
class FakeProviderAPIRequester(requester.ProviderAPIRequester):
@@ -275,6 +287,26 @@ def mock_app_for_modelmgr():
app.llm_model_service = AsyncMock()
app.embedding_models_service = AsyncMock()
app.monitoring_service = AsyncMock()
app.workspace_service = SimpleNamespace(
get_execution_binding=AsyncMock(
return_value=WorkspaceExecutionBinding(
instance_uuid=TEST_INSTANCE_UUID,
workspace_uuid=TEST_WORKSPACE_UUID,
placement_generation=TEST_GENERATION,
write_fenced=False,
state='active',
)
),
get_local_execution_binding=AsyncMock(
return_value=WorkspaceExecutionBinding(
instance_uuid=TEST_INSTANCE_UUID,
workspace_uuid=TEST_WORKSPACE_UUID,
placement_generation=TEST_GENERATION,
write_fenced=False,
state='active',
)
),
)
return app
@@ -302,6 +334,7 @@ def fake_persistence_data():
providers = [
persistence_model.ModelProvider(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid=provider_uuid,
name='Test Provider',
requester='fake-requester',
@@ -309,6 +342,7 @@ def fake_persistence_data():
api_keys=['test-api-key-1', 'test-api-key-2'],
),
persistence_model.ModelProvider(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid=provider_uuid2,
name='Test Provider 2',
requester='another-fake-requester',
@@ -319,6 +353,7 @@ def fake_persistence_data():
llm_models = [
persistence_model.LLMModel(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid='test-llm-uuid-1',
name='TestLLM-1',
provider_uuid=provider_uuid,
@@ -326,6 +361,7 @@ def fake_persistence_data():
extra_args={'temperature': 0.7},
),
persistence_model.LLMModel(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid='test-llm-uuid-2',
name='TestLLM-2',
provider_uuid=provider_uuid,
@@ -336,6 +372,7 @@ def fake_persistence_data():
embedding_models = [
persistence_model.EmbeddingModel(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid='test-embedding-uuid-1',
name='TestEmbedding-1',
provider_uuid=provider_uuid,
@@ -345,6 +382,7 @@ def fake_persistence_data():
rerank_models = [
persistence_model.RerankModel(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid='test-rerank-uuid-1',
name='TestRerank-1',
provider_uuid=provider_uuid2,
@@ -370,6 +408,7 @@ def runtime_provider(fake_persistence_data, mock_app_for_modelmgr):
requester_inst = FakeProviderAPIRequester(mock_app_for_modelmgr, {'base_url': provider_entity.base_url})
return requester.RuntimeProvider(
execution_context=TEST_EXECUTION_CONTEXT,
provider_entity=provider_entity,
token_mgr=token_mgr,
requester=requester_inst,
@@ -381,6 +420,7 @@ def runtime_llm_model(fake_persistence_data, runtime_provider):
"""Provides a RuntimeLLMModel instance for testing."""
model_entity = fake_persistence_data['llm_models'][0]
return requester.RuntimeLLMModel(
execution_context=TEST_EXECUTION_CONTEXT,
model_entity=model_entity,
provider=runtime_provider,
)
@@ -391,6 +431,7 @@ def runtime_embedding_model(fake_persistence_data, runtime_provider):
"""Provides a RuntimeEmbeddingModel instance for testing."""
model_entity = fake_persistence_data['embedding_models'][0]
return requester.RuntimeEmbeddingModel(
execution_context=TEST_EXECUTION_CONTEXT,
model_entity=model_entity,
provider=runtime_provider,
)
@@ -404,6 +445,7 @@ def runtime_rerank_model(fake_persistence_data, mock_app_for_modelmgr):
requester_inst = AnotherFakeRequester(mock_app_for_modelmgr, {'base_url': provider_entity.base_url})
provider = requester.RuntimeProvider(
execution_context=TEST_EXECUTION_CONTEXT,
provider_entity=provider_entity,
token_mgr=token_mgr,
requester=requester_inst,
@@ -411,6 +453,7 @@ def runtime_rerank_model(fake_persistence_data, mock_app_for_modelmgr):
model_entity = fake_persistence_data['rerank_models'][0]
return requester.RuntimeRerankModel(
execution_context=TEST_EXECUTION_CONTEXT,
model_entity=model_entity,
provider=provider,
)
@@ -6,6 +6,7 @@ import pytest
from langbot.pkg.entity.persistence import model as persistence_model
from langbot.pkg.provider.modelmgr import requester
from tests.unit_tests.provider.conftest import TEST_EXECUTION_CONTEXT
from langbot_plugin.api.entities.builtin.provider import message as provider_message
from langbot_plugin.api.entities.builtin.resource import tool as resource_tool
@@ -18,10 +19,12 @@ async def test_fake_requester_counts_messages_and_tools(runtime_provider):
uuid='fake-count-model',
name='fake-count-model',
provider_uuid=runtime_provider.provider_entity.uuid,
workspace_uuid=TEST_EXECUTION_CONTEXT.workspace_uuid,
abilities=['func_call'],
extra_args={},
),
provider=runtime_provider,
execution_context=TEST_EXECUTION_CONTEXT,
)
async def _placeholder_func(**kwargs):
@@ -1148,6 +1148,46 @@ class TestInvokeRerank:
assert results[0]['relevance_score'] == 1.0
assert results[1]['relevance_score'] == 0.0
@pytest.mark.asyncio
@pytest.mark.parametrize(
('model_extra_args', 'expected_url'),
[
({'rerank_path': 'reranks'}, 'https://gateway.example.com/v1/reranks'),
({'rerank_url': 'https://rerank.example.com/api/rerank'}, 'https://rerank.example.com/api/rerank'),
],
)
async def test_invoke_rerank_openai_compatible_endpoint_override(self, model_extra_args, expected_url):
"""Endpoint configuration controls routing and is not sent in the Cohere body."""
requester = litellmchat.LiteLLMRequester(
ap=Mock(),
config={
'base_url': 'https://gateway.example.com/v1/',
'custom_llm_provider': 'openai',
},
)
model = MockRuntimeRerankModel('Qwen3-Reranker-8B', 'test-api-key')
model.model_entity.extra_args = model_extra_args
mock_resp = Mock()
mock_resp.raise_for_status = Mock()
mock_resp.json = Mock(return_value={'results': [{'index': 0, 'relevance_score': 0.8}]})
mock_client = AsyncMock()
mock_client.post = AsyncMock(return_value=mock_resp)
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock(return_value=False)
with patch('httpx.AsyncClient', return_value=mock_client):
await requester.invoke_rerank(model=model, query='query', documents=['document'])
assert mock_client.post.call_args.args[0] == expected_url
payload = mock_client.post.call_args.kwargs['json']
assert payload == {
'model': 'Qwen3-Reranker-8B',
'query': 'query',
'documents': ['document'],
'top_n': 1,
}
class TestConvertMessages:
"""Test _convert_messages method"""
@@ -6,7 +6,6 @@ triggering the circular import chain through the app module.
from __future__ import annotations
import asyncio
import importlib
import importlib.util
import os
@@ -155,7 +154,17 @@ def mcp_module():
def _make_ap():
ap = Mock()
ap.logger = Mock()
ap.instance_config = SimpleNamespace(data={'mcp': {'stdio': {'enabled': True}}})
ap.workspace_service = Mock()
ap.workspace_service.get_execution_binding = AsyncMock(
return_value=SimpleNamespace(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
)
ap.box_service = Mock()
ap.box_service.get_managed_process_websocket_connection = AsyncMock(return_value=('ws://box.example/process', {}))
return ap
@@ -167,6 +176,11 @@ def _make_session(mcp_module, server_config: dict, ap=None):
server_config=server_config,
enable=True,
ap=ap,
execution_context=mcp_module.ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
),
)
@@ -418,7 +432,7 @@ class TestBuildBoxSessionPayload:
payload = s._build_box_session_payload('session-123')
assert payload['image'] == 'node:20'
assert payload['cpus'] == 2.0
assert payload["memory_mb"] == 1024
assert payload['memory_mb'] == 1024
assert payload['pids_limit'] == 256
def test_none_fields_excluded(self, mcp_module):
@@ -592,6 +606,26 @@ class TestGetRuntimeInfoDict:
assert info['status'] == 'connecting'
assert 'box_session_id' not in info
def test_runtime_error_detail_never_echoes_secret_config(self, mcp_module):
s = _make_session(
mcp_module,
{
'name': 'test',
'uuid': 'test-uuid',
'mode': 'invalid',
'headers': {'Authorization': 'Bearer TOPSECRET'},
'env': {'API_KEY': 'TOPSECRET'},
},
)
s.status = mcp_module.MCPSessionStatus.ERROR
s.error_message = f'Unknown MCP server mode: {s.server_config}'
info = s.get_runtime_info_dict()
assert info['error_message'] == 'MCP runtime failed'
assert info['error_code'] == 'runtime_error'
assert 'TOPSECRET' not in str(info)
def test_runtime_tools_include_parameters(self, mcp_module):
s = _make_session(
mcp_module,
@@ -681,7 +715,52 @@ class TestGetRuntimeInfoDict:
# ... but are isolated by distinct process_ids within that session.
assert transient._box_stdio_runtime.process_id != live._box_stdio_runtime.process_id
def test_stdio_session_keeps_box_transport_while_runtime_reconnects(self, mcp_module):
def test_different_resource_profiles_use_different_box_sessions(self, mcp_module):
ap = _make_ap()
ap.box_service.available = True
default = _make_session(
mcp_module,
{
'name': 'default',
'uuid': 'default-uuid',
'mode': 'stdio',
'command': 'uvx',
'args': ['mcp-server-time'],
},
ap=ap,
)
constrained = _make_session(
mcp_module,
{
'name': 'constrained',
'uuid': 'constrained-uuid',
'mode': 'stdio',
'command': 'uvx',
'args': ['mcp-server-time'],
'box': {'memory_mb': 2048},
},
ap=ap,
)
assert default._build_box_session_id() == 'mcp-shared'
assert constrained._build_box_session_id().startswith('mcp-shared-')
assert constrained._build_box_session_id() != default._build_box_session_id()
writable = _make_session(
mcp_module,
{
'name': 'writable',
'uuid': 'writable-uuid',
'mode': 'stdio',
'command': 'uvx',
'args': ['mcp-server-time'],
'box': {'host_path_mode': 'rw'},
},
ap=ap,
)
assert writable._build_box_session_id() != default._build_box_session_id()
def test_stdio_session_waits_for_enabled_box_when_temporarily_unavailable(self, mcp_module):
ap = _make_ap()
ap.box_service.available = False
ap.box_service.enabled = True
@@ -715,10 +794,107 @@ class TestGetRuntimeInfoDict:
},
ap=ap,
)
info = s.get_runtime_info_dict()
assert 'box_session_id' not in info
assert 'box_enabled' not in info
@pytest.mark.asyncio
async def test_stdio_session_waits_until_box_reconnects(self, mcp_module, monkeypatch):
mcp_stdio_module = sys.modules['langbot.pkg.provider.tools.loaders.mcp_stdio']
ap = _make_ap()
ap.box_service.available = False
ap.box_service.enabled = True
session = _make_session(
mcp_module,
{
'name': 'test',
'uuid': 'test-uuid',
'mode': 'stdio',
'command': 'python',
'args': [],
'box': {'startup_timeout_sec': 1},
},
ap=ap,
)
async def reconnect_box(_delay):
ap.box_service.available = True
monkeypatch.setattr(mcp_stdio_module.asyncio, 'sleep', reconnect_box)
await session._box_stdio_runtime._wait_for_box_runtime()
assert ap.box_service.available is True
@pytest.mark.asyncio
async def test_enabled_box_timeout_does_not_exhaust_mcp_retry_budget(self, mcp_module):
ap = _make_ap()
ap.box_service.available = False
ap.box_service.enabled = True
session = _make_session(
mcp_module,
{
'name': 'test',
'uuid': 'test-uuid',
'mode': 'stdio',
'command': 'python',
'args': [],
},
ap=ap,
)
attempts = 0
async def lifecycle():
nonlocal attempts
attempts += 1
if attempts == 1:
session.error_phase = mcp_module.MCPSessionErrorPhase.BOX_UNAVAILABLE
raise RuntimeError('Box runtime is not available after 1 seconds')
session._shutdown_event.set()
session._lifecycle_loop = lifecycle
session._sleep_with_execution_fence = AsyncMock()
await session._lifecycle_loop_with_retry()
assert attempts == 2
assert session.retry_count == 1
session._sleep_with_execution_fence.assert_awaited_once_with(1)
@pytest.mark.asyncio
async def test_disabled_box_still_stops_mcp_retry_loop(self, mcp_module):
ap = _make_ap()
ap.box_service.available = False
ap.box_service.enabled = False
session = _make_session(
mcp_module,
{
'name': 'test',
'uuid': 'test-uuid',
'mode': 'stdio',
'command': 'python',
'args': [],
},
ap=ap,
)
attempts = 0
async def lifecycle():
nonlocal attempts
attempts += 1
session.error_phase = mcp_module.MCPSessionErrorPhase.BOX_UNAVAILABLE
raise RuntimeError('box_disabled_in_config')
session._lifecycle_loop = lifecycle
await session._lifecycle_loop_with_retry()
assert attempts == 1
assert session.status == mcp_module.MCPSessionStatus.ERROR
assert session.retry_count == 1
def test_stdio_session_without_box_service_uses_local_stdio(self, mcp_module):
ap = _make_ap()
del ap.box_service
@@ -774,213 +950,28 @@ class TestBoxConfigParsing:
@pytest.mark.asyncio
async def test_lifecycle_cleanup_timeout_does_not_block_box_stdio_retry(mcp_module, monkeypatch):
class HangingExitStack:
async def aclose(self):
await asyncio.Event().wait()
async def test_stdio_instance_gate_runs_before_box_transport(mcp_module):
ap = _make_ap()
ap.instance_config.data['mcp']['stdio']['enabled'] = False
ap.box_service.available = True
session = _make_session(
mcp_module,
{
'name': 'cleanup-timeout',
'uuid': 'cleanup-timeout-uuid',
'name': 'blocked',
'uuid': 'blocked-uuid',
'mode': 'stdio',
'command': 'python',
'args': [],
'env': {},
},
ap=ap,
)
session.exit_stack = HangingExitStack()
session.functions.append(Mock())
session.resources.append({'uri': 'test://stale'})
session.session = Mock()
session._cleanup_box_stdio_session = AsyncMock()
monkeypatch.setattr(mcp_module, 'MCP_TRANSPORT_CLEANUP_TIMEOUT_SECONDS', 0.01)
session._box_stdio_runtime.initialize = AsyncMock()
await asyncio.wait_for(session._cleanup_lifecycle_attempt(), timeout=1)
with pytest.raises(RuntimeError, match='disabled by instance policy'):
await session._init_stdio_python_server()
assert isinstance(session.exit_stack, mcp_module.AsyncExitStack)
assert session.functions == []
assert session.resources == []
assert session.session is None
session._cleanup_box_stdio_session.assert_awaited_once()
session.ap.logger.warning.assert_called_once_with(
'Timed out cleaning up MCP transport for cleanup-timeout; continuing lifecycle recovery'
)
@pytest.mark.asyncio
async def test_lifecycle_transport_cleanup_error_is_recovery_warning(mcp_module):
class FailingExitStack:
async def aclose(self):
raise RuntimeError('transport already closed')
session = _make_session(
mcp_module,
{
'name': 'cleanup-error',
'uuid': 'cleanup-error-uuid',
'mode': 'stdio',
'command': 'python',
'args': [],
},
)
session.exit_stack = FailingExitStack()
session._cleanup_box_stdio_session = AsyncMock()
await session._cleanup_lifecycle_attempt()
session.ap.logger.warning.assert_called_once_with(
'Error cleaning up MCP transport for cleanup-error; '
'continuing lifecycle recovery: RuntimeError: transport already closed'
)
session.ap.logger.error.assert_not_called()
@pytest.mark.asyncio
async def test_unexpected_box_stdio_transport_cancellation_becomes_retryable_error(mcp_module, monkeypatch):
session = _make_session(
mcp_module,
{
'name': 'cancelled-transport',
'uuid': 'cancelled-transport-uuid',
'mode': 'stdio',
'command': 'python',
'args': [],
},
)
session._uses_box_stdio = Mock(return_value=True)
lifecycle_calls = 0
async def cancelled_then_shutdown():
nonlocal lifecycle_calls
lifecycle_calls += 1
if lifecycle_calls == 1:
raise asyncio.CancelledError()
session._shutdown_event.set()
session._lifecycle_loop = AsyncMock(side_effect=cancelled_then_shutdown)
session._cleanup_box_stdio_session = AsyncMock()
session.functions.append(Mock())
session.resources.append({'uri': 'test://stale'})
session.resource_templates.append({'uri_template': 'test://{id}'})
session.resource_capabilities = {'subscribe': True}
session._resource_cache[('test://stale', 1, None, False)] = {'content': 'stale'}
session.session = Mock()
monkeypatch.setattr(session, '_RETRY_DELAYS', [0, 0, 0])
await session._lifecycle_loop_with_retry()
assert session._lifecycle_loop.await_count == 2
assert session.functions == []
assert session.resources == []
assert session.resource_templates == []
assert session.resource_capabilities == {}
assert session._resource_cache == {}
assert session.session is None
session.ap.logger.error.assert_called_once_with(
'Error in MCP session lifecycle cancelled-transport: '
'Box MCP transport task was cancelled unexpectedly'
)
@pytest.mark.asyncio
async def test_cancelled_lifecycle_closes_transport_in_its_own_task(mcp_module):
class TrackingExitStack:
def __init__(self):
self.close_task = None
async def aclose(self):
self.close_task = asyncio.current_task()
session = _make_session(
mcp_module,
{
'name': 'same-task-cleanup',
'uuid': 'same-task-cleanup-uuid',
'mode': 'stdio',
'command': 'python',
'args': [],
},
)
transport_stack = TrackingExitStack()
session.exit_stack = transport_stack
session._init_stdio_python_server = AsyncMock(side_effect=asyncio.CancelledError())
session._cleanup_box_stdio_session = AsyncMock()
lifecycle_task = asyncio.create_task(session._lifecycle_loop())
with pytest.raises(asyncio.CancelledError):
await lifecycle_task
assert transport_stack.close_task is lifecycle_task
session._cleanup_box_stdio_session.assert_awaited_once()
@pytest.mark.asyncio
async def test_connected_lifecycle_failures_receive_fresh_retry_budgets(mcp_module, monkeypatch):
session = _make_session(
mcp_module,
{
'name': 'repeated-runtime-recovery',
'uuid': 'repeated-runtime-recovery-uuid',
'mode': 'stdio',
'command': 'python',
'args': [],
},
)
session._uses_box_stdio = Mock(return_value=True)
session._cleanup_box_stdio_session = AsyncMock()
lifecycle_calls = 0
async def connected_then_failed_repeatedly():
nonlocal lifecycle_calls
lifecycle_calls += 1
if lifecycle_calls <= 4:
session._connection_generation += 1
raise RuntimeError(f'runtime failure {lifecycle_calls}')
session._shutdown_event.set()
session._lifecycle_loop = AsyncMock(side_effect=connected_then_failed_repeatedly)
monkeypatch.setattr(session, '_MAX_RETRIES', 1)
monkeypatch.setattr(session, '_RETRY_DELAYS', [0])
await session._lifecycle_loop_with_retry()
assert session._lifecycle_loop.await_count == 5
assert session.retry_count == 4
assert session.status != mcp_module.MCPSessionStatus.ERROR
@pytest.mark.asyncio
async def test_unexpected_box_stdio_lifecycle_return_is_retried(mcp_module, monkeypatch):
session = _make_session(
mcp_module,
{
'name': 'ended-transport',
'uuid': 'ended-transport-uuid',
'mode': 'stdio',
'command': 'python',
'args': [],
},
)
session._uses_box_stdio = Mock(return_value=True)
session._cleanup_box_stdio_session = AsyncMock()
lifecycle_calls = 0
async def ended_then_shutdown():
nonlocal lifecycle_calls
lifecycle_calls += 1
if lifecycle_calls == 2:
session._shutdown_event.set()
session._lifecycle_loop = AsyncMock(side_effect=ended_then_shutdown)
monkeypatch.setattr(session, '_RETRY_DELAYS', [0, 0, 0])
await session._lifecycle_loop_with_retry()
assert session._lifecycle_loop.await_count == 2
session.ap.logger.error.assert_called_once_with(
'Error in MCP session lifecycle ended-transport: Box MCP lifecycle ended unexpectedly'
)
session._box_stdio_runtime.initialize.assert_not_awaited()
@pytest.mark.asyncio
@@ -1000,12 +991,16 @@ async def test_init_box_stdio_server_stages_host_path_in_shared_workspace(mcp_mo
async def initialize(self):
return None
captured_transport = {}
@asynccontextmanager
async def fake_websocket_client(_url: str):
async def fake_authenticated_websocket_client(url: str, headers: dict[str, str]):
captured_transport['url'] = url
captured_transport['headers'] = headers
yield ('read-stream', 'write-stream')
mcp_stdio_module.ClientSession = FakeClientSession
mcp_stdio_module.websocket_client = fake_websocket_client
mcp_stdio_module.authenticated_websocket_client = fake_authenticated_websocket_client
ap = _make_ap()
ap.box_service.available = True
@@ -1016,7 +1011,17 @@ async def test_init_box_stdio_server_stages_host_path_in_shared_workspace(mcp_mo
execute=AsyncMock(return_value=SimpleNamespace(ok=True, stderr='', exit_code=0))
)
ap.box_service.start_managed_process = AsyncMock(return_value={})
ap.box_service.get_managed_process_websocket_url = Mock(return_value='ws://box.example/process')
ap.box_service.get_managed_process_websocket_connection = AsyncMock(
return_value=(
'ws://box.example/process',
{
'X-LangBot-Box-Control-Token': 'secret-token',
'X-LangBot-Instance-Id': 'instance-a',
'X-LangBot-Workspace-Id': 'workspace-a',
'X-LangBot-Placement-Generation': '1',
},
)
)
host_path = tmp_path / 'mcp-source'
host_path.mkdir()
@@ -1040,7 +1045,8 @@ async def test_init_box_stdio_server_stages_host_path_in_shared_workspace(mcp_mo
await session.exit_stack.aclose()
assert ap.box_service.create_session.await_count == 1
session_payload = ap.box_service.create_session.await_args.args[0]
assert ap.box_service.create_session.await_args.args[0] == session.execution_context
session_payload = ap.box_service.create_session.await_args.args[1]
assert session_payload['session_id'] == 'mcp-shared'
assert 'host_path' not in session_payload
assert ap.box_service.build_spec.call_count == 1
@@ -1050,11 +1056,22 @@ async def test_init_box_stdio_server_stages_host_path_in_shared_workspace(mcp_mo
staged_file = tmp_path / 'shared-box-workspace' / '.mcp' / 'u1' / 'workspace' / 'server.py'
assert staged_file.read_text(encoding='utf-8') == 'print("hello")\n'
process_payload = ap.box_service.start_managed_process.await_args.args[1]
assert ap.box_service.start_managed_process.await_args.args[0] == session.execution_context
process_payload = ap.box_service.start_managed_process.await_args.args[2]
assert process_payload['process_id'] == 'u1'
assert process_payload['command'] == 'python'
assert process_payload['args'] == ['/workspace/.mcp/u1/workspace/server.py']
assert process_payload['cwd'] == '/workspace/.mcp/u1/workspace'
assert captured_transport == {
'url': 'ws://box.example/process',
'headers': {
'X-LangBot-Box-Control-Token': 'secret-token',
'X-LangBot-Instance-Id': 'instance-a',
'X-LangBot-Workspace-Id': 'workspace-a',
'X-LangBot-Placement-Generation': '1',
},
}
assert 'secret-token' not in captured_transport['url']
@pytest.mark.asyncio
@@ -1093,7 +1110,7 @@ async def test_stdio_handshake_raises_coldstart_retry_while_process_alive(mcp_mo
ap.box_service.available = True
ap.box_service.create_session = AsyncMock(return_value={})
ap.box_service.start_managed_process = AsyncMock(return_value={})
ap.box_service.get_managed_process_websocket_url = Mock(return_value='ws://box/p')
ap.box_service.get_managed_process_websocket_connection = AsyncMock(return_value=('ws://box/p', {}))
session = _make_session(
mcp_module,
@@ -1159,7 +1176,7 @@ async def test_stdio_handshake_raises_fatal_when_process_exited(mcp_module, tmp_
ap.box_service.available = True
ap.box_service.create_session = AsyncMock(return_value={})
ap.box_service.start_managed_process = AsyncMock(return_value={})
ap.box_service.get_managed_process_websocket_url = Mock(return_value='ws://box/p')
ap.box_service.get_managed_process_websocket_connection = AsyncMock(return_value=('ws://box/p', {}))
session = _make_session(
mcp_module,
@@ -12,7 +12,15 @@ import pytest
from aiohttp import web
from mcp import types as mcp_types
from langbot.pkg.provider.tools.loaders.mcp import RuntimeMCPSession
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.provider.tools.loaders.mcp import MCPToolCallTimeoutError, RuntimeMCPSession
TEST_EXECUTION_CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
@pytest.fixture(autouse=True)
@@ -79,6 +87,25 @@ class _TransportProbe:
},
}
)
if method == 'tools/call':
tool_name = message.get('params', {}).get('name')
if tool_name == 'hang':
return web.Response(status=202)
return web.json_response(
{
'jsonrpc': '2.0',
'id': message['id'],
'result': {
'content': [
{
'type': 'text',
'text': 'healthy',
}
],
'isError': False,
},
}
)
return web.Response(status=202)
return web.Response(status=self.streamable_status)
@@ -140,13 +167,39 @@ async def _transport_server(streamable_status: int | None):
await runner.cleanup()
def _session(url: str, *, timeout: float = 2) -> RuntimeMCPSession:
app = cast(Any, SimpleNamespace(logger=Mock()))
def _session(
url: str,
*,
timeout: float = 2,
tool_call_timeout_sec: float = 300,
) -> RuntimeMCPSession:
app = cast(
Any,
SimpleNamespace(
logger=Mock(),
workspace_service=SimpleNamespace(
get_execution_binding=AsyncMock(
return_value=SimpleNamespace(
instance_uuid=TEST_EXECUTION_CONTEXT.instance_uuid,
workspace_uuid=TEST_EXECUTION_CONTEXT.workspace_uuid,
placement_generation=TEST_EXECUTION_CONTEXT.placement_generation,
)
)
),
),
)
return RuntimeMCPSession(
'remote-transport-test',
{'uuid': 'srv-1', 'mode': 'remote', 'url': url, 'timeout': timeout},
{
'uuid': 'srv-1',
'mode': 'remote',
'url': url,
'timeout': timeout,
'tool_call_timeout_sec': tool_call_timeout_sec,
},
True,
app,
TEST_EXECUTION_CONTEXT,
)
@@ -178,6 +231,24 @@ async def test_remote_transport_real_streamable_http_success_keeps_session_usabl
await _close_session(session)
@pytest.mark.asyncio
async def test_remote_transport_tool_timeout_does_not_poison_session():
async with _transport_server(200) as (probe, url):
session = _session(url, tool_call_timeout_sec=0.05)
try:
await session._init_remote_server()
with pytest.raises(MCPToolCallTimeoutError, match='timed out after 0.05 seconds'):
await session.invoke_mcp_tool('hang', {})
result = await session.invoke_mcp_tool('health_check', {})
assert result[0].text == 'healthy'
assert probe.streamable_messages.count('tools/call') == 2
finally:
await _close_session(session)
@pytest.mark.asyncio
@pytest.mark.parametrize('status_code', [400, 404, 405])
async def test_remote_transport_real_streamable_http_error_falls_back_to_legacy_sse(status_code: int):
+514 -120
View File
@@ -1,30 +1,54 @@
from __future__ import annotations
import asyncio
import base64
from datetime import timedelta
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import httpx
import pytest
from mcp import types as mcp_types
from mcp.shared.exceptions import McpError
from langbot.pkg.agent.runner.execution_context import project_mcp_resource_config
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.provider.tools.errors import ToolExecutionDeniedError
from langbot.pkg.provider.tools.loaders.mcp import (
MCP_RESOURCE_CONTEXT_QUERY_KEY,
MCP_RESOURCE_TRACE_QUERY_KEY,
MCP_READ_RESOURCE_SCHEMA,
MCP_TOOL_CALL_TIMEOUT_DEFAULT_SECONDS,
MCP_TOOL_LIST_RESOURCES,
MCP_TOOL_READ_RESOURCE,
MCPLoader,
MCPSessionStatus,
MCPToolCallTimeoutError,
RuntimeMCPSession,
)
from langbot.pkg.telemetry import features as telemetry_features
from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError
TEST_EXECUTION_CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
query_uuid='query-a',
)
def _app() -> SimpleNamespace:
return SimpleNamespace(logger=Mock())
return SimpleNamespace(
logger=Mock(),
workspace_service=SimpleNamespace(
get_execution_binding=AsyncMock(
return_value=SimpleNamespace(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
)
),
)
def _connected_session(
@@ -33,8 +57,15 @@ def _connected_session(
uuid: str = 'srv-1',
resources: list[dict] | None = None,
templates: list[dict] | None = None,
execution_context: ExecutionContext = TEST_EXECUTION_CONTEXT,
) -> RuntimeMCPSession:
session = RuntimeMCPSession(name, {'uuid': uuid, 'mode': 'remote'}, True, _app())
session = RuntimeMCPSession(
name,
{'uuid': uuid, 'mode': 'remote'},
True,
_app(),
execution_context,
)
session.status = MCPSessionStatus.CONNECTED
session.session = SimpleNamespace(read_resource=AsyncMock())
session.resources = resources or [
@@ -54,8 +85,24 @@ def _connected_session(
return session
def _query() -> SimpleNamespace:
return SimpleNamespace(variables={})
def _query(variables: dict | None = None, context: ExecutionContext = TEST_EXECUTION_CONTEXT) -> SimpleNamespace:
return SimpleNamespace(
instance_uuid=context.instance_uuid,
workspace_uuid=context.workspace_uuid,
placement_generation=context.placement_generation,
bot_uuid=context.bot_uuid,
pipeline_uuid=context.pipeline_uuid,
query_uuid=context.query_uuid,
variables=variables or {},
)
def _register_session(loader: MCPLoader, session: RuntimeMCPSession) -> None:
loader._register_session(
session.execution_context,
session.server_name,
session,
)
def _http_status_error(status_code: int) -> httpx.HTTPStatusError:
@@ -64,6 +111,115 @@ def _http_status_error(status_code: int) -> httpx.HTTPStatusError:
return httpx.HTTPStatusError(f'HTTP {status_code}', request=request, response=response)
def _tool_result(text: str = 'ok') -> mcp_types.CallToolResult:
return mcp_types.CallToolResult(
content=[mcp_types.TextContent(type='text', text=text)],
isError=False,
)
@pytest.mark.asyncio
async def test_invoke_mcp_tool_uses_configurable_request_timeout():
session = RuntimeMCPSession(
'slow-tools',
{
'uuid': 'srv-1',
'mode': 'remote',
'tool_call_timeout_sec': 900,
},
True,
_app(),
TEST_EXECUTION_CONTEXT,
)
session.session = SimpleNamespace(call_tool=AsyncMock(return_value=_tool_result()))
result = await session.invoke_mcp_tool('render_video', {'quality': 'high'})
assert result[0].text == 'ok'
session.session.call_tool.assert_awaited_once_with(
'render_video',
{'quality': 'high'},
read_timeout_seconds=timedelta(seconds=900),
)
@pytest.mark.asyncio
async def test_invoke_mcp_tool_zero_timeout_disables_request_deadline():
session = RuntimeMCPSession(
'unbounded-tools',
{
'uuid': 'srv-1',
'mode': 'remote',
'tool_call_timeout_sec': 0,
},
True,
_app(),
TEST_EXECUTION_CONTEXT,
)
session.session = SimpleNamespace(call_tool=AsyncMock(return_value=_tool_result()))
await session.invoke_mcp_tool('long_job', {})
session.session.call_tool.assert_awaited_once_with(
'long_job',
{},
read_timeout_seconds=None,
)
@pytest.mark.asyncio
async def test_invoke_mcp_tool_timeout_is_not_retried_and_session_remains_usable():
session = RuntimeMCPSession(
'recoverable-tools',
{
'uuid': 'srv-1',
'mode': 'remote',
'tool_call_timeout_sec': 5,
},
True,
_app(),
TEST_EXECUTION_CONTEXT,
)
timeout = McpError(
mcp_types.ErrorData(
code=httpx.codes.REQUEST_TIMEOUT,
message='Timed out while waiting for response to ClientRequest. Waited 5 seconds.',
)
)
call_tool = AsyncMock(side_effect=[timeout, _tool_result('recovered')])
session.session = SimpleNamespace(call_tool=call_tool)
with pytest.raises(
MCPToolCallTimeoutError,
match="MCP tool 'long_job' on server 'recoverable-tools' timed out after 5 seconds",
):
await session.invoke_mcp_tool('long_job', {})
assert call_tool.await_count == 1
second_result = await session.invoke_mcp_tool('health_check', {})
assert second_result[0].text == 'recovered'
assert call_tool.await_count == 2
@pytest.mark.parametrize('invalid_timeout', [-1, float('inf'), 1e300, True, 'not-a-number'])
def test_invalid_tool_call_timeout_falls_back_to_default(invalid_timeout):
ap = _app()
session = RuntimeMCPSession(
'invalid-timeout',
{
'uuid': 'srv-1',
'mode': 'remote',
'tool_call_timeout_sec': invalid_timeout,
},
True,
ap,
TEST_EXECUTION_CONTEXT,
)
assert session.tool_call_timeout_sec == MCP_TOOL_CALL_TIMEOUT_DEFAULT_SECONDS
ap.logger.warning.assert_called_once()
@pytest.mark.asyncio
async def test_remote_transport_falls_back_to_sse_for_compatible_http_status_in_exception_group():
session = RuntimeMCPSession(
@@ -71,6 +227,7 @@ async def test_remote_transport_falls_back_to_sse_for_compatible_http_status_in_
{'uuid': 'srv-1', 'mode': 'remote', 'url': 'https://example.com/mcp'},
True,
_app(),
TEST_EXECUTION_CONTEXT,
)
session._init_streamable_http_server = AsyncMock(
side_effect=ExceptionGroup('transport failed', [_http_status_error(405)])
@@ -90,6 +247,7 @@ async def test_remote_transport_does_not_fallback_for_auth_http_status():
{'uuid': 'srv-1', 'mode': 'remote', 'url': 'https://example.com/mcp'},
True,
_app(),
TEST_EXECUTION_CONTEXT,
)
error = _http_status_error(403)
session._init_streamable_http_server = AsyncMock(side_effect=error)
@@ -247,10 +405,18 @@ def test_resource_uri_allowed_supports_listed_templates_conservatively():
async def test_mcp_loader_can_hide_synthetic_resource_tools():
loader = MCPLoader(_app())
session = _connected_session()
loader.sessions = {'docs': session}
_register_session(loader, session)
with_resource_tools = await loader.get_tools(['srv-1'], include_resource_tools=True)
without_resource_tools = await loader.get_tools(['srv-1'], include_resource_tools=False)
with_resource_tools = await loader.get_tools(
TEST_EXECUTION_CONTEXT,
['srv-1'],
include_resource_tools=True,
)
without_resource_tools = await loader.get_tools(
TEST_EXECUTION_CONTEXT,
['srv-1'],
include_resource_tools=False,
)
assert {tool.name for tool in with_resource_tools} == {
MCP_TOOL_LIST_RESOURCES,
@@ -260,42 +426,23 @@ async def test_mcp_loader_can_hide_synthetic_resource_tools():
@pytest.mark.asyncio
async def test_mcp_loader_get_tool_returns_synthetic_resource_schema():
loader = MCPLoader(_app())
loader.sessions = {'docs': _connected_session()}
tool = await loader.get_tool(MCP_TOOL_READ_RESOURCE)
assert tool is not None
assert tool.name == MCP_TOOL_READ_RESOURCE
assert tool.parameters == MCP_READ_RESOURCE_SCHEMA
@pytest.mark.asyncio
@pytest.mark.parametrize('read_enabled', [False, 0, None, 'false'])
@pytest.mark.parametrize(
('tool_name', 'parameters'),
[
(MCP_TOOL_LIST_RESOURCES, {'server_name': 'docs'}),
(MCP_TOOL_READ_RESOURCE, {'server_name': 'docs', 'uri': 'file:///README.md'}),
],
)
async def test_mcp_loader_refuses_resource_tool_calls_when_agent_read_disabled(
read_enabled,
tool_name,
parameters,
):
async def test_mcp_loader_refuses_resource_tool_calls_when_agent_read_disabled():
loader = MCPLoader(_app())
session = _connected_session()
loader.sessions = {'docs': session}
query = SimpleNamespace(variables={'_pipeline_bound_mcp_servers': ['srv-1']})
project_mcp_resource_config(
query,
{'mcp-resource-agent-read-enabled': read_enabled},
_register_session(loader, session)
query = _query(
{
'_pipeline_bound_mcp_servers': ['srv-1'],
'_pipeline_mcp_resource_agent_read_enabled': False,
}
)
with pytest.raises(ToolExecutionDeniedError, match='MCP resource agent reads are disabled'):
await loader.invoke_tool(tool_name, parameters, query)
await loader.invoke_tool(
MCP_TOOL_READ_RESOURCE,
{'server_name': 'docs', 'uri': 'file:///README.md'},
query,
)
session.session.read_resource.assert_not_called()
@@ -323,9 +470,10 @@ async def test_build_resource_context_for_query_uses_only_bound_attached_text_re
)
]
)
loader.sessions = {'docs': docs, 'other': other}
query = SimpleNamespace(
variables={
_register_session(loader, docs)
_register_session(loader, other)
query = _query(
{
'_pipeline_bound_mcp_servers': ['srv-1'],
'_pipeline_mcp_resource_attachments': [
{'server_uuid': 'srv-1', 'server_name': 'docs', 'uri': 'file:///README.md', 'mode': 'pinned'},
@@ -345,91 +493,337 @@ async def test_build_resource_context_for_query_uses_only_bound_attached_text_re
other.session.read_resource.assert_not_called()
@pytest.mark.asyncio
@pytest.mark.parametrize(
('enabled_config', 'expected_enabled'),
[
pytest.param({}, True, id='missing'),
pytest.param({'enabled': True}, True, id='true'),
pytest.param({'enabled': False}, False, id='false'),
pytest.param({'enabled': 0}, False, id='zero'),
pytest.param({'enabled': None}, False, id='none'),
pytest.param({'enabled': 'false'}, False, id='string-false'),
],
)
async def test_build_resource_context_attachment_enabled_fails_closed(enabled_config, expected_enabled):
def test_mcp_loader_session_keys_do_not_collide_between_workspaces():
loader = MCPLoader(_app())
session = _connected_session(name='docs', uuid='srv-1')
session.session.read_resource.return_value = mcp_types.ReadResourceResult(
contents=[
mcp_types.TextResourceContents(
uri='file:///README.md',
mimeType='text/markdown',
text='enabled attachment',
)
]
workspace_b = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-b',
placement_generation=1,
query_uuid='query-b',
)
loader.sessions = {'docs': session}
attachment = {
'server_uuid': 'srv-1',
'server_name': 'docs',
'uri': 'file:///README.md',
'mode': 'pinned',
**enabled_config,
}
query = SimpleNamespace(
variables={
'_pipeline_bound_mcp_servers': ['srv-1'],
'_pipeline_mcp_resource_attachments': [attachment],
}
session_a = _connected_session(name='docs', uuid='srv-a')
session_b = _connected_session(
name='docs',
uuid='srv-b',
execution_context=workspace_b,
)
_register_session(loader, session_a)
_register_session(loader, session_b)
context = await loader.build_resource_context_for_query(query)
if expected_enabled:
assert 'enabled attachment' in context
session.session.read_resource.assert_awaited_once()
else:
assert context == ''
session.session.read_resource.assert_not_awaited()
assert len(loader.sessions) == 2
assert loader.get_session(TEST_EXECUTION_CONTEXT, 'docs') is session_a
assert loader.get_session(workspace_b, 'docs') is session_b
assert loader.get_session(TEST_EXECUTION_CONTEXT, 'docs') is not loader.get_session(workspace_b, 'docs')
@pytest.mark.asyncio
@pytest.mark.parametrize(
('field_name', 'invalid_value'),
[
pytest.param('max_tokens', '100', id='string-tokens'),
pytest.param('max_tokens', 0, id='zero-tokens'),
pytest.param('max_tokens', -1, id='negative-tokens'),
pytest.param('max_tokens', True, id='boolean-tokens'),
pytest.param('max_bytes', '4096', id='string-bytes'),
pytest.param('max_bytes', 0, id='zero-bytes'),
pytest.param('max_bytes', -1, id='negative-bytes'),
pytest.param('max_bytes', True, id='boolean-bytes'),
],
)
async def test_build_resource_context_invalid_attachment_limits_fail_closed(field_name, invalid_value):
ap = _app()
loader = MCPLoader(ap)
session = _connected_session(name='docs', uuid='srv-1')
loader.sessions = {'docs': session}
query = SimpleNamespace(
variables={
'_pipeline_bound_mcp_servers': ['srv-1'],
'_pipeline_mcp_resource_attachments': [
{
'server_uuid': 'srv-1',
'server_name': 'docs',
'uri': 'file:///README.md',
'mode': 'pinned',
field_name: invalid_value,
}
],
}
async def test_mcp_tool_result_is_discarded_when_generation_changes_during_call():
session = _connected_session()
binding = SimpleNamespace(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
session.ap.workspace_service.get_execution_binding.side_effect = [
binding,
binding,
WorkspaceGenerationMismatchError('generation changed during tool call'),
]
session.session = SimpleNamespace(call_tool=AsyncMock(return_value=SimpleNamespace(isError=False, content=[])))
with pytest.raises(WorkspaceGenerationMismatchError):
await session.invoke_mcp_tool('side_effecting_tool', {})
session.session.call_tool.assert_awaited_once_with(
'side_effecting_tool',
{},
read_timeout_seconds=timedelta(seconds=MCP_TOOL_CALL_TIMEOUT_DEFAULT_SECONDS),
)
context = await loader.build_resource_context_for_query(query)
assert context == ''
@pytest.mark.asyncio
async def test_mcp_resource_cache_is_not_served_to_stale_generation():
session = _connected_session()
session._resource_cache[('file:///README.md', 10, None, False)] = {
'cached_at': 0,
'envelope': {'contents': [{'type': 'text', 'text': 'stale'}]},
}
session.ap.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError(
'stale generation'
)
with pytest.raises(WorkspaceGenerationMismatchError):
await session.read_resource_envelope('file:///README.md', max_bytes=10)
session.session.read_resource.assert_not_awaited()
ap.logger.warning.assert_called_once()
@pytest.mark.asyncio
async def test_directory_projection_retires_idle_mcp_scope_without_db_poll():
loader = MCPLoader(_app())
sessions = []
for index in range(100):
context = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid=f'workspace-{index}',
placement_generation=1,
)
session = RuntimeMCPSession(
f'server-{index}',
{'uuid': f'srv-{index}', 'mode': 'remote'},
True,
loader.ap,
context,
)
session.shutdown = AsyncMock()
loader._register_session(context, session.server_name, session)
sessions.append(session)
loader.reconcile_execution_projection('instance-a', {})
reconcile_task = loader._projection_reconcile_task
assert reconcile_task is not None
assert len(loader._pending_projection_retirements) == 100
# A second projection coalesces into the same worker instead of creating
# one timer or task per Workspace.
loader.reconcile_execution_projection('instance-a', {})
assert loader._projection_reconcile_task is reconcile_task
await asyncio.wait_for(reconcile_task, timeout=1)
assert loader.sessions == {}
assert loader._scope_generations == {}
assert loader._pending_projection_retirements == set()
assert sum(session.shutdown.await_count for session in sessions) == 100
loader.ap.workspace_service.get_execution_binding.assert_not_awaited()
@pytest.mark.asyncio
async def test_directory_projection_keeps_matching_and_unaffected_mcp_scopes():
loader = MCPLoader(_app())
matching = _connected_session()
other_context = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-b',
placement_generation=1,
)
unaffected = _connected_session(
name='other',
uuid='srv-2',
execution_context=other_context,
)
_register_session(loader, matching)
_register_session(loader, unaffected)
loader.reconcile_execution_projection(
'instance-a',
{'workspace-a': 1},
affected_workspace_uuids={'workspace-a'},
)
await asyncio.sleep(0)
assert loader.get_session(TEST_EXECUTION_CONTEXT, 'docs') is matching
assert loader.get_session(other_context, 'other') is unaffected
assert loader._projection_reconcile_task is None
@pytest.mark.asyncio
async def test_mcp_loader_shutdown_cancels_startup_tasks_and_closes_sessions_concurrently():
loader = MCPLoader(_app())
hosted_cancelled = asyncio.Event()
async def pending_host():
try:
await asyncio.Event().wait()
finally:
hosted_cancelled.set()
hosted_task = asyncio.create_task(pending_host())
await asyncio.sleep(0)
loader._hosted_mcp_tasks = [hosted_task]
started: set[str] = set()
all_started = asyncio.Event()
class Session:
def __init__(self, name: str):
self.name = name
self.server_name = name
async def shutdown(self):
started.add(self.name)
if len(started) == 2:
all_started.set()
await all_started.wait()
loader.sessions = {'one': Session('one'), 'two': Session('two')}
await asyncio.wait_for(loader.shutdown(), timeout=1)
assert hosted_cancelled.is_set()
assert hosted_task.cancelled()
assert started == {'one', 'two'}
assert loader._hosted_mcp_tasks == []
assert loader.sessions == {}
@pytest.mark.asyncio
async def test_completed_mcp_host_tasks_do_not_accumulate():
loader = MCPLoader(_app())
task = asyncio.create_task(asyncio.sleep(0))
loader.track_hosted_task(task, TEST_EXECUTION_CONTEXT)
await task
await asyncio.sleep(0)
assert loader._hosted_mcp_tasks == []
assert loader._hosted_mcp_tasks_by_scope == {}
assert loader._scope_generations == {}
@pytest.mark.asyncio
async def test_generation_advance_cancels_host_tasks_and_closes_old_sessions():
loader = MCPLoader(_app())
old_session = SimpleNamespace(
server_name='old',
shutdown=AsyncMock(),
)
loader._register_session(
TEST_EXECUTION_CONTEXT,
old_session.server_name,
old_session,
)
async def pending_host():
await asyncio.Event().wait()
hosted_task = asyncio.create_task(pending_host())
loader.track_hosted_task(hosted_task, TEST_EXECUTION_CONTEXT)
await asyncio.sleep(0)
next_context = ExecutionContext(
instance_uuid=TEST_EXECUTION_CONTEXT.instance_uuid,
workspace_uuid=TEST_EXECUTION_CONTEXT.workspace_uuid,
placement_generation=2,
)
loader.ap.workspace_service.get_execution_binding = AsyncMock(
return_value=SimpleNamespace(
instance_uuid=next_context.instance_uuid,
workspace_uuid=next_context.workspace_uuid,
placement_generation=next_context.placement_generation,
)
)
await loader._assert_execution_active(next_context)
assert hosted_task.cancelled()
old_session.shutdown.assert_awaited_once_with()
assert loader.sessions == {}
assert loader._session_keys_by_scope == {}
assert loader._hosted_mcp_tasks_by_scope == {}
assert loader._scope_generations == {}
def test_session_lookup_uses_scope_index_without_global_iteration():
class NoGlobalIterationDict(dict):
def __iter__(self):
raise AssertionError('MCP lookup scanned every tenant session')
def items(self):
raise AssertionError('MCP lookup scanned every tenant session')
def values(self):
raise AssertionError('MCP lookup scanned every tenant session')
loader = MCPLoader(_app())
target_context = None
target_session = None
for index in range(1_000):
context = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid=f'workspace-{index}',
placement_generation=1,
)
session = SimpleNamespace(server_name=f'server-{index}')
loader._register_session(context, session.server_name, session)
if index == 777:
target_context = context
target_session = session
loader._sessions = NoGlobalIterationDict(loader._sessions)
assert loader._sessions_for_context(target_context) == [target_session]
@pytest.mark.asyncio
async def test_mcp_startup_concurrency_is_instance_bounded():
app = _app()
app.instance_config = SimpleNamespace(data={'mcp': {'lifecycle_concurrency': 2}})
loader = MCPLoader(app)
active = 0
maximum_active = 0
release = asyncio.Event()
async def fake_host(_context, _config):
nonlocal active, maximum_active
active += 1
maximum_active = max(maximum_active, active)
if maximum_active == 2:
release.set()
await release.wait()
await asyncio.sleep(0)
active -= 1
loader._host_mcp_server = fake_host
await asyncio.gather(
*(
loader.host_mcp_server(
TEST_EXECUTION_CONTEXT,
{'name': f'server-{index}'},
)
for index in range(20)
)
)
assert loader._lifecycle_concurrency == 2
assert maximum_active == 2
@pytest.mark.asyncio
async def test_mcp_startup_dispatcher_does_not_create_every_server_task_at_once():
app = _app()
app.instance_config = SimpleNamespace(data={'mcp': {'lifecycle_concurrency': 2}})
loader = MCPLoader(app)
started = 0
first_batch_started = asyncio.Event()
release = asyncio.Event()
async def fake_host(_context, _config):
nonlocal started
started += 1
if started == 2:
first_batch_started.set()
await release.wait()
loader.host_mcp_server = fake_host
configs = [(TEST_EXECUTION_CONTEXT, {'name': f'server-{index}'}) for index in range(20)]
dispatch_task = asyncio.create_task(loader._host_server_configs_bounded(configs))
await asyncio.wait_for(first_batch_started.wait(), timeout=1)
assert started == 2
assert len(loader._hosted_mcp_tasks) == 2
release.set()
await dispatch_task
await asyncio.sleep(0)
assert started == 20
assert loader._hosted_mcp_tasks == []
assert loader._hosted_mcp_tasks_by_scope == {}
def test_invalid_mcp_lifecycle_concurrency_uses_safe_default():
app = _app()
app.instance_config = SimpleNamespace(data={'mcp': {'lifecycle_concurrency': True}})
assert MCPLoader(app)._lifecycle_concurrency == 16
@@ -0,0 +1,70 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.provider.tools.loaders.mcp_policy import (
MCPStdioDisabledError,
require_stdio_mcp_enabled,
stdio_mcp_enabled,
)
from langbot.pkg.provider.tools.loaders.mcp import MCPLoader
def _app(config: dict) -> SimpleNamespace:
return SimpleNamespace(instance_config=SimpleNamespace(data=config))
def test_oss_default_remains_enabled_when_key_is_absent():
assert stdio_mcp_enabled(_app({})) is True
@pytest.mark.parametrize(
'value',
[False, 'false', 0, None, {}, []],
)
def test_disabled_or_invalid_values_fail_closed(value):
ap = _app({'mcp': {'stdio': {'enabled': value}}})
assert stdio_mcp_enabled(ap) is False
with pytest.raises(MCPStdioDisabledError, match='disabled by instance policy'):
require_stdio_mcp_enabled(ap, {'mode': 'stdio'})
def test_remote_transport_is_independent_of_stdio_gate():
ap = _app({'mcp': {'stdio': {'enabled': False}}})
require_stdio_mcp_enabled(ap, {'mode': 'remote'})
@pytest.mark.asyncio
async def test_bootstrap_retains_but_does_not_launch_disabled_stdio_rows():
server = SimpleNamespace(uuid='server-a', workspace_uuid='workspace-a')
result = Mock()
result.all.return_value = [server]
ap = _app({'mcp': {'stdio': {'enabled': False}}})
ap.logger = Mock()
ap.persistence_mgr = SimpleNamespace(
execute_async=AsyncMock(return_value=result),
serialize_model=Mock(
return_value={
'uuid': 'server-a',
'workspace_uuid': 'workspace-a',
'name': 'local',
'mode': 'stdio',
'enable': True,
'extra_args': {},
}
),
)
ap.workspace_service = SimpleNamespace(get_execution_binding=AsyncMock())
loader = MCPLoader(ap)
loader.host_mcp_server = AsyncMock()
await loader.load_mcp_servers_from_db()
loader.host_mcp_server.assert_not_awaited()
ap.workspace_service.get_execution_binding.assert_not_awaited()
assert loader.sessions == {}
+360 -45
View File
@@ -7,15 +7,25 @@ and error handling without calling real LLM APIs.
from __future__ import annotations
import dataclasses
import pytest
from unittest.mock import Mock
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
from langbot.pkg.provider.modelmgr.modelmgr import ModelManager
from langbot.pkg.provider.modelmgr import requester
from langbot.pkg.entity.persistence import model as persistence_model
from langbot.pkg.entity.errors import provider as provider_errors
from langbot.pkg.provider.modelmgr import token
from tests.unit_tests.provider.conftest import _make_mock_result, _make_row_mock
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.workspace.entities import WorkspaceExecutionBinding
from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError, WorkspaceInvariantError
from tests.unit_tests.provider.conftest import (
TEST_EXECUTION_CONTEXT,
TEST_WORKSPACE_UUID,
_make_mock_result,
_make_row_mock,
)
# ============================================================================
@@ -62,6 +72,63 @@ async def test_model_manager_skips_space_sync_when_disabled(mock_app_for_modelmg
app.space_service.get_models.assert_not_called()
@pytest.mark.asyncio
async def test_model_manager_skips_legacy_space_sync_in_cloud_runtime(mock_app_for_modelmgr):
"""Cloud startup must not resolve an OSS-local Workspace for legacy model sync."""
app = mock_app_for_modelmgr
app.instance_config.data = {'space': {'disable_models_service': False}}
app.persistence_mgr.mode = SimpleNamespace(value='cloud_runtime')
model_mgr = ModelManager(app)
model_mgr.load_models_from_db = AsyncMock()
await model_mgr.initialize()
app.workspace_service.get_local_execution_binding.assert_not_awaited()
@pytest.mark.asyncio
async def test_sync_new_models_from_space_creates_rerank_models(mock_app_for_modelmgr):
"""Space rerank entries are discovered and persisted under the shared provider."""
app = mock_app_for_modelmgr
provider = persistence_model.ModelProvider(
uuid='space-provider',
name='LangBot Space',
requester='space-chat-completions',
base_url='https://api.langbot.cloud/v1',
api_keys=['space-key'],
)
app.persistence_mgr.execute_async = AsyncMock(return_value=_make_mock_result([provider], first_item=provider))
app.space_service.get_models = AsyncMock(
return_value=[
SimpleNamespace(
uuid='rerank-model-uuid',
model_id='Qwen3-Reranker-8B',
category='rerank',
featured_order=10,
)
]
)
app.llm_model_service.get_llm_models = AsyncMock(return_value=[])
app.embedding_models_service.get_embedding_models = AsyncMock(return_value=[])
app.rerank_models_service = AsyncMock()
app.rerank_models_service.get_rerank_models = AsyncMock(return_value=[])
model_mgr = ModelManager(app)
await model_mgr.sync_new_models_from_space(TEST_EXECUTION_CONTEXT)
app.rerank_models_service.create_rerank_model.assert_awaited_once_with(
TEST_EXECUTION_CONTEXT,
{
'uuid': 'rerank-model-uuid',
'name': 'Qwen3-Reranker-8B',
'provider_uuid': 'space-provider',
'extra_args': {},
'prefered_ranking': 10,
},
preserve_uuid=True,
)
# ============================================================================
# Model Loading Tests
# ============================================================================
@@ -91,13 +158,33 @@ async def test_model_manager_load_models_from_db(fake_requester_registry, fake_p
# Check providers loaded
assert len(model_mgr.provider_dict) == 2
assert fake_persistence_data['provider_uuid'] in model_mgr.provider_dict
assert fake_persistence_data['provider_uuid2'] in model_mgr.provider_dict
assert {provider.provider_entity.uuid for provider in model_mgr.provider_dict.values()} == {
fake_persistence_data['provider_uuid'],
fake_persistence_data['provider_uuid2'],
}
# Check models loaded
assert len(model_mgr.llm_models) == 2
assert len(model_mgr.embedding_models) == 1
assert len(model_mgr.rerank_models) == 1
assert len(model_mgr.llm_model_dict) == 2
assert len(model_mgr.embedding_model_dict) == 1
assert len(model_mgr.rerank_model_dict) == 1
@pytest.mark.asyncio
async def test_empty_cloud_workspace_does_not_retain_generation(
mock_app_for_modelmgr,
):
model_mgr = ModelManager(mock_app_for_modelmgr)
await model_mgr._load_workspace_models(TEST_EXECUTION_CONTEXT)
assert model_mgr.provider_dict == {}
assert model_mgr.llm_model_dict == {}
assert model_mgr.embedding_model_dict == {}
assert model_mgr.rerank_model_dict == {}
assert model_mgr._scope_generations == {}
await model_mgr.resolve_execution_context(TEST_EXECUTION_CONTEXT)
assert model_mgr._scope_generations == {}
@pytest.mark.asyncio
@@ -118,7 +205,7 @@ async def test_model_manager_load_provider_unknown_requester(mock_app_for_modelm
}
with pytest.raises(provider_errors.RequesterNotFoundError) as exc_info:
await model_mgr.load_provider(provider_info)
await model_mgr.load_provider(TEST_EXECUTION_CONTEXT, provider_info)
assert exc_info.value.requester_name == 'non-existent-requester'
@@ -137,7 +224,7 @@ async def test_model_manager_load_provider_from_dict(fake_requester_registry):
'api_keys': ['dict-key'],
}
runtime_provider = await model_mgr.load_provider(provider_info)
runtime_provider = await model_mgr.load_provider(TEST_EXECUTION_CONTEXT, provider_info)
assert runtime_provider.provider_entity.uuid == 'dict-provider-uuid'
assert runtime_provider.provider_entity.name == 'Dict Provider'
@@ -154,7 +241,7 @@ async def test_model_manager_load_provider_from_entity(fake_requester_registry,
provider_entity = fake_persistence_data['providers'][0]
runtime_provider = await model_mgr.load_provider(provider_entity)
runtime_provider = await model_mgr.load_provider(TEST_EXECUTION_CONTEXT, provider_entity)
assert runtime_provider.provider_entity.uuid == provider_entity.uuid
assert runtime_provider.requester is not None
@@ -181,7 +268,7 @@ async def test_model_manager_get_model_by_uuid(fake_requester_registry, fake_per
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
model = await model_mgr.get_model_by_uuid('test-llm-uuid-1')
model = await model_mgr.get_model_by_uuid(TEST_EXECUTION_CONTEXT, 'test-llm-uuid-1')
assert model.model_entity.uuid == 'test-llm-uuid-1'
assert model.model_entity.name == 'TestLLM-1'
@@ -194,7 +281,7 @@ async def test_model_manager_get_model_by_uuid_not_found(fake_requester_registry
await model_mgr.initialize()
with pytest.raises(ValueError) as exc_info:
await model_mgr.get_model_by_uuid('unknown-model-uuid')
await model_mgr.get_model_by_uuid(TEST_EXECUTION_CONTEXT, 'unknown-model-uuid')
assert 'unknown-model-uuid' in str(exc_info.value)
@@ -215,7 +302,10 @@ async def test_model_manager_get_embedding_model_by_uuid(fake_requester_registry
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
model = await model_mgr.get_embedding_model_by_uuid('test-embedding-uuid-1')
model = await model_mgr.get_embedding_model_by_uuid(
TEST_EXECUTION_CONTEXT,
'test-embedding-uuid-1',
)
assert model.model_entity.uuid == 'test-embedding-uuid-1'
@@ -227,7 +317,10 @@ async def test_model_manager_get_embedding_model_by_uuid_not_found(fake_requeste
await model_mgr.initialize()
with pytest.raises(ValueError):
await model_mgr.get_embedding_model_by_uuid('unknown-embedding-uuid')
await model_mgr.get_embedding_model_by_uuid(
TEST_EXECUTION_CONTEXT,
'unknown-embedding-uuid',
)
@pytest.mark.asyncio
@@ -246,7 +339,7 @@ async def test_model_manager_get_rerank_model_by_uuid(fake_requester_registry, f
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
model = await model_mgr.get_rerank_model_by_uuid('test-rerank-uuid-1')
model = await model_mgr.get_rerank_model_by_uuid(TEST_EXECUTION_CONTEXT, 'test-rerank-uuid-1')
assert model.model_entity.uuid == 'test-rerank-uuid-1'
@@ -258,7 +351,7 @@ async def test_model_manager_get_rerank_model_by_uuid_not_found(fake_requester_r
await model_mgr.initialize()
with pytest.raises(ValueError):
await model_mgr.get_rerank_model_by_uuid('unknown-rerank-uuid')
await model_mgr.get_rerank_model_by_uuid(TEST_EXECUTION_CONTEXT, 'unknown-rerank-uuid')
# ============================================================================
@@ -282,12 +375,12 @@ async def test_model_manager_remove_llm_model(fake_requester_registry, fake_pers
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
assert len(model_mgr.llm_models) == 2
assert len(model_mgr.llm_model_dict) == 2
await model_mgr.remove_llm_model('test-llm-uuid-1')
await model_mgr.remove_llm_model(TEST_EXECUTION_CONTEXT, 'test-llm-uuid-1')
assert len(model_mgr.llm_models) == 1
assert model_mgr.llm_models[0].model_entity.uuid == 'test-llm-uuid-2'
assert len(model_mgr.llm_model_dict) == 1
assert next(iter(model_mgr.llm_model_dict.values())).model_entity.uuid == 'test-llm-uuid-2'
@pytest.mark.asyncio
@@ -306,12 +399,12 @@ async def test_model_manager_remove_llm_model_not_found(fake_requester_registry,
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
original_count = len(model_mgr.llm_models)
original_count = len(model_mgr.llm_model_dict)
# Removing unknown model should do nothing (no error)
await model_mgr.remove_llm_model('unknown-model-uuid')
await model_mgr.remove_llm_model(TEST_EXECUTION_CONTEXT, 'unknown-model-uuid')
assert len(model_mgr.llm_models) == original_count
assert len(model_mgr.llm_model_dict) == original_count
@pytest.mark.asyncio
@@ -330,11 +423,11 @@ async def test_model_manager_remove_embedding_model(fake_requester_registry, fak
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
assert len(model_mgr.embedding_models) == 1
assert len(model_mgr.embedding_model_dict) == 1
await model_mgr.remove_embedding_model('test-embedding-uuid-1')
await model_mgr.remove_embedding_model(TEST_EXECUTION_CONTEXT, 'test-embedding-uuid-1')
assert len(model_mgr.embedding_models) == 0
assert len(model_mgr.embedding_model_dict) == 0
@pytest.mark.asyncio
@@ -353,11 +446,11 @@ async def test_model_manager_remove_rerank_model(fake_requester_registry, fake_p
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
assert len(model_mgr.rerank_models) == 1
assert len(model_mgr.rerank_model_dict) == 1
await model_mgr.remove_rerank_model('test-rerank-uuid-1')
await model_mgr.remove_rerank_model(TEST_EXECUTION_CONTEXT, 'test-rerank-uuid-1')
assert len(model_mgr.rerank_models) == 0
assert len(model_mgr.rerank_model_dict) == 0
@pytest.mark.asyncio
@@ -376,11 +469,17 @@ async def test_model_manager_remove_provider(fake_requester_registry, fake_persi
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
assert fake_persistence_data['provider_uuid'] in model_mgr.provider_dict
assert any(
provider.provider_entity.uuid == fake_persistence_data['provider_uuid']
for provider in model_mgr.provider_dict.values()
)
await model_mgr.remove_provider(fake_persistence_data['provider_uuid'])
await model_mgr.remove_provider(TEST_EXECUTION_CONTEXT, fake_persistence_data['provider_uuid'])
assert fake_persistence_data['provider_uuid'] not in model_mgr.provider_dict
assert all(
provider.provider_entity.uuid != fake_persistence_data['provider_uuid']
for provider in model_mgr.provider_dict.values()
)
# ============================================================================
@@ -498,7 +597,7 @@ async def test_model_manager_init_temporary_runtime_llm_model(fake_requester_reg
'extra_args': {'temperature': 0.5},
}
runtime_model = await model_mgr.init_temporary_runtime_llm_model(model_info)
runtime_model = await model_mgr.init_temporary_runtime_llm_model(TEST_EXECUTION_CONTEXT, model_info)
assert runtime_model.model_entity.uuid == 'temp-model-uuid'
assert runtime_model.model_entity.name == 'TempModel'
@@ -528,7 +627,10 @@ async def test_model_manager_init_temporary_runtime_embedding_model(fake_request
'extra_args': {'dimensions': 512},
}
runtime_model = await model_mgr.init_temporary_runtime_embedding_model(model_info)
runtime_model = await model_mgr.init_temporary_runtime_embedding_model(
TEST_EXECUTION_CONTEXT,
model_info,
)
assert runtime_model.model_entity.uuid == 'temp-embedding-uuid'
assert runtime_model.model_entity.name == 'TempEmbedding'
@@ -553,7 +655,10 @@ async def test_model_manager_init_temporary_runtime_rerank_model(fake_requester_
'extra_args': {},
}
runtime_model = await model_mgr.init_temporary_runtime_rerank_model(model_info)
runtime_model = await model_mgr.init_temporary_runtime_rerank_model(
TEST_EXECUTION_CONTEXT,
model_info,
)
assert runtime_model.model_entity.uuid == 'temp-rerank-uuid'
assert runtime_model.model_entity.name == 'TempRerank'
@@ -589,12 +694,16 @@ async def test_model_manager_reload_provider(fake_requester_registry, fake_persi
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
original_provider = model_mgr.provider_dict[fake_persistence_data['provider_uuid']]
original_provider = await model_mgr.get_provider_by_uuid(
TEST_EXECUTION_CONTEXT,
fake_persistence_data['provider_uuid'],
)
original_base_url = original_provider.provider_entity.base_url
# Setup for reload - return updated provider
async def reload_execute(query):
updated_provider = persistence_model.ModelProvider(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid=fake_persistence_data['provider_uuid'],
name='Updated Provider',
requester='fake-requester',
@@ -605,9 +714,12 @@ async def test_model_manager_reload_provider(fake_requester_registry, fake_persi
model_mgr.ap.persistence_mgr.execute_async = reload_execute
await model_mgr.reload_provider(fake_persistence_data['provider_uuid'])
await model_mgr.reload_provider(TEST_EXECUTION_CONTEXT, fake_persistence_data['provider_uuid'])
updated_provider = model_mgr.provider_dict[fake_persistence_data['provider_uuid']]
updated_provider = await model_mgr.get_provider_by_uuid(
TEST_EXECUTION_CONTEXT,
fake_persistence_data['provider_uuid'],
)
assert updated_provider.provider_entity.base_url == 'https://updated.example.com'
assert updated_provider.provider_entity.base_url != original_base_url
@@ -624,7 +736,7 @@ async def test_model_manager_reload_provider_not_found(fake_requester_registry):
model_mgr.ap.persistence_mgr.execute_async = fake_execute
with pytest.raises(provider_errors.ProviderNotFoundError) as exc_info:
await model_mgr.reload_provider('unknown-provider-uuid')
await model_mgr.reload_provider(TEST_EXECUTION_CONTEXT, 'unknown-provider-uuid')
assert exc_info.value.provider_name == 'unknown-provider-uuid'
@@ -643,7 +755,11 @@ async def test_model_manager_load_llm_model_with_provider(
model_entity = fake_persistence_data['llm_models'][0]
runtime_model = await model_mgr.load_llm_model_with_provider(model_entity, runtime_provider)
runtime_model = await model_mgr.load_llm_model_with_provider(
TEST_EXECUTION_CONTEXT,
model_entity,
runtime_provider,
)
assert runtime_model.model_entity.uuid == model_entity.uuid
assert runtime_model.provider is runtime_provider
@@ -659,7 +775,11 @@ async def test_model_manager_load_llm_model_with_provider_from_row(
model_entity = fake_persistence_data['llm_models'][0]
row_mock = _make_row_mock(model_entity)
runtime_model = await model_mgr.load_llm_model_with_provider(row_mock, runtime_provider)
runtime_model = await model_mgr.load_llm_model_with_provider(
TEST_EXECUTION_CONTEXT,
row_mock,
runtime_provider,
)
assert runtime_model.model_entity.uuid == model_entity.uuid
@@ -673,7 +793,11 @@ async def test_model_manager_load_embedding_model_with_provider(
model_entity = fake_persistence_data['embedding_models'][0]
runtime_model = await model_mgr.load_embedding_model_with_provider(model_entity, runtime_provider)
runtime_model = await model_mgr.load_embedding_model_with_provider(
TEST_EXECUTION_CONTEXT,
model_entity,
runtime_provider,
)
assert runtime_model.model_entity.uuid == model_entity.uuid
assert runtime_model.provider is runtime_provider
@@ -692,6 +816,7 @@ async def test_model_manager_load_rerank_model_with_provider(fake_requester_regi
)
await requester_inst.initialize()
provider = requester.RuntimeProvider(
execution_context=TEST_EXECUTION_CONTEXT,
provider_entity=provider_entity,
token_mgr=token_mgr,
requester=requester_inst,
@@ -699,7 +824,11 @@ async def test_model_manager_load_rerank_model_with_provider(fake_requester_regi
model_entity = fake_persistence_data['rerank_models'][0]
runtime_model = await model_mgr.load_rerank_model_with_provider(model_entity, provider)
runtime_model = await model_mgr.load_rerank_model_with_provider(
TEST_EXECUTION_CONTEXT,
model_entity,
provider,
)
assert runtime_model.model_entity.uuid == model_entity.uuid
assert runtime_model.provider is provider
@@ -723,6 +852,7 @@ async def test_model_manager_logs_warning_for_missing_provider(fake_requester_re
elif 'llm_models' in query_str:
# Return model with missing provider
fake_model = persistence_model.LLMModel(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid='model-with-missing-provider',
name='MissingProviderModel',
provider_uuid='missing-provider-uuid',
@@ -736,7 +866,7 @@ async def test_model_manager_logs_warning_for_missing_provider(fake_requester_re
await model_mgr.initialize()
# Should have logged warning and skipped the model
assert len(model_mgr.llm_models) == 0
assert len(model_mgr.llm_model_dict) == 0
model_mgr.ap.logger.warning.assert_called()
@@ -750,6 +880,7 @@ async def test_model_manager_handles_requester_not_found_gracefully(fake_request
if 'model_providers' in query_str:
# Return provider with unknown requester
fake_provider = persistence_model.ModelProvider(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid='provider-with-unknown-requester',
name='Unknown Requester Provider',
requester='unknown-requester-name',
@@ -759,6 +890,7 @@ async def test_model_manager_handles_requester_not_found_gracefully(fake_request
return _make_mock_result([_make_row_mock(fake_provider)])
elif 'llm_models' in query_str:
fake_model = persistence_model.LLMModel(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid='model-uuid',
name='Model',
provider_uuid='provider-with-unknown-requester',
@@ -773,7 +905,7 @@ async def test_model_manager_handles_requester_not_found_gracefully(fake_request
# Provider should be skipped
assert len(model_mgr.provider_dict) == 0
assert len(model_mgr.llm_models) == 0
assert len(model_mgr.llm_model_dict) == 0
model_mgr.ap.logger.warning.assert_called()
@@ -790,6 +922,189 @@ def test_requester_not_found_error_str():
assert error.requester_name == 'test-requester'
@pytest.mark.asyncio
async def test_runtime_cache_isolates_same_resource_uuid_between_workspaces(fake_requester_registry):
"""A UUID collision cannot select another Workspace's runtime object."""
model_mgr = fake_requester_registry
await model_mgr.initialize()
contexts = {
workspace_uuid: ExecutionContext(
instance_uuid='test-instance',
workspace_uuid=workspace_uuid,
placement_generation=1,
)
for workspace_uuid in ('workspace-a', 'workspace-b')
}
async def resolve_binding(workspace_uuid, *, expected_generation=None):
assert expected_generation in (None, 1)
return WorkspaceExecutionBinding(
instance_uuid='test-instance',
workspace_uuid=workspace_uuid,
placement_generation=1,
write_fenced=False,
state='active',
)
model_mgr.ap.workspace_service.get_execution_binding = AsyncMock(side_effect=resolve_binding)
for workspace_uuid, context in contexts.items():
provider = await model_mgr.load_provider(
context,
{
'uuid': 'shared-provider',
'name': f'Provider {workspace_uuid}',
'requester': 'fake-requester',
'base_url': f'https://{workspace_uuid}.example.com',
'api_keys': [],
},
)
await model_mgr.cache_provider(context, provider)
runtime_model = await model_mgr.load_llm_model_with_provider(
context,
persistence_model.LLMModel(
workspace_uuid=workspace_uuid,
uuid='shared-model',
name=f'Model {workspace_uuid}',
provider_uuid='shared-provider',
abilities=[],
extra_args={},
),
provider,
)
await model_mgr.cache_llm_model(context, runtime_model)
workspace_a_model = await model_mgr.get_model_by_uuid(contexts['workspace-a'], 'shared-model')
workspace_b_model = await model_mgr.get_model_by_uuid(contexts['workspace-b'], 'shared-model')
assert workspace_a_model.model_entity.name == 'Model workspace-a'
assert workspace_b_model.model_entity.name == 'Model workspace-b'
assert workspace_a_model is not workspace_b_model
@pytest.mark.asyncio
async def test_runtime_cache_rejects_stale_placement_generation(fake_requester_registry):
"""A stale generation is fenced before any cached model can be returned."""
model_mgr = fake_requester_registry
await model_mgr.initialize()
stale_context = TEST_EXECUTION_CONTEXT
async def reject_stale(_workspace_uuid, *, expected_generation=None):
if expected_generation == stale_context.placement_generation:
raise WorkspaceGenerationMismatchError('stale generation')
raise AssertionError('lookup must include the supplied generation')
model_mgr.ap.workspace_service.get_execution_binding = AsyncMock(side_effect=reject_stale)
with pytest.raises(WorkspaceGenerationMismatchError, match='stale generation'):
await model_mgr.get_model_by_uuid(stale_context, 'any-model')
def test_generation_advance_prunes_superseded_model_runtime_objects():
class NoGlobalIterationDict(dict):
def __iter__(self):
raise AssertionError('generation advance scanned every model runtime')
def items(self):
raise AssertionError('generation advance scanned every model runtime')
def keys(self):
raise AssertionError('generation advance scanned every model runtime')
model_mgr = ModelManager(Mock())
old_context = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
new_context = dataclasses.replace(old_context, placement_generation=2)
model_mgr._observe_execution_context(old_context)
for cache in (
model_mgr.provider_dict,
model_mgr.llm_model_dict,
model_mgr.embedding_model_dict,
model_mgr.rerank_model_dict,
):
model_mgr._cache_set(
cache,
('instance-a', 'workspace-a', 1, 'resource-a'),
object(),
)
model_mgr._cache_set(
cache,
('instance-a', 'workspace-b', 1, 'resource-b'),
object(),
)
model_mgr.provider_dict = NoGlobalIterationDict(model_mgr.provider_dict)
model_mgr.llm_model_dict = NoGlobalIterationDict(model_mgr.llm_model_dict)
model_mgr.embedding_model_dict = NoGlobalIterationDict(model_mgr.embedding_model_dict)
model_mgr.rerank_model_dict = NoGlobalIterationDict(model_mgr.rerank_model_dict)
model_mgr._observe_execution_context(new_context)
for cache in (
model_mgr.provider_dict,
model_mgr.llm_model_dict,
model_mgr.embedding_model_dict,
model_mgr.rerank_model_dict,
):
assert ('instance-a', 'workspace-a', 1, 'resource-a') not in cache
assert ('instance-a', 'workspace-b', 1, 'resource-b') in cache
with pytest.raises(WorkspaceInvariantError, match='rolled back'):
model_mgr._observe_execution_context(old_context)
@pytest.mark.asyncio
async def test_generation_advance_closes_retired_provider_requester(
fake_requester_registry,
runtime_provider,
):
model_mgr = fake_requester_registry
runtime_provider.requester.aclose = AsyncMock()
await model_mgr.cache_provider(TEST_EXECUTION_CONTEXT, runtime_provider)
next_context = dataclasses.replace(
TEST_EXECUTION_CONTEXT,
placement_generation=2,
)
model_mgr.ap.workspace_service.get_execution_binding = AsyncMock(
return_value=WorkspaceExecutionBinding(
instance_uuid=next_context.instance_uuid,
workspace_uuid=next_context.workspace_uuid,
placement_generation=next_context.placement_generation,
write_fenced=False,
state='active',
)
)
await model_mgr.resolve_execution_context(next_context)
runtime_provider.requester.aclose.assert_awaited_once_with()
assert model_mgr.provider_dict == {}
assert model_mgr._scope_generations == {}
@pytest.mark.asyncio
async def test_model_manager_shutdown_closes_all_requesters_once(
fake_requester_registry,
runtime_provider,
):
model_mgr = fake_requester_registry
runtime_provider.requester.aclose = AsyncMock()
await model_mgr.cache_provider(TEST_EXECUTION_CONTEXT, runtime_provider)
await model_mgr.shutdown()
await model_mgr.shutdown()
runtime_provider.requester.aclose.assert_awaited_once_with()
assert model_mgr.provider_dict == {}
assert model_mgr.llm_model_dict == {}
assert model_mgr.embedding_model_dict == {}
assert model_mgr.rerank_model_dict == {}
def test_provider_not_found_error_str():
"""Test ProviderNotFoundError string representation."""
error = provider_errors.ProviderNotFoundError('test-provider')
+77 -122
View File
@@ -5,46 +5,15 @@ from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
import langbot_plugin.api.entities.builtin.platform.entities as platform_entities
import langbot_plugin.api.entities.builtin.platform.events as platform_events
import langbot_plugin.api.entities.builtin.platform.message as platform_message
import langbot_plugin.api.entities.builtin.provider.session as provider_session
from langbot.pkg.api.http.service.model import _runtime_model_data
from langbot.pkg.agent.runner.descriptor import AgentRunnerDescriptor
from langbot.pkg.api.http.service.provider import ModelProviderService
from langbot.pkg.entity.persistence import model as persistence_model
from langbot.pkg.pipeline.preproc.preproc import PreProcessor
from langbot.pkg.provider.modelmgr import requester
from langbot.pkg.provider.modelmgr.modelmgr import ModelManager
from langbot.pkg.provider.modelmgr.token import TokenManager
DEFAULT_RUNNER_ID = 'plugin:langbot-team/LocalAgent/default'
class FakeAgentRunnerRegistry:
async def get(self, runner_id, bound_plugins=None):
return AgentRunnerDescriptor(
id=runner_id,
source='plugin',
label={'en_US': 'Local Agent'},
plugin_author='langbot-team',
plugin_name='LocalAgent',
runner_name='default',
config_schema=[
{'name': 'model', 'type': 'model-fallback-selector'},
{'name': 'prompt', 'type': 'prompt-editor', 'default': []},
{'name': 'knowledge-bases', 'type': 'knowledge-base-multi-selector', 'default': []},
],
capabilities={'tool_calling': True, 'knowledge_retrieval': True, 'multimodal_input': True},
permissions={
'models': ['invoke', 'stream'],
'tools': ['detail', 'call'],
'knowledge_bases': ['list', 'retrieve'],
},
)
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.workspace.entities import WorkspaceExecutionBinding
def test_runtime_llm_model_data_preserves_uuid_after_update_payload_uuid_removed():
@@ -121,11 +90,22 @@ async def test_model_manager_initialize_skips_space_sync_after_timeout():
ap.discover = SimpleNamespace(get_components_by_kind=Mock(return_value=[]))
ap.instance_config = SimpleNamespace(data={'space': {'models_sync_timeout': 0.01}})
ap.logger = Mock()
binding = WorkspaceExecutionBinding(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
write_fenced=False,
state='active',
)
ap.workspace_service = SimpleNamespace(
get_local_execution_binding=AsyncMock(return_value=binding),
get_execution_binding=AsyncMock(return_value=binding),
)
mgr = ModelManager(ap)
mgr.load_models_from_db = AsyncMock()
async def slow_sync():
async def slow_sync(_context):
await asyncio.sleep(1)
mgr.sync_new_models_from_space = AsyncMock(side_effect=slow_sync)
@@ -138,39 +118,80 @@ async def test_model_manager_initialize_skips_space_sync_after_timeout():
@pytest.mark.asyncio
async def test_updated_llm_model_is_immediately_usable_by_local_agent_pipeline():
async def test_updated_llm_model_immediately_refreshes_runtime_cache():
from langbot.pkg.api.http.service.model import LLMModelsService
model_uuid = 'qwen-model-uuid'
provider_uuid = 'ollama-provider-uuid'
workspace_uuid = 'workspace-test'
execution_context = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid=workspace_uuid,
placement_generation=1,
bot_uuid='bot-uuid',
pipeline_uuid='pipeline-uuid',
)
ap = SimpleNamespace()
ap.logger = Mock()
ap.agent_runner_registry = FakeAgentRunnerRegistry()
ap.persistence_mgr = SimpleNamespace(execute_async=AsyncMock())
ap.tool_mgr = SimpleNamespace(get_all_tools=AsyncMock(return_value=[]))
ap.skill_mgr = None # PreProcessor only uses skill_mgr for the local-agent skill-binding branch
ap.plugin_connector = SimpleNamespace(
emit_event=AsyncMock(return_value=SimpleNamespace(event=SimpleNamespace(default_prompt=[], prompt=[])))
binding = WorkspaceExecutionBinding(
instance_uuid='instance-test',
workspace_uuid=workspace_uuid,
placement_generation=1,
write_fenced=False,
state='active',
)
ap.workspace_service = SimpleNamespace(get_execution_binding=AsyncMock(return_value=binding))
ap.model_mgr = ModelManager(ap)
runtime_provider = Mock()
ap.model_mgr.provider_dict = {provider_uuid: runtime_provider}
ap.model_mgr.llm_models = [
requester.RuntimeLLMModel(
model_entity=persistence_model.LLMModel(
uuid=model_uuid,
name='old-qwen-name',
provider_uuid=provider_uuid,
abilities=[],
extra_args={},
),
provider=runtime_provider,
)
]
runtime_provider = Mock(
execution_context=execution_context,
provider_entity=persistence_model.ModelProvider(
workspace_uuid=workspace_uuid,
uuid=provider_uuid,
name='Ollama',
requester='ollama',
base_url='http://localhost:11434',
api_keys=[],
),
)
cache_key = ('instance-test', workspace_uuid, 1, provider_uuid)
ap.model_mgr.provider_dict = {cache_key: runtime_provider}
runtime_model = requester.RuntimeLLMModel(
execution_context=execution_context,
model_entity=persistence_model.LLMModel(
workspace_uuid=workspace_uuid,
uuid=model_uuid,
name='old-qwen-name',
provider_uuid=provider_uuid,
abilities=[],
extra_args={},
),
provider=runtime_provider,
)
ap.model_mgr.llm_model_dict = {
('instance-test', workspace_uuid, 1, model_uuid): runtime_model,
}
await LLMModelsService(ap).update_llm_model(
ap.provider_service = SimpleNamespace(
get_provider=AsyncMock(return_value={'uuid': provider_uuid, 'workspace_uuid': workspace_uuid})
)
model_service = LLMModelsService(ap)
model_service.get_llm_model = AsyncMock(
return_value={
'uuid': model_uuid,
'workspace_uuid': workspace_uuid,
'name': 'old-qwen-name',
'provider_uuid': provider_uuid,
'abilities': [],
'context_length': None,
'extra_args': {},
'prefered_ranking': 0,
}
)
await model_service.update_llm_model(
workspace_uuid,
model_uuid,
{
'name': 'Qwen3.5-27B',
@@ -180,72 +201,6 @@ async def test_updated_llm_model_is_immediately_usable_by_local_agent_pipeline()
},
)
runtime_model = await ap.model_mgr.get_model_by_uuid(model_uuid)
runtime_model = await ap.model_mgr.get_model_by_uuid(execution_context, model_uuid)
assert runtime_model.model_entity.uuid == model_uuid
assert runtime_model.model_entity.name == 'Qwen3.5-27B'
session = SimpleNamespace(
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
)
conversation = SimpleNamespace(
uuid='conversation-uuid',
create_time=None,
update_time=None,
prompt=SimpleNamespace(messages=[], copy=Mock(return_value=SimpleNamespace(messages=[]))),
messages=[],
)
ap.sess_mgr = SimpleNamespace(
get_session=AsyncMock(return_value=session),
get_conversation=AsyncMock(return_value=conversation),
)
message_chain = platform_message.MessageChain([platform_message.Plain(text='hello')])
sender = platform_entities.Friend(id=12345, nickname='Tester', remark=None)
message_event = platform_events.FriendMessage(
type='FriendMessage',
sender=sender,
message_chain=message_chain,
time=1710000000,
)
pipeline_config = {
'ai': {
'runner': {'id': DEFAULT_RUNNER_ID},
'runner_config': {
DEFAULT_RUNNER_ID: {
'model': {'primary': model_uuid, 'fallbacks': []},
'prompt': [],
'knowledge-bases': [],
},
},
},
'trigger': {'misc': {'combine-quote-message': False}},
'output': {'misc': {'remove-think': False}},
}
query = pipeline_query.Query.model_construct(
query_id='query-id',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
sender_id=12345,
message_chain=message_chain,
message_event=message_event,
adapter=AsyncMock(),
pipeline_uuid='pipeline-uuid',
bot_uuid='bot-uuid',
pipeline_config=pipeline_config,
session=None,
prompt=None,
messages=[],
user_message=None,
use_funcs=[],
use_llm_model_uuid=None,
variables={},
resp_messages=[],
resp_message_chain=None,
current_stage_name=None,
)
result = await PreProcessor(ap).process(query, 'PreProcessor')
processed_query = result.new_query
assert processed_query.use_llm_model_uuid == model_uuid
@@ -15,6 +15,7 @@ from langbot.pkg.provider.modelmgr import requester
from langbot.pkg.provider.modelmgr import token
from langbot.pkg.entity.persistence import model as persistence_model
from langbot.pkg.provider.modelmgr.errors import RequesterError
from tests.unit_tests.provider.conftest import TEST_EXECUTION_CONTEXT, TEST_WORKSPACE_UUID
# ============================================================================
@@ -134,6 +135,7 @@ async def test_requester_invoke_rerank_not_implemented():
# Create fake model
fake_provider_entity = persistence_model.ModelProvider(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid='provider-uuid',
name='Provider',
requester='test',
@@ -143,17 +145,20 @@ async def test_requester_invoke_rerank_not_implemented():
fake_token_mgr = token.TokenManager(name='test', tokens=[])
fake_requester = inst
fake_provider = requester.RuntimeProvider(
execution_context=TEST_EXECUTION_CONTEXT,
provider_entity=fake_provider_entity,
token_mgr=fake_token_mgr,
requester=fake_requester,
)
fake_model_entity = persistence_model.RerankModel(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid='model-uuid',
name='Model',
provider_uuid='provider-uuid',
extra_args={},
)
fake_model = requester.RuntimeRerankModel(
execution_context=TEST_EXECUTION_CONTEXT,
model_entity=fake_model_entity,
provider=fake_provider,
)
@@ -289,6 +294,7 @@ async def test_runtime_provider_invoke_llm_delegates(runtime_provider, runtime_l
resp_message_chain=None,
current_stage_name=None,
)
object.__setattr__(query, '_execution_context', TEST_EXECUTION_CONTEXT)
messages = [
provider_message.Message(role='user', content=[provider_message.ContentElement(type='text', text='Hello')])
@@ -332,6 +338,7 @@ async def test_runtime_provider_invoke_llm_stashes_usage(runtime_provider, runti
resp_message_chain=None,
current_stage_name=None,
)
object.__setattr__(query, '_execution_context', TEST_EXECUTION_CONTEXT)
usage = {
'prompt_tokens': 11,
'completion_tokens': 7,
@@ -385,6 +392,7 @@ async def test_runtime_provider_invoke_llm_stream_yields_chunks(runtime_provider
resp_message_chain=None,
current_stage_name=None,
)
object.__setattr__(query, '_execution_context', TEST_EXECUTION_CONTEXT)
messages = [
provider_message.Message(role='user', content=[provider_message.ContentElement(type='text', text='Hello')])
@@ -428,6 +436,7 @@ async def test_runtime_provider_invoke_llm_stream_stashes_usage(runtime_provider
resp_message_chain=None,
current_stage_name=None,
)
object.__setattr__(query, '_execution_context', TEST_EXECUTION_CONTEXT)
usage = {
'prompt_tokens': 13,
'completion_tokens': 2,
@@ -458,7 +467,11 @@ async def test_runtime_provider_invoke_embedding_returns_vectors(runtime_provide
"""Test RuntimeProvider.invoke_embedding returns embedding vectors."""
provider = runtime_provider
result = await provider.invoke_embedding(runtime_embedding_model, ['text1', 'text2'])
result = await provider.invoke_embedding(
runtime_embedding_model,
['text1', 'text2'],
execution_context=TEST_EXECUTION_CONTEXT,
)
assert len(result) == 2
assert result[0] == [0.1, 0.2, 0.3]
@@ -470,7 +483,12 @@ async def test_runtime_provider_invoke_rerank_returns_scores(runtime_provider, r
# Need to use the correct provider for rerank model
provider = runtime_rerank_model.provider
result = await provider.invoke_rerank(runtime_rerank_model, 'query', ['doc1', 'doc2', 'doc3'])
result = await provider.invoke_rerank(
runtime_rerank_model,
'query',
['doc1', 'doc2', 'doc3'],
execution_context=TEST_EXECUTION_CONTEXT,
)
assert len(result) == 3
assert result[0]['index'] == 0
@@ -640,6 +658,7 @@ async def test_runtime_provider_invoke_llm_propagates_error(mock_app_for_modelmg
await requester_inst.initialize()
provider_entity = persistence_model.ModelProvider(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid='error-provider',
name='Error Provider',
requester='error-requester',
@@ -649,19 +668,25 @@ async def test_runtime_provider_invoke_llm_propagates_error(mock_app_for_modelmg
token_mgr = token.TokenManager(name='error-provider', tokens=['error-key'])
provider = requester.RuntimeProvider(
execution_context=TEST_EXECUTION_CONTEXT,
provider_entity=provider_entity,
token_mgr=token_mgr,
requester=requester_inst,
)
model_entity = persistence_model.LLMModel(
workspace_uuid=TEST_WORKSPACE_UUID,
uuid='error-model',
name='Error Model',
provider_uuid='error-provider',
abilities=[],
extra_args={},
)
model = requester.RuntimeLLMModel(model_entity=model_entity, provider=provider)
model = requester.RuntimeLLMModel(
execution_context=TEST_EXECUTION_CONTEXT,
model_entity=model_entity,
provider=provider,
)
import langbot_plugin.api.entities.builtin.provider.message as provider_message
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
@@ -688,6 +713,7 @@ async def test_runtime_provider_invoke_llm_propagates_error(mock_app_for_modelmg
resp_message_chain=None,
current_stage_name=None,
)
object.__setattr__(query, '_execution_context', TEST_EXECUTION_CONTEXT)
messages = [
provider_message.Message(role='user', content=[provider_message.ContentElement(type='text', text='Hello')])
@@ -10,11 +10,74 @@ from __future__ import annotations
import pytest
import asyncio
from types import SimpleNamespace
from unittest.mock import Mock
from importlib import import_module
import langbot_plugin.api.entities.builtin.provider.session as provider_session
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
import langbot_plugin.api.entities.builtin.provider.message as provider_message
import langbot_plugin.api.entities.builtin.provider.prompt as provider_prompt
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.pipeline.pool import (
ExecutionContextMismatchError,
ExecutionContextRequiredError,
)
TEST_CONTEXT = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
)
TEST_BOT_UUID = 'bot-123'
def bind_query_context(query, *, context=TEST_CONTEXT, bot_uuid=TEST_BOT_UUID):
"""Attach the trusted runtime scope expected by the session manager."""
query.bot_uuid = bot_uuid
query._execution_context = context
return query
def bind_session_context(session, query):
"""Make a mocked legacy Session belong to the Query execution scope."""
session.bot_uuid = query.bot_uuid
session._langbot_session_key = (
TEST_CONTEXT.instance_uuid,
TEST_CONTEXT.workspace_uuid,
TEST_CONTEXT.placement_generation,
query.bot_uuid,
query.launcher_type.value,
query.launcher_id,
)
return session
def scoped_query(
*,
workspace_uuid='workspace-test',
bot_uuid=TEST_BOT_UUID,
placement_generation=1,
pipeline_uuid=None,
):
"""Create a small Query-like object with a complete trusted scope."""
context = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid=workspace_uuid,
placement_generation=placement_generation,
bot_uuid=bot_uuid,
pipeline_uuid=pipeline_uuid,
query_uuid=f'query-{workspace_uuid}-{bot_uuid}',
)
return SimpleNamespace(
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id='same-launcher',
sender_id='same-sender',
bot_uuid=bot_uuid,
_execution_context=context,
)
def get_session_module():
@@ -71,7 +134,7 @@ class TestSessionManagerGetSession:
query.launcher_type = provider_session.LauncherTypes.PERSON
query.launcher_id = '12345'
query.sender_id = '12345'
return query
return bind_query_context(query)
@pytest.mark.asyncio
async def test_creates_new_session_when_not_found(self, mock_app_with_config, sample_query):
@@ -126,11 +189,13 @@ class TestSessionManagerGetSession:
query1.launcher_type = provider_session.LauncherTypes.PERSON
query1.launcher_id = 'user1'
query1.sender_id = 'user1'
bind_query_context(query1)
query2 = Mock(spec=pipeline_query.Query)
query2.launcher_type = provider_session.LauncherTypes.PERSON
query2.launcher_id = 'user2'
query2.sender_id = 'user2'
bind_query_context(query2)
session1 = await manager.get_session(query1)
session2 = await manager.get_session(query2)
@@ -149,11 +214,13 @@ class TestSessionManagerGetSession:
query1.launcher_type = provider_session.LauncherTypes.PERSON
query1.launcher_id = 'same_id'
query1.sender_id = 'same_id'
bind_query_context(query1)
query2 = Mock(spec=pipeline_query.Query)
query2.launcher_type = provider_session.LauncherTypes.GROUP
query2.launcher_id = 'same_id'
query2.sender_id = 'same_id'
bind_query_context(query2)
session1 = await manager.get_session(query1)
session2 = await manager.get_session(query2)
@@ -191,7 +258,7 @@ class TestSessionManagerGetConversation:
query.launcher_type = provider_session.LauncherTypes.PERSON
query.launcher_id = '12345'
query.sender_id = '12345'
return query
return bind_query_context(query)
@pytest.mark.asyncio
async def test_creates_conversation_with_prompt(self, mock_app_with_config, sample_query, sample_session):
@@ -199,6 +266,7 @@ class TestSessionManagerGetConversation:
sessionmgr = get_session_module()
manager = sessionmgr.SessionManager(mock_app_with_config)
bind_session_context(sample_session, sample_query)
prompt_config = [{'role': 'system', 'content': 'You are a helpful assistant.'}]
pipeline_uuid = 'pipeline-123'
@@ -222,6 +290,7 @@ class TestSessionManagerGetConversation:
sessionmgr = get_session_module()
manager = sessionmgr.SessionManager(mock_app_with_config)
bind_session_context(sample_session, sample_query)
prompt_config = [{'role': 'system', 'content': 'You are a helpful assistant.'}]
pipeline_uuid = 'pipeline-123'
@@ -244,14 +313,15 @@ class TestSessionManagerGetConversation:
sessionmgr = get_session_module()
manager = sessionmgr.SessionManager(mock_app_with_config)
bind_session_context(sample_session, sample_query)
prompt_config = [{'role': 'system', 'content': 'You are a helpful assistant.'}]
# First call with pipeline1
conv1 = await manager.get_conversation(sample_query, sample_session, prompt_config, 'pipeline-1', 'bot-1')
conv1 = await manager.get_conversation(sample_query, sample_session, prompt_config, 'pipeline-1', TEST_BOT_UUID)
# Second call with different pipeline should create new conversation
conv2 = await manager.get_conversation(sample_query, sample_session, prompt_config, 'pipeline-2', 'bot-2')
conv2 = await manager.get_conversation(sample_query, sample_session, prompt_config, 'pipeline-2', TEST_BOT_UUID)
assert conv1 is not conv2
assert len(sample_session.conversations) == 2
@@ -263,6 +333,7 @@ class TestSessionManagerGetConversation:
sessionmgr = get_session_module()
manager = sessionmgr.SessionManager(mock_app_with_config)
bind_session_context(sample_session, sample_query)
prompt_config = [{'role': 'system', 'content': 'You are a helpful assistant.'}]
@@ -278,6 +349,7 @@ class TestSessionManagerGetConversation:
sessionmgr = get_session_module()
manager = sessionmgr.SessionManager(mock_app_with_config)
bind_session_context(sample_session, sample_query)
prompt_config = [{'role': 'system', 'content': 'System message'}, {'role': 'user', 'content': 'User message'}]
@@ -287,3 +359,200 @@ class TestSessionManagerGetConversation:
assert conversation.prompt.name == 'default'
assert len(conversation.prompt.messages) == 2
class TestSessionManagerWorkspaceIsolation:
"""Regression coverage for workspace, bot, and placement fencing."""
@staticmethod
def manager():
mock_app = Mock()
mock_app.instance_config.data = {'concurrency': {'session': 5}}
return get_session_module().SessionManager(mock_app)
@pytest.mark.asyncio
async def test_get_session_requires_trusted_query_scope(self):
query = SimpleNamespace(
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id='same-launcher',
sender_id='same-sender',
bot_uuid=TEST_BOT_UUID,
)
with pytest.raises(ExecutionContextRequiredError):
await self.manager().get_session(query)
@pytest.mark.asyncio
async def test_same_launcher_in_two_workspaces_does_not_share_session(self):
manager = self.manager()
first = await manager.get_session(scoped_query(workspace_uuid='workspace-a'))
second = await manager.get_session(scoped_query(workspace_uuid='workspace-b'))
assert first is not second
assert len(manager.session_list) == 2
@pytest.mark.asyncio
async def test_same_launcher_in_two_bots_does_not_share_session(self):
manager = self.manager()
first = await manager.get_session(scoped_query(bot_uuid='bot-a'))
second = await manager.get_session(scoped_query(bot_uuid='bot-b'))
assert first is not second
@pytest.mark.asyncio
async def test_new_placement_generation_does_not_reuse_old_session(self):
manager = self.manager()
first = await manager.get_session(scoped_query(placement_generation=1))
second = await manager.get_session(scoped_query(placement_generation=2))
assert first is not second
@pytest.mark.asyncio
async def test_conversation_rejects_session_from_another_workspace(self):
manager = self.manager()
query_a = scoped_query(workspace_uuid='workspace-a')
query_b = scoped_query(workspace_uuid='workspace-b')
session_a = await manager.get_session(query_a)
with pytest.raises(ExecutionContextMismatchError):
await manager.get_conversation(query_b, session_a, [], 'pipeline-1', TEST_BOT_UUID)
@pytest.mark.asyncio
async def test_conversation_rejects_substituted_bot_argument(self):
manager = self.manager()
query = scoped_query(bot_uuid='bot-a')
session = await manager.get_session(query)
with pytest.raises(ExecutionContextMismatchError):
await manager.get_conversation(query, session, [], 'pipeline-1', 'bot-b')
@pytest.mark.asyncio
async def test_per_workspace_capacity_evicts_oldest_idle_session(self):
manager = self.manager()
manager.ap.instance_config.data['system'] = {
'session_retention': {
'max_entries': 10,
'max_entries_per_workspace': 2,
}
}
queries = []
sessions = []
for index in range(3):
query = scoped_query()
query.launcher_id = f'launcher-{index}'
queries.append(query)
sessions.append(await manager.get_session(query))
assert len(manager.session_list) == 2
assert sessions[0] not in manager.session_list
assert sessions[1:] == manager.session_list
assert await manager.get_session(queries[-1]) is manager.session_list[-1]
@pytest.mark.asyncio
async def test_new_session_does_not_scan_other_workspace_sessions(self):
manager = self.manager()
manager.ap.instance_config.data['system'] = {
'session_retention': {
'max_entries': 600,
'max_entries_per_workspace': 2,
}
}
for index in range(512):
query = scoped_query(workspace_uuid=f'workspace-{index}')
query.launcher_id = f'launcher-{index}'
await manager.get_session(query)
class NoGlobalIterationDict(dict):
def __iter__(self):
raise AssertionError('global session index iteration is forbidden')
def items(self):
raise AssertionError('global session index iteration is forbidden')
def values(self):
raise AssertionError('global session index iteration is forbidden')
manager._session_index = NoGlobalIterationDict(manager._session_index)
query = scoped_query(workspace_uuid='workspace-new')
query.launcher_id = 'launcher-new'
session = await manager.get_session(query)
assert session.workspace_uuid == 'workspace-new'
assert len(manager._session_index) == 513
@pytest.mark.asyncio
async def test_stale_expiry_revision_does_not_evict_recent_session(
self,
monkeypatch,
):
sessionmgr = get_session_module()
manager = self.manager()
manager.ap.instance_config.data['system'] = {
'session_retention': {
'max_entries': 10,
'max_entries_per_workspace': 10,
'idle_ttl_seconds': 1,
}
}
clock = [0.0]
monkeypatch.setattr(sessionmgr.time, 'monotonic', lambda: clock[0])
first_query = scoped_query()
first_query.launcher_id = 'first'
first = await manager.get_session(first_query)
clock[0] = 0.5
assert await manager.get_session(first_query) is first
clock[0] = 1.25
second_query = scoped_query()
second_query.launcher_id = 'second'
await manager.get_session(second_query)
assert first in manager.session_list
clock[0] = 2.0
third_query = scoped_query()
third_query.launcher_id = 'third'
await manager.get_session(third_query)
assert first not in manager.session_list
@pytest.mark.asyncio
async def test_access_revision_heap_stays_bounded(self):
manager = self.manager()
query = scoped_query()
await manager.get_session(query)
for _ in range(1000):
await manager.get_session(query)
assert len(manager._session_expiry_heap) <= 64
def test_trim_conversation_drops_retained_binary_payloads(self):
manager = self.manager()
conversation = provider_session.Conversation(
prompt=provider_prompt.Prompt(name='test', messages=[]),
messages=[
provider_message.Message(
role='user',
content=[
provider_message.ContentElement.from_text('hello'),
provider_message.ContentElement.from_image_base64('x' * 1000000),
provider_message.ContentElement.from_file_base64(
'y' * 1000000,
'large.bin',
),
],
)
],
pipeline_uuid='pipeline-1',
bot_uuid=TEST_BOT_UUID,
)
manager.trim_conversation_messages(conversation, max_rounds=10)
content = conversation.messages[0].content
assert content[1].image_base64 is None
assert content[2].file_base64 is None
+220 -192
View File
@@ -7,6 +7,37 @@ from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.api.http.context import ExecutionContext
_CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
query_uuid='query-a',
)
def _make_query(*, variables=None, **kwargs):
return SimpleNamespace(
query_id=kwargs.pop('query_id', 'query-a'),
query_uuid=kwargs.pop('query_uuid', 'query-a'),
instance_uuid=kwargs.pop('instance_uuid', _CONTEXT.instance_uuid),
workspace_uuid=kwargs.pop('workspace_uuid', _CONTEXT.workspace_uuid),
placement_generation=kwargs.pop('placement_generation', _CONTEXT.placement_generation),
variables={} if variables is None else variables,
**kwargs,
)
def _make_skill_manager(skills: dict[str, dict], **kwargs):
return SimpleNamespace(
skills=skills,
get_skills=Mock(return_value=skills),
get_skill_by_name=Mock(side_effect=lambda _context, name: skills.get(name)),
**kwargs,
)
def _make_ap(logger=None):
ap = SimpleNamespace()
@@ -51,14 +82,14 @@ class TestSkillManagerCache:
mgr = SkillManager(ap)
# Empty cache → returns False
assert mgr.refresh_skill_from_disk('test-skill') is False
assert mgr.refresh_skill_from_disk(_CONTEXT, 'test-skill') is False
# Cache populated → returns True; method does NOT mutate the cache
cached = _make_skill_data(name='test-skill', instructions='Cached')
mgr.skills['test-skill'] = cached
assert mgr.refresh_skill_from_disk('test-skill') is True
assert mgr.skills['test-skill'] is cached
assert mgr.refresh_skill_from_disk('') is False
mgr._skills_by_scope[mgr._scope_key(_CONTEXT)] = {'test-skill': cached}
assert mgr.refresh_skill_from_disk(_CONTEXT, 'test-skill') is True
assert mgr.get_skills(_CONTEXT)['test-skill'] is cached
assert mgr.refresh_skill_from_disk(_CONTEXT, '') is False
@pytest.mark.asyncio
async def test_reload_skills_drops_box_skills_with_missing_package_root(self):
@@ -85,9 +116,9 @@ class TestSkillManagerCache:
ap.box_service = box_service
mgr = SkillManager(ap)
await mgr.reload_skills()
await mgr.reload_skills(_CONTEXT)
assert list(mgr.skills) == ['alive']
assert list(mgr.get_skills(_CONTEXT)) == ['alive']
# Warning fired with the dropped skill name so operators can see it.
warning_messages = [str(call.args[0]) for call in ap.logger.warning.call_args_list]
assert any('ghost' in msg and 'package_root missing' in msg for msg in warning_messages)
@@ -116,9 +147,9 @@ class TestSkillManagerCache:
ap.box_service = box_service
mgr = SkillManager(ap)
await mgr.reload_skills()
await mgr.reload_skills(_CONTEXT)
assert sorted(mgr.skills) == ['alpha', 'beta']
assert sorted(mgr.get_skills(_CONTEXT)) == ['alpha', 'beta']
# No skill dropped → no "package_root missing" warning.
warning_messages = [str(call.args[0]) for call in ap.logger.warning.call_args_list]
assert not any('package_root missing' in msg for msg in warning_messages)
@@ -141,12 +172,12 @@ class TestSkillActivationHelper:
ap = _make_ap()
mgr = SkillManager(ap)
mgr.skills = {
mgr._skills_by_scope[mgr._scope_key(_CONTEXT)] = {
'primary': _make_skill_data(name='primary', instructions='Primary instructions'),
}
ap.skill_mgr = mgr
query = SimpleNamespace(variables={})
query = _make_query()
assert register_activated_skill(ap, query, 'primary') is True
assert set(query.variables[ACTIVATED_SKILLS_KEY].keys()) == {'primary'}
@@ -159,10 +190,10 @@ class TestSkillActivationHelper:
ap = _make_ap()
mgr = SkillManager(ap)
mgr.skills = {'primary': _make_skill_data(name='primary')}
mgr._skills_by_scope[mgr._scope_key(_CONTEXT)] = {'primary': _make_skill_data(name='primary')}
ap.skill_mgr = mgr
query = SimpleNamespace(variables={})
query = _make_query()
assert register_activated_skill(ap, query, 'missing') is False
assert ACTIVATED_SKILLS_KEY not in query.variables
@@ -171,107 +202,46 @@ class TestSkillActivationHelper:
from langbot.pkg.skill.activation import register_activated_skill
ap = _make_ap() # no skill_mgr attribute
query = SimpleNamespace(variables={})
query = _make_query()
assert register_activated_skill(ap, query, 'primary') is False
class TestPersistActivatedSkill:
"""Host-side persistence of activated skills into conversation state (S-01/S-02)."""
@pytest.mark.asyncio
async def test_persist_writes_conversation_state(self):
from unittest.mock import patch
from langbot.pkg.provider.tools.loaders.skill import (
persist_activated_skill,
ACTIVATED_SKILLS_KEY,
ACTIVATED_SKILL_NAMES_STATE_KEY,
)
ap = _make_ap()
ap.persistence_mgr.get_db_engine = Mock(return_value=Mock())
query = SimpleNamespace(variables={ACTIVATED_SKILLS_KEY: {'pdf': {'name': 'pdf'}}})
query._agent_run_session = {
'runner_id': 'plugin:test/runner/default',
'authorization': {
'state_context': {
'scope_keys': {'conversation': 'conv-scope-key'},
'binding_identity': 'binding-1',
'conversation_id': 'c1',
},
},
}
store = SimpleNamespace(state_set=AsyncMock(return_value=(True, None)))
with patch(
'langbot.pkg.agent.runner.persistent_state_store.get_persistent_state_store',
return_value=store,
):
await persist_activated_skill(ap, query, 'pdf')
store.state_set.assert_awaited_once()
kwargs = store.state_set.await_args.kwargs
assert kwargs['scope_key'] == 'conv-scope-key'
assert kwargs['state_key'] == ACTIVATED_SKILL_NAMES_STATE_KEY
assert kwargs['value'] == ['pdf']
assert kwargs['scope'] == 'conversation'
assert kwargs['runner_id'] == 'plugin:test/runner/default'
assert kwargs['binding_identity'] == 'binding-1'
@pytest.mark.asyncio
async def test_persist_noop_without_run_session(self):
from unittest.mock import patch
from langbot.pkg.provider.tools.loaders.skill import persist_activated_skill
ap = _make_ap()
query = SimpleNamespace(variables={'_activated_skills': {'pdf': {'name': 'pdf'}}})
with patch(
'langbot.pkg.agent.runner.persistent_state_store.get_persistent_state_store',
) as mock_factory:
await persist_activated_skill(ap, query, 'pdf')
mock_factory.assert_not_called()
class TestSkillPathHelpers:
def test_get_visible_skills_filters_by_bound_names(self):
from langbot.pkg.provider.tools.loaders.skill import PIPELINE_BOUND_SKILLS_KEY, get_visible_skills
ap = _make_ap()
ap.skill_mgr = SimpleNamespace(
skills={
ap.skill_mgr = _make_skill_manager(
{
'visible': _make_skill_data(name='visible'),
'hidden': _make_skill_data(name='hidden'),
}
)
query = SimpleNamespace(variables={PIPELINE_BOUND_SKILLS_KEY: ['visible']})
query = _make_query(variables={PIPELINE_BOUND_SKILLS_KEY: ['visible']})
result = get_visible_skills(ap, query)
assert list(result.keys()) == ['visible']
def test_restore_activated_skills_from_state_filters_by_visibility(self):
def test_restore_activated_skills_uses_caller_provided_names_and_visibility(self):
from langbot.pkg.provider.tools.loaders.skill import (
ACTIVATED_SKILLS_KEY,
ACTIVATED_SKILL_NAMES_STATE_KEY,
PIPELINE_BOUND_SKILLS_KEY,
get_activated_skill_names,
restore_activated_skills_from_state,
restore_activated_skills,
)
ap = _make_ap()
ap.skill_mgr = SimpleNamespace(
skills={
ap.skill_mgr = _make_skill_manager(
{
'visible': _make_skill_data(name='visible'),
'hidden': _make_skill_data(name='hidden'),
}
)
query = SimpleNamespace(variables={PIPELINE_BOUND_SKILLS_KEY: ['visible']})
state = {'conversation': {ACTIVATED_SKILL_NAMES_STATE_KEY: ['visible', 'hidden', 'visible', '']}}
query = _make_query(variables={PIPELINE_BOUND_SKILLS_KEY: ['visible']})
restored = restore_activated_skills_from_state(ap, query, state)
restored = restore_activated_skills(ap, query, ['visible', 'hidden', 'visible', ''])
assert restored == ['visible']
assert list(query.variables[ACTIVATED_SKILLS_KEY].keys()) == ['visible']
@@ -284,8 +254,8 @@ class TestSkillPathHelpers:
)
ap = _make_ap()
ap.skill_mgr = SimpleNamespace(skills={'demo': _make_skill_data(name='demo')})
query = SimpleNamespace(variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']})
ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo')})
query = _make_query(variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']})
skill, rewritten = resolve_virtual_skill_path(
ap,
@@ -334,6 +304,22 @@ class TestSkillPathHelpers:
assert 'export VIRTUAL_ENV="$_LB_VENV_DIR"' in command
assert command.rstrip().endswith('python scripts/run.py')
def test_wrap_skill_python_env_keeps_state_outside_read_only_source(self):
from langbot.pkg.provider.tools.loaders.skill import wrap_skill_command_with_python_env
command = wrap_skill_command_with_python_env(
'python scripts/run.py',
mount_path='/workspace/.skills/demo',
state_path='/workspace/.skill-envs/demo',
)
assert '_LB_VENV_DIR="/workspace/.skill-envs/demo/.venv"' in command
assert '_LB_META_DIR="/workspace/.skill-envs/demo/.langbot"' in command
assert '_LB_TMP_DIR="/workspace/.skill-envs/demo/.tmp"' in command
assert '_LB_PIP_CACHE_DIR="/workspace/.skill-envs/demo/.cache/pip"' in command
assert 'root = "/workspace/.skills/demo"' in command
assert 'pip install "/workspace/.skills/demo"' in command
class TestSkillToolLoader:
"""The skill tool surface is now just ``activate`` + ``register_skill``.
@@ -353,13 +339,10 @@ class TestSkillToolLoader:
skill = _make_skill_data(name='demo', package_root='/data/skills/demo', instructions='Step 1')
ap = _make_ap()
ap.skill_mgr = SimpleNamespace(
skills={'demo': skill},
get_skill_by_name=lambda name: skill if name == 'demo' else None,
)
ap.skill_mgr = _make_skill_manager({'demo': skill})
loader = SkillToolLoader(ap)
query = SimpleNamespace(variables={})
query = _make_query()
result = await loader.invoke_tool(ACTIVATE_SKILL_TOOL_NAME, {'skill_name': 'demo'}, query)
@@ -378,10 +361,7 @@ class TestSkillToolLoader:
)
ap = _make_ap()
ap.skill_mgr = SimpleNamespace(
skills={'demo': _make_skill_data(name='demo')},
get_skill_by_name=lambda name: None,
)
ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo')})
loader = SkillToolLoader(ap)
@@ -389,7 +369,7 @@ class TestSkillToolLoader:
await loader.invoke_tool(
ACTIVATE_SKILL_TOOL_NAME,
{'skill_name': 'ghost'},
SimpleNamespace(variables={}),
_make_query(),
)
@pytest.mark.asyncio
@@ -423,65 +403,24 @@ class TestSkillToolLoader:
result = await loader.invoke_tool(
REGISTER_SKILL_TOOL_NAME,
{'path': '/workspace/repo'},
SimpleNamespace(),
_make_query(),
)
ap.skill_service.scan_directory_async.assert_awaited_once_with(os.path.realpath(repo_dir))
ap.skill_service.scan_directory_async.assert_awaited_once_with(_CONTEXT, os.path.realpath(repo_dir))
ap.skill_service.create_skill.assert_awaited_once_with(
_CONTEXT,
{
'name': 'cloned-skill',
'display_name': 'Cloned Skill',
'description': 'Imported from clone',
'instructions': 'Do work',
'package_root': os.path.realpath(repo_dir),
}
},
)
assert result['registered'] is True
assert result['skill_name'] == 'cloned-skill'
assert result['source_path'] == '/workspace/repo'
@pytest.mark.asyncio
async def test_registered_skill_can_be_activated_in_same_query(self):
from langbot.pkg.provider.tools.loaders.skill import PIPELINE_BOUND_SKILLS_KEY
from langbot.pkg.provider.tools.loaders.skill_authoring import (
ACTIVATE_SKILL_TOOL_NAME,
REGISTER_SKILL_TOOL_NAME,
SkillToolLoader,
)
with tempfile.TemporaryDirectory() as tmpdir:
repo_dir = os.path.join(tmpdir, 'repo')
os.makedirs(repo_dir)
created = _make_skill_data(name='cloned-skill', package_root=os.path.realpath(repo_dir))
ap = _make_ap()
ap.box_service = SimpleNamespace(default_workspace=tmpdir, available=True)
ap.skill_mgr = SimpleNamespace(skills={'existing': _make_skill_data(name='existing')})
async def create_skill(_data):
ap.skill_mgr.skills['cloned-skill'] = created
return created
ap.skill_service = SimpleNamespace(
scan_directory_async=AsyncMock(return_value=created),
create_skill=AsyncMock(side_effect=create_skill),
)
query = SimpleNamespace(variables={PIPELINE_BOUND_SKILLS_KEY: ['existing']})
loader = SkillToolLoader(ap)
await loader.invoke_tool(
REGISTER_SKILL_TOOL_NAME,
{'path': '/workspace/repo', 'name': 'cloned-skill'},
query,
)
activated = await loader.invoke_tool(
ACTIVATE_SKILL_TOOL_NAME,
{'skill_name': 'cloned-skill'},
query,
)
assert query.variables[PIPELINE_BOUND_SKILLS_KEY] == ['existing', 'cloned-skill']
assert activated['activated'] is True
@pytest.mark.asyncio
async def test_register_skill_rejects_workspace_escape(self):
from langbot.pkg.provider.tools.loaders.skill_authoring import (
@@ -500,7 +439,7 @@ class TestSkillToolLoader:
await loader.invoke_tool(
REGISTER_SKILL_TOOL_NAME,
{'path': '/workspace/../../etc'},
SimpleNamespace(),
_make_query(),
)
@pytest.mark.asyncio
@@ -520,7 +459,7 @@ class TestSkillToolLoader:
await loader.invoke_tool(
REGISTER_SKILL_TOOL_NAME,
{'path': '/workspace/foo'},
SimpleNamespace(),
_make_query(),
)
@pytest.mark.asyncio
@@ -531,7 +470,7 @@ class TestSkillToolLoader:
ap.skill_mgr = SimpleNamespace(skills={})
ap.box_service = SimpleNamespace(
available=True,
get_status=AsyncMock(return_value={'backend': {'available': False}}),
get_backend_status=AsyncMock(return_value={'backend': {'available': False}}),
)
loader = SkillToolLoader(ap)
@@ -546,10 +485,10 @@ class TestSkillToolLoader:
from langbot.pkg.provider.tools.loaders.skill_authoring import SkillToolLoader
ap = _make_ap()
ap.skill_mgr = SimpleNamespace(skills={'demo': _make_skill_data(name='demo')})
ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo')})
ap.box_service = SimpleNamespace(
available=True,
get_status=AsyncMock(return_value={'backend': {'available': True}}),
get_backend_status=AsyncMock(return_value={'backend': {'available': True}}),
)
loader = SkillToolLoader(ap)
@@ -561,47 +500,27 @@ class TestSkillToolLoader:
assert await loader.has_tool('activate') is True
assert await loader.has_tool('register_skill') is True
@pytest.mark.asyncio
async def test_tools_reappear_after_box_backend_recovers(self):
from langbot.pkg.provider.tools.loaders.skill_authoring import SkillToolLoader
ap = _make_ap()
ap.skill_mgr = SimpleNamespace(skills={'demo': _make_skill_data(name='demo')})
ap.box_service = SimpleNamespace(
available=False,
get_backend_status=AsyncMock(return_value={'backend': {'available': True}}),
)
loader = SkillToolLoader(ap)
await loader.initialize()
assert await loader.get_tools() == []
ap.box_service.available = True
assert sorted(tool.name for tool in await loader.get_tools()) == ['activate', 'register_skill']
class TestNativeToolLoaderSkillPaths:
@pytest.mark.asyncio
async def test_glob_skill_root_lists_visible_and_activated_mounts(self):
from langbot.pkg.provider.tools.loaders.native import NativeToolLoader
from langbot.pkg.provider.tools.loaders.skill import PIPELINE_BOUND_SKILLS_KEY, register_activated_skill
with tempfile.TemporaryDirectory() as tmpdir:
ap = _make_ap()
ap.box_service = SimpleNamespace(available=True, default_workspace=tmpdir)
ap.skill_mgr = SimpleNamespace(
skills={
'visible-skill': _make_skill_data(name='visible-skill', package_root=tmpdir),
'hidden-skill': _make_skill_data(name='hidden-skill', package_root=tmpdir),
}
)
loader = NativeToolLoader(ap)
query = SimpleNamespace(
query_id='q1',
variables={PIPELINE_BOUND_SKILLS_KEY: ['visible-skill']},
)
register_activated_skill(query, _make_skill_data(name='activated-skill', package_root=tmpdir))
result = await loader.invoke_tool(
'glob',
{'path': '/workspace/.skills', 'pattern': '*'},
query,
)
assert result == {
'ok': True,
'matches': [
'/workspace/.skills/activated-skill',
'/workspace/.skills/visible-skill',
],
'preview': '/workspace/.skills/activated-skill\n/workspace/.skills/visible-skill',
'total': 2,
'truncated': False,
'truncated_by': None,
}
@pytest.mark.asyncio
async def test_read_visible_skill_file(self):
from langbot.pkg.provider.tools.loaders.native import NativeToolLoader
@@ -613,20 +532,88 @@ class TestNativeToolLoaderSkillPaths:
f.write('demo instructions')
ap = _make_ap()
ap.box_service = SimpleNamespace(available=True, default_workspace=tmpdir)
ap.skill_mgr = SimpleNamespace(skills={'demo': _make_skill_data(name='demo', package_root=tmpdir)})
ap.box_service = SimpleNamespace(
available=True,
default_workspace=tmpdir,
shares_filesystem_with_box=True,
)
ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo', package_root=tmpdir)})
loader = NativeToolLoader(ap)
result = await loader.invoke_tool(
'read',
{'path': '/workspace/.skills/demo/SKILL.md'},
SimpleNamespace(query_id='q1', variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']}),
_make_query(query_id='q1', variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']}),
)
assert result['ok'] is True
assert result['content'] == 'demo instructions'
assert result['truncated'] is False
@pytest.mark.asyncio
async def test_external_runtime_read_never_interprets_package_root_on_core_host(self):
from langbot.pkg.provider.tools.loaders.native import NativeToolLoader
from langbot.pkg.provider.tools.loaders.skill import PIPELINE_BOUND_SKILLS_KEY
with tempfile.TemporaryDirectory() as tmpdir:
with open(os.path.join(tmpdir, 'SKILL.md'), 'w', encoding='utf-8') as file_obj:
file_obj.write('core-host-secret')
ap = _make_ap()
ap.box_service = SimpleNamespace(
available=True,
shares_filesystem_with_box=False,
read_skill_file=AsyncMock(return_value={'content': 'runtime-owned-content'}),
)
ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo', package_root=tmpdir)})
loader = NativeToolLoader(ap)
query = _make_query(
query_id='q-external-read',
variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']},
)
result = await loader.invoke_tool(
'read',
{'path': '/workspace/.skills/demo/SKILL.md'},
query,
)
assert result['ok'] is True
assert result['content'] == 'runtime-owned-content'
assert 'core-host-secret' not in repr(result)
ap.box_service.read_skill_file.assert_awaited_once_with(_CONTEXT, 'demo', 'SKILL.md')
@pytest.mark.asyncio
async def test_external_runtime_rejects_skill_host_fallback_without_protocol_capability(self):
from langbot.pkg.provider.tools.loaders.native import NativeToolLoader
from langbot.pkg.provider.tools.loaders.skill import PIPELINE_BOUND_SKILLS_KEY
with tempfile.TemporaryDirectory() as tmpdir:
with open(os.path.join(tmpdir, 'secret.txt'), 'w', encoding='utf-8') as file_obj:
file_obj.write('core-host-secret')
ap = _make_ap()
ap.box_service = SimpleNamespace(
available=True,
shares_filesystem_with_box=False,
)
ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo', package_root=tmpdir)})
loader = NativeToolLoader(ap)
query = _make_query(
query_id='q-external-no-protocol',
variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']},
)
with pytest.raises(ValueError, match='owned by the Box Runtime'):
await loader.invoke_tool(
'grep',
{
'path': '/workspace/.skills/demo',
'pattern': 'core-host-secret',
},
query,
)
@pytest.mark.asyncio
async def test_exec_in_activated_skill_mount_rewrites_command_and_refreshes(self):
from langbot.pkg.provider.tools.loaders.native import NativeToolLoader
@@ -642,7 +629,7 @@ class TestNativeToolLoaderSkillPaths:
ap.skill_mgr = SimpleNamespace(refresh_skill_from_disk=Mock())
loader = NativeToolLoader(ap)
query = SimpleNamespace(query_id='q1', launcher_type='person', launcher_id='123', variables={})
query = _make_query(query_id='q1', launcher_type='person', launcher_id='123')
register_activated_skill(query, _make_skill_data(name='demo', package_root=tmpdir))
result = await loader.invoke_tool(
@@ -658,7 +645,48 @@ class TestNativeToolLoaderSkillPaths:
tool_parameters = ap.box_service.execute_tool.await_args.args[0]
assert tool_parameters['command'] == 'python /workspace/.skills/demo/scripts/run.py'
assert tool_parameters['workdir'] == '/workspace/.skills/demo'
ap.skill_mgr.refresh_skill_from_disk.assert_called_once_with('demo')
assert ap.box_service.execute_tool.await_args.kwargs['skill_name'] == 'demo'
ap.skill_mgr.refresh_skill_from_disk.assert_called_once_with(_CONTEXT, 'demo')
@pytest.mark.asyncio
async def test_external_runtime_python_skill_uses_trusted_metadata_and_writable_env(self):
from langbot.pkg.provider.tools.loaders.native import NativeToolLoader
from langbot.pkg.provider.tools.loaders.skill import register_activated_skill
ap = _make_ap()
ap.box_service = SimpleNamespace(
available=True,
shares_filesystem_with_box=False,
execute_tool=AsyncMock(return_value={'ok': True}),
)
ap.skill_mgr = SimpleNamespace(refresh_skill_from_disk=Mock())
loader = NativeToolLoader(ap)
query = _make_query(query_id='q-external', launcher_type='person', launcher_id='123')
register_activated_skill(
query,
_make_skill_data(
name='demo',
package_root='/box-runtime/skills/tenants/workspace/demo',
python_project=True,
),
)
result = await loader.invoke_tool(
'exec',
{
'command': 'python /workspace/.skills/demo/scripts/run.py',
'workdir': '/workspace/.skills/demo',
},
query,
)
assert result['ok'] is True
tool_parameters = ap.box_service.execute_tool.await_args.args[0]
wrapped = tool_parameters['command']
assert '_LB_VENV_DIR="/workspace/.skill-envs/demo/.venv"' in wrapped
assert 'root = "/workspace/.skills/demo"' in wrapped
assert '/box-runtime/skills/tenants/workspace/demo' not in wrapped
assert ap.box_service.execute_tool.await_args.kwargs['skill_name'] == 'demo'
@pytest.mark.asyncio
async def test_write_requires_skill_activation(self):
@@ -668,10 +696,10 @@ class TestNativeToolLoaderSkillPaths:
with tempfile.TemporaryDirectory() as tmpdir:
ap = _make_ap()
ap.box_service = SimpleNamespace(available=True, default_workspace=tmpdir)
ap.skill_mgr = SimpleNamespace(skills={'demo': _make_skill_data(name='demo', package_root=tmpdir)})
ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo', package_root=tmpdir)})
loader = NativeToolLoader(ap)
query = SimpleNamespace(query_id='q1', variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']})
query = _make_query(query_id='q1', variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']})
with pytest.raises(ValueError, match='Skill "demo" is not available at this path'):
await loader.invoke_tool(
+27 -2
View File
@@ -17,6 +17,15 @@ import pytest
from langbot.pkg.provider.tools.errors import ToolExecutionDeniedError
from langbot.pkg.api.http.context import ExecutionContext
_CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
def get_toolmgr_module():
"""Lazy import to avoid circular import issues."""
@@ -191,6 +200,10 @@ class TestToolManagerExecuteFuncCall:
def sample_query(self):
"""Create sample query for testing."""
query = Mock(spec=pipeline_query.Query)
query.bot_uuid = None
query.pipeline_uuid = None
query.query_uuid = None
query._execution_context = _CONTEXT
return query
@pytest.mark.asyncio
@@ -271,6 +284,10 @@ class TestToolManagerExecuteFuncCall:
manager = toolmgr.ToolManager(mock_app)
self._wire_loaders(manager, mock_app, mock_plugin_loader, mock_mcp_loader)
query = Mock(variables={})
query._execution_context = _CONTEXT
query.bot_uuid = None
query.pipeline_uuid = None
query.query_uuid = None
result = await manager.execute_func_call(
'shared_tool',
@@ -282,7 +299,11 @@ class TestToolManagerExecuteFuncCall:
assert result == 'mcp_result'
mock_plugin_loader.has_tool.assert_not_awaited()
mock_plugin_loader.invoke_tool.assert_not_awaited()
mock_mcp_loader.has_tool.assert_awaited_once_with('shared_tool', source_id='bound-mcp')
mock_mcp_loader.has_tool.assert_awaited_once_with(
_CONTEXT,
'shared_tool',
source_id='bound-mcp',
)
mock_mcp_loader.invoke_tool.assert_awaited_once_with(
'shared_tool',
{'value': 1},
@@ -299,6 +320,10 @@ class TestToolManagerExecuteFuncCall:
manager = toolmgr.ToolManager(mock_app)
self._wire_loaders(manager, mock_app, mock_plugin_loader, mock_mcp_loader)
query = Mock(variables={})
query._execution_context = _CONTEXT
query.bot_uuid = None
query.pipeline_uuid = None
query.query_uuid = None
result = await manager.execute_func_call(
'shared_tool',
@@ -370,7 +395,7 @@ class TestToolManagerSourceResolution:
]
)
catalog = await manager.get_resolved_tool_catalog()
catalog = await manager.get_resolved_tool_catalog(_CONTEXT)
assert [item['name'] for item in catalog] == ['unique_tool']
app.logger.warning.assert_called_once()
@@ -1,17 +1,28 @@
from __future__ import annotations
import base64
import contextlib
import os
import tempfile
import threading
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock, patch
from unittest.mock import AsyncMock, Mock
import pytest
import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
from langbot.pkg.provider.tools.loaders.native import NativeToolLoader
from langbot.pkg.provider.tools.loaders import native as native_loader
from langbot.pkg.provider.tools.toolmgr import ToolManager
from langbot.pkg.api.http.context import ExecutionContext
_CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
class StubLoader:
@@ -44,7 +55,8 @@ class StubLoader:
for tool in self._tools
]
async def has_tool(self, name: str) -> bool:
async def has_tool(self, *args) -> bool:
name = args[-1]
return any(tool.name == name for tool in self._tools)
async def invoke_tool(self, name: str, parameters: dict, query):
@@ -65,31 +77,29 @@ def make_tool(name: str) -> resource_tool.LLMTool:
@pytest.mark.asyncio
async def test_tool_manager_includes_skill_tools_by_default():
"""Skill tools are exposed like native tools; the SkillToolLoader self-gates."""
async def test_tool_manager_omits_skill_authoring_tools_by_default():
manager = ToolManager(SimpleNamespace())
manager.native_tool_loader = StubLoader([make_tool('exec')])
manager.skill_tool_loader = StubLoader([make_tool('activate')])
manager.plugin_tool_loader = StubLoader([make_tool('plugin_tool')])
manager.mcp_tool_loader = StubLoader([make_tool('mcp_tool')])
tools = await manager.get_all_tools()
tools = await manager.get_all_tools(_CONTEXT)
assert [tool.name for tool in tools] == ['exec', 'activate', 'plugin_tool', 'mcp_tool']
assert [tool.name for tool in tools] == ['exec', 'plugin_tool', 'mcp_tool']
@pytest.mark.asyncio
async def test_tool_manager_omits_skill_tools_when_loader_unavailable():
"""When the SkillToolLoader gate is closed (no sandbox / skill_mgr) it returns no tools."""
async def test_tool_manager_includes_skill_authoring_tools_when_requested():
manager = ToolManager(SimpleNamespace())
manager.native_tool_loader = StubLoader([make_tool('exec')])
manager.skill_tool_loader = StubLoader([])
manager.skill_tool_loader = StubLoader([make_tool('activate')])
manager.plugin_tool_loader = StubLoader([make_tool('plugin_tool')])
manager.mcp_tool_loader = StubLoader([make_tool('mcp_tool')])
tools = await manager.get_all_tools()
tools = await manager.get_all_tools(_CONTEXT, include_skill_authoring=True)
assert [tool.name for tool in tools] == ['exec', 'plugin_tool', 'mcp_tool']
assert [tool.name for tool in tools] == ['exec', 'activate', 'plugin_tool', 'mcp_tool']
@pytest.mark.asyncio
@@ -104,7 +114,7 @@ async def test_tool_manager_catalog_labels_tool_sources():
)
manager.mcp_tool_loader = StubLoader([make_tool('mcp_tool')])
catalog = await manager.get_tool_catalog(include_skill_authoring=True)
catalog = await manager.get_tool_catalog(_CONTEXT, include_skill_authoring=True)
assert [(item['name'], item['source'], item['source_name']) for item in catalog] == [
('exec', 'builtin', 'LangBot'),
@@ -123,11 +133,55 @@ async def test_tool_manager_routes_native_tool_calls():
manager.plugin_tool_loader = StubLoader([make_tool('plugin_tool')])
manager.mcp_tool_loader = StubLoader([make_tool('mcp_tool')])
result = await manager.execute_func_call('exec', {'command': 'pwd'}, query=Mock())
query = SimpleNamespace(
_execution_context=_CONTEXT,
bot_uuid=None,
pipeline_uuid=None,
query_uuid=None,
)
result = await manager.execute_func_call('exec', {'command': 'pwd'}, query=query)
assert result == {'backend': 'fake'}
@pytest.mark.asyncio
async def test_tool_manager_hides_sandbox_and_skill_tools_without_workspace_entitlement():
box_service = SimpleNamespace(is_workspace_sandbox_available=AsyncMock(return_value=False))
manager = ToolManager(SimpleNamespace(box_service=box_service))
manager.native_tool_loader = StubLoader([make_tool('exec')])
manager.skill_tool_loader = StubLoader([make_tool('activate')])
manager.plugin_tool_loader = StubLoader([make_tool('plugin_tool')])
manager.mcp_tool_loader = StubLoader([make_tool('mcp_tool')])
tools = await manager.get_all_tools(_CONTEXT, include_skill_authoring=True)
catalog = await manager.get_tool_catalog(_CONTEXT, include_skill_authoring=True)
assert [tool.name for tool in tools] == ['plugin_tool', 'mcp_tool']
assert [item['name'] for item in catalog] == ['plugin_tool', 'mcp_tool']
assert box_service.is_workspace_sandbox_available.await_count == 2
@pytest.mark.asyncio
async def test_tool_manager_rechecks_workspace_entitlement_before_native_invocation():
box_service = SimpleNamespace(is_workspace_sandbox_available=AsyncMock(return_value=False))
manager = ToolManager(SimpleNamespace(box_service=box_service))
manager.native_tool_loader = StubLoader([make_tool('exec')], invoke_result={'unexpected': True})
manager.skill_tool_loader = StubLoader([])
manager.plugin_tool_loader = StubLoader([])
manager.mcp_tool_loader = StubLoader([])
query = SimpleNamespace(
_execution_context=_CONTEXT,
bot_uuid=None,
pipeline_uuid=None,
query_uuid=None,
)
with pytest.raises(Exception, match='exec'):
await manager.execute_func_call('exec', {'command': 'pwd'}, query=query)
box_service.is_workspace_sandbox_available.assert_awaited_once_with(_CONTEXT)
@pytest.mark.asyncio
async def test_native_tool_loader_hides_tools_when_box_unavailable():
loader = NativeToolLoader(SimpleNamespace(box_service=SimpleNamespace(available=False)))
@@ -141,7 +195,7 @@ async def test_native_tool_loader_hides_tools_when_box_unavailable():
async def test_native_tool_loader_exposes_all_tools_when_box_available():
box_service = SimpleNamespace(
available=True,
get_status=AsyncMock(return_value={'backend': {'available': True}}),
get_backend_status=AsyncMock(return_value={'backend': {'available': True}}),
)
loader = NativeToolLoader(SimpleNamespace(box_service=box_service, logger=Mock()))
await loader.initialize()
@@ -153,20 +207,66 @@ async def test_native_tool_loader_exposes_all_tools_when_box_available():
assert await loader.has_tool(tool_name) is True
@pytest.mark.asyncio
async def test_native_tool_loader_refreshes_after_box_recovers():
box_service = SimpleNamespace(
available=False,
get_backend_status=AsyncMock(return_value={'backend': {'available': True}}),
)
loader = NativeToolLoader(SimpleNamespace(box_service=box_service, logger=Mock()))
await loader.initialize()
assert await loader.get_tools() == []
box_service.available = True
assert [tool.name for tool in await loader.get_tools()] == ['exec', 'read', 'write', 'edit', 'glob', 'grep']
@pytest.mark.asyncio
async def test_native_tool_loader_rechecks_admission_at_the_final_invoke_boundary():
box_service = SimpleNamespace(
available=True,
require_workspace_sandbox=AsyncMock(side_effect=RuntimeError('entitlement expired')),
)
loader = NativeToolLoader(SimpleNamespace(box_service=box_service, logger=Mock()))
query = SimpleNamespace(
_execution_context=_CONTEXT,
bot_uuid=None,
pipeline_uuid=None,
query_uuid=None,
)
with pytest.raises(RuntimeError, match='entitlement expired'):
await loader.invoke_tool('read', {'path': '/workspace/private.txt'}, query)
box_service.require_workspace_sandbox.assert_awaited_once_with(_CONTEXT)
# ── read/write/edit file tool tests ─────────────────────────────
def _make_loader_with_workspace(tmpdir: str) -> tuple[NativeToolLoader, Mock]:
logger = Mock()
box_service = SimpleNamespace(available=True, default_workspace=tmpdir)
box_service = SimpleNamespace(
available=True,
default_workspace=tmpdir,
_tenant_workspace=Mock(return_value=tmpdir),
)
ap = SimpleNamespace(box_service=box_service, logger=logger)
return NativeToolLoader(ap), logger
def _make_query() -> Mock:
q = Mock()
q.query_id = 'test-query-1'
return q
def _make_query() -> SimpleNamespace:
return SimpleNamespace(
query_id='test-query-1',
query_uuid='test-query-1',
instance_uuid=_CONTEXT.instance_uuid,
workspace_uuid=_CONTEXT.workspace_uuid,
placement_generation=_CONTEXT.placement_generation,
bot_uuid=None,
pipeline_uuid=None,
variables={},
)
@pytest.mark.asyncio
@@ -236,32 +336,6 @@ async def test_write_creates_subdirectories():
assert f.read() == 'nested'
@pytest.mark.asyncio
async def test_write_falls_back_to_box_for_container_owned_workspace_file():
with tempfile.TemporaryDirectory() as tmpdir:
box_service = SimpleNamespace(
available=True,
default_workspace=tmpdir,
execute_tool=AsyncMock(
return_value={
'ok': True,
'stdout': '{"ok": true, "path": "/workspace/root-owned.txt"}',
}
),
)
loader = NativeToolLoader(SimpleNamespace(box_service=box_service, logger=Mock()))
loader._write_host_file = Mock(side_effect=PermissionError)
result = await loader.invoke_tool(
'write',
{'path': '/workspace/root-owned.txt', 'content': 'updated'},
_make_query(),
)
assert result == {'ok': True, 'path': '/workspace/root-owned.txt'}
box_service.execute_tool.assert_awaited_once()
@pytest.mark.asyncio
async def test_read_binary_file_as_base64_chunk():
with tempfile.TemporaryDirectory() as tmpdir:
@@ -352,35 +426,6 @@ async def test_edit_replaces_unique_string():
assert f.read() == 'def foo():\n return 42\n'
@pytest.mark.asyncio
async def test_edit_falls_back_to_box_for_container_owned_workspace_file():
with tempfile.TemporaryDirectory() as tmpdir:
path = os.path.join(tmpdir, 'root-owned.txt')
with open(path, 'w') as f:
f.write('before')
box_service = SimpleNamespace(
available=True,
default_workspace=tmpdir,
execute_tool=AsyncMock(
return_value={
'ok': True,
'stdout': '{"ok": true, "path": "/workspace/root-owned.txt"}',
}
),
)
loader = NativeToolLoader(SimpleNamespace(box_service=box_service, logger=Mock()))
with patch('builtins.open', side_effect=PermissionError):
result = await loader.invoke_tool(
'edit',
{'path': '/workspace/root-owned.txt', 'old_string': 'before', 'new_string': 'after'},
_make_query(),
)
assert result == {'ok': True, 'path': '/workspace/root-owned.txt'}
box_service.execute_tool.assert_awaited_once()
@pytest.mark.asyncio
async def test_edit_rejects_ambiguous_match():
with tempfile.TemporaryDirectory() as tmpdir:
@@ -415,6 +460,24 @@ async def test_edit_rejects_missing_string():
assert 'not found' in result['error'].lower()
@pytest.mark.asyncio
async def test_edit_rejects_oversized_host_file(monkeypatch):
with tempfile.TemporaryDirectory() as tmpdir:
loader, _ = _make_loader_with_workspace(tmpdir)
with open(os.path.join(tmpdir, 'large.txt'), 'wb') as f:
f.write(b'12345')
monkeypatch.setattr(native_loader, '_MAX_HOST_EDIT_FILE_BYTES', 4)
result = await loader.invoke_tool(
'edit',
{'path': '/workspace/large.txt', 'old_string': '1', 'new_string': 'x'},
_make_query(),
)
assert result['ok'] is False
assert 'edit limit' in result['error']
@pytest.mark.asyncio
async def test_path_escape_blocked():
with tempfile.TemporaryDirectory() as tmpdir:
@@ -424,6 +487,137 @@ async def test_path_escape_blocked():
await loader.invoke_tool('read', {'path': '/workspace/../../etc/passwd'}, _make_query())
@pytest.mark.parametrize(
('tool_name', 'parameters'),
[
('read', {'path': '/workspace/shared/tenant-b-only.txt'}),
(
'write',
{'path': '/workspace/shared/tenant-b-only.txt', 'content': 'overwritten by tenant a'},
),
(
'edit',
{
'path': '/workspace/shared/tenant-b-only.txt',
'old_string': 'tenant-b-secret',
'new_string': 'overwritten by tenant a',
},
),
('glob', {'path': '/workspace/shared', 'pattern': '*'}),
('grep', {'path': '/workspace/shared', 'pattern': 'tenant-b-secret'}),
],
)
@pytest.mark.asyncio
async def test_host_workspace_operations_do_not_follow_a_swapped_ancestor(
monkeypatch,
tool_name: str,
parameters: dict,
):
with tempfile.TemporaryDirectory() as tmpdir:
tenant_a = os.path.join(tmpdir, 'tenant-a')
tenant_b = os.path.join(tmpdir, 'tenant-b')
os.makedirs(os.path.join(tenant_a, 'shared'))
os.makedirs(os.path.join(tenant_b, 'shared'))
tenant_b_file = os.path.join(tenant_b, 'shared', 'tenant-b-only.txt')
with open(os.path.join(tenant_a, 'shared', 'tenant-a-only.txt'), 'w', encoding='utf-8') as file_obj:
file_obj.write('tenant-a-content')
with open(tenant_b_file, 'w', encoding='utf-8') as file_obj:
file_obj.write('tenant-b-secret')
box_service = SimpleNamespace(
available=True,
default_workspace=tmpdir,
_tenant_workspace=Mock(return_value=tenant_a),
)
loader = NativeToolLoader(SimpleNamespace(box_service=box_service, logger=Mock()))
original_open_host_root = native_loader._open_host_root
@contextlib.contextmanager
def open_host_root_after_swap(location, *, create):
with original_open_host_root(location, create=create) as root_fd:
original_ancestor = os.path.join(tenant_a, 'shared-original')
os.rename(os.path.join(tenant_a, 'shared'), original_ancestor)
os.symlink(os.path.join(tenant_b, 'shared'), os.path.join(tenant_a, 'shared'))
yield root_fd
monkeypatch.setattr(native_loader, '_open_host_root', open_host_root_after_swap)
try:
result = await loader.invoke_tool(tool_name, parameters, _make_query())
except ValueError as exc:
result = {'ok': False, 'error': str(exc)}
assert result.get('ok') is False
assert 'tenant-b-secret' not in repr(result)
assert 'tenant-b-only.txt' not in repr(result)
with open(tenant_b_file, encoding='utf-8') as file_obj:
assert file_obj.read() == 'tenant-b-secret'
@pytest.mark.asyncio
async def test_host_file_api_falls_back_to_tenant_box_when_openat_is_unavailable(monkeypatch):
with tempfile.TemporaryDirectory() as tmpdir:
box_service = SimpleNamespace(
available=True,
default_workspace=tmpdir,
_tenant_workspace=Mock(return_value=tmpdir),
execute_tool=AsyncMock(
return_value={
'ok': True,
'stdout': '{"ok": true, "content": "box-owned", "truncated": false}',
'stderr': '',
}
),
)
loader = NativeToolLoader(SimpleNamespace(box_service=box_service, logger=Mock()))
monkeypatch.setattr(native_loader, '_SECURE_HOST_FILE_OPS_AVAILABLE', False)
result = await loader.invoke_tool(
'read',
{'path': '/workspace/file.txt'},
_make_query(),
)
assert result['ok'] is True
assert result['content'] == 'box-owned'
command = box_service.execute_tool.await_args.args[0]['command']
assert 'path = "/workspace/file.txt"' in command
@pytest.mark.asyncio
async def test_box_workspace_edit_script_bounds_file_read_and_replacement(monkeypatch):
with tempfile.TemporaryDirectory() as tmpdir:
box_service = SimpleNamespace(
available=True,
default_workspace=tmpdir,
_tenant_workspace=Mock(return_value=tmpdir),
execute_tool=AsyncMock(
return_value={
'ok': True,
'stdout': '{"ok": false, "error": "File exceeds limit"}',
'stderr': '',
}
),
)
loader = NativeToolLoader(SimpleNamespace(box_service=box_service, logger=Mock()))
monkeypatch.setattr(native_loader, '_SECURE_HOST_FILE_OPS_AVAILABLE', False)
await loader.invoke_tool(
'edit',
{
'path': '/workspace/file.txt',
'old_string': 'old',
'new_string': 'new',
},
_make_query(),
)
command = box_service.execute_tool.await_args.args[0]['command']
assert f'os.path.getsize(path) > {native_loader._MAX_HOST_EDIT_FILE_BYTES}' in command
assert f'f.read({native_loader._MAX_HOST_EDIT_FILE_BYTES + 1})' in command
assert f"len(new_content.encode('utf-8')) > {native_loader._MAX_HOST_EDIT_FILE_BYTES}" in command
@pytest.mark.asyncio
async def test_box_availability_helper_handles_unavailable_and_errors():
from langbot.pkg.provider.tools.loaders.availability import is_box_backend_available
@@ -433,13 +627,13 @@ async def test_box_availability_helper_handles_unavailable_and_errors():
unavailable_backend = SimpleNamespace(
available=True,
get_status=AsyncMock(return_value={'backend': {'available': False}}),
get_backend_status=AsyncMock(return_value={'backend': {'available': False}}),
)
assert await is_box_backend_available(SimpleNamespace(box_service=unavailable_backend)) is False
failing_backend = SimpleNamespace(
available=True,
get_status=AsyncMock(side_effect=RuntimeError('box unavailable')),
get_backend_status=AsyncMock(side_effect=RuntimeError('box unavailable')),
)
assert await is_box_backend_available(SimpleNamespace(box_service=failing_backend)) is False
@@ -537,6 +731,34 @@ async def test_glob_caps_match_count_and_returns_preview():
assert result['truncated_by'] == 'matches'
@pytest.mark.asyncio
async def test_glob_runs_off_event_loop_and_caps_directory_walk(monkeypatch):
monkeypatch.setattr(native_loader, '_FILE_WALK_MAX_ENTRIES', 10)
with tempfile.TemporaryDirectory() as tmpdir:
loader, _ = _make_loader_with_workspace(tmpdir)
event_loop_thread = threading.get_ident()
observed_threads: list[int] = []
original = loader._glob_host_location
def observe(*args, **kwargs):
observed_threads.append(threading.get_ident())
return original(*args, **kwargs)
monkeypatch.setattr(loader, '_glob_host_location', observe)
for index in range(12):
with open(os.path.join(tmpdir, f'file-{index:03d}.txt'), 'w', encoding='utf-8') as f:
f.write(str(index))
result = await loader.invoke_tool('glob', {'path': '/workspace', 'pattern': '*.txt'}, _make_query())
assert result['ok'] is True
assert result['total'] == 10
assert result['truncated'] is True
assert result['truncated_by'] == 'scan'
assert observed_threads and observed_threads[0] != event_loop_thread
@pytest.mark.asyncio
async def test_grep_reports_invalid_regex_and_truncates_long_matching_lines():
with tempfile.TemporaryDirectory() as tmpdir:
@@ -554,3 +776,21 @@ async def test_grep_reports_invalid_regex_and_truncates_long_matching_lines():
assert result['truncated_by'] == 'line'
assert result['matches'][0]['file'] == '/workspace/data.txt'
assert result['matches'][0]['content'].endswith('... [truncated]')
@pytest.mark.asyncio
async def test_grep_interrupts_catastrophic_regex(monkeypatch):
monkeypatch.setattr(native_loader, '_GREP_REGEX_TIMEOUT_SECONDS', 0.001)
with tempfile.TemporaryDirectory() as tmpdir:
loader, _ = _make_loader_with_workspace(tmpdir)
with open(os.path.join(tmpdir, 'data.txt'), 'w', encoding='utf-8') as f:
f.write(('a' * 100_000) + '!')
result = await loader.invoke_tool(
'grep',
{'path': '/workspace', 'pattern': r'(a+)+$'},
_make_query(),
)
assert result == {'ok': False, 'error': 'Regex search timed out'}
@@ -5,6 +5,7 @@ from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.provider.tools.loaders.mcp import (
MCP_TOOL_LIST_RESOURCES,
MCP_TOOL_READ_RESOURCE,
@@ -15,6 +16,41 @@ from langbot.pkg.provider.tools.loaders.plugin import PluginToolLoader
from langbot_plugin.api.entities.builtin.resource.tool import LLMTool
_CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
query_uuid='query-a',
)
def _mcp_app() -> SimpleNamespace:
return SimpleNamespace(
logger=Mock(),
workspace_service=SimpleNamespace(
get_execution_binding=AsyncMock(
return_value=SimpleNamespace(
instance_uuid=_CONTEXT.instance_uuid,
workspace_uuid=_CONTEXT.workspace_uuid,
placement_generation=_CONTEXT.placement_generation,
)
)
),
)
def _query() -> SimpleNamespace:
return SimpleNamespace(
instance_uuid=_CONTEXT.instance_uuid,
workspace_uuid=_CONTEXT.workspace_uuid,
placement_generation=_CONTEXT.placement_generation,
query_uuid=_CONTEXT.query_uuid,
bot_uuid=None,
pipeline_uuid=None,
variables={},
)
def make_tool(name: str) -> LLMTool:
return LLMTool(
name=name,
@@ -38,9 +74,10 @@ async def test_two_mcp_servers_with_same_tool_route_to_authorized_server():
get_tools=Mock(return_value=[tool]),
invoke_mcp_tool=AsyncMock(return_value='from-b'),
)
loader = MCPLoader(SimpleNamespace(logger=Mock()))
loader.sessions = {'first': first, 'second': second}
query = SimpleNamespace(variables={})
loader = MCPLoader(_mcp_app())
loader._register_session(_CONTEXT, 'first', first)
loader._register_session(_CONTEXT, 'second', second)
query = _query()
result = await loader.invoke_tool(
'shared_tool',
@@ -67,12 +104,12 @@ async def test_reserved_name_routes_to_real_mcp_tool_when_source_id_is_present(r
get_tools=Mock(return_value=[real_tool]),
invoke_mcp_tool=AsyncMock(return_value='real-result'),
)
loader = MCPLoader(SimpleNamespace(logger=Mock()))
loader.sessions = {'real': session}
query = SimpleNamespace(variables={})
loader = MCPLoader(_mcp_app())
loader._register_session(_CONTEXT, 'real', session)
query = _query()
assert await loader.has_tool(reserved_name, source_id='srv-real') is True
assert await loader.get_tool(reserved_name, source_id='srv-real') is real_tool
assert await loader.has_tool(_CONTEXT, reserved_name, source_id='srv-real') is True
assert await loader.get_tool(_CONTEXT, reserved_name, source_id='srv-real') is real_tool
result = await loader.invoke_tool(reserved_name, {}, query, source_id='srv-real')
assert result == 'real-result'
@@ -102,14 +139,14 @@ async def test_reserved_name_uses_host_synthetic_tool_when_source_id_is_none(
has_resource_support=Mock(return_value=True),
invoke_mcp_tool=AsyncMock(return_value='real-result'),
)
loader = MCPLoader(SimpleNamespace(logger=Mock()))
loader.sessions = {'real': session}
loader = MCPLoader(_mcp_app())
loader._register_session(_CONTEXT, 'real', session)
synthetic_invoke = AsyncMock(return_value='synthetic-result')
setattr(loader, invoke_method, synthetic_invoke)
query = SimpleNamespace(variables={})
query = _query()
assert await loader.has_tool(reserved_name, source_id=None) is True
tool = await loader.get_tool(reserved_name, source_id=None)
assert await loader.has_tool(_CONTEXT, reserved_name, source_id=None) is True
tool = await loader.get_tool(_CONTEXT, reserved_name, source_id=None)
result = await loader.invoke_tool(reserved_name, {}, query, source_id=None)
assert tool is not None
@@ -137,7 +174,7 @@ async def test_two_plugins_with_same_tool_forward_only_authorized_plugin():
call_tool=AsyncMock(return_value='authorized-result'),
)
loader = PluginToolLoader(SimpleNamespace(plugin_connector=connector, logger=Mock()))
query = SimpleNamespace(session=SimpleNamespace(), query_id=7)
query = SimpleNamespace(session=SimpleNamespace(), query_id=7, query_uuid='query-a')
query.session.model_dump = Mock(return_value={})
result = await loader.invoke_tool(
@@ -154,4 +191,5 @@ async def test_two_plugins_with_same_tool_forward_only_authorized_plugin():
session=query.session,
query_id=7,
bound_plugins=['authorized/plugin'],
query_uuid='query-a',
)