mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-08-09 04:40:57 +00:00
chore(merge): sync master into dev/4.11.x
This commit is contained in:
@@ -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):
|
||||
|
||||
@@ -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 == {}
|
||||
@@ -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')
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -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',
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user