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

This commit is contained in:
huanghuoguoguo
2026-07-31 19:29:38 +08:00
502 changed files with 77975 additions and 12729 deletions
+41
View File
@@ -3,6 +3,47 @@
from __future__ import annotations
import typing
from types import SimpleNamespace
from unittest.mock import AsyncMock
from langbot_plugin.entities.io.context import InstallationBinding
TEST_RUNTIME_BINDING = InstallationBinding(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
installation_uuid='00000000-0000-4000-8000-000000000001',
runtime_revision=1,
artifact_digest='a' * 64,
)
def bind_runtime_action_context(
handler,
application,
*,
plugin_identity: str = 'test/runner',
):
"""Simulate the trusted Runtime envelope used around direct action calls."""
application.workspace_service = SimpleNamespace(
get_execution_binding=AsyncMock(
return_value=SimpleNamespace(
instance_uuid=TEST_RUNTIME_BINDING.instance_uuid,
workspace_uuid=TEST_RUNTIME_BINDING.workspace_uuid,
placement_generation=TEST_RUNTIME_BINDING.placement_generation,
)
)
)
plugin_author, plugin_name = plugin_identity.split('/', 1)
handler.register_installation_binding(
TEST_RUNTIME_BINDING,
plugin_author=plugin_author,
plugin_name=plugin_name,
)
handler._current_action_context.set(TEST_RUNTIME_BINDING)
return handler
def make_resources(
@@ -350,6 +350,7 @@ def mock_query():
"""Create a mock query for testing."""
query = Mock()
query.query_id = 123
query.workspace_uuid = 'workspace-test'
query.bot_uuid = 'bot-uuid-123'
query.pipeline_uuid = 'pipeline-uuid-456'
query.launcher_type = Mock(value='person')
@@ -398,6 +399,7 @@ def mock_query_no_session():
"""Create a mock Query without session."""
query = Mock()
query.query_id = 456
query.workspace_uuid = 'workspace-test'
query.bot_uuid = 'bot-uuid-456'
query.pipeline_uuid = 'pipeline-uuid-789'
query.launcher_type = Mock(value='person')
+15 -4
View File
@@ -19,6 +19,7 @@ from unittest.mock import AsyncMock, MagicMock
from langbot.pkg.agent.runner.descriptor import AgentRunnerDescriptor
from langbot.pkg.agent.runner.session_registry import AgentRunSessionRegistry
from langbot.pkg.plugin.handler import _get_pipeline_knowledge_base_uuids
from langbot.pkg.api.http.context import ExecutionContext
# Import shared test fixtures from conftest.py
from .conftest import make_resources, make_session
@@ -69,6 +70,16 @@ class MockQuery:
self.session.launcher_id = 'group_123'
self.sender_id = 'user_001'
self.bot_uuid = 'bot_001'
self.pipeline_uuid = 'pipeline-001'
self.query_uuid = f'query-{query_id}'
self._execution_context = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
bot_uuid=self.bot_uuid,
pipeline_uuid=self.pipeline_uuid,
query_uuid=self.query_uuid,
)
self.pipeline_config = {
'ai': {
'runner': {
@@ -112,7 +123,7 @@ class MockApplication:
class FakeAgentRunnerRegistry:
async def get(self, runner_id, bound_plugins=None):
async def get(self, context, runner_id, bound_plugins=None):
return AgentRunnerDescriptor(
id=runner_id,
source='plugin',
@@ -306,19 +317,19 @@ async def test_tool_manager_get_tool_detail_returns_uniform_schema():
mgr = ToolManager.__new__(ToolManager)
async def fake_get_tool_by_name(name):
async def fake_get_tool_by_name(context, name):
return tool if name == 'search' else None
mgr.get_tool_by_name = fake_get_tool_by_name
detail = await mgr.get_tool_detail('search')
detail = await mgr.get_tool_detail('workspace-test', 'search')
assert detail == {
'name': 'search',
'description': 'Search test data',
'human_desc': 'Search public data',
'parameters': {'type': 'object', 'properties': {'q': {'type': 'string'}}},
}
assert await mgr.get_tool_detail('missing') is None
assert await mgr.get_tool_detail('workspace-test', 'missing') is None
class TestCallToolAuthorization:
@@ -17,7 +17,7 @@ from langbot_plugin.api.entities.builtin.agent_runner.page_results import (
)
from langbot_plugin.entities.io.actions.enums import PluginToRuntimeAction
from .conftest import make_resources
from .conftest import bind_runtime_action_context, make_resources
class FakeConnection:
@@ -56,7 +56,10 @@ def _handler(db_engine, session_registry):
return True
fake_app = FakeApplication(db_engine)
return RuntimeConnectionHandler(FakeConnection(), fake_disconnect, fake_app)
return bind_runtime_action_context(
RuntimeConnectionHandler(FakeConnection(), fake_disconnect, fake_app),
fake_app,
)
async def _register_session(
@@ -19,6 +19,7 @@ from langbot.pkg.agent.runner.session_registry import get_session_registry
from langbot.pkg.agent.runner.run_ledger_store import RunLedgerStore
from langbot.pkg.agent.runner.interaction_store import InteractionStore
from langbot.pkg.agent.runner.persistent_state_store import reset_persistent_state_store
from langbot.pkg.api.http.context import ExecutionContext
from langbot_plugin.api.entities.builtin.platform import entities as platform_entities
from langbot_plugin.api.entities.builtin.platform import events as platform_events
from langbot_plugin.api.entities.builtin.platform import message as platform_message
@@ -28,6 +29,11 @@ from langbot_plugin.api.entities.builtin.resource import tool as resource_tool
RUNNER_ID = 'plugin:langbot-team/LocalAgent/default'
TEST_CONTEXT = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
)
class FakeLogger:
@@ -103,8 +109,15 @@ class FakeRegistry:
self.descriptor = descriptor
self.calls: list[dict] = []
async def get(self, runner_id, bound_plugins=None):
self.calls.append({'runner_id': runner_id, 'bound_plugins': bound_plugins})
async def get(self, context, runner_id, bound_plugins=None):
self.calls.append(
{
'context': context,
'runner_id': runner_id,
'bound_plugins': bound_plugins,
}
)
assert context.workspace_uuid == TEST_CONTEXT.workspace_uuid
assert runner_id == self.descriptor.id
return self.descriptor
@@ -129,7 +142,7 @@ class FakeApplication:
get_knowledge_base_by_uuid=AsyncMock(return_value=FakeKnowledgeBase('kb_001'))
)
self.skill_mgr = types.SimpleNamespace(
skills={
get_skills=lambda context: {
'demo': {
'name': 'demo',
'display_name': 'Demo Skill',
@@ -201,7 +214,7 @@ def make_query():
using_conversation=FakeConversation(),
)
return types.SimpleNamespace(
query = types.SimpleNamespace(
query_id=1001,
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id='user_001',
@@ -254,6 +267,8 @@ def make_query():
)
],
)
query._execution_context = TEST_CONTEXT
return query
def test_context_builder_includes_consumable_base64_attachments():
@@ -903,6 +918,7 @@ class TestQueryEntrySessionQueryId:
def __init__(self):
self.resolver = object.__new__(BoxService)
self.resolver._cloud_managed = False
self.materialize_session_id = None
async def materialize_inbound_attachments(self, query):
@@ -1043,7 +1059,14 @@ class TestQueryEntrySessionQueryId:
enabled=True,
)
messages = [message async for message in orchestrator.run(event, binding)]
messages = [
message
async for message in orchestrator.run(
event,
binding,
adapter_context={'_execution_context': TEST_CONTEXT},
)
]
assert len(messages) == 1
# Verify session during run has query_id=None
+34 -13
View File
@@ -7,6 +7,14 @@ import pytest
from langbot.pkg.agent.runner.registry import AgentRunnerRegistry
from langbot.pkg.agent.runner.descriptor import AgentRunnerDescriptor
from langbot.pkg.agent.runner.errors import RunnerNotFoundError, RunnerNotAuthorizedError
from langbot.pkg.api.http.context import ExecutionContext
TEST_CONTEXT = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
)
class FakeApplication:
@@ -31,6 +39,9 @@ class FakeApplication:
class FakePluginConnector:
is_enable_plugin = True
async def require_workspace_context(self, context):
return context
async def list_agent_runners(self, bound_plugins=None):
# Return sample runner data
return [
@@ -96,7 +107,7 @@ class TestRegistryDiscovery:
ap = FakeApplication()
registry = AgentRunnerRegistry(ap)
runners = await registry.list_runners(use_cache=False)
runners = await registry.list_runners(TEST_CONTEXT, use_cache=False)
# Should find 2 valid runners (langbot-team/LocalAgent and alice/my-agent)
assert len(runners) == 2
@@ -112,10 +123,10 @@ class TestRegistryDiscovery:
registry = AgentRunnerRegistry(ap)
# First discovery
runners1 = await registry.list_runners(use_cache=True)
runners1 = await registry.list_runners(TEST_CONTEXT, use_cache=True)
# Second call should use cache
runners2 = await registry.list_runners(use_cache=True)
runners2 = await registry.list_runners(TEST_CONTEXT, use_cache=True)
assert registry._cache is not None
assert len(runners1) == len(runners2)
@@ -127,7 +138,7 @@ class TestRegistryDiscovery:
ap.plugin_connector.is_enable_plugin = False
registry = AgentRunnerRegistry(ap)
runners = await registry.list_runners(use_cache=False)
runners = await registry.list_runners(TEST_CONTEXT, use_cache=False)
assert runners == []
@@ -143,23 +154,28 @@ class TestRegistryDiscovery:
# First: get with bound_plugins filter (should not pollute cache)
descriptor = await registry.get(
TEST_CONTEXT,
'plugin:langbot-team/LocalAgent/default',
bound_plugins=['langbot-team/LocalAgent'],
)
assert descriptor.id == 'plugin:langbot-team/LocalAgent/default'
# Cache should contain ALL runners (both langbot and alice)
assert registry._cache is not None
assert len(registry._cache) == 2 # Both runners in cache
assert 'plugin:langbot-team/LocalAgent/default' in registry._cache
assert 'plugin:alice/my-agent/custom' in registry._cache
scoped_cache = registry._cache[('instance-test', 'workspace-test', 1)]
assert len(scoped_cache) == 2
assert 'plugin:langbot-team/LocalAgent/default' in scoped_cache
assert 'plugin:alice/my-agent/custom' in scoped_cache
# Second: list_runners without filter should return ALL runners
all_runners = await registry.list_runners(bound_plugins=None, use_cache=True)
all_runners = await registry.list_runners(TEST_CONTEXT, bound_plugins=None, use_cache=True)
assert len(all_runners) == 2 # Both runners returned
# Third: list_runners with different filter should work correctly
alice_runners = await registry.list_runners(bound_plugins=['alice/my-agent'], use_cache=True)
alice_runners = await registry.list_runners(
TEST_CONTEXT,
bound_plugins=['alice/my-agent'],
use_cache=True,
)
assert len(alice_runners) == 1
assert alice_runners[0].id == 'plugin:alice/my-agent/custom'
@@ -173,7 +189,10 @@ class TestRegistryGet:
ap = FakeApplication()
registry = AgentRunnerRegistry(ap)
descriptor = await registry.get('plugin:langbot-team/LocalAgent/default')
descriptor = await registry.get(
TEST_CONTEXT,
'plugin:langbot-team/LocalAgent/default',
)
assert descriptor.id == 'plugin:langbot-team/LocalAgent/default'
assert descriptor.plugin_author == 'langbot-team'
@@ -187,7 +206,7 @@ class TestRegistryGet:
registry = AgentRunnerRegistry(ap)
with pytest.raises(RunnerNotFoundError) as exc_info:
await registry.get('plugin:notexist/unknown/default')
await registry.get(TEST_CONTEXT, 'plugin:notexist/unknown/default')
assert exc_info.value.runner_id == 'plugin:notexist/unknown/default'
@@ -199,6 +218,7 @@ class TestRegistryGet:
# Authorized - langbot plugin in bound list
descriptor = await registry.get(
TEST_CONTEXT,
'plugin:langbot-team/LocalAgent/default',
bound_plugins=['langbot-team/LocalAgent'],
)
@@ -207,6 +227,7 @@ class TestRegistryGet:
# Not authorized - plugin not in bound list
with pytest.raises(RunnerNotAuthorizedError):
await registry.get(
TEST_CONTEXT,
'plugin:alice/my-agent/custom',
bound_plugins=['langbot-team/LocalAgent'],
)
@@ -221,7 +242,7 @@ class TestRegistryMetadataForPipeline:
ap = FakeApplication()
registry = AgentRunnerRegistry(ap)
options, stages = await registry.get_runner_metadata_for_pipeline()
options, stages = await registry.get_runner_metadata_for_pipeline(TEST_CONTEXT)
# Should have options for each runner
assert len(options) == 2
@@ -12,9 +12,15 @@ from langbot.pkg.agent.runner.binding_resolver import AgentBindingResolver
from langbot.pkg.agent.runner.query_entry_adapter import QueryEntryAdapter
from langbot.pkg.agent.runner.resource_builder import AgentResourceBuilder
from langbot.pkg.agent.runner.host_models import AgentBinding, BindingScope, ResourcePolicy
from langbot.pkg.api.http.context import ExecutionContext
RUNNER_ID = 'plugin:test/runner/default'
TEST_CONTEXT = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
)
FULL_PERMISSIONS = {
'models': ['count_tokens', 'invoke', 'stream', 'rerank'],
'tools': ['detail', 'call'],
@@ -86,6 +92,7 @@ async def build_resources(app, query, descriptor):
agent_config = QueryEntryAdapter.config_to_agent_config(query, descriptor.id)
binding = AgentBindingResolver().resolve_one(event, [agent_config])
return await AgentResourceBuilder(app).build_resources_from_binding(
execution_context=TEST_CONTEXT,
event=event,
binding=binding,
descriptor=descriptor,
@@ -118,10 +125,12 @@ async def test_build_models_authorizes_config_declared_llm_and_rerank_models(app
'rerank': make_model(model_type='rerank', provider='rerank-provider'),
}
async def get_model_by_uuid(model_uuid):
async def get_model_by_uuid(context, model_uuid):
assert context == TEST_CONTEXT
return llm_models.get(model_uuid)
async def get_rerank_model_by_uuid(model_uuid):
async def get_rerank_model_by_uuid(context, model_uuid):
assert context == TEST_CONTEXT
return rerank_models.get(model_uuid)
app.model_mgr.get_model_by_uuid = AsyncMock(side_effect=get_model_by_uuid)
@@ -228,7 +237,8 @@ async def test_build_resources_accepts_dynamic_form_type_aliases(app):
"""Frontend DynamicForm aliases should resolve to runtime resource grants."""
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=make_model())
async def get_kb(kb_uuid):
async def get_kb(context, kb_uuid):
assert context == TEST_CONTEXT
return SimpleNamespace(
uuid=kb_uuid,
get_name=lambda: f'name-{kb_uuid}',
@@ -322,7 +332,7 @@ async def test_build_tools_authorizes_query_declared_tools(app):
"""Tools discovered by Pipeline preprocessing become run-scoped authorized
resources, with full parameters schema prefilled by the host."""
app.tool_mgr.get_tool_schema = AsyncMock(
side_effect=lambda name, source_ref=None: {
side_effect=lambda context, name, source_ref=None: {
'qa_plugin_echo': (
'Echo test tool',
{'type': 'object', 'properties': {'text': {'type': 'string'}}},
@@ -411,6 +421,7 @@ async def test_build_tools_materializes_independent_agent_all_tools_policy(app):
)
resources = await AgentResourceBuilder(app).build_resources_from_binding(
execution_context=TEST_CONTEXT,
event=QueryEntryAdapter.query_to_event(make_query({})),
binding=binding,
descriptor=descriptor,
@@ -418,6 +429,7 @@ async def test_build_tools_materializes_independent_agent_all_tools_policy(app):
assert [tool['tool_name'] for tool in resources['tools']] == ['exec', 'plugin_tool']
app.tool_mgr.get_resolved_tool_catalog.assert_awaited_once_with(
TEST_CONTEXT,
include_skill_authoring=True,
include_mcp_resource_tools=True,
)
@@ -442,6 +454,7 @@ async def test_build_tools_denies_mcp_resource_tools_when_agent_reads_disabled(a
)
resources = await AgentResourceBuilder(app).build_resources_from_binding(
execution_context=TEST_CONTEXT,
event=QueryEntryAdapter.query_to_event(make_query({})),
binding=binding,
descriptor=descriptor,
@@ -471,6 +484,7 @@ async def test_build_tools_keeps_plugin_using_synthetic_mcp_tool_name_when_reads
)
resources = await AgentResourceBuilder(app).build_resources_from_binding(
execution_context=TEST_CONTEXT,
event=QueryEntryAdapter.query_to_event(make_query({})),
binding=binding,
descriptor=descriptor,
@@ -502,7 +516,8 @@ async def test_build_knowledge_bases_unions_config_and_policy_grants(app):
variables={'_knowledge_base_uuids': ['kb_policy']},
)
async def get_kb(kb_uuid):
async def get_kb(context, kb_uuid):
assert context == TEST_CONTEXT
return SimpleNamespace(
uuid=kb_uuid,
get_name=lambda: f'name-{kb_uuid}',
@@ -25,7 +25,7 @@ from langbot_plugin.api.entities.builtin.agent_runner.run_ledger import (
)
from langbot_plugin.entities.io.actions.enums import PluginToRuntimeAction
from .conftest import make_resources
from .conftest import bind_runtime_action_context, make_resources
class FakeConnection:
@@ -72,17 +72,32 @@ class FakeRunnerRegistry:
self.runners = runners
self.calls = []
async def list_runners(self, *, bound_plugins=None, use_cache=True):
self.calls.append({'bound_plugins': bound_plugins, 'use_cache': use_cache})
async def list_runners(self, context, *, bound_plugins=None, use_cache=True):
self.calls.append(
{
'workspace_uuid': context.workspace_uuid,
'bound_plugins': bound_plugins,
'use_cache': use_cache,
}
)
return self.runners
def _handler(db_engine, admin_plugins=None, runner_registry=None):
def _handler(
db_engine,
admin_plugins=None,
runner_registry=None,
plugin_identity='test/runner',
):
async def fake_disconnect():
return True
fake_app = FakeApplication(db_engine, admin_plugins=admin_plugins, runner_registry=runner_registry)
return RuntimeConnectionHandler(FakeConnection(), fake_disconnect, fake_app)
return bind_runtime_action_context(
RuntimeConnectionHandler(FakeConnection(), fake_disconnect, fake_app),
fake_app,
plugin_identity=plugin_identity,
)
async def _register_session(
@@ -505,6 +520,7 @@ async def test_agent_run_admin_can_list_runner_registry_without_run_id(db_engine
}
],
runner_registry=runner_registry,
plugin_identity='langbot/control',
)
runner_list = handler.actions['runner_list']
@@ -519,6 +535,7 @@ async def test_agent_run_admin_can_list_runner_registry_without_run_id(db_engine
assert result.data['items'][0]['id'] == 'plugin:test/runner/default'
assert runner_registry.calls == [
{
'workspace_uuid': 'workspace-test',
'bound_plugins': ['test/runner'],
'use_cache': True,
}
@@ -602,6 +619,7 @@ async def test_agent_run_admin_can_get_and_page_cross_scope_without_run_id(db_en
'permissions': ['agent_run:admin'],
}
],
plugin_identity='langbot/control',
)
run_get = handler.actions[PluginToRuntimeAction.RUN_GET.value]
run_events_page = handler.actions[PluginToRuntimeAction.RUN_EVENTS_PAGE.value]
@@ -720,7 +738,7 @@ async def test_configured_admin_identity_cannot_be_spoofed_with_other_run_sessio
)
assert result.code != 0
assert 'mismatch' in result.message.lower()
assert 'does not match' in result.message.lower()
@pytest.mark.asyncio
@@ -834,6 +852,7 @@ async def test_runtime_admin_can_register_list_and_claim_without_run_id(db_engin
'permissions': ['runtime:admin'],
}
],
plugin_identity='langbot/control',
)
runtime_register = handler.actions[PluginToRuntimeAction.RUNTIME_REGISTER.value]
runtime_list = handler.actions[PluginToRuntimeAction.RUNTIME_LIST.value]
@@ -917,6 +936,7 @@ async def test_runtime_admin_can_reconcile_without_run_id(db_engine):
'permissions': ['runtime:admin'],
}
],
plugin_identity='langbot/control',
)
runtime_reconcile = handler.actions['runtime_reconcile']
+20 -7
View File
@@ -25,11 +25,11 @@ from sqlalchemy.ext.asyncio import create_async_engine
from langbot.pkg.agent.runner.session_registry import AgentRunSessionRegistry
from langbot.pkg.agent.runner.persistent_state_store import PersistentStateStore, reset_persistent_state_store
from langbot.pkg.plugin.handler import RuntimeConnectionHandler
from langbot.pkg.plugin.handler import RuntimeConnectionHandler as HostRuntimeConnectionHandler
from langbot_plugin.entities.io.actions.enums import PluginToRuntimeAction
# Import shared test fixtures
from .conftest import make_resources
from .conftest import bind_runtime_action_context, make_resources
class FakeConnection:
@@ -48,6 +48,14 @@ class FakeApplication:
self.persistence_mgr.get_db_engine = MagicMock(return_value=db_engine)
class RuntimeConnectionHandler(HostRuntimeConnectionHandler):
"""Host handler with the trusted Runtime envelope installed for direct calls."""
def __init__(self, connection, disconnect, application):
super().__init__(connection, disconnect, application)
bind_runtime_action_context(self, application)
@pytest.fixture
def session_registry():
"""Create a fresh session registry for each test."""
@@ -126,8 +134,13 @@ class TestStateAPIHandlerAuthorization:
assert 'not found' in result.message.lower()
@pytest.mark.asyncio
async def test_state_get_missing_caller_plugin_identity_returns_error(self, session_registry, db_engine, persistent_store):
"""STATE_GET: missing caller_plugin_identity when session has plugin_identity returns error."""
async def test_state_get_uses_installation_identity_when_payload_omits_caller(
self,
session_registry,
db_engine,
persistent_store,
):
"""STATE_GET derives caller identity from the trusted installation binding."""
fake_app = FakeApplication(db_engine)
fake_app.persistence_mgr.get_db_engine = MagicMock(return_value=db_engine)
@@ -157,8 +170,8 @@ class TestStateAPIHandlerAuthorization:
'key': 'test_key',
})
assert result.code != 0
assert 'caller_plugin_identity is required' in result.message
assert result.code == 0
assert result.data == {'value': None}
await session_registry.unregister('run_test_missing_identity')
@@ -195,7 +208,7 @@ class TestStateAPIHandlerAuthorization:
})
assert result.code != 0
assert 'mismatch' in result.message.lower()
assert 'does not match' in result.message.lower()
await session_registry.unregister('run_test_mismatch')
@@ -6,6 +6,9 @@ from sqlalchemy.sql.dml import Update
from langbot.pkg.api.http.service.bot import BotService
WORKSPACE_UUID = 'workspace-a'
class _FakeResult:
def __init__(self, value):
self.value = value
@@ -21,7 +24,9 @@ class _PersistenceManager:
async def execute_async(self, statement):
if isinstance(statement, Update):
self.update_values = {
key: value for key, value in statement.compile().params.items() if not key.startswith('uuid_')
key: value
for key, value in statement.compile().params.items()
if not key.startswith(('uuid_', 'workspace_uuid_'))
}
return None
@@ -49,7 +54,7 @@ async def test_update_bot_copies_input_before_filtering_legacy_routing_fields():
'pipeline_routing_rules': [{'type': 'launcher_type'}],
}
await service.update_bot('bot-1', payload)
await service.update_bot(WORKSPACE_UUID, 'bot-1', payload)
# caller's dict must not be mutated
assert payload == {
@@ -0,0 +1,34 @@
import pytest
import sqlalchemy
from langbot.pkg.api.http.authz import WorkspaceRequiredError
from langbot.pkg.api.http.context import ExecutionContext, PrincipalContext, PrincipalType
from langbot.pkg.api.http.service.tenant import require_workspace_uuid, scope_statement
class _TenantRow:
workspace_uuid = sqlalchemy.column('workspace_uuid')
def test_require_workspace_uuid_accepts_execution_context():
context = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
)
assert require_workspace_uuid(context) == 'workspace-test'
@pytest.mark.parametrize('context', [None, '', ' '])
def test_require_workspace_uuid_rejects_missing_context(context):
with pytest.raises(WorkspaceRequiredError):
require_workspace_uuid(context)
def test_scope_statement_adds_workspace_predicate():
statement = scope_statement(sqlalchemy.select(_TenantRow.workspace_uuid), _TenantRow, 'workspace-test')
assert 'workspace_uuid = :workspace_uuid_1' in str(statement)
assert statement.compile().params == {'workspace_uuid_1': 'workspace-test'}
@@ -0,0 +1,381 @@
from __future__ import annotations
import datetime
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import sqlalchemy
from sqlalchemy.ext.asyncio import create_async_engine
from langbot.pkg.api.http.service.bot import BotService
from langbot.pkg.api.http.service.model import LLMModelsService
from langbot.pkg.api.http.service.pipeline import PipelineService
from langbot.pkg.api.http.service.provider import ModelProviderService
from langbot.pkg.api.http.service.tenant import require_workspace_uuid
from langbot.pkg.api.http.authz import WorkspaceRequiredError
from langbot.pkg.entity.persistence.base import Base
from langbot.pkg.entity.persistence.bot import Bot
from langbot.pkg.entity.persistence.model import LLMModel, ModelProvider
from langbot.pkg.entity.persistence.pipeline import LegacyPipeline
from langbot.pkg.entity.persistence.workspace import Workspace
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
pytestmark = pytest.mark.asyncio
WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
class _PersistenceManager:
def __init__(self, engine):
self.engine = engine
async def execute_async(self, *args, **kwargs):
async with self.engine.connect() as connection:
result = await connection.execute(*args, **kwargs)
await connection.commit()
return result
@staticmethod
def serialize_model(model, data, masked_columns=None):
masked_columns = masked_columns or []
return {
column.name: (
getattr(data, column.name).isoformat()
if isinstance(getattr(data, column.name), datetime.datetime)
else getattr(data, column.name)
)
for column in model.__table__.columns
if column.name not in masked_columns
}
@pytest.fixture
async def tenant_services(tmp_path):
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "tenant-resources.db"}')
async with engine.begin() as connection:
await connection.run_sync(Base.metadata.create_all)
await connection.execute(
sqlalchemy.insert(Workspace),
[
{
'uuid': WORKSPACE_A,
'instance_uuid': 'instance-a',
'name': 'Workspace A',
'slug': 'workspace-a',
'source': 'cloud_projection',
},
{
'uuid': WORKSPACE_B,
'instance_uuid': 'instance-b',
'name': 'Workspace B',
'slug': 'workspace-b',
'source': 'cloud_projection',
},
],
)
await connection.execute(
sqlalchemy.insert(ModelProvider),
[
{
'uuid': 'provider-a',
'workspace_uuid': WORKSPACE_A,
'name': 'Same Provider',
'requester': 'chatcmpl',
'base_url': 'https://a.invalid',
'api_keys': ['secret-a'],
},
{
'uuid': 'provider-b',
'workspace_uuid': WORKSPACE_B,
'name': 'Same Provider',
'requester': 'chatcmpl',
'base_url': 'https://b.invalid',
'api_keys': ['secret-b'],
},
],
)
await connection.execute(
sqlalchemy.insert(LLMModel),
[
{
'uuid': 'model-a',
'workspace_uuid': WORKSPACE_A,
'name': 'Same Model',
'provider_uuid': 'provider-a',
'abilities': [],
'extra_args': {},
'prefered_ranking': 0,
},
{
'uuid': 'model-b',
'workspace_uuid': WORKSPACE_B,
'name': 'Same Model',
'provider_uuid': 'provider-b',
'abilities': [],
'extra_args': {},
'prefered_ranking': 0,
},
],
)
await connection.execute(
sqlalchemy.insert(LegacyPipeline),
[
{
'uuid': 'pipeline-a',
'workspace_uuid': WORKSPACE_A,
'name': 'Same Pipeline',
'description': 'A',
'for_version': 'test',
'is_default': False,
'stages': [],
'config': {},
'extensions_preferences': {},
},
{
'uuid': 'pipeline-b',
'workspace_uuid': WORKSPACE_B,
'name': 'Same Pipeline',
'description': 'B',
'for_version': 'test',
'is_default': False,
'stages': [],
'config': {},
'extensions_preferences': {},
},
],
)
await connection.execute(
sqlalchemy.insert(Bot),
[
{
'uuid': 'bot-a',
'workspace_uuid': WORKSPACE_A,
'name': 'Same Bot',
'description': 'A',
'adapter': 'test',
'adapter_config': {},
'enable': False,
'use_pipeline_uuid': 'pipeline-a',
'use_pipeline_name': 'Same Pipeline',
'pipeline_routing_rules': [],
},
{
'uuid': 'bot-b',
'workspace_uuid': WORKSPACE_B,
'name': 'Same Bot',
'description': 'B',
'adapter': 'test',
'adapter_config': {},
'enable': False,
'use_pipeline_uuid': 'pipeline-b',
'use_pipeline_name': 'Same Pipeline',
'pipeline_routing_rules': [],
},
],
)
runtime_provider_a = SimpleNamespace(provider_entity=SimpleNamespace(uuid='provider-a'))
runtime_provider_b = SimpleNamespace(provider_entity=SimpleNamespace(uuid='provider-b'))
application = SimpleNamespace(
persistence_mgr=_PersistenceManager(engine),
instance_config=SimpleNamespace(data={'system': {'limitation': {}}, 'api': {}}),
ver_mgr=SimpleNamespace(get_current_version=lambda: 'test'),
platform_mgr=SimpleNamespace(
load_bot=AsyncMock(return_value=SimpleNamespace(enable=False)),
remove_bot=AsyncMock(),
get_bot_by_uuid=AsyncMock(return_value=None),
),
pipeline_mgr=SimpleNamespace(
load_pipeline=AsyncMock(),
remove_pipeline=AsyncMock(),
),
model_mgr=SimpleNamespace(
provider_dict={'provider-a': runtime_provider_a, 'provider-b': runtime_provider_b},
llm_models=[],
embedding_models=[],
rerank_models=[],
load_provider=AsyncMock(),
cache_provider=AsyncMock(),
get_provider_by_uuid=AsyncMock(return_value=runtime_provider_a),
reload_provider=AsyncMock(),
remove_provider=AsyncMock(),
load_llm_model_with_provider=AsyncMock(return_value=SimpleNamespace()),
cache_llm_model=AsyncMock(),
remove_llm_model=AsyncMock(),
),
sess_mgr=SimpleNamespace(session_list=[]),
)
application.provider_service = ModelProviderService(application)
application.llm_model_service = LLMModelsService(application)
application.pipeline_service = PipelineService(application)
application.bot_service = BotService(application)
yield application, engine
await engine.dispose()
async def test_context_is_mandatory_and_fails_closed(tenant_services):
application, _engine = tenant_services
with pytest.raises(WorkspaceRequiredError):
require_workspace_uuid(None)
with pytest.raises(WorkspaceRequiredError):
await application.bot_service.get_bots(None)
with pytest.raises(WorkspaceRequiredError):
await application.provider_service.get_providers(None)
with pytest.raises(WorkspaceRequiredError):
await application.pipeline_service.get_pipelines(None)
with pytest.raises(WorkspaceRequiredError):
await application.llm_model_service.get_llm_models(None)
async def test_lists_and_same_names_are_isolated(tenant_services):
application, _engine = tenant_services
assert [item['uuid'] for item in await application.bot_service.get_bots(WORKSPACE_A)] == ['bot-a']
assert [item['uuid'] for item in await application.pipeline_service.get_pipelines(WORKSPACE_A)] == ['pipeline-a']
assert [item['uuid'] for item in await application.provider_service.get_providers(WORKSPACE_A)] == ['provider-a']
assert [item['uuid'] for item in await application.llm_model_service.get_llm_models(WORKSPACE_A)] == ['model-a']
async def test_cross_workspace_uuid_guessing_cannot_read_update_or_delete(tenant_services):
application, engine = tenant_services
assert await application.bot_service.get_bot(WORKSPACE_A, 'bot-b') is None
assert await application.pipeline_service.get_pipeline(WORKSPACE_A, 'pipeline-b') is None
assert await application.provider_service.get_provider(WORKSPACE_A, 'provider-b') is None
assert await application.llm_model_service.get_llm_model(WORKSPACE_A, 'model-b') is None
with pytest.raises(WorkspaceNotFoundError):
await application.bot_service.update_bot(WORKSPACE_A, 'bot-b', {'name': 'stolen'})
with pytest.raises(WorkspaceNotFoundError):
await application.pipeline_service.update_pipeline(
WORKSPACE_A,
'pipeline-b',
{'description': 'stolen'},
)
with pytest.raises(WorkspaceNotFoundError):
await application.provider_service.update_provider(WORKSPACE_A, 'provider-b', {'name': 'stolen'})
with pytest.raises(WorkspaceNotFoundError):
await application.llm_model_service.update_llm_model(
WORKSPACE_A,
'model-b',
{'name': 'stolen'},
)
with pytest.raises(WorkspaceNotFoundError):
await application.bot_service.delete_bot(WORKSPACE_A, 'bot-b')
with pytest.raises(WorkspaceNotFoundError):
await application.pipeline_service.delete_pipeline(WORKSPACE_A, 'pipeline-b')
with pytest.raises(WorkspaceNotFoundError):
await application.provider_service.delete_provider(WORKSPACE_A, 'provider-b')
with pytest.raises(WorkspaceNotFoundError):
await application.llm_model_service.delete_llm_model(WORKSPACE_A, 'model-b')
async with engine.connect() as connection:
assert await connection.scalar(sqlalchemy.select(Bot.name).where(Bot.uuid == 'bot-b')) == 'Same Bot'
assert (
await connection.scalar(sqlalchemy.select(LegacyPipeline.uuid).where(LegacyPipeline.uuid == 'pipeline-b'))
== 'pipeline-b'
)
assert (
await connection.scalar(sqlalchemy.select(ModelProvider.name).where(ModelProvider.uuid == 'provider-b'))
== 'Same Provider'
)
assert await connection.scalar(sqlalchemy.select(LLMModel.uuid).where(LLMModel.uuid == 'model-b')) == 'model-b'
async def test_cross_workspace_parent_references_are_rejected(tenant_services):
application, _engine = tenant_services
with pytest.raises(ValueError, match='Pipeline not found'):
await application.bot_service.update_bot(
WORKSPACE_A,
'bot-a',
{
'event_bindings': [
{
'event_pattern': 'message.*',
'target_type': 'pipeline',
'target_uuid': 'pipeline-b',
}
]
},
)
with pytest.raises(WorkspaceNotFoundError):
await application.llm_model_service.create_llm_model(
WORKSPACE_A,
{
'name': 'Cross reference',
'provider_uuid': 'provider-b',
'abilities': [],
'extra_args': {},
'prefered_ranking': 0,
},
auto_set_to_default_pipeline=False,
)
async def test_created_resources_are_bound_to_callers_workspace(tenant_services):
application, engine = tenant_services
runtime_provider = SimpleNamespace(provider_entity=SimpleNamespace(uuid='provider-created'))
application.model_mgr.load_provider.return_value = runtime_provider
provider_uuid = await application.provider_service.create_provider(
WORKSPACE_A,
{
'name': 'Created Provider',
'requester': 'chatcmpl',
'base_url': 'https://created.invalid',
'api_keys': [],
},
)
pipeline_uuid = await application.pipeline_service.create_pipeline(
WORKSPACE_A,
{'name': 'Created Pipeline', 'description': 'created'},
)
bot_uuid = await application.bot_service.create_bot(
WORKSPACE_A,
{
'name': 'Created Bot',
'description': 'created',
'adapter': 'test',
'adapter_config': {},
'enable': False,
'pipeline_routing_rules': [],
},
)
model_uuid = await application.llm_model_service.create_llm_model(
WORKSPACE_A,
{
'name': 'Created Model',
'provider_uuid': 'provider-a',
'abilities': [],
'extra_args': {},
'prefered_ranking': 0,
},
auto_set_to_default_pipeline=False,
)
async with engine.connect() as connection:
assert (
await connection.scalar(
sqlalchemy.select(ModelProvider.workspace_uuid).where(ModelProvider.uuid == provider_uuid)
)
== WORKSPACE_A
)
assert (
await connection.scalar(
sqlalchemy.select(LegacyPipeline.workspace_uuid).where(LegacyPipeline.uuid == pipeline_uuid)
)
== WORKSPACE_A
)
assert await connection.scalar(sqlalchemy.select(Bot.workspace_uuid).where(Bot.uuid == bot_uuid)) == WORKSPACE_A
assert (
await connection.scalar(sqlalchemy.select(LLMModel.workspace_uuid).where(LLMModel.uuid == model_uuid))
== WORKSPACE_A
)
+74
View File
@@ -0,0 +1,74 @@
from langbot.pkg.api.http import authz
from langbot.pkg.api.http.context import PrincipalContext, PrincipalType, RequestContext, WorkspaceContext
def _context(role: authz.WorkspaceRole) -> RequestContext:
return RequestContext(
instance_uuid='instance-test',
placement_generation=1,
request_id='request-test',
auth_type='user-token',
principal=PrincipalContext(
principal_type=PrincipalType.ACCOUNT,
account_uuid='account-test',
),
workspace=WorkspaceContext(
workspace_uuid='workspace-test',
membership_uuid='membership-test',
role=role.value,
permissions=authz.permissions_for_role(role),
),
)
def test_owner_has_every_fixed_permission():
ctx = _context(authz.WorkspaceRole.OWNER)
assert ctx.workspace.permissions == frozenset(permission.value for permission in authz.Permission)
def test_admin_cannot_transfer_owner_delete_workspace_or_link_billing():
ctx = _context(authz.WorkspaceRole.ADMIN)
assert not authz.has_permission(ctx, authz.Permission.OWNER_TRANSFER)
assert not authz.has_permission(ctx, authz.Permission.WORKSPACE_DELETE)
assert not authz.has_permission(ctx, authz.Permission.BILLING_LINK_MANAGE)
assert authz.has_permission(ctx, authz.Permission.MEMBER_INVITE)
def test_operator_can_run_but_cannot_manage_resources_or_secrets():
ctx = _context(authz.WorkspaceRole.OPERATOR)
assert authz.has_permission(ctx, authz.Permission.RUNTIME_OPERATE)
assert not authz.has_permission(ctx, authz.Permission.RESOURCE_MANAGE)
assert not authz.has_permission(ctx, authz.Permission.PROVIDER_SECRET_MANAGE)
def test_unknown_role_has_no_permissions():
assert authz.permissions_for_role('unknown') == frozenset()
def test_require_permission_reports_stable_permission():
ctx = _context(authz.WorkspaceRole.VIEWER)
try:
authz.require_permission(ctx, authz.Permission.RESOURCE_MANAGE)
except authz.PermissionDeniedError as exc:
assert exc.permission == authz.Permission.RESOURCE_MANAGE.value
assert exc.error_code == 'permission_denied'
else:
raise AssertionError('PermissionDeniedError was not raised')
def test_execution_context_preserves_workspace_and_generation():
from langbot.pkg.api.http.context import ExecutionContext
ctx = _context(authz.WorkspaceRole.DEVELOPER)
execution = ExecutionContext.from_request(ctx, bot_uuid='bot-test', pipeline_uuid='pipeline-test')
assert execution.instance_uuid == 'instance-test'
assert execution.workspace_uuid == 'workspace-test'
assert execution.placement_generation == 1
assert execution.bot_uuid == 'bot-test'
assert execution.pipeline_uuid == 'pipeline-test'
assert execution.trigger_principal == ctx.principal
@@ -0,0 +1,39 @@
from __future__ import annotations
import quart
from langbot.pkg.api.http.controller import main as controller_main
from langbot.pkg.utils import bounded_executor
async def test_bounded_json_request_decodes_off_loop_in_workspace_scope(
monkeypatch,
):
app = quart.Quart(__name__)
app.request_class = controller_main.BoundedJSONRequest
observed_scopes: list[str | None] = []
async def fake_to_thread(fn, *args, **kwargs):
observed_scopes.append(bounded_executor.current_blocking_work_scope())
return fn(*args, **kwargs)
monkeypatch.setattr(
controller_main.asyncio,
'to_thread',
fake_to_thread,
)
@app.post('/json')
async def parse_json():
with bounded_executor.blocking_work_scope('workspace-a'):
payload = await quart.request.get_json()
return quart.jsonify(payload)
response = await app.test_client().post(
'/json',
json={'nested': {'value': 1}},
)
assert response.status_code == 200
assert await response.get_json() == {'nested': {'value': 1}}
assert observed_scopes == ['workspace-a']
@@ -0,0 +1,78 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import quart
from langbot.pkg.api.http.context import (
ExecutionContext,
PrincipalContext,
PrincipalType,
RequestContext,
WorkspaceContext,
)
from langbot.pkg.api.http.controller.group import RouterGroup
from langbot.pkg.cloud.entitlements import EntitlementSnapshot, EntitlementUnavailableError
from langbot.pkg.cloud.entitlements import EntitlementResolver
class _Group(RouterGroup):
async def initialize(self) -> None:
return None
def _router(deployment) -> _Group:
provider = getattr(deployment, 'entitlement_provider', None)
resolver = EntitlementResolver('instance-a', provider) if provider is not None else None
ap = SimpleNamespace(deployment=deployment, entitlement_resolver=resolver)
return _Group(ap, quart.Quart(__name__))
@pytest.mark.asyncio
async def test_cloud_request_resolves_verified_entitlement_revision():
snapshot = EntitlementSnapshot(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
entitlement_revision=9,
status='active',
not_before=1,
expires_at=4_000_000_000,
features={},
limits={},
)
provider = SimpleNamespace(get_workspace_entitlement=AsyncMock(return_value=snapshot))
router = _router(SimpleNamespace(multi_workspace_enabled=True, entitlement_provider=provider))
revision = await router._resolve_entitlement_revision('instance-a', 'workspace-a')
assert revision == 9
provider.get_workspace_entitlement.assert_awaited_once_with('workspace-a')
@pytest.mark.asyncio
async def test_cloud_request_fails_closed_without_entitlement_provider():
router = _router(SimpleNamespace(multi_workspace_enabled=True, entitlement_provider=None))
with pytest.raises(EntitlementUnavailableError):
await router._resolve_entitlement_revision('instance-a', 'workspace-a')
def test_execution_context_preserves_entitlement_revision():
request = RequestContext(
instance_uuid='instance-a',
placement_generation=1,
request_id='request-a',
auth_type='user-token',
principal=PrincipalContext(PrincipalType.ACCOUNT, account_uuid='account-a'),
workspace=WorkspaceContext(
workspace_uuid='workspace-a',
membership_uuid='membership-a',
role='owner',
permissions=frozenset(),
),
entitlement_revision=11,
)
assert ExecutionContext.from_request(request).entitlement_revision == 11
@@ -0,0 +1,257 @@
from __future__ import annotations
import contextlib
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
import quart
from langbot.pkg.api.http.controller import group
from langbot.pkg.api.http.controller.groups.webhooks import WebhookRouterGroup
from langbot.pkg.utils.bounded_executor import (
BlockingWorkCapacityError,
current_blocking_work_scope,
)
pytestmark = pytest.mark.asyncio
class _FailingRouterGroup(group.RouterGroup):
name = 'failing-test'
path = '/failing-test'
async def initialize(self) -> None:
@self.route('', methods=['GET'], auth_type=group.AuthType.NONE)
async def _():
raise RuntimeError('database password=do-not-return')
class _AuthenticatedRouterGroup(group.RouterGroup):
name = 'authenticated-test'
path = '/authenticated-test'
async def initialize(self) -> None:
@self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
async def _():
return self.success()
class _BlockingCapacityRouterGroup(group.RouterGroup):
name = 'blocking-capacity-test'
path = '/blocking-capacity-test'
async def initialize(self) -> None:
@self.route('', methods=['GET'], auth_type=group.AuthType.NONE)
async def _():
raise BlockingWorkCapacityError('Workspace blocking executor capacity reached')
class _InvalidAccountRouterGroup(group.RouterGroup):
name = 'invalid-account-test'
path = '/invalid-account-test'
async def initialize(self) -> None:
@self.route(
'',
methods=['GET'],
auth_type=group.AuthType.ACCOUNT_TOKEN,
permission='workspace.view',
)
async def _():
return self.success()
async def test_unhandled_http_error_returns_generic_body_and_correlated_request_id():
logger = Mock()
application = SimpleNamespace(logger=logger)
quart_app = quart.Quart(__name__)
await _FailingRouterGroup(application, quart_app).initialize()
response = await quart_app.test_client().get(
'/failing-test',
headers={'X-Request-Id': 'request-http-test'},
)
assert response.status_code == 500
assert await response.get_json() == {
'code': 'internal_error',
'msg': 'Internal server error',
'request_id': 'request-http-test',
}
assert response.headers['X-Request-Id'] == 'request-http-test'
log_message = logger.error.call_args.args[0]
assert 'request_id=request-http-test' in log_message
assert 'database password=do-not-return' in log_message
assert 'do-not-return' not in (await response.get_data(as_text=True))
async def test_public_webhook_error_uses_same_generic_error_contract():
logger = Mock()
application = SimpleNamespace(
logger=logger,
platform_mgr=SimpleNamespace(
resolve_public_bot=AsyncMock(side_effect=RuntimeError('adapter credential=do-not-return'))
),
)
quart_app = quart.Quart(__name__)
await WebhookRouterGroup(application, quart_app).initialize()
response = await quart_app.test_client().post(
'/bots/11111111-1111-4111-8111-111111111111',
headers={'X-Request-Id': 'request-webhook-test'},
)
assert response.status_code == 500
assert await response.get_json() == {
'code': 'internal_error',
'msg': 'Internal server error',
'request_id': 'request-webhook-test',
}
assert response.headers['X-Request-Id'] == 'request-webhook-test'
log_message = logger.error.call_args.args[0]
assert 'request_id=request-webhook-test' in log_message
assert 'adapter credential=do-not-return' in log_message
assert 'do-not-return' not in (await response.get_data(as_text=True))
async def test_blocking_work_capacity_maps_to_retryable_http_response():
application = SimpleNamespace(logger=Mock())
quart_app = quart.Quart(__name__)
await _BlockingCapacityRouterGroup(application, quart_app).initialize()
response = await quart_app.test_client().get('/blocking-capacity-test')
assert response.status_code == 429
assert await response.get_json() == {
'code': 'blocking_work_capacity_exceeded',
'msg': 'Workspace blocking executor capacity reached',
}
async def test_public_webhook_carries_scope_without_holding_database_session():
class ScopeOnlyPersistenceManager:
mode = SimpleNamespace(value='cloud_runtime')
def __init__(self):
self.active_workspace = None
@contextlib.asynccontextmanager
async def tenant_scope(self, workspace_uuid):
self.active_workspace = workspace_uuid
try:
yield
finally:
self.active_workspace = None
def current_session(self):
return None
persistence_mgr = ScopeOnlyPersistenceManager()
workspace_uuid = '00000000-0000-0000-0000-00000000000a'
bot_uuid = '11111111-1111-4111-8111-111111111111'
class Adapter:
async def handle_unified_webhook(self, **_kwargs):
assert persistence_mgr.active_workspace == workspace_uuid
assert persistence_mgr.current_session() is None
assert current_blocking_work_scope() == workspace_uuid
return {'ok': True}
async def get_execution_binding(resolved_workspace_uuid, expected_generation=None):
assert resolved_workspace_uuid == workspace_uuid
assert expected_generation == 4
runtime_bot = SimpleNamespace(
workspace_uuid=workspace_uuid,
placement_generation=4,
enable=True,
adapter=Adapter(),
)
application = SimpleNamespace(
logger=Mock(),
persistence_mgr=persistence_mgr,
platform_mgr=SimpleNamespace(resolve_public_bot=AsyncMock(return_value=runtime_bot)),
workspace_service=SimpleNamespace(get_execution_binding=get_execution_binding),
)
quart_app = quart.Quart(__name__)
await WebhookRouterGroup(application, quart_app).initialize()
response = await quart_app.test_client().post(f'/bots/{bot_uuid}')
assert response.status_code == 200
assert await response.get_json() == {'ok': True}
assert persistence_mgr.active_workspace is None
async def test_public_webhook_blocking_capacity_is_retryable():
workspace_uuid = '00000000-0000-0000-0000-00000000000a'
bot_uuid = '11111111-1111-4111-8111-111111111111'
class Adapter:
async def handle_unified_webhook(self, **_kwargs):
raise BlockingWorkCapacityError(
'Workspace blocking executor capacity reached',
scope=workspace_uuid,
)
runtime_bot = SimpleNamespace(
workspace_uuid=workspace_uuid,
placement_generation=4,
enable=True,
adapter=Adapter(),
)
application = SimpleNamespace(
logger=Mock(),
persistence_mgr=SimpleNamespace(mode=SimpleNamespace(value='oss')),
platform_mgr=SimpleNamespace(resolve_public_bot=AsyncMock(return_value=runtime_bot)),
workspace_service=SimpleNamespace(get_execution_binding=AsyncMock(return_value=None)),
)
quart_app = quart.Quart(__name__)
await WebhookRouterGroup(application, quart_app).initialize()
response = await quart_app.test_client().post(f'/bots/{bot_uuid}')
assert response.status_code == 429
assert await response.get_json() == {
'code': 'blocking_work_capacity_exceeded',
'msg': 'Workspace blocking executor capacity reached',
}
async def test_authentication_failure_does_not_return_internal_exception_text():
logger = Mock()
application = SimpleNamespace(
logger=logger,
user_service=SimpleNamespace(
get_authenticated_account=AsyncMock(side_effect=RuntimeError('database password=do-not-return'))
),
)
quart_app = quart.Quart(__name__)
await _AuthenticatedRouterGroup(application, quart_app).initialize()
response = await quart_app.test_client().get(
'/authenticated-test',
headers={
'Authorization': 'Bearer invalid',
'X-Request-Id': 'request-auth-test',
},
)
assert response.status_code == 401
assert await response.get_json() == {
'code': 'invalid_authentication',
'msg': 'Invalid authentication credentials',
}
assert 'do-not-return' not in (await response.get_data(as_text=True))
assert 'request_id=request-auth-test' in logger.warning.call_args.args[0]
assert 'database password=do-not-return' in logger.warning.call_args.args[0]
async def test_account_token_route_cannot_declare_workspace_permission():
application = SimpleNamespace(logger=Mock())
quart_app = quart.Quart(__name__)
with pytest.raises(ValueError, match='cannot declare Workspace permissions'):
await _InvalidAccountRouterGroup(application, quart_app).initialize()
@@ -0,0 +1,73 @@
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
import quart
from sqlalchemy.ext.asyncio import create_async_engine
from langbot.pkg.api.http.controller import group
from langbot.pkg.persistence.mgr import PersistenceManager, PersistenceMode
pytestmark = pytest.mark.asyncio
async def test_authenticated_route_does_not_hold_database_session_during_external_wait():
entered = asyncio.Event()
release = asyncio.Event()
observations: list[bool] = []
engine = create_async_engine('sqlite+aiosqlite:///:memory:')
persistence = PersistenceManager(object(), mode=PersistenceMode.CLOUD_RUNTIME)
persistence.db = SimpleNamespace(get_engine=lambda: engine)
class BlockingRouter(group.RouterGroup):
name = 'blocking-route-test'
path = '/blocking-route-test'
async def initialize(self) -> None:
@self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
async def _():
observations.append(persistence.current_session() is None)
entered.set()
await release.wait()
observations.append(persistence.current_session() is None)
return self.success(data={})
account = SimpleNamespace(uuid='account-a', user='owner@example.com')
access = SimpleNamespace(
execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=1),
workspace=SimpleNamespace(uuid='workspace-a'),
membership=SimpleNamespace(uuid='membership-a', role='owner', projection_revision=1),
)
application = SimpleNamespace(
persistence_mgr=persistence,
deployment=SimpleNamespace(multi_workspace_enabled=False),
user_service=SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account)),
workspace_collaboration_service=SimpleNamespace(resolve_account_workspace=AsyncMock(return_value=access)),
logger=Mock(),
)
quart_app = quart.Quart(__name__)
await BlockingRouter(application, quart_app).initialize()
client = quart_app.test_client()
request = asyncio.create_task(
client.get(
'/blocking-route-test',
headers={'Authorization': 'Bearer token', 'X-Workspace-Id': 'workspace-a'},
)
)
try:
await entered.wait()
assert observations == [True]
release.set()
response = await request
assert response.status_code == 200
assert observations == [True, True]
finally:
release.set()
if not request.done():
await request
await engine.dispose()
@@ -16,6 +16,7 @@ from langbot.pkg.api.http.service.agent import (
pytestmark = pytest.mark.asyncio
WORKSPACE_UUID = 'workspace-test'
def _result(items: list | None = None, first_item=None):
@@ -33,6 +34,7 @@ def _agent_row(
supported_event_patterns: list[str] | None = None,
):
return SimpleNamespace(
workspace_uuid=WORKSPACE_UUID,
uuid=agent_uuid,
name=name,
description='Agent description',
@@ -53,6 +55,7 @@ def _agent_row(
def _serialize_agent(model_cls, entity, masked_columns=None):
return {
'workspace_uuid': entity.workspace_uuid,
'uuid': entity.uuid,
'name': entity.name,
'description': entity.description,
@@ -72,7 +75,11 @@ def _compiled_params(statement):
def _compiled_update_values(statement):
return {key: value for key, value in statement.compile().params.items() if not key.startswith('uuid_')}
return {
key: value
for key, value in statement.compile().params.items()
if not key.startswith(('uuid_', 'workspace_uuid_'))
}
def _make_app():
@@ -103,7 +110,10 @@ class TestAgentServiceMetadata:
return_value=[{'name': 'trigger'}, ai_metadata, {'name': 'output'}]
)
metadata = await AgentService(app).get_agent_metadata()
metadata = await AgentService(app).get_agent_metadata(WORKSPACE_UUID)
app.pipeline_service.get_pipeline_metadata.assert_awaited_once_with(
WORKSPACE_UUID
)
assert metadata['runner_config'] == ai_metadata
assert metadata['kinds'] == [
@@ -148,7 +158,11 @@ class TestAgentServiceListAndLookup:
]
)
agents = await AgentService(app).get_agents(sort_by='updated_at', sort_order='DESC')
agents = await AgentService(app).get_agents(
WORKSPACE_UUID,
sort_by='updated_at',
sort_order='DESC',
)
assert [agent['uuid'] for agent in agents] == ['pipeline-1', 'agent-1']
assert agents[0]['kind'] == AGENT_KIND_PIPELINE
@@ -169,7 +183,7 @@ class TestAgentServiceListAndLookup:
agent = _agent_row(agent_uuid='agent-1')
app.persistence_mgr.execute_async = AsyncMock(return_value=_result(first_item=agent))
result = await AgentService(app).get_agent('agent-1')
result = await AgentService(app).get_agent(WORKSPACE_UUID, 'agent-1')
assert result['uuid'] == 'agent-1'
assert result['kind'] == AGENT_KIND_AGENT
@@ -191,7 +205,7 @@ class TestAgentServiceListAndLookup:
}
)
result = await AgentService(app).get_agent('pipeline-1')
result = await AgentService(app).get_agent(WORKSPACE_UUID, 'pipeline-1')
assert result['kind'] == AGENT_KIND_PIPELINE
assert result['enabled'] is True
@@ -217,6 +231,7 @@ class TestAgentServiceCreateUpdateDelete:
app.persistence_mgr.execute_async = AsyncMock(return_value=Mock())
result = await AgentService(app).create_agent(
WORKSPACE_UUID,
{
'name': 'Support Agent',
'description': 'Handles support events',
@@ -259,7 +274,10 @@ class TestAgentServiceCreateUpdateDelete:
app = _make_app()
with pytest.raises(ValueError, match='Agent config|runner_config'):
await AgentService(app).create_agent({'name': 'Invalid Agent', 'config': config})
await AgentService(app).create_agent(
WORKSPACE_UUID,
{'name': 'Invalid Agent', 'config': config},
)
app.persistence_mgr.execute_async.assert_not_awaited()
@@ -284,6 +302,7 @@ class TestAgentServiceCreateUpdateDelete:
with pytest.raises(ValueError, match=f'{field_name}.*boolean'):
await AgentService(app).create_agent(
WORKSPACE_UUID,
{
'name': 'Invalid Agent',
'config': {
@@ -302,6 +321,7 @@ class TestAgentServiceCreateUpdateDelete:
with pytest.raises(ValueError, match=r'mcp-resources\[0\]\.enabled.*boolean'):
await AgentService(app).create_agent(
WORKSPACE_UUID,
{
'name': 'Invalid Agent',
'config': {
@@ -324,6 +344,7 @@ class TestAgentServiceCreateUpdateDelete:
app.persistence_mgr.execute_async = AsyncMock(return_value=Mock())
await AgentService(app).create_agent(
WORKSPACE_UUID,
{
'name': 'Unconfigured Agent',
'component_ref': 'plugin:caller/must-not-win/default',
@@ -343,6 +364,7 @@ class TestAgentServiceCreateUpdateDelete:
with pytest.raises(ValueError, match='runner_config'):
await AgentService(app).update_agent(
WORKSPACE_UUID,
'agent-1',
{
'config': {
@@ -368,6 +390,7 @@ class TestAgentServiceCreateUpdateDelete:
}
await AgentService(app).update_agent(
WORKSPACE_UUID,
'agent-1',
{
'uuid': 'caller-owned-uuid',
@@ -400,6 +423,7 @@ class TestAgentServiceCreateUpdateDelete:
)
await AgentService(app).update_agent(
WORKSPACE_UUID,
'agent-1',
{
'name': 'Updated Agent',
@@ -423,6 +447,7 @@ class TestAgentServiceCreateUpdateDelete:
)
await AgentService(app).update_agent(
WORKSPACE_UUID,
'agent-1',
{'component_ref': 'plugin:caller/must-not-win/default'},
)
@@ -441,6 +466,7 @@ class TestAgentServiceCreateUpdateDelete:
config = {'runner': {'id': ''}, 'runner_config': {}}
await AgentService(app).update_agent(
WORKSPACE_UUID,
'agent-1',
{
'component_ref': 'plugin:caller/must-not-win/default',
@@ -459,6 +485,7 @@ class TestAgentServiceCreateUpdateDelete:
service = AgentService(app)
created = await service.create_agent(
WORKSPACE_UUID,
{
'kind': AGENT_KIND_PIPELINE,
'name': 'Pipeline Agent',
@@ -466,11 +493,16 @@ class TestAgentServiceCreateUpdateDelete:
'emoji': 'P',
}
)
await service.update_agent('pipeline-1', {'name': 'Updated Pipeline'})
await service.delete_agent('pipeline-1')
await service.update_agent(
WORKSPACE_UUID,
'pipeline-1',
{'name': 'Updated Pipeline'},
)
await service.delete_agent(WORKSPACE_UUID, 'pipeline-1')
assert created == {'uuid': 'pipeline-created', 'kind': AGENT_KIND_PIPELINE}
app.pipeline_service.create_pipeline.assert_awaited_once_with(
WORKSPACE_UUID,
{
'name': 'Pipeline Agent',
'description': 'Legacy pipeline',
@@ -479,7 +511,11 @@ class TestAgentServiceCreateUpdateDelete:
}
)
app.pipeline_service.update_pipeline.assert_awaited_once_with(
WORKSPACE_UUID,
'pipeline-1',
{'name': 'Updated Pipeline'},
)
app.pipeline_service.delete_pipeline.assert_awaited_once_with('pipeline-1')
app.pipeline_service.delete_pipeline.assert_awaited_once_with(
WORKSPACE_UUID,
'pipeline-1',
)
@@ -1,482 +1,466 @@
"""
Unit tests for ApiKeyService.
Tests API key CRUD operations with mocked persistence layer.
Source: src/langbot/pkg/api/http/service/apikey.py
"""
from __future__ import annotations
import pytest
from unittest.mock import AsyncMock, Mock, patch
import datetime
import hashlib
import logging
import uuid
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import sqlalchemy
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
from langbot.pkg.api.http.authz import Permission, PermissionDeniedError
from langbot.pkg.api.http.context import PrincipalContext, PrincipalType, RequestContext, WorkspaceContext
from langbot.pkg.api.http.service.apikey import ApiKeyService
from langbot.pkg.entity.persistence.apikey import ApiKey
from langbot.pkg.entity.persistence.base import Base
from langbot.pkg.entity.persistence.user import User
from langbot.pkg.entity.persistence.workspace import (
Workspace,
WorkspaceExecutionSource,
WorkspaceExecutionState,
WorkspaceSource,
)
from langbot.pkg.workspace.policy import SingleWorkspacePolicy
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
from langbot.pkg.workspace.service import WorkspaceService
pytestmark = pytest.mark.asyncio
class _PersistenceManager:
def __init__(self, engine):
self.engine = engine
def get_db_engine(self):
return self.engine
async def execute_async(self, *args, **kwargs):
async with self.engine.connect() as connection:
result = await connection.execute(*args, **kwargs)
await connection.commit()
return result
@staticmethod
def serialize_model(model, row, masked_columns=()):
return {
column.name: (
getattr(row, column.name).isoformat()
if isinstance(getattr(row, column.name), datetime.datetime)
else getattr(row, column.name)
)
for column in model.__table__.columns
if column.name not in masked_columns
}
def _context(workspace_uuid: str, account_uuid: str, permissions: set[Permission]) -> RequestContext:
return RequestContext(
instance_uuid='api-key-instance',
placement_generation=1,
request_id=str(uuid.uuid4()),
auth_type='user-token',
principal=PrincipalContext(PrincipalType.ACCOUNT, account_uuid=account_uuid),
workspace=WorkspaceContext(
workspace_uuid=workspace_uuid,
membership_uuid=str(uuid.uuid4()),
role='owner',
permissions=frozenset(permission.value for permission in permissions),
),
)
@pytest.fixture
async def api_key_context(tmp_path):
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "api-keys.db"}')
async with engine.begin() as connection:
await connection.run_sync(Base.metadata.create_all)
application = SimpleNamespace(
persistence_mgr=_PersistenceManager(engine),
instance_config=SimpleNamespace(data={'api': {'global_api_key': ''}}),
logger=logging.getLogger('api-key-test'),
)
application.workspace_service = WorkspaceService(application, instance_uuid='api-key-instance')
workspace = await application.workspace_service.ensure_singleton_workspace()
account_uuid = str(uuid.uuid4())
session_factory = async_sessionmaker(engine, expire_on_commit=False)
async with session_factory.begin() as session:
session.add(
User(
uuid=account_uuid,
user='owner@example.com',
normalized_email='owner@example.com',
password='hash',
account_type='local',
)
)
service = ApiKeyService(application)
context = _context(workspace.uuid, account_uuid, set(Permission))
yield application, service, context, engine
await engine.dispose()
async def test_secret_is_returned_once_and_only_hash_is_persisted(api_key_context):
_application, service, context, engine = api_key_context
created = await service.create_api_key(context, 'Automation', 'CI key')
secret = created['key']
assert secret.startswith('lbk_')
assert created['secret_available'] is True
assert 'key_hash' not in created
listed = await service.get_api_keys(context)
assert len(listed) == 1
assert 'key' not in listed[0]
assert 'key_hash' not in listed[0]
assert listed[0]['secret_available'] is False
async with engine.connect() as connection:
stored = await connection.scalar(sqlalchemy.select(ApiKey.key_hash))
assert stored == hashlib.sha256(secret.encode()).hexdigest()
assert secret not in stored
async def test_authentication_derives_workspace_scopes_and_updates_usage(api_key_context):
_application, service, context, engine = api_key_context
created = await service.create_api_key(
context,
'Read only',
scopes=[Permission.RESOURCE_VIEW.value],
)
identity = await service.authenticate_api_key(created['key'])
assert identity is not None
assert identity.workspace_uuid == context.workspace_uuid
assert identity.permissions == frozenset({Permission.RESOURCE_VIEW.value})
async with engine.connect() as connection:
last_used_at = await connection.scalar(sqlalchemy.select(ApiKey.last_used_at))
assert last_used_at is not None
async def test_revoked_expired_and_unknown_keys_fail_closed(api_key_context):
_application, service, context, _engine = api_key_context
created = await service.create_api_key(context, 'Revocable')
await service.delete_api_key(context, created['id'])
assert await service.authenticate_api_key(created['key']) is None
assert await service.verify_api_key('') is False
assert await service.verify_api_key('plain-secret') is False
assert await service.verify_api_key('lbk_unknown') is False
expired_secret = 'lbk_expired'
await service.ap.persistence_mgr.execute_async(
sqlalchemy.insert(ApiKey).values(
workspace_uuid=context.workspace_uuid,
name='Expired',
key_hash=hashlib.sha256(expired_secret.encode()).hexdigest(),
scopes=[Permission.RESOURCE_VIEW.value],
status='active',
expires_at=datetime.datetime.now(datetime.UTC).replace(tzinfo=None) - datetime.timedelta(seconds=1),
)
)
assert await service.authenticate_api_key(expired_secret) is None
async def test_revoke_winning_last_used_update_race_fails_authentication(api_key_context):
application, service, context, _engine = api_key_context
created = await service.create_api_key(context, 'Racing revoke')
original_execute = application.persistence_mgr.execute_async
injected_revoke = False
async def execute_with_revoke(statement, *args, **kwargs):
nonlocal injected_revoke
if (
not injected_revoke
and isinstance(statement, sqlalchemy.sql.dml.Update)
and statement.table.name == ApiKey.__tablename__
):
injected_revoke = True
await original_execute(sqlalchemy.update(ApiKey).where(ApiKey.id == created['id']).values(status='revoked'))
return await original_execute(statement, *args, **kwargs)
application.persistence_mgr.execute_async = execute_with_revoke
assert await service.authenticate_api_key(created['key']) is None
assert injected_revoke is True
async def test_cross_workspace_crud_and_secret_guessing_are_isolated(api_key_context):
application, service, first_context, engine = api_key_context
second_workspace_uuid = str(uuid.uuid4())
async with async_sessionmaker(engine, expire_on_commit=False).begin() as session:
session.add(
Workspace(
uuid=second_workspace_uuid,
instance_uuid='api-key-instance',
name='Second',
slug='second',
source=WorkspaceSource.CLOUD_PROJECTION.value,
)
)
session.add(
WorkspaceExecutionState(
workspace_uuid=second_workspace_uuid,
instance_uuid='api-key-instance',
active_generation=3,
state='active',
write_fenced=False,
source=WorkspaceExecutionSource.CLOUD.value,
)
)
second_context = _context(second_workspace_uuid, first_context.account_uuid or '', set(Permission))
created = await service.create_api_key(first_context, 'First only')
assert await service.get_api_key(second_context, created['id']) is None
assert await service.get_api_keys(second_context) == []
identity = await service.authenticate_api_key(created['key'])
assert identity is not None
assert identity.workspace_uuid == first_context.workspace_uuid
assert identity.workspace_uuid != second_workspace_uuid
# Prove the explicit multi-Workspace policy does not change key-derived routing.
application.workspace_service.policy = SingleWorkspacePolicy(workspace_limit=10, multi_workspace_enabled=True)
identity = await service.authenticate_api_key(created['key'])
assert identity is not None
assert identity.workspace_uuid == first_context.workspace_uuid
async def test_global_config_key_is_oss_singleton_only(api_key_context):
application, service, _context_value, _engine = api_key_context
application.instance_config.data['api']['global_api_key'] = 'configured-secret'
identity = await service.authenticate_api_key('configured-secret')
assert identity is not None
assert identity.api_key_uuid == 'global-oss-api-key'
application.workspace_service.policy = SingleWorkspacePolicy(workspace_limit=10, multi_workspace_enabled=True)
assert await service.authenticate_api_key('configured-secret') is None
async def test_explicit_scopes_cannot_exceed_callers_workspace_permissions(api_key_context):
_application, service, context, _engine = api_key_context
limited_context = _context(
context.workspace_uuid,
context.account_uuid or '',
{Permission.API_KEY_MANAGE, Permission.RESOURCE_VIEW},
)
created = await service.create_api_key(
limited_context,
'Read only',
scopes=[Permission.RESOURCE_VIEW.value],
)
identity = await service.authenticate_api_key(created['key'])
assert identity is not None
assert identity.permissions == frozenset({Permission.RESOURCE_VIEW.value})
with pytest.raises(PermissionDeniedError) as exc_info:
await service.create_api_key(
limited_context,
'Escalated',
scopes=[Permission.WORKSPACE_DELETE.value],
)
assert exc_info.value.permission == Permission.WORKSPACE_DELETE.value
# Preserve the pre-tenancy CRUD and verification regression matrix while
# exercising it through the new Workspace-bound API. The assertions reflect
# intentional security changes: secrets are returned once, deletion revokes,
# and missing Workspace resources are reported as not found.
class TestApiKeyServiceGetApiKeys:
"""Tests for get_api_keys method."""
async def test_get_api_keys_empty_list(self, api_key_context):
_application, service, context, _engine = api_key_context
async def test_get_api_keys_empty_list(self):
"""Returns empty list when no API keys exist."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
mock_result = Mock()
mock_result.all = Mock(return_value=[])
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
ap.persistence_mgr.serialize_model = Mock(
side_effect=lambda model_cls, entity: {
'id': entity.id,
'name': entity.name,
'key': entity.key,
'description': entity.description,
}
if entity
else {}
)
assert await service.get_api_keys(context) == []
service = ApiKeyService(ap)
async def test_get_api_keys_returns_serialized_list(self, api_key_context):
_application, service, context, _engine = api_key_context
await service.create_api_key(context, 'Test Key 1', 'First test key')
await service.create_api_key(context, 'Test Key 2', 'Second test key')
# Execute
result = await service.get_api_keys()
result = await service.get_api_keys(context)
# Verify
assert result == []
ap.persistence_mgr.execute_async.assert_called_once()
async def test_get_api_keys_returns_serialized_list(self):
"""Returns serialized list of API keys."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
# Create mock API key entities
key1 = Mock(spec=ApiKey)
key1.id = 1
key1.name = 'Test Key 1'
key1.key = 'lbk_test_key_1'
key1.description = 'First test key'
key2 = Mock(spec=ApiKey)
key2.id = 2
key2.name = 'Test Key 2'
key2.key = 'lbk_test_key_2'
key2.description = 'Second test key'
mock_result = Mock()
mock_result.all = Mock(return_value=[key1, key2])
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
ap.persistence_mgr.serialize_model = Mock(
side_effect=lambda model_cls, entity: {
'id': entity.id,
'name': entity.name,
'key': entity.key,
'description': entity.description,
}
)
service = ApiKeyService(ap)
# Execute
result = await service.get_api_keys()
# Verify
assert len(result) == 2
assert result[0]['name'] == 'Test Key 1'
assert result[1]['name'] == 'Test Key 2'
assert [item['name'] for item in result] == ['Test Key 1', 'Test Key 2']
assert [item['description'] for item in result] == ['First test key', 'Second test key']
assert all('key' not in item and 'key_hash' not in item for item in result)
class TestApiKeyServiceCreateApiKey:
"""Tests for create_api_key method."""
async def test_create_api_key_generates_key_with_prefix(self, api_key_context):
_application, service, context, _engine = api_key_context
async def test_create_api_key_generates_key_with_prefix(self):
"""Creates API key with 'lbk_' prefix."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
with pytest.MonkeyPatch.context() as monkeypatch:
monkeypatch.setattr(
'langbot.pkg.api.http.service.apikey.secrets.token_urlsafe', lambda _size: 'fixed-token'
)
result = await service.create_api_key(context, 'New Key', 'Test description')
created_key = Mock(spec=ApiKey)
created_key.id = 1
created_key.name = 'New Key'
created_key.key = 'lbk_fixed-token'
created_key.description = 'Test description'
select_result = Mock()
select_result.first = Mock(return_value=created_key)
insert_params = []
async def mock_execute(query):
params = query.compile().params
if {'name', 'key', 'description'}.issubset(params):
insert_params.append(params)
return Mock()
return select_result
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
ap.persistence_mgr.serialize_model = Mock(
side_effect=lambda model_cls, entity: {
'id': 1,
'name': entity.name,
'key': entity.key,
'description': entity.description,
}
)
service = ApiKeyService(ap)
with patch('langbot.pkg.api.http.service.apikey.secrets.token_urlsafe', return_value='fixed-token'):
result = await service.create_api_key('New Key', 'Test description')
assert insert_params == [{'name': 'New Key', 'key': 'lbk_fixed-token', 'description': 'Test description'}]
assert result['key'].startswith('lbk_')
assert result['key'] == 'lbk_fixed-token'
assert result['name'] == 'New Key'
assert result['description'] == 'Test description'
assert result['secret_available'] is True
async def test_create_api_key_without_description(self):
"""Creates API key with empty description when not provided."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
async def test_create_api_key_without_description(self, api_key_context):
_application, service, context, _engine = api_key_context
created_key = Mock(spec=ApiKey)
created_key.id = 1
created_key.name = 'No Desc Key'
created_key.key = 'lbk_no_desc_key'
created_key.description = ''
result = await service.create_api_key(context, 'No Desc Key')
select_result = Mock()
select_result.first = Mock(return_value=created_key)
insert_result = Mock()
async def mock_execute(query):
if hasattr(query, 'values'):
return insert_result
return select_result
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
ap.persistence_mgr.serialize_model = Mock(
return_value={
'id': 1,
'name': 'No Desc Key',
'key': 'lbk_no_desc_key',
'description': '',
}
)
service = ApiKeyService(ap)
# Execute
result = await service.create_api_key('No Desc Key')
# Verify
assert result['description'] == ''
class TestApiKeyServiceGetApiKey:
"""Tests for get_api_key method."""
async def test_get_api_key_by_id_found(self, api_key_context):
_application, service, context, _engine = api_key_context
created = await service.create_api_key(context, 'Found Key', 'Found')
async def test_get_api_key_by_id_found(self):
"""Returns API key when found by ID."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
result = await service.get_api_key(context, created['id'])
key = Mock(spec=ApiKey)
key.id = 1
key.name = 'Found Key'
key.key = 'lbk_found_key'
key.description = 'Found'
mock_result = Mock()
mock_result.first = Mock(return_value=key)
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
ap.persistence_mgr.serialize_model = Mock(
return_value={
'id': 1,
'name': 'Found Key',
'key': 'lbk_found_key',
'description': 'Found',
}
)
service = ApiKeyService(ap)
# Execute
result = await service.get_api_key(1)
# Verify
assert result is not None
assert result['id'] == 1
assert result['id'] == created['id']
assert result['name'] == 'Found Key'
assert 'key' not in result and 'key_hash' not in result
async def test_get_api_key_by_id_not_found(self):
"""Returns None when API key not found."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
async def test_get_api_key_by_id_not_found(self, api_key_context):
_application, service, context, _engine = api_key_context
mock_result = Mock()
mock_result.first = Mock(return_value=None)
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
assert await service.get_api_key(context, 999) is None
service = ApiKeyService(ap)
async def test_get_api_key_by_id_zero(self, api_key_context):
_application, service, context, _engine = api_key_context
# Execute
result = await service.get_api_key(999)
# Verify
assert result is None
async def test_get_api_key_by_id_zero(self):
"""Handles ID=0 (edge case) correctly."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
mock_result = Mock()
mock_result.first = Mock(return_value=None)
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
service = ApiKeyService(ap)
# Execute
result = await service.get_api_key(0)
# Verify - should return None (no key with ID 0)
assert result is None
assert await service.get_api_key(context, 0) is None
class TestApiKeyServiceVerifyApiKey:
"""Tests for verify_api_key method."""
async def test_verify_api_key_valid(self, api_key_context):
_application, service, context, _engine = api_key_context
created = await service.create_api_key(context, 'Valid')
@staticmethod
def _make_ap(db_key=None, global_api_key=''):
"""Build a mock Application with persistence + instance_config."""
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
mock_result = Mock()
mock_result.first = Mock(return_value=db_key)
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
ap.instance_config = SimpleNamespace(data={'api': {'global_api_key': global_api_key}})
return ap
assert await service.verify_api_key(created['key']) is True
async def test_verify_api_key_valid(self):
"""Returns True for valid API key."""
# Setup
key = Mock(spec=ApiKey)
ap = self._make_ap(db_key=key)
async def test_verify_api_key_invalid(self, api_key_context):
_application, service, _context, _engine = api_key_context
service = ApiKeyService(ap)
assert await service.verify_api_key('lbk_invalid_key') is False
# Execute
result = await service.verify_api_key('lbk_valid_key')
async def test_verify_api_key_empty_string(self, api_key_context):
_application, service, _context, _engine = api_key_context
# Verify
assert result is True
assert await service.verify_api_key('') is False
async def test_verify_api_key_invalid(self):
"""Returns False for invalid API key."""
# Setup
ap = self._make_ap(db_key=None)
async def test_verify_api_key_unknown_key(self, api_key_context):
_application, service, _context, _engine = api_key_context
service = ApiKeyService(ap)
assert await service.verify_api_key('unknown_key') is False
# Execute
result = await service.verify_api_key('lbk_invalid_key')
async def test_verify_global_api_key_match(self, api_key_context):
application, service, context, _engine = api_key_context
application.instance_config.data['api']['global_api_key'] = 'my-global-secret'
# Verify
assert result is False
identity = await service.authenticate_api_key('my-global-secret')
async def test_verify_api_key_empty_string(self):
"""Returns False for empty key string."""
# Setup
ap = self._make_ap(db_key=None)
assert identity is not None
assert identity.workspace_uuid == context.workspace_uuid
assert identity.api_key_uuid == 'global-oss-api-key'
service = ApiKeyService(ap)
async def test_verify_global_api_key_no_prefix_required(self, api_key_context):
application, service, _context, _engine = api_key_context
application.instance_config.data['api']['global_api_key'] = 'plainsecret123'
# Execute
result = await service.verify_api_key('')
assert await service.verify_api_key('plainsecret123') is True
# Verify
assert result is False
async def test_verify_global_api_key_mismatch_falls_back_to_db(self, api_key_context):
application, service, context, _engine = api_key_context
application.instance_config.data['api']['global_api_key'] = 'my-global-secret'
created = await service.create_api_key(context, 'DB key')
async def test_verify_api_key_unknown_key(self):
"""Returns False when the key is not present in persistence."""
# Setup
ap = self._make_ap(db_key=None)
identity = await service.authenticate_api_key(created['key'])
service = ApiKeyService(ap)
assert identity is not None
assert identity.api_key_uuid == created['uuid']
# Execute
result = await service.verify_api_key('unknown_key')
async def test_verify_empty_global_api_key_disabled(self, api_key_context):
application, service, _context, _engine = api_key_context
application.instance_config.data['api']['global_api_key'] = ''
# Verify
assert result is False
async def test_verify_global_api_key_match(self):
"""Returns True when key matches the config.yaml global API key (no DB lookup)."""
# Setup: no DB record, but a global key is configured
ap = self._make_ap(db_key=None, global_api_key='my-global-secret')
service = ApiKeyService(ap)
# Execute
result = await service.verify_api_key('my-global-secret')
# Verify: accepted purely on config match
assert result is True
# DB should not have been consulted for the global-key path
ap.persistence_mgr.execute_async.assert_not_called()
async def test_verify_global_api_key_no_prefix_required(self):
"""Global API key is accepted even without the lbk_ prefix."""
ap = self._make_ap(db_key=None, global_api_key='plainsecret123')
service = ApiKeyService(ap)
result = await service.verify_api_key('plainsecret123')
assert result is True
async def test_verify_global_api_key_mismatch_falls_back_to_db(self):
"""A non-matching key still falls through to the DB lookup."""
# Global key set, but request uses a different lbk_ key that IS in DB
key = Mock(spec=ApiKey)
ap = self._make_ap(db_key=key, global_api_key='my-global-secret')
service = ApiKeyService(ap)
result = await service.verify_api_key('lbk_db_key')
assert result is True
ap.persistence_mgr.execute_async.assert_called_once()
async def test_verify_empty_global_api_key_disabled(self):
"""An empty global_api_key must never authenticate an empty/blank request."""
ap = self._make_ap(db_key=None, global_api_key='')
service = ApiKeyService(ap)
# Empty request key is rejected, and a blank global key never matches
assert await service.verify_api_key('') is False
assert await service.verify_api_key(' ') is False
async def test_verify_api_key_missing_global_config_key(self):
"""Works even when api.global_api_key is absent (existing installs)."""
# instance_config without the global_api_key field at all
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
mock_result = Mock()
mock_result.first = Mock(return_value=None)
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
ap.instance_config = SimpleNamespace(data={'api': {}})
async def test_verify_api_key_missing_global_config_key(self, api_key_context):
application, service, _context, _engine = api_key_context
application.instance_config.data = {'api': {}}
service = ApiKeyService(ap)
result = await service.verify_api_key('lbk_some_key')
assert result is False
assert await service.verify_api_key('lbk_some_key') is False
class TestApiKeyServiceDeleteApiKey:
"""Tests for delete_api_key method."""
async def test_delete_api_key_by_id(self, api_key_context):
_application, service, context, _engine = api_key_context
created = await service.create_api_key(context, 'Delete me')
async def test_delete_api_key_by_id(self):
"""Deletes API key by ID."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
await service.delete_api_key(context, created['id'])
service = ApiKeyService(ap)
stored = await service.get_api_key(context, created['id'])
assert stored is not None
assert stored['status'] == 'revoked'
assert await service.verify_api_key(created['key']) is False
# Execute
await service.delete_api_key(1)
async def test_delete_api_key_nonexistent_id(self, api_key_context):
_application, service, context, _engine = api_key_context
# Verify - execute_async was called (delete operation)
ap.persistence_mgr.execute_async.assert_called_once()
async def test_delete_api_key_nonexistent_id(self):
"""Delete operation completes even for nonexistent ID (no error raised)."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
service = ApiKeyService(ap)
# Execute - should not raise error
await service.delete_api_key(999)
# Verify - execute_async was called regardless
ap.persistence_mgr.execute_async.assert_called_once()
with pytest.raises(WorkspaceNotFoundError, match='API key not found'):
await service.delete_api_key(context, 999)
class TestApiKeyServiceUpdateApiKey:
"""Tests for update_api_key method."""
async def test_update_api_key_name_only(self, api_key_context):
_application, service, context, _engine = api_key_context
created = await service.create_api_key(context, 'Original', 'Description')
async def test_update_api_key_name_only(self):
"""Updates only the name field."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
await service.update_api_key(context, created['id'], name='Updated Name')
service = ApiKeyService(ap)
stored = await service.get_api_key(context, created['id'])
assert stored is not None
assert stored['name'] == 'Updated Name'
assert stored['description'] == 'Description'
# Execute
await service.update_api_key(1, name='Updated Name')
async def test_update_api_key_description_only(self, api_key_context):
_application, service, context, _engine = api_key_context
created = await service.create_api_key(context, 'Original', 'Description')
# Verify - execute_async was called with update
ap.persistence_mgr.execute_async.assert_called_once()
await service.update_api_key(context, created['id'], description='Updated description')
async def test_update_api_key_description_only(self):
"""Updates only the description field."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
stored = await service.get_api_key(context, created['id'])
assert stored is not None
assert stored['name'] == 'Original'
assert stored['description'] == 'Updated description'
service = ApiKeyService(ap)
async def test_update_api_key_both_fields(self, api_key_context):
_application, service, context, _engine = api_key_context
created = await service.create_api_key(context, 'Original', 'Description')
# Execute
await service.update_api_key(1, description='Updated description')
await service.update_api_key(
context,
created['id'],
name='New Name',
description='New description',
)
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
stored = await service.get_api_key(context, created['id'])
assert stored is not None
assert stored['name'] == 'New Name'
assert stored['description'] == 'New description'
async def test_update_api_key_both_fields(self):
"""Updates both name and description."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
async def test_update_api_key_no_fields(self, api_key_context):
application, service, context, _engine = api_key_context
created = await service.create_api_key(context, 'Original')
original_execute = application.persistence_mgr.execute_async
application.persistence_mgr.execute_async = AsyncMock(wraps=original_execute)
service = ApiKeyService(ap)
await service.update_api_key(context, created['id'])
# Execute
await service.update_api_key(1, name='New Name', description='New description')
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
async def test_update_api_key_no_fields(self):
"""Does nothing when no fields provided."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
service = ApiKeyService(ap)
# Execute
await service.update_api_key(1)
# Verify - no execute call since no update_data
ap.persistence_mgr.execute_async.assert_not_called()
application.persistence_mgr.execute_async.assert_not_awaited()
File diff suppressed because it is too large Load Diff
@@ -1,389 +1,581 @@
"""Unit tests for API knowledge service.
Tests cover:
- Knowledge base CRUD operations
- Capability checking
- Knowledge engine discovery
- File operations
"""
"""Tests for the tenant-aware knowledge service facade."""
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
from unittest.mock import Mock, AsyncMock
from importlib import import_module
from langbot.pkg.api.http.authz import WorkspaceRequiredError
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.api.http.service.knowledge import KnowledgeService
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
def get_knowledge_service_module():
"""Lazy import to avoid circular import issues."""
return import_module('langbot.pkg.api.http.service.knowledge')
CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=2,
)
def create_mock_app():
"""Create mock Application for testing."""
mock_app = Mock()
mock_app.logger = Mock()
mock_app.rag_mgr = AsyncMock()
mock_app.persistence_mgr = AsyncMock()
mock_app.persistence_mgr.execute_async = AsyncMock()
mock_app.persistence_mgr.serialize_model = Mock(return_value={})
mock_app.plugin_connector = AsyncMock()
mock_app.plugin_connector.is_enable_plugin = True
return mock_app
class _Rows:
def __init__(self, rows=()):
self.rows = list(rows)
def all(self):
return self.rows
def __iter__(self):
return iter(self.rows)
def _app():
return SimpleNamespace(
logger=Mock(),
instance_config=SimpleNamespace(data={}),
rag_mgr=SimpleNamespace(
get_all_knowledge_base_details=AsyncMock(return_value=[]),
get_knowledge_base_details=AsyncMock(return_value=None),
create_knowledge_base=AsyncMock(),
remove_knowledge_base_from_runtime=AsyncMock(),
load_knowledge_base=AsyncMock(),
get_knowledge_base_by_uuid=AsyncMock(return_value=None),
delete_knowledge_base=AsyncMock(),
),
persistence_mgr=SimpleNamespace(
execute_async=AsyncMock(return_value=_Rows()),
serialize_model=Mock(return_value={}),
),
plugin_connector=SimpleNamespace(
is_enable_plugin=True,
require_workspace_context=AsyncMock(side_effect=lambda context: context),
get_rag_creation_schema=AsyncMock(return_value={}),
get_rag_retrieval_schema=AsyncMock(return_value={}),
list_knowledge_engines=AsyncMock(return_value=[]),
list_parsers=AsyncMock(return_value=[]),
),
)
@pytest.mark.asyncio
async def test_list_and_get_forward_explicit_context():
app = _app()
app.rag_mgr.get_all_knowledge_base_details.return_value = [{'uuid': 'kb-a'}]
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb-a'}
service = KnowledgeService(app)
assert await service.get_knowledge_bases(CONTEXT) == [{'uuid': 'kb-a'}]
assert await service.get_knowledge_base(CONTEXT, 'kb-a') == {'uuid': 'kb-a'}
app.rag_mgr.get_all_knowledge_base_details.assert_awaited_once_with(CONTEXT)
app.rag_mgr.get_knowledge_base_details.assert_awaited_once_with(CONTEXT, 'kb-a')
@pytest.mark.asyncio
async def test_none_context_fails_closed_before_plugin_or_manager_access():
app = _app()
service = KnowledgeService(app)
with pytest.raises(WorkspaceRequiredError):
await service.get_knowledge_bases(None)
with pytest.raises(WorkspaceRequiredError):
await service.create_knowledge_base(None, {'knowledge_engine_plugin_id': 'author/engine'})
app.plugin_connector.get_rag_creation_schema.assert_not_awaited()
app.rag_mgr.create_knowledge_base.assert_not_awaited()
@pytest.mark.asyncio
async def test_create_validates_schema_and_binds_context():
app = _app()
app.plugin_connector.get_rag_creation_schema.return_value = {
'schema': [{'name': 'endpoint', 'label': {'en_US': 'Endpoint'}, 'required': True}]
}
app.rag_mgr.create_knowledge_base.return_value = SimpleNamespace(uuid='kb-created')
service = KnowledgeService(app)
with pytest.raises(ValueError, match='Endpoint is required'):
await service.create_knowledge_base(
CONTEXT,
{'knowledge_engine_plugin_id': 'author/engine'},
)
result = await service.create_knowledge_base(
CONTEXT,
{
'name': 'KB',
'description': 'desc',
'knowledge_engine_plugin_id': 'author/engine',
'creation_settings': {'endpoint': 'https://example.invalid'},
},
)
assert result == 'kb-created'
app.rag_mgr.create_knowledge_base.assert_awaited_once_with(
CONTEXT,
name='KB',
knowledge_engine_plugin_id='author/engine',
creation_settings={'endpoint': 'https://example.invalid'},
retrieval_settings={},
description='desc',
)
@pytest.mark.asyncio
async def test_create_enforces_workspace_knowledge_base_limit():
app = _app()
app.instance_config.data = {'system': {'limitation': {'max_knowledge_bases': 2}}}
app.rag_mgr.get_all_knowledge_base_details.return_value = [{'uuid': 'kb-a'}, {'uuid': 'kb-b'}]
service = KnowledgeService(app)
with pytest.raises(ValueError, match=r'Maximum number of knowledge bases \(2\) reached'):
await service.create_knowledge_base(
CONTEXT,
{'knowledge_engine_plugin_id': 'author/engine'},
)
app.rag_mgr.create_knowledge_base.assert_not_awaited()
@pytest.mark.asyncio
async def test_update_rejects_guessed_uuid_and_scopes_reload():
app = _app()
service = KnowledgeService(app)
with pytest.raises(WorkspaceNotFoundError):
await service.update_knowledge_base(CONTEXT, 'kb-other', {'name': 'stolen'})
app.persistence_mgr.execute_async.assert_not_awaited()
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb-a', 'workspace_uuid': 'workspace-a'}
await service.update_knowledge_base(CONTEXT, 'kb-a', {'name': 'updated', 'uuid': 'ignored'})
app.rag_mgr.remove_knowledge_base_from_runtime.assert_awaited_once_with(CONTEXT, 'kb-a')
app.rag_mgr.load_knowledge_base.assert_awaited_once_with(
CONTEXT,
{'uuid': 'kb-a', 'workspace_uuid': 'workspace-a'},
)
@pytest.mark.asyncio
async def test_runtime_retrieve_uses_execution_context():
app = _app()
entry = SimpleNamespace(model_dump=Mock(return_value={'id': 'entry-a'}))
runtime_kb = SimpleNamespace(retrieve=AsyncMock(return_value=[entry]))
app.rag_mgr.get_knowledge_base_by_uuid.return_value = runtime_kb
service = KnowledgeService(app)
assert await service.retrieve_knowledge_base(CONTEXT, 'kb-a', 'query', {'top_k': 3}) == [{'id': 'entry-a'}]
runtime_kb.retrieve.assert_awaited_once_with(CONTEXT, 'query', settings={'top_k': 3})
@pytest.mark.asyncio
async def test_runtime_retrieve_cross_workspace_uuid_is_not_found():
app = _app()
service = KnowledgeService(app)
with pytest.raises(WorkspaceNotFoundError):
await service.retrieve_knowledge_base(CONTEXT, 'kb-other', 'query')
@pytest.mark.asyncio
async def test_file_listing_checks_parent_knowledge_base_first():
app = _app()
service = KnowledgeService(app)
with pytest.raises(WorkspaceNotFoundError):
await service.get_files_by_knowledge_base(CONTEXT, 'kb-other')
app.persistence_mgr.execute_async.assert_not_awaited()
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb-a'}
row = SimpleNamespace(uuid='file-a')
app.persistence_mgr.execute_async.return_value = _Rows([row])
app.persistence_mgr.serialize_model.return_value = {'uuid': 'file-a'}
assert await service.get_files_by_knowledge_base(CONTEXT, 'kb-a') == [{'uuid': 'file-a'}]
@pytest.mark.asyncio
async def test_store_and_delete_file_require_runtime_parent_and_capability():
app = _app()
runtime_kb = SimpleNamespace(
store_file=AsyncMock(return_value='task-a'),
delete_file=AsyncMock(),
)
app.rag_mgr.get_knowledge_base_by_uuid.return_value = runtime_kb
app.rag_mgr.get_knowledge_base_details.return_value = {'knowledge_engine': {'capabilities': ['doc_ingestion']}}
service = KnowledgeService(app)
assert await service.store_file(CONTEXT, 'kb-a', 'upload.pdf', 'author/parser') == 'task-a'
runtime_kb.store_file.assert_awaited_once_with(CONTEXT, 'upload.pdf', parser_plugin_id='author/parser')
await service.delete_file(CONTEXT, 'kb-a', 'file-a')
runtime_kb.delete_file.assert_awaited_once_with(CONTEXT, 'file-a')
@pytest.mark.asyncio
async def test_delete_knowledge_base_rejects_cross_workspace_uuid():
app = _app()
service = KnowledgeService(app)
with pytest.raises(WorkspaceNotFoundError):
await service.delete_knowledge_base(CONTEXT, 'kb-other')
app.rag_mgr.delete_knowledge_base.assert_not_awaited()
@pytest.mark.asyncio
async def test_engine_and_parser_discovery_require_context_and_filter_results():
app = _app()
app.plugin_connector.list_knowledge_engines.return_value = [{'plugin_id': 'author/engine'}]
app.plugin_connector.list_parsers.return_value = [
{'id': 'text', 'supported_mime_types': ['text/plain']},
{'id': 'pdf', 'supported_mime_types': ['application/pdf']},
]
service = KnowledgeService(app)
assert await service.list_knowledge_engines(CONTEXT) == [{'plugin_id': 'author/engine'}]
assert await service.list_parsers(CONTEXT, 'application/pdf') == [
{'id': 'pdf', 'supported_mime_types': ['application/pdf']}
]
with pytest.raises(WorkspaceRequiredError):
await service.list_parsers(None)
@pytest.mark.asyncio
async def test_engine_discovery_rejects_connector_workspace_or_generation_mismatch():
app = _app()
app.plugin_connector.require_workspace_context.side_effect = WorkspaceNotFoundError('Plugin resource not found')
service = KnowledgeService(app)
with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
await service.list_knowledge_engines(CONTEXT)
app.plugin_connector.list_knowledge_engines.assert_not_awaited()
@pytest.mark.asyncio
async def test_schema_validation_refences_before_second_runtime_call():
app = _app()
app.plugin_connector.require_workspace_context.side_effect = [
CONTEXT,
WorkspaceNotFoundError('Plugin resource not found'),
]
service = KnowledgeService(app)
with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
await service.create_knowledge_base(
CONTEXT,
{'knowledge_engine_plugin_id': 'author/engine'},
)
app.plugin_connector.get_rag_creation_schema.assert_awaited_once()
app.plugin_connector.get_rag_retrieval_schema.assert_not_awaited()
app.rag_mgr.create_knowledge_base.assert_not_awaited()
@pytest.mark.asyncio
async def test_engine_schemas_are_context_gated_and_fail_soft_on_connector_error():
app = _app()
app.plugin_connector.get_rag_creation_schema.return_value = {'schema': ['creation']}
app.plugin_connector.get_rag_retrieval_schema.side_effect = RuntimeError('offline')
service = KnowledgeService(app)
assert await service.get_engine_creation_schema(CONTEXT, 'author/engine') == {'schema': ['creation']}
assert await service.get_engine_retrieval_schema(CONTEXT, 'author/engine') == {}
with pytest.raises(WorkspaceRequiredError):
await service.get_engine_creation_schema(None, 'author/engine')
# Preserve the original service regression matrix with the new explicit
# Workspace context. These intentionally overlap a few isolation-focused
# tests above so legacy business behavior cannot disappear behind new guards.
class TestKnowledgeServiceInit:
"""Tests for KnowledgeService initialization."""
def test_init_stores_app_reference(self):
"""Test that __init__ stores Application reference."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
app = _app()
service = knowledge_module.KnowledgeService(mock_app)
service = KnowledgeService(app)
assert service.ap is mock_app
assert service.ap is app
class TestGetKnowledgeBases:
"""Tests for get_knowledge_bases method."""
@pytest.mark.asyncio
async def test_returns_all_kb_details(self):
"""Test that it returns all knowledge base details."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.rag_mgr.get_all_knowledge_base_details = AsyncMock(return_value=[{'uuid': 'kb1', 'name': 'KB1'}])
app = _app()
app.rag_mgr.get_all_knowledge_base_details.return_value = [{'uuid': 'kb1', 'name': 'KB1'}]
service = knowledge_module.KnowledgeService(mock_app)
result = await service.get_knowledge_bases()
result = await KnowledgeService(app).get_knowledge_bases(CONTEXT)
assert len(result) == 1
assert result[0]['uuid'] == 'kb1'
assert result == [{'uuid': 'kb1', 'name': 'KB1'}]
app.rag_mgr.get_all_knowledge_base_details.assert_awaited_once_with(CONTEXT)
@pytest.mark.asyncio
async def test_returns_empty_list_when_no_kbs(self):
"""Test that it returns empty list when no knowledge bases."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.rag_mgr.get_all_knowledge_base_details = AsyncMock(return_value=[])
app = _app()
service = knowledge_module.KnowledgeService(mock_app)
result = await service.get_knowledge_bases()
assert result == []
assert await KnowledgeService(app).get_knowledge_bases(CONTEXT) == []
class TestGetKnowledgeBase:
"""Tests for get_knowledge_base method."""
@pytest.mark.asyncio
async def test_returns_kb_details_by_uuid(self):
"""Test that it returns specific KB details."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value={'uuid': 'kb1', 'name': 'KB1'})
app = _app()
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb1', 'name': 'KB1'}
service = knowledge_module.KnowledgeService(mock_app)
result = await service.get_knowledge_base('kb1')
result = await KnowledgeService(app).get_knowledge_base(CONTEXT, 'kb1')
assert result['uuid'] == 'kb1'
assert result == {'uuid': 'kb1', 'name': 'KB1'}
@pytest.mark.asyncio
async def test_returns_none_when_not_found(self):
"""Test that it returns None when KB not found."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value=None)
app = _app()
service = knowledge_module.KnowledgeService(mock_app)
result = await service.get_knowledge_base('nonexistent')
assert result is None
assert await KnowledgeService(app).get_knowledge_base(CONTEXT, 'nonexistent') is None
class TestCreateKnowledgeBase:
"""Tests for create_knowledge_base method."""
@pytest.mark.asyncio
async def test_creates_kb_with_required_fields(self):
"""Test creating KB with required plugin ID."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_kb = Mock()
mock_kb.uuid = 'new_kb_uuid'
mock_app.rag_mgr.create_knowledge_base = AsyncMock(return_value=mock_kb)
service = knowledge_module.KnowledgeService(mock_app)
app = _app()
app.rag_mgr.create_knowledge_base.return_value = SimpleNamespace(uuid='new_kb_uuid')
service = KnowledgeService(app)
kb_data = {
'name': 'Test KB',
'knowledge_engine_plugin_id': 'author/engine',
'description': 'Test description',
}
result = await service.create_knowledge_base(kb_data)
result = await service.create_knowledge_base(CONTEXT, kb_data)
assert result == 'new_kb_uuid'
mock_app.rag_mgr.create_knowledge_base.assert_called_once()
app.rag_mgr.create_knowledge_base.assert_awaited_once_with(
CONTEXT,
name='Test KB',
knowledge_engine_plugin_id='author/engine',
creation_settings={},
retrieval_settings={},
description='Test description',
)
@pytest.mark.asyncio
async def test_raises_when_missing_plugin_id(self):
"""Test that ValueError is raised when plugin ID missing."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
app = _app()
service = knowledge_module.KnowledgeService(mock_app)
with pytest.raises(ValueError, match='knowledge_engine_plugin_id is required'):
await KnowledgeService(app).create_knowledge_base(CONTEXT, {'name': 'Test'})
with pytest.raises(ValueError) as exc_info:
await service.create_knowledge_base({'name': 'Test'})
assert 'knowledge_engine_plugin_id is required' in str(exc_info.value)
app.rag_mgr.create_knowledge_base.assert_not_awaited()
@pytest.mark.asyncio
async def test_creates_with_default_name(self):
"""Test that KB is created with default name if not provided."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_kb = Mock()
mock_kb.uuid = 'new_kb_uuid'
mock_app.rag_mgr.create_knowledge_base = AsyncMock(return_value=mock_kb)
app = _app()
app.rag_mgr.create_knowledge_base.return_value = SimpleNamespace(uuid='new_kb_uuid')
service = knowledge_module.KnowledgeService(mock_app)
await KnowledgeService(app).create_knowledge_base(
CONTEXT,
{'knowledge_engine_plugin_id': 'author/engine'},
)
await service.create_knowledge_base({'knowledge_engine_plugin_id': 'author/engine'})
# Check that default name 'Untitled' was used
call_args = mock_app.rag_mgr.create_knowledge_base.call_args
assert call_args.kwargs['name'] == 'Untitled'
assert app.rag_mgr.create_knowledge_base.await_args.kwargs['name'] == 'Untitled'
class TestUpdateKnowledgeBase:
"""Tests for update_knowledge_base method."""
@pytest.mark.asyncio
async def test_updates_mutable_fields_only(self):
"""Test that only mutable fields are updated."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value={'uuid': 'kb1', 'name': 'Updated'})
mock_app.rag_mgr.remove_knowledge_base_from_runtime = AsyncMock()
mock_app.rag_mgr.load_knowledge_base = AsyncMock()
app = _app()
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb1', 'name': 'Updated'}
service = KnowledgeService(app)
service = knowledge_module.KnowledgeService(mock_app)
# Pass both mutable and immutable fields
await service.update_knowledge_base(
CONTEXT,
'kb1',
{
'name': 'New Name',
'description': 'New desc',
'uuid': 'should_be_filtered', # immutable
'uuid': 'should_be_filtered',
},
)
# Check that only mutable fields were passed to update
call_args = mock_app.persistence_mgr.execute_async.call_args
assert call_args is not None
update_statement = app.persistence_mgr.execute_async.await_args_list[0].args[0]
params = update_statement.compile().params
assert params['name'] == 'New Name'
assert params['description'] == 'New desc'
assert 'uuid' not in params
app.rag_mgr.remove_knowledge_base_from_runtime.assert_awaited_once_with(CONTEXT, 'kb1')
@pytest.mark.asyncio
async def test_returns_early_when_no_mutable_fields(self):
"""Test that update returns early when no mutable fields provided."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
app = _app()
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb1'}
service = knowledge_module.KnowledgeService(mock_app)
await KnowledgeService(app).update_knowledge_base(
CONTEXT,
'kb1',
{'uuid': 'should_be_filtered'},
)
# Pass only immutable fields
await service.update_knowledge_base('kb1', {'uuid': 'should_be_filtered'})
# No DB update should be called
mock_app.persistence_mgr.execute_async.assert_not_called()
app.persistence_mgr.execute_async.assert_not_awaited()
app.rag_mgr.remove_knowledge_base_from_runtime.assert_not_awaited()
class TestCheckDocCapability:
"""Tests for _check_doc_capability method."""
@pytest.mark.asyncio
async def test_passes_when_capability_supported(self):
"""Test that check passes when doc_ingestion capability exists."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(
return_value={'knowledge_engine': {'capabilities': ['doc_ingestion']}}
)
app = _app()
app.rag_mgr.get_knowledge_base_details.return_value = {'knowledge_engine': {'capabilities': ['doc_ingestion']}}
service = knowledge_module.KnowledgeService(mock_app)
await service._check_doc_capability('kb1', 'document upload')
# No exception raised means success
await KnowledgeService(app)._check_doc_capability(CONTEXT, 'kb1', 'document upload')
@pytest.mark.asyncio
async def test_raises_when_kb_not_found(self):
"""Test that Exception is raised when KB not found."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value=None)
app = _app()
service = knowledge_module.KnowledgeService(mock_app)
with pytest.raises(Exception) as exc_info:
await service._check_doc_capability('nonexistent', 'test operation')
assert 'Knowledge base not found' in str(exc_info.value)
with pytest.raises(WorkspaceNotFoundError, match='Knowledge base not found'):
await KnowledgeService(app)._check_doc_capability(
CONTEXT,
'nonexistent',
'test operation',
)
@pytest.mark.asyncio
async def test_raises_when_capability_not_supported(self):
"""Test that Exception is raised when doc_ingestion not in capabilities."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(
return_value={'knowledge_engine': {'capabilities': ['other_capability']}}
)
app = _app()
app.rag_mgr.get_knowledge_base_details.return_value = {
'knowledge_engine': {'capabilities': ['other_capability']}
}
service = knowledge_module.KnowledgeService(mock_app)
with pytest.raises(Exception) as exc_info:
await service._check_doc_capability('kb1', 'document upload')
assert 'does not support document upload' in str(exc_info.value)
with pytest.raises(Exception, match='does not support document upload'):
await KnowledgeService(app)._check_doc_capability(
CONTEXT,
'kb1',
'document upload',
)
class TestListKnowledgeEngines:
"""Tests for list_knowledge_engines method."""
@pytest.mark.asyncio
async def test_returns_engines_from_plugin_connector(self):
"""Test that it returns knowledge engines from plugin connector."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.plugin_connector.list_knowledge_engines = AsyncMock(
return_value=[{'id': 'engine1', 'name': 'Engine 1'}]
)
app = _app()
app.plugin_connector.list_knowledge_engines.return_value = [{'id': 'engine1', 'name': 'Engine 1'}]
service = knowledge_module.KnowledgeService(mock_app)
result = await service.list_knowledge_engines()
result = await KnowledgeService(app).list_knowledge_engines(CONTEXT)
assert len(result) == 1
assert result[0]['id'] == 'engine1'
assert result == [{'id': 'engine1', 'name': 'Engine 1'}]
@pytest.mark.asyncio
async def test_returns_empty_when_plugin_disabled(self):
"""Test that it returns empty list when plugin disabled."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.plugin_connector.is_enable_plugin = False
app = _app()
app.plugin_connector.is_enable_plugin = False
service = knowledge_module.KnowledgeService(mock_app)
result = await service.list_knowledge_engines()
assert result == []
assert await KnowledgeService(app).list_knowledge_engines(CONTEXT) == []
app.plugin_connector.list_knowledge_engines.assert_not_awaited()
@pytest.mark.asyncio
async def test_returns_empty_on_exception(self):
"""Test that it returns empty list and logs warning on exception."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.plugin_connector.list_knowledge_engines = AsyncMock(side_effect=Exception('Connection error'))
app = _app()
app.plugin_connector.list_knowledge_engines.side_effect = RuntimeError('Connection error')
service = knowledge_module.KnowledgeService(mock_app)
result = await service.list_knowledge_engines()
assert result == []
mock_app.logger.warning.assert_called_once()
assert await KnowledgeService(app).list_knowledge_engines(CONTEXT) == []
app.logger.warning.assert_called_once()
class TestListParsers:
"""Tests for list_parsers method."""
@pytest.mark.asyncio
async def test_returns_all_parsers(self):
"""Test that it returns all parsers when no MIME type filter."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.plugin_connector.list_parsers = AsyncMock(
return_value=[
{'id': 'parser1', 'supported_mime_types': ['text/plain']},
{'id': 'parser2', 'supported_mime_types': ['application/pdf']},
]
)
app = _app()
app.plugin_connector.list_parsers.return_value = [
{'id': 'parser1', 'supported_mime_types': ['text/plain']},
{'id': 'parser2', 'supported_mime_types': ['application/pdf']},
]
service = knowledge_module.KnowledgeService(mock_app)
result = await service.list_parsers()
result = await KnowledgeService(app).list_parsers(CONTEXT)
assert len(result) == 2
@pytest.mark.asyncio
async def test_filters_by_mime_type(self):
"""Test that it filters parsers by MIME type."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.plugin_connector.list_parsers = AsyncMock(
return_value=[
{'id': 'parser1', 'supported_mime_types': ['text/plain']},
{'id': 'parser2', 'supported_mime_types': ['application/pdf']},
]
)
app = _app()
app.plugin_connector.list_parsers.return_value = [
{'id': 'parser1', 'supported_mime_types': ['text/plain']},
{'id': 'parser2', 'supported_mime_types': ['application/pdf']},
]
service = knowledge_module.KnowledgeService(mock_app)
result = await service.list_parsers(mime_type='application/pdf')
result = await KnowledgeService(app).list_parsers(CONTEXT, 'application/pdf')
assert len(result) == 1
assert result[0]['id'] == 'parser2'
assert result == [{'id': 'parser2', 'supported_mime_types': ['application/pdf']}]
@pytest.mark.asyncio
async def test_returns_empty_when_plugin_disabled(self):
"""Test that it returns empty list when plugin disabled."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.plugin_connector.is_enable_plugin = False
app = _app()
app.plugin_connector.is_enable_plugin = False
service = knowledge_module.KnowledgeService(mock_app)
result = await service.list_parsers()
assert result == []
assert await KnowledgeService(app).list_parsers(CONTEXT) == []
app.plugin_connector.list_parsers.assert_not_awaited()
class TestGetEngineSchemas:
"""Tests for get_engine_creation_schema and get_engine_retrieval_schema."""
@pytest.mark.asyncio
async def test_returns_creation_schema(self):
"""Test that it returns creation schema for engine."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.plugin_connector.get_rag_creation_schema = AsyncMock(
return_value={'properties': {'name': {'type': 'string'}}}
)
app = _app()
app.plugin_connector.get_rag_creation_schema.return_value = {'properties': {'name': {'type': 'string'}}}
service = knowledge_module.KnowledgeService(mock_app)
result = await service.get_engine_creation_schema('author/engine')
result = await KnowledgeService(app).get_engine_creation_schema(
CONTEXT,
'author/engine',
)
assert 'properties' in result
@pytest.mark.asyncio
async def test_returns_retrieval_schema(self):
"""Test that it returns retrieval schema for engine."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.plugin_connector.get_rag_retrieval_schema = AsyncMock(
return_value={'properties': {'top_k': {'type': 'integer'}}}
)
app = _app()
app.plugin_connector.get_rag_retrieval_schema.return_value = {'properties': {'top_k': {'type': 'integer'}}}
service = knowledge_module.KnowledgeService(mock_app)
result = await service.get_engine_retrieval_schema('author/engine')
result = await KnowledgeService(app).get_engine_retrieval_schema(
CONTEXT,
'author/engine',
)
assert 'properties' in result
@pytest.mark.asyncio
async def test_returns_empty_dict_on_exception(self):
"""Test that it returns empty dict and logs warning on exception."""
knowledge_module = get_knowledge_service_module()
mock_app = create_mock_app()
mock_app.plugin_connector.get_rag_creation_schema = AsyncMock(side_effect=Exception('Plugin error'))
app = _app()
app.plugin_connector.get_rag_creation_schema.side_effect = RuntimeError('Plugin error')
service = knowledge_module.KnowledgeService(mock_app)
result = await service.get_engine_creation_schema('author/engine')
result = await KnowledgeService(app).get_engine_creation_schema(
CONTEXT,
'author/engine',
)
assert result == {}
mock_app.logger.warning.assert_called_once()
app.logger.warning.assert_called_once()
class TestKnowledgeBaseSecretViews:
@pytest.mark.asyncio
async def test_creation_settings_are_redacted_for_resource_view_only(self):
app = _app()
raw = {
'uuid': 'kb-secret',
'creation_settings': {
'dify_apikey': 'dify-secret',
'headers': {'Authorization': 'Bearer secret'},
},
}
app.rag_mgr.get_all_knowledge_base_details.return_value = [raw]
service = KnowledgeService(app)
redacted = await service.get_knowledge_bases(CONTEXT)
manager_view = await service.get_knowledge_bases(CONTEXT, include_secret=True)
assert redacted[0]['creation_settings']['dify_apikey'] == '***'
assert redacted[0]['creation_settings']['headers']['Authorization'] == '***'
assert manager_view[0]['creation_settings']['dify_apikey'] == 'dify-secret'
assert raw['creation_settings']['dify_apikey'] == 'dify-secret'
@pytest.mark.asyncio
async def test_new_masked_creation_secret_is_rejected(self):
app = _app()
with pytest.raises(ValueError, match='no existing value'):
await KnowledgeService(app).create_knowledge_base(
CONTEXT,
{
'knowledge_engine_plugin_id': 'author/engine',
'creation_settings': {'dify_apikey': '***'},
},
)
app.rag_mgr.create_knowledge_base.assert_not_awaited()
@@ -13,16 +13,39 @@ Source: src/langbot/pkg/api/http/service/maintenance.py
from __future__ import annotations
import contextlib
import pytest
from unittest.mock import AsyncMock, Mock, patch, MagicMock
from types import SimpleNamespace
import datetime
from pathlib import Path
import sqlalchemy
from sqlalchemy.ext.asyncio import create_async_engine
from langbot.pkg.api.http.authz import WorkspaceRequiredError
from langbot.pkg.api.http.service.maintenance import MaintenanceService
from langbot.pkg.api.http.context import ExecutionContext, PrincipalContext, PrincipalType
from langbot.pkg.entity.persistence.base import Base
from langbot.pkg.entity.persistence.bstorage import BinaryStorage
from langbot.pkg.entity.persistence.monitoring import MonitoringMessage
from langbot.pkg.entity.persistence.workspace import Workspace
pytestmark = pytest.mark.asyncio
TEST_CONTEXT = ExecutionContext(
instance_uuid='test-instance',
workspace_uuid='test-workspace',
placement_generation=1,
)
@pytest.fixture(autouse=True)
def assume_oss_singleton(monkeypatch):
async def is_oss_singleton(_self, _context):
return True
monkeypatch.setattr(MaintenanceService, '_is_oss_singleton', is_oss_singleton)
def _create_mock_result(scalar_value=None):
@@ -32,6 +55,14 @@ def _create_mock_result(scalar_value=None):
return result
def _scoped_storage_manager():
prefix = 'instances/i/workspaces/w/generations/1/owners/upload/o/'
return SimpleNamespace(
scoped_prefix=Mock(return_value=prefix),
is_scoped_object_key=Mock(side_effect=lambda key, **_: key == f'{prefix}uploaded_file.txt'),
)
class TestMaintenanceServiceCleanupExpiredFiles:
"""Tests for cleanup_expired_files method."""
@@ -39,6 +70,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
"""Uses default retention days when config not set."""
# Setup
ap = SimpleNamespace()
ap.storage_mgr = _scoped_storage_manager()
ap.instance_config = SimpleNamespace()
ap.instance_config.data = {}
ap.storage_mgr = SimpleNamespace()
@@ -58,7 +90,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
service._cleanup_expired_log_files = Mock(return_value=0) # NOT async!
# Execute
result = await service.cleanup_expired_files()
result = await service.cleanup_expired_files(TEST_CONTEXT)
# Verify - returns counts
assert 'uploaded_files' in result
@@ -95,7 +127,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
service._cleanup_expired_log_files = Mock(return_value=3) # NOT async
# Execute
result = await service.cleanup_expired_files()
result = await service.cleanup_expired_files(TEST_CONTEXT)
# Verify
assert result['uploaded_files'] == 2
@@ -124,7 +156,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
service._cleanup_expired_log_files = Mock(return_value=0) # NOT async
# Execute
result = await service.cleanup_expired_files()
result = await service.cleanup_expired_files(TEST_CONTEXT)
# Verify
assert result['uploaded_files'] == 1
@@ -159,12 +191,57 @@ class TestMaintenanceServiceCleanupExpiredFiles:
service._cleanup_expired_log_files = Mock(return_value=0) # NOT async
# Execute
result = await service.cleanup_expired_files()
result = await service.cleanup_expired_files(TEST_CONTEXT)
# Verify - warning logged, defaults used
assert ap.logger.warning.called
assert 'uploaded_files' in result
async def test_cloud_cleanup_carries_scope_without_holding_database_session(self):
class ScopeOnlyPersistenceManager:
mode = SimpleNamespace(value='cloud_runtime')
def __init__(self):
self.active_workspace = None
@contextlib.asynccontextmanager
async def tenant_scope(self, workspace_uuid):
self.active_workspace = workspace_uuid
try:
yield
finally:
self.active_workspace = None
def current_session(self):
return None
persistence_mgr = ScopeOnlyPersistenceManager()
application = SimpleNamespace(
persistence_mgr=persistence_mgr,
instance_config=SimpleNamespace(data={}),
logger=SimpleNamespace(warning=Mock()),
)
service = MaintenanceService(application)
async def cleanup_uploads(_context, _retention_days):
assert persistence_mgr.active_workspace == TEST_CONTEXT.workspace_uuid
assert persistence_mgr.current_session() is None
return 2
def cleanup_logs(_retention_days):
assert persistence_mgr.active_workspace == TEST_CONTEXT.workspace_uuid
assert persistence_mgr.current_session() is None
return 1
service._cleanup_expired_uploaded_files = cleanup_uploads
service._cleanup_expired_log_files = cleanup_logs
assert await service.cleanup_expired_files(TEST_CONTEXT) == {
'uploaded_files': 2,
'log_files': 1,
}
assert persistence_mgr.active_workspace is None
class TestMaintenanceServiceGetStorageAnalysis:
"""Tests for get_storage_analysis method."""
@@ -196,7 +273,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
service._expired_log_candidates = Mock(return_value=[])
# Execute
result = await service.get_storage_analysis()
result = await service.get_storage_analysis(TEST_CONTEXT)
# Verify
assert 'generated_at' in result
@@ -229,7 +306,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
service._expired_log_candidates = Mock(return_value=[])
# Execute
result = await service.get_storage_analysis()
result = await service.get_storage_analysis(TEST_CONTEXT)
# Verify - all sections present
sections = {s['key'] for s in result['sections']}
@@ -265,7 +342,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
service._expired_log_candidates = Mock(return_value=[])
# Execute
result = await service.get_storage_analysis()
result = await service.get_storage_analysis(TEST_CONTEXT)
# Verify
assert result['database']['type'] == 'postgresql'
@@ -294,7 +371,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
service._expired_log_candidates = Mock(return_value=[{'name': 'old_log', 'size_bytes': 50}])
# Execute
result = await service.get_storage_analysis()
result = await service.get_storage_analysis(TEST_CONTEXT)
# Verify
assert len(result['cleanup_candidates']['uploaded_files']) == 1
@@ -316,7 +393,7 @@ class TestMaintenanceServiceMonitoringCounts:
service = MaintenanceService(ap)
# Execute
result = await service._monitoring_counts()
result = await service._monitoring_counts(TEST_CONTEXT)
# Verify - all table keys present
assert 'messages' in result
@@ -338,7 +415,7 @@ class TestMaintenanceServiceMonitoringCounts:
service = MaintenanceService(ap)
# Execute
result = await service._monitoring_counts()
result = await service._monitoring_counts(TEST_CONTEXT)
# Verify - all zero
assert all(v == 0 for v in result.values())
@@ -374,7 +451,7 @@ class TestMaintenanceServiceBinaryStorageStats:
service = MaintenanceService(ap)
# Execute
result = await service._binary_storage_stats()
result = await service._binary_storage_stats(TEST_CONTEXT)
# Verify
assert result['count'] == 10
@@ -404,7 +481,7 @@ class TestMaintenanceServiceBinaryStorageStats:
service = MaintenanceService(ap)
# Execute
result = await service._binary_storage_stats()
result = await service._binary_storage_stats(TEST_CONTEXT)
# Verify - warning logged, size_bytes None or 0
assert ap.logger.warning.called
@@ -618,11 +695,13 @@ class TestMaintenanceServiceIsUploadedFileKey:
"""Returns True for valid upload file key."""
# Setup
ap = SimpleNamespace()
ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
# Execute - simple filename without path
result = service._is_uploaded_file_key('uploaded_file.txt')
key = f'{ap.storage_mgr.scoped_prefix(TEST_CONTEXT, owner_type="upload")}uploaded_file.txt'
result = service._is_uploaded_file_key(TEST_CONTEXT, key)
# Verify
assert result is True
@@ -631,11 +710,12 @@ class TestMaintenanceServiceIsUploadedFileKey:
"""Returns False for key with path separator."""
# Setup
ap = SimpleNamespace()
ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
# Execute - key with path
result = service._is_uploaded_file_key('path/to/file.txt')
result = service._is_uploaded_file_key(TEST_CONTEXT, 'path/to/file.txt')
# Verify
assert result is False
@@ -644,11 +724,12 @@ class TestMaintenanceServiceIsUploadedFileKey:
"""Returns False for plugin config prefix."""
# Setup
ap = SimpleNamespace()
ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
# Execute - plugin config file
result = service._is_uploaded_file_key('plugin_config_some_plugin.json')
result = service._is_uploaded_file_key(TEST_CONTEXT, 'plugin_config_some_plugin.json')
# Verify
assert result is False
@@ -662,6 +743,7 @@ class TestMaintenanceServiceExpiredLogCandidates:
# Setup
ap = SimpleNamespace()
ap.logger = SimpleNamespace()
ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
@@ -748,11 +830,12 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
# Setup
ap = SimpleNamespace()
ap.logger = SimpleNamespace()
ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
with patch.object(Path, 'exists', return_value=False):
result = service._expired_local_upload_candidates(7)
result = service._expired_local_upload_candidates(TEST_CONTEXT, 7)
# Verify
assert result == []
@@ -762,12 +845,10 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
# Setup
ap = SimpleNamespace()
ap.logger = SimpleNamespace()
ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
# Mock _is_uploaded_file_key
service._is_uploaded_file_key = Mock(side_effect=lambda key: 'plugin_config_' not in key and '/' not in key)
# Create mock files - one valid, one plugin config
# Create one file and one non-file entry under the scoped upload root.
mock_entry_valid = Mock(spec=Path)
mock_entry_valid.is_file = Mock(return_value=True)
mock_entry_valid.name = 'valid_upload.txt'
@@ -775,9 +856,10 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
mock_stat.st_size = 100
mock_stat.st_mtime = 0 # Very old
mock_entry_valid.stat = Mock(return_value=mock_stat)
mock_entry_valid.relative_to = Mock(return_value=Path('scoped/valid_upload.txt'))
mock_entry_plugin = Mock(spec=Path)
mock_entry_plugin.is_file = Mock(return_value=True)
mock_entry_plugin.is_file = Mock(return_value=False)
mock_entry_plugin.name = 'plugin_config_test.json'
mock_stat2 = Mock()
mock_stat2.st_size = 200
@@ -785,23 +867,22 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
mock_entry_plugin.stat = Mock(return_value=mock_stat2)
with patch.object(Path, 'exists', return_value=True):
with patch.object(Path, 'iterdir') as mock_iterdir:
mock_iterdir.return_value = [mock_entry_valid, mock_entry_plugin]
result = service._expired_local_upload_candidates(7)
with patch.object(Path, 'rglob') as mock_rglob:
mock_rglob.return_value = [mock_entry_valid, mock_entry_plugin]
result = service._expired_local_upload_candidates(TEST_CONTEXT, 7)
# Verify - only valid upload included
assert len(result) == 1
assert result[0]['key'] == 'valid_upload.txt'
assert result[0]['key'] == 'scoped/valid_upload.txt'
def test_expired_local_upload_candidates_includes_path(self):
"""Includes path when include_paths=True."""
# Setup
ap = SimpleNamespace()
ap.logger = SimpleNamespace()
ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
service._is_uploaded_file_key = Mock(return_value=True)
mock_entry = Mock(spec=Path)
mock_entry.is_file = Mock(return_value=True)
mock_entry.name = 'old_file.txt'
@@ -810,11 +891,178 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
mock_stat.st_size = 100
mock_stat.st_mtime = 0
mock_entry.stat = Mock(return_value=mock_stat)
mock_entry.relative_to = Mock(return_value=Path('scoped/old_file.txt'))
with patch.object(Path, 'exists', return_value=True):
with patch.object(Path, 'iterdir') as mock_iterdir:
mock_iterdir.return_value = [mock_entry]
result = service._expired_local_upload_candidates(7, include_paths=True)
with patch.object(Path, 'rglob') as mock_rglob:
mock_rglob.return_value = [mock_entry]
result = service._expired_local_upload_candidates(
TEST_CONTEXT,
7,
include_paths=True,
)
# Verify - path included
assert 'path' in result[0]
def test_expired_local_upload_candidates_respects_run_limit(self):
ap = SimpleNamespace(
logger=SimpleNamespace(warning=Mock()),
storage_mgr=_scoped_storage_manager(),
instance_config=SimpleNamespace(data={'storage': {'cleanup': {'max_files_per_run': 2}}}),
)
service = MaintenanceService(ap)
entries = []
for index in range(3):
entry = Mock(spec=Path)
entry.is_file = Mock(return_value=True)
entry.stat = Mock(return_value=SimpleNamespace(st_size=100, st_mtime=0))
entry.relative_to = Mock(return_value=Path(f'scoped/old-{index}.txt'))
entries.append(entry)
with patch.object(Path, 'exists', return_value=True):
with patch.object(Path, 'rglob', return_value=entries):
result = service._expired_local_upload_candidates(TEST_CONTEXT, 7)
assert [item['key'] for item in result] == [
'scoped/old-0.txt',
'scoped/old-1.txt',
]
ap.instance_config.data['storage']['cleanup']['max_files_per_run'] = 999999
assert service._max_files_per_run() == 10000
ISOLATION_WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
ISOLATION_WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
def _tenant_context(workspace_uuid: str) -> ExecutionContext:
return ExecutionContext(
instance_uuid='instance',
workspace_uuid=workspace_uuid,
placement_generation=1,
trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
)
class _RealPersistenceManager:
def __init__(self, engine):
self.engine = engine
async def execute_async(self, *args, **kwargs):
async with self.engine.connect() as connection:
result = await connection.execute(*args, **kwargs)
await connection.commit()
return result
@pytest.fixture
async def tenant_maintenance_service(tmp_path):
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "maintenance.db"}')
async with engine.begin() as connection:
await connection.run_sync(Base.metadata.create_all)
await connection.execute(
sqlalchemy.insert(Workspace),
[
{
'uuid': ISOLATION_WORKSPACE_A,
'instance_uuid': 'instance',
'name': 'A',
'slug': 'a',
'source': 'cloud_projection',
},
{
'uuid': ISOLATION_WORKSPACE_B,
'instance_uuid': 'instance',
'name': 'B',
'slug': 'b',
'source': 'cloud_projection',
},
],
)
now = datetime.datetime.now(datetime.UTC).replace(tzinfo=None)
await connection.execute(
sqlalchemy.insert(MonitoringMessage),
[
{
'id': 'message-a',
'workspace_uuid': ISOLATION_WORKSPACE_A,
'timestamp': now,
'bot_id': 'bot',
'bot_name': 'Bot',
'pipeline_id': 'pipeline',
'pipeline_name': 'Pipeline',
'message_content': 'A',
'session_id': 'same-session',
'status': 'success',
'level': 'info',
},
{
'id': 'message-b',
'workspace_uuid': ISOLATION_WORKSPACE_B,
'timestamp': now,
'bot_id': 'bot',
'bot_name': 'Bot',
'pipeline_id': 'pipeline',
'pipeline_name': 'Pipeline',
'message_content': 'B',
'session_id': 'same-session',
'status': 'success',
'level': 'info',
},
],
)
await connection.execute(
sqlalchemy.insert(BinaryStorage),
[
{
'workspace_uuid': ISOLATION_WORKSPACE_A,
'unique_key': 'a',
'key': 'same',
'owner_type': 'plugin',
'owner': 'same',
'value': b'aaa',
},
{
'workspace_uuid': ISOLATION_WORKSPACE_B,
'unique_key': 'b',
'key': 'same',
'owner_type': 'plugin',
'owner': 'same',
'value': b'bbbbb',
},
],
)
application = SimpleNamespace(
persistence_mgr=_RealPersistenceManager(engine),
instance_config=SimpleNamespace(data={}),
logger=SimpleNamespace(warning=lambda *_: None),
)
yield MaintenanceService(application)
await engine.dispose()
async def test_cleanup_requires_execution_context(tenant_maintenance_service):
with pytest.raises(WorkspaceRequiredError):
await tenant_maintenance_service.cleanup_expired_files(None)
async def test_monitoring_counts_are_workspace_scoped(tenant_maintenance_service):
counts_a = await tenant_maintenance_service._monitoring_counts(_tenant_context(ISOLATION_WORKSPACE_A))
counts_b = await tenant_maintenance_service._monitoring_counts(_tenant_context(ISOLATION_WORKSPACE_B))
assert counts_a['messages'] == 1
assert counts_b['messages'] == 1
async def test_binary_storage_stats_are_workspace_scoped(tenant_maintenance_service):
stats_a = await tenant_maintenance_service._binary_storage_stats(_tenant_context(ISOLATION_WORKSPACE_A))
stats_b = await tenant_maintenance_service._binary_storage_stats(_tenant_context(ISOLATION_WORKSPACE_B))
assert stats_a == {'count': 1, 'size_bytes': 3}
assert stats_b == {'count': 1, 'size_bytes': 5}
async def test_path_helpers_handle_missing_paths(tenant_maintenance_service, tmp_path):
missing = tmp_path / 'missing'
assert tenant_maintenance_service._path_size(missing) == 0
assert tenant_maintenance_service._file_count(missing) == 0
+427 -69
View File
@@ -13,17 +13,62 @@ Source: src/langbot/pkg/api/http/service/mcp.py
from __future__ import annotations
import asyncio
import copy
import pytest
from unittest.mock import AsyncMock, Mock, MagicMock
from types import SimpleNamespace
import uuid
from langbot.pkg.api.http.service.mcp import MCPService
from langbot.pkg.api.http.authz import Permission
from langbot.pkg.api.http.context import (
ExecutionContext,
PrincipalContext,
PrincipalType,
RequestContext,
WorkspaceContext,
)
from langbot.pkg.api.http.service.mcp import MCPService, redact_mcp_secrets, restore_mcp_secret_placeholders
from langbot.pkg.core.taskmgr import TaskCapacityError
from langbot.pkg.entity.persistence.mcp import MCPServer
from langbot.pkg.provider.tools.loaders.mcp_policy import MCPStdioDisabledError
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
pytestmark = pytest.mark.asyncio
_CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
_VIEWER_CONTEXT = RequestContext(
instance_uuid='instance-a',
placement_generation=1,
request_id='request-a',
auth_type='user_token',
principal=PrincipalContext(
principal_type=PrincipalType.ACCOUNT,
account_uuid='account-a',
),
workspace=WorkspaceContext(
workspace_uuid='workspace-a',
membership_uuid='membership-a',
role='viewer',
permissions=frozenset({Permission.RESOURCE_VIEW.value}),
),
)
def _service(ap: SimpleNamespace) -> MCPService:
ap.workspace_service = SimpleNamespace(
get_execution_binding=AsyncMock(return_value=SimpleNamespace(instance_uuid=_CONTEXT.instance_uuid))
)
if not hasattr(ap, 'logger'):
ap.logger = Mock()
return MCPService(ap)
def _create_mock_mcp_server(
server_uuid: str = None,
@@ -42,11 +87,13 @@ def _create_mock_mcp_server(
return server
def _create_mock_result(items: list = None, first_item=None):
def _create_mock_result(items: list = None, first_item=None, *, scalar_value=0, rowcount=1):
"""Create mock result object for persistence queries."""
result = Mock()
result.all = Mock(return_value=items or [])
result.first = Mock(return_value=first_item)
result.scalar = Mock(return_value=scalar_value)
result.rowcount = rowcount
return result
@@ -64,10 +111,10 @@ class TestMCPServiceGetRuntimeInfo:
mock_session.get_runtime_info_dict = Mock(return_value={'status': 'running', 'tools': 5})
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=mock_session)
service = MCPService(ap)
service = _service(ap)
# Execute
result = await service.get_runtime_info('test-server')
result = await service.get_runtime_info(_CONTEXT, 'test-server')
# Verify
assert result is not None
@@ -81,10 +128,10 @@ class TestMCPServiceGetRuntimeInfo:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=None)
service = MCPService(ap)
service = _service(ap)
# Execute
result = await service.get_runtime_info('nonexistent-server')
result = await service.get_runtime_info(_CONTEXT, 'nonexistent-server')
# Verify
assert result is None
@@ -101,12 +148,13 @@ class TestMCPServiceResources:
return_value=[{'uri_template': 'file:///{path}', 'name': 'files'}]
)
service = MCPService(ap)
service = _service(ap)
service._require_server = AsyncMock(return_value=(_CONTEXT, {'name': 'docs'}))
result = await service.get_mcp_server_resource_templates('docs')
result = await service.get_mcp_server_resource_templates(_CONTEXT, 'docs')
assert result == [{'uri_template': 'file:///{path}', 'name': 'files'}]
ap.tool_mgr.mcp_tool_loader.get_resource_templates.assert_awaited_once_with('docs')
ap.tool_mgr.mcp_tool_loader.get_resource_templates.assert_awaited_once_with(_CONTEXT, 'docs')
async def test_read_resource_envelope_uses_ui_preview_source(self):
ap = SimpleNamespace()
@@ -121,9 +169,11 @@ class TestMCPServiceResources:
}
)
service = MCPService(ap)
service = _service(ap)
service._require_server = AsyncMock(return_value=(_CONTEXT, {'name': 'docs'}))
result = await service.read_mcp_server_resource_envelope(
_CONTEXT,
'docs',
'file:///README.md',
max_bytes=4096,
@@ -132,6 +182,7 @@ class TestMCPServiceResources:
assert result['source'] == 'ui_preview'
ap.tool_mgr.mcp_tool_loader.read_resource_envelope.assert_awaited_once_with(
_CONTEXT,
'docs',
'file:///README.md',
include_blob=True,
@@ -156,12 +207,12 @@ class TestMCPServiceGetMCPServers:
'name': entity.name,
}
)
ap.tool_mgr = None
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
service = MCPService(ap)
service = _service(ap)
# Execute
result = await service.get_mcp_servers()
result = await service.get_mcp_servers(_CONTEXT)
# Verify
assert result == []
@@ -185,12 +236,12 @@ class TestMCPServiceGetMCPServers:
'mode': entity.mode,
}
)
ap.tool_mgr = None
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
service = MCPService(ap)
service = _service(ap)
# Execute
result = await service.get_mcp_servers()
result = await service.get_mcp_servers(_CONTEXT)
# Verify
assert len(result) == 2
@@ -215,21 +266,115 @@ class TestMCPServiceGetMCPServers:
)
ap.tool_mgr = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=None)
runtime_session = SimpleNamespace(get_runtime_info_dict=Mock(return_value={'status': 'connected'}))
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=runtime_session)
service = MCPService(ap)
service.get_runtime_info = AsyncMock(return_value={'status': 'connected'})
service = _service(ap)
# Execute
result = await service.get_mcp_servers(contain_runtime_info=True)
result = await service.get_mcp_servers(_CONTEXT, contain_runtime_info=True)
# Verify - runtime info included
assert result[0]['runtime_info'] == {'status': 'connected'}
async def test_resource_view_list_and_detail_redact_secrets_without_mutating_raw_data(self):
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
server = _create_mock_mcp_server(name='Secret Server')
serialized = {
'uuid': 'secret-uuid',
'name': 'Secret Server',
'enable': True,
'extra_args': {
'url': (
'https://mcp-user:mcp-password@mcp.invalid/connect'
'?token=url-secret&transport=streamable&sig=signed-secret'
),
'headers': {
'Authorization': 'Bearer top-secret',
'X-API-Key': 'api-secret',
'Accept': 'application/json',
},
'env': {
'ACCESS_TOKEN': 'access-secret',
'TOKENIZER': 'public-model-name',
},
'credentials': {
'username': 'service-user',
'password': 'password-secret',
},
'public_key': 'public-value',
},
}
original = copy.deepcopy(serialized)
ap.persistence_mgr.execute_async = AsyncMock(
side_effect=[
_create_mock_result([server]),
_create_mock_result(first_item=server),
]
)
ap.persistence_mgr.serialize_model = Mock(return_value=serialized)
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
service = _service(ap)
listed = await service.get_mcp_servers(_VIEWER_CONTEXT)
detail = await service.get_mcp_server_by_name(_VIEWER_CONTEXT, 'Secret Server')
for response in (listed[0], detail):
assert response['extra_args']['url'] == (
'https://***@mcp.invalid/connect?token=***&transport=streamable&sig=***'
)
assert response['extra_args']['headers'] == {
'Authorization': '***',
'X-API-Key': '***',
'Accept': 'application/json',
}
assert response['extra_args']['env'] == {
'ACCESS_TOKEN': '***',
'TOKENIZER': 'public-model-name',
}
assert response['extra_args']['credentials'] == {
'username': '***',
'password': '***',
}
assert response['extra_args']['public_key'] == 'public-value'
assert serialized == original
async def test_redacted_url_roundtrip_restores_persisted_credentials(self):
persisted = {
'extra_args': {'url': 'https://mcp-user:mcp-password@mcp.invalid/connect?token=url-secret&transport=http'}
}
submitted = redact_mcp_secrets(persisted)
assert submitted['extra_args']['url'] == 'https://***@mcp.invalid/connect?token=***&transport=http'
assert restore_mcp_secret_placeholders(submitted, persisted) == persisted
class TestMCPServiceCreateMCPServer:
"""Tests for create_mcp_server method."""
async def test_create_stdio_rejected_by_independent_instance_gate(self):
ap = SimpleNamespace(
instance_config=SimpleNamespace(
data={
'mcp': {'stdio': {'enabled': False}},
'system': {'limitation': {'max_extensions': -1}},
}
),
persistence_mgr=SimpleNamespace(execute_async=AsyncMock()),
tool_mgr=None,
)
service = _service(ap)
with pytest.raises(MCPStdioDisabledError, match='disabled by instance policy'):
await service.create_mcp_server(
_CONTEXT,
{'name': 'local', 'mode': 'stdio', 'enable': True, 'extra_args': {}},
)
ap.persistence_mgr.execute_async.assert_not_awaited()
async def test_create_mcp_server_max_extensions_reached_raises(self):
"""Raises ValueError when max_extensions limit reached."""
# Setup
@@ -241,16 +386,20 @@ class TestMCPServiceCreateMCPServer:
ap.plugin_connector.list_plugins = AsyncMock(return_value=[Mock(), Mock()]) # 2 plugins
# Mock get_mcp_servers to return 0 servers (2 plugins already)
mock_result = _create_mock_result([])
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
ap.persistence_mgr.execute_async = AsyncMock(
side_effect=[
_create_mock_result(scalar_value=0),
_create_mock_result(scalar_value=2),
]
)
ap.persistence_mgr.serialize_model = Mock(return_value={})
ap.tool_mgr = None
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
service = MCPService(ap)
service = _service(ap)
# Execute & Verify - 2 plugins + new server would exceed limit
with pytest.raises(ValueError, match='Maximum number of extensions'):
await service.create_mcp_server({'name': 'New Server'})
await service.create_mcp_server(_CONTEXT, {'name': 'New Server'})
async def test_create_mcp_server_no_limit(self):
"""Creates MCP server without limit when max_extensions=-1."""
@@ -271,10 +420,10 @@ class TestMCPServiceCreateMCPServer:
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
ap.persistence_mgr.serialize_model = Mock(return_value={'uuid': 'new-uuid'})
service = MCPService(ap)
service = _service(ap)
# Execute
server_uuid = await service.create_mcp_server({'name': 'New Server'})
server_uuid = await service.create_mcp_server(_CONTEXT, {'name': 'New Server'})
# Verify
assert server_uuid is not None
@@ -293,11 +442,11 @@ class TestMCPServiceCreateMCPServer:
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result(first_item=existing_server))
ap.persistence_mgr.serialize_model = Mock(return_value={})
service = MCPService(ap)
service = _service(ap)
# Execute & Verify
with pytest.raises(ValueError, match='MCP server already exists: Existing Server'):
await service.create_mcp_server({'name': 'Existing Server'})
await service.create_mcp_server(_CONTEXT, {'name': 'Existing Server'})
async def test_create_mcp_server_loads_server(self):
"""Loads server into tool_mgr when enabled."""
@@ -330,14 +479,62 @@ class TestMCPServiceCreateMCPServer:
return_value={'uuid': 'new-uuid', 'name': 'New Server', 'enable': True}
)
service = MCPService(ap)
service = _service(ap)
# Execute
await service.create_mcp_server({'name': 'New Server', 'enable': True})
await service.create_mcp_server(_CONTEXT, {'name': 'New Server', 'enable': True})
# Verify - host_mcp_server was called
ap.tool_mgr.mcp_tool_loader.host_mcp_server.assert_called_once()
async def test_create_mcp_server_does_not_start_host_until_transaction_commits(self):
"""The Runtime must not observe a server row that can still roll back."""
gate = asyncio.get_running_loop().create_future()
class PersistenceManagerStub:
def create_after_commit_gate(self):
return gate
ap = SimpleNamespace()
ap.persistence_mgr = PersistenceManagerStub()
ap.instance_config = SimpleNamespace(data={'system': {'limitation': {'max_extensions': -1}}})
observed = []
async def host_mcp_server(context, config):
observed.append((context, config))
ap.tool_mgr = SimpleNamespace(
mcp_tool_loader=SimpleNamespace(
host_mcp_server=host_mcp_server,
_hosted_mcp_tasks=[],
)
)
server_entity = _create_mock_mcp_server(server_uuid='new-uuid', enable=True)
results = [
_create_mock_result([]),
Mock(),
_create_mock_result(first_item=server_entity),
]
ap.persistence_mgr.execute_async = AsyncMock(side_effect=results)
ap.persistence_mgr.serialize_model = Mock(
return_value={'uuid': 'new-uuid', 'name': 'New Server', 'enable': True}
)
service = _service(ap)
await service.create_mcp_server(_CONTEXT, {'name': 'New Server', 'enable': True})
await asyncio.sleep(0)
assert observed == []
gate.set_result(None)
await ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks[0]
assert observed == [
(
_CONTEXT,
{'uuid': 'new-uuid', 'name': 'New Server', 'enable': True},
)
]
async def test_create_mcp_server_disabled_no_load(self):
"""Does not load server when disabled."""
# Setup
@@ -351,10 +548,10 @@ class TestMCPServiceCreateMCPServer:
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
ap.persistence_mgr.serialize_model = Mock(return_value={'uuid': 'new-uuid'})
service = MCPService(ap)
service = _service(ap)
# Execute with enable=False
server_uuid = await service.create_mcp_server({'name': 'New Server', 'enable': False})
server_uuid = await service.create_mcp_server(_CONTEXT, {'name': 'New Server', 'enable': False})
# Verify - no tool_mgr load attempt
assert server_uuid is not None
@@ -379,13 +576,11 @@ class TestMCPServiceGetMCPServerByName:
'runtime_info': None,
}
)
ap.tool_mgr = None
service = MCPService(ap)
service.get_runtime_info = AsyncMock(return_value=None)
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
service = _service(ap)
# Execute
result = await service.get_mcp_server_by_name('Found Server')
result = await service.get_mcp_server_by_name(_CONTEXT, 'Found Server')
# Verify
assert result is not None
@@ -400,10 +595,10 @@ class TestMCPServiceGetMCPServerByName:
mock_result = _create_mock_result(first_item=None)
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
service = MCPService(ap)
service = _service(ap)
# Execute
result = await service.get_mcp_server_by_name('Nonexistent Server')
result = await service.get_mcp_server_by_name(_CONTEXT, 'Nonexistent Server')
# Verify
assert result is None
@@ -421,8 +616,10 @@ class TestMCPServiceUpdateMCPServer:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.sessions = {'Old Server': Mock()}
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=True)
old_server = _create_mock_mcp_server(name='Old Server', enable=True)
updated_server = _create_mock_mcp_server(name='Old Server', enable=False)
call_count = 0
@@ -431,14 +628,23 @@ class TestMCPServiceUpdateMCPServer:
call_count += 1
if call_count == 1:
return _create_mock_result(first_item=old_server)
return Mock() # Update
if call_count == 2:
return _create_mock_result()
return _create_mock_result(first_item=updated_server)
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
ap.persistence_mgr.serialize_model = Mock(
side_effect=lambda _model, entity: {
'uuid': 'test-uuid',
'name': entity.name,
'enable': entity.enable,
}
)
service = MCPService(ap)
service = _service(ap)
# Execute - disable server
await service.update_mcp_server('test-uuid', {'enable': False})
await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': False})
# Verify - server was removed
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once()
@@ -453,6 +659,7 @@ class TestMCPServiceUpdateMCPServer:
ap.tool_mgr.mcp_tool_loader.sessions = {}
ap.tool_mgr.mcp_tool_loader.host_mcp_server = AsyncMock()
ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks = []
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=False)
old_server = _create_mock_mcp_server(name='Old Server', enable=False)
@@ -474,10 +681,10 @@ class TestMCPServiceUpdateMCPServer:
return_value={'uuid': 'test-uuid', 'name': 'Old Server', 'enable': True}
)
service = MCPService(ap)
service = _service(ap)
# Execute - enable server
await service.update_mcp_server('test-uuid', {'enable': True})
await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': True})
# Verify - server was loaded
ap.tool_mgr.mcp_tool_loader.host_mcp_server.assert_called_once()
@@ -493,6 +700,7 @@ class TestMCPServiceUpdateMCPServer:
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
ap.tool_mgr.mcp_tool_loader.host_mcp_server = AsyncMock()
ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks = []
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=True)
old_server = _create_mock_mcp_server(name='Old Server', enable=True)
@@ -510,13 +718,13 @@ class TestMCPServiceUpdateMCPServer:
return_value={'uuid': 'test-uuid', 'name': 'Old Server', 'enable': True}
)
service = MCPService(ap)
service = _service(ap)
# Execute - update enabled server (keep enabled, update extra_args)
await service.update_mcp_server('test-uuid', {'enable': True, 'extra_args': {'new': 'args'}})
await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': True, 'extra_args': {'new': 'args'}})
# Verify - remove and reload
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with('Old Server')
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with(_CONTEXT, 'Old Server')
ap.tool_mgr.mcp_tool_loader.host_mcp_server.assert_called_once()
async def test_update_mcp_server_no_tool_mgr(self):
@@ -541,15 +749,99 @@ class TestMCPServiceUpdateMCPServer:
return Mock() # Update
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
ap.persistence_mgr.serialize_model = Mock(
return_value={
'uuid': 'test-uuid',
'name': 'Server',
'enable': True,
}
)
service = MCPService(ap)
service = _service(ap)
# Execute - should not raise
await service.update_mcp_server('test-uuid', {'name': 'New Name'})
await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': False})
# Verify - persistence was called
assert ap.persistence_mgr.execute_async.call_count >= 2
async def test_update_restores_existing_masked_secrets_and_preserves_explicit_changes(self):
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=None)
old_server = _create_mock_mcp_server(name='Server', enable=True)
old_data = {
'uuid': 'test-uuid',
'name': 'Server',
'enable': True,
'mode': 'streamable_http',
'extra_args': {
'headers': {
'Authorization': 'Bearer original-secret',
'X-API-Key': 'original-api-key',
'Cookie': 'original-cookie',
}
},
}
captured_updates = []
async def mock_execute(statement):
if not captured_updates:
captured_updates.append(None)
return _create_mock_result(first_item=old_server)
captured_updates[0] = statement
return _create_mock_result()
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
ap.persistence_mgr.serialize_model = Mock(return_value=old_data)
service = _service(ap)
await service.update_mcp_server(
_CONTEXT,
'test-uuid',
{
'extra_args': {
'headers': {
'Authorization': '***',
'X-API-Key': 'replacement-api-key',
'Cookie': '',
}
}
},
)
persisted = captured_updates[0].compile().params['extra_args']
assert persisted['headers'] == {
'Authorization': 'Bearer original-secret',
'X-API-Key': 'replacement-api-key',
'Cookie': '',
}
async def test_update_rejects_masked_secret_without_existing_value(self):
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=None)
old_server = _create_mock_mcp_server(name='Server', enable=True)
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result(first_item=old_server))
ap.persistence_mgr.serialize_model = Mock(
return_value={
'uuid': 'test-uuid',
'name': 'Server',
'enable': True,
'extra_args': {'headers': {'Accept': 'application/json'}},
}
)
service = _service(ap)
with pytest.raises(ValueError, match='Masked MCP secret has no existing value'):
await service.update_mcp_server(
_CONTEXT,
'test-uuid',
{'extra_args': {'headers': {'Authorization': '***'}}},
)
assert ap.persistence_mgr.execute_async.await_count == 1
class TestMCPServiceDeleteMCPServer:
"""Tests for delete_mcp_server method."""
@@ -563,6 +855,7 @@ class TestMCPServiceDeleteMCPServer:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.sessions = {'Server to Delete': Mock()}
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=True)
server = _create_mock_mcp_server(name='Server to Delete')
@@ -576,14 +869,21 @@ class TestMCPServiceDeleteMCPServer:
return Mock() # Delete
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
ap.persistence_mgr.serialize_model = Mock(
return_value={
'uuid': 'test-uuid',
'name': 'Server to Delete',
'enable': True,
}
)
service = MCPService(ap)
service = _service(ap)
# Execute
await service.delete_mcp_server('test-uuid')
await service.delete_mcp_server(_CONTEXT, 'test-uuid')
# Verify
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with('Server to Delete')
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with(_CONTEXT, 'Server to Delete')
ap.persistence_mgr.execute_async.assert_called()
async def test_delete_mcp_server_not_in_sessions(self):
@@ -595,6 +895,7 @@ class TestMCPServiceDeleteMCPServer:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.sessions = {} # Server not in sessions
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=False)
server = _create_mock_mcp_server(name='Not in Sessions')
@@ -608,11 +909,18 @@ class TestMCPServiceDeleteMCPServer:
return Mock()
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
ap.persistence_mgr.serialize_model = Mock(
return_value={
'uuid': 'test-uuid',
'name': 'Not in Sessions',
'enable': True,
}
)
service = MCPService(ap)
service = _service(ap)
# Execute
await service.delete_mcp_server('test-uuid')
await service.delete_mcp_server(_CONTEXT, 'test-uuid')
# Verify - remove not called (server not in sessions)
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_not_called()
@@ -626,6 +934,7 @@ class TestMCPServiceDeleteMCPServer:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.sessions = {}
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=False)
# No server found
call_count = 0
@@ -639,18 +948,35 @@ class TestMCPServiceDeleteMCPServer:
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
service = MCPService(ap)
service = _service(ap)
# Execute - should not raise
await service.delete_mcp_server('nonexistent-uuid')
with pytest.raises(WorkspaceNotFoundError, match='MCP server not found'):
await service.delete_mcp_server(_CONTEXT, 'nonexistent-uuid')
# Verify - delete was called regardless
ap.persistence_mgr.execute_async.assert_called()
assert ap.persistence_mgr.execute_async.await_count == 1
class TestMCPServiceTestMCPServer:
"""Tests for test_mcp_server method."""
async def test_transient_stdio_test_rejected_by_instance_gate(self):
ap = SimpleNamespace(
instance_config=SimpleNamespace(data={'mcp': {'stdio': {'enabled': False}}}),
tool_mgr=SimpleNamespace(mcp_tool_loader=SimpleNamespace(load_mcp_server=AsyncMock())),
task_mgr=SimpleNamespace(create_user_task=Mock()),
)
service = _service(ap)
with pytest.raises(MCPStdioDisabledError, match='disabled by instance policy'):
await service.test_mcp_server(
_CONTEXT,
'_',
{'name': 'local', 'mode': 'stdio', 'enable': True, 'extra_args': {}},
)
ap.tool_mgr.mcp_tool_loader.load_mcp_server.assert_not_awaited()
ap.task_mgr.create_user_task.assert_not_called()
async def test_test_mcp_server_existing_server(self):
"""Tests existing MCP server connection."""
# Setup
@@ -667,12 +993,18 @@ class TestMCPServiceTestMCPServer:
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=mock_session)
ap.task_mgr = SimpleNamespace()
ap.task_mgr.create_user_task = Mock(return_value=SimpleNamespace(id=123))
service = MCPService(ap)
service = _service(ap)
service._require_server = AsyncMock(return_value=(_CONTEXT, {'name': 'existing-server'}))
def create_user_task(coroutine, **_kwargs):
coroutine.close()
return SimpleNamespace(id=123)
ap.task_mgr.create_user_task = Mock(side_effect=create_user_task)
# Execute
task_id = await service.test_mcp_server('existing-server', {})
task_id = await service.test_mcp_server(_CONTEXT, 'existing-server', {})
# Verify - returns task ID
assert task_id == 123
@@ -685,11 +1017,12 @@ class TestMCPServiceTestMCPServer:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=None)
service = MCPService(ap)
service = _service(ap)
service._require_server = AsyncMock(side_effect=WorkspaceNotFoundError('MCP server not found'))
# Execute & Verify
with pytest.raises(ValueError, match='Server not found'):
await service.test_mcp_server('nonexistent-server', {})
with pytest.raises(WorkspaceNotFoundError, match='MCP server not found'):
await service.test_mcp_server(_CONTEXT, 'nonexistent-server', {})
async def test_test_mcp_server_new_server(self):
"""Tests new MCP server with underscore name."""
@@ -703,13 +1036,38 @@ class TestMCPServiceTestMCPServer:
ap.tool_mgr.mcp_tool_loader.load_mcp_server = AsyncMock(return_value=mock_session)
ap.task_mgr = SimpleNamespace()
ap.task_mgr.create_user_task = Mock(return_value=SimpleNamespace(id=456))
service = MCPService(ap)
service = _service(ap)
def create_user_task(coroutine, **_kwargs):
coroutine.close()
return SimpleNamespace(id=456)
ap.task_mgr.create_user_task = Mock(side_effect=create_user_task)
# Execute with '_' name (new server)
task_id = await service.test_mcp_server('_', {'name': 'New Server'})
task_id = await service.test_mcp_server(_CONTEXT, '_', {'name': 'New Server'})
# Verify - load_mcp_server called
ap.tool_mgr.mcp_tool_loader.load_mcp_server.assert_called_once()
assert task_id == 456
async def test_rejected_transient_test_session_is_shut_down(self):
ap = SimpleNamespace()
mock_session = MagicMock()
mock_session.shutdown = AsyncMock()
ap.tool_mgr = SimpleNamespace(
mcp_tool_loader=SimpleNamespace(load_mcp_server=AsyncMock(return_value=mock_session))
)
def reject(coroutine, **_kwargs):
coroutine.close()
raise TaskCapacityError('capacity')
ap.task_mgr = SimpleNamespace(create_user_task=Mock(side_effect=reject))
service = _service(ap)
with pytest.raises(TaskCapacityError, match='capacity'):
await service.test_mcp_server(_CONTEXT, '_', {'name': 'New Server'})
mock_session.shutdown.assert_awaited_once_with()
+203 -124
View File
@@ -13,13 +13,10 @@ Source: src/langbot/pkg/api/http/service/model.py
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
from unittest.mock import AsyncMock, Mock
from types import SimpleNamespace
from langbot.pkg.agent.runner.default_config import AgentRunnerDefaultConfigService
from langbot.pkg.agent.runner.descriptor import AgentRunnerDescriptor
from langbot.pkg.api.http.service.model import (
LLMModelsService,
EmbeddingModelsService,
@@ -28,11 +25,36 @@ from langbot.pkg.api.http.service.model import (
_runtime_model_data,
_validate_provider_supports,
)
from langbot.pkg.api.http.service import model as model_service_module
from langbot.pkg.entity.persistence.model import LLMModel, EmbeddingModel, RerankModel, ModelProvider
pytestmark = pytest.mark.asyncio
RUNNER_ID = 'plugin:test/runner/default'
WORKSPACE_UUID = 'workspace-a'
@pytest.fixture(autouse=True)
def assume_test_provider_belongs_to_workspace(monkeypatch):
"""Keep legacy runtime-focused tests isolated from the new ownership lookup."""
async def _allow_provider(_ap, _context, provider_uuid):
return {'uuid': provider_uuid}
monkeypatch.setattr(model_service_module, '_require_workspace_provider', _allow_provider)
def _existing_llm_data(provider_uuid: str = 'provider-uuid') -> dict:
return {
'uuid': 'existing-uuid',
'workspace_uuid': WORKSPACE_UUID,
'name': 'Existing Model',
'provider_uuid': provider_uuid,
'abilities': [],
'context_length': None,
'extra_args': {},
'prefered_ranking': 0,
}
def _create_mock_llm_model(
@@ -105,20 +127,33 @@ def _create_mock_result(items: list = None, first_item=None):
return result
class FakeAgentRunnerRegistry:
async def get(self, runner_id, bound_plugins=None):
return AgentRunnerDescriptor(
id=runner_id,
source='plugin',
label={'en_US': 'Test Runner'},
plugin_author='test',
plugin_name='runner',
runner_name='default',
config_schema=[
{'name': 'model', 'type': 'model-fallback-selector', 'default': {'primary': '', 'fallbacks': []}},
],
permissions={'models': ['invoke']},
)
def _create_runtime_model_mgr() -> SimpleNamespace:
"""Build a context-aware runtime-manager double for service tests."""
manager = SimpleNamespace(
provider_dict={},
llm_models=[],
embedding_models=[],
rerank_models=[],
load_llm_model_with_provider=AsyncMock(return_value=Mock()),
load_embedding_model_with_provider=AsyncMock(return_value=Mock()),
load_rerank_model_with_provider=AsyncMock(return_value=Mock()),
cache_llm_model=AsyncMock(),
cache_embedding_model=AsyncMock(),
cache_rerank_model=AsyncMock(),
remove_llm_model=AsyncMock(),
remove_embedding_model=AsyncMock(),
remove_rerank_model=AsyncMock(),
)
async def get_provider(_context, provider_uuid):
provider = manager.provider_dict.get(provider_uuid)
if provider is None:
raise ValueError(f'Model provider {provider_uuid} not found')
return provider
manager.get_provider_by_uuid = AsyncMock(side_effect=get_provider)
return manager
class TestParseProviderApiKeys:
@@ -203,7 +238,9 @@ class TestLLMModelsServiceGetLLMModels:
service = LLMModelsService(ap)
# Execute
result = await service.get_llm_models()
result = await service.get_llm_models(
WORKSPACE_UUID,
)
# Verify
assert result == []
@@ -241,7 +278,9 @@ class TestLLMModelsServiceGetLLMModels:
service = LLMModelsService(ap)
# Execute
result = await service.get_llm_models()
result = await service.get_llm_models(
WORKSPACE_UUID,
)
# Verify
assert len(result) == 1
@@ -280,7 +319,7 @@ class TestLLMModelsServiceGetLLMModels:
service = LLMModelsService(ap)
# Execute
result = await service.get_llm_models(include_secret=False)
result = await service.get_llm_models(WORKSPACE_UUID, include_secret=False)
# Verify - keys should be masked
assert result[0]['provider']['api_keys'] == ['***', '***']
@@ -322,7 +361,7 @@ class TestLLMModelsServiceGetLLMModel:
service = LLMModelsService(ap)
# Execute
result = await service.get_llm_model('found-uuid')
result = await service.get_llm_model(WORKSPACE_UUID, 'found-uuid')
# Verify
assert result is not None
@@ -341,7 +380,7 @@ class TestLLMModelsServiceGetLLMModel:
service = LLMModelsService(ap)
# Execute
result = await service.get_llm_model('nonexistent-uuid')
result = await service.get_llm_model(WORKSPACE_UUID, 'nonexistent-uuid')
# Verify
assert result is None
@@ -366,7 +405,7 @@ class TestLLMModelsServiceGetLLMModelsByProvider:
service = LLMModelsService(ap)
# Execute
result = await service.get_llm_models_by_provider('target-provider')
result = await service.get_llm_models_by_provider(WORKSPACE_UUID, 'target-provider')
# Verify
assert len(result) == 2
@@ -380,7 +419,7 @@ class TestLLMModelsServiceCreateLLMModel:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {'provider-uuid': Mock()}
ap.model_mgr.llm_models = []
ap.model_mgr.load_llm_model_with_provider = AsyncMock(return_value=Mock())
@@ -394,12 +433,13 @@ class TestLLMModelsServiceCreateLLMModel:
# Execute
model_uuid = await service.create_llm_model(
WORKSPACE_UUID,
{
'name': 'New LLM',
'provider_uuid': 'provider-uuid',
'abilities': [],
'extra_args': {},
}
},
)
# Verify
@@ -411,7 +451,7 @@ class TestLLMModelsServiceCreateLLMModel:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {'provider-uuid': Mock()}
ap.model_mgr.llm_models = []
ap.model_mgr.load_llm_model_with_provider = AsyncMock(return_value=Mock())
@@ -425,6 +465,7 @@ class TestLLMModelsServiceCreateLLMModel:
# Execute
model_uuid = await service.create_llm_model(
WORKSPACE_UUID,
{
'uuid': 'preserved-uuid',
'name': 'Preserved UUID Model',
@@ -442,7 +483,7 @@ class TestLLMModelsServiceCreateLLMModel:
"""Creates LLM model with context_length outside extra_args."""
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {'provider-uuid': Mock()}
ap.model_mgr.llm_models = []
ap.model_mgr.load_llm_model_with_provider = AsyncMock(return_value=Mock())
@@ -454,6 +495,7 @@ class TestLLMModelsServiceCreateLLMModel:
service = LLMModelsService(ap)
await service.create_llm_model(
WORKSPACE_UUID,
{
'uuid': 'model-with-context',
'name': 'Context Model',
@@ -466,63 +508,17 @@ class TestLLMModelsServiceCreateLLMModel:
auto_set_to_default_pipeline=False,
)
runtime_entity = ap.model_mgr.load_llm_model_with_provider.await_args.args[0]
runtime_entity = ap.model_mgr.load_llm_model_with_provider.await_args.args[1]
assert runtime_entity.context_length == 128000
assert runtime_entity.extra_args == {'temperature': 0.2}
assert 'context_length' not in runtime_entity.extra_args
async def test_create_llm_model_auto_sets_schema_defined_default_pipeline_model(self):
"""Auto-default model selection should use runner schema, not legacy field names."""
ap = SimpleNamespace()
ap.logger = Mock()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr.provider_dict = {'provider-uuid': Mock()}
ap.model_mgr.llm_models = []
ap.model_mgr.load_llm_model_with_provider = AsyncMock(return_value=Mock())
ap.pipeline_service = SimpleNamespace(update_pipeline=AsyncMock())
ap.agent_runner_registry = FakeAgentRunnerRegistry()
ap.agent_runner_default_config_service = AgentRunnerDefaultConfigService(ap)
pipeline = SimpleNamespace(
uuid='pipeline-uuid',
config={
'ai': {
'runner': {'id': RUNNER_ID},
'runner_config': {
RUNNER_ID: {
'model': {'primary': '', 'fallbacks': []},
},
},
},
},
)
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result(first_item=pipeline))
service = LLMModelsService(ap)
model_uuid = await service.create_llm_model({
'uuid': 'new-model-uuid',
'name': 'New LLM',
'provider_uuid': 'provider-uuid',
'abilities': [],
'extra_args': {},
}, preserve_uuid=True)
assert model_uuid == 'new-model-uuid'
ap.pipeline_service.update_pipeline.assert_awaited_once()
updated_config = ap.pipeline_service.update_pipeline.await_args.args[1]['config']
assert updated_config['ai']['runner_config'][RUNNER_ID]['model'] == {
'primary': 'new-model-uuid',
'fallbacks': [],
}
async def test_create_llm_model_provider_not_found_raises_error(self):
"""Raises Exception when provider not found in runtime."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {} # Empty - no provider
mock_result = _create_mock_result([])
@@ -533,12 +529,13 @@ class TestLLMModelsServiceCreateLLMModel:
# Execute & Verify
with pytest.raises(Exception, match='provider not found'):
await service.create_llm_model(
WORKSPACE_UUID,
{
'name': 'No Provider Model',
'provider_uuid': 'nonexistent-provider',
'abilities': [],
'extra_args': {},
}
},
)
async def test_create_llm_model_with_provider_data(self):
@@ -546,7 +543,7 @@ class TestLLMModelsServiceCreateLLMModel:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {}
ap.model_mgr.llm_models = []
ap.model_mgr.load_llm_model_with_provider = AsyncMock(return_value=Mock())
@@ -566,6 +563,7 @@ class TestLLMModelsServiceCreateLLMModel:
# Execute - with provider data (no UUID)
result_uuid = await service.create_llm_model(
WORKSPACE_UUID,
{
'name': 'Model with New Provider',
'provider': {
@@ -575,7 +573,7 @@ class TestLLMModelsServiceCreateLLMModel:
},
'abilities': [],
'extra_args': {},
}
},
)
# Verify - provider_service was called and UUID generated
@@ -591,7 +589,7 @@ class TestLLMModelsServiceUpdateLLMModel:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {'provider-uuid': Mock()}
ap.model_mgr.llm_models = []
ap.model_mgr.remove_llm_model = AsyncMock()
@@ -600,9 +598,11 @@ class TestLLMModelsServiceUpdateLLMModel:
ap.persistence_mgr.execute_async = AsyncMock()
service = LLMModelsService(ap)
service.get_llm_model = AsyncMock(return_value=_existing_llm_data())
# Execute
await service.update_llm_model(
WORKSPACE_UUID,
'existing-uuid',
{
'uuid': 'should-be-removed',
@@ -612,24 +612,26 @@ class TestLLMModelsServiceUpdateLLMModel:
)
# Verify - remove and load called
ap.model_mgr.remove_llm_model.assert_called_once_with('existing-uuid')
ap.model_mgr.remove_llm_model.assert_called_once_with(WORKSPACE_UUID, 'existing-uuid')
async def test_update_llm_model_provider_not_found_raises_error(self):
"""Raises Exception when provider not found after update."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {} # Empty
ap.model_mgr.remove_llm_model = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock()
service = LLMModelsService(ap)
service.get_llm_model = AsyncMock(return_value=_existing_llm_data('nonexistent-provider'))
# Execute & Verify
with pytest.raises(Exception, match='provider not found'):
await service.update_llm_model(
WORKSPACE_UUID,
'model-uuid',
{
'name': 'Update',
@@ -641,15 +643,17 @@ class TestLLMModelsServiceUpdateLLMModel:
"""Updates runtime model with context_length outside extra_args."""
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace(execute_async=AsyncMock())
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {'provider-uuid': Mock()}
ap.model_mgr.llm_models = []
ap.model_mgr.remove_llm_model = AsyncMock()
ap.model_mgr.load_llm_model_with_provider = AsyncMock(return_value=Mock())
service = LLMModelsService(ap)
service.get_llm_model = AsyncMock(return_value=_existing_llm_data())
await service.update_llm_model(
WORKSPACE_UUID,
'existing-uuid',
{
'name': 'Updated Name',
@@ -660,7 +664,7 @@ class TestLLMModelsServiceUpdateLLMModel:
},
)
runtime_entity = ap.model_mgr.load_llm_model_with_provider.await_args.args[0]
runtime_entity = ap.model_mgr.load_llm_model_with_provider.await_args.args[1]
assert runtime_entity.uuid == 'existing-uuid'
assert runtime_entity.context_length == 64000
assert runtime_entity.extra_args == {'temperature': 0.4}
@@ -675,7 +679,7 @@ class TestLLMModelsServiceDeleteLLMModel:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.remove_llm_model = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock()
@@ -683,11 +687,11 @@ class TestLLMModelsServiceDeleteLLMModel:
service = LLMModelsService(ap)
# Execute
await service.delete_llm_model('delete-uuid')
await service.delete_llm_model(WORKSPACE_UUID, 'delete-uuid')
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
ap.model_mgr.remove_llm_model.assert_called_once_with('delete-uuid')
ap.model_mgr.remove_llm_model.assert_called_once_with(WORKSPACE_UUID, 'delete-uuid')
class TestEmbeddingModelsServiceGetEmbeddingModels:
@@ -706,7 +710,9 @@ class TestEmbeddingModelsServiceGetEmbeddingModels:
service = EmbeddingModelsService(ap)
# Execute
result = await service.get_embedding_models()
result = await service.get_embedding_models(
WORKSPACE_UUID,
)
# Verify
assert result == []
@@ -743,7 +749,9 @@ class TestEmbeddingModelsServiceGetEmbeddingModels:
service = EmbeddingModelsService(ap)
# Execute
result = await service.get_embedding_models()
result = await service.get_embedding_models(
WORKSPACE_UUID,
)
# Verify
assert len(result) == 1
@@ -783,7 +791,7 @@ class TestEmbeddingModelsServiceGetEmbeddingModel:
service = EmbeddingModelsService(ap)
# Execute
result = await service.get_embedding_model('found-embedding')
result = await service.get_embedding_model(WORKSPACE_UUID, 'found-embedding')
# Verify
assert result is not None
@@ -800,7 +808,7 @@ class TestEmbeddingModelsServiceGetEmbeddingModel:
service = EmbeddingModelsService(ap)
# Execute
result = await service.get_embedding_model('nonexistent-embedding')
result = await service.get_embedding_model(WORKSPACE_UUID, 'nonexistent-embedding')
# Verify
assert result is None
@@ -814,7 +822,7 @@ class TestEmbeddingModelsServiceCreateEmbeddingModel:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {'provider-uuid': Mock()}
ap.model_mgr.embedding_models = []
ap.model_mgr.load_embedding_model_with_provider = AsyncMock(return_value=Mock())
@@ -826,11 +834,12 @@ class TestEmbeddingModelsServiceCreateEmbeddingModel:
# Execute
model_uuid = await service.create_embedding_model(
WORKSPACE_UUID,
{
'name': 'New Embedding',
'provider_uuid': 'provider-uuid',
'extra_args': {},
}
},
)
# Verify
@@ -842,7 +851,7 @@ class TestEmbeddingModelsServiceCreateEmbeddingModel:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {} # Empty
mock_result = _create_mock_result([])
@@ -853,11 +862,12 @@ class TestEmbeddingModelsServiceCreateEmbeddingModel:
# Execute & Verify
with pytest.raises(Exception, match='provider not found'):
await service.create_embedding_model(
WORKSPACE_UUID,
{
'name': 'No Provider Embedding',
'provider_uuid': 'nonexistent',
'extra_args': {},
}
},
)
@@ -869,7 +879,7 @@ class TestEmbeddingModelsServiceDeleteEmbeddingModel:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.remove_embedding_model = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock()
@@ -877,7 +887,7 @@ class TestEmbeddingModelsServiceDeleteEmbeddingModel:
service = EmbeddingModelsService(ap)
# Execute
await service.delete_embedding_model('delete-embedding-uuid')
await service.delete_embedding_model(WORKSPACE_UUID, 'delete-embedding-uuid')
# Verify
ap.model_mgr.remove_embedding_model.assert_called_once()
@@ -898,7 +908,9 @@ class TestRerankModelsServiceGetRerankModels:
service = RerankModelsService(ap)
# Execute
result = await service.get_rerank_models()
result = await service.get_rerank_models(
WORKSPACE_UUID,
)
# Verify
assert result == []
@@ -935,7 +947,9 @@ class TestRerankModelsServiceGetRerankModels:
service = RerankModelsService(ap)
# Execute
result = await service.get_rerank_models()
result = await service.get_rerank_models(
WORKSPACE_UUID,
)
# Verify
assert len(result) == 1
@@ -975,7 +989,7 @@ class TestRerankModelsServiceGetRerankModel:
service = RerankModelsService(ap)
# Execute
result = await service.get_rerank_model('found-rerank')
result = await service.get_rerank_model(WORKSPACE_UUID, 'found-rerank')
# Verify
assert result is not None
@@ -992,7 +1006,7 @@ class TestRerankModelsServiceGetRerankModel:
service = RerankModelsService(ap)
# Execute
result = await service.get_rerank_model('nonexistent-rerank')
result = await service.get_rerank_model(WORKSPACE_UUID, 'nonexistent-rerank')
# Verify
assert result is None
@@ -1006,7 +1020,7 @@ class TestRerankModelsServiceCreateRerankModel:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {'provider-uuid': Mock()}
ap.model_mgr.rerank_models = []
ap.model_mgr.load_rerank_model_with_provider = AsyncMock(return_value=Mock())
@@ -1018,11 +1032,12 @@ class TestRerankModelsServiceCreateRerankModel:
# Execute
model_uuid = await service.create_rerank_model(
WORKSPACE_UUID,
{
'name': 'New Rerank',
'provider_uuid': 'provider-uuid',
'extra_args': {},
}
},
)
# Verify
@@ -1033,7 +1048,7 @@ class TestRerankModelsServiceCreateRerankModel:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.provider_dict = {}
mock_result = _create_mock_result([])
@@ -1044,11 +1059,12 @@ class TestRerankModelsServiceCreateRerankModel:
# Execute & Verify
with pytest.raises(Exception, match='provider not found'):
await service.create_rerank_model(
WORKSPACE_UUID,
{
'name': 'No Provider Rerank',
'provider_uuid': 'nonexistent',
'extra_args': {},
}
},
)
@@ -1060,7 +1076,7 @@ class TestRerankModelsServiceDeleteRerankModel:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.model_mgr = _create_runtime_model_mgr()
ap.model_mgr.remove_rerank_model = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock()
@@ -1068,7 +1084,7 @@ class TestRerankModelsServiceDeleteRerankModel:
service = RerankModelsService(ap)
# Execute
await service.delete_rerank_model('delete-rerank-uuid')
await service.delete_rerank_model(WORKSPACE_UUID, 'delete-rerank-uuid')
# Verify
ap.model_mgr.remove_rerank_model.assert_called_once()
@@ -1093,7 +1109,7 @@ class TestEmbeddingModelsServiceGetEmbeddingModelsByProvider:
service = EmbeddingModelsService(ap)
# Execute
result = await service.get_embedding_models_by_provider('provider-uuid')
result = await service.get_embedding_models_by_provider(WORKSPACE_UUID, 'provider-uuid')
# Verify
assert len(result) == 2
@@ -1118,7 +1134,7 @@ class TestRerankModelsServiceGetRerankModelsByProvider:
service = RerankModelsService(ap)
# Execute
result = await service.get_rerank_models_by_provider('provider-uuid')
result = await service.get_rerank_models_by_provider(WORKSPACE_UUID, 'provider-uuid')
# Verify
assert len(result) == 2
@@ -1132,39 +1148,102 @@ class TestValidateProviderSupports:
"""Build a fake ap whose model_mgr resolves a manifest with support_type."""
manifest = SimpleNamespace(spec={'support_type': support_type})
runtime_provider = SimpleNamespace(provider_entity=SimpleNamespace(requester=requester_name))
model_mgr = SimpleNamespace(
provider_dict={'p1': runtime_provider},
get_available_requester_manifest_by_name=lambda name: manifest if name == requester_name else None,
)
model_mgr = _create_runtime_model_mgr()
model_mgr.provider_dict = {'p1': runtime_provider}
model_mgr.get_available_requester_manifest_by_name = lambda name: manifest if name == requester_name else None
return SimpleNamespace(model_mgr=model_mgr)
async def test_allows_supported_type(self):
ap = self._make_ap('cohere-rerank', ['rerank'])
# Should not raise
await _validate_provider_supports(ap, 'p1', 'rerank')
await _validate_provider_supports(ap, WORKSPACE_UUID, 'p1', 'rerank')
async def test_rejects_unsupported_type(self):
ap = self._make_ap('cohere-rerank', ['rerank'])
with pytest.raises(ValueError, match='does not support llm'):
await _validate_provider_supports(ap, 'p1', 'llm')
await _validate_provider_supports(ap, WORKSPACE_UUID, 'p1', 'llm')
async def test_allows_when_support_type_missing(self):
# Manifest without support_type must not block (backward compatible)
manifest = SimpleNamespace(spec={})
runtime_provider = SimpleNamespace(provider_entity=SimpleNamespace(requester='legacy'))
model_mgr = SimpleNamespace(
provider_dict={'p1': runtime_provider},
get_available_requester_manifest_by_name=lambda name: manifest,
)
model_mgr = _create_runtime_model_mgr()
model_mgr.provider_dict = {'p1': runtime_provider}
model_mgr.get_available_requester_manifest_by_name = lambda name: manifest
ap = SimpleNamespace(model_mgr=model_mgr)
await _validate_provider_supports(ap, 'p1', 'rerank')
await _validate_provider_supports(ap, WORKSPACE_UUID, 'p1', 'rerank')
async def test_allows_when_provider_unknown(self):
ap = self._make_ap('cohere-rerank', ['rerank'])
# Unknown provider uuid -> no entry -> no block
await _validate_provider_supports(ap, 'missing', 'llm')
await _validate_provider_supports(ap, WORKSPACE_UUID, 'missing', 'llm')
async def test_degrades_when_model_mgr_incomplete(self):
# A bare ap without a usable model_mgr must not raise (defensive)
ap = SimpleNamespace(model_mgr=SimpleNamespace())
await _validate_provider_supports(ap, 'p1', 'llm')
await _validate_provider_supports(ap, WORKSPACE_UUID, 'p1', 'llm')
class TestModelSecretRoundtrip:
async def test_provider_filtered_list_redacts_extra_args_without_mutating_source(self):
model = _create_mock_llm_model(extra_args={'headers': {'Authorization': 'Bearer secret'}})
raw = {
'uuid': model.uuid,
'provider_uuid': model.provider_uuid,
'extra_args': {'headers': {'Authorization': 'Bearer secret'}},
}
ap = SimpleNamespace(
persistence_mgr=SimpleNamespace(
execute_async=AsyncMock(return_value=_create_mock_result([model])),
serialize_model=Mock(return_value=raw),
)
)
service = LLMModelsService(ap)
redacted = await service.get_llm_models_by_provider(WORKSPACE_UUID, model.provider_uuid)
unredacted = await service.get_llm_models_by_provider(
WORKSPACE_UUID,
model.provider_uuid,
include_secret=True,
)
assert redacted[0]['extra_args']['headers']['Authorization'] == '***'
assert unredacted[0]['extra_args']['headers']['Authorization'] == 'Bearer secret'
assert raw['extra_args']['headers']['Authorization'] == 'Bearer secret'
async def test_masked_extra_args_update_restores_existing_header(self):
existing = _existing_llm_data()
existing['extra_args'] = {
'headers': {'Authorization': 'Bearer secret', 'X-API-Key': 'key-secret'},
'timeout': 30,
}
runtime_provider = SimpleNamespace(provider_entity=SimpleNamespace(requester=None))
write_result = Mock(rowcount=1)
model_mgr = _create_runtime_model_mgr()
model_mgr.provider_dict = {'provider-uuid': runtime_provider}
ap = SimpleNamespace(
persistence_mgr=SimpleNamespace(execute_async=AsyncMock(return_value=write_result)),
model_mgr=model_mgr,
)
service = LLMModelsService(ap)
service.get_llm_model = AsyncMock(return_value=existing)
await service.update_llm_model(
WORKSPACE_UUID,
'existing-uuid',
{
'extra_args': {
'headers': {'Authorization': '***', 'X-API-Key': ''},
'timeout': 60,
}
},
)
statement = ap.persistence_mgr.execute_async.await_args.args[0]
stored_extra_args = next(
value.value for column, value in statement._values.items() if column.key == 'extra_args'
)
assert stored_extra_args == {
'headers': {'Authorization': 'Bearer secret', 'X-API-Key': ''},
'timeout': 60,
}
@@ -5,6 +5,7 @@ from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.api.http.service.monitoring import MonitoringService
@@ -34,6 +35,11 @@ async def test_record_tool_call_uses_full_monitoring_message_row_for_context():
service = MonitoringService(SimpleNamespace(persistence_mgr=persistence_mgr))
await service.record_tool_call(
ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
),
tool_name='exec',
tool_source='native',
duration=12,
@@ -42,6 +48,7 @@ async def test_record_tool_call_uses_full_monitoring_message_row_for_context():
insert_statement = persistence_mgr.execute_async.await_args_list[1].args[0]
values = insert_statement.compile().params
assert values['workspace_uuid'] == 'workspace-test'
assert values['bot_id'] == 'bot-1'
assert values['pipeline_id'] == 'pipeline-1'
assert values['session_id'] == 'person_1'
@@ -0,0 +1,392 @@
from __future__ import annotations
import datetime
from types import SimpleNamespace
import pytest
import sqlalchemy
from sqlalchemy.ext.asyncio import create_async_engine
from langbot.pkg.api.http.authz import WorkspaceRequiredError
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.api.http.service.monitoring import MonitoringService
from langbot.pkg.entity.persistence.base import Base
from langbot.pkg.entity.persistence.monitoring import MonitoringLLMCall, MonitoringMessage
from langbot.pkg.entity.persistence.workspace import Workspace
from langbot.pkg.persistence.mgr import PersistenceManager
pytestmark = pytest.mark.asyncio
WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
def _context(workspace_uuid: str) -> ExecutionContext:
return ExecutionContext(
instance_uuid='instance',
workspace_uuid=workspace_uuid,
placement_generation=3,
bot_uuid='same-bot',
pipeline_uuid='same-pipeline',
)
class _PersistenceManager:
def __init__(self, engine):
self.engine = engine
async def execute_async(self, *args, **kwargs):
async with self.engine.connect() as connection:
result = await connection.execute(*args, **kwargs)
await connection.commit()
return result
def get_db_engine(self):
return self.engine
@staticmethod
def serialize_model(model, data, masked_columns=None):
return {
column.name: (
getattr(data, column.name).isoformat()
if isinstance(getattr(data, column.name), datetime.datetime)
else getattr(data, column.name)
)
for column in model.__table__.columns
if column.name not in (masked_columns or [])
}
@pytest.fixture
async def service(tmp_path):
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "monitoring.db"}')
async with engine.begin() as connection:
await connection.run_sync(Base.metadata.create_all)
await connection.execute(
sqlalchemy.insert(Workspace),
[
{
'uuid': WORKSPACE_A,
'instance_uuid': 'instance',
'name': 'A',
'slug': 'a',
'source': 'cloud_projection',
},
{
'uuid': WORKSPACE_B,
'instance_uuid': 'instance',
'name': 'B',
'slug': 'b',
'source': 'cloud_projection',
},
],
)
application = SimpleNamespace(
persistence_mgr=_PersistenceManager(engine),
instance_config=SimpleNamespace(data={'database': {'use': 'sqlite'}}),
)
yield MonitoringService(application)
await engine.dispose()
async def _record_message(service, context, content):
return await service.record_message(
context,
bot_id='same-bot',
bot_name='Same Bot',
pipeline_id='same-pipeline',
pipeline_name='Same Pipeline',
message_content=content,
session_id='same-session',
)
async def test_monitoring_write_without_execution_context_fails_closed(service):
with pytest.raises(WorkspaceRequiredError):
await _record_message(service, None, 'unscoped')
async def test_same_session_and_resource_ids_do_not_collide(service):
context_a = _context(WORKSPACE_A)
context_b = _context(WORKSPACE_B)
message_a = await _record_message(service, context_a, 'tenant-a')
message_b = await _record_message(service, context_b, 'tenant-b')
await service.record_session_start(
context_a,
session_id='same-session',
bot_id='same-bot',
bot_name='Same Bot',
pipeline_id='same-pipeline',
pipeline_name='Same Pipeline',
)
await service.record_session_start(
context_b,
session_id='same-session',
bot_id='same-bot',
bot_name='Same Bot',
pipeline_id='same-pipeline',
pipeline_name='Same Pipeline',
)
messages_a, total_a = await service.get_messages(context_a)
messages_b, total_b = await service.get_messages(context_b)
assert total_a == total_b == 1
assert messages_a[0]['message_content'] == 'tenant-a'
assert messages_b[0]['message_content'] == 'tenant-b'
assert (await service.get_message_details(context_b, message_a))['found'] is False
assert (await service.get_message_details(context_a, message_b))['found'] is False
async def test_tool_call_inherits_context_from_connection_message_row(service):
context = _context(WORKSPACE_A)
message_id = await _record_message(service, context, 'tool context')
await service.record_tool_call(
context,
tool_name='search',
tool_source='native',
duration=12,
message_id=message_id,
)
tool_calls, total = await service.get_tool_calls(context)
assert total == 1
assert tool_calls[0]['bot_id'] == 'same-bot'
assert tool_calls[0]['pipeline_id'] == 'same-pipeline'
assert tool_calls[0]['session_id'] == 'same-session'
assert tool_calls[0]['message_id'] == message_id
async def test_feedback_upsert_and_cancel_are_workspace_scoped(service):
context_a = _context(WORKSPACE_A)
context_b = _context(WORKSPACE_B)
await service.record_feedback(context_a, feedback_id='same-feedback', feedback_type=1)
await service.record_feedback(context_b, feedback_id='same-feedback', feedback_type=2)
stats_a = await service.get_feedback_stats(context_a)
stats_b = await service.get_feedback_stats(context_b)
assert stats_a['total_likes'] == 1
assert stats_a['total_dislikes'] == 0
assert stats_b['total_likes'] == 0
assert stats_b['total_dislikes'] == 1
await service.record_feedback(context_a, feedback_id='same-feedback', feedback_type=3)
assert (await service.get_feedback_stats(context_a))['total_feedback'] == 0
assert (await service.get_feedback_stats(context_b))['total_feedback'] == 1
async def test_monitoring_queries_and_detail_views_are_strictly_bounded(service):
context = _context(WORKSPACE_A)
service.ap.instance_config.data['monitoring'] = {
'query_limits': {
'page_rows': 2,
'export_rows': 2,
'detail_rows': 2,
'timeseries_buckets': 2,
'max_offset': 10,
}
}
await service.record_session_start(
context,
session_id='same-session',
bot_id='same-bot',
bot_name='Same Bot',
pipeline_id='same-pipeline',
pipeline_name='Same Pipeline',
)
message_ids = [await _record_message(service, context, f'message-{index}') for index in range(4)]
for index in range(3):
await service.record_llm_call(
context,
bot_id='same-bot',
bot_name='Same Bot',
pipeline_id='same-pipeline',
pipeline_name='Same Pipeline',
session_id='same-session',
model_name='model',
input_tokens=1,
output_tokens=2,
duration=10,
message_id=message_ids[0],
)
await service.record_tool_call(
context,
tool_name=f'tool-{index}',
tool_source='native',
duration=5,
session_id='same-session',
message_id=message_ids[0],
)
await service.record_error(
context,
bot_id='same-bot',
bot_name='Same Bot',
pipeline_id='same-pipeline',
pipeline_name='Same Pipeline',
error_type='Failure',
error_message=f'error-{index}',
session_id='same-session',
message_id=message_ids[0],
)
page, total = await service.get_messages(context, limit=100000, offset=-5)
exported = await service.export_messages(context, limit=100000)
session_detail = await service.get_session_analysis(context, 'same-session')
message_detail = await service.get_message_details(context, message_ids[0])
assert total == 4
assert len(page) == 2
assert len(exported) == 2
assert session_detail['message_stats']['total'] == 4
assert session_detail['llm_stats']['total_calls'] == 3
assert session_detail['tool_stats']['total_calls'] == 3
assert len(session_detail['tool_calls']) == 2
assert len(session_detail['errors']) == 2
assert session_detail['detail_truncated'] == {
'tool_calls': True,
'errors': True,
}
assert message_detail['llm_stats']['total_calls'] == 3
assert len(message_detail['llm_calls']) == 2
assert len(message_detail['errors']) == 2
assert message_detail['detail_truncated'] == {
'llm_calls': True,
'errors': True,
}
service.ap.instance_config.data['monitoring']['query_limits'] = {
'page_rows': 999999,
'export_rows': 999999,
'detail_rows': 999999,
'timeseries_buckets': 999999,
'max_offset': 99999999,
}
assert service.normalize_page_window(999999, 99999999) == (5000, 10000000)
assert service.normalize_export_limit(999999) == 50000
assert service._detail_limit() == 10000
assert service._timeseries_bucket_limit() == 10000
async def test_token_statistics_aggregate_and_limit_groups_in_database(service):
context = _context(WORKSPACE_A)
service.ap.instance_config.data['monitoring'] = {
'query_limits': {
'page_rows': 1,
'timeseries_buckets': 2,
}
}
first_hour = datetime.datetime(2026, 7, 28, 10, 0)
rows = [
{
'id': f'llm-{index}',
'workspace_uuid': WORKSPACE_A,
'timestamp': first_hour + datetime.timedelta(hours=hour, minutes=index),
'model_name': model,
'input_tokens': input_tokens,
'output_tokens': output_tokens,
'total_tokens': input_tokens + output_tokens,
'duration': 100,
'cost': 0.01,
'status': 'success',
'bot_id': 'same-bot',
'bot_name': 'Same Bot',
'pipeline_id': 'same-pipeline',
'pipeline_name': 'Same Pipeline',
'session_id': 'same-session',
}
for index, (hour, model, input_tokens, output_tokens) in enumerate(
[
(0, 'small-model', 1, 2),
(1, 'large-model', 3, 4),
(2, 'large-model', 5, 6),
(2, 'large-model', 7, 8),
]
)
]
await service.ap.persistence_mgr.execute_async(sqlalchemy.insert(MonitoringLLMCall), rows)
stats = await service.get_token_statistics(context, bucket='hour')
assert stats['summary']['total_calls'] == 4
assert stats['summary']['total_tokens'] == 36
assert stats['by_model_truncated'] is True
assert [model['model_name'] for model in stats['by_model']] == ['large-model']
assert stats['timeseries_truncated'] is True
assert stats['timeseries'] == [
{
'bucket': '2026-07-28 11:00',
'input_tokens': 3,
'output_tokens': 4,
'total_tokens': 7,
'calls': 1,
},
{
'bucket': '2026-07-28 12:00',
'input_tokens': 12,
'output_tokens': 14,
'total_tokens': 26,
'calls': 2,
},
]
async def test_cleanup_commits_sqlite_delete_before_vacuum(tmp_path):
engine = create_async_engine(
f'sqlite+aiosqlite:///{tmp_path / "monitoring-cleanup.db"}',
connect_args={'timeout': 0.1},
)
application = SimpleNamespace(
instance_config=SimpleNamespace(data={'database': {'use': 'sqlite'}}),
)
manager = PersistenceManager(application)
manager.db = SimpleNamespace(get_engine=lambda: engine)
application.persistence_mgr = manager
try:
async with engine.begin() as connection:
await connection.run_sync(Base.metadata.create_all)
await connection.execute(
sqlalchemy.insert(Workspace).values(
uuid=WORKSPACE_A,
instance_uuid='instance',
name='A',
slug='a',
source='cloud_projection',
)
)
await connection.execute(
sqlalchemy.insert(MonitoringMessage),
[
{
'id': f'expired-message-{index}',
'workspace_uuid': WORKSPACE_A,
'timestamp': datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)
- datetime.timedelta(days=30),
'bot_id': 'bot',
'bot_name': 'Bot',
'pipeline_id': 'pipeline',
'pipeline_name': 'Pipeline',
'message_content': 'expired',
'session_id': 'session',
'status': 'success',
'level': 'info',
}
for index in range(5)
],
)
deleted = await MonitoringService(application).cleanup_expired_records(
_context(WORKSPACE_A),
retention_days=1,
batch_size=2,
max_batches_per_table=1,
)
assert deleted['monitoring_messages'] == 2
async with engine.connect() as connection:
remaining = await connection.scalar(
sqlalchemy.select(sqlalchemy.func.count()).select_from(MonitoringMessage)
)
assert remaining == 3
finally:
await engine.dispose()
@@ -21,10 +21,13 @@ import json
from langbot.pkg.api.http.service.pipeline import PipelineService, default_stage_order
from langbot.pkg.entity.persistence.pipeline import LegacyPipeline
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
pytestmark = pytest.mark.asyncio
WORKSPACE_UUID = 'workspace-a'
def _create_mock_pipeline(
pipeline_uuid: str = None,
@@ -77,7 +80,9 @@ class TestPipelineServiceGetPipelineMetadata:
service = PipelineService(ap)
# Execute
result = await service.get_pipeline_metadata()
result = await service.get_pipeline_metadata(
WORKSPACE_UUID,
)
# Verify
assert len(result) == 4
@@ -107,7 +112,9 @@ class TestPipelineServiceGetPipelines:
service = PipelineService(ap)
# Execute
result = await service.get_pipelines()
result = await service.get_pipelines(
WORKSPACE_UUID,
)
# Verify
assert result == []
@@ -133,7 +140,9 @@ class TestPipelineServiceGetPipelines:
service = PipelineService(ap)
# Execute
result = await service.get_pipelines()
result = await service.get_pipelines(
WORKSPACE_UUID,
)
# Verify
assert len(result) == 2
@@ -152,7 +161,7 @@ class TestPipelineServiceGetPipelines:
service = PipelineService(ap)
# Execute
await service.get_pipelines(sort_by='updated_at', sort_order='ASC')
await service.get_pipelines(WORKSPACE_UUID, sort_by='updated_at', sort_order='ASC')
# Verify - execute was called with sort parameters
ap.persistence_mgr.execute_async.assert_called_once()
@@ -181,7 +190,7 @@ class TestPipelineServiceGetPipeline:
service = PipelineService(ap)
# Execute
result = await service.get_pipeline('test-uuid')
result = await service.get_pipeline(WORKSPACE_UUID, 'test-uuid')
# Verify
assert result is not None
@@ -200,7 +209,7 @@ class TestPipelineServiceGetPipeline:
service = PipelineService(ap)
# Execute
result = await service.get_pipeline('nonexistent-uuid')
result = await service.get_pipeline(WORKSPACE_UUID, 'nonexistent-uuid')
# Verify
assert result is None
@@ -229,64 +238,7 @@ class TestPipelineServiceCreatePipeline:
# Execute & Verify
with pytest.raises(ValueError, match='Maximum number of pipelines'):
await service.create_pipeline({'name': 'New Pipeline'})
@pytest.mark.parametrize('invalid_preferences', [None, [], 'all', 0, False])
async def test_create_pipeline_rejects_non_object_extension_preferences(
self,
invalid_preferences,
):
service = PipelineService(SimpleNamespace())
with pytest.raises(ValueError, match='extensions_preferences must be an object'):
await service.create_pipeline(
{
'name': 'Invalid Pipeline',
'extensions_preferences': invalid_preferences,
}
)
@pytest.mark.parametrize('invalid_value', [0, None, 'false'])
async def test_create_pipeline_rejects_non_boolean_runner_security_field(self, invalid_value):
service = PipelineService(SimpleNamespace())
runner_id = 'plugin:test/runner/default'
with pytest.raises(ValueError, match='enable-all-tools.*boolean'):
await service.create_pipeline(
{
'name': 'Invalid Pipeline',
'config': {
'ai': {
'runner': {'id': runner_id},
'runner_config': {runner_id: {'enable-all-tools': invalid_value}},
}
},
}
)
@pytest.mark.parametrize('invalid_value', [0, None, 'false'])
async def test_create_pipeline_rejects_non_boolean_mcp_resource_enabled(self, invalid_value):
service = PipelineService(SimpleNamespace())
runner_id = 'plugin:test/runner/default'
with pytest.raises(ValueError, match=r'mcp-resources\[0\]\.enabled.*boolean'):
await service.create_pipeline(
{
'name': 'Invalid Pipeline',
'config': {
'ai': {
'runner': {'id': runner_id},
'runner_config': {
runner_id: {
'mcp-resources': [
{'uri': 'file:///README.md', 'enabled': invalid_value},
]
}
},
}
},
}
)
await service.create_pipeline(WORKSPACE_UUID, {'name': 'New Pipeline'})
async def test_create_pipeline_no_limit(self):
"""Creates pipeline without limit when max_pipelines=-1."""
@@ -315,7 +267,7 @@ class TestPipelineServiceCreatePipeline:
with patch(
'langbot.pkg.utils.paths.get_resource_path', return_value='templates/default-pipeline-config.json'
):
bot_uuid = await service.create_pipeline({'name': 'New Pipeline'})
bot_uuid = await service.create_pipeline(WORKSPACE_UUID, {'name': 'New Pipeline'})
# Verify
assert bot_uuid is not None
@@ -350,7 +302,7 @@ class TestPipelineServiceCreatePipeline:
with patch(
'langbot.pkg.utils.paths.get_resource_path', return_value='templates/default-pipeline-config.json'
):
await service.create_pipeline({'name': 'Default Pipeline'}, default=True)
await service.create_pipeline(WORKSPACE_UUID, {'name': 'Default Pipeline'}, default=True)
# Verify - execute was called
ap.persistence_mgr.execute_async.assert_called()
@@ -397,7 +349,7 @@ class TestPipelineServiceCreatePipeline:
with patch(
'langbot.pkg.utils.paths.get_resource_path', return_value='templates/default-pipeline-config.json'
):
await service.create_pipeline({'name': 'New Pipeline'})
await service.create_pipeline(WORKSPACE_UUID, {'name': 'New Pipeline'})
assert len(insert_params) == 1
assert insert_params[0]['extensions_preferences'] == {
@@ -410,6 +362,19 @@ class TestPipelineServiceCreatePipeline:
}
class _MockResultWithBots:
"""Helper class to mock SQLAlchemy result with iterable .all() method."""
def __init__(self, bots_list):
self._bots_list = bots_list
def all(self):
return self._bots_list
def first(self):
return self._bots_list[0] if self._bots_list else None
class TestPipelineServiceUpdatePipeline:
"""Tests for update_pipeline method."""
@@ -423,20 +388,22 @@ class TestPipelineServiceUpdatePipeline:
ap.pipeline_mgr.load_pipeline = AsyncMock()
ap.sess_mgr = SimpleNamespace()
ap.sess_mgr.session_list = []
ap.bot_service = None # No bot_service when not updating name
ap.persistence_mgr.execute_async = AsyncMock()
service = PipelineService(ap)
service.get_pipeline = AsyncMock(return_value={'uuid': 'test-uuid', 'name': 'Updated'})
# Execute with protected fields.
# Execute with protected fields - no name change, so no bot sync
pipeline_data = {
'uuid': 'should-be-removed',
'for_version': 'should-be-removed',
'stages': ['should-be-removed'],
'is_default': True,
'description': 'New description',
'description': 'New description', # Not name change, so no bot_service needed
}
await service.update_pipeline('test-uuid', pipeline_data)
await service.update_pipeline(WORKSPACE_UUID, 'test-uuid', pipeline_data)
update_params = ap.persistence_mgr.execute_async.await_args_list[0].args[0].compile().params
assert update_params['description'] == 'New description'
@@ -444,45 +411,8 @@ class TestPipelineServiceUpdatePipeline:
assert ['should-be-removed'] not in update_params.values()
assert not any(value is True for value in update_params.values())
@pytest.mark.parametrize('invalid_preferences', [None, [], 'all', 0, False])
async def test_update_pipeline_rejects_non_object_extension_preferences_before_write(
self,
invalid_preferences,
):
ap = SimpleNamespace(persistence_mgr=SimpleNamespace(execute_async=AsyncMock()))
service = PipelineService(ap)
with pytest.raises(ValueError, match='extensions_preferences must be an object'):
await service.update_pipeline(
'test-uuid',
{'extensions_preferences': invalid_preferences},
)
ap.persistence_mgr.execute_async.assert_not_awaited()
@pytest.mark.parametrize('invalid_value', [0, None, 'false'])
async def test_update_pipeline_rejects_non_boolean_runner_security_field_before_write(self, invalid_value):
ap = SimpleNamespace(persistence_mgr=SimpleNamespace(execute_async=AsyncMock()))
service = PipelineService(ap)
runner_id = 'plugin:test/runner/default'
with pytest.raises(ValueError, match='mcp-resource-agent-read-enabled.*boolean'):
await service.update_pipeline(
'test-uuid',
{
'config': {
'ai': {
'runner': {'id': runner_id},
'runner_config': {runner_id: {'mcp-resource-agent-read-enabled': invalid_value}},
}
}
},
)
ap.persistence_mgr.execute_async.assert_not_awaited()
async def test_update_pipeline_name_does_not_rewrite_bot_routes(self):
"""Bot event bindings remain independent from pipeline display names."""
async def test_update_pipeline_syncs_bot_names(self):
"""Updates bot use_pipeline_name when pipeline name changes."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
@@ -494,13 +424,44 @@ class TestPipelineServiceUpdatePipeline:
ap.bot_service = SimpleNamespace()
ap.bot_service.update_bot = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock(return_value=Mock())
# Create proper mock Bot entities with uuid attribute
mock_bot1 = Mock()
mock_bot1.uuid = 'bot-uuid-1'
mock_bot2 = Mock()
mock_bot2.uuid = 'bot-uuid-2'
# Create bot list
bot_list = [mock_bot1, mock_bot2]
# Create mock result using helper class
bot_result = _MockResultWithBots(bot_list)
# The order of calls in update_pipeline:
# 1. UPDATE (line 125) - returns Mock (no result needed)
# 2. SELECT bots (line 136) - returns bot_result with .all()
call_count = 0
async def mock_execute(query):
nonlocal call_count
call_count += 1
if call_count == 1:
# First call is the UPDATE - just return a Mock
return Mock()
elif call_count == 2:
# Second call is the SELECT bots - return proper result
return bot_result
return Mock() # Any additional calls
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
ap.persistence_mgr.serialize_model = Mock(return_value={})
service = PipelineService(ap)
service.get_pipeline = AsyncMock(return_value={'uuid': 'test-uuid', 'name': 'New Name'})
await service.update_pipeline('test-uuid', {'name': 'New Name'})
# Execute with name change
await service.update_pipeline(WORKSPACE_UUID, 'test-uuid', {'name': 'New Name'})
# Bots no longer denormalize Pipeline names in the EventBindings model.
ap.bot_service.update_bot.assert_not_awaited()
async def test_update_pipeline_clears_conversations(self):
@@ -526,7 +487,7 @@ class TestPipelineServiceUpdatePipeline:
service.get_pipeline = AsyncMock(return_value={'uuid': 'test-uuid'})
# Execute
await service.update_pipeline('test-uuid', {'description': 'Updated'})
await service.update_pipeline(WORKSPACE_UUID, 'test-uuid', {'description': 'Updated'})
# Verify - conversation was cleared
assert session.using_conversation is None
@@ -547,10 +508,10 @@ class TestPipelineServiceDeletePipeline:
service = PipelineService(ap)
# Execute
await service.delete_pipeline('test-uuid')
await service.delete_pipeline(WORKSPACE_UUID, 'test-uuid')
# Verify
ap.pipeline_mgr.remove_pipeline.assert_called_once_with('test-uuid')
ap.pipeline_mgr.remove_pipeline.assert_called_once_with(WORKSPACE_UUID, 'test-uuid')
ap.persistence_mgr.execute_async.assert_called_once()
async def test_delete_pipeline_nonexistent_uuid(self):
@@ -565,7 +526,7 @@ class TestPipelineServiceDeletePipeline:
service = PipelineService(ap)
# Execute - should not raise
await service.delete_pipeline('nonexistent-uuid')
await service.delete_pipeline(WORKSPACE_UUID, 'nonexistent-uuid')
# Verify
ap.pipeline_mgr.remove_pipeline.assert_called_once()
@@ -597,7 +558,7 @@ class TestPipelineServiceCopyPipeline:
# Execute & Verify
with pytest.raises(ValueError, match='Maximum number of pipelines'):
await service.copy_pipeline('original-uuid')
await service.copy_pipeline(WORKSPACE_UUID, 'original-uuid')
async def test_copy_pipeline_not_found_raises(self):
"""Raises ValueError when original pipeline not found."""
@@ -618,8 +579,8 @@ class TestPipelineServiceCopyPipeline:
ap.persistence_mgr.serialize_model = Mock(return_value={})
# Execute & Verify
with pytest.raises(ValueError, match='Pipeline original-uuid not found'):
await service.copy_pipeline('original-uuid')
with pytest.raises(WorkspaceNotFoundError, match='Pipeline original-uuid not found'):
await service.copy_pipeline(WORKSPACE_UUID, 'original-uuid')
async def test_copy_pipeline_creates_copy(self):
"""Creates a copy with (Copy) suffix."""
@@ -662,7 +623,7 @@ class TestPipelineServiceCopyPipeline:
)
# Execute
new_uuid = await service.copy_pipeline('original-uuid')
new_uuid = await service.copy_pipeline(WORKSPACE_UUID, 'original-uuid')
# Verify
assert new_uuid is not None
@@ -695,7 +656,7 @@ class TestPipelineServiceCopyPipeline:
service.get_pipeline = AsyncMock(return_value={'uuid': 'copy-uuid', 'is_default': False})
# Execute
await service.copy_pipeline('original-uuid')
await service.copy_pipeline(WORKSPACE_UUID, 'original-uuid')
# Verify - pipeline_mgr.load_pipeline called (copy created)
ap.pipeline_mgr.load_pipeline.assert_called_once()
@@ -715,100 +676,8 @@ class TestPipelineServiceUpdatePipelineExtensions:
service = PipelineService(ap)
# Execute & Verify
with pytest.raises(ValueError, match='Pipeline nonexistent-uuid not found'):
await service.update_pipeline_extensions('nonexistent-uuid', [])
@pytest.mark.parametrize(
('field', 'invalid_value'),
[
('bound_plugins', 'author/plugin'),
('bound_plugins', [{'author': 'author'}]),
('bound_mcp_servers', 'server-1'),
('bound_mcp_servers', ['server-1', 2]),
('bound_skills', 'skill-1'),
('bound_skills', ['skill-1', None]),
('bound_mcp_resources', {'uri': 'file:///README.md'}),
('bound_mcp_resources', [{'uri': 'file:///README.md'}, 'bad']),
],
)
async def test_update_extensions_rejects_malformed_binding_lists_before_query(
self,
field,
invalid_value,
):
ap = SimpleNamespace(persistence_mgr=SimpleNamespace(execute_async=AsyncMock()))
service = PipelineService(ap)
kwargs = {field: invalid_value}
if field != 'bound_plugins':
kwargs['bound_plugins'] = []
with pytest.raises(ValueError, match=field):
await service.update_pipeline_extensions('test-uuid', **kwargs)
ap.persistence_mgr.execute_async.assert_not_awaited()
@pytest.mark.parametrize('invalid_value', [0, 'false'])
async def test_update_extensions_rejects_non_boolean_resource_read_before_query(self, invalid_value):
ap = SimpleNamespace(persistence_mgr=SimpleNamespace(execute_async=AsyncMock()))
service = PipelineService(ap)
with pytest.raises(ValueError, match='mcp_resource_agent_read_enabled.*boolean'):
await service.update_pipeline_extensions(
'test-uuid',
[],
mcp_resource_agent_read_enabled=invalid_value,
)
ap.persistence_mgr.execute_async.assert_not_awaited()
@pytest.mark.parametrize(
'field',
[
'enable_all_plugins',
'enable_all_mcp_servers',
'enable_all_skills',
],
)
@pytest.mark.parametrize('invalid_value', [0, None, 'false'])
async def test_update_extensions_rejects_non_boolean_enable_all_flags_before_query(
self,
field,
invalid_value,
):
ap = SimpleNamespace(persistence_mgr=SimpleNamespace(execute_async=AsyncMock()))
service = PipelineService(ap)
with pytest.raises(ValueError, match=rf'{field}.*boolean'):
await service.update_pipeline_extensions(
'test-uuid',
[],
**{field: invalid_value},
)
ap.persistence_mgr.execute_async.assert_not_awaited()
@pytest.mark.parametrize('invalid_value', [0, None, 'false'])
async def test_update_extensions_rejects_non_boolean_attachment_enabled_before_query(
self,
invalid_value,
):
ap = SimpleNamespace(persistence_mgr=SimpleNamespace(execute_async=AsyncMock()))
service = PipelineService(ap)
with pytest.raises(ValueError, match=r'bound_mcp_resources.*enabled.*boolean'):
await service.update_pipeline_extensions(
'test-uuid',
[],
bound_mcp_resources=[
{
'server_uuid': 'server-1',
'uri': 'file:///README.md',
'enabled': invalid_value,
}
],
)
ap.persistence_mgr.execute_async.assert_not_awaited()
with pytest.raises(WorkspaceNotFoundError, match='Pipeline nonexistent-uuid not found'):
await service.update_pipeline_extensions(WORKSPACE_UUID, 'nonexistent-uuid', [])
async def test_update_extensions_sets_plugins(self):
"""Updates plugins in extensions_preferences."""
@@ -836,7 +705,7 @@ class TestPipelineServiceUpdatePipelineExtensions:
'uuid': 'test-uuid',
'extensions_preferences': {
'enable_all_plugins': False,
'plugins': [{'plugin_uuid': 'plugin-1'}],
'plugins': [{'author': 'test', 'name': 'plugin-1'}],
},
}
)
@@ -847,7 +716,7 @@ class TestPipelineServiceUpdatePipelineExtensions:
'uuid': 'test-uuid',
'extensions_preferences': {
'enable_all_plugins': False,
'plugins': [{'plugin_uuid': 'plugin-1'}],
'plugins': [{'author': 'test', 'name': 'plugin-1'}],
},
}
)
@@ -855,6 +724,7 @@ class TestPipelineServiceUpdatePipelineExtensions:
# Execute
bound_plugins = [{'author': 'test', 'name': 'plugin-1'}]
await service.update_pipeline_extensions(
WORKSPACE_UUID,
'test-uuid',
bound_plugins=bound_plugins,
enable_all_plugins=False,
@@ -904,6 +774,7 @@ class TestPipelineServiceUpdatePipelineExtensions:
# Execute
await service.update_pipeline_extensions(
WORKSPACE_UUID,
'test-uuid',
bound_plugins=[],
bound_mcp_servers=['mcp-server-1'],
@@ -951,7 +822,7 @@ class TestPipelineServiceUpdatePipelineExtensions:
)
# Execute - bound_mcp_servers is None (not provided)
await service.update_pipeline_extensions('test-uuid', bound_plugins=[])
await service.update_pipeline_extensions(WORKSPACE_UUID, 'test-uuid', bound_plugins=[])
# Verify - persistence was called
ap.persistence_mgr.execute_async.assert_called()
@@ -990,7 +861,7 @@ class TestPipelineServiceUpdatePipelineExtensions:
service = PipelineService(ap)
service.get_pipeline = AsyncMock(return_value={'uuid': 'test-uuid'})
await service.update_pipeline_extensions('test-uuid', bound_plugins=[])
await service.update_pipeline_extensions(WORKSPACE_UUID, 'test-uuid', bound_plugins=[])
assert original_pipeline.extensions_preferences['mcp_resource_agent_read_enabled'] is False
assert original_pipeline.extensions_preferences['mcp_resources'] == [
@@ -998,6 +869,82 @@ class TestPipelineServiceUpdatePipelineExtensions:
]
class TestPipelineSecretRoundtrip:
async def test_resource_view_redacts_runner_secrets_without_mutating_serialized_data(self):
raw = {
'uuid': 'pipeline-secret',
'config': {
'ai': {
'n8n': {
'webhook-url': 'https://hook.invalid/bearer-secret',
'headers': {'Authorization': 'Bearer secret'},
}
}
},
}
pipeline = _create_mock_pipeline(pipeline_uuid='pipeline-secret')
ap = SimpleNamespace(
persistence_mgr=SimpleNamespace(
execute_async=AsyncMock(return_value=_create_mock_result([pipeline])),
serialize_model=Mock(return_value=raw),
)
)
redacted = await PipelineService(ap).get_pipelines(WORKSPACE_UUID)
assert redacted[0]['config']['ai']['n8n']['webhook-url'] == '***'
assert redacted[0]['config']['ai']['n8n']['headers']['Authorization'] == '***'
assert raw['config']['ai']['n8n']['webhook-url'] == 'https://hook.invalid/bearer-secret'
async def test_masked_runner_config_update_restores_existing_secret(self):
raw_config = {
'ai': {
'n8n': {
'webhook-url': 'https://hook.invalid/bearer-secret',
'headers': {'Authorization': 'Bearer secret'},
'timeout': 30,
}
}
}
current_pipeline = {'uuid': 'pipeline-secret', 'config': raw_config}
write_result = Mock(rowcount=1)
ap = SimpleNamespace(
persistence_mgr=SimpleNamespace(execute_async=AsyncMock(return_value=write_result)),
pipeline_mgr=SimpleNamespace(remove_pipeline=AsyncMock(), load_pipeline=AsyncMock()),
sess_mgr=SimpleNamespace(session_list=[]),
)
service = PipelineService(ap)
service.get_pipeline = AsyncMock(side_effect=[current_pipeline, current_pipeline])
await service.update_pipeline(
WORKSPACE_UUID,
'pipeline-secret',
{
'config': {
'ai': {
'n8n': {
'webhook-url': '***',
'headers': {'Authorization': '***'},
'timeout': 60,
}
}
}
},
)
statement = ap.persistence_mgr.execute_async.await_args.args[0]
stored_config = next(value.value for column, value in statement._values.items() if column.key == 'config')
assert stored_config == {
'ai': {
'n8n': {
'webhook-url': 'https://hook.invalid/bearer-secret',
'headers': {'Authorization': 'Bearer secret'},
'timeout': 60,
}
}
}
class TestDefaultStageOrder:
"""Tests for default_stage_order constant."""
@@ -19,10 +19,13 @@ from types import SimpleNamespace
from langbot.pkg.api.http.service.provider import ModelProviderService
from langbot.pkg.entity.persistence.model import ModelProvider, LLMModel, EmbeddingModel, RerankModel
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
pytestmark = pytest.mark.asyncio
WORKSPACE_UUID = 'workspace-a'
def _create_mock_provider(
provider_uuid: str = 'test-provider-uuid',
@@ -86,7 +89,9 @@ class TestModelProviderServiceGetProviders:
service = ModelProviderService(ap)
# Execute
result = await service.get_providers()
result = await service.get_providers(
WORKSPACE_UUID,
)
# Verify
assert result == []
@@ -115,7 +120,9 @@ class TestModelProviderServiceGetProviders:
service = ModelProviderService(ap)
# Execute
result = await service.get_providers()
result = await service.get_providers(
WORKSPACE_UUID,
)
# Verify
assert len(result) == 2
@@ -143,7 +150,10 @@ class TestModelProviderServiceGetProviders:
service = ModelProviderService(ap)
# Execute
result = await service.get_providers()
result = await service.get_providers(
WORKSPACE_UUID,
include_secret=True,
)
# Verify - api_keys should be parsed from string
assert result[0]['api_keys'] == ['key1', 'key2']
@@ -169,11 +179,41 @@ class TestModelProviderServiceGetProviders:
service = ModelProviderService(ap)
# Execute
result = await service.get_providers()
result = await service.get_providers(
WORKSPACE_UUID,
)
# Verify - invalid JSON returns empty list
assert result[0]['api_keys'] == []
async def test_get_providers_masks_api_keys_for_resource_view(self):
ap = SimpleNamespace()
provider = _create_mock_provider(
api_keys=['first', 'second'],
base_url=(
'https://provider-user:provider-password@api.provider.invalid/v1?access_token=url-secret&region=sg'
),
)
ap.persistence_mgr = SimpleNamespace(
execute_async=AsyncMock(return_value=_create_mock_result([provider])),
serialize_model=Mock(
return_value={
'uuid': provider.uuid,
'name': provider.name,
'base_url': provider.base_url,
'api_keys': provider.api_keys,
}
),
)
result = await ModelProviderService(ap).get_providers(
WORKSPACE_UUID,
include_secret=False,
)
assert result[0]['api_keys'] == ['***', '***']
assert result[0]['base_url'] == ('https://***@api.provider.invalid/v1?access_token=***&region=sg')
class TestModelProviderServiceGetProvider:
"""Tests for get_provider method."""
@@ -199,7 +239,7 @@ class TestModelProviderServiceGetProvider:
service = ModelProviderService(ap)
# Execute
result = await service.get_provider('found-uuid')
result = await service.get_provider(WORKSPACE_UUID, 'found-uuid')
# Verify
assert result is not None
@@ -217,7 +257,7 @@ class TestModelProviderServiceGetProvider:
service = ModelProviderService(ap)
# Execute
result = await service.get_provider('nonexistent-uuid')
result = await service.get_provider(WORKSPACE_UUID, 'nonexistent-uuid')
# Verify
assert result is None
@@ -239,6 +279,7 @@ class TestModelProviderServiceCreateProvider:
runtime_provider.provider_entity = Mock()
runtime_provider.provider_entity.uuid = 'generated-uuid'
ap.model_mgr.load_provider = AsyncMock(return_value=runtime_provider)
ap.model_mgr.cache_provider = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock()
@@ -246,12 +287,13 @@ class TestModelProviderServiceCreateProvider:
# Execute
provider_uuid = await service.create_provider(
WORKSPACE_UUID,
{
'name': 'New Provider',
'requester': 'openai',
'base_url': 'https://api.openai.com',
'api_keys': ['key'],
}
},
)
# Verify - UUID is generated
@@ -270,6 +312,7 @@ class TestModelProviderServiceCreateProvider:
runtime_provider.provider_entity = Mock()
runtime_provider.provider_entity.uuid = 'runtime-uuid'
ap.model_mgr.load_provider = AsyncMock(return_value=runtime_provider)
ap.model_mgr.cache_provider = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock()
@@ -277,12 +320,13 @@ class TestModelProviderServiceCreateProvider:
# Execute
result_uuid = await service.create_provider(
WORKSPACE_UUID,
{
'name': 'Runtime Provider',
'requester': 'openai',
'base_url': 'https://api.openai.com',
'api_keys': ['key'],
}
},
)
# Verify - provider added to runtime dict and UUID generated
@@ -307,6 +351,7 @@ class TestModelProviderServiceUpdateProvider:
# Execute
await service.update_provider(
WORKSPACE_UUID,
'existing-uuid',
{
'uuid': 'should-be-removed', # Will be removed
@@ -315,7 +360,7 @@ class TestModelProviderServiceUpdateProvider:
)
# Verify - reload called
ap.model_mgr.reload_provider.assert_called_once_with('existing-uuid')
ap.model_mgr.reload_provider.assert_called_once_with(WORKSPACE_UUID, 'existing-uuid')
async def test_update_provider_reloads_runtime(self):
"""Reloads provider in runtime after update."""
@@ -330,7 +375,7 @@ class TestModelProviderServiceUpdateProvider:
service = ModelProviderService(ap)
# Execute
await service.update_provider('update-uuid', {'name': 'New Name'})
await service.update_provider(WORKSPACE_UUID, 'update-uuid', {'name': 'New Name'})
# Verify
ap.model_mgr.reload_provider.assert_called_once()
@@ -354,7 +399,7 @@ class TestModelProviderServiceDeleteProvider:
# Execute & Verify
with pytest.raises(ValueError, match='Cannot delete provider: LLM models'):
await service.delete_provider('provider-with-llm')
await service.delete_provider(WORKSPACE_UUID, 'provider-with-llm')
async def test_delete_provider_with_embedding_models_raises_error(self):
"""Raises ValueError when Embedding models reference provider."""
@@ -387,7 +432,7 @@ class TestModelProviderServiceDeleteProvider:
# Execute & Verify - should raise embedding error (LLM check passes, embedding check fails)
with pytest.raises(ValueError, match='Cannot delete provider: Embedding models'):
await service.delete_provider('provider-with-embedding')
await service.delete_provider(WORKSPACE_UUID, 'provider-with-embedding')
async def test_delete_provider_with_rerank_models_raises_error(self):
"""Raises ValueError when Rerank models reference provider."""
@@ -420,7 +465,7 @@ class TestModelProviderServiceDeleteProvider:
# Execute & Verify - should raise rerank error (LLM and embedding checks pass, rerank check fails)
with pytest.raises(ValueError, match='Cannot delete provider: Rerank models'):
await service.delete_provider('provider-with-rerank')
await service.delete_provider(WORKSPACE_UUID, 'provider-with-rerank')
async def test_delete_provider_no_models_success(self):
"""Deletes provider when no models reference it."""
@@ -439,10 +484,10 @@ class TestModelProviderServiceDeleteProvider:
service = ModelProviderService(ap)
# Execute
await service.delete_provider('provider-no-models')
await service.delete_provider(WORKSPACE_UUID, 'provider-no-models')
# Verify - delete and remove called
ap.model_mgr.remove_provider.assert_called_once_with('provider-no-models')
ap.model_mgr.remove_provider.assert_called_once_with(WORKSPACE_UUID, 'provider-no-models')
class TestModelProviderServiceGetProviderModelCounts:
@@ -476,9 +521,10 @@ class TestModelProviderServiceGetProviderModelCounts:
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
service = ModelProviderService(ap)
service.get_provider = AsyncMock(return_value={'uuid': 'provider-uuid'})
# Execute
result = await service.get_provider_model_counts('provider-uuid')
result = await service.get_provider_model_counts(WORKSPACE_UUID, 'provider-uuid')
# Verify
assert result['llm_count'] == 3
@@ -497,9 +543,10 @@ class TestModelProviderServiceGetProviderModelCounts:
ap.persistence_mgr.execute_async = AsyncMock(return_value=zero_result)
service = ModelProviderService(ap)
service.get_provider = AsyncMock(return_value={'uuid': 'empty-provider'})
# Execute
result = await service.get_provider_model_counts('empty-provider')
result = await service.get_provider_model_counts(WORKSPACE_UUID, 'empty-provider')
# Verify
assert result['llm_count'] == 0
@@ -530,6 +577,7 @@ class TestModelProviderServiceFindOrCreateProvider:
# Execute
result = await service.find_or_create_provider(
WORKSPACE_UUID,
requester='openai',
base_url='https://api.openai.com',
api_keys=['key1', 'key2'], # Same keys (sorted)
@@ -558,6 +606,7 @@ class TestModelProviderServiceFindOrCreateProvider:
# Execute with reversed key order
result = await service.find_or_create_provider(
WORKSPACE_UUID,
requester='openai',
base_url='https://api.openai.com',
api_keys=['key2', 'key1'], # Different order, should still match
@@ -578,6 +627,7 @@ class TestModelProviderServiceFindOrCreateProvider:
runtime_provider.provider_entity = Mock()
runtime_provider.provider_entity.uuid = None # Will be set by uuid.uuid4()
ap.model_mgr.load_provider = AsyncMock(return_value=runtime_provider)
ap.model_mgr.cache_provider = AsyncMock()
# Mock no existing providers
mock_result = _create_mock_result([])
@@ -587,6 +637,7 @@ class TestModelProviderServiceFindOrCreateProvider:
# Execute
result = await service.find_or_create_provider(
WORKSPACE_UUID,
requester='new-requester',
base_url='https://new.api.com',
api_keys=['new-key'],
@@ -610,6 +661,7 @@ class TestModelProviderServiceFindOrCreateProvider:
runtime_provider.provider_entity = Mock()
runtime_provider.provider_entity.uuid = 'parsed-url-uuid'
ap.model_mgr.load_provider = AsyncMock(return_value=runtime_provider)
ap.model_mgr.cache_provider = AsyncMock()
mock_result = _create_mock_result([])
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
@@ -618,6 +670,7 @@ class TestModelProviderServiceFindOrCreateProvider:
# Execute
result_uuid = await service.find_or_create_provider(
WORKSPACE_UUID,
requester='custom',
base_url='https://api.example.com/v1',
api_keys=['key'],
@@ -644,17 +697,20 @@ class TestModelProviderServiceUpdateSpaceModelProviderApiKeys:
service = ModelProviderService(ap)
# Execute
await service.update_space_model_provider_api_keys('space-api-key')
await service.update_space_model_provider_api_keys(WORKSPACE_UUID, 'space-api-key')
# Verify - update and reload called for Space provider UUID
ap.model_mgr.reload_provider.assert_called_once_with('00000000-0000-0000-0000-000000000000')
ap.model_mgr.reload_provider.assert_called_once_with(
WORKSPACE_UUID,
'00000000-0000-0000-0000-000000000000',
)
class TestModelProviderServiceScanProviderModels:
"""Tests for scan_provider_models method."""
async def test_scan_provider_not_found_raises_error(self):
"""Raises ValueError when provider not found."""
"""Raises a non-enumerating not-found error when provider is outside the Workspace."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
@@ -665,8 +721,8 @@ class TestModelProviderServiceScanProviderModels:
service = ModelProviderService(ap)
# Execute & Verify
with pytest.raises(ValueError, match='provider not found'):
await service.scan_provider_models('nonexistent-uuid')
with pytest.raises(WorkspaceNotFoundError, match='Provider not found'):
await service.scan_provider_models(WORKSPACE_UUID, 'nonexistent-uuid')
async def test_scan_provider_returns_models_list(self):
"""Returns scanned models list."""
@@ -718,7 +774,7 @@ class TestModelProviderServiceScanProviderModels:
service = ModelProviderService(ap)
# Execute
result = await service.scan_provider_models('scan-uuid')
result = await service.scan_provider_models(WORKSPACE_UUID, 'scan-uuid')
# Verify
assert 'models' in result
@@ -771,12 +827,50 @@ class TestModelProviderServiceScanProviderModels:
service = ModelProviderService(ap)
# Execute - filter for LLM only
result = await service.scan_provider_models('filter-uuid', model_type='llm')
result = await service.scan_provider_models(WORKSPACE_UUID, 'filter-uuid', model_type='llm')
# Verify - only LLM models returned
assert len(result['models']) == 1
assert result['models'][0]['type'] == 'llm'
async def test_scan_provider_marks_existing_rerank_model(self):
"""Rerank scan results use the rerank service when computing already_added."""
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.model_mgr = SimpleNamespace()
ap.llm_model_service = SimpleNamespace()
ap.embedding_models_service = SimpleNamespace()
ap.rerank_models_service = SimpleNamespace()
provider = _create_mock_provider(provider_uuid='rerank-scan-uuid')
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result([], first_item=provider))
ap.persistence_mgr.serialize_model = Mock(
return_value={
'uuid': 'rerank-scan-uuid',
'name': 'New API',
'requester': 'new-api-chat-completions',
'base_url': 'https://new-api.example.com/v1',
'api_keys': ['key'],
}
)
runtime_provider = Mock()
runtime_provider.token_mgr.get_token.return_value = 'token'
runtime_provider.requester.scan_models = AsyncMock(
return_value={'models': [{'id': 'Qwen3-Reranker-8B', 'type': 'rerank'}]}
)
ap.model_mgr.load_provider = AsyncMock(return_value=runtime_provider)
ap.llm_model_service.get_llm_models_by_provider = AsyncMock(return_value=[])
ap.embedding_models_service.get_embedding_models_by_provider = AsyncMock(return_value=[])
ap.rerank_models_service.get_rerank_models_by_provider = AsyncMock(return_value=[{'name': 'Qwen3-Reranker-8B'}])
result = await ModelProviderService(ap).scan_provider_models(
WORKSPACE_UUID, 'rerank-scan-uuid', model_type='rerank'
)
assert result['models'][0]['type'] == 'rerank'
assert result['models'][0]['already_added'] is True
async def test_scan_provider_not_implemented_raises_error(self):
"""Raises ValueError when scan not implemented."""
# Setup
@@ -810,7 +904,7 @@ class TestModelProviderServiceScanProviderModels:
# Execute & Verify
with pytest.raises(ValueError, match='current provider does not support model scanning'):
await service.scan_provider_models('no-scan-uuid')
await service.scan_provider_models(WORKSPACE_UUID, 'no-scan-uuid')
async def test_scan_provider_marks_already_added_models(self):
"""Marks models that are already added."""
@@ -860,7 +954,7 @@ class TestModelProviderServiceScanProviderModels:
service = ModelProviderService(ap)
# Execute
result = await service.scan_provider_models('already-added-uuid')
result = await service.scan_provider_models(WORKSPACE_UUID, 'already-added-uuid')
# Verify - existing model marked as already_added
existing_model = next(m for m in result['models'] if m['name'] == 'Existing Model')
@@ -868,3 +962,46 @@ class TestModelProviderServiceScanProviderModels:
new_model = next(m for m in result['models'] if m['name'] == 'New Model')
assert new_model['already_added'] is False
class TestProviderSecretRoundtrip:
async def test_masked_api_keys_update_preserves_existing_values(self):
write_result = Mock(rowcount=1)
ap = SimpleNamespace(
persistence_mgr=SimpleNamespace(execute_async=AsyncMock(return_value=write_result)),
model_mgr=SimpleNamespace(reload_provider=AsyncMock()),
)
service = ModelProviderService(ap)
service.get_provider = AsyncMock(
return_value={
'uuid': 'provider-secret',
'api_keys': ['first-secret', 'second-secret'],
}
)
await service.update_provider(
WORKSPACE_UUID,
'provider-secret',
{'name': 'Updated', 'api_keys': ['***', 'replacement-secret']},
)
statement = ap.persistence_mgr.execute_async.await_args.args[0]
stored_api_keys = next(value.value for column, value in statement._values.items() if column.key == 'api_keys')
assert stored_api_keys == ['first-secret', 'replacement-secret']
async def test_extra_masked_api_key_is_rejected(self):
ap = SimpleNamespace(
persistence_mgr=SimpleNamespace(execute_async=AsyncMock()),
model_mgr=SimpleNamespace(reload_provider=AsyncMock()),
)
service = ModelProviderService(ap)
service.get_provider = AsyncMock(return_value={'uuid': 'provider-secret', 'api_keys': ['only-secret']})
with pytest.raises(ValueError, match='no existing value'):
await service.update_provider(
WORKSPACE_UUID,
'provider-secret',
{'api_keys': ['***', '***']},
)
ap.persistence_mgr.execute_async.assert_not_awaited()
@@ -0,0 +1,88 @@
from __future__ import annotations
import copy
import pytest
from langbot.pkg.api.http.service.secrets import (
contains_secret_placeholder,
redact_secrets,
restore_secret_placeholders,
)
RAW_CONFIG = {
'apiKey': 'api-secret',
'dify_apikey': 'dify-secret',
'base_url': (
'https://service-user:service-password@api.invalid/v1'
'?api_key=query-secret&region=sg&X-Amz-Signature=signed-secret'
),
'nested': {
'headers': {
'Authorization': 'Bearer nested-secret',
'X-API-Key': 'header-secret',
'Accept': 'application/json',
},
'webhook-url': 'https://hooks.invalid/path?token=secret',
'public_key': 'public-material',
'tokenizer': 'not-a-secret',
},
'credentials': {'username': 'service-user', 'password': 'service-password'},
'secret_list': ['first-secret', {'value': 'second-secret'}],
'empty_secret': '',
'enabled': True,
}
def test_recursive_redaction_is_shape_preserving_and_does_not_mutate_source():
source = copy.deepcopy(RAW_CONFIG)
redacted = redact_secrets(source)
assert redacted['apiKey'] == '***'
assert redacted['dify_apikey'] == '***'
assert redacted['base_url'] == ('https://***@api.invalid/v1?api_key=***&region=sg&X-Amz-Signature=***')
assert redacted['nested']['headers'] == {
'Authorization': '***',
'X-API-Key': '***',
'Accept': 'application/json',
}
assert redacted['nested']['webhook-url'] == '***'
assert redacted['nested']['public_key'] == 'public-material'
assert redacted['nested']['tokenizer'] == 'not-a-secret'
assert redacted['credentials'] == {'username': '***', 'password': '***'}
assert redacted['secret_list'] == ['***', {'value': '***'}]
assert redacted['empty_secret'] == ''
assert redacted['enabled'] is True
assert source == RAW_CONFIG
def test_masked_roundtrip_preserves_existing_secrets_and_accepts_replace_and_clear():
submitted = redact_secrets(RAW_CONFIG)
submitted['enabled'] = False
submitted['apiKey'] = 'replacement-secret'
submitted['nested']['headers']['X-API-Key'] = ''
restored = restore_secret_placeholders(submitted, RAW_CONFIG)
assert restored['apiKey'] == 'replacement-secret'
assert restored['dify_apikey'] == 'dify-secret'
assert restored['nested']['headers']['Authorization'] == 'Bearer nested-secret'
assert restored['nested']['headers']['X-API-Key'] == ''
assert restored['nested']['webhook-url'] == RAW_CONFIG['nested']['webhook-url']
assert restored['base_url'] == RAW_CONFIG['base_url']
assert restored['enabled'] is False
assert RAW_CONFIG['apiKey'] == 'api-secret'
def test_new_or_extra_masked_secret_fails_closed():
assert contains_secret_placeholder({'headers': {'Authorization': '***'}})
assert contains_secret_placeholder({'base_url': 'https://***@api.invalid?token=***'})
with pytest.raises(ValueError, match='no existing value'):
restore_secret_placeholders({'api_key': '***'})
with pytest.raises(ValueError, match='no existing value'):
restore_secret_placeholders(
{'api_keys': ['***', '***']},
{'api_keys': ['existing']},
)
@@ -13,6 +13,10 @@ Source: src/langbot/pkg/api/http/service/space.py
from __future__ import annotations
from collections import OrderedDict
import json
from urllib.parse import parse_qs, urlsplit
import pytest
from unittest.mock import AsyncMock, Mock, patch, MagicMock
from types import SimpleNamespace
@@ -26,6 +30,23 @@ from langbot.pkg.entity.persistence.user import User
pytestmark = pytest.mark.asyncio
def _set_response_body(response: MagicMock, body: dict | str) -> None:
"""Configure an aiohttp-like streaming body on an HTTP response mock."""
raw_body = body.encode() if isinstance(body, str) else json.dumps(body).encode()
class Content:
async def iter_chunked(self, _chunk_size: int):
midpoint = max(len(raw_body) // 2, 1)
yield raw_body[:midpoint]
if midpoint < len(raw_body):
yield raw_body[midpoint:]
response.headers = {}
response.content = Content()
response.charset = 'utf-8'
def _create_mock_user(
email: str = 'test@example.com',
account_type: str = 'space',
@@ -73,7 +94,7 @@ class TestSpaceServiceGetOAuthAuthorizeUrl:
result = service.get_oauth_authorize_url('http://localhost/callback')
# Verify
assert 'redirect_uri=http://localhost/callback' in result
assert parse_qs(urlsplit(result).query)['redirect_uri'] == ['http://localhost/callback']
assert 'https://space.langbot.app/auth/authorize' in result
def test_get_oauth_authorize_url_with_state(self):
@@ -93,8 +114,9 @@ class TestSpaceServiceGetOAuthAuthorizeUrl:
result = service.get_oauth_authorize_url('http://localhost/callback', state='random_state')
# Verify
assert 'redirect_uri=http://localhost/callback' in result
assert 'state=random_state' in result
params = parse_qs(urlsplit(result).query)
assert params['redirect_uri'] == ['http://localhost/callback']
assert params['state'] == ['random_state']
def test_get_oauth_authorize_url_default_config(self):
"""Uses default OAuth URL when config not set."""
@@ -289,6 +311,40 @@ class TestSpaceServiceGetCredits:
# Verify - returns cached value without API call
assert result == 100
async def test_cached_credit_lookup_does_not_scan_all_users(self):
ap = SimpleNamespace()
ap.instance_config = SimpleNamespace(data={})
ap.persistence_mgr = SimpleNamespace()
service = SpaceService(ap)
class AtMostOneStepOrderedDict(OrderedDict):
def __iter__(self):
iterator = super().__iter__()
yielded = False
def next_entry():
nonlocal yielded
if yielded:
raise AssertionError('credits cache scanned all users')
yielded = True
return next(iterator)
class AtMostOneStepIterator:
def __iter__(self):
return self
def __next__(self):
return next_entry()
return AtMostOneStepIterator()
now = time.time()
service._credits_cache = AtMostOneStepOrderedDict(
(f'user-{index}@example.com', (index, now)) for index in range(512)
)
assert await service.get_credits('user-511@example.com') == 511
async def test_get_credits_cache_expired_refreshes(self):
"""Refreshes expired cache."""
# Setup
@@ -403,6 +459,7 @@ class TestSpaceServiceRefreshToken:
},
}
)
_set_response_body(mock_response, mock_response.json.return_value)
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
mock_session_obj = MagicMock()
@@ -438,6 +495,7 @@ class TestSpaceServiceRefreshToken:
}
)
mock_response.text = AsyncMock(return_value='{"code":1,"msg":"Invalid refresh token"}')
_set_response_body(mock_response, mock_response.json.return_value)
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
mock_session_obj = MagicMock()
@@ -464,6 +522,7 @@ class TestSpaceServiceRefreshToken:
mock_response = MagicMock()
mock_response.status = 500
mock_response.text = AsyncMock(return_value='Internal Server Error')
_set_response_body(mock_response, mock_response.text.return_value)
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
mock_session_obj = MagicMock()
@@ -503,6 +562,7 @@ class TestSpaceServiceExchangeOAuthCode:
},
}
)
_set_response_body(mock_response, mock_response.json.return_value)
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
mock_session_obj = MagicMock()
@@ -532,6 +592,7 @@ class TestSpaceServiceExchangeOAuthCode:
mock_response.status = 200
mock_response.json = AsyncMock(return_value={'code': 1, 'msg': 'Invalid code'})
mock_response.text = AsyncMock(return_value='{"code":1,"msg":"Invalid code"}')
_set_response_body(mock_response, mock_response.json.return_value)
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
mock_session_obj = MagicMock()
@@ -570,6 +631,7 @@ class TestSpaceServiceGetUserInfoRaw:
},
}
)
_set_response_body(mock_response, mock_response.json.return_value)
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
mock_session_obj = MagicMock()
@@ -600,6 +662,7 @@ class TestSpaceServiceGetUserInfoRaw:
mock_response.status = 200
mock_response.json = AsyncMock(return_value={'code': 1, 'msg': 'Unauthorized'})
mock_response.text = AsyncMock(return_value='{"code":1,"msg":"Unauthorized"}')
_set_response_body(mock_response, mock_response.json.return_value)
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
mock_session_obj = MagicMock()
@@ -693,13 +756,14 @@ class TestSpaceServiceGetModels:
'uuid': 'uuid-2',
'model_id': 'model-2',
'provider': 'provider-2',
'category': 'chat',
'category': 'rerank',
'status': 'active',
},
]
},
}
)
_set_response_body(mock_response, mock_response.json.return_value)
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
mock_session_obj = MagicMock()
@@ -714,6 +778,7 @@ class TestSpaceServiceGetModels:
# Verify
assert len(result) == 2
assert result[1].category == 'rerank'
async def test_get_models_api_error(self):
"""Raises ValueError on API error."""
@@ -729,6 +794,7 @@ class TestSpaceServiceGetModels:
mock_response.status = 200
mock_response.json = AsyncMock(return_value={'code': 1, 'msg': 'Unauthorized'})
mock_response.text = AsyncMock(return_value='{"code":1,"msg":"Unauthorized"}')
_set_response_body(mock_response, mock_response.json.return_value)
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
mock_session_obj = MagicMock()
@@ -14,17 +14,124 @@ Source: src/langbot/pkg/api/http/service/user.py
from __future__ import annotations
import pytest
import jwt
import datetime
from unittest.mock import AsyncMock, Mock
from types import SimpleNamespace
from langbot.pkg.api.http.service.user import UserService
from langbot.pkg.entity.persistence.user import User
from langbot.pkg.entity.errors.account import AccountEmailMismatchError
from langbot.pkg.api.http.service.user import (
ControlPlaneDirectoryRequiredError,
UserService,
)
from langbot.pkg.entity.persistence.user import AccountSource, AccountStatus, User
from langbot.pkg.entity.errors.account import (
AccountEmailMismatchError,
SpaceAccountBindingRequiredError,
SpaceAccountNotRegisteredError,
)
from langbot.pkg.utils.bounded_executor import BlockingWorkCapacityError
pytestmark = pytest.mark.asyncio
async def test_password_hashing_rejects_concurrent_waiters() -> None:
service = UserService(SimpleNamespace())
await service._password_hash_lock.acquire()
try:
with pytest.raises(
BlockingWorkCapacityError,
match='Password hashing capacity reached',
):
await service._hash_password('secret')
finally:
service._password_hash_lock.release()
class TestSpaceOAuthState:
async def test_login_state_is_opaque_single_use(self):
service = UserService(SimpleNamespace())
state = await service.issue_space_oauth_state('login')
assert state.count('.') == 0
assert await service.consume_space_oauth_state(state, 'login') is None
with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
await service.consume_space_oauth_state(state, 'login')
async def test_bind_state_resolves_only_bound_active_account(self):
service = UserService(SimpleNamespace())
account = SimpleNamespace(uuid='account-a', status=AccountStatus.ACTIVE.value)
service.get_user_by_uuid = AsyncMock(return_value=account)
state = await service.issue_space_oauth_state('bind', account_uuid='account-a')
assert await service.consume_space_oauth_state(state, 'bind') is account
service.get_user_by_uuid.assert_awaited_once_with('account-a')
async def test_state_purpose_mismatch_is_rejected_and_consumed(self):
service = UserService(SimpleNamespace())
state = await service.issue_space_oauth_state('login')
with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
await service.consume_space_oauth_state(state, 'bind')
with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
await service.consume_space_oauth_state(state, 'login')
async def test_expired_state_is_rejected(self):
service = UserService(SimpleNamespace())
state = await service.issue_space_oauth_state('login')
digest = service._space_oauth_state_digest(state)
purpose, account_uuid, _, launch_workspace_uuid = service._space_oauth_states[digest]
service._space_oauth_states[digest] = (purpose, account_uuid, 0, launch_workspace_uuid)
with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
await service.consume_space_oauth_state(state, 'login')
async def test_login_state_can_carry_launch_workspace_without_changing_normal_return(self):
service = UserService(SimpleNamespace())
state = await service.issue_space_oauth_state(
'login',
launch_workspace_uuid='workspace-a',
)
assert await service.consume_space_oauth_state(state, 'login') is None
state = await service.issue_space_oauth_state(
'login',
launch_workspace_uuid='workspace-a',
)
consumed = await service.consume_space_oauth_state_details(state, 'login')
assert consumed.account is None
assert consumed.launch_workspace_uuid == 'workspace-a'
async def test_issue_state_does_not_scan_all_live_states(self, monkeypatch):
service = UserService(SimpleNamespace())
for _ in range(512):
await service.issue_space_oauth_state('login')
class NoGlobalIterationDict(dict):
def __iter__(self):
raise AssertionError('OAuth state issuance scanned all live states')
def keys(self):
raise AssertionError('OAuth state issuance scanned all live states')
def items(self):
raise AssertionError('OAuth state issuance scanned all live states')
def values(self):
raise AssertionError('OAuth state issuance scanned all live states')
guarded_states = NoGlobalIterationDict(service._space_oauth_states)
monkeypatch.setattr(service, '_space_oauth_states', guarded_states)
state = await service.issue_space_oauth_state('login')
assert await service.consume_space_oauth_state(state, 'login') is None
assert len(guarded_states) == 512
def _create_mock_user(
email: str = 'test@example.com',
password: str = 'hashed_password',
@@ -34,6 +141,7 @@ def _create_mock_user(
"""Helper to create mock User entity."""
user = Mock(spec=User)
user.user = email
user.uuid = f'account-{email}'
user.password = password
user.account_type = account_type
user.space_account_uuid = space_account_uuid
@@ -102,6 +210,41 @@ class TestUserServiceIsInitialized:
assert result is False
class TestUserServiceGetLoginCapabilities:
"""Tests for public login capability discovery."""
async def test_uses_explicit_identity_discovery_scope(self):
discovery_result = Mock()
discovery_result.one = Mock(return_value=(1, 2))
discovery_session = SimpleNamespace(execute=AsyncMock(return_value=discovery_result))
class DiscoveryContext:
async def __aenter__(self):
return SimpleNamespace(session=discovery_session)
async def __aexit__(self, exc_type, exc, tb):
return False
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace(
current_session=Mock(return_value=None),
identity_discovery_uow=Mock(return_value=DiscoveryContext()),
execute_async=AsyncMock(side_effect=AssertionError('unscoped persistence access')),
)
ap.workspace_service = SimpleNamespace(instance_uuid='instance-a')
service = UserService(ap)
result = await service.get_login_capabilities()
assert result == {
'password_login_enabled': True,
'space_login_enabled': True,
}
ap.persistence_mgr.identity_discovery_uow.assert_called_once()
discovery_session.execute.assert_awaited_once()
ap.persistence_mgr.execute_async.assert_not_awaited()
class TestUserServiceGetUserByEmail:
"""Tests for get_user_by_email method."""
@@ -309,6 +452,50 @@ class TestUserServiceVerifyJwtToken:
with pytest.raises(Exception): # jwt.DecodeError or similar
await service.verify_jwt_token('invalid.token.here')
async def test_verify_jwt_token_rejects_foreign_audience(self):
ap = SimpleNamespace()
ap.instance_config = SimpleNamespace()
ap.instance_config.data = {'system': {'jwt': {'secret': 'test_secret', 'expire': 3600}}}
service = UserService(ap)
token = jwt.encode(
{
'user': 'verify@example.com',
'iss': 'langbot-core',
'aud': 'langbot-instance:another-instance',
'exp': datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(hours=1),
},
'test_secret',
algorithm='HS256',
)
with pytest.raises(jwt.InvalidAudienceError):
await service.verify_jwt_token(token)
async def test_verify_jwt_token_accepts_legacy_community_token_only_in_oss(self):
ap = SimpleNamespace()
ap.instance_config = SimpleNamespace()
ap.instance_config.data = {'system': {'jwt': {'secret': 'test_secret', 'expire': 3600}}}
ap.workspace_service = SimpleNamespace(
instance_uuid='instance-a',
policy=SimpleNamespace(multi_workspace_enabled=False),
)
service = UserService(ap)
legacy_token = jwt.encode(
{
'user': 'legacy@example.com',
'iss': 'LangBot-community',
'exp': datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(hours=1),
},
'test_secret',
algorithm='HS256',
)
assert await service.verify_jwt_token(legacy_token) == 'legacy@example.com'
ap.workspace_service.policy.multi_workspace_enabled = True
with pytest.raises(jwt.MissingRequiredClaimError):
await service.verify_jwt_token(legacy_token)
class TestUserServiceResetPassword:
"""Tests for reset_password method."""
@@ -476,6 +663,71 @@ class TestUserServiceCreateOrUpdateSpaceUser:
ap.persistence_mgr.execute_async.assert_called()
assert updated_user.space_account_uuid == 'existing-space-uuid'
async def test_cloud_login_updates_only_the_projected_space_account(self):
projected = SimpleNamespace(
uuid='projected-space-uuid',
user='Cloud Owner',
normalized_email='owner@example.com',
password='',
account_type='space',
status=AccountStatus.ACTIVE.value,
source=AccountSource.CLOUD_PROJECTION.value,
projection_revision=7,
space_account_uuid='projected-space-uuid',
)
persistence = SimpleNamespace(execute_async=AsyncMock())
ap = SimpleNamespace(
persistence_mgr=persistence,
workspace_service=SimpleNamespace(policy=SimpleNamespace(multi_workspace_enabled=True)),
)
service = UserService(ap)
service.get_user_by_space_account_uuid = AsyncMock(side_effect=[projected, projected])
result = await service.create_or_update_space_user(
space_account_uuid='projected-space-uuid',
email='OWNER@example.com',
access_token='access-token',
refresh_token='refresh-token',
api_key='api-key',
expires_in=3600,
)
assert result is projected
persistence.execute_async.assert_awaited_once()
async def test_cloud_login_never_creates_an_unprojected_account(self):
persistence = SimpleNamespace(execute_async=AsyncMock())
ap = SimpleNamespace(
persistence_mgr=persistence,
workspace_service=SimpleNamespace(policy=SimpleNamespace(multi_workspace_enabled=True)),
)
service = UserService(ap)
service.get_user_by_space_account_uuid = AsyncMock(return_value=None)
with pytest.raises(
ControlPlaneDirectoryRequiredError,
match='verified Cloud directory',
):
await service.create_or_update_space_user(
space_account_uuid='unknown-space-uuid',
email='unknown@example.com',
access_token='access-token',
refresh_token='refresh-token',
api_key='api-key',
expires_in=3600,
)
persistence.execute_async.assert_not_awaited()
async def test_cloud_invitation_registration_requires_space_identity(self):
ap = SimpleNamespace(
workspace_service=SimpleNamespace(policy=SimpleNamespace(multi_workspace_enabled=True)),
)
service = UserService(ap)
with pytest.raises(ControlPlaneDirectoryRequiredError, match='Space account'):
await service.register_invited_account('invite-token', 'member@example.com', 'password')
async def test_create_or_update_new_space_user_first_init(self):
"""Creates new Space user on first initialization."""
# Setup
@@ -522,8 +774,8 @@ class TestUserServiceCreateOrUpdateSpaceUser:
# Verify
assert result.space_account_uuid == 'new-space-uuid'
async def test_create_or_update_space_user_already_initialized_raises_error(self):
"""Raises AccountEmailMismatchError when system already initialized and user not found."""
async def test_create_or_update_space_user_already_initialized_reports_unknown_space_email(self):
"""Unknown Space email is distinct from an existing local Account collision."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
@@ -538,7 +790,7 @@ class TestUserServiceCreateOrUpdateSpaceUser:
service.is_initialized = AsyncMock(return_value=True) # Already initialized
# Execute & Verify
with pytest.raises(AccountEmailMismatchError):
with pytest.raises(SpaceAccountNotRegisteredError):
await service.create_or_update_space_user(
space_account_uuid='unknown-space-uuid',
email='unknown@example.com',
@@ -548,6 +800,78 @@ class TestUserServiceCreateOrUpdateSpaceUser:
expires_in=3600,
)
async def test_unknown_space_subject_cannot_claim_existing_account_by_email(self):
"""An OAuth login collision requires the explicit account-bound bind flow."""
existing_user = _create_mock_user(
email='owner@example.com',
account_type='local',
space_account_uuid=None,
)
ap = SimpleNamespace(
persistence_mgr=SimpleNamespace(execute_async=AsyncMock()),
provider_service=SimpleNamespace(update_space_model_provider_api_keys=AsyncMock()),
space_service=SimpleNamespace(
get_user_info_raw=AsyncMock(
return_value={
'account': {
'uuid': 'attacker-space-subject',
'email': 'owner@example.com',
},
'api_key': 'attacker-api-key',
}
)
),
)
service = UserService(ap)
service.get_user_by_space_account_uuid = AsyncMock(return_value=None)
service.get_user_by_email = AsyncMock(return_value=existing_user)
service.generate_jwt_token = AsyncMock(return_value='must-not-be-issued')
with pytest.raises(SpaceAccountBindingRequiredError):
await service.authenticate_space_user(
'attacker-access-token',
'attacker-refresh-token',
3600,
)
ap.persistence_mgr.execute_async.assert_not_awaited()
ap.provider_service.update_space_model_provider_api_keys.assert_not_awaited()
service.generate_jwt_token.assert_not_awaited()
async def test_oss_space_provider_refresh_requires_workspace_owner(self):
member_account = _create_mock_user(email='member@example.com', space_account_uuid='space-member')
access = SimpleNamespace(
workspace=SimpleNamespace(uuid='workspace-a'),
membership=SimpleNamespace(role='admin'),
)
provider_service = SimpleNamespace(update_space_model_provider_api_keys=AsyncMock())
ap = SimpleNamespace(
workspace_service=SimpleNamespace(policy=SimpleNamespace(multi_workspace_enabled=False)),
workspace_collaboration_service=SimpleNamespace(list_account_workspaces=AsyncMock(return_value=[access])),
provider_service=provider_service,
)
await UserService(ap)._update_space_provider_for_account(member_account, 'member-api-key')
provider_service.update_space_model_provider_api_keys.assert_not_awaited()
async def test_oss_space_provider_refresh_uses_workspace_owner_credentials(self):
owner_account = _create_mock_user(email='owner@example.com', space_account_uuid='space-owner')
access = SimpleNamespace(
workspace=SimpleNamespace(uuid='workspace-a'),
membership=SimpleNamespace(role='owner'),
)
provider_service = SimpleNamespace(update_space_model_provider_api_keys=AsyncMock())
ap = SimpleNamespace(
workspace_service=SimpleNamespace(policy=SimpleNamespace(multi_workspace_enabled=False)),
workspace_collaboration_service=SimpleNamespace(list_account_workspaces=AsyncMock(return_value=[access])),
provider_service=provider_service,
)
await UserService(ap)._update_space_provider_for_account(owner_account, 'owner-api-key')
provider_service.update_space_model_provider_api_keys.assert_awaited_once_with('workspace-a', 'owner-api-key')
async def test_create_or_update_space_user_no_expiry(self):
"""Creates Space user without token expiry."""
# Setup
@@ -594,6 +918,58 @@ class TestUserServiceCreateOrUpdateSpaceUser:
assert result is not None
assert result.space_account_uuid == 'noexpiry-uuid'
async def test_bind_space_account_rejects_different_email(self):
service = UserService(SimpleNamespace())
service.get_user_by_email = AsyncMock(return_value=_create_mock_user(email='invited@example.com'))
service.ap.space_service = SimpleNamespace(
exchange_oauth_code=AsyncMock(
return_value={'access_token': 'access', 'refresh_token': 'refresh', 'expires_in': 3600}
),
get_user_info_raw=AsyncMock(
return_value={
'account': {'uuid': 'space-other', 'email': 'other@example.com'},
'api_key': 'key',
}
),
)
service.get_user_by_space_account_uuid = AsyncMock(return_value=None)
service._identity_execute = AsyncMock()
with pytest.raises(AccountEmailMismatchError):
await service.bind_space_account('invited@example.com', 'code')
service._identity_execute.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_workspace_owner_returns_user_object_from_core_connection_result(self):
service = UserService(SimpleNamespace())
owner = _create_mock_user('owner@example.com', password='pw')
service.ap.persistence_mgr = SimpleNamespace(
current_session=lambda: SimpleNamespace(scalar=AsyncMock(return_value=owner)),
)
resolved = await service.get_workspace_owner('workspace-1')
assert resolved is owner
class TestUserServiceLoginCapabilities:
async def test_capabilities_are_derived_from_all_accounts(self):
result = SimpleNamespace(one=lambda: (2, 1))
ap = SimpleNamespace(persistence_mgr=SimpleNamespace(execute_async=AsyncMock(return_value=result)))
capabilities = await UserService(ap).get_login_capabilities()
assert capabilities == {'password_login_enabled': True, 'space_login_enabled': True}
async def test_capabilities_disable_absent_login_methods(self):
result = SimpleNamespace(one=lambda: (0, 0))
ap = SimpleNamespace(persistence_mgr=SimpleNamespace(execute_async=AsyncMock(return_value=result)))
capabilities = await UserService(ap).get_login_capabilities()
assert capabilities == {'password_login_enabled': False, 'space_login_enabled': False}
class TestUserServiceCreateUserLock:
"""Tests for create_user_lock attribute."""
@@ -14,15 +14,23 @@ Source: src/langbot/pkg/api/http/service/webhook.py
from __future__ import annotations
import datetime
import pytest
import sqlalchemy
from sqlalchemy.ext.asyncio import create_async_engine
from unittest.mock import AsyncMock, Mock
from types import SimpleNamespace
from langbot.pkg.api.http.authz import WorkspaceRequiredError
from langbot.pkg.api.http.service.webhook import WebhookService
from langbot.pkg.entity.persistence.base import Base
from langbot.pkg.entity.persistence.webhook import Webhook
from langbot.pkg.entity.persistence.workspace import Workspace
pytestmark = pytest.mark.asyncio
WORKSPACE_UUID = 'workspace-a'
def _create_mock_webhook(
@@ -42,11 +50,20 @@ def _create_mock_webhook(
return webhook
def _create_mock_result(items: list = None, first_item=None):
def _create_mock_result(items: list = None, first_item=None, scalar_value=None):
"""Create mock result object for persistence queries."""
result = Mock()
result.all = Mock(return_value=items or [])
result.first = Mock(return_value=first_item)
result.scalar = Mock(return_value=scalar_value)
result.rowcount = 1
return result
def _create_write_result(rowcount: int = 1, inserted_id: int = 1):
result = Mock()
result.rowcount = rowcount
result.inserted_primary_key = [inserted_id]
return result
@@ -71,7 +88,7 @@ class TestWebhookServiceGetWebhooks:
service = WebhookService(ap)
# Execute
result = await service.get_webhooks()
result = await service.get_webhooks(WORKSPACE_UUID)
# Verify
assert result == []
@@ -100,7 +117,7 @@ class TestWebhookServiceGetWebhooks:
service = WebhookService(ap)
# Execute
result = await service.get_webhooks()
result = await service.get_webhooks(WORKSPACE_UUID)
# Verify
assert len(result) == 2
@@ -119,6 +136,7 @@ class TestWebhookServiceCreateWebhook:
# Mock insert result
insert_result = Mock()
insert_result.inserted_primary_key = [1]
# Mock select result for retrieving created webhook
created_webhook = _create_mock_webhook(
@@ -137,6 +155,8 @@ class TestWebhookServiceCreateWebhook:
nonlocal call_count
call_count += 1
if call_count == 1:
return _create_mock_result(scalar_value=0) # Count
if call_count == 2:
return insert_result # Insert
return select_result # Select
@@ -155,6 +175,7 @@ class TestWebhookServiceCreateWebhook:
# Execute
result = await service.create_webhook(
WORKSPACE_UUID,
name='New Webhook',
url='http://new.example.com/webhook',
description='New Description',
@@ -187,7 +208,9 @@ class TestWebhookServiceCreateWebhook:
nonlocal call_count
call_count += 1
if call_count == 1:
return Mock() # Insert
return _create_mock_result(scalar_value=0)
if call_count == 2:
return _create_write_result() # Insert
return _create_mock_result(first_item=created_webhook)
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
@@ -204,7 +227,11 @@ class TestWebhookServiceCreateWebhook:
service = WebhookService(ap)
# Execute - only name and url required
result = await service.create_webhook(name='Minimal Webhook', url='http://minimal.example.com')
result = await service.create_webhook(
WORKSPACE_UUID,
name='Minimal Webhook',
url='http://minimal.example.com',
)
# Verify defaults
assert result['description'] == ''
@@ -224,7 +251,9 @@ class TestWebhookServiceCreateWebhook:
nonlocal call_count
call_count += 1
if call_count == 1:
return Mock()
return _create_mock_result(scalar_value=0)
if call_count == 2:
return _create_write_result()
return _create_mock_result(first_item=created_webhook)
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
@@ -233,11 +262,52 @@ class TestWebhookServiceCreateWebhook:
service = WebhookService(ap)
# Execute
result = await service.create_webhook(name='Disabled', url='http://disabled.com', enabled=False)
result = await service.create_webhook(
WORKSPACE_UUID,
name='Disabled',
url='http://disabled.com',
enabled=False,
)
# Verify
assert result['enabled'] is False
async def test_create_webhook_rejects_workspace_at_capacity(self):
ap = SimpleNamespace(
instance_config=SimpleNamespace(
data={'webhooks': {'max_per_workspace': 2}},
),
persistence_mgr=SimpleNamespace(
execute_async=AsyncMock(return_value=_create_mock_result(scalar_value=2)),
),
)
service = WebhookService(ap)
with pytest.raises(ValueError, match=r'Maximum number of webhooks \(2\) reached'):
await service.create_webhook(
WORKSPACE_UUID,
name='Too many',
url='https://example.invalid',
)
ap.persistence_mgr.execute_async.assert_awaited_once()
async def test_max_per_workspace_clamps_invalid_and_oversized_values(self):
ap = SimpleNamespace(
instance_config=SimpleNamespace(
data={'webhooks': {'max_per_workspace': 999999}},
)
)
service = WebhookService(ap)
assert service.max_per_workspace() == 64
ap.instance_config.data['webhooks']['max_per_workspace'] = 0
assert service.max_per_workspace() == 1
ap.instance_config.data['webhooks']['max_per_workspace'] = 'invalid'
assert service.max_per_workspace() == 16
class TestWebhookServiceGetWebhook:
"""Tests for get_webhook method."""
@@ -262,7 +332,7 @@ class TestWebhookServiceGetWebhook:
service = WebhookService(ap)
# Execute
result = await service.get_webhook(1)
result = await service.get_webhook(WORKSPACE_UUID, 1)
# Verify
assert result is not None
@@ -281,7 +351,7 @@ class TestWebhookServiceGetWebhook:
service = WebhookService(ap)
# Execute
result = await service.get_webhook(999)
result = await service.get_webhook(WORKSPACE_UUID, 999)
# Verify
assert result is None
@@ -298,7 +368,7 @@ class TestWebhookServiceGetWebhook:
service = WebhookService(ap)
# Execute
result = await service.get_webhook(0)
result = await service.get_webhook(WORKSPACE_UUID, 0)
# Verify - should return None (no webhook with ID 0)
assert result is None
@@ -312,12 +382,12 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
await service.update_webhook(1, name='Updated Name')
await service.update_webhook(WORKSPACE_UUID, 1, name='Updated Name')
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
@@ -327,12 +397,12 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
await service.update_webhook(1, url='http://updated.example.com')
await service.update_webhook(WORKSPACE_UUID, 1, url='http://updated.example.com')
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
@@ -342,12 +412,12 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
await service.update_webhook(1, description='Updated description')
await service.update_webhook(WORKSPACE_UUID, 1, description='Updated description')
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
@@ -357,12 +427,12 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
await service.update_webhook(1, enabled=False)
await service.update_webhook(WORKSPACE_UUID, 1, enabled=False)
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
@@ -372,12 +442,13 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
await service.update_webhook(
WORKSPACE_UUID,
1,
name='All Updated',
url='http://all.updated.com',
@@ -393,15 +464,17 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
existing = _create_mock_webhook(webhook_id=1)
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result(first_item=existing))
ap.persistence_mgr.serialize_model = Mock(return_value={'id': 1})
service = WebhookService(ap)
# Execute - no update parameters
await service.update_webhook(1)
await service.update_webhook(WORKSPACE_UUID, 1)
# Verify - no execute call since no update_data
ap.persistence_mgr.execute_async.assert_not_called()
# No write is issued; one scoped existence lookup is performed.
ap.persistence_mgr.execute_async.assert_called_once()
class TestWebhookServiceDeleteWebhook:
@@ -412,12 +485,12 @@ class TestWebhookServiceDeleteWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
await service.delete_webhook(1)
await service.delete_webhook(WORKSPACE_UUID, 1)
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
@@ -427,12 +500,12 @@ class TestWebhookServiceDeleteWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
ap.persistence_mgr.execute_async = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result(rowcount=0))
service = WebhookService(ap)
# Execute - should not raise
await service.delete_webhook(999)
await service.delete_webhook(WORKSPACE_UUID, 999)
# Verify - still called
ap.persistence_mgr.execute_async.assert_called_once()
@@ -453,7 +526,7 @@ class TestWebhookServiceGetEnabledWebhooks:
service = WebhookService(ap)
# Execute
result = await service.get_enabled_webhooks()
result = await service.get_enabled_webhooks(WORKSPACE_UUID)
# Verify
assert result == []
@@ -481,7 +554,7 @@ class TestWebhookServiceGetEnabledWebhooks:
service = WebhookService(ap)
# Execute
result = await service.get_enabled_webhooks()
result = await service.get_enabled_webhooks(WORKSPACE_UUID)
# Verify
assert len(result) == 2
@@ -501,7 +574,170 @@ class TestWebhookServiceGetEnabledWebhooks:
service = WebhookService(ap)
# Execute
result = await service.get_enabled_webhooks()
result = await service.get_enabled_webhooks(WORKSPACE_UUID)
# Verify - should be empty (SQL would filter disabled)
assert result == []
ISOLATION_WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
ISOLATION_WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
class _RealPersistenceManager:
def __init__(self, engine):
self.engine = engine
async def execute_async(self, *args, **kwargs):
async with self.engine.connect() as connection:
result = await connection.execute(*args, **kwargs)
await connection.commit()
return result
@staticmethod
def serialize_model(model, data, masked_columns=None):
return {
column.name: (
getattr(data, column.name).isoformat()
if isinstance(getattr(data, column.name), datetime.datetime)
else getattr(data, column.name)
)
for column in model.__table__.columns
if column.name not in (masked_columns or [])
}
@pytest.fixture
async def tenant_webhook_service(tmp_path):
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "webhooks.db"}')
async with engine.begin() as connection:
await connection.run_sync(Base.metadata.create_all)
await connection.execute(
sqlalchemy.insert(Workspace),
[
{
'uuid': ISOLATION_WORKSPACE_A,
'instance_uuid': 'instance',
'name': 'A',
'slug': 'a',
'source': 'cloud_projection',
},
{
'uuid': ISOLATION_WORKSPACE_B,
'instance_uuid': 'instance',
'name': 'B',
'slug': 'b',
'source': 'cloud_projection',
},
],
)
service = WebhookService(SimpleNamespace(persistence_mgr=_RealPersistenceManager(engine)))
yield service
await engine.dispose()
async def test_webhook_service_requires_workspace(tenant_webhook_service):
with pytest.raises(WorkspaceRequiredError):
await tenant_webhook_service.get_webhooks(None)
async def test_same_name_webhooks_are_isolated(tenant_webhook_service):
created_a = await tenant_webhook_service.create_webhook(
ISOLATION_WORKSPACE_A,
'deploy',
'https://a.invalid',
)
created_b = await tenant_webhook_service.create_webhook(
ISOLATION_WORKSPACE_B,
'deploy',
'https://b.invalid',
)
assert created_a['workspace_uuid'] == ISOLATION_WORKSPACE_A
assert created_b['workspace_uuid'] == ISOLATION_WORKSPACE_B
assert [item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_A)] == ['***']
assert [item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_B)] == ['***']
assert [
item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_A, include_secret=True)
] == ['https://a.invalid']
assert [
item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_B, include_secret=True)
] == ['https://b.invalid']
async def test_cross_workspace_id_guessing_is_not_found(tenant_webhook_service):
created = await tenant_webhook_service.create_webhook(
ISOLATION_WORKSPACE_A,
'secret',
'https://a.invalid/hook',
)
webhook_id = created['id']
assert await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_B, webhook_id) is None
assert not await tenant_webhook_service.update_webhook(
ISOLATION_WORKSPACE_B,
webhook_id,
name='stolen',
)
assert not await tenant_webhook_service.delete_webhook(ISOLATION_WORKSPACE_B, webhook_id)
assert (await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_A, webhook_id))['name'] == 'secret'
async def test_update_and_delete_are_scoped(tenant_webhook_service):
created = await tenant_webhook_service.create_webhook(
ISOLATION_WORKSPACE_A,
'old',
'https://a.invalid/old',
)
assert await tenant_webhook_service.update_webhook(
ISOLATION_WORKSPACE_A,
created['id'],
name='new',
enabled=False,
)
assert await tenant_webhook_service.get_enabled_webhooks(ISOLATION_WORKSPACE_A) == []
assert await tenant_webhook_service.delete_webhook(ISOLATION_WORKSPACE_A, created['id'])
assert await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_A, created['id']) is None
async def test_masked_webhook_url_roundtrip_preserves_replace_and_clear(tenant_webhook_service):
created = await tenant_webhook_service.create_webhook(
ISOLATION_WORKSPACE_A,
'roundtrip',
'https://a.invalid/bearer-secret',
)
masked = await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_A, created['id'])
assert masked['url'] == '***'
assert await tenant_webhook_service.update_webhook(
ISOLATION_WORKSPACE_A,
created['id'],
name='preserved',
url=masked['url'],
)
preserved = await tenant_webhook_service.get_webhook(
ISOLATION_WORKSPACE_A,
created['id'],
include_secret=True,
)
assert preserved['url'] == 'https://a.invalid/bearer-secret'
assert await tenant_webhook_service.update_webhook(
ISOLATION_WORKSPACE_A,
created['id'],
url='https://a.invalid/replacement',
)
replaced = await tenant_webhook_service.get_webhook(
ISOLATION_WORKSPACE_A,
created['id'],
include_secret=True,
)
assert replaced['url'] == 'https://a.invalid/replacement'
assert await tenant_webhook_service.update_webhook(ISOLATION_WORKSPACE_A, created['id'], url='')
cleared = await tenant_webhook_service.get_webhook(
ISOLATION_WORKSPACE_A,
created['id'],
include_secret=True,
)
assert cleared['url'] == ''
@@ -0,0 +1,254 @@
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import lark_oapi
import pytest
import quart
from langbot.pkg.api.http.context import (
PrincipalContext,
PrincipalType,
RequestContext,
WorkspaceContext,
)
from langbot.pkg.api.http.controller.groups.platform.adapters import (
AdaptersRouterGroup,
_AdapterSessionScope,
_bind_session_scope,
_get_owned_session,
_make_room_for_session,
_pop_owned_session,
_start_adapter_session_task,
)
pytestmark = pytest.mark.asyncio
SENSITIVE_ADAPTER_ROUTES = (
('post', '/api/v1/platform/adapters/lark/create-app'),
('get', '/api/v1/platform/adapters/lark/create-app/status/missing'),
('delete', '/api/v1/platform/adapters/lark/create-app/missing'),
('post', '/api/v1/platform/adapters/weixin/login'),
('get', '/api/v1/platform/adapters/weixin/login/status/missing'),
('delete', '/api/v1/platform/adapters/weixin/login/missing'),
('post', '/api/v1/platform/adapters/dingtalk/create-app'),
('get', '/api/v1/platform/adapters/dingtalk/create-app/status/missing'),
('delete', '/api/v1/platform/adapters/dingtalk/create-app/missing'),
('post', '/api/v1/platform/adapters/wecombot/create-bot'),
('get', '/api/v1/platform/adapters/wecombot/create-bot/status/missing'),
('delete', '/api/v1/platform/adapters/wecombot/create-bot/missing'),
('post', '/api/v1/platform/adapters/qqofficial/bind'),
('get', '/api/v1/platform/adapters/qqofficial/bind/status/missing'),
('delete', '/api/v1/platform/adapters/qqofficial/bind/missing'),
)
def _request_context(
*,
account_uuid: str = 'account-a',
workspace_uuid: str = 'workspace-a',
placement_generation: int = 1,
) -> RequestContext:
return RequestContext(
instance_uuid='instance-test',
placement_generation=placement_generation,
request_id='request-test',
auth_type='user-token',
principal=PrincipalContext(
principal_type=PrincipalType.ACCOUNT,
account_uuid=account_uuid,
),
workspace=WorkspaceContext(
workspace_uuid=workspace_uuid,
membership_uuid='membership-test',
role='developer',
permissions=frozenset({'resource.manage'}),
),
)
async def _create_client(*, role: str = 'developer'):
quart_app = quart.Quart(__name__)
accounts = {
'owner-token': SimpleNamespace(uuid='account-a', user='owner@example.com'),
'other-token': SimpleNamespace(uuid='account-b', user='other@example.com'),
}
async def get_authenticated_account(token: str):
return accounts[token]
async def resolve_account_workspace(account_uuid: str, requested_workspace_uuid: str | None):
workspace_uuid = requested_workspace_uuid or 'workspace-a'
return SimpleNamespace(
execution=SimpleNamespace(
instance_uuid='instance-test',
placement_generation=1,
),
workspace=SimpleNamespace(uuid=workspace_uuid),
membership=SimpleNamespace(
uuid=f'membership-{account_uuid}-{workspace_uuid}',
role=role,
projection_revision=1,
),
)
class TestTaskManager:
def create_user_task(self, coro, **_kwargs):
return SimpleNamespace(task=asyncio.create_task(coro))
application = SimpleNamespace(
user_service=SimpleNamespace(
get_authenticated_account=AsyncMock(side_effect=get_authenticated_account),
),
workspace_collaboration_service=SimpleNamespace(
resolve_account_workspace=AsyncMock(side_effect=resolve_account_workspace),
),
platform_mgr=SimpleNamespace(),
task_mgr=TestTaskManager(),
)
router = AdaptersRouterGroup(application, quart_app)
await router.initialize()
return quart_app.test_client()
@pytest.mark.parametrize(('method', 'path'), SENSITIVE_ADAPTER_ROUTES)
async def test_sensitive_adapter_flows_require_resource_manage(method: str, path: str):
client = await _create_client(role='viewer')
response = await getattr(client, method)(
path,
headers={'Authorization': 'Bearer owner-token'},
)
assert response.status_code == 403
assert (await response.get_json())['code'] == 'permission_denied'
async def test_session_scope_matches_exact_tenant_placement_and_principal():
owner_context = _request_context()
sessions: dict[str, dict] = {'session-test': {'status': 'waiting'}}
_bind_session_scope(sessions['session-test'], owner_context)
assert sessions['session-test']['scope'] == _AdapterSessionScope.from_request_context(owner_context)
assert _get_owned_session(sessions, 'session-test', owner_context) is sessions['session-test']
for other_context in (
_request_context(account_uuid='account-b'),
_request_context(workspace_uuid='workspace-b'),
_request_context(placement_generation=2),
):
assert _get_owned_session(sessions, 'session-test', other_context) is None
assert _pop_owned_session(sessions, 'session-test', other_context) is None
assert 'session-test' in sessions
assert _pop_owned_session(sessions, 'session-test', owner_context) is not None
assert sessions == {}
async def test_session_capacity_evicts_oldest_session_in_same_workspace():
owner_context = _request_context()
sessions: dict[str, dict] = {}
tasks = []
for index in range(10):
task = SimpleNamespace(done=Mock(return_value=False), cancel=Mock())
tasks.append(task)
session = {'created_at': float(index), 'task': task}
_bind_session_scope(session, owner_context)
sessions[f'session-{index}'] = session
_make_room_for_session(sessions, owner_context)
assert 'session-0' not in sessions
assert len(sessions) == 9
tasks[0].cancel.assert_called_once_with()
async def test_adapter_session_task_uses_tenant_task_admission():
blocker = asyncio.Event()
async def credential_exchange():
await blocker.wait()
task_manager = SimpleNamespace(create_user_task=Mock())
def create_user_task(coro, **_kwargs):
return SimpleNamespace(task=asyncio.create_task(coro))
task_manager.create_user_task.side_effect = create_user_task
application = SimpleNamespace(task_mgr=task_manager)
request_context = _request_context()
returned = _start_adapter_session_task(
application,
credential_exchange(),
adapter='lark',
session_id='session-test',
request_context=request_context,
)
assert returned is not None
task_manager.create_user_task.assert_called_once()
kwargs = task_manager.create_user_task.call_args.kwargs
assert kwargs['kind'] == 'platform-adapter-credential-exchange'
assert kwargs['instance_uuid'] == request_context.instance_uuid
assert kwargs['workspace_uuid'] == request_context.workspace_uuid
assert kwargs['placement_generation'] == request_context.placement_generation
blocker.set()
await returned
async def test_lark_session_status_and_delete_hide_cross_scope_sessions(monkeypatch):
registration_blocker = asyncio.Event()
async def fake_register_app(*, on_qr_code, source: str):
assert source == 'langbot'
on_qr_code({'url': 'https://example.test/lark-qr'})
await registration_blocker.wait()
raise AssertionError('registration should have been cancelled')
monkeypatch.setattr(lark_oapi, 'aregister_app', fake_register_app)
client = await _create_client()
owner_headers = {
'Authorization': 'Bearer owner-token',
'X-Workspace-Id': 'workspace-a',
}
create_response = await client.post(
'/api/v1/platform/adapters/lark/create-app',
headers=owner_headers,
)
assert create_response.status_code == 200
session_id = (await create_response.get_json())['data']['session_id']
status_path = f'/api/v1/platform/adapters/lark/create-app/status/{session_id}'
delete_path = f'/api/v1/platform/adapters/lark/create-app/{session_id}'
for headers in (
{
'Authorization': 'Bearer other-token',
'X-Workspace-Id': 'workspace-a',
},
{
'Authorization': 'Bearer owner-token',
'X-Workspace-Id': 'workspace-b',
},
):
status_response = await client.get(status_path, headers=headers)
delete_response = await client.delete(delete_path, headers=headers)
assert status_response.status_code == 404
assert delete_response.status_code == 404
assert (await status_response.get_json())['msg'] == 'Session not found'
assert (await delete_response.get_json())['msg'] == 'Session not found'
owner_status_response = await client.get(status_path, headers=owner_headers)
assert owner_status_response.status_code == 200
assert (await owner_status_response.get_json())['data']['status'] == 'waiting'
owner_delete_response = await client.delete(delete_path, headers=owner_headers)
assert owner_delete_response.status_code == 200
missing_delete_response = await client.delete(delete_path, headers=owner_headers)
assert missing_delete_response.status_code == 404
await asyncio.sleep(0)
+31 -5
View File
@@ -4,7 +4,7 @@ import sys
import types
from importlib import import_module
from types import SimpleNamespace
from unittest.mock import AsyncMock
from unittest.mock import ANY, AsyncMock
import pytest
import quart
@@ -19,11 +19,35 @@ pytestmark = pytest.mark.asyncio
async def _create_test_client(agent_service: SimpleNamespace):
app = quart.Quart(__name__)
user_service = SimpleNamespace(
verify_jwt_token=AsyncMock(return_value='test@example.com'),
get_user_by_email=AsyncMock(return_value=SimpleNamespace(user='test@example.com')),
account = SimpleNamespace(
uuid='account-test',
user='test@example.com',
)
user_service = SimpleNamespace(
get_authenticated_account=AsyncMock(return_value=account),
)
access = SimpleNamespace(
workspace=SimpleNamespace(uuid='workspace-test'),
membership=SimpleNamespace(
uuid='membership-test',
role='developer',
projection_revision=1,
),
execution=SimpleNamespace(
instance_uuid='instance-test',
placement_generation=1,
),
)
ap = SimpleNamespace(
agent_service=agent_service,
user_service=user_service,
apikey_service=SimpleNamespace(
authenticate_api_key=AsyncMock(return_value=None)
),
workspace_collaboration_service=SimpleNamespace(
resolve_account_workspace=AsyncMock(return_value=access)
),
)
ap = SimpleNamespace(agent_service=agent_service, user_service=user_service)
AgentsRouterGroup = import_module('langbot.pkg.api.http.controller.groups.agents').AgentsRouterGroup
group = AgentsRouterGroup(ap, app)
await group.initialize()
@@ -44,6 +68,7 @@ async def test_create_agent_returns_bad_request_for_invalid_runner_config():
assert response.status_code == 400
assert await response.get_json() == {'code': -1, 'msg': message}
agent_service.create_agent.assert_awaited_once_with(
ANY,
{'name': 'Invalid Agent', 'config': {'runner_config': []}},
)
@@ -62,6 +87,7 @@ async def test_update_agent_returns_bad_request_for_invalid_runner_config():
assert response.status_code == 400
assert await response.get_json() == {'code': -1, 'msg': message}
agent_service.update_agent.assert_awaited_once_with(
ANY,
'agent-1',
{'config': {'runner': {'id': 7}}},
)
+56 -15
View File
@@ -6,6 +6,7 @@ from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.api.http.service.apikey import ApiKeyService
from langbot.pkg.entity.persistence.apikey import ApiKeyStatus
@pytest.mark.asyncio
@@ -13,30 +14,70 @@ from langbot.pkg.api.http.service.apikey import ApiKeyService
async def test_verify_api_key_rejects_non_lbk_keys_without_db_query(api_key):
persistence_mgr = SimpleNamespace(execute_async=AsyncMock())
instance_config = SimpleNamespace(data={'api': {'global_api_key': ''}})
service = ApiKeyService(SimpleNamespace(persistence_mgr=persistence_mgr, instance_config=instance_config))
workspace_service = SimpleNamespace(get_execution_binding=AsyncMock())
service = ApiKeyService(
SimpleNamespace(
persistence_mgr=persistence_mgr,
instance_config=instance_config,
workspace_service=workspace_service,
)
)
result = await service.verify_api_key(api_key)
assert result is False
persistence_mgr.execute_async.assert_not_awaited()
workspace_service.get_execution_binding.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
('db_row', 'expected'),
[
(object(), True),
(None, False),
],
)
async def test_verify_api_key_keeps_db_validation_for_lbk_keys(db_row, expected):
query_result = Mock()
query_result.first.return_value = db_row
persistence_mgr = SimpleNamespace(execute_async=AsyncMock(return_value=query_result))
@pytest.mark.parametrize('key_exists', [True, False])
async def test_verify_api_key_keeps_db_validation_for_lbk_keys(key_exists):
key = (
SimpleNamespace(
id=1,
uuid='key-uuid',
workspace_uuid='workspace-a',
status=ApiKeyStatus.ACTIVE.value,
expires_at=None,
scopes=[],
)
if key_exists
else None
)
discovery_result = Mock()
discovery_result.first.return_value = key
query_results = [discovery_result]
if key_exists:
scoped_result = Mock()
scoped_result.first.return_value = key
update_result = Mock()
update_result.scalar_one_or_none.return_value = key.id
query_results.extend([scoped_result, update_result])
persistence_mgr = SimpleNamespace(execute_async=AsyncMock(side_effect=query_results))
instance_config = SimpleNamespace(data={'api': {'global_api_key': ''}})
service = ApiKeyService(SimpleNamespace(persistence_mgr=persistence_mgr, instance_config=instance_config))
workspace_service = SimpleNamespace(
get_execution_binding=AsyncMock(
return_value=SimpleNamespace(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
)
)
service = ApiKeyService(
SimpleNamespace(
persistence_mgr=persistence_mgr,
instance_config=instance_config,
workspace_service=workspace_service,
)
)
result = await service.verify_api_key('lbk_valid_format')
assert result is expected
persistence_mgr.execute_async.assert_awaited_once()
assert result is key_exists
assert persistence_mgr.execute_async.await_count == (3 if key_exists else 1)
if key_exists:
workspace_service.get_execution_binding.assert_awaited_once_with('workspace-a')
else:
workspace_service.get_execution_binding.assert_not_awaited()
@@ -0,0 +1,113 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import quart
from langbot.pkg.api.http.controller.groups.platform.bots import BotsRouterGroup
pytestmark = pytest.mark.asyncio
SECRET_CONFIG = {'token': 'tenant-secret', 'app_secret': 'also-secret'}
async def create_client(*, role: str):
quart_app = quart.Quart(__name__)
account = SimpleNamespace(uuid='account-test', user='test@example.com')
user_service = SimpleNamespace(
get_authenticated_account=AsyncMock(return_value=account),
)
access = SimpleNamespace(
execution=SimpleNamespace(
instance_uuid='instance-test',
placement_generation=1,
),
workspace=SimpleNamespace(uuid='workspace-test'),
membership=SimpleNamespace(
uuid='membership-test',
role=role,
projection_revision=1,
),
)
async def get_bots(_context, *, include_secret=False):
bot = {'uuid': 'bot-test', 'name': 'Test Bot'}
if include_secret:
bot['adapter_config'] = SECRET_CONFIG
return [bot]
async def get_runtime_bot_info(_context, _bot_uuid, *, include_secret=False):
bot = {'uuid': 'bot-test', 'name': 'Test Bot'}
if include_secret:
bot['adapter_config'] = SECRET_CONFIG
return bot
bot_service = SimpleNamespace(
get_bots=AsyncMock(side_effect=get_bots),
get_runtime_bot_info=AsyncMock(side_effect=get_runtime_bot_info),
update_bot=AsyncMock(),
)
application = SimpleNamespace(
user_service=user_service,
apikey_service=SimpleNamespace(
authenticate_api_key=AsyncMock(
return_value=SimpleNamespace(
instance_uuid='instance-test',
placement_generation=1,
api_key_uuid='api-key-test',
workspace_uuid='workspace-test',
permissions=frozenset({'resource.view'}),
)
)
),
workspace_collaboration_service=SimpleNamespace(resolve_account_workspace=AsyncMock(return_value=access)),
bot_service=bot_service,
)
router = BotsRouterGroup(application, quart_app)
await router.initialize()
return quart_app.test_client(), bot_service
async def test_viewer_list_and_detail_never_receive_adapter_credentials():
client, bot_service = await create_client(role='viewer')
headers = {'Authorization': 'Bearer test-token'}
list_response = await client.get('/api/v1/platform/bots', headers=headers)
detail_response = await client.get('/api/v1/platform/bots/bot-test', headers=headers)
assert list_response.status_code == 200
assert detail_response.status_code == 200
assert 'adapter_config' not in (await list_response.get_json())['data']['bots'][0]
assert 'adapter_config' not in (await detail_response.get_json())['data']['bot']
assert bot_service.get_bots.await_args.kwargs['include_secret'] is False
assert bot_service.get_runtime_bot_info.await_args.kwargs['include_secret'] is False
async def test_resource_manager_can_read_adapter_credentials():
client, bot_service = await create_client(role='developer')
headers = {'Authorization': 'Bearer test-token'}
list_response = await client.get('/api/v1/platform/bots', headers=headers)
detail_response = await client.get('/api/v1/platform/bots/bot-test', headers=headers)
assert (await list_response.get_json())['data']['bots'][0]['adapter_config'] == SECRET_CONFIG
assert (await detail_response.get_json())['data']['bot']['adapter_config'] == SECRET_CONFIG
assert bot_service.get_bots.await_args.kwargs['include_secret'] is True
assert bot_service.get_runtime_bot_info.await_args.kwargs['include_secret'] is True
async def test_viewer_cannot_write_adapter_credentials():
client, bot_service = await create_client(role='viewer')
response = await client.put(
'/api/v1/platform/bots/bot-test',
headers={'Authorization': 'Bearer test-token'},
json={'adapter_config': SECRET_CONFIG},
)
assert response.status_code == 403
assert (await response.get_json())['code'] == 'permission_denied'
bot_service.update_bot.assert_not_awaited()
@@ -0,0 +1,142 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import quart
import sqlalchemy
from sqlalchemy.ext.asyncio import create_async_engine
from langbot.pkg.api.http.controller.groups.extensions import ExtensionsRouterGroup
from langbot.pkg.persistence.mgr import PersistenceManager
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
@pytest.mark.asyncio
async def test_extensions_route_hides_runtime_bound_to_another_workspace():
account = SimpleNamespace(uuid='account-a', user='owner@example.com')
connector = SimpleNamespace(
is_enable_plugin=True,
require_workspace_context=AsyncMock(side_effect=WorkspaceNotFoundError('Plugin resource not found')),
list_plugins=AsyncMock(return_value=[]),
)
ap = SimpleNamespace(
user_service=SimpleNamespace(
get_authenticated_account=AsyncMock(return_value=account),
),
workspace_collaboration_service=SimpleNamespace(
resolve_account_workspace=AsyncMock(
return_value=SimpleNamespace(
workspace=SimpleNamespace(uuid='workspace-a'),
membership=SimpleNamespace(uuid='membership-a', role='owner', projection_revision=0),
execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=2),
)
)
),
plugin_connector=connector,
mcp_service=SimpleNamespace(get_mcp_servers=AsyncMock(return_value=[])),
skill_service=SimpleNamespace(list_skills=AsyncMock(return_value=[])),
)
quart_app = quart.Quart(__name__)
router = ExtensionsRouterGroup(ap, quart_app)
await router.initialize()
response = await quart_app.test_client().get(
'/api/v1/extensions',
headers={'Authorization': 'Bearer token'},
)
assert response.status_code == 404
connector.list_plugins.assert_not_awaited()
ap.mcp_service.get_mcp_servers.assert_not_awaited()
ap.skill_service.list_skills.assert_not_awaited()
@pytest.mark.asyncio
async def test_extensions_route_redacts_plugin_secrets_without_mutating_runtime_data():
account = SimpleNamespace(uuid='account-a', user='viewer@example.com')
raw_plugin = {
'plugin_config': {'apiKey': 'plugin-secret', 'nested': {'token': 'nested-secret'}},
'debug': {'plugin_debug_key': 'debug-secret'},
}
connector = SimpleNamespace(
is_enable_plugin=True,
require_workspace_context=AsyncMock(),
list_plugins=AsyncMock(return_value=[raw_plugin]),
)
ap = SimpleNamespace(
user_service=SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account)),
workspace_collaboration_service=SimpleNamespace(
resolve_account_workspace=AsyncMock(
return_value=SimpleNamespace(
workspace=SimpleNamespace(uuid='workspace-a'),
membership=SimpleNamespace(uuid='membership-a', role='viewer', projection_revision=0),
execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=2),
)
)
),
plugin_connector=connector,
mcp_service=SimpleNamespace(get_mcp_servers=AsyncMock(return_value=[])),
skill_service=SimpleNamespace(list_skills=AsyncMock(return_value=[])),
)
quart_app = quart.Quart(__name__)
router = ExtensionsRouterGroup(ap, quart_app)
await router.initialize()
response = await quart_app.test_client().get(
'/api/v1/extensions',
headers={'Authorization': 'Bearer token', 'X-Workspace-Id': 'workspace-a'},
)
assert response.status_code == 200
plugin = (await response.get_json())['data']['extensions'][0]['plugin']
assert plugin['plugin_config']['apiKey'] == '***'
assert plugin['plugin_config']['nested']['token'] == '***'
assert plugin['debug']['plugin_debug_key'] == '***'
assert raw_plugin['plugin_config']['apiKey'] == 'plugin-secret'
@pytest.mark.asyncio
async def test_extensions_parallel_reads_open_explicit_child_task_scopes():
engine = create_async_engine('sqlite+aiosqlite:///:memory:')
account = SimpleNamespace(uuid='account-a', user='owner@example.com')
ap = SimpleNamespace()
ap.persistence_mgr = PersistenceManager(ap)
ap.persistence_mgr.db = SimpleNamespace(get_engine=lambda: engine)
ap.user_service = SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account))
ap.workspace_collaboration_service = SimpleNamespace(
resolve_account_workspace=AsyncMock(
return_value=SimpleNamespace(
workspace=SimpleNamespace(uuid='workspace-a'),
membership=SimpleNamespace(uuid='membership-a', role='owner', projection_revision=0),
execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=2),
)
)
)
ap.plugin_connector = SimpleNamespace(
is_enable_plugin=False,
list_plugins=AsyncMock(return_value=[]),
)
async def list_mcp_servers(_context, *, contain_runtime_info):
await ap.persistence_mgr.execute_async(sqlalchemy.select(sqlalchemy.literal(1)))
assert contain_runtime_info is True
return [{'name': 'Scoped MCP'}]
ap.mcp_service = SimpleNamespace(get_mcp_servers=AsyncMock(side_effect=list_mcp_servers))
ap.skill_service = SimpleNamespace(list_skills=AsyncMock(return_value=[]))
quart_app = quart.Quart(__name__)
router = ExtensionsRouterGroup(ap, quart_app)
await router.initialize()
try:
response = await quart_app.test_client().get(
'/api/v1/extensions',
headers={'Authorization': 'Bearer token', 'X-Workspace-Id': 'workspace-a'},
)
finally:
await engine.dispose()
assert response.status_code == 200
assert (await response.get_json())['data']['extensions'] == [{'type': 'mcp', 'server': {'name': 'Scoped MCP'}}]
@@ -0,0 +1,59 @@
from __future__ import annotations
import io
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import quart
from quart.datastructures import FileStorage
from langbot.pkg.api.http.controller.groups.files import FilesRouterGroup
pytestmark = pytest.mark.asyncio
async def test_document_upload_uses_dedicated_scoped_owner_type():
quart_app = quart.Quart(__name__)
account = SimpleNamespace(uuid='account-test', user='test@example.com')
access = SimpleNamespace(
execution=SimpleNamespace(
instance_uuid='instance-test',
placement_generation=3,
),
workspace=SimpleNamespace(uuid='00000000-0000-0000-0000-00000000000a'),
membership=SimpleNamespace(
uuid='membership-test',
role='developer',
projection_revision=1,
),
)
storage_mgr = SimpleNamespace(save_scoped=AsyncMock(return_value='scoped-document-key'))
application = SimpleNamespace(
user_service=SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account)),
workspace_collaboration_service=SimpleNamespace(resolve_account_workspace=AsyncMock(return_value=access)),
storage_mgr=storage_mgr,
)
router = FilesRouterGroup(application, quart_app)
await router.initialize()
client = quart_app.test_client()
response = await client.post(
'/api/v1/files/documents',
headers={'Authorization': 'Bearer test-token'},
files={
'file': FileStorage(
stream=io.BytesIO(b'document bytes'),
filename='report.pdf',
)
},
)
assert response.status_code == 200
assert (await response.get_json())['data']['file_id'] == 'scoped-document-key'
kwargs = storage_mgr.save_scoped.await_args.kwargs
assert kwargs['owner_type'] == 'upload_document'
assert kwargs['owner'] == 'account:account-test'
assert kwargs['key'].endswith('.pdf')
assert kwargs['value'] == b'document bytes'
@@ -0,0 +1,293 @@
from __future__ import annotations
from datetime import datetime
import json
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
import sqlalchemy
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.api.http.controller.groups.knowledge.migration import KnowledgeMigrationRouterGroup
from langbot.pkg.persistence.tenant_uow import _validate_scoped_statement_call
from langbot.pkg.workspace.errors import WorkspaceInvariantError, WorkspaceNotFoundError
CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=3,
)
@pytest.mark.asyncio
async def test_background_migration_propagates_generation_change_before_runtime_call():
connector = SimpleNamespace(
require_workspace_context=AsyncMock(side_effect=[CONTEXT, WorkspaceNotFoundError('Plugin resource not found')]),
list_knowledge_engines=AsyncMock(return_value=[]),
)
router = object.__new__(KnowledgeMigrationRouterGroup)
router.ap = SimpleNamespace(
plugin_connector=connector,
workspace_service=SimpleNamespace(
get_local_execution_binding=AsyncMock(return_value=CONTEXT),
),
logger=Mock(),
)
router._table_exists = AsyncMock(return_value=False)
router._set_migration_flag = AsyncMock()
task_context = SimpleNamespace(trace=Mock())
with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
await router._execute_rag_migration(
CONTEXT,
task_context,
install_plugin=False,
)
assert connector.require_workspace_context.await_count == 2
connector.list_knowledge_engines.assert_not_awaited()
router._set_migration_flag.assert_not_awaited()
@pytest.mark.asyncio
async def test_cloud_migration_is_rejected_before_legacy_table_access():
router = object.__new__(KnowledgeMigrationRouterGroup)
router.ap = SimpleNamespace(
workspace_service=SimpleNamespace(
get_local_execution_binding=AsyncMock(side_effect=WorkspaceInvariantError('not an OSS local workspace')),
),
plugin_connector=SimpleNamespace(require_workspace_context=AsyncMock()),
logger=Mock(),
)
router._table_exists = AsyncMock()
router._set_migration_flag = AsyncMock()
task_context = SimpleNamespace(trace=Mock())
with pytest.raises(WorkspaceNotFoundError, match='migration is unavailable'):
await router._execute_rag_migration(CONTEXT, task_context, install_plugin=False)
router._table_exists.assert_not_awaited()
router.ap.plugin_connector.require_workspace_context.assert_not_awaited()
router._set_migration_flag.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize('database_name, exists', [('postgresql', True), ('sqlite', False)])
async def test_legacy_table_discovery_uses_scoped_structured_queries(database_name: str, exists: bool):
result = Mock()
result.first.return_value = ('knowledge_bases_backup',) if exists else None
execute_async = AsyncMock(return_value=result)
router = object.__new__(KnowledgeMigrationRouterGroup)
router.ap = SimpleNamespace(
persistence_mgr=SimpleNamespace(
db=SimpleNamespace(name=database_name),
execute_async=execute_async,
)
)
assert await router._table_exists('knowledge_bases_backup') is exists
statement = execute_async.await_args.args[0]
assert isinstance(statement, sqlalchemy.sql.selectable.SelectBase)
_validate_scoped_statement_call((statement,), {})
@pytest.mark.asyncio
async def test_legacy_restore_emits_only_scoped_structured_statements():
missing_table_result = Mock()
missing_table_result.first.return_value = None
existing_table_result = Mock()
existing_table_result.first.return_value = ('knowledge_bases_backup',)
backup_result = Mock()
backup_result.keys.return_value = [
'uuid',
'name',
'description',
'emoji',
'embedding_model_uuid',
'top_k',
'created_at',
'updated_at',
]
now = datetime.now()
backup_result.fetchall.return_value = [
('kb-legacy', 'Legacy KB', 'Description', 'U0001f4da', 'embedding-model', 7, now, now)
]
execute_async = AsyncMock(
side_effect=[
missing_table_result,
existing_table_result,
backup_result,
Mock(),
Mock(),
]
)
connector = SimpleNamespace(
require_workspace_context=AsyncMock(return_value=CONTEXT),
list_knowledge_engines=AsyncMock(return_value=[]),
rag_on_kb_create=AsyncMock(),
)
router = object.__new__(KnowledgeMigrationRouterGroup)
router.ap = SimpleNamespace(
workspace_service=SimpleNamespace(
get_local_execution_binding=AsyncMock(return_value=CONTEXT),
),
plugin_connector=connector,
persistence_mgr=SimpleNamespace(
db=SimpleNamespace(name='sqlite'),
execute_async=execute_async,
),
rag_mgr=SimpleNamespace(load_knowledge_bases_from_db=AsyncMock()),
logger=Mock(),
)
task_context = SimpleNamespace(trace=Mock())
await router._execute_rag_migration(CONTEXT, task_context, install_plugin=False)
statements = [call.args[0] for call in execute_async.await_args_list]
assert any(isinstance(statement, sqlalchemy.sql.dml.Insert) for statement in statements)
assert any(isinstance(statement, sqlalchemy.sql.dml.Update) for statement in statements)
for statement in statements:
assert not isinstance(statement, sqlalchemy.sql.elements.TextClause)
_validate_scoped_statement_call((statement,), {})
connector.rag_on_kb_create.assert_awaited_once_with(
'langbot-team/LangRAG',
'kb-legacy',
{'embedding_model_uuid': 'embedding-model'},
)
@pytest.mark.asyncio
async def test_legacy_restore_accepts_sqlite_string_dates_and_text_json_columns():
engine = sqlalchemy.ext.asyncio.create_async_engine('sqlite+aiosqlite:///:memory:')
try:
async with engine.begin() as connection:
await connection.exec_driver_sql(
"""
CREATE TABLE knowledge_bases_backup (
uuid TEXT PRIMARY KEY,
name TEXT,
description TEXT,
emoji TEXT,
embedding_model_uuid TEXT,
top_k INTEGER,
created_at DATETIME,
updated_at DATETIME
)
"""
)
await connection.exec_driver_sql(
"""
CREATE TABLE external_knowledge_bases (
uuid TEXT PRIMARY KEY,
name TEXT,
description TEXT,
emoji TEXT,
plugin_author TEXT,
plugin_name TEXT,
retriever_config TEXT,
created_at DATETIME
)
"""
)
await connection.exec_driver_sql(
"""
CREATE TABLE knowledge_bases (
uuid TEXT PRIMARY KEY,
workspace_uuid TEXT NOT NULL,
name TEXT,
description TEXT,
emoji TEXT,
created_at DATETIME,
updated_at DATETIME,
knowledge_engine_plugin_id TEXT,
collection_id TEXT,
creation_settings TEXT,
retrieval_settings TEXT
)
"""
)
await connection.exec_driver_sql(
'CREATE TABLE workspace_metadata (workspace_uuid TEXT, key TEXT, value TEXT)'
)
legacy_timestamp = '2026-07-20 03:00:00.123456'
await connection.exec_driver_sql(
"""
INSERT INTO knowledge_bases_backup
(uuid, name, description, emoji, embedding_model_uuid, top_k, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
""",
(
'kb-internal',
'Internal',
'Internal legacy KB',
'U0001f4da',
'embedding-model',
5,
legacy_timestamp,
legacy_timestamp,
),
)
await connection.exec_driver_sql(
"""
INSERT INTO external_knowledge_bases
(uuid, name, description, emoji, plugin_author, plugin_name, retriever_config, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
""",
(
'kb-external',
'External',
'External legacy KB',
'U0001f517',
'langbot-team',
'DifyDatasetsRetriever',
json.dumps({'api_base_url': 'https://example.invalid', 'top_k': 8}),
legacy_timestamp,
),
)
connector = SimpleNamespace(
require_workspace_context=AsyncMock(return_value=CONTEXT),
list_knowledge_engines=AsyncMock(return_value=[]),
rag_on_kb_create=AsyncMock(),
)
router = object.__new__(KnowledgeMigrationRouterGroup)
router.ap = SimpleNamespace(
workspace_service=SimpleNamespace(
get_local_execution_binding=AsyncMock(return_value=CONTEXT),
),
plugin_connector=connector,
persistence_mgr=SimpleNamespace(
db=SimpleNamespace(name='sqlite'),
execute_async=connection.execute,
),
rag_mgr=SimpleNamespace(load_knowledge_bases_from_db=AsyncMock()),
logger=Mock(),
)
await router._execute_rag_migration(
CONTEXT,
SimpleNamespace(trace=Mock()),
install_plugin=False,
)
restored = (
await connection.exec_driver_sql(
"""
SELECT uuid, created_at, updated_at, creation_settings, retrieval_settings
FROM knowledge_bases
ORDER BY uuid
"""
)
).all()
assert [row.uuid for row in restored] == ['kb-external', 'kb-internal']
assert all(row.created_at == legacy_timestamp for row in restored)
assert all(row.updated_at == legacy_timestamp for row in restored)
assert json.loads(restored[0].creation_settings)['api_base_url'] == 'https://example.invalid'
assert json.loads(restored[0].retrieval_settings) == {'top_k': 8}
assert json.loads(restored[1].creation_settings) == {'embedding_model_uuid': 'embedding-model'}
assert json.loads(restored[1].retrieval_settings) == {'top_k': 5}
finally:
await engine.dispose()
+63 -7
View File
@@ -17,19 +17,63 @@ sys.modules.setdefault('langbot.pkg.core.app', core_app_module)
pytestmark = pytest.mark.asyncio
async def _create_test_client(mcp_service: SimpleNamespace):
async def _create_test_client(mcp_service: SimpleNamespace, *, role: str = 'owner'):
app = quart.Quart(__name__)
user_service = SimpleNamespace(
verify_jwt_token=AsyncMock(return_value='test@example.com'),
get_user_by_email=AsyncMock(return_value=SimpleNamespace(user='test@example.com')),
get_user_by_email=AsyncMock(
return_value=SimpleNamespace(
user='test@example.com',
uuid='account-a',
)
),
)
workspace_collaboration_service = SimpleNamespace(
resolve_account_workspace=AsyncMock(
return_value=SimpleNamespace(
execution=SimpleNamespace(
instance_uuid='instance-a',
placement_generation=1,
),
workspace=SimpleNamespace(uuid='workspace-a'),
membership=SimpleNamespace(
uuid='membership-a',
role=role,
projection_revision=1,
),
)
)
)
ap = SimpleNamespace(
mcp_service=mcp_service,
user_service=user_service,
workspace_collaboration_service=workspace_collaboration_service,
)
ap = SimpleNamespace(mcp_service=mcp_service, user_service=user_service)
MCPRouterGroup = import_module('langbot.pkg.api.http.controller.groups.resources.mcp').MCPRouterGroup
group = MCPRouterGroup(ap, app)
await group.initialize()
return app.test_client()
async def test_viewer_cannot_read_mcp_runtime_logs():
mcp_service = SimpleNamespace(
get_mcp_server_logs=AsyncMock(return_value=['private runtime line']),
)
client = await _create_test_client(mcp_service, role='viewer')
response = await client.get(
'/api/v1/mcp/servers/example/logs',
headers={
'Authorization': 'Bearer test-token',
'X-Workspace-Id': 'workspace-a',
},
)
assert response.status_code == 403
assert (await response.get_json())['code'] == 'permission_denied'
mcp_service.get_mcp_server_logs.assert_not_awaited()
async def test_mcp_server_route_accepts_encoded_slash_name():
mcp_service = SimpleNamespace(
get_mcp_server_by_name=AsyncMock(
@@ -46,11 +90,17 @@ async def test_mcp_server_route_accepts_encoded_slash_name():
response = await client.get(
'/api/v1/mcp/servers/pab1it0%2Fprometheus',
headers={'Authorization': 'Bearer test-token'},
headers={
'Authorization': 'Bearer test-token',
'X-Workspace-Id': 'workspace-a',
},
)
assert response.status_code == 200
mcp_service.get_mcp_server_by_name.assert_awaited_once_with('pab1it0/prometheus')
mcp_service.get_mcp_server_by_name.assert_awaited_once()
context, server_name = mcp_service.get_mcp_server_by_name.await_args.args
assert context.workspace_uuid == 'workspace-a'
assert server_name == 'pab1it0/prometheus'
payload = await response.get_json()
assert payload['data']['server']['name'] == 'pab1it0/prometheus'
@@ -66,11 +116,17 @@ async def test_mcp_resource_route_accepts_encoded_slash_name():
response = await client.get(
'/api/v1/mcp/servers/pab1it0%2Fprometheus/resources',
headers={'Authorization': 'Bearer test-token'},
headers={
'Authorization': 'Bearer test-token',
'X-Workspace-Id': 'workspace-a',
},
)
assert response.status_code == 200
mcp_service.get_mcp_server_by_name.assert_not_awaited()
mcp_service.get_mcp_server_resources.assert_awaited_once_with('pab1it0/prometheus')
mcp_service.get_mcp_server_resources.assert_awaited_once()
context, server_name = mcp_service.get_mcp_server_resources.await_args.args
assert context.workspace_uuid == 'workspace-a'
assert server_name == 'pab1it0/prometheus'
payload = await response.get_json()
assert payload['data']['resource_capabilities'] == {'subscribe': False}
@@ -0,0 +1,121 @@
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import sqlalchemy as sa
from sqlalchemy.ext.asyncio import create_async_engine
from langbot.pkg.api.http.service.apikey import ApiKeyIdentity
from langbot.pkg.api.mcp.context import get_request_context
from langbot.pkg.api.mcp.mount import MCPMount
from langbot.pkg.persistence.mgr import PersistenceManager, PersistenceMode
from langbot.pkg.persistence.tenant_uow import PersistenceScopeKind
@pytest.mark.asyncio
async def test_mcp_mount_keeps_request_context_but_no_session_during_stream_wait(tmp_path) -> None:
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "mcp-short-scope.db"}')
table = sa.Table('mcp_scope_probe', sa.MetaData(), sa.Column('id', sa.Integer, primary_key=True))
manager = PersistenceManager(object(), mode=PersistenceMode.CLOUD_RUNTIME)
manager.db = SimpleNamespace(get_engine=lambda: engine)
checked_out = 0
def on_checkout(*_args):
nonlocal checked_out
checked_out += 1
def on_checkin(*_args):
nonlocal checked_out
checked_out -= 1
sa.event.listen(engine.sync_engine, 'checkout', on_checkout)
sa.event.listen(engine.sync_engine, 'checkin', on_checkin)
try:
async with engine.begin() as conn:
await conn.run_sync(table.metadata.create_all)
identity = ApiKeyIdentity(
instance_uuid='instance-1',
workspace_uuid='workspace-1',
placement_generation=7,
api_key_uuid='key-1',
permissions=frozenset({'pipelines:read'}),
)
app = SimpleNamespace(
apikey_service=SimpleNamespace(authenticate_api_key=AsyncMock(return_value=identity)),
persistence_mgr=manager,
deployment_admission=None,
deployment=None,
)
stream_waiting = asyncio.Event()
release_stream = asyncio.Event()
observations: list[tuple[str, str, bool]] = []
async def fake_mcp_asgi(scope, receive, send):
del scope, receive
context = get_request_context()
assert manager.current_scope().kind is PersistenceScopeKind.WORKSPACE
assert manager.current_session() is None
await manager.execute_async(sa.select(table.c.id))
assert manager.current_session() is None
observations.append((context.request_id, context.workspace_uuid, manager.current_session() is None))
stream_waiting.set()
await release_stream.wait()
preserved_context = get_request_context()
observations.append(
(
preserved_context.request_id,
preserved_context.workspace_uuid,
manager.current_session() is None,
)
)
await manager.execute_async(sa.select(table.c.id))
assert manager.current_session() is None
await send({'type': 'http.response.start', 'status': 200, 'headers': []})
await send({'type': 'http.response.body', 'body': b'{}'})
async def unused_quart_asgi(scope, receive, send):
del scope, receive, send
raise AssertionError('MCP request was routed to Quart')
mount = MCPMount.__new__(MCPMount)
mount.ap = app
mount._mcp_asgi = fake_mcp_asgi
sent_messages: list[dict] = []
async def receive():
return {'type': 'http.request', 'body': b'', 'more_body': False}
async def send(message):
sent_messages.append(message)
async def release_after_observation() -> None:
await asyncio.wait_for(stream_waiting.wait(), timeout=2)
assert checked_out == 0
release_stream.set()
release_task = asyncio.create_task(release_after_observation())
await mount.wrap(unused_quart_asgi)(
{
'type': 'http',
'path': '/mcp',
'headers': [(b'x-api-key', b'secret')],
},
receive,
send,
)
await release_task
assert sent_messages[0]['status'] == 200
assert len(observations) == 2
assert observations[0] == observations[1]
assert observations[0][1:] == ('workspace-1', True)
assert checked_out == 0
assert manager.current_scope() is None
with pytest.raises(RuntimeError, match='context is unavailable'):
get_request_context()
finally:
await engine.dispose()
@@ -19,7 +19,7 @@ class FakeRegistry:
def __init__(self, runners):
self.runners = runners
async def list_runners(self, bound_plugins=None):
async def list_runners(self, context, bound_plugins=None):
return self.runners
@@ -54,7 +54,7 @@ async def test_default_pipeline_config_uses_first_installed_runner_schema():
agent_runner_registry=FakeRegistry([custom_agent, local_agent]),
)
config = await PipelineService(ap).get_default_pipeline_config()
config = await PipelineService(ap).get_default_pipeline_config('workspace-test')
assert config['ai']['runner']['id'] == 'plugin:alice/custom-agent/default'
assert config['ai']['runner_config'] == {
@@ -71,7 +71,7 @@ async def test_default_pipeline_config_stays_neutral_without_installed_runners()
agent_runner_registry=FakeRegistry([]),
)
config = await PipelineService(ap).get_default_pipeline_config()
config = await PipelineService(ap).get_default_pipeline_config('workspace-test')
assert config['ai']['runner']['id'] == ''
assert config['ai']['runner_config'] == {}
@@ -4,7 +4,7 @@ import sys
import types
from importlib import import_module
from types import SimpleNamespace
from unittest.mock import AsyncMock
from unittest.mock import AsyncMock, Mock
import pytest
import quart
@@ -19,11 +19,37 @@ pytestmark = pytest.mark.asyncio
async def _create_test_client(pipeline_service: SimpleNamespace, **extra_ap):
app = quart.Quart(__name__)
account = SimpleNamespace(uuid='account-test', user='test@example.com')
user_service = SimpleNamespace(
verify_jwt_token=AsyncMock(return_value='test@example.com'),
get_user_by_email=AsyncMock(return_value=SimpleNamespace(user='test@example.com')),
get_authenticated_account=AsyncMock(return_value=account),
)
access = SimpleNamespace(
workspace=SimpleNamespace(uuid='workspace-test'),
membership=SimpleNamespace(
uuid='membership-test',
role='owner',
projection_revision=1,
),
execution=SimpleNamespace(
instance_uuid='instance-test',
placement_generation=1,
),
)
plugin_connector = extra_ap.get('plugin_connector')
if plugin_connector is not None and not hasattr(plugin_connector, 'is_enable_plugin'):
plugin_connector.is_enable_plugin = False
extra_ap.setdefault('logger', SimpleNamespace(warning=Mock()))
ap = SimpleNamespace(
pipeline_service=pipeline_service,
user_service=user_service,
apikey_service=SimpleNamespace(
authenticate_api_key=AsyncMock(return_value=None),
),
workspace_collaboration_service=SimpleNamespace(
resolve_account_workspace=AsyncMock(return_value=access),
),
**extra_ap,
)
ap = SimpleNamespace(pipeline_service=pipeline_service, user_service=user_service, **extra_ap)
router_class = import_module('langbot.pkg.api.http.controller.groups.pipelines.pipelines').PipelinesRouterGroup
group = router_class(ap, app)
await group.initialize()
@@ -0,0 +1,139 @@
from __future__ import annotations
from contextlib import asynccontextmanager
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
import quart
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=4,
)
@pytest.fixture(scope='module')
def plugin_router_cls():
from tests.utils.import_isolation import MockLifecycleControlScope, isolated_sys_modules
class FakeMinimalApplication:
pass
mock_app = Mock(Application=FakeMinimalApplication)
mock_entities = Mock(LifecycleControlScope=MockLifecycleControlScope)
clear = [
'langbot.pkg.core.taskmgr',
'langbot.pkg.api.http.controller.group',
'langbot.pkg.api.http.controller.groups',
'langbot.pkg.api.http.controller.groups.plugins',
'langbot.pkg.api.http.controller.main',
]
with isolated_sys_modules(
mocks={
'langbot.pkg.core.app': mock_app,
'langbot.pkg.core.entities': mock_entities,
},
clear=clear,
):
from langbot.pkg.api.http.controller.groups.plugins import PluginsRouterGroup
yield PluginsRouterGroup
@pytest.mark.asyncio
async def test_public_plugin_asset_route_is_disabled_for_multi_workspace_policy(plugin_router_cls):
connector = SimpleNamespace(
get_plugin_icon=AsyncMock(),
require_workspace_context=AsyncMock(),
)
ap = SimpleNamespace(
plugin_connector=connector,
workspace_service=SimpleNamespace(
policy=SimpleNamespace(multi_workspace_enabled=True),
),
)
quart_app = quart.Quart(__name__)
router = plugin_router_cls(ap, quart_app)
await router.initialize()
response = await quart_app.test_client().get('/api/v1/plugins/author/plugin/icon')
assert response.status_code == 404
connector.require_workspace_context.assert_not_awaited()
connector.get_plugin_icon.assert_not_awaited()
@pytest.mark.asyncio
async def test_public_plugin_asset_uses_trusted_oss_singleton_binding(plugin_router_cls):
connector = SimpleNamespace(
require_workspace_context=AsyncMock(side_effect=lambda context: context),
)
binding = SimpleNamespace(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=4,
)
router = object.__new__(plugin_router_cls)
router.ap = SimpleNamespace(
plugin_connector=connector,
workspace_service=SimpleNamespace(
policy=SimpleNamespace(multi_workspace_enabled=False),
get_local_execution_binding=AsyncMock(return_value=binding),
),
)
result = await router._require_public_plugin_runtime_context()
assert result == CONTEXT
connector.require_workspace_context.assert_awaited_once_with(CONTEXT)
@pytest.mark.asyncio
async def test_background_plugin_operation_refences_captured_generation(plugin_router_cls):
operation = AsyncMock()
connector = SimpleNamespace(
require_workspace_context=AsyncMock(side_effect=WorkspaceNotFoundError('Plugin resource not found')),
)
router = object.__new__(plugin_router_cls)
router.ap = SimpleNamespace(plugin_connector=connector)
with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
await router._run_fenced_plugin_operation(CONTEXT, operation)
operation.assert_not_awaited()
@pytest.mark.asyncio
async def test_background_plugin_operation_revalidates_inside_short_tenant_uow(plugin_router_cls):
scopes = []
@asynccontextmanager
async def tenant_uow(workspace_uuid):
scopes.append(workspace_uuid)
yield
connector = SimpleNamespace(
require_workspace_context=AsyncMock(side_effect=lambda context: context),
)
operation = AsyncMock(return_value='done')
router = object.__new__(plugin_router_cls)
router.ap = SimpleNamespace(
plugin_connector=connector,
persistence_mgr=SimpleNamespace(
mode=SimpleNamespace(value='cloud_runtime'),
tenant_uow=tenant_uow,
),
)
result = await router._run_fenced_plugin_operation(CONTEXT, operation)
assert result == 'done'
assert scopes == [CONTEXT.workspace_uuid]
connector.require_workspace_context.assert_awaited_once_with(CONTEXT)
operation.assert_awaited_once()
@@ -0,0 +1,196 @@
from __future__ import annotations
import copy
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import quart
from langbot.pkg.api.http.controller.groups.knowledge.base import KnowledgeBaseRouterGroup
from langbot.pkg.api.http.controller.groups.pipelines.pipelines import PipelinesRouterGroup
from langbot.pkg.api.http.controller.groups.provider.models import LLMModelsRouterGroup
from langbot.pkg.api.http.controller.groups.provider.providers import ModelProvidersRouterGroup
from langbot.pkg.api.http.controller.groups.resources.mcp import MCPRouterGroup
from langbot.pkg.api.http.controller.groups.webhook_mgmt import WebhookManagementRouterGroup
from langbot.pkg.api.http.service.secrets import mask_secret_value, redact_secrets
pytestmark = pytest.mark.asyncio
WORKSPACE_UUID = '11111111-1111-4111-8111-111111111111'
RAW_PIPELINE = {
'uuid': 'pipeline-test',
'config': {'ai': {'n8n': {'webhook-url': 'https://hook.invalid/bearer-secret'}}},
}
RAW_MODEL = {
'uuid': 'model-test',
'provider_uuid': 'provider-test',
'extra_args': {'headers': {'Authorization': 'Bearer model-secret'}},
}
RAW_PROVIDER = {
'uuid': 'provider-test',
'base_url': 'https://provider-user:provider-password@provider.invalid/v1?token=url-secret&region=sg',
'api_keys': ['provider-secret'],
}
RAW_MCP_SERVER = {
'uuid': 'mcp-test',
'name': 'MCP Test',
'extra_args': {'url': 'https://mcp-user:mcp-password@mcp.invalid/connect?api_key=url-secret&transport=http'},
}
RAW_KNOWLEDGE_BASE = {
'uuid': 'kb-test',
'creation_settings': {'dify_apikey': 'knowledge-secret'},
}
RAW_WEBHOOK = {'id': 1, 'url': 'https://hook.invalid/path?token=webhook-secret'}
def _access(role: str):
return SimpleNamespace(
workspace=SimpleNamespace(uuid=WORKSPACE_UUID),
membership=SimpleNamespace(uuid='membership-test', role=role, projection_revision=1),
execution=SimpleNamespace(instance_uuid='instance-test', placement_generation=1),
)
async def _create_client(role: str):
application = SimpleNamespace()
account = SimpleNamespace(uuid='account-test', user='test@example.com')
application.user_service = SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account))
application.apikey_service = SimpleNamespace(authenticate_api_key=AsyncMock(return_value=None))
application.workspace_collaboration_service = SimpleNamespace(
resolve_account_workspace=AsyncMock(return_value=_access(role))
)
async def get_pipelines(_context, *_args, include_secret=False):
value = copy.deepcopy(RAW_PIPELINE)
return [value] if include_secret else [redact_secrets(value)]
async def get_pipeline(_context, _uuid, *, include_secret=False):
value = copy.deepcopy(RAW_PIPELINE)
return value if include_secret else redact_secrets(value)
application.pipeline_service = SimpleNamespace(
get_pipelines=AsyncMock(side_effect=get_pipelines),
get_pipeline=AsyncMock(side_effect=get_pipeline),
)
application.plugin_connector = SimpleNamespace(list_plugins=AsyncMock(return_value=[]))
application.mcp_service = SimpleNamespace(
get_mcp_servers=AsyncMock(return_value=[redact_secrets(copy.deepcopy(RAW_MCP_SERVER))])
)
application.skill_service = SimpleNamespace(list_skills=AsyncMock(return_value=[]))
async def get_models_by_provider(_context, _provider_uuid, *, include_secret=False):
value = copy.deepcopy(RAW_MODEL)
return [value] if include_secret else [redact_secrets(value)]
application.llm_model_service = SimpleNamespace(
get_llm_models_by_provider=AsyncMock(side_effect=get_models_by_provider)
)
async def get_providers(_context, *, include_secret=False):
value = copy.deepcopy(RAW_PROVIDER)
return [value] if include_secret else [redact_secrets(value)]
application.provider_service = SimpleNamespace(
get_providers=AsyncMock(side_effect=get_providers),
get_provider_model_counts=AsyncMock(return_value={'llm_count': 0, 'embedding_count': 0, 'rerank_count': 0}),
)
async def get_knowledge_bases(_context, *, include_secret=False):
value = copy.deepcopy(RAW_KNOWLEDGE_BASE)
return [value] if include_secret else [redact_secrets(value)]
application.knowledge_service = SimpleNamespace(get_knowledge_bases=AsyncMock(side_effect=get_knowledge_bases))
async def get_webhooks(_context, *, include_secret=False):
value = copy.deepcopy(RAW_WEBHOOK)
if not include_secret:
value['url'] = mask_secret_value(value['url'])
return [value]
application.webhook_service = SimpleNamespace(get_webhooks=AsyncMock(side_effect=get_webhooks))
quart_app = quart.Quart(__name__)
for router_type in (
PipelinesRouterGroup,
LLMModelsRouterGroup,
ModelProvidersRouterGroup,
MCPRouterGroup,
KnowledgeBaseRouterGroup,
WebhookManagementRouterGroup,
):
await router_type(application, quart_app).initialize()
return application, quart_app.test_client()
def _headers() -> dict[str, str]:
return {'Authorization': 'Bearer test-token', 'X-Workspace-Id': WORKSPACE_UUID}
@pytest.mark.parametrize('role', ['viewer', 'operator'])
async def test_viewer_and_operator_resource_reads_are_redacted(role: str):
application, client = await _create_client(role)
pipeline = (await (await client.get('/api/v1/pipelines', headers=_headers())).get_json())['data']['pipelines'][0]
model = (
await (
await client.get(
'/api/v1/provider/models/llm?provider_uuid=provider-test',
headers=_headers(),
)
).get_json()
)['data']['models'][0]
provider = (await (await client.get('/api/v1/provider/providers', headers=_headers())).get_json())['data'][
'providers'
][0]
mcp_server = (await (await client.get('/api/v1/mcp/servers', headers=_headers())).get_json())['data']['servers'][0]
knowledge_base = (await (await client.get('/api/v1/knowledge/bases', headers=_headers())).get_json())['data'][
'bases'
][0]
webhook = (await (await client.get('/api/v1/webhooks', headers=_headers())).get_json())['data']['webhooks'][0]
assert pipeline['config']['ai']['n8n']['webhook-url'] == '***'
assert model['extra_args']['headers']['Authorization'] == '***'
assert provider['api_keys'] == ['***']
assert provider['base_url'] == 'https://***@provider.invalid/v1?token=***&region=sg'
assert mcp_server['extra_args']['url'] == 'https://***@mcp.invalid/connect?api_key=***&transport=http'
assert knowledge_base['creation_settings']['dify_apikey'] == '***'
assert webhook['url'] == '***'
assert application.pipeline_service.get_pipelines.await_args.kwargs['include_secret'] is False
assert application.llm_model_service.get_llm_models_by_provider.await_args.kwargs['include_secret'] is False
assert application.provider_service.get_providers.await_args.kwargs['include_secret'] is False
assert application.knowledge_service.get_knowledge_bases.await_args.kwargs['include_secret'] is False
assert application.webhook_service.get_webhooks.await_args.kwargs['include_secret'] is False
async def test_resource_manager_receives_credentials_needed_for_management():
application, client = await _create_client('developer')
pipeline = (await (await client.get('/api/v1/pipelines', headers=_headers())).get_json())['data']['pipelines'][0]
model = (
await (
await client.get(
'/api/v1/provider/models/llm?provider_uuid=provider-test',
headers=_headers(),
)
).get_json()
)['data']['models'][0]
provider = (await (await client.get('/api/v1/provider/providers', headers=_headers())).get_json())['data'][
'providers'
][0]
knowledge_base = (await (await client.get('/api/v1/knowledge/bases', headers=_headers())).get_json())['data'][
'bases'
][0]
webhook = (await (await client.get('/api/v1/webhooks', headers=_headers())).get_json())['data']['webhooks'][0]
assert pipeline == RAW_PIPELINE
assert model == RAW_MODEL
assert provider['api_keys'] == ['provider-secret']
assert knowledge_base == RAW_KNOWLEDGE_BASE
assert webhook == RAW_WEBHOOK
assert application.pipeline_service.get_pipelines.await_args.kwargs['include_secret'] is True
assert application.llm_model_service.get_llm_models_by_provider.await_args.kwargs['include_secret'] is True
assert application.provider_service.get_providers.await_args.kwargs['include_secret'] is True
assert application.knowledge_service.get_knowledge_bases.await_args.kwargs['include_secret'] is True
assert application.webhook_service.get_webhooks.await_args.kwargs['include_secret'] is True
@@ -0,0 +1,110 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import quart
from langbot.pkg.api.http.controller.groups.stats import StatsRouterGroup
pytestmark = pytest.mark.asyncio
def session(
workspace_uuid: str,
*,
placement_generation: int = 1,
conversation_count: int = 0,
):
return SimpleNamespace(
instance_uuid='instance-test',
workspace_uuid=workspace_uuid,
placement_generation=placement_generation,
conversations=[object() for _ in range(conversation_count)],
)
async def create_client(*, role='viewer'):
quart_app = quart.Quart(__name__)
account = SimpleNamespace(uuid='account-test', user='test@example.com')
user_service = SimpleNamespace(
get_authenticated_account=AsyncMock(return_value=account),
)
access = SimpleNamespace(
execution=SimpleNamespace(
instance_uuid='instance-test',
placement_generation=1,
),
workspace=SimpleNamespace(uuid='workspace-a'),
membership=SimpleNamespace(
uuid='membership-test',
role=role,
projection_revision=1,
),
)
collaboration_service = SimpleNamespace(
resolve_account_workspace=AsyncMock(return_value=access),
)
def get_query_count(context):
assert context.instance_uuid == 'instance-test'
assert context.workspace_uuid == 'workspace-a'
assert context.placement_generation == 1
return 7
ap = SimpleNamespace(
user_service=user_service,
workspace_collaboration_service=collaboration_service,
sess_mgr=SimpleNamespace(
session_list=[
session('workspace-a', conversation_count=2),
session('workspace-b', conversation_count=5),
session(
'workspace-a',
placement_generation=2,
conversation_count=3,
),
SimpleNamespace(conversations=[object()] * 11),
]
),
query_pool=SimpleNamespace(get_query_count=get_query_count),
)
router = StatsRouterGroup(ap, quart_app)
await router.initialize()
return quart_app.test_client(), collaboration_service
async def test_basic_stats_are_scoped_to_selected_workspace_placement():
client, collaboration_service = await create_client()
response = await client.get(
'/api/v1/stats/basic',
headers={
'Authorization': 'Bearer test-token',
'X-Workspace-Id': 'workspace-a',
},
)
assert response.status_code == 200
payload = await response.get_json()
assert payload['data'] == {
'active_session_count': 1,
'conversation_count': 2,
'query_count': 7,
}
collaboration_service.resolve_account_workspace.assert_awaited_once_with('account-test', 'workspace-a')
async def test_basic_stats_requires_resource_view_permission():
client, _ = await create_client(role='unknown-role')
response = await client.get(
'/api/v1/stats/basic',
headers={'Authorization': 'Bearer test-token'},
)
assert response.status_code == 403
payload = await response.get_json()
assert payload['code'] == 'permission_denied'
+34 -6
View File
@@ -4,7 +4,7 @@ import sys
import types
from importlib import import_module
from types import SimpleNamespace
from unittest.mock import AsyncMock
from unittest.mock import ANY, AsyncMock
import pytest
import quart
@@ -19,14 +19,32 @@ pytestmark = pytest.mark.asyncio
async def _create_test_client(tool_mgr: SimpleNamespace, pipeline_service: SimpleNamespace):
app = quart.Quart(__name__)
account = SimpleNamespace(uuid='account-test', user='test@example.com')
user_service = SimpleNamespace(
verify_jwt_token=AsyncMock(return_value='test@example.com'),
get_user_by_email=AsyncMock(return_value=SimpleNamespace(user='test@example.com')),
get_authenticated_account=AsyncMock(return_value=account),
)
access = SimpleNamespace(
workspace=SimpleNamespace(uuid='workspace-test'),
membership=SimpleNamespace(
uuid='membership-test',
role='developer',
projection_revision=1,
),
execution=SimpleNamespace(
instance_uuid='instance-test',
placement_generation=1,
),
)
ap = SimpleNamespace(
tool_mgr=tool_mgr,
pipeline_service=pipeline_service,
user_service=user_service,
apikey_service=SimpleNamespace(
authenticate_api_key=AsyncMock(return_value=None)
),
workspace_collaboration_service=SimpleNamespace(
resolve_account_workspace=AsyncMock(return_value=access)
),
)
router_class = import_module('langbot.pkg.api.http.controller.groups.resources.tools').ToolsRouterGroup
group = router_class(ap, app)
@@ -50,6 +68,7 @@ async def test_global_tool_selector_uses_unambiguous_host_catalog():
payload = await response.get_json()
assert payload['data']['tools'] == [{'name': 'unique_tool', 'source': 'builtin'}]
tool_mgr.get_resolved_tool_catalog.assert_awaited_once_with(
ANY,
None,
None,
include_skill_authoring=True,
@@ -91,8 +110,9 @@ async def test_pipeline_tool_selector_resolves_only_bound_sources():
assert response.status_code == 200
payload = await response.get_json()
assert payload['data']['tools'][0]['source_id'] == 'bound-mcp'
pipeline_service.get_pipeline.assert_awaited_once_with('pipeline-1')
pipeline_service.get_pipeline.assert_awaited_once_with(ANY, 'pipeline-1')
tool_mgr.get_resolved_tool_catalog.assert_awaited_once_with(
ANY,
['allowed/plugin'],
['bound-mcp'],
include_skill_authoring=True,
@@ -123,6 +143,7 @@ async def test_pipeline_tool_selector_malformed_enable_all_flags_fail_closed(inv
assert response.status_code == 200
tool_mgr.get_resolved_tool_catalog.assert_awaited_once_with(
ANY,
['allowed/plugin'],
['bound-mcp'],
include_skill_authoring=True,
@@ -148,6 +169,7 @@ async def test_pipeline_tool_selector_malformed_extension_root_uses_empty_allowl
assert response.status_code == 200
tool_mgr.get_resolved_tool_catalog.assert_awaited_once_with(
ANY,
[],
[],
include_skill_authoring=True,
@@ -177,6 +199,7 @@ async def test_pipeline_tool_selector_malformed_binding_lists_use_empty_allowlis
assert response.status_code == 200
tool_mgr.get_resolved_tool_catalog.assert_awaited_once_with(
ANY,
[],
[],
include_skill_authoring=True,
@@ -230,8 +253,9 @@ async def test_tool_detail_uses_pipeline_scoped_catalog_and_path_tool_name(pipel
'source_name': 'allowed/plugin',
'source_id': 'allowed/plugin',
}
pipeline_service.get_pipeline.assert_awaited_once_with('pipeline-1')
pipeline_service.get_pipeline.assert_awaited_once_with(ANY, 'pipeline-1')
tool_mgr.get_resolved_tool_catalog.assert_awaited_once_with(
ANY,
['allowed/plugin'],
['bound-mcp'],
include_skill_authoring=True,
@@ -279,6 +303,7 @@ async def test_tool_detail_hides_ambiguous_or_missing_name():
assert response.status_code == 404
tool_mgr.get_resolved_tool_catalog.assert_awaited_once_with(
ANY,
None,
None,
include_skill_authoring=True,
@@ -297,5 +322,8 @@ async def test_tool_detail_returns_pipeline_not_found_before_catalog_lookup():
assert response.status_code == 404
assert await response.get_json() == {'code': -1, 'msg': 'pipeline not found'}
pipeline_service.get_pipeline.assert_awaited_once_with('missing-pipeline')
pipeline_service.get_pipeline.assert_awaited_once_with(
ANY,
'missing-pipeline',
)
tool_mgr.get_resolved_tool_catalog.assert_not_awaited()
@@ -0,0 +1,144 @@
from __future__ import annotations
import asyncio
from contextlib import asynccontextmanager
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.api.http.context import (
PrincipalContext,
PrincipalType,
RequestContext,
WorkspaceContext,
)
from langbot.pkg.api.http.controller.groups.pipelines.websocket_chat import WebSocketChatRouterGroup
from langbot.pkg.api.http.controller.groups.pipelines.websocket_chat import (
create_scoped_duplex_tasks,
)
from langbot.pkg.api.http.controller.groups.pipelines.websocket_chat import wait_for_duplex_tasks
from langbot.pkg.utils.bounded_executor import current_blocking_work_scope
@pytest.mark.asyncio
async def test_websocket_pipeline_lookup_opens_workspace_uow_after_auth_scope_closed() -> None:
workspace_uuid = 'workspace-a'
scopes: list[str] = []
in_scope = False
@asynccontextmanager
async def tenant_uow(selected_workspace_uuid: str):
nonlocal in_scope
assert not in_scope
in_scope = True
scopes.append(selected_workspace_uuid)
try:
yield
finally:
in_scope = False
async def get_pipeline(_context, _pipeline_uuid):
assert in_scope
return {'uuid': 'pipeline-a'}
adapter = Mock()
router = object.__new__(WebSocketChatRouterGroup)
router.ap = SimpleNamespace(
persistence_mgr=SimpleNamespace(
mode=SimpleNamespace(value='cloud_runtime'),
tenant_uow=tenant_uow,
),
pipeline_service=SimpleNamespace(get_pipeline=AsyncMock(side_effect=get_pipeline)),
platform_mgr=SimpleNamespace(get_websocket_proxy_bot=AsyncMock(return_value=SimpleNamespace(adapter=adapter))),
)
request_context = RequestContext(
instance_uuid='instance-a',
placement_generation=1,
request_id='request-a',
auth_type='user_token',
principal=PrincipalContext(
principal_type=PrincipalType.ACCOUNT,
account_uuid='account-a',
),
workspace=WorkspaceContext(
workspace_uuid=workspace_uuid,
membership_uuid='membership-a',
role='owner',
permissions=frozenset(),
),
)
result = await router._get_scoped_adapter(request_context, 'pipeline-a')
assert result is adapter
assert scopes == [workspace_uuid]
@pytest.mark.asyncio
async def test_duplex_websocket_tasks_cancel_blocked_peer_when_one_direction_ends() -> None:
blocked = asyncio.Event()
async def receive_forever() -> None:
blocked.set()
await asyncio.Future()
async def send_finishes() -> None:
await blocked.wait()
receive_task = asyncio.create_task(receive_forever())
send_task = asyncio.create_task(send_finishes())
await asyncio.wait_for(
wait_for_duplex_tasks(receive_task, send_task),
timeout=1,
)
assert receive_task.cancelled()
assert send_task.done()
@pytest.mark.asyncio
async def test_duplex_websocket_tasks_allow_terminal_send_to_drain() -> None:
receive_finished = asyncio.Event()
send_drained = asyncio.Event()
async def receive_finishes() -> None:
receive_finished.set()
async def send_terminal_frame() -> None:
await receive_finished.wait()
await asyncio.sleep(0)
send_drained.set()
receive_task = asyncio.create_task(receive_finishes())
send_task = asyncio.create_task(send_terminal_frame())
await wait_for_duplex_tasks(receive_task, send_task)
assert send_drained.is_set()
assert send_task.done()
assert not send_task.cancelled()
@pytest.mark.asyncio
async def test_duplex_websocket_tasks_share_trusted_workspace_budget() -> None:
observed: list[tuple[str, str | None]] = []
async def observe(direction: str) -> None:
await asyncio.sleep(0)
observed.append((direction, current_blocking_work_scope()))
receive_task, send_task = create_scoped_duplex_tasks(
observe('receive'),
observe('send'),
'workspace-a',
)
await asyncio.gather(receive_task, send_task)
assert sorted(observed) == [
('receive', 'workspace-a'),
('send', 'workspace-a'),
]
assert current_blocking_work_scope() is None
+230
View File
@@ -0,0 +1,230 @@
from __future__ import annotations
import datetime as dt
import hashlib
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
from langbot_plugin.box.models import SandboxAdmissionPolicy
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.box.admission import (
BoxAdmissionError,
SandboxAdmissionController,
require_cloud_admission_policy,
)
from langbot.pkg.box.service import BoxService
from langbot.pkg.cloud.entitlements import (
EntitlementResolver,
EntitlementSnapshot,
EntitlementUnavailableError,
)
_UTC = dt.timezone.utc
_CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=3,
entitlement_revision=7,
)
def _snapshot(
*,
revision: int = 7,
managed: bool = True,
sessions: int = 1,
expires_at: int = 2_000,
) -> EntitlementSnapshot:
return EntitlementSnapshot(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
entitlement_revision=revision,
status='active',
not_before=1,
expires_at=expires_at,
features={'managed_sandbox': managed},
limits={'managed_sandbox_sessions': sessions},
)
def _controller(snapshot: EntitlementSnapshot, *, now: float = 1_000.25):
provider = SimpleNamespace(get_workspace_entitlement=AsyncMock(return_value=snapshot))
resolver = EntitlementResolver('instance-a', provider)
client = SimpleNamespace(
upsert_sandbox_admission_grant=AsyncMock(
side_effect=lambda grant: {
'installed': True,
'workspace_uuid': grant.workspace_uuid,
'execution_generation': grant.execution_generation,
'entitlement_revision': grant.entitlement_revision,
'max_sessions': grant.max_sessions,
'max_managed_processes': grant.max_managed_processes,
}
),
revoke_sandbox_admission_grant=AsyncMock(
side_effect=lambda revocation: {
'revoked': True,
'workspace_uuid': revocation.workspace_uuid,
'entitlement_revision': revocation.entitlement_revision,
}
),
)
app = SimpleNamespace(entitlement_resolver=resolver, logger=Mock())
controller = SandboxAdmissionController(
app,
client,
policy=SandboxAdmissionPolicy(required=True, max_grant_ttl_sec=300),
wall_time=lambda: now,
)
return controller, client, provider
def test_cloud_admission_policy_requires_positive_workspace_quota():
with pytest.raises(BoxAdmissionError, match='workspace quota must be a positive integer'):
require_cloud_admission_policy(
{
'required': True,
'workspace_quota_mb': 0,
}
)
policy = require_cloud_admission_policy(
{
'required': True,
'workspace_quota_mb': 32,
}
)
assert policy.workspace_quota_mb == 32
@pytest.mark.asyncio
async def test_active_generic_entitlement_installs_short_lived_numeric_grant():
controller, client, provider = _controller(_snapshot())
grant = await controller.require(_CONTEXT)
assert grant.instance_uuid == _CONTEXT.instance_uuid
assert grant.workspace_uuid == _CONTEXT.workspace_uuid
assert grant.execution_generation == _CONTEXT.placement_generation
assert grant.entitlement_revision == 7
assert grant.max_sessions == 1
assert grant.max_managed_processes == 0
assert grant.expires_at == dt.datetime.fromtimestamp(1_300, tz=_UTC)
assert (grant.expires_at - dt.datetime.fromtimestamp(1_000.25, tz=_UTC)).total_seconds() < 300
provider.get_workspace_entitlement.assert_awaited_once_with('workspace-a')
client.upsert_sandbox_admission_grant.assert_awaited_once_with(grant)
client.revoke_sandbox_admission_grant.assert_not_awaited()
@pytest.mark.asyncio
@pytest.mark.parametrize(
'snapshot',
[
_snapshot(managed=False),
_snapshot(sessions=0),
_snapshot(sessions=2),
],
)
async def test_non_eligible_entitlement_revokes_and_fails_closed(snapshot):
controller, client, _provider = _controller(snapshot)
with pytest.raises(EntitlementUnavailableError):
await controller.require(_CONTEXT)
client.upsert_sandbox_admission_grant.assert_not_awaited()
revocation = client.revoke_sandbox_admission_grant.await_args.args[0]
assert revocation.entitlement_revision == snapshot.entitlement_revision
@pytest.mark.asyncio
async def test_transient_entitlement_failure_does_not_tombstone_valid_revision():
controller, client, provider = _controller(_snapshot())
await controller.require(_CONTEXT)
provider.get_workspace_entitlement.side_effect = RuntimeError('control plane unavailable')
with pytest.raises(RuntimeError, match='control plane unavailable'):
await controller.require(_CONTEXT)
client.revoke_sandbox_admission_grant.assert_not_awaited()
provider.get_workspace_entitlement.side_effect = None
provider.get_workspace_entitlement.return_value = _snapshot()
recovered = await controller.require(_CONTEXT)
assert recovered.entitlement_revision == 7
@pytest.mark.asyncio
async def test_runtime_receipt_mismatch_is_revoked_and_never_admitted():
controller, client, _provider = _controller(_snapshot())
client.upsert_sandbox_admission_grant.return_value = {'installed': True, 'workspace_uuid': 'other'}
client.upsert_sandbox_admission_grant.side_effect = None
with pytest.raises(Exception, match='invalid sandbox admission receipt'):
await controller.require(_CONTEXT)
client.revoke_sandbox_admission_grant.assert_not_awaited()
@pytest.mark.asyncio
async def test_authoritative_cancelled_revision_is_revoked():
cancelled = _snapshot(revision=8).model_copy(update={'status': 'cancelled'})
controller, client, _provider = _controller(cancelled)
with pytest.raises(EntitlementUnavailableError, match='not active'):
await controller.require(_CONTEXT)
revocation = client.revoke_sandbox_admission_grant.await_args.args[0]
assert revocation.entitlement_revision == 8
@pytest.mark.asyncio
async def test_cloud_box_readiness_failure_aborts_service_initialization(tmp_path):
workspace_root = tmp_path / 'box' / 'workspaces'
workspace_root.mkdir(parents=True)
box_config = {
'enabled': True,
'backend': 'nsjail',
'runtime': {'endpoint': 'ws://box:5410'},
'local': {
'host_root': str(tmp_path / 'box'),
'default_workspace': str(workspace_root),
'allowed_mount_roots': [str(tmp_path / 'box')],
},
'admission': {
'required': True,
'logical_session_id': 'global',
'required_backend': 'nsjail',
'max_sessions': 1,
'max_managed_processes': 0,
'max_grant_ttl_sec': 300,
'workspace_quota_mb': 32,
},
}
client = SimpleNamespace(
initialize=AsyncMock(),
verify_shared_workspace=AsyncMock(
side_effect=lambda marker_name: {
'marker_name': marker_name,
'size': (workspace_root / marker_name).stat().st_size,
'sha256': hashlib.sha256((workspace_root / marker_name).read_bytes()).hexdigest(),
}
),
get_backend_info=AsyncMock(return_value={'name': 'docker', 'available': True}),
)
app = SimpleNamespace(
logger=Mock(),
deployment=SimpleNamespace(multi_workspace_enabled=True),
entitlement_resolver=Mock(),
workspace_service=SimpleNamespace(instance_uuid='instance-a'),
instance_config=SimpleNamespace(data={'box': box_config}),
)
service = BoxService(app, client=client)
with pytest.raises(Exception, match='nsjail isolation readiness failed'):
await service.initialize()
assert service.available is False
+263 -29
View File
@@ -6,13 +6,28 @@ from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.box import connector as connector_module
from langbot_plugin.box.client import ActionRPCBoxClient
from langbot_plugin.box.errors import BoxRuntimeUnavailableError
from langbot_plugin.box.security import (
BOX_CONTROL_TOKEN_ENV,
BOX_CONTROL_TOKEN_HEADER,
BOX_INSTANCE_HEADER,
BOX_PLACEMENT_GENERATION_HEADER,
BOX_TRUSTED_INSTANCE_ENV,
BOX_WORKSPACE_HEADER,
)
from langbot_plugin.entities.io.context import ActionContext
from langbot.pkg.box.connector import BoxRuntimeConnector
_CONTROL_TOKEN = 'box-control-token-that-is-longer-than-32-bytes'
def make_app(logger: Mock, runtime_endpoint: str = ''):
return SimpleNamespace(
logger=logger,
workspace_service=SimpleNamespace(instance_uuid='instance-a'),
instance_config=SimpleNamespace(
data={
'box': {
@@ -108,41 +123,260 @@ def test_box_runtime_connector_dispose_terminates_subprocess(monkeypatch: pytest
@pytest.mark.asyncio
async def test_box_runtime_connector_heartbeat_failure_requests_reconnect(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setattr('langbot.pkg.box.connector._HEARTBEAT_INTERVAL_SEC', 0)
disconnected = asyncio.Event()
callbacks = []
async def test_box_runtime_connector_cleans_partial_transport_on_connect_failure(
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.setattr('langbot.pkg.utils.platform.get_platform', lambda: 'linux')
monkeypatch.setattr('langbot.pkg.utils.platform.standalone_box', False)
connector = BoxRuntimeConnector(make_app(Mock()))
connector._start_local_stdio = AsyncMock(side_effect=RuntimeError('bind failed'))
connector._stop_transport = AsyncMock()
connector._close_managed_subprocess = AsyncMock()
async def on_disconnect(connector):
callbacks.append(connector)
disconnected.set()
with pytest.raises(RuntimeError, match='bind failed'):
await connector.initialize()
connector = BoxRuntimeConnector(make_app(Mock()), runtime_disconnect_callback=on_disconnect)
connector._handler = SimpleNamespace(call_action=AsyncMock(side_effect=TimeoutError('silent runtime')))
heartbeat = asyncio.create_task(connector._heartbeat_loop())
await asyncio.wait_for(disconnected.wait(), timeout=1)
await heartbeat
assert callbacks == [connector]
assert connector._stop_transport.await_count == 2
connector._close_managed_subprocess.assert_awaited_once()
@pytest.mark.asyncio
async def test_box_runtime_stdio_disconnect_requests_reconnect(monkeypatch: pytest.MonkeyPatch):
async def test_box_runtime_connector_starts_heartbeat_after_reconnect(
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.setattr('langbot.pkg.utils.platform.get_platform', lambda: 'linux')
monkeypatch.setattr('langbot.pkg.utils.platform.standalone_box', False)
logger = Mock()
on_disconnect = AsyncMock()
handler = SimpleNamespace(
call_action=AsyncMock(return_value={}),
run=AsyncMock(return_value=None),
connector = BoxRuntimeConnector(make_app(Mock()))
connector._start_local_stdio = AsyncMock(side_effect=[RuntimeError('bind failed'), None])
connector._stop_transport = AsyncMock()
connector._close_managed_subprocess = AsyncMock()
with pytest.raises(RuntimeError, match='bind failed'):
await connector.initialize()
assert connector._heartbeat_task is None
await connector.reconnect()
assert connector._heartbeat_task is not None
assert not connector._heartbeat_task.done()
await connector.aclose()
@pytest.mark.asyncio
async def test_box_stdio_connection_does_not_capture_unconsumed_stderr(
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.setattr('langbot.pkg.utils.platform.get_platform', lambda: 'linux')
monkeypatch.setattr('langbot.pkg.utils.platform.standalone_box', False)
created = {}
class FakeHandler:
def __init__(self, connection):
self.release = asyncio.Event()
async def call_action(self, action, data):
return None
async def run(self):
await self.release.wait()
async def close(self):
self.release.set()
class FakeController:
def __init__(self, **kwargs):
created.update(kwargs)
self.process = SimpleNamespace(returncode=0)
async def run(self, callback):
await callback(object())
async def close(self):
return None
monkeypatch.setattr(connector_module, 'Handler', FakeHandler)
monkeypatch.setattr(
'langbot_plugin.runtime.io.controllers.stdio.client.StdioClientController',
FakeController,
)
monkeypatch.setattr('langbot.pkg.box.connector.Handler', Mock(return_value=handler))
connector = BoxRuntimeConnector(make_app(logger), runtime_disconnect_callback=on_disconnect)
connected = asyncio.Event()
connector = BoxRuntimeConnector(make_app(Mock()))
callback = connector._make_connection_callback('stdio', connected, [])
await callback(Mock())
await connector.initialize()
assert connected.is_set()
on_disconnect.assert_awaited_once_with(connector)
logger.error.assert_called_once_with('Disconnected from Box runtime, trying to reconnect...')
assert created['capture_stderr'] is False
assert connector._handler is not None
await connector.aclose()
@pytest.mark.asyncio
async def test_box_disconnect_notifies_once_and_clears_handler(
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.setattr('langbot.pkg.utils.platform.get_platform', lambda: 'linux')
monkeypatch.setattr('langbot.pkg.utils.platform.standalone_box', False)
disconnect = AsyncMock()
class FakeHandler:
def __init__(self, connection):
pass
async def call_action(self, action, data):
return None
async def run(self):
return None
async def close(self):
return None
class FakeController:
def __init__(self, **kwargs):
self.process = SimpleNamespace(returncode=0)
async def run(self, callback):
await callback(object())
async def close(self):
return None
monkeypatch.setattr(connector_module, 'Handler', FakeHandler)
monkeypatch.setattr(
'langbot_plugin.runtime.io.controllers.stdio.client.StdioClientController',
FakeController,
)
connector = BoxRuntimeConnector(make_app(Mock()), runtime_disconnect_callback=disconnect)
await connector.initialize()
await asyncio.sleep(0)
disconnect.assert_awaited_once_with(connector)
assert connector._handler is None
await connector.aclose()
def test_box_runtime_connector_builds_host_control_headers(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv(BOX_CONTROL_TOKEN_ENV, _CONTROL_TOKEN)
connector = BoxRuntimeConnector(make_app(Mock(), runtime_endpoint='http://box-runtime:5410'))
headers = connector.get_control_headers()
assert headers == {
BOX_CONTROL_TOKEN_HEADER: _CONTROL_TOKEN,
BOX_INSTANCE_HEADER: 'instance-a',
}
assert _CONTROL_TOKEN not in connector._resolve_rpc_ws_url()
def test_box_runtime_connector_builds_placement_scoped_relay_headers(
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.setenv(BOX_CONTROL_TOKEN_ENV, _CONTROL_TOKEN)
connector = BoxRuntimeConnector(make_app(Mock(), runtime_endpoint='http://box-runtime:5410'))
headers = connector.get_relay_headers(
ActionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=7,
)
)
assert headers == {
BOX_CONTROL_TOKEN_HEADER: _CONTROL_TOKEN,
BOX_INSTANCE_HEADER: 'instance-a',
BOX_WORKSPACE_HEADER: 'workspace-a',
BOX_PLACEMENT_GENERATION_HEADER: '7',
}
def test_box_runtime_connector_rejects_relay_context_from_other_instance(
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.setenv(BOX_CONTROL_TOKEN_ENV, _CONTROL_TOKEN)
connector = BoxRuntimeConnector(make_app(Mock(), runtime_endpoint='http://box-runtime:5410'))
with pytest.raises(BoxRuntimeUnavailableError, match='another LangBot instance'):
connector.get_relay_headers(
ActionContext(
instance_uuid='instance-b',
workspace_uuid='workspace-a',
placement_generation=1,
)
)
def test_external_box_runtime_fails_closed_without_control_token(monkeypatch: pytest.MonkeyPatch):
monkeypatch.delenv(BOX_CONTROL_TOKEN_ENV, raising=False)
connector = BoxRuntimeConnector(make_app(Mock(), runtime_endpoint='http://box-runtime:5410'))
with pytest.raises(BoxRuntimeUnavailableError, match=BOX_CONTROL_TOKEN_ENV):
connector.get_control_headers()
async def test_local_stdio_injects_generated_token_and_trusted_instance(
monkeypatch: pytest.MonkeyPatch,
):
monkeypatch.delenv(BOX_CONTROL_TOKEN_ENV, raising=False)
captured = {}
class FakeStdioClientController:
def __init__(self, **kwargs):
captured.update(kwargs)
self.process = Mock()
async def run(self, callback):
await callback(None)
monkeypatch.setattr(
'langbot_plugin.runtime.io.controllers.stdio.client.StdioClientController',
FakeStdioClientController,
)
connector = BoxRuntimeConnector(make_app(Mock()))
def fake_callback(_transport_name, connected, _connect_error, _generation):
async def callback(_connection):
connected.set()
return callback
monkeypatch.setattr(connector, '_make_connection_callback', fake_callback)
await connector._start_local_stdio()
assert len(captured['env'][BOX_CONTROL_TOKEN_ENV]) >= 32
assert captured['env'][BOX_TRUSTED_INSTANCE_ENV] == 'instance-a'
async def test_websocket_controller_receives_control_headers(monkeypatch: pytest.MonkeyPatch):
monkeypatch.setenv(BOX_CONTROL_TOKEN_ENV, _CONTROL_TOKEN)
captured = {}
class FakeWebSocketClientController:
def __init__(self, **kwargs):
captured.update(kwargs)
async def run(self, callback):
await callback(None)
monkeypatch.setattr(
'langbot_plugin.runtime.io.controllers.ws.client.WebSocketClientController',
FakeWebSocketClientController,
)
connector = BoxRuntimeConnector(make_app(Mock(), runtime_endpoint='http://box-runtime:5410'))
def fake_callback(_transport_name, connected, _connect_error, _generation):
async def callback(_connection):
connected.set()
return callback
monkeypatch.setattr(connector, '_make_connection_callback', fake_callback)
await connector._connect_ws('ws://box-runtime:5410/rpc/ws', 'WebSocket')
assert captured['additional_headers'] == {
BOX_CONTROL_TOKEN_HEADER: _CONTROL_TOKEN,
BOX_INSTANCE_HEADER: 'instance-a',
}
assert _CONTROL_TOKEN not in captured['ws_url']
@@ -0,0 +1,65 @@
from pathlib import Path
_REPO_ROOT = Path(__file__).resolve().parents[3]
def test_compose_injects_the_same_box_control_token_into_host_and_runtime():
compose = (_REPO_ROOT / 'docker' / 'docker-compose.yaml').read_text(encoding='utf-8')
box_service = compose.split(' langbot_box:', 1)[1].split(' langbot:', 1)[0]
langbot_service = compose.split(' langbot:', 1)[1]
token_env = 'LANGBOT_BOX_CONTROL_TOKEN=${LANGBOT_BOX_CONTROL_TOKEN:-}'
assert token_env in box_service
assert token_env in langbot_service
def test_kubernetes_uses_one_secret_for_box_runtime_and_langbot():
manifest = (_REPO_ROOT / 'docker' / 'kubernetes.yaml').read_text(encoding='utf-8')
box_deployment = manifest.split('name: langbot-box', 1)[1].split('# Service for LangBot Box runtime', 1)[0]
langbot_deployment = manifest.split('# Deployment for LangBot\n', 1)[1]
secret_reference = '\n'.join(
[
'- name: LANGBOT_BOX_CONTROL_TOKEN',
' valueFrom:',
' secretKeyRef:',
' name: langbot-box-control',
' key: token',
]
)
assert secret_reference in box_deployment
assert secret_reference in langbot_deployment
assert '--from-literal=token="$(openssl rand -hex 32)"' in manifest
def test_compose_injects_same_plugin_runtime_control_token_into_both_services():
compose = (_REPO_ROOT / 'docker' / 'docker-compose.yaml').read_text(encoding='utf-8')
runtime_service = compose.split(' langbot_plugin_runtime:', 1)[1].split(' langbot_box:', 1)[0]
langbot_service = compose.split(' langbot:', 1)[1]
token_env = 'LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN=${LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN:-}'
assert token_env in runtime_service
assert token_env in langbot_service
def test_kubernetes_uses_one_secret_for_plugin_runtime_and_langbot():
manifest = (_REPO_ROOT / 'docker' / 'kubernetes.yaml').read_text(encoding='utf-8')
runtime_deployment = manifest.split('# Deployment for LangBot Plugin Runtime', 1)[1].split(
'# Service for LangBot Plugin Runtime',
1,
)[0]
langbot_deployment = manifest.split('# Deployment for LangBot\n', 1)[1]
secret_reference = '\n'.join(
[
'- name: LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN',
' valueFrom:',
' secretKeyRef:',
' name: langbot-plugin-runtime-control',
' key: token',
]
)
assert secret_reference in runtime_deployment
assert secret_reference in langbot_deployment
assert 'create secret generic langbot-plugin-runtime-control' in manifest
File diff suppressed because it is too large Load Diff
+47 -2
View File
@@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.box.workspace import (
BoxWorkspaceSession,
classify_python_workspace,
@@ -16,6 +17,13 @@ from langbot.pkg.box.workspace import (
)
_CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
def test_rewrite_mounted_path_translates_host_prefix():
result = rewrite_mounted_path('/tmp/demo/project/app.py', '/tmp/demo/project')
assert result == '/workspace/app.py'
@@ -57,6 +65,9 @@ def test_wrap_python_command_with_env_contains_bootstrap_and_command():
assert '_LB_SYSTEM_PYTHON="$(command -v python3 || command -v python || true)"' in command
assert '"$_LB_SYSTEM_PYTHON" -m venv "$_LB_VENV_DIR"' in command
assert 'kill -0 "$_LB_LOCK_OWNER"' in command
assert 'max_manifest_bytes = 10 * 1024 * 1024' in command
assert 'handle.read(1024 * 1024)' in command
assert 'digest.update(handle.read())' not in command
assert 'export VIRTUAL_ENV="$_LB_VENV_DIR"' in command
assert command.rstrip().endswith('python script.py')
@@ -66,6 +77,7 @@ async def test_workspace_session_execute_for_query_uses_session_payload():
box_service = SimpleNamespace(execute_spec_payload=AsyncMock(return_value={'ok': True}))
workspace = BoxWorkspaceSession(
box_service,
_CONTEXT,
'skill-person_123-demo',
host_path='/tmp/project',
host_path_mode='rw',
@@ -94,6 +106,7 @@ async def test_workspace_session_start_managed_process_rewrites_command_and_args
box_service = SimpleNamespace(start_managed_process=AsyncMock(return_value={'status': 'running'}))
workspace = BoxWorkspaceSession(
box_service,
_CONTEXT,
'mcp-u1',
host_path='/tmp/project',
host_path_mode='ro',
@@ -106,8 +119,10 @@ async def test_workspace_session_start_managed_process_rewrites_command_and_args
)
assert result == {'status': 'running'}
session_id = box_service.start_managed_process.await_args.args[0]
payload = box_service.start_managed_process.await_args.args[1]
execution_context = box_service.start_managed_process.await_args.args[0]
session_id = box_service.start_managed_process.await_args.args[1]
payload = box_service.start_managed_process.await_args.args[2]
assert execution_context == _CONTEXT
assert session_id == 'mcp-u1'
assert payload == {
'command': 'python',
@@ -118,9 +133,39 @@ async def test_workspace_session_start_managed_process_rewrites_command_and_args
}
@pytest.mark.asyncio
async def test_workspace_session_relay_connection_keeps_execution_context():
box_service = SimpleNamespace(
get_managed_process_websocket_connection=AsyncMock(
return_value=(
'ws://box/relay',
{'X-LangBot-Placement-Generation': '1'},
)
)
)
workspace = BoxWorkspaceSession(
box_service,
_CONTEXT,
'mcp-shared',
)
connection = await workspace.get_managed_process_websocket_connection('server-a')
assert connection == (
'ws://box/relay',
{'X-LangBot-Placement-Generation': '1'},
)
box_service.get_managed_process_websocket_connection.assert_awaited_once_with(
_CONTEXT,
'mcp-shared',
'server-a',
)
def test_workspace_session_build_session_payload_keeps_generic_workspace_shape():
workspace = BoxWorkspaceSession(
Mock(),
_CONTEXT,
'workspace-1',
host_path='/tmp/project',
host_path_mode='rw',
+380
View File
@@ -0,0 +1,380 @@
from __future__ import annotations
import dataclasses
from types import SimpleNamespace
import pytest
from langbot.pkg.cloud.bootstrap import (
CloudBootstrapError,
CloudManifestRefreshService,
CloudRuntimeUnavailableError,
DeploymentAdmissionGuard,
OpenSourceDeployment,
VerifiedCloudDeployment,
resolve_deployment,
)
from langbot.pkg.cloud.entitlements import EntitlementSnapshot
pytestmark = pytest.mark.asyncio
class _Entitlements:
async def get_workspace_entitlement(self, workspace_uuid: str) -> EntitlementSnapshot:
return EntitlementSnapshot(
instance_uuid='instance-a',
workspace_uuid=workspace_uuid,
entitlement_revision=1,
status='active',
not_before=1,
expires_at=4_000_000_000,
features={'managed_sandbox': True},
limits={'managed_sandbox_sessions': 1},
)
class _Directory:
async def fetch_snapshot(self, instance_uuid: str):
del instance_uuid
raise AssertionError('not used by bootstrap contract tests')
async def fetch_events(self, instance_uuid: str, after_cursor: int, limit: int):
del instance_uuid, after_cursor, limit
raise AssertionError('not used by bootstrap contract tests')
async def fetch_workspaces(self, instance_uuid: str, workspace_uuids: tuple[str, ...]):
del instance_uuid, workspace_uuids
raise AssertionError('not used by bootstrap contract tests')
class _Manifest:
def __init__(self):
self.candidate = None
self.closed = False
async def refresh_manifest(self):
if self.candidate is None:
raise AssertionError('no refreshed Manifest was configured')
return self.candidate
async def aclose(self) -> None:
self.closed = True
class _Provider:
def __init__(self):
self.manifest_provider = _Manifest()
def bootstrap(self, *, instance_uuid: str, instance_config: dict):
del instance_config
return VerifiedCloudDeployment(
instance_uuid=instance_uuid,
manifest_jti='manifest-a',
manifest_generation=3,
expires_at=4_000_000_000,
release='cloud-v2',
capabilities=frozenset({'multi_workspace_v2'}),
tenant_isolation_version=2,
entitlement_provider=_Entitlements(),
directory_provider=_Directory(),
manifest_provider=self.manifest_provider,
verification_key_id='root-2026',
)
class _EntryPoint:
def __init__(self, value):
self.value = value
def load(self):
return self.value
class _EntryPoints(list):
def select(self, *, group: str):
return self if group == 'langbot.cloud_bootstrap' else []
def _cloud_config() -> dict:
return {
'database': {'use': 'postgresql'},
'vdb': {
'use': 'pgvector',
'pgvector': {
'use_business_database': True,
'allowed_dimensions': [384, 768, 1536],
},
},
'mcp': {'stdio': {'enabled': False}},
'plugin': {'worker': {'require_hard_limits': True}},
'box': {
'enabled': True,
'backend': 'nsjail',
'runtime': {'endpoint': 'ws://langbot-box:5410'},
'admission': {
'required': True,
'logical_session_id': 'global',
'required_backend': 'nsjail',
'max_sessions': 1,
'max_managed_processes': 0,
'max_grant_ttl_sec': 300,
'workspace_quota_mb': 32,
},
'local': {
'host_root': '/var/lib/langbot/box',
'default_workspace': '/var/lib/langbot/box/workspaces',
'allowed_mount_roots': ['/var/lib/langbot/box'],
},
},
# Proves mutable product metadata does not participate in selection.
'system': {'edition': 'community'},
}
async def test_no_closed_entry_point_selects_oss_singleton_even_if_edition_says_cloud():
deployment = await resolve_deployment(
instance_uuid='instance-a',
instance_config={'system': {'edition': 'cloud'}},
entry_points=lambda: _EntryPoints(),
)
assert isinstance(deployment, OpenSourceDeployment)
assert deployment.multi_workspace_enabled is False
async def test_verified_closed_entry_point_activates_cloud_policy():
deployment = await resolve_deployment(
instance_uuid='instance-a',
instance_config=_cloud_config(),
entry_points=lambda: _EntryPoints([_EntryPoint(_Provider)]),
now=1_000,
)
assert isinstance(deployment, VerifiedCloudDeployment)
assert deployment.multi_workspace_enabled is True
assert deployment.persistence_mode == 'cloud_runtime'
@pytest.mark.parametrize(
('field', 'value', 'message'),
[
('database', {'use': 'sqlite'}, 'database.use=postgresql'),
('vdb', {'use': 'chroma'}, 'vdb.use=pgvector'),
('mcp', {'stdio': {'enabled': True}}, 'mcp.stdio.enabled=false'),
('plugin', {'worker': {'require_hard_limits': False}}, 'plugin.worker.require_hard_limits=true'),
],
)
async def test_cloud_runtime_config_is_fail_closed(field, value, message):
config = _cloud_config()
config[field] = value
with pytest.raises(CloudBootstrapError, match=message):
await resolve_deployment(
instance_uuid='instance-a',
instance_config=config,
entry_points=lambda: _EntryPoints([_EntryPoint(_Provider())]),
now=1_000,
)
@pytest.mark.parametrize(
('directory_config', 'message'),
[
({'max_active_workspaces': 0}, 'greater than or equal to 1'),
({'max_active_workspaces': True}, 'must be an integer'),
(
{
'max_active_workspaces': 10,
'max_snapshot_workspaces': 9,
},
'max_snapshot_workspaces',
),
({'max_response_bytes': 64 * 1024 * 1024 + 1}, 'less than or equal to'),
],
)
async def test_cloud_directory_capacity_contract_is_fail_closed(directory_config, message):
config = _cloud_config()
config['cloud'] = {'directory': directory_config}
with pytest.raises(CloudBootstrapError, match=message):
await resolve_deployment(
instance_uuid='instance-a',
instance_config=config,
entry_points=lambda: _EntryPoints([_EntryPoint(_Provider())]),
now=1_000,
)
@pytest.mark.parametrize(
('pgvector_config', 'message'),
[
({'use_business_database': False, 'allowed_dimensions': [1536]}, 'use_business_database=true'),
({'use_business_database': True, 'allowed_dimensions': []}, 'allowed_dimensions'),
({'use_business_database': True, 'allowed_dimensions': [3072]}, 'allowed_dimensions'),
({'use_business_database': True, 'allowed_dimensions': [True]}, 'allowed_dimensions'),
],
)
async def test_cloud_pgvector_contract_is_fail_closed(pgvector_config, message):
config = _cloud_config()
config['vdb']['pgvector'] = pgvector_config
with pytest.raises(CloudBootstrapError, match=message):
await resolve_deployment(
instance_uuid='instance-a',
instance_config=config,
entry_points=lambda: _EntryPoints([_EntryPoint(_Provider())]),
now=1_000,
)
@pytest.mark.parametrize(
('mutate', 'message'),
[
(lambda config: config['box'].update(enabled=False), 'box.enabled=true'),
(lambda config: config['box'].update(backend='docker'), 'box.backend=nsjail'),
(lambda config: config['box']['runtime'].update(endpoint=''), 'box.runtime.endpoint'),
(
lambda config: config['box']['admission'].update(max_sessions=2),
'grant-enforced Box admission',
),
(
lambda config: config['box']['admission'].update(max_managed_processes=1),
'zero managed processes',
),
(
lambda config: config['box']['admission'].update(max_grant_ttl_sec=301),
'max_grant_ttl_sec',
),
(
lambda config: config['box']['admission'].update(workspace_quota_mb=0),
'workspace_quota_mb must be a positive integer',
),
(
lambda config: config['box']['admission'].update(workspace_quota_mb=True),
'workspace_quota_mb must be a positive integer',
),
(
lambda config: config['box']['local'].update(default_workspace='relative/workspaces'),
'default_workspace must be an absolute',
),
(
lambda config: config['box']['local'].update(
default_workspace='/other/workspaces',
),
'under allowed_mount_roots',
),
],
)
async def test_cloud_box_contract_is_fail_closed(mutate, message):
config = _cloud_config()
mutate(config)
with pytest.raises(CloudBootstrapError, match=message):
await resolve_deployment(
instance_uuid='instance-a',
instance_config=config,
entry_points=lambda: _EntryPoints([_EntryPoint(_Provider())]),
now=1_000,
)
async def test_invalid_provider_never_falls_back_to_oss():
provider = SimpleNamespace(bootstrap=lambda **_: object())
with pytest.raises(CloudBootstrapError, match='must return VerifiedCloudDeployment'):
await resolve_deployment(
instance_uuid='instance-a',
instance_config=_cloud_config(),
entry_points=lambda: _EntryPoints([_EntryPoint(provider)]),
now=1_000,
)
async def test_duplicate_closed_providers_fail_closed():
with pytest.raises(CloudBootstrapError, match='Exactly one'):
await resolve_deployment(
instance_uuid='instance-a',
instance_config=_cloud_config(),
entry_points=lambda: _EntryPoints([_EntryPoint(_Provider()), _EntryPoint(_Provider())]),
now=1_000,
)
async def test_deployment_admission_expires_even_after_wall_clock_rollback():
wall = [1_000.0]
monotonic = [50.0]
deployment = dataclasses.replace(
_Provider().bootstrap(instance_uuid='instance-a', instance_config={}),
expires_at=1_010,
)
guard = DeploymentAdmissionGuard(
'instance-a',
deployment,
wall_time=lambda: wall[0],
monotonic_time=lambda: monotonic[0],
)
assert guard.require_active() is deployment
wall[0] = 900.0
monotonic[0] = 60.0
with pytest.raises(CloudRuntimeUnavailableError, match='expired'):
guard.require_active()
async def test_deployment_admission_accepts_only_monotonic_non_conflicting_renewal():
wall = [1_000.0]
monotonic = [50.0]
current = dataclasses.replace(
_Provider().bootstrap(instance_uuid='instance-a', instance_config={}),
expires_at=1_010,
)
guard = DeploymentAdmissionGuard(
'instance-a',
current,
wall_time=lambda: wall[0],
monotonic_time=lambda: monotonic[0],
)
renewed = dataclasses.replace(
current,
manifest_jti='manifest-b',
manifest_generation=4,
expires_at=2_000,
)
guard.replace(renewed)
assert guard.require_active() is renewed
rollback = dataclasses.replace(current, manifest_generation=2)
with pytest.raises(CloudRuntimeUnavailableError, match='rolled back'):
guard.replace(rollback)
conflicting = dataclasses.replace(renewed, manifest_jti='different')
with pytest.raises(CloudRuntimeUnavailableError, match='conflicting'):
guard.replace(conflicting)
async def test_manifest_refresh_replaces_receipt_before_short_ttl_expires():
wall = [1_000.0]
provider = _Provider()
current = dataclasses.replace(
provider.bootstrap(instance_uuid='instance-a', instance_config={}),
expires_at=1_300,
)
guard = DeploymentAdmissionGuard('instance-a', current, wall_time=lambda: wall[0])
renewed = dataclasses.replace(
current,
manifest_jti='manifest-renewed',
manifest_generation=current.manifest_generation + 1,
expires_at=2_000,
)
provider.manifest_provider.candidate = renewed
service = CloudManifestRefreshService(
guard,
provider.manifest_provider,
SimpleNamespace(exception=lambda *_: None),
wall_time=lambda: wall[0],
)
assert service.next_refresh_delay() == 120
assert await service.refresh_once() is renewed
assert guard.deployment is renewed
File diff suppressed because it is too large Load Diff
+161
View File
@@ -0,0 +1,161 @@
from __future__ import annotations
import asyncio
import pytest
from unittest.mock import AsyncMock
from langbot.pkg.cloud.entitlements import EntitlementResolver, EntitlementSnapshot, EntitlementUnavailableError
def _snapshot(**overrides) -> EntitlementSnapshot:
values = {
'instance_uuid': 'instance-a',
'workspace_uuid': 'workspace-a',
'entitlement_revision': 7,
'status': 'active',
'not_before': 100,
'expires_at': 200,
'features': {'managed_sandbox': True, 'mcp_stdio': False},
'limits': {'managed_sandbox_sessions': 1},
}
values.update(overrides)
return EntitlementSnapshot(**values)
def test_active_snapshot_exposes_only_generic_features_and_limits():
snapshot = _snapshot().require_active(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
now=150,
)
snapshot.require_feature('managed_sandbox')
assert snapshot.limit('managed_sandbox_sessions') == 1
assert 'plan' not in snapshot.model_fields
@pytest.mark.parametrize(
'snapshot,now',
[
(_snapshot(status='suspended'), 150),
(_snapshot(), 99),
(_snapshot(), 200),
],
)
def test_inactive_or_expired_snapshot_fails_closed(snapshot, now):
with pytest.raises(EntitlementUnavailableError):
snapshot.require_active(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
now=now,
)
def test_scope_mismatch_fails_closed():
with pytest.raises(EntitlementUnavailableError, match='scope'):
_snapshot().require_active(
instance_uuid='instance-a',
workspace_uuid='workspace-b',
now=150,
)
@pytest.mark.asyncio
async def test_resolver_rejects_revision_rollback():
provider = AsyncMock()
provider.get_workspace_entitlement = AsyncMock(side_effect=[_snapshot(), _snapshot(entitlement_revision=6)])
resolver = EntitlementResolver('instance-a', provider)
await resolver.resolve('workspace-a', now=150)
with pytest.raises(EntitlementUnavailableError, match='rolled back'):
await resolver.resolve('workspace-a', now=150)
@pytest.mark.asyncio
async def test_resolver_rejects_same_revision_with_different_contents():
provider = AsyncMock()
provider.get_workspace_entitlement = AsyncMock(
side_effect=[
_snapshot(),
_snapshot(features={'managed_sandbox': False}),
]
)
resolver = EntitlementResolver('instance-a', provider)
await resolver.resolve('workspace-a', now=150)
with pytest.raises(EntitlementUnavailableError, match='conflicting contents'):
await resolver.resolve('workspace-a', now=150)
@pytest.mark.asyncio
async def test_resolver_checks_deployment_admission_before_and_after_provider_call():
checks = 0
def require_admission() -> None:
nonlocal checks
checks += 1
if checks == 2:
raise RuntimeError('manifest expired during provider call')
provider = AsyncMock()
provider.get_workspace_entitlement = AsyncMock(return_value=_snapshot())
resolver = EntitlementResolver(
'instance-a',
provider,
deployment_admission=require_admission,
)
with pytest.raises(RuntimeError, match='expired during provider call'):
await resolver.resolve('workspace-a', now=150)
assert checks == 2
@pytest.mark.asyncio
async def test_directory_activity_reconciliation_drops_historical_snapshots():
provider = AsyncMock()
provider.get_workspace_entitlement = AsyncMock(return_value=_snapshot())
resolver = EntitlementResolver('instance-a', provider)
await resolver.reconcile_active_workspaces({'workspace-a', 'workspace-b'})
await resolver.resolve('workspace-a', now=150)
await resolver.reconcile_active_workspaces({'workspace-b'})
assert resolver.snapshot_counts() == {
'active_workspaces': 1,
'cached_snapshots': 0,
}
with pytest.raises(EntitlementUnavailableError, match='directory projection'):
await resolver.resolve('workspace-a', now=150)
provider.get_workspace_entitlement.assert_awaited_once()
@pytest.mark.asyncio
async def test_directory_fence_wins_race_with_inflight_entitlement_fetch():
provider_started = asyncio.Event()
release_provider = asyncio.Event()
async def fetch(_workspace_uuid: str) -> EntitlementSnapshot:
provider_started.set()
await release_provider.wait()
return _snapshot()
provider = AsyncMock()
provider.get_workspace_entitlement = AsyncMock(side_effect=fetch)
resolver = EntitlementResolver('instance-a', provider)
await resolver.reconcile_active_workspaces({'workspace-a'})
resolve_task = asyncio.create_task(resolver.resolve('workspace-a', now=150))
await provider_started.wait()
await resolver.update_workspace_activity(
active_workspace_uuids=set(),
inactive_workspace_uuids={'workspace-a'},
)
release_provider.set()
with pytest.raises(EntitlementUnavailableError, match='directory projection'):
await resolve_task
assert resolver.snapshot_counts() == {
'active_workspaces': 0,
'cached_snapshots': 0,
}
+164
View File
@@ -0,0 +1,164 @@
from __future__ import annotations
import base64
import json
import time
import uuid
from types import SimpleNamespace
import pytest
from cryptography.hazmat.primitives import serialization
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
from langbot.pkg.cloud.launch import SpaceLaunchError, SpaceLaunchService
pytestmark = pytest.mark.asyncio
INSTANCE_UUID = 'instance-test'
ACCOUNT_UUID = '11111111-1111-4111-8111-111111111111'
WORKSPACE_UUID = '22222222-2222-4222-8222-222222222222'
KEY_ID = 'space-key-1'
def _base64url(raw: bytes) -> str:
return base64.urlsafe_b64encode(raw).rstrip(b'=').decode('ascii')
def _sign(private_key: Ed25519PrivateKey, claims: dict, *, key_id: str = KEY_ID) -> str:
header = {'alg': 'EdDSA', 'kid': key_id, 'typ': 'langbot-control-plane+jwt'}
encoded_header = _base64url(json.dumps(header, separators=(',', ':')).encode('utf-8'))
encoded_claims = _base64url(json.dumps(claims, separators=(',', ':')).encode('utf-8'))
signing_input = f'{encoded_header}.{encoded_claims}'
return f'{signing_input}.{_base64url(private_key.sign(signing_input.encode("ascii")))}'
def _claims(*, now: int, jti: str | None = None, workspace_uuid: str = WORKSPACE_UUID) -> dict:
return {
'iss': 'langbot-space',
'aud': 'langbot-cloud-runtime',
'sub': f'langbot-instance:{INSTANCE_UUID}',
'jti': jti or str(uuid.uuid4()),
'iat': now,
'nbf': now - 5,
'exp': now + 90,
'instance_uuid': INSTANCE_UUID,
'kind': 'workspace.launch',
'payload': {
'account_uuid': ACCOUNT_UUID,
'workspace_uuid': workspace_uuid,
},
}
def _service(private_key: Ed25519PrivateKey, *, now: int) -> SpaceLaunchService:
public_key = private_key.public_key().public_bytes(
encoding=serialization.Encoding.Raw,
format=serialization.PublicFormat.Raw,
)
app = SimpleNamespace(
deployment=SimpleNamespace(multi_workspace_enabled=True, verification_key_id=KEY_ID),
workspace_service=SimpleNamespace(instance_uuid=INSTANCE_UUID),
instance_config=SimpleNamespace(
data={
'space': {
'launch': {
'control_plane_public_key': _base64url(public_key),
}
}
}
),
)
return SpaceLaunchService(app, wall_time=lambda: now)
async def test_consumes_valid_workspace_launch_assertion_once():
private_key = Ed25519PrivateKey.generate()
now = int(time.time())
service = _service(private_key, now=now)
token = _sign(private_key, _claims(now=now))
launch = await service.consume_assertion(token, expected_workspace_uuid=WORKSPACE_UUID)
assert launch == {'account_uuid': ACCOUNT_UUID, 'workspace_uuid': WORKSPACE_UUID}
with pytest.raises(SpaceLaunchError, match='already been consumed'):
await service.consume_assertion(token, expected_workspace_uuid=WORKSPACE_UUID)
async def test_replay_cache_does_not_scan_all_live_assertions(monkeypatch):
private_key = Ed25519PrivateKey.generate()
now = int(time.time())
service = _service(private_key, now=now)
for index in range(512):
await service._consume_jti(f'jti-{index}', now + 90)
class NoGlobalIterationDict(dict):
def __iter__(self):
raise AssertionError('replay admission scanned all live assertions')
def keys(self):
raise AssertionError('replay admission scanned all live assertions')
def items(self):
raise AssertionError('replay admission scanned all live assertions')
def values(self):
raise AssertionError('replay admission scanned all live assertions')
guarded_jtis = NoGlobalIterationDict(service._consumed_jtis)
monkeypatch.setattr(service, '_consumed_jtis', guarded_jtis)
await service._consume_jti('jti-new', now + 90)
assert len(guarded_jtis) == 513
async def test_replay_cache_fails_closed_at_capacity(monkeypatch):
from langbot.pkg.cloud import launch
private_key = Ed25519PrivateKey.generate()
now = int(time.time())
service = _service(private_key, now=now)
monkeypatch.setattr(launch, '_CONSUMED_JTI_MAX_ENTRIES', 2)
await service._consume_jti('jti-1', now + 90)
await service._consume_jti('jti-2', now + 90)
with pytest.raises(SpaceLaunchError, match='replay cache capacity'):
await service._consume_jti('jti-3', now + 90)
with pytest.raises(SpaceLaunchError, match='already been consumed'):
await service._consume_jti('jti-1', now + 90)
async def test_rejects_expired_wrong_workspace_and_wrong_instance_assertions():
private_key = Ed25519PrivateKey.generate()
now = int(time.time())
service = _service(private_key, now=now)
expired = _claims(now=now)
expired['exp'] = now - 60
with pytest.raises(SpaceLaunchError, match='expired'):
await service.consume_assertion(_sign(private_key, expired), expected_workspace_uuid=WORKSPACE_UUID)
wrong_workspace = _sign(private_key, _claims(now=now, workspace_uuid='33333333-3333-4333-8333-333333333333'))
with pytest.raises(SpaceLaunchError, match='another Workspace'):
await service.consume_assertion(wrong_workspace, expected_workspace_uuid=WORKSPACE_UUID)
wrong_instance = _claims(now=now)
wrong_instance['instance_uuid'] = 'other-instance'
with pytest.raises(SpaceLaunchError, match='instance UUID'):
await service.consume_assertion(_sign(private_key, wrong_instance), expected_workspace_uuid=WORKSPACE_UUID)
async def test_rejects_invalid_signature_and_non_cloud_mode():
private_key = Ed25519PrivateKey.generate()
now = int(time.time())
token = _sign(private_key, _claims(now=now))
service = _service(Ed25519PrivateKey.generate(), now=now)
with pytest.raises(SpaceLaunchError, match='signature'):
await service.consume_assertion(token, expected_workspace_uuid=WORKSPACE_UUID)
oss_service = _service(private_key, now=now)
oss_service.ap.deployment.multi_workspace_enabled = False
with pytest.raises(SpaceLaunchError, match='verified Cloud mode'):
await oss_service.consume_assertion(token, expected_workspace_uuid=WORKSPACE_UUID)
@@ -0,0 +1,113 @@
from __future__ import annotations
import asyncio
from types import SimpleNamespace
import pytest
from langbot.pkg.core.app import Application
pytestmark = pytest.mark.asyncio
class _TaskManager:
def __init__(self, stop: asyncio.Event) -> None:
self.stop = stop
self.tasks: list[asyncio.Task] = []
def create_task(self, coro, *, name='', **_kwargs):
task = asyncio.create_task(coro, name=name)
self.tasks.append(task)
return SimpleNamespace(task=task)
async def wait_all(self) -> None:
await self.stop.wait()
for task in self.tasks:
task.cancel()
await asyncio.gather(*self.tasks, return_exceptions=True)
async def _wait_forever() -> None:
await asyncio.Event().wait()
async def test_resource_maintenance_waits_and_shares_workspace_discovery() -> None:
stop = asyncio.Event()
completed = asyncio.Event()
discovery_calls = 0
job_calls: list[str] = []
async def list_bindings():
nonlocal discovery_calls
discovery_calls += 1
return [
SimpleNamespace(
instance_uuid='instance',
workspace_uuid='workspace',
placement_generation=1,
)
]
async def cleanup_monitoring(_context, _retention_days, *, batch_size):
assert batch_size == 10
job_calls.append('monitoring')
return {}
async def cleanup_storage(_context):
job_calls.append('storage')
completed.set()
return {}
application = Application()
application.event_loop = asyncio.get_running_loop()
application.event_loop_monitor = SimpleNamespace(start=lambda: None)
application.task_mgr = _TaskManager(stop)
application.plugin_connector = SimpleNamespace(initialize_plugins=lambda: asyncio.sleep(0))
application.platform_mgr = SimpleNamespace(run=_wait_forever)
application.ctrl = SimpleNamespace(run=_wait_forever)
application.http_ctrl = SimpleNamespace(run=_wait_forever)
application.telemetry = None
application.workspace_collaboration_service = None
application.workspace_service = SimpleNamespace(list_active_execution_bindings=list_bindings)
application.monitoring_service = SimpleNamespace(cleanup_expired_records=cleanup_monitoring)
application.maintenance_service = SimpleNamespace(cleanup_expired_files=cleanup_storage)
application.instance_config = SimpleNamespace(
data={
'monitoring': {
'auto_cleanup': {
'enabled': True,
'retention_days': 30,
'delete_batch_size': 10,
'check_interval_hours': 0.00002,
}
},
'storage': {
'cleanup': {
'enabled': True,
'check_interval_hours': 0.00002,
}
},
}
)
application.logger = SimpleNamespace(
info=lambda *_args, **_kwargs: None,
warning=lambda *_args, **_kwargs: None,
error=lambda *_args, **_kwargs: None,
debug=lambda *_args, **_kwargs: None,
)
async def no_web_info() -> None:
return None
application.print_web_access_info = no_web_info
run_task = asyncio.create_task(application.run())
try:
await asyncio.sleep(0.01)
assert discovery_calls == 0
await asyncio.wait_for(completed.wait(), timeout=1)
assert discovery_calls == 1
assert job_calls == ['monitoring', 'storage']
finally:
stop.set()
await asyncio.wait_for(run_task, timeout=1)
+142
View File
@@ -0,0 +1,142 @@
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
from langbot.pkg.core.app import Application
@pytest.mark.asyncio
async def test_shutdown_closes_mcp_session_manager_once() -> None:
app = Application()
stop_session_manager = AsyncMock()
app.platform_mgr = SimpleNamespace(shutdown=AsyncMock())
app.tool_mgr = SimpleNamespace(shutdown=AsyncMock())
app.model_mgr = SimpleNamespace(shutdown=AsyncMock())
app.box_service = SimpleNamespace(shutdown=AsyncMock())
app.plugin_connector = SimpleNamespace(aclose=AsyncMock())
app.telemetry = SimpleNamespace(shutdown=AsyncMock())
app.vector_db_mgr = SimpleNamespace(shutdown=AsyncMock())
app.storage_mgr = SimpleNamespace(shutdown=AsyncMock())
manifest_provider = SimpleNamespace(aclose=AsyncMock())
app.deployment = SimpleNamespace(manifest_provider=manifest_provider)
persistence_engine = SimpleNamespace(dispose=AsyncMock())
app.persistence_mgr = SimpleNamespace(db=SimpleNamespace(engine=persistence_engine))
app.http_ctrl = SimpleNamespace(mcp_mount=SimpleNamespace(stop_session_manager=stop_session_manager))
await app.shutdown()
await app.shutdown()
stop_session_manager.assert_awaited_once()
app.platform_mgr.shutdown.assert_awaited_once()
app.tool_mgr.shutdown.assert_awaited_once()
app.model_mgr.shutdown.assert_awaited_once()
app.box_service.shutdown.assert_awaited_once()
app.plugin_connector.aclose.assert_awaited_once()
app.telemetry.shutdown.assert_awaited_once()
app.vector_db_mgr.shutdown.assert_awaited_once()
app.storage_mgr.shutdown.assert_awaited_once()
manifest_provider.aclose.assert_awaited_once()
persistence_engine.dispose.assert_awaited_once()
@pytest.mark.asyncio
async def test_dispose_tracks_only_one_shutdown_task() -> None:
app = Application()
app.event_loop = asyncio.get_running_loop()
app.dispose()
shutdown_task = app._shutdown_task
app.dispose()
assert shutdown_task is not None
assert app._shutdown_task is shutdown_task
await shutdown_task
app.dispose()
assert app._shutdown_task is shutdown_task
@pytest.mark.asyncio
async def test_runtime_resource_stats_are_aggregate_and_constant_time() -> None:
app = Application()
app.event_loop = asyncio.get_running_loop()
app.blocking_executor = SimpleNamespace(
snapshot=lambda: {
'inflight': 3,
'running': 2,
'pending': 1,
'rejected_total': 4,
}
)
app.task_mgr = SimpleNamespace(get_stats=lambda: {'total': 5, 'completed': 2})
app.query_pool = SimpleNamespace(
queries=[object()],
cached_queries={},
active_query_count_by_workspace={'workspace-a': 1},
)
app.model_mgr = SimpleNamespace(
provider_dict={'provider': object()},
llm_model_dict={},
embedding_model_dict={},
rerank_model_dict={},
)
app.platform_mgr = SimpleNamespace(_bots_by_key={})
app.pipeline_mgr = SimpleNamespace(_pipelines_by_key={})
app.rag_mgr = SimpleNamespace(knowledge_bases={})
app.plugin_connector = SimpleNamespace(_known_desired_states={'installation': object()})
app.persistence_mgr = SimpleNamespace(
get_resource_stats=lambda: {
'configured_capacity': 20,
'checked_out': 3,
}
)
app.directory_projection_service = SimpleNamespace(
resource_snapshot=lambda: {
'active_workspaces': 10,
'max_active_workspaces': 1000,
}
)
app.tool_mgr = SimpleNamespace(
mcp_tool_loader=SimpleNamespace(
_sessions={},
_hosted_mcp_tasks=[],
_host_dispatch_tasks=set(),
)
)
app.telemetry = SimpleNamespace(send_tasks=[])
stats = app.get_runtime_resource_stats()
assert stats['asyncio_tasks'] >= 1
assert stats['event_loop'] == {
'running': False,
'samples_total': 0,
'last_lag_ms': 0,
'recent_p95_lag_ms': 0,
'recent_max_lag_ms': 0,
'max_lag_ms': 0,
}
assert stats['blocking_executor']['rejected_total'] == 4
assert stats['application_tasks'] == {
'total': 5,
'completed': 2,
}
assert stats['database_pool'] == {
'configured_capacity': 20,
'checked_out': 3,
}
assert stats['directory'] == {
'active_workspaces': 10,
'max_active_workspaces': 1000,
}
assert stats['query_pool'] == {
'queued': 1,
'cached': 0,
'active_workspaces': 1,
}
assert stats['models']['providers'] == 1
assert stats['runtimes']['plugin_installations'] == 1
+58 -20
View File
@@ -2,12 +2,37 @@ from __future__ import annotations
import signal
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.core import boot
@pytest.mark.asyncio
async def test_make_app_shuts_down_partially_built_application(monkeypatch):
app_inst = SimpleNamespace(
event_loop=None,
shutdown=AsyncMock(),
initialize=AsyncMock(),
)
class FailingStage:
async def run(self, ap):
assert ap is app_inst
raise RuntimeError('startup failed')
monkeypatch.setattr(boot.app, 'Application', lambda: app_inst)
monkeypatch.setattr(boot, 'stage_order', ['FailingStage'])
monkeypatch.setitem(boot.stage.preregistered_stages, 'FailingStage', FailingStage)
with pytest.raises(RuntimeError, match='startup failed'):
await boot.make_app(SimpleNamespace())
app_inst.shutdown.assert_awaited_once()
app_inst.initialize.assert_not_awaited()
@pytest.mark.asyncio
async def test_main_signal_handler_handles_sigint_before_app_created(monkeypatch):
captured_handler = {}
@@ -18,47 +43,60 @@ async def test_main_signal_handler_handles_sigint_before_app_created(monkeypatch
async def fake_make_app(loop):
captured_handler[signal.SIGINT](signal.SIGINT, None)
def fake_exit(code):
raise SystemExit(code)
monkeypatch.setattr(signal, 'signal', fake_signal)
monkeypatch.setattr(boot, 'make_app', fake_make_app)
monkeypatch.setattr(boot.os, '_exit', fake_exit)
with pytest.raises(SystemExit) as exc_info:
await boot.main(SimpleNamespace())
assert exc_info.value.code == 0
await boot.main(SimpleNamespace())
@pytest.mark.asyncio
async def test_main_signal_handler_disposes_created_app(monkeypatch):
captured_handler = {}
app_inst = SimpleNamespace(disposed=False)
app_inst = SimpleNamespace(shutdown_called=False)
def fake_signal(sig, handler):
captured_handler[sig] = handler
def dispose():
app_inst.disposed = True
async def shutdown():
app_inst.shutdown_called = True
async def run():
captured_handler[signal.SIGINT](signal.SIGINT, None)
async def fake_make_app(loop):
app_inst.dispose = dispose
app_inst.shutdown = shutdown
app_inst.run = run
return app_inst
def fake_exit(code):
raise SystemExit(code)
monkeypatch.setattr(signal, 'signal', fake_signal)
monkeypatch.setattr(boot, 'make_app', fake_make_app)
monkeypatch.setattr(boot.os, '_exit', fake_exit)
with pytest.raises(SystemExit) as exc_info:
await boot.main(SimpleNamespace())
await boot.main(SimpleNamespace())
assert exc_info.value.code == 0
assert app_inst.disposed is True
assert app_inst.shutdown_called is True
@pytest.mark.asyncio
async def test_main_reports_app_run_failure_and_still_shuts_down(monkeypatch):
app_inst = SimpleNamespace(shutdown_called=False)
async def shutdown():
app_inst.shutdown_called = True
async def run():
raise RuntimeError('run failed')
async def fake_make_app(loop):
app_inst.shutdown = shutdown
app_inst.run = run
return app_inst
print_exc = Mock()
monkeypatch.setattr(signal, 'signal', lambda *_args: None)
monkeypatch.setattr(boot, 'make_app', fake_make_app)
monkeypatch.setattr(boot.traceback, 'print_exc', print_exc)
await boot.main(SimpleNamespace())
print_exc.assert_called_once()
assert app_inst.shutdown_called is True
+134
View File
@@ -35,6 +35,22 @@ class TestApplyEnvOverridesToConfig:
assert result['system']['name'] == 'custom_name'
def test_override_log_never_prints_secret_value(self, capsys):
"""Environment-backed credentials must not be copied into logs."""
load_config = get_load_config_module()
secret = 'database-password-that-must-not-leak'
cfg = {'database': {'postgresql': {'password': ''}}}
env = {'DATABASE__POSTGRESQL__PASSWORD': secret}
with patch.dict(os.environ, env, clear=True):
result = load_config._apply_env_overrides_to_config(cfg)
captured = capsys.readouterr().out
assert result['database']['postgresql']['password'] == secret
assert 'DATABASE__POSTGRESQL__PASSWORD' in captured
assert secret not in captured
def test_override_int_value(self):
"""Test overriding an int value with proper conversion."""
load_config = get_load_config_module()
@@ -48,6 +64,20 @@ class TestApplyEnvOverridesToConfig:
assert result['concurrency']['pipeline'] == 10
assert isinstance(result['concurrency']['pipeline'], int)
def test_cloud_directory_limit_override_keeps_integer_type_on_upgraded_config(self):
load_config = get_load_config_module()
cfg = load_config._complete_runtime_policy_defaults({})
with patch.dict(
os.environ,
{'CLOUD__DIRECTORY__MAX_ACTIVE_WORKSPACES': '250'},
clear=True,
):
result = load_config._apply_env_overrides_to_config(cfg)
assert result['cloud']['directory']['max_active_workspaces'] == 250
assert isinstance(result['cloud']['directory']['max_active_workspaces'], int)
def test_override_int_value_invalid_conversion(self):
"""Test that invalid int conversion keeps string value."""
load_config = get_load_config_module()
@@ -122,6 +152,19 @@ class TestApplyEnvOverridesToConfig:
assert result['system']['disabled_adapters'] == ['aiocqhttp', 'dingtalk', 'telegram']
def test_override_integer_list_preserves_item_type(self):
"""Comma-separated overrides inherit the existing list item type."""
load_config = get_load_config_module()
cfg = {'vdb': {'pgvector': {'allowed_dimensions': [384, 512]}}}
env = {'VDB__PGVECTOR__ALLOWED_DIMENSIONS': '384,512,768'}
with patch.dict(os.environ, env, clear=True):
result = load_config._apply_env_overrides_to_config(cfg)
assert result['vdb']['pgvector']['allowed_dimensions'] == [384, 512, 768]
assert all(isinstance(item, int) for item in result['vdb']['pgvector']['allowed_dimensions'])
def test_override_list_value_empty_items(self):
"""Test that empty items in comma-separated list are filtered."""
load_config = get_load_config_module()
@@ -196,6 +239,19 @@ class TestApplyEnvOverridesToConfig:
assert result['system']['name'] == 'default'
def test_skip_env_vars_with_empty_path_segments(self, capsys):
"""Platform variables such as __CF_USER_TEXT_ENCODING are not config."""
load_config = get_load_config_module()
cfg = {'system': {'name': 'default'}}
env = {'__CF_USER_TEXT_ENCODING': '0x1F5:0x0:0x64'}
with patch.dict(os.environ, env, clear=True):
result = load_config._apply_env_overrides_to_config(cfg)
assert result == cfg
assert capsys.readouterr().out == ''
def test_nested_config_path(self):
"""Test overriding deeply nested config."""
load_config = get_load_config_module()
@@ -259,6 +315,84 @@ class TestApplyEnvOverridesToConfig:
assert result['system']['enable'] is False
assert result['concurrency']['pipeline'] == 10
def test_plugin_worker_and_stdio_policy_native_env_overrides(self):
load_config = get_load_config_module()
cfg = {
'plugin': {
'worker': {
'max_cpus': 1.0,
'max_memory_mb': 512,
'max_pids': 128,
'max_open_files': 256,
'max_file_size_mb': 512,
'max_concurrent_restarts': 1,
'restart_failure_threshold': 8,
'restart_failure_window_seconds': 30.0,
'restart_circuit_open_seconds': 60.0,
}
},
'mcp': {'stdio': {'enabled': True}},
}
env = {
'PLUGIN__WORKER__MAX_CPUS': '2.5',
'PLUGIN__WORKER__MAX_MEMORY_MB': '1024',
'PLUGIN__WORKER__MAX_PIDS': '64',
'PLUGIN__WORKER__MAX_OPEN_FILES': '128',
'PLUGIN__WORKER__MAX_FILE_SIZE_MB': '256',
'PLUGIN__WORKER__MAX_CONCURRENT_RESTARTS': '2',
'PLUGIN__WORKER__RESTART_FAILURE_THRESHOLD': '12',
'PLUGIN__WORKER__RESTART_FAILURE_WINDOW_SECONDS': '45.5',
'PLUGIN__WORKER__RESTART_CIRCUIT_OPEN_SECONDS': '90.0',
'MCP__STDIO__ENABLED': 'false',
}
with patch.dict(os.environ, env, clear=True):
result = load_config._apply_env_overrides_to_config(cfg)
assert result['plugin']['worker'] == {
'max_cpus': 2.5,
'max_memory_mb': 1024,
'max_pids': 64,
'max_open_files': 128,
'max_file_size_mb': 256,
'max_concurrent_restarts': 2,
'restart_failure_threshold': 12,
'restart_failure_window_seconds': 45.5,
'restart_circuit_open_seconds': 90.0,
}
assert result['mcp']['stdio']['enabled'] is False
def test_runtime_policy_defaults_preserve_env_types_for_upgraded_config(self):
load_config = get_load_config_module()
cfg = {'plugin': {'enable': True}}
completed = load_config._complete_runtime_policy_defaults(cfg)
with patch.dict(
os.environ,
{
'PLUGIN__WORKER__MAX_MEMORY_MB': '768',
'MCP__STDIO__ENABLED': 'false',
'SYSTEM__BLOCKING_EXECUTOR__MAX_WORKERS': '12',
'SYSTEM__BLOCKING_EXECUTOR__MAX_PENDING': '256',
'SYSTEM__BLOCKING_EXECUTOR__MAX_INFLIGHT_PER_SCOPE': '3',
},
clear=True,
):
result = load_config._apply_env_overrides_to_config(completed)
assert result['system']['blocking_executor'] == {
'max_workers': 12,
'max_pending': 256,
'max_inflight_per_scope': 3,
}
assert isinstance(
result['system']['blocking_executor']['max_workers'],
int,
)
assert result['plugin']['worker']['max_memory_mb'] == 768
assert isinstance(result['plugin']['worker']['max_memory_mb'], int)
assert result['mcp']['stdio']['enabled'] is False
def test_webhook_prefix_override(self):
"""Test overriding webhook_prefix via environment variable."""
load_config = get_load_config_module()
+121
View File
@@ -12,6 +12,8 @@ from __future__ import annotations
import pytest
import asyncio
import contextvars
import inspect
import sys
from unittest.mock import Mock, MagicMock
from contextlib import contextmanager
@@ -264,6 +266,28 @@ class TestTaskWrapper:
wrapper.cancel()
@pytest.mark.asyncio
async def test_workspace_task_sets_blocking_work_scope(self):
"""Detached tasks recover tenant fairness from durable ownership."""
_, TaskWrapper, _ = get_taskmgr_classes()
from langbot.pkg.utils.bounded_executor import (
current_blocking_work_scope,
)
mock_app = create_mock_app()
async def read_scope():
return current_blocking_work_scope()
wrapper = TaskWrapper(
mock_app,
read_scope(),
workspace_uuid='workspace-a',
)
assert await wrapper.task == 'workspace-a'
assert current_blocking_work_scope() is None
@pytest.mark.asyncio
async def test_to_dict_serialization(self):
"""Test TaskWrapper.to_dict serialization."""
@@ -360,6 +384,53 @@ class TestAsyncTaskManager:
wrapper.cancel()
@pytest.mark.asyncio
async def test_create_task_does_not_inherit_request_context(self):
"""Long-lived tasks must receive identity through explicit arguments."""
_, _, AsyncTaskManager = get_taskmgr_classes()
mock_app = create_mock_app()
manager = AsyncTaskManager(mock_app)
request_value = contextvars.ContextVar('request_value', default=None)
token = request_value.set('request-scoped-transaction')
observed = []
async def detached_task(captured_workspace: str) -> None:
observed.append((request_value.get(), captured_workspace))
try:
wrapper = manager.create_task(detached_task('workspace-a'))
await wrapper.task
finally:
request_value.reset(token)
assert observed == [(None, 'workspace-a')]
@pytest.mark.asyncio
async def test_create_task_waits_for_registered_transaction_commit(self):
_, _, AsyncTaskManager = get_taskmgr_classes()
mock_app = create_mock_app()
gate = asyncio.get_running_loop().create_future()
class PersistenceManagerStub:
def create_after_commit_gate(self):
return gate
mock_app.persistence_mgr = PersistenceManagerStub()
manager = AsyncTaskManager(mock_app)
observed = []
async def background_work() -> None:
observed.append('started')
wrapper = manager.create_task(background_work())
await asyncio.sleep(0)
assert observed == []
gate.set_result(None)
await wrapper.task
assert observed == ['started']
@pytest.mark.asyncio
async def test_get_stats_counts_correctly(self):
"""Test get_stats returns correct counts."""
@@ -482,6 +553,56 @@ class TestAsyncTaskManager:
wrapper.cancel()
@pytest.mark.asyncio
async def test_create_user_task_enforces_workspace_active_limit_and_closes_rejected_coroutine(self):
"""A noisy Workspace cannot accumulate unbounded background work."""
_, _, AsyncTaskManager = get_taskmgr_classes()
mock_app = create_mock_app()
mock_app.instance_config.data['system']['task_retention'].update(
{
'max_active_user_tasks': 10,
'max_active_user_tasks_per_workspace': 1,
}
)
manager = AsyncTaskManager(mock_app)
async def long_coro():
await asyncio.sleep(10)
first = manager.create_user_task(long_coro(), workspace_uuid='workspace-a')
rejected = long_coro()
with pytest.raises(RuntimeError, match='Workspace has too many active user operations'):
manager.create_user_task(rejected, workspace_uuid='workspace-a')
assert inspect.getcoroutinestate(rejected) == inspect.CORO_CLOSED
other_workspace = manager.create_user_task(long_coro(), workspace_uuid='workspace-b')
first.cancel()
other_workspace.cancel()
@pytest.mark.asyncio
async def test_create_user_task_enforces_instance_active_limit(self):
"""The shared process retains a hard cap even across Workspaces."""
_, _, AsyncTaskManager = get_taskmgr_classes()
mock_app = create_mock_app()
mock_app.instance_config.data['system']['task_retention'].update(
{
'max_active_user_tasks': 1,
'max_active_user_tasks_per_workspace': 10,
}
)
manager = AsyncTaskManager(mock_app)
async def long_coro():
await asyncio.sleep(10)
first = manager.create_user_task(long_coro(), workspace_uuid='workspace-a')
rejected = long_coro()
with pytest.raises(RuntimeError, match='instance has too many active user operations'):
manager.create_user_task(rejected, workspace_uuid='workspace-b')
assert inspect.getcoroutinestate(rejected) == inspect.CORO_CLOSED
first.cancel()
@pytest.mark.asyncio
async def test_get_task_by_id(self):
"""Test get_task_by_id returns correct task."""
@@ -11,9 +11,19 @@ Note: Uses import isolation to break circular import chains.
from __future__ import annotations
import sys
from unittest.mock import Mock, MagicMock
from contextlib import contextmanager
from typing import Generator
from unittest.mock import MagicMock, Mock
import pytest
@pytest.fixture(autouse=True)
def isolate_database_manager_registry(monkeypatch: pytest.MonkeyPatch) -> None:
"""Keep decorator tests from mutating the process-wide manager registry."""
from langbot.pkg.persistence import database
monkeypatch.setattr(database, 'preregistered_managers', list(database.preregistered_managers))
@contextmanager
@@ -0,0 +1,233 @@
from __future__ import annotations
from types import SimpleNamespace
import pytest
# Persistence manager performs the package's database-manager registration;
# importing a concrete manager first would enter the historical app/mgr cycle.
from langbot.pkg.persistence import mgr as _persistence_mgr # noqa: F401
from langbot.pkg.persistence.databases import postgresql
@pytest.mark.asyncio
async def test_postgresql_manager_parses_explicit_url_without_string_reassembly(monkeypatch) -> None:
captured = None
captured_options = None
sentinel_engine = object()
def create_engine(url, **options):
nonlocal captured, captured_options
captured = url
captured_options = options
return sentinel_engine
monkeypatch.setattr(postgresql.sqlalchemy_asyncio, 'create_async_engine', create_engine)
ap = SimpleNamespace(
instance_config=SimpleNamespace(
data={
'database': {
'postgresql': {
'url': 'postgresql://runtime:p%40ss@db.internal:5432/langbot?sslmode=require',
}
}
}
)
)
manager = postgresql.PostgreSQLDatabaseManager(ap)
await manager.initialize()
assert captured.drivername == 'postgresql+asyncpg'
assert captured.password == 'p@ss'
assert captured.query['ssl'] == 'require'
assert 'sslmode' not in captured.query
assert captured_options == {
'pool_size': 10,
'max_overflow': 10,
'pool_timeout': 30,
'pool_recycle': 1800,
'pool_pre_ping': True,
}
assert manager.engine is sentinel_engine
@pytest.mark.asyncio
async def test_postgresql_manager_builds_structured_url_with_special_password(monkeypatch) -> None:
captured = None
def create_engine(url, **_options):
nonlocal captured
captured = url
return object()
monkeypatch.setattr(postgresql.sqlalchemy_asyncio, 'create_async_engine', create_engine)
ap = SimpleNamespace(
instance_config=SimpleNamespace(
data={
'database': {
'postgresql': {
'host': 'db.internal',
'port': 5432,
'user': 'runtime',
'password': 'p@ss:/?#word',
'database': 'langbot',
}
}
}
)
)
await postgresql.PostgreSQLDatabaseManager(ap).initialize()
assert captured.password == 'p@ss:/?#word'
assert captured.host == 'db.internal'
assert captured.database == 'langbot'
@pytest.mark.asyncio
async def test_postgresql_manager_applies_explicit_bounded_pool_options(monkeypatch) -> None:
captured_options = None
def create_engine(_url, **options):
nonlocal captured_options
captured_options = options
return object()
monkeypatch.setattr(postgresql.sqlalchemy_asyncio, 'create_async_engine', create_engine)
ap = SimpleNamespace(
instance_config=SimpleNamespace(
data={
'database': {
'postgresql': {
'pool_size': 24,
'max_overflow': 0,
'pool_timeout_seconds': 7,
'pool_recycle_seconds': 600,
}
}
}
)
)
await postgresql.PostgreSQLDatabaseManager(ap).initialize()
assert captured_options == {
'pool_size': 24,
'max_overflow': 0,
'pool_timeout': 7,
'pool_recycle': 600,
'pool_pre_ping': True,
}
@pytest.mark.asyncio
async def test_cloud_postgresql_manager_applies_bounded_server_timeouts(monkeypatch) -> None:
captured_options = None
def create_engine(_url, **options):
nonlocal captured_options
captured_options = options
return object()
monkeypatch.setattr(postgresql.sqlalchemy_asyncio, 'create_async_engine', create_engine)
ap = SimpleNamespace(
instance_config=SimpleNamespace(
data={
'database': {
'postgresql': {
'statement_timeout_ms': 45_000,
'lock_timeout_ms': 4_000,
'idle_in_transaction_session_timeout_ms': 55_000,
}
}
}
)
)
manager = postgresql.PostgreSQLDatabaseManager(ap)
manager.persistence_mode = 'cloud_runtime'
await manager.initialize()
assert captured_options['connect_args'] == {
'server_settings': {
'statement_timeout': '45000',
'lock_timeout': '4000',
'idle_in_transaction_session_timeout': '55000',
}
}
@pytest.mark.asyncio
@pytest.mark.parametrize(
('name', 'value'),
[
('pool_size', 0),
('pool_size', True),
('max_overflow', -1),
('pool_size', 101),
('max_overflow', 101),
('pool_timeout_seconds', 0),
('pool_timeout_seconds', 301),
('pool_recycle_seconds', '1800'),
('pool_recycle_seconds', 86401),
],
)
async def test_postgresql_manager_rejects_invalid_pool_options(name, value) -> None:
ap = SimpleNamespace(instance_config=SimpleNamespace(data={'database': {'postgresql': {name: value}}}))
with pytest.raises(ValueError, match=rf'database\.postgresql\.{name}'):
await postgresql.PostgreSQLDatabaseManager(ap).initialize()
@pytest.mark.asyncio
async def test_postgresql_manager_rejects_combined_pool_capacity_above_hard_ceiling() -> None:
ap = SimpleNamespace(
instance_config=SimpleNamespace(
data={
'database': {
'postgresql': {
'pool_size': 60,
'max_overflow': 41,
}
}
}
)
)
with pytest.raises(ValueError, match=r'pool_size \+ max_overflow'):
await postgresql.PostgreSQLDatabaseManager(ap).initialize()
@pytest.mark.asyncio
@pytest.mark.parametrize(
('name', 'value'),
[
('statement_timeout_ms', 0),
('statement_timeout_ms', 300_001),
('lock_timeout_ms', 60_001),
('idle_in_transaction_session_timeout_ms', True),
('idle_in_transaction_session_timeout_ms', 300_001),
],
)
async def test_cloud_postgresql_manager_rejects_unsafe_server_timeouts(name, value) -> None:
ap = SimpleNamespace(instance_config=SimpleNamespace(data={'database': {'postgresql': {name: value}}}))
with pytest.raises(ValueError, match=rf'database\.postgresql\.{name}'):
manager = postgresql.PostgreSQLDatabaseManager(ap)
manager.persistence_mode = 'cloud_runtime'
await manager.initialize()
@pytest.mark.asyncio
async def test_postgresql_manager_rejects_non_postgresql_url_without_echoing_secret() -> None:
ap = SimpleNamespace(
instance_config=SimpleNamespace(
data={'database': {'postgresql': {'url': 'sqlite:///operator-super-secret.db'}}}
)
)
manager = postgresql.PostgreSQLDatabaseManager(ap)
with pytest.raises(ValueError, match='valid PostgreSQL') as exc_info:
await manager.initialize()
assert 'operator-super-secret' not in str(exc_info.value)
@@ -0,0 +1,155 @@
from __future__ import annotations
import logging
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import sqlalchemy
from langbot.__main__ import _build_parser
from langbot.pkg.persistence import release_migration
from langbot.pkg.persistence.mgr import PersistenceManager, PersistenceMode
def _cloud_config(*, database_use: str = 'postgresql', runtime_user: str = 'langbot_runtime') -> dict:
return {
'database': {
'use': database_use,
'postgresql': {
'host': 'runtime-db',
'port': 5432,
'user': runtime_user,
'password': 'runtime-secret',
'database': 'langbot',
},
'cloud_migration': {
'operator_dsn_env': 'TEST_LANGBOT_OPERATOR_DSN',
},
},
'vdb': {
'use': 'pgvector',
'pgvector': {
'use_business_database': True,
'allowed_dimensions': [384, 1536],
},
},
}
def _operator_environ(
*,
user: str = 'langbot_migrator',
database: str = 'langbot',
host: str = 'runtime-db',
port: int = 5432,
) -> dict[str, str]:
return {
'TEST_LANGBOT_OPERATOR_DSN': (f'postgresql://{user}:operator%40secret@{host}:{port}/{database}?sslmode=require')
}
def test_cloud_migration_cli_is_explicit() -> None:
args = _build_parser().parse_args(['migrate', '--cloud'])
assert args.command == 'migrate'
assert args.cloud is True
with pytest.raises(SystemExit) as exc_info:
_build_parser().parse_args(['migrate'])
assert exc_info.value.code == 2
def test_operator_url_is_separate_and_preserves_escaped_secret() -> None:
url = release_migration._operator_database_url(
_cloud_config(),
environ=_operator_environ(),
)
assert url.drivername == 'postgresql+asyncpg'
assert url.username == 'langbot_migrator'
assert url.password == 'operator@secret'
assert url.host == 'runtime-db'
assert url.port == 5432
assert url.database == 'langbot'
assert url.query['ssl'] == 'require'
assert 'sslmode' not in url.query
@pytest.mark.parametrize(
('config', 'environ', 'message'),
[
(_cloud_config(database_use='sqlite'), _operator_environ(), 'SQLite fallback is forbidden'),
(_cloud_config(), {}, 'requires the operator DSN'),
(_cloud_config(), {'TEST_LANGBOT_OPERATOR_DSN': 'not a secret://operator-password'}, 'DSN is invalid'),
(
_cloud_config(),
{'TEST_LANGBOT_OPERATOR_DSN': 'postgresql://operator:secret@runtime-db:not-a-port/langbot'},
'DSN is invalid',
),
(_cloud_config(), _operator_environ(user='langbot_runtime'), 'distinct operator role'),
(_cloud_config(), _operator_environ(database='another_database'), 'configured runtime database'),
(_cloud_config(), _operator_environ(host='other-cluster'), 'runtime PostgreSQL endpoint'),
(_cloud_config(), _operator_environ(port=6432), 'runtime PostgreSQL endpoint'),
],
)
def test_operator_url_rejects_unsafe_configuration(config: dict, environ: dict[str, str], message: str) -> None:
with pytest.raises(release_migration.CloudReleaseMigrationConfigurationError, match=message) as exc_info:
release_migration._operator_database_url(config, environ=environ)
assert 'operator-password' not in str(exc_info.value)
@pytest.mark.asyncio
async def test_release_migration_disposes_operator_engine_on_failure(monkeypatch) -> None:
engine = SimpleNamespace(dispose=AsyncMock())
manager = SimpleNamespace(
db=SimpleNamespace(engine=engine),
initialize=AsyncMock(side_effect=RuntimeError('migration failed')),
shutdown=AsyncMock(side_effect=engine.dispose),
)
def manager_factory(*args, **kwargs):
del args, kwargs
return manager
monkeypatch.setattr(release_migration, 'PersistenceManager', manager_factory)
ap = SimpleNamespace(
instance_config=SimpleNamespace(data=_cloud_config()),
logger=logging.getLogger('release-migration-disposal-test'),
persistence_mgr=None,
)
with pytest.raises(RuntimeError, match='migration failed'):
await release_migration.run_cloud_release_migration(ap, environ=_operator_environ())
assert ap.persistence_mgr is manager
manager.shutdown.assert_awaited_once()
engine.dispose.assert_awaited_once()
@pytest.mark.asyncio
async def test_release_mode_rejects_sqlite_before_schema_changes(tmp_path, monkeypatch) -> None:
from langbot.pkg.persistence import mgr as persistence_mgr_module
from langbot.pkg.persistence.databases.sqlite import SQLiteDatabaseManager
monkeypatch.setattr(persistence_mgr_module.database, 'preregistered_managers', [SQLiteDatabaseManager])
ap = SimpleNamespace(
instance_config=SimpleNamespace(
data={
'database': {
'use': 'sqlite',
'sqlite': {'path': str(tmp_path / 'must-not-migrate.db')},
}
}
),
logger=logging.getLogger('release-migration-sqlite-rejection-test'),
)
manager = PersistenceManager(ap, mode=PersistenceMode.RELEASE_MIGRATION)
with pytest.raises(RuntimeError, match='requires PostgreSQL'):
await manager.initialize()
await manager.get_db_engine().dispose()
engine = sqlalchemy.create_engine(f'sqlite:///{tmp_path / "must-not-migrate.db"}')
try:
assert sqlalchemy.inspect(engine).get_table_names() == []
finally:
engine.dispose()
File diff suppressed because it is too large Load Diff
+24
View File
@@ -17,6 +17,7 @@ from unittest.mock import AsyncMock, Mock
# this, running a stage test in isolation triggers a circular-import error:
# stage.py → core.app → pipelinemgr → stage.stage_class (not yet bound).
import langbot.pkg.pipeline.pipelinemgr # noqa: F401
from langbot.pkg.api.http.context import ExecutionContext
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
import langbot_plugin.api.entities.builtin.platform.message as platform_message
@@ -43,6 +44,14 @@ class MockApplication:
self.query_pool = self._create_mock_query_pool()
self.instance_config = self._create_mock_instance_config()
self.task_mgr = self._create_mock_task_manager()
self.workspace_service = AsyncMock()
self.workspace_service.get_execution_binding = AsyncMock(
return_value=Mock(
instance_uuid='test-instance',
workspace_uuid='test-workspace',
placement_generation=1,
)
)
# Skill manager is optional; PreProcessor only touches it for the
# local-agent runner. None keeps the skill-binding branch inert.
self.skill_mgr = None
@@ -86,6 +95,7 @@ class MockApplication:
query_pool.cached_queries = {}
query_pool.queries = []
query_pool.condition = AsyncMock()
query_pool.remove_query = AsyncMock(return_value=True)
return query_pool
def _create_mock_instance_config(self):
@@ -194,6 +204,9 @@ def sample_query(sample_message_chain, sample_message_event, mock_adapter):
# Use model_construct to bypass Pydantic validation for test purposes
query = pipeline_query.Query.model_construct(
instance_uuid='test-instance',
workspace_uuid='test-workspace',
placement_generation=1,
query_id='test-query-id',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -227,6 +240,17 @@ def sample_query(sample_message_chain, sample_message_event, mock_adapter):
resp_message_chain=None,
current_stage_name=None,
)
object.__setattr__(
query,
'_execution_context',
ExecutionContext(
instance_uuid='test-instance',
workspace_uuid='test-workspace',
placement_generation=1,
bot_uuid='test-bot-uuid',
pipeline_uuid='test-pipeline-uuid',
),
)
return query
+423 -24
View File
@@ -13,8 +13,11 @@ from __future__ import annotations
import pytest
import asyncio
import contextvars
from contextlib import asynccontextmanager
from unittest.mock import Mock, AsyncMock
from importlib import import_module
from types import SimpleNamespace
from tests.factories import (
FakeApp,
@@ -25,6 +28,49 @@ from tests.factories import (
import langbot_plugin.api.entities.builtin.provider.session as provider_session
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.pipeline.pool import (
ExecutionContextMismatchError,
ExecutionContextRequiredError,
bind_execution_context,
)
from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError
def execution_context(
workspace_uuid='workspace-test',
*,
bot_uuid='test-bot',
pipeline_uuid=None,
placement_generation=1,
):
return ExecutionContext(
instance_uuid='instance-test',
workspace_uuid=workspace_uuid,
placement_generation=placement_generation,
bot_uuid=bot_uuid,
pipeline_uuid=pipeline_uuid,
)
def aggregation_key(
context,
*,
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
bot_uuid='test-bot',
pipeline_uuid=None,
):
return (
context.instance_uuid,
context.workspace_uuid,
context.placement_generation,
bot_uuid,
pipeline_uuid,
launcher_type.value,
launcher_id,
)
def get_aggregator_module():
"""Lazy import to avoid circular import issues."""
@@ -36,12 +82,66 @@ def make_aggregator_app():
app = FakeApp()
# Ensure query_pool has add_query method
app.query_pool.add_query = AsyncMock()
async def resolve_context(
context,
*,
bot_uuid,
pipeline_uuid,
query_uuid=None,
):
if context is None:
raise ExecutionContextRequiredError('ExecutionContext required in test')
return bind_execution_context(
context,
bot_uuid=bot_uuid,
pipeline_uuid=pipeline_uuid,
query_uuid=query_uuid,
)
app.query_pool.resolve_execution_context = AsyncMock(side_effect=resolve_context)
# Add pipeline_mgr mock
app.pipeline_mgr = AsyncMock()
app.pipeline_mgr.get_pipeline_by_uuid = AsyncMock(return_value=None)
app.workspace_service = Mock()
app.workspace_service.get_execution_binding = AsyncMock(
return_value=Mock(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
)
)
return app
def enable_aggregation(app, *, delay=10.0):
pipeline = Mock()
pipeline.pipeline_entity.config = {
'trigger': {
'message-aggregation': {
'enabled': True,
'delay': delay,
}
}
}
app.pipeline_mgr.get_pipeline_by_uuid = AsyncMock(return_value=pipeline)
def scoped_message_kwargs(context, *, launcher_id=12345, text='hello'):
chain = text_chain(text)
return {
'execution_context': context,
'bot_uuid': context.bot_uuid,
'launcher_type': provider_session.LauncherTypes.PERSON,
'launcher_id': launcher_id,
'sender_id': launcher_id,
'message_event': friend_message_event(chain),
'message_chain': chain,
'adapter': mock_adapter(),
'pipeline_uuid': context.pipeline_uuid,
}
class TestPendingMessage:
"""Tests for PendingMessage dataclass."""
@@ -54,6 +154,7 @@ class TestPendingMessage:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
execution_context=execution_context(pipeline_uuid='test-pipeline'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -77,9 +178,14 @@ class TestSessionBuffer:
"""SessionBuffer should be created with correct fields."""
aggregator = get_aggregator_module()
buffer = aggregator.SessionBuffer(session_id='test-session')
context = execution_context()
key = aggregation_key(context)
buffer = aggregator.SessionBuffer(
aggregation_key=key,
execution_context=context,
)
assert buffer.session_id == 'test-session'
assert buffer.aggregation_key == key
assert buffer.messages == []
assert buffer.timer_task is None
assert buffer.last_message_time is not None
@@ -93,6 +199,7 @@ class TestSessionBuffer:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -103,8 +210,10 @@ class TestSessionBuffer:
pipeline_uuid=None,
)
context = execution_context()
buffer = aggregator.SessionBuffer(
session_id='test-session',
aggregation_key=aggregation_key(context),
execution_context=context,
messages=[pending],
)
@@ -127,7 +236,7 @@ class TestMessageAggregatorInit:
class TestMessageAggregatorSessionId:
"""Tests for session ID generation."""
"""Tests for scoped aggregation key generation."""
def test_session_id_format(self):
"""Session ID should be correctly formatted."""
@@ -136,13 +245,24 @@ class TestMessageAggregatorSessionId:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
session_id = agg._get_session_id(
context = execution_context()
session_id = agg._get_aggregation_key(
context,
bot_uuid='bot-123',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=45678,
pipeline_uuid=None,
)
assert session_id == 'bot-123:person:45678'
assert session_id == (
'instance-test',
'workspace-test',
1,
'bot-123',
None,
'person',
45678,
)
def test_session_id_different_launchers(self):
"""Different launcher types should produce different IDs."""
@@ -151,16 +271,21 @@ class TestMessageAggregatorSessionId:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
person_id = agg._get_session_id(
context = execution_context()
person_id = agg._get_aggregation_key(
context,
bot_uuid='bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=123,
pipeline_uuid=None,
)
group_id = agg._get_session_id(
group_id = agg._get_aggregation_key(
context,
bot_uuid='bot',
launcher_type=provider_session.LauncherTypes.GROUP,
launcher_id=123,
pipeline_uuid=None,
)
assert person_id != group_id
@@ -177,7 +302,7 @@ class TestMessageAggregatorConfig:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config(None)
enabled, delay = await agg._get_aggregation_config(execution_context(), None)
assert enabled == False
assert delay == 1.5
@@ -191,7 +316,10 @@ class TestMessageAggregatorConfig:
app.pipeline_mgr.get_pipeline_by_uuid = AsyncMock(return_value=None)
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config('unknown-pipeline')
enabled, delay = await agg._get_aggregation_config(
execution_context(pipeline_uuid='unknown-pipeline'),
'unknown-pipeline',
)
assert enabled == False
assert delay == 1.5
@@ -217,7 +345,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config('test-pipeline')
enabled, delay = await agg._get_aggregation_config(
execution_context(pipeline_uuid='test-pipeline'),
'test-pipeline',
)
assert enabled == True
assert delay == 2.0
@@ -243,7 +374,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config('test-pipeline')
enabled, delay = await agg._get_aggregation_config(
execution_context(pipeline_uuid='test-pipeline'),
'test-pipeline',
)
assert delay == 1.0 # Clamped to minimum
@@ -268,7 +402,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config('test-pipeline')
enabled, delay = await agg._get_aggregation_config(
execution_context(pipeline_uuid='test-pipeline'),
'test-pipeline',
)
assert delay == 10.0 # Clamped to maximum
@@ -293,7 +430,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
enabled, delay = await agg._get_aggregation_config('test-pipeline')
enabled, delay = await agg._get_aggregation_config(
execution_context(pipeline_uuid='test-pipeline'),
'test-pipeline',
)
assert delay == 1.5 # Default
@@ -314,6 +454,7 @@ class TestMessageAggregatorAddMessage:
adapter = mock_adapter()
await agg.add_message(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -353,6 +494,7 @@ class TestMessageAggregatorAddMessage:
adapter = mock_adapter()
await agg.add_message(
execution_context=execution_context(pipeline_uuid='test-pipeline'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -379,6 +521,7 @@ class TestMessageAggregatorAddMessage:
chain = text_chain('approve')
await agg.add_message(
execution_context=execution_context(pipeline_uuid='test-pipeline'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -424,6 +567,7 @@ class TestMessageAggregatorAddMessage:
# Add messages up to MAX_BUFFER_MESSAGES
for i in range(aggregator.MAX_BUFFER_MESSAGES):
await agg.add_message(
execution_context=execution_context(pipeline_uuid='test-pipeline'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -435,7 +579,14 @@ class TestMessageAggregatorAddMessage:
)
# Buffer should be flushed (empty or no buffer)
session_id = agg._get_session_id('test-bot', provider_session.LauncherTypes.PERSON, 12345)
context = execution_context(pipeline_uuid='test-pipeline')
session_id = agg._get_aggregation_key(
context,
'test-bot',
provider_session.LauncherTypes.PERSON,
12345,
'test-pipeline',
)
assert session_id not in agg.buffers or len(agg.buffers[session_id].messages) == 0
@@ -454,6 +605,7 @@ class TestMessageAggregatorMerge:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -481,6 +633,7 @@ class TestMessageAggregatorMerge:
adapter = mock_adapter()
pending1 = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -492,6 +645,7 @@ class TestMessageAggregatorMerge:
)
pending2 = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -522,6 +676,7 @@ class TestMessageAggregatorMerge:
adapter = mock_adapter()
pending1 = aggregator.PendingMessage(
execution_context=execution_context(pipeline_uuid='test-pipeline-uuid'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -534,6 +689,7 @@ class TestMessageAggregatorMerge:
)
pending2 = aggregator.PendingMessage(
execution_context=execution_context(pipeline_uuid='test-pipeline-uuid'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -562,7 +718,8 @@ class TestMessageAggregatorFlush:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
await agg._flush_buffer('nonexistent-session')
context = execution_context()
await agg._flush_buffer(aggregation_key(context), context)
# Should not call query_pool
assert not app.query_pool.add_query.called
@@ -580,6 +737,7 @@ class TestMessageAggregatorFlush:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -590,17 +748,57 @@ class TestMessageAggregatorFlush:
pipeline_uuid=None,
)
context = execution_context()
key = aggregation_key(context)
buffer = aggregator.SessionBuffer(
session_id='test-session',
aggregation_key=key,
execution_context=context,
messages=[pending],
)
agg.buffers['test-session'] = buffer
agg.buffers[key] = buffer
await agg._flush_buffer('test-session')
await agg._flush_buffer(key, context)
assert app.query_pool.add_query.called
assert 'test-session' not in agg.buffers
assert key not in agg.buffers
@pytest.mark.asyncio
async def test_flush_drops_buffer_when_placement_generation_is_stale(self):
"""A debounce timer cannot enqueue work after its placement is fenced."""
aggregator = get_aggregator_module()
app = make_aggregator_app()
app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError('stale generation')
agg = aggregator.MessageAggregator(app)
context = execution_context(placement_generation=3)
pending = aggregator.PendingMessage(
execution_context=context,
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
sender_id=12345,
message_event=friend_message_event(text_chain('stale')),
message_chain=text_chain('stale'),
adapter=mock_adapter(),
pipeline_uuid=None,
)
key = aggregation_key(context)
agg.buffers[key] = aggregator.SessionBuffer(
aggregation_key=key,
execution_context=context,
messages=[pending],
)
with pytest.raises(WorkspaceGenerationMismatchError):
await agg._flush_buffer(key, context)
app.workspace_service.get_execution_binding.assert_awaited_once_with(
'workspace-test',
expected_generation=3,
)
app.query_pool.add_query.assert_not_awaited()
assert key not in agg.buffers
class TestMessageAggregatorFlushAll:
@@ -633,6 +831,7 @@ class TestMessageAggregatorFlushAll:
# Create two buffers
pending1 = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -644,6 +843,7 @@ class TestMessageAggregatorFlushAll:
)
pending2 = aggregator.PendingMessage(
execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=67890,
@@ -654,14 +854,213 @@ class TestMessageAggregatorFlushAll:
pipeline_uuid=None,
)
buffer1 = aggregator.SessionBuffer(session_id='session-1', messages=[pending1])
buffer2 = aggregator.SessionBuffer(session_id='session-2', messages=[pending2])
context = execution_context()
key1 = aggregation_key(context, launcher_id=12345)
key2 = aggregation_key(context, launcher_id=67890)
buffer1 = aggregator.SessionBuffer(
aggregation_key=key1,
execution_context=context,
messages=[pending1],
)
buffer2 = aggregator.SessionBuffer(
aggregation_key=key2,
execution_context=context,
messages=[pending2],
)
agg.buffers['session-1'] = buffer1
agg.buffers['session-2'] = buffer2
agg.buffers[key1] = buffer1
agg.buffers[key2] = buffer2
await agg.flush_all()
# Both buffers should be flushed
assert len(agg.buffers) == 0
assert app.query_pool.add_query.call_count == 2
class TestMessageAggregatorWorkspaceIsolation:
"""Regression coverage for fail-closed and cross-workspace behavior."""
@pytest.mark.asyncio
async def test_missing_execution_context_fails_closed(self):
app = make_aggregator_app()
agg = get_aggregator_module().MessageAggregator(app)
kwargs = scoped_message_kwargs(execution_context())
kwargs['execution_context'] = None
with pytest.raises(ExecutionContextRequiredError):
await agg.add_message(**kwargs)
app.query_pool.add_query.assert_not_awaited()
@pytest.mark.asyncio
async def test_same_launcher_in_two_workspaces_uses_separate_buffers(self):
app = make_aggregator_app()
enable_aggregation(app)
agg = get_aggregator_module().MessageAggregator(app)
await agg.add_message(**scoped_message_kwargs(execution_context('workspace-a', pipeline_uuid='test-pipeline')))
await agg.add_message(**scoped_message_kwargs(execution_context('workspace-b', pipeline_uuid='test-pipeline')))
assert len(agg.buffers) == 2
assert {key[1] for key in agg.buffers} == {'workspace-a', 'workspace-b'}
await agg.flush_all()
@pytest.mark.asyncio
async def test_new_buffer_uses_scope_counter_without_global_scan(self):
class NoGlobalIterationDict(dict):
def __iter__(self):
raise AssertionError('aggregation admission scanned all buffers')
def items(self):
raise AssertionError('aggregation admission scanned all buffers')
def values(self):
raise AssertionError('aggregation admission scanned all buffers')
app = make_aggregator_app()
enable_aggregation(app)
agg = get_aggregator_module().MessageAggregator(app)
agg.max_buffers = 2_000
agg.max_buffers_per_workspace = 2_000
existing = {
(
'instance-test',
f'workspace-{index}',
1,
'bot',
'pipeline',
'person',
index,
): object()
for index in range(1_000)
}
agg.buffers = NoGlobalIterationDict(existing)
agg._buffer_counts_by_scope = {key[:3]: 1 for key in existing}
context = execution_context(
'workspace-target',
pipeline_uuid='test-pipeline',
)
await agg.add_message(**scoped_message_kwargs(context))
key = aggregation_key(
context,
pipeline_uuid='test-pipeline',
)
assert key in agg.buffers
assert agg._buffer_counts_by_scope[key[:3]] == 1
timer_task = agg.buffers[key].timer_task
assert timer_task is not None
timer_task.cancel()
await asyncio.gather(timer_task, return_exceptions=True)
await agg._flush_buffer(key, context)
assert key[:3] not in agg._buffer_counts_by_scope
@pytest.mark.asyncio
async def test_same_launcher_in_two_bots_uses_separate_buffers(self):
app = make_aggregator_app()
enable_aggregation(app)
agg = get_aggregator_module().MessageAggregator(app)
await agg.add_message(
**scoped_message_kwargs(execution_context(bot_uuid='bot-a', pipeline_uuid='test-pipeline'))
)
await agg.add_message(
**scoped_message_kwargs(execution_context(bot_uuid='bot-b', pipeline_uuid='test-pipeline'))
)
assert len(agg.buffers) == 2
assert {key[3] for key in agg.buffers} == {'bot-a', 'bot-b'}
await agg.flush_all()
@pytest.mark.asyncio
async def test_timer_receives_exact_captured_execution_context(self, monkeypatch):
app = make_aggregator_app()
enable_aggregation(app)
agg = get_aggregator_module().MessageAggregator(app)
request_value = contextvars.ContextVar('aggregator_request_value', default=None)
token = request_value.set('request-scope')
observed = []
async def delayed_flush(*args):
observed.append((request_value.get(), args[2]))
monkeypatch.setattr(agg, '_delayed_flush', delayed_flush)
context = execution_context(pipeline_uuid='test-pipeline')
try:
await agg.add_message(**scoped_message_kwargs(context))
await asyncio.sleep(0)
finally:
request_value.reset(token)
assert observed == [(None, context)]
await agg.flush_all()
@pytest.mark.asyncio
async def test_delayed_flush_opens_explicit_workspace_uow(self, monkeypatch):
app = make_aggregator_app()
app.persistence_mgr.mode = SimpleNamespace(value='cloud_runtime')
scopes = []
@asynccontextmanager
async def tenant_uow(workspace_uuid):
scopes.append(workspace_uuid)
yield
app.persistence_mgr.tenant_uow = tenant_uow
agg = get_aggregator_module().MessageAggregator(app)
flush = AsyncMock()
monkeypatch.setattr(agg, '_flush_buffer', flush)
context = execution_context('workspace-a', pipeline_uuid='test-pipeline')
key = aggregation_key(context, pipeline_uuid='test-pipeline')
await agg._delayed_flush(key, 0, context)
assert scopes == ['workspace-a']
flush.assert_awaited_once_with(key, context)
@pytest.mark.asyncio
async def test_flush_rejects_context_from_another_workspace(self):
app = make_aggregator_app()
enable_aggregation(app)
agg = get_aggregator_module().MessageAggregator(app)
context_a = execution_context('workspace-a', pipeline_uuid='test-pipeline')
context_b = execution_context('workspace-b', pipeline_uuid='test-pipeline')
await agg.add_message(**scoped_message_kwargs(context_a))
key = next(iter(agg.buffers))
with pytest.raises(ExecutionContextMismatchError):
await agg._flush_buffer(key, context_b)
assert key in agg.buffers
await agg.flush_all()
def test_merge_rejects_messages_from_different_workspaces(self):
app = make_aggregator_app()
agg = get_aggregator_module().MessageAggregator(app)
aggregator = get_aggregator_module()
with pytest.raises(ExecutionContextMismatchError):
agg._merge_messages(
[
aggregator.PendingMessage(**scoped_message_kwargs(execution_context('workspace-a'))),
aggregator.PendingMessage(**scoped_message_kwargs(execution_context('workspace-b'))),
]
)
@pytest.mark.asyncio
async def test_flush_all_preserves_each_workspace_context(self):
app = make_aggregator_app()
enable_aggregation(app)
agg = get_aggregator_module().MessageAggregator(app)
await agg.add_message(**scoped_message_kwargs(execution_context('workspace-a', pipeline_uuid='test-pipeline')))
await agg.add_message(**scoped_message_kwargs(execution_context('workspace-b', pipeline_uuid='test-pipeline')))
await agg.flush_all()
forwarded_workspaces = {
call.kwargs['execution_context'].workspace_uuid for call in app.query_pool.add_query.await_args_list
}
assert forwarded_workspaces == {'workspace-a', 'workspace-b'}
@@ -486,3 +486,13 @@ class TestChatHandlerHelper:
handler = chat.ChatMessageHandler(fake_app)
result = handler.cut_str('first line\nsecond line')
assert '...' in result
def test_response_size_limit_uses_instance_config(self, fake_app):
from langbot_plugin.api.entities.builtin.provider.message import Message
fake_app.instance_config.data['system'] = {'response_limits': {'max_generated_chars': 4}}
chat = get_chat_handler()
handler = chat.ChatMessageHandler(fake_app)
with pytest.raises(RuntimeError, match='configured limit'):
handler._check_response_size(Message(role='assistant', content='12345'))
@@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, Mock
import pytest
import yaml
from langbot_plugin.api.entities.builtin.provider import session as provider_session
def _preproc_module():
@@ -43,7 +44,12 @@ def _prompt_preprocessing_context(default_prompt=None, prompt=None):
async def _run_preprocessor(mock_app, sample_query, conversation):
session = SimpleNamespace(launcher_type=sample_query.launcher_type, launcher_id=sample_query.launcher_id)
session = provider_session.Session(
launcher_type=sample_query.launcher_type,
launcher_id=sample_query.launcher_id,
sender_id=sample_query.sender_id,
bot_uuid=sample_query.bot_uuid,
)
mock_app.sess_mgr.get_session = AsyncMock(return_value=session)
mock_app.sess_mgr.get_conversation = AsyncMock(return_value=conversation)
mock_app.plugin_connector.emit_event = AsyncMock(return_value=_prompt_preprocessing_context())
@@ -158,17 +158,21 @@ class TestCommandHandlerReal:
@pytest.mark.asyncio
async def test_admin_privilege_check(self, fake_app, mock_event_ctx, mock_execute_factory):
"""Admin users get privilege level 2."""
"""A per-bot admin from the database is marked as admin in command events."""
from langbot_plugin.api.entities.builtin.provider.session import LauncherTypes
command = get_command_handler()
fake_app.instance_config.data = {'admins': ['person_12345']}
admin_result = Mock()
admin_result.first.return_value = Mock()
fake_app.persistence_mgr.execute_async = AsyncMock(return_value=admin_result)
fake_app.instance_config.data = {}
fake_app.plugin_connector.emit_event = AsyncMock(return_value=mock_event_ctx)
fake_app.cmd_mgr.execute = mock_execute_factory()
handler = command.CommandHandler(fake_app)
query = command_query('status')
query.bot_uuid = 'bot-1'
query.launcher_type = LauncherTypes.PERSON
query.launcher_id = 12345
@@ -176,23 +180,28 @@ class TestCommandHandlerReal:
async for result in handler.handle(query):
results.append(result)
fake_app.persistence_mgr.execute_async.assert_awaited_once()
call_args = fake_app.plugin_connector.emit_event.call_args
event = call_args[0][0]
assert event.is_admin is True
@pytest.mark.asyncio
async def test_non_admin_privilege_check(self, fake_app, mock_event_ctx, mock_execute_factory):
"""Non-admin users get privilege level 1."""
"""A launcher absent from the per-bot admin table is not an admin."""
from langbot_plugin.api.entities.builtin.provider.session import LauncherTypes
command = get_command_handler()
fake_app.instance_config.data = {'admins': ['person_12345']}
admin_result = Mock()
admin_result.first.return_value = None
fake_app.persistence_mgr.execute_async = AsyncMock(return_value=admin_result)
fake_app.instance_config.data = {}
fake_app.plugin_connector.emit_event = AsyncMock(return_value=mock_event_ctx)
fake_app.cmd_mgr.execute = mock_execute_factory()
handler = command.CommandHandler(fake_app)
query = command_query('status')
query.bot_uuid = 'bot-1'
query.launcher_type = LauncherTypes.PERSON
query.launcher_id = 67890
@@ -200,6 +209,7 @@ class TestCommandHandlerReal:
async for result in handler.handle(query):
results.append(result)
fake_app.persistence_mgr.execute_async.assert_awaited_once()
call_args = fake_app.plugin_connector.emit_event.call_args
event = call_args[0][0]
assert event.is_admin is False
+18 -2
View File
@@ -5,6 +5,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.agent.runner.errors import RunnerNotFoundError
from langbot.pkg.pipeline.controller import Controller
@@ -30,6 +31,21 @@ def make_pipeline():
)
def make_query(query_id: int, pipeline_uuid: str):
context = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
pipeline_uuid=pipeline_uuid,
)
return SimpleNamespace(
query_id=query_id,
pipeline_uuid=pipeline_uuid,
variables={},
_execution_context=context,
)
@pytest.mark.asyncio
async def test_try_claim_steering_returns_false_when_runner_lookup_fails():
app = make_app()
@@ -38,7 +54,7 @@ async def test_try_claim_steering_returns_false_when_runner_lookup_fails():
'plugin:missing/runner/default'
)
controller = Controller(app)
query = SimpleNamespace(query_id=1, pipeline_uuid='pipeline-001', variables={})
query = make_query(1, 'pipeline-001')
claimed = await controller._try_claim_steering_before_session_slot(query)
@@ -53,7 +69,7 @@ async def test_try_claim_steering_sets_pipeline_context_before_claiming():
app.pipeline_mgr.get_pipeline_by_uuid.return_value = pipeline
app.agent_run_orchestrator.try_claim_steering_from_query.return_value = True
controller = Controller(app)
query = SimpleNamespace(query_id=2, pipeline_uuid='pipeline-002', variables={})
query = make_query(2, 'pipeline-002')
claimed = await controller._try_claim_steering_before_session_slot(query)
@@ -0,0 +1,145 @@
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock, Mock
import pytest
import sqlalchemy as sa
from sqlalchemy.ext.asyncio import create_async_engine
from langbot.pkg.persistence.mgr import PersistenceManager, PersistenceMode
from langbot.pkg.persistence.tenant_uow import PersistenceScopeKind
from langbot.pkg.pipeline.controller import Controller
from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError
def _prepare_scheduler(mock_app):
query_pool = MagicMock()
query_pool.remove_query = AsyncMock(return_value=True)
query_pool.__aenter__ = AsyncMock(return_value=query_pool)
query_pool.__aexit__ = AsyncMock(return_value=None)
query_pool.condition = SimpleNamespace(notify_all=Mock())
mock_app.query_pool = query_pool
session = SimpleNamespace(_semaphore=SimpleNamespace(release=Mock()))
mock_app.sess_mgr.get_session = AsyncMock(return_value=session)
mock_app.pipeline_mgr = SimpleNamespace(get_pipeline_by_uuid=AsyncMock())
return query_pool, session
@pytest.mark.asyncio
async def test_controller_drops_stale_query_before_pipeline_lookup(
mock_app,
sample_query,
):
query_pool, session = _prepare_scheduler(mock_app)
mock_app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError('stale generation')
controller = Controller(mock_app)
initial_slots = controller.semaphore._value
await controller._process_query(sample_query)
mock_app.workspace_service.get_execution_binding.assert_awaited_once_with(
'test-workspace',
expected_generation=1,
)
mock_app.pipeline_mgr.get_pipeline_by_uuid.assert_not_awaited()
query_pool.remove_query.assert_awaited_once_with(sample_query)
session._semaphore.release.assert_called_once_with()
query_pool.condition.notify_all.assert_called_once_with()
assert controller.semaphore._value == initial_slots
@pytest.mark.asyncio
async def test_cloud_controller_releases_database_connection_during_pipeline_wait(
tmp_path,
mock_app,
sample_query,
):
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "pipeline-short-scope.db"}')
table = sa.Table('pipeline_scope_probe', sa.MetaData(), sa.Column('id', sa.Integer, primary_key=True))
manager = PersistenceManager(object(), mode=PersistenceMode.CLOUD_RUNTIME)
manager.db = SimpleNamespace(get_engine=lambda: engine)
checked_out = 0
def on_checkout(*_args):
nonlocal checked_out
checked_out += 1
def on_checkin(*_args):
nonlocal checked_out
checked_out -= 1
sa.event.listen(engine.sync_engine, 'checkout', on_checkout)
sa.event.listen(engine.sync_engine, 'checkin', on_checkin)
try:
async with engine.begin() as conn:
await conn.run_sync(table.metadata.create_all)
_prepare_scheduler(mock_app)
mock_app.persistence_mgr = manager
pipeline_waiting = asyncio.Event()
release_pipeline = asyncio.Event()
async def get_binding(*_args, **_kwargs):
assert manager.current_scope().kind is PersistenceScopeKind.WORKSPACE
assert manager.current_session() is None
await manager.execute_async(sa.select(table.c.id))
assert manager.current_session() is None
return SimpleNamespace(
instance_uuid='test-instance',
workspace_uuid='test-workspace',
placement_generation=1,
)
async def run_pipeline(_query):
await manager.execute_async(sa.select(table.c.id))
assert manager.current_session() is None
pipeline_waiting.set()
await release_pipeline.wait()
assert manager.current_scope().kind is PersistenceScopeKind.WORKSPACE
assert manager.current_session() is None
runtime_pipeline = SimpleNamespace(run=AsyncMock(side_effect=run_pipeline))
async def get_pipeline(*_args, **_kwargs):
await manager.execute_async(sa.select(table.c.id))
assert manager.current_session() is None
return runtime_pipeline
mock_app.workspace_service.get_execution_binding = AsyncMock(side_effect=get_binding)
mock_app.pipeline_mgr.get_pipeline_by_uuid = AsyncMock(side_effect=get_pipeline)
controller = Controller(mock_app)
task = asyncio.create_task(controller._process_query(sample_query))
await asyncio.wait_for(pipeline_waiting.wait(), timeout=2)
assert checked_out == 0
assert not task.done()
release_pipeline.set()
await asyncio.wait_for(task, timeout=2)
assert checked_out == 0
runtime_pipeline.run.assert_awaited_once_with(sample_query)
finally:
await engine.dispose()
@pytest.mark.asyncio
async def test_controller_revalidates_generation_before_running_pipeline(
mock_app,
sample_query,
):
query_pool, session = _prepare_scheduler(mock_app)
runtime_pipeline = SimpleNamespace(run=AsyncMock())
mock_app.pipeline_mgr.get_pipeline_by_uuid.return_value = runtime_pipeline
controller = Controller(mock_app)
await controller._process_query(sample_query)
mock_app.workspace_service.get_execution_binding.assert_awaited_once_with(
'test-workspace',
expected_generation=1,
)
runtime_pipeline.run.assert_awaited_once_with(sample_query)
query_pool.remove_query.assert_awaited_once_with(sample_query)
session._semaphore.release.assert_called_once_with()
@@ -0,0 +1,50 @@
from types import SimpleNamespace
from unittest.mock import Mock
import pytest
import langbot_plugin.api.entities.builtin.platform.message as platform_message
from langbot.pkg.pipeline.longtext.strategies.image import Text2ImageStrategy
from langbot.pkg.pipeline.longtext.strategies import image
class _WideFont:
def getlength(self, text: str) -> int:
return len(text) * 100
def test_image_strategy_line_split_always_consumes_input():
strategy = Text2ImageStrategy(Mock())
lines = strategy._split_text_lines('abc', 1, _WideFont())
assert lines == ['a', 'b', 'c']
assert ''.join(lines) == 'abc'
def test_image_strategy_numeric_boundaries_are_found_in_linear_order():
strategy = Text2ImageStrategy(Mock())
assert strategy.indexNumber('a12-b12-c345') == [['12', 1], ['12', 5], ['345', 9]]
def test_image_strategy_rejects_unbounded_line_count_before_allocating_canvas(monkeypatch):
strategy = Text2ImageStrategy(Mock())
monkeypatch.setattr(image, '_MAX_TEXT_TO_IMAGE_LINES', 2)
with pytest.raises(ValueError, match='2 lines'):
strategy._split_text_lines('one\ntwo\nthree', 1000, _WideFont())
@pytest.mark.asyncio
async def test_image_strategy_falls_back_to_forward_for_oversized_text(monkeypatch):
app = Mock()
strategy = Text2ImageStrategy(app)
monkeypatch.setattr(image, '_MAX_TEXT_TO_IMAGE_CHARS', 4)
query = SimpleNamespace(adapter=SimpleNamespace(bot_account_id='bot'))
components = await strategy.process('12345', query)
assert len(components) == 1
assert isinstance(components[0], platform_message.Forward)
app.logger.warning.assert_called_once()
@@ -5,13 +5,19 @@ import pytest
from langbot.pkg.api.http.service.pipeline import PipelineService
WORKSPACE_UUID = 'workspace-a'
@pytest.mark.asyncio
async def test_update_pipeline_filters_protected_fields_without_mutating_input(mock_app):
service = PipelineService(mock_app)
loaded_pipeline = Mock()
service.get_pipeline = AsyncMock(return_value=loaded_pipeline)
mock_app.persistence_mgr.execute_async = AsyncMock(return_value=None)
bot = Mock(uuid='bot-uuid')
bot_result = Mock(all=Mock(return_value=[bot]))
mock_app.persistence_mgr.execute_async = AsyncMock(side_effect=[None, bot_result])
mock_app.bot_service = Mock(update_bot=AsyncMock())
mock_app.pipeline_mgr = Mock(remove_pipeline=AsyncMock(), load_pipeline=AsyncMock())
mock_app.sess_mgr.session_list = []
@@ -24,7 +30,7 @@ async def test_update_pipeline_filters_protected_fields_without_mutating_input(m
}
original_pipeline_data = pipeline_data.copy()
await service.update_pipeline('pipeline-uuid', pipeline_data)
await service.update_pipeline(WORKSPACE_UUID, 'pipeline-uuid', pipeline_data)
assert pipeline_data == original_pipeline_data
@@ -32,5 +38,6 @@ async def test_update_pipeline_filters_protected_fields_without_mutating_input(m
updated_fields = {getattr(field, 'key', str(field)) for field in update_stmt._values}
assert updated_fields == {'name'}
mock_app.pipeline_mgr.remove_pipeline.assert_awaited_once_with('pipeline-uuid')
mock_app.pipeline_mgr.load_pipeline.assert_awaited_once_with(loaded_pipeline)
mock_app.bot_service.update_bot.assert_not_awaited()
mock_app.pipeline_mgr.remove_pipeline.assert_awaited_once_with('workspace-a', 'pipeline-uuid')
mock_app.pipeline_mgr.load_pipeline.assert_awaited_once_with('workspace-a', loaded_pipeline)
+193 -183
View File
@@ -3,9 +3,26 @@ PipelineManager unit tests
"""
import pytest
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
from importlib import import_module
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.workspace.entities import WorkspaceExecutionBinding
from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError, WorkspaceInvariantError
RUNNER_ID = 'plugin:langbot-team/LocalAgent/default'
def _context(pipeline_uuid: str = 'test-uuid') -> ExecutionContext:
return ExecutionContext(
instance_uuid='test-instance',
workspace_uuid='test-workspace',
placement_generation=1,
pipeline_uuid=pipeline_uuid,
)
def get_pipelinemgr_module():
return import_module('langbot.pkg.pipeline.pipelinemgr')
@@ -37,6 +54,95 @@ async def test_pipeline_manager_initialize(mock_app):
assert len(manager.pipelines) == 0
@pytest.mark.asyncio
async def test_cloud_startup_reuses_validated_pipeline_binding(mock_app):
class TenantUow:
async def __aenter__(self):
return self
async def __aexit__(self, *_args):
return False
binding = WorkspaceExecutionBinding(
instance_uuid='test-instance',
workspace_uuid='test-workspace',
placement_generation=1,
write_fenced=False,
state='active',
)
pipeline_entity = Mock(
uuid='test-uuid',
workspace_uuid='test-workspace',
stages=[],
config={},
extensions_preferences={},
)
mock_app.persistence_mgr.mode = SimpleNamespace(value='cloud_runtime')
mock_app.persistence_mgr.tenant_uow = lambda _workspace_uuid: TenantUow()
mock_app.persistence_mgr.execute_async = AsyncMock(return_value=Mock(all=Mock(return_value=[pipeline_entity])))
mock_app.workspace_service.list_active_execution_bindings = AsyncMock(return_value=[binding])
mock_app.workspace_service.get_execution_binding = AsyncMock(
side_effect=AssertionError('startup pipeline loader repeated a validated binding lookup')
)
manager = get_pipelinemgr_module().PipelineManager(mock_app)
manager.stage_dict = {}
await manager.load_pipelines_from_db()
assert len(manager.pipelines) == 1
mock_app.workspace_service.get_execution_binding.assert_not_awaited()
def test_generation_advance_prunes_superseded_workspace_pipelines(mock_app):
class NoGlobalIterationDict(dict):
def __iter__(self):
raise AssertionError('generation advance scanned every pipeline')
def items(self):
raise AssertionError('generation advance scanned every pipeline')
def values(self):
raise AssertionError('generation advance scanned every pipeline')
pipelinemgr = get_pipelinemgr_module()
manager = pipelinemgr.PipelineManager(mock_app)
old_context = _context()
next_context = ExecutionContext(
instance_uuid=old_context.instance_uuid,
workspace_uuid=old_context.workspace_uuid,
placement_generation=2,
pipeline_uuid=old_context.pipeline_uuid,
)
old_pipeline = SimpleNamespace(
execution_context=old_context,
workspace_uuid=old_context.workspace_uuid,
placement_generation=old_context.placement_generation,
)
other_pipelines = [
SimpleNamespace(
execution_context=ExecutionContext(
instance_uuid='test-instance',
workspace_uuid=f'workspace-{index}',
placement_generation=1,
pipeline_uuid=f'pipeline-{index}',
),
workspace_uuid=f'workspace-{index}',
placement_generation=1,
)
for index in range(1_000)
]
manager.pipelines = [old_pipeline, *other_pipelines]
manager._observe_execution_context(old_context)
manager._pipelines_by_key = NoGlobalIterationDict(manager._pipelines_by_key)
manager._observe_execution_context(next_context)
manager._pipelines_by_key = dict(manager._pipelines_by_key)
assert manager.pipelines == other_pipelines
with pytest.raises(WorkspaceInvariantError, match='rolled back'):
manager._observe_execution_context(old_context)
@pytest.mark.asyncio
async def test_load_pipeline(mock_app):
"""Test loading a single pipeline"""
@@ -51,11 +157,12 @@ async def test_load_pipeline(mock_app):
# Create test pipeline entity
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.uuid = 'test-uuid'
pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.stages = []
pipeline_entity.config = {'test': 'config'}
pipeline_entity.extensions_preferences = {'plugins': []}
await manager.load_pipeline(pipeline_entity)
await manager.load_pipeline(_context(), pipeline_entity)
assert len(manager.pipelines) == 1
assert manager.pipelines[0].pipeline_entity.uuid == 'test-uuid'
@@ -75,19 +182,20 @@ async def test_get_pipeline_by_uuid(mock_app):
# Create and add test pipeline
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.uuid = 'test-uuid'
pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.stages = []
pipeline_entity.config = {}
pipeline_entity.extensions_preferences = {'plugins': []}
await manager.load_pipeline(pipeline_entity)
await manager.load_pipeline(_context(), pipeline_entity)
# Test retrieval
result = await manager.get_pipeline_by_uuid('test-uuid')
result = await manager.get_pipeline_by_uuid(_context(), 'test-uuid')
assert result is not None
assert result.pipeline_entity.uuid == 'test-uuid'
# Test non-existent UUID
result = await manager.get_pipeline_by_uuid('non-existent')
result = await manager.get_pipeline_by_uuid(_context('non-existent'), 'non-existent')
assert result is None
@@ -105,15 +213,16 @@ async def test_remove_pipeline(mock_app):
# Create and add test pipeline
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.uuid = 'test-uuid'
pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.stages = []
pipeline_entity.config = {}
pipeline_entity.extensions_preferences = {'plugins': []}
await manager.load_pipeline(pipeline_entity)
await manager.load_pipeline(_context(), pipeline_entity)
assert len(manager.pipelines) == 1
# Remove pipeline
await manager.remove_pipeline('test-uuid')
await manager.remove_pipeline(_context(), 'test-uuid')
assert len(manager.pipelines) == 0
@@ -143,128 +252,119 @@ async def test_runtime_pipeline_execute(mock_app, sample_query):
# Create pipeline entity
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.uuid = 'test-pipeline-uuid'
pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.config = sample_query.pipeline_config
pipeline_entity.extensions_preferences = {'plugins': []}
# Create runtime pipeline
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [stage_container])
runtime_pipeline = pipelinemgr.RuntimePipeline(
mock_app,
pipeline_entity,
[stage_container],
_context('test-pipeline-uuid'),
)
# Mock plugin connector
event_ctx = Mock()
event_ctx.is_prevented_default = Mock(return_value=False)
mock_app.plugin_connector.emit_event = AsyncMock(return_value=event_ctx)
# Add query to cached_queries to prevent KeyError in finally block
mock_app.query_pool.cached_queries[sample_query.query_id] = sample_query
# Execute pipeline
await runtime_pipeline.run(sample_query)
# Verify stage was called
mock_stage.process.assert_called_once()
mock_app.query_pool.remove_query.assert_awaited_once_with(sample_query)
@pytest.mark.asyncio
async def test_runtime_pipeline_delivers_latest_chunk_as_final(mock_app, sample_query):
"""The terminal chunk, not the first chunk, controls final stream delivery."""
async def test_runtime_pipeline_rejects_stale_generation_before_side_effects(
mock_app,
sample_query,
):
pipelinemgr = get_pipelinemgr_module()
persistence_pipeline = get_persistence_pipeline_module()
entities = get_entities_module()
provider_message = import_module('langbot_plugin.api.entities.builtin.provider.message')
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.uuid = 'test-pipeline-uuid'
pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.config = sample_query.pipeline_config
pipeline_entity.extensions_preferences = {'plugins': []}
runtime_pipeline = pipelinemgr.RuntimePipeline(
mock_app,
pipeline_entity,
[],
_context('test-pipeline-uuid'),
)
mock_app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError('stale generation')
with pytest.raises(WorkspaceGenerationMismatchError):
await runtime_pipeline.run(sample_query)
mock_app.plugin_connector.emit_event.assert_not_awaited()
sample_query.adapter.reply_message.assert_not_awaited()
sample_query.adapter.reply_message_chunk.assert_not_awaited()
@pytest.mark.asyncio
async def test_runtime_pipeline_revalidates_after_awaited_stage(
mock_app,
sample_query,
):
pipelinemgr = get_pipelinemgr_module()
stage = get_stage_module()
persistence_pipeline = get_persistence_pipeline_module()
entities = get_entities_module()
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.uuid = 'test-pipeline-uuid'
pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.config = sample_query.pipeline_config
pipeline_entity.extensions_preferences = {'plugins': []}
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [])
first_chunk = provider_message.MessageChunk(role='assistant', content='Starting', is_final=False)
final_chunk = provider_message.MessageChunk(role='assistant', content='Done', is_final=True)
sample_query.resp_messages = [first_chunk, final_chunk]
sample_query.adapter.is_stream_output_supported = AsyncMock(return_value=True)
result = entities.StageProcessResult(
result_type=entities.ResultType.CONTINUE,
new_query=sample_query,
user_notice='StartingDone',
user_notice='must not be sent',
console_notice='',
debug_notice='',
error_notice='',
)
await runtime_pipeline._check_output(sample_query, result)
async def stage_process(*_args):
mock_app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError(
'generation changed during stage'
)
return result
sample_query.adapter.reply_message_chunk.assert_awaited_once()
call = sample_query.adapter.reply_message_chunk.await_args.kwargs
assert call['bot_message'] is final_chunk
assert call['is_final'] is True
mock_stage = Mock(spec=stage.PipelineStage)
mock_stage.process = Mock(side_effect=stage_process)
runtime_pipeline = pipelinemgr.RuntimePipeline(
mock_app,
pipeline_entity,
[pipelinemgr.StageInstContainer(inst_name='TestStage', inst=mock_stage)],
_context('test-pipeline-uuid'),
)
with pytest.raises(WorkspaceGenerationMismatchError):
await runtime_pipeline._execute_from_stage(0, sample_query)
@pytest.mark.asyncio
async def test_response_back_stage_delivers_latest_chunk_as_final(mock_app, sample_query):
respback = import_module('langbot.pkg.pipeline.respback.respback')
provider_message = import_module('langbot_plugin.api.entities.builtin.provider.message')
platform_message = import_module('langbot_plugin.api.entities.builtin.platform.message')
first_chunk = provider_message.MessageChunk(role='assistant', content='Starting', is_final=False)
final_chunk = provider_message.MessageChunk(role='assistant', content='Done', is_final=True)
sample_query.resp_messages = [first_chunk, final_chunk]
sample_query.resp_message_chain = [platform_message.MessageChain([platform_message.Plain(text='StartingDone')])]
sample_query.pipeline_config['output']['force-delay'] = {'min': 0, 'max': 0}
sample_query.adapter.is_stream_output_supported = AsyncMock(return_value=True)
await respback.SendResponseBackStage(mock_app).process(sample_query, 'response-back')
sample_query.adapter.reply_message_chunk.assert_awaited_once()
call = sample_query.adapter.reply_message_chunk.await_args.kwargs
assert call['bot_message'] is final_chunk
assert call['is_final'] is True
@pytest.mark.asyncio
async def test_response_back_stage_keeps_consuming_after_stream_delivery_failure(mock_app, sample_query):
respback = import_module('langbot.pkg.pipeline.respback.respback')
provider_message = import_module('langbot_plugin.api.entities.builtin.provider.message')
platform_message = import_module('langbot_plugin.api.entities.builtin.platform.message')
chunk = provider_message.MessageChunk(role='assistant', content='Progress', is_final=False)
sample_query.resp_messages = [chunk]
sample_query.resp_message_chain = [platform_message.MessageChain([platform_message.Plain(text='Progress')])]
sample_query.pipeline_config['output']['force-delay'] = {'min': 0, 'max': 0}
sample_query.adapter.is_stream_output_supported = AsyncMock(return_value=True)
sample_query.adapter.reply_message_chunk.side_effect = RuntimeError('stream update failed')
result = await respback.SendResponseBackStage(mock_app).process(sample_query, 'response-back')
assert result.result_type.name == 'CONTINUE'
sample_query.adapter.reply_message.assert_not_awaited()
sample_query.adapter.reply_message_chunk.assert_not_awaited()
@pytest.mark.asyncio
async def test_response_back_stage_falls_back_to_plain_message_for_failed_final_chunk(mock_app, sample_query):
respback = import_module('langbot.pkg.pipeline.respback.respback')
provider_message = import_module('langbot_plugin.api.entities.builtin.provider.message')
platform_message = import_module('langbot_plugin.api.entities.builtin.platform.message')
chunk = provider_message.MessageChunk(role='assistant', content='Final answer', is_final=True)
sample_query.resp_messages = [chunk]
sample_query.resp_message_chain = [platform_message.MessageChain([platform_message.Plain(text='Final answer')])]
sample_query.pipeline_config['output']['force-delay'] = {'min': 0, 'max': 0}
sample_query.adapter.is_stream_output_supported = AsyncMock(return_value=True)
sample_query.adapter.reply_message_chunk.side_effect = RuntimeError('stream update failed')
result = await respback.SendResponseBackStage(mock_app).process(sample_query, 'response-back')
assert result.result_type.name == 'CONTINUE'
sample_query.adapter.reply_message.assert_awaited_once()
def test_runtime_pipeline_prefers_runner_mcp_resources(mock_app):
"""Runner resource selection should override extension preferences."""
def test_runtime_pipeline_prefers_local_agent_mcp_resources(mock_app):
"""AgentRunner resource selection should override legacy extension prefs."""
pipelinemgr = get_pipelinemgr_module()
persistence_pipeline = get_persistence_pipeline_module()
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.uuid = 'test-uuid'
pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.config = {
'ai': {
'runner': {'id': 'plugin:langbot-team/LocalAgent/default'},
'runner': {'id': RUNNER_ID},
'runner_config': {
'plugin:langbot-team/LocalAgent/default': {
RUNNER_ID: {
'mcp-resources': [{'server_uuid': 'srv-new', 'uri': 'file:///new.md'}],
'mcp-resource-agent-read-enabled': False,
},
@@ -276,22 +376,24 @@ def test_runtime_pipeline_prefers_runner_mcp_resources(mock_app):
'mcp_resource_agent_read_enabled': True,
}
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [])
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [], _context())
assert runtime_pipeline.mcp_resource_attachments == [{'server_uuid': 'srv-new', 'uri': 'file:///new.md'}]
assert runtime_pipeline.mcp_resource_agent_read_enabled is False
def test_runtime_pipeline_falls_back_to_extension_mcp_resources(mock_app):
"""Extension preferences apply when the current runner has no override."""
"""Existing extension prefs remain compatible until a runner value exists."""
pipelinemgr = get_pipelinemgr_module()
persistence_pipeline = get_persistence_pipeline_module()
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.uuid = 'test-uuid'
pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.config = {
'ai': {
'runner': {'id': 'plugin:langbot-team/LocalAgent/default'},
'runner_config': {'plugin:langbot-team/LocalAgent/default': {}},
'runner': {'id': RUNNER_ID},
'runner_config': {RUNNER_ID: {}},
}
}
pipeline_entity.extensions_preferences = {
@@ -299,99 +401,7 @@ def test_runtime_pipeline_falls_back_to_extension_mcp_resources(mock_app):
'mcp_resource_agent_read_enabled': False,
}
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [])
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [], _context())
assert runtime_pipeline.mcp_resource_attachments == [{'server_uuid': 'srv-old', 'uri': 'file:///old.md'}]
assert runtime_pipeline.mcp_resource_agent_read_enabled is False
@pytest.mark.parametrize('invalid_value', [0, None, 'false', [], {}])
def test_runtime_pipeline_mcp_resource_read_flag_fails_closed(mock_app, invalid_value):
pipelinemgr = get_pipelinemgr_module()
persistence_pipeline = get_persistence_pipeline_module()
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.config = {
'ai': {
'runner': {'id': 'plugin:test/runner/default'},
'runner_config': {
'plugin:test/runner/default': {
'mcp-resource-agent-read-enabled': invalid_value,
}
},
}
}
pipeline_entity.extensions_preferences = {'mcp_resource_agent_read_enabled': True}
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [])
assert runtime_pipeline.mcp_resource_agent_read_enabled is False
@pytest.mark.parametrize('invalid_value', [0, None, 'false', [], {}])
def test_runtime_pipeline_extension_enable_all_flags_fail_closed(mock_app, invalid_value):
pipelinemgr = get_pipelinemgr_module()
persistence_pipeline = get_persistence_pipeline_module()
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.config = {}
pipeline_entity.extensions_preferences = {
'enable_all_plugins': invalid_value,
'plugins': [{'author': 'allowed', 'name': 'plugin'}],
'enable_all_mcp_servers': invalid_value,
'mcp_servers': ['bound-mcp'],
}
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [])
assert runtime_pipeline.enable_all_plugins is False
assert runtime_pipeline.bound_plugins == ['allowed/plugin']
assert runtime_pipeline.enable_all_mcp_servers is False
assert runtime_pipeline.bound_mcp_servers == ['bound-mcp']
@pytest.mark.parametrize('invalid_preferences', [None, [], '', 0, False])
def test_runtime_pipeline_malformed_extension_root_disables_all_extensions(
mock_app,
invalid_preferences,
):
pipelinemgr = get_pipelinemgr_module()
persistence_pipeline = get_persistence_pipeline_module()
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.config = {}
pipeline_entity.extensions_preferences = invalid_preferences
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [])
assert runtime_pipeline.enable_all_plugins is False
assert runtime_pipeline.bound_plugins == []
assert runtime_pipeline.enable_all_mcp_servers is False
assert runtime_pipeline.bound_mcp_servers == []
assert runtime_pipeline.mcp_resource_attachments == []
assert runtime_pipeline.mcp_resource_agent_read_enabled is False
def test_runtime_pipeline_malformed_extension_lists_are_empty_allowlists(mock_app):
pipelinemgr = get_pipelinemgr_module()
persistence_pipeline = get_persistence_pipeline_module()
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.config = {}
pipeline_entity.extensions_preferences = {
'enable_all_plugins': True,
'plugins': 'allowed/plugin',
'enable_all_mcp_servers': True,
'mcp_servers': 'bound-mcp',
'mcp_resources': 'file:///README.md',
'mcp_resource_agent_read_enabled': True,
}
runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [])
assert runtime_pipeline.enable_all_plugins is False
assert runtime_pipeline.bound_plugins == []
assert runtime_pipeline.enable_all_mcp_servers is False
assert runtime_pipeline.bound_mcp_servers == []
assert runtime_pipeline.mcp_resource_attachments == []
assert runtime_pipeline.mcp_resource_agent_read_enabled is False
+213 -15
View File
@@ -6,10 +6,52 @@ Tests query management, ID generation, and async context handling.
from __future__ import annotations
import uuid
from types import SimpleNamespace
import pytest
from unittest.mock import Mock, patch
from langbot.pkg.pipeline.pool import QueryPool
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.pipeline.pool import (
ExecutionContextMismatchError,
ExecutionContextRequiredError,
QueryNotFoundError,
QueryPool,
QueryPoolCapacityError,
get_query_execution_context,
)
TEST_CONTEXT = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
)
def oss_pool():
"""Build the explicit singleton resolver used by the OSS compatibility path."""
return QueryPool(singleton_context_resolver=lambda: TEST_CONTEXT)
async def add_scoped_mock_query(pool, context, *, bot_uuid='bot-a'):
"""Create a Query through the real pool while keeping SDK details mocked."""
query = Mock()
query.bot_uuid = bot_uuid
query.pipeline_uuid = None
query.query_id = pool.query_id_counter
with patch('langbot.pkg.pipeline.pool.pipeline_query.Query', return_value=query):
return await pool.add_query(
bot_uuid=bot_uuid,
launcher_type=Mock(),
launcher_id='launcher-1',
sender_id='sender-1',
message_event=Mock(),
message_chain=Mock(),
adapter=Mock(),
execution_context=context,
)
pytestmark = pytest.mark.asyncio
@@ -39,7 +81,7 @@ class TestQueryPoolAddQuery:
async def test_add_query_adds_query_with_id(self):
"""add_query creates, stores, and caches a Query with the correct ID."""
pool = QueryPool()
pool = oss_pool()
# Mock Query creation
mock_query = Mock()
@@ -62,12 +104,12 @@ class TestQueryPoolAddQuery:
# Query is added to list and cache
assert pool.queries[0] is mock_query
assert pool.cached_queries[0] is mock_query
assert pool.cached_queries[('workspace-test', mock_query.query_uuid)] is mock_query
assert mock_query.query_id == 0
async def test_add_query_increments_counter(self):
"""Each add_query increments the counter."""
pool = QueryPool()
pool = oss_pool()
mock_query1 = Mock()
mock_query1.query_id = 0
@@ -103,7 +145,7 @@ class TestQueryPoolAddQuery:
async def test_add_query_appends_to_list(self):
"""Query is appended to queries list."""
pool = QueryPool()
pool = oss_pool()
mock_query = Mock()
mock_query.query_id = 0
@@ -126,7 +168,7 @@ class TestQueryPoolAddQuery:
async def test_add_query_caches_query(self):
"""Query is cached by query_id."""
pool = QueryPool()
pool = oss_pool()
mock_query = Mock()
mock_query.query_id = 0
@@ -144,12 +186,13 @@ class TestQueryPoolAddQuery:
adapter=Mock(),
)
assert 0 in pool.cached_queries
assert pool.cached_queries[0] is mock_query
cache_key = ('workspace-test', mock_query.query_uuid)
assert cache_key in pool.cached_queries
assert pool.cached_queries[cache_key] is mock_query
async def test_add_query_with_pipeline_uuid(self):
"""Query can have pipeline_uuid set."""
pool = QueryPool()
pool = oss_pool()
mock_query = Mock()
mock_query.query_id = 0
@@ -175,7 +218,7 @@ class TestQueryPoolAddQuery:
async def test_add_query_sets_routed_by_rule_variable(self):
"""Query has _routed_by_rule variable."""
pool = QueryPool()
pool = oss_pool()
mock_query = Mock()
mock_query.query_id = 0
@@ -201,7 +244,7 @@ class TestQueryPoolAddQuery:
async def test_add_query_merges_control_variables(self):
"""Caller-provided control variables are preserved with routing metadata."""
pool = QueryPool()
pool = oss_pool()
mock_query = Mock(query_id=0)
with patch('langbot.pkg.pipeline.pool.pipeline_query.Query') as MockQuery:
@@ -224,7 +267,7 @@ class TestQueryPoolAddQuery:
async def test_add_query_notifier_condition(self):
"""add_query notifies waiting consumers."""
pool = QueryPool()
pool = oss_pool()
mock_query = Mock()
mock_query.query_id = 0
@@ -260,7 +303,7 @@ class TestQueryPoolContext:
async def test_aenter_acquires_lock(self):
"""__aenter__ acquires the pool lock."""
pool = QueryPool()
pool = oss_pool()
async with pool as p:
# Lock is acquired
@@ -283,7 +326,7 @@ class TestQueryPoolEdgeCases:
async def test_multiple_queries_cached_correctly(self):
"""Multiple queries are cached separately."""
pool = QueryPool()
pool = oss_pool()
mock_queries = []
for i in range(5):
@@ -310,4 +353,159 @@ class TestQueryPoolEdgeCases:
# Each query is cached by its ID
for i in range(5):
assert pool.cached_queries[i] is mock_queries[i]
query = mock_queries[i]
assert pool.cached_queries[('workspace-test', query.query_uuid)] is query
class TestQueryPoolWorkspaceIsolation:
"""Regression coverage for trusted scope and scoped cache indexes."""
async def test_add_query_requires_execution_context_by_default(self):
with pytest.raises(ExecutionContextRequiredError):
await QueryPool().add_query(
bot_uuid='bot-a',
launcher_type=Mock(),
launcher_id='launcher-1',
sender_id='sender-1',
message_event=Mock(),
message_chain=Mock(),
adapter=Mock(),
)
async def test_serialized_scope_fields_are_not_trusted_context(self):
forged_query = SimpleNamespace(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
bot_uuid='bot-a',
pipeline_uuid=None,
query_uuid='forged-query',
)
with pytest.raises(ExecutionContextRequiredError):
get_query_execution_context(forged_query)
async def test_query_lookup_is_workspace_scoped(self):
pool = QueryPool()
query = await add_scoped_mock_query(pool, TEST_CONTEXT)
uuid.UUID(query.query_uuid)
assert await pool.get_query('workspace-test', query.query_uuid) is query
assert await pool.get_query('workspace-other', query.query_uuid) is None
assert await pool.get_query_by_legacy_id('workspace-test', 0) is query
assert await pool.get_query_by_legacy_id('workspace-other', 0) is None
with pytest.raises(QueryNotFoundError):
await pool.require_query('workspace-other', query.query_uuid)
async def test_cache_separates_same_opaque_id_between_workspaces(self, monkeypatch):
fixed_uuid = uuid.UUID('11111111-1111-4111-8111-111111111111')
monkeypatch.setattr('langbot.pkg.pipeline.pool.uuid.uuid4', lambda: fixed_uuid)
pool = QueryPool()
context_a = TEST_CONTEXT
context_b = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-other',
placement_generation=1,
)
query_a = await add_scoped_mock_query(pool, context_a)
query_b = await add_scoped_mock_query(pool, context_b)
assert query_a.query_uuid == query_b.query_uuid
assert await pool.get_query('workspace-test', query_a.query_uuid) is query_a
assert await pool.get_query('workspace-other', query_b.query_uuid) is query_b
async def test_remove_query_cleans_both_scoped_indexes(self):
pool = QueryPool()
query = await add_scoped_mock_query(pool, TEST_CONTEXT)
assert await pool.remove_query(query) is True
assert await pool.get_query('workspace-test', query.query_uuid) is None
assert await pool.get_query_by_legacy_id('workspace-test', query.query_id) is None
assert await pool.remove_query(query) is False
async def test_context_cannot_substitute_bot_identity(self):
context = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
bot_uuid='bot-b',
)
with pytest.raises(ExecutionContextMismatchError):
await add_scoped_mock_query(QueryPool(), context, bot_uuid='bot-a')
async def test_query_counter_is_scoped_by_workspace_and_generation(self):
pool = QueryPool()
workspace_a = TEST_CONTEXT
workspace_b = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-other',
placement_generation=1,
)
next_generation = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=2,
)
await add_scoped_mock_query(pool, workspace_a)
await add_scoped_mock_query(pool, workspace_a)
await add_scoped_mock_query(pool, workspace_b)
assert pool.get_query_count(workspace_a) == 2
assert pool.get_query_count(workspace_b) == 1
assert pool.get_query_count(next_generation) == 0
assert pool.query_id_counter == 3
async def test_workspace_capacity_discards_oldest_queued_query(self):
pool = QueryPool(max_queries=3, max_queries_per_workspace=2)
first = await add_scoped_mock_query(pool, TEST_CONTEXT)
second = await add_scoped_mock_query(pool, TEST_CONTEXT)
third = await add_scoped_mock_query(pool, TEST_CONTEXT)
assert await pool.get_query(TEST_CONTEXT.workspace_uuid, first.query_uuid) is None
assert await pool.get_query(TEST_CONTEXT.workspace_uuid, second.query_uuid) is second
assert await pool.get_query(TEST_CONTEXT.workspace_uuid, third.query_uuid) is third
assert pool.active_query_count_by_workspace == {TEST_CONTEXT.workspace_uuid: 2}
assert pool.get_dropped_query_count(TEST_CONTEXT) == 1
async def test_capacity_rejects_when_every_query_is_already_running(self):
pool = QueryPool(max_queries=1, max_queries_per_workspace=1)
running = await add_scoped_mock_query(pool, TEST_CONTEXT)
async with pool:
pool.mark_query_running_locked(running)
with pytest.raises(QueryPoolCapacityError):
await add_scoped_mock_query(pool, TEST_CONTEXT)
assert pool.active_query_count_by_workspace == {TEST_CONTEXT.workspace_uuid: 1}
async def test_mark_query_running_keeps_active_indexes_but_removes_queue_entry(self):
pool = QueryPool(max_queries=1, max_queries_per_workspace=1)
running = await add_scoped_mock_query(pool, TEST_CONTEXT)
async with pool:
pool.mark_query_running_locked(running)
assert running not in pool.queries
assert await pool.get_query(TEST_CONTEXT.workspace_uuid, running.query_uuid) is running
assert pool.active_query_count_by_workspace == {TEST_CONTEXT.workspace_uuid: 1}
async def test_historical_workspace_counters_are_bounded(self):
pool = QueryPool(max_queries=2, max_queries_per_workspace=1)
contexts = [
ExecutionContext(
instance_uuid='instance-test',
workspace_uuid=f'workspace-{index}',
placement_generation=1,
)
for index in range(3)
]
for context in contexts:
query = await add_scoped_mock_query(pool, context)
await pool.remove_query(query)
assert len(pool.query_count_by_scope) == 2
assert (contexts[0].instance_uuid, contexts[0].workspace_uuid, 1) not in pool.query_count_by_scope
+103 -218
View File
@@ -15,6 +15,9 @@ import pytest
from unittest.mock import AsyncMock, Mock
from importlib import import_module
from langbot_plugin.api.entities.builtin.provider import session as provider_session
from langbot.pkg.agent.runner.descriptor import AgentRunnerDescriptor
from tests.factories import (
FakeApp,
text_query,
@@ -24,6 +27,28 @@ from tests.factories import (
)
RUNNER_ID = 'plugin:langbot-team/LocalAgent/default'
def attach_agent_runner_descriptor(app):
descriptor = 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': []},
],
capabilities={'tool_calling': True, 'multimodal_input': True},
)
app.agent_runner_registry = Mock()
app.agent_runner_registry.get = AsyncMock(return_value=descriptor)
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(return_value=[])
def get_preproc_module():
"""Lazy import to avoid circular import issues."""
return import_module('langbot.pkg.pipeline.preproc.preproc')
@@ -34,48 +59,18 @@ def get_entities_module():
return import_module('langbot.pkg.pipeline.entities')
RUNNER_ID = 'plugin:langbot-team/LocalAgent/default'
def make_session(
launcher_type: provider_session.LauncherTypes = provider_session.LauncherTypes.PERSON,
launcher_id: int = 12345,
) -> provider_session.Session:
"""Build a scope-aware Session that matches the shared Query factory."""
def attach_agent_runner_descriptor(app, *, multimodal_input=True, tool_calling=True):
"""Attach a schema-backed AgentRunner descriptor to a FakeApp."""
from langbot.pkg.agent.runner.descriptor import AgentRunnerDescriptor
descriptor = 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': []},
],
capabilities={
'tool_calling': tool_calling,
'multimodal_input': multimodal_input,
},
return provider_session.Session(
launcher_type=launcher_type,
launcher_id=launcher_id,
sender_id=12345,
bot_uuid='test-bot-uuid',
)
app.agent_runner_registry = Mock()
app.agent_runner_registry.get = AsyncMock(return_value=descriptor)
return descriptor
def agent_runner_pipeline_config(model_config, *, prompt='default'):
return {
'ai': {
'runner': {'id': RUNNER_ID},
'runner_config': {
RUNNER_ID: {
'model': model_config,
'prompt': prompt,
},
},
},
'output': {'misc': {'at-sender': False}},
'trigger': {'misc': {}},
}
class TestPreProcessorNormalText:
@@ -89,9 +84,7 @@ class TestPreProcessorNormalText:
app = FakeApp()
# Mock session manager to return a session
mock_session = Mock()
mock_session.launcher_type = Mock(value='person')
mock_session.launcher_id = 12345
mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
# Mock conversation
@@ -112,7 +105,7 @@ class TestPreProcessorNormalText:
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=mock_model)
# Mock tool manager
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(return_value=[])
app.tool_mgr.get_all_tools = AsyncMock(return_value=[])
# Mock plugin connector
mock_event_ctx = Mock()
@@ -135,9 +128,7 @@ class TestPreProcessorNormalText:
preproc = get_preproc_module()
app = FakeApp()
mock_session = Mock()
mock_session.launcher_type = Mock(value='person')
mock_session.launcher_id = 12345
mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -150,7 +141,7 @@ class TestPreProcessorNormalText:
mock_model = Mock()
mock_model.model_entity = Mock(uuid='test-model', abilities=['func_call'])
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=mock_model)
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(return_value=[])
app.tool_mgr.get_all_tools = AsyncMock(return_value=[])
mock_event_ctx = Mock()
mock_event_ctx.event = Mock(default_prompt=[], prompt=[])
@@ -175,9 +166,7 @@ class TestPreProcessorEmptyMessage:
entities = get_entities_module()
app = FakeApp()
mock_session = Mock()
mock_session.launcher_type = Mock(value='person')
mock_session.launcher_id = 12345
mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -188,7 +177,7 @@ class TestPreProcessorEmptyMessage:
app.sess_mgr.get_conversation = AsyncMock(return_value=mock_conversation)
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=None)
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(return_value=[])
app.tool_mgr.get_all_tools = AsyncMock(return_value=[])
mock_event_ctx = Mock()
mock_event_ctx.event = Mock(default_prompt=[], prompt=[])
@@ -214,9 +203,7 @@ class TestPreProcessorImageSegment:
preproc = get_preproc_module()
app = FakeApp()
mock_session = Mock()
mock_session.launcher_type = Mock(value='person')
mock_session.launcher_id = 12345
mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -230,7 +217,7 @@ class TestPreProcessorImageSegment:
mock_model = Mock()
mock_model.model_entity = Mock(uuid='vision-model', abilities=['func_call', 'vision'])
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=mock_model)
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(return_value=[])
app.tool_mgr.get_all_tools = AsyncMock(return_value=[])
mock_event_ctx = Mock()
mock_event_ctx.event = Mock(default_prompt=[], prompt=[])
@@ -262,9 +249,7 @@ class TestPreProcessorImageSegment:
preproc = get_preproc_module()
app = FakeApp()
mock_session = Mock()
mock_session.launcher_type = Mock(value='person')
mock_session.launcher_id = 12345
mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -278,7 +263,7 @@ class TestPreProcessorImageSegment:
mock_model = Mock()
mock_model.model_entity = Mock(uuid='text-only-model', abilities=['func_call'])
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=mock_model)
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(return_value=[])
app.tool_mgr.get_all_tools = AsyncMock(return_value=[])
mock_event_ctx = Mock()
mock_event_ctx.event = Mock(default_prompt=[], prompt=[])
@@ -301,9 +286,7 @@ class TestPreProcessorModelSelection:
preproc = get_preproc_module()
app = FakeApp()
mock_session = Mock()
mock_session.launcher_type = Mock(value='person')
mock_session.launcher_id = 12345
mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -316,7 +299,6 @@ class TestPreProcessorModelSelection:
mock_model = Mock()
mock_model.model_entity = Mock(uuid='primary-model-uuid', abilities=['func_call'])
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=mock_model)
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(return_value=[])
attach_agent_runner_descriptor(app)
mock_event_ctx = Mock()
@@ -327,9 +309,19 @@ class TestPreProcessorModelSelection:
query = text_query('hello')
# Set pipeline config with primary model
query.pipeline_config = agent_runner_pipeline_config(
{'primary': 'primary-model-uuid', 'fallbacks': []},
)
query.pipeline_config = {
'ai': {
'runner': {'id': RUNNER_ID},
'runner_config': {
RUNNER_ID: {
'model': {'primary': 'primary-model-uuid', 'fallbacks': []},
'prompt': [],
},
},
},
'output': {'misc': {'at-sender': False}},
'trigger': {'misc': {}},
}
result = await stage.process(query, 'PreProcessor')
@@ -341,9 +333,7 @@ class TestPreProcessorModelSelection:
preproc = get_preproc_module()
app = FakeApp()
mock_session = Mock()
mock_session.launcher_type = Mock(value='person')
mock_session.launcher_id = 12345
mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -360,7 +350,7 @@ class TestPreProcessorModelSelection:
mock_fallback = Mock()
mock_fallback.model_entity = Mock(uuid='fallback-uuid', abilities=['func_call'])
async def mock_get_model(uuid):
async def mock_get_model(_context, uuid):
if uuid == 'primary-uuid':
return mock_primary
elif uuid == 'fallback-uuid':
@@ -368,7 +358,6 @@ class TestPreProcessorModelSelection:
raise ValueError(f'Model {uuid} not found')
app.model_mgr.get_model_by_uuid = AsyncMock(side_effect=mock_get_model)
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(return_value=[])
attach_agent_runner_descriptor(app)
mock_event_ctx = Mock()
@@ -378,9 +367,19 @@ class TestPreProcessorModelSelection:
stage = preproc.PreProcessor(app)
query = text_query('hello')
query.pipeline_config = agent_runner_pipeline_config(
{'primary': 'primary-uuid', 'fallbacks': ['fallback-uuid']},
)
query.pipeline_config = {
'ai': {
'runner': {'id': RUNNER_ID},
'runner_config': {
RUNNER_ID: {
'model': {'primary': 'primary-uuid', 'fallbacks': ['fallback-uuid']},
'prompt': [],
},
},
},
'output': {'misc': {'at-sender': False}},
'trigger': {'misc': {}},
}
result = await stage.process(query, 'PreProcessor')
@@ -397,9 +396,7 @@ class TestPreProcessorVariables:
preproc = get_preproc_module()
app = FakeApp()
mock_session = Mock()
mock_session.launcher_type = Mock(value='person')
mock_session.launcher_id = 12345
mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -410,7 +407,7 @@ class TestPreProcessorVariables:
app.sess_mgr.get_conversation = AsyncMock(return_value=mock_conversation)
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=None)
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(return_value=[])
app.tool_mgr.get_all_tools = AsyncMock(return_value=[])
mock_event_ctx = Mock()
mock_event_ctx.event = Mock(default_prompt=[], prompt=[])
@@ -434,9 +431,10 @@ class TestPreProcessorVariables:
preproc = get_preproc_module()
app = FakeApp()
mock_session = Mock()
mock_session.launcher_type = Mock(value='group')
mock_session.launcher_id = 99999
mock_session = make_session(
provider_session.LauncherTypes.GROUP,
99999,
)
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -447,7 +445,7 @@ class TestPreProcessorVariables:
app.sess_mgr.get_conversation = AsyncMock(return_value=mock_conversation)
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=None)
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(return_value=[])
app.tool_mgr.get_all_tools = AsyncMock(return_value=[])
mock_event_ctx = Mock()
mock_event_ctx.event = Mock(default_prompt=[], prompt=[])
@@ -462,69 +460,17 @@ class TestPreProcessorVariables:
assert 'group_name' in variables
assert 'sender_name' in variables
@pytest.mark.asyncio
@pytest.mark.parametrize('invalid_value', [0, None, 'false'])
@pytest.mark.parametrize(
('configured_skills', 'expected_skills'),
[
(['bound-skill'], ['bound-skill']),
(None, []),
('bound-skill', []),
],
)
async def test_malformed_enable_all_skills_flag_uses_bound_skills(
self,
invalid_value,
configured_skills,
expected_skills,
):
preproc = get_preproc_module()
app = FakeApp()
mock_session = Mock()
mock_session.launcher_type = Mock(value='person')
mock_session.launcher_id = 12345
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
mock_conversation.prompt = Mock(messages=[])
mock_conversation.prompt.copy = Mock(return_value=Mock(messages=[]))
mock_conversation.messages = []
mock_conversation.uuid = None
app.sess_mgr.get_conversation = AsyncMock(return_value=mock_conversation)
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=None)
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(return_value=[])
app.pipeline_service.get_pipeline = AsyncMock(
return_value={
'extensions_preferences': {
'enable_all_skills': invalid_value,
'skills': configured_skills,
}
}
)
mock_event_ctx = Mock()
mock_event_ctx.event = Mock(default_prompt=[], prompt=[])
app.plugin_connector.emit_event = AsyncMock(return_value=mock_event_ctx)
result = await preproc.PreProcessor(app).process(text_query('hello'), 'PreProcessor')
assert result.new_query.variables['_pipeline_bound_skills'] == expected_skills
class TestPreProcessorToolSelection:
"""Tests for generic AgentRunner tool selection."""
"""Tests for Local Agent tool selection."""
@pytest.mark.asyncio
async def test_agent_runner_filters_selected_tools(self):
async def test_local_agent_filters_selected_tools(self):
"""Only selected tools should be exposed when all-tools mode is off."""
preproc = get_preproc_module()
app = FakeApp()
mock_session = Mock()
mock_session.launcher_type = Mock(value='person')
mock_session.launcher_id = 12345
mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -537,98 +483,37 @@ class TestPreProcessorToolSelection:
mock_model = Mock()
mock_model.model_entity = Mock(uuid='primary-model-uuid', abilities=['func_call'])
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=mock_model)
attach_agent_runner_descriptor(app)
app.tool_mgr.get_resolved_tool_catalog = AsyncMock(
return_value=[
{
'name': 'exec',
'source': 'builtin',
'description': 'Execute',
'parameters': {},
},
{
'name': 'plugin_tool',
'source': 'plugin',
'source_id': 'test/plugin',
'description': 'Plugin tool',
'parameters': {},
},
{
'name': 'mcp_tool',
'source': 'mcp',
'source_id': 'mcp-server',
'description': 'MCP tool',
'parameters': {},
},
{'name': 'exec', 'source': 'builtin'},
{'name': 'plugin_tool', 'source': 'plugin', 'source_id': 'test/plugin'},
{'name': 'mcp_tool', 'source': 'mcp', 'source_id': 'test-mcp'},
]
)
mock_event_ctx = Mock()
mock_event_ctx.event = Mock(default_prompt=[], prompt=[])
app.plugin_connector.emit_event = AsyncMock(return_value=mock_event_ctx)
attach_agent_runner_descriptor(app)
stage = preproc.PreProcessor(app)
query = text_query('hello')
query.pipeline_config = agent_runner_pipeline_config(
{'primary': 'primary-model-uuid', 'fallbacks': []},
)
query.pipeline_config['ai']['runner_config'][RUNNER_ID].update(
{
'enable-all-tools': False,
'tools': ['plugin_tool'],
}
)
query.pipeline_config = {
'ai': {
'runner': {'id': RUNNER_ID},
'runner_config': {
RUNNER_ID: {
'model': {'primary': 'primary-model-uuid', 'fallbacks': []},
'prompt': [],
'enable-all-tools': False,
'tools': ['plugin_tool'],
},
},
},
'output': {'misc': {'at-sender': False}},
'trigger': {'misc': {}},
}
result = await stage.process(query, 'PreProcessor')
assert [tool.name for tool in result.new_query.use_funcs] == ['plugin_tool']
assert result.new_query.variables['_host_tool_source_refs'] == {
'plugin_tool': {'source': 'plugin', 'source_id': 'test/plugin'},
}
class TestPreProcessorMCPResourceContext:
"""Tests for deferring MCP context until the run-scoped execution input."""
@pytest.mark.asyncio
async def test_pinned_context_does_not_mutate_preprocessed_input(self):
preproc = get_preproc_module()
from langbot.pkg.agent.runner.query_entry_adapter import QueryEntryAdapter
app = FakeApp()
mock_session = Mock()
mock_session.launcher_type = Mock(value='person')
mock_session.launcher_id = 12345
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
mock_conversation.prompt = Mock(messages=[])
mock_conversation.prompt.copy = Mock(return_value=Mock(messages=[]))
mock_conversation.messages = []
mock_conversation.uuid = 'conversation-1'
app.sess_mgr.get_conversation = AsyncMock(return_value=mock_conversation)
mock_model = Mock()
mock_model.model_entity = Mock(uuid='primary-model-uuid', abilities=[])
app.model_mgr.get_model_by_uuid = AsyncMock(return_value=mock_model)
mcp_loader = Mock()
mcp_loader.build_resource_context_for_query = AsyncMock(return_value='Pinned documentation')
app.tool_mgr.mcp_tool_loader = mcp_loader
mock_event_ctx = Mock()
mock_event_ctx.event = Mock(default_prompt=[], prompt=[])
app.plugin_connector.emit_event = AsyncMock(return_value=mock_event_ctx)
attach_agent_runner_descriptor(app, tool_calling=False)
query = text_query('hello')
query.launcher_id = '12345'
query.pipeline_config = agent_runner_pipeline_config(
{'primary': 'primary-model-uuid', 'fallbacks': []},
)
result = await preproc.PreProcessor(app).process(query, 'PreProcessor')
event = QueryEntryAdapter.query_to_event(result.new_query)
assert event.input.text == 'hello'
assert 'Pinned documentation' not in str(event.input.contents)
mcp_loader.build_resource_context_for_query.assert_not_awaited()
@@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, Mock
import pytest
from langbot_plugin.api.entities.builtin.platform import message as platform_message
from langbot_plugin.api.entities.builtin.provider import session as provider_session
RUNNER_ID = 'plugin:langbot-team/LocalAgent/default'
@@ -73,7 +74,12 @@ async def test_preprocessor_keeps_image_placeholder_for_text_only_local_agent(mo
mock_app.model_mgr.get_model_by_uuid = AsyncMock(return_value=model)
_attach_agent_runner_descriptor(mock_app)
mock_app.sess_mgr.get_session = AsyncMock(
return_value=SimpleNamespace(launcher_type=sample_query.launcher_type, launcher_id=sample_query.launcher_id)
return_value=provider_session.Session(
launcher_type=sample_query.launcher_type,
launcher_id=sample_query.launcher_id,
sender_id=sample_query.sender_id,
bot_uuid=sample_query.bot_uuid,
)
)
mock_app.sess_mgr.get_conversation = AsyncMock(return_value=_conversation())
mock_app.plugin_connector.emit_event = AsyncMock(return_value=_prompt_preprocessing_context())
+8 -1
View File
@@ -9,6 +9,7 @@ import langbot_plugin.api.definition.abstract.platform.adapter as abstract_platf
import langbot_plugin.api.definition.abstract.platform.event_logger as abstract_platform_logger
from langbot.pkg.pipeline.pool import QueryPool
from langbot.pkg.api.http.context import ExecutionContext
class DummyEventLogger(abstract_platform_logger.AbstractEventLogger):
@@ -64,12 +65,18 @@ async def test_add_query_returns_created_query_and_preserves_side_effects(
adapter=adapter,
pipeline_uuid='test-pipeline-uuid',
routed_by_rule=True,
execution_context=ExecutionContext(
instance_uuid='test-instance-uuid',
workspace_uuid='test-workspace-uuid',
placement_generation=1,
),
)
assert query is query_pool.queries[0]
assert query_pool.cached_queries[0] is query
assert query_pool.cached_queries[('test-workspace-uuid', query.query_uuid)] is query
assert query_pool.query_id_counter == 1
assert query.query_id == 0
assert query.bot_uuid == 'test-bot-uuid'
assert query.pipeline_uuid == 'test-pipeline-uuid'
assert query.workspace_uuid == 'test-workspace-uuid'
assert query.variables == {'_routed_by_rule': True}
+24 -4
View File
@@ -152,8 +152,18 @@ class TestFixedWindowAlgo:
# First request creates container
await algo.require_access(sample_query_with_rate_limit, provider_session.LauncherTypes.PERSON, '12345')
# Key format: 'LauncherTypes.PERSON_12345' (enum string representation)
expected_key = 'LauncherTypes.PERSON_12345'
context = sample_query_with_rate_limit._execution_context
expected_key = ':'.join(
(
context.instance_uuid,
context.workspace_uuid,
str(context.placement_generation),
str(sample_query_with_rate_limit.bot_uuid),
str(sample_query_with_rate_limit.pipeline_uuid),
str(provider_session.LauncherTypes.PERSON),
'12345',
)
)
assert expected_key in algo.containers
container = algo.containers[expected_key]
@@ -191,8 +201,18 @@ class TestFixedWindowAlgo:
for i in range(5):
await algo.require_access(sample_query, provider_session.LauncherTypes.PERSON, 'test')
# Key format: 'LauncherTypes.PERSON_test'
expected_key = 'LauncherTypes.PERSON_test'
context = sample_query._execution_context
expected_key = ':'.join(
(
context.instance_uuid,
context.workspace_uuid,
str(context.placement_generation),
str(sample_query.bot_uuid),
str(sample_query.pipeline_uuid),
str(provider_session.LauncherTypes.PERSON),
'test',
)
)
container = algo.containers[expected_key]
assert window_start in container.records
assert container.records[window_start] == 5
@@ -68,6 +68,25 @@ async def test_connection_listener_only_suppresses_exact_duplicates():
]
@pytest.mark.asyncio
async def test_connection_event_cache_is_bounded():
adapter, _ = _make_adapter()
for index in range(150):
await adapter._on_websocket_connection(aiocqhttp.Event({'self_id': index, 'time': index}))
assert len(adapter.on_websocket_connection_event_cache) == 100
def test_group_lookup_caches_are_bounded():
converter = AiocqhttpEventConverter()
converter._group_name_cache = {index: (str(index), 10_000.0) for index in range(5000)}
converter._prune_caches(1.0)
assert len(converter._group_name_cache) == 4096
def test_unregister_listener_removes_registered_wrapper():
adapter, _ = _make_adapter()
@@ -0,0 +1,400 @@
from __future__ import annotations
import asyncio
import contextlib
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
from langbot.pkg.api.http.authz import WorkspaceRequiredError
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.entity.persistence.bot import Bot
from langbot.pkg.platform.botmgr import PlatformManager, RuntimeBot
from langbot.pkg.workspace.entities import WorkspaceExecutionBinding
from langbot.pkg.workspace.errors import WorkspaceInvariantError
import langbot_plugin.api.entities.builtin.platform.events as platform_events
from tests.factories import friend_message_event, text_chain
WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
BOT_A = '10000000-0000-0000-0000-00000000000a'
BOT_B = '10000000-0000-0000-0000-00000000000b'
def _context(workspace_uuid: str, bot_uuid: str, generation: int = 4) -> ExecutionContext:
return ExecutionContext(
instance_uuid='instance',
workspace_uuid=workspace_uuid,
placement_generation=generation,
bot_uuid=bot_uuid,
)
def _runtime(application, workspace_uuid: str, bot_uuid: str) -> RuntimeBot:
entity = SimpleNamespace(
uuid=bot_uuid,
workspace_uuid=workspace_uuid,
name='Same Name',
enable=True,
event_bindings=[],
)
return RuntimeBot(
ap=application,
bot_entity=entity,
adapter=SimpleNamespace(),
logger=SimpleNamespace(),
execution_context=_context(workspace_uuid, bot_uuid),
)
class _WorkspaceService:
async def get_execution_binding(self, workspace_uuid, expected_generation=None):
if workspace_uuid not in {WORKSPACE_A, WORKSPACE_B} or expected_generation != 4:
raise ValueError('stale')
return SimpleNamespace(
instance_uuid='instance',
workspace_uuid=workspace_uuid,
placement_generation=4,
)
@pytest.fixture
def manager():
application = SimpleNamespace(workspace_service=_WorkspaceService())
platform_manager = PlatformManager(application)
platform_manager.bots = [
_runtime(application, WORKSPACE_A, BOT_A),
_runtime(application, WORKSPACE_B, BOT_B),
]
return platform_manager
@pytest.mark.asyncio
async def test_runtime_lookup_cannot_guess_another_workspace_bot(manager):
assert await manager.get_bot_by_uuid(_context(WORKSPACE_A, BOT_A), BOT_A) is manager.bots[0]
assert await manager.get_bot_by_uuid(_context(WORKSPACE_B, BOT_A), BOT_A) is None
assert await manager.get_bot_by_uuid(_context(WORKSPACE_A, BOT_B), BOT_B) is None
@pytest.mark.asyncio
async def test_public_route_key_resolves_bound_runtime_and_rejects_non_opaque_input(manager):
assert await manager.resolve_public_bot(BOT_A) is manager.bots[0]
assert await manager.resolve_public_bot('Same Name') is None
assert await manager.resolve_public_bot('not-a-uuid') is None
@pytest.mark.asyncio
async def test_stale_runtime_generation_is_not_returned(manager):
with pytest.raises(ValueError, match='stale'):
await manager.get_bot_by_uuid(_context(WORKSPACE_A, BOT_A, generation=5), BOT_A)
@pytest.mark.asyncio
async def test_generation_advance_shuts_down_and_prunes_old_workspace_bots():
class NoGlobalIterationDict(dict):
def __iter__(self):
raise AssertionError('generation advance scanned every bot runtime')
def items(self):
raise AssertionError('generation advance scanned every bot runtime')
def values(self):
raise AssertionError('generation advance scanned every bot runtime')
manager = PlatformManager(SimpleNamespace())
old_bot = SimpleNamespace(
workspace_uuid=WORKSPACE_A,
placement_generation=4,
enable=True,
shutdown=AsyncMock(),
)
other_bot = SimpleNamespace(
workspace_uuid=WORKSPACE_B,
placement_generation=4,
enable=True,
shutdown=AsyncMock(),
)
unrelated_bots = [
SimpleNamespace(
workspace_uuid=f'workspace-{index}',
placement_generation=4,
enable=False,
shutdown=AsyncMock(),
)
for index in range(1_000)
]
manager.bots = [old_bot, other_bot, *unrelated_bots]
old_context = _context(WORKSPACE_A, BOT_A, generation=4)
next_context = _context(WORKSPACE_A, BOT_A, generation=5)
await manager._observe_execution_context(old_context)
manager._bots_by_key = NoGlobalIterationDict(manager._bots_by_key)
await manager._observe_execution_context(next_context)
manager._bots_by_key = dict(manager._bots_by_key)
old_bot.shutdown.assert_awaited_once_with()
assert manager.bots == [other_bot, *unrelated_bots]
with pytest.raises(WorkspaceInvariantError, match='rolled back'):
await manager._observe_execution_context(old_context)
@pytest.mark.asyncio
async def test_concurrent_websocket_proxy_creation_reuses_one_runtime():
created_adapters = []
class WebsocketAdapter:
def __init__(self, *_args, **_kwargs):
created_adapters.append(self)
def register_listener(self, *_args):
pass
application = SimpleNamespace(workspace_service=_WorkspaceService())
manager = PlatformManager(application)
manager.adapter_dict = {'websocket': WebsocketAdapter}
context = ExecutionContext(
instance_uuid='instance',
workspace_uuid=WORKSPACE_A,
placement_generation=4,
)
runtimes = await asyncio.gather(*(manager.get_websocket_proxy_bot(context) for _ in range(20)))
assert len(created_adapters) == 1
assert len({id(runtime) for runtime in runtimes}) == 1
assert manager.websocket_proxy_bots == {WORKSPACE_A: runtimes[0]}
@pytest.mark.asyncio
async def test_websocket_proxy_cache_evicts_oldest_idle_workspace():
created_adapters = []
class WebsocketAdapter:
def __init__(self, *_args, **_kwargs):
self.kill = AsyncMock()
self.inbound_listener_tasks = set()
created_adapters.append(self)
def register_listener(self, *_args):
pass
application = SimpleNamespace(
workspace_service=_WorkspaceService(),
instance_config=SimpleNamespace(
data={
'system': {
'websocket_retention': {'max_workspace_proxies': 1},
}
}
),
)
manager = PlatformManager(application)
manager.adapter_dict = {'websocket': WebsocketAdapter}
await manager.get_websocket_proxy_bot(
ExecutionContext(
instance_uuid='instance',
workspace_uuid=WORKSPACE_A,
placement_generation=4,
)
)
second = await manager.get_websocket_proxy_bot(
ExecutionContext(
instance_uuid='instance',
workspace_uuid=WORKSPACE_B,
placement_generation=4,
)
)
created_adapters[0].kill.assert_awaited_once_with()
assert manager.websocket_proxy_bots == {WORKSPACE_B: second}
assert WORKSPACE_A not in manager._proxy_last_accessed
@pytest.mark.asyncio
async def test_reload_stops_and_drops_existing_platform_runtimes():
old_bot = SimpleNamespace(enable=True, shutdown=AsyncMock())
old_proxy = SimpleNamespace(enable=True, shutdown=AsyncMock())
persistence_mgr = SimpleNamespace(
execute_async=AsyncMock(return_value=SimpleNamespace(all=lambda: [])),
)
application = SimpleNamespace(
logger=SimpleNamespace(info=lambda *_args: None, warning=lambda *_args: None),
persistence_mgr=persistence_mgr,
workspace_service=SimpleNamespace(),
)
manager = PlatformManager(application)
manager.bots = [old_bot]
manager.websocket_proxy_bots = {WORKSPACE_A: old_proxy}
manager._scope_generations = {('instance', WORKSPACE_A): 4}
await manager.load_bots_from_db()
old_bot.shutdown.assert_awaited_once_with()
old_proxy.shutdown.assert_awaited_once_with()
assert manager.bots == []
assert manager.websocket_proxy_bots == {}
assert manager._scope_generations == {}
@pytest.mark.asyncio
async def test_cloud_startup_reuses_validated_platform_binding():
class TenantUow:
async def __aenter__(self):
return self
async def __aexit__(self, *_args):
return False
class ProbeAdapter:
def __init__(self, _config, _logger):
self.listeners = []
def register_listener(self, event_type, listener):
self.listeners.append((event_type, listener))
async def kill(self):
return None
binding = WorkspaceExecutionBinding(
instance_uuid='instance',
workspace_uuid=WORKSPACE_A,
placement_generation=4,
write_fenced=False,
state='active',
)
bot = Bot(
uuid=BOT_A,
workspace_uuid=WORKSPACE_A,
name='Probe',
description='',
adapter='probe',
adapter_config={},
enable=False,
event_bindings=[],
)
workspace_service = SimpleNamespace(
list_active_execution_bindings=AsyncMock(return_value=[binding]),
get_execution_binding=AsyncMock(
side_effect=AssertionError('startup platform loader repeated a validated binding lookup')
),
)
application = SimpleNamespace(
logger=SimpleNamespace(
info=lambda *_args, **_kwargs: None,
warning=lambda *_args, **_kwargs: None,
error=lambda *_args, **_kwargs: None,
),
persistence_mgr=SimpleNamespace(
mode=SimpleNamespace(value='cloud_runtime'),
tenant_uow=lambda _workspace_uuid: TenantUow(),
execute_async=AsyncMock(return_value=SimpleNamespace(all=lambda: [bot])),
),
workspace_service=workspace_service,
)
manager = PlatformManager(application)
manager.adapter_dict = {'probe': ProbeAdapter}
await manager.load_bots_from_db()
assert len(manager.bots) == 1
workspace_service.get_execution_binding.assert_not_awaited()
@pytest.mark.asyncio
async def test_runtime_bot_revalidates_its_generation_before_handling_events(manager):
runtime_bot = manager.bots[0]
await runtime_bot.assert_execution_active()
runtime_bot.placement_generation = 5
with pytest.raises(ValueError, match='stale'):
await runtime_bot.assert_execution_active()
def test_runtime_bot_rejects_workspace_mismatch():
application = SimpleNamespace()
entity = SimpleNamespace(
uuid=BOT_A,
workspace_uuid=WORKSPACE_A,
name='Bot',
enable=True,
event_bindings=[],
)
with pytest.raises(WorkspaceRequiredError):
RuntimeBot(
ap=application,
bot_entity=entity,
adapter=SimpleNamespace(),
logger=SimpleNamespace(),
execution_context=_context(WORKSPACE_B, BOT_A),
)
class _ScopeOnlyPersistenceManager:
mode = SimpleNamespace(value='cloud_runtime')
def __init__(self):
self.active_workspace = None
@contextlib.asynccontextmanager
async def tenant_scope(self, workspace_uuid: str):
assert self.active_workspace is None
self.active_workspace = workspace_uuid
try:
yield
finally:
self.active_workspace = None
def current_session(self):
return None
class _ListenerAdapter:
def __init__(self):
self.listeners = {}
def register_listener(self, event_type, listener):
self.listeners[event_type] = listener
@pytest.mark.asyncio
async def test_platform_callback_carries_scope_without_holding_database_session():
persistence_mgr = _ScopeOnlyPersistenceManager()
adapter = _ListenerAdapter()
async def push_person_message(*_args, **_kwargs):
assert persistence_mgr.active_workspace == WORKSPACE_A
assert persistence_mgr.current_session() is None
return True
application = SimpleNamespace(
persistence_mgr=persistence_mgr,
workspace_service=_WorkspaceService(),
webhook_pusher=SimpleNamespace(push_person_message=push_person_message),
)
entity = SimpleNamespace(
uuid=BOT_A,
workspace_uuid=WORKSPACE_A,
name='Bot',
enable=True,
event_bindings=[],
)
logger = SimpleNamespace(info=AsyncMock(), error=AsyncMock())
runtime = RuntimeBot(
ap=application,
bot_entity=entity,
adapter=adapter,
logger=logger,
execution_context=_context(WORKSPACE_A, BOT_A),
)
await runtime.initialize()
listener = adapter.listeners[platform_events.FriendMessage]
event = friend_message_event(text_chain('hello'), sender_id='user')
await listener(event, adapter)
assert persistence_mgr.active_workspace is None
logger.info.assert_awaited()
@@ -5,6 +5,7 @@ from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from langbot.pkg.platform import botmgr as _botmgr # noqa: F401
from langbot.pkg.platform.sources.dingtalk import (
DingTalkAdapter,
_dingtalk_card_markdown,
@@ -17,6 +18,17 @@ from langbot.pkg.platform.sources.dingtalk import (
)
def test_dingtalk_auxiliary_tasks_are_bounded():
adapter = DingTalkAdapter.model_construct()
adapter._background_tasks = {MagicMock(done=MagicMock(return_value=False)) for _ in range(100)}
async def callback():
raise AssertionError('rejected callback must not run')
assert adapter._start_background_task(callback()) is False
assert len(adapter._background_tasks) == 100
def test_dingtalk_select_component_params_expose_options():
params = _dingtalk_form_component_params(
{
@@ -0,0 +1,12 @@
from __future__ import annotations
import pytest
from langbot.pkg.platform.sources import discord
def test_discord_base64_decode_is_bounded(monkeypatch):
monkeypatch.setattr(discord, '_MAX_DISCORD_MEDIA_BYTES', 4)
with pytest.raises(ValueError, match='exceeds'):
discord._decode_discord_base64_limited('A' * 12)
@@ -0,0 +1,150 @@
from __future__ import annotations
import asyncio
import time
from types import SimpleNamespace
import pytest
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.platform.sources.http_bot import HttpBotAdapter
from langbot.pkg.platform.sources import http_bot as http_bot_module
def _session(key):
session = SimpleNamespace()
session._langbot_session_key = key
return session
def _adapter(app, execution_context) -> HttpBotAdapter:
adapter = HttpBotAdapter.model_construct(
config={'signature_required': False},
logger=SimpleNamespace(execution_context=execution_context),
bot_uuid='bot-a',
outbound_states={},
idempotency_cache={},
sync_waiters={},
inbound_tasks=set(),
)
object.__setattr__(adapter, 'ap', app)
return adapter
@pytest.mark.asyncio
async def test_http_bot_reset_removes_only_exact_execution_scope():
context = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=3,
bot_uuid='bot-a',
)
target_key = ('instance-a', 'workspace-a', 3, 'bot-a', 'person', 'shared-session')
retained_keys = [
('instance-b', 'workspace-a', 3, 'bot-a', 'person', 'shared-session'),
('instance-a', 'workspace-b', 3, 'bot-a', 'person', 'shared-session'),
('instance-a', 'workspace-a', 4, 'bot-a', 'person', 'shared-session'),
('instance-a', 'workspace-a', 3, 'bot-b', 'person', 'shared-session'),
('instance-a', 'workspace-a', 3, 'bot-a', 'group', 'shared-session'),
('instance-a', 'workspace-a', 3, 'bot-a', 'person', 'other-session'),
]
sessions = [_session(target_key), *[_session(key) for key in retained_keys], SimpleNamespace()]
app = SimpleNamespace(sess_mgr=SimpleNamespace(session_list=sessions))
adapter = _adapter(app, context)
removed = await adapter._reset_session('person', 'shared-session')
assert removed is True
assert [getattr(session, '_langbot_session_key', None) for session in app.sess_mgr.session_list] == [
*retained_keys,
None,
]
@pytest.mark.asyncio
async def test_http_bot_reset_fails_closed_without_trusted_scope():
app = SimpleNamespace(sess_mgr=SimpleNamespace(session_list=[]))
adapter = _adapter(app, None)
with pytest.raises(RuntimeError, match='trusted execution scope'):
await adapter._reset_session('person', 'shared-session')
@pytest.mark.asyncio
async def test_http_bot_bounds_inbound_listener_tasks(monkeypatch):
monkeypatch.setattr(http_bot_module, '_INBOUND_TASK_MAX', 1)
adapter = _adapter(SimpleNamespace(), None)
started = asyncio.Event()
release = asyncio.Event()
async def blocking_listener():
started.set()
await release.wait()
first = adapter._start_inbound_task(blocking_listener())
await started.wait()
rejected = adapter._start_inbound_task(blocking_listener())
assert first is not None
assert rejected is None
assert len(adapter.inbound_tasks) == 1
release.set()
await first
await asyncio.sleep(0)
assert adapter.inbound_tasks == set()
def test_http_bot_outbound_state_has_a_hard_capacity(monkeypatch):
monkeypatch.setattr(http_bot_module, '_OUTBOUND_STATE_MAX', 2)
monkeypatch.setattr(http_bot_module, '_OUTBOUND_PRUNE_SCAN_MAX', 2)
adapter = _adapter(SimpleNamespace(), None)
first = adapter._outbound_state('first')
second = adapter._outbound_state('second')
first.queue.put_nowait({})
second.queue.put_nowait({})
with pytest.raises(RuntimeError, match='outbound session capacity reached'):
adapter._next_sequence('third', is_final=True)
assert len(adapter.outbound_states) == 2
assert adapter._next_sequence('first', is_final=True) == 1
def test_http_bot_outbound_state_pruning_is_bounded_and_reclaims_stale(monkeypatch):
monkeypatch.setattr(http_bot_module, '_OUTBOUND_STATE_MAX', 2)
monkeypatch.setattr(http_bot_module, '_OUTBOUND_PRUNE_SCAN_MAX', 1)
monkeypatch.setattr(http_bot_module, '_OUTBOUND_IDLE_SECONDS', 10)
adapter = _adapter(SimpleNamespace(), None)
stale = adapter._outbound_state('stale')
stale.last_active = time.monotonic() - 11
adapter._outbound_state('active')
assert adapter._next_sequence('replacement', is_final=True) == 1
assert set(adapter.outbound_states) == {'active', 'replacement'}
def test_http_bot_idempotency_cache_has_a_hard_capacity(monkeypatch):
monkeypatch.setattr(http_bot_module, '_IDEMPOTENCY_MAX', 2)
monkeypatch.setattr(http_bot_module, '_IDEMPOTENCY_PRUNE_SCAN_MAX', 1)
adapter = _adapter(SimpleNamespace(), None)
assert adapter._reserve_idempotency_key('first') == 'accepted'
assert adapter._reserve_idempotency_key('second') == 'accepted'
assert adapter._reserve_idempotency_key('third') == 'overloaded'
assert len(adapter.idempotency_cache) == 2
assert adapter._reserve_idempotency_key('first') == 'duplicate'
def test_http_bot_idempotency_cache_reclaims_expired_oldest(monkeypatch):
monkeypatch.setattr(http_bot_module, '_IDEMPOTENCY_MAX', 2)
monkeypatch.setattr(http_bot_module, '_IDEMPOTENCY_PRUNE_SCAN_MAX', 1)
monkeypatch.setattr(http_bot_module, '_IDEMPOTENCY_TTL', 10)
adapter = _adapter(SimpleNamespace(), None)
adapter.idempotency_cache = {
'expired': time.monotonic() - 11,
'active': time.monotonic(),
}
assert adapter._reserve_idempotency_key('replacement') == 'accepted'
assert set(adapter.idempotency_cache) == {'active', 'replacement'}
@@ -0,0 +1,24 @@
from __future__ import annotations
import json
import zlib
import pytest
from langbot.pkg.platform.sources import kook
def test_kook_gateway_decoder_accepts_raw_and_compressed_json():
payload = {'s': 1, 'd': {'session_id': 'session-a'}}
encoded = json.dumps(payload).encode()
assert kook._decode_gateway_message(encoded) == payload
assert kook._decode_gateway_message(zlib.compress(encoded)) == payload
def test_kook_gateway_decoder_rejects_decompression_bomb(monkeypatch):
monkeypatch.setattr(kook, '_KOOK_MAX_GATEWAY_MESSAGE_BYTES', 1024)
compressed = zlib.compress(b'x' * 1025)
with pytest.raises(ValueError, match='decompressed size limit'):
kook._decode_gateway_message(compressed)
@@ -1,7 +1,13 @@
"""Tests for Lark adapter helper behavior."""
import threading
from unittest.mock import MagicMock
import pytest
from langbot.pkg.platform.sources.lark import (
LarkAdapter,
_decode_lark_base64_limited,
_lark_clean_form_content,
_lark_completed_input_lines,
_lark_current_input_defs,
@@ -11,6 +17,27 @@ from langbot.pkg.platform.sources.lark import (
)
def test_lark_base64_decode_is_bounded(monkeypatch):
import langbot.pkg.platform.sources.lark as lark_module
monkeypatch.setattr(lark_module, '_MAX_LARK_MEDIA_BYTES', 4)
with pytest.raises(ValueError, match='exceeds'):
_decode_lark_base64_limited('A' * 12)
def test_lark_threadsafe_callbacks_are_bounded():
adapter = LarkAdapter.model_construct()
adapter.threadsafe_event_lock = threading.Lock()
adapter.threadsafe_event_futures = {MagicMock(done=MagicMock(return_value=False)) for _ in range(100)}
async def callback():
raise AssertionError('rejected callback must not run')
assert adapter._schedule_threadsafe_event(callback()) is None
assert len(adapter.threadsafe_event_futures) == 100
def test_lark_current_input_defs_only_returns_active_stage():
input_defs = [
{'output_variable_name': 'us_input', 'type': 'paragraph'},
@@ -115,6 +115,17 @@ class DummyWSClient:
self._reconnect = AsyncMock()
class DummyExpiringCache:
def __init__(self, clear_interval=60):
self.clear_interval = clear_interval
def get(self, key):
return None
def set(self, key, value, ttl):
return None
def manifest() -> dict:
path = (
pathlib.Path(__file__).parents[3]
@@ -130,18 +141,19 @@ def manifest() -> dict:
def make_adapter(config: dict | None = None) -> LarkAdapter:
adapter = LarkAdapter(
{
'app_id': 'cli_xxx',
'app_secret': 'secret',
'bot_name': 'LangBotDev',
'enable-webhook': False,
'enable-stream-reply': False,
'app_type': 'self',
**(config or {}),
},
DummyLogger(),
)
with patch('lark_oapi.ws.client.ExpiringCache', DummyExpiringCache):
adapter = LarkAdapter(
{
'app_id': 'cli_xxx',
'app_secret': 'secret',
'bot_name': 'LangBotDev',
'enable-webhook': False,
'enable-stream-reply': False,
'app_type': 'self',
**(config or {}),
},
DummyLogger(),
)
adapter.api_client = DummyAPIClient()
adapter.bot = DummyWSClient()
return adapter
@@ -181,6 +193,17 @@ def test_lark_platform_api_map_matches_manifest():
assert set(PLATFORM_API_MAP) == manifest_actions
@pytest.mark.asyncio
async def test_lark_kill_cancels_sdk_cache_task():
adapter = make_adapter()
cache_task = asyncio.create_task(asyncio.sleep(60))
adapter.bot._cache = SimpleNamespace(_cron=cache_task)
assert await adapter.kill() is True
assert cache_task.cancelled()
adapter.bot._disconnect.assert_awaited_once_with()
@pytest.mark.asyncio
async def test_lark_message_converter_maps_outbound_components():
with (
@@ -0,0 +1,30 @@
from __future__ import annotations
import pytest
from unittest.mock import MagicMock
from langbot.pkg.platform import botmgr as _botmgr # noqa: F401
from langbot.pkg.platform.sources import line
def test_line_media_content_accepts_limit_boundary(monkeypatch) -> None:
monkeypatch.setattr(line, 'MAX_LINE_MEDIA_BYTES', 4)
content = b'1234'
assert line._validate_line_media_content(content) is content
def test_line_media_content_rejects_oversized_payload(monkeypatch) -> None:
monkeypatch.setattr(line, 'MAX_LINE_MEDIA_BYTES', 4)
with pytest.raises(ValueError, match='LINE media exceeds'):
line._validate_line_media_content(b'12345')
@pytest.mark.asyncio
async def test_line_kill_closes_api_client() -> None:
api_client = MagicMock()
adapter = line.LINEAdapter.model_construct(api_client=api_client)
assert await adapter.kill() is True
api_client.close.assert_called_once_with()
@@ -0,0 +1,21 @@
from __future__ import annotations
import pytest
from langbot.pkg.platform.sources import matrix
def test_matrix_base64_decode_is_bounded(monkeypatch):
monkeypatch.setattr(matrix, '_MAX_MATRIX_MEDIA_BYTES', 4)
with pytest.raises(ValueError, match='exceeds'):
matrix._decode_matrix_base64_limited('A' * 12)
def test_matrix_local_file_read_is_bounded(tmp_path, monkeypatch):
monkeypatch.setattr(matrix, '_MAX_MATRIX_MEDIA_BYTES', 4)
path = tmp_path / 'large.bin'
path.write_bytes(b'12345')
with pytest.raises(ValueError, match='exceeds'):
matrix._read_matrix_file_limited(str(path))
@@ -0,0 +1,31 @@
from __future__ import annotations
import pytest
from langbot.libs.openclaw_weixin_api.client import (
MAX_CDN_MEDIA_BYTES,
OpenClawWeixinClient,
_decrypt_cdn_payload,
_encrypt_cdn_payload,
)
from langbot.libs.openclaw_weixin_api.types import ApiError
def test_cdn_crypto_helpers_round_trip():
original = b'tenant-media' * 128
aes_key_hex, _encoded_key, encrypted, _raw_md5 = _encrypt_cdn_payload(original)
assert _decrypt_cdn_payload(encrypted, bytes.fromhex(aes_key_hex)) == original
@pytest.mark.asyncio
async def test_upload_media_rejects_oversized_input_before_network_access():
client = OpenClawWeixinClient('https://example.invalid', 'token')
with pytest.raises(ApiError, match='exceeds the size limit'):
await client.upload_media(
b'x' * (MAX_CDN_MEDIA_BYTES + 1),
'recipient',
3,
)
@@ -0,0 +1,95 @@
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
import langbot_plugin.api.entities.builtin.platform.message as platform_message
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.platform.sources import openclaw_weixin
from langbot.pkg.platform.sources.openclaw_weixin import OpenClawWeixinAdapter
def make_adapter(*, execution_context: ExecutionContext | None):
app = SimpleNamespace(
persistence_mgr=SimpleNamespace(execute_async=AsyncMock()),
workspace_service=SimpleNamespace(
get_execution_binding=AsyncMock(
return_value=SimpleNamespace(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
)
),
)
logger = SimpleNamespace(
ap=app,
execution_context=execution_context,
warning=AsyncMock(),
)
adapter = OpenClawWeixinAdapter.model_construct(
config={'token': 'refreshed-token'},
logger=logger,
client=Mock(),
bot_account_id='',
listeners={},
name='openclaw-weixin',
)
adapter._bot_uuid = 'shared-bot-uuid'
return adapter, app, logger
@pytest.mark.asyncio
async def test_persist_config_scopes_duplicate_bot_uuid_to_workspace():
adapter, app, _ = make_adapter(
execution_context=ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
bot_uuid='shared-bot-uuid',
)
)
await adapter._persist_config()
app.workspace_service.get_execution_binding.assert_awaited_once_with(
'workspace-a',
expected_generation=1,
)
statement = app.persistence_mgr.execute_async.await_args.args[0]
params = statement.compile().params
assert 'workspace-a' in params.values()
assert 'shared-bot-uuid' in params.values()
assert {'workspace_uuid', 'uuid'} <= {comparison.left.name for comparison in statement._where_criteria}
@pytest.mark.asyncio
@pytest.mark.parametrize(
'execution_context',
[
None,
ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
bot_uuid='another-bot-uuid',
),
],
ids=['missing-context', 'mismatched-bot'],
)
async def test_persist_config_fails_closed_without_matching_execution_context(execution_context):
adapter, app, logger = make_adapter(execution_context=execution_context)
await adapter._persist_config()
app.persistence_mgr.execute_async.assert_not_awaited()
logger.warning.assert_awaited_once()
@pytest.mark.asyncio
async def test_component_base64_decode_is_bounded(monkeypatch):
monkeypatch.setattr(openclaw_weixin, '_MAX_OPENCLAW_COMPONENT_BYTES', 4)
component = platform_message.File(base64='MTIzNDU=')
with pytest.raises(ValueError, match='exceeds'):
await OpenClawWeixinAdapter._get_component_bytes(component)
@@ -10,6 +10,7 @@ import langbot_plugin.api.entities.builtin.platform.message as platform_message
from langbot.libs.qq_official_api.api import (
QQ_SELECT_ACTION_PREFIX,
QQOfficialClient,
build_keyboard_from_select_field,
get_select_field_options,
resolve_select_button_action,
@@ -49,6 +50,28 @@ def test_qq_select_button_resolves_field_and_value():
assert resolve_select_button_action(form_data, f'{QQ_SELECT_ACTION_PREFIX}99') is None
@pytest.mark.asyncio
async def test_qq_seed_rejects_empty_secret_without_spinning():
client = QQOfficialClient('', 'token', 'app-id', AsyncMock())
with pytest.raises(ValueError, match='must not be empty'):
await asyncio.wait_for(client.repeat_seed(''), timeout=0.1)
def test_qq_auxiliary_tasks_are_bounded():
import langbot.pkg.core.app # noqa: F401
from langbot.pkg.platform.sources.qqofficial import QQOfficialAdapter
adapter = QQOfficialAdapter.model_construct()
adapter._background_tasks = {MagicMock(done=MagicMock(return_value=False)) for _ in range(100)}
async def callback():
raise AssertionError('rejected callback must not run')
assert adapter._start_background_task(callback()) is False
assert len(adapter._background_tasks) == 100
def test_qq_select_keyboard_fits_twenty_five_options():
form_data = _select_form_data()
form_data['input_defs'][0]['option_source']['value'] = [f'Option {idx}' for idx in range(25)]
@@ -5,6 +5,28 @@ from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.api.http.context import ExecutionContext
TEST_CONTEXT = ExecutionContext(
instance_uuid='instance-test',
workspace_uuid='workspace-test',
placement_generation=1,
bot_uuid='bot-1',
)
def active_workspace_service():
return SimpleNamespace(
get_execution_binding=AsyncMock(
return_value=SimpleNamespace(
instance_uuid=TEST_CONTEXT.instance_uuid,
workspace_uuid=TEST_CONTEXT.workspace_uuid,
placement_generation=TEST_CONTEXT.placement_generation,
)
)
)
class TestEventRouteTrace:
"""Test structured event route trace logging."""
@@ -14,7 +36,14 @@ class TestEventRouteTrace:
from langbot.pkg.platform.botmgr import RuntimeBot
bot = object.__new__(RuntimeBot)
bot.bot_entity = SimpleNamespace(uuid='bot-1', event_bindings=event_bindings)
bot.bot_entity = SimpleNamespace(
uuid='bot-1',
workspace_uuid=TEST_CONTEXT.workspace_uuid,
event_bindings=event_bindings,
)
bot.execution_context = TEST_CONTEXT
bot.workspace_uuid = TEST_CONTEXT.workspace_uuid
bot.placement_generation = TEST_CONTEXT.placement_generation
bot.logger = SimpleNamespace(
info=AsyncMock(),
warning=AsyncMock(),
@@ -91,6 +120,7 @@ class TestEventRouteTrace:
]
)
bot.ap = SimpleNamespace(
workspace_service=active_workspace_service(),
agent_service=SimpleNamespace(
get_agent=AsyncMock(
return_value={
@@ -163,6 +193,7 @@ class TestEventRouteTrace:
yield None
bot.ap = SimpleNamespace(
workspace_service=active_workspace_service(),
agent_service=SimpleNamespace(get_agent=AsyncMock(side_effect=[malformed_agent, valid_agent])),
agent_run_orchestrator=SimpleNamespace(run=fake_run),
)
@@ -194,6 +225,7 @@ class TestEventRouteTrace:
]
)
bot.ap = SimpleNamespace(
workspace_service=active_workspace_service(),
msg_aggregator=SimpleNamespace(add_message=AsyncMock()),
)
bot.adapter = SimpleNamespace(
@@ -342,7 +374,12 @@ class TestEventLoggerMetadata:
"""Metadata is optional and no_throw remains the fourth positional argument."""
from langbot.pkg.platform.logger import EventLogger
logger = EventLogger(name='test', ap=SimpleNamespace())
logger = EventLogger(
name='test',
ap=SimpleNamespace(),
execution_context=TEST_CONTEXT,
owner='bot-1',
)
await logger.info('plain log', None, None, False)
await logger.info(
@@ -368,9 +405,14 @@ class TestRuntimeBotLifecycle:
task_mgr = SimpleNamespace(cancel_task=Mock())
bot = RuntimeBot(
ap=SimpleNamespace(task_mgr=task_mgr),
bot_entity=SimpleNamespace(enable=True),
bot_entity=SimpleNamespace(
uuid='bot-1',
workspace_uuid=TEST_CONTEXT.workspace_uuid,
enable=True,
),
adapter=SimpleNamespace(kill=AsyncMock()),
logger=Mock(),
execution_context=TEST_CONTEXT,
)
await bot.shutdown()
@@ -391,9 +433,14 @@ class TestRuntimeBotLifecycle:
)
bot = RuntimeBot(
ap=SimpleNamespace(),
bot_entity=SimpleNamespace(enable=True),
bot_entity=SimpleNamespace(
uuid='bot-1',
workspace_uuid=TEST_CONTEXT.workspace_uuid,
enable=True,
),
adapter=adapter,
logger=Mock(),
execution_context=TEST_CONTEXT,
)
await bot.initialize()
@@ -414,9 +461,14 @@ class TestRuntimeBotLifecycle:
)
bot = RuntimeBot(
ap=SimpleNamespace(),
bot_entity=SimpleNamespace(enable=True),
bot_entity=SimpleNamespace(
uuid='bot-1',
workspace_uuid=TEST_CONTEXT.workspace_uuid,
enable=True,
),
adapter=adapter,
logger=Mock(),
execution_context=TEST_CONTEXT,
)
await bot.initialize()
@@ -613,6 +665,9 @@ class TestInteractionResumeRouting:
bot = object.__new__(RuntimeBot)
bot.bot_entity = SimpleNamespace(uuid='bot-1', name='Test', event_bindings=[])
bot.execution_context = TEST_CONTEXT
bot.workspace_uuid = TEST_CONTEXT.workspace_uuid
bot.placement_generation = TEST_CONTEXT.placement_generation
interaction_manager = SimpleNamespace(
consume_callback=AsyncMock(return_value=record),
acknowledge_submission=AsyncMock(),
@@ -750,3 +805,23 @@ class TestInteractionResumeRouting:
)
assert binding is None
def test_websocket_task_override_does_not_mutate_bot_default():
from langbot.pkg.platform.botmgr import RuntimeBot
bot = object.__new__(RuntimeBot)
bot.bot_entity = Mock(use_pipeline_uuid='default-uuid')
adapter = Mock()
adapter.get_pipeline_uuid_override.return_value = 'connection-pipeline'
pipeline_uuid, routed = bot.resolve_event_pipeline_uuid(
adapter,
'person',
'launcher',
'hello',
)
assert pipeline_uuid == 'connection-pipeline'
assert routed is False
assert bot.bot_entity.use_pipeline_uuid == 'default-uuid'
@@ -11,11 +11,21 @@ import langbot_plugin.api.entities.builtin.platform.events as platform_events
import langbot_plugin.api.entities.builtin.platform.message as platform_message
from langbot.pkg.platform.sources.telegram import (
TelegramAdapter,
_decode_telegram_base64_limited,
_telegram_form_action_from_callback,
_telegram_select_field_options,
)
def test_telegram_base64_decode_is_bounded(monkeypatch):
import langbot.pkg.platform.sources.telegram as telegram_module
monkeypatch.setattr(telegram_module, '_MAX_TELEGRAM_MEDIA_BYTES', 4)
with pytest.raises(ValueError, match='exceeds'):
_decode_telegram_base64_limited('A' * 12)
def _select_form_data() -> dict:
return {
'_current_input_field': 'choice',
@@ -88,6 +98,18 @@ def test_telegram_form_callback_cache_preserves_pipeline_uuid():
)
def test_telegram_form_callback_cache_is_bounded():
adapter = TelegramAdapter.model_construct()
adapter._form_action_titles = {}
adapter._cache_form_action_titles(
{f'callback-{index}': str(index) for index in range(5000)},
now=100.0,
)
assert len(adapter._form_action_titles) == adapter._MAX_FORM_ACTION_TITLES
@pytest.mark.asyncio
async def test_telegram_select_field_sends_two_column_inline_keyboard():
bot = MagicMock()
@@ -0,0 +1,99 @@
from __future__ import annotations
import asyncio
import logging
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
from langbot.pkg.platform.webhook_pusher import WebhookPusher
pytestmark = pytest.mark.asyncio
def _application(max_inflight_requests: object) -> SimpleNamespace:
return SimpleNamespace(
instance_config=SimpleNamespace(
data={
'webhooks': {
'max_inflight_requests': max_inflight_requests,
}
}
),
logger=logging.getLogger(__name__),
)
async def test_delivery_admission_never_queues_above_instance_limit():
pusher = WebhookPusher(_application(2))
release = asyncio.Event()
both_started = asyncio.Event()
calls = 0
active = 0
peak_active = 0
async def fake_push(url: str, payload: dict) -> dict:
nonlocal calls, active, peak_active
calls += 1
active += 1
peak_active = max(peak_active, active)
if active == 2:
both_started.set()
try:
await release.wait()
return {'url': url}
finally:
active -= 1
pusher._push_to_webhook = fake_push
webhooks = [{'url': f'https://example.invalid/{index}'} for index in range(5)]
first_delivery = asyncio.create_task(pusher._push_to_webhooks(webhooks, {}))
await asyncio.wait_for(both_started.wait(), timeout=1)
second_results = await pusher._push_to_webhooks(webhooks, {})
release.set()
first_results = await first_delivery
assert len(first_results) == 2
assert second_results == []
assert calls == 2
assert peak_active == 2
assert pusher._inflight_requests == 0
async def test_cancelled_delivery_reaps_children_and_releases_slots():
pusher = WebhookPusher(_application(1))
started = asyncio.Event()
never = asyncio.Event()
async def blocking_push(url: str, payload: dict) -> dict:
started.set()
await never.wait()
return {}
pusher._push_to_webhook = blocking_push
delivery = asyncio.create_task(
pusher._push_to_webhooks([{'url': 'https://example.invalid'}], {}),
)
await asyncio.wait_for(started.wait(), timeout=1)
delivery.cancel()
with pytest.raises(asyncio.CancelledError):
await delivery
assert pusher._inflight_requests == 0
pusher._push_to_webhook = AsyncMock(return_value={})
assert await pusher._push_to_webhooks([{'url': 'https://example.invalid'}], {}) == [{}]
async def test_max_inflight_requests_clamps_config():
pusher = WebhookPusher(_application(999999))
assert pusher._max_inflight_requests() == 128
pusher.ap.instance_config.data['webhooks']['max_inflight_requests'] = 0
assert pusher._max_inflight_requests() == 1
pusher.ap.instance_config.data['webhooks']['max_inflight_requests'] = 'invalid'
assert pusher._max_inflight_requests() == 16
@@ -4,157 +4,116 @@ The web debug client uploads Image / Voice / File components carrying a storage
key in ``path``. This helper resolves each to a base64 data URI (so multimodal
LLM input and the Box sandbox inbox have usable bytes), then deletes the
consumed storage object and clears ``path``. Covers mimetype selection per
type and graceful error handling.
type and fail-closed error handling.
"""
from __future__ import annotations
import base64
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
import langbot_plugin.api.entities.builtin.platform.events as platform_events
import langbot_plugin.api.entities.builtin.platform.entities as platform_entities
import langbot_plugin.api.entities.builtin.platform.message as platform_message
from langbot.pkg.platform.botmgr import RuntimeBot
from langbot.pkg.platform.sources.websocket_adapter import WebSocketAdapter, WebSocketSession
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.platform.sources.websocket_adapter import WebSocketAdapter
_CONTEXT = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
pipeline_uuid='pipeline-a',
)
_UPLOAD_PREFIX = 'v1/instance-a/workspace-a/1/upload_image/'
def _make_connection():
return SimpleNamespace(execution_context=_CONTEXT)
def _make_adapter(load_return=b'hello', load_side_effect=None):
provider = Mock()
provider.load = AsyncMock(return_value=load_return, side_effect=load_side_effect)
provider.delete = AsyncMock()
storage_mgr = Mock()
storage_mgr.storage_provider = provider
storage_mgr.load_scoped_object_key = AsyncMock(return_value=load_return, side_effect=load_side_effect)
storage_mgr.scoped_prefix.return_value = _UPLOAD_PREFIX
storage_mgr.is_scoped_object_key.return_value = True
storage_mgr.delete_scoped_object_key = AsyncMock()
ap = Mock()
ap.storage_mgr.storage_provider = provider
ap.storage_mgr = storage_mgr
logger = Mock()
logger.error = AsyncMock()
logger.warning = AsyncMock()
# WebSocketAdapter is a pydantic model; bypass full __init__/validation.
adapter = WebSocketAdapter.model_construct(ap=ap, logger=logger)
return adapter, provider
return adapter, storage_mgr, provider
@pytest.mark.asyncio
async def test_image_jpeg_mimetype_and_cleanup():
adapter, provider = _make_adapter(load_return=b'\xff\xd8\xff')
chain = [{'type': 'Image', 'path': 'storage://abc/photo.jpg'}]
adapter, storage_mgr, _ = _make_adapter(load_return=b'\xff\xd8\xff')
path = f'{_UPLOAD_PREFIX}photo.jpg'
chain = [{'type': 'Image', 'path': path}]
await adapter._process_image_components(chain)
await adapter._process_image_components(_make_connection(), chain)
expected_b64 = base64.b64encode(b'\xff\xd8\xff').decode('utf-8')
assert chain[0]['base64'] == f'data:image/jpeg;base64,{expected_b64}'
assert chain[0]['path'] == '' # consumed
provider.delete.assert_awaited_once_with('storage://abc/photo.jpg')
storage_mgr.delete_scoped_object_key.assert_awaited_once_with(
_CONTEXT,
path,
expected_owner_type='upload_image',
)
@pytest.mark.asyncio
async def test_image_defaults_to_png():
adapter, _ = _make_adapter()
chain = [{'type': 'Image', 'path': 'storage://abc/blob'}]
await adapter._process_image_components(chain)
adapter, _, _ = _make_adapter()
chain = [{'type': 'Image', 'path': f'{_UPLOAD_PREFIX}blob'}]
await adapter._process_image_components(_make_connection(), chain)
assert chain[0]['base64'].startswith('data:image/png;base64,')
@pytest.mark.asyncio
async def test_voice_uses_guessed_or_wav_mimetype():
adapter, _ = _make_adapter()
chain = [{'type': 'Voice', 'path': 'storage://abc/clip.wav'}]
await adapter._process_image_components(chain)
adapter, _, _ = _make_adapter()
chain = [{'type': 'Voice', 'path': f'{_UPLOAD_PREFIX}clip.wav'}]
await adapter._process_image_components(_make_connection(), chain)
assert chain[0]['base64'].startswith('data:audio/')
@pytest.mark.asyncio
async def test_file_uses_octet_stream_fallback():
adapter, _ = _make_adapter()
chain = [{'type': 'File', 'path': 'storage://abc/unknownblob'}]
await adapter._process_image_components(chain)
adapter, _, _ = _make_adapter()
chain = [{'type': 'File', 'path': f'{_UPLOAD_PREFIX}unknownblob'}]
await adapter._process_image_components(_make_connection(), chain)
assert chain[0]['base64'].startswith('data:application/octet-stream;base64,')
@pytest.mark.asyncio
async def test_skips_components_without_path_or_unknown_type():
adapter, provider = _make_adapter()
adapter, storage_mgr, provider = _make_adapter()
chain = [
{'type': 'Image', 'path': ''}, # no path
{'type': 'Plain', 'path': 'storage://abc/x'}, # not a file component
{'type': 'At', 'target': '123'}, # no path key at all
]
await adapter._process_image_components(chain)
await adapter._process_image_components(_make_connection(), chain)
provider.load.assert_not_awaited()
storage_mgr.load_scoped_object_key.assert_not_awaited()
assert 'base64' not in chain[0]
assert 'base64' not in chain[1]
@pytest.mark.asyncio
async def test_load_failure_is_logged_not_raised():
adapter, _ = _make_adapter(load_side_effect=RuntimeError('storage down'))
chain = [{'type': 'File', 'path': 'storage://abc/doc.pdf'}]
async def test_load_failure_is_logged_and_aborts_processing():
adapter, _, _ = _make_adapter(load_side_effect=RuntimeError('storage down'))
chain = [{'type': 'File', 'path': f'{_UPLOAD_PREFIX}doc.pdf'}]
# must not raise
await adapter._process_image_components(chain)
with pytest.raises(RuntimeError, match='storage down'):
await adapter._process_image_components(_make_connection(), chain)
assert 'base64' not in chain[0]
adapter.logger.error.assert_awaited_once()
@pytest.mark.asyncio
async def test_handle_websocket_message_marks_event_with_pipeline_uuid():
adapter, _ = _make_adapter()
adapter.websocket_person_session = WebSocketSession(id='websocketperson')
adapter.listeners = {}
adapter.listeners[platform_events.FriendMessage] = AsyncMock()
adapter.ap.platform_mgr.websocket_proxy_bot.bot_entity.use_pipeline_uuid = ''
connection = SimpleNamespace(
pipeline_uuid='pipeline-123',
session_type='person',
session_id=None,
connection_id='conn-1',
)
await adapter.handle_websocket_message(
connection,
{'message': [{'type': 'Plain', 'text': 'hello'}], 'stream': True},
)
await asyncio.sleep(0)
event = adapter.listeners[platform_events.FriendMessage].await_args.args[0]
assert getattr(event, '_langbot_pipeline_uuid') == 'pipeline-123'
@pytest.mark.asyncio
async def test_runtime_bot_websocket_listener_uses_event_pipeline_uuid():
app = Mock()
app.msg_aggregator.add_message = AsyncMock()
app.webhook_pusher = None
logger = Mock()
logger.info = AsyncMock()
logger.warning = AsyncMock()
logger.error = AsyncMock()
bot_entity = Mock()
bot_entity.uuid = 'websocket-proxy-bot'
bot_entity.enable = True
bot_entity.use_pipeline_uuid = ''
adapter = WebSocketAdapter.model_construct(
ap=app,
logger=Mock(error=AsyncMock()),
listeners={},
websocket_person_session=WebSocketSession(id='websocketperson'),
websocket_group_session=WebSocketSession(id='websocketgroup'),
)
bot = RuntimeBot(ap=app, bot_entity=bot_entity, adapter=adapter, logger=logger)
await bot.initialize()
event = platform_events.FriendMessage(
sender=platform_entities.Friend(id='sender-1', nickname='User', remark='User'),
message_chain=platform_message.MessageChain([platform_message.Plain(text='hello')]),
time=1,
)
object.__setattr__(event, '_langbot_pipeline_uuid', 'pipeline-123')
await adapter.listeners[platform_events.FriendMessage](event, adapter)
app.msg_aggregator.add_message.assert_awaited_once()
assert app.msg_aggregator.add_message.await_args.kwargs['pipeline_uuid'] == 'pipeline-123'
@@ -12,7 +12,25 @@ import langbot_plugin.api.entities.builtin.platform.message as platform_message
import langbot_plugin.api.entities.builtin.provider.message as provider_message
from langbot.pkg.platform.sources import websocket_adapter as websocket_adapter_module
from langbot.pkg.platform.sources.websocket_adapter import WebSocketAdapter, WebSocketMessage, WebSocketSession
from langbot.pkg.platform.sources.websocket_manager import WebSocketConnectionManager, is_valid_session_id
from langbot.pkg.platform.sources.websocket_manager import (
WebSocketConnectionManager,
WebSocketScope,
is_valid_session_id,
)
SCOPE_A = WebSocketScope('instance-a', 'workspace-a', 1)
SCOPE_B = WebSocketScope('instance-a', 'workspace-b', 1)
def _adapter_logger(scope: WebSocketScope = SCOPE_A):
logger = AsyncMock()
logger.execution_context = Mock(
instance_uuid=scope.instance_uuid,
workspace_uuid=scope.workspace_uuid,
placement_generation=scope.placement_generation,
)
return logger
@pytest.mark.asyncio
@@ -20,18 +38,21 @@ async def test_broadcast_only_reaches_connections_in_same_browser_session():
manager = WebSocketConnectionManager()
first = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
session_id='session-a',
)
second = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
session_id='session-b',
)
dashboard = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
)
@@ -39,6 +60,7 @@ async def test_broadcast_only_reaches_connections_in_same_browser_session():
await manager.broadcast_to_pipeline(
'pipeline-1',
{'type': 'response'},
scope=SCOPE_A,
session_type='person',
session_id='session-a',
)
@@ -50,6 +72,7 @@ async def test_broadcast_only_reaches_connections_in_same_browser_session():
await manager.broadcast_to_pipeline(
'pipeline-1',
{'type': 'dashboard-response'},
scope=SCOPE_A,
session_type='person',
session_id=None,
)
@@ -59,19 +82,114 @@ async def test_broadcast_only_reaches_connections_in_same_browser_session():
assert second.send_queue.empty()
@pytest.mark.asyncio
async def test_pipeline_indexes_and_broadcasts_are_workspace_scoped():
manager = WebSocketConnectionManager()
workspace_a = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='shared-pipeline',
session_type='person',
)
workspace_b = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_B,
pipeline_uuid='shared-pipeline',
session_type='person',
)
await manager.broadcast_to_pipeline(
'shared-pipeline',
{'type': 'workspace-a'},
scope=SCOPE_A,
)
assert await workspace_a.send_queue.get() == {'type': 'workspace-a'}
assert workspace_b.send_queue.empty()
assert await manager.get_connection(workspace_b.connection_id, scope=SCOPE_A) is None
assert await manager.get_connection(workspace_b.connection_id, scope=SCOPE_B) is workspace_b
assert manager.get_stats(scope=SCOPE_A)['total_connections'] == 1
@pytest.mark.asyncio
async def test_connection_admission_is_bounded_globally_and_per_workspace():
manager = WebSocketConnectionManager()
await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
max_connections=2,
max_connections_per_workspace=1,
)
with pytest.raises(RuntimeError, match='Workspace WebSocket'):
await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-2',
session_type='person',
max_connections=2,
max_connections_per_workspace=1,
)
await manager.add_connection(
websocket=Mock(),
scope=SCOPE_B,
pipeline_uuid='pipeline-1',
session_type='person',
max_connections=2,
max_connections_per_workspace=1,
)
with pytest.raises(RuntimeError, match='WebSocket connection capacity'):
await manager.add_connection(
websocket=Mock(),
scope=WebSocketScope('instance-a', 'workspace-c', 1),
pipeline_uuid='pipeline-1',
session_type='person',
max_connections=2,
max_connections_per_workspace=1,
)
@pytest.mark.asyncio
async def test_close_scope_closes_and_removes_only_matching_connections():
manager = WebSocketConnectionManager()
websocket_a = Mock(close=AsyncMock())
connection_a = await manager.add_connection(
websocket=websocket_a,
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
)
connection_b = await manager.add_connection(
websocket=Mock(close=AsyncMock()),
scope=SCOPE_B,
pipeline_uuid='pipeline-1',
session_type='person',
)
await manager.close_scope(SCOPE_A)
websocket_a.close.assert_awaited_once()
assert await manager.get_connection(connection_a.connection_id, scope=SCOPE_A) is None
assert await manager.get_connection(connection_b.connection_id, scope=SCOPE_B) is connection_b
@pytest.mark.asyncio
async def test_embed_event_uses_stable_session_launcher(monkeypatch):
manager = WebSocketConnectionManager()
session_id = '31c0f2e9-b115-4ee6-8f15-3e624d6456b1'
connection = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
session_id=session_id,
)
monkeypatch.setattr(websocket_adapter_module, 'ws_connection_manager', manager)
adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=AsyncMock())
adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=_adapter_logger())
adapter.websocket_person_session = WebSocketSession(id='person')
adapter.websocket_group_session = WebSocketSession(id='group')
received = []
@@ -95,13 +213,14 @@ async def test_embed_group_event_uses_stable_session_launcher(monkeypatch):
session_id = '31c0f2e9-b115-4ee6-8f15-3e624d6456b1'
connection = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='group',
session_id=session_id,
)
monkeypatch.setattr(websocket_adapter_module, 'ws_connection_manager', manager)
adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=AsyncMock())
adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=_adapter_logger())
adapter.websocket_person_session = WebSocketSession(id='person')
adapter.websocket_group_session = WebSocketSession(id='group')
received = []
@@ -121,6 +240,7 @@ async def test_embed_group_event_uses_stable_session_launcher(monkeypatch):
dashboard = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='group',
)
@@ -141,30 +261,46 @@ async def test_stable_session_launcher_resolves_to_active_connection(monkeypatch
session_id = '31c0f2e9-b115-4ee6-8f15-3e624d6456b1'
await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-2',
session_type='person',
session_id=session_id,
)
connection = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
session_id=session_id,
)
monkeypatch.setattr(websocket_adapter_module, 'ws_connection_manager', manager)
adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=AsyncMock())
adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=_adapter_logger())
message_source = Mock()
message_source.sender.id = f'websocket_pipeline-1:{session_id}'
assert await adapter._get_message_context(message_source) == ('pipeline-1', session_id)
assert await adapter._get_connection_from_target(f'websocketgroup_pipeline-1:{session_id}') is connection
assert await manager.get_connection_by_session_id(session_id, 'pipeline-1') is connection
assert (
await manager.get_connection_by_session_id(
session_id,
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
)
is connection
)
await manager.remove_connection(connection.connection_id)
assert await adapter._get_message_context(message_source) == ('pipeline-1', session_id)
assert await manager.get_connection_by_session_id(session_id, 'pipeline-1') is None
assert (
await manager.get_connection_by_session_id(
session_id,
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
)
is None
)
@pytest.mark.asyncio
@@ -172,6 +308,7 @@ async def test_dashboard_reply_uses_event_pipeline_after_connection_closes(monke
manager = WebSocketConnectionManager()
connection = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
)
@@ -200,12 +337,13 @@ async def test_late_final_events_update_one_stream_message(monkeypatch):
manager = WebSocketConnectionManager()
connection = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
)
monkeypatch.setattr(websocket_adapter_module, 'ws_connection_manager', manager)
adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=AsyncMock())
adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=_adapter_logger())
adapter.websocket_person_session = WebSocketSession(id='person')
adapter.websocket_group_session = WebSocketSession(id='group')
message_source = platform_events.FriendMessage(
@@ -256,7 +394,7 @@ def test_session_ids_must_be_canonical_random_uuids():
def test_history_read_does_not_allocate_unknown_session():
adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=AsyncMock())
adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=_adapter_logger())
adapter.websocket_person_session = WebSocketSession(id='person')
adapter.websocket_group_session = WebSocketSession(id='group')
@@ -264,16 +402,75 @@ def test_history_read_does_not_allocate_unknown_session():
assert adapter.websocket_person_session.message_lists == {}
@pytest.mark.asyncio
async def test_attachment_key_must_belong_to_connection_upload_scope():
manager = WebSocketConnectionManager()
connection = await manager.add_connection(
websocket=Mock(),
scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
)
storage_mgr = Mock()
storage_mgr.scoped_prefix.return_value = 'v1/current/upload_image/'
storage_mgr.is_scoped_object_key.return_value = True
storage_mgr.load_scoped_object_key = AsyncMock(return_value=b'image')
storage_mgr.delete_scoped_object_key = AsyncMock()
adapter = WebSocketAdapter.model_construct(
ap=Mock(storage_mgr=storage_mgr),
logger=_adapter_logger(),
)
message_chain = [{'type': 'Image', 'path': 'v1/current/upload_image/key.png'}]
await adapter._process_image_components(connection, message_chain)
assert message_chain[0]['base64'].startswith('data:image/png;base64,')
assert message_chain[0]['path'] == ''
storage_mgr.scoped_prefix.assert_called_once_with(
connection.execution_context,
owner_type='upload_image',
)
storage_mgr.is_scoped_object_key.assert_called_once_with(
'v1/current/upload_image/key.png',
expected_owner_type='upload_image',
)
storage_mgr.load_scoped_object_key.assert_awaited_once_with(
connection.execution_context,
'v1/current/upload_image/key.png',
expected_owner_type='upload_image',
)
storage_mgr.delete_scoped_object_key.assert_awaited_once_with(
connection.execution_context,
'v1/current/upload_image/key.png',
expected_owner_type='upload_image',
)
with pytest.raises(ValueError, match='does not belong'):
await adapter._process_image_components(
connection,
[{'type': 'File', 'path': 'v1/other/upload/key.txt'}],
)
def test_history_and_reset_are_scoped_to_browser_session():
matching_provider_session = Mock(
instance_uuid=SCOPE_A.instance_uuid,
workspace_uuid=SCOPE_A.workspace_uuid,
placement_generation=SCOPE_A.placement_generation,
launcher_type=Mock(value='person'),
launcher_id='websocket_pipeline-1:session-a',
)
matching_group_provider_session = Mock(
instance_uuid=SCOPE_A.instance_uuid,
workspace_uuid=SCOPE_A.workspace_uuid,
placement_generation=SCOPE_A.placement_generation,
launcher_type=Mock(value='group'),
launcher_id='websocketgroup_pipeline-1:session-a',
)
other_session = Mock(
instance_uuid=SCOPE_A.instance_uuid,
workspace_uuid=SCOPE_A.workspace_uuid,
placement_generation=SCOPE_A.placement_generation,
launcher_type=Mock(value='person'),
launcher_id='websocket_pipeline-1:session-b',
)
@@ -285,7 +482,7 @@ def test_history_and_reset_are_scoped_to_browser_session():
]
adapter = WebSocketAdapter.model_construct(
ap=ap,
logger=AsyncMock(),
logger=_adapter_logger(),
)
adapter.websocket_person_session = Mock()
adapter.websocket_group_session = Mock()
@@ -0,0 +1,39 @@
from __future__ import annotations
import httpx
import pytest
from langbot.libs.wechatpad_api.api import downloadpai
from langbot.libs.wechatpad_api.util import http_util
class _Response:
headers = {}
def __init__(self, chunks: list[bytes]):
self._chunks = chunks
def iter_content(self, chunk_size=None):
del chunk_size
yield from self._chunks
def test_wechatpad_response_reader_is_bounded(monkeypatch):
monkeypatch.setattr(http_util, '_MAX_WECHATPAD_RESPONSE_BYTES', 4)
with pytest.raises(RuntimeError, match='exceeds the runtime limit'):
http_util._read_requests_response_limited(_Response([b'1234', b'5']))
def test_wechatpad_response_reader_requires_json_object():
with pytest.raises(RuntimeError, match='non-object'):
http_util._read_requests_response_limited(_Response([b'[]']))
@pytest.mark.asyncio
async def test_wechatpad_media_reader_is_bounded(monkeypatch):
monkeypatch.setattr(downloadpai, '_MAX_WECHATPAD_MEDIA_BYTES', 4)
response = httpx.Response(200, content=b'oversized')
with pytest.raises(RuntimeError, match='exceeds'):
await downloadpai._read_media_limited(response)

Some files were not shown because too many files have changed in this diff Show More