mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-08-09 04:40:57 +00:00
feat(tenancy): add Workspace multi-tenant foundation (#2353)
* Document multi-tenant workspace architecture * Add OSS and commercial workspace boundaries * docs: redesign multi-tenant workspace architecture * feat(tenancy): implement workspace isolation * docs(tenancy): record verification evidence * docs(tenancy): revise single-instance SaaS topology * docs(tenancy): refine architecture options * docs: finalize cloud v2 multi-tenant decisions * feat(tenancy): establish cloud isolation foundations * feat(tenancy): harden shared cloud runtime boundaries * docs(tenancy): record final isolation verification * fix(tenancy): close isolation and permission gaps * docs(tenancy): record final isolation verification * feat(tenancy): connect cloud workspace control plane * fix(build): install git for pinned SDK * docs(cloud): update control plane verification * chore: update multi-tenant SDK pin * fix(cloud): skip legacy model sync during startup * test(cloud): preserve minimal model manager fixtures * fix(cloud): preserve authenticated account context * fix(cloud): reuse authenticated account for user info * feat(cloud): complete Workspace settings navigation * test(web): cover Workspace dropdown menu * feat(web): place workspace controls in sidebar * refactor(web): streamline workspace controls * style(web): format workspace layout test * fix(cloud): surface runtime and workspace plan status * fix(plugin): keep runtime identity stable across restarts * fix(ui): widen and center workspace switcher * fix(ui): hide roles from workspace switcher * fix(ui): align workspace switcher with sidebar entries * feat(workspace): add in-product collaboration and direct Cloud launch * style: format collaboration changes * fix(workspace): bind collaboration APIs to tenant UoW * fix(cloud): preserve Core-owned collaboration state * test(cloud): require Space identity for invite registration * feat(cloud): complete secure invitation experience * style(web): format invitation flows * fix(cloud): recover box runtime without unscoped skill reload * feat(oss): enforce invitation account and owner billing flows * style: format OSS account service * test(oss): cover invitation logout handoff * fix(oss): resolve workspace owner in scoped session * feat(cloud): harden multi-tenant runtime resources * fix(cloud): bound runtime restart storms * fix(cloud): eliminate periodic runtime CPU spikes * fix(cloud): enforce instance capacity ceilings * fix(cloud): scope public login capability discovery * fix(cloud): bound tenant maintenance and monitoring work * fix(runtime): bound tenant resource amplification * fix(deps): pin green multi-tenant plugin SDK * fix(cloud): handle unavailable skill capability * fix(security): require authentication for image file endpoint (H-2) - Changed /api/v1/files/image from AuthType.NONE to USER_TOKEN_OR_API_KEY - Added Permission.RESOURCE_VIEW requirement - Prevents unauthenticated cross-tenant file access via leaked keys - Fixes HIGH severity finding from multi-tenant security review docs: add comprehensive database migration guide - Complete migration steps for OSS → multi-tenant - Backup, execution, verification procedures - Rollback scenarios and recovery plans - Performance tuning recommendations * test: add comprehensive cross-tenant isolation tests Added 7 critical test scenarios for multi-tenant boundaries: - Cross-tenant bot access prevention - Viewer role read-only enforcement - Removed member immediate access revocation - Model provider credential isolation - WebSocket message isolation - Invitation token workspace scoping - Multi-workspace context validation These tests address P0-2 coverage gaps for: - workspaces.py (membership & invitation flows) - user.py (authentication & authorization) - websocket_chat.py (real-time isolation) - plugins.py (resource access control) docs: finalize database migration guide * fix(security): resolve M-1, M-2, M-3 security findings M-1: WebSocket authorization TOCTOU race (FIXED) - Changed _revalidate_websocket_authorization to return RequestContext - Ensures validated context is used immediately without race window - Prevents removed members from sending messages during revalidation gap M-2: Model Manager cache workspace isolation (VERIFIED) - Confirmed _CacheKey already uses 4-tuple: (instance, workspace, generation, resource) - Cache is properly scoped per workspace, no cross-tenant leakage possible - No code change needed, documented as working correctly M-3: Invitation lock workspace scoping (FIXED) - Changed lock key from token_digest to workspace_uuid:token_digest - Prevents DoS where attacker locks token in Workspace A to block Workspace B - Locks now isolated per workspace All MEDIUM severity findings from security review now resolved. * fix(cloud): unblock tenant CI and enforce knowledge quotas * fix(tenancy): scope rerank model sync --------- Co-authored-by: dadachann <185672915+dadachann@users.noreply.github.com>
This commit is contained in:
@@ -6,6 +6,9 @@ from sqlalchemy.sql.dml import Update
|
||||
from langbot.pkg.api.http.service.bot import BotService
|
||||
|
||||
|
||||
WORKSPACE_UUID = 'workspace-a'
|
||||
|
||||
|
||||
class _FakeResult:
|
||||
def __init__(self, value):
|
||||
self.value = value
|
||||
@@ -21,7 +24,9 @@ class _PersistenceManager:
|
||||
async def execute_async(self, statement):
|
||||
if isinstance(statement, Update):
|
||||
self.update_values = {
|
||||
key: value for key, value in statement.compile().params.items() if not key.startswith('uuid_')
|
||||
key: value
|
||||
for key, value in statement.compile().params.items()
|
||||
if not key.startswith(('uuid_', 'workspace_uuid_'))
|
||||
}
|
||||
return None
|
||||
|
||||
@@ -48,7 +53,7 @@ async def test_update_bot_copies_input_before_filtering_and_setting_pipeline_nam
|
||||
'use_pipeline_uuid': 'pipeline-1',
|
||||
}
|
||||
|
||||
await service.update_bot('bot-1', payload)
|
||||
await service.update_bot(WORKSPACE_UUID, 'bot-1', payload)
|
||||
|
||||
assert payload == {
|
||||
'uuid': 'caller-owned-uuid',
|
||||
|
||||
@@ -0,0 +1,34 @@
|
||||
import pytest
|
||||
import sqlalchemy
|
||||
|
||||
from langbot.pkg.api.http.authz import WorkspaceRequiredError
|
||||
from langbot.pkg.api.http.context import ExecutionContext, PrincipalContext, PrincipalType
|
||||
from langbot.pkg.api.http.service.tenant import require_workspace_uuid, scope_statement
|
||||
|
||||
|
||||
class _TenantRow:
|
||||
workspace_uuid = sqlalchemy.column('workspace_uuid')
|
||||
|
||||
|
||||
def test_require_workspace_uuid_accepts_execution_context():
|
||||
context = ExecutionContext(
|
||||
instance_uuid='instance-test',
|
||||
workspace_uuid='workspace-test',
|
||||
placement_generation=1,
|
||||
trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
|
||||
)
|
||||
|
||||
assert require_workspace_uuid(context) == 'workspace-test'
|
||||
|
||||
|
||||
@pytest.mark.parametrize('context', [None, '', ' '])
|
||||
def test_require_workspace_uuid_rejects_missing_context(context):
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
require_workspace_uuid(context)
|
||||
|
||||
|
||||
def test_scope_statement_adds_workspace_predicate():
|
||||
statement = scope_statement(sqlalchemy.select(_TenantRow.workspace_uuid), _TenantRow, 'workspace-test')
|
||||
|
||||
assert 'workspace_uuid = :workspace_uuid_1' in str(statement)
|
||||
assert statement.compile().params == {'workspace_uuid_1': 'workspace-test'}
|
||||
@@ -0,0 +1,373 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import sqlalchemy
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
from langbot.pkg.api.http.service.bot import BotService
|
||||
from langbot.pkg.api.http.service.model import LLMModelsService
|
||||
from langbot.pkg.api.http.service.pipeline import PipelineService
|
||||
from langbot.pkg.api.http.service.provider import ModelProviderService
|
||||
from langbot.pkg.api.http.service.tenant import require_workspace_uuid
|
||||
from langbot.pkg.api.http.authz import WorkspaceRequiredError
|
||||
from langbot.pkg.entity.persistence.base import Base
|
||||
from langbot.pkg.entity.persistence.bot import Bot
|
||||
from langbot.pkg.entity.persistence.model import LLMModel, ModelProvider
|
||||
from langbot.pkg.entity.persistence.pipeline import LegacyPipeline
|
||||
from langbot.pkg.entity.persistence.workspace import Workspace
|
||||
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
|
||||
WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
|
||||
|
||||
|
||||
class _PersistenceManager:
|
||||
def __init__(self, engine):
|
||||
self.engine = engine
|
||||
|
||||
async def execute_async(self, *args, **kwargs):
|
||||
async with self.engine.connect() as connection:
|
||||
result = await connection.execute(*args, **kwargs)
|
||||
await connection.commit()
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def serialize_model(model, data, masked_columns=None):
|
||||
masked_columns = masked_columns or []
|
||||
return {
|
||||
column.name: (
|
||||
getattr(data, column.name).isoformat()
|
||||
if isinstance(getattr(data, column.name), datetime.datetime)
|
||||
else getattr(data, column.name)
|
||||
)
|
||||
for column in model.__table__.columns
|
||||
if column.name not in masked_columns
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def tenant_services(tmp_path):
|
||||
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "tenant-resources.db"}')
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(Workspace),
|
||||
[
|
||||
{
|
||||
'uuid': WORKSPACE_A,
|
||||
'instance_uuid': 'instance-a',
|
||||
'name': 'Workspace A',
|
||||
'slug': 'workspace-a',
|
||||
'source': 'cloud_projection',
|
||||
},
|
||||
{
|
||||
'uuid': WORKSPACE_B,
|
||||
'instance_uuid': 'instance-b',
|
||||
'name': 'Workspace B',
|
||||
'slug': 'workspace-b',
|
||||
'source': 'cloud_projection',
|
||||
},
|
||||
],
|
||||
)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(ModelProvider),
|
||||
[
|
||||
{
|
||||
'uuid': 'provider-a',
|
||||
'workspace_uuid': WORKSPACE_A,
|
||||
'name': 'Same Provider',
|
||||
'requester': 'chatcmpl',
|
||||
'base_url': 'https://a.invalid',
|
||||
'api_keys': ['secret-a'],
|
||||
},
|
||||
{
|
||||
'uuid': 'provider-b',
|
||||
'workspace_uuid': WORKSPACE_B,
|
||||
'name': 'Same Provider',
|
||||
'requester': 'chatcmpl',
|
||||
'base_url': 'https://b.invalid',
|
||||
'api_keys': ['secret-b'],
|
||||
},
|
||||
],
|
||||
)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(LLMModel),
|
||||
[
|
||||
{
|
||||
'uuid': 'model-a',
|
||||
'workspace_uuid': WORKSPACE_A,
|
||||
'name': 'Same Model',
|
||||
'provider_uuid': 'provider-a',
|
||||
'abilities': [],
|
||||
'extra_args': {},
|
||||
'prefered_ranking': 0,
|
||||
},
|
||||
{
|
||||
'uuid': 'model-b',
|
||||
'workspace_uuid': WORKSPACE_B,
|
||||
'name': 'Same Model',
|
||||
'provider_uuid': 'provider-b',
|
||||
'abilities': [],
|
||||
'extra_args': {},
|
||||
'prefered_ranking': 0,
|
||||
},
|
||||
],
|
||||
)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(LegacyPipeline),
|
||||
[
|
||||
{
|
||||
'uuid': 'pipeline-a',
|
||||
'workspace_uuid': WORKSPACE_A,
|
||||
'name': 'Same Pipeline',
|
||||
'description': 'A',
|
||||
'for_version': 'test',
|
||||
'is_default': False,
|
||||
'stages': [],
|
||||
'config': {},
|
||||
'extensions_preferences': {},
|
||||
},
|
||||
{
|
||||
'uuid': 'pipeline-b',
|
||||
'workspace_uuid': WORKSPACE_B,
|
||||
'name': 'Same Pipeline',
|
||||
'description': 'B',
|
||||
'for_version': 'test',
|
||||
'is_default': False,
|
||||
'stages': [],
|
||||
'config': {},
|
||||
'extensions_preferences': {},
|
||||
},
|
||||
],
|
||||
)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(Bot),
|
||||
[
|
||||
{
|
||||
'uuid': 'bot-a',
|
||||
'workspace_uuid': WORKSPACE_A,
|
||||
'name': 'Same Bot',
|
||||
'description': 'A',
|
||||
'adapter': 'test',
|
||||
'adapter_config': {},
|
||||
'enable': False,
|
||||
'use_pipeline_uuid': 'pipeline-a',
|
||||
'use_pipeline_name': 'Same Pipeline',
|
||||
'pipeline_routing_rules': [],
|
||||
},
|
||||
{
|
||||
'uuid': 'bot-b',
|
||||
'workspace_uuid': WORKSPACE_B,
|
||||
'name': 'Same Bot',
|
||||
'description': 'B',
|
||||
'adapter': 'test',
|
||||
'adapter_config': {},
|
||||
'enable': False,
|
||||
'use_pipeline_uuid': 'pipeline-b',
|
||||
'use_pipeline_name': 'Same Pipeline',
|
||||
'pipeline_routing_rules': [],
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
runtime_provider_a = SimpleNamespace(provider_entity=SimpleNamespace(uuid='provider-a'))
|
||||
runtime_provider_b = SimpleNamespace(provider_entity=SimpleNamespace(uuid='provider-b'))
|
||||
application = SimpleNamespace(
|
||||
persistence_mgr=_PersistenceManager(engine),
|
||||
instance_config=SimpleNamespace(data={'system': {'limitation': {}}, 'api': {}}),
|
||||
ver_mgr=SimpleNamespace(get_current_version=lambda: 'test'),
|
||||
platform_mgr=SimpleNamespace(
|
||||
load_bot=AsyncMock(return_value=SimpleNamespace(enable=False)),
|
||||
remove_bot=AsyncMock(),
|
||||
get_bot_by_uuid=AsyncMock(return_value=None),
|
||||
),
|
||||
pipeline_mgr=SimpleNamespace(
|
||||
load_pipeline=AsyncMock(),
|
||||
remove_pipeline=AsyncMock(),
|
||||
),
|
||||
model_mgr=SimpleNamespace(
|
||||
provider_dict={'provider-a': runtime_provider_a, 'provider-b': runtime_provider_b},
|
||||
llm_models=[],
|
||||
embedding_models=[],
|
||||
rerank_models=[],
|
||||
load_provider=AsyncMock(),
|
||||
cache_provider=AsyncMock(),
|
||||
get_provider_by_uuid=AsyncMock(return_value=runtime_provider_a),
|
||||
reload_provider=AsyncMock(),
|
||||
remove_provider=AsyncMock(),
|
||||
load_llm_model_with_provider=AsyncMock(return_value=SimpleNamespace()),
|
||||
cache_llm_model=AsyncMock(),
|
||||
remove_llm_model=AsyncMock(),
|
||||
),
|
||||
sess_mgr=SimpleNamespace(session_list=[]),
|
||||
)
|
||||
application.provider_service = ModelProviderService(application)
|
||||
application.llm_model_service = LLMModelsService(application)
|
||||
application.pipeline_service = PipelineService(application)
|
||||
application.bot_service = BotService(application)
|
||||
|
||||
yield application, engine
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def test_context_is_mandatory_and_fails_closed(tenant_services):
|
||||
application, _engine = tenant_services
|
||||
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
require_workspace_uuid(None)
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
await application.bot_service.get_bots(None)
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
await application.provider_service.get_providers(None)
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
await application.pipeline_service.get_pipelines(None)
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
await application.llm_model_service.get_llm_models(None)
|
||||
|
||||
|
||||
async def test_lists_and_same_names_are_isolated(tenant_services):
|
||||
application, _engine = tenant_services
|
||||
|
||||
assert [item['uuid'] for item in await application.bot_service.get_bots(WORKSPACE_A)] == ['bot-a']
|
||||
assert [item['uuid'] for item in await application.pipeline_service.get_pipelines(WORKSPACE_A)] == ['pipeline-a']
|
||||
assert [item['uuid'] for item in await application.provider_service.get_providers(WORKSPACE_A)] == ['provider-a']
|
||||
assert [item['uuid'] for item in await application.llm_model_service.get_llm_models(WORKSPACE_A)] == ['model-a']
|
||||
|
||||
|
||||
async def test_cross_workspace_uuid_guessing_cannot_read_update_or_delete(tenant_services):
|
||||
application, engine = tenant_services
|
||||
|
||||
assert await application.bot_service.get_bot(WORKSPACE_A, 'bot-b') is None
|
||||
assert await application.pipeline_service.get_pipeline(WORKSPACE_A, 'pipeline-b') is None
|
||||
assert await application.provider_service.get_provider(WORKSPACE_A, 'provider-b') is None
|
||||
assert await application.llm_model_service.get_llm_model(WORKSPACE_A, 'model-b') is None
|
||||
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await application.bot_service.update_bot(WORKSPACE_A, 'bot-b', {'name': 'stolen'})
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await application.pipeline_service.update_pipeline(
|
||||
WORKSPACE_A,
|
||||
'pipeline-b',
|
||||
{'description': 'stolen'},
|
||||
)
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await application.provider_service.update_provider(WORKSPACE_A, 'provider-b', {'name': 'stolen'})
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await application.llm_model_service.update_llm_model(
|
||||
WORKSPACE_A,
|
||||
'model-b',
|
||||
{'name': 'stolen'},
|
||||
)
|
||||
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await application.bot_service.delete_bot(WORKSPACE_A, 'bot-b')
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await application.pipeline_service.delete_pipeline(WORKSPACE_A, 'pipeline-b')
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await application.provider_service.delete_provider(WORKSPACE_A, 'provider-b')
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await application.llm_model_service.delete_llm_model(WORKSPACE_A, 'model-b')
|
||||
|
||||
async with engine.connect() as connection:
|
||||
assert await connection.scalar(sqlalchemy.select(Bot.name).where(Bot.uuid == 'bot-b')) == 'Same Bot'
|
||||
assert (
|
||||
await connection.scalar(sqlalchemy.select(LegacyPipeline.uuid).where(LegacyPipeline.uuid == 'pipeline-b'))
|
||||
== 'pipeline-b'
|
||||
)
|
||||
assert (
|
||||
await connection.scalar(sqlalchemy.select(ModelProvider.name).where(ModelProvider.uuid == 'provider-b'))
|
||||
== 'Same Provider'
|
||||
)
|
||||
assert await connection.scalar(sqlalchemy.select(LLMModel.uuid).where(LLMModel.uuid == 'model-b')) == 'model-b'
|
||||
|
||||
|
||||
async def test_cross_workspace_parent_references_are_rejected(tenant_services):
|
||||
application, _engine = tenant_services
|
||||
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await application.bot_service.update_bot(
|
||||
WORKSPACE_A,
|
||||
'bot-a',
|
||||
{'use_pipeline_uuid': 'pipeline-b'},
|
||||
)
|
||||
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await application.llm_model_service.create_llm_model(
|
||||
WORKSPACE_A,
|
||||
{
|
||||
'name': 'Cross reference',
|
||||
'provider_uuid': 'provider-b',
|
||||
'abilities': [],
|
||||
'extra_args': {},
|
||||
'prefered_ranking': 0,
|
||||
},
|
||||
auto_set_to_default_pipeline=False,
|
||||
)
|
||||
|
||||
|
||||
async def test_created_resources_are_bound_to_callers_workspace(tenant_services):
|
||||
application, engine = tenant_services
|
||||
|
||||
runtime_provider = SimpleNamespace(provider_entity=SimpleNamespace(uuid='provider-created'))
|
||||
application.model_mgr.load_provider.return_value = runtime_provider
|
||||
provider_uuid = await application.provider_service.create_provider(
|
||||
WORKSPACE_A,
|
||||
{
|
||||
'name': 'Created Provider',
|
||||
'requester': 'chatcmpl',
|
||||
'base_url': 'https://created.invalid',
|
||||
'api_keys': [],
|
||||
},
|
||||
)
|
||||
pipeline_uuid = await application.pipeline_service.create_pipeline(
|
||||
WORKSPACE_A,
|
||||
{'name': 'Created Pipeline', 'description': 'created'},
|
||||
)
|
||||
bot_uuid = await application.bot_service.create_bot(
|
||||
WORKSPACE_A,
|
||||
{
|
||||
'name': 'Created Bot',
|
||||
'description': 'created',
|
||||
'adapter': 'test',
|
||||
'adapter_config': {},
|
||||
'enable': False,
|
||||
'pipeline_routing_rules': [],
|
||||
},
|
||||
)
|
||||
model_uuid = await application.llm_model_service.create_llm_model(
|
||||
WORKSPACE_A,
|
||||
{
|
||||
'name': 'Created Model',
|
||||
'provider_uuid': 'provider-a',
|
||||
'abilities': [],
|
||||
'extra_args': {},
|
||||
'prefered_ranking': 0,
|
||||
},
|
||||
auto_set_to_default_pipeline=False,
|
||||
)
|
||||
|
||||
async with engine.connect() as connection:
|
||||
assert (
|
||||
await connection.scalar(
|
||||
sqlalchemy.select(ModelProvider.workspace_uuid).where(ModelProvider.uuid == provider_uuid)
|
||||
)
|
||||
== WORKSPACE_A
|
||||
)
|
||||
assert (
|
||||
await connection.scalar(
|
||||
sqlalchemy.select(LegacyPipeline.workspace_uuid).where(LegacyPipeline.uuid == pipeline_uuid)
|
||||
)
|
||||
== WORKSPACE_A
|
||||
)
|
||||
assert await connection.scalar(sqlalchemy.select(Bot.workspace_uuid).where(Bot.uuid == bot_uuid)) == WORKSPACE_A
|
||||
assert (
|
||||
await connection.scalar(sqlalchemy.select(LLMModel.workspace_uuid).where(LLMModel.uuid == model_uuid))
|
||||
== WORKSPACE_A
|
||||
)
|
||||
@@ -0,0 +1,74 @@
|
||||
from langbot.pkg.api.http import authz
|
||||
from langbot.pkg.api.http.context import PrincipalContext, PrincipalType, RequestContext, WorkspaceContext
|
||||
|
||||
|
||||
def _context(role: authz.WorkspaceRole) -> RequestContext:
|
||||
return RequestContext(
|
||||
instance_uuid='instance-test',
|
||||
placement_generation=1,
|
||||
request_id='request-test',
|
||||
auth_type='user-token',
|
||||
principal=PrincipalContext(
|
||||
principal_type=PrincipalType.ACCOUNT,
|
||||
account_uuid='account-test',
|
||||
),
|
||||
workspace=WorkspaceContext(
|
||||
workspace_uuid='workspace-test',
|
||||
membership_uuid='membership-test',
|
||||
role=role.value,
|
||||
permissions=authz.permissions_for_role(role),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def test_owner_has_every_fixed_permission():
|
||||
ctx = _context(authz.WorkspaceRole.OWNER)
|
||||
|
||||
assert ctx.workspace.permissions == frozenset(permission.value for permission in authz.Permission)
|
||||
|
||||
|
||||
def test_admin_cannot_transfer_owner_delete_workspace_or_link_billing():
|
||||
ctx = _context(authz.WorkspaceRole.ADMIN)
|
||||
|
||||
assert not authz.has_permission(ctx, authz.Permission.OWNER_TRANSFER)
|
||||
assert not authz.has_permission(ctx, authz.Permission.WORKSPACE_DELETE)
|
||||
assert not authz.has_permission(ctx, authz.Permission.BILLING_LINK_MANAGE)
|
||||
assert authz.has_permission(ctx, authz.Permission.MEMBER_INVITE)
|
||||
|
||||
|
||||
def test_operator_can_run_but_cannot_manage_resources_or_secrets():
|
||||
ctx = _context(authz.WorkspaceRole.OPERATOR)
|
||||
|
||||
assert authz.has_permission(ctx, authz.Permission.RUNTIME_OPERATE)
|
||||
assert not authz.has_permission(ctx, authz.Permission.RESOURCE_MANAGE)
|
||||
assert not authz.has_permission(ctx, authz.Permission.PROVIDER_SECRET_MANAGE)
|
||||
|
||||
|
||||
def test_unknown_role_has_no_permissions():
|
||||
assert authz.permissions_for_role('unknown') == frozenset()
|
||||
|
||||
|
||||
def test_require_permission_reports_stable_permission():
|
||||
ctx = _context(authz.WorkspaceRole.VIEWER)
|
||||
|
||||
try:
|
||||
authz.require_permission(ctx, authz.Permission.RESOURCE_MANAGE)
|
||||
except authz.PermissionDeniedError as exc:
|
||||
assert exc.permission == authz.Permission.RESOURCE_MANAGE.value
|
||||
assert exc.error_code == 'permission_denied'
|
||||
else:
|
||||
raise AssertionError('PermissionDeniedError was not raised')
|
||||
|
||||
|
||||
def test_execution_context_preserves_workspace_and_generation():
|
||||
from langbot.pkg.api.http.context import ExecutionContext
|
||||
|
||||
ctx = _context(authz.WorkspaceRole.DEVELOPER)
|
||||
execution = ExecutionContext.from_request(ctx, bot_uuid='bot-test', pipeline_uuid='pipeline-test')
|
||||
|
||||
assert execution.instance_uuid == 'instance-test'
|
||||
assert execution.workspace_uuid == 'workspace-test'
|
||||
assert execution.placement_generation == 1
|
||||
assert execution.bot_uuid == 'bot-test'
|
||||
assert execution.pipeline_uuid == 'pipeline-test'
|
||||
assert execution.trigger_principal == ctx.principal
|
||||
@@ -0,0 +1,39 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import quart
|
||||
|
||||
from langbot.pkg.api.http.controller import main as controller_main
|
||||
from langbot.pkg.utils import bounded_executor
|
||||
|
||||
|
||||
async def test_bounded_json_request_decodes_off_loop_in_workspace_scope(
|
||||
monkeypatch,
|
||||
):
|
||||
app = quart.Quart(__name__)
|
||||
app.request_class = controller_main.BoundedJSONRequest
|
||||
observed_scopes: list[str | None] = []
|
||||
|
||||
async def fake_to_thread(fn, *args, **kwargs):
|
||||
observed_scopes.append(bounded_executor.current_blocking_work_scope())
|
||||
return fn(*args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(
|
||||
controller_main.asyncio,
|
||||
'to_thread',
|
||||
fake_to_thread,
|
||||
)
|
||||
|
||||
@app.post('/json')
|
||||
async def parse_json():
|
||||
with bounded_executor.blocking_work_scope('workspace-a'):
|
||||
payload = await quart.request.get_json()
|
||||
return quart.jsonify(payload)
|
||||
|
||||
response = await app.test_client().post(
|
||||
'/json',
|
||||
json={'nested': {'value': 1}},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert await response.get_json() == {'nested': {'value': 1}}
|
||||
assert observed_scopes == ['workspace-a']
|
||||
@@ -0,0 +1,78 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import quart
|
||||
|
||||
from langbot.pkg.api.http.context import (
|
||||
ExecutionContext,
|
||||
PrincipalContext,
|
||||
PrincipalType,
|
||||
RequestContext,
|
||||
WorkspaceContext,
|
||||
)
|
||||
from langbot.pkg.api.http.controller.group import RouterGroup
|
||||
from langbot.pkg.cloud.entitlements import EntitlementSnapshot, EntitlementUnavailableError
|
||||
from langbot.pkg.cloud.entitlements import EntitlementResolver
|
||||
|
||||
|
||||
class _Group(RouterGroup):
|
||||
async def initialize(self) -> None:
|
||||
return None
|
||||
|
||||
|
||||
def _router(deployment) -> _Group:
|
||||
provider = getattr(deployment, 'entitlement_provider', None)
|
||||
resolver = EntitlementResolver('instance-a', provider) if provider is not None else None
|
||||
ap = SimpleNamespace(deployment=deployment, entitlement_resolver=resolver)
|
||||
return _Group(ap, quart.Quart(__name__))
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cloud_request_resolves_verified_entitlement_revision():
|
||||
snapshot = EntitlementSnapshot(
|
||||
instance_uuid='instance-a',
|
||||
workspace_uuid='workspace-a',
|
||||
entitlement_revision=9,
|
||||
status='active',
|
||||
not_before=1,
|
||||
expires_at=4_000_000_000,
|
||||
features={},
|
||||
limits={},
|
||||
)
|
||||
provider = SimpleNamespace(get_workspace_entitlement=AsyncMock(return_value=snapshot))
|
||||
router = _router(SimpleNamespace(multi_workspace_enabled=True, entitlement_provider=provider))
|
||||
|
||||
revision = await router._resolve_entitlement_revision('instance-a', 'workspace-a')
|
||||
|
||||
assert revision == 9
|
||||
provider.get_workspace_entitlement.assert_awaited_once_with('workspace-a')
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cloud_request_fails_closed_without_entitlement_provider():
|
||||
router = _router(SimpleNamespace(multi_workspace_enabled=True, entitlement_provider=None))
|
||||
|
||||
with pytest.raises(EntitlementUnavailableError):
|
||||
await router._resolve_entitlement_revision('instance-a', 'workspace-a')
|
||||
|
||||
|
||||
def test_execution_context_preserves_entitlement_revision():
|
||||
request = RequestContext(
|
||||
instance_uuid='instance-a',
|
||||
placement_generation=1,
|
||||
request_id='request-a',
|
||||
auth_type='user-token',
|
||||
principal=PrincipalContext(PrincipalType.ACCOUNT, account_uuid='account-a'),
|
||||
workspace=WorkspaceContext(
|
||||
workspace_uuid='workspace-a',
|
||||
membership_uuid='membership-a',
|
||||
role='owner',
|
||||
permissions=frozenset(),
|
||||
),
|
||||
entitlement_revision=11,
|
||||
)
|
||||
|
||||
assert ExecutionContext.from_request(request).entitlement_revision == 11
|
||||
@@ -0,0 +1,257 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
import quart
|
||||
|
||||
from langbot.pkg.api.http.controller import group
|
||||
from langbot.pkg.api.http.controller.groups.webhooks import WebhookRouterGroup
|
||||
from langbot.pkg.utils.bounded_executor import (
|
||||
BlockingWorkCapacityError,
|
||||
current_blocking_work_scope,
|
||||
)
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
class _FailingRouterGroup(group.RouterGroup):
|
||||
name = 'failing-test'
|
||||
path = '/failing-test'
|
||||
|
||||
async def initialize(self) -> None:
|
||||
@self.route('', methods=['GET'], auth_type=group.AuthType.NONE)
|
||||
async def _():
|
||||
raise RuntimeError('database password=do-not-return')
|
||||
|
||||
|
||||
class _AuthenticatedRouterGroup(group.RouterGroup):
|
||||
name = 'authenticated-test'
|
||||
path = '/authenticated-test'
|
||||
|
||||
async def initialize(self) -> None:
|
||||
@self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
|
||||
async def _():
|
||||
return self.success()
|
||||
|
||||
|
||||
class _BlockingCapacityRouterGroup(group.RouterGroup):
|
||||
name = 'blocking-capacity-test'
|
||||
path = '/blocking-capacity-test'
|
||||
|
||||
async def initialize(self) -> None:
|
||||
@self.route('', methods=['GET'], auth_type=group.AuthType.NONE)
|
||||
async def _():
|
||||
raise BlockingWorkCapacityError('Workspace blocking executor capacity reached')
|
||||
|
||||
|
||||
class _InvalidAccountRouterGroup(group.RouterGroup):
|
||||
name = 'invalid-account-test'
|
||||
path = '/invalid-account-test'
|
||||
|
||||
async def initialize(self) -> None:
|
||||
@self.route(
|
||||
'',
|
||||
methods=['GET'],
|
||||
auth_type=group.AuthType.ACCOUNT_TOKEN,
|
||||
permission='workspace.view',
|
||||
)
|
||||
async def _():
|
||||
return self.success()
|
||||
|
||||
|
||||
async def test_unhandled_http_error_returns_generic_body_and_correlated_request_id():
|
||||
logger = Mock()
|
||||
application = SimpleNamespace(logger=logger)
|
||||
quart_app = quart.Quart(__name__)
|
||||
await _FailingRouterGroup(application, quart_app).initialize()
|
||||
|
||||
response = await quart_app.test_client().get(
|
||||
'/failing-test',
|
||||
headers={'X-Request-Id': 'request-http-test'},
|
||||
)
|
||||
|
||||
assert response.status_code == 500
|
||||
assert await response.get_json() == {
|
||||
'code': 'internal_error',
|
||||
'msg': 'Internal server error',
|
||||
'request_id': 'request-http-test',
|
||||
}
|
||||
assert response.headers['X-Request-Id'] == 'request-http-test'
|
||||
log_message = logger.error.call_args.args[0]
|
||||
assert 'request_id=request-http-test' in log_message
|
||||
assert 'database password=do-not-return' in log_message
|
||||
assert 'do-not-return' not in (await response.get_data(as_text=True))
|
||||
|
||||
|
||||
async def test_public_webhook_error_uses_same_generic_error_contract():
|
||||
logger = Mock()
|
||||
application = SimpleNamespace(
|
||||
logger=logger,
|
||||
platform_mgr=SimpleNamespace(
|
||||
resolve_public_bot=AsyncMock(side_effect=RuntimeError('adapter credential=do-not-return'))
|
||||
),
|
||||
)
|
||||
quart_app = quart.Quart(__name__)
|
||||
await WebhookRouterGroup(application, quart_app).initialize()
|
||||
|
||||
response = await quart_app.test_client().post(
|
||||
'/bots/11111111-1111-4111-8111-111111111111',
|
||||
headers={'X-Request-Id': 'request-webhook-test'},
|
||||
)
|
||||
|
||||
assert response.status_code == 500
|
||||
assert await response.get_json() == {
|
||||
'code': 'internal_error',
|
||||
'msg': 'Internal server error',
|
||||
'request_id': 'request-webhook-test',
|
||||
}
|
||||
assert response.headers['X-Request-Id'] == 'request-webhook-test'
|
||||
log_message = logger.error.call_args.args[0]
|
||||
assert 'request_id=request-webhook-test' in log_message
|
||||
assert 'adapter credential=do-not-return' in log_message
|
||||
assert 'do-not-return' not in (await response.get_data(as_text=True))
|
||||
|
||||
|
||||
async def test_blocking_work_capacity_maps_to_retryable_http_response():
|
||||
application = SimpleNamespace(logger=Mock())
|
||||
quart_app = quart.Quart(__name__)
|
||||
await _BlockingCapacityRouterGroup(application, quart_app).initialize()
|
||||
|
||||
response = await quart_app.test_client().get('/blocking-capacity-test')
|
||||
|
||||
assert response.status_code == 429
|
||||
assert await response.get_json() == {
|
||||
'code': 'blocking_work_capacity_exceeded',
|
||||
'msg': 'Workspace blocking executor capacity reached',
|
||||
}
|
||||
|
||||
|
||||
async def test_public_webhook_carries_scope_without_holding_database_session():
|
||||
class ScopeOnlyPersistenceManager:
|
||||
mode = SimpleNamespace(value='cloud_runtime')
|
||||
|
||||
def __init__(self):
|
||||
self.active_workspace = None
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def tenant_scope(self, workspace_uuid):
|
||||
self.active_workspace = workspace_uuid
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self.active_workspace = None
|
||||
|
||||
def current_session(self):
|
||||
return None
|
||||
|
||||
persistence_mgr = ScopeOnlyPersistenceManager()
|
||||
workspace_uuid = '00000000-0000-0000-0000-00000000000a'
|
||||
bot_uuid = '11111111-1111-4111-8111-111111111111'
|
||||
|
||||
class Adapter:
|
||||
async def handle_unified_webhook(self, **_kwargs):
|
||||
assert persistence_mgr.active_workspace == workspace_uuid
|
||||
assert persistence_mgr.current_session() is None
|
||||
assert current_blocking_work_scope() == workspace_uuid
|
||||
return {'ok': True}
|
||||
|
||||
async def get_execution_binding(resolved_workspace_uuid, expected_generation=None):
|
||||
assert resolved_workspace_uuid == workspace_uuid
|
||||
assert expected_generation == 4
|
||||
|
||||
runtime_bot = SimpleNamespace(
|
||||
workspace_uuid=workspace_uuid,
|
||||
placement_generation=4,
|
||||
enable=True,
|
||||
adapter=Adapter(),
|
||||
)
|
||||
application = SimpleNamespace(
|
||||
logger=Mock(),
|
||||
persistence_mgr=persistence_mgr,
|
||||
platform_mgr=SimpleNamespace(resolve_public_bot=AsyncMock(return_value=runtime_bot)),
|
||||
workspace_service=SimpleNamespace(get_execution_binding=get_execution_binding),
|
||||
)
|
||||
quart_app = quart.Quart(__name__)
|
||||
await WebhookRouterGroup(application, quart_app).initialize()
|
||||
|
||||
response = await quart_app.test_client().post(f'/bots/{bot_uuid}')
|
||||
|
||||
assert response.status_code == 200
|
||||
assert await response.get_json() == {'ok': True}
|
||||
assert persistence_mgr.active_workspace is None
|
||||
|
||||
|
||||
async def test_public_webhook_blocking_capacity_is_retryable():
|
||||
workspace_uuid = '00000000-0000-0000-0000-00000000000a'
|
||||
bot_uuid = '11111111-1111-4111-8111-111111111111'
|
||||
|
||||
class Adapter:
|
||||
async def handle_unified_webhook(self, **_kwargs):
|
||||
raise BlockingWorkCapacityError(
|
||||
'Workspace blocking executor capacity reached',
|
||||
scope=workspace_uuid,
|
||||
)
|
||||
|
||||
runtime_bot = SimpleNamespace(
|
||||
workspace_uuid=workspace_uuid,
|
||||
placement_generation=4,
|
||||
enable=True,
|
||||
adapter=Adapter(),
|
||||
)
|
||||
application = SimpleNamespace(
|
||||
logger=Mock(),
|
||||
persistence_mgr=SimpleNamespace(mode=SimpleNamespace(value='oss')),
|
||||
platform_mgr=SimpleNamespace(resolve_public_bot=AsyncMock(return_value=runtime_bot)),
|
||||
workspace_service=SimpleNamespace(get_execution_binding=AsyncMock(return_value=None)),
|
||||
)
|
||||
quart_app = quart.Quart(__name__)
|
||||
await WebhookRouterGroup(application, quart_app).initialize()
|
||||
|
||||
response = await quart_app.test_client().post(f'/bots/{bot_uuid}')
|
||||
|
||||
assert response.status_code == 429
|
||||
assert await response.get_json() == {
|
||||
'code': 'blocking_work_capacity_exceeded',
|
||||
'msg': 'Workspace blocking executor capacity reached',
|
||||
}
|
||||
|
||||
|
||||
async def test_authentication_failure_does_not_return_internal_exception_text():
|
||||
logger = Mock()
|
||||
application = SimpleNamespace(
|
||||
logger=logger,
|
||||
user_service=SimpleNamespace(
|
||||
get_authenticated_account=AsyncMock(side_effect=RuntimeError('database password=do-not-return'))
|
||||
),
|
||||
)
|
||||
quart_app = quart.Quart(__name__)
|
||||
await _AuthenticatedRouterGroup(application, quart_app).initialize()
|
||||
|
||||
response = await quart_app.test_client().get(
|
||||
'/authenticated-test',
|
||||
headers={
|
||||
'Authorization': 'Bearer invalid',
|
||||
'X-Request-Id': 'request-auth-test',
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 401
|
||||
assert await response.get_json() == {
|
||||
'code': 'invalid_authentication',
|
||||
'msg': 'Invalid authentication credentials',
|
||||
}
|
||||
assert 'do-not-return' not in (await response.get_data(as_text=True))
|
||||
assert 'request_id=request-auth-test' in logger.warning.call_args.args[0]
|
||||
assert 'database password=do-not-return' in logger.warning.call_args.args[0]
|
||||
|
||||
|
||||
async def test_account_token_route_cannot_declare_workspace_permission():
|
||||
application = SimpleNamespace(logger=Mock())
|
||||
quart_app = quart.Quart(__name__)
|
||||
|
||||
with pytest.raises(ValueError, match='cannot declare Workspace permissions'):
|
||||
await _InvalidAccountRouterGroup(application, quart_app).initialize()
|
||||
@@ -0,0 +1,73 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
import quart
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
from langbot.pkg.api.http.controller import group
|
||||
from langbot.pkg.persistence.mgr import PersistenceManager, PersistenceMode
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
async def test_authenticated_route_does_not_hold_database_session_during_external_wait():
|
||||
entered = asyncio.Event()
|
||||
release = asyncio.Event()
|
||||
observations: list[bool] = []
|
||||
engine = create_async_engine('sqlite+aiosqlite:///:memory:')
|
||||
persistence = PersistenceManager(object(), mode=PersistenceMode.CLOUD_RUNTIME)
|
||||
persistence.db = SimpleNamespace(get_engine=lambda: engine)
|
||||
|
||||
class BlockingRouter(group.RouterGroup):
|
||||
name = 'blocking-route-test'
|
||||
path = '/blocking-route-test'
|
||||
|
||||
async def initialize(self) -> None:
|
||||
@self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
|
||||
async def _():
|
||||
observations.append(persistence.current_session() is None)
|
||||
entered.set()
|
||||
await release.wait()
|
||||
observations.append(persistence.current_session() is None)
|
||||
return self.success(data={})
|
||||
|
||||
account = SimpleNamespace(uuid='account-a', user='owner@example.com')
|
||||
access = SimpleNamespace(
|
||||
execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=1),
|
||||
workspace=SimpleNamespace(uuid='workspace-a'),
|
||||
membership=SimpleNamespace(uuid='membership-a', role='owner', projection_revision=1),
|
||||
)
|
||||
application = SimpleNamespace(
|
||||
persistence_mgr=persistence,
|
||||
deployment=SimpleNamespace(multi_workspace_enabled=False),
|
||||
user_service=SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account)),
|
||||
workspace_collaboration_service=SimpleNamespace(resolve_account_workspace=AsyncMock(return_value=access)),
|
||||
logger=Mock(),
|
||||
)
|
||||
quart_app = quart.Quart(__name__)
|
||||
await BlockingRouter(application, quart_app).initialize()
|
||||
client = quart_app.test_client()
|
||||
|
||||
request = asyncio.create_task(
|
||||
client.get(
|
||||
'/blocking-route-test',
|
||||
headers={'Authorization': 'Bearer token', 'X-Workspace-Id': 'workspace-a'},
|
||||
)
|
||||
)
|
||||
try:
|
||||
await entered.wait()
|
||||
assert observations == [True]
|
||||
release.set()
|
||||
response = await request
|
||||
assert response.status_code == 200
|
||||
assert observations == [True, True]
|
||||
finally:
|
||||
release.set()
|
||||
if not request.done():
|
||||
await request
|
||||
await engine.dispose()
|
||||
@@ -1,482 +1,466 @@
|
||||
"""
|
||||
Unit tests for ApiKeyService.
|
||||
|
||||
Tests API key CRUD operations with mocked persistence layer.
|
||||
|
||||
Source: src/langbot/pkg/api/http/service/apikey.py
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, Mock, patch
|
||||
import datetime
|
||||
import hashlib
|
||||
import logging
|
||||
import uuid
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import sqlalchemy
|
||||
from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
|
||||
|
||||
from langbot.pkg.api.http.authz import Permission, PermissionDeniedError
|
||||
from langbot.pkg.api.http.context import PrincipalContext, PrincipalType, RequestContext, WorkspaceContext
|
||||
from langbot.pkg.api.http.service.apikey import ApiKeyService
|
||||
from langbot.pkg.entity.persistence.apikey import ApiKey
|
||||
from langbot.pkg.entity.persistence.base import Base
|
||||
from langbot.pkg.entity.persistence.user import User
|
||||
from langbot.pkg.entity.persistence.workspace import (
|
||||
Workspace,
|
||||
WorkspaceExecutionSource,
|
||||
WorkspaceExecutionState,
|
||||
WorkspaceSource,
|
||||
)
|
||||
from langbot.pkg.workspace.policy import SingleWorkspacePolicy
|
||||
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
|
||||
from langbot.pkg.workspace.service import WorkspaceService
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
class _PersistenceManager:
|
||||
def __init__(self, engine):
|
||||
self.engine = engine
|
||||
|
||||
def get_db_engine(self):
|
||||
return self.engine
|
||||
|
||||
async def execute_async(self, *args, **kwargs):
|
||||
async with self.engine.connect() as connection:
|
||||
result = await connection.execute(*args, **kwargs)
|
||||
await connection.commit()
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def serialize_model(model, row, masked_columns=()):
|
||||
return {
|
||||
column.name: (
|
||||
getattr(row, column.name).isoformat()
|
||||
if isinstance(getattr(row, column.name), datetime.datetime)
|
||||
else getattr(row, column.name)
|
||||
)
|
||||
for column in model.__table__.columns
|
||||
if column.name not in masked_columns
|
||||
}
|
||||
|
||||
|
||||
def _context(workspace_uuid: str, account_uuid: str, permissions: set[Permission]) -> RequestContext:
|
||||
return RequestContext(
|
||||
instance_uuid='api-key-instance',
|
||||
placement_generation=1,
|
||||
request_id=str(uuid.uuid4()),
|
||||
auth_type='user-token',
|
||||
principal=PrincipalContext(PrincipalType.ACCOUNT, account_uuid=account_uuid),
|
||||
workspace=WorkspaceContext(
|
||||
workspace_uuid=workspace_uuid,
|
||||
membership_uuid=str(uuid.uuid4()),
|
||||
role='owner',
|
||||
permissions=frozenset(permission.value for permission in permissions),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def api_key_context(tmp_path):
|
||||
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "api-keys.db"}')
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
|
||||
application = SimpleNamespace(
|
||||
persistence_mgr=_PersistenceManager(engine),
|
||||
instance_config=SimpleNamespace(data={'api': {'global_api_key': ''}}),
|
||||
logger=logging.getLogger('api-key-test'),
|
||||
)
|
||||
application.workspace_service = WorkspaceService(application, instance_uuid='api-key-instance')
|
||||
workspace = await application.workspace_service.ensure_singleton_workspace()
|
||||
account_uuid = str(uuid.uuid4())
|
||||
session_factory = async_sessionmaker(engine, expire_on_commit=False)
|
||||
async with session_factory.begin() as session:
|
||||
session.add(
|
||||
User(
|
||||
uuid=account_uuid,
|
||||
user='owner@example.com',
|
||||
normalized_email='owner@example.com',
|
||||
password='hash',
|
||||
account_type='local',
|
||||
)
|
||||
)
|
||||
service = ApiKeyService(application)
|
||||
context = _context(workspace.uuid, account_uuid, set(Permission))
|
||||
yield application, service, context, engine
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def test_secret_is_returned_once_and_only_hash_is_persisted(api_key_context):
|
||||
_application, service, context, engine = api_key_context
|
||||
|
||||
created = await service.create_api_key(context, 'Automation', 'CI key')
|
||||
secret = created['key']
|
||||
assert secret.startswith('lbk_')
|
||||
assert created['secret_available'] is True
|
||||
assert 'key_hash' not in created
|
||||
|
||||
listed = await service.get_api_keys(context)
|
||||
assert len(listed) == 1
|
||||
assert 'key' not in listed[0]
|
||||
assert 'key_hash' not in listed[0]
|
||||
assert listed[0]['secret_available'] is False
|
||||
|
||||
async with engine.connect() as connection:
|
||||
stored = await connection.scalar(sqlalchemy.select(ApiKey.key_hash))
|
||||
assert stored == hashlib.sha256(secret.encode()).hexdigest()
|
||||
assert secret not in stored
|
||||
|
||||
|
||||
async def test_authentication_derives_workspace_scopes_and_updates_usage(api_key_context):
|
||||
_application, service, context, engine = api_key_context
|
||||
created = await service.create_api_key(
|
||||
context,
|
||||
'Read only',
|
||||
scopes=[Permission.RESOURCE_VIEW.value],
|
||||
)
|
||||
|
||||
identity = await service.authenticate_api_key(created['key'])
|
||||
assert identity is not None
|
||||
assert identity.workspace_uuid == context.workspace_uuid
|
||||
assert identity.permissions == frozenset({Permission.RESOURCE_VIEW.value})
|
||||
|
||||
async with engine.connect() as connection:
|
||||
last_used_at = await connection.scalar(sqlalchemy.select(ApiKey.last_used_at))
|
||||
assert last_used_at is not None
|
||||
|
||||
|
||||
async def test_revoked_expired_and_unknown_keys_fail_closed(api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
created = await service.create_api_key(context, 'Revocable')
|
||||
await service.delete_api_key(context, created['id'])
|
||||
assert await service.authenticate_api_key(created['key']) is None
|
||||
assert await service.verify_api_key('') is False
|
||||
assert await service.verify_api_key('plain-secret') is False
|
||||
assert await service.verify_api_key('lbk_unknown') is False
|
||||
|
||||
expired_secret = 'lbk_expired'
|
||||
await service.ap.persistence_mgr.execute_async(
|
||||
sqlalchemy.insert(ApiKey).values(
|
||||
workspace_uuid=context.workspace_uuid,
|
||||
name='Expired',
|
||||
key_hash=hashlib.sha256(expired_secret.encode()).hexdigest(),
|
||||
scopes=[Permission.RESOURCE_VIEW.value],
|
||||
status='active',
|
||||
expires_at=datetime.datetime.now(datetime.UTC).replace(tzinfo=None) - datetime.timedelta(seconds=1),
|
||||
)
|
||||
)
|
||||
assert await service.authenticate_api_key(expired_secret) is None
|
||||
|
||||
|
||||
async def test_revoke_winning_last_used_update_race_fails_authentication(api_key_context):
|
||||
application, service, context, _engine = api_key_context
|
||||
created = await service.create_api_key(context, 'Racing revoke')
|
||||
original_execute = application.persistence_mgr.execute_async
|
||||
injected_revoke = False
|
||||
|
||||
async def execute_with_revoke(statement, *args, **kwargs):
|
||||
nonlocal injected_revoke
|
||||
if (
|
||||
not injected_revoke
|
||||
and isinstance(statement, sqlalchemy.sql.dml.Update)
|
||||
and statement.table.name == ApiKey.__tablename__
|
||||
):
|
||||
injected_revoke = True
|
||||
await original_execute(sqlalchemy.update(ApiKey).where(ApiKey.id == created['id']).values(status='revoked'))
|
||||
return await original_execute(statement, *args, **kwargs)
|
||||
|
||||
application.persistence_mgr.execute_async = execute_with_revoke
|
||||
|
||||
assert await service.authenticate_api_key(created['key']) is None
|
||||
assert injected_revoke is True
|
||||
|
||||
|
||||
async def test_cross_workspace_crud_and_secret_guessing_are_isolated(api_key_context):
|
||||
application, service, first_context, engine = api_key_context
|
||||
second_workspace_uuid = str(uuid.uuid4())
|
||||
async with async_sessionmaker(engine, expire_on_commit=False).begin() as session:
|
||||
session.add(
|
||||
Workspace(
|
||||
uuid=second_workspace_uuid,
|
||||
instance_uuid='api-key-instance',
|
||||
name='Second',
|
||||
slug='second',
|
||||
source=WorkspaceSource.CLOUD_PROJECTION.value,
|
||||
)
|
||||
)
|
||||
session.add(
|
||||
WorkspaceExecutionState(
|
||||
workspace_uuid=second_workspace_uuid,
|
||||
instance_uuid='api-key-instance',
|
||||
active_generation=3,
|
||||
state='active',
|
||||
write_fenced=False,
|
||||
source=WorkspaceExecutionSource.CLOUD.value,
|
||||
)
|
||||
)
|
||||
second_context = _context(second_workspace_uuid, first_context.account_uuid or '', set(Permission))
|
||||
created = await service.create_api_key(first_context, 'First only')
|
||||
|
||||
assert await service.get_api_key(second_context, created['id']) is None
|
||||
assert await service.get_api_keys(second_context) == []
|
||||
identity = await service.authenticate_api_key(created['key'])
|
||||
assert identity is not None
|
||||
assert identity.workspace_uuid == first_context.workspace_uuid
|
||||
assert identity.workspace_uuid != second_workspace_uuid
|
||||
|
||||
# Prove the explicit multi-Workspace policy does not change key-derived routing.
|
||||
application.workspace_service.policy = SingleWorkspacePolicy(workspace_limit=10, multi_workspace_enabled=True)
|
||||
identity = await service.authenticate_api_key(created['key'])
|
||||
assert identity is not None
|
||||
assert identity.workspace_uuid == first_context.workspace_uuid
|
||||
|
||||
|
||||
async def test_global_config_key_is_oss_singleton_only(api_key_context):
|
||||
application, service, _context_value, _engine = api_key_context
|
||||
application.instance_config.data['api']['global_api_key'] = 'configured-secret'
|
||||
|
||||
identity = await service.authenticate_api_key('configured-secret')
|
||||
assert identity is not None
|
||||
assert identity.api_key_uuid == 'global-oss-api-key'
|
||||
|
||||
application.workspace_service.policy = SingleWorkspacePolicy(workspace_limit=10, multi_workspace_enabled=True)
|
||||
assert await service.authenticate_api_key('configured-secret') is None
|
||||
|
||||
|
||||
async def test_explicit_scopes_cannot_exceed_callers_workspace_permissions(api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
limited_context = _context(
|
||||
context.workspace_uuid,
|
||||
context.account_uuid or '',
|
||||
{Permission.API_KEY_MANAGE, Permission.RESOURCE_VIEW},
|
||||
)
|
||||
|
||||
created = await service.create_api_key(
|
||||
limited_context,
|
||||
'Read only',
|
||||
scopes=[Permission.RESOURCE_VIEW.value],
|
||||
)
|
||||
identity = await service.authenticate_api_key(created['key'])
|
||||
assert identity is not None
|
||||
assert identity.permissions == frozenset({Permission.RESOURCE_VIEW.value})
|
||||
|
||||
with pytest.raises(PermissionDeniedError) as exc_info:
|
||||
await service.create_api_key(
|
||||
limited_context,
|
||||
'Escalated',
|
||||
scopes=[Permission.WORKSPACE_DELETE.value],
|
||||
)
|
||||
assert exc_info.value.permission == Permission.WORKSPACE_DELETE.value
|
||||
|
||||
|
||||
# Preserve the pre-tenancy CRUD and verification regression matrix while
|
||||
# exercising it through the new Workspace-bound API. The assertions reflect
|
||||
# intentional security changes: secrets are returned once, deletion revokes,
|
||||
# and missing Workspace resources are reported as not found.
|
||||
class TestApiKeyServiceGetApiKeys:
|
||||
"""Tests for get_api_keys method."""
|
||||
async def test_get_api_keys_empty_list(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
|
||||
async def test_get_api_keys_empty_list(self):
|
||||
"""Returns empty list when no API keys exist."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
mock_result = Mock()
|
||||
mock_result.all = Mock(return_value=[])
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
|
||||
ap.persistence_mgr.serialize_model = Mock(
|
||||
side_effect=lambda model_cls, entity: {
|
||||
'id': entity.id,
|
||||
'name': entity.name,
|
||||
'key': entity.key,
|
||||
'description': entity.description,
|
||||
}
|
||||
if entity
|
||||
else {}
|
||||
)
|
||||
assert await service.get_api_keys(context) == []
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
async def test_get_api_keys_returns_serialized_list(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
await service.create_api_key(context, 'Test Key 1', 'First test key')
|
||||
await service.create_api_key(context, 'Test Key 2', 'Second test key')
|
||||
|
||||
# Execute
|
||||
result = await service.get_api_keys()
|
||||
result = await service.get_api_keys(context)
|
||||
|
||||
# Verify
|
||||
assert result == []
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
|
||||
async def test_get_api_keys_returns_serialized_list(self):
|
||||
"""Returns serialized list of API keys."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
|
||||
# Create mock API key entities
|
||||
key1 = Mock(spec=ApiKey)
|
||||
key1.id = 1
|
||||
key1.name = 'Test Key 1'
|
||||
key1.key = 'lbk_test_key_1'
|
||||
key1.description = 'First test key'
|
||||
|
||||
key2 = Mock(spec=ApiKey)
|
||||
key2.id = 2
|
||||
key2.name = 'Test Key 2'
|
||||
key2.key = 'lbk_test_key_2'
|
||||
key2.description = 'Second test key'
|
||||
|
||||
mock_result = Mock()
|
||||
mock_result.all = Mock(return_value=[key1, key2])
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
|
||||
ap.persistence_mgr.serialize_model = Mock(
|
||||
side_effect=lambda model_cls, entity: {
|
||||
'id': entity.id,
|
||||
'name': entity.name,
|
||||
'key': entity.key,
|
||||
'description': entity.description,
|
||||
}
|
||||
)
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_api_keys()
|
||||
|
||||
# Verify
|
||||
assert len(result) == 2
|
||||
assert result[0]['name'] == 'Test Key 1'
|
||||
assert result[1]['name'] == 'Test Key 2'
|
||||
assert [item['name'] for item in result] == ['Test Key 1', 'Test Key 2']
|
||||
assert [item['description'] for item in result] == ['First test key', 'Second test key']
|
||||
assert all('key' not in item and 'key_hash' not in item for item in result)
|
||||
|
||||
|
||||
class TestApiKeyServiceCreateApiKey:
|
||||
"""Tests for create_api_key method."""
|
||||
async def test_create_api_key_generates_key_with_prefix(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
|
||||
async def test_create_api_key_generates_key_with_prefix(self):
|
||||
"""Creates API key with 'lbk_' prefix."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
with pytest.MonkeyPatch.context() as monkeypatch:
|
||||
monkeypatch.setattr(
|
||||
'langbot.pkg.api.http.service.apikey.secrets.token_urlsafe', lambda _size: 'fixed-token'
|
||||
)
|
||||
result = await service.create_api_key(context, 'New Key', 'Test description')
|
||||
|
||||
created_key = Mock(spec=ApiKey)
|
||||
created_key.id = 1
|
||||
created_key.name = 'New Key'
|
||||
created_key.key = 'lbk_fixed-token'
|
||||
created_key.description = 'Test description'
|
||||
select_result = Mock()
|
||||
select_result.first = Mock(return_value=created_key)
|
||||
insert_params = []
|
||||
|
||||
async def mock_execute(query):
|
||||
params = query.compile().params
|
||||
if {'name', 'key', 'description'}.issubset(params):
|
||||
insert_params.append(params)
|
||||
return Mock()
|
||||
return select_result
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
ap.persistence_mgr.serialize_model = Mock(
|
||||
side_effect=lambda model_cls, entity: {
|
||||
'id': 1,
|
||||
'name': entity.name,
|
||||
'key': entity.key,
|
||||
'description': entity.description,
|
||||
}
|
||||
)
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
with patch('langbot.pkg.api.http.service.apikey.secrets.token_urlsafe', return_value='fixed-token'):
|
||||
result = await service.create_api_key('New Key', 'Test description')
|
||||
|
||||
assert insert_params == [{'name': 'New Key', 'key': 'lbk_fixed-token', 'description': 'Test description'}]
|
||||
assert result['key'].startswith('lbk_')
|
||||
assert result['key'] == 'lbk_fixed-token'
|
||||
assert result['name'] == 'New Key'
|
||||
assert result['description'] == 'Test description'
|
||||
assert result['secret_available'] is True
|
||||
|
||||
async def test_create_api_key_without_description(self):
|
||||
"""Creates API key with empty description when not provided."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
async def test_create_api_key_without_description(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
|
||||
created_key = Mock(spec=ApiKey)
|
||||
created_key.id = 1
|
||||
created_key.name = 'No Desc Key'
|
||||
created_key.key = 'lbk_no_desc_key'
|
||||
created_key.description = ''
|
||||
result = await service.create_api_key(context, 'No Desc Key')
|
||||
|
||||
select_result = Mock()
|
||||
select_result.first = Mock(return_value=created_key)
|
||||
insert_result = Mock()
|
||||
|
||||
async def mock_execute(query):
|
||||
if hasattr(query, 'values'):
|
||||
return insert_result
|
||||
return select_result
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
ap.persistence_mgr.serialize_model = Mock(
|
||||
return_value={
|
||||
'id': 1,
|
||||
'name': 'No Desc Key',
|
||||
'key': 'lbk_no_desc_key',
|
||||
'description': '',
|
||||
}
|
||||
)
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.create_api_key('No Desc Key')
|
||||
|
||||
# Verify
|
||||
assert result['description'] == ''
|
||||
|
||||
|
||||
class TestApiKeyServiceGetApiKey:
|
||||
"""Tests for get_api_key method."""
|
||||
async def test_get_api_key_by_id_found(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
created = await service.create_api_key(context, 'Found Key', 'Found')
|
||||
|
||||
async def test_get_api_key_by_id_found(self):
|
||||
"""Returns API key when found by ID."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
result = await service.get_api_key(context, created['id'])
|
||||
|
||||
key = Mock(spec=ApiKey)
|
||||
key.id = 1
|
||||
key.name = 'Found Key'
|
||||
key.key = 'lbk_found_key'
|
||||
key.description = 'Found'
|
||||
|
||||
mock_result = Mock()
|
||||
mock_result.first = Mock(return_value=key)
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
|
||||
ap.persistence_mgr.serialize_model = Mock(
|
||||
return_value={
|
||||
'id': 1,
|
||||
'name': 'Found Key',
|
||||
'key': 'lbk_found_key',
|
||||
'description': 'Found',
|
||||
}
|
||||
)
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_api_key(1)
|
||||
|
||||
# Verify
|
||||
assert result is not None
|
||||
assert result['id'] == 1
|
||||
assert result['id'] == created['id']
|
||||
assert result['name'] == 'Found Key'
|
||||
assert 'key' not in result and 'key_hash' not in result
|
||||
|
||||
async def test_get_api_key_by_id_not_found(self):
|
||||
"""Returns None when API key not found."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
async def test_get_api_key_by_id_not_found(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
|
||||
mock_result = Mock()
|
||||
mock_result.first = Mock(return_value=None)
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
|
||||
assert await service.get_api_key(context, 999) is None
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
async def test_get_api_key_by_id_zero(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
|
||||
# Execute
|
||||
result = await service.get_api_key(999)
|
||||
|
||||
# Verify
|
||||
assert result is None
|
||||
|
||||
async def test_get_api_key_by_id_zero(self):
|
||||
"""Handles ID=0 (edge case) correctly."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
|
||||
mock_result = Mock()
|
||||
mock_result.first = Mock(return_value=None)
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_api_key(0)
|
||||
|
||||
# Verify - should return None (no key with ID 0)
|
||||
assert result is None
|
||||
assert await service.get_api_key(context, 0) is None
|
||||
|
||||
|
||||
class TestApiKeyServiceVerifyApiKey:
|
||||
"""Tests for verify_api_key method."""
|
||||
async def test_verify_api_key_valid(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
created = await service.create_api_key(context, 'Valid')
|
||||
|
||||
@staticmethod
|
||||
def _make_ap(db_key=None, global_api_key=''):
|
||||
"""Build a mock Application with persistence + instance_config."""
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
mock_result = Mock()
|
||||
mock_result.first = Mock(return_value=db_key)
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
|
||||
ap.instance_config = SimpleNamespace(data={'api': {'global_api_key': global_api_key}})
|
||||
return ap
|
||||
assert await service.verify_api_key(created['key']) is True
|
||||
|
||||
async def test_verify_api_key_valid(self):
|
||||
"""Returns True for valid API key."""
|
||||
# Setup
|
||||
key = Mock(spec=ApiKey)
|
||||
ap = self._make_ap(db_key=key)
|
||||
async def test_verify_api_key_invalid(self, api_key_context):
|
||||
_application, service, _context, _engine = api_key_context
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
assert await service.verify_api_key('lbk_invalid_key') is False
|
||||
|
||||
# Execute
|
||||
result = await service.verify_api_key('lbk_valid_key')
|
||||
async def test_verify_api_key_empty_string(self, api_key_context):
|
||||
_application, service, _context, _engine = api_key_context
|
||||
|
||||
# Verify
|
||||
assert result is True
|
||||
assert await service.verify_api_key('') is False
|
||||
|
||||
async def test_verify_api_key_invalid(self):
|
||||
"""Returns False for invalid API key."""
|
||||
# Setup
|
||||
ap = self._make_ap(db_key=None)
|
||||
async def test_verify_api_key_unknown_key(self, api_key_context):
|
||||
_application, service, _context, _engine = api_key_context
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
assert await service.verify_api_key('unknown_key') is False
|
||||
|
||||
# Execute
|
||||
result = await service.verify_api_key('lbk_invalid_key')
|
||||
async def test_verify_global_api_key_match(self, api_key_context):
|
||||
application, service, context, _engine = api_key_context
|
||||
application.instance_config.data['api']['global_api_key'] = 'my-global-secret'
|
||||
|
||||
# Verify
|
||||
assert result is False
|
||||
identity = await service.authenticate_api_key('my-global-secret')
|
||||
|
||||
async def test_verify_api_key_empty_string(self):
|
||||
"""Returns False for empty key string."""
|
||||
# Setup
|
||||
ap = self._make_ap(db_key=None)
|
||||
assert identity is not None
|
||||
assert identity.workspace_uuid == context.workspace_uuid
|
||||
assert identity.api_key_uuid == 'global-oss-api-key'
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
async def test_verify_global_api_key_no_prefix_required(self, api_key_context):
|
||||
application, service, _context, _engine = api_key_context
|
||||
application.instance_config.data['api']['global_api_key'] = 'plainsecret123'
|
||||
|
||||
# Execute
|
||||
result = await service.verify_api_key('')
|
||||
assert await service.verify_api_key('plainsecret123') is True
|
||||
|
||||
# Verify
|
||||
assert result is False
|
||||
async def test_verify_global_api_key_mismatch_falls_back_to_db(self, api_key_context):
|
||||
application, service, context, _engine = api_key_context
|
||||
application.instance_config.data['api']['global_api_key'] = 'my-global-secret'
|
||||
created = await service.create_api_key(context, 'DB key')
|
||||
|
||||
async def test_verify_api_key_unknown_key(self):
|
||||
"""Returns False when the key is not present in persistence."""
|
||||
# Setup
|
||||
ap = self._make_ap(db_key=None)
|
||||
identity = await service.authenticate_api_key(created['key'])
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
assert identity is not None
|
||||
assert identity.api_key_uuid == created['uuid']
|
||||
|
||||
# Execute
|
||||
result = await service.verify_api_key('unknown_key')
|
||||
async def test_verify_empty_global_api_key_disabled(self, api_key_context):
|
||||
application, service, _context, _engine = api_key_context
|
||||
application.instance_config.data['api']['global_api_key'] = ''
|
||||
|
||||
# Verify
|
||||
assert result is False
|
||||
|
||||
async def test_verify_global_api_key_match(self):
|
||||
"""Returns True when key matches the config.yaml global API key (no DB lookup)."""
|
||||
# Setup: no DB record, but a global key is configured
|
||||
ap = self._make_ap(db_key=None, global_api_key='my-global-secret')
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.verify_api_key('my-global-secret')
|
||||
|
||||
# Verify: accepted purely on config match
|
||||
assert result is True
|
||||
# DB should not have been consulted for the global-key path
|
||||
ap.persistence_mgr.execute_async.assert_not_called()
|
||||
|
||||
async def test_verify_global_api_key_no_prefix_required(self):
|
||||
"""Global API key is accepted even without the lbk_ prefix."""
|
||||
ap = self._make_ap(db_key=None, global_api_key='plainsecret123')
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
result = await service.verify_api_key('plainsecret123')
|
||||
|
||||
assert result is True
|
||||
|
||||
async def test_verify_global_api_key_mismatch_falls_back_to_db(self):
|
||||
"""A non-matching key still falls through to the DB lookup."""
|
||||
# Global key set, but request uses a different lbk_ key that IS in DB
|
||||
key = Mock(spec=ApiKey)
|
||||
ap = self._make_ap(db_key=key, global_api_key='my-global-secret')
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
result = await service.verify_api_key('lbk_db_key')
|
||||
|
||||
assert result is True
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
|
||||
async def test_verify_empty_global_api_key_disabled(self):
|
||||
"""An empty global_api_key must never authenticate an empty/blank request."""
|
||||
ap = self._make_ap(db_key=None, global_api_key='')
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
# Empty request key is rejected, and a blank global key never matches
|
||||
assert await service.verify_api_key('') is False
|
||||
assert await service.verify_api_key(' ') is False
|
||||
|
||||
async def test_verify_api_key_missing_global_config_key(self):
|
||||
"""Works even when api.global_api_key is absent (existing installs)."""
|
||||
# instance_config without the global_api_key field at all
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
mock_result = Mock()
|
||||
mock_result.first = Mock(return_value=None)
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
|
||||
ap.instance_config = SimpleNamespace(data={'api': {}})
|
||||
async def test_verify_api_key_missing_global_config_key(self, api_key_context):
|
||||
application, service, _context, _engine = api_key_context
|
||||
application.instance_config.data = {'api': {}}
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
result = await service.verify_api_key('lbk_some_key')
|
||||
|
||||
assert result is False
|
||||
assert await service.verify_api_key('lbk_some_key') is False
|
||||
|
||||
|
||||
class TestApiKeyServiceDeleteApiKey:
|
||||
"""Tests for delete_api_key method."""
|
||||
async def test_delete_api_key_by_id(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
created = await service.create_api_key(context, 'Delete me')
|
||||
|
||||
async def test_delete_api_key_by_id(self):
|
||||
"""Deletes API key by ID."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
await service.delete_api_key(context, created['id'])
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
stored = await service.get_api_key(context, created['id'])
|
||||
assert stored is not None
|
||||
assert stored['status'] == 'revoked'
|
||||
assert await service.verify_api_key(created['key']) is False
|
||||
|
||||
# Execute
|
||||
await service.delete_api_key(1)
|
||||
async def test_delete_api_key_nonexistent_id(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
|
||||
# Verify - execute_async was called (delete operation)
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
|
||||
async def test_delete_api_key_nonexistent_id(self):
|
||||
"""Delete operation completes even for nonexistent ID (no error raised)."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
# Execute - should not raise error
|
||||
await service.delete_api_key(999)
|
||||
|
||||
# Verify - execute_async was called regardless
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
with pytest.raises(WorkspaceNotFoundError, match='API key not found'):
|
||||
await service.delete_api_key(context, 999)
|
||||
|
||||
|
||||
class TestApiKeyServiceUpdateApiKey:
|
||||
"""Tests for update_api_key method."""
|
||||
async def test_update_api_key_name_only(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
created = await service.create_api_key(context, 'Original', 'Description')
|
||||
|
||||
async def test_update_api_key_name_only(self):
|
||||
"""Updates only the name field."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
await service.update_api_key(context, created['id'], name='Updated Name')
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
stored = await service.get_api_key(context, created['id'])
|
||||
assert stored is not None
|
||||
assert stored['name'] == 'Updated Name'
|
||||
assert stored['description'] == 'Description'
|
||||
|
||||
# Execute
|
||||
await service.update_api_key(1, name='Updated Name')
|
||||
async def test_update_api_key_description_only(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
created = await service.create_api_key(context, 'Original', 'Description')
|
||||
|
||||
# Verify - execute_async was called with update
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
await service.update_api_key(context, created['id'], description='Updated description')
|
||||
|
||||
async def test_update_api_key_description_only(self):
|
||||
"""Updates only the description field."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
stored = await service.get_api_key(context, created['id'])
|
||||
assert stored is not None
|
||||
assert stored['name'] == 'Original'
|
||||
assert stored['description'] == 'Updated description'
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
async def test_update_api_key_both_fields(self, api_key_context):
|
||||
_application, service, context, _engine = api_key_context
|
||||
created = await service.create_api_key(context, 'Original', 'Description')
|
||||
|
||||
# Execute
|
||||
await service.update_api_key(1, description='Updated description')
|
||||
await service.update_api_key(
|
||||
context,
|
||||
created['id'],
|
||||
name='New Name',
|
||||
description='New description',
|
||||
)
|
||||
|
||||
# Verify
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
stored = await service.get_api_key(context, created['id'])
|
||||
assert stored is not None
|
||||
assert stored['name'] == 'New Name'
|
||||
assert stored['description'] == 'New description'
|
||||
|
||||
async def test_update_api_key_both_fields(self):
|
||||
"""Updates both name and description."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
async def test_update_api_key_no_fields(self, api_key_context):
|
||||
application, service, context, _engine = api_key_context
|
||||
created = await service.create_api_key(context, 'Original')
|
||||
original_execute = application.persistence_mgr.execute_async
|
||||
application.persistence_mgr.execute_async = AsyncMock(wraps=original_execute)
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
await service.update_api_key(context, created['id'])
|
||||
|
||||
# Execute
|
||||
await service.update_api_key(1, name='New Name', description='New description')
|
||||
|
||||
# Verify
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
|
||||
async def test_update_api_key_no_fields(self):
|
||||
"""Does nothing when no fields provided."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
|
||||
service = ApiKeyService(ap)
|
||||
|
||||
# Execute
|
||||
await service.update_api_key(1)
|
||||
|
||||
# Verify - no execute call since no update_data
|
||||
ap.persistence_mgr.execute_async.assert_not_called()
|
||||
application.persistence_mgr.execute_async.assert_not_awaited()
|
||||
|
||||
@@ -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,581 @@
|
||||
"""Unit tests for API knowledge service.
|
||||
|
||||
Tests cover:
|
||||
- Knowledge base CRUD operations
|
||||
- Capability checking
|
||||
- Knowledge engine discovery
|
||||
- File operations
|
||||
"""
|
||||
"""Tests for the tenant-aware knowledge service facade."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
from unittest.mock import Mock, AsyncMock
|
||||
from importlib import import_module
|
||||
|
||||
from langbot.pkg.api.http.authz import WorkspaceRequiredError
|
||||
from langbot.pkg.api.http.context import ExecutionContext
|
||||
from langbot.pkg.api.http.service.knowledge import KnowledgeService
|
||||
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
|
||||
|
||||
|
||||
def get_knowledge_service_module():
|
||||
"""Lazy import to avoid circular import issues."""
|
||||
return import_module('langbot.pkg.api.http.service.knowledge')
|
||||
CONTEXT = ExecutionContext(
|
||||
instance_uuid='instance-a',
|
||||
workspace_uuid='workspace-a',
|
||||
placement_generation=2,
|
||||
)
|
||||
|
||||
|
||||
def create_mock_app():
|
||||
"""Create mock Application for testing."""
|
||||
mock_app = Mock()
|
||||
mock_app.logger = Mock()
|
||||
mock_app.rag_mgr = AsyncMock()
|
||||
mock_app.persistence_mgr = AsyncMock()
|
||||
mock_app.persistence_mgr.execute_async = AsyncMock()
|
||||
mock_app.persistence_mgr.serialize_model = Mock(return_value={})
|
||||
mock_app.plugin_connector = AsyncMock()
|
||||
mock_app.plugin_connector.is_enable_plugin = True
|
||||
return mock_app
|
||||
class _Rows:
|
||||
def __init__(self, rows=()):
|
||||
self.rows = list(rows)
|
||||
|
||||
def all(self):
|
||||
return self.rows
|
||||
|
||||
def __iter__(self):
|
||||
return iter(self.rows)
|
||||
|
||||
|
||||
def _app():
|
||||
return SimpleNamespace(
|
||||
logger=Mock(),
|
||||
instance_config=SimpleNamespace(data={}),
|
||||
rag_mgr=SimpleNamespace(
|
||||
get_all_knowledge_base_details=AsyncMock(return_value=[]),
|
||||
get_knowledge_base_details=AsyncMock(return_value=None),
|
||||
create_knowledge_base=AsyncMock(),
|
||||
remove_knowledge_base_from_runtime=AsyncMock(),
|
||||
load_knowledge_base=AsyncMock(),
|
||||
get_knowledge_base_by_uuid=AsyncMock(return_value=None),
|
||||
delete_knowledge_base=AsyncMock(),
|
||||
),
|
||||
persistence_mgr=SimpleNamespace(
|
||||
execute_async=AsyncMock(return_value=_Rows()),
|
||||
serialize_model=Mock(return_value={}),
|
||||
),
|
||||
plugin_connector=SimpleNamespace(
|
||||
is_enable_plugin=True,
|
||||
require_workspace_context=AsyncMock(side_effect=lambda context: context),
|
||||
get_rag_creation_schema=AsyncMock(return_value={}),
|
||||
get_rag_retrieval_schema=AsyncMock(return_value={}),
|
||||
list_knowledge_engines=AsyncMock(return_value=[]),
|
||||
list_parsers=AsyncMock(return_value=[]),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_list_and_get_forward_explicit_context():
|
||||
app = _app()
|
||||
app.rag_mgr.get_all_knowledge_base_details.return_value = [{'uuid': 'kb-a'}]
|
||||
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb-a'}
|
||||
service = KnowledgeService(app)
|
||||
|
||||
assert await service.get_knowledge_bases(CONTEXT) == [{'uuid': 'kb-a'}]
|
||||
assert await service.get_knowledge_base(CONTEXT, 'kb-a') == {'uuid': 'kb-a'}
|
||||
app.rag_mgr.get_all_knowledge_base_details.assert_awaited_once_with(CONTEXT)
|
||||
app.rag_mgr.get_knowledge_base_details.assert_awaited_once_with(CONTEXT, 'kb-a')
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_none_context_fails_closed_before_plugin_or_manager_access():
|
||||
app = _app()
|
||||
service = KnowledgeService(app)
|
||||
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
await service.get_knowledge_bases(None)
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
await service.create_knowledge_base(None, {'knowledge_engine_plugin_id': 'author/engine'})
|
||||
app.plugin_connector.get_rag_creation_schema.assert_not_awaited()
|
||||
app.rag_mgr.create_knowledge_base.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_validates_schema_and_binds_context():
|
||||
app = _app()
|
||||
app.plugin_connector.get_rag_creation_schema.return_value = {
|
||||
'schema': [{'name': 'endpoint', 'label': {'en_US': 'Endpoint'}, 'required': True}]
|
||||
}
|
||||
app.rag_mgr.create_knowledge_base.return_value = SimpleNamespace(uuid='kb-created')
|
||||
service = KnowledgeService(app)
|
||||
|
||||
with pytest.raises(ValueError, match='Endpoint is required'):
|
||||
await service.create_knowledge_base(
|
||||
CONTEXT,
|
||||
{'knowledge_engine_plugin_id': 'author/engine'},
|
||||
)
|
||||
|
||||
result = await service.create_knowledge_base(
|
||||
CONTEXT,
|
||||
{
|
||||
'name': 'KB',
|
||||
'description': 'desc',
|
||||
'knowledge_engine_plugin_id': 'author/engine',
|
||||
'creation_settings': {'endpoint': 'https://example.invalid'},
|
||||
},
|
||||
)
|
||||
assert result == 'kb-created'
|
||||
app.rag_mgr.create_knowledge_base.assert_awaited_once_with(
|
||||
CONTEXT,
|
||||
name='KB',
|
||||
knowledge_engine_plugin_id='author/engine',
|
||||
creation_settings={'endpoint': 'https://example.invalid'},
|
||||
retrieval_settings={},
|
||||
description='desc',
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_create_enforces_workspace_knowledge_base_limit():
|
||||
app = _app()
|
||||
app.instance_config.data = {'system': {'limitation': {'max_knowledge_bases': 2}}}
|
||||
app.rag_mgr.get_all_knowledge_base_details.return_value = [{'uuid': 'kb-a'}, {'uuid': 'kb-b'}]
|
||||
service = KnowledgeService(app)
|
||||
|
||||
with pytest.raises(ValueError, match=r'Maximum number of knowledge bases \(2\) reached'):
|
||||
await service.create_knowledge_base(
|
||||
CONTEXT,
|
||||
{'knowledge_engine_plugin_id': 'author/engine'},
|
||||
)
|
||||
app.rag_mgr.create_knowledge_base.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_update_rejects_guessed_uuid_and_scopes_reload():
|
||||
app = _app()
|
||||
service = KnowledgeService(app)
|
||||
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await service.update_knowledge_base(CONTEXT, 'kb-other', {'name': 'stolen'})
|
||||
app.persistence_mgr.execute_async.assert_not_awaited()
|
||||
|
||||
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb-a', 'workspace_uuid': 'workspace-a'}
|
||||
await service.update_knowledge_base(CONTEXT, 'kb-a', {'name': 'updated', 'uuid': 'ignored'})
|
||||
app.rag_mgr.remove_knowledge_base_from_runtime.assert_awaited_once_with(CONTEXT, 'kb-a')
|
||||
app.rag_mgr.load_knowledge_base.assert_awaited_once_with(
|
||||
CONTEXT,
|
||||
{'uuid': 'kb-a', 'workspace_uuid': 'workspace-a'},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_retrieve_uses_execution_context():
|
||||
app = _app()
|
||||
entry = SimpleNamespace(model_dump=Mock(return_value={'id': 'entry-a'}))
|
||||
runtime_kb = SimpleNamespace(retrieve=AsyncMock(return_value=[entry]))
|
||||
app.rag_mgr.get_knowledge_base_by_uuid.return_value = runtime_kb
|
||||
service = KnowledgeService(app)
|
||||
|
||||
assert await service.retrieve_knowledge_base(CONTEXT, 'kb-a', 'query', {'top_k': 3}) == [{'id': 'entry-a'}]
|
||||
runtime_kb.retrieve.assert_awaited_once_with(CONTEXT, 'query', settings={'top_k': 3})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_runtime_retrieve_cross_workspace_uuid_is_not_found():
|
||||
app = _app()
|
||||
service = KnowledgeService(app)
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await service.retrieve_knowledge_base(CONTEXT, 'kb-other', 'query')
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_file_listing_checks_parent_knowledge_base_first():
|
||||
app = _app()
|
||||
service = KnowledgeService(app)
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await service.get_files_by_knowledge_base(CONTEXT, 'kb-other')
|
||||
app.persistence_mgr.execute_async.assert_not_awaited()
|
||||
|
||||
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb-a'}
|
||||
row = SimpleNamespace(uuid='file-a')
|
||||
app.persistence_mgr.execute_async.return_value = _Rows([row])
|
||||
app.persistence_mgr.serialize_model.return_value = {'uuid': 'file-a'}
|
||||
assert await service.get_files_by_knowledge_base(CONTEXT, 'kb-a') == [{'uuid': 'file-a'}]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_store_and_delete_file_require_runtime_parent_and_capability():
|
||||
app = _app()
|
||||
runtime_kb = SimpleNamespace(
|
||||
store_file=AsyncMock(return_value='task-a'),
|
||||
delete_file=AsyncMock(),
|
||||
)
|
||||
app.rag_mgr.get_knowledge_base_by_uuid.return_value = runtime_kb
|
||||
app.rag_mgr.get_knowledge_base_details.return_value = {'knowledge_engine': {'capabilities': ['doc_ingestion']}}
|
||||
service = KnowledgeService(app)
|
||||
|
||||
assert await service.store_file(CONTEXT, 'kb-a', 'upload.pdf', 'author/parser') == 'task-a'
|
||||
runtime_kb.store_file.assert_awaited_once_with(CONTEXT, 'upload.pdf', parser_plugin_id='author/parser')
|
||||
await service.delete_file(CONTEXT, 'kb-a', 'file-a')
|
||||
runtime_kb.delete_file.assert_awaited_once_with(CONTEXT, 'file-a')
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_delete_knowledge_base_rejects_cross_workspace_uuid():
|
||||
app = _app()
|
||||
service = KnowledgeService(app)
|
||||
with pytest.raises(WorkspaceNotFoundError):
|
||||
await service.delete_knowledge_base(CONTEXT, 'kb-other')
|
||||
app.rag_mgr.delete_knowledge_base.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_engine_and_parser_discovery_require_context_and_filter_results():
|
||||
app = _app()
|
||||
app.plugin_connector.list_knowledge_engines.return_value = [{'plugin_id': 'author/engine'}]
|
||||
app.plugin_connector.list_parsers.return_value = [
|
||||
{'id': 'text', 'supported_mime_types': ['text/plain']},
|
||||
{'id': 'pdf', 'supported_mime_types': ['application/pdf']},
|
||||
]
|
||||
service = KnowledgeService(app)
|
||||
|
||||
assert await service.list_knowledge_engines(CONTEXT) == [{'plugin_id': 'author/engine'}]
|
||||
assert await service.list_parsers(CONTEXT, 'application/pdf') == [
|
||||
{'id': 'pdf', 'supported_mime_types': ['application/pdf']}
|
||||
]
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
await service.list_parsers(None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_engine_discovery_rejects_connector_workspace_or_generation_mismatch():
|
||||
app = _app()
|
||||
app.plugin_connector.require_workspace_context.side_effect = WorkspaceNotFoundError('Plugin resource not found')
|
||||
service = KnowledgeService(app)
|
||||
|
||||
with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
|
||||
await service.list_knowledge_engines(CONTEXT)
|
||||
|
||||
app.plugin_connector.list_knowledge_engines.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_schema_validation_refences_before_second_runtime_call():
|
||||
app = _app()
|
||||
app.plugin_connector.require_workspace_context.side_effect = [
|
||||
CONTEXT,
|
||||
WorkspaceNotFoundError('Plugin resource not found'),
|
||||
]
|
||||
service = KnowledgeService(app)
|
||||
|
||||
with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
|
||||
await service.create_knowledge_base(
|
||||
CONTEXT,
|
||||
{'knowledge_engine_plugin_id': 'author/engine'},
|
||||
)
|
||||
|
||||
app.plugin_connector.get_rag_creation_schema.assert_awaited_once()
|
||||
app.plugin_connector.get_rag_retrieval_schema.assert_not_awaited()
|
||||
app.rag_mgr.create_knowledge_base.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_engine_schemas_are_context_gated_and_fail_soft_on_connector_error():
|
||||
app = _app()
|
||||
app.plugin_connector.get_rag_creation_schema.return_value = {'schema': ['creation']}
|
||||
app.plugin_connector.get_rag_retrieval_schema.side_effect = RuntimeError('offline')
|
||||
service = KnowledgeService(app)
|
||||
|
||||
assert await service.get_engine_creation_schema(CONTEXT, 'author/engine') == {'schema': ['creation']}
|
||||
assert await service.get_engine_retrieval_schema(CONTEXT, 'author/engine') == {}
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
await service.get_engine_creation_schema(None, 'author/engine')
|
||||
|
||||
|
||||
# Preserve the original service regression matrix with the new explicit
|
||||
# Workspace context. These intentionally overlap a few isolation-focused
|
||||
# tests above so legacy business behavior cannot disappear behind new guards.
|
||||
class TestKnowledgeServiceInit:
|
||||
"""Tests for KnowledgeService initialization."""
|
||||
|
||||
def test_init_stores_app_reference(self):
|
||||
"""Test that __init__ stores Application reference."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
app = _app()
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
service = KnowledgeService(app)
|
||||
|
||||
assert service.ap is mock_app
|
||||
assert service.ap is app
|
||||
|
||||
|
||||
class TestGetKnowledgeBases:
|
||||
"""Tests for get_knowledge_bases method."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_all_kb_details(self):
|
||||
"""Test that it returns all knowledge base details."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.rag_mgr.get_all_knowledge_base_details = AsyncMock(return_value=[{'uuid': 'kb1', 'name': 'KB1'}])
|
||||
app = _app()
|
||||
app.rag_mgr.get_all_knowledge_base_details.return_value = [{'uuid': 'kb1', 'name': 'KB1'}]
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.get_knowledge_bases()
|
||||
result = await KnowledgeService(app).get_knowledge_bases(CONTEXT)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0]['uuid'] == 'kb1'
|
||||
assert result == [{'uuid': 'kb1', 'name': 'KB1'}]
|
||||
app.rag_mgr.get_all_knowledge_base_details.assert_awaited_once_with(CONTEXT)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_empty_list_when_no_kbs(self):
|
||||
"""Test that it returns empty list when no knowledge bases."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.rag_mgr.get_all_knowledge_base_details = AsyncMock(return_value=[])
|
||||
app = _app()
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.get_knowledge_bases()
|
||||
|
||||
assert result == []
|
||||
assert await KnowledgeService(app).get_knowledge_bases(CONTEXT) == []
|
||||
|
||||
|
||||
class TestGetKnowledgeBase:
|
||||
"""Tests for get_knowledge_base method."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_kb_details_by_uuid(self):
|
||||
"""Test that it returns specific KB details."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value={'uuid': 'kb1', 'name': 'KB1'})
|
||||
app = _app()
|
||||
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb1', 'name': 'KB1'}
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.get_knowledge_base('kb1')
|
||||
result = await KnowledgeService(app).get_knowledge_base(CONTEXT, 'kb1')
|
||||
|
||||
assert result['uuid'] == 'kb1'
|
||||
assert result == {'uuid': 'kb1', 'name': 'KB1'}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_none_when_not_found(self):
|
||||
"""Test that it returns None when KB not found."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value=None)
|
||||
app = _app()
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.get_knowledge_base('nonexistent')
|
||||
|
||||
assert result is None
|
||||
assert await KnowledgeService(app).get_knowledge_base(CONTEXT, 'nonexistent') is None
|
||||
|
||||
|
||||
class TestCreateKnowledgeBase:
|
||||
"""Tests for create_knowledge_base method."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_creates_kb_with_required_fields(self):
|
||||
"""Test creating KB with required plugin ID."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_kb = Mock()
|
||||
mock_kb.uuid = 'new_kb_uuid'
|
||||
mock_app.rag_mgr.create_knowledge_base = AsyncMock(return_value=mock_kb)
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
app = _app()
|
||||
app.rag_mgr.create_knowledge_base.return_value = SimpleNamespace(uuid='new_kb_uuid')
|
||||
service = KnowledgeService(app)
|
||||
kb_data = {
|
||||
'name': 'Test KB',
|
||||
'knowledge_engine_plugin_id': 'author/engine',
|
||||
'description': 'Test description',
|
||||
}
|
||||
|
||||
result = await service.create_knowledge_base(kb_data)
|
||||
result = await service.create_knowledge_base(CONTEXT, kb_data)
|
||||
|
||||
assert result == 'new_kb_uuid'
|
||||
mock_app.rag_mgr.create_knowledge_base.assert_called_once()
|
||||
app.rag_mgr.create_knowledge_base.assert_awaited_once_with(
|
||||
CONTEXT,
|
||||
name='Test KB',
|
||||
knowledge_engine_plugin_id='author/engine',
|
||||
creation_settings={},
|
||||
retrieval_settings={},
|
||||
description='Test description',
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raises_when_missing_plugin_id(self):
|
||||
"""Test that ValueError is raised when plugin ID missing."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
app = _app()
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
with pytest.raises(ValueError, match='knowledge_engine_plugin_id is required'):
|
||||
await KnowledgeService(app).create_knowledge_base(CONTEXT, {'name': 'Test'})
|
||||
|
||||
with pytest.raises(ValueError) as exc_info:
|
||||
await service.create_knowledge_base({'name': 'Test'})
|
||||
|
||||
assert 'knowledge_engine_plugin_id is required' in str(exc_info.value)
|
||||
app.rag_mgr.create_knowledge_base.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_creates_with_default_name(self):
|
||||
"""Test that KB is created with default name if not provided."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_kb = Mock()
|
||||
mock_kb.uuid = 'new_kb_uuid'
|
||||
mock_app.rag_mgr.create_knowledge_base = AsyncMock(return_value=mock_kb)
|
||||
app = _app()
|
||||
app.rag_mgr.create_knowledge_base.return_value = SimpleNamespace(uuid='new_kb_uuid')
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
await KnowledgeService(app).create_knowledge_base(
|
||||
CONTEXT,
|
||||
{'knowledge_engine_plugin_id': 'author/engine'},
|
||||
)
|
||||
|
||||
await service.create_knowledge_base({'knowledge_engine_plugin_id': 'author/engine'})
|
||||
|
||||
# Check that default name 'Untitled' was used
|
||||
call_args = mock_app.rag_mgr.create_knowledge_base.call_args
|
||||
assert call_args.kwargs['name'] == 'Untitled'
|
||||
assert app.rag_mgr.create_knowledge_base.await_args.kwargs['name'] == 'Untitled'
|
||||
|
||||
|
||||
class TestUpdateKnowledgeBase:
|
||||
"""Tests for update_knowledge_base method."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_updates_mutable_fields_only(self):
|
||||
"""Test that only mutable fields are updated."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value={'uuid': 'kb1', 'name': 'Updated'})
|
||||
mock_app.rag_mgr.remove_knowledge_base_from_runtime = AsyncMock()
|
||||
mock_app.rag_mgr.load_knowledge_base = AsyncMock()
|
||||
app = _app()
|
||||
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb1', 'name': 'Updated'}
|
||||
service = KnowledgeService(app)
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
|
||||
# Pass both mutable and immutable fields
|
||||
await service.update_knowledge_base(
|
||||
CONTEXT,
|
||||
'kb1',
|
||||
{
|
||||
'name': 'New Name',
|
||||
'description': 'New desc',
|
||||
'uuid': 'should_be_filtered', # immutable
|
||||
'uuid': 'should_be_filtered',
|
||||
},
|
||||
)
|
||||
|
||||
# Check that only mutable fields were passed to update
|
||||
call_args = mock_app.persistence_mgr.execute_async.call_args
|
||||
assert call_args is not None
|
||||
update_statement = app.persistence_mgr.execute_async.await_args_list[0].args[0]
|
||||
params = update_statement.compile().params
|
||||
assert params['name'] == 'New Name'
|
||||
assert params['description'] == 'New desc'
|
||||
assert 'uuid' not in params
|
||||
app.rag_mgr.remove_knowledge_base_from_runtime.assert_awaited_once_with(CONTEXT, 'kb1')
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_early_when_no_mutable_fields(self):
|
||||
"""Test that update returns early when no mutable fields provided."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
app = _app()
|
||||
app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb1'}
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
await KnowledgeService(app).update_knowledge_base(
|
||||
CONTEXT,
|
||||
'kb1',
|
||||
{'uuid': 'should_be_filtered'},
|
||||
)
|
||||
|
||||
# Pass only immutable fields
|
||||
await service.update_knowledge_base('kb1', {'uuid': 'should_be_filtered'})
|
||||
|
||||
# No DB update should be called
|
||||
mock_app.persistence_mgr.execute_async.assert_not_called()
|
||||
app.persistence_mgr.execute_async.assert_not_awaited()
|
||||
app.rag_mgr.remove_knowledge_base_from_runtime.assert_not_awaited()
|
||||
|
||||
|
||||
class TestCheckDocCapability:
|
||||
"""Tests for _check_doc_capability method."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_passes_when_capability_supported(self):
|
||||
"""Test that check passes when doc_ingestion capability exists."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(
|
||||
return_value={'knowledge_engine': {'capabilities': ['doc_ingestion']}}
|
||||
)
|
||||
app = _app()
|
||||
app.rag_mgr.get_knowledge_base_details.return_value = {'knowledge_engine': {'capabilities': ['doc_ingestion']}}
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
|
||||
await service._check_doc_capability('kb1', 'document upload')
|
||||
|
||||
# No exception raised means success
|
||||
await KnowledgeService(app)._check_doc_capability(CONTEXT, 'kb1', 'document upload')
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raises_when_kb_not_found(self):
|
||||
"""Test that Exception is raised when KB not found."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value=None)
|
||||
app = _app()
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
await service._check_doc_capability('nonexistent', 'test operation')
|
||||
|
||||
assert 'Knowledge base not found' in str(exc_info.value)
|
||||
with pytest.raises(WorkspaceNotFoundError, match='Knowledge base not found'):
|
||||
await KnowledgeService(app)._check_doc_capability(
|
||||
CONTEXT,
|
||||
'nonexistent',
|
||||
'test operation',
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raises_when_capability_not_supported(self):
|
||||
"""Test that Exception is raised when doc_ingestion not in capabilities."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(
|
||||
return_value={'knowledge_engine': {'capabilities': ['other_capability']}}
|
||||
)
|
||||
app = _app()
|
||||
app.rag_mgr.get_knowledge_base_details.return_value = {
|
||||
'knowledge_engine': {'capabilities': ['other_capability']}
|
||||
}
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
|
||||
with pytest.raises(Exception) as exc_info:
|
||||
await service._check_doc_capability('kb1', 'document upload')
|
||||
|
||||
assert 'does not support document upload' in str(exc_info.value)
|
||||
with pytest.raises(Exception, match='does not support document upload'):
|
||||
await KnowledgeService(app)._check_doc_capability(
|
||||
CONTEXT,
|
||||
'kb1',
|
||||
'document upload',
|
||||
)
|
||||
|
||||
|
||||
class TestListKnowledgeEngines:
|
||||
"""Tests for list_knowledge_engines method."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_engines_from_plugin_connector(self):
|
||||
"""Test that it returns knowledge engines from plugin connector."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.plugin_connector.list_knowledge_engines = AsyncMock(
|
||||
return_value=[{'id': 'engine1', 'name': 'Engine 1'}]
|
||||
)
|
||||
app = _app()
|
||||
app.plugin_connector.list_knowledge_engines.return_value = [{'id': 'engine1', 'name': 'Engine 1'}]
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.list_knowledge_engines()
|
||||
result = await KnowledgeService(app).list_knowledge_engines(CONTEXT)
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0]['id'] == 'engine1'
|
||||
assert result == [{'id': 'engine1', 'name': 'Engine 1'}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_empty_when_plugin_disabled(self):
|
||||
"""Test that it returns empty list when plugin disabled."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.plugin_connector.is_enable_plugin = False
|
||||
app = _app()
|
||||
app.plugin_connector.is_enable_plugin = False
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.list_knowledge_engines()
|
||||
|
||||
assert result == []
|
||||
assert await KnowledgeService(app).list_knowledge_engines(CONTEXT) == []
|
||||
app.plugin_connector.list_knowledge_engines.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_empty_on_exception(self):
|
||||
"""Test that it returns empty list and logs warning on exception."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.plugin_connector.list_knowledge_engines = AsyncMock(side_effect=Exception('Connection error'))
|
||||
app = _app()
|
||||
app.plugin_connector.list_knowledge_engines.side_effect = RuntimeError('Connection error')
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.list_knowledge_engines()
|
||||
|
||||
assert result == []
|
||||
mock_app.logger.warning.assert_called_once()
|
||||
assert await KnowledgeService(app).list_knowledge_engines(CONTEXT) == []
|
||||
app.logger.warning.assert_called_once()
|
||||
|
||||
|
||||
class TestListParsers:
|
||||
"""Tests for list_parsers method."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_all_parsers(self):
|
||||
"""Test that it returns all parsers when no MIME type filter."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.plugin_connector.list_parsers = AsyncMock(
|
||||
return_value=[
|
||||
{'id': 'parser1', 'supported_mime_types': ['text/plain']},
|
||||
{'id': 'parser2', 'supported_mime_types': ['application/pdf']},
|
||||
]
|
||||
)
|
||||
app = _app()
|
||||
app.plugin_connector.list_parsers.return_value = [
|
||||
{'id': 'parser1', 'supported_mime_types': ['text/plain']},
|
||||
{'id': 'parser2', 'supported_mime_types': ['application/pdf']},
|
||||
]
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.list_parsers()
|
||||
result = await KnowledgeService(app).list_parsers(CONTEXT)
|
||||
|
||||
assert len(result) == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_filters_by_mime_type(self):
|
||||
"""Test that it filters parsers by MIME type."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.plugin_connector.list_parsers = AsyncMock(
|
||||
return_value=[
|
||||
{'id': 'parser1', 'supported_mime_types': ['text/plain']},
|
||||
{'id': 'parser2', 'supported_mime_types': ['application/pdf']},
|
||||
]
|
||||
)
|
||||
app = _app()
|
||||
app.plugin_connector.list_parsers.return_value = [
|
||||
{'id': 'parser1', 'supported_mime_types': ['text/plain']},
|
||||
{'id': 'parser2', 'supported_mime_types': ['application/pdf']},
|
||||
]
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.list_parsers(mime_type='application/pdf')
|
||||
result = await KnowledgeService(app).list_parsers(CONTEXT, 'application/pdf')
|
||||
|
||||
assert len(result) == 1
|
||||
assert result[0]['id'] == 'parser2'
|
||||
assert result == [{'id': 'parser2', 'supported_mime_types': ['application/pdf']}]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_empty_when_plugin_disabled(self):
|
||||
"""Test that it returns empty list when plugin disabled."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.plugin_connector.is_enable_plugin = False
|
||||
app = _app()
|
||||
app.plugin_connector.is_enable_plugin = False
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.list_parsers()
|
||||
|
||||
assert result == []
|
||||
assert await KnowledgeService(app).list_parsers(CONTEXT) == []
|
||||
app.plugin_connector.list_parsers.assert_not_awaited()
|
||||
|
||||
|
||||
class TestGetEngineSchemas:
|
||||
"""Tests for get_engine_creation_schema and get_engine_retrieval_schema."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_creation_schema(self):
|
||||
"""Test that it returns creation schema for engine."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.plugin_connector.get_rag_creation_schema = AsyncMock(
|
||||
return_value={'properties': {'name': {'type': 'string'}}}
|
||||
)
|
||||
app = _app()
|
||||
app.plugin_connector.get_rag_creation_schema.return_value = {'properties': {'name': {'type': 'string'}}}
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.get_engine_creation_schema('author/engine')
|
||||
result = await KnowledgeService(app).get_engine_creation_schema(
|
||||
CONTEXT,
|
||||
'author/engine',
|
||||
)
|
||||
|
||||
assert 'properties' in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_retrieval_schema(self):
|
||||
"""Test that it returns retrieval schema for engine."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.plugin_connector.get_rag_retrieval_schema = AsyncMock(
|
||||
return_value={'properties': {'top_k': {'type': 'integer'}}}
|
||||
)
|
||||
app = _app()
|
||||
app.plugin_connector.get_rag_retrieval_schema.return_value = {'properties': {'top_k': {'type': 'integer'}}}
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.get_engine_retrieval_schema('author/engine')
|
||||
result = await KnowledgeService(app).get_engine_retrieval_schema(
|
||||
CONTEXT,
|
||||
'author/engine',
|
||||
)
|
||||
|
||||
assert 'properties' in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_returns_empty_dict_on_exception(self):
|
||||
"""Test that it returns empty dict and logs warning on exception."""
|
||||
knowledge_module = get_knowledge_service_module()
|
||||
mock_app = create_mock_app()
|
||||
mock_app.plugin_connector.get_rag_creation_schema = AsyncMock(side_effect=Exception('Plugin error'))
|
||||
app = _app()
|
||||
app.plugin_connector.get_rag_creation_schema.side_effect = RuntimeError('Plugin error')
|
||||
|
||||
service = knowledge_module.KnowledgeService(mock_app)
|
||||
result = await service.get_engine_creation_schema('author/engine')
|
||||
result = await KnowledgeService(app).get_engine_creation_schema(
|
||||
CONTEXT,
|
||||
'author/engine',
|
||||
)
|
||||
|
||||
assert result == {}
|
||||
mock_app.logger.warning.assert_called_once()
|
||||
app.logger.warning.assert_called_once()
|
||||
|
||||
|
||||
class TestKnowledgeBaseSecretViews:
|
||||
@pytest.mark.asyncio
|
||||
async def test_creation_settings_are_redacted_for_resource_view_only(self):
|
||||
app = _app()
|
||||
raw = {
|
||||
'uuid': 'kb-secret',
|
||||
'creation_settings': {
|
||||
'dify_apikey': 'dify-secret',
|
||||
'headers': {'Authorization': 'Bearer secret'},
|
||||
},
|
||||
}
|
||||
app.rag_mgr.get_all_knowledge_base_details.return_value = [raw]
|
||||
service = KnowledgeService(app)
|
||||
|
||||
redacted = await service.get_knowledge_bases(CONTEXT)
|
||||
manager_view = await service.get_knowledge_bases(CONTEXT, include_secret=True)
|
||||
|
||||
assert redacted[0]['creation_settings']['dify_apikey'] == '***'
|
||||
assert redacted[0]['creation_settings']['headers']['Authorization'] == '***'
|
||||
assert manager_view[0]['creation_settings']['dify_apikey'] == 'dify-secret'
|
||||
assert raw['creation_settings']['dify_apikey'] == 'dify-secret'
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_new_masked_creation_secret_is_rejected(self):
|
||||
app = _app()
|
||||
|
||||
with pytest.raises(ValueError, match='no existing value'):
|
||||
await KnowledgeService(app).create_knowledge_base(
|
||||
CONTEXT,
|
||||
{
|
||||
'knowledge_engine_plugin_id': 'author/engine',
|
||||
'creation_settings': {'dify_apikey': '***'},
|
||||
},
|
||||
)
|
||||
|
||||
app.rag_mgr.create_knowledge_base.assert_not_awaited()
|
||||
|
||||
@@ -13,16 +13,39 @@ Source: src/langbot/pkg/api/http/service/maintenance.py
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import contextlib
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, Mock, patch, MagicMock
|
||||
from types import SimpleNamespace
|
||||
import datetime
|
||||
from pathlib import Path
|
||||
|
||||
import sqlalchemy
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
from langbot.pkg.api.http.authz import WorkspaceRequiredError
|
||||
from langbot.pkg.api.http.service.maintenance import MaintenanceService
|
||||
from langbot.pkg.api.http.context import ExecutionContext, PrincipalContext, PrincipalType
|
||||
from langbot.pkg.entity.persistence.base import Base
|
||||
from langbot.pkg.entity.persistence.bstorage import BinaryStorage
|
||||
from langbot.pkg.entity.persistence.monitoring import MonitoringMessage
|
||||
from langbot.pkg.entity.persistence.workspace import Workspace
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
TEST_CONTEXT = ExecutionContext(
|
||||
instance_uuid='test-instance',
|
||||
workspace_uuid='test-workspace',
|
||||
placement_generation=1,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def assume_oss_singleton(monkeypatch):
|
||||
async def is_oss_singleton(_self, _context):
|
||||
return True
|
||||
|
||||
monkeypatch.setattr(MaintenanceService, '_is_oss_singleton', is_oss_singleton)
|
||||
|
||||
|
||||
def _create_mock_result(scalar_value=None):
|
||||
@@ -32,6 +55,14 @@ def _create_mock_result(scalar_value=None):
|
||||
return result
|
||||
|
||||
|
||||
def _scoped_storage_manager():
|
||||
prefix = 'instances/i/workspaces/w/generations/1/owners/upload/o/'
|
||||
return SimpleNamespace(
|
||||
scoped_prefix=Mock(return_value=prefix),
|
||||
is_scoped_object_key=Mock(side_effect=lambda key, **_: key == f'{prefix}uploaded_file.txt'),
|
||||
)
|
||||
|
||||
|
||||
class TestMaintenanceServiceCleanupExpiredFiles:
|
||||
"""Tests for cleanup_expired_files method."""
|
||||
|
||||
@@ -39,6 +70,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
|
||||
"""Uses default retention days when config not set."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.storage_mgr = _scoped_storage_manager()
|
||||
ap.instance_config = SimpleNamespace()
|
||||
ap.instance_config.data = {}
|
||||
ap.storage_mgr = SimpleNamespace()
|
||||
@@ -58,7 +90,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
|
||||
service._cleanup_expired_log_files = Mock(return_value=0) # NOT async!
|
||||
|
||||
# Execute
|
||||
result = await service.cleanup_expired_files()
|
||||
result = await service.cleanup_expired_files(TEST_CONTEXT)
|
||||
|
||||
# Verify - returns counts
|
||||
assert 'uploaded_files' in result
|
||||
@@ -95,7 +127,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
|
||||
service._cleanup_expired_log_files = Mock(return_value=3) # NOT async
|
||||
|
||||
# Execute
|
||||
result = await service.cleanup_expired_files()
|
||||
result = await service.cleanup_expired_files(TEST_CONTEXT)
|
||||
|
||||
# Verify
|
||||
assert result['uploaded_files'] == 2
|
||||
@@ -124,7 +156,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
|
||||
service._cleanup_expired_log_files = Mock(return_value=0) # NOT async
|
||||
|
||||
# Execute
|
||||
result = await service.cleanup_expired_files()
|
||||
result = await service.cleanup_expired_files(TEST_CONTEXT)
|
||||
|
||||
# Verify
|
||||
assert result['uploaded_files'] == 1
|
||||
@@ -159,12 +191,57 @@ class TestMaintenanceServiceCleanupExpiredFiles:
|
||||
service._cleanup_expired_log_files = Mock(return_value=0) # NOT async
|
||||
|
||||
# Execute
|
||||
result = await service.cleanup_expired_files()
|
||||
result = await service.cleanup_expired_files(TEST_CONTEXT)
|
||||
|
||||
# Verify - warning logged, defaults used
|
||||
assert ap.logger.warning.called
|
||||
assert 'uploaded_files' in result
|
||||
|
||||
async def test_cloud_cleanup_carries_scope_without_holding_database_session(self):
|
||||
class ScopeOnlyPersistenceManager:
|
||||
mode = SimpleNamespace(value='cloud_runtime')
|
||||
|
||||
def __init__(self):
|
||||
self.active_workspace = None
|
||||
|
||||
@contextlib.asynccontextmanager
|
||||
async def tenant_scope(self, workspace_uuid):
|
||||
self.active_workspace = workspace_uuid
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
self.active_workspace = None
|
||||
|
||||
def current_session(self):
|
||||
return None
|
||||
|
||||
persistence_mgr = ScopeOnlyPersistenceManager()
|
||||
application = SimpleNamespace(
|
||||
persistence_mgr=persistence_mgr,
|
||||
instance_config=SimpleNamespace(data={}),
|
||||
logger=SimpleNamespace(warning=Mock()),
|
||||
)
|
||||
service = MaintenanceService(application)
|
||||
|
||||
async def cleanup_uploads(_context, _retention_days):
|
||||
assert persistence_mgr.active_workspace == TEST_CONTEXT.workspace_uuid
|
||||
assert persistence_mgr.current_session() is None
|
||||
return 2
|
||||
|
||||
def cleanup_logs(_retention_days):
|
||||
assert persistence_mgr.active_workspace == TEST_CONTEXT.workspace_uuid
|
||||
assert persistence_mgr.current_session() is None
|
||||
return 1
|
||||
|
||||
service._cleanup_expired_uploaded_files = cleanup_uploads
|
||||
service._cleanup_expired_log_files = cleanup_logs
|
||||
|
||||
assert await service.cleanup_expired_files(TEST_CONTEXT) == {
|
||||
'uploaded_files': 2,
|
||||
'log_files': 1,
|
||||
}
|
||||
assert persistence_mgr.active_workspace is None
|
||||
|
||||
|
||||
class TestMaintenanceServiceGetStorageAnalysis:
|
||||
"""Tests for get_storage_analysis method."""
|
||||
@@ -196,7 +273,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
|
||||
service._expired_log_candidates = Mock(return_value=[])
|
||||
|
||||
# Execute
|
||||
result = await service.get_storage_analysis()
|
||||
result = await service.get_storage_analysis(TEST_CONTEXT)
|
||||
|
||||
# Verify
|
||||
assert 'generated_at' in result
|
||||
@@ -229,7 +306,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
|
||||
service._expired_log_candidates = Mock(return_value=[])
|
||||
|
||||
# Execute
|
||||
result = await service.get_storage_analysis()
|
||||
result = await service.get_storage_analysis(TEST_CONTEXT)
|
||||
|
||||
# Verify - all sections present
|
||||
sections = {s['key'] for s in result['sections']}
|
||||
@@ -265,7 +342,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
|
||||
service._expired_log_candidates = Mock(return_value=[])
|
||||
|
||||
# Execute
|
||||
result = await service.get_storage_analysis()
|
||||
result = await service.get_storage_analysis(TEST_CONTEXT)
|
||||
|
||||
# Verify
|
||||
assert result['database']['type'] == 'postgresql'
|
||||
@@ -294,7 +371,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
|
||||
service._expired_log_candidates = Mock(return_value=[{'name': 'old_log', 'size_bytes': 50}])
|
||||
|
||||
# Execute
|
||||
result = await service.get_storage_analysis()
|
||||
result = await service.get_storage_analysis(TEST_CONTEXT)
|
||||
|
||||
# Verify
|
||||
assert len(result['cleanup_candidates']['uploaded_files']) == 1
|
||||
@@ -316,7 +393,7 @@ class TestMaintenanceServiceMonitoringCounts:
|
||||
service = MaintenanceService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service._monitoring_counts()
|
||||
result = await service._monitoring_counts(TEST_CONTEXT)
|
||||
|
||||
# Verify - all table keys present
|
||||
assert 'messages' in result
|
||||
@@ -338,7 +415,7 @@ class TestMaintenanceServiceMonitoringCounts:
|
||||
service = MaintenanceService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service._monitoring_counts()
|
||||
result = await service._monitoring_counts(TEST_CONTEXT)
|
||||
|
||||
# Verify - all zero
|
||||
assert all(v == 0 for v in result.values())
|
||||
@@ -374,7 +451,7 @@ class TestMaintenanceServiceBinaryStorageStats:
|
||||
service = MaintenanceService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service._binary_storage_stats()
|
||||
result = await service._binary_storage_stats(TEST_CONTEXT)
|
||||
|
||||
# Verify
|
||||
assert result['count'] == 10
|
||||
@@ -404,7 +481,7 @@ class TestMaintenanceServiceBinaryStorageStats:
|
||||
service = MaintenanceService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service._binary_storage_stats()
|
||||
result = await service._binary_storage_stats(TEST_CONTEXT)
|
||||
|
||||
# Verify - warning logged, size_bytes None or 0
|
||||
assert ap.logger.warning.called
|
||||
@@ -618,11 +695,13 @@ class TestMaintenanceServiceIsUploadedFileKey:
|
||||
"""Returns True for valid upload file key."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.storage_mgr = _scoped_storage_manager()
|
||||
|
||||
service = MaintenanceService(ap)
|
||||
|
||||
# Execute - simple filename without path
|
||||
result = service._is_uploaded_file_key('uploaded_file.txt')
|
||||
key = f'{ap.storage_mgr.scoped_prefix(TEST_CONTEXT, owner_type="upload")}uploaded_file.txt'
|
||||
result = service._is_uploaded_file_key(TEST_CONTEXT, key)
|
||||
|
||||
# Verify
|
||||
assert result is True
|
||||
@@ -631,11 +710,12 @@ class TestMaintenanceServiceIsUploadedFileKey:
|
||||
"""Returns False for key with path separator."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.storage_mgr = _scoped_storage_manager()
|
||||
|
||||
service = MaintenanceService(ap)
|
||||
|
||||
# Execute - key with path
|
||||
result = service._is_uploaded_file_key('path/to/file.txt')
|
||||
result = service._is_uploaded_file_key(TEST_CONTEXT, 'path/to/file.txt')
|
||||
|
||||
# Verify
|
||||
assert result is False
|
||||
@@ -644,11 +724,12 @@ class TestMaintenanceServiceIsUploadedFileKey:
|
||||
"""Returns False for plugin config prefix."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.storage_mgr = _scoped_storage_manager()
|
||||
|
||||
service = MaintenanceService(ap)
|
||||
|
||||
# Execute - plugin config file
|
||||
result = service._is_uploaded_file_key('plugin_config_some_plugin.json')
|
||||
result = service._is_uploaded_file_key(TEST_CONTEXT, 'plugin_config_some_plugin.json')
|
||||
|
||||
# Verify
|
||||
assert result is False
|
||||
@@ -662,6 +743,7 @@ class TestMaintenanceServiceExpiredLogCandidates:
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.logger = SimpleNamespace()
|
||||
ap.storage_mgr = _scoped_storage_manager()
|
||||
|
||||
service = MaintenanceService(ap)
|
||||
|
||||
@@ -748,11 +830,12 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.logger = SimpleNamespace()
|
||||
ap.storage_mgr = _scoped_storage_manager()
|
||||
|
||||
service = MaintenanceService(ap)
|
||||
|
||||
with patch.object(Path, 'exists', return_value=False):
|
||||
result = service._expired_local_upload_candidates(7)
|
||||
result = service._expired_local_upload_candidates(TEST_CONTEXT, 7)
|
||||
|
||||
# Verify
|
||||
assert result == []
|
||||
@@ -762,12 +845,10 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.logger = SimpleNamespace()
|
||||
ap.storage_mgr = _scoped_storage_manager()
|
||||
|
||||
service = MaintenanceService(ap)
|
||||
# Mock _is_uploaded_file_key
|
||||
service._is_uploaded_file_key = Mock(side_effect=lambda key: 'plugin_config_' not in key and '/' not in key)
|
||||
|
||||
# Create mock files - one valid, one plugin config
|
||||
# Create one file and one non-file entry under the scoped upload root.
|
||||
mock_entry_valid = Mock(spec=Path)
|
||||
mock_entry_valid.is_file = Mock(return_value=True)
|
||||
mock_entry_valid.name = 'valid_upload.txt'
|
||||
@@ -775,9 +856,10 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
|
||||
mock_stat.st_size = 100
|
||||
mock_stat.st_mtime = 0 # Very old
|
||||
mock_entry_valid.stat = Mock(return_value=mock_stat)
|
||||
mock_entry_valid.relative_to = Mock(return_value=Path('scoped/valid_upload.txt'))
|
||||
|
||||
mock_entry_plugin = Mock(spec=Path)
|
||||
mock_entry_plugin.is_file = Mock(return_value=True)
|
||||
mock_entry_plugin.is_file = Mock(return_value=False)
|
||||
mock_entry_plugin.name = 'plugin_config_test.json'
|
||||
mock_stat2 = Mock()
|
||||
mock_stat2.st_size = 200
|
||||
@@ -785,23 +867,22 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
|
||||
mock_entry_plugin.stat = Mock(return_value=mock_stat2)
|
||||
|
||||
with patch.object(Path, 'exists', return_value=True):
|
||||
with patch.object(Path, 'iterdir') as mock_iterdir:
|
||||
mock_iterdir.return_value = [mock_entry_valid, mock_entry_plugin]
|
||||
result = service._expired_local_upload_candidates(7)
|
||||
with patch.object(Path, 'rglob') as mock_rglob:
|
||||
mock_rglob.return_value = [mock_entry_valid, mock_entry_plugin]
|
||||
result = service._expired_local_upload_candidates(TEST_CONTEXT, 7)
|
||||
|
||||
# Verify - only valid upload included
|
||||
assert len(result) == 1
|
||||
assert result[0]['key'] == 'valid_upload.txt'
|
||||
assert result[0]['key'] == 'scoped/valid_upload.txt'
|
||||
|
||||
def test_expired_local_upload_candidates_includes_path(self):
|
||||
"""Includes path when include_paths=True."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.logger = SimpleNamespace()
|
||||
ap.storage_mgr = _scoped_storage_manager()
|
||||
|
||||
service = MaintenanceService(ap)
|
||||
service._is_uploaded_file_key = Mock(return_value=True)
|
||||
|
||||
mock_entry = Mock(spec=Path)
|
||||
mock_entry.is_file = Mock(return_value=True)
|
||||
mock_entry.name = 'old_file.txt'
|
||||
@@ -810,11 +891,178 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
|
||||
mock_stat.st_size = 100
|
||||
mock_stat.st_mtime = 0
|
||||
mock_entry.stat = Mock(return_value=mock_stat)
|
||||
mock_entry.relative_to = Mock(return_value=Path('scoped/old_file.txt'))
|
||||
|
||||
with patch.object(Path, 'exists', return_value=True):
|
||||
with patch.object(Path, 'iterdir') as mock_iterdir:
|
||||
mock_iterdir.return_value = [mock_entry]
|
||||
result = service._expired_local_upload_candidates(7, include_paths=True)
|
||||
with patch.object(Path, 'rglob') as mock_rglob:
|
||||
mock_rglob.return_value = [mock_entry]
|
||||
result = service._expired_local_upload_candidates(
|
||||
TEST_CONTEXT,
|
||||
7,
|
||||
include_paths=True,
|
||||
)
|
||||
|
||||
# Verify - path included
|
||||
assert 'path' in result[0]
|
||||
|
||||
def test_expired_local_upload_candidates_respects_run_limit(self):
|
||||
ap = SimpleNamespace(
|
||||
logger=SimpleNamespace(warning=Mock()),
|
||||
storage_mgr=_scoped_storage_manager(),
|
||||
instance_config=SimpleNamespace(data={'storage': {'cleanup': {'max_files_per_run': 2}}}),
|
||||
)
|
||||
service = MaintenanceService(ap)
|
||||
entries = []
|
||||
for index in range(3):
|
||||
entry = Mock(spec=Path)
|
||||
entry.is_file = Mock(return_value=True)
|
||||
entry.stat = Mock(return_value=SimpleNamespace(st_size=100, st_mtime=0))
|
||||
entry.relative_to = Mock(return_value=Path(f'scoped/old-{index}.txt'))
|
||||
entries.append(entry)
|
||||
|
||||
with patch.object(Path, 'exists', return_value=True):
|
||||
with patch.object(Path, 'rglob', return_value=entries):
|
||||
result = service._expired_local_upload_candidates(TEST_CONTEXT, 7)
|
||||
|
||||
assert [item['key'] for item in result] == [
|
||||
'scoped/old-0.txt',
|
||||
'scoped/old-1.txt',
|
||||
]
|
||||
ap.instance_config.data['storage']['cleanup']['max_files_per_run'] = 999999
|
||||
assert service._max_files_per_run() == 10000
|
||||
|
||||
|
||||
ISOLATION_WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
|
||||
ISOLATION_WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
|
||||
|
||||
|
||||
def _tenant_context(workspace_uuid: str) -> ExecutionContext:
|
||||
return ExecutionContext(
|
||||
instance_uuid='instance',
|
||||
workspace_uuid=workspace_uuid,
|
||||
placement_generation=1,
|
||||
trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
|
||||
)
|
||||
|
||||
|
||||
class _RealPersistenceManager:
|
||||
def __init__(self, engine):
|
||||
self.engine = engine
|
||||
|
||||
async def execute_async(self, *args, **kwargs):
|
||||
async with self.engine.connect() as connection:
|
||||
result = await connection.execute(*args, **kwargs)
|
||||
await connection.commit()
|
||||
return result
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def tenant_maintenance_service(tmp_path):
|
||||
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "maintenance.db"}')
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(Workspace),
|
||||
[
|
||||
{
|
||||
'uuid': ISOLATION_WORKSPACE_A,
|
||||
'instance_uuid': 'instance',
|
||||
'name': 'A',
|
||||
'slug': 'a',
|
||||
'source': 'cloud_projection',
|
||||
},
|
||||
{
|
||||
'uuid': ISOLATION_WORKSPACE_B,
|
||||
'instance_uuid': 'instance',
|
||||
'name': 'B',
|
||||
'slug': 'b',
|
||||
'source': 'cloud_projection',
|
||||
},
|
||||
],
|
||||
)
|
||||
now = datetime.datetime.now(datetime.UTC).replace(tzinfo=None)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(MonitoringMessage),
|
||||
[
|
||||
{
|
||||
'id': 'message-a',
|
||||
'workspace_uuid': ISOLATION_WORKSPACE_A,
|
||||
'timestamp': now,
|
||||
'bot_id': 'bot',
|
||||
'bot_name': 'Bot',
|
||||
'pipeline_id': 'pipeline',
|
||||
'pipeline_name': 'Pipeline',
|
||||
'message_content': 'A',
|
||||
'session_id': 'same-session',
|
||||
'status': 'success',
|
||||
'level': 'info',
|
||||
},
|
||||
{
|
||||
'id': 'message-b',
|
||||
'workspace_uuid': ISOLATION_WORKSPACE_B,
|
||||
'timestamp': now,
|
||||
'bot_id': 'bot',
|
||||
'bot_name': 'Bot',
|
||||
'pipeline_id': 'pipeline',
|
||||
'pipeline_name': 'Pipeline',
|
||||
'message_content': 'B',
|
||||
'session_id': 'same-session',
|
||||
'status': 'success',
|
||||
'level': 'info',
|
||||
},
|
||||
],
|
||||
)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(BinaryStorage),
|
||||
[
|
||||
{
|
||||
'workspace_uuid': ISOLATION_WORKSPACE_A,
|
||||
'unique_key': 'a',
|
||||
'key': 'same',
|
||||
'owner_type': 'plugin',
|
||||
'owner': 'same',
|
||||
'value': b'aaa',
|
||||
},
|
||||
{
|
||||
'workspace_uuid': ISOLATION_WORKSPACE_B,
|
||||
'unique_key': 'b',
|
||||
'key': 'same',
|
||||
'owner_type': 'plugin',
|
||||
'owner': 'same',
|
||||
'value': b'bbbbb',
|
||||
},
|
||||
],
|
||||
)
|
||||
|
||||
application = SimpleNamespace(
|
||||
persistence_mgr=_RealPersistenceManager(engine),
|
||||
instance_config=SimpleNamespace(data={}),
|
||||
logger=SimpleNamespace(warning=lambda *_: None),
|
||||
)
|
||||
yield MaintenanceService(application)
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def test_cleanup_requires_execution_context(tenant_maintenance_service):
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
await tenant_maintenance_service.cleanup_expired_files(None)
|
||||
|
||||
|
||||
async def test_monitoring_counts_are_workspace_scoped(tenant_maintenance_service):
|
||||
counts_a = await tenant_maintenance_service._monitoring_counts(_tenant_context(ISOLATION_WORKSPACE_A))
|
||||
counts_b = await tenant_maintenance_service._monitoring_counts(_tenant_context(ISOLATION_WORKSPACE_B))
|
||||
assert counts_a['messages'] == 1
|
||||
assert counts_b['messages'] == 1
|
||||
|
||||
|
||||
async def test_binary_storage_stats_are_workspace_scoped(tenant_maintenance_service):
|
||||
stats_a = await tenant_maintenance_service._binary_storage_stats(_tenant_context(ISOLATION_WORKSPACE_A))
|
||||
stats_b = await tenant_maintenance_service._binary_storage_stats(_tenant_context(ISOLATION_WORKSPACE_B))
|
||||
assert stats_a == {'count': 1, 'size_bytes': 3}
|
||||
assert stats_b == {'count': 1, 'size_bytes': 5}
|
||||
|
||||
|
||||
async def test_path_helpers_handle_missing_paths(tenant_maintenance_service, tmp_path):
|
||||
missing = tmp_path / 'missing'
|
||||
assert tenant_maintenance_service._path_size(missing) == 0
|
||||
assert tenant_maintenance_service._file_count(missing) == 0
|
||||
|
||||
@@ -13,17 +13,62 @@ Source: src/langbot/pkg/api/http/service/mcp.py
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import copy
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, Mock, MagicMock
|
||||
from types import SimpleNamespace
|
||||
import uuid
|
||||
|
||||
from langbot.pkg.api.http.service.mcp import MCPService
|
||||
from langbot.pkg.api.http.authz import Permission
|
||||
from langbot.pkg.api.http.context import (
|
||||
ExecutionContext,
|
||||
PrincipalContext,
|
||||
PrincipalType,
|
||||
RequestContext,
|
||||
WorkspaceContext,
|
||||
)
|
||||
from langbot.pkg.api.http.service.mcp import MCPService, redact_mcp_secrets, restore_mcp_secret_placeholders
|
||||
from langbot.pkg.core.taskmgr import TaskCapacityError
|
||||
from langbot.pkg.entity.persistence.mcp import MCPServer
|
||||
from langbot.pkg.provider.tools.loaders.mcp_policy import MCPStdioDisabledError
|
||||
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
_CONTEXT = ExecutionContext(
|
||||
instance_uuid='instance-a',
|
||||
workspace_uuid='workspace-a',
|
||||
placement_generation=1,
|
||||
)
|
||||
|
||||
_VIEWER_CONTEXT = RequestContext(
|
||||
instance_uuid='instance-a',
|
||||
placement_generation=1,
|
||||
request_id='request-a',
|
||||
auth_type='user_token',
|
||||
principal=PrincipalContext(
|
||||
principal_type=PrincipalType.ACCOUNT,
|
||||
account_uuid='account-a',
|
||||
),
|
||||
workspace=WorkspaceContext(
|
||||
workspace_uuid='workspace-a',
|
||||
membership_uuid='membership-a',
|
||||
role='viewer',
|
||||
permissions=frozenset({Permission.RESOURCE_VIEW.value}),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
def _service(ap: SimpleNamespace) -> MCPService:
|
||||
ap.workspace_service = SimpleNamespace(
|
||||
get_execution_binding=AsyncMock(return_value=SimpleNamespace(instance_uuid=_CONTEXT.instance_uuid))
|
||||
)
|
||||
if not hasattr(ap, 'logger'):
|
||||
ap.logger = Mock()
|
||||
return MCPService(ap)
|
||||
|
||||
|
||||
def _create_mock_mcp_server(
|
||||
server_uuid: str = None,
|
||||
@@ -42,11 +87,13 @@ def _create_mock_mcp_server(
|
||||
return server
|
||||
|
||||
|
||||
def _create_mock_result(items: list = None, first_item=None):
|
||||
def _create_mock_result(items: list = None, first_item=None, *, scalar_value=0, rowcount=1):
|
||||
"""Create mock result object for persistence queries."""
|
||||
result = Mock()
|
||||
result.all = Mock(return_value=items or [])
|
||||
result.first = Mock(return_value=first_item)
|
||||
result.scalar = Mock(return_value=scalar_value)
|
||||
result.rowcount = rowcount
|
||||
return result
|
||||
|
||||
|
||||
@@ -64,10 +111,10 @@ class TestMCPServiceGetRuntimeInfo:
|
||||
mock_session.get_runtime_info_dict = Mock(return_value={'status': 'running', 'tools': 5})
|
||||
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=mock_session)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_runtime_info('test-server')
|
||||
result = await service.get_runtime_info(_CONTEXT, 'test-server')
|
||||
|
||||
# Verify
|
||||
assert result is not None
|
||||
@@ -81,10 +128,10 @@ class TestMCPServiceGetRuntimeInfo:
|
||||
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
|
||||
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=None)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_runtime_info('nonexistent-server')
|
||||
result = await service.get_runtime_info(_CONTEXT, 'nonexistent-server')
|
||||
|
||||
# Verify
|
||||
assert result is None
|
||||
@@ -101,12 +148,13 @@ class TestMCPServiceResources:
|
||||
return_value=[{'uri_template': 'file:///{path}', 'name': 'files'}]
|
||||
)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
service._require_server = AsyncMock(return_value=(_CONTEXT, {'name': 'docs'}))
|
||||
|
||||
result = await service.get_mcp_server_resource_templates('docs')
|
||||
result = await service.get_mcp_server_resource_templates(_CONTEXT, 'docs')
|
||||
|
||||
assert result == [{'uri_template': 'file:///{path}', 'name': 'files'}]
|
||||
ap.tool_mgr.mcp_tool_loader.get_resource_templates.assert_awaited_once_with('docs')
|
||||
ap.tool_mgr.mcp_tool_loader.get_resource_templates.assert_awaited_once_with(_CONTEXT, 'docs')
|
||||
|
||||
async def test_read_resource_envelope_uses_ui_preview_source(self):
|
||||
ap = SimpleNamespace()
|
||||
@@ -121,9 +169,11 @@ class TestMCPServiceResources:
|
||||
}
|
||||
)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
service._require_server = AsyncMock(return_value=(_CONTEXT, {'name': 'docs'}))
|
||||
|
||||
result = await service.read_mcp_server_resource_envelope(
|
||||
_CONTEXT,
|
||||
'docs',
|
||||
'file:///README.md',
|
||||
max_bytes=4096,
|
||||
@@ -132,6 +182,7 @@ class TestMCPServiceResources:
|
||||
|
||||
assert result['source'] == 'ui_preview'
|
||||
ap.tool_mgr.mcp_tool_loader.read_resource_envelope.assert_awaited_once_with(
|
||||
_CONTEXT,
|
||||
'docs',
|
||||
'file:///README.md',
|
||||
include_blob=True,
|
||||
@@ -156,12 +207,12 @@ class TestMCPServiceGetMCPServers:
|
||||
'name': entity.name,
|
||||
}
|
||||
)
|
||||
ap.tool_mgr = None
|
||||
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_mcp_servers()
|
||||
result = await service.get_mcp_servers(_CONTEXT)
|
||||
|
||||
# Verify
|
||||
assert result == []
|
||||
@@ -185,12 +236,12 @@ class TestMCPServiceGetMCPServers:
|
||||
'mode': entity.mode,
|
||||
}
|
||||
)
|
||||
ap.tool_mgr = None
|
||||
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_mcp_servers()
|
||||
result = await service.get_mcp_servers(_CONTEXT)
|
||||
|
||||
# Verify
|
||||
assert len(result) == 2
|
||||
@@ -215,21 +266,115 @@ class TestMCPServiceGetMCPServers:
|
||||
)
|
||||
ap.tool_mgr = SimpleNamespace()
|
||||
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
|
||||
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=None)
|
||||
runtime_session = SimpleNamespace(get_runtime_info_dict=Mock(return_value={'status': 'connected'}))
|
||||
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=runtime_session)
|
||||
|
||||
service = MCPService(ap)
|
||||
service.get_runtime_info = AsyncMock(return_value={'status': 'connected'})
|
||||
service = _service(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_mcp_servers(contain_runtime_info=True)
|
||||
result = await service.get_mcp_servers(_CONTEXT, contain_runtime_info=True)
|
||||
|
||||
# Verify - runtime info included
|
||||
assert result[0]['runtime_info'] == {'status': 'connected'}
|
||||
|
||||
async def test_resource_view_list_and_detail_redact_secrets_without_mutating_raw_data(self):
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
server = _create_mock_mcp_server(name='Secret Server')
|
||||
serialized = {
|
||||
'uuid': 'secret-uuid',
|
||||
'name': 'Secret Server',
|
||||
'enable': True,
|
||||
'extra_args': {
|
||||
'url': (
|
||||
'https://mcp-user:mcp-password@mcp.invalid/connect'
|
||||
'?token=url-secret&transport=streamable&sig=signed-secret'
|
||||
),
|
||||
'headers': {
|
||||
'Authorization': 'Bearer top-secret',
|
||||
'X-API-Key': 'api-secret',
|
||||
'Accept': 'application/json',
|
||||
},
|
||||
'env': {
|
||||
'ACCESS_TOKEN': 'access-secret',
|
||||
'TOKENIZER': 'public-model-name',
|
||||
},
|
||||
'credentials': {
|
||||
'username': 'service-user',
|
||||
'password': 'password-secret',
|
||||
},
|
||||
'public_key': 'public-value',
|
||||
},
|
||||
}
|
||||
original = copy.deepcopy(serialized)
|
||||
ap.persistence_mgr.execute_async = AsyncMock(
|
||||
side_effect=[
|
||||
_create_mock_result([server]),
|
||||
_create_mock_result(first_item=server),
|
||||
]
|
||||
)
|
||||
ap.persistence_mgr.serialize_model = Mock(return_value=serialized)
|
||||
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
|
||||
service = _service(ap)
|
||||
|
||||
listed = await service.get_mcp_servers(_VIEWER_CONTEXT)
|
||||
detail = await service.get_mcp_server_by_name(_VIEWER_CONTEXT, 'Secret Server')
|
||||
|
||||
for response in (listed[0], detail):
|
||||
assert response['extra_args']['url'] == (
|
||||
'https://***@mcp.invalid/connect?token=***&transport=streamable&sig=***'
|
||||
)
|
||||
assert response['extra_args']['headers'] == {
|
||||
'Authorization': '***',
|
||||
'X-API-Key': '***',
|
||||
'Accept': 'application/json',
|
||||
}
|
||||
assert response['extra_args']['env'] == {
|
||||
'ACCESS_TOKEN': '***',
|
||||
'TOKENIZER': 'public-model-name',
|
||||
}
|
||||
assert response['extra_args']['credentials'] == {
|
||||
'username': '***',
|
||||
'password': '***',
|
||||
}
|
||||
assert response['extra_args']['public_key'] == 'public-value'
|
||||
assert serialized == original
|
||||
|
||||
async def test_redacted_url_roundtrip_restores_persisted_credentials(self):
|
||||
persisted = {
|
||||
'extra_args': {'url': 'https://mcp-user:mcp-password@mcp.invalid/connect?token=url-secret&transport=http'}
|
||||
}
|
||||
|
||||
submitted = redact_mcp_secrets(persisted)
|
||||
|
||||
assert submitted['extra_args']['url'] == 'https://***@mcp.invalid/connect?token=***&transport=http'
|
||||
assert restore_mcp_secret_placeholders(submitted, persisted) == persisted
|
||||
|
||||
|
||||
class TestMCPServiceCreateMCPServer:
|
||||
"""Tests for create_mcp_server method."""
|
||||
|
||||
async def test_create_stdio_rejected_by_independent_instance_gate(self):
|
||||
ap = SimpleNamespace(
|
||||
instance_config=SimpleNamespace(
|
||||
data={
|
||||
'mcp': {'stdio': {'enabled': False}},
|
||||
'system': {'limitation': {'max_extensions': -1}},
|
||||
}
|
||||
),
|
||||
persistence_mgr=SimpleNamespace(execute_async=AsyncMock()),
|
||||
tool_mgr=None,
|
||||
)
|
||||
service = _service(ap)
|
||||
|
||||
with pytest.raises(MCPStdioDisabledError, match='disabled by instance policy'):
|
||||
await service.create_mcp_server(
|
||||
_CONTEXT,
|
||||
{'name': 'local', 'mode': 'stdio', 'enable': True, 'extra_args': {}},
|
||||
)
|
||||
|
||||
ap.persistence_mgr.execute_async.assert_not_awaited()
|
||||
|
||||
async def test_create_mcp_server_max_extensions_reached_raises(self):
|
||||
"""Raises ValueError when max_extensions limit reached."""
|
||||
# Setup
|
||||
@@ -241,16 +386,20 @@ class TestMCPServiceCreateMCPServer:
|
||||
ap.plugin_connector.list_plugins = AsyncMock(return_value=[Mock(), Mock()]) # 2 plugins
|
||||
|
||||
# Mock get_mcp_servers to return 0 servers (2 plugins already)
|
||||
mock_result = _create_mock_result([])
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
|
||||
ap.persistence_mgr.execute_async = AsyncMock(
|
||||
side_effect=[
|
||||
_create_mock_result(scalar_value=0),
|
||||
_create_mock_result(scalar_value=2),
|
||||
]
|
||||
)
|
||||
ap.persistence_mgr.serialize_model = Mock(return_value={})
|
||||
ap.tool_mgr = None
|
||||
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute & Verify - 2 plugins + new server would exceed limit
|
||||
with pytest.raises(ValueError, match='Maximum number of extensions'):
|
||||
await service.create_mcp_server({'name': 'New Server'})
|
||||
await service.create_mcp_server(_CONTEXT, {'name': 'New Server'})
|
||||
|
||||
async def test_create_mcp_server_no_limit(self):
|
||||
"""Creates MCP server without limit when max_extensions=-1."""
|
||||
@@ -271,10 +420,10 @@ class TestMCPServiceCreateMCPServer:
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
|
||||
ap.persistence_mgr.serialize_model = Mock(return_value={'uuid': 'new-uuid'})
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute
|
||||
server_uuid = await service.create_mcp_server({'name': 'New Server'})
|
||||
server_uuid = await service.create_mcp_server(_CONTEXT, {'name': 'New Server'})
|
||||
|
||||
# Verify
|
||||
assert server_uuid is not None
|
||||
@@ -293,11 +442,11 @@ class TestMCPServiceCreateMCPServer:
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result(first_item=existing_server))
|
||||
ap.persistence_mgr.serialize_model = Mock(return_value={})
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute & Verify
|
||||
with pytest.raises(ValueError, match='MCP server already exists: Existing Server'):
|
||||
await service.create_mcp_server({'name': 'Existing Server'})
|
||||
await service.create_mcp_server(_CONTEXT, {'name': 'Existing Server'})
|
||||
|
||||
async def test_create_mcp_server_loads_server(self):
|
||||
"""Loads server into tool_mgr when enabled."""
|
||||
@@ -330,14 +479,62 @@ class TestMCPServiceCreateMCPServer:
|
||||
return_value={'uuid': 'new-uuid', 'name': 'New Server', 'enable': True}
|
||||
)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute
|
||||
await service.create_mcp_server({'name': 'New Server', 'enable': True})
|
||||
await service.create_mcp_server(_CONTEXT, {'name': 'New Server', 'enable': True})
|
||||
|
||||
# Verify - host_mcp_server was called
|
||||
ap.tool_mgr.mcp_tool_loader.host_mcp_server.assert_called_once()
|
||||
|
||||
async def test_create_mcp_server_does_not_start_host_until_transaction_commits(self):
|
||||
"""The Runtime must not observe a server row that can still roll back."""
|
||||
|
||||
gate = asyncio.get_running_loop().create_future()
|
||||
|
||||
class PersistenceManagerStub:
|
||||
def create_after_commit_gate(self):
|
||||
return gate
|
||||
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = PersistenceManagerStub()
|
||||
ap.instance_config = SimpleNamespace(data={'system': {'limitation': {'max_extensions': -1}}})
|
||||
observed = []
|
||||
|
||||
async def host_mcp_server(context, config):
|
||||
observed.append((context, config))
|
||||
|
||||
ap.tool_mgr = SimpleNamespace(
|
||||
mcp_tool_loader=SimpleNamespace(
|
||||
host_mcp_server=host_mcp_server,
|
||||
_hosted_mcp_tasks=[],
|
||||
)
|
||||
)
|
||||
server_entity = _create_mock_mcp_server(server_uuid='new-uuid', enable=True)
|
||||
results = [
|
||||
_create_mock_result([]),
|
||||
Mock(),
|
||||
_create_mock_result(first_item=server_entity),
|
||||
]
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=results)
|
||||
ap.persistence_mgr.serialize_model = Mock(
|
||||
return_value={'uuid': 'new-uuid', 'name': 'New Server', 'enable': True}
|
||||
)
|
||||
service = _service(ap)
|
||||
|
||||
await service.create_mcp_server(_CONTEXT, {'name': 'New Server', 'enable': True})
|
||||
await asyncio.sleep(0)
|
||||
assert observed == []
|
||||
|
||||
gate.set_result(None)
|
||||
await ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks[0]
|
||||
assert observed == [
|
||||
(
|
||||
_CONTEXT,
|
||||
{'uuid': 'new-uuid', 'name': 'New Server', 'enable': True},
|
||||
)
|
||||
]
|
||||
|
||||
async def test_create_mcp_server_disabled_no_load(self):
|
||||
"""Does not load server when disabled."""
|
||||
# Setup
|
||||
@@ -351,10 +548,10 @@ class TestMCPServiceCreateMCPServer:
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
|
||||
ap.persistence_mgr.serialize_model = Mock(return_value={'uuid': 'new-uuid'})
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute with enable=False
|
||||
server_uuid = await service.create_mcp_server({'name': 'New Server', 'enable': False})
|
||||
server_uuid = await service.create_mcp_server(_CONTEXT, {'name': 'New Server', 'enable': False})
|
||||
|
||||
# Verify - no tool_mgr load attempt
|
||||
assert server_uuid is not None
|
||||
@@ -379,13 +576,11 @@ class TestMCPServiceGetMCPServerByName:
|
||||
'runtime_info': None,
|
||||
}
|
||||
)
|
||||
ap.tool_mgr = None
|
||||
|
||||
service = MCPService(ap)
|
||||
service.get_runtime_info = AsyncMock(return_value=None)
|
||||
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
|
||||
|
||||
service = _service(ap)
|
||||
# Execute
|
||||
result = await service.get_mcp_server_by_name('Found Server')
|
||||
result = await service.get_mcp_server_by_name(_CONTEXT, 'Found Server')
|
||||
|
||||
# Verify
|
||||
assert result is not None
|
||||
@@ -400,10 +595,10 @@ class TestMCPServiceGetMCPServerByName:
|
||||
mock_result = _create_mock_result(first_item=None)
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_mcp_server_by_name('Nonexistent Server')
|
||||
result = await service.get_mcp_server_by_name(_CONTEXT, 'Nonexistent Server')
|
||||
|
||||
# Verify
|
||||
assert result is None
|
||||
@@ -421,8 +616,10 @@ class TestMCPServiceUpdateMCPServer:
|
||||
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
|
||||
ap.tool_mgr.mcp_tool_loader.sessions = {'Old Server': Mock()}
|
||||
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
|
||||
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=True)
|
||||
|
||||
old_server = _create_mock_mcp_server(name='Old Server', enable=True)
|
||||
updated_server = _create_mock_mcp_server(name='Old Server', enable=False)
|
||||
|
||||
call_count = 0
|
||||
|
||||
@@ -431,14 +628,23 @@ class TestMCPServiceUpdateMCPServer:
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
return _create_mock_result(first_item=old_server)
|
||||
return Mock() # Update
|
||||
if call_count == 2:
|
||||
return _create_mock_result()
|
||||
return _create_mock_result(first_item=updated_server)
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
ap.persistence_mgr.serialize_model = Mock(
|
||||
side_effect=lambda _model, entity: {
|
||||
'uuid': 'test-uuid',
|
||||
'name': entity.name,
|
||||
'enable': entity.enable,
|
||||
}
|
||||
)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute - disable server
|
||||
await service.update_mcp_server('test-uuid', {'enable': False})
|
||||
await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': False})
|
||||
|
||||
# Verify - server was removed
|
||||
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once()
|
||||
@@ -453,6 +659,7 @@ class TestMCPServiceUpdateMCPServer:
|
||||
ap.tool_mgr.mcp_tool_loader.sessions = {}
|
||||
ap.tool_mgr.mcp_tool_loader.host_mcp_server = AsyncMock()
|
||||
ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks = []
|
||||
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=False)
|
||||
|
||||
old_server = _create_mock_mcp_server(name='Old Server', enable=False)
|
||||
|
||||
@@ -474,10 +681,10 @@ class TestMCPServiceUpdateMCPServer:
|
||||
return_value={'uuid': 'test-uuid', 'name': 'Old Server', 'enable': True}
|
||||
)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute - enable server
|
||||
await service.update_mcp_server('test-uuid', {'enable': True})
|
||||
await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': True})
|
||||
|
||||
# Verify - server was loaded
|
||||
ap.tool_mgr.mcp_tool_loader.host_mcp_server.assert_called_once()
|
||||
@@ -493,6 +700,7 @@ class TestMCPServiceUpdateMCPServer:
|
||||
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
|
||||
ap.tool_mgr.mcp_tool_loader.host_mcp_server = AsyncMock()
|
||||
ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks = []
|
||||
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=True)
|
||||
|
||||
old_server = _create_mock_mcp_server(name='Old Server', enable=True)
|
||||
|
||||
@@ -510,13 +718,13 @@ class TestMCPServiceUpdateMCPServer:
|
||||
return_value={'uuid': 'test-uuid', 'name': 'Old Server', 'enable': True}
|
||||
)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute - update enabled server (keep enabled, update extra_args)
|
||||
await service.update_mcp_server('test-uuid', {'enable': True, 'extra_args': {'new': 'args'}})
|
||||
await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': True, 'extra_args': {'new': 'args'}})
|
||||
|
||||
# Verify - remove and reload
|
||||
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with('Old Server')
|
||||
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with(_CONTEXT, 'Old Server')
|
||||
ap.tool_mgr.mcp_tool_loader.host_mcp_server.assert_called_once()
|
||||
|
||||
async def test_update_mcp_server_no_tool_mgr(self):
|
||||
@@ -541,15 +749,99 @@ class TestMCPServiceUpdateMCPServer:
|
||||
return Mock() # Update
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
ap.persistence_mgr.serialize_model = Mock(
|
||||
return_value={
|
||||
'uuid': 'test-uuid',
|
||||
'name': 'Server',
|
||||
'enable': True,
|
||||
}
|
||||
)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute - should not raise
|
||||
await service.update_mcp_server('test-uuid', {'name': 'New Name'})
|
||||
await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': False})
|
||||
|
||||
# Verify - persistence was called
|
||||
assert ap.persistence_mgr.execute_async.call_count >= 2
|
||||
|
||||
async def test_update_restores_existing_masked_secrets_and_preserves_explicit_changes(self):
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=None)
|
||||
old_server = _create_mock_mcp_server(name='Server', enable=True)
|
||||
old_data = {
|
||||
'uuid': 'test-uuid',
|
||||
'name': 'Server',
|
||||
'enable': True,
|
||||
'mode': 'streamable_http',
|
||||
'extra_args': {
|
||||
'headers': {
|
||||
'Authorization': 'Bearer original-secret',
|
||||
'X-API-Key': 'original-api-key',
|
||||
'Cookie': 'original-cookie',
|
||||
}
|
||||
},
|
||||
}
|
||||
captured_updates = []
|
||||
|
||||
async def mock_execute(statement):
|
||||
if not captured_updates:
|
||||
captured_updates.append(None)
|
||||
return _create_mock_result(first_item=old_server)
|
||||
captured_updates[0] = statement
|
||||
return _create_mock_result()
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
ap.persistence_mgr.serialize_model = Mock(return_value=old_data)
|
||||
service = _service(ap)
|
||||
|
||||
await service.update_mcp_server(
|
||||
_CONTEXT,
|
||||
'test-uuid',
|
||||
{
|
||||
'extra_args': {
|
||||
'headers': {
|
||||
'Authorization': '***',
|
||||
'X-API-Key': 'replacement-api-key',
|
||||
'Cookie': '',
|
||||
}
|
||||
}
|
||||
},
|
||||
)
|
||||
|
||||
persisted = captured_updates[0].compile().params['extra_args']
|
||||
assert persisted['headers'] == {
|
||||
'Authorization': 'Bearer original-secret',
|
||||
'X-API-Key': 'replacement-api-key',
|
||||
'Cookie': '',
|
||||
}
|
||||
|
||||
async def test_update_rejects_masked_secret_without_existing_value(self):
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.tool_mgr = SimpleNamespace(mcp_tool_loader=None)
|
||||
old_server = _create_mock_mcp_server(name='Server', enable=True)
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result(first_item=old_server))
|
||||
ap.persistence_mgr.serialize_model = Mock(
|
||||
return_value={
|
||||
'uuid': 'test-uuid',
|
||||
'name': 'Server',
|
||||
'enable': True,
|
||||
'extra_args': {'headers': {'Accept': 'application/json'}},
|
||||
}
|
||||
)
|
||||
service = _service(ap)
|
||||
|
||||
with pytest.raises(ValueError, match='Masked MCP secret has no existing value'):
|
||||
await service.update_mcp_server(
|
||||
_CONTEXT,
|
||||
'test-uuid',
|
||||
{'extra_args': {'headers': {'Authorization': '***'}}},
|
||||
)
|
||||
|
||||
assert ap.persistence_mgr.execute_async.await_count == 1
|
||||
|
||||
|
||||
class TestMCPServiceDeleteMCPServer:
|
||||
"""Tests for delete_mcp_server method."""
|
||||
@@ -563,6 +855,7 @@ class TestMCPServiceDeleteMCPServer:
|
||||
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
|
||||
ap.tool_mgr.mcp_tool_loader.sessions = {'Server to Delete': Mock()}
|
||||
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
|
||||
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=True)
|
||||
|
||||
server = _create_mock_mcp_server(name='Server to Delete')
|
||||
|
||||
@@ -576,14 +869,21 @@ class TestMCPServiceDeleteMCPServer:
|
||||
return Mock() # Delete
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
ap.persistence_mgr.serialize_model = Mock(
|
||||
return_value={
|
||||
'uuid': 'test-uuid',
|
||||
'name': 'Server to Delete',
|
||||
'enable': True,
|
||||
}
|
||||
)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute
|
||||
await service.delete_mcp_server('test-uuid')
|
||||
await service.delete_mcp_server(_CONTEXT, 'test-uuid')
|
||||
|
||||
# Verify
|
||||
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with('Server to Delete')
|
||||
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with(_CONTEXT, 'Server to Delete')
|
||||
ap.persistence_mgr.execute_async.assert_called()
|
||||
|
||||
async def test_delete_mcp_server_not_in_sessions(self):
|
||||
@@ -595,6 +895,7 @@ class TestMCPServiceDeleteMCPServer:
|
||||
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
|
||||
ap.tool_mgr.mcp_tool_loader.sessions = {} # Server not in sessions
|
||||
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
|
||||
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=False)
|
||||
|
||||
server = _create_mock_mcp_server(name='Not in Sessions')
|
||||
|
||||
@@ -608,11 +909,18 @@ class TestMCPServiceDeleteMCPServer:
|
||||
return Mock()
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
ap.persistence_mgr.serialize_model = Mock(
|
||||
return_value={
|
||||
'uuid': 'test-uuid',
|
||||
'name': 'Not in Sessions',
|
||||
'enable': True,
|
||||
}
|
||||
)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute
|
||||
await service.delete_mcp_server('test-uuid')
|
||||
await service.delete_mcp_server(_CONTEXT, 'test-uuid')
|
||||
|
||||
# Verify - remove not called (server not in sessions)
|
||||
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_not_called()
|
||||
@@ -626,6 +934,7 @@ class TestMCPServiceDeleteMCPServer:
|
||||
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
|
||||
ap.tool_mgr.mcp_tool_loader.sessions = {}
|
||||
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
|
||||
ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=False)
|
||||
|
||||
# No server found
|
||||
call_count = 0
|
||||
@@ -639,18 +948,35 @@ class TestMCPServiceDeleteMCPServer:
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
# Execute - should not raise
|
||||
await service.delete_mcp_server('nonexistent-uuid')
|
||||
with pytest.raises(WorkspaceNotFoundError, match='MCP server not found'):
|
||||
await service.delete_mcp_server(_CONTEXT, 'nonexistent-uuid')
|
||||
|
||||
# Verify - delete was called regardless
|
||||
ap.persistence_mgr.execute_async.assert_called()
|
||||
assert ap.persistence_mgr.execute_async.await_count == 1
|
||||
|
||||
|
||||
class TestMCPServiceTestMCPServer:
|
||||
"""Tests for test_mcp_server method."""
|
||||
|
||||
async def test_transient_stdio_test_rejected_by_instance_gate(self):
|
||||
ap = SimpleNamespace(
|
||||
instance_config=SimpleNamespace(data={'mcp': {'stdio': {'enabled': False}}}),
|
||||
tool_mgr=SimpleNamespace(mcp_tool_loader=SimpleNamespace(load_mcp_server=AsyncMock())),
|
||||
task_mgr=SimpleNamespace(create_user_task=Mock()),
|
||||
)
|
||||
service = _service(ap)
|
||||
|
||||
with pytest.raises(MCPStdioDisabledError, match='disabled by instance policy'):
|
||||
await service.test_mcp_server(
|
||||
_CONTEXT,
|
||||
'_',
|
||||
{'name': 'local', 'mode': 'stdio', 'enable': True, 'extra_args': {}},
|
||||
)
|
||||
|
||||
ap.tool_mgr.mcp_tool_loader.load_mcp_server.assert_not_awaited()
|
||||
ap.task_mgr.create_user_task.assert_not_called()
|
||||
|
||||
async def test_test_mcp_server_existing_server(self):
|
||||
"""Tests existing MCP server connection."""
|
||||
# Setup
|
||||
@@ -667,12 +993,18 @@ class TestMCPServiceTestMCPServer:
|
||||
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=mock_session)
|
||||
|
||||
ap.task_mgr = SimpleNamespace()
|
||||
ap.task_mgr.create_user_task = Mock(return_value=SimpleNamespace(id=123))
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
service._require_server = AsyncMock(return_value=(_CONTEXT, {'name': 'existing-server'}))
|
||||
|
||||
def create_user_task(coroutine, **_kwargs):
|
||||
coroutine.close()
|
||||
return SimpleNamespace(id=123)
|
||||
|
||||
ap.task_mgr.create_user_task = Mock(side_effect=create_user_task)
|
||||
|
||||
# Execute
|
||||
task_id = await service.test_mcp_server('existing-server', {})
|
||||
task_id = await service.test_mcp_server(_CONTEXT, 'existing-server', {})
|
||||
|
||||
# Verify - returns task ID
|
||||
assert task_id == 123
|
||||
@@ -685,11 +1017,12 @@ class TestMCPServiceTestMCPServer:
|
||||
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
|
||||
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=None)
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
service._require_server = AsyncMock(side_effect=WorkspaceNotFoundError('MCP server not found'))
|
||||
|
||||
# Execute & Verify
|
||||
with pytest.raises(ValueError, match='Server not found'):
|
||||
await service.test_mcp_server('nonexistent-server', {})
|
||||
with pytest.raises(WorkspaceNotFoundError, match='MCP server not found'):
|
||||
await service.test_mcp_server(_CONTEXT, 'nonexistent-server', {})
|
||||
|
||||
async def test_test_mcp_server_new_server(self):
|
||||
"""Tests new MCP server with underscore name."""
|
||||
@@ -703,13 +1036,38 @@ class TestMCPServiceTestMCPServer:
|
||||
ap.tool_mgr.mcp_tool_loader.load_mcp_server = AsyncMock(return_value=mock_session)
|
||||
|
||||
ap.task_mgr = SimpleNamespace()
|
||||
ap.task_mgr.create_user_task = Mock(return_value=SimpleNamespace(id=456))
|
||||
|
||||
service = MCPService(ap)
|
||||
service = _service(ap)
|
||||
|
||||
def create_user_task(coroutine, **_kwargs):
|
||||
coroutine.close()
|
||||
return SimpleNamespace(id=456)
|
||||
|
||||
ap.task_mgr.create_user_task = Mock(side_effect=create_user_task)
|
||||
|
||||
# Execute with '_' name (new server)
|
||||
task_id = await service.test_mcp_server('_', {'name': 'New Server'})
|
||||
task_id = await service.test_mcp_server(_CONTEXT, '_', {'name': 'New Server'})
|
||||
|
||||
# Verify - load_mcp_server called
|
||||
ap.tool_mgr.mcp_tool_loader.load_mcp_server.assert_called_once()
|
||||
assert task_id == 456
|
||||
|
||||
async def test_rejected_transient_test_session_is_shut_down(self):
|
||||
ap = SimpleNamespace()
|
||||
mock_session = MagicMock()
|
||||
mock_session.shutdown = AsyncMock()
|
||||
ap.tool_mgr = SimpleNamespace(
|
||||
mcp_tool_loader=SimpleNamespace(load_mcp_server=AsyncMock(return_value=mock_session))
|
||||
)
|
||||
|
||||
def reject(coroutine, **_kwargs):
|
||||
coroutine.close()
|
||||
raise TaskCapacityError('capacity')
|
||||
|
||||
ap.task_mgr = SimpleNamespace(create_user_task=Mock(side_effect=reject))
|
||||
service = _service(ap)
|
||||
|
||||
with pytest.raises(TaskCapacityError, match='capacity'):
|
||||
await service.test_mcp_server(_CONTEXT, '_', {'name': 'New Server'})
|
||||
|
||||
mock_session.shutdown.assert_awaited_once_with()
|
||||
|
||||
@@ -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,392 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import sqlalchemy
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
from langbot.pkg.api.http.authz import WorkspaceRequiredError
|
||||
from langbot.pkg.api.http.context import ExecutionContext
|
||||
from langbot.pkg.api.http.service.monitoring import MonitoringService
|
||||
from langbot.pkg.entity.persistence.base import Base
|
||||
from langbot.pkg.entity.persistence.monitoring import MonitoringLLMCall, MonitoringMessage
|
||||
from langbot.pkg.entity.persistence.workspace import Workspace
|
||||
from langbot.pkg.persistence.mgr import PersistenceManager
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
|
||||
WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
|
||||
|
||||
|
||||
def _context(workspace_uuid: str) -> ExecutionContext:
|
||||
return ExecutionContext(
|
||||
instance_uuid='instance',
|
||||
workspace_uuid=workspace_uuid,
|
||||
placement_generation=3,
|
||||
bot_uuid='same-bot',
|
||||
pipeline_uuid='same-pipeline',
|
||||
)
|
||||
|
||||
|
||||
class _PersistenceManager:
|
||||
def __init__(self, engine):
|
||||
self.engine = engine
|
||||
|
||||
async def execute_async(self, *args, **kwargs):
|
||||
async with self.engine.connect() as connection:
|
||||
result = await connection.execute(*args, **kwargs)
|
||||
await connection.commit()
|
||||
return result
|
||||
|
||||
def get_db_engine(self):
|
||||
return self.engine
|
||||
|
||||
@staticmethod
|
||||
def serialize_model(model, data, masked_columns=None):
|
||||
return {
|
||||
column.name: (
|
||||
getattr(data, column.name).isoformat()
|
||||
if isinstance(getattr(data, column.name), datetime.datetime)
|
||||
else getattr(data, column.name)
|
||||
)
|
||||
for column in model.__table__.columns
|
||||
if column.name not in (masked_columns or [])
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def service(tmp_path):
|
||||
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "monitoring.db"}')
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(Workspace),
|
||||
[
|
||||
{
|
||||
'uuid': WORKSPACE_A,
|
||||
'instance_uuid': 'instance',
|
||||
'name': 'A',
|
||||
'slug': 'a',
|
||||
'source': 'cloud_projection',
|
||||
},
|
||||
{
|
||||
'uuid': WORKSPACE_B,
|
||||
'instance_uuid': 'instance',
|
||||
'name': 'B',
|
||||
'slug': 'b',
|
||||
'source': 'cloud_projection',
|
||||
},
|
||||
],
|
||||
)
|
||||
application = SimpleNamespace(
|
||||
persistence_mgr=_PersistenceManager(engine),
|
||||
instance_config=SimpleNamespace(data={'database': {'use': 'sqlite'}}),
|
||||
)
|
||||
yield MonitoringService(application)
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def _record_message(service, context, content):
|
||||
return await service.record_message(
|
||||
context,
|
||||
bot_id='same-bot',
|
||||
bot_name='Same Bot',
|
||||
pipeline_id='same-pipeline',
|
||||
pipeline_name='Same Pipeline',
|
||||
message_content=content,
|
||||
session_id='same-session',
|
||||
)
|
||||
|
||||
|
||||
async def test_monitoring_write_without_execution_context_fails_closed(service):
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
await _record_message(service, None, 'unscoped')
|
||||
|
||||
|
||||
async def test_same_session_and_resource_ids_do_not_collide(service):
|
||||
context_a = _context(WORKSPACE_A)
|
||||
context_b = _context(WORKSPACE_B)
|
||||
message_a = await _record_message(service, context_a, 'tenant-a')
|
||||
message_b = await _record_message(service, context_b, 'tenant-b')
|
||||
await service.record_session_start(
|
||||
context_a,
|
||||
session_id='same-session',
|
||||
bot_id='same-bot',
|
||||
bot_name='Same Bot',
|
||||
pipeline_id='same-pipeline',
|
||||
pipeline_name='Same Pipeline',
|
||||
)
|
||||
await service.record_session_start(
|
||||
context_b,
|
||||
session_id='same-session',
|
||||
bot_id='same-bot',
|
||||
bot_name='Same Bot',
|
||||
pipeline_id='same-pipeline',
|
||||
pipeline_name='Same Pipeline',
|
||||
)
|
||||
|
||||
messages_a, total_a = await service.get_messages(context_a)
|
||||
messages_b, total_b = await service.get_messages(context_b)
|
||||
assert total_a == total_b == 1
|
||||
assert messages_a[0]['message_content'] == 'tenant-a'
|
||||
assert messages_b[0]['message_content'] == 'tenant-b'
|
||||
assert (await service.get_message_details(context_b, message_a))['found'] is False
|
||||
assert (await service.get_message_details(context_a, message_b))['found'] is False
|
||||
|
||||
|
||||
async def test_tool_call_inherits_context_from_connection_message_row(service):
|
||||
context = _context(WORKSPACE_A)
|
||||
message_id = await _record_message(service, context, 'tool context')
|
||||
|
||||
await service.record_tool_call(
|
||||
context,
|
||||
tool_name='search',
|
||||
tool_source='native',
|
||||
duration=12,
|
||||
message_id=message_id,
|
||||
)
|
||||
|
||||
tool_calls, total = await service.get_tool_calls(context)
|
||||
assert total == 1
|
||||
assert tool_calls[0]['bot_id'] == 'same-bot'
|
||||
assert tool_calls[0]['pipeline_id'] == 'same-pipeline'
|
||||
assert tool_calls[0]['session_id'] == 'same-session'
|
||||
assert tool_calls[0]['message_id'] == message_id
|
||||
|
||||
|
||||
async def test_feedback_upsert_and_cancel_are_workspace_scoped(service):
|
||||
context_a = _context(WORKSPACE_A)
|
||||
context_b = _context(WORKSPACE_B)
|
||||
await service.record_feedback(context_a, feedback_id='same-feedback', feedback_type=1)
|
||||
await service.record_feedback(context_b, feedback_id='same-feedback', feedback_type=2)
|
||||
|
||||
stats_a = await service.get_feedback_stats(context_a)
|
||||
stats_b = await service.get_feedback_stats(context_b)
|
||||
assert stats_a['total_likes'] == 1
|
||||
assert stats_a['total_dislikes'] == 0
|
||||
assert stats_b['total_likes'] == 0
|
||||
assert stats_b['total_dislikes'] == 1
|
||||
|
||||
await service.record_feedback(context_a, feedback_id='same-feedback', feedback_type=3)
|
||||
assert (await service.get_feedback_stats(context_a))['total_feedback'] == 0
|
||||
assert (await service.get_feedback_stats(context_b))['total_feedback'] == 1
|
||||
|
||||
|
||||
async def test_monitoring_queries_and_detail_views_are_strictly_bounded(service):
|
||||
context = _context(WORKSPACE_A)
|
||||
service.ap.instance_config.data['monitoring'] = {
|
||||
'query_limits': {
|
||||
'page_rows': 2,
|
||||
'export_rows': 2,
|
||||
'detail_rows': 2,
|
||||
'timeseries_buckets': 2,
|
||||
'max_offset': 10,
|
||||
}
|
||||
}
|
||||
await service.record_session_start(
|
||||
context,
|
||||
session_id='same-session',
|
||||
bot_id='same-bot',
|
||||
bot_name='Same Bot',
|
||||
pipeline_id='same-pipeline',
|
||||
pipeline_name='Same Pipeline',
|
||||
)
|
||||
message_ids = [await _record_message(service, context, f'message-{index}') for index in range(4)]
|
||||
for index in range(3):
|
||||
await service.record_llm_call(
|
||||
context,
|
||||
bot_id='same-bot',
|
||||
bot_name='Same Bot',
|
||||
pipeline_id='same-pipeline',
|
||||
pipeline_name='Same Pipeline',
|
||||
session_id='same-session',
|
||||
model_name='model',
|
||||
input_tokens=1,
|
||||
output_tokens=2,
|
||||
duration=10,
|
||||
message_id=message_ids[0],
|
||||
)
|
||||
await service.record_tool_call(
|
||||
context,
|
||||
tool_name=f'tool-{index}',
|
||||
tool_source='native',
|
||||
duration=5,
|
||||
session_id='same-session',
|
||||
message_id=message_ids[0],
|
||||
)
|
||||
await service.record_error(
|
||||
context,
|
||||
bot_id='same-bot',
|
||||
bot_name='Same Bot',
|
||||
pipeline_id='same-pipeline',
|
||||
pipeline_name='Same Pipeline',
|
||||
error_type='Failure',
|
||||
error_message=f'error-{index}',
|
||||
session_id='same-session',
|
||||
message_id=message_ids[0],
|
||||
)
|
||||
|
||||
page, total = await service.get_messages(context, limit=100000, offset=-5)
|
||||
exported = await service.export_messages(context, limit=100000)
|
||||
session_detail = await service.get_session_analysis(context, 'same-session')
|
||||
message_detail = await service.get_message_details(context, message_ids[0])
|
||||
|
||||
assert total == 4
|
||||
assert len(page) == 2
|
||||
assert len(exported) == 2
|
||||
assert session_detail['message_stats']['total'] == 4
|
||||
assert session_detail['llm_stats']['total_calls'] == 3
|
||||
assert session_detail['tool_stats']['total_calls'] == 3
|
||||
assert len(session_detail['tool_calls']) == 2
|
||||
assert len(session_detail['errors']) == 2
|
||||
assert session_detail['detail_truncated'] == {
|
||||
'tool_calls': True,
|
||||
'errors': True,
|
||||
}
|
||||
assert message_detail['llm_stats']['total_calls'] == 3
|
||||
assert len(message_detail['llm_calls']) == 2
|
||||
assert len(message_detail['errors']) == 2
|
||||
assert message_detail['detail_truncated'] == {
|
||||
'llm_calls': True,
|
||||
'errors': True,
|
||||
}
|
||||
|
||||
service.ap.instance_config.data['monitoring']['query_limits'] = {
|
||||
'page_rows': 999999,
|
||||
'export_rows': 999999,
|
||||
'detail_rows': 999999,
|
||||
'timeseries_buckets': 999999,
|
||||
'max_offset': 99999999,
|
||||
}
|
||||
assert service.normalize_page_window(999999, 99999999) == (5000, 10000000)
|
||||
assert service.normalize_export_limit(999999) == 50000
|
||||
assert service._detail_limit() == 10000
|
||||
assert service._timeseries_bucket_limit() == 10000
|
||||
|
||||
|
||||
async def test_token_statistics_aggregate_and_limit_groups_in_database(service):
|
||||
context = _context(WORKSPACE_A)
|
||||
service.ap.instance_config.data['monitoring'] = {
|
||||
'query_limits': {
|
||||
'page_rows': 1,
|
||||
'timeseries_buckets': 2,
|
||||
}
|
||||
}
|
||||
first_hour = datetime.datetime(2026, 7, 28, 10, 0)
|
||||
rows = [
|
||||
{
|
||||
'id': f'llm-{index}',
|
||||
'workspace_uuid': WORKSPACE_A,
|
||||
'timestamp': first_hour + datetime.timedelta(hours=hour, minutes=index),
|
||||
'model_name': model,
|
||||
'input_tokens': input_tokens,
|
||||
'output_tokens': output_tokens,
|
||||
'total_tokens': input_tokens + output_tokens,
|
||||
'duration': 100,
|
||||
'cost': 0.01,
|
||||
'status': 'success',
|
||||
'bot_id': 'same-bot',
|
||||
'bot_name': 'Same Bot',
|
||||
'pipeline_id': 'same-pipeline',
|
||||
'pipeline_name': 'Same Pipeline',
|
||||
'session_id': 'same-session',
|
||||
}
|
||||
for index, (hour, model, input_tokens, output_tokens) in enumerate(
|
||||
[
|
||||
(0, 'small-model', 1, 2),
|
||||
(1, 'large-model', 3, 4),
|
||||
(2, 'large-model', 5, 6),
|
||||
(2, 'large-model', 7, 8),
|
||||
]
|
||||
)
|
||||
]
|
||||
await service.ap.persistence_mgr.execute_async(sqlalchemy.insert(MonitoringLLMCall), rows)
|
||||
|
||||
stats = await service.get_token_statistics(context, bucket='hour')
|
||||
|
||||
assert stats['summary']['total_calls'] == 4
|
||||
assert stats['summary']['total_tokens'] == 36
|
||||
assert stats['by_model_truncated'] is True
|
||||
assert [model['model_name'] for model in stats['by_model']] == ['large-model']
|
||||
assert stats['timeseries_truncated'] is True
|
||||
assert stats['timeseries'] == [
|
||||
{
|
||||
'bucket': '2026-07-28 11:00',
|
||||
'input_tokens': 3,
|
||||
'output_tokens': 4,
|
||||
'total_tokens': 7,
|
||||
'calls': 1,
|
||||
},
|
||||
{
|
||||
'bucket': '2026-07-28 12:00',
|
||||
'input_tokens': 12,
|
||||
'output_tokens': 14,
|
||||
'total_tokens': 26,
|
||||
'calls': 2,
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
async def test_cleanup_commits_sqlite_delete_before_vacuum(tmp_path):
|
||||
engine = create_async_engine(
|
||||
f'sqlite+aiosqlite:///{tmp_path / "monitoring-cleanup.db"}',
|
||||
connect_args={'timeout': 0.1},
|
||||
)
|
||||
application = SimpleNamespace(
|
||||
instance_config=SimpleNamespace(data={'database': {'use': 'sqlite'}}),
|
||||
)
|
||||
manager = PersistenceManager(application)
|
||||
manager.db = SimpleNamespace(get_engine=lambda: engine)
|
||||
application.persistence_mgr = manager
|
||||
try:
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(Workspace).values(
|
||||
uuid=WORKSPACE_A,
|
||||
instance_uuid='instance',
|
||||
name='A',
|
||||
slug='a',
|
||||
source='cloud_projection',
|
||||
)
|
||||
)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(MonitoringMessage),
|
||||
[
|
||||
{
|
||||
'id': f'expired-message-{index}',
|
||||
'workspace_uuid': WORKSPACE_A,
|
||||
'timestamp': datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)
|
||||
- datetime.timedelta(days=30),
|
||||
'bot_id': 'bot',
|
||||
'bot_name': 'Bot',
|
||||
'pipeline_id': 'pipeline',
|
||||
'pipeline_name': 'Pipeline',
|
||||
'message_content': 'expired',
|
||||
'session_id': 'session',
|
||||
'status': 'success',
|
||||
'level': 'info',
|
||||
}
|
||||
for index in range(5)
|
||||
],
|
||||
)
|
||||
|
||||
deleted = await MonitoringService(application).cleanup_expired_records(
|
||||
_context(WORKSPACE_A),
|
||||
retention_days=1,
|
||||
batch_size=2,
|
||||
max_batches_per_table=1,
|
||||
)
|
||||
|
||||
assert deleted['monitoring_messages'] == 2
|
||||
async with engine.connect() as connection:
|
||||
remaining = await connection.scalar(
|
||||
sqlalchemy.select(sqlalchemy.func.count()).select_from(MonitoringMessage)
|
||||
)
|
||||
assert remaining == 3
|
||||
finally:
|
||||
await engine.dispose()
|
||||
@@ -21,10 +21,13 @@ import json
|
||||
|
||||
from langbot.pkg.api.http.service.pipeline import PipelineService, default_stage_order
|
||||
from langbot.pkg.entity.persistence.pipeline import LegacyPipeline
|
||||
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
WORKSPACE_UUID = 'workspace-a'
|
||||
|
||||
|
||||
def _create_mock_pipeline(
|
||||
pipeline_uuid: str = None,
|
||||
@@ -77,7 +80,9 @@ class TestPipelineServiceGetPipelineMetadata:
|
||||
service = PipelineService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_pipeline_metadata()
|
||||
result = await service.get_pipeline_metadata(
|
||||
WORKSPACE_UUID,
|
||||
)
|
||||
|
||||
# Verify
|
||||
assert len(result) == 4
|
||||
@@ -107,7 +112,9 @@ class TestPipelineServiceGetPipelines:
|
||||
service = PipelineService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_pipelines()
|
||||
result = await service.get_pipelines(
|
||||
WORKSPACE_UUID,
|
||||
)
|
||||
|
||||
# Verify
|
||||
assert result == []
|
||||
@@ -133,7 +140,9 @@ class TestPipelineServiceGetPipelines:
|
||||
service = PipelineService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_pipelines()
|
||||
result = await service.get_pipelines(
|
||||
WORKSPACE_UUID,
|
||||
)
|
||||
|
||||
# Verify
|
||||
assert len(result) == 2
|
||||
@@ -152,7 +161,7 @@ class TestPipelineServiceGetPipelines:
|
||||
service = PipelineService(ap)
|
||||
|
||||
# Execute
|
||||
await service.get_pipelines(sort_by='updated_at', sort_order='ASC')
|
||||
await service.get_pipelines(WORKSPACE_UUID, sort_by='updated_at', sort_order='ASC')
|
||||
|
||||
# Verify - execute was called with sort parameters
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
@@ -181,7 +190,7 @@ class TestPipelineServiceGetPipeline:
|
||||
service = PipelineService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_pipeline('test-uuid')
|
||||
result = await service.get_pipeline(WORKSPACE_UUID, 'test-uuid')
|
||||
|
||||
# Verify
|
||||
assert result is not None
|
||||
@@ -200,7 +209,7 @@ class TestPipelineServiceGetPipeline:
|
||||
service = PipelineService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_pipeline('nonexistent-uuid')
|
||||
result = await service.get_pipeline(WORKSPACE_UUID, 'nonexistent-uuid')
|
||||
|
||||
# Verify
|
||||
assert result is None
|
||||
@@ -229,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
|
||||
@@ -806,11 +862,11 @@ class TestModelProviderServiceScanProviderModels:
|
||||
ap.model_mgr.load_provider = AsyncMock(return_value=runtime_provider)
|
||||
ap.llm_model_service.get_llm_models_by_provider = AsyncMock(return_value=[])
|
||||
ap.embedding_models_service.get_embedding_models_by_provider = AsyncMock(return_value=[])
|
||||
ap.rerank_models_service.get_rerank_models_by_provider = AsyncMock(
|
||||
return_value=[{'name': 'Qwen3-Reranker-8B'}]
|
||||
)
|
||||
ap.rerank_models_service.get_rerank_models_by_provider = AsyncMock(return_value=[{'name': 'Qwen3-Reranker-8B'}])
|
||||
|
||||
result = await ModelProviderService(ap).scan_provider_models('rerank-scan-uuid', model_type='rerank')
|
||||
result = await ModelProviderService(ap).scan_provider_models(
|
||||
WORKSPACE_UUID, 'rerank-scan-uuid', model_type='rerank'
|
||||
)
|
||||
|
||||
assert result['models'][0]['type'] == 'rerank'
|
||||
assert result['models'][0]['already_added'] is True
|
||||
@@ -848,7 +904,7 @@ class TestModelProviderServiceScanProviderModels:
|
||||
|
||||
# Execute & Verify
|
||||
with pytest.raises(ValueError, match='current provider does not support model scanning'):
|
||||
await service.scan_provider_models('no-scan-uuid')
|
||||
await service.scan_provider_models(WORKSPACE_UUID, 'no-scan-uuid')
|
||||
|
||||
async def test_scan_provider_marks_already_added_models(self):
|
||||
"""Marks models that are already added."""
|
||||
@@ -898,7 +954,7 @@ class TestModelProviderServiceScanProviderModels:
|
||||
service = ModelProviderService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.scan_provider_models('already-added-uuid')
|
||||
result = await service.scan_provider_models(WORKSPACE_UUID, 'already-added-uuid')
|
||||
|
||||
# Verify - existing model marked as already_added
|
||||
existing_model = next(m for m in result['models'] if m['name'] == 'Existing Model')
|
||||
@@ -906,3 +962,46 @@ class TestModelProviderServiceScanProviderModels:
|
||||
|
||||
new_model = next(m for m in result['models'] if m['name'] == 'New Model')
|
||||
assert new_model['already_added'] is False
|
||||
|
||||
|
||||
class TestProviderSecretRoundtrip:
|
||||
async def test_masked_api_keys_update_preserves_existing_values(self):
|
||||
write_result = Mock(rowcount=1)
|
||||
ap = SimpleNamespace(
|
||||
persistence_mgr=SimpleNamespace(execute_async=AsyncMock(return_value=write_result)),
|
||||
model_mgr=SimpleNamespace(reload_provider=AsyncMock()),
|
||||
)
|
||||
service = ModelProviderService(ap)
|
||||
service.get_provider = AsyncMock(
|
||||
return_value={
|
||||
'uuid': 'provider-secret',
|
||||
'api_keys': ['first-secret', 'second-secret'],
|
||||
}
|
||||
)
|
||||
|
||||
await service.update_provider(
|
||||
WORKSPACE_UUID,
|
||||
'provider-secret',
|
||||
{'name': 'Updated', 'api_keys': ['***', 'replacement-secret']},
|
||||
)
|
||||
|
||||
statement = ap.persistence_mgr.execute_async.await_args.args[0]
|
||||
stored_api_keys = next(value.value for column, value in statement._values.items() if column.key == 'api_keys')
|
||||
assert stored_api_keys == ['first-secret', 'replacement-secret']
|
||||
|
||||
async def test_extra_masked_api_key_is_rejected(self):
|
||||
ap = SimpleNamespace(
|
||||
persistence_mgr=SimpleNamespace(execute_async=AsyncMock()),
|
||||
model_mgr=SimpleNamespace(reload_provider=AsyncMock()),
|
||||
)
|
||||
service = ModelProviderService(ap)
|
||||
service.get_provider = AsyncMock(return_value={'uuid': 'provider-secret', 'api_keys': ['only-secret']})
|
||||
|
||||
with pytest.raises(ValueError, match='no existing value'):
|
||||
await service.update_provider(
|
||||
WORKSPACE_UUID,
|
||||
'provider-secret',
|
||||
{'api_keys': ['***', '***']},
|
||||
)
|
||||
|
||||
ap.persistence_mgr.execute_async.assert_not_awaited()
|
||||
|
||||
@@ -0,0 +1,88 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
|
||||
import pytest
|
||||
|
||||
from langbot.pkg.api.http.service.secrets import (
|
||||
contains_secret_placeholder,
|
||||
redact_secrets,
|
||||
restore_secret_placeholders,
|
||||
)
|
||||
|
||||
|
||||
RAW_CONFIG = {
|
||||
'apiKey': 'api-secret',
|
||||
'dify_apikey': 'dify-secret',
|
||||
'base_url': (
|
||||
'https://service-user:service-password@api.invalid/v1'
|
||||
'?api_key=query-secret®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,10 @@ Source: src/langbot/pkg/api/http/service/space.py
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections import OrderedDict
|
||||
import json
|
||||
from urllib.parse import parse_qs, urlsplit
|
||||
|
||||
import pytest
|
||||
from unittest.mock import AsyncMock, Mock, patch, MagicMock
|
||||
from types import SimpleNamespace
|
||||
@@ -26,6 +30,23 @@ from langbot.pkg.entity.persistence.user import User
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
def _set_response_body(response: MagicMock, body: dict | str) -> None:
|
||||
"""Configure an aiohttp-like streaming body on an HTTP response mock."""
|
||||
|
||||
raw_body = body.encode() if isinstance(body, str) else json.dumps(body).encode()
|
||||
|
||||
class Content:
|
||||
async def iter_chunked(self, _chunk_size: int):
|
||||
midpoint = max(len(raw_body) // 2, 1)
|
||||
yield raw_body[:midpoint]
|
||||
if midpoint < len(raw_body):
|
||||
yield raw_body[midpoint:]
|
||||
|
||||
response.headers = {}
|
||||
response.content = Content()
|
||||
response.charset = 'utf-8'
|
||||
|
||||
|
||||
def _create_mock_user(
|
||||
email: str = 'test@example.com',
|
||||
account_type: str = 'space',
|
||||
@@ -73,7 +94,7 @@ class TestSpaceServiceGetOAuthAuthorizeUrl:
|
||||
result = service.get_oauth_authorize_url('http://localhost/callback')
|
||||
|
||||
# Verify
|
||||
assert 'redirect_uri=http://localhost/callback' in result
|
||||
assert parse_qs(urlsplit(result).query)['redirect_uri'] == ['http://localhost/callback']
|
||||
assert 'https://space.langbot.app/auth/authorize' in result
|
||||
|
||||
def test_get_oauth_authorize_url_with_state(self):
|
||||
@@ -93,8 +114,9 @@ class TestSpaceServiceGetOAuthAuthorizeUrl:
|
||||
result = service.get_oauth_authorize_url('http://localhost/callback', state='random_state')
|
||||
|
||||
# Verify
|
||||
assert 'redirect_uri=http://localhost/callback' in result
|
||||
assert 'state=random_state' in result
|
||||
params = parse_qs(urlsplit(result).query)
|
||||
assert params['redirect_uri'] == ['http://localhost/callback']
|
||||
assert params['state'] == ['random_state']
|
||||
|
||||
def test_get_oauth_authorize_url_default_config(self):
|
||||
"""Uses default OAuth URL when config not set."""
|
||||
@@ -289,6 +311,40 @@ class TestSpaceServiceGetCredits:
|
||||
# Verify - returns cached value without API call
|
||||
assert result == 100
|
||||
|
||||
async def test_cached_credit_lookup_does_not_scan_all_users(self):
|
||||
ap = SimpleNamespace()
|
||||
ap.instance_config = SimpleNamespace(data={})
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
service = SpaceService(ap)
|
||||
|
||||
class AtMostOneStepOrderedDict(OrderedDict):
|
||||
def __iter__(self):
|
||||
iterator = super().__iter__()
|
||||
yielded = False
|
||||
|
||||
def next_entry():
|
||||
nonlocal yielded
|
||||
if yielded:
|
||||
raise AssertionError('credits cache scanned all users')
|
||||
yielded = True
|
||||
return next(iterator)
|
||||
|
||||
class AtMostOneStepIterator:
|
||||
def __iter__(self):
|
||||
return self
|
||||
|
||||
def __next__(self):
|
||||
return next_entry()
|
||||
|
||||
return AtMostOneStepIterator()
|
||||
|
||||
now = time.time()
|
||||
service._credits_cache = AtMostOneStepOrderedDict(
|
||||
(f'user-{index}@example.com', (index, now)) for index in range(512)
|
||||
)
|
||||
|
||||
assert await service.get_credits('user-511@example.com') == 511
|
||||
|
||||
async def test_get_credits_cache_expired_refreshes(self):
|
||||
"""Refreshes expired cache."""
|
||||
# Setup
|
||||
@@ -403,6 +459,7 @@ class TestSpaceServiceRefreshToken:
|
||||
},
|
||||
}
|
||||
)
|
||||
_set_response_body(mock_response, mock_response.json.return_value)
|
||||
|
||||
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
|
||||
mock_session_obj = MagicMock()
|
||||
@@ -438,6 +495,7 @@ class TestSpaceServiceRefreshToken:
|
||||
}
|
||||
)
|
||||
mock_response.text = AsyncMock(return_value='{"code":1,"msg":"Invalid refresh token"}')
|
||||
_set_response_body(mock_response, mock_response.json.return_value)
|
||||
|
||||
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
|
||||
mock_session_obj = MagicMock()
|
||||
@@ -464,6 +522,7 @@ class TestSpaceServiceRefreshToken:
|
||||
mock_response = MagicMock()
|
||||
mock_response.status = 500
|
||||
mock_response.text = AsyncMock(return_value='Internal Server Error')
|
||||
_set_response_body(mock_response, mock_response.text.return_value)
|
||||
|
||||
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
|
||||
mock_session_obj = MagicMock()
|
||||
@@ -503,6 +562,7 @@ class TestSpaceServiceExchangeOAuthCode:
|
||||
},
|
||||
}
|
||||
)
|
||||
_set_response_body(mock_response, mock_response.json.return_value)
|
||||
|
||||
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
|
||||
mock_session_obj = MagicMock()
|
||||
@@ -532,6 +592,7 @@ class TestSpaceServiceExchangeOAuthCode:
|
||||
mock_response.status = 200
|
||||
mock_response.json = AsyncMock(return_value={'code': 1, 'msg': 'Invalid code'})
|
||||
mock_response.text = AsyncMock(return_value='{"code":1,"msg":"Invalid code"}')
|
||||
_set_response_body(mock_response, mock_response.json.return_value)
|
||||
|
||||
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
|
||||
mock_session_obj = MagicMock()
|
||||
@@ -570,6 +631,7 @@ class TestSpaceServiceGetUserInfoRaw:
|
||||
},
|
||||
}
|
||||
)
|
||||
_set_response_body(mock_response, mock_response.json.return_value)
|
||||
|
||||
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
|
||||
mock_session_obj = MagicMock()
|
||||
@@ -600,6 +662,7 @@ class TestSpaceServiceGetUserInfoRaw:
|
||||
mock_response.status = 200
|
||||
mock_response.json = AsyncMock(return_value={'code': 1, 'msg': 'Unauthorized'})
|
||||
mock_response.text = AsyncMock(return_value='{"code":1,"msg":"Unauthorized"}')
|
||||
_set_response_body(mock_response, mock_response.json.return_value)
|
||||
|
||||
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
|
||||
mock_session_obj = MagicMock()
|
||||
@@ -700,6 +763,7 @@ class TestSpaceServiceGetModels:
|
||||
},
|
||||
}
|
||||
)
|
||||
_set_response_body(mock_response, mock_response.json.return_value)
|
||||
|
||||
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
|
||||
mock_session_obj = MagicMock()
|
||||
@@ -730,6 +794,7 @@ class TestSpaceServiceGetModels:
|
||||
mock_response.status = 200
|
||||
mock_response.json = AsyncMock(return_value={'code': 1, 'msg': 'Unauthorized'})
|
||||
mock_response.text = AsyncMock(return_value='{"code":1,"msg":"Unauthorized"}')
|
||||
_set_response_body(mock_response, mock_response.json.return_value)
|
||||
|
||||
with patch('langbot.pkg.api.http.service.space.httpclient.get_session') as mock_session:
|
||||
mock_session_obj = MagicMock()
|
||||
|
||||
@@ -14,17 +14,124 @@ Source: src/langbot/pkg/api/http/service/user.py
|
||||
from __future__ import annotations
|
||||
|
||||
import pytest
|
||||
import jwt
|
||||
import datetime
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
from types import SimpleNamespace
|
||||
|
||||
from langbot.pkg.api.http.service.user import UserService
|
||||
from langbot.pkg.entity.persistence.user import User
|
||||
from langbot.pkg.entity.errors.account import AccountEmailMismatchError
|
||||
from langbot.pkg.api.http.service.user import (
|
||||
ControlPlaneDirectoryRequiredError,
|
||||
UserService,
|
||||
)
|
||||
from langbot.pkg.entity.persistence.user import AccountSource, AccountStatus, User
|
||||
from langbot.pkg.entity.errors.account import (
|
||||
AccountEmailMismatchError,
|
||||
SpaceAccountBindingRequiredError,
|
||||
SpaceAccountNotRegisteredError,
|
||||
)
|
||||
from langbot.pkg.utils.bounded_executor import BlockingWorkCapacityError
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
async def test_password_hashing_rejects_concurrent_waiters() -> None:
|
||||
service = UserService(SimpleNamespace())
|
||||
await service._password_hash_lock.acquire()
|
||||
try:
|
||||
with pytest.raises(
|
||||
BlockingWorkCapacityError,
|
||||
match='Password hashing capacity reached',
|
||||
):
|
||||
await service._hash_password('secret')
|
||||
finally:
|
||||
service._password_hash_lock.release()
|
||||
|
||||
|
||||
class TestSpaceOAuthState:
|
||||
async def test_login_state_is_opaque_single_use(self):
|
||||
service = UserService(SimpleNamespace())
|
||||
|
||||
state = await service.issue_space_oauth_state('login')
|
||||
|
||||
assert state.count('.') == 0
|
||||
assert await service.consume_space_oauth_state(state, 'login') is None
|
||||
with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
|
||||
await service.consume_space_oauth_state(state, 'login')
|
||||
|
||||
async def test_bind_state_resolves_only_bound_active_account(self):
|
||||
service = UserService(SimpleNamespace())
|
||||
account = SimpleNamespace(uuid='account-a', status=AccountStatus.ACTIVE.value)
|
||||
service.get_user_by_uuid = AsyncMock(return_value=account)
|
||||
|
||||
state = await service.issue_space_oauth_state('bind', account_uuid='account-a')
|
||||
|
||||
assert await service.consume_space_oauth_state(state, 'bind') is account
|
||||
service.get_user_by_uuid.assert_awaited_once_with('account-a')
|
||||
|
||||
async def test_state_purpose_mismatch_is_rejected_and_consumed(self):
|
||||
service = UserService(SimpleNamespace())
|
||||
state = await service.issue_space_oauth_state('login')
|
||||
|
||||
with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
|
||||
await service.consume_space_oauth_state(state, 'bind')
|
||||
with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
|
||||
await service.consume_space_oauth_state(state, 'login')
|
||||
|
||||
async def test_expired_state_is_rejected(self):
|
||||
service = UserService(SimpleNamespace())
|
||||
state = await service.issue_space_oauth_state('login')
|
||||
digest = service._space_oauth_state_digest(state)
|
||||
purpose, account_uuid, _, launch_workspace_uuid = service._space_oauth_states[digest]
|
||||
service._space_oauth_states[digest] = (purpose, account_uuid, 0, launch_workspace_uuid)
|
||||
|
||||
with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
|
||||
await service.consume_space_oauth_state(state, 'login')
|
||||
|
||||
async def test_login_state_can_carry_launch_workspace_without_changing_normal_return(self):
|
||||
service = UserService(SimpleNamespace())
|
||||
state = await service.issue_space_oauth_state(
|
||||
'login',
|
||||
launch_workspace_uuid='workspace-a',
|
||||
)
|
||||
|
||||
assert await service.consume_space_oauth_state(state, 'login') is None
|
||||
|
||||
state = await service.issue_space_oauth_state(
|
||||
'login',
|
||||
launch_workspace_uuid='workspace-a',
|
||||
)
|
||||
consumed = await service.consume_space_oauth_state_details(state, 'login')
|
||||
assert consumed.account is None
|
||||
assert consumed.launch_workspace_uuid == 'workspace-a'
|
||||
|
||||
async def test_issue_state_does_not_scan_all_live_states(self, monkeypatch):
|
||||
service = UserService(SimpleNamespace())
|
||||
for _ in range(512):
|
||||
await service.issue_space_oauth_state('login')
|
||||
|
||||
class NoGlobalIterationDict(dict):
|
||||
def __iter__(self):
|
||||
raise AssertionError('OAuth state issuance scanned all live states')
|
||||
|
||||
def keys(self):
|
||||
raise AssertionError('OAuth state issuance scanned all live states')
|
||||
|
||||
def items(self):
|
||||
raise AssertionError('OAuth state issuance scanned all live states')
|
||||
|
||||
def values(self):
|
||||
raise AssertionError('OAuth state issuance scanned all live states')
|
||||
|
||||
guarded_states = NoGlobalIterationDict(service._space_oauth_states)
|
||||
monkeypatch.setattr(service, '_space_oauth_states', guarded_states)
|
||||
|
||||
state = await service.issue_space_oauth_state('login')
|
||||
|
||||
assert await service.consume_space_oauth_state(state, 'login') is None
|
||||
assert len(guarded_states) == 512
|
||||
|
||||
|
||||
def _create_mock_user(
|
||||
email: str = 'test@example.com',
|
||||
password: str = 'hashed_password',
|
||||
@@ -34,6 +141,7 @@ def _create_mock_user(
|
||||
"""Helper to create mock User entity."""
|
||||
user = Mock(spec=User)
|
||||
user.user = email
|
||||
user.uuid = f'account-{email}'
|
||||
user.password = password
|
||||
user.account_type = account_type
|
||||
user.space_account_uuid = space_account_uuid
|
||||
@@ -102,6 +210,41 @@ class TestUserServiceIsInitialized:
|
||||
assert result is False
|
||||
|
||||
|
||||
class TestUserServiceGetLoginCapabilities:
|
||||
"""Tests for public login capability discovery."""
|
||||
|
||||
async def test_uses_explicit_identity_discovery_scope(self):
|
||||
discovery_result = Mock()
|
||||
discovery_result.one = Mock(return_value=(1, 2))
|
||||
discovery_session = SimpleNamespace(execute=AsyncMock(return_value=discovery_result))
|
||||
|
||||
class DiscoveryContext:
|
||||
async def __aenter__(self):
|
||||
return SimpleNamespace(session=discovery_session)
|
||||
|
||||
async def __aexit__(self, exc_type, exc, tb):
|
||||
return False
|
||||
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace(
|
||||
current_session=Mock(return_value=None),
|
||||
identity_discovery_uow=Mock(return_value=DiscoveryContext()),
|
||||
execute_async=AsyncMock(side_effect=AssertionError('unscoped persistence access')),
|
||||
)
|
||||
ap.workspace_service = SimpleNamespace(instance_uuid='instance-a')
|
||||
service = UserService(ap)
|
||||
|
||||
result = await service.get_login_capabilities()
|
||||
|
||||
assert result == {
|
||||
'password_login_enabled': True,
|
||||
'space_login_enabled': True,
|
||||
}
|
||||
ap.persistence_mgr.identity_discovery_uow.assert_called_once()
|
||||
discovery_session.execute.assert_awaited_once()
|
||||
ap.persistence_mgr.execute_async.assert_not_awaited()
|
||||
|
||||
|
||||
class TestUserServiceGetUserByEmail:
|
||||
"""Tests for get_user_by_email method."""
|
||||
|
||||
@@ -309,6 +452,50 @@ class TestUserServiceVerifyJwtToken:
|
||||
with pytest.raises(Exception): # jwt.DecodeError or similar
|
||||
await service.verify_jwt_token('invalid.token.here')
|
||||
|
||||
async def test_verify_jwt_token_rejects_foreign_audience(self):
|
||||
ap = SimpleNamespace()
|
||||
ap.instance_config = SimpleNamespace()
|
||||
ap.instance_config.data = {'system': {'jwt': {'secret': 'test_secret', 'expire': 3600}}}
|
||||
service = UserService(ap)
|
||||
token = jwt.encode(
|
||||
{
|
||||
'user': 'verify@example.com',
|
||||
'iss': 'langbot-core',
|
||||
'aud': 'langbot-instance:another-instance',
|
||||
'exp': datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(hours=1),
|
||||
},
|
||||
'test_secret',
|
||||
algorithm='HS256',
|
||||
)
|
||||
|
||||
with pytest.raises(jwt.InvalidAudienceError):
|
||||
await service.verify_jwt_token(token)
|
||||
|
||||
async def test_verify_jwt_token_accepts_legacy_community_token_only_in_oss(self):
|
||||
ap = SimpleNamespace()
|
||||
ap.instance_config = SimpleNamespace()
|
||||
ap.instance_config.data = {'system': {'jwt': {'secret': 'test_secret', 'expire': 3600}}}
|
||||
ap.workspace_service = SimpleNamespace(
|
||||
instance_uuid='instance-a',
|
||||
policy=SimpleNamespace(multi_workspace_enabled=False),
|
||||
)
|
||||
service = UserService(ap)
|
||||
legacy_token = jwt.encode(
|
||||
{
|
||||
'user': 'legacy@example.com',
|
||||
'iss': 'LangBot-community',
|
||||
'exp': datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(hours=1),
|
||||
},
|
||||
'test_secret',
|
||||
algorithm='HS256',
|
||||
)
|
||||
|
||||
assert await service.verify_jwt_token(legacy_token) == 'legacy@example.com'
|
||||
|
||||
ap.workspace_service.policy.multi_workspace_enabled = True
|
||||
with pytest.raises(jwt.MissingRequiredClaimError):
|
||||
await service.verify_jwt_token(legacy_token)
|
||||
|
||||
|
||||
class TestUserServiceResetPassword:
|
||||
"""Tests for reset_password method."""
|
||||
@@ -476,6 +663,71 @@ class TestUserServiceCreateOrUpdateSpaceUser:
|
||||
ap.persistence_mgr.execute_async.assert_called()
|
||||
assert updated_user.space_account_uuid == 'existing-space-uuid'
|
||||
|
||||
async def test_cloud_login_updates_only_the_projected_space_account(self):
|
||||
projected = SimpleNamespace(
|
||||
uuid='projected-space-uuid',
|
||||
user='Cloud Owner',
|
||||
normalized_email='owner@example.com',
|
||||
password='',
|
||||
account_type='space',
|
||||
status=AccountStatus.ACTIVE.value,
|
||||
source=AccountSource.CLOUD_PROJECTION.value,
|
||||
projection_revision=7,
|
||||
space_account_uuid='projected-space-uuid',
|
||||
)
|
||||
persistence = SimpleNamespace(execute_async=AsyncMock())
|
||||
ap = SimpleNamespace(
|
||||
persistence_mgr=persistence,
|
||||
workspace_service=SimpleNamespace(policy=SimpleNamespace(multi_workspace_enabled=True)),
|
||||
)
|
||||
service = UserService(ap)
|
||||
service.get_user_by_space_account_uuid = AsyncMock(side_effect=[projected, projected])
|
||||
|
||||
result = await service.create_or_update_space_user(
|
||||
space_account_uuid='projected-space-uuid',
|
||||
email='OWNER@example.com',
|
||||
access_token='access-token',
|
||||
refresh_token='refresh-token',
|
||||
api_key='api-key',
|
||||
expires_in=3600,
|
||||
)
|
||||
|
||||
assert result is projected
|
||||
persistence.execute_async.assert_awaited_once()
|
||||
|
||||
async def test_cloud_login_never_creates_an_unprojected_account(self):
|
||||
persistence = SimpleNamespace(execute_async=AsyncMock())
|
||||
ap = SimpleNamespace(
|
||||
persistence_mgr=persistence,
|
||||
workspace_service=SimpleNamespace(policy=SimpleNamespace(multi_workspace_enabled=True)),
|
||||
)
|
||||
service = UserService(ap)
|
||||
service.get_user_by_space_account_uuid = AsyncMock(return_value=None)
|
||||
|
||||
with pytest.raises(
|
||||
ControlPlaneDirectoryRequiredError,
|
||||
match='verified Cloud directory',
|
||||
):
|
||||
await service.create_or_update_space_user(
|
||||
space_account_uuid='unknown-space-uuid',
|
||||
email='unknown@example.com',
|
||||
access_token='access-token',
|
||||
refresh_token='refresh-token',
|
||||
api_key='api-key',
|
||||
expires_in=3600,
|
||||
)
|
||||
|
||||
persistence.execute_async.assert_not_awaited()
|
||||
|
||||
async def test_cloud_invitation_registration_requires_space_identity(self):
|
||||
ap = SimpleNamespace(
|
||||
workspace_service=SimpleNamespace(policy=SimpleNamespace(multi_workspace_enabled=True)),
|
||||
)
|
||||
service = UserService(ap)
|
||||
|
||||
with pytest.raises(ControlPlaneDirectoryRequiredError, match='Space account'):
|
||||
await service.register_invited_account('invite-token', 'member@example.com', 'password')
|
||||
|
||||
async def test_create_or_update_new_space_user_first_init(self):
|
||||
"""Creates new Space user on first initialization."""
|
||||
# Setup
|
||||
@@ -522,8 +774,8 @@ class TestUserServiceCreateOrUpdateSpaceUser:
|
||||
# Verify
|
||||
assert result.space_account_uuid == 'new-space-uuid'
|
||||
|
||||
async def test_create_or_update_space_user_already_initialized_raises_error(self):
|
||||
"""Raises AccountEmailMismatchError when system already initialized and user not found."""
|
||||
async def test_create_or_update_space_user_already_initialized_reports_unknown_space_email(self):
|
||||
"""Unknown Space email is distinct from an existing local Account collision."""
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
@@ -538,7 +790,7 @@ class TestUserServiceCreateOrUpdateSpaceUser:
|
||||
service.is_initialized = AsyncMock(return_value=True) # Already initialized
|
||||
|
||||
# Execute & Verify
|
||||
with pytest.raises(AccountEmailMismatchError):
|
||||
with pytest.raises(SpaceAccountNotRegisteredError):
|
||||
await service.create_or_update_space_user(
|
||||
space_account_uuid='unknown-space-uuid',
|
||||
email='unknown@example.com',
|
||||
@@ -548,6 +800,78 @@ class TestUserServiceCreateOrUpdateSpaceUser:
|
||||
expires_in=3600,
|
||||
)
|
||||
|
||||
async def test_unknown_space_subject_cannot_claim_existing_account_by_email(self):
|
||||
"""An OAuth login collision requires the explicit account-bound bind flow."""
|
||||
existing_user = _create_mock_user(
|
||||
email='owner@example.com',
|
||||
account_type='local',
|
||||
space_account_uuid=None,
|
||||
)
|
||||
ap = SimpleNamespace(
|
||||
persistence_mgr=SimpleNamespace(execute_async=AsyncMock()),
|
||||
provider_service=SimpleNamespace(update_space_model_provider_api_keys=AsyncMock()),
|
||||
space_service=SimpleNamespace(
|
||||
get_user_info_raw=AsyncMock(
|
||||
return_value={
|
||||
'account': {
|
||||
'uuid': 'attacker-space-subject',
|
||||
'email': 'owner@example.com',
|
||||
},
|
||||
'api_key': 'attacker-api-key',
|
||||
}
|
||||
)
|
||||
),
|
||||
)
|
||||
service = UserService(ap)
|
||||
service.get_user_by_space_account_uuid = AsyncMock(return_value=None)
|
||||
service.get_user_by_email = AsyncMock(return_value=existing_user)
|
||||
service.generate_jwt_token = AsyncMock(return_value='must-not-be-issued')
|
||||
|
||||
with pytest.raises(SpaceAccountBindingRequiredError):
|
||||
await service.authenticate_space_user(
|
||||
'attacker-access-token',
|
||||
'attacker-refresh-token',
|
||||
3600,
|
||||
)
|
||||
|
||||
ap.persistence_mgr.execute_async.assert_not_awaited()
|
||||
ap.provider_service.update_space_model_provider_api_keys.assert_not_awaited()
|
||||
service.generate_jwt_token.assert_not_awaited()
|
||||
|
||||
async def test_oss_space_provider_refresh_requires_workspace_owner(self):
|
||||
member_account = _create_mock_user(email='member@example.com', space_account_uuid='space-member')
|
||||
access = SimpleNamespace(
|
||||
workspace=SimpleNamespace(uuid='workspace-a'),
|
||||
membership=SimpleNamespace(role='admin'),
|
||||
)
|
||||
provider_service = SimpleNamespace(update_space_model_provider_api_keys=AsyncMock())
|
||||
ap = SimpleNamespace(
|
||||
workspace_service=SimpleNamespace(policy=SimpleNamespace(multi_workspace_enabled=False)),
|
||||
workspace_collaboration_service=SimpleNamespace(list_account_workspaces=AsyncMock(return_value=[access])),
|
||||
provider_service=provider_service,
|
||||
)
|
||||
|
||||
await UserService(ap)._update_space_provider_for_account(member_account, 'member-api-key')
|
||||
|
||||
provider_service.update_space_model_provider_api_keys.assert_not_awaited()
|
||||
|
||||
async def test_oss_space_provider_refresh_uses_workspace_owner_credentials(self):
|
||||
owner_account = _create_mock_user(email='owner@example.com', space_account_uuid='space-owner')
|
||||
access = SimpleNamespace(
|
||||
workspace=SimpleNamespace(uuid='workspace-a'),
|
||||
membership=SimpleNamespace(role='owner'),
|
||||
)
|
||||
provider_service = SimpleNamespace(update_space_model_provider_api_keys=AsyncMock())
|
||||
ap = SimpleNamespace(
|
||||
workspace_service=SimpleNamespace(policy=SimpleNamespace(multi_workspace_enabled=False)),
|
||||
workspace_collaboration_service=SimpleNamespace(list_account_workspaces=AsyncMock(return_value=[access])),
|
||||
provider_service=provider_service,
|
||||
)
|
||||
|
||||
await UserService(ap)._update_space_provider_for_account(owner_account, 'owner-api-key')
|
||||
|
||||
provider_service.update_space_model_provider_api_keys.assert_awaited_once_with('workspace-a', 'owner-api-key')
|
||||
|
||||
async def test_create_or_update_space_user_no_expiry(self):
|
||||
"""Creates Space user without token expiry."""
|
||||
# Setup
|
||||
@@ -594,6 +918,58 @@ class TestUserServiceCreateOrUpdateSpaceUser:
|
||||
assert result is not None
|
||||
assert result.space_account_uuid == 'noexpiry-uuid'
|
||||
|
||||
async def test_bind_space_account_rejects_different_email(self):
|
||||
service = UserService(SimpleNamespace())
|
||||
service.get_user_by_email = AsyncMock(return_value=_create_mock_user(email='invited@example.com'))
|
||||
service.ap.space_service = SimpleNamespace(
|
||||
exchange_oauth_code=AsyncMock(
|
||||
return_value={'access_token': 'access', 'refresh_token': 'refresh', 'expires_in': 3600}
|
||||
),
|
||||
get_user_info_raw=AsyncMock(
|
||||
return_value={
|
||||
'account': {'uuid': 'space-other', 'email': 'other@example.com'},
|
||||
'api_key': 'key',
|
||||
}
|
||||
),
|
||||
)
|
||||
service.get_user_by_space_account_uuid = AsyncMock(return_value=None)
|
||||
service._identity_execute = AsyncMock()
|
||||
|
||||
with pytest.raises(AccountEmailMismatchError):
|
||||
await service.bind_space_account('invited@example.com', 'code')
|
||||
|
||||
service._identity_execute.assert_not_awaited()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_workspace_owner_returns_user_object_from_core_connection_result(self):
|
||||
service = UserService(SimpleNamespace())
|
||||
owner = _create_mock_user('owner@example.com', password='pw')
|
||||
service.ap.persistence_mgr = SimpleNamespace(
|
||||
current_session=lambda: SimpleNamespace(scalar=AsyncMock(return_value=owner)),
|
||||
)
|
||||
|
||||
resolved = await service.get_workspace_owner('workspace-1')
|
||||
|
||||
assert resolved is owner
|
||||
|
||||
|
||||
class TestUserServiceLoginCapabilities:
|
||||
async def test_capabilities_are_derived_from_all_accounts(self):
|
||||
result = SimpleNamespace(one=lambda: (2, 1))
|
||||
ap = SimpleNamespace(persistence_mgr=SimpleNamespace(execute_async=AsyncMock(return_value=result)))
|
||||
|
||||
capabilities = await UserService(ap).get_login_capabilities()
|
||||
|
||||
assert capabilities == {'password_login_enabled': True, 'space_login_enabled': True}
|
||||
|
||||
async def test_capabilities_disable_absent_login_methods(self):
|
||||
result = SimpleNamespace(one=lambda: (0, 0))
|
||||
ap = SimpleNamespace(persistence_mgr=SimpleNamespace(execute_async=AsyncMock(return_value=result)))
|
||||
|
||||
capabilities = await UserService(ap).get_login_capabilities()
|
||||
|
||||
assert capabilities == {'password_login_enabled': False, 'space_login_enabled': False}
|
||||
|
||||
|
||||
class TestUserServiceCreateUserLock:
|
||||
"""Tests for create_user_lock attribute."""
|
||||
|
||||
@@ -14,15 +14,23 @@ Source: src/langbot/pkg/api/http/service/webhook.py
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import datetime
|
||||
|
||||
import pytest
|
||||
import sqlalchemy
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
from types import SimpleNamespace
|
||||
|
||||
from langbot.pkg.api.http.authz import WorkspaceRequiredError
|
||||
from langbot.pkg.api.http.service.webhook import WebhookService
|
||||
from langbot.pkg.entity.persistence.base import Base
|
||||
from langbot.pkg.entity.persistence.webhook import Webhook
|
||||
from langbot.pkg.entity.persistence.workspace import Workspace
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
WORKSPACE_UUID = 'workspace-a'
|
||||
|
||||
|
||||
def _create_mock_webhook(
|
||||
@@ -42,11 +50,20 @@ def _create_mock_webhook(
|
||||
return webhook
|
||||
|
||||
|
||||
def _create_mock_result(items: list = None, first_item=None):
|
||||
def _create_mock_result(items: list = None, first_item=None, scalar_value=None):
|
||||
"""Create mock result object for persistence queries."""
|
||||
result = Mock()
|
||||
result.all = Mock(return_value=items or [])
|
||||
result.first = Mock(return_value=first_item)
|
||||
result.scalar = Mock(return_value=scalar_value)
|
||||
result.rowcount = 1
|
||||
return result
|
||||
|
||||
|
||||
def _create_write_result(rowcount: int = 1, inserted_id: int = 1):
|
||||
result = Mock()
|
||||
result.rowcount = rowcount
|
||||
result.inserted_primary_key = [inserted_id]
|
||||
return result
|
||||
|
||||
|
||||
@@ -71,7 +88,7 @@ class TestWebhookServiceGetWebhooks:
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_webhooks()
|
||||
result = await service.get_webhooks(WORKSPACE_UUID)
|
||||
|
||||
# Verify
|
||||
assert result == []
|
||||
@@ -100,7 +117,7 @@ class TestWebhookServiceGetWebhooks:
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_webhooks()
|
||||
result = await service.get_webhooks(WORKSPACE_UUID)
|
||||
|
||||
# Verify
|
||||
assert len(result) == 2
|
||||
@@ -119,6 +136,7 @@ class TestWebhookServiceCreateWebhook:
|
||||
|
||||
# Mock insert result
|
||||
insert_result = Mock()
|
||||
insert_result.inserted_primary_key = [1]
|
||||
|
||||
# Mock select result for retrieving created webhook
|
||||
created_webhook = _create_mock_webhook(
|
||||
@@ -137,6 +155,8 @@ class TestWebhookServiceCreateWebhook:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
return _create_mock_result(scalar_value=0) # Count
|
||||
if call_count == 2:
|
||||
return insert_result # Insert
|
||||
return select_result # Select
|
||||
|
||||
@@ -155,6 +175,7 @@ class TestWebhookServiceCreateWebhook:
|
||||
|
||||
# Execute
|
||||
result = await service.create_webhook(
|
||||
WORKSPACE_UUID,
|
||||
name='New Webhook',
|
||||
url='http://new.example.com/webhook',
|
||||
description='New Description',
|
||||
@@ -187,7 +208,9 @@ class TestWebhookServiceCreateWebhook:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
return Mock() # Insert
|
||||
return _create_mock_result(scalar_value=0)
|
||||
if call_count == 2:
|
||||
return _create_write_result() # Insert
|
||||
return _create_mock_result(first_item=created_webhook)
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
@@ -204,7 +227,11 @@ class TestWebhookServiceCreateWebhook:
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute - only name and url required
|
||||
result = await service.create_webhook(name='Minimal Webhook', url='http://minimal.example.com')
|
||||
result = await service.create_webhook(
|
||||
WORKSPACE_UUID,
|
||||
name='Minimal Webhook',
|
||||
url='http://minimal.example.com',
|
||||
)
|
||||
|
||||
# Verify defaults
|
||||
assert result['description'] == ''
|
||||
@@ -224,7 +251,9 @@ class TestWebhookServiceCreateWebhook:
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
if call_count == 1:
|
||||
return Mock()
|
||||
return _create_mock_result(scalar_value=0)
|
||||
if call_count == 2:
|
||||
return _create_write_result()
|
||||
return _create_mock_result(first_item=created_webhook)
|
||||
|
||||
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
|
||||
@@ -233,11 +262,52 @@ class TestWebhookServiceCreateWebhook:
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.create_webhook(name='Disabled', url='http://disabled.com', enabled=False)
|
||||
result = await service.create_webhook(
|
||||
WORKSPACE_UUID,
|
||||
name='Disabled',
|
||||
url='http://disabled.com',
|
||||
enabled=False,
|
||||
)
|
||||
|
||||
# Verify
|
||||
assert result['enabled'] is False
|
||||
|
||||
async def test_create_webhook_rejects_workspace_at_capacity(self):
|
||||
ap = SimpleNamespace(
|
||||
instance_config=SimpleNamespace(
|
||||
data={'webhooks': {'max_per_workspace': 2}},
|
||||
),
|
||||
persistence_mgr=SimpleNamespace(
|
||||
execute_async=AsyncMock(return_value=_create_mock_result(scalar_value=2)),
|
||||
),
|
||||
)
|
||||
|
||||
service = WebhookService(ap)
|
||||
|
||||
with pytest.raises(ValueError, match=r'Maximum number of webhooks \(2\) reached'):
|
||||
await service.create_webhook(
|
||||
WORKSPACE_UUID,
|
||||
name='Too many',
|
||||
url='https://example.invalid',
|
||||
)
|
||||
|
||||
ap.persistence_mgr.execute_async.assert_awaited_once()
|
||||
|
||||
async def test_max_per_workspace_clamps_invalid_and_oversized_values(self):
|
||||
ap = SimpleNamespace(
|
||||
instance_config=SimpleNamespace(
|
||||
data={'webhooks': {'max_per_workspace': 999999}},
|
||||
)
|
||||
)
|
||||
service = WebhookService(ap)
|
||||
assert service.max_per_workspace() == 64
|
||||
|
||||
ap.instance_config.data['webhooks']['max_per_workspace'] = 0
|
||||
assert service.max_per_workspace() == 1
|
||||
|
||||
ap.instance_config.data['webhooks']['max_per_workspace'] = 'invalid'
|
||||
assert service.max_per_workspace() == 16
|
||||
|
||||
|
||||
class TestWebhookServiceGetWebhook:
|
||||
"""Tests for get_webhook method."""
|
||||
@@ -262,7 +332,7 @@ class TestWebhookServiceGetWebhook:
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_webhook(1)
|
||||
result = await service.get_webhook(WORKSPACE_UUID, 1)
|
||||
|
||||
# Verify
|
||||
assert result is not None
|
||||
@@ -281,7 +351,7 @@ class TestWebhookServiceGetWebhook:
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_webhook(999)
|
||||
result = await service.get_webhook(WORKSPACE_UUID, 999)
|
||||
|
||||
# Verify
|
||||
assert result is None
|
||||
@@ -298,7 +368,7 @@ class TestWebhookServiceGetWebhook:
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_webhook(0)
|
||||
result = await service.get_webhook(WORKSPACE_UUID, 0)
|
||||
|
||||
# Verify - should return None (no webhook with ID 0)
|
||||
assert result is None
|
||||
@@ -312,12 +382,12 @@ class TestWebhookServiceUpdateWebhook:
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
|
||||
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
await service.update_webhook(1, name='Updated Name')
|
||||
await service.update_webhook(WORKSPACE_UUID, 1, name='Updated Name')
|
||||
|
||||
# Verify
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
@@ -327,12 +397,12 @@ class TestWebhookServiceUpdateWebhook:
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
|
||||
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
await service.update_webhook(1, url='http://updated.example.com')
|
||||
await service.update_webhook(WORKSPACE_UUID, 1, url='http://updated.example.com')
|
||||
|
||||
# Verify
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
@@ -342,12 +412,12 @@ class TestWebhookServiceUpdateWebhook:
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
|
||||
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
await service.update_webhook(1, description='Updated description')
|
||||
await service.update_webhook(WORKSPACE_UUID, 1, description='Updated description')
|
||||
|
||||
# Verify
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
@@ -357,12 +427,12 @@ class TestWebhookServiceUpdateWebhook:
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
|
||||
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
await service.update_webhook(1, enabled=False)
|
||||
await service.update_webhook(WORKSPACE_UUID, 1, enabled=False)
|
||||
|
||||
# Verify
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
@@ -372,12 +442,13 @@ class TestWebhookServiceUpdateWebhook:
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
|
||||
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
await service.update_webhook(
|
||||
WORKSPACE_UUID,
|
||||
1,
|
||||
name='All Updated',
|
||||
url='http://all.updated.com',
|
||||
@@ -393,15 +464,17 @@ class TestWebhookServiceUpdateWebhook:
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
existing = _create_mock_webhook(webhook_id=1)
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result(first_item=existing))
|
||||
ap.persistence_mgr.serialize_model = Mock(return_value={'id': 1})
|
||||
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute - no update parameters
|
||||
await service.update_webhook(1)
|
||||
await service.update_webhook(WORKSPACE_UUID, 1)
|
||||
|
||||
# Verify - no execute call since no update_data
|
||||
ap.persistence_mgr.execute_async.assert_not_called()
|
||||
# No write is issued; one scoped existence lookup is performed.
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
|
||||
|
||||
class TestWebhookServiceDeleteWebhook:
|
||||
@@ -412,12 +485,12 @@ class TestWebhookServiceDeleteWebhook:
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
|
||||
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
await service.delete_webhook(1)
|
||||
await service.delete_webhook(WORKSPACE_UUID, 1)
|
||||
|
||||
# Verify
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
@@ -427,12 +500,12 @@ class TestWebhookServiceDeleteWebhook:
|
||||
# Setup
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = SimpleNamespace()
|
||||
ap.persistence_mgr.execute_async = AsyncMock()
|
||||
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result(rowcount=0))
|
||||
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute - should not raise
|
||||
await service.delete_webhook(999)
|
||||
await service.delete_webhook(WORKSPACE_UUID, 999)
|
||||
|
||||
# Verify - still called
|
||||
ap.persistence_mgr.execute_async.assert_called_once()
|
||||
@@ -453,7 +526,7 @@ class TestWebhookServiceGetEnabledWebhooks:
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_enabled_webhooks()
|
||||
result = await service.get_enabled_webhooks(WORKSPACE_UUID)
|
||||
|
||||
# Verify
|
||||
assert result == []
|
||||
@@ -481,7 +554,7 @@ class TestWebhookServiceGetEnabledWebhooks:
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_enabled_webhooks()
|
||||
result = await service.get_enabled_webhooks(WORKSPACE_UUID)
|
||||
|
||||
# Verify
|
||||
assert len(result) == 2
|
||||
@@ -501,7 +574,170 @@ class TestWebhookServiceGetEnabledWebhooks:
|
||||
service = WebhookService(ap)
|
||||
|
||||
# Execute
|
||||
result = await service.get_enabled_webhooks()
|
||||
result = await service.get_enabled_webhooks(WORKSPACE_UUID)
|
||||
|
||||
# Verify - should be empty (SQL would filter disabled)
|
||||
assert result == []
|
||||
|
||||
|
||||
ISOLATION_WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
|
||||
ISOLATION_WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
|
||||
|
||||
|
||||
class _RealPersistenceManager:
|
||||
def __init__(self, engine):
|
||||
self.engine = engine
|
||||
|
||||
async def execute_async(self, *args, **kwargs):
|
||||
async with self.engine.connect() as connection:
|
||||
result = await connection.execute(*args, **kwargs)
|
||||
await connection.commit()
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def serialize_model(model, data, masked_columns=None):
|
||||
return {
|
||||
column.name: (
|
||||
getattr(data, column.name).isoformat()
|
||||
if isinstance(getattr(data, column.name), datetime.datetime)
|
||||
else getattr(data, column.name)
|
||||
)
|
||||
for column in model.__table__.columns
|
||||
if column.name not in (masked_columns or [])
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
async def tenant_webhook_service(tmp_path):
|
||||
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "webhooks.db"}')
|
||||
async with engine.begin() as connection:
|
||||
await connection.run_sync(Base.metadata.create_all)
|
||||
await connection.execute(
|
||||
sqlalchemy.insert(Workspace),
|
||||
[
|
||||
{
|
||||
'uuid': ISOLATION_WORKSPACE_A,
|
||||
'instance_uuid': 'instance',
|
||||
'name': 'A',
|
||||
'slug': 'a',
|
||||
'source': 'cloud_projection',
|
||||
},
|
||||
{
|
||||
'uuid': ISOLATION_WORKSPACE_B,
|
||||
'instance_uuid': 'instance',
|
||||
'name': 'B',
|
||||
'slug': 'b',
|
||||
'source': 'cloud_projection',
|
||||
},
|
||||
],
|
||||
)
|
||||
service = WebhookService(SimpleNamespace(persistence_mgr=_RealPersistenceManager(engine)))
|
||||
yield service
|
||||
await engine.dispose()
|
||||
|
||||
|
||||
async def test_webhook_service_requires_workspace(tenant_webhook_service):
|
||||
with pytest.raises(WorkspaceRequiredError):
|
||||
await tenant_webhook_service.get_webhooks(None)
|
||||
|
||||
|
||||
async def test_same_name_webhooks_are_isolated(tenant_webhook_service):
|
||||
created_a = await tenant_webhook_service.create_webhook(
|
||||
ISOLATION_WORKSPACE_A,
|
||||
'deploy',
|
||||
'https://a.invalid',
|
||||
)
|
||||
created_b = await tenant_webhook_service.create_webhook(
|
||||
ISOLATION_WORKSPACE_B,
|
||||
'deploy',
|
||||
'https://b.invalid',
|
||||
)
|
||||
|
||||
assert created_a['workspace_uuid'] == ISOLATION_WORKSPACE_A
|
||||
assert created_b['workspace_uuid'] == ISOLATION_WORKSPACE_B
|
||||
assert [item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_A)] == ['***']
|
||||
assert [item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_B)] == ['***']
|
||||
assert [
|
||||
item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_A, include_secret=True)
|
||||
] == ['https://a.invalid']
|
||||
assert [
|
||||
item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_B, include_secret=True)
|
||||
] == ['https://b.invalid']
|
||||
|
||||
|
||||
async def test_cross_workspace_id_guessing_is_not_found(tenant_webhook_service):
|
||||
created = await tenant_webhook_service.create_webhook(
|
||||
ISOLATION_WORKSPACE_A,
|
||||
'secret',
|
||||
'https://a.invalid/hook',
|
||||
)
|
||||
webhook_id = created['id']
|
||||
|
||||
assert await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_B, webhook_id) is None
|
||||
assert not await tenant_webhook_service.update_webhook(
|
||||
ISOLATION_WORKSPACE_B,
|
||||
webhook_id,
|
||||
name='stolen',
|
||||
)
|
||||
assert not await tenant_webhook_service.delete_webhook(ISOLATION_WORKSPACE_B, webhook_id)
|
||||
assert (await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_A, webhook_id))['name'] == 'secret'
|
||||
|
||||
|
||||
async def test_update_and_delete_are_scoped(tenant_webhook_service):
|
||||
created = await tenant_webhook_service.create_webhook(
|
||||
ISOLATION_WORKSPACE_A,
|
||||
'old',
|
||||
'https://a.invalid/old',
|
||||
)
|
||||
assert await tenant_webhook_service.update_webhook(
|
||||
ISOLATION_WORKSPACE_A,
|
||||
created['id'],
|
||||
name='new',
|
||||
enabled=False,
|
||||
)
|
||||
assert await tenant_webhook_service.get_enabled_webhooks(ISOLATION_WORKSPACE_A) == []
|
||||
assert await tenant_webhook_service.delete_webhook(ISOLATION_WORKSPACE_A, created['id'])
|
||||
assert await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_A, created['id']) is None
|
||||
|
||||
|
||||
async def test_masked_webhook_url_roundtrip_preserves_replace_and_clear(tenant_webhook_service):
|
||||
created = await tenant_webhook_service.create_webhook(
|
||||
ISOLATION_WORKSPACE_A,
|
||||
'roundtrip',
|
||||
'https://a.invalid/bearer-secret',
|
||||
)
|
||||
|
||||
masked = await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_A, created['id'])
|
||||
assert masked['url'] == '***'
|
||||
assert await tenant_webhook_service.update_webhook(
|
||||
ISOLATION_WORKSPACE_A,
|
||||
created['id'],
|
||||
name='preserved',
|
||||
url=masked['url'],
|
||||
)
|
||||
preserved = await tenant_webhook_service.get_webhook(
|
||||
ISOLATION_WORKSPACE_A,
|
||||
created['id'],
|
||||
include_secret=True,
|
||||
)
|
||||
assert preserved['url'] == 'https://a.invalid/bearer-secret'
|
||||
|
||||
assert await tenant_webhook_service.update_webhook(
|
||||
ISOLATION_WORKSPACE_A,
|
||||
created['id'],
|
||||
url='https://a.invalid/replacement',
|
||||
)
|
||||
replaced = await tenant_webhook_service.get_webhook(
|
||||
ISOLATION_WORKSPACE_A,
|
||||
created['id'],
|
||||
include_secret=True,
|
||||
)
|
||||
assert replaced['url'] == 'https://a.invalid/replacement'
|
||||
|
||||
assert await tenant_webhook_service.update_webhook(ISOLATION_WORKSPACE_A, created['id'], url='')
|
||||
cleared = await tenant_webhook_service.get_webhook(
|
||||
ISOLATION_WORKSPACE_A,
|
||||
created['id'],
|
||||
include_secret=True,
|
||||
)
|
||||
assert cleared['url'] == ''
|
||||
|
||||
@@ -0,0 +1,254 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import lark_oapi
|
||||
import pytest
|
||||
import quart
|
||||
|
||||
from langbot.pkg.api.http.context import (
|
||||
PrincipalContext,
|
||||
PrincipalType,
|
||||
RequestContext,
|
||||
WorkspaceContext,
|
||||
)
|
||||
from langbot.pkg.api.http.controller.groups.platform.adapters import (
|
||||
AdaptersRouterGroup,
|
||||
_AdapterSessionScope,
|
||||
_bind_session_scope,
|
||||
_get_owned_session,
|
||||
_make_room_for_session,
|
||||
_pop_owned_session,
|
||||
_start_adapter_session_task,
|
||||
)
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
SENSITIVE_ADAPTER_ROUTES = (
|
||||
('post', '/api/v1/platform/adapters/lark/create-app'),
|
||||
('get', '/api/v1/platform/adapters/lark/create-app/status/missing'),
|
||||
('delete', '/api/v1/platform/adapters/lark/create-app/missing'),
|
||||
('post', '/api/v1/platform/adapters/weixin/login'),
|
||||
('get', '/api/v1/platform/adapters/weixin/login/status/missing'),
|
||||
('delete', '/api/v1/platform/adapters/weixin/login/missing'),
|
||||
('post', '/api/v1/platform/adapters/dingtalk/create-app'),
|
||||
('get', '/api/v1/platform/adapters/dingtalk/create-app/status/missing'),
|
||||
('delete', '/api/v1/platform/adapters/dingtalk/create-app/missing'),
|
||||
('post', '/api/v1/platform/adapters/wecombot/create-bot'),
|
||||
('get', '/api/v1/platform/adapters/wecombot/create-bot/status/missing'),
|
||||
('delete', '/api/v1/platform/adapters/wecombot/create-bot/missing'),
|
||||
('post', '/api/v1/platform/adapters/qqofficial/bind'),
|
||||
('get', '/api/v1/platform/adapters/qqofficial/bind/status/missing'),
|
||||
('delete', '/api/v1/platform/adapters/qqofficial/bind/missing'),
|
||||
)
|
||||
|
||||
|
||||
def _request_context(
|
||||
*,
|
||||
account_uuid: str = 'account-a',
|
||||
workspace_uuid: str = 'workspace-a',
|
||||
placement_generation: int = 1,
|
||||
) -> RequestContext:
|
||||
return RequestContext(
|
||||
instance_uuid='instance-test',
|
||||
placement_generation=placement_generation,
|
||||
request_id='request-test',
|
||||
auth_type='user-token',
|
||||
principal=PrincipalContext(
|
||||
principal_type=PrincipalType.ACCOUNT,
|
||||
account_uuid=account_uuid,
|
||||
),
|
||||
workspace=WorkspaceContext(
|
||||
workspace_uuid=workspace_uuid,
|
||||
membership_uuid='membership-test',
|
||||
role='developer',
|
||||
permissions=frozenset({'resource.manage'}),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
async def _create_client(*, role: str = 'developer'):
|
||||
quart_app = quart.Quart(__name__)
|
||||
accounts = {
|
||||
'owner-token': SimpleNamespace(uuid='account-a', user='owner@example.com'),
|
||||
'other-token': SimpleNamespace(uuid='account-b', user='other@example.com'),
|
||||
}
|
||||
|
||||
async def get_authenticated_account(token: str):
|
||||
return accounts[token]
|
||||
|
||||
async def resolve_account_workspace(account_uuid: str, requested_workspace_uuid: str | None):
|
||||
workspace_uuid = requested_workspace_uuid or 'workspace-a'
|
||||
return SimpleNamespace(
|
||||
execution=SimpleNamespace(
|
||||
instance_uuid='instance-test',
|
||||
placement_generation=1,
|
||||
),
|
||||
workspace=SimpleNamespace(uuid=workspace_uuid),
|
||||
membership=SimpleNamespace(
|
||||
uuid=f'membership-{account_uuid}-{workspace_uuid}',
|
||||
role=role,
|
||||
projection_revision=1,
|
||||
),
|
||||
)
|
||||
|
||||
class TestTaskManager:
|
||||
def create_user_task(self, coro, **_kwargs):
|
||||
return SimpleNamespace(task=asyncio.create_task(coro))
|
||||
|
||||
application = SimpleNamespace(
|
||||
user_service=SimpleNamespace(
|
||||
get_authenticated_account=AsyncMock(side_effect=get_authenticated_account),
|
||||
),
|
||||
workspace_collaboration_service=SimpleNamespace(
|
||||
resolve_account_workspace=AsyncMock(side_effect=resolve_account_workspace),
|
||||
),
|
||||
platform_mgr=SimpleNamespace(),
|
||||
task_mgr=TestTaskManager(),
|
||||
)
|
||||
router = AdaptersRouterGroup(application, quart_app)
|
||||
await router.initialize()
|
||||
return quart_app.test_client()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(('method', 'path'), SENSITIVE_ADAPTER_ROUTES)
|
||||
async def test_sensitive_adapter_flows_require_resource_manage(method: str, path: str):
|
||||
client = await _create_client(role='viewer')
|
||||
|
||||
response = await getattr(client, method)(
|
||||
path,
|
||||
headers={'Authorization': 'Bearer owner-token'},
|
||||
)
|
||||
|
||||
assert response.status_code == 403
|
||||
assert (await response.get_json())['code'] == 'permission_denied'
|
||||
|
||||
|
||||
async def test_session_scope_matches_exact_tenant_placement_and_principal():
|
||||
owner_context = _request_context()
|
||||
sessions: dict[str, dict] = {'session-test': {'status': 'waiting'}}
|
||||
_bind_session_scope(sessions['session-test'], owner_context)
|
||||
|
||||
assert sessions['session-test']['scope'] == _AdapterSessionScope.from_request_context(owner_context)
|
||||
assert _get_owned_session(sessions, 'session-test', owner_context) is sessions['session-test']
|
||||
|
||||
for other_context in (
|
||||
_request_context(account_uuid='account-b'),
|
||||
_request_context(workspace_uuid='workspace-b'),
|
||||
_request_context(placement_generation=2),
|
||||
):
|
||||
assert _get_owned_session(sessions, 'session-test', other_context) is None
|
||||
assert _pop_owned_session(sessions, 'session-test', other_context) is None
|
||||
assert 'session-test' in sessions
|
||||
|
||||
assert _pop_owned_session(sessions, 'session-test', owner_context) is not None
|
||||
assert sessions == {}
|
||||
|
||||
|
||||
async def test_session_capacity_evicts_oldest_session_in_same_workspace():
|
||||
owner_context = _request_context()
|
||||
sessions: dict[str, dict] = {}
|
||||
tasks = []
|
||||
for index in range(10):
|
||||
task = SimpleNamespace(done=Mock(return_value=False), cancel=Mock())
|
||||
tasks.append(task)
|
||||
session = {'created_at': float(index), 'task': task}
|
||||
_bind_session_scope(session, owner_context)
|
||||
sessions[f'session-{index}'] = session
|
||||
|
||||
_make_room_for_session(sessions, owner_context)
|
||||
|
||||
assert 'session-0' not in sessions
|
||||
assert len(sessions) == 9
|
||||
tasks[0].cancel.assert_called_once_with()
|
||||
|
||||
|
||||
async def test_adapter_session_task_uses_tenant_task_admission():
|
||||
blocker = asyncio.Event()
|
||||
|
||||
async def credential_exchange():
|
||||
await blocker.wait()
|
||||
|
||||
task_manager = SimpleNamespace(create_user_task=Mock())
|
||||
|
||||
def create_user_task(coro, **_kwargs):
|
||||
return SimpleNamespace(task=asyncio.create_task(coro))
|
||||
|
||||
task_manager.create_user_task.side_effect = create_user_task
|
||||
application = SimpleNamespace(task_mgr=task_manager)
|
||||
request_context = _request_context()
|
||||
|
||||
returned = _start_adapter_session_task(
|
||||
application,
|
||||
credential_exchange(),
|
||||
adapter='lark',
|
||||
session_id='session-test',
|
||||
request_context=request_context,
|
||||
)
|
||||
|
||||
assert returned is not None
|
||||
task_manager.create_user_task.assert_called_once()
|
||||
kwargs = task_manager.create_user_task.call_args.kwargs
|
||||
assert kwargs['kind'] == 'platform-adapter-credential-exchange'
|
||||
assert kwargs['instance_uuid'] == request_context.instance_uuid
|
||||
assert kwargs['workspace_uuid'] == request_context.workspace_uuid
|
||||
assert kwargs['placement_generation'] == request_context.placement_generation
|
||||
blocker.set()
|
||||
await returned
|
||||
|
||||
|
||||
async def test_lark_session_status_and_delete_hide_cross_scope_sessions(monkeypatch):
|
||||
registration_blocker = asyncio.Event()
|
||||
|
||||
async def fake_register_app(*, on_qr_code, source: str):
|
||||
assert source == 'langbot'
|
||||
on_qr_code({'url': 'https://example.test/lark-qr'})
|
||||
await registration_blocker.wait()
|
||||
raise AssertionError('registration should have been cancelled')
|
||||
|
||||
monkeypatch.setattr(lark_oapi, 'aregister_app', fake_register_app)
|
||||
client = await _create_client()
|
||||
owner_headers = {
|
||||
'Authorization': 'Bearer owner-token',
|
||||
'X-Workspace-Id': 'workspace-a',
|
||||
}
|
||||
|
||||
create_response = await client.post(
|
||||
'/api/v1/platform/adapters/lark/create-app',
|
||||
headers=owner_headers,
|
||||
)
|
||||
assert create_response.status_code == 200
|
||||
session_id = (await create_response.get_json())['data']['session_id']
|
||||
status_path = f'/api/v1/platform/adapters/lark/create-app/status/{session_id}'
|
||||
delete_path = f'/api/v1/platform/adapters/lark/create-app/{session_id}'
|
||||
|
||||
for headers in (
|
||||
{
|
||||
'Authorization': 'Bearer other-token',
|
||||
'X-Workspace-Id': 'workspace-a',
|
||||
},
|
||||
{
|
||||
'Authorization': 'Bearer owner-token',
|
||||
'X-Workspace-Id': 'workspace-b',
|
||||
},
|
||||
):
|
||||
status_response = await client.get(status_path, headers=headers)
|
||||
delete_response = await client.delete(delete_path, headers=headers)
|
||||
assert status_response.status_code == 404
|
||||
assert delete_response.status_code == 404
|
||||
assert (await status_response.get_json())['msg'] == 'Session not found'
|
||||
assert (await delete_response.get_json())['msg'] == 'Session not found'
|
||||
|
||||
owner_status_response = await client.get(status_path, headers=owner_headers)
|
||||
assert owner_status_response.status_code == 200
|
||||
assert (await owner_status_response.get_json())['data']['status'] == 'waiting'
|
||||
|
||||
owner_delete_response = await client.delete(delete_path, headers=owner_headers)
|
||||
assert owner_delete_response.status_code == 200
|
||||
missing_delete_response = await client.delete(delete_path, headers=owner_headers)
|
||||
assert missing_delete_response.status_code == 404
|
||||
await asyncio.sleep(0)
|
||||
@@ -6,6 +6,7 @@ from unittest.mock import AsyncMock, Mock
|
||||
import pytest
|
||||
|
||||
from langbot.pkg.api.http.service.apikey import ApiKeyService
|
||||
from langbot.pkg.entity.persistence.apikey import ApiKeyStatus
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@@ -13,30 +14,70 @@ from langbot.pkg.api.http.service.apikey import ApiKeyService
|
||||
async def test_verify_api_key_rejects_non_lbk_keys_without_db_query(api_key):
|
||||
persistence_mgr = SimpleNamespace(execute_async=AsyncMock())
|
||||
instance_config = SimpleNamespace(data={'api': {'global_api_key': ''}})
|
||||
service = ApiKeyService(SimpleNamespace(persistence_mgr=persistence_mgr, instance_config=instance_config))
|
||||
workspace_service = SimpleNamespace(get_execution_binding=AsyncMock())
|
||||
service = ApiKeyService(
|
||||
SimpleNamespace(
|
||||
persistence_mgr=persistence_mgr,
|
||||
instance_config=instance_config,
|
||||
workspace_service=workspace_service,
|
||||
)
|
||||
)
|
||||
|
||||
result = await service.verify_api_key(api_key)
|
||||
|
||||
assert result is False
|
||||
persistence_mgr.execute_async.assert_not_awaited()
|
||||
workspace_service.get_execution_binding.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
('db_row', 'expected'),
|
||||
[
|
||||
(object(), True),
|
||||
(None, False),
|
||||
],
|
||||
)
|
||||
async def test_verify_api_key_keeps_db_validation_for_lbk_keys(db_row, expected):
|
||||
query_result = Mock()
|
||||
query_result.first.return_value = db_row
|
||||
persistence_mgr = SimpleNamespace(execute_async=AsyncMock(return_value=query_result))
|
||||
@pytest.mark.parametrize('key_exists', [True, False])
|
||||
async def test_verify_api_key_keeps_db_validation_for_lbk_keys(key_exists):
|
||||
key = (
|
||||
SimpleNamespace(
|
||||
id=1,
|
||||
uuid='key-uuid',
|
||||
workspace_uuid='workspace-a',
|
||||
status=ApiKeyStatus.ACTIVE.value,
|
||||
expires_at=None,
|
||||
scopes=[],
|
||||
)
|
||||
if key_exists
|
||||
else None
|
||||
)
|
||||
discovery_result = Mock()
|
||||
discovery_result.first.return_value = key
|
||||
query_results = [discovery_result]
|
||||
if key_exists:
|
||||
scoped_result = Mock()
|
||||
scoped_result.first.return_value = key
|
||||
update_result = Mock()
|
||||
update_result.scalar_one_or_none.return_value = key.id
|
||||
query_results.extend([scoped_result, update_result])
|
||||
persistence_mgr = SimpleNamespace(execute_async=AsyncMock(side_effect=query_results))
|
||||
instance_config = SimpleNamespace(data={'api': {'global_api_key': ''}})
|
||||
service = ApiKeyService(SimpleNamespace(persistence_mgr=persistence_mgr, instance_config=instance_config))
|
||||
workspace_service = SimpleNamespace(
|
||||
get_execution_binding=AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
instance_uuid='instance-a',
|
||||
workspace_uuid='workspace-a',
|
||||
placement_generation=1,
|
||||
)
|
||||
)
|
||||
)
|
||||
service = ApiKeyService(
|
||||
SimpleNamespace(
|
||||
persistence_mgr=persistence_mgr,
|
||||
instance_config=instance_config,
|
||||
workspace_service=workspace_service,
|
||||
)
|
||||
)
|
||||
|
||||
result = await service.verify_api_key('lbk_valid_format')
|
||||
|
||||
assert result is expected
|
||||
persistence_mgr.execute_async.assert_awaited_once()
|
||||
assert result is key_exists
|
||||
assert persistence_mgr.execute_async.await_count == (3 if key_exists else 1)
|
||||
if key_exists:
|
||||
workspace_service.get_execution_binding.assert_awaited_once_with('workspace-a')
|
||||
else:
|
||||
workspace_service.get_execution_binding.assert_not_awaited()
|
||||
|
||||
@@ -0,0 +1,113 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import quart
|
||||
|
||||
from langbot.pkg.api.http.controller.groups.platform.bots import BotsRouterGroup
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
SECRET_CONFIG = {'token': 'tenant-secret', 'app_secret': 'also-secret'}
|
||||
|
||||
|
||||
async def create_client(*, role: str):
|
||||
quart_app = quart.Quart(__name__)
|
||||
account = SimpleNamespace(uuid='account-test', user='test@example.com')
|
||||
user_service = SimpleNamespace(
|
||||
get_authenticated_account=AsyncMock(return_value=account),
|
||||
)
|
||||
access = SimpleNamespace(
|
||||
execution=SimpleNamespace(
|
||||
instance_uuid='instance-test',
|
||||
placement_generation=1,
|
||||
),
|
||||
workspace=SimpleNamespace(uuid='workspace-test'),
|
||||
membership=SimpleNamespace(
|
||||
uuid='membership-test',
|
||||
role=role,
|
||||
projection_revision=1,
|
||||
),
|
||||
)
|
||||
|
||||
async def get_bots(_context, *, include_secret=False):
|
||||
bot = {'uuid': 'bot-test', 'name': 'Test Bot'}
|
||||
if include_secret:
|
||||
bot['adapter_config'] = SECRET_CONFIG
|
||||
return [bot]
|
||||
|
||||
async def get_runtime_bot_info(_context, _bot_uuid, *, include_secret=False):
|
||||
bot = {'uuid': 'bot-test', 'name': 'Test Bot'}
|
||||
if include_secret:
|
||||
bot['adapter_config'] = SECRET_CONFIG
|
||||
return bot
|
||||
|
||||
bot_service = SimpleNamespace(
|
||||
get_bots=AsyncMock(side_effect=get_bots),
|
||||
get_runtime_bot_info=AsyncMock(side_effect=get_runtime_bot_info),
|
||||
update_bot=AsyncMock(),
|
||||
)
|
||||
application = SimpleNamespace(
|
||||
user_service=user_service,
|
||||
apikey_service=SimpleNamespace(
|
||||
authenticate_api_key=AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
instance_uuid='instance-test',
|
||||
placement_generation=1,
|
||||
api_key_uuid='api-key-test',
|
||||
workspace_uuid='workspace-test',
|
||||
permissions=frozenset({'resource.view'}),
|
||||
)
|
||||
)
|
||||
),
|
||||
workspace_collaboration_service=SimpleNamespace(resolve_account_workspace=AsyncMock(return_value=access)),
|
||||
bot_service=bot_service,
|
||||
)
|
||||
router = BotsRouterGroup(application, quart_app)
|
||||
await router.initialize()
|
||||
return quart_app.test_client(), bot_service
|
||||
|
||||
|
||||
async def test_viewer_list_and_detail_never_receive_adapter_credentials():
|
||||
client, bot_service = await create_client(role='viewer')
|
||||
headers = {'Authorization': 'Bearer test-token'}
|
||||
|
||||
list_response = await client.get('/api/v1/platform/bots', headers=headers)
|
||||
detail_response = await client.get('/api/v1/platform/bots/bot-test', headers=headers)
|
||||
|
||||
assert list_response.status_code == 200
|
||||
assert detail_response.status_code == 200
|
||||
assert 'adapter_config' not in (await list_response.get_json())['data']['bots'][0]
|
||||
assert 'adapter_config' not in (await detail_response.get_json())['data']['bot']
|
||||
assert bot_service.get_bots.await_args.kwargs['include_secret'] is False
|
||||
assert bot_service.get_runtime_bot_info.await_args.kwargs['include_secret'] is False
|
||||
|
||||
|
||||
async def test_resource_manager_can_read_adapter_credentials():
|
||||
client, bot_service = await create_client(role='developer')
|
||||
headers = {'Authorization': 'Bearer test-token'}
|
||||
|
||||
list_response = await client.get('/api/v1/platform/bots', headers=headers)
|
||||
detail_response = await client.get('/api/v1/platform/bots/bot-test', headers=headers)
|
||||
|
||||
assert (await list_response.get_json())['data']['bots'][0]['adapter_config'] == SECRET_CONFIG
|
||||
assert (await detail_response.get_json())['data']['bot']['adapter_config'] == SECRET_CONFIG
|
||||
assert bot_service.get_bots.await_args.kwargs['include_secret'] is True
|
||||
assert bot_service.get_runtime_bot_info.await_args.kwargs['include_secret'] is True
|
||||
|
||||
|
||||
async def test_viewer_cannot_write_adapter_credentials():
|
||||
client, bot_service = await create_client(role='viewer')
|
||||
|
||||
response = await client.put(
|
||||
'/api/v1/platform/bots/bot-test',
|
||||
headers={'Authorization': 'Bearer test-token'},
|
||||
json={'adapter_config': SECRET_CONFIG},
|
||||
)
|
||||
|
||||
assert response.status_code == 403
|
||||
assert (await response.get_json())['code'] == 'permission_denied'
|
||||
bot_service.update_bot.assert_not_awaited()
|
||||
@@ -0,0 +1,142 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import quart
|
||||
import sqlalchemy
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
from langbot.pkg.api.http.controller.groups.extensions import ExtensionsRouterGroup
|
||||
from langbot.pkg.persistence.mgr import PersistenceManager
|
||||
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extensions_route_hides_runtime_bound_to_another_workspace():
|
||||
account = SimpleNamespace(uuid='account-a', user='owner@example.com')
|
||||
connector = SimpleNamespace(
|
||||
is_enable_plugin=True,
|
||||
require_workspace_context=AsyncMock(side_effect=WorkspaceNotFoundError('Plugin resource not found')),
|
||||
list_plugins=AsyncMock(return_value=[]),
|
||||
)
|
||||
ap = SimpleNamespace(
|
||||
user_service=SimpleNamespace(
|
||||
get_authenticated_account=AsyncMock(return_value=account),
|
||||
),
|
||||
workspace_collaboration_service=SimpleNamespace(
|
||||
resolve_account_workspace=AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
workspace=SimpleNamespace(uuid='workspace-a'),
|
||||
membership=SimpleNamespace(uuid='membership-a', role='owner', projection_revision=0),
|
||||
execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=2),
|
||||
)
|
||||
)
|
||||
),
|
||||
plugin_connector=connector,
|
||||
mcp_service=SimpleNamespace(get_mcp_servers=AsyncMock(return_value=[])),
|
||||
skill_service=SimpleNamespace(list_skills=AsyncMock(return_value=[])),
|
||||
)
|
||||
quart_app = quart.Quart(__name__)
|
||||
router = ExtensionsRouterGroup(ap, quart_app)
|
||||
await router.initialize()
|
||||
|
||||
response = await quart_app.test_client().get(
|
||||
'/api/v1/extensions',
|
||||
headers={'Authorization': 'Bearer token'},
|
||||
)
|
||||
|
||||
assert response.status_code == 404
|
||||
connector.list_plugins.assert_not_awaited()
|
||||
ap.mcp_service.get_mcp_servers.assert_not_awaited()
|
||||
ap.skill_service.list_skills.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extensions_route_redacts_plugin_secrets_without_mutating_runtime_data():
|
||||
account = SimpleNamespace(uuid='account-a', user='viewer@example.com')
|
||||
raw_plugin = {
|
||||
'plugin_config': {'apiKey': 'plugin-secret', 'nested': {'token': 'nested-secret'}},
|
||||
'debug': {'plugin_debug_key': 'debug-secret'},
|
||||
}
|
||||
connector = SimpleNamespace(
|
||||
is_enable_plugin=True,
|
||||
require_workspace_context=AsyncMock(),
|
||||
list_plugins=AsyncMock(return_value=[raw_plugin]),
|
||||
)
|
||||
ap = SimpleNamespace(
|
||||
user_service=SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account)),
|
||||
workspace_collaboration_service=SimpleNamespace(
|
||||
resolve_account_workspace=AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
workspace=SimpleNamespace(uuid='workspace-a'),
|
||||
membership=SimpleNamespace(uuid='membership-a', role='viewer', projection_revision=0),
|
||||
execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=2),
|
||||
)
|
||||
)
|
||||
),
|
||||
plugin_connector=connector,
|
||||
mcp_service=SimpleNamespace(get_mcp_servers=AsyncMock(return_value=[])),
|
||||
skill_service=SimpleNamespace(list_skills=AsyncMock(return_value=[])),
|
||||
)
|
||||
quart_app = quart.Quart(__name__)
|
||||
router = ExtensionsRouterGroup(ap, quart_app)
|
||||
await router.initialize()
|
||||
|
||||
response = await quart_app.test_client().get(
|
||||
'/api/v1/extensions',
|
||||
headers={'Authorization': 'Bearer token', 'X-Workspace-Id': 'workspace-a'},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
plugin = (await response.get_json())['data']['extensions'][0]['plugin']
|
||||
assert plugin['plugin_config']['apiKey'] == '***'
|
||||
assert plugin['plugin_config']['nested']['token'] == '***'
|
||||
assert plugin['debug']['plugin_debug_key'] == '***'
|
||||
assert raw_plugin['plugin_config']['apiKey'] == 'plugin-secret'
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extensions_parallel_reads_open_explicit_child_task_scopes():
|
||||
engine = create_async_engine('sqlite+aiosqlite:///:memory:')
|
||||
account = SimpleNamespace(uuid='account-a', user='owner@example.com')
|
||||
ap = SimpleNamespace()
|
||||
ap.persistence_mgr = PersistenceManager(ap)
|
||||
ap.persistence_mgr.db = SimpleNamespace(get_engine=lambda: engine)
|
||||
ap.user_service = SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account))
|
||||
ap.workspace_collaboration_service = SimpleNamespace(
|
||||
resolve_account_workspace=AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
workspace=SimpleNamespace(uuid='workspace-a'),
|
||||
membership=SimpleNamespace(uuid='membership-a', role='owner', projection_revision=0),
|
||||
execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=2),
|
||||
)
|
||||
)
|
||||
)
|
||||
ap.plugin_connector = SimpleNamespace(
|
||||
is_enable_plugin=False,
|
||||
list_plugins=AsyncMock(return_value=[]),
|
||||
)
|
||||
|
||||
async def list_mcp_servers(_context, *, contain_runtime_info):
|
||||
await ap.persistence_mgr.execute_async(sqlalchemy.select(sqlalchemy.literal(1)))
|
||||
assert contain_runtime_info is True
|
||||
return [{'name': 'Scoped MCP'}]
|
||||
|
||||
ap.mcp_service = SimpleNamespace(get_mcp_servers=AsyncMock(side_effect=list_mcp_servers))
|
||||
ap.skill_service = SimpleNamespace(list_skills=AsyncMock(return_value=[]))
|
||||
quart_app = quart.Quart(__name__)
|
||||
router = ExtensionsRouterGroup(ap, quart_app)
|
||||
await router.initialize()
|
||||
|
||||
try:
|
||||
response = await quart_app.test_client().get(
|
||||
'/api/v1/extensions',
|
||||
headers={'Authorization': 'Bearer token', 'X-Workspace-Id': 'workspace-a'},
|
||||
)
|
||||
finally:
|
||||
await engine.dispose()
|
||||
|
||||
assert response.status_code == 200
|
||||
assert (await response.get_json())['data']['extensions'] == [{'type': 'mcp', 'server': {'name': 'Scoped MCP'}}]
|
||||
@@ -0,0 +1,59 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import io
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import quart
|
||||
from quart.datastructures import FileStorage
|
||||
|
||||
from langbot.pkg.api.http.controller.groups.files import FilesRouterGroup
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
async def test_document_upload_uses_dedicated_scoped_owner_type():
|
||||
quart_app = quart.Quart(__name__)
|
||||
account = SimpleNamespace(uuid='account-test', user='test@example.com')
|
||||
access = SimpleNamespace(
|
||||
execution=SimpleNamespace(
|
||||
instance_uuid='instance-test',
|
||||
placement_generation=3,
|
||||
),
|
||||
workspace=SimpleNamespace(uuid='00000000-0000-0000-0000-00000000000a'),
|
||||
membership=SimpleNamespace(
|
||||
uuid='membership-test',
|
||||
role='developer',
|
||||
projection_revision=1,
|
||||
),
|
||||
)
|
||||
storage_mgr = SimpleNamespace(save_scoped=AsyncMock(return_value='scoped-document-key'))
|
||||
application = SimpleNamespace(
|
||||
user_service=SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account)),
|
||||
workspace_collaboration_service=SimpleNamespace(resolve_account_workspace=AsyncMock(return_value=access)),
|
||||
storage_mgr=storage_mgr,
|
||||
)
|
||||
router = FilesRouterGroup(application, quart_app)
|
||||
await router.initialize()
|
||||
client = quart_app.test_client()
|
||||
|
||||
response = await client.post(
|
||||
'/api/v1/files/documents',
|
||||
headers={'Authorization': 'Bearer test-token'},
|
||||
files={
|
||||
'file': FileStorage(
|
||||
stream=io.BytesIO(b'document bytes'),
|
||||
filename='report.pdf',
|
||||
)
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
assert (await response.get_json())['data']['file_id'] == 'scoped-document-key'
|
||||
kwargs = storage_mgr.save_scoped.await_args.kwargs
|
||||
assert kwargs['owner_type'] == 'upload_document'
|
||||
assert kwargs['owner'] == 'account:account-test'
|
||||
assert kwargs['key'].endswith('.pdf')
|
||||
assert kwargs['value'] == b'document bytes'
|
||||
@@ -0,0 +1,293 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
import json
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
import sqlalchemy
|
||||
|
||||
from langbot.pkg.api.http.context import ExecutionContext
|
||||
from langbot.pkg.api.http.controller.groups.knowledge.migration import KnowledgeMigrationRouterGroup
|
||||
from langbot.pkg.persistence.tenant_uow import _validate_scoped_statement_call
|
||||
from langbot.pkg.workspace.errors import WorkspaceInvariantError, WorkspaceNotFoundError
|
||||
|
||||
|
||||
CONTEXT = ExecutionContext(
|
||||
instance_uuid='instance-a',
|
||||
workspace_uuid='workspace-a',
|
||||
placement_generation=3,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_background_migration_propagates_generation_change_before_runtime_call():
|
||||
connector = SimpleNamespace(
|
||||
require_workspace_context=AsyncMock(side_effect=[CONTEXT, WorkspaceNotFoundError('Plugin resource not found')]),
|
||||
list_knowledge_engines=AsyncMock(return_value=[]),
|
||||
)
|
||||
router = object.__new__(KnowledgeMigrationRouterGroup)
|
||||
router.ap = SimpleNamespace(
|
||||
plugin_connector=connector,
|
||||
workspace_service=SimpleNamespace(
|
||||
get_local_execution_binding=AsyncMock(return_value=CONTEXT),
|
||||
),
|
||||
logger=Mock(),
|
||||
)
|
||||
router._table_exists = AsyncMock(return_value=False)
|
||||
router._set_migration_flag = AsyncMock()
|
||||
task_context = SimpleNamespace(trace=Mock())
|
||||
|
||||
with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
|
||||
await router._execute_rag_migration(
|
||||
CONTEXT,
|
||||
task_context,
|
||||
install_plugin=False,
|
||||
)
|
||||
|
||||
assert connector.require_workspace_context.await_count == 2
|
||||
connector.list_knowledge_engines.assert_not_awaited()
|
||||
router._set_migration_flag.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cloud_migration_is_rejected_before_legacy_table_access():
|
||||
router = object.__new__(KnowledgeMigrationRouterGroup)
|
||||
router.ap = SimpleNamespace(
|
||||
workspace_service=SimpleNamespace(
|
||||
get_local_execution_binding=AsyncMock(side_effect=WorkspaceInvariantError('not an OSS local workspace')),
|
||||
),
|
||||
plugin_connector=SimpleNamespace(require_workspace_context=AsyncMock()),
|
||||
logger=Mock(),
|
||||
)
|
||||
router._table_exists = AsyncMock()
|
||||
router._set_migration_flag = AsyncMock()
|
||||
task_context = SimpleNamespace(trace=Mock())
|
||||
|
||||
with pytest.raises(WorkspaceNotFoundError, match='migration is unavailable'):
|
||||
await router._execute_rag_migration(CONTEXT, task_context, install_plugin=False)
|
||||
|
||||
router._table_exists.assert_not_awaited()
|
||||
router.ap.plugin_connector.require_workspace_context.assert_not_awaited()
|
||||
router._set_migration_flag.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize('database_name, exists', [('postgresql', True), ('sqlite', False)])
|
||||
async def test_legacy_table_discovery_uses_scoped_structured_queries(database_name: str, exists: bool):
|
||||
result = Mock()
|
||||
result.first.return_value = ('knowledge_bases_backup',) if exists else None
|
||||
execute_async = AsyncMock(return_value=result)
|
||||
router = object.__new__(KnowledgeMigrationRouterGroup)
|
||||
router.ap = SimpleNamespace(
|
||||
persistence_mgr=SimpleNamespace(
|
||||
db=SimpleNamespace(name=database_name),
|
||||
execute_async=execute_async,
|
||||
)
|
||||
)
|
||||
|
||||
assert await router._table_exists('knowledge_bases_backup') is exists
|
||||
|
||||
statement = execute_async.await_args.args[0]
|
||||
assert isinstance(statement, sqlalchemy.sql.selectable.SelectBase)
|
||||
_validate_scoped_statement_call((statement,), {})
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_legacy_restore_emits_only_scoped_structured_statements():
|
||||
missing_table_result = Mock()
|
||||
missing_table_result.first.return_value = None
|
||||
existing_table_result = Mock()
|
||||
existing_table_result.first.return_value = ('knowledge_bases_backup',)
|
||||
backup_result = Mock()
|
||||
backup_result.keys.return_value = [
|
||||
'uuid',
|
||||
'name',
|
||||
'description',
|
||||
'emoji',
|
||||
'embedding_model_uuid',
|
||||
'top_k',
|
||||
'created_at',
|
||||
'updated_at',
|
||||
]
|
||||
now = datetime.now()
|
||||
backup_result.fetchall.return_value = [
|
||||
('kb-legacy', 'Legacy KB', 'Description', 'U0001f4da', 'embedding-model', 7, now, now)
|
||||
]
|
||||
execute_async = AsyncMock(
|
||||
side_effect=[
|
||||
missing_table_result,
|
||||
existing_table_result,
|
||||
backup_result,
|
||||
Mock(),
|
||||
Mock(),
|
||||
]
|
||||
)
|
||||
connector = SimpleNamespace(
|
||||
require_workspace_context=AsyncMock(return_value=CONTEXT),
|
||||
list_knowledge_engines=AsyncMock(return_value=[]),
|
||||
rag_on_kb_create=AsyncMock(),
|
||||
)
|
||||
router = object.__new__(KnowledgeMigrationRouterGroup)
|
||||
router.ap = SimpleNamespace(
|
||||
workspace_service=SimpleNamespace(
|
||||
get_local_execution_binding=AsyncMock(return_value=CONTEXT),
|
||||
),
|
||||
plugin_connector=connector,
|
||||
persistence_mgr=SimpleNamespace(
|
||||
db=SimpleNamespace(name='sqlite'),
|
||||
execute_async=execute_async,
|
||||
),
|
||||
rag_mgr=SimpleNamespace(load_knowledge_bases_from_db=AsyncMock()),
|
||||
logger=Mock(),
|
||||
)
|
||||
task_context = SimpleNamespace(trace=Mock())
|
||||
|
||||
await router._execute_rag_migration(CONTEXT, task_context, install_plugin=False)
|
||||
|
||||
statements = [call.args[0] for call in execute_async.await_args_list]
|
||||
assert any(isinstance(statement, sqlalchemy.sql.dml.Insert) for statement in statements)
|
||||
assert any(isinstance(statement, sqlalchemy.sql.dml.Update) for statement in statements)
|
||||
for statement in statements:
|
||||
assert not isinstance(statement, sqlalchemy.sql.elements.TextClause)
|
||||
_validate_scoped_statement_call((statement,), {})
|
||||
connector.rag_on_kb_create.assert_awaited_once_with(
|
||||
'langbot-team/LangRAG',
|
||||
'kb-legacy',
|
||||
{'embedding_model_uuid': 'embedding-model'},
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_legacy_restore_accepts_sqlite_string_dates_and_text_json_columns():
|
||||
engine = sqlalchemy.ext.asyncio.create_async_engine('sqlite+aiosqlite:///:memory:')
|
||||
try:
|
||||
async with engine.begin() as connection:
|
||||
await connection.exec_driver_sql(
|
||||
"""
|
||||
CREATE TABLE knowledge_bases_backup (
|
||||
uuid TEXT PRIMARY KEY,
|
||||
name TEXT,
|
||||
description TEXT,
|
||||
emoji TEXT,
|
||||
embedding_model_uuid TEXT,
|
||||
top_k INTEGER,
|
||||
created_at DATETIME,
|
||||
updated_at DATETIME
|
||||
)
|
||||
"""
|
||||
)
|
||||
await connection.exec_driver_sql(
|
||||
"""
|
||||
CREATE TABLE external_knowledge_bases (
|
||||
uuid TEXT PRIMARY KEY,
|
||||
name TEXT,
|
||||
description TEXT,
|
||||
emoji TEXT,
|
||||
plugin_author TEXT,
|
||||
plugin_name TEXT,
|
||||
retriever_config TEXT,
|
||||
created_at DATETIME
|
||||
)
|
||||
"""
|
||||
)
|
||||
await connection.exec_driver_sql(
|
||||
"""
|
||||
CREATE TABLE knowledge_bases (
|
||||
uuid TEXT PRIMARY KEY,
|
||||
workspace_uuid TEXT NOT NULL,
|
||||
name TEXT,
|
||||
description TEXT,
|
||||
emoji TEXT,
|
||||
created_at DATETIME,
|
||||
updated_at DATETIME,
|
||||
knowledge_engine_plugin_id TEXT,
|
||||
collection_id TEXT,
|
||||
creation_settings TEXT,
|
||||
retrieval_settings TEXT
|
||||
)
|
||||
"""
|
||||
)
|
||||
await connection.exec_driver_sql(
|
||||
'CREATE TABLE workspace_metadata (workspace_uuid TEXT, key TEXT, value TEXT)'
|
||||
)
|
||||
legacy_timestamp = '2026-07-20 03:00:00.123456'
|
||||
await connection.exec_driver_sql(
|
||||
"""
|
||||
INSERT INTO knowledge_bases_backup
|
||||
(uuid, name, description, emoji, embedding_model_uuid, top_k, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
'kb-internal',
|
||||
'Internal',
|
||||
'Internal legacy KB',
|
||||
'U0001f4da',
|
||||
'embedding-model',
|
||||
5,
|
||||
legacy_timestamp,
|
||||
legacy_timestamp,
|
||||
),
|
||||
)
|
||||
await connection.exec_driver_sql(
|
||||
"""
|
||||
INSERT INTO external_knowledge_bases
|
||||
(uuid, name, description, emoji, plugin_author, plugin_name, retriever_config, created_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
""",
|
||||
(
|
||||
'kb-external',
|
||||
'External',
|
||||
'External legacy KB',
|
||||
'U0001f517',
|
||||
'langbot-team',
|
||||
'DifyDatasetsRetriever',
|
||||
json.dumps({'api_base_url': 'https://example.invalid', 'top_k': 8}),
|
||||
legacy_timestamp,
|
||||
),
|
||||
)
|
||||
|
||||
connector = SimpleNamespace(
|
||||
require_workspace_context=AsyncMock(return_value=CONTEXT),
|
||||
list_knowledge_engines=AsyncMock(return_value=[]),
|
||||
rag_on_kb_create=AsyncMock(),
|
||||
)
|
||||
router = object.__new__(KnowledgeMigrationRouterGroup)
|
||||
router.ap = SimpleNamespace(
|
||||
workspace_service=SimpleNamespace(
|
||||
get_local_execution_binding=AsyncMock(return_value=CONTEXT),
|
||||
),
|
||||
plugin_connector=connector,
|
||||
persistence_mgr=SimpleNamespace(
|
||||
db=SimpleNamespace(name='sqlite'),
|
||||
execute_async=connection.execute,
|
||||
),
|
||||
rag_mgr=SimpleNamespace(load_knowledge_bases_from_db=AsyncMock()),
|
||||
logger=Mock(),
|
||||
)
|
||||
|
||||
await router._execute_rag_migration(
|
||||
CONTEXT,
|
||||
SimpleNamespace(trace=Mock()),
|
||||
install_plugin=False,
|
||||
)
|
||||
|
||||
restored = (
|
||||
await connection.exec_driver_sql(
|
||||
"""
|
||||
SELECT uuid, created_at, updated_at, creation_settings, retrieval_settings
|
||||
FROM knowledge_bases
|
||||
ORDER BY uuid
|
||||
"""
|
||||
)
|
||||
).all()
|
||||
assert [row.uuid for row in restored] == ['kb-external', 'kb-internal']
|
||||
assert all(row.created_at == legacy_timestamp for row in restored)
|
||||
assert all(row.updated_at == legacy_timestamp for row in restored)
|
||||
assert json.loads(restored[0].creation_settings)['api_base_url'] == 'https://example.invalid'
|
||||
assert json.loads(restored[0].retrieval_settings) == {'top_k': 8}
|
||||
assert json.loads(restored[1].creation_settings) == {'embedding_model_uuid': 'embedding-model'}
|
||||
assert json.loads(restored[1].retrieval_settings) == {'top_k': 5}
|
||||
finally:
|
||||
await engine.dispose()
|
||||
@@ -17,19 +17,63 @@ sys.modules.setdefault('langbot.pkg.core.app', core_app_module)
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
async def _create_test_client(mcp_service: SimpleNamespace):
|
||||
async def _create_test_client(mcp_service: SimpleNamespace, *, role: str = 'owner'):
|
||||
app = quart.Quart(__name__)
|
||||
user_service = SimpleNamespace(
|
||||
verify_jwt_token=AsyncMock(return_value='test@example.com'),
|
||||
get_user_by_email=AsyncMock(return_value=SimpleNamespace(user='test@example.com')),
|
||||
get_user_by_email=AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
user='test@example.com',
|
||||
uuid='account-a',
|
||||
)
|
||||
),
|
||||
)
|
||||
workspace_collaboration_service = SimpleNamespace(
|
||||
resolve_account_workspace=AsyncMock(
|
||||
return_value=SimpleNamespace(
|
||||
execution=SimpleNamespace(
|
||||
instance_uuid='instance-a',
|
||||
placement_generation=1,
|
||||
),
|
||||
workspace=SimpleNamespace(uuid='workspace-a'),
|
||||
membership=SimpleNamespace(
|
||||
uuid='membership-a',
|
||||
role=role,
|
||||
projection_revision=1,
|
||||
),
|
||||
)
|
||||
)
|
||||
)
|
||||
ap = SimpleNamespace(
|
||||
mcp_service=mcp_service,
|
||||
user_service=user_service,
|
||||
workspace_collaboration_service=workspace_collaboration_service,
|
||||
)
|
||||
ap = SimpleNamespace(mcp_service=mcp_service, user_service=user_service)
|
||||
MCPRouterGroup = import_module('langbot.pkg.api.http.controller.groups.resources.mcp').MCPRouterGroup
|
||||
group = MCPRouterGroup(ap, app)
|
||||
await group.initialize()
|
||||
return app.test_client()
|
||||
|
||||
|
||||
async def test_viewer_cannot_read_mcp_runtime_logs():
|
||||
mcp_service = SimpleNamespace(
|
||||
get_mcp_server_logs=AsyncMock(return_value=['private runtime line']),
|
||||
)
|
||||
client = await _create_test_client(mcp_service, role='viewer')
|
||||
|
||||
response = await client.get(
|
||||
'/api/v1/mcp/servers/example/logs',
|
||||
headers={
|
||||
'Authorization': 'Bearer test-token',
|
||||
'X-Workspace-Id': 'workspace-a',
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 403
|
||||
assert (await response.get_json())['code'] == 'permission_denied'
|
||||
mcp_service.get_mcp_server_logs.assert_not_awaited()
|
||||
|
||||
|
||||
async def test_mcp_server_route_accepts_encoded_slash_name():
|
||||
mcp_service = SimpleNamespace(
|
||||
get_mcp_server_by_name=AsyncMock(
|
||||
@@ -46,11 +90,17 @@ async def test_mcp_server_route_accepts_encoded_slash_name():
|
||||
|
||||
response = await client.get(
|
||||
'/api/v1/mcp/servers/pab1it0%2Fprometheus',
|
||||
headers={'Authorization': 'Bearer test-token'},
|
||||
headers={
|
||||
'Authorization': 'Bearer test-token',
|
||||
'X-Workspace-Id': 'workspace-a',
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
mcp_service.get_mcp_server_by_name.assert_awaited_once_with('pab1it0/prometheus')
|
||||
mcp_service.get_mcp_server_by_name.assert_awaited_once()
|
||||
context, server_name = mcp_service.get_mcp_server_by_name.await_args.args
|
||||
assert context.workspace_uuid == 'workspace-a'
|
||||
assert server_name == 'pab1it0/prometheus'
|
||||
payload = await response.get_json()
|
||||
assert payload['data']['server']['name'] == 'pab1it0/prometheus'
|
||||
|
||||
@@ -66,11 +116,17 @@ async def test_mcp_resource_route_accepts_encoded_slash_name():
|
||||
|
||||
response = await client.get(
|
||||
'/api/v1/mcp/servers/pab1it0%2Fprometheus/resources',
|
||||
headers={'Authorization': 'Bearer test-token'},
|
||||
headers={
|
||||
'Authorization': 'Bearer test-token',
|
||||
'X-Workspace-Id': 'workspace-a',
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
mcp_service.get_mcp_server_by_name.assert_not_awaited()
|
||||
mcp_service.get_mcp_server_resources.assert_awaited_once_with('pab1it0/prometheus')
|
||||
mcp_service.get_mcp_server_resources.assert_awaited_once()
|
||||
context, server_name = mcp_service.get_mcp_server_resources.await_args.args
|
||||
assert context.workspace_uuid == 'workspace-a'
|
||||
assert server_name == 'pab1it0/prometheus'
|
||||
payload = await response.get_json()
|
||||
assert payload['data']['resource_capabilities'] == {'subscribe': False}
|
||||
|
||||
@@ -0,0 +1,121 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import sqlalchemy as sa
|
||||
from sqlalchemy.ext.asyncio import create_async_engine
|
||||
|
||||
from langbot.pkg.api.http.service.apikey import ApiKeyIdentity
|
||||
from langbot.pkg.api.mcp.context import get_request_context
|
||||
from langbot.pkg.api.mcp.mount import MCPMount
|
||||
from langbot.pkg.persistence.mgr import PersistenceManager, PersistenceMode
|
||||
from langbot.pkg.persistence.tenant_uow import PersistenceScopeKind
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_mcp_mount_keeps_request_context_but_no_session_during_stream_wait(tmp_path) -> None:
|
||||
engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "mcp-short-scope.db"}')
|
||||
table = sa.Table('mcp_scope_probe', sa.MetaData(), sa.Column('id', sa.Integer, primary_key=True))
|
||||
manager = PersistenceManager(object(), mode=PersistenceMode.CLOUD_RUNTIME)
|
||||
manager.db = SimpleNamespace(get_engine=lambda: engine)
|
||||
checked_out = 0
|
||||
|
||||
def on_checkout(*_args):
|
||||
nonlocal checked_out
|
||||
checked_out += 1
|
||||
|
||||
def on_checkin(*_args):
|
||||
nonlocal checked_out
|
||||
checked_out -= 1
|
||||
|
||||
sa.event.listen(engine.sync_engine, 'checkout', on_checkout)
|
||||
sa.event.listen(engine.sync_engine, 'checkin', on_checkin)
|
||||
try:
|
||||
async with engine.begin() as conn:
|
||||
await conn.run_sync(table.metadata.create_all)
|
||||
|
||||
identity = ApiKeyIdentity(
|
||||
instance_uuid='instance-1',
|
||||
workspace_uuid='workspace-1',
|
||||
placement_generation=7,
|
||||
api_key_uuid='key-1',
|
||||
permissions=frozenset({'pipelines:read'}),
|
||||
)
|
||||
app = SimpleNamespace(
|
||||
apikey_service=SimpleNamespace(authenticate_api_key=AsyncMock(return_value=identity)),
|
||||
persistence_mgr=manager,
|
||||
deployment_admission=None,
|
||||
deployment=None,
|
||||
)
|
||||
stream_waiting = asyncio.Event()
|
||||
release_stream = asyncio.Event()
|
||||
observations: list[tuple[str, str, bool]] = []
|
||||
|
||||
async def fake_mcp_asgi(scope, receive, send):
|
||||
del scope, receive
|
||||
context = get_request_context()
|
||||
assert manager.current_scope().kind is PersistenceScopeKind.WORKSPACE
|
||||
assert manager.current_session() is None
|
||||
await manager.execute_async(sa.select(table.c.id))
|
||||
assert manager.current_session() is None
|
||||
observations.append((context.request_id, context.workspace_uuid, manager.current_session() is None))
|
||||
stream_waiting.set()
|
||||
await release_stream.wait()
|
||||
preserved_context = get_request_context()
|
||||
observations.append(
|
||||
(
|
||||
preserved_context.request_id,
|
||||
preserved_context.workspace_uuid,
|
||||
manager.current_session() is None,
|
||||
)
|
||||
)
|
||||
await manager.execute_async(sa.select(table.c.id))
|
||||
assert manager.current_session() is None
|
||||
await send({'type': 'http.response.start', 'status': 200, 'headers': []})
|
||||
await send({'type': 'http.response.body', 'body': b'{}'})
|
||||
|
||||
async def unused_quart_asgi(scope, receive, send):
|
||||
del scope, receive, send
|
||||
raise AssertionError('MCP request was routed to Quart')
|
||||
|
||||
mount = MCPMount.__new__(MCPMount)
|
||||
mount.ap = app
|
||||
mount._mcp_asgi = fake_mcp_asgi
|
||||
sent_messages: list[dict] = []
|
||||
|
||||
async def receive():
|
||||
return {'type': 'http.request', 'body': b'', 'more_body': False}
|
||||
|
||||
async def send(message):
|
||||
sent_messages.append(message)
|
||||
|
||||
async def release_after_observation() -> None:
|
||||
await asyncio.wait_for(stream_waiting.wait(), timeout=2)
|
||||
assert checked_out == 0
|
||||
release_stream.set()
|
||||
|
||||
release_task = asyncio.create_task(release_after_observation())
|
||||
await mount.wrap(unused_quart_asgi)(
|
||||
{
|
||||
'type': 'http',
|
||||
'path': '/mcp',
|
||||
'headers': [(b'x-api-key', b'secret')],
|
||||
},
|
||||
receive,
|
||||
send,
|
||||
)
|
||||
await release_task
|
||||
|
||||
assert sent_messages[0]['status'] == 200
|
||||
assert len(observations) == 2
|
||||
assert observations[0] == observations[1]
|
||||
assert observations[0][1:] == ('workspace-1', True)
|
||||
assert checked_out == 0
|
||||
assert manager.current_scope() is None
|
||||
with pytest.raises(RuntimeError, match='context is unavailable'):
|
||||
get_request_context()
|
||||
finally:
|
||||
await engine.dispose()
|
||||
@@ -0,0 +1,139 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import asynccontextmanager
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
import quart
|
||||
|
||||
from langbot.pkg.api.http.context import ExecutionContext
|
||||
from langbot.pkg.workspace.errors import WorkspaceNotFoundError
|
||||
|
||||
|
||||
CONTEXT = ExecutionContext(
|
||||
instance_uuid='instance-a',
|
||||
workspace_uuid='workspace-a',
|
||||
placement_generation=4,
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture(scope='module')
|
||||
def plugin_router_cls():
|
||||
from tests.utils.import_isolation import MockLifecycleControlScope, isolated_sys_modules
|
||||
|
||||
class FakeMinimalApplication:
|
||||
pass
|
||||
|
||||
mock_app = Mock(Application=FakeMinimalApplication)
|
||||
mock_entities = Mock(LifecycleControlScope=MockLifecycleControlScope)
|
||||
clear = [
|
||||
'langbot.pkg.core.taskmgr',
|
||||
'langbot.pkg.api.http.controller.group',
|
||||
'langbot.pkg.api.http.controller.groups',
|
||||
'langbot.pkg.api.http.controller.groups.plugins',
|
||||
'langbot.pkg.api.http.controller.main',
|
||||
]
|
||||
with isolated_sys_modules(
|
||||
mocks={
|
||||
'langbot.pkg.core.app': mock_app,
|
||||
'langbot.pkg.core.entities': mock_entities,
|
||||
},
|
||||
clear=clear,
|
||||
):
|
||||
from langbot.pkg.api.http.controller.groups.plugins import PluginsRouterGroup
|
||||
|
||||
yield PluginsRouterGroup
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_public_plugin_asset_route_is_disabled_for_multi_workspace_policy(plugin_router_cls):
|
||||
connector = SimpleNamespace(
|
||||
get_plugin_icon=AsyncMock(),
|
||||
require_workspace_context=AsyncMock(),
|
||||
)
|
||||
ap = SimpleNamespace(
|
||||
plugin_connector=connector,
|
||||
workspace_service=SimpleNamespace(
|
||||
policy=SimpleNamespace(multi_workspace_enabled=True),
|
||||
),
|
||||
)
|
||||
quart_app = quart.Quart(__name__)
|
||||
router = plugin_router_cls(ap, quart_app)
|
||||
await router.initialize()
|
||||
|
||||
response = await quart_app.test_client().get('/api/v1/plugins/author/plugin/icon')
|
||||
|
||||
assert response.status_code == 404
|
||||
connector.require_workspace_context.assert_not_awaited()
|
||||
connector.get_plugin_icon.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_public_plugin_asset_uses_trusted_oss_singleton_binding(plugin_router_cls):
|
||||
connector = SimpleNamespace(
|
||||
require_workspace_context=AsyncMock(side_effect=lambda context: context),
|
||||
)
|
||||
binding = SimpleNamespace(
|
||||
instance_uuid='instance-a',
|
||||
workspace_uuid='workspace-a',
|
||||
placement_generation=4,
|
||||
)
|
||||
router = object.__new__(plugin_router_cls)
|
||||
router.ap = SimpleNamespace(
|
||||
plugin_connector=connector,
|
||||
workspace_service=SimpleNamespace(
|
||||
policy=SimpleNamespace(multi_workspace_enabled=False),
|
||||
get_local_execution_binding=AsyncMock(return_value=binding),
|
||||
),
|
||||
)
|
||||
|
||||
result = await router._require_public_plugin_runtime_context()
|
||||
|
||||
assert result == CONTEXT
|
||||
connector.require_workspace_context.assert_awaited_once_with(CONTEXT)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_background_plugin_operation_refences_captured_generation(plugin_router_cls):
|
||||
operation = AsyncMock()
|
||||
connector = SimpleNamespace(
|
||||
require_workspace_context=AsyncMock(side_effect=WorkspaceNotFoundError('Plugin resource not found')),
|
||||
)
|
||||
router = object.__new__(plugin_router_cls)
|
||||
router.ap = SimpleNamespace(plugin_connector=connector)
|
||||
|
||||
with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
|
||||
await router._run_fenced_plugin_operation(CONTEXT, operation)
|
||||
|
||||
operation.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_background_plugin_operation_revalidates_inside_short_tenant_uow(plugin_router_cls):
|
||||
scopes = []
|
||||
|
||||
@asynccontextmanager
|
||||
async def tenant_uow(workspace_uuid):
|
||||
scopes.append(workspace_uuid)
|
||||
yield
|
||||
|
||||
connector = SimpleNamespace(
|
||||
require_workspace_context=AsyncMock(side_effect=lambda context: context),
|
||||
)
|
||||
operation = AsyncMock(return_value='done')
|
||||
router = object.__new__(plugin_router_cls)
|
||||
router.ap = SimpleNamespace(
|
||||
plugin_connector=connector,
|
||||
persistence_mgr=SimpleNamespace(
|
||||
mode=SimpleNamespace(value='cloud_runtime'),
|
||||
tenant_uow=tenant_uow,
|
||||
),
|
||||
)
|
||||
|
||||
result = await router._run_fenced_plugin_operation(CONTEXT, operation)
|
||||
|
||||
assert result == 'done'
|
||||
assert scopes == [CONTEXT.workspace_uuid]
|
||||
connector.require_workspace_context.assert_awaited_once_with(CONTEXT)
|
||||
operation.assert_awaited_once()
|
||||
@@ -0,0 +1,196 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import quart
|
||||
|
||||
from langbot.pkg.api.http.controller.groups.knowledge.base import KnowledgeBaseRouterGroup
|
||||
from langbot.pkg.api.http.controller.groups.pipelines.pipelines import PipelinesRouterGroup
|
||||
from langbot.pkg.api.http.controller.groups.provider.models import LLMModelsRouterGroup
|
||||
from langbot.pkg.api.http.controller.groups.provider.providers import ModelProvidersRouterGroup
|
||||
from langbot.pkg.api.http.controller.groups.resources.mcp import MCPRouterGroup
|
||||
from langbot.pkg.api.http.controller.groups.webhook_mgmt import WebhookManagementRouterGroup
|
||||
from langbot.pkg.api.http.service.secrets import mask_secret_value, redact_secrets
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
WORKSPACE_UUID = '11111111-1111-4111-8111-111111111111'
|
||||
|
||||
RAW_PIPELINE = {
|
||||
'uuid': 'pipeline-test',
|
||||
'config': {'ai': {'n8n': {'webhook-url': 'https://hook.invalid/bearer-secret'}}},
|
||||
}
|
||||
RAW_MODEL = {
|
||||
'uuid': 'model-test',
|
||||
'provider_uuid': 'provider-test',
|
||||
'extra_args': {'headers': {'Authorization': 'Bearer model-secret'}},
|
||||
}
|
||||
RAW_PROVIDER = {
|
||||
'uuid': 'provider-test',
|
||||
'base_url': 'https://provider-user:provider-password@provider.invalid/v1?token=url-secret®ion=sg',
|
||||
'api_keys': ['provider-secret'],
|
||||
}
|
||||
RAW_MCP_SERVER = {
|
||||
'uuid': 'mcp-test',
|
||||
'name': 'MCP Test',
|
||||
'extra_args': {'url': 'https://mcp-user:mcp-password@mcp.invalid/connect?api_key=url-secret&transport=http'},
|
||||
}
|
||||
RAW_KNOWLEDGE_BASE = {
|
||||
'uuid': 'kb-test',
|
||||
'creation_settings': {'dify_apikey': 'knowledge-secret'},
|
||||
}
|
||||
RAW_WEBHOOK = {'id': 1, 'url': 'https://hook.invalid/path?token=webhook-secret'}
|
||||
|
||||
|
||||
def _access(role: str):
|
||||
return SimpleNamespace(
|
||||
workspace=SimpleNamespace(uuid=WORKSPACE_UUID),
|
||||
membership=SimpleNamespace(uuid='membership-test', role=role, projection_revision=1),
|
||||
execution=SimpleNamespace(instance_uuid='instance-test', placement_generation=1),
|
||||
)
|
||||
|
||||
|
||||
async def _create_client(role: str):
|
||||
application = SimpleNamespace()
|
||||
account = SimpleNamespace(uuid='account-test', user='test@example.com')
|
||||
application.user_service = SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account))
|
||||
application.apikey_service = SimpleNamespace(authenticate_api_key=AsyncMock(return_value=None))
|
||||
application.workspace_collaboration_service = SimpleNamespace(
|
||||
resolve_account_workspace=AsyncMock(return_value=_access(role))
|
||||
)
|
||||
|
||||
async def get_pipelines(_context, *_args, include_secret=False):
|
||||
value = copy.deepcopy(RAW_PIPELINE)
|
||||
return [value] if include_secret else [redact_secrets(value)]
|
||||
|
||||
async def get_pipeline(_context, _uuid, *, include_secret=False):
|
||||
value = copy.deepcopy(RAW_PIPELINE)
|
||||
return value if include_secret else redact_secrets(value)
|
||||
|
||||
application.pipeline_service = SimpleNamespace(
|
||||
get_pipelines=AsyncMock(side_effect=get_pipelines),
|
||||
get_pipeline=AsyncMock(side_effect=get_pipeline),
|
||||
)
|
||||
application.plugin_connector = SimpleNamespace(list_plugins=AsyncMock(return_value=[]))
|
||||
application.mcp_service = SimpleNamespace(
|
||||
get_mcp_servers=AsyncMock(return_value=[redact_secrets(copy.deepcopy(RAW_MCP_SERVER))])
|
||||
)
|
||||
application.skill_service = SimpleNamespace(list_skills=AsyncMock(return_value=[]))
|
||||
|
||||
async def get_models_by_provider(_context, _provider_uuid, *, include_secret=False):
|
||||
value = copy.deepcopy(RAW_MODEL)
|
||||
return [value] if include_secret else [redact_secrets(value)]
|
||||
|
||||
application.llm_model_service = SimpleNamespace(
|
||||
get_llm_models_by_provider=AsyncMock(side_effect=get_models_by_provider)
|
||||
)
|
||||
|
||||
async def get_providers(_context, *, include_secret=False):
|
||||
value = copy.deepcopy(RAW_PROVIDER)
|
||||
return [value] if include_secret else [redact_secrets(value)]
|
||||
|
||||
application.provider_service = SimpleNamespace(
|
||||
get_providers=AsyncMock(side_effect=get_providers),
|
||||
get_provider_model_counts=AsyncMock(return_value={'llm_count': 0, 'embedding_count': 0, 'rerank_count': 0}),
|
||||
)
|
||||
|
||||
async def get_knowledge_bases(_context, *, include_secret=False):
|
||||
value = copy.deepcopy(RAW_KNOWLEDGE_BASE)
|
||||
return [value] if include_secret else [redact_secrets(value)]
|
||||
|
||||
application.knowledge_service = SimpleNamespace(get_knowledge_bases=AsyncMock(side_effect=get_knowledge_bases))
|
||||
|
||||
async def get_webhooks(_context, *, include_secret=False):
|
||||
value = copy.deepcopy(RAW_WEBHOOK)
|
||||
if not include_secret:
|
||||
value['url'] = mask_secret_value(value['url'])
|
||||
return [value]
|
||||
|
||||
application.webhook_service = SimpleNamespace(get_webhooks=AsyncMock(side_effect=get_webhooks))
|
||||
|
||||
quart_app = quart.Quart(__name__)
|
||||
for router_type in (
|
||||
PipelinesRouterGroup,
|
||||
LLMModelsRouterGroup,
|
||||
ModelProvidersRouterGroup,
|
||||
MCPRouterGroup,
|
||||
KnowledgeBaseRouterGroup,
|
||||
WebhookManagementRouterGroup,
|
||||
):
|
||||
await router_type(application, quart_app).initialize()
|
||||
return application, quart_app.test_client()
|
||||
|
||||
|
||||
def _headers() -> dict[str, str]:
|
||||
return {'Authorization': 'Bearer test-token', 'X-Workspace-Id': WORKSPACE_UUID}
|
||||
|
||||
|
||||
@pytest.mark.parametrize('role', ['viewer', 'operator'])
|
||||
async def test_viewer_and_operator_resource_reads_are_redacted(role: str):
|
||||
application, client = await _create_client(role)
|
||||
|
||||
pipeline = (await (await client.get('/api/v1/pipelines', headers=_headers())).get_json())['data']['pipelines'][0]
|
||||
model = (
|
||||
await (
|
||||
await client.get(
|
||||
'/api/v1/provider/models/llm?provider_uuid=provider-test',
|
||||
headers=_headers(),
|
||||
)
|
||||
).get_json()
|
||||
)['data']['models'][0]
|
||||
provider = (await (await client.get('/api/v1/provider/providers', headers=_headers())).get_json())['data'][
|
||||
'providers'
|
||||
][0]
|
||||
mcp_server = (await (await client.get('/api/v1/mcp/servers', headers=_headers())).get_json())['data']['servers'][0]
|
||||
knowledge_base = (await (await client.get('/api/v1/knowledge/bases', headers=_headers())).get_json())['data'][
|
||||
'bases'
|
||||
][0]
|
||||
webhook = (await (await client.get('/api/v1/webhooks', headers=_headers())).get_json())['data']['webhooks'][0]
|
||||
|
||||
assert pipeline['config']['ai']['n8n']['webhook-url'] == '***'
|
||||
assert model['extra_args']['headers']['Authorization'] == '***'
|
||||
assert provider['api_keys'] == ['***']
|
||||
assert provider['base_url'] == 'https://***@provider.invalid/v1?token=***®ion=sg'
|
||||
assert mcp_server['extra_args']['url'] == 'https://***@mcp.invalid/connect?api_key=***&transport=http'
|
||||
assert knowledge_base['creation_settings']['dify_apikey'] == '***'
|
||||
assert webhook['url'] == '***'
|
||||
assert application.pipeline_service.get_pipelines.await_args.kwargs['include_secret'] is False
|
||||
assert application.llm_model_service.get_llm_models_by_provider.await_args.kwargs['include_secret'] is False
|
||||
assert application.provider_service.get_providers.await_args.kwargs['include_secret'] is False
|
||||
assert application.knowledge_service.get_knowledge_bases.await_args.kwargs['include_secret'] is False
|
||||
assert application.webhook_service.get_webhooks.await_args.kwargs['include_secret'] is False
|
||||
|
||||
|
||||
async def test_resource_manager_receives_credentials_needed_for_management():
|
||||
application, client = await _create_client('developer')
|
||||
|
||||
pipeline = (await (await client.get('/api/v1/pipelines', headers=_headers())).get_json())['data']['pipelines'][0]
|
||||
model = (
|
||||
await (
|
||||
await client.get(
|
||||
'/api/v1/provider/models/llm?provider_uuid=provider-test',
|
||||
headers=_headers(),
|
||||
)
|
||||
).get_json()
|
||||
)['data']['models'][0]
|
||||
provider = (await (await client.get('/api/v1/provider/providers', headers=_headers())).get_json())['data'][
|
||||
'providers'
|
||||
][0]
|
||||
knowledge_base = (await (await client.get('/api/v1/knowledge/bases', headers=_headers())).get_json())['data'][
|
||||
'bases'
|
||||
][0]
|
||||
webhook = (await (await client.get('/api/v1/webhooks', headers=_headers())).get_json())['data']['webhooks'][0]
|
||||
|
||||
assert pipeline == RAW_PIPELINE
|
||||
assert model == RAW_MODEL
|
||||
assert provider['api_keys'] == ['provider-secret']
|
||||
assert knowledge_base == RAW_KNOWLEDGE_BASE
|
||||
assert webhook == RAW_WEBHOOK
|
||||
assert application.pipeline_service.get_pipelines.await_args.kwargs['include_secret'] is True
|
||||
assert application.llm_model_service.get_llm_models_by_provider.await_args.kwargs['include_secret'] is True
|
||||
assert application.provider_service.get_providers.await_args.kwargs['include_secret'] is True
|
||||
assert application.knowledge_service.get_knowledge_bases.await_args.kwargs['include_secret'] is True
|
||||
assert application.webhook_service.get_webhooks.await_args.kwargs['include_secret'] is True
|
||||
@@ -0,0 +1,110 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
import pytest
|
||||
import quart
|
||||
|
||||
from langbot.pkg.api.http.controller.groups.stats import StatsRouterGroup
|
||||
|
||||
|
||||
pytestmark = pytest.mark.asyncio
|
||||
|
||||
|
||||
def session(
|
||||
workspace_uuid: str,
|
||||
*,
|
||||
placement_generation: int = 1,
|
||||
conversation_count: int = 0,
|
||||
):
|
||||
return SimpleNamespace(
|
||||
instance_uuid='instance-test',
|
||||
workspace_uuid=workspace_uuid,
|
||||
placement_generation=placement_generation,
|
||||
conversations=[object() for _ in range(conversation_count)],
|
||||
)
|
||||
|
||||
|
||||
async def create_client(*, role='viewer'):
|
||||
quart_app = quart.Quart(__name__)
|
||||
account = SimpleNamespace(uuid='account-test', user='test@example.com')
|
||||
user_service = SimpleNamespace(
|
||||
get_authenticated_account=AsyncMock(return_value=account),
|
||||
)
|
||||
access = SimpleNamespace(
|
||||
execution=SimpleNamespace(
|
||||
instance_uuid='instance-test',
|
||||
placement_generation=1,
|
||||
),
|
||||
workspace=SimpleNamespace(uuid='workspace-a'),
|
||||
membership=SimpleNamespace(
|
||||
uuid='membership-test',
|
||||
role=role,
|
||||
projection_revision=1,
|
||||
),
|
||||
)
|
||||
collaboration_service = SimpleNamespace(
|
||||
resolve_account_workspace=AsyncMock(return_value=access),
|
||||
)
|
||||
|
||||
def get_query_count(context):
|
||||
assert context.instance_uuid == 'instance-test'
|
||||
assert context.workspace_uuid == 'workspace-a'
|
||||
assert context.placement_generation == 1
|
||||
return 7
|
||||
|
||||
ap = SimpleNamespace(
|
||||
user_service=user_service,
|
||||
workspace_collaboration_service=collaboration_service,
|
||||
sess_mgr=SimpleNamespace(
|
||||
session_list=[
|
||||
session('workspace-a', conversation_count=2),
|
||||
session('workspace-b', conversation_count=5),
|
||||
session(
|
||||
'workspace-a',
|
||||
placement_generation=2,
|
||||
conversation_count=3,
|
||||
),
|
||||
SimpleNamespace(conversations=[object()] * 11),
|
||||
]
|
||||
),
|
||||
query_pool=SimpleNamespace(get_query_count=get_query_count),
|
||||
)
|
||||
router = StatsRouterGroup(ap, quart_app)
|
||||
await router.initialize()
|
||||
return quart_app.test_client(), collaboration_service
|
||||
|
||||
|
||||
async def test_basic_stats_are_scoped_to_selected_workspace_placement():
|
||||
client, collaboration_service = await create_client()
|
||||
|
||||
response = await client.get(
|
||||
'/api/v1/stats/basic',
|
||||
headers={
|
||||
'Authorization': 'Bearer test-token',
|
||||
'X-Workspace-Id': 'workspace-a',
|
||||
},
|
||||
)
|
||||
|
||||
assert response.status_code == 200
|
||||
payload = await response.get_json()
|
||||
assert payload['data'] == {
|
||||
'active_session_count': 1,
|
||||
'conversation_count': 2,
|
||||
'query_count': 7,
|
||||
}
|
||||
collaboration_service.resolve_account_workspace.assert_awaited_once_with('account-test', 'workspace-a')
|
||||
|
||||
|
||||
async def test_basic_stats_requires_resource_view_permission():
|
||||
client, _ = await create_client(role='unknown-role')
|
||||
|
||||
response = await client.get(
|
||||
'/api/v1/stats/basic',
|
||||
headers={'Authorization': 'Bearer test-token'},
|
||||
)
|
||||
|
||||
assert response.status_code == 403
|
||||
payload = await response.get_json()
|
||||
assert payload['code'] == 'permission_denied'
|
||||
@@ -0,0 +1,144 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from contextlib import asynccontextmanager
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
import pytest
|
||||
|
||||
from langbot.pkg.api.http.context import (
|
||||
PrincipalContext,
|
||||
PrincipalType,
|
||||
RequestContext,
|
||||
WorkspaceContext,
|
||||
)
|
||||
from langbot.pkg.api.http.controller.groups.pipelines.websocket_chat import WebSocketChatRouterGroup
|
||||
from langbot.pkg.api.http.controller.groups.pipelines.websocket_chat import (
|
||||
create_scoped_duplex_tasks,
|
||||
)
|
||||
from langbot.pkg.api.http.controller.groups.pipelines.websocket_chat import wait_for_duplex_tasks
|
||||
from langbot.pkg.utils.bounded_executor import current_blocking_work_scope
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_websocket_pipeline_lookup_opens_workspace_uow_after_auth_scope_closed() -> None:
|
||||
workspace_uuid = 'workspace-a'
|
||||
scopes: list[str] = []
|
||||
in_scope = False
|
||||
|
||||
@asynccontextmanager
|
||||
async def tenant_uow(selected_workspace_uuid: str):
|
||||
nonlocal in_scope
|
||||
assert not in_scope
|
||||
in_scope = True
|
||||
scopes.append(selected_workspace_uuid)
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
in_scope = False
|
||||
|
||||
async def get_pipeline(_context, _pipeline_uuid):
|
||||
assert in_scope
|
||||
return {'uuid': 'pipeline-a'}
|
||||
|
||||
adapter = Mock()
|
||||
router = object.__new__(WebSocketChatRouterGroup)
|
||||
router.ap = SimpleNamespace(
|
||||
persistence_mgr=SimpleNamespace(
|
||||
mode=SimpleNamespace(value='cloud_runtime'),
|
||||
tenant_uow=tenant_uow,
|
||||
),
|
||||
pipeline_service=SimpleNamespace(get_pipeline=AsyncMock(side_effect=get_pipeline)),
|
||||
platform_mgr=SimpleNamespace(get_websocket_proxy_bot=AsyncMock(return_value=SimpleNamespace(adapter=adapter))),
|
||||
)
|
||||
request_context = RequestContext(
|
||||
instance_uuid='instance-a',
|
||||
placement_generation=1,
|
||||
request_id='request-a',
|
||||
auth_type='user_token',
|
||||
principal=PrincipalContext(
|
||||
principal_type=PrincipalType.ACCOUNT,
|
||||
account_uuid='account-a',
|
||||
),
|
||||
workspace=WorkspaceContext(
|
||||
workspace_uuid=workspace_uuid,
|
||||
membership_uuid='membership-a',
|
||||
role='owner',
|
||||
permissions=frozenset(),
|
||||
),
|
||||
)
|
||||
|
||||
result = await router._get_scoped_adapter(request_context, 'pipeline-a')
|
||||
|
||||
assert result is adapter
|
||||
assert scopes == [workspace_uuid]
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplex_websocket_tasks_cancel_blocked_peer_when_one_direction_ends() -> None:
|
||||
blocked = asyncio.Event()
|
||||
|
||||
async def receive_forever() -> None:
|
||||
blocked.set()
|
||||
await asyncio.Future()
|
||||
|
||||
async def send_finishes() -> None:
|
||||
await blocked.wait()
|
||||
|
||||
receive_task = asyncio.create_task(receive_forever())
|
||||
send_task = asyncio.create_task(send_finishes())
|
||||
|
||||
await asyncio.wait_for(
|
||||
wait_for_duplex_tasks(receive_task, send_task),
|
||||
timeout=1,
|
||||
)
|
||||
|
||||
assert receive_task.cancelled()
|
||||
assert send_task.done()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplex_websocket_tasks_allow_terminal_send_to_drain() -> None:
|
||||
receive_finished = asyncio.Event()
|
||||
send_drained = asyncio.Event()
|
||||
|
||||
async def receive_finishes() -> None:
|
||||
receive_finished.set()
|
||||
|
||||
async def send_terminal_frame() -> None:
|
||||
await receive_finished.wait()
|
||||
await asyncio.sleep(0)
|
||||
send_drained.set()
|
||||
|
||||
receive_task = asyncio.create_task(receive_finishes())
|
||||
send_task = asyncio.create_task(send_terminal_frame())
|
||||
|
||||
await wait_for_duplex_tasks(receive_task, send_task)
|
||||
|
||||
assert send_drained.is_set()
|
||||
assert send_task.done()
|
||||
assert not send_task.cancelled()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplex_websocket_tasks_share_trusted_workspace_budget() -> None:
|
||||
observed: list[tuple[str, str | None]] = []
|
||||
|
||||
async def observe(direction: str) -> None:
|
||||
await asyncio.sleep(0)
|
||||
observed.append((direction, current_blocking_work_scope()))
|
||||
|
||||
receive_task, send_task = create_scoped_duplex_tasks(
|
||||
observe('receive'),
|
||||
observe('send'),
|
||||
'workspace-a',
|
||||
)
|
||||
|
||||
await asyncio.gather(receive_task, send_task)
|
||||
|
||||
assert sorted(observed) == [
|
||||
('receive', 'workspace-a'),
|
||||
('send', 'workspace-a'),
|
||||
]
|
||||
assert current_blocking_work_scope() is None
|
||||
Reference in New Issue
Block a user