mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-07-21 20:06:06 +00:00
feat(tenancy): implement workspace isolation
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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®ion=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=***®ion=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®ion=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=***®ion=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'] == ''
|
||||
|
||||
Reference in New Issue
Block a user