diff --git a/src/langbot/pkg/api/http/controller/groups/provider/models.py b/src/langbot/pkg/api/http/controller/groups/provider/models.py index 236000d9f..fed754201 100644 --- a/src/langbot/pkg/api/http/controller/groups/provider/models.py +++ b/src/langbot/pkg/api/http/controller/groups/provider/models.py @@ -3,6 +3,7 @@ import quart from ....authz import Permission, has_permission from ....context import RequestContext from ... import group +from .query import resolve_include_secret @group.group_class('models/llm', '/api/v1/provider/models/llm') @@ -16,7 +17,12 @@ class LLMModelsRouterGroup(group.RouterGroup): ) async def _(request_context: RequestContext) -> str: provider_uuid = quart.request.args.get('provider_uuid') - include_secret = has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE) + include_secret, error = resolve_include_secret( + quart.request.args.get('include_secret'), + permitted=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + ) + if error: + return self.http_status(400, -1, error) if provider_uuid: models = await self.ap.llm_model_service.get_llm_models_by_provider( request_context, @@ -53,10 +59,16 @@ class LLMModelsRouterGroup(group.RouterGroup): permission=Permission.RESOURCE_VIEW, ) async def _(model_uuid: str, request_context: RequestContext) -> str: + include_secret, error = resolve_include_secret( + quart.request.args.get('include_secret'), + permitted=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + ) + if error: + return self.http_status(400, -1, error) model = await self.ap.llm_model_service.get_llm_model( request_context, model_uuid, - include_secret=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + include_secret=include_secret, ) if model is None: return self.http_status(404, -1, 'model not found') @@ -111,7 +123,12 @@ class EmbeddingModelsRouterGroup(group.RouterGroup): ) async def _(request_context: RequestContext) -> str: provider_uuid = quart.request.args.get('provider_uuid') - include_secret = has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE) + include_secret, error = resolve_include_secret( + quart.request.args.get('include_secret'), + permitted=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + ) + if error: + return self.http_status(400, -1, error) if provider_uuid: models = await self.ap.embedding_models_service.get_embedding_models_by_provider( request_context, @@ -148,10 +165,16 @@ class EmbeddingModelsRouterGroup(group.RouterGroup): permission=Permission.RESOURCE_VIEW, ) async def _(model_uuid: str, request_context: RequestContext) -> str: + include_secret, error = resolve_include_secret( + quart.request.args.get('include_secret'), + permitted=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + ) + if error: + return self.http_status(400, -1, error) model = await self.ap.embedding_models_service.get_embedding_model( request_context, model_uuid, - include_secret=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + include_secret=include_secret, ) if model is None: return self.http_status(404, -1, 'model not found') @@ -208,7 +231,12 @@ class RerankModelsRouterGroup(group.RouterGroup): ) async def _(request_context: RequestContext) -> str: provider_uuid = quart.request.args.get('provider_uuid') - include_secret = has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE) + include_secret, error = resolve_include_secret( + quart.request.args.get('include_secret'), + permitted=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + ) + if error: + return self.http_status(400, -1, error) if provider_uuid: models = await self.ap.rerank_models_service.get_rerank_models_by_provider( request_context, @@ -245,10 +273,16 @@ class RerankModelsRouterGroup(group.RouterGroup): permission=Permission.RESOURCE_VIEW, ) async def _(model_uuid: str, request_context: RequestContext) -> str: + include_secret, error = resolve_include_secret( + quart.request.args.get('include_secret'), + permitted=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + ) + if error: + return self.http_status(400, -1, error) model = await self.ap.rerank_models_service.get_rerank_model( request_context, model_uuid, - include_secret=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + include_secret=include_secret, ) if model is None: return self.http_status(404, -1, 'model not found') diff --git a/src/langbot/pkg/api/http/controller/groups/provider/providers.py b/src/langbot/pkg/api/http/controller/groups/provider/providers.py index a7097745c..107cb9b7f 100644 --- a/src/langbot/pkg/api/http/controller/groups/provider/providers.py +++ b/src/langbot/pkg/api/http/controller/groups/provider/providers.py @@ -3,6 +3,7 @@ import quart from ....authz import Permission, has_permission from ....context import RequestContext from ... import group +from .query import resolve_include_secret @group.group_class('models/providers', '/api/v1/provider/providers') @@ -89,9 +90,15 @@ class ModelProvidersRouterGroup(group.RouterGroup): permission=Permission.RESOURCE_VIEW, ) async def _(request_context: RequestContext) -> str: + include_secret, error = resolve_include_secret( + quart.request.args.get('include_secret'), + permitted=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + ) + if error: + return self.http_status(400, -1, error) providers = await self.ap.provider_service.get_providers( request_context, - include_secret=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + include_secret=include_secret, ) for provider in providers: counts = await self.ap.provider_service.get_provider_model_counts(request_context, provider['uuid']) @@ -121,10 +128,16 @@ class ModelProvidersRouterGroup(group.RouterGroup): permission=Permission.RESOURCE_VIEW, ) async def _(provider_uuid: str, request_context: RequestContext) -> str: + include_secret, error = resolve_include_secret( + quart.request.args.get('include_secret'), + permitted=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + ) + if error: + return self.http_status(400, -1, error) provider = await self.ap.provider_service.get_provider( request_context, provider_uuid, - include_secret=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE), + include_secret=include_secret, ) if provider is None: return self.http_status(404, -1, 'provider not found') diff --git a/src/langbot/pkg/api/http/controller/groups/provider/query.py b/src/langbot/pkg/api/http/controller/groups/provider/query.py new file mode 100644 index 000000000..bd1793fe2 --- /dev/null +++ b/src/langbot/pkg/api/http/controller/groups/provider/query.py @@ -0,0 +1,15 @@ +from __future__ import annotations + + +def resolve_include_secret(raw_value: str | None, *, permitted: bool) -> tuple[bool, str | None]: + """Resolve the optional secret projection query parameter.""" + + if raw_value is None: + return permitted, None + + value = raw_value.strip().lower() + if value == 'false': + return False, None + if value == 'true': + return permitted, None + return False, 'include_secret must be either true or false' diff --git a/src/langbot/pkg/api/http/controller/groups/system.py b/src/langbot/pkg/api/http/controller/groups/system.py index a8be6ba22..57feaecaf 100644 --- a/src/langbot/pkg/api/http/controller/groups/system.py +++ b/src/langbot/pkg/api/http/controller/groups/system.py @@ -7,11 +7,93 @@ from .. import group from .....utils import constants from .....entity.persistence.metadata import WorkspaceMetadata from ...authz import Permission -from ...context import RequestContext +from ...context import PrincipalType, RequestContext from .....provider.tools.loaders.mcp_policy import stdio_mcp_enabled from .....workspace.invitation_delivery import InvitationDeliveryService +SYSTEM_CAPABILITY_OPERATIONS = ( + 'bot.list', + 'bot.get', + 'bot.create', + 'bot.update', + 'bot.delete', + 'pipeline.list', + 'pipeline.get', + 'pipeline.create', + 'pipeline.update', + 'pipeline.delete', + 'pipeline.copy', + 'task.list', + 'task.get', + 'knowledge_base.list', + 'knowledge_base.get', + 'knowledge_base.create', + 'knowledge_base.update', + 'knowledge_base.delete', + 'knowledge_base.file.list', + 'knowledge_base.file.store', + 'knowledge_base.file.delete', + 'knowledge_base.retrieve', + 'file.document.upload', + 'plugin.install.github', + 'plugin.install.marketplace', + 'plugin.install.local', + 'plugin.upgrade', + 'plugin.get', + 'plugin.list', + 'plugin.config.get', + 'plugin.config.update', + 'plugin.logs', + 'plugin.delete', + 'provider.list', + 'provider.get', + 'provider.create', + 'provider.update', + 'provider.delete', + 'provider.scan_models', + 'model.llm.list', + 'model.llm.get', + 'model.llm.create', + 'model.llm.update', + 'model.llm.delete', + 'model.llm.test', + 'model.embedding.list', + 'model.embedding.get', + 'model.embedding.create', + 'model.embedding.update', + 'model.embedding.delete', + 'model.embedding.test', + 'model.rerank.list', + 'model.rerank.get', + 'model.rerank.create', + 'model.rerank.update', + 'model.rerank.delete', + 'model.rerank.test', + 'skill.list', + 'skill.get', + 'skill.create', + 'skill.update', + 'skill.delete', + 'skill.files.list', + 'skill.files.read', + 'skill.files.write', + 'skill.preview', + 'skill.install.github', + 'skill.install.upload', + 'mcp_server.list', + 'mcp_server.get', + 'mcp_server.create', + 'mcp_server.update', + 'mcp_server.delete', + 'mcp_server.resources', + 'mcp_server.resource_templates', + 'mcp_server.resource_read', + 'mcp_server.logs', + 'mcp_server.test', +) + + @group.group_class('system', '/api/v1/system') class SystemRouterGroup(group.RouterGroup): async def initialize(self) -> None: @@ -26,6 +108,15 @@ class SystemRouterGroup(group.RouterGroup): } ) + @self.route('/capabilities', methods=['GET'], auth_type=group.AuthType.API_KEY) + async def _() -> str: + return self.success( + data={ + 'schema_version': 1, + 'operations': {operation: {'supported': True} for operation in SYSTEM_CAPABILITY_OPERATIONS}, + } + ) + @self.route('/info', methods=['GET'], auth_type=group.AuthType.NONE) async def _() -> str: # Read wizard_status and wizard_progress from metadata table @@ -234,7 +325,7 @@ class SystemRouterGroup(group.RouterGroup): @self.route( '/tasks', methods=['GET'], - auth_type=group.AuthType.USER_TOKEN, + auth_type=group.AuthType.USER_TOKEN_OR_API_KEY, permission=Permission.RESOURCE_VIEW, ) async def _(request_context: RequestContext) -> str: @@ -253,18 +344,23 @@ class SystemRouterGroup(group.RouterGroup): instance_uuid=request_context.instance_uuid, workspace_uuid=request_context.workspace_uuid, placement_generation=request_context.placement_generation, + public=request_context.principal.principal_type == PrincipalType.API_KEY, ) ) @self.route( '/tasks/', methods=['GET'], - auth_type=group.AuthType.USER_TOKEN, + auth_type=group.AuthType.USER_TOKEN_OR_API_KEY, permission=Permission.RESOURCE_VIEW, ) async def _(task_id: str, request_context: RequestContext) -> str: + try: + task_index = int(task_id) + except (TypeError, ValueError): + return self.http_status(404, 404, 'Task not found') task = self.ap.task_mgr.get_task_by_id( - int(task_id), + task_index, instance_uuid=request_context.instance_uuid, workspace_uuid=request_context.workspace_uuid, placement_generation=request_context.placement_generation, @@ -273,6 +369,8 @@ class SystemRouterGroup(group.RouterGroup): if task is None: return self.http_status(404, 404, 'Task not found') + if request_context.principal.principal_type == PrincipalType.API_KEY: + return self.success(data=task.to_public_dict()) return self.success(data=task.to_dict()) @self.route( diff --git a/src/langbot/pkg/core/taskmgr.py b/src/langbot/pkg/core/taskmgr.py index 25cc38e0d..6c5df7b81 100644 --- a/src/langbot/pkg/core/taskmgr.py +++ b/src/langbot/pkg/core/taskmgr.py @@ -1,6 +1,7 @@ from __future__ import annotations import asyncio +import json import typing import datetime import time @@ -197,6 +198,41 @@ class TaskWrapper: }, } + def to_public_dict(self) -> dict: + """Return the stable task projection exposed to API-key callers.""" + if self.task.cancelled(): + status = 'cancelled' + error = {'type': 'task_cancelled', 'message': 'Task was cancelled'} + result = None + elif not self.task.done(): + status = 'running' + error = None + result = None + else: + exception = self.assume_exception() + if exception is not None: + status = 'failed' + error = {'type': 'task_failed', 'message': 'Task execution failed'} + result = None + else: + status = 'succeeded' + error = None + result = self.assume_result() + try: + json.dumps(result) + except (TypeError, ValueError): + result = None + + return { + 'id': self.id, + 'task_type': self.task_type, + 'kind': self.kind, + 'status': status, + 'error': error, + 'result': result, + 'created_at': self.created_at, + } + def cancel(self): self.task.cancel() @@ -325,19 +361,20 @@ class AsyncTaskManager: instance_uuid: str | None = None, workspace_uuid: str | None = None, placement_generation: int | None = None, + public: bool = False, ) -> dict: - return { - 'tasks': [ - t.to_dict() - for t in self.tasks - if (type is None or t.task_type == type) - and (kind is None or t.kind == kind) - and (instance_uuid is None or t.instance_uuid == instance_uuid) - and (workspace_uuid is None or t.workspace_uuid == workspace_uuid) - and (placement_generation is None or t.placement_generation == placement_generation) - ], - 'id_index': TaskWrapper._id_index, - } + tasks = [ + t.to_public_dict() if public else t.to_dict() + for t in self.tasks + if (type is None or t.task_type == type) + and (kind is None or t.kind == kind) + and (instance_uuid is None or t.instance_uuid == instance_uuid) + and (workspace_uuid is None or t.workspace_uuid == workspace_uuid) + and (placement_generation is None or t.placement_generation == placement_generation) + ] + if public: + return {'tasks': tasks} + return {'tasks': tasks, 'id_index': TaskWrapper._id_index} def get_stats(self) -> dict: completed = sum(1 for t in self.tasks if t.task.done()) diff --git a/tests/integration/api/test_workspaces.py b/tests/integration/api/test_workspaces.py index 511607f0d..ed3ca22bd 100644 --- a/tests/integration/api/test_workspaces.py +++ b/tests/integration/api/test_workspaces.py @@ -1,5 +1,6 @@ from __future__ import annotations +import datetime import json import logging from types import SimpleNamespace @@ -22,6 +23,7 @@ from langbot.pkg.api.http.service.apikey import ApiKeyService from langbot.pkg.api.http.service.user import ControlPlaneDirectoryRequiredError, UserService from langbot.pkg.entity.persistence.base import Base from langbot.pkg.entity.persistence.metadata import WorkspaceMetadata +from langbot.pkg.entity.persistence import apikey from langbot.pkg.entity.persistence.user import User from langbot.pkg.entity.persistence.workspace import ( Workspace, @@ -462,6 +464,12 @@ async def test_api_key_context_returns_bound_identity_without_workspace_permissi ) assert invalid_auth.status_code == 401 + invalid_capabilities = await client.get( + '/api/v1/system/capabilities', + headers={'X-API-Key': 'lbk_invalid'}, + ) + assert invalid_capabilities.status_code == 401 + response = await client.get( '/api/v1/system/context', headers={ @@ -478,6 +486,101 @@ async def test_api_key_context_returns_bound_identity_without_workspace_permissi 'permissions': [], } + capabilities_response = await client.get( + '/api/v1/system/capabilities', + headers={ + 'X-API-Key': created['key'], + 'X-Workspace-Id': 'caller-selected-workspace-must-be-ignored', + }, + ) + assert capabilities_response.status_code == 200 + capabilities = (await capabilities_response.get_json())['data'] + assert capabilities['schema_version'] == 1 + assert sorted(capabilities['operations']) == sorted( + [ + 'bot.list', + 'bot.get', + 'bot.create', + 'bot.update', + 'bot.delete', + 'pipeline.list', + 'pipeline.get', + 'pipeline.create', + 'pipeline.update', + 'pipeline.delete', + 'pipeline.copy', + 'task.list', + 'task.get', + 'knowledge_base.list', + 'knowledge_base.get', + 'knowledge_base.create', + 'knowledge_base.update', + 'knowledge_base.delete', + 'knowledge_base.file.list', + 'knowledge_base.file.store', + 'knowledge_base.file.delete', + 'knowledge_base.retrieve', + 'file.document.upload', + 'plugin.install.github', + 'plugin.install.marketplace', + 'plugin.install.local', + 'plugin.upgrade', + 'plugin.get', + 'plugin.list', + 'plugin.config.get', + 'plugin.config.update', + 'plugin.logs', + 'plugin.delete', + 'provider.list', + 'provider.get', + 'provider.create', + 'provider.update', + 'provider.delete', + 'provider.scan_models', + 'model.llm.list', + 'model.llm.get', + 'model.llm.create', + 'model.llm.update', + 'model.llm.delete', + 'model.llm.test', + 'model.embedding.list', + 'model.embedding.get', + 'model.embedding.create', + 'model.embedding.update', + 'model.embedding.delete', + 'model.embedding.test', + 'model.rerank.list', + 'model.rerank.get', + 'model.rerank.create', + 'model.rerank.update', + 'model.rerank.delete', + 'model.rerank.test', + 'skill.list', + 'skill.get', + 'skill.create', + 'skill.update', + 'skill.delete', + 'skill.files.list', + 'skill.files.read', + 'skill.files.write', + 'skill.preview', + 'skill.install.github', + 'skill.install.upload', + 'mcp_server.list', + 'mcp_server.get', + 'mcp_server.create', + 'mcp_server.update', + 'mcp_server.delete', + 'mcp_server.resources', + 'mcp_server.resource_templates', + 'mcp_server.resource_read', + 'mcp_server.logs', + 'mcp_server.test', + ] + ) + assert all(item == {'supported': True} for item in capabilities['operations'].values()) + assert created['key'] not in await capabilities_response.get_data(as_text=True) + bearer_response = await client.get( '/api/v1/system/context', headers={'Authorization': f'Bearer {created["key"]}'}, @@ -491,6 +594,17 @@ async def test_api_key_context_returns_bound_identity_without_workspace_permissi ) assert jwt_response.status_code == 401 + await application.persistence_mgr.execute_async( + sqlalchemy.update(apikey.ApiKey) + .where(apikey.ApiKey.uuid == created['uuid']) + .values(expires_at=datetime.datetime.now(datetime.UTC).replace(tzinfo=None) - datetime.timedelta(seconds=1)) + ) + expired_capabilities = await client.get( + '/api/v1/system/capabilities', + headers={'X-API-Key': created['key']}, + ) + assert expired_capabilities.status_code == 401 + revoke_response = await client.delete( f'/api/v1/apikeys/{created["id"]}', headers=_auth(owner_token, workspace_uuid), @@ -502,6 +616,88 @@ async def test_api_key_context_returns_bound_identity_without_workspace_permissi headers={'X-API-Key': created['key']}, ) assert revoked_response.status_code == 401 + revoked_capabilities = await client.get( + '/api/v1/system/capabilities', + headers={'X-API-Key': created['key']}, + ) + assert revoked_capabilities.status_code == 401 + + +async def test_api_key_can_query_tasks_with_public_contract_and_resource_permission(workspace_api): + application, client, _, owner_token = workspace_api + task_query = {} + task_lookup = {} + fake_task = SimpleNamespace( + to_public_dict=lambda: {'id': 7, 'status': 'running', 'error': None, 'result': None}, + to_dict=lambda: {'id': 7, 'runtime': {'state': 'PENDING'}}, + ) + + def get_tasks_dict(*args, **kwargs): + task_query.update(kwargs) + if kwargs.get('public'): + return {'tasks': []} + return {'tasks': [], 'id_index': 1} + + def get_task_by_id(*args, **kwargs): + task_lookup.update(kwargs) + return fake_task if args and args[0] == 7 else None + + application.task_mgr = SimpleNamespace( + get_tasks_dict=get_tasks_dict, + get_task_by_id=get_task_by_id, + ) + current_response = await client.get('/api/v1/workspaces/current', headers=_auth(owner_token)) + workspace_uuid = (await current_response.get_json())['data']['workspace']['uuid'] + create_response = await client.post( + '/api/v1/apikeys', + headers=_auth(owner_token, workspace_uuid), + json={'name': 'Task reader', 'scopes': ['resource.view']}, + ) + assert create_response.status_code == 200 + key = (await create_response.get_json())['data']['key']['key'] + + listing = await client.get('/api/v1/system/tasks', headers={'X-API-Key': key}) + assert listing.status_code == 200 + assert (await listing.get_json())['data'] == {'tasks': []} + assert task_query['instance_uuid'] == application.workspace_service.instance_uuid + assert task_query['workspace_uuid'] == workspace_uuid + assert task_query['placement_generation'] == 1 + assert task_query['public'] is True + + bearer_listing = await client.get('/api/v1/system/tasks', headers=_auth(owner_token, workspace_uuid)) + assert bearer_listing.status_code == 200 + assert (await bearer_listing.get_json())['data'] == {'tasks': [], 'id_index': 1} + + public_task = await client.get('/api/v1/system/tasks/7', headers={'X-API-Key': key}) + assert public_task.status_code == 200 + assert (await public_task.get_json())['data'] == { + 'id': 7, + 'status': 'running', + 'error': None, + 'result': None, + } + assert task_lookup == { + 'instance_uuid': application.workspace_service.instance_uuid, + 'workspace_uuid': workspace_uuid, + 'placement_generation': 1, + } + + legacy_task = await client.get('/api/v1/system/tasks/7', headers=_auth(owner_token, workspace_uuid)) + assert legacy_task.status_code == 200 + assert (await legacy_task.get_json())['data'] == {'id': 7, 'runtime': {'state': 'PENDING'}} + + missing = await client.get('/api/v1/system/tasks/not-an-id', headers={'X-API-Key': key}) + assert missing.status_code == 404 + + no_permission_response = await client.post( + '/api/v1/apikeys', + headers=_auth(owner_token, workspace_uuid), + json={'name': 'Task denied', 'scopes': []}, + ) + assert no_permission_response.status_code == 200 + no_permission_key = (await no_permission_response.get_json())['data']['key']['key'] + denied = await client.get('/api/v1/system/tasks', headers={'X-API-Key': no_permission_key}) + assert denied.status_code == 403 async def test_cloud_projection_is_selected_explicitly_and_collaboration_runs_in_core( diff --git a/tests/unit_tests/api/test_provider_controller_secrets.py b/tests/unit_tests/api/test_provider_controller_secrets.py new file mode 100644 index 000000000..78c72ba8e --- /dev/null +++ b/tests/unit_tests/api/test_provider_controller_secrets.py @@ -0,0 +1,182 @@ +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.provider.models import ( + EmbeddingModelsRouterGroup, + LLMModelsRouterGroup, + RerankModelsRouterGroup, +) +from langbot.pkg.api.http.controller.groups.provider.providers import ModelProvidersRouterGroup +from langbot.pkg.api.http.controller.groups.provider.query import resolve_include_secret +from langbot.pkg.api.http.service.secrets import redact_secrets + + +pytestmark = pytest.mark.asyncio + +RAW_PROVIDER = { + 'uuid': 'provider-test', + 'name': 'Test Provider', + 'api_keys': ['provider-secret'], +} +RAW_MODEL = { + 'uuid': 'model-test', + 'name': 'Test Model', + 'extra_args': {'headers': {'Authorization': 'Bearer model-secret'}}, +} + + +def _access(role: str): + return 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), + ) + + +def _project(value: dict, include_secret: bool) -> dict: + value = copy.deepcopy(value) + return value if include_secret else redact_secrets(value) + + +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_providers(_context, *, include_secret=False): + return [_project(RAW_PROVIDER, include_secret)] + + async def get_provider(_context, _uuid, *, include_secret=False): + return _project(RAW_PROVIDER, include_secret) + + application.provider_service = SimpleNamespace( + get_providers=AsyncMock(side_effect=get_providers), + get_provider=AsyncMock(side_effect=get_provider), + get_provider_model_counts=AsyncMock( + return_value={'llm_count': 1, 'embedding_count': 1, 'rerank_count': 1} + ), + ) + + def model_service(list_name: str, get_name: str): + async def get_models(_context, *, include_secret=False): + return [_project(RAW_MODEL, include_secret)] + + async def get_model(_context, _uuid, *, include_secret=False): + return _project(RAW_MODEL, include_secret) + + return SimpleNamespace( + **{ + list_name: AsyncMock(side_effect=get_models), + get_name: AsyncMock(side_effect=get_model), + } + ) + + application.llm_model_service = model_service('get_llm_models', 'get_llm_model') + application.embedding_models_service = model_service('get_embedding_models', 'get_embedding_model') + application.rerank_models_service = model_service('get_rerank_models', 'get_rerank_model') + + quart_app = quart.Quart(__name__) + for router_type in ( + ModelProvidersRouterGroup, + LLMModelsRouterGroup, + EmbeddingModelsRouterGroup, + RerankModelsRouterGroup, + ): + await router_type(application, quart_app).initialize() + return application, quart_app.test_client() + + +def _headers() -> dict[str, str]: + return {'Authorization': 'Bearer test-token'} + + +@pytest.mark.parametrize( + ('raw_value', 'permitted', 'expected', 'error'), + [ + (None, True, True, None), + (None, False, False, None), + ('false', True, False, None), + ('true', True, True, None), + ('true', False, False, None), + ('invalid', True, False, 'include_secret must be either true or false'), + ], +) +def test_resolve_include_secret(raw_value, permitted, expected, error): + assert resolve_include_secret(raw_value, permitted=permitted) == (expected, error) + + +@pytest.mark.parametrize( + 'endpoint', + [ + '/api/v1/provider/providers', + '/api/v1/provider/models/llm', + '/api/v1/provider/models/embedding', + '/api/v1/provider/models/rerank', + ], +) +async def test_default_preserves_secrets_and_explicit_false_redacts_high_permission_reads(endpoint): + application, client = await _create_client('developer') + + default_response = await client.get(endpoint, headers=_headers()) + false_response = await client.get(f'{endpoint}?include_secret=false', headers=_headers()) + + assert default_response.status_code == 200 + assert false_response.status_code == 200 + default_data = await default_response.get_json() + false_data = await false_response.get_json() + default_value = default_data['data'].get('providers', default_data['data'].get('models'))[0] + false_value = false_data['data'].get('providers', false_data['data'].get('models'))[0] + assert '***' not in str(default_value) + assert '***' in str(false_value) + + +@pytest.mark.parametrize( + 'endpoint', + [ + '/api/v1/provider/providers', + '/api/v1/provider/models/llm', + '/api/v1/provider/models/embedding', + '/api/v1/provider/models/rerank', + ], +) +async def test_explicit_true_does_not_grant_low_permission_reads(endpoint): + _application, client = await _create_client('viewer') + + response = await client.get(f'{endpoint}?include_secret=true', headers=_headers()) + + assert response.status_code == 200 + data = await response.get_json() + value = data['data'].get('providers', data['data'].get('models'))[0] + assert '***' in str(value) + + +@pytest.mark.parametrize( + 'endpoint', + [ + '/api/v1/provider/providers', + '/api/v1/provider/providers/provider-test', + '/api/v1/provider/models/llm', + '/api/v1/provider/models/llm/model-test', + '/api/v1/provider/models/embedding', + '/api/v1/provider/models/embedding/model-test', + '/api/v1/provider/models/rerank', + '/api/v1/provider/models/rerank/model-test', + ], +) +async def test_invalid_include_secret_returns_bad_request(endpoint): + _application, client = await _create_client('developer') + + response = await client.get(f'{endpoint}?include_secret=maybe', headers=_headers()) + + assert response.status_code == 400 + assert (await response.get_json())['msg'] == 'include_secret must be either true or false' diff --git a/tests/unit_tests/core/test_taskmgr.py b/tests/unit_tests/core/test_taskmgr.py index 44503de9a..3477478bb 100644 --- a/tests/unit_tests/core/test_taskmgr.py +++ b/tests/unit_tests/core/test_taskmgr.py @@ -338,6 +338,70 @@ class TestTaskWrapper: assert result['runtime']['exception'] == 'Test error' assert 'exception_traceback' in result['runtime'] + @pytest.mark.asyncio + async def test_public_dict_has_stable_success_projection(self): + _, TaskWrapper, _ = get_taskmgr_classes() + mock_app = create_mock_app() + + async def successful_coro(): + return {'file_id': 'file-a'} + + wrapper = TaskWrapper(mock_app, successful_coro(), kind='knowledge_base.store') + await wrapper.task + + result = wrapper.to_public_dict() + + assert result == { + 'id': wrapper.id, + 'task_type': 'system', + 'kind': 'knowledge_base.store', + 'status': 'succeeded', + 'error': None, + 'result': {'file_id': 'file-a'}, + 'created_at': result['created_at'], + } + assert 'runtime' not in result + assert 'traceback' not in str(result).lower() + + @pytest.mark.asyncio + async def test_public_dict_hides_exception_traceback(self): + _, TaskWrapper, _ = get_taskmgr_classes() + mock_app = create_mock_app() + + async def failing_coro(): + raise ValueError('private failure') + + wrapper = TaskWrapper(mock_app, failing_coro()) + try: + await wrapper.task + except ValueError: + # Expected failure: task must complete in failed state for public serialization checks. + pass + + result = wrapper.to_public_dict() + + assert result['status'] == 'failed' + assert result['error'] == {'type': 'task_failed', 'message': 'Task execution failed'} + assert 'runtime' not in result + assert 'traceback' not in str(result).lower() + + @pytest.mark.asyncio + async def test_public_dict_does_not_change_success_when_result_is_not_json_serializable(self): + _, TaskWrapper, _ = get_taskmgr_classes() + mock_app = create_mock_app() + + async def successful_coro(): + return object() + + wrapper = TaskWrapper(mock_app, successful_coro()) + await wrapper.task + + result = wrapper.to_public_dict() + + assert result['status'] == 'succeeded' + assert result['error'] is None + assert result['result'] is None + @pytest.mark.asyncio async def test_cancel_task(self): """Test cancel method cancels the asyncio task.""" @@ -487,6 +551,66 @@ class TestAsyncTaskManager: w2.cancel() w3.cancel() + @pytest.mark.asyncio + async def test_public_task_queries_keep_workspace_and_generation_isolation(self): + _, _, AsyncTaskManager = get_taskmgr_classes() + mock_app = create_mock_app() + manager = AsyncTaskManager(mock_app) + + async def dummy_coro(): + await asyncio.sleep(10) + + current = manager.create_user_task( + dummy_coro(), + instance_uuid='instance-a', + workspace_uuid='workspace-a', + placement_generation=2, + ) + other_workspace = manager.create_user_task( + dummy_coro(), + instance_uuid='instance-a', + workspace_uuid='workspace-b', + placement_generation=2, + ) + stale_generation = manager.create_user_task( + dummy_coro(), + instance_uuid='instance-a', + workspace_uuid='workspace-a', + placement_generation=1, + ) + + result = manager.get_tasks_dict( + instance_uuid='instance-a', + workspace_uuid='workspace-a', + placement_generation=2, + public=True, + ) + + assert [task['id'] for task in result['tasks']] == [current.id] + assert 'id_index' not in result + assert ( + manager.get_task_by_id( + other_workspace.id, + instance_uuid='instance-a', + workspace_uuid='workspace-a', + placement_generation=2, + ) + is None + ) + assert ( + manager.get_task_by_id( + stale_generation.id, + instance_uuid='instance-a', + workspace_uuid='workspace-a', + placement_generation=2, + ) + is None + ) + + current.cancel() + other_workspace.cancel() + stale_generation.cancel() + @pytest.mark.asyncio async def test_cancel_by_scope(self): """Test cancel_by_scope cancels matching tasks."""