feat(tenancy): implement workspace isolation

This commit is contained in:
Junyan Qin
2026-07-19 09:58:59 +08:00
parent 9eb292992d
commit c6f826fe2d
271 changed files with 31162 additions and 6106 deletions
@@ -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
@@ -48,7 +53,7 @@ async def test_update_bot_copies_input_before_filtering_and_setting_pipeline_nam
'use_pipeline_uuid': 'pipeline-1',
}
await service.update_bot('bot-1', payload)
await service.update_bot(WORKSPACE_UUID, 'bot-1', payload)
assert payload == {
'uuid': 'caller-owned-uuid',
@@ -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,373 @@
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(WorkspaceNotFoundError):
await application.bot_service.update_bot(
WORKSPACE_A,
'bot-a',
{'use_pipeline_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,138 @@
from __future__ import annotations
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
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 _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_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()
@@ -1,482 +1,443 @@
"""
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_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()
@@ -19,6 +19,8 @@ from langbot.pkg.entity.persistence.bot import Bot
pytestmark = pytest.mark.asyncio
WORKSPACE_UUID = 'workspace-a'
def _create_mock_bot(
bot_uuid: str = None,
@@ -73,7 +75,9 @@ class TestBotServiceGetBots:
service = BotService(ap)
# Execute
result = await service.get_bots()
result = await service.get_bots(
WORKSPACE_UUID,
)
# Verify
assert result == []
@@ -101,7 +105,7 @@ class TestBotServiceGetBots:
service = BotService(ap)
# Execute
result = await service.get_bots(include_secret=True)
result = await service.get_bots(WORKSPACE_UUID, include_secret=True)
# Verify
assert len(result) == 2
@@ -130,7 +134,7 @@ class TestBotServiceGetBots:
service = BotService(ap)
# Execute
result = await service.get_bots(include_secret=False)
result = await service.get_bots(WORKSPACE_UUID, include_secret=False)
# Verify - adapter_config should be masked
assert result[0]['adapter_config'] is None
@@ -159,7 +163,7 @@ class TestBotServiceGetBot:
service = BotService(ap)
# Execute
result = await service.get_bot('test-uuid')
result = await service.get_bot(WORKSPACE_UUID, 'test-uuid')
# Verify
assert result is not None
@@ -178,7 +182,7 @@ class TestBotServiceGetBot:
service = BotService(ap)
# Execute
result = await service.get_bot('nonexistent-uuid')
result = await service.get_bot(WORKSPACE_UUID, 'nonexistent-uuid')
# Verify
assert result is None
@@ -203,7 +207,7 @@ class TestBotServiceGetRuntimeBotInfo:
# Execute & Verify
with pytest.raises(Exception, match='Bot not found'):
await service.get_runtime_bot_info('nonexistent-uuid')
await service.get_runtime_bot_info(WORKSPACE_UUID, 'nonexistent-uuid')
async def test_get_runtime_bot_info_returns_webhook_for_wecom(self):
"""Returns webhook URL for wecom adapter."""
@@ -231,7 +235,7 @@ class TestBotServiceGetRuntimeBotInfo:
service.get_bot = AsyncMock(return_value=bot_data)
# Execute
result = await service.get_runtime_bot_info('wecom-uuid')
result = await service.get_runtime_bot_info(WORKSPACE_UUID, 'wecom-uuid')
# Verify
assert result['adapter_runtime_values']['webhook_url'] == '/bots/wecom-uuid'
@@ -257,7 +261,7 @@ class TestBotServiceGetRuntimeBotInfo:
service.get_bot = AsyncMock(return_value=bot_data)
# Execute
result = await service.get_runtime_bot_info('telegram-uuid')
result = await service.get_runtime_bot_info(WORKSPACE_UUID, 'telegram-uuid')
# Verify - no webhook for telegram
assert result['adapter_runtime_values']['webhook_url'] is None
@@ -288,7 +292,7 @@ class TestBotServiceGetRuntimeBotInfo:
service.get_bot = AsyncMock(return_value=bot_data)
# Execute
result = await service.get_runtime_bot_info('runtime-uuid')
result = await service.get_runtime_bot_info(WORKSPACE_UUID, 'runtime-uuid')
# Verify
assert result['adapter_runtime_values']['bot_account_id'] == 'runtime-account-123'
@@ -318,7 +322,7 @@ class TestBotServiceCreateBot:
# Execute & Verify
with pytest.raises(ValueError, match='Maximum number of bots'):
await service.create_bot({'name': 'New Bot'})
await service.create_bot(WORKSPACE_UUID, {'name': 'New Bot'})
async def test_create_bot_no_limit(self):
"""Creates bot without limit check when max_bots=-1."""
@@ -360,7 +364,9 @@ class TestBotServiceCreateBot:
service = BotService(ap)
# Execute
bot_uuid = await service.create_bot({'name': 'New Bot', 'adapter': 'telegram', 'adapter_config': {}})
bot_uuid = await service.create_bot(
WORKSPACE_UUID, {'name': 'New Bot', 'adapter': 'telegram', 'adapter_config': {}}
)
# Verify
assert bot_uuid is not None
@@ -412,11 +418,15 @@ class TestBotServiceCreateBot:
# Execute
bot_data = {'name': 'New Bot', 'adapter': 'telegram', 'adapter_config': {}}
bot_uuid = await service.create_bot(bot_data)
bot_uuid = await service.create_bot(WORKSPACE_UUID, bot_data)
# Verify - pipeline uuid and name were set
assert 'use_pipeline_uuid' in bot_data
assert 'use_pipeline_name' in bot_data
# The service owns a copy and cannot mutate caller input while adding tenant data.
assert bot_data == {'name': 'New Bot', 'adapter': 'telegram', 'adapter_config': {}}
insert_statement = ap.persistence_mgr.execute_async.await_args_list[1].args[0]
insert_values = insert_statement.compile().params
assert insert_values['workspace_uuid'] == WORKSPACE_UUID
assert insert_values['use_pipeline_uuid'] == 'default-pipeline-uuid'
assert insert_values['use_pipeline_name'] == 'Default Pipeline'
assert bot_uuid is not None # Verify UUID was returned
@@ -446,7 +456,7 @@ class TestBotServiceUpdateBot:
# Execute
update_data = {'uuid': 'should-be-removed', 'name': 'Updated Name'}
await service.update_bot('test-uuid', update_data)
await service.update_bot(WORKSPACE_UUID, 'test-uuid', update_data)
update_params = ap.persistence_mgr.execute_async.await_args_list[0].args[0].compile().params
assert update_params['name'] == 'Updated Name'
@@ -467,7 +477,7 @@ class TestBotServiceUpdateBot:
# Execute & Verify
with pytest.raises(Exception, match='Pipeline not found'):
await service.update_bot('test-uuid', {'use_pipeline_uuid': 'nonexistent-pipeline'})
await service.update_bot(WORKSPACE_UUID, 'test-uuid', {'use_pipeline_uuid': 'nonexistent-pipeline'})
async def test_update_bot_sets_pipeline_name(self):
"""Sets use_pipeline_name when updating use_pipeline_uuid."""
@@ -504,7 +514,7 @@ class TestBotServiceUpdateBot:
ap.platform_mgr.load_bot = AsyncMock(return_value=runtime_bot)
# Execute
await service.update_bot('test-uuid', {'use_pipeline_uuid': 'pipeline-uuid'})
await service.update_bot(WORKSPACE_UUID, 'test-uuid', {'use_pipeline_uuid': 'pipeline-uuid'})
update_params = ap.persistence_mgr.execute_async.await_args_list[1].args[0].compile().params
assert update_params['use_pipeline_uuid'] == 'pipeline-uuid'
@@ -524,12 +534,13 @@ class TestBotServiceDeleteBot:
ap.platform_mgr.remove_bot = AsyncMock()
service = BotService(ap)
service.get_bot = AsyncMock(return_value={'uuid': 'bot-uuid'})
# Execute
await service.delete_bot('test-uuid')
await service.delete_bot(WORKSPACE_UUID, 'test-uuid')
# Verify
ap.platform_mgr.remove_bot.assert_called_once_with('test-uuid')
ap.platform_mgr.remove_bot.assert_called_once_with(WORKSPACE_UUID, 'test-uuid')
ap.persistence_mgr.execute_async.assert_called_once()
async def test_delete_bot_nonexistent_uuid(self):
@@ -542,9 +553,10 @@ class TestBotServiceDeleteBot:
ap.platform_mgr.remove_bot = AsyncMock()
service = BotService(ap)
service.get_bot = AsyncMock(return_value={'uuid': 'bot-uuid'})
# Execute - should not raise
await service.delete_bot('nonexistent-uuid')
await service.delete_bot(WORKSPACE_UUID, 'nonexistent-uuid')
# Verify - both called regardless
ap.platform_mgr.remove_bot.assert_called_once()
@@ -561,10 +573,11 @@ class TestBotServiceListEventLogs:
ap.platform_mgr.get_bot_by_uuid = AsyncMock(return_value=None)
service = BotService(ap)
service.get_bot = AsyncMock(return_value={'uuid': 'nonexistent-uuid'})
# Execute & Verify
with pytest.raises(Exception, match='Bot not found'):
await service.list_event_logs('nonexistent-uuid', 0, 10)
await service.list_event_logs(WORKSPACE_UUID, 'nonexistent-uuid', 0, 10)
async def test_list_event_logs_returns_logs(self):
"""Returns logs from runtime bot logger."""
@@ -581,9 +594,10 @@ class TestBotServiceListEventLogs:
ap.platform_mgr.get_bot_by_uuid = AsyncMock(return_value=runtime_bot)
service = BotService(ap)
service.get_bot = AsyncMock(return_value={'uuid': 'bot-uuid'})
# Execute
logs, total = await service.list_event_logs('bot-uuid', 0, 10)
logs, total = await service.list_event_logs(WORKSPACE_UUID, 'bot-uuid', 0, 10)
# Verify
assert len(logs) == 1
@@ -602,10 +616,11 @@ class TestBotServiceSendMessage:
ap.platform_mgr.get_bot_by_uuid = AsyncMock(return_value=None)
service = BotService(ap)
service.get_bot = AsyncMock(return_value={'uuid': 'nonexistent-uuid'})
# Execute & Verify
with pytest.raises(Exception, match='Bot not found'):
await service.send_message('nonexistent-uuid', 'group', '123', {'test': 'data'})
await service.send_message(WORKSPACE_UUID, 'nonexistent-uuid', 'group', '123', {'test': 'data'})
async def test_send_message_invalid_message_chain_raises(self):
"""Raises Exception when message_chain_data is invalid."""
@@ -619,10 +634,11 @@ class TestBotServiceSendMessage:
ap.platform_mgr.get_bot_by_uuid = AsyncMock(return_value=runtime_bot)
service = BotService(ap)
service.get_bot = AsyncMock(return_value={'uuid': 'bot-uuid'})
# Execute & Verify - invalid format should raise
with pytest.raises(Exception, match='Invalid message_chain format'):
await service.send_message('bot-uuid', 'group', '123', {'invalid': 'format'})
await service.send_message(WORKSPACE_UUID, 'bot-uuid', 'group', '123', {'invalid': 'format'})
async def test_send_message_valid_call(self):
"""Sends message through adapter when all valid."""
@@ -636,6 +652,7 @@ class TestBotServiceSendMessage:
ap.platform_mgr.get_bot_by_uuid = AsyncMock(return_value=runtime_bot)
service = BotService(ap)
service.get_bot = AsyncMock(return_value={'uuid': 'bot-uuid'})
# Execute with valid message chain format
message_chain_data = {'messages': [{'type': 'text', 'data': {'text': 'Hello'}}]}
@@ -644,7 +661,7 @@ class TestBotServiceSendMessage:
with patch('langbot_plugin.api.entities.builtin.platform.message.MessageChain') as MockMessageChain:
mock_chain = Mock()
MockMessageChain.model_validate = Mock(return_value=mock_chain)
await service.send_message('bot-uuid', 'group', '123', message_chain_data)
await service.send_message(WORKSPACE_UUID, 'bot-uuid', 'group', '123', message_chain_data)
# Verify adapter.send_message was called
runtime_bot.adapter.send_message.assert_called_once_with('group', '123', mock_chain)
@@ -1,389 +1,565 @@
"""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(),
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_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()
@@ -19,10 +19,32 @@ 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 +54,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 +69,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 +89,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 +126,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 +155,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,7 +190,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 - warning logged, defaults used
assert ap.logger.warning.called
@@ -196,7 +227,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 +260,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 +296,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 +325,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 +347,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 +369,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 +405,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 +435,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 +649,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 +664,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 +678,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 +697,7 @@ class TestMaintenanceServiceExpiredLogCandidates:
# Setup
ap = SimpleNamespace()
ap.logger = SimpleNamespace()
ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
@@ -748,11 +784,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 +799,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 +810,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 +821,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 +845,152 @@ 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]
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
+317 -69
View File
@@ -13,17 +13,59 @@ Source: src/langbot/pkg/api/http/service/mcp.py
from __future__ import annotations
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.entity.persistence.mcp import MCPServer
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 +84,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 +108,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 +125,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 +145,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 +166,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 +179,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 +204,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 +233,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,17 +263,90 @@ 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."""
@@ -241,16 +362,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 +396,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 +418,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,10 +455,10 @@ 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()
@@ -351,10 +476,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 +504,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 +523,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 +544,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 +556,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 +587,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 +609,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 +628,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 +646,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 +677,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 +783,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 +797,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 +823,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 +837,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 +862,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,13 +876,12 @@ 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:
@@ -667,12 +903,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 +927,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,12 +946,17 @@ 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()
@@ -25,11 +25,37 @@ 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
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(
model_uuid: str = 'llm-uuid',
@@ -101,6 +127,35 @@ def _create_mock_result(items: list = None, first_item=None):
return result
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:
"""Tests for _parse_provider_api_keys helper function."""
@@ -183,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 == []
@@ -221,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
@@ -260,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'] == ['***', '***']
@@ -302,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
@@ -321,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
@@ -346,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
@@ -360,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())
@@ -374,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
@@ -391,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())
@@ -405,6 +465,7 @@ class TestLLMModelsServiceCreateLLMModel:
# Execute
model_uuid = await service.create_llm_model(
WORKSPACE_UUID,
{
'uuid': 'preserved-uuid',
'name': 'Preserved UUID Model',
@@ -422,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())
@@ -434,6 +495,7 @@ class TestLLMModelsServiceCreateLLMModel:
service = LLMModelsService(ap)
await service.create_llm_model(
WORKSPACE_UUID,
{
'uuid': 'model-with-context',
'name': 'Context Model',
@@ -446,7 +508,7 @@ 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
@@ -456,7 +518,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 = {} # Empty - no provider
mock_result = _create_mock_result([])
@@ -467,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):
@@ -480,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())
@@ -500,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': {
@@ -509,7 +573,7 @@ class TestLLMModelsServiceCreateLLMModel:
},
'abilities': [],
'extra_args': {},
}
},
)
# Verify - provider_service was called and UUID generated
@@ -525,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()
@@ -534,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',
@@ -546,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',
@@ -575,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',
@@ -594,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}
@@ -609,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()
@@ -617,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:
@@ -640,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 == []
@@ -677,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
@@ -717,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
@@ -734,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
@@ -748,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())
@@ -760,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
@@ -776,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([])
@@ -787,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': {},
}
},
)
@@ -803,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()
@@ -811,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()
@@ -832,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 == []
@@ -869,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
@@ -909,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
@@ -926,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
@@ -940,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())
@@ -952,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
@@ -967,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([])
@@ -978,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': {},
}
},
)
@@ -994,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()
@@ -1002,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()
@@ -1027,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
@@ -1052,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
@@ -1066,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,
}
@@ -0,0 +1,154 @@
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.workspace import Workspace
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_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
@@ -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,7 +238,7 @@ class TestPipelineServiceCreatePipeline:
# Execute & Verify
with pytest.raises(ValueError, match='Maximum number of pipelines'):
await service.create_pipeline({'name': 'New Pipeline'})
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."""
@@ -258,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
@@ -293,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()
@@ -340,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'] == {
@@ -394,7 +403,7 @@ class TestPipelineServiceUpdatePipeline:
'is_default': True,
'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'
@@ -450,7 +459,7 @@ class TestPipelineServiceUpdatePipeline:
service.get_pipeline = AsyncMock(return_value={'uuid': 'test-uuid', 'name': 'New Name'})
# Execute with name change
await service.update_pipeline('test-uuid', {'name': 'New Name'})
await service.update_pipeline(WORKSPACE_UUID, 'test-uuid', {'name': 'New Name'})
# Verify - bot_service.update_bot was called for each bot
assert ap.bot_service.update_bot.call_count == 2
@@ -478,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
@@ -499,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):
@@ -517,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()
@@ -549,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."""
@@ -570,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."""
@@ -614,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
@@ -647,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()
@@ -667,8 +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', [])
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."""
@@ -715,6 +724,7 @@ class TestPipelineServiceUpdatePipelineExtensions:
# Execute
bound_plugins = [{'plugin_uuid': 'plugin-1'}]
await service.update_pipeline_extensions(
WORKSPACE_UUID,
'test-uuid',
bound_plugins=bound_plugins,
enable_all_plugins=False,
@@ -764,6 +774,7 @@ class TestPipelineServiceUpdatePipelineExtensions:
# Execute
await service.update_pipeline_extensions(
WORKSPACE_UUID,
'test-uuid',
bound_plugins=[],
bound_mcp_servers=['mcp-server-1'],
@@ -811,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()
@@ -850,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'] == [
@@ -858,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,7 +827,7 @@ 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
@@ -810,7 +866,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 +916,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 +924,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,8 @@ Source: src/langbot/pkg/api/http/service/space.py
from __future__ import annotations
from urllib.parse import parse_qs, urlsplit
import pytest
from unittest.mock import AsyncMock, Mock, patch, MagicMock
from types import SimpleNamespace
@@ -73,7 +75,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 +95,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."""
@@ -14,17 +14,60 @@ 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.persistence.user import AccountStatus, User
from langbot.pkg.entity.errors.account import AccountEmailMismatchError
pytestmark = pytest.mark.asyncio
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, _ = service._space_oauth_states[digest]
service._space_oauth_states[digest] = (purpose, account_uuid, 0)
with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
await service.consume_space_oauth_state(state, 'login')
def _create_mock_user(
email: str = 'test@example.com',
password: str = 'hashed_password',
@@ -309,6 +352,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."""
@@ -548,6 +635,44 @@ 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(AccountEmailMismatchError):
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_create_or_update_space_user_no_expiry(self):
"""Creates Space user without token expiry."""
# Setup
@@ -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(
@@ -47,6 +55,14 @@ def _create_mock_result(items: list = None, first_item=None):
result = Mock()
result.all = Mock(return_value=items or [])
result.first = Mock(return_value=first_item)
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 +87,7 @@ class TestWebhookServiceGetWebhooks:
service = WebhookService(ap)
# Execute
result = await service.get_webhooks()
result = await service.get_webhooks(WORKSPACE_UUID)
# Verify
assert result == []
@@ -100,7 +116,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 +135,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(
@@ -155,6 +172,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 +205,7 @@ class TestWebhookServiceCreateWebhook:
nonlocal call_count
call_count += 1
if call_count == 1:
return Mock() # Insert
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 +222,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 +246,7 @@ class TestWebhookServiceCreateWebhook:
nonlocal call_count
call_count += 1
if call_count == 1:
return Mock()
return _create_write_result()
return _create_mock_result(first_item=created_webhook)
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
@@ -233,7 +255,12 @@ 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
@@ -262,7 +289,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 +308,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 +325,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 +339,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 +354,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 +369,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 +384,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 +399,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 +421,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 +442,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 +457,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 +483,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 +511,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 +531,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,195 @@
from __future__ import annotations
import asyncio
from types import SimpleNamespace
from unittest.mock import AsyncMock
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,
_pop_owned_session,
)
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,
),
)
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(),
)
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_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)
+48 -14
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,63 @@ 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):
@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
)
query_result = Mock()
query_result.first.return_value = db_row
persistence_mgr = SimpleNamespace(execute_async=AsyncMock(return_value=query_result))
query_result.first.return_value = key
persistence_mgr = SimpleNamespace(execute_async=AsyncMock(side_effect=[query_result, Mock(rowcount=1)]))
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 == (2 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,94 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock
import pytest
import quart
from langbot.pkg.api.http.controller.groups.extensions import ExtensionsRouterGroup
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'
@@ -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,69 @@
from __future__ import annotations
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.api.http.controller.groups.knowledge.migration import KnowledgeMigrationRouterGroup
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()
+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,108 @@
from __future__ import annotations
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()
@@ -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'