From c6f826fe2d2e12cc73479155e8f25566048c1c56 Mon Sep 17 00:00:00 2001
From: Junyan Qin
Date: Sun, 19 Jul 2026 09:58:59 +0800
Subject: [PATCH] feat(tenancy): implement workspace isolation
---
ARCHITECTURE.md | 6 +
docker/docker-compose.yaml | 20 +-
docker/kubernetes.yaml | 30 +-
docs/API_KEY_AUTH.md | 45 +-
docs/multi-tenant/implementation-checklist.md | 246 ++++
docs/multi-tenant/implementation-decisions.md | 175 +++
pyproject.toml | 2 +-
skills/skills/langbot-deploy/SKILL.md | 11 +
skills/skills/langbot-dev/SKILL.md | 21 +-
skills/skills/langbot-mcp-ops/SKILL.md | 23 +-
src/langbot/pkg/api/http/authz.py | 116 ++
src/langbot/pkg/api/http/context.py | 92 ++
src/langbot/pkg/api/http/controller/group.py | 296 ++++-
.../pkg/api/http/controller/groups/apikeys.py | 85 +-
.../pkg/api/http/controller/groups/box.py | 35 +-
.../api/http/controller/groups/extensions.py | 20 +-
.../pkg/api/http/controller/groups/files.py | 66 +-
.../http/controller/groups/knowledge/base.py | 175 ++-
.../controller/groups/knowledge/engines.py | 38 +-
.../controller/groups/knowledge/migration.py | 144 ++-
.../controller/groups/knowledge/parsers.py | 14 +-
.../pkg/api/http/controller/groups/logs.py | 13 +-
.../api/http/controller/groups/monitoring.py | 85 +-
.../http/controller/groups/pipelines/embed.py | 141 ++-
.../controller/groups/pipelines/pipelines.py | 228 ++--
.../groups/pipelines/websocket_chat.py | 405 ++++---
.../controller/groups/platform/adapters.py | 209 +++-
.../http/controller/groups/platform/bots.py | 178 ++-
.../pkg/api/http/controller/groups/plugins.py | 493 ++++++--
.../http/controller/groups/provider/models.py | 371 ++++--
.../controller/groups/provider/providers.py | 120 +-
.../http/controller/groups/resources/mcp.py | 193 +--
.../http/controller/groups/resources/tools.py | 25 +-
.../pkg/api/http/controller/groups/skills.py | 185 ++-
.../pkg/api/http/controller/groups/stats.py | 46 +-
.../pkg/api/http/controller/groups/system.py | 157 ++-
.../pkg/api/http/controller/groups/user.py | 128 +-
.../http/controller/groups/webhook_mgmt.py | 101 +-
.../api/http/controller/groups/webhooks.py | 14 +-
.../api/http/controller/groups/workspaces.py | 305 +++++
src/langbot/pkg/api/http/service/apikey.py | 283 ++++-
src/langbot/pkg/api/http/service/bot.py | 152 ++-
src/langbot/pkg/api/http/service/knowledge.py | 186 ++-
.../pkg/api/http/service/maintenance.py | 224 ++--
src/langbot/pkg/api/http/service/mcp.py | 499 +++++---
src/langbot/pkg/api/http/service/model.py | 599 +++++++---
.../pkg/api/http/service/monitoring.py | 217 +++-
src/langbot/pkg/api/http/service/pipeline.py | 175 ++-
src/langbot/pkg/api/http/service/provider.py | 213 +++-
src/langbot/pkg/api/http/service/secrets.py | 336 ++++++
src/langbot/pkg/api/http/service/skill.py | 155 ++-
src/langbot/pkg/api/http/service/space.py | 8 +-
src/langbot/pkg/api/http/service/tenant.py | 34 +
src/langbot/pkg/api/http/service/user.py | 431 ++++++-
src/langbot/pkg/api/http/service/webhook.py | 106 +-
src/langbot/pkg/api/mcp/context.py | 30 +
src/langbot/pkg/api/mcp/mount.py | 33 +-
src/langbot/pkg/api/mcp/server.py | 70 +-
src/langbot/pkg/box/connector.py | 62 +-
src/langbot/pkg/box/service.py | 387 ++++--
src/langbot/pkg/box/workspace.py | 28 +-
src/langbot/pkg/command/cmdmgr.py | 5 +
src/langbot/pkg/core/app.py | 67 +-
src/langbot/pkg/core/stages/build_app.py | 47 +-
src/langbot/pkg/core/taskmgr.py | 76 +-
src/langbot/pkg/entity/persistence/apikey.py | 46 +-
src/langbot/pkg/entity/persistence/bot.py | 31 +-
.../pkg/entity/persistence/bstorage.py | 14 +
src/langbot/pkg/entity/persistence/mcp.py | 10 +
.../pkg/entity/persistence/metadata.py | 14 +
src/langbot/pkg/entity/persistence/model.py | 56 +
.../pkg/entity/persistence/monitoring.py | 85 +-
.../pkg/entity/persistence/pipeline.py | 39 +
src/langbot/pkg/entity/persistence/plugin.py | 7 +
src/langbot/pkg/entity/persistence/rag.py | 66 +-
src/langbot/pkg/entity/persistence/user.py | 51 +
src/langbot/pkg/entity/persistence/webhook.py | 7 +
.../pkg/entity/persistence/workspace.py | 271 +++++
.../versions/0009_workspace_tenancy_kernel.py | 542 +++++++++
.../versions/0010_scope_tenant_resources.py | 884 ++++++++++++++
src/langbot/pkg/persistence/alembic_runner.py | 21 +
src/langbot/pkg/persistence/mgr.py | 207 +++-
.../persistence/sqlite_migration_backup.py | 272 +++++
src/langbot/pkg/pipeline/aggregator.py | 310 ++---
src/langbot/pkg/pipeline/controller.py | 80 +-
src/langbot/pkg/pipeline/monitoring_helper.py | 10 +
src/langbot/pkg/pipeline/pipelinemgr.py | 162 ++-
src/langbot/pkg/pipeline/pool.py | 267 ++++-
src/langbot/pkg/pipeline/preproc/preproc.py | 26 +-
.../pkg/pipeline/process/handlers/chat.py | 6 +-
src/langbot/pkg/platform/botmgr.py | 305 ++++-
src/langbot/pkg/platform/logger.py | 25 +-
src/langbot/pkg/platform/sources/http_bot.py | 34 +-
.../pkg/platform/sources/openclaw_weixin.py | 16 +
.../pkg/platform/sources/websocket_adapter.py | 110 +-
.../pkg/platform/sources/websocket_manager.py | 182 ++-
src/langbot/pkg/platform/webhook_pusher.py | 21 +-
src/langbot/pkg/plugin/connector.py | 199 +++-
src/langbot/pkg/plugin/handler.py | 742 ++++++++++--
src/langbot/pkg/provider/modelmgr/modelmgr.py | 727 +++++++-----
.../pkg/provider/modelmgr/requester.py | 88 +-
src/langbot/pkg/provider/runners/difysvapi.py | 19 +-
.../pkg/provider/runners/localagent.py | 21 +-
.../pkg/provider/session/sessionmgr.py | 79 +-
.../provider/tools/loaders/availability.py | 2 +-
src/langbot/pkg/provider/tools/loaders/mcp.py | 362 +++++-
.../pkg/provider/tools/loaders/mcp_stdio.py | 75 +-
.../pkg/provider/tools/loaders/native.py | 61 +-
.../pkg/provider/tools/loaders/plugin.py | 6 +-
.../pkg/provider/tools/loaders/skill.py | 11 +-
.../provider/tools/loaders/skill_authoring.py | 18 +-
src/langbot/pkg/provider/tools/toolmgr.py | 15 +-
src/langbot/pkg/rag/knowledge/base.py | 11 +-
src/langbot/pkg/rag/knowledge/kbmgr.py | 393 +++++--
src/langbot/pkg/rag/service/runtime.py | 113 +-
src/langbot/pkg/skill/activation.py | 11 +-
src/langbot/pkg/skill/manager.py | 164 ++-
src/langbot/pkg/storage/mgr.py | 301 +++++
src/langbot/pkg/telemetry/heartbeat.py | 4 +-
src/langbot/pkg/utils/httpclient.py | 8 +-
src/langbot/pkg/utils/managed_runtime.py | 8 +-
src/langbot/pkg/vector/mgr.py | 178 ++-
src/langbot/pkg/workspace/__init__.py | 25 +
src/langbot/pkg/workspace/collaboration.py | 684 +++++++++++
src/langbot/pkg/workspace/entities.py | 14 +
src/langbot/pkg/workspace/errors.py | 31 +
src/langbot/pkg/workspace/policy.py | 50 +
src/langbot/pkg/workspace/repository.py | 84 ++
src/langbot/pkg/workspace/service.py | 417 +++++++
src/langbot/templates/config.yaml | 8 +
src/langbot/templates/embed/widget.js | 9 +-
tests/factories/app.py | 2 +
tests/factories/message.py | 16 +-
tests/integration/api/test_bots.py | 53 +-
tests/integration/api/test_box_security.py | 83 ++
tests/integration/api/test_embed.py | 116 +-
tests/integration/api/test_knowledge.py | 24 +-
tests/integration/api/test_monitoring.py | 22 +-
tests/integration/api/test_pipelines.py | 174 ++-
.../integration/api/test_plugins_security.py | 248 ++++
tests/integration/api/test_providers.py | 18 +-
tests/integration/api/test_smoke.py | 18 +
.../integration/api/test_user_space_oauth.py | 234 ++++
tests/integration/api/test_workspaces.py | 444 +++++++
.../persistence/resource_migration_support.py | 253 ++++
.../persistence/test_migrations_postgres.py | 202 ++++
.../test_resource_tenancy_migration.py | 286 +++++
.../test_sqlite_migration_backup.py | 107 ++
.../persistence/test_workspace_migration.py | 379 ++++++
tests/integration/pipeline/test_full_flow.py | 10 +-
.../box/test_box_integration.py | 52 +-
.../box/test_box_mcp_integration.py | 122 +-
tests/manual/mcp_smoke.py | 18 +-
.../api/http/service/test_bot_service.py | 9 +-
.../api/http/service/test_tenant.py | 34 +
.../service/test_tenant_resource_isolation.py | 373 ++++++
tests/unit_tests/api/http/test_authz.py | 74 ++
.../api/http/test_internal_error_responses.py | 138 +++
.../api/service/test_apikey_service.py | 757 ++++++------
.../api/service/test_bot_service.py | 69 +-
.../api/service/test_knowledge_service.py | 666 +++++++----
.../api/service/test_maintenance_service.py | 236 +++-
.../api/service/test_mcp_service.py | 386 ++++--
.../api/service/test_model_service.py | 261 ++++-
.../api/service/test_monitoring_tenancy.py | 154 +++
.../api/service/test_pipeline_service.py | 137 ++-
.../api/service/test_provider_service.py | 151 ++-
.../api/service/test_secret_redaction.py | 88 ++
.../api/service/test_space_service.py | 9 +-
.../api/service/test_user_service.py | 127 +-
.../api/service/test_webhook_service.py | 251 +++-
.../api/test_adapter_session_scoping.py | 195 +++
tests/unit_tests/api/test_apikey_service.py | 62 +-
.../api/test_bot_controller_secrets.py | 113 ++
.../api/test_extensions_runtime_fence.py | 94 ++
.../api/test_file_upload_scoping.py | 59 +
.../test_knowledge_migration_runtime_fence.py | 69 ++
tests/unit_tests/api/test_mcp_controller.py | 70 +-
.../api/test_plugin_runtime_route_fence.py | 108 ++
.../api/test_resource_secret_permissions.py | 196 ++++
tests/unit_tests/api/test_stats_controller.py | 110 ++
tests/unit_tests/box/test_box_connector.py | 141 +++
.../box/test_box_deployment_config.py | 65 +
tests/unit_tests/box/test_box_service.py | 292 ++++-
tests/unit_tests/box/test_workspace.py | 46 +-
tests/unit_tests/pipeline/conftest.py | 24 +
tests/unit_tests/pipeline/test_aggregator.py | 361 +++++-
.../pipeline/test_chat_session_limit.py | 8 +-
.../pipeline/test_controller_tenancy.py | 65 +
.../pipeline/test_pipeline_service.py | 10 +-
tests/unit_tests/pipeline/test_pipelinemgr.py | 122 +-
tests/unit_tests/pipeline/test_pool.py | 173 ++-
tests/unit_tests/pipeline/test_preproc.py | 61 +-
tests/unit_tests/pipeline/test_query_pool.py | 9 +-
.../platform/test_botmgr_tenancy.py | 114 ++
.../platform/test_http_bot_tenancy.py | 66 ++
.../platform/test_openclaw_weixin_tenancy.py | 84 ++
.../unit_tests/platform/test_routing_rules.py | 16 +
.../test_websocket_adapter_attachments.py | 73 +-
.../test_websocket_session_isolation.py | 141 ++-
.../plugin/test_connector_methods.py | 13 +
.../unit_tests/plugin/test_connector_ping.py | 196 ++++
tests/unit_tests/plugin/test_handler.py | 45 +-
.../unit_tests/plugin/test_handler_actions.py | 243 +++-
.../unit_tests/plugin/test_handler_tenancy.py | 303 +++++
tests/unit_tests/provider/conftest.py | 43 +
.../provider/runners/test_difysvapi_runner.py | 25 +-
.../provider/test_localagent_sandbox_exec.py | 17 +-
.../provider/test_mcp_box_integration.py | 76 +-
.../provider/test_mcp_remote_transport.py | 9 +
.../unit_tests/provider/test_mcp_resources.py | 162 ++-
.../unit_tests/provider/test_model_manager.py | 224 +++-
.../unit_tests/provider/test_model_service.py | 104 +-
.../provider/test_requester_base.py | 30 +-
.../provider/test_session_manager.py | 147 ++-
tests/unit_tests/provider/test_skill_tools.py | 118 +-
.../unit_tests/provider/test_tool_manager.py | 13 +
.../provider/test_tool_manager_native.py | 41 +-
tests/unit_tests/rag/test_file_storage.py | 145 ++-
tests/unit_tests/rag/test_kbmgr.py | 1043 +++++++++--------
tests/unit_tests/rag/test_runtime_service.py | 789 ++++++-------
tests/unit_tests/rag/test_tenant_isolation.py | 350 ++++++
.../storage/test_workspace_scoping.py | 235 ++++
tests/unit_tests/telemetry/test_heartbeat.py | 1 +
tests/unit_tests/test_preproc.py | 30 +-
tests/unit_tests/test_skill_service.py | 48 +-
tests/unit_tests/utils/test_httpclient.py | 10 +
tests/unit_tests/workspace/__init__.py | 0
.../workspace/test_workspace_collaboration.py | 310 +++++
.../workspace/test_workspace_policy.py | 19 +
.../workspace/test_workspace_service.py | 237 ++++
tests/utils/import_isolation.py | 41 +-
uv.lock | 10 +-
web/src/app/auth/space/callback/page.tsx | 55 +-
web/src/app/home/bots/BotDetailContent.tsx | 163 +--
.../home/bots/components/bot-form/BotForm.tsx | 511 ++++----
.../AccountSettingsPanel.tsx | 12 +-
.../ApiIntegrationPanel.tsx | 38 +-
.../dynamic-form/DynamicFormItemComponent.tsx | 7 +-
.../components/home-sidebar/HomeSidebar.tsx | 80 +-
.../components/models-dialog/ModelsPanel.tsx | 25 +-
.../models-dialog/components/ModelItem.tsx | 6 +-
.../models-dialog/components/ProviderCard.tsx | 11 +-
.../settings-dialog/SettingsDialog.tsx | 48 +-
.../WorkspaceSettingsPanel.tsx | 398 +++++++
.../workspace-settings/WorkspaceSwitcher.tsx | 56 +
.../app/home/knowledge/KBDetailContent.tsx | 130 +-
web/src/app/home/layout.tsx | 103 +-
web/src/app/home/mcp/MCPDetailContent.tsx | 129 +-
.../home/pipelines/PipelineDetailContent.tsx | 154 ++-
.../app/home/skills/SkillDetailContent.tsx | 82 +-
web/src/app/home/skills/page.tsx | 26 +-
web/src/app/infra/entities/workspace.ts | 51 +
web/src/app/infra/http/BackendClient.ts | 225 +++-
web/src/app/infra/http/BaseHttpClient.ts | 31 +
.../app/infra/http/currentWorkspaceStore.ts | 30 +
web/src/app/infra/http/index.ts | 229 +++-
.../app/infra/http/workspaceBootstrapStore.ts | 35 +
web/src/app/infra/http/workspaceContext.ts | 54 +
.../app/infra/websocket/WebSocketClient.ts | 20 +-
web/src/app/invitations/accept/page.tsx | 339 ++++++
web/src/app/login/page.tsx | 69 +-
web/src/app/wizard/page.tsx | 16 +-
web/src/app/workspaces/select/page.tsx | 191 +++
web/src/i18n/locales/en-US.ts | 73 ++
web/src/i18n/locales/ja-JP.ts | 72 ++
web/src/i18n/locales/zh-Hans.ts | 69 ++
web/src/router.tsx | 10 +
web/tests/e2e/fixtures/langbot-api.ts | 118 +-
web/tests/e2e/invitations.spec.ts | 55 +
web/tests/e2e/login.spec.ts | 127 +-
271 files changed, 31162 insertions(+), 6106 deletions(-)
create mode 100644 docs/multi-tenant/implementation-checklist.md
create mode 100644 docs/multi-tenant/implementation-decisions.md
create mode 100644 src/langbot/pkg/api/http/authz.py
create mode 100644 src/langbot/pkg/api/http/context.py
create mode 100644 src/langbot/pkg/api/http/controller/groups/workspaces.py
create mode 100644 src/langbot/pkg/api/http/service/secrets.py
create mode 100644 src/langbot/pkg/api/http/service/tenant.py
create mode 100644 src/langbot/pkg/api/mcp/context.py
create mode 100644 src/langbot/pkg/entity/persistence/workspace.py
create mode 100644 src/langbot/pkg/persistence/alembic/versions/0009_workspace_tenancy_kernel.py
create mode 100644 src/langbot/pkg/persistence/alembic/versions/0010_scope_tenant_resources.py
create mode 100644 src/langbot/pkg/persistence/sqlite_migration_backup.py
create mode 100644 src/langbot/pkg/workspace/__init__.py
create mode 100644 src/langbot/pkg/workspace/collaboration.py
create mode 100644 src/langbot/pkg/workspace/entities.py
create mode 100644 src/langbot/pkg/workspace/errors.py
create mode 100644 src/langbot/pkg/workspace/policy.py
create mode 100644 src/langbot/pkg/workspace/repository.py
create mode 100644 src/langbot/pkg/workspace/service.py
create mode 100644 tests/integration/api/test_box_security.py
create mode 100644 tests/integration/api/test_plugins_security.py
create mode 100644 tests/integration/api/test_user_space_oauth.py
create mode 100644 tests/integration/api/test_workspaces.py
create mode 100644 tests/integration/persistence/resource_migration_support.py
create mode 100644 tests/integration/persistence/test_resource_tenancy_migration.py
create mode 100644 tests/integration/persistence/test_sqlite_migration_backup.py
create mode 100644 tests/integration/persistence/test_workspace_migration.py
create mode 100644 tests/unit_tests/api/http/service/test_tenant.py
create mode 100644 tests/unit_tests/api/http/service/test_tenant_resource_isolation.py
create mode 100644 tests/unit_tests/api/http/test_authz.py
create mode 100644 tests/unit_tests/api/http/test_internal_error_responses.py
create mode 100644 tests/unit_tests/api/service/test_monitoring_tenancy.py
create mode 100644 tests/unit_tests/api/service/test_secret_redaction.py
create mode 100644 tests/unit_tests/api/test_adapter_session_scoping.py
create mode 100644 tests/unit_tests/api/test_bot_controller_secrets.py
create mode 100644 tests/unit_tests/api/test_extensions_runtime_fence.py
create mode 100644 tests/unit_tests/api/test_file_upload_scoping.py
create mode 100644 tests/unit_tests/api/test_knowledge_migration_runtime_fence.py
create mode 100644 tests/unit_tests/api/test_plugin_runtime_route_fence.py
create mode 100644 tests/unit_tests/api/test_resource_secret_permissions.py
create mode 100644 tests/unit_tests/api/test_stats_controller.py
create mode 100644 tests/unit_tests/box/test_box_deployment_config.py
create mode 100644 tests/unit_tests/pipeline/test_controller_tenancy.py
create mode 100644 tests/unit_tests/platform/test_botmgr_tenancy.py
create mode 100644 tests/unit_tests/platform/test_http_bot_tenancy.py
create mode 100644 tests/unit_tests/platform/test_openclaw_weixin_tenancy.py
create mode 100644 tests/unit_tests/plugin/test_handler_tenancy.py
create mode 100644 tests/unit_tests/rag/test_tenant_isolation.py
create mode 100644 tests/unit_tests/storage/test_workspace_scoping.py
create mode 100644 tests/unit_tests/workspace/__init__.py
create mode 100644 tests/unit_tests/workspace/test_workspace_collaboration.py
create mode 100644 tests/unit_tests/workspace/test_workspace_policy.py
create mode 100644 tests/unit_tests/workspace/test_workspace_service.py
create mode 100644 web/src/app/home/components/workspace-settings/WorkspaceSettingsPanel.tsx
create mode 100644 web/src/app/home/components/workspace-settings/WorkspaceSwitcher.tsx
create mode 100644 web/src/app/infra/entities/workspace.ts
create mode 100644 web/src/app/infra/http/currentWorkspaceStore.ts
create mode 100644 web/src/app/infra/http/workspaceBootstrapStore.ts
create mode 100644 web/src/app/infra/http/workspaceContext.ts
create mode 100644 web/src/app/invitations/accept/page.tsx
create mode 100644 web/src/app/workspaces/select/page.tsx
create mode 100644 web/tests/e2e/invitations.spec.ts
diff --git a/ARCHITECTURE.md b/ARCHITECTURE.md
index ded90e7ea..2b5defb82 100644
--- a/ARCHITECTURE.md
+++ b/ARCHITECTURE.md
@@ -178,6 +178,12 @@ In this repo:
- `pkg/provider/tools/loaders/native.py`, `mcp_stdio.py`, and skill loaders depend on Box availability.
- `pkg/skill/manager.py` loads skills from the Box runtime, falling back to local `data/skills` when needed.
+Durable Box Workspace storage is shared across placement generations, but
+sandbox sessions and managed processes are generation-scoped. LangBot validates
+the current execution binding before an MCP stdio relay attach and sends the
+Workspace/generation binding in authenticated headers, so a placement cutover
+retires stale processes and closes already-attached relays.
+
In `langbot-plugin-sdk`:
- `src/langbot_plugin/box/server.py` implements `lbp box` and the WebSocket endpoints on `:5410`.
diff --git a/docker/docker-compose.yaml b/docker/docker-compose.yaml
index bdd347021..dd2444c49 100644
--- a/docker/docker-compose.yaml
+++ b/docker/docker-compose.yaml
@@ -14,6 +14,9 @@ services:
restart: on-failure
environment:
- TZ=Asia/Shanghai
+ # Shared with the langbot service and sent only as a WebSocket handshake
+ # header. Generate with: openssl rand -hex 32
+ - LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN=${LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN:-}
command: ["uv", "run", "--no-sync", "-m", "langbot_plugin.cli.__init__", "rt"]
networks:
- langbot_network
@@ -40,9 +43,16 @@ services:
restart: on-failure
environment:
- TZ=Asia/Shanghai
+ # Shared control-plane secret used to authenticate both the RPC socket
+ # and managed-process relay. Generate once (for example with
+ # ``openssl rand -hex 32``) and export it before enabling this profile.
+ # An empty value is accepted by Compose so Box can remain optional, but
+ # the Box runtime itself fails closed when the profile is started.
+ - LANGBOT_BOX_CONTROL_TOKEN=${LANGBOT_BOX_CONTROL_TOKEN:-}
# The Box runtime does NOT read box.local.* from config.yaml or env; it
- # receives its configuration from LangBot via the INIT RPC action.
- # Do not add LANGBOT_BOX_* / BOX__* here — they would be silently ignored.
+ # receives its functional configuration from LangBot via the INIT RPC
+ # action. LANGBOT_BOX_CONTROL_TOKEN is the intentional security-only
+ # exception; do not add BOX__* here because those would be ignored.
# Launched through the same CLI entry point as the plugin runtime
# (`langbot_plugin.cli.__init__ `). WebSocket is the default
# control transport — mirrors `rt`, which also runs with no flag. Pass
@@ -60,6 +70,12 @@ services:
restart: on-failure
environment:
- TZ=Asia/Shanghai
+ # Must match langbot_plugin_runtime. Empty/missing values make the
+ # external control channel fail closed.
+ - LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN=${LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN:-}
+ # Must match the value supplied to langbot_box. The token is sent only
+ # in WebSocket handshake headers, never in URLs or action payloads.
+ - LANGBOT_BOX_CONTROL_TOKEN=${LANGBOT_BOX_CONTROL_TOKEN:-}
# Unified env-override convention: SECTION__SUBSECTION__KEY overrides the
# matching config.yaml field (see LoadConfigStage). These map onto
# box.* and are forwarded to the Box runtime via INIT RPC.
diff --git a/docker/kubernetes.yaml b/docker/kubernetes.yaml
index 6adc50510..b142468ba 100644
--- a/docker/kubernetes.yaml
+++ b/docker/kubernetes.yaml
@@ -4,6 +4,10 @@
# Full deployment guide (zh/en/ja): https://docs.langbot.app -> Installation -> Kubernetes
#
# Usage:
+# kubectl -n langbot create secret generic langbot-plugin-runtime-control \
+# --from-literal=token="$(openssl rand -hex 32)"
+# kubectl -n langbot create secret generic langbot-box-control \
+# --from-literal=token="$(openssl rand -hex 32)"
# kubectl apply -f kubernetes.yaml
#
# Prerequisites:
@@ -127,6 +131,11 @@ spec:
configMapKeyRef:
name: langbot-config
key: TZ
+ - name: LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN
+ valueFrom:
+ secretKeyRef:
+ name: langbot-plugin-runtime-control
+ key: token
volumeMounts:
- name: plugin-data
mountPath: /app/data/plugins
@@ -246,9 +255,14 @@ spec:
configMapKeyRef:
name: langbot-config
key: TZ
+ - name: LANGBOT_BOX_CONTROL_TOKEN
+ valueFrom:
+ secretKeyRef:
+ name: langbot-box-control
+ key: token
# The Box runtime does NOT read box.local.* / BOX__* from its own env;
- # it receives its configuration from LangBot via the INIT RPC action.
- # Do not add BOX__* here — they would be silently ignored.
+ # it receives its functional configuration from LangBot via INIT.
+ # LANGBOT_BOX_CONTROL_TOKEN is the security-only exception.
volumeMounts:
# Box workspace root — identical path on node, box, and sandbox
# containers (see the IMPORTANT note above).
@@ -352,6 +366,11 @@ spec:
configMapKeyRef:
name: langbot-config
key: PLUGIN__RUNTIME_WS_URL
+ - name: LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN
+ valueFrom:
+ secretKeyRef:
+ name: langbot-plugin-runtime-control
+ key: token
# Box (sandbox) runtime endpoint. Connects LangBot to the langbot-box
# Service over WebSocket. Remove this (and the langbot-box Deployment)
# and set BOX__ENABLED=false if you do not want the sandbox.
@@ -360,6 +379,13 @@ spec:
configMapKeyRef:
name: langbot-config
key: BOX__RUNTIME__ENDPOINT
+ # Same Secret as langbot-box. It authenticates the RPC and managed-
+ # process relay handshakes and is never put in a URL or RPC payload.
+ - name: LANGBOT_BOX_CONTROL_TOKEN
+ valueFrom:
+ secretKeyRef:
+ name: langbot-box-control
+ key: token
# box.local.* config — forwarded to the Box runtime via INIT RPC. The
# host_root MUST match the box-root hostPath mountPath below AND the box
# Deployment's box-root mountPath, so that skill package paths resolve
diff --git a/docs/API_KEY_AUTH.md b/docs/API_KEY_AUTH.md
index ad3bb6693..49d80b6f9 100644
--- a/docs/API_KEY_AUTH.md
+++ b/docs/API_KEY_AUTH.md
@@ -8,13 +8,21 @@ API keys can be managed through the web interface:
1. Log in to the LangBot web interface
2. Click the "API Keys" button at the bottom of the sidebar
-3. Create, view, copy, or delete API keys as needed
+3. Create an API key and copy its secret immediately
+4. Revoke keys that are no longer needed
+
+Database-backed API-key secrets are returned exactly once. LangBot stores only
+a SHA-256 lookup hash, so an existing secret cannot be displayed or recovered
+later. Each key belongs to one Workspace, has explicit permission scopes, and
+may have an expiry. The Workspace is derived from the authenticated key; an
+`X-Workspace-Id` header cannot redirect it to another tenant.
## Global API Key (config.yaml)
In addition to web-UI-created keys (stored in the database, prefixed `lbk_`),
LangBot supports a **global API key** defined directly in `data/config.yaml`.
-This is useful for automated deployments, infrastructure-as-code, and AI agents
+This is a Community-edition bootstrap option for automated deployments,
+infrastructure-as-code, and AI agents
that need API/MCP access **without a login session and without creating a
database record first**.
@@ -27,10 +35,12 @@ api:
Behavior:
-- When `api.global_api_key` is a non-empty string, that exact value is accepted
- anywhere a normal API key is accepted — the `X-API-Key` header or
- `Authorization: Bearer ` — across the HTTP service API **and the MCP
- server**.
+- In Community edition's singleton Workspace, a non-empty
+ `api.global_api_key` is bound to that Workspace and accepted across the HTTP
+ service API and the MCP server.
+- The global config key is rejected when multi-Workspace SaaS mode is enabled;
+ SaaS automation must use a database-backed Workspace key or a closed control
+ plane credential.
- The global key does **not** require the `lbk_` prefix; use any sufficiently
strong secret.
- Leave it empty (`''`, the default) to disable it entirely; only database-backed
@@ -38,9 +48,10 @@ Behavior:
- Existing installs are unaffected until you add the key — config completion only
backfills top-level keys, and the lookup is defensive when the field is absent.
-> **Security:** the global key is stored in plaintext in `config.yaml`. Only
-> enable it on trusted/internal deployments, keep the file permissions tight,
-> always serve over HTTPS, and rotate the value if it may have leaked.
+> **Security:** the global key is stored in plaintext in `config.yaml` and has
+> the singleton Workspace's full fixed permission set. Only enable it on
+> trusted/internal Community deployments, keep file permissions tight, always
+> serve over HTTPS, and rotate it if it may have leaked.
## Using API Keys
@@ -60,7 +71,9 @@ Authorization: Bearer lbk_your_api_key_here
## Available APIs
-All existing LangBot APIs now support **both user token and API key authentication**. This means you can use API keys to access:
+Endpoints that declare API-key authentication accept either a user token or a
+Workspace API key. The key must include the permission required by the route.
+This includes:
- **Model Management** - `/api/v1/provider/models/llm` and `/api/v1/provider/models/embedding`
- **Bot Management** - `/api/v1/platform/bots`
@@ -227,6 +240,11 @@ or
}
```
+### 403 Forbidden
+
+The key is valid for its Workspace but does not include the fixed permission
+required by the route.
+
### 500 Internal Server Error
```json
@@ -240,7 +258,7 @@ or
1. **Keep API keys secure**: Store them securely and never commit them to version control
2. **Use HTTPS**: Always use HTTPS in production to encrypt API key transmission
-3. **Rotate keys regularly**: Create new API keys periodically and delete old ones
+3. **Rotate keys regularly**: Create new API keys periodically and revoke old ones
4. **Use descriptive names**: Give your API keys meaningful names to track their usage
5. **Delete unused keys**: Remove API keys that are no longer needed
6. **Use X-API-Key header**: Prefer using the `X-API-Key` header for clarity
@@ -317,7 +335,6 @@ curl -X POST \
## Notes
-- The same endpoints work for both the web UI (with user tokens) and external services (with API keys)
+- API-key-enabled endpoints use the same resource shapes as the web UI
- No need to learn different API paths - use the existing API documentation with API key authentication
-- All endpoints that previously required user authentication now also accept API keys
-
+- API keys never select a Workspace from a request header; their persisted binding is authoritative
diff --git a/docs/multi-tenant/implementation-checklist.md b/docs/multi-tenant/implementation-checklist.md
new file mode 100644
index 000000000..d73fb022b
--- /dev/null
+++ b/docs/multi-tenant/implementation-checklist.md
@@ -0,0 +1,246 @@
+# Multi-tenant implementation checklist
+
+This checklist turns the Workspace architecture into implementation and verification gates.
+
+## Scope guard
+
+- [x] LangBot uses branch feat/multi-tenants.
+- [x] langbot-plugin-sdk uses branch feat/multi-tenants.
+- [x] langbot-space has no changes made by this implementation; Cloud v2 does not extend the legacy Space deployment scheme.
+- [x] Unrelated untracked files in either repository remain untouched.
+- [x] Open-source startup cannot enable SaaS multi-workspace through edition flags or unsigned configuration.
+
+## SaaS-only release gates
+
+These items intentionally remain incomplete. The feature branch delivers the Core isolation kernel, not the closed SaaS product or a production Cloud v2 deployment.
+
+- [ ] The closed Control Plane owns the global Account, Workspace, Membership, and Invitation directory.
+- [ ] The closed placement service issues monotonic generations and leases for projected Workspaces.
+- [ ] Core verifies a signed `InstanceManifest` before the closed bootstrap can inject `CloudWorkspacePolicy`.
+- [ ] Tenant database writes hold a generation-aware shared transaction fence through commit, while placement cutovers take the exclusive fence.
+- [ ] Business writes and non-transactional side effects use a generation-stamped outbox or equivalent publish fence.
+- [ ] Durable object references survive a placement-generation change through stable published keys or an explicitly atomic key/reference migration.
+- [ ] SaaS cells enforce tenant-safe egress and SSRF controls for Webhooks, providers, MCP servers, and every tenant-configurable outbound URL.
+- [ ] Entitlement checks, usage aggregation, subscription lifecycle, and billing are implemented in the closed Control Plane.
+- [ ] OAuth state and directory projection use an atomic shared store suitable for horizontally scaled SaaS services.
+- [ ] A greenfield Cloud v2 deployment is designed and validated independently of the legacy Space deployment scheme.
+- [ ] Multi-workspace is enabled in SaaS only after all closed Control Plane, deployment, and security gates pass.
+
+## 1. Persistence foundation
+
+### Account and directory
+
+- [x] User has a stable, unique account UUID and explicit status.
+- [x] Existing email and password behavior remains compatible during migration.
+- [x] Workspace table represents the instance-local tenant.
+- [x] WorkspaceMembership has a unique Workspace and Account pair.
+- [x] WorkspaceInvitation stores only a token hash and supports expiry, revoke, and one-time accept.
+- [x] WorkspaceExecutionState stores generation, state, source, and write fence.
+- [x] OSS initialization creates exactly one Workspace and one owner membership atomically.
+- [x] OSS refuses a second Workspace while allowing multiple members.
+
+### Migration
+
+- [x] Alembic migration upgrades SQLite.
+- [x] Alembic migration upgrades PostgreSQL.
+- [x] Existing first user becomes owner of the default Workspace.
+- [x] Existing tenant resources are backfilled with the default Workspace UUID.
+- [x] SQLite destructive boundaries create verified, revision-aware backups and atomically restore after failure.
+- [x] Migration can resume safely after interruption.
+- [x] New installs and upgraded installs produce the same tenancy-kernel schema.
+
+## 2. Authentication and authorization
+
+### Identity
+
+- [x] JWT sub uses account UUID, with a bounded compatibility path for legacy email tokens.
+- [x] Disabled or deleted accounts cannot authenticate.
+- [x] Local password and Space-linked account flows support more than one local Account.
+- [x] Public registration closes after initialization by default.
+- [x] Invitation registration works without requiring SMTP.
+- [x] An unknown Space OAuth subject cannot claim an existing Account by email; explicit account-bound binding is required.
+
+### Request context
+
+- [x] PrincipalContext identifies Account, API Key, or trusted runtime principal.
+- [x] WorkspaceContext contains Workspace, Membership, role, permissions, and revision.
+- [x] RequestContext contains instance UUID, Workspace context, auth type, request ID, and generation.
+- [x] ExecutionContext propagates Workspace and generation to runtime work.
+- [x] SaaS-style requests never fall back to the first or most recent Workspace.
+- [x] OSS may resolve the single Workspace when the selector is omitted.
+- [x] Account-token bootstrap can list only the authenticated Account's active memberships before a Workspace selector exists.
+
+### Fixed RBAC
+
+- [x] owner, admin, developer, operator, and viewer permissions match the architecture matrix.
+- [x] Invitation cannot grant owner.
+- [x] The last owner cannot be removed or demoted.
+- [x] Cross-Workspace resources return 404.
+- [x] Same-Workspace permission failures return 403.
+
+## 3. Workspace and member APIs
+
+- [x] GET /api/v1/workspaces returns the OSS singleton Workspace.
+- [x] POST /api/v1/workspaces returns edition_limit in OSS.
+- [x] Current Workspace endpoint returns the authenticated Membership.
+- [x] Member list is permission scoped.
+- [x] Invitation create, revoke, inspect, and accept are atomic.
+- [x] Member role update and removal enforce owner rules.
+- [x] Invitation tokens travel in a request body and are redacted from logs.
+- [x] Relevant MCP tools and in-repo skills are updated with the same contract.
+
+## 4. Tenant-scoped persistence and services
+
+Each row type must have a non-null Workspace UUID, scoped indexes, scoped uniqueness, and scoped CRUD tests.
+
+- [x] Bots and bot admins.
+- [x] Legacy pipelines and pipeline run records.
+- [x] Model providers.
+- [x] LLM models.
+- [x] Embedding models.
+- [x] Rerank models.
+- [x] Plugin installations, settings, and configuration.
+- [x] MCP servers and resource preferences.
+- [x] Knowledge bases, files, and chunks.
+- [x] Vector collections and handles.
+- [x] Monitoring messages, calls, sessions, errors, embeddings, and feedback.
+- [x] API keys and scopes.
+- [x] Webhooks and public route resolution.
+- [x] Binary storage and Workspace storage.
+- [x] Workspace metadata, separated from system metadata.
+
+### Service and API rules
+
+- [x] Every tenant Service receives RequestContext or an explicit Workspace UUID.
+- [x] No tenant Service treats context None as global access.
+- [x] Every applicable get, list, create, update, delete, copy, export, and bulk operation is scoped.
+- [x] Parent-child references use the same Workspace.
+- [x] API Key authentication derives Workspace from the key, not a header.
+- [x] Webhook and Bot public routes derive Workspace from a trusted resource.
+- [x] Background jobs carry Workspace and generation explicitly.
+
+## 5. Runtime isolation
+
+### Core runtime
+
+- [x] RuntimeBot carries Workspace UUID and placement generation.
+- [x] RuntimePipeline carries Workspace UUID and placement generation.
+- [x] Query and Event carry Workspace UUID without making it an authorization source.
+- [x] Session key includes Workspace UUID, Bot UUID, launcher type, and launcher ID.
+- [x] QueryPool and manager indexes cannot collide across Workspaces.
+- [x] Query and aggregation cache keys and locks include Workspace UUID.
+- [x] Runtime transports, cached results, object operations, and long-lived tasks revalidate WorkspaceExecutionState generation at side-effect boundaries.
+- [ ] Ordinary tenant database writes hold the generation fence in the same transaction until commit; this remains a SaaS activation gate.
+
+### Plugin
+
+- [x] Plugin installation and configuration are Workspace scoped.
+- [x] Runtime control actions carry trusted Workspace binding and placement generation.
+- [x] A plugin process or supervisor never serves multiple Workspaces in SaaS mode.
+- [x] Host API derives Workspace from the connection, installation, and trusted action context, not plugin input.
+- [x] Plugin get_bots, models, tools, vector, RAG, configuration, and messaging calls are scoped.
+- [x] Plugin Workspace storage no longer uses owner default.
+- [x] Plugin page APIs check Membership and installation ownership.
+- [x] Local plugin launches use short-lived, one-use registration capabilities bound to manifest identity.
+
+### MCP, RAG, and Box
+
+- [x] MCP runtime key contains instance UUID, Workspace UUID, placement generation, and server UUID.
+- [x] Same-named MCP servers in two Workspaces do not share sessions.
+- [x] Pipeline cannot reference another Workspace's MCP resource.
+- [x] RAG collection names and handles are server-derived and Workspace scoped.
+- [x] Legacy global vector migration is available only to the local OSS singleton Workspace.
+- [x] Object storage paths include instance, Workspace, and placement generation for the fixed-generation OSS runtime.
+- [x] Object storage revalidates generation before touching a provider or resolving an opaque key.
+- [ ] Cloud cutover uses generation-scoped staging plus stable published object references, rather than making the staging generation the durable identity.
+- [x] Box persistent and ephemeral namespaces include the required instance, Workspace, and generation scope.
+- [x] Same-named Box sessions and processes cannot collide across Workspaces or placement generations.
+- [x] Box relay and process I/O reject or retire stale generations.
+- [x] External paths and privileged mounts cannot be supplied by an untrusted plugin.
+
+## 6. SDK and protocol
+
+- [x] Public Query, Event, Session, and context entities carry backward-compatible Workspace data.
+- [x] Action RPC request models carry trusted Workspace binding where required.
+- [x] Action enums and callers remain consistent.
+- [x] Old plugins continue to deserialize compatible events.
+- [x] Plugins cannot select an arbitrary Workspace through a Host API argument.
+- [x] Runtime storage uses the bound Workspace UUID.
+- [x] SDK API tests pass.
+- [x] Runtime tests pass.
+- [x] Action consistency script passes.
+
+## 7. Frontend
+
+- [x] Every browser tenant API request carries the current Workspace selector after bootstrap.
+- [x] OSS automatically selects the singleton Workspace.
+- [x] OSS does not show Create Workspace or a misleading switcher.
+- [x] Workspace settings show current Workspace information.
+- [x] Members page lists roles and permissions.
+- [x] Invitation creation shows a one-time link when SMTP is unavailable.
+- [x] Invitation acceptance supports a signed-out user flow.
+- [x] Role controls are hidden or disabled consistently with backend permissions.
+- [x] Switching accounts clears stale Workspace query cache and local state.
+- [x] User-facing strings support en_US, zh_Hans, and ja_JP.
+
+## 8. Automated verification
+
+### Persistence and authorization
+
+- [x] SQLite fresh install.
+- [x] SQLite upgrade from pre-tenant schema, including verified failure recovery.
+- [x] PostgreSQL fresh install.
+- [x] PostgreSQL upgrade from pre-tenant schema.
+- [x] All fixed roles have positive and negative permission-matrix tests.
+- [x] Concurrent invitation acceptance creates one Membership.
+- [x] Concurrent owner changes never leave zero owners.
+
+### Cross-tenant isolation
+
+- [x] Two Workspaces are created through a test-only policy.
+- [x] Applicable resource operations and parent-child references have cross-Workspace negative coverage.
+- [x] Resource UUID guessing cannot cross Workspace.
+- [x] API Key cannot cross Workspace.
+- [x] Plugin cannot enumerate or invoke another Workspace's resources.
+- [x] Sessions, caches, locks, MCP, RAG, Box, storage, and monitoring do not collide.
+- [x] Background jobs cannot execute without an explicit Workspace and placement generation.
+
+### Security and revocation
+
+- [x] Space login and binding use purpose-bound, one-time opaque OAuth state; caller-supplied state is rejected.
+- [x] OAuth redirects trust only server-configured WebUI or webhook origins, never request `Host` or `Origin` headers.
+- [x] Dashboard WebSockets revalidate authentication, Membership, resource, permission, and generation per message.
+- [x] Public embed WebSockets re-resolve Bot availability and execution binding per message.
+- [x] Runtime, storage, Plugin Runtime, MCP, RAG, and Box reject a stale placement generation.
+- [x] Unhandled API and webhook failures return a generic error plus request ID without exception text.
+- [x] URL user information and sensitive query parameters are redacted before configuration is serialized or logged.
+
+### Regression
+
+- [x] LangBot unit tests pass.
+- [x] LangBot integration tests pass.
+- [x] Frontend lint completes without errors and the production build passes.
+- [x] SDK focused and full relevant tests pass.
+- [x] Local SDK is installed into LangBot from the exact pushed SDK commit and cross-repo tests pass with no sync.
+
+## 9. Real browser E2E
+
+- [x] Start from a clean local data directory.
+- [x] First user initializes the singleton Workspace as owner.
+- [x] Owner creates an invitation link.
+- [x] A second signed-out browser identity accepts the invitation and registers.
+- [x] owner, admin, developer, operator, and viewer UI permissions match backend enforcement.
+- [x] Direct API calls cannot bypass hidden controls.
+- [x] Account switch does not expose prior account or Workspace data.
+- [x] Refresh and a new browser tab recover the correct Workspace safely.
+- [x] OSS rejects a second Workspace with `edition_limit`; same-name and same-identifier isolation is covered by the test-only multi-Workspace policy because OSS deliberately has no multi-Workspace browser surface.
+- [x] Explicit error states are visible for expired, revoked, reused, and email-mismatched invitations.
+
+## 10. Completion evidence
+
+- [x] LangBot and SDK branch refs are recorded in the verification report.
+- [x] Space git diff is empty relative to the pre-work snapshot.
+- [x] Migration output is captured for SQLite and PostgreSQL.
+- [x] Test commands and results are recorded.
+- [x] Browser E2E actions and observed results are recorded.
+- [x] No remaining tenant table, global Service query, owner default, or unscoped runtime key is found by the final audit.
diff --git a/docs/multi-tenant/implementation-decisions.md b/docs/multi-tenant/implementation-decisions.md
new file mode 100644
index 000000000..c343c6969
--- /dev/null
+++ b/docs/multi-tenant/implementation-decisions.md
@@ -0,0 +1,175 @@
+# Multi-tenant implementation decisions
+
+This log records implementation choices made while delivering the Workspace architecture. It is intended to make trade-offs auditable without interrupting implementation for routine decisions.
+
+## 2026-07-18
+
+### OSS remains a singleton Workspace with multiple Accounts
+
+- Decision: Community builds create exactly one Workspace per LangBot instance and allow multiple Accounts through invitations.
+- Reason: This preserves a simple self-hosted deployment while making authorization and ownership explicit. Creating a second Workspace is an edition error, not a hidden fallback.
+- SaaS boundary: Multi-Workspace directory, placement, entitlement, and billing are the responsibility of a separate closed SaaS Control Plane. Core consumes a validated projection and remains the final isolation and authorization enforcement point; it does not become the SaaS system of record or billing engine.
+- Deployment boundary: Cloud v2 is a greenfield deployment design. The previous Space instance/pod deployment scheme is not migrated, preserved for compatibility, or extended by this implementation. Existing Space OAuth, marketplace, and payment concepts may be reused only through explicit adapters where they still fit; `langbot-space` itself is not changed.
+
+### Workspace selection is trusted only after authentication
+
+- Decision: Browser requests carry `X-Workspace-Id`, but the server resolves it against the authenticated Account membership. API keys, public Bot routes, webhooks, jobs, and plugin calls derive Workspace from their trusted owning resource or binding instead of trusting the header.
+- Reason: A selector is routing input, not authorization evidence.
+- Compatibility: Community builds may select the singleton Workspace when the header is omitted. A multi-Workspace-capable build must reject an omitted selector.
+
+### Stable Account UUID is the token subject
+
+- Decision: New JWTs use the stable Account UUID as `sub`; a bounded compatibility path accepts legacy email-subject tokens and rotates them when checked.
+- Reason: Email can change and therefore cannot be a durable authorization identity.
+
+### Fixed roles are authoritative in Core
+
+- Decision: `owner`, `admin`, `developer`, `operator`, and `viewer` map to a fixed permission matrix in LangBot Core. The last owner cannot be removed or demoted, and invitations cannot create an owner directly.
+- Reason: Core must remain the final authorization boundary in both OSS and SaaS deployments.
+
+### Cross-Workspace access is indistinguishable from absence
+
+- Decision: Resource lookups always include Workspace UUID. A guessed UUID belonging to another Workspace returns 404; a visible resource with insufficient same-Workspace permission returns 403.
+- Reason: This avoids leaking resource existence across tenants while preserving actionable same-tenant errors.
+
+### Plugin Runtime binds to exactly one Workspace
+
+- Decision: A trusted LangBot control connection binds a Plugin Runtime to one `instance_uuid`, `workspace_uuid`, `placement_generation`, and optional installation. Repeating the same binding is idempotent; rebinding is rejected. Plugin-supplied context fields are stripped.
+- Reason: One untrusted plugin process must never become a cross-Workspace router.
+- Compatibility: Older SDK payloads without context still deserialize, but Workspace storage and tenant Host APIs fail explicitly until a trusted binding is established.
+- Startup fence: The Runtime does not launch or register plugins until `SET_RUNTIME_CONFIG` establishes the trusted Workspace binding. A cloud or multi-Workspace connector must be constructed for one explicit projected Workspace and generation; it never falls back to a migration-created local compatibility Workspace.
+
+### Invitation delivery does not require SMTP
+
+- Decision: Core returns an invitation secret once for copy-and-share, persists only its hash, and supports expiry, revocation, and one-time acceptance.
+- Reason: Self-hosted OSS must support adding users without an email service while avoiding recoverable invitation secrets at rest.
+- Browser handling: The copyable invitation URL carries the secret in its fragment, which browsers do not send in HTTP requests or Referer headers. The acceptance page immediately removes the fragment and keeps the secret only in `sessionStorage` until login or acceptance completes; it is never placed in a path, query string, analytics event, or persistent local storage.
+
+### Schema rollout is additive before enforcement
+
+- Decision: Add Account/Workspace directory tables first, then add non-null Workspace ownership to every tenant resource with a deterministic default-Workspace backfill. Runtime and service enforcement is enabled only with matching migration and isolation tests.
+- Reason: A Workspace column alone is not isolation, and enforcing queries before data backfill would break upgraded installations.
+
+### Login capability discovery is instance-scoped, not account-scoped
+
+- Decision: The unauthenticated login bootstrap endpoint reports only which login mechanisms the instance supports. It does not inspect the first Account or expose whether that Account has a password. Both password and Space OAuth entry points are available on a multi-user instance; the submitted identity determines which mechanism is valid.
+- Reason: This avoids projecting the original owner's authentication type onto invited users and removes a public Account-state disclosure.
+
+### Space OAuth identity does not choose a SaaS Workspace
+
+- Decision: Space OAuth tokens remain Account credentials. In OSS singleton mode, an OAuth refresh may update the singleton Workspace's Space provider only when that Account's role can manage provider secrets. In SaaS multi-Workspace mode an OAuth callback without an authenticated Workspace selector never guesses which Workspace to mutate; explicit Workspace configuration or the closed control plane owns that linkage.
+- Reason: An Account may belong to several Workspaces, and authentication must not silently mutate a shared tenant secret.
+
+### SaaS execution state is a validated Core projection
+
+- Decision: Core can resolve both local and `cloud_projection` Workspaces, but only from an explicit Workspace UUID and an active, unfenced `WorkspaceExecutionState` for the current instance and matching source. OSS-only bootstrap paths additionally require `source=local`.
+- Reason: The closed control plane owns placement decisions, while Core remains the enforcement point for instance binding, generation, and write fences.
+
+### API-key secrets are one-time and Workspace-bound
+
+- Decision: Database API keys persist only a globally unique SHA-256 hash, an opaque UUID, one Workspace UUID, explicit fixed-permission scopes, status, expiry, creator, and last-used time. The raw secret is returned once. Authentication derives Workspace and generation from the key record and ignores Workspace selectors. Legacy plaintext keys are hashed during migration and receive a compatibility `*` scope. The plaintext config key works only for the OSS singleton Workspace and is disabled in multi-Workspace mode.
+- Reason: A bearer key is an identity and routing credential, not merely a password layered on top of caller-controlled tenant selection.
+
+### MCP tools inherit the authenticated API-key context
+
+- Decision: The MCP ASGI mount authenticates the API key once, binds an immutable per-request `RequestContext`, and every tool checks a fixed permission before calling tenant services with that same context.
+- Reason: Authenticating the transport without propagating Workspace identity into tool calls would leave the direct service path globally scoped.
+
+### Unreleased SDK protocol is pinned reproducibly without publishing
+
+- Decision: The SDK tenancy protocol is versioned as 0.4.15. This task does not create a GitHub release or publish PyPI because the user authorized pushing code, not a package release. After the SDK feature branch is final, LangBot's feature branch temporarily pins the exact pushed SDK Git commit. Before merging to master, the release gate is to publish `langbot-plugin==0.4.15` and replace the Git pin with the registry pin.
+- Reason: PyPI 0.4.14 does not contain `ActionContext`; pinning it makes clean installs fail, while pinning an unpublished 0.4.15 makes dependency sync impossible. An exact Git commit is reproducible and keeps the feature branch testable without expanding release authority.
+
+### Cloud directory writes stay outside Core
+
+- Decision: The open-source Core startup always installs `SingleWorkspacePolicy`, creates or repairs one local Workspace, and permits local membership/invitation workflows. Changing mutable configuration such as `system.edition` cannot activate multi-Workspace routing. The future closed Cloud bootstrap will install `CloudWorkspacePolicy` only after verifying a signed `InstanceManifest`; that policy requires an explicit projected Workspace selector, does not create Workspaces, and rejects invitation or membership mutations with `control_plane_required`; member reads use the versioned local projection.
+- Ownership split: The closed Control Plane owns the global Account/Workspace/Membership/Invitation directory, placement and lease generations, entitlements, subscription state, usage aggregation, and billing decisions. Core owns request authorization, resource scoping, execution-generation validation, and fail-closed enforcement. Provisioning and invoice computation do not belong in open-source Core.
+- Reason: The closed control plane is authoritative for SaaS Account, Workspace, Membership, and Invitation state. Allowing Core to mutate the same directory would create split-brain ownership and would make an ownerless compatibility Workspace a dangerous fallback.
+- Release gate: Multi-Workspace activation is deliberately unavailable in the open-source bootstrap. Production Cloud v2 must implement the signed `InstanceManifest` verifier and closed bootstrap described in the architecture document before it can inject `CloudWorkspacePolicy`; `edition=cloud`, an environment variable, or any unsigned local configuration is never a valid activation credential.
+
+### Workspace bootstrap is reactive and ordered before browser resource calls
+
+- Decision: The web application blocks Workspace-owned pages until Account and current Workspace bootstrap completes. A `useSyncExternalStore` Workspace store publishes permission changes to React consumers; direct mutation-only routes and controls are hidden or disabled when the fixed role lacks the required permission.
+- Reason: Mutating a module-level variable after the initial React render did not reliably re-render permission controls, and mounting resource pages before the selector was established could issue tenant requests without `X-Workspace-Id`.
+
+### JWTs are bound to one LangBot instance
+
+- Decision: New Core JWTs require `iss=langbot-core`, an audience derived from the immutable instance UUID, and an expiry. Legacy community tokens are accepted only when they have the historical issuer, carry no audience, and the active policy is the OSS singleton policy.
+- Reason: A token issued by one instance must not authenticate against another instance that happens to share a secret, and a compatibility decoder must not become an alternate path around the SaaS trust boundary.
+
+### Runtime control transports authenticate before protocol dispatch
+
+- Decision: External Plugin Runtime and Box WebSocket control channels require independent strong shared secrets in handshake headers. Locally managed child processes receive ephemeral secrets through their environment; secrets are not placed in URLs, process arguments, request payloads, or logs. Box additionally binds the first authenticated control channel to one trusted instance. Plugin Runtime debug and control credentials remain separate.
+- Reason: Workspace context inside an RPC payload is not trustworthy until the transport peer itself is authenticated. Separating control and debug credentials also limits accidental privilege reuse.
+- Deployment consequence: Docker Compose and Kubernetes wire one shared secret to each host/runtime pair. An empty external-runtime secret fails startup instead of silently exposing an unauthenticated socket.
+
+### Dashboard WebSocket sessions are tenant runtime objects
+
+- Decision: A dashboard WebSocket sends an authentication frame immediately after upgrade. The server validates Account, Membership, permission, Pipeline ownership, instance, Workspace, and placement generation before registering the connection. Connection indexes, sessions, broadcasts, attachments, and resets include the complete execution scope.
+- Reason: Browser WebSocket APIs cannot attach the normal authorization headers, and a process-global `pipeline_uuid` or `session_type` index can collide across Workspaces.
+
+### Read permissions never imply secret permissions
+
+- Decision: `resource.view` responses recursively redact Bot, Plugin, MCP, and provider credentials. Provider secrets require `provider_secret.manage`; Bot and Plugin configuration writes require `resource.manage`. Masked Plugin values can be round-tripped by a manager without overwriting the stored secret. Plugin Runtime debug credentials require `resource.manage`, not the operator-only `runtime.operate` permission.
+- Reason: A multi-user Workspace needs useful viewer access without turning every visible configuration endpoint into credential export. Plugin debug attachment can register executable code and is therefore a resource-management operation.
+
+### Temporary credential exchanges are bound to their initiator
+
+- Decision: Lark, Weixin, DingTalk, WeComBot, and QQOfficial one-click registration sessions require `resource.manage` and store the initiating instance, Workspace, placement generation, and principal. Status and cancellation by any other scope return the same 404 as an unknown session.
+- Reason: Random session IDs reduce guessing probability but do not authorize access to credentials returned by a completed exchange.
+
+### Uploaded images and documents use different storage capabilities
+
+- Decision: Browser images use the scoped `upload_image` owner type and may be resolved only through the opaque public-image route. RAG documents use `upload_document` and can be read, sized, or deleted only by an exact instance, Workspace, generation, and owner-type match. Legacy `upload` objects are cleanup-only.
+- Reason: Treating every upload as a public image made a leaked document key sufficient to bypass authenticated RAG access.
+
+## 2026-07-19
+
+### Space OAuth state is server-issued and single-use
+
+- Decision: Core issues an opaque, cryptographically random OAuth state for each Space login or Account-binding attempt, stores only its digest, and consumes it exactly once within a short expiry. Login and binding states are different capabilities; a binding state is additionally bound to the authenticated Account. Caller-supplied state, including a LangBot JWT, is rejected.
+- Redirect boundary: Callback redirects are accepted only for the known callback path and an origin declared by the server-side `api.webui_url` or `api.webhook_prefix`. Request `Host` and `Origin` headers never expand this allowlist.
+- Current deployment: The OSS state store is bounded and process-local, so a Core restart safely invalidates outstanding attempts. A horizontally scaled SaaS deployment must move this exchange to an atomic, shared Control Plane store before enabling the closed Cloud bootstrap.
+- Reason: OAuth state is a narrow, one-time CSRF and flow-binding capability. Reusing a bearer JWT or trusting caller-controlled Host or Origin data would turn an authorization redirect into an Account-token theft or open-redirect primitive.
+
+### OAuth provider subjects, not email addresses, bind Accounts
+
+- Decision: A known Space `account_uuid` may refresh the credentials of its already-bound local Account. An unknown provider subject that presents an email belonging to an existing Account is rejected, even when the normalized emails match. The Account owner must authenticate locally and use the one-time, account-bound binding flow.
+- Reason: Email is contact and display data, not a stable federated identity key. Email-only auto-linking would let provider verification drift or identity reassignment become a local Account takeover.
+
+### Workspace discovery is an account-only bootstrap capability
+
+- Decision: `ACCOUNT_TOKEN` validates the active Account JWT but intentionally cannot resolve a Workspace, receive `RequestContext`, or declare Workspace permissions. Its narrow bootstrap endpoint returns only active Workspace memberships belonging to that Account and never chooses the first Workspace when several exist. All tenant resource routes still require the explicit selector in multi-Workspace mode.
+- Reason: Requiring a Workspace header to discover the Account's Workspaces creates an authentication deadlock; allowing the bootstrap route to perform tenant actions would create an authorization bypass. Separating the two capabilities resolves the cycle without weakening tenant routes.
+
+### SQLite tenancy migrations have a verified recovery boundary
+
+- Decision: Before each destructive tenant-schema boundary, a file-backed SQLite installation creates an online-consistent backup with its source and target revisions, runs `PRAGMA quick_check`, writes a durable manifest, and fsyncs restrictive-permission files and directories. A failed boundary disposes the engine, removes stale journal sidecars, atomically restores the verified source revision, and verifies the restored database before startup continues.
+- Compatibility: In-memory SQLite cannot provide this recovery guarantee and is rejected for destructive production migration boundaries; it remains usable in tests that create the final schema directly.
+- Reason: SQLite batch table rebuilds can leave an installation between schemas if a process or migration fails. A verified pre-boundary image makes retry behavior recoverable instead of merely idempotent in the happy path.
+
+### Placement generation is an execution revocation capability
+
+- Decision: RuntimeBot, RuntimePipeline, background tasks, object storage, Plugin Runtime, MCP, RAG, and Box operations carry the complete instance, Workspace, and placement-generation scope. They revalidate the active execution binding before accessing a provider or transport; long-running calls validate again before accepting results. A stale generation is fenced before it can read, write, or reuse a cached object.
+- Plugin boundary: Each locally launched plugin receives a short-lived, one-use registration capability bound to the expected manifest identity and execution scope. The production child environment does not inherit the reusable debug credential, and Host APIs derive scope from the trusted connection and action context.
+- Box boundary: Persistent skill content remains Workspace-scoped, while session/process state and relay requests also include placement generation. A generation change retires matching live sessions and closes a stale relay before further stdin, stdout, or file operations.
+- Transaction boundary: Request admission and runtime side effects are fenced in this branch, but ordinary tenant database mutations do not yet hold a generation-aware lock through commit. The closed Cloud bootstrap must remain disabled until Core provides the shared-write/exclusive-cutover transaction primitive and a generation-stamped outbox (or an equivalent atomic publish fence). The OSS singleton policy has a fixed local generation and cannot trigger a placement cutover.
+- Durable-object boundary: Current opaque storage keys include generation and therefore fail closed after a generation change. That is safe for OSS's fixed generation, but a Cloud cutover must not strand durable KB files, images, or plugin references. Cloud v2 must publish stable final object identities from generation-scoped staging, or perform an atomic object-and-reference migration before activating the new generation.
+- Reason: Workspace UUID prevents cross-tenant collisions, but it cannot revoke work after a Workspace is moved or fenced. Placement generation is the monotonic lease that makes old runtimes unusable.
+
+### Long-lived WebSockets continuously revalidate authority
+
+- Decision: Dashboard WebSockets re-authenticate the Account, Membership, permission, resource ownership, instance, Workspace, and placement generation for every inbound message, not only during the initial frame. A changed role, removed Membership, or fenced placement takes effect without waiting for reconnect.
+- Public embed boundary: The embed connection re-resolves its Bot before every message and rejects a Bot that was disabled, deleted, moved, or rebound. The public connection may identify a Bot, but it cannot make the initial Bot object an indefinite authorization capability.
+- Reason: Authorization and resource state can change while a socket remains open. Connection-time validation alone leaves a revocation gap.
+
+### Legacy vector migration is an OSS-local compatibility path
+
+- Decision: Status, backup, execute, dismiss, and background entry points for legacy global vector collections require an active local Workspace binding under `SingleWorkspacePolicy`. A `cloud_projection` Workspace cannot observe or migrate the old global collection, even when it carries a legacy marker.
+- Reason: The legacy collection predates tenant ownership. Treating it as a SaaS fallback would expose one installation's historical vectors to an arbitrary projected Workspace.
+
+### External errors and persisted URLs are redacted centrally
+
+- Decision: Unhandled HTTP and webhook failures return a stable `internal_error` response and request ID, and expose that ID in `X-Request-Id`; the detailed exception is retained only in server logs correlated by the same ID. Explicit domain and validation errors keep their documented status and code.
+- Secret boundary: Shared sanitization removes URL user information and masks sensitive query parameters before provider or MCP configuration is serialized, logged, or shown to a reader. Masked placeholders can be round-tripped by an authorized manager without replacing the stored secret.
+- Reason: Tenant isolation is incomplete if framework exceptions, connection URLs, or configuration reads can export credentials across otherwise authorized interfaces.
diff --git a/pyproject.toml b/pyproject.toml
index 1f1c3e438..c391b1351 100644
--- a/pyproject.toml
+++ b/pyproject.toml
@@ -70,7 +70,7 @@ dependencies = [
"chromadb>=1.0.0,<2.0.0",
"qdrant-client (>=1.15.1,<2.0.0)",
"pyseekdb==1.1.0.post3",
- "langbot-plugin==0.4.13",
+ "langbot-plugin @ git+https://github.com/langbot-app/langbot-plugin-sdk.git@a1544b6b38a37ba72e3284f2836618144f0742c1",
"asyncpg>=0.30.0",
"line-bot-sdk>=3.19.0",
"matrix-nio>=0.25.2",
diff --git a/skills/skills/langbot-deploy/SKILL.md b/skills/skills/langbot-deploy/SKILL.md
index 3a314fbfa..bfde2bd3e 100644
--- a/skills/skills/langbot-deploy/SKILL.md
+++ b/skills/skills/langbot-deploy/SKILL.md
@@ -27,6 +27,17 @@ The `all` / `box` profile starts three services:
- `langbot_box` — Box sandbox runtime (`:5410`). Uses the host Docker socket to
spawn sandbox containers, so the **Box root host path and in-container path
must be identical** (`BOX__LOCAL__HOST_ROOT=${LANGBOT_BOX_ROOT:-${PWD}/data/box}`).
+ Its RPC and managed-process relay require a shared
+ `LANGBOT_BOX_CONTROL_TOKEN` (at least 32 non-whitespace characters) in both
+ the LangBot and Box containers. Generate it once with `openssl rand -hex 32`;
+ never put it in `box.runtime.endpoint` or commit it to config.
+
+Every Compose deployment also needs one
+`LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN` shared by `langbot` and
+`langbot_plugin_runtime`. Generate it with `openssl rand -hex 32` and export it
+before `docker compose up`; the external Plugin Runtime fails closed when the
+token is empty or weak. Kubernetes uses the `langbot-plugin-runtime-control`
+Secret shown in `docker/kubernetes.yaml`.
With Box off, the dashboard/skills list stays visible (read-only) but sandbox
tools, skill add/edit, and stdio MCP are disabled. Set `box.enabled: false`
diff --git a/skills/skills/langbot-dev/SKILL.md b/skills/skills/langbot-dev/SKILL.md
index 01ee06fd1..77a988ce8 100644
--- a/skills/skills/langbot-dev/SKILL.md
+++ b/skills/skills/langbot-dev/SKILL.md
@@ -65,10 +65,23 @@ Route auth is declared per-route via `AuthType` in
- `API_KEY` — `X-API-Key` or `Authorization: Bearer `.
- `USER_TOKEN_OR_API_KEY` — either.
-API keys are verified by `apikey_service.verify_api_key()`, which accepts:
-1. the **global key** from `config.yaml` `api.global_api_key` (no DB, no login,
- no `lbk_` prefix required), then
-2. **web-UI keys** (DB-stored, `lbk_` prefix).
+Authenticated routes receive an immutable `RequestContext` containing the
+principal, authorized Workspace membership, fixed-role permissions, instance,
+request id, and placement generation. A browser's `X-Workspace-Id` is only a
+selector and is always checked against the Account membership. Tenant services
+must accept this context (or an explicit trusted execution context) and fail
+closed when it is absent.
+
+API-key authentication accepts:
+
+1. the **global key** from `config.yaml` `api.global_api_key` only for a
+ community instance with exactly one local Workspace, then
+2. **web-UI keys** whose one-time `lbk_` secret is stored only as a hash and is
+ bound to one Workspace, explicit scopes, status, and optional expiry.
+
+An API key derives its Workspace from the key record and ignores a caller's
+Workspace selector. Public Bot/Webhook routes similarly derive Workspace from
+the opaque owning resource rather than a header.
Route groups self-register via `@group.group_class(name, path)` and are
discovered by `importutil.import_modules_in_pkg`.
diff --git a/skills/skills/langbot-mcp-ops/SKILL.md b/skills/skills/langbot-mcp-ops/SKILL.md
index 94f25a3a4..fb1152fa4 100644
--- a/skills/skills/langbot-mcp-ops/SKILL.md
+++ b/skills/skills/langbot-mcp-ops/SKILL.md
@@ -29,13 +29,19 @@ Authorization: Bearer
Two kinds of key are accepted:
-1. **Web-UI key** — created in the web UI (sidebar → API Keys), prefixed `lbk_`,
- stored in the database.
+1. **Web-UI key** — created in the web UI (sidebar → API Keys), prefixed `lbk_`.
+ The secret is shown once; only its SHA-256 hash is stored. Each key is bound
+ to one Workspace and has explicit scopes, status, optional expiry, and
+ last-used metadata. The key determines the Workspace; callers cannot switch
+ it with `X-Workspace-Id`.
2. **Global API key** — set in `data/config.yaml` under `api.global_api_key`.
Requires no login session and no DB record; does not need the `lbk_` prefix.
- Leave empty to disable. See the `langbot-deploy` skill for config details.
+ It is accepted only by a community instance with exactly one local
+ Workspace and is disabled for SaaS multi-Workspace operation. Leave empty to
+ disable. See the `langbot-deploy` skill for config details.
-Requests without a valid key get `401 Unauthorized`.
+Invalid, revoked, or expired keys get `401 Unauthorized`. A valid key whose
+scopes do not authorize a tool gets `403 Forbidden`.
## Client configuration
@@ -66,7 +72,9 @@ The tools wrap the LangBot service layer. Current tools (v1):
Mutating tools (`create_*`, `update_*`) take a JSON object matching the same
shape as the corresponding HTTP API request body. Discover resources with the
-`list_*` / `get_*` tools before mutating; identifiers are UUIDs.
+`list_*` / `get_*` tools before mutating; identifiers are UUIDs. Reads require
+`resource.view`; mutations require `resource.manage`. All service calls inherit
+the immutable Workspace context authenticated at the MCP transport boundary.
## How to use
@@ -93,7 +101,8 @@ shape as the corresponding HTTP API request body. Discover resources with the
- `/mcp` is the **server** LangBot exposes. The `/api/v1/mcp` routes are the
**client** side (managing external MCP servers LangBot connects to). Don't
confuse them.
-- A `401` means the key is wrong, missing, or (for the global key)
- `api.global_api_key` is empty in config.yaml.
+- A `401` means the key is wrong, missing, revoked, expired, or (for the global
+ key) `api.global_api_key` is empty or the instance is not an OSS singleton.
+- A `403` means the key is valid but lacks the permission required by the tool.
- The global key is plaintext in config.yaml — only enable it on trusted/internal
deployments and serve over HTTPS.
diff --git a/src/langbot/pkg/api/http/authz.py b/src/langbot/pkg/api/http/authz.py
new file mode 100644
index 000000000..36b796fe7
--- /dev/null
+++ b/src/langbot/pkg/api/http/authz.py
@@ -0,0 +1,116 @@
+from __future__ import annotations
+
+import enum
+import types
+import typing
+
+from .context import RequestContext
+
+
+class WorkspaceRole(enum.StrEnum):
+ OWNER = 'owner'
+ ADMIN = 'admin'
+ DEVELOPER = 'developer'
+ OPERATOR = 'operator'
+ VIEWER = 'viewer'
+
+
+class Permission(enum.StrEnum):
+ WORKSPACE_VIEW = 'workspace.view'
+ WORKSPACE_UPDATE = 'workspace.update'
+ WORKSPACE_DELETE = 'workspace.delete'
+ OWNER_TRANSFER = 'owner.transfer'
+ MEMBER_VIEW = 'member.view'
+ MEMBER_INVITE = 'member.invite'
+ MEMBER_UPDATE_ROLE = 'member.update_role'
+ MEMBER_REMOVE = 'member.remove'
+ RESOURCE_VIEW = 'resource.view'
+ RESOURCE_MANAGE = 'resource.manage'
+ RUNTIME_OPERATE = 'runtime.operate'
+ PROVIDER_SECRET_MANAGE = 'provider_secret.manage'
+ API_KEY_MANAGE = 'api_key.manage'
+ AUDIT_VIEW = 'audit.view'
+ DATA_EXPORT = 'data.export'
+ BILLING_LINK_MANAGE = 'billing_link.manage'
+
+
+_VIEW_PERMISSIONS = {
+ Permission.WORKSPACE_VIEW,
+ Permission.MEMBER_VIEW,
+ Permission.RESOURCE_VIEW,
+}
+
+_ROLE_PERMISSIONS: typing.Final = types.MappingProxyType(
+ {
+ WorkspaceRole.OWNER: frozenset(Permission),
+ WorkspaceRole.ADMIN: frozenset(
+ permission
+ for permission in Permission
+ if permission
+ not in {
+ Permission.WORKSPACE_DELETE,
+ Permission.OWNER_TRANSFER,
+ Permission.BILLING_LINK_MANAGE,
+ }
+ ),
+ WorkspaceRole.DEVELOPER: frozenset(
+ _VIEW_PERMISSIONS
+ | {
+ Permission.RESOURCE_MANAGE,
+ Permission.RUNTIME_OPERATE,
+ Permission.PROVIDER_SECRET_MANAGE,
+ }
+ ),
+ WorkspaceRole.OPERATOR: frozenset(_VIEW_PERMISSIONS | {Permission.RUNTIME_OPERATE}),
+ WorkspaceRole.VIEWER: frozenset(_VIEW_PERMISSIONS),
+ }
+)
+
+
+class AuthorizationError(Exception):
+ """Base class for errors that map to an HTTP authorization response."""
+
+ status_code = 403
+ error_code = 'forbidden'
+
+
+class WorkspaceRequiredError(AuthorizationError):
+ status_code = 400
+ error_code = 'workspace_required'
+
+
+class PermissionDeniedError(AuthorizationError):
+ error_code = 'permission_denied'
+
+ def __init__(self, permission: str) -> None:
+ super().__init__(f'Missing Workspace permission: {permission}')
+ self.permission = permission
+
+
+class EditionLimitError(AuthorizationError):
+ error_code = 'edition_limit'
+
+
+def permissions_for_role(role: str | WorkspaceRole) -> frozenset[str]:
+ """Return the canonical fixed permissions for a Workspace role."""
+
+ try:
+ parsed_role = WorkspaceRole(role)
+ except ValueError:
+ return frozenset()
+ return frozenset(permission.value for permission in _ROLE_PERMISSIONS[parsed_role])
+
+
+def has_permission(ctx: RequestContext, permission: str | Permission) -> bool:
+ """Return whether the context contains one effective permission."""
+
+ permission_value = permission.value if isinstance(permission, Permission) else permission
+ return permission_value in ctx.workspace.permissions
+
+
+def require_permission(ctx: RequestContext, permission: str | Permission) -> None:
+ """Raise a stable authorization error when a permission is missing."""
+
+ permission_value = permission.value if isinstance(permission, Permission) else permission
+ if not has_permission(ctx, permission_value):
+ raise PermissionDeniedError(permission_value)
diff --git a/src/langbot/pkg/api/http/context.py b/src/langbot/pkg/api/http/context.py
new file mode 100644
index 000000000..8236c8dc7
--- /dev/null
+++ b/src/langbot/pkg/api/http/context.py
@@ -0,0 +1,92 @@
+from __future__ import annotations
+
+import dataclasses
+import enum
+
+
+class PrincipalType(enum.StrEnum):
+ """Kinds of authenticated principals accepted by LangBot."""
+
+ ACCOUNT = 'account'
+ API_KEY = 'api_key'
+ SYSTEM = 'system'
+ PUBLIC_BOT = 'public_bot'
+
+
+@dataclasses.dataclass(frozen=True, slots=True)
+class PrincipalContext:
+ """Authenticated identity before Workspace authorization is applied."""
+
+ principal_type: PrincipalType
+ account_uuid: str | None = None
+ api_key_uuid: str | None = None
+
+
+@dataclasses.dataclass(frozen=True, slots=True)
+class WorkspaceContext:
+ """Workspace membership and effective permissions for one request."""
+
+ workspace_uuid: str
+ membership_uuid: str | None
+ role: str | None
+ permissions: frozenset[str]
+ membership_revision: int = 0
+
+
+@dataclasses.dataclass(frozen=True, slots=True)
+class RequestContext:
+ """Trusted authorization context passed to HTTP services."""
+
+ instance_uuid: str
+ placement_generation: int
+ request_id: str
+ auth_type: str
+ principal: PrincipalContext
+ workspace: WorkspaceContext
+ entitlement_revision: int = 0
+
+ @property
+ def workspace_uuid(self) -> str:
+ """Return the selected Workspace UUID."""
+
+ return self.workspace.workspace_uuid
+
+ @property
+ def account_uuid(self) -> str | None:
+ """Return the Account UUID when the principal is an Account."""
+
+ return self.principal.account_uuid
+
+
+@dataclasses.dataclass(frozen=True, slots=True)
+class ExecutionContext:
+ """Workspace context propagated to asynchronous and runtime work."""
+
+ instance_uuid: str
+ workspace_uuid: str
+ placement_generation: int
+ bot_uuid: str | None = None
+ pipeline_uuid: str | None = None
+ query_uuid: str | None = None
+ trigger_principal: PrincipalContext | None = None
+
+ @classmethod
+ def from_request(
+ cls,
+ ctx: RequestContext,
+ *,
+ bot_uuid: str | None = None,
+ pipeline_uuid: str | None = None,
+ query_uuid: str | None = None,
+ ) -> ExecutionContext:
+ """Create a runtime context without losing the tenant generation."""
+
+ return cls(
+ instance_uuid=ctx.instance_uuid,
+ workspace_uuid=ctx.workspace_uuid,
+ placement_generation=ctx.placement_generation,
+ bot_uuid=bot_uuid,
+ pipeline_uuid=pipeline_uuid,
+ query_uuid=query_uuid,
+ trigger_principal=ctx.principal,
+ )
diff --git a/src/langbot/pkg/api/http/controller/group.py b/src/langbot/pkg/api/http/controller/group.py
index 2ed55187d..2d8c95bbc 100644
--- a/src/langbot/pkg/api/http/controller/group.py
+++ b/src/langbot/pkg/api/http/controller/group.py
@@ -5,9 +5,18 @@ import typing
import enum
import quart
import traceback
+import inspect
+import uuid
from quart.typing import RouteCallable
-from ....core import app
+from ....utils import constants
+from ....workspace.collaboration import MembershipPermissionError, WorkspaceCollaborationError
+from ....workspace.errors import WorkspaceNotFoundError
+from ..authz import AuthorizationError, Permission, permissions_for_role, require_permission
+from ..context import PrincipalContext, PrincipalType, RequestContext, WorkspaceContext
+
+if typing.TYPE_CHECKING:
+ from ....core.app import Application
# Maximum file upload size limit (10MB)
MAX_FILE_SIZE = 10 * 1024 * 1024 # 10MB
@@ -33,6 +42,7 @@ class AuthType(enum.Enum):
"""Authentication type"""
NONE = 'none'
+ ACCOUNT_TOKEN = 'account-token'
USER_TOKEN = 'user-token'
API_KEY = 'api-key'
USER_TOKEN_OR_API_KEY = 'user-token-or-api-key'
@@ -43,11 +53,11 @@ class RouterGroup(abc.ABC):
path: str
- ap: app.Application
+ ap: Application
quart_app: quart.Quart
- def __init__(self, ap: app.Application, quart_app: quart.Quart) -> None:
+ def __init__(self, ap: Application, quart_app: quart.Quart) -> None:
self.ap = ap
self.quart_app = quart_app
@@ -59,16 +69,37 @@ class RouterGroup(abc.ABC):
self,
rule: str,
auth_type: AuthType = AuthType.USER_TOKEN,
+ permission: Permission | str | None = None,
**options: typing.Any,
) -> typing.Callable[[RouteCallable], RouteCallable]: # decorator
"""Register a route"""
+ if auth_type == AuthType.ACCOUNT_TOKEN and permission is not None:
+ raise ValueError('Account-token routes cannot declare Workspace permissions')
+
def decorator(f: RouteCallable) -> RouteCallable:
nonlocal rule
rule = self.path + rule
async def handler_error(*args, **kwargs):
- if auth_type == AuthType.USER_TOKEN:
+ if auth_type == AuthType.ACCOUNT_TOKEN:
+ authorization = quart.request.headers.get('Authorization', '')
+ if not authorization.startswith('Bearer '):
+ return self.http_status(401, -1, 'No valid user token provided')
+ token = authorization.removeprefix('Bearer ')
+ if not token:
+ return self.http_status(401, -1, 'No valid user token provided')
+
+ try:
+ _, user_email = await self._authenticate_account(token)
+ # Account-token routes deliberately stop before Workspace
+ # selection. They may bootstrap a selector, but cannot
+ # receive RequestContext or enforce Workspace permissions.
+ self._inject_handler_context(f, kwargs, user_email, None)
+ except Exception as e:
+ return self._auth_error_response(e)
+
+ elif auth_type == AuthType.USER_TOKEN:
# get token from Authorization header
token = quart.request.headers.get('Authorization', '').replace('Bearer ', '')
@@ -76,18 +107,15 @@ class RouterGroup(abc.ABC):
return self.http_status(401, -1, 'No valid user token provided')
try:
- user_email = await self.ap.user_service.verify_jwt_token(token)
-
- # check if this account exists
- user = await self.ap.user_service.get_user_by_email(user_email)
- if not user:
- return self.http_status(401, -1, 'User not found')
-
- # check if f accepts user_email parameter
- if 'user_email' in f.__code__.co_varnames:
- kwargs['user_email'] = user_email
+ account, user_email = await self._authenticate_account(token)
+ request_context = await self._resolve_account_context(account, auth_type)
+ if permission is not None:
+ if request_context is None:
+ raise AuthorizationError('Workspace authorization is unavailable')
+ require_permission(request_context, permission)
+ self._inject_handler_context(f, kwargs, user_email, request_context)
except Exception as e:
- return self.http_status(401, -1, str(e))
+ return self._auth_error_response(e)
elif auth_type == AuthType.API_KEY:
# get API key from Authorization header or X-API-Key header
@@ -101,11 +129,12 @@ class RouterGroup(abc.ABC):
return self.http_status(401, -1, 'No valid API key provided')
try:
- is_valid = await self.ap.apikey_service.verify_api_key(api_key)
- if not is_valid:
- return self.http_status(401, -1, 'Invalid API key')
+ request_context = await self._authenticate_api_key(api_key, auth_type)
+ if permission is not None:
+ require_permission(request_context, permission)
+ self._inject_handler_context(f, kwargs, None, request_context)
except Exception as e:
- return self.http_status(401, -1, str(e))
+ return self._auth_error_response(e)
elif auth_type == AuthType.USER_TOKEN_OR_API_KEY:
# Try API key first (check X-API-Key header)
@@ -114,11 +143,12 @@ class RouterGroup(abc.ABC):
if api_key:
# API key authentication
try:
- is_valid = await self.ap.apikey_service.verify_api_key(api_key)
- if not is_valid:
- return self.http_status(401, -1, 'Invalid API key')
+ request_context = await self._authenticate_api_key(api_key, auth_type)
+ if permission is not None:
+ require_permission(request_context, permission)
+ self._inject_handler_context(f, kwargs, None, request_context)
except Exception as e:
- return self.http_status(401, -1, str(e))
+ return self._auth_error_response(e)
else:
# Try user token authentication (Authorization header)
token = quart.request.headers.get('Authorization', '').replace('Bearer ', '')
@@ -129,35 +159,56 @@ class RouterGroup(abc.ABC):
)
try:
- user_email = await self.ap.user_service.verify_jwt_token(token)
-
- # check if this account exists
- user = await self.ap.user_service.get_user_by_email(user_email)
- if not user:
- return self.http_status(401, -1, 'User not found')
-
- # check if f accepts user_email parameter
- if 'user_email' in f.__code__.co_varnames:
- kwargs['user_email'] = user_email
+ account, user_email = await self._authenticate_account(token)
+ request_context = await self._resolve_account_context(account, auth_type)
+ if permission is not None:
+ if request_context is None:
+ raise AuthorizationError('Workspace authorization is unavailable')
+ require_permission(request_context, permission)
+ self._inject_handler_context(f, kwargs, user_email, request_context)
+ except (AuthorizationError, WorkspaceNotFoundError, MembershipPermissionError) as e:
+ # Authentication succeeded and authorization was
+ # evaluated. Do not reinterpret a denied user token
+ # as an API key, which would mask the stable 403/404.
+ return self._auth_error_response(e)
except Exception:
# If user token fails, maybe it's an API key in Authorization header
try:
- is_valid = await self.ap.apikey_service.verify_api_key(token)
- if not is_valid:
- return self.http_status(401, -1, 'Invalid authentication credentials')
+ request_context = await self._authenticate_api_key(token, auth_type)
+ if permission is not None:
+ require_permission(request_context, permission)
+ self._inject_handler_context(f, kwargs, None, request_context)
except Exception as e:
- return self.http_status(401, -1, str(e))
+ return self._auth_error_response(e)
try:
return await f(*args, **kwargs)
except Exception as e: # 自动 500
- traceback.print_exc()
- # return self.http_status(500, -2, str(e))
- return self.http_status(500, -2, str(e))
+ if isinstance(e, AuthorizationError):
+ return self.http_status(e.status_code, e.error_code, str(e))
+ if isinstance(e, WorkspaceNotFoundError):
+ return self.http_status(404, 'resource_not_found', 'Resource not found')
+ if isinstance(e, MembershipPermissionError):
+ return self.http_status(403, e.code, str(e))
+ if isinstance(e, WorkspaceCollaborationError):
+ return self.http_status(400, e.code, str(e))
+ request_id = self.request_id()
+ logger = getattr(self.ap, 'logger', self.quart_app.logger)
+ logger.error(
+ f'Unhandled HTTP error request_id={request_id} '
+ f'method={quart.request.method} path={quart.request.path}\n{traceback.format_exc()}'
+ )
+ return self.internal_error_response(request_id)
new_f = handler_error
- new_f.__name__ = (self.name + rule).replace('/', '__')
+ # Quart/Flask requires a unique endpoint name even when the same URL
+ # intentionally has separate handlers for different HTTP methods.
+ # Include the method set so CRUD routes can declare distinct
+ # permissions without colliding during application startup.
+ methods = options.get('methods') or ['GET']
+ method_suffix = '__'.join(sorted(str(method).upper() for method in methods))
+ new_f.__name__ = (self.name + rule + '__' + method_suffix).replace('/', '__')
new_f.__doc__ = f.__doc__
self.quart_app.route(rule, **options)(new_f)
@@ -165,6 +216,165 @@ class RouterGroup(abc.ABC):
return decorator
+ async def _authenticate_account(self, token: str) -> tuple[typing.Any, str]:
+ account: typing.Any = None
+ resolver = getattr(self.ap.user_service, 'get_authenticated_account', None)
+ if callable(resolver):
+ resolved = resolver(token)
+ if inspect.isawaitable(resolved):
+ account = await resolved
+
+ if isinstance(account, str) or account is None:
+ user_email = account or await self.ap.user_service.verify_jwt_token(token)
+ account = await self.ap.user_service.get_user_by_email(user_email)
+ if account is None:
+ raise ValueError('User not found')
+ return account, account.user
+
+ async def _resolve_account_context(
+ self,
+ account: typing.Any,
+ auth_type: AuthType,
+ ) -> RequestContext | None:
+ collaboration_service = getattr(self.ap, 'workspace_collaboration_service', None)
+ account_uuid = getattr(account, 'uuid', None)
+ # Compatibility for isolated controller tests that do not wire the tenancy kernel.
+ if collaboration_service is None or not isinstance(account_uuid, str):
+ return None
+
+ requested_workspace_uuid = quart.request.headers.get('X-Workspace-Id')
+ access = await collaboration_service.resolve_account_workspace(account_uuid, requested_workspace_uuid)
+ request_context = RequestContext(
+ instance_uuid=access.execution.instance_uuid,
+ placement_generation=access.execution.placement_generation,
+ request_id=self.request_id(),
+ auth_type=auth_type.value,
+ principal=PrincipalContext(
+ principal_type=PrincipalType.ACCOUNT,
+ account_uuid=account_uuid,
+ ),
+ workspace=WorkspaceContext(
+ workspace_uuid=access.workspace.uuid,
+ membership_uuid=access.membership.uuid,
+ role=access.membership.role,
+ permissions=permissions_for_role(access.membership.role),
+ membership_revision=access.membership.projection_revision,
+ ),
+ )
+ quart.g.request_context = request_context
+ quart.g.workspace_membership = access.membership
+ return request_context
+
+ async def _authenticate_api_key(self, api_key: str, auth_type: AuthType) -> RequestContext:
+ authenticator = getattr(self.ap.apikey_service, 'authenticate_api_key', None)
+ if callable(authenticator):
+ authenticated = authenticator(api_key)
+ if inspect.isawaitable(authenticated):
+ identity = await authenticated
+ if identity is not None:
+ request_context = RequestContext(
+ instance_uuid=identity.instance_uuid,
+ placement_generation=identity.placement_generation,
+ request_id=self.request_id(),
+ auth_type=auth_type.value,
+ principal=PrincipalContext(
+ principal_type=PrincipalType.API_KEY,
+ api_key_uuid=identity.api_key_uuid,
+ ),
+ workspace=WorkspaceContext(
+ workspace_uuid=identity.workspace_uuid,
+ membership_uuid=None,
+ role=None,
+ permissions=identity.permissions,
+ ),
+ )
+ quart.g.request_context = request_context
+ return request_context
+
+ if not await self.ap.apikey_service.verify_api_key(api_key):
+ raise ValueError('Invalid API key')
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ if workspace_service is None:
+ raise ValueError('API key Workspace binding is unavailable')
+ binding = await workspace_service.get_local_execution_binding()
+ request_context = RequestContext(
+ instance_uuid=binding.instance_uuid or constants.instance_id,
+ placement_generation=binding.placement_generation,
+ request_id=self.request_id(),
+ auth_type=auth_type.value,
+ principal=PrincipalContext(
+ principal_type=PrincipalType.API_KEY,
+ api_key_uuid='legacy-oss-api-key',
+ ),
+ workspace=WorkspaceContext(
+ workspace_uuid=binding.workspace_uuid,
+ membership_uuid=None,
+ role=None,
+ permissions=frozenset(item.value for item in Permission),
+ ),
+ )
+ quart.g.request_context = request_context
+ return request_context
+
+ @staticmethod
+ def _inject_handler_context(
+ handler: RouteCallable,
+ kwargs: dict[str, typing.Any],
+ user_email: str | None,
+ request_context: RequestContext | None,
+ ) -> None:
+ parameters = handler.__code__.co_varnames
+ if user_email is not None and 'user_email' in parameters:
+ kwargs['user_email'] = user_email
+ if request_context is not None:
+ if 'request_context' in parameters:
+ kwargs['request_context'] = request_context
+ elif 'ctx' in parameters:
+ kwargs['ctx'] = request_context
+
+ def _auth_error_response(self, error: Exception) -> typing.Any:
+ if isinstance(error, AuthorizationError):
+ return self.http_status(error.status_code, error.error_code, str(error))
+ if isinstance(error, WorkspaceNotFoundError):
+ return self.http_status(404, 'resource_not_found', 'Resource not found')
+ if isinstance(error, MembershipPermissionError):
+ return self.http_status(403, error.code, str(error))
+ request_id = self.request_id()
+ logger = getattr(self.ap, 'logger', self.quart_app.logger)
+ logger.warning(f'Authentication failed request_id={request_id} error_type={type(error).__name__}: {error}')
+ return self.http_status(
+ 401,
+ 'invalid_authentication',
+ 'Invalid authentication credentials',
+ )
+
+ def request_id(self) -> str:
+ """Return one stable request ID for authentication, logs, and errors."""
+
+ request_context = getattr(quart.g, 'request_context', None)
+ request_id = getattr(request_context, 'request_id', None) or getattr(quart.g, 'request_id', None)
+ if not request_id:
+ candidate = str(quart.request.headers.get('X-Request-Id') or '').strip()
+ if not candidate or len(candidate) > 128 or any(ord(char) < 32 for char in candidate):
+ candidate = str(uuid.uuid4())
+ request_id = candidate
+ quart.g.request_id = request_id
+ return str(request_id)
+
+ def internal_error_response(self, request_id: str | None = None) -> typing.Tuple[quart.Response, int]:
+ """Return a stable 500 response without exposing the underlying exception."""
+
+ resolved_request_id = request_id or self.request_id()
+ response = quart.jsonify(
+ {
+ 'code': 'internal_error',
+ 'msg': 'Internal server error',
+ 'request_id': resolved_request_id,
+ }
+ )
+ response.headers['X-Request-Id'] = resolved_request_id
+ return response, 500
+
def success(self, data: typing.Any = None) -> quart.Response:
"""Return a 200 response"""
return quart.jsonify(
@@ -175,7 +385,7 @@ class RouterGroup(abc.ABC):
}
)
- def fail(self, code: int, msg: str) -> quart.Response:
+ def fail(self, code: int | str, msg: str) -> quart.Response:
"""Return an error response"""
return quart.jsonify(
@@ -185,6 +395,6 @@ class RouterGroup(abc.ABC):
}
)
- def http_status(self, status: int, code: int, msg: str) -> typing.Tuple[quart.Response, int]:
+ def http_status(self, status: int, code: int | str, msg: str) -> typing.Tuple[quart.Response, int]:
"""返回一个指定状态码的响应"""
return (self.fail(code, msg), status)
diff --git a/src/langbot/pkg/api/http/controller/groups/apikeys.py b/src/langbot/pkg/api/http/controller/groups/apikeys.py
index f53728bf0..6c8b4a656 100644
--- a/src/langbot/pkg/api/http/controller/groups/apikeys.py
+++ b/src/langbot/pkg/api/http/controller/groups/apikeys.py
@@ -1,43 +1,66 @@
+from __future__ import annotations
+
+import datetime
+
import quart
+from ...authz import Permission
+from ...context import RequestContext
from .. import group
@group.group_class('apikeys', '/api/v1/apikeys')
class ApiKeysRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET', 'POST'])
- async def _() -> str:
- if quart.request.method == 'GET':
- keys = await self.ap.apikey_service.get_api_keys()
- return self.success(data={'keys': keys})
- elif quart.request.method == 'POST':
- json_data = await quart.request.json
- name = json_data.get('name', '')
- description = json_data.get('description', '')
+ @self.route('', methods=['GET'], permission=Permission.API_KEY_MANAGE)
+ async def _(request_context: RequestContext) -> str:
+ keys = await self.ap.apikey_service.get_api_keys(request_context)
+ return self.success(data={'keys': keys})
- if not name:
- return self.http_status(400, -1, 'Name is required')
+ @self.route('', methods=['POST'], permission=Permission.API_KEY_MANAGE)
+ async def _(request_context: RequestContext) -> str:
+ json_data = await quart.request.json
+ expires_at = json_data.get('expires_at')
+ parsed_expiry = None
+ if expires_at:
+ try:
+ parsed_expiry = datetime.datetime.fromisoformat(str(expires_at).replace('Z', '+00:00'))
+ except ValueError:
+ return self.http_status(400, 'invalid_expiry', 'Invalid API key expiry')
+ try:
+ key = await self.ap.apikey_service.create_api_key(
+ request_context,
+ json_data.get('name', ''),
+ json_data.get('description', ''),
+ scopes=json_data.get('scopes'),
+ expires_at=parsed_expiry,
+ )
+ except ValueError as error:
+ return self.http_status(400, 'invalid_api_key', str(error))
+ return self.success(data={'key': key})
- key = await self.ap.apikey_service.create_api_key(name, description)
- return self.success(data={'key': key})
+ @self.route('/', methods=['GET'], permission=Permission.API_KEY_MANAGE)
+ async def _(key_id: int, request_context: RequestContext) -> str:
+ key = await self.ap.apikey_service.get_api_key(request_context, key_id)
+ if key is None:
+ return self.http_status(404, 'resource_not_found', 'API key not found')
+ return self.success(data={'key': key})
- @self.route('/', methods=['GET', 'PUT', 'DELETE'])
- async def _(key_id: int) -> str:
- if quart.request.method == 'GET':
- key = await self.ap.apikey_service.get_api_key(key_id)
- if key is None:
- return self.http_status(404, -1, 'API key not found')
- return self.success(data={'key': key})
+ @self.route('/', methods=['PUT'], permission=Permission.API_KEY_MANAGE)
+ async def _(key_id: int, request_context: RequestContext) -> str:
+ json_data = await quart.request.json
+ try:
+ await self.ap.apikey_service.update_api_key(
+ request_context,
+ key_id,
+ json_data.get('name'),
+ json_data.get('description'),
+ )
+ except ValueError as error:
+ return self.http_status(400, 'invalid_api_key', str(error))
+ return self.success()
- elif quart.request.method == 'PUT':
- json_data = await quart.request.json
- name = json_data.get('name')
- description = json_data.get('description')
-
- await self.ap.apikey_service.update_api_key(key_id, name, description)
- return self.success()
-
- elif quart.request.method == 'DELETE':
- await self.ap.apikey_service.delete_api_key(key_id)
- return self.success()
+ @self.route('/', methods=['DELETE'], permission=Permission.API_KEY_MANAGE)
+ async def _(key_id: int, request_context: RequestContext) -> str:
+ await self.ap.apikey_service.delete_api_key(request_context, key_id)
+ return self.success()
diff --git a/src/langbot/pkg/api/http/controller/groups/box.py b/src/langbot/pkg/api/http/controller/groups/box.py
index d8c961e7a..ba10c4b19 100644
--- a/src/langbot/pkg/api/http/controller/groups/box.py
+++ b/src/langbot/pkg/api/http/controller/groups/box.py
@@ -2,6 +2,8 @@ from __future__ import annotations
from langbot.pkg.utils import constants
+from ...authz import Permission
+from ...context import RequestContext
from .. import group
from .box_visibility import should_hide_box_runtime_status
@@ -9,18 +11,33 @@ from .box_visibility import should_hide_box_runtime_status
@group.group_class('box', '/api/v1/box')
class BoxRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('/status', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
- status = await self.ap.box_service.get_status()
+ @self.route(
+ '/status',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ status = await self.ap.box_service.get_status(request_context)
status['hidden'] = should_hide_box_runtime_status(constants.edition, status.get('enabled'))
return self.success(data=status)
- @self.route('/sessions', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
- sessions = await self.ap.box_service.get_sessions()
+ @self.route(
+ '/sessions',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.AUDIT_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ sessions = await self.ap.box_service.get_sessions(request_context)
return self.success(data=sessions)
- @self.route('/errors', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
- errors = self.ap.box_service.get_recent_errors()
+ @self.route(
+ '/errors',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.AUDIT_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ errors = self.ap.box_service.get_recent_errors(request_context)
return self.success(data=errors)
diff --git a/src/langbot/pkg/api/http/controller/groups/extensions.py b/src/langbot/pkg/api/http/controller/groups/extensions.py
index ac8463c90..690e6f0a1 100644
--- a/src/langbot/pkg/api/http/controller/groups/extensions.py
+++ b/src/langbot/pkg/api/http/controller/groups/extensions.py
@@ -3,6 +3,9 @@ from __future__ import annotations
import asyncio
import quart
+from ...authz import Permission
+from ...context import RequestContext
+from ...service.secrets import redact_secrets
from .. import group
@@ -11,12 +14,19 @@ class ExtensionsRouterGroup(group.RouterGroup):
"""Unified API for installed extensions (plugins, MCP servers, skills)."""
async def initialize(self) -> None:
- @self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> quart.Response:
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> quart.Response:
+ if self.ap.plugin_connector.is_enable_plugin:
+ await self.ap.plugin_connector.require_workspace_context(request_context)
plugins, mcp_servers, skills = await asyncio.gather(
self.ap.plugin_connector.list_plugins(),
- self.ap.mcp_service.get_mcp_servers(contain_runtime_info=True),
- self.ap.skill_service.list_skills(),
+ self.ap.mcp_service.get_mcp_servers(request_context, contain_runtime_info=True),
+ self.ap.skill_service.list_skills(request_context),
return_exceptions=True,
)
@@ -39,7 +49,7 @@ class ExtensionsRouterGroup(group.RouterGroup):
extensions: list[dict] = []
if isinstance(plugins, list):
for plugin in plugins:
- extensions.append({'type': 'plugin', 'plugin': plugin})
+ extensions.append({'type': 'plugin', 'plugin': redact_secrets(plugin)})
if isinstance(mcp_servers, list):
for server in mcp_servers:
extensions.append({'type': 'mcp', 'server': server})
diff --git a/src/langbot/pkg/api/http/controller/groups/files.py b/src/langbot/pkg/api/http/controller/groups/files.py
index 439cc57e3..57ba60107 100644
--- a/src/langbot/pkg/api/http/controller/groups/files.py
+++ b/src/langbot/pkg/api/http/controller/groups/files.py
@@ -7,29 +7,48 @@ import asyncio
import quart.datastructures
+from ...authz import Permission
+from ...context import RequestContext
from .. import group
+def _storage_owner(context: RequestContext) -> str:
+ if context.principal.account_uuid:
+ return f'account:{context.principal.account_uuid}'
+ if context.principal.api_key_uuid:
+ return f'api-key:{context.principal.api_key_uuid}'
+ return f'principal:{context.principal.principal_type.value}'
+
+
@group.group_class('files', '/api/v1/files')
class FilesRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
@self.route('/image/', methods=['GET'], auth_type=group.AuthType.NONE)
async def _(image_key: str) -> quart.Response:
- if '..' in image_key or '\\' in image_key:
+ image_bytes = await self.ap.storage_mgr.resolve_public_object(
+ image_key,
+ expected_owner_type='upload_image',
+ )
+ if image_bytes is None:
+ image_bytes = await self.ap.storage_mgr.resolve_public_object(
+ image_key,
+ expected_owner_type='bot_log',
+ )
+ if image_bytes is None:
return quart.Response(status=404)
-
- if not await self.ap.storage_mgr.storage_provider.exists(image_key):
- return quart.Response(status=404)
-
- image_bytes = await self.ap.storage_mgr.storage_provider.load(image_key)
mime_type = mimetypes.guess_type(image_key)[0]
if mime_type is None:
mime_type = 'image/jpeg'
return quart.Response(image_bytes, mimetype=mime_type)
- @self.route('/images', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def upload_image() -> quart.Response:
+ @self.route(
+ '/images',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def upload_image(request_context: RequestContext) -> quart.Response:
request = quart.request
# Check file size limit before reading the file
@@ -66,18 +85,29 @@ class FilesRouterGroup(group.RouterGroup):
if '/' in file_name or '\\' in file_name:
return self.fail(400, 'File name contains invalid characters')
- file_key = file_name + '_' + str(uuid.uuid4())[:8] + '.' + extension
+ logical_key = f'{uuid.uuid4()}.{extension}'
# save file to storage
- await self.ap.storage_mgr.storage_provider.save(file_key, file_bytes)
+ file_key = await self.ap.storage_mgr.save_scoped(
+ request_context,
+ owner_type='upload_image',
+ owner=_storage_owner(request_context),
+ key=logical_key,
+ value=file_bytes,
+ )
return self.success(
data={
'file_key': file_key,
}
)
- @self.route('/documents', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def upload_document() -> quart.Response:
+ @self.route(
+ '/documents',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def upload_document(request_context: RequestContext) -> quart.Response:
request = quart.request
# Check file size limit before reading the file
@@ -110,12 +140,18 @@ class FilesRouterGroup(group.RouterGroup):
if '/' in file_name or '\\' in file_name:
return self.fail(400, 'File name contains invalid characters')
- file_key = file_name + '_' + str(uuid.uuid4())[:8]
+ logical_key = str(uuid.uuid4())
if extension:
- file_key += '.' + extension
+ logical_key += '.' + extension
# save file to storage
- await self.ap.storage_mgr.storage_provider.save(file_key, file_bytes)
+ file_key = await self.ap.storage_mgr.save_scoped(
+ request_context,
+ owner_type='upload_document',
+ owner=_storage_owner(request_context),
+ key=logical_key,
+ value=file_bytes,
+ )
return self.success(
data={
'file_id': file_key,
diff --git a/src/langbot/pkg/api/http/controller/groups/knowledge/base.py b/src/langbot/pkg/api/http/controller/groups/knowledge/base.py
index 4f9bb5b4f..87dc2ac30 100644
--- a/src/langbot/pkg/api/http/controller/groups/knowledge/base.py
+++ b/src/langbot/pkg/api/http/controller/groups/knowledge/base.py
@@ -1,100 +1,146 @@
import quart
+
+from ....authz import Permission, has_permission
+from ....context import RequestContext
from ... import group
@group.group_class('knowledge_base', '/api/v1/knowledge/bases')
class KnowledgeBaseRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['POST', 'GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def handle_knowledge_bases() -> quart.Response:
- if quart.request.method == 'GET':
- knowledge_bases = await self.ap.knowledge_service.get_knowledge_bases()
- return self.success(data={'bases': knowledge_bases})
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def handle_knowledge_bases(request_context: RequestContext) -> quart.Response:
+ knowledge_bases = await self.ap.knowledge_service.get_knowledge_bases(
+ request_context,
+ include_secret=has_permission(request_context, Permission.RESOURCE_MANAGE),
+ )
+ return self.success(data={'bases': knowledge_bases})
- elif quart.request.method == 'POST':
- json_data = await quart.request.json
- try:
- knowledge_base_uuid = await self.ap.knowledge_service.create_knowledge_base(json_data)
- except ValueError as e:
- return self.http_status(400, -1, str(e))
- return self.success(data={'uuid': knowledge_base_uuid})
-
- return self.http_status(405, -1, 'Method not allowed')
+ @self.route(
+ '',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def create_knowledge_base(request_context: RequestContext) -> quart.Response:
+ json_data = await quart.request.json
+ try:
+ knowledge_base_uuid = await self.ap.knowledge_service.create_knowledge_base(
+ request_context,
+ json_data,
+ )
+ except ValueError as e:
+ return self.http_status(400, -1, str(e))
+ return self.success(data={'uuid': knowledge_base_uuid})
@self.route(
'/',
- methods=['GET', 'DELETE', 'PUT'],
+ methods=['GET'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def handle_specific_knowledge_base(knowledge_base_uuid: str) -> quart.Response:
- if quart.request.method == 'GET':
- knowledge_base = await self.ap.knowledge_service.get_knowledge_base(knowledge_base_uuid)
+ async def get_specific_knowledge_base(
+ knowledge_base_uuid: str,
+ request_context: RequestContext,
+ ) -> quart.Response:
+ knowledge_base = await self.ap.knowledge_service.get_knowledge_base(
+ request_context,
+ knowledge_base_uuid,
+ include_secret=has_permission(request_context, Permission.RESOURCE_MANAGE),
+ )
+ if knowledge_base is None:
+ return self.http_status(404, 'resource_not_found', 'knowledge base not found')
+ return self.success(data={'base': knowledge_base})
- if knowledge_base is None:
- return self.http_status(404, -1, 'knowledge base not found')
-
- return self.success(
- data={
- 'base': knowledge_base,
- }
- )
-
- elif quart.request.method == 'PUT':
+ @self.route(
+ '/',
+ methods=['DELETE', 'PUT'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def mutate_specific_knowledge_base(
+ knowledge_base_uuid: str,
+ request_context: RequestContext,
+ ) -> quart.Response:
+ if quart.request.method == 'PUT':
json_data = await quart.request.json
- await self.ap.knowledge_service.update_knowledge_base(knowledge_base_uuid, json_data)
+ await self.ap.knowledge_service.update_knowledge_base(
+ request_context,
+ knowledge_base_uuid,
+ json_data,
+ )
return self.success(data={'uuid': knowledge_base_uuid})
-
- elif quart.request.method == 'DELETE':
- await self.ap.knowledge_service.delete_knowledge_base(knowledge_base_uuid)
- return self.success({})
+ await self.ap.knowledge_service.delete_knowledge_base(request_context, knowledge_base_uuid)
+ return self.success({})
@self.route(
'//files',
- methods=['GET', 'POST'],
+ methods=['GET'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def get_knowledge_base_files(knowledge_base_uuid: str) -> str:
- if quart.request.method == 'GET':
- files = await self.ap.knowledge_service.get_files_by_knowledge_base(knowledge_base_uuid)
- return self.success(
- data={
- 'files': files,
- }
- )
+ async def get_knowledge_base_files(
+ knowledge_base_uuid: str,
+ request_context: RequestContext,
+ ) -> str:
+ files = await self.ap.knowledge_service.get_files_by_knowledge_base(
+ request_context,
+ knowledge_base_uuid,
+ )
+ return self.success(data={'files': files})
- elif quart.request.method == 'POST':
- json_data = await quart.request.json
- file_id = json_data.get('file_id')
- if not file_id:
- return self.http_status(400, -1, 'File ID is required')
-
- parser_plugin_id = json_data.get('parser_plugin_id')
-
- # 调用服务层方法将文件与知识库关联
- task_id = await self.ap.knowledge_service.store_file(
- knowledge_base_uuid, file_id, parser_plugin_id=parser_plugin_id
- )
- return self.success(
- {
- 'task_id': task_id,
- }
- )
+ @self.route(
+ '//files',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def add_knowledge_base_file(
+ knowledge_base_uuid: str,
+ request_context: RequestContext,
+ ) -> str:
+ json_data = await quart.request.json
+ file_id = json_data.get('file_id')
+ if not file_id:
+ return self.http_status(400, -1, 'File ID is required')
+ parser_plugin_id = json_data.get('parser_plugin_id')
+ task_id = await self.ap.knowledge_service.store_file(
+ request_context,
+ knowledge_base_uuid,
+ file_id,
+ parser_plugin_id=parser_plugin_id,
+ )
+ return self.success({'task_id': task_id})
@self.route(
'//files/',
methods=['DELETE'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
)
- async def delete_specific_file_in_kb(file_id: str, knowledge_base_uuid: str) -> str:
- await self.ap.knowledge_service.delete_file(knowledge_base_uuid, file_id)
+ async def delete_specific_file_in_kb(
+ file_id: str,
+ knowledge_base_uuid: str,
+ request_context: RequestContext,
+ ) -> str:
+ await self.ap.knowledge_service.delete_file(request_context, knowledge_base_uuid, file_id)
return self.success({})
@self.route(
'//retrieve',
methods=['POST'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def retrieve_knowledge_base(knowledge_base_uuid: str) -> str:
+ async def retrieve_knowledge_base(
+ knowledge_base_uuid: str,
+ request_context: RequestContext,
+ ) -> str:
json_data = await quart.request.json
query = json_data.get('query')
@@ -104,6 +150,9 @@ class KnowledgeBaseRouterGroup(group.RouterGroup):
# Extract retrieval_settings to allow dynamic control over Knowledge Engine behavior (e.g. top_k, filters)
retrieval_settings = json_data.get('retrieval_settings', {})
results = await self.ap.knowledge_service.retrieve_knowledge_base(
- knowledge_base_uuid, query, retrieval_settings
+ request_context,
+ knowledge_base_uuid,
+ query,
+ retrieval_settings,
)
return self.success(data={'results': results})
diff --git a/src/langbot/pkg/api/http/controller/groups/knowledge/engines.py b/src/langbot/pkg/api/http/controller/groups/knowledge/engines.py
index 28f0710e8..02d047f15 100644
--- a/src/langbot/pkg/api/http/controller/groups/knowledge/engines.py
+++ b/src/langbot/pkg/api/http/controller/groups/knowledge/engines.py
@@ -1,25 +1,39 @@
import quart
from urllib.parse import unquote
+
+from ....authz import Permission
+from ....context import RequestContext
from ... import group
@group.group_class('knowledge_engines', '/api/v1/knowledge/engines')
class KnowledgeEnginesRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def list_knowledge_engines() -> quart.Response:
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def list_knowledge_engines(request_context: RequestContext) -> quart.Response:
"""List all available Knowledge Engines from plugins.
Returns a list of Knowledge Engines with their capabilities and configuration schemas.
This is used by the frontend to render the knowledge base creation wizard.
"""
- engines = await self.ap.knowledge_service.list_knowledge_engines()
+ engines = await self.ap.knowledge_service.list_knowledge_engines(request_context)
return self.success(data={'engines': engines})
@self.route(
- '//creation-schema', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY
+ '//creation-schema',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def get_engine_creation_schema(plugin_id: str) -> quart.Response:
+ async def get_engine_creation_schema(
+ plugin_id: str,
+ request_context: RequestContext,
+ ) -> quart.Response:
"""Get creation settings schema for a specific Knowledge Engine.
plugin_id is in 'author/name' format, captured via converter.
@@ -27,13 +41,19 @@ class KnowledgeEnginesRouterGroup(group.RouterGroup):
plugin_id = unquote(plugin_id)
if '/' not in plugin_id:
return self.http_status(400, -1, 'Invalid plugin_id format. Expected author/name.')
- schema = await self.ap.knowledge_service.get_engine_creation_schema(plugin_id)
+ schema = await self.ap.knowledge_service.get_engine_creation_schema(request_context, plugin_id)
return self.success(data={'schema': schema})
@self.route(
- '//retrieval-schema', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY
+ '//retrieval-schema',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def get_engine_retrieval_schema(plugin_id: str) -> quart.Response:
+ async def get_engine_retrieval_schema(
+ plugin_id: str,
+ request_context: RequestContext,
+ ) -> quart.Response:
"""Get retrieval settings schema for a specific Knowledge Engine.
plugin_id is in 'author/name' format, captured via converter.
@@ -41,5 +61,5 @@ class KnowledgeEnginesRouterGroup(group.RouterGroup):
plugin_id = unquote(plugin_id)
if '/' not in plugin_id:
return self.http_status(400, -1, 'Invalid plugin_id format. Expected author/name.')
- schema = await self.ap.knowledge_service.get_engine_retrieval_schema(plugin_id)
+ schema = await self.ap.knowledge_service.get_engine_retrieval_schema(request_context, plugin_id)
return self.success(data={'schema': schema})
diff --git a/src/langbot/pkg/api/http/controller/groups/knowledge/migration.py b/src/langbot/pkg/api/http/controller/groups/knowledge/migration.py
index 2db835d89..f93c759d6 100644
--- a/src/langbot/pkg/api/http/controller/groups/knowledge/migration.py
+++ b/src/langbot/pkg/api/http/controller/groups/knowledge/migration.py
@@ -6,8 +6,11 @@ import quart
import sqlalchemy
from ... import group
+from ....authz import Permission
+from ....context import ExecutionContext, RequestContext
from ......core import taskmgr
from ......entity.persistence import metadata as persistence_metadata
+from ......workspace.errors import WorkspaceError, WorkspaceNotFoundError
from langbot_plugin.runtime.plugin.mgr import PluginInstallSource
LANGRAG_PLUGIN_AUTHOR = 'langbot-team'
@@ -34,21 +37,49 @@ EXTERNAL_PLUGIN_CREATION_FIELDS: dict[str, set[str] | None] = {
@group.group_class('knowledge/migration', '/api/v1/knowledge/migration')
class KnowledgeMigrationRouterGroup(group.RouterGroup):
- async def _get_migration_flag(self) -> bool:
+ async def _require_local_migration_context(
+ self,
+ execution_context: ExecutionContext,
+ ) -> ExecutionContext:
+ """Fence legacy-table migration to the OSS singleton Workspace.
+
+ The backup tables predate Workspace scoping and are deliberately
+ instance-global. A cloud projection must therefore never be allowed
+ to inspect or restore them, even when it has a valid execution lease.
+ """
+ try:
+ binding = await self.ap.workspace_service.get_local_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ except WorkspaceNotFoundError:
+ raise
+ except WorkspaceError as exc:
+ raise WorkspaceNotFoundError('RAG migration is unavailable') from exc
+
+ if binding.instance_uuid != execution_context.instance_uuid:
+ raise WorkspaceNotFoundError('RAG migration is unavailable')
+ return ExecutionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ )
+
+ async def _get_migration_flag(self, execution_context: ExecutionContext) -> bool:
"""Check if rag_plugin_migration_needed flag is set."""
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_metadata.Metadata).where(
- persistence_metadata.Metadata.key == 'rag_plugin_migration_needed'
- )
+ sqlalchemy.select(persistence_metadata.WorkspaceMetadata.value)
+ .where(persistence_metadata.WorkspaceMetadata.workspace_uuid == execution_context.workspace_uuid)
+ .where(persistence_metadata.WorkspaceMetadata.key == 'rag_plugin_migration_needed')
)
- row = result.first()
- return row is not None and row.value == 'true'
+ return result.scalar_one_or_none() == 'true'
- async def _set_migration_flag(self, value: str):
+ async def _set_migration_flag(self, execution_context: ExecutionContext, value: str):
"""Set rag_plugin_migration_needed flag."""
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(persistence_metadata.Metadata)
- .where(persistence_metadata.Metadata.key == 'rag_plugin_migration_needed')
+ sqlalchemy.update(persistence_metadata.WorkspaceMetadata)
+ .where(persistence_metadata.WorkspaceMetadata.workspace_uuid == execution_context.workspace_uuid)
+ .where(persistence_metadata.WorkspaceMetadata.key == 'rag_plugin_migration_needed')
.values(value=value)
)
@@ -70,7 +101,11 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
return result.first() is not None
async def _install_plugin_from_marketplace(
- self, plugin_id: str, task_context: taskmgr.TaskContext, space_url: str
+ self,
+ execution_context: ExecutionContext,
+ plugin_id: str,
+ task_context: taskmgr.TaskContext,
+ space_url: str,
) -> None:
"""Install a single plugin from the marketplace."""
p_author, p_name = plugin_id.split('/', 1)
@@ -85,6 +120,7 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
if not p_version:
raise Exception(f'Could not determine latest version for {plugin_id}')
+ await self.ap.plugin_connector.require_workspace_context(execution_context)
await self.ap.plugin_connector.install_plugin(
PluginInstallSource.MARKETPLACE,
{
@@ -96,8 +132,15 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
)
self.ap.logger.info(f'RAG migration: plugin {plugin_id} install request sent.')
- async def _execute_rag_migration(self, task_context: taskmgr.TaskContext, install_plugin: bool = True):
+ async def _execute_rag_migration(
+ self,
+ execution_context: ExecutionContext,
+ task_context: taskmgr.TaskContext,
+ install_plugin: bool = True,
+ ):
"""Execute RAG migration: install required plugins and restore backup data."""
+ execution_context = await self._require_local_migration_context(execution_context)
+ execution_context = await self.ap.plugin_connector.require_workspace_context(execution_context)
warnings = []
# Collect all plugins we need: LangRAG (always) + connector plugins (from external KBs)
@@ -127,7 +170,14 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
for plugin_id in needed_plugins:
try:
- await self._install_plugin_from_marketplace(plugin_id, task_context, space_url)
+ await self._install_plugin_from_marketplace(
+ execution_context,
+ plugin_id,
+ task_context,
+ space_url,
+ )
+ except WorkspaceNotFoundError:
+ raise
except Exception as e:
self.ap.logger.warning(f'RAG migration: plugin {plugin_id} install returned: {e}')
task_context.trace(f'Plugin install note ({plugin_id}): {e}')
@@ -141,8 +191,11 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
engine_id_set: set[str] = set()
for i in range(max_retries):
try:
+ await self.ap.plugin_connector.require_workspace_context(execution_context)
engines = await self.ap.plugin_connector.list_knowledge_engines()
engine_id_set = {e.get('plugin_id') for e in engines}
+ except WorkspaceNotFoundError:
+ raise
except Exception:
pass
if all(pid in engine_id_set for pid in needed_plugins):
@@ -158,8 +211,11 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
await asyncio.sleep(2)
else:
try:
+ await self.ap.plugin_connector.require_workspace_context(execution_context)
engines = await self.ap.plugin_connector.list_knowledge_engines()
engine_id_set = {e.get('plugin_id') for e in engines}
+ except WorkspaceNotFoundError:
+ raise
except Exception:
engine_id_set = set()
@@ -189,12 +245,13 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
await self.ap.persistence_mgr.execute_async(
sqlalchemy.text(
'INSERT INTO knowledge_bases '
- '(uuid, name, description, emoji, created_at, updated_at, '
+ '(uuid, workspace_uuid, name, description, emoji, created_at, updated_at, '
'knowledge_engine_plugin_id, collection_id, creation_settings, retrieval_settings) '
- 'VALUES (:uuid, :name, :description, :emoji, :created_at, :updated_at, '
+ 'VALUES (:uuid, :workspace_uuid, :name, :description, :emoji, :created_at, :updated_at, '
':plugin_id, :collection_id, :creation_settings, :retrieval_settings);'
).bindparams(
uuid=kb_uuid,
+ workspace_uuid=execution_context.workspace_uuid,
name=name,
description=description,
emoji=emoji,
@@ -207,6 +264,7 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
)
)
+ await self.ap.plugin_connector.require_workspace_context(execution_context)
try:
config = {'embedding_model_uuid': embedding_model_uuid}
await self.ap.plugin_connector.rag_on_kb_create(LANGRAG_PLUGIN_ID, kb_uuid, config)
@@ -268,12 +326,13 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
await self.ap.persistence_mgr.execute_async(
sqlalchemy.text(
'INSERT INTO knowledge_bases '
- '(uuid, name, description, emoji, created_at, updated_at, '
+ '(uuid, workspace_uuid, name, description, emoji, created_at, updated_at, '
'knowledge_engine_plugin_id, collection_id, creation_settings, retrieval_settings) '
- 'VALUES (:uuid, :name, :description, :emoji, :created_at, :updated_at, '
+ 'VALUES (:uuid, :workspace_uuid, :name, :description, :emoji, :created_at, :updated_at, '
':plugin_id, :collection_id, :creation_settings, :retrieval_settings);'
).bindparams(
uuid=kb_uuid,
+ workspace_uuid=execution_context.workspace_uuid,
name=name,
description=description,
emoji=emoji,
@@ -294,6 +353,7 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
warnings.append(warning)
task_context.trace(warning)
else:
+ await self.ap.plugin_connector.require_workspace_context(execution_context)
try:
await self.ap.plugin_connector.rag_on_kb_create(
external_plugin_id, kb_uuid, creation_settings_dict
@@ -307,16 +367,23 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
await self.ap.rag_mgr.load_knowledge_bases_from_db()
# Step 5: Clear migration flag
- await self._set_migration_flag('false')
+ await self._set_migration_flag(execution_context, 'false')
task_context.trace('RAG migration completed.', action='done')
if warnings:
task_context.trace(f'Completed with {len(warnings)} warning(s).')
async def initialize(self) -> None:
- @self.route('/status', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
- needed = await self._get_migration_flag()
+ @self.route(
+ '/status',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ execution_context = ExecutionContext.from_request(request_context)
+ execution_context = await self._require_local_migration_context(execution_context)
+ needed = await self._get_migration_flag(execution_context)
internal_kb_count = 0
external_kb_count = 0
@@ -342,9 +409,16 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
}
)
- @self.route('/execute', methods=['POST'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
- needed = await self._get_migration_flag()
+ @self.route(
+ '/execute',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
+ execution_context = ExecutionContext.from_request(request_context)
+ execution_context = await self._require_local_migration_context(execution_context)
+ needed = await self._get_migration_flag(execution_context)
if not needed:
return self.http_status(400, -1, 'RAG migration is not needed')
@@ -353,20 +427,34 @@ class KnowledgeMigrationRouterGroup(group.RouterGroup):
ctx = taskmgr.TaskContext.new()
wrapper = self.ap.task_mgr.create_user_task(
- self._execute_rag_migration(task_context=ctx, install_plugin=install_plugin),
+ self._execute_rag_migration(
+ execution_context,
+ task_context=ctx,
+ install_plugin=install_plugin,
+ ),
kind='rag-migration',
name='rag-migration-execute',
label='Migrating knowledge bases to plugin architecture',
context=ctx,
+ instance_uuid=execution_context.instance_uuid,
+ workspace_uuid=execution_context.workspace_uuid,
+ placement_generation=execution_context.placement_generation,
)
return self.success(data={'task_id': wrapper.id})
- @self.route('/dismiss', methods=['POST'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
- needed = await self._get_migration_flag()
+ @self.route(
+ '/dismiss',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
+ execution_context = ExecutionContext.from_request(request_context)
+ execution_context = await self._require_local_migration_context(execution_context)
+ needed = await self._get_migration_flag(execution_context)
if not needed:
return self.http_status(400, -1, 'RAG migration is not needed')
- await self._set_migration_flag('false')
+ await self._set_migration_flag(execution_context, 'false')
return self.success()
diff --git a/src/langbot/pkg/api/http/controller/groups/knowledge/parsers.py b/src/langbot/pkg/api/http/controller/groups/knowledge/parsers.py
index a5e853cb6..495539307 100644
--- a/src/langbot/pkg/api/http/controller/groups/knowledge/parsers.py
+++ b/src/langbot/pkg/api/http/controller/groups/knowledge/parsers.py
@@ -1,16 +1,24 @@
import quart
+
+from ....authz import Permission
+from ....context import RequestContext
from ... import group
@group.group_class('parsers', '/api/v1/knowledge/parsers')
class ParsersRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def list_parsers() -> quart.Response:
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def list_parsers(request_context: RequestContext) -> quart.Response:
"""List all available parsers from plugins.
Optional query parameter `mime_type` to filter parsers by supported MIME type.
"""
mime_type = quart.request.args.get('mime_type')
- parsers = await self.ap.knowledge_service.list_parsers(mime_type)
+ parsers = await self.ap.knowledge_service.list_parsers(request_context, mime_type)
return self.success(data={'parsers': parsers})
diff --git a/src/langbot/pkg/api/http/controller/groups/logs.py b/src/langbot/pkg/api/http/controller/groups/logs.py
index e3bff9db4..7adc14883 100644
--- a/src/langbot/pkg/api/http/controller/groups/logs.py
+++ b/src/langbot/pkg/api/http/controller/groups/logs.py
@@ -3,14 +3,23 @@ from __future__ import annotations
import quart
+from ...authz import Permission
+from ...context import RequestContext
from .. import group
@group.group_class('logs', '/api/v1/logs')
class LogsRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
+ @self.route('', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def _(request_context: RequestContext) -> str:
+ # The process log is instance-global. It is safe to expose only in
+ # the OSS singleton Workspace; SaaS must use Workspace-scoped
+ # observability records instead of leaking another tenant's lines.
+ await self.ap.workspace_service.get_local_execution_binding(
+ request_context.workspace_uuid,
+ expected_generation=request_context.placement_generation,
+ )
start_page_number = int(quart.request.args.get('start_page_number', 0))
start_offset = int(quart.request.args.get('start_offset', 0))
diff --git a/src/langbot/pkg/api/http/controller/groups/monitoring.py b/src/langbot/pkg/api/http/controller/groups/monitoring.py
index 29dedcafb..60f79ad01 100644
--- a/src/langbot/pkg/api/http/controller/groups/monitoring.py
+++ b/src/langbot/pkg/api/http/controller/groups/monitoring.py
@@ -3,6 +3,8 @@ from __future__ import annotations
import datetime
import quart
+from ...authz import Permission
+from ...context import RequestContext
from .. import group
@@ -24,8 +26,8 @@ def parse_iso_datetime(datetime_str: str | None) -> datetime.datetime | None:
@group.group_class('monitoring', '/api/v1/monitoring')
class MonitoringRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('/overview', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_overview() -> str:
+ @self.route('/overview', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_overview(request_context: RequestContext) -> str:
"""Get overview metrics"""
# Parse query parameters
bot_ids = quart.request.args.getlist('botId')
@@ -38,6 +40,7 @@ class MonitoringRouterGroup(group.RouterGroup):
end_time = parse_iso_datetime(end_time_str)
metrics = await self.ap.monitoring_service.get_overview_metrics(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -46,8 +49,8 @@ class MonitoringRouterGroup(group.RouterGroup):
return self.success(data=metrics)
- @self.route('/token-statistics', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_token_statistics() -> str:
+ @self.route('/token-statistics', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_token_statistics(request_context: RequestContext) -> str:
"""Get detailed token usage statistics (summary, per-model, timeseries)."""
bot_ids = quart.request.args.getlist('botId')
pipeline_ids = quart.request.args.getlist('pipelineId')
@@ -61,6 +64,7 @@ class MonitoringRouterGroup(group.RouterGroup):
end_time = parse_iso_datetime(end_time_str)
stats = await self.ap.monitoring_service.get_token_statistics(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -70,8 +74,8 @@ class MonitoringRouterGroup(group.RouterGroup):
return self.success(data=stats)
- @self.route('/messages', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_messages() -> str:
+ @self.route('/messages', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_messages(request_context: RequestContext) -> str:
"""Get message logs"""
# Parse query parameters
bot_ids = quart.request.args.getlist('botId')
@@ -87,6 +91,7 @@ class MonitoringRouterGroup(group.RouterGroup):
end_time = parse_iso_datetime(end_time_str)
messages, total = await self.ap.monitoring_service.get_messages(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
session_ids=session_ids if session_ids else None,
@@ -105,8 +110,8 @@ class MonitoringRouterGroup(group.RouterGroup):
}
)
- @self.route('/llm-calls', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_llm_calls() -> str:
+ @self.route('/llm-calls', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_llm_calls(request_context: RequestContext) -> str:
"""Get LLM call records"""
# Parse query parameters
bot_ids = quart.request.args.getlist('botId')
@@ -121,6 +126,7 @@ class MonitoringRouterGroup(group.RouterGroup):
end_time = parse_iso_datetime(end_time_str)
llm_calls, total = await self.ap.monitoring_service.get_llm_calls(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -138,8 +144,8 @@ class MonitoringRouterGroup(group.RouterGroup):
}
)
- @self.route('/tool-calls', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_tool_calls() -> str:
+ @self.route('/tool-calls', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_tool_calls(request_context: RequestContext) -> str:
"""Get tool call records"""
bot_ids = quart.request.args.getlist('botId')
pipeline_ids = quart.request.args.getlist('pipelineId')
@@ -153,6 +159,7 @@ class MonitoringRouterGroup(group.RouterGroup):
end_time = parse_iso_datetime(end_time_str)
tool_calls, total = await self.ap.monitoring_service.get_tool_calls(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
session_ids=session_ids if session_ids else None,
@@ -171,8 +178,8 @@ class MonitoringRouterGroup(group.RouterGroup):
}
)
- @self.route('/embedding-calls', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_embedding_calls() -> str:
+ @self.route('/embedding-calls', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_embedding_calls(request_context: RequestContext) -> str:
"""Get embedding call records"""
# Parse query parameters
start_time_str = quart.request.args.get('startTime')
@@ -186,6 +193,7 @@ class MonitoringRouterGroup(group.RouterGroup):
end_time = parse_iso_datetime(end_time_str)
embedding_calls, total = await self.ap.monitoring_service.get_embedding_calls(
+ request_context,
start_time=start_time,
end_time=end_time,
knowledge_base_id=knowledge_base_id if knowledge_base_id else None,
@@ -202,8 +210,8 @@ class MonitoringRouterGroup(group.RouterGroup):
}
)
- @self.route('/sessions', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_sessions() -> str:
+ @self.route('/sessions', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_sessions(request_context: RequestContext) -> str:
"""Get session information"""
# Parse query parameters
bot_ids = quart.request.args.getlist('botId')
@@ -224,6 +232,7 @@ class MonitoringRouterGroup(group.RouterGroup):
is_active = is_active_str.lower() == 'true'
sessions, total = await self.ap.monitoring_service.get_sessions(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -242,8 +251,8 @@ class MonitoringRouterGroup(group.RouterGroup):
}
)
- @self.route('/errors', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_errors() -> str:
+ @self.route('/errors', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_errors(request_context: RequestContext) -> str:
"""Get error logs"""
# Parse query parameters
bot_ids = quart.request.args.getlist('botId')
@@ -258,6 +267,7 @@ class MonitoringRouterGroup(group.RouterGroup):
end_time = parse_iso_datetime(end_time_str)
errors, total = await self.ap.monitoring_service.get_errors(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -275,8 +285,8 @@ class MonitoringRouterGroup(group.RouterGroup):
}
)
- @self.route('/data', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_all_data() -> str:
+ @self.route('/data', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_all_data(request_context: RequestContext) -> str:
"""Get all monitoring data in a single request"""
# Parse query parameters
bot_ids = quart.request.args.getlist('botId')
@@ -291,6 +301,7 @@ class MonitoringRouterGroup(group.RouterGroup):
# Get overview metrics
overview = await self.ap.monitoring_service.get_overview_metrics(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -299,6 +310,7 @@ class MonitoringRouterGroup(group.RouterGroup):
# Get messages
messages, messages_total = await self.ap.monitoring_service.get_messages(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -309,6 +321,7 @@ class MonitoringRouterGroup(group.RouterGroup):
# Get LLM calls
llm_calls, llm_calls_total = await self.ap.monitoring_service.get_llm_calls(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -319,6 +332,7 @@ class MonitoringRouterGroup(group.RouterGroup):
# Get tool calls
tool_calls, tool_calls_total = await self.ap.monitoring_service.get_tool_calls(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -329,6 +343,7 @@ class MonitoringRouterGroup(group.RouterGroup):
# Get sessions
sessions, sessions_total = await self.ap.monitoring_service.get_sessions(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -340,6 +355,7 @@ class MonitoringRouterGroup(group.RouterGroup):
# Get errors
errors, errors_total = await self.ap.monitoring_service.get_errors(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -350,6 +366,7 @@ class MonitoringRouterGroup(group.RouterGroup):
# Get embedding calls
embedding_calls, embedding_calls_total = await self.ap.monitoring_service.get_embedding_calls(
+ request_context,
start_time=start_time,
end_time=end_time,
limit=limit,
@@ -376,27 +393,27 @@ class MonitoringRouterGroup(group.RouterGroup):
}
)
- @self.route('/sessions//analysis', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_session_analysis(session_id: str) -> str:
+ @self.route('/sessions//analysis', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_session_analysis(session_id: str, request_context: RequestContext) -> str:
"""Get detailed analysis for a specific session"""
- analysis = await self.ap.monitoring_service.get_session_analysis(session_id)
+ analysis = await self.ap.monitoring_service.get_session_analysis(request_context, session_id)
# Always return success with the analysis data
# The frontend will handle the 'found: false' case
return self.success(data=analysis)
- @self.route('/messages//details', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_message_details(message_id: str) -> str:
+ @self.route('/messages//details', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_message_details(message_id: str, request_context: RequestContext) -> str:
"""Get detailed information for a specific message"""
- details = await self.ap.monitoring_service.get_message_details(message_id)
+ details = await self.ap.monitoring_service.get_message_details(request_context, message_id)
if not details.get('found'):
return self.error(message=f'Message {message_id} not found', code=404)
return self.success(data=details)
- @self.route('/export', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def export_data() -> tuple[str, int]:
+ @self.route('/export', methods=['GET'], permission=Permission.DATA_EXPORT)
+ async def export_data(request_context: RequestContext) -> tuple[str, int]:
"""Export monitoring data as CSV"""
# Parse query parameters
export_type = quart.request.args.get('type', 'messages')
@@ -413,6 +430,7 @@ class MonitoringRouterGroup(group.RouterGroup):
# Get data based on export type
if export_type == 'messages':
data = await self.ap.monitoring_service.export_messages(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -437,6 +455,7 @@ class MonitoringRouterGroup(group.RouterGroup):
]
elif export_type == 'llm-calls':
data = await self.ap.monitoring_service.export_llm_calls(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -463,6 +482,7 @@ class MonitoringRouterGroup(group.RouterGroup):
]
elif export_type == 'embedding-calls':
data = await self.ap.monitoring_service.export_embedding_calls(
+ request_context,
start_time=start_time,
end_time=end_time,
limit=limit,
@@ -485,6 +505,7 @@ class MonitoringRouterGroup(group.RouterGroup):
]
elif export_type == 'errors':
data = await self.ap.monitoring_service.export_errors(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -506,6 +527,7 @@ class MonitoringRouterGroup(group.RouterGroup):
]
elif export_type == 'sessions':
data = await self.ap.monitoring_service.export_sessions(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -527,6 +549,7 @@ class MonitoringRouterGroup(group.RouterGroup):
]
elif export_type == 'feedback':
data = await self.ap.monitoring_service.export_feedback(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -581,8 +604,8 @@ class MonitoringRouterGroup(group.RouterGroup):
return response, 200
- @self.route('/feedback/stats', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_feedback_stats() -> str:
+ @self.route('/feedback/stats', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_feedback_stats(request_context: RequestContext) -> str:
"""Get feedback statistics"""
# Parse query parameters
bot_ids = quart.request.args.getlist('botId')
@@ -595,6 +618,7 @@ class MonitoringRouterGroup(group.RouterGroup):
end_time = parse_iso_datetime(end_time_str)
stats = await self.ap.monitoring_service.get_feedback_stats(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
start_time=start_time,
@@ -603,8 +627,8 @@ class MonitoringRouterGroup(group.RouterGroup):
return self.success(data=stats)
- @self.route('/feedback', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def get_feedback() -> str:
+ @self.route('/feedback', methods=['GET'], permission=Permission.AUDIT_VIEW)
+ async def get_feedback(request_context: RequestContext) -> str:
"""Get feedback list"""
# Parse query parameters
bot_ids = quart.request.args.getlist('botId')
@@ -623,6 +647,7 @@ class MonitoringRouterGroup(group.RouterGroup):
feedback_type = int(feedback_type_str) if feedback_type_str else None
feedback_list, total = await self.ap.monitoring_service.get_feedback_list(
+ request_context,
bot_ids=bot_ids if bot_ids else None,
pipeline_ids=pipeline_ids if pipeline_ids else None,
feedback_type=feedback_type,
diff --git a/src/langbot/pkg/api/http/controller/groups/pipelines/embed.py b/src/langbot/pkg/api/http/controller/groups/pipelines/embed.py
index 50e9112b7..599a09ba3 100644
--- a/src/langbot/pkg/api/http/controller/groups/pipelines/embed.py
+++ b/src/langbot/pkg/api/http/controller/groups/pipelines/embed.py
@@ -21,9 +21,10 @@ import quart
from ... import group
from ......utils import paths
-from ......platform.sources.websocket_manager import is_valid_session_id, ws_connection_manager
+from ......platform.sources.websocket_manager import WebSocketScope, is_valid_session_id, ws_connection_manager
logger = logging.getLogger(__name__)
+_AUTH_TIMEOUT_SECONDS = 10.0
# Cache the widget template content
_widget_template_cache: str | None = None
@@ -58,37 +59,31 @@ def _get_logo_bytes() -> bytes:
class EmbedRouterGroup(group.RouterGroup):
# -- helpers -------------------------------------------------------------
- def _resolve_bot(self, bot_uuid: str):
+ async def _resolve_bot(self, bot_uuid: str):
"""Resolve *bot_uuid* to ``(runtime_bot, pipeline_uuid)``.
Returns ``(None, None)`` when the bot does not exist, is not a
``web_page_bot``, is disabled, or has no pipeline bound.
"""
- for bot in self.ap.platform_mgr.bots:
- if (
- bot.bot_entity.uuid == bot_uuid
- and bot.bot_entity.adapter == 'web_page_bot'
- and bot.bot_entity.enable
- and bot.bot_entity.use_pipeline_uuid
- ):
- return bot, bot.bot_entity.use_pipeline_uuid
+ bot = await self.ap.platform_mgr.resolve_public_bot(bot_uuid)
+ if (
+ bot is not None
+ and bot.bot_entity.adapter == 'web_page_bot'
+ and bot.bot_entity.enable
+ and bot.bot_entity.use_pipeline_uuid
+ ):
+ return bot, bot.bot_entity.use_pipeline_uuid
return None, None
- def _get_bot_config(self, bot_uuid: str) -> dict:
- for bot in self.ap.platform_mgr.bots:
- if bot.bot_entity.uuid == bot_uuid and bot.bot_entity.adapter == 'web_page_bot':
- return bot.bot_entity.adapter_config
- return {}
+ @staticmethod
+ def _get_bot_config(runtime_bot) -> dict:
+ return runtime_bot.bot_entity.adapter_config
- async def _verify_session_token(self, request, bot_uuid: str) -> bool:
- config = self._get_bot_config(bot_uuid)
+ def _verify_session_token_value(self, token: str, runtime_bot) -> bool:
+ config = self._get_bot_config(runtime_bot)
secret = config.get('turnstile_secret_key', '')
if not secret:
return True
- auth_header = request.headers.get('Authorization', '')
- if not auth_header.startswith('Bearer '):
- return False
- token = auth_header[7:]
try:
ts_str, mac = token.split('.', 1)
ts = float(ts_str)
@@ -99,6 +94,50 @@ class EmbedRouterGroup(group.RouterGroup):
except Exception:
return False
+ async def _verify_session_token(self, request, runtime_bot) -> bool:
+ auth_header = request.headers.get('Authorization', '')
+ token = auth_header[7:] if auth_header.startswith('Bearer ') else ''
+ return self._verify_session_token_value(token, runtime_bot)
+
+ async def _authenticate_websocket(self, runtime_bot) -> None:
+ """Require the embed session token as the first WebSocket frame."""
+
+ raw_message = await asyncio.wait_for(quart.websocket.receive(), timeout=_AUTH_TIMEOUT_SECONDS)
+ payload = json.loads(raw_message)
+ if not isinstance(payload, dict) or payload.get('type') != 'authenticate':
+ raise ValueError('Authentication is required')
+ token = str(payload.get('token') or '')
+ if not self._verify_session_token_value(token, runtime_bot):
+ raise ValueError('Authentication is required')
+
+ async def _assert_execution_active(self, runtime_bot) -> None:
+ context = runtime_bot.execution_context
+ await self.ap.workspace_service.get_execution_binding(
+ context.workspace_uuid,
+ expected_generation=context.placement_generation,
+ )
+
+ async def _resolve_connected_bot(self, owner_bot, pipeline_uuid: str):
+ """Re-resolve mutable bot state before every public message."""
+ current_bot, current_pipeline_uuid = await self._resolve_bot(owner_bot.bot_entity.uuid)
+ if current_bot is None or current_pipeline_uuid != pipeline_uuid:
+ raise RuntimeError('Bot is unavailable')
+
+ owner_context = owner_bot.execution_context
+ current_context = current_bot.execution_context
+ if (
+ current_context.instance_uuid,
+ current_context.workspace_uuid,
+ current_context.placement_generation,
+ ) != (
+ owner_context.instance_uuid,
+ owner_context.workspace_uuid,
+ owner_context.placement_generation,
+ ):
+ raise RuntimeError('Bot is unavailable')
+ await self._assert_execution_active(current_bot)
+ return current_bot
+
# -- routes --------------------------------------------------------------
async def initialize(self) -> None:
@@ -106,7 +145,7 @@ class EmbedRouterGroup(group.RouterGroup):
async def verify_turnstile(bot_uuid: str) -> str:
if not _is_valid_uuid(bot_uuid):
return self.http_status(400, -1, 'Invalid bot_uuid format')
- runtime_bot, pipeline_uuid = self._resolve_bot(bot_uuid)
+ runtime_bot, pipeline_uuid = await self._resolve_bot(bot_uuid)
if runtime_bot is None:
return self.http_status(404, -1, 'Bot not found or not available')
try:
@@ -115,7 +154,7 @@ class EmbedRouterGroup(group.RouterGroup):
if not token:
return self.http_status(400, -1, 'Token is required')
- config = self._get_bot_config(bot_uuid)
+ config = self._get_bot_config(runtime_bot)
secret = config.get('turnstile_secret_key', '')
if not secret:
ts = time.time()
@@ -146,7 +185,7 @@ class EmbedRouterGroup(group.RouterGroup):
"""Serve the embed widget JavaScript with injected configuration."""
if not _is_valid_uuid(bot_uuid):
return self.http_status(400, -1, 'Invalid bot_uuid format')
- runtime_bot, pipeline_uuid = self._resolve_bot(bot_uuid)
+ runtime_bot, pipeline_uuid = await self._resolve_bot(bot_uuid)
if runtime_bot is None:
return quart.Response(
'// Bot not found or not available', status=404, content_type='application/javascript'
@@ -164,7 +203,7 @@ class EmbedRouterGroup(group.RouterGroup):
if not re.match(r'^https?://[a-zA-Z0-9._:/-]+$', base_url):
base_url = quart.request.host_url.rstrip('/')
- config = self._get_bot_config(bot_uuid)
+ config = self._get_bot_config(runtime_bot)
site_key = config.get('turnstile_site_key', '')
locale = config.get('language', 'en_US') or 'en_US'
bubble_icon = config.get('bubble_icon', 'logo') or 'logo'
@@ -194,10 +233,10 @@ class EmbedRouterGroup(group.RouterGroup):
async def get_embed_messages(bot_uuid: str, session_type: str) -> str:
if not _is_valid_uuid(bot_uuid):
return self.http_status(400, -1, 'Invalid bot_uuid format')
- runtime_bot, pipeline_uuid = self._resolve_bot(bot_uuid)
+ runtime_bot, pipeline_uuid = await self._resolve_bot(bot_uuid)
if runtime_bot is None:
return self.http_status(404, -1, 'Bot not found or not available')
- if not await self._verify_session_token(quart.request, bot_uuid):
+ if not await self._verify_session_token(quart.request, runtime_bot):
return self.http_status(403, -1, 'Unauthorized or session expired')
try:
if session_type not in ['person', 'group']:
@@ -207,7 +246,8 @@ class EmbedRouterGroup(group.RouterGroup):
if not is_valid_session_id(session_id):
return self.http_status(400, -1, 'Valid session_id is required')
- websocket_adapter = self.ap.platform_mgr.websocket_proxy_bot.adapter
+ proxy_bot = await self.ap.platform_mgr.get_websocket_proxy_bot(runtime_bot.execution_context)
+ websocket_adapter = proxy_bot.adapter
if not websocket_adapter:
return self.http_status(404, -1, 'WebSocket adapter not found')
@@ -222,10 +262,10 @@ class EmbedRouterGroup(group.RouterGroup):
async def reset_embed_session(bot_uuid: str, session_type: str) -> str:
if not _is_valid_uuid(bot_uuid):
return self.http_status(400, -1, 'Invalid bot_uuid format')
- runtime_bot, pipeline_uuid = self._resolve_bot(bot_uuid)
+ runtime_bot, pipeline_uuid = await self._resolve_bot(bot_uuid)
if runtime_bot is None:
return self.http_status(404, -1, 'Bot not found or not available')
- if not await self._verify_session_token(quart.request, bot_uuid):
+ if not await self._verify_session_token(quart.request, runtime_bot):
return self.http_status(403, -1, 'Unauthorized or session expired')
try:
if session_type not in ['person', 'group']:
@@ -235,7 +275,8 @@ class EmbedRouterGroup(group.RouterGroup):
if not is_valid_session_id(session_id):
return self.http_status(400, -1, 'Valid session_id is required')
- websocket_adapter = self.ap.platform_mgr.websocket_proxy_bot.adapter
+ proxy_bot = await self.ap.platform_mgr.get_websocket_proxy_bot(runtime_bot.execution_context)
+ websocket_adapter = proxy_bot.adapter
if not websocket_adapter:
return self.http_status(404, -1, 'WebSocket adapter not found')
@@ -250,10 +291,10 @@ class EmbedRouterGroup(group.RouterGroup):
async def submit_feedback(bot_uuid: str) -> str:
if not _is_valid_uuid(bot_uuid):
return self.http_status(400, -1, 'Invalid bot_uuid format')
- runtime_bot, pipeline_uuid = self._resolve_bot(bot_uuid)
+ runtime_bot, pipeline_uuid = await self._resolve_bot(bot_uuid)
if runtime_bot is None:
return self.http_status(404, -1, 'Bot not found or not available')
- if not await self._verify_session_token(quart.request, bot_uuid):
+ if not await self._verify_session_token(quart.request, runtime_bot):
return self.http_status(403, -1, 'Unauthorized or session expired')
try:
data = await quart.request.get_json()
@@ -266,6 +307,7 @@ class EmbedRouterGroup(group.RouterGroup):
feedback_id = f'embed_{uuid.uuid4().hex[:12]}'
await self.ap.monitoring_service.record_feedback(
+ runtime_bot.execution_context,
feedback_id=feedback_id,
feedback_type=feedback_type,
bot_id=runtime_bot.bot_entity.uuid,
@@ -286,11 +328,12 @@ class EmbedRouterGroup(group.RouterGroup):
@self.quart_app.websocket(self.path + '//ws/connect')
async def embed_websocket_connect(bot_uuid: str):
"""WebSocket connection for embed widget, keyed by bot_uuid."""
+ await quart.websocket.accept()
if not _is_valid_uuid(bot_uuid):
await quart.websocket.send(json.dumps({'type': 'error', 'message': 'Invalid bot_uuid format'}))
return
- runtime_bot, pipeline_uuid = self._resolve_bot(bot_uuid)
+ runtime_bot, pipeline_uuid = await self._resolve_bot(bot_uuid)
if runtime_bot is None:
await quart.websocket.send(json.dumps({'type': 'error', 'message': 'Bot not found or not available'}))
return
@@ -307,14 +350,23 @@ class EmbedRouterGroup(group.RouterGroup):
await quart.websocket.send(json.dumps({'type': 'error', 'message': 'Valid session_id is required'}))
return
- websocket_adapter = self.ap.platform_mgr.websocket_proxy_bot.adapter
- if not websocket_adapter:
- await quart.websocket.send(json.dumps({'type': 'error', 'message': 'WebSocket adapter not found'}))
+ try:
+ await self._authenticate_websocket(runtime_bot)
+ await self._assert_execution_active(runtime_bot)
+ except Exception:
+ await quart.websocket.send(json.dumps({'type': 'error', 'message': 'Unauthorized'}))
return
try:
+ proxy_bot = await self.ap.platform_mgr.get_websocket_proxy_bot(runtime_bot.execution_context)
+ websocket_adapter = proxy_bot.adapter
+ if not websocket_adapter:
+ await quart.websocket.send(json.dumps({'type': 'error', 'message': 'WebSocket adapter not found'}))
+ return
+
connection = await ws_connection_manager.add_connection(
websocket=quart.websocket._get_current_object(),
+ scope=WebSocketScope.from_context(runtime_bot.execution_context),
pipeline_uuid=pipeline_uuid,
session_type=session_type,
session_id=session_id,
@@ -338,7 +390,9 @@ class EmbedRouterGroup(group.RouterGroup):
f'(bot={bot_uuid}, pipeline={pipeline_uuid}, session_type={session_type})'
)
- receive_task = asyncio.create_task(self._handle_receive(connection, websocket_adapter, runtime_bot))
+ receive_task = asyncio.create_task(
+ self._handle_receive(connection, websocket_adapter, runtime_bot, pipeline_uuid)
+ )
send_task = asyncio.create_task(self._handle_send(connection))
try:
@@ -357,7 +411,7 @@ class EmbedRouterGroup(group.RouterGroup):
# -- WebSocket receive/send helpers --------------------------------------
- async def _handle_receive(self, connection, websocket_adapter, owner_bot):
+ async def _handle_receive(self, connection, websocket_adapter, owner_bot, pipeline_uuid: str):
try:
while connection.is_active:
message = await quart.websocket.receive()
@@ -372,7 +426,12 @@ class EmbedRouterGroup(group.RouterGroup):
{'type': 'pong', 'timestamp': datetime.datetime.now().isoformat()}
)
elif message_type == 'message':
- await websocket_adapter.handle_websocket_message(connection, data, owner_bot=owner_bot)
+ try:
+ current_bot = await self._resolve_connected_bot(owner_bot, pipeline_uuid)
+ except Exception:
+ await connection.send_queue.put({'type': 'error', 'message': 'Bot is unavailable'})
+ break
+ await websocket_adapter.handle_websocket_message(connection, data, owner_bot=current_bot)
elif message_type == 'disconnect':
break
@@ -386,7 +445,7 @@ class EmbedRouterGroup(group.RouterGroup):
async def _handle_send(self, connection):
try:
- while connection.is_active:
+ while connection.is_active or not connection.send_queue.empty():
try:
message = await asyncio.wait_for(connection.send_queue.get(), timeout=1.0)
await quart.websocket.send(json.dumps(message))
diff --git a/src/langbot/pkg/api/http/controller/groups/pipelines/pipelines.py b/src/langbot/pkg/api/http/controller/groups/pipelines/pipelines.py
index 2e45add77..69189d2ee 100644
--- a/src/langbot/pkg/api/http/controller/groups/pipelines/pipelines.py
+++ b/src/langbot/pkg/api/http/controller/groups/pipelines/pipelines.py
@@ -2,120 +2,156 @@ from __future__ import annotations
import quart
+from ....authz import Permission, has_permission
+from ....context import RequestContext
+from ....service.secrets import redact_secrets
from ... import group
@group.group_class('pipelines', '/api/v1/pipelines')
class PipelinesRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET', 'POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
- if quart.request.method == 'GET':
- sort_by = quart.request.args.get('sort_by', 'created_at')
- sort_order = quart.request.args.get('sort_order', 'DESC')
- return self.success(
- data={'pipelines': await self.ap.pipeline_service.get_pipelines(sort_by, sort_order)}
- )
- elif quart.request.method == 'POST':
- json_data = await quart.request.json
-
- pipeline_uuid = await self.ap.pipeline_service.create_pipeline(json_data)
-
- return self.success(data={'uuid': pipeline_uuid})
-
- @self.route('/_/metadata', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
- return self.success(data={'configs': await self.ap.pipeline_service.get_pipeline_metadata()})
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ sort_by = quart.request.args.get('sort_by', 'created_at')
+ sort_order = quart.request.args.get('sort_order', 'DESC')
+ include_secret = has_permission(request_context, Permission.RESOURCE_MANAGE)
+ return self.success(
+ data={
+ 'pipelines': await self.ap.pipeline_service.get_pipelines(
+ request_context,
+ sort_by,
+ sort_order,
+ include_secret=include_secret,
+ )
+ }
+ )
@self.route(
- '/', methods=['GET', 'PUT', 'DELETE'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY
+ '',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
)
- async def _(pipeline_uuid: str) -> str:
- if quart.request.method == 'GET':
- pipeline = await self.ap.pipeline_service.get_pipeline(pipeline_uuid)
+ async def _(request_context: RequestContext) -> str:
+ pipeline_uuid = await self.ap.pipeline_service.create_pipeline(request_context, await quart.request.json)
+ return self.success(data={'uuid': pipeline_uuid})
- if pipeline is None:
- return self.http_status(404, -1, 'pipeline not found')
+ @self.route(
+ '/_/metadata',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ return self.success(data={'configs': await self.ap.pipeline_service.get_pipeline_metadata(request_context)})
- return self.success(data={'pipeline': pipeline})
- elif quart.request.method == 'PUT':
- json_data = await quart.request.json
+ @self.route(
+ '/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(pipeline_uuid: str, request_context: RequestContext) -> str:
+ pipeline = await self.ap.pipeline_service.get_pipeline(
+ request_context,
+ pipeline_uuid,
+ include_secret=has_permission(request_context, Permission.RESOURCE_MANAGE),
+ )
+ if pipeline is None:
+ return self.http_status(404, -1, 'pipeline not found')
+ return self.success(data={'pipeline': pipeline})
- await self.ap.pipeline_service.update_pipeline(pipeline_uuid, json_data)
+ @self.route(
+ '/',
+ methods=['PUT', 'DELETE'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(pipeline_uuid: str, request_context: RequestContext) -> str:
+ if quart.request.method == 'PUT':
+ try:
+ await self.ap.pipeline_service.update_pipeline(
+ request_context,
+ pipeline_uuid,
+ await quart.request.json,
+ )
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ else:
+ await self.ap.pipeline_service.delete_pipeline(request_context, pipeline_uuid)
+ return self.success()
- return self.success()
- elif quart.request.method == 'DELETE':
- await self.ap.pipeline_service.delete_pipeline(pipeline_uuid)
-
- return self.success()
-
- @self.route('//copy', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _(pipeline_uuid: str) -> str:
+ @self.route(
+ '//copy',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(pipeline_uuid: str, request_context: RequestContext) -> str:
try:
- new_uuid = await self.ap.pipeline_service.copy_pipeline(pipeline_uuid)
+ new_uuid = await self.ap.pipeline_service.copy_pipeline(request_context, pipeline_uuid)
return self.success(data={'uuid': new_uuid})
except ValueError as e:
- return self.http_status(404, -1, str(e))
+ return self.http_status(400, -1, str(e))
@self.route(
- '//extensions', methods=['GET', 'PUT'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY
+ '//extensions',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def _(pipeline_uuid: str) -> str:
- if quart.request.method == 'GET':
- # Get current extensions and available plugins
- pipeline = await self.ap.pipeline_service.get_pipeline(pipeline_uuid)
- if pipeline is None:
- return self.http_status(404, -1, 'pipeline not found')
+ async def _(pipeline_uuid: str, request_context: RequestContext) -> str:
+ pipeline = await self.ap.pipeline_service.get_pipeline(request_context, pipeline_uuid)
+ if pipeline is None:
+ return self.http_status(404, -1, 'pipeline not found')
- # Only include plugins with pipeline-related components (Command, EventListener, Tool)
- # Plugins that only have KnowledgeEngine components are not suitable for pipeline extensions
- pipeline_component_kinds = ['Command', 'EventListener', 'Tool']
- plugins = await self.ap.plugin_connector.list_plugins(component_kinds=pipeline_component_kinds)
- mcp_servers = await self.ap.mcp_service.get_mcp_servers(contain_runtime_info=True)
+ pipeline_component_kinds = ['Command', 'EventListener', 'Tool']
+ if self.ap.plugin_connector.is_enable_plugin:
+ await self.ap.plugin_connector.require_workspace_context(request_context)
+ plugins = await self.ap.plugin_connector.list_plugins(component_kinds=pipeline_component_kinds)
+ mcp_servers = await self.ap.mcp_service.get_mcp_servers(request_context, contain_runtime_info=True)
+ available_skills = await self.ap.skill_service.list_skills(request_context)
+ extensions_prefs = pipeline.get('extensions_preferences', {})
+ return self.success(
+ data={
+ 'enable_all_plugins': extensions_prefs.get('enable_all_plugins', True),
+ 'enable_all_mcp_servers': extensions_prefs.get('enable_all_mcp_servers', True),
+ 'enable_all_skills': extensions_prefs.get('enable_all_skills', True),
+ 'bound_plugins': extensions_prefs.get('plugins', []),
+ 'available_plugins': redact_secrets(plugins),
+ 'bound_mcp_servers': extensions_prefs.get('mcp_servers', []),
+ 'available_mcp_servers': mcp_servers,
+ 'bound_mcp_resources': extensions_prefs.get('mcp_resources', []),
+ 'mcp_resource_agent_read_enabled': extensions_prefs.get('mcp_resource_agent_read_enabled', True),
+ 'bound_skills': extensions_prefs.get('skills', []),
+ 'available_skills': available_skills,
+ }
+ )
- # Get available skills
- available_skills = await self.ap.skill_service.list_skills()
-
- extensions_prefs = pipeline.get('extensions_preferences', {})
- return self.success(
- data={
- 'enable_all_plugins': extensions_prefs.get('enable_all_plugins', True),
- 'enable_all_mcp_servers': extensions_prefs.get('enable_all_mcp_servers', True),
- 'enable_all_skills': extensions_prefs.get('enable_all_skills', True),
- 'bound_plugins': extensions_prefs.get('plugins', []),
- 'available_plugins': plugins,
- 'bound_mcp_servers': extensions_prefs.get('mcp_servers', []),
- 'available_mcp_servers': mcp_servers,
- 'bound_mcp_resources': extensions_prefs.get('mcp_resources', []),
- 'mcp_resource_agent_read_enabled': extensions_prefs.get(
- 'mcp_resource_agent_read_enabled', True
- ),
- 'bound_skills': extensions_prefs.get('skills', []),
- 'available_skills': available_skills,
- }
- )
- elif quart.request.method == 'PUT':
- # Update bound plugins and MCP servers for this pipeline
- json_data = await quart.request.json
- enable_all_plugins = json_data.get('enable_all_plugins', True)
- enable_all_mcp_servers = json_data.get('enable_all_mcp_servers', True)
- enable_all_skills = json_data.get('enable_all_skills', True)
- bound_plugins = json_data.get('bound_plugins', [])
- bound_mcp_servers = json_data.get('bound_mcp_servers', [])
- bound_skills = json_data.get('bound_skills', [])
- bound_mcp_resources = json_data.get('bound_mcp_resources')
- mcp_resource_agent_read_enabled = json_data.get('mcp_resource_agent_read_enabled')
-
- await self.ap.pipeline_service.update_pipeline_extensions(
- pipeline_uuid,
- bound_plugins,
- bound_mcp_servers,
- enable_all_plugins,
- enable_all_mcp_servers,
- bound_skills=bound_skills,
- enable_all_skills=enable_all_skills,
- bound_mcp_resources=bound_mcp_resources,
- mcp_resource_agent_read_enabled=mcp_resource_agent_read_enabled,
- )
-
- return self.success()
+ @self.route(
+ '//extensions',
+ methods=['PUT'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(pipeline_uuid: str, request_context: RequestContext) -> str:
+ json_data = await quart.request.json
+ await self.ap.pipeline_service.update_pipeline_extensions(
+ request_context,
+ pipeline_uuid,
+ json_data.get('bound_plugins', []),
+ json_data.get('bound_mcp_servers', []),
+ json_data.get('enable_all_plugins', True),
+ json_data.get('enable_all_mcp_servers', True),
+ bound_skills=json_data.get('bound_skills', []),
+ enable_all_skills=json_data.get('enable_all_skills', True),
+ bound_mcp_resources=json_data.get('bound_mcp_resources'),
+ mcp_resource_agent_read_enabled=json_data.get('mcp_resource_agent_read_enabled'),
+ )
+ return self.success()
diff --git a/src/langbot/pkg/api/http/controller/groups/pipelines/websocket_chat.py b/src/langbot/pkg/api/http/controller/groups/pipelines/websocket_chat.py
index ebe46b8fe..ad16c2ca5 100644
--- a/src/langbot/pkg/api/http/controller/groups/pipelines/websocket_chat.py
+++ b/src/langbot/pkg/api/http/controller/groups/pipelines/websocket_chat.py
@@ -1,64 +1,157 @@
-"""WebSocket聊天路由 - 支持双向实时通信"""
+"""Authenticated dashboard WebSocket chat routes."""
+
+from __future__ import annotations
import asyncio
import datetime
import json
import logging
+import uuid
import quart
+from ....authz import Permission, permissions_for_role, require_permission
+from ....context import PrincipalContext, PrincipalType, RequestContext, WorkspaceContext
from ... import group
-from ......platform.sources.websocket_manager import ws_connection_manager
+from ......platform.sources.websocket_manager import WebSocketScope, ws_connection_manager
logger = logging.getLogger(__name__)
+_AUTH_TIMEOUT_SECONDS = 10.0
@group.group_class('websocket_chat', '/api/v1/pipelines//ws')
class WebSocketChatRouterGroup(group.RouterGroup):
+ async def _authenticate_websocket(self) -> tuple[RequestContext, str]:
+ """Authenticate the first dashboard WebSocket message.
+
+ Browsers cannot attach the normal Authorization/X-Workspace-Id headers
+ to a WebSocket handshake. The client therefore sends one auth frame
+ immediately after opening the socket; no connection is registered and
+ no runtime object is resolved before this method succeeds.
+ """
+
+ raw_message = await asyncio.wait_for(quart.websocket.receive(), timeout=_AUTH_TIMEOUT_SECONDS)
+ payload = json.loads(raw_message)
+ if not isinstance(payload, dict) or payload.get('type') != 'authenticate':
+ raise ValueError('Authentication is required')
+
+ token = str(payload.get('token') or '').strip()
+ workspace_uuid = str(payload.get('workspace_uuid') or '').strip()
+ if not token or not workspace_uuid:
+ raise ValueError('Authentication is required')
+
+ account, _ = await self._authenticate_account(token)
+ account_uuid = getattr(account, 'uuid', None)
+ collaboration_service = getattr(self.ap, 'workspace_collaboration_service', None)
+ if not isinstance(account_uuid, str) or collaboration_service is None:
+ raise ValueError('Workspace authentication is unavailable')
+
+ access = await collaboration_service.resolve_account_workspace(account_uuid, workspace_uuid)
+ request_context = RequestContext(
+ instance_uuid=access.execution.instance_uuid,
+ placement_generation=access.execution.placement_generation,
+ request_id=quart.websocket.headers.get('X-Request-Id') or str(uuid.uuid4()),
+ auth_type=group.AuthType.USER_TOKEN.value,
+ principal=PrincipalContext(
+ principal_type=PrincipalType.ACCOUNT,
+ account_uuid=account_uuid,
+ ),
+ workspace=WorkspaceContext(
+ workspace_uuid=access.workspace.uuid,
+ membership_uuid=access.membership.uuid,
+ role=access.membership.role,
+ permissions=permissions_for_role(access.membership.role),
+ membership_revision=access.membership.projection_revision,
+ ),
+ )
+ require_permission(request_context, Permission.RUNTIME_OPERATE)
+ return request_context, token
+
+ async def _revalidate_websocket_authorization(
+ self,
+ request_context: RequestContext,
+ token: str,
+ ) -> None:
+ """Recheck revocable account, membership, permission, and placement state."""
+
+ account, _ = await self._authenticate_account(token)
+ account_uuid = getattr(account, 'uuid', None)
+ if account_uuid != request_context.account_uuid:
+ raise ValueError('WebSocket account changed')
+
+ collaboration_service = getattr(self.ap, 'workspace_collaboration_service', None)
+ if collaboration_service is None or not isinstance(account_uuid, str):
+ raise ValueError('Workspace authentication is unavailable')
+ access = await collaboration_service.resolve_account_workspace(
+ account_uuid,
+ request_context.workspace_uuid,
+ )
+ if (
+ access.workspace.uuid != request_context.workspace_uuid
+ or access.membership.uuid != request_context.workspace.membership_uuid
+ or access.membership.projection_revision != request_context.workspace.membership_revision
+ or access.execution.instance_uuid != request_context.instance_uuid
+ or access.execution.placement_generation != request_context.placement_generation
+ ):
+ raise ValueError('WebSocket authorization changed')
+
+ current_context = RequestContext(
+ instance_uuid=access.execution.instance_uuid,
+ placement_generation=access.execution.placement_generation,
+ request_id=request_context.request_id,
+ auth_type=request_context.auth_type,
+ principal=request_context.principal,
+ workspace=WorkspaceContext(
+ workspace_uuid=access.workspace.uuid,
+ membership_uuid=access.membership.uuid,
+ role=access.membership.role,
+ permissions=permissions_for_role(access.membership.role),
+ membership_revision=access.membership.projection_revision,
+ ),
+ entitlement_revision=request_context.entitlement_revision,
+ )
+ require_permission(current_context, Permission.RUNTIME_OPERATE)
+
+ async def _get_scoped_adapter(self, request_context: RequestContext, pipeline_uuid: str):
+ pipeline = await self.ap.pipeline_service.get_pipeline(request_context, pipeline_uuid)
+ if pipeline is None:
+ return None
+ proxy_bot = await self.ap.platform_mgr.get_websocket_proxy_bot(request_context)
+ return proxy_bot.adapter
+
async def initialize(self) -> None:
- # 直接使用 quart_app 注册 WebSocket 路由
@self.quart_app.websocket(self.path + '/connect')
async def websocket_connect(pipeline_uuid: str):
- """
- 建立WebSocket连接
+ """Open one authenticated dashboard debug connection."""
- URL参数:
- - pipeline_uuid: 流水线UUID
- - session_type: 会话类型 (person/group)
- """
+ await quart.websocket.accept()
try:
- # 获取参数 - 在WebSocket上下文中使用 quart.websocket.args
- session_type = quart.websocket.args.get('session_type', 'person')
+ request_context, token = await self._authenticate_websocket()
+ except Exception:
+ await quart.websocket.send(json.dumps({'type': 'error', 'message': 'Unauthorized'}))
+ return
- if session_type not in ['person', 'group']:
- await quart.websocket.send(
- json.dumps({'type': 'error', 'message': 'session_type must be person or group'})
- )
+ session_type = quart.websocket.args.get('session_type', 'person')
+ if session_type not in ['person', 'group']:
+ await quart.websocket.send(
+ json.dumps({'type': 'error', 'message': 'session_type must be person or group'})
+ )
+ return
+
+ try:
+ websocket_adapter = await self._get_scoped_adapter(request_context, pipeline_uuid)
+ if websocket_adapter is None:
+ await quart.websocket.send(json.dumps({'type': 'error', 'message': 'Pipeline not found'}))
return
- # 获取WebSocket适配器
- websocket_adapter = self.ap.platform_mgr.websocket_proxy_bot.adapter
-
- if not websocket_adapter:
- await quart.websocket.send(json.dumps({'type': 'error', 'message': 'WebSocket adapter not found'}))
- return
-
- # Dashboard pipeline-debug sessions must always run under the
- # built-in websocket_proxy_bot identity. We deliberately do NOT
- # resolve a web_page_bot owner here — even if one is bound to
- # the same pipeline, debug requests must not be attributed to
- # it. The embed widget path (`/api/v1/embed//ws/connect`)
- # is the one that carries the page-bot identity.
-
- # 注册连接
connection = await ws_connection_manager.add_connection(
websocket=quart.websocket._get_current_object(),
+ scope=WebSocketScope.from_context(request_context),
pipeline_uuid=pipeline_uuid,
session_type=session_type,
metadata={'user_agent': quart.websocket.headers.get('User-Agent', '')},
)
- # 发送连接成功消息
await quart.websocket.send(
json.dumps(
{
@@ -72,182 +165,180 @@ class WebSocketChatRouterGroup(group.RouterGroup):
)
logger.debug(
- f'WebSocket connection established: {connection.connection_id} '
- f'(pipeline={pipeline_uuid}, session_type={session_type})'
+ f'Dashboard WebSocket connected: {connection.connection_id} '
+ f'(workspace={connection.workspace_uuid}, pipeline={pipeline_uuid}, '
+ f'session_type={session_type})'
)
- # 创建接收和发送任务
- receive_task = asyncio.create_task(self._handle_receive(connection, websocket_adapter))
+ receive_task = asyncio.create_task(
+ self._handle_receive(
+ connection,
+ websocket_adapter,
+ request_context,
+ token,
+ )
+ )
send_task = asyncio.create_task(self._handle_send(connection))
-
- # 等待任务完成
try:
await asyncio.gather(receive_task, send_task)
- except Exception as e:
- logger.error(f'WebSocket task execution error: {e}')
+ except Exception as exc:
+ logger.error(f'WebSocket task execution error: {exc}')
finally:
- # 清理连接
await ws_connection_manager.remove_connection(connection.connection_id)
- logger.debug(f'WebSocket connection cleaned: {connection.connection_id}')
- except Exception as e:
- logger.error(f'WebSocket connection error: {e}', exc_info=True)
+ except Exception:
+ logger.error('Dashboard WebSocket connection error', exc_info=True)
try:
- await quart.websocket.send(json.dumps({'type': 'error', 'message': str(e)}))
- except:
+ await quart.websocket.send(json.dumps({'type': 'error', 'message': 'Internal server error'}))
+ except Exception:
pass
- @self.route('/messages/', methods=['GET'])
- async def get_messages(pipeline_uuid: str, session_type: str) -> str:
- """获取消息历史"""
- try:
- if session_type not in ['person', 'group']:
- return self.http_status(400, -1, 'session_type must be person or group')
+ @self.route(
+ '/messages/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RUNTIME_OPERATE,
+ )
+ async def get_messages(
+ pipeline_uuid: str,
+ session_type: str,
+ request_context: RequestContext,
+ ) -> str:
+ if session_type not in ['person', 'group']:
+ return self.http_status(400, -1, 'session_type must be person or group')
- websocket_adapter = self.ap.platform_mgr.websocket_proxy_bot.adapter
+ websocket_adapter = await self._get_scoped_adapter(request_context, pipeline_uuid)
+ if websocket_adapter is None:
+ return self.http_status(404, -1, 'Pipeline not found')
+ messages = websocket_adapter.get_websocket_messages(pipeline_uuid, session_type)
+ return self.success(data={'messages': messages})
- if not websocket_adapter:
- return self.http_status(404, -1, 'WebSocket adapter not found')
+ @self.route(
+ '/reset/',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RUNTIME_OPERATE,
+ )
+ async def reset_session(
+ pipeline_uuid: str,
+ session_type: str,
+ request_context: RequestContext,
+ ) -> str:
+ if session_type not in ['person', 'group']:
+ return self.http_status(400, -1, 'session_type must be person or group')
- messages = websocket_adapter.get_websocket_messages(pipeline_uuid, session_type)
+ websocket_adapter = await self._get_scoped_adapter(request_context, pipeline_uuid)
+ if websocket_adapter is None:
+ return self.http_status(404, -1, 'Pipeline not found')
+ websocket_adapter.reset_session(pipeline_uuid, session_type)
+ return self.success(data={'message': 'Session reset successfully'})
- return self.success(data={'messages': messages})
+ @self.route(
+ '/connections',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RUNTIME_OPERATE,
+ )
+ async def get_connections(pipeline_uuid: str, request_context: RequestContext) -> str:
+ if await self.ap.pipeline_service.get_pipeline(request_context, pipeline_uuid) is None:
+ return self.http_status(404, -1, 'Pipeline not found')
- except Exception as e:
- return self.http_status(500, -1, f'Internal server error: {str(e)}')
-
- @self.route('/reset/', methods=['POST'])
- async def reset_session(pipeline_uuid: str, session_type: str) -> str:
- """重置会话"""
- try:
- if session_type not in ['person', 'group']:
- return self.http_status(400, -1, 'session_type must be person or group')
-
- websocket_adapter = self.ap.platform_mgr.websocket_proxy_bot.adapter
-
- if not websocket_adapter:
- return self.http_status(404, -1, 'WebSocket adapter not found')
-
- websocket_adapter.reset_session(pipeline_uuid, session_type)
-
- return self.success(data={'message': 'Session reset successfully'})
-
- except Exception as e:
- return self.http_status(500, -1, f'Internal server error: {str(e)}')
-
- @self.route('/connections', methods=['GET'])
- async def get_connections(pipeline_uuid: str) -> str:
- """获取当前连接统计"""
- try:
- stats = ws_connection_manager.get_stats()
- connections = await ws_connection_manager.get_connections_by_pipeline(pipeline_uuid)
-
- return self.success(
- data={
- 'stats': stats,
- 'connections': [
- {
- 'connection_id': conn.connection_id,
- 'session_type': conn.session_type,
- 'created_at': conn.created_at.isoformat(),
- 'last_active': conn.last_active.isoformat(),
- 'is_active': conn.is_active,
- }
- for conn in connections
- ],
- }
- )
-
- except Exception as e:
- return self.http_status(500, -1, f'Internal server error: {str(e)}')
-
- @self.route('/broadcast', methods=['POST'])
- async def broadcast_message(pipeline_uuid: str) -> str:
- """向所有连接广播消息(后端主动推送)"""
- try:
- data = await quart.request.get_json()
- message = data.get('message')
-
- if not message:
- return self.http_status(400, -1, 'message is required')
-
- # 广播消息
- broadcast_data = {
- 'type': 'broadcast',
- 'message': message,
- 'timestamp': datetime.datetime.now().isoformat(),
+ scope = WebSocketScope.from_context(request_context)
+ stats = ws_connection_manager.get_stats(scope=scope)
+ connections = await ws_connection_manager.get_connections_by_pipeline(
+ pipeline_uuid,
+ scope=scope,
+ )
+ return self.success(
+ data={
+ 'stats': stats,
+ 'connections': [
+ {
+ 'connection_id': connection.connection_id,
+ 'session_type': connection.session_type,
+ 'created_at': connection.created_at.isoformat(),
+ 'last_active': connection.last_active.isoformat(),
+ 'is_active': connection.is_active,
+ }
+ for connection in connections
+ ],
}
+ )
- await ws_connection_manager.broadcast_to_pipeline(pipeline_uuid, broadcast_data)
+ @self.route(
+ '/broadcast',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RUNTIME_OPERATE,
+ )
+ async def broadcast_message(pipeline_uuid: str, request_context: RequestContext) -> str:
+ if await self.ap.pipeline_service.get_pipeline(request_context, pipeline_uuid) is None:
+ return self.http_status(404, -1, 'Pipeline not found')
- return self.success(data={'message': 'Broadcast sent successfully'})
+ data = await quart.request.get_json()
+ message = data.get('message')
+ if not message:
+ return self.http_status(400, -1, 'message is required')
- except Exception as e:
- return self.http_status(500, -1, f'Internal server error: {str(e)}')
+ broadcast_data = {
+ 'type': 'broadcast',
+ 'message': message,
+ 'timestamp': datetime.datetime.now().isoformat(),
+ }
+ await ws_connection_manager.broadcast_to_pipeline(
+ pipeline_uuid,
+ broadcast_data,
+ scope=WebSocketScope.from_context(request_context),
+ )
+ return self.success(data={'message': 'Broadcast sent successfully'})
- async def _handle_receive(self, connection, websocket_adapter):
- """处理接收消息的任务"""
+ async def _handle_receive(
+ self,
+ connection,
+ websocket_adapter,
+ request_context: RequestContext,
+ token: str,
+ ):
try:
while connection.is_active:
- # 接收消息
message = await quart.websocket.receive()
-
- # 更新活跃时间
await ws_connection_manager.update_activity(connection.connection_id)
try:
data = json.loads(message)
message_type = data.get('type', 'message')
-
if message_type == 'ping':
- # 心跳响应
await connection.send_queue.put(
{'type': 'pong', 'timestamp': datetime.datetime.now().isoformat()}
)
-
elif message_type == 'message':
- # 处理用户消息
- logger.debug(f'收到消息: {data} from {connection.connection_id}')
-
- # 处理消息(不等待响应,响应会通过broadcast异步发送)
- # owner_bot is intentionally NOT passed: the dashboard
- # debug WebSocket must always run under the proxy bot,
- # never under a coincidentally-bound web_page_bot.
+ try:
+ await self._revalidate_websocket_authorization(request_context, token)
+ except Exception:
+ await connection.send_queue.put({'type': 'error', 'message': 'Unauthorized'})
+ break
await websocket_adapter.handle_websocket_message(connection, data)
-
elif message_type == 'disconnect':
- # 客户端主动断开
- logger.debug(f'Client disconnected: {connection.connection_id}')
break
-
else:
- logger.warning(f'Unknown message type: {message_type}')
-
+ logger.warning(f'Unknown WebSocket message type: {message_type}')
except json.JSONDecodeError:
- logger.error(f'Invalid JSON message: {message}')
await connection.send_queue.put({'type': 'error', 'message': 'Invalid JSON format'})
- except Exception as e:
- logger.error(f'Receive message error: {e}', exc_info=True)
+ except Exception:
+ logger.error('Dashboard WebSocket receive error', exc_info=True)
finally:
connection.is_active = False
async def _handle_send(self, connection):
- """处理发送消息的任务"""
try:
- while connection.is_active:
- # 从队列获取消息
+ while connection.is_active or not connection.send_queue.empty():
try:
message = await asyncio.wait_for(connection.send_queue.get(), timeout=1.0)
-
- # 发送消息
await quart.websocket.send(json.dumps(message))
-
except asyncio.TimeoutError:
- # 超时继续循环
continue
-
- except Exception as e:
- logger.error(f'Send message error: {e}', exc_info=True)
+ except Exception:
+ logger.error('Dashboard WebSocket send error', exc_info=True)
finally:
connection.is_active = False
diff --git a/src/langbot/pkg/api/http/controller/groups/platform/adapters.py b/src/langbot/pkg/api/http/controller/groups/platform/adapters.py
index 0e32f9d29..fe8d8275d 100644
--- a/src/langbot/pkg/api/http/controller/groups/platform/adapters.py
+++ b/src/langbot/pkg/api/http/controller/groups/platform/adapters.py
@@ -1,9 +1,75 @@
-import quart
-import mimetypes
import asyncio
-from ... import group
+import dataclasses
+import mimetypes
+
+import quart
+
+from langbot.pkg.api.http.authz import Permission
+from langbot.pkg.api.http.context import RequestContext
from langbot.pkg.utils import importutil
+from ... import group
+
+
+@dataclasses.dataclass(frozen=True, slots=True)
+class _AdapterSessionScope:
+ """Immutable tenant and principal binding for a credential exchange."""
+
+ instance_uuid: str
+ workspace_uuid: str
+ placement_generation: int
+ principal_type: str
+ account_uuid: str | None
+ api_key_uuid: str | None
+
+ @classmethod
+ def from_request_context(cls, request_context: RequestContext) -> '_AdapterSessionScope':
+ principal = request_context.principal
+ return cls(
+ instance_uuid=request_context.instance_uuid,
+ workspace_uuid=request_context.workspace_uuid,
+ placement_generation=request_context.placement_generation,
+ principal_type=principal.principal_type.value,
+ account_uuid=principal.account_uuid,
+ api_key_uuid=principal.api_key_uuid,
+ )
+
+ def matches(self, request_context: RequestContext) -> bool:
+ """Return whether a request is from the exact initiating tenant principal."""
+
+ return self == self.from_request_context(request_context)
+
+
+def _bind_session_scope(session: dict, request_context: RequestContext) -> None:
+ session['scope'] = _AdapterSessionScope.from_request_context(request_context)
+
+
+def _get_owned_session(
+ sessions: dict[str, dict],
+ session_id: str,
+ request_context: RequestContext,
+) -> dict | None:
+ """Resolve a session without revealing sessions owned by another scope."""
+
+ session = sessions.get(session_id)
+ scope = session.get('scope') if session is not None else None
+ if not isinstance(scope, _AdapterSessionScope) or not scope.matches(request_context):
+ return None
+ return session
+
+
+def _pop_owned_session(
+ sessions: dict[str, dict],
+ session_id: str,
+ request_context: RequestContext,
+) -> dict | None:
+ """Remove an owned session without allowing cross-scope cancellation."""
+
+ session = _get_owned_session(sessions, session_id, request_context)
+ if session is None:
+ return None
+ return sessions.pop(session_id, None)
+
def _decrypt_qqofficial_secret(encrypted_b64: str, key: bytes) -> str:
"""Decrypt the AppSecret returned by the QQ Official QR binding endpoint.
@@ -84,8 +150,8 @@ class AdaptersRouterGroup(group.RouterGroup):
if session and session.get('task') and not session['task'].done():
session['task'].cancel()
- @self.route('/lark/create-app', methods=['POST'])
- async def _() -> str:
+ @self.route('/lark/create-app', methods=['POST'], permission=Permission.RESOURCE_MANAGE)
+ async def _(request_context: RequestContext) -> str:
"""Start Feishu one-click app registration. Returns session_id + QR code URL."""
import uuid
import time
@@ -106,6 +172,7 @@ class AdaptersRouterGroup(group.RouterGroup):
'error': None,
'created_at': time.time(),
}
+ _bind_session_scope(session, request_context)
_create_app_sessions[session_id] = session
def on_qr_code(info):
@@ -160,10 +227,15 @@ class AdaptersRouterGroup(group.RouterGroup):
}
)
- @self.route('/lark/create-app/status/', methods=['GET'])
- async def _(session_id: str) -> str:
+ @self.route(
+ '/lark/create-app/status/',
+ methods=['GET'],
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(session_id: str, request_context: RequestContext) -> str:
"""Poll registration status."""
- session = _create_app_sessions.get(session_id)
+ _cleanup_expired_sessions()
+ session = _get_owned_session(_create_app_sessions, session_id, request_context)
if not session:
return self.http_status(404, -1, 'Session not found')
@@ -179,10 +251,16 @@ class AdaptersRouterGroup(group.RouterGroup):
return self.success(data=data)
- @self.route('/lark/create-app/', methods=['DELETE'])
- async def _(session_id: str) -> str:
+ @self.route(
+ '/lark/create-app/',
+ methods=['DELETE'],
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(session_id: str, request_context: RequestContext) -> str:
"""Cancel and clean up a registration session."""
- session = _create_app_sessions.pop(session_id, None)
+ session = _pop_owned_session(_create_app_sessions, session_id, request_context)
+ if session is None:
+ return self.http_status(404, -1, 'Session not found')
if session and session.get('task') and not session['task'].done():
session['task'].cancel()
return self.success(data={})
@@ -206,8 +284,8 @@ class AdaptersRouterGroup(group.RouterGroup):
if session and session.get('task') and not session['task'].done():
session['task'].cancel()
- @self.route('/weixin/login', methods=['POST'])
- async def _() -> str:
+ @self.route('/weixin/login', methods=['POST'], permission=Permission.RESOURCE_MANAGE)
+ async def _(request_context: RequestContext) -> str:
"""Start WeChat QR code login. Returns session_id + QR code data URL."""
import uuid
import time
@@ -229,6 +307,7 @@ class AdaptersRouterGroup(group.RouterGroup):
'error': None,
'created_at': time.time(),
}
+ _bind_session_scope(session, request_context)
_weixin_login_sessions[session_id] = session
client = OpenClawWeixinClient(
@@ -290,10 +369,15 @@ class AdaptersRouterGroup(group.RouterGroup):
}
)
- @self.route('/weixin/login/status/', methods=['GET'])
- async def _(session_id: str) -> str:
+ @self.route(
+ '/weixin/login/status/',
+ methods=['GET'],
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(session_id: str, request_context: RequestContext) -> str:
"""Poll WeChat login status."""
- session = _weixin_login_sessions.get(session_id)
+ _cleanup_expired_weixin_sessions()
+ session = _get_owned_session(_weixin_login_sessions, session_id, request_context)
if not session:
return self.http_status(404, -1, 'Session not found')
@@ -317,10 +401,16 @@ class AdaptersRouterGroup(group.RouterGroup):
return self.success(data=data)
- @self.route('/weixin/login/', methods=['DELETE'])
- async def _(session_id: str) -> str:
+ @self.route(
+ '/weixin/login/',
+ methods=['DELETE'],
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(session_id: str, request_context: RequestContext) -> str:
"""Cancel and clean up a WeChat login session."""
- session = _weixin_login_sessions.pop(session_id, None)
+ session = _pop_owned_session(_weixin_login_sessions, session_id, request_context)
+ if session is None:
+ return self.http_status(404, -1, 'Session not found')
if session and session.get('task') and not session['task'].done():
session['task'].cancel()
return self.success(data={})
@@ -344,8 +434,8 @@ class AdaptersRouterGroup(group.RouterGroup):
if session and session.get('task') and not session['task'].done():
session['task'].cancel()
- @self.route('/dingtalk/create-app', methods=['POST'])
- async def _() -> str:
+ @self.route('/dingtalk/create-app', methods=['POST'], permission=Permission.RESOURCE_MANAGE)
+ async def _(request_context: RequestContext) -> str:
"""Start DingTalk one-click app creation via Device Flow. Returns session_id + QR code URL."""
import uuid
import time
@@ -368,6 +458,7 @@ class AdaptersRouterGroup(group.RouterGroup):
'device_code': None,
'interval': 5,
}
+ _bind_session_scope(session, request_context)
_dingtalk_sessions[session_id] = session
async def run_device_flow():
@@ -491,11 +582,15 @@ class AdaptersRouterGroup(group.RouterGroup):
}
)
- @self.route('/dingtalk/create-app/status/', methods=['GET'])
- async def _(session_id: str) -> str:
+ @self.route(
+ '/dingtalk/create-app/status/',
+ methods=['GET'],
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(session_id: str, request_context: RequestContext) -> str:
"""Poll DingTalk Device Flow status."""
_cleanup_expired_dingtalk_sessions()
- session = _dingtalk_sessions.get(session_id)
+ session = _get_owned_session(_dingtalk_sessions, session_id, request_context)
if not session:
return self.http_status(404, -1, 'Session not found')
@@ -511,10 +606,16 @@ class AdaptersRouterGroup(group.RouterGroup):
return self.success(data=data)
- @self.route('/dingtalk/create-app/', methods=['DELETE'])
- async def _(session_id: str) -> str:
+ @self.route(
+ '/dingtalk/create-app/',
+ methods=['DELETE'],
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(session_id: str, request_context: RequestContext) -> str:
"""Cancel and clean up a DingTalk Device Flow session."""
- session = _dingtalk_sessions.pop(session_id, None)
+ session = _pop_owned_session(_dingtalk_sessions, session_id, request_context)
+ if session is None:
+ return self.http_status(404, -1, 'Session not found')
if session and session.get('task') and not session['task'].done():
session['task'].cancel()
return self.success(data={})
@@ -538,8 +639,8 @@ class AdaptersRouterGroup(group.RouterGroup):
if session and session.get('task') and not session['task'].done():
session['task'].cancel()
- @self.route('/wecombot/create-bot', methods=['POST'])
- async def _() -> str:
+ @self.route('/wecombot/create-bot', methods=['POST'], permission=Permission.RESOURCE_MANAGE)
+ async def _(request_context: RequestContext) -> str:
"""Start WeComBot one-click creation via QR code. Returns session_id + QR code URL."""
import uuid
import time
@@ -563,6 +664,7 @@ class AdaptersRouterGroup(group.RouterGroup):
'scode': None,
'task': None,
}
+ _bind_session_scope(session, request_context)
_wecombot_sessions[session_id] = session
async def run_qr_flow():
@@ -655,11 +757,15 @@ class AdaptersRouterGroup(group.RouterGroup):
}
)
- @self.route('/wecombot/create-bot/status/', methods=['GET'])
- async def _(session_id: str) -> str:
+ @self.route(
+ '/wecombot/create-bot/status/',
+ methods=['GET'],
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(session_id: str, request_context: RequestContext) -> str:
"""Poll WeComBot creation status."""
_cleanup_expired_wecombot_sessions()
- session = _wecombot_sessions.get(session_id)
+ session = _get_owned_session(_wecombot_sessions, session_id, request_context)
if not session:
return self.http_status(404, -1, 'Session not found')
@@ -675,10 +781,16 @@ class AdaptersRouterGroup(group.RouterGroup):
return self.success(data=data)
- @self.route('/wecombot/create-bot/', methods=['DELETE'])
- async def _(session_id: str) -> str:
+ @self.route(
+ '/wecombot/create-bot/',
+ methods=['DELETE'],
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(session_id: str, request_context: RequestContext) -> str:
"""Cancel and clean up a WeComBot creation session."""
- session = _wecombot_sessions.pop(session_id, None)
+ session = _pop_owned_session(_wecombot_sessions, session_id, request_context)
+ if session is None:
+ return self.http_status(404, -1, 'Session not found')
if session and session.get('task') and not session['task'].done():
session['task'].cancel()
return self.success(data={})
@@ -702,8 +814,8 @@ class AdaptersRouterGroup(group.RouterGroup):
if session and session.get('task') and not session['task'].done():
session['task'].cancel()
- @self.route('/qqofficial/bind', methods=['POST'])
- async def _() -> str:
+ @self.route('/qqofficial/bind', methods=['POST'], permission=Permission.RESOURCE_MANAGE)
+ async def _(request_context: RequestContext) -> str:
"""Start QQ Official QR binding. Returns session_id + QR URL.
Flow: generate a local AES-256 key, register it with
@@ -739,6 +851,7 @@ class AdaptersRouterGroup(group.RouterGroup):
'bind_key_bytes': bind_key_bytes,
'interval': 2,
}
+ _bind_session_scope(session, request_context)
_qqofficial_sessions[session_id] = session
async def run_qr_binding():
@@ -870,11 +983,15 @@ class AdaptersRouterGroup(group.RouterGroup):
}
)
- @self.route('/qqofficial/bind/status/', methods=['GET'])
- async def _(session_id: str) -> str:
+ @self.route(
+ '/qqofficial/bind/status/',
+ methods=['GET'],
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(session_id: str, request_context: RequestContext) -> str:
"""Poll QQ Official QR binding status."""
_cleanup_expired_qqofficial_sessions()
- session = _qqofficial_sessions.get(session_id)
+ session = _get_owned_session(_qqofficial_sessions, session_id, request_context)
if not session:
return self.http_status(404, -1, 'Session not found')
@@ -892,10 +1009,16 @@ class AdaptersRouterGroup(group.RouterGroup):
return self.success(data=data)
- @self.route('/qqofficial/bind/', methods=['DELETE'])
- async def _(session_id: str) -> str:
+ @self.route(
+ '/qqofficial/bind/',
+ methods=['DELETE'],
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(session_id: str, request_context: RequestContext) -> str:
"""Cancel and clean up a QQ Official QR binding session."""
- session = _qqofficial_sessions.pop(session_id, None)
+ session = _pop_owned_session(_qqofficial_sessions, session_id, request_context)
+ if session is None:
+ return self.http_status(404, -1, 'Session not found')
if session and session.get('task') and not session['task'].done():
session['task'].cancel()
return self.success(data={})
diff --git a/src/langbot/pkg/api/http/controller/groups/platform/bots.py b/src/langbot/pkg/api/http/controller/groups/platform/bots.py
index e3a13b789..373088024 100644
--- a/src/langbot/pkg/api/http/controller/groups/platform/bots.py
+++ b/src/langbot/pkg/api/http/controller/groups/platform/bots.py
@@ -1,45 +1,95 @@
import quart
+from sqlalchemy.exc import IntegrityError
+from ....authz import Permission, has_permission
+from ....context import RequestContext
from ... import group
@group.group_class('bots', '/api/v1/platform/bots')
class BotsRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET', 'POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
- if quart.request.method == 'GET':
- return self.success(data={'bots': await self.ap.bot_service.get_bots()})
- elif quart.request.method == 'POST':
- json_data = await quart.request.json
- bot_uuid = await self.ap.bot_service.create_bot(json_data)
- return self.success(data={'uuid': bot_uuid})
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ include_secret = has_permission(request_context, Permission.RESOURCE_MANAGE)
+ return self.success(
+ data={
+ 'bots': await self.ap.bot_service.get_bots(
+ request_context,
+ include_secret=include_secret,
+ )
+ }
+ )
- @self.route('/', methods=['GET', 'PUT', 'DELETE'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _(bot_uuid: str) -> str:
- if quart.request.method == 'GET':
- bot = await self.ap.bot_service.get_runtime_bot_info(bot_uuid)
- if bot is None:
- return self.http_status(404, -1, 'bot not found')
- return self.success(data={'bot': bot})
- elif quart.request.method == 'PUT':
- json_data = await quart.request.json
- await self.ap.bot_service.update_bot(bot_uuid, json_data)
- return self.success()
- elif quart.request.method == 'DELETE':
- await self.ap.bot_service.delete_bot(bot_uuid)
- return self.success()
+ @self.route(
+ '',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
+ json_data = await quart.request.json
+ bot_uuid = await self.ap.bot_service.create_bot(request_context, json_data)
+ return self.success(data={'uuid': bot_uuid})
- @self.route('//logs', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _(bot_uuid: str) -> str:
+ @self.route(
+ '/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(bot_uuid: str, request_context: RequestContext) -> str:
+ include_secret = has_permission(request_context, Permission.RESOURCE_MANAGE)
+ bot = await self.ap.bot_service.get_runtime_bot_info(
+ request_context,
+ bot_uuid,
+ include_secret=include_secret,
+ )
+ if bot is None:
+ return self.http_status(404, -1, 'bot not found')
+ return self.success(data={'bot': bot})
+
+ @self.route(
+ '/',
+ methods=['PUT', 'DELETE'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(bot_uuid: str, request_context: RequestContext) -> str:
+ if quart.request.method == 'PUT':
+ json_data = await quart.request.json
+ await self.ap.bot_service.update_bot(request_context, bot_uuid, json_data)
+ else:
+ await self.ap.bot_service.delete_bot(request_context, bot_uuid)
+ return self.success()
+
+ @self.route(
+ '//logs',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.AUDIT_VIEW,
+ )
+ async def _(bot_uuid: str, request_context: RequestContext) -> str:
json_data = await quart.request.json
from_index = json_data.get('from_index', -1)
max_count = json_data.get('max_count', 10)
- logs, total_count = await self.ap.bot_service.list_event_logs(bot_uuid, from_index, max_count)
+ logs, total_count = await self.ap.bot_service.list_event_logs(
+ request_context, bot_uuid, from_index, max_count
+ )
return self.success(data={'logs': logs, 'total_count': total_count})
- @self.route('//send_message', methods=['POST'], auth_type=group.AuthType.API_KEY)
- async def _(bot_uuid: str) -> str:
+ @self.route(
+ '//send_message',
+ methods=['POST'],
+ auth_type=group.AuthType.API_KEY,
+ permission=Permission.RUNTIME_OPERATE,
+ )
+ async def _(bot_uuid: str, request_context: RequestContext) -> str:
json_data = await quart.request.json
target_type = json_data.get('target_type')
target_id = json_data.get('target_id')
@@ -54,37 +104,51 @@ class BotsRouterGroup(group.RouterGroup):
if target_type not in ['person', 'group']:
return self.http_status(400, -1, 'target_type must be either "person" or "group"')
- try:
- await self.ap.bot_service.send_message(bot_uuid, target_type, target_id, message_chain_data)
- return self.success(data={'sent': True})
- except Exception as e:
- import traceback
-
- traceback.print_exc()
- return self.http_status(500, -1, f'Failed to send message: {str(e)}')
-
- # ============ Bot Admins ============
-
- @self.route('//admins', methods=['GET', 'POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _(bot_uuid: str) -> str:
- if quart.request.method == 'GET':
- admins = await self.ap.bot_service.get_bot_admins(bot_uuid)
- return self.success(data={'admins': admins})
- elif quart.request.method == 'POST':
- json_data = await quart.request.json
- launcher_type = json_data.get('launcher_type', '').strip()
- launcher_id = str(json_data.get('launcher_id', '')).strip()
- if not launcher_type or not launcher_id:
- return self.http_status(400, -1, 'launcher_type and launcher_id are required')
- try:
- admin_id = await self.ap.bot_service.add_bot_admin(bot_uuid, launcher_type, launcher_id)
- return self.success(data={'id': admin_id})
- except Exception as e:
- return self.http_status(409, -1, str(e))
+ await self.ap.bot_service.send_message(
+ request_context,
+ bot_uuid,
+ target_type,
+ target_id,
+ message_chain_data,
+ )
+ return self.success(data={'sent': True})
@self.route(
- '//admins/', methods=['DELETE'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY
+ '//admins',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def _(bot_uuid: str, admin_id: int) -> str:
- await self.ap.bot_service.delete_bot_admin(bot_uuid, admin_id)
+ async def _(bot_uuid: str, request_context: RequestContext) -> str:
+ admins = await self.ap.bot_service.get_bot_admins(request_context, bot_uuid)
+ return self.success(data={'admins': admins})
+
+ @self.route(
+ '//admins',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(bot_uuid: str, request_context: RequestContext) -> str:
+ json_data = await quart.request.json
+ launcher_type = json_data.get('launcher_type', '').strip()
+ launcher_id = str(json_data.get('launcher_id', '')).strip()
+ if not launcher_type or not launcher_id:
+ return self.http_status(400, -1, 'launcher_type and launcher_id are required')
+ try:
+ admin_id = await self.ap.bot_service.add_bot_admin(
+ request_context, bot_uuid, launcher_type, launcher_id
+ )
+ return self.success(data={'id': admin_id})
+ except IntegrityError as e:
+ return self.http_status(409, -1, str(e))
+
+ @self.route(
+ '//admins/',
+ methods=['DELETE'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(bot_uuid: str, admin_id: int, request_context: RequestContext) -> str:
+ await self.ap.bot_service.delete_bot_admin(request_context, bot_uuid, admin_id)
return self.success()
diff --git a/src/langbot/pkg/api/http/controller/groups/plugins.py b/src/langbot/pkg/api/http/controller/groups/plugins.py
index c291c1232..e52677336 100644
--- a/src/langbot/pkg/api/http/controller/groups/plugins.py
+++ b/src/langbot/pkg/api/http/controller/groups/plugins.py
@@ -1,6 +1,8 @@
from __future__ import annotations
import base64
+import collections.abc
+import copy
import io
import quart
import re
@@ -15,9 +17,139 @@ import sqlalchemy
from .....core import taskmgr
from .....entity.persistence import plugin as persistence_plugin
+from ...authz import Permission
+from ...context import ExecutionContext, RequestContext
from .. import group
+from .....workspace.errors import WorkspaceNotFoundError
from langbot_plugin.runtime.plugin.mgr import PluginInstallSource
+
+_SECRET_MASK = '***'
+_MISSING_SECRET = object()
+_SENSITIVE_CONFIG_NAMES = frozenset(
+ {
+ 'api_key',
+ 'apikey',
+ 'auth',
+ 'authorization',
+ 'cookie',
+ 'credentials',
+ 'database_url',
+ 'dsn',
+ 'key',
+ 'proxy_authorization',
+ 'set_cookie',
+ }
+)
+_SENSITIVE_CONFIG_TOKENS = frozenset(
+ {
+ 'credential',
+ 'credentials',
+ 'passwd',
+ 'password',
+ 'secret',
+ 'token',
+ }
+)
+_SENSITIVE_KEY_QUALIFIERS = frozenset(
+ {
+ 'access',
+ 'api',
+ 'auth',
+ 'bearer',
+ 'client',
+ 'debug',
+ 'encryption',
+ 'private',
+ 'signing',
+ }
+)
+
+
+def _normalize_config_key(key: object) -> str:
+ value = re.sub(r'([a-z0-9])([A-Z])', r'\1_\2', str(key or ''))
+ return re.sub(r'[^a-zA-Z0-9]+', '_', value).strip('_').lower()
+
+
+def _is_sensitive_config_key(key: object) -> bool:
+ normalized = _normalize_config_key(key)
+ if normalized in _SENSITIVE_CONFIG_NAMES:
+ return True
+ tokens = frozenset(token for token in normalized.split('_') if token)
+ if tokens & _SENSITIVE_CONFIG_TOKENS:
+ return True
+ return 'key' in tokens and bool(tokens & _SENSITIVE_KEY_QUALIFIERS)
+
+
+def _mask_secret_structure(value):
+ """Mask every non-empty leaf while preserving container structure."""
+
+ if isinstance(value, dict):
+ return {key: _mask_secret_structure(item) for key, item in value.items()}
+ if isinstance(value, list):
+ return [_mask_secret_structure(item) for item in value]
+ if isinstance(value, tuple):
+ return tuple(_mask_secret_structure(item) for item in value)
+ if value is None or value == '':
+ return value
+ return _SECRET_MASK
+
+
+def redact_plugin_secrets(value):
+ """Return a recursively redacted copy of plugin-facing data."""
+
+ if isinstance(value, dict):
+ return {
+ key: (_mask_secret_structure(item) if _is_sensitive_config_key(key) else redact_plugin_secrets(item))
+ for key, item in value.items()
+ }
+ if isinstance(value, list):
+ return [redact_plugin_secrets(item) for item in value]
+ if isinstance(value, tuple):
+ return tuple(redact_plugin_secrets(item) for item in value)
+ return value
+
+
+def restore_plugin_secret_placeholders(value, current_value=_MISSING_SECRET, *, sensitive: bool = False):
+ """Restore masked leaves from the current config before a management write."""
+
+ if sensitive and value == _SECRET_MASK:
+ if current_value is _MISSING_SECRET:
+ raise ValueError('Masked plugin secret has no existing value')
+ return copy.deepcopy(current_value)
+ if isinstance(value, dict):
+ current_mapping = current_value if isinstance(current_value, dict) else {}
+ return {
+ key: restore_plugin_secret_placeholders(
+ item,
+ current_mapping.get(key, _MISSING_SECRET),
+ sensitive=sensitive or _is_sensitive_config_key(key),
+ )
+ for key, item in value.items()
+ }
+ if isinstance(value, list):
+ current_items = current_value if isinstance(current_value, (list, tuple)) else ()
+ return [
+ restore_plugin_secret_placeholders(
+ item,
+ current_items[index] if index < len(current_items) else _MISSING_SECRET,
+ sensitive=sensitive,
+ )
+ for index, item in enumerate(value)
+ ]
+ if isinstance(value, tuple):
+ current_items = current_value if isinstance(current_value, (list, tuple)) else ()
+ return tuple(
+ restore_plugin_secret_placeholders(
+ item,
+ current_items[index] if index < len(current_items) else _MISSING_SECRET,
+ sensitive=sensitive,
+ )
+ for index, item in enumerate(value)
+ )
+ return value
+
+
# Resolve the built-in page SDK JS from the langbot_plugin package
_PAGE_SDK_PATH = None
try:
@@ -148,18 +280,74 @@ class PluginsRouterGroup(group.RouterGroup):
'subdir': subdir,
}
- async def _check_extensions_limit(self) -> str | None:
+ async def _check_extensions_limit(self, request_context: RequestContext) -> str | None:
"""Check if extensions limit is reached. Returns error response if limit exceeded, None otherwise."""
+ await self.ap.plugin_connector.require_workspace_context(request_context)
limitation = self.ap.instance_config.data.get('system', {}).get('limitation', {})
max_extensions = limitation.get('max_extensions', -1)
if max_extensions >= 0:
plugins = await self.ap.plugin_connector.list_plugins()
- mcp_servers = await self.ap.mcp_service.get_mcp_servers()
+ mcp_servers = await self.ap.mcp_service.get_mcp_servers(request_context)
total_extensions = len(plugins) + len(mcp_servers)
if total_extensions >= max_extensions:
return self.http_status(400, -1, f'Maximum number of extensions ({max_extensions}) reached')
return None
+ @staticmethod
+ def _task_scope(request_context: RequestContext) -> dict[str, str | int]:
+ return {
+ 'instance_uuid': request_context.instance_uuid,
+ 'workspace_uuid': request_context.workspace_uuid,
+ 'placement_generation': request_context.placement_generation,
+ }
+
+ async def _run_fenced_plugin_operation(
+ self,
+ execution_context: ExecutionContext,
+ operation: collections.abc.Callable[[], collections.abc.Awaitable],
+ ):
+ """Revalidate a captured task context immediately before Runtime I/O."""
+
+ await self.ap.plugin_connector.require_workspace_context(execution_context)
+ return await operation()
+
+ async def _require_public_plugin_runtime_context(self) -> ExecutionContext:
+ """Resolve public assets only for the OSS singleton Workspace.
+
+ Public image and iframe requests cannot carry the WebUI bearer token.
+ They therefore remain available for the one-Workspace Core deployment,
+ but fail closed instead of guessing a Workspace when multi-Workspace
+ policy is active.
+ """
+
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ policy = getattr(workspace_service, 'policy', None)
+ if workspace_service is None or policy is None or getattr(policy, 'multi_workspace_enabled', False):
+ raise WorkspaceNotFoundError('Plugin resource not found')
+ binding = await workspace_service.get_local_execution_binding()
+ execution_context = ExecutionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ )
+ return await self.ap.plugin_connector.require_workspace_context(execution_context)
+
+ async def _get_stored_plugin_config(
+ self,
+ request_context: RequestContext,
+ author: str,
+ plugin_name: str,
+ plugin: dict,
+ ):
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(persistence_plugin.PluginSetting.config)
+ .where(persistence_plugin.PluginSetting.workspace_uuid == request_context.workspace_uuid)
+ .where(persistence_plugin.PluginSetting.plugin_author == author)
+ .where(persistence_plugin.PluginSetting.plugin_name == plugin_name)
+ )
+ persisted_config = result.scalar_one_or_none()
+ return persisted_config if persisted_config is not None else plugin['plugin_config']
+
async def initialize(self) -> None:
@self.route('/_sdk/page-sdk.js', methods=['GET'], auth_type=group.AuthType.NONE)
async def _() -> quart.Response:
@@ -170,15 +358,27 @@ class PluginsRouterGroup(group.RouterGroup):
return quart.Response(content, mimetype='application/javascript')
return quart.Response('// SDK not found', status=404, mimetype='application/javascript')
- @self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ await self.ap.plugin_connector.require_workspace_context(request_context)
plugins = await self.ap.plugin_connector.list_plugins()
- return self.success(data={'plugins': plugins})
+ return self.success(data={'plugins': redact_plugin_secrets(plugins)})
- @self.route('/debug-info', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
+ @self.route(
+ '/debug-info',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
"""Get plugin debug information including debug URL and key"""
+ await self.ap.plugin_connector.require_workspace_context(request_context)
debug_info = await self.ap.plugin_connector.get_debug_info()
# Get debug URL from config
@@ -196,77 +396,121 @@ class PluginsRouterGroup(group.RouterGroup):
'///upgrade',
methods=['POST'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
)
- async def _(author: str, plugin_name: str) -> str:
+ async def _(author: str, plugin_name: str, request_context: RequestContext) -> str:
+ execution_context = await self.ap.plugin_connector.require_workspace_context(request_context)
ctx = taskmgr.TaskContext.new()
wrapper = self.ap.task_mgr.create_user_task(
- self.ap.plugin_connector.upgrade_plugin(author, plugin_name, task_context=ctx),
+ self._run_fenced_plugin_operation(
+ execution_context,
+ lambda: self.ap.plugin_connector.upgrade_plugin(author, plugin_name, task_context=ctx),
+ ),
kind='plugin-operation',
name=f'plugin-upgrade-{plugin_name}',
label=f'Upgrading plugin {plugin_name}',
context=ctx,
+ **self._task_scope(request_context),
)
return self.success(data={'task_id': wrapper.id})
@self.route(
'//',
- methods=['GET', 'DELETE'],
+ methods=['GET'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def _(author: str, plugin_name: str) -> str:
- if quart.request.method == 'GET':
- plugin = await self.ap.plugin_connector.get_plugin_info(author, plugin_name)
- if plugin is None:
- return self.http_status(404, -1, 'plugin not found')
- return self.success(data={'plugin': plugin})
- elif quart.request.method == 'DELETE':
- delete_data = quart.request.args.get('delete_data', 'false').lower() == 'true'
- ctx = taskmgr.TaskContext.new()
- wrapper = self.ap.task_mgr.create_user_task(
- self.ap.plugin_connector.delete_plugin(
- author, plugin_name, delete_data=delete_data, task_context=ctx
- ),
- kind='plugin-operation',
- name=f'plugin-remove-{plugin_name}',
- label=f'Removing plugin {plugin_name}',
- context=ctx,
- )
+ async def _(author: str, plugin_name: str, request_context: RequestContext) -> str:
+ await self.ap.plugin_connector.require_workspace_context(request_context)
+ plugin = await self.ap.plugin_connector.get_plugin_info(author, plugin_name)
+ if plugin is None:
+ return self.http_status(404, -1, 'plugin not found')
+ return self.success(data={'plugin': redact_plugin_secrets(plugin)})
- return self.success(data={'task_id': wrapper.id})
+ @self.route(
+ '//',
+ methods=['DELETE'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(author: str, plugin_name: str, request_context: RequestContext) -> str:
+ execution_context = await self.ap.plugin_connector.require_workspace_context(request_context)
+ delete_data = quart.request.args.get('delete_data', 'false').lower() == 'true'
+ ctx = taskmgr.TaskContext.new()
+ wrapper = self.ap.task_mgr.create_user_task(
+ self._run_fenced_plugin_operation(
+ execution_context,
+ lambda: self.ap.plugin_connector.delete_plugin(
+ author,
+ plugin_name,
+ delete_data=delete_data,
+ task_context=ctx,
+ ),
+ ),
+ kind='plugin-operation',
+ name=f'plugin-remove-{plugin_name}',
+ label=f'Removing plugin {plugin_name}',
+ context=ctx,
+ **self._task_scope(request_context),
+ )
+ return self.success(data={'task_id': wrapper.id})
@self.route(
'///config',
- methods=['GET', 'PUT'],
+ methods=['GET'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def _(author: str, plugin_name: str) -> quart.Response:
+ async def _(author: str, plugin_name: str, request_context: RequestContext) -> quart.Response:
+ await self.ap.plugin_connector.require_workspace_context(request_context)
plugin = await self.ap.plugin_connector.get_plugin_info(author, plugin_name)
if plugin is None:
return self.http_status(404, -1, 'plugin not found')
- if quart.request.method == 'GET':
- result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_plugin.PluginSetting.config)
- .where(persistence_plugin.PluginSetting.plugin_author == author)
- .where(persistence_plugin.PluginSetting.plugin_name == plugin_name)
+ config = await self._get_stored_plugin_config(
+ request_context,
+ author,
+ plugin_name,
+ plugin,
+ )
+ return self.success(data={'config': redact_plugin_secrets(config)})
+
+ @self.route(
+ '///config',
+ methods=['PUT'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(author: str, plugin_name: str, request_context: RequestContext) -> quart.Response:
+ await self.ap.plugin_connector.require_workspace_context(request_context)
+ plugin = await self.ap.plugin_connector.get_plugin_info(author, plugin_name)
+ if plugin is None:
+ return self.http_status(404, -1, 'plugin not found')
+ current_config = await self._get_stored_plugin_config(
+ request_context,
+ author,
+ plugin_name,
+ plugin,
+ )
+ try:
+ config = restore_plugin_secret_placeholders(
+ await quart.request.json,
+ current_config,
)
- persisted_config = result.scalar_one_or_none()
-
- config = persisted_config if persisted_config is not None else plugin['plugin_config']
- return self.success(data={'config': config})
- elif quart.request.method == 'PUT':
- data = await quart.request.json
-
- await self.ap.plugin_connector.set_plugin_config(author, plugin_name, data)
-
- return self.success(data={})
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ await self.ap.plugin_connector.require_workspace_context(request_context)
+ await self.ap.plugin_connector.set_plugin_config(author, plugin_name, config)
+ return self.success(data={})
@self.route(
'///readme',
methods=['GET'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def _(author: str, plugin_name: str) -> quart.Response:
+ async def _(author: str, plugin_name: str, request_context: RequestContext) -> quart.Response:
+ await self.ap.plugin_connector.require_workspace_context(request_context)
language = quart.request.args.get('language', 'en')
readme = await self.ap.plugin_connector.get_plugin_readme(author, plugin_name, language=language)
return self.success(data={'readme': readme})
@@ -275,8 +519,10 @@ class PluginsRouterGroup(group.RouterGroup):
'///logs',
methods=['GET'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.AUDIT_VIEW,
)
- async def _(author: str, plugin_name: str) -> quart.Response:
+ async def _(author: str, plugin_name: str, request_context: RequestContext) -> quart.Response:
+ await self.ap.plugin_connector.require_workspace_context(request_context)
try:
limit = int(quart.request.args.get('limit', 200))
except (TypeError, ValueError):
@@ -291,6 +537,7 @@ class PluginsRouterGroup(group.RouterGroup):
auth_type=group.AuthType.NONE,
)
async def _(author: str, plugin_name: str) -> quart.Response:
+ await self._require_public_plugin_runtime_context()
icon_data = await self.ap.plugin_connector.get_plugin_icon(author, plugin_name)
icon_base64 = icon_data['plugin_icon_base64']
mime_type = icon_data['mime_type']
@@ -305,6 +552,7 @@ class PluginsRouterGroup(group.RouterGroup):
auth_type=group.AuthType.NONE,
)
async def _(author: str, plugin_name: str, filepath: str) -> quart.Response:
+ await self._require_public_plugin_runtime_context()
asset_path = _normalize_plugin_asset_path(filepath)
if asset_path is None:
return quart.Response('Asset not found', status=404)
@@ -334,9 +582,11 @@ class PluginsRouterGroup(group.RouterGroup):
'///page-api',
methods=['POST'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
)
- async def _(author: str, plugin_name: str) -> str:
+ async def _(author: str, plugin_name: str, request_context: RequestContext) -> str:
"""Forward a page API request to the plugin."""
+ await self.ap.plugin_connector.require_workspace_context(request_context)
data = await quart.request.json
if not isinstance(data, dict):
return self.http_status(400, -1, 'invalid request body')
@@ -357,9 +607,15 @@ class PluginsRouterGroup(group.RouterGroup):
return self.http_status(400, -1, result['error'])
return self.success(data=result.get('data'))
- @self.route('/github/releases', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
+ @self.route(
+ '/github/releases',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
"""Get releases from a GitHub repository URL"""
+ await self.ap.plugin_connector.require_workspace_context(request_context)
data = await quart.request.json
repo_url = data.get('repo_url', '')
@@ -427,16 +683,18 @@ class PluginsRouterGroup(group.RouterGroup):
'source_subdir': requested_subdir,
}
)
- except httpx.RequestError as e:
- return self.http_status(500, -1, f'Failed to fetch releases: {str(e)}')
+ except httpx.RequestError:
+ raise
@self.route(
'/github/release-assets',
methods=['POST'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def _() -> str:
+ async def _(request_context: RequestContext) -> str:
"""Get assets from a specific GitHub release"""
+ await self.ap.plugin_connector.require_workspace_context(request_context)
data = await quart.request.json
owner = data.get('owner', '')
repo = data.get('repo', '')
@@ -484,13 +742,18 @@ class PluginsRouterGroup(group.RouterGroup):
# )
return self.success(data={'assets': formatted_assets})
- except httpx.RequestError as e:
- return self.http_status(500, -1, f'Failed to fetch release assets: {str(e)}')
+ except httpx.RequestError:
+ raise
- @self.route('/install/github', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
+ @self.route(
+ '/install/github',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
"""Install plugin from GitHub release asset"""
- limit_error = await self._check_extensions_limit()
+ limit_error = await self._check_extensions_limit(request_context)
if limit_error is not None:
return limit_error
@@ -503,6 +766,8 @@ class PluginsRouterGroup(group.RouterGroup):
if not asset_url:
return self.http_status(400, -1, 'Missing asset_url parameter')
+ execution_context = await self.ap.plugin_connector.require_workspace_context(request_context)
+
ctx = taskmgr.TaskContext.new()
ctx.metadata['plugin_name'] = f'{owner}/{repo}'
ctx.metadata['install_source'] = 'github'
@@ -515,11 +780,19 @@ class PluginsRouterGroup(group.RouterGroup):
}
wrapper = self.ap.task_mgr.create_user_task(
- self.ap.plugin_connector.install_plugin(PluginInstallSource.GITHUB, install_info, task_context=ctx),
+ self._run_fenced_plugin_operation(
+ execution_context,
+ lambda: self.ap.plugin_connector.install_plugin(
+ PluginInstallSource.GITHUB,
+ install_info,
+ task_context=ctx,
+ ),
+ ),
kind='plugin-operation',
name='plugin-install-github',
label=f'Installing plugin from GitHub {owner}/{repo}@{release_tag}',
context=ctx,
+ **self._task_scope(request_context),
)
return self.success(data={'task_id': wrapper.id})
@@ -528,9 +801,10 @@ class PluginsRouterGroup(group.RouterGroup):
'/install/marketplace',
methods=['POST'],
auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
)
- async def _() -> str:
- limit_error = await self._check_extensions_limit()
+ async def _(request_context: RequestContext) -> str:
+ limit_error = await self._check_extensions_limit(request_context)
if limit_error is not None:
return limit_error
@@ -538,23 +812,37 @@ class PluginsRouterGroup(group.RouterGroup):
plugin_author = data.get('plugin_author', '')
plugin_name = data.get('plugin_name', '')
+ execution_context = await self.ap.plugin_connector.require_workspace_context(request_context)
ctx = taskmgr.TaskContext.new()
ctx.metadata['plugin_name'] = f'{plugin_author}/{plugin_name}'
ctx.metadata['install_source'] = 'marketplace'
wrapper = self.ap.task_mgr.create_user_task(
- self.ap.plugin_connector.install_plugin(PluginInstallSource.MARKETPLACE, data, task_context=ctx),
+ self._run_fenced_plugin_operation(
+ execution_context,
+ lambda: self.ap.plugin_connector.install_plugin(
+ PluginInstallSource.MARKETPLACE,
+ data,
+ task_context=ctx,
+ ),
+ ),
kind='plugin-operation',
name='plugin-install-marketplace',
label=f'Installing plugin from marketplace {plugin_author}/{plugin_name}',
context=ctx,
+ **self._task_scope(request_context),
)
return self.success(data={'task_id': wrapper.id})
- @self.route('/install/local', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
- limit_error = await self._check_extensions_limit()
+ @self.route(
+ '/install/local',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
+ limit_error = await self._check_extensions_limit(request_context)
if limit_error is not None:
return limit_error
@@ -563,6 +851,7 @@ class PluginsRouterGroup(group.RouterGroup):
return self.http_status(400, -1, 'file is required')
file_bytes = file.read()
+ execution_context = await self.ap.plugin_connector.require_workspace_context(request_context)
data = {
'plugin_file': file_bytes,
@@ -572,17 +861,31 @@ class PluginsRouterGroup(group.RouterGroup):
ctx.metadata['plugin_name'] = file.filename or 'local plugin'
ctx.metadata['install_source'] = 'local'
wrapper = self.ap.task_mgr.create_user_task(
- self.ap.plugin_connector.install_plugin(PluginInstallSource.LOCAL, data, task_context=ctx),
+ self._run_fenced_plugin_operation(
+ execution_context,
+ lambda: self.ap.plugin_connector.install_plugin(
+ PluginInstallSource.LOCAL,
+ data,
+ task_context=ctx,
+ ),
+ ),
kind='plugin-operation',
name='plugin-install-local',
label=f'Installing plugin from local {file.filename}',
context=ctx,
+ **self._task_scope(request_context),
)
return self.success(data={'task_id': wrapper.id})
- @self.route('/install/local/preview', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
+ @self.route(
+ '/install/local/preview',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
+ await self.ap.plugin_connector.require_workspace_context(request_context)
file = (await quart.request.files).get('file')
if file is None:
return self.http_status(400, -1, 'file is required')
@@ -634,12 +937,18 @@ class PluginsRouterGroup(group.RouterGroup):
)
except zipfile.BadZipFile:
return self.http_status(400, -1, 'invalid .lbpkg file')
- except Exception as exc:
- return self.http_status(500, -1, f'Failed to preview plugin package: {exc}')
+ except Exception:
+ raise
- @self.route('/config-files', methods=['POST'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
+ @self.route(
+ '/config-files',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
"""Upload a file for plugin configuration"""
+ await self.ap.plugin_connector.require_workspace_context(request_context)
file = (await quart.request.files).get('file')
if file is None:
return self.http_status(400, -1, 'file is required')
@@ -650,25 +959,37 @@ class PluginsRouterGroup(group.RouterGroup):
if len(file_bytes) > MAX_FILE_SIZE:
return self.http_status(400, -1, 'file size exceeds 10MB limit')
- # Generate unique file key with original extension
- original_filename = file.filename
+ original_filename = file.filename or 'config.bin'
_, ext = os.path.splitext(original_filename)
- file_key = f'plugin_config_{uuid.uuid4().hex}{ext}'
-
- # Save file using storage manager
- await self.ap.storage_mgr.storage_provider.save(file_key, file_bytes)
+ logical_key = f'plugin_config_{uuid.uuid4().hex}{ext}'
+ file_key = await self.ap.storage_mgr.save_scoped(
+ request_context,
+ owner_type='plugin_config',
+ owner=request_context.workspace_uuid,
+ key=logical_key,
+ value=file_bytes,
+ )
return self.success(data={'file_key': file_key})
- @self.route('/config-files/', methods=['DELETE'], auth_type=group.AuthType.USER_TOKEN)
- async def _(file_key: str) -> str:
+ @self.route(
+ '/config-files/',
+ methods=['DELETE'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(file_key: str, request_context: RequestContext) -> str:
"""Delete a plugin configuration file"""
- # Only allow deletion of files with plugin_config_ prefix for security
- if not file_key.startswith('plugin_config_'):
+ await self.ap.plugin_connector.require_workspace_context(request_context)
+ if not self.ap.storage_mgr.is_scoped_object_key(file_key, expected_owner_type='plugin_config'):
return self.http_status(400, -1, 'invalid file key')
try:
- await self.ap.storage_mgr.storage_provider.delete(file_key)
+ await self.ap.storage_mgr.delete_scoped_object_key(
+ request_context,
+ file_key,
+ expected_owner_type='plugin_config',
+ )
return self.success(data={'deleted': True})
- except Exception as e:
- return self.http_status(500, -1, f'failed to delete file: {str(e)}')
+ except Exception:
+ raise
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 f683c98fc..236000d9f 100644
--- a/src/langbot/pkg/api/http/controller/groups/provider/models.py
+++ b/src/langbot/pkg/api/http/controller/groups/provider/models.py
@@ -1,147 +1,292 @@
import quart
+from ....authz import Permission, has_permission
+from ....context import RequestContext
from ... import group
@group.group_class('models/llm', '/api/v1/provider/models/llm')
class LLMModelsRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET', 'POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
- if quart.request.method == 'GET':
- provider_uuid = quart.request.args.get('provider_uuid')
- if provider_uuid:
- return self.success(
- data={'models': await self.ap.llm_model_service.get_llm_models_by_provider(provider_uuid)}
- )
- return self.success(data={'models': await self.ap.llm_model_service.get_llm_models()})
- elif quart.request.method == 'POST':
- json_data = await quart.request.json
- model_uuid = await self.ap.llm_model_service.create_llm_model(json_data)
- return self.success(data={'uuid': model_uuid})
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ provider_uuid = quart.request.args.get('provider_uuid')
+ include_secret = has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE)
+ if provider_uuid:
+ models = await self.ap.llm_model_service.get_llm_models_by_provider(
+ request_context,
+ provider_uuid,
+ include_secret=include_secret,
+ )
+ else:
+ models = await self.ap.llm_model_service.get_llm_models(
+ request_context,
+ include_secret=include_secret,
+ )
+ return self.success(data={'models': models})
- @self.route('/', methods=['GET', 'PUT', 'DELETE'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _(model_uuid: str) -> str:
- if quart.request.method == 'GET':
- model = await self.ap.llm_model_service.get_llm_model(model_uuid)
+ @self.route(
+ '',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
+ try:
+ model_uuid = await self.ap.llm_model_service.create_llm_model(
+ request_context,
+ await quart.request.json,
+ )
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ return self.success(data={'uuid': model_uuid})
- if model is None:
- return self.http_status(404, -1, 'model not found')
+ @self.route(
+ '/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ model = await self.ap.llm_model_service.get_llm_model(
+ request_context,
+ model_uuid,
+ include_secret=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE),
+ )
+ if model is None:
+ return self.http_status(404, -1, 'model not found')
+ return self.success(data={'model': model})
- return self.success(data={'model': model})
- elif quart.request.method == 'PUT':
- json_data = await quart.request.json
+ @self.route(
+ '/',
+ methods=['PUT'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ try:
+ await self.ap.llm_model_service.update_llm_model(
+ request_context,
+ model_uuid,
+ await quart.request.json,
+ )
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ return self.success()
- await self.ap.llm_model_service.update_llm_model(model_uuid, json_data)
-
- return self.success()
- elif quart.request.method == 'DELETE':
- await self.ap.llm_model_service.delete_llm_model(model_uuid)
-
- return self.success()
-
- @self.route('//test', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _(model_uuid: str) -> str:
- json_data = await quart.request.json
-
- await self.ap.llm_model_service.test_llm_model(model_uuid, json_data)
+ @self.route(
+ '/',
+ methods=['DELETE'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ await self.ap.llm_model_service.delete_llm_model(request_context, model_uuid)
+ return self.success()
+ @self.route(
+ '//test',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ await self.ap.llm_model_service.test_llm_model(request_context, model_uuid, await quart.request.json)
return self.success()
@group.group_class('models/embedding', '/api/v1/provider/models/embedding')
class EmbeddingModelsRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET', 'POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
- if quart.request.method == 'GET':
- provider_uuid = quart.request.args.get('provider_uuid')
- if provider_uuid:
- return self.success(
- data={
- 'models': await self.ap.embedding_models_service.get_embedding_models_by_provider(
- provider_uuid
- )
- }
- )
- return self.success(data={'models': await self.ap.embedding_models_service.get_embedding_models()})
- elif quart.request.method == 'POST':
- json_data = await quart.request.json
- model_uuid = await self.ap.embedding_models_service.create_embedding_model(json_data)
- return self.success(data={'uuid': model_uuid})
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ provider_uuid = quart.request.args.get('provider_uuid')
+ include_secret = has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE)
+ if provider_uuid:
+ models = await self.ap.embedding_models_service.get_embedding_models_by_provider(
+ request_context,
+ provider_uuid,
+ include_secret=include_secret,
+ )
+ else:
+ models = await self.ap.embedding_models_service.get_embedding_models(
+ request_context,
+ include_secret=include_secret,
+ )
+ return self.success(data={'models': models})
- @self.route('/', methods=['GET', 'PUT', 'DELETE'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _(model_uuid: str) -> str:
- if quart.request.method == 'GET':
- model = await self.ap.embedding_models_service.get_embedding_model(model_uuid)
+ @self.route(
+ '',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
+ try:
+ model_uuid = await self.ap.embedding_models_service.create_embedding_model(
+ request_context,
+ await quart.request.json,
+ )
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ return self.success(data={'uuid': model_uuid})
- if model is None:
- return self.http_status(404, -1, 'model not found')
+ @self.route(
+ '/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ model = await self.ap.embedding_models_service.get_embedding_model(
+ request_context,
+ model_uuid,
+ include_secret=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE),
+ )
+ if model is None:
+ return self.http_status(404, -1, 'model not found')
+ return self.success(data={'model': model})
- return self.success(data={'model': model})
- elif quart.request.method == 'PUT':
- json_data = await quart.request.json
+ @self.route(
+ '/',
+ methods=['PUT'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ try:
+ await self.ap.embedding_models_service.update_embedding_model(
+ request_context,
+ model_uuid,
+ await quart.request.json,
+ )
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ return self.success()
- await self.ap.embedding_models_service.update_embedding_model(model_uuid, json_data)
-
- return self.success()
- elif quart.request.method == 'DELETE':
- await self.ap.embedding_models_service.delete_embedding_model(model_uuid)
-
- return self.success()
-
- @self.route('//test', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _(model_uuid: str) -> str:
- json_data = await quart.request.json
-
- await self.ap.embedding_models_service.test_embedding_model(model_uuid, json_data)
+ @self.route(
+ '/',
+ methods=['DELETE'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ await self.ap.embedding_models_service.delete_embedding_model(request_context, model_uuid)
+ return self.success()
+ @self.route(
+ '//test',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ await self.ap.embedding_models_service.test_embedding_model(
+ request_context, model_uuid, await quart.request.json
+ )
return self.success()
@group.group_class('models/rerank', '/api/v1/provider/models/rerank')
class RerankModelsRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET', 'POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
- if quart.request.method == 'GET':
- provider_uuid = quart.request.args.get('provider_uuid')
- if provider_uuid:
- return self.success(
- data={
- 'models': await self.ap.rerank_models_service.get_rerank_models_by_provider(provider_uuid)
- }
- )
- return self.success(data={'models': await self.ap.rerank_models_service.get_rerank_models()})
- elif quart.request.method == 'POST':
- json_data = await quart.request.json
- model_uuid = await self.ap.rerank_models_service.create_rerank_model(json_data)
- return self.success(data={'uuid': model_uuid})
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ provider_uuid = quart.request.args.get('provider_uuid')
+ include_secret = has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE)
+ if provider_uuid:
+ models = await self.ap.rerank_models_service.get_rerank_models_by_provider(
+ request_context,
+ provider_uuid,
+ include_secret=include_secret,
+ )
+ else:
+ models = await self.ap.rerank_models_service.get_rerank_models(
+ request_context,
+ include_secret=include_secret,
+ )
+ return self.success(data={'models': models})
- @self.route('/', methods=['GET', 'PUT', 'DELETE'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _(model_uuid: str) -> str:
- if quart.request.method == 'GET':
- model = await self.ap.rerank_models_service.get_rerank_model(model_uuid)
+ @self.route(
+ '',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
+ try:
+ model_uuid = await self.ap.rerank_models_service.create_rerank_model(
+ request_context,
+ await quart.request.json,
+ )
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ return self.success(data={'uuid': model_uuid})
- if model is None:
- return self.http_status(404, -1, 'model not found')
-
- return self.success(data={'model': model})
- elif quart.request.method == 'PUT':
- json_data = await quart.request.json
-
- await self.ap.rerank_models_service.update_rerank_model(model_uuid, json_data)
-
- return self.success()
- elif quart.request.method == 'DELETE':
- await self.ap.rerank_models_service.delete_rerank_model(model_uuid)
-
- return self.success()
-
- @self.route('//test', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _(model_uuid: str) -> str:
- json_data = await quart.request.json
-
- await self.ap.rerank_models_service.test_rerank_model(model_uuid, json_data)
+ @self.route(
+ '/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ model = await self.ap.rerank_models_service.get_rerank_model(
+ request_context,
+ model_uuid,
+ include_secret=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE),
+ )
+ if model is None:
+ return self.http_status(404, -1, 'model not found')
+ return self.success(data={'model': model})
+ @self.route(
+ '/',
+ methods=['PUT'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ try:
+ await self.ap.rerank_models_service.update_rerank_model(
+ request_context,
+ model_uuid,
+ await quart.request.json,
+ )
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ return self.success()
+
+ @self.route(
+ '/',
+ methods=['DELETE'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ await self.ap.rerank_models_service.delete_rerank_model(request_context, model_uuid)
+ return self.success()
+
+ @self.route(
+ '//test',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(model_uuid: str, request_context: RequestContext) -> str:
+ await self.ap.rerank_models_service.test_rerank_model(request_context, model_uuid, await quart.request.json)
return self.success()
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 fcea598f9..bf8a195ae 100644
--- a/src/langbot/pkg/api/http/controller/groups/provider/providers.py
+++ b/src/langbot/pkg/api/http/controller/groups/provider/providers.py
@@ -1,56 +1,102 @@
import quart
+from ....authz import Permission, has_permission
+from ....context import RequestContext
from ... import group
@group.group_class('models/providers', '/api/v1/provider/providers')
class ModelProvidersRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET', 'POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _() -> str:
- if quart.request.method == 'GET':
- providers = await self.ap.provider_service.get_providers()
- # Add model counts
- for provider in providers:
- counts = await self.ap.provider_service.get_provider_model_counts(provider['uuid'])
- provider['llm_count'] = counts['llm_count']
- provider['embedding_count'] = counts['embedding_count']
- provider['rerank_count'] = counts['rerank_count']
- return self.success(data={'providers': providers})
- elif quart.request.method == 'POST':
- json_data = await quart.request.json
- provider_uuid = await self.ap.provider_service.create_provider(json_data)
- return self.success(data={'uuid': provider_uuid})
-
@self.route(
- '/', methods=['GET', 'PUT', 'DELETE'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def _(provider_uuid: str) -> str:
- if quart.request.method == 'GET':
- provider = await self.ap.provider_service.get_provider(provider_uuid)
- if provider is None:
- return self.http_status(404, -1, 'provider not found')
- counts = await self.ap.provider_service.get_provider_model_counts(provider_uuid)
+ async def _(request_context: RequestContext) -> str:
+ providers = await self.ap.provider_service.get_providers(
+ request_context,
+ include_secret=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE),
+ )
+ for provider in providers:
+ counts = await self.ap.provider_service.get_provider_model_counts(request_context, provider['uuid'])
provider['llm_count'] = counts['llm_count']
provider['embedding_count'] = counts['embedding_count']
provider['rerank_count'] = counts['rerank_count']
- return self.success(data={'provider': provider})
- elif quart.request.method == 'PUT':
- json_data = await quart.request.json
- await self.ap.provider_service.update_provider(provider_uuid, json_data)
- return self.success()
- elif quart.request.method == 'DELETE':
- try:
- await self.ap.provider_service.delete_provider(provider_uuid)
- return self.success()
- except ValueError as e:
- return self.http_status(400, -1, str(e))
+ return self.success(data={'providers': providers})
- @self.route('//scan-models', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def _(provider_uuid: str) -> str:
+ @self.route(
+ '',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(request_context: RequestContext) -> str:
+ json_data = await quart.request.json
+ try:
+ provider_uuid = await self.ap.provider_service.create_provider(request_context, json_data)
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ return self.success(data={'uuid': provider_uuid})
+
+ @self.route(
+ '/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(provider_uuid: str, request_context: RequestContext) -> str:
+ provider = await self.ap.provider_service.get_provider(
+ request_context,
+ provider_uuid,
+ include_secret=has_permission(request_context, Permission.PROVIDER_SECRET_MANAGE),
+ )
+ if provider is None:
+ return self.http_status(404, -1, 'provider not found')
+ counts = await self.ap.provider_service.get_provider_model_counts(request_context, provider_uuid)
+ provider['llm_count'] = counts['llm_count']
+ provider['embedding_count'] = counts['embedding_count']
+ provider['rerank_count'] = counts['rerank_count']
+ return self.success(data={'provider': provider})
+
+ @self.route(
+ '/',
+ methods=['PUT'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(provider_uuid: str, request_context: RequestContext) -> str:
+ json_data = await quart.request.json
+ try:
+ await self.ap.provider_service.update_provider(request_context, provider_uuid, json_data)
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ return self.success()
+
+ @self.route(
+ '/',
+ methods=['DELETE'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(provider_uuid: str, request_context: RequestContext) -> str:
+ try:
+ await self.ap.provider_service.delete_provider(request_context, provider_uuid)
+ return self.success()
+ except ValueError as e:
+ return self.http_status(400, -1, str(e))
+
+ @self.route(
+ '//scan-models',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.PROVIDER_SECRET_MANAGE,
+ )
+ async def _(provider_uuid: str, request_context: RequestContext) -> str:
try:
model_type = quart.request.args.get('type')
- result = await self.ap.provider_service.scan_provider_models(provider_uuid, model_type)
+ result = await self.ap.provider_service.scan_provider_models(request_context, provider_uuid, model_type)
return self.success(data=result)
except ValueError as e:
return self.http_status(400, -1, str(e))
diff --git a/src/langbot/pkg/api/http/controller/groups/resources/mcp.py b/src/langbot/pkg/api/http/controller/groups/resources/mcp.py
index 27654e70e..3bd8a3813 100644
--- a/src/langbot/pkg/api/http/controller/groups/resources/mcp.py
+++ b/src/langbot/pkg/api/http/controller/groups/resources/mcp.py
@@ -1,103 +1,130 @@
from __future__ import annotations
import quart
-import traceback
from urllib.parse import unquote
-
+from ....authz import Permission
+from ....context import RequestContext
from ... import group
@group.group_class('mcp', '/api/v1/mcp')
class MCPRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('/servers', methods=['GET', 'POST'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
- """获取MCP服务器列表"""
- if quart.request.method == 'GET':
- servers = await self.ap.mcp_service.get_mcp_servers(contain_runtime_info=True)
-
- return self.success(data={'servers': servers})
-
- elif quart.request.method == 'POST':
- data = await quart.request.json
-
- try:
- uuid = await self.ap.mcp_service.create_mcp_server(data)
- return self.success(data={'uuid': uuid})
- except Exception as e:
- traceback.print_exc()
- return self.http_status(500, -1, f'Failed to create MCP server: {str(e)}')
+ @self.route(
+ '/servers',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ servers = await self.ap.mcp_service.get_mcp_servers(request_context, contain_runtime_info=True)
+ return self.success(data={'servers': servers})
@self.route(
- '/servers/', methods=['GET', 'PUT', 'DELETE'], auth_type=group.AuthType.USER_TOKEN
+ '/servers',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
)
- async def _(server_name: str) -> str:
- """获取、更新或删除MCP服务器配置"""
- server_name = unquote(server_name)
+ async def _(request_context: RequestContext) -> str:
+ data = await quart.request.json
+ try:
+ server_uuid = await self.ap.mcp_service.create_mcp_server(request_context, data)
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ return self.success(data={'uuid': server_uuid})
- server_data = await self.ap.mcp_service.get_mcp_server_by_name(server_name)
+ @self.route(
+ '/servers/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(server_name: str, request_context: RequestContext) -> str:
+ server_name = unquote(server_name)
+ server_data = await self.ap.mcp_service.get_mcp_server_by_name(request_context, server_name)
if server_data is None:
return self.http_status(404, -1, 'Server not found')
+ return self.success(data={'server': server_data})
- if quart.request.method == 'GET':
- return self.success(data={'server': server_data})
-
- elif quart.request.method == 'PUT':
+ @self.route(
+ '/servers/',
+ methods=['PUT', 'DELETE'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(server_name: str, request_context: RequestContext) -> str:
+ server_name = unquote(server_name)
+ server_data = await self.ap.mcp_service.get_mcp_server_by_name(request_context, server_name)
+ if server_data is None:
+ return self.http_status(404, -1, 'Server not found')
+ if quart.request.method == 'PUT':
data = await quart.request.json
try:
- await self.ap.mcp_service.update_mcp_server(server_data['uuid'], data)
- return self.success()
- except Exception as e:
- return self.http_status(500, -1, f'Failed to update MCP server: {str(e)}')
+ await self.ap.mcp_service.update_mcp_server(request_context, server_data['uuid'], data)
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ else:
+ await self.ap.mcp_service.delete_mcp_server(request_context, server_data['uuid'])
+ return self.success()
- elif quart.request.method == 'DELETE':
- try:
- await self.ap.mcp_service.delete_mcp_server(server_data['uuid'])
- return self.success()
- except Exception as e:
- return self.http_status(500, -1, f'Failed to delete MCP server: {str(e)}')
-
- @self.route('/servers//test', methods=['POST'], auth_type=group.AuthType.USER_TOKEN)
- async def _(server_name: str) -> str:
+ @self.route(
+ '/servers//test',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(server_name: str, request_context: RequestContext) -> str:
"""测试MCP服务器连接"""
server_name = unquote(server_name)
server_data = await quart.request.json
- task_id = await self.ap.mcp_service.test_mcp_server(server_name=server_name, server_data=server_data)
+ task_id = await self.ap.mcp_service.test_mcp_server(
+ request_context,
+ server_name=server_name,
+ server_data=server_data,
+ )
return self.success(data={'task_id': task_id})
- @self.route('/servers//resources', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _(server_name: str) -> str:
+ @self.route(
+ '/servers//resources',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(server_name: str, request_context: RequestContext) -> str:
"""Get resources from an MCP server"""
server_name = unquote(server_name)
- try:
- resources = await self.ap.mcp_service.get_mcp_server_resources(server_name)
- templates = await self.ap.mcp_service.get_mcp_server_resource_templates(server_name)
- runtime_info = await self.ap.mcp_service.get_runtime_info(server_name)
- return self.success(
- data={
- 'resources': resources,
- 'resource_templates': templates,
- 'resource_capabilities': (runtime_info or {}).get('resource_capabilities', {}),
- }
- )
- except Exception as e:
- return self.http_status(500, -1, f'Failed to get resources: {str(e)}')
+ resources = await self.ap.mcp_service.get_mcp_server_resources(request_context, server_name)
+ templates = await self.ap.mcp_service.get_mcp_server_resource_templates(request_context, server_name)
+ runtime_info = await self.ap.mcp_service.get_runtime_info(request_context, server_name)
+ return self.success(
+ data={
+ 'resources': resources,
+ 'resource_templates': templates,
+ 'resource_capabilities': (runtime_info or {}).get('resource_capabilities', {}),
+ }
+ )
@self.route(
- '/servers//resource-templates', methods=['GET'], auth_type=group.AuthType.USER_TOKEN
+ '/servers//resource-templates',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def _(server_name: str) -> str:
+ async def _(server_name: str, request_context: RequestContext) -> str:
"""Get resource templates from an MCP server"""
server_name = unquote(server_name)
- try:
- templates = await self.ap.mcp_service.get_mcp_server_resource_templates(server_name)
- return self.success(data={'resource_templates': templates})
- except Exception as e:
- return self.http_status(500, -1, f'Failed to get resource templates: {str(e)}')
+ templates = await self.ap.mcp_service.get_mcp_server_resource_templates(request_context, server_name)
+ return self.success(data={'resource_templates': templates})
- @self.route('/servers//logs', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _(server_name: str) -> str:
+ @self.route(
+ '/servers//logs',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.AUDIT_VIEW,
+ )
+ async def _(server_name: str, request_context: RequestContext) -> str:
"""Get logs from an MCP server"""
server_name = unquote(server_name)
try:
@@ -106,24 +133,32 @@ class MCPRouterGroup(group.RouterGroup):
limit = 200
limit = min(limit, 500)
level = quart.request.args.get('level') or None
- logs = await self.ap.mcp_service.get_mcp_server_logs(server_name, limit=limit, level=level)
+ logs = await self.ap.mcp_service.get_mcp_server_logs(
+ request_context,
+ server_name,
+ limit=limit,
+ level=level,
+ )
return self.success(data={'logs': logs})
- @self.route('/servers//resources/read', methods=['POST'], auth_type=group.AuthType.USER_TOKEN)
- async def _(server_name: str) -> str:
+ @self.route(
+ '/servers//resources/read',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(server_name: str, request_context: RequestContext) -> str:
"""Read a resource from an MCP server"""
server_name = unquote(server_name)
data = await quart.request.json
uri = data.get('uri')
if not uri:
return self.http_status(400, -1, 'URI is required')
- try:
- envelope = await self.ap.mcp_service.read_mcp_server_resource_envelope(
- server_name,
- uri,
- max_bytes=data.get('max_bytes'),
- include_blob=bool(data.get('include_blob', False)),
- )
- return self.success(data=envelope)
- except Exception as e:
- return self.http_status(500, -1, f'Failed to read resource: {str(e)}')
+ envelope = await self.ap.mcp_service.read_mcp_server_resource_envelope(
+ request_context,
+ server_name,
+ uri,
+ max_bytes=data.get('max_bytes'),
+ include_blob=bool(data.get('include_blob', False)),
+ )
+ return self.success(data=envelope)
diff --git a/src/langbot/pkg/api/http/controller/groups/resources/tools.py b/src/langbot/pkg/api/http/controller/groups/resources/tools.py
index 128a0647d..87ec2c1b6 100644
--- a/src/langbot/pkg/api/http/controller/groups/resources/tools.py
+++ b/src/langbot/pkg/api/http/controller/groups/resources/tools.py
@@ -2,21 +2,28 @@ from __future__ import annotations
import quart
+from ....authz import Permission
+from ....context import RequestContext
from ... import group
@group.group_class('tools', '/api/v1/tools')
class ToolsRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
"""获取所有可用工具列表"""
pipeline_uuid = quart.request.args.get('pipeline_uuid') or quart.request.args.get('pipeline_id')
bound_plugins: list[str] | None = None
bound_mcp_servers: list[str] | None = None
if pipeline_uuid:
- pipeline = await self.ap.pipeline_service.get_pipeline(pipeline_uuid)
+ pipeline = await self.ap.pipeline_service.get_pipeline(request_context, pipeline_uuid)
if pipeline is None:
return self.http_status(404, -1, 'pipeline not found')
@@ -35,6 +42,7 @@ class ToolsRouterGroup(group.RouterGroup):
return self.success(
data={
'tools': await self.ap.tool_mgr.get_tool_catalog(
+ request_context,
bound_plugins,
bound_mcp_servers,
include_skill_authoring=True,
@@ -42,10 +50,15 @@ class ToolsRouterGroup(group.RouterGroup):
}
)
- @self.route('/', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _(tool_name: str) -> str:
+ @self.route(
+ '/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(tool_name: str, request_context: RequestContext) -> str:
"""获取特定工具详情"""
- tools = await self.ap.tool_mgr.get_all_tools(include_skill_authoring=True)
+ tools = await self.ap.tool_mgr.get_all_tools(request_context, include_skill_authoring=True)
for tool in tools:
if tool.name == tool_name:
diff --git a/src/langbot/pkg/api/http/controller/groups/skills.py b/src/langbot/pkg/api/http/controller/groups/skills.py
index 946741d76..9bcddb6e2 100644
--- a/src/langbot/pkg/api/http/controller/groups/skills.py
+++ b/src/langbot/pkg/api/http/controller/groups/skills.py
@@ -4,6 +4,8 @@ import quart
from langbot_plugin.box.errors import BoxError
+from ...authz import Permission
+from ...context import RequestContext
from .. import group
@@ -12,58 +14,86 @@ class SkillsRouterGroup(group.RouterGroup):
"""Skills management API endpoints."""
async def initialize(self) -> None:
- @self.route('', methods=['GET', 'POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def list_or_create_skills() -> quart.Response:
- if quart.request.method == 'GET':
- try:
- skills = await self.ap.skill_service.list_skills()
- except (ValueError, BoxError) as exc:
- return self.http_status(400, -1, str(exc))
- return self.success(data={'skills': skills})
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def list_skills(request_context: RequestContext) -> quart.Response:
+ try:
+ skills = await self.ap.skill_service.list_skills(request_context)
+ except (ValueError, BoxError) as exc:
+ return self.http_status(400, -1, str(exc))
+ return self.success(data={'skills': skills})
+ @self.route(
+ '',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def create_skill(request_context: RequestContext) -> quart.Response:
data = await quart.request.json
if 'name' not in data or not data['name']:
return self.http_status(400, -1, 'Missing required field: name')
try:
- skill = await self.ap.skill_service.create_skill(data)
+ skill = await self.ap.skill_service.create_skill(request_context, data)
return self.success(data={'skill': skill})
except (ValueError, BoxError) as exc:
return self.http_status(400, -1, str(exc))
- @self.route('/', methods=['GET', 'PUT', 'DELETE'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def get_update_delete_skill(skill_name: str) -> quart.Response:
- if quart.request.method == 'GET':
- try:
- skill = await self.ap.skill_service.get_skill(skill_name)
- except (ValueError, BoxError) as exc:
- return self.http_status(400, -1, str(exc))
- if not skill:
- return self.http_status(404, -1, 'Skill not found')
- return self.success(data={'skill': skill})
+ @self.route(
+ '/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def get_skill(skill_name: str, request_context: RequestContext) -> quart.Response:
+ try:
+ skill = await self.ap.skill_service.get_skill(request_context, skill_name)
+ except (ValueError, BoxError) as exc:
+ return self.http_status(400, -1, str(exc))
+ if not skill:
+ return self.http_status(404, -1, 'Skill not found')
+ return self.success(data={'skill': skill})
+ @self.route(
+ '/',
+ methods=['PUT', 'DELETE'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def update_delete_skill(skill_name: str, request_context: RequestContext) -> quart.Response:
if quart.request.method == 'PUT':
data = await quart.request.json
try:
- skill = await self.ap.skill_service.update_skill(skill_name, data)
+ skill = await self.ap.skill_service.update_skill(request_context, skill_name, data)
return self.success(data={'skill': skill})
except (ValueError, BoxError) as exc:
return self.http_status(400, -1, str(exc))
try:
- await self.ap.skill_service.delete_skill(skill_name)
+ await self.ap.skill_service.delete_skill(request_context, skill_name)
return self.success()
except (ValueError, BoxError) as exc:
return self.http_status(400, -1, str(exc))
- @self.route('//files', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def list_skill_files(skill_name: str) -> quart.Response:
+ @self.route(
+ '//files',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def list_skill_files(skill_name: str, request_context: RequestContext) -> quart.Response:
"""List files in skill package directory."""
path = quart.request.args.get('path', '.').strip()
include_hidden = quart.request.args.get('include_hidden', 'false').lower() == 'true'
try:
result = await self.ap.skill_service.list_skill_files(
+ request_context,
skill_name,
path=path,
include_hidden=include_hidden,
@@ -73,38 +103,55 @@ class SkillsRouterGroup(group.RouterGroup):
return self.http_status(400, -1, str(exc))
@self.route(
- '//files/', methods=['GET', 'PUT'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY
+ '//files/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
)
- async def read_or_write_skill_file(skill_name: str, path: str) -> quart.Response:
- """Read or write a file in skill package."""
- if quart.request.method == 'GET':
- try:
- result = await self.ap.skill_service.read_skill_file(skill_name, path)
- return self.success(data=result)
- except (ValueError, BoxError) as exc:
- return self.http_status(400, -1, str(exc))
+ async def read_skill_file(skill_name: str, path: str, request_context: RequestContext) -> quart.Response:
+ try:
+ result = await self.ap.skill_service.read_skill_file(request_context, skill_name, path)
+ return self.success(data=result)
+ except (ValueError, BoxError) as exc:
+ return self.http_status(400, -1, str(exc))
- # PUT - write file
+ @self.route(
+ '//files/',
+ methods=['PUT'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def write_skill_file(skill_name: str, path: str, request_context: RequestContext) -> quart.Response:
data = await quart.request.json
content = data.get('content', '')
if content is None:
return self.http_status(400, -1, 'Missing required field: content')
try:
- result = await self.ap.skill_service.write_skill_file(skill_name, path, content)
+ result = await self.ap.skill_service.write_skill_file(request_context, skill_name, path, content)
return self.success(data=result)
except (ValueError, BoxError) as exc:
return self.http_status(400, -1, str(exc))
- @self.route('//preview', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def preview_skill(skill_name: str) -> quart.Response:
- skill = self.ap.skill_mgr.get_skill_by_name(skill_name)
+ @self.route(
+ '//preview',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def preview_skill(skill_name: str, request_context: RequestContext) -> quart.Response:
+ skill = await self.ap.skill_service.get_skill(request_context, skill_name)
if not skill:
return self.http_status(404, -1, 'Skill not found')
return self.success(data={'instructions': skill.get('instructions', '')})
- @self.route('/install/github', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def install_skill_from_github() -> quart.Response:
+ @self.route(
+ '/install/github',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def install_skill_from_github(request_context: RequestContext) -> quart.Response:
data = await quart.request.json
required_fields = ['asset_url', 'owner', 'repo']
for field in required_fields:
@@ -115,15 +162,20 @@ class SkillsRouterGroup(group.RouterGroup):
return self.http_status(400, -1, 'Missing required field: release_tag')
try:
- skill = await self.ap.skill_service.install_from_github(data)
+ skill = await self.ap.skill_service.install_from_github(request_context, data)
return self.success(data={'skills': skill})
except (ValueError, BoxError) as exc:
return self.http_status(400, -1, str(exc))
- except Exception as exc:
- return self.http_status(500, -1, f'Failed to install skill: {exc}')
+ except Exception:
+ raise
- @self.route('/install/github/preview', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def preview_skill_from_github() -> quart.Response:
+ @self.route(
+ '/install/github/preview',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def preview_skill_from_github(request_context: RequestContext) -> quart.Response:
data = await quart.request.json
required_fields = ['asset_url', 'owner', 'repo']
for field in required_fields:
@@ -134,15 +186,20 @@ class SkillsRouterGroup(group.RouterGroup):
return self.http_status(400, -1, 'Missing required field: release_tag')
try:
- preview = await self.ap.skill_service.preview_install_from_github(data)
+ preview = await self.ap.skill_service.preview_install_from_github(request_context, data)
return self.success(data={'skills': preview})
except (ValueError, BoxError) as exc:
return self.http_status(400, -1, str(exc))
- except Exception as exc:
- return self.http_status(500, -1, f'Failed to preview skill: {exc}')
+ except Exception:
+ raise
- @self.route('/install/upload', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def install_skill_from_upload() -> quart.Response:
+ @self.route(
+ '/install/upload',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def install_skill_from_upload(request_context: RequestContext) -> quart.Response:
file = (await quart.request.files).get('file')
if file is None:
return self.http_status(400, -1, 'file is required')
@@ -150,6 +207,7 @@ class SkillsRouterGroup(group.RouterGroup):
try:
skill = await self.ap.skill_service.install_from_zip_upload(
+ request_context,
file_bytes=file.read(),
filename=file.filename or '',
source_paths=form.getlist('source_paths'),
@@ -157,34 +215,45 @@ class SkillsRouterGroup(group.RouterGroup):
return self.success(data={'skills': skill})
except (ValueError, BoxError) as exc:
return self.http_status(400, -1, str(exc))
- except Exception as exc:
- return self.http_status(500, -1, f'Failed to install skill: {exc}')
+ except Exception:
+ raise
- @self.route('/install/upload/preview', methods=['POST'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def preview_skill_from_upload() -> quart.Response:
+ @self.route(
+ '/install/upload/preview',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def preview_skill_from_upload(request_context: RequestContext) -> quart.Response:
file = (await quart.request.files).get('file')
if file is None:
return self.http_status(400, -1, 'file is required')
try:
preview = await self.ap.skill_service.preview_install_from_zip_upload(
+ request_context,
file_bytes=file.read(),
filename=file.filename or '',
)
return self.success(data={'skills': preview})
except (ValueError, BoxError) as exc:
return self.http_status(400, -1, str(exc))
- except Exception as exc:
- return self.http_status(500, -1, f'Failed to preview skill: {exc}')
+ except Exception:
+ raise
- @self.route('/scan', methods=['GET'], auth_type=group.AuthType.USER_TOKEN_OR_API_KEY)
- async def scan_skill_directory() -> quart.Response:
+ @self.route(
+ '/scan',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN_OR_API_KEY,
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def scan_skill_directory(request_context: RequestContext) -> quart.Response:
path = quart.request.args.get('path', '').strip()
if not path:
return self.http_status(400, -1, 'Missing required parameter: path')
try:
- result = await self.ap.skill_service.scan_directory_async(path)
+ result = await self.ap.skill_service.scan_directory_async(request_context, path)
return self.success(data=result)
except (ValueError, BoxError) as exc:
return self.http_status(400, -1, str(exc))
diff --git a/src/langbot/pkg/api/http/controller/groups/stats.py b/src/langbot/pkg/api/http/controller/groups/stats.py
index 8c8e9113c..e4c8dbdae 100644
--- a/src/langbot/pkg/api/http/controller/groups/stats.py
+++ b/src/langbot/pkg/api/http/controller/groups/stats.py
@@ -1,19 +1,39 @@
from .. import group
+from ...authz import Permission
+from ...context import ExecutionContext, RequestContext
+
+
+def collect_basic_stats(ap, request_context: RequestContext) -> dict[str, int]:
+ """Collect runtime counters only from the selected Workspace placement."""
+
+ execution_context = ExecutionContext.from_request(request_context)
+ sessions = [
+ session
+ for session in ap.sess_mgr.session_list
+ if (
+ getattr(session, 'instance_uuid', None) == execution_context.instance_uuid
+ and getattr(session, 'workspace_uuid', None) == execution_context.workspace_uuid
+ and getattr(session, 'placement_generation', None) == execution_context.placement_generation
+ )
+ ]
+ conversation_count = sum(
+ len(session.conversations if session.conversations is not None else []) for session in sessions
+ )
+ return {
+ 'active_session_count': len(sessions),
+ 'conversation_count': conversation_count,
+ 'query_count': ap.query_pool.get_query_count(execution_context),
+ }
@group.group_class('stats', '/api/v1/stats')
class StatsRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('/basic', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
- conv_count = 0
- for session in self.ap.sess_mgr.session_list:
- conv_count += len(session.conversations if session.conversations is not None else [])
-
- return self.success(
- data={
- 'active_session_count': len(self.ap.sess_mgr.session_list),
- 'conversation_count': conv_count,
- 'query_count': self.ap.query_pool.query_id_counter,
- }
- )
+ @self.route(
+ '/basic',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ return self.success(data=collect_basic_stats(self.ap, request_context))
diff --git a/src/langbot/pkg/api/http/controller/groups/system.py b/src/langbot/pkg/api/http/controller/groups/system.py
index 236a23582..002c037af 100644
--- a/src/langbot/pkg/api/http/controller/groups/system.py
+++ b/src/langbot/pkg/api/http/controller/groups/system.py
@@ -5,7 +5,9 @@ import sqlalchemy
from .. import group
from .....utils import constants
-from .....entity.persistence.metadata import Metadata
+from .....entity.persistence.metadata import WorkspaceMetadata
+from ...authz import Permission
+from ...context import RequestContext
@group.group_class('system', '/api/v1/system')
@@ -17,17 +19,25 @@ class SystemRouterGroup(group.RouterGroup):
wizard_status = 'none'
wizard_progress = None
try:
- result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(Metadata).where(Metadata.key.in_(['wizard_status', 'wizard_progress']))
- )
- for row in result:
- if row.key == 'wizard_status':
- wizard_status = row.value
- elif row.key == 'wizard_progress':
- try:
- wizard_progress = json.loads(row.value)
- except (json.JSONDecodeError, TypeError):
- wizard_progress = None
+ authorization = quart.request.headers.get('Authorization', '')
+ if authorization.startswith('Bearer '):
+ account, _ = await self._authenticate_account(authorization.removeprefix('Bearer '))
+ request_context = await self._resolve_account_context(account, group.AuthType.USER_TOKEN)
+ if request_context is not None:
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(WorkspaceMetadata).where(
+ WorkspaceMetadata.workspace_uuid == request_context.workspace_uuid,
+ WorkspaceMetadata.key.in_(['wizard_status', 'wizard_progress']),
+ )
+ )
+ for row in result:
+ if row.key == 'wizard_status':
+ wizard_status = row.value
+ elif row.key == 'wizard_progress':
+ try:
+ wizard_progress = json.loads(row.value)
+ except (json.JSONDecodeError, TypeError):
+ wizard_progress = None
except Exception:
pass
@@ -67,8 +77,13 @@ class SystemRouterGroup(group.RouterGroup):
}
)
- @self.route('/wizard/completed', methods=['POST'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
+ @self.route(
+ '/wizard/completed',
+ methods=['POST'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.WORKSPACE_UPDATE,
+ )
+ async def _(request_context: RequestContext) -> str:
"""Mark wizard status in metadata table and clear progress.
Accepts JSON body: { "status": "skipped" | "completed" }
@@ -80,28 +95,48 @@ class SystemRouterGroup(group.RouterGroup):
try:
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(Metadata).where(Metadata.key == 'wizard_status')
+ sqlalchemy.select(WorkspaceMetadata).where(
+ WorkspaceMetadata.workspace_uuid == request_context.workspace_uuid,
+ WorkspaceMetadata.key == 'wizard_status',
+ )
)
if result.first():
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(Metadata).where(Metadata.key == 'wizard_status').values(value=status)
+ sqlalchemy.update(WorkspaceMetadata)
+ .where(
+ WorkspaceMetadata.workspace_uuid == request_context.workspace_uuid,
+ WorkspaceMetadata.key == 'wizard_status',
+ )
+ .values(value=status)
)
else:
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.insert(Metadata).values(key='wizard_status', value=status)
+ sqlalchemy.insert(WorkspaceMetadata).values(
+ workspace_uuid=request_context.workspace_uuid,
+ key='wizard_status',
+ value=status,
+ )
)
# Clear wizard progress when wizard is completed/skipped
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(Metadata).where(Metadata.key == 'wizard_progress')
+ sqlalchemy.delete(WorkspaceMetadata).where(
+ WorkspaceMetadata.workspace_uuid == request_context.workspace_uuid,
+ WorkspaceMetadata.key == 'wizard_progress',
+ )
)
- except Exception as e:
- return self.http_status(500, 500, f'Failed to update wizard status: {e}')
+ except Exception:
+ raise
return self.success(data={})
- @self.route('/wizard/progress', methods=['PUT'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
+ @self.route(
+ '/wizard/progress',
+ methods=['PUT'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.WORKSPACE_UPDATE,
+ )
+ async def _(request_context: RequestContext) -> str:
"""Save wizard progress to metadata table.
Accepts JSON body with wizard state fields:
@@ -113,23 +148,40 @@ class SystemRouterGroup(group.RouterGroup):
try:
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(Metadata).where(Metadata.key == 'wizard_progress')
+ sqlalchemy.select(WorkspaceMetadata).where(
+ WorkspaceMetadata.workspace_uuid == request_context.workspace_uuid,
+ WorkspaceMetadata.key == 'wizard_progress',
+ )
)
if result.first():
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(Metadata).where(Metadata.key == 'wizard_progress').values(value=progress_json)
+ sqlalchemy.update(WorkspaceMetadata)
+ .where(
+ WorkspaceMetadata.workspace_uuid == request_context.workspace_uuid,
+ WorkspaceMetadata.key == 'wizard_progress',
+ )
+ .values(value=progress_json)
)
else:
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.insert(Metadata).values(key='wizard_progress', value=progress_json)
+ sqlalchemy.insert(WorkspaceMetadata).values(
+ workspace_uuid=request_context.workspace_uuid,
+ key='wizard_progress',
+ value=progress_json,
+ )
)
- except Exception as e:
- return self.http_status(500, 500, f'Failed to save wizard progress: {e}')
+ except Exception:
+ raise
return self.success(data={})
- @self.route('/tasks', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
+ @self.route(
+ '/tasks',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
task_type = quart.request.args.get('type')
task_kind = quart.request.args.get('kind')
@@ -138,30 +190,56 @@ class SystemRouterGroup(group.RouterGroup):
if task_kind == '':
task_kind = None
- return self.success(data=self.ap.task_mgr.get_tasks_dict(task_type, task_kind))
+ return self.success(
+ data=self.ap.task_mgr.get_tasks_dict(
+ task_type,
+ task_kind,
+ instance_uuid=request_context.instance_uuid,
+ workspace_uuid=request_context.workspace_uuid,
+ placement_generation=request_context.placement_generation,
+ )
+ )
- @self.route('/tasks/', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _(task_id: str) -> str:
- task = self.ap.task_mgr.get_task_by_id(int(task_id))
+ @self.route(
+ '/tasks/',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.RESOURCE_VIEW,
+ )
+ async def _(task_id: str, request_context: RequestContext) -> str:
+ task = self.ap.task_mgr.get_task_by_id(
+ int(task_id),
+ instance_uuid=request_context.instance_uuid,
+ workspace_uuid=request_context.workspace_uuid,
+ placement_generation=request_context.placement_generation,
+ )
if task is None:
return self.http_status(404, 404, 'Task not found')
return self.success(data=task.to_dict())
- @self.route('/storage-analysis', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
- async def _() -> str:
- return self.success(data=await self.ap.maintenance_service.get_storage_analysis())
+ @self.route(
+ '/storage-analysis',
+ methods=['GET'],
+ auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.AUDIT_VIEW,
+ )
+ async def _(request_context: RequestContext) -> str:
+ return self.success(data=await self.ap.maintenance_service.get_storage_analysis(request_context))
@self.route(
'/debug/plugin/action',
methods=['POST'],
auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.RUNTIME_OPERATE,
)
- async def _() -> str:
+ async def _(request_context: RequestContext) -> str:
if not constants.debug_mode:
return self.http_status(403, 403, 'Forbidden')
+ await self.ap.plugin_connector.require_workspace_context(request_context)
+
data = await quart.request.json
class AnoymousAction:
@@ -174,6 +252,7 @@ class SystemRouterGroup(group.RouterGroup):
AnoymousAction(data['action']),
data['data'],
timeout=data.get('timeout', 10),
+ action_context=self.ap.plugin_connector.handler.require_bound_action_context().without_installation(),
)
return self.success(data=resp)
@@ -182,8 +261,10 @@ class SystemRouterGroup(group.RouterGroup):
'/status/plugin-system',
methods=['GET'],
auth_type=group.AuthType.USER_TOKEN,
+ permission=Permission.RESOURCE_VIEW,
)
- async def _() -> str:
+ async def _(request_context: RequestContext) -> str:
+ await self.ap.plugin_connector.require_workspace_context(request_context)
plugin_connector_error = 'ok'
is_connected = True
diff --git a/src/langbot/pkg/api/http/controller/groups/user.py b/src/langbot/pkg/api/http/controller/groups/user.py
index 886dc5d0d..ce5c8c591 100644
--- a/src/langbot/pkg/api/http/controller/groups/user.py
+++ b/src/langbot/pkg/api/http/controller/groups/user.py
@@ -1,14 +1,53 @@
import quart
import argon2
import asyncio
-import traceback
+from urllib.parse import parse_qs, urlsplit
from .. import group
from .....entity.errors import account as account_errors
+from ...context import RequestContext
+from ...service.user import ControlPlaneDirectoryRequiredError, PublicRegistrationClosedError
@group.group_class('user', '/api/v1/user')
class UserRouterGroup(group.RouterGroup):
+ @staticmethod
+ def _origin(value: str) -> tuple[str, str, int | None] | None:
+ parsed = urlsplit(value)
+ if parsed.scheme not in {'http', 'https'} or not parsed.hostname:
+ return None
+ return parsed.scheme, parsed.hostname.casefold(), parsed.port
+
+ def _validate_space_redirect_uri(self, redirect_uri: str, *, bind: bool) -> str:
+ parsed = urlsplit(redirect_uri)
+ if (
+ parsed.scheme not in {'http', 'https'}
+ or not parsed.hostname
+ or parsed.username is not None
+ or parsed.password is not None
+ or parsed.fragment
+ or parsed.path != '/auth/space/callback'
+ ):
+ raise ValueError('Invalid redirect_uri parameter')
+
+ query = parse_qs(parsed.query, keep_blank_values=True)
+ if bind:
+ if query != {'mode': ['bind']}:
+ raise ValueError('Invalid Space binding redirect_uri')
+ elif query:
+ raise ValueError('Invalid Space login redirect_uri')
+
+ redirect_origin = self._origin(redirect_uri)
+ api_config = self.ap.instance_config.data.get('api', {})
+ trusted_origins = {
+ self._origin(str(api_config.get(config_key, '') or '').strip())
+ for config_key in ('webui_url', 'webhook_prefix')
+ }
+ trusted_origins.discard(None)
+ if redirect_origin not in trusted_origins:
+ raise ValueError('Untrusted redirect_uri origin')
+ return redirect_uri
+
async def initialize(self) -> None:
@self.route('/init', methods=['GET', 'POST'], auth_type=group.AuthType.NONE)
async def _() -> str:
@@ -23,7 +62,12 @@ class UserRouterGroup(group.RouterGroup):
user_email = json_data['user']
password = json_data['password']
- await self.ap.user_service.create_user(user_email, password)
+ try:
+ await self.ap.user_service.create_user(user_email, password)
+ except ControlPlaneDirectoryRequiredError as exc:
+ return self.http_status(409, exc.code, str(exc))
+ except PublicRegistrationClosedError:
+ return self.http_status(409, 'registration_closed', 'System already initialized')
return self.success()
@@ -40,7 +84,7 @@ class UserRouterGroup(group.RouterGroup):
return self.success(data={'token': token})
- @self.route('/check-token', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
+ @self.route('/check-token', methods=['GET'], auth_type=group.AuthType.ACCOUNT_TOKEN)
async def _(user_email: str) -> str:
token = await self.ap.user_service.generate_jwt_token(user_email)
@@ -101,15 +145,37 @@ class UserRouterGroup(group.RouterGroup):
async def _() -> str:
"""Get Space OAuth authorization URL for redirect"""
redirect_uri = quart.request.args.get('redirect_uri', '')
- state = quart.request.args.get('state', '')
if not redirect_uri:
return self.fail(1, 'Missing redirect_uri parameter')
+ if 'state' in quart.request.args:
+ return self.fail(1, 'Caller-supplied OAuth state is not allowed')
try:
+ redirect_uri = self._validate_space_redirect_uri(redirect_uri, bind=False)
+ state = await self.ap.user_service.issue_space_oauth_state('login')
authorize_url = self.ap.space_service.get_oauth_authorize_url(redirect_uri, state)
return self.success(data={'authorize_url': authorize_url})
- except Exception as e:
+ except ValueError as e:
+ return self.fail(1, str(e))
+
+ @self.route('/space/bind-authorize-url', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
+ async def _(request_context: RequestContext) -> str:
+ """Issue an account-bound, one-time Space OAuth redirect."""
+ redirect_uri = quart.request.args.get('redirect_uri', '')
+ if not redirect_uri:
+ return self.fail(1, 'Missing redirect_uri parameter')
+ if not request_context.account_uuid:
+ return self.http_status(403, 'account_required', 'An Account is required')
+ try:
+ redirect_uri = self._validate_space_redirect_uri(redirect_uri, bind=True)
+ state = await self.ap.user_service.issue_space_oauth_state(
+ 'bind',
+ account_uuid=request_context.account_uuid,
+ )
+ authorize_url = self.ap.space_service.get_oauth_authorize_url(redirect_uri, state)
+ return self.success(data={'authorize_url': authorize_url})
+ except ValueError as e:
return self.fail(1, str(e))
@self.route('/space/callback', methods=['POST'], auth_type=group.AuthType.NONE)
@@ -117,11 +183,15 @@ class UserRouterGroup(group.RouterGroup):
"""Handle OAuth callback - exchange code for tokens and authenticate"""
json_data = await quart.request.json
code = json_data.get('code')
+ state = json_data.get('state')
if not code:
return self.fail(1, 'Missing authorization code')
+ if not state:
+ return self.fail(1, 'Missing state parameter')
try:
+ await self.ap.user_service.consume_space_oauth_state(state, 'login')
# Exchange code for tokens
token_data = await self.ap.space_service.exchange_oauth_code(code)
access_token = token_data.get('access_token')
@@ -142,15 +212,15 @@ class UserRouterGroup(group.RouterGroup):
'user': user_obj.user,
}
)
+ except ControlPlaneDirectoryRequiredError as e:
+ return self.http_status(409, e.code, str(e))
except account_errors.AccountEmailMismatchError as e:
return self.fail(3, str(e))
- except ValueError as e:
- traceback.print_exc()
- self.ap.logger.warning(f'Space OAuth callback failed: {e}')
- return self.fail(1, str(e))
- except Exception as e:
- traceback.print_exc()
- return self.fail(2, f'OAuth callback failed: {str(e)}')
+ except ValueError:
+ self.ap.logger.exception('Space OAuth callback failed')
+ return self.fail(1, 'Space OAuth failed')
+ except Exception:
+ raise
@self.route('/info', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
async def _(user_email: str) -> str:
@@ -162,6 +232,7 @@ class UserRouterGroup(group.RouterGroup):
return self.success(
data={
+ 'account_uuid': user_obj.uuid,
'user': user_obj.user,
'account_type': user_obj.account_type,
'has_password': bool(user_obj.password and user_obj.password.strip()),
@@ -176,19 +247,18 @@ class UserRouterGroup(group.RouterGroup):
@self.route('/account-info', methods=['GET'], auth_type=group.AuthType.NONE)
async def _() -> str:
- """Get account info for login page (account type and has_password)"""
+ """Return instance login capabilities without disclosing an account."""
if not await self.ap.user_service.is_initialized():
return self.success(data={'initialized': False})
- user_obj = await self.ap.user_service.get_first_user()
- if user_obj is None:
- return self.success(data={'initialized': False})
-
return self.success(
data={
'initialized': True,
- 'account_type': user_obj.account_type,
- 'has_password': bool(user_obj.password and user_obj.password.strip()),
+ # Login is selected per account in a multi-user instance. A public
+ # bootstrap endpoint must never project one user's authentication
+ # methods onto every other user or disclose that user's state.
+ 'password_login_enabled': True,
+ 'space_login_enabled': True,
}
)
@@ -233,7 +303,7 @@ class UserRouterGroup(group.RouterGroup):
json_data = await quart.request.json
code = json_data.get('code')
- state = json_data.get('state') # JWT token passed as state
+ state = json_data.get('state')
if not code:
return self.http_status(400, -1, 'Missing authorization code')
@@ -241,13 +311,10 @@ class UserRouterGroup(group.RouterGroup):
if not state:
return self.http_status(400, -1, 'Missing state parameter')
- # Verify state is a valid JWT token
try:
- user_email = await self.ap.user_service.verify_jwt_token(state)
+ user_obj = await self.ap.user_service.consume_space_oauth_state(state, 'bind')
except Exception:
return self.http_status(401, -1, 'Invalid or expired state')
-
- user_obj = await self.ap.user_service.get_user_by_email(user_email)
if user_obj is None:
return self.http_status(404, -1, 'User not found')
@@ -255,8 +322,8 @@ class UserRouterGroup(group.RouterGroup):
return self.http_status(400, -1, 'Only local accounts can bind to Space')
try:
- updated_user = await self.ap.user_service.bind_space_account(user_email, code)
- jwt_token = await self.ap.user_service.generate_jwt_token(updated_user.user)
+ updated_user = await self.ap.user_service.bind_space_account(user_obj.user, code)
+ jwt_token = await self.ap.user_service.generate_jwt_token(updated_user)
return self.success(
data={
'token': jwt_token,
@@ -264,7 +331,8 @@ class UserRouterGroup(group.RouterGroup):
'account_type': updated_user.account_type,
}
)
- except ValueError as e:
- return self.http_status(400, -1, str(e))
- except Exception as e:
- return self.http_status(500, -1, f'Failed to bind Space account: {str(e)}')
+ except ValueError:
+ self.ap.logger.exception('Space account binding failed')
+ return self.http_status(400, -1, 'Space account binding failed')
+ except Exception:
+ raise
diff --git a/src/langbot/pkg/api/http/controller/groups/webhook_mgmt.py b/src/langbot/pkg/api/http/controller/groups/webhook_mgmt.py
index f82184c18..727a3540d 100644
--- a/src/langbot/pkg/api/http/controller/groups/webhook_mgmt.py
+++ b/src/langbot/pkg/api/http/controller/groups/webhook_mgmt.py
@@ -1,49 +1,80 @@
+from __future__ import annotations
+
import quart
+from ...authz import Permission, has_permission
+from ...context import RequestContext
from .. import group
@group.group_class('webhook_mgmt', '/api/v1/webhooks')
class WebhookManagementRouterGroup(group.RouterGroup):
async def initialize(self) -> None:
- @self.route('', methods=['GET', 'POST'])
- async def _() -> str:
- if quart.request.method == 'GET':
- webhooks = await self.ap.webhook_service.get_webhooks()
- return self.success(data={'webhooks': webhooks})
- elif quart.request.method == 'POST':
- json_data = await quart.request.json
- name = json_data.get('name', '')
- url = json_data.get('url', '')
- description = json_data.get('description', '')
- enabled = json_data.get('enabled', True)
+ @self.route('', methods=['GET'], permission=Permission.RESOURCE_VIEW)
+ async def _(request_context: RequestContext) -> str:
+ webhooks = await self.ap.webhook_service.get_webhooks(
+ request_context,
+ include_secret=has_permission(request_context, Permission.RESOURCE_MANAGE),
+ )
+ return self.success(data={'webhooks': webhooks})
- if not name:
- return self.http_status(400, -1, 'Name is required')
- if not url:
- return self.http_status(400, -1, 'URL is required')
+ @self.route('', methods=['POST'], permission=Permission.RESOURCE_MANAGE)
+ async def _(request_context: RequestContext) -> str:
+ json_data = await quart.request.get_json(silent=True) or {}
+ name = json_data.get('name', '')
+ url = json_data.get('url', '')
+ description = json_data.get('description', '')
+ enabled = json_data.get('enabled', True)
- webhook = await self.ap.webhook_service.create_webhook(name, url, description, enabled)
- return self.success(data={'webhook': webhook})
+ if not name:
+ return self.http_status(400, -1, 'Name is required')
+ if not url:
+ return self.http_status(400, -1, 'URL is required')
- @self.route('/', methods=['GET', 'PUT', 'DELETE'])
- async def _(webhook_id: int) -> str:
- if quart.request.method == 'GET':
- webhook = await self.ap.webhook_service.get_webhook(webhook_id)
- if webhook is None:
+ try:
+ webhook = await self.ap.webhook_service.create_webhook(
+ request_context,
+ name,
+ url,
+ description,
+ enabled,
+ )
+ except ValueError as exc:
+ return self.http_status(400, -1, str(exc))
+ return self.success(data={'webhook': webhook})
+
+ @self.route('/', methods=['GET'], permission=Permission.RESOURCE_VIEW)
+ async def _(webhook_id: int, request_context: RequestContext) -> str:
+ webhook = await self.ap.webhook_service.get_webhook(
+ request_context,
+ webhook_id,
+ include_secret=has_permission(request_context, Permission.RESOURCE_MANAGE),
+ )
+ if webhook is None:
+ return self.http_status(404, -1, 'Webhook not found')
+ return self.success(data={'webhook': webhook})
+
+ @self.route(
+ '/',
+ methods=['PUT', 'DELETE'],
+ permission=Permission.RESOURCE_MANAGE,
+ )
+ async def _(webhook_id: int, request_context: RequestContext) -> str:
+ if quart.request.method == 'PUT':
+ json_data = await quart.request.get_json(silent=True) or {}
+ updated = await self.ap.webhook_service.update_webhook(
+ request_context,
+ webhook_id,
+ json_data.get('name'),
+ json_data.get('url'),
+ json_data.get('description'),
+ json_data.get('enabled'),
+ )
+ if not updated:
return self.http_status(404, -1, 'Webhook not found')
- return self.success(data={'webhook': webhook})
-
- elif quart.request.method == 'PUT':
- json_data = await quart.request.json
- name = json_data.get('name')
- url = json_data.get('url')
- description = json_data.get('description')
- enabled = json_data.get('enabled')
-
- await self.ap.webhook_service.update_webhook(webhook_id, name, url, description, enabled)
return self.success()
- elif quart.request.method == 'DELETE':
- await self.ap.webhook_service.delete_webhook(webhook_id)
- return self.success()
+ deleted = await self.ap.webhook_service.delete_webhook(request_context, webhook_id)
+ if not deleted:
+ return self.http_status(404, -1, 'Webhook not found')
+ return self.success()
diff --git a/src/langbot/pkg/api/http/controller/groups/webhooks.py b/src/langbot/pkg/api/http/controller/groups/webhooks.py
index ec46c7447..7bd315ca4 100644
--- a/src/langbot/pkg/api/http/controller/groups/webhooks.py
+++ b/src/langbot/pkg/api/http/controller/groups/webhooks.py
@@ -30,7 +30,10 @@ class WebhookRouterGroup(group.RouterGroup):
适配器返回的响应
"""
try:
- runtime_bot = await self.ap.platform_mgr.get_bot_by_uuid(bot_uuid)
+ # Public ingress never accepts X-Workspace-Id. The opaque bot UUID
+ # is resolved against the already-bound runtime resource, which
+ # carries the trusted Workspace and placement generation.
+ runtime_bot = await self.ap.platform_mgr.resolve_public_bot(bot_uuid)
if not runtime_bot:
return quart.jsonify({'error': 'Bot not found'}), 404
@@ -49,6 +52,9 @@ class WebhookRouterGroup(group.RouterGroup):
return response
- except Exception as e:
- self.ap.logger.error(f'Webhook dispatch error for bot {bot_uuid}: {traceback.format_exc()}')
- return quart.jsonify({'error': str(e)}), 500
+ except Exception:
+ request_id = self.request_id()
+ self.ap.logger.error(
+ f'Webhook dispatch error request_id={request_id} bot={bot_uuid}: {traceback.format_exc()}'
+ )
+ return self.internal_error_response(request_id)
diff --git a/src/langbot/pkg/api/http/controller/groups/workspaces.py b/src/langbot/pkg/api/http/controller/groups/workspaces.py
new file mode 100644
index 000000000..273cadf05
--- /dev/null
+++ b/src/langbot/pkg/api/http/controller/groups/workspaces.py
@@ -0,0 +1,305 @@
+from __future__ import annotations
+
+import typing
+
+import quart
+
+from ...authz import Permission, permissions_for_role
+from ...context import RequestContext
+from ...service.user import AccountExistsLoginRequiredError, ControlPlaneDirectoryRequiredError
+from .....entity.persistence.workspace import Workspace, WorkspaceInvitation, WorkspaceMembership
+from .....entity.persistence.workspace import WorkspaceSource
+from .....workspace.collaboration import WorkspaceMemberView
+from .....workspace.errors import WorkspaceNotFoundError
+from .. import group
+
+
+def _workspace_payload(workspace: Workspace) -> dict[str, typing.Any]:
+ return {
+ 'uuid': workspace.uuid,
+ 'instance_uuid': workspace.instance_uuid,
+ 'name': workspace.name,
+ 'slug': workspace.slug,
+ 'type': workspace.type,
+ 'status': workspace.status,
+ 'source': workspace.source,
+ }
+
+
+def _membership_payload(
+ membership: WorkspaceMembership,
+ *,
+ email: str,
+) -> dict[str, typing.Any]:
+ return {
+ 'uuid': membership.uuid,
+ 'workspace_uuid': membership.workspace_uuid,
+ 'account_uuid': membership.account_uuid,
+ 'email': email,
+ 'role': membership.role,
+ 'status': membership.status,
+ 'joined_at': membership.joined_at.isoformat() if membership.joined_at else None,
+ 'created_at': membership.created_at.isoformat() if membership.created_at else None,
+ }
+
+
+def _invitation_payload(invitation: WorkspaceInvitation) -> dict[str, typing.Any]:
+ """Serialize an invitation without its bearer-secret hash."""
+
+ return {
+ 'uuid': invitation.uuid,
+ 'workspace_uuid': invitation.workspace_uuid,
+ 'normalized_email': invitation.normalized_email,
+ 'role': invitation.role,
+ 'status': invitation.status,
+ 'expires_at': invitation.expires_at.isoformat(),
+ 'created_at': invitation.created_at.isoformat() if invitation.created_at else None,
+ }
+
+
+@group.group_class('workspaces', '/api/v1/workspaces')
+class WorkspacesRouterGroup(group.RouterGroup):
+ async def initialize(self) -> None:
+ @self.route('/bootstrap', methods=['GET'], auth_type=group.AuthType.ACCOUNT_TOKEN)
+ async def _(user_email: str) -> typing.Any:
+ """List the active Workspaces available to an authenticated Account.
+
+ This account-only endpoint intentionally runs before Workspace
+ selection. It never accepts a selector as authority and does not
+ choose a default Workspace for a multi-membership Account.
+ """
+
+ account = await self.ap.user_service.get_user_by_email(user_email)
+ if account is None:
+ return self.http_status(401, 'invalid_authentication', 'Account not found')
+ accesses = await self.ap.workspace_collaboration_service.list_account_workspaces(account.uuid)
+ return self.success(
+ data={
+ 'workspaces': [
+ {
+ 'workspace': _workspace_payload(access.workspace),
+ 'membership': _membership_payload(access.membership, email=account.user),
+ 'permissions': sorted(permissions_for_role(access.membership.role)),
+ 'placement_generation': access.execution.placement_generation,
+ }
+ for access in accesses
+ ]
+ }
+ )
+
+ @self.route('', methods=['GET', 'POST'], permission=Permission.WORKSPACE_VIEW)
+ async def _(request_context: RequestContext) -> typing.Any:
+ if quart.request.method == 'POST':
+ if self.ap.workspace_service.policy.multi_workspace_enabled:
+ return self.http_status(
+ 409,
+ 'control_plane_required',
+ 'Cloud Workspaces are created by the SaaS control plane',
+ )
+ return self.http_status(403, 'edition_limit', 'This edition supports one Workspace per instance')
+
+ accesses = await self.ap.workspace_collaboration_service.list_account_workspaces(
+ request_context.account_uuid
+ )
+ return self.success(data={'workspaces': [_workspace_payload(access.workspace) for access in accesses]})
+
+ @self.route('/current', methods=['GET'], permission=Permission.WORKSPACE_VIEW)
+ async def _(request_context: RequestContext) -> typing.Any:
+ membership = quart.g.workspace_membership
+ account = await self.ap.user_service.get_user_by_uuid(request_context.account_uuid)
+ if account is None:
+ return self.http_status(401, 'invalid_authentication', 'Account not found')
+ workspace = await self.ap.workspace_service.get_workspace(request_context.workspace_uuid)
+ return self.success(
+ data={
+ 'workspace': _workspace_payload(workspace),
+ 'membership': _membership_payload(membership, email=account.user),
+ 'permissions': sorted(request_context.workspace.permissions),
+ 'placement_generation': request_context.placement_generation,
+ }
+ )
+
+ @self.route('/', methods=['GET'], permission=Permission.WORKSPACE_VIEW)
+ async def _(workspace_uuid: str, request_context: RequestContext) -> typing.Any:
+ self._require_current_workspace(workspace_uuid, request_context)
+ workspace = await self.ap.workspace_service.get_workspace(workspace_uuid)
+ return self.success(data={'workspace': _workspace_payload(workspace)})
+
+ @self.route('//members', methods=['GET'], permission=Permission.MEMBER_VIEW)
+ async def _(workspace_uuid: str, request_context: RequestContext) -> typing.Any:
+ self._require_current_workspace(workspace_uuid, request_context)
+ members = await self.ap.workspace_collaboration_service.list_members(
+ workspace_uuid,
+ quart.g.workspace_membership,
+ )
+ return self.success(data={'members': [self._member_view_payload(item) for item in members]})
+
+ @self.route(
+ '//invitations',
+ methods=['GET', 'POST'],
+ permission=Permission.MEMBER_INVITE,
+ )
+ async def _(workspace_uuid: str, request_context: RequestContext) -> typing.Any:
+ self._require_current_workspace(workspace_uuid, request_context)
+ if await self._requires_control_plane(workspace_uuid):
+ return self._control_plane_required()
+ if quart.request.method == 'GET':
+ invitations = await self.ap.workspace_collaboration_service.list_invitations(
+ workspace_uuid,
+ quart.g.workspace_membership,
+ )
+ return self.success(data={'invitations': [_invitation_payload(item) for item in invitations]})
+
+ data = await quart.request.get_json(silent=True) or {}
+ created = await self.ap.workspace_collaboration_service.create_invitation(
+ workspace_uuid,
+ quart.g.workspace_membership,
+ str(data.get('email', '')),
+ str(data.get('role', 'viewer')),
+ )
+ return self.success(
+ data={
+ 'invitation': _invitation_payload(created.invitation),
+ 'token': created.token,
+ }
+ )
+
+ @self.route(
+ '//invitations/',
+ methods=['DELETE'],
+ permission=Permission.MEMBER_INVITE,
+ )
+ async def _(
+ workspace_uuid: str,
+ invitation_uuid: str,
+ request_context: RequestContext,
+ ) -> typing.Any:
+ self._require_current_workspace(workspace_uuid, request_context)
+ if await self._requires_control_plane(workspace_uuid):
+ return self._control_plane_required()
+ invitation = await self.ap.workspace_collaboration_service.revoke_invitation(
+ workspace_uuid,
+ invitation_uuid,
+ quart.g.workspace_membership,
+ )
+ return self.success(data={'invitation': _invitation_payload(invitation)})
+
+ @self.route(
+ '//members/',
+ methods=['PATCH', 'DELETE'],
+ permission=Permission.MEMBER_UPDATE_ROLE,
+ )
+ async def _(
+ workspace_uuid: str,
+ account_uuid: str,
+ request_context: RequestContext,
+ ) -> typing.Any:
+ self._require_current_workspace(workspace_uuid, request_context)
+ if await self._requires_control_plane(workspace_uuid):
+ return self._control_plane_required()
+ if quart.request.method == 'DELETE':
+ if Permission.MEMBER_REMOVE.value not in request_context.workspace.permissions:
+ return self.http_status(403, 'permission_denied', 'Member removal permission is required')
+ member = await self.ap.workspace_collaboration_service.remove_member(
+ workspace_uuid,
+ account_uuid,
+ quart.g.workspace_membership,
+ )
+ return self.success(data={'account_uuid': member.account_uuid})
+
+ data = await quart.request.get_json(silent=True) or {}
+ member = await self.ap.workspace_collaboration_service.update_member_role(
+ workspace_uuid,
+ account_uuid,
+ str(data.get('role', '')),
+ quart.g.workspace_membership,
+ )
+ account = await self.ap.user_service.get_user_by_uuid(member.account_uuid)
+ return self.success(
+ data={
+ 'member': _membership_payload(
+ member,
+ email=account.user if account is not None else '',
+ )
+ }
+ )
+
+ @staticmethod
+ def _require_current_workspace(workspace_uuid: str, request_context: RequestContext) -> None:
+ if workspace_uuid != request_context.workspace_uuid:
+ raise WorkspaceNotFoundError('Workspace not found')
+
+ async def _requires_control_plane(self, workspace_uuid: str) -> bool:
+ workspace = await self.ap.workspace_service.get_workspace(workspace_uuid)
+ return workspace.source == WorkspaceSource.CLOUD_PROJECTION.value
+
+ def _control_plane_required(self) -> typing.Any:
+ return self.http_status(
+ 409,
+ 'control_plane_required',
+ 'Cloud Workspace membership and invitations are managed by the SaaS control plane',
+ )
+
+ @staticmethod
+ def _member_view_payload(view: WorkspaceMemberView) -> dict[str, typing.Any]:
+ return _membership_payload(view.membership, email=view.email)
+
+
+@group.group_class('invitations', '/api/v1/invitations')
+class InvitationsRouterGroup(group.RouterGroup):
+ async def initialize(self) -> None:
+ @self.route('/inspect', methods=['POST'], auth_type=group.AuthType.NONE)
+ async def _() -> typing.Any:
+ data = await quart.request.get_json(silent=True) or {}
+ invitation, workspace = await self.ap.workspace_collaboration_service.inspect_invitation(
+ str(data.get('token', ''))
+ )
+ return self.success(
+ data={
+ 'invitation': _invitation_payload(invitation),
+ 'workspace': _workspace_payload(workspace),
+ }
+ )
+
+ @self.route('/accept', methods=['POST'], auth_type=group.AuthType.NONE)
+ async def _() -> typing.Any:
+ data = await quart.request.get_json(silent=True) or {}
+ invitation_token = str(data.get('token', ''))
+ if not invitation_token:
+ return self.http_status(400, 'invitation_invalid', 'Invitation token is required')
+
+ authorization = quart.request.headers.get('Authorization', '')
+ if authorization.startswith('Bearer '):
+ account = await self.ap.user_service.get_authenticated_account(authorization.removeprefix('Bearer '))
+ if isinstance(account, str):
+ account = await self.ap.user_service.get_user_by_email(account)
+ if account is None:
+ return self.http_status(401, 'invalid_authentication', 'Account not found')
+ membership = await self.ap.workspace_collaboration_service.accept_invitation(
+ invitation_token,
+ account.uuid,
+ )
+ token = await self.ap.user_service.generate_jwt_token(account)
+ return self.success(data={'token': token, 'workspace_uuid': membership.workspace_uuid})
+
+ registration = data.get('registration')
+ if not isinstance(registration, dict):
+ return self.http_status(
+ 401,
+ 'account_exists_login_required',
+ 'Sign in or provide registration details to accept this invitation',
+ )
+ password = registration.get('password')
+ if not isinstance(password, str) or len(password) < 8:
+ return self.http_status(400, 'invalid_password', 'Password must contain at least 8 characters')
+ try:
+ _, membership, token = await self.ap.user_service.register_invited_account(
+ invitation_token,
+ str(registration.get('email', '')),
+ password,
+ )
+ except ControlPlaneDirectoryRequiredError as exc:
+ return self.http_status(409, exc.code, str(exc))
+ except AccountExistsLoginRequiredError as exc:
+ return self.http_status(409, exc.code, str(exc))
+ return self.success(data={'token': token, 'workspace_uuid': membership.workspace_uuid})
diff --git a/src/langbot/pkg/api/http/service/apikey.py b/src/langbot/pkg/api/http/service/apikey.py
index 207254351..9ae91dcc9 100644
--- a/src/langbot/pkg/api/http/service/apikey.py
+++ b/src/langbot/pkg/api/http/service/apikey.py
@@ -1,97 +1,252 @@
from __future__ import annotations
+import dataclasses
+import datetime
+import hashlib
import secrets
+import typing
+import uuid
+
import sqlalchemy
-from ....core import app
from ....entity.persistence import apikey
+from ....workspace.errors import WorkspaceNotFoundError
+from ..authz import Permission, PermissionDeniedError
+from .tenant import TenantContext, require_workspace_uuid, scope_statement
+
+if typing.TYPE_CHECKING:
+ from ....core.app import Application
+
+
+@dataclasses.dataclass(frozen=True, slots=True)
+class ApiKeyIdentity:
+ """Trusted Workspace identity derived from an API-key secret."""
+
+ instance_uuid: str
+ workspace_uuid: str
+ placement_generation: int
+ api_key_uuid: str
+ permissions: frozenset[str]
class ApiKeyService:
- ap: app.Application
+ """Manage hashed, Workspace-bound API keys."""
- def __init__(self, ap: app.Application) -> None:
+ def __init__(self, ap: Application) -> None:
self.ap = ap
- async def get_api_keys(self) -> list[dict]:
- """Get all API keys"""
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(apikey.ApiKey))
+ @staticmethod
+ def _hash_secret(secret: str) -> str:
+ return hashlib.sha256(secret.encode('utf-8')).hexdigest()
- keys = result.all()
- return [self.ap.persistence_mgr.serialize_model(apikey.ApiKey, key) for key in keys]
+ @staticmethod
+ def _utcnow() -> datetime.datetime:
+ return datetime.datetime.now(datetime.UTC).replace(tzinfo=None)
- async def create_api_key(self, name: str, description: str = '') -> dict:
- """Create a new API key"""
- # Generate a secure random API key
- key = f'lbk_{secrets.token_urlsafe(32)}'
+ @staticmethod
+ def _normalize_scopes(
+ scopes: typing.Iterable[str] | None,
+ *,
+ default: typing.Iterable[str] = (),
+ ) -> list[str]:
+ requested = list(default if scopes is None else scopes)
+ valid = {permission.value for permission in Permission}
+ normalized: list[str] = []
+ for scope in requested:
+ if not isinstance(scope, str):
+ raise ValueError('API key scopes must be strings')
+ value = scope.strip()
+ if value not in valid:
+ raise ValueError(f'Unknown API key scope: {value}')
+ if value not in normalized:
+ normalized.append(value)
+ return normalized
- key_data = {'name': name, 'key': key, 'description': description}
+ def _serialize(self, row: typing.Any) -> dict[str, typing.Any]:
+ value = self.ap.persistence_mgr.serialize_model(apikey.ApiKey, row)
+ value.pop('key_hash', None)
+ # The secret is deliberately unrecoverable after creation.
+ value['secret_available'] = False
+ return value
- await self.ap.persistence_mgr.execute_async(sqlalchemy.insert(apikey.ApiKey).values(**key_data))
-
- # Retrieve the created key
+ async def get_api_keys(self, context: TenantContext) -> list[dict[str, typing.Any]]:
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(apikey.ApiKey).where(apikey.ApiKey.key == key)
+ scope_statement(
+ sqlalchemy.select(apikey.ApiKey).order_by(apikey.ApiKey.created_at, apikey.ApiKey.id),
+ apikey.ApiKey,
+ context,
+ )
)
- created_key = result.first()
+ return [self._serialize(key) for key in result.all()]
- return self.ap.persistence_mgr.serialize_model(apikey.ApiKey, created_key)
+ async def create_api_key(
+ self,
+ context: TenantContext,
+ name: str,
+ description: str = '',
+ *,
+ scopes: typing.Iterable[str] | None = None,
+ expires_at: datetime.datetime | None = None,
+ ) -> dict[str, typing.Any]:
+ workspace_uuid = require_workspace_uuid(context)
+ normalized_name = name.strip()
+ if not normalized_name:
+ raise ValueError('Name is required')
+ if expires_at is not None:
+ if expires_at.tzinfo is not None:
+ expires_at = expires_at.astimezone(datetime.UTC).replace(tzinfo=None)
+ if expires_at <= self._utcnow():
+ raise ValueError('API key expiry must be in the future')
- async def get_api_key(self, key_id: int) -> dict | None:
- """Get a specific API key by ID"""
+ default_scopes = getattr(getattr(context, 'workspace', None), 'permissions', frozenset())
+ normalized_scopes = self._normalize_scopes(scopes, default=default_scopes)
+ allowed_scopes = frozenset(default_scopes)
+ unauthorized_scopes = sorted(set(normalized_scopes) - allowed_scopes)
+ if unauthorized_scopes:
+ # API-key management delegates the caller's authority; it must not
+ # become a path for minting a stronger principal.
+ raise PermissionDeniedError(unauthorized_scopes[0])
+ secret = f'lbk_{secrets.token_urlsafe(32)}'
+ key_uuid = str(uuid.uuid4())
+ await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.insert(apikey.ApiKey).values(
+ uuid=key_uuid,
+ workspace_uuid=workspace_uuid,
+ created_by_account_uuid=getattr(context, 'account_uuid', None),
+ name=normalized_name,
+ key_hash=self._hash_secret(secret),
+ scopes=normalized_scopes,
+ status=apikey.ApiKeyStatus.ACTIVE.value,
+ expires_at=expires_at,
+ description=description.strip(),
+ )
+ )
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(apikey.ApiKey).where(apikey.ApiKey.id == key_id)
+ scope_statement(
+ sqlalchemy.select(apikey.ApiKey).where(apikey.ApiKey.uuid == key_uuid),
+ apikey.ApiKey,
+ workspace_uuid,
+ )
)
+ created = result.first()
+ if created is None:
+ raise RuntimeError('Created API key could not be loaded')
+ value = self._serialize(created)
+ value['key'] = secret
+ value['secret_available'] = True
+ return value
+ async def get_api_key(self, context: TenantContext, key_id: int) -> dict[str, typing.Any] | None:
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.select(apikey.ApiKey).where(apikey.ApiKey.id == key_id),
+ apikey.ApiKey,
+ context,
+ )
+ )
key = result.first()
+ return None if key is None else self._serialize(key)
- if key is None:
+ async def authenticate_api_key(self, secret: str) -> ApiKeyIdentity | None:
+ """Authenticate a secret and derive its Workspace without trusting headers."""
+
+ if not isinstance(secret, str) or not secret.strip():
return None
- return self.ap.persistence_mgr.serialize_model(apikey.ApiKey, key)
-
- async def verify_api_key(self, key: str) -> bool:
- """Verify if an API key is valid.
-
- A key is accepted if it matches the global API key configured in
- ``config.yaml`` (``api.global_api_key``) — which requires no login
- session and no database record — or if it matches a key created via
- the web UI (stored in the database, prefixed with ``lbk_``).
- """
- if not isinstance(key, str) or not key:
- return False
-
- # 1. Global API key from config.yaml (no DB lookup, no login state).
- # Note: config completion only backfills top-level keys, so existing
- # installs may not have this key — access it defensively.
- global_api_key = self.ap.instance_config.data.get('api', {}).get('global_api_key', '')
- if global_api_key and secrets.compare_digest(key, global_api_key):
- return True
-
- # 2. Web-UI-created keys are stored in the database and prefixed lbk_.
- if not key.startswith('lbk_'):
- return False
+ global_secret = self.ap.instance_config.data.get('api', {}).get('global_api_key', '')
+ if global_secret and secrets.compare_digest(secret, global_secret):
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ if workspace_service is None or workspace_service.policy.multi_workspace_enabled:
+ return None
+ binding = await workspace_service.get_local_execution_binding()
+ return ApiKeyIdentity(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ api_key_uuid='global-oss-api-key',
+ permissions=frozenset(permission.value for permission in Permission),
+ )
+ if not secret.startswith('lbk_'):
+ return None
+ secret_hash = self._hash_secret(secret)
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(apikey.ApiKey).where(apikey.ApiKey.key == key)
+ sqlalchemy.select(apikey.ApiKey).where(apikey.ApiKey.key_hash == secret_hash)
+ )
+ key = result.first()
+ if key is None or key.status != apikey.ApiKeyStatus.ACTIVE.value:
+ return None
+ now = self._utcnow()
+ if key.expires_at is not None and key.expires_at <= now:
+ return None
+
+ raw_scopes = list(key.scopes or [])
+ permissions = (
+ frozenset(permission.value for permission in Permission)
+ if '*' in raw_scopes
+ else frozenset(self._normalize_scopes(raw_scopes))
+ )
+ binding = await self.ap.workspace_service.get_execution_binding(key.workspace_uuid)
+ await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.update(apikey.ApiKey)
+ .where(
+ apikey.ApiKey.id == key.id,
+ apikey.ApiKey.workspace_uuid == key.workspace_uuid,
+ apikey.ApiKey.key_hash == secret_hash,
+ )
+ .values(last_used_at=now)
+ )
+ return ApiKeyIdentity(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ api_key_uuid=key.uuid,
+ permissions=permissions,
)
- key_obj = result.first()
- return key_obj is not None
+ async def verify_api_key(self, secret: str) -> bool:
+ try:
+ return await self.authenticate_api_key(secret) is not None
+ except Exception:
+ return False
- async def delete_api_key(self, key_id: int) -> None:
- """Delete an API key"""
- await self.ap.persistence_mgr.execute_async(sqlalchemy.delete(apikey.ApiKey).where(apikey.ApiKey.id == key_id))
-
- async def update_api_key(self, key_id: int, name: str = None, description: str = None) -> None:
- """Update an API key's metadata (name, description)"""
- update_data = {}
- if name is not None:
- update_data['name'] = name
- if description is not None:
- update_data['description'] = description
-
- if update_data:
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(apikey.ApiKey).where(apikey.ApiKey.id == key_id).values(**update_data)
+ async def delete_api_key(self, context: TenantContext, key_id: int) -> None:
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.update(apikey.ApiKey)
+ .where(apikey.ApiKey.id == key_id)
+ .values(status=apikey.ApiKeyStatus.REVOKED.value),
+ apikey.ApiKey,
+ context,
)
+ )
+ if getattr(result, 'rowcount', 0) == 0:
+ raise WorkspaceNotFoundError('API key not found')
+
+ async def update_api_key(
+ self,
+ context: TenantContext,
+ key_id: int,
+ name: str | None = None,
+ description: str | None = None,
+ ) -> None:
+ update_data: dict[str, typing.Any] = {}
+ if name is not None:
+ normalized_name = name.strip()
+ if not normalized_name:
+ raise ValueError('Name is required')
+ update_data['name'] = normalized_name
+ if description is not None:
+ update_data['description'] = description.strip()
+ if not update_data:
+ return
+
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.update(apikey.ApiKey).where(apikey.ApiKey.id == key_id).values(**update_data),
+ apikey.ApiKey,
+ context,
+ )
+ )
+ if getattr(result, 'rowcount', 0) == 0:
+ raise WorkspaceNotFoundError('API key not found')
diff --git a/src/langbot/pkg/api/http/service/bot.py b/src/langbot/pkg/api/http/service/bot.py
index 995267cf5..9d9211be3 100644
--- a/src/langbot/pkg/api/http/service/bot.py
+++ b/src/langbot/pkg/api/http/service/bot.py
@@ -2,11 +2,12 @@ from __future__ import annotations
import uuid
import sqlalchemy
-import typing
from ....core import app
from ....entity.persistence import bot as persistence_bot
from ....entity.persistence import pipeline as persistence_pipeline
+from ....workspace.errors import WorkspaceNotFoundError
+from .tenant import TenantContext, require_workspace_uuid, scope_statement
class BotService:
@@ -17,9 +18,11 @@ class BotService:
def __init__(self, ap: app.Application) -> None:
self.ap = ap
- async def get_bots(self, include_secret: bool = True) -> list[dict]:
+ async def get_bots(self, context: TenantContext, include_secret: bool = False) -> list[dict]:
"""获取所有机器人"""
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_bot.Bot))
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(sqlalchemy.select(persistence_bot.Bot), persistence_bot.Bot, context)
+ )
bots = result.all()
@@ -29,10 +32,14 @@ class BotService:
return [self.ap.persistence_mgr.serialize_model(persistence_bot.Bot, bot, masked_columns) for bot in bots]
- async def get_bot(self, bot_uuid: str, include_secret: bool = True) -> dict | None:
+ async def get_bot(self, context: TenantContext, bot_uuid: str, include_secret: bool = False) -> dict | None:
"""获取机器人"""
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_bot.Bot).where(persistence_bot.Bot.uuid == bot_uuid)
+ scope_statement(
+ sqlalchemy.select(persistence_bot.Bot).where(persistence_bot.Bot.uuid == bot_uuid),
+ persistence_bot.Bot,
+ context,
+ )
)
bot = result.first()
@@ -46,15 +53,20 @@ class BotService:
return self.ap.persistence_mgr.serialize_model(persistence_bot.Bot, bot, masked_columns)
- async def get_runtime_bot_info(self, bot_uuid: str, include_secret: bool = True) -> dict:
+ async def get_runtime_bot_info(
+ self,
+ context: TenantContext,
+ bot_uuid: str,
+ include_secret: bool = False,
+ ) -> dict:
"""获取机器人运行时信息"""
- persistence_bot = await self.get_bot(bot_uuid, include_secret)
+ persistence_bot = await self.get_bot(context, bot_uuid, include_secret)
if persistence_bot is None:
- raise Exception('Bot not found')
+ raise WorkspaceNotFoundError('Bot not found')
adapter_runtime_values = {}
- runtime_bot = await self.ap.platform_mgr.get_bot_by_uuid(bot_uuid)
+ runtime_bot = await self.ap.platform_mgr.get_bot_by_uuid(context, bot_uuid)
if runtime_bot is not None:
adapter_runtime_values['bot_account_id'] = runtime_bot.adapter.bot_account_id
@@ -86,22 +98,29 @@ class BotService:
return persistence_bot
- async def create_bot(self, bot_data: dict) -> str:
+ async def create_bot(self, context: TenantContext, bot_data: dict) -> str:
"""Create bot"""
+ workspace_uuid = require_workspace_uuid(context)
# Check limitation
limitation = self.ap.instance_config.data.get('system', {}).get('limitation', {})
max_bots = limitation.get('max_bots', -1)
if max_bots >= 0:
- existing_bots = await self.get_bots()
+ existing_bots = await self.get_bots(context)
if len(existing_bots) >= max_bots:
raise ValueError(f'Maximum number of bots ({max_bots}) reached')
# TODO: 检查配置信息格式
+ bot_data = bot_data.copy()
bot_data['uuid'] = str(uuid.uuid4())
+ bot_data['workspace_uuid'] = workspace_uuid
# bind the most recently updated pipeline if any exist
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_pipeline.LegacyPipeline)
+ scope_statement(
+ sqlalchemy.select(persistence_pipeline.LegacyPipeline),
+ persistence_pipeline.LegacyPipeline,
+ context,
+ )
.order_by(persistence_pipeline.LegacyPipeline.updated_at.desc())
.limit(1)
)
@@ -112,61 +131,84 @@ class BotService:
await self.ap.persistence_mgr.execute_async(sqlalchemy.insert(persistence_bot.Bot).values(bot_data))
- bot = await self.get_bot(bot_data['uuid'])
+ bot = await self.get_bot(context, bot_data['uuid'], include_secret=True)
- await self.ap.platform_mgr.load_bot(bot)
+ await self.ap.platform_mgr.load_bot(context, bot)
return bot_data['uuid']
- async def update_bot(self, bot_uuid: str, bot_data: dict) -> None:
+ async def update_bot(self, context: TenantContext, bot_uuid: str, bot_data: dict) -> None:
"""Update bot"""
+ workspace_uuid = require_workspace_uuid(context)
update_data = bot_data.copy()
- if 'uuid' in update_data:
- del update_data['uuid']
+ update_data.pop('uuid', None)
+ update_data.pop('workspace_uuid', None)
# set use_pipeline_name
if 'use_pipeline_uuid' in update_data:
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
- persistence_pipeline.LegacyPipeline.uuid == update_data['use_pipeline_uuid']
+ scope_statement(
+ sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
+ persistence_pipeline.LegacyPipeline.uuid == update_data['use_pipeline_uuid']
+ ),
+ persistence_pipeline.LegacyPipeline,
+ workspace_uuid,
)
)
pipeline = result.first()
if pipeline is not None:
update_data['use_pipeline_name'] = pipeline.name
else:
- raise Exception('Pipeline not found')
+ raise WorkspaceNotFoundError('Pipeline not found')
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(persistence_bot.Bot).values(update_data).where(persistence_bot.Bot.uuid == bot_uuid)
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.update(persistence_bot.Bot).values(update_data).where(persistence_bot.Bot.uuid == bot_uuid),
+ persistence_bot.Bot,
+ workspace_uuid,
+ )
)
- await self.ap.platform_mgr.remove_bot(bot_uuid)
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Bot not found')
+ await self.ap.platform_mgr.remove_bot(context, bot_uuid)
# select from db
- bot = await self.get_bot(bot_uuid)
+ bot = await self.get_bot(context, bot_uuid, include_secret=True)
- runtime_bot = await self.ap.platform_mgr.load_bot(bot)
+ runtime_bot = await self.ap.platform_mgr.load_bot(context, bot)
if runtime_bot.enable:
await runtime_bot.run()
# update all conversation that use this bot
for session in self.ap.sess_mgr.session_list:
- if session.using_conversation is not None and session.using_conversation.bot_uuid == bot_uuid:
+ if (
+ session.using_conversation is not None
+ and session.using_conversation.bot_uuid == bot_uuid
+ and getattr(session, 'workspace_uuid', workspace_uuid) == workspace_uuid
+ ):
session.using_conversation = None
- async def delete_bot(self, bot_uuid: str) -> None:
+ async def delete_bot(self, context: TenantContext, bot_uuid: str) -> None:
"""Delete bot"""
- await self.ap.platform_mgr.remove_bot(bot_uuid)
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_bot.Bot).where(persistence_bot.Bot.uuid == bot_uuid)
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.delete(persistence_bot.Bot).where(persistence_bot.Bot.uuid == bot_uuid),
+ persistence_bot.Bot,
+ context,
+ )
)
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Bot not found')
+ await self.ap.platform_mgr.remove_bot(context, bot_uuid)
async def list_event_logs(
- self, bot_uuid: str, from_index: int, max_count: int
- ) -> typing.Tuple[list[dict], int, int, int]:
- runtime_bot = await self.ap.platform_mgr.get_bot_by_uuid(bot_uuid)
+ self, context: TenantContext, bot_uuid: str, from_index: int, max_count: int
+ ) -> tuple[list[dict], int]:
+ if await self.get_bot(context, bot_uuid, include_secret=False) is None:
+ raise WorkspaceNotFoundError('Bot not found')
+ runtime_bot = await self.ap.platform_mgr.get_bot_by_uuid(context, bot_uuid)
if runtime_bot is None:
raise Exception('Bot not found')
@@ -174,7 +216,14 @@ class BotService:
return [log.to_json() for log in logs], total_count
- async def send_message(self, bot_uuid: str, target_type: str, target_id: str, message_chain_data: dict) -> None:
+ async def send_message(
+ self,
+ context: TenantContext,
+ bot_uuid: str,
+ target_type: str,
+ target_id: str,
+ message_chain_data: dict,
+ ) -> None:
"""Send message to a specific target via bot
Args:
@@ -183,11 +232,14 @@ class BotService:
target_id: The ID of the target
message_chain_data: The message chain data in dict format
"""
+ if await self.get_bot(context, bot_uuid, include_secret=False) is None:
+ raise WorkspaceNotFoundError('Bot not found')
+
# Import here to avoid circular imports
import langbot_plugin.api.entities.builtin.platform.message as platform_message
# Get runtime bot
- runtime_bot = await self.ap.platform_mgr.get_bot_by_uuid(bot_uuid)
+ runtime_bot = await self.ap.platform_mgr.get_bot_by_uuid(context, bot_uuid)
if runtime_bot is None:
raise Exception(f'Bot not found: {bot_uuid}')
@@ -202,19 +254,29 @@ class BotService:
# ============ Bot Admins ============
- async def get_bot_admins(self, bot_uuid: str) -> list[dict]:
+ async def get_bot_admins(self, context: TenantContext, bot_uuid: str) -> list[dict]:
from ....entity.persistence import bot as persistence_bot
+ if await self.get_bot(context, bot_uuid, include_secret=False) is None:
+ raise WorkspaceNotFoundError('Bot not found')
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_bot.BotAdmin).where(persistence_bot.BotAdmin.bot_uuid == bot_uuid)
+ scope_statement(
+ sqlalchemy.select(persistence_bot.BotAdmin).where(persistence_bot.BotAdmin.bot_uuid == bot_uuid),
+ persistence_bot.BotAdmin,
+ context,
+ )
)
return [{'id': r.id, 'launcher_type': r.launcher_type, 'launcher_id': r.launcher_id} for r in result.all()]
- async def add_bot_admin(self, bot_uuid: str, launcher_type: str, launcher_id: str) -> int:
+ async def add_bot_admin(self, context: TenantContext, bot_uuid: str, launcher_type: str, launcher_id: str) -> int:
from ....entity.persistence import bot as persistence_bot
+ workspace_uuid = require_workspace_uuid(context)
+ if await self.get_bot(context, bot_uuid, include_secret=False) is None:
+ raise WorkspaceNotFoundError('Bot not found')
result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.insert(persistence_bot.BotAdmin).values(
+ workspace_uuid=workspace_uuid,
bot_uuid=bot_uuid,
launcher_type=launcher_type,
launcher_id=launcher_id,
@@ -222,12 +284,18 @@ class BotService:
)
return result.inserted_primary_key[0]
- async def delete_bot_admin(self, bot_uuid: str, admin_id: int) -> None:
+ async def delete_bot_admin(self, context: TenantContext, bot_uuid: str, admin_id: int) -> None:
from ....entity.persistence import bot as persistence_bot
+ if await self.get_bot(context, bot_uuid, include_secret=False) is None:
+ raise WorkspaceNotFoundError('Bot not found')
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_bot.BotAdmin).where(
- persistence_bot.BotAdmin.bot_uuid == bot_uuid,
- persistence_bot.BotAdmin.id == admin_id,
+ scope_statement(
+ sqlalchemy.delete(persistence_bot.BotAdmin).where(
+ persistence_bot.BotAdmin.bot_uuid == bot_uuid,
+ persistence_bot.BotAdmin.id == admin_id,
+ ),
+ persistence_bot.BotAdmin,
+ context,
)
)
diff --git a/src/langbot/pkg/api/http/service/knowledge.py b/src/langbot/pkg/api/http/service/knowledge.py
index 48cb7cace..09fb499ff 100644
--- a/src/langbot/pkg/api/http/service/knowledge.py
+++ b/src/langbot/pkg/api/http/service/knowledge.py
@@ -2,8 +2,13 @@ from __future__ import annotations
import sqlalchemy
+from ....api.http.authz import WorkspaceRequiredError
+from ....api.http.context import ExecutionContext, RequestContext
from ....core import app
from ....entity.persistence import rag as persistence_rag
+from ....workspace.errors import WorkspaceNotFoundError
+from .secrets import redact_secrets, restore_secret_placeholders
+from .tenant import TenantContext, require_workspace_uuid
class KnowledgeService:
@@ -14,16 +19,41 @@ class KnowledgeService:
def __init__(self, ap: app.Application) -> None:
self.ap = ap
- async def get_knowledge_bases(self) -> list[dict]:
+ @staticmethod
+ def _execution_context(context: RequestContext | ExecutionContext) -> ExecutionContext:
+ if isinstance(context, RequestContext):
+ return ExecutionContext.from_request(context)
+ if isinstance(context, ExecutionContext):
+ return context
+ raise WorkspaceRequiredError('RequestContext or ExecutionContext is required')
+
+ async def get_knowledge_bases(self, context: TenantContext, *, include_secret: bool = False) -> list[dict]:
"""获取所有知识库"""
- return await self.ap.rag_mgr.get_all_knowledge_base_details()
+ require_workspace_uuid(context)
+ knowledge_bases = await self.ap.rag_mgr.get_all_knowledge_base_details(context)
+ return knowledge_bases if include_secret else [redact_secrets(base) for base in knowledge_bases]
- async def get_knowledge_base(self, kb_uuid: str) -> dict | None:
+ async def get_knowledge_base(
+ self,
+ context: TenantContext,
+ kb_uuid: str,
+ *,
+ include_secret: bool = False,
+ ) -> dict | None:
"""获取知识库"""
- return await self.ap.rag_mgr.get_knowledge_base_details(kb_uuid)
+ require_workspace_uuid(context)
+ knowledge_base = await self.ap.rag_mgr.get_knowledge_base_details(context, kb_uuid)
+ if knowledge_base is None or include_secret:
+ return knowledge_base
+ return redact_secrets(knowledge_base)
- async def create_knowledge_base(self, kb_data: dict) -> str:
+ async def create_knowledge_base(
+ self,
+ context: RequestContext | ExecutionContext,
+ kb_data: dict,
+ ) -> str:
"""创建知识库"""
+ require_workspace_uuid(context)
# In new architecture, we delegate entirely to RAGManager which uses plugins.
# Legacy internal KB creation is removed.
@@ -31,17 +61,19 @@ class KnowledgeService:
if not knowledge_engine_plugin_id:
raise ValueError('knowledge_engine_plugin_id is required')
- creation_settings = kb_data.get('creation_settings', {})
+ creation_settings = restore_secret_placeholders(kb_data.get('creation_settings', {}))
retrieval_settings = kb_data.get('retrieval_settings', {})
# Validate required fields based on plugin's creation_schema and retrieval_schema
await self._validate_schema_required_fields(
+ context,
knowledge_engine_plugin_id,
creation_settings,
retrieval_settings,
)
kb = await self.ap.rag_mgr.create_knowledge_base(
+ context,
name=kb_data.get('name', 'Untitled'),
knowledge_engine_plugin_id=knowledge_engine_plugin_id,
creation_settings=creation_settings,
@@ -52,6 +84,7 @@ class KnowledgeService:
async def _validate_schema_required_fields(
self,
+ context: RequestContext | ExecutionContext,
plugin_id: str,
creation_settings: dict,
retrieval_settings: dict,
@@ -69,7 +102,11 @@ class KnowledgeService:
Raises:
ValueError: If any required field is missing or empty.
"""
+ if not self.ap.plugin_connector.is_enable_plugin:
+ return
+
# Validate creation_schema
+ await self.ap.plugin_connector.require_workspace_context(context)
try:
creation_schema = await self.ap.plugin_connector.get_rag_creation_schema(plugin_id)
self._check_required_fields(creation_schema, creation_settings, 'creation_settings')
@@ -79,6 +116,7 @@ class KnowledgeService:
self.ap.logger.warning(f'Failed to get creation_schema for validation: {e}')
# Validate retrieval_schema
+ await self.ap.plugin_connector.require_workspace_context(context)
try:
retrieval_schema = await self.ap.plugin_connector.get_rag_retrieval_schema(plugin_id)
self._check_required_fields(retrieval_schema, retrieval_settings, 'retrieval_settings')
@@ -151,8 +189,16 @@ class KnowledgeService:
)
raise ValueError(f'{field_label} is required ({context}.{field_name})')
- async def update_knowledge_base(self, kb_uuid: str, kb_data: dict) -> None:
+ async def update_knowledge_base(
+ self,
+ context: RequestContext | ExecutionContext,
+ kb_uuid: str,
+ kb_data: dict,
+ ) -> None:
"""更新知识库"""
+ workspace_uuid = require_workspace_uuid(context)
+ if await self.get_knowledge_base(context, kb_uuid) is None:
+ raise WorkspaceNotFoundError('Knowledge base not found')
# Filter to only mutable fields
filtered_data = {k: v for k, v in kb_data.items() if k in persistence_rag.KnowledgeBase.MUTABLE_FIELDS}
@@ -162,17 +208,18 @@ class KnowledgeService:
await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(persistence_rag.KnowledgeBase)
.values(filtered_data)
+ .where(persistence_rag.KnowledgeBase.workspace_uuid == workspace_uuid)
.where(persistence_rag.KnowledgeBase.uuid == kb_uuid)
)
- await self.ap.rag_mgr.remove_knowledge_base_from_runtime(kb_uuid)
+ await self.ap.rag_mgr.remove_knowledge_base_from_runtime(context, kb_uuid)
- kb = await self.get_knowledge_base(kb_uuid)
+ kb = await self.get_knowledge_base(context, kb_uuid, include_secret=True)
if kb is None:
- raise Exception('Knowledge base not found after update')
+ raise WorkspaceNotFoundError('Knowledge base not found')
- await self.ap.rag_mgr.load_knowledge_base(kb)
+ await self.ap.rag_mgr.load_knowledge_base(context, kb)
- async def _check_doc_capability(self, kb_uuid: str, operation: str) -> None:
+ async def _check_doc_capability(self, context: TenantContext, kb_uuid: str, operation: str) -> None:
"""Check if the KB's Knowledge Engine supports document operations.
Args:
@@ -182,104 +229,145 @@ class KnowledgeService:
Raises:
Exception: If the KB does not support doc_ingestion.
"""
- kb_info = await self.ap.rag_mgr.get_knowledge_base_details(kb_uuid)
+ kb_info = await self.ap.rag_mgr.get_knowledge_base_details(context, kb_uuid)
if not kb_info:
- raise Exception('Knowledge base not found')
+ raise WorkspaceNotFoundError('Knowledge base not found')
capabilities = kb_info.get('knowledge_engine', {}).get('capabilities', [])
if 'doc_ingestion' not in capabilities:
raise Exception(f'This knowledge base does not support {operation}')
- async def store_file(self, kb_uuid: str, file_id: str, parser_plugin_id: str | None = None) -> str:
+ async def store_file(
+ self,
+ context: RequestContext | ExecutionContext,
+ kb_uuid: str,
+ file_id: str,
+ parser_plugin_id: str | None = None,
+ ) -> str:
"""存储文件"""
- runtime_kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(kb_uuid)
+ execution_context = self._execution_context(context)
+ runtime_kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(execution_context, kb_uuid)
if runtime_kb is None:
- raise Exception('Knowledge base not found')
+ raise WorkspaceNotFoundError('Knowledge base not found')
- await self._check_doc_capability(kb_uuid, 'document upload')
+ await self._check_doc_capability(context, kb_uuid, 'document upload')
- result = await runtime_kb.store_file(file_id, parser_plugin_id=parser_plugin_id)
+ result = await runtime_kb.store_file(execution_context, file_id, parser_plugin_id=parser_plugin_id)
# Update the KB's updated_at timestamp
await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(persistence_rag.KnowledgeBase)
.values(updated_at=sqlalchemy.func.now())
+ .where(persistence_rag.KnowledgeBase.workspace_uuid == execution_context.workspace_uuid)
.where(persistence_rag.KnowledgeBase.uuid == kb_uuid)
)
return result
async def retrieve_knowledge_base(
- self, kb_uuid: str, query: str, retrieval_settings: dict | None = None
+ self,
+ context: RequestContext | ExecutionContext,
+ kb_uuid: str,
+ query: str,
+ retrieval_settings: dict | None = None,
) -> list[dict]:
"""检索知识库"""
- runtime_kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(kb_uuid)
+ execution_context = self._execution_context(context)
+ runtime_kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(execution_context, kb_uuid)
if runtime_kb is None:
- raise Exception('Knowledge base not found')
+ raise WorkspaceNotFoundError('Knowledge base not found')
# Pass retrieval_settings
- results = await runtime_kb.retrieve(query, settings=retrieval_settings)
+ results = await runtime_kb.retrieve(execution_context, query, settings=retrieval_settings)
return [result.model_dump() for result in results]
- async def get_files_by_knowledge_base(self, kb_uuid: str) -> list[dict]:
+ async def get_files_by_knowledge_base(self, context: TenantContext, kb_uuid: str) -> list[dict]:
"""获取知识库文件"""
+ workspace_uuid = require_workspace_uuid(context)
+ if await self.get_knowledge_base(context, kb_uuid) is None:
+ raise WorkspaceNotFoundError('Knowledge base not found')
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_rag.File).where(persistence_rag.File.kb_id == kb_uuid)
+ sqlalchemy.select(persistence_rag.File)
+ .where(persistence_rag.File.workspace_uuid == workspace_uuid)
+ .where(persistence_rag.File.kb_id == kb_uuid)
)
files = result.all()
return [self.ap.persistence_mgr.serialize_model(persistence_rag.File, file) for file in files]
- async def delete_file(self, kb_uuid: str, file_id: str) -> None:
+ async def delete_file(
+ self,
+ context: RequestContext | ExecutionContext,
+ kb_uuid: str,
+ file_id: str,
+ ) -> None:
"""删除文件"""
- runtime_kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(kb_uuid)
+ execution_context = self._execution_context(context)
+ runtime_kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(execution_context, kb_uuid)
if runtime_kb is None:
- raise Exception('Knowledge base not found')
+ raise WorkspaceNotFoundError('Knowledge base not found')
- await self._check_doc_capability(kb_uuid, 'document deletion')
+ await self._check_doc_capability(context, kb_uuid, 'document deletion')
- await runtime_kb.delete_file(file_id)
+ await runtime_kb.delete_file(execution_context, file_id)
# Update the KB's updated_at timestamp
await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(persistence_rag.KnowledgeBase)
.values(updated_at=sqlalchemy.func.now())
+ .where(persistence_rag.KnowledgeBase.workspace_uuid == execution_context.workspace_uuid)
.where(persistence_rag.KnowledgeBase.uuid == kb_uuid)
)
- async def delete_knowledge_base(self, kb_uuid: str) -> None:
+ async def delete_knowledge_base(
+ self,
+ context: RequestContext | ExecutionContext,
+ kb_uuid: str,
+ ) -> None:
"""删除知识库"""
- # Delete from DB first to commit the deletion, then clean up runtime/plugin (best-effort)
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_rag.KnowledgeBase).where(persistence_rag.KnowledgeBase.uuid == kb_uuid)
- )
+ workspace_uuid = require_workspace_uuid(context)
+ if await self.get_knowledge_base(context, kb_uuid) is None:
+ raise WorkspaceNotFoundError('Knowledge base not found')
# delete files
# NOTE: Chunk cleanup is for legacy (pre-plugin) KBs that stored chunks locally.
# For plugin-based Knowledge Engines, the Chunk table is not populated, so this is a no-op.
files = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_rag.File).where(persistence_rag.File.kb_id == kb_uuid)
+ sqlalchemy.select(persistence_rag.File)
+ .where(persistence_rag.File.workspace_uuid == workspace_uuid)
+ .where(persistence_rag.File.kb_id == kb_uuid)
)
for file in files:
# delete chunks
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_rag.Chunk).where(persistence_rag.Chunk.file_id == file.uuid)
+ sqlalchemy.delete(persistence_rag.Chunk)
+ .where(persistence_rag.Chunk.workspace_uuid == workspace_uuid)
+ .where(persistence_rag.Chunk.file_id == file.uuid)
)
# delete file
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_rag.File).where(persistence_rag.File.uuid == file.uuid)
+ sqlalchemy.delete(persistence_rag.File)
+ .where(persistence_rag.File.workspace_uuid == workspace_uuid)
+ .where(persistence_rag.File.uuid == file.uuid)
)
- # Remove from runtime and notify plugin (best-effort, DB is already cleaned up)
- await self.ap.rag_mgr.delete_knowledge_base(kb_uuid)
+ # Remove from runtime and notify plugin before deleting the owning row.
+ await self.ap.rag_mgr.delete_knowledge_base(context, kb_uuid)
+ await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.delete(persistence_rag.KnowledgeBase)
+ .where(persistence_rag.KnowledgeBase.workspace_uuid == workspace_uuid)
+ .where(persistence_rag.KnowledgeBase.uuid == kb_uuid)
+ )
# ================= Knowledge Engine Discovery =================
- async def list_knowledge_engines(self) -> list[dict]:
+ async def list_knowledge_engines(self, context: TenantContext) -> list[dict]:
"""List all available Knowledge Engines from plugins."""
+ require_workspace_uuid(context)
engines = []
if not self.ap.plugin_connector.is_enable_plugin:
return engines
+ await self.ap.plugin_connector.require_workspace_context(context)
# Get KnowledgeEngine plugins
try:
@@ -290,10 +378,12 @@ class KnowledgeService:
return engines
- async def list_parsers(self, mime_type: str | None = None) -> list[dict]:
+ async def list_parsers(self, context: TenantContext, mime_type: str | None = None) -> list[dict]:
"""List available parsers, optionally filtered by MIME type."""
+ require_workspace_uuid(context)
if not self.ap.plugin_connector.is_enable_plugin:
return []
+ await self.ap.plugin_connector.require_workspace_context(context)
try:
parsers = await self.ap.plugin_connector.list_parsers()
if mime_type:
@@ -303,16 +393,24 @@ class KnowledgeService:
self.ap.logger.warning(f'Failed to list parsers: {e}')
return []
- async def get_engine_creation_schema(self, plugin_id: str) -> dict:
+ async def get_engine_creation_schema(self, context: TenantContext, plugin_id: str) -> dict:
"""Get creation settings schema for a specific Knowledge Engine."""
+ require_workspace_uuid(context)
+ if not self.ap.plugin_connector.is_enable_plugin:
+ return {}
+ await self.ap.plugin_connector.require_workspace_context(context)
try:
return await self.ap.plugin_connector.get_rag_creation_schema(plugin_id)
except Exception as e:
self.ap.logger.warning(f'Failed to get creation schema for {plugin_id}: {e}')
return {}
- async def get_engine_retrieval_schema(self, plugin_id: str) -> dict:
+ async def get_engine_retrieval_schema(self, context: TenantContext, plugin_id: str) -> dict:
"""Get retrieval settings schema for a specific Knowledge Engine."""
+ require_workspace_uuid(context)
+ if not self.ap.plugin_connector.is_enable_plugin:
+ return {}
+ await self.ap.plugin_connector.require_workspace_context(context)
try:
return await self.ap.plugin_connector.get_rag_retrieval_schema(plugin_id)
except Exception as e:
diff --git a/src/langbot/pkg/api/http/service/maintenance.py b/src/langbot/pkg/api/http/service/maintenance.py
index fa7359cba..59c4e1e3d 100644
--- a/src/langbot/pkg/api/http/service/maintenance.py
+++ b/src/langbot/pkg/api/http/service/maintenance.py
@@ -11,11 +11,15 @@ import sqlalchemy
from ....core import app
from ....entity.persistence import bstorage as persistence_bstorage
from ....entity.persistence import monitoring as persistence_monitoring
+from ..authz import WorkspaceRequiredError
+from ..context import ExecutionContext
+from .tenant import TenantContext, require_workspace_uuid
LOG_FILE_PATTERN = re.compile(r'^langbot-(\d{4}-\d{2}-\d{2})\.log(?:\.\d+)?$')
DEFAULT_UPLOAD_FILE_RETENTION_DAYS = 7
DEFAULT_LOG_RETENTION_DAYS = 3
+UPLOAD_OWNER_TYPES = ('upload_image', 'upload_document', 'upload')
class MaintenanceService:
@@ -26,7 +30,10 @@ class MaintenanceService:
def __init__(self, ap: app.Application) -> None:
self.ap = ap
- async def cleanup_expired_files(self) -> dict[str, int]:
+ async def cleanup_expired_files(self, context: ExecutionContext) -> dict[str, int]:
+ if not isinstance(context, ExecutionContext):
+ raise WorkspaceRequiredError('Storage cleanup requires an ExecutionContext')
+ require_workspace_uuid(context)
cleanup_cfg = self.ap.instance_config.data.get('storage', {}).get('cleanup', {})
upload_retention_days = self._positive_int(
cleanup_cfg.get('uploaded_file_retention_days'),
@@ -40,11 +47,14 @@ class MaintenanceService:
)
return {
- 'uploaded_files': await self._cleanup_expired_uploaded_files(upload_retention_days),
- 'log_files': self._cleanup_expired_log_files(log_retention_days),
+ 'uploaded_files': await self._cleanup_expired_uploaded_files(context, upload_retention_days),
+ 'log_files': self._cleanup_expired_log_files(log_retention_days)
+ if await self._is_oss_singleton(context)
+ else 0,
}
- async def get_storage_analysis(self) -> dict[str, Any]:
+ async def get_storage_analysis(self, context: TenantContext) -> dict[str, Any]:
+ require_workspace_uuid(context)
cleanup_cfg = self.ap.instance_config.data.get('storage', {}).get('cleanup', {})
upload_retention_days = self._positive_int(
cleanup_cfg.get('uploaded_file_retention_days'),
@@ -62,15 +72,20 @@ class MaintenanceService:
database_path = (
Path(database_cfg.get('sqlite', {}).get('path', 'data/langbot.db')) if database_type == 'sqlite' else None
)
- roots: list[tuple[str, Path | None]] = [
- ('database', database_path),
- ('logs', Path('data/logs')),
- ('storage', Path('data/storage')),
- ('vector_store', Path('data/chroma')),
- ('plugins', Path('data/plugins')),
- ('mcp', Path('data/mcp')),
- ('temp', Path('data/temp')),
- ]
+ is_oss_singleton = await self._is_oss_singleton(context)
+ if is_oss_singleton:
+ roots: list[tuple[str, Path | None]] = [
+ ('database', database_path),
+ ('logs', Path('data/logs')),
+ ('storage', Path('data/storage')),
+ ('vector_store', Path('data/chroma')),
+ ('plugins', Path('data/plugins')),
+ ('mcp', Path('data/mcp')),
+ ('temp', Path('data/temp')),
+ ]
+ else:
+ scoped_storage_path = Path('data/storage') / self.ap.storage_mgr.scoped_prefix(context)
+ roots = [('storage', scoped_storage_path)]
sections = []
for key, path in roots:
@@ -84,10 +99,10 @@ class MaintenanceService:
}
)
- monitoring_counts = await self._monitoring_counts()
- binary_storage = await self._binary_storage_stats()
- upload_candidates = await self._expired_uploaded_candidates(upload_retention_days)
- log_candidates = self._expired_log_candidates(log_retention_days)
+ monitoring_counts = await self._monitoring_counts(context)
+ binary_storage = await self._binary_storage_stats(context)
+ upload_candidates = await self._expired_uploaded_candidates(context, upload_retention_days)
+ log_candidates = self._expired_log_candidates(log_retention_days) if is_oss_singleton else []
return {
'generated_at': datetime.datetime.now(datetime.timezone.utc).isoformat(),
@@ -105,14 +120,32 @@ class MaintenanceService:
'uploaded_files': upload_candidates,
'log_files': log_candidates,
},
- 'tasks': self.ap.task_mgr.get_stats() if self.ap.task_mgr else {},
+ 'tasks': self.ap.task_mgr.get_stats() if is_oss_singleton and self.ap.task_mgr else {},
}
- async def _cleanup_expired_uploaded_files(self, retention_days: int) -> int:
+ async def _is_oss_singleton(self, context: TenantContext) -> bool:
+ try:
+ await self.ap.workspace_service.get_local_execution_binding(
+ require_workspace_uuid(context),
+ expected_generation=getattr(context, 'placement_generation', None),
+ )
+ except Exception:
+ return False
+ return True
+
+ async def _cleanup_expired_uploaded_files(
+ self,
+ context: ExecutionContext,
+ retention_days: int,
+ ) -> int:
provider = self.ap.storage_mgr.storage_provider
provider_name = provider.__class__.__name__
if provider_name == 'LocalStorageProvider':
- candidates = self._expired_local_upload_candidates(retention_days, include_paths=True)
+ candidates = self._expired_local_upload_candidates(
+ context,
+ retention_days,
+ include_paths=True,
+ )
deleted = 0
for item in candidates:
try:
@@ -125,47 +158,65 @@ class MaintenanceService:
return deleted
if provider_name == 'S3StorageProvider':
- return await self._cleanup_expired_s3_uploaded_files(retention_days)
+ return await self._cleanup_expired_s3_uploaded_files(context, retention_days)
return 0
- async def _expired_uploaded_candidates(self, retention_days: int) -> list[dict[str, Any]]:
+ async def _expired_uploaded_candidates(
+ self,
+ context: TenantContext,
+ retention_days: int,
+ ) -> list[dict[str, Any]]:
provider_name = self.ap.storage_mgr.storage_provider.__class__.__name__
if provider_name == 'LocalStorageProvider':
- return self._expired_local_upload_candidates(retention_days)
+ return self._expired_local_upload_candidates(context, retention_days)
if provider_name == 'S3StorageProvider':
- return await self._expired_s3_upload_candidates(retention_days)
+ return await self._expired_s3_upload_candidates(context, retention_days)
return []
- async def _cleanup_expired_s3_uploaded_files(self, retention_days: int) -> int:
+ async def _cleanup_expired_s3_uploaded_files(
+ self,
+ context: ExecutionContext,
+ retention_days: int,
+ ) -> int:
provider = self.ap.storage_mgr.storage_provider
- candidates = await self._expired_s3_upload_candidates(retention_days)
+ candidates = await self._expired_s3_upload_candidates(context, retention_days)
deleted = 0
for item in candidates:
await provider.delete(item['key'])
deleted += 1
return deleted
- async def _expired_s3_upload_candidates(self, retention_days: int) -> list[dict[str, Any]]:
+ async def _expired_s3_upload_candidates(
+ self,
+ context: TenantContext,
+ retention_days: int,
+ ) -> list[dict[str, Any]]:
provider = self.ap.storage_mgr.storage_provider
cutoff = datetime.datetime.now(datetime.timezone.utc) - datetime.timedelta(days=retention_days)
candidates = []
paginator = provider.s3_client.get_paginator('list_objects_v2')
- for page in paginator.paginate(Bucket=provider.bucket_name):
- for obj in page.get('Contents', []):
- key = obj.get('Key', '')
- last_modified = obj.get('LastModified')
- if not self._is_uploaded_file_key(key):
- continue
- if last_modified and last_modified < cutoff:
- candidates.append(
- {
- 'key': key,
- 'size_bytes': obj.get('Size', 0),
- 'modified_at': last_modified.isoformat(),
- }
- )
+ seen_prefixes: set[str] = set()
+ for owner_type in UPLOAD_OWNER_TYPES:
+ prefix = self.ap.storage_mgr.scoped_prefix(context, owner_type=owner_type)
+ if prefix in seen_prefixes:
+ continue
+ seen_prefixes.add(prefix)
+ for page in paginator.paginate(Bucket=provider.bucket_name, Prefix=prefix):
+ for obj in page.get('Contents', []):
+ key = obj.get('Key', '')
+ last_modified = obj.get('LastModified')
+ if not self._is_uploaded_file_key(context, key):
+ continue
+ if last_modified and last_modified < cutoff:
+ candidates.append(
+ {
+ 'key': key,
+ 'size_bytes': obj.get('Size', 0),
+ 'modified_at': last_modified.isoformat(),
+ }
+ )
return candidates
@@ -182,28 +233,39 @@ class MaintenanceService:
return deleted
def _expired_local_upload_candidates(
- self, retention_days: int, include_paths: bool = False
+ self,
+ context: TenantContext,
+ retention_days: int,
+ include_paths: bool = False,
) -> list[dict[str, Any]]:
storage_root = Path('data/storage')
- if not storage_root.exists():
- return []
-
cutoff = datetime.datetime.now().timestamp() - retention_days * 86400
candidates = []
- for entry in storage_root.iterdir():
- if not entry.is_file() or not self._is_uploaded_file_key(entry.name):
+ seen_roots: set[Path] = set()
+ for owner_type in UPLOAD_OWNER_TYPES:
+ scoped_root = storage_root / self.ap.storage_mgr.scoped_prefix(context, owner_type=owner_type)
+ if scoped_root in seen_roots:
continue
- stat = entry.stat()
- if stat.st_mtime >= cutoff:
+ seen_roots.add(scoped_root)
+ if not scoped_root.exists():
continue
- item = {
- 'key': entry.name,
- 'size_bytes': stat.st_size,
- 'modified_at': datetime.datetime.fromtimestamp(stat.st_mtime, datetime.timezone.utc).isoformat(),
- }
- if include_paths:
- item['path'] = str(entry)
- candidates.append(item)
+ for entry in scoped_root.rglob('*'):
+ if not entry.is_file():
+ continue
+ stat = entry.stat()
+ if stat.st_mtime >= cutoff:
+ continue
+ item = {
+ 'key': entry.relative_to(storage_root).as_posix(),
+ 'size_bytes': stat.st_size,
+ 'modified_at': datetime.datetime.fromtimestamp(
+ stat.st_mtime,
+ datetime.timezone.utc,
+ ).isoformat(),
+ }
+ if include_paths:
+ item['path'] = str(entry)
+ candidates.append(item)
return candidates
def _expired_log_candidates(self, retention_days: int, include_paths: bool = False) -> list[dict[str, Any]]:
@@ -236,33 +298,51 @@ class MaintenanceService:
candidates.append(item)
return candidates
- def _is_uploaded_file_key(self, key: str) -> bool:
- return '/' not in key and not key.startswith('plugin_config_')
+ def _is_uploaded_file_key(self, context: TenantContext, key: str) -> bool:
+ return any(
+ key.startswith(self.ap.storage_mgr.scoped_prefix(context, owner_type=owner_type))
+ and self.ap.storage_mgr.is_scoped_object_key(key, expected_owner_type=owner_type)
+ for owner_type in UPLOAD_OWNER_TYPES
+ )
- async def _monitoring_counts(self) -> dict[str, int]:
+ async def _monitoring_counts(self, context: TenantContext) -> dict[str, int]:
+ workspace_uuid = require_workspace_uuid(context)
tables = {
- 'messages': persistence_monitoring.MonitoringMessage.id,
- 'llm_calls': persistence_monitoring.MonitoringLLMCall.id,
- 'tool_calls': persistence_monitoring.MonitoringToolCall.id,
- 'embedding_calls': persistence_monitoring.MonitoringEmbeddingCall.id,
- 'errors': persistence_monitoring.MonitoringError.id,
- 'sessions': persistence_monitoring.MonitoringSession.session_id,
- 'feedback': persistence_monitoring.MonitoringFeedback.id,
+ 'messages': (persistence_monitoring.MonitoringMessage, persistence_monitoring.MonitoringMessage.id),
+ 'llm_calls': (persistence_monitoring.MonitoringLLMCall, persistence_monitoring.MonitoringLLMCall.id),
+ 'tool_calls': (persistence_monitoring.MonitoringToolCall, persistence_monitoring.MonitoringToolCall.id),
+ 'embedding_calls': (
+ persistence_monitoring.MonitoringEmbeddingCall,
+ persistence_monitoring.MonitoringEmbeddingCall.id,
+ ),
+ 'errors': (persistence_monitoring.MonitoringError, persistence_monitoring.MonitoringError.id),
+ 'sessions': (
+ persistence_monitoring.MonitoringSession,
+ persistence_monitoring.MonitoringSession.session_id,
+ ),
+ 'feedback': (persistence_monitoring.MonitoringFeedback, persistence_monitoring.MonitoringFeedback.id),
}
counts: dict[str, int] = {}
- for key, column in tables.items():
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(sqlalchemy.func.count(column)))
+ for key, (model, column) in tables.items():
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(sqlalchemy.func.count(column)).where(model.workspace_uuid == workspace_uuid)
+ )
counts[key] = result.scalar() or 0
return counts
- async def _binary_storage_stats(self) -> dict[str, Any]:
+ async def _binary_storage_stats(self, context: TenantContext) -> dict[str, Any]:
+ workspace_uuid = require_workspace_uuid(context)
count_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(sqlalchemy.func.count(persistence_bstorage.BinaryStorage.unique_key))
+ sqlalchemy.select(sqlalchemy.func.count(persistence_bstorage.BinaryStorage.unique_key)).where(
+ persistence_bstorage.BinaryStorage.workspace_uuid == workspace_uuid
+ )
)
size_bytes = None
try:
size_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(sqlalchemy.func.sum(sqlalchemy.func.length(persistence_bstorage.BinaryStorage.value)))
+ sqlalchemy.select(
+ sqlalchemy.func.sum(sqlalchemy.func.length(persistence_bstorage.BinaryStorage.value))
+ ).where(persistence_bstorage.BinaryStorage.workspace_uuid == workspace_uuid)
)
size_bytes = size_result.scalar() or 0
except Exception as e:
diff --git a/src/langbot/pkg/api/http/service/mcp.py b/src/langbot/pkg/api/http/service/mcp.py
index 1dbceb6e5..9443de74c 100644
--- a/src/langbot/pkg/api/http/service/mcp.py
+++ b/src/langbot/pkg/api/http/service/mcp.py
@@ -1,198 +1,417 @@
from __future__ import annotations
-import sqlalchemy
-import uuid
import asyncio
+import copy
+import re
+import uuid
-from ....core import app
+import sqlalchemy
+
+from ....core import app, taskmgr
from ....entity.persistence import mcp as persistence_mcp
-from ....core import taskmgr
-from ....provider.tools.loaders.mcp import RuntimeMCPSession, MCPSessionStatus
+from ....entity.persistence import plugin as persistence_plugin
+from ....provider.tools.loaders.mcp import MCPSessionStatus, RuntimeMCPSession
+from ....workspace.errors import WorkspaceNotFoundError
+from ..context import ExecutionContext
+from .secrets import is_url_key, redact_url_secrets, restore_url_secret_placeholders
+from .tenant import TenantContext, require_workspace_uuid, scope_statement
+
+
+_SECRET_MASK = '***'
+_MISSING_SECRET = object()
+_SENSITIVE_CONFIG_NAMES = frozenset(
+ {
+ 'api_key',
+ 'apikey',
+ 'auth',
+ 'authorization',
+ 'cookie',
+ 'credentials',
+ 'database_url',
+ 'dsn',
+ 'key',
+ 'proxy_authorization',
+ 'set_cookie',
+ }
+)
+_SENSITIVE_CONFIG_TOKENS = frozenset(
+ {
+ 'credential',
+ 'credentials',
+ 'passwd',
+ 'password',
+ 'secret',
+ 'token',
+ }
+)
+_SENSITIVE_KEY_QUALIFIERS = frozenset(
+ {
+ 'access',
+ 'api',
+ 'auth',
+ 'bearer',
+ 'client',
+ 'debug',
+ 'encryption',
+ 'private',
+ 'signing',
+ }
+)
+
+
+def _normalize_config_key(key: object) -> str:
+ value = re.sub(r'([a-z0-9])([A-Z])', r'\1_\2', str(key or ''))
+ return re.sub(r'[^a-zA-Z0-9]+', '_', value).strip('_').lower()
+
+
+def _is_sensitive_config_key(key: object) -> bool:
+ normalized = _normalize_config_key(key)
+ if normalized in _SENSITIVE_CONFIG_NAMES:
+ return True
+ tokens = frozenset(token for token in normalized.split('_') if token)
+ if tokens & _SENSITIVE_CONFIG_TOKENS:
+ return True
+ return 'key' in tokens and bool(tokens & _SENSITIVE_KEY_QUALIFIERS)
+
+
+def _mask_secret_structure(value):
+ if isinstance(value, dict):
+ return {key: _mask_secret_structure(item) for key, item in value.items()}
+ if isinstance(value, list):
+ return [_mask_secret_structure(item) for item in value]
+ if isinstance(value, tuple):
+ return tuple(_mask_secret_structure(item) for item in value)
+ if value is None or value == '':
+ return value
+ return _SECRET_MASK
+
+
+def redact_mcp_secrets(value):
+ """Return a recursively redacted copy of MCP configuration data."""
+
+ if isinstance(value, dict):
+ return {
+ key: (
+ _mask_secret_structure(item)
+ if _is_sensitive_config_key(key)
+ else redact_url_secrets(item)
+ if is_url_key(key)
+ else redact_mcp_secrets(item)
+ )
+ for key, item in value.items()
+ }
+ if isinstance(value, list):
+ return [redact_mcp_secrets(item) for item in value]
+ if isinstance(value, tuple):
+ return tuple(redact_mcp_secrets(item) for item in value)
+ return value
+
+
+def restore_mcp_secret_placeholders(value, current_value=_MISSING_SECRET, *, sensitive: bool = False):
+ """Restore masked leaves from the current MCP config before a write."""
+
+ if sensitive and value == _SECRET_MASK:
+ if current_value is _MISSING_SECRET:
+ raise ValueError('Masked MCP secret has no existing value')
+ return copy.deepcopy(current_value)
+ if isinstance(value, dict):
+ current_mapping = current_value if isinstance(current_value, dict) else {}
+ return {
+ key: (
+ restore_url_secret_placeholders(
+ item,
+ current_mapping.get(key, _MISSING_SECRET),
+ )
+ if not sensitive and not _is_sensitive_config_key(key) and is_url_key(key)
+ else restore_mcp_secret_placeholders(
+ item,
+ current_mapping.get(key, _MISSING_SECRET),
+ sensitive=sensitive or _is_sensitive_config_key(key),
+ )
+ )
+ for key, item in value.items()
+ }
+ if isinstance(value, list):
+ current_items = current_value if isinstance(current_value, (list, tuple)) else ()
+ return [
+ restore_mcp_secret_placeholders(
+ item,
+ current_items[index] if index < len(current_items) else _MISSING_SECRET,
+ sensitive=sensitive,
+ )
+ for index, item in enumerate(value)
+ ]
+ if isinstance(value, tuple):
+ current_items = current_value if isinstance(current_value, (list, tuple)) else ()
+ return tuple(
+ restore_mcp_secret_placeholders(
+ item,
+ current_items[index] if index < len(current_items) else _MISSING_SECRET,
+ sensitive=sensitive,
+ )
+ for index, item in enumerate(value)
+ )
+ return value
class MCPService:
+ """Workspace-scoped MCP configuration and runtime facade."""
+
ap: app.Application
def __init__(self, ap: app.Application) -> None:
self.ap = ap
- async def get_runtime_info(self, server_name: str) -> dict | None:
- session = self.ap.tool_mgr.mcp_tool_loader.get_session(server_name)
- if session:
- return session.get_runtime_info_dict()
- return None
+ async def _execution_context(self, context: TenantContext) -> ExecutionContext:
+ workspace_uuid = require_workspace_uuid(context)
+ instance_uuid = str(getattr(context, 'instance_uuid', '') or '').strip()
+ generation = getattr(context, 'placement_generation', None)
+ if not instance_uuid or not isinstance(generation, int) or isinstance(generation, bool) or generation <= 0:
+ raise ValueError('MCP operations require an explicit fenced execution context')
+ binding = await self.ap.workspace_service.get_execution_binding(
+ workspace_uuid,
+ expected_generation=generation,
+ )
+ if binding.instance_uuid != instance_uuid:
+ raise ValueError('MCP execution context belongs to another LangBot instance')
+ return ExecutionContext(
+ instance_uuid=instance_uuid,
+ workspace_uuid=workspace_uuid,
+ placement_generation=generation,
+ bot_uuid=getattr(context, 'bot_uuid', None),
+ pipeline_uuid=getattr(context, 'pipeline_uuid', None),
+ query_uuid=getattr(context, 'query_uuid', None),
+ )
- async def get_mcp_servers(self, contain_runtime_info: bool = False) -> list[dict]:
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_mcp.MCPServer))
+ async def get_runtime_info(self, context: TenantContext, server_name: str) -> dict | None:
+ execution_context = await self._execution_context(context)
+ session = self.ap.tool_mgr.mcp_tool_loader.get_session(execution_context, server_name)
+ return session.get_runtime_info_dict() if session else None
- servers = result.all()
+ async def get_mcp_servers(self, context: TenantContext, contain_runtime_info: bool = False) -> list[dict]:
+ execution_context = await self._execution_context(context)
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(sqlalchemy.select(persistence_mcp.MCPServer), persistence_mcp.MCPServer, context)
+ )
serialized_servers = [
- self.ap.persistence_mgr.serialize_model(persistence_mcp.MCPServer, server) for server in servers
+ redact_mcp_secrets(self.ap.persistence_mgr.serialize_model(persistence_mcp.MCPServer, server))
+ for server in result.all()
]
if contain_runtime_info:
for server in serialized_servers:
- runtime_info = await self.get_runtime_info(server['name'])
-
- server['runtime_info'] = runtime_info if runtime_info else None
-
+ session = self.ap.tool_mgr.mcp_tool_loader.get_session(execution_context, server['name'])
+ server['runtime_info'] = session.get_runtime_info_dict() if session else None
return serialized_servers
- async def create_mcp_server(self, server_data: dict) -> str:
- # Check limitation (extensions = MCP servers + plugins)
+ async def create_mcp_server(self, context: TenantContext, server_data: dict) -> str:
+ execution_context = await self._execution_context(context)
+ workspace_uuid = execution_context.workspace_uuid
+
limitation = self.ap.instance_config.data.get('system', {}).get('limitation', {})
max_extensions = limitation.get('max_extensions', -1)
if max_extensions >= 0:
- existing_mcp_servers = await self.get_mcp_servers()
- plugins = await self.ap.plugin_connector.list_plugins()
- total_extensions = len(existing_mcp_servers) + len(plugins)
- if total_extensions >= max_extensions:
+ mcp_count_result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(sqlalchemy.func.count(persistence_mcp.MCPServer.uuid)).where(
+ persistence_mcp.MCPServer.workspace_uuid == workspace_uuid
+ )
+ )
+ plugin_count_result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(sqlalchemy.func.count())
+ .select_from(persistence_plugin.PluginSetting)
+ .where(persistence_plugin.PluginSetting.workspace_uuid == workspace_uuid)
+ )
+ if (mcp_count_result.scalar() or 0) + (plugin_count_result.scalar() or 0) >= max_extensions:
raise ValueError(f'Maximum number of extensions ({max_extensions}) reached')
- server_name = str(server_data.get('name') or '').strip()
+ payload = dict(server_data)
+ payload.pop('workspace_uuid', None)
+ server_name = str(payload.get('name') or '').strip()
if not server_name:
raise ValueError('MCP server name is required')
- server_data['name'] = server_name
+ payload['name'] = server_name
+ payload['workspace_uuid'] = workspace_uuid
+ payload['uuid'] = str(uuid.uuid4())
existing_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.name == server_name)
+ sqlalchemy.select(persistence_mcp.MCPServer).where(
+ persistence_mcp.MCPServer.workspace_uuid == workspace_uuid,
+ persistence_mcp.MCPServer.name == server_name,
+ )
)
if existing_result.first() is not None:
raise ValueError(f'MCP server already exists: {server_name}')
- server_data['uuid'] = str(uuid.uuid4())
- await self.ap.persistence_mgr.execute_async(sqlalchemy.insert(persistence_mcp.MCPServer).values(server_data))
+ await self.ap.persistence_mgr.execute_async(sqlalchemy.insert(persistence_mcp.MCPServer).values(payload))
+ created = await self._get_mcp_server_by_uuid_raw(execution_context, payload['uuid'])
+ if created and self.ap.tool_mgr.mcp_tool_loader:
+ task = asyncio.create_task(self.ap.tool_mgr.mcp_tool_loader.host_mcp_server(execution_context, created))
+ self.ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks.append(task)
+ return payload['uuid']
+ async def get_mcp_server_by_uuid(self, context: TenantContext, server_uuid: str) -> dict | None:
+ execution_context = await self._execution_context(context)
+ server_data = await self._get_mcp_server_by_uuid_raw(execution_context, server_uuid)
+ return redact_mcp_secrets(server_data) if server_data is not None else None
+
+ async def _get_mcp_server_by_uuid_raw(
+ self,
+ execution_context: ExecutionContext,
+ server_uuid: str,
+ ) -> dict | None:
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.uuid == server_data['uuid'])
+ scope_statement(
+ sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.uuid == server_uuid),
+ persistence_mcp.MCPServer,
+ execution_context,
+ )
)
- server_entity = result.first()
- if server_entity:
- server_config = self.ap.persistence_mgr.serialize_model(persistence_mcp.MCPServer, server_entity)
- if self.ap.tool_mgr.mcp_tool_loader:
- task = asyncio.create_task(self.ap.tool_mgr.mcp_tool_loader.host_mcp_server(server_config))
- self.ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks.append(task)
+ server = result.first()
+ return self.ap.persistence_mgr.serialize_model(persistence_mcp.MCPServer, server) if server else None
- return server_data['uuid']
+ async def get_mcp_server_by_name(self, context: TenantContext, server_name: str) -> dict | None:
+ execution_context = await self._execution_context(context)
+ server_data = await self._get_mcp_server_by_name_raw(execution_context, server_name)
+ if server_data is None:
+ return None
+ session = self.ap.tool_mgr.mcp_tool_loader.get_session(execution_context, server_name)
+ response_data = {
+ **server_data,
+ 'runtime_info': session.get_runtime_info_dict() if session else None,
+ }
+ return redact_mcp_secrets(response_data)
- async def get_mcp_server_by_name(self, server_name: str) -> dict | None:
+ async def _get_mcp_server_by_name_raw(
+ self,
+ execution_context: ExecutionContext,
+ server_name: str,
+ ) -> dict | None:
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.name == server_name)
+ scope_statement(
+ sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.name == server_name),
+ persistence_mcp.MCPServer,
+ execution_context,
+ )
)
server = result.first()
if server is None:
return None
+ return self.ap.persistence_mgr.serialize_model(persistence_mcp.MCPServer, server)
- runtime_info = await self.get_runtime_info(server.name)
- server_data = self.ap.persistence_mgr.serialize_model(persistence_mcp.MCPServer, server)
- server_data['runtime_info'] = runtime_info if runtime_info else None
- return server_data
+ async def update_mcp_server(self, context: TenantContext, server_uuid: str, server_data: dict) -> None:
+ execution_context = await self._execution_context(context)
+ old_server = await self._get_mcp_server_by_uuid_raw(execution_context, server_uuid)
+ if old_server is None:
+ raise WorkspaceNotFoundError('MCP server not found')
+
+ payload = dict(server_data)
+ payload.pop('uuid', None)
+ payload.pop('workspace_uuid', None)
+ payload = restore_mcp_secret_placeholders(payload, old_server)
+ if 'name' in payload:
+ payload['name'] = str(payload['name'] or '').strip()
+ if not payload['name']:
+ raise ValueError('MCP server name is required')
+ duplicate = await self._get_mcp_server_by_name_raw(execution_context, payload['name'])
+ if duplicate is not None and duplicate['uuid'] != server_uuid:
+ raise ValueError(f'MCP server already exists: {payload["name"]}')
- async def update_mcp_server(self, server_uuid: str, server_data: dict) -> None:
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.uuid == server_uuid)
+ scope_statement(
+ sqlalchemy.update(persistence_mcp.MCPServer)
+ .where(persistence_mcp.MCPServer.uuid == server_uuid)
+ .values(payload),
+ persistence_mcp.MCPServer,
+ execution_context,
+ )
)
- old_server = result.first()
- old_server_name = old_server.name if old_server else None
- old_enable = old_server.enable if old_server else False
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('MCP server not found')
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(persistence_mcp.MCPServer)
- .where(persistence_mcp.MCPServer.uuid == server_uuid)
- .values(server_data)
- )
+ loader = self.ap.tool_mgr.mcp_tool_loader
+ if loader is None:
+ return
+ old_name = old_server['name']
+ old_enable = bool(old_server['enable'])
+ updated = await self._get_mcp_server_by_uuid_raw(execution_context, server_uuid)
+ if updated is None:
+ raise WorkspaceNotFoundError('MCP server not found')
+ new_enable = bool(updated['enable'])
+ if old_enable and loader.has_session(execution_context, old_name):
+ await loader.remove_mcp_server(execution_context, old_name)
+ if new_enable:
+ task = asyncio.create_task(loader.host_mcp_server(execution_context, updated))
+ loader._hosted_mcp_tasks.append(task)
- if self.ap.tool_mgr.mcp_tool_loader:
- new_enable = server_data.get('enable', False)
-
- need_remove = old_server_name and old_server_name in self.ap.tool_mgr.mcp_tool_loader.sessions
-
- if old_enable and not new_enable:
- if need_remove:
- await self.ap.tool_mgr.mcp_tool_loader.remove_mcp_server(old_server_name)
-
- elif not old_enable and new_enable:
- result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.uuid == server_uuid)
- )
- updated_server = result.first()
- if updated_server:
- server_config = self.ap.persistence_mgr.serialize_model(persistence_mcp.MCPServer, updated_server)
- task = asyncio.create_task(self.ap.tool_mgr.mcp_tool_loader.host_mcp_server(server_config))
- self.ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks.append(task)
-
- elif old_enable and new_enable:
- if need_remove:
- await self.ap.tool_mgr.mcp_tool_loader.remove_mcp_server(old_server_name)
- result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.uuid == server_uuid)
- )
- updated_server = result.first()
- if updated_server:
- server_config = self.ap.persistence_mgr.serialize_model(persistence_mcp.MCPServer, updated_server)
- task = asyncio.create_task(self.ap.tool_mgr.mcp_tool_loader.host_mcp_server(server_config))
- self.ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks.append(task)
-
- async def delete_mcp_server(self, server_uuid: str) -> None:
+ async def delete_mcp_server(self, context: TenantContext, server_uuid: str) -> None:
+ execution_context = await self._execution_context(context)
+ server = await self._get_mcp_server_by_uuid_raw(execution_context, server_uuid)
+ if server is None:
+ raise WorkspaceNotFoundError('MCP server not found')
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.uuid == server_uuid)
+ scope_statement(
+ sqlalchemy.delete(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.uuid == server_uuid),
+ persistence_mcp.MCPServer,
+ execution_context,
+ )
)
- server = result.first()
- server_name = server.name if server else None
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('MCP server not found')
+ loader = self.ap.tool_mgr.mcp_tool_loader
+ if loader and loader.has_session(execution_context, server['name']):
+ await loader.remove_mcp_server(execution_context, server['name'])
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.uuid == server_uuid)
- )
+ async def _require_server(self, context: TenantContext, server_name: str) -> tuple[ExecutionContext, dict]:
+ execution_context = await self._execution_context(context)
+ server = await self._get_mcp_server_by_name_raw(execution_context, server_name)
+ if server is None:
+ raise WorkspaceNotFoundError('MCP server not found')
+ return execution_context, server
- if server_name and self.ap.tool_mgr.mcp_tool_loader:
- if server_name in self.ap.tool_mgr.mcp_tool_loader.sessions:
- await self.ap.tool_mgr.mcp_tool_loader.remove_mcp_server(server_name)
+ async def get_mcp_server_resources(self, context: TenantContext, server_name: str) -> list[dict]:
+ execution_context, _ = await self._require_server(context, server_name)
+ return await self.ap.tool_mgr.mcp_tool_loader.get_resources(execution_context, server_name)
- async def get_mcp_server_resources(self, server_name: str) -> list[dict]:
- """Get resources from a specific MCP server."""
- return await self.ap.tool_mgr.mcp_tool_loader.get_resources(server_name)
-
- async def get_mcp_server_resource_templates(self, server_name: str) -> list[dict]:
- """Get resource templates from a specific MCP server."""
- return await self.ap.tool_mgr.mcp_tool_loader.get_resource_templates(server_name)
+ async def get_mcp_server_resource_templates(self, context: TenantContext, server_name: str) -> list[dict]:
+ execution_context, _ = await self._require_server(context, server_name)
+ return await self.ap.tool_mgr.mcp_tool_loader.get_resource_templates(execution_context, server_name)
async def read_mcp_server_resource_envelope(
self,
+ context: TenantContext,
server_name: str,
uri: str,
*,
max_bytes: int | None = None,
include_blob: bool = False,
) -> dict:
- """Read a resource from a specific MCP server with metadata."""
+ execution_context, _ = await self._require_server(context, server_name)
kwargs = {'include_blob': include_blob, 'source': 'ui_preview'}
if max_bytes is not None:
kwargs['max_bytes'] = max_bytes
- return await self.ap.tool_mgr.mcp_tool_loader.read_resource_envelope(server_name, uri, **kwargs)
+ return await self.ap.tool_mgr.mcp_tool_loader.read_resource_envelope(
+ execution_context,
+ server_name,
+ uri,
+ **kwargs,
+ )
- async def read_mcp_server_resource(self, server_name: str, uri: str) -> list[dict]:
- """Read a resource from a specific MCP server."""
- return await self.ap.tool_mgr.mcp_tool_loader.read_resource(server_name, uri)
-
- async def test_mcp_server(self, server_name: str, server_data: dict) -> int:
- """测试 MCP 服务器连接并返回任务 ID"""
+ async def read_mcp_server_resource(self, context: TenantContext, server_name: str, uri: str) -> list[dict]:
+ execution_context, _ = await self._require_server(context, server_name)
+ return await self.ap.tool_mgr.mcp_tool_loader.read_resource(execution_context, server_name, uri)
+ async def test_mcp_server(self, context: TenantContext, server_name: str, server_data: dict) -> int:
+ execution_context = await self._execution_context(context)
runtime_mcp_session: RuntimeMCPSession | None = None
-
ctx = taskmgr.TaskContext.new()
if server_name != '_':
- runtime_mcp_session = self.ap.tool_mgr.mcp_tool_loader.get_session(server_name)
+ await self._require_server(execution_context, server_name)
+ runtime_mcp_session = self.ap.tool_mgr.mcp_tool_loader.get_session(execution_context, server_name)
if runtime_mcp_session is None:
- raise ValueError(f'Server not found: {server_name}')
-
+ raise WorkspaceNotFoundError('MCP server not found')
persisted_session = runtime_mcp_session
async def _refresh_and_report() -> None:
- # Testing a persisted server should REUSE its live shared-session
- # process, not rebuild it. Try a lightweight refresh (a real
- # list_tools probe over the existing connection) first; only fall
- # back to a full start() when the session has no live connection
- # to probe (never connected, or the process is actually gone).
needs_start = persisted_session.status == MCPSessionStatus.ERROR or persisted_session.session is None
if needs_start:
await persisted_session.start()
@@ -200,30 +419,23 @@ class MCPService:
try:
await persisted_session.refresh()
except Exception:
- # The live connection was stale/dropped: reconnect once
- # (reusing the live managed process where possible) and
- # re-probe, instead of reporting a false failure.
await persisted_session.start()
- # Surface the discovered tools so the config page can render them
- # even for an already-hosted server.
ctx.metadata['runtime_info'] = persisted_session.get_runtime_info_dict()
coroutine = _refresh_and_report()
else:
- runtime_mcp_session = await self.ap.tool_mgr.mcp_tool_loader.load_mcp_server(server_config=server_data)
-
- # A transient test owns an isolated Box session. Always tear it down
- # after the test completes (success or failure) so it does not leak.
+ payload = dict(server_data)
+ payload.pop('workspace_uuid', None)
+ payload['workspace_uuid'] = execution_context.workspace_uuid
+ runtime_mcp_session = await self.ap.tool_mgr.mcp_tool_loader.load_mcp_server(
+ execution_context,
+ payload,
+ )
test_session = runtime_mcp_session
async def _run_and_cleanup() -> None:
try:
await test_session.start()
- # Capture the runtime info (status + discovered tools) BEFORE
- # shutting the transient session down. The create/edit config
- # page has no persisted server to reload from, so without this
- # a successful test could only show "no tools found". The
- # frontend reads ctx.metadata.runtime_info to render the tools.
ctx.metadata['runtime_info'] = test_session.get_runtime_info_dict()
finally:
try:
@@ -239,24 +451,27 @@ class MCPService:
wrapper = self.ap.task_mgr.create_user_task(
coroutine,
kind='mcp-operation',
- name=f'mcp-test-{server_name}',
+ name=f'mcp-test-{execution_context.workspace_uuid}-{server_name}',
label=f'Testing MCP server {server_name}',
context=ctx,
+ instance_uuid=execution_context.instance_uuid,
+ workspace_uuid=execution_context.workspace_uuid,
+ placement_generation=execution_context.placement_generation,
)
return wrapper.id
- async def get_mcp_server_logs(self, server_name: str, limit: int = 200, level: str | None = None) -> list[dict]:
- """Get recent log lines captured from the MCP server's stderr."""
- session = self.ap.tool_mgr.mcp_tool_loader.get_session(server_name)
+ async def get_mcp_server_logs(
+ self,
+ context: TenantContext,
+ server_name: str,
+ limit: int = 200,
+ level: str | None = None,
+ ) -> list[dict]:
+ execution_context, _ = await self._require_server(context, server_name)
+ session = self.ap.tool_mgr.mcp_tool_loader.get_session(execution_context, server_name)
if not session:
return []
-
- # Get logs from the session's buffer
logs = list(session._log_buffer)
-
- # Filter by level if specified
if level:
logs = [log for log in logs if log.get('level') == level]
-
- # Return the most recent 'limit' logs
return logs[-limit:]
diff --git a/src/langbot/pkg/api/http/service/model.py b/src/langbot/pkg/api/http/service/model.py
index 87298c084..88801c6b8 100644
--- a/src/langbot/pkg/api/http/service/model.py
+++ b/src/langbot/pkg/api/http/service/model.py
@@ -9,6 +9,9 @@ from ....core import app
from ....entity.persistence import model as persistence_model
from ....entity.persistence import pipeline as persistence_pipeline
from ....provider.modelmgr import requester as model_requester
+from ....workspace.errors import WorkspaceNotFoundError
+from .secrets import mask_secret_value, redact_secrets, restore_secret_placeholders
+from .tenant import TenantContext, require_workspace_uuid, scope_statement
def _parse_provider_api_keys(provider_dict: dict) -> dict:
@@ -34,7 +37,29 @@ def _runtime_model_data(model_uuid: str, model_data: dict) -> dict:
return {**model_data, 'uuid': model_uuid}
-async def _validate_provider_supports(ap: app.Application, provider_uuid: str, model_type: str) -> None:
+def _redact_model_secrets(model_data: dict) -> dict:
+ """Return a copy with model args and embedded provider credentials masked."""
+
+ redacted = model_data.copy()
+ if 'extra_args' in redacted:
+ redacted['extra_args'] = redact_secrets(redacted['extra_args'])
+ if isinstance(redacted.get('provider'), dict):
+ provider = redacted['provider'].copy()
+ # ModelProvider never contains another provider. Dropping this key also
+ # makes the serializer robust to a reused/self-referential test double.
+ provider.pop('provider', None)
+ if 'api_keys' in provider:
+ provider['api_keys'] = mask_secret_value(provider['api_keys'])
+ redacted['provider'] = provider
+ return redacted
+
+
+async def _validate_provider_supports(
+ ap: app.Application,
+ context: TenantContext,
+ provider_uuid: str,
+ model_type: str,
+) -> None:
"""Validate that the provider's requester declares support for ``model_type``.
``model_type`` is one of the manifest ``support_type`` values:
@@ -47,11 +72,12 @@ async def _validate_provider_supports(ap: app.Application, provider_uuid: str, m
if model_mgr is None:
return
- provider_dict = getattr(model_mgr, 'provider_dict', None)
- if not provider_dict:
+ get_provider = getattr(model_mgr, 'get_provider_by_uuid', None)
+ if not callable(get_provider):
return
- runtime_provider = provider_dict.get(provider_uuid)
- if runtime_provider is None:
+ try:
+ runtime_provider = await get_provider(context, provider_uuid)
+ except ValueError:
return
requester_name = getattr(getattr(runtime_provider, 'provider_entity', None), 'requester', None)
@@ -74,20 +100,48 @@ async def _validate_provider_supports(ap: app.Application, provider_uuid: str, m
raise ValueError(f'Provider requester "{requester_name}" does not support {model_type} models')
+async def _require_workspace_provider(
+ ap: app.Application,
+ context: TenantContext,
+ provider_uuid: str,
+) -> dict:
+ """Require the referenced provider to belong to the active Workspace."""
+
+ provider = await ap.provider_service.get_provider(context, provider_uuid)
+ if provider is None:
+ raise WorkspaceNotFoundError('Provider not found')
+ return provider
+
+
+async def _require_runtime_provider(
+ ap: app.Application,
+ context: TenantContext,
+ provider_uuid: str,
+) -> model_requester.RuntimeProvider:
+ try:
+ return await ap.model_mgr.get_provider_by_uuid(context, provider_uuid)
+ except ValueError as exc:
+ raise Exception('provider not found') from exc
+
+
class LLMModelsService:
ap: app.Application
def __init__(self, ap: app.Application) -> None:
self.ap = ap
- async def get_llm_models(self, include_secret: bool = True) -> list[dict]:
+ async def get_llm_models(self, context: TenantContext, include_secret: bool = False) -> list[dict]:
"""Get all LLM models with provider info"""
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_model.LLMModel))
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(sqlalchemy.select(persistence_model.LLMModel), persistence_model.LLMModel, context)
+ )
models = result.all()
# Get all providers for lookup
providers_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.ModelProvider)
+ scope_statement(
+ sqlalchemy.select(persistence_model.ModelProvider), persistence_model.ModelProvider, context
+ )
)
providers = {p.uuid: p for p in providers_result.all()}
@@ -98,29 +152,50 @@ class LLMModelsService:
if provider:
provider_dict = self.ap.persistence_mgr.serialize_model(persistence_model.ModelProvider, provider)
provider_dict = _parse_provider_api_keys(provider_dict)
- if not include_secret:
- provider_dict['api_keys'] = ['***'] * len(provider_dict.get('api_keys', []))
model_dict['provider'] = provider_dict
+ if not include_secret:
+ model_dict = _redact_model_secrets(model_dict)
models_list.append(model_dict)
return models_list
- async def get_llm_models_by_provider(self, provider_uuid: str) -> list[dict]:
+ async def get_llm_models_by_provider(
+ self,
+ context: TenantContext,
+ provider_uuid: str,
+ *,
+ include_secret: bool = False,
+ ) -> list[dict]:
"""Get LLM models by provider UUID"""
+ await _require_workspace_provider(self.ap, context, provider_uuid)
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.LLMModel).where(
- persistence_model.LLMModel.provider_uuid == provider_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_model.LLMModel).where(
+ persistence_model.LLMModel.provider_uuid == provider_uuid
+ ),
+ persistence_model.LLMModel,
+ context,
)
)
models = result.all()
- return [self.ap.persistence_mgr.serialize_model(persistence_model.LLMModel, m) for m in models]
+ serialized = [self.ap.persistence_mgr.serialize_model(persistence_model.LLMModel, m) for m in models]
+ return serialized if include_secret else [_redact_model_secrets(model) for model in serialized]
async def create_llm_model(
- self, model_data: dict, preserve_uuid: bool = False, auto_set_to_default_pipeline: bool = True
+ self,
+ context: TenantContext,
+ model_data: dict,
+ preserve_uuid: bool = False,
+ auto_set_to_default_pipeline: bool = True,
) -> str:
"""Create a new LLM model"""
+ workspace_uuid = require_workspace_uuid(context)
+ model_data = model_data.copy()
if not preserve_uuid:
model_data['uuid'] = str(uuid.uuid4())
+ model_data['workspace_uuid'] = workspace_uuid
+ if 'extra_args' in model_data:
+ model_data['extra_args'] = restore_secret_placeholders(model_data['extra_args'])
# Handle provider creation if needed
if 'provider' in model_data:
@@ -130,31 +205,35 @@ class LLMModelsService:
else:
# Create new provider
provider_uuid = await self.ap.provider_service.find_or_create_provider(
+ context,
requester=provider_data.get('requester', ''),
base_url=provider_data.get('base_url', ''),
api_keys=provider_data.get('api_keys', []),
)
model_data['provider_uuid'] = provider_uuid
- await _validate_provider_supports(self.ap, model_data['provider_uuid'], 'llm')
+ await _require_workspace_provider(self.ap, context, model_data['provider_uuid'])
+ await _validate_provider_supports(self.ap, context, model_data['provider_uuid'], 'llm')
await self.ap.persistence_mgr.execute_async(sqlalchemy.insert(persistence_model.LLMModel).values(**model_data))
- runtime_provider = self.ap.model_mgr.provider_dict.get(model_data['provider_uuid'])
- if runtime_provider is None:
- raise Exception('provider not found')
-
+ runtime_provider = await _require_runtime_provider(self.ap, context, model_data['provider_uuid'])
runtime_llm_model = await self.ap.model_mgr.load_llm_model_with_provider(
+ context,
persistence_model.LLMModel(**model_data),
runtime_provider,
)
- self.ap.model_mgr.llm_models.append(runtime_llm_model)
+ await self.ap.model_mgr.cache_llm_model(context, runtime_llm_model)
if auto_set_to_default_pipeline:
# set the default pipeline model to this model
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
- persistence_pipeline.LegacyPipeline.is_default == True
+ scope_statement(
+ sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
+ persistence_pipeline.LegacyPipeline.is_default == True
+ ),
+ persistence_pipeline.LegacyPipeline,
+ workspace_uuid,
)
)
pipeline = result.first()
@@ -167,14 +246,23 @@ class LLMModelsService:
'fallbacks': [],
}
pipeline_data = {'config': pipeline_config}
- await self.ap.pipeline_service.update_pipeline(pipeline.uuid, pipeline_data)
+ await self.ap.pipeline_service.update_pipeline(context, pipeline.uuid, pipeline_data)
return model_data['uuid']
- async def get_llm_model(self, model_uuid: str) -> dict | None:
+ async def get_llm_model(
+ self,
+ context: TenantContext,
+ model_uuid: str,
+ include_secret: bool = False,
+ ) -> dict | None:
"""Get a single LLM model with provider info"""
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.LLMModel).where(persistence_model.LLMModel.uuid == model_uuid)
+ scope_statement(
+ sqlalchemy.select(persistence_model.LLMModel).where(persistence_model.LLMModel.uuid == model_uuid),
+ persistence_model.LLMModel,
+ context,
+ )
)
model = result.first()
if model is None:
@@ -184,21 +272,38 @@ class LLMModelsService:
# Get provider
provider_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.ModelProvider).where(
- persistence_model.ModelProvider.uuid == model.provider_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_model.ModelProvider).where(
+ persistence_model.ModelProvider.uuid == model.provider_uuid
+ ),
+ persistence_model.ModelProvider,
+ context,
)
)
provider = provider_result.first()
if provider:
provider_dict = self.ap.persistence_mgr.serialize_model(persistence_model.ModelProvider, provider)
- model_dict['provider'] = _parse_provider_api_keys(provider_dict)
+ provider_dict = _parse_provider_api_keys(provider_dict)
+ model_dict['provider'] = provider_dict
+
+ if not include_secret:
+ model_dict = _redact_model_secrets(model_dict)
return model_dict
- async def update_llm_model(self, model_uuid: str, model_data: dict) -> None:
+ async def update_llm_model(self, context: TenantContext, model_uuid: str, model_data: dict) -> None:
"""Update an existing LLM model"""
- if 'uuid' in model_data:
- del model_data['uuid']
+ existing_model = await self.get_llm_model(context, model_uuid, include_secret=True)
+ if existing_model is None:
+ raise WorkspaceNotFoundError('Model not found')
+ model_data = model_data.copy()
+ model_data.pop('uuid', None)
+ model_data.pop('workspace_uuid', None)
+ if 'extra_args' in model_data:
+ model_data['extra_args'] = restore_secret_placeholders(
+ model_data['extra_args'],
+ existing_model.get('extra_args', {}),
+ )
# Handle provider update if needed
if 'provider' in model_data:
@@ -207,50 +312,71 @@ class LLMModelsService:
model_data['provider_uuid'] = provider_data['uuid']
else:
provider_uuid = await self.ap.provider_service.find_or_create_provider(
+ context,
requester=provider_data.get('requester', ''),
base_url=provider_data.get('base_url', ''),
api_keys=provider_data.get('api_keys', []),
)
model_data['provider_uuid'] = provider_uuid
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(persistence_model.LLMModel)
- .where(persistence_model.LLMModel.uuid == model_uuid)
- .values(**model_data)
+ provider_uuid = model_data.get('provider_uuid', existing_model['provider_uuid'])
+ await _require_workspace_provider(self.ap, context, provider_uuid)
+ await _validate_provider_supports(self.ap, context, provider_uuid, 'llm')
+
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.update(persistence_model.LLMModel)
+ .where(persistence_model.LLMModel.uuid == model_uuid)
+ .values(**model_data),
+ persistence_model.LLMModel,
+ context,
+ )
)
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Model not found')
- await self.ap.model_mgr.remove_llm_model(model_uuid)
-
- runtime_provider = self.ap.model_mgr.provider_dict.get(model_data['provider_uuid'])
- if runtime_provider is None:
- raise Exception('provider not found')
-
+ await self.ap.model_mgr.remove_llm_model(context, model_uuid)
+ runtime_provider = await _require_runtime_provider(self.ap, context, provider_uuid)
runtime_llm_model = await self.ap.model_mgr.load_llm_model_with_provider(
- persistence_model.LLMModel(**_runtime_model_data(model_uuid, model_data)),
+ context,
+ persistence_model.LLMModel(
+ **_runtime_model_data(
+ model_uuid,
+ {
+ key: value
+ for key, value in {**existing_model, **model_data, 'provider_uuid': provider_uuid}.items()
+ if key not in {'provider', 'created_at', 'updated_at'}
+ },
+ )
+ ),
runtime_provider,
)
- self.ap.model_mgr.llm_models.append(runtime_llm_model)
+ await self.ap.model_mgr.cache_llm_model(context, runtime_llm_model)
- async def delete_llm_model(self, model_uuid: str) -> None:
+ async def delete_llm_model(self, context: TenantContext, model_uuid: str) -> None:
"""Delete an LLM model"""
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_model.LLMModel).where(persistence_model.LLMModel.uuid == model_uuid)
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.delete(persistence_model.LLMModel).where(persistence_model.LLMModel.uuid == model_uuid),
+ persistence_model.LLMModel,
+ context,
+ )
)
- await self.ap.model_mgr.remove_llm_model(model_uuid)
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Model not found')
+ await self.ap.model_mgr.remove_llm_model(context, model_uuid)
- async def test_llm_model(self, model_uuid: str, model_data: dict) -> None:
+ async def test_llm_model(self, context: TenantContext, model_uuid: str, model_data: dict) -> None:
"""Test an LLM model"""
+ require_workspace_uuid(context)
runtime_llm_model: model_requester.RuntimeLLMModel | None = None
if model_uuid != '_':
- for model in self.ap.model_mgr.llm_models:
- if model.model_entity.uuid == model_uuid:
- runtime_llm_model = model
- break
- if runtime_llm_model is None:
- raise Exception('model not found')
+ if await self.get_llm_model(context, model_uuid) is None:
+ raise WorkspaceNotFoundError('Model not found')
+ runtime_llm_model = await self.ap.model_mgr.get_model_by_uuid(context, model_uuid)
else:
- runtime_llm_model = await self.ap.model_mgr.init_temporary_runtime_llm_model(model_data)
+ runtime_llm_model = await self.ap.model_mgr.init_temporary_runtime_llm_model(context, model_data)
extra_args = model_data.get('extra_args', {})
await runtime_llm_model.provider.invoke_llm(
@@ -259,6 +385,7 @@ class LLMModelsService:
messages=[provider_message.Message(role='user', content='Hello, world! Please just reply a "Hello".')],
funcs=[],
extra_args=extra_args,
+ execution_context=runtime_llm_model.execution_context,
)
@@ -268,13 +395,19 @@ class EmbeddingModelsService:
def __init__(self, ap: app.Application) -> None:
self.ap = ap
- async def get_embedding_models(self) -> list[dict]:
+ async def get_embedding_models(self, context: TenantContext, include_secret: bool = False) -> list[dict]:
"""Get all embedding models with provider info"""
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_model.EmbeddingModel))
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.select(persistence_model.EmbeddingModel), persistence_model.EmbeddingModel, context
+ )
+ )
models = result.all()
providers_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.ModelProvider)
+ scope_statement(
+ sqlalchemy.select(persistence_model.ModelProvider), persistence_model.ModelProvider, context
+ )
)
providers = {p.uuid: p for p in providers_result.all()}
@@ -284,25 +417,46 @@ class EmbeddingModelsService:
provider = providers.get(model.provider_uuid)
if provider:
provider_dict = self.ap.persistence_mgr.serialize_model(persistence_model.ModelProvider, provider)
- model_dict['provider'] = _parse_provider_api_keys(provider_dict)
+ provider_dict = _parse_provider_api_keys(provider_dict)
+ model_dict['provider'] = provider_dict
+ if not include_secret:
+ model_dict = _redact_model_secrets(model_dict)
models_list.append(model_dict)
return models_list
- async def get_embedding_models_by_provider(self, provider_uuid: str) -> list[dict]:
+ async def get_embedding_models_by_provider(
+ self,
+ context: TenantContext,
+ provider_uuid: str,
+ *,
+ include_secret: bool = False,
+ ) -> list[dict]:
"""Get embedding models by provider UUID"""
+ await _require_workspace_provider(self.ap, context, provider_uuid)
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.EmbeddingModel).where(
- persistence_model.EmbeddingModel.provider_uuid == provider_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_model.EmbeddingModel).where(
+ persistence_model.EmbeddingModel.provider_uuid == provider_uuid
+ ),
+ persistence_model.EmbeddingModel,
+ context,
)
)
models = result.all()
- return [self.ap.persistence_mgr.serialize_model(persistence_model.EmbeddingModel, m) for m in models]
+ serialized = [self.ap.persistence_mgr.serialize_model(persistence_model.EmbeddingModel, m) for m in models]
+ return serialized if include_secret else [_redact_model_secrets(model) for model in serialized]
- async def create_embedding_model(self, model_data: dict, preserve_uuid: bool = False) -> str:
+ async def create_embedding_model(
+ self, context: TenantContext, model_data: dict, preserve_uuid: bool = False
+ ) -> str:
"""Create a new embedding model"""
+ model_data = model_data.copy()
if not preserve_uuid:
model_data['uuid'] = str(uuid.uuid4())
+ model_data['workspace_uuid'] = require_workspace_uuid(context)
+ if 'extra_args' in model_data:
+ model_data['extra_args'] = restore_secret_placeholders(model_data['extra_args'])
if 'provider' in model_data:
provider_data = model_data.pop('provider')
@@ -310,35 +464,44 @@ class EmbeddingModelsService:
model_data['provider_uuid'] = provider_data['uuid']
else:
provider_uuid = await self.ap.provider_service.find_or_create_provider(
+ context,
requester=provider_data.get('requester', ''),
base_url=provider_data.get('base_url', ''),
api_keys=provider_data.get('api_keys', []),
)
model_data['provider_uuid'] = provider_uuid
- await _validate_provider_supports(self.ap, model_data['provider_uuid'], 'text-embedding')
+ await _require_workspace_provider(self.ap, context, model_data['provider_uuid'])
+ await _validate_provider_supports(self.ap, context, model_data['provider_uuid'], 'text-embedding')
await self.ap.persistence_mgr.execute_async(
sqlalchemy.insert(persistence_model.EmbeddingModel).values(**model_data)
)
- runtime_provider = self.ap.model_mgr.provider_dict.get(model_data['provider_uuid'])
- if runtime_provider is None:
- raise Exception('provider not found')
-
+ runtime_provider = await _require_runtime_provider(self.ap, context, model_data['provider_uuid'])
runtime_embedding_model = await self.ap.model_mgr.load_embedding_model_with_provider(
+ context,
persistence_model.EmbeddingModel(**model_data),
runtime_provider,
)
- self.ap.model_mgr.embedding_models.append(runtime_embedding_model)
+ await self.ap.model_mgr.cache_embedding_model(context, runtime_embedding_model)
return model_data['uuid']
- async def get_embedding_model(self, model_uuid: str) -> dict | None:
+ async def get_embedding_model(
+ self,
+ context: TenantContext,
+ model_uuid: str,
+ include_secret: bool = False,
+ ) -> dict | None:
"""Get a single embedding model with provider info"""
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.EmbeddingModel).where(
- persistence_model.EmbeddingModel.uuid == model_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_model.EmbeddingModel).where(
+ persistence_model.EmbeddingModel.uuid == model_uuid
+ ),
+ persistence_model.EmbeddingModel,
+ context,
)
)
model = result.first()
@@ -348,21 +511,38 @@ class EmbeddingModelsService:
model_dict = self.ap.persistence_mgr.serialize_model(persistence_model.EmbeddingModel, model)
provider_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.ModelProvider).where(
- persistence_model.ModelProvider.uuid == model.provider_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_model.ModelProvider).where(
+ persistence_model.ModelProvider.uuid == model.provider_uuid
+ ),
+ persistence_model.ModelProvider,
+ context,
)
)
provider = provider_result.first()
if provider:
provider_dict = self.ap.persistence_mgr.serialize_model(persistence_model.ModelProvider, provider)
- model_dict['provider'] = _parse_provider_api_keys(provider_dict)
+ provider_dict = _parse_provider_api_keys(provider_dict)
+ model_dict['provider'] = provider_dict
+
+ if not include_secret:
+ model_dict = _redact_model_secrets(model_dict)
return model_dict
- async def update_embedding_model(self, model_uuid: str, model_data: dict) -> None:
+ async def update_embedding_model(self, context: TenantContext, model_uuid: str, model_data: dict) -> None:
"""Update an existing embedding model"""
- if 'uuid' in model_data:
- del model_data['uuid']
+ existing_model = await self.get_embedding_model(context, model_uuid, include_secret=True)
+ if existing_model is None:
+ raise WorkspaceNotFoundError('Model not found')
+ model_data = model_data.copy()
+ model_data.pop('uuid', None)
+ model_data.pop('workspace_uuid', None)
+ if 'extra_args' in model_data:
+ model_data['extra_args'] = restore_secret_placeholders(
+ model_data['extra_args'],
+ existing_model.get('extra_args', {}),
+ )
if 'provider' in model_data:
provider_data = model_data.pop('provider')
@@ -370,57 +550,82 @@ class EmbeddingModelsService:
model_data['provider_uuid'] = provider_data['uuid']
else:
provider_uuid = await self.ap.provider_service.find_or_create_provider(
+ context,
requester=provider_data.get('requester', ''),
base_url=provider_data.get('base_url', ''),
api_keys=provider_data.get('api_keys', []),
)
model_data['provider_uuid'] = provider_uuid
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(persistence_model.EmbeddingModel)
- .where(persistence_model.EmbeddingModel.uuid == model_uuid)
- .values(**model_data)
- )
+ provider_uuid = model_data.get('provider_uuid', existing_model['provider_uuid'])
+ await _require_workspace_provider(self.ap, context, provider_uuid)
+ await _validate_provider_supports(self.ap, context, provider_uuid, 'text-embedding')
- await self.ap.model_mgr.remove_embedding_model(model_uuid)
-
- runtime_provider = self.ap.model_mgr.provider_dict.get(model_data['provider_uuid'])
- if runtime_provider is None:
- raise Exception('provider not found')
-
- runtime_embedding_model = await self.ap.model_mgr.load_embedding_model_with_provider(
- persistence_model.EmbeddingModel(**_runtime_model_data(model_uuid, model_data)),
- runtime_provider,
- )
- self.ap.model_mgr.embedding_models.append(runtime_embedding_model)
-
- async def delete_embedding_model(self, model_uuid: str) -> None:
- """Delete an embedding model"""
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_model.EmbeddingModel).where(
- persistence_model.EmbeddingModel.uuid == model_uuid
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.update(persistence_model.EmbeddingModel)
+ .where(persistence_model.EmbeddingModel.uuid == model_uuid)
+ .values(**model_data),
+ persistence_model.EmbeddingModel,
+ context,
)
)
- await self.ap.model_mgr.remove_embedding_model(model_uuid)
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Model not found')
- async def test_embedding_model(self, model_uuid: str, model_data: dict) -> None:
+ await self.ap.model_mgr.remove_embedding_model(context, model_uuid)
+ runtime_provider = await _require_runtime_provider(self.ap, context, provider_uuid)
+ runtime_embedding_model = await self.ap.model_mgr.load_embedding_model_with_provider(
+ context,
+ persistence_model.EmbeddingModel(
+ **_runtime_model_data(
+ model_uuid,
+ {
+ key: value
+ for key, value in {**existing_model, **model_data, 'provider_uuid': provider_uuid}.items()
+ if key not in {'provider', 'created_at', 'updated_at'}
+ },
+ )
+ ),
+ runtime_provider,
+ )
+ await self.ap.model_mgr.cache_embedding_model(context, runtime_embedding_model)
+
+ async def delete_embedding_model(self, context: TenantContext, model_uuid: str) -> None:
+ """Delete an embedding model"""
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.delete(persistence_model.EmbeddingModel).where(
+ persistence_model.EmbeddingModel.uuid == model_uuid
+ ),
+ persistence_model.EmbeddingModel,
+ context,
+ )
+ )
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Model not found')
+ await self.ap.model_mgr.remove_embedding_model(context, model_uuid)
+
+ async def test_embedding_model(self, context: TenantContext, model_uuid: str, model_data: dict) -> None:
"""Test an embedding model"""
+ require_workspace_uuid(context)
runtime_embedding_model: model_requester.RuntimeEmbeddingModel | None = None
if model_uuid != '_':
- for model in self.ap.model_mgr.embedding_models:
- if model.model_entity.uuid == model_uuid:
- runtime_embedding_model = model
- break
- if runtime_embedding_model is None:
- raise Exception('model not found')
+ if await self.get_embedding_model(context, model_uuid) is None:
+ raise WorkspaceNotFoundError('Model not found')
+ runtime_embedding_model = await self.ap.model_mgr.get_embedding_model_by_uuid(context, model_uuid)
else:
- runtime_embedding_model = await self.ap.model_mgr.init_temporary_runtime_embedding_model(model_data)
+ runtime_embedding_model = await self.ap.model_mgr.init_temporary_runtime_embedding_model(
+ context,
+ model_data,
+ )
await runtime_embedding_model.provider.invoke_embedding(
model=runtime_embedding_model,
input_text=['Hello, world!'],
extra_args={},
+ execution_context=runtime_embedding_model.execution_context,
)
@@ -430,13 +635,17 @@ class RerankModelsService:
def __init__(self, ap: app.Application) -> None:
self.ap = ap
- async def get_rerank_models(self) -> list[dict]:
+ async def get_rerank_models(self, context: TenantContext, include_secret: bool = False) -> list[dict]:
"""Get all rerank models with provider info"""
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_model.RerankModel))
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(sqlalchemy.select(persistence_model.RerankModel), persistence_model.RerankModel, context)
+ )
models = result.all()
providers_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.ModelProvider)
+ scope_statement(
+ sqlalchemy.select(persistence_model.ModelProvider), persistence_model.ModelProvider, context
+ )
)
providers = {p.uuid: p for p in providers_result.all()}
@@ -446,25 +655,44 @@ class RerankModelsService:
provider = providers.get(model.provider_uuid)
if provider:
provider_dict = self.ap.persistence_mgr.serialize_model(persistence_model.ModelProvider, provider)
- model_dict['provider'] = _parse_provider_api_keys(provider_dict)
+ provider_dict = _parse_provider_api_keys(provider_dict)
+ model_dict['provider'] = provider_dict
+ if not include_secret:
+ model_dict = _redact_model_secrets(model_dict)
models_list.append(model_dict)
return models_list
- async def get_rerank_models_by_provider(self, provider_uuid: str) -> list[dict]:
+ async def get_rerank_models_by_provider(
+ self,
+ context: TenantContext,
+ provider_uuid: str,
+ *,
+ include_secret: bool = False,
+ ) -> list[dict]:
"""Get rerank models by provider UUID"""
+ await _require_workspace_provider(self.ap, context, provider_uuid)
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.RerankModel).where(
- persistence_model.RerankModel.provider_uuid == provider_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_model.RerankModel).where(
+ persistence_model.RerankModel.provider_uuid == provider_uuid
+ ),
+ persistence_model.RerankModel,
+ context,
)
)
models = result.all()
- return [self.ap.persistence_mgr.serialize_model(persistence_model.RerankModel, m) for m in models]
+ serialized = [self.ap.persistence_mgr.serialize_model(persistence_model.RerankModel, m) for m in models]
+ return serialized if include_secret else [_redact_model_secrets(model) for model in serialized]
- async def create_rerank_model(self, model_data: dict, preserve_uuid: bool = False) -> str:
+ async def create_rerank_model(self, context: TenantContext, model_data: dict, preserve_uuid: bool = False) -> str:
"""Create a new rerank model"""
+ model_data = model_data.copy()
if not preserve_uuid:
model_data['uuid'] = str(uuid.uuid4())
+ model_data['workspace_uuid'] = require_workspace_uuid(context)
+ if 'extra_args' in model_data:
+ model_data['extra_args'] = restore_secret_placeholders(model_data['extra_args'])
if 'provider' in model_data:
provider_data = model_data.pop('provider')
@@ -472,34 +700,45 @@ class RerankModelsService:
model_data['provider_uuid'] = provider_data['uuid']
else:
provider_uuid = await self.ap.provider_service.find_or_create_provider(
+ context,
requester=provider_data.get('requester', ''),
base_url=provider_data.get('base_url', ''),
api_keys=provider_data.get('api_keys', []),
)
model_data['provider_uuid'] = provider_uuid
- await _validate_provider_supports(self.ap, model_data['provider_uuid'], 'rerank')
+ await _require_workspace_provider(self.ap, context, model_data['provider_uuid'])
+ await _validate_provider_supports(self.ap, context, model_data['provider_uuid'], 'rerank')
await self.ap.persistence_mgr.execute_async(
sqlalchemy.insert(persistence_model.RerankModel).values(**model_data)
)
- runtime_provider = self.ap.model_mgr.provider_dict.get(model_data['provider_uuid'])
- if runtime_provider is None:
- raise Exception('provider not found')
-
+ runtime_provider = await _require_runtime_provider(self.ap, context, model_data['provider_uuid'])
runtime_rerank_model = await self.ap.model_mgr.load_rerank_model_with_provider(
+ context,
persistence_model.RerankModel(**model_data),
runtime_provider,
)
- self.ap.model_mgr.rerank_models.append(runtime_rerank_model)
+ await self.ap.model_mgr.cache_rerank_model(context, runtime_rerank_model)
return model_data['uuid']
- async def get_rerank_model(self, model_uuid: str) -> dict | None:
+ async def get_rerank_model(
+ self,
+ context: TenantContext,
+ model_uuid: str,
+ include_secret: bool = False,
+ ) -> dict | None:
"""Get a single rerank model with provider info"""
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.RerankModel).where(persistence_model.RerankModel.uuid == model_uuid)
+ scope_statement(
+ sqlalchemy.select(persistence_model.RerankModel).where(
+ persistence_model.RerankModel.uuid == model_uuid
+ ),
+ persistence_model.RerankModel,
+ context,
+ )
)
model = result.first()
if model is None:
@@ -508,21 +747,38 @@ class RerankModelsService:
model_dict = self.ap.persistence_mgr.serialize_model(persistence_model.RerankModel, model)
provider_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.ModelProvider).where(
- persistence_model.ModelProvider.uuid == model.provider_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_model.ModelProvider).where(
+ persistence_model.ModelProvider.uuid == model.provider_uuid
+ ),
+ persistence_model.ModelProvider,
+ context,
)
)
provider = provider_result.first()
if provider:
provider_dict = self.ap.persistence_mgr.serialize_model(persistence_model.ModelProvider, provider)
- model_dict['provider'] = _parse_provider_api_keys(provider_dict)
+ provider_dict = _parse_provider_api_keys(provider_dict)
+ model_dict['provider'] = provider_dict
+
+ if not include_secret:
+ model_dict = _redact_model_secrets(model_dict)
return model_dict
- async def update_rerank_model(self, model_uuid: str, model_data: dict) -> None:
+ async def update_rerank_model(self, context: TenantContext, model_uuid: str, model_data: dict) -> None:
"""Update an existing rerank model"""
- if 'uuid' in model_data:
- del model_data['uuid']
+ existing_model = await self.get_rerank_model(context, model_uuid, include_secret=True)
+ if existing_model is None:
+ raise WorkspaceNotFoundError('Model not found')
+ model_data = model_data.copy()
+ model_data.pop('uuid', None)
+ model_data.pop('workspace_uuid', None)
+ if 'extra_args' in model_data:
+ model_data['extra_args'] = restore_secret_placeholders(
+ model_data['extra_args'],
+ existing_model.get('extra_args', {}),
+ )
if 'provider' in model_data:
provider_data = model_data.pop('provider')
@@ -530,50 +786,76 @@ class RerankModelsService:
model_data['provider_uuid'] = provider_data['uuid']
else:
provider_uuid = await self.ap.provider_service.find_or_create_provider(
+ context,
requester=provider_data.get('requester', ''),
base_url=provider_data.get('base_url', ''),
api_keys=provider_data.get('api_keys', []),
)
model_data['provider_uuid'] = provider_uuid
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(persistence_model.RerankModel)
- .where(persistence_model.RerankModel.uuid == model_uuid)
- .values(**model_data)
+ provider_uuid = model_data.get('provider_uuid', existing_model['provider_uuid'])
+ await _require_workspace_provider(self.ap, context, provider_uuid)
+ await _validate_provider_supports(self.ap, context, provider_uuid, 'rerank')
+
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.update(persistence_model.RerankModel)
+ .where(persistence_model.RerankModel.uuid == model_uuid)
+ .values(**model_data),
+ persistence_model.RerankModel,
+ context,
+ )
)
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Model not found')
- await self.ap.model_mgr.remove_rerank_model(model_uuid)
-
- runtime_provider = self.ap.model_mgr.provider_dict.get(model_data['provider_uuid'])
- if runtime_provider is None:
- raise Exception('provider not found')
-
+ await self.ap.model_mgr.remove_rerank_model(context, model_uuid)
+ runtime_provider = await _require_runtime_provider(self.ap, context, provider_uuid)
runtime_rerank_model = await self.ap.model_mgr.load_rerank_model_with_provider(
- persistence_model.RerankModel(**_runtime_model_data(model_uuid, model_data)),
+ context,
+ persistence_model.RerankModel(
+ **_runtime_model_data(
+ model_uuid,
+ {
+ key: value
+ for key, value in {**existing_model, **model_data, 'provider_uuid': provider_uuid}.items()
+ if key not in {'provider', 'created_at', 'updated_at'}
+ },
+ )
+ ),
runtime_provider,
)
- self.ap.model_mgr.rerank_models.append(runtime_rerank_model)
+ await self.ap.model_mgr.cache_rerank_model(context, runtime_rerank_model)
- async def delete_rerank_model(self, model_uuid: str) -> None:
+ async def delete_rerank_model(self, context: TenantContext, model_uuid: str) -> None:
"""Delete a rerank model"""
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_model.RerankModel).where(persistence_model.RerankModel.uuid == model_uuid)
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.delete(persistence_model.RerankModel).where(
+ persistence_model.RerankModel.uuid == model_uuid
+ ),
+ persistence_model.RerankModel,
+ context,
+ )
)
- await self.ap.model_mgr.remove_rerank_model(model_uuid)
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Model not found')
+ await self.ap.model_mgr.remove_rerank_model(context, model_uuid)
- async def test_rerank_model(self, model_uuid: str, model_data: dict) -> None:
+ async def test_rerank_model(self, context: TenantContext, model_uuid: str, model_data: dict) -> None:
"""Test a rerank model"""
+ require_workspace_uuid(context)
runtime_rerank_model: model_requester.RuntimeRerankModel | None = None
if model_uuid != '_':
- for model in self.ap.model_mgr.rerank_models:
- if model.model_entity.uuid == model_uuid:
- runtime_rerank_model = model
- break
- if runtime_rerank_model is None:
- raise Exception('model not found')
+ if await self.get_rerank_model(context, model_uuid) is None:
+ raise WorkspaceNotFoundError('Model not found')
+ runtime_rerank_model = await self.ap.model_mgr.get_rerank_model_by_uuid(context, model_uuid)
else:
- runtime_rerank_model = await self.ap.model_mgr.init_temporary_runtime_rerank_model(model_data)
+ runtime_rerank_model = await self.ap.model_mgr.init_temporary_runtime_rerank_model(
+ context,
+ model_data,
+ )
await runtime_rerank_model.provider.invoke_rerank(
model=runtime_rerank_model,
@@ -582,4 +864,5 @@ class RerankModelsService:
'Artificial intelligence is a branch of computer science.',
'The weather is nice today.',
],
+ execution_context=runtime_rerank_model.execution_context,
)
diff --git a/src/langbot/pkg/api/http/service/monitoring.py b/src/langbot/pkg/api/http/service/monitoring.py
index 46a352ad1..85ca16dd6 100644
--- a/src/langbot/pkg/api/http/service/monitoring.py
+++ b/src/langbot/pkg/api/http/service/monitoring.py
@@ -7,6 +7,9 @@ import sqlalchemy
from ....core import app
from ....entity.persistence import monitoring as persistence_monitoring
+from ..authz import WorkspaceRequiredError
+from ..context import ExecutionContext
+from .tenant import TenantContext, require_workspace_uuid
class MonitoringService:
@@ -17,9 +20,26 @@ class MonitoringService:
def __init__(self, ap: app.Application) -> None:
self.ap = ap
+ @staticmethod
+ def _require_write_context(context: ExecutionContext | None) -> str:
+ """Reject background/runtime writes that lost their execution fence."""
+
+ if not isinstance(context, ExecutionContext):
+ raise WorkspaceRequiredError('Monitoring writes require an ExecutionContext')
+ if not context.instance_uuid.strip() or not context.workspace_uuid.strip():
+ raise WorkspaceRequiredError('Monitoring writes require an instance and Workspace')
+ if context.placement_generation <= 0:
+ raise WorkspaceRequiredError('Monitoring writes require a positive placement generation')
+ return context.workspace_uuid
+
# ========== Cleanup Methods ==========
- async def cleanup_expired_records(self, retention_days: int, batch_size: int = 1000) -> dict[str, int]:
+ async def cleanup_expired_records(
+ self,
+ context: ExecutionContext,
+ retention_days: int,
+ batch_size: int = 1000,
+ ) -> dict[str, int]:
"""Delete monitoring records older than the specified retention period.
Args:
@@ -29,6 +49,7 @@ class MonitoringService:
Returns:
A dict mapping table name to the number of deleted rows.
"""
+ self._require_write_context(context)
if retention_days < 1:
raise ValueError('retention_days must be >= 1')
if batch_size < 1:
@@ -87,6 +108,7 @@ class MonitoringService:
for table_name, model_cls, ts_column, pk_column in tables_and_columns:
deleted_counts[table_name] = await self._delete_expired_in_batches(
+ context=context,
model_cls=model_cls,
ts_column=ts_column,
pk_column=pk_column,
@@ -101,24 +123,31 @@ class MonitoringService:
async def _delete_expired_in_batches(
self,
+ context: ExecutionContext,
model_cls: type,
ts_column: sqlalchemy.Column,
pk_column: sqlalchemy.Column,
cutoff: datetime.datetime,
batch_size: int,
) -> int:
+ workspace_uuid = self._require_write_context(context)
deleted_total = 0
while True:
select_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(pk_column).where(ts_column < cutoff).limit(batch_size)
+ sqlalchemy.select(pk_column)
+ .where(model_cls.workspace_uuid == workspace_uuid, ts_column < cutoff)
+ .limit(batch_size)
)
pk_values = list(select_result.scalars().all())
if not pk_values:
break
delete_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(model_cls).where(pk_column.in_(pk_values))
+ sqlalchemy.delete(model_cls).where(
+ model_cls.workspace_uuid == workspace_uuid,
+ pk_column.in_(pk_values),
+ )
)
deleted = delete_result.rowcount or 0
deleted_total += deleted
@@ -158,13 +187,16 @@ class MonitoringService:
async def _get_message_for_tool_context(
self,
+ context: ExecutionContext,
message_id: str | None = None,
session_id: str | None = None,
):
+ workspace_uuid = self._require_write_context(context)
if message_id:
result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_monitoring.MonitoringMessage).where(
- persistence_monitoring.MonitoringMessage.id == message_id
+ persistence_monitoring.MonitoringMessage.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringMessage.id == message_id,
)
)
row = result.first()
@@ -180,6 +212,7 @@ class MonitoringService:
sqlalchemy.and_(
persistence_monitoring.MonitoringMessage.session_id == session_id,
persistence_monitoring.MonitoringMessage.role == 'user',
+ persistence_monitoring.MonitoringMessage.workspace_uuid == workspace_uuid,
)
)
.order_by(persistence_monitoring.MonitoringMessage.timestamp.desc())
@@ -192,7 +225,10 @@ class MonitoringService:
any_query = (
sqlalchemy.select(persistence_monitoring.MonitoringMessage)
- .where(persistence_monitoring.MonitoringMessage.session_id == session_id)
+ .where(
+ persistence_monitoring.MonitoringMessage.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringMessage.session_id == session_id,
+ )
.order_by(persistence_monitoring.MonitoringMessage.timestamp.desc())
.limit(1)
)
@@ -204,6 +240,7 @@ class MonitoringService:
async def record_message(
self,
+ context: ExecutionContext,
bot_id: str,
bot_name: str,
pipeline_id: str,
@@ -220,9 +257,11 @@ class MonitoringService:
role: str = 'user',
) -> str:
"""Record a message"""
+ workspace_uuid = self._require_write_context(context)
message_id = str(uuid.uuid4())
message_data = {
'id': message_id,
+ 'workspace_uuid': workspace_uuid,
'timestamp': datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
'bot_id': bot_id,
'bot_name': bot_name,
@@ -248,6 +287,7 @@ class MonitoringService:
async def record_llm_call(
self,
+ context: ExecutionContext,
bot_id: str,
bot_name: str,
pipeline_id: str,
@@ -263,9 +303,11 @@ class MonitoringService:
message_id: str | None = None,
) -> str:
"""Record an LLM call"""
+ workspace_uuid = self._require_write_context(context)
call_id = str(uuid.uuid4())
call_data = {
'id': call_id,
+ 'workspace_uuid': workspace_uuid,
'timestamp': datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
'model_name': model_name,
'input_tokens': input_tokens,
@@ -291,6 +333,7 @@ class MonitoringService:
async def record_tool_call(
self,
+ context: ExecutionContext,
tool_name: str,
tool_source: str,
duration: int,
@@ -306,7 +349,12 @@ class MonitoringService:
error_message: str | None = None,
) -> str:
"""Record a tool call."""
- context_message = await self._get_message_for_tool_context(message_id=message_id, session_id=session_id)
+ workspace_uuid = self._require_write_context(context)
+ context_message = await self._get_message_for_tool_context(
+ context,
+ message_id=message_id,
+ session_id=session_id,
+ )
if context_message:
bot_id = bot_id or context_message.bot_id
bot_name = bot_name or context_message.bot_name
@@ -318,6 +366,7 @@ class MonitoringService:
call_id = str(uuid.uuid4())
call_data = {
'id': call_id,
+ 'workspace_uuid': workspace_uuid,
'timestamp': datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
'tool_name': tool_name,
'tool_source': tool_source,
@@ -342,6 +391,7 @@ class MonitoringService:
async def record_embedding_call(
self,
+ context: ExecutionContext,
model_name: str,
prompt_tokens: int,
total_tokens: int,
@@ -356,9 +406,11 @@ class MonitoringService:
call_type: str | None = None,
) -> str:
"""Record an embedding call"""
+ workspace_uuid = self._require_write_context(context)
call_id = str(uuid.uuid4())
call_data = {
'id': call_id,
+ 'workspace_uuid': workspace_uuid,
'timestamp': datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
'model_name': model_name,
'prompt_tokens': prompt_tokens,
@@ -382,6 +434,7 @@ class MonitoringService:
async def record_session_start(
self,
+ context: ExecutionContext,
session_id: str,
bot_id: str,
bot_name: str,
@@ -392,7 +445,9 @@ class MonitoringService:
user_name: str | None = None,
) -> None:
"""Record a new session"""
+ workspace_uuid = self._require_write_context(context)
session_data = {
+ 'workspace_uuid': workspace_uuid,
'session_id': session_id,
'bot_id': bot_id,
'bot_name': bot_name,
@@ -413,6 +468,7 @@ class MonitoringService:
async def update_session_activity(
self,
+ context: ExecutionContext,
session_id: str,
pipeline_id: str | None = None,
pipeline_name: str | None = None,
@@ -424,6 +480,7 @@ class MonitoringService:
Returns:
True if session was found and updated, False if session doesn't exist.
"""
+ workspace_uuid = self._require_write_context(context)
update_values = {
'last_activity': datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
'message_count': persistence_monitoring.MonitoringSession.message_count + 1,
@@ -437,7 +494,10 @@ class MonitoringService:
result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(persistence_monitoring.MonitoringSession)
- .where(persistence_monitoring.MonitoringSession.session_id == session_id)
+ .where(
+ persistence_monitoring.MonitoringSession.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringSession.session_id == session_id,
+ )
.values(update_values)
)
# Check if any rows were updated
@@ -445,6 +505,7 @@ class MonitoringService:
async def record_error(
self,
+ context: ExecutionContext,
bot_id: str,
bot_name: str,
pipeline_id: str,
@@ -456,9 +517,11 @@ class MonitoringService:
message_id: str | None = None,
) -> str:
"""Record an error"""
+ workspace_uuid = self._require_write_context(context)
error_id = str(uuid.uuid4())
error_data = {
'id': error_id,
+ 'workspace_uuid': workspace_uuid,
'timestamp': datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None),
'error_type': error_type,
'error_message': error_message,
@@ -479,12 +542,14 @@ class MonitoringService:
async def update_message_status(
self,
+ context: ExecutionContext,
message_id: str,
status: str,
level: str | None = None,
variables: str | None = None,
) -> None:
"""Update message status and optionally variables"""
+ workspace_uuid = self._require_write_context(context)
update_values = {'status': status}
if level is not None:
update_values['level'] = level
@@ -493,7 +558,10 @@ class MonitoringService:
await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(persistence_monitoring.MonitoringMessage)
- .where(persistence_monitoring.MonitoringMessage.id == message_id)
+ .where(
+ persistence_monitoring.MonitoringMessage.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringMessage.id == message_id,
+ )
.values(update_values)
)
@@ -501,17 +569,19 @@ class MonitoringService:
async def get_overview_metrics(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
) -> dict:
"""Get overview metrics"""
+ workspace_uuid = require_workspace_uuid(context)
# Build base query conditions
- message_conditions = []
- llm_conditions = []
- embedding_conditions = []
- session_conditions = []
+ message_conditions = [persistence_monitoring.MonitoringMessage.workspace_uuid == workspace_uuid]
+ llm_conditions = [persistence_monitoring.MonitoringLLMCall.workspace_uuid == workspace_uuid]
+ embedding_conditions = [persistence_monitoring.MonitoringEmbeddingCall.workspace_uuid == workspace_uuid]
+ session_conditions = [persistence_monitoring.MonitoringSession.workspace_uuid == workspace_uuid]
if bot_ids:
message_conditions.append(persistence_monitoring.MonitoringMessage.bot_id.in_(bot_ids))
@@ -594,6 +664,7 @@ class MonitoringService:
async def get_token_statistics(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
start_time: datetime.datetime | None = None,
@@ -612,8 +683,9 @@ class MonitoringService:
token accounting.
"""
LLMCall = persistence_monitoring.MonitoringLLMCall
+ workspace_uuid = require_workspace_uuid(context)
- conditions = []
+ conditions = [LLMCall.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(LLMCall.bot_id.in_(bot_ids))
if pipeline_ids:
@@ -767,6 +839,7 @@ class MonitoringService:
async def get_messages(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
session_ids: list[str] | None = None,
@@ -776,7 +849,8 @@ class MonitoringService:
offset: int = 0,
) -> tuple[list[dict], int]:
"""Get messages with filters"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringMessage.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringMessage.bot_id.in_(bot_ids))
@@ -820,6 +894,7 @@ class MonitoringService:
async def get_llm_calls(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
start_time: datetime.datetime | None = None,
@@ -828,7 +903,8 @@ class MonitoringService:
offset: int = 0,
) -> tuple[list[dict], int]:
"""Get LLM calls with filters"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringLLMCall.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringLLMCall.bot_id.in_(bot_ids))
@@ -871,6 +947,7 @@ class MonitoringService:
async def get_tool_calls(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
session_ids: list[str] | None = None,
@@ -880,7 +957,8 @@ class MonitoringService:
offset: int = 0,
) -> tuple[list[dict], int]:
"""Get tool calls with filters"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringToolCall.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringToolCall.bot_id.in_(bot_ids))
@@ -923,6 +1001,7 @@ class MonitoringService:
async def get_embedding_calls(
self,
+ context: TenantContext,
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
knowledge_base_id: str | None = None,
@@ -930,7 +1009,8 @@ class MonitoringService:
offset: int = 0,
) -> tuple[list[dict], int]:
"""Get embedding calls with filters"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringEmbeddingCall.workspace_uuid == workspace_uuid]
if start_time:
conditions.append(persistence_monitoring.MonitoringEmbeddingCall.timestamp >= start_time)
@@ -971,6 +1051,7 @@ class MonitoringService:
async def get_sessions(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
start_time: datetime.datetime | None = None,
@@ -980,7 +1061,8 @@ class MonitoringService:
offset: int = 0,
) -> tuple[list[dict], int]:
"""Get sessions with filters"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringSession.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringSession.bot_id.in_(bot_ids))
@@ -1025,6 +1107,7 @@ class MonitoringService:
async def get_errors(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
start_time: datetime.datetime | None = None,
@@ -1033,7 +1116,8 @@ class MonitoringService:
offset: int = 0,
) -> tuple[list[dict], int]:
"""Get errors with filters"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringError.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringError.bot_id.in_(bot_ids))
@@ -1076,12 +1160,15 @@ class MonitoringService:
async def get_session_analysis(
self,
+ context: TenantContext,
session_id: str,
) -> dict:
"""Get detailed analysis for a specific session"""
+ workspace_uuid = require_workspace_uuid(context)
# Get session info
session_query = sqlalchemy.select(persistence_monitoring.MonitoringSession).where(
- persistence_monitoring.MonitoringSession.session_id == session_id
+ persistence_monitoring.MonitoringSession.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringSession.session_id == session_id,
)
session_result = await self.ap.persistence_mgr.execute_async(session_query)
session_row = session_result.first()
@@ -1097,7 +1184,10 @@ class MonitoringService:
# Get messages for this session
messages_query = (
sqlalchemy.select(persistence_monitoring.MonitoringMessage)
- .where(persistence_monitoring.MonitoringMessage.session_id == session_id)
+ .where(
+ persistence_monitoring.MonitoringMessage.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringMessage.session_id == session_id,
+ )
.order_by(persistence_monitoring.MonitoringMessage.timestamp.asc())
)
messages_result = await self.ap.persistence_mgr.execute_async(messages_query)
@@ -1118,7 +1208,8 @@ class MonitoringService:
# Get LLM calls for this session
llm_query = sqlalchemy.select(persistence_monitoring.MonitoringLLMCall).where(
- persistence_monitoring.MonitoringLLMCall.session_id == session_id
+ persistence_monitoring.MonitoringLLMCall.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringLLMCall.session_id == session_id,
)
llm_result = await self.ap.persistence_mgr.execute_async(llm_query)
llm_rows = llm_result.all()
@@ -1146,7 +1237,10 @@ class MonitoringService:
# Get tool calls for this session
tool_query = (
sqlalchemy.select(persistence_monitoring.MonitoringToolCall)
- .where(persistence_monitoring.MonitoringToolCall.session_id == session_id)
+ .where(
+ persistence_monitoring.MonitoringToolCall.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringToolCall.session_id == session_id,
+ )
.order_by(persistence_monitoring.MonitoringToolCall.timestamp.asc())
)
tool_result = await self.ap.persistence_mgr.execute_async(tool_query)
@@ -1174,7 +1268,10 @@ class MonitoringService:
# Get errors for this session
error_query = (
sqlalchemy.select(persistence_monitoring.MonitoringError)
- .where(persistence_monitoring.MonitoringError.session_id == session_id)
+ .where(
+ persistence_monitoring.MonitoringError.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringError.session_id == session_id,
+ )
.order_by(persistence_monitoring.MonitoringError.timestamp.desc())
)
error_result = await self.ap.persistence_mgr.execute_async(error_query)
@@ -1228,12 +1325,15 @@ class MonitoringService:
async def get_message_details(
self,
+ context: TenantContext,
message_id: str,
) -> dict:
"""Get detailed information for a specific message including associated LLM calls and errors"""
+ workspace_uuid = require_workspace_uuid(context)
# Get message info
message_query = sqlalchemy.select(persistence_monitoring.MonitoringMessage).where(
- persistence_monitoring.MonitoringMessage.id == message_id
+ persistence_monitoring.MonitoringMessage.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringMessage.id == message_id,
)
message_result = await self.ap.persistence_mgr.execute_async(message_query)
message_row = message_result.first()
@@ -1249,7 +1349,10 @@ class MonitoringService:
# Get LLM calls for this message
llm_query = (
sqlalchemy.select(persistence_monitoring.MonitoringLLMCall)
- .where(persistence_monitoring.MonitoringLLMCall.message_id == message_id)
+ .where(
+ persistence_monitoring.MonitoringLLMCall.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringLLMCall.message_id == message_id,
+ )
.order_by(persistence_monitoring.MonitoringLLMCall.timestamp.asc())
)
llm_result = await self.ap.persistence_mgr.execute_async(llm_query)
@@ -1271,7 +1374,10 @@ class MonitoringService:
# Get errors for this message
error_query = (
sqlalchemy.select(persistence_monitoring.MonitoringError)
- .where(persistence_monitoring.MonitoringError.message_id == message_id)
+ .where(
+ persistence_monitoring.MonitoringError.workspace_uuid == workspace_uuid,
+ persistence_monitoring.MonitoringError.message_id == message_id,
+ )
.order_by(persistence_monitoring.MonitoringError.timestamp.asc())
)
error_result = await self.ap.persistence_mgr.execute_async(error_query)
@@ -1379,6 +1485,7 @@ class MonitoringService:
async def export_messages(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
start_time: datetime.datetime | None = None,
@@ -1386,7 +1493,8 @@ class MonitoringService:
limit: int = 100000,
) -> list[dict]:
"""Export messages as list of dictionaries for CSV conversion"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringMessage.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringMessage.bot_id.in_(bot_ids))
@@ -1432,6 +1540,7 @@ class MonitoringService:
async def export_llm_calls(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
start_time: datetime.datetime | None = None,
@@ -1439,7 +1548,8 @@ class MonitoringService:
limit: int = 100000,
) -> list[dict]:
"""Export LLM calls as list of dictionaries for CSV conversion"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringLLMCall.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringLLMCall.bot_id.in_(bot_ids))
@@ -1485,13 +1595,15 @@ class MonitoringService:
async def export_embedding_calls(
self,
+ context: TenantContext,
start_time: datetime.datetime | None = None,
end_time: datetime.datetime | None = None,
knowledge_base_id: str | None = None,
limit: int = 100000,
) -> list[dict]:
"""Export embedding calls as list of dictionaries for CSV conversion"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringEmbeddingCall.workspace_uuid == workspace_uuid]
if start_time:
conditions.append(persistence_monitoring.MonitoringEmbeddingCall.timestamp >= start_time)
@@ -1533,6 +1645,7 @@ class MonitoringService:
async def export_errors(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
start_time: datetime.datetime | None = None,
@@ -1540,7 +1653,8 @@ class MonitoringService:
limit: int = 100000,
) -> list[dict]:
"""Export errors as list of dictionaries for CSV conversion"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringError.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringError.bot_id.in_(bot_ids))
@@ -1581,6 +1695,7 @@ class MonitoringService:
async def export_sessions(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
start_time: datetime.datetime | None = None,
@@ -1588,7 +1703,8 @@ class MonitoringService:
limit: int = 100000,
) -> list[dict]:
"""Export sessions as list of dictionaries for CSV conversion"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringSession.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringSession.bot_id.in_(bot_ids))
@@ -1633,6 +1749,7 @@ class MonitoringService:
async def record_feedback(
self,
+ context: ExecutionContext,
feedback_id: str,
feedback_type: int,
feedback_content: str | None = None,
@@ -1646,7 +1763,7 @@ class MonitoringService:
stream_id: str | None = None,
user_id: str | None = None,
platform: str | None = None,
- ) -> str:
+ ) -> str | None:
"""Record user feedback (like/dislike) from AI Bot conversation.
Args:
@@ -1669,6 +1786,7 @@ class MonitoringService:
"""
import json
+ workspace_uuid = self._require_write_context(context)
now = datetime.datetime.now(datetime.timezone.utc).replace(tzinfo=None)
reasons_json = json.dumps(inaccurate_reasons, ensure_ascii=False) if inaccurate_reasons else None
@@ -1677,13 +1795,19 @@ class MonitoringService:
# Handle cancel feedback (type=3): delete existing record
if feedback_type == 3:
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(MonitoringFeedback).where(MonitoringFeedback.feedback_id == feedback_id)
+ sqlalchemy.delete(MonitoringFeedback).where(
+ MonitoringFeedback.workspace_uuid == workspace_uuid,
+ MonitoringFeedback.feedback_id == feedback_id,
+ )
)
return None
# Check if record with this feedback_id already exists
existing_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(MonitoringFeedback).where(MonitoringFeedback.feedback_id == feedback_id)
+ sqlalchemy.select(MonitoringFeedback).where(
+ MonitoringFeedback.workspace_uuid == workspace_uuid,
+ MonitoringFeedback.feedback_id == feedback_id,
+ )
)
existing_row = existing_result.first()
@@ -1692,7 +1816,10 @@ class MonitoringService:
existing = existing_row[0] if isinstance(existing_row, tuple) else existing_row
await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(MonitoringFeedback)
- .where(MonitoringFeedback.feedback_id == feedback_id)
+ .where(
+ MonitoringFeedback.workspace_uuid == workspace_uuid,
+ MonitoringFeedback.feedback_id == feedback_id,
+ )
.values(
timestamp=now,
feedback_type=feedback_type,
@@ -1715,6 +1842,7 @@ class MonitoringService:
record_id = str(uuid.uuid4())
record_data = {
'id': record_id,
+ 'workspace_uuid': workspace_uuid,
'timestamp': now,
'feedback_id': feedback_id,
'feedback_type': feedback_type,
@@ -1737,7 +1865,10 @@ class MonitoringService:
# UNIQUE constraint conflict (concurrent feedback for same feedback_id)
await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(MonitoringFeedback)
- .where(MonitoringFeedback.feedback_id == feedback_id)
+ .where(
+ MonitoringFeedback.workspace_uuid == workspace_uuid,
+ MonitoringFeedback.feedback_id == feedback_id,
+ )
.values(
timestamp=now,
feedback_type=feedback_type,
@@ -1749,6 +1880,7 @@ class MonitoringService:
async def get_feedback_stats(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
start_time: datetime.datetime | None = None,
@@ -1759,7 +1891,8 @@ class MonitoringService:
Returns:
Dictionary with total likes, dislikes, and breakdown by bot/pipeline
"""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringFeedback.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringFeedback.bot_id.in_(bot_ids))
@@ -1837,6 +1970,7 @@ class MonitoringService:
async def get_feedback_list(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
feedback_type: int | None = None,
@@ -1846,7 +1980,8 @@ class MonitoringService:
offset: int = 0,
) -> tuple[list[dict], int]:
"""Get feedback list with filters."""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringFeedback.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringFeedback.bot_id.in_(bot_ids))
@@ -1889,6 +2024,7 @@ class MonitoringService:
async def export_feedback(
self,
+ context: TenantContext,
bot_ids: list[str] | None = None,
pipeline_ids: list[str] | None = None,
start_time: datetime.datetime | None = None,
@@ -1896,7 +2032,8 @@ class MonitoringService:
limit: int = 100000,
) -> list[dict]:
"""Export feedback as list of dictionaries for CSV conversion."""
- conditions = []
+ workspace_uuid = require_workspace_uuid(context)
+ conditions = [persistence_monitoring.MonitoringFeedback.workspace_uuid == workspace_uuid]
if bot_ids:
conditions.append(persistence_monitoring.MonitoringFeedback.bot_id.in_(bot_ids))
diff --git a/src/langbot/pkg/api/http/service/pipeline.py b/src/langbot/pkg/api/http/service/pipeline.py
index 2a6451c8b..d02fda6aa 100644
--- a/src/langbot/pkg/api/http/service/pipeline.py
+++ b/src/langbot/pkg/api/http/service/pipeline.py
@@ -6,6 +6,9 @@ import sqlalchemy
from ....core import app
from ....entity.persistence import pipeline as persistence_pipeline
+from ....workspace.errors import WorkspaceNotFoundError
+from .secrets import contains_secret_placeholder, redact_secrets, restore_secret_placeholders
+from .tenant import TenantContext, require_workspace_uuid, scope_statement
default_stage_order = [
@@ -30,7 +33,8 @@ class PipelineService:
def __init__(self, ap: app.Application) -> None:
self.ap = ap
- async def get_pipeline_metadata(self) -> list[dict]:
+ async def get_pipeline_metadata(self, context: TenantContext) -> list[dict]:
+ require_workspace_uuid(context)
return [
self.ap.pipeline_config_meta_trigger,
self.ap.pipeline_config_meta_safety,
@@ -38,8 +42,19 @@ class PipelineService:
self.ap.pipeline_config_meta_output,
]
- async def get_pipelines(self, sort_by: str = 'created_at', sort_order: str = 'DESC') -> list[dict]:
- query = sqlalchemy.select(persistence_pipeline.LegacyPipeline)
+ async def get_pipelines(
+ self,
+ context: TenantContext,
+ sort_by: str = 'created_at',
+ sort_order: str = 'DESC',
+ *,
+ include_secret: bool = False,
+ ) -> list[dict]:
+ query = scope_statement(
+ sqlalchemy.select(persistence_pipeline.LegacyPipeline),
+ persistence_pipeline.LegacyPipeline,
+ context,
+ )
if sort_by == 'created_at':
if sort_order == 'DESC':
@@ -54,15 +69,26 @@ class PipelineService:
result = await self.ap.persistence_mgr.execute_async(query)
pipelines = result.all()
- return [
+ serialized = [
self.ap.persistence_mgr.serialize_model(persistence_pipeline.LegacyPipeline, pipeline)
for pipeline in pipelines
]
+ return serialized if include_secret else [redact_secrets(pipeline) for pipeline in serialized]
- async def get_pipeline(self, pipeline_uuid: str) -> dict | None:
+ async def get_pipeline(
+ self,
+ context: TenantContext,
+ pipeline_uuid: str,
+ *,
+ include_secret: bool = False,
+ ) -> dict | None:
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
- persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
+ persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid
+ ),
+ persistence_pipeline.LegacyPipeline,
+ context,
)
)
@@ -71,20 +97,24 @@ class PipelineService:
if pipeline is None:
return None
- return self.ap.persistence_mgr.serialize_model(persistence_pipeline.LegacyPipeline, pipeline)
+ serialized = self.ap.persistence_mgr.serialize_model(persistence_pipeline.LegacyPipeline, pipeline)
+ return serialized if include_secret else redact_secrets(serialized)
- async def create_pipeline(self, pipeline_data: dict, default: bool = False) -> str:
+ async def create_pipeline(self, context: TenantContext, pipeline_data: dict, default: bool = False) -> str:
from ....utils import paths as path_utils
+ workspace_uuid = require_workspace_uuid(context)
# Check limitation
limitation = self.ap.instance_config.data.get('system', {}).get('limitation', {})
max_pipelines = limitation.get('max_pipelines', -1)
if max_pipelines >= 0:
- existing_pipelines = await self.get_pipelines()
+ existing_pipelines = await self.get_pipelines(context)
if len(existing_pipelines) >= max_pipelines:
raise ValueError(f'Maximum number of pipelines ({max_pipelines}) reached')
+ pipeline_data = pipeline_data.copy()
pipeline_data['uuid'] = str(uuid.uuid4())
+ pipeline_data['workspace_uuid'] = workspace_uuid
pipeline_data['for_version'] = self.ap.ver_mgr.get_current_version()
pipeline_data['stages'] = default_stage_order.copy()
pipeline_data['is_default'] = default
@@ -108,79 +138,122 @@ class PipelineService:
sqlalchemy.insert(persistence_pipeline.LegacyPipeline).values(**pipeline_data)
)
- pipeline = await self.get_pipeline(pipeline_data['uuid'])
+ pipeline = await self.get_pipeline(context, pipeline_data['uuid'], include_secret=True)
- await self.ap.pipeline_mgr.load_pipeline(pipeline)
+ await self.ap.pipeline_mgr.load_pipeline(context, pipeline)
return pipeline_data['uuid']
- async def update_pipeline(self, pipeline_uuid: str, pipeline_data: dict) -> None:
+ async def update_pipeline(self, context: TenantContext, pipeline_uuid: str, pipeline_data: dict) -> None:
+ workspace_uuid = require_workspace_uuid(context)
pipeline_data = pipeline_data.copy()
- for protected_field in ('uuid', 'for_version', 'stages', 'is_default'):
+ for protected_field in ('uuid', 'workspace_uuid', 'for_version', 'stages', 'is_default'):
pipeline_data.pop(protected_field, None)
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(persistence_pipeline.LegacyPipeline)
- .where(persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid)
- .values(**pipeline_data)
- )
+ if 'config' in pipeline_data:
+ current_config = None
+ if contains_secret_placeholder(pipeline_data['config']):
+ current_pipeline = await self.get_pipeline(context, pipeline_uuid, include_secret=True)
+ if current_pipeline is None:
+ raise WorkspaceNotFoundError('Pipeline not found')
+ current_config = current_pipeline.get('config', {})
+ pipeline_data['config'] = restore_secret_placeholders(
+ pipeline_data['config'],
+ current_config if current_config is not None else {},
+ )
- pipeline = await self.get_pipeline(pipeline_uuid)
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.update(persistence_pipeline.LegacyPipeline)
+ .where(persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid)
+ .values(**pipeline_data),
+ persistence_pipeline.LegacyPipeline,
+ workspace_uuid,
+ )
+ )
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Pipeline not found')
+
+ pipeline = await self.get_pipeline(context, pipeline_uuid, include_secret=True)
+ if pipeline is None:
+ raise WorkspaceNotFoundError('Pipeline not found')
if 'name' in pipeline_data:
from ....entity.persistence import bot as persistence_bot
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_bot.Bot).where(persistence_bot.Bot.use_pipeline_uuid == pipeline_uuid)
+ scope_statement(
+ sqlalchemy.select(persistence_bot.Bot).where(
+ persistence_bot.Bot.use_pipeline_uuid == pipeline_uuid
+ ),
+ persistence_bot.Bot,
+ workspace_uuid,
+ )
)
bots = result.all()
for bot in bots:
bot_data = {'use_pipeline_name': pipeline_data['name']}
- await self.ap.bot_service.update_bot(bot.uuid, bot_data)
+ await self.ap.bot_service.update_bot(context, bot.uuid, bot_data)
- await self.ap.pipeline_mgr.remove_pipeline(pipeline_uuid)
- await self.ap.pipeline_mgr.load_pipeline(pipeline)
+ await self.ap.pipeline_mgr.remove_pipeline(context, pipeline_uuid)
+ await self.ap.pipeline_mgr.load_pipeline(context, pipeline)
# update all conversation that use this pipeline
for session in self.ap.sess_mgr.session_list:
- if session.using_conversation is not None and session.using_conversation.pipeline_uuid == pipeline_uuid:
+ if (
+ session.using_conversation is not None
+ and session.using_conversation.pipeline_uuid == pipeline_uuid
+ and getattr(session, 'workspace_uuid', workspace_uuid) == workspace_uuid
+ ):
session.using_conversation = None
- async def delete_pipeline(self, pipeline_uuid: str) -> None:
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_pipeline.LegacyPipeline).where(
- persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid
+ async def delete_pipeline(self, context: TenantContext, pipeline_uuid: str) -> None:
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.delete(persistence_pipeline.LegacyPipeline).where(
+ persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid
+ ),
+ persistence_pipeline.LegacyPipeline,
+ context,
)
)
- await self.ap.pipeline_mgr.remove_pipeline(pipeline_uuid)
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Pipeline not found')
+ await self.ap.pipeline_mgr.remove_pipeline(context, pipeline_uuid)
- async def copy_pipeline(self, pipeline_uuid: str) -> str:
+ async def copy_pipeline(self, context: TenantContext, pipeline_uuid: str) -> str:
"""Copy a pipeline with all its configurations"""
+ workspace_uuid = require_workspace_uuid(context)
# Check limitation
limitation = self.ap.instance_config.data.get('system', {}).get('limitation', {})
max_pipelines = limitation.get('max_pipelines', -1)
if max_pipelines >= 0:
- existing_pipelines = await self.get_pipelines()
+ existing_pipelines = await self.get_pipelines(context)
if len(existing_pipelines) >= max_pipelines:
raise ValueError(f'Maximum number of pipelines ({max_pipelines}) reached')
# Get the original pipeline
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
- persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
+ persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid
+ ),
+ persistence_pipeline.LegacyPipeline,
+ workspace_uuid,
)
)
original_pipeline = result.first()
if original_pipeline is None:
- raise ValueError(f'Pipeline {pipeline_uuid} not found')
+ raise WorkspaceNotFoundError(f'Pipeline {pipeline_uuid} not found')
# Create new pipeline data
new_uuid = str(uuid.uuid4())
new_pipeline_data = {
'uuid': new_uuid,
+ 'workspace_uuid': workspace_uuid,
'name': f'{original_pipeline.name} (Copy)',
'description': original_pipeline.description,
'for_version': self.ap.ver_mgr.get_current_version(),
@@ -207,13 +280,14 @@ class PipelineService:
)
# Load the new pipeline
- pipeline = await self.get_pipeline(new_uuid)
- await self.ap.pipeline_mgr.load_pipeline(pipeline)
+ pipeline = await self.get_pipeline(context, new_uuid, include_secret=True)
+ await self.ap.pipeline_mgr.load_pipeline(context, pipeline)
return new_uuid
async def update_pipeline_extensions(
self,
+ context: TenantContext,
pipeline_uuid: str,
bound_plugins: list[dict],
bound_mcp_servers: list[str] = None,
@@ -225,16 +299,21 @@ class PipelineService:
mcp_resource_agent_read_enabled: bool | None = None,
) -> None:
"""Update the bound plugins and MCP servers for a pipeline"""
+ workspace_uuid = require_workspace_uuid(context)
# Get current pipeline
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
- persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_pipeline.LegacyPipeline).where(
+ persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid
+ ),
+ persistence_pipeline.LegacyPipeline,
+ workspace_uuid,
)
)
pipeline = result.first()
if pipeline is None:
- raise ValueError(f'Pipeline {pipeline_uuid} not found')
+ raise WorkspaceNotFoundError(f'Pipeline {pipeline_uuid} not found')
# Update extensions_preferences
extensions_preferences = pipeline.extensions_preferences or {}
@@ -252,12 +331,16 @@ class PipelineService:
extensions_preferences['mcp_resources'] = bound_mcp_resources
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(persistence_pipeline.LegacyPipeline)
- .where(persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid)
- .values(extensions_preferences=extensions_preferences)
+ scope_statement(
+ sqlalchemy.update(persistence_pipeline.LegacyPipeline)
+ .where(persistence_pipeline.LegacyPipeline.uuid == pipeline_uuid)
+ .values(extensions_preferences=extensions_preferences),
+ persistence_pipeline.LegacyPipeline,
+ workspace_uuid,
+ )
)
# Reload pipeline to apply changes
- await self.ap.pipeline_mgr.remove_pipeline(pipeline_uuid)
- pipeline = await self.get_pipeline(pipeline_uuid)
- await self.ap.pipeline_mgr.load_pipeline(pipeline)
+ await self.ap.pipeline_mgr.remove_pipeline(context, pipeline_uuid)
+ pipeline = await self.get_pipeline(context, pipeline_uuid, include_secret=True)
+ await self.ap.pipeline_mgr.load_pipeline(context, pipeline)
diff --git a/src/langbot/pkg/api/http/service/provider.py b/src/langbot/pkg/api/http/service/provider.py
index 598d72e8d..0a3a1cc75 100644
--- a/src/langbot/pkg/api/http/service/provider.py
+++ b/src/langbot/pkg/api/http/service/provider.py
@@ -7,6 +7,9 @@ import sqlalchemy
from ....core import app
from ....entity.persistence import model as persistence_model
+from ....workspace.errors import WorkspaceNotFoundError
+from .secrets import contains_secret_placeholder, redact_secrets, restore_secret_placeholders
+from .tenant import TenantContext, require_workspace_uuid, scope_statement
class ModelProviderService:
@@ -35,9 +38,15 @@ class ModelProviderService:
return normalized_keys
- async def get_providers(self) -> list[dict]:
+ async def get_providers(self, context: TenantContext, include_secret: bool = False) -> list[dict]:
"""Get all providers"""
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_model.ModelProvider))
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.select(persistence_model.ModelProvider),
+ persistence_model.ModelProvider,
+ context,
+ )
+ )
providers = result.all()
providers_list = []
for p in providers:
@@ -50,14 +59,25 @@ class ModelProviderService:
provider_dict['api_keys'] = json.loads(provider_dict['api_keys'])
except Exception:
provider_dict['api_keys'] = []
+ if not include_secret:
+ provider_dict = redact_secrets(provider_dict)
providers_list.append(provider_dict)
return providers_list
- async def get_provider(self, provider_uuid: str) -> dict | None:
+ async def get_provider(
+ self,
+ context: TenantContext,
+ provider_uuid: str,
+ include_secret: bool = False,
+ ) -> dict | None:
"""Get a single provider by UUID"""
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.ModelProvider).where(
- persistence_model.ModelProvider.uuid == provider_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_model.ModelProvider).where(
+ persistence_model.ModelProvider.uuid == provider_uuid
+ ),
+ persistence_model.ModelProvider,
+ context,
)
)
provider = result.first()
@@ -72,103 +92,171 @@ class ModelProviderService:
provider_dict['api_keys'] = json.loads(provider_dict['api_keys'])
except Exception:
provider_dict['api_keys'] = []
+ if not include_secret:
+ provider_dict = redact_secrets(provider_dict)
return provider_dict
- async def create_provider(self, provider_data: dict) -> str:
+ async def create_provider(self, context: TenantContext, provider_data: dict) -> str:
"""Create a new provider"""
+ provider_data = provider_data.copy()
provider_data['uuid'] = str(uuid.uuid4())
- provider_data['api_keys'] = self._normalize_api_keys(provider_data.get('api_keys'))
+ provider_data['workspace_uuid'] = require_workspace_uuid(context)
+ provider_data['api_keys'] = self._normalize_api_keys(
+ restore_secret_placeholders(provider_data.get('api_keys'), sensitive=True)
+ )
await self.ap.persistence_mgr.execute_async(
sqlalchemy.insert(persistence_model.ModelProvider).values(**provider_data)
)
# load to runtime
- runtime_provider = await self.ap.model_mgr.load_provider(provider_data)
- self.ap.model_mgr.provider_dict[runtime_provider.provider_entity.uuid] = runtime_provider
+ runtime_provider = await self.ap.model_mgr.load_provider(context, provider_data)
+ await self.ap.model_mgr.cache_provider(context, runtime_provider)
return provider_data['uuid']
- async def update_provider(self, provider_uuid: str, provider_data: dict) -> None:
+ async def update_provider(self, context: TenantContext, provider_uuid: str, provider_data: dict) -> None:
"""Update an existing provider"""
- if 'uuid' in provider_data:
- del provider_data['uuid']
+ provider_data = provider_data.copy()
+ provider_data.pop('uuid', None)
+ provider_data.pop('workspace_uuid', None)
if 'api_keys' in provider_data:
- provider_data['api_keys'] = self._normalize_api_keys(provider_data.get('api_keys'))
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(persistence_model.ModelProvider)
- .where(persistence_model.ModelProvider.uuid == provider_uuid)
- .values(**provider_data)
+ submitted_keys = provider_data.get('api_keys')
+ if contains_secret_placeholder(submitted_keys, sensitive=True):
+ current_provider = await self.get_provider(context, provider_uuid, include_secret=True)
+ if current_provider is None:
+ raise WorkspaceNotFoundError('Provider not found')
+ submitted_keys = restore_secret_placeholders(
+ submitted_keys,
+ current_provider.get('api_keys', []),
+ sensitive=True,
+ )
+ provider_data['api_keys'] = self._normalize_api_keys(submitted_keys)
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.update(persistence_model.ModelProvider)
+ .where(persistence_model.ModelProvider.uuid == provider_uuid)
+ .values(**provider_data),
+ persistence_model.ModelProvider,
+ context,
+ )
)
- await self.ap.model_mgr.reload_provider(provider_uuid)
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Provider not found')
+ await self.ap.model_mgr.reload_provider(context, provider_uuid)
- async def delete_provider(self, provider_uuid: str) -> None:
+ async def delete_provider(self, context: TenantContext, provider_uuid: str) -> None:
"""Delete a provider (only if no models reference it)"""
+ workspace_uuid = require_workspace_uuid(context)
# Check if any models use this provider
llm_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.LLMModel).where(
- persistence_model.LLMModel.provider_uuid == provider_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_model.LLMModel).where(
+ persistence_model.LLMModel.provider_uuid == provider_uuid
+ ),
+ persistence_model.LLMModel,
+ workspace_uuid,
)
)
if llm_result.first() is not None:
raise ValueError('Cannot delete provider: LLM models still reference it')
embedding_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.EmbeddingModel).where(
- persistence_model.EmbeddingModel.provider_uuid == provider_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_model.EmbeddingModel).where(
+ persistence_model.EmbeddingModel.provider_uuid == provider_uuid
+ ),
+ persistence_model.EmbeddingModel,
+ workspace_uuid,
)
)
if embedding_result.first() is not None:
raise ValueError('Cannot delete provider: Embedding models still reference it')
rerank_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.RerankModel).where(
- persistence_model.RerankModel.provider_uuid == provider_uuid
+ scope_statement(
+ sqlalchemy.select(persistence_model.RerankModel).where(
+ persistence_model.RerankModel.provider_uuid == provider_uuid
+ ),
+ persistence_model.RerankModel,
+ workspace_uuid,
)
)
if rerank_result.first() is not None:
raise ValueError('Cannot delete provider: Rerank models still reference it')
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_model.ModelProvider).where(
- persistence_model.ModelProvider.uuid == provider_uuid
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.delete(persistence_model.ModelProvider).where(
+ persistence_model.ModelProvider.uuid == provider_uuid
+ ),
+ persistence_model.ModelProvider,
+ workspace_uuid,
)
)
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Provider not found')
- await self.ap.model_mgr.remove_provider(provider_uuid)
+ await self.ap.model_mgr.remove_provider(context, provider_uuid)
- async def get_provider_model_counts(self, provider_uuid: str) -> dict:
+ async def get_provider_model_counts(self, context: TenantContext, provider_uuid: str) -> dict:
"""Get count of models using this provider"""
+ workspace_uuid = require_workspace_uuid(context)
+ if await self.get_provider(context, provider_uuid) is None:
+ raise WorkspaceNotFoundError('Provider not found')
llm_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(sqlalchemy.func.count())
- .select_from(persistence_model.LLMModel)
- .where(persistence_model.LLMModel.provider_uuid == provider_uuid)
+ scope_statement(
+ sqlalchemy.select(sqlalchemy.func.count())
+ .select_from(persistence_model.LLMModel)
+ .where(persistence_model.LLMModel.provider_uuid == provider_uuid),
+ persistence_model.LLMModel,
+ workspace_uuid,
+ )
)
llm_count = llm_result.scalar() or 0
embedding_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(sqlalchemy.func.count())
- .select_from(persistence_model.EmbeddingModel)
- .where(persistence_model.EmbeddingModel.provider_uuid == provider_uuid)
+ scope_statement(
+ sqlalchemy.select(sqlalchemy.func.count())
+ .select_from(persistence_model.EmbeddingModel)
+ .where(persistence_model.EmbeddingModel.provider_uuid == provider_uuid),
+ persistence_model.EmbeddingModel,
+ workspace_uuid,
+ )
)
embedding_count = embedding_result.scalar() or 0
rerank_result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(sqlalchemy.func.count())
- .select_from(persistence_model.RerankModel)
- .where(persistence_model.RerankModel.provider_uuid == provider_uuid)
+ scope_statement(
+ sqlalchemy.select(sqlalchemy.func.count())
+ .select_from(persistence_model.RerankModel)
+ .where(persistence_model.RerankModel.provider_uuid == provider_uuid),
+ persistence_model.RerankModel,
+ workspace_uuid,
+ )
)
rerank_count = rerank_result.scalar() or 0
return {'llm_count': llm_count, 'embedding_count': embedding_count, 'rerank_count': rerank_count}
- async def find_or_create_provider(self, requester: str, base_url: str, api_keys: list) -> str:
+ async def find_or_create_provider(
+ self,
+ context: TenantContext,
+ requester: str,
+ base_url: str,
+ api_keys: list,
+ ) -> str:
"""Find existing provider or create new one"""
- api_keys = self._normalize_api_keys(api_keys)
+ workspace_uuid = require_workspace_uuid(context)
+ api_keys = self._normalize_api_keys(restore_secret_placeholders(api_keys, sensitive=True))
# Try to find existing provider with same config
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_model.ModelProvider).where(
- persistence_model.ModelProvider.requester == requester,
- persistence_model.ModelProvider.base_url == base_url,
+ scope_statement(
+ sqlalchemy.select(persistence_model.ModelProvider).where(
+ persistence_model.ModelProvider.requester == requester,
+ persistence_model.ModelProvider.base_url == base_url,
+ ),
+ persistence_model.ModelProvider,
+ workspace_uuid,
)
)
for provider in result.all():
@@ -187,29 +275,38 @@ class ModelProviderService:
pass
return await self.create_provider(
+ context,
{
'name': provider_name,
'requester': requester,
'base_url': base_url,
'api_keys': api_keys,
- }
+ },
)
- async def update_space_model_provider_api_keys(self, api_key: str) -> None:
+ async def update_space_model_provider_api_keys(self, context: TenantContext, api_key: str) -> None:
"""Update Space model provider API keys"""
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(persistence_model.ModelProvider)
- .where(persistence_model.ModelProvider.uuid == '00000000-0000-0000-0000-000000000000')
- .values(api_keys=self._normalize_api_keys(api_key))
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.update(persistence_model.ModelProvider)
+ .where(persistence_model.ModelProvider.uuid == '00000000-0000-0000-0000-000000000000')
+ .values(api_keys=self._normalize_api_keys(api_key)),
+ persistence_model.ModelProvider,
+ context,
+ )
)
- await self.ap.model_mgr.reload_provider('00000000-0000-0000-0000-000000000000')
+ if getattr(result, 'rowcount', None) == 0:
+ raise WorkspaceNotFoundError('Provider not found')
+ await self.ap.model_mgr.reload_provider(context, '00000000-0000-0000-0000-000000000000')
- async def scan_provider_models(self, provider_uuid: str, model_type: str | None = None) -> dict:
- provider = await self.get_provider(provider_uuid)
+ async def scan_provider_models(
+ self, context: TenantContext, provider_uuid: str, model_type: str | None = None
+ ) -> dict:
+ provider = await self.get_provider(context, provider_uuid, include_secret=True)
if provider is None:
- raise ValueError('provider not found')
+ raise WorkspaceNotFoundError('Provider not found')
- runtime_provider = await self.ap.model_mgr.load_provider(provider)
+ runtime_provider = await self.ap.model_mgr.load_provider(context, provider)
try:
scan_result = await runtime_provider.requester.scan_models(
@@ -230,8 +327,10 @@ class ModelProviderService:
scanned_models = scan_result
debug_info = None
- llm_models = await self.ap.llm_model_service.get_llm_models_by_provider(provider_uuid)
- embedding_models = await self.ap.embedding_models_service.get_embedding_models_by_provider(provider_uuid)
+ llm_models = await self.ap.llm_model_service.get_llm_models_by_provider(context, provider_uuid)
+ embedding_models = await self.ap.embedding_models_service.get_embedding_models_by_provider(
+ context, provider_uuid
+ )
existing_llm_names = {model['name'] for model in llm_models}
existing_embedding_names = {model['name'] for model in embedding_models}
diff --git a/src/langbot/pkg/api/http/service/secrets.py b/src/langbot/pkg/api/http/service/secrets.py
new file mode 100644
index 000000000..bd0b82ac7
--- /dev/null
+++ b/src/langbot/pkg/api/http/service/secrets.py
@@ -0,0 +1,336 @@
+from __future__ import annotations
+
+import copy
+import re
+from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
+
+
+SECRET_MASK = '***'
+_MISSING_SECRET = object()
+
+_SENSITIVE_NAMES = frozenset(
+ {
+ 'api_key',
+ 'api_keys',
+ 'apikey',
+ 'apikeys',
+ 'auth',
+ 'authorization',
+ 'cookie',
+ 'credentials',
+ 'database_url',
+ 'dsn',
+ 'header_value',
+ 'key',
+ 'proxy_authorization',
+ 'set_cookie',
+ 'webhook_url',
+ }
+)
+_SENSITIVE_TOKENS = frozenset(
+ {
+ 'apikey',
+ 'credential',
+ 'credentials',
+ 'passwd',
+ 'password',
+ 'secret',
+ 'token',
+ }
+)
+_KEY_QUALIFIERS = frozenset(
+ {
+ 'access',
+ 'api',
+ 'auth',
+ 'bearer',
+ 'client',
+ 'debug',
+ 'encryption',
+ 'private',
+ 'signing',
+ }
+)
+_SENSITIVE_URL_QUERY_NAMES = frozenset(
+ {
+ 'code',
+ 'credential',
+ 'credentials',
+ 'password',
+ 'passwd',
+ 'sig',
+ 'signature',
+ }
+)
+
+
+def _normalize_key(key: object) -> str:
+ value = re.sub(r'([a-z0-9])([A-Z])', r'\1_\2', str(key or ''))
+ return re.sub(r'[^a-zA-Z0-9]+', '_', value).strip('_').lower()
+
+
+def is_sensitive_key(key: object) -> bool:
+ """Return whether a configuration key conventionally carries a secret."""
+
+ normalized = _normalize_key(key)
+ if normalized in _SENSITIVE_NAMES:
+ return True
+ tokens = frozenset(token for token in normalized.split('_') if token)
+ if tokens & _SENSITIVE_TOKENS:
+ return True
+ return bool(tokens & {'key', 'keys'}) and bool(tokens & _KEY_QUALIFIERS)
+
+
+def is_url_key(key: object) -> bool:
+ """Return whether a configuration field conventionally carries a URL."""
+
+ normalized = _normalize_key(key)
+ return normalized == 'url' or normalized.endswith('_url')
+
+
+def _is_sensitive_url_query_key(key: object) -> bool:
+ normalized = _normalize_key(key)
+ return (
+ is_sensitive_key(key) or normalized in _SENSITIVE_URL_QUERY_NAMES or normalized.endswith(('_sig', '_signature'))
+ )
+
+
+def _redact_url_string(value: str) -> str:
+ if not value:
+ return value
+ try:
+ parsed = urlsplit(value)
+ netloc = parsed.netloc
+ if '@' in netloc:
+ _, host = netloc.rsplit('@', 1)
+ netloc = f'{SECRET_MASK}@{host}'
+ query = urlencode(
+ [
+ (key, SECRET_MASK if _is_sensitive_url_query_key(key) and item else item)
+ for key, item in parse_qsl(parsed.query, keep_blank_values=True)
+ ],
+ doseq=True,
+ safe='*',
+ )
+ return urlunsplit((parsed.scheme, netloc, parsed.path, query, parsed.fragment))
+ except (TypeError, ValueError):
+ # A malformed URL cannot be safely decomposed, so fail closed.
+ return SECRET_MASK
+
+
+def redact_url_secrets(value):
+ """Redact URL userinfo and credential-like query values."""
+
+ if isinstance(value, str):
+ return _redact_url_string(value)
+ if isinstance(value, list):
+ return [redact_url_secrets(item) for item in value]
+ if isinstance(value, tuple):
+ return tuple(redact_url_secrets(item) for item in value)
+ return copy.deepcopy(value)
+
+
+def _contains_url_secret_placeholder(value) -> bool:
+ if isinstance(value, str):
+ if value == SECRET_MASK:
+ return True
+ try:
+ parsed = urlsplit(value)
+ if '@' in parsed.netloc and SECRET_MASK in parsed.netloc.rsplit('@', 1)[0]:
+ return True
+ return any(
+ item == SECRET_MASK and _is_sensitive_url_query_key(key)
+ for key, item in parse_qsl(parsed.query, keep_blank_values=True)
+ )
+ except (TypeError, ValueError):
+ return False
+ if isinstance(value, (list, tuple)):
+ return any(_contains_url_secret_placeholder(item) for item in value)
+ return False
+
+
+def _restore_url_string(value: str, current_value) -> str:
+ if value == SECRET_MASK:
+ if current_value is _MISSING_SECRET:
+ raise ValueError('Masked URL secret has no existing value')
+ return copy.deepcopy(current_value)
+
+ try:
+ submitted = urlsplit(value)
+ except (TypeError, ValueError):
+ return value
+
+ current = None
+ if isinstance(current_value, str):
+ try:
+ current = urlsplit(current_value)
+ except (TypeError, ValueError):
+ current = None
+
+ netloc = submitted.netloc
+ if '@' in netloc:
+ submitted_userinfo, host = netloc.rsplit('@', 1)
+ if SECRET_MASK in submitted_userinfo:
+ if current is None or '@' not in current.netloc:
+ raise ValueError('Masked URL userinfo has no existing value')
+ current_userinfo, _ = current.netloc.rsplit('@', 1)
+ netloc = f'{current_userinfo}@{host}'
+
+ current_query: dict[str, list[str]] = {}
+ if current is not None:
+ for key, item in parse_qsl(current.query, keep_blank_values=True):
+ current_query.setdefault(_normalize_key(key), []).append(item)
+ consumed: dict[str, int] = {}
+ restored_query: list[tuple[str, str]] = []
+ for key, item in parse_qsl(submitted.query, keep_blank_values=True):
+ normalized = _normalize_key(key)
+ if item == SECRET_MASK and _is_sensitive_url_query_key(key):
+ index = consumed.get(normalized, 0)
+ candidates = current_query.get(normalized, [])
+ if index >= len(candidates):
+ raise ValueError('Masked URL query secret has no existing value')
+ item = candidates[index]
+ consumed[normalized] = index + 1
+ restored_query.append((key, item))
+
+ return urlunsplit(
+ (
+ submitted.scheme,
+ netloc,
+ submitted.path,
+ urlencode(restored_query, doseq=True, safe='*'),
+ submitted.fragment,
+ )
+ )
+
+
+def restore_url_secret_placeholders(value, current_value=_MISSING_SECRET):
+ """Restore URL placeholders from the corresponding persisted URL."""
+
+ if isinstance(value, str):
+ return _restore_url_string(value, current_value)
+ if isinstance(value, list):
+ current_items = current_value if isinstance(current_value, (list, tuple)) else ()
+ return [
+ restore_url_secret_placeholders(
+ item,
+ current_items[index] if index < len(current_items) else _MISSING_SECRET,
+ )
+ for index, item in enumerate(value)
+ ]
+ if isinstance(value, tuple):
+ current_items = current_value if isinstance(current_value, (list, tuple)) else ()
+ return tuple(
+ restore_url_secret_placeholders(
+ item,
+ current_items[index] if index < len(current_items) else _MISSING_SECRET,
+ )
+ for index, item in enumerate(value)
+ )
+ return copy.deepcopy(value)
+
+
+def mask_secret_value(value):
+ """Return a shape-preserving copy whose non-empty leaves are masked."""
+
+ if isinstance(value, dict):
+ return {key: mask_secret_value(item) for key, item in value.items()}
+ if isinstance(value, list):
+ return [mask_secret_value(item) for item in value]
+ if isinstance(value, tuple):
+ return tuple(mask_secret_value(item) for item in value)
+ if value is None or value == '':
+ return value
+ return SECRET_MASK
+
+
+def redact_secrets(value):
+ """Return a recursively redacted copy without mutating the source value."""
+
+ if isinstance(value, dict):
+ return {
+ key: (
+ mask_secret_value(item)
+ if is_sensitive_key(key)
+ else redact_url_secrets(item)
+ if is_url_key(key)
+ else redact_secrets(item)
+ )
+ for key, item in value.items()
+ }
+ if isinstance(value, list):
+ return [redact_secrets(item) for item in value]
+ if isinstance(value, tuple):
+ return tuple(redact_secrets(item) for item in value)
+ return copy.deepcopy(value)
+
+
+def restore_secret_placeholders(value, current_value=_MISSING_SECRET, *, sensitive: bool = False):
+ """Restore masked leaves from existing data before a management write.
+
+ ``***`` is a reserved placeholder only inside a sensitive field. A masked
+ leaf without an existing counterpart is rejected so it can never become a
+ persisted credential. Empty values and explicit replacements pass through.
+ """
+
+ if sensitive and value == SECRET_MASK:
+ if current_value is _MISSING_SECRET:
+ raise ValueError('Masked secret has no existing value')
+ return copy.deepcopy(current_value)
+ if isinstance(value, dict):
+ current_mapping = current_value if isinstance(current_value, dict) else {}
+ return {
+ key: (
+ restore_url_secret_placeholders(
+ item,
+ current_mapping.get(key, _MISSING_SECRET),
+ )
+ if not sensitive and not is_sensitive_key(key) and is_url_key(key)
+ else restore_secret_placeholders(
+ item,
+ current_mapping.get(key, _MISSING_SECRET),
+ sensitive=sensitive or is_sensitive_key(key),
+ )
+ )
+ for key, item in value.items()
+ }
+ if isinstance(value, list):
+ current_items = current_value if isinstance(current_value, (list, tuple)) else ()
+ return [
+ restore_secret_placeholders(
+ item,
+ current_items[index] if index < len(current_items) else _MISSING_SECRET,
+ sensitive=sensitive,
+ )
+ for index, item in enumerate(value)
+ ]
+ if isinstance(value, tuple):
+ current_items = current_value if isinstance(current_value, (list, tuple)) else ()
+ return tuple(
+ restore_secret_placeholders(
+ item,
+ current_items[index] if index < len(current_items) else _MISSING_SECRET,
+ sensitive=sensitive,
+ )
+ for index, item in enumerate(value)
+ )
+ return copy.deepcopy(value)
+
+
+def contains_secret_placeholder(value, *, sensitive: bool = False) -> bool:
+ """Return whether ``value`` contains a meaningful masked secret leaf."""
+
+ if sensitive and value == SECRET_MASK:
+ return True
+ if isinstance(value, dict):
+ return any(
+ (
+ _contains_url_secret_placeholder(item)
+ if not sensitive and not is_sensitive_key(key) and is_url_key(key)
+ else contains_secret_placeholder(item, sensitive=sensitive or is_sensitive_key(key))
+ )
+ for key, item in value.items()
+ )
+ if isinstance(value, (list, tuple)):
+ return any(contains_secret_placeholder(item, sensitive=sensitive) for item in value)
+ return False
diff --git a/src/langbot/pkg/api/http/service/skill.py b/src/langbot/pkg/api/http/service/skill.py
index 94b926975..3c2661dca 100644
--- a/src/langbot/pkg/api/http/service/skill.py
+++ b/src/langbot/pkg/api/http/service/skill.py
@@ -12,6 +12,8 @@ import httpx
from ....core import app
from ....skill.utils import parse_frontmatter
+from ..context import ExecutionContext
+from .tenant import TenantContext, require_workspace_uuid
_PUBLIC_SKILL_FIELDS = (
@@ -75,75 +77,112 @@ class SkillService:
"""Backwards-compatible alias preserved for clarity at call sites."""
self._require_box(action)
+ async def _execution_context(self, context: TenantContext) -> ExecutionContext:
+ workspace_uuid = require_workspace_uuid(context)
+ instance_uuid = str(getattr(context, 'instance_uuid', '') or '').strip()
+ generation = getattr(context, 'placement_generation', None)
+ if not instance_uuid or isinstance(generation, bool) or not isinstance(generation, int) or generation <= 0:
+ raise ValueError('Skill operations require an explicit fenced execution context')
+ binding = await self.ap.workspace_service.get_execution_binding(
+ workspace_uuid,
+ expected_generation=generation,
+ )
+ if binding.instance_uuid != instance_uuid:
+ raise ValueError('Skill execution context belongs to another LangBot instance')
+ return ExecutionContext(
+ instance_uuid=instance_uuid,
+ workspace_uuid=workspace_uuid,
+ placement_generation=generation,
+ bot_uuid=getattr(context, 'bot_uuid', None),
+ pipeline_uuid=getattr(context, 'pipeline_uuid', None),
+ query_uuid=getattr(context, 'query_uuid', None),
+ )
+
@staticmethod
def _serialize_skill(skill: dict) -> dict:
return {field: skill.get(field) for field in _PUBLIC_SKILL_FIELDS if field in skill}
- async def list_skills(self) -> list[dict]:
+ async def list_skills(self, context: TenantContext) -> list[dict]:
+ execution_context = await self._execution_context(context)
# When Box is unavailable, surface an empty list rather than raising —
# the skills page should render cleanly, and the UI separately renders
# a "Box disabled / unavailable" banner via useBoxStatus.
box_service = self._box_service()
if box_service is None:
return []
- return [self._serialize_skill(skill) for skill in await box_service.list_skills()]
+ return [self._serialize_skill(skill) for skill in await box_service.list_skills(execution_context)]
- async def get_skill(self, skill_name: str) -> Optional[dict]:
+ async def get_skill(self, context: TenantContext, skill_name: str) -> Optional[dict]:
+ execution_context = await self._execution_context(context)
box_service = self._box_service()
if box_service is None:
return None
- skill = await box_service.get_skill(skill_name)
+ skill = await box_service.get_skill(execution_context, skill_name)
return self._serialize_skill(skill) if skill else None
- async def get_skill_by_name(self, name: str) -> Optional[dict]:
- return await self.get_skill(name)
+ async def get_skill_by_name(self, context: TenantContext, name: str) -> Optional[dict]:
+ return await self.get_skill(context, name)
- async def create_skill(self, data: dict) -> dict:
+ async def create_skill(self, context: TenantContext, data: dict) -> dict:
+ execution_context = await self._execution_context(context)
box_service = self._require_box('Creating a skill')
- created = await box_service.create_skill(data)
- await self._reload_skills()
+ created = await box_service.create_skill(execution_context, data)
+ await self._reload_skills(execution_context)
return self._serialize_skill(created)
- async def update_skill(self, skill_name: str, data: dict) -> dict:
+ async def update_skill(self, context: TenantContext, skill_name: str, data: dict) -> dict:
+ execution_context = await self._execution_context(context)
box_service = self._require_box('Editing a skill')
- updated = await box_service.update_skill(skill_name, data)
- await self._reload_skills()
+ updated = await box_service.update_skill(execution_context, skill_name, data)
+ await self._reload_skills(execution_context)
return self._serialize_skill(updated)
- async def delete_skill(self, skill_name: str) -> bool:
+ async def delete_skill(self, context: TenantContext, skill_name: str) -> bool:
+ execution_context = await self._execution_context(context)
box_service = self._require_box('Deleting a skill')
- await box_service.delete_skill(skill_name)
- await self._reload_skills()
+ await box_service.delete_skill(execution_context, skill_name)
+ await self._reload_skills(execution_context)
return True
async def list_skill_files(
self,
+ context: TenantContext,
skill_name: str,
path: str = '.',
include_hidden: bool = False,
max_entries: int = 200,
) -> dict:
+ execution_context = await self._execution_context(context)
box_service = self._require_box('Browsing skill files')
- return await box_service.list_skill_files(skill_name, path, include_hidden, max_entries)
+ return await box_service.list_skill_files(execution_context, skill_name, path, include_hidden, max_entries)
- async def read_skill_file(self, skill_name: str, path: str) -> dict:
+ async def read_skill_file(self, context: TenantContext, skill_name: str, path: str) -> dict:
+ execution_context = await self._execution_context(context)
box_service = self._require_box('Reading a skill file')
- return await box_service.read_skill_file(skill_name, path)
+ return await box_service.read_skill_file(execution_context, skill_name, path)
- async def write_skill_file(self, skill_name: str, path: str, content: str) -> dict:
+ async def write_skill_file(self, context: TenantContext, skill_name: str, path: str, content: str) -> dict:
+ execution_context = await self._execution_context(context)
box_service = self._require_box('Editing skill files')
- result = await box_service.write_skill_file(skill_name, path, content)
- await self._reload_skills()
+ result = await box_service.write_skill_file(execution_context, skill_name, path, content)
+ await self._reload_skills(execution_context)
return result
- async def install_from_github(self, data: dict) -> list[dict]:
+ async def install_from_github(self, context: TenantContext, data: dict) -> list[dict]:
+ execution_context = await self._execution_context(context)
box_service = self._require_box('Installing a skill from GitHub')
owner = str(data['owner']).strip()
repo = str(data['repo']).strip()
release_tag = str(data.get('release_tag', '')).strip()
raw_asset_url = str(data['asset_url']).strip()
if self._is_github_skill_md_url(raw_asset_url):
- return await self._install_github_skill_md(raw_asset_url, owner=owner, repo=repo, data=data)
+ return await self._install_github_skill_md(
+ execution_context,
+ raw_asset_url,
+ owner=owner,
+ repo=repo,
+ data=data,
+ )
asset_url = self._validate_github_asset_url(raw_asset_url, owner=owner, repo=repo, release_tag=release_tag)
source_subdir = str(data.get('source_subdir', '') or '').strip()
@@ -151,29 +190,37 @@ class SkillService:
zip_bytes = await self._download_github_asset(asset_url)
filename = f'{repo}-{release_tag.lstrip("v").replace("/", "-") or "source"}.zip'
installed = await box_service.install_skill_zip(
+ execution_context,
zip_bytes,
filename,
source_paths=data.get('source_paths') or [],
source_path=str(data.get('source_path', '') or ''),
source_subdir=source_subdir,
)
- await self._reload_skills()
+ await self._reload_skills(execution_context)
return [self._serialize_skill(skill) for skill in installed]
- async def preview_install_from_github(self, data: dict) -> list[dict]:
+ async def preview_install_from_github(self, context: TenantContext, data: dict) -> list[dict]:
+ execution_context = await self._execution_context(context)
box_service = self._require_box('Previewing a skill from GitHub')
owner = str(data['owner']).strip()
repo = str(data['repo']).strip()
release_tag = str(data.get('release_tag', '')).strip()
raw_asset_url = str(data['asset_url']).strip()
if self._is_github_skill_md_url(raw_asset_url):
- return await self._preview_github_skill_md(raw_asset_url, owner=owner, repo=repo)
+ return await self._preview_github_skill_md(
+ execution_context,
+ raw_asset_url,
+ owner=owner,
+ repo=repo,
+ )
asset_url = self._validate_github_asset_url(raw_asset_url, owner=owner, repo=repo, release_tag=release_tag)
source_subdir = str(data.get('source_subdir', '') or '').strip()
zip_bytes = await self._download_github_asset(asset_url)
return await box_service.preview_skill_zip(
+ execution_context,
zip_bytes,
f'{repo}-{release_tag.lstrip("v").replace("/", "-") or "source"}.zip',
source_subdir=source_subdir,
@@ -181,27 +228,45 @@ class SkillService:
async def install_from_zip_upload(
self,
+ context: TenantContext,
*,
file_bytes: bytes,
filename: str,
source_paths: list[str] | None = None,
source_path: str = '',
) -> list[dict]:
+ execution_context = await self._execution_context(context)
box_service = self._require_box('Installing a skill from upload')
installed = await box_service.install_skill_zip(
+ execution_context,
file_bytes,
filename,
source_paths=source_paths or [],
source_path=source_path,
)
- await self._reload_skills()
+ await self._reload_skills(execution_context)
return [self._serialize_skill(skill) for skill in installed]
- async def preview_install_from_zip_upload(self, *, file_bytes: bytes, filename: str) -> list[dict]:
+ async def preview_install_from_zip_upload(
+ self,
+ context: TenantContext,
+ *,
+ file_bytes: bytes,
+ filename: str,
+ ) -> list[dict]:
+ execution_context = await self._execution_context(context)
box_service = self._require_box('Previewing a skill upload')
- return await box_service.preview_skill_zip(file_bytes, filename)
+ return await box_service.preview_skill_zip(execution_context, file_bytes, filename)
- async def _install_github_skill_md(self, asset_url: str, *, owner: str, repo: str, data: dict) -> list[dict]:
+ async def _install_github_skill_md(
+ self,
+ context: TenantContext,
+ asset_url: str,
+ *,
+ owner: str,
+ repo: str,
+ data: dict,
+ ) -> list[dict]:
box_service = self._require_box('Installing a skill from GitHub')
zip_bytes, filename, _package_name = await self._download_github_skill_directory_as_zip(
asset_url,
@@ -210,38 +275,48 @@ class SkillService:
)
installed = await box_service.install_skill_zip(
+ context,
zip_bytes,
filename,
source_paths=data.get('source_paths') or [],
source_path=str(data.get('source_path', '') or ''),
target_suffix='',
)
- await self._reload_skills()
+ await self._reload_skills(context)
return [self._serialize_skill(skill) for skill in installed]
- async def _preview_github_skill_md(self, asset_url: str, *, owner: str, repo: str) -> list[dict]:
+ async def _preview_github_skill_md(
+ self,
+ context: TenantContext,
+ asset_url: str,
+ *,
+ owner: str,
+ repo: str,
+ ) -> list[dict]:
box_service = self._require_box('Previewing a skill from GitHub')
zip_bytes, _filename, package_name = await self._download_github_skill_directory_as_zip(
asset_url,
owner=owner,
repo=repo,
)
- return await box_service.preview_skill_zip(zip_bytes, f'{package_name}.zip', target_suffix='')
+ return await box_service.preview_skill_zip(context, zip_bytes, f'{package_name}.zip', target_suffix='')
- async def reload_skills(self) -> list[dict]:
- await self._reload_skills()
- return await self.list_skills()
+ async def reload_skills(self, context: TenantContext) -> list[dict]:
+ execution_context = await self._execution_context(context)
+ await self._reload_skills(execution_context)
+ return await self.list_skills(execution_context)
- async def scan_directory_async(self, path: str) -> dict:
+ async def scan_directory_async(self, context: TenantContext, path: str) -> dict:
+ execution_context = await self._execution_context(context)
box_service = self._require_box('Scanning a skill directory')
- return await box_service.scan_skill_directory(path)
+ return await box_service.scan_skill_directory(execution_context, path)
- async def _reload_skills(self) -> None:
+ async def _reload_skills(self, context: TenantContext) -> None:
skill_mgr = getattr(self.ap, 'skill_mgr', None)
reload_skills = getattr(skill_mgr, 'reload_skills', None)
if not callable(reload_skills):
return
- result = reload_skills()
+ result = reload_skills(context)
if inspect.isawaitable(result):
await result
diff --git a/src/langbot/pkg/api/http/service/space.py b/src/langbot/pkg/api/http/service/space.py
index 6de259321..d280d5ff1 100644
--- a/src/langbot/pkg/api/http/service/space.py
+++ b/src/langbot/pkg/api/http/service/space.py
@@ -85,12 +85,14 @@ class SpaceService:
def get_oauth_authorize_url(self, redirect_uri: str, state: str = '') -> str:
"""Get the Space OAuth authorization URL for redirect"""
+ from urllib.parse import urlencode
+
space_config = self._get_space_config()
authorize_url = space_config['oauth_authorize_url']
- params = f'redirect_uri={redirect_uri}'
+ params = {'redirect_uri': redirect_uri}
if state:
- params += f'&state={state}'
- return f'{authorize_url}?{params}'
+ params['state'] = state
+ return f'{authorize_url}?{urlencode(params)}'
async def exchange_oauth_code(self, code: str) -> typing.Dict:
"""Exchange OAuth authorization code for tokens"""
diff --git a/src/langbot/pkg/api/http/service/tenant.py b/src/langbot/pkg/api/http/service/tenant.py
new file mode 100644
index 000000000..7b0fc3805
--- /dev/null
+++ b/src/langbot/pkg/api/http/service/tenant.py
@@ -0,0 +1,34 @@
+from __future__ import annotations
+
+import typing
+
+from ..authz import WorkspaceRequiredError
+from ..context import ExecutionContext, RequestContext, WorkspaceContext
+
+TenantContext: typing.TypeAlias = RequestContext | ExecutionContext | WorkspaceContext | str
+
+
+def require_workspace_uuid(context: TenantContext | None) -> str:
+ """Resolve an explicit Workspace UUID without allowing a global fallback."""
+
+ if isinstance(context, str):
+ workspace_uuid = context
+ elif isinstance(context, RequestContext):
+ workspace_uuid = context.workspace_uuid
+ elif isinstance(context, ExecutionContext):
+ workspace_uuid = context.workspace_uuid
+ elif isinstance(context, WorkspaceContext):
+ workspace_uuid = context.workspace_uuid
+ else:
+ raise WorkspaceRequiredError('Workspace context is required')
+
+ normalized = workspace_uuid.strip()
+ if not normalized:
+ raise WorkspaceRequiredError('Workspace context is required')
+ return normalized
+
+
+def scope_statement(statement: typing.Any, model: typing.Any, context: TenantContext) -> typing.Any:
+ """Add the mandatory Workspace predicate to a SQLAlchemy statement."""
+
+ return statement.where(model.workspace_uuid == require_workspace_uuid(context))
diff --git a/src/langbot/pkg/api/http/service/user.py b/src/langbot/pkg/api/http/service/user.py
index a9185f9bc..9e33f4e67 100644
--- a/src/langbot/pkg/api/http/service/user.py
+++ b/src/langbot/pkg/api/http/service/user.py
@@ -6,51 +6,260 @@ import jwt
import datetime
import typing
import asyncio
+import hashlib
+import secrets
+import time
+import uuid
+
+from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
-from ....core import app
from ....entity.persistence import user
from ....utils import constants
from ....entity.errors import account as account_errors
+from ....workspace.collaboration import normalize_email
+from ..authz import Permission, permissions_for_role
+
+if typing.TYPE_CHECKING:
+ from ....core.app import Application
+
+
+class AccountExistsLoginRequiredError(ValueError):
+ code = 'account_exists_login_required'
+
+
+class PublicRegistrationClosedError(ValueError):
+ code = 'registration_closed'
+
+
+class ControlPlaneDirectoryRequiredError(PublicRegistrationClosedError):
+ code = 'control_plane_required'
+
+
+class AccountDisabledError(ValueError):
+ code = 'account_disabled'
class UserService:
- ap: app.Application
+ ap: Application
_create_user_lock: asyncio.Lock
- def __init__(self, ap: app.Application) -> None:
+ def __init__(self, ap: Application) -> None:
self.ap = ap
self._create_user_lock = asyncio.Lock()
self._password_hash_lock = asyncio.Semaphore(1)
+ self._space_oauth_state_lock = asyncio.Lock()
+ self._space_oauth_states: dict[str, tuple[str, str | None, float]] = {}
+
+ @staticmethod
+ def _space_oauth_state_digest(state: str) -> str:
+ return hashlib.sha256(state.encode('utf-8')).hexdigest()
+
+ async def issue_space_oauth_state(
+ self,
+ purpose: typing.Literal['login', 'bind'],
+ *,
+ account_uuid: str | None = None,
+ ttl_seconds: int = 600,
+ ) -> str:
+ """Issue an opaque, single-use OAuth state without exposing a JWT."""
+ if purpose == 'bind' and not account_uuid:
+ raise ValueError('An Account is required for Space binding')
+ if purpose == 'login' and account_uuid is not None:
+ raise ValueError('Login state cannot be bound to an Account')
+ if ttl_seconds <= 0:
+ raise ValueError('OAuth state lifetime must be positive')
+
+ raw_state = secrets.token_urlsafe(32)
+ digest = self._space_oauth_state_digest(raw_state)
+ expires_at = time.monotonic() + min(ttl_seconds, 600)
+ async with self._space_oauth_state_lock:
+ now = time.monotonic()
+ self._space_oauth_states = {key: value for key, value in self._space_oauth_states.items() if value[2] > now}
+ if len(self._space_oauth_states) >= 4096:
+ oldest = min(self._space_oauth_states, key=lambda key: self._space_oauth_states[key][2])
+ self._space_oauth_states.pop(oldest, None)
+ self._space_oauth_states[digest] = (purpose, account_uuid, expires_at)
+ return raw_state
+
+ async def consume_space_oauth_state(
+ self,
+ raw_state: str,
+ purpose: typing.Literal['login', 'bind'],
+ ) -> user.User | None:
+ """Atomically consume OAuth state and resolve its active bind Account."""
+ if not isinstance(raw_state, str) or not raw_state:
+ raise ValueError('Invalid or expired OAuth state')
+ digest = self._space_oauth_state_digest(raw_state)
+ async with self._space_oauth_state_lock:
+ entry = self._space_oauth_states.pop(digest, None)
+ if entry is None or entry[0] != purpose or entry[2] <= time.monotonic():
+ raise ValueError('Invalid or expired OAuth state')
+ if purpose == 'login':
+ return None
+
+ account_uuid = entry[1]
+ account = await self.get_user_by_uuid(account_uuid or '')
+ if account is None:
+ raise ValueError('Invalid or expired OAuth state')
+ self._require_active_account(account)
+ return account
async def _hash_password(self, password: str) -> str:
async with self._password_hash_lock:
return await asyncio.to_thread(argon2.PasswordHasher().hash, password)
+ def _require_local_directory(self) -> None:
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ if workspace_service is not None and workspace_service.policy.multi_workspace_enabled:
+ raise ControlPlaneDirectoryRequiredError(
+ 'Cloud Accounts and directory changes are managed by the SaaS control plane'
+ )
+
async def _verify_password(self, hashed_password: str, password: str) -> None:
async with self._password_hash_lock:
await asyncio.to_thread(argon2.PasswordHasher().verify, hashed_password, password)
+ async def _update_space_provider_for_account(self, account: typing.Any, api_key: str) -> None:
+ """Refresh the OSS Workspace Space provider without guessing a SaaS Workspace.
+
+ Space OAuth credentials belong to an Account, while model-provider secrets
+ belong to a Workspace. Community edition has one unambiguous Workspace, so
+ the historical automatic refresh remains available to members allowed to
+ manage provider secrets. In multi-Workspace SaaS mode the OAuth callback has
+ no trusted Workspace selector; the closed control plane or an explicit
+ Workspace settings action must perform that linkage instead.
+ """
+
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ collaboration_service = getattr(self.ap, 'workspace_collaboration_service', None)
+ account_uuid = getattr(account, 'uuid', None)
+ if workspace_service is None or collaboration_service is None or not isinstance(account_uuid, str):
+ # Never turn a missing tenant kernel into a global secret mutation.
+ return
+ if workspace_service.policy.multi_workspace_enabled:
+ return
+
+ accesses = await collaboration_service.list_account_workspaces(account_uuid)
+ if len(accesses) != 1:
+ return
+ access = accesses[0]
+ if Permission.PROVIDER_SECRET_MANAGE.value not in permissions_for_role(access.membership.role):
+ return
+ await self.ap.provider_service.update_space_model_provider_api_keys(
+ access.workspace.uuid,
+ api_key,
+ )
+
async def is_initialized(self) -> bool:
result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(user.User).limit(1))
result_list = result.all()
return result_list is not None and len(result_list) > 0
+ def _session_factory(self) -> async_sessionmaker[AsyncSession]:
+ return async_sessionmaker(self.ap.persistence_mgr.get_db_engine(), expire_on_commit=False)
+
+ def _jwt_identity(self) -> tuple[str, str]:
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ instance_uuid = str(getattr(workspace_service, 'instance_uuid', '') or constants.instance_id).strip()
+ # UserService is constructed only after config/bootstrap in production.
+ # The fallback keeps lightweight isolated unit tests deterministic.
+ if not instance_uuid:
+ instance_uuid = 'uninitialized-test-instance'
+ return 'langbot-core', f'langbot-instance:{instance_uuid}'
+
+ def _legacy_local_tokens_allowed(self) -> bool:
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ policy = getattr(workspace_service, 'policy', None)
+ return getattr(policy, 'multi_workspace_enabled', False) is not True
+
async def create_user(self, user_email: str, password: str) -> None:
+ """Create the first local Account and Workspace owner atomically."""
+
+ await self.create_initial_account(user_email, password)
+
+ async def create_initial_account(self, user_email: str, password: str) -> user.User:
+ self._require_local_directory()
+ normalized_email = normalize_email(user_email)
hashed_password = await self._hash_password(password)
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.insert(user.User).values(user=user_email, password=hashed_password, account_type='local')
+ async with self._create_user_lock:
+ async with self._session_factory()() as session:
+ async with session.begin():
+ existing_count = int(
+ (await session.scalar(sqlalchemy.select(sqlalchemy.func.count()).select_from(user.User))) or 0
+ )
+ if existing_count:
+ raise PublicRegistrationClosedError('System already initialized')
+ account = self._new_account(normalized_email, hashed_password)
+ session.add(account)
+ await session.flush()
+ await self.ap.workspace_service.bootstrap_local_account(account.uuid, session=session)
+ return account
+
+ async def register_invited_account(
+ self,
+ invitation_token: str,
+ user_email: str,
+ password: str,
+ ) -> tuple[user.User, typing.Any, str]:
+ """Create an invited Account and accept its Membership in one transaction."""
+
+ self._require_local_directory()
+ normalized_email = normalize_email(user_email)
+ invitation, _ = await self.ap.workspace_collaboration_service.inspect_invitation(invitation_token)
+ if invitation.normalized_email != normalized_email:
+ from ....workspace.collaboration import InvitationEmailMismatchError
+
+ raise InvitationEmailMismatchError('Invitation email does not match the Account')
+ hashed_password = await self._hash_password(password)
+
+ async with self._create_user_lock:
+ async with self._session_factory()() as session:
+ async with session.begin():
+ existing = await session.scalar(
+ sqlalchemy.select(user.User).where(user.User.normalized_email == normalized_email)
+ )
+ if existing is not None:
+ raise AccountExistsLoginRequiredError('An Account already exists for this email')
+ account = self._new_account(normalized_email, hashed_password)
+ session.add(account)
+ await session.flush()
+ membership = await self.ap.workspace_collaboration_service.accept_invitation(
+ invitation_token,
+ account.uuid,
+ session=session,
+ )
+ token = await self.generate_jwt_token(account)
+ return account, membership, token
+
+ def _new_account(self, normalized_email: str, hashed_password: str) -> user.User:
+ return user.User(
+ uuid=str(uuid.uuid4()),
+ user=normalized_email,
+ normalized_email=normalized_email,
+ password=hashed_password,
+ account_type='local',
+ status=user.AccountStatus.ACTIVE.value,
+ source=user.AccountSource.LOCAL.value,
+ projection_revision=0,
)
async def get_user_by_email(self, user_email: str) -> user.User | None:
+ normalized_email = user_email.strip().casefold()
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(user.User).where(user.User.user == user_email)
+ sqlalchemy.select(user.User).where(user.User.normalized_email == normalized_email)
)
result_list = result.all()
return result_list[0] if result_list is not None and len(result_list) > 0 else None
+ async def get_user_by_uuid(self, account_uuid: str) -> user.User | None:
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(user.User).where(user.User.uuid == account_uuid)
+ )
+ return result.first()
+
async def get_user_by_space_account_uuid(self, space_account_uuid: str) -> user.User | None:
"""Get user by Space account UUID"""
result = await self.ap.persistence_mgr.execute_async(
@@ -61,16 +270,10 @@ class UserService:
return result_list[0] if result_list is not None and len(result_list) > 0 else None
async def authenticate(self, user_email: str, password: str) -> str | None:
- result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(user.User).where(user.User.user == user_email)
- )
-
- result_list = result.all()
-
- if result_list is None or len(result_list) == 0:
+ user_obj = await self.get_user_by_email(user_email)
+ if user_obj is None:
raise ValueError('用户不存在')
-
- user_obj = result_list[0]
+ self._require_active_account(user_obj)
# Check if this user has a local password set
if not user_obj.password:
@@ -78,30 +281,119 @@ class UserService:
await self._verify_password(user_obj.password, password)
- return await self.generate_jwt_token(user_email)
+ return await self.generate_jwt_token(user_obj)
- async def generate_jwt_token(self, user_email: str) -> str:
+ async def generate_jwt_token(self, account: user.User | str) -> str:
jwt_secret = self.ap.instance_config.data['system']['jwt']['secret']
jwt_expire = self.ap.instance_config.data['system']['jwt']['expire']
+ account_obj: user.User | None = account if not isinstance(account, str) and hasattr(account, 'user') else None
+ user_email = account_obj.user if account_obj is not None else account
+ if account_obj is None and hasattr(self.ap, 'persistence_mgr'):
+ try:
+ account_obj = await self.get_user_by_email(user_email)
+ except (AttributeError, TypeError):
+ # Lightweight unit-test and bootstrap callers may not have persistence wired.
+ account_obj = None
+
payload = {
'user': user_email,
- 'iss': 'LangBot-' + constants.edition,
+ 'iss': self._jwt_identity()[0],
+ 'aud': self._jwt_identity()[1],
'exp': datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(seconds=jwt_expire),
}
+ if account_obj is not None:
+ self._require_active_account(account_obj)
+ payload.update(
+ {
+ 'sub': account_obj.uuid,
+ 'account_revision': account_obj.projection_revision,
+ }
+ )
return jwt.encode(payload, jwt_secret, algorithm='HS256')
async def verify_jwt_token(self, token: str) -> str:
- jwt_secret = self.ap.instance_config.data['system']['jwt']['secret']
+ account = await self.get_authenticated_account(token, allow_unresolved_legacy=True)
+ if isinstance(account, str):
+ return account
+ return account.user
- return jwt.decode(token, jwt_secret, algorithms=['HS256'])['user']
+ async def get_authenticated_account(
+ self,
+ token: str,
+ *,
+ allow_unresolved_legacy: bool = False,
+ ) -> user.User | str:
+ """Resolve a JWT to an active Account, accepting bounded legacy email tokens."""
+
+ jwt_secret = self.ap.instance_config.data['system']['jwt']['secret']
+ issuer, audience = self._jwt_identity()
+ try:
+ payload = jwt.decode(
+ token,
+ jwt_secret,
+ algorithms=['HS256'],
+ issuer=issuer,
+ audience=audience,
+ options={'require': ['exp', 'iss', 'aud']},
+ )
+ except jwt.MissingRequiredClaimError:
+ # Preserve one bounded OSS upgrade path for previously issued
+ # community tokens. SaaS/Cloud policy never accepts these tokens,
+ # and a token carrying a new-style or foreign audience cannot fall
+ # back into the legacy decoder.
+ unverified = jwt.decode(token, options={'verify_signature': False})
+ if (
+ not self._legacy_local_tokens_allowed()
+ or 'aud' in unverified
+ or unverified.get('iss') != 'LangBot-community'
+ ):
+ raise
+ payload = jwt.decode(
+ token,
+ jwt_secret,
+ algorithms=['HS256'],
+ options={'require': ['exp'], 'verify_aud': False, 'verify_iss': False},
+ )
+ account_obj: user.User | None = None
+ account_uuid = payload.get('sub')
+ if isinstance(account_uuid, str) and account_uuid:
+ try:
+ account_obj = await self.get_user_by_uuid(account_uuid)
+ except AttributeError:
+ account_obj = None
+ if account_obj is None:
+ legacy_email = payload.get('user')
+ if not isinstance(legacy_email, str) or not legacy_email:
+ raise ValueError('JWT Account identity is missing')
+ try:
+ account_obj = await self.get_user_by_email(legacy_email)
+ except AttributeError:
+ account_obj = None
+ if account_obj is None and allow_unresolved_legacy:
+ return legacy_email
+ if account_obj is None:
+ raise ValueError('Account not found')
+ self._require_active_account(account_obj)
+ token_revision = payload.get('account_revision')
+ if token_revision is not None and int(token_revision) != account_obj.projection_revision:
+ raise ValueError('Account token revision is stale')
+ return account_obj
+
+ @staticmethod
+ def _require_active_account(account: user.User) -> None:
+ status = getattr(account, 'status', user.AccountStatus.ACTIVE.value)
+ if isinstance(status, str) and status != user.AccountStatus.ACTIVE.value:
+ raise AccountDisabledError('Account is disabled')
async def reset_password(self, user_email: str, new_password: str) -> None:
hashed_password = await self._hash_password(new_password)
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(user.User).where(user.User.user == user_email).values(password=hashed_password)
+ sqlalchemy.update(user.User)
+ .where(user.User.normalized_email == normalize_email(user_email))
+ .values(password=hashed_password)
)
async def change_password(self, user_email: str, current_password: str, new_password: str) -> None:
@@ -117,7 +409,9 @@ class UserService:
hashed_password = await self._hash_password(new_password)
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(user.User).where(user.User.user == user_email).values(password=hashed_password)
+ sqlalchemy.update(user.User)
+ .where(user.User.normalized_email == normalize_email(user_email))
+ .values(password=hashed_password)
)
# Space user management
@@ -132,6 +426,7 @@ class UserService:
expires_in: int = 0,
) -> user.User:
"""Create or update a Space user account (only if system not initialized or user exists)"""
+ self._require_local_directory()
expires_at = datetime.datetime.now() + datetime.timedelta(seconds=expires_in) if expires_in > 0 else None
async with self._create_user_lock:
@@ -150,17 +445,53 @@ class UserService:
space_access_token_expires_at=expires_at,
)
)
- await self.ap.provider_service.update_space_model_provider_api_keys(api_key)
+ await self._update_space_provider_for_account(existing_user, api_key)
return await self.get_user_by_space_account_uuid(space_account_uuid)
# Check if user with same email exists
existing_email_user = await self.get_user_by_email(email)
if existing_email_user:
- # Update existing user to link with Space account
+ # Email is display/contact identity, not an OAuth subject. An
+ # unknown Space subject must never take over an existing local
+ # Account merely by presenting the same email. The Account
+ # owner must first authenticate locally and use the explicit,
+ # account-bound bind flow.
+ raise account_errors.AccountEmailMismatchError()
+
+ # Check if system is already initialized
+ is_initialized = await self.is_initialized()
+ if is_initialized:
+ raise account_errors.AccountEmailMismatchError()
+
+ # Create new Space user (first time initialization)
+ if hasattr(self.ap.persistence_mgr, 'get_db_engine') and hasattr(self.ap, 'workspace_service'):
+ async with self._session_factory()() as session:
+ async with session.begin():
+ account = user.User(
+ uuid=str(uuid.uuid4()),
+ user=normalize_email(email),
+ normalized_email=normalize_email(email),
+ password='',
+ account_type='space',
+ status=user.AccountStatus.ACTIVE.value,
+ source=user.AccountSource.LOCAL.value,
+ projection_revision=0,
+ space_account_uuid=space_account_uuid,
+ space_access_token=access_token,
+ space_refresh_token=refresh_token,
+ space_api_key=api_key,
+ space_access_token_expires_at=expires_at,
+ )
+ session.add(account)
+ await session.flush()
+ await self.ap.workspace_service.bootstrap_local_account(account.uuid, session=session)
+ else:
+ # Compatibility path for lightweight service tests without a real engine.
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(user.User)
- .where(user.User.user == email)
- .values(
+ sqlalchemy.insert(user.User).values(
+ user=normalize_email(email),
+ normalized_email=normalize_email(email),
+ password='',
account_type='space',
space_account_uuid=space_account_uuid,
space_access_token=access_token,
@@ -169,30 +500,10 @@ class UserService:
space_access_token_expires_at=expires_at,
)
)
- await self.ap.provider_service.update_space_model_provider_api_keys(api_key)
- return await self.get_user_by_email(email)
-
- # Check if system is already initialized
- is_initialized = await self.is_initialized()
- if is_initialized:
- raise account_errors.AccountEmailMismatchError()
-
- # Create new Space user (first time initialization)
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.insert(user.User).values(
- user=email,
- password='', # Space users don't have local password
- account_type='space',
- space_account_uuid=space_account_uuid,
- space_access_token=access_token,
- space_refresh_token=refresh_token,
- space_api_key=api_key,
- space_access_token_expires_at=expires_at,
- )
- )
- await self.ap.provider_service.update_space_model_provider_api_keys(api_key)
-
- return await self.get_user_by_space_account_uuid(space_account_uuid)
+ created_user = await self.get_user_by_space_account_uuid(space_account_uuid)
+ if created_user is not None:
+ await self._update_space_provider_for_account(created_user, api_key)
+ return created_user
async def authenticate_space_user(
self, access_token: str, refresh_token: str, expires_in: int = 0
@@ -221,7 +532,7 @@ class UserService:
)
# Generate JWT token
- jwt_token = await self.generate_jwt_token(email)
+ jwt_token = await self.generate_jwt_token(user_obj)
return jwt_token, user_obj
@@ -247,11 +558,16 @@ class UserService:
hashed_password = await self._hash_password(new_password)
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(user.User).where(user.User.user == user_email).values(password=hashed_password)
+ sqlalchemy.update(user.User)
+ .where(user.User.normalized_email == normalize_email(user_email))
+ .values(password=hashed_password)
)
async def bind_space_account(self, user_email: str, code: str) -> user.User:
"""Bind Space account to existing local account"""
+ local_account = await self.get_user_by_email(user_email)
+ if local_account is None:
+ raise ValueError('User not found')
# Exchange code for tokens
token_data = await self.ap.space_service.exchange_oauth_code(code)
access_token = token_data.get('access_token')
@@ -276,15 +592,16 @@ class UserService:
# Check if this Space account is already bound to another user
existing_space_user = await self.get_user_by_space_account_uuid(space_account_uuid)
- if existing_space_user and existing_space_user.user != user_email:
+ if existing_space_user and existing_space_user.normalized_email != normalize_email(user_email):
raise ValueError('This Space account is already bound to another user')
# Update local account to Space account
await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(user.User)
- .where(user.User.user == user_email)
+ .where(user.User.normalized_email == normalize_email(user_email))
.values(
- user=space_email, # Update email to Space email
+ user=normalize_email(space_email), # Update email to Space email
+ normalized_email=normalize_email(space_email),
account_type='space',
space_account_uuid=space_account_uuid,
space_access_token=access_token,
@@ -295,6 +612,6 @@ class UserService:
)
# Update Space model provider API keys
- await self.ap.provider_service.update_space_model_provider_api_keys(api_key)
+ await self._update_space_provider_for_account(local_account, api_key)
return await self.get_user_by_email(space_email)
diff --git a/src/langbot/pkg/api/http/service/webhook.py b/src/langbot/pkg/api/http/service/webhook.py
index b3a671189..14b029864 100644
--- a/src/langbot/pkg/api/http/service/webhook.py
+++ b/src/langbot/pkg/api/http/service/webhook.py
@@ -4,6 +4,8 @@ import sqlalchemy
from ....core import app
from ....entity.persistence import webhook
+from .secrets import SECRET_MASK, mask_secret_value, restore_secret_placeholders
+from .tenant import TenantContext, require_workspace_uuid, scope_statement
class WebhookService:
@@ -12,31 +14,71 @@ class WebhookService:
def __init__(self, ap: app.Application) -> None:
self.ap = ap
- async def get_webhooks(self) -> list[dict]:
+ def _serialize_webhook(self, entity, *, include_secret: bool) -> dict:
+ serialized = self.ap.persistence_mgr.serialize_model(webhook.Webhook, entity)
+ if not include_secret:
+ serialized = serialized.copy()
+ serialized['url'] = mask_secret_value(serialized.get('url'))
+ return serialized
+
+ async def get_webhooks(self, context: TenantContext, *, include_secret: bool = False) -> list[dict]:
"""Get all webhooks"""
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(webhook.Webhook))
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(sqlalchemy.select(webhook.Webhook), webhook.Webhook, context)
+ )
webhooks = result.all()
- return [self.ap.persistence_mgr.serialize_model(webhook.Webhook, wh) for wh in webhooks]
+ return [self._serialize_webhook(wh, include_secret=include_secret) for wh in webhooks]
- async def create_webhook(self, name: str, url: str, description: str = '', enabled: bool = True) -> dict:
+ async def create_webhook(
+ self,
+ context: TenantContext,
+ name: str,
+ url: str,
+ description: str = '',
+ enabled: bool = True,
+ ) -> dict:
"""Create a new webhook"""
- webhook_data = {'name': name, 'url': url, 'description': description, 'enabled': enabled}
+ workspace_uuid = require_workspace_uuid(context)
+ url = restore_secret_placeholders(url, sensitive=True)
+ webhook_data = {
+ 'workspace_uuid': workspace_uuid,
+ 'name': name,
+ 'url': url,
+ 'description': description,
+ 'enabled': enabled,
+ }
- await self.ap.persistence_mgr.execute_async(sqlalchemy.insert(webhook.Webhook).values(**webhook_data))
+ insert_result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.insert(webhook.Webhook).values(**webhook_data)
+ )
# Retrieve the created webhook
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(webhook.Webhook).where(webhook.Webhook.url == url).order_by(webhook.Webhook.id.desc())
+ scope_statement(
+ sqlalchemy.select(webhook.Webhook).where(webhook.Webhook.id == insert_result.inserted_primary_key[0]),
+ webhook.Webhook,
+ workspace_uuid,
+ )
)
created_webhook = result.first()
return self.ap.persistence_mgr.serialize_model(webhook.Webhook, created_webhook)
- async def get_webhook(self, webhook_id: int) -> dict | None:
+ async def get_webhook(
+ self,
+ context: TenantContext,
+ webhook_id: int,
+ *,
+ include_secret: bool = False,
+ ) -> dict | None:
"""Get a specific webhook by ID"""
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(webhook.Webhook).where(webhook.Webhook.id == webhook_id)
+ scope_statement(
+ sqlalchemy.select(webhook.Webhook).where(webhook.Webhook.id == webhook_id),
+ webhook.Webhook,
+ context,
+ )
)
wh = result.first()
@@ -44,16 +86,27 @@ class WebhookService:
if wh is None:
return None
- return self.ap.persistence_mgr.serialize_model(webhook.Webhook, wh)
+ return self._serialize_webhook(wh, include_secret=include_secret)
async def update_webhook(
- self, webhook_id: int, name: str = None, url: str = None, description: str = None, enabled: bool = None
- ) -> None:
+ self,
+ context: TenantContext,
+ webhook_id: int,
+ name: str | None = None,
+ url: str | None = None,
+ description: str | None = None,
+ enabled: bool | None = None,
+ ) -> bool:
"""Update a webhook's metadata"""
update_data = {}
if name is not None:
update_data['name'] = name
if url is not None:
+ if url == SECRET_MASK:
+ current = await self.get_webhook(context, webhook_id, include_secret=True)
+ if current is None:
+ return False
+ url = restore_secret_placeholders(url, current.get('url'), sensitive=True)
update_data['url'] = url
if description is not None:
update_data['description'] = description
@@ -61,20 +114,35 @@ class WebhookService:
update_data['enabled'] = enabled
if update_data:
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(webhook.Webhook).where(webhook.Webhook.id == webhook_id).values(**update_data)
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.update(webhook.Webhook).where(webhook.Webhook.id == webhook_id).values(**update_data),
+ webhook.Webhook,
+ context,
+ )
)
+ return (result.rowcount or 0) > 0
+ return await self.get_webhook(context, webhook_id) is not None
- async def delete_webhook(self, webhook_id: int) -> None:
+ async def delete_webhook(self, context: TenantContext, webhook_id: int) -> bool:
"""Delete a webhook"""
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(webhook.Webhook).where(webhook.Webhook.id == webhook_id)
+ result = await self.ap.persistence_mgr.execute_async(
+ scope_statement(
+ sqlalchemy.delete(webhook.Webhook).where(webhook.Webhook.id == webhook_id),
+ webhook.Webhook,
+ context,
+ )
)
+ return (result.rowcount or 0) > 0
- async def get_enabled_webhooks(self) -> list[dict]:
+ async def get_enabled_webhooks(self, context: TenantContext) -> list[dict]:
"""Get all enabled webhooks"""
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(webhook.Webhook).where(webhook.Webhook.enabled == True)
+ scope_statement(
+ sqlalchemy.select(webhook.Webhook).where(webhook.Webhook.enabled == True),
+ webhook.Webhook,
+ context,
+ )
)
webhooks = result.all()
diff --git a/src/langbot/pkg/api/mcp/context.py b/src/langbot/pkg/api/mcp/context.py
new file mode 100644
index 000000000..56fbe0ea1
--- /dev/null
+++ b/src/langbot/pkg/api/mcp/context.py
@@ -0,0 +1,30 @@
+from __future__ import annotations
+
+import contextvars
+
+from ..http.context import RequestContext
+
+
+_request_context: contextvars.ContextVar[RequestContext | None] = contextvars.ContextVar(
+ 'langbot_mcp_request_context',
+ default=None,
+)
+
+
+def bind_request_context(context: RequestContext) -> contextvars.Token[RequestContext | None]:
+ """Bind the authenticated MCP request while its ASGI request is executing."""
+
+ return _request_context.set(context)
+
+
+def reset_request_context(token: contextvars.Token[RequestContext | None]) -> None:
+ _request_context.reset(token)
+
+
+def get_request_context() -> RequestContext:
+ """Return the current trusted MCP context or fail closed."""
+
+ context = _request_context.get()
+ if context is None:
+ raise RuntimeError('MCP Workspace context is unavailable')
+ return context
diff --git a/src/langbot/pkg/api/mcp/mount.py b/src/langbot/pkg/api/mcp/mount.py
index d113b2501..d3eb2c76c 100644
--- a/src/langbot/pkg/api/mcp/mount.py
+++ b/src/langbot/pkg/api/mcp/mount.py
@@ -19,7 +19,10 @@ from __future__ import annotations
import contextlib
import typing
+import uuid
+from ..http.context import PrincipalContext, PrincipalType, RequestContext, WorkspaceContext
+from .context import bind_request_context, reset_request_context
from .server import LangBotMCPServer
if typing.TYPE_CHECKING:
@@ -76,7 +79,7 @@ class MCPMount:
def wrap(self, quart_asgi: typing.Callable) -> typing.Callable:
"""Return a dispatcher ASGI app fronting ``quart_asgi``."""
mcp_asgi = self._mcp_asgi
- verify_api_key = self.ap.apikey_service.verify_api_key
+ authenticate_api_key = self.ap.apikey_service.authenticate_api_key
is_mcp_path = self._is_mcp_path
async def dispatcher(scope, receive, send): # type: ignore[no-untyped-def]
@@ -88,12 +91,12 @@ class MCPMount:
# Authenticate MCP HTTP requests with a LangBot API key.
api_key = _extract_api_key(scope.get('headers', []))
- authorized = False
+ identity = None
if api_key:
with contextlib.suppress(Exception):
- authorized = await verify_api_key(api_key)
+ identity = await authenticate_api_key(api_key)
- if not authorized:
+ if identity is None:
await send(
{
'type': 'http.response.start',
@@ -107,6 +110,26 @@ class MCPMount:
await send({'type': 'http.response.body', 'body': _UNAUTHORIZED_BODY})
return
- await mcp_asgi(scope, receive, send)
+ request_context = RequestContext(
+ instance_uuid=identity.instance_uuid,
+ placement_generation=identity.placement_generation,
+ request_id=str(uuid.uuid4()),
+ auth_type='api-key',
+ principal=PrincipalContext(
+ principal_type=PrincipalType.API_KEY,
+ api_key_uuid=identity.api_key_uuid,
+ ),
+ workspace=WorkspaceContext(
+ workspace_uuid=identity.workspace_uuid,
+ membership_uuid=None,
+ role=None,
+ permissions=identity.permissions,
+ ),
+ )
+ token = bind_request_context(request_context)
+ try:
+ await mcp_asgi(scope, receive, send)
+ finally:
+ reset_request_context(token)
return dispatcher
diff --git a/src/langbot/pkg/api/mcp/server.py b/src/langbot/pkg/api/mcp/server.py
index 95630bbaf..4cf4e33ad 100644
--- a/src/langbot/pkg/api/mcp/server.py
+++ b/src/langbot/pkg/api/mcp/server.py
@@ -22,6 +22,9 @@ import typing
from mcp.server.fastmcp import FastMCP
+from ..http.authz import Permission, require_permission
+from .context import get_request_context
+
if typing.TYPE_CHECKING:
from ...core import app as app_module
@@ -46,6 +49,12 @@ def _dump(value: typing.Any) -> str:
return json.dumps(value, ensure_ascii=False, default=str)
+def _authorized(permission: Permission):
+ context = get_request_context()
+ require_permission(context, permission)
+ return context
+
+
class LangBotMCPServer:
"""Builds and owns the FastMCP instance for LangBot."""
@@ -72,6 +81,7 @@ class LangBotMCPServer:
# ----- System (read-only) -------------------------------------- #
@mcp.tool(description='Get basic LangBot system/runtime information (version, edition).')
async def get_system_info() -> str:
+ _authorized(Permission.WORKSPACE_VIEW)
version = None
try:
version = ap.ver_mgr.get_current_version()
@@ -87,11 +97,13 @@ class LangBotMCPServer:
# ----- Bots ---------------------------------------------------- #
@mcp.tool(description='List all messaging-platform bots. Secrets are redacted.')
async def list_bots() -> str:
- return _dump(await ap.bot_service.get_bots(include_secret=False))
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.bot_service.get_bots(context, include_secret=False))
@mcp.tool(description='Get a single bot by its UUID. Secrets are redacted.')
async def get_bot(bot_uuid: str) -> str:
- return _dump(await ap.bot_service.get_bot(bot_uuid, include_secret=False))
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.bot_service.get_bot(context, bot_uuid, include_secret=False))
@mcp.tool(
description=(
@@ -101,26 +113,31 @@ class LangBotMCPServer:
)
)
async def create_bot(bot_data: dict) -> str:
- return _dump({'uuid': await ap.bot_service.create_bot(bot_data)})
+ context = _authorized(Permission.RESOURCE_MANAGE)
+ return _dump({'uuid': await ap.bot_service.create_bot(context, bot_data)})
@mcp.tool(description='Update a bot by UUID. `bot_data` matches the PUT bot body.')
async def update_bot(bot_uuid: str, bot_data: dict) -> str:
- await ap.bot_service.update_bot(bot_uuid, bot_data)
+ context = _authorized(Permission.RESOURCE_MANAGE)
+ await ap.bot_service.update_bot(context, bot_uuid, bot_data)
return _dump({'ok': True})
@mcp.tool(description='Delete a bot by UUID.')
async def delete_bot(bot_uuid: str) -> str:
- await ap.bot_service.delete_bot(bot_uuid)
+ context = _authorized(Permission.RESOURCE_MANAGE)
+ await ap.bot_service.delete_bot(context, bot_uuid)
return _dump({'ok': True})
# ----- Pipelines ----------------------------------------------- #
@mcp.tool(description='List all pipelines.')
async def list_pipelines() -> str:
- return _dump(await ap.pipeline_service.get_pipelines())
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.pipeline_service.get_pipelines(context))
@mcp.tool(description='Get a single pipeline by UUID.')
async def get_pipeline(pipeline_uuid: str) -> str:
- return _dump(await ap.pipeline_service.get_pipeline(pipeline_uuid))
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.pipeline_service.get_pipeline(context, pipeline_uuid))
@mcp.tool(
description=(
@@ -129,49 +146,59 @@ class LangBotMCPServer:
)
)
async def create_pipeline(pipeline_data: dict) -> str:
- return _dump({'uuid': await ap.pipeline_service.create_pipeline(pipeline_data)})
+ context = _authorized(Permission.RESOURCE_MANAGE)
+ return _dump({'uuid': await ap.pipeline_service.create_pipeline(context, pipeline_data)})
@mcp.tool(description='Update a pipeline by UUID. `pipeline_data` matches the PUT body.')
async def update_pipeline(pipeline_uuid: str, pipeline_data: dict) -> str:
- await ap.pipeline_service.update_pipeline(pipeline_uuid, pipeline_data)
+ context = _authorized(Permission.RESOURCE_MANAGE)
+ await ap.pipeline_service.update_pipeline(context, pipeline_uuid, pipeline_data)
return _dump({'ok': True})
@mcp.tool(description='Delete a pipeline by UUID.')
async def delete_pipeline(pipeline_uuid: str) -> str:
- await ap.pipeline_service.delete_pipeline(pipeline_uuid)
+ context = _authorized(Permission.RESOURCE_MANAGE)
+ await ap.pipeline_service.delete_pipeline(context, pipeline_uuid)
return _dump({'ok': True})
# ----- Models -------------------------------------------------- #
@mcp.tool(description='List all configured LLM models. Secrets are redacted.')
async def list_llm_models() -> str:
- return _dump(await ap.llm_model_service.get_llm_models(include_secret=False))
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.llm_model_service.get_llm_models(context, include_secret=False))
@mcp.tool(description='Get a single LLM model by UUID.')
async def get_llm_model(model_uuid: str) -> str:
- return _dump(await ap.llm_model_service.get_llm_model(model_uuid))
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.llm_model_service.get_llm_model(context, model_uuid, include_secret=False))
@mcp.tool(description='List all configured embedding models.')
async def list_embedding_models() -> str:
- return _dump(await ap.embedding_models_service.get_embedding_models())
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.embedding_models_service.get_embedding_models(context, include_secret=False))
@mcp.tool(description='List all model providers (OpenAI-compatible, Anthropic, etc.).')
async def list_model_providers() -> str:
- return _dump(await ap.provider_service.get_providers())
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.provider_service.get_providers(context, include_secret=False))
# ----- Knowledge bases ----------------------------------------- #
@mcp.tool(description='List all knowledge bases (RAG).')
async def list_knowledge_bases() -> str:
- return _dump(await ap.knowledge_service.get_knowledge_bases())
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.knowledge_service.get_knowledge_bases(context))
@mcp.tool(description='Get a single knowledge base by UUID.')
async def get_knowledge_base(kb_uuid: str) -> str:
- return _dump(await ap.knowledge_service.get_knowledge_base(kb_uuid))
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.knowledge_service.get_knowledge_base(context, kb_uuid))
@mcp.tool(
description=('Retrieve (semantic search) from a knowledge base. Returns the matched chunks for `query`.')
)
async def retrieve_knowledge_base(kb_uuid: str, query: str) -> str:
- return _dump(await ap.knowledge_service.retrieve_knowledge_base(kb_uuid, query))
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.knowledge_service.retrieve_knowledge_base(context, kb_uuid, query))
# ----- MCP servers (LangBot as MCP client) --------------------- #
@mcp.tool(
@@ -180,16 +207,19 @@ class LangBotMCPServer:
)
)
async def list_mcp_servers() -> str:
- return _dump(await ap.mcp_service.get_mcp_servers())
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.mcp_service.get_mcp_servers(context))
# ----- Skills -------------------------------------------------- #
@mcp.tool(description='List installed skills.')
async def list_skills() -> str:
- return _dump(await ap.skill_service.list_skills())
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.skill_service.list_skills(context))
@mcp.tool(description='Get a single skill by name.')
async def get_skill(skill_name: str) -> str:
- return _dump(await ap.skill_service.get_skill(skill_name))
+ context = _authorized(Permission.RESOURCE_VIEW)
+ return _dump(await ap.skill_service.get_skill(context, skill_name))
# ------------------------------------------------------------------ #
# ASGI app
diff --git a/src/langbot/pkg/box/connector.py b/src/langbot/pkg/box/connector.py
index 2257910d1..ced6f015f 100644
--- a/src/langbot/pkg/box/connector.py
+++ b/src/langbot/pkg/box/connector.py
@@ -3,6 +3,7 @@ from __future__ import annotations
import asyncio
import json
import os
+import secrets
import sys
import typing
from typing import TYPE_CHECKING
@@ -15,6 +16,17 @@ from langbot_plugin.runtime.io.connection import Connection
from langbot_plugin.box.client import ActionRPCBoxClient
from langbot_plugin.box.errors import BoxRuntimeUnavailableError
from langbot_plugin.box.actions import LangBotToBoxAction
+from langbot_plugin.box.security import (
+ BOX_CONTROL_TOKEN_ENV,
+ BOX_CONTROL_TOKEN_HEADER,
+ BOX_INSTANCE_HEADER,
+ BOX_PLACEMENT_GENERATION_HEADER,
+ BOX_TRUSTED_INSTANCE_ENV,
+ BOX_WORKSPACE_HEADER,
+ normalize_instance_uuid,
+ validate_control_token,
+)
+from langbot_plugin.entities.io.context import ActionContext
from ..utils import platform
from ..utils.managed_runtime import ManagedRuntimeConnector
@@ -119,6 +131,8 @@ class BoxRuntimeConnector(ManagedRuntimeConnector):
self._relay_host = parsed.hostname or '127.0.0.1'
self._relay_port = parsed.port or _DEFAULT_PORT
self._filtered_box_config = _filter_config_for_runtime(_get_box_config(ap))
+ self._trusted_instance_uuid = normalize_instance_uuid(self.ap.workspace_service.instance_uuid)
+ self._control_token = str(os.environ.get(BOX_CONTROL_TOKEN_ENV) or '').strip()
def uses_websocket(self) -> bool:
"""Whether the connector should use WebSocket to reach the Box runtime.
@@ -181,8 +195,11 @@ class BoxRuntimeConnector(ManagedRuntimeConnector):
from langbot_plugin.runtime.io.controllers.stdio.client import StdioClientController
self.ap.logger.info('Use stdio to connect to box runtime')
+ self._ensure_control_token(allow_generate=True)
python_path = sys.executable
env = os.environ.copy()
+ env[BOX_CONTROL_TOKEN_ENV] = self._control_token
+ env[BOX_TRUSTED_INSTANCE_ENV] = self._trusted_instance_uuid
if self._filtered_box_config:
env['LANGBOT_BOX_CONFIG'] = json.dumps(self._filtered_box_config)
@@ -215,7 +232,10 @@ class BoxRuntimeConnector(ManagedRuntimeConnector):
"""Launch box server as detached subprocess, then connect via WS (Windows)."""
self.ap.logger.info('(windows) Use cmd to launch box runtime and communicate via ws')
+ self._ensure_control_token(allow_generate=True)
env = os.environ.copy()
+ env[BOX_CONTROL_TOKEN_ENV] = self._control_token
+ env[BOX_TRUSTED_INSTANCE_ENV] = self._trusted_instance_uuid
if self._filtered_box_config:
env['LANGBOT_BOX_CONFIG'] = json.dumps(self._filtered_box_config)
@@ -238,6 +258,7 @@ class BoxRuntimeConnector(ManagedRuntimeConnector):
async def _connect_remote_ws(self) -> None:
"""Connect to a remote (or Docker) box server via WebSocket."""
+ self._ensure_control_token(allow_generate=False)
ws_url = self._resolve_rpc_ws_url()
self.ap.logger.info(f'Use WebSocket to connect to box runtime ({ws_url})')
await self._connect_ws(ws_url, 'WebSocket')
@@ -281,7 +302,11 @@ class BoxRuntimeConnector(ManagedRuntimeConnector):
if self.runtime_disconnect_callback is not None:
await self.runtime_disconnect_callback(self)
- ctrl = WebSocketClientController(ws_url=ws_url, make_connection_failed_callback=on_connect_failed)
+ ctrl = WebSocketClientController(
+ ws_url=ws_url,
+ make_connection_failed_callback=on_connect_failed,
+ additional_headers=self.get_control_headers(),
+ )
self._ctrl_task = asyncio.create_task(
ctrl.run(self._make_connection_callback(transport_name, connected, connect_error))
)
@@ -294,6 +319,41 @@ class BoxRuntimeConnector(ManagedRuntimeConnector):
if connect_error:
raise BoxRuntimeUnavailableError(f'box runtime connection failed: {connect_error[0]}')
+ def _ensure_control_token(self, *, allow_generate: bool) -> str:
+ if not self._control_token and allow_generate:
+ self._control_token = secrets.token_urlsafe(48)
+ try:
+ self._control_token = validate_control_token(self._control_token)
+ except ValueError as exc:
+ raise BoxRuntimeUnavailableError(
+ f'{BOX_CONTROL_TOKEN_ENV} must be configured with a strong shared secret for an external Box runtime'
+ ) from exc
+ return self._control_token
+
+ def get_control_headers(self) -> dict[str, str]:
+ """Headers for the instance-authenticated RPC control handshake."""
+
+ self._ensure_control_token(allow_generate=False)
+ return {
+ BOX_CONTROL_TOKEN_HEADER: self._control_token,
+ BOX_INSTANCE_HEADER: self._trusted_instance_uuid,
+ }
+
+ def get_relay_headers(
+ self,
+ action_context: ActionContext,
+ ) -> dict[str, str]:
+ """Return authenticated, placement-scoped relay handshake headers."""
+
+ context = ActionContext.model_validate(action_context).without_installation()
+ if context.instance_uuid != self._trusted_instance_uuid:
+ raise BoxRuntimeUnavailableError('Box relay context belongs to another LangBot instance')
+ return {
+ **self.get_control_headers(),
+ BOX_WORKSPACE_HEADER: context.workspace_uuid,
+ BOX_PLACEMENT_GENERATION_HEADER: str(context.placement_generation),
+ }
+
def _make_connection_callback(
self,
transport_name: str,
diff --git a/src/langbot/pkg/box/service.py b/src/langbot/pkg/box/service.py
index ca30eb930..199ad61cf 100644
--- a/src/langbot/pkg/box/service.py
+++ b/src/langbot/pkg/box/service.py
@@ -11,8 +11,12 @@ from typing import TYPE_CHECKING
import pydantic
from langbot_plugin.box.client import BoxRuntimeClient
+from langbot_plugin.entities.io.context import ActionContext
+from langbot_plugin.box.tenancy import box_namespace
from .connector import BoxRuntimeConnector, _get_box_config
from ..telemetry import features as telemetry_features
+from ..api.http.context import ExecutionContext
+from ..api.http.service.tenant import TenantContext, require_workspace_uuid
from langbot_plugin.box.errors import BoxError, BoxValidationError
from langbot_plugin.box.models import (
BUILTIN_PROFILES,
@@ -180,6 +184,69 @@ class BoxService:
return False
return not self._runtime_connector.uses_websocket()
+ @staticmethod
+ def _execution_context(context: TenantContext) -> ExecutionContext:
+ workspace_uuid = require_workspace_uuid(context)
+ instance_uuid = str(getattr(context, 'instance_uuid', '') or '').strip()
+ generation = getattr(context, 'placement_generation', None)
+ if not instance_uuid:
+ raise BoxValidationError('Box operations require an explicit instance UUID')
+ if isinstance(generation, bool) or not isinstance(generation, int) or generation <= 0:
+ raise BoxValidationError('Box operations require a positive placement generation')
+ return ExecutionContext(
+ instance_uuid=instance_uuid,
+ workspace_uuid=workspace_uuid,
+ placement_generation=generation,
+ bot_uuid=getattr(context, 'bot_uuid', None),
+ pipeline_uuid=getattr(context, 'pipeline_uuid', None),
+ query_uuid=getattr(context, 'query_uuid', None),
+ )
+
+ @classmethod
+ def _query_execution_context(cls, query: pipeline_query.Query) -> ExecutionContext:
+ return cls._execution_context(
+ ExecutionContext(
+ instance_uuid=str(getattr(query, 'instance_uuid', '') or ''),
+ workspace_uuid=str(getattr(query, 'workspace_uuid', '') or ''),
+ placement_generation=getattr(query, 'placement_generation', 0) or 0,
+ bot_uuid=getattr(query, 'bot_uuid', None),
+ pipeline_uuid=getattr(query, 'pipeline_uuid', None),
+ query_uuid=getattr(query, 'query_uuid', None),
+ )
+ )
+
+ @classmethod
+ def _action_context(cls, context: TenantContext) -> ActionContext:
+ execution_context = cls._execution_context(context)
+ return ActionContext(
+ instance_uuid=execution_context.instance_uuid,
+ workspace_uuid=execution_context.workspace_uuid,
+ placement_generation=execution_context.placement_generation,
+ )
+
+ async def _validated_execution_context(self, context: TenantContext) -> ExecutionContext:
+ """Resolve and fence a tenant context before touching shared Box state."""
+
+ execution_context = self._execution_context(context)
+ binding = await self.ap.workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ if binding.instance_uuid != execution_context.instance_uuid:
+ raise BoxValidationError('Box execution context belongs to another LangBot instance')
+ if (
+ str(getattr(binding, 'workspace_uuid', '') or '') != execution_context.workspace_uuid
+ or getattr(binding, 'placement_generation', None) != execution_context.placement_generation
+ ):
+ raise BoxValidationError('Box execution context belongs to a stale Workspace placement')
+ return execution_context
+
+ def _tenant_workspace(self, context: TenantContext) -> str | None:
+ if self.default_workspace is None:
+ return None
+ namespace = box_namespace(self._action_context(context))
+ return os.path.join(self.default_workspace, 'tenants', namespace)
+
async def execute_spec_payload(
self,
spec_payload: dict,
@@ -189,6 +256,14 @@ class BoxService:
) -> dict:
if not self._available:
raise BoxError('Box runtime is not available. Install and start Docker to use sandbox features.')
+ execution_context = await self._validated_execution_context(self._query_execution_context(query))
+ spec_payload = dict(spec_payload)
+ if spec_payload.get('host_path') in (None, ''):
+ tenant_workspace = self._tenant_workspace(execution_context)
+ if tenant_workspace is not None:
+ spec_payload['host_path'] = tenant_workspace
+ if self.shares_filesystem_with_box:
+ os.makedirs(tenant_workspace, exist_ok=True)
try:
spec = self.build_spec(spec_payload, skip_host_mount_validation=skip_host_mount_validation)
except BoxError as exc:
@@ -205,14 +280,17 @@ class BoxService:
self._record_error(exc, query)
raise
try:
- result = await self.client.execute(spec)
+ result = await self.client.execute(
+ spec,
+ action_context=self._action_context(execution_context),
+ )
except BoxError as exc:
self._record_error(exc, query)
raise
try:
await self._enforce_workspace_quota(spec, phase='after execution')
except BoxError as exc:
- await self._cleanup_exceeded_session(spec)
+ await self._cleanup_exceeded_session(execution_context, spec)
self._record_error(exc, query)
raise
self.ap.logger.info(
@@ -336,6 +414,26 @@ class BoxService:
return await self.execute_spec_payload(spec_payload, query)
+ async def execute_in_context(
+ self,
+ context: TenantContext,
+ spec_payload: dict,
+ *,
+ skip_host_mount_validation: bool = False,
+ ) -> BoxExecutionResult:
+ """Execute trusted internal Box work inside one Workspace namespace."""
+
+ execution_context = await self._validated_execution_context(context)
+ payload = dict(spec_payload)
+ if payload.get('host_path') in (None, ''):
+ tenant_workspace = self._tenant_workspace(execution_context)
+ if tenant_workspace is not None:
+ payload['host_path'] = tenant_workspace
+ if self.shares_filesystem_with_box:
+ os.makedirs(tenant_workspace, exist_ok=True)
+ spec = self.build_spec(payload, skip_host_mount_validation=skip_host_mount_validation)
+ return await self.client.execute(spec, action_context=self._action_context(execution_context))
+
# ── Attachment passthrough (inbound / outbound) ──────────────────
#
# IM/webchat attachments (images, voices, files) reach the LLM as
@@ -368,7 +466,7 @@ class BoxService:
# truncation). The host-filesystem path has no such limit.
_EXEC_FALLBACK_MAX_BYTES = 256 * 1024
- def _host_query_dir(self, subdir: str, query_id) -> str | None:
+ def _host_query_dir(self, subdir: str, query: pipeline_query.Query) -> str | None:
"""Host path for ``/workspace//`` when LangBot can
access the bind-mounted workspace directly, else ``None``.
@@ -378,10 +476,10 @@ class BoxService:
to the sandbox (and vice-versa). It is ``None`` / not a local dir for
E2B and remote runtimes, where we must fall back to the exec channel.
"""
- root = self.default_workspace
+ root = self._tenant_workspace(self._query_execution_context(query))
if not root or not os.path.isdir(root):
return None
- return os.path.join(root, subdir, str(query_id))
+ return os.path.join(root, subdir, str(query.query_id))
async def _purge_attachment_dirs(self) -> None:
"""Remove leftover inbox/outbox directories on startup.
@@ -391,12 +489,10 @@ class BoxService:
a previous process would otherwise be silently reused — leaking a prior
run's inbound files and re-sending stale outbound files.
- Outbox files are written by the sandbox **container**, which runs as
- root over the bind-mount, so the LangBot host process (a non-root user)
- cannot ``rmtree`` them. We therefore try a host-side delete first (fast,
- works for host-owned inbox files) and, for anything that survives,
- delete from *inside* the sandbox via exec where the container's root can
- remove its own files. Best-effort: never block startup.
+ Tenant workspaces live below ``default_workspace/tenants``. Startup has
+ no authenticated Workspace context, so cleanup is deliberately limited
+ to direct host-filesystem deletion. It must never issue an unscoped Box
+ exec merely to remove root-owned container output.
"""
root = self.default_workspace
if not root or not os.path.isdir(root):
@@ -407,14 +503,29 @@ class BoxService:
host_survivors: list[str] = []
def _host_purge() -> list[str]:
+ candidates = [
+ os.path.join(root, self.INBOX_SUBDIR),
+ os.path.join(root, self.OUTBOX_SUBDIR),
+ ]
+ tenants_root = os.path.join(root, 'tenants')
+ if os.path.isdir(tenants_root):
+ with os.scandir(tenants_root) as tenant_entries:
+ for tenant_entry in tenant_entries:
+ if not tenant_entry.is_dir(follow_symlinks=False):
+ continue
+ candidates.extend(
+ [
+ os.path.join(tenant_entry.path, self.INBOX_SUBDIR),
+ os.path.join(tenant_entry.path, self.OUTBOX_SUBDIR),
+ ]
+ )
survivors: list[str] = []
- for subdir in (self.INBOX_SUBDIR, self.OUTBOX_SUBDIR):
- path = os.path.join(root, subdir)
+ for path in candidates:
if not os.path.isdir(path):
continue
shutil.rmtree(path, ignore_errors=True)
if os.path.exists(path):
- survivors.append(subdir)
+ survivors.append(path)
return survivors
try:
@@ -427,18 +538,11 @@ class BoxService:
self.ap.logger.info('Purged leftover sandbox attachment dirs from a previous process.')
return
- # Root-owned leftovers (container output): delete from inside the box.
- targets = ' '.join(f'/workspace/{sub}' for sub in host_survivors)
- try:
- spec = self.build_spec({'cmd': f'rm -rf {targets}', 'session_id': '__startup_purge__', 'timeout_sec': 30})
- await self.client.execute(spec)
- self.ap.logger.info(
- f'Purged root-owned leftover sandbox attachment dirs via sandbox exec: {host_survivors}'
- )
- except Exception as exc:
- self.ap.logger.warning(
- f'Failed to purge root-owned sandbox attachment dirs {host_survivors} via exec: {exc}'
- )
+ self.ap.logger.warning(
+ 'Could not purge root-owned sandbox attachment directories from the host; '
+ 'skipping an unsafe unscoped Box exec because startup has no trusted '
+ f'Workspace context: {host_survivors}'
+ )
@staticmethod
def _sanitize_attachment_name(name: str, fallback: str) -> str:
@@ -518,7 +622,7 @@ class BoxService:
if not files:
return []
- host_dir = self._host_query_dir(subdir, query.query_id)
+ host_dir = self._host_query_dir(subdir, query)
if host_dir is not None:
return await asyncio.to_thread(self._write_files_host, host_dir, target_mount_dir, files)
@@ -691,7 +795,7 @@ class BoxService:
if not self._available:
return []
- host_dir = self._host_query_dir(self.OUTBOX_SUBDIR, query.query_id)
+ host_dir = self._host_query_dir(self.OUTBOX_SUBDIR, query)
if host_dir is not None:
entries = await asyncio.to_thread(self._read_outbox_host, host_dir)
else:
@@ -847,11 +951,12 @@ class BoxService:
if loop is not None and not loop.is_closed() and (self._shutdown_task is None or self._shutdown_task.done()):
self._shutdown_task = loop.create_task(self.shutdown())
- async def get_sessions(self) -> list[dict]:
+ async def get_sessions(self, context: TenantContext) -> list[dict]:
+ execution_context = await self._validated_execution_context(context)
if not self._available:
return []
try:
- return await self.client.get_sessions()
+ return await self.client.get_sessions(action_context=self._action_context(execution_context))
except Exception:
return []
@@ -879,21 +984,70 @@ class BoxService:
self._validate_host_mount(spec)
return spec
- async def create_session(self, spec_payload: dict, *, skip_host_mount_validation: bool = False) -> dict:
+ async def create_session(
+ self,
+ context: TenantContext,
+ spec_payload: dict,
+ *,
+ skip_host_mount_validation: bool = False,
+ ) -> dict:
+ execution_context = await self._validated_execution_context(context)
+ spec_payload = dict(spec_payload)
+ if spec_payload.get('host_path') in (None, ''):
+ tenant_workspace = self._tenant_workspace(execution_context)
+ if tenant_workspace is not None:
+ spec_payload['host_path'] = tenant_workspace
+ if self.shares_filesystem_with_box:
+ os.makedirs(tenant_workspace, exist_ok=True)
spec = self.build_spec(spec_payload, skip_host_mount_validation=skip_host_mount_validation)
- return await self.client.create_session(spec)
+ return await self.client.create_session(spec, action_context=self._action_context(execution_context))
- async def start_managed_process(self, session_id: str, process_payload: dict) -> BoxManagedProcessInfo:
+ async def start_managed_process(
+ self,
+ context: TenantContext,
+ session_id: str,
+ process_payload: dict,
+ ) -> BoxManagedProcessInfo:
+ execution_context = await self._validated_execution_context(context)
process_spec = BoxManagedProcessSpec.model_validate(process_payload)
- return await self.client.start_managed_process(session_id, process_spec)
+ return await self.client.start_managed_process(
+ session_id,
+ process_spec,
+ action_context=self._action_context(execution_context),
+ )
- async def get_managed_process(self, session_id: str, process_id: str = 'default') -> BoxManagedProcessInfo:
- return await self.client.get_managed_process(session_id, process_id)
+ async def get_managed_process(
+ self,
+ context: TenantContext,
+ session_id: str,
+ process_id: str = 'default',
+ ) -> BoxManagedProcessInfo:
+ execution_context = await self._validated_execution_context(context)
+ return await self.client.get_managed_process(
+ session_id,
+ process_id,
+ action_context=self._action_context(execution_context),
+ )
- async def stop_managed_process(self, session_id: str, process_id: str = 'default') -> None:
- return await self.client.stop_managed_process(session_id, process_id)
+ async def stop_managed_process(
+ self,
+ context: TenantContext,
+ session_id: str,
+ process_id: str = 'default',
+ ) -> None:
+ execution_context = await self._validated_execution_context(context)
+ return await self.client.stop_managed_process(
+ session_id,
+ process_id,
+ action_context=self._action_context(execution_context),
+ )
- def get_managed_process_websocket_url(self, session_id: str, process_id: str = 'default') -> str:
+ def _get_managed_process_websocket_url(
+ self,
+ context: TenantContext,
+ session_id: str,
+ process_id: str = 'default',
+ ) -> str:
getter = getattr(self.client, 'get_managed_process_websocket_url', None)
if getter is None:
raise BoxValidationError('box runtime client does not support managed process websocket attach')
@@ -902,52 +1056,127 @@ class BoxService:
if self._runtime_connector is not None
else 'http://127.0.0.1:5410'
)
- return getter(session_id, ws_relay_base_url, process_id)
+ return getter(
+ session_id,
+ ws_relay_base_url,
+ process_id,
+ action_context=self._action_context(context),
+ )
- async def list_skills(self) -> list[dict]:
- return await self.client.list_skills()
+ async def get_managed_process_websocket_connection(
+ self,
+ context: TenantContext,
+ session_id: str,
+ process_id: str = 'default',
+ ) -> tuple[str, dict[str, str]]:
+ """Resolve a relay URL and headers after fencing the placement.
- async def get_skill(self, name: str) -> dict | None:
- return await self.client.get_skill(name)
+ The shared Box control secret is transported only in headers. The
+ Workspace and generation headers bind the relay to the same trusted
+ execution context used by the action RPC that created the process.
+ """
- async def create_skill(self, skill: dict) -> dict:
- return await self.client.create_skill(skill)
+ execution_context = await self._validated_execution_context(context)
+ if self._runtime_connector is None:
+ raise BoxValidationError(
+ 'box runtime connector does not support authenticated managed process websocket attach'
+ )
+ action_context = self._action_context(execution_context)
+ return (
+ self._get_managed_process_websocket_url(
+ execution_context,
+ session_id,
+ process_id,
+ ),
+ self._runtime_connector.get_relay_headers(action_context),
+ )
- async def update_skill(self, name: str, skill: dict) -> dict:
- return await self.client.update_skill(name, skill)
+ async def list_skills(self, context: TenantContext) -> list[dict]:
+ execution_context = await self._validated_execution_context(context)
+ return await self.client.list_skills(action_context=self._action_context(execution_context))
- async def delete_skill(self, name: str) -> None:
- await self.client.delete_skill(name)
+ async def get_skill(self, context: TenantContext, name: str) -> dict | None:
+ execution_context = await self._validated_execution_context(context)
+ return await self.client.get_skill(name, action_context=self._action_context(execution_context))
- async def scan_skill_directory(self, path: str) -> dict:
- return await self.client.scan_skill_directory(path)
+ async def create_skill(self, context: TenantContext, skill: dict) -> dict:
+ execution_context = await self._validated_execution_context(context)
+ payload = dict(skill)
+ payload.pop('workspace_uuid', None)
+ return await self.client.create_skill(payload, action_context=self._action_context(execution_context))
+
+ async def update_skill(self, context: TenantContext, name: str, skill: dict) -> dict:
+ execution_context = await self._validated_execution_context(context)
+ payload = dict(skill)
+ payload.pop('workspace_uuid', None)
+ return await self.client.update_skill(
+ name,
+ payload,
+ action_context=self._action_context(execution_context),
+ )
+
+ async def delete_skill(self, context: TenantContext, name: str) -> None:
+ execution_context = await self._validated_execution_context(context)
+ await self.client.delete_skill(name, action_context=self._action_context(execution_context))
+
+ async def scan_skill_directory(self, context: TenantContext, path: str) -> dict:
+ execution_context = await self._validated_execution_context(context)
+ return await self.client.scan_skill_directory(path, action_context=self._action_context(execution_context))
async def list_skill_files(
self,
+ context: TenantContext,
name: str,
path: str = '.',
include_hidden: bool = False,
max_entries: int = 200,
) -> dict:
- return await self.client.list_skill_files(name, path, include_hidden, max_entries)
+ execution_context = await self._validated_execution_context(context)
+ return await self.client.list_skill_files(
+ name,
+ path,
+ include_hidden,
+ max_entries,
+ action_context=self._action_context(execution_context),
+ )
- async def read_skill_file(self, name: str, path: str) -> dict:
- return await self.client.read_skill_file(name, path)
+ async def read_skill_file(self, context: TenantContext, name: str, path: str) -> dict:
+ execution_context = await self._validated_execution_context(context)
+ return await self.client.read_skill_file(
+ name,
+ path,
+ action_context=self._action_context(execution_context),
+ )
- async def write_skill_file(self, name: str, path: str, content: str) -> dict:
- return await self.client.write_skill_file(name, path, content)
+ async def write_skill_file(self, context: TenantContext, name: str, path: str, content: str) -> dict:
+ execution_context = await self._validated_execution_context(context)
+ return await self.client.write_skill_file(
+ name,
+ path,
+ content,
+ action_context=self._action_context(execution_context),
+ )
async def preview_skill_zip(
self,
+ context: TenantContext,
file_bytes: bytes,
filename: str,
source_subdir: str = '',
target_suffix: str = 'upload',
) -> list[dict]:
- return await self.client.preview_skill_zip(file_bytes, filename, source_subdir, target_suffix)
+ execution_context = await self._validated_execution_context(context)
+ return await self.client.preview_skill_zip(
+ file_bytes,
+ filename,
+ source_subdir,
+ target_suffix,
+ action_context=self._action_context(execution_context),
+ )
async def install_skill_zip(
self,
+ context: TenantContext,
file_bytes: bytes,
filename: str,
source_paths: list[str] | None = None,
@@ -955,6 +1184,7 @@ class BoxService:
source_subdir: str = '',
target_suffix: str = 'upload',
) -> list[dict]:
+ execution_context = await self._validated_execution_context(context)
return await self.client.install_skill_zip(
file_bytes,
filename,
@@ -962,6 +1192,7 @@ class BoxService:
source_path,
source_subdir,
target_suffix,
+ action_context=self._action_context(execution_context),
)
def _serialize_result(self, result: BoxExecutionResult) -> dict:
@@ -1280,9 +1511,12 @@ class BoxService:
f'host_path={host_path} session_id={spec.session_id}'
)
- async def _cleanup_exceeded_session(self, spec: BoxSpec) -> None:
+ async def _cleanup_exceeded_session(self, context: TenantContext, spec: BoxSpec) -> None:
try:
- await self.client.delete_session(spec.session_id)
+ await self.client.delete_session(
+ spec.session_id,
+ action_context=self._action_context(context),
+ )
except Exception as exc:
self.ap.logger.warning(
'Failed to clean up Box session after workspace quota was exceeded: '
@@ -1299,11 +1533,19 @@ class BoxService:
'type': type(exc).__name__,
'message': str(exc),
'query_id': str(query.query_id),
+ 'instance_uuid': str(getattr(query, 'instance_uuid', '') or ''),
+ 'workspace_uuid': str(getattr(query, 'workspace_uuid', '') or ''),
}
)
- def get_recent_errors(self) -> list[dict]:
- return list(self._recent_errors)
+ def get_recent_errors(self, context: TenantContext) -> list[dict]:
+ execution_context = self._execution_context(context)
+ return [
+ error
+ for error in self._recent_errors
+ if error.get('instance_uuid') == execution_context.instance_uuid
+ and error.get('workspace_uuid') == execution_context.workspace_uuid
+ ]
def get_system_guidance(self, query_id=None) -> str:
"""Return LLM system-prompt guidance for the exec tool.
@@ -1341,17 +1583,28 @@ class BoxService:
)
return guidance
- async def get_status(self) -> dict:
+ async def get_backend_status(self) -> dict:
+ """Return instance-level backend readiness without tenant resource data."""
+
+ if not self._available:
+ return {'available': False, 'enabled': self._enabled, 'connector_error': self._connector_error}
+ backend = await self.client.get_backend_info()
+ return {'available': bool(backend.get('available', False)), 'enabled': self._enabled, 'backend': backend}
+
+ async def get_status(self, context: TenantContext) -> dict:
+ execution_context = await self._validated_execution_context(context)
+ action_context = self._action_context(execution_context)
+ recent_error_count = len(self.get_recent_errors(execution_context))
if not self._available:
return {
'available': False,
'enabled': self._enabled,
'profile': self.profile.name,
- 'recent_error_count': len(self._recent_errors),
+ 'recent_error_count': recent_error_count,
'connector_error': self._connector_error,
}
try:
- runtime_status = await self.client.get_status()
+ runtime_status = await self.client.get_status(action_context=action_context)
except Exception as exc:
# RPC failed — the runtime likely just disconnected and the
# heartbeat hasn't flipped _available yet.
@@ -1359,7 +1612,7 @@ class BoxService:
'available': False,
'enabled': self._enabled,
'profile': self.profile.name,
- 'recent_error_count': len(self._recent_errors),
+ 'recent_error_count': recent_error_count,
'connector_error': str(exc),
}
# Backend state can be unavailable even when the connector is healthy
@@ -1377,7 +1630,7 @@ class BoxService:
'available': backend_ok,
'enabled': self._enabled,
'profile': self.profile.name,
- 'recent_error_count': len(self._recent_errors),
+ 'recent_error_count': recent_error_count,
}
if not backend_ok and 'connector_error' not in payload:
backend_name = backend_info.get('name') if backend_info else None
diff --git a/src/langbot/pkg/box/workspace.py b/src/langbot/pkg/box/workspace.py
index 26d1a41e9..6f260ae7b 100644
--- a/src/langbot/pkg/box/workspace.py
+++ b/src/langbot/pkg/box/workspace.py
@@ -274,6 +274,7 @@ class BoxWorkspaceSession:
def __init__(
self,
box_service,
+ execution_context,
session_id: str,
*,
host_path: str | None = None,
@@ -290,6 +291,7 @@ class BoxWorkspaceSession:
persistent: bool = False,
):
self.box_service = box_service
+ self.execution_context = execution_context
self.session_id = session_id
self.host_path = host_path
self.host_path_mode = host_path_mode
@@ -363,7 +365,7 @@ class BoxWorkspaceSession:
timeout_sec: int | None = None,
):
payload = self.build_exec_payload(cmd, workdir=workdir, env=env, timeout_sec=timeout_sec)
- return await self.box_service.client.execute(self.box_service.build_spec(payload))
+ return await self.box_service.execute_in_context(self.execution_context, payload)
async def execute_for_query(
self,
@@ -378,7 +380,7 @@ class BoxWorkspaceSession:
return await self.box_service.execute_spec_payload(payload, query)
async def create_session(self):
- return await self.box_service.create_session(self.build_session_payload())
+ return await self.box_service.create_session(self.execution_context, self.build_session_payload())
def build_process_payload(
self,
@@ -415,16 +417,26 @@ class BoxWorkspaceSession:
):
payload = self.build_process_payload(command, args, env=env, cwd=cwd)
payload['process_id'] = process_id
- return await self.box_service.start_managed_process(self.session_id, payload)
+ return await self.box_service.start_managed_process(self.execution_context, self.session_id, payload)
async def get_managed_process(self, process_id: str = 'default'):
- return await self.box_service.get_managed_process(self.session_id, process_id)
+ return await self.box_service.get_managed_process(self.execution_context, self.session_id, process_id)
async def stop_managed_process(self, process_id: str = 'default') -> None:
- await self.box_service.stop_managed_process(self.session_id, process_id)
+ await self.box_service.stop_managed_process(self.execution_context, self.session_id, process_id)
- def get_managed_process_websocket_url(self, process_id: str = 'default') -> str:
- return self.box_service.get_managed_process_websocket_url(self.session_id, process_id)
+ async def get_managed_process_websocket_connection(
+ self,
+ process_id: str = 'default',
+ ) -> tuple[str, dict[str, str]]:
+ return await self.box_service.get_managed_process_websocket_connection(
+ self.execution_context,
+ self.session_id,
+ process_id,
+ )
async def cleanup(self) -> None:
- await self.box_service.client.delete_session(self.session_id)
+ await self.box_service.client.delete_session(
+ self.session_id,
+ action_context=self.box_service._action_context(self.execution_context),
+ )
diff --git a/src/langbot/pkg/command/cmdmgr.py b/src/langbot/pkg/command/cmdmgr.py
index ee064c219..7fe9966a7 100644
--- a/src/langbot/pkg/command/cmdmgr.py
+++ b/src/langbot/pkg/command/cmdmgr.py
@@ -89,6 +89,7 @@ class CommandManager:
_admins = await self.ap.persistence_mgr.execute_async(
_sa.select(_BotAdmin).where(
+ _BotAdmin.workspace_uuid == query.workspace_uuid,
_BotAdmin.bot_uuid == (query.bot_uuid or ''),
_BotAdmin.launcher_type == query.launcher_type.value,
_BotAdmin.launcher_id == str(query.launcher_id),
@@ -98,7 +99,11 @@ class CommandManager:
privilege = 2
ctx = command_context.ExecuteContext(
+ instance_uuid=query.instance_uuid,
+ workspace_uuid=query.workspace_uuid,
+ placement_generation=query.placement_generation,
query_id=query.query_id,
+ query_uuid=query.query_uuid,
session=session,
command_text=command_text,
full_command_text=full_command_text,
diff --git a/src/langbot/pkg/core/app.py b/src/langbot/pkg/core/app.py
index b0adb5594..1952b775b 100644
--- a/src/langbot/pkg/core/app.py
+++ b/src/langbot/pkg/core/app.py
@@ -4,6 +4,7 @@ import logging
import asyncio
import traceback
import os
+import sqlalchemy
from ..platform import botmgr as im_mgr
from ..platform.webhook_pusher import WebhookPusher
@@ -45,6 +46,10 @@ from ..vector import mgr as vectordb_mgr
from ..telemetry import telemetry as telemetry_module
from ..survey import manager as survey_module
from ..skill import manager as skill_mgr
+from ..workspace import service as workspace_service_module
+from ..workspace import collaboration as workspace_collaboration_module
+from ..api.http.context import ExecutionContext, PrincipalContext, PrincipalType
+from ..entity.persistence.workspace import WorkspaceExecutionState, WorkspaceExecutionStatus
class Application:
@@ -119,6 +124,10 @@ class Application:
persistence_mgr: persistencemgr.PersistenceManager = None
+ workspace_service: workspace_service_module.WorkspaceService = None
+
+ workspace_collaboration_service: workspace_collaboration_module.WorkspaceCollaborationService = None
+
vector_db_mgr: vectordb_mgr.VectorDBManager = None
http_ctrl: http_controller.HTTPController = None
@@ -235,16 +244,31 @@ class Application:
check_interval_seconds = check_interval_hours * 3600
while True:
try:
- deleted = await self.monitoring_service.cleanup_expired_records(
- retention_days,
- batch_size=delete_batch_size,
- )
- total_deleted = sum(deleted.values())
- if total_deleted > 0:
- self.logger.info(
- f'Monitoring auto-cleanup: deleted {total_deleted} expired records '
- f'(retention={retention_days}d): {deleted}'
+ execution_states = await self.persistence_mgr.execute_async(
+ sqlalchemy.select(WorkspaceExecutionState).where(
+ WorkspaceExecutionState.instance_uuid == self.workspace_service.instance_uuid,
+ WorkspaceExecutionState.state == WorkspaceExecutionStatus.ACTIVE.value,
+ WorkspaceExecutionState.write_fenced == sqlalchemy.false(),
)
+ )
+ for execution_state in execution_states.all():
+ context = ExecutionContext(
+ instance_uuid=execution_state.instance_uuid,
+ workspace_uuid=execution_state.workspace_uuid,
+ placement_generation=execution_state.active_generation,
+ trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
+ )
+ deleted = await self.monitoring_service.cleanup_expired_records(
+ context,
+ retention_days,
+ batch_size=delete_batch_size,
+ )
+ total_deleted = sum(deleted.values())
+ if total_deleted > 0:
+ self.logger.info(
+ f'Monitoring auto-cleanup: deleted {total_deleted} expired records '
+ f'for Workspace {context.workspace_uuid} (retention={retention_days}d): {deleted}'
+ )
except Exception as e:
self.logger.warning(f'Monitoring auto-cleanup error: {e}')
await asyncio.sleep(check_interval_seconds)
@@ -268,10 +292,27 @@ class Application:
check_interval_seconds = check_interval_hours * 3600
while True:
try:
- deleted = await self.maintenance_service.cleanup_expired_files()
- total_deleted = sum(deleted.values())
- if total_deleted > 0:
- self.logger.info(f'Storage maintenance: deleted expired files: {deleted}')
+ execution_states = await self.persistence_mgr.execute_async(
+ sqlalchemy.select(WorkspaceExecutionState).where(
+ WorkspaceExecutionState.instance_uuid == self.workspace_service.instance_uuid,
+ WorkspaceExecutionState.state == WorkspaceExecutionStatus.ACTIVE.value,
+ WorkspaceExecutionState.write_fenced == sqlalchemy.false(),
+ )
+ )
+ for execution_state in execution_states.all():
+ context = ExecutionContext(
+ instance_uuid=execution_state.instance_uuid,
+ workspace_uuid=execution_state.workspace_uuid,
+ placement_generation=execution_state.active_generation,
+ trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
+ )
+ deleted = await self.maintenance_service.cleanup_expired_files(context)
+ total_deleted = sum(deleted.values())
+ if total_deleted > 0:
+ self.logger.info(
+ f'Storage maintenance for Workspace {context.workspace_uuid}: '
+ f'deleted expired files: {deleted}'
+ )
except Exception as e:
self.logger.warning(f'Storage maintenance error: {e}')
await asyncio.sleep(check_interval_seconds)
diff --git a/src/langbot/pkg/core/stages/build_app.py b/src/langbot/pkg/core/stages/build_app.py
index a8d53b7b3..3fe94223b 100644
--- a/src/langbot/pkg/core/stages/build_app.py
+++ b/src/langbot/pkg/core/stages/build_app.py
@@ -39,6 +39,11 @@ from ...vector import mgr as vectordb_mgr
from .. import taskmgr
from ...telemetry import telemetry as telemetry_module
from ...survey import manager as survey_module
+from ...workspace import service as workspace_service_module
+from ...workspace import collaboration as workspace_collaboration_module
+from ...workspace import policy as workspace_policy_module
+from ...api.http.context import ExecutionContext, PrincipalContext, PrincipalType
+from ...api.http.authz import WorkspaceRequiredError
@stage.stage_class('BuildAppStage')
@@ -53,9 +58,6 @@ class BuildAppStage(stage.BootingStage):
discover.discover_blueprint('templates/components.yaml')
ap.discover = discover
- user_service_inst = user_service.UserService(ap)
- ap.user_service = user_service_inst
-
space_service_inst = space_service.SpaceService(ap)
ap.space_service = space_service_inst
@@ -100,8 +102,6 @@ class BuildAppStage(stage.BootingStage):
await ver_mgr.initialize()
ap.ver_mgr = ver_mgr
- ap.query_pool = pool.QueryPool()
-
log_cache = logcache.LogCache()
ap.log_cache = log_cache
@@ -113,6 +113,43 @@ class BuildAppStage(stage.BootingStage):
ap.persistence_mgr = persistence_mgr_inst
await persistence_mgr_inst.initialize()
+ # The open-source Core is intentionally single-Workspace. A mutable
+ # config value such as ``system.edition`` is product metadata, not a
+ # trust credential, and must never activate SaaS routing. The closed
+ # Cloud bootstrap will install its verified policy only after checking
+ # a signed InstanceManifest; until that bootstrap exists, fail closed.
+ workspace_policy = workspace_policy_module.open_core_workspace_policy()
+ workspace_service_inst = workspace_service_module.WorkspaceService(
+ ap,
+ policy=workspace_policy,
+ )
+ if not workspace_policy.multi_workspace_enabled:
+ await workspace_service_inst.ensure_singleton_workspace()
+ ap.workspace_service = workspace_service_inst
+
+ ap.workspace_collaboration_service = workspace_collaboration_module.WorkspaceCollaborationService(
+ ap,
+ workspace_service_inst,
+ )
+
+ user_service_inst = user_service.UserService(ap)
+ ap.user_service = user_service_inst
+
+ async def resolve_singleton_execution_context() -> ExecutionContext:
+ if workspace_policy.multi_workspace_enabled:
+ raise WorkspaceRequiredError('Cloud runtime work requires an explicit Workspace context')
+ binding = await workspace_service_inst.get_local_execution_binding()
+ return ExecutionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
+ )
+
+ ap.query_pool = pool.QueryPool(
+ singleton_context_resolver=resolve_singleton_execution_context,
+ )
+
# Telemetry manager: attach to app so other components can call via self.ap.telemetry
telemetry_inst = telemetry_module.TelemetryManager(ap)
await telemetry_inst.initialize()
diff --git a/src/langbot/pkg/core/taskmgr.py b/src/langbot/pkg/core/taskmgr.py
index 8bf8784a1..7272bdfb5 100644
--- a/src/langbot/pkg/core/taskmgr.py
+++ b/src/langbot/pkg/core/taskmgr.py
@@ -98,6 +98,15 @@ class TaskWrapper:
scopes: list[core_entities.LifecycleControlScope]
"""Task scope"""
+ instance_uuid: str | None
+ """Owning LangBot instance for a tenant user task."""
+
+ workspace_uuid: str | None
+ """Owning Workspace for a tenant user task."""
+
+ placement_generation: int | None
+ """Workspace execution fence captured when the task was created."""
+
def __init__(
self,
ap: app.Application,
@@ -108,6 +117,9 @@ class TaskWrapper:
label: str = '',
context: TaskContext = None,
scopes: list[core_entities.LifecycleControlScope] = [core_entities.LifecycleControlScope.APPLICATION],
+ instance_uuid: str | None = None,
+ workspace_uuid: str | None = None,
+ placement_generation: int | None = None,
):
self.id = TaskWrapper._id_index
TaskWrapper._id_index += 1
@@ -120,6 +132,9 @@ class TaskWrapper:
self.label = label if label != '' else name
self.task.set_name(name)
self.scopes = scopes
+ self.instance_uuid = instance_uuid
+ self.workspace_uuid = workspace_uuid
+ self.placement_generation = placement_generation
self.created_at = time.time()
def assume_exception(self):
@@ -155,6 +170,8 @@ class TaskWrapper:
'kind': self.kind,
'name': self.name,
'label': self.label,
+ 'workspace_uuid': self.workspace_uuid,
+ 'placement_generation': self.placement_generation,
'scopes': [scope.value for scope in self.scopes],
'created_at': self.created_at,
'task_context': self.task_context.to_dict(),
@@ -193,8 +210,23 @@ class AsyncTaskManager:
label: str = '',
context: TaskContext = None,
scopes: list[core_entities.LifecycleControlScope] = [core_entities.LifecycleControlScope.APPLICATION],
+ instance_uuid: str | None = None,
+ workspace_uuid: str | None = None,
+ placement_generation: int | None = None,
) -> TaskWrapper:
- wrapper = TaskWrapper(self.ap, coro, task_type, kind, name, label, context, scopes)
+ wrapper = TaskWrapper(
+ self.ap,
+ coro,
+ task_type,
+ kind,
+ name,
+ label,
+ context,
+ scopes,
+ instance_uuid,
+ workspace_uuid,
+ placement_generation,
+ )
self.tasks.append(wrapper)
wrapper.task.add_done_callback(lambda _: self._prune_completed_tasks())
self._prune_completed_tasks()
@@ -208,8 +240,22 @@ class AsyncTaskManager:
label: str = '',
context: TaskContext = None,
scopes: list[core_entities.LifecycleControlScope] = [core_entities.LifecycleControlScope.APPLICATION],
+ instance_uuid: str | None = None,
+ workspace_uuid: str | None = None,
+ placement_generation: int | None = None,
) -> TaskWrapper:
- return self.create_task(coro, 'user', kind, name, label, context, scopes)
+ return self.create_task(
+ coro,
+ 'user',
+ kind,
+ name,
+ label,
+ context,
+ scopes,
+ instance_uuid,
+ workspace_uuid,
+ placement_generation,
+ )
async def wait_all(self):
await asyncio.gather(*[t.task for t in self.tasks], return_exceptions=True)
@@ -221,12 +267,20 @@ class AsyncTaskManager:
self,
type: str = None,
kind: str = None,
+ *,
+ instance_uuid: str | None = None,
+ workspace_uuid: str | None = None,
+ placement_generation: int | None = None,
) -> 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)
+ 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,
}
@@ -240,9 +294,21 @@ class AsyncTaskManager:
'id_index': TaskWrapper._id_index,
}
- def get_task_by_id(self, id: int) -> TaskWrapper | None:
+ def get_task_by_id(
+ self,
+ id: int,
+ *,
+ instance_uuid: str | None = None,
+ workspace_uuid: str | None = None,
+ placement_generation: int | None = None,
+ ) -> TaskWrapper | None:
for t in self.tasks:
- if t.id == id:
+ if (
+ t.id == id
+ 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)
+ ):
return t
return None
diff --git a/src/langbot/pkg/entity/persistence/apikey.py b/src/langbot/pkg/entity/persistence/apikey.py
index 488c03241..f15bb9a41 100644
--- a/src/langbot/pkg/entity/persistence/apikey.py
+++ b/src/langbot/pkg/entity/persistence/apikey.py
@@ -1,16 +1,47 @@
+import enum
+import uuid as uuid_lib
+
import sqlalchemy
from .base import Base
+class ApiKeyStatus(enum.StrEnum):
+ ACTIVE = 'active'
+ REVOKED = 'revoked'
+
+
+def _new_uuid() -> str:
+ return str(uuid_lib.uuid4())
+
+
class ApiKey(Base):
"""API Key for external service authentication"""
__tablename__ = 'api_keys'
id = sqlalchemy.Column(sqlalchemy.Integer, primary_key=True, autoincrement=True)
+ uuid = sqlalchemy.Column(sqlalchemy.String(36), nullable=False, default=_new_uuid)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
+ created_by_account_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('users.uuid', ondelete='SET NULL'),
+ nullable=True,
+ )
name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
- key = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, unique=True)
+ key_hash = sqlalchemy.Column(sqlalchemy.String(64), nullable=False)
+ scopes = sqlalchemy.Column(sqlalchemy.JSON, nullable=False, default=list, server_default='[]')
+ status = sqlalchemy.Column(
+ sqlalchemy.String(32),
+ nullable=False,
+ server_default=ApiKeyStatus.ACTIVE.value,
+ )
+ expires_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=True)
+ last_used_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=True)
description = sqlalchemy.Column(sqlalchemy.String(512), nullable=True, default='')
created_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, server_default=sqlalchemy.func.now())
updated_at = sqlalchemy.Column(
@@ -19,3 +50,16 @@ class ApiKey(Base):
server_default=sqlalchemy.func.now(),
onupdate=sqlalchemy.func.now(),
)
+
+ __table_args__ = (
+ sqlalchemy.Index('uq_api_keys_uuid', 'uuid', unique=True),
+ # Authentication begins with the presented secret, before a Workspace
+ # can be trusted, so hashes remain globally unique.
+ sqlalchemy.Index('uq_api_keys_key_hash', 'key_hash', unique=True),
+ sqlalchemy.Index('ix_api_keys_workspace_name', 'workspace_uuid', 'name'),
+ sqlalchemy.Index('ix_api_keys_workspace_status', 'workspace_uuid', 'status'),
+ sqlalchemy.CheckConstraint(
+ "status IN ('active', 'revoked')",
+ name='ck_api_keys_status',
+ ),
+ )
diff --git a/src/langbot/pkg/entity/persistence/bot.py b/src/langbot/pkg/entity/persistence/bot.py
index 9043b7560..d95169f33 100644
--- a/src/langbot/pkg/entity/persistence/bot.py
+++ b/src/langbot/pkg/entity/persistence/bot.py
@@ -9,12 +9,31 @@ class BotAdmin(Base):
__tablename__ = 'bot_admins'
id = sqlalchemy.Column(sqlalchemy.Integer, primary_key=True, autoincrement=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
bot_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
launcher_type = sqlalchemy.Column(sqlalchemy.String(64), nullable=False)
launcher_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
created_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, server_default=sqlalchemy.func.now())
- __table_args__ = (sqlalchemy.UniqueConstraint('bot_uuid', 'launcher_type', 'launcher_id', name='uq_bot_admin'),)
+ __table_args__ = (
+ sqlalchemy.UniqueConstraint(
+ 'workspace_uuid',
+ 'bot_uuid',
+ 'launcher_type',
+ 'launcher_id',
+ name='uq_bot_admin',
+ ),
+ sqlalchemy.ForeignKeyConstraint(
+ ['workspace_uuid', 'bot_uuid'],
+ ['bots.workspace_uuid', 'bots.uuid'],
+ name='fk_bot_admins_workspace_bot',
+ ondelete='CASCADE',
+ ),
+ )
class Bot(Base):
@@ -23,6 +42,11 @@ class Bot(Base):
__tablename__ = 'bots'
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
description = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
adapter = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
@@ -38,3 +62,8 @@ class Bot(Base):
server_default=sqlalchemy.func.now(),
onupdate=sqlalchemy.func.now(),
)
+
+ __table_args__ = (
+ sqlalchemy.UniqueConstraint('workspace_uuid', 'uuid', name='uq_bots_workspace_uuid'),
+ sqlalchemy.Index('ix_bots_workspace_name', 'workspace_uuid', 'name'),
+ )
diff --git a/src/langbot/pkg/entity/persistence/bstorage.py b/src/langbot/pkg/entity/persistence/bstorage.py
index 674dee29b..eb24b5cb3 100644
--- a/src/langbot/pkg/entity/persistence/bstorage.py
+++ b/src/langbot/pkg/entity/persistence/bstorage.py
@@ -8,6 +8,11 @@ class BinaryStorage(Base):
__tablename__ = 'binary_storages'
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ primary_key=True,
+ )
unique_key = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
key = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
owner_type = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
@@ -20,3 +25,12 @@ class BinaryStorage(Base):
server_default=sqlalchemy.func.now(),
onupdate=sqlalchemy.func.now(),
)
+
+ __table_args__ = (
+ sqlalchemy.Index(
+ 'ix_binary_storages_workspace_owner',
+ 'workspace_uuid',
+ 'owner_type',
+ 'owner',
+ ),
+ )
diff --git a/src/langbot/pkg/entity/persistence/mcp.py b/src/langbot/pkg/entity/persistence/mcp.py
index 983fdce53..1e7f93a4a 100644
--- a/src/langbot/pkg/entity/persistence/mcp.py
+++ b/src/langbot/pkg/entity/persistence/mcp.py
@@ -7,6 +7,11 @@ class MCPServer(Base):
__tablename__ = 'mcp_servers'
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
enable = sqlalchemy.Column(sqlalchemy.Boolean, nullable=False, default=False)
mode = sqlalchemy.Column(sqlalchemy.String(255), nullable=False) # stdio, remote (legacy: sse, http)
@@ -22,3 +27,8 @@ class MCPServer(Base):
server_default=sqlalchemy.func.now(),
onupdate=sqlalchemy.func.now(),
)
+
+ __table_args__ = (
+ sqlalchemy.UniqueConstraint('workspace_uuid', 'name', name='uq_mcp_servers_workspace_name'),
+ sqlalchemy.Index('ix_mcp_servers_workspace_enable', 'workspace_uuid', 'enable'),
+ )
diff --git a/src/langbot/pkg/entity/persistence/metadata.py b/src/langbot/pkg/entity/persistence/metadata.py
index ac3b4602f..70b1790d7 100644
--- a/src/langbot/pkg/entity/persistence/metadata.py
+++ b/src/langbot/pkg/entity/persistence/metadata.py
@@ -19,3 +19,17 @@ class Metadata(Base):
key = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
value = sqlalchemy.Column(sqlalchemy.String(255))
+
+
+class WorkspaceMetadata(Base):
+ """Metadata owned by one workspace rather than by the LangBot instance."""
+
+ __tablename__ = 'workspace_metadata'
+
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ primary_key=True,
+ )
+ key = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
+ value = sqlalchemy.Column(sqlalchemy.String(255))
diff --git a/src/langbot/pkg/entity/persistence/model.py b/src/langbot/pkg/entity/persistence/model.py
index 5b5f1fe2f..c04a532b5 100644
--- a/src/langbot/pkg/entity/persistence/model.py
+++ b/src/langbot/pkg/entity/persistence/model.py
@@ -9,6 +9,11 @@ class ModelProvider(Base):
__tablename__ = 'model_providers'
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
requester = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
base_url = sqlalchemy.Column(sqlalchemy.String(512), nullable=False)
@@ -21,6 +26,12 @@ class ModelProvider(Base):
onupdate=sqlalchemy.func.now(),
)
+ __table_args__ = (
+ sqlalchemy.UniqueConstraint('workspace_uuid', 'uuid', name='uq_model_providers_workspace_uuid'),
+ sqlalchemy.Index('ix_model_providers_workspace_name', 'workspace_uuid', 'name'),
+ sqlalchemy.Index('ix_model_providers_workspace_requester', 'workspace_uuid', 'requester'),
+ )
+
class LLMModel(Base):
"""LLM model"""
@@ -28,6 +39,11 @@ class LLMModel(Base):
__tablename__ = 'llm_models'
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
provider_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
abilities = sqlalchemy.Column(sqlalchemy.JSON, nullable=False, default=[])
@@ -42,6 +58,16 @@ class LLMModel(Base):
onupdate=sqlalchemy.func.now(),
)
+ __table_args__ = (
+ sqlalchemy.ForeignKeyConstraint(
+ ['workspace_uuid', 'provider_uuid'],
+ ['model_providers.workspace_uuid', 'model_providers.uuid'],
+ name='fk_llm_models_workspace_provider',
+ ),
+ sqlalchemy.Index('ix_llm_models_workspace_provider', 'workspace_uuid', 'provider_uuid'),
+ sqlalchemy.Index('ix_llm_models_workspace_name', 'workspace_uuid', 'name'),
+ )
+
class EmbeddingModel(Base):
"""Embedding model"""
@@ -49,6 +75,11 @@ class EmbeddingModel(Base):
__tablename__ = 'embedding_models'
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
provider_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
extra_args = sqlalchemy.Column(sqlalchemy.JSON, nullable=False, default={})
@@ -61,6 +92,16 @@ class EmbeddingModel(Base):
onupdate=sqlalchemy.func.now(),
)
+ __table_args__ = (
+ sqlalchemy.ForeignKeyConstraint(
+ ['workspace_uuid', 'provider_uuid'],
+ ['model_providers.workspace_uuid', 'model_providers.uuid'],
+ name='fk_embedding_models_workspace_provider',
+ ),
+ sqlalchemy.Index('ix_embedding_models_workspace_provider', 'workspace_uuid', 'provider_uuid'),
+ sqlalchemy.Index('ix_embedding_models_workspace_name', 'workspace_uuid', 'name'),
+ )
+
class RerankModel(Base):
"""Rerank model"""
@@ -68,6 +109,11 @@ class RerankModel(Base):
__tablename__ = 'rerank_models'
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
provider_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
extra_args = sqlalchemy.Column(sqlalchemy.JSON, nullable=False, default={})
@@ -79,3 +125,13 @@ class RerankModel(Base):
server_default=sqlalchemy.func.now(),
onupdate=sqlalchemy.func.now(),
)
+
+ __table_args__ = (
+ sqlalchemy.ForeignKeyConstraint(
+ ['workspace_uuid', 'provider_uuid'],
+ ['model_providers.workspace_uuid', 'model_providers.uuid'],
+ name='fk_rerank_models_workspace_provider',
+ ),
+ sqlalchemy.Index('ix_rerank_models_workspace_provider', 'workspace_uuid', 'provider_uuid'),
+ sqlalchemy.Index('ix_rerank_models_workspace_name', 'workspace_uuid', 'name'),
+ )
diff --git a/src/langbot/pkg/entity/persistence/monitoring.py b/src/langbot/pkg/entity/persistence/monitoring.py
index f594b8187..35ebe161a 100644
--- a/src/langbot/pkg/entity/persistence/monitoring.py
+++ b/src/langbot/pkg/entity/persistence/monitoring.py
@@ -9,6 +9,11 @@ class MonitoringMessage(Base):
__tablename__ = 'monitoring_messages'
id = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
timestamp = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, index=True)
bot_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, index=True)
bot_name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
@@ -25,6 +30,11 @@ class MonitoringMessage(Base):
variables = sqlalchemy.Column(sqlalchemy.Text, nullable=True) # Query variables as JSON string
role = sqlalchemy.Column(sqlalchemy.String(50), nullable=True, default='user') # user, assistant
+ __table_args__ = (
+ sqlalchemy.Index('ix_monitoring_messages_workspace_timestamp', 'workspace_uuid', 'timestamp'),
+ sqlalchemy.Index('ix_monitoring_messages_workspace_session', 'workspace_uuid', 'session_id'),
+ )
+
class MonitoringLLMCall(Base):
"""LLM call records"""
@@ -32,6 +42,11 @@ class MonitoringLLMCall(Base):
__tablename__ = 'monitoring_llm_calls'
id = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
timestamp = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, index=True)
model_name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
input_tokens = sqlalchemy.Column(sqlalchemy.Integer, nullable=False)
@@ -48,6 +63,11 @@ class MonitoringLLMCall(Base):
error_message = sqlalchemy.Column(sqlalchemy.Text, nullable=True)
message_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=True, index=True) # Associated message ID
+ __table_args__ = (
+ sqlalchemy.Index('ix_monitoring_llm_calls_workspace_timestamp', 'workspace_uuid', 'timestamp'),
+ sqlalchemy.Index('ix_monitoring_llm_calls_workspace_session', 'workspace_uuid', 'session_id'),
+ )
+
class MonitoringToolCall(Base):
"""Tool call records"""
@@ -55,6 +75,11 @@ class MonitoringToolCall(Base):
__tablename__ = 'monitoring_tool_calls'
id = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
timestamp = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, index=True)
tool_name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
tool_source = sqlalchemy.Column(sqlalchemy.String(50), nullable=False) # native, plugin, mcp, skill
@@ -70,12 +95,22 @@ class MonitoringToolCall(Base):
result = sqlalchemy.Column(sqlalchemy.Text, nullable=True)
error_message = sqlalchemy.Column(sqlalchemy.Text, nullable=True)
+ __table_args__ = (
+ sqlalchemy.Index('ix_monitoring_tool_calls_workspace_timestamp', 'workspace_uuid', 'timestamp'),
+ sqlalchemy.Index('ix_monitoring_tool_calls_workspace_session', 'workspace_uuid', 'session_id'),
+ )
+
class MonitoringSession(Base):
"""Session tracking records"""
__tablename__ = 'monitoring_sessions'
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ primary_key=True,
+ )
session_id = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
bot_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, index=True)
bot_name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
@@ -89,6 +124,11 @@ class MonitoringSession(Base):
user_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=True)
user_name = sqlalchemy.Column(sqlalchemy.String(255), nullable=True) # User display name
+ __table_args__ = (
+ sqlalchemy.Index('ix_monitoring_sessions_workspace_activity', 'workspace_uuid', 'last_activity'),
+ sqlalchemy.Index('ix_monitoring_sessions_workspace_active', 'workspace_uuid', 'is_active'),
+ )
+
class MonitoringError(Base):
"""Error log records"""
@@ -96,6 +136,11 @@ class MonitoringError(Base):
__tablename__ = 'monitoring_errors'
id = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
timestamp = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, index=True)
error_type = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
error_message = sqlalchemy.Column(sqlalchemy.Text, nullable=False)
@@ -107,6 +152,11 @@ class MonitoringError(Base):
stack_trace = sqlalchemy.Column(sqlalchemy.Text, nullable=True)
message_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=True, index=True) # Associated message ID
+ __table_args__ = (
+ sqlalchemy.Index('ix_monitoring_errors_workspace_timestamp', 'workspace_uuid', 'timestamp'),
+ sqlalchemy.Index('ix_monitoring_errors_workspace_session', 'workspace_uuid', 'session_id'),
+ )
+
class MonitoringEmbeddingCall(Base):
"""Embedding call records"""
@@ -114,6 +164,11 @@ class MonitoringEmbeddingCall(Base):
__tablename__ = 'monitoring_embedding_calls'
id = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
timestamp = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, index=True)
model_name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
prompt_tokens = sqlalchemy.Column(sqlalchemy.Integer, nullable=False)
@@ -129,6 +184,19 @@ class MonitoringEmbeddingCall(Base):
message_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=True, index=True)
call_type = sqlalchemy.Column(sqlalchemy.String(50), nullable=True) # embedding, retrieve
+ __table_args__ = (
+ sqlalchemy.Index(
+ 'ix_monitoring_embedding_calls_workspace_timestamp',
+ 'workspace_uuid',
+ 'timestamp',
+ ),
+ sqlalchemy.Index(
+ 'ix_monitoring_embedding_calls_workspace_kb',
+ 'workspace_uuid',
+ 'knowledge_base_id',
+ ),
+ )
+
class MonitoringFeedback(Base):
"""User feedback records (like/dislike) from AI Bot conversations"""
@@ -136,8 +204,13 @@ class MonitoringFeedback(Base):
__tablename__ = 'monitoring_feedback'
id = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
timestamp = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, index=True)
- feedback_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, unique=True, index=True)
+ feedback_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=False, index=True)
feedback_type = sqlalchemy.Column(sqlalchemy.Integer, nullable=False) # 1=like, 2=dislike
feedback_content = sqlalchemy.Column(sqlalchemy.Text, nullable=True) # User feedback text
inaccurate_reasons = sqlalchemy.Column(sqlalchemy.Text, nullable=True) # JSON list of inaccurate reasons
@@ -151,3 +224,13 @@ class MonitoringFeedback(Base):
stream_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=True, index=True)
user_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=True)
platform = sqlalchemy.Column(sqlalchemy.String(255), nullable=True) # e.g., wecom
+
+ __table_args__ = (
+ sqlalchemy.UniqueConstraint(
+ 'workspace_uuid',
+ 'feedback_id',
+ name='uq_monitoring_feedback_workspace_feedback_id',
+ ),
+ sqlalchemy.Index('ix_monitoring_feedback_workspace_timestamp', 'workspace_uuid', 'timestamp'),
+ sqlalchemy.Index('ix_monitoring_feedback_workspace_session', 'workspace_uuid', 'session_id'),
+ )
diff --git a/src/langbot/pkg/entity/persistence/pipeline.py b/src/langbot/pkg/entity/persistence/pipeline.py
index d74cf78ee..6451ed6b7 100644
--- a/src/langbot/pkg/entity/persistence/pipeline.py
+++ b/src/langbot/pkg/entity/persistence/pipeline.py
@@ -9,6 +9,11 @@ class LegacyPipeline(Base):
__tablename__ = 'legacy_pipelines'
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
description = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
emoji = sqlalchemy.Column(sqlalchemy.String(10), nullable=True, default='⚙️')
@@ -36,6 +41,16 @@ class LegacyPipeline(Base):
},
)
+ __table_args__ = (
+ sqlalchemy.UniqueConstraint(
+ 'workspace_uuid',
+ 'uuid',
+ name='uq_legacy_pipelines_workspace_uuid',
+ ),
+ sqlalchemy.Index('ix_legacy_pipelines_workspace_name', 'workspace_uuid', 'name'),
+ sqlalchemy.Index('ix_legacy_pipelines_workspace_default', 'workspace_uuid', 'is_default'),
+ )
+
class PipelineRunRecord(Base):
"""Pipeline run record"""
@@ -43,6 +58,11 @@ class PipelineRunRecord(Base):
__tablename__ = 'pipeline_run_records'
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
pipeline_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
status = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
created_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, server_default=sqlalchemy.func.now())
@@ -56,3 +76,22 @@ class PipelineRunRecord(Base):
finished_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False)
result = sqlalchemy.Column(sqlalchemy.JSON, nullable=False)
knowledge_base_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=True)
+
+ __table_args__ = (
+ sqlalchemy.ForeignKeyConstraint(
+ ['workspace_uuid', 'pipeline_uuid'],
+ ['legacy_pipelines.workspace_uuid', 'legacy_pipelines.uuid'],
+ name='fk_pipeline_run_records_workspace_pipeline',
+ ondelete='CASCADE',
+ ),
+ sqlalchemy.Index(
+ 'ix_pipeline_run_records_workspace_pipeline',
+ 'workspace_uuid',
+ 'pipeline_uuid',
+ ),
+ sqlalchemy.Index(
+ 'ix_pipeline_run_records_workspace_created',
+ 'workspace_uuid',
+ 'created_at',
+ ),
+ )
diff --git a/src/langbot/pkg/entity/persistence/plugin.py b/src/langbot/pkg/entity/persistence/plugin.py
index 61629586d..b1b394c43 100644
--- a/src/langbot/pkg/entity/persistence/plugin.py
+++ b/src/langbot/pkg/entity/persistence/plugin.py
@@ -8,6 +8,11 @@ class PluginSetting(Base):
__tablename__ = 'plugin_settings'
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ primary_key=True,
+ )
plugin_author = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
plugin_name = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True)
enabled = sqlalchemy.Column(sqlalchemy.Boolean, nullable=False, default=True)
@@ -22,3 +27,5 @@ class PluginSetting(Base):
server_default=sqlalchemy.func.now(),
onupdate=sqlalchemy.func.now(),
)
+
+ __table_args__ = (sqlalchemy.Index('ix_plugin_settings_workspace_enabled', 'workspace_uuid', 'enabled'),)
diff --git a/src/langbot/pkg/entity/persistence/rag.py b/src/langbot/pkg/entity/persistence/rag.py
index cfb1f0a5c..702659f10 100644
--- a/src/langbot/pkg/entity/persistence/rag.py
+++ b/src/langbot/pkg/entity/persistence/rag.py
@@ -5,6 +5,11 @@ from .base import Base
class KnowledgeBase(Base):
__tablename__ = 'knowledge_bases'
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
name = sqlalchemy.Column(sqlalchemy.String, index=True)
description = sqlalchemy.Column(sqlalchemy.Text)
emoji = sqlalchemy.Column(sqlalchemy.String(10), nullable=True, default='📚')
@@ -13,6 +18,15 @@ class KnowledgeBase(Base):
# New fields for plugin-based RAG
knowledge_engine_plugin_id = sqlalchemy.Column(sqlalchemy.String, nullable=True)
collection_id = sqlalchemy.Column(sqlalchemy.String, nullable=True)
+ # Server-managed compatibility marker. Pre-tenancy installations stored
+ # vectors directly under ``collection_id``; new knowledge bases use a
+ # tenant-derived opaque physical collection instead.
+ legacy_vector_collection = sqlalchemy.Column(
+ sqlalchemy.Boolean,
+ nullable=False,
+ default=False,
+ server_default=sqlalchemy.false(),
+ )
creation_settings = sqlalchemy.Column(sqlalchemy.JSON, nullable=True, default=None)
retrieval_settings = sqlalchemy.Column(sqlalchemy.JSON, nullable=True, default=None)
@@ -23,22 +37,72 @@ class KnowledgeBase(Base):
CREATE_FIELDS = MUTABLE_FIELDS | {'uuid', 'knowledge_engine_plugin_id', 'collection_id', 'creation_settings'}
"""Fields used when creating a new knowledge base."""
- ALL_DB_FIELDS = CREATE_FIELDS | {'emoji', 'created_at', 'updated_at'}
+ ALL_DB_FIELDS = CREATE_FIELDS | {
+ 'workspace_uuid',
+ 'legacy_vector_collection',
+ 'emoji',
+ 'created_at',
+ 'updated_at',
+ }
"""All fields stored in database (for loading from DB row)."""
+ __table_args__ = (
+ sqlalchemy.UniqueConstraint('workspace_uuid', 'uuid', name='uq_knowledge_bases_workspace_uuid'),
+ sqlalchemy.Index('ix_knowledge_bases_workspace_name', 'workspace_uuid', 'name'),
+ sqlalchemy.Index(
+ 'uq_knowledge_bases_workspace_collection',
+ 'workspace_uuid',
+ 'collection_id',
+ unique=True,
+ sqlite_where=sqlalchemy.text('collection_id IS NOT NULL'),
+ postgresql_where=sqlalchemy.text('collection_id IS NOT NULL'),
+ ),
+ )
+
class File(Base):
__tablename__ = 'knowledge_base_files'
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
kb_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=True)
file_name = sqlalchemy.Column(sqlalchemy.String)
extension = sqlalchemy.Column(sqlalchemy.String)
created_at = sqlalchemy.Column(sqlalchemy.DateTime, default=sqlalchemy.func.now())
status = sqlalchemy.Column(sqlalchemy.String, default='pending') # pending, processing, completed, failed
+ __table_args__ = (
+ sqlalchemy.UniqueConstraint('workspace_uuid', 'uuid', name='uq_knowledge_base_files_workspace_uuid'),
+ sqlalchemy.ForeignKeyConstraint(
+ ['workspace_uuid', 'kb_id'],
+ ['knowledge_bases.workspace_uuid', 'knowledge_bases.uuid'],
+ name='fk_knowledge_base_files_workspace_kb',
+ ondelete='CASCADE',
+ ),
+ sqlalchemy.Index('ix_knowledge_base_files_workspace_kb', 'workspace_uuid', 'kb_id'),
+ )
+
class Chunk(Base):
__tablename__ = 'knowledge_base_chunks'
uuid = sqlalchemy.Column(sqlalchemy.String(255), primary_key=True, unique=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
file_id = sqlalchemy.Column(sqlalchemy.String(255), nullable=True)
text = sqlalchemy.Column(sqlalchemy.Text)
+
+ __table_args__ = (
+ sqlalchemy.ForeignKeyConstraint(
+ ['workspace_uuid', 'file_id'],
+ ['knowledge_base_files.workspace_uuid', 'knowledge_base_files.uuid'],
+ name='fk_knowledge_base_chunks_workspace_file',
+ ondelete='CASCADE',
+ ),
+ sqlalchemy.Index('ix_knowledge_base_chunks_workspace_file', 'workspace_uuid', 'file_id'),
+ )
diff --git a/src/langbot/pkg/entity/persistence/user.py b/src/langbot/pkg/entity/persistence/user.py
index 00ea73803..de3a97143 100644
--- a/src/langbot/pkg/entity/persistence/user.py
+++ b/src/langbot/pkg/entity/persistence/user.py
@@ -1,15 +1,47 @@
+import enum
+import uuid as uuid_lib
+
import sqlalchemy
from .base import Base
+class AccountStatus(enum.StrEnum):
+ ACTIVE = 'active'
+ DISABLED = 'disabled'
+ DELETED = 'deleted'
+
+
+class AccountSource(enum.StrEnum):
+ LOCAL = 'local'
+ CLOUD_PROJECTION = 'cloud_projection'
+
+
class User(Base):
__tablename__ = 'users'
id = sqlalchemy.Column(sqlalchemy.Integer, primary_key=True)
+ uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ nullable=False,
+ default=lambda: str(uuid_lib.uuid4()),
+ )
user = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
+ normalized_email = sqlalchemy.Column(sqlalchemy.String(320), nullable=False)
password = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
+ status = sqlalchemy.Column(
+ sqlalchemy.String(32),
+ nullable=False,
+ server_default=AccountStatus.ACTIVE.value,
+ )
+ source = sqlalchemy.Column(
+ sqlalchemy.String(32),
+ nullable=False,
+ server_default=AccountSource.LOCAL.value,
+ )
+ projection_revision = sqlalchemy.Column(sqlalchemy.BigInteger, nullable=False, server_default='0')
+
# Account type: 'local' (default) or 'space'
account_type = sqlalchemy.Column(sqlalchemy.String(32), nullable=False, server_default='local')
@@ -27,3 +59,22 @@ class User(Base):
server_default=sqlalchemy.func.now(),
onupdate=sqlalchemy.func.now(),
)
+
+ __table_args__ = (
+ sqlalchemy.Index('uq_users_uuid', 'uuid', unique=True),
+ sqlalchemy.Index('uq_users_normalized_email', 'normalized_email', unique=True),
+ sqlalchemy.CheckConstraint(
+ 'normalized_email = trim(normalized_email) '
+ 'AND length(normalized_email) > 0 '
+ 'AND length(normalized_email) <= 320',
+ name='ck_users_normalized_email',
+ ),
+ sqlalchemy.CheckConstraint(
+ "status IN ('active', 'disabled', 'deleted')",
+ name='ck_users_status',
+ ),
+ sqlalchemy.CheckConstraint(
+ "source IN ('local', 'cloud_projection')",
+ name='ck_users_source',
+ ),
+ )
diff --git a/src/langbot/pkg/entity/persistence/webhook.py b/src/langbot/pkg/entity/persistence/webhook.py
index 326ab6c47..e5d120132 100644
--- a/src/langbot/pkg/entity/persistence/webhook.py
+++ b/src/langbot/pkg/entity/persistence/webhook.py
@@ -9,6 +9,11 @@ class Webhook(Base):
__tablename__ = 'webhooks'
id = sqlalchemy.Column(sqlalchemy.Integer, primary_key=True, autoincrement=True)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
url = sqlalchemy.Column(sqlalchemy.String(1024), nullable=False)
description = sqlalchemy.Column(sqlalchemy.String(512), nullable=True, default='')
@@ -20,3 +25,5 @@ class Webhook(Base):
server_default=sqlalchemy.func.now(),
onupdate=sqlalchemy.func.now(),
)
+
+ __table_args__ = (sqlalchemy.Index('ix_webhooks_workspace_name', 'workspace_uuid', 'name'),)
diff --git a/src/langbot/pkg/entity/persistence/workspace.py b/src/langbot/pkg/entity/persistence/workspace.py
new file mode 100644
index 000000000..ca3743d4b
--- /dev/null
+++ b/src/langbot/pkg/entity/persistence/workspace.py
@@ -0,0 +1,271 @@
+from __future__ import annotations
+
+import enum
+import uuid as uuid_lib
+
+import sqlalchemy
+
+from .base import Base
+
+
+class WorkspaceType(enum.StrEnum):
+ PERSONAL = 'personal'
+ TEAM = 'team'
+
+
+class WorkspaceStatus(enum.StrEnum):
+ PROVISIONING = 'provisioning'
+ ACTIVE = 'active'
+ SUSPENDED = 'suspended'
+ ARCHIVED = 'archived'
+ DELETED = 'deleted'
+
+
+class WorkspaceSource(enum.StrEnum):
+ LOCAL = 'local'
+ CLOUD_PROJECTION = 'cloud_projection'
+
+
+class MembershipRole(enum.StrEnum):
+ OWNER = 'owner'
+ ADMIN = 'admin'
+ DEVELOPER = 'developer'
+ OPERATOR = 'operator'
+ VIEWER = 'viewer'
+
+
+class MembershipStatus(enum.StrEnum):
+ ACTIVE = 'active'
+ DISABLED = 'disabled'
+ REMOVED = 'removed'
+
+
+class InvitationStatus(enum.StrEnum):
+ PENDING = 'pending'
+ ACCEPTED = 'accepted'
+ REVOKED = 'revoked'
+ EXPIRED = 'expired'
+
+
+class WorkspaceExecutionStatus(enum.StrEnum):
+ PROVISIONING = 'provisioning'
+ ACTIVE = 'active'
+ MIGRATING = 'migrating'
+ DRAINING = 'draining'
+ INACTIVE = 'inactive'
+
+
+class WorkspaceExecutionSource(enum.StrEnum):
+ LOCAL = 'local'
+ CLOUD = 'cloud'
+
+
+def _new_uuid() -> str:
+ return str(uuid_lib.uuid4())
+
+
+class Workspace(Base):
+ __tablename__ = 'workspaces'
+
+ uuid = sqlalchemy.Column(sqlalchemy.String(36), primary_key=True, default=_new_uuid)
+ instance_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
+ name = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
+ slug = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
+ type = sqlalchemy.Column(
+ sqlalchemy.String(32),
+ nullable=False,
+ server_default=WorkspaceType.TEAM.value,
+ )
+ status = sqlalchemy.Column(
+ sqlalchemy.String(32),
+ nullable=False,
+ server_default=WorkspaceStatus.ACTIVE.value,
+ )
+ created_by_account_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('users.uuid', ondelete='SET NULL'),
+ nullable=True,
+ )
+ source = sqlalchemy.Column(
+ sqlalchemy.String(32),
+ nullable=False,
+ server_default=WorkspaceSource.LOCAL.value,
+ )
+ projection_revision = sqlalchemy.Column(sqlalchemy.BigInteger, nullable=False, server_default='0')
+ created_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, server_default=sqlalchemy.func.now())
+ updated_at = sqlalchemy.Column(
+ sqlalchemy.DateTime,
+ nullable=False,
+ server_default=sqlalchemy.func.now(),
+ onupdate=sqlalchemy.func.now(),
+ )
+
+ __table_args__ = (
+ sqlalchemy.UniqueConstraint('instance_uuid', 'slug', name='uq_workspaces_instance_slug'),
+ sqlalchemy.Index('ix_workspaces_instance_status', 'instance_uuid', 'status'),
+ sqlalchemy.Index(
+ 'uq_workspaces_local_instance',
+ 'instance_uuid',
+ unique=True,
+ sqlite_where=sqlalchemy.text("source = 'local'"),
+ postgresql_where=sqlalchemy.text("source = 'local'"),
+ ),
+ sqlalchemy.CheckConstraint(
+ "type IN ('personal', 'team')",
+ name='ck_workspaces_type',
+ ),
+ sqlalchemy.CheckConstraint(
+ "status IN ('provisioning', 'active', 'suspended', 'archived', 'deleted')",
+ name='ck_workspaces_status',
+ ),
+ sqlalchemy.CheckConstraint(
+ "source IN ('local', 'cloud_projection')",
+ name='ck_workspaces_source',
+ ),
+ )
+
+
+class WorkspaceMembership(Base):
+ __tablename__ = 'workspace_memberships'
+
+ uuid = sqlalchemy.Column(sqlalchemy.String(36), primary_key=True, default=_new_uuid)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
+ account_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('users.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
+ role = sqlalchemy.Column(sqlalchemy.String(32), nullable=False)
+ status = sqlalchemy.Column(
+ sqlalchemy.String(32),
+ nullable=False,
+ server_default=MembershipStatus.ACTIVE.value,
+ )
+ invited_by_account_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('users.uuid', ondelete='SET NULL'),
+ nullable=True,
+ )
+ joined_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=True)
+ projection_revision = sqlalchemy.Column(sqlalchemy.BigInteger, nullable=False, server_default='0')
+ created_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, server_default=sqlalchemy.func.now())
+ updated_at = sqlalchemy.Column(
+ sqlalchemy.DateTime,
+ nullable=False,
+ server_default=sqlalchemy.func.now(),
+ onupdate=sqlalchemy.func.now(),
+ )
+
+ __table_args__ = (
+ sqlalchemy.UniqueConstraint('workspace_uuid', 'account_uuid', name='uq_workspace_membership_account'),
+ sqlalchemy.Index('ix_workspace_memberships_account_status', 'account_uuid', 'status'),
+ sqlalchemy.CheckConstraint(
+ "role IN ('owner', 'admin', 'developer', 'operator', 'viewer')",
+ name='ck_workspace_memberships_role',
+ ),
+ sqlalchemy.CheckConstraint(
+ "status IN ('active', 'disabled', 'removed')",
+ name='ck_workspace_memberships_status',
+ ),
+ )
+
+
+class WorkspaceInvitation(Base):
+ __tablename__ = 'workspace_invitations'
+
+ uuid = sqlalchemy.Column(sqlalchemy.String(36), primary_key=True, default=_new_uuid)
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
+ normalized_email = sqlalchemy.Column(sqlalchemy.String(320), nullable=False)
+ role = sqlalchemy.Column(sqlalchemy.String(32), nullable=False)
+ token_hash = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
+ status = sqlalchemy.Column(
+ sqlalchemy.String(32),
+ nullable=False,
+ server_default=InvitationStatus.PENDING.value,
+ )
+ expires_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False)
+ accepted_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=True)
+ revoked_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=True)
+ created_by_account_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('users.uuid', ondelete='CASCADE'),
+ nullable=False,
+ )
+ created_at = sqlalchemy.Column(sqlalchemy.DateTime, nullable=False, server_default=sqlalchemy.func.now())
+ updated_at = sqlalchemy.Column(
+ sqlalchemy.DateTime,
+ nullable=False,
+ server_default=sqlalchemy.func.now(),
+ onupdate=sqlalchemy.func.now(),
+ )
+
+ __table_args__ = (
+ sqlalchemy.Index('uq_workspace_invitations_token_hash', 'token_hash', unique=True),
+ sqlalchemy.Index(
+ 'uq_workspace_invitations_pending_email',
+ 'workspace_uuid',
+ 'normalized_email',
+ unique=True,
+ sqlite_where=sqlalchemy.text("status = 'pending'"),
+ postgresql_where=sqlalchemy.text("status = 'pending'"),
+ ),
+ sqlalchemy.CheckConstraint(
+ "role IN ('admin', 'developer', 'operator', 'viewer')",
+ name='ck_workspace_invitations_role',
+ ),
+ sqlalchemy.CheckConstraint(
+ "status IN ('pending', 'accepted', 'revoked', 'expired')",
+ name='ck_workspace_invitations_status',
+ ),
+ )
+
+
+class WorkspaceExecutionState(Base):
+ __tablename__ = 'workspace_execution_states'
+
+ workspace_uuid = sqlalchemy.Column(
+ sqlalchemy.String(36),
+ sqlalchemy.ForeignKey('workspaces.uuid', ondelete='CASCADE'),
+ primary_key=True,
+ )
+ instance_uuid = sqlalchemy.Column(sqlalchemy.String(255), nullable=False)
+ active_generation = sqlalchemy.Column(sqlalchemy.BigInteger, nullable=False, server_default='1')
+ state = sqlalchemy.Column(
+ sqlalchemy.String(32),
+ nullable=False,
+ server_default=WorkspaceExecutionStatus.ACTIVE.value,
+ )
+ write_fenced = sqlalchemy.Column(sqlalchemy.Boolean, nullable=False, server_default=sqlalchemy.false())
+ source = sqlalchemy.Column(
+ sqlalchemy.String(32),
+ nullable=False,
+ server_default=WorkspaceExecutionSource.LOCAL.value,
+ )
+ desired_state_revision = sqlalchemy.Column(sqlalchemy.BigInteger, nullable=False, server_default='0')
+ updated_at = sqlalchemy.Column(
+ sqlalchemy.DateTime,
+ nullable=False,
+ server_default=sqlalchemy.func.now(),
+ onupdate=sqlalchemy.func.now(),
+ )
+
+ __table_args__ = (
+ sqlalchemy.Index('ix_workspace_execution_states_instance_state', 'instance_uuid', 'state'),
+ sqlalchemy.CheckConstraint('active_generation > 0', name='ck_workspace_execution_generation'),
+ sqlalchemy.CheckConstraint(
+ "state IN ('provisioning', 'active', 'migrating', 'draining', 'inactive')",
+ name='ck_workspace_execution_state',
+ ),
+ sqlalchemy.CheckConstraint(
+ "source IN ('local', 'cloud')",
+ name='ck_workspace_execution_source',
+ ),
+ )
diff --git a/src/langbot/pkg/persistence/alembic/versions/0009_workspace_tenancy_kernel.py b/src/langbot/pkg/persistence/alembic/versions/0009_workspace_tenancy_kernel.py
new file mode 100644
index 000000000..d0a6c233e
--- /dev/null
+++ b/src/langbot/pkg/persistence/alembic/versions/0009_workspace_tenancy_kernel.py
@@ -0,0 +1,542 @@
+"""add the workspace tenancy persistence kernel
+
+Revision ID: 0009_workspace_tenancy
+Revises: 0008_mcp_resource_prefs
+Create Date: 2026-07-18
+"""
+
+from __future__ import annotations
+
+import datetime
+import uuid
+
+import sqlalchemy as sa
+from alembic import op
+
+revision = '0009_workspace_tenancy'
+down_revision = '0008_mcp_resource_prefs'
+branch_labels = None
+depends_on = None
+
+
+def _table_names(conn: sa.Connection) -> set[str]:
+ return set(sa.inspect(conn).get_table_names())
+
+
+def _column_map(conn: sa.Connection, table_name: str) -> dict[str, dict]:
+ return {column['name']: column for column in sa.inspect(conn).get_columns(table_name)}
+
+
+def _constraint_names(conn: sa.Connection, table_name: str) -> set[str]:
+ inspector = sa.inspect(conn)
+ names = {
+ constraint['name']
+ for constraint in inspector.get_check_constraints(table_name)
+ if constraint.get('name') is not None
+ }
+ names.update(
+ constraint['name']
+ for constraint in inspector.get_unique_constraints(table_name)
+ if constraint.get('name') is not None
+ )
+ return names
+
+
+def _index_names(conn: sa.Connection, table_name: str) -> set[str]:
+ return {index['name'] for index in sa.inspect(conn).get_indexes(table_name)}
+
+
+def _upgrade_users(conn: sa.Connection) -> None:
+ if 'users' not in _table_names(conn):
+ return
+
+ columns = _column_map(conn, 'users')
+ if 'uuid' not in columns:
+ op.add_column('users', sa.Column('uuid', sa.String(36), nullable=True))
+ if 'status' not in columns:
+ op.add_column('users', sa.Column('status', sa.String(32), nullable=True, server_default='active'))
+ if 'source' not in columns:
+ op.add_column('users', sa.Column('source', sa.String(32), nullable=True, server_default='local'))
+ if 'projection_revision' not in columns:
+ op.add_column(
+ 'users',
+ sa.Column('projection_revision', sa.BigInteger(), nullable=True, server_default='0'),
+ )
+
+ users = sa.table(
+ 'users',
+ sa.column('id', sa.Integer()),
+ sa.column('uuid', sa.String(36)),
+ sa.column('status', sa.String(32)),
+ sa.column('source', sa.String(32)),
+ sa.column('projection_revision', sa.BigInteger()),
+ )
+
+ seen_uuids: set[str] = set()
+ for user_id, account_uuid in conn.execute(sa.select(users.c.id, users.c.uuid).order_by(users.c.id)).all():
+ normalized_uuid = account_uuid.strip() if isinstance(account_uuid, str) else ''
+ try:
+ normalized_uuid = str(uuid.UUID(normalized_uuid))
+ except (ValueError, AttributeError):
+ normalized_uuid = ''
+ if not normalized_uuid or normalized_uuid in seen_uuids:
+ normalized_uuid = str(uuid.uuid4())
+ if normalized_uuid != account_uuid:
+ conn.execute(users.update().where(users.c.id == user_id).values(uuid=normalized_uuid))
+ seen_uuids.add(normalized_uuid)
+
+ conn.execute(users.update().where(users.c.status.is_(None)).values(status='active'))
+ conn.execute(users.update().where(users.c.source.is_(None)).values(source='local'))
+ conn.execute(users.update().where(users.c.projection_revision.is_(None)).values(projection_revision=0))
+
+ columns = _column_map(conn, 'users')
+ constraint_names = _constraint_names(conn, 'users')
+ needs_batch_alter = any(
+ columns[column_name]['nullable'] for column_name in ('uuid', 'status', 'source', 'projection_revision')
+ ) or not {'ck_users_status', 'ck_users_source'}.issubset(constraint_names)
+
+ if needs_batch_alter:
+ with op.batch_alter_table('users') as batch_op:
+ if columns['uuid']['nullable']:
+ batch_op.alter_column('uuid', existing_type=sa.String(36), nullable=False)
+ if columns['status']['nullable']:
+ batch_op.alter_column(
+ 'status',
+ existing_type=sa.String(32),
+ nullable=False,
+ server_default='active',
+ )
+ if columns['source']['nullable']:
+ batch_op.alter_column(
+ 'source',
+ existing_type=sa.String(32),
+ nullable=False,
+ server_default='local',
+ )
+ if columns['projection_revision']['nullable']:
+ batch_op.alter_column(
+ 'projection_revision',
+ existing_type=sa.BigInteger(),
+ nullable=False,
+ server_default='0',
+ )
+ if 'ck_users_status' not in constraint_names:
+ batch_op.create_check_constraint(
+ 'ck_users_status',
+ "status IN ('active', 'disabled', 'deleted')",
+ )
+ if 'ck_users_source' not in constraint_names:
+ batch_op.create_check_constraint(
+ 'ck_users_source',
+ "source IN ('local', 'cloud_projection')",
+ )
+
+ if 'uq_users_uuid' not in _index_names(conn, 'users'):
+ op.create_index('uq_users_uuid', 'users', ['uuid'], unique=True)
+
+
+def _create_workspace_tables(conn: sa.Connection) -> None:
+ tables = _table_names(conn)
+ if 'users' not in tables:
+ # LangBot's supported startup path creates the baseline schema before
+ # Alembic runs. Keep direct Alembic probes on an empty database safe.
+ return
+
+ if 'workspaces' not in tables:
+ op.create_table(
+ 'workspaces',
+ sa.Column('uuid', sa.String(36), primary_key=True),
+ sa.Column('instance_uuid', sa.String(255), nullable=False),
+ sa.Column('name', sa.String(255), nullable=False),
+ sa.Column('slug', sa.String(255), nullable=False),
+ sa.Column('type', sa.String(32), nullable=False, server_default='team'),
+ sa.Column('status', sa.String(32), nullable=False, server_default='active'),
+ sa.Column('created_by_account_uuid', sa.String(36), nullable=True),
+ sa.Column('source', sa.String(32), nullable=False, server_default='local'),
+ sa.Column('projection_revision', sa.BigInteger(), nullable=False, server_default='0'),
+ sa.Column('created_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
+ sa.Column('updated_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
+ sa.ForeignKeyConstraint(
+ ['created_by_account_uuid'],
+ ['users.uuid'],
+ name='fk_workspaces_created_by_account',
+ ondelete='SET NULL',
+ ),
+ sa.UniqueConstraint('instance_uuid', 'slug', name='uq_workspaces_instance_slug'),
+ sa.CheckConstraint("type IN ('personal', 'team')", name='ck_workspaces_type'),
+ sa.CheckConstraint(
+ "status IN ('provisioning', 'active', 'suspended', 'archived', 'deleted')",
+ name='ck_workspaces_status',
+ ),
+ sa.CheckConstraint(
+ "source IN ('local', 'cloud_projection')",
+ name='ck_workspaces_source',
+ ),
+ )
+ workspace_indexes = _index_names(conn, 'workspaces')
+ if 'ix_workspaces_instance_status' not in workspace_indexes:
+ op.create_index(
+ 'ix_workspaces_instance_status',
+ 'workspaces',
+ ['instance_uuid', 'status'],
+ )
+ if 'uq_workspaces_local_instance' not in workspace_indexes:
+ op.create_index(
+ 'uq_workspaces_local_instance',
+ 'workspaces',
+ ['instance_uuid'],
+ unique=True,
+ sqlite_where=sa.text("source = 'local'"),
+ postgresql_where=sa.text("source = 'local'"),
+ )
+
+ tables = _table_names(conn)
+ if 'workspace_memberships' not in tables:
+ op.create_table(
+ 'workspace_memberships',
+ sa.Column('uuid', sa.String(36), primary_key=True),
+ sa.Column('workspace_uuid', sa.String(36), nullable=False),
+ sa.Column('account_uuid', sa.String(36), nullable=False),
+ sa.Column('role', sa.String(32), nullable=False),
+ sa.Column('status', sa.String(32), nullable=False, server_default='active'),
+ sa.Column('invited_by_account_uuid', sa.String(36), nullable=True),
+ sa.Column('joined_at', sa.DateTime(), nullable=True),
+ sa.Column('projection_revision', sa.BigInteger(), nullable=False, server_default='0'),
+ sa.Column('created_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
+ sa.Column('updated_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
+ sa.ForeignKeyConstraint(
+ ['workspace_uuid'],
+ ['workspaces.uuid'],
+ name='fk_workspace_memberships_workspace',
+ ondelete='CASCADE',
+ ),
+ sa.ForeignKeyConstraint(
+ ['account_uuid'],
+ ['users.uuid'],
+ name='fk_workspace_memberships_account',
+ ondelete='CASCADE',
+ ),
+ sa.ForeignKeyConstraint(
+ ['invited_by_account_uuid'],
+ ['users.uuid'],
+ name='fk_workspace_memberships_invited_by_account',
+ ondelete='SET NULL',
+ ),
+ sa.UniqueConstraint(
+ 'workspace_uuid',
+ 'account_uuid',
+ name='uq_workspace_membership_account',
+ ),
+ sa.CheckConstraint(
+ "role IN ('owner', 'admin', 'developer', 'operator', 'viewer')",
+ name='ck_workspace_memberships_role',
+ ),
+ sa.CheckConstraint(
+ "status IN ('active', 'disabled', 'removed')",
+ name='ck_workspace_memberships_status',
+ ),
+ )
+ membership_indexes = _index_names(conn, 'workspace_memberships')
+ if 'ix_workspace_memberships_account_status' not in membership_indexes:
+ op.create_index(
+ 'ix_workspace_memberships_account_status',
+ 'workspace_memberships',
+ ['account_uuid', 'status'],
+ )
+
+ tables = _table_names(conn)
+ if 'workspace_invitations' not in tables:
+ op.create_table(
+ 'workspace_invitations',
+ sa.Column('uuid', sa.String(36), primary_key=True),
+ sa.Column('workspace_uuid', sa.String(36), nullable=False),
+ sa.Column('normalized_email', sa.String(320), nullable=False),
+ sa.Column('role', sa.String(32), nullable=False),
+ sa.Column('token_hash', sa.String(255), nullable=False),
+ sa.Column('status', sa.String(32), nullable=False, server_default='pending'),
+ sa.Column('expires_at', sa.DateTime(), nullable=False),
+ sa.Column('accepted_at', sa.DateTime(), nullable=True),
+ sa.Column('revoked_at', sa.DateTime(), nullable=True),
+ sa.Column('created_by_account_uuid', sa.String(36), nullable=False),
+ sa.Column('created_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
+ sa.Column('updated_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
+ sa.ForeignKeyConstraint(
+ ['workspace_uuid'],
+ ['workspaces.uuid'],
+ name='fk_workspace_invitations_workspace',
+ ondelete='CASCADE',
+ ),
+ sa.ForeignKeyConstraint(
+ ['created_by_account_uuid'],
+ ['users.uuid'],
+ name='fk_workspace_invitations_created_by_account',
+ ondelete='CASCADE',
+ ),
+ sa.CheckConstraint(
+ "role IN ('admin', 'developer', 'operator', 'viewer')",
+ name='ck_workspace_invitations_role',
+ ),
+ sa.CheckConstraint(
+ "status IN ('pending', 'accepted', 'revoked', 'expired')",
+ name='ck_workspace_invitations_status',
+ ),
+ )
+ invitation_indexes = _index_names(conn, 'workspace_invitations')
+ if 'uq_workspace_invitations_token_hash' not in invitation_indexes:
+ op.create_index(
+ 'uq_workspace_invitations_token_hash',
+ 'workspace_invitations',
+ ['token_hash'],
+ unique=True,
+ )
+ if 'uq_workspace_invitations_pending_email' not in invitation_indexes:
+ op.create_index(
+ 'uq_workspace_invitations_pending_email',
+ 'workspace_invitations',
+ ['workspace_uuid', 'normalized_email'],
+ unique=True,
+ sqlite_where=sa.text("status = 'pending'"),
+ postgresql_where=sa.text("status = 'pending'"),
+ )
+
+ tables = _table_names(conn)
+ if 'workspace_execution_states' not in tables:
+ op.create_table(
+ 'workspace_execution_states',
+ sa.Column('workspace_uuid', sa.String(36), primary_key=True),
+ sa.Column('instance_uuid', sa.String(255), nullable=False),
+ sa.Column('active_generation', sa.BigInteger(), nullable=False, server_default='1'),
+ sa.Column('state', sa.String(32), nullable=False, server_default='active'),
+ sa.Column('write_fenced', sa.Boolean(), nullable=False, server_default=sa.false()),
+ sa.Column('source', sa.String(32), nullable=False, server_default='local'),
+ sa.Column('desired_state_revision', sa.BigInteger(), nullable=False, server_default='0'),
+ sa.Column('updated_at', sa.DateTime(), nullable=False, server_default=sa.func.now()),
+ sa.ForeignKeyConstraint(
+ ['workspace_uuid'],
+ ['workspaces.uuid'],
+ name='fk_workspace_execution_states_workspace',
+ ondelete='CASCADE',
+ ),
+ sa.CheckConstraint('active_generation > 0', name='ck_workspace_execution_generation'),
+ sa.CheckConstraint(
+ "state IN ('provisioning', 'active', 'migrating', 'draining', 'inactive')",
+ name='ck_workspace_execution_state',
+ ),
+ sa.CheckConstraint(
+ "source IN ('local', 'cloud')",
+ name='ck_workspace_execution_source',
+ ),
+ )
+ execution_indexes = _index_names(conn, 'workspace_execution_states')
+ if 'ix_workspace_execution_states_instance_state' not in execution_indexes:
+ op.create_index(
+ 'ix_workspace_execution_states_instance_state',
+ 'workspace_execution_states',
+ ['instance_uuid', 'state'],
+ )
+
+
+def _load_instance_uuid(conn: sa.Connection) -> str | None:
+ if 'metadata' not in _table_names(conn):
+ return None
+
+ metadata = sa.table(
+ 'metadata',
+ sa.column('key', sa.String(255)),
+ sa.column('value', sa.String(255)),
+ )
+ value = conn.execute(sa.select(metadata.c.value).where(metadata.c.key == 'instance_uuid')).scalar_one_or_none()
+ if not isinstance(value, str) or not value.strip():
+ return None
+ return value.strip()
+
+
+def _bootstrap_default_workspace(conn: sa.Connection) -> None:
+ required_tables = {'users', 'workspaces', 'workspace_memberships', 'workspace_execution_states'}
+ if not required_tables.issubset(_table_names(conn)):
+ return
+
+ instance_uuid = _load_instance_uuid(conn)
+ users_exist = 'users' in _table_names(conn) and bool(conn.execute(sa.text('SELECT 1 FROM users LIMIT 1')).first())
+ if instance_uuid is None:
+ if users_exist:
+ raise RuntimeError("Cannot bootstrap the default workspace without metadata['instance_uuid']")
+ return
+
+ workspaces = sa.table(
+ 'workspaces',
+ sa.column('uuid', sa.String(36)),
+ sa.column('instance_uuid', sa.String(255)),
+ sa.column('name', sa.String(255)),
+ sa.column('slug', sa.String(255)),
+ sa.column('type', sa.String(32)),
+ sa.column('status', sa.String(32)),
+ sa.column('created_by_account_uuid', sa.String(36)),
+ sa.column('source', sa.String(32)),
+ sa.column('projection_revision', sa.BigInteger()),
+ )
+ local_rows = conn.execute(
+ sa.select(workspaces.c.uuid).where(
+ workspaces.c.instance_uuid == instance_uuid,
+ workspaces.c.source == 'local',
+ )
+ ).all()
+ if len(local_rows) > 1:
+ raise RuntimeError(f'Multiple local workspaces already exist for instance {instance_uuid!r}')
+
+ users = sa.table(
+ 'users',
+ sa.column('id', sa.Integer()),
+ sa.column('uuid', sa.String(36)),
+ )
+ owner_account_uuid = None
+ if 'users' in _table_names(conn):
+ owner_account_uuid = conn.execute(sa.select(users.c.uuid).order_by(users.c.id).limit(1)).scalar_one_or_none()
+
+ if local_rows:
+ workspace_uuid = local_rows[0][0]
+ if owner_account_uuid is not None:
+ conn.execute(
+ workspaces.update()
+ .where(workspaces.c.uuid == workspace_uuid)
+ .where(workspaces.c.created_by_account_uuid.is_(None))
+ .values(created_by_account_uuid=owner_account_uuid)
+ )
+ else:
+ workspace_uuid = str(uuid.uuid4())
+ conn.execute(
+ workspaces.insert().values(
+ uuid=workspace_uuid,
+ instance_uuid=instance_uuid,
+ name='Default Workspace',
+ slug='default',
+ type='team',
+ status='active',
+ created_by_account_uuid=owner_account_uuid,
+ source='local',
+ projection_revision=0,
+ )
+ )
+
+ execution_states = sa.table(
+ 'workspace_execution_states',
+ sa.column('workspace_uuid', sa.String(36)),
+ sa.column('instance_uuid', sa.String(255)),
+ sa.column('active_generation', sa.BigInteger()),
+ sa.column('state', sa.String(32)),
+ sa.column('write_fenced', sa.Boolean()),
+ sa.column('source', sa.String(32)),
+ sa.column('desired_state_revision', sa.BigInteger()),
+ )
+ execution_state = conn.execute(
+ sa.select(
+ execution_states.c.instance_uuid,
+ execution_states.c.active_generation,
+ execution_states.c.state,
+ execution_states.c.write_fenced,
+ execution_states.c.source,
+ ).where(execution_states.c.workspace_uuid == workspace_uuid)
+ ).first()
+ if execution_state is None:
+ conn.execute(
+ execution_states.insert().values(
+ workspace_uuid=workspace_uuid,
+ instance_uuid=instance_uuid,
+ active_generation=1,
+ state='active',
+ write_fenced=False,
+ source='local',
+ desired_state_revision=0,
+ )
+ )
+ elif (
+ execution_state.instance_uuid != instance_uuid
+ or execution_state.active_generation != 1
+ or execution_state.state != 'active'
+ or execution_state.write_fenced
+ or execution_state.source != 'local'
+ ):
+ raise RuntimeError(f'Default workspace {workspace_uuid!r} has an invalid local execution state')
+
+ if owner_account_uuid is None:
+ return
+
+ memberships = sa.table(
+ 'workspace_memberships',
+ sa.column('uuid', sa.String(36)),
+ sa.column('workspace_uuid', sa.String(36)),
+ sa.column('account_uuid', sa.String(36)),
+ sa.column('role', sa.String(32)),
+ sa.column('status', sa.String(32)),
+ sa.column('joined_at', sa.DateTime()),
+ sa.column('projection_revision', sa.BigInteger()),
+ )
+ membership = conn.execute(
+ sa.select(memberships.c.uuid, memberships.c.joined_at).where(
+ memberships.c.workspace_uuid == workspace_uuid,
+ memberships.c.account_uuid == owner_account_uuid,
+ )
+ ).first()
+ now = datetime.datetime.now(datetime.UTC).replace(tzinfo=None)
+ if membership is None:
+ conn.execute(
+ memberships.insert().values(
+ uuid=str(uuid.uuid4()),
+ workspace_uuid=workspace_uuid,
+ account_uuid=owner_account_uuid,
+ role='owner',
+ status='active',
+ joined_at=now,
+ projection_revision=0,
+ )
+ )
+ else:
+ conn.execute(
+ memberships.update()
+ .where(memberships.c.uuid == membership.uuid)
+ .values(
+ role='owner',
+ status='active',
+ joined_at=membership.joined_at or now,
+ )
+ )
+
+
+def upgrade() -> None:
+ conn = op.get_bind()
+ _upgrade_users(conn)
+ _create_workspace_tables(conn)
+ _bootstrap_default_workspace(conn)
+
+
+def downgrade() -> None:
+ conn = op.get_bind()
+ tables = _table_names(conn)
+ for table_name in (
+ 'workspace_execution_states',
+ 'workspace_invitations',
+ 'workspace_memberships',
+ 'workspaces',
+ ):
+ if table_name in tables:
+ op.drop_table(table_name)
+
+ if 'users' not in _table_names(conn):
+ return
+
+ indexes = _index_names(conn, 'users')
+ if 'uq_users_uuid' in indexes:
+ op.drop_index('uq_users_uuid', table_name='users')
+
+ columns = _column_map(conn, 'users')
+ constraint_names = _constraint_names(conn, 'users')
+ with op.batch_alter_table('users') as batch_op:
+ # SQLite batch recreation otherwise preserves the named checks while
+ # dropping their referenced columns, producing ``no such column`` only
+ # after the Workspace directory tables have already been removed.
+ for constraint_name in ('ck_users_source', 'ck_users_status'):
+ if constraint_name in constraint_names:
+ batch_op.drop_constraint(constraint_name, type_='check')
+ for column_name in ('projection_revision', 'source', 'status', 'uuid'):
+ if column_name in columns:
+ batch_op.drop_column(column_name)
diff --git a/src/langbot/pkg/persistence/alembic/versions/0010_scope_tenant_resources.py b/src/langbot/pkg/persistence/alembic/versions/0010_scope_tenant_resources.py
new file mode 100644
index 000000000..9bcc1aac9
--- /dev/null
+++ b/src/langbot/pkg/persistence/alembic/versions/0010_scope_tenant_resources.py
@@ -0,0 +1,884 @@
+"""scope every tenant-owned resource to a workspace
+
+Revision ID: 0010_scope_resources
+Revises: 0009_workspace_tenancy
+Create Date: 2026-07-19
+
+This migration is intentionally expand/backfill/contract. Existing rows are
+bound to the single local Workspace created by revision 0009 before any
+non-null or scoped-key constraint is installed.
+"""
+
+from __future__ import annotations
+
+import hashlib
+import uuid
+
+import sqlalchemy as sa
+from alembic import op
+
+revision = '0010_scope_resources'
+down_revision = '0009_workspace_tenancy'
+branch_labels = None
+depends_on = None
+
+
+_TENANT_TABLES = (
+ 'api_keys',
+ 'bots',
+ 'bot_admins',
+ 'binary_storages',
+ 'mcp_servers',
+ 'model_providers',
+ 'llm_models',
+ 'embedding_models',
+ 'rerank_models',
+ 'legacy_pipelines',
+ 'pipeline_run_records',
+ 'plugin_settings',
+ 'knowledge_bases',
+ 'knowledge_base_files',
+ 'knowledge_base_chunks',
+ 'webhooks',
+ 'monitoring_messages',
+ 'monitoring_llm_calls',
+ 'monitoring_tool_calls',
+ 'monitoring_sessions',
+ 'monitoring_errors',
+ 'monitoring_embedding_calls',
+ 'monitoring_feedback',
+)
+
+_COMPOSITE_PRIMARY_KEYS = {
+ 'binary_storages': ('workspace_uuid', 'unique_key'),
+ 'plugin_settings': ('workspace_uuid', 'plugin_author', 'plugin_name'),
+ 'monitoring_sessions': ('workspace_uuid', 'session_id'),
+}
+
+_COMPOSITE_FOREIGN_KEYS = {
+ 'bot_admins': (
+ (
+ 'fk_bot_admins_workspace_bot',
+ ('workspace_uuid', 'bot_uuid'),
+ 'bots',
+ ('workspace_uuid', 'uuid'),
+ 'CASCADE',
+ ),
+ ),
+ 'llm_models': (
+ (
+ 'fk_llm_models_workspace_provider',
+ ('workspace_uuid', 'provider_uuid'),
+ 'model_providers',
+ ('workspace_uuid', 'uuid'),
+ None,
+ ),
+ ),
+ 'embedding_models': (
+ (
+ 'fk_embedding_models_workspace_provider',
+ ('workspace_uuid', 'provider_uuid'),
+ 'model_providers',
+ ('workspace_uuid', 'uuid'),
+ None,
+ ),
+ ),
+ 'rerank_models': (
+ (
+ 'fk_rerank_models_workspace_provider',
+ ('workspace_uuid', 'provider_uuid'),
+ 'model_providers',
+ ('workspace_uuid', 'uuid'),
+ None,
+ ),
+ ),
+ 'pipeline_run_records': (
+ (
+ 'fk_pipeline_run_records_workspace_pipeline',
+ ('workspace_uuid', 'pipeline_uuid'),
+ 'legacy_pipelines',
+ ('workspace_uuid', 'uuid'),
+ 'CASCADE',
+ ),
+ ),
+ 'knowledge_base_files': (
+ (
+ 'fk_knowledge_base_files_workspace_kb',
+ ('workspace_uuid', 'kb_id'),
+ 'knowledge_bases',
+ ('workspace_uuid', 'uuid'),
+ 'CASCADE',
+ ),
+ ),
+ 'knowledge_base_chunks': (
+ (
+ 'fk_knowledge_base_chunks_workspace_file',
+ ('workspace_uuid', 'file_id'),
+ 'knowledge_base_files',
+ ('workspace_uuid', 'uuid'),
+ 'CASCADE',
+ ),
+ ),
+}
+
+_SCOPED_INDEXES: dict[str, tuple[tuple[str, tuple[str, ...], bool, sa.TextClause | None], ...]] = {
+ 'api_keys': (
+ ('uq_api_keys_uuid', ('uuid',), True, None),
+ ('uq_api_keys_key_hash', ('key_hash',), True, None),
+ ('ix_api_keys_workspace_name', ('workspace_uuid', 'name'), False, None),
+ ('ix_api_keys_workspace_status', ('workspace_uuid', 'status'), False, None),
+ ),
+ 'bots': (
+ ('uq_bots_workspace_uuid', ('workspace_uuid', 'uuid'), True, None),
+ ('ix_bots_workspace_name', ('workspace_uuid', 'name'), False, None),
+ ('ix_bots_workspace_updated', ('workspace_uuid', 'updated_at'), False, None),
+ ),
+ 'bot_admins': (
+ (
+ 'uq_bot_admin',
+ ('workspace_uuid', 'bot_uuid', 'launcher_type', 'launcher_id'),
+ True,
+ None,
+ ),
+ ('ix_bot_admins_workspace_bot', ('workspace_uuid', 'bot_uuid'), False, None),
+ ),
+ 'binary_storages': (
+ (
+ 'ix_binary_storages_workspace_owner',
+ ('workspace_uuid', 'owner_type', 'owner'),
+ False,
+ None,
+ ),
+ ),
+ 'mcp_servers': (
+ ('uq_mcp_servers_workspace_name', ('workspace_uuid', 'name'), True, None),
+ ('ix_mcp_servers_workspace_enable', ('workspace_uuid', 'enable'), False, None),
+ ('ix_mcp_servers_workspace_updated', ('workspace_uuid', 'updated_at'), False, None),
+ ),
+ 'model_providers': (
+ ('uq_model_providers_workspace_uuid', ('workspace_uuid', 'uuid'), True, None),
+ ('ix_model_providers_workspace_name', ('workspace_uuid', 'name'), False, None),
+ ('ix_model_providers_workspace_requester', ('workspace_uuid', 'requester'), False, None),
+ ),
+ 'llm_models': (
+ ('ix_llm_models_workspace_provider', ('workspace_uuid', 'provider_uuid'), False, None),
+ ('ix_llm_models_workspace_name', ('workspace_uuid', 'name'), False, None),
+ ),
+ 'embedding_models': (
+ ('ix_embedding_models_workspace_provider', ('workspace_uuid', 'provider_uuid'), False, None),
+ ('ix_embedding_models_workspace_name', ('workspace_uuid', 'name'), False, None),
+ ),
+ 'rerank_models': (
+ ('ix_rerank_models_workspace_provider', ('workspace_uuid', 'provider_uuid'), False, None),
+ ('ix_rerank_models_workspace_name', ('workspace_uuid', 'name'), False, None),
+ ),
+ 'legacy_pipelines': (
+ ('uq_legacy_pipelines_workspace_uuid', ('workspace_uuid', 'uuid'), True, None),
+ ('ix_legacy_pipelines_workspace_name', ('workspace_uuid', 'name'), False, None),
+ ('ix_legacy_pipelines_workspace_default', ('workspace_uuid', 'is_default'), False, None),
+ ('ix_legacy_pipelines_workspace_updated', ('workspace_uuid', 'updated_at'), False, None),
+ ),
+ 'pipeline_run_records': (
+ (
+ 'ix_pipeline_run_records_workspace_pipeline',
+ ('workspace_uuid', 'pipeline_uuid'),
+ False,
+ None,
+ ),
+ (
+ 'ix_pipeline_run_records_workspace_created',
+ ('workspace_uuid', 'created_at'),
+ False,
+ None,
+ ),
+ ),
+ 'plugin_settings': (('ix_plugin_settings_workspace_enabled', ('workspace_uuid', 'enabled'), False, None),),
+ 'knowledge_bases': (
+ ('uq_knowledge_bases_workspace_uuid', ('workspace_uuid', 'uuid'), True, None),
+ ('ix_knowledge_bases_workspace_name', ('workspace_uuid', 'name'), False, None),
+ (
+ 'uq_knowledge_bases_workspace_collection',
+ ('workspace_uuid', 'collection_id'),
+ True,
+ sa.text('collection_id IS NOT NULL'),
+ ),
+ ),
+ 'knowledge_base_files': (
+ ('uq_knowledge_base_files_workspace_uuid', ('workspace_uuid', 'uuid'), True, None),
+ ('ix_knowledge_base_files_workspace_kb', ('workspace_uuid', 'kb_id'), False, None),
+ ),
+ 'knowledge_base_chunks': (('ix_knowledge_base_chunks_workspace_file', ('workspace_uuid', 'file_id'), False, None),),
+ 'webhooks': (
+ ('ix_webhooks_workspace_name', ('workspace_uuid', 'name'), False, None),
+ ('ix_webhooks_workspace_enabled', ('workspace_uuid', 'enabled'), False, None),
+ ('ix_webhooks_workspace_created', ('workspace_uuid', 'created_at'), False, None),
+ ),
+ 'monitoring_messages': (
+ ('ix_monitoring_messages_workspace_timestamp', ('workspace_uuid', 'timestamp'), False, None),
+ ('ix_monitoring_messages_workspace_bot', ('workspace_uuid', 'bot_id', 'timestamp'), False, None),
+ (
+ 'ix_monitoring_messages_workspace_pipeline',
+ ('workspace_uuid', 'pipeline_id', 'timestamp'),
+ False,
+ None,
+ ),
+ ('ix_monitoring_messages_workspace_session', ('workspace_uuid', 'session_id'), False, None),
+ ),
+ 'monitoring_llm_calls': (
+ ('ix_monitoring_llm_calls_workspace_timestamp', ('workspace_uuid', 'timestamp'), False, None),
+ ('ix_monitoring_llm_calls_workspace_session', ('workspace_uuid', 'session_id'), False, None),
+ ('ix_monitoring_llm_calls_workspace_message', ('workspace_uuid', 'message_id'), False, None),
+ ),
+ 'monitoring_tool_calls': (
+ ('ix_monitoring_tool_calls_workspace_timestamp', ('workspace_uuid', 'timestamp'), False, None),
+ ('ix_monitoring_tool_calls_workspace_session', ('workspace_uuid', 'session_id'), False, None),
+ ('ix_monitoring_tool_calls_workspace_message', ('workspace_uuid', 'message_id'), False, None),
+ ),
+ 'monitoring_sessions': (
+ ('ix_monitoring_sessions_workspace_activity', ('workspace_uuid', 'last_activity'), False, None),
+ ('ix_monitoring_sessions_workspace_active', ('workspace_uuid', 'is_active'), False, None),
+ ('ix_monitoring_sessions_workspace_bot', ('workspace_uuid', 'bot_id', 'last_activity'), False, None),
+ ),
+ 'monitoring_errors': (
+ ('ix_monitoring_errors_workspace_timestamp', ('workspace_uuid', 'timestamp'), False, None),
+ ('ix_monitoring_errors_workspace_session', ('workspace_uuid', 'session_id'), False, None),
+ ('ix_monitoring_errors_workspace_message', ('workspace_uuid', 'message_id'), False, None),
+ ),
+ 'monitoring_embedding_calls': (
+ (
+ 'ix_monitoring_embedding_calls_workspace_timestamp',
+ ('workspace_uuid', 'timestamp'),
+ False,
+ None,
+ ),
+ (
+ 'ix_monitoring_embedding_calls_workspace_kb',
+ ('workspace_uuid', 'knowledge_base_id'),
+ False,
+ None,
+ ),
+ (
+ 'ix_monitoring_embedding_calls_workspace_session',
+ ('workspace_uuid', 'session_id'),
+ False,
+ None,
+ ),
+ ),
+ 'monitoring_feedback': (
+ (
+ 'uq_monitoring_feedback_workspace_feedback_id',
+ ('workspace_uuid', 'feedback_id'),
+ True,
+ None,
+ ),
+ ('ix_monitoring_feedback_workspace_timestamp', ('workspace_uuid', 'timestamp'), False, None),
+ ('ix_monitoring_feedback_workspace_session', ('workspace_uuid', 'session_id'), False, None),
+ ('ix_monitoring_feedback_workspace_message', ('workspace_uuid', 'message_id'), False, None),
+ ),
+}
+
+
+def _inspector(conn: sa.Connection) -> sa.Inspector:
+ return sa.inspect(conn)
+
+
+def _table_names(conn: sa.Connection) -> set[str]:
+ return set(_inspector(conn).get_table_names())
+
+
+def _columns(conn: sa.Connection, table_name: str) -> dict[str, dict]:
+ return {column['name']: column for column in _inspector(conn).get_columns(table_name)}
+
+
+def _index_names(conn: sa.Connection, table_name: str) -> set[str]:
+ return {index['name'] for index in _inspector(conn).get_indexes(table_name)}
+
+
+def _unique_column_sets(conn: sa.Connection, table_name: str) -> set[tuple[str, ...]]:
+ inspector = _inspector(conn)
+ result = {
+ tuple(constraint.get('column_names') or ()) for constraint in inspector.get_unique_constraints(table_name)
+ }
+ result.update(
+ tuple(index.get('column_names') or ()) for index in inspector.get_indexes(table_name) if index.get('unique')
+ )
+ return result
+
+
+def _foreign_key_exists(
+ conn: sa.Connection,
+ table_name: str,
+ local_columns: tuple[str, ...],
+ referred_table: str,
+ referred_columns: tuple[str, ...],
+) -> bool:
+ return any(
+ tuple(foreign_key.get('constrained_columns') or ()) == local_columns
+ and foreign_key.get('referred_table') == referred_table
+ and tuple(foreign_key.get('referred_columns') or ()) == referred_columns
+ for foreign_key in _inspector(conn).get_foreign_keys(table_name)
+ )
+
+
+def _metadata_value(conn: sa.Connection, key: str) -> str | None:
+ if 'metadata' not in _table_names(conn):
+ return None
+ metadata = sa.table(
+ 'metadata',
+ sa.column('key', sa.String(255)),
+ sa.column('value', sa.String(255)),
+ )
+ value = conn.execute(sa.select(metadata.c.value).where(metadata.c.key == key)).scalar_one_or_none()
+ return value.strip() if isinstance(value, str) and value.strip() else None
+
+
+def _default_workspace_uuid(conn: sa.Connection) -> str | None:
+ if 'workspaces' not in _table_names(conn):
+ return None
+ workspaces = sa.table(
+ 'workspaces',
+ sa.column('uuid', sa.String(36)),
+ sa.column('instance_uuid', sa.String(255)),
+ sa.column('source', sa.String(32)),
+ )
+ instance_uuid = _metadata_value(conn, 'instance_uuid')
+ query = sa.select(workspaces.c.uuid).where(workspaces.c.source == 'local')
+ if instance_uuid is not None:
+ query = query.where(workspaces.c.instance_uuid == instance_uuid)
+ rows = conn.execute(query).all()
+ if len(rows) > 1:
+ raise RuntimeError('Cannot backfill tenant resources: multiple local Workspaces exist')
+ return rows[0][0] if rows else None
+
+
+def _upgrade_normalized_email(conn: sa.Connection) -> None:
+ if 'users' not in _table_names(conn):
+ return
+ columns = _columns(conn, 'users')
+ if 'normalized_email' not in columns:
+ op.add_column('users', sa.Column('normalized_email', sa.String(320), nullable=True))
+
+ users = sa.table(
+ 'users',
+ sa.column('id', sa.Integer()),
+ sa.column('user', sa.String(255)),
+ sa.column('normalized_email', sa.String(320)),
+ )
+ # Use the exact same normalization algorithm as the runtime. Database
+ # ``lower()`` is ASCII-only on SQLite and is not equivalent to Python
+ # ``casefold()`` (for example, Straße -> strasse). Recompute every row so
+ # an interrupted expand/backfill attempt using an older migration body can
+ # be resumed safely.
+ seen_emails: dict[str, int] = {}
+ for user_id, email in conn.execute(sa.select(users.c.id, users.c.user).order_by(users.c.id)).all():
+ normalized_email = str(email or '').strip().casefold()
+ if not normalized_email:
+ raise RuntimeError(f'Cannot normalize empty account identity for user row {user_id}')
+ if len(normalized_email) > 320:
+ raise RuntimeError(
+ f'Cannot normalize account identity for user row {user_id}: canonical value exceeds 320 characters'
+ )
+ duplicate_user_id = seen_emails.get(normalized_email)
+ if duplicate_user_id is not None:
+ raise RuntimeError(
+ f'Cannot create normalized account identity: user rows '
+ f'{duplicate_user_id} and {user_id} both normalize to {normalized_email!r}'
+ )
+ seen_emails[normalized_email] = user_id
+ conn.execute(users.update().where(users.c.id == user_id).values(normalized_email=normalized_email))
+
+ columns = _columns(conn, 'users')
+ checks = {
+ constraint.get('name'): constraint
+ for constraint in _inspector(conn).get_check_constraints('users')
+ if constraint.get('name') is not None
+ }
+ identity_check = checks.get('ck_users_normalized_email')
+ identity_check_sql = str((identity_check or {}).get('sqltext') or '').casefold()
+ # Python casefold is the canonical identity algorithm. SQL ``lower`` is
+ # dialect/locale dependent (notably Cherokee folds to uppercase in Python
+ # but PostgreSQL lowercases it), so the database validates only portable
+ # structural invariants and uniqueness.
+ replace_legacy_identity_check = identity_check is not None and 'lower' in identity_check_sql
+ needs_contract = columns['normalized_email']['nullable'] or identity_check is None or replace_legacy_identity_check
+ if needs_contract:
+ with op.batch_alter_table('users') as batch_op:
+ if replace_legacy_identity_check:
+ batch_op.drop_constraint('ck_users_normalized_email', type_='check')
+ if columns['normalized_email']['nullable']:
+ batch_op.alter_column(
+ 'normalized_email',
+ existing_type=columns['normalized_email']['type'],
+ nullable=False,
+ )
+ if identity_check is None or replace_legacy_identity_check:
+ batch_op.create_check_constraint(
+ 'ck_users_normalized_email',
+ 'normalized_email = trim(normalized_email) '
+ 'AND length(normalized_email) > 0 '
+ 'AND length(normalized_email) <= 320',
+ )
+ if 'uq_users_normalized_email' not in _index_names(conn, 'users'):
+ op.create_index('uq_users_normalized_email', 'users', ['normalized_email'], unique=True)
+
+
+def _api_key_owner(conn: sa.Connection, workspace_uuid: str | None) -> str | None:
+ if workspace_uuid is None or 'workspace_memberships' not in _table_names(conn):
+ return None
+ memberships = sa.table(
+ 'workspace_memberships',
+ sa.column('workspace_uuid', sa.String(36)),
+ sa.column('account_uuid', sa.String(36)),
+ sa.column('role', sa.String(32)),
+ )
+ return conn.execute(
+ sa.select(memberships.c.account_uuid)
+ .where(
+ memberships.c.workspace_uuid == workspace_uuid,
+ memberships.c.role == 'owner',
+ )
+ .limit(1)
+ ).scalar_one_or_none()
+
+
+def _expand_and_hash_api_keys(conn: sa.Connection, workspace_uuid: str | None) -> None:
+ if 'api_keys' not in _table_names(conn):
+ return
+ columns = _columns(conn, 'api_keys')
+ additions = (
+ ('uuid', sa.Column('uuid', sa.String(36), nullable=True)),
+ ('created_by_account_uuid', sa.Column('created_by_account_uuid', sa.String(36), nullable=True)),
+ ('key_hash', sa.Column('key_hash', sa.String(64), nullable=True)),
+ ('scopes', sa.Column('scopes', sa.JSON(), nullable=True)),
+ ('status', sa.Column('status', sa.String(32), nullable=True, server_default='active')),
+ ('expires_at', sa.Column('expires_at', sa.DateTime(), nullable=True)),
+ ('last_used_at', sa.Column('last_used_at', sa.DateTime(), nullable=True)),
+ )
+ for name, column in additions:
+ if name not in columns:
+ op.add_column('api_keys', column)
+
+ columns = _columns(conn, 'api_keys')
+ api_key_columns = [sa.column('id', sa.Integer())]
+ for name in ('uuid', 'created_by_account_uuid', 'key_hash', 'scopes', 'status', 'key'):
+ if name in columns:
+ api_key_columns.append(sa.column(name, columns[name]['type']))
+ api_keys = sa.table('api_keys', *api_key_columns)
+ owner_uuid = _api_key_owner(conn, workspace_uuid)
+ selected_columns = [api_keys.c.id, api_keys.c.uuid, api_keys.c.key_hash]
+ if 'key' in api_keys.c:
+ selected_columns.append(api_keys.c.key)
+ rows = conn.execute(sa.select(*selected_columns).order_by(api_keys.c.id)).mappings().all()
+ for row in rows:
+ values: dict[str, object] = {}
+ if not row['uuid']:
+ values['uuid'] = str(uuid.uuid4())
+ if not row['key_hash']:
+ plaintext = row.get('key')
+ if not isinstance(plaintext, str) or not plaintext:
+ raise RuntimeError(f'API key row {row["id"]} has no secret to hash')
+ values['key_hash'] = hashlib.sha256(plaintext.encode()).hexdigest()
+ if values:
+ conn.execute(api_keys.update().where(api_keys.c.id == row['id']).values(**values))
+ # Pre-tenancy API keys historically had unrestricted instance access. A
+ # wildcard preserves that behavior while binding it to the backfilled
+ # Workspace; new keys must store their requested explicit scopes.
+ conn.execute(api_keys.update().where(api_keys.c.scopes.is_(None)).values(scopes=['*']))
+ conn.execute(api_keys.update().where(api_keys.c.status.is_(None)).values(status='active'))
+ if owner_uuid is not None:
+ conn.execute(
+ api_keys.update()
+ .where(api_keys.c.created_by_account_uuid.is_(None))
+ .values(created_by_account_uuid=owner_uuid)
+ )
+
+ columns = _columns(conn, 'api_keys')
+ check_names = {constraint.get('name') for constraint in _inspector(conn).get_check_constraints('api_keys')}
+ existing_fks = _inspector(conn).get_foreign_keys('api_keys')
+ creator_fk_exists = any(
+ tuple(foreign_key.get('constrained_columns') or ()) == ('created_by_account_uuid',)
+ and foreign_key.get('referred_table') == 'users'
+ and tuple(foreign_key.get('referred_columns') or ()) == ('uuid',)
+ for foreign_key in existing_fks
+ )
+ has_legacy_key = 'key' in columns
+ needs_contract = (
+ any(columns[name]['nullable'] for name in ('uuid', 'key_hash', 'scopes', 'status'))
+ or 'ck_api_keys_status' not in check_names
+ or not creator_fk_exists
+ or has_legacy_key
+ )
+ if needs_contract:
+ naming = {'uq': 'uq_%(table_name)s_%(column_0_name)s'}
+ with op.batch_alter_table('api_keys', naming_convention=naming) as batch_op:
+ for name in ('uuid', 'key_hash', 'scopes', 'status'):
+ if columns[name]['nullable']:
+ batch_op.alter_column(name, existing_type=columns[name]['type'], nullable=False)
+ if 'ck_api_keys_status' not in check_names:
+ batch_op.create_check_constraint('ck_api_keys_status', "status IN ('active', 'revoked')")
+ if not creator_fk_exists:
+ batch_op.create_foreign_key(
+ 'fk_api_keys_created_by_account',
+ 'users',
+ ['created_by_account_uuid'],
+ ['uuid'],
+ ondelete='SET NULL',
+ )
+ if has_legacy_key:
+ # Dropping the column also removes its old global plaintext
+ # unique constraint/index during SQLite's batch rebuild.
+ batch_op.drop_column('key')
+
+
+def _expand_workspace_columns(conn: sa.Connection, workspace_uuid: str | None) -> None:
+ tables = _table_names(conn)
+ for table_name in _TENANT_TABLES:
+ if table_name not in tables:
+ continue
+ columns = _columns(conn, table_name)
+ if 'workspace_uuid' not in columns:
+ op.add_column(table_name, sa.Column('workspace_uuid', sa.String(36), nullable=True))
+ tenant_table = sa.table(table_name, sa.column('workspace_uuid', sa.String(36)))
+ null_count = conn.scalar(
+ sa.select(sa.func.count()).select_from(tenant_table).where(tenant_table.c.workspace_uuid.is_(None))
+ )
+ if null_count:
+ if workspace_uuid is None:
+ raise RuntimeError(f'Cannot backfill {table_name}: the instance has no unique local Workspace')
+ conn.execute(
+ tenant_table.update()
+ .where(tenant_table.c.workspace_uuid.is_(None))
+ .values(workspace_uuid=workspace_uuid)
+ )
+
+
+def _mark_legacy_vector_collections(conn: sa.Connection, workspace_uuid: str | None) -> None:
+ """Persist which pre-tenancy KBs must keep using ``collection_id``.
+
+ The marker is backfilled only when this migration introduces the column,
+ or resumes while that newly added column is still nullable. A fresh
+ schema already contains the non-null column with ``false`` as its default,
+ so knowledge bases created under the scoped-vector contract can never be
+ mistaken for legacy data during a later migration retry.
+ """
+
+ if 'knowledge_bases' not in _table_names(conn):
+ return
+ columns = _columns(conn, 'knowledge_bases')
+ introduced = 'legacy_vector_collection' not in columns
+ if introduced:
+ op.add_column(
+ 'knowledge_bases',
+ sa.Column('legacy_vector_collection', sa.Boolean(), nullable=True),
+ )
+ columns = _columns(conn, 'knowledge_bases')
+ needs_legacy_backfill = introduced or columns['legacy_vector_collection']['nullable']
+
+ knowledge_bases = sa.table(
+ 'knowledge_bases',
+ sa.column('collection_id', columns['collection_id']['type']),
+ sa.column('legacy_vector_collection', sa.Boolean()),
+ *((sa.column('workspace_uuid', columns['workspace_uuid']['type']),) if 'workspace_uuid' in columns else ()),
+ )
+ if needs_legacy_backfill and workspace_uuid is not None:
+ legacy_filter = sa.and_(
+ knowledge_bases.c.collection_id.is_not(None),
+ sa.func.length(sa.func.trim(knowledge_bases.c.collection_id)) > 0,
+ )
+ if 'workspace_uuid' in knowledge_bases.c:
+ # A partially migrated database may already have Workspace
+ # columns. Never mark a projected cloud row as legacy.
+ legacy_filter = sa.and_(
+ legacy_filter,
+ sa.or_(
+ knowledge_bases.c.workspace_uuid.is_(None),
+ knowledge_bases.c.workspace_uuid == workspace_uuid,
+ ),
+ )
+ conn.execute(knowledge_bases.update().where(legacy_filter).values(legacy_vector_collection=True))
+
+ conn.execute(
+ knowledge_bases.update()
+ .where(knowledge_bases.c.legacy_vector_collection.is_(None))
+ .values(legacy_vector_collection=False)
+ )
+ columns = _columns(conn, 'knowledge_bases')
+ if columns['legacy_vector_collection']['nullable']:
+ with op.batch_alter_table('knowledge_bases') as batch_op:
+ batch_op.alter_column(
+ 'legacy_vector_collection',
+ existing_type=columns['legacy_vector_collection']['type'],
+ nullable=False,
+ server_default=sa.false(),
+ )
+
+
+def _drop_legacy_uniqueness(conn: sa.Connection) -> None:
+ if 'bot_admins' in _table_names(conn):
+ for constraint in _inspector(conn).get_unique_constraints('bot_admins'):
+ if tuple(constraint.get('column_names') or ()) == ('bot_uuid', 'launcher_type', 'launcher_id'):
+ with op.batch_alter_table('bot_admins') as batch_op:
+ batch_op.drop_constraint(constraint['name'], type_='unique')
+ break
+
+ if 'monitoring_feedback' in _table_names(conn):
+ dropped_constraint = False
+ for constraint in _inspector(conn).get_unique_constraints('monitoring_feedback'):
+ if tuple(constraint.get('column_names') or ()) == ('feedback_id',):
+ convention = {'uq': 'uq_%(table_name)s_%(column_0_name)s'}
+ constraint_name = constraint.get('name') or 'uq_monitoring_feedback_feedback_id'
+ with op.batch_alter_table(
+ 'monitoring_feedback',
+ naming_convention=convention,
+ ) as batch_op:
+ batch_op.drop_constraint(constraint_name, type_='unique')
+ dropped_constraint = True
+ break
+ if not dropped_constraint:
+ for index in _inspector(conn).get_indexes('monitoring_feedback'):
+ if index.get('unique') and tuple(index.get('column_names') or ()) == ('feedback_id',):
+ op.drop_index(index['name'], table_name='monitoring_feedback')
+
+
+def _create_index_if_missing(
+ conn: sa.Connection,
+ table_name: str,
+ name: str,
+ columns: tuple[str, ...],
+ unique: bool,
+ predicate: sa.TextClause | None,
+) -> None:
+ if name in _index_names(conn, table_name):
+ return
+ if unique and predicate is None and columns in _unique_column_sets(conn, table_name):
+ return
+ kwargs = {}
+ if predicate is not None:
+ kwargs = {'sqlite_where': predicate, 'postgresql_where': predicate}
+ op.create_index(name, table_name, list(columns), unique=unique, **kwargs)
+
+
+def _validate_scoped_unique_data(conn: sa.Connection) -> None:
+ checks = (
+ ('mcp_servers', ('workspace_uuid', 'name')),
+ ('knowledge_bases', ('workspace_uuid', 'collection_id')),
+ )
+ for table_name, column_names in checks:
+ if table_name not in _table_names(conn):
+ continue
+ columns = _columns(conn, table_name)
+ if not all(column_name in columns for column_name in column_names):
+ continue
+ table = sa.table(
+ table_name,
+ *(sa.column(column_name, columns[column_name]['type']) for column_name in column_names),
+ )
+ group_columns = [table.c[column_name] for column_name in column_names]
+ query = sa.select(*group_columns, sa.func.count()).group_by(*group_columns).having(sa.func.count() > 1)
+ if column_names[-1] == 'collection_id':
+ query = query.where(group_columns[-1].is_not(None))
+ duplicate = conn.execute(query.limit(1)).first()
+ if duplicate is not None:
+ raise RuntimeError(
+ f'Cannot create scoped unique key on {table_name}{column_names}: duplicate {duplicate!r}'
+ )
+
+
+def _create_parent_and_scoped_indexes(conn: sa.Connection) -> None:
+ _validate_scoped_unique_data(conn)
+ tables = _table_names(conn)
+ for table_name, indexes in _SCOPED_INDEXES.items():
+ if table_name not in tables:
+ continue
+ available_columns = _columns(conn, table_name)
+ for name, columns, unique, predicate in indexes:
+ if all(column in available_columns for column in columns):
+ _create_index_if_missing(conn, table_name, name, columns, unique, predicate)
+
+
+def _contract_table(conn: sa.Connection, table_name: str) -> None:
+ columns = _columns(conn, table_name)
+ if 'workspace_uuid' not in columns:
+ return
+ direct_workspace_fk = _foreign_key_exists(
+ conn,
+ table_name,
+ ('workspace_uuid',),
+ 'workspaces',
+ ('uuid',),
+ )
+ current_pk = tuple(_inspector(conn).get_pk_constraint(table_name).get('constrained_columns') or ())
+ desired_pk = _COMPOSITE_PRIMARY_KEYS.get(table_name)
+ missing_composite_fks = [
+ foreign_key
+ for foreign_key in _COMPOSITE_FOREIGN_KEYS.get(table_name, ())
+ if not _foreign_key_exists(
+ conn,
+ table_name,
+ foreign_key[1],
+ foreign_key[2],
+ foreign_key[3],
+ )
+ ]
+ needs_contract = (
+ columns['workspace_uuid']['nullable']
+ or not direct_workspace_fk
+ or (desired_pk is not None and current_pk != desired_pk)
+ or bool(missing_composite_fks)
+ )
+ if not needs_contract:
+ return
+
+ naming = {
+ 'pk': 'pk_%(table_name)s',
+ 'fk': 'fk_%(table_name)s_%(column_0_name)s_%(referred_table_name)s',
+ }
+ pk_name = _inspector(conn).get_pk_constraint(table_name).get('name') or f'pk_{table_name}'
+ with op.batch_alter_table(table_name, naming_convention=naming) as batch_op:
+ if columns['workspace_uuid']['nullable']:
+ batch_op.alter_column(
+ 'workspace_uuid',
+ existing_type=columns['workspace_uuid']['type'],
+ nullable=False,
+ )
+ if desired_pk is not None and current_pk != desired_pk:
+ batch_op.drop_constraint(pk_name, type_='primary')
+ batch_op.create_primary_key(f'pk_{table_name}', list(desired_pk))
+ if not direct_workspace_fk:
+ batch_op.create_foreign_key(
+ f'fk_{table_name}_workspace',
+ 'workspaces',
+ ['workspace_uuid'],
+ ['uuid'],
+ ondelete='CASCADE',
+ )
+ for name, local_columns, referred_table, referred_columns, ondelete in missing_composite_fks:
+ batch_op.create_foreign_key(
+ name,
+ referred_table,
+ list(local_columns),
+ list(referred_columns),
+ ondelete=ondelete,
+ )
+
+
+def _contract_workspace_columns(conn: sa.Connection) -> None:
+ tables = _table_names(conn)
+ # Parents must be contracted before their children so SQLite can validate
+ # the exact composite target key during a batch-table rebuild.
+ order = (
+ 'api_keys',
+ 'bots',
+ 'bot_admins',
+ 'binary_storages',
+ 'mcp_servers',
+ 'model_providers',
+ 'llm_models',
+ 'embedding_models',
+ 'rerank_models',
+ 'legacy_pipelines',
+ 'pipeline_run_records',
+ 'plugin_settings',
+ 'knowledge_bases',
+ 'knowledge_base_files',
+ 'knowledge_base_chunks',
+ 'webhooks',
+ 'monitoring_messages',
+ 'monitoring_llm_calls',
+ 'monitoring_tool_calls',
+ 'monitoring_sessions',
+ 'monitoring_errors',
+ 'monitoring_embedding_calls',
+ 'monitoring_feedback',
+ )
+ for table_name in order:
+ if table_name in tables:
+ _contract_table(conn, table_name)
+
+
+def _migrate_workspace_metadata(conn: sa.Connection, workspace_uuid: str | None) -> None:
+ tables = _table_names(conn)
+ if 'workspaces' not in tables:
+ return
+ if 'workspace_metadata' not in tables:
+ op.create_table(
+ 'workspace_metadata',
+ sa.Column('workspace_uuid', sa.String(36), nullable=False),
+ sa.Column('key', sa.String(255), nullable=False),
+ sa.Column('value', sa.String(255), nullable=True),
+ sa.ForeignKeyConstraint(
+ ['workspace_uuid'],
+ ['workspaces.uuid'],
+ name='fk_workspace_metadata_workspace',
+ ondelete='CASCADE',
+ ),
+ sa.PrimaryKeyConstraint('workspace_uuid', 'key', name='pk_workspace_metadata'),
+ )
+ if workspace_uuid is None or 'metadata' not in tables:
+ return
+ metadata = sa.table(
+ 'metadata',
+ sa.column('key', sa.String(255)),
+ sa.column('value', sa.String(255)),
+ )
+ workspace_metadata = sa.table(
+ 'workspace_metadata',
+ sa.column('workspace_uuid', sa.String(36)),
+ sa.column('key', sa.String(255)),
+ sa.column('value', sa.String(255)),
+ )
+ tenant_keys = ('wizard_status', 'wizard_progress', 'rag_plugin_migration_needed')
+ rows = conn.execute(sa.select(metadata.c.key, metadata.c.value).where(metadata.c.key.in_(tenant_keys))).all()
+ for key, value in rows:
+ exists = conn.execute(
+ sa.select(workspace_metadata.c.key).where(
+ workspace_metadata.c.workspace_uuid == workspace_uuid,
+ workspace_metadata.c.key == key,
+ )
+ ).first()
+ if exists is None:
+ conn.execute(
+ workspace_metadata.insert().values(
+ workspace_uuid=workspace_uuid,
+ key=key,
+ value=value,
+ )
+ )
+ if rows:
+ conn.execute(metadata.delete().where(metadata.c.key.in_(tenant_keys)))
+
+
+def _validate_contract(conn: sa.Connection) -> None:
+ for table_name in _TENANT_TABLES:
+ if table_name not in _table_names(conn):
+ continue
+ columns = _columns(conn, table_name)
+ if 'workspace_uuid' not in columns or columns['workspace_uuid']['nullable']:
+ raise RuntimeError(f'{table_name}.workspace_uuid was not contracted to NOT NULL')
+ table = sa.table(table_name, sa.column('workspace_uuid', sa.String(36)))
+ if conn.scalar(sa.select(sa.func.count()).select_from(table).where(table.c.workspace_uuid.is_(None))):
+ raise RuntimeError(f'{table_name} still contains unscoped rows')
+ if conn.dialect.name == 'sqlite':
+ violations = conn.execute(sa.text('PRAGMA foreign_key_check')).all()
+ if violations:
+ raise RuntimeError(f'SQLite foreign key validation failed: {violations[:5]!r}')
+
+
+def upgrade() -> None:
+ conn = op.get_bind()
+ _upgrade_normalized_email(conn)
+ workspace_uuid = _default_workspace_uuid(conn)
+ _mark_legacy_vector_collections(conn, workspace_uuid)
+ _expand_workspace_columns(conn, workspace_uuid)
+ _expand_and_hash_api_keys(conn, workspace_uuid)
+ _drop_legacy_uniqueness(conn)
+ _create_parent_and_scoped_indexes(conn)
+ _contract_workspace_columns(conn)
+ _migrate_workspace_metadata(conn, workspace_uuid)
+ _validate_contract(conn)
+
+
+def downgrade() -> None:
+ raise RuntimeError(
+ '0010_scope_resources is intentionally irreversible because plaintext API key secrets were securely removed'
+ )
diff --git a/src/langbot/pkg/persistence/alembic_runner.py b/src/langbot/pkg/persistence/alembic_runner.py
index 74c2bac3d..b5daa174f 100644
--- a/src/langbot/pkg/persistence/alembic_runner.py
+++ b/src/langbot/pkg/persistence/alembic_runner.py
@@ -47,6 +47,12 @@ def _do_stamp(connection: Connection, revision: str = 'head') -> None:
command.stamp(cfg, revision)
+def _do_downgrade(connection: Connection, revision: str) -> None:
+ """Synchronous downgrade — runs inside run_sync."""
+ cfg = _build_config(connection)
+ command.downgrade(cfg, revision)
+
+
def _do_get_current(connection: Connection) -> str | None:
"""Get current alembic revision synchronously."""
ctx = MigrationContext.configure(connection)
@@ -73,6 +79,13 @@ async def run_alembic_stamp(async_engine: AsyncEngine, revision: str = 'head') -
await conn.commit()
+async def run_alembic_downgrade(async_engine: AsyncEngine, revision: str) -> None:
+ """Run Alembic downgrade to the given revision."""
+ async with async_engine.connect() as conn:
+ await conn.run_sync(_do_downgrade, revision)
+ await conn.commit()
+
+
async def get_alembic_current(async_engine: AsyncEngine) -> str | None:
"""Get current alembic revision, or None if not stamped."""
async with async_engine.connect() as conn:
@@ -121,6 +134,7 @@ if __name__ == '__main__':
print('Commands:')
print(' autogenerate "message" — Generate migration from ORM model diff')
print(' upgrade [revision] — Upgrade database (default: head)')
+ print(' downgrade — Downgrade database to a revision')
print(' stamp [revision] — Stamp revision without running (default: head)')
print(' current — Show current revision')
sys.exit(1)
@@ -140,6 +154,13 @@ if __name__ == '__main__':
rev = sys.argv[2] if len(sys.argv) > 2 else 'head'
asyncio.run(run_alembic_stamp(engine, rev))
print(f'Stamped: {rev}')
+ elif cmd == 'downgrade':
+ if len(sys.argv) < 3:
+ print('Usage: python -m langbot.pkg.persistence.alembic_runner downgrade ')
+ sys.exit(1)
+ rev = sys.argv[2]
+ asyncio.run(run_alembic_downgrade(engine, rev))
+ print(f'Downgraded to: {rev}')
elif cmd == 'current':
rev = asyncio.run(get_alembic_current(engine))
print(f'Current revision: {rev}')
diff --git a/src/langbot/pkg/persistence/mgr.py b/src/langbot/pkg/persistence/mgr.py
index 7ad7b9683..8454e83fd 100644
--- a/src/langbot/pkg/persistence/mgr.py
+++ b/src/langbot/pkg/persistence/mgr.py
@@ -1,14 +1,16 @@
from __future__ import annotations
import datetime
+import sqlite3
import typing
import sqlalchemy.ext.asyncio as sqlalchemy_asyncio
import sqlalchemy
-from . import database, migration
+from . import database, migration, sqlite_migration_backup
from ..entity.persistence import base, metadata, model as persistence_model
+from ..entity.persistence import workspace as persistence_workspace
from ..entity import persistence
from ..core import app
from ..utils import constants, importutil
@@ -19,6 +21,51 @@ importutil.import_modules_in_pkg(migrations)
importutil.import_modules_in_pkg(persistence)
+_ALEMBIC_TENANT_TABLES = {
+ 'workspaces',
+ 'workspace_memberships',
+ 'workspace_invitations',
+ 'workspace_execution_states',
+ 'workspace_metadata',
+ 'api_keys',
+ 'bots',
+ 'bot_admins',
+ 'binary_storages',
+ 'mcp_servers',
+ 'model_providers',
+ 'llm_models',
+ 'embedding_models',
+ 'rerank_models',
+ 'legacy_pipelines',
+ 'pipeline_run_records',
+ 'plugin_settings',
+ 'knowledge_bases',
+ 'knowledge_base_files',
+ 'knowledge_base_chunks',
+ 'webhooks',
+ 'monitoring_messages',
+ 'monitoring_llm_calls',
+ 'monitoring_tool_calls',
+ 'monitoring_sessions',
+ 'monitoring_errors',
+ 'monitoring_embedding_calls',
+ 'monitoring_feedback',
+}
+
+_PRE_WORKSPACE_ALEMBIC_REVISIONS = {
+ '0001_baseline',
+ '0002_sample',
+ '0003_add_rerank_models',
+ '0004_add_mcp_readme',
+ '0005_add_llm_context_length',
+ '0006_normalize_mcp_remote_mode',
+ '0007_add_bot_admins',
+ '0008_mcp_resource_prefs',
+}
+_WORKSPACE_ALEMBIC_REVISION = '0009_workspace_tenancy'
+_RESOURCE_SCOPE_ALEMBIC_REVISION = '0010_scope_resources'
+
+
class PersistenceManager:
"""Persistence module manager"""
@@ -42,6 +89,8 @@ class PersistenceManager:
await self.db.initialize()
break
+ self._enable_sqlite_foreign_keys()
+
await self.create_tables()
# run migrations
@@ -79,12 +128,33 @@ class PersistenceManager:
# Run Alembic migrations (new migration system)
await self._run_alembic_migrations()
+ # A legacy database may not contain tenant tables introduced by a
+ # newer release. They were deliberately deferred before 0009 because
+ # their Workspace/account FK targets did not exist yet; create them
+ # now that the tenancy contract is in place.
+ await self.create_tables()
+
await self.write_space_model_providers()
async def create_tables(self):
- # create tables
async with self.get_db_engine().connect() as conn:
- await conn.run_sync(self.meta.create_all)
+
+ def create_compatible_tables(sync_conn: sqlalchemy.Connection) -> None:
+ inspector = sqlalchemy.inspect(sync_conn)
+ existing_tables = set(inspector.get_table_names())
+ legacy_users = 'users' in existing_tables and (
+ 'uuid' not in {column['name'] for column in inspector.get_columns('users')}
+ or 'workspaces' not in existing_tables
+ )
+ # On a legacy installation, resource tables already exist
+ # without workspace_uuid and Workspace itself references the
+ # account UUID introduced by 0009. Alembic must expand those
+ # tables before SQLAlchemy may create any new tenant table.
+ excluded_tables = _ALEMBIC_TENANT_TABLES if legacy_users else set()
+ tables_to_create = [table for table in self.meta.sorted_tables if table.name not in excluded_tables]
+ self.meta.create_all(sync_conn, tables=tables_to_create)
+
+ await conn.run_sync(create_compatible_tables)
await conn.commit()
@@ -101,15 +171,80 @@ class PersistenceManager:
if row is None:
await self.execute_async(sqlalchemy.insert(metadata.Metadata).values(item))
+ await self._ensure_instance_uuid_metadata()
+
+ def _enable_sqlite_foreign_keys(self) -> None:
+ """Enable SQLite FK enforcement for every pooled runtime connection."""
+ engine = self.get_db_engine()
+ if engine.dialect.name != 'sqlite':
+ return
+ if getattr(self, '_sqlite_fk_listener_installed', False):
+ return
+
+ def set_sqlite_pragma(dbapi_connection, _connection_record) -> None:
+ # aiosqlite exposes the normal sqlite cursor API through its
+ # SQLAlchemy adapter. Guard the direct sqlite type too for tests.
+ if isinstance(dbapi_connection, sqlite3.Connection) or hasattr(dbapi_connection, 'cursor'):
+ cursor = dbapi_connection.cursor()
+ cursor.execute('PRAGMA foreign_keys=ON')
+ cursor.close()
+
+ sqlalchemy.event.listen(engine.sync_engine, 'connect', set_sqlite_pragma)
+ self._sqlite_fk_listener_installed = True
+
+ async def _ensure_instance_uuid_metadata(self) -> None:
+ """Persist the runtime instance identifier before tenant migrations run."""
+ runtime_instance_uuid = constants.instance_id.strip()
+ if not runtime_instance_uuid:
+ raise RuntimeError('LangBot instance UUID is empty before persistence initialization')
+
+ result = await self.execute_async(
+ sqlalchemy.select(metadata.Metadata.value).where(metadata.Metadata.key == 'instance_uuid')
+ )
+ persisted_instance_uuid = result.scalar_one_or_none()
+
+ if persisted_instance_uuid is None:
+ await self.execute_async(
+ sqlalchemy.insert(metadata.Metadata).values(key='instance_uuid', value=runtime_instance_uuid)
+ )
+ return
+
+ if persisted_instance_uuid != runtime_instance_uuid:
+ raise RuntimeError(
+ 'LangBot instance UUID does not match the value bound to this database: '
+ f'{runtime_instance_uuid!r} != {persisted_instance_uuid!r}'
+ )
+
async def write_space_model_providers(self):
+ if constants.edition != 'community':
+ # SaaS Workspace/provider linkage is explicit control-plane state;
+ # a process-level compatibility provider must never be projected
+ # into an arbitrary cloud Workspace.
+ return
+
space_models_gateway_api_url = self.ap.instance_config.data.get('space', {}).get(
'models_gateway_api_url', 'https://api.langbot.cloud/v1'
)
- # write space model providers
+ workspace_result = await self.execute_async(
+ sqlalchemy.select(persistence_workspace.Workspace.uuid).where(
+ persistence_workspace.Workspace.instance_uuid == constants.instance_id,
+ persistence_workspace.Workspace.source == persistence_workspace.WorkspaceSource.LOCAL.value,
+ )
+ )
+ workspace_uuids = workspace_result.scalars().all()
+ if len(workspace_uuids) != 1:
+ raise RuntimeError(
+ f'The fixed LangBot Models provider requires exactly one local Workspace; found {len(workspace_uuids)}'
+ )
+ workspace_uuid = workspace_uuids[0]
+
+ # The compatibility Space provider belongs to the OSS singleton
+ # Workspace. It must never be discovered or inserted globally.
result = await self.execute_async(
sqlalchemy.select(persistence_model.ModelProvider).where(
- persistence_model.ModelProvider.requester == 'space-chat-completions'
+ persistence_model.ModelProvider.workspace_uuid == workspace_uuid,
+ persistence_model.ModelProvider.requester == 'space-chat-completions',
)
)
exists_space_chat_completions_model_provider = result.first()
@@ -119,6 +254,7 @@ class PersistenceManager:
self.ap.logger.info('Creating space model providers...')
space_chat_completions_model_provider = {
'uuid': '00000000-0000-0000-0000-000000000000',
+ 'workspace_uuid': workspace_uuid,
'name': 'LangBot Models',
'requester': 'space-chat-completions',
'base_url': space_models_gateway_api_url,
@@ -132,7 +268,10 @@ class PersistenceManager:
if exists_space_chat_completions_model_provider.base_url != space_models_gateway_api_url:
await self.execute_async(
sqlalchemy.update(persistence_model.ModelProvider)
- .where(persistence_model.ModelProvider.uuid == exists_space_chat_completions_model_provider.uuid)
+ .where(
+ persistence_model.ModelProvider.workspace_uuid == workspace_uuid,
+ persistence_model.ModelProvider.uuid == exists_space_chat_completions_model_provider.uuid,
+ )
.values({'base_url': space_models_gateway_api_url})
)
@@ -153,13 +292,67 @@ class PersistenceManager:
await alembic_runner.run_alembic_stamp(engine, '0001_baseline')
current_rev = '0001_baseline'
- # Upgrade to head
+ if engine.dialect.name == 'sqlite':
+ if current_rev in _PRE_WORKSPACE_ALEMBIC_REVISIONS:
+ await self._run_verified_sqlite_migration(
+ engine,
+ source_revision=current_rev,
+ target_revision=_WORKSPACE_ALEMBIC_REVISION,
+ )
+ current_rev = await alembic_runner.get_alembic_current(engine)
+ if current_rev == _WORKSPACE_ALEMBIC_REVISION:
+ await self._run_verified_sqlite_migration(
+ engine,
+ source_revision=current_rev,
+ target_revision=_RESOURCE_SCOPE_ALEMBIC_REVISION,
+ )
+
+ # PostgreSQL has transactional DDL. SQLite has already crossed the
+ # two destructive tenancy boundaries under verified backups; this
+ # final call is a no-op today and applies future migrations.
await alembic_runner.run_alembic_upgrade(engine, 'head')
self.ap.logger.info('Alembic migrations completed.')
except Exception as e:
self.ap.logger.error(f'Alembic migration failed: {e}', exc_info=True)
raise
+ async def _run_verified_sqlite_migration(
+ self,
+ engine: sqlalchemy_asyncio.AsyncEngine,
+ *,
+ source_revision: str,
+ target_revision: str,
+ ) -> None:
+ from . import alembic_runner
+
+ backup = await sqlite_migration_backup.create_verified_backup(
+ engine,
+ source_revision=source_revision,
+ target_revision=target_revision,
+ )
+ self.ap.logger.info(f'Created verified SQLite migration backup {backup.backup_path} before {target_revision}.')
+ try:
+ await alembic_runner.run_alembic_upgrade(engine, target_revision)
+ completed_revision = await alembic_runner.get_alembic_current(engine)
+ if completed_revision != target_revision:
+ raise RuntimeError(f'Alembic stopped at {completed_revision!r}, expected {target_revision!r}')
+ await sqlite_migration_backup.mark_migration_succeeded(
+ backup,
+ completed_revision=completed_revision,
+ )
+ except BaseException:
+ await sqlite_migration_backup.restore_verified_backup(engine, backup)
+ restored_revision = await alembic_runner.get_alembic_current(engine)
+ if restored_revision != source_revision:
+ raise RuntimeError(
+ f'SQLite migration recovery restored revision {restored_revision!r}, expected {source_revision!r}'
+ )
+ self.ap.logger.error(
+ f'SQLite migration to {target_revision} failed; restored verified backup '
+ f'{backup.backup_path} at revision {source_revision}.'
+ )
+ raise
+
async def execute_async(self, *args, **kwargs) -> sqlalchemy.engine.cursor.CursorResult:
async with self.get_db_engine().connect() as conn:
result = await conn.execute(*args, **kwargs)
diff --git a/src/langbot/pkg/persistence/sqlite_migration_backup.py b/src/langbot/pkg/persistence/sqlite_migration_backup.py
new file mode 100644
index 000000000..5e1f7e683
--- /dev/null
+++ b/src/langbot/pkg/persistence/sqlite_migration_backup.py
@@ -0,0 +1,272 @@
+"""Durable SQLite backups for destructive Alembic migration boundaries."""
+
+from __future__ import annotations
+
+import asyncio
+import dataclasses
+import datetime
+import json
+import os
+import pathlib
+import re
+import secrets
+import sqlite3
+import tempfile
+import typing
+
+from sqlalchemy.ext.asyncio import AsyncEngine
+
+
+class SQLiteMigrationBackupError(RuntimeError):
+ """A verified migration backup could not be created or restored."""
+
+
+@dataclasses.dataclass(frozen=True, slots=True)
+class SQLiteMigrationBackup:
+ database_path: pathlib.Path
+ backup_path: pathlib.Path
+ manifest_path: pathlib.Path
+ source_revision: str
+ target_revision: str
+ created_at: str
+
+
+def _safe_label(value: str) -> str:
+ label = re.sub(r'[^A-Za-z0-9_.-]+', '-', value).strip('-')
+ return label or 'unknown'
+
+
+def _database_path(engine: AsyncEngine) -> pathlib.Path:
+ if engine.dialect.name != 'sqlite':
+ raise SQLiteMigrationBackupError('SQLite migration backups require a SQLite engine')
+ database = engine.url.database
+ if not database or database == ':memory:' or engine.url.query.get('mode') == 'memory':
+ raise SQLiteMigrationBackupError('Tenant schema migrations require a file-backed SQLite database for recovery')
+ database_path = pathlib.Path(database).expanduser()
+ if not database_path.is_absolute():
+ database_path = pathlib.Path.cwd() / database_path
+ database_path = database_path.resolve()
+ if not database_path.is_file():
+ raise SQLiteMigrationBackupError(f'SQLite database does not exist: {database_path}')
+ return database_path
+
+
+def _open_read_only(path: pathlib.Path) -> sqlite3.Connection:
+ return sqlite3.connect(f'{path.as_uri()}?mode=ro', uri=True, timeout=30)
+
+
+def _read_revision(connection: sqlite3.Connection) -> str | None:
+ has_version_table = connection.execute(
+ "SELECT 1 FROM sqlite_master WHERE type = 'table' AND name = 'alembic_version'"
+ ).fetchone()
+ if has_version_table is None:
+ return None
+ rows = connection.execute('SELECT version_num FROM alembic_version').fetchall()
+ if not rows:
+ return None
+ if len(rows) != 1 or not isinstance(rows[0][0], str):
+ raise SQLiteMigrationBackupError('SQLite backup has an invalid Alembic revision table')
+ return rows[0][0]
+
+
+def _verify_connection(connection: sqlite3.Connection, expected_revision: str) -> None:
+ quick_check = connection.execute('PRAGMA quick_check').fetchall()
+ if quick_check != [('ok',)]:
+ raise SQLiteMigrationBackupError(f'SQLite quick_check failed: {quick_check[:5]!r}')
+ actual_revision = _read_revision(connection)
+ if actual_revision != expected_revision:
+ raise SQLiteMigrationBackupError(
+ f'SQLite backup revision mismatch: {actual_revision!r} != {expected_revision!r}'
+ )
+
+
+def _verify_file(path: pathlib.Path, expected_revision: str) -> None:
+ with _open_read_only(path) as connection:
+ _verify_connection(connection, expected_revision)
+
+
+def _write_manifest(backup: SQLiteMigrationBackup, status: str, **extra: typing.Any) -> None:
+ payload: dict[str, typing.Any] = {
+ 'version': 1,
+ 'status': status,
+ 'created_at': backup.created_at,
+ 'database_path': str(backup.database_path),
+ 'backup_path': str(backup.backup_path),
+ 'source_revision': backup.source_revision,
+ 'target_revision': backup.target_revision,
+ 'quick_check': 'ok',
+ **extra,
+ }
+ backup.manifest_path.parent.mkdir(mode=0o700, parents=True, exist_ok=True)
+ descriptor, temporary_name = tempfile.mkstemp(
+ prefix=f'.{backup.manifest_path.name}.',
+ suffix='.tmp',
+ dir=backup.manifest_path.parent,
+ )
+ temporary_path = pathlib.Path(temporary_name)
+ try:
+ with os.fdopen(descriptor, 'w', encoding='utf-8') as file:
+ json.dump(payload, file, ensure_ascii=False, indent=2, sort_keys=True)
+ file.write('\n')
+ file.flush()
+ os.fsync(file.fileno())
+ os.chmod(temporary_path, 0o600)
+ os.replace(temporary_path, backup.manifest_path)
+ _fsync_directory(backup.manifest_path.parent)
+ finally:
+ temporary_path.unlink(missing_ok=True)
+
+
+def _fsync_file(path: pathlib.Path) -> None:
+ descriptor = os.open(path, os.O_RDONLY)
+ try:
+ os.fsync(descriptor)
+ finally:
+ os.close(descriptor)
+
+
+def _fsync_directory(path: pathlib.Path) -> None:
+ descriptor = os.open(path, os.O_RDONLY)
+ try:
+ os.fsync(descriptor)
+ finally:
+ os.close(descriptor)
+
+
+def _create_backup(
+ database_path: pathlib.Path,
+ source_revision: str,
+ target_revision: str,
+) -> SQLiteMigrationBackup:
+ backup_directory = database_path.parent / 'migration-backups'
+ backup_directory.mkdir(mode=0o700, parents=True, exist_ok=True)
+ os.chmod(backup_directory, 0o700)
+ created_at = datetime.datetime.now(datetime.UTC).strftime('%Y-%m-%dT%H-%M-%S.%fZ')
+ stem = (
+ f'{database_path.stem}-pre-{_safe_label(target_revision)}-'
+ f'from-{_safe_label(source_revision)}-{created_at}-{secrets.token_hex(4)}'
+ )
+ backup_path = backup_directory / f'{stem}.sqlite3'
+ manifest_path = backup_directory / f'{stem}.json'
+ descriptor, temporary_name = tempfile.mkstemp(
+ prefix=f'.{stem}.',
+ suffix='.creating',
+ dir=backup_directory,
+ )
+ os.close(descriptor)
+ temporary_path = pathlib.Path(temporary_name)
+ try:
+ with (
+ _open_read_only(database_path) as source,
+ sqlite3.connect(
+ temporary_path,
+ timeout=30,
+ ) as destination,
+ ):
+ source.execute('PRAGMA busy_timeout = 30000')
+ source.backup(destination)
+ destination.commit()
+ _verify_connection(destination, source_revision)
+ os.chmod(temporary_path, 0o600)
+ _fsync_file(temporary_path)
+ os.replace(temporary_path, backup_path)
+ _fsync_file(backup_path)
+ _fsync_directory(backup_directory)
+ backup = SQLiteMigrationBackup(
+ database_path=database_path,
+ backup_path=backup_path,
+ manifest_path=manifest_path,
+ source_revision=source_revision,
+ target_revision=target_revision,
+ created_at=created_at,
+ )
+ _write_manifest(backup, 'verified')
+ return backup
+ except Exception:
+ backup_path.unlink(missing_ok=True)
+ manifest_path.unlink(missing_ok=True)
+ raise
+ finally:
+ temporary_path.unlink(missing_ok=True)
+
+
+async def create_verified_backup(
+ engine: AsyncEngine,
+ *,
+ source_revision: str,
+ target_revision: str,
+) -> SQLiteMigrationBackup:
+ """Create and verify an online-consistent backup next to instance data."""
+
+ database_path = _database_path(engine)
+ return await asyncio.to_thread(
+ _create_backup,
+ database_path,
+ source_revision,
+ target_revision,
+ )
+
+
+def _restore_backup(backup: SQLiteMigrationBackup) -> None:
+ _verify_file(backup.backup_path, backup.source_revision)
+ descriptor, temporary_name = tempfile.mkstemp(
+ prefix=f'.{backup.database_path.name}.',
+ suffix='.restoring',
+ dir=backup.database_path.parent,
+ )
+ os.close(descriptor)
+ temporary_path = pathlib.Path(temporary_name)
+ try:
+ with (
+ _open_read_only(backup.backup_path) as source,
+ sqlite3.connect(
+ temporary_path,
+ timeout=30,
+ ) as destination,
+ ):
+ source.backup(destination)
+ destination.commit()
+ _verify_connection(destination, backup.source_revision)
+ os.chmod(temporary_path, 0o600)
+ _fsync_file(temporary_path)
+
+ # A stale WAL could replay pages from the failed migration after the
+ # main database file is replaced. The engine is disposed before this
+ # function runs, so these exact sidecars are safe to remove.
+ for suffix in ('-wal', '-shm', '-journal'):
+ pathlib.Path(f'{backup.database_path}{suffix}').unlink(missing_ok=True)
+ os.replace(temporary_path, backup.database_path)
+ _fsync_file(backup.database_path)
+ _fsync_directory(backup.database_path.parent)
+ _verify_file(backup.database_path, backup.source_revision)
+ finally:
+ temporary_path.unlink(missing_ok=True)
+
+
+async def restore_verified_backup(engine: AsyncEngine, backup: SQLiteMigrationBackup) -> None:
+ """Atomically restore a verified backup after a migration failure."""
+
+ await engine.dispose()
+ await asyncio.to_thread(_restore_backup, backup)
+ await asyncio.to_thread(
+ _write_manifest,
+ backup,
+ 'restored_after_failure',
+ restored_at=datetime.datetime.now(datetime.UTC).isoformat(),
+ )
+
+
+async def mark_migration_succeeded(
+ backup: SQLiteMigrationBackup,
+ *,
+ completed_revision: str,
+) -> None:
+ """Mark a retained verified backup after its migration boundary succeeds."""
+
+ await asyncio.to_thread(
+ _write_manifest,
+ backup,
+ 'migration_succeeded',
+ completed_at=datetime.datetime.now(datetime.UTC).isoformat(),
+ completed_revision=completed_revision,
+ )
diff --git a/src/langbot/pkg/pipeline/aggregator.py b/src/langbot/pkg/pipeline/aggregator.py
index 96358e329..c1b05e6e8 100644
--- a/src/langbot/pkg/pipeline/aggregator.py
+++ b/src/langbot/pkg/pipeline/aggregator.py
@@ -1,10 +1,4 @@
-"""Message Aggregator Module
-
-This module provides message aggregation/debounce functionality.
-When users send multiple messages consecutively, the aggregator will wait
-for a configurable delay period and merge them into a single message
-before processing.
-"""
+"""Workspace-scoped message aggregation and debounce support."""
from __future__ import annotations
@@ -13,96 +7,114 @@ import time
import typing
from dataclasses import dataclass, field
-import langbot_plugin.api.entities.builtin.platform.message as platform_message
-import langbot_plugin.api.entities.builtin.platform.events as platform_events
-import langbot_plugin.api.entities.builtin.provider.session as provider_session
import langbot_plugin.api.definition.abstract.platform.adapter as abstract_platform_adapter
+import langbot_plugin.api.entities.builtin.platform.events as platform_events
+import langbot_plugin.api.entities.builtin.platform.message as platform_message
+import langbot_plugin.api.entities.builtin.provider.session as provider_session
+
+from ..api.http.context import ExecutionContext
+from .pool import ExecutionContextMismatchError
+from ..workspace.errors import WorkspaceError, WorkspaceInvariantError
if typing.TYPE_CHECKING:
from ..core import app
-# Maximum number of messages to buffer before forcing a flush
MAX_BUFFER_MESSAGES = 10
+AggregationKey = tuple[
+ str,
+ str,
+ int,
+ str,
+ str | None,
+ str,
+ int | str,
+]
+
@dataclass
class PendingMessage:
- """A pending message waiting to be aggregated"""
+ """A pending message carrying its trusted execution scope."""
+ execution_context: ExecutionContext
bot_uuid: str
launcher_type: provider_session.LauncherTypes
- launcher_id: typing.Union[int, str]
- sender_id: typing.Union[int, str]
+ launcher_id: int | str
+ sender_id: int | str
message_event: platform_events.MessageEvent
message_chain: platform_message.MessageChain
adapter: abstract_platform_adapter.AbstractMessagePlatformAdapter
- pipeline_uuid: typing.Optional[str]
+ pipeline_uuid: str | None
routed_by_rule: bool = False
timestamp: float = field(default_factory=time.time)
@dataclass
class SessionBuffer:
- """Buffer for a single session's pending messages"""
+ """Pending messages for one scoped aggregation key."""
- session_id: str
+ aggregation_key: AggregationKey
+ execution_context: ExecutionContext
messages: list[PendingMessage] = field(default_factory=list)
- timer_task: typing.Optional[asyncio.Task] = None
+ timer_task: asyncio.Task | None = None
last_message_time: float = field(default_factory=time.time)
class MessageAggregator:
- """Message aggregator that buffers and merges consecutive messages
-
- This class implements a debounce mechanism for incoming messages.
- When a message arrives, it starts a timer. If more messages arrive
- before the timer expires, they are buffered. When the timer expires,
- all buffered messages are merged and sent to the query pool.
- """
+ """Debounce consecutive messages without crossing Workspace boundaries."""
ap: app.Application
-
- buffers: dict[str, SessionBuffer]
- """Session ID -> SessionBuffer mapping"""
-
+ buffers: dict[AggregationKey, SessionBuffer]
lock: asyncio.Lock
- """Lock for thread-safe buffer operations"""
def __init__(self, ap: app.Application):
self.ap = ap
self.buffers = {}
self.lock = asyncio.Lock()
- def _get_session_id(
+ def _get_aggregation_key(
self,
+ execution_context: ExecutionContext,
bot_uuid: str,
launcher_type: provider_session.LauncherTypes,
- launcher_id: typing.Union[int, str],
- ) -> str:
- """Generate a unique session ID"""
- return f'{bot_uuid}:{launcher_type.value}:{launcher_id}'
+ launcher_id: int | str,
+ pipeline_uuid: str | None,
+ ) -> AggregationKey:
+ """Build a key that cannot alias another Workspace, bot, or pipeline."""
- async def _get_aggregation_config(self, pipeline_uuid: typing.Optional[str]) -> tuple[bool, float]:
- """Get aggregation configuration for a pipeline
+ return (
+ execution_context.instance_uuid,
+ execution_context.workspace_uuid,
+ execution_context.placement_generation,
+ bot_uuid,
+ pipeline_uuid,
+ launcher_type.value,
+ launcher_id,
+ )
+
+ async def _get_aggregation_config(
+ self,
+ execution_context: ExecutionContext,
+ pipeline_uuid: str | None,
+ ) -> tuple[bool, float]:
+ """Return aggregation enablement and a clamped debounce delay."""
- Returns:
- tuple: (enabled, delay_seconds)
- """
default_enabled = False
default_delay = 1.5
if pipeline_uuid is None:
return default_enabled, default_delay
- # Get pipeline from pipeline manager
- pipeline = await self.ap.pipeline_mgr.get_pipeline_by_uuid(pipeline_uuid)
+ pipeline = await self.ap.pipeline_mgr.get_pipeline_by_uuid(
+ execution_context,
+ pipeline_uuid,
+ )
if pipeline is None:
return default_enabled, default_delay
config = pipeline.pipeline_entity.config or {}
trigger_config = config.get('trigger', {})
aggregation_config = trigger_config.get('message-aggregation', {})
-
enabled = aggregation_config.get('enabled', default_enabled)
delay_raw = aggregation_config.get('delay', default_delay)
@@ -111,33 +123,31 @@ class MessageAggregator:
except (TypeError, ValueError):
delay = default_delay
- # Clamp delay to valid range
- delay = max(1.0, min(10.0, delay))
-
- return enabled, delay
+ return enabled, max(1.0, min(10.0, delay))
async def add_message(
self,
bot_uuid: str,
launcher_type: provider_session.LauncherTypes,
- launcher_id: typing.Union[int, str],
- sender_id: typing.Union[int, str],
+ launcher_id: int | str,
+ sender_id: int | str,
message_event: platform_events.MessageEvent,
message_chain: platform_message.MessageChain,
adapter: abstract_platform_adapter.AbstractMessagePlatformAdapter,
- pipeline_uuid: typing.Optional[str] = None,
+ pipeline_uuid: str | None = None,
routed_by_rule: bool = False,
+ execution_context: ExecutionContext | None = None,
) -> None:
- """Add a message to the aggregation buffer
+ """Buffer or directly enqueue a message in its trusted Workspace."""
- If aggregation is disabled for the pipeline, the message is sent
- directly to the query pool. Otherwise, it's buffered and will be
- merged with other messages from the same session.
- """
- enabled, delay = await self._get_aggregation_config(pipeline_uuid)
+ execution_context = await self.ap.query_pool.resolve_execution_context(
+ execution_context,
+ bot_uuid=bot_uuid,
+ pipeline_uuid=pipeline_uuid,
+ )
+ enabled, delay = await self._get_aggregation_config(execution_context, pipeline_uuid)
if not enabled:
- # Aggregation disabled, send directly to query pool
await self.ap.query_pool.add_query(
bot_uuid=bot_uuid,
launcher_type=launcher_type,
@@ -148,12 +158,19 @@ class MessageAggregator:
adapter=adapter,
pipeline_uuid=pipeline_uuid,
routed_by_rule=routed_by_rule,
+ execution_context=execution_context,
)
return
- session_id = self._get_session_id(bot_uuid, launcher_type, launcher_id)
-
+ aggregation_key = self._get_aggregation_key(
+ execution_context,
+ bot_uuid,
+ launcher_type,
+ launcher_id,
+ pipeline_uuid,
+ )
pending_msg = PendingMessage(
+ execution_context=execution_context,
bot_uuid=bot_uuid,
launcher_type=launcher_type,
launcher_id=launcher_id,
@@ -167,106 +184,122 @@ class MessageAggregator:
force_flush = False
async with self.lock:
- if session_id in self.buffers:
- buffer = self.buffers[session_id]
- # Cancel existing timer (just cancel, don't await inside lock)
+ buffer = self.buffers.get(aggregation_key)
+ if buffer is None:
+ buffer = SessionBuffer(
+ aggregation_key=aggregation_key,
+ execution_context=execution_context,
+ messages=[pending_msg],
+ )
+ self.buffers[aggregation_key] = buffer
+ else:
+ if buffer.execution_context != execution_context:
+ raise ExecutionContextMismatchError('Aggregation buffer ExecutionContext changed for the same key')
if buffer.timer_task and not buffer.timer_task.done():
buffer.timer_task.cancel()
buffer.messages.append(pending_msg)
- else:
- buffer = SessionBuffer(
- session_id=session_id,
- messages=[pending_msg],
- )
- self.buffers[session_id] = buffer
buffer.last_message_time = time.time()
-
- # Check if buffer reached max capacity
if len(buffer.messages) >= MAX_BUFFER_MESSAGES:
force_flush = True
else:
- # Start new timer
- buffer.timer_task = asyncio.create_task(self._delayed_flush(session_id, delay))
+ buffer.timer_task = asyncio.create_task(self._delayed_flush(aggregation_key, delay, execution_context))
if force_flush:
- await self._flush_buffer(session_id)
+ await self._flush_buffer(aggregation_key, execution_context)
+
+ async def _delayed_flush(
+ self,
+ aggregation_key: AggregationKey,
+ delay: float,
+ execution_context: ExecutionContext,
+ ) -> None:
+ """Flush after the debounce delay using the captured context."""
- async def _delayed_flush(self, session_id: str, delay: float) -> None:
- """Wait for delay then flush the buffer"""
try:
await asyncio.sleep(delay)
- await self._flush_buffer(session_id)
+ await self._flush_buffer(aggregation_key, execution_context)
except asyncio.CancelledError:
- # Timer was cancelled, new message arrived
pass
-
- async def _flush_buffer(self, session_id: str) -> None:
- """Flush the buffer for a session, merging all messages"""
- async with self.lock:
- buffer = self.buffers.pop(session_id, None)
-
- if buffer is None or not buffer.messages:
- return
-
- if len(buffer.messages) == 1:
- # Only one message, no need to merge
- msg = buffer.messages[0]
- await self.ap.query_pool.add_query(
- bot_uuid=msg.bot_uuid,
- launcher_type=msg.launcher_type,
- launcher_id=msg.launcher_id,
- sender_id=msg.sender_id,
- message_event=msg.message_event,
- message_chain=msg.message_chain,
- adapter=msg.adapter,
- pipeline_uuid=msg.pipeline_uuid,
- routed_by_rule=msg.routed_by_rule,
+ except WorkspaceError as exc:
+ self.ap.logger.info(
+ f'Dropped an aggregated message because its Workspace execution binding is stale: {exc}'
)
+
+ async def _flush_buffer(
+ self,
+ aggregation_key: AggregationKey,
+ execution_context: ExecutionContext,
+ ) -> None:
+ """Flush one buffer only when the captured scope still matches."""
+
+ async with self.lock:
+ buffer = self.buffers.get(aggregation_key)
+ if buffer is None:
+ return
+ if buffer.execution_context != execution_context:
+ raise ExecutionContextMismatchError('Timer ExecutionContext does not match the aggregation buffer')
+ self.buffers.pop(aggregation_key)
+
+ if not buffer.messages:
return
- # Merge multiple messages
- merged_msg = self._merge_messages(buffer.messages)
+ message = buffer.messages[0] if len(buffer.messages) == 1 else self._merge_messages(buffer.messages)
+ binding = await self.ap.workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ if binding.instance_uuid != execution_context.instance_uuid:
+ raise WorkspaceInvariantError('Aggregation buffer instance does not match the active Workspace binding')
await self.ap.query_pool.add_query(
- bot_uuid=merged_msg.bot_uuid,
- launcher_type=merged_msg.launcher_type,
- launcher_id=merged_msg.launcher_id,
- sender_id=merged_msg.sender_id,
- message_event=merged_msg.message_event,
- message_chain=merged_msg.message_chain,
- adapter=merged_msg.adapter,
- pipeline_uuid=merged_msg.pipeline_uuid,
- routed_by_rule=merged_msg.routed_by_rule,
+ bot_uuid=message.bot_uuid,
+ launcher_type=message.launcher_type,
+ launcher_id=message.launcher_id,
+ sender_id=message.sender_id,
+ message_event=message.message_event,
+ message_chain=message.message_chain,
+ adapter=message.adapter,
+ pipeline_uuid=message.pipeline_uuid,
+ routed_by_rule=message.routed_by_rule,
+ execution_context=message.execution_context,
)
def _merge_messages(self, messages: list[PendingMessage]) -> PendingMessage:
- """Merge multiple messages into one
+ """Merge message chains after proving all messages share one scope."""
- The merged message uses the first message as base and combines
- all message chains with newline separators.
- The original message_event is kept unmodified to preserve
- message metadata (message_id, etc.) for reply/quote.
- """
+ if not messages:
+ raise ValueError('At least one pending message is required')
if len(messages) == 1:
return messages[0]
base_msg = messages[0]
+ base_key = self._get_aggregation_key(
+ base_msg.execution_context,
+ base_msg.bot_uuid,
+ base_msg.launcher_type,
+ base_msg.launcher_id,
+ base_msg.pipeline_uuid,
+ )
+ for message in messages[1:]:
+ message_key = self._get_aggregation_key(
+ message.execution_context,
+ message.bot_uuid,
+ message.launcher_type,
+ message.launcher_id,
+ message.pipeline_uuid,
+ )
+ if message_key != base_key or message.execution_context != base_msg.execution_context:
+ raise ExecutionContextMismatchError('Cannot merge pending messages from different execution scopes')
- # Build merged message chain
merged_chain = platform_message.MessageChain([])
-
- for i, msg in enumerate(messages):
- if i > 0:
- # Add newline separator between messages
+ for index, message in enumerate(messages):
+ if index > 0:
merged_chain.append(platform_message.Plain(text='\n'))
-
- # Copy all components from this message
- for component in msg.message_chain:
+ for component in message.message_chain:
merged_chain.append(component)
- # Keep message_event unmodified (preserves original message_id and
- # metadata for reply/quote), only pass merged chain separately
return PendingMessage(
+ execution_context=base_msg.execution_context,
bot_uuid=base_msg.bot_uuid,
launcher_type=base_msg.launcher_type,
launcher_id=base_msg.launcher_id,
@@ -275,22 +308,23 @@ class MessageAggregator:
message_chain=merged_chain,
adapter=base_msg.adapter,
pipeline_uuid=base_msg.pipeline_uuid,
- routed_by_rule=any(msg.routed_by_rule for msg in messages),
+ routed_by_rule=any(message.routed_by_rule for message in messages),
)
async def flush_all(self) -> None:
- """Flush all pending buffers immediately
+ """Flush all pending buffers without dropping their captured scopes."""
- This is useful during shutdown to ensure no messages are lost.
- """
- # Snapshot session IDs and cancel all timers under lock
async with self.lock:
- session_ids = list(self.buffers.keys())
- for sid in session_ids:
- buffer = self.buffers.get(sid)
- if buffer and buffer.timer_task and not buffer.timer_task.done():
+ pending_buffers = [(key, buffer.execution_context) for key, buffer in self.buffers.items()]
+ for buffer in self.buffers.values():
+ if buffer.timer_task and not buffer.timer_task.done():
buffer.timer_task.cancel()
- # Flush each buffer outside the lock
- for session_id in session_ids:
- await self._flush_buffer(session_id)
+ for aggregation_key, execution_context in pending_buffers:
+ try:
+ await self._flush_buffer(aggregation_key, execution_context)
+ except WorkspaceError as exc:
+ self.ap.logger.info(
+ 'Dropped an aggregated message during shutdown because its '
+ f'Workspace execution binding is stale: {exc}'
+ )
diff --git a/src/langbot/pkg/pipeline/controller.py b/src/langbot/pkg/pipeline/controller.py
index 09d18a582..a275e5ff1 100644
--- a/src/langbot/pkg/pipeline/controller.py
+++ b/src/langbot/pkg/pipeline/controller.py
@@ -5,8 +5,10 @@ import traceback
from ..core import app
from ..core import entities as core_entities
+from ..workspace.errors import WorkspaceError, WorkspaceInvariantError
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
+from .pool import get_query_execution_context
class Controller:
@@ -21,6 +23,52 @@ class Controller:
self.ap = ap
self.semaphore = asyncio.Semaphore(self.ap.instance_config.data['concurrency']['pipeline'])
+ async def _assert_query_execution_active(
+ self,
+ query: pipeline_query.Query,
+ ):
+ """Revalidate a queued query immediately before runtime work starts."""
+
+ execution_context = get_query_execution_context(query)
+ binding = await self.ap.workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ if binding.instance_uuid != execution_context.instance_uuid:
+ raise WorkspaceInvariantError('Queued query instance does not match the active Workspace binding')
+ return execution_context
+
+ async def _process_query(self, selected_query: pipeline_query.Query) -> None:
+ """Run one selected query and always release its scheduling slot."""
+
+ try:
+ async with self.semaphore:
+ execution_context = await self._assert_query_execution_active(selected_query)
+ pipeline_uuid = selected_query.pipeline_uuid
+
+ if pipeline_uuid:
+ pipeline = await self.ap.pipeline_mgr.get_pipeline_by_uuid(
+ execution_context,
+ pipeline_uuid,
+ )
+ if pipeline:
+ await pipeline.run(selected_query)
+ else:
+ self.ap.logger.warning(
+ f'Pipeline {pipeline_uuid} not found for query {selected_query.query_id}, query dropped'
+ )
+ else:
+ self.ap.logger.warning(f'No pipeline_uuid for query {selected_query.query_id}, query dropped')
+ except WorkspaceError as exc:
+ self.ap.logger.info(
+ f'Dropped query {selected_query.query_id} because its Workspace execution binding is stale: {exc}'
+ )
+ finally:
+ await self.ap.query_pool.remove_query(selected_query)
+ async with self.ap.query_pool:
+ (await self.ap.sess_mgr.get_session(selected_query))._semaphore.release()
+ self.ap.query_pool.condition.notify_all()
+
async def consumer(self):
"""事件处理循环"""
try:
@@ -51,40 +99,18 @@ class Controller:
continue
if selected_query:
-
- async def _process_query(selected_query: pipeline_query.Query):
- async with self.semaphore: # 总并发上限
- # find pipeline
- # Here firstly find the bot, then find the pipeline, in case the bot adapter's config is not the latest one.
- # Like aiocqhttp, once a client is connected, even the adapter was updated and restarted, the existing client connection will not be affected.
- pipeline_uuid = selected_query.pipeline_uuid
-
- if pipeline_uuid:
- pipeline = await self.ap.pipeline_mgr.get_pipeline_by_uuid(pipeline_uuid)
- if pipeline:
- await pipeline.run(selected_query)
- else:
- self.ap.logger.warning(
- f'Pipeline {pipeline_uuid} not found for query {selected_query.query_id}, query dropped'
- )
- else:
- self.ap.logger.warning(
- f'No pipeline_uuid for query {selected_query.query_id}, query dropped'
- )
-
- async with self.ap.query_pool:
- (await self.ap.sess_mgr.get_session(selected_query))._semaphore.release()
- # 通知其他协程,有新的请求可以处理了
- self.ap.query_pool.condition.notify_all()
-
+ execution_context = get_query_execution_context(selected_query)
self.ap.task_mgr.create_task(
- _process_query(selected_query),
+ self._process_query(selected_query),
kind='query',
name=f'query-{selected_query.query_id}',
scopes=[
core_entities.LifecycleControlScope.APPLICATION,
core_entities.LifecycleControlScope.PLATFORM,
],
+ instance_uuid=execution_context.instance_uuid,
+ workspace_uuid=execution_context.workspace_uuid,
+ placement_generation=execution_context.placement_generation,
)
except Exception as e:
diff --git a/src/langbot/pkg/pipeline/monitoring_helper.py b/src/langbot/pkg/pipeline/monitoring_helper.py
index a3a9654bc..0728b0c83 100644
--- a/src/langbot/pkg/pipeline/monitoring_helper.py
+++ b/src/langbot/pkg/pipeline/monitoring_helper.py
@@ -15,6 +15,8 @@ if typing.TYPE_CHECKING:
from ..core import app
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
+from .pool import get_query_execution_context
+
class MonitoringHelper:
"""Helper class for monitoring operations"""
@@ -54,6 +56,7 @@ class MonitoringHelper:
# Here we just record None, the full variables will be set when query completes
message_id = await ap.monitoring_service.record_message(
+ get_query_execution_context(query),
bot_id=bot_id,
bot_name=bot_name,
pipeline_id=pipeline_id,
@@ -74,6 +77,7 @@ class MonitoringHelper:
# Update session activity or create new session if it doesn't exist
# Always pass pipeline info to handle pipeline switches
session_updated = await ap.monitoring_service.update_session_activity(
+ get_query_execution_context(query),
session_id,
pipeline_id=pipeline_id,
pipeline_name=pipeline_name,
@@ -81,6 +85,7 @@ class MonitoringHelper:
if not session_updated:
# Session doesn't exist, create it
await ap.monitoring_service.record_session_start(
+ get_query_execution_context(query),
session_id=session_id,
bot_id=bot_id,
bot_name=bot_name,
@@ -118,6 +123,7 @@ class MonitoringHelper:
pass
await ap.monitoring_service.update_message_status(
+ get_query_execution_context(query),
message_id=message_id,
status='success',
variables=query_variables_str,
@@ -170,6 +176,7 @@ class MonitoringHelper:
return # No response to record
await ap.monitoring_service.record_message(
+ get_query_execution_context(query),
bot_id=bot_id,
bot_name=bot_name,
pipeline_id=pipeline_id,
@@ -215,6 +222,7 @@ class MonitoringHelper:
# Record error message
message_id = await ap.monitoring_service.record_message(
+ get_query_execution_context(query),
bot_id=bot_id,
bot_name=bot_name,
pipeline_id=pipeline_id,
@@ -233,6 +241,7 @@ class MonitoringHelper:
# Record error log
await ap.monitoring_service.record_error(
+ get_query_execution_context(query),
bot_id=bot_id,
bot_name=bot_name,
pipeline_id=pipeline_id,
@@ -271,6 +280,7 @@ class MonitoringHelper:
session_id = f'{query.launcher_type.value if hasattr(query.launcher_type, "value") else query.launcher_type}_{query.launcher_id}'
await ap.monitoring_service.record_llm_call(
+ get_query_execution_context(query),
bot_id=bot_id,
bot_name=bot_name,
pipeline_id=pipeline_id,
diff --git a/src/langbot/pkg/pipeline/pipelinemgr.py b/src/langbot/pkg/pipeline/pipelinemgr.py
index ef189beaf..7f094efa1 100644
--- a/src/langbot/pkg/pipeline/pipelinemgr.py
+++ b/src/langbot/pkg/pipeline/pipelinemgr.py
@@ -1,5 +1,6 @@
from __future__ import annotations
+import dataclasses
import typing
import traceback
@@ -13,7 +14,11 @@ import langbot_plugin.api.entities.builtin.platform.message as platform_message
import langbot_plugin.api.entities.builtin.platform.events as platform_events
import langbot_plugin.api.entities.events as events
from ..utils import importutil
+from ..api.http.authz import WorkspaceRequiredError
+from ..api.http.context import ExecutionContext, PrincipalContext, PrincipalType, RequestContext
+from ..workspace.errors import WorkspaceError, WorkspaceInvariantError
from .config_coercion import coerce_pipeline_config
+from .pool import get_query_execution_context
import langbot_plugin.api.entities.builtin.provider.session as provider_session
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
@@ -82,15 +87,39 @@ class RuntimePipeline:
enable_all_mcp_servers: bool
"""是否启用所有MCP服务器"""
+ execution_context: ExecutionContext
+
+ workspace_uuid: str
+
+ placement_generation: int
+
def __init__(
self,
ap: app.Application,
pipeline_entity: persistence_pipeline.LegacyPipeline,
stage_containers: list[StageInstContainer],
+ execution_context: ExecutionContext,
):
+ if not isinstance(execution_context, ExecutionContext):
+ raise WorkspaceRequiredError('RuntimePipeline requires an ExecutionContext')
+ if not execution_context.instance_uuid.strip() or not execution_context.workspace_uuid.strip():
+ raise WorkspaceRequiredError('RuntimePipeline requires an instance and Workspace')
+ if execution_context.placement_generation <= 0:
+ raise WorkspaceRequiredError('RuntimePipeline requires a positive placement generation')
+ if pipeline_entity.workspace_uuid != execution_context.workspace_uuid:
+ raise WorkspaceRequiredError('RuntimePipeline entity Workspace does not match its ExecutionContext')
+ if execution_context.pipeline_uuid not in (None, pipeline_entity.uuid):
+ raise WorkspaceRequiredError('RuntimePipeline UUID does not match its ExecutionContext')
+
self.ap = ap
self.pipeline_entity = pipeline_entity
self.stage_containers = stage_containers
+ self.execution_context = dataclasses.replace(
+ execution_context,
+ pipeline_uuid=pipeline_entity.uuid,
+ )
+ self.workspace_uuid = self.execution_context.workspace_uuid
+ self.placement_generation = self.execution_context.placement_generation
# Extract bound plugins and MCP servers from extensions_preferences
extensions_prefs = pipeline_entity.extensions_preferences or {}
@@ -120,7 +149,37 @@ class RuntimePipeline:
mcp_server_list = extensions_prefs.get('mcp_servers', [])
self.bound_mcp_servers = mcp_server_list if mcp_server_list else []
+ async def _assert_execution_active(
+ self,
+ query: pipeline_query.Query | None = None,
+ ) -> ExecutionContext:
+ """Fail closed when this runtime or query belongs to a stale placement."""
+
+ execution_context = self.execution_context if query is None else get_query_execution_context(query)
+ if (
+ execution_context.instance_uuid != self.execution_context.instance_uuid
+ or execution_context.workspace_uuid != self.workspace_uuid
+ or execution_context.placement_generation != self.placement_generation
+ or execution_context.pipeline_uuid != self.pipeline_entity.uuid
+ ):
+ raise WorkspaceInvariantError('Query execution scope does not match RuntimePipeline')
+ binding = await self.ap.workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ if binding.instance_uuid != execution_context.instance_uuid:
+ raise WorkspaceInvariantError('RuntimePipeline instance does not match the active Workspace binding')
+ return execution_context
+
async def run(self, query: pipeline_query.Query):
+ if (
+ query.instance_uuid != self.execution_context.instance_uuid
+ or query.workspace_uuid != self.workspace_uuid
+ or query.placement_generation != self.placement_generation
+ or query.pipeline_uuid != self.pipeline_entity.uuid
+ ):
+ raise WorkspaceRequiredError('Query execution scope does not match RuntimePipeline')
+ await self._assert_execution_active(query)
query.pipeline_config = self.pipeline_entity.config
# Store bound plugins and MCP servers in query for filtering
query.variables['_pipeline_bound_plugins'] = self.bound_plugins
@@ -134,7 +193,11 @@ class RuntimePipeline:
bot_name = 'WebChat'
if query.bot_uuid:
try:
- bot = await self.ap.bot_service.get_bot(query.bot_uuid, include_secret=False)
+ bot = await self.ap.bot_service.get_bot(
+ query.workspace_uuid,
+ query.bot_uuid,
+ include_secret=False,
+ )
if bot:
bot_name = bot.get('name', 'Unknown')
except Exception:
@@ -150,6 +213,7 @@ class RuntimePipeline:
async def _check_output(self, query: pipeline_query.Query, result: pipeline_entities.StageProcessResult):
"""检查输出"""
+ await self._assert_execution_active(query)
if result.user_notice:
# 处理str类型
@@ -162,7 +226,9 @@ class RuntimePipeline:
query.message_event, platform_events.GroupMessage
):
result.user_notice.insert(0, platform_message.At(target=query.message_event.sender.id))
- if await query.adapter.is_stream_output_supported() and query.resp_messages:
+ stream_output_supported = await query.adapter.is_stream_output_supported()
+ await self._assert_execution_active(query)
+ if stream_output_supported and query.resp_messages:
await query.adapter.reply_message_chunk(
message_source=query.message_event,
bot_message=query.resp_messages[-1],
@@ -186,6 +252,7 @@ class RuntimePipeline:
query.variables['_monitoring_has_error'] = True
# Record error to monitoring system
try:
+ await self._assert_execution_active(query)
bot_name = query.variables.get('_monitoring_bot_name', 'Unknown')
pipeline_name = query.variables.get('_monitoring_pipeline_name', 'Unknown')
message_id = query.variables.get('_monitoring_message_id', '')
@@ -194,6 +261,7 @@ class RuntimePipeline:
# Update message status to error
if message_id:
await self.ap.monitoring_service.update_message_status(
+ get_query_execution_context(query),
message_id=message_id,
status='error',
level='error',
@@ -201,6 +269,7 @@ class RuntimePipeline:
# Record error log
await self.ap.monitoring_service.record_error(
+ get_query_execution_context(query),
bot_id=query.bot_uuid or 'unknown',
bot_name=bot_name,
pipeline_id=self.pipeline_entity.uuid,
@@ -242,6 +311,7 @@ class RuntimePipeline:
i = stage_index
while i < len(self.stage_containers):
+ await self._assert_execution_active(query)
stage_container = self.stage_containers[i]
query.current_stage_name = stage_container.inst_name # 标记到 Query 对象里
@@ -250,6 +320,7 @@ class RuntimePipeline:
if isinstance(result, typing.Coroutine):
result = await result
+ await self._assert_execution_active(query)
if isinstance(result, pipeline_entities.StageProcessResult): # 直接返回结果
self.ap.logger.debug(
@@ -265,7 +336,14 @@ class RuntimePipeline:
elif isinstance(result, typing.AsyncGenerator): # 生成器
self.ap.logger.debug(f'Stage {stage_container.inst_name} processed query {query.query_id} gen')
- async for sub_result in result:
+ iterator = result.__aiter__()
+ while True:
+ await self._assert_execution_active(query)
+ try:
+ sub_result = await anext(iterator)
+ except StopAsyncIteration:
+ break
+ await self._assert_execution_active(query)
self.ap.logger.debug(
f'Stage {stage_container.inst_name} processed query {query.query_id} res {sub_result.result_type}'
)
@@ -283,6 +361,7 @@ class RuntimePipeline:
async def process_query(self, query: pipeline_query.Query):
"""处理请求"""
+ await self._assert_execution_active(query)
# Get monitoring metadata
bot_name = query.variables.get('_monitoring_bot_name', 'Unknown')
pipeline_name = query.variables.get('_monitoring_pipeline_name', 'Unknown')
@@ -310,6 +389,7 @@ class RuntimePipeline:
query.variables['_monitoring_message_id'] = message_id
# Notify adapter so it can map platform-specific IDs to monitoring message ID
if hasattr(query.adapter, 'on_monitoring_message_created'):
+ await self._assert_execution_active(query)
await query.adapter.on_monitoring_message_created(query, message_id)
except Exception as e:
self.ap.logger.error(f'Failed to record query start: {e}')
@@ -334,7 +414,9 @@ class RuntimePipeline:
message_chain=query.message_chain,
)
+ await self._assert_execution_active(query)
event_ctx = await self.ap.plugin_connector.emit_event(event_obj, bound_plugins)
+ await self._assert_execution_active(query)
if event_ctx.is_prevented_default():
self.ap.logger.debug(
@@ -349,6 +431,7 @@ class RuntimePipeline:
# Record query success only if no error occurred during processing
if not query.variables.get('_monitoring_has_error', False):
try:
+ await self._assert_execution_active(query)
await monitoring_helper.MonitoringHelper.record_query_success(
ap=self.ap,
message_id=message_id,
@@ -359,6 +442,7 @@ class RuntimePipeline:
# Record bot response message
try:
+ await self._assert_execution_active(query)
await monitoring_helper.MonitoringHelper.record_query_response(
ap=self.ap,
query=query,
@@ -371,6 +455,8 @@ class RuntimePipeline:
except Exception as e:
self.ap.logger.error(f'Failed to record query response: {e}')
+ except WorkspaceError as e:
+ self.ap.logger.info(f'Dropped query {query.query_id} because its Workspace execution binding is stale: {e}')
except Exception as e:
inst_name = query.current_stage_name if query.current_stage_name else 'unknown'
self.ap.logger.error(f'Error processing query {query.query_id} stage={inst_name} : {e}')
@@ -380,6 +466,7 @@ class RuntimePipeline:
try:
from . import monitoring_helper
+ await self._assert_execution_active(query)
await monitoring_helper.MonitoringHelper.record_query_error(
ap=self.ap,
query=query,
@@ -395,7 +482,7 @@ class RuntimePipeline:
finally:
self.ap.logger.debug(f'Query {query.query_id} processed')
- del self.ap.query_pool.cached_queries[query.query_id]
+ await self.ap.query_pool.remove_query(query)
class PipelineManager:
@@ -425,10 +512,38 @@ class PipelineManager:
# load pipelines
for pipeline in pipelines:
- await self.load_pipeline(pipeline)
+ binding = await self.ap.workspace_service.get_execution_binding(pipeline.workspace_uuid)
+ await self.load_pipeline(
+ ExecutionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ pipeline_uuid=pipeline.uuid,
+ trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
+ ),
+ pipeline,
+ )
+
+ @staticmethod
+ def _normalize_execution_context(
+ context: ExecutionContext | RequestContext,
+ pipeline_uuid: str,
+ ) -> ExecutionContext:
+ if isinstance(context, RequestContext):
+ return ExecutionContext.from_request(context, pipeline_uuid=pipeline_uuid)
+ if not isinstance(context, ExecutionContext):
+ raise WorkspaceRequiredError('Pipeline runtime operations require an ExecutionContext')
+ if not context.instance_uuid.strip() or not context.workspace_uuid.strip():
+ raise WorkspaceRequiredError('Pipeline runtime operations require an instance and Workspace')
+ if context.placement_generation <= 0:
+ raise WorkspaceRequiredError('Pipeline runtime operations require a positive placement generation')
+ if context.pipeline_uuid not in (None, pipeline_uuid):
+ raise WorkspaceRequiredError('Pipeline UUID does not match its ExecutionContext')
+ return dataclasses.replace(context, pipeline_uuid=pipeline_uuid)
async def load_pipeline(
self,
+ context: ExecutionContext | RequestContext,
pipeline_entity: persistence_pipeline.LegacyPipeline
| sqlalchemy.Row[persistence_pipeline.LegacyPipeline]
| dict,
@@ -438,6 +553,14 @@ class PipelineManager:
elif isinstance(pipeline_entity, dict):
pipeline_entity = persistence_pipeline.LegacyPipeline(**pipeline_entity)
+ execution_context = self._normalize_execution_context(context, pipeline_entity.uuid)
+ if pipeline_entity.workspace_uuid != execution_context.workspace_uuid:
+ raise WorkspaceRequiredError('Pipeline entity Workspace does not match its runtime context')
+ await self.ap.workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+
coerce_pipeline_config(
pipeline_entity.config,
getattr(self.ap, 'pipeline_config_meta_trigger', {'name': 'trigger', 'stages': []}),
@@ -454,17 +577,36 @@ class PipelineManager:
for stage_container in stage_containers:
await stage_container.inst.initialize(pipeline_entity.config)
- runtime_pipeline = RuntimePipeline(self.ap, pipeline_entity, stage_containers)
+ runtime_pipeline = RuntimePipeline(
+ self.ap,
+ pipeline_entity,
+ stage_containers,
+ execution_context,
+ )
self.pipelines.append(runtime_pipeline)
- async def get_pipeline_by_uuid(self, uuid: str) -> RuntimePipeline | None:
+ async def get_pipeline_by_uuid(
+ self,
+ context: ExecutionContext | RequestContext,
+ uuid: str,
+ ) -> RuntimePipeline | None:
+ execution_context = self._normalize_execution_context(context, uuid)
for pipeline in self.pipelines:
- if pipeline.pipeline_entity.uuid == uuid:
+ if (
+ pipeline.workspace_uuid == execution_context.workspace_uuid
+ and pipeline.placement_generation == execution_context.placement_generation
+ and pipeline.pipeline_entity.uuid == uuid
+ ):
return pipeline
return None
- async def remove_pipeline(self, uuid: str):
+ async def remove_pipeline(
+ self,
+ context: ExecutionContext | RequestContext,
+ uuid: str,
+ ) -> None:
+ execution_context = self._normalize_execution_context(context, uuid)
for pipeline in self.pipelines:
- if pipeline.pipeline_entity.uuid == uuid:
+ if pipeline.workspace_uuid == execution_context.workspace_uuid and pipeline.pipeline_entity.uuid == uuid:
self.pipelines.remove(pipeline)
return
diff --git a/src/langbot/pkg/pipeline/pool.py b/src/langbot/pkg/pipeline/pool.py
index 55ce7fe12..fecdf84d3 100644
--- a/src/langbot/pkg/pipeline/pool.py
+++ b/src/langbot/pkg/pipeline/pool.py
@@ -1,57 +1,196 @@
from __future__ import annotations
import asyncio
+import dataclasses
+import inspect
import typing
+import uuid
-import langbot_plugin.api.entities.builtin.platform.message as platform_message
-import langbot_plugin.api.entities.builtin.platform.events as platform_events
-import langbot_plugin.api.entities.builtin.provider.session as provider_session
-import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
import langbot_plugin.api.definition.abstract.platform.adapter as abstract_platform_adapter
+import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
+import langbot_plugin.api.entities.builtin.platform.events as platform_events
+import langbot_plugin.api.entities.builtin.platform.message as platform_message
+import langbot_plugin.api.entities.builtin.provider.session as provider_session
+
+from ..api.http.context import ExecutionContext
+
+QueryCacheKey = tuple[str, str]
+LegacyQueryKey = tuple[str, int]
+QueryCounterKey = tuple[str, str, int]
+SingletonContextResolver = typing.Callable[
+ [],
+ ExecutionContext | typing.Awaitable[ExecutionContext],
+]
+
+
+class ExecutionContextRequiredError(ValueError):
+ """Raised when runtime work is created without a trusted Workspace scope."""
+
+
+class ExecutionContextMismatchError(ValueError):
+ """Raised when entity fields conflict with their trusted execution scope."""
+
+
+class QueryNotFoundError(LookupError):
+ """Raised when a query does not exist inside the requested Workspace."""
+
+
+def _validate_execution_context(execution_context: ExecutionContext) -> None:
+ if not isinstance(execution_context, ExecutionContext):
+ raise ExecutionContextRequiredError('A trusted ExecutionContext is required')
+ if not isinstance(execution_context.instance_uuid, str) or not execution_context.instance_uuid.strip():
+ raise ExecutionContextRequiredError('ExecutionContext.instance_uuid is required')
+ if not isinstance(execution_context.workspace_uuid, str) or not execution_context.workspace_uuid.strip():
+ raise ExecutionContextRequiredError('ExecutionContext.workspace_uuid is required')
+ if (
+ isinstance(execution_context.placement_generation, bool)
+ or not isinstance(execution_context.placement_generation, int)
+ or execution_context.placement_generation <= 0
+ ):
+ raise ExecutionContextRequiredError('ExecutionContext.placement_generation must be a positive integer')
+ for field_name in ('bot_uuid', 'pipeline_uuid', 'query_uuid'):
+ value = getattr(execution_context, field_name)
+ if value is not None and (not isinstance(value, str) or not value.strip()):
+ raise ExecutionContextRequiredError(f'ExecutionContext.{field_name} must be a non-empty string when set')
+
+
+def bind_execution_context(
+ execution_context: ExecutionContext,
+ *,
+ bot_uuid: str | None = None,
+ pipeline_uuid: str | None = None,
+ query_uuid: str | None = None,
+) -> ExecutionContext:
+ """Bind runtime entity identifiers without allowing scope substitution."""
+
+ _validate_execution_context(execution_context)
+
+ requested_fields = {
+ 'bot_uuid': bot_uuid,
+ 'pipeline_uuid': pipeline_uuid,
+ 'query_uuid': query_uuid,
+ }
+ updates: dict[str, str] = {}
+ for field_name, requested_value in requested_fields.items():
+ if requested_value is None:
+ continue
+ if not isinstance(requested_value, str) or not requested_value.strip():
+ raise ExecutionContextRequiredError(f'{field_name} must be a non-empty string')
+ current_value = getattr(execution_context, field_name)
+ if current_value is not None and current_value != requested_value:
+ raise ExecutionContextMismatchError(f'ExecutionContext.{field_name} does not match the runtime entity')
+ if current_value is None:
+ updates[field_name] = requested_value
+
+ if not updates:
+ return execution_context
+ return dataclasses.replace(execution_context, **updates)
+
+
+def get_query_execution_context(query: pipeline_query.Query) -> ExecutionContext:
+ """Return and validate the trusted context attached to a Query."""
+
+ attached_context = getattr(query, '_execution_context', None)
+ bot_uuid = getattr(query, 'bot_uuid', None)
+ pipeline_uuid = getattr(query, 'pipeline_uuid', None)
+ query_uuid = getattr(query, 'query_uuid', None)
+
+ if isinstance(attached_context, ExecutionContext):
+ return bind_execution_context(
+ attached_context,
+ bot_uuid=bot_uuid,
+ pipeline_uuid=pipeline_uuid,
+ query_uuid=query_uuid,
+ )
+
+ raise ExecutionContextRequiredError('Query is missing its trusted ExecutionContext')
class QueryPool:
- """请求池,请求获得调度进入pipeline之前,保存在这里"""
-
- query_id_counter: int = 0
+ """Workspace-scoped queue of requests waiting for pipeline scheduling."""
+ query_id_counter: int
pool_lock: asyncio.Lock
-
queries: list[pipeline_query.Query]
-
- cached_queries: dict[int, pipeline_query.Query]
- """Cached queries, used for plugin backward api call, will be removed after the query completely processed"""
-
+ cached_queries: dict[QueryCacheKey, pipeline_query.Query]
+ legacy_query_index: dict[LegacyQueryKey, str]
+ query_count_by_scope: dict[QueryCounterKey, int]
condition: asyncio.Condition
- def __init__(self):
+ def __init__(
+ self,
+ singleton_context_resolver: SingletonContextResolver | None = None,
+ ):
self.query_id_counter = 0
self.pool_lock = asyncio.Lock()
self.queries = []
self.cached_queries = {}
+ self.legacy_query_index = {}
+ self.query_count_by_scope = {}
self.condition = asyncio.Condition(self.pool_lock)
+ self._singleton_context_resolver = singleton_context_resolver
+
+ async def resolve_execution_context(
+ self,
+ execution_context: ExecutionContext | None,
+ *,
+ bot_uuid: str,
+ pipeline_uuid: str | None,
+ query_uuid: str | None = None,
+ ) -> ExecutionContext:
+ """Resolve an explicit scope or the opt-in OSS singleton scope."""
+
+ if execution_context is None:
+ if self._singleton_context_resolver is None:
+ raise ExecutionContextRequiredError('ExecutionContext is required; no singleton resolver is configured')
+ resolved_context = self._singleton_context_resolver()
+ if inspect.isawaitable(resolved_context):
+ resolved_context = await resolved_context
+ execution_context = resolved_context
+
+ return bind_execution_context(
+ execution_context,
+ bot_uuid=bot_uuid,
+ pipeline_uuid=pipeline_uuid,
+ query_uuid=query_uuid,
+ )
async def add_query(
self,
bot_uuid: str,
launcher_type: provider_session.LauncherTypes,
- launcher_id: typing.Union[int, str],
- sender_id: typing.Union[int, str],
+ launcher_id: int | str,
+ sender_id: int | str,
message_event: platform_events.MessageEvent,
message_chain: platform_message.MessageChain,
adapter: abstract_platform_adapter.AbstractMessagePlatformAdapter,
- pipeline_uuid: typing.Optional[str] = None,
+ pipeline_uuid: str | None = None,
routed_by_rule: bool = False,
- variables: typing.Optional[dict[str, typing.Any]] = None,
+ variables: dict[str, typing.Any] | None = None,
+ execution_context: ExecutionContext | None = None,
) -> pipeline_query.Query:
+ """Create a query and cache it under an opaque, Workspace-scoped key."""
+
+ query_uuid = str(uuid.uuid4())
+ execution_context = await self.resolve_execution_context(
+ execution_context,
+ bot_uuid=bot_uuid,
+ pipeline_uuid=pipeline_uuid,
+ query_uuid=query_uuid,
+ )
+
async with self.condition:
query_id = self.query_id_counter
initial_variables: dict[str, typing.Any] = {'_routed_by_rule': routed_by_rule}
if variables:
initial_variables.update(variables)
query = pipeline_query.Query(
+ instance_uuid=execution_context.instance_uuid,
+ workspace_uuid=execution_context.workspace_uuid,
+ placement_generation=execution_context.placement_generation,
bot_uuid=bot_uuid,
query_id=query_id,
+ query_uuid=query_uuid,
launcher_type=launcher_type,
launcher_id=launcher_id,
sender_id=sender_id,
@@ -63,12 +202,104 @@ class QueryPool:
adapter=adapter,
pipeline_uuid=pipeline_uuid,
)
+
+ # langbot-plugin 0.4.13 ignores these forward-compatible fields.
+ # Attach them explicitly until the Workspace-aware SDK is released.
+ object.__setattr__(query, 'instance_uuid', execution_context.instance_uuid)
+ object.__setattr__(query, 'workspace_uuid', execution_context.workspace_uuid)
+ object.__setattr__(
+ query,
+ 'placement_generation',
+ execution_context.placement_generation,
+ )
+ object.__setattr__(query, 'query_uuid', query_uuid)
+ object.__setattr__(query, '_execution_context', execution_context)
+
self.queries.append(query)
- self.cached_queries[query_id] = query
+ self.cached_queries[(execution_context.workspace_uuid, query_uuid)] = query
+ self.legacy_query_index[(execution_context.workspace_uuid, query_id)] = query_uuid
self.query_id_counter += 1
+ counter_key = (
+ execution_context.instance_uuid,
+ execution_context.workspace_uuid,
+ execution_context.placement_generation,
+ )
+ self.query_count_by_scope[counter_key] = self.query_count_by_scope.get(counter_key, 0) + 1
self.condition.notify_all()
return query
+ def get_query_count(self, execution_context: ExecutionContext) -> int:
+ """Return the lifetime query count for one active placement scope."""
+
+ _validate_execution_context(execution_context)
+ return self.query_count_by_scope.get(
+ (
+ execution_context.instance_uuid,
+ execution_context.workspace_uuid,
+ execution_context.placement_generation,
+ ),
+ 0,
+ )
+
+ async def get_query(
+ self,
+ workspace_uuid: str,
+ query_uuid: str,
+ ) -> pipeline_query.Query | None:
+ """Return a query only from the explicitly selected Workspace."""
+
+ async with self.pool_lock:
+ return self.cached_queries.get((workspace_uuid, query_uuid))
+
+ async def require_query(
+ self,
+ workspace_uuid: str,
+ query_uuid: str,
+ ) -> pipeline_query.Query:
+ """Return a scoped query or raise without checking other Workspaces."""
+
+ query = await self.get_query(workspace_uuid, query_uuid)
+ if query is None:
+ raise QueryNotFoundError(f'Query {query_uuid!r} was not found in Workspace {workspace_uuid!r}')
+ return query
+
+ async def get_query_by_legacy_id(
+ self,
+ workspace_uuid: str,
+ query_id: int,
+ ) -> pipeline_query.Query | None:
+ """Resolve a legacy integer ID within one explicit Workspace."""
+
+ async with self.pool_lock:
+ query_uuid = self.legacy_query_index.get((workspace_uuid, query_id))
+ if query_uuid is None:
+ return None
+ return self.cached_queries.get((workspace_uuid, query_uuid))
+
+ async def remove_query(self, query: pipeline_query.Query) -> bool:
+ """Remove a query and both of its Workspace-scoped indexes."""
+
+ execution_context = get_query_execution_context(query)
+ query_uuid = execution_context.query_uuid
+ if query_uuid is None:
+ raise ExecutionContextRequiredError('Query.query_uuid is required for removal')
+
+ async with self.pool_lock:
+ cache_key = (execution_context.workspace_uuid, query_uuid)
+ cached_query = self.cached_queries.get(cache_key)
+ if cached_query is not query:
+ return False
+ del self.cached_queries[cache_key]
+ self.legacy_query_index.pop(
+ (execution_context.workspace_uuid, query.query_id),
+ None,
+ )
+ for index, queued_query in enumerate(self.queries):
+ if queued_query is query:
+ self.queries.pop(index)
+ break
+ return True
+
async def __aenter__(self):
await self.pool_lock.acquire()
return self
diff --git a/src/langbot/pkg/pipeline/preproc/preproc.py b/src/langbot/pkg/pipeline/preproc/preproc.py
index b14d0a827..5b2201b2e 100644
--- a/src/langbot/pkg/pipeline/preproc/preproc.py
+++ b/src/langbot/pkg/pipeline/preproc/preproc.py
@@ -8,6 +8,7 @@ import langbot_plugin.api.entities.events as events
import langbot_plugin.api.entities.builtin.platform.message as platform_message
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
import langbot_plugin.api.entities.builtin.platform.events as platform_events
+from ...pipeline.pool import get_query_execution_context
@stage.stage_class('PreProcessor')
@@ -70,7 +71,10 @@ class PreProcessor(stage.PipelineStage):
if primary_uuid:
try:
- llm_model = await self.ap.model_mgr.get_model_by_uuid(primary_uuid)
+ llm_model = await self.ap.model_mgr.get_model_by_uuid(
+ get_query_execution_context(query),
+ primary_uuid,
+ )
except ValueError:
self.ap.logger.warning(f'LLM model {primary_uuid} not found or not configured')
@@ -79,7 +83,10 @@ class PreProcessor(stage.PipelineStage):
valid_fallbacks = []
for fb_uuid in fallback_uuids:
try:
- await self.ap.model_mgr.get_model_by_uuid(fb_uuid)
+ await self.ap.model_mgr.get_model_by_uuid(
+ get_query_execution_context(query),
+ fb_uuid,
+ )
valid_fallbacks.append(fb_uuid)
except ValueError:
self.ap.logger.warning(f'Fallback model {fb_uuid} not found, skipping')
@@ -131,6 +138,7 @@ class PreProcessor(stage.PipelineStage):
bound_mcp_servers = query.variables.get('_pipeline_bound_mcp_servers', None)
include_mcp_resource_tools = query.variables.get('_pipeline_mcp_resource_agent_read_enabled', True)
all_tools = await self.ap.tool_mgr.get_all_tools(
+ get_query_execution_context(query),
bound_plugins,
bound_mcp_servers,
include_skill_authoring=include_skill_authoring,
@@ -149,6 +157,7 @@ class PreProcessor(stage.PipelineStage):
bound_mcp_servers = query.variables.get('_pipeline_bound_mcp_servers', None)
include_mcp_resource_tools = query.variables.get('_pipeline_mcp_resource_agent_read_enabled', True)
all_tools = await self.ap.tool_mgr.get_all_tools(
+ get_query_execution_context(query),
bound_plugins,
bound_mcp_servers,
include_skill_authoring=include_skill_authoring,
@@ -279,7 +288,13 @@ class PreProcessor(stage.PipelineStage):
# relied on this injection; without it the LLM never discovers
# the skills are there and just calls native tools instead.
if selected_runner == 'local-agent' and self.ap.skill_mgr:
- pipeline_data = await self.ap.pipeline_service.get_pipeline(query.pipeline_uuid)
+ skill_execution_context = get_query_execution_context(query)
+ await self.ap.skill_mgr.ensure_loaded(skill_execution_context)
+ pipeline_data = await self.ap.pipeline_service.get_pipeline(
+ query.workspace_uuid,
+ query.pipeline_uuid,
+ include_secret=True,
+ )
extensions_prefs = (pipeline_data or {}).get('extensions_preferences', {})
enable_all_skills = extensions_prefs.get('enable_all_skills', True)
@@ -291,6 +306,7 @@ class PreProcessor(stage.PipelineStage):
query.variables['_pipeline_bound_skills'] = bound_skills
skill_addition = self.ap.skill_mgr.build_skill_aware_prompt_addition(
+ skill_execution_context,
bound_skills=bound_skills,
)
if skill_addition:
@@ -319,13 +335,13 @@ class PreProcessor(stage.PipelineStage):
f'Skill index injected into system prompt: '
f'pipeline={query.pipeline_uuid} '
f'bound_skills={bound_skills or "all"} '
- f'loaded_skills={len(self.ap.skill_mgr.skills)}'
+ f'loaded_skills={len(self.ap.skill_mgr.get_skills(skill_execution_context))}'
)
else:
self.ap.logger.debug(
f'No skills available for prompt injection: '
f'pipeline={query.pipeline_uuid} '
- f'loaded_skills={len(self.ap.skill_mgr.skills)} '
+ f'loaded_skills={len(self.ap.skill_mgr.get_skills(skill_execution_context))} '
f'bound_skills={bound_skills}'
)
diff --git a/src/langbot/pkg/pipeline/process/handlers/chat.py b/src/langbot/pkg/pipeline/process/handlers/chat.py
index 4e5c14ea4..a15226a9a 100644
--- a/src/langbot/pkg/pipeline/process/handlers/chat.py
+++ b/src/langbot/pkg/pipeline/process/handlers/chat.py
@@ -19,6 +19,7 @@ from ....provider import runners
import langbot_plugin.api.entities.builtin.provider.session as provider_session
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
import langbot_plugin.api.entities.builtin.provider.message as provider_message
+from ...pool import get_query_execution_context
importutil.import_modules_in_pkg(runners)
@@ -198,7 +199,10 @@ class ChatMessageHandler(handler.MessageHandler):
model_name = None
try:
if runner_name == 'local-agent' and getattr(query, 'use_llm_model_uuid', None):
- m = await self.ap.model_mgr.get_model_by_uuid(query.use_llm_model_uuid)
+ m = await self.ap.model_mgr.get_model_by_uuid(
+ get_query_execution_context(query),
+ query.use_llm_model_uuid,
+ )
if m and getattr(m, 'model_entity', None):
model_name = getattr(m.model_entity, 'name', None)
except Exception:
diff --git a/src/langbot/pkg/platform/botmgr.py b/src/langbot/pkg/platform/botmgr.py
index 6e995206f..5681a9fd0 100644
--- a/src/langbot/pkg/platform/botmgr.py
+++ b/src/langbot/pkg/platform/botmgr.py
@@ -1,9 +1,11 @@
from __future__ import annotations
import asyncio
+import dataclasses
import json
import re
import traceback
+import uuid
import sqlalchemy
from ..core import app, entities as core_entities, taskmgr
@@ -14,6 +16,8 @@ from ..entity.persistence import bot as persistence_bot
from ..entity.persistence import pipeline as persistence_pipeline
from ..entity.errors import platform as platform_errors
+from ..api.http.context import ExecutionContext, PrincipalContext, PrincipalType, RequestContext
+from ..api.http.authz import WorkspaceRequiredError
from .logger import EventLogger
@@ -40,20 +44,50 @@ class RuntimeBot:
logger: EventLogger
+ execution_context: ExecutionContext
+
+ workspace_uuid: str
+
+ placement_generation: int
+
def __init__(
self,
ap: app.Application,
bot_entity: persistence_bot.Bot,
adapter: abstract_platform_adapter.AbstractMessagePlatformAdapter,
logger: EventLogger,
+ execution_context: ExecutionContext,
):
+ if not isinstance(execution_context, ExecutionContext):
+ raise WorkspaceRequiredError('RuntimeBot requires an ExecutionContext')
+ if not execution_context.instance_uuid.strip() or not execution_context.workspace_uuid.strip():
+ raise WorkspaceRequiredError('RuntimeBot requires an instance and Workspace')
+ if execution_context.placement_generation <= 0:
+ raise WorkspaceRequiredError('RuntimeBot requires a positive placement generation')
+ entity_workspace_uuid = getattr(bot_entity, 'workspace_uuid', None)
+ if entity_workspace_uuid != execution_context.workspace_uuid:
+ raise WorkspaceRequiredError('RuntimeBot entity Workspace does not match its ExecutionContext')
+ if execution_context.bot_uuid not in (None, bot_entity.uuid):
+ raise WorkspaceRequiredError('RuntimeBot bot UUID does not match its ExecutionContext')
+
self.ap = ap
self.bot_entity = bot_entity
+ self.execution_context = dataclasses.replace(execution_context, bot_uuid=bot_entity.uuid)
+ self.workspace_uuid = self.execution_context.workspace_uuid
+ self.placement_generation = self.execution_context.placement_generation
self.enable = bot_entity.enable
self.adapter = adapter
self.task_context = taskmgr.TaskContext()
self.logger = logger
+ async def assert_execution_active(self) -> None:
+ """Fail closed when this long-lived adapter belongs to a stale placement."""
+
+ await self.ap.workspace_service.get_execution_binding(
+ self.workspace_uuid,
+ expected_generation=self.placement_generation,
+ )
+
@staticmethod
def _match_operator(actual: str, operator: str, expected: str) -> bool:
"""Evaluate a single operator condition."""
@@ -135,6 +169,28 @@ class RuntimeBot:
return self.bot_entity.use_pipeline_uuid, False
+ def resolve_event_pipeline_uuid(
+ self,
+ adapter: abstract_platform_adapter.AbstractMessagePlatformAdapter,
+ launcher_type: str,
+ launcher_id: str,
+ message_text: str,
+ message_element_types: list[str] | None = None,
+ ) -> tuple[str | None, bool]:
+ """Resolve a pipeline, honoring a trusted per-task adapter override."""
+
+ get_override = getattr(adapter, 'get_pipeline_uuid_override', None)
+ if callable(get_override):
+ override = get_override()
+ if override:
+ return str(override), False
+ return self.resolve_pipeline_uuid(
+ launcher_type,
+ launcher_id,
+ message_text,
+ message_element_types,
+ )
+
async def _record_discarded_message(
self,
launcher_type: provider_session.LauncherTypes,
@@ -162,6 +218,7 @@ class RuntimeBot:
platform = launcher_type.value if hasattr(launcher_type, 'value') else str(launcher_type)
await self.ap.monitoring_service.record_message(
+ self.execution_context,
bot_id=self.bot_entity.uuid,
bot_name=self.bot_entity.name or self.bot_entity.uuid,
pipeline_id=self.PIPELINE_DISCARD,
@@ -179,11 +236,13 @@ class RuntimeBot:
# Don't overwrite pipeline info — a session may have messages from
# multiple pipelines; discarding shouldn't change the displayed pipeline.
session_updated = await self.ap.monitoring_service.update_session_activity(
+ self.execution_context,
session_id,
)
if not session_updated:
# No session yet (first message for this launcher was discarded).
await self.ap.monitoring_service.record_session_start(
+ self.execution_context,
session_id=session_id,
bot_id=self.bot_entity.uuid,
bot_name=self.bot_entity.name or self.bot_entity.uuid,
@@ -201,6 +260,7 @@ class RuntimeBot:
event: platform_events.FriendMessage,
adapter: abstract_platform_adapter.AbstractMessagePlatformAdapter,
):
+ await self.assert_execution_active()
image_components = [
component for component in event.message_chain if isinstance(component, platform_message.Image)
]
@@ -215,7 +275,10 @@ class RuntimeBot:
skip_pipeline = False
if hasattr(self.ap, 'webhook_pusher') and self.ap.webhook_pusher:
skip_pipeline = await self.ap.webhook_pusher.push_person_message(
- event, self.bot_entity.uuid, adapter.__class__.__name__
+ self.execution_context,
+ event,
+ self.bot_entity.uuid,
+ adapter.__class__.__name__,
)
# Only add to query pool if no webhook requested to skip pipeline
@@ -229,8 +292,12 @@ class RuntimeBot:
message_text = str(event.message_chain)
element_types = [comp.type for comp in event.message_chain]
- pipeline_uuid, routed_by_rule = self.resolve_pipeline_uuid(
- 'person', launcher_id, message_text, element_types
+ pipeline_uuid, routed_by_rule = self.resolve_event_pipeline_uuid(
+ adapter,
+ 'person',
+ launcher_id,
+ message_text,
+ element_types,
)
if pipeline_uuid == self.PIPELINE_DISCARD:
@@ -254,6 +321,7 @@ class RuntimeBot:
adapter=adapter,
pipeline_uuid=pipeline_uuid,
routed_by_rule=routed_by_rule,
+ execution_context=self.execution_context,
)
else:
await self.logger.info('Pipeline skipped for person message due to webhook response')
@@ -262,6 +330,7 @@ class RuntimeBot:
event: platform_events.GroupMessage,
adapter: abstract_platform_adapter.AbstractMessagePlatformAdapter,
):
+ await self.assert_execution_active()
image_components = [
component for component in event.message_chain if isinstance(component, platform_message.Image)
]
@@ -276,7 +345,10 @@ class RuntimeBot:
skip_pipeline = False
if hasattr(self.ap, 'webhook_pusher') and self.ap.webhook_pusher:
skip_pipeline = await self.ap.webhook_pusher.push_group_message(
- event, self.bot_entity.uuid, adapter.__class__.__name__
+ self.execution_context,
+ event,
+ self.bot_entity.uuid,
+ adapter.__class__.__name__,
)
# Only add to query pool if no webhook requested to skip pipeline
@@ -290,8 +362,12 @@ class RuntimeBot:
message_text = str(event.message_chain)
element_types = [comp.type for comp in event.message_chain]
- pipeline_uuid, routed_by_rule = self.resolve_pipeline_uuid(
- 'group', launcher_id, message_text, element_types
+ pipeline_uuid, routed_by_rule = self.resolve_event_pipeline_uuid(
+ adapter,
+ 'group',
+ launcher_id,
+ message_text,
+ element_types,
)
if pipeline_uuid == self.PIPELINE_DISCARD:
@@ -315,6 +391,7 @@ class RuntimeBot:
adapter=adapter,
pipeline_uuid=pipeline_uuid,
routed_by_rule=routed_by_rule,
+ execution_context=self.execution_context,
)
else:
await self.logger.info('Pipeline skipped for group message due to webhook response')
@@ -328,13 +405,15 @@ class RuntimeBot:
adapter: abstract_platform_adapter.AbstractMessagePlatformAdapter,
):
try:
+ await self.assert_execution_active()
# Resolve pipeline name
pipeline_name = ''
if self.bot_entity.use_pipeline_uuid:
try:
pipeline_result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_pipeline.LegacyPipeline.name).where(
- persistence_pipeline.LegacyPipeline.uuid == self.bot_entity.use_pipeline_uuid
+ persistence_pipeline.LegacyPipeline.workspace_uuid == self.workspace_uuid,
+ persistence_pipeline.LegacyPipeline.uuid == self.bot_entity.use_pipeline_uuid,
)
)
pipeline_row = pipeline_result.first()
@@ -344,6 +423,7 @@ class RuntimeBot:
pass
await self.ap.monitoring_service.record_feedback(
+ self.execution_context,
feedback_id=event.feedback_id,
feedback_type=event.feedback_type,
feedback_content=event.feedback_content,
@@ -405,7 +485,7 @@ class PlatformManager:
bots: list[RuntimeBot]
- websocket_proxy_bot: RuntimeBot
+ websocket_proxy_bots: dict[str, RuntimeBot]
adapter_components: list[engine.Component]
@@ -414,6 +494,7 @@ class PlatformManager:
def __init__(self, ap: app.Application = None):
self.ap = ap
self.bots = []
+ self.websocket_proxy_bots = {}
self.adapter_components = []
self.adapter_dict = {}
@@ -435,19 +516,104 @@ class PlatformManager:
if disabled_adapters:
self.adapter_components = [c for c in self.adapter_components if c.metadata.name not in disabled_adapters]
- # initialize websocket adapter
- websocket_adapter_class = self.adapter_dict['websocket']
- websocket_logger = EventLogger(name='websocket-adapter', ap=self.ap)
- websocket_adapter_inst = websocket_adapter_class(
- {},
- websocket_logger,
- ap=self.ap,
- )
+ await self.load_bots_from_db()
- self.websocket_proxy_bot = RuntimeBot(
+ # OSS may have no persisted bots. Its singleton Workspace still needs
+ # a debug WebSocket proxy. SaaS creates proxies lazily from an explicit
+ # request/runtime context instead of guessing among Workspaces.
+ if not self.websocket_proxy_bots:
+ try:
+ binding = await self.ap.workspace_service.get_execution_binding()
+ except Exception:
+ pass
+ else:
+ await self.get_websocket_proxy_bot(
+ ExecutionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
+ )
+ )
+
+ @property
+ def websocket_proxy_bot(self) -> RuntimeBot:
+ """Compatibility accessor that is safe only for a singleton Workspace."""
+
+ if len(self.websocket_proxy_bots) != 1:
+ raise WorkspaceRequiredError('An explicit Workspace is required for the WebSocket proxy bot')
+ return next(iter(self.websocket_proxy_bots.values()))
+
+ @websocket_proxy_bot.setter
+ def websocket_proxy_bot(self, runtime_bot: RuntimeBot) -> None:
+ """Keep isolated tests that inject one proxy bot working."""
+
+ workspace_uuid = getattr(runtime_bot, 'workspace_uuid', '__test_singleton__')
+ self.websocket_proxy_bots = {workspace_uuid: runtime_bot}
+
+ @staticmethod
+ def _normalize_execution_context(
+ context: ExecutionContext | RequestContext,
+ *,
+ bot_uuid: str | None = None,
+ pipeline_uuid: str | None = None,
+ ) -> ExecutionContext:
+ if isinstance(context, RequestContext):
+ return ExecutionContext.from_request(
+ context,
+ bot_uuid=bot_uuid,
+ pipeline_uuid=pipeline_uuid,
+ )
+ if not isinstance(context, ExecutionContext):
+ raise WorkspaceRequiredError('Runtime operations require an ExecutionContext')
+ if not context.instance_uuid.strip() or not context.workspace_uuid.strip():
+ raise WorkspaceRequiredError('Runtime operations require an instance and Workspace')
+ if context.placement_generation <= 0:
+ raise WorkspaceRequiredError('Runtime operations require a positive placement generation')
+ updates = {}
+ if bot_uuid is not None:
+ if context.bot_uuid not in (None, bot_uuid):
+ raise WorkspaceRequiredError('Runtime bot UUID does not match its ExecutionContext')
+ updates['bot_uuid'] = bot_uuid
+ if pipeline_uuid is not None:
+ if context.pipeline_uuid not in (None, pipeline_uuid):
+ raise WorkspaceRequiredError('Runtime pipeline UUID does not match its ExecutionContext')
+ updates['pipeline_uuid'] = pipeline_uuid
+ return dataclasses.replace(context, **updates) if updates else context
+
+ async def get_websocket_proxy_bot(
+ self,
+ context: ExecutionContext | RequestContext,
+ ) -> RuntimeBot:
+ execution_context = self._normalize_execution_context(context)
+ existing = self.websocket_proxy_bots.get(execution_context.workspace_uuid)
+ if existing is not None:
+ if existing.placement_generation != execution_context.placement_generation:
+ raise WorkspaceRequiredError('WebSocket proxy placement generation is stale')
+ return existing
+
+ binding = await self.ap.workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ websocket_adapter_class = self.adapter_dict['websocket']
+ websocket_logger = EventLogger(
+ name='websocket-adapter',
+ ap=self.ap,
+ execution_context=execution_context,
+ owner='websocket-proxy-bot',
+ )
+ websocket_adapter_inst = websocket_adapter_class({}, websocket_logger, ap=self.ap)
+ proxy_context = dataclasses.replace(
+ execution_context,
+ instance_uuid=binding.instance_uuid,
+ bot_uuid='websocket-proxy-bot',
+ )
+ runtime_bot = RuntimeBot(
ap=self.ap,
bot_entity=persistence_bot.Bot(
uuid='websocket-proxy-bot',
+ workspace_uuid=binding.workspace_uuid,
name='WebSocket',
description='',
adapter='websocket',
@@ -456,13 +622,20 @@ class PlatformManager:
),
adapter=websocket_adapter_inst,
logger=websocket_logger,
+ execution_context=proxy_context,
)
- await self.websocket_proxy_bot.initialize()
+ await runtime_bot.initialize()
+ self.websocket_proxy_bots[binding.workspace_uuid] = runtime_bot
+ return runtime_bot
- await self.load_bots_from_db()
-
- def get_running_adapters(self) -> list[abstract_platform_adapter.AbstractMessagePlatformAdapter]:
- return [bot.adapter for bot in self.bots if bot.enable]
+ def get_running_adapters(
+ self,
+ context: ExecutionContext | RequestContext,
+ ) -> list[abstract_platform_adapter.AbstractMessagePlatformAdapter]:
+ execution_context = self._normalize_execution_context(context)
+ return [
+ bot.adapter for bot in self.bots if bot.enable and bot.workspace_uuid == execution_context.workspace_uuid
+ ]
async def load_bots_from_db(self):
self.ap.logger.info('Loading bots from db...')
@@ -476,7 +649,15 @@ class PlatformManager:
for bot in bots:
# load all bots here, enable or disable will be handled in runtime
try:
- await self.load_bot(bot)
+ binding = await self.ap.workspace_service.get_execution_binding(bot.workspace_uuid)
+ execution_context = ExecutionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ bot_uuid=bot.uuid,
+ trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
+ )
+ await self.load_bot(execution_context, bot)
except platform_errors.AdapterNotFoundError as e:
self.ap.logger.warning(f'Adapter {e.adapter_name} not found, skipping bot {bot.uuid}')
except Exception as e:
@@ -484,6 +665,7 @@ class PlatformManager:
async def load_bot(
self,
+ context: ExecutionContext | RequestContext,
bot_entity: persistence_bot.Bot | sqlalchemy.Row[persistence_bot.Bot] | dict,
) -> RuntimeBot:
"""加载机器人"""
@@ -492,7 +674,20 @@ class PlatformManager:
elif isinstance(bot_entity, dict):
bot_entity = persistence_bot.Bot(**bot_entity)
- logger = EventLogger(name=f'platform-adapter-{bot_entity.name}', ap=self.ap)
+ execution_context = self._normalize_execution_context(context, bot_uuid=bot_entity.uuid)
+ if bot_entity.workspace_uuid != execution_context.workspace_uuid:
+ raise WorkspaceRequiredError('Bot entity Workspace does not match its runtime context')
+ await self.ap.workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+
+ logger = EventLogger(
+ name=f'platform-adapter-{bot_entity.name}',
+ ap=self.ap,
+ execution_context=execution_context,
+ owner=bot_entity.uuid,
+ )
if bot_entity.adapter not in self.adapter_dict:
raise platform_errors.AdapterNotFoundError(bot_entity.adapter)
@@ -508,7 +703,13 @@ class PlatformManager:
if hasattr(adapter_inst, 'set_bot_uuid'):
adapter_inst.set_bot_uuid(bot_entity.uuid)
- runtime_bot = RuntimeBot(ap=self.ap, bot_entity=bot_entity, adapter=adapter_inst, logger=logger)
+ runtime_bot = RuntimeBot(
+ ap=self.ap,
+ bot_entity=bot_entity,
+ adapter=adapter_inst,
+ logger=logger,
+ execution_context=execution_context,
+ )
await runtime_bot.initialize()
@@ -516,17 +717,53 @@ class PlatformManager:
return runtime_bot
- async def get_bot_by_uuid(self, bot_uuid: str) -> RuntimeBot | None:
- if self.websocket_proxy_bot and self.websocket_proxy_bot.bot_entity.uuid == bot_uuid:
- return self.websocket_proxy_bot
+ async def get_bot_by_uuid(
+ self,
+ context: ExecutionContext | RequestContext,
+ bot_uuid: str,
+ ) -> RuntimeBot | None:
+ execution_context = self._normalize_execution_context(context, bot_uuid=bot_uuid)
+ proxy_bot = self.websocket_proxy_bots.get(execution_context.workspace_uuid)
+ if proxy_bot and proxy_bot.bot_entity.uuid == bot_uuid:
+ if proxy_bot.placement_generation != execution_context.placement_generation:
+ return None
+ return proxy_bot
for bot in self.bots:
- if bot.bot_entity.uuid == bot_uuid:
+ if (
+ bot.workspace_uuid == execution_context.workspace_uuid
+ and bot.placement_generation == execution_context.placement_generation
+ and bot.bot_entity.uuid == bot_uuid
+ ):
return bot
return None
- async def remove_bot(self, bot_uuid: str):
+ async def resolve_public_bot(self, route_key: str) -> RuntimeBot | None:
+ """Resolve an opaque public bot UUID without consulting request headers."""
+
+ try:
+ normalized = str(uuid.UUID(route_key))
+ except (ValueError, AttributeError, TypeError):
+ return None
+ for bot in self.bots:
+ if bot.bot_entity.uuid == normalized:
+ try:
+ await self.ap.workspace_service.get_execution_binding(
+ bot.workspace_uuid,
+ expected_generation=bot.placement_generation,
+ )
+ except Exception:
+ return None
+ return bot
+ return None
+
+ async def remove_bot(
+ self,
+ context: ExecutionContext | RequestContext,
+ bot_uuid: str,
+ ) -> None:
+ execution_context = self._normalize_execution_context(context, bot_uuid=bot_uuid)
for bot in self.bots[:]:
- if bot.bot_entity.uuid == bot_uuid:
+ if bot.workspace_uuid == execution_context.workspace_uuid and bot.bot_entity.uuid == bot_uuid:
if bot.enable:
await bot.shutdown()
self.bots.remove(bot)
@@ -551,13 +788,17 @@ class PlatformManager:
async def run(self):
# This method will only be called when the application launching
- await self.websocket_proxy_bot.run()
+ for proxy_bot in self.websocket_proxy_bots.values():
+ await proxy_bot.run()
for bot in self.bots:
if bot.enable:
await bot.run()
async def shutdown(self):
+ for proxy_bot in self.websocket_proxy_bots.values():
+ if proxy_bot.enable:
+ await proxy_bot.shutdown()
for bot in self.bots:
if bot.enable:
await bot.shutdown()
diff --git a/src/langbot/pkg/platform/logger.py b/src/langbot/pkg/platform/logger.py
index 681648656..8fcdc1795 100644
--- a/src/langbot/pkg/platform/logger.py
+++ b/src/langbot/pkg/platform/logger.py
@@ -9,6 +9,7 @@ import traceback
import uuid
from ..core import app
+from ..api.http.context import ExecutionContext
import langbot_plugin.api.entities.builtin.platform.message as platform_message
import langbot_plugin.api.definition.abstract.platform.event_logger as abstract_platform_event_logger
@@ -65,13 +66,21 @@ class EventLogger(abstract_platform_event_logger.AbstractEventLogger):
logs: list[EventLog]
+ execution_context: ExecutionContext
+
+ owner: str
+
def __init__(
self,
name: str,
ap: app.Application,
+ execution_context: ExecutionContext,
+ owner: str,
):
self.name = name
self.ap = ap
+ self.execution_context = execution_context
+ self.owner = owner
self.logs = []
self.seq_id_inc = 0
@@ -121,7 +130,11 @@ class EventLogger(abstract_platform_event_logger.AbstractEventLogger):
if len(self.logs) > MAX_LOG_COUNT:
for i in range(DELETE_COUNT_PER_TIME):
for image_key in self.logs[i].images: # type: ignore
- await self.ap.storage_mgr.storage_provider.delete(image_key)
+ await self.ap.storage_mgr.delete_scoped_object_key(
+ self.execution_context,
+ image_key,
+ expected_owner_type='bot_log',
+ )
self.logs = self.logs[DELETE_COUNT_PER_TIME:]
async def _add_log(
@@ -149,8 +162,14 @@ class EventLogger(abstract_platform_event_logger.AbstractEventLogger):
extension = mimetypes.guess_extension(mime_type)
if extension is None:
extension = '.jpg'
- image_key = f'bot_log_images/{message_session_id}-{uuid.uuid4()}{extension}'
- await self.ap.storage_mgr.storage_provider.save(image_key, img_bytes)
+ logical_key = f'{message_session_id}-{uuid.uuid4()}{extension}'
+ image_key = await self.ap.storage_mgr.save_scoped(
+ self.execution_context,
+ owner_type='bot_log',
+ owner=self.owner,
+ key=logical_key,
+ value=img_bytes,
+ )
image_keys.append(image_key)
self.logs.append(
diff --git a/src/langbot/pkg/platform/sources/http_bot.py b/src/langbot/pkg/platform/sources/http_bot.py
index 16a891991..6a95bb460 100644
--- a/src/langbot/pkg/platform/sources/http_bot.py
+++ b/src/langbot/pkg/platform/sources/http_bot.py
@@ -311,18 +311,40 @@ class HttpBotAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter):
async def _reset_session(self, launcher_type: str, launcher_id: str) -> bool:
"""Drop the matching session so the next message starts a fresh conversation."""
+ execution_context = getattr(self.logger, 'execution_context', None)
+ if (
+ execution_context is None
+ or not execution_context.instance_uuid
+ or not execution_context.workspace_uuid
+ or execution_context.placement_generation <= 0
+ or not self.bot_uuid
+ ):
+ raise RuntimeError('http_bot reset requires a trusted execution scope')
+ expected_prefix = (
+ execution_context.instance_uuid,
+ execution_context.workspace_uuid,
+ execution_context.placement_generation,
+ self.bot_uuid,
+ launcher_type,
+ )
+
sess_mgr = self.ap.sess_mgr
before = len(sess_mgr.session_list)
sess_mgr.session_list = [
- s
- for s in sess_mgr.session_list
- if not (
- str(s.launcher_type.value if hasattr(s.launcher_type, 'value') else s.launcher_type) == launcher_type
- and str(s.launcher_id) == launcher_id
- )
+ s for s in sess_mgr.session_list if not self._matches_session_scope(s, expected_prefix, launcher_id)
]
return len(sess_mgr.session_list) < before
+ @staticmethod
+ def _matches_session_scope(session, expected_prefix: tuple[str, str, int, str, str], launcher_id: str) -> bool:
+ session_key = getattr(session, '_langbot_session_key', None)
+ return (
+ isinstance(session_key, tuple)
+ and len(session_key) == 6
+ and session_key[:5] == expected_prefix
+ and str(session_key[5]) == launcher_id
+ )
+
# -- outbound -------------------------------------------------------------
@staticmethod
diff --git a/src/langbot/pkg/platform/sources/openclaw_weixin.py b/src/langbot/pkg/platform/sources/openclaw_weixin.py
index 9253f90e4..69b9f8fc2 100644
--- a/src/langbot/pkg/platform/sources/openclaw_weixin.py
+++ b/src/langbot/pkg/platform/sources/openclaw_weixin.py
@@ -33,6 +33,8 @@ import langbot_plugin.api.entities.builtin.platform.entities as platform_entitie
import langbot_plugin.api.entities.builtin.platform.events as platform_events
import langbot_plugin.api.entities.builtin.platform.message as platform_message
+from langbot.pkg.api.http.context import ExecutionContext
+
class OpenClawWeixinMessageConverter(abstract_platform_adapter.AbstractMessageConverter):
"""Converts between LangBot MessageChain and OpenClaw WeChat message items."""
@@ -278,8 +280,22 @@ class OpenClawWeixinAdapter(abstract_platform_adapter.AbstractMessagePlatformAda
return
try:
ap = self.logger.ap
+ execution_context = getattr(self.logger, 'execution_context', None)
+ if not isinstance(execution_context, ExecutionContext):
+ raise RuntimeError('Weixin Bot config persistence requires an ExecutionContext')
+ if execution_context.bot_uuid != self._bot_uuid:
+ raise RuntimeError('Weixin Bot UUID does not match its ExecutionContext')
+
+ binding = await ap.workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ if binding.instance_uuid != execution_context.instance_uuid:
+ raise RuntimeError('Weixin Bot Workspace belongs to another LangBot instance')
+
await ap.persistence_mgr.execute_async(
sqlalchemy.update(persistence_bot.Bot)
+ .where(persistence_bot.Bot.workspace_uuid == execution_context.workspace_uuid)
.where(persistence_bot.Bot.uuid == self._bot_uuid)
.values(adapter_config=self.config)
)
diff --git a/src/langbot/pkg/platform/sources/websocket_adapter.py b/src/langbot/pkg/platform/sources/websocket_adapter.py
index a66d4ac9a..9f9d9758d 100644
--- a/src/langbot/pkg/platform/sources/websocket_adapter.py
+++ b/src/langbot/pkg/platform/sources/websocket_adapter.py
@@ -1,6 +1,7 @@
"""WebSocket适配器 - 支持双向通信的IM系统"""
import asyncio
+import contextvars
import logging
import typing
from datetime import datetime
@@ -13,9 +14,13 @@ import langbot_plugin.api.entities.builtin.platform.events as platform_events
import langbot_plugin.api.entities.builtin.platform.entities as platform_entities
import langbot_plugin.api.definition.abstract.platform.event_logger as abstract_platform_logger
from ...core import app
-from .websocket_manager import WebSocketConnection, is_valid_session_id, ws_connection_manager
+from .websocket_manager import WebSocketConnection, WebSocketScope, is_valid_session_id, ws_connection_manager
logger = logging.getLogger(__name__)
+_current_pipeline_uuid: contextvars.ContextVar[str | None] = contextvars.ContextVar(
+ 'websocket_pipeline_uuid',
+ default=None,
+)
class WebSocketMessage(pydantic.BaseModel):
@@ -113,9 +118,19 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
return None
return pipeline_uuid, session_id
- @classmethod
- async def _get_connection_from_target(cls, target_id: str):
+ def _scope(self) -> WebSocketScope:
+ """Return this adapter's immutable runtime placement."""
+
+ return WebSocketScope.from_context(self.logger.execution_context)
+
+ def get_pipeline_uuid_override(self) -> str | None:
+ """Return the connection pipeline propagated into the listener task."""
+
+ return _current_pipeline_uuid.get()
+
+ async def _get_connection_from_target(self, target_id: str):
"""Resolve a person or group WebSocket launcher to its connection."""
+ scope = self._scope()
target_value = str(target_id)
for prefix in ('websocket_', 'websocketgroup_'):
if target_value.startswith(prefix):
@@ -123,14 +138,18 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
break
else:
return None
- connection = await ws_connection_manager.get_connection(target)
+ connection = await ws_connection_manager.get_connection(target, scope=scope)
if connection is not None:
return connection
- embed_target = cls._parse_embed_target(target_id)
+ embed_target = self._parse_embed_target(target_id)
if embed_target is not None:
pipeline_uuid, session_id = embed_target
- return await ws_connection_manager.get_connection_by_session_id(session_id, pipeline_uuid)
- return await ws_connection_manager.get_connection_by_session_id(target)
+ return await ws_connection_manager.get_connection_by_session_id(
+ session_id,
+ scope=scope,
+ pipeline_uuid=pipeline_uuid,
+ )
+ return await ws_connection_manager.get_connection_by_session_id(target, scope=scope)
async def _get_message_context(self, message_source) -> tuple[str, str | None]:
"""Resolve the originating pipeline and browser session for a reply."""
@@ -142,7 +161,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
embed_target = self._parse_embed_target(sender_id)
if embed_target is not None:
return embed_target
- return typing.cast(str, self.ap.platform_mgr.websocket_proxy_bot.bot_entity.use_pipeline_uuid), None
+ raise ValueError('WebSocket reply target is not bound to this adapter scope')
async def send_message(
self,
@@ -160,16 +179,17 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
if connection is not None:
pipeline_uuid = connection.pipeline_uuid
session_id = connection.session_id
+ scope = connection.scope
else:
embed_target = self._parse_embed_target(target_id)
if embed_target is not None:
pipeline_uuid, session_id = embed_target
else:
- pipeline_uuid = typing.cast(
- str,
- self.ap.platform_mgr.websocket_proxy_bot.bot_entity.use_pipeline_uuid,
- )
+ pipeline_uuid = str(target_id).strip()
+ if not pipeline_uuid:
+ raise ValueError('WebSocket target pipeline is required')
session_id = None
+ scope = self._scope()
session_type = 'group' if target_type == 'group' else 'person'
conversation_key = self._conversation_key(pipeline_uuid, session_id)
@@ -195,6 +215,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
'session_type': session_type,
'data': message_data.model_dump(),
},
+ scope=scope,
session_type=session_type,
session_id=session_id,
)
@@ -216,6 +237,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
)
pipeline_uuid, session_id = await self._get_message_context(message_source)
+ scope = self._scope()
session_type = 'group' if isinstance(message_source, platform_events.GroupMessage) else 'person'
conversation_key = self._conversation_key(pipeline_uuid, session_id)
@@ -239,6 +261,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
'session_type': session_type,
'data': message_data.model_dump(),
},
+ scope=scope,
session_type=session_type,
session_id=session_id,
)
@@ -262,6 +285,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
)
pipeline_uuid, session_id = await self._get_message_context(message_source)
+ scope = self._scope()
session_type = 'group' if isinstance(message_source, platform_events.GroupMessage) else 'person'
conversation_key = self._conversation_key(pipeline_uuid, session_id)
message_list = session.get_message_list(conversation_key)
@@ -316,6 +340,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
'session_type': session_type,
'data': message_data.model_dump(),
},
+ scope=scope,
session_type=session_type,
session_id=session_id,
)
@@ -360,7 +385,11 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
message = await asyncio.wait_for(self.outbound_message_queue.get(), timeout=0.1)
# 广播到所有相关连接
target_id = message.get('target_id', '')
- await ws_connection_manager.broadcast_to_pipeline(target_id, message)
+ await ws_connection_manager.broadcast_to_pipeline(
+ target_id,
+ message,
+ scope=self._scope(),
+ )
except asyncio.TimeoutError:
pass
@@ -372,7 +401,11 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
"""停止适配器"""
pass
- async def _process_image_components(self, message_chain_obj: list):
+ async def _process_image_components(
+ self,
+ connection: WebSocketConnection,
+ message_chain_obj: list,
+ ):
"""
处理消息链中的图片、语音和文件组件,将 path 转换为 base64
@@ -387,14 +420,28 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
import base64
import mimetypes
- storage_mgr = self.ap.storage_mgr
+ attachments = [
+ component
+ for component in message_chain_obj
+ if component.get('path') and component.get('type') in ('Image', 'Voice', 'File')
+ ]
+ if not attachments:
+ return
- for component in message_chain_obj:
+ storage_mgr = self.ap.storage_mgr
+ execution_context = connection.execution_context
+ expected_prefix = storage_mgr.scoped_prefix(execution_context, owner_type='upload_image')
+
+ for component in attachments:
comp_type = component.get('type', '')
comp_path = component.get('path', '')
- if not comp_path or comp_type not in ('Image', 'Voice', 'File'):
- continue
+ if not comp_path.startswith(expected_prefix) or not storage_mgr.is_scoped_object_key(
+ comp_path,
+ expected_owner_type='upload_image',
+ ):
+ await self.logger.warning(f'Rejected {comp_type} attachment outside the WebSocket connection scope')
+ raise ValueError('Attachment key does not belong to this WebSocket connection')
try:
file_content = await storage_mgr.storage_provider.load(comp_path)
@@ -416,10 +463,15 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
mime_type = mimetypes.guess_type(comp_path)[0] or 'application/octet-stream'
component['base64'] = f'data:{mime_type};base64,{base64_str}'
- await storage_mgr.storage_provider.delete(comp_path)
+ await storage_mgr.delete_scoped_object_key(
+ execution_context,
+ comp_path,
+ expected_owner_type='upload_image',
+ )
component['path'] = ''
except Exception as e:
await self.logger.error(f'Failed to load {comp_type} file {comp_path}: {e}')
+ raise
async def handle_websocket_message(
self,
@@ -451,7 +503,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
message_chain_obj = message_data.get('message', [])
- await self._process_image_components(message_chain_obj)
+ await self._process_image_components(connection, message_chain_obj)
message_chain = platform_message.MessageChain.model_validate(message_chain_obj)
@@ -476,6 +528,7 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
'session_type': session_type,
'data': user_message.model_dump(),
},
+ scope=connection.scope,
session_type=session_type,
session_id=connection.session_id,
)
@@ -506,11 +559,6 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
sender=sender, message_chain=message_chain, time=datetime.now().timestamp()
)
- # 设置流水线UUID (proxy bot always needs it for reply_message routing)
- self.ap.platform_mgr.websocket_proxy_bot.bot_entity.use_pipeline_uuid = pipeline_uuid
- if owner_bot is not None:
- owner_bot.bot_entity.use_pipeline_uuid = pipeline_uuid
-
# 异步触发事件处理
# Use owner_bot's listeners if available, otherwise fall back to proxy bot
listeners = (
@@ -525,7 +573,11 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
owner_bot.adapter.set_ws_adapter(self)
callback_adapter = owner_bot.adapter if (owner_bot and hasattr(owner_bot, 'adapter')) else self
if event.__class__ in listeners:
- asyncio.create_task(listeners[event.__class__](event, callback_adapter))
+ token = _current_pipeline_uuid.set(pipeline_uuid)
+ try:
+ asyncio.create_task(listeners[event.__class__](event, callback_adapter))
+ finally:
+ _current_pipeline_uuid.reset(token)
def get_websocket_messages(
self,
@@ -558,11 +610,15 @@ class WebSocketAdapter(abstract_platform_adapter.AbstractMessagePlatformAdapter)
if session_type == 'group'
else f'websocket_{pipeline_uuid}:{session_id}'
)
+ scope = self._scope()
self.ap.sess_mgr.session_list = [
candidate_session
for candidate_session in self.ap.sess_mgr.session_list
if not (
- str(
+ getattr(candidate_session, 'instance_uuid', None) == scope.instance_uuid
+ and getattr(candidate_session, 'workspace_uuid', None) == scope.workspace_uuid
+ and getattr(candidate_session, 'placement_generation', None) == scope.placement_generation
+ and str(
candidate_session.launcher_type.value
if hasattr(candidate_session.launcher_type, 'value')
else candidate_session.launcher_type
diff --git a/src/langbot/pkg/platform/sources/websocket_manager.py b/src/langbot/pkg/platform/sources/websocket_manager.py
index c0c4d1807..216043e9a 100644
--- a/src/langbot/pkg/platform/sources/websocket_manager.py
+++ b/src/langbot/pkg/platform/sources/websocket_manager.py
@@ -1,6 +1,7 @@
"""WebSocket连接管理器 - 管理多个并发WebSocket连接"""
import asyncio
+import dataclasses
import logging
import typing
import uuid
@@ -8,10 +9,35 @@ from datetime import datetime
import pydantic
+from ...api.http.context import ExecutionContext
+
logger = logging.getLogger(__name__)
_SESSION_FILTER_UNSET = object()
+@dataclasses.dataclass(frozen=True, slots=True)
+class WebSocketScope:
+ """Trusted runtime placement carried by every WebSocket connection."""
+
+ instance_uuid: str
+ workspace_uuid: str
+ placement_generation: int
+
+ def __post_init__(self) -> None:
+ if not self.instance_uuid.strip() or not self.workspace_uuid.strip():
+ raise ValueError('WebSocket scope requires an instance and Workspace')
+ if self.placement_generation <= 0:
+ raise ValueError('WebSocket scope requires a positive placement generation')
+
+ @classmethod
+ def from_context(cls, context: typing.Any) -> 'WebSocketScope':
+ return cls(
+ instance_uuid=str(getattr(context, 'instance_uuid', '')),
+ workspace_uuid=str(getattr(context, 'workspace_uuid', '')),
+ placement_generation=int(getattr(context, 'placement_generation', 0)),
+ )
+
+
def is_valid_session_id(value: str) -> bool:
"""Accept only canonical random UUIDs for client conversation identifiers."""
try:
@@ -29,6 +55,15 @@ class WebSocketConnection(pydantic.BaseModel):
connection_id: str = pydantic.Field(default_factory=lambda: str(uuid.uuid4()))
"""连接唯一ID"""
+ instance_uuid: str
+ """Owning LangBot instance."""
+
+ workspace_uuid: str
+ """Owning Workspace."""
+
+ placement_generation: int
+ """Workspace placement generation captured at connect time."""
+
pipeline_uuid: str
"""关联的流水线UUID"""
@@ -56,6 +91,25 @@ class WebSocketConnection(pydantic.BaseModel):
metadata: dict = pydantic.Field(default_factory=dict)
"""连接元数据(可存储额外信息)"""
+ @property
+ def scope(self) -> WebSocketScope:
+ return WebSocketScope(
+ instance_uuid=self.instance_uuid,
+ workspace_uuid=self.workspace_uuid,
+ placement_generation=self.placement_generation,
+ )
+
+ @property
+ def execution_context(self) -> ExecutionContext:
+ """Return the storage/runtime context captured for this connection."""
+
+ return ExecutionContext(
+ instance_uuid=self.instance_uuid,
+ workspace_uuid=self.workspace_uuid,
+ placement_generation=self.placement_generation,
+ pipeline_uuid=self.pipeline_uuid,
+ )
+
class WebSocketConnectionManager:
"""WebSocket连接管理器 - 支持多连接并发"""
@@ -64,11 +118,11 @@ class WebSocketConnectionManager:
self.connections: dict[str, WebSocketConnection] = {}
"""所有活跃连接 {connection_id: connection}"""
- self.pipeline_connections: dict[str, set[str]] = {}
- """流水线到连接的映射 {pipeline_uuid: {connection_id, ...}}"""
+ self.pipeline_connections: dict[tuple[str, str, int, str], set[str]] = {}
+ """Scoped pipeline to connection mapping."""
- self.session_connections: dict[str, set[str]] = {}
- """会话类型到连接的映射 {session_type: {connection_id, ...}}"""
+ self.session_connections: dict[tuple[str, str, int, str], set[str]] = {}
+ """Scoped session-type to connection mapping."""
self._lock = asyncio.Lock()
"""线程锁,保护并发访问"""
@@ -76,6 +130,7 @@ class WebSocketConnectionManager:
async def add_connection(
self,
websocket: typing.Any,
+ scope: WebSocketScope,
pipeline_uuid: str,
session_type: str,
metadata: dict | None = None,
@@ -84,6 +139,9 @@ class WebSocketConnectionManager:
"""Register a WebSocket connection and its optional embed session."""
async with self._lock:
connection = WebSocketConnection(
+ instance_uuid=scope.instance_uuid,
+ workspace_uuid=scope.workspace_uuid,
+ placement_generation=scope.placement_generation,
pipeline_uuid=pipeline_uuid,
session_type=session_type,
session_id=session_id,
@@ -94,18 +152,21 @@ class WebSocketConnectionManager:
self.connections[connection.connection_id] = connection
# 更新流水线映射
- if pipeline_uuid not in self.pipeline_connections:
- self.pipeline_connections[pipeline_uuid] = set()
- self.pipeline_connections[pipeline_uuid].add(connection.connection_id)
+ pipeline_key = self._pipeline_key(scope, pipeline_uuid)
+ if pipeline_key not in self.pipeline_connections:
+ self.pipeline_connections[pipeline_key] = set()
+ self.pipeline_connections[pipeline_key].add(connection.connection_id)
# 更新会话类型映射
- if session_type not in self.session_connections:
- self.session_connections[session_type] = set()
- self.session_connections[session_type].add(connection.connection_id)
+ session_key = self._session_key(scope, session_type)
+ if session_key not in self.session_connections:
+ self.session_connections[session_key] = set()
+ self.session_connections[session_key].add(connection.connection_id)
logger.debug(
f'WebSocket connection established: {connection.connection_id} '
- f'(pipeline={pipeline_uuid}, session_type={session_type})'
+ f'(workspace={scope.workspace_uuid}, generation={scope.placement_generation}, '
+ f'pipeline={pipeline_uuid}, session_type={session_type})'
)
return connection
@@ -120,28 +181,59 @@ class WebSocketConnectionManager:
connection.is_active = False
# 从流水线映射中移除
- if connection.pipeline_uuid in self.pipeline_connections:
- self.pipeline_connections[connection.pipeline_uuid].discard(connection_id)
- if not self.pipeline_connections[connection.pipeline_uuid]:
- del self.pipeline_connections[connection.pipeline_uuid]
+ pipeline_key = self._pipeline_key(connection.scope, connection.pipeline_uuid)
+ if pipeline_key in self.pipeline_connections:
+ self.pipeline_connections[pipeline_key].discard(connection_id)
+ if not self.pipeline_connections[pipeline_key]:
+ del self.pipeline_connections[pipeline_key]
# 从会话类型映射中移除
- if connection.session_type in self.session_connections:
- self.session_connections[connection.session_type].discard(connection_id)
- if not self.session_connections[connection.session_type]:
- del self.session_connections[connection.session_type]
+ session_key = self._session_key(connection.scope, connection.session_type)
+ if session_key in self.session_connections:
+ self.session_connections[session_key].discard(connection_id)
+ if not self.session_connections[session_key]:
+ del self.session_connections[session_key]
del self.connections[connection_id]
logger.debug(f'WebSocket connection disconnected: {connection_id}')
- async def get_connection(self, connection_id: str) -> WebSocketConnection | None:
- """Get a connection by its transport identifier."""
- return self.connections.get(connection_id)
+ @staticmethod
+ def _pipeline_key(scope: WebSocketScope, pipeline_uuid: str) -> tuple[str, str, int, str]:
+ return (
+ scope.instance_uuid,
+ scope.workspace_uuid,
+ scope.placement_generation,
+ pipeline_uuid,
+ )
+
+ @staticmethod
+ def _session_key(scope: WebSocketScope, session_type: str) -> tuple[str, str, int, str]:
+ return (
+ scope.instance_uuid,
+ scope.workspace_uuid,
+ scope.placement_generation,
+ session_type,
+ )
+
+ async def get_connection(
+ self,
+ connection_id: str,
+ *,
+ scope: WebSocketScope,
+ ) -> WebSocketConnection | None:
+ """Get a connection only when it belongs to the expected placement."""
+
+ connection = self.connections.get(connection_id)
+ if connection is None or connection.scope != scope:
+ return None
+ return connection
async def get_connection_by_session_id(
self,
session_id: str,
+ *,
+ scope: WebSocketScope,
pipeline_uuid: str | None = None,
) -> WebSocketConnection | None:
"""Get an active embed connection by its stable browser session identifier."""
@@ -149,25 +241,38 @@ class WebSocketConnectionManager:
if (
connection.session_id == session_id
and connection.is_active
+ and connection.scope == scope
and (pipeline_uuid is None or connection.pipeline_uuid == pipeline_uuid)
):
return connection
return None
- async def get_connections_by_pipeline(self, pipeline_uuid: str) -> list[WebSocketConnection]:
+ async def get_connections_by_pipeline(
+ self,
+ pipeline_uuid: str,
+ *,
+ scope: WebSocketScope,
+ ) -> list[WebSocketConnection]:
"""获取指定流水线的所有连接"""
- connection_ids = self.pipeline_connections.get(pipeline_uuid, set())
+ connection_ids = self.pipeline_connections.get(self._pipeline_key(scope, pipeline_uuid), set())
return [self.connections[cid] for cid in connection_ids if cid in self.connections]
- async def get_connections_by_session_type(self, session_type: str) -> list[WebSocketConnection]:
+ async def get_connections_by_session_type(
+ self,
+ session_type: str,
+ *,
+ scope: WebSocketScope,
+ ) -> list[WebSocketConnection]:
"""获取指定会话类型的所有连接"""
- connection_ids = self.session_connections.get(session_type, set())
+ connection_ids = self.session_connections.get(self._session_key(scope, session_type), set())
return [self.connections[cid] for cid in connection_ids if cid in self.connections]
async def broadcast_to_pipeline(
self,
pipeline_uuid: str,
message: dict,
+ *,
+ scope: WebSocketScope,
session_type: str | None = None,
session_id: typing.Any = _SESSION_FILTER_UNSET,
):
@@ -180,7 +285,7 @@ class WebSocketConnectionManager:
session_id: Embed conversation filter. Omit it to broadcast across
conversations; pass ``None`` to target non-embed connections.
"""
- connections = await self.get_connections_by_pipeline(pipeline_uuid)
+ connections = await self.get_connections_by_pipeline(pipeline_uuid, scope=scope)
if session_type is not None:
connections = [conn for conn in connections if conn.session_type == session_type]
@@ -196,7 +301,7 @@ class WebSocketConnectionManager:
async def send_to_connection(self, connection_id: str, message: dict):
"""向指定连接发送消息"""
- connection = await self.get_connection(connection_id)
+ connection = self.connections.get(connection_id)
if not connection or not connection.is_active:
logger.warning(f'Attempt to send message to invalid connection: {connection_id}')
return
@@ -210,17 +315,24 @@ class WebSocketConnectionManager:
async def update_activity(self, connection_id: str):
"""更新连接活跃时间"""
- connection = await self.get_connection(connection_id)
+ connection = self.connections.get(connection_id)
if connection:
connection.last_active = datetime.now()
- def get_stats(self) -> dict:
- """获取连接统计信息"""
+ def get_stats(self, *, scope: WebSocketScope) -> dict:
+ """Return connection statistics for one trusted placement."""
+
+ scoped_connections = [connection for connection in self.connections.values() if connection.scope == scope]
+ pipelines: dict[str, int] = {}
+ session_types: dict[str, int] = {}
+ for connection in scoped_connections:
+ pipelines[connection.pipeline_uuid] = pipelines.get(connection.pipeline_uuid, 0) + 1
+ session_types[connection.session_type] = session_types.get(connection.session_type, 0) + 1
return {
- 'total_connections': len(self.connections),
- 'pipelines': len(self.pipeline_connections),
- 'connections_by_pipeline': {k: len(v) for k, v in self.pipeline_connections.items()},
- 'connections_by_session_type': {k: len(v) for k, v in self.session_connections.items()},
+ 'total_connections': len(scoped_connections),
+ 'pipelines': len(pipelines),
+ 'connections_by_pipeline': pipelines,
+ 'connections_by_session_type': session_types,
}
diff --git a/src/langbot/pkg/platform/webhook_pusher.py b/src/langbot/pkg/platform/webhook_pusher.py
index f3cf39b27..79eb1de24 100644
--- a/src/langbot/pkg/platform/webhook_pusher.py
+++ b/src/langbot/pkg/platform/webhook_pusher.py
@@ -5,6 +5,7 @@ import logging
import aiohttp
from langbot.pkg.utils import httpclient
+from langbot.pkg.api.http.context import ExecutionContext
import uuid
from typing import TYPE_CHECKING
@@ -24,14 +25,20 @@ class WebhookPusher:
self.ap = ap
self.logger = self.ap.logger
- async def push_person_message(self, event: platform_events.FriendMessage, bot_uuid: str, adapter_name: str) -> bool:
+ async def push_person_message(
+ self,
+ execution_context: ExecutionContext,
+ event: platform_events.FriendMessage,
+ bot_uuid: str,
+ adapter_name: str,
+ ) -> bool:
"""Push person message event to webhooks
Returns:
bool: True if any webhook responded with skip_pipeline=true, False otherwise
"""
try:
- webhooks = await self.ap.webhook_service.get_enabled_webhooks()
+ webhooks = await self.ap.webhook_service.get_enabled_webhooks(execution_context)
if not webhooks:
return False
@@ -67,14 +74,20 @@ class WebhookPusher:
self.logger.error(f'Failed to push person message to webhooks: {e}')
return False
- async def push_group_message(self, event: platform_events.GroupMessage, bot_uuid: str, adapter_name: str) -> bool:
+ async def push_group_message(
+ self,
+ execution_context: ExecutionContext,
+ event: platform_events.GroupMessage,
+ bot_uuid: str,
+ adapter_name: str,
+ ) -> bool:
"""Push group message event to webhooks
Returns:
bool: True if any webhook responded with skip_pipeline=true, False otherwise
"""
try:
- webhooks = await self.ap.webhook_service.get_enabled_webhooks()
+ webhooks = await self.ap.webhook_service.get_enabled_webhooks(execution_context)
if not webhooks:
return False
diff --git a/src/langbot/pkg/plugin/connector.py b/src/langbot/pkg/plugin/connector.py
index 6075d4b68..0cf9dc302 100644
--- a/src/langbot/pkg/plugin/connector.py
+++ b/src/langbot/pkg/plugin/connector.py
@@ -8,6 +8,7 @@ import zipfile
from typing import Any
import typing
import os
+import secrets
import sys
import httpx
import sqlalchemy
@@ -32,8 +33,17 @@ from langbot_plugin.api.entities.builtin.command import (
errors as command_errors,
)
from langbot_plugin.runtime.plugin.mgr import PluginInstallSource
+from langbot_plugin.runtime.security import (
+ PLUGIN_RUNTIME_CONTROL_TOKEN_ENV,
+ PLUGIN_RUNTIME_CONTROL_TOKEN_HEADER,
+ validate_runtime_secret,
+)
+from langbot_plugin.entities.io.context import ActionContext
from ..core import taskmgr
from ..entity.persistence import plugin as persistence_plugin
+from ..api.http.context import ExecutionContext
+from ..api.http.service.tenant import TenantContext, require_workspace_uuid
+from ..workspace.errors import WorkspaceNotFoundError
class PluginRuntimeNotConnectedError(RuntimeError):
@@ -66,10 +76,68 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
runtime_disconnect_callback: typing.Callable[
[PluginRuntimeConnector], typing.Coroutine[typing.Any, typing.Any, None]
],
+ action_context: ActionContext | None = None,
):
super().__init__(ap)
self.runtime_disconnect_callback = runtime_disconnect_callback
self.is_enable_plugin = self.ap.instance_config.data.get('plugin', {}).get('enable', True)
+ self._configured_action_context = (
+ ActionContext.model_validate(action_context).without_installation() if action_context is not None else None
+ )
+ self._control_token = str(os.environ.get(PLUGIN_RUNTIME_CONTROL_TOKEN_ENV) or '').strip()
+
+ def _requires_explicit_workspace_binding(self) -> bool:
+ """Return whether this process may host more than the OSS singleton."""
+
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ policy = getattr(workspace_service, 'policy', None)
+ return getattr(policy, 'multi_workspace_enabled', False) is True
+
+ def _control_headers(self, *, allow_generate: bool) -> dict[str, str]:
+ if not self._control_token and allow_generate:
+ self._control_token = secrets.token_urlsafe(48)
+ try:
+ self._control_token = validate_runtime_secret(
+ self._control_token,
+ name=PLUGIN_RUNTIME_CONTROL_TOKEN_ENV,
+ )
+ except ValueError as exc:
+ raise PluginRuntimeNotConnectedError(
+ f'{PLUGIN_RUNTIME_CONTROL_TOKEN_ENV} must be configured with a strong shared secret '
+ 'for an external Plugin Runtime'
+ ) from exc
+ return {PLUGIN_RUNTIME_CONTROL_TOKEN_HEADER: self._control_token}
+
+ async def _resolve_action_context(self) -> ActionContext:
+ """Resolve the trusted connector binding; never use plugin input."""
+
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ if workspace_service is None:
+ raise RuntimeError('Plugin Runtime Workspace binding is unavailable')
+
+ if self._configured_action_context is not None:
+ configured = self._configured_action_context
+ binding = await workspace_service.get_execution_binding(
+ configured.workspace_uuid,
+ expected_generation=configured.placement_generation,
+ )
+ if binding.instance_uuid != configured.instance_uuid:
+ raise RuntimeError('Plugin Runtime Workspace binding belongs to another instance')
+ return ActionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ )
+
+ if self._requires_explicit_workspace_binding():
+ raise RuntimeError('Cloud plugin Runtime connectors require an explicit projected Workspace binding')
+
+ binding = await workspace_service.get_local_execution_binding()
+ return ActionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ )
async def heartbeat_loop(self):
while True:
@@ -84,6 +152,8 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
if not self.is_enable_plugin:
self.ap.logger.info('Plugin system is disabled.')
return
+ if self._configured_action_context is None and self._requires_explicit_workspace_binding():
+ raise RuntimeError('Cloud plugin Runtime connectors require an explicit projected Workspace binding')
async def new_connection_callback(connection: base_connection.Connection):
async def disconnect_callback(
@@ -99,7 +169,13 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
)
return False
- self.handler = handler.RuntimeConnectionHandler(connection, disconnect_callback, self.ap)
+ action_context = await self._resolve_action_context()
+ self.handler = handler.RuntimeConnectionHandler(
+ connection,
+ disconnect_callback,
+ self.ap,
+ action_context,
+ )
self.handler_task = asyncio.create_task(self.handler.run())
_ = await self.handler.ping()
@@ -107,12 +183,13 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
# downloads plugins from the same Space LangBot is bound to, rather
# than relying on the runtime's own env/default.
space_url = self.ap.instance_config.data.get('space', {}).get('url', '').rstrip('/')
- if space_url:
- try:
- await self.handler.set_runtime_config(cloud_service_url=space_url)
+ try:
+ await self.handler.set_runtime_config(cloud_service_url=space_url or None)
+ if space_url:
self.ap.logger.info(f'Pushed marketplace URL to plugin runtime: {space_url}')
- except Exception as e:
- self.ap.logger.warning(f'Failed to push runtime config: {e}')
+ except Exception as e:
+ self.ap.logger.warning(f'Failed to bind plugin runtime config: {e}')
+ raise
self.ap.logger.info('Connected to plugin runtime.')
await self.handler_task
@@ -120,6 +197,7 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
if platform.get_platform() == 'docker' or platform.use_websocket_to_connect_plugin_runtime(): # use websocket
self.ap.logger.info('use websocket to connect to plugin runtime')
+ control_headers = self._control_headers(allow_generate=False)
ws_url = self.ap.instance_config.data.get('plugin', {}).get(
'runtime_ws_url', 'ws://langbot_plugin_runtime:5400/control/ws'
)
@@ -137,6 +215,7 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
self.ctrl = ws_client_controller.WebSocketClientController(
ws_url=ws_url,
make_connection_failed_callback=make_connection_failed_callback,
+ additional_headers=control_headers,
)
task = self.ctrl.run(new_connection_callback)
elif platform.get_platform() == 'win32':
@@ -145,7 +224,13 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
# We have to launch runtime via cmd but communicate via ws.
self.ap.logger.info('(windows) use cmd to launch plugin runtime and communicate via ws')
- await self._start_runtime_subprocess('-m', 'langbot_plugin.cli.__init__', 'rt')
+ control_headers = self._control_headers(allow_generate=True)
+ await self._start_runtime_subprocess(
+ '-m',
+ 'langbot_plugin.cli.__init__',
+ 'rt',
+ env_overrides={PLUGIN_RUNTIME_CONTROL_TOKEN_ENV: self._control_token},
+ )
ws_url = 'ws://localhost:5400/control/ws'
@@ -164,6 +249,7 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
self.ctrl = ws_client_controller.WebSocketClientController(
ws_url=ws_url,
make_connection_failed_callback=make_connection_failed_callback,
+ additional_headers=control_headers,
)
task = self.ctrl.run(new_connection_callback)
@@ -193,6 +279,48 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
return await self.handler.ping()
+ async def require_workspace_context(self, context: TenantContext) -> ExecutionContext:
+ """Fence an HTTP/runtime caller to this connector's one Workspace.
+
+ A Plugin Runtime connection is deliberately not a cross-Workspace
+ router. Calls from another Workspace therefore look like an absent
+ resource instead of being forwarded to the bound Runtime.
+ """
+
+ workspace_uuid = require_workspace_uuid(context)
+ instance_uuid = str(getattr(context, 'instance_uuid', '') or '').strip()
+ generation = getattr(context, 'placement_generation', None)
+ if not instance_uuid or isinstance(generation, bool) or not isinstance(generation, int) or generation <= 0:
+ raise WorkspaceNotFoundError('Plugin resource not found')
+
+ binding = await self.ap.workspace_service.get_execution_binding(
+ workspace_uuid,
+ expected_generation=generation,
+ )
+ if binding.instance_uuid != instance_uuid:
+ raise WorkspaceNotFoundError('Plugin resource not found')
+
+ execution_context = ExecutionContext(
+ instance_uuid=instance_uuid,
+ workspace_uuid=workspace_uuid,
+ placement_generation=generation,
+ trigger_principal=getattr(context, 'principal', None),
+ )
+ if not self.is_enable_plugin:
+ return execution_context
+ if not hasattr(self, 'handler'):
+ raise PluginRuntimeNotConnectedError('Plugin runtime is not connected')
+
+ bound_context = self.handler.require_bound_action_context().without_installation()
+ if (
+ bound_context.instance_uuid != instance_uuid
+ or bound_context.workspace_uuid != workspace_uuid
+ or bound_context.placement_generation != generation
+ ):
+ raise WorkspaceNotFoundError('Plugin resource not found')
+
+ return execution_context
+
def _inspect_plugin_package(
self,
file_bytes: bytes,
@@ -231,6 +359,7 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
async def _install_mcp_from_marketplace(
self,
+ execution_context: ExecutionContext,
mcp_data: dict[str, Any],
task_context: taskmgr.TaskContext | None = None,
):
@@ -243,9 +372,6 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
for ``http``/``sse`` it preserves ``url``/``headers``/``timeout``/
``ssereadtimeout``.
"""
- from ..entity.persistence import mcp as persistence_mcp
- import uuid
-
mode = mcp_data.get('mode') or 'stdio'
extra_args = mcp_data.get('extra_args') or {}
# The MCP transport selection was simplified to two modes: 'stdio'
@@ -263,18 +389,12 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
# Use __ instead of / to avoid URL routing issues with slashes
name = f'{mcp_data.get("author", "")}__{mcp_data.get("name", "")}'
- # Check if MCP server already exists
- existing = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.name == name)
- )
- if existing.scalar_one_or_none():
+ existing = await self.ap.mcp_service.get_mcp_server_by_name(execution_context, name)
+ if existing is not None:
self.ap.logger.info(f'MCP server {name} already exists, skipping installation')
return
- # Create MCP server record
- server_uuid = str(uuid.uuid4())
server_data = {
- 'uuid': server_uuid,
'name': name,
'enable': True,
'mode': mode,
@@ -282,23 +402,13 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
'readme': readme,
}
- await self.ap.persistence_mgr.execute_async(sqlalchemy.insert(persistence_mcp.MCPServer).values(server_data))
-
- # Start the MCP server
- result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_mcp.MCPServer).where(persistence_mcp.MCPServer.uuid == server_uuid)
- )
- server_entity = result.first()
- if server_entity:
- server_config = self.ap.persistence_mgr.serialize_model(persistence_mcp.MCPServer, server_entity)
- if self.ap.tool_mgr.mcp_tool_loader:
- mcp_task = asyncio.create_task(self.ap.tool_mgr.mcp_tool_loader.host_mcp_server(server_config))
- self.ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks.append(mcp_task)
+ await self.ap.mcp_service.create_mcp_server(execution_context, server_data)
self.ap.logger.info(f'Installed MCP server {name} from marketplace')
async def _install_skill_from_zip(
self,
+ execution_context: ExecutionContext,
file_bytes: bytes,
filename: str,
task_context: taskmgr.TaskContext | None = None,
@@ -312,6 +422,7 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
# Install from ZIP using skill service
result = await skill_service.install_from_zip_upload(
+ execution_context,
file_bytes=file_bytes,
filename=filename + '.zip',
)
@@ -380,6 +491,12 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
plugin_name = install_info.get('plugin_name')
if install_source == PluginInstallSource.MARKETPLACE:
+ action_context = self.handler.require_bound_action_context()
+ execution_context = ExecutionContext(
+ instance_uuid=action_context.instance_uuid,
+ workspace_uuid=action_context.workspace_uuid,
+ placement_generation=action_context.placement_generation,
+ )
# Handle marketplace plugin/mcp/skill installation
plugin_author = install_info.get('plugin_author', '')
plugin_name = install_info.get('plugin_name', '')
@@ -397,7 +514,7 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
self.ap.logger.info(f'Installing MCP from marketplace: {plugin_author}/{plugin_name}')
if task_context:
task_context.set_current_action('installing mcp server')
- await self._install_mcp_from_marketplace(mcp_data, task_context)
+ await self._install_mcp_from_marketplace(execution_context, mcp_data, task_context)
# Best-effort install report (bumps marketplace install_count).
try:
await client.post(
@@ -440,7 +557,12 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
self.ap.logger.info(f'Downloaded skill ZIP ({file_size} bytes)')
# Install skill from ZIP using skill service
- await self._install_skill_from_zip(file_bytes, f'{plugin_author}-{plugin_name}', task_context)
+ await self._install_skill_from_zip(
+ execution_context,
+ file_bytes,
+ f'{plugin_author}-{plugin_name}',
+ task_context,
+ )
return
elif skill_resp.status_code == 404:
# Try plugin endpoint - get versions and download
@@ -645,6 +767,7 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
# Fetch all timestamps in a single query using OR conditions
if plugin_ids:
+ action_context = self.handler.require_bound_action_context()
conditions = [
sqlalchemy.and_(
persistence_plugin.PluginSetting.plugin_author == author,
@@ -658,7 +781,9 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
persistence_plugin.PluginSetting.plugin_author,
persistence_plugin.PluginSetting.plugin_name,
persistence_plugin.PluginSetting.created_at,
- ).where(sqlalchemy.or_(*conditions))
+ )
+ .where(persistence_plugin.PluginSetting.workspace_uuid == action_context.workspace_uuid)
+ .where(sqlalchemy.or_(*conditions))
)
for row in result:
@@ -780,13 +905,19 @@ class PluginRuntimeConnector(ManagedRuntimeConnector):
session: provider_session.Session,
query_id: int,
bound_plugins: list[str] | None = None,
+ query_uuid: str | None = None,
) -> dict[str, Any]:
if not self.is_enable_plugin:
return {'error': 'Tool not found: plugin system is disabled'}
# Pass include_plugins to runtime for validation
return await self.handler.call_tool(
- tool_name, parameters, session.model_dump(serialize_as_any=True), query_id, include_plugins=bound_plugins
+ tool_name,
+ parameters,
+ session.model_dump(serialize_as_any=True),
+ query_id,
+ query_uuid=query_uuid,
+ include_plugins=bound_plugins,
)
async def list_commands(self, bound_plugins: list[str] | None = None) -> list[ComponentManifest]:
diff --git a/src/langbot/pkg/plugin/handler.py b/src/langbot/pkg/plugin/handler.py
index dcfb006b5..86f8c59ad 100644
--- a/src/langbot/pkg/plugin/handler.py
+++ b/src/langbot/pkg/plugin/handler.py
@@ -1,14 +1,18 @@
from __future__ import annotations
+import inspect
import typing
from typing import Any
import base64
import traceback
+import uuid
+from dataclasses import dataclass
import sqlalchemy
from langbot_plugin.runtime.io import handler
from langbot_plugin.runtime.io.connection import Connection
+from langbot_plugin.entities.io.context import ActionContext
from langbot_plugin.entities.io.actions.enums import (
CommonAction,
RuntimeToLangBotAction,
@@ -19,8 +23,11 @@ import langbot_plugin.api.entities.builtin.platform.message as platform_message
import langbot_plugin.api.entities.builtin.provider.message as provider_message
import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
+from ..api.http.context import ExecutionContext
from ..entity.persistence import plugin as persistence_plugin
from ..entity.persistence import bstorage as persistence_bstorage
+from ..entity.persistence import bot as persistence_bot
+from ..entity.persistence import model as persistence_model
from ..core import app
from ..utils import constants
@@ -49,23 +56,317 @@ def _make_rag_error_response(error: Exception, error_type: str, **extra_context)
return handler.ActionResponse.error(message=message)
+@dataclass(frozen=True, slots=True)
+class _PluginInstallationIdentity:
+ workspace_uuid: str
+ plugin_author: str
+ plugin_name: str
+
+
+_UNTRUSTED_SCOPE_FIELDS = frozenset(
+ {
+ 'context',
+ 'action_context',
+ 'instance_uuid',
+ 'workspace_uuid',
+ 'placement_generation',
+ 'installation_uuid',
+ }
+)
+
+_RUNTIME_SCOPED_ACTIONS = frozenset(
+ {
+ CommonAction.PING.value,
+ CommonAction.FILE_CHUNK.value,
+ RuntimeToLangBotAction.INITIALIZE_PLUGIN_SETTINGS.value,
+ RuntimeToLangBotAction.GET_PLUGIN_SETTINGS.value,
+ }
+)
+
+
class RuntimeConnectionHandler(handler.Handler):
"""Runtime connection handler"""
ap: app.Application
+ @staticmethod
+ def derive_installation_uuid(
+ action_context: ActionContext,
+ plugin_author: str,
+ plugin_name: str,
+ ) -> str:
+ """Derive the current stable installation capability without trusting plugin data."""
+
+ return str(
+ uuid.uuid5(
+ uuid.NAMESPACE_URL,
+ 'langbot:plugin-installation:'
+ f'{action_context.instance_uuid}:{action_context.workspace_uuid}:'
+ f'{plugin_author}/{plugin_name}',
+ )
+ )
+
+ def validate_inbound_action_context(
+ self,
+ action: str,
+ action_context: ActionContext | None,
+ ) -> ActionContext | None:
+ """Keep a validated child installation capability from the request envelope."""
+
+ validated = super().validate_inbound_action_context(action, action_context)
+ if action_context is not None and action_context.installation_uuid is not None:
+ # The base handler already checked instance, Workspace and placement
+ # generation against this connector's immutable binding. The
+ # installation is resolved against trusted Host state asynchronously
+ # before dispatch.
+ return action_context
+ return validated
+
+ def _require_runtime_action_context(self) -> ActionContext:
+ action_context = self.current_action_context or self.bound_action_context
+ if action_context is None:
+ raise ValueError('Plugin Runtime action is missing a trusted Workspace context')
+
+ bound = self.require_bound_action_context()
+ if not bound.same_workspace(action_context):
+ raise ValueError('Plugin Runtime action context does not match its connector binding')
+ return action_context
+
+ def _remember_installation(
+ self,
+ action_context: ActionContext,
+ plugin_author: str,
+ plugin_name: str,
+ ) -> str:
+ installation_uuid = self.derive_installation_uuid(
+ action_context,
+ plugin_author,
+ plugin_name,
+ )
+ self._installation_bindings[installation_uuid] = _PluginInstallationIdentity(
+ workspace_uuid=action_context.workspace_uuid,
+ plugin_author=plugin_author,
+ plugin_name=plugin_name,
+ )
+ return installation_uuid
+
+ async def _resolve_installation_identity(
+ self,
+ action_context: ActionContext,
+ ) -> _PluginInstallationIdentity:
+ installation_uuid = action_context.installation_uuid
+ if installation_uuid is None:
+ raise ValueError('Plugin action is missing installation_uuid')
+
+ identity = self._installation_bindings.get(installation_uuid)
+ if identity is not None:
+ if identity.workspace_uuid != action_context.workspace_uuid:
+ raise ValueError('Plugin installation belongs to another Workspace')
+ return identity
+
+ # A control connection may reconnect while the Runtime keeps plugin
+ # processes alive. Rebuild the capability map from Workspace-scoped
+ # settings instead of accepting an installation asserted by the peer.
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(
+ persistence_plugin.PluginSetting.plugin_author,
+ persistence_plugin.PluginSetting.plugin_name,
+ ).where(persistence_plugin.PluginSetting.workspace_uuid == action_context.workspace_uuid)
+ )
+ for setting in result.all():
+ candidate_uuid = self.derive_installation_uuid(
+ action_context,
+ setting.plugin_author,
+ setting.plugin_name,
+ )
+ if candidate_uuid == installation_uuid:
+ return self._installation_bindings.setdefault(
+ installation_uuid,
+ _PluginInstallationIdentity(
+ workspace_uuid=action_context.workspace_uuid,
+ plugin_author=setting.plugin_author,
+ plugin_name=setting.plugin_name,
+ ),
+ )
+
+ raise ValueError('Plugin installation is not registered in this Workspace')
+
+ async def _require_plugin_action_context(
+ self,
+ ) -> tuple[ActionContext, _PluginInstallationIdentity]:
+ action_context = self._require_runtime_action_context()
+ identity = await self._resolve_installation_identity(action_context)
+ return action_context, identity
+
+ async def _require_active_action_context(self, action_context: ActionContext) -> None:
+ """Fence stale Runtime generations against Core's active projection."""
+
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ if workspace_service is None:
+ raise ValueError('Workspace execution service is unavailable')
+ binding = await workspace_service.get_execution_binding(
+ action_context.workspace_uuid,
+ expected_generation=action_context.placement_generation,
+ )
+ if (
+ binding.instance_uuid != action_context.instance_uuid
+ or binding.workspace_uuid != action_context.workspace_uuid
+ or binding.placement_generation != action_context.placement_generation
+ ):
+ raise ValueError('Plugin Runtime action uses a stale Workspace execution binding')
+
+ def _secure_plugin_actions(self) -> None:
+ """Wrap plugin-origin actions with installation validation and payload scrubbing."""
+
+ for action_name, action_handler in list(self.actions.items()):
+ if action_name in _RUNTIME_SCOPED_ACTIONS:
+ continue
+
+ async def secured_action(
+ data: dict[str, Any],
+ *,
+ _action_handler=action_handler,
+ ) -> handler.ActionResponse:
+ action_context, _ = await self._require_plugin_action_context()
+ await self._require_active_action_context(action_context)
+ safe_data = {key: value for key, value in data.items() if key not in _UNTRUSTED_SCOPE_FIELDS}
+ response = _action_handler(safe_data)
+ if inspect.isawaitable(response):
+ response = await response
+ return response
+
+ self.actions[action_name] = secured_action
+
+ async def _get_plugin_setting(
+ self,
+ action_context: ActionContext,
+ identity: _PluginInstallationIdentity,
+ ):
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(persistence_plugin.PluginSetting)
+ .where(persistence_plugin.PluginSetting.workspace_uuid == action_context.workspace_uuid)
+ .where(persistence_plugin.PluginSetting.plugin_author == identity.plugin_author)
+ .where(persistence_plugin.PluginSetting.plugin_name == identity.plugin_name)
+ )
+ setting = result.first()
+ if setting is None:
+ raise ValueError('Plugin installation setting was not found')
+ return setting
+
+ async def _resolve_query(
+ self,
+ data: dict[str, Any],
+ action_context: ActionContext,
+ ):
+ query_uuid = data.get('query_uuid')
+ if query_uuid is not None:
+ query = await self.ap.query_pool.get_query(
+ action_context.workspace_uuid,
+ query_uuid,
+ )
+ else:
+ query = await self.ap.query_pool.get_query_by_legacy_id(
+ action_context.workspace_uuid,
+ data['query_id'],
+ )
+
+ if query is None:
+ return None
+ if (
+ getattr(query, 'instance_uuid', None) != action_context.instance_uuid
+ or getattr(query, 'workspace_uuid', None) != action_context.workspace_uuid
+ or getattr(query, 'placement_generation', None) != action_context.placement_generation
+ ):
+ return None
+ return query
+
+ async def _resource_exists(
+ self,
+ model,
+ uuid_column,
+ resource_uuid: str,
+ workspace_uuid: str,
+ ) -> bool:
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(uuid_column)
+ .select_from(model)
+ .where(model.workspace_uuid == workspace_uuid)
+ .where(uuid_column == resource_uuid)
+ .limit(1)
+ )
+ return result.first() is not None
+
+ @staticmethod
+ def _config_contains_file_key(value: Any, file_key: str) -> bool:
+ if isinstance(value, dict):
+ return any(RuntimeConnectionHandler._config_contains_file_key(item, file_key) for item in value.values())
+ if isinstance(value, list):
+ return any(RuntimeConnectionHandler._config_contains_file_key(item, file_key) for item in value)
+ return isinstance(value, str) and value == file_key
+
+ @staticmethod
+ def _execution_context(action_context: ActionContext) -> ExecutionContext:
+ """Project the trusted wire binding into Core's runtime context."""
+
+ return ExecutionContext(
+ instance_uuid=action_context.instance_uuid,
+ workspace_uuid=action_context.workspace_uuid,
+ placement_generation=action_context.placement_generation,
+ )
+
+ @staticmethod
+ def _binary_storage_owner(
+ action_context: ActionContext,
+ identity: _PluginInstallationIdentity,
+ owner_type: str,
+ ) -> str:
+ """Resolve storage ownership exclusively from the trusted binding."""
+
+ if owner_type == 'workspace':
+ return action_context.workspace_uuid
+ if owner_type == 'plugin':
+ return f'{identity.plugin_author}/{identity.plugin_name}'
+ raise ValueError(f'Unsupported binary storage owner_type {owner_type!r}')
+
+ @classmethod
+ def _binary_storage_key(
+ cls,
+ action_context: ActionContext,
+ *,
+ owner_type: str,
+ owner: str,
+ key: str,
+ ) -> str:
+ """Use Core's canonical key across every persistent owner dimension."""
+
+ # Import lazily: StorageMgr references the Application type, whose
+ # module wires the plugin connector during startup.
+ from ..storage.mgr import StorageMgr
+
+ return StorageMgr.canonical_binary_storage_key(
+ cls._execution_context(action_context),
+ owner_type=owner_type,
+ owner=owner,
+ key=key,
+ )
+
def __init__(
self,
connection: Connection,
disconnect_callback: typing.Callable[[], typing.Coroutine[typing.Any, typing.Any, bool]],
ap: app.Application,
+ action_context: ActionContext,
):
super().__init__(connection, disconnect_callback)
self.ap = ap
+ self.bind_action_context(action_context.without_installation())
+ self._installation_bindings: dict[str, _PluginInstallationIdentity] = {}
@self.action(RuntimeToLangBotAction.INITIALIZE_PLUGIN_SETTINGS)
async def initialize_plugin_settings(data: dict[str, Any]) -> handler.ActionResponse:
"""Initialize plugin settings"""
+ action_context = self._require_runtime_action_context()
+ await self._require_active_action_context(action_context)
# check if exists plugin setting
plugin_author = data['plugin_author']
plugin_name = data['plugin_name']
@@ -75,6 +376,7 @@ class RuntimeConnectionHandler(handler.Handler):
try:
result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_plugin.PluginSetting)
+ .where(persistence_plugin.PluginSetting.workspace_uuid == action_context.workspace_uuid)
.where(persistence_plugin.PluginSetting.plugin_author == plugin_author)
.where(persistence_plugin.PluginSetting.plugin_name == plugin_name)
)
@@ -85,6 +387,7 @@ class RuntimeConnectionHandler(handler.Handler):
# delete plugin setting
await self.ap.persistence_mgr.execute_async(
sqlalchemy.delete(persistence_plugin.PluginSetting)
+ .where(persistence_plugin.PluginSetting.workspace_uuid == action_context.workspace_uuid)
.where(persistence_plugin.PluginSetting.plugin_author == plugin_author)
.where(persistence_plugin.PluginSetting.plugin_name == plugin_name)
)
@@ -92,6 +395,7 @@ class RuntimeConnectionHandler(handler.Handler):
# create plugin setting
await self.ap.persistence_mgr.execute_async(
sqlalchemy.insert(persistence_plugin.PluginSetting).values(
+ workspace_uuid=action_context.workspace_uuid,
plugin_author=plugin_author,
plugin_name=plugin_name,
install_source=install_source,
@@ -116,11 +420,14 @@ class RuntimeConnectionHandler(handler.Handler):
async def get_plugin_settings(data: dict[str, Any]) -> handler.ActionResponse:
"""Get plugin settings"""
+ action_context = self._require_runtime_action_context()
+ await self._require_active_action_context(action_context)
plugin_author = data['plugin_author']
plugin_name = data['plugin_name']
result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_plugin.PluginSetting)
+ .where(persistence_plugin.PluginSetting.workspace_uuid == action_context.workspace_uuid)
.where(persistence_plugin.PluginSetting.plugin_author == plugin_author)
.where(persistence_plugin.PluginSetting.plugin_name == plugin_name)
)
@@ -131,6 +438,11 @@ class RuntimeConnectionHandler(handler.Handler):
'plugin_config': {},
'install_source': 'local',
'install_info': {},
+ 'installation_uuid': self._remember_installation(
+ action_context,
+ plugin_author,
+ plugin_name,
+ ),
}
setting = result.first()
@@ -149,17 +461,17 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.REPLY_MESSAGE)
async def reply_message(data: dict[str, Any]) -> handler.ActionResponse:
"""Reply message"""
+ action_context, _ = await self._require_plugin_action_context()
query_id = data['query_id']
message_chain = data['message_chain']
quote_origin = data['quote_origin']
- if query_id not in self.ap.query_pool.cached_queries:
+ query = await self._resolve_query(data, action_context)
+ if query is None:
return handler.ActionResponse.error(
message=f'Query with query_id {query_id} not found',
)
- query = self.ap.query_pool.cached_queries[query_id]
-
message_chain_obj = platform_message.MessageChain.model_validate(message_chain)
self.ap.logger.debug(f'Reply message: {message_chain_obj.model_dump(serialize_as_any=False)}')
@@ -177,14 +489,14 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.GET_BOT_UUID)
async def get_bot_uuid(data: dict[str, Any]) -> handler.ActionResponse:
"""Get bot uuid"""
+ action_context, _ = await self._require_plugin_action_context()
query_id = data['query_id']
- if query_id not in self.ap.query_pool.cached_queries:
+ query = await self._resolve_query(data, action_context)
+ if query is None:
return handler.ActionResponse.error(
message=f'Query with query_id {query_id} not found',
)
- query = self.ap.query_pool.cached_queries[query_id]
-
return handler.ActionResponse.success(
data={
'bot_uuid': query.bot_uuid,
@@ -194,17 +506,17 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.SET_QUERY_VAR)
async def set_query_var(data: dict[str, Any]) -> handler.ActionResponse:
"""Set query var"""
+ action_context, _ = await self._require_plugin_action_context()
query_id = data['query_id']
key = data['key']
value = data['value']
- if query_id not in self.ap.query_pool.cached_queries:
+ query = await self._resolve_query(data, action_context)
+ if query is None:
return handler.ActionResponse.error(
message=f'Query with query_id {query_id} not found',
)
- query = self.ap.query_pool.cached_queries[query_id]
-
query.variables[key] = value
return handler.ActionResponse.success(
@@ -214,16 +526,16 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.GET_QUERY_VAR)
async def get_query_var(data: dict[str, Any]) -> handler.ActionResponse:
"""Get query var"""
+ action_context, _ = await self._require_plugin_action_context()
query_id = data['query_id']
key = data['key']
- if query_id not in self.ap.query_pool.cached_queries:
+ query = await self._resolve_query(data, action_context)
+ if query is None:
return handler.ActionResponse.error(
message=f'Query with query_id {query_id} not found',
)
- query = self.ap.query_pool.cached_queries[query_id]
-
return handler.ActionResponse.success(
data={
'value': query.variables[key],
@@ -233,14 +545,14 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.GET_QUERY_VARS)
async def get_query_vars(data: dict[str, Any]) -> handler.ActionResponse:
"""Get query vars"""
+ action_context, _ = await self._require_plugin_action_context()
query_id = data['query_id']
- if query_id not in self.ap.query_pool.cached_queries:
+ query = await self._resolve_query(data, action_context)
+ if query is None:
return handler.ActionResponse.error(
message=f'Query with query_id {query_id} not found',
)
- query = self.ap.query_pool.cached_queries[query_id]
-
return handler.ActionResponse.success(
data={
'vars': query.variables,
@@ -250,14 +562,14 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.CREATE_NEW_CONVERSATION)
async def create_new_conversation(data: dict[str, Any]) -> handler.ActionResponse:
"""Create new conversation"""
+ action_context, _ = await self._require_plugin_action_context()
query_id = data['query_id']
- if query_id not in self.ap.query_pool.cached_queries:
+ query = await self._resolve_query(data, action_context)
+ if query is None:
return handler.ActionResponse.error(
message=f'Query with query_id {query_id} not found',
)
- query = self.ap.query_pool.cached_queries[query_id]
-
query.session.using_conversation = None
return handler.ActionResponse.success(
@@ -276,7 +588,20 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.GET_BOTS)
async def get_bots(data: dict[str, Any]) -> handler.ActionResponse:
"""Get bots"""
- bots = await self.ap.bot_service.get_bots(include_secret=False)
+ action_context, _ = await self._require_plugin_action_context()
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(persistence_bot.Bot).where(
+ persistence_bot.Bot.workspace_uuid == action_context.workspace_uuid
+ )
+ )
+ bots = [
+ self.ap.persistence_mgr.serialize_model(
+ persistence_bot.Bot,
+ bot,
+ ['adapter_config'],
+ )
+ for bot in result.all()
+ ]
return handler.ActionResponse.success(
data={
'bots': bots,
@@ -286,8 +611,23 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.GET_BOT_INFO)
async def get_bot_info(data: dict[str, Any]) -> handler.ActionResponse:
"""Get bot info"""
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
bot_uuid = data['bot_uuid']
- bot = await self.ap.bot_service.get_runtime_bot_info(bot_uuid, include_secret=False)
+ if not await self._resource_exists(
+ persistence_bot.Bot,
+ persistence_bot.Bot.uuid,
+ bot_uuid,
+ action_context.workspace_uuid,
+ ):
+ return handler.ActionResponse.error(
+ message=f'Bot with bot_uuid {bot_uuid} not found',
+ )
+ bot = await self.ap.bot_service.get_runtime_bot_info(
+ execution_context,
+ bot_uuid,
+ include_secret=False,
+ )
return handler.ActionResponse.success(
data={
'bot': bot,
@@ -297,15 +637,30 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.SEND_MESSAGE)
async def send_message(data: dict[str, Any]) -> handler.ActionResponse:
"""Send message"""
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
bot_uuid = data['bot_uuid']
target_type = data['target_type']
target_id = data['target_id']
message_chain = data['message_chain']
+ if not await self._resource_exists(
+ persistence_bot.Bot,
+ persistence_bot.Bot.uuid,
+ bot_uuid,
+ action_context.workspace_uuid,
+ ):
+ return handler.ActionResponse.error(
+ message=f'Bot with bot_uuid {bot_uuid} not found',
+ )
+
# Use custom deserializer that properly handles Forward messages
message_chain_obj = platform_message.MessageChain.model_validate(message_chain)
- bot = await self.ap.platform_mgr.get_bot_by_uuid(bot_uuid)
+ bot = await self.ap.platform_mgr.get_bot_by_uuid(
+ execution_context,
+ bot_uuid,
+ )
if bot is None:
return handler.ActionResponse.error(
message=f'Bot with bot_uuid {bot_uuid} not found',
@@ -324,23 +679,52 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.GET_LLM_MODELS)
async def get_llm_models(data: dict[str, Any]) -> handler.ActionResponse:
"""Get llm models, returns list of UUID strings"""
- llm_models = await self.ap.llm_model_service.get_llm_models(include_secret=False)
+ action_context, _ = await self._require_plugin_action_context()
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(persistence_model.LLMModel.uuid).where(
+ persistence_model.LLMModel.workspace_uuid == action_context.workspace_uuid
+ )
+ )
return handler.ActionResponse.success(
data={
- 'llm_models': [m['uuid'] for m in llm_models],
+ 'llm_models': list(result.scalars().all()),
},
)
@self.action(PluginToRuntimeAction.INVOKE_LLM)
async def invoke_llm(data: dict[str, Any]) -> handler.ActionResponse:
"""Invoke llm"""
+ action_context, _ = await self._require_plugin_action_context()
llm_model_uuid = data['llm_model_uuid']
messages = data['messages']
funcs = data.get('funcs', [])
extra_args = data.get('extra_args', {})
- llm_model = await self.ap.model_mgr.get_model_by_uuid(llm_model_uuid)
- if llm_model is None:
+ if not await self._resource_exists(
+ persistence_model.LLMModel,
+ persistence_model.LLMModel.uuid,
+ llm_model_uuid,
+ action_context.workspace_uuid,
+ ):
+ return handler.ActionResponse.error(
+ message=f'LLM model with llm_model_uuid {llm_model_uuid} not found',
+ )
+ try:
+ execution_context = self._execution_context(action_context)
+ llm_model = await self.ap.model_mgr.get_model_by_uuid(
+ execution_context,
+ llm_model_uuid,
+ )
+ except ValueError:
+ return handler.ActionResponse.error(
+ message=f'LLM model with llm_model_uuid {llm_model_uuid} not found',
+ )
+ runtime_workspace_uuid = getattr(
+ getattr(llm_model, 'model_entity', None),
+ 'workspace_uuid',
+ None,
+ )
+ if runtime_workspace_uuid not in (None, action_context.workspace_uuid):
return handler.ActionResponse.error(
message=f'LLM model with llm_model_uuid {llm_model_uuid} not found',
)
@@ -361,6 +745,7 @@ class RuntimeConnectionHandler(handler.Handler):
messages=messages_obj,
funcs=funcs_obj,
extra_args=extra_args,
+ execution_context=execution_context,
)
return handler.ActionResponse.success(
@@ -372,9 +757,21 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(RuntimeToLangBotAction.SET_BINARY_STORAGE)
async def set_binary_storage(data: dict[str, Any]) -> handler.ActionResponse:
"""Set binary storage"""
+ action_context, identity = await self._require_plugin_action_context()
key = data['key']
owner_type = data['owner_type']
- owner = data['owner']
+ try:
+ owner = self._binary_storage_owner(action_context, identity, owner_type)
+ unique_key = self._binary_storage_key(
+ action_context,
+ owner_type=owner_type,
+ owner=owner,
+ key=key,
+ )
+ except ValueError as e:
+ return handler.ActionResponse.error(
+ message=str(e),
+ )
value = base64.b64decode(data['value_base64'])
max_value_bytes = (
self.ap.instance_config.data.get('plugin', {})
@@ -395,23 +792,22 @@ class RuntimeConnectionHandler(handler.Handler):
result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_bstorage.BinaryStorage)
- .where(persistence_bstorage.BinaryStorage.key == key)
- .where(persistence_bstorage.BinaryStorage.owner_type == owner_type)
- .where(persistence_bstorage.BinaryStorage.owner == owner)
+ .where(persistence_bstorage.BinaryStorage.workspace_uuid == action_context.workspace_uuid)
+ .where(persistence_bstorage.BinaryStorage.unique_key == unique_key)
)
if result.first() is not None:
await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(persistence_bstorage.BinaryStorage)
- .where(persistence_bstorage.BinaryStorage.key == key)
- .where(persistence_bstorage.BinaryStorage.owner_type == owner_type)
- .where(persistence_bstorage.BinaryStorage.owner == owner)
+ .where(persistence_bstorage.BinaryStorage.workspace_uuid == action_context.workspace_uuid)
+ .where(persistence_bstorage.BinaryStorage.unique_key == unique_key)
.values(value=value)
)
else:
await self.ap.persistence_mgr.execute_async(
sqlalchemy.insert(persistence_bstorage.BinaryStorage).values(
- unique_key=f'{owner_type}:{owner}:{key}',
+ workspace_uuid=action_context.workspace_uuid,
+ unique_key=unique_key,
key=key,
owner_type=owner_type,
owner=owner,
@@ -426,15 +822,26 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(RuntimeToLangBotAction.GET_BINARY_STORAGE)
async def get_binary_storage(data: dict[str, Any]) -> handler.ActionResponse:
"""Get binary storage"""
+ action_context, identity = await self._require_plugin_action_context()
key = data['key']
owner_type = data['owner_type']
- owner = data['owner']
+ try:
+ owner = self._binary_storage_owner(action_context, identity, owner_type)
+ unique_key = self._binary_storage_key(
+ action_context,
+ owner_type=owner_type,
+ owner=owner,
+ key=key,
+ )
+ except ValueError as e:
+ return handler.ActionResponse.error(
+ message=str(e),
+ )
result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_bstorage.BinaryStorage)
- .where(persistence_bstorage.BinaryStorage.key == key)
- .where(persistence_bstorage.BinaryStorage.owner_type == owner_type)
- .where(persistence_bstorage.BinaryStorage.owner == owner)
+ .where(persistence_bstorage.BinaryStorage.workspace_uuid == action_context.workspace_uuid)
+ .where(persistence_bstorage.BinaryStorage.unique_key == unique_key)
)
storage = result.first()
@@ -452,15 +859,26 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(RuntimeToLangBotAction.DELETE_BINARY_STORAGE)
async def delete_binary_storage(data: dict[str, Any]) -> handler.ActionResponse:
"""Delete binary storage"""
+ action_context, identity = await self._require_plugin_action_context()
key = data['key']
owner_type = data['owner_type']
- owner = data['owner']
+ try:
+ owner = self._binary_storage_owner(action_context, identity, owner_type)
+ unique_key = self._binary_storage_key(
+ action_context,
+ owner_type=owner_type,
+ owner=owner,
+ key=key,
+ )
+ except ValueError as e:
+ return handler.ActionResponse.error(
+ message=str(e),
+ )
await self.ap.persistence_mgr.execute_async(
sqlalchemy.delete(persistence_bstorage.BinaryStorage)
- .where(persistence_bstorage.BinaryStorage.key == key)
- .where(persistence_bstorage.BinaryStorage.owner_type == owner_type)
- .where(persistence_bstorage.BinaryStorage.owner == owner)
+ .where(persistence_bstorage.BinaryStorage.workspace_uuid == action_context.workspace_uuid)
+ .where(persistence_bstorage.BinaryStorage.unique_key == unique_key)
)
return handler.ActionResponse.success(
@@ -470,11 +888,18 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(RuntimeToLangBotAction.GET_BINARY_STORAGE_KEYS)
async def get_binary_storage_keys(data: dict[str, Any]) -> handler.ActionResponse:
"""Get binary storage keys"""
+ action_context, identity = await self._require_plugin_action_context()
owner_type = data['owner_type']
- owner = data['owner']
+ try:
+ owner = self._binary_storage_owner(action_context, identity, owner_type)
+ except ValueError as e:
+ return handler.ActionResponse.error(
+ message=str(e),
+ )
result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_bstorage.BinaryStorage.key)
+ .where(persistence_bstorage.BinaryStorage.workspace_uuid == action_context.workspace_uuid)
.where(persistence_bstorage.BinaryStorage.owner_type == owner_type)
.where(persistence_bstorage.BinaryStorage.owner == owner)
)
@@ -488,11 +913,22 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.GET_CONFIG_FILE)
async def get_config_file(data: dict[str, Any]) -> handler.ActionResponse:
"""Get a config file by file key"""
+ action_context, identity = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
file_key = data['file_key']
try:
- # Load file from storage
- file_bytes = await self.ap.storage_mgr.storage_provider.load(file_key)
+ setting = await self._get_plugin_setting(action_context, identity)
+ if not self._config_contains_file_key(setting.config, file_key):
+ raise ValueError('Config file does not belong to this plugin installation')
+ # The persisted config is user-controlled and therefore cannot
+ # turn an arbitrary opaque object key into authority. Validate
+ # every trusted scope dimension before touching the provider.
+ file_bytes = await self.ap.storage_mgr.load_scoped_object_key(
+ execution_context,
+ file_key,
+ expected_owner_type='plugin_config',
+ )
return handler.ActionResponse.success(
data={
@@ -508,35 +944,86 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.INVOKE_EMBEDDING)
async def invoke_embedding(data: dict[str, Any]) -> handler.ActionResponse:
+ action_context, _ = await self._require_plugin_action_context()
embedding_model_uuid = data['embedding_model_uuid']
texts = data['texts']
- embedding_model = await self.ap.model_mgr.get_embedding_model_by_uuid(embedding_model_uuid)
- if embedding_model is None:
+ if not await self._resource_exists(
+ persistence_model.EmbeddingModel,
+ persistence_model.EmbeddingModel.uuid,
+ embedding_model_uuid,
+ action_context.workspace_uuid,
+ ):
+ return handler.ActionResponse.error(
+ message=f'Embedding model with embedding_model_uuid {embedding_model_uuid} not found',
+ )
+ try:
+ execution_context = self._execution_context(action_context)
+ embedding_model = await self.ap.model_mgr.get_embedding_model_by_uuid(
+ execution_context,
+ embedding_model_uuid,
+ )
+ except ValueError:
+ return handler.ActionResponse.error(
+ message=f'Embedding model with embedding_model_uuid {embedding_model_uuid} not found',
+ )
+ runtime_workspace_uuid = getattr(
+ getattr(embedding_model, 'model_entity', None),
+ 'workspace_uuid',
+ None,
+ )
+ if runtime_workspace_uuid not in (None, action_context.workspace_uuid):
return handler.ActionResponse.error(
message=f'Embedding model with embedding_model_uuid {embedding_model_uuid} not found',
)
try:
- vectors = await embedding_model.provider.invoke_embedding(embedding_model, texts)
+ vectors = await embedding_model.provider.invoke_embedding(
+ embedding_model,
+ texts,
+ execution_context=execution_context,
+ )
return handler.ActionResponse.success(data={'vectors': vectors})
except Exception as e:
return _make_rag_error_response(e, 'EmbeddingError', embedding_model_uuid=embedding_model_uuid)
@self.action(PluginToRuntimeAction.INVOKE_RERANK)
async def invoke_rerank(data: dict[str, Any]) -> handler.ActionResponse:
+ action_context, _ = await self._require_plugin_action_context()
rerank_model_uuid = data['rerank_model_uuid']
query = data['query']
documents = data['documents']
top_k = data.get('top_k')
extra_args = data.get('extra_args', {})
+ if not await self._resource_exists(
+ persistence_model.RerankModel,
+ persistence_model.RerankModel.uuid,
+ rerank_model_uuid,
+ action_context.workspace_uuid,
+ ):
+ return handler.ActionResponse.error(
+ message=f'Rerank model with rerank_model_uuid {rerank_model_uuid} not found',
+ )
try:
- rerank_model = await self.ap.model_mgr.get_rerank_model_by_uuid(rerank_model_uuid)
+ execution_context = self._execution_context(action_context)
+ rerank_model = await self.ap.model_mgr.get_rerank_model_by_uuid(
+ execution_context,
+ rerank_model_uuid,
+ )
except ValueError:
return handler.ActionResponse.error(
message=f'Rerank model with rerank_model_uuid {rerank_model_uuid} not found',
)
+ runtime_workspace_uuid = getattr(
+ getattr(rerank_model, 'model_entity', None),
+ 'workspace_uuid',
+ None,
+ )
+ if runtime_workspace_uuid not in (None, action_context.workspace_uuid):
+ return handler.ActionResponse.error(
+ message=f'Rerank model with rerank_model_uuid {rerank_model_uuid} not found',
+ )
try:
scores = await rerank_model.provider.invoke_rerank(
@@ -544,6 +1031,7 @@ class RuntimeConnectionHandler(handler.Handler):
query=query,
documents=documents[:64],
extra_args=extra_args,
+ execution_context=execution_context,
)
scored = sorted(scores, key=lambda x: x.get('relevance_score', 0), reverse=True)
if top_k is not None:
@@ -554,6 +1042,8 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.VECTOR_UPSERT)
async def vector_upsert(data: dict[str, Any]) -> handler.ActionResponse:
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
collection_id = data['collection_id']
vectors = data['vectors']
ids = data['ids']
@@ -567,6 +1057,7 @@ class RuntimeConnectionHandler(handler.Handler):
return handler.ActionResponse.error(message='documents must match vectors length')
try:
await self.ap.rag_runtime_service.vector_upsert(
+ execution_context,
collection_id,
vectors,
ids,
@@ -575,10 +1066,16 @@ class RuntimeConnectionHandler(handler.Handler):
)
return handler.ActionResponse.success(data={})
except Exception as e:
- return _make_rag_error_response(e, 'VectorStoreError', collection_id=collection_id)
+ return _make_rag_error_response(
+ e,
+ 'VectorStoreError',
+ collection_id=collection_id,
+ )
@self.action(PluginToRuntimeAction.VECTOR_SEARCH)
async def vector_search(data: dict[str, Any]) -> handler.ActionResponse:
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
collection_id = data['collection_id']
query_vector = data['query_vector']
top_k = data['top_k']
@@ -588,6 +1085,7 @@ class RuntimeConnectionHandler(handler.Handler):
vector_weight = data.get('vector_weight')
try:
results = await self.ap.rag_runtime_service.vector_search(
+ execution_context,
collection_id,
query_vector,
top_k,
@@ -598,36 +1096,68 @@ class RuntimeConnectionHandler(handler.Handler):
)
return handler.ActionResponse.success(data={'results': results})
except Exception as e:
- return _make_rag_error_response(e, 'VectorStoreError', collection_id=collection_id)
+ return _make_rag_error_response(
+ e,
+ 'VectorStoreError',
+ collection_id=collection_id,
+ )
@self.action(PluginToRuntimeAction.VECTOR_DELETE)
async def vector_delete(data: dict[str, Any]) -> handler.ActionResponse:
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
collection_id = data['collection_id']
file_ids = data.get('file_ids')
filters = data.get('filters')
try:
- count = await self.ap.rag_runtime_service.vector_delete(collection_id, file_ids, filters)
+ count = await self.ap.rag_runtime_service.vector_delete(
+ execution_context,
+ collection_id,
+ file_ids,
+ filters,
+ )
return handler.ActionResponse.success(data={'count': count})
except Exception as e:
- return _make_rag_error_response(e, 'VectorStoreError', collection_id=collection_id)
+ return _make_rag_error_response(
+ e,
+ 'VectorStoreError',
+ collection_id=collection_id,
+ )
@self.action(PluginToRuntimeAction.VECTOR_LIST)
async def vector_list(data: dict[str, Any]) -> handler.ActionResponse:
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
collection_id = data['collection_id']
filters = data.get('filters')
limit = data.get('limit', 20)
offset = data.get('offset', 0)
try:
- items, total = await self.ap.rag_runtime_service.vector_list(collection_id, filters, limit, offset)
+ items, total = await self.ap.rag_runtime_service.vector_list(
+ execution_context,
+ collection_id,
+ filters,
+ limit,
+ offset,
+ )
return handler.ActionResponse.success(data={'items': items, 'total': total})
except Exception as e:
- return _make_rag_error_response(e, 'VectorStoreError', collection_id=collection_id)
+ return _make_rag_error_response(
+ e,
+ 'VectorStoreError',
+ collection_id=collection_id,
+ )
@self.action(PluginToRuntimeAction.GET_KNOWLEDEGE_FILE_STREAM)
async def get_knowledge_file_stream(data: dict[str, Any]) -> handler.ActionResponse:
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
storage_path = data['storage_path']
try:
- content_bytes = await self.ap.rag_runtime_service.get_file_stream(storage_path)
+ content_bytes = await self.ap.rag_runtime_service.get_file_stream(
+ execution_context,
+ storage_path,
+ )
file_key = await self.send_file(content_bytes, '')
return handler.ActionResponse.success(data={'file_key': file_key})
except Exception as e:
@@ -636,9 +1166,14 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.LIST_PARSERS)
async def list_parsers(data: dict[str, Any]) -> handler.ActionResponse:
"""Plugin requests host to list available parser plugins."""
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
mime_type = data.get('mime_type')
try:
- parsers = await self.ap.knowledge_service.list_parsers(mime_type)
+ parsers = await self.ap.knowledge_service.list_parsers(
+ execution_context,
+ mime_type,
+ )
return handler.ActionResponse.success(data={'parsers': parsers})
except Exception as e:
return _make_rag_error_response(e, 'ParserDiscoveryError', mime_type=mime_type)
@@ -646,6 +1181,8 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.INVOKE_PARSER)
async def invoke_parser(data: dict[str, Any]) -> handler.ActionResponse:
"""Plugin requests host to invoke a parser plugin."""
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
plugin_author = data['plugin_author']
plugin_name = data['plugin_name']
storage_path = data['storage_path']
@@ -654,12 +1191,16 @@ class RuntimeConnectionHandler(handler.Handler):
metadata = data.get('metadata', {})
try:
# Read file from storage
- file_bytes = await self.ap.rag_runtime_service.get_file_stream(storage_path)
+ file_bytes = await self.ap.rag_runtime_service.get_file_stream(
+ execution_context,
+ storage_path,
+ )
context_data = {
'mime_type': mime_type,
'filename': filename,
'metadata': metadata,
}
+ await self.ap.plugin_connector.require_workspace_context(execution_context)
result = await self.ap.plugin_connector.call_parser(
f'{plugin_author}/{plugin_name}', context_data, file_bytes
)
@@ -671,27 +1212,34 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.LIST_KNOWLEDGE_BASES)
async def list_knowledge_bases(data: dict[str, Any]) -> handler.ActionResponse:
- """List all knowledge bases available in the LangBot instance (unrestricted)."""
- knowledge_bases = []
- for kb_uuid, kb in self.ap.rag_mgr.knowledge_bases.items():
- knowledge_bases.append(
- {
- 'uuid': kb.get_uuid(),
- 'name': kb.get_name(),
- 'description': kb.knowledge_base_entity.description or '',
- }
- )
+ """List knowledge bases visible to the bound Workspace."""
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
+ details = await self.ap.rag_mgr.get_all_knowledge_base_details(execution_context)
+ knowledge_bases = [
+ {
+ 'uuid': kb['uuid'],
+ 'name': kb['name'],
+ 'description': kb.get('description') or '',
+ }
+ for kb in details
+ ]
return handler.ActionResponse.success(data={'knowledge_bases': knowledge_bases})
@self.action(PluginToRuntimeAction.RETRIEVE_KNOWLEDGE)
async def retrieve_knowledge(data: dict[str, Any]) -> handler.ActionResponse:
- """Retrieve documents from any knowledge base (unrestricted)."""
+ """Retrieve documents from a knowledge base in the bound Workspace."""
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
kb_id = data['kb_id']
query_text = data['query_text']
top_k = data.get('top_k', 5)
filters = data.get('filters', {})
- kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(kb_id)
+ kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(
+ execution_context,
+ kb_id,
+ )
if not kb:
return handler.ActionResponse.error(
message=f'Knowledge base {kb_id} not found',
@@ -699,6 +1247,7 @@ class RuntimeConnectionHandler(handler.Handler):
try:
entries = await kb.retrieve(
+ execution_context,
query_text,
settings={
'top_k': top_k,
@@ -713,15 +1262,16 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.LIST_PIPELINE_KNOWLEDGE_BASES)
async def list_pipeline_knowledge_bases(data: dict[str, Any]) -> handler.ActionResponse:
"""List knowledge bases configured for the current query's pipeline."""
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
query_id = data['query_id']
- if query_id not in self.ap.query_pool.cached_queries:
+ query = await self._resolve_query(data, action_context)
+ if query is None:
return handler.ActionResponse.error(
message=f'Query with query_id {query_id} not found',
)
- query = self.ap.query_pool.cached_queries[query_id]
-
kb_uuids = []
if query.pipeline_config:
local_agent_config = query.pipeline_config.get('ai', {}).get('local-agent', {})
@@ -734,7 +1284,10 @@ class RuntimeConnectionHandler(handler.Handler):
knowledge_bases = []
for kb_uuid in kb_uuids:
- kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(kb_uuid)
+ kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(
+ execution_context,
+ kb_uuid,
+ )
if kb:
knowledge_bases.append(
{
@@ -749,19 +1302,20 @@ class RuntimeConnectionHandler(handler.Handler):
@self.action(PluginToRuntimeAction.RETRIEVE_KNOWLEDGE_BASE)
async def retrieve_knowledge_base(data: dict[str, Any]) -> handler.ActionResponse:
"""Retrieve documents from a knowledge base within the pipeline's scope."""
+ action_context, _ = await self._require_plugin_action_context()
+ execution_context = self._execution_context(action_context)
query_id = data['query_id']
kb_id = data['kb_id']
query_text = data['query_text']
top_k = data.get('top_k', 5)
filters = data.get('filters', {})
- if query_id not in self.ap.query_pool.cached_queries:
+ query = await self._resolve_query(data, action_context)
+ if query is None:
return handler.ActionResponse.error(
message=f'Query with query_id {query_id} not found',
)
- query = self.ap.query_pool.cached_queries[query_id]
-
# Validate kb_id is in pipeline's allowed list
allowed_kb_uuids = []
if query.pipeline_config:
@@ -777,7 +1331,10 @@ class RuntimeConnectionHandler(handler.Handler):
message=f'Knowledge base {kb_id} is not configured for this pipeline',
)
- kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(kb_id)
+ kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(
+ execution_context,
+ kb_id,
+ )
if not kb:
return handler.ActionResponse.error(
message=f'Knowledge base {kb_id} not found',
@@ -786,6 +1343,7 @@ class RuntimeConnectionHandler(handler.Handler):
try:
session_name = f'{query.session.launcher_type.value}_{query.session.launcher_id}'
entries = await kb.retrieve(
+ execution_context,
query_text,
settings={
'top_k': top_k,
@@ -809,6 +1367,8 @@ class RuntimeConnectionHandler(handler.Handler):
},
)
+ self._secure_plugin_actions()
+
async def ping(self) -> dict[str, Any]:
"""Ping the runtime"""
return await self.call_action(
@@ -817,13 +1377,14 @@ class RuntimeConnectionHandler(handler.Handler):
timeout=10,
)
- async def set_runtime_config(self, cloud_service_url: str) -> dict[str, Any]:
+ async def set_runtime_config(self, cloud_service_url: str | None) -> dict[str, Any]:
"""Push runtime configuration (e.g. marketplace URL) to the runtime."""
+ data = {}
+ if cloud_service_url:
+ data['cloud_service_url'] = cloud_service_url
return await self.call_action(
LangBotToRuntimeAction.SET_RUNTIME_CONFIG,
- {
- 'cloud_service_url': cloud_service_url,
- },
+ data,
timeout=10,
)
@@ -894,9 +1455,11 @@ class RuntimeConnectionHandler(handler.Handler):
async def set_plugin_config(self, plugin_author: str, plugin_name: str, config: dict[str, Any]) -> dict[str, Any]:
"""Set plugin config"""
+ action_context = self.require_bound_action_context()
# update plugin setting
await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(persistence_plugin.PluginSetting)
+ .where(persistence_plugin.PluginSetting.workspace_uuid == action_context.workspace_uuid)
.where(persistence_plugin.PluginSetting.plugin_author == plugin_author)
.where(persistence_plugin.PluginSetting.plugin_name == plugin_name)
.values(config=config)
@@ -1079,9 +1642,16 @@ class RuntimeConnectionHandler(handler.Handler):
async def cleanup_plugin_data(self, plugin_author: str, plugin_name: str) -> None:
"""Cleanup plugin settings and binary storage"""
+ action_context = self.require_bound_action_context()
+ installation_uuid = self.derive_installation_uuid(
+ action_context,
+ plugin_author,
+ plugin_name,
+ )
# Delete plugin settings
await self.ap.persistence_mgr.execute_async(
sqlalchemy.delete(persistence_plugin.PluginSetting)
+ .where(persistence_plugin.PluginSetting.workspace_uuid == action_context.workspace_uuid)
.where(persistence_plugin.PluginSetting.plugin_author == plugin_author)
.where(persistence_plugin.PluginSetting.plugin_name == plugin_name)
)
@@ -1090,9 +1660,13 @@ class RuntimeConnectionHandler(handler.Handler):
owner = f'{plugin_author}/{plugin_name}'
await self.ap.persistence_mgr.execute_async(
sqlalchemy.delete(persistence_bstorage.BinaryStorage)
+ .where(persistence_bstorage.BinaryStorage.workspace_uuid == action_context.workspace_uuid)
.where(persistence_bstorage.BinaryStorage.owner_type == 'plugin')
.where(persistence_bstorage.BinaryStorage.owner == owner)
)
+ installation_bindings = getattr(self, '_installation_bindings', None)
+ if installation_bindings is not None:
+ installation_bindings.pop(installation_uuid, None)
async def call_tool(
self,
@@ -1101,15 +1675,19 @@ class RuntimeConnectionHandler(handler.Handler):
session: dict[str, Any],
query_id: int,
include_plugins: list[str] | None = None,
+ query_uuid: str | None = None,
) -> dict[str, Any]:
"""Call tool"""
+ query_ref: dict[str, Any] = {'query_id': query_id}
+ if query_uuid is not None:
+ query_ref['query_uuid'] = query_uuid
result = await self.call_action(
LangBotToRuntimeAction.CALL_TOOL,
{
'tool_name': tool_name,
'tool_parameters': parameters,
'session': session,
- 'query_id': query_id,
+ **query_ref,
'include_plugins': include_plugins,
},
timeout=180,
diff --git a/src/langbot/pkg/provider/modelmgr/modelmgr.py b/src/langbot/pkg/provider/modelmgr/modelmgr.py
index e3e20e026..94f400143 100644
--- a/src/langbot/pkg/provider/modelmgr/modelmgr.py
+++ b/src/langbot/pkg/provider/modelmgr/modelmgr.py
@@ -1,40 +1,55 @@
from __future__ import annotations
import asyncio
-import sqlalchemy
import traceback
+from typing import TypeVar
-from . import requester
+import sqlalchemy
+
+from ...api.http.context import (
+ ExecutionContext,
+ PrincipalContext,
+ PrincipalType,
+ RequestContext,
+)
+from ...api.http.service.tenant import TenantContext, require_workspace_uuid
from ...core import app
from ...discover import engine
-from . import token
-from ...entity.persistence import model as persistence_model
from ...entity.errors import provider as provider_errors
+from ...entity.persistence import model as persistence_model
+from ...workspace.entities import WorkspaceExecutionBinding
+from ...workspace.errors import WorkspaceError, WorkspaceInvariantError
+from . import requester, token
+
+
+_CacheKey = tuple[str, str, int, str]
+_ModelEntity = TypeVar(
+ '_ModelEntity',
+ persistence_model.LLMModel,
+ persistence_model.EmbeddingModel,
+ persistence_model.RerankModel,
+)
class ModelManager:
- """Model manager"""
+ """Workspace-scoped runtime provider and model cache."""
ap: app.Application
- provider_dict: dict[str, requester.RuntimeProvider]
- """运行时模型提供商字典, uuid -> RuntimeProvider"""
-
- llm_models: list[requester.RuntimeLLMModel]
-
- embedding_models: list[requester.RuntimeEmbeddingModel]
-
- rerank_models: list[requester.RuntimeRerankModel]
+ provider_dict: dict[_CacheKey, requester.RuntimeProvider]
+ llm_model_dict: dict[_CacheKey, requester.RuntimeLLMModel]
+ embedding_model_dict: dict[_CacheKey, requester.RuntimeEmbeddingModel]
+ rerank_model_dict: dict[_CacheKey, requester.RuntimeRerankModel]
requester_components: list[engine.Component]
-
requester_dict: dict[str, type[requester.ProviderAPIRequester]]
def __init__(self, ap: app.Application):
self.ap = ap
- self.llm_models = []
- self.embedding_models = []
- self.rerank_models = []
+ self.provider_dict = {}
+ self.llm_model_dict = {}
+ self.embedding_model_dict = {}
+ self.rerank_model_dict = {}
self.requester_components = []
self.requester_dict = {}
@@ -60,12 +75,75 @@ class ModelManager:
return litellm_provider
return None
- async def initialize(self):
+ @staticmethod
+ def _context_from_binding(
+ binding: WorkspaceExecutionBinding,
+ *,
+ trigger_principal: PrincipalContext | None = None,
+ ) -> ExecutionContext:
+ return ExecutionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ trigger_principal=trigger_principal,
+ )
+
+ @staticmethod
+ def _cache_key(context: ExecutionContext, resource_uuid: str) -> _CacheKey:
+ return (
+ context.instance_uuid,
+ context.workspace_uuid,
+ context.placement_generation,
+ resource_uuid,
+ )
+
+ @staticmethod
+ def _ensure_same_scope(
+ expected: ExecutionContext,
+ actual: ExecutionContext,
+ *,
+ resource: str,
+ ) -> None:
+ if (
+ actual.instance_uuid != expected.instance_uuid
+ or actual.workspace_uuid != expected.workspace_uuid
+ or actual.placement_generation != expected.placement_generation
+ ):
+ raise WorkspaceInvariantError(f'{resource} runtime belongs to another Workspace execution scope')
+
+ @staticmethod
+ def _ensure_entity_workspace(entity: object, context: ExecutionContext, *, resource: str) -> None:
+ workspace_uuid = getattr(entity, 'workspace_uuid', None)
+ if workspace_uuid != context.workspace_uuid:
+ raise WorkspaceInvariantError(f'{resource} belongs to another Workspace')
+
+ async def resolve_execution_context(self, context: TenantContext) -> ExecutionContext:
+ """Resolve and fence-check an explicit tenant context for runtime access."""
+
+ workspace_uuid = require_workspace_uuid(context)
+ expected_generation = None
+ supplied_instance_uuid = None
+ trigger_principal = None
+
+ if isinstance(context, (RequestContext, ExecutionContext)):
+ expected_generation = context.placement_generation
+ supplied_instance_uuid = context.instance_uuid
+ trigger_principal = context.principal if isinstance(context, RequestContext) else context.trigger_principal
+
+ binding = await self.ap.workspace_service.get_execution_binding(
+ workspace_uuid,
+ expected_generation=expected_generation,
+ )
+ if supplied_instance_uuid is not None and supplied_instance_uuid != binding.instance_uuid:
+ raise WorkspaceInvariantError('Runtime context belongs to another LangBot instance')
+
+ return self._context_from_binding(binding, trigger_principal=trigger_principal)
+
+ async def initialize(self) -> None:
self.requester_components = self.ap.discover.get_components_by_kind('LLMAPIRequester')
requester_dict: dict[str, type[requester.ProviderAPIRequester]] = {}
for component in self.requester_components:
- # Skip components that use litellm_provider (they will use litellmchat.py instead)
litellm_provider = self._get_litellm_provider_from_manifest(component)
if litellm_provider:
self.ap.logger.debug(
@@ -76,133 +154,151 @@ class ModelManager:
requester_dict[component.metadata.name] = component.get_python_component_class()
self.requester_dict = requester_dict
-
await self.load_models_from_db()
- # Check if space models service is disabled
space_config = self.ap.instance_config.data.get('space', {})
if space_config.get('disable_models_service', False):
self.ap.logger.info('LangBot Space Models service is disabled, skipping sync.')
return
+ # Space model synchronization is a legacy OSS-singleton facility. A
+ # cloud instance must receive tenant model projections from its control
+ # plane and must never infer one Workspace for this global operation.
+ try:
+ binding = await self.ap.workspace_service.get_local_execution_binding()
+ except WorkspaceError as exc:
+ self.ap.logger.info(f'Skipping LangBot Space model sync outside an OSS local Workspace: {exc}')
+ return
+
+ sync_context = self._context_from_binding(
+ binding,
+ trigger_principal=PrincipalContext(principal_type=PrincipalType.SYSTEM),
+ )
sync_timeout = space_config.get('models_sync_timeout')
try:
if sync_timeout:
await asyncio.wait_for(
- self.sync_new_models_from_space(),
+ self.sync_new_models_from_space(sync_context),
timeout=float(sync_timeout),
)
else:
- await self.sync_new_models_from_space()
+ await self.sync_new_models_from_space(sync_context)
except asyncio.TimeoutError:
self.ap.logger.warning(f'LangBot Space model sync timed out after {sync_timeout}s, skipping startup sync.')
- except Exception as e:
+ except Exception as exc:
self.ap.logger.warning('Failed to sync new models from LangBot Space, model list may not be updated.')
- self.ap.logger.warning(f' - Error: {e}')
+ self.ap.logger.warning(f' - Error: {exc}')
+
+ async def load_models_from_db(self) -> None:
+ """Load every active projected Workspace into isolated runtime caches."""
- async def load_models_from_db(self):
- """Load models from database"""
self.ap.logger.info('Loading models from db...')
-
- self.llm_models = []
- self.embedding_models = []
- self.rerank_models = []
self.provider_dict = {}
+ self.llm_model_dict = {}
+ self.embedding_model_dict = {}
+ self.rerank_model_dict = {}
+ contexts: dict[str, ExecutionContext] = {}
+
+ async def context_for(workspace_uuid: str | None) -> ExecutionContext:
+ if not workspace_uuid:
+ raise WorkspaceInvariantError('Runtime model resource has no Workspace')
+ cached = contexts.get(workspace_uuid)
+ if cached is not None:
+ return cached
+ binding = await self.ap.workspace_service.get_execution_binding(workspace_uuid)
+ resolved = self._context_from_binding(
+ binding,
+ trigger_principal=PrincipalContext(principal_type=PrincipalType.SYSTEM),
+ )
+ contexts[workspace_uuid] = resolved
+ return resolved
+
providers_result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_model.ModelProvider)
)
- for provider in providers_result.all():
+ for provider_entity in providers_result.all():
try:
- runtime_provider = await self.load_provider(provider)
- self.provider_dict[provider.uuid] = runtime_provider
- except provider_errors.RequesterNotFoundError as e:
- self.ap.logger.warning(f'Requester {e.requester_name} not found, skipping provider {provider.uuid}')
- continue
- except Exception as e:
- self.ap.logger.error(f'Failed to load provider {provider.uuid}: {e}\n{traceback.format_exc()}')
+ context = await context_for(provider_entity.workspace_uuid)
+ runtime_provider = await self._build_provider(context, provider_entity)
+ self.provider_dict[self._cache_key(context, provider_entity.uuid)] = runtime_provider
+ except provider_errors.RequesterNotFoundError as exc:
+ self.ap.logger.warning(
+ f'Requester {exc.requester_name} not found, skipping provider {provider_entity.uuid}'
+ )
+ except Exception as exc:
+ self.ap.logger.error(f'Failed to load provider {provider_entity.uuid}: {exc}\n{traceback.format_exc()}')
- # Load LLM models
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_model.LLMModel))
- llm_models = result.all()
- for llm_model in llm_models:
- try:
- provider = self.provider_dict.get(llm_model.provider_uuid)
- if provider is None:
- self.ap.logger.warning(f'Provider {llm_model.provider_uuid} not found for model {llm_model.uuid}')
- continue
- runtime_llm_model = await self.load_llm_model_with_provider(llm_model, provider)
- self.llm_models.append(runtime_llm_model)
- except Exception as e:
- self.ap.logger.error(f'Failed to load model {llm_model.uuid}: {e}\n{traceback.format_exc()}')
+ await self._load_model_kind(
+ persistence_model.LLMModel,
+ self.llm_model_dict,
+ self._build_llm_model,
+ context_for,
+ )
+ await self._load_model_kind(
+ persistence_model.EmbeddingModel,
+ self.embedding_model_dict,
+ self._build_embedding_model,
+ context_for,
+ )
+ await self._load_model_kind(
+ persistence_model.RerankModel,
+ self.rerank_model_dict,
+ self._build_rerank_model,
+ context_for,
+ )
- # Load embedding models
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_model.EmbeddingModel))
- embedding_models = result.all()
- for embedding_model in embedding_models:
+ async def _load_model_kind(self, entity_type, cache: dict, builder, context_for) -> None:
+ result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(entity_type))
+ for model_entity in result.all():
try:
- provider = self.provider_dict.get(embedding_model.provider_uuid)
+ context = await context_for(model_entity.workspace_uuid)
+ provider = self.provider_dict.get(self._cache_key(context, model_entity.provider_uuid))
if provider is None:
self.ap.logger.warning(
- f'Provider {embedding_model.provider_uuid} not found for model {embedding_model.uuid}'
+ f'Provider {model_entity.provider_uuid} not found for model {model_entity.uuid}'
)
continue
- runtime_embedding_model = await self.load_embedding_model_with_provider(embedding_model, provider)
- self.embedding_models.append(runtime_embedding_model)
- except Exception as e:
- self.ap.logger.error(f'Failed to load model {embedding_model.uuid}: {e}\n{traceback.format_exc()}')
+ runtime_model = builder(context, model_entity, provider)
+ cache[self._cache_key(context, model_entity.uuid)] = runtime_model
+ except Exception as exc:
+ self.ap.logger.error(f'Failed to load model {model_entity.uuid}: {exc}\n{traceback.format_exc()}')
- # Load rerank models
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_model.RerankModel))
- rerank_models = result.all()
- for rerank_model in rerank_models:
- try:
- provider = self.provider_dict.get(rerank_model.provider_uuid)
- if provider is None:
- self.ap.logger.warning(
- f'Provider {rerank_model.provider_uuid} not found for model {rerank_model.uuid}'
- )
- continue
- runtime_rerank_model = await self.load_rerank_model_with_provider(rerank_model, provider)
- self.rerank_models.append(runtime_rerank_model)
- except Exception as e:
- self.ap.logger.error(f'Failed to load model {rerank_model.uuid}: {e}\n{traceback.format_exc()}')
+ async def sync_new_models_from_space(self, context: ExecutionContext) -> None:
+ """Sync legacy Space models for the explicitly selected OSS Workspace."""
- async def sync_new_models_from_space(self):
- """Sync models from Space"""
- space_model_provider = await self.ap.persistence_mgr.execute_async(
+ context = await self.resolve_execution_context(context)
+ await self.ap.workspace_service.get_local_execution_binding(
+ context.workspace_uuid,
+ expected_generation=context.placement_generation,
+ )
+ space_model_provider_result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_model.ModelProvider).where(
- persistence_model.ModelProvider.requester == 'space-chat-completions'
+ persistence_model.ModelProvider.workspace_uuid == context.workspace_uuid,
+ persistence_model.ModelProvider.requester == 'space-chat-completions',
)
)
- result = space_model_provider.first()
- if result is None:
+ space_model_provider = space_model_provider_result.first()
+ if space_model_provider is None:
raise provider_errors.ProviderNotFoundError('LangBot Models')
- space_model_provider = result
-
- # get the latest models from space
space_models = await self.ap.space_service.get_models()
-
- # Index existing models by uuid. Space reuses a model's uuid across
- # renames / re-specs (e.g. the uuid that used to be ``claude-opus-4-6``
- # may later become ``claude-opus-4-7``). So for Space-managed models we
- # upsert: create when the uuid is new, otherwise update name/abilities/
- # ranking to track Space. Models owned by other providers are never
- # touched, even on an (unexpected) uuid collision.
- existing_llm_models = {m['uuid']: m for m in await self.ap.llm_model_service.get_llm_models()}
+ existing_llm_models = {
+ model['uuid']: model
+ for model in await self.ap.llm_model_service.get_llm_models(context, include_secret=True)
+ }
existing_embedding_models = {
- m['uuid']: m for m in await self.ap.embedding_models_service.get_embedding_models()
+ model['uuid']: model
+ for model in await self.ap.embedding_models_service.get_embedding_models(context, include_secret=True)
}
created = 0
updated = 0
-
for space_model in space_models:
if space_model.category == 'chat':
existing = existing_llm_models.get(space_model.uuid)
if existing is None:
- # model will be automatically loaded
await self.ap.llm_model_service.create_llm_model(
+ context,
{
'uuid': space_model.uuid,
'name': space_model.model_id,
@@ -227,14 +323,14 @@ class ModelManager:
or list(existing.get('abilities') or []) != list(desired['abilities'])
or existing.get('prefered_ranking') != desired['prefered_ranking']
):
- await self.ap.llm_model_service.update_llm_model(space_model.uuid, dict(desired))
+ await self.ap.llm_model_service.update_llm_model(context, space_model.uuid, dict(desired))
updated += 1
elif space_model.category == 'embedding':
existing = existing_embedding_models.get(space_model.uuid)
if existing is None:
- # model will be automatically loaded
await self.ap.embedding_models_service.create_embedding_model(
+ context,
{
'uuid': space_model.uuid,
'name': space_model.model_id,
@@ -255,7 +351,11 @@ class ModelManager:
existing.get('name') != desired['name']
or existing.get('prefered_ranking') != desired['prefered_ranking']
):
- await self.ap.embedding_models_service.update_embedding_model(space_model.uuid, dict(desired))
+ await self.ap.embedding_models_service.update_embedding_model(
+ context,
+ space_model.uuid,
+ dict(desired),
+ )
updated += 1
if created or updated:
@@ -263,313 +363,330 @@ class ModelManager:
async def init_temporary_runtime_llm_model(
self,
+ context: TenantContext,
model_info: dict,
) -> requester.RuntimeLLMModel:
- """Initialize runtime LLM model from dict (for testing)"""
- provider_info = model_info.get('provider', {})
-
- runtime_provider = await self.load_provider(provider_info)
-
- runtime_llm_model = requester.RuntimeLLMModel(
- model_entity=persistence_model.LLMModel(
- uuid=model_info.get('uuid', ''),
- name=model_info.get('name', ''),
- provider_uuid='',
- abilities=model_info.get('abilities', []),
- context_length=model_info.get('context_length'),
- extra_args=model_info.get('extra_args', {}),
- ),
- provider=runtime_provider,
+ execution_context = await self.resolve_execution_context(context)
+ provider_info = {**model_info.get('provider', {}), 'workspace_uuid': execution_context.workspace_uuid}
+ runtime_provider = await self._build_provider(
+ execution_context,
+ persistence_model.ModelProvider(**provider_info),
)
-
- return runtime_llm_model
+ model_entity = persistence_model.LLMModel(
+ workspace_uuid=execution_context.workspace_uuid,
+ uuid=model_info.get('uuid', ''),
+ name=model_info.get('name', ''),
+ provider_uuid=runtime_provider.provider_entity.uuid,
+ abilities=model_info.get('abilities', []),
+ context_length=model_info.get('context_length'),
+ extra_args=model_info.get('extra_args', {}),
+ )
+ return self._build_llm_model(execution_context, model_entity, runtime_provider)
async def init_temporary_runtime_embedding_model(
self,
+ context: TenantContext,
model_info: dict,
) -> requester.RuntimeEmbeddingModel:
- """Initialize runtime embedding model from dict (for testing)"""
- provider_info = model_info.get('provider', {})
- runtime_provider = await self.load_provider(provider_info)
-
- runtime_embedding_model = requester.RuntimeEmbeddingModel(
- model_entity=persistence_model.EmbeddingModel(
- uuid=model_info.get('uuid', ''),
- name=model_info.get('name', ''),
- provider_uuid='',
- extra_args=model_info.get('extra_args', {}),
- ),
- provider=runtime_provider,
+ execution_context = await self.resolve_execution_context(context)
+ provider_info = {**model_info.get('provider', {}), 'workspace_uuid': execution_context.workspace_uuid}
+ runtime_provider = await self._build_provider(
+ execution_context,
+ persistence_model.ModelProvider(**provider_info),
)
-
- return runtime_embedding_model
+ model_entity = persistence_model.EmbeddingModel(
+ workspace_uuid=execution_context.workspace_uuid,
+ uuid=model_info.get('uuid', ''),
+ name=model_info.get('name', ''),
+ provider_uuid=runtime_provider.provider_entity.uuid,
+ extra_args=model_info.get('extra_args', {}),
+ )
+ return self._build_embedding_model(execution_context, model_entity, runtime_provider)
async def init_temporary_runtime_rerank_model(
self,
+ context: TenantContext,
model_info: dict,
) -> requester.RuntimeRerankModel:
- """Initialize runtime rerank model from dict (for testing)"""
- provider_info = model_info.get('provider', {})
- runtime_provider = await self.load_provider(provider_info)
-
- runtime_rerank_model = requester.RuntimeRerankModel(
- model_entity=persistence_model.RerankModel(
- uuid=model_info.get('uuid', ''),
- name=model_info.get('name', ''),
- provider_uuid='',
- extra_args=model_info.get('extra_args', {}),
- ),
- provider=runtime_provider,
+ execution_context = await self.resolve_execution_context(context)
+ provider_info = {**model_info.get('provider', {}), 'workspace_uuid': execution_context.workspace_uuid}
+ runtime_provider = await self._build_provider(
+ execution_context,
+ persistence_model.ModelProvider(**provider_info),
)
+ model_entity = persistence_model.RerankModel(
+ workspace_uuid=execution_context.workspace_uuid,
+ uuid=model_info.get('uuid', ''),
+ name=model_info.get('name', ''),
+ provider_uuid=runtime_provider.provider_entity.uuid,
+ extra_args=model_info.get('extra_args', {}),
+ )
+ return self._build_rerank_model(execution_context, model_entity, runtime_provider)
- return runtime_rerank_model
-
- async def load_provider(
- self, provider_info: persistence_model.ModelProvider | sqlalchemy.Row | dict
- ) -> requester.RuntimeProvider:
- """Load provider from dict"""
+ @staticmethod
+ def _coerce_provider(
+ provider_info: persistence_model.ModelProvider | sqlalchemy.Row | dict,
+ context: ExecutionContext,
+ ) -> persistence_model.ModelProvider:
if isinstance(provider_info, sqlalchemy.Row):
provider_entity = persistence_model.ModelProvider(**provider_info._mapping)
elif isinstance(provider_info, dict):
- provider_entity = persistence_model.ModelProvider(**provider_info)
+ provider_entity = persistence_model.ModelProvider(
+ **{**provider_info, 'workspace_uuid': context.workspace_uuid}
+ )
else:
provider_entity = provider_info
+ ModelManager._ensure_entity_workspace(provider_entity, context, resource='Provider')
+ return provider_entity
- # Get requester manifest to check for litellm_provider
+ async def _build_provider(
+ self,
+ context: ExecutionContext,
+ provider_info: persistence_model.ModelProvider | sqlalchemy.Row | dict,
+ ) -> requester.RuntimeProvider:
+ provider_entity = self._coerce_provider(provider_info, context)
requester_manifest = self.get_available_requester_manifest_by_name(provider_entity.requester)
litellm_provider = self._get_litellm_provider_from_manifest(requester_manifest)
-
- # Build config from base_url
config = {'base_url': provider_entity.base_url}
- # Check if requester manifest specifies litellm_provider
if litellm_provider:
from .requesters import litellmchat
- # Use unified LiteLLMRequester with provider prefix
- # Map litellm_provider (YAML spec) to custom_llm_provider (config)
config['custom_llm_provider'] = litellm_provider
- requester_inst = litellmchat.LiteLLMRequester(
- ap=self.ap,
- config=config,
- )
+ requester_inst = litellmchat.LiteLLMRequester(ap=self.ap, config=config)
self.ap.logger.debug(
f'Using LiteLLMRequester for {provider_entity.requester} '
f'with custom_llm_provider={config["custom_llm_provider"]}'
)
else:
- # Use original requester class (for backward compatibility)
if provider_entity.requester not in self.requester_dict:
raise provider_errors.RequesterNotFoundError(provider_entity.requester)
- requester_inst = self.requester_dict[provider_entity.requester](
- ap=self.ap,
- config=config,
- )
+ requester_inst = self.requester_dict[provider_entity.requester](ap=self.ap, config=config)
await requester_inst.initialize()
-
token_mgr = token.TokenManager(name=provider_entity.uuid, tokens=provider_entity.api_keys or [])
-
- provider = requester.RuntimeProvider(
+ return requester.RuntimeProvider(
+ execution_context=context,
provider_entity=provider_entity,
token_mgr=token_mgr,
requester=requester_inst,
)
+
+ async def load_provider(
+ self,
+ context: TenantContext,
+ provider_info: persistence_model.ModelProvider | sqlalchemy.Row | dict,
+ ) -> requester.RuntimeProvider:
+ execution_context = await self.resolve_execution_context(context)
+ return await self._build_provider(execution_context, provider_info)
+
+ async def cache_provider(self, context: TenantContext, provider: requester.RuntimeProvider) -> None:
+ execution_context = await self.resolve_execution_context(context)
+ self._ensure_same_scope(execution_context, provider.execution_context, resource='Provider')
+ self._ensure_entity_workspace(provider.provider_entity, execution_context, resource='Provider')
+ self.provider_dict[self._cache_key(execution_context, provider.provider_entity.uuid)] = provider
+
+ async def get_provider_by_uuid(
+ self,
+ context: TenantContext,
+ provider_uuid: str,
+ ) -> requester.RuntimeProvider:
+ execution_context = await self.resolve_execution_context(context)
+ provider = self.provider_dict.get(self._cache_key(execution_context, provider_uuid))
+ if provider is None:
+ raise ValueError(f'Model provider {provider_uuid} not found')
+ self._ensure_same_scope(execution_context, provider.execution_context, resource='Provider')
return provider
- async def remove_provider(self, provider_uuid: str):
- """Remove provider
+ async def remove_provider(self, context: TenantContext, provider_uuid: str) -> None:
+ execution_context = await self.resolve_execution_context(context)
+ self.provider_dict.pop(self._cache_key(execution_context, provider_uuid), None)
- This method will not consider the models using this provider,
- because the models should be removed by the caller.
- """
- del self.provider_dict[provider_uuid]
-
- async def reload_provider(self, provider_uuid: str):
- """Reload provider"""
- provider_entity = await self.ap.persistence_mgr.execute_async(
+ async def reload_provider(self, context: TenantContext, provider_uuid: str) -> None:
+ execution_context = await self.resolve_execution_context(context)
+ result = await self.ap.persistence_mgr.execute_async(
sqlalchemy.select(persistence_model.ModelProvider).where(
- persistence_model.ModelProvider.uuid == provider_uuid
+ persistence_model.ModelProvider.workspace_uuid == execution_context.workspace_uuid,
+ persistence_model.ModelProvider.uuid == provider_uuid,
)
)
- provider_entity = provider_entity.first()
+ provider_entity = result.first()
if provider_entity is None:
raise provider_errors.ProviderNotFoundError(provider_uuid)
- new_runtime_provider = await self.load_provider(provider_entity)
+ new_provider = await self._build_provider(execution_context, provider_entity)
+ cache_prefix = self._cache_key(execution_context, '')[:3]
+ for cache in (self.llm_model_dict, self.embedding_model_dict, self.rerank_model_dict):
+ for key, model in cache.items():
+ if key[:3] == cache_prefix and model.provider.provider_entity.uuid == provider_uuid:
+ model.provider = new_provider
+ self.provider_dict[self._cache_key(execution_context, provider_uuid)] = new_provider
- # update refs in runtime models
- for model in self.llm_models:
- if model.provider.provider_entity.uuid == provider_uuid:
- model.provider = new_runtime_provider
- for model in self.embedding_models:
- if model.provider.provider_entity.uuid == provider_uuid:
- model.provider = new_runtime_provider
- for model in self.rerank_models:
- if model.provider.provider_entity.uuid == provider_uuid:
- model.provider = new_runtime_provider
+ @staticmethod
+ def _coerce_model(model_info: _ModelEntity | sqlalchemy.Row, entity_type: type[_ModelEntity]) -> _ModelEntity:
+ if isinstance(model_info, sqlalchemy.Row):
+ return entity_type(**model_info._mapping)
+ return model_info
- # update ref in provider dict
- self.provider_dict[provider_uuid] = new_runtime_provider
-
- async def load_llm_model_with_provider(
+ def _validate_model_provider(
self,
+ context: ExecutionContext,
+ model_entity: _ModelEntity,
+ provider: requester.RuntimeProvider,
+ ) -> None:
+ self._ensure_entity_workspace(model_entity, context, resource='Model')
+ self._ensure_same_scope(context, provider.execution_context, resource='Provider')
+ if model_entity.provider_uuid != provider.provider_entity.uuid:
+ raise WorkspaceInvariantError('Model references a different provider')
+
+ def _build_llm_model(
+ self,
+ context: ExecutionContext,
model_info: persistence_model.LLMModel | sqlalchemy.Row,
provider: requester.RuntimeProvider,
) -> requester.RuntimeLLMModel:
- """Load LLM model with provider info"""
- if isinstance(model_info, sqlalchemy.Row):
- model_info = persistence_model.LLMModel(**model_info._mapping)
-
- runtime_llm_model = requester.RuntimeLLMModel(
- model_entity=model_info,
+ model_entity = self._coerce_model(model_info, persistence_model.LLMModel)
+ self._validate_model_provider(context, model_entity, provider)
+ return requester.RuntimeLLMModel(
+ execution_context=context,
+ model_entity=model_entity,
provider=provider,
)
- return runtime_llm_model
-
- async def load_embedding_model_with_provider(
+ def _build_embedding_model(
self,
+ context: ExecutionContext,
model_info: persistence_model.EmbeddingModel | sqlalchemy.Row,
provider: requester.RuntimeProvider,
) -> requester.RuntimeEmbeddingModel:
- """Load embedding model with provider info"""
- if isinstance(model_info, sqlalchemy.Row):
- model_info = persistence_model.EmbeddingModel(**model_info._mapping)
-
- runtime_embedding_model = requester.RuntimeEmbeddingModel(
- model_entity=model_info,
+ model_entity = self._coerce_model(model_info, persistence_model.EmbeddingModel)
+ self._validate_model_provider(context, model_entity, provider)
+ return requester.RuntimeEmbeddingModel(
+ execution_context=context,
+ model_entity=model_entity,
provider=provider,
)
- return runtime_embedding_model
-
- async def load_rerank_model_with_provider(
+ def _build_rerank_model(
self,
+ context: ExecutionContext,
model_info: persistence_model.RerankModel | sqlalchemy.Row,
provider: requester.RuntimeProvider,
) -> requester.RuntimeRerankModel:
- """Load rerank model with provider info"""
- if isinstance(model_info, sqlalchemy.Row):
- model_info = persistence_model.RerankModel(**model_info._mapping)
-
- runtime_rerank_model = requester.RuntimeRerankModel(
- model_entity=model_info,
+ model_entity = self._coerce_model(model_info, persistence_model.RerankModel)
+ self._validate_model_provider(context, model_entity, provider)
+ return requester.RuntimeRerankModel(
+ execution_context=context,
+ model_entity=model_entity,
provider=provider,
)
- return runtime_rerank_model
+ async def load_llm_model_with_provider(
+ self,
+ context: TenantContext,
+ model_info: persistence_model.LLMModel | sqlalchemy.Row,
+ provider: requester.RuntimeProvider,
+ ) -> requester.RuntimeLLMModel:
+ execution_context = await self.resolve_execution_context(context)
+ return self._build_llm_model(execution_context, model_info, provider)
- async def load_llm_model(self, model_info: dict):
- """Load LLM model from dict (with provider info)"""
- provider_info = model_info.get('provider', {})
- if not provider_info:
- raise ValueError('Provider info is required')
+ async def load_embedding_model_with_provider(
+ self,
+ context: TenantContext,
+ model_info: persistence_model.EmbeddingModel | sqlalchemy.Row,
+ provider: requester.RuntimeProvider,
+ ) -> requester.RuntimeEmbeddingModel:
+ execution_context = await self.resolve_execution_context(context)
+ return self._build_embedding_model(execution_context, model_info, provider)
- model_entity = persistence_model.LLMModel(
- uuid=model_info.get('uuid', ''),
- name=model_info.get('name', ''),
- provider_uuid=model_info.get('provider_uuid', ''),
- abilities=model_info.get('abilities', []),
- context_length=model_info.get('context_length'),
- extra_args=model_info.get('extra_args', {}),
- )
+ async def load_rerank_model_with_provider(
+ self,
+ context: TenantContext,
+ model_info: persistence_model.RerankModel | sqlalchemy.Row,
+ provider: requester.RuntimeProvider,
+ ) -> requester.RuntimeRerankModel:
+ execution_context = await self.resolve_execution_context(context)
+ return self._build_rerank_model(execution_context, model_info, provider)
- provider_entity = persistence_model.ModelProvider(
- uuid=provider_info.get('uuid', ''),
- name=provider_info.get('name', ''),
- requester=provider_info.get('requester', ''),
- base_url=provider_info.get('base_url', ''),
- api_keys=provider_info.get('api_keys', []),
- )
+ async def cache_llm_model(self, context: TenantContext, model: requester.RuntimeLLMModel) -> None:
+ execution_context = await self.resolve_execution_context(context)
+ self._ensure_same_scope(execution_context, model.execution_context, resource='LLM model')
+ self.llm_model_dict[self._cache_key(execution_context, model.model_entity.uuid)] = model
- await self.load_llm_model_with_provider(model_entity, provider_entity)
+ async def cache_embedding_model(
+ self,
+ context: TenantContext,
+ model: requester.RuntimeEmbeddingModel,
+ ) -> None:
+ execution_context = await self.resolve_execution_context(context)
+ self._ensure_same_scope(execution_context, model.execution_context, resource='Embedding model')
+ self.embedding_model_dict[self._cache_key(execution_context, model.model_entity.uuid)] = model
- async def load_embedding_model(self, model_info: dict):
- """Load embedding model from dict (with provider info)"""
- provider_info = model_info.get('provider', {})
- if not provider_info:
- raise ValueError('Provider info is required')
+ async def cache_rerank_model(self, context: TenantContext, model: requester.RuntimeRerankModel) -> None:
+ execution_context = await self.resolve_execution_context(context)
+ self._ensure_same_scope(execution_context, model.execution_context, resource='Rerank model')
+ self.rerank_model_dict[self._cache_key(execution_context, model.model_entity.uuid)] = model
- model_entity = persistence_model.EmbeddingModel(
- uuid=model_info.get('uuid', ''),
- name=model_info.get('name', ''),
- provider_uuid=model_info.get('provider_uuid', ''),
- extra_args=model_info.get('extra_args', {}),
- )
+ async def get_model_by_uuid(self, context: TenantContext, model_uuid: str) -> requester.RuntimeLLMModel:
+ execution_context = await self.resolve_execution_context(context)
+ model = self.llm_model_dict.get(self._cache_key(execution_context, model_uuid))
+ if model is None:
+ raise ValueError(f'LLM model {model_uuid} not found')
+ self._ensure_same_scope(execution_context, model.execution_context, resource='LLM model')
+ return model
- provider_entity = persistence_model.ModelProvider(
- uuid=provider_info.get('uuid', ''),
- name=provider_info.get('name', ''),
- requester=provider_info.get('requester', ''),
- base_url=provider_info.get('base_url', ''),
- api_keys=provider_info.get('api_keys', []),
- )
+ async def get_embedding_model_by_uuid(
+ self,
+ context: TenantContext,
+ model_uuid: str,
+ ) -> requester.RuntimeEmbeddingModel:
+ execution_context = await self.resolve_execution_context(context)
+ model = self.embedding_model_dict.get(self._cache_key(execution_context, model_uuid))
+ if model is None:
+ raise ValueError(f'Embedding model {model_uuid} not found')
+ self._ensure_same_scope(execution_context, model.execution_context, resource='Embedding model')
+ return model
- await self.load_embedding_model_with_provider(model_entity, provider_entity)
+ async def get_rerank_model_by_uuid(
+ self,
+ context: TenantContext,
+ model_uuid: str,
+ ) -> requester.RuntimeRerankModel:
+ execution_context = await self.resolve_execution_context(context)
+ model = self.rerank_model_dict.get(self._cache_key(execution_context, model_uuid))
+ if model is None:
+ raise ValueError(f'Rerank model {model_uuid} not found')
+ self._ensure_same_scope(execution_context, model.execution_context, resource='Rerank model')
+ return model
- async def get_model_by_uuid(self, uuid: str) -> requester.RuntimeLLMModel:
- """Get LLM model by uuid"""
- for model in self.llm_models:
- if model.model_entity.uuid == uuid:
- return model
- raise ValueError(f'LLM model {uuid} not found')
+ async def remove_llm_model(self, context: TenantContext, model_uuid: str) -> None:
+ execution_context = await self.resolve_execution_context(context)
+ self.llm_model_dict.pop(self._cache_key(execution_context, model_uuid), None)
- async def get_embedding_model_by_uuid(self, uuid: str) -> requester.RuntimeEmbeddingModel:
- """Get embedding model by uuid"""
- for model in self.embedding_models:
- if model.model_entity.uuid == uuid:
- return model
- raise ValueError(f'Embedding model {uuid} not found')
+ async def remove_embedding_model(self, context: TenantContext, model_uuid: str) -> None:
+ execution_context = await self.resolve_execution_context(context)
+ self.embedding_model_dict.pop(self._cache_key(execution_context, model_uuid), None)
- async def get_rerank_model_by_uuid(self, uuid: str) -> requester.RuntimeRerankModel:
- """Get rerank model by uuid"""
- for model in self.rerank_models:
- if model.model_entity.uuid == uuid:
- return model
- raise ValueError(f'Rerank model {uuid} not found')
-
- async def remove_llm_model(self, model_uuid: str):
- """Remove LLM model"""
- for model in self.llm_models:
- if model.model_entity.uuid == model_uuid:
- self.llm_models.remove(model)
- return
-
- async def remove_embedding_model(self, model_uuid: str):
- """Remove embedding model"""
- for model in self.embedding_models:
- if model.model_entity.uuid == model_uuid:
- self.embedding_models.remove(model)
- return
-
- async def remove_rerank_model(self, model_uuid: str):
- """Remove rerank model"""
- for model in self.rerank_models:
- if model.model_entity.uuid == model_uuid:
- self.rerank_models.remove(model)
- return
+ async def remove_rerank_model(self, context: TenantContext, model_uuid: str) -> None:
+ execution_context = await self.resolve_execution_context(context)
+ self.rerank_model_dict.pop(self._cache_key(execution_context, model_uuid), None)
def get_available_requesters_info(self, model_type: str) -> list[dict]:
- """Get all available requesters"""
- if model_type != '':
+ if model_type:
return [
component.to_plain_dict()
for component in self.requester_components
if model_type in component.spec['support_type']
]
- else:
- return [component.to_plain_dict() for component in self.requester_components]
+ return [component.to_plain_dict() for component in self.requester_components]
def get_available_requester_info_by_name(self, name: str) -> dict | None:
- """Get requester info by name"""
for component in self.requester_components:
if component.metadata.name == name:
return component.to_plain_dict()
return None
def get_available_requester_manifest_by_name(self, name: str) -> engine.Component | None:
- """Get requester manifest by name"""
for component in self.requester_components:
if component.metadata.name == name:
return component
diff --git a/src/langbot/pkg/provider/modelmgr/requester.py b/src/langbot/pkg/provider/modelmgr/requester.py
index 377f7d4a8..7bc6b43f2 100644
--- a/src/langbot/pkg/provider/modelmgr/requester.py
+++ b/src/langbot/pkg/provider/modelmgr/requester.py
@@ -5,7 +5,9 @@ import typing
import time
from ...core import app
+from ...api.http.context import ExecutionContext
from ...entity.persistence import model as persistence_model
+from ...workspace.errors import WorkspaceInvariantError
import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
from . import token
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
@@ -16,6 +18,20 @@ LLM_USAGE_QUERY_VARIABLE = '_llm_usage'
STREAM_USAGE_QUERY_VARIABLE = '_stream_usage'
+def _ensure_same_execution_scope(
+ expected: ExecutionContext,
+ actual: ExecutionContext,
+ *,
+ resource: str,
+) -> None:
+ if (
+ actual.instance_uuid != expected.instance_uuid
+ or actual.workspace_uuid != expected.workspace_uuid
+ or actual.placement_generation != expected.placement_generation
+ ):
+ raise WorkspaceInvariantError(f'{resource} belongs to another Workspace execution scope')
+
+
def _store_llm_usage(query: pipeline_query.Query | None, usage_info: dict | None) -> None:
"""Store the latest provider usage on the query for upstream action handlers."""
if query is None or not usage_info:
@@ -39,24 +55,61 @@ class RuntimeProvider:
def __init__(
self,
+ execution_context: ExecutionContext,
provider_entity: persistence_model.ModelProvider,
token_mgr: token.TokenManager,
requester: ProviderAPIRequester,
):
+ if provider_entity.workspace_uuid != execution_context.workspace_uuid:
+ raise WorkspaceInvariantError('Provider belongs to another Workspace')
+ self.execution_context = execution_context
self.provider_entity = provider_entity
self.token_mgr = token_mgr
self.requester = requester
+ def _validate_invocation(
+ self,
+ model: RuntimeLLMModel | RuntimeEmbeddingModel | RuntimeRerankModel,
+ execution_context: ExecutionContext,
+ ) -> None:
+ _ensure_same_execution_scope(self.execution_context, execution_context, resource='Provider invocation')
+ _ensure_same_execution_scope(self.execution_context, model.execution_context, resource='Runtime model')
+ if model.provider is not self:
+ raise WorkspaceInvariantError('Runtime model is attached to another provider')
+
+ def _resolve_llm_execution_context(
+ self,
+ query: pipeline_query.Query | None,
+ execution_context: ExecutionContext | None,
+ ) -> ExecutionContext:
+ if query is not None:
+ from ...pipeline.pool import get_query_execution_context
+
+ query_context = get_query_execution_context(query)
+ if execution_context is not None:
+ _ensure_same_execution_scope(
+ query_context,
+ execution_context,
+ resource='Explicit LLM invocation context',
+ )
+ return query_context
+ if execution_context is None:
+ raise WorkspaceInvariantError('LLM invocation requires an ExecutionContext when query is absent')
+ return execution_context
+
async def invoke_llm(
self,
- query: pipeline_query.Query,
+ query: pipeline_query.Query | None,
model: RuntimeLLMModel,
messages: typing.List[provider_message.Message],
funcs: typing.List[resource_tool.LLMTool] = None,
extra_args: dict[str, typing.Any] = {},
remove_think: bool = False,
+ execution_context: ExecutionContext | None = None,
) -> provider_message.Message:
"""Bridge method for invoking LLM with monitoring"""
+ invocation_context = self._resolve_llm_execution_context(query, execution_context)
+ self._validate_invocation(model, invocation_context)
# Start timing for monitoring
start_time = time.time()
input_tokens = 0
@@ -130,14 +183,17 @@ class RuntimeProvider:
async def invoke_llm_stream(
self,
- query: pipeline_query.Query,
+ query: pipeline_query.Query | None,
model: RuntimeLLMModel,
messages: typing.List[provider_message.Message],
funcs: typing.List[resource_tool.LLMTool] = None,
extra_args: dict[str, typing.Any] = {},
remove_think: bool = False,
+ execution_context: ExecutionContext | None = None,
) -> provider_message.MessageChunk:
"""Bridge method for invoking LLM stream with monitoring"""
+ invocation_context = self._resolve_llm_execution_context(query, execution_context)
+ self._validate_invocation(model, invocation_context)
# Start timing for monitoring
start_time = time.time()
status = 'success'
@@ -212,6 +268,8 @@ class RuntimeProvider:
model: RuntimeEmbeddingModel,
input_text: typing.List[str],
extra_args: dict[str, typing.Any] = {},
+ *,
+ execution_context: ExecutionContext,
knowledge_base_id: str | None = None,
query_text: str | None = None,
session_id: str | None = None,
@@ -219,6 +277,7 @@ class RuntimeProvider:
call_type: str | None = None,
) -> typing.List[typing.List[float]]:
"""Bridge method for invoking embedding with monitoring"""
+ self._validate_invocation(model, execution_context)
# Start timing for monitoring
start_time = time.time()
prompt_tokens = 0
@@ -254,6 +313,7 @@ class RuntimeProvider:
try:
await self.requester.ap.monitoring_service.record_embedding_call(
+ execution_context,
model_name=model.model_entity.name,
prompt_tokens=prompt_tokens,
total_tokens=total_tokens,
@@ -276,8 +336,11 @@ class RuntimeProvider:
query: str,
documents: typing.List[str],
extra_args: dict[str, typing.Any] = {},
+ *,
+ execution_context: ExecutionContext,
) -> typing.List[dict]:
"""Bridge method for invoking rerank with monitoring"""
+ self._validate_invocation(model, execution_context)
start_time = time.time()
status = 'success'
@@ -316,9 +379,16 @@ class RuntimeLLMModel:
def __init__(
self,
+ execution_context: ExecutionContext,
model_entity: persistence_model.LLMModel,
provider: RuntimeProvider,
):
+ _ensure_same_execution_scope(provider.execution_context, execution_context, resource='LLM model')
+ if model_entity.workspace_uuid != execution_context.workspace_uuid:
+ raise WorkspaceInvariantError('LLM model belongs to another Workspace')
+ if model_entity.provider_uuid != provider.provider_entity.uuid:
+ raise WorkspaceInvariantError('LLM model references another provider')
+ self.execution_context = execution_context
self.model_entity = model_entity
self.provider = provider
@@ -334,9 +404,16 @@ class RuntimeEmbeddingModel:
def __init__(
self,
+ execution_context: ExecutionContext,
model_entity: persistence_model.EmbeddingModel,
provider: RuntimeProvider,
):
+ _ensure_same_execution_scope(provider.execution_context, execution_context, resource='Embedding model')
+ if model_entity.workspace_uuid != execution_context.workspace_uuid:
+ raise WorkspaceInvariantError('Embedding model belongs to another Workspace')
+ if model_entity.provider_uuid != provider.provider_entity.uuid:
+ raise WorkspaceInvariantError('Embedding model references another provider')
+ self.execution_context = execution_context
self.model_entity = model_entity
self.provider = provider
@@ -352,9 +429,16 @@ class RuntimeRerankModel:
def __init__(
self,
+ execution_context: ExecutionContext,
model_entity: persistence_model.RerankModel,
provider: RuntimeProvider,
):
+ _ensure_same_execution_scope(provider.execution_context, execution_context, resource='Rerank model')
+ if model_entity.workspace_uuid != execution_context.workspace_uuid:
+ raise WorkspaceInvariantError('Rerank model belongs to another Workspace')
+ if model_entity.provider_uuid != provider.provider_entity.uuid:
+ raise WorkspaceInvariantError('Rerank model references another provider')
+ self.execution_context = execution_context
self.model_entity = model_entity
self.provider = provider
diff --git a/src/langbot/pkg/provider/runners/difysvapi.py b/src/langbot/pkg/provider/runners/difysvapi.py
index 7566d20aa..6846e587a 100644
--- a/src/langbot/pkg/provider/runners/difysvapi.py
+++ b/src/langbot/pkg/provider/runners/difysvapi.py
@@ -22,10 +22,12 @@ from langbot.libs.dify_service_api.v1 import client, errors
import httpx
-# Module-level store for paused-workflow form state. The key isolates the bot,
-# pipeline, adapter, and launcher; each value holds an insertion-ordered map of
-# form_token -> form_data so one conversation can pause multiple workflows.
-PendingFormKey = tuple[str, str, str, str, str]
+# Module-level store for paused-workflow form state. The key includes the full
+# execution scope before the bot, pipeline, adapter, and launcher dimensions;
+# each value holds an insertion-ordered map of form_token -> form_data so one
+# conversation can pause multiple workflows without crossing Workspaces or
+# placement generations.
+PendingFormKey = tuple[str, str, int, str, str, str, str, str]
_PENDING_FORMS: dict[PendingFormKey, 'OrderedDict[str, dict[str, typing.Any]]'] = {}
_PENDING_FORM_DEFAULT_TTL = 30 * 60 # 30 minutes safety cap
_STREAM_FORM_PLACEHOLDER = '\u200b'
@@ -48,10 +50,13 @@ def _dify_user_from_query(query: pipeline_query.Query) -> str:
def _session_key_from_query(query: pipeline_query.Query) -> PendingFormKey:
- """Build a process-local pending-form key isolated by bot and pipeline."""
+ """Build a process-local pending-form key isolated by execution scope."""
adapter = getattr(query, 'adapter', None)
adapter_type = f'{type(adapter).__module__}.{type(adapter).__qualname__}'
return (
+ str(getattr(query, 'instance_uuid', '') or ''),
+ str(getattr(query, 'workspace_uuid', '') or ''),
+ int(getattr(query, 'placement_generation', 0) or 0),
str(getattr(query, 'bot_uuid', '') or ''),
str(getattr(query, 'pipeline_uuid', '') or ''),
adapter_type,
@@ -74,8 +79,8 @@ def _prune_pending_forms(now: float | None = None) -> None:
def _set_pending_form(session_key: PendingFormKey, form_data: dict[str, typing.Any]) -> None:
_prune_pending_forms()
- if isinstance(session_key, tuple) and len(session_key) > 1:
- form_data['pipeline_uuid'] = session_key[1]
+ if isinstance(session_key, tuple) and len(session_key) == 8:
+ form_data['pipeline_uuid'] = session_key[4]
stored = dict(form_data)
expiration_time = stored.get('expiration_time')
try:
diff --git a/src/langbot/pkg/provider/runners/localagent.py b/src/langbot/pkg/provider/runners/localagent.py
index 6c877239c..af46a442e 100644
--- a/src/langbot/pkg/provider/runners/localagent.py
+++ b/src/langbot/pkg/provider/runners/localagent.py
@@ -11,6 +11,7 @@ import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
import langbot_plugin.api.entities.builtin.provider.message as provider_message
import langbot_plugin.api.entities.builtin.rag.context as rag_context
+from ...pipeline.pool import get_query_execution_context
rag_combined_prompt_template = """
The following are relevant context entries retrieved from the knowledge base.
@@ -227,7 +228,10 @@ class LocalAgentRunner(runner.RequestRunner):
# Primary model
if query.use_llm_model_uuid:
try:
- primary = await self.ap.model_mgr.get_model_by_uuid(query.use_llm_model_uuid)
+ primary = await self.ap.model_mgr.get_model_by_uuid(
+ get_query_execution_context(query),
+ query.use_llm_model_uuid,
+ )
candidates.append(primary)
except ValueError:
self.ap.logger.warning(f'Primary model {query.use_llm_model_uuid} not found')
@@ -236,7 +240,10 @@ class LocalAgentRunner(runner.RequestRunner):
fallback_uuids = (query.variables or {}).get('_fallback_model_uuids', [])
for fb_uuid in fallback_uuids:
try:
- fb_model = await self.ap.model_mgr.get_model_by_uuid(fb_uuid)
+ fb_model = await self.ap.model_mgr.get_model_by_uuid(
+ get_query_execution_context(query),
+ fb_uuid,
+ )
candidates.append(fb_model)
except ValueError:
self.ap.logger.warning(f'Fallback model {fb_uuid} not found, skipping')
@@ -346,12 +353,13 @@ class LocalAgentRunner(runner.RequestRunner):
if kb_uuids and user_message_text:
# only support text for now
all_results: list[rag_context.RetrievalResultEntry] = []
+ execution_context = get_query_execution_context(query)
kb_engine_plugins: set[str] = set()
# Retrieve from each knowledge base
for kb_uuid in kb_uuids:
- kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(kb_uuid)
+ kb = await self.ap.rag_mgr.get_knowledge_base_by_uuid(execution_context, kb_uuid)
if not kb:
self.ap.logger.warning(f'Knowledge base {kb_uuid} not found, skipping')
@@ -364,6 +372,7 @@ class LocalAgentRunner(runner.RequestRunner):
kb_engine_plugins.add(engine_plugin_id)
result = await kb.retrieve(
+ execution_context,
user_message_text,
settings={
'bot_uuid': query.bot_uuid or '',
@@ -398,7 +407,10 @@ class LocalAgentRunner(runner.RequestRunner):
)
if all_results and rerank_model_uuid:
try:
- rerank_model = await self.ap.model_mgr.get_rerank_model_by_uuid(rerank_model_uuid)
+ rerank_model = await self.ap.model_mgr.get_rerank_model_by_uuid(
+ execution_context,
+ rerank_model_uuid,
+ )
rerank_top_k = int(local_agent_config.get('rerank-top-k', 5))
doc_texts = []
@@ -411,6 +423,7 @@ class LocalAgentRunner(runner.RequestRunner):
model=rerank_model,
query=user_message_text,
documents=doc_texts_capped,
+ execution_context=execution_context,
)
scored = sorted(scores, key=lambda x: x.get('relevance_score', 0), reverse=True)
diff --git a/src/langbot/pkg/provider/session/sessionmgr.py b/src/langbot/pkg/provider/session/sessionmgr.py
index 8d7823651..4beda1661 100644
--- a/src/langbot/pkg/provider/session/sessionmgr.py
+++ b/src/langbot/pkg/provider/session/sessionmgr.py
@@ -1,12 +1,48 @@
from __future__ import annotations
import asyncio
+import dataclasses
-from ...core import app
from langbot_plugin.api.entities.builtin.provider import message as provider_message, prompt as provider_prompt
import langbot_plugin.api.entities.builtin.provider.session as provider_session
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
+from ...api.http.context import ExecutionContext
+from ...core import app
+from ...pipeline.pool import (
+ ExecutionContextMismatchError,
+ ExecutionContextRequiredError,
+ bind_execution_context,
+ get_query_execution_context,
+)
+
+SessionKey = tuple[
+ str,
+ str,
+ int,
+ str,
+ str,
+ int | str,
+]
+
+
+def _query_session_key(query: pipeline_query.Query) -> tuple[SessionKey, ExecutionContext]:
+ execution_context = get_query_execution_context(query)
+ bot_uuid = getattr(query, 'bot_uuid', None)
+ if not isinstance(bot_uuid, str) or not bot_uuid.strip():
+ raise ExecutionContextRequiredError('Query.bot_uuid is required for session lookup')
+
+ execution_context = bind_execution_context(execution_context, bot_uuid=bot_uuid)
+ key: SessionKey = (
+ execution_context.instance_uuid,
+ execution_context.workspace_uuid,
+ execution_context.placement_generation,
+ bot_uuid,
+ query.launcher_type.value,
+ query.launcher_id,
+ )
+ return key, execution_context
+
class SessionManager:
"""会话管理器"""
@@ -24,17 +60,39 @@ class SessionManager:
async def get_session(self, query: pipeline_query.Query) -> provider_session.Session:
"""获取会话"""
+ session_key, execution_context = _query_session_key(query)
for session in self.session_list:
- if query.launcher_type == session.launcher_type and query.launcher_id == session.launcher_id:
+ if getattr(session, '_langbot_session_key', None) == session_key:
return session
session_concurrency = self.ap.instance_config.data['concurrency']['session']
session = provider_session.Session(
+ instance_uuid=execution_context.instance_uuid,
+ workspace_uuid=execution_context.workspace_uuid,
+ placement_generation=execution_context.placement_generation,
+ bot_uuid=query.bot_uuid,
launcher_type=query.launcher_type,
launcher_id=query.launcher_id,
sender_id=query.sender_id,
)
+ session_context = dataclasses.replace(
+ execution_context,
+ pipeline_uuid=None,
+ query_uuid=None,
+ )
+ # langbot-plugin 0.4.13 ignores Workspace fields. Preserve them until
+ # the Workspace-aware SDK becomes the minimum supported version.
+ object.__setattr__(session, 'instance_uuid', session_context.instance_uuid)
+ object.__setattr__(session, 'workspace_uuid', session_context.workspace_uuid)
+ object.__setattr__(
+ session,
+ 'placement_generation',
+ session_context.placement_generation,
+ )
+ object.__setattr__(session, 'bot_uuid', query.bot_uuid)
+ object.__setattr__(session, '_execution_context', session_context)
+ object.__setattr__(session, '_langbot_session_key', session_key)
session._semaphore = asyncio.Semaphore(session_concurrency)
self.session_list.append(session)
return session
@@ -49,6 +107,17 @@ class SessionManager:
) -> provider_session.Conversation:
"""获取对话或创建对话"""
+ session_key, execution_context = _query_session_key(query)
+ if getattr(session, '_langbot_session_key', None) != session_key:
+ raise ExecutionContextMismatchError('Session does not belong to the Query execution scope')
+ execution_context = bind_execution_context(
+ execution_context,
+ bot_uuid=bot_uuid,
+ pipeline_uuid=pipeline_uuid,
+ )
+ if execution_context.bot_uuid != getattr(session, 'bot_uuid', None):
+ raise ExecutionContextMismatchError('Session bot_uuid does not match the Query execution scope')
+
if not session.conversations:
session.conversations = []
@@ -63,7 +132,11 @@ class SessionManager:
messages=prompt_messages,
)
- if session.using_conversation is None or session.using_conversation.pipeline_uuid != pipeline_uuid:
+ if (
+ session.using_conversation is None
+ or session.using_conversation.pipeline_uuid != pipeline_uuid
+ or session.using_conversation.bot_uuid != bot_uuid
+ ):
conversation = provider_session.Conversation(
prompt=prompt,
messages=[],
diff --git a/src/langbot/pkg/provider/tools/loaders/availability.py b/src/langbot/pkg/provider/tools/loaders/availability.py
index 58d795864..1b9293c3d 100644
--- a/src/langbot/pkg/provider/tools/loaders/availability.py
+++ b/src/langbot/pkg/provider/tools/loaders/availability.py
@@ -11,7 +11,7 @@ async def is_box_backend_available(ap: Any) -> bool:
if not getattr(box_service, 'available', False):
return False
try:
- status = await box_service.get_status()
+ status = await box_service.get_backend_status()
backend_info = status.get('backend', {})
return bool(backend_info.get('available', False))
except Exception:
diff --git a/src/langbot/pkg/provider/tools/loaders/mcp.py b/src/langbot/pkg/provider/tools/loaders/mcp.py
index 11a943545..062bee4af 100644
--- a/src/langbot/pkg/provider/tools/loaders/mcp.py
+++ b/src/langbot/pkg/provider/tools/loaders/mcp.py
@@ -23,6 +23,9 @@ from pydantic import AnyUrl
from .. import loader
from ....core import app
+from ....api.http.context import ExecutionContext
+from ....api.http.service.tenant import TenantContext, require_workspace_uuid
+from ....workspace.errors import WorkspaceError, WorkspaceInvariantError
import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
import langbot_plugin.api.entities.builtin.provider.message as provider_message
from ....entity.persistence import mcp as persistence_mcp
@@ -209,6 +212,8 @@ class _CallerReconnect(Exception):
class RuntimeMCPSession:
"""运行时 MCP 会话"""
+ _FENCE_POLL_INTERVAL = 5.0
+
ap: app.Application
server_name: str
@@ -248,11 +253,19 @@ class RuntimeMCPSession:
_box_stdio_runtime: BoxStdioSessionRuntime
- def __init__(self, server_name: str, server_config: dict, enable: bool, ap: app.Application):
+ def __init__(
+ self,
+ server_name: str,
+ server_config: dict,
+ enable: bool,
+ ap: app.Application,
+ execution_context: ExecutionContext,
+ ):
self.server_name = server_name
self.server_uuid = server_config.get('uuid', '')
self.server_config = server_config
self.ap = ap
+ self.execution_context = execution_context
self.enable = enable
self.session = None
@@ -295,6 +308,47 @@ class RuntimeMCPSession:
self._box_stdio_runtime = BoxStdioSessionRuntime(self)
self.box_config = self._box_stdio_runtime.config
+ async def _assert_execution_active(self) -> None:
+ """Fail closed when this long-lived session belongs to a stale placement."""
+
+ binding = await self.ap.workspace_service.get_execution_binding(
+ self.execution_context.workspace_uuid,
+ expected_generation=self.execution_context.placement_generation,
+ )
+ if binding.instance_uuid != self.execution_context.instance_uuid:
+ raise WorkspaceInvariantError('MCP session instance does not match the active Workspace binding')
+
+ async def _monitor_execution_fence(self) -> None:
+ """Poll the placement fence while an MCP transport is idle."""
+
+ while not self._shutdown_event.is_set():
+ await asyncio.sleep(self._FENCE_POLL_INTERVAL)
+ if self._shutdown_event.is_set():
+ return
+ await self._assert_execution_active()
+
+ async def _sleep_with_execution_fence(self, delay: float) -> None:
+ """Back off without reconnecting after the captured placement expires."""
+
+ await self._assert_execution_active()
+ try:
+ await asyncio.wait_for(self._shutdown_event.wait(), timeout=delay)
+ except asyncio.TimeoutError:
+ pass
+ if not self._shutdown_event.is_set():
+ await self._assert_execution_active()
+
+ def _stop_for_stale_execution(self, error: WorkspaceError) -> None:
+ """Mark the session terminal without retrying a fenced placement."""
+
+ self.status = MCPSessionStatus.ERROR
+ self.error_message = 'Workspace execution binding is stale'
+ self._shutdown_event.set()
+ self._ready_event.set()
+ self.ap.logger.info(
+ f'MCP session {self.server_name} stopped because its Workspace execution binding is stale: {error}'
+ )
+
async def _init_stdio_python_server(self):
if self._uses_box_stdio():
await self._box_stdio_runtime.initialize()
@@ -423,6 +477,7 @@ class RuntimeMCPSession:
async def _lifecycle_loop(self):
"""Manage the full MCP session lifecycle in a background task."""
try:
+ await self._assert_execution_active()
if self.server_config['mode'] == 'stdio':
await self._init_stdio_python_server()
elif self.server_config['mode'] == 'remote':
@@ -432,9 +487,11 @@ class RuntimeMCPSession:
elif self.server_config['mode'] == 'http':
await self._init_streamable_http_server()
else:
- raise ValueError(f'Unknown MCP server mode: {self.server_name}: {self.server_config}')
+ raise ValueError(f'Unknown MCP server mode for {self.server_name}')
+ await self._assert_execution_active()
await self.refresh()
+ await self._assert_execution_active()
self.status = MCPSessionStatus.CONNECTED
@@ -446,12 +503,16 @@ class RuntimeMCPSession:
monitor_task = asyncio.create_task(self._box_stdio_runtime.monitor_process_health())
shutdown_task = asyncio.create_task(self._shutdown_event.wait())
reconnect_task = asyncio.create_task(self._reconnect_event.wait())
+ fence_task = asyncio.create_task(self._monitor_execution_fence())
done, pending = await asyncio.wait(
- [shutdown_task, monitor_task, reconnect_task],
+ [shutdown_task, monitor_task, reconnect_task, fence_task],
return_when=asyncio.FIRST_COMPLETED,
)
for task in pending:
task.cancel()
+ await asyncio.gather(*pending, return_exceptions=True)
+ if fence_task in done and not self._shutdown_event.is_set():
+ fence_task.result()
if reconnect_task in done and not self._shutdown_event.is_set():
self._reconnect_event.clear()
self.ap.logger.info(
@@ -487,12 +548,16 @@ class RuntimeMCPSession:
else:
shutdown_task = asyncio.create_task(self._shutdown_event.wait())
reconnect_task = asyncio.create_task(self._reconnect_event.wait())
+ fence_task = asyncio.create_task(self._monitor_execution_fence())
done, pending = await asyncio.wait(
- [shutdown_task, reconnect_task],
+ [shutdown_task, reconnect_task, fence_task],
return_when=asyncio.FIRST_COMPLETED,
)
for task in pending:
task.cancel()
+ await asyncio.gather(*pending, return_exceptions=True)
+ if fence_task in done and not self._shutdown_event.is_set():
+ fence_task.result()
if reconnect_task in done and not self._shutdown_event.is_set():
self._reconnect_event.clear()
self.ap.logger.info(
@@ -555,7 +620,11 @@ class RuntimeMCPSession:
self.status = MCPSessionStatus.CONNECTING
self.error_message = None
self.error_phase = None
- await asyncio.sleep(1)
+ try:
+ await self._sleep_with_execution_fence(1)
+ except WorkspaceError as fence_error:
+ self._stop_for_stale_execution(fence_error)
+ return
continue
except _CallerReconnect:
# A tool/resource call hit a server-expired session and asked us
@@ -572,6 +641,7 @@ class RuntimeMCPSession:
self.error_message = None
self.error_phase = None
try:
+ await self._assert_execution_active()
if self.server_config['mode'] == 'stdio':
await self._init_stdio_python_server()
elif self.server_config['mode'] == 'remote':
@@ -581,8 +651,12 @@ class RuntimeMCPSession:
elif self.server_config['mode'] == 'http':
await self._init_streamable_http_server()
await self.refresh()
+ await self._assert_execution_active()
self.status = MCPSessionStatus.CONNECTED
self.ap.logger.info(f'MCP session {self.server_name} reconnected successfully after session expiry')
+ except WorkspaceError as reconnect_err:
+ self._stop_for_stale_execution(reconnect_err)
+ return
except Exception as reconnect_err:
self.status = MCPSessionStatus.ERROR
self.error_message = str(reconnect_err)
@@ -610,8 +684,15 @@ class RuntimeMCPSession:
self.status = MCPSessionStatus.CONNECTING
self.error_message = None
self.error_phase = None
- await asyncio.sleep(2)
+ try:
+ await self._sleep_with_execution_fence(2)
+ except WorkspaceError as fence_error:
+ self._stop_for_stale_execution(fence_error)
+ return
continue
+ except WorkspaceError as e:
+ self._stop_for_stale_execution(e)
+ return
except Exception as e:
self.retry_count = attempt + 1
if self._shutdown_event.is_set():
@@ -639,7 +720,11 @@ class RuntimeMCPSession:
self.status = MCPSessionStatus.CONNECTING
self.error_message = None
self.error_phase = None
- await asyncio.sleep(delay)
+ try:
+ await self._sleep_with_execution_fence(delay)
+ except WorkspaceError as fence_error:
+ self._stop_for_stale_execution(fence_error)
+ return
attempt += 1
@staticmethod
@@ -722,6 +807,7 @@ class RuntimeMCPSession:
Returns True if reconnection succeeded within the timeout.
"""
+ await self._assert_execution_active()
if self._shutdown_event.is_set():
return False
@@ -732,6 +818,7 @@ class RuntimeMCPSession:
try:
await asyncio.wait_for(reconnected_event.wait(), timeout=self._RECONNECT_WAIT_TIMEOUT)
+ await self._assert_execution_active()
return self.status == MCPSessionStatus.CONNECTED
except asyncio.TimeoutError:
self.ap.logger.warning(f'MCP session {self.server_name} reconnect timed out')
@@ -747,6 +834,7 @@ class RuntimeMCPSession:
if not self.enable:
return
+ await self._assert_execution_active()
# Create background task for lifecycle management with retry
self._lifecycle_task = asyncio.create_task(self._lifecycle_loop_with_retry())
@@ -758,11 +846,13 @@ class RuntimeMCPSession:
self.status = MCPSessionStatus.ERROR
raise Exception(f'Connection timeout after {startup_timeout} seconds')
+ await self._assert_execution_active()
# Check for errors
if self.status == MCPSessionStatus.ERROR:
raise Exception('Connection failed, please check URL')
async def refresh(self):
+ await self._assert_execution_active()
if not self.session:
return
@@ -778,6 +868,7 @@ class RuntimeMCPSession:
self.resource_capabilities = {}
tools = await self.session.list_tools()
+ await self._assert_execution_active()
self.ap.logger.debug(f'Refresh MCP tools: {tools}')
@@ -799,34 +890,44 @@ class RuntimeMCPSession:
)
await self._refresh_resources()
+ await self._assert_execution_active()
async def _refresh_resources(self):
+ await self._assert_execution_active()
if not self.session:
return
try:
cursor: str | None = None
for _ in range(MCP_RESOURCE_DISCOVERY_MAX_PAGES):
+ await self._assert_execution_active()
resources_result = await self.session.list_resources(cursor)
+ await self._assert_execution_active()
for resource in resources_result.resources:
self.resources.append(_resource_to_dict(resource))
cursor = getattr(resources_result, 'nextCursor', None)
if not cursor:
break
self.ap.logger.debug(f'Refresh MCP resources: {len(self.resources)} resources found')
+ except WorkspaceError:
+ raise
except Exception as e:
self.ap.logger.debug(f'MCP server {self.server_name} does not support resources or failed to list: {e}')
try:
cursor = None
for _ in range(MCP_RESOURCE_DISCOVERY_MAX_PAGES):
+ await self._assert_execution_active()
templates_result = await self.session.list_resource_templates(cursor)
+ await self._assert_execution_active()
for template in templates_result.resourceTemplates:
self.resource_templates.append(_resource_template_to_dict(template))
cursor = getattr(templates_result, 'nextCursor', None)
if not cursor:
break
self.ap.logger.debug(f'Refresh MCP resource templates: {len(self.resource_templates)} templates found')
+ except WorkspaceError:
+ raise
except Exception as e:
self.ap.logger.debug(
f'MCP server {self.server_name} does not support resource templates or failed to list: {e}'
@@ -945,12 +1046,15 @@ class RuntimeMCPSession:
arguments: dict,
query: pipeline_query.Query | None = None,
) -> list[provider_message.ContentElement]:
+ await self._assert_execution_active()
for attempt in range(2):
if not self.session:
raise Exception('MCP session is not connected')
try:
+ await self._assert_execution_active()
result = await self.session.call_tool(tool_name, arguments)
+ await self._assert_execution_active()
except Exception as e:
if attempt == 0 and self._is_session_terminated(e):
self.ap.logger.warning(
@@ -1016,6 +1120,7 @@ class RuntimeMCPSession:
query: pipeline_query.Query | None = None,
) -> dict:
"""Read a resource by URI with safety limits and audit metadata."""
+ await self._assert_execution_active()
if not self.session:
raise Exception('MCP session is not connected')
@@ -1042,7 +1147,9 @@ class RuntimeMCPSession:
if not self.session:
raise Exception('MCP session is not connected')
try:
+ await self._assert_execution_active()
result = await self.session.read_resource(AnyUrl(uri))
+ await self._assert_execution_active()
break
except Exception as e:
if attempt == 0 and self._is_session_terminated(e):
@@ -1123,6 +1230,7 @@ class RuntimeMCPSession:
'cache_hit': False,
'warnings': warnings,
}
+ await self._assert_execution_active()
self._resource_cache[cache_key] = {'cached_at': now, 'envelope': envelope}
self._record_resource_read_trace(query, envelope)
return envelope
@@ -1157,7 +1265,11 @@ class RuntimeMCPSession:
def get_runtime_info_dict(self) -> dict:
info = {
'status': self.status.value,
- 'error_message': self.error_message,
+ # Raw transport exceptions may echo command arguments, headers, or
+ # environment values. Detailed diagnostics belong in AUDIT_VIEW
+ # logs; resource-list responses expose only a stable status.
+ 'error_message': 'MCP runtime failed' if self.error_message else None,
+ 'error_code': 'runtime_error' if self.error_message else None,
'error_phase': self.error_phase.value if self.error_phase else None,
'retry_count': self.retry_count,
'tool_count': len(self.get_tools()),
@@ -1246,6 +1358,37 @@ class RuntimeMCPSession:
await self._box_stdio_runtime.cleanup_session()
+def _execution_context_from_tenant(context: TenantContext) -> ExecutionContext:
+ workspace_uuid = require_workspace_uuid(context)
+ instance_uuid = str(getattr(context, 'instance_uuid', '') or '').strip()
+ generation = getattr(context, 'placement_generation', None)
+ if not instance_uuid:
+ raise ValueError('MCP runtime requires an explicit instance UUID')
+ if isinstance(generation, bool) or not isinstance(generation, int) or generation <= 0:
+ raise ValueError('MCP runtime requires a positive placement generation')
+ return ExecutionContext(
+ instance_uuid=instance_uuid,
+ workspace_uuid=workspace_uuid,
+ placement_generation=generation,
+ bot_uuid=getattr(context, 'bot_uuid', None),
+ pipeline_uuid=getattr(context, 'pipeline_uuid', None),
+ query_uuid=getattr(context, 'query_uuid', None),
+ )
+
+
+def _execution_context_from_query(query: pipeline_query.Query) -> ExecutionContext:
+ return _execution_context_from_tenant(
+ ExecutionContext(
+ instance_uuid=str(getattr(query, 'instance_uuid', '') or ''),
+ workspace_uuid=str(getattr(query, 'workspace_uuid', '') or ''),
+ placement_generation=getattr(query, 'placement_generation', 0) or 0,
+ bot_uuid=getattr(query, 'bot_uuid', None),
+ pipeline_uuid=getattr(query, 'pipeline_uuid', None),
+ query_uuid=getattr(query, 'query_uuid', None),
+ )
+ )
+
+
# @loader.loader_class('mcp')
class MCPLoader(loader.ToolLoader):
"""MCP 工具加载器。
@@ -1253,7 +1396,7 @@ class MCPLoader(loader.ToolLoader):
在此加载器中管理所有与 MCP Server 的连接。
"""
- sessions: dict[str, RuntimeMCPSession]
+ sessions: dict[tuple[str, str, int, str], RuntimeMCPSession]
_last_listed_functions: list[resource_tool.LLMTool]
@@ -1265,6 +1408,21 @@ class MCPLoader(loader.ToolLoader):
self._last_listed_functions = []
self._hosted_mcp_tasks = []
+ async def _assert_execution_active(
+ self,
+ context: TenantContext,
+ ) -> ExecutionContext:
+ """Validate a caller's placement before accessing an MCP session."""
+
+ execution_context = _execution_context_from_tenant(context)
+ binding = await self.ap.workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ if binding.instance_uuid != execution_context.instance_uuid:
+ raise WorkspaceInvariantError('MCP caller instance does not match the active Workspace binding')
+ return execution_context
+
async def initialize(self):
await self.load_mcp_servers_from_db()
@@ -1278,15 +1436,51 @@ class MCPLoader(loader.ToolLoader):
for server in servers:
config = self.ap.persistence_mgr.serialize_model(persistence_mcp.MCPServer, server)
+ try:
+ binding = await self.ap.workspace_service.get_execution_binding(server.workspace_uuid)
+ execution_context = ExecutionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ )
+ except Exception as exc:
+ self.ap.logger.warning(
+ f'Skipping MCP server {server.uuid}: Workspace execution binding is unavailable: {exc}'
+ )
+ continue
- task = asyncio.create_task(self.host_mcp_server(config))
+ task = asyncio.create_task(self.host_mcp_server(execution_context, config))
self._hosted_mcp_tasks.append(task)
- async def host_mcp_server(self, server_config: dict):
+ @staticmethod
+ def _scope_key(context: TenantContext) -> tuple[str, str, int]:
+ execution_context = _execution_context_from_tenant(context)
+ return (
+ execution_context.instance_uuid,
+ execution_context.workspace_uuid,
+ execution_context.placement_generation,
+ )
+
+ @classmethod
+ def _session_key(cls, context: TenantContext, server_name: str) -> tuple[str, str, int, str]:
+ return (*cls._scope_key(context), server_name)
+
+ def _sessions_for_context(self, context: TenantContext) -> list[RuntimeMCPSession]:
+ scope_key = self._scope_key(context)
+ return [session for key, session in self.sessions.items() if key[:3] == scope_key]
+
+ async def host_mcp_server(self, context: TenantContext, server_config: dict):
+ execution_context = await self._assert_execution_active(context)
+ configured_workspace = str(server_config.get('workspace_uuid') or '').strip()
+ if configured_workspace and configured_workspace != execution_context.workspace_uuid:
+ raise ValueError('MCP server configuration belongs to another Workspace')
+ server_config = dict(server_config)
+ server_config['workspace_uuid'] = execution_context.workspace_uuid
self.ap.logger.debug(f'Loading MCP server {server_config}')
try:
- session = await self.load_mcp_server(server_config)
- self.sessions[server_config['name']] = session
+ session = await self.load_mcp_server(execution_context, server_config)
+ await self._assert_execution_active(execution_context)
+ self.sessions[self._session_key(execution_context, server_config['name'])] = session
except Exception as e:
self.ap.logger.error(
f'Failed to load MCP server from db: {server_config["name"]}({server_config["uuid"]}): {e}\n{traceback.format_exc()}'
@@ -1295,6 +1489,7 @@ class MCPLoader(loader.ToolLoader):
self.ap.logger.debug(f'Starting MCP server {server_config["name"]}({server_config["uuid"]})')
try:
+ await self._assert_execution_active(execution_context)
await session.start()
except Exception as e:
self.ap.logger.error(
@@ -1304,7 +1499,7 @@ class MCPLoader(loader.ToolLoader):
self.ap.logger.debug(f'Started MCP server {server_config["name"]}({server_config["uuid"]})')
- async def load_mcp_server(self, server_config: dict) -> RuntimeMCPSession:
+ async def load_mcp_server(self, context: TenantContext, server_config: dict) -> RuntimeMCPSession:
"""加载 MCP 服务器到运行时
Args:
@@ -1314,6 +1509,13 @@ class MCPLoader(loader.ToolLoader):
- enable: 是否启用
- extra_args: 额外的配置参数 (可选)
"""
+ execution_context = await self._assert_execution_active(context)
+ server_config = dict(server_config)
+ configured_workspace = str(server_config.get('workspace_uuid') or '').strip()
+ if configured_workspace and configured_workspace != execution_context.workspace_uuid:
+ raise ValueError('MCP server configuration belongs to another Workspace')
+ server_config['workspace_uuid'] = execution_context.workspace_uuid
+
uuid_ = server_config.get('uuid')
is_transient = False
if not uuid_:
@@ -1339,7 +1541,7 @@ class MCPLoader(loader.ToolLoader):
**extra_args,
}
- session = RuntimeMCPSession(name, mixed_config, enable, self.ap)
+ session = RuntimeMCPSession(name, mixed_config, enable, self.ap, execution_context)
return session
@@ -1348,9 +1550,13 @@ class MCPLoader(loader.ToolLoader):
v = getattr(query, 'variables', None) or {}
return v.get('_pipeline_bound_mcp_servers', None)
- def _eligible_sessions_for_bound(self, bound_mcp_servers: list[str] | None) -> list[RuntimeMCPSession]:
+ def _eligible_sessions_for_bound(
+ self,
+ context: TenantContext,
+ bound_mcp_servers: list[str] | None,
+ ) -> list[RuntimeMCPSession]:
out: list[RuntimeMCPSession] = []
- for session in self.sessions.values():
+ for session in self._sessions_for_context(context):
if not session.enable:
continue
if session.status != MCPSessionStatus.CONNECTED:
@@ -1362,10 +1568,14 @@ class MCPLoader(loader.ToolLoader):
out.append(session)
return out
- def _eligible_resource_sessions_for_bound(self, bound_mcp_servers: list[str] | None) -> list[RuntimeMCPSession]:
+ def _eligible_resource_sessions_for_bound(
+ self,
+ context: TenantContext,
+ bound_mcp_servers: list[str] | None,
+ ) -> list[RuntimeMCPSession]:
return [
session
- for session in self._eligible_sessions_for_bound(bound_mcp_servers)
+ for session in self._eligible_sessions_for_bound(context, bound_mcp_servers)
if session.has_resource_support()
]
@@ -1396,12 +1606,13 @@ class MCPLoader(loader.ToolLoader):
]
async def _invoke_mcp_list_resources(self, parameters: dict, query: pipeline_query.Query) -> typing.Any:
+ execution_context = _execution_context_from_query(query)
server_name = parameters.get('server_name') if parameters else None
if not server_name or not isinstance(server_name, str):
return [provider_message.ContentElement.from_text('Error: "server_name" (string) is required.')]
bound = self._get_bound_mcp_from_query(query)
- allowed = {s.server_name for s in self._eligible_resource_sessions_for_bound(bound)}
+ allowed = {s.server_name for s in self._eligible_resource_sessions_for_bound(execution_context, bound)}
if server_name not in allowed:
return [
provider_message.ContentElement.from_text(
@@ -1411,7 +1622,7 @@ class MCPLoader(loader.ToolLoader):
)
]
- session = self.get_session(server_name)
+ session = self.get_session(execution_context, server_name)
if session is None or session.status != MCPSessionStatus.CONNECTED:
return [provider_message.ContentElement.from_text(f'Error: MCP server not connected: {server_name!r}')]
@@ -1428,6 +1639,7 @@ class MCPLoader(loader.ToolLoader):
return [provider_message.ContentElement.from_text(json.dumps(body, ensure_ascii=False, indent=2))]
async def _invoke_mcp_read_resource(self, parameters: dict, query: pipeline_query.Query) -> typing.Any:
+ execution_context = _execution_context_from_query(query)
server_name = parameters.get('server_name') if parameters else None
uri = parameters.get('uri') if parameters else None
if not server_name or not isinstance(server_name, str):
@@ -1436,7 +1648,7 @@ class MCPLoader(loader.ToolLoader):
return [provider_message.ContentElement.from_text('Error: "uri" (string) is required.')]
bound = self._get_bound_mcp_from_query(query)
- allowed = {s.server_name for s in self._eligible_resource_sessions_for_bound(bound)}
+ allowed = {s.server_name for s in self._eligible_resource_sessions_for_bound(execution_context, bound)}
if server_name not in allowed:
return [
provider_message.ContentElement.from_text(
@@ -1445,7 +1657,7 @@ class MCPLoader(loader.ToolLoader):
)
]
- session = self.get_session(server_name)
+ session = self.get_session(execution_context, server_name)
if session is None or session.status != MCPSessionStatus.CONNECTED:
return [provider_message.ContentElement.from_text(f'Error: MCP server not connected: {server_name!r}')]
@@ -1496,13 +1708,15 @@ class MCPLoader(loader.ToolLoader):
async def get_tools(
self,
+ context: TenantContext,
bound_mcp_servers: list[str] | None = None,
*,
include_resource_tools: bool = True,
) -> list[resource_tool.LLMTool]:
+ await self._assert_execution_active(context)
all_functions: list[resource_tool.LLMTool] = []
- for session in self.sessions.values():
+ for session in self._sessions_for_context(context):
# If bound_mcp_servers is specified, only include tools from those servers
if bound_mcp_servers is not None:
if session.server_uuid in bound_mcp_servers:
@@ -1511,7 +1725,7 @@ class MCPLoader(loader.ToolLoader):
# If no bound servers specified, include all tools
all_functions.extend(session.get_tools())
- if include_resource_tools and self._eligible_resource_sessions_for_bound(bound_mcp_servers):
+ if include_resource_tools and self._eligible_resource_sessions_for_bound(context, bound_mcp_servers):
all_functions.extend(self._mcp_synthetic_resource_tools())
self._last_listed_functions = all_functions
@@ -1520,13 +1734,15 @@ class MCPLoader(loader.ToolLoader):
async def get_tool_catalog(
self,
+ context: TenantContext,
bound_mcp_servers: list[str] | None = None,
*,
include_resource_tools: bool = False,
) -> list[dict[str, typing.Any]]:
+ await self._assert_execution_active(context)
items: list[dict[str, typing.Any]] = []
- for session in self.sessions.values():
+ for session in self._sessions_for_context(context):
if bound_mcp_servers is not None and session.server_uuid not in bound_mcp_servers:
continue
for tool in session.get_tools():
@@ -1542,7 +1758,7 @@ class MCPLoader(loader.ToolLoader):
}
)
- if include_resource_tools and self._eligible_resource_sessions_for_bound(bound_mcp_servers):
+ if include_resource_tools and self._eligible_resource_sessions_for_bound(context, bound_mcp_servers):
for tool in self._mcp_synthetic_resource_tools():
items.append(
{
@@ -1558,18 +1774,20 @@ class MCPLoader(loader.ToolLoader):
return items
- async def has_tool(self, name: str) -> bool:
+ async def has_tool(self, context: TenantContext, name: str) -> bool:
"""检查工具是否存在"""
+ await self._assert_execution_active(context)
if name in (MCP_TOOL_LIST_RESOURCES, MCP_TOOL_READ_RESOURCE):
- return bool(self._eligible_resource_sessions_for_bound(None))
- for session in self.sessions.values():
+ return bool(self._eligible_resource_sessions_for_bound(context, None))
+ for session in self._sessions_for_context(context):
for function in session.get_tools():
if function.name == name:
return True
return False
- async def get_tool(self, name: str) -> resource_tool.LLMTool | None:
- for session in self.sessions.values():
+ async def get_tool(self, context: TenantContext, name: str) -> resource_tool.LLMTool | None:
+ await self._assert_execution_active(context)
+ for session in self._sessions_for_context(context):
for function in session.get_tools():
if function.name == name:
return function
@@ -1577,6 +1795,7 @@ class MCPLoader(loader.ToolLoader):
async def invoke_tool(self, name: str, parameters: dict, query: pipeline_query.Query) -> typing.Any:
"""执行工具调用"""
+ execution_context = await self._assert_execution_active(_execution_context_from_query(query))
if name == MCP_TOOL_LIST_RESOURCES:
if getattr(query, 'variables', {}).get('_pipeline_mcp_resource_agent_read_enabled', True) is False:
return [provider_message.ContentElement.from_text('Error: MCP resource agent reads are disabled.')]
@@ -1586,7 +1805,7 @@ class MCPLoader(loader.ToolLoader):
return [provider_message.ContentElement.from_text('Error: MCP resource agent reads are disabled.')]
return await self._invoke_mcp_read_resource(parameters, query)
- for session in self.sessions.values():
+ for session in self._sessions_for_context(execution_context):
for function in session.get_tools():
if function.name == name:
self.ap.logger.debug(f'Invoking MCP tool: {name} with parameters: {parameters}')
@@ -1600,22 +1819,25 @@ class MCPLoader(loader.ToolLoader):
raise ValueError(f'Tool not found: {name}')
- async def get_resources(self, server_name: str) -> list[dict]:
+ async def get_resources(self, context: TenantContext, server_name: str) -> list[dict]:
"""Get resources from a specific MCP server."""
- session = self.get_session(server_name)
+ await self._assert_execution_active(context)
+ session = self.get_session(context, server_name)
if session is None:
raise ValueError(f'MCP server not found: {server_name}')
return session.get_resources()
- async def get_resource_templates(self, server_name: str) -> list[dict]:
+ async def get_resource_templates(self, context: TenantContext, server_name: str) -> list[dict]:
"""Get resource templates from a specific MCP server."""
- session = self.get_session(server_name)
+ await self._assert_execution_active(context)
+ session = self.get_session(context, server_name)
if session is None:
raise ValueError(f'MCP server not found: {server_name}')
return session.get_resource_templates()
async def read_resource_envelope(
self,
+ context: TenantContext,
server_name: str,
uri: str,
*,
@@ -1626,7 +1848,8 @@ class MCPLoader(loader.ToolLoader):
query: pipeline_query.Query | None = None,
) -> dict:
"""Read a resource from a specific MCP server and return metadata plus contents."""
- session = self.get_session(server_name)
+ await self._assert_execution_active(context)
+ session = self.get_session(context, server_name)
if session is None:
raise ValueError(f'MCP server not found: {server_name}')
return await session.read_resource_envelope(
@@ -1638,24 +1861,28 @@ class MCPLoader(loader.ToolLoader):
query=query,
)
- async def read_resource(self, server_name: str, uri: str) -> list[dict]:
+ async def read_resource(self, context: TenantContext, server_name: str, uri: str) -> list[dict]:
"""Read a resource from a specific MCP server."""
- envelope = await self.read_resource_envelope(server_name, uri)
+ envelope = await self.read_resource_envelope(context, server_name, uri)
return envelope['contents']
- def get_session_by_uuid(self, server_uuid: str) -> RuntimeMCPSession | None:
- for session in self.sessions.values():
+ def get_session_by_uuid(self, context: TenantContext, server_uuid: str) -> RuntimeMCPSession | None:
+ for session in self._sessions_for_context(context):
if session.server_uuid == server_uuid:
return session
return None
- def _resolve_attachment_session(self, attachment: dict) -> RuntimeMCPSession | None:
+ def _resolve_attachment_session(
+ self,
+ context: TenantContext,
+ attachment: dict,
+ ) -> RuntimeMCPSession | None:
server_uuid = attachment.get('server_uuid') or attachment.get('server_id')
server_name = attachment.get('server_name')
if server_uuid:
- return self.get_session_by_uuid(server_uuid)
+ return self.get_session_by_uuid(context, server_uuid)
if server_name:
- return self.get_session(server_name)
+ return self.get_session(context, server_name)
return None
async def build_resource_context_for_query(
@@ -1666,6 +1893,7 @@ class MCPLoader(loader.ToolLoader):
default_max_bytes: int = MCP_RESOURCE_CONTEXT_MAX_BYTES,
) -> str:
"""Build host-controlled MCP resource context for the current query."""
+ execution_context = await self._assert_execution_active(_execution_context_from_query(query))
if getattr(query, 'variables', {}).get('_pipeline_mcp_resource_agent_read_enabled', True) is False:
return ''
@@ -1674,7 +1902,7 @@ class MCPLoader(loader.ToolLoader):
return ''
bound = self._get_bound_mcp_from_query(query)
- eligible = self._eligible_resource_sessions_for_bound(bound)
+ eligible = self._eligible_resource_sessions_for_bound(execution_context, bound)
eligible_by_uuid = {session.server_uuid: session for session in eligible}
eligible_by_name = {session.server_name: session for session in eligible}
@@ -1682,6 +1910,7 @@ class MCPLoader(loader.ToolLoader):
remaining_tokens = default_max_tokens
for raw_attachment in attachments:
+ await self._assert_execution_active(execution_context)
if remaining_tokens <= 0:
break
if not isinstance(raw_attachment, dict) or raw_attachment.get('enabled') is False:
@@ -1696,7 +1925,7 @@ class MCPLoader(loader.ToolLoader):
if not uri or not isinstance(uri, str):
continue
- session = self._resolve_attachment_session(attachment)
+ session = self._resolve_attachment_session(execution_context, attachment)
if session is None:
continue
if session.server_uuid not in eligible_by_uuid and session.server_name not in eligible_by_name:
@@ -1714,6 +1943,8 @@ class MCPLoader(loader.ToolLoader):
source='preloaded',
query=query,
)
+ except WorkspaceError:
+ raise
except Exception as e:
self.ap.logger.warning(f'Failed to preload MCP resource {uri!r} from {session.server_name!r}: {e}')
continue
@@ -1753,37 +1984,40 @@ class MCPLoader(loader.ToolLoader):
pass
return context
- async def remove_mcp_server(self, server_name: str):
+ async def remove_mcp_server(self, context: TenantContext, server_name: str):
"""移除 MCP 服务器"""
- if server_name not in self.sessions:
+ await self._assert_execution_active(context)
+ key = self._session_key(context, server_name)
+ if key not in self.sessions:
self.ap.logger.warning(f'MCP server {server_name} not found in sessions, skipping removal')
return
- session = self.sessions.pop(server_name)
+ session = self.sessions.pop(key)
await session.shutdown()
self.ap.logger.info(f'Removed MCP server: {server_name}')
- def get_session(self, server_name: str) -> RuntimeMCPSession | None:
+ def get_session(self, context: TenantContext, server_name: str) -> RuntimeMCPSession | None:
"""获取指定名称的 MCP 会话"""
- return self.sessions.get(server_name)
+ return self.sessions.get(self._session_key(context, server_name))
- def has_session(self, server_name: str) -> bool:
+ def has_session(self, context: TenantContext, server_name: str) -> bool:
"""检查是否存在指定名称的 MCP 会话"""
- return server_name in self.sessions
+ return self._session_key(context, server_name) in self.sessions
- def get_all_server_names(self) -> list[str]:
+ def get_all_server_names(self, context: TenantContext) -> list[str]:
"""获取所有已加载的 MCP 服务器名称"""
- return list(self.sessions.keys())
+ return [session.server_name for session in self._sessions_for_context(context)]
- def get_server_tool_count(self, server_name: str) -> int:
+ def get_server_tool_count(self, context: TenantContext, server_name: str) -> int:
"""获取指定服务器的工具数量"""
- session = self.get_session(server_name)
+ session = self.get_session(context, server_name)
return len(session.get_tools()) if session else 0
- def get_all_servers_info(self) -> dict[str, dict]:
+ def get_all_servers_info(self, context: TenantContext) -> dict[str, dict]:
"""获取所有服务器的信息"""
info = {}
- for server_name, session in self.sessions.items():
+ for session in self._sessions_for_context(context):
+ server_name = session.server_name
tools = session.get_tools()
info[server_name] = {
'name': server_name,
@@ -1797,11 +2031,13 @@ class MCPLoader(loader.ToolLoader):
async def shutdown(self):
"""关闭所有工具"""
self.ap.logger.info('Shutting down all MCP sessions...')
- for server_name, session in list(self.sessions.items()):
+ for key, session in list(self.sessions.items()):
try:
await session.shutdown()
- self.ap.logger.debug(f'Shutdown MCP session: {server_name}')
+ self.ap.logger.debug(f'Shutdown MCP session: {session.server_name}')
except Exception as e:
- self.ap.logger.error(f'Error shutting down MCP session {server_name}: {e}\n{traceback.format_exc()}')
+ self.ap.logger.error(
+ f'Error shutting down MCP session {session.server_name}: {e}\n{traceback.format_exc()}'
+ )
self.sessions.clear()
self.ap.logger.info('All MCP sessions shutdown complete')
diff --git a/src/langbot/pkg/provider/tools/loaders/mcp_stdio.py b/src/langbot/pkg/provider/tools/loaders/mcp_stdio.py
index 134c60a88..590a11479 100644
--- a/src/langbot/pkg/provider/tools/loaders/mcp_stdio.py
+++ b/src/langbot/pkg/provider/tools/loaders/mcp_stdio.py
@@ -6,7 +6,7 @@ import os
import shutil
import shlex
import threading
-from contextlib import suppress, AsyncExitStack
+from contextlib import suppress, AsyncExitStack, asynccontextmanager
from typing import TYPE_CHECKING, Any
import pydantic
@@ -94,6 +94,60 @@ class MCPServerBoxConfig(pydantic.BaseModel):
_HANDSHAKE_ATTEMPT_TIMEOUT_SEC = 10.0
+@asynccontextmanager
+async def authenticated_websocket_client(url: str, headers: dict[str, str]):
+ """MCP WebSocket transport with host-only Box relay headers.
+
+ The upstream MCP helper does not expose WebSocket handshake headers. This
+ mirrors that transport while keeping the Box control token out of the URL,
+ JSON-RPC payloads, and logs.
+ """
+
+ import json
+
+ import anyio
+ import mcp.types as mcp_types
+ from mcp.shared.message import SessionMessage
+ from pydantic import ValidationError
+ from websockets.asyncio.client import connect as ws_connect
+ from websockets.typing import Subprotocol
+
+ read_stream_writer, read_stream = anyio.create_memory_object_stream(0)
+ write_stream, write_stream_reader = anyio.create_memory_object_stream(0)
+
+ async with ws_connect(
+ url,
+ subprotocols=[Subprotocol('mcp')],
+ additional_headers=dict(headers),
+ proxy=None,
+ ) as websocket:
+
+ async def ws_reader():
+ async with read_stream_writer:
+ async for raw_text in websocket:
+ try:
+ message = mcp_types.JSONRPCMessage.model_validate_json(raw_text)
+ await read_stream_writer.send(SessionMessage(message))
+ except ValidationError as exc: # pragma: no cover - upstream parity
+ await read_stream_writer.send(exc)
+
+ async def ws_writer():
+ async with write_stream_reader:
+ async for session_message in write_stream_reader:
+ payload = session_message.message.model_dump(
+ by_alias=True,
+ mode='json',
+ exclude_none=True,
+ )
+ await websocket.send(json.dumps(payload))
+
+ async with anyio.create_task_group() as task_group:
+ task_group.start_soon(ws_reader)
+ task_group.start_soon(ws_writer)
+ yield read_stream, write_stream
+ task_group.cancel_scope.cancel()
+
+
class _TransferredStack:
"""Adapts an already-populated AsyncExitStack into an async context manager
so ownership of its resources can be transferred into another exit stack.
@@ -149,6 +203,7 @@ class BoxStdioSessionRuntime:
resolved_host_path = self.resolve_host_path() if host_path is ... else host_path
return BoxWorkspaceSession(
self.ap.box_service,
+ self.owner.execution_context,
self.owner._build_box_session_id(),
host_path=resolved_host_path,
host_path_mode=self.config.host_path_mode,
@@ -250,7 +305,11 @@ class BoxStdioSessionRuntime:
if install_cmd:
payload = self._wrap_process_payload_with_python_env(payload, process_cwd)
payload['process_id'] = self.process_id
- await workspace.box_service.start_managed_process(workspace.session_id, payload)
+ await workspace.box_service.start_managed_process(
+ workspace.execution_context,
+ workspace.session_id,
+ payload,
+ )
except Exception:
self.owner.error_phase = MCPSessionErrorPhase.PROCESS_START
raise
@@ -260,7 +319,10 @@ class BoxStdioSessionRuntime:
f'process_id={self.process_id} (transport reconnect)'
)
- websocket_url = workspace.get_managed_process_websocket_url(self.process_id)
+ (
+ websocket_url,
+ websocket_headers,
+ ) = await workspace.get_managed_process_websocket_connection(self.process_id)
# Attach the WS transport + MCP session ONCE, on the owner's exit stack,
# in the same task as the serve loop that follows. websocket_client and
@@ -278,7 +340,12 @@ class BoxStdioSessionRuntime:
# attempt re-attaches to the same live process; once it has finished
# cold start the handshake succeeds and stays healthy.
try:
- transport = await self.owner.exit_stack.enter_async_context(websocket_client(websocket_url))
+ transport_context = (
+ authenticated_websocket_client(websocket_url, websocket_headers)
+ if websocket_headers
+ else websocket_client(websocket_url)
+ )
+ transport = await self.owner.exit_stack.enter_async_context(transport_context)
read_stream, write_stream = transport
self.owner.session = await self.owner.exit_stack.enter_async_context(
ClientSession(read_stream, write_stream)
diff --git a/src/langbot/pkg/provider/tools/loaders/native.py b/src/langbot/pkg/provider/tools/loaders/native.py
index 7f5ee4226..5084b347c 100644
--- a/src/langbot/pkg/provider/tools/loaders/native.py
+++ b/src/langbot/pkg/provider/tools/loaders/native.py
@@ -11,6 +11,7 @@ from .. import loader
from ..errors import ToolNotFoundError
from .availability import is_box_backend_available
from . import skill as skill_loader
+from ....api.http.context import ExecutionContext
EXEC_TOOL_NAME = 'exec'
READ_TOOL_NAME = 'read'
@@ -56,6 +57,17 @@ class NativeToolLoader(loader.ToolLoader):
"""Check if the box backend is truly available (not just the runtime)."""
return await is_box_backend_available(self.ap)
+ @staticmethod
+ def _execution_context(query: pipeline_query.Query) -> ExecutionContext:
+ return ExecutionContext(
+ instance_uuid=str(getattr(query, 'instance_uuid', '') or ''),
+ workspace_uuid=str(getattr(query, 'workspace_uuid', '') or ''),
+ placement_generation=getattr(query, 'placement_generation', 0) or 0,
+ bot_uuid=getattr(query, 'bot_uuid', None),
+ pipeline_uuid=getattr(query, 'pipeline_uuid', None),
+ query_uuid=getattr(query, 'query_uuid', None),
+ )
+
async def get_tools(self, bound_plugins: list[str] | None = None) -> list[resource_tool.LLMTool]:
if not self._is_sandbox_available():
return []
@@ -142,7 +154,7 @@ class NativeToolLoader(loader.ToolLoader):
result = self._normalize_exec_result(result)
if selected_skill is not None:
- self._refresh_skill_from_disk(selected_skill)
+ self._refresh_skill_from_disk(query, selected_skill)
return result
def _resolve_host_path(
@@ -162,7 +174,11 @@ class NativeToolLoader(loader.ToolLoader):
)
box_service = self.ap.box_service
- host_root = selected_skill.get('package_root') if selected_skill is not None else box_service.default_workspace
+ host_root = (
+ selected_skill.get('package_root')
+ if selected_skill is not None
+ else box_service._tenant_workspace(self._execution_context(query))
+ )
if not host_root:
raise ValueError('No host workspace configured for file operations.')
@@ -522,11 +538,19 @@ else:
return self._read_text_file_preview(host_path, parameters)
try:
- result = await self.ap.box_service.read_skill_file(selected_skill['name'], relative)
+ result = await self.ap.box_service.read_skill_file(
+ self._execution_context(query),
+ selected_skill['name'],
+ relative,
+ )
return self._build_read_result_from_text(str(result.get('content', '')), parameters)
except Exception:
try:
- result = await self.ap.box_service.list_skill_files(selected_skill['name'], relative)
+ result = await self.ap.box_service.list_skill_files(
+ self._execution_context(query),
+ selected_skill['name'],
+ relative,
+ )
entries = [entry['name'] for entry in result.get('entries', [])]
return self._build_directory_result(entries)
except Exception as exc:
@@ -562,8 +586,9 @@ else:
if encoding != 'text':
return {'ok': False, 'error': 'base64 writes to skill packages are not supported.'}
selected_skill, relative = skill_request
- await self.ap.box_service.write_skill_file(selected_skill['name'], relative, content)
- await self.ap.skill_mgr.reload_skills()
+ execution_context = self._execution_context(query)
+ await self.ap.box_service.write_skill_file(execution_context, selected_skill['name'], relative, content)
+ await self.ap.skill_mgr.reload_skills(execution_context)
return {'ok': True, 'path': path}
host_path, selected_skill = self._resolve_host_path(
@@ -579,7 +604,7 @@ else:
self._write_host_file(host_path, content, parameters)
except ValueError as exc:
return {'ok': False, 'error': str(exc)}
- self._refresh_skill_from_disk(selected_skill)
+ self._refresh_skill_from_disk(query, selected_skill)
return {'ok': True, 'path': path}
async def _invoke_edit(self, parameters: dict, query: pipeline_query.Query) -> dict:
@@ -603,7 +628,11 @@ else:
):
selected_skill, relative = skill_request
try:
- result = await self.ap.box_service.read_skill_file(selected_skill['name'], relative)
+ result = await self.ap.box_service.read_skill_file(
+ self._execution_context(query),
+ selected_skill['name'],
+ relative,
+ )
except Exception:
return {'ok': False, 'error': f'File not found: {path}'}
content = result.get('content', '')
@@ -613,8 +642,14 @@ else:
if count > 1:
return {'ok': False, 'error': f'old_string matches {count} locations; provide a more unique string.'}
new_content = content.replace(old_string, new_string, 1)
- await self.ap.box_service.write_skill_file(selected_skill['name'], relative, new_content)
- await self.ap.skill_mgr.reload_skills()
+ execution_context = self._execution_context(query)
+ await self.ap.box_service.write_skill_file(
+ execution_context,
+ selected_skill['name'],
+ relative,
+ new_content,
+ )
+ await self.ap.skill_mgr.reload_skills(execution_context)
return {'ok': True, 'path': path}
host_path, selected_skill = self._resolve_host_path(
@@ -637,10 +672,10 @@ else:
new_content = content.replace(old_string, new_string, 1)
with open(host_path, 'w', encoding='utf-8') as f:
f.write(new_content)
- self._refresh_skill_from_disk(selected_skill)
+ self._refresh_skill_from_disk(query, selected_skill)
return {'ok': True, 'path': path}
- def _refresh_skill_from_disk(self, selected_skill: dict | None) -> None:
+ def _refresh_skill_from_disk(self, query: pipeline_query.Query, selected_skill: dict | None) -> None:
if selected_skill is None:
return
@@ -650,7 +685,7 @@ else:
refresh_skill = getattr(skill_mgr, 'refresh_skill_from_disk', None)
if callable(refresh_skill):
- refresh_skill(selected_skill.get('name', ''))
+ refresh_skill(self._execution_context(query), selected_skill.get('name', ''))
def _is_sandbox_available(self) -> bool:
"""Check if sandbox backend is available.
diff --git a/src/langbot/pkg/provider/tools/loaders/plugin.py b/src/langbot/pkg/provider/tools/loaders/plugin.py
index baac91d1d..80595049c 100644
--- a/src/langbot/pkg/provider/tools/loaders/plugin.py
+++ b/src/langbot/pkg/provider/tools/loaders/plugin.py
@@ -67,7 +67,11 @@ class PluginToolLoader(loader.ToolLoader):
async def invoke_tool(self, name: str, parameters: dict, query: pipeline_query.Query) -> typing.Any:
try:
return await self.ap.plugin_connector.call_tool(
- name, parameters, session=query.session, query_id=query.query_id
+ name,
+ parameters,
+ session=query.session,
+ query_id=query.query_id,
+ query_uuid=query.query_uuid,
)
except Exception as e:
self.ap.logger.error(f'执行函数 {name} 时发生错误: {e}')
diff --git a/src/langbot/pkg/provider/tools/loaders/skill.py b/src/langbot/pkg/provider/tools/loaders/skill.py
index b62f3e7d5..489cc2a1d 100644
--- a/src/langbot/pkg/provider/tools/loaders/skill.py
+++ b/src/langbot/pkg/provider/tools/loaders/skill.py
@@ -4,6 +4,7 @@ import re
import typing
from ....box import workspace as box_workspace
+from ....api.http.context import ExecutionContext
if typing.TYPE_CHECKING:
from ....core import app
@@ -36,7 +37,15 @@ def get_visible_skills(ap: app.Application, query: pipeline_query.Query) -> dict
if skill_mgr is None:
return {}
- visible_skills = getattr(skill_mgr, 'skills', {})
+ execution_context = ExecutionContext(
+ instance_uuid=str(getattr(query, 'instance_uuid', '') or ''),
+ workspace_uuid=str(getattr(query, 'workspace_uuid', '') or ''),
+ placement_generation=getattr(query, 'placement_generation', 0) or 0,
+ bot_uuid=getattr(query, 'bot_uuid', None),
+ pipeline_uuid=getattr(query, 'pipeline_uuid', None),
+ query_uuid=getattr(query, 'query_uuid', None),
+ )
+ visible_skills = skill_mgr.get_skills(execution_context)
bound_skills = get_bound_skill_names(query)
if bound_skills is None:
return visible_skills
diff --git a/src/langbot/pkg/provider/tools/loaders/skill_authoring.py b/src/langbot/pkg/provider/tools/loaders/skill_authoring.py
index d53721785..25c1c0249 100644
--- a/src/langbot/pkg/provider/tools/loaders/skill_authoring.py
+++ b/src/langbot/pkg/provider/tools/loaders/skill_authoring.py
@@ -7,6 +7,7 @@ import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
from .. import loader
from .availability import is_box_backend_available
+from ....api.http.context import ExecutionContext
# Align with Claude Code's Skill tool design:
# - activate: Activate a skill via Tool Call, returns SKILL.md content
@@ -67,7 +68,7 @@ class SkillToolLoader(loader.ToolLoader):
if name == ACTIVATE_SKILL_TOOL_NAME:
return await self._invoke_activate_skill(parameters, query)
if name == REGISTER_SKILL_TOOL_NAME:
- return await self._invoke_register_skill(parameters)
+ return await self._invoke_register_skill(parameters, query)
raise ValueError(f'Unknown skill tool: {name}')
async def shutdown(self):
@@ -120,7 +121,7 @@ class SkillToolLoader(loader.ToolLoader):
'content': result_content,
}
- async def _invoke_register_skill(self, parameters: dict) -> typing.Any:
+ async def _invoke_register_skill(self, parameters: dict, query) -> typing.Any:
"""Register a skill from sandbox directory to data/skills/."""
sandbox_path = str(parameters.get('path', '') or '').strip()
if not sandbox_path:
@@ -135,7 +136,15 @@ class SkillToolLoader(loader.ToolLoader):
raise ValueError('Skill service not available')
# Scan and register the skill
- scanned = await skill_service.scan_directory_async(host_path)
+ execution_context = ExecutionContext(
+ instance_uuid=str(getattr(query, 'instance_uuid', '') or ''),
+ workspace_uuid=str(getattr(query, 'workspace_uuid', '') or ''),
+ placement_generation=getattr(query, 'placement_generation', 0) or 0,
+ bot_uuid=getattr(query, 'bot_uuid', None),
+ pipeline_uuid=getattr(query, 'pipeline_uuid', None),
+ query_uuid=getattr(query, 'query_uuid', None),
+ )
+ scanned = await skill_service.scan_directory_async(execution_context, host_path)
# Override name if provided
skill_name = str(parameters.get('name') or scanned['name']).strip()
@@ -144,13 +153,14 @@ class SkillToolLoader(loader.ToolLoader):
# Create the skill
created = await skill_service.create_skill(
+ execution_context,
{
'name': skill_name,
'display_name': str(parameters.get('display_name') or scanned.get('display_name', '')).strip(),
'description': str(parameters.get('description') or scanned.get('description', '')).strip(),
'instructions': str(parameters.get('instructions') or scanned.get('instructions', '')),
'package_root': host_path,
- }
+ },
)
return {
diff --git a/src/langbot/pkg/provider/tools/toolmgr.py b/src/langbot/pkg/provider/tools/toolmgr.py
index 60a16ce7e..777d8d1c8 100644
--- a/src/langbot/pkg/provider/tools/toolmgr.py
+++ b/src/langbot/pkg/provider/tools/toolmgr.py
@@ -9,6 +9,8 @@ from langbot_plugin.api.entities.events import pipeline_query
from . import loader as tool_loader
from .errors import ToolNotFoundError
+from ...pipeline.pool import get_query_execution_context
+from ...api.http.service.tenant import TenantContext
if TYPE_CHECKING:
from ...core import app
@@ -57,6 +59,7 @@ class ToolManager:
async def get_all_tools(
self,
+ context: TenantContext,
bound_plugins: list[str] | None = None,
bound_mcp_servers: list[str] | None = None,
include_skill_authoring: bool = False,
@@ -70,6 +73,7 @@ class ToolManager:
all_functions.extend(await self.plugin_tool_loader.get_tools(bound_plugins))
all_functions.extend(
await self.mcp_tool_loader.get_tools(
+ context,
bound_mcp_servers,
include_resource_tools=include_mcp_resource_tools,
)
@@ -79,6 +83,7 @@ class ToolManager:
async def get_tool_catalog(
self,
+ context: TenantContext,
bound_plugins: list[str] | None = None,
bound_mcp_servers: list[str] | None = None,
include_skill_authoring: bool = False,
@@ -106,6 +111,7 @@ class ToolManager:
if self.mcp_tool_loader:
for item in await self.mcp_tool_loader.get_tool_catalog(
+ context,
bound_mcp_servers,
include_resource_tools=include_mcp_resource_tools,
):
@@ -113,19 +119,18 @@ class ToolManager:
return catalog
- async def get_tool_by_name(self, name: str) -> tool_loader.ToolLookupResult | None:
+ async def get_tool_by_name(self, context: TenantContext, name: str) -> tool_loader.ToolLookupResult | None:
"""Get tool by name from any active loader."""
for active_loader in (
self.native_tool_loader,
self.plugin_tool_loader,
- self.mcp_tool_loader,
self.skill_tool_loader,
):
tool = await active_loader.get_tool(name)
if tool:
return tool
- return None
+ return await self.mcp_tool_loader.get_tool(context, name)
async def generate_tools_for_openai(self, use_funcs: list[resource_tool.LLMTool]) -> list:
tools = []
@@ -175,6 +180,7 @@ class ToolManager:
try:
await monitoring_service.record_tool_call(
+ get_query_execution_context(query),
tool_name=name,
tool_source=source,
duration=duration_ms,
@@ -249,7 +255,8 @@ class ToolManager:
query=query,
invoke=lambda: self.plugin_tool_loader.invoke_tool(name, parameters, query),
)
- if await self.mcp_tool_loader.has_tool(name):
+ execution_context = get_query_execution_context(query)
+ if await self.mcp_tool_loader.has_tool(execution_context, name):
telemetry_features.increment(query, 'tool_calls', 'mcp')
return await self._invoke_tool_with_monitoring(
source='mcp',
diff --git a/src/langbot/pkg/rag/knowledge/base.py b/src/langbot/pkg/rag/knowledge/base.py
index 28d010fef..16d383989 100644
--- a/src/langbot/pkg/rag/knowledge/base.py
+++ b/src/langbot/pkg/rag/knowledge/base.py
@@ -5,6 +5,7 @@ from __future__ import annotations
import abc
from langbot.pkg.core import app
+from langbot.pkg.api.http.context import ExecutionContext
from langbot_plugin.api.entities.builtin.rag import context as rag_context
@@ -22,10 +23,16 @@ class KnowledgeBaseInterface(metaclass=abc.ABCMeta):
pass
@abc.abstractmethod
- async def retrieve(self, query: str, settings: dict | None = None) -> list[rag_context.RetrievalResultEntry]:
+ async def retrieve(
+ self,
+ execution_context: ExecutionContext,
+ query: str,
+ settings: dict | None = None,
+ ) -> list[rag_context.RetrievalResultEntry]:
"""Retrieve relevant documents from the knowledge base
Args:
+ execution_context: Trusted active Workspace placement.
query: The query string
settings: Optional per-request retrieval settings overrides
@@ -50,6 +57,6 @@ class KnowledgeBaseInterface(metaclass=abc.ABCMeta):
pass
@abc.abstractmethod
- async def dispose(self):
+ async def dispose(self, execution_context: ExecutionContext):
"""Clean up resources"""
pass
diff --git a/src/langbot/pkg/rag/knowledge/kbmgr.py b/src/langbot/pkg/rag/knowledge/kbmgr.py
index cd37994c4..5d03b565a 100644
--- a/src/langbot/pkg/rag/knowledge/kbmgr.py
+++ b/src/langbot/pkg/rag/knowledge/kbmgr.py
@@ -1,18 +1,22 @@
from __future__ import annotations
+import io
import mimetypes
import os.path
import traceback
import uuid
import zipfile
-import io
from typing import Any
-from langbot.pkg.core import app
+
import sqlalchemy
-
-
-from langbot.pkg.entity.persistence import rag as persistence_rag
-from langbot.pkg.core import taskmgr
from langbot_plugin.api.entities.builtin.rag import context as rag_context
+
+from langbot.pkg.api.http.authz import WorkspaceRequiredError
+from langbot.pkg.api.http.context import ExecutionContext, RequestContext
+from langbot.pkg.api.http.service.tenant import TenantContext, require_workspace_uuid
+from langbot.pkg.core import app, taskmgr
+from langbot.pkg.entity.persistence import rag as persistence_rag
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
+
from .base import KnowledgeBaseInterface
@@ -21,20 +25,78 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
knowledge_base_entity: persistence_rag.KnowledgeBase
- def __init__(self, ap: app.Application, knowledge_base_entity: persistence_rag.KnowledgeBase):
+ def __init__(
+ self,
+ ap: app.Application,
+ knowledge_base_entity: persistence_rag.KnowledgeBase,
+ execution_context: ExecutionContext,
+ ):
super().__init__(ap)
self.knowledge_base_entity = knowledge_base_entity
+ self.execution_context = execution_context
async def initialize(self):
pass
+ async def _assert_execution_context(self, execution_context: ExecutionContext) -> None:
+ """Reject stale or cross-Workspace runtime access."""
+
+ if not isinstance(execution_context, ExecutionContext):
+ raise WorkspaceRequiredError('ExecutionContext is required for knowledge runtime access')
+ if (
+ execution_context.instance_uuid != self.execution_context.instance_uuid
+ or execution_context.workspace_uuid != self.execution_context.workspace_uuid
+ or execution_context.placement_generation != self.execution_context.placement_generation
+ ):
+ raise WorkspaceNotFoundError('Knowledge base not found')
+ if self.knowledge_base_entity.workspace_uuid != execution_context.workspace_uuid:
+ raise WorkspaceNotFoundError('Knowledge base not found')
+ binding = await self.ap.workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ if binding.instance_uuid != execution_context.instance_uuid:
+ raise WorkspaceNotFoundError('Knowledge base not found')
+
+ async def _require_plugin_runtime_context(
+ self,
+ execution_context: ExecutionContext,
+ ) -> ExecutionContext:
+ """Fence every singleton Plugin Runtime call to this runtime KB."""
+
+ await self._assert_execution_context(execution_context)
+ return await self.ap.plugin_connector.require_workspace_context(execution_context)
+
+ def _require_upload_object_key(
+ self,
+ execution_context: ExecutionContext,
+ object_key: str,
+ ) -> None:
+ """Reject raw, cross-Workspace, stale, or non-upload object keys."""
+
+ try:
+ self.ap.storage_mgr.require_scoped_object_key(
+ execution_context,
+ object_key,
+ expected_owner_type='upload_document',
+ )
+ except (WorkspaceRequiredError, ValueError) as exc:
+ raise WorkspaceNotFoundError('Upload not found') from exc
+
async def _store_file_task(
- self, file: persistence_rag.File, task_context: taskmgr.TaskContext, parser_plugin_id: str | None = None
+ self,
+ execution_context: ExecutionContext,
+ file: persistence_rag.File,
+ task_context: taskmgr.TaskContext,
+ parser_plugin_id: str | None = None,
):
+ await self._assert_execution_context(execution_context)
+ self._require_upload_object_key(execution_context, file.file_name)
try:
# set file status to processing
await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(persistence_rag.File)
+ .where(persistence_rag.File.workspace_uuid == execution_context.workspace_uuid)
.where(persistence_rag.File.uuid == file.uuid)
.values(status='processing')
)
@@ -42,7 +104,11 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
task_context.set_current_action('Processing file')
# Get file size from storage
- file_size = await self.ap.storage_mgr.storage_provider.size(file.file_name)
+ file_size = await self.ap.storage_mgr.size_scoped_object_key(
+ execution_context,
+ file.file_name,
+ expected_owner_type='upload_document',
+ )
# Detect MIME type from extension
mime_type, _ = mimetypes.guess_type(file.file_name)
@@ -53,16 +119,22 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
parsed_content = None
if parser_plugin_id:
task_context.set_current_action('Parsing file')
- file_bytes = await self.ap.storage_mgr.storage_provider.load(file.file_name)
+ file_bytes = await self.ap.storage_mgr.load_scoped_object_key(
+ execution_context,
+ file.file_name,
+ expected_owner_type='upload_document',
+ )
parse_context = {
'mime_type': mime_type,
'filename': file.file_name,
'metadata': {},
}
+ await self._require_plugin_runtime_context(execution_context)
parsed_content = await self.ap.plugin_connector.call_parser(parser_plugin_id, parse_context, file_bytes)
# Call plugin to ingest document
result = await self._ingest_document(
+ execution_context,
{
'document_id': file.uuid,
'filename': file.file_name,
@@ -80,8 +152,10 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
raise Exception(error_msg)
# set file status to completed
+ await self._assert_execution_context(execution_context)
await self.ap.persistence_mgr.execute_async(
sqlalchemy.update(persistence_rag.File)
+ .where(persistence_rag.File.workspace_uuid == execution_context.workspace_uuid)
.where(persistence_rag.File.uuid == file.uuid)
.values(status='completed')
)
@@ -89,35 +163,63 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
except Exception as e:
self.ap.logger.error(f'Error storing file {file.uuid}: {e}')
traceback.print_exc()
- # set file status to failed
- await self.ap.persistence_mgr.execute_async(
- sqlalchemy.update(persistence_rag.File)
- .where(persistence_rag.File.uuid == file.uuid)
- .values(status='failed')
- )
+ # A stale placement is fenced from all writes, including failure
+ # status updates from an old background task.
+ try:
+ await self._assert_execution_context(execution_context)
+ except Exception:
+ self.ap.logger.warning(f'Skipping stale RAG task status update for file {file.uuid}')
+ else:
+ await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.update(persistence_rag.File)
+ .where(persistence_rag.File.workspace_uuid == execution_context.workspace_uuid)
+ .where(persistence_rag.File.uuid == file.uuid)
+ .values(status='failed')
+ )
raise
finally:
- # delete file from storage
- await self.ap.storage_mgr.storage_provider.delete(file.file_name)
+ # An old background task must not touch an upload after its
+ # placement generation has been fenced off.
+ try:
+ await self._assert_execution_context(execution_context)
+ await self.ap.storage_mgr.delete_scoped_object_key(
+ execution_context,
+ file.file_name,
+ expected_owner_type='upload_document',
+ )
+ except (WorkspaceRequiredError, WorkspaceNotFoundError):
+ self.ap.logger.warning(f'Skipping stale RAG upload cleanup for file {file.uuid}')
- async def store_file(self, file_id: str, parser_plugin_id: str | None = None) -> str:
+ async def store_file(
+ self,
+ execution_context: ExecutionContext,
+ file_id: str,
+ parser_plugin_id: str | None = None,
+ ) -> str:
+ await self._assert_execution_context(execution_context)
+ self._require_upload_object_key(execution_context, file_id)
# pre checking
- if not await self.ap.storage_mgr.storage_provider.exists(file_id):
- raise Exception(f'File {file_id} not found')
+ if not await self.ap.storage_mgr.exists_scoped_object_key(
+ execution_context,
+ file_id,
+ expected_owner_type='upload_document',
+ ):
+ raise WorkspaceNotFoundError('Upload not found')
file_name = file_id
_, ext = os.path.splitext(file_name)
extension = ext.lstrip('.').lower() if ext else ''
if extension == 'zip':
- return await self._store_zip_file(file_id, parser_plugin_id=parser_plugin_id)
+ return await self._store_zip_file(execution_context, file_id, parser_plugin_id=parser_plugin_id)
file_uuid = str(uuid.uuid4())
kb_id = self.knowledge_base_entity.uuid
file_obj_data = {
'uuid': file_uuid,
+ 'workspace_uuid': execution_context.workspace_uuid,
'kb_id': kb_id,
'file_name': file_name,
'extension': extension,
@@ -131,19 +233,38 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
# run background task asynchronously
ctx = taskmgr.TaskContext.new()
wrapper = self.ap.task_mgr.create_user_task(
- self._store_file_task(file_obj, task_context=ctx, parser_plugin_id=parser_plugin_id),
+ self._store_file_task(
+ execution_context,
+ file_obj,
+ task_context=ctx,
+ parser_plugin_id=parser_plugin_id,
+ ),
kind='knowledge-operation',
name=f'knowledge-store-file-{file_id}',
label=f'Store file {file_id}',
context=ctx,
+ instance_uuid=execution_context.instance_uuid,
+ workspace_uuid=execution_context.workspace_uuid,
+ placement_generation=execution_context.placement_generation,
)
return wrapper.id
- async def _store_zip_file(self, zip_file_id: str, parser_plugin_id: str | None = None) -> str:
+ async def _store_zip_file(
+ self,
+ execution_context: ExecutionContext,
+ zip_file_id: str,
+ parser_plugin_id: str | None = None,
+ ) -> str:
"""Handle ZIP file by extracting each document and storing them separately."""
+ await self._assert_execution_context(execution_context)
+ self._require_upload_object_key(execution_context, zip_file_id)
self.ap.logger.info(f'Processing ZIP file: {zip_file_id}')
- zip_bytes = await self.ap.storage_mgr.storage_provider.load(zip_file_id)
+ zip_bytes = await self.ap.storage_mgr.load_scoped_object_key(
+ execution_context,
+ zip_file_id,
+ expected_owner_type='upload_document',
+ )
supported_extensions = {'txt', 'pdf', 'docx', 'md', 'html'}
stored_file_tasks = []
@@ -173,15 +294,31 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
continue
extracted_file_id = file_stem + '_' + str(uuid.uuid4())[:8] + '.' + extension
- # save file to storage
+ extracted_object_key = await self.ap.storage_mgr.save_scoped(
+ execution_context,
+ owner_type='upload_document',
+ owner=f'knowledge-base:{self.knowledge_base_entity.uuid}',
+ key=extracted_file_id,
+ value=file_content,
+ )
- await self.ap.storage_mgr.storage_provider.save(extracted_file_id, file_content)
-
- task_id = await self.store_file(extracted_file_id, parser_plugin_id=parser_plugin_id)
+ try:
+ task_id = await self.store_file(
+ execution_context,
+ extracted_object_key,
+ parser_plugin_id=parser_plugin_id,
+ )
+ except Exception:
+ await self.ap.storage_mgr.delete_scoped_object_key(
+ execution_context,
+ extracted_object_key,
+ expected_owner_type='upload_document',
+ )
+ raise
stored_file_tasks.append(task_id)
self.ap.logger.info(
- f'Extracted and stored file from ZIP: {file_info.filename} -> {extracted_file_id}'
+ f'Extracted and stored file from ZIP: {file_info.filename} -> {extracted_object_key}'
)
except Exception as e:
@@ -197,20 +334,33 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
return stored_file_tasks[0] if stored_file_tasks else ''
finally:
try:
- await self.ap.storage_mgr.storage_provider.delete(zip_file_id)
+ await self._assert_execution_context(execution_context)
+ await self.ap.storage_mgr.delete_scoped_object_key(
+ execution_context,
+ zip_file_id,
+ expected_owner_type='upload_document',
+ )
except FileNotFoundError:
pass
+ except (WorkspaceRequiredError, WorkspaceNotFoundError):
+ self.ap.logger.warning(f'Skipping stale RAG ZIP cleanup for upload {zip_file_id}')
except Exception as e:
self.ap.logger.warning(f'Failed to cleanup ZIP file {zip_file_id}: {e}')
- async def retrieve(self, query: str, settings: dict | None = None) -> list[rag_context.RetrievalResultEntry]:
+ async def retrieve(
+ self,
+ execution_context: ExecutionContext,
+ query: str,
+ settings: dict | None = None,
+ ) -> list[rag_context.RetrievalResultEntry]:
+ await self._assert_execution_context(execution_context)
# Merge stored retrieval_settings with per-request overrides
stored = self.knowledge_base_entity.retrieval_settings or {}
merged = {**stored, **(settings or {})}
if 'top_k' not in merged:
merged['top_k'] = 5 # fallback default
- response = await self._retrieve(query, merged)
+ response = await self._retrieve(execution_context, query, merged)
results_data = response.get('results', [])
entries = []
@@ -221,12 +371,25 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
entries.append(r)
return entries
- async def delete_file(self, file_id: str):
- await self._delete_document(file_id)
+ async def delete_file(self, execution_context: ExecutionContext, file_id: str):
+ await self._assert_execution_context(execution_context)
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(persistence_rag.File.uuid)
+ .where(persistence_rag.File.workspace_uuid == execution_context.workspace_uuid)
+ .where(persistence_rag.File.kb_id == self.knowledge_base_entity.uuid)
+ .where(persistence_rag.File.uuid == file_id)
+ .limit(1)
+ )
+ if result.first() is None:
+ raise WorkspaceNotFoundError('Knowledge file not found')
+ await self._delete_document(execution_context, file_id)
# Also cleanup DB record
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_rag.File).where(persistence_rag.File.uuid == file_id)
+ sqlalchemy.delete(persistence_rag.File)
+ .where(persistence_rag.File.workspace_uuid == execution_context.workspace_uuid)
+ .where(persistence_rag.File.kb_id == self.knowledge_base_entity.uuid)
+ .where(persistence_rag.File.uuid == file_id)
)
def get_uuid(self) -> str:
@@ -241,14 +404,16 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
"""Get the Knowledge Engine plugin ID"""
return self.knowledge_base_entity.knowledge_engine_plugin_id or ''
- async def dispose(self):
+ async def dispose(self, execution_context: ExecutionContext):
"""Dispose the knowledge base, notifying the plugin to cleanup."""
- await self._on_kb_delete()
+ await self._assert_execution_context(execution_context)
+ await self._on_kb_delete(execution_context)
# ========== Plugin Communication Methods ==========
- async def _on_kb_create(self) -> None:
+ async def _on_kb_create(self, execution_context: ExecutionContext) -> None:
"""Notify plugin about KB creation."""
+ await self._assert_execution_context(execution_context)
plugin_id = self.knowledge_base_entity.knowledge_engine_plugin_id
if not plugin_id:
return
@@ -258,17 +423,20 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
self.ap.logger.info(
f'Calling RAG plugin {plugin_id}: on_knowledge_base_create(kb_id={self.knowledge_base_entity.uuid})'
)
+ await self._require_plugin_runtime_context(execution_context)
await self.ap.plugin_connector.rag_on_kb_create(plugin_id, self.knowledge_base_entity.uuid, config)
except Exception as e:
self.ap.logger.error(f'Failed to notify plugin {plugin_id} on KB create: {e}')
raise
- async def _on_kb_delete(self) -> None:
+ async def _on_kb_delete(self, execution_context: ExecutionContext) -> None:
"""Notify plugin about KB deletion."""
+ await self._assert_execution_context(execution_context)
plugin_id = self.knowledge_base_entity.knowledge_engine_plugin_id
if not plugin_id:
return
+ await self._require_plugin_runtime_context(execution_context)
try:
self.ap.logger.info(
f'Calling RAG plugin {plugin_id}: on_knowledge_base_delete(kb_id={self.knowledge_base_entity.uuid})'
@@ -279,11 +447,13 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
async def _ingest_document(
self,
+ execution_context: ExecutionContext,
file_metadata: dict[str, Any],
storage_path: str,
parsed_content: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Call plugin to ingest document."""
+ await self._assert_execution_context(execution_context)
kb = self.knowledge_base_entity
plugin_id = kb.knowledge_engine_plugin_id
if not plugin_id:
@@ -306,6 +476,7 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
'parsed_content': parsed_content,
}
+ await self._require_plugin_runtime_context(execution_context)
try:
result = await self.ap.plugin_connector.call_rag_ingest(plugin_id, context_data)
return result
@@ -315,6 +486,7 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
async def _retrieve(
self,
+ execution_context: ExecutionContext,
query: str,
settings: dict[str, Any],
) -> dict[str, Any]:
@@ -324,6 +496,7 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
ValueError: If no RAG plugin is configured for this KB.
Exception: If the plugin retrieval call fails.
"""
+ await self._assert_execution_context(execution_context)
kb = self.knowledge_base_entity
plugin_id = kb.knowledge_engine_plugin_id
if not plugin_id:
@@ -333,25 +506,28 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
# for plugins that need it. Do NOT move them into filters, as filters
# are passed directly to vector_search by some plugins (e.g. LangRAG)
# and would cause empty results when the metadata field doesn't exist.
- filters = settings.pop('filters', {})
+ plugin_settings = dict(settings)
+ filters = plugin_settings.pop('filters', {})
retrieval_context = {
'query': query,
'knowledge_base_id': kb.uuid,
'collection_id': kb.collection_id or kb.uuid,
- 'retrieval_settings': settings,
+ 'retrieval_settings': plugin_settings,
'creation_settings': kb.creation_settings or {},
'filters': filters,
}
+ await self._require_plugin_runtime_context(execution_context)
result = await self.ap.plugin_connector.call_rag_retrieve(
plugin_id,
retrieval_context,
)
return result
- async def _delete_document(self, document_id: str) -> bool:
+ async def _delete_document(self, execution_context: ExecutionContext, document_id: str) -> bool:
"""Call plugin to delete document."""
+ await self._assert_execution_context(execution_context)
kb = self.knowledge_base_entity
plugin_id = kb.knowledge_engine_plugin_id
if not plugin_id:
@@ -359,6 +535,7 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
self.ap.logger.info(f'Calling RAG plugin {plugin_id}: delete_document(doc_id={document_id})')
+ await self._require_plugin_runtime_context(execution_context)
try:
return await self.ap.plugin_connector.call_rag_delete_document(plugin_id, document_id, kb.uuid)
except Exception as e:
@@ -369,7 +546,7 @@ class RuntimeKnowledgeBase(KnowledgeBaseInterface):
class RAGManager:
ap: app.Application
- knowledge_bases: dict[str, KnowledgeBaseInterface]
+ knowledge_bases: dict[tuple[str, str], RuntimeKnowledgeBase]
def __init__(self, ap: app.Application):
self.ap = ap
@@ -378,20 +555,50 @@ class RAGManager:
async def initialize(self):
await self.load_knowledge_bases_from_db()
- async def get_all_knowledge_base_details(self) -> list[dict]:
+ async def _to_execution_context(
+ self,
+ context: RequestContext | ExecutionContext,
+ ) -> ExecutionContext:
+ if isinstance(context, RequestContext):
+ execution_context = ExecutionContext.from_request(context)
+ elif isinstance(context, ExecutionContext):
+ execution_context = context
+ else:
+ raise WorkspaceRequiredError('RequestContext or ExecutionContext is required')
+
+ binding = await self.ap.workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ if binding.instance_uuid != execution_context.instance_uuid:
+ raise WorkspaceNotFoundError('Workspace not found')
+ return execution_context
+
+ async def _get_engine_map(self, context: TenantContext) -> dict[str, dict]:
+ engine_map: dict[str, dict] = {}
+ connector = getattr(self.ap, 'plugin_connector', None)
+ if connector is not None and connector.is_enable_plugin:
+ await connector.require_workspace_context(context)
+ try:
+ engines = await connector.list_knowledge_engines()
+ engine_map = {engine['plugin_id']: engine for engine in engines}
+ except Exception as e:
+ self.ap.logger.warning(f'Failed to list Knowledge Engines: {e}')
+ return engine_map
+
+ async def get_all_knowledge_base_details(self, context: TenantContext) -> list[dict]:
"""Get all knowledge bases with enriched Knowledge Engine details."""
+ workspace_uuid = require_workspace_uuid(context)
# 1. Get raw KBs from DB
- result = await self.ap.persistence_mgr.execute_async(sqlalchemy.select(persistence_rag.KnowledgeBase))
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(persistence_rag.KnowledgeBase).where(
+ persistence_rag.KnowledgeBase.workspace_uuid == workspace_uuid
+ )
+ )
knowledge_bases = result.all()
# 2. Get all available Knowledge Engines for enrichment
- engine_map = {}
- if self.ap.plugin_connector.is_enable_plugin:
- try:
- engines = await self.ap.plugin_connector.list_knowledge_engines()
- engine_map = {e['plugin_id']: e for e in engines}
- except Exception as e:
- self.ap.logger.warning(f'Failed to list Knowledge Engines: {e}')
+ engine_map = await self._get_engine_map(context)
# 3. Serialize and enrich
kb_list = []
@@ -402,10 +609,13 @@ class RAGManager:
return kb_list
- async def get_knowledge_base_details(self, kb_uuid: str) -> dict | None:
+ async def get_knowledge_base_details(self, context: TenantContext, kb_uuid: str) -> dict | None:
"""Get specific knowledge base with enriched Knowledge Engine details."""
+ workspace_uuid = require_workspace_uuid(context)
result = await self.ap.persistence_mgr.execute_async(
- sqlalchemy.select(persistence_rag.KnowledgeBase).where(persistence_rag.KnowledgeBase.uuid == kb_uuid)
+ sqlalchemy.select(persistence_rag.KnowledgeBase)
+ .where(persistence_rag.KnowledgeBase.workspace_uuid == workspace_uuid)
+ .where(persistence_rag.KnowledgeBase.uuid == kb_uuid)
)
kb = result.first()
if not kb:
@@ -414,13 +624,7 @@ class RAGManager:
kb_dict = self.ap.persistence_mgr.serialize_model(persistence_rag.KnowledgeBase, kb)
# Fetch engines
- engine_map = {}
- if self.ap.plugin_connector.is_enable_plugin:
- try:
- engines = await self.ap.plugin_connector.list_knowledge_engines()
- engine_map = {e['plugin_id']: e for e in engines}
- except Exception as e:
- self.ap.logger.warning(f'Failed to list Knowledge Engines: {e}')
+ engine_map = await self._get_engine_map(context)
self._enrich_kb_dict(kb_dict, engine_map)
return kb_dict
@@ -465,6 +669,7 @@ class RAGManager:
async def create_knowledge_base(
self,
+ context: RequestContext | ExecutionContext,
name: str,
knowledge_engine_plugin_id: str,
creation_settings: dict,
@@ -472,8 +677,10 @@ class RAGManager:
description: str = '',
) -> persistence_rag.KnowledgeBase:
"""Create a new knowledge base using a RAG plugin."""
+ execution_context = await self._to_execution_context(context)
# Validate that the Knowledge Engine plugin exists
if self.ap.plugin_connector.is_enable_plugin:
+ await self.ap.plugin_connector.require_workspace_context(execution_context)
try:
engines = await self.ap.plugin_connector.list_knowledge_engines()
engine_ids = [e.get('plugin_id') for e in engines]
@@ -490,6 +697,7 @@ class RAGManager:
kb_data = {
'uuid': kb_uuid,
+ 'workspace_uuid': execution_context.workspace_uuid,
'name': name,
'description': description,
'knowledge_engine_plugin_id': knowledge_engine_plugin_id,
@@ -505,15 +713,17 @@ class RAGManager:
await self.ap.persistence_mgr.execute_async(sqlalchemy.insert(persistence_rag.KnowledgeBase).values(kb_data))
# Load into Runtime
- runtime_kb = await self.load_knowledge_base(kb)
+ runtime_kb = await self.load_knowledge_base(execution_context, kb)
# Notify Plugin — rollback DB record and runtime entry on failure
try:
- await runtime_kb._on_kb_create()
+ await runtime_kb._on_kb_create(execution_context)
except Exception:
- self.knowledge_bases.pop(kb_uuid, None)
+ self.knowledge_bases.pop((execution_context.workspace_uuid, kb_uuid), None)
await self.ap.persistence_mgr.execute_async(
- sqlalchemy.delete(persistence_rag.KnowledgeBase).where(persistence_rag.KnowledgeBase.uuid == kb_uuid)
+ sqlalchemy.delete(persistence_rag.KnowledgeBase)
+ .where(persistence_rag.KnowledgeBase.workspace_uuid == execution_context.workspace_uuid)
+ .where(persistence_rag.KnowledgeBase.uuid == kb_uuid)
)
raise
@@ -531,7 +741,13 @@ class RAGManager:
for knowledge_base in knowledge_bases:
try:
- await self.load_knowledge_base(knowledge_base)
+ binding = await self.ap.workspace_service.get_execution_binding(knowledge_base.workspace_uuid)
+ execution_context = ExecutionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ )
+ await self.load_knowledge_base(execution_context, knowledge_base)
except Exception as e:
self.ap.logger.error(
f'Error loading knowledge base {knowledge_base.uuid}: {e}\n{traceback.format_exc()}'
@@ -539,6 +755,7 @@ class RAGManager:
async def load_knowledge_base(
self,
+ context: RequestContext | ExecutionContext,
knowledge_base_entity: persistence_rag.KnowledgeBase | sqlalchemy.Row | dict,
) -> RuntimeKnowledgeBase:
if isinstance(knowledge_base_entity, sqlalchemy.Row):
@@ -551,23 +768,47 @@ class RAGManager:
}
knowledge_base_entity = persistence_rag.KnowledgeBase(**filtered_dict)
- runtime_knowledge_base = RuntimeKnowledgeBase(ap=self.ap, knowledge_base_entity=knowledge_base_entity)
+ execution_context = await self._to_execution_context(context)
+ if knowledge_base_entity.workspace_uuid != execution_context.workspace_uuid:
+ raise WorkspaceNotFoundError('Knowledge base not found')
+ runtime_knowledge_base = RuntimeKnowledgeBase(
+ ap=self.ap,
+ knowledge_base_entity=knowledge_base_entity,
+ execution_context=execution_context,
+ )
await runtime_knowledge_base.initialize()
- self.knowledge_bases[runtime_knowledge_base.get_uuid()] = runtime_knowledge_base
+ self.knowledge_bases[(execution_context.workspace_uuid, runtime_knowledge_base.get_uuid())] = (
+ runtime_knowledge_base
+ )
return runtime_knowledge_base
- async def get_knowledge_base_by_uuid(self, kb_uuid: str) -> KnowledgeBaseInterface | None:
- return self.knowledge_bases.get(kb_uuid)
+ async def get_knowledge_base_by_uuid(
+ self,
+ context: RequestContext | ExecutionContext,
+ kb_uuid: str,
+ ) -> RuntimeKnowledgeBase | None:
+ execution_context = await self._to_execution_context(context)
+ return self.knowledge_bases.get((execution_context.workspace_uuid, kb_uuid))
- async def remove_knowledge_base_from_runtime(self, kb_uuid: str):
- self.knowledge_bases.pop(kb_uuid, None)
+ async def remove_knowledge_base_from_runtime(
+ self,
+ context: RequestContext | ExecutionContext,
+ kb_uuid: str,
+ ) -> None:
+ execution_context = await self._to_execution_context(context)
+ self.knowledge_bases.pop((execution_context.workspace_uuid, kb_uuid), None)
- async def delete_knowledge_base(self, kb_uuid: str):
- kb = self.knowledge_bases.pop(kb_uuid, None)
+ async def delete_knowledge_base(
+ self,
+ context: RequestContext | ExecutionContext,
+ kb_uuid: str,
+ ) -> None:
+ execution_context = await self._to_execution_context(context)
+ kb = self.knowledge_bases.pop((execution_context.workspace_uuid, kb_uuid), None)
if kb is not None:
- await kb.dispose()
+ await kb.dispose(execution_context)
else:
self.ap.logger.warning(f'Knowledge base {kb_uuid} not found in runtime, skipping plugin notification')
diff --git a/src/langbot/pkg/rag/service/runtime.py b/src/langbot/pkg/rag/service/runtime.py
index 0de1ae885..a685dc521 100644
--- a/src/langbot/pkg/rag/service/runtime.py
+++ b/src/langbot/pkg/rag/service/runtime.py
@@ -5,6 +5,13 @@ import re
from typing import TYPE_CHECKING, Any
from urllib.parse import unquote
+import sqlalchemy
+
+from langbot.pkg.api.http.authz import WorkspaceRequiredError
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.entity.persistence import rag as persistence_rag
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
+
if TYPE_CHECKING:
from langbot.pkg.core import app
@@ -19,8 +26,54 @@ class RAGRuntimeService:
def __init__(self, ap: app.Application):
self.ap = ap
+ async def _validate_execution_context(self, execution_context: ExecutionContext) -> None:
+ if not isinstance(execution_context, ExecutionContext):
+ raise WorkspaceRequiredError('ExecutionContext is required for RAG runtime access')
+ if (
+ not execution_context.instance_uuid.strip()
+ or not execution_context.workspace_uuid.strip()
+ or execution_context.placement_generation <= 0
+ ):
+ raise WorkspaceRequiredError('A complete active ExecutionContext is required')
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ if workspace_service is None:
+ raise WorkspaceRequiredError('Workspace execution service is unavailable')
+ binding = await workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ if binding.instance_uuid != execution_context.instance_uuid:
+ raise WorkspaceRequiredError('ExecutionContext belongs to another LangBot instance')
+
+ async def _resolve_knowledge_base_uuid(
+ self,
+ execution_context: ExecutionContext,
+ collection_id: str,
+ ) -> str:
+ """Resolve a plugin logical handle to a Workspace-owned KB UUID."""
+
+ await self._validate_execution_context(execution_context)
+ if not isinstance(collection_id, str) or not collection_id.strip():
+ raise WorkspaceNotFoundError('Knowledge base not found')
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(persistence_rag.KnowledgeBase.uuid)
+ .where(persistence_rag.KnowledgeBase.workspace_uuid == execution_context.workspace_uuid)
+ .where(
+ sqlalchemy.or_(
+ persistence_rag.KnowledgeBase.uuid == collection_id,
+ persistence_rag.KnowledgeBase.collection_id == collection_id,
+ )
+ )
+ .limit(1)
+ )
+ kb_uuid = result.scalar_one_or_none()
+ if kb_uuid is None:
+ raise WorkspaceNotFoundError('Knowledge base not found')
+ return kb_uuid
+
async def vector_upsert(
self,
+ execution_context: ExecutionContext,
collection_id: str,
vectors: list[list[float]],
ids: list[str],
@@ -28,9 +81,17 @@ class RAGRuntimeService:
documents: list[str] | None = None,
) -> None:
"""Handle VECTOR_UPSERT action."""
+ knowledge_base_uuid = await self._resolve_knowledge_base_uuid(execution_context, collection_id)
+ if len(vectors) != len(ids):
+ raise ValueError('vectors and ids must have the same length')
+ if metadata is not None and len(metadata) != len(vectors):
+ raise ValueError('metadata must have the same length as vectors')
+ if documents is not None and len(documents) != len(vectors):
+ raise ValueError('documents must have the same length as vectors')
metadatas = metadata if metadata else [{} for _ in vectors]
await self.ap.vector_db_mgr.upsert(
- collection_name=collection_id,
+ execution_context=execution_context,
+ knowledge_base_uuid=knowledge_base_uuid,
vectors=vectors,
ids=ids,
metadata=metadatas,
@@ -39,6 +100,7 @@ class RAGRuntimeService:
async def vector_search(
self,
+ execution_context: ExecutionContext,
collection_id: str,
query_vector: list[float],
top_k: int,
@@ -48,8 +110,10 @@ class RAGRuntimeService:
vector_weight: float | None = None,
) -> list[dict[str, Any]]:
"""Handle VECTOR_SEARCH action."""
+ knowledge_base_uuid = await self._resolve_knowledge_base_uuid(execution_context, collection_id)
return await self.ap.vector_db_mgr.search(
- collection_name=collection_id,
+ execution_context=execution_context,
+ knowledge_base_uuid=knowledge_base_uuid,
query_vector=query_vector,
limit=top_k,
filter=filters,
@@ -59,7 +123,11 @@ class RAGRuntimeService:
)
async def vector_delete(
- self, collection_id: str, file_ids: list[str] | None = None, filters: dict[str, Any] | None = None
+ self,
+ execution_context: ExecutionContext,
+ collection_id: str,
+ file_ids: list[str] | None = None,
+ filters: dict[str, Any] | None = None,
) -> int:
"""Handle VECTOR_DELETE action.
@@ -73,16 +141,26 @@ class RAGRuntimeService:
in their metadata.
filters: Filter-based deletion (not yet supported, will raise).
"""
+ knowledge_base_uuid = await self._resolve_knowledge_base_uuid(execution_context, collection_id)
count = 0
if file_ids:
- await self.ap.vector_db_mgr.delete_by_file_id(collection_name=collection_id, file_ids=file_ids)
+ await self.ap.vector_db_mgr.delete_by_file_id(
+ execution_context=execution_context,
+ knowledge_base_uuid=knowledge_base_uuid,
+ file_ids=file_ids,
+ )
count = len(file_ids)
elif filters:
- count = await self.ap.vector_db_mgr.delete_by_filter(collection_name=collection_id, filter=filters)
+ count = await self.ap.vector_db_mgr.delete_by_filter(
+ execution_context=execution_context,
+ knowledge_base_uuid=knowledge_base_uuid,
+ filter=filters,
+ )
return count
async def vector_list(
self,
+ execution_context: ExecutionContext,
collection_id: str,
filters: dict[str, Any] | None = None,
limit: int = 20,
@@ -99,14 +177,20 @@ class RAGRuntimeService:
Returns:
Tuple of (items, total).
"""
+ knowledge_base_uuid = await self._resolve_knowledge_base_uuid(execution_context, collection_id)
return await self.ap.vector_db_mgr.list_by_filter(
- collection_name=collection_id,
+ execution_context=execution_context,
+ knowledge_base_uuid=knowledge_base_uuid,
filter=filters,
limit=limit,
offset=offset,
)
- async def get_file_stream(self, storage_path: str) -> bytes:
+ async def get_file_stream(
+ self,
+ execution_context: ExecutionContext,
+ storage_path: str,
+ ) -> bytes:
"""Handle GET_KNOWLEDEGE_FILE_STREAM action.
Uses the storage manager abstraction to load file content,
@@ -125,5 +209,18 @@ class RAGRuntimeService:
or re.match(r'^[A-Za-z]:/', normalized)
):
raise ValueError('Invalid storage path')
- content_bytes = await self.ap.storage_mgr.storage_provider.load(normalized)
+ await self._validate_execution_context(execution_context)
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(persistence_rag.File.uuid)
+ .where(persistence_rag.File.workspace_uuid == execution_context.workspace_uuid)
+ .where(persistence_rag.File.file_name == normalized)
+ .limit(1)
+ )
+ if result.first() is None:
+ raise WorkspaceNotFoundError('Knowledge file not found')
+ content_bytes = await self.ap.storage_mgr.load_scoped_object_key(
+ execution_context,
+ normalized,
+ expected_owner_type='upload_document',
+ )
return content_bytes if content_bytes else b''
diff --git a/src/langbot/pkg/skill/activation.py b/src/langbot/pkg/skill/activation.py
index 706747060..27ce1b3da 100644
--- a/src/langbot/pkg/skill/activation.py
+++ b/src/langbot/pkg/skill/activation.py
@@ -3,6 +3,7 @@ from __future__ import annotations
import typing
from ..provider.tools.loaders import skill as skill_loader
+from ..api.http.context import ExecutionContext
if typing.TYPE_CHECKING:
from ..core import app
@@ -27,7 +28,15 @@ def register_activated_skill(
if skill_mgr is None:
return False
- skill_data = skill_mgr.get_skill_by_name(skill_name)
+ execution_context = ExecutionContext(
+ instance_uuid=str(getattr(query, 'instance_uuid', '') or ''),
+ workspace_uuid=str(getattr(query, 'workspace_uuid', '') or ''),
+ placement_generation=getattr(query, 'placement_generation', 0) or 0,
+ bot_uuid=getattr(query, 'bot_uuid', None),
+ pipeline_uuid=getattr(query, 'pipeline_uuid', None),
+ query_uuid=getattr(query, 'query_uuid', None),
+ )
+ skill_data = skill_mgr.get_skill_by_name(execution_context, skill_name)
if skill_data is None:
return False
diff --git a/src/langbot/pkg/skill/manager.py b/src/langbot/pkg/skill/manager.py
index ddb2125c3..42ce7acda 100644
--- a/src/langbot/pkg/skill/manager.py
+++ b/src/langbot/pkg/skill/manager.py
@@ -1,133 +1,126 @@
from __future__ import annotations
import os
-import typing
+from ..api.http.context import ExecutionContext
+from ..api.http.service.tenant import TenantContext, require_workspace_uuid
from ..core import app
-if typing.TYPE_CHECKING:
- pass
-
class SkillManager:
- """Skill manager backed by Box-managed or local filesystem packages.
-
- In sandbox deployments, skills are loaded from the Box runtime. Local
- data/skills remains as the fallback for non-Box development.
-
- Skills are activated through the `activate` tool (Tool Call mechanism),
- aligned with Claude Code's design. This protects KV Cache and follows
- industry standard.
- """
+ """Workspace-scoped in-memory view of Box-managed skill packages."""
ap: app.Application
- skills: dict[str, dict]
def __init__(self, ap: app.Application):
self.ap = ap
- self.skills = {}
+ self._skills_by_scope: dict[tuple[str, str, int], dict[str, dict]] = {}
+
+ @staticmethod
+ def _execution_context(context: TenantContext) -> ExecutionContext:
+ workspace_uuid = require_workspace_uuid(context)
+ instance_uuid = str(getattr(context, 'instance_uuid', '') or '').strip()
+ generation = getattr(context, 'placement_generation', None)
+ if not instance_uuid or isinstance(generation, bool) or not isinstance(generation, int) or generation <= 0:
+ raise ValueError('Skill cache requires an explicit fenced execution context')
+ return ExecutionContext(
+ instance_uuid=instance_uuid,
+ workspace_uuid=workspace_uuid,
+ placement_generation=generation,
+ bot_uuid=getattr(context, 'bot_uuid', None),
+ pipeline_uuid=getattr(context, 'pipeline_uuid', None),
+ query_uuid=getattr(context, 'query_uuid', None),
+ )
+
+ @classmethod
+ def _scope_key(cls, context: TenantContext) -> tuple[str, str, int]:
+ execution_context = cls._execution_context(context)
+ return (
+ execution_context.instance_uuid,
+ execution_context.workspace_uuid,
+ execution_context.placement_generation,
+ )
async def initialize(self):
- await self.reload_skills()
+ try:
+ binding = await self.ap.workspace_service.get_execution_binding()
+ except Exception:
+ self.ap.logger.info('No unambiguous Workspace binding; skill caches will load on demand.')
+ return
+ await self.reload_skills(
+ ExecutionContext(
+ instance_uuid=binding.instance_uuid,
+ workspace_uuid=binding.workspace_uuid,
+ placement_generation=binding.placement_generation,
+ )
+ )
- async def reload_skills(self):
- """Reload all skills from the Box runtime.
-
- Box is the only source of truth for skills. When Box is unavailable
- (disabled in config or unreachable) the cache is emptied — there is
- no local filesystem fallback. Skills whose ``package_root`` is no
- longer visible on the LangBot-side filesystem are dropped so they
- don't surface as stale ``extra_mounts``.
- """
- self.skills = {}
+ async def reload_skills(self, context: TenantContext) -> None:
+ execution_context = self._execution_context(context)
+ key = self._scope_key(execution_context)
+ self._skills_by_scope[key] = {}
box_service = getattr(self.ap, 'box_service', None)
if box_service is None or not getattr(box_service, 'available', False):
- self.ap.logger.info('Box runtime unavailable; skill cache is empty.')
+ self.ap.logger.info(
+ f'Box runtime unavailable; skill cache is empty for Workspace {execution_context.workspace_uuid}.'
+ )
return
- # LangBot may only validate Box-reported paths against its own
- # filesystem when the two share one (local stdio mode). In separated
- # deployments (Docker Compose, k8s sidecar, --standalone-box, remote
- # endpoint) the package_root lives on the Box runtime's filesystem and
- # is not resolvable here, so we trust what Box reports.
validate_locally = bool(getattr(box_service, 'shares_filesystem_with_box', False))
-
try:
dropped = 0
- for skill_data in await box_service.list_skills():
+ skills: dict[str, dict] = {}
+ for skill_data in await box_service.list_skills(execution_context):
skill_name = skill_data.get('name')
if not skill_name:
continue
package_root = str(skill_data.get('package_root', '') or '').strip()
if validate_locally and package_root and not os.path.isdir(package_root):
self.ap.logger.warning(
- f'Skill "{skill_name}" reported by Box runtime but '
- f'package_root missing on LangBot filesystem '
- f'({package_root}); dropping from in-memory cache.'
+ f'Skill "{skill_name}" reported by Box runtime but package_root '
+ f'missing on LangBot filesystem ({package_root}); dropping from cache.'
)
dropped += 1
continue
- self.skills[skill_name] = skill_data
- if dropped:
- self.ap.logger.warning(
- f'Loaded {len(self.skills)} skills from Box runtime '
- f'({dropped} dropped due to missing package_root).'
- )
- else:
- self.ap.logger.info(f'Loaded {len(self.skills)} skills from Box runtime')
+ skills[skill_name] = skill_data
+ self._skills_by_scope[key] = skills
+ suffix = f' ({dropped} dropped due to missing package_root)' if dropped else ''
+ self.ap.logger.info(f'Loaded {len(skills)} skills for Workspace {execution_context.workspace_uuid}{suffix}')
except Exception as exc:
- self.ap.logger.warning(f'Failed to load skills from Box runtime: {exc}')
+ self.ap.logger.warning(f'Failed to load skills for Workspace {execution_context.workspace_uuid}: {exc}')
- def refresh_skill_from_disk(self, skill_name: str) -> bool:
- """Confirm a single skill is present in the cache.
+ async def ensure_loaded(self, context: TenantContext) -> None:
+ key = self._scope_key(context)
+ if key not in self._skills_by_scope:
+ await self.reload_skills(context)
- With Box as the only source of truth, the actual reload is driven by
- SkillService callers awaiting ``reload_skills``; this method only
- reports whether the cache still has the skill.
- """
- if not skill_name:
- return False
- return skill_name in self.skills
+ def get_skills(self, context: TenantContext) -> dict[str, dict]:
+ return self._skills_by_scope.get(self._scope_key(context), {})
- def get_skill_by_name(self, name: str) -> dict | None:
- """Get skill data by name."""
- return self.skills.get(name)
+ def refresh_skill_from_disk(self, context: TenantContext, skill_name: str) -> bool:
+ return bool(skill_name) and skill_name in self.get_skills(context)
- def get_skill_index(self, bound_skills: list[str] | None = None) -> str:
- """Render the pipeline-visible skills as a short ``name: description``
- index suitable for the system prompt.
+ def get_skill_by_name(self, context: TenantContext, name: str) -> dict | None:
+ return self.get_skills(context).get(name)
- ``bound_skills`` follows the same convention as
- ``query.variables['_pipeline_bound_skills']``: ``None`` means every
- loaded skill is exposed; an explicit list filters to that subset.
- Returns an empty string when no skills are visible.
- """
+ def get_skill_index(self, context: TenantContext, bound_skills: list[str] | None = None) -> str:
lines: list[str] = []
- for skill in self.skills.values():
+ for skill in self.get_skills(context).values():
name = skill.get('name')
- if not name:
- continue
- if bound_skills is not None and name not in bound_skills:
+ if not name or (bound_skills is not None and name not in bound_skills):
continue
display = skill.get('display_name') or name
description = (skill.get('description') or '').strip().replace('\n', ' ')
lines.append(f'- {name} ({display}): {description}')
+ return 'Available Skills:\n' + '\n'.join(lines) if lines else ''
- if not lines:
- return ''
- return 'Available Skills:\n' + '\n'.join(lines)
-
- def build_skill_aware_prompt_addition(self, bound_skills: list[str] | None = None) -> str:
- """Build the system-prompt addendum that makes the LLM aware of the
- pipeline-visible skills.
-
- Only metadata (name + description) is injected — the full SKILL.md is
- loaded later via the ``activate`` Tool Call, protecting KV cache and
- matching Claude Code's progressive disclosure pattern. Returns an
- empty string when no skills are visible (no prompt change at all).
- """
- skill_index = self.get_skill_index(bound_skills)
+ def build_skill_aware_prompt_addition(
+ self,
+ context: TenantContext,
+ bound_skills: list[str] | None = None,
+ ) -> str:
+ skill_index = self.get_skill_index(context, bound_skills)
if not skill_index:
return ''
return (
@@ -140,3 +133,6 @@ class SkillManager:
'the tool result. If no skill is a clear match, respond normally '
'without activating any skill.'
)
+
+ def total_cached_skill_count(self) -> int:
+ return sum(len(skills) for skills in self._skills_by_scope.values())
diff --git a/src/langbot/pkg/storage/mgr.py b/src/langbot/pkg/storage/mgr.py
index 08f7c8d78..354bd2f29 100644
--- a/src/langbot/pkg/storage/mgr.py
+++ b/src/langbot/pkg/storage/mgr.py
@@ -1,11 +1,28 @@
from __future__ import annotations
+import hashlib
+import json
+import re
+from pathlib import PurePath
from ..core import app
+from ..api.http.authz import WorkspaceRequiredError
+from ..api.http.context import ExecutionContext, RequestContext
from . import provider
from .providers import localstorage
+_SAFE_OWNER_TYPE = re.compile(r'^[a-z][a-z0-9_-]{0,63}$')
+_SCOPED_KEY = re.compile(
+ r'^v1/(?P[a-f0-9]{24})/'
+ r'(?P[0-9a-fA-F-]{36})/'
+ r'(?P[1-9][0-9]*)/'
+ r'(?P[a-z][a-z0-9_-]{0,63})/'
+ r'(?P[a-f0-9]{32})/'
+ r'(?P[a-f0-9]{64})(?P\.[a-zA-Z0-9]{1,16})?$'
+)
+
+
class StorageMgr:
"""Storage manager"""
@@ -16,6 +33,290 @@ class StorageMgr:
def __init__(self, ap: app.Application):
self.ap = ap
+ @staticmethod
+ def _require_execution_scope(
+ context: ExecutionContext | RequestContext,
+ ) -> tuple[str, str, int]:
+ if not isinstance(context, (ExecutionContext, RequestContext)):
+ raise WorkspaceRequiredError('Storage operations require an explicit Workspace context')
+ instance_uuid = context.instance_uuid.strip()
+ workspace_uuid = context.workspace_uuid.strip()
+ generation = context.placement_generation
+ if not instance_uuid or not workspace_uuid:
+ raise WorkspaceRequiredError('Storage operations require an instance and Workspace')
+ if generation <= 0:
+ raise WorkspaceRequiredError('Storage operations require a positive placement generation')
+ return instance_uuid, workspace_uuid, generation
+
+ @staticmethod
+ def _digest(value: str, length: int) -> str:
+ return hashlib.sha256(value.encode('utf-8')).hexdigest()[:length]
+
+ async def _require_active_execution_scope(
+ self,
+ context: ExecutionContext | RequestContext,
+ ) -> None:
+ """Revalidate the captured generation before touching object storage."""
+ instance_uuid, workspace_uuid, generation = self._require_execution_scope(context)
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ if workspace_service is None:
+ raise WorkspaceRequiredError('Storage execution scope is unavailable')
+ try:
+ binding = await workspace_service.get_execution_binding(
+ workspace_uuid,
+ expected_generation=generation,
+ )
+ except Exception as exc:
+ raise WorkspaceRequiredError('Storage execution scope is unavailable') from exc
+ if (
+ getattr(binding, 'instance_uuid', None) != instance_uuid
+ or getattr(binding, 'workspace_uuid', None) != workspace_uuid
+ or getattr(binding, 'placement_generation', None) != generation
+ ):
+ raise WorkspaceRequiredError('Storage execution scope is unavailable')
+
+ @classmethod
+ def canonical_binary_storage_key(
+ cls,
+ context: ExecutionContext | RequestContext,
+ *,
+ owner_type: str,
+ owner: str,
+ key: str,
+ ) -> str:
+ """Return a bounded canonical key over every BinaryStorage owner dimension."""
+
+ instance_uuid, workspace_uuid, _ = cls._require_execution_scope(context)
+ if not _SAFE_OWNER_TYPE.fullmatch(owner_type):
+ raise ValueError('Invalid storage owner_type')
+ if not owner or not key:
+ raise ValueError('Storage owner and key are required')
+ canonical = json.dumps(
+ [instance_uuid, workspace_uuid, owner_type, owner, key],
+ ensure_ascii=False,
+ separators=(',', ':'),
+ )
+ return f'v1:{cls._digest(instance_uuid, 24)}:{workspace_uuid}:{owner_type}:{hashlib.sha256(canonical.encode()).hexdigest()}'
+
+ @classmethod
+ def scoped_object_key(
+ cls,
+ context: ExecutionContext | RequestContext,
+ *,
+ owner_type: str,
+ owner: str,
+ key: str,
+ preserve_suffix: bool = True,
+ ) -> str:
+ """Build a non-enumerable object key with an explicit tenant boundary."""
+
+ instance_uuid, workspace_uuid, generation = cls._require_execution_scope(context)
+ if not _SAFE_OWNER_TYPE.fullmatch(owner_type):
+ raise ValueError('Invalid storage owner_type')
+ if not owner or not key:
+ raise ValueError('Storage owner and key are required')
+ suffix = PurePath(key).suffix.lower() if preserve_suffix else ''
+ if not re.fullmatch(r'\.[a-z0-9]{1,16}', suffix):
+ suffix = ''
+ return (
+ f'v1/{cls._digest(instance_uuid, 24)}/{workspace_uuid}/{generation}/'
+ f'{owner_type}/{cls._digest(owner, 32)}/{cls._digest(key, 64)}{suffix}'
+ )
+
+ @classmethod
+ def scoped_prefix(
+ cls,
+ context: ExecutionContext | RequestContext,
+ *,
+ owner_type: str | None = None,
+ ) -> str:
+ instance_uuid, workspace_uuid, generation = cls._require_execution_scope(context)
+ prefix = f'v1/{cls._digest(instance_uuid, 24)}/{workspace_uuid}/{generation}/'
+ if owner_type is not None:
+ if not _SAFE_OWNER_TYPE.fullmatch(owner_type):
+ raise ValueError('Invalid storage owner_type')
+ prefix += f'{owner_type}/'
+ return prefix
+
+ async def save_scoped(
+ self,
+ context: ExecutionContext | RequestContext,
+ *,
+ owner_type: str,
+ owner: str,
+ key: str,
+ value: bytes,
+ preserve_suffix: bool = True,
+ ) -> str:
+ await self._require_active_execution_scope(context)
+ object_key = self.scoped_object_key(
+ context,
+ owner_type=owner_type,
+ owner=owner,
+ key=key,
+ preserve_suffix=preserve_suffix,
+ )
+ await self.storage_provider.save(object_key, value)
+ return object_key
+
+ async def load_scoped(
+ self,
+ context: ExecutionContext | RequestContext,
+ *,
+ owner_type: str,
+ owner: str,
+ key: str,
+ preserve_suffix: bool = True,
+ ) -> bytes:
+ await self._require_active_execution_scope(context)
+ object_key = self.scoped_object_key(
+ context,
+ owner_type=owner_type,
+ owner=owner,
+ key=key,
+ preserve_suffix=preserve_suffix,
+ )
+ return await self.storage_provider.load(object_key)
+
+ async def delete_scoped(
+ self,
+ context: ExecutionContext | RequestContext,
+ *,
+ owner_type: str,
+ owner: str,
+ key: str,
+ preserve_suffix: bool = True,
+ ) -> None:
+ await self._require_active_execution_scope(context)
+ object_key = self.scoped_object_key(
+ context,
+ owner_type=owner_type,
+ owner=owner,
+ key=key,
+ preserve_suffix=preserve_suffix,
+ )
+ await self.storage_provider.delete(object_key)
+
+ async def resolve_public_object(
+ self,
+ object_key: str,
+ *,
+ expected_owner_type: str,
+ ) -> bytes | None:
+ """Load an opaque public object after validating its trusted scope."""
+
+ match = _SCOPED_KEY.fullmatch(object_key)
+ if match is None or match.group('owner_type') != expected_owner_type:
+ return None
+ if match.group('instance') != self._digest(self.ap.workspace_service.instance_uuid, 24):
+ return None
+ workspace_uuid = match.group('workspace')
+ generation = int(match.group('generation'))
+ try:
+ await self.ap.workspace_service.get_execution_binding(
+ workspace_uuid,
+ expected_generation=generation,
+ )
+ except Exception:
+ return None
+ if not await self.storage_provider.exists(object_key):
+ return None
+ return await self.storage_provider.load(object_key)
+
+ @classmethod
+ def require_scoped_object_key(
+ cls,
+ context: ExecutionContext | RequestContext,
+ object_key: str,
+ *,
+ expected_owner_type: str,
+ ) -> None:
+ """Validate an opaque object key against every captured scope field."""
+
+ instance_uuid, workspace_uuid, generation = cls._require_execution_scope(context)
+ match = _SCOPED_KEY.fullmatch(object_key)
+ if (
+ match is None
+ or match.group('instance') != cls._digest(instance_uuid, 24)
+ or match.group('workspace') != workspace_uuid
+ or int(match.group('generation')) != generation
+ or match.group('owner_type') != expected_owner_type
+ ):
+ raise WorkspaceRequiredError('Object key does not belong to the execution scope')
+
+ async def exists_scoped_object_key(
+ self,
+ context: ExecutionContext | RequestContext,
+ object_key: str,
+ *,
+ expected_owner_type: str,
+ ) -> bool:
+ await self._require_active_execution_scope(context)
+ self.require_scoped_object_key(
+ context,
+ object_key,
+ expected_owner_type=expected_owner_type,
+ )
+ return await self.storage_provider.exists(object_key)
+
+ async def load_scoped_object_key(
+ self,
+ context: ExecutionContext | RequestContext,
+ object_key: str,
+ *,
+ expected_owner_type: str,
+ ) -> bytes:
+ await self._require_active_execution_scope(context)
+ self.require_scoped_object_key(
+ context,
+ object_key,
+ expected_owner_type=expected_owner_type,
+ )
+ return await self.storage_provider.load(object_key)
+
+ async def size_scoped_object_key(
+ self,
+ context: ExecutionContext | RequestContext,
+ object_key: str,
+ *,
+ expected_owner_type: str,
+ ) -> int:
+ await self._require_active_execution_scope(context)
+ self.require_scoped_object_key(
+ context,
+ object_key,
+ expected_owner_type=expected_owner_type,
+ )
+ return await self.storage_provider.size(object_key)
+
+ async def delete_scoped_object_key(
+ self,
+ context: ExecutionContext | RequestContext,
+ object_key: str,
+ *,
+ expected_owner_type: str,
+ ) -> None:
+ """Delete a previously returned key only inside the captured scope."""
+
+ await self._require_active_execution_scope(context)
+ self.require_scoped_object_key(
+ context,
+ object_key,
+ expected_owner_type=expected_owner_type,
+ )
+ if await self.storage_provider.exists(object_key):
+ await self.storage_provider.delete(object_key)
+
+ @classmethod
+ def is_scoped_object_key(
+ cls,
+ object_key: str,
+ *,
+ expected_owner_type: str | None = None,
+ ) -> bool:
+ match = _SCOPED_KEY.fullmatch(object_key)
+ return match is not None and (expected_owner_type is None or match.group('owner_type') == expected_owner_type)
+
async def initialize(self):
storage_config = self.ap.instance_config.data.get('storage', {})
storage_type = storage_config.get('use', 'local')
diff --git a/src/langbot/pkg/telemetry/heartbeat.py b/src/langbot/pkg/telemetry/heartbeat.py
index 34b616733..538c480bf 100644
--- a/src/langbot/pkg/telemetry/heartbeat.py
+++ b/src/langbot/pkg/telemetry/heartbeat.py
@@ -99,8 +99,8 @@ async def build_heartbeat_payload(ap: core_app.Application) -> dict:
# Skill count (from Box runtime via skill manager)
try:
skill_mgr = getattr(ap, 'skill_mgr', None)
- if skill_mgr is not None and getattr(skill_mgr, 'skills', None) is not None:
- features['skill_count'] = len(skill_mgr.skills)
+ if skill_mgr is not None:
+ features['skill_count'] = skill_mgr.total_cached_skill_count()
except Exception:
pass
diff --git a/src/langbot/pkg/utils/httpclient.py b/src/langbot/pkg/utils/httpclient.py
index e9c04b346..3c7df9afe 100644
--- a/src/langbot/pkg/utils/httpclient.py
+++ b/src/langbot/pkg/utils/httpclient.py
@@ -29,7 +29,13 @@ def get_session(*, trust_env: bool = False) -> aiohttp.ClientSession:
session = _sessions.get(key)
if session is None or session.closed:
- session = aiohttp.ClientSession(trust_env=trust_env)
+ # Shared transport pools must never share upstream cookie state across
+ # Workspace-scoped requests. Callers that need a stateful cookie jar
+ # must own a dedicated session instead of using this global pool.
+ session = aiohttp.ClientSession(
+ trust_env=trust_env,
+ cookie_jar=aiohttp.DummyCookieJar(),
+ )
_sessions[key] = session
return session
diff --git a/src/langbot/pkg/utils/managed_runtime.py b/src/langbot/pkg/utils/managed_runtime.py
index 77f59be4c..bf0511d0a 100644
--- a/src/langbot/pkg/utils/managed_runtime.py
+++ b/src/langbot/pkg/utils/managed_runtime.py
@@ -28,7 +28,11 @@ class ManagedRuntimeConnector:
self.runtime_subprocess = None
self.runtime_subprocess_task = None
- async def _start_runtime_subprocess(self, *args: str) -> None:
+ async def _start_runtime_subprocess(
+ self,
+ *args: str,
+ env_overrides: dict[str, str] | None = None,
+ ) -> None:
"""Launch a local runtime as a subprocess of the current Python interpreter.
If a subprocess is already running (no *returncode* yet), this is a no-op.
@@ -38,6 +42,8 @@ class ManagedRuntimeConnector:
python_path = sys.executable
env = os.environ.copy()
+ if env_overrides:
+ env.update(env_overrides)
self.runtime_subprocess = await asyncio.create_subprocess_exec(
python_path,
*args,
diff --git a/src/langbot/pkg/vector/mgr.py b/src/langbot/pkg/vector/mgr.py
index 765c259f9..88ef12dac 100644
--- a/src/langbot/pkg/vector/mgr.py
+++ b/src/langbot/pkg/vector/mgr.py
@@ -1,6 +1,15 @@
from __future__ import annotations
+import uuid
+
+import sqlalchemy
+
+from ..api.http.authz import WorkspaceRequiredError
+from ..api.http.context import ExecutionContext
from ..core import app
+from ..entity.persistence import rag as persistence_rag
+from ..entity.persistence import workspace as persistence_workspace
+from ..workspace.errors import WorkspaceNotFoundError
from .vdb import VectorDatabase, SearchType
@@ -87,26 +96,141 @@ class VectorDBManager:
return [SearchType.VECTOR.value]
return [st.value for st in self.vector_db.supported_search_types()]
+ @staticmethod
+ def physical_collection_name(
+ execution_context: ExecutionContext,
+ knowledge_base_uuid: str,
+ ) -> str:
+ """Derive an opaque physical collection from trusted tenant identity.
+
+ Vector backends have different collection-name constraints, so the
+ instance, Workspace and knowledge-base identifiers are encoded through
+ UUIDv5 instead of being concatenated into a client-visible handle.
+ Placement generation is deliberately not part of the name: generation
+ fencing rejects stale work while preserving data across placements.
+ """
+
+ if not isinstance(execution_context, ExecutionContext):
+ raise WorkspaceRequiredError('ExecutionContext is required for vector access')
+ instance_uuid = execution_context.instance_uuid.strip()
+ workspace_uuid = execution_context.workspace_uuid.strip()
+ kb_uuid = knowledge_base_uuid.strip() if isinstance(knowledge_base_uuid, str) else ''
+ if not instance_uuid or not workspace_uuid or not kb_uuid:
+ raise WorkspaceRequiredError('Instance, Workspace and knowledge-base context are required')
+ if execution_context.placement_generation <= 0:
+ raise WorkspaceRequiredError('A positive placement generation is required')
+
+ collection_uuid = uuid.uuid5(
+ uuid.NAMESPACE_URL,
+ f'langbot:knowledge-vector:{instance_uuid}:{workspace_uuid}:{kb_uuid}',
+ )
+ return f'lb_{collection_uuid.hex}'
+
+ async def _validate_execution_context(self, execution_context: ExecutionContext) -> None:
+ """Validate the active placement before touching a vector backend."""
+
+ # Also performs structural validation before accessing app services.
+ self.physical_collection_name(execution_context, 'context-validation')
+ workspace_service = getattr(self.ap, 'workspace_service', None)
+ if workspace_service is None:
+ raise WorkspaceRequiredError('Workspace execution service is unavailable')
+ binding = await workspace_service.get_execution_binding(
+ execution_context.workspace_uuid,
+ expected_generation=execution_context.placement_generation,
+ )
+ if binding.instance_uuid != execution_context.instance_uuid:
+ raise WorkspaceRequiredError('ExecutionContext belongs to another LangBot instance')
+
+ async def _resolve_physical_collection_name(
+ self,
+ execution_context: ExecutionContext,
+ knowledge_base_uuid: str,
+ ) -> str:
+ """Resolve a scoped collection or an explicitly migrated OSS handle.
+
+ Legacy handles are server-owned migration state, not caller input.
+ They are honored only for the one local Workspace under the OSS
+ single-Workspace policy. A projected/cloud Workspace always gets the
+ opaque tenant-derived collection, even if its database row was
+ incorrectly marked as legacy.
+ """
+
+ await self._validate_execution_context(execution_context)
+ result = await self.ap.persistence_mgr.execute_async(
+ sqlalchemy.select(
+ persistence_rag.KnowledgeBase.collection_id,
+ persistence_rag.KnowledgeBase.legacy_vector_collection,
+ persistence_workspace.Workspace.source,
+ )
+ .join(
+ persistence_workspace.Workspace,
+ persistence_workspace.Workspace.uuid == persistence_rag.KnowledgeBase.workspace_uuid,
+ )
+ .where(
+ persistence_rag.KnowledgeBase.workspace_uuid == execution_context.workspace_uuid,
+ persistence_rag.KnowledgeBase.uuid == knowledge_base_uuid,
+ persistence_workspace.Workspace.instance_uuid == execution_context.instance_uuid,
+ )
+ .limit(1)
+ )
+ row = result.first()
+ if row is None:
+ raise WorkspaceNotFoundError('Knowledge base not found')
+
+ collection_id, legacy_vector_collection, workspace_source = row
+ if legacy_vector_collection:
+ policy = getattr(self.ap, 'workspace_policy', None)
+ is_single_workspace = policy is not None and not getattr(
+ policy,
+ 'multi_workspace_enabled',
+ True,
+ )
+ is_local_workspace = workspace_source == persistence_workspace.WorkspaceSource.LOCAL.value
+ if is_single_workspace and is_local_workspace and isinstance(collection_id, str) and collection_id.strip():
+ return collection_id
+ self.ap.logger.warning(
+ 'Ignored a legacy vector collection marker outside the local single-Workspace compatibility boundary.'
+ )
+
+ return self.physical_collection_name(execution_context, knowledge_base_uuid)
+
async def upsert(
self,
- collection_name: str,
+ execution_context: ExecutionContext,
+ knowledge_base_uuid: str,
vectors: list[list[float]],
ids: list[str],
metadata: list[dict] | None = None,
documents: list[str] | None = None,
):
- """Proxy: Upsert vectors"""
+ """Upsert vectors into a server-derived tenant collection."""
+
+ collection_name = await self._resolve_physical_collection_name(
+ execution_context,
+ knowledge_base_uuid,
+ )
+ source_metadata = metadata or [{} for _ in vectors]
+ scoped_metadata = [
+ {
+ **item,
+ '_langbot_instance_uuid': execution_context.instance_uuid,
+ '_langbot_workspace_uuid': execution_context.workspace_uuid,
+ '_langbot_knowledge_base_uuid': knowledge_base_uuid,
+ }
+ for item in source_metadata
+ ]
await self.vector_db.add_embeddings(
collection=collection_name,
ids=ids,
embeddings_list=vectors,
- metadatas=metadata or [{} for _ in vectors],
+ metadatas=scoped_metadata,
documents=documents,
)
async def search(
self,
- collection_name: str,
+ execution_context: ExecutionContext,
+ knowledge_base_uuid: str,
query_vector: list[float],
limit: int,
filter: dict | None = None,
@@ -120,6 +244,10 @@ class VectorDBManager:
The underlying VectorDatabase.search returns Chroma-style format:
{ 'ids': [['id1']], 'distances': [[0.1]], 'metadatas': [[{}]] }
"""
+ collection_name = await self._resolve_physical_collection_name(
+ execution_context,
+ knowledge_base_uuid,
+ )
results = await self.vector_db.search(
collection=collection_name,
query_embedding=query_vector,
@@ -154,30 +282,58 @@ class VectorDBManager:
return parsed_results
- async def delete_by_file_id(self, collection_name: str, file_ids: list[str]):
+ async def delete_by_file_id(
+ self,
+ execution_context: ExecutionContext,
+ knowledge_base_uuid: str,
+ file_ids: list[str],
+ ):
"""Proxy: Delete vectors by file_id (metadata-level identifier).
This delegates to VectorDatabase.delete_by_file_id which removes
all vectors associated with the given file IDs.
"""
+ collection_name = await self._resolve_physical_collection_name(
+ execution_context,
+ knowledge_base_uuid,
+ )
for file_id in file_ids:
await self.vector_db.delete_by_file_id(collection_name, file_id)
- async def delete_collection(self, collection_name: str):
- """Proxy: Delete an entire collection."""
+ async def delete_collection(
+ self,
+ execution_context: ExecutionContext,
+ knowledge_base_uuid: str,
+ ):
+ """Delete one server-derived tenant collection."""
+
+ collection_name = await self._resolve_physical_collection_name(
+ execution_context,
+ knowledge_base_uuid,
+ )
await self.vector_db.delete_collection(collection_name)
- async def delete_by_filter(self, collection_name: str, filter: dict) -> int:
+ async def delete_by_filter(
+ self,
+ execution_context: ExecutionContext,
+ knowledge_base_uuid: str,
+ filter: dict,
+ ) -> int:
"""Proxy: Delete vectors by metadata filter.
Returns:
Number of deleted vectors (best-effort; some backends return 0).
"""
+ collection_name = await self._resolve_physical_collection_name(
+ execution_context,
+ knowledge_base_uuid,
+ )
return await self.vector_db.delete_by_filter(collection_name, filter)
async def list_by_filter(
self,
- collection_name: str,
+ execution_context: ExecutionContext,
+ knowledge_base_uuid: str,
filter: dict | None = None,
limit: int = 20,
offset: int = 0,
@@ -187,4 +343,8 @@ class VectorDBManager:
Returns:
Tuple of (items, total).
"""
+ collection_name = await self._resolve_physical_collection_name(
+ execution_context,
+ knowledge_base_uuid,
+ )
return await self.vector_db.list_by_filter(collection_name, filter, limit, offset)
diff --git a/src/langbot/pkg/workspace/__init__.py b/src/langbot/pkg/workspace/__init__.py
new file mode 100644
index 000000000..c8db78b9a
--- /dev/null
+++ b/src/langbot/pkg/workspace/__init__.py
@@ -0,0 +1,25 @@
+from .errors import (
+ WorkspaceExecutionUnavailableError,
+ WorkspaceGenerationMismatchError,
+ WorkspaceInvariantError,
+ WorkspaceLimitExceededError,
+ WorkspaceNotFoundError,
+ WorkspaceOwnerAlreadyExistsError,
+)
+from .entities import WorkspaceExecutionBinding
+from .policy import SingleWorkspacePolicy
+from .repository import WorkspaceRepository
+from .service import WorkspaceService
+
+__all__ = [
+ 'SingleWorkspacePolicy',
+ 'WorkspaceExecutionBinding',
+ 'WorkspaceExecutionUnavailableError',
+ 'WorkspaceGenerationMismatchError',
+ 'WorkspaceInvariantError',
+ 'WorkspaceLimitExceededError',
+ 'WorkspaceNotFoundError',
+ 'WorkspaceOwnerAlreadyExistsError',
+ 'WorkspaceRepository',
+ 'WorkspaceService',
+]
diff --git a/src/langbot/pkg/workspace/collaboration.py b/src/langbot/pkg/workspace/collaboration.py
new file mode 100644
index 000000000..69eda0a9c
--- /dev/null
+++ b/src/langbot/pkg/workspace/collaboration.py
@@ -0,0 +1,684 @@
+from __future__ import annotations
+
+import dataclasses
+import datetime
+import hashlib
+import secrets
+import typing
+import uuid
+import asyncio
+from contextlib import asynccontextmanager
+from collections.abc import Awaitable, Callable
+
+import sqlalchemy
+from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
+
+from ..entity.persistence.user import AccountStatus, User
+from ..entity.persistence.workspace import (
+ InvitationStatus,
+ MembershipRole,
+ MembershipStatus,
+ Workspace,
+ WorkspaceInvitation,
+ WorkspaceMembership,
+ WorkspaceSource,
+ WorkspaceStatus,
+)
+from .entities import WorkspaceExecutionBinding
+from .errors import WorkspaceNotFoundError
+from .policy import CloudWorkspacePolicy, SingleWorkspacePolicy
+from .service import WorkspaceService
+
+if typing.TYPE_CHECKING:
+ from ..core.app import Application
+
+
+class WorkspaceCollaborationError(Exception):
+ """Stable collaboration error surfaced by the Workspace API."""
+
+ code = 'workspace_collaboration_error'
+
+
+class MembershipNotFoundError(WorkspaceCollaborationError):
+ code = 'membership_not_found'
+
+
+class MembershipPermissionError(WorkspaceCollaborationError):
+ code = 'permission_denied'
+
+
+class LastOwnerError(WorkspaceCollaborationError):
+ code = 'last_owner_required'
+
+
+class InvitationError(WorkspaceCollaborationError):
+ code = 'invitation_invalid'
+
+
+class InvitationExpiredError(InvitationError):
+ code = 'invitation_expired'
+
+
+class InvitationRevokedError(InvitationError):
+ code = 'invitation_revoked'
+
+
+class InvitationUsedError(InvitationError):
+ code = 'invitation_used'
+
+
+class InvitationEmailMismatchError(InvitationError):
+ code = 'invitation_email_mismatch'
+
+
+class InvitationRoleError(InvitationError):
+ code = 'invitation_role_invalid'
+
+
+class AlreadyMemberError(InvitationError):
+ code = 'already_a_member'
+
+
+@dataclasses.dataclass(frozen=True, slots=True)
+class ResolvedWorkspaceAccess:
+ workspace: Workspace
+ membership: WorkspaceMembership
+ execution: WorkspaceExecutionBinding
+
+
+@dataclasses.dataclass(frozen=True, slots=True)
+class WorkspaceMemberView:
+ membership: WorkspaceMembership
+ email: str
+
+
+@dataclasses.dataclass(frozen=True, slots=True)
+class CreatedInvitation:
+ invitation: WorkspaceInvitation
+ token: str
+
+
+@dataclasses.dataclass(slots=True)
+class _InvitationLockEntry:
+ lock: asyncio.Lock
+ users: int = 0
+
+
+T = typing.TypeVar('T')
+
+
+def normalize_email(email: str) -> str:
+ """Return the canonical email identity used by invitations."""
+
+ normalized = email.strip().casefold()
+ if not normalized or '@' not in normalized:
+ raise ValueError('A valid email address is required')
+ if len(normalized) > 320:
+ raise ValueError('Email address exceeds the normalized identity limit')
+ return normalized
+
+
+def hash_invitation_token(token: str) -> str:
+ """Hash an invitation bearer secret for lookup and at-rest storage."""
+
+ return hashlib.sha256(token.encode('utf-8')).hexdigest()
+
+
+class WorkspaceCollaborationService:
+ """Membership and invitation operations for the local Workspace directory."""
+
+ def __init__(
+ self,
+ ap: Application,
+ workspace_service: WorkspaceService,
+ *,
+ policy: SingleWorkspacePolicy | CloudWorkspacePolicy | None = None,
+ ) -> None:
+ self.ap = ap
+ self.workspace_service = workspace_service
+ self.policy = policy or workspace_service.policy
+ self._invitation_locks: dict[str, _InvitationLockEntry] = {}
+ self._invitation_locks_guard = asyncio.Lock()
+
+ def _session_factory(self) -> async_sessionmaker[AsyncSession]:
+ return async_sessionmaker(
+ self.ap.persistence_mgr.get_db_engine(),
+ expire_on_commit=False,
+ )
+
+ async def resolve_account_workspace(
+ self,
+ account_uuid: str,
+ requested_workspace_uuid: str | None,
+ *,
+ session: AsyncSession | None = None,
+ ) -> ResolvedWorkspaceAccess:
+ """Resolve a selector against an active Account membership."""
+
+ async def operation(active_session: AsyncSession) -> ResolvedWorkspaceAccess:
+ workspace_uuid = requested_workspace_uuid.strip() if requested_workspace_uuid else None
+ if workspace_uuid is None:
+ if self.policy.multi_workspace_enabled:
+ raise WorkspaceNotFoundError('A Workspace selector is required')
+ workspace = await self.workspace_service.get_singleton_workspace(session=active_session)
+ else:
+ workspace = await active_session.get(Workspace, workspace_uuid)
+ if (
+ workspace is None
+ or workspace.instance_uuid != self.workspace_service.instance_uuid
+ or workspace.status != WorkspaceStatus.ACTIVE.value
+ ):
+ raise WorkspaceNotFoundError('Workspace not found')
+
+ membership = await active_session.scalar(
+ sqlalchemy.select(WorkspaceMembership).where(
+ WorkspaceMembership.workspace_uuid == workspace.uuid,
+ WorkspaceMembership.account_uuid == account_uuid,
+ WorkspaceMembership.status == MembershipStatus.ACTIVE.value,
+ )
+ )
+ if membership is None:
+ # Deliberately hide Workspace existence across Accounts.
+ raise WorkspaceNotFoundError('Workspace not found')
+
+ execution = await self.workspace_service.get_execution_binding(
+ workspace.uuid,
+ session=active_session,
+ )
+ return ResolvedWorkspaceAccess(workspace, membership, execution)
+
+ return await self._run(operation, session=session, read_only=True)
+
+ async def list_account_workspaces(
+ self,
+ account_uuid: str,
+ *,
+ session: AsyncSession | None = None,
+ ) -> list[ResolvedWorkspaceAccess]:
+ async def operation(active_session: AsyncSession) -> list[ResolvedWorkspaceAccess]:
+ statement = (
+ sqlalchemy.select(WorkspaceMembership, Workspace)
+ .join(Workspace, Workspace.uuid == WorkspaceMembership.workspace_uuid)
+ .where(
+ WorkspaceMembership.account_uuid == account_uuid,
+ WorkspaceMembership.status == MembershipStatus.ACTIVE.value,
+ Workspace.instance_uuid == self.workspace_service.instance_uuid,
+ Workspace.status == WorkspaceStatus.ACTIVE.value,
+ )
+ .order_by(Workspace.created_at, Workspace.uuid)
+ )
+ rows = (await active_session.execute(statement)).all()
+ accesses: list[ResolvedWorkspaceAccess] = []
+ for membership, workspace in rows:
+ execution = await self.workspace_service.get_execution_binding(
+ workspace.uuid,
+ session=active_session,
+ )
+ accesses.append(ResolvedWorkspaceAccess(workspace, membership, execution))
+ return accesses
+
+ return await self._run(operation, session=session, read_only=True)
+
+ async def list_members(
+ self,
+ workspace_uuid: str,
+ actor: WorkspaceMembership,
+ *,
+ session: AsyncSession | None = None,
+ ) -> list[WorkspaceMemberView]:
+ async def operation(active_session: AsyncSession) -> list[WorkspaceMemberView]:
+ await self._load_actor(active_session, workspace_uuid, actor)
+ statement = (
+ sqlalchemy.select(WorkspaceMembership, User.user)
+ .join(User, User.uuid == WorkspaceMembership.account_uuid)
+ .where(
+ WorkspaceMembership.workspace_uuid == workspace_uuid,
+ WorkspaceMembership.status == MembershipStatus.ACTIVE.value,
+ User.status == AccountStatus.ACTIVE.value,
+ )
+ .order_by(WorkspaceMembership.created_at, WorkspaceMembership.uuid)
+ )
+ return [
+ WorkspaceMemberView(membership=membership, email=email)
+ for membership, email in (await active_session.execute(statement)).all()
+ ]
+
+ return await self._run(operation, session=session, read_only=True)
+
+ async def list_invitations(
+ self,
+ workspace_uuid: str,
+ actor: WorkspaceMembership,
+ *,
+ session: AsyncSession | None = None,
+ ) -> list[WorkspaceInvitation]:
+ async def operation(active_session: AsyncSession) -> list[WorkspaceInvitation]:
+ await self._require_local_workspace(active_session, workspace_uuid)
+ persisted_actor = await self._load_actor(active_session, workspace_uuid, actor)
+ self._require_member_manager(persisted_actor, workspace_uuid)
+ await self._expire_pending_invitations(active_session, workspace_uuid=workspace_uuid)
+ statement = (
+ sqlalchemy.select(WorkspaceInvitation)
+ .where(
+ WorkspaceInvitation.workspace_uuid == workspace_uuid,
+ WorkspaceInvitation.status == InvitationStatus.PENDING.value,
+ )
+ .order_by(WorkspaceInvitation.created_at, WorkspaceInvitation.uuid)
+ )
+ return list((await active_session.scalars(statement)).all())
+
+ return await self._run(operation, session=session)
+
+ async def create_invitation(
+ self,
+ workspace_uuid: str,
+ actor: WorkspaceMembership,
+ email: str,
+ role: str,
+ *,
+ expires_in: datetime.timedelta = datetime.timedelta(days=7),
+ session: AsyncSession | None = None,
+ ) -> CreatedInvitation:
+ if role not in {
+ MembershipRole.ADMIN.value,
+ MembershipRole.DEVELOPER.value,
+ MembershipRole.OPERATOR.value,
+ MembershipRole.VIEWER.value,
+ }:
+ raise InvitationRoleError('Invitations cannot grant this role')
+ normalized_email = normalize_email(email)
+ if expires_in <= datetime.timedelta(0):
+ raise InvitationError('Invitation expiry must be in the future')
+
+ async def operation(active_session: AsyncSession) -> CreatedInvitation:
+ await self._require_local_workspace(active_session, workspace_uuid)
+ persisted_actor = await self._load_actor(active_session, workspace_uuid, actor, for_update=True)
+ self._require_member_manager(persisted_actor, workspace_uuid)
+ existing_account = await active_session.scalar(
+ sqlalchemy.select(User).where(User.normalized_email == normalized_email)
+ )
+ if existing_account is not None:
+ existing_membership = await active_session.scalar(
+ sqlalchemy.select(WorkspaceMembership).where(
+ WorkspaceMembership.workspace_uuid == workspace_uuid,
+ WorkspaceMembership.account_uuid == existing_account.uuid,
+ WorkspaceMembership.status == MembershipStatus.ACTIVE.value,
+ )
+ )
+ if existing_membership is not None:
+ raise AlreadyMemberError('This Account is already a Workspace member')
+
+ existing_pending = await active_session.scalar(
+ sqlalchemy.select(WorkspaceInvitation)
+ .where(
+ WorkspaceInvitation.workspace_uuid == workspace_uuid,
+ WorkspaceInvitation.normalized_email == normalized_email,
+ WorkspaceInvitation.status == InvitationStatus.PENDING.value,
+ )
+ .with_for_update()
+ )
+ now = self._utcnow()
+ if existing_pending is not None:
+ existing_pending.status = InvitationStatus.REVOKED.value
+ existing_pending.revoked_at = now
+ await active_session.flush()
+
+ token = f'lbi_{secrets.token_urlsafe(32)}'
+ invitation = WorkspaceInvitation(
+ uuid=str(uuid.uuid4()),
+ workspace_uuid=workspace_uuid,
+ normalized_email=normalized_email,
+ role=role,
+ token_hash=hash_invitation_token(token),
+ status=InvitationStatus.PENDING.value,
+ expires_at=now + expires_in,
+ created_by_account_uuid=persisted_actor.account_uuid,
+ )
+ active_session.add(invitation)
+ await active_session.flush()
+ return CreatedInvitation(invitation, token)
+
+ return await self._run(operation, session=session)
+
+ async def inspect_invitation(
+ self,
+ token: str,
+ *,
+ session: AsyncSession | None = None,
+ ) -> tuple[WorkspaceInvitation, Workspace]:
+ async def operation(active_session: AsyncSession) -> tuple[WorkspaceInvitation, Workspace]:
+ invitation = await self._get_invitation_by_token(active_session, token, for_update=True)
+ self._validate_invitation_state(invitation)
+ workspace = await active_session.get(Workspace, invitation.workspace_uuid)
+ if workspace is None or workspace.status != WorkspaceStatus.ACTIVE.value:
+ raise InvitationError('The invitation Workspace is unavailable')
+ if workspace.source != WorkspaceSource.LOCAL.value:
+ raise InvitationError('Cloud invitations are managed by the SaaS control plane')
+ return invitation, workspace
+
+ return await self._run(operation, session=session)
+
+ async def accept_invitation(
+ self,
+ token: str,
+ account_uuid: str,
+ *,
+ session: AsyncSession | None = None,
+ ) -> WorkspaceMembership:
+ async def operation(active_session: AsyncSession) -> WorkspaceMembership:
+ invitation = await self._get_invitation_by_token(active_session, token, for_update=True)
+ self._validate_invitation_state(invitation)
+ await self._require_local_workspace(active_session, invitation.workspace_uuid)
+ account = await active_session.scalar(sqlalchemy.select(User).where(User.uuid == account_uuid))
+ if account is None or account.status != AccountStatus.ACTIVE.value:
+ raise MembershipNotFoundError('Account not found')
+ if account.normalized_email != invitation.normalized_email:
+ raise InvitationEmailMismatchError('Invitation email does not match the Account')
+
+ membership = await active_session.scalar(
+ sqlalchemy.select(WorkspaceMembership)
+ .where(
+ WorkspaceMembership.workspace_uuid == invitation.workspace_uuid,
+ WorkspaceMembership.account_uuid == account_uuid,
+ )
+ .with_for_update()
+ )
+ now = self._utcnow()
+ if membership is None:
+ membership = WorkspaceMembership(
+ uuid=str(uuid.uuid4()),
+ workspace_uuid=invitation.workspace_uuid,
+ account_uuid=account_uuid,
+ role=invitation.role,
+ status=MembershipStatus.ACTIVE.value,
+ invited_by_account_uuid=invitation.created_by_account_uuid,
+ joined_at=now,
+ projection_revision=0,
+ )
+ active_session.add(membership)
+ elif membership.status != MembershipStatus.ACTIVE.value:
+ membership.role = invitation.role
+ membership.status = MembershipStatus.ACTIVE.value
+ membership.invited_by_account_uuid = invitation.created_by_account_uuid
+ membership.joined_at = now
+
+ invitation.status = InvitationStatus.ACCEPTED.value
+ invitation.accepted_at = now
+ await active_session.flush()
+ return membership
+
+ token_digest = hash_invitation_token(token)
+ async with self._invitation_lock(token_digest):
+ return await self._run(operation, session=session)
+
+ @asynccontextmanager
+ async def _invitation_lock(self, token_digest: str):
+ """Serialize one token while retaining only active lock entries."""
+
+ async with self._invitation_locks_guard:
+ entry = self._invitation_locks.get(token_digest)
+ if entry is None:
+ entry = _InvitationLockEntry(lock=asyncio.Lock())
+ self._invitation_locks[token_digest] = entry
+ entry.users += 1
+
+ await entry.lock.acquire()
+ try:
+ yield
+ finally:
+ entry.lock.release()
+ async with self._invitation_locks_guard:
+ entry.users -= 1
+ if entry.users == 0:
+ self._invitation_locks.pop(token_digest, None)
+
+ async def revoke_invitation(
+ self,
+ workspace_uuid: str,
+ invitation_uuid: str,
+ actor: WorkspaceMembership,
+ *,
+ session: AsyncSession | None = None,
+ ) -> WorkspaceInvitation:
+ async def operation(active_session: AsyncSession) -> WorkspaceInvitation:
+ await self._require_local_workspace(active_session, workspace_uuid)
+ persisted_actor = await self._load_actor(active_session, workspace_uuid, actor, for_update=True)
+ self._require_member_manager(persisted_actor, workspace_uuid)
+ invitation = await active_session.scalar(
+ sqlalchemy.select(WorkspaceInvitation)
+ .where(
+ WorkspaceInvitation.uuid == invitation_uuid,
+ WorkspaceInvitation.workspace_uuid == workspace_uuid,
+ )
+ .with_for_update()
+ )
+ if invitation is None:
+ raise InvitationError('Invitation not found')
+ if invitation.status == InvitationStatus.REVOKED.value:
+ return invitation
+ if invitation.status != InvitationStatus.PENDING.value:
+ self._validate_invitation_state(invitation)
+ invitation.status = InvitationStatus.REVOKED.value
+ invitation.revoked_at = self._utcnow()
+ await active_session.flush()
+ return invitation
+
+ return await self._run(operation, session=session)
+
+ async def update_member_role(
+ self,
+ workspace_uuid: str,
+ target_account_uuid: str,
+ role: str,
+ actor: WorkspaceMembership,
+ *,
+ session: AsyncSession | None = None,
+ ) -> WorkspaceMembership:
+ if role not in {item.value for item in MembershipRole}:
+ raise MembershipPermissionError('Unknown Workspace role')
+
+ async def operation(active_session: AsyncSession) -> WorkspaceMembership:
+ await self._require_local_workspace(active_session, workspace_uuid)
+ persisted_actor = await self._load_actor(active_session, workspace_uuid, actor, for_update=True)
+ self._require_member_manager(persisted_actor, workspace_uuid)
+ target = await self._get_active_member_for_update(
+ active_session,
+ workspace_uuid,
+ target_account_uuid,
+ )
+ self._require_can_manage_target(persisted_actor, target, new_role=role)
+ if target.role == MembershipRole.OWNER.value and role != MembershipRole.OWNER.value:
+ await self._require_another_owner(active_session, workspace_uuid, target.account_uuid)
+ target.role = role
+ await active_session.flush()
+ return target
+
+ return await self._run(operation, session=session)
+
+ async def remove_member(
+ self,
+ workspace_uuid: str,
+ target_account_uuid: str,
+ actor: WorkspaceMembership,
+ *,
+ session: AsyncSession | None = None,
+ ) -> WorkspaceMembership:
+ async def operation(active_session: AsyncSession) -> WorkspaceMembership:
+ await self._require_local_workspace(active_session, workspace_uuid)
+ persisted_actor = await self._load_actor(active_session, workspace_uuid, actor, for_update=True)
+ self._require_member_manager(persisted_actor, workspace_uuid)
+ target = await self._get_active_member_for_update(
+ active_session,
+ workspace_uuid,
+ target_account_uuid,
+ )
+ self._require_can_manage_target(persisted_actor, target)
+ if target.role == MembershipRole.OWNER.value:
+ await self._require_another_owner(active_session, workspace_uuid, target.account_uuid)
+ target.status = MembershipStatus.REMOVED.value
+ await active_session.flush()
+ return target
+
+ return await self._run(operation, session=session)
+
+ async def _get_invitation_by_token(
+ self,
+ session: AsyncSession,
+ token: str,
+ *,
+ for_update: bool,
+ ) -> WorkspaceInvitation:
+ if not isinstance(token, str) or not token.startswith('lbi_'):
+ raise InvitationError('Invitation not found')
+ statement = sqlalchemy.select(WorkspaceInvitation).where(
+ WorkspaceInvitation.token_hash == hash_invitation_token(token)
+ )
+ if for_update:
+ statement = statement.with_for_update()
+ invitation = await session.scalar(statement)
+ if invitation is None:
+ raise InvitationError('Invitation not found')
+ return invitation
+
+ async def _require_local_workspace(
+ self,
+ session: AsyncSession,
+ workspace_uuid: str,
+ ) -> Workspace:
+ workspace = await session.get(Workspace, workspace_uuid)
+ if (
+ workspace is None
+ or workspace.instance_uuid != self.workspace_service.instance_uuid
+ or workspace.status != WorkspaceStatus.ACTIVE.value
+ ):
+ raise WorkspaceNotFoundError('Workspace not found')
+ if workspace.source != WorkspaceSource.LOCAL.value:
+ raise MembershipPermissionError('Cloud Workspace directory changes are managed by the SaaS control plane')
+ return workspace
+
+ def _validate_invitation_state(self, invitation: WorkspaceInvitation) -> None:
+ if invitation.status == InvitationStatus.REVOKED.value:
+ raise InvitationRevokedError('Invitation was revoked')
+ if invitation.status == InvitationStatus.ACCEPTED.value:
+ raise InvitationUsedError('Invitation was already accepted')
+ if invitation.status == InvitationStatus.EXPIRED.value or invitation.expires_at <= self._utcnow():
+ invitation.status = InvitationStatus.EXPIRED.value
+ raise InvitationExpiredError('Invitation has expired')
+ if invitation.status != InvitationStatus.PENDING.value:
+ raise InvitationError('Invitation is not pending')
+
+ async def _expire_pending_invitations(
+ self,
+ session: AsyncSession,
+ *,
+ workspace_uuid: str,
+ ) -> None:
+ await session.execute(
+ sqlalchemy.update(WorkspaceInvitation)
+ .where(
+ WorkspaceInvitation.workspace_uuid == workspace_uuid,
+ WorkspaceInvitation.status == InvitationStatus.PENDING.value,
+ WorkspaceInvitation.expires_at <= self._utcnow(),
+ )
+ .values(status=InvitationStatus.EXPIRED.value)
+ )
+
+ async def _get_active_member_for_update(
+ self,
+ session: AsyncSession,
+ workspace_uuid: str,
+ account_uuid: str,
+ ) -> WorkspaceMembership:
+ membership = await session.scalar(
+ sqlalchemy.select(WorkspaceMembership)
+ .where(
+ WorkspaceMembership.workspace_uuid == workspace_uuid,
+ WorkspaceMembership.account_uuid == account_uuid,
+ WorkspaceMembership.status == MembershipStatus.ACTIVE.value,
+ )
+ .with_for_update()
+ )
+ if membership is None:
+ raise MembershipNotFoundError('Workspace member not found')
+ return membership
+
+ async def _load_actor(
+ self,
+ session: AsyncSession,
+ workspace_uuid: str,
+ actor: WorkspaceMembership,
+ *,
+ for_update: bool = False,
+ ) -> WorkspaceMembership:
+ self._require_actor_workspace(actor, workspace_uuid)
+ statement = sqlalchemy.select(WorkspaceMembership).where(
+ WorkspaceMembership.workspace_uuid == workspace_uuid,
+ WorkspaceMembership.account_uuid == actor.account_uuid,
+ WorkspaceMembership.status == MembershipStatus.ACTIVE.value,
+ )
+ if for_update:
+ statement = statement.with_for_update()
+ persisted_actor = await session.scalar(statement)
+ if persisted_actor is None:
+ raise WorkspaceNotFoundError('Workspace not found')
+ return persisted_actor
+
+ async def _require_another_owner(
+ self,
+ session: AsyncSession,
+ workspace_uuid: str,
+ excluded_account_uuid: str,
+ ) -> None:
+ owners = (
+ await session.scalars(
+ sqlalchemy.select(WorkspaceMembership)
+ .where(
+ WorkspaceMembership.workspace_uuid == workspace_uuid,
+ WorkspaceMembership.status == MembershipStatus.ACTIVE.value,
+ WorkspaceMembership.role == MembershipRole.OWNER.value,
+ )
+ .with_for_update()
+ )
+ ).all()
+ if not any(owner.account_uuid != excluded_account_uuid for owner in owners):
+ raise LastOwnerError('The last Workspace owner cannot be removed or demoted')
+
+ def _require_actor_workspace(self, actor: WorkspaceMembership, workspace_uuid: str) -> None:
+ if actor.workspace_uuid != workspace_uuid or actor.status != MembershipStatus.ACTIVE.value:
+ raise WorkspaceNotFoundError('Workspace not found')
+
+ def _require_member_manager(self, actor: WorkspaceMembership, workspace_uuid: str) -> None:
+ self._require_actor_workspace(actor, workspace_uuid)
+ if actor.role not in {MembershipRole.OWNER.value, MembershipRole.ADMIN.value}:
+ raise MembershipPermissionError('Member management permission is required')
+
+ def _require_can_manage_target(
+ self,
+ actor: WorkspaceMembership,
+ target: WorkspaceMembership,
+ *,
+ new_role: str | None = None,
+ ) -> None:
+ if actor.role == MembershipRole.ADMIN.value and (
+ target.role == MembershipRole.OWNER.value or new_role == MembershipRole.OWNER.value
+ ):
+ raise MembershipPermissionError('Admins cannot manage Workspace owners')
+
+ @staticmethod
+ def _utcnow() -> datetime.datetime:
+ return datetime.datetime.now(datetime.UTC).replace(tzinfo=None)
+
+ async def _run(
+ self,
+ operation: Callable[[AsyncSession], Awaitable[T]],
+ *,
+ session: AsyncSession | None,
+ read_only: bool = False,
+ ) -> T:
+ if session is not None:
+ return await operation(session)
+ async with self._session_factory()() as owned_session:
+ if read_only:
+ return await operation(owned_session)
+ async with owned_session.begin():
+ return await operation(owned_session)
diff --git a/src/langbot/pkg/workspace/entities.py b/src/langbot/pkg/workspace/entities.py
new file mode 100644
index 000000000..77ef4914f
--- /dev/null
+++ b/src/langbot/pkg/workspace/entities.py
@@ -0,0 +1,14 @@
+from __future__ import annotations
+
+from dataclasses import dataclass
+
+
+@dataclass(frozen=True, slots=True)
+class WorkspaceExecutionBinding:
+ """Core-neutral binding for one validated Workspace execution generation."""
+
+ instance_uuid: str
+ workspace_uuid: str
+ placement_generation: int
+ write_fenced: bool
+ state: str
diff --git a/src/langbot/pkg/workspace/errors.py b/src/langbot/pkg/workspace/errors.py
new file mode 100644
index 000000000..5e93171ac
--- /dev/null
+++ b/src/langbot/pkg/workspace/errors.py
@@ -0,0 +1,31 @@
+from __future__ import annotations
+
+
+class WorkspaceError(Exception):
+ """Base error for workspace directory operations."""
+
+
+class WorkspaceNotFoundError(WorkspaceError):
+ """Raised when the instance does not have its required local workspace."""
+
+
+class WorkspaceInvariantError(WorkspaceError):
+ """Raised when persisted workspace state violates a tenancy invariant."""
+
+
+class WorkspaceLimitExceededError(WorkspaceError):
+ """Raised when OSS code attempts to create a second local workspace."""
+
+ code = 'edition_limit'
+
+
+class WorkspaceOwnerAlreadyExistsError(WorkspaceError):
+ """Raised when another account already owns the singleton workspace."""
+
+
+class WorkspaceExecutionUnavailableError(WorkspaceError):
+ """Raised when a Workspace cannot accept work in its current execution state."""
+
+
+class WorkspaceGenerationMismatchError(WorkspaceExecutionUnavailableError):
+ """Raised when a caller holds a stale Workspace placement generation."""
diff --git a/src/langbot/pkg/workspace/policy.py b/src/langbot/pkg/workspace/policy.py
new file mode 100644
index 000000000..bfdca864d
--- /dev/null
+++ b/src/langbot/pkg/workspace/policy.py
@@ -0,0 +1,50 @@
+from __future__ import annotations
+
+from dataclasses import dataclass
+
+from .errors import WorkspaceLimitExceededError
+
+
+@dataclass(frozen=True, slots=True)
+class SingleWorkspacePolicy:
+ """OSS edition policy: one local workspace with unrestricted membership count."""
+
+ workspace_limit: int = 1
+ members_enabled: bool = True
+ invitations_enabled: bool = True
+ fixed_rbac_enabled: bool = True
+ multi_workspace_enabled: bool = False
+
+ def require_workspace_creation_allowed(self, current_workspace_count: int) -> None:
+ if current_workspace_count >= self.workspace_limit:
+ raise WorkspaceLimitExceededError(f'This LangBot edition allows at most {self.workspace_limit} workspace')
+
+
+@dataclass(frozen=True, slots=True)
+class CloudWorkspacePolicy:
+ """SaaS data-plane policy backed by closed control-plane projections.
+
+ Core never creates a cloud Workspace or mutates its directory. The policy
+ only enables explicit selection among already projected Workspaces.
+ """
+
+ workspace_limit: int = 0
+ members_enabled: bool = True
+ invitations_enabled: bool = False
+ fixed_rbac_enabled: bool = True
+ multi_workspace_enabled: bool = True
+
+ def require_workspace_creation_allowed(self, current_workspace_count: int) -> None:
+ del current_workspace_count
+ raise WorkspaceLimitExceededError('Cloud Workspaces are created by the SaaS control plane')
+
+
+def open_core_workspace_policy() -> SingleWorkspacePolicy:
+ """Return the only policy the open-source bootstrap may activate.
+
+ ``system.edition`` and other local configuration are deliberately absent
+ from this boundary. A future closed Cloud bootstrap must first verify its
+ signed InstanceManifest and then explicitly construct the cloud policy.
+ """
+
+ return SingleWorkspacePolicy()
diff --git a/src/langbot/pkg/workspace/repository.py b/src/langbot/pkg/workspace/repository.py
new file mode 100644
index 000000000..374fa9b0b
--- /dev/null
+++ b/src/langbot/pkg/workspace/repository.py
@@ -0,0 +1,84 @@
+from __future__ import annotations
+
+import sqlalchemy
+from sqlalchemy.ext.asyncio import AsyncSession
+
+from ..entity.persistence.workspace import (
+ MembershipRole,
+ MembershipStatus,
+ Workspace,
+ WorkspaceExecutionState,
+ WorkspaceMembership,
+ WorkspaceSource,
+)
+
+
+class WorkspaceRepository:
+ """Transaction-bound persistence operations for the workspace directory."""
+
+ def __init__(self, session: AsyncSession) -> None:
+ self.session = session
+
+ async def count_local_workspaces(self, instance_uuid: str) -> int:
+ statement = (
+ sqlalchemy.select(sqlalchemy.func.count())
+ .select_from(Workspace)
+ .where(
+ Workspace.instance_uuid == instance_uuid,
+ Workspace.source == WorkspaceSource.LOCAL.value,
+ )
+ )
+ return int((await self.session.scalar(statement)) or 0)
+
+ async def list_local_workspaces(self, instance_uuid: str, *, for_update: bool = False) -> list[Workspace]:
+ statement = (
+ sqlalchemy.select(Workspace)
+ .where(
+ Workspace.instance_uuid == instance_uuid,
+ Workspace.source == WorkspaceSource.LOCAL.value,
+ )
+ .order_by(Workspace.created_at, Workspace.uuid)
+ )
+ if for_update:
+ statement = statement.with_for_update()
+ return list((await self.session.scalars(statement)).all())
+
+ async def get_workspace(self, workspace_uuid: str) -> Workspace | None:
+ return await self.session.get(Workspace, workspace_uuid)
+
+ def add_workspace(self, workspace: Workspace) -> None:
+ self.session.add(workspace)
+
+ async def get_execution_state(self, workspace_uuid: str) -> WorkspaceExecutionState | None:
+ return await self.session.get(WorkspaceExecutionState, workspace_uuid)
+
+ def add_execution_state(self, execution_state: WorkspaceExecutionState) -> None:
+ self.session.add(execution_state)
+
+ async def get_membership(self, workspace_uuid: str, account_uuid: str) -> WorkspaceMembership | None:
+ statement = sqlalchemy.select(WorkspaceMembership).where(
+ WorkspaceMembership.workspace_uuid == workspace_uuid,
+ WorkspaceMembership.account_uuid == account_uuid,
+ )
+ return await self.session.scalar(statement)
+
+ async def get_active_owner(self, workspace_uuid: str, *, for_update: bool = False) -> WorkspaceMembership | None:
+ statement = (
+ sqlalchemy.select(WorkspaceMembership)
+ .where(
+ WorkspaceMembership.workspace_uuid == workspace_uuid,
+ WorkspaceMembership.role == MembershipRole.OWNER.value,
+ WorkspaceMembership.status == MembershipStatus.ACTIVE.value,
+ )
+ .order_by(WorkspaceMembership.created_at, WorkspaceMembership.uuid)
+ .limit(1)
+ )
+ if for_update:
+ statement = statement.with_for_update()
+ return await self.session.scalar(statement)
+
+ def add_membership(self, membership: WorkspaceMembership) -> None:
+ self.session.add(membership)
+
+ async def flush(self) -> None:
+ await self.session.flush()
diff --git a/src/langbot/pkg/workspace/service.py b/src/langbot/pkg/workspace/service.py
new file mode 100644
index 000000000..e1bd1ec5f
--- /dev/null
+++ b/src/langbot/pkg/workspace/service.py
@@ -0,0 +1,417 @@
+from __future__ import annotations
+
+import datetime
+import uuid
+import typing
+from collections.abc import Awaitable, Callable
+from typing import TypeVar
+
+from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker
+
+from ..entity.persistence.workspace import (
+ MembershipRole,
+ MembershipStatus,
+ Workspace,
+ WorkspaceExecutionSource,
+ WorkspaceExecutionState,
+ WorkspaceExecutionStatus,
+ WorkspaceMembership,
+ WorkspaceSource,
+ WorkspaceStatus,
+ WorkspaceType,
+)
+from ..utils import constants
+from .errors import (
+ WorkspaceExecutionUnavailableError,
+ WorkspaceGenerationMismatchError,
+ WorkspaceInvariantError,
+ WorkspaceNotFoundError,
+ WorkspaceOwnerAlreadyExistsError,
+)
+from .entities import WorkspaceExecutionBinding
+from .policy import CloudWorkspacePolicy, SingleWorkspacePolicy
+from .repository import WorkspaceRepository
+
+
+T = TypeVar('T')
+
+if typing.TYPE_CHECKING:
+ from ..core.app import Application
+
+
+class WorkspaceService:
+ """Local workspace lifecycle service used by OSS bootstrap and account flows."""
+
+ def __init__(
+ self,
+ ap: Application,
+ *,
+ policy: SingleWorkspacePolicy | CloudWorkspacePolicy | None = None,
+ instance_uuid: str | None = None,
+ ) -> None:
+ self.ap = ap
+ self.policy = policy or SingleWorkspacePolicy()
+ self._instance_uuid = instance_uuid
+
+ @property
+ def instance_uuid(self) -> str:
+ instance_uuid = (self._instance_uuid or constants.instance_id).strip()
+ if not instance_uuid:
+ raise WorkspaceInvariantError('LangBot instance UUID is empty')
+ return instance_uuid
+
+ async def get_workspace(
+ self,
+ workspace_uuid: str,
+ *,
+ session: AsyncSession | None = None,
+ ) -> Workspace:
+ """Load one Workspace projected onto this LangBot instance."""
+
+ async def operation(repository: WorkspaceRepository) -> Workspace:
+ workspace = await repository.get_workspace(workspace_uuid)
+ if workspace is None or workspace.instance_uuid != self.instance_uuid:
+ raise WorkspaceNotFoundError('Workspace not found')
+ return workspace
+
+ return await self._run(operation, session=session)
+
+ async def get_singleton_workspace(self, *, session: AsyncSession | None = None) -> Workspace:
+ async def operation(repository: WorkspaceRepository) -> Workspace:
+ workspaces = await repository.list_local_workspaces(self.instance_uuid)
+ if not workspaces:
+ raise WorkspaceNotFoundError('The local workspace has not been initialized')
+ if len(workspaces) != 1:
+ raise WorkspaceInvariantError(
+ f'Expected one local workspace for {self.instance_uuid!r}, found {len(workspaces)}'
+ )
+ return workspaces[0]
+
+ return await self._run(operation, session=session)
+
+ async def get_execution_state(
+ self,
+ workspace_uuid: str,
+ *,
+ session: AsyncSession | None = None,
+ ) -> WorkspaceExecutionState:
+ """Load a Workspace execution state and validate its instance binding."""
+
+ async def operation(repository: WorkspaceRepository) -> WorkspaceExecutionState:
+ execution_state = await repository.get_execution_state(workspace_uuid)
+ if execution_state is None:
+ raise WorkspaceExecutionUnavailableError(f'Workspace {workspace_uuid!r} has no execution state')
+ if execution_state.instance_uuid != self.instance_uuid:
+ raise WorkspaceInvariantError(
+ f'Workspace {workspace_uuid!r} execution state belongs to another instance'
+ )
+ return execution_state
+
+ return await self._run(operation, session=session)
+
+ async def get_execution_binding(
+ self,
+ workspace_uuid: str | None = None,
+ *,
+ expected_generation: int | None = None,
+ session: AsyncSession | None = None,
+ _require_local: bool = False,
+ ) -> WorkspaceExecutionBinding:
+ """Resolve an active, unfenced binding projected onto this instance.
+
+ SaaS Workspaces are projected by a closed control plane, but Core still
+ validates the local projection and execution fence. Callers never infer
+ a Workspace from source or recency.
+ """
+
+ async def operation(repository: WorkspaceRepository) -> WorkspaceExecutionBinding:
+ if workspace_uuid is None:
+ workspaces = await repository.list_local_workspaces(self.instance_uuid)
+ if not workspaces:
+ raise WorkspaceNotFoundError('The local workspace has not been initialized')
+ if len(workspaces) != 1:
+ raise WorkspaceInvariantError(
+ f'Expected one local workspace for {self.instance_uuid!r}, found {len(workspaces)}'
+ )
+ workspace = workspaces[0]
+ else:
+ workspace = await repository.get_workspace(workspace_uuid)
+ if workspace is None:
+ raise WorkspaceNotFoundError(f'Workspace {workspace_uuid!r} does not exist')
+
+ if workspace.instance_uuid != self.instance_uuid:
+ raise WorkspaceInvariantError(f'Workspace {workspace.uuid!r} belongs to another instance')
+ if _require_local and workspace.source != WorkspaceSource.LOCAL.value:
+ raise WorkspaceInvariantError(f'Workspace {workspace.uuid!r} is not an OSS local workspace')
+ if workspace.status != WorkspaceStatus.ACTIVE.value:
+ raise WorkspaceExecutionUnavailableError(f'Workspace {workspace.uuid!r} is not active')
+
+ execution_state = await repository.get_execution_state(workspace.uuid)
+ if execution_state is None:
+ raise WorkspaceExecutionUnavailableError(f'Workspace {workspace.uuid!r} has no execution state')
+ if execution_state.instance_uuid != self.instance_uuid:
+ raise WorkspaceInvariantError(
+ f'Workspace {workspace.uuid!r} execution state belongs to another instance'
+ )
+ expected_source = (
+ WorkspaceExecutionSource.LOCAL.value
+ if workspace.source == WorkspaceSource.LOCAL.value
+ else WorkspaceExecutionSource.CLOUD.value
+ )
+ if execution_state.source != expected_source:
+ raise WorkspaceInvariantError(
+ f'Workspace {workspace.uuid!r} execution source does not match its directory source'
+ )
+ if execution_state.state != WorkspaceExecutionStatus.ACTIVE.value or execution_state.write_fenced:
+ raise WorkspaceExecutionUnavailableError(f'Workspace {workspace.uuid!r} execution is unavailable')
+ if execution_state.active_generation <= 0:
+ raise WorkspaceInvariantError(f'Workspace {workspace.uuid!r} has an invalid execution generation')
+ if expected_generation is not None and execution_state.active_generation != expected_generation:
+ raise WorkspaceGenerationMismatchError(
+ f'Workspace {workspace.uuid!r} generation {execution_state.active_generation} '
+ f'does not match expected generation {expected_generation}'
+ )
+
+ return WorkspaceExecutionBinding(
+ instance_uuid=self.instance_uuid,
+ workspace_uuid=workspace.uuid,
+ placement_generation=execution_state.active_generation,
+ write_fenced=execution_state.write_fenced,
+ state=execution_state.state,
+ )
+
+ return await self._run(operation, session=session)
+
+ async def get_local_execution_binding(
+ self,
+ workspace_uuid: str | None = None,
+ *,
+ expected_generation: int | None = None,
+ session: AsyncSession | None = None,
+ ) -> WorkspaceExecutionBinding:
+ """Resolve an active binding and require an OSS-local Workspace."""
+
+ return await self.get_execution_binding(
+ workspace_uuid,
+ expected_generation=expected_generation,
+ session=session,
+ _require_local=True,
+ )
+
+ async def get_local_execution_context(
+ self,
+ workspace_uuid: str | None = None,
+ *,
+ expected_generation: int | None = None,
+ session: AsyncSession | None = None,
+ ) -> WorkspaceExecutionBinding:
+ """Compatibility alias for callers introduced during the tenancy rollout."""
+ return await self.get_local_execution_binding(
+ workspace_uuid,
+ expected_generation=expected_generation,
+ session=session,
+ )
+
+ async def ensure_singleton_workspace(
+ self,
+ *,
+ session: AsyncSession | None = None,
+ name: str = 'Default Workspace',
+ slug: str = 'default',
+ ) -> Workspace:
+ """Create or repair the instance's single local workspace and execution state."""
+
+ async def operation(repository: WorkspaceRepository) -> Workspace:
+ workspaces = await repository.list_local_workspaces(self.instance_uuid, for_update=True)
+ if len(workspaces) > 1:
+ raise WorkspaceInvariantError(f'Multiple local workspaces exist for instance {self.instance_uuid!r}')
+ if workspaces:
+ workspace = workspaces[0]
+ else:
+ self.policy.require_workspace_creation_allowed(0)
+ workspace = self._new_local_workspace(name=name, slug=slug)
+ repository.add_workspace(workspace)
+ await repository.flush()
+
+ await self._ensure_execution_state(repository, workspace)
+ return workspace
+
+ return await self._run(operation, session=session)
+
+ async def create_local_workspace(
+ self,
+ *,
+ name: str,
+ slug: str,
+ created_by_account_uuid: str | None = None,
+ session: AsyncSession | None = None,
+ ) -> Workspace:
+ """Create the OSS workspace, enforcing the one-workspace edition limit."""
+
+ async def operation(repository: WorkspaceRepository) -> Workspace:
+ current_count = await repository.count_local_workspaces(self.instance_uuid)
+ self.policy.require_workspace_creation_allowed(current_count)
+
+ workspace = self._new_local_workspace(
+ name=name,
+ slug=slug,
+ created_by_account_uuid=created_by_account_uuid,
+ )
+ repository.add_workspace(workspace)
+ await repository.flush()
+ await self._ensure_execution_state(repository, workspace)
+ if created_by_account_uuid is not None:
+ await self._claim_initial_owner(repository, workspace, created_by_account_uuid)
+ return workspace
+
+ return await self._run(operation, session=session)
+
+ async def claim_initial_owner(
+ self,
+ account_uuid: str,
+ *,
+ session: AsyncSession | None = None,
+ ) -> WorkspaceMembership:
+ """Atomically claim an ownerless singleton workspace for the first account."""
+
+ async def operation(repository: WorkspaceRepository) -> WorkspaceMembership:
+ workspaces = await repository.list_local_workspaces(self.instance_uuid, for_update=True)
+ if not workspaces:
+ self.policy.require_workspace_creation_allowed(0)
+ workspace = self._new_local_workspace(name='Default Workspace', slug='default')
+ repository.add_workspace(workspace)
+ await repository.flush()
+ elif len(workspaces) == 1:
+ workspace = workspaces[0]
+ else:
+ raise WorkspaceInvariantError(f'Multiple local workspaces exist for instance {self.instance_uuid!r}')
+ await self._ensure_execution_state(repository, workspace)
+ return await self._claim_initial_owner(repository, workspace, account_uuid)
+
+ return await self._run(operation, session=session)
+
+ async def bootstrap_local_account(
+ self,
+ account_uuid: str,
+ *,
+ session: AsyncSession | None = None,
+ ) -> tuple[Workspace, WorkspaceMembership]:
+ """Bind the first local account to the singleton workspace as owner."""
+
+ async def operation(repository: WorkspaceRepository) -> tuple[Workspace, WorkspaceMembership]:
+ workspaces = await repository.list_local_workspaces(self.instance_uuid, for_update=True)
+ if not workspaces:
+ self.policy.require_workspace_creation_allowed(0)
+ workspace = self._new_local_workspace(name='Default Workspace', slug='default')
+ repository.add_workspace(workspace)
+ await repository.flush()
+ elif len(workspaces) == 1:
+ workspace = workspaces[0]
+ else:
+ raise WorkspaceInvariantError(f'Multiple local workspaces exist for instance {self.instance_uuid!r}')
+
+ await self._ensure_execution_state(repository, workspace)
+ membership = await self._claim_initial_owner(repository, workspace, account_uuid)
+ return workspace, membership
+
+ return await self._run(operation, session=session)
+
+ async def _claim_initial_owner(
+ self,
+ repository: WorkspaceRepository,
+ workspace: Workspace,
+ account_uuid: str,
+ ) -> WorkspaceMembership:
+ active_owner = await repository.get_active_owner(workspace.uuid, for_update=True)
+ if active_owner is not None and active_owner.account_uuid != account_uuid:
+ raise WorkspaceOwnerAlreadyExistsError(f'Workspace {workspace.uuid!r} already has an owner')
+ if active_owner is not None:
+ if workspace.created_by_account_uuid is None:
+ workspace.created_by_account_uuid = account_uuid
+ await repository.flush()
+ return active_owner
+
+ membership = await repository.get_membership(workspace.uuid, account_uuid)
+ joined_at = datetime.datetime.now(datetime.UTC).replace(tzinfo=None)
+ if membership is None:
+ membership = WorkspaceMembership(
+ uuid=str(uuid.uuid4()),
+ workspace_uuid=workspace.uuid,
+ account_uuid=account_uuid,
+ role=MembershipRole.OWNER.value,
+ status=MembershipStatus.ACTIVE.value,
+ joined_at=joined_at,
+ projection_revision=0,
+ )
+ repository.add_membership(membership)
+ else:
+ membership.role = MembershipRole.OWNER.value
+ membership.status = MembershipStatus.ACTIVE.value
+ membership.joined_at = membership.joined_at or joined_at
+
+ if workspace.created_by_account_uuid is None:
+ workspace.created_by_account_uuid = account_uuid
+ await repository.flush()
+ return membership
+
+ async def _ensure_execution_state(
+ self,
+ repository: WorkspaceRepository,
+ workspace: Workspace,
+ ) -> WorkspaceExecutionState:
+ execution_state = await repository.get_execution_state(workspace.uuid)
+ if execution_state is not None:
+ if execution_state.instance_uuid != self.instance_uuid:
+ raise WorkspaceInvariantError(
+ f'Workspace {workspace.uuid!r} execution state belongs to another instance'
+ )
+ return execution_state
+
+ execution_state = WorkspaceExecutionState(
+ workspace_uuid=workspace.uuid,
+ instance_uuid=self.instance_uuid,
+ active_generation=1,
+ state=WorkspaceExecutionStatus.ACTIVE.value,
+ write_fenced=False,
+ source=WorkspaceExecutionSource.LOCAL.value,
+ desired_state_revision=0,
+ )
+ repository.add_execution_state(execution_state)
+ await repository.flush()
+ return execution_state
+
+ def _new_local_workspace(
+ self,
+ *,
+ name: str,
+ slug: str,
+ created_by_account_uuid: str | None = None,
+ ) -> Workspace:
+ return Workspace(
+ uuid=str(uuid.uuid4()),
+ instance_uuid=self.instance_uuid,
+ name=name,
+ slug=slug,
+ type=WorkspaceType.TEAM.value,
+ status=WorkspaceStatus.ACTIVE.value,
+ created_by_account_uuid=created_by_account_uuid,
+ source=WorkspaceSource.LOCAL.value,
+ projection_revision=0,
+ )
+
+ async def _run(
+ self,
+ operation: Callable[[WorkspaceRepository], Awaitable[T]],
+ *,
+ session: AsyncSession | None,
+ ) -> T:
+ if session is not None:
+ return await operation(WorkspaceRepository(session))
+
+ session_factory = async_sessionmaker(
+ self.ap.persistence_mgr.get_db_engine(),
+ expire_on_commit=False,
+ )
+ async with session_factory() as owned_session:
+ async with owned_session.begin():
+ return await operation(WorkspaceRepository(owned_session))
diff --git a/src/langbot/templates/config.yaml b/src/langbot/templates/config.yaml
index f4ae79bf8..944c24d13 100644
--- a/src/langbot/templates/config.yaml
+++ b/src/langbot/templates/config.yaml
@@ -2,6 +2,12 @@ api:
port: 5300
webhook_prefix: 'http://127.0.0.1:5300'
extra_webhook_prefix: ''
+ # Canonical browser origin when WebUI and API use different origins in
+ # development (for example http://localhost:3000). Production bundled UI
+ # may leave this empty when webhook_prefix already has the browser origin.
+ # OAuth redirects trust only these server-side values, never request Host
+ # or Origin headers.
+ webui_url: ''
# Global API key for the HTTP service API and the MCP server. When set to a
# non-empty string, this key is accepted anywhere a web-UI-created API key is
# accepted (X-API-Key header or "Authorization: Bearer "), WITHOUT any
@@ -142,6 +148,8 @@ box:
enabled: true
backend: 'local' # 'local' (Docker/nsjail), 'docker', 'nsjail', or 'e2b'. Can be written via BOX__BACKEND.
runtime:
+ # External WebSocket runtimes also require LANGBOT_BOX_CONTROL_TOKEN in
+ # both LangBot and Box. Keep the shared secret out of this config file.
endpoint: '' # External Box Runtime base URL, e.g. 'ws://127.0.0.1:5410'. Leave empty for local auto-managed runtime.
local:
profile: 'default'
diff --git a/src/langbot/templates/embed/widget.js b/src/langbot/templates/embed/widget.js
index e08ffe180..2a710711b 100644
--- a/src/langbot/templates/embed/widget.js
+++ b/src/langbot/templates/embed/widget.js
@@ -356,6 +356,7 @@
isConnected: false,
ws: null,
connectionId: null,
+ sessionToken: null,
sessionId: getOrCreateSessionId(),
reconnectAttempts: 0,
heartbeatTimer: null,
@@ -538,7 +539,12 @@
state.ws.onopen = function () {
state.reconnectAttempts = 0;
- startHeartbeat();
+ state.ws.send(
+ JSON.stringify({
+ type: "authenticate",
+ token: state.sessionToken || "",
+ }),
+ );
};
state.ws.onmessage = function (event) {
@@ -576,6 +582,7 @@
state.connectionId = data.connection_id;
if (state.hasConnected) loadHistory(true);
state.hasConnected = true;
+ startHeartbeat();
updateStatusDot();
updateSendBtn();
break;
diff --git a/tests/factories/app.py b/tests/factories/app.py
index d1edf56a2..3832210af 100644
--- a/tests/factories/app.py
+++ b/tests/factories/app.py
@@ -125,6 +125,8 @@ class FakeApp:
"""Mock SkillManager that returns no skill index addition by default."""
skill_mgr = Mock()
skill_mgr.skills = {}
+ skill_mgr.ensure_loaded = AsyncMock()
+ skill_mgr.get_skills = Mock(return_value=[])
skill_mgr.build_skill_aware_prompt_addition = Mock(return_value='')
skill_mgr.get_skill_index = Mock(return_value=[])
return skill_mgr
diff --git a/tests/factories/message.py b/tests/factories/message.py
index 9b3cc3602..619354e2c 100644
--- a/tests/factories/message.py
+++ b/tests/factories/message.py
@@ -14,6 +14,7 @@ import langbot_plugin.api.entities.builtin.platform.message as platform_message
import langbot_plugin.api.entities.builtin.platform.events as platform_events
import langbot_plugin.api.entities.builtin.platform.entities as platform_entities
import langbot_plugin.api.entities.builtin.provider.session as provider_session
+from langbot.pkg.api.http.context import ExecutionContext
# Counter for generating unique IDs
@@ -194,7 +195,20 @@ def _base_query(
for key, value in overrides.items():
base_data[key] = value
- return pipeline_query.Query.model_construct(**base_data)
+ query = pipeline_query.Query.model_construct(**base_data)
+ object.__setattr__(
+ query,
+ '_execution_context',
+ ExecutionContext(
+ instance_uuid='test-instance',
+ workspace_uuid='test-workspace',
+ placement_generation=1,
+ bot_uuid=query.bot_uuid,
+ pipeline_uuid=query.pipeline_uuid,
+ query_uuid=query.query_uuid,
+ ),
+ )
+ return query
def text_query(
diff --git a/tests/integration/api/test_bots.py b/tests/integration/api/test_bots.py
index 0e6854bf9..fe40b6291 100644
--- a/tests/integration/api/test_bots.py
+++ b/tests/integration/api/test_bots.py
@@ -10,6 +10,7 @@ from __future__ import annotations
import pytest
from unittest.mock import MagicMock, AsyncMock, Mock
+from types import SimpleNamespace
from tests.factories import FakeApp
@@ -65,12 +66,43 @@ def fake_bot_app():
)
# Auth services
+ account = SimpleNamespace(uuid='account-test', user='test@example.com')
app.user_service = Mock()
app.user_service.is_initialized = AsyncMock(return_value=True)
app.user_service.verify_jwt_token = AsyncMock(return_value='test@example.com')
- app.user_service.get_user_by_email = AsyncMock(return_value=Mock(email='test@example.com'))
+ app.user_service.get_user_by_email = AsyncMock(return_value=account)
+ app.user_service.get_authenticated_account = AsyncMock(return_value=account)
+ app.workspace_collaboration_service = SimpleNamespace(
+ resolve_account_workspace=AsyncMock(
+ return_value=SimpleNamespace(
+ workspace=SimpleNamespace(uuid='workspace-test'),
+ membership=SimpleNamespace(
+ uuid='membership-test',
+ role='owner',
+ projection_revision=0,
+ ),
+ execution=SimpleNamespace(instance_uuid='instance-test', placement_generation=1),
+ )
+ )
+ )
app.apikey_service = Mock()
app.apikey_service.verify_api_key = AsyncMock(return_value=True)
+ app.apikey_service.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',
+ 'resource.manage',
+ 'runtime.operate',
+ 'provider_secret.manage',
+ }
+ ),
+ )
+ )
# Bot service
app.bot_service = Mock()
@@ -198,6 +230,25 @@ class TestBotLogsEndpoint:
assert 'logs' in data['data']
assert 'total_count' in data['data']
+ @pytest.mark.asyncio
+ async def test_viewer_cannot_read_message_logs(self, quart_test_client, fake_bot_app):
+ access = fake_bot_app.workspace_collaboration_service.resolve_account_workspace.return_value
+ original_role = access.membership.role
+ access.membership.role = 'viewer'
+ fake_bot_app.bot_service.list_event_logs.reset_mock()
+ try:
+ response = await quart_test_client.post(
+ '/api/v1/platform/bots/test-bot-uuid/logs',
+ headers={'Authorization': 'Bearer test_token'},
+ json={'from_index': -1, 'max_count': 10},
+ )
+ finally:
+ access.membership.role = original_role
+
+ assert response.status_code == 403
+ assert (await response.get_json())['code'] == 'permission_denied'
+ fake_bot_app.bot_service.list_event_logs.assert_not_awaited()
+
@pytest.mark.usefixtures('mock_circular_import_chain')
class TestBotSendMessageEndpoint:
diff --git a/tests/integration/api/test_box_security.py b/tests/integration/api/test_box_security.py
new file mode 100644
index 000000000..79f803b28
--- /dev/null
+++ b/tests/integration/api/test_box_security.py
@@ -0,0 +1,83 @@
+"""Authorization tests for sensitive Box runtime observability."""
+
+from __future__ import annotations
+
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, Mock
+
+import pytest
+import quart
+
+from langbot.pkg.api.http.controller.groups.box import BoxRouterGroup
+
+
+pytestmark = pytest.mark.integration
+WORKSPACE_UUID = '11111111-1111-4111-8111-111111111111'
+
+
+def _access(account_uuid: str):
+ return SimpleNamespace(
+ workspace=SimpleNamespace(uuid=WORKSPACE_UUID),
+ membership=SimpleNamespace(
+ uuid=f'member-{account_uuid}',
+ role='viewer' if account_uuid == 'viewer-account' else 'owner',
+ projection_revision=1,
+ ),
+ execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=1),
+ )
+
+
+@pytest.fixture
+async def box_security_api():
+ accounts = {
+ 'viewer-token': SimpleNamespace(uuid='viewer-account', user='viewer@example.com'),
+ 'owner-token': SimpleNamespace(uuid='owner-account', user='owner@example.com'),
+ }
+ application = Mock()
+ application.user_service.get_authenticated_account = AsyncMock(side_effect=lambda token: accounts[token])
+ application.workspace_collaboration_service.resolve_account_workspace = AsyncMock(
+ side_effect=lambda account_uuid, _workspace_uuid: _access(account_uuid)
+ )
+ application.box_service.get_status = AsyncMock(return_value={'enabled': True})
+ application.box_service.get_sessions = AsyncMock(return_value=[{'session_id': 'private-session'}])
+ application.box_service.get_recent_errors = Mock(return_value=[{'error': 'private error'}])
+
+ quart_app = quart.Quart(__name__)
+ router = BoxRouterGroup(application, quart_app)
+ await router.initialize()
+ return application, quart_app.test_client()
+
+
+def _headers(token: str) -> dict[str, str]:
+ return {
+ 'Authorization': f'Bearer {token}',
+ 'X-Workspace-Id': WORKSPACE_UUID,
+ }
+
+
+@pytest.mark.asyncio
+async def test_viewer_can_read_status_but_not_sessions_or_errors(box_security_api):
+ application, client = box_security_api
+
+ status = await client.get('/api/v1/box/status', headers=_headers('viewer-token'))
+ sessions = await client.get('/api/v1/box/sessions', headers=_headers('viewer-token'))
+ errors = await client.get('/api/v1/box/errors', headers=_headers('viewer-token'))
+
+ assert status.status_code == 200
+ assert sessions.status_code == 403
+ assert errors.status_code == 403
+ application.box_service.get_sessions.assert_not_awaited()
+ application.box_service.get_recent_errors.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_owner_can_audit_box_sessions_and_errors(box_security_api):
+ application, client = box_security_api
+
+ sessions = await client.get('/api/v1/box/sessions', headers=_headers('owner-token'))
+ errors = await client.get('/api/v1/box/errors', headers=_headers('owner-token'))
+
+ assert sessions.status_code == 200
+ assert errors.status_code == 200
+ application.box_service.get_sessions.assert_awaited_once()
+ application.box_service.get_recent_errors.assert_called_once()
diff --git a/tests/integration/api/test_embed.py b/tests/integration/api/test_embed.py
index 2862337f5..b31d6fbe2 100644
--- a/tests/integration/api/test_embed.py
+++ b/tests/integration/api/test_embed.py
@@ -8,8 +8,11 @@ Run: uv run pytest tests/integration/api/test_embed.py -q
from __future__ import annotations
+import json
+
import pytest
from unittest.mock import MagicMock, AsyncMock, Mock
+from types import SimpleNamespace
from tests.factories import FakeApp
@@ -80,10 +83,18 @@ def fake_embed_app():
mock_runtime_bot = Mock()
mock_runtime_bot.bot_entity = mock_bot_entity
+ mock_runtime_bot.execution_context = SimpleNamespace(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-test',
+ placement_generation=1,
+ )
# Platform manager with bots
app.platform_mgr = Mock()
app.platform_mgr.bots = [mock_runtime_bot]
+ app.platform_mgr.resolve_public_bot = AsyncMock(
+ side_effect=lambda route_key: mock_runtime_bot if route_key == mock_bot_entity.uuid else None
+ )
# WebSocket proxy bot with adapter
mock_websocket_adapter = Mock()
@@ -94,6 +105,16 @@ def fake_embed_app():
mock_ws_proxy_bot = Mock()
mock_ws_proxy_bot.adapter = mock_websocket_adapter
app.platform_mgr.websocket_proxy_bot = mock_ws_proxy_bot
+ app.platform_mgr.get_websocket_proxy_bot = AsyncMock(return_value=mock_ws_proxy_bot)
+ app.workspace_service = SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-test',
+ placement_generation=1,
+ )
+ )
+ )
# Monitoring service for feedback
app.monitoring_service = Mock()
@@ -117,12 +138,13 @@ class TestEmbedWidgetEndpoint:
"""Tests for widget.js endpoint."""
@pytest.mark.asyncio
- async def test_get_widget_js_success(self, quart_test_client):
+ async def test_get_widget_js_success(self, quart_test_client, fake_embed_app):
"""GET /api/v1/embed/{bot_uuid}/widget.js returns JS."""
response = await quart_test_client.get('/api/v1/embed/a1b2c3d4-5678-90ab-cdef-123456789abc/widget.js')
assert response.status_code == 200
assert 'javascript' in response.content_type
+ fake_embed_app.platform_mgr.resolve_public_bot.assert_any_await('a1b2c3d4-5678-90ab-cdef-123456789abc')
@pytest.mark.asyncio
async def test_get_widget_js_invalid_uuid(self, quart_test_client):
@@ -203,9 +225,8 @@ class TestEmbedMessagesEndpoint:
data = await response.get_json()
assert data['code'] == 0
assert 'messages' in data['data']
- fake_embed_app.platform_mgr.websocket_proxy_bot.adapter.get_websocket_messages.assert_called_with(
- 'test-pipeline-uuid', 'person', SESSION_ID
- )
+ proxy_bot = fake_embed_app.platform_mgr.get_websocket_proxy_bot.return_value
+ proxy_bot.adapter.get_websocket_messages.assert_called_with('test-pipeline-uuid', 'person', SESSION_ID)
@pytest.mark.asyncio
async def test_get_messages_group_success(self, quart_test_client):
@@ -253,9 +274,8 @@ class TestEmbedResetEndpoint:
assert response.status_code == 200
data = await response.get_json()
assert data['code'] == 0
- fake_embed_app.platform_mgr.websocket_proxy_bot.adapter.reset_session.assert_called_with(
- 'test-pipeline-uuid', 'person', SESSION_ID
- )
+ proxy_bot = fake_embed_app.platform_mgr.get_websocket_proxy_bot.return_value
+ proxy_bot.adapter.reset_session.assert_called_with('test-pipeline-uuid', 'person', SESSION_ID)
@pytest.mark.asyncio
async def test_reset_session_requires_session_id(self, quart_test_client):
@@ -316,3 +336,85 @@ class TestEmbedFeedbackEndpoint:
)
assert response.status_code == 400
+
+
+@pytest.mark.usefixtures('mock_circular_import_chain')
+class TestEmbedWebSocketEndpoint:
+ """The public socket authenticates before resolving shared runtime state."""
+
+ @pytest.mark.asyncio
+ async def test_authenticates_before_connecting(self, quart_test_client, fake_embed_app):
+ async with quart_test_client.websocket(
+ f'/api/v1/embed/a1b2c3d4-5678-90ab-cdef-123456789abc/ws/connect'
+ f'?session_type=person&session_id={SESSION_ID}',
+ headers={'Origin': 'http://localhost'},
+ ) as websocket:
+ await websocket.send(json.dumps({'type': 'authenticate', 'token': ''}))
+ connected = json.loads(await websocket.receive())
+ assert connected['type'] == 'connected'
+ assert connected['bot_uuid'] == 'a1b2c3d4-5678-90ab-cdef-123456789abc'
+ await websocket.send(json.dumps({'type': 'disconnect'}))
+
+ fake_embed_app.workspace_service.get_execution_binding.assert_awaited_with(
+ 'workspace-test',
+ expected_generation=1,
+ )
+
+ @pytest.mark.asyncio
+ async def test_rejects_non_auth_first_frame_before_runtime_lookup(self, quart_test_client, fake_embed_app):
+ fake_embed_app.platform_mgr.get_websocket_proxy_bot.reset_mock()
+
+ async with quart_test_client.websocket(
+ f'/api/v1/embed/a1b2c3d4-5678-90ab-cdef-123456789abc/ws/connect'
+ f'?session_type=person&session_id={SESSION_ID}',
+ headers={'Origin': 'http://localhost'},
+ ) as websocket:
+ await websocket.send(json.dumps({'type': 'message', 'message': []}))
+ response = json.loads(await websocket.receive())
+ assert response == {'type': 'error', 'message': 'Unauthorized'}
+
+ fake_embed_app.platform_mgr.get_websocket_proxy_bot.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_rejects_invalid_turnstile_session_before_runtime_lookup(self, quart_test_client, fake_embed_app):
+ fake_embed_app.platform_mgr.get_websocket_proxy_bot.reset_mock()
+ config = fake_embed_app.platform_mgr.resolve_public_bot.side_effect(
+ 'a1b2c3d4-5678-90ab-cdef-123456789abc'
+ ).bot_entity.adapter_config
+ config['turnstile_secret_key'] = 'test-secret'
+ try:
+ async with quart_test_client.websocket(
+ f'/api/v1/embed/a1b2c3d4-5678-90ab-cdef-123456789abc/ws/connect'
+ f'?session_type=person&session_id={SESSION_ID}',
+ headers={'Origin': 'http://localhost'},
+ ) as websocket:
+ await websocket.send(json.dumps({'type': 'authenticate', 'token': 'invalid'}))
+ response = json.loads(await websocket.receive())
+ assert response == {'type': 'error', 'message': 'Unauthorized'}
+ finally:
+ config['turnstile_secret_key'] = ''
+
+ fake_embed_app.platform_mgr.get_websocket_proxy_bot.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_rejects_message_when_bot_is_disabled_after_connect(self, quart_test_client, fake_embed_app):
+ runtime_bot = fake_embed_app.platform_mgr.resolve_public_bot.side_effect('a1b2c3d4-5678-90ab-cdef-123456789abc')
+ adapter = fake_embed_app.platform_mgr.get_websocket_proxy_bot.return_value.adapter
+ adapter.handle_websocket_message.reset_mock()
+
+ async with quart_test_client.websocket(
+ f'/api/v1/embed/a1b2c3d4-5678-90ab-cdef-123456789abc/ws/connect'
+ f'?session_type=person&session_id={SESSION_ID}',
+ headers={'Origin': 'http://localhost'},
+ ) as websocket:
+ await websocket.send(json.dumps({'type': 'authenticate', 'token': ''}))
+ assert json.loads(await websocket.receive())['type'] == 'connected'
+ runtime_bot.bot_entity.enable = False
+ try:
+ await websocket.send(json.dumps({'type': 'message', 'message': [{'type': 'text', 'text': 'hi'}]}))
+ response = json.loads(await websocket.receive())
+ assert response == {'type': 'error', 'message': 'Bot is unavailable'}
+ finally:
+ runtime_bot.bot_entity.enable = True
+
+ adapter.handle_websocket_message.assert_not_awaited()
diff --git a/tests/integration/api/test_knowledge.py b/tests/integration/api/test_knowledge.py
index 973356c3e..13b0785c7 100644
--- a/tests/integration/api/test_knowledge.py
+++ b/tests/integration/api/test_knowledge.py
@@ -10,6 +10,7 @@ from __future__ import annotations
import pytest
from unittest.mock import MagicMock, AsyncMock, Mock
+from types import SimpleNamespace
from tests.factories import FakeApp
@@ -69,7 +70,28 @@ def fake_knowledge_app():
app.user_service = Mock()
app.user_service.is_initialized = AsyncMock(return_value=True)
app.user_service.verify_jwt_token = AsyncMock(return_value='test@example.com')
- app.user_service.get_user_by_email = AsyncMock(return_value=Mock(email='test@example.com'))
+ account = SimpleNamespace(
+ uuid='00000000-0000-0000-0000-000000000001',
+ user='test@example.com',
+ )
+ app.user_service.get_authenticated_account = AsyncMock(return_value=account)
+ app.user_service.get_user_by_email = AsyncMock(return_value=account)
+ app.workspace_collaboration_service = SimpleNamespace(
+ resolve_account_workspace=AsyncMock(
+ return_value=SimpleNamespace(
+ execution=SimpleNamespace(
+ instance_uuid='instance-knowledge-api',
+ placement_generation=1,
+ ),
+ workspace=SimpleNamespace(uuid='00000000-0000-0000-0000-00000000000a'),
+ membership=SimpleNamespace(
+ uuid='00000000-0000-0000-0000-000000000010',
+ role='owner',
+ projection_revision=0,
+ ),
+ )
+ )
+ )
app.apikey_service = Mock()
app.apikey_service.verify_api_key = AsyncMock(return_value=True)
diff --git a/tests/integration/api/test_monitoring.py b/tests/integration/api/test_monitoring.py
index 64c155840..9edc1e4ba 100644
--- a/tests/integration/api/test_monitoring.py
+++ b/tests/integration/api/test_monitoring.py
@@ -10,6 +10,7 @@ from __future__ import annotations
import pytest
from unittest.mock import MagicMock, AsyncMock, Mock
+from types import SimpleNamespace
from tests.factories import FakeApp
@@ -66,7 +67,26 @@ def fake_monitoring_app():
app.user_service = Mock()
app.user_service.is_initialized = AsyncMock(return_value=True)
app.user_service.verify_jwt_token = AsyncMock(return_value='test@example.com')
- app.user_service.get_user_by_email = AsyncMock(return_value=Mock(email='test@example.com'))
+ app.user_service.get_user_by_email = AsyncMock(
+ return_value=SimpleNamespace(
+ uuid='account-uuid',
+ user='test@example.com',
+ email='test@example.com',
+ )
+ )
+ app.workspace_collaboration_service = SimpleNamespace(
+ resolve_account_workspace=AsyncMock(
+ return_value=SimpleNamespace(
+ execution=SimpleNamespace(instance_uuid='instance', placement_generation=1),
+ workspace=SimpleNamespace(uuid='00000000-0000-0000-0000-00000000000a'),
+ membership=SimpleNamespace(
+ uuid='membership-uuid',
+ role='owner',
+ projection_revision=1,
+ ),
+ )
+ )
+ )
# Monitoring service
app.monitoring_service = Mock()
diff --git a/tests/integration/api/test_pipelines.py b/tests/integration/api/test_pipelines.py
index 50ac37bc5..80fce9747 100644
--- a/tests/integration/api/test_pipelines.py
+++ b/tests/integration/api/test_pipelines.py
@@ -9,10 +9,13 @@ Run: uv run pytest tests/integration/api/test_pipelines.py -q
from __future__ import annotations
+import json
import pytest
from unittest.mock import MagicMock, AsyncMock, Mock
+from types import SimpleNamespace
from tests.factories import FakeApp
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
pytestmark = pytest.mark.integration
@@ -54,6 +57,7 @@ def mock_circular_import_chain():
):
# Import groups after mocking to populate preregistered_groups
import langbot.pkg.api.http.controller.groups.pipelines.pipelines as _pipelines # noqa: E402, F401
+ import langbot.pkg.api.http.controller.groups.pipelines.websocket_chat as _websocket_chat # noqa: E402, F401
yield
@@ -75,10 +79,25 @@ def fake_pipeline_app():
)
# Auth services
+ account = SimpleNamespace(uuid='account-test', user='test@example.com')
app.user_service = Mock()
app.user_service.is_initialized = AsyncMock(return_value=True)
app.user_service.verify_jwt_token = AsyncMock(return_value='test@example.com')
- app.user_service.get_user_by_email = AsyncMock(return_value=Mock(email='test@example.com'))
+ app.user_service.get_user_by_email = AsyncMock(return_value=account)
+ app.user_service.get_authenticated_account = AsyncMock(return_value=account)
+ app.workspace_collaboration_service = SimpleNamespace(
+ resolve_account_workspace=AsyncMock(
+ return_value=SimpleNamespace(
+ workspace=SimpleNamespace(uuid='workspace-test'),
+ membership=SimpleNamespace(
+ uuid='membership-test',
+ role='owner',
+ projection_revision=0,
+ ),
+ execution=SimpleNamespace(instance_uuid='instance-test', placement_generation=1),
+ )
+ )
+ )
app.apikey_service = Mock()
app.apikey_service.verify_api_key = AsyncMock(return_value=True)
@@ -119,6 +138,15 @@ def fake_pipeline_app():
app.bot_service.get_bots = AsyncMock(return_value=[])
app.bot_service.create_bot = AsyncMock(return_value={'uuid': 'new-bot-uuid'})
+ # Workspace-scoped dashboard WebSocket proxy
+ websocket_adapter = Mock()
+ websocket_adapter.get_websocket_messages = Mock(return_value=[])
+ websocket_adapter.reset_session = Mock()
+ websocket_adapter.handle_websocket_message = AsyncMock()
+ websocket_proxy_bot = Mock(adapter=websocket_adapter)
+ app.platform_mgr = Mock()
+ app.platform_mgr.get_websocket_proxy_bot = AsyncMock(return_value=websocket_proxy_bot)
+
# MCP service (for extensions endpoint)
app.mcp_service = Mock()
app.mcp_service.get_mcp_servers = AsyncMock(return_value=[])
@@ -278,3 +306,147 @@ class TestPipelineExtensionsEndpoint:
assert response.status_code == 200
data = await response.get_json()
assert data['code'] == 0
+
+ @pytest.mark.asyncio
+ async def test_get_extensions_redacts_available_plugin_secrets(
+ self,
+ quart_test_client,
+ fake_pipeline_app,
+ ):
+ connector = fake_pipeline_app.plugin_connector
+ raw_plugin = {
+ 'plugin_config': {'apiKey': 'plugin-secret'},
+ 'debug': {'plugin_debug_key': 'debug-secret'},
+ }
+ connector.list_plugins.return_value = [raw_plugin]
+ try:
+ response = await quart_test_client.get(
+ '/api/v1/pipelines/test-pipeline-uuid/extensions',
+ headers={'Authorization': 'Bearer test_token'},
+ )
+ finally:
+ connector.list_plugins.return_value = []
+
+ assert response.status_code == 200
+ plugin = (await response.get_json())['data']['available_plugins'][0]
+ assert plugin['plugin_config']['apiKey'] == '***'
+ assert plugin['debug']['plugin_debug_key'] == '***'
+ assert raw_plugin['plugin_config']['apiKey'] == 'plugin-secret'
+
+ @pytest.mark.asyncio
+ async def test_get_extensions_hides_connector_bound_to_another_workspace(
+ self,
+ quart_test_client,
+ fake_pipeline_app,
+ ):
+ connector = fake_pipeline_app.plugin_connector
+ original_enabled = connector.is_enable_plugin
+ connector.is_enable_plugin = True
+ connector.require_workspace_context.reset_mock()
+ connector.list_plugins.reset_mock()
+ connector.require_workspace_context.side_effect = WorkspaceNotFoundError('Plugin resource not found')
+ try:
+ response = await quart_test_client.get(
+ '/api/v1/pipelines/test-pipeline-uuid/extensions',
+ headers={'Authorization': 'Bearer test_token'},
+ )
+ finally:
+ connector.require_workspace_context.side_effect = None
+ connector.is_enable_plugin = original_enabled
+
+ assert response.status_code == 404
+ connector.list_plugins.assert_not_awaited()
+
+
+@pytest.mark.usefixtures('mock_circular_import_chain')
+class TestPipelineDashboardWebSocket:
+ @pytest.mark.asyncio
+ async def test_websocket_authenticates_before_registering(self, quart_test_client, fake_pipeline_app):
+ async with quart_test_client.websocket(
+ '/api/v1/pipelines/test-pipeline-uuid/ws/connect?session_type=person',
+ headers={'Origin': 'http://localhost'},
+ ) as websocket:
+ await websocket.send(
+ json.dumps(
+ {
+ 'type': 'authenticate',
+ 'token': 'test_token',
+ 'workspace_uuid': 'workspace-test',
+ }
+ )
+ )
+ connected = json.loads(await websocket.receive())
+ assert connected['type'] == 'connected'
+ assert connected['pipeline_uuid'] == 'test-pipeline-uuid'
+ await websocket.send(json.dumps({'type': 'disconnect'}))
+
+ fake_pipeline_app.workspace_collaboration_service.resolve_account_workspace.assert_awaited_with(
+ 'account-test',
+ 'workspace-test',
+ )
+ fake_pipeline_app.platform_mgr.get_websocket_proxy_bot.assert_awaited()
+
+ @pytest.mark.asyncio
+ async def test_websocket_rejects_non_auth_first_frame(self, quart_test_client):
+ async with quart_test_client.websocket(
+ '/api/v1/pipelines/test-pipeline-uuid/ws/connect?session_type=person',
+ headers={'Origin': 'http://localhost'},
+ ) as websocket:
+ await websocket.send(json.dumps({'type': 'message', 'message': []}))
+ response = json.loads(await websocket.receive())
+ assert response == {'type': 'error', 'message': 'Unauthorized'}
+
+ @pytest.mark.asyncio
+ async def test_websocket_rechecks_revocable_membership_before_each_message(
+ self,
+ quart_test_client,
+ fake_pipeline_app,
+ ):
+ access = fake_pipeline_app.workspace_collaboration_service.resolve_account_workspace.return_value
+ adapter = fake_pipeline_app.platform_mgr.get_websocket_proxy_bot.return_value.adapter
+ original_role = access.membership.role
+ adapter.handle_websocket_message.reset_mock()
+ try:
+ async with quart_test_client.websocket(
+ '/api/v1/pipelines/test-pipeline-uuid/ws/connect?session_type=person',
+ headers={'Origin': 'http://localhost'},
+ ) as websocket:
+ await websocket.send(
+ json.dumps(
+ {
+ 'type': 'authenticate',
+ 'token': 'test_token',
+ 'workspace_uuid': 'workspace-test',
+ }
+ )
+ )
+ assert json.loads(await websocket.receive())['type'] == 'connected'
+
+ access.membership.role = 'viewer'
+ await websocket.send(json.dumps({'type': 'message', 'message': []}))
+ assert json.loads(await websocket.receive()) == {
+ 'type': 'error',
+ 'message': 'Unauthorized',
+ }
+ finally:
+ access.membership.role = original_role
+
+ adapter.handle_websocket_message.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_dashboard_history_requires_runtime_permission(self, quart_test_client, fake_pipeline_app):
+ access = fake_pipeline_app.workspace_collaboration_service.resolve_account_workspace.return_value
+ original_role = access.membership.role
+ access.membership.role = 'viewer'
+ try:
+ response = await quart_test_client.get(
+ '/api/v1/pipelines/test-pipeline-uuid/ws/messages/person',
+ headers={
+ 'Authorization': 'Bearer test_token',
+ 'X-Workspace-Id': 'workspace-test',
+ },
+ )
+ finally:
+ access.membership.role = original_role
+
+ assert response.status_code == 403
diff --git a/tests/integration/api/test_plugins_security.py b/tests/integration/api/test_plugins_security.py
new file mode 100644
index 000000000..89f89aeb4
--- /dev/null
+++ b/tests/integration/api/test_plugins_security.py
@@ -0,0 +1,248 @@
+"""Security regression tests for plugin configuration HTTP responses."""
+
+from __future__ import annotations
+
+import copy
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, Mock, call
+
+import pytest
+import quart
+
+
+pytestmark = pytest.mark.integration
+WORKSPACE_UUID = '11111111-1111-4111-8111-111111111111'
+RAW_CONFIG = {
+ 'apiKey': 'api-secret',
+ 'nested': {
+ 'headers': {'Authorization': 'Bearer nested-secret', 'Accept': 'application/json'},
+ 'refresh_token': 'refresh-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 _access(account_uuid: str):
+ roles = {
+ 'viewer-account': 'viewer',
+ 'operator-account': 'operator',
+ 'manager-account': 'developer',
+ }
+ return SimpleNamespace(
+ workspace=SimpleNamespace(uuid=WORKSPACE_UUID),
+ membership=SimpleNamespace(
+ uuid=f'membership-{account_uuid}',
+ role=roles[account_uuid],
+ projection_revision=1,
+ ),
+ execution=SimpleNamespace(instance_uuid='instance-test', placement_generation=1),
+ )
+
+
+@pytest.fixture(scope='module')
+def plugin_module():
+ """Import the plugin router without following the core HTTP cycle."""
+
+ 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,
+ ):
+ import langbot.pkg.api.http.controller.groups.plugins as plugins
+
+ yield plugins
+
+
+@pytest.fixture
+async def plugin_security_api(plugin_module):
+ viewer = SimpleNamespace(uuid='viewer-account', user='viewer@example.com')
+ operator = SimpleNamespace(uuid='operator-account', user='operator@example.com')
+ manager = SimpleNamespace(uuid='manager-account', user='manager@example.com')
+ accounts = {
+ 'viewer-token': viewer,
+ 'operator-token': operator,
+ 'manager-token': manager,
+ }
+
+ application = Mock()
+ application.user_service.get_authenticated_account = AsyncMock(side_effect=lambda token: accounts[token])
+ application.workspace_collaboration_service.resolve_account_workspace = AsyncMock(
+ side_effect=lambda account_uuid, _workspace_uuid: _access(account_uuid)
+ )
+ application.apikey_service.verify_api_key = AsyncMock(return_value=False)
+ application.instance_config.data = {
+ 'plugin': {'display_plugin_debug_url': 'http://localhost:5401'},
+ 'system': {'limitation': {}},
+ }
+
+ raw_plugin = {
+ 'author': 'example',
+ 'name': 'secure-plugin',
+ 'plugin_config': RAW_CONFIG,
+ 'debug': {'plugin_debug_key': 'list-debug-secret'},
+ }
+ application.plugin_connector.require_workspace_context = AsyncMock()
+ application.plugin_connector.list_plugins = AsyncMock(return_value=[raw_plugin])
+ application.plugin_connector.get_plugin_info = AsyncMock(return_value=raw_plugin)
+ application.plugin_connector.get_debug_info = AsyncMock(return_value={'plugin_debug_key': 'runtime-debug-secret'})
+ application.plugin_connector.get_plugin_logs = AsyncMock(return_value=['private runtime line'])
+ application.plugin_connector.set_plugin_config = AsyncMock()
+
+ persistence_result = Mock()
+ persistence_result.scalar_one_or_none.return_value = RAW_CONFIG
+ application.persistence_mgr.execute_async = AsyncMock(return_value=persistence_result)
+
+ quart_app = quart.Quart(__name__)
+ router = plugin_module.PluginsRouterGroup(application, quart_app)
+ await router.initialize()
+ return application, quart_app.test_client(), raw_plugin
+
+
+def _headers(token: str) -> dict[str, str]:
+ return {
+ 'Authorization': f'Bearer {token}',
+ 'X-Workspace-Id': WORKSPACE_UUID,
+ }
+
+
+def test_recursive_redaction_preserves_structure_without_mutating_input(plugin_module):
+ redacted = plugin_module.redact_plugin_secrets(RAW_CONFIG)
+
+ assert redacted['apiKey'] == '***'
+ assert redacted['nested']['headers'] == {
+ 'Authorization': '***',
+ 'Accept': 'application/json',
+ }
+ assert redacted['nested']['refresh_token'] == '***'
+ 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 RAW_CONFIG['apiKey'] == 'api-secret'
+ assert RAW_CONFIG['nested']['headers']['Authorization'] == 'Bearer nested-secret'
+ with pytest.raises(ValueError, match='no existing value'):
+ plugin_module.restore_plugin_secret_placeholders({'api_key': '***'}, {})
+
+
+@pytest.mark.asyncio
+async def test_viewer_plugin_reads_are_recursively_redacted(plugin_security_api):
+ application, client, raw_plugin = plugin_security_api
+
+ list_response = await client.get('/api/v1/plugins', headers=_headers('viewer-token'))
+ detail_response = await client.get(
+ '/api/v1/plugins/example/secure-plugin',
+ headers=_headers('viewer-token'),
+ )
+ config_response = await client.get(
+ '/api/v1/plugins/example/secure-plugin/config',
+ headers=_headers('viewer-token'),
+ )
+
+ assert list_response.status_code == 200
+ assert detail_response.status_code == 200
+ assert config_response.status_code == 200
+ listed_plugin = (await list_response.get_json())['data']['plugins'][0]
+ detailed_plugin = (await detail_response.get_json())['data']['plugin']
+ config = (await config_response.get_json())['data']['config']
+ for plugin in (listed_plugin, detailed_plugin):
+ assert plugin['plugin_config']['apiKey'] == '***'
+ assert plugin['debug']['plugin_debug_key'] == '***'
+ assert config['apiKey'] == '***'
+ assert config['nested']['headers']['Authorization'] == '***'
+ assert raw_plugin['plugin_config']['apiKey'] == 'api-secret'
+ application.plugin_connector.require_workspace_context.assert_awaited()
+
+
+@pytest.mark.asyncio
+async def test_manager_read_is_redacted_but_write_preserves_or_replaces_secrets(
+ plugin_security_api,
+ plugin_module,
+):
+ application, client, _ = plugin_security_api
+
+ read_response = await client.get(
+ '/api/v1/plugins/example/secure-plugin/config',
+ headers=_headers('manager-token'),
+ )
+ masked_update = plugin_module.redact_plugin_secrets(RAW_CONFIG)
+ masked_update['enabled'] = False
+ preserved_write = await client.put(
+ '/api/v1/plugins/example/secure-plugin/config',
+ headers=_headers('manager-token'),
+ json=masked_update,
+ )
+ replacement = copy.deepcopy(RAW_CONFIG)
+ replacement['apiKey'] = 'replacement-secret'
+ replaced_write = await client.put(
+ '/api/v1/plugins/example/secure-plugin/config',
+ headers=_headers('manager-token'),
+ json=replacement,
+ )
+
+ preserved = copy.deepcopy(RAW_CONFIG)
+ preserved['enabled'] = False
+
+ assert read_response.status_code == 200
+ assert (await read_response.get_json())['data']['config']['apiKey'] == '***'
+ assert preserved_write.status_code == 200
+ assert replaced_write.status_code == 200
+ assert application.plugin_connector.set_plugin_config.await_args_list == [
+ call('example', 'secure-plugin', preserved),
+ call('example', 'secure-plugin', replacement),
+ ]
+
+
+@pytest.mark.asyncio
+async def test_debug_key_requires_resource_manage_permission(plugin_security_api):
+ application, client, _ = plugin_security_api
+
+ viewer_denied = await client.get('/api/v1/plugins/debug-info', headers=_headers('viewer-token'))
+ operator_denied = await client.get('/api/v1/plugins/debug-info', headers=_headers('operator-token'))
+ application.plugin_connector.get_debug_info.assert_not_awaited()
+ allowed = await client.get('/api/v1/plugins/debug-info', headers=_headers('manager-token'))
+
+ assert viewer_denied.status_code == 403
+ assert operator_denied.status_code == 403
+ assert allowed.status_code == 200
+ assert (await allowed.get_json())['data'] == {
+ 'debug_url': 'http://localhost:5401',
+ 'plugin_debug_key': 'runtime-debug-secret',
+ }
+ application.plugin_connector.get_debug_info.assert_awaited_once_with()
+
+
+@pytest.mark.asyncio
+async def test_viewer_cannot_read_plugin_runtime_logs(plugin_security_api):
+ application, client, _ = plugin_security_api
+
+ response = await client.get(
+ '/api/v1/plugins/example/secure-plugin/logs',
+ headers=_headers('viewer-token'),
+ )
+
+ assert response.status_code == 403
+ assert (await response.get_json())['code'] == 'permission_denied'
+ application.plugin_connector.get_plugin_logs.assert_not_awaited()
diff --git a/tests/integration/api/test_providers.py b/tests/integration/api/test_providers.py
index a42a99428..4aa2e1342 100644
--- a/tests/integration/api/test_providers.py
+++ b/tests/integration/api/test_providers.py
@@ -10,6 +10,7 @@ from __future__ import annotations
import pytest
from unittest.mock import MagicMock, AsyncMock, Mock
+from types import SimpleNamespace
from tests.factories import FakeApp
@@ -66,10 +67,25 @@ def fake_provider_app():
)
# Auth services
+ account = SimpleNamespace(uuid='account-test', user='test@example.com')
app.user_service = Mock()
app.user_service.is_initialized = AsyncMock(return_value=True)
app.user_service.verify_jwt_token = AsyncMock(return_value='test@example.com')
- app.user_service.get_user_by_email = AsyncMock(return_value=Mock(email='test@example.com'))
+ app.user_service.get_user_by_email = AsyncMock(return_value=account)
+ app.user_service.get_authenticated_account = AsyncMock(return_value=account)
+ app.workspace_collaboration_service = SimpleNamespace(
+ resolve_account_workspace=AsyncMock(
+ return_value=SimpleNamespace(
+ workspace=SimpleNamespace(uuid='workspace-test'),
+ membership=SimpleNamespace(
+ uuid='membership-test',
+ role='owner',
+ projection_revision=0,
+ ),
+ execution=SimpleNamespace(instance_uuid='instance-test', placement_generation=1),
+ )
+ )
+ )
app.apikey_service = Mock()
app.apikey_service.verify_api_key = AsyncMock(return_value=True)
diff --git a/tests/integration/api/test_smoke.py b/tests/integration/api/test_smoke.py
index 9f611bb6c..58a7c2a01 100644
--- a/tests/integration/api/test_smoke.py
+++ b/tests/integration/api/test_smoke.py
@@ -288,6 +288,24 @@ class TestUserInitEndpoint:
assert data['msg'] == 'ok'
assert data['data']['initialized'] is False
+ @pytest.mark.asyncio
+ async def test_account_info_exposes_instance_capabilities_not_first_account(self, quart_test_client, fake_api_app):
+ fake_api_app.user_service.is_initialized.return_value = True
+ fake_api_app.user_service.get_first_user = AsyncMock(
+ side_effect=AssertionError('public login bootstrap must not inspect an account')
+ )
+
+ response = await quart_test_client.get('/api/v1/user/account-info')
+
+ assert response.status_code == 200
+ data = await response.get_json()
+ assert data['data'] == {
+ 'initialized': True,
+ 'password_login_enabled': True,
+ 'space_login_enabled': True,
+ }
+ fake_api_app.user_service.get_first_user.assert_not_awaited()
+
@pytest.mark.usefixtures('mock_circular_import_chain')
class TestRealImports:
diff --git a/tests/integration/api/test_user_space_oauth.py b/tests/integration/api/test_user_space_oauth.py
new file mode 100644
index 000000000..7138cf8bc
--- /dev/null
+++ b/tests/integration/api/test_user_space_oauth.py
@@ -0,0 +1,234 @@
+"""Security tests for the LangBot-to-Space OAuth redirect boundary."""
+
+from __future__ import annotations
+
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, Mock
+from urllib.parse import parse_qs, urlsplit
+
+import pytest
+import quart
+
+from langbot.pkg.api.http.controller.groups.user import UserRouterGroup
+
+
+pytestmark = pytest.mark.integration
+WORKSPACE_UUID = '11111111-1111-4111-8111-111111111111'
+
+
+@pytest.fixture
+async def space_oauth_api():
+ account = SimpleNamespace(uuid='account-a', user='owner@example.com')
+ access = SimpleNamespace(
+ workspace=SimpleNamespace(uuid=WORKSPACE_UUID),
+ membership=SimpleNamespace(uuid='member-a', role='owner', projection_revision=1),
+ execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=1),
+ )
+ application = Mock()
+ application.user_service.get_authenticated_account = AsyncMock(return_value=account)
+ application.user_service.issue_space_oauth_state = AsyncMock(
+ side_effect=lambda purpose, **_: f'opaque-{purpose}-state'
+ )
+ local_account = SimpleNamespace(
+ uuid='account-a',
+ user='owner@example.com',
+ account_type='local',
+ )
+ bound_account = SimpleNamespace(
+ uuid='account-a',
+ user='owner@example.com',
+ account_type='space',
+ )
+ application.user_service.consume_space_oauth_state = AsyncMock(
+ side_effect=lambda state, purpose: (
+ local_account if (state, purpose) == ('opaque-bind-state', 'bind') else None
+ )
+ )
+ application.user_service.bind_space_account = AsyncMock(return_value=bound_account)
+ application.user_service.generate_jwt_token = AsyncMock(return_value='rotated-account-token')
+ application.user_service.authenticate_space_user = AsyncMock(return_value=('space-login-token', bound_account))
+ application.user_service.verify_jwt_token = AsyncMock()
+ application.workspace_collaboration_service.resolve_account_workspace = AsyncMock(return_value=access)
+ application.space_service.get_oauth_authorize_url = Mock(
+ side_effect=lambda redirect_uri, state: f'https://space.example/authorize?state={state}'
+ )
+ application.space_service.exchange_oauth_code = AsyncMock(
+ return_value={
+ 'access_token': 'space-access-token',
+ 'refresh_token': 'space-refresh-token',
+ 'expires_in': 3600,
+ }
+ )
+ application.instance_config.data = {
+ 'api': {'webui_url': 'http://localhost'},
+ 'system': {'allow_modify_login_info': True},
+ }
+
+ quart_app = quart.Quart(__name__)
+ router = UserRouterGroup(application, quart_app)
+ await router.initialize()
+ return application, quart_app.test_client()
+
+
+@pytest.mark.asyncio
+async def test_public_login_state_is_server_issued(space_oauth_api):
+ application, client = space_oauth_api
+ response = await client.get(
+ '/api/v1/user/space/authorize-url',
+ query_string={'redirect_uri': 'http://localhost/auth/space/callback'},
+ headers={'Origin': 'http://localhost'},
+ )
+
+ assert response.status_code == 200
+ authorize_url = (await response.get_json())['data']['authorize_url']
+ assert parse_qs(urlsplit(authorize_url).query)['state'] == ['opaque-login-state']
+ application.user_service.issue_space_oauth_state.assert_awaited_once_with('login')
+
+
+@pytest.mark.asyncio
+async def test_public_login_rejects_caller_supplied_state(space_oauth_api):
+ application, client = space_oauth_api
+ response = await client.get(
+ '/api/v1/user/space/authorize-url',
+ query_string={
+ 'redirect_uri': 'http://localhost/auth/space/callback',
+ 'state': 'jwt.must-not-be-used',
+ },
+ headers={'Origin': 'http://localhost'},
+ )
+
+ assert response.status_code == 200
+ assert (await response.get_json())['code'] == 1
+ application.space_service.get_oauth_authorize_url.assert_not_called()
+
+
+@pytest.mark.asyncio
+async def test_bind_state_is_account_bound_and_requires_authentication(space_oauth_api):
+ application, client = space_oauth_api
+ path = '/api/v1/user/space/bind-authorize-url'
+ query = {'redirect_uri': 'http://localhost/auth/space/callback?mode=bind'}
+
+ unauthorized = await client.get(path, query_string=query, headers={'Origin': 'http://localhost'})
+ response = await client.get(
+ path,
+ query_string=query,
+ headers={
+ 'Origin': 'http://localhost',
+ 'Authorization': 'Bearer user-token',
+ 'X-Workspace-Id': WORKSPACE_UUID,
+ },
+ )
+
+ assert unauthorized.status_code == 401
+ assert response.status_code == 200
+ application.user_service.issue_space_oauth_state.assert_awaited_once_with('bind', account_uuid='account-a')
+
+
+@pytest.mark.asyncio
+async def test_redirect_origin_and_callback_path_are_restricted(space_oauth_api):
+ _, client = space_oauth_api
+
+ wrong_origin = await client.get(
+ '/api/v1/user/space/authorize-url',
+ query_string={'redirect_uri': 'https://evil.example/auth/space/callback'},
+ headers={'Origin': 'http://localhost'},
+ )
+ wrong_path = await client.get(
+ '/api/v1/user/space/authorize-url',
+ query_string={'redirect_uri': 'http://localhost/arbitrary'},
+ headers={'Origin': 'http://localhost'},
+ )
+ forged_origin = await client.get(
+ '/api/v1/user/space/authorize-url',
+ query_string={'redirect_uri': 'https://evil.example/auth/space/callback'},
+ headers={'Origin': 'https://evil.example'},
+ )
+ forged_host = await client.get(
+ '/api/v1/user/space/authorize-url',
+ query_string={'redirect_uri': 'https://evil.example/auth/space/callback'},
+ headers={'Host': 'evil.example'},
+ )
+
+ assert (await wrong_origin.get_json())['code'] == 1
+ assert (await wrong_path.get_json())['code'] == 1
+ assert (await forged_origin.get_json())['code'] == 1
+ assert (await forged_host.get_json())['code'] == 1
+
+
+@pytest.mark.asyncio
+async def test_explicit_server_side_webui_origin_supports_split_dev_server(space_oauth_api):
+ application, client = space_oauth_api
+ application.instance_config.data['api'] = {'webui_url': 'http://localhost:5173'}
+
+ response = await client.get(
+ '/api/v1/user/space/authorize-url',
+ query_string={'redirect_uri': 'http://localhost:5173/auth/space/callback'},
+ headers={'Origin': 'https://irrelevant.example'},
+ )
+
+ assert response.status_code == 200
+ assert (await response.get_json())['code'] == 0
+
+
+@pytest.mark.asyncio
+async def test_server_side_webhook_origin_supports_bundled_ui(space_oauth_api):
+ application, client = space_oauth_api
+ application.instance_config.data['api'] = {
+ 'webui_url': '',
+ 'webhook_prefix': 'https://langbot.example/base/path',
+ }
+
+ response = await client.get(
+ '/api/v1/user/space/authorize-url',
+ query_string={'redirect_uri': 'https://langbot.example/auth/space/callback'},
+ headers={'Host': 'attacker.example'},
+ )
+
+ assert response.status_code == 200
+ assert (await response.get_json())['code'] == 0
+
+
+@pytest.mark.asyncio
+async def test_login_callback_requires_and_consumes_server_state(space_oauth_api):
+ application, client = space_oauth_api
+
+ missing = await client.post('/api/v1/user/space/callback', json={'code': 'oauth-code'})
+ response = await client.post(
+ '/api/v1/user/space/callback',
+ json={'code': 'oauth-code', 'state': 'opaque-login-state'},
+ )
+
+ assert (await missing.get_json())['code'] == 1
+ assert response.status_code == 200
+ assert (await response.get_json())['data']['token'] == 'space-login-token'
+ application.user_service.consume_space_oauth_state.assert_awaited_once_with('opaque-login-state', 'login')
+ application.space_service.exchange_oauth_code.assert_awaited_once_with('oauth-code')
+
+
+@pytest.mark.asyncio
+async def test_bind_callback_uses_opaque_state_and_never_treats_it_as_jwt(space_oauth_api):
+ application, client = space_oauth_api
+ application.user_service.consume_space_oauth_state.reset_mock()
+ application.user_service.consume_space_oauth_state.side_effect = [
+ ValueError('invalid state'),
+ SimpleNamespace(
+ uuid='account-a',
+ user='owner@example.com',
+ account_type='local',
+ ),
+ ]
+
+ rejected = await client.post(
+ '/api/v1/user/bind-space',
+ json={'code': 'attacker-code', 'state': 'jwt.must-not-be-used'},
+ )
+ response = await client.post(
+ '/api/v1/user/bind-space',
+ json={'code': 'oauth-code', 'state': 'opaque-bind-state'},
+ )
+
+ assert rejected.status_code == 401
+ assert response.status_code == 200
+ assert (await response.get_json())['data']['token'] == 'rotated-account-token'
+ application.user_service.verify_jwt_token.assert_not_awaited()
+ application.user_service.bind_space_account.assert_awaited_once_with('owner@example.com', 'oauth-code')
diff --git a/tests/integration/api/test_workspaces.py b/tests/integration/api/test_workspaces.py
new file mode 100644
index 000000000..85864f1a4
--- /dev/null
+++ b/tests/integration/api/test_workspaces.py
@@ -0,0 +1,444 @@
+from __future__ import annotations
+
+import logging
+from types import SimpleNamespace
+
+import pytest
+import sqlalchemy
+import jwt
+from quart import Quart
+from sqlalchemy.ext.asyncio import create_async_engine
+
+from langbot.pkg.api.http.controller.groups.workspaces import (
+ InvitationsRouterGroup,
+ WorkspacesRouterGroup,
+)
+from langbot.pkg.api.http.controller.groups.apikeys import ApiKeysRouterGroup
+from langbot.pkg.api.http.controller.groups.user import UserRouterGroup
+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.user import User
+from langbot.pkg.entity.persistence.workspace import (
+ Workspace,
+ WorkspaceExecutionState,
+ WorkspaceInvitation,
+ WorkspaceMembership,
+)
+from langbot.pkg.workspace.collaboration import WorkspaceCollaborationService
+from langbot.pkg.workspace.service import WorkspaceService
+from langbot.pkg.workspace.policy import CloudWorkspacePolicy
+
+
+pytestmark = [pytest.mark.integration, 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)
+ for column in model.__table__.columns
+ if column.name not in masked_columns
+ }
+
+
+@pytest.fixture
+async def workspace_api(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "workspace-api.db"}')
+ async with engine.begin() as connection:
+ await connection.run_sync(Base.metadata.create_all)
+
+ application = SimpleNamespace()
+ application.persistence_mgr = _PersistenceManager(engine)
+ application.instance_config = SimpleNamespace(
+ data={
+ 'system': {
+ 'jwt': {'secret': 'workspace-api-secret', 'expire': 3600},
+ 'allow_modify_login_info': True,
+ },
+ 'api': {'global_api_key': ''},
+ }
+ )
+ application.logger = logging.getLogger('workspace-api-test')
+ application.workspace_service = WorkspaceService(
+ application,
+ instance_uuid='instance-workspace-api',
+ )
+ await application.workspace_service.ensure_singleton_workspace()
+ application.workspace_collaboration_service = WorkspaceCollaborationService(
+ application,
+ application.workspace_service,
+ )
+ application.user_service = UserService(application)
+ application.apikey_service = ApiKeyService(application)
+
+ owner = await application.user_service.create_initial_account(
+ 'owner@example.com',
+ 'owner-password',
+ )
+ owner_token = await application.user_service.generate_jwt_token(owner)
+
+ quart_app = Quart(__name__)
+ await WorkspacesRouterGroup(application, quart_app).initialize()
+ await InvitationsRouterGroup(application, quart_app).initialize()
+ await ApiKeysRouterGroup(application, quart_app).initialize()
+ await UserRouterGroup(application, quart_app).initialize()
+
+ yield application, quart_app.test_client(), engine, owner_token
+ await engine.dispose()
+
+
+def _auth(token: str, workspace_uuid: str | None = None) -> dict[str, str]:
+ headers = {'Authorization': f'Bearer {token}'}
+ if workspace_uuid is not None:
+ headers['X-Workspace-Id'] = workspace_uuid
+ return headers
+
+
+async def test_owner_invites_second_account_and_secret_is_not_persisted(workspace_api):
+ application, client, engine, owner_token = workspace_api
+
+ current_response = await client.get('/api/v1/workspaces/current', headers=_auth(owner_token))
+ assert current_response.status_code == 200
+ current = (await current_response.get_json())['data']
+ workspace_uuid = current['workspace']['uuid']
+ assert current['membership']['role'] == 'owner'
+ assert 'member.invite' in current['permissions']
+
+ invite_response = await client.post(
+ f'/api/v1/workspaces/{workspace_uuid}/invitations',
+ headers=_auth(owner_token, workspace_uuid),
+ json={'email': 'member@example.com', 'role': 'viewer'},
+ )
+ assert invite_response.status_code == 200
+ invite_data = (await invite_response.get_json())['data']
+ invitation_token = invite_data['token']
+ assert invitation_token.startswith('lbi_')
+ assert 'token_hash' not in invite_data['invitation']
+
+ async with engine.connect() as connection:
+ persisted_token_hash = await connection.scalar(
+ sqlalchemy.select(WorkspaceInvitation.token_hash).where(
+ WorkspaceInvitation.uuid == invite_data['invitation']['uuid']
+ )
+ )
+ assert persisted_token_hash is not None
+ assert persisted_token_hash != invitation_token
+
+ inspect_response = await client.post(
+ '/api/v1/invitations/inspect',
+ json={'token': invitation_token},
+ )
+ assert inspect_response.status_code == 200
+ inspected = (await inspect_response.get_json())['data']
+ assert inspected['workspace']['uuid'] == workspace_uuid
+ assert inspected['invitation']['normalized_email'] == 'member@example.com'
+
+ accept_response = await client.post(
+ '/api/v1/invitations/accept',
+ json={
+ 'token': invitation_token,
+ 'registration': {
+ 'email': 'member@example.com',
+ 'password': 'member-password',
+ },
+ },
+ )
+ assert accept_response.status_code == 200
+ member_auth = (await accept_response.get_json())['data']
+ assert member_auth['workspace_uuid'] == workspace_uuid
+
+ reused_response = await client.post(
+ '/api/v1/invitations/accept',
+ json={
+ 'token': invitation_token,
+ 'registration': {
+ 'email': 'member@example.com',
+ 'password': 'member-password',
+ },
+ },
+ )
+ assert reused_response.status_code == 400
+ assert (await reused_response.get_json())['code'] == 'invitation_used'
+
+ member_current_response = await client.get(
+ '/api/v1/workspaces/current',
+ headers=_auth(member_auth['token'], workspace_uuid),
+ )
+ assert member_current_response.status_code == 200
+ member_current = (await member_current_response.get_json())['data']
+ assert member_current['membership']['role'] == 'viewer'
+ assert 'member.invite' not in member_current['permissions']
+
+ forbidden_invite = await client.post(
+ f'/api/v1/workspaces/{workspace_uuid}/invitations',
+ headers=_auth(member_auth['token'], workspace_uuid),
+ json={'email': 'third@example.com', 'role': 'viewer'},
+ )
+ assert forbidden_invite.status_code == 403
+ assert (await forbidden_invite.get_json())['code'] == 'permission_denied'
+
+
+async def test_workspace_selector_and_path_cannot_escape_membership(workspace_api):
+ _, client, _, owner_token = workspace_api
+
+ unknown_uuid = '00000000-0000-0000-0000-000000000099'
+ selector_response = await client.get(
+ '/api/v1/workspaces/current',
+ headers=_auth(owner_token, unknown_uuid),
+ )
+ assert selector_response.status_code == 404
+ assert (await selector_response.get_json())['code'] == 'resource_not_found'
+
+ path_response = await client.get(
+ f'/api/v1/workspaces/{unknown_uuid}',
+ headers=_auth(owner_token),
+ )
+ assert path_response.status_code == 404
+ assert (await path_response.get_json())['code'] == 'resource_not_found'
+
+
+async def test_oss_rejects_second_workspace(workspace_api):
+ _, client, _, owner_token = workspace_api
+
+ response = await client.post('/api/v1/workspaces', headers=_auth(owner_token), json={'name': 'Second'})
+ assert response.status_code == 403
+ assert (await response.get_json())['code'] == 'edition_limit'
+
+
+async def test_jwt_uses_account_uuid_and_disabled_account_is_rejected(workspace_api):
+ _, client, engine, owner_token = workspace_api
+ payload = jwt.decode(
+ owner_token,
+ 'workspace-api-secret',
+ algorithms=['HS256'],
+ audience='langbot-instance:instance-workspace-api',
+ issuer='langbot-core',
+ )
+ assert payload['sub']
+ assert payload['sub'] != payload['user']
+
+ async with engine.begin() as connection:
+ await connection.execute(sqlalchemy.update(User).where(User.uuid == payload['sub']).values(status='disabled'))
+
+ response = await client.get('/api/v1/workspaces/current', headers=_auth(owner_token))
+ assert response.status_code == 401
+ assert (await response.get_json())['code'] == 'invalid_authentication'
+
+
+async def test_api_key_secret_is_one_time_and_viewer_cannot_manage_keys(workspace_api):
+ application, client, _engine, owner_token = workspace_api
+ 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': 'E2E automation', 'scopes': ['resource.view']},
+ )
+ assert create_response.status_code == 200
+ created = (await create_response.get_json())['data']['key']
+ assert created['key'].startswith('lbk_')
+ assert created['secret_available'] is True
+ assert 'key_hash' not in created
+
+ list_response = await client.get('/api/v1/apikeys', headers=_auth(owner_token, workspace_uuid))
+ listed = (await list_response.get_json())['data']['keys']
+ assert len(listed) == 1
+ assert 'key' not in listed[0]
+ assert 'key_hash' not in listed[0]
+ assert listed[0]['secret_available'] is False
+ identity = await application.apikey_service.authenticate_api_key(created['key'])
+ assert identity is not None
+ assert identity.workspace_uuid == workspace_uuid
+ assert identity.permissions == frozenset({'resource.view'})
+
+ invite_response = await client.post(
+ f'/api/v1/workspaces/{workspace_uuid}/invitations',
+ headers=_auth(owner_token, workspace_uuid),
+ json={'email': 'viewer@example.com', 'role': 'viewer'},
+ )
+ invitation_token = (await invite_response.get_json())['data']['token']
+ accept_response = await client.post(
+ '/api/v1/invitations/accept',
+ json={
+ 'token': invitation_token,
+ 'registration': {'email': 'viewer@example.com', 'password': 'viewer-password'},
+ },
+ )
+ viewer_token = (await accept_response.get_json())['data']['token']
+ forbidden = await client.post(
+ '/api/v1/apikeys',
+ headers=_auth(viewer_token, workspace_uuid),
+ json={'name': 'forbidden'},
+ )
+ assert forbidden.status_code == 403
+ assert (await forbidden.get_json())['code'] == 'permission_denied'
+
+
+async def test_cloud_projection_is_selected_explicitly_and_directory_writes_use_control_plane(
+ workspace_api,
+):
+ application, client, engine, owner_token = workspace_api
+ owner_uuid = jwt.decode(
+ owner_token,
+ 'workspace-api-secret',
+ algorithms=['HS256'],
+ audience='langbot-instance:instance-workspace-api',
+ issuer='langbot-core',
+ )['sub']
+ cloud_workspace_uuid = '00000000-0000-0000-0000-000000000777'
+
+ async with engine.begin() as connection:
+ await connection.execute(
+ sqlalchemy.insert(Workspace).values(
+ uuid=cloud_workspace_uuid,
+ instance_uuid='instance-workspace-api',
+ name='Cloud Team',
+ slug='cloud-team',
+ type='team',
+ status='active',
+ source='cloud_projection',
+ projection_revision=12,
+ )
+ )
+ await connection.execute(
+ sqlalchemy.insert(WorkspaceExecutionState).values(
+ workspace_uuid=cloud_workspace_uuid,
+ instance_uuid='instance-workspace-api',
+ active_generation=12,
+ state='active',
+ write_fenced=False,
+ source='cloud',
+ desired_state_revision=12,
+ )
+ )
+ await connection.execute(
+ sqlalchemy.insert(WorkspaceMembership).values(
+ uuid='00000000-0000-0000-0000-000000000778',
+ workspace_uuid=cloud_workspace_uuid,
+ account_uuid=owner_uuid,
+ role='owner',
+ status='active',
+ projection_revision=12,
+ )
+ )
+
+ policy = CloudWorkspacePolicy()
+ application.workspace_service.policy = policy
+ application.workspace_collaboration_service.policy = policy
+
+ with pytest.raises(ControlPlaneDirectoryRequiredError):
+ await application.user_service.create_initial_account(
+ 'forbidden-cloud-local@example.com',
+ 'password',
+ )
+
+ omitted = await client.get('/api/v1/workspaces/current', headers=_auth(owner_token))
+ assert omitted.status_code == 404
+
+ refreshed_token = await client.get('/api/v1/user/check-token', headers=_auth(owner_token))
+ assert refreshed_token.status_code == 200
+ assert (await refreshed_token.get_json())['data']['token']
+
+ bootstrap_response = await client.get(
+ '/api/v1/workspaces/bootstrap',
+ headers=_auth(owner_token),
+ )
+ assert bootstrap_response.status_code == 200
+ bootstrap = (await bootstrap_response.get_json())['data']
+ singleton_uuid = (await application.workspace_service.get_singleton_workspace()).uuid
+ workspace_uuids = [item['workspace']['uuid'] for item in bootstrap['workspaces']]
+ assert set(workspace_uuids) == {singleton_uuid, cloud_workspace_uuid}
+ repeated = await client.get('/api/v1/workspaces/bootstrap', headers=_auth(owner_token))
+ assert [item['workspace']['uuid'] for item in (await repeated.get_json())['data']['workspaces']] == workspace_uuids
+ by_uuid = {item['workspace']['uuid']: item for item in bootstrap['workspaces']}
+ assert by_uuid[singleton_uuid]['membership']['account_uuid'] == owner_uuid
+ assert by_uuid[singleton_uuid]['membership']['email'] == 'owner@example.com'
+ assert by_uuid[singleton_uuid]['permissions']
+ assert by_uuid[cloud_workspace_uuid]['placement_generation'] == 12
+
+ current_response = await client.get(
+ '/api/v1/workspaces/current',
+ headers=_auth(owner_token, cloud_workspace_uuid),
+ )
+ assert current_response.status_code == 200
+ current = (await current_response.get_json())['data']
+ assert current['workspace']['uuid'] == cloud_workspace_uuid
+ assert current['workspace']['source'] == 'cloud_projection'
+ assert current['placement_generation'] == 12
+
+ create_workspace = await client.post(
+ '/api/v1/workspaces',
+ headers=_auth(owner_token, cloud_workspace_uuid),
+ json={'name': 'Not in Core'},
+ )
+ assert create_workspace.status_code == 409
+ assert (await create_workspace.get_json())['code'] == 'control_plane_required'
+
+ create_invitation = await client.post(
+ f'/api/v1/workspaces/{cloud_workspace_uuid}/invitations',
+ headers=_auth(owner_token, cloud_workspace_uuid),
+ json={'email': 'member@example.com', 'role': 'viewer'},
+ )
+ assert create_invitation.status_code == 409
+ assert (await create_invitation.get_json())['code'] == 'control_plane_required'
+
+
+async def test_account_bootstrap_does_not_disclose_non_member_workspaces(workspace_api):
+ application, client, engine, owner_token = workspace_api
+ foreign_workspace_uuid = '00000000-0000-0000-0000-000000000880'
+
+ async with engine.begin() as connection:
+ await connection.execute(
+ sqlalchemy.insert(Workspace).values(
+ uuid=foreign_workspace_uuid,
+ instance_uuid='instance-workspace-api',
+ name='Foreign Team',
+ slug='foreign-team',
+ type='team',
+ status='active',
+ source='cloud_projection',
+ projection_revision=1,
+ )
+ )
+ await connection.execute(
+ sqlalchemy.insert(WorkspaceExecutionState).values(
+ workspace_uuid=foreign_workspace_uuid,
+ instance_uuid='instance-workspace-api',
+ active_generation=1,
+ state='active',
+ write_fenced=False,
+ source='cloud',
+ desired_state_revision=1,
+ )
+ )
+
+ policy = CloudWorkspacePolicy()
+ application.workspace_service.policy = policy
+ application.workspace_collaboration_service.policy = policy
+
+ response = await client.get('/api/v1/workspaces/bootstrap', headers=_auth(owner_token))
+ assert response.status_code == 200
+ workspace_uuids = {item['workspace']['uuid'] for item in (await response.get_json())['data']['workspaces']}
+ assert foreign_workspace_uuid not in workspace_uuids
+
+ current = await client.get(
+ '/api/v1/workspaces/current',
+ headers=_auth(owner_token, foreign_workspace_uuid),
+ )
+ assert current.status_code == 404
+ assert (await current.get_json())['code'] == 'resource_not_found'
diff --git a/tests/integration/persistence/resource_migration_support.py b/tests/integration/persistence/resource_migration_support.py
new file mode 100644
index 000000000..f3eea0f30
--- /dev/null
+++ b/tests/integration/persistence/resource_migration_support.py
@@ -0,0 +1,253 @@
+from __future__ import annotations
+
+import datetime
+
+import sqlalchemy as sa
+
+
+TENANT_TABLES = (
+ 'api_keys',
+ 'bots',
+ 'bot_admins',
+ 'binary_storages',
+ 'mcp_servers',
+ 'model_providers',
+ 'llm_models',
+ 'embedding_models',
+ 'rerank_models',
+ 'legacy_pipelines',
+ 'pipeline_run_records',
+ 'plugin_settings',
+ 'knowledge_bases',
+ 'knowledge_base_files',
+ 'knowledge_base_chunks',
+ 'webhooks',
+ 'monitoring_messages',
+ 'monitoring_llm_calls',
+ 'monitoring_tool_calls',
+ 'monitoring_sessions',
+ 'monitoring_errors',
+ 'monitoring_embedding_calls',
+ 'monitoring_feedback',
+)
+
+
+def _uuid_table(metadata: sa.MetaData, name: str, *columns: sa.Column) -> sa.Table:
+ return sa.Table(name, metadata, sa.Column('uuid', sa.String(255), primary_key=True), *columns)
+
+
+async def create_legacy_resource_schema(engine, *, instance_uuid: str) -> None:
+ """Create the smallest representative pre-0010 schema with one row/table."""
+ metadata = sa.MetaData()
+ system_metadata = sa.Table(
+ 'metadata',
+ metadata,
+ sa.Column('key', sa.String(255), primary_key=True),
+ sa.Column('value', sa.String(255)),
+ )
+ users = sa.Table(
+ 'users',
+ metadata,
+ sa.Column('id', sa.Integer, primary_key=True),
+ sa.Column('user', sa.String(255), nullable=False),
+ sa.Column('password', sa.String(255), nullable=False),
+ )
+ api_keys = sa.Table(
+ 'api_keys',
+ metadata,
+ sa.Column('id', sa.Integer, primary_key=True, autoincrement=True),
+ sa.Column('name', sa.String(255), nullable=False),
+ sa.Column('key', sa.String(255), nullable=False, unique=True),
+ )
+ bots = _uuid_table(
+ metadata,
+ 'bots',
+ sa.Column('name', sa.String(255), nullable=False),
+ sa.Column('updated_at', sa.DateTime, nullable=False),
+ )
+ bot_admins = sa.Table(
+ 'bot_admins',
+ metadata,
+ sa.Column('id', sa.Integer, primary_key=True, autoincrement=True),
+ sa.Column('bot_uuid', sa.String(255), nullable=False),
+ sa.Column('launcher_type', sa.String(64), nullable=False),
+ sa.Column('launcher_id', sa.String(255), nullable=False),
+ sa.UniqueConstraint('bot_uuid', 'launcher_type', 'launcher_id', name='uq_bot_admin'),
+ )
+ binary_storages = sa.Table(
+ 'binary_storages',
+ metadata,
+ sa.Column('unique_key', sa.String(255), primary_key=True),
+ sa.Column('key', sa.String(255), nullable=False),
+ sa.Column('owner_type', sa.String(255), nullable=False),
+ sa.Column('owner', sa.String(255), nullable=False),
+ )
+ mcp_servers = _uuid_table(
+ metadata,
+ 'mcp_servers',
+ sa.Column('name', sa.String(255), nullable=False),
+ sa.Column('enable', sa.Boolean, nullable=False),
+ sa.Column('updated_at', sa.DateTime, nullable=False),
+ )
+ model_providers = _uuid_table(
+ metadata,
+ 'model_providers',
+ sa.Column('name', sa.String(255), nullable=False),
+ sa.Column('requester', sa.String(255), nullable=False),
+ )
+ llm_models = _uuid_table(
+ metadata,
+ 'llm_models',
+ sa.Column('name', sa.String(255), nullable=False),
+ sa.Column('provider_uuid', sa.String(255), nullable=False),
+ )
+ embedding_models = _uuid_table(
+ metadata,
+ 'embedding_models',
+ sa.Column('name', sa.String(255), nullable=False),
+ sa.Column('provider_uuid', sa.String(255), nullable=False),
+ )
+ rerank_models = _uuid_table(
+ metadata,
+ 'rerank_models',
+ sa.Column('name', sa.String(255), nullable=False),
+ sa.Column('provider_uuid', sa.String(255), nullable=False),
+ )
+ legacy_pipelines = _uuid_table(
+ metadata,
+ 'legacy_pipelines',
+ sa.Column('name', sa.String(255), nullable=False),
+ sa.Column('is_default', sa.Boolean, nullable=False),
+ sa.Column('updated_at', sa.DateTime, nullable=False),
+ )
+ pipeline_run_records = _uuid_table(
+ metadata,
+ 'pipeline_run_records',
+ sa.Column('pipeline_uuid', sa.String(255), nullable=False),
+ sa.Column('created_at', sa.DateTime, nullable=False),
+ )
+ plugin_settings = sa.Table(
+ 'plugin_settings',
+ metadata,
+ sa.Column('plugin_author', sa.String(255), primary_key=True),
+ sa.Column('plugin_name', sa.String(255), primary_key=True),
+ sa.Column('enabled', sa.Boolean, nullable=False),
+ )
+ knowledge_bases = _uuid_table(
+ metadata,
+ 'knowledge_bases',
+ sa.Column('name', sa.String(255), nullable=False),
+ sa.Column('collection_id', sa.String(255), nullable=True),
+ )
+ knowledge_base_files = _uuid_table(
+ metadata,
+ 'knowledge_base_files',
+ sa.Column('kb_id', sa.String(255), nullable=True),
+ )
+ knowledge_base_chunks = _uuid_table(
+ metadata,
+ 'knowledge_base_chunks',
+ sa.Column('file_id', sa.String(255), nullable=True),
+ )
+ webhooks = sa.Table(
+ 'webhooks',
+ metadata,
+ sa.Column('id', sa.Integer, primary_key=True, autoincrement=True),
+ sa.Column('name', sa.String(255), nullable=False),
+ sa.Column('enabled', sa.Boolean, nullable=False),
+ sa.Column('created_at', sa.DateTime, nullable=False),
+ )
+
+ monitoring_tables: dict[str, sa.Table] = {}
+ for table_name in (
+ 'monitoring_messages',
+ 'monitoring_llm_calls',
+ 'monitoring_tool_calls',
+ 'monitoring_errors',
+ 'monitoring_embedding_calls',
+ ):
+ monitoring_tables[table_name] = sa.Table(
+ table_name,
+ metadata,
+ sa.Column('id', sa.String(255), primary_key=True),
+ sa.Column('timestamp', sa.DateTime, nullable=False),
+ sa.Column('session_id', sa.String(255), nullable=True),
+ sa.Column('message_id', sa.String(255), nullable=True),
+ )
+ monitoring_tables['monitoring_sessions'] = sa.Table(
+ 'monitoring_sessions',
+ metadata,
+ sa.Column('session_id', sa.String(255), primary_key=True),
+ sa.Column('bot_id', sa.String(255), nullable=False),
+ sa.Column('last_activity', sa.DateTime, nullable=False),
+ sa.Column('is_active', sa.Boolean, nullable=False),
+ )
+ monitoring_tables['monitoring_feedback'] = sa.Table(
+ 'monitoring_feedback',
+ metadata,
+ sa.Column('id', sa.String(255), primary_key=True),
+ sa.Column('feedback_id', sa.String(255), nullable=False, unique=True),
+ sa.Column('timestamp', sa.DateTime, nullable=False),
+ sa.Column('session_id', sa.String(255), nullable=True),
+ sa.Column('message_id', sa.String(255), nullable=True),
+ )
+
+ now = datetime.datetime(2026, 1, 1)
+ async with engine.begin() as conn:
+ await conn.run_sync(metadata.create_all)
+ await conn.execute(
+ system_metadata.insert(),
+ [
+ {'key': 'database_version', 'value': '25'},
+ {'key': 'instance_uuid', 'value': instance_uuid},
+ {'key': 'wizard_status', 'value': 'completed'},
+ {'key': 'wizard_progress', 'value': '3'},
+ {'key': 'rag_plugin_migration_needed', 'value': 'true'},
+ ],
+ )
+ await conn.execute(users.insert().values(user='Owner@Example.COM', password='hash'))
+ await conn.execute(api_keys.insert().values(name='legacy', key='lbk_legacy-secret'))
+ await conn.execute(bots.insert().values(uuid='bot-1', name='bot', updated_at=now))
+ await conn.execute(bot_admins.insert().values(bot_uuid='bot-1', launcher_type='person', launcher_id='owner'))
+ await conn.execute(
+ binary_storages.insert().values(unique_key='plugin:demo:key', key='key', owner_type='plugin', owner='demo')
+ )
+ await conn.execute(mcp_servers.insert().values(uuid='mcp-1', name='shared-name', enable=True, updated_at=now))
+ await conn.execute(model_providers.insert().values(uuid='provider-1', name='provider', requester='openai'))
+ for table in (llm_models, embedding_models, rerank_models):
+ await conn.execute(table.insert().values(uuid=f'{table.name}-1', name='model', provider_uuid='provider-1'))
+ await conn.execute(
+ legacy_pipelines.insert().values(uuid='pipeline-1', name='pipeline', is_default=True, updated_at=now)
+ )
+ await conn.execute(
+ pipeline_run_records.insert().values(uuid='run-1', pipeline_uuid='pipeline-1', created_at=now)
+ )
+ await conn.execute(plugin_settings.insert().values(plugin_author='author', plugin_name='plugin', enabled=True))
+ await conn.execute(knowledge_bases.insert().values(uuid='kb-1', name='knowledge', collection_id='collection-1'))
+ await conn.execute(knowledge_base_files.insert().values(uuid='file-1', kb_id='kb-1'))
+ await conn.execute(knowledge_base_chunks.insert().values(uuid='chunk-1', file_id='file-1'))
+ await conn.execute(webhooks.insert().values(name='hook', enabled=True, created_at=now))
+ for table_name, table in monitoring_tables.items():
+ if table_name == 'monitoring_sessions':
+ values = {
+ 'session_id': 'session-1',
+ 'bot_id': 'bot-1',
+ 'last_activity': now,
+ 'is_active': True,
+ }
+ elif table_name == 'monitoring_feedback':
+ values = {
+ 'id': 'feedback-row-1',
+ 'feedback_id': 'feedback-1',
+ 'timestamp': now,
+ 'session_id': 'session-1',
+ 'message_id': 'message-1',
+ }
+ else:
+ values = {
+ 'id': f'{table_name}-1',
+ 'timestamp': now,
+ 'session_id': 'session-1',
+ 'message_id': 'message-1',
+ }
+ await conn.execute(table.insert().values(**values))
diff --git a/tests/integration/persistence/test_migrations_postgres.py b/tests/integration/persistence/test_migrations_postgres.py
index 7eee23785..eee2fd825 100644
--- a/tests/integration/persistence/test_migrations_postgres.py
+++ b/tests/integration/persistence/test_migrations_postgres.py
@@ -13,12 +13,17 @@ CI runs automatically with PostgreSQL service container.
from __future__ import annotations
+import logging
import os
import pytest
+import sqlalchemy as sa
+from sqlalchemy.exc import IntegrityError
from sqlalchemy.ext.asyncio import create_async_engine
from sqlalchemy import text
from langbot.pkg.entity.persistence.base import Base
+from langbot.pkg.entity.persistence.user import User
+from langbot.pkg.persistence.mgr import PersistenceManager
from langbot.pkg.persistence.alembic_runner import (
run_alembic_upgrade,
run_alembic_stamp,
@@ -27,6 +32,10 @@ from langbot.pkg.persistence.alembic_runner import (
)
from alembic.config import Config
from alembic.script import ScriptDirectory
+from langbot.pkg.utils import constants
+from langbot.pkg.workspace.collaboration import normalize_email
+
+from .resource_migration_support import TENANT_TABLES, create_legacy_resource_schema
def _get_script_head() -> str:
@@ -130,6 +139,27 @@ class TestPostgreSQLMigrationBaseline:
rev = await get_alembic_current(postgres_engine)
assert rev == '0001_baseline'
+ @pytest.mark.asyncio
+ async def test_fresh_postgres_schema_accepts_application_casefold_identity(
+ self,
+ postgres_engine,
+ clean_tables,
+ clean_alembic_version,
+ ):
+ canonical_email = normalize_email('Ꭰ@Example.COM')
+ async with postgres_engine.begin() as conn:
+ await conn.run_sync(Base.metadata.create_all)
+ await conn.execute(
+ sa.insert(User).values(
+ uuid='00000000-0000-0000-0000-000000000099',
+ user=canonical_email,
+ normalized_email=canonical_email,
+ password='hash',
+ )
+ )
+ async with postgres_engine.connect() as conn:
+ assert await conn.scalar(sa.select(User.normalized_email)) == 'Ꭰ@example.com'
+
class TestPostgreSQLMigrationUpgrade:
"""Tests for upgrade to head workflow on PostgreSQL."""
@@ -223,3 +253,175 @@ class TestPostgreSQLMigrationGetCurrent:
rev = await get_alembic_current(postgres_engine)
assert rev == '0001_baseline'
+
+
+class TestPostgreSQLWorkspaceMigration:
+ """Focused coverage for upgrading a pre-tenancy PostgreSQL instance."""
+
+ @pytest.mark.asyncio
+ async def test_postgres_legacy_instance_gets_default_workspace(
+ self,
+ postgres_engine,
+ clean_tables,
+ clean_alembic_version,
+ monkeypatch,
+ ):
+ legacy_metadata = sa.MetaData()
+ metadata_table = sa.Table(
+ 'metadata',
+ legacy_metadata,
+ sa.Column('key', sa.String(255), primary_key=True),
+ sa.Column('value', sa.String(255)),
+ )
+ users = sa.Table(
+ 'users',
+ legacy_metadata,
+ sa.Column('id', sa.Integer, primary_key=True),
+ sa.Column('user', sa.String(255), nullable=False),
+ sa.Column('password', sa.String(255), nullable=False),
+ sa.Column(
+ 'account_type',
+ sa.String(32),
+ nullable=False,
+ server_default='local',
+ ),
+ sa.Column('created_at', sa.DateTime, server_default=text('now()')),
+ sa.Column('updated_at', sa.DateTime, server_default=text('now()')),
+ )
+ async with postgres_engine.begin() as conn:
+ await conn.run_sync(legacy_metadata.create_all)
+ await conn.execute(metadata_table.insert().values(key='database_version', value='25'))
+ await conn.execute(metadata_table.insert().values(key='instance_uuid', value='instance_postgres_test'))
+ await conn.execute(users.insert().values(user='owner@example.com', password='owner-hash'))
+
+ await run_alembic_stamp(postgres_engine, '0008_mcp_resource_prefs')
+ monkeypatch.setattr(constants, 'instance_id', 'instance_postgres_test')
+ database = type('Database', (), {'get_engine': lambda self: postgres_engine})()
+ application = type('Application', (), {})()
+ application.logger = logging.getLogger('postgres-workspace-startup-test')
+ manager = PersistenceManager(application)
+ manager.db = database
+
+ await manager.create_tables()
+ async with postgres_engine.connect() as conn:
+ tables_before_migration = set(
+ await conn.run_sync(lambda sync_conn: sa.inspect(sync_conn).get_table_names())
+ )
+ assert 'workspaces' not in tables_before_migration
+
+ await manager._run_alembic_migrations()
+
+ async with postgres_engine.connect() as conn:
+ account = (await conn.execute(text('SELECT uuid, status, source FROM users'))).mappings().one()
+ workspace = (
+ (await conn.execute(text('SELECT * FROM workspaces WHERE source = :source'), {'source': 'local'}))
+ .mappings()
+ .one()
+ )
+ membership = (await conn.execute(text('SELECT * FROM workspace_memberships'))).mappings().one()
+ execution_state = (await conn.execute(text('SELECT * FROM workspace_execution_states'))).mappings().one()
+
+ assert account['status'] == 'active'
+ assert account['source'] == 'local'
+ assert workspace['instance_uuid'] == 'instance_postgres_test'
+ assert workspace['created_by_account_uuid'] == account['uuid']
+ assert membership['account_uuid'] == account['uuid']
+ assert membership['role'] == 'owner'
+ assert execution_state['active_generation'] == 1
+ assert execution_state['write_fenced'] is False
+
+
+class TestPostgreSQLResourceTenancyMigration:
+ """Legacy backfill and scoped-key enforcement on real PostgreSQL."""
+
+ @pytest.mark.asyncio
+ async def test_postgres_resources_are_backfilled_and_scoped(
+ self,
+ postgres_engine,
+ clean_tables,
+ clean_alembic_version,
+ ):
+ await create_legacy_resource_schema(
+ postgres_engine,
+ instance_uuid='postgres-resource-migration-test',
+ )
+ async with postgres_engine.begin() as conn:
+ await conn.execute(text('UPDATE users SET "user" = \'Straße@Example.COM\''))
+ await conn.execute(
+ text('INSERT INTO users ("user", password) VALUES (:email, :password)'),
+ {'email': 'Ꭰ@Example.COM', 'password': 'cherokee-hash'},
+ )
+ await run_alembic_stamp(postgres_engine, '0008_mcp_resource_prefs')
+ await run_alembic_upgrade(postgres_engine, 'head')
+
+ async with postgres_engine.connect() as conn:
+ workspace_uuid = await conn.scalar(text("SELECT uuid FROM workspaces WHERE source = 'local'"))
+ for table_name in TENANT_TABLES:
+ count, distinct_workspaces = (
+ await conn.execute(text(f'SELECT COUNT(*), COUNT(DISTINCT workspace_uuid) FROM {table_name}'))
+ ).one()
+ assert (count, distinct_workspaces) == (1, 1), table_name
+ columns = await conn.run_sync(
+ lambda sync_conn, name=table_name: {
+ column['name']: column for column in sa.inspect(sync_conn).get_columns(name)
+ }
+ )
+ assert columns['workspace_uuid']['nullable'] is False, table_name
+
+ api_columns = await conn.run_sync(
+ lambda sync_conn: {column['name'] for column in sa.inspect(sync_conn).get_columns('api_keys')}
+ )
+ assert 'key' not in api_columns
+ assert await conn.scalar(text('SELECT scopes FROM api_keys')) == ['*']
+ assert (await conn.execute(text('SELECT normalized_email FROM users ORDER BY id'))).scalars().all() == [
+ 'strasse@example.com',
+ 'Ꭰ@example.com',
+ ]
+
+ second_workspace_uuid = '00000000-0000-0000-0000-000000000002'
+ async with postgres_engine.begin() as conn:
+ await conn.execute(
+ text(
+ 'INSERT INTO workspaces '
+ '(uuid, instance_uuid, name, slug, type, status, source, projection_revision) '
+ "VALUES (:uuid, 'postgres-resource-migration-test', 'Second', 'second', "
+ "'team', 'active', 'cloud_projection', 0)"
+ ),
+ {'uuid': second_workspace_uuid},
+ )
+ await conn.execute(
+ text(
+ 'INSERT INTO mcp_servers (uuid, workspace_uuid, name, enable, updated_at) '
+ "VALUES ('mcp-2', :workspace_uuid, 'shared-name', true, now())"
+ ),
+ {'workspace_uuid': second_workspace_uuid},
+ )
+ await conn.execute(
+ text(
+ 'INSERT INTO plugin_settings '
+ '(workspace_uuid, plugin_author, plugin_name, enabled) '
+ "VALUES (:workspace_uuid, 'author', 'plugin', true)"
+ ),
+ {'workspace_uuid': second_workspace_uuid},
+ )
+
+ with pytest.raises(IntegrityError):
+ async with postgres_engine.begin() as conn:
+ await conn.execute(
+ text(
+ 'INSERT INTO mcp_servers '
+ '(uuid, workspace_uuid, name, enable, updated_at) '
+ "VALUES ('mcp-duplicate', :workspace_uuid, 'shared-name', true, now())"
+ ),
+ {'workspace_uuid': workspace_uuid},
+ )
+
+ with pytest.raises(IntegrityError):
+ async with postgres_engine.begin() as conn:
+ await conn.execute(
+ text(
+ 'INSERT INTO llm_models (uuid, workspace_uuid, name, provider_uuid) '
+ "VALUES ('cross-workspace-model', :workspace_uuid, 'model', 'provider-1')"
+ ),
+ {'workspace_uuid': second_workspace_uuid},
+ )
diff --git a/tests/integration/persistence/test_resource_tenancy_migration.py b/tests/integration/persistence/test_resource_tenancy_migration.py
new file mode 100644
index 000000000..3f7183a68
--- /dev/null
+++ b/tests/integration/persistence/test_resource_tenancy_migration.py
@@ -0,0 +1,286 @@
+from __future__ import annotations
+
+import hashlib
+import json
+import uuid
+
+import pytest
+import sqlalchemy as sa
+from sqlalchemy.exc import IntegrityError
+from sqlalchemy.ext.asyncio import create_async_engine
+
+from langbot.pkg.entity import persistence
+from langbot.pkg.entity.persistence.base import Base
+from langbot.pkg.persistence.alembic_runner import run_alembic_stamp, run_alembic_upgrade
+from langbot.pkg.utils import importutil
+
+from .resource_migration_support import TENANT_TABLES, create_legacy_resource_schema
+
+
+pytestmark = [pytest.mark.integration, pytest.mark.asyncio]
+
+
+async def _inspect(engine, callback):
+ async with engine.connect() as conn:
+ return await conn.run_sync(callback)
+
+
+async def test_legacy_sqlite_resources_are_backfilled_and_contracted(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "legacy-resources.db"}')
+ try:
+ await create_legacy_resource_schema(engine, instance_uuid='resource-migration-test')
+ await run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+ await run_alembic_upgrade(engine, 'head')
+
+ async with engine.connect() as conn:
+ workspace_uuid = await conn.scalar(sa.text("SELECT uuid FROM workspaces WHERE source = 'local'"))
+ assert workspace_uuid is not None
+ for table_name in TENANT_TABLES:
+ count, distinct_workspaces = (
+ await conn.execute(sa.text(f'SELECT COUNT(*), COUNT(DISTINCT workspace_uuid) FROM {table_name}'))
+ ).one()
+ assert count == 1, table_name
+ assert distinct_workspaces == 1, table_name
+ assert (
+ await conn.scalar(
+ sa.text(f'SELECT COUNT(*) FROM {table_name} WHERE workspace_uuid != :workspace_uuid'),
+ {'workspace_uuid': workspace_uuid},
+ )
+ == 0
+ )
+
+ api_key = (
+ (await conn.execute(sa.text('SELECT key_hash, scopes, status, created_by_account_uuid FROM api_keys')))
+ .mappings()
+ .one()
+ )
+ assert api_key['key_hash'] == hashlib.sha256(b'lbk_legacy-secret').hexdigest()
+ stored_scopes = api_key['scopes']
+ if isinstance(stored_scopes, str):
+ stored_scopes = json.loads(stored_scopes)
+ assert stored_scopes == ['*']
+ assert api_key['status'] == 'active'
+ assert api_key['created_by_account_uuid'] is not None
+ assert await conn.scalar(sa.text('SELECT normalized_email FROM users')) == 'owner@example.com'
+ legacy_kb = (
+ (
+ await conn.execute(
+ sa.text(
+ 'SELECT collection_id, legacy_vector_collection FROM knowledge_bases WHERE uuid = :uuid'
+ ),
+ {'uuid': 'kb-1'},
+ )
+ )
+ .mappings()
+ .one()
+ )
+ assert legacy_kb['collection_id'] == 'collection-1'
+ assert legacy_kb['legacy_vector_collection'] == 1
+ assert (
+ await conn.scalar(
+ sa.text(
+ 'SELECT COUNT(*) FROM metadata '
+ "WHERE key IN ('wizard_status', 'wizard_progress', 'rag_plugin_migration_needed')"
+ )
+ )
+ == 0
+ )
+ assert (
+ await conn.scalar(
+ sa.text('SELECT COUNT(*) FROM workspace_metadata WHERE workspace_uuid = :workspace_uuid'),
+ {'workspace_uuid': workspace_uuid},
+ )
+ == 3
+ )
+
+ api_columns = await _inspect(
+ engine,
+ lambda conn: {column['name'] for column in sa.inspect(conn).get_columns('api_keys')},
+ )
+ assert 'key' not in api_columns
+ assert {'uuid', 'key_hash', 'scopes', 'status', 'expires_at', 'last_used_at'} <= api_columns
+ for table_name in TENANT_TABLES:
+ columns = await _inspect(
+ engine,
+ lambda conn, name=table_name: {column['name']: column for column in sa.inspect(conn).get_columns(name)},
+ )
+ assert columns['workspace_uuid']['nullable'] is False, table_name
+ if table_name == 'knowledge_bases':
+ assert columns['legacy_vector_collection']['nullable'] is False
+
+ pk_columns = {
+ table_name: tuple(
+ (
+ await _inspect(
+ engine,
+ lambda conn, name=table_name: sa.inspect(conn).get_pk_constraint(name),
+ )
+ )['constrained_columns']
+ )
+ for table_name in ('binary_storages', 'plugin_settings', 'monitoring_sessions')
+ }
+ assert pk_columns == {
+ 'binary_storages': ('workspace_uuid', 'unique_key'),
+ 'plugin_settings': ('workspace_uuid', 'plugin_author', 'plugin_name'),
+ 'monitoring_sessions': ('workspace_uuid', 'session_id'),
+ }
+
+ pipeline_run_foreign_keys = await _inspect(
+ engine,
+ lambda conn: sa.inspect(conn).get_foreign_keys('pipeline_run_records'),
+ )
+ assert any(
+ tuple(foreign_key['constrained_columns']) == ('workspace_uuid', 'pipeline_uuid')
+ and foreign_key['referred_table'] == 'legacy_pipelines'
+ and tuple(foreign_key['referred_columns']) == ('workspace_uuid', 'uuid')
+ for foreign_key in pipeline_run_foreign_keys
+ )
+ finally:
+ await engine.dispose()
+
+
+async def test_legacy_vector_marker_backfill_resumes_from_nullable_expand_step(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "legacy-vector-retry.db"}')
+ try:
+ await create_legacy_resource_schema(engine, instance_uuid='legacy-vector-retry')
+ async with engine.begin() as conn:
+ await conn.execute(sa.text('ALTER TABLE knowledge_bases ADD COLUMN legacy_vector_collection BOOLEAN NULL'))
+ await run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+ await run_alembic_upgrade(engine, 'head')
+
+ async with engine.connect() as conn:
+ assert (
+ await conn.scalar(
+ sa.text('SELECT legacy_vector_collection FROM knowledge_bases WHERE uuid = :uuid'),
+ {'uuid': 'kb-1'},
+ )
+ == 1
+ )
+ columns = await _inspect(
+ engine,
+ lambda conn: {column['name']: column for column in sa.inspect(conn).get_columns('knowledge_bases')},
+ )
+ assert columns['legacy_vector_collection']['nullable'] is False
+ finally:
+ await engine.dispose()
+
+
+async def test_sqlite_scoped_keys_allow_cross_workspace_but_reject_same_workspace(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "scoped-keys.db"}')
+ try:
+ await create_legacy_resource_schema(engine, instance_uuid='scoped-key-test')
+ await run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+ await run_alembic_upgrade(engine, 'head')
+
+ second_workspace_uuid = str(uuid.uuid4())
+ async with engine.begin() as conn:
+ await conn.execute(sa.text('PRAGMA foreign_keys=ON'))
+ first_workspace_uuid = await conn.scalar(sa.text("SELECT uuid FROM workspaces WHERE source = 'local'"))
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO workspaces '
+ '(uuid, instance_uuid, name, slug, type, status, source, projection_revision) '
+ "VALUES (:uuid, 'scoped-key-test', 'Second', 'second', 'team', 'active', "
+ "'cloud_projection', 0)"
+ ),
+ {'uuid': second_workspace_uuid},
+ )
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO mcp_servers (uuid, workspace_uuid, name, enable, updated_at) '
+ "VALUES ('mcp-2', :workspace_uuid, 'shared-name', 1, CURRENT_TIMESTAMP)"
+ ),
+ {'workspace_uuid': second_workspace_uuid},
+ )
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO plugin_settings '
+ '(workspace_uuid, plugin_author, plugin_name, enabled) '
+ "VALUES (:workspace_uuid, 'author', 'plugin', 1)"
+ ),
+ {'workspace_uuid': second_workspace_uuid},
+ )
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO binary_storages '
+ '(workspace_uuid, unique_key, key, owner_type, owner) '
+ "VALUES (:workspace_uuid, 'plugin:demo:key', 'key', 'plugin', 'demo')"
+ ),
+ {'workspace_uuid': second_workspace_uuid},
+ )
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO monitoring_sessions '
+ '(workspace_uuid, session_id, bot_id, last_activity, is_active) '
+ "VALUES (:workspace_uuid, 'session-1', 'bot-2', CURRENT_TIMESTAMP, 1)"
+ ),
+ {'workspace_uuid': second_workspace_uuid},
+ )
+
+ with pytest.raises(IntegrityError):
+ async with engine.begin() as conn:
+ await conn.execute(sa.text('PRAGMA foreign_keys=ON'))
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO mcp_servers (uuid, workspace_uuid, name, enable, updated_at) '
+ "VALUES ('mcp-duplicate', :workspace_uuid, 'shared-name', 1, CURRENT_TIMESTAMP)"
+ ),
+ {'workspace_uuid': first_workspace_uuid},
+ )
+
+ with pytest.raises(IntegrityError):
+ async with engine.begin() as conn:
+ await conn.execute(sa.text('PRAGMA foreign_keys=ON'))
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO llm_models (uuid, workspace_uuid, name, provider_uuid) '
+ "VALUES ('cross-workspace-model', :workspace_uuid, 'model', 'provider-1')"
+ ),
+ {'workspace_uuid': second_workspace_uuid},
+ )
+
+ with pytest.raises(IntegrityError):
+ async with engine.begin() as conn:
+ await conn.execute(sa.text('PRAGMA foreign_keys=ON'))
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO pipeline_run_records '
+ '(uuid, workspace_uuid, pipeline_uuid, created_at) '
+ "VALUES ('cross-workspace-run', :workspace_uuid, 'pipeline-1', CURRENT_TIMESTAMP)"
+ ),
+ {'workspace_uuid': second_workspace_uuid},
+ )
+
+ with pytest.raises(IntegrityError):
+ async with engine.begin() as conn:
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO mcp_servers (uuid, name, enable, updated_at) '
+ "VALUES ('unscoped-mcp', 'unscoped', 1, CURRENT_TIMESTAMP)"
+ )
+ )
+ finally:
+ await engine.dispose()
+
+
+async def test_fresh_sqlite_schema_matches_resource_tenancy_contract(tmp_path):
+ importutil.import_modules_in_pkg(persistence)
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "fresh-resources.db"}')
+ try:
+ async with engine.begin() as conn:
+ await conn.run_sync(Base.metadata.create_all)
+ await run_alembic_stamp(engine, '0001_baseline')
+ await run_alembic_upgrade(engine, 'head')
+
+ tables = await _inspect(engine, lambda conn: set(sa.inspect(conn).get_table_names()))
+ assert set(TENANT_TABLES) | {'workspace_metadata'} <= tables
+ for table_name in TENANT_TABLES:
+ columns = await _inspect(
+ engine,
+ lambda conn, name=table_name: {column['name']: column for column in sa.inspect(conn).get_columns(name)},
+ )
+ assert columns['workspace_uuid']['nullable'] is False, table_name
+ if table_name == 'knowledge_bases':
+ assert columns['legacy_vector_collection']['nullable'] is False
+ finally:
+ await engine.dispose()
diff --git a/tests/integration/persistence/test_sqlite_migration_backup.py b/tests/integration/persistence/test_sqlite_migration_backup.py
new file mode 100644
index 000000000..3e9ac276b
--- /dev/null
+++ b/tests/integration/persistence/test_sqlite_migration_backup.py
@@ -0,0 +1,107 @@
+from __future__ import annotations
+
+import json
+import logging
+import pathlib
+import sqlite3
+
+import pytest
+import sqlalchemy as sa
+from sqlalchemy.ext.asyncio import create_async_engine
+
+from langbot.pkg.persistence import alembic_runner
+from langbot.pkg.persistence.mgr import PersistenceManager
+
+from .resource_migration_support import create_legacy_resource_schema
+
+
+pytestmark = [pytest.mark.integration, pytest.mark.asyncio]
+
+
+def _manager(engine) -> PersistenceManager:
+ database = type('Database', (), {'get_engine': lambda self: engine})()
+ application = type('Application', (), {})()
+ application.logger = logging.getLogger('sqlite-migration-backup-test')
+ manager = PersistenceManager(application)
+ manager.db = database
+ return manager
+
+
+def _manifest_payloads(backup_directory) -> list[dict]:
+ return [json.loads(path.read_text(encoding='utf-8')) for path in sorted(backup_directory.glob('*.json'))]
+
+
+def _assert_verified_backup(payload: dict) -> None:
+ backup_path = pathlib.Path(payload['backup_path'])
+ with sqlite3.connect(f'{backup_path.as_uri()}?mode=ro', uri=True) as connection:
+ assert connection.execute('PRAGMA quick_check').fetchall() == [('ok',)]
+ assert connection.execute('SELECT version_num FROM alembic_version').fetchone()[0] == payload['source_revision']
+
+
+async def test_tenancy_migrations_retain_verified_boundary_backups(tmp_path):
+ database_path = tmp_path / 'legacy-with-backups.db'
+ engine = create_async_engine(f'sqlite+aiosqlite:///{database_path}')
+ try:
+ await create_legacy_resource_schema(engine, instance_uuid='backup-success')
+ await alembic_runner.run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+
+ await _manager(engine)._run_alembic_migrations()
+
+ assert await alembic_runner.get_alembic_current(engine) == '0010_scope_resources'
+ payloads = _manifest_payloads(tmp_path / 'migration-backups')
+ assert len(payloads) == 2
+ assert {
+ (payload['source_revision'], payload['target_revision'], payload['status']) for payload in payloads
+ } == {
+ ('0008_mcp_resource_prefs', '0009_workspace_tenancy', 'migration_succeeded'),
+ ('0009_workspace_tenancy', '0010_scope_resources', 'migration_succeeded'),
+ }
+ for payload in payloads:
+ _assert_verified_backup(payload)
+ finally:
+ await engine.dispose()
+
+
+async def test_failed_tenancy_migration_restores_backup_and_revision(
+ tmp_path,
+ monkeypatch,
+):
+ database_path = tmp_path / 'legacy-fault-injection.db'
+ engine = create_async_engine(f'sqlite+aiosqlite:///{database_path}')
+ real_upgrade = alembic_runner.run_alembic_upgrade
+
+ async def injected_upgrade(async_engine, revision='head'):
+ if revision != '0010_scope_resources':
+ return await real_upgrade(async_engine, revision)
+ async with async_engine.begin() as connection:
+ await connection.execute(sa.text('CREATE TABLE injected_partial_migration (value TEXT NOT NULL)'))
+ await alembic_runner.run_alembic_stamp(async_engine, '0010_scope_resources')
+ raise RuntimeError('injected migration failure after a fake revision stamp')
+
+ try:
+ await create_legacy_resource_schema(engine, instance_uuid='backup-failure')
+ await alembic_runner.run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+ monkeypatch.setattr(alembic_runner, 'run_alembic_upgrade', injected_upgrade)
+
+ with pytest.raises(RuntimeError, match='injected migration failure'):
+ await _manager(engine)._run_alembic_migrations()
+
+ assert await alembic_runner.get_alembic_current(engine) == '0009_workspace_tenancy'
+ async with engine.connect() as connection:
+ tables = set(
+ await connection.run_sync(lambda sync_connection: sa.inspect(sync_connection).get_table_names())
+ )
+ assert 'injected_partial_migration' not in tables
+
+ payloads = _manifest_payloads(tmp_path / 'migration-backups')
+ restored = [payload for payload in payloads if payload['target_revision'] == '0010_scope_resources']
+ assert len(restored) == 1
+ assert restored[0]['status'] == 'restored_after_failure'
+ assert restored[0]['source_revision'] == '0009_workspace_tenancy'
+ _assert_verified_backup(restored[0])
+
+ monkeypatch.setattr(alembic_runner, 'run_alembic_upgrade', real_upgrade)
+ await _manager(engine)._run_alembic_migrations()
+ assert await alembic_runner.get_alembic_current(engine) == '0010_scope_resources'
+ finally:
+ await engine.dispose()
diff --git a/tests/integration/persistence/test_workspace_migration.py b/tests/integration/persistence/test_workspace_migration.py
new file mode 100644
index 000000000..0c1147f65
--- /dev/null
+++ b/tests/integration/persistence/test_workspace_migration.py
@@ -0,0 +1,379 @@
+from __future__ import annotations
+
+import logging
+import uuid
+
+import pytest
+import sqlalchemy as sa
+from sqlalchemy.exc import IntegrityError
+from sqlalchemy.ext.asyncio import create_async_engine
+
+from langbot.pkg.entity import persistence
+from langbot.pkg.entity.persistence.base import Base
+from langbot.pkg.entity.persistence.user import User
+from langbot.pkg.persistence.mgr import PersistenceManager
+from langbot.pkg.persistence.alembic_runner import (
+ get_alembic_current,
+ run_alembic_downgrade,
+ run_alembic_stamp,
+ run_alembic_upgrade,
+)
+from langbot.pkg.utils import constants
+from langbot.pkg.utils import importutil
+from langbot.pkg.workspace.collaboration import normalize_email
+
+
+pytestmark = [pytest.mark.integration, pytest.mark.asyncio]
+
+
+async def _create_legacy_schema(
+ engine,
+ *,
+ include_instance_uuid: bool = True,
+ include_users: bool = True,
+) -> None:
+ legacy_metadata = sa.MetaData()
+ metadata_table = sa.Table(
+ 'metadata',
+ legacy_metadata,
+ sa.Column('key', sa.String(255), primary_key=True),
+ sa.Column('value', sa.String(255)),
+ )
+ users = sa.Table(
+ 'users',
+ legacy_metadata,
+ sa.Column('id', sa.Integer, primary_key=True),
+ sa.Column('user', sa.String(255), nullable=False),
+ sa.Column('password', sa.String(255), nullable=False),
+ sa.Column('account_type', sa.String(32), nullable=False, server_default='local'),
+ sa.Column('space_account_uuid', sa.String(255), nullable=True),
+ sa.Column('space_access_token', sa.Text, nullable=True),
+ sa.Column('space_refresh_token', sa.Text, nullable=True),
+ sa.Column('space_access_token_expires_at', sa.DateTime, nullable=True),
+ sa.Column('space_api_key', sa.String(255), nullable=True),
+ sa.Column('created_at', sa.DateTime, nullable=False, server_default=sa.func.now()),
+ sa.Column('updated_at', sa.DateTime, nullable=False, server_default=sa.func.now()),
+ )
+ async with engine.begin() as conn:
+ await conn.run_sync(legacy_metadata.create_all)
+ await conn.execute(metadata_table.insert().values(key='database_version', value='25'))
+ if include_instance_uuid:
+ await conn.execute(metadata_table.insert().values(key='instance_uuid', value='instance_migration_test'))
+ if include_users:
+ await conn.execute(
+ users.insert(),
+ [
+ {'user': 'owner@example.com', 'password': 'owner-hash'},
+ {'user': 'member@example.com', 'password': 'member-hash'},
+ ],
+ )
+
+
+@pytest.fixture
+async def legacy_engine(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "legacy-workspace.db"}')
+ await _create_legacy_schema(engine)
+ await run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+ yield engine
+ await engine.dispose()
+
+
+async def test_legacy_instance_gets_stable_accounts_and_default_workspace(legacy_engine):
+ await run_alembic_upgrade(legacy_engine, 'head')
+
+ async with legacy_engine.connect() as conn:
+ tables = set(await conn.run_sync(lambda sync_conn: sa.inspect(sync_conn).get_table_names()))
+ assert {
+ 'workspaces',
+ 'workspace_memberships',
+ 'workspace_invitations',
+ 'workspace_execution_states',
+ }.issubset(tables)
+
+ accounts = (
+ (await conn.execute(sa.text('SELECT id, uuid, status, source, projection_revision FROM users ORDER BY id')))
+ .mappings()
+ .all()
+ )
+ assert len(accounts) == 2
+ assert len({account['uuid'] for account in accounts}) == 2
+ for account in accounts:
+ uuid.UUID(account['uuid'])
+ assert account['status'] == 'active'
+ assert account['source'] == 'local'
+ assert account['projection_revision'] == 0
+
+ workspace = (
+ (await conn.execute(sa.text('SELECT * FROM workspaces WHERE source = :source'), {'source': 'local'}))
+ .mappings()
+ .one()
+ )
+ assert workspace['instance_uuid'] == 'instance_migration_test'
+ assert workspace['slug'] == 'default'
+ assert workspace['status'] == 'active'
+ assert workspace['created_by_account_uuid'] == accounts[0]['uuid']
+
+ membership = (await conn.execute(sa.text('SELECT * FROM workspace_memberships'))).mappings().one()
+ assert membership['workspace_uuid'] == workspace['uuid']
+ assert membership['account_uuid'] == accounts[0]['uuid']
+ assert membership['role'] == 'owner'
+ assert membership['status'] == 'active'
+
+ execution_state = (await conn.execute(sa.text('SELECT * FROM workspace_execution_states'))).mappings().one()
+ assert execution_state['workspace_uuid'] == workspace['uuid']
+ assert execution_state['instance_uuid'] == 'instance_migration_test'
+ assert execution_state['active_generation'] == 1
+ assert execution_state['state'] == 'active'
+ assert execution_state['write_fenced'] in (False, 0)
+
+ assert await get_alembic_current(legacy_engine) == '0010_scope_resources'
+
+
+async def test_workspace_upgrade_is_idempotent_and_preserves_identifiers(legacy_engine):
+ await run_alembic_upgrade(legacy_engine, 'head')
+ async with legacy_engine.connect() as conn:
+ account_uuids_before = (await conn.execute(sa.text('SELECT uuid FROM users ORDER BY id'))).scalars().all()
+ workspace_uuid_before = (
+ await conn.execute(sa.text("SELECT uuid FROM workspaces WHERE source = 'local'"))
+ ).scalar_one()
+
+ await run_alembic_upgrade(legacy_engine, 'head')
+
+ async with legacy_engine.connect() as conn:
+ account_uuids_after = (await conn.execute(sa.text('SELECT uuid FROM users ORDER BY id'))).scalars().all()
+ workspace_uuid_after = (
+ await conn.execute(sa.text("SELECT uuid FROM workspaces WHERE source = 'local'"))
+ ).scalar_one()
+ assert account_uuids_after == account_uuids_before
+ assert workspace_uuid_after == workspace_uuid_before
+
+
+async def test_workspace_kernel_upgrade_downgrade_upgrade_round_trip(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "workspace-round-trip.db"}')
+ try:
+ await _create_legacy_schema(engine)
+ await run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+ await run_alembic_upgrade(engine, '0009_workspace_tenancy')
+ assert await get_alembic_current(engine) == '0009_workspace_tenancy'
+
+ await run_alembic_downgrade(engine, '0008_mcp_resource_prefs')
+ assert await get_alembic_current(engine) == '0008_mcp_resource_prefs'
+ async with engine.connect() as conn:
+ tables = set(await conn.run_sync(lambda sync_conn: sa.inspect(sync_conn).get_table_names()))
+ user_columns = {
+ column['name']
+ for column in await conn.run_sync(lambda sync_conn: sa.inspect(sync_conn).get_columns('users'))
+ }
+ accounts = (await conn.execute(sa.text('SELECT user, password FROM users ORDER BY id'))).all()
+ assert (
+ not {
+ 'workspaces',
+ 'workspace_memberships',
+ 'workspace_invitations',
+ 'workspace_execution_states',
+ }
+ & tables
+ )
+ assert not {'uuid', 'status', 'source', 'projection_revision'} & user_columns
+ assert accounts == [
+ ('owner@example.com', 'owner-hash'),
+ ('member@example.com', 'member-hash'),
+ ]
+
+ await run_alembic_upgrade(engine, '0009_workspace_tenancy')
+ assert await get_alembic_current(engine) == '0009_workspace_tenancy'
+ async with engine.connect() as conn:
+ assert await conn.scalar(sa.text('SELECT COUNT(*) FROM workspaces')) == 1
+ assert await conn.scalar(sa.text('SELECT COUNT(*) FROM workspace_memberships')) == 1
+ finally:
+ await engine.dispose()
+
+
+@pytest.mark.parametrize(
+ ('raw_email', 'expected_email'),
+ [
+ ('Straße@Example.COM', 'strasse@example.com'),
+ ('Ꭰ@Example.COM', 'Ꭰ@example.com'),
+ ],
+)
+async def test_workspace_upgrade_uses_runtime_unicode_email_normalization(
+ tmp_path,
+ raw_email,
+ expected_email,
+):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "unicode-email.db"}')
+ try:
+ await _create_legacy_schema(engine, include_users=False)
+ async with engine.begin() as conn:
+ await conn.execute(
+ sa.text('INSERT INTO users (user, password, account_type) VALUES (:email, :password, :type)'),
+ {'email': raw_email, 'password': 'owner-hash', 'type': 'local'},
+ )
+ await run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+ await run_alembic_upgrade(engine, 'head')
+
+ async with engine.connect() as conn:
+ assert await conn.scalar(sa.text('SELECT normalized_email FROM users')) == expected_email
+ finally:
+ await engine.dispose()
+
+
+async def test_fresh_sqlite_schema_accepts_application_casefold_identity(tmp_path):
+ importutil.import_modules_in_pkg(persistence)
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "fresh-unicode-email.db"}')
+ canonical_email = normalize_email('Ꭰ@Example.COM')
+ try:
+ async with engine.begin() as conn:
+ await conn.run_sync(Base.metadata.create_all)
+ await conn.execute(
+ sa.insert(User).values(
+ uuid='00000000-0000-0000-0000-000000000099',
+ user=canonical_email,
+ normalized_email=canonical_email,
+ password='hash',
+ )
+ )
+ async with engine.connect() as conn:
+ assert await conn.scalar(sa.select(User.normalized_email)) == 'Ꭰ@example.com'
+ finally:
+ await engine.dispose()
+
+
+async def test_workspace_upgrade_rejects_unicode_casefold_duplicate_accounts(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "unicode-email-duplicate.db"}')
+ try:
+ await _create_legacy_schema(engine, include_users=False)
+ async with engine.begin() as conn:
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO users (user, password, account_type) VALUES '
+ "('Straße@Example.COM', 'first-hash', 'local'), "
+ "('STRASSE@example.com', 'second-hash', 'local')"
+ )
+ )
+ await run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+ with pytest.raises(RuntimeError, match='both normalize'):
+ await run_alembic_upgrade(engine, 'head')
+ finally:
+ await engine.dispose()
+
+
+async def test_uninitialized_instance_gets_ownerless_default_workspace(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "uninitialized-instance.db"}')
+ try:
+ await _create_legacy_schema(engine, include_users=False)
+ await run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+ await run_alembic_upgrade(engine, 'head')
+
+ async with engine.connect() as conn:
+ workspace = (await conn.execute(sa.text('SELECT * FROM workspaces'))).mappings().one()
+ membership_count = await conn.scalar(sa.text('SELECT COUNT(*) FROM workspace_memberships'))
+ execution_state = (await conn.execute(sa.text('SELECT * FROM workspace_execution_states'))).mappings().one()
+
+ assert workspace['created_by_account_uuid'] is None
+ assert membership_count == 0
+ assert execution_state['workspace_uuid'] == workspace['uuid']
+ assert execution_state['active_generation'] == 1
+ finally:
+ await engine.dispose()
+
+
+async def test_local_workspace_unique_index_allows_cloud_projections(legacy_engine):
+ await run_alembic_upgrade(legacy_engine, 'head')
+
+ async with legacy_engine.begin() as conn:
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO workspaces '
+ '(uuid, instance_uuid, name, slug, type, status, source, projection_revision) '
+ 'VALUES (:uuid, :instance_uuid, :name, :slug, :type, :status, :source, 0)'
+ ),
+ {
+ 'uuid': str(uuid.uuid4()),
+ 'instance_uuid': 'instance_migration_test',
+ 'name': 'Cloud Projection',
+ 'slug': 'cloud-projection',
+ 'type': 'team',
+ 'status': 'active',
+ 'source': 'cloud_projection',
+ },
+ )
+
+ with pytest.raises(IntegrityError):
+ async with legacy_engine.begin() as conn:
+ await conn.execute(
+ sa.text(
+ 'INSERT INTO workspaces '
+ '(uuid, instance_uuid, name, slug, type, status, source, projection_revision) '
+ 'VALUES (:uuid, :instance_uuid, :name, :slug, :type, :status, :source, 0)'
+ ),
+ {
+ 'uuid': str(uuid.uuid4()),
+ 'instance_uuid': 'instance_migration_test',
+ 'name': 'Second Local',
+ 'slug': 'second-local',
+ 'type': 'team',
+ 'status': 'active',
+ 'source': 'local',
+ },
+ )
+
+
+async def test_legacy_instance_without_bound_instance_uuid_fails_closed(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "missing-instance.db"}')
+ try:
+ await _create_legacy_schema(engine, include_instance_uuid=False)
+ await run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+ with pytest.raises(RuntimeError, match='instance_uuid'):
+ await run_alembic_upgrade(engine, 'head')
+ finally:
+ await engine.dispose()
+
+
+async def test_persistence_startup_defers_workspace_tables_until_account_upgrade(tmp_path, monkeypatch):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "startup-order.db"}')
+ try:
+ await _create_legacy_schema(engine)
+ await run_alembic_stamp(engine, '0008_mcp_resource_prefs')
+ monkeypatch.setattr(constants, 'instance_id', 'instance_migration_test')
+
+ database = type('Database', (), {'get_engine': lambda self: engine})()
+ application = type('Application', (), {})()
+ application.logger = logging.getLogger('workspace-startup-test')
+ manager = PersistenceManager(application)
+ manager.db = database
+
+ await manager.create_tables()
+ async with engine.connect() as conn:
+ tables_before_migration = set(
+ await conn.run_sync(lambda sync_conn: sa.inspect(sync_conn).get_table_names())
+ )
+ assert 'workspaces' not in tables_before_migration
+
+ await manager._run_alembic_migrations()
+
+ async with engine.connect() as conn:
+ workspace = (
+ (await conn.execute(sa.text("SELECT * FROM workspaces WHERE source = 'local'"))).mappings().one()
+ )
+ assert workspace['instance_uuid'] == 'instance_migration_test'
+ finally:
+ await engine.dispose()
+
+
+async def test_persistence_startup_rejects_instance_uuid_drift(tmp_path, monkeypatch):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "instance-drift.db"}')
+ try:
+ await _create_legacy_schema(engine)
+ monkeypatch.setattr(constants, 'instance_id', 'different_instance')
+
+ database = type('Database', (), {'get_engine': lambda self: engine})()
+ application = type('Application', (), {})()
+ application.logger = logging.getLogger('workspace-instance-drift-test')
+ manager = PersistenceManager(application)
+ manager.db = database
+
+ with pytest.raises(RuntimeError, match='does not match'):
+ await manager.create_tables()
+ finally:
+ await engine.dispose()
diff --git a/tests/integration/pipeline/test_full_flow.py b/tests/integration/pipeline/test_full_flow.py
index 767594c33..712fdd3b6 100644
--- a/tests/integration/pipeline/test_full_flow.py
+++ b/tests/integration/pipeline/test_full_flow.py
@@ -210,7 +210,15 @@ def pipeline_app():
mock_conversation.update_time = None
mock_conversation.create_time = None
- app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
+ async def get_scoped_session(query):
+ context = query._execution_context
+ mock_session.instance_uuid = context.instance_uuid
+ mock_session.workspace_uuid = context.workspace_uuid
+ mock_session.placement_generation = context.placement_generation
+ mock_session.bot_uuid = query.bot_uuid
+ return mock_session
+
+ app.sess_mgr.get_session = AsyncMock(side_effect=get_scoped_session)
app.sess_mgr.get_conversation = AsyncMock(return_value=mock_conversation)
# Model mock for PreProcessor
diff --git a/tests/integration_tests/box/test_box_integration.py b/tests/integration_tests/box/test_box_integration.py
index c20a1d87f..cda85815a 100644
--- a/tests/integration_tests/box/test_box_integration.py
+++ b/tests/integration_tests/box/test_box_integration.py
@@ -18,6 +18,7 @@ import shutil
import socket
import subprocess
from types import SimpleNamespace
+from unittest.mock import AsyncMock
import pytest
@@ -27,7 +28,8 @@ from langbot_plugin.box.client import ActionRPCBoxClient
from langbot_plugin.box.errors import BoxBackendUnavailableError
from langbot_plugin.box.models import BoxExecutionStatus, BoxNetworkMode, BoxSpec
from langbot_plugin.box.runtime import BoxRuntime
-from langbot_plugin.box.server import BoxServerHandler
+from langbot_plugin.box.server import BoxGenerationFence, BoxServerHandler
+from langbot_plugin.entities.io.context import ActionContext
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
@@ -35,6 +37,11 @@ _logger = logging.getLogger('test.box.integration')
# Default image for integration tests — small and fast to pull.
_TEST_IMAGE = 'alpine:latest'
+_ACTION_CONTEXT = ActionContext(
+ instance_uuid='box-integration-instance',
+ workspace_uuid='box-integration-workspace',
+ placement_generation=1,
+)
# ── Skip helpers ──────────────────────────────────────────────────────
@@ -97,6 +104,22 @@ class _QueueConnection:
pass
+class _TenantBoxClient(ActionRPCBoxClient):
+ async def _call(
+ self,
+ action,
+ data,
+ timeout=15.0,
+ action_context=None,
+ ):
+ return await super()._call(
+ action,
+ data,
+ timeout=timeout,
+ action_context=action_context or _ACTION_CONTEXT,
+ )
+
+
async def _make_rpc_pair(runtime: BoxRuntime):
"""Create an in-process (ActionRPCBoxClient, server_task, client_task) connected via queues."""
from langbot_plugin.runtime.io.handler import Handler
@@ -106,14 +129,20 @@ async def _make_rpc_pair(runtime: BoxRuntime):
client_conn = _QueueConnection(rx=s2c, tx=c2s)
server_conn = _QueueConnection(rx=c2s, tx=s2c)
- server_handler = BoxServerHandler(server_conn, runtime)
+ server_handler = BoxServerHandler(
+ server_conn,
+ runtime,
+ host_control_authenticated=True,
+ trusted_instance_uuid=_ACTION_CONTEXT.instance_uuid,
+ generation_fence=BoxGenerationFence(),
+ )
server_task = asyncio.create_task(server_handler.run())
client_handler = Handler.__new__(Handler)
Handler.__init__(client_handler, client_conn)
client_task = asyncio.create_task(client_handler.run())
- client = ActionRPCBoxClient(logger=_logger)
+ client = _TenantBoxClient(logger=_logger)
client.set_handler(client_handler)
return client, server_task, client_task
@@ -294,6 +323,16 @@ async def test_full_service_to_remote_runtime(tmp_path):
mock_ap = SimpleNamespace(
logger=_logger,
+ workspace_service=SimpleNamespace(
+ instance_uuid=_ACTION_CONTEXT.instance_uuid,
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid=_ACTION_CONTEXT.instance_uuid,
+ workspace_uuid=_ACTION_CONTEXT.workspace_uuid,
+ placement_generation=_ACTION_CONTEXT.placement_generation,
+ )
+ ),
+ ),
instance_config=SimpleNamespace(
data={
'box': {
@@ -313,7 +352,12 @@ async def test_full_service_to_remote_runtime(tmp_path):
service = BoxService(mock_ap, client=client)
await service.initialize()
- query = pipeline_query.Query.model_construct(query_id=42)
+ query = pipeline_query.Query.model_construct(
+ query_id=42,
+ instance_uuid=_ACTION_CONTEXT.instance_uuid,
+ workspace_uuid=_ACTION_CONTEXT.workspace_uuid,
+ placement_generation=_ACTION_CONTEXT.placement_generation,
+ )
result = await service.execute_tool(
{'command': 'echo service-path'},
query,
diff --git a/tests/integration_tests/box/test_box_mcp_integration.py b/tests/integration_tests/box/test_box_mcp_integration.py
index 2fcfcb934..325910a1b 100644
--- a/tests/integration_tests/box/test_box_mcp_integration.py
+++ b/tests/integration_tests/box/test_box_mcp_integration.py
@@ -23,14 +23,41 @@ import pytest
from aiohttp.test_utils import TestServer
from langbot_plugin.box.client import ActionRPCBoxClient
-from langbot_plugin.box.errors import BoxManagedProcessNotFoundError, BoxSessionNotFoundError
+from langbot_plugin.box.errors import (
+ BoxError,
+ BoxManagedProcessNotFoundError,
+ BoxSessionNotFoundError,
+)
from langbot_plugin.box.models import BoxManagedProcessSpec, BoxManagedProcessStatus, BoxSpec
from langbot_plugin.box.runtime import BoxRuntime
-from langbot_plugin.box.server import BoxServerHandler, create_ws_relay_app
+from langbot_plugin.box.security import (
+ BOX_CONTROL_TOKEN_HEADER,
+ BOX_INSTANCE_HEADER,
+ BOX_PLACEMENT_GENERATION_HEADER,
+ BOX_WORKSPACE_HEADER,
+)
+from langbot_plugin.box.server import (
+ BoxGenerationFence,
+ BoxServerHandler,
+ create_ws_relay_app,
+)
+from langbot_plugin.entities.io.context import ActionContext
_logger = logging.getLogger('test.box.mcp_integration')
_TEST_IMAGE = 'alpine:latest'
+_ACTION_CONTEXT = ActionContext(
+ instance_uuid='box-integration-instance',
+ workspace_uuid='box-integration-workspace',
+ placement_generation=1,
+)
+_CONTROL_TOKEN = 'box-integration-control-token-longer-than-32-bytes'
+_RELAY_HEADERS = {
+ BOX_CONTROL_TOKEN_HEADER: _CONTROL_TOKEN,
+ BOX_INSTANCE_HEADER: _ACTION_CONTEXT.instance_uuid,
+ BOX_WORKSPACE_HEADER: _ACTION_CONTEXT.workspace_uuid,
+ BOX_PLACEMENT_GENERATION_HEADER: str(_ACTION_CONTEXT.placement_generation),
+}
# ── Skip helpers ──────────────────────────────────────────────────────
@@ -89,7 +116,26 @@ class _QueueConnection:
pass
-async def _make_rpc_pair(runtime: BoxRuntime):
+class _TenantBoxClient(ActionRPCBoxClient):
+ async def _call(
+ self,
+ action,
+ data,
+ timeout=15.0,
+ action_context=None,
+ ):
+ return await super()._call(
+ action,
+ data,
+ timeout=timeout,
+ action_context=action_context or _ACTION_CONTEXT,
+ )
+
+
+async def _make_rpc_pair(
+ runtime: BoxRuntime,
+ generation_fence: BoxGenerationFence,
+):
"""Create an in-process RPC pair connected via queues."""
from langbot_plugin.runtime.io.handler import Handler
@@ -98,14 +144,20 @@ async def _make_rpc_pair(runtime: BoxRuntime):
client_conn = _QueueConnection(rx=s2c, tx=c2s)
server_conn = _QueueConnection(rx=c2s, tx=s2c)
- server_handler = BoxServerHandler(server_conn, runtime)
+ server_handler = BoxServerHandler(
+ server_conn,
+ runtime,
+ host_control_authenticated=True,
+ trusted_instance_uuid=_ACTION_CONTEXT.instance_uuid,
+ generation_fence=generation_fence,
+ )
server_task = asyncio.create_task(server_handler.run())
client_handler = Handler.__new__(Handler)
Handler.__init__(client_handler, client_conn)
client_task = asyncio.create_task(client_handler.run())
- client = ActionRPCBoxClient(logger=_logger)
+ client = _TenantBoxClient(logger=_logger)
client.set_handler(client_handler)
return client, server_task, client_task
@@ -119,13 +171,22 @@ async def box_server():
"""Yield a (ws_relay_url, ActionRPCBoxClient) backed by a real BoxRuntime."""
runtime = BoxRuntime(logger=_logger)
await runtime.initialize()
+ generation_fence = BoxGenerationFence()
# Start ws relay for managed process attach
- ws_app = create_ws_relay_app(runtime)
+ ws_app = create_ws_relay_app(
+ runtime,
+ control_token=_CONTROL_TOKEN,
+ trusted_instance_uuid=_ACTION_CONTEXT.instance_uuid,
+ generation_fence=generation_fence,
+ )
ws_server = TestServer(ws_app)
await ws_server.start_server()
- client, server_task, client_task = await _make_rpc_pair(runtime)
+ client, server_task, client_task = await _make_rpc_pair(
+ runtime,
+ generation_fence,
+ )
ws_relay_url = str(ws_server.make_url(''))
yield ws_relay_url, client
@@ -207,10 +268,14 @@ async def test_ws_stdio_attach_echo(box_server):
await client.start_managed_process('mcp-int-ws', proc_spec)
# Connect via WebSocket (ws relay)
- ws_url = client.get_managed_process_websocket_url('mcp-int-ws', ws_relay_url)
+ ws_url = client.get_managed_process_websocket_url(
+ 'mcp-int-ws',
+ ws_relay_url,
+ action_context=_ACTION_CONTEXT,
+ )
session = aiohttp.ClientSession()
try:
- async with session.ws_connect(ws_url) as ws:
+ async with session.ws_connect(ws_url, headers=_RELAY_HEADERS) as ws:
# Send a line
await ws.send_str('hello from test')
@@ -224,6 +289,45 @@ async def test_ws_stdio_attach_echo(box_server):
await client.delete_session('mcp-int-ws')
+@requires_container
+@requires_socket
+@pytest.mark.asyncio
+async def test_ws_stdio_attach_closes_on_generation_advance(box_server):
+ """A real attached relay is revoked by the next placement RPC."""
+
+ ws_relay_url, client = box_server
+ spec = BoxSpec(
+ cmd='',
+ session_id='mcp-int-generation',
+ workdir='/tmp',
+ image=_TEST_IMAGE,
+ )
+ await client.create_session(spec)
+ await client.start_managed_process(
+ 'mcp-int-generation',
+ BoxManagedProcessSpec(command='cat', args=[], cwd='/tmp'),
+ )
+ ws_url = client.get_managed_process_websocket_url(
+ 'mcp-int-generation',
+ ws_relay_url,
+ action_context=_ACTION_CONTEXT,
+ )
+ second_context = _ACTION_CONTEXT.model_copy(update={'placement_generation': 2})
+
+ async with aiohttp.ClientSession() as session:
+ async with session.ws_connect(ws_url, headers=_RELAY_HEADERS) as ws:
+ assert await client.get_sessions(action_context=second_context) == []
+ close_message = await asyncio.wait_for(ws.receive(), timeout=5)
+ assert close_message.type in {
+ aiohttp.WSMsgType.CLOSE,
+ aiohttp.WSMsgType.CLOSING,
+ aiohttp.WSMsgType.CLOSED,
+ }
+
+ with pytest.raises(BoxError, match='Stale Box placement generation'):
+ await client.get_sessions(action_context=_ACTION_CONTEXT)
+
+
# ── 3. Session cleanup removes container ─────────────────────────────
diff --git a/tests/manual/mcp_smoke.py b/tests/manual/mcp_smoke.py
index 197724fff..13125a640 100644
--- a/tests/manual/mcp_smoke.py
+++ b/tests/manual/mcp_smoke.py
@@ -18,6 +18,8 @@ from hypercorn.config import Config
from quart import Quart
from langbot.pkg.api.mcp.mount import MCPMount
+from langbot.pkg.api.http.authz import Permission
+from langbot.pkg.api.http.service.apikey import ApiKeyIdentity
PORT = 5399
GLOBAL_KEY = 'test-global-key-123'
@@ -31,11 +33,19 @@ def build_ap() -> SimpleNamespace:
ap.ver_mgr = SimpleNamespace(get_current_version=lambda: '4.5.0-test')
ap.logger = SimpleNamespace(info=print, error=print, warning=print)
- # API key verification: reuse real logic shape (global key match)
- async def verify_api_key(key: str) -> bool:
- return bool(key) and key == GLOBAL_KEY
+ # API key authentication derives the trusted Workspace carried into tools.
+ async def authenticate_api_key(key: str) -> ApiKeyIdentity | None:
+ if key != GLOBAL_KEY:
+ return None
+ return ApiKeyIdentity(
+ instance_uuid='inst-1',
+ workspace_uuid='workspace-1',
+ placement_generation=1,
+ api_key_uuid='global-test-key',
+ permissions=frozenset(permission.value for permission in Permission),
+ )
- ap.apikey_service = SimpleNamespace(verify_api_key=verify_api_key)
+ ap.apikey_service = SimpleNamespace(authenticate_api_key=authenticate_api_key)
ap.bot_service = SimpleNamespace(
get_bots=AsyncMock(return_value=[{'uuid': 'bot-1', 'name': 'Demo Bot', 'adapter': 'telegram'}])
)
diff --git a/tests/unit_tests/api/http/service/test_bot_service.py b/tests/unit_tests/api/http/service/test_bot_service.py
index 6fdc2342f..5bfacec5a 100644
--- a/tests/unit_tests/api/http/service/test_bot_service.py
+++ b/tests/unit_tests/api/http/service/test_bot_service.py
@@ -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',
diff --git a/tests/unit_tests/api/http/service/test_tenant.py b/tests/unit_tests/api/http/service/test_tenant.py
new file mode 100644
index 000000000..b0ed29e1c
--- /dev/null
+++ b/tests/unit_tests/api/http/service/test_tenant.py
@@ -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'}
diff --git a/tests/unit_tests/api/http/service/test_tenant_resource_isolation.py b/tests/unit_tests/api/http/service/test_tenant_resource_isolation.py
new file mode 100644
index 000000000..0dc7bb5e2
--- /dev/null
+++ b/tests/unit_tests/api/http/service/test_tenant_resource_isolation.py
@@ -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
+ )
diff --git a/tests/unit_tests/api/http/test_authz.py b/tests/unit_tests/api/http/test_authz.py
new file mode 100644
index 000000000..705891a16
--- /dev/null
+++ b/tests/unit_tests/api/http/test_authz.py
@@ -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
diff --git a/tests/unit_tests/api/http/test_internal_error_responses.py b/tests/unit_tests/api/http/test_internal_error_responses.py
new file mode 100644
index 000000000..3fcb10823
--- /dev/null
+++ b/tests/unit_tests/api/http/test_internal_error_responses.py
@@ -0,0 +1,138 @@
+from __future__ import annotations
+
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, Mock
+
+import pytest
+import quart
+
+from langbot.pkg.api.http.controller import group
+from langbot.pkg.api.http.controller.groups.webhooks import WebhookRouterGroup
+
+
+pytestmark = pytest.mark.asyncio
+
+
+class _FailingRouterGroup(group.RouterGroup):
+ name = 'failing-test'
+ path = '/failing-test'
+
+ async def initialize(self) -> None:
+ @self.route('', methods=['GET'], auth_type=group.AuthType.NONE)
+ async def _():
+ raise RuntimeError('database password=do-not-return')
+
+
+class _AuthenticatedRouterGroup(group.RouterGroup):
+ name = 'authenticated-test'
+ path = '/authenticated-test'
+
+ async def initialize(self) -> None:
+ @self.route('', methods=['GET'], auth_type=group.AuthType.USER_TOKEN)
+ async def _():
+ return self.success()
+
+
+class _InvalidAccountRouterGroup(group.RouterGroup):
+ name = 'invalid-account-test'
+ path = '/invalid-account-test'
+
+ async def initialize(self) -> None:
+ @self.route(
+ '',
+ methods=['GET'],
+ auth_type=group.AuthType.ACCOUNT_TOKEN,
+ permission='workspace.view',
+ )
+ async def _():
+ return self.success()
+
+
+async def test_unhandled_http_error_returns_generic_body_and_correlated_request_id():
+ logger = Mock()
+ application = SimpleNamespace(logger=logger)
+ quart_app = quart.Quart(__name__)
+ await _FailingRouterGroup(application, quart_app).initialize()
+
+ response = await quart_app.test_client().get(
+ '/failing-test',
+ headers={'X-Request-Id': 'request-http-test'},
+ )
+
+ assert response.status_code == 500
+ assert await response.get_json() == {
+ 'code': 'internal_error',
+ 'msg': 'Internal server error',
+ 'request_id': 'request-http-test',
+ }
+ assert response.headers['X-Request-Id'] == 'request-http-test'
+ log_message = logger.error.call_args.args[0]
+ assert 'request_id=request-http-test' in log_message
+ assert 'database password=do-not-return' in log_message
+ assert 'do-not-return' not in (await response.get_data(as_text=True))
+
+
+async def test_public_webhook_error_uses_same_generic_error_contract():
+ logger = Mock()
+ application = SimpleNamespace(
+ logger=logger,
+ platform_mgr=SimpleNamespace(
+ resolve_public_bot=AsyncMock(side_effect=RuntimeError('adapter credential=do-not-return'))
+ ),
+ )
+ quart_app = quart.Quart(__name__)
+ await WebhookRouterGroup(application, quart_app).initialize()
+
+ response = await quart_app.test_client().post(
+ '/bots/11111111-1111-4111-8111-111111111111',
+ headers={'X-Request-Id': 'request-webhook-test'},
+ )
+
+ assert response.status_code == 500
+ assert await response.get_json() == {
+ 'code': 'internal_error',
+ 'msg': 'Internal server error',
+ 'request_id': 'request-webhook-test',
+ }
+ assert response.headers['X-Request-Id'] == 'request-webhook-test'
+ log_message = logger.error.call_args.args[0]
+ assert 'request_id=request-webhook-test' in log_message
+ assert 'adapter credential=do-not-return' in log_message
+ assert 'do-not-return' not in (await response.get_data(as_text=True))
+
+
+async def test_authentication_failure_does_not_return_internal_exception_text():
+ logger = Mock()
+ application = SimpleNamespace(
+ logger=logger,
+ user_service=SimpleNamespace(
+ get_authenticated_account=AsyncMock(side_effect=RuntimeError('database password=do-not-return'))
+ ),
+ )
+ quart_app = quart.Quart(__name__)
+ await _AuthenticatedRouterGroup(application, quart_app).initialize()
+
+ response = await quart_app.test_client().get(
+ '/authenticated-test',
+ headers={
+ 'Authorization': 'Bearer invalid',
+ 'X-Request-Id': 'request-auth-test',
+ },
+ )
+
+ assert response.status_code == 401
+ assert await response.get_json() == {
+ 'code': 'invalid_authentication',
+ 'msg': 'Invalid authentication credentials',
+ }
+ assert 'do-not-return' not in (await response.get_data(as_text=True))
+ assert 'request_id=request-auth-test' in logger.warning.call_args.args[0]
+ assert 'database password=do-not-return' in logger.warning.call_args.args[0]
+
+
+async def test_account_token_route_cannot_declare_workspace_permission():
+ application = SimpleNamespace(logger=Mock())
+ quart_app = quart.Quart(__name__)
+
+ with pytest.raises(ValueError, match='cannot declare Workspace permissions'):
+ await _InvalidAccountRouterGroup(application, quart_app).initialize()
diff --git a/tests/unit_tests/api/service/test_apikey_service.py b/tests/unit_tests/api/service/test_apikey_service.py
index 9726888eb..0738eb20d 100644
--- a/tests/unit_tests/api/service/test_apikey_service.py
+++ b/tests/unit_tests/api/service/test_apikey_service.py
@@ -1,482 +1,443 @@
-"""
-Unit tests for ApiKeyService.
-
-Tests API key CRUD operations with mocked persistence layer.
-
-Source: src/langbot/pkg/api/http/service/apikey.py
-"""
-
from __future__ import annotations
-import pytest
-from unittest.mock import AsyncMock, Mock, patch
+import datetime
+import hashlib
+import logging
+import uuid
from types import SimpleNamespace
+from unittest.mock import AsyncMock
+import pytest
+import sqlalchemy
+from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
+
+from langbot.pkg.api.http.authz import Permission, PermissionDeniedError
+from langbot.pkg.api.http.context import PrincipalContext, PrincipalType, RequestContext, WorkspaceContext
from langbot.pkg.api.http.service.apikey import ApiKeyService
from langbot.pkg.entity.persistence.apikey import ApiKey
+from langbot.pkg.entity.persistence.base import Base
+from langbot.pkg.entity.persistence.user import User
+from langbot.pkg.entity.persistence.workspace import (
+ Workspace,
+ WorkspaceExecutionSource,
+ WorkspaceExecutionState,
+ WorkspaceSource,
+)
+from langbot.pkg.workspace.policy import SingleWorkspacePolicy
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
+from langbot.pkg.workspace.service import WorkspaceService
pytestmark = pytest.mark.asyncio
+class _PersistenceManager:
+ def __init__(self, engine):
+ self.engine = engine
+
+ def get_db_engine(self):
+ return self.engine
+
+ async def execute_async(self, *args, **kwargs):
+ async with self.engine.connect() as connection:
+ result = await connection.execute(*args, **kwargs)
+ await connection.commit()
+ return result
+
+ @staticmethod
+ def serialize_model(model, row, masked_columns=()):
+ return {
+ column.name: (
+ getattr(row, column.name).isoformat()
+ if isinstance(getattr(row, column.name), datetime.datetime)
+ else getattr(row, column.name)
+ )
+ for column in model.__table__.columns
+ if column.name not in masked_columns
+ }
+
+
+def _context(workspace_uuid: str, account_uuid: str, permissions: set[Permission]) -> RequestContext:
+ return RequestContext(
+ instance_uuid='api-key-instance',
+ placement_generation=1,
+ request_id=str(uuid.uuid4()),
+ auth_type='user-token',
+ principal=PrincipalContext(PrincipalType.ACCOUNT, account_uuid=account_uuid),
+ workspace=WorkspaceContext(
+ workspace_uuid=workspace_uuid,
+ membership_uuid=str(uuid.uuid4()),
+ role='owner',
+ permissions=frozenset(permission.value for permission in permissions),
+ ),
+ )
+
+
+@pytest.fixture
+async def api_key_context(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "api-keys.db"}')
+ async with engine.begin() as connection:
+ await connection.run_sync(Base.metadata.create_all)
+
+ application = SimpleNamespace(
+ persistence_mgr=_PersistenceManager(engine),
+ instance_config=SimpleNamespace(data={'api': {'global_api_key': ''}}),
+ logger=logging.getLogger('api-key-test'),
+ )
+ application.workspace_service = WorkspaceService(application, instance_uuid='api-key-instance')
+ workspace = await application.workspace_service.ensure_singleton_workspace()
+ account_uuid = str(uuid.uuid4())
+ session_factory = async_sessionmaker(engine, expire_on_commit=False)
+ async with session_factory.begin() as session:
+ session.add(
+ User(
+ uuid=account_uuid,
+ user='owner@example.com',
+ normalized_email='owner@example.com',
+ password='hash',
+ account_type='local',
+ )
+ )
+ service = ApiKeyService(application)
+ context = _context(workspace.uuid, account_uuid, set(Permission))
+ yield application, service, context, engine
+ await engine.dispose()
+
+
+async def test_secret_is_returned_once_and_only_hash_is_persisted(api_key_context):
+ _application, service, context, engine = api_key_context
+
+ created = await service.create_api_key(context, 'Automation', 'CI key')
+ secret = created['key']
+ assert secret.startswith('lbk_')
+ assert created['secret_available'] is True
+ assert 'key_hash' not in created
+
+ listed = await service.get_api_keys(context)
+ assert len(listed) == 1
+ assert 'key' not in listed[0]
+ assert 'key_hash' not in listed[0]
+ assert listed[0]['secret_available'] is False
+
+ async with engine.connect() as connection:
+ stored = await connection.scalar(sqlalchemy.select(ApiKey.key_hash))
+ assert stored == hashlib.sha256(secret.encode()).hexdigest()
+ assert secret not in stored
+
+
+async def test_authentication_derives_workspace_scopes_and_updates_usage(api_key_context):
+ _application, service, context, engine = api_key_context
+ created = await service.create_api_key(
+ context,
+ 'Read only',
+ scopes=[Permission.RESOURCE_VIEW.value],
+ )
+
+ identity = await service.authenticate_api_key(created['key'])
+ assert identity is not None
+ assert identity.workspace_uuid == context.workspace_uuid
+ assert identity.permissions == frozenset({Permission.RESOURCE_VIEW.value})
+
+ async with engine.connect() as connection:
+ last_used_at = await connection.scalar(sqlalchemy.select(ApiKey.last_used_at))
+ assert last_used_at is not None
+
+
+async def test_revoked_expired_and_unknown_keys_fail_closed(api_key_context):
+ _application, service, context, _engine = api_key_context
+ created = await service.create_api_key(context, 'Revocable')
+ await service.delete_api_key(context, created['id'])
+ assert await service.authenticate_api_key(created['key']) is None
+ assert await service.verify_api_key('') is False
+ assert await service.verify_api_key('plain-secret') is False
+ assert await service.verify_api_key('lbk_unknown') is False
+
+ expired_secret = 'lbk_expired'
+ await service.ap.persistence_mgr.execute_async(
+ sqlalchemy.insert(ApiKey).values(
+ workspace_uuid=context.workspace_uuid,
+ name='Expired',
+ key_hash=hashlib.sha256(expired_secret.encode()).hexdigest(),
+ scopes=[Permission.RESOURCE_VIEW.value],
+ status='active',
+ expires_at=datetime.datetime.now(datetime.UTC).replace(tzinfo=None) - datetime.timedelta(seconds=1),
+ )
+ )
+ assert await service.authenticate_api_key(expired_secret) is None
+
+
+async def test_cross_workspace_crud_and_secret_guessing_are_isolated(api_key_context):
+ application, service, first_context, engine = api_key_context
+ second_workspace_uuid = str(uuid.uuid4())
+ async with async_sessionmaker(engine, expire_on_commit=False).begin() as session:
+ session.add(
+ Workspace(
+ uuid=second_workspace_uuid,
+ instance_uuid='api-key-instance',
+ name='Second',
+ slug='second',
+ source=WorkspaceSource.CLOUD_PROJECTION.value,
+ )
+ )
+ session.add(
+ WorkspaceExecutionState(
+ workspace_uuid=second_workspace_uuid,
+ instance_uuid='api-key-instance',
+ active_generation=3,
+ state='active',
+ write_fenced=False,
+ source=WorkspaceExecutionSource.CLOUD.value,
+ )
+ )
+ second_context = _context(second_workspace_uuid, first_context.account_uuid or '', set(Permission))
+ created = await service.create_api_key(first_context, 'First only')
+
+ assert await service.get_api_key(second_context, created['id']) is None
+ assert await service.get_api_keys(second_context) == []
+ identity = await service.authenticate_api_key(created['key'])
+ assert identity is not None
+ assert identity.workspace_uuid == first_context.workspace_uuid
+ assert identity.workspace_uuid != second_workspace_uuid
+
+ # Prove the explicit multi-Workspace policy does not change key-derived routing.
+ application.workspace_service.policy = SingleWorkspacePolicy(workspace_limit=10, multi_workspace_enabled=True)
+ identity = await service.authenticate_api_key(created['key'])
+ assert identity is not None
+ assert identity.workspace_uuid == first_context.workspace_uuid
+
+
+async def test_global_config_key_is_oss_singleton_only(api_key_context):
+ application, service, _context_value, _engine = api_key_context
+ application.instance_config.data['api']['global_api_key'] = 'configured-secret'
+
+ identity = await service.authenticate_api_key('configured-secret')
+ assert identity is not None
+ assert identity.api_key_uuid == 'global-oss-api-key'
+
+ application.workspace_service.policy = SingleWorkspacePolicy(workspace_limit=10, multi_workspace_enabled=True)
+ assert await service.authenticate_api_key('configured-secret') is None
+
+
+async def test_explicit_scopes_cannot_exceed_callers_workspace_permissions(api_key_context):
+ _application, service, context, _engine = api_key_context
+ limited_context = _context(
+ context.workspace_uuid,
+ context.account_uuid or '',
+ {Permission.API_KEY_MANAGE, Permission.RESOURCE_VIEW},
+ )
+
+ created = await service.create_api_key(
+ limited_context,
+ 'Read only',
+ scopes=[Permission.RESOURCE_VIEW.value],
+ )
+ identity = await service.authenticate_api_key(created['key'])
+ assert identity is not None
+ assert identity.permissions == frozenset({Permission.RESOURCE_VIEW.value})
+
+ with pytest.raises(PermissionDeniedError) as exc_info:
+ await service.create_api_key(
+ limited_context,
+ 'Escalated',
+ scopes=[Permission.WORKSPACE_DELETE.value],
+ )
+ assert exc_info.value.permission == Permission.WORKSPACE_DELETE.value
+
+
+# Preserve the pre-tenancy CRUD and verification regression matrix while
+# exercising it through the new Workspace-bound API. The assertions reflect
+# intentional security changes: secrets are returned once, deletion revokes,
+# and missing Workspace resources are reported as not found.
class TestApiKeyServiceGetApiKeys:
- """Tests for get_api_keys method."""
+ async def test_get_api_keys_empty_list(self, api_key_context):
+ _application, service, context, _engine = api_key_context
- async def test_get_api_keys_empty_list(self):
- """Returns empty list when no API keys exist."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
- mock_result = Mock()
- mock_result.all = Mock(return_value=[])
- ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
- ap.persistence_mgr.serialize_model = Mock(
- side_effect=lambda model_cls, entity: {
- 'id': entity.id,
- 'name': entity.name,
- 'key': entity.key,
- 'description': entity.description,
- }
- if entity
- else {}
- )
+ assert await service.get_api_keys(context) == []
- service = ApiKeyService(ap)
+ async def test_get_api_keys_returns_serialized_list(self, api_key_context):
+ _application, service, context, _engine = api_key_context
+ await service.create_api_key(context, 'Test Key 1', 'First test key')
+ await service.create_api_key(context, 'Test Key 2', 'Second test key')
- # Execute
- result = await service.get_api_keys()
+ result = await service.get_api_keys(context)
- # Verify
- assert result == []
- ap.persistence_mgr.execute_async.assert_called_once()
-
- async def test_get_api_keys_returns_serialized_list(self):
- """Returns serialized list of API keys."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
-
- # Create mock API key entities
- key1 = Mock(spec=ApiKey)
- key1.id = 1
- key1.name = 'Test Key 1'
- key1.key = 'lbk_test_key_1'
- key1.description = 'First test key'
-
- key2 = Mock(spec=ApiKey)
- key2.id = 2
- key2.name = 'Test Key 2'
- key2.key = 'lbk_test_key_2'
- key2.description = 'Second test key'
-
- mock_result = Mock()
- mock_result.all = Mock(return_value=[key1, key2])
- ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
- ap.persistence_mgr.serialize_model = Mock(
- side_effect=lambda model_cls, entity: {
- 'id': entity.id,
- 'name': entity.name,
- 'key': entity.key,
- 'description': entity.description,
- }
- )
-
- service = ApiKeyService(ap)
-
- # Execute
- result = await service.get_api_keys()
-
- # Verify
- assert len(result) == 2
- assert result[0]['name'] == 'Test Key 1'
- assert result[1]['name'] == 'Test Key 2'
+ assert [item['name'] for item in result] == ['Test Key 1', 'Test Key 2']
+ assert [item['description'] for item in result] == ['First test key', 'Second test key']
+ assert all('key' not in item and 'key_hash' not in item for item in result)
class TestApiKeyServiceCreateApiKey:
- """Tests for create_api_key method."""
+ async def test_create_api_key_generates_key_with_prefix(self, api_key_context):
+ _application, service, context, _engine = api_key_context
- async def test_create_api_key_generates_key_with_prefix(self):
- """Creates API key with 'lbk_' prefix."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
+ with pytest.MonkeyPatch.context() as monkeypatch:
+ monkeypatch.setattr(
+ 'langbot.pkg.api.http.service.apikey.secrets.token_urlsafe', lambda _size: 'fixed-token'
+ )
+ result = await service.create_api_key(context, 'New Key', 'Test description')
- created_key = Mock(spec=ApiKey)
- created_key.id = 1
- created_key.name = 'New Key'
- created_key.key = 'lbk_fixed-token'
- created_key.description = 'Test description'
- select_result = Mock()
- select_result.first = Mock(return_value=created_key)
- insert_params = []
-
- async def mock_execute(query):
- params = query.compile().params
- if {'name', 'key', 'description'}.issubset(params):
- insert_params.append(params)
- return Mock()
- return select_result
-
- ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
- ap.persistence_mgr.serialize_model = Mock(
- side_effect=lambda model_cls, entity: {
- 'id': 1,
- 'name': entity.name,
- 'key': entity.key,
- 'description': entity.description,
- }
- )
-
- service = ApiKeyService(ap)
-
- with patch('langbot.pkg.api.http.service.apikey.secrets.token_urlsafe', return_value='fixed-token'):
- result = await service.create_api_key('New Key', 'Test description')
-
- assert insert_params == [{'name': 'New Key', 'key': 'lbk_fixed-token', 'description': 'Test description'}]
- assert result['key'].startswith('lbk_')
assert result['key'] == 'lbk_fixed-token'
assert result['name'] == 'New Key'
assert result['description'] == 'Test description'
+ assert result['secret_available'] is True
- async def test_create_api_key_without_description(self):
- """Creates API key with empty description when not provided."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
+ async def test_create_api_key_without_description(self, api_key_context):
+ _application, service, context, _engine = api_key_context
- created_key = Mock(spec=ApiKey)
- created_key.id = 1
- created_key.name = 'No Desc Key'
- created_key.key = 'lbk_no_desc_key'
- created_key.description = ''
+ result = await service.create_api_key(context, 'No Desc Key')
- select_result = Mock()
- select_result.first = Mock(return_value=created_key)
- insert_result = Mock()
-
- async def mock_execute(query):
- if hasattr(query, 'values'):
- return insert_result
- return select_result
-
- ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
- ap.persistence_mgr.serialize_model = Mock(
- return_value={
- 'id': 1,
- 'name': 'No Desc Key',
- 'key': 'lbk_no_desc_key',
- 'description': '',
- }
- )
-
- service = ApiKeyService(ap)
-
- # Execute
- result = await service.create_api_key('No Desc Key')
-
- # Verify
assert result['description'] == ''
class TestApiKeyServiceGetApiKey:
- """Tests for get_api_key method."""
+ async def test_get_api_key_by_id_found(self, api_key_context):
+ _application, service, context, _engine = api_key_context
+ created = await service.create_api_key(context, 'Found Key', 'Found')
- async def test_get_api_key_by_id_found(self):
- """Returns API key when found by ID."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
+ result = await service.get_api_key(context, created['id'])
- key = Mock(spec=ApiKey)
- key.id = 1
- key.name = 'Found Key'
- key.key = 'lbk_found_key'
- key.description = 'Found'
-
- mock_result = Mock()
- mock_result.first = Mock(return_value=key)
- ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
- ap.persistence_mgr.serialize_model = Mock(
- return_value={
- 'id': 1,
- 'name': 'Found Key',
- 'key': 'lbk_found_key',
- 'description': 'Found',
- }
- )
-
- service = ApiKeyService(ap)
-
- # Execute
- result = await service.get_api_key(1)
-
- # Verify
assert result is not None
- assert result['id'] == 1
+ assert result['id'] == created['id']
assert result['name'] == 'Found Key'
+ assert 'key' not in result and 'key_hash' not in result
- async def test_get_api_key_by_id_not_found(self):
- """Returns None when API key not found."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
+ async def test_get_api_key_by_id_not_found(self, api_key_context):
+ _application, service, context, _engine = api_key_context
- mock_result = Mock()
- mock_result.first = Mock(return_value=None)
- ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
+ assert await service.get_api_key(context, 999) is None
- service = ApiKeyService(ap)
+ async def test_get_api_key_by_id_zero(self, api_key_context):
+ _application, service, context, _engine = api_key_context
- # Execute
- result = await service.get_api_key(999)
-
- # Verify
- assert result is None
-
- async def test_get_api_key_by_id_zero(self):
- """Handles ID=0 (edge case) correctly."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
-
- mock_result = Mock()
- mock_result.first = Mock(return_value=None)
- ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
-
- service = ApiKeyService(ap)
-
- # Execute
- result = await service.get_api_key(0)
-
- # Verify - should return None (no key with ID 0)
- assert result is None
+ assert await service.get_api_key(context, 0) is None
class TestApiKeyServiceVerifyApiKey:
- """Tests for verify_api_key method."""
+ async def test_verify_api_key_valid(self, api_key_context):
+ _application, service, context, _engine = api_key_context
+ created = await service.create_api_key(context, 'Valid')
- @staticmethod
- def _make_ap(db_key=None, global_api_key=''):
- """Build a mock Application with persistence + instance_config."""
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
- mock_result = Mock()
- mock_result.first = Mock(return_value=db_key)
- ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
- ap.instance_config = SimpleNamespace(data={'api': {'global_api_key': global_api_key}})
- return ap
+ assert await service.verify_api_key(created['key']) is True
- async def test_verify_api_key_valid(self):
- """Returns True for valid API key."""
- # Setup
- key = Mock(spec=ApiKey)
- ap = self._make_ap(db_key=key)
+ async def test_verify_api_key_invalid(self, api_key_context):
+ _application, service, _context, _engine = api_key_context
- service = ApiKeyService(ap)
+ assert await service.verify_api_key('lbk_invalid_key') is False
- # Execute
- result = await service.verify_api_key('lbk_valid_key')
+ async def test_verify_api_key_empty_string(self, api_key_context):
+ _application, service, _context, _engine = api_key_context
- # Verify
- assert result is True
+ assert await service.verify_api_key('') is False
- async def test_verify_api_key_invalid(self):
- """Returns False for invalid API key."""
- # Setup
- ap = self._make_ap(db_key=None)
+ async def test_verify_api_key_unknown_key(self, api_key_context):
+ _application, service, _context, _engine = api_key_context
- service = ApiKeyService(ap)
+ assert await service.verify_api_key('unknown_key') is False
- # Execute
- result = await service.verify_api_key('lbk_invalid_key')
+ async def test_verify_global_api_key_match(self, api_key_context):
+ application, service, context, _engine = api_key_context
+ application.instance_config.data['api']['global_api_key'] = 'my-global-secret'
- # Verify
- assert result is False
+ identity = await service.authenticate_api_key('my-global-secret')
- async def test_verify_api_key_empty_string(self):
- """Returns False for empty key string."""
- # Setup
- ap = self._make_ap(db_key=None)
+ assert identity is not None
+ assert identity.workspace_uuid == context.workspace_uuid
+ assert identity.api_key_uuid == 'global-oss-api-key'
- service = ApiKeyService(ap)
+ async def test_verify_global_api_key_no_prefix_required(self, api_key_context):
+ application, service, _context, _engine = api_key_context
+ application.instance_config.data['api']['global_api_key'] = 'plainsecret123'
- # Execute
- result = await service.verify_api_key('')
+ assert await service.verify_api_key('plainsecret123') is True
- # Verify
- assert result is False
+ async def test_verify_global_api_key_mismatch_falls_back_to_db(self, api_key_context):
+ application, service, context, _engine = api_key_context
+ application.instance_config.data['api']['global_api_key'] = 'my-global-secret'
+ created = await service.create_api_key(context, 'DB key')
- async def test_verify_api_key_unknown_key(self):
- """Returns False when the key is not present in persistence."""
- # Setup
- ap = self._make_ap(db_key=None)
+ identity = await service.authenticate_api_key(created['key'])
- service = ApiKeyService(ap)
+ assert identity is not None
+ assert identity.api_key_uuid == created['uuid']
- # Execute
- result = await service.verify_api_key('unknown_key')
+ async def test_verify_empty_global_api_key_disabled(self, api_key_context):
+ application, service, _context, _engine = api_key_context
+ application.instance_config.data['api']['global_api_key'] = ''
- # Verify
- assert result is False
-
- async def test_verify_global_api_key_match(self):
- """Returns True when key matches the config.yaml global API key (no DB lookup)."""
- # Setup: no DB record, but a global key is configured
- ap = self._make_ap(db_key=None, global_api_key='my-global-secret')
-
- service = ApiKeyService(ap)
-
- # Execute
- result = await service.verify_api_key('my-global-secret')
-
- # Verify: accepted purely on config match
- assert result is True
- # DB should not have been consulted for the global-key path
- ap.persistence_mgr.execute_async.assert_not_called()
-
- async def test_verify_global_api_key_no_prefix_required(self):
- """Global API key is accepted even without the lbk_ prefix."""
- ap = self._make_ap(db_key=None, global_api_key='plainsecret123')
-
- service = ApiKeyService(ap)
-
- result = await service.verify_api_key('plainsecret123')
-
- assert result is True
-
- async def test_verify_global_api_key_mismatch_falls_back_to_db(self):
- """A non-matching key still falls through to the DB lookup."""
- # Global key set, but request uses a different lbk_ key that IS in DB
- key = Mock(spec=ApiKey)
- ap = self._make_ap(db_key=key, global_api_key='my-global-secret')
-
- service = ApiKeyService(ap)
-
- result = await service.verify_api_key('lbk_db_key')
-
- assert result is True
- ap.persistence_mgr.execute_async.assert_called_once()
-
- async def test_verify_empty_global_api_key_disabled(self):
- """An empty global_api_key must never authenticate an empty/blank request."""
- ap = self._make_ap(db_key=None, global_api_key='')
-
- service = ApiKeyService(ap)
-
- # Empty request key is rejected, and a blank global key never matches
assert await service.verify_api_key('') is False
assert await service.verify_api_key(' ') is False
- async def test_verify_api_key_missing_global_config_key(self):
- """Works even when api.global_api_key is absent (existing installs)."""
- # instance_config without the global_api_key field at all
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
- mock_result = Mock()
- mock_result.first = Mock(return_value=None)
- ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
- ap.instance_config = SimpleNamespace(data={'api': {}})
+ async def test_verify_api_key_missing_global_config_key(self, api_key_context):
+ application, service, _context, _engine = api_key_context
+ application.instance_config.data = {'api': {}}
- service = ApiKeyService(ap)
-
- result = await service.verify_api_key('lbk_some_key')
-
- assert result is False
+ assert await service.verify_api_key('lbk_some_key') is False
class TestApiKeyServiceDeleteApiKey:
- """Tests for delete_api_key method."""
+ async def test_delete_api_key_by_id(self, api_key_context):
+ _application, service, context, _engine = api_key_context
+ created = await service.create_api_key(context, 'Delete me')
- async def test_delete_api_key_by_id(self):
- """Deletes API key by ID."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ await service.delete_api_key(context, created['id'])
- service = ApiKeyService(ap)
+ stored = await service.get_api_key(context, created['id'])
+ assert stored is not None
+ assert stored['status'] == 'revoked'
+ assert await service.verify_api_key(created['key']) is False
- # Execute
- await service.delete_api_key(1)
+ async def test_delete_api_key_nonexistent_id(self, api_key_context):
+ _application, service, context, _engine = api_key_context
- # Verify - execute_async was called (delete operation)
- ap.persistence_mgr.execute_async.assert_called_once()
-
- async def test_delete_api_key_nonexistent_id(self):
- """Delete operation completes even for nonexistent ID (no error raised)."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
-
- service = ApiKeyService(ap)
-
- # Execute - should not raise error
- await service.delete_api_key(999)
-
- # Verify - execute_async was called regardless
- ap.persistence_mgr.execute_async.assert_called_once()
+ with pytest.raises(WorkspaceNotFoundError, match='API key not found'):
+ await service.delete_api_key(context, 999)
class TestApiKeyServiceUpdateApiKey:
- """Tests for update_api_key method."""
+ async def test_update_api_key_name_only(self, api_key_context):
+ _application, service, context, _engine = api_key_context
+ created = await service.create_api_key(context, 'Original', 'Description')
- async def test_update_api_key_name_only(self):
- """Updates only the name field."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ await service.update_api_key(context, created['id'], name='Updated Name')
- service = ApiKeyService(ap)
+ stored = await service.get_api_key(context, created['id'])
+ assert stored is not None
+ assert stored['name'] == 'Updated Name'
+ assert stored['description'] == 'Description'
- # Execute
- await service.update_api_key(1, name='Updated Name')
+ async def test_update_api_key_description_only(self, api_key_context):
+ _application, service, context, _engine = api_key_context
+ created = await service.create_api_key(context, 'Original', 'Description')
- # Verify - execute_async was called with update
- ap.persistence_mgr.execute_async.assert_called_once()
+ await service.update_api_key(context, created['id'], description='Updated description')
- async def test_update_api_key_description_only(self):
- """Updates only the description field."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ stored = await service.get_api_key(context, created['id'])
+ assert stored is not None
+ assert stored['name'] == 'Original'
+ assert stored['description'] == 'Updated description'
- service = ApiKeyService(ap)
+ async def test_update_api_key_both_fields(self, api_key_context):
+ _application, service, context, _engine = api_key_context
+ created = await service.create_api_key(context, 'Original', 'Description')
- # Execute
- await service.update_api_key(1, description='Updated description')
+ await service.update_api_key(
+ context,
+ created['id'],
+ name='New Name',
+ description='New description',
+ )
- # Verify
- ap.persistence_mgr.execute_async.assert_called_once()
+ stored = await service.get_api_key(context, created['id'])
+ assert stored is not None
+ assert stored['name'] == 'New Name'
+ assert stored['description'] == 'New description'
- async def test_update_api_key_both_fields(self):
- """Updates both name and description."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ async def test_update_api_key_no_fields(self, api_key_context):
+ application, service, context, _engine = api_key_context
+ created = await service.create_api_key(context, 'Original')
+ original_execute = application.persistence_mgr.execute_async
+ application.persistence_mgr.execute_async = AsyncMock(wraps=original_execute)
- service = ApiKeyService(ap)
+ await service.update_api_key(context, created['id'])
- # Execute
- await service.update_api_key(1, name='New Name', description='New description')
-
- # Verify
- ap.persistence_mgr.execute_async.assert_called_once()
-
- async def test_update_api_key_no_fields(self):
- """Does nothing when no fields provided."""
- # Setup
- ap = SimpleNamespace()
- ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
-
- service = ApiKeyService(ap)
-
- # Execute
- await service.update_api_key(1)
-
- # Verify - no execute call since no update_data
- ap.persistence_mgr.execute_async.assert_not_called()
+ application.persistence_mgr.execute_async.assert_not_awaited()
diff --git a/tests/unit_tests/api/service/test_bot_service.py b/tests/unit_tests/api/service/test_bot_service.py
index 8a6d0ad2a..dea5763f2 100644
--- a/tests/unit_tests/api/service/test_bot_service.py
+++ b/tests/unit_tests/api/service/test_bot_service.py
@@ -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)
diff --git a/tests/unit_tests/api/service/test_knowledge_service.py b/tests/unit_tests/api/service/test_knowledge_service.py
index 1e0592b01..21c275d8f 100644
--- a/tests/unit_tests/api/service/test_knowledge_service.py
+++ b/tests/unit_tests/api/service/test_knowledge_service.py
@@ -1,389 +1,565 @@
-"""Unit tests for API knowledge service.
-
-Tests cover:
-- Knowledge base CRUD operations
-- Capability checking
-- Knowledge engine discovery
-- File operations
-"""
+"""Tests for the tenant-aware knowledge service facade."""
from __future__ import annotations
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, Mock
+
import pytest
-from unittest.mock import Mock, AsyncMock
-from importlib import import_module
+
+from langbot.pkg.api.http.authz import WorkspaceRequiredError
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.api.http.service.knowledge import KnowledgeService
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
-def get_knowledge_service_module():
- """Lazy import to avoid circular import issues."""
- return import_module('langbot.pkg.api.http.service.knowledge')
+CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=2,
+)
-def create_mock_app():
- """Create mock Application for testing."""
- mock_app = Mock()
- mock_app.logger = Mock()
- mock_app.rag_mgr = AsyncMock()
- mock_app.persistence_mgr = AsyncMock()
- mock_app.persistence_mgr.execute_async = AsyncMock()
- mock_app.persistence_mgr.serialize_model = Mock(return_value={})
- mock_app.plugin_connector = AsyncMock()
- mock_app.plugin_connector.is_enable_plugin = True
- return mock_app
+class _Rows:
+ def __init__(self, rows=()):
+ self.rows = list(rows)
+
+ def all(self):
+ return self.rows
+
+ def __iter__(self):
+ return iter(self.rows)
+def _app():
+ return SimpleNamespace(
+ logger=Mock(),
+ rag_mgr=SimpleNamespace(
+ get_all_knowledge_base_details=AsyncMock(return_value=[]),
+ get_knowledge_base_details=AsyncMock(return_value=None),
+ create_knowledge_base=AsyncMock(),
+ remove_knowledge_base_from_runtime=AsyncMock(),
+ load_knowledge_base=AsyncMock(),
+ get_knowledge_base_by_uuid=AsyncMock(return_value=None),
+ delete_knowledge_base=AsyncMock(),
+ ),
+ persistence_mgr=SimpleNamespace(
+ execute_async=AsyncMock(return_value=_Rows()),
+ serialize_model=Mock(return_value={}),
+ ),
+ plugin_connector=SimpleNamespace(
+ is_enable_plugin=True,
+ require_workspace_context=AsyncMock(side_effect=lambda context: context),
+ get_rag_creation_schema=AsyncMock(return_value={}),
+ get_rag_retrieval_schema=AsyncMock(return_value={}),
+ list_knowledge_engines=AsyncMock(return_value=[]),
+ list_parsers=AsyncMock(return_value=[]),
+ ),
+ )
+
+
+@pytest.mark.asyncio
+async def test_list_and_get_forward_explicit_context():
+ app = _app()
+ app.rag_mgr.get_all_knowledge_base_details.return_value = [{'uuid': 'kb-a'}]
+ app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb-a'}
+ service = KnowledgeService(app)
+
+ assert await service.get_knowledge_bases(CONTEXT) == [{'uuid': 'kb-a'}]
+ assert await service.get_knowledge_base(CONTEXT, 'kb-a') == {'uuid': 'kb-a'}
+ app.rag_mgr.get_all_knowledge_base_details.assert_awaited_once_with(CONTEXT)
+ app.rag_mgr.get_knowledge_base_details.assert_awaited_once_with(CONTEXT, 'kb-a')
+
+
+@pytest.mark.asyncio
+async def test_none_context_fails_closed_before_plugin_or_manager_access():
+ app = _app()
+ service = KnowledgeService(app)
+
+ with pytest.raises(WorkspaceRequiredError):
+ await service.get_knowledge_bases(None)
+ with pytest.raises(WorkspaceRequiredError):
+ await service.create_knowledge_base(None, {'knowledge_engine_plugin_id': 'author/engine'})
+ app.plugin_connector.get_rag_creation_schema.assert_not_awaited()
+ app.rag_mgr.create_knowledge_base.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_create_validates_schema_and_binds_context():
+ app = _app()
+ app.plugin_connector.get_rag_creation_schema.return_value = {
+ 'schema': [{'name': 'endpoint', 'label': {'en_US': 'Endpoint'}, 'required': True}]
+ }
+ app.rag_mgr.create_knowledge_base.return_value = SimpleNamespace(uuid='kb-created')
+ service = KnowledgeService(app)
+
+ with pytest.raises(ValueError, match='Endpoint is required'):
+ await service.create_knowledge_base(
+ CONTEXT,
+ {'knowledge_engine_plugin_id': 'author/engine'},
+ )
+
+ result = await service.create_knowledge_base(
+ CONTEXT,
+ {
+ 'name': 'KB',
+ 'description': 'desc',
+ 'knowledge_engine_plugin_id': 'author/engine',
+ 'creation_settings': {'endpoint': 'https://example.invalid'},
+ },
+ )
+ assert result == 'kb-created'
+ app.rag_mgr.create_knowledge_base.assert_awaited_once_with(
+ CONTEXT,
+ name='KB',
+ knowledge_engine_plugin_id='author/engine',
+ creation_settings={'endpoint': 'https://example.invalid'},
+ retrieval_settings={},
+ description='desc',
+ )
+
+
+@pytest.mark.asyncio
+async def test_update_rejects_guessed_uuid_and_scopes_reload():
+ app = _app()
+ service = KnowledgeService(app)
+
+ with pytest.raises(WorkspaceNotFoundError):
+ await service.update_knowledge_base(CONTEXT, 'kb-other', {'name': 'stolen'})
+ app.persistence_mgr.execute_async.assert_not_awaited()
+
+ app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb-a', 'workspace_uuid': 'workspace-a'}
+ await service.update_knowledge_base(CONTEXT, 'kb-a', {'name': 'updated', 'uuid': 'ignored'})
+ app.rag_mgr.remove_knowledge_base_from_runtime.assert_awaited_once_with(CONTEXT, 'kb-a')
+ app.rag_mgr.load_knowledge_base.assert_awaited_once_with(
+ CONTEXT,
+ {'uuid': 'kb-a', 'workspace_uuid': 'workspace-a'},
+ )
+
+
+@pytest.mark.asyncio
+async def test_runtime_retrieve_uses_execution_context():
+ app = _app()
+ entry = SimpleNamespace(model_dump=Mock(return_value={'id': 'entry-a'}))
+ runtime_kb = SimpleNamespace(retrieve=AsyncMock(return_value=[entry]))
+ app.rag_mgr.get_knowledge_base_by_uuid.return_value = runtime_kb
+ service = KnowledgeService(app)
+
+ assert await service.retrieve_knowledge_base(CONTEXT, 'kb-a', 'query', {'top_k': 3}) == [{'id': 'entry-a'}]
+ runtime_kb.retrieve.assert_awaited_once_with(CONTEXT, 'query', settings={'top_k': 3})
+
+
+@pytest.mark.asyncio
+async def test_runtime_retrieve_cross_workspace_uuid_is_not_found():
+ app = _app()
+ service = KnowledgeService(app)
+ with pytest.raises(WorkspaceNotFoundError):
+ await service.retrieve_knowledge_base(CONTEXT, 'kb-other', 'query')
+
+
+@pytest.mark.asyncio
+async def test_file_listing_checks_parent_knowledge_base_first():
+ app = _app()
+ service = KnowledgeService(app)
+ with pytest.raises(WorkspaceNotFoundError):
+ await service.get_files_by_knowledge_base(CONTEXT, 'kb-other')
+ app.persistence_mgr.execute_async.assert_not_awaited()
+
+ app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb-a'}
+ row = SimpleNamespace(uuid='file-a')
+ app.persistence_mgr.execute_async.return_value = _Rows([row])
+ app.persistence_mgr.serialize_model.return_value = {'uuid': 'file-a'}
+ assert await service.get_files_by_knowledge_base(CONTEXT, 'kb-a') == [{'uuid': 'file-a'}]
+
+
+@pytest.mark.asyncio
+async def test_store_and_delete_file_require_runtime_parent_and_capability():
+ app = _app()
+ runtime_kb = SimpleNamespace(
+ store_file=AsyncMock(return_value='task-a'),
+ delete_file=AsyncMock(),
+ )
+ app.rag_mgr.get_knowledge_base_by_uuid.return_value = runtime_kb
+ app.rag_mgr.get_knowledge_base_details.return_value = {'knowledge_engine': {'capabilities': ['doc_ingestion']}}
+ service = KnowledgeService(app)
+
+ assert await service.store_file(CONTEXT, 'kb-a', 'upload.pdf', 'author/parser') == 'task-a'
+ runtime_kb.store_file.assert_awaited_once_with(CONTEXT, 'upload.pdf', parser_plugin_id='author/parser')
+ await service.delete_file(CONTEXT, 'kb-a', 'file-a')
+ runtime_kb.delete_file.assert_awaited_once_with(CONTEXT, 'file-a')
+
+
+@pytest.mark.asyncio
+async def test_delete_knowledge_base_rejects_cross_workspace_uuid():
+ app = _app()
+ service = KnowledgeService(app)
+ with pytest.raises(WorkspaceNotFoundError):
+ await service.delete_knowledge_base(CONTEXT, 'kb-other')
+ app.rag_mgr.delete_knowledge_base.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_engine_and_parser_discovery_require_context_and_filter_results():
+ app = _app()
+ app.plugin_connector.list_knowledge_engines.return_value = [{'plugin_id': 'author/engine'}]
+ app.plugin_connector.list_parsers.return_value = [
+ {'id': 'text', 'supported_mime_types': ['text/plain']},
+ {'id': 'pdf', 'supported_mime_types': ['application/pdf']},
+ ]
+ service = KnowledgeService(app)
+
+ assert await service.list_knowledge_engines(CONTEXT) == [{'plugin_id': 'author/engine'}]
+ assert await service.list_parsers(CONTEXT, 'application/pdf') == [
+ {'id': 'pdf', 'supported_mime_types': ['application/pdf']}
+ ]
+ with pytest.raises(WorkspaceRequiredError):
+ await service.list_parsers(None)
+
+
+@pytest.mark.asyncio
+async def test_engine_discovery_rejects_connector_workspace_or_generation_mismatch():
+ app = _app()
+ app.plugin_connector.require_workspace_context.side_effect = WorkspaceNotFoundError('Plugin resource not found')
+ service = KnowledgeService(app)
+
+ with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
+ await service.list_knowledge_engines(CONTEXT)
+
+ app.plugin_connector.list_knowledge_engines.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_schema_validation_refences_before_second_runtime_call():
+ app = _app()
+ app.plugin_connector.require_workspace_context.side_effect = [
+ CONTEXT,
+ WorkspaceNotFoundError('Plugin resource not found'),
+ ]
+ service = KnowledgeService(app)
+
+ with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
+ await service.create_knowledge_base(
+ CONTEXT,
+ {'knowledge_engine_plugin_id': 'author/engine'},
+ )
+
+ app.plugin_connector.get_rag_creation_schema.assert_awaited_once()
+ app.plugin_connector.get_rag_retrieval_schema.assert_not_awaited()
+ app.rag_mgr.create_knowledge_base.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_engine_schemas_are_context_gated_and_fail_soft_on_connector_error():
+ app = _app()
+ app.plugin_connector.get_rag_creation_schema.return_value = {'schema': ['creation']}
+ app.plugin_connector.get_rag_retrieval_schema.side_effect = RuntimeError('offline')
+ service = KnowledgeService(app)
+
+ assert await service.get_engine_creation_schema(CONTEXT, 'author/engine') == {'schema': ['creation']}
+ assert await service.get_engine_retrieval_schema(CONTEXT, 'author/engine') == {}
+ with pytest.raises(WorkspaceRequiredError):
+ await service.get_engine_creation_schema(None, 'author/engine')
+
+
+# Preserve the original service regression matrix with the new explicit
+# Workspace context. These intentionally overlap a few isolation-focused
+# tests above so legacy business behavior cannot disappear behind new guards.
class TestKnowledgeServiceInit:
- """Tests for KnowledgeService initialization."""
-
def test_init_stores_app_reference(self):
- """Test that __init__ stores Application reference."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
+ app = _app()
- service = knowledge_module.KnowledgeService(mock_app)
+ service = KnowledgeService(app)
- assert service.ap is mock_app
+ assert service.ap is app
class TestGetKnowledgeBases:
- """Tests for get_knowledge_bases method."""
-
@pytest.mark.asyncio
async def test_returns_all_kb_details(self):
- """Test that it returns all knowledge base details."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.rag_mgr.get_all_knowledge_base_details = AsyncMock(return_value=[{'uuid': 'kb1', 'name': 'KB1'}])
+ app = _app()
+ app.rag_mgr.get_all_knowledge_base_details.return_value = [{'uuid': 'kb1', 'name': 'KB1'}]
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.get_knowledge_bases()
+ result = await KnowledgeService(app).get_knowledge_bases(CONTEXT)
- assert len(result) == 1
- assert result[0]['uuid'] == 'kb1'
+ assert result == [{'uuid': 'kb1', 'name': 'KB1'}]
+ app.rag_mgr.get_all_knowledge_base_details.assert_awaited_once_with(CONTEXT)
@pytest.mark.asyncio
async def test_returns_empty_list_when_no_kbs(self):
- """Test that it returns empty list when no knowledge bases."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.rag_mgr.get_all_knowledge_base_details = AsyncMock(return_value=[])
+ app = _app()
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.get_knowledge_bases()
-
- assert result == []
+ assert await KnowledgeService(app).get_knowledge_bases(CONTEXT) == []
class TestGetKnowledgeBase:
- """Tests for get_knowledge_base method."""
-
@pytest.mark.asyncio
async def test_returns_kb_details_by_uuid(self):
- """Test that it returns specific KB details."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value={'uuid': 'kb1', 'name': 'KB1'})
+ app = _app()
+ app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb1', 'name': 'KB1'}
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.get_knowledge_base('kb1')
+ result = await KnowledgeService(app).get_knowledge_base(CONTEXT, 'kb1')
- assert result['uuid'] == 'kb1'
+ assert result == {'uuid': 'kb1', 'name': 'KB1'}
@pytest.mark.asyncio
async def test_returns_none_when_not_found(self):
- """Test that it returns None when KB not found."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value=None)
+ app = _app()
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.get_knowledge_base('nonexistent')
-
- assert result is None
+ assert await KnowledgeService(app).get_knowledge_base(CONTEXT, 'nonexistent') is None
class TestCreateKnowledgeBase:
- """Tests for create_knowledge_base method."""
-
@pytest.mark.asyncio
async def test_creates_kb_with_required_fields(self):
- """Test creating KB with required plugin ID."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_kb = Mock()
- mock_kb.uuid = 'new_kb_uuid'
- mock_app.rag_mgr.create_knowledge_base = AsyncMock(return_value=mock_kb)
-
- service = knowledge_module.KnowledgeService(mock_app)
+ app = _app()
+ app.rag_mgr.create_knowledge_base.return_value = SimpleNamespace(uuid='new_kb_uuid')
+ service = KnowledgeService(app)
kb_data = {
'name': 'Test KB',
'knowledge_engine_plugin_id': 'author/engine',
'description': 'Test description',
}
- result = await service.create_knowledge_base(kb_data)
+ result = await service.create_knowledge_base(CONTEXT, kb_data)
assert result == 'new_kb_uuid'
- mock_app.rag_mgr.create_knowledge_base.assert_called_once()
+ app.rag_mgr.create_knowledge_base.assert_awaited_once_with(
+ CONTEXT,
+ name='Test KB',
+ knowledge_engine_plugin_id='author/engine',
+ creation_settings={},
+ retrieval_settings={},
+ description='Test description',
+ )
@pytest.mark.asyncio
async def test_raises_when_missing_plugin_id(self):
- """Test that ValueError is raised when plugin ID missing."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
+ app = _app()
- service = knowledge_module.KnowledgeService(mock_app)
+ with pytest.raises(ValueError, match='knowledge_engine_plugin_id is required'):
+ await KnowledgeService(app).create_knowledge_base(CONTEXT, {'name': 'Test'})
- with pytest.raises(ValueError) as exc_info:
- await service.create_knowledge_base({'name': 'Test'})
-
- assert 'knowledge_engine_plugin_id is required' in str(exc_info.value)
+ app.rag_mgr.create_knowledge_base.assert_not_awaited()
@pytest.mark.asyncio
async def test_creates_with_default_name(self):
- """Test that KB is created with default name if not provided."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_kb = Mock()
- mock_kb.uuid = 'new_kb_uuid'
- mock_app.rag_mgr.create_knowledge_base = AsyncMock(return_value=mock_kb)
+ app = _app()
+ app.rag_mgr.create_knowledge_base.return_value = SimpleNamespace(uuid='new_kb_uuid')
- service = knowledge_module.KnowledgeService(mock_app)
+ await KnowledgeService(app).create_knowledge_base(
+ CONTEXT,
+ {'knowledge_engine_plugin_id': 'author/engine'},
+ )
- await service.create_knowledge_base({'knowledge_engine_plugin_id': 'author/engine'})
-
- # Check that default name 'Untitled' was used
- call_args = mock_app.rag_mgr.create_knowledge_base.call_args
- assert call_args.kwargs['name'] == 'Untitled'
+ assert app.rag_mgr.create_knowledge_base.await_args.kwargs['name'] == 'Untitled'
class TestUpdateKnowledgeBase:
- """Tests for update_knowledge_base method."""
-
@pytest.mark.asyncio
async def test_updates_mutable_fields_only(self):
- """Test that only mutable fields are updated."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value={'uuid': 'kb1', 'name': 'Updated'})
- mock_app.rag_mgr.remove_knowledge_base_from_runtime = AsyncMock()
- mock_app.rag_mgr.load_knowledge_base = AsyncMock()
+ app = _app()
+ app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb1', 'name': 'Updated'}
+ service = KnowledgeService(app)
- service = knowledge_module.KnowledgeService(mock_app)
-
- # Pass both mutable and immutable fields
await service.update_knowledge_base(
+ CONTEXT,
'kb1',
{
'name': 'New Name',
'description': 'New desc',
- 'uuid': 'should_be_filtered', # immutable
+ 'uuid': 'should_be_filtered',
},
)
- # Check that only mutable fields were passed to update
- call_args = mock_app.persistence_mgr.execute_async.call_args
- assert call_args is not None
+ update_statement = app.persistence_mgr.execute_async.await_args_list[0].args[0]
+ params = update_statement.compile().params
+ assert params['name'] == 'New Name'
+ assert params['description'] == 'New desc'
+ assert 'uuid' not in params
+ app.rag_mgr.remove_knowledge_base_from_runtime.assert_awaited_once_with(CONTEXT, 'kb1')
@pytest.mark.asyncio
async def test_returns_early_when_no_mutable_fields(self):
- """Test that update returns early when no mutable fields provided."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
+ app = _app()
+ app.rag_mgr.get_knowledge_base_details.return_value = {'uuid': 'kb1'}
- service = knowledge_module.KnowledgeService(mock_app)
+ await KnowledgeService(app).update_knowledge_base(
+ CONTEXT,
+ 'kb1',
+ {'uuid': 'should_be_filtered'},
+ )
- # Pass only immutable fields
- await service.update_knowledge_base('kb1', {'uuid': 'should_be_filtered'})
-
- # No DB update should be called
- mock_app.persistence_mgr.execute_async.assert_not_called()
+ app.persistence_mgr.execute_async.assert_not_awaited()
+ app.rag_mgr.remove_knowledge_base_from_runtime.assert_not_awaited()
class TestCheckDocCapability:
- """Tests for _check_doc_capability method."""
-
@pytest.mark.asyncio
async def test_passes_when_capability_supported(self):
- """Test that check passes when doc_ingestion capability exists."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(
- return_value={'knowledge_engine': {'capabilities': ['doc_ingestion']}}
- )
+ app = _app()
+ app.rag_mgr.get_knowledge_base_details.return_value = {'knowledge_engine': {'capabilities': ['doc_ingestion']}}
- service = knowledge_module.KnowledgeService(mock_app)
-
- await service._check_doc_capability('kb1', 'document upload')
-
- # No exception raised means success
+ await KnowledgeService(app)._check_doc_capability(CONTEXT, 'kb1', 'document upload')
@pytest.mark.asyncio
async def test_raises_when_kb_not_found(self):
- """Test that Exception is raised when KB not found."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(return_value=None)
+ app = _app()
- service = knowledge_module.KnowledgeService(mock_app)
-
- with pytest.raises(Exception) as exc_info:
- await service._check_doc_capability('nonexistent', 'test operation')
-
- assert 'Knowledge base not found' in str(exc_info.value)
+ with pytest.raises(WorkspaceNotFoundError, match='Knowledge base not found'):
+ await KnowledgeService(app)._check_doc_capability(
+ CONTEXT,
+ 'nonexistent',
+ 'test operation',
+ )
@pytest.mark.asyncio
async def test_raises_when_capability_not_supported(self):
- """Test that Exception is raised when doc_ingestion not in capabilities."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.rag_mgr.get_knowledge_base_details = AsyncMock(
- return_value={'knowledge_engine': {'capabilities': ['other_capability']}}
- )
+ app = _app()
+ app.rag_mgr.get_knowledge_base_details.return_value = {
+ 'knowledge_engine': {'capabilities': ['other_capability']}
+ }
- service = knowledge_module.KnowledgeService(mock_app)
-
- with pytest.raises(Exception) as exc_info:
- await service._check_doc_capability('kb1', 'document upload')
-
- assert 'does not support document upload' in str(exc_info.value)
+ with pytest.raises(Exception, match='does not support document upload'):
+ await KnowledgeService(app)._check_doc_capability(
+ CONTEXT,
+ 'kb1',
+ 'document upload',
+ )
class TestListKnowledgeEngines:
- """Tests for list_knowledge_engines method."""
-
@pytest.mark.asyncio
async def test_returns_engines_from_plugin_connector(self):
- """Test that it returns knowledge engines from plugin connector."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.plugin_connector.list_knowledge_engines = AsyncMock(
- return_value=[{'id': 'engine1', 'name': 'Engine 1'}]
- )
+ app = _app()
+ app.plugin_connector.list_knowledge_engines.return_value = [{'id': 'engine1', 'name': 'Engine 1'}]
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.list_knowledge_engines()
+ result = await KnowledgeService(app).list_knowledge_engines(CONTEXT)
- assert len(result) == 1
- assert result[0]['id'] == 'engine1'
+ assert result == [{'id': 'engine1', 'name': 'Engine 1'}]
@pytest.mark.asyncio
async def test_returns_empty_when_plugin_disabled(self):
- """Test that it returns empty list when plugin disabled."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.plugin_connector.is_enable_plugin = False
+ app = _app()
+ app.plugin_connector.is_enable_plugin = False
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.list_knowledge_engines()
-
- assert result == []
+ assert await KnowledgeService(app).list_knowledge_engines(CONTEXT) == []
+ app.plugin_connector.list_knowledge_engines.assert_not_awaited()
@pytest.mark.asyncio
async def test_returns_empty_on_exception(self):
- """Test that it returns empty list and logs warning on exception."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.plugin_connector.list_knowledge_engines = AsyncMock(side_effect=Exception('Connection error'))
+ app = _app()
+ app.plugin_connector.list_knowledge_engines.side_effect = RuntimeError('Connection error')
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.list_knowledge_engines()
-
- assert result == []
- mock_app.logger.warning.assert_called_once()
+ assert await KnowledgeService(app).list_knowledge_engines(CONTEXT) == []
+ app.logger.warning.assert_called_once()
class TestListParsers:
- """Tests for list_parsers method."""
-
@pytest.mark.asyncio
async def test_returns_all_parsers(self):
- """Test that it returns all parsers when no MIME type filter."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.plugin_connector.list_parsers = AsyncMock(
- return_value=[
- {'id': 'parser1', 'supported_mime_types': ['text/plain']},
- {'id': 'parser2', 'supported_mime_types': ['application/pdf']},
- ]
- )
+ app = _app()
+ app.plugin_connector.list_parsers.return_value = [
+ {'id': 'parser1', 'supported_mime_types': ['text/plain']},
+ {'id': 'parser2', 'supported_mime_types': ['application/pdf']},
+ ]
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.list_parsers()
+ result = await KnowledgeService(app).list_parsers(CONTEXT)
assert len(result) == 2
@pytest.mark.asyncio
async def test_filters_by_mime_type(self):
- """Test that it filters parsers by MIME type."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.plugin_connector.list_parsers = AsyncMock(
- return_value=[
- {'id': 'parser1', 'supported_mime_types': ['text/plain']},
- {'id': 'parser2', 'supported_mime_types': ['application/pdf']},
- ]
- )
+ app = _app()
+ app.plugin_connector.list_parsers.return_value = [
+ {'id': 'parser1', 'supported_mime_types': ['text/plain']},
+ {'id': 'parser2', 'supported_mime_types': ['application/pdf']},
+ ]
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.list_parsers(mime_type='application/pdf')
+ result = await KnowledgeService(app).list_parsers(CONTEXT, 'application/pdf')
- assert len(result) == 1
- assert result[0]['id'] == 'parser2'
+ assert result == [{'id': 'parser2', 'supported_mime_types': ['application/pdf']}]
@pytest.mark.asyncio
async def test_returns_empty_when_plugin_disabled(self):
- """Test that it returns empty list when plugin disabled."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.plugin_connector.is_enable_plugin = False
+ app = _app()
+ app.plugin_connector.is_enable_plugin = False
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.list_parsers()
-
- assert result == []
+ assert await KnowledgeService(app).list_parsers(CONTEXT) == []
+ app.plugin_connector.list_parsers.assert_not_awaited()
class TestGetEngineSchemas:
- """Tests for get_engine_creation_schema and get_engine_retrieval_schema."""
-
@pytest.mark.asyncio
async def test_returns_creation_schema(self):
- """Test that it returns creation schema for engine."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.plugin_connector.get_rag_creation_schema = AsyncMock(
- return_value={'properties': {'name': {'type': 'string'}}}
- )
+ app = _app()
+ app.plugin_connector.get_rag_creation_schema.return_value = {'properties': {'name': {'type': 'string'}}}
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.get_engine_creation_schema('author/engine')
+ result = await KnowledgeService(app).get_engine_creation_schema(
+ CONTEXT,
+ 'author/engine',
+ )
assert 'properties' in result
@pytest.mark.asyncio
async def test_returns_retrieval_schema(self):
- """Test that it returns retrieval schema for engine."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.plugin_connector.get_rag_retrieval_schema = AsyncMock(
- return_value={'properties': {'top_k': {'type': 'integer'}}}
- )
+ app = _app()
+ app.plugin_connector.get_rag_retrieval_schema.return_value = {'properties': {'top_k': {'type': 'integer'}}}
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.get_engine_retrieval_schema('author/engine')
+ result = await KnowledgeService(app).get_engine_retrieval_schema(
+ CONTEXT,
+ 'author/engine',
+ )
assert 'properties' in result
@pytest.mark.asyncio
async def test_returns_empty_dict_on_exception(self):
- """Test that it returns empty dict and logs warning on exception."""
- knowledge_module = get_knowledge_service_module()
- mock_app = create_mock_app()
- mock_app.plugin_connector.get_rag_creation_schema = AsyncMock(side_effect=Exception('Plugin error'))
+ app = _app()
+ app.plugin_connector.get_rag_creation_schema.side_effect = RuntimeError('Plugin error')
- service = knowledge_module.KnowledgeService(mock_app)
- result = await service.get_engine_creation_schema('author/engine')
+ result = await KnowledgeService(app).get_engine_creation_schema(
+ CONTEXT,
+ 'author/engine',
+ )
assert result == {}
- mock_app.logger.warning.assert_called_once()
+ app.logger.warning.assert_called_once()
+
+
+class TestKnowledgeBaseSecretViews:
+ @pytest.mark.asyncio
+ async def test_creation_settings_are_redacted_for_resource_view_only(self):
+ app = _app()
+ raw = {
+ 'uuid': 'kb-secret',
+ 'creation_settings': {
+ 'dify_apikey': 'dify-secret',
+ 'headers': {'Authorization': 'Bearer secret'},
+ },
+ }
+ app.rag_mgr.get_all_knowledge_base_details.return_value = [raw]
+ service = KnowledgeService(app)
+
+ redacted = await service.get_knowledge_bases(CONTEXT)
+ manager_view = await service.get_knowledge_bases(CONTEXT, include_secret=True)
+
+ assert redacted[0]['creation_settings']['dify_apikey'] == '***'
+ assert redacted[0]['creation_settings']['headers']['Authorization'] == '***'
+ assert manager_view[0]['creation_settings']['dify_apikey'] == 'dify-secret'
+ assert raw['creation_settings']['dify_apikey'] == 'dify-secret'
+
+ @pytest.mark.asyncio
+ async def test_new_masked_creation_secret_is_rejected(self):
+ app = _app()
+
+ with pytest.raises(ValueError, match='no existing value'):
+ await KnowledgeService(app).create_knowledge_base(
+ CONTEXT,
+ {
+ 'knowledge_engine_plugin_id': 'author/engine',
+ 'creation_settings': {'dify_apikey': '***'},
+ },
+ )
+
+ app.rag_mgr.create_knowledge_base.assert_not_awaited()
diff --git a/tests/unit_tests/api/service/test_maintenance_service.py b/tests/unit_tests/api/service/test_maintenance_service.py
index 8d5b5b0df..bd03c1b70 100644
--- a/tests/unit_tests/api/service/test_maintenance_service.py
+++ b/tests/unit_tests/api/service/test_maintenance_service.py
@@ -19,10 +19,32 @@ from types import SimpleNamespace
import datetime
from pathlib import Path
+import sqlalchemy
+from sqlalchemy.ext.asyncio import create_async_engine
+
+from langbot.pkg.api.http.authz import WorkspaceRequiredError
from langbot.pkg.api.http.service.maintenance import MaintenanceService
+from langbot.pkg.api.http.context import ExecutionContext, PrincipalContext, PrincipalType
+from langbot.pkg.entity.persistence.base import Base
+from langbot.pkg.entity.persistence.bstorage import BinaryStorage
+from langbot.pkg.entity.persistence.monitoring import MonitoringMessage
+from langbot.pkg.entity.persistence.workspace import Workspace
pytestmark = pytest.mark.asyncio
+TEST_CONTEXT = ExecutionContext(
+ instance_uuid='test-instance',
+ workspace_uuid='test-workspace',
+ placement_generation=1,
+)
+
+
+@pytest.fixture(autouse=True)
+def assume_oss_singleton(monkeypatch):
+ async def is_oss_singleton(_self, _context):
+ return True
+
+ monkeypatch.setattr(MaintenanceService, '_is_oss_singleton', is_oss_singleton)
def _create_mock_result(scalar_value=None):
@@ -32,6 +54,14 @@ def _create_mock_result(scalar_value=None):
return result
+def _scoped_storage_manager():
+ prefix = 'instances/i/workspaces/w/generations/1/owners/upload/o/'
+ return SimpleNamespace(
+ scoped_prefix=Mock(return_value=prefix),
+ is_scoped_object_key=Mock(side_effect=lambda key, **_: key == f'{prefix}uploaded_file.txt'),
+ )
+
+
class TestMaintenanceServiceCleanupExpiredFiles:
"""Tests for cleanup_expired_files method."""
@@ -39,6 +69,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
"""Uses default retention days when config not set."""
# Setup
ap = SimpleNamespace()
+ ap.storage_mgr = _scoped_storage_manager()
ap.instance_config = SimpleNamespace()
ap.instance_config.data = {}
ap.storage_mgr = SimpleNamespace()
@@ -58,7 +89,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
service._cleanup_expired_log_files = Mock(return_value=0) # NOT async!
# Execute
- result = await service.cleanup_expired_files()
+ result = await service.cleanup_expired_files(TEST_CONTEXT)
# Verify - returns counts
assert 'uploaded_files' in result
@@ -95,7 +126,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
service._cleanup_expired_log_files = Mock(return_value=3) # NOT async
# Execute
- result = await service.cleanup_expired_files()
+ result = await service.cleanup_expired_files(TEST_CONTEXT)
# Verify
assert result['uploaded_files'] == 2
@@ -124,7 +155,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
service._cleanup_expired_log_files = Mock(return_value=0) # NOT async
# Execute
- result = await service.cleanup_expired_files()
+ result = await service.cleanup_expired_files(TEST_CONTEXT)
# Verify
assert result['uploaded_files'] == 1
@@ -159,7 +190,7 @@ class TestMaintenanceServiceCleanupExpiredFiles:
service._cleanup_expired_log_files = Mock(return_value=0) # NOT async
# Execute
- result = await service.cleanup_expired_files()
+ result = await service.cleanup_expired_files(TEST_CONTEXT)
# Verify - warning logged, defaults used
assert ap.logger.warning.called
@@ -196,7 +227,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
service._expired_log_candidates = Mock(return_value=[])
# Execute
- result = await service.get_storage_analysis()
+ result = await service.get_storage_analysis(TEST_CONTEXT)
# Verify
assert 'generated_at' in result
@@ -229,7 +260,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
service._expired_log_candidates = Mock(return_value=[])
# Execute
- result = await service.get_storage_analysis()
+ result = await service.get_storage_analysis(TEST_CONTEXT)
# Verify - all sections present
sections = {s['key'] for s in result['sections']}
@@ -265,7 +296,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
service._expired_log_candidates = Mock(return_value=[])
# Execute
- result = await service.get_storage_analysis()
+ result = await service.get_storage_analysis(TEST_CONTEXT)
# Verify
assert result['database']['type'] == 'postgresql'
@@ -294,7 +325,7 @@ class TestMaintenanceServiceGetStorageAnalysis:
service._expired_log_candidates = Mock(return_value=[{'name': 'old_log', 'size_bytes': 50}])
# Execute
- result = await service.get_storage_analysis()
+ result = await service.get_storage_analysis(TEST_CONTEXT)
# Verify
assert len(result['cleanup_candidates']['uploaded_files']) == 1
@@ -316,7 +347,7 @@ class TestMaintenanceServiceMonitoringCounts:
service = MaintenanceService(ap)
# Execute
- result = await service._monitoring_counts()
+ result = await service._monitoring_counts(TEST_CONTEXT)
# Verify - all table keys present
assert 'messages' in result
@@ -338,7 +369,7 @@ class TestMaintenanceServiceMonitoringCounts:
service = MaintenanceService(ap)
# Execute
- result = await service._monitoring_counts()
+ result = await service._monitoring_counts(TEST_CONTEXT)
# Verify - all zero
assert all(v == 0 for v in result.values())
@@ -374,7 +405,7 @@ class TestMaintenanceServiceBinaryStorageStats:
service = MaintenanceService(ap)
# Execute
- result = await service._binary_storage_stats()
+ result = await service._binary_storage_stats(TEST_CONTEXT)
# Verify
assert result['count'] == 10
@@ -404,7 +435,7 @@ class TestMaintenanceServiceBinaryStorageStats:
service = MaintenanceService(ap)
# Execute
- result = await service._binary_storage_stats()
+ result = await service._binary_storage_stats(TEST_CONTEXT)
# Verify - warning logged, size_bytes None or 0
assert ap.logger.warning.called
@@ -618,11 +649,13 @@ class TestMaintenanceServiceIsUploadedFileKey:
"""Returns True for valid upload file key."""
# Setup
ap = SimpleNamespace()
+ ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
# Execute - simple filename without path
- result = service._is_uploaded_file_key('uploaded_file.txt')
+ key = f'{ap.storage_mgr.scoped_prefix(TEST_CONTEXT, owner_type="upload")}uploaded_file.txt'
+ result = service._is_uploaded_file_key(TEST_CONTEXT, key)
# Verify
assert result is True
@@ -631,11 +664,12 @@ class TestMaintenanceServiceIsUploadedFileKey:
"""Returns False for key with path separator."""
# Setup
ap = SimpleNamespace()
+ ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
# Execute - key with path
- result = service._is_uploaded_file_key('path/to/file.txt')
+ result = service._is_uploaded_file_key(TEST_CONTEXT, 'path/to/file.txt')
# Verify
assert result is False
@@ -644,11 +678,12 @@ class TestMaintenanceServiceIsUploadedFileKey:
"""Returns False for plugin config prefix."""
# Setup
ap = SimpleNamespace()
+ ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
# Execute - plugin config file
- result = service._is_uploaded_file_key('plugin_config_some_plugin.json')
+ result = service._is_uploaded_file_key(TEST_CONTEXT, 'plugin_config_some_plugin.json')
# Verify
assert result is False
@@ -662,6 +697,7 @@ class TestMaintenanceServiceExpiredLogCandidates:
# Setup
ap = SimpleNamespace()
ap.logger = SimpleNamespace()
+ ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
@@ -748,11 +784,12 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
# Setup
ap = SimpleNamespace()
ap.logger = SimpleNamespace()
+ ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
with patch.object(Path, 'exists', return_value=False):
- result = service._expired_local_upload_candidates(7)
+ result = service._expired_local_upload_candidates(TEST_CONTEXT, 7)
# Verify
assert result == []
@@ -762,12 +799,10 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
# Setup
ap = SimpleNamespace()
ap.logger = SimpleNamespace()
+ ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
- # Mock _is_uploaded_file_key
- service._is_uploaded_file_key = Mock(side_effect=lambda key: 'plugin_config_' not in key and '/' not in key)
-
- # Create mock files - one valid, one plugin config
+ # Create one file and one non-file entry under the scoped upload root.
mock_entry_valid = Mock(spec=Path)
mock_entry_valid.is_file = Mock(return_value=True)
mock_entry_valid.name = 'valid_upload.txt'
@@ -775,9 +810,10 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
mock_stat.st_size = 100
mock_stat.st_mtime = 0 # Very old
mock_entry_valid.stat = Mock(return_value=mock_stat)
+ mock_entry_valid.relative_to = Mock(return_value=Path('scoped/valid_upload.txt'))
mock_entry_plugin = Mock(spec=Path)
- mock_entry_plugin.is_file = Mock(return_value=True)
+ mock_entry_plugin.is_file = Mock(return_value=False)
mock_entry_plugin.name = 'plugin_config_test.json'
mock_stat2 = Mock()
mock_stat2.st_size = 200
@@ -785,23 +821,22 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
mock_entry_plugin.stat = Mock(return_value=mock_stat2)
with patch.object(Path, 'exists', return_value=True):
- with patch.object(Path, 'iterdir') as mock_iterdir:
- mock_iterdir.return_value = [mock_entry_valid, mock_entry_plugin]
- result = service._expired_local_upload_candidates(7)
+ with patch.object(Path, 'rglob') as mock_rglob:
+ mock_rglob.return_value = [mock_entry_valid, mock_entry_plugin]
+ result = service._expired_local_upload_candidates(TEST_CONTEXT, 7)
# Verify - only valid upload included
assert len(result) == 1
- assert result[0]['key'] == 'valid_upload.txt'
+ assert result[0]['key'] == 'scoped/valid_upload.txt'
def test_expired_local_upload_candidates_includes_path(self):
"""Includes path when include_paths=True."""
# Setup
ap = SimpleNamespace()
ap.logger = SimpleNamespace()
+ ap.storage_mgr = _scoped_storage_manager()
service = MaintenanceService(ap)
- service._is_uploaded_file_key = Mock(return_value=True)
-
mock_entry = Mock(spec=Path)
mock_entry.is_file = Mock(return_value=True)
mock_entry.name = 'old_file.txt'
@@ -810,11 +845,152 @@ class TestMaintenanceServiceExpiredLocalUploadCandidates:
mock_stat.st_size = 100
mock_stat.st_mtime = 0
mock_entry.stat = Mock(return_value=mock_stat)
+ mock_entry.relative_to = Mock(return_value=Path('scoped/old_file.txt'))
with patch.object(Path, 'exists', return_value=True):
- with patch.object(Path, 'iterdir') as mock_iterdir:
- mock_iterdir.return_value = [mock_entry]
- result = service._expired_local_upload_candidates(7, include_paths=True)
+ with patch.object(Path, 'rglob') as mock_rglob:
+ mock_rglob.return_value = [mock_entry]
+ result = service._expired_local_upload_candidates(
+ TEST_CONTEXT,
+ 7,
+ include_paths=True,
+ )
# Verify - path included
assert 'path' in result[0]
+
+
+ISOLATION_WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
+ISOLATION_WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
+
+
+def _tenant_context(workspace_uuid: str) -> ExecutionContext:
+ return ExecutionContext(
+ instance_uuid='instance',
+ workspace_uuid=workspace_uuid,
+ placement_generation=1,
+ trigger_principal=PrincipalContext(PrincipalType.SYSTEM),
+ )
+
+
+class _RealPersistenceManager:
+ def __init__(self, engine):
+ self.engine = engine
+
+ async def execute_async(self, *args, **kwargs):
+ async with self.engine.connect() as connection:
+ result = await connection.execute(*args, **kwargs)
+ await connection.commit()
+ return result
+
+
+@pytest.fixture
+async def tenant_maintenance_service(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "maintenance.db"}')
+ async with engine.begin() as connection:
+ await connection.run_sync(Base.metadata.create_all)
+ await connection.execute(
+ sqlalchemy.insert(Workspace),
+ [
+ {
+ 'uuid': ISOLATION_WORKSPACE_A,
+ 'instance_uuid': 'instance',
+ 'name': 'A',
+ 'slug': 'a',
+ 'source': 'cloud_projection',
+ },
+ {
+ 'uuid': ISOLATION_WORKSPACE_B,
+ 'instance_uuid': 'instance',
+ 'name': 'B',
+ 'slug': 'b',
+ 'source': 'cloud_projection',
+ },
+ ],
+ )
+ now = datetime.datetime.now(datetime.UTC).replace(tzinfo=None)
+ await connection.execute(
+ sqlalchemy.insert(MonitoringMessage),
+ [
+ {
+ 'id': 'message-a',
+ 'workspace_uuid': ISOLATION_WORKSPACE_A,
+ 'timestamp': now,
+ 'bot_id': 'bot',
+ 'bot_name': 'Bot',
+ 'pipeline_id': 'pipeline',
+ 'pipeline_name': 'Pipeline',
+ 'message_content': 'A',
+ 'session_id': 'same-session',
+ 'status': 'success',
+ 'level': 'info',
+ },
+ {
+ 'id': 'message-b',
+ 'workspace_uuid': ISOLATION_WORKSPACE_B,
+ 'timestamp': now,
+ 'bot_id': 'bot',
+ 'bot_name': 'Bot',
+ 'pipeline_id': 'pipeline',
+ 'pipeline_name': 'Pipeline',
+ 'message_content': 'B',
+ 'session_id': 'same-session',
+ 'status': 'success',
+ 'level': 'info',
+ },
+ ],
+ )
+ await connection.execute(
+ sqlalchemy.insert(BinaryStorage),
+ [
+ {
+ 'workspace_uuid': ISOLATION_WORKSPACE_A,
+ 'unique_key': 'a',
+ 'key': 'same',
+ 'owner_type': 'plugin',
+ 'owner': 'same',
+ 'value': b'aaa',
+ },
+ {
+ 'workspace_uuid': ISOLATION_WORKSPACE_B,
+ 'unique_key': 'b',
+ 'key': 'same',
+ 'owner_type': 'plugin',
+ 'owner': 'same',
+ 'value': b'bbbbb',
+ },
+ ],
+ )
+
+ application = SimpleNamespace(
+ persistence_mgr=_RealPersistenceManager(engine),
+ instance_config=SimpleNamespace(data={}),
+ logger=SimpleNamespace(warning=lambda *_: None),
+ )
+ yield MaintenanceService(application)
+ await engine.dispose()
+
+
+async def test_cleanup_requires_execution_context(tenant_maintenance_service):
+ with pytest.raises(WorkspaceRequiredError):
+ await tenant_maintenance_service.cleanup_expired_files(None)
+
+
+async def test_monitoring_counts_are_workspace_scoped(tenant_maintenance_service):
+ counts_a = await tenant_maintenance_service._monitoring_counts(_tenant_context(ISOLATION_WORKSPACE_A))
+ counts_b = await tenant_maintenance_service._monitoring_counts(_tenant_context(ISOLATION_WORKSPACE_B))
+ assert counts_a['messages'] == 1
+ assert counts_b['messages'] == 1
+
+
+async def test_binary_storage_stats_are_workspace_scoped(tenant_maintenance_service):
+ stats_a = await tenant_maintenance_service._binary_storage_stats(_tenant_context(ISOLATION_WORKSPACE_A))
+ stats_b = await tenant_maintenance_service._binary_storage_stats(_tenant_context(ISOLATION_WORKSPACE_B))
+ assert stats_a == {'count': 1, 'size_bytes': 3}
+ assert stats_b == {'count': 1, 'size_bytes': 5}
+
+
+async def test_path_helpers_handle_missing_paths(tenant_maintenance_service, tmp_path):
+ missing = tmp_path / 'missing'
+ assert tenant_maintenance_service._path_size(missing) == 0
+ assert tenant_maintenance_service._file_count(missing) == 0
diff --git a/tests/unit_tests/api/service/test_mcp_service.py b/tests/unit_tests/api/service/test_mcp_service.py
index ea08897a1..4b3504c46 100644
--- a/tests/unit_tests/api/service/test_mcp_service.py
+++ b/tests/unit_tests/api/service/test_mcp_service.py
@@ -13,17 +13,59 @@ Source: src/langbot/pkg/api/http/service/mcp.py
from __future__ import annotations
+import copy
import pytest
from unittest.mock import AsyncMock, Mock, MagicMock
from types import SimpleNamespace
import uuid
-from langbot.pkg.api.http.service.mcp import MCPService
+from langbot.pkg.api.http.authz import Permission
+from langbot.pkg.api.http.context import (
+ ExecutionContext,
+ PrincipalContext,
+ PrincipalType,
+ RequestContext,
+ WorkspaceContext,
+)
+from langbot.pkg.api.http.service.mcp import MCPService, redact_mcp_secrets, restore_mcp_secret_placeholders
from langbot.pkg.entity.persistence.mcp import MCPServer
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
pytestmark = pytest.mark.asyncio
+_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+)
+
+_VIEWER_CONTEXT = RequestContext(
+ instance_uuid='instance-a',
+ placement_generation=1,
+ request_id='request-a',
+ auth_type='user_token',
+ principal=PrincipalContext(
+ principal_type=PrincipalType.ACCOUNT,
+ account_uuid='account-a',
+ ),
+ workspace=WorkspaceContext(
+ workspace_uuid='workspace-a',
+ membership_uuid='membership-a',
+ role='viewer',
+ permissions=frozenset({Permission.RESOURCE_VIEW.value}),
+ ),
+)
+
+
+def _service(ap: SimpleNamespace) -> MCPService:
+ ap.workspace_service = SimpleNamespace(
+ get_execution_binding=AsyncMock(return_value=SimpleNamespace(instance_uuid=_CONTEXT.instance_uuid))
+ )
+ if not hasattr(ap, 'logger'):
+ ap.logger = Mock()
+ return MCPService(ap)
+
def _create_mock_mcp_server(
server_uuid: str = None,
@@ -42,11 +84,13 @@ def _create_mock_mcp_server(
return server
-def _create_mock_result(items: list = None, first_item=None):
+def _create_mock_result(items: list = None, first_item=None, *, scalar_value=0, rowcount=1):
"""Create mock result object for persistence queries."""
result = Mock()
result.all = Mock(return_value=items or [])
result.first = Mock(return_value=first_item)
+ result.scalar = Mock(return_value=scalar_value)
+ result.rowcount = rowcount
return result
@@ -64,10 +108,10 @@ class TestMCPServiceGetRuntimeInfo:
mock_session.get_runtime_info_dict = Mock(return_value={'status': 'running', 'tools': 5})
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=mock_session)
- service = MCPService(ap)
+ service = _service(ap)
# Execute
- result = await service.get_runtime_info('test-server')
+ result = await service.get_runtime_info(_CONTEXT, 'test-server')
# Verify
assert result is not None
@@ -81,10 +125,10 @@ class TestMCPServiceGetRuntimeInfo:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=None)
- service = MCPService(ap)
+ service = _service(ap)
# Execute
- result = await service.get_runtime_info('nonexistent-server')
+ result = await service.get_runtime_info(_CONTEXT, 'nonexistent-server')
# Verify
assert result is None
@@ -101,12 +145,13 @@ class TestMCPServiceResources:
return_value=[{'uri_template': 'file:///{path}', 'name': 'files'}]
)
- service = MCPService(ap)
+ service = _service(ap)
+ service._require_server = AsyncMock(return_value=(_CONTEXT, {'name': 'docs'}))
- result = await service.get_mcp_server_resource_templates('docs')
+ result = await service.get_mcp_server_resource_templates(_CONTEXT, 'docs')
assert result == [{'uri_template': 'file:///{path}', 'name': 'files'}]
- ap.tool_mgr.mcp_tool_loader.get_resource_templates.assert_awaited_once_with('docs')
+ ap.tool_mgr.mcp_tool_loader.get_resource_templates.assert_awaited_once_with(_CONTEXT, 'docs')
async def test_read_resource_envelope_uses_ui_preview_source(self):
ap = SimpleNamespace()
@@ -121,9 +166,11 @@ class TestMCPServiceResources:
}
)
- service = MCPService(ap)
+ service = _service(ap)
+ service._require_server = AsyncMock(return_value=(_CONTEXT, {'name': 'docs'}))
result = await service.read_mcp_server_resource_envelope(
+ _CONTEXT,
'docs',
'file:///README.md',
max_bytes=4096,
@@ -132,6 +179,7 @@ class TestMCPServiceResources:
assert result['source'] == 'ui_preview'
ap.tool_mgr.mcp_tool_loader.read_resource_envelope.assert_awaited_once_with(
+ _CONTEXT,
'docs',
'file:///README.md',
include_blob=True,
@@ -156,12 +204,12 @@ class TestMCPServiceGetMCPServers:
'name': entity.name,
}
)
- ap.tool_mgr = None
+ ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
- service = MCPService(ap)
+ service = _service(ap)
# Execute
- result = await service.get_mcp_servers()
+ result = await service.get_mcp_servers(_CONTEXT)
# Verify
assert result == []
@@ -185,12 +233,12 @@ class TestMCPServiceGetMCPServers:
'mode': entity.mode,
}
)
- ap.tool_mgr = None
+ ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
- service = MCPService(ap)
+ service = _service(ap)
# Execute
- result = await service.get_mcp_servers()
+ result = await service.get_mcp_servers(_CONTEXT)
# Verify
assert len(result) == 2
@@ -215,17 +263,90 @@ class TestMCPServiceGetMCPServers:
)
ap.tool_mgr = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
- ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=None)
+ runtime_session = SimpleNamespace(get_runtime_info_dict=Mock(return_value={'status': 'connected'}))
+ ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=runtime_session)
- service = MCPService(ap)
- service.get_runtime_info = AsyncMock(return_value={'status': 'connected'})
+ service = _service(ap)
# Execute
- result = await service.get_mcp_servers(contain_runtime_info=True)
+ result = await service.get_mcp_servers(_CONTEXT, contain_runtime_info=True)
# Verify - runtime info included
assert result[0]['runtime_info'] == {'status': 'connected'}
+ async def test_resource_view_list_and_detail_redact_secrets_without_mutating_raw_data(self):
+ ap = SimpleNamespace()
+ ap.persistence_mgr = SimpleNamespace()
+ server = _create_mock_mcp_server(name='Secret Server')
+ serialized = {
+ 'uuid': 'secret-uuid',
+ 'name': 'Secret Server',
+ 'enable': True,
+ 'extra_args': {
+ 'url': (
+ 'https://mcp-user:mcp-password@mcp.invalid/connect'
+ '?token=url-secret&transport=streamable&sig=signed-secret'
+ ),
+ 'headers': {
+ 'Authorization': 'Bearer top-secret',
+ 'X-API-Key': 'api-secret',
+ 'Accept': 'application/json',
+ },
+ 'env': {
+ 'ACCESS_TOKEN': 'access-secret',
+ 'TOKENIZER': 'public-model-name',
+ },
+ 'credentials': {
+ 'username': 'service-user',
+ 'password': 'password-secret',
+ },
+ 'public_key': 'public-value',
+ },
+ }
+ original = copy.deepcopy(serialized)
+ ap.persistence_mgr.execute_async = AsyncMock(
+ side_effect=[
+ _create_mock_result([server]),
+ _create_mock_result(first_item=server),
+ ]
+ )
+ ap.persistence_mgr.serialize_model = Mock(return_value=serialized)
+ ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
+ service = _service(ap)
+
+ listed = await service.get_mcp_servers(_VIEWER_CONTEXT)
+ detail = await service.get_mcp_server_by_name(_VIEWER_CONTEXT, 'Secret Server')
+
+ for response in (listed[0], detail):
+ assert response['extra_args']['url'] == (
+ 'https://***@mcp.invalid/connect?token=***&transport=streamable&sig=***'
+ )
+ assert response['extra_args']['headers'] == {
+ 'Authorization': '***',
+ 'X-API-Key': '***',
+ 'Accept': 'application/json',
+ }
+ assert response['extra_args']['env'] == {
+ 'ACCESS_TOKEN': '***',
+ 'TOKENIZER': 'public-model-name',
+ }
+ assert response['extra_args']['credentials'] == {
+ 'username': '***',
+ 'password': '***',
+ }
+ assert response['extra_args']['public_key'] == 'public-value'
+ assert serialized == original
+
+ async def test_redacted_url_roundtrip_restores_persisted_credentials(self):
+ persisted = {
+ 'extra_args': {'url': 'https://mcp-user:mcp-password@mcp.invalid/connect?token=url-secret&transport=http'}
+ }
+
+ submitted = redact_mcp_secrets(persisted)
+
+ assert submitted['extra_args']['url'] == 'https://***@mcp.invalid/connect?token=***&transport=http'
+ assert restore_mcp_secret_placeholders(submitted, persisted) == persisted
+
class TestMCPServiceCreateMCPServer:
"""Tests for create_mcp_server method."""
@@ -241,16 +362,20 @@ class TestMCPServiceCreateMCPServer:
ap.plugin_connector.list_plugins = AsyncMock(return_value=[Mock(), Mock()]) # 2 plugins
# Mock get_mcp_servers to return 0 servers (2 plugins already)
- mock_result = _create_mock_result([])
- ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
+ ap.persistence_mgr.execute_async = AsyncMock(
+ side_effect=[
+ _create_mock_result(scalar_value=0),
+ _create_mock_result(scalar_value=2),
+ ]
+ )
ap.persistence_mgr.serialize_model = Mock(return_value={})
- ap.tool_mgr = None
+ ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
- service = MCPService(ap)
+ service = _service(ap)
# Execute & Verify - 2 plugins + new server would exceed limit
with pytest.raises(ValueError, match='Maximum number of extensions'):
- await service.create_mcp_server({'name': 'New Server'})
+ await service.create_mcp_server(_CONTEXT, {'name': 'New Server'})
async def test_create_mcp_server_no_limit(self):
"""Creates MCP server without limit when max_extensions=-1."""
@@ -271,10 +396,10 @@ class TestMCPServiceCreateMCPServer:
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
ap.persistence_mgr.serialize_model = Mock(return_value={'uuid': 'new-uuid'})
- service = MCPService(ap)
+ service = _service(ap)
# Execute
- server_uuid = await service.create_mcp_server({'name': 'New Server'})
+ server_uuid = await service.create_mcp_server(_CONTEXT, {'name': 'New Server'})
# Verify
assert server_uuid is not None
@@ -293,11 +418,11 @@ class TestMCPServiceCreateMCPServer:
ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result(first_item=existing_server))
ap.persistence_mgr.serialize_model = Mock(return_value={})
- service = MCPService(ap)
+ service = _service(ap)
# Execute & Verify
with pytest.raises(ValueError, match='MCP server already exists: Existing Server'):
- await service.create_mcp_server({'name': 'Existing Server'})
+ await service.create_mcp_server(_CONTEXT, {'name': 'Existing Server'})
async def test_create_mcp_server_loads_server(self):
"""Loads server into tool_mgr when enabled."""
@@ -330,10 +455,10 @@ class TestMCPServiceCreateMCPServer:
return_value={'uuid': 'new-uuid', 'name': 'New Server', 'enable': True}
)
- service = MCPService(ap)
+ service = _service(ap)
# Execute
- await service.create_mcp_server({'name': 'New Server', 'enable': True})
+ await service.create_mcp_server(_CONTEXT, {'name': 'New Server', 'enable': True})
# Verify - host_mcp_server was called
ap.tool_mgr.mcp_tool_loader.host_mcp_server.assert_called_once()
@@ -351,10 +476,10 @@ class TestMCPServiceCreateMCPServer:
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
ap.persistence_mgr.serialize_model = Mock(return_value={'uuid': 'new-uuid'})
- service = MCPService(ap)
+ service = _service(ap)
# Execute with enable=False
- server_uuid = await service.create_mcp_server({'name': 'New Server', 'enable': False})
+ server_uuid = await service.create_mcp_server(_CONTEXT, {'name': 'New Server', 'enable': False})
# Verify - no tool_mgr load attempt
assert server_uuid is not None
@@ -379,13 +504,11 @@ class TestMCPServiceGetMCPServerByName:
'runtime_info': None,
}
)
- ap.tool_mgr = None
-
- service = MCPService(ap)
- service.get_runtime_info = AsyncMock(return_value=None)
+ ap.tool_mgr = SimpleNamespace(mcp_tool_loader=SimpleNamespace(get_session=Mock(return_value=None)))
+ service = _service(ap)
# Execute
- result = await service.get_mcp_server_by_name('Found Server')
+ result = await service.get_mcp_server_by_name(_CONTEXT, 'Found Server')
# Verify
assert result is not None
@@ -400,10 +523,10 @@ class TestMCPServiceGetMCPServerByName:
mock_result = _create_mock_result(first_item=None)
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
- service = MCPService(ap)
+ service = _service(ap)
# Execute
- result = await service.get_mcp_server_by_name('Nonexistent Server')
+ result = await service.get_mcp_server_by_name(_CONTEXT, 'Nonexistent Server')
# Verify
assert result is None
@@ -421,8 +544,10 @@ class TestMCPServiceUpdateMCPServer:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.sessions = {'Old Server': Mock()}
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
+ ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=True)
old_server = _create_mock_mcp_server(name='Old Server', enable=True)
+ updated_server = _create_mock_mcp_server(name='Old Server', enable=False)
call_count = 0
@@ -431,14 +556,23 @@ class TestMCPServiceUpdateMCPServer:
call_count += 1
if call_count == 1:
return _create_mock_result(first_item=old_server)
- return Mock() # Update
+ if call_count == 2:
+ return _create_mock_result()
+ return _create_mock_result(first_item=updated_server)
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
+ ap.persistence_mgr.serialize_model = Mock(
+ side_effect=lambda _model, entity: {
+ 'uuid': 'test-uuid',
+ 'name': entity.name,
+ 'enable': entity.enable,
+ }
+ )
- service = MCPService(ap)
+ service = _service(ap)
# Execute - disable server
- await service.update_mcp_server('test-uuid', {'enable': False})
+ await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': False})
# Verify - server was removed
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once()
@@ -453,6 +587,7 @@ class TestMCPServiceUpdateMCPServer:
ap.tool_mgr.mcp_tool_loader.sessions = {}
ap.tool_mgr.mcp_tool_loader.host_mcp_server = AsyncMock()
ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks = []
+ ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=False)
old_server = _create_mock_mcp_server(name='Old Server', enable=False)
@@ -474,10 +609,10 @@ class TestMCPServiceUpdateMCPServer:
return_value={'uuid': 'test-uuid', 'name': 'Old Server', 'enable': True}
)
- service = MCPService(ap)
+ service = _service(ap)
# Execute - enable server
- await service.update_mcp_server('test-uuid', {'enable': True})
+ await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': True})
# Verify - server was loaded
ap.tool_mgr.mcp_tool_loader.host_mcp_server.assert_called_once()
@@ -493,6 +628,7 @@ class TestMCPServiceUpdateMCPServer:
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
ap.tool_mgr.mcp_tool_loader.host_mcp_server = AsyncMock()
ap.tool_mgr.mcp_tool_loader._hosted_mcp_tasks = []
+ ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=True)
old_server = _create_mock_mcp_server(name='Old Server', enable=True)
@@ -510,13 +646,13 @@ class TestMCPServiceUpdateMCPServer:
return_value={'uuid': 'test-uuid', 'name': 'Old Server', 'enable': True}
)
- service = MCPService(ap)
+ service = _service(ap)
# Execute - update enabled server (keep enabled, update extra_args)
- await service.update_mcp_server('test-uuid', {'enable': True, 'extra_args': {'new': 'args'}})
+ await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': True, 'extra_args': {'new': 'args'}})
# Verify - remove and reload
- ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with('Old Server')
+ ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with(_CONTEXT, 'Old Server')
ap.tool_mgr.mcp_tool_loader.host_mcp_server.assert_called_once()
async def test_update_mcp_server_no_tool_mgr(self):
@@ -541,15 +677,99 @@ class TestMCPServiceUpdateMCPServer:
return Mock() # Update
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
+ ap.persistence_mgr.serialize_model = Mock(
+ return_value={
+ 'uuid': 'test-uuid',
+ 'name': 'Server',
+ 'enable': True,
+ }
+ )
- service = MCPService(ap)
+ service = _service(ap)
# Execute - should not raise
- await service.update_mcp_server('test-uuid', {'name': 'New Name'})
+ await service.update_mcp_server(_CONTEXT, 'test-uuid', {'enable': False})
# Verify - persistence was called
assert ap.persistence_mgr.execute_async.call_count >= 2
+ async def test_update_restores_existing_masked_secrets_and_preserves_explicit_changes(self):
+ ap = SimpleNamespace()
+ ap.persistence_mgr = SimpleNamespace()
+ ap.tool_mgr = SimpleNamespace(mcp_tool_loader=None)
+ old_server = _create_mock_mcp_server(name='Server', enable=True)
+ old_data = {
+ 'uuid': 'test-uuid',
+ 'name': 'Server',
+ 'enable': True,
+ 'mode': 'streamable_http',
+ 'extra_args': {
+ 'headers': {
+ 'Authorization': 'Bearer original-secret',
+ 'X-API-Key': 'original-api-key',
+ 'Cookie': 'original-cookie',
+ }
+ },
+ }
+ captured_updates = []
+
+ async def mock_execute(statement):
+ if not captured_updates:
+ captured_updates.append(None)
+ return _create_mock_result(first_item=old_server)
+ captured_updates[0] = statement
+ return _create_mock_result()
+
+ ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
+ ap.persistence_mgr.serialize_model = Mock(return_value=old_data)
+ service = _service(ap)
+
+ await service.update_mcp_server(
+ _CONTEXT,
+ 'test-uuid',
+ {
+ 'extra_args': {
+ 'headers': {
+ 'Authorization': '***',
+ 'X-API-Key': 'replacement-api-key',
+ 'Cookie': '',
+ }
+ }
+ },
+ )
+
+ persisted = captured_updates[0].compile().params['extra_args']
+ assert persisted['headers'] == {
+ 'Authorization': 'Bearer original-secret',
+ 'X-API-Key': 'replacement-api-key',
+ 'Cookie': '',
+ }
+
+ async def test_update_rejects_masked_secret_without_existing_value(self):
+ ap = SimpleNamespace()
+ ap.persistence_mgr = SimpleNamespace()
+ ap.tool_mgr = SimpleNamespace(mcp_tool_loader=None)
+ old_server = _create_mock_mcp_server(name='Server', enable=True)
+ ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result(first_item=old_server))
+ ap.persistence_mgr.serialize_model = Mock(
+ return_value={
+ 'uuid': 'test-uuid',
+ 'name': 'Server',
+ 'enable': True,
+ 'extra_args': {'headers': {'Accept': 'application/json'}},
+ }
+ )
+ service = _service(ap)
+
+ with pytest.raises(ValueError, match='Masked MCP secret has no existing value'):
+ await service.update_mcp_server(
+ _CONTEXT,
+ 'test-uuid',
+ {'extra_args': {'headers': {'Authorization': '***'}}},
+ )
+
+ assert ap.persistence_mgr.execute_async.await_count == 1
+
class TestMCPServiceDeleteMCPServer:
"""Tests for delete_mcp_server method."""
@@ -563,6 +783,7 @@ class TestMCPServiceDeleteMCPServer:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.sessions = {'Server to Delete': Mock()}
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
+ ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=True)
server = _create_mock_mcp_server(name='Server to Delete')
@@ -576,14 +797,21 @@ class TestMCPServiceDeleteMCPServer:
return Mock() # Delete
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
+ ap.persistence_mgr.serialize_model = Mock(
+ return_value={
+ 'uuid': 'test-uuid',
+ 'name': 'Server to Delete',
+ 'enable': True,
+ }
+ )
- service = MCPService(ap)
+ service = _service(ap)
# Execute
- await service.delete_mcp_server('test-uuid')
+ await service.delete_mcp_server(_CONTEXT, 'test-uuid')
# Verify
- ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with('Server to Delete')
+ ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_called_once_with(_CONTEXT, 'Server to Delete')
ap.persistence_mgr.execute_async.assert_called()
async def test_delete_mcp_server_not_in_sessions(self):
@@ -595,6 +823,7 @@ class TestMCPServiceDeleteMCPServer:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.sessions = {} # Server not in sessions
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
+ ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=False)
server = _create_mock_mcp_server(name='Not in Sessions')
@@ -608,11 +837,18 @@ class TestMCPServiceDeleteMCPServer:
return Mock()
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
+ ap.persistence_mgr.serialize_model = Mock(
+ return_value={
+ 'uuid': 'test-uuid',
+ 'name': 'Not in Sessions',
+ 'enable': True,
+ }
+ )
- service = MCPService(ap)
+ service = _service(ap)
# Execute
- await service.delete_mcp_server('test-uuid')
+ await service.delete_mcp_server(_CONTEXT, 'test-uuid')
# Verify - remove not called (server not in sessions)
ap.tool_mgr.mcp_tool_loader.remove_mcp_server.assert_not_called()
@@ -626,6 +862,7 @@ class TestMCPServiceDeleteMCPServer:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.sessions = {}
ap.tool_mgr.mcp_tool_loader.remove_mcp_server = AsyncMock()
+ ap.tool_mgr.mcp_tool_loader.has_session = Mock(return_value=False)
# No server found
call_count = 0
@@ -639,13 +876,12 @@ class TestMCPServiceDeleteMCPServer:
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
- service = MCPService(ap)
+ service = _service(ap)
- # Execute - should not raise
- await service.delete_mcp_server('nonexistent-uuid')
+ with pytest.raises(WorkspaceNotFoundError, match='MCP server not found'):
+ await service.delete_mcp_server(_CONTEXT, 'nonexistent-uuid')
- # Verify - delete was called regardless
- ap.persistence_mgr.execute_async.assert_called()
+ assert ap.persistence_mgr.execute_async.await_count == 1
class TestMCPServiceTestMCPServer:
@@ -667,12 +903,18 @@ class TestMCPServiceTestMCPServer:
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=mock_session)
ap.task_mgr = SimpleNamespace()
- ap.task_mgr.create_user_task = Mock(return_value=SimpleNamespace(id=123))
- service = MCPService(ap)
+ service = _service(ap)
+ service._require_server = AsyncMock(return_value=(_CONTEXT, {'name': 'existing-server'}))
+
+ def create_user_task(coroutine, **_kwargs):
+ coroutine.close()
+ return SimpleNamespace(id=123)
+
+ ap.task_mgr.create_user_task = Mock(side_effect=create_user_task)
# Execute
- task_id = await service.test_mcp_server('existing-server', {})
+ task_id = await service.test_mcp_server(_CONTEXT, 'existing-server', {})
# Verify - returns task ID
assert task_id == 123
@@ -685,11 +927,12 @@ class TestMCPServiceTestMCPServer:
ap.tool_mgr.mcp_tool_loader = SimpleNamespace()
ap.tool_mgr.mcp_tool_loader.get_session = Mock(return_value=None)
- service = MCPService(ap)
+ service = _service(ap)
+ service._require_server = AsyncMock(side_effect=WorkspaceNotFoundError('MCP server not found'))
# Execute & Verify
- with pytest.raises(ValueError, match='Server not found'):
- await service.test_mcp_server('nonexistent-server', {})
+ with pytest.raises(WorkspaceNotFoundError, match='MCP server not found'):
+ await service.test_mcp_server(_CONTEXT, 'nonexistent-server', {})
async def test_test_mcp_server_new_server(self):
"""Tests new MCP server with underscore name."""
@@ -703,12 +946,17 @@ class TestMCPServiceTestMCPServer:
ap.tool_mgr.mcp_tool_loader.load_mcp_server = AsyncMock(return_value=mock_session)
ap.task_mgr = SimpleNamespace()
- ap.task_mgr.create_user_task = Mock(return_value=SimpleNamespace(id=456))
- service = MCPService(ap)
+ service = _service(ap)
+
+ def create_user_task(coroutine, **_kwargs):
+ coroutine.close()
+ return SimpleNamespace(id=456)
+
+ ap.task_mgr.create_user_task = Mock(side_effect=create_user_task)
# Execute with '_' name (new server)
- task_id = await service.test_mcp_server('_', {'name': 'New Server'})
+ task_id = await service.test_mcp_server(_CONTEXT, '_', {'name': 'New Server'})
# Verify - load_mcp_server called
ap.tool_mgr.mcp_tool_loader.load_mcp_server.assert_called_once()
diff --git a/tests/unit_tests/api/service/test_model_service.py b/tests/unit_tests/api/service/test_model_service.py
index 42129ed3b..7e6af718a 100644
--- a/tests/unit_tests/api/service/test_model_service.py
+++ b/tests/unit_tests/api/service/test_model_service.py
@@ -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,
+ }
diff --git a/tests/unit_tests/api/service/test_monitoring_tenancy.py b/tests/unit_tests/api/service/test_monitoring_tenancy.py
new file mode 100644
index 000000000..38f706cff
--- /dev/null
+++ b/tests/unit_tests/api/service/test_monitoring_tenancy.py
@@ -0,0 +1,154 @@
+from __future__ import annotations
+
+import datetime
+from types import SimpleNamespace
+
+import pytest
+import sqlalchemy
+from sqlalchemy.ext.asyncio import create_async_engine
+
+from langbot.pkg.api.http.authz import WorkspaceRequiredError
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.api.http.service.monitoring import MonitoringService
+from langbot.pkg.entity.persistence.base import Base
+from langbot.pkg.entity.persistence.workspace import Workspace
+
+
+pytestmark = pytest.mark.asyncio
+
+WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
+WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
+
+
+def _context(workspace_uuid: str) -> ExecutionContext:
+ return ExecutionContext(
+ instance_uuid='instance',
+ workspace_uuid=workspace_uuid,
+ placement_generation=3,
+ bot_uuid='same-bot',
+ pipeline_uuid='same-pipeline',
+ )
+
+
+class _PersistenceManager:
+ def __init__(self, engine):
+ self.engine = engine
+
+ async def execute_async(self, *args, **kwargs):
+ async with self.engine.connect() as connection:
+ result = await connection.execute(*args, **kwargs)
+ await connection.commit()
+ return result
+
+ def get_db_engine(self):
+ return self.engine
+
+ @staticmethod
+ def serialize_model(model, data, masked_columns=None):
+ return {
+ column.name: (
+ getattr(data, column.name).isoformat()
+ if isinstance(getattr(data, column.name), datetime.datetime)
+ else getattr(data, column.name)
+ )
+ for column in model.__table__.columns
+ if column.name not in (masked_columns or [])
+ }
+
+
+@pytest.fixture
+async def service(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "monitoring.db"}')
+ async with engine.begin() as connection:
+ await connection.run_sync(Base.metadata.create_all)
+ await connection.execute(
+ sqlalchemy.insert(Workspace),
+ [
+ {
+ 'uuid': WORKSPACE_A,
+ 'instance_uuid': 'instance',
+ 'name': 'A',
+ 'slug': 'a',
+ 'source': 'cloud_projection',
+ },
+ {
+ 'uuid': WORKSPACE_B,
+ 'instance_uuid': 'instance',
+ 'name': 'B',
+ 'slug': 'b',
+ 'source': 'cloud_projection',
+ },
+ ],
+ )
+ application = SimpleNamespace(
+ persistence_mgr=_PersistenceManager(engine),
+ instance_config=SimpleNamespace(data={'database': {'use': 'sqlite'}}),
+ )
+ yield MonitoringService(application)
+ await engine.dispose()
+
+
+async def _record_message(service, context, content):
+ return await service.record_message(
+ context,
+ bot_id='same-bot',
+ bot_name='Same Bot',
+ pipeline_id='same-pipeline',
+ pipeline_name='Same Pipeline',
+ message_content=content,
+ session_id='same-session',
+ )
+
+
+async def test_monitoring_write_without_execution_context_fails_closed(service):
+ with pytest.raises(WorkspaceRequiredError):
+ await _record_message(service, None, 'unscoped')
+
+
+async def test_same_session_and_resource_ids_do_not_collide(service):
+ context_a = _context(WORKSPACE_A)
+ context_b = _context(WORKSPACE_B)
+ message_a = await _record_message(service, context_a, 'tenant-a')
+ message_b = await _record_message(service, context_b, 'tenant-b')
+ await service.record_session_start(
+ context_a,
+ session_id='same-session',
+ bot_id='same-bot',
+ bot_name='Same Bot',
+ pipeline_id='same-pipeline',
+ pipeline_name='Same Pipeline',
+ )
+ await service.record_session_start(
+ context_b,
+ session_id='same-session',
+ bot_id='same-bot',
+ bot_name='Same Bot',
+ pipeline_id='same-pipeline',
+ pipeline_name='Same Pipeline',
+ )
+
+ messages_a, total_a = await service.get_messages(context_a)
+ messages_b, total_b = await service.get_messages(context_b)
+ assert total_a == total_b == 1
+ assert messages_a[0]['message_content'] == 'tenant-a'
+ assert messages_b[0]['message_content'] == 'tenant-b'
+ assert (await service.get_message_details(context_b, message_a))['found'] is False
+ assert (await service.get_message_details(context_a, message_b))['found'] is False
+
+
+async def test_feedback_upsert_and_cancel_are_workspace_scoped(service):
+ context_a = _context(WORKSPACE_A)
+ context_b = _context(WORKSPACE_B)
+ await service.record_feedback(context_a, feedback_id='same-feedback', feedback_type=1)
+ await service.record_feedback(context_b, feedback_id='same-feedback', feedback_type=2)
+
+ stats_a = await service.get_feedback_stats(context_a)
+ stats_b = await service.get_feedback_stats(context_b)
+ assert stats_a['total_likes'] == 1
+ assert stats_a['total_dislikes'] == 0
+ assert stats_b['total_likes'] == 0
+ assert stats_b['total_dislikes'] == 1
+
+ await service.record_feedback(context_a, feedback_id='same-feedback', feedback_type=3)
+ assert (await service.get_feedback_stats(context_a))['total_feedback'] == 0
+ assert (await service.get_feedback_stats(context_b))['total_feedback'] == 1
diff --git a/tests/unit_tests/api/service/test_pipeline_service.py b/tests/unit_tests/api/service/test_pipeline_service.py
index fade30372..a4ef851db 100644
--- a/tests/unit_tests/api/service/test_pipeline_service.py
+++ b/tests/unit_tests/api/service/test_pipeline_service.py
@@ -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."""
diff --git a/tests/unit_tests/api/service/test_provider_service.py b/tests/unit_tests/api/service/test_provider_service.py
index 8b308af8d..81f5ee6c1 100644
--- a/tests/unit_tests/api/service/test_provider_service.py
+++ b/tests/unit_tests/api/service/test_provider_service.py
@@ -19,10 +19,13 @@ from types import SimpleNamespace
from langbot.pkg.api.http.service.provider import ModelProviderService
from langbot.pkg.entity.persistence.model import ModelProvider, LLMModel, EmbeddingModel, RerankModel
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
pytestmark = pytest.mark.asyncio
+WORKSPACE_UUID = 'workspace-a'
+
def _create_mock_provider(
provider_uuid: str = 'test-provider-uuid',
@@ -86,7 +89,9 @@ class TestModelProviderServiceGetProviders:
service = ModelProviderService(ap)
# Execute
- result = await service.get_providers()
+ result = await service.get_providers(
+ WORKSPACE_UUID,
+ )
# Verify
assert result == []
@@ -115,7 +120,9 @@ class TestModelProviderServiceGetProviders:
service = ModelProviderService(ap)
# Execute
- result = await service.get_providers()
+ result = await service.get_providers(
+ WORKSPACE_UUID,
+ )
# Verify
assert len(result) == 2
@@ -143,7 +150,10 @@ class TestModelProviderServiceGetProviders:
service = ModelProviderService(ap)
# Execute
- result = await service.get_providers()
+ result = await service.get_providers(
+ WORKSPACE_UUID,
+ include_secret=True,
+ )
# Verify - api_keys should be parsed from string
assert result[0]['api_keys'] == ['key1', 'key2']
@@ -169,11 +179,41 @@ class TestModelProviderServiceGetProviders:
service = ModelProviderService(ap)
# Execute
- result = await service.get_providers()
+ result = await service.get_providers(
+ WORKSPACE_UUID,
+ )
# Verify - invalid JSON returns empty list
assert result[0]['api_keys'] == []
+ async def test_get_providers_masks_api_keys_for_resource_view(self):
+ ap = SimpleNamespace()
+ provider = _create_mock_provider(
+ api_keys=['first', 'second'],
+ base_url=(
+ 'https://provider-user:provider-password@api.provider.invalid/v1?access_token=url-secret®ion=sg'
+ ),
+ )
+ ap.persistence_mgr = SimpleNamespace(
+ execute_async=AsyncMock(return_value=_create_mock_result([provider])),
+ serialize_model=Mock(
+ return_value={
+ 'uuid': provider.uuid,
+ 'name': provider.name,
+ 'base_url': provider.base_url,
+ 'api_keys': provider.api_keys,
+ }
+ ),
+ )
+
+ result = await ModelProviderService(ap).get_providers(
+ WORKSPACE_UUID,
+ include_secret=False,
+ )
+
+ assert result[0]['api_keys'] == ['***', '***']
+ assert result[0]['base_url'] == ('https://***@api.provider.invalid/v1?access_token=***®ion=sg')
+
class TestModelProviderServiceGetProvider:
"""Tests for get_provider method."""
@@ -199,7 +239,7 @@ class TestModelProviderServiceGetProvider:
service = ModelProviderService(ap)
# Execute
- result = await service.get_provider('found-uuid')
+ result = await service.get_provider(WORKSPACE_UUID, 'found-uuid')
# Verify
assert result is not None
@@ -217,7 +257,7 @@ class TestModelProviderServiceGetProvider:
service = ModelProviderService(ap)
# Execute
- result = await service.get_provider('nonexistent-uuid')
+ result = await service.get_provider(WORKSPACE_UUID, 'nonexistent-uuid')
# Verify
assert result is None
@@ -239,6 +279,7 @@ class TestModelProviderServiceCreateProvider:
runtime_provider.provider_entity = Mock()
runtime_provider.provider_entity.uuid = 'generated-uuid'
ap.model_mgr.load_provider = AsyncMock(return_value=runtime_provider)
+ ap.model_mgr.cache_provider = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock()
@@ -246,12 +287,13 @@ class TestModelProviderServiceCreateProvider:
# Execute
provider_uuid = await service.create_provider(
+ WORKSPACE_UUID,
{
'name': 'New Provider',
'requester': 'openai',
'base_url': 'https://api.openai.com',
'api_keys': ['key'],
- }
+ },
)
# Verify - UUID is generated
@@ -270,6 +312,7 @@ class TestModelProviderServiceCreateProvider:
runtime_provider.provider_entity = Mock()
runtime_provider.provider_entity.uuid = 'runtime-uuid'
ap.model_mgr.load_provider = AsyncMock(return_value=runtime_provider)
+ ap.model_mgr.cache_provider = AsyncMock()
ap.persistence_mgr.execute_async = AsyncMock()
@@ -277,12 +320,13 @@ class TestModelProviderServiceCreateProvider:
# Execute
result_uuid = await service.create_provider(
+ WORKSPACE_UUID,
{
'name': 'Runtime Provider',
'requester': 'openai',
'base_url': 'https://api.openai.com',
'api_keys': ['key'],
- }
+ },
)
# Verify - provider added to runtime dict and UUID generated
@@ -307,6 +351,7 @@ class TestModelProviderServiceUpdateProvider:
# Execute
await service.update_provider(
+ WORKSPACE_UUID,
'existing-uuid',
{
'uuid': 'should-be-removed', # Will be removed
@@ -315,7 +360,7 @@ class TestModelProviderServiceUpdateProvider:
)
# Verify - reload called
- ap.model_mgr.reload_provider.assert_called_once_with('existing-uuid')
+ ap.model_mgr.reload_provider.assert_called_once_with(WORKSPACE_UUID, 'existing-uuid')
async def test_update_provider_reloads_runtime(self):
"""Reloads provider in runtime after update."""
@@ -330,7 +375,7 @@ class TestModelProviderServiceUpdateProvider:
service = ModelProviderService(ap)
# Execute
- await service.update_provider('update-uuid', {'name': 'New Name'})
+ await service.update_provider(WORKSPACE_UUID, 'update-uuid', {'name': 'New Name'})
# Verify
ap.model_mgr.reload_provider.assert_called_once()
@@ -354,7 +399,7 @@ class TestModelProviderServiceDeleteProvider:
# Execute & Verify
with pytest.raises(ValueError, match='Cannot delete provider: LLM models'):
- await service.delete_provider('provider-with-llm')
+ await service.delete_provider(WORKSPACE_UUID, 'provider-with-llm')
async def test_delete_provider_with_embedding_models_raises_error(self):
"""Raises ValueError when Embedding models reference provider."""
@@ -387,7 +432,7 @@ class TestModelProviderServiceDeleteProvider:
# Execute & Verify - should raise embedding error (LLM check passes, embedding check fails)
with pytest.raises(ValueError, match='Cannot delete provider: Embedding models'):
- await service.delete_provider('provider-with-embedding')
+ await service.delete_provider(WORKSPACE_UUID, 'provider-with-embedding')
async def test_delete_provider_with_rerank_models_raises_error(self):
"""Raises ValueError when Rerank models reference provider."""
@@ -420,7 +465,7 @@ class TestModelProviderServiceDeleteProvider:
# Execute & Verify - should raise rerank error (LLM and embedding checks pass, rerank check fails)
with pytest.raises(ValueError, match='Cannot delete provider: Rerank models'):
- await service.delete_provider('provider-with-rerank')
+ await service.delete_provider(WORKSPACE_UUID, 'provider-with-rerank')
async def test_delete_provider_no_models_success(self):
"""Deletes provider when no models reference it."""
@@ -439,10 +484,10 @@ class TestModelProviderServiceDeleteProvider:
service = ModelProviderService(ap)
# Execute
- await service.delete_provider('provider-no-models')
+ await service.delete_provider(WORKSPACE_UUID, 'provider-no-models')
# Verify - delete and remove called
- ap.model_mgr.remove_provider.assert_called_once_with('provider-no-models')
+ ap.model_mgr.remove_provider.assert_called_once_with(WORKSPACE_UUID, 'provider-no-models')
class TestModelProviderServiceGetProviderModelCounts:
@@ -476,9 +521,10 @@ class TestModelProviderServiceGetProviderModelCounts:
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
service = ModelProviderService(ap)
+ service.get_provider = AsyncMock(return_value={'uuid': 'provider-uuid'})
# Execute
- result = await service.get_provider_model_counts('provider-uuid')
+ result = await service.get_provider_model_counts(WORKSPACE_UUID, 'provider-uuid')
# Verify
assert result['llm_count'] == 3
@@ -497,9 +543,10 @@ class TestModelProviderServiceGetProviderModelCounts:
ap.persistence_mgr.execute_async = AsyncMock(return_value=zero_result)
service = ModelProviderService(ap)
+ service.get_provider = AsyncMock(return_value={'uuid': 'empty-provider'})
# Execute
- result = await service.get_provider_model_counts('empty-provider')
+ result = await service.get_provider_model_counts(WORKSPACE_UUID, 'empty-provider')
# Verify
assert result['llm_count'] == 0
@@ -530,6 +577,7 @@ class TestModelProviderServiceFindOrCreateProvider:
# Execute
result = await service.find_or_create_provider(
+ WORKSPACE_UUID,
requester='openai',
base_url='https://api.openai.com',
api_keys=['key1', 'key2'], # Same keys (sorted)
@@ -558,6 +606,7 @@ class TestModelProviderServiceFindOrCreateProvider:
# Execute with reversed key order
result = await service.find_or_create_provider(
+ WORKSPACE_UUID,
requester='openai',
base_url='https://api.openai.com',
api_keys=['key2', 'key1'], # Different order, should still match
@@ -578,6 +627,7 @@ class TestModelProviderServiceFindOrCreateProvider:
runtime_provider.provider_entity = Mock()
runtime_provider.provider_entity.uuid = None # Will be set by uuid.uuid4()
ap.model_mgr.load_provider = AsyncMock(return_value=runtime_provider)
+ ap.model_mgr.cache_provider = AsyncMock()
# Mock no existing providers
mock_result = _create_mock_result([])
@@ -587,6 +637,7 @@ class TestModelProviderServiceFindOrCreateProvider:
# Execute
result = await service.find_or_create_provider(
+ WORKSPACE_UUID,
requester='new-requester',
base_url='https://new.api.com',
api_keys=['new-key'],
@@ -610,6 +661,7 @@ class TestModelProviderServiceFindOrCreateProvider:
runtime_provider.provider_entity = Mock()
runtime_provider.provider_entity.uuid = 'parsed-url-uuid'
ap.model_mgr.load_provider = AsyncMock(return_value=runtime_provider)
+ ap.model_mgr.cache_provider = AsyncMock()
mock_result = _create_mock_result([])
ap.persistence_mgr.execute_async = AsyncMock(return_value=mock_result)
@@ -618,6 +670,7 @@ class TestModelProviderServiceFindOrCreateProvider:
# Execute
result_uuid = await service.find_or_create_provider(
+ WORKSPACE_UUID,
requester='custom',
base_url='https://api.example.com/v1',
api_keys=['key'],
@@ -644,17 +697,20 @@ class TestModelProviderServiceUpdateSpaceModelProviderApiKeys:
service = ModelProviderService(ap)
# Execute
- await service.update_space_model_provider_api_keys('space-api-key')
+ await service.update_space_model_provider_api_keys(WORKSPACE_UUID, 'space-api-key')
# Verify - update and reload called for Space provider UUID
- ap.model_mgr.reload_provider.assert_called_once_with('00000000-0000-0000-0000-000000000000')
+ ap.model_mgr.reload_provider.assert_called_once_with(
+ WORKSPACE_UUID,
+ '00000000-0000-0000-0000-000000000000',
+ )
class TestModelProviderServiceScanProviderModels:
"""Tests for scan_provider_models method."""
async def test_scan_provider_not_found_raises_error(self):
- """Raises ValueError when provider not found."""
+ """Raises a non-enumerating not-found error when provider is outside the Workspace."""
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
@@ -665,8 +721,8 @@ class TestModelProviderServiceScanProviderModels:
service = ModelProviderService(ap)
# Execute & Verify
- with pytest.raises(ValueError, match='provider not found'):
- await service.scan_provider_models('nonexistent-uuid')
+ with pytest.raises(WorkspaceNotFoundError, match='Provider not found'):
+ await service.scan_provider_models(WORKSPACE_UUID, 'nonexistent-uuid')
async def test_scan_provider_returns_models_list(self):
"""Returns scanned models list."""
@@ -718,7 +774,7 @@ class TestModelProviderServiceScanProviderModels:
service = ModelProviderService(ap)
# Execute
- result = await service.scan_provider_models('scan-uuid')
+ result = await service.scan_provider_models(WORKSPACE_UUID, 'scan-uuid')
# Verify
assert 'models' in result
@@ -771,7 +827,7 @@ class TestModelProviderServiceScanProviderModels:
service = ModelProviderService(ap)
# Execute - filter for LLM only
- result = await service.scan_provider_models('filter-uuid', model_type='llm')
+ result = await service.scan_provider_models(WORKSPACE_UUID, 'filter-uuid', model_type='llm')
# Verify - only LLM models returned
assert len(result['models']) == 1
@@ -810,7 +866,7 @@ class TestModelProviderServiceScanProviderModels:
# Execute & Verify
with pytest.raises(ValueError, match='current provider does not support model scanning'):
- await service.scan_provider_models('no-scan-uuid')
+ await service.scan_provider_models(WORKSPACE_UUID, 'no-scan-uuid')
async def test_scan_provider_marks_already_added_models(self):
"""Marks models that are already added."""
@@ -860,7 +916,7 @@ class TestModelProviderServiceScanProviderModels:
service = ModelProviderService(ap)
# Execute
- result = await service.scan_provider_models('already-added-uuid')
+ result = await service.scan_provider_models(WORKSPACE_UUID, 'already-added-uuid')
# Verify - existing model marked as already_added
existing_model = next(m for m in result['models'] if m['name'] == 'Existing Model')
@@ -868,3 +924,46 @@ class TestModelProviderServiceScanProviderModels:
new_model = next(m for m in result['models'] if m['name'] == 'New Model')
assert new_model['already_added'] is False
+
+
+class TestProviderSecretRoundtrip:
+ async def test_masked_api_keys_update_preserves_existing_values(self):
+ write_result = Mock(rowcount=1)
+ ap = SimpleNamespace(
+ persistence_mgr=SimpleNamespace(execute_async=AsyncMock(return_value=write_result)),
+ model_mgr=SimpleNamespace(reload_provider=AsyncMock()),
+ )
+ service = ModelProviderService(ap)
+ service.get_provider = AsyncMock(
+ return_value={
+ 'uuid': 'provider-secret',
+ 'api_keys': ['first-secret', 'second-secret'],
+ }
+ )
+
+ await service.update_provider(
+ WORKSPACE_UUID,
+ 'provider-secret',
+ {'name': 'Updated', 'api_keys': ['***', 'replacement-secret']},
+ )
+
+ statement = ap.persistence_mgr.execute_async.await_args.args[0]
+ stored_api_keys = next(value.value for column, value in statement._values.items() if column.key == 'api_keys')
+ assert stored_api_keys == ['first-secret', 'replacement-secret']
+
+ async def test_extra_masked_api_key_is_rejected(self):
+ ap = SimpleNamespace(
+ persistence_mgr=SimpleNamespace(execute_async=AsyncMock()),
+ model_mgr=SimpleNamespace(reload_provider=AsyncMock()),
+ )
+ service = ModelProviderService(ap)
+ service.get_provider = AsyncMock(return_value={'uuid': 'provider-secret', 'api_keys': ['only-secret']})
+
+ with pytest.raises(ValueError, match='no existing value'):
+ await service.update_provider(
+ WORKSPACE_UUID,
+ 'provider-secret',
+ {'api_keys': ['***', '***']},
+ )
+
+ ap.persistence_mgr.execute_async.assert_not_awaited()
diff --git a/tests/unit_tests/api/service/test_secret_redaction.py b/tests/unit_tests/api/service/test_secret_redaction.py
new file mode 100644
index 000000000..432613a06
--- /dev/null
+++ b/tests/unit_tests/api/service/test_secret_redaction.py
@@ -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']},
+ )
diff --git a/tests/unit_tests/api/service/test_space_service.py b/tests/unit_tests/api/service/test_space_service.py
index f48b18937..b9a833bd3 100644
--- a/tests/unit_tests/api/service/test_space_service.py
+++ b/tests/unit_tests/api/service/test_space_service.py
@@ -13,6 +13,8 @@ Source: src/langbot/pkg/api/http/service/space.py
from __future__ import annotations
+from urllib.parse import parse_qs, urlsplit
+
import pytest
from unittest.mock import AsyncMock, Mock, patch, MagicMock
from types import SimpleNamespace
@@ -73,7 +75,7 @@ class TestSpaceServiceGetOAuthAuthorizeUrl:
result = service.get_oauth_authorize_url('http://localhost/callback')
# Verify
- assert 'redirect_uri=http://localhost/callback' in result
+ assert parse_qs(urlsplit(result).query)['redirect_uri'] == ['http://localhost/callback']
assert 'https://space.langbot.app/auth/authorize' in result
def test_get_oauth_authorize_url_with_state(self):
@@ -93,8 +95,9 @@ class TestSpaceServiceGetOAuthAuthorizeUrl:
result = service.get_oauth_authorize_url('http://localhost/callback', state='random_state')
# Verify
- assert 'redirect_uri=http://localhost/callback' in result
- assert 'state=random_state' in result
+ params = parse_qs(urlsplit(result).query)
+ assert params['redirect_uri'] == ['http://localhost/callback']
+ assert params['state'] == ['random_state']
def test_get_oauth_authorize_url_default_config(self):
"""Uses default OAuth URL when config not set."""
diff --git a/tests/unit_tests/api/service/test_user_service.py b/tests/unit_tests/api/service/test_user_service.py
index c5d37f167..8a580a41d 100644
--- a/tests/unit_tests/api/service/test_user_service.py
+++ b/tests/unit_tests/api/service/test_user_service.py
@@ -14,17 +14,60 @@ Source: src/langbot/pkg/api/http/service/user.py
from __future__ import annotations
import pytest
+import jwt
+import datetime
from unittest.mock import AsyncMock, Mock
from types import SimpleNamespace
from langbot.pkg.api.http.service.user import UserService
-from langbot.pkg.entity.persistence.user import User
+from langbot.pkg.entity.persistence.user import AccountStatus, User
from langbot.pkg.entity.errors.account import AccountEmailMismatchError
pytestmark = pytest.mark.asyncio
+class TestSpaceOAuthState:
+ async def test_login_state_is_opaque_single_use(self):
+ service = UserService(SimpleNamespace())
+
+ state = await service.issue_space_oauth_state('login')
+
+ assert state.count('.') == 0
+ assert await service.consume_space_oauth_state(state, 'login') is None
+ with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
+ await service.consume_space_oauth_state(state, 'login')
+
+ async def test_bind_state_resolves_only_bound_active_account(self):
+ service = UserService(SimpleNamespace())
+ account = SimpleNamespace(uuid='account-a', status=AccountStatus.ACTIVE.value)
+ service.get_user_by_uuid = AsyncMock(return_value=account)
+
+ state = await service.issue_space_oauth_state('bind', account_uuid='account-a')
+
+ assert await service.consume_space_oauth_state(state, 'bind') is account
+ service.get_user_by_uuid.assert_awaited_once_with('account-a')
+
+ async def test_state_purpose_mismatch_is_rejected_and_consumed(self):
+ service = UserService(SimpleNamespace())
+ state = await service.issue_space_oauth_state('login')
+
+ with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
+ await service.consume_space_oauth_state(state, 'bind')
+ with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
+ await service.consume_space_oauth_state(state, 'login')
+
+ async def test_expired_state_is_rejected(self):
+ service = UserService(SimpleNamespace())
+ state = await service.issue_space_oauth_state('login')
+ digest = service._space_oauth_state_digest(state)
+ purpose, account_uuid, _ = service._space_oauth_states[digest]
+ service._space_oauth_states[digest] = (purpose, account_uuid, 0)
+
+ with pytest.raises(ValueError, match='Invalid or expired OAuth state'):
+ await service.consume_space_oauth_state(state, 'login')
+
+
def _create_mock_user(
email: str = 'test@example.com',
password: str = 'hashed_password',
@@ -309,6 +352,50 @@ class TestUserServiceVerifyJwtToken:
with pytest.raises(Exception): # jwt.DecodeError or similar
await service.verify_jwt_token('invalid.token.here')
+ async def test_verify_jwt_token_rejects_foreign_audience(self):
+ ap = SimpleNamespace()
+ ap.instance_config = SimpleNamespace()
+ ap.instance_config.data = {'system': {'jwt': {'secret': 'test_secret', 'expire': 3600}}}
+ service = UserService(ap)
+ token = jwt.encode(
+ {
+ 'user': 'verify@example.com',
+ 'iss': 'langbot-core',
+ 'aud': 'langbot-instance:another-instance',
+ 'exp': datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(hours=1),
+ },
+ 'test_secret',
+ algorithm='HS256',
+ )
+
+ with pytest.raises(jwt.InvalidAudienceError):
+ await service.verify_jwt_token(token)
+
+ async def test_verify_jwt_token_accepts_legacy_community_token_only_in_oss(self):
+ ap = SimpleNamespace()
+ ap.instance_config = SimpleNamespace()
+ ap.instance_config.data = {'system': {'jwt': {'secret': 'test_secret', 'expire': 3600}}}
+ ap.workspace_service = SimpleNamespace(
+ instance_uuid='instance-a',
+ policy=SimpleNamespace(multi_workspace_enabled=False),
+ )
+ service = UserService(ap)
+ legacy_token = jwt.encode(
+ {
+ 'user': 'legacy@example.com',
+ 'iss': 'LangBot-community',
+ 'exp': datetime.datetime.now(datetime.timezone.utc) + datetime.timedelta(hours=1),
+ },
+ 'test_secret',
+ algorithm='HS256',
+ )
+
+ assert await service.verify_jwt_token(legacy_token) == 'legacy@example.com'
+
+ ap.workspace_service.policy.multi_workspace_enabled = True
+ with pytest.raises(jwt.MissingRequiredClaimError):
+ await service.verify_jwt_token(legacy_token)
+
class TestUserServiceResetPassword:
"""Tests for reset_password method."""
@@ -548,6 +635,44 @@ class TestUserServiceCreateOrUpdateSpaceUser:
expires_in=3600,
)
+ async def test_unknown_space_subject_cannot_claim_existing_account_by_email(self):
+ """An OAuth login collision requires the explicit account-bound bind flow."""
+ existing_user = _create_mock_user(
+ email='owner@example.com',
+ account_type='local',
+ space_account_uuid=None,
+ )
+ ap = SimpleNamespace(
+ persistence_mgr=SimpleNamespace(execute_async=AsyncMock()),
+ provider_service=SimpleNamespace(update_space_model_provider_api_keys=AsyncMock()),
+ space_service=SimpleNamespace(
+ get_user_info_raw=AsyncMock(
+ return_value={
+ 'account': {
+ 'uuid': 'attacker-space-subject',
+ 'email': 'owner@example.com',
+ },
+ 'api_key': 'attacker-api-key',
+ }
+ )
+ ),
+ )
+ service = UserService(ap)
+ service.get_user_by_space_account_uuid = AsyncMock(return_value=None)
+ service.get_user_by_email = AsyncMock(return_value=existing_user)
+ service.generate_jwt_token = AsyncMock(return_value='must-not-be-issued')
+
+ with pytest.raises(AccountEmailMismatchError):
+ await service.authenticate_space_user(
+ 'attacker-access-token',
+ 'attacker-refresh-token',
+ 3600,
+ )
+
+ ap.persistence_mgr.execute_async.assert_not_awaited()
+ ap.provider_service.update_space_model_provider_api_keys.assert_not_awaited()
+ service.generate_jwt_token.assert_not_awaited()
+
async def test_create_or_update_space_user_no_expiry(self):
"""Creates Space user without token expiry."""
# Setup
diff --git a/tests/unit_tests/api/service/test_webhook_service.py b/tests/unit_tests/api/service/test_webhook_service.py
index 7a5a075ef..0948c123b 100644
--- a/tests/unit_tests/api/service/test_webhook_service.py
+++ b/tests/unit_tests/api/service/test_webhook_service.py
@@ -14,15 +14,23 @@ Source: src/langbot/pkg/api/http/service/webhook.py
from __future__ import annotations
+import datetime
+
import pytest
+import sqlalchemy
+from sqlalchemy.ext.asyncio import create_async_engine
from unittest.mock import AsyncMock, Mock
from types import SimpleNamespace
+from langbot.pkg.api.http.authz import WorkspaceRequiredError
from langbot.pkg.api.http.service.webhook import WebhookService
+from langbot.pkg.entity.persistence.base import Base
from langbot.pkg.entity.persistence.webhook import Webhook
+from langbot.pkg.entity.persistence.workspace import Workspace
pytestmark = pytest.mark.asyncio
+WORKSPACE_UUID = 'workspace-a'
def _create_mock_webhook(
@@ -47,6 +55,14 @@ def _create_mock_result(items: list = None, first_item=None):
result = Mock()
result.all = Mock(return_value=items or [])
result.first = Mock(return_value=first_item)
+ result.rowcount = 1
+ return result
+
+
+def _create_write_result(rowcount: int = 1, inserted_id: int = 1):
+ result = Mock()
+ result.rowcount = rowcount
+ result.inserted_primary_key = [inserted_id]
return result
@@ -71,7 +87,7 @@ class TestWebhookServiceGetWebhooks:
service = WebhookService(ap)
# Execute
- result = await service.get_webhooks()
+ result = await service.get_webhooks(WORKSPACE_UUID)
# Verify
assert result == []
@@ -100,7 +116,7 @@ class TestWebhookServiceGetWebhooks:
service = WebhookService(ap)
# Execute
- result = await service.get_webhooks()
+ result = await service.get_webhooks(WORKSPACE_UUID)
# Verify
assert len(result) == 2
@@ -119,6 +135,7 @@ class TestWebhookServiceCreateWebhook:
# Mock insert result
insert_result = Mock()
+ insert_result.inserted_primary_key = [1]
# Mock select result for retrieving created webhook
created_webhook = _create_mock_webhook(
@@ -155,6 +172,7 @@ class TestWebhookServiceCreateWebhook:
# Execute
result = await service.create_webhook(
+ WORKSPACE_UUID,
name='New Webhook',
url='http://new.example.com/webhook',
description='New Description',
@@ -187,7 +205,7 @@ class TestWebhookServiceCreateWebhook:
nonlocal call_count
call_count += 1
if call_count == 1:
- return Mock() # Insert
+ return _create_write_result() # Insert
return _create_mock_result(first_item=created_webhook)
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
@@ -204,7 +222,11 @@ class TestWebhookServiceCreateWebhook:
service = WebhookService(ap)
# Execute - only name and url required
- result = await service.create_webhook(name='Minimal Webhook', url='http://minimal.example.com')
+ result = await service.create_webhook(
+ WORKSPACE_UUID,
+ name='Minimal Webhook',
+ url='http://minimal.example.com',
+ )
# Verify defaults
assert result['description'] == ''
@@ -224,7 +246,7 @@ class TestWebhookServiceCreateWebhook:
nonlocal call_count
call_count += 1
if call_count == 1:
- return Mock()
+ return _create_write_result()
return _create_mock_result(first_item=created_webhook)
ap.persistence_mgr.execute_async = AsyncMock(side_effect=mock_execute)
@@ -233,7 +255,12 @@ class TestWebhookServiceCreateWebhook:
service = WebhookService(ap)
# Execute
- result = await service.create_webhook(name='Disabled', url='http://disabled.com', enabled=False)
+ result = await service.create_webhook(
+ WORKSPACE_UUID,
+ name='Disabled',
+ url='http://disabled.com',
+ enabled=False,
+ )
# Verify
assert result['enabled'] is False
@@ -262,7 +289,7 @@ class TestWebhookServiceGetWebhook:
service = WebhookService(ap)
# Execute
- result = await service.get_webhook(1)
+ result = await service.get_webhook(WORKSPACE_UUID, 1)
# Verify
assert result is not None
@@ -281,7 +308,7 @@ class TestWebhookServiceGetWebhook:
service = WebhookService(ap)
# Execute
- result = await service.get_webhook(999)
+ result = await service.get_webhook(WORKSPACE_UUID, 999)
# Verify
assert result is None
@@ -298,7 +325,7 @@ class TestWebhookServiceGetWebhook:
service = WebhookService(ap)
# Execute
- result = await service.get_webhook(0)
+ result = await service.get_webhook(WORKSPACE_UUID, 0)
# Verify - should return None (no webhook with ID 0)
assert result is None
@@ -312,12 +339,12 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
- await service.update_webhook(1, name='Updated Name')
+ await service.update_webhook(WORKSPACE_UUID, 1, name='Updated Name')
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
@@ -327,12 +354,12 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
- await service.update_webhook(1, url='http://updated.example.com')
+ await service.update_webhook(WORKSPACE_UUID, 1, url='http://updated.example.com')
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
@@ -342,12 +369,12 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
- await service.update_webhook(1, description='Updated description')
+ await service.update_webhook(WORKSPACE_UUID, 1, description='Updated description')
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
@@ -357,12 +384,12 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
- await service.update_webhook(1, enabled=False)
+ await service.update_webhook(WORKSPACE_UUID, 1, enabled=False)
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
@@ -372,12 +399,13 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
await service.update_webhook(
+ WORKSPACE_UUID,
1,
name='All Updated',
url='http://all.updated.com',
@@ -393,15 +421,17 @@ class TestWebhookServiceUpdateWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ existing = _create_mock_webhook(webhook_id=1)
+ ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_mock_result(first_item=existing))
+ ap.persistence_mgr.serialize_model = Mock(return_value={'id': 1})
service = WebhookService(ap)
# Execute - no update parameters
- await service.update_webhook(1)
+ await service.update_webhook(WORKSPACE_UUID, 1)
- # Verify - no execute call since no update_data
- ap.persistence_mgr.execute_async.assert_not_called()
+ # No write is issued; one scoped existence lookup is performed.
+ ap.persistence_mgr.execute_async.assert_called_once()
class TestWebhookServiceDeleteWebhook:
@@ -412,12 +442,12 @@ class TestWebhookServiceDeleteWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result())
service = WebhookService(ap)
# Execute
- await service.delete_webhook(1)
+ await service.delete_webhook(WORKSPACE_UUID, 1)
# Verify
ap.persistence_mgr.execute_async.assert_called_once()
@@ -427,12 +457,12 @@ class TestWebhookServiceDeleteWebhook:
# Setup
ap = SimpleNamespace()
ap.persistence_mgr = SimpleNamespace()
- ap.persistence_mgr.execute_async = AsyncMock()
+ ap.persistence_mgr.execute_async = AsyncMock(return_value=_create_write_result(rowcount=0))
service = WebhookService(ap)
# Execute - should not raise
- await service.delete_webhook(999)
+ await service.delete_webhook(WORKSPACE_UUID, 999)
# Verify - still called
ap.persistence_mgr.execute_async.assert_called_once()
@@ -453,7 +483,7 @@ class TestWebhookServiceGetEnabledWebhooks:
service = WebhookService(ap)
# Execute
- result = await service.get_enabled_webhooks()
+ result = await service.get_enabled_webhooks(WORKSPACE_UUID)
# Verify
assert result == []
@@ -481,7 +511,7 @@ class TestWebhookServiceGetEnabledWebhooks:
service = WebhookService(ap)
# Execute
- result = await service.get_enabled_webhooks()
+ result = await service.get_enabled_webhooks(WORKSPACE_UUID)
# Verify
assert len(result) == 2
@@ -501,7 +531,170 @@ class TestWebhookServiceGetEnabledWebhooks:
service = WebhookService(ap)
# Execute
- result = await service.get_enabled_webhooks()
+ result = await service.get_enabled_webhooks(WORKSPACE_UUID)
# Verify - should be empty (SQL would filter disabled)
assert result == []
+
+
+ISOLATION_WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
+ISOLATION_WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
+
+
+class _RealPersistenceManager:
+ def __init__(self, engine):
+ self.engine = engine
+
+ async def execute_async(self, *args, **kwargs):
+ async with self.engine.connect() as connection:
+ result = await connection.execute(*args, **kwargs)
+ await connection.commit()
+ return result
+
+ @staticmethod
+ def serialize_model(model, data, masked_columns=None):
+ return {
+ column.name: (
+ getattr(data, column.name).isoformat()
+ if isinstance(getattr(data, column.name), datetime.datetime)
+ else getattr(data, column.name)
+ )
+ for column in model.__table__.columns
+ if column.name not in (masked_columns or [])
+ }
+
+
+@pytest.fixture
+async def tenant_webhook_service(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "webhooks.db"}')
+ async with engine.begin() as connection:
+ await connection.run_sync(Base.metadata.create_all)
+ await connection.execute(
+ sqlalchemy.insert(Workspace),
+ [
+ {
+ 'uuid': ISOLATION_WORKSPACE_A,
+ 'instance_uuid': 'instance',
+ 'name': 'A',
+ 'slug': 'a',
+ 'source': 'cloud_projection',
+ },
+ {
+ 'uuid': ISOLATION_WORKSPACE_B,
+ 'instance_uuid': 'instance',
+ 'name': 'B',
+ 'slug': 'b',
+ 'source': 'cloud_projection',
+ },
+ ],
+ )
+ service = WebhookService(SimpleNamespace(persistence_mgr=_RealPersistenceManager(engine)))
+ yield service
+ await engine.dispose()
+
+
+async def test_webhook_service_requires_workspace(tenant_webhook_service):
+ with pytest.raises(WorkspaceRequiredError):
+ await tenant_webhook_service.get_webhooks(None)
+
+
+async def test_same_name_webhooks_are_isolated(tenant_webhook_service):
+ created_a = await tenant_webhook_service.create_webhook(
+ ISOLATION_WORKSPACE_A,
+ 'deploy',
+ 'https://a.invalid',
+ )
+ created_b = await tenant_webhook_service.create_webhook(
+ ISOLATION_WORKSPACE_B,
+ 'deploy',
+ 'https://b.invalid',
+ )
+
+ assert created_a['workspace_uuid'] == ISOLATION_WORKSPACE_A
+ assert created_b['workspace_uuid'] == ISOLATION_WORKSPACE_B
+ assert [item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_A)] == ['***']
+ assert [item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_B)] == ['***']
+ assert [
+ item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_A, include_secret=True)
+ ] == ['https://a.invalid']
+ assert [
+ item['url'] for item in await tenant_webhook_service.get_webhooks(ISOLATION_WORKSPACE_B, include_secret=True)
+ ] == ['https://b.invalid']
+
+
+async def test_cross_workspace_id_guessing_is_not_found(tenant_webhook_service):
+ created = await tenant_webhook_service.create_webhook(
+ ISOLATION_WORKSPACE_A,
+ 'secret',
+ 'https://a.invalid/hook',
+ )
+ webhook_id = created['id']
+
+ assert await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_B, webhook_id) is None
+ assert not await tenant_webhook_service.update_webhook(
+ ISOLATION_WORKSPACE_B,
+ webhook_id,
+ name='stolen',
+ )
+ assert not await tenant_webhook_service.delete_webhook(ISOLATION_WORKSPACE_B, webhook_id)
+ assert (await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_A, webhook_id))['name'] == 'secret'
+
+
+async def test_update_and_delete_are_scoped(tenant_webhook_service):
+ created = await tenant_webhook_service.create_webhook(
+ ISOLATION_WORKSPACE_A,
+ 'old',
+ 'https://a.invalid/old',
+ )
+ assert await tenant_webhook_service.update_webhook(
+ ISOLATION_WORKSPACE_A,
+ created['id'],
+ name='new',
+ enabled=False,
+ )
+ assert await tenant_webhook_service.get_enabled_webhooks(ISOLATION_WORKSPACE_A) == []
+ assert await tenant_webhook_service.delete_webhook(ISOLATION_WORKSPACE_A, created['id'])
+ assert await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_A, created['id']) is None
+
+
+async def test_masked_webhook_url_roundtrip_preserves_replace_and_clear(tenant_webhook_service):
+ created = await tenant_webhook_service.create_webhook(
+ ISOLATION_WORKSPACE_A,
+ 'roundtrip',
+ 'https://a.invalid/bearer-secret',
+ )
+
+ masked = await tenant_webhook_service.get_webhook(ISOLATION_WORKSPACE_A, created['id'])
+ assert masked['url'] == '***'
+ assert await tenant_webhook_service.update_webhook(
+ ISOLATION_WORKSPACE_A,
+ created['id'],
+ name='preserved',
+ url=masked['url'],
+ )
+ preserved = await tenant_webhook_service.get_webhook(
+ ISOLATION_WORKSPACE_A,
+ created['id'],
+ include_secret=True,
+ )
+ assert preserved['url'] == 'https://a.invalid/bearer-secret'
+
+ assert await tenant_webhook_service.update_webhook(
+ ISOLATION_WORKSPACE_A,
+ created['id'],
+ url='https://a.invalid/replacement',
+ )
+ replaced = await tenant_webhook_service.get_webhook(
+ ISOLATION_WORKSPACE_A,
+ created['id'],
+ include_secret=True,
+ )
+ assert replaced['url'] == 'https://a.invalid/replacement'
+
+ assert await tenant_webhook_service.update_webhook(ISOLATION_WORKSPACE_A, created['id'], url='')
+ cleared = await tenant_webhook_service.get_webhook(
+ ISOLATION_WORKSPACE_A,
+ created['id'],
+ include_secret=True,
+ )
+ assert cleared['url'] == ''
diff --git a/tests/unit_tests/api/test_adapter_session_scoping.py b/tests/unit_tests/api/test_adapter_session_scoping.py
new file mode 100644
index 000000000..cdbcf8cca
--- /dev/null
+++ b/tests/unit_tests/api/test_adapter_session_scoping.py
@@ -0,0 +1,195 @@
+from __future__ import annotations
+
+import asyncio
+from types import SimpleNamespace
+from unittest.mock import AsyncMock
+
+import lark_oapi
+import pytest
+import quart
+
+from langbot.pkg.api.http.context import (
+ PrincipalContext,
+ PrincipalType,
+ RequestContext,
+ WorkspaceContext,
+)
+from langbot.pkg.api.http.controller.groups.platform.adapters import (
+ AdaptersRouterGroup,
+ _AdapterSessionScope,
+ _bind_session_scope,
+ _get_owned_session,
+ _pop_owned_session,
+)
+
+
+pytestmark = pytest.mark.asyncio
+
+
+SENSITIVE_ADAPTER_ROUTES = (
+ ('post', '/api/v1/platform/adapters/lark/create-app'),
+ ('get', '/api/v1/platform/adapters/lark/create-app/status/missing'),
+ ('delete', '/api/v1/platform/adapters/lark/create-app/missing'),
+ ('post', '/api/v1/platform/adapters/weixin/login'),
+ ('get', '/api/v1/platform/adapters/weixin/login/status/missing'),
+ ('delete', '/api/v1/platform/adapters/weixin/login/missing'),
+ ('post', '/api/v1/platform/adapters/dingtalk/create-app'),
+ ('get', '/api/v1/platform/adapters/dingtalk/create-app/status/missing'),
+ ('delete', '/api/v1/platform/adapters/dingtalk/create-app/missing'),
+ ('post', '/api/v1/platform/adapters/wecombot/create-bot'),
+ ('get', '/api/v1/platform/adapters/wecombot/create-bot/status/missing'),
+ ('delete', '/api/v1/platform/adapters/wecombot/create-bot/missing'),
+ ('post', '/api/v1/platform/adapters/qqofficial/bind'),
+ ('get', '/api/v1/platform/adapters/qqofficial/bind/status/missing'),
+ ('delete', '/api/v1/platform/adapters/qqofficial/bind/missing'),
+)
+
+
+def _request_context(
+ *,
+ account_uuid: str = 'account-a',
+ workspace_uuid: str = 'workspace-a',
+ placement_generation: int = 1,
+) -> RequestContext:
+ return RequestContext(
+ instance_uuid='instance-test',
+ placement_generation=placement_generation,
+ request_id='request-test',
+ auth_type='user-token',
+ principal=PrincipalContext(
+ principal_type=PrincipalType.ACCOUNT,
+ account_uuid=account_uuid,
+ ),
+ workspace=WorkspaceContext(
+ workspace_uuid=workspace_uuid,
+ membership_uuid='membership-test',
+ role='developer',
+ permissions=frozenset({'resource.manage'}),
+ ),
+ )
+
+
+async def _create_client(*, role: str = 'developer'):
+ quart_app = quart.Quart(__name__)
+ accounts = {
+ 'owner-token': SimpleNamespace(uuid='account-a', user='owner@example.com'),
+ 'other-token': SimpleNamespace(uuid='account-b', user='other@example.com'),
+ }
+
+ async def get_authenticated_account(token: str):
+ return accounts[token]
+
+ async def resolve_account_workspace(account_uuid: str, requested_workspace_uuid: str | None):
+ workspace_uuid = requested_workspace_uuid or 'workspace-a'
+ return SimpleNamespace(
+ execution=SimpleNamespace(
+ instance_uuid='instance-test',
+ placement_generation=1,
+ ),
+ workspace=SimpleNamespace(uuid=workspace_uuid),
+ membership=SimpleNamespace(
+ uuid=f'membership-{account_uuid}-{workspace_uuid}',
+ role=role,
+ projection_revision=1,
+ ),
+ )
+
+ application = SimpleNamespace(
+ user_service=SimpleNamespace(
+ get_authenticated_account=AsyncMock(side_effect=get_authenticated_account),
+ ),
+ workspace_collaboration_service=SimpleNamespace(
+ resolve_account_workspace=AsyncMock(side_effect=resolve_account_workspace),
+ ),
+ platform_mgr=SimpleNamespace(),
+ )
+ router = AdaptersRouterGroup(application, quart_app)
+ await router.initialize()
+ return quart_app.test_client()
+
+
+@pytest.mark.parametrize(('method', 'path'), SENSITIVE_ADAPTER_ROUTES)
+async def test_sensitive_adapter_flows_require_resource_manage(method: str, path: str):
+ client = await _create_client(role='viewer')
+
+ response = await getattr(client, method)(
+ path,
+ headers={'Authorization': 'Bearer owner-token'},
+ )
+
+ assert response.status_code == 403
+ assert (await response.get_json())['code'] == 'permission_denied'
+
+
+async def test_session_scope_matches_exact_tenant_placement_and_principal():
+ owner_context = _request_context()
+ sessions: dict[str, dict] = {'session-test': {'status': 'waiting'}}
+ _bind_session_scope(sessions['session-test'], owner_context)
+
+ assert sessions['session-test']['scope'] == _AdapterSessionScope.from_request_context(owner_context)
+ assert _get_owned_session(sessions, 'session-test', owner_context) is sessions['session-test']
+
+ for other_context in (
+ _request_context(account_uuid='account-b'),
+ _request_context(workspace_uuid='workspace-b'),
+ _request_context(placement_generation=2),
+ ):
+ assert _get_owned_session(sessions, 'session-test', other_context) is None
+ assert _pop_owned_session(sessions, 'session-test', other_context) is None
+ assert 'session-test' in sessions
+
+ assert _pop_owned_session(sessions, 'session-test', owner_context) is not None
+ assert sessions == {}
+
+
+async def test_lark_session_status_and_delete_hide_cross_scope_sessions(monkeypatch):
+ registration_blocker = asyncio.Event()
+
+ async def fake_register_app(*, on_qr_code, source: str):
+ assert source == 'langbot'
+ on_qr_code({'url': 'https://example.test/lark-qr'})
+ await registration_blocker.wait()
+ raise AssertionError('registration should have been cancelled')
+
+ monkeypatch.setattr(lark_oapi, 'aregister_app', fake_register_app)
+ client = await _create_client()
+ owner_headers = {
+ 'Authorization': 'Bearer owner-token',
+ 'X-Workspace-Id': 'workspace-a',
+ }
+
+ create_response = await client.post(
+ '/api/v1/platform/adapters/lark/create-app',
+ headers=owner_headers,
+ )
+ assert create_response.status_code == 200
+ session_id = (await create_response.get_json())['data']['session_id']
+ status_path = f'/api/v1/platform/adapters/lark/create-app/status/{session_id}'
+ delete_path = f'/api/v1/platform/adapters/lark/create-app/{session_id}'
+
+ for headers in (
+ {
+ 'Authorization': 'Bearer other-token',
+ 'X-Workspace-Id': 'workspace-a',
+ },
+ {
+ 'Authorization': 'Bearer owner-token',
+ 'X-Workspace-Id': 'workspace-b',
+ },
+ ):
+ status_response = await client.get(status_path, headers=headers)
+ delete_response = await client.delete(delete_path, headers=headers)
+ assert status_response.status_code == 404
+ assert delete_response.status_code == 404
+ assert (await status_response.get_json())['msg'] == 'Session not found'
+ assert (await delete_response.get_json())['msg'] == 'Session not found'
+
+ owner_status_response = await client.get(status_path, headers=owner_headers)
+ assert owner_status_response.status_code == 200
+ assert (await owner_status_response.get_json())['data']['status'] == 'waiting'
+
+ owner_delete_response = await client.delete(delete_path, headers=owner_headers)
+ assert owner_delete_response.status_code == 200
+ missing_delete_response = await client.delete(delete_path, headers=owner_headers)
+ assert missing_delete_response.status_code == 404
+ await asyncio.sleep(0)
diff --git a/tests/unit_tests/api/test_apikey_service.py b/tests/unit_tests/api/test_apikey_service.py
index 2065ae3b7..5c4f16c6e 100644
--- a/tests/unit_tests/api/test_apikey_service.py
+++ b/tests/unit_tests/api/test_apikey_service.py
@@ -6,6 +6,7 @@ from unittest.mock import AsyncMock, Mock
import pytest
from langbot.pkg.api.http.service.apikey import ApiKeyService
+from langbot.pkg.entity.persistence.apikey import ApiKeyStatus
@pytest.mark.asyncio
@@ -13,30 +14,63 @@ from langbot.pkg.api.http.service.apikey import ApiKeyService
async def test_verify_api_key_rejects_non_lbk_keys_without_db_query(api_key):
persistence_mgr = SimpleNamespace(execute_async=AsyncMock())
instance_config = SimpleNamespace(data={'api': {'global_api_key': ''}})
- service = ApiKeyService(SimpleNamespace(persistence_mgr=persistence_mgr, instance_config=instance_config))
+ workspace_service = SimpleNamespace(get_execution_binding=AsyncMock())
+ service = ApiKeyService(
+ SimpleNamespace(
+ persistence_mgr=persistence_mgr,
+ instance_config=instance_config,
+ workspace_service=workspace_service,
+ )
+ )
result = await service.verify_api_key(api_key)
assert result is False
persistence_mgr.execute_async.assert_not_awaited()
+ workspace_service.get_execution_binding.assert_not_awaited()
@pytest.mark.asyncio
-@pytest.mark.parametrize(
- ('db_row', 'expected'),
- [
- (object(), True),
- (None, False),
- ],
-)
-async def test_verify_api_key_keeps_db_validation_for_lbk_keys(db_row, expected):
+@pytest.mark.parametrize('key_exists', [True, False])
+async def test_verify_api_key_keeps_db_validation_for_lbk_keys(key_exists):
+ key = (
+ SimpleNamespace(
+ id=1,
+ uuid='key-uuid',
+ workspace_uuid='workspace-a',
+ status=ApiKeyStatus.ACTIVE.value,
+ expires_at=None,
+ scopes=[],
+ )
+ if key_exists
+ else None
+ )
query_result = Mock()
- query_result.first.return_value = db_row
- persistence_mgr = SimpleNamespace(execute_async=AsyncMock(return_value=query_result))
+ query_result.first.return_value = key
+ persistence_mgr = SimpleNamespace(execute_async=AsyncMock(side_effect=[query_result, Mock(rowcount=1)]))
instance_config = SimpleNamespace(data={'api': {'global_api_key': ''}})
- service = ApiKeyService(SimpleNamespace(persistence_mgr=persistence_mgr, instance_config=instance_config))
+ workspace_service = SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ )
+ )
+ )
+ service = ApiKeyService(
+ SimpleNamespace(
+ persistence_mgr=persistence_mgr,
+ instance_config=instance_config,
+ workspace_service=workspace_service,
+ )
+ )
result = await service.verify_api_key('lbk_valid_format')
- assert result is expected
- persistence_mgr.execute_async.assert_awaited_once()
+ assert result is key_exists
+ assert persistence_mgr.execute_async.await_count == (2 if key_exists else 1)
+ if key_exists:
+ workspace_service.get_execution_binding.assert_awaited_once_with('workspace-a')
+ else:
+ workspace_service.get_execution_binding.assert_not_awaited()
diff --git a/tests/unit_tests/api/test_bot_controller_secrets.py b/tests/unit_tests/api/test_bot_controller_secrets.py
new file mode 100644
index 000000000..971129315
--- /dev/null
+++ b/tests/unit_tests/api/test_bot_controller_secrets.py
@@ -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()
diff --git a/tests/unit_tests/api/test_extensions_runtime_fence.py b/tests/unit_tests/api/test_extensions_runtime_fence.py
new file mode 100644
index 000000000..0ca854ce3
--- /dev/null
+++ b/tests/unit_tests/api/test_extensions_runtime_fence.py
@@ -0,0 +1,94 @@
+from __future__ import annotations
+
+from types import SimpleNamespace
+from unittest.mock import AsyncMock
+
+import pytest
+import quart
+
+from langbot.pkg.api.http.controller.groups.extensions import ExtensionsRouterGroup
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
+
+
+@pytest.mark.asyncio
+async def test_extensions_route_hides_runtime_bound_to_another_workspace():
+ account = SimpleNamespace(uuid='account-a', user='owner@example.com')
+ connector = SimpleNamespace(
+ is_enable_plugin=True,
+ require_workspace_context=AsyncMock(side_effect=WorkspaceNotFoundError('Plugin resource not found')),
+ list_plugins=AsyncMock(return_value=[]),
+ )
+ ap = SimpleNamespace(
+ user_service=SimpleNamespace(
+ get_authenticated_account=AsyncMock(return_value=account),
+ ),
+ workspace_collaboration_service=SimpleNamespace(
+ resolve_account_workspace=AsyncMock(
+ return_value=SimpleNamespace(
+ workspace=SimpleNamespace(uuid='workspace-a'),
+ membership=SimpleNamespace(uuid='membership-a', role='owner', projection_revision=0),
+ execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=2),
+ )
+ )
+ ),
+ plugin_connector=connector,
+ mcp_service=SimpleNamespace(get_mcp_servers=AsyncMock(return_value=[])),
+ skill_service=SimpleNamespace(list_skills=AsyncMock(return_value=[])),
+ )
+ quart_app = quart.Quart(__name__)
+ router = ExtensionsRouterGroup(ap, quart_app)
+ await router.initialize()
+
+ response = await quart_app.test_client().get(
+ '/api/v1/extensions',
+ headers={'Authorization': 'Bearer token'},
+ )
+
+ assert response.status_code == 404
+ connector.list_plugins.assert_not_awaited()
+ ap.mcp_service.get_mcp_servers.assert_not_awaited()
+ ap.skill_service.list_skills.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_extensions_route_redacts_plugin_secrets_without_mutating_runtime_data():
+ account = SimpleNamespace(uuid='account-a', user='viewer@example.com')
+ raw_plugin = {
+ 'plugin_config': {'apiKey': 'plugin-secret', 'nested': {'token': 'nested-secret'}},
+ 'debug': {'plugin_debug_key': 'debug-secret'},
+ }
+ connector = SimpleNamespace(
+ is_enable_plugin=True,
+ require_workspace_context=AsyncMock(),
+ list_plugins=AsyncMock(return_value=[raw_plugin]),
+ )
+ ap = SimpleNamespace(
+ user_service=SimpleNamespace(get_authenticated_account=AsyncMock(return_value=account)),
+ workspace_collaboration_service=SimpleNamespace(
+ resolve_account_workspace=AsyncMock(
+ return_value=SimpleNamespace(
+ workspace=SimpleNamespace(uuid='workspace-a'),
+ membership=SimpleNamespace(uuid='membership-a', role='viewer', projection_revision=0),
+ execution=SimpleNamespace(instance_uuid='instance-a', placement_generation=2),
+ )
+ )
+ ),
+ plugin_connector=connector,
+ mcp_service=SimpleNamespace(get_mcp_servers=AsyncMock(return_value=[])),
+ skill_service=SimpleNamespace(list_skills=AsyncMock(return_value=[])),
+ )
+ quart_app = quart.Quart(__name__)
+ router = ExtensionsRouterGroup(ap, quart_app)
+ await router.initialize()
+
+ response = await quart_app.test_client().get(
+ '/api/v1/extensions',
+ headers={'Authorization': 'Bearer token', 'X-Workspace-Id': 'workspace-a'},
+ )
+
+ assert response.status_code == 200
+ plugin = (await response.get_json())['data']['extensions'][0]['plugin']
+ assert plugin['plugin_config']['apiKey'] == '***'
+ assert plugin['plugin_config']['nested']['token'] == '***'
+ assert plugin['debug']['plugin_debug_key'] == '***'
+ assert raw_plugin['plugin_config']['apiKey'] == 'plugin-secret'
diff --git a/tests/unit_tests/api/test_file_upload_scoping.py b/tests/unit_tests/api/test_file_upload_scoping.py
new file mode 100644
index 000000000..0463a4c84
--- /dev/null
+++ b/tests/unit_tests/api/test_file_upload_scoping.py
@@ -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'
diff --git a/tests/unit_tests/api/test_knowledge_migration_runtime_fence.py b/tests/unit_tests/api/test_knowledge_migration_runtime_fence.py
new file mode 100644
index 000000000..5750b7374
--- /dev/null
+++ b/tests/unit_tests/api/test_knowledge_migration_runtime_fence.py
@@ -0,0 +1,69 @@
+from __future__ import annotations
+
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, Mock
+
+import pytest
+
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.api.http.controller.groups.knowledge.migration import KnowledgeMigrationRouterGroup
+from langbot.pkg.workspace.errors import WorkspaceInvariantError, WorkspaceNotFoundError
+
+
+CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=3,
+)
+
+
+@pytest.mark.asyncio
+async def test_background_migration_propagates_generation_change_before_runtime_call():
+ connector = SimpleNamespace(
+ require_workspace_context=AsyncMock(side_effect=[CONTEXT, WorkspaceNotFoundError('Plugin resource not found')]),
+ list_knowledge_engines=AsyncMock(return_value=[]),
+ )
+ router = object.__new__(KnowledgeMigrationRouterGroup)
+ router.ap = SimpleNamespace(
+ plugin_connector=connector,
+ workspace_service=SimpleNamespace(
+ get_local_execution_binding=AsyncMock(return_value=CONTEXT),
+ ),
+ logger=Mock(),
+ )
+ router._table_exists = AsyncMock(return_value=False)
+ router._set_migration_flag = AsyncMock()
+ task_context = SimpleNamespace(trace=Mock())
+
+ with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
+ await router._execute_rag_migration(
+ CONTEXT,
+ task_context,
+ install_plugin=False,
+ )
+
+ assert connector.require_workspace_context.await_count == 2
+ connector.list_knowledge_engines.assert_not_awaited()
+ router._set_migration_flag.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_cloud_migration_is_rejected_before_legacy_table_access():
+ router = object.__new__(KnowledgeMigrationRouterGroup)
+ router.ap = SimpleNamespace(
+ workspace_service=SimpleNamespace(
+ get_local_execution_binding=AsyncMock(side_effect=WorkspaceInvariantError('not an OSS local workspace')),
+ ),
+ plugin_connector=SimpleNamespace(require_workspace_context=AsyncMock()),
+ logger=Mock(),
+ )
+ router._table_exists = AsyncMock()
+ router._set_migration_flag = AsyncMock()
+ task_context = SimpleNamespace(trace=Mock())
+
+ with pytest.raises(WorkspaceNotFoundError, match='migration is unavailable'):
+ await router._execute_rag_migration(CONTEXT, task_context, install_plugin=False)
+
+ router._table_exists.assert_not_awaited()
+ router.ap.plugin_connector.require_workspace_context.assert_not_awaited()
+ router._set_migration_flag.assert_not_awaited()
diff --git a/tests/unit_tests/api/test_mcp_controller.py b/tests/unit_tests/api/test_mcp_controller.py
index 47f2247f2..d5e9f6da1 100644
--- a/tests/unit_tests/api/test_mcp_controller.py
+++ b/tests/unit_tests/api/test_mcp_controller.py
@@ -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}
diff --git a/tests/unit_tests/api/test_plugin_runtime_route_fence.py b/tests/unit_tests/api/test_plugin_runtime_route_fence.py
new file mode 100644
index 000000000..353301f4c
--- /dev/null
+++ b/tests/unit_tests/api/test_plugin_runtime_route_fence.py
@@ -0,0 +1,108 @@
+from __future__ import annotations
+
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, Mock
+
+import pytest
+import quart
+
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
+
+
+CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=4,
+)
+
+
+@pytest.fixture(scope='module')
+def plugin_router_cls():
+ from tests.utils.import_isolation import MockLifecycleControlScope, isolated_sys_modules
+
+ class FakeMinimalApplication:
+ pass
+
+ mock_app = Mock(Application=FakeMinimalApplication)
+ mock_entities = Mock(LifecycleControlScope=MockLifecycleControlScope)
+ clear = [
+ 'langbot.pkg.core.taskmgr',
+ 'langbot.pkg.api.http.controller.group',
+ 'langbot.pkg.api.http.controller.groups',
+ 'langbot.pkg.api.http.controller.groups.plugins',
+ 'langbot.pkg.api.http.controller.main',
+ ]
+ with isolated_sys_modules(
+ mocks={
+ 'langbot.pkg.core.app': mock_app,
+ 'langbot.pkg.core.entities': mock_entities,
+ },
+ clear=clear,
+ ):
+ from langbot.pkg.api.http.controller.groups.plugins import PluginsRouterGroup
+
+ yield PluginsRouterGroup
+
+
+@pytest.mark.asyncio
+async def test_public_plugin_asset_route_is_disabled_for_multi_workspace_policy(plugin_router_cls):
+ connector = SimpleNamespace(
+ get_plugin_icon=AsyncMock(),
+ require_workspace_context=AsyncMock(),
+ )
+ ap = SimpleNamespace(
+ plugin_connector=connector,
+ workspace_service=SimpleNamespace(
+ policy=SimpleNamespace(multi_workspace_enabled=True),
+ ),
+ )
+ quart_app = quart.Quart(__name__)
+ router = plugin_router_cls(ap, quart_app)
+ await router.initialize()
+
+ response = await quart_app.test_client().get('/api/v1/plugins/author/plugin/icon')
+
+ assert response.status_code == 404
+ connector.require_workspace_context.assert_not_awaited()
+ connector.get_plugin_icon.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_public_plugin_asset_uses_trusted_oss_singleton_binding(plugin_router_cls):
+ connector = SimpleNamespace(
+ require_workspace_context=AsyncMock(side_effect=lambda context: context),
+ )
+ binding = SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=4,
+ )
+ router = object.__new__(plugin_router_cls)
+ router.ap = SimpleNamespace(
+ plugin_connector=connector,
+ workspace_service=SimpleNamespace(
+ policy=SimpleNamespace(multi_workspace_enabled=False),
+ get_local_execution_binding=AsyncMock(return_value=binding),
+ ),
+ )
+
+ result = await router._require_public_plugin_runtime_context()
+
+ assert result == CONTEXT
+ connector.require_workspace_context.assert_awaited_once_with(CONTEXT)
+
+
+@pytest.mark.asyncio
+async def test_background_plugin_operation_refences_captured_generation(plugin_router_cls):
+ operation = AsyncMock()
+ connector = SimpleNamespace(
+ require_workspace_context=AsyncMock(side_effect=WorkspaceNotFoundError('Plugin resource not found')),
+ )
+ router = object.__new__(plugin_router_cls)
+ router.ap = SimpleNamespace(plugin_connector=connector)
+
+ with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
+ await router._run_fenced_plugin_operation(CONTEXT, operation)
+
+ operation.assert_not_awaited()
diff --git a/tests/unit_tests/api/test_resource_secret_permissions.py b/tests/unit_tests/api/test_resource_secret_permissions.py
new file mode 100644
index 000000000..9e4e31d98
--- /dev/null
+++ b/tests/unit_tests/api/test_resource_secret_permissions.py
@@ -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
diff --git a/tests/unit_tests/api/test_stats_controller.py b/tests/unit_tests/api/test_stats_controller.py
new file mode 100644
index 000000000..2fdfcf731
--- /dev/null
+++ b/tests/unit_tests/api/test_stats_controller.py
@@ -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'
diff --git a/tests/unit_tests/box/test_box_connector.py b/tests/unit_tests/box/test_box_connector.py
index ddd4899b0..0034d0199 100644
--- a/tests/unit_tests/box/test_box_connector.py
+++ b/tests/unit_tests/box/test_box_connector.py
@@ -6,12 +6,26 @@ from unittest.mock import Mock
import pytest
from langbot_plugin.box.client import ActionRPCBoxClient
+from langbot_plugin.box.errors import BoxRuntimeUnavailableError
+from langbot_plugin.box.security import (
+ BOX_CONTROL_TOKEN_ENV,
+ BOX_CONTROL_TOKEN_HEADER,
+ BOX_INSTANCE_HEADER,
+ BOX_PLACEMENT_GENERATION_HEADER,
+ BOX_TRUSTED_INSTANCE_ENV,
+ BOX_WORKSPACE_HEADER,
+)
+from langbot_plugin.entities.io.context import ActionContext
from langbot.pkg.box.connector import BoxRuntimeConnector
+_CONTROL_TOKEN = 'box-control-token-that-is-longer-than-32-bytes'
+
+
def make_app(logger: Mock, runtime_endpoint: str = ''):
return SimpleNamespace(
logger=logger,
+ workspace_service=SimpleNamespace(instance_uuid='instance-a'),
instance_config=SimpleNamespace(
data={
'box': {
@@ -104,3 +118,130 @@ def test_box_runtime_connector_dispose_terminates_subprocess(monkeypatch: pytest
ctrl_task.cancel.assert_called_once()
assert connector._handler_task is None
assert connector._ctrl_task is None
+
+
+def test_box_runtime_connector_builds_host_control_headers(monkeypatch: pytest.MonkeyPatch):
+ monkeypatch.setenv(BOX_CONTROL_TOKEN_ENV, _CONTROL_TOKEN)
+ connector = BoxRuntimeConnector(make_app(Mock(), runtime_endpoint='http://box-runtime:5410'))
+
+ headers = connector.get_control_headers()
+
+ assert headers == {
+ BOX_CONTROL_TOKEN_HEADER: _CONTROL_TOKEN,
+ BOX_INSTANCE_HEADER: 'instance-a',
+ }
+ assert _CONTROL_TOKEN not in connector._resolve_rpc_ws_url()
+
+
+def test_box_runtime_connector_builds_placement_scoped_relay_headers(
+ monkeypatch: pytest.MonkeyPatch,
+):
+ monkeypatch.setenv(BOX_CONTROL_TOKEN_ENV, _CONTROL_TOKEN)
+ connector = BoxRuntimeConnector(make_app(Mock(), runtime_endpoint='http://box-runtime:5410'))
+
+ headers = connector.get_relay_headers(
+ ActionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=7,
+ )
+ )
+
+ assert headers == {
+ BOX_CONTROL_TOKEN_HEADER: _CONTROL_TOKEN,
+ BOX_INSTANCE_HEADER: 'instance-a',
+ BOX_WORKSPACE_HEADER: 'workspace-a',
+ BOX_PLACEMENT_GENERATION_HEADER: '7',
+ }
+
+
+def test_box_runtime_connector_rejects_relay_context_from_other_instance(
+ monkeypatch: pytest.MonkeyPatch,
+):
+ monkeypatch.setenv(BOX_CONTROL_TOKEN_ENV, _CONTROL_TOKEN)
+ connector = BoxRuntimeConnector(make_app(Mock(), runtime_endpoint='http://box-runtime:5410'))
+
+ with pytest.raises(BoxRuntimeUnavailableError, match='another LangBot instance'):
+ connector.get_relay_headers(
+ ActionContext(
+ instance_uuid='instance-b',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ )
+ )
+
+
+def test_external_box_runtime_fails_closed_without_control_token(monkeypatch: pytest.MonkeyPatch):
+ monkeypatch.delenv(BOX_CONTROL_TOKEN_ENV, raising=False)
+ connector = BoxRuntimeConnector(make_app(Mock(), runtime_endpoint='http://box-runtime:5410'))
+
+ with pytest.raises(BoxRuntimeUnavailableError, match=BOX_CONTROL_TOKEN_ENV):
+ connector.get_control_headers()
+
+
+async def test_local_stdio_injects_generated_token_and_trusted_instance(
+ monkeypatch: pytest.MonkeyPatch,
+):
+ monkeypatch.delenv(BOX_CONTROL_TOKEN_ENV, raising=False)
+ captured = {}
+
+ class FakeStdioClientController:
+ def __init__(self, **kwargs):
+ captured.update(kwargs)
+ self.process = Mock()
+
+ async def run(self, callback):
+ await callback(None)
+
+ monkeypatch.setattr(
+ 'langbot_plugin.runtime.io.controllers.stdio.client.StdioClientController',
+ FakeStdioClientController,
+ )
+ connector = BoxRuntimeConnector(make_app(Mock()))
+
+ def fake_callback(_transport_name, connected, _connect_error):
+ async def callback(_connection):
+ connected.set()
+
+ return callback
+
+ monkeypatch.setattr(connector, '_make_connection_callback', fake_callback)
+
+ await connector._start_local_stdio()
+
+ assert len(captured['env'][BOX_CONTROL_TOKEN_ENV]) >= 32
+ assert captured['env'][BOX_TRUSTED_INSTANCE_ENV] == 'instance-a'
+
+
+async def test_websocket_controller_receives_control_headers(monkeypatch: pytest.MonkeyPatch):
+ monkeypatch.setenv(BOX_CONTROL_TOKEN_ENV, _CONTROL_TOKEN)
+ captured = {}
+
+ class FakeWebSocketClientController:
+ def __init__(self, **kwargs):
+ captured.update(kwargs)
+
+ async def run(self, callback):
+ await callback(None)
+
+ monkeypatch.setattr(
+ 'langbot_plugin.runtime.io.controllers.ws.client.WebSocketClientController',
+ FakeWebSocketClientController,
+ )
+ connector = BoxRuntimeConnector(make_app(Mock(), runtime_endpoint='http://box-runtime:5410'))
+
+ def fake_callback(_transport_name, connected, _connect_error):
+ async def callback(_connection):
+ connected.set()
+
+ return callback
+
+ monkeypatch.setattr(connector, '_make_connection_callback', fake_callback)
+
+ await connector._connect_ws('ws://box-runtime:5410/rpc/ws', 'WebSocket')
+
+ assert captured['additional_headers'] == {
+ BOX_CONTROL_TOKEN_HEADER: _CONTROL_TOKEN,
+ BOX_INSTANCE_HEADER: 'instance-a',
+ }
+ assert _CONTROL_TOKEN not in captured['ws_url']
diff --git a/tests/unit_tests/box/test_box_deployment_config.py b/tests/unit_tests/box/test_box_deployment_config.py
new file mode 100644
index 000000000..f4b1d4d3c
--- /dev/null
+++ b/tests/unit_tests/box/test_box_deployment_config.py
@@ -0,0 +1,65 @@
+from pathlib import Path
+
+
+_REPO_ROOT = Path(__file__).resolve().parents[3]
+
+
+def test_compose_injects_the_same_box_control_token_into_host_and_runtime():
+ compose = (_REPO_ROOT / 'docker' / 'docker-compose.yaml').read_text(encoding='utf-8')
+ box_service = compose.split(' langbot_box:', 1)[1].split(' langbot:', 1)[0]
+ langbot_service = compose.split(' langbot:', 1)[1]
+ token_env = 'LANGBOT_BOX_CONTROL_TOKEN=${LANGBOT_BOX_CONTROL_TOKEN:-}'
+
+ assert token_env in box_service
+ assert token_env in langbot_service
+
+
+def test_kubernetes_uses_one_secret_for_box_runtime_and_langbot():
+ manifest = (_REPO_ROOT / 'docker' / 'kubernetes.yaml').read_text(encoding='utf-8')
+ box_deployment = manifest.split('name: langbot-box', 1)[1].split('# Service for LangBot Box runtime', 1)[0]
+ langbot_deployment = manifest.split('# Deployment for LangBot\n', 1)[1]
+ secret_reference = '\n'.join(
+ [
+ '- name: LANGBOT_BOX_CONTROL_TOKEN',
+ ' valueFrom:',
+ ' secretKeyRef:',
+ ' name: langbot-box-control',
+ ' key: token',
+ ]
+ )
+
+ assert secret_reference in box_deployment
+ assert secret_reference in langbot_deployment
+ assert '--from-literal=token="$(openssl rand -hex 32)"' in manifest
+
+
+def test_compose_injects_same_plugin_runtime_control_token_into_both_services():
+ compose = (_REPO_ROOT / 'docker' / 'docker-compose.yaml').read_text(encoding='utf-8')
+ runtime_service = compose.split(' langbot_plugin_runtime:', 1)[1].split(' langbot_box:', 1)[0]
+ langbot_service = compose.split(' langbot:', 1)[1]
+ token_env = 'LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN=${LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN:-}'
+
+ assert token_env in runtime_service
+ assert token_env in langbot_service
+
+
+def test_kubernetes_uses_one_secret_for_plugin_runtime_and_langbot():
+ manifest = (_REPO_ROOT / 'docker' / 'kubernetes.yaml').read_text(encoding='utf-8')
+ runtime_deployment = manifest.split('# Deployment for LangBot Plugin Runtime', 1)[1].split(
+ '# Service for LangBot Plugin Runtime',
+ 1,
+ )[0]
+ langbot_deployment = manifest.split('# Deployment for LangBot\n', 1)[1]
+ secret_reference = '\n'.join(
+ [
+ '- name: LANGBOT_PLUGIN_RUNTIME_CONTROL_TOKEN',
+ ' valueFrom:',
+ ' secretKeyRef:',
+ ' name: langbot-plugin-runtime-control',
+ ' key: token',
+ ]
+ )
+
+ assert secret_reference in runtime_deployment
+ assert secret_reference in langbot_deployment
+ assert 'create secret generic langbot-plugin-runtime-control' in manifest
diff --git a/tests/unit_tests/box/test_box_service.py b/tests/unit_tests/box/test_box_service.py
index 4d66ec8f8..43631504a 100644
--- a/tests/unit_tests/box/test_box_service.py
+++ b/tests/unit_tests/box/test_box_service.py
@@ -15,6 +15,7 @@ from langbot_plugin.box.backend import BaseSandboxBackend
from langbot_plugin.box.client import BoxRuntimeClient, ActionRPCBoxClient
from langbot_plugin.box.errors import (
BoxBackendUnavailableError,
+ BoxError,
BoxSessionConflictError,
BoxSessionNotFoundError,
BoxValidationError,
@@ -30,9 +31,27 @@ from langbot_plugin.box.models import (
BoxSpec,
)
from langbot_plugin.box.runtime import BoxRuntime
+from langbot_plugin.box.security import (
+ BOX_CONTROL_TOKEN_HEADER,
+ BOX_INSTANCE_HEADER,
+ BOX_PLACEMENT_GENERATION_HEADER,
+ BOX_WORKSPACE_HEADER,
+)
+from langbot_plugin.entities.io.context import ActionContext
+from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.box.service import BoxService
_UTC = dt.timezone.utc
+_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+)
+_ACTION_CONTEXT = ActionContext(
+ instance_uuid=_CONTEXT.instance_uuid,
+ workspace_uuid=_CONTEXT.workspace_uuid,
+ placement_generation=_CONTEXT.placement_generation,
+)
class _InProcessBoxRuntimeClient(BoxRuntimeClient):
@@ -44,37 +63,55 @@ class _InProcessBoxRuntimeClient(BoxRuntimeClient):
async def initialize(self):
await self._runtime.initialize()
- async def execute(self, spec):
+ async def execute(self, spec, *, action_context=None):
return await self._runtime.execute(spec)
async def shutdown(self):
await self._runtime.shutdown()
- async def get_status(self):
+ async def get_status(self, *, action_context=None):
return await self._runtime.get_status()
- async def get_sessions(self):
+ async def get_sessions(self, *, action_context=None):
return self._runtime.get_sessions()
async def get_backend_info(self):
return await self._runtime.get_backend_info()
- async def delete_session(self, session_id):
+ async def delete_session(self, session_id, *, action_context=None):
await self._runtime.delete_session(session_id)
- async def create_session(self, spec):
+ async def create_session(self, spec, *, action_context=None):
return await self._runtime.create_session(spec)
- async def start_managed_process(self, session_id: str, spec: BoxManagedProcessSpec):
+ async def start_managed_process(
+ self,
+ session_id: str,
+ spec: BoxManagedProcessSpec,
+ *,
+ action_context=None,
+ ):
return await self._runtime.start_managed_process(session_id, spec)
- async def get_managed_process(self, session_id: str, process_id: str = 'default'):
+ async def get_managed_process(
+ self,
+ session_id: str,
+ process_id: str = 'default',
+ *,
+ action_context=None,
+ ):
return self._runtime.get_managed_process(session_id, process_id)
- async def stop_managed_process(self, session_id: str, process_id: str = 'default'):
+ async def stop_managed_process(
+ self,
+ session_id: str,
+ process_id: str = 'default',
+ *,
+ action_context=None,
+ ):
await self._runtime.stop_managed_process(session_id, process_id)
- async def get_session(self, session_id: str):
+ async def get_session(self, session_id: str, *, action_context=None):
return self._runtime.get_session(session_id)
async def init(self, config: dict) -> None:
@@ -134,6 +171,12 @@ class FakeBackend(BaseSandboxBackend):
def make_query(query_id: int = 42) -> pipeline_query.Query:
return pipeline_query.Query.model_construct(
query_id=query_id,
+ query_uuid=f'query-{query_id}',
+ instance_uuid=_CONTEXT.instance_uuid,
+ workspace_uuid=_CONTEXT.workspace_uuid,
+ placement_generation=_CONTEXT.placement_generation,
+ bot_uuid='bot-a',
+ pipeline_uuid='pipeline-a',
launcher_type='person',
launcher_id='test_user',
sender_id='test_user',
@@ -170,8 +213,19 @@ def make_app(
if workspace_quota_mb is not None:
box_config['local']['workspace_quota_mb'] = workspace_quota_mb
+ workspace_service = SimpleNamespace(
+ instance_uuid=_CONTEXT.instance_uuid,
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid=_CONTEXT.instance_uuid,
+ workspace_uuid=_CONTEXT.workspace_uuid,
+ placement_generation=_CONTEXT.placement_generation,
+ )
+ ),
+ )
return SimpleNamespace(
logger=logger,
+ workspace_service=workspace_service,
instance_config=SimpleNamespace(
data={
'box': box_config,
@@ -289,12 +343,77 @@ async def test_box_service_get_sessions_delegates_to_client():
service = BoxService(make_app(Mock()), client=client)
service._available = True
- sessions = await service.get_sessions()
+ sessions = await service.get_sessions(_CONTEXT)
assert sessions == [{'session_id': 'test-session'}]
client.get_sessions.assert_awaited_once()
+@pytest.mark.asyncio
+async def test_box_service_relay_connection_is_binding_checked_and_scoped():
+ app = make_app(Mock())
+ client = Mock()
+ client.get_managed_process_websocket_url = Mock(
+ return_value='ws://box/v1/sessions/physical/managed-process/server-a/ws'
+ )
+ connector = Mock()
+ connector.ws_relay_base_url = 'http://box:5410'
+ connector.get_relay_headers = Mock(
+ return_value={
+ BOX_CONTROL_TOKEN_HEADER: 'secret',
+ BOX_INSTANCE_HEADER: _CONTEXT.instance_uuid,
+ BOX_WORKSPACE_HEADER: _CONTEXT.workspace_uuid,
+ BOX_PLACEMENT_GENERATION_HEADER: '1',
+ }
+ )
+ service = BoxService(app, client=client)
+ service._runtime_connector = connector
+
+ url, headers = await service.get_managed_process_websocket_connection(
+ _CONTEXT,
+ 'mcp-shared',
+ 'server-a',
+ )
+
+ assert url == 'ws://box/v1/sessions/physical/managed-process/server-a/ws'
+ assert headers[BOX_WORKSPACE_HEADER] == _CONTEXT.workspace_uuid
+ assert headers[BOX_PLACEMENT_GENERATION_HEADER] == '1'
+ app.workspace_service.get_execution_binding.assert_awaited_once_with(
+ _CONTEXT.workspace_uuid,
+ expected_generation=_CONTEXT.placement_generation,
+ )
+ action_context = connector.get_relay_headers.call_args.args[0]
+ assert action_context == _ACTION_CONTEXT
+ client.get_managed_process_websocket_url.assert_called_once_with(
+ 'mcp-shared',
+ 'http://box:5410',
+ 'server-a',
+ action_context=_ACTION_CONTEXT,
+ )
+
+
+@pytest.mark.asyncio
+async def test_box_service_relay_connection_rejects_stale_binding():
+ app = make_app(Mock())
+ app.workspace_service.get_execution_binding.return_value = SimpleNamespace(
+ instance_uuid=_CONTEXT.instance_uuid,
+ workspace_uuid=_CONTEXT.workspace_uuid,
+ placement_generation=2,
+ )
+ client = Mock()
+ service = BoxService(app, client=client)
+ service._runtime_connector = Mock()
+
+ with pytest.raises(BoxValidationError, match='stale Workspace placement'):
+ await service.get_managed_process_websocket_connection(
+ _CONTEXT,
+ 'mcp-shared',
+ 'server-a',
+ )
+
+ service._runtime_connector.get_relay_headers.assert_not_called()
+
+
def test_box_service_dispose_delegates_to_internal_connector(monkeypatch: pytest.MonkeyPatch):
connector = Mock()
connector.client = Mock()
@@ -369,7 +488,14 @@ async def test_box_service_session_id_uses_query_attributes_without_variables():
service = BoxService(make_app(logger), client=_InProcessBoxRuntimeClient(logger, runtime))
await service.initialize()
- query = pipeline_query.Query.model_construct(query_id=7, launcher_type='group', launcher_id='room-1')
+ query = pipeline_query.Query.model_construct(
+ query_id=7,
+ instance_uuid=_CONTEXT.instance_uuid,
+ workspace_uuid=_CONTEXT.workspace_uuid,
+ placement_generation=_CONTEXT.placement_generation,
+ launcher_type='group',
+ launcher_id='room-1',
+ )
result = await service.execute_tool({'command': 'pwd'}, query)
assert result['session_id'] == 'group_room-1'
@@ -385,7 +511,12 @@ async def test_box_service_session_id_falls_back_to_query_id_for_synthetic_queri
service = BoxService(make_app(logger), client=_InProcessBoxRuntimeClient(logger, runtime))
await service.initialize()
- query = pipeline_query.Query.model_construct(query_id=7)
+ query = pipeline_query.Query.model_construct(
+ query_id=7,
+ instance_uuid=_CONTEXT.instance_uuid,
+ workspace_uuid=_CONTEXT.workspace_uuid,
+ placement_generation=_CONTEXT.placement_generation,
+ )
result = await service.execute_tool({'command': 'pwd'}, query)
assert result['session_id'] == 'query_7'
@@ -407,8 +538,22 @@ async def test_box_service_forced_global_scope_overrides_pipeline_template():
await service.initialize()
# Two distinct callers that would otherwise get separate sandboxes.
- q1 = pipeline_query.Query.model_construct(query_id=1, launcher_type='group', launcher_id='room-1')
- q2 = pipeline_query.Query.model_construct(query_id=2, launcher_type='person', launcher_id='alice')
+ q1 = pipeline_query.Query.model_construct(
+ query_id=1,
+ instance_uuid=_CONTEXT.instance_uuid,
+ workspace_uuid=_CONTEXT.workspace_uuid,
+ placement_generation=_CONTEXT.placement_generation,
+ launcher_type='group',
+ launcher_id='room-1',
+ )
+ q2 = pipeline_query.Query.model_construct(
+ query_id=2,
+ instance_uuid=_CONTEXT.instance_uuid,
+ workspace_uuid=_CONTEXT.workspace_uuid,
+ placement_generation=_CONTEXT.placement_generation,
+ launcher_type='person',
+ launcher_id='alice',
+ )
r1 = await service.execute_tool({'command': 'pwd'}, q1)
r2 = await service.execute_tool({'command': 'pwd'}, q2)
@@ -511,7 +656,7 @@ async def test_box_service_uses_default_workspace_when_host_path_omitted(tmp_pat
assert result['ok'] is True
assert backend.start_calls == ['person_test_user']
assert backend.exec_calls == [('person_test_user', 'pwd')]
- assert backend.start_specs[0].host_path == os.path.realpath(host_dir)
+ assert backend.start_specs[0].host_path == service._tenant_workspace(_CONTEXT)
@pytest.mark.asyncio
@@ -954,11 +1099,16 @@ async def test_box_service_rejects_execution_when_workspace_already_exceeds_quot
runtime = BoxRuntime(logger=logger, backends=[backend], session_ttl_sec=300)
host_dir = tmp_path / 'quota-workspace'
host_dir.mkdir()
- (host_dir / 'already-too-large.bin').write_bytes(b'x' * (2 * 1024 * 1024))
app = make_app(logger, [str(tmp_path)], workspace_quota_mb=1)
app.instance_config.data['box']['local']['default_workspace'] = str(host_dir)
service = BoxService(app, client=_InProcessBoxRuntimeClient(logger, runtime))
+ tenant_host_dir = service._tenant_workspace(_CONTEXT)
+ assert tenant_host_dir is not None
+ os.makedirs(tenant_host_dir, exist_ok=True)
+ with open(os.path.join(tenant_host_dir, 'already-too-large.bin'), 'wb') as handle:
+ handle.write(b'x' * (2 * 1024 * 1024))
+
await service.initialize()
with pytest.raises(BoxValidationError, match='workspace quota exceeded before execution'):
@@ -1097,7 +1247,7 @@ async def test_service_records_errors_on_failure():
with pytest.raises(Exception):
await service.execute_tool({'command': 'echo hello'}, make_query(50))
- errors = service.get_recent_errors()
+ errors = service.get_recent_errors(_CONTEXT)
assert len(errors) == 1
assert errors[0]['type'] == 'BoxBackendUnavailableError'
assert errors[0]['query_id'] == '50'
@@ -1116,7 +1266,7 @@ async def test_service_error_ring_buffer_capped():
with pytest.raises(Exception):
await service.execute_tool({'command': 'fail'}, make_query(100 + i))
- errors = service.get_recent_errors()
+ errors = service.get_recent_errors(_CONTEXT)
assert len(errors) == 50
# Oldest should have been evicted, newest kept
assert errors[0]['query_id'] == '110'
@@ -1131,7 +1281,7 @@ async def test_service_get_status_aggregates_runtime_and_profile():
service = BoxService(make_app(logger), client=_InProcessBoxRuntimeClient(logger, runtime))
await service.initialize()
- status = await service.get_status()
+ status = await service.get_status(_CONTEXT)
assert status['profile'] == 'default'
assert status['backend']['name'] == 'fake'
assert status['backend']['available'] is True
@@ -1175,7 +1325,12 @@ async def _make_rpc_pair(runtime: BoxRuntime):
client_conn, server_conn = _make_queue_connection_pair()
- server_handler = BoxServerHandler(server_conn, runtime)
+ server_handler = BoxServerHandler(
+ server_conn,
+ runtime,
+ host_control_authenticated=True,
+ trusted_instance_uuid=_CONTEXT.instance_uuid,
+ )
server_task = asyncio.create_task(server_handler.run())
client_handler = Handler.__new__(Handler)
@@ -1199,7 +1354,7 @@ async def test_rpc_client_execute():
client, server_task, client_task = await _make_rpc_pair(runtime)
try:
spec = BoxSpec.model_validate({'cmd': 'echo remote', 'session_id': 'r-1'})
- result = await client.execute(spec)
+ result = await client.execute(spec, action_context=_ACTION_CONTEXT)
assert result.session_id == 'r-1'
assert result.status == BoxExecutionStatus.COMPLETED
@@ -1221,9 +1376,9 @@ async def test_rpc_client_get_sessions():
client, server_task, client_task = await _make_rpc_pair(runtime)
try:
spec = BoxSpec.model_validate({'cmd': 'echo hi', 'session_id': 'r-2'})
- await client.execute(spec)
+ await client.execute(spec, action_context=_ACTION_CONTEXT)
- sessions = await client.get_sessions()
+ sessions = await client.get_sessions(action_context=_ACTION_CONTEXT)
assert len(sessions) == 1
assert sessions[0]['session_id'] == 'r-2'
finally:
@@ -1232,6 +1387,34 @@ async def test_rpc_client_get_sessions():
await runtime.shutdown()
+@pytest.mark.asyncio
+async def test_rpc_generation_advance_retires_old_session_and_rejects_old_context():
+ logger = Mock()
+ backend = FakeBackend(logger)
+ runtime = BoxRuntime(logger=logger, backends=[backend], session_ttl_sec=300)
+ await runtime.initialize()
+ second_context = _ACTION_CONTEXT.model_copy(update={'placement_generation': 2})
+
+ client, server_task, client_task = await _make_rpc_pair(runtime)
+ try:
+ spec = BoxSpec.model_validate({'cmd': 'echo generation', 'session_id': 'shared'})
+ await client.execute(spec, action_context=_ACTION_CONTEXT)
+ await client.execute(spec, action_context=second_context)
+
+ assert len(backend.start_calls) == 2
+ assert backend.start_calls[0] != backend.start_calls[1]
+ assert backend.stop_calls == [backend.start_calls[0]]
+ assert [session['session_id'] for session in await client.get_sessions(action_context=second_context)] == [
+ 'shared'
+ ]
+ with pytest.raises(BoxError, match='Stale Box placement generation'):
+ await client.execute(spec, action_context=_ACTION_CONTEXT)
+ finally:
+ server_task.cancel()
+ client_task.cancel()
+ await runtime.shutdown()
+
+
@pytest.mark.asyncio
async def test_rpc_client_get_status():
logger = Mock()
@@ -1241,7 +1424,7 @@ async def test_rpc_client_get_status():
client, server_task, client_task = await _make_rpc_pair(runtime)
try:
- status = await client.get_status()
+ status = await client.get_status(action_context=_ACTION_CONTEXT)
assert 'backend' in status
assert 'active_sessions' in status
@@ -1283,11 +1466,11 @@ async def test_rpc_client_delete_session():
client, server_task, client_task = await _make_rpc_pair(runtime)
try:
spec = BoxSpec.model_validate({'cmd': 'echo hi', 'session_id': 'r-del-1'})
- await client.execute(spec)
+ await client.execute(spec, action_context=_ACTION_CONTEXT)
- await client.delete_session('r-del-1')
+ await client.delete_session('r-del-1', action_context=_ACTION_CONTEXT)
- sessions = await client.get_sessions()
+ sessions = await client.get_sessions(action_context=_ACTION_CONTEXT)
assert len(sessions) == 0
finally:
server_task.cancel()
@@ -1305,7 +1488,7 @@ async def test_rpc_client_delete_session_raises_not_found():
client, server_task, client_task = await _make_rpc_pair(runtime)
try:
with pytest.raises(BoxSessionNotFoundError):
- await client.delete_session('nonexistent')
+ await client.delete_session('nonexistent', action_context=_ACTION_CONTEXT)
finally:
server_task.cancel()
client_task.cancel()
@@ -1322,11 +1505,11 @@ async def test_rpc_client_create_session():
client, server_task, client_task = await _make_rpc_pair(runtime)
try:
spec = BoxSpec.model_validate({'cmd': 'placeholder', 'session_id': 'r-create-1'})
- info = await client.create_session(spec)
+ info = await client.create_session(spec, action_context=_ACTION_CONTEXT)
assert info['session_id'] == 'r-create-1'
assert info['backend_name'] == 'fake'
- sessions = await client.get_sessions()
+ sessions = await client.get_sessions(action_context=_ACTION_CONTEXT)
assert len(sessions) == 1
finally:
server_task.cancel()
@@ -1344,11 +1527,11 @@ async def test_rpc_client_exec_raises_conflict_error():
client, server_task, client_task = await _make_rpc_pair(runtime)
try:
spec1 = BoxSpec.model_validate({'cmd': 'echo first', 'session_id': 'r-conflict-1', 'network': 'off'})
- await client.execute(spec1)
+ await client.execute(spec1, action_context=_ACTION_CONTEXT)
spec2 = BoxSpec.model_validate({'cmd': 'echo second', 'session_id': 'r-conflict-1', 'network': 'on'})
with pytest.raises(BoxSessionConflictError):
- await client.execute(spec2)
+ await client.execute(spec2, action_context=_ACTION_CONTEXT)
finally:
server_task.cancel()
client_task.cancel()
@@ -1447,7 +1630,7 @@ class TestBoxDisabledByConfig:
service = BoxService(make_app(logger, enabled=False), client=Mock(spec=BoxRuntimeClient))
await service.initialize()
- status = await service.get_status()
+ status = await service.get_status(_CONTEXT)
assert status['available'] is False
assert status['enabled'] is False
@@ -1462,7 +1645,7 @@ class TestBoxDisabledByConfig:
await service.initialize()
- status = await service.get_status()
+ status = await service.get_status(_CONTEXT)
assert status['available'] is False
assert status['enabled'] is True
assert 'docker daemon' in status['connector_error']
@@ -1486,7 +1669,7 @@ class TestBoxDisabledByConfig:
service = BoxService(make_app(logger, enabled=True), client=client)
await service.initialize()
- status = await service.get_status()
+ status = await service.get_status(_CONTEXT)
assert status['available'] is False
assert status['enabled'] is True
# The detailed backend object is preserved for the dialog
@@ -1507,7 +1690,7 @@ class TestBoxDisabledByConfig:
service = BoxService(make_app(logger, enabled=True), client=client)
await service.initialize()
- status = await service.get_status()
+ status = await service.get_status(_CONTEXT)
assert status['available'] is True
assert status['backend'] == {'name': 'docker', 'available': True}
# No spurious connector_error overlay when everything is healthy
@@ -1537,7 +1720,7 @@ class TestBuildSkillExtraMounts:
def _make_service(self, logger, skills, *, shares_filesystem=True):
app = make_app(logger)
- app.skill_mgr = SimpleNamespace(skills=skills)
+ app.skill_mgr = SimpleNamespace(skills=skills, get_skills=Mock(return_value=skills))
client = Mock(spec=BoxRuntimeClient)
service = BoxService(app, client=client)
# Tests construct BoxService with an injected client (no connector), so
@@ -1795,14 +1978,17 @@ class TestAttachmentHostPath:
"""
def _service_with_workspace(self, tmp_path):
- ws = str(tmp_path / 'box' / 'default')
- os.makedirs(ws, exist_ok=True)
+ default_workspace = str(tmp_path / 'box' / 'default')
+ os.makedirs(default_workspace, exist_ok=True)
app = make_app(Mock(), allowed_mount_roots=[str(tmp_path)], host_root=str(tmp_path / 'box'))
service = BoxService(app, client=Mock(spec=BoxRuntimeClient))
service._available = True
# Force the default_workspace to our tmp dir so _host_query_dir resolves.
- service.default_workspace = ws
- return service, ws
+ service.default_workspace = default_workspace
+ tenant_workspace = service._tenant_workspace(_CONTEXT)
+ assert tenant_workspace is not None
+ os.makedirs(tenant_workspace, exist_ok=True)
+ return service, tenant_workspace
@pytest.mark.asyncio
async def test_inbound_writes_to_host_no_exec(self, tmp_path):
@@ -1923,9 +2109,10 @@ class TestAttachmentHostPath:
assert os.path.isdir(ws)
@pytest.mark.asyncio
- async def test_purge_attachment_dirs_falls_back_to_exec_for_root_owned(self, tmp_path, monkeypatch):
- # When the host delete cannot remove a dir (root-owned container output),
- # purge must fall back to deleting from inside the sandbox via exec.
+ async def test_purge_attachment_dirs_never_uses_unscoped_exec_for_root_owned(self, tmp_path, monkeypatch):
+ # Startup has no trusted Workspace context. If host deletion cannot
+ # remove root-owned output, cleanup must fail closed instead of issuing
+ # an unscoped Box exec that could cross a tenant boundary.
service, ws = self._service_with_workspace(tmp_path)
outbox = os.path.join(ws, 'outbox')
os.makedirs(os.path.join(outbox, '0'), exist_ok=True)
@@ -1942,18 +2129,17 @@ class TestAttachmentHostPath:
monkeypatch.setattr(_shutil, 'rmtree', fake_rmtree)
- executed = {}
- spec_obj = object()
- service.build_spec = Mock(return_value=spec_obj)
- service.client.execute = AsyncMock(side_effect=lambda s: executed.setdefault('spec', s))
+ service.build_spec = Mock()
+ service.client.execute = AsyncMock()
await service._purge_attachment_dirs()
- # build_spec was asked to rm the surviving outbox via exec.
- cmd = service.build_spec.call_args.args[0]['cmd']
- assert 'rm -rf' in cmd and '/workspace/outbox' in cmd
- assert '/workspace/inbox' not in cmd # inbox was host-deletable
- service.client.execute.assert_awaited_once_with(spec_obj)
+ assert os.path.isdir(outbox)
+ service.build_spec.assert_not_called()
+ service.client.execute.assert_not_awaited()
+ assert any(
+ 'no trusted Workspace context' in str(call.args[0]) for call in service.ap.logger.warning.call_args_list
+ )
@pytest.mark.asyncio
async def test_purge_attachment_dirs_noop_without_workspace(self):
diff --git a/tests/unit_tests/box/test_workspace.py b/tests/unit_tests/box/test_workspace.py
index e4620ad32..1e41fb757 100644
--- a/tests/unit_tests/box/test_workspace.py
+++ b/tests/unit_tests/box/test_workspace.py
@@ -7,6 +7,7 @@ from unittest.mock import AsyncMock, Mock
import pytest
+from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.box.workspace import (
BoxWorkspaceSession,
classify_python_workspace,
@@ -16,6 +17,13 @@ from langbot.pkg.box.workspace import (
)
+_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+)
+
+
def test_rewrite_mounted_path_translates_host_prefix():
result = rewrite_mounted_path('/tmp/demo/project/app.py', '/tmp/demo/project')
assert result == '/workspace/app.py'
@@ -66,6 +74,7 @@ async def test_workspace_session_execute_for_query_uses_session_payload():
box_service = SimpleNamespace(execute_spec_payload=AsyncMock(return_value={'ok': True}))
workspace = BoxWorkspaceSession(
box_service,
+ _CONTEXT,
'skill-person_123-demo',
host_path='/tmp/project',
host_path_mode='rw',
@@ -94,6 +103,7 @@ async def test_workspace_session_start_managed_process_rewrites_command_and_args
box_service = SimpleNamespace(start_managed_process=AsyncMock(return_value={'status': 'running'}))
workspace = BoxWorkspaceSession(
box_service,
+ _CONTEXT,
'mcp-u1',
host_path='/tmp/project',
host_path_mode='ro',
@@ -106,8 +116,10 @@ async def test_workspace_session_start_managed_process_rewrites_command_and_args
)
assert result == {'status': 'running'}
- session_id = box_service.start_managed_process.await_args.args[0]
- payload = box_service.start_managed_process.await_args.args[1]
+ execution_context = box_service.start_managed_process.await_args.args[0]
+ session_id = box_service.start_managed_process.await_args.args[1]
+ payload = box_service.start_managed_process.await_args.args[2]
+ assert execution_context == _CONTEXT
assert session_id == 'mcp-u1'
assert payload == {
'command': 'python',
@@ -118,9 +130,39 @@ async def test_workspace_session_start_managed_process_rewrites_command_and_args
}
+@pytest.mark.asyncio
+async def test_workspace_session_relay_connection_keeps_execution_context():
+ box_service = SimpleNamespace(
+ get_managed_process_websocket_connection=AsyncMock(
+ return_value=(
+ 'ws://box/relay',
+ {'X-LangBot-Placement-Generation': '1'},
+ )
+ )
+ )
+ workspace = BoxWorkspaceSession(
+ box_service,
+ _CONTEXT,
+ 'mcp-shared',
+ )
+
+ connection = await workspace.get_managed_process_websocket_connection('server-a')
+
+ assert connection == (
+ 'ws://box/relay',
+ {'X-LangBot-Placement-Generation': '1'},
+ )
+ box_service.get_managed_process_websocket_connection.assert_awaited_once_with(
+ _CONTEXT,
+ 'mcp-shared',
+ 'server-a',
+ )
+
+
def test_workspace_session_build_session_payload_keeps_generic_workspace_shape():
workspace = BoxWorkspaceSession(
Mock(),
+ _CONTEXT,
'workspace-1',
host_path='/tmp/project',
host_path_mode='rw',
diff --git a/tests/unit_tests/pipeline/conftest.py b/tests/unit_tests/pipeline/conftest.py
index ce8ee7eb0..ebd3b5363 100644
--- a/tests/unit_tests/pipeline/conftest.py
+++ b/tests/unit_tests/pipeline/conftest.py
@@ -17,6 +17,7 @@ from unittest.mock import AsyncMock, Mock
# this, running a stage test in isolation triggers a circular-import error:
# stage.py → core.app → pipelinemgr → stage.stage_class (not yet bound).
import langbot.pkg.pipeline.pipelinemgr # noqa: F401
+from langbot.pkg.api.http.context import ExecutionContext
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
import langbot_plugin.api.entities.builtin.platform.message as platform_message
@@ -40,6 +41,14 @@ class MockApplication:
self.query_pool = self._create_mock_query_pool()
self.instance_config = self._create_mock_instance_config()
self.task_mgr = self._create_mock_task_manager()
+ self.workspace_service = AsyncMock()
+ self.workspace_service.get_execution_binding = AsyncMock(
+ return_value=Mock(
+ instance_uuid='test-instance',
+ workspace_uuid='test-workspace',
+ placement_generation=1,
+ )
+ )
# Skill manager is optional; PreProcessor only touches it for the
# local-agent runner. None keeps the skill-binding branch inert.
self.skill_mgr = None
@@ -83,6 +92,7 @@ class MockApplication:
query_pool.cached_queries = {}
query_pool.queries = []
query_pool.condition = AsyncMock()
+ query_pool.remove_query = AsyncMock(return_value=True)
return query_pool
def _create_mock_instance_config(self):
@@ -191,6 +201,9 @@ def sample_query(sample_message_chain, sample_message_event, mock_adapter):
# Use model_construct to bypass Pydantic validation for test purposes
query = pipeline_query.Query.model_construct(
+ instance_uuid='test-instance',
+ workspace_uuid='test-workspace',
+ placement_generation=1,
query_id='test-query-id',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -219,6 +232,17 @@ def sample_query(sample_message_chain, sample_message_event, mock_adapter):
resp_message_chain=None,
current_stage_name=None,
)
+ object.__setattr__(
+ query,
+ '_execution_context',
+ ExecutionContext(
+ instance_uuid='test-instance',
+ workspace_uuid='test-workspace',
+ placement_generation=1,
+ bot_uuid='test-bot-uuid',
+ pipeline_uuid='test-pipeline-uuid',
+ ),
+ )
return query
diff --git a/tests/unit_tests/pipeline/test_aggregator.py b/tests/unit_tests/pipeline/test_aggregator.py
index 9eab7615f..7bea148ec 100644
--- a/tests/unit_tests/pipeline/test_aggregator.py
+++ b/tests/unit_tests/pipeline/test_aggregator.py
@@ -25,6 +25,49 @@ from tests.factories import (
import langbot_plugin.api.entities.builtin.provider.session as provider_session
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.pipeline.pool import (
+ ExecutionContextMismatchError,
+ ExecutionContextRequiredError,
+ bind_execution_context,
+)
+from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError
+
+
+def execution_context(
+ workspace_uuid='workspace-test',
+ *,
+ bot_uuid='test-bot',
+ pipeline_uuid=None,
+ placement_generation=1,
+):
+ return ExecutionContext(
+ instance_uuid='instance-test',
+ workspace_uuid=workspace_uuid,
+ placement_generation=placement_generation,
+ bot_uuid=bot_uuid,
+ pipeline_uuid=pipeline_uuid,
+ )
+
+
+def aggregation_key(
+ context,
+ *,
+ launcher_type=provider_session.LauncherTypes.PERSON,
+ launcher_id=12345,
+ bot_uuid='test-bot',
+ pipeline_uuid=None,
+):
+ return (
+ context.instance_uuid,
+ context.workspace_uuid,
+ context.placement_generation,
+ bot_uuid,
+ pipeline_uuid,
+ launcher_type.value,
+ launcher_id,
+ )
+
def get_aggregator_module():
"""Lazy import to avoid circular import issues."""
@@ -36,12 +79,66 @@ def make_aggregator_app():
app = FakeApp()
# Ensure query_pool has add_query method
app.query_pool.add_query = AsyncMock()
+
+ async def resolve_context(
+ context,
+ *,
+ bot_uuid,
+ pipeline_uuid,
+ query_uuid=None,
+ ):
+ if context is None:
+ raise ExecutionContextRequiredError('ExecutionContext required in test')
+ return bind_execution_context(
+ context,
+ bot_uuid=bot_uuid,
+ pipeline_uuid=pipeline_uuid,
+ query_uuid=query_uuid,
+ )
+
+ app.query_pool.resolve_execution_context = AsyncMock(side_effect=resolve_context)
# Add pipeline_mgr mock
app.pipeline_mgr = AsyncMock()
app.pipeline_mgr.get_pipeline_by_uuid = AsyncMock(return_value=None)
+ app.workspace_service = Mock()
+ app.workspace_service.get_execution_binding = AsyncMock(
+ return_value=Mock(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-test',
+ placement_generation=1,
+ )
+ )
return app
+def enable_aggregation(app, *, delay=10.0):
+ pipeline = Mock()
+ pipeline.pipeline_entity.config = {
+ 'trigger': {
+ 'message-aggregation': {
+ 'enabled': True,
+ 'delay': delay,
+ }
+ }
+ }
+ app.pipeline_mgr.get_pipeline_by_uuid = AsyncMock(return_value=pipeline)
+
+
+def scoped_message_kwargs(context, *, launcher_id=12345, text='hello'):
+ chain = text_chain(text)
+ return {
+ 'execution_context': context,
+ 'bot_uuid': context.bot_uuid,
+ 'launcher_type': provider_session.LauncherTypes.PERSON,
+ 'launcher_id': launcher_id,
+ 'sender_id': launcher_id,
+ 'message_event': friend_message_event(chain),
+ 'message_chain': chain,
+ 'adapter': mock_adapter(),
+ 'pipeline_uuid': context.pipeline_uuid,
+ }
+
+
class TestPendingMessage:
"""Tests for PendingMessage dataclass."""
@@ -54,6 +151,7 @@ class TestPendingMessage:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
+ execution_context=execution_context(pipeline_uuid='test-pipeline'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -77,9 +175,14 @@ class TestSessionBuffer:
"""SessionBuffer should be created with correct fields."""
aggregator = get_aggregator_module()
- buffer = aggregator.SessionBuffer(session_id='test-session')
+ context = execution_context()
+ key = aggregation_key(context)
+ buffer = aggregator.SessionBuffer(
+ aggregation_key=key,
+ execution_context=context,
+ )
- assert buffer.session_id == 'test-session'
+ assert buffer.aggregation_key == key
assert buffer.messages == []
assert buffer.timer_task is None
assert buffer.last_message_time is not None
@@ -93,6 +196,7 @@ class TestSessionBuffer:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
+ execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -103,8 +207,10 @@ class TestSessionBuffer:
pipeline_uuid=None,
)
+ context = execution_context()
buffer = aggregator.SessionBuffer(
- session_id='test-session',
+ aggregation_key=aggregation_key(context),
+ execution_context=context,
messages=[pending],
)
@@ -127,7 +233,7 @@ class TestMessageAggregatorInit:
class TestMessageAggregatorSessionId:
- """Tests for session ID generation."""
+ """Tests for scoped aggregation key generation."""
def test_session_id_format(self):
"""Session ID should be correctly formatted."""
@@ -136,13 +242,24 @@ class TestMessageAggregatorSessionId:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
- session_id = agg._get_session_id(
+ context = execution_context()
+ session_id = agg._get_aggregation_key(
+ context,
bot_uuid='bot-123',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=45678,
+ pipeline_uuid=None,
)
- assert session_id == 'bot-123:person:45678'
+ assert session_id == (
+ 'instance-test',
+ 'workspace-test',
+ 1,
+ 'bot-123',
+ None,
+ 'person',
+ 45678,
+ )
def test_session_id_different_launchers(self):
"""Different launcher types should produce different IDs."""
@@ -151,16 +268,21 @@ class TestMessageAggregatorSessionId:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
- person_id = agg._get_session_id(
+ context = execution_context()
+ person_id = agg._get_aggregation_key(
+ context,
bot_uuid='bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=123,
+ pipeline_uuid=None,
)
- group_id = agg._get_session_id(
+ group_id = agg._get_aggregation_key(
+ context,
bot_uuid='bot',
launcher_type=provider_session.LauncherTypes.GROUP,
launcher_id=123,
+ pipeline_uuid=None,
)
assert person_id != group_id
@@ -177,7 +299,7 @@ class TestMessageAggregatorConfig:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
- enabled, delay = await agg._get_aggregation_config(None)
+ enabled, delay = await agg._get_aggregation_config(execution_context(), None)
assert enabled == False
assert delay == 1.5
@@ -191,7 +313,10 @@ class TestMessageAggregatorConfig:
app.pipeline_mgr.get_pipeline_by_uuid = AsyncMock(return_value=None)
agg = aggregator.MessageAggregator(app)
- enabled, delay = await agg._get_aggregation_config('unknown-pipeline')
+ enabled, delay = await agg._get_aggregation_config(
+ execution_context(pipeline_uuid='unknown-pipeline'),
+ 'unknown-pipeline',
+ )
assert enabled == False
assert delay == 1.5
@@ -217,7 +342,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
- enabled, delay = await agg._get_aggregation_config('test-pipeline')
+ enabled, delay = await agg._get_aggregation_config(
+ execution_context(pipeline_uuid='test-pipeline'),
+ 'test-pipeline',
+ )
assert enabled == True
assert delay == 2.0
@@ -243,7 +371,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
- enabled, delay = await agg._get_aggregation_config('test-pipeline')
+ enabled, delay = await agg._get_aggregation_config(
+ execution_context(pipeline_uuid='test-pipeline'),
+ 'test-pipeline',
+ )
assert delay == 1.0 # Clamped to minimum
@@ -268,7 +399,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
- enabled, delay = await agg._get_aggregation_config('test-pipeline')
+ enabled, delay = await agg._get_aggregation_config(
+ execution_context(pipeline_uuid='test-pipeline'),
+ 'test-pipeline',
+ )
assert delay == 10.0 # Clamped to maximum
@@ -293,7 +427,10 @@ class TestMessageAggregatorConfig:
agg = aggregator.MessageAggregator(app)
- enabled, delay = await agg._get_aggregation_config('test-pipeline')
+ enabled, delay = await agg._get_aggregation_config(
+ execution_context(pipeline_uuid='test-pipeline'),
+ 'test-pipeline',
+ )
assert delay == 1.5 # Default
@@ -314,6 +451,7 @@ class TestMessageAggregatorAddMessage:
adapter = mock_adapter()
await agg.add_message(
+ execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -353,6 +491,7 @@ class TestMessageAggregatorAddMessage:
adapter = mock_adapter()
await agg.add_message(
+ execution_context=execution_context(pipeline_uuid='test-pipeline'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -394,6 +533,7 @@ class TestMessageAggregatorAddMessage:
# Add messages up to MAX_BUFFER_MESSAGES
for i in range(aggregator.MAX_BUFFER_MESSAGES):
await agg.add_message(
+ execution_context=execution_context(pipeline_uuid='test-pipeline'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -405,7 +545,14 @@ class TestMessageAggregatorAddMessage:
)
# Buffer should be flushed (empty or no buffer)
- session_id = agg._get_session_id('test-bot', provider_session.LauncherTypes.PERSON, 12345)
+ context = execution_context(pipeline_uuid='test-pipeline')
+ session_id = agg._get_aggregation_key(
+ context,
+ 'test-bot',
+ provider_session.LauncherTypes.PERSON,
+ 12345,
+ 'test-pipeline',
+ )
assert session_id not in agg.buffers or len(agg.buffers[session_id].messages) == 0
@@ -424,6 +571,7 @@ class TestMessageAggregatorMerge:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
+ execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -451,6 +599,7 @@ class TestMessageAggregatorMerge:
adapter = mock_adapter()
pending1 = aggregator.PendingMessage(
+ execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -462,6 +611,7 @@ class TestMessageAggregatorMerge:
)
pending2 = aggregator.PendingMessage(
+ execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -492,6 +642,7 @@ class TestMessageAggregatorMerge:
adapter = mock_adapter()
pending1 = aggregator.PendingMessage(
+ execution_context=execution_context(pipeline_uuid='test-pipeline-uuid'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -504,6 +655,7 @@ class TestMessageAggregatorMerge:
)
pending2 = aggregator.PendingMessage(
+ execution_context=execution_context(pipeline_uuid='test-pipeline-uuid'),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -532,7 +684,8 @@ class TestMessageAggregatorFlush:
app = make_aggregator_app()
agg = aggregator.MessageAggregator(app)
- await agg._flush_buffer('nonexistent-session')
+ context = execution_context()
+ await agg._flush_buffer(aggregation_key(context), context)
# Should not call query_pool
assert not app.query_pool.add_query.called
@@ -550,6 +703,7 @@ class TestMessageAggregatorFlush:
adapter = mock_adapter()
pending = aggregator.PendingMessage(
+ execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -560,17 +714,57 @@ class TestMessageAggregatorFlush:
pipeline_uuid=None,
)
+ context = execution_context()
+ key = aggregation_key(context)
buffer = aggregator.SessionBuffer(
- session_id='test-session',
+ aggregation_key=key,
+ execution_context=context,
messages=[pending],
)
- agg.buffers['test-session'] = buffer
+ agg.buffers[key] = buffer
- await agg._flush_buffer('test-session')
+ await agg._flush_buffer(key, context)
assert app.query_pool.add_query.called
- assert 'test-session' not in agg.buffers
+ assert key not in agg.buffers
+
+ @pytest.mark.asyncio
+ async def test_flush_drops_buffer_when_placement_generation_is_stale(self):
+ """A debounce timer cannot enqueue work after its placement is fenced."""
+ aggregator = get_aggregator_module()
+
+ app = make_aggregator_app()
+ app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError('stale generation')
+ agg = aggregator.MessageAggregator(app)
+ context = execution_context(placement_generation=3)
+ pending = aggregator.PendingMessage(
+ execution_context=context,
+ bot_uuid='test-bot',
+ launcher_type=provider_session.LauncherTypes.PERSON,
+ launcher_id=12345,
+ sender_id=12345,
+ message_event=friend_message_event(text_chain('stale')),
+ message_chain=text_chain('stale'),
+ adapter=mock_adapter(),
+ pipeline_uuid=None,
+ )
+ key = aggregation_key(context)
+ agg.buffers[key] = aggregator.SessionBuffer(
+ aggregation_key=key,
+ execution_context=context,
+ messages=[pending],
+ )
+
+ with pytest.raises(WorkspaceGenerationMismatchError):
+ await agg._flush_buffer(key, context)
+
+ app.workspace_service.get_execution_binding.assert_awaited_once_with(
+ 'workspace-test',
+ expected_generation=3,
+ )
+ app.query_pool.add_query.assert_not_awaited()
+ assert key not in agg.buffers
class TestMessageAggregatorFlushAll:
@@ -603,6 +797,7 @@ class TestMessageAggregatorFlushAll:
# Create two buffers
pending1 = aggregator.PendingMessage(
+ execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -614,6 +809,7 @@ class TestMessageAggregatorFlushAll:
)
pending2 = aggregator.PendingMessage(
+ execution_context=execution_context(),
bot_uuid='test-bot',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=67890,
@@ -624,14 +820,131 @@ class TestMessageAggregatorFlushAll:
pipeline_uuid=None,
)
- buffer1 = aggregator.SessionBuffer(session_id='session-1', messages=[pending1])
- buffer2 = aggregator.SessionBuffer(session_id='session-2', messages=[pending2])
+ context = execution_context()
+ key1 = aggregation_key(context, launcher_id=12345)
+ key2 = aggregation_key(context, launcher_id=67890)
+ buffer1 = aggregator.SessionBuffer(
+ aggregation_key=key1,
+ execution_context=context,
+ messages=[pending1],
+ )
+ buffer2 = aggregator.SessionBuffer(
+ aggregation_key=key2,
+ execution_context=context,
+ messages=[pending2],
+ )
- agg.buffers['session-1'] = buffer1
- agg.buffers['session-2'] = buffer2
+ agg.buffers[key1] = buffer1
+ agg.buffers[key2] = buffer2
await agg.flush_all()
# Both buffers should be flushed
assert len(agg.buffers) == 0
assert app.query_pool.add_query.call_count == 2
+
+
+class TestMessageAggregatorWorkspaceIsolation:
+ """Regression coverage for fail-closed and cross-workspace behavior."""
+
+ @pytest.mark.asyncio
+ async def test_missing_execution_context_fails_closed(self):
+ app = make_aggregator_app()
+ agg = get_aggregator_module().MessageAggregator(app)
+ kwargs = scoped_message_kwargs(execution_context())
+ kwargs['execution_context'] = None
+
+ with pytest.raises(ExecutionContextRequiredError):
+ await agg.add_message(**kwargs)
+
+ app.query_pool.add_query.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_same_launcher_in_two_workspaces_uses_separate_buffers(self):
+ app = make_aggregator_app()
+ enable_aggregation(app)
+ agg = get_aggregator_module().MessageAggregator(app)
+
+ await agg.add_message(**scoped_message_kwargs(execution_context('workspace-a', pipeline_uuid='test-pipeline')))
+ await agg.add_message(**scoped_message_kwargs(execution_context('workspace-b', pipeline_uuid='test-pipeline')))
+
+ assert len(agg.buffers) == 2
+ assert {key[1] for key in agg.buffers} == {'workspace-a', 'workspace-b'}
+ await agg.flush_all()
+
+ @pytest.mark.asyncio
+ async def test_same_launcher_in_two_bots_uses_separate_buffers(self):
+ app = make_aggregator_app()
+ enable_aggregation(app)
+ agg = get_aggregator_module().MessageAggregator(app)
+
+ await agg.add_message(
+ **scoped_message_kwargs(execution_context(bot_uuid='bot-a', pipeline_uuid='test-pipeline'))
+ )
+ await agg.add_message(
+ **scoped_message_kwargs(execution_context(bot_uuid='bot-b', pipeline_uuid='test-pipeline'))
+ )
+
+ assert len(agg.buffers) == 2
+ assert {key[3] for key in agg.buffers} == {'bot-a', 'bot-b'}
+ await agg.flush_all()
+
+ @pytest.mark.asyncio
+ async def test_timer_receives_exact_captured_execution_context(self, monkeypatch):
+ app = make_aggregator_app()
+ enable_aggregation(app)
+ agg = get_aggregator_module().MessageAggregator(app)
+ delayed_flush = AsyncMock()
+ monkeypatch.setattr(agg, '_delayed_flush', delayed_flush)
+ context = execution_context(pipeline_uuid='test-pipeline')
+
+ await agg.add_message(**scoped_message_kwargs(context))
+ await asyncio.sleep(0)
+
+ delayed_flush.assert_awaited_once()
+ assert delayed_flush.await_args.args[2] is context
+ await agg.flush_all()
+
+ @pytest.mark.asyncio
+ async def test_flush_rejects_context_from_another_workspace(self):
+ app = make_aggregator_app()
+ enable_aggregation(app)
+ agg = get_aggregator_module().MessageAggregator(app)
+ context_a = execution_context('workspace-a', pipeline_uuid='test-pipeline')
+ context_b = execution_context('workspace-b', pipeline_uuid='test-pipeline')
+ await agg.add_message(**scoped_message_kwargs(context_a))
+ key = next(iter(agg.buffers))
+
+ with pytest.raises(ExecutionContextMismatchError):
+ await agg._flush_buffer(key, context_b)
+
+ assert key in agg.buffers
+ await agg.flush_all()
+
+ def test_merge_rejects_messages_from_different_workspaces(self):
+ app = make_aggregator_app()
+ agg = get_aggregator_module().MessageAggregator(app)
+ aggregator = get_aggregator_module()
+
+ with pytest.raises(ExecutionContextMismatchError):
+ agg._merge_messages(
+ [
+ aggregator.PendingMessage(**scoped_message_kwargs(execution_context('workspace-a'))),
+ aggregator.PendingMessage(**scoped_message_kwargs(execution_context('workspace-b'))),
+ ]
+ )
+
+ @pytest.mark.asyncio
+ async def test_flush_all_preserves_each_workspace_context(self):
+ app = make_aggregator_app()
+ enable_aggregation(app)
+ agg = get_aggregator_module().MessageAggregator(app)
+ await agg.add_message(**scoped_message_kwargs(execution_context('workspace-a', pipeline_uuid='test-pipeline')))
+ await agg.add_message(**scoped_message_kwargs(execution_context('workspace-b', pipeline_uuid='test-pipeline')))
+
+ await agg.flush_all()
+
+ forwarded_workspaces = {
+ call.kwargs['execution_context'].workspace_uuid for call in app.query_pool.add_query.await_args_list
+ }
+ assert forwarded_workspaces == {'workspace-a', 'workspace-b'}
diff --git a/tests/unit_tests/pipeline/test_chat_session_limit.py b/tests/unit_tests/pipeline/test_chat_session_limit.py
index ef351b29f..1739c9623 100644
--- a/tests/unit_tests/pipeline/test_chat_session_limit.py
+++ b/tests/unit_tests/pipeline/test_chat_session_limit.py
@@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, Mock
import pytest
import yaml
+from langbot_plugin.api.entities.builtin.provider import session as provider_session
def _preproc_module():
@@ -43,7 +44,12 @@ def _prompt_preprocessing_context(default_prompt=None, prompt=None):
async def _run_preprocessor(mock_app, sample_query, conversation):
- session = SimpleNamespace(launcher_type=sample_query.launcher_type, launcher_id=sample_query.launcher_id)
+ session = provider_session.Session(
+ launcher_type=sample_query.launcher_type,
+ launcher_id=sample_query.launcher_id,
+ sender_id=sample_query.sender_id,
+ bot_uuid=sample_query.bot_uuid,
+ )
mock_app.sess_mgr.get_session = AsyncMock(return_value=session)
mock_app.sess_mgr.get_conversation = AsyncMock(return_value=conversation)
mock_app.plugin_connector.emit_event = AsyncMock(return_value=_prompt_preprocessing_context())
diff --git a/tests/unit_tests/pipeline/test_controller_tenancy.py b/tests/unit_tests/pipeline/test_controller_tenancy.py
new file mode 100644
index 000000000..c5669455b
--- /dev/null
+++ b/tests/unit_tests/pipeline/test_controller_tenancy.py
@@ -0,0 +1,65 @@
+from __future__ import annotations
+
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, MagicMock, Mock
+
+import pytest
+
+from langbot.pkg.pipeline.controller import Controller
+from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError
+
+
+def _prepare_scheduler(mock_app):
+ query_pool = MagicMock()
+ query_pool.remove_query = AsyncMock(return_value=True)
+ query_pool.__aenter__ = AsyncMock(return_value=query_pool)
+ query_pool.__aexit__ = AsyncMock(return_value=None)
+ query_pool.condition = SimpleNamespace(notify_all=Mock())
+ mock_app.query_pool = query_pool
+
+ session = SimpleNamespace(_semaphore=SimpleNamespace(release=Mock()))
+ mock_app.sess_mgr.get_session = AsyncMock(return_value=session)
+ mock_app.pipeline_mgr = SimpleNamespace(get_pipeline_by_uuid=AsyncMock())
+ return query_pool, session
+
+
+@pytest.mark.asyncio
+async def test_controller_drops_stale_query_before_pipeline_lookup(
+ mock_app,
+ sample_query,
+):
+ query_pool, session = _prepare_scheduler(mock_app)
+ mock_app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError('stale generation')
+ controller = Controller(mock_app)
+
+ await controller._process_query(sample_query)
+
+ mock_app.workspace_service.get_execution_binding.assert_awaited_once_with(
+ 'test-workspace',
+ expected_generation=1,
+ )
+ mock_app.pipeline_mgr.get_pipeline_by_uuid.assert_not_awaited()
+ query_pool.remove_query.assert_awaited_once_with(sample_query)
+ session._semaphore.release.assert_called_once_with()
+ query_pool.condition.notify_all.assert_called_once_with()
+
+
+@pytest.mark.asyncio
+async def test_controller_revalidates_generation_before_running_pipeline(
+ mock_app,
+ sample_query,
+):
+ query_pool, session = _prepare_scheduler(mock_app)
+ runtime_pipeline = SimpleNamespace(run=AsyncMock())
+ mock_app.pipeline_mgr.get_pipeline_by_uuid.return_value = runtime_pipeline
+ controller = Controller(mock_app)
+
+ await controller._process_query(sample_query)
+
+ mock_app.workspace_service.get_execution_binding.assert_awaited_once_with(
+ 'test-workspace',
+ expected_generation=1,
+ )
+ runtime_pipeline.run.assert_awaited_once_with(sample_query)
+ query_pool.remove_query.assert_awaited_once_with(sample_query)
+ session._semaphore.release.assert_called_once_with()
diff --git a/tests/unit_tests/pipeline/test_pipeline_service.py b/tests/unit_tests/pipeline/test_pipeline_service.py
index b862c3ff4..305d640b0 100644
--- a/tests/unit_tests/pipeline/test_pipeline_service.py
+++ b/tests/unit_tests/pipeline/test_pipeline_service.py
@@ -5,6 +5,9 @@ import pytest
from langbot.pkg.api.http.service.pipeline import PipelineService
+WORKSPACE_UUID = 'workspace-a'
+
+
@pytest.mark.asyncio
async def test_update_pipeline_filters_protected_fields_without_mutating_input(mock_app):
service = PipelineService(mock_app)
@@ -27,7 +30,7 @@ async def test_update_pipeline_filters_protected_fields_without_mutating_input(m
}
original_pipeline_data = pipeline_data.copy()
- await service.update_pipeline('pipeline-uuid', pipeline_data)
+ await service.update_pipeline(WORKSPACE_UUID, 'pipeline-uuid', pipeline_data)
assert pipeline_data == original_pipeline_data
@@ -36,8 +39,9 @@ async def test_update_pipeline_filters_protected_fields_without_mutating_input(m
assert updated_fields == {'name'}
mock_app.bot_service.update_bot.assert_awaited_once_with(
+ WORKSPACE_UUID,
'bot-uuid',
{'use_pipeline_name': 'Updated pipeline'},
)
- mock_app.pipeline_mgr.remove_pipeline.assert_awaited_once_with('pipeline-uuid')
- mock_app.pipeline_mgr.load_pipeline.assert_awaited_once_with(loaded_pipeline)
+ mock_app.pipeline_mgr.remove_pipeline.assert_awaited_once_with('workspace-a', 'pipeline-uuid')
+ mock_app.pipeline_mgr.load_pipeline.assert_awaited_once_with('workspace-a', loaded_pipeline)
diff --git a/tests/unit_tests/pipeline/test_pipelinemgr.py b/tests/unit_tests/pipeline/test_pipelinemgr.py
index 49984542c..2033df402 100644
--- a/tests/unit_tests/pipeline/test_pipelinemgr.py
+++ b/tests/unit_tests/pipeline/test_pipelinemgr.py
@@ -6,6 +6,18 @@ import pytest
from unittest.mock import AsyncMock, Mock
from importlib import import_module
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError
+
+
+def _context(pipeline_uuid: str = 'test-uuid') -> ExecutionContext:
+ return ExecutionContext(
+ instance_uuid='test-instance',
+ workspace_uuid='test-workspace',
+ placement_generation=1,
+ pipeline_uuid=pipeline_uuid,
+ )
+
def get_pipelinemgr_module():
return import_module('langbot.pkg.pipeline.pipelinemgr')
@@ -51,11 +63,12 @@ async def test_load_pipeline(mock_app):
# Create test pipeline entity
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.uuid = 'test-uuid'
+ pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.stages = []
pipeline_entity.config = {'test': 'config'}
pipeline_entity.extensions_preferences = {'plugins': []}
- await manager.load_pipeline(pipeline_entity)
+ await manager.load_pipeline(_context(), pipeline_entity)
assert len(manager.pipelines) == 1
assert manager.pipelines[0].pipeline_entity.uuid == 'test-uuid'
@@ -75,19 +88,20 @@ async def test_get_pipeline_by_uuid(mock_app):
# Create and add test pipeline
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.uuid = 'test-uuid'
+ pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.stages = []
pipeline_entity.config = {}
pipeline_entity.extensions_preferences = {'plugins': []}
- await manager.load_pipeline(pipeline_entity)
+ await manager.load_pipeline(_context(), pipeline_entity)
# Test retrieval
- result = await manager.get_pipeline_by_uuid('test-uuid')
+ result = await manager.get_pipeline_by_uuid(_context(), 'test-uuid')
assert result is not None
assert result.pipeline_entity.uuid == 'test-uuid'
# Test non-existent UUID
- result = await manager.get_pipeline_by_uuid('non-existent')
+ result = await manager.get_pipeline_by_uuid(_context('non-existent'), 'non-existent')
assert result is None
@@ -105,15 +119,16 @@ async def test_remove_pipeline(mock_app):
# Create and add test pipeline
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
pipeline_entity.uuid = 'test-uuid'
+ pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.stages = []
pipeline_entity.config = {}
pipeline_entity.extensions_preferences = {'plugins': []}
- await manager.load_pipeline(pipeline_entity)
+ await manager.load_pipeline(_context(), pipeline_entity)
assert len(manager.pipelines) == 1
# Remove pipeline
- await manager.remove_pipeline('test-uuid')
+ await manager.remove_pipeline(_context(), 'test-uuid')
assert len(manager.pipelines) == 0
@@ -143,25 +158,104 @@ async def test_runtime_pipeline_execute(mock_app, sample_query):
# Create pipeline entity
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
+ pipeline_entity.uuid = 'test-pipeline-uuid'
+ pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.config = sample_query.pipeline_config
pipeline_entity.extensions_preferences = {'plugins': []}
# Create runtime pipeline
- runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [stage_container])
+ runtime_pipeline = pipelinemgr.RuntimePipeline(
+ mock_app,
+ pipeline_entity,
+ [stage_container],
+ _context('test-pipeline-uuid'),
+ )
# Mock plugin connector
event_ctx = Mock()
event_ctx.is_prevented_default = Mock(return_value=False)
mock_app.plugin_connector.emit_event = AsyncMock(return_value=event_ctx)
- # Add query to cached_queries to prevent KeyError in finally block
- mock_app.query_pool.cached_queries[sample_query.query_id] = sample_query
-
# Execute pipeline
await runtime_pipeline.run(sample_query)
# Verify stage was called
mock_stage.process.assert_called_once()
+ mock_app.query_pool.remove_query.assert_awaited_once_with(sample_query)
+
+
+@pytest.mark.asyncio
+async def test_runtime_pipeline_rejects_stale_generation_before_side_effects(
+ mock_app,
+ sample_query,
+):
+ pipelinemgr = get_pipelinemgr_module()
+ persistence_pipeline = get_persistence_pipeline_module()
+ pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
+ pipeline_entity.uuid = 'test-pipeline-uuid'
+ pipeline_entity.workspace_uuid = 'test-workspace'
+ pipeline_entity.config = sample_query.pipeline_config
+ pipeline_entity.extensions_preferences = {'plugins': []}
+ runtime_pipeline = pipelinemgr.RuntimePipeline(
+ mock_app,
+ pipeline_entity,
+ [],
+ _context('test-pipeline-uuid'),
+ )
+ mock_app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError('stale generation')
+
+ with pytest.raises(WorkspaceGenerationMismatchError):
+ await runtime_pipeline.run(sample_query)
+
+ mock_app.plugin_connector.emit_event.assert_not_awaited()
+ sample_query.adapter.reply_message.assert_not_awaited()
+ sample_query.adapter.reply_message_chunk.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_runtime_pipeline_revalidates_after_awaited_stage(
+ mock_app,
+ sample_query,
+):
+ pipelinemgr = get_pipelinemgr_module()
+ stage = get_stage_module()
+ persistence_pipeline = get_persistence_pipeline_module()
+ entities = get_entities_module()
+ pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
+ pipeline_entity.uuid = 'test-pipeline-uuid'
+ pipeline_entity.workspace_uuid = 'test-workspace'
+ pipeline_entity.config = sample_query.pipeline_config
+ pipeline_entity.extensions_preferences = {'plugins': []}
+
+ result = entities.StageProcessResult(
+ result_type=entities.ResultType.CONTINUE,
+ new_query=sample_query,
+ user_notice='must not be sent',
+ console_notice='',
+ debug_notice='',
+ error_notice='',
+ )
+
+ async def stage_process(*_args):
+ mock_app.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError(
+ 'generation changed during stage'
+ )
+ return result
+
+ mock_stage = Mock(spec=stage.PipelineStage)
+ mock_stage.process = Mock(side_effect=stage_process)
+ runtime_pipeline = pipelinemgr.RuntimePipeline(
+ mock_app,
+ pipeline_entity,
+ [pipelinemgr.StageInstContainer(inst_name='TestStage', inst=mock_stage)],
+ _context('test-pipeline-uuid'),
+ )
+
+ with pytest.raises(WorkspaceGenerationMismatchError):
+ await runtime_pipeline._execute_from_stage(0, sample_query)
+
+ sample_query.adapter.reply_message.assert_not_awaited()
+ sample_query.adapter.reply_message_chunk.assert_not_awaited()
def test_runtime_pipeline_prefers_local_agent_mcp_resources(mock_app):
@@ -170,6 +264,8 @@ def test_runtime_pipeline_prefers_local_agent_mcp_resources(mock_app):
persistence_pipeline = get_persistence_pipeline_module()
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
+ pipeline_entity.uuid = 'test-uuid'
+ pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.config = {
'ai': {
'local-agent': {
@@ -183,7 +279,7 @@ def test_runtime_pipeline_prefers_local_agent_mcp_resources(mock_app):
'mcp_resource_agent_read_enabled': True,
}
- runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [])
+ runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [], _context())
assert runtime_pipeline.mcp_resource_attachments == [{'server_uuid': 'srv-new', 'uri': 'file:///new.md'}]
assert runtime_pipeline.mcp_resource_agent_read_enabled is False
@@ -195,13 +291,15 @@ def test_runtime_pipeline_falls_back_to_extension_mcp_resources(mock_app):
persistence_pipeline = get_persistence_pipeline_module()
pipeline_entity = Mock(spec=persistence_pipeline.LegacyPipeline)
+ pipeline_entity.uuid = 'test-uuid'
+ pipeline_entity.workspace_uuid = 'test-workspace'
pipeline_entity.config = {'ai': {'local-agent': {}}}
pipeline_entity.extensions_preferences = {
'mcp_resources': [{'server_uuid': 'srv-old', 'uri': 'file:///old.md'}],
'mcp_resource_agent_read_enabled': False,
}
- runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [])
+ runtime_pipeline = pipelinemgr.RuntimePipeline(mock_app, pipeline_entity, [], _context())
assert runtime_pipeline.mcp_resource_attachments == [{'server_uuid': 'srv-old', 'uri': 'file:///old.md'}]
assert runtime_pipeline.mcp_resource_agent_read_enabled is False
diff --git a/tests/unit_tests/pipeline/test_pool.py b/tests/unit_tests/pipeline/test_pool.py
index 86515e7f0..15d4bffb9 100644
--- a/tests/unit_tests/pipeline/test_pool.py
+++ b/tests/unit_tests/pipeline/test_pool.py
@@ -6,10 +6,51 @@ Tests query management, ID generation, and async context handling.
from __future__ import annotations
+import uuid
+from types import SimpleNamespace
+
import pytest
from unittest.mock import Mock, patch
-from langbot.pkg.pipeline.pool import QueryPool
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.pipeline.pool import (
+ ExecutionContextMismatchError,
+ ExecutionContextRequiredError,
+ QueryNotFoundError,
+ QueryPool,
+ get_query_execution_context,
+)
+
+
+TEST_CONTEXT = ExecutionContext(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-test',
+ placement_generation=1,
+)
+
+
+def oss_pool():
+ """Build the explicit singleton resolver used by the OSS compatibility path."""
+ return QueryPool(singleton_context_resolver=lambda: TEST_CONTEXT)
+
+
+async def add_scoped_mock_query(pool, context, *, bot_uuid='bot-a'):
+ """Create a Query through the real pool while keeping SDK details mocked."""
+ query = Mock()
+ query.bot_uuid = bot_uuid
+ query.pipeline_uuid = None
+ query.query_id = pool.query_id_counter
+ with patch('langbot.pkg.pipeline.pool.pipeline_query.Query', return_value=query):
+ return await pool.add_query(
+ bot_uuid=bot_uuid,
+ launcher_type=Mock(),
+ launcher_id='launcher-1',
+ sender_id='sender-1',
+ message_event=Mock(),
+ message_chain=Mock(),
+ adapter=Mock(),
+ execution_context=context,
+ )
pytestmark = pytest.mark.asyncio
@@ -39,7 +80,7 @@ class TestQueryPoolAddQuery:
async def test_add_query_adds_query_with_id(self):
"""add_query creates, stores, and caches a Query with the correct ID."""
- pool = QueryPool()
+ pool = oss_pool()
# Mock Query creation
mock_query = Mock()
@@ -62,12 +103,12 @@ class TestQueryPoolAddQuery:
# Query is added to list and cache
assert pool.queries[0] is mock_query
- assert pool.cached_queries[0] is mock_query
+ assert pool.cached_queries[('workspace-test', mock_query.query_uuid)] is mock_query
assert mock_query.query_id == 0
async def test_add_query_increments_counter(self):
"""Each add_query increments the counter."""
- pool = QueryPool()
+ pool = oss_pool()
mock_query1 = Mock()
mock_query1.query_id = 0
@@ -103,7 +144,7 @@ class TestQueryPoolAddQuery:
async def test_add_query_appends_to_list(self):
"""Query is appended to queries list."""
- pool = QueryPool()
+ pool = oss_pool()
mock_query = Mock()
mock_query.query_id = 0
@@ -126,7 +167,7 @@ class TestQueryPoolAddQuery:
async def test_add_query_caches_query(self):
"""Query is cached by query_id."""
- pool = QueryPool()
+ pool = oss_pool()
mock_query = Mock()
mock_query.query_id = 0
@@ -144,12 +185,13 @@ class TestQueryPoolAddQuery:
adapter=Mock(),
)
- assert 0 in pool.cached_queries
- assert pool.cached_queries[0] is mock_query
+ cache_key = ('workspace-test', mock_query.query_uuid)
+ assert cache_key in pool.cached_queries
+ assert pool.cached_queries[cache_key] is mock_query
async def test_add_query_with_pipeline_uuid(self):
"""Query can have pipeline_uuid set."""
- pool = QueryPool()
+ pool = oss_pool()
mock_query = Mock()
mock_query.query_id = 0
@@ -175,7 +217,7 @@ class TestQueryPoolAddQuery:
async def test_add_query_sets_routed_by_rule_variable(self):
"""Query has _routed_by_rule variable."""
- pool = QueryPool()
+ pool = oss_pool()
mock_query = Mock()
mock_query.query_id = 0
@@ -201,7 +243,7 @@ class TestQueryPoolAddQuery:
async def test_add_query_notifier_condition(self):
"""add_query notifies waiting consumers."""
- pool = QueryPool()
+ pool = oss_pool()
mock_query = Mock()
mock_query.query_id = 0
@@ -237,7 +279,7 @@ class TestQueryPoolContext:
async def test_aenter_acquires_lock(self):
"""__aenter__ acquires the pool lock."""
- pool = QueryPool()
+ pool = oss_pool()
async with pool as p:
# Lock is acquired
@@ -260,7 +302,7 @@ class TestQueryPoolEdgeCases:
async def test_multiple_queries_cached_correctly(self):
"""Multiple queries are cached separately."""
- pool = QueryPool()
+ pool = oss_pool()
mock_queries = []
for i in range(5):
@@ -287,4 +329,107 @@ class TestQueryPoolEdgeCases:
# Each query is cached by its ID
for i in range(5):
- assert pool.cached_queries[i] is mock_queries[i]
+ query = mock_queries[i]
+ assert pool.cached_queries[('workspace-test', query.query_uuid)] is query
+
+
+class TestQueryPoolWorkspaceIsolation:
+ """Regression coverage for trusted scope and scoped cache indexes."""
+
+ async def test_add_query_requires_execution_context_by_default(self):
+ with pytest.raises(ExecutionContextRequiredError):
+ await QueryPool().add_query(
+ bot_uuid='bot-a',
+ launcher_type=Mock(),
+ launcher_id='launcher-1',
+ sender_id='sender-1',
+ message_event=Mock(),
+ message_chain=Mock(),
+ adapter=Mock(),
+ )
+
+ async def test_serialized_scope_fields_are_not_trusted_context(self):
+ forged_query = SimpleNamespace(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-test',
+ placement_generation=1,
+ bot_uuid='bot-a',
+ pipeline_uuid=None,
+ query_uuid='forged-query',
+ )
+
+ with pytest.raises(ExecutionContextRequiredError):
+ get_query_execution_context(forged_query)
+
+ async def test_query_lookup_is_workspace_scoped(self):
+ pool = QueryPool()
+ query = await add_scoped_mock_query(pool, TEST_CONTEXT)
+
+ uuid.UUID(query.query_uuid)
+ assert await pool.get_query('workspace-test', query.query_uuid) is query
+ assert await pool.get_query('workspace-other', query.query_uuid) is None
+ assert await pool.get_query_by_legacy_id('workspace-test', 0) is query
+ assert await pool.get_query_by_legacy_id('workspace-other', 0) is None
+ with pytest.raises(QueryNotFoundError):
+ await pool.require_query('workspace-other', query.query_uuid)
+
+ async def test_cache_separates_same_opaque_id_between_workspaces(self, monkeypatch):
+ fixed_uuid = uuid.UUID('11111111-1111-4111-8111-111111111111')
+ monkeypatch.setattr('langbot.pkg.pipeline.pool.uuid.uuid4', lambda: fixed_uuid)
+ pool = QueryPool()
+ context_a = TEST_CONTEXT
+ context_b = ExecutionContext(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-other',
+ placement_generation=1,
+ )
+
+ query_a = await add_scoped_mock_query(pool, context_a)
+ query_b = await add_scoped_mock_query(pool, context_b)
+
+ assert query_a.query_uuid == query_b.query_uuid
+ assert await pool.get_query('workspace-test', query_a.query_uuid) is query_a
+ assert await pool.get_query('workspace-other', query_b.query_uuid) is query_b
+
+ async def test_remove_query_cleans_both_scoped_indexes(self):
+ pool = QueryPool()
+ query = await add_scoped_mock_query(pool, TEST_CONTEXT)
+
+ assert await pool.remove_query(query) is True
+ assert await pool.get_query('workspace-test', query.query_uuid) is None
+ assert await pool.get_query_by_legacy_id('workspace-test', query.query_id) is None
+ assert await pool.remove_query(query) is False
+
+ async def test_context_cannot_substitute_bot_identity(self):
+ context = ExecutionContext(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-test',
+ placement_generation=1,
+ bot_uuid='bot-b',
+ )
+
+ with pytest.raises(ExecutionContextMismatchError):
+ await add_scoped_mock_query(QueryPool(), context, bot_uuid='bot-a')
+
+ async def test_query_counter_is_scoped_by_workspace_and_generation(self):
+ pool = QueryPool()
+ workspace_a = TEST_CONTEXT
+ workspace_b = ExecutionContext(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-other',
+ placement_generation=1,
+ )
+ next_generation = ExecutionContext(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-test',
+ placement_generation=2,
+ )
+
+ await add_scoped_mock_query(pool, workspace_a)
+ await add_scoped_mock_query(pool, workspace_a)
+ await add_scoped_mock_query(pool, workspace_b)
+
+ assert pool.get_query_count(workspace_a) == 2
+ assert pool.get_query_count(workspace_b) == 1
+ assert pool.get_query_count(next_generation) == 0
+ assert pool.query_id_counter == 3
diff --git a/tests/unit_tests/pipeline/test_preproc.py b/tests/unit_tests/pipeline/test_preproc.py
index 15bf60801..c28858959 100644
--- a/tests/unit_tests/pipeline/test_preproc.py
+++ b/tests/unit_tests/pipeline/test_preproc.py
@@ -16,6 +16,8 @@ from unittest.mock import AsyncMock, Mock
from importlib import import_module
from types import SimpleNamespace
+from langbot_plugin.api.entities.builtin.provider import session as provider_session
+
from tests.factories import (
FakeApp,
text_query,
@@ -35,6 +37,20 @@ def get_entities_module():
return import_module('langbot.pkg.pipeline.entities')
+def make_session(
+ launcher_type: provider_session.LauncherTypes = provider_session.LauncherTypes.PERSON,
+ launcher_id: int = 12345,
+) -> provider_session.Session:
+ """Build a scope-aware Session that matches the shared Query factory."""
+
+ return provider_session.Session(
+ launcher_type=launcher_type,
+ launcher_id=launcher_id,
+ sender_id=12345,
+ bot_uuid='test-bot-uuid',
+ )
+
+
class TestPreProcessorNormalText:
"""Tests for normal text message preprocessing."""
@@ -46,9 +62,7 @@ class TestPreProcessorNormalText:
app = FakeApp()
# Mock session manager to return a session
- mock_session = Mock()
- mock_session.launcher_type = Mock(value='person')
- mock_session.launcher_id = 12345
+ mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
# Mock conversation
@@ -92,9 +106,7 @@ class TestPreProcessorNormalText:
preproc = get_preproc_module()
app = FakeApp()
- mock_session = Mock()
- mock_session.launcher_type = Mock(value='person')
- mock_session.launcher_id = 12345
+ mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -132,9 +144,7 @@ class TestPreProcessorEmptyMessage:
entities = get_entities_module()
app = FakeApp()
- mock_session = Mock()
- mock_session.launcher_type = Mock(value='person')
- mock_session.launcher_id = 12345
+ mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -171,9 +181,7 @@ class TestPreProcessorImageSegment:
preproc = get_preproc_module()
app = FakeApp()
- mock_session = Mock()
- mock_session.launcher_type = Mock(value='person')
- mock_session.launcher_id = 12345
+ mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -219,9 +227,7 @@ class TestPreProcessorImageSegment:
preproc = get_preproc_module()
app = FakeApp()
- mock_session = Mock()
- mock_session.launcher_type = Mock(value='person')
- mock_session.launcher_id = 12345
+ mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -258,9 +264,7 @@ class TestPreProcessorModelSelection:
preproc = get_preproc_module()
app = FakeApp()
- mock_session = Mock()
- mock_session.launcher_type = Mock(value='person')
- mock_session.launcher_id = 12345
+ mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -305,9 +309,7 @@ class TestPreProcessorModelSelection:
preproc = get_preproc_module()
app = FakeApp()
- mock_session = Mock()
- mock_session.launcher_type = Mock(value='person')
- mock_session.launcher_id = 12345
+ mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -324,7 +326,7 @@ class TestPreProcessorModelSelection:
mock_fallback = Mock()
mock_fallback.model_entity = Mock(uuid='fallback-uuid', abilities=['func_call'])
- async def mock_get_model(uuid):
+ async def mock_get_model(_context, uuid):
if uuid == 'primary-uuid':
return mock_primary
elif uuid == 'fallback-uuid':
@@ -368,9 +370,7 @@ class TestPreProcessorVariables:
preproc = get_preproc_module()
app = FakeApp()
- mock_session = Mock()
- mock_session.launcher_type = Mock(value='person')
- mock_session.launcher_id = 12345
+ mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -405,9 +405,10 @@ class TestPreProcessorVariables:
preproc = get_preproc_module()
app = FakeApp()
- mock_session = Mock()
- mock_session.launcher_type = Mock(value='group')
- mock_session.launcher_id = 99999
+ mock_session = make_session(
+ provider_session.LauncherTypes.GROUP,
+ 99999,
+ )
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
@@ -443,9 +444,7 @@ class TestPreProcessorToolSelection:
preproc = get_preproc_module()
app = FakeApp()
- mock_session = Mock()
- mock_session.launcher_type = Mock(value='person')
- mock_session.launcher_id = 12345
+ mock_session = make_session()
app.sess_mgr.get_session = AsyncMock(return_value=mock_session)
mock_conversation = Mock()
diff --git a/tests/unit_tests/pipeline/test_query_pool.py b/tests/unit_tests/pipeline/test_query_pool.py
index 228be093e..a5c526372 100644
--- a/tests/unit_tests/pipeline/test_query_pool.py
+++ b/tests/unit_tests/pipeline/test_query_pool.py
@@ -9,6 +9,7 @@ import langbot_plugin.api.definition.abstract.platform.adapter as abstract_platf
import langbot_plugin.api.definition.abstract.platform.event_logger as abstract_platform_logger
from langbot.pkg.pipeline.pool import QueryPool
+from langbot.pkg.api.http.context import ExecutionContext
class DummyEventLogger(abstract_platform_logger.AbstractEventLogger):
@@ -64,12 +65,18 @@ async def test_add_query_returns_created_query_and_preserves_side_effects(
adapter=adapter,
pipeline_uuid='test-pipeline-uuid',
routed_by_rule=True,
+ execution_context=ExecutionContext(
+ instance_uuid='test-instance-uuid',
+ workspace_uuid='test-workspace-uuid',
+ placement_generation=1,
+ ),
)
assert query is query_pool.queries[0]
- assert query_pool.cached_queries[0] is query
+ assert query_pool.cached_queries[('test-workspace-uuid', query.query_uuid)] is query
assert query_pool.query_id_counter == 1
assert query.query_id == 0
assert query.bot_uuid == 'test-bot-uuid'
assert query.pipeline_uuid == 'test-pipeline-uuid'
+ assert query.workspace_uuid == 'test-workspace-uuid'
assert query.variables == {'_routed_by_rule': True}
diff --git a/tests/unit_tests/platform/test_botmgr_tenancy.py b/tests/unit_tests/platform/test_botmgr_tenancy.py
new file mode 100644
index 000000000..fb2c61f86
--- /dev/null
+++ b/tests/unit_tests/platform/test_botmgr_tenancy.py
@@ -0,0 +1,114 @@
+from __future__ import annotations
+
+from types import SimpleNamespace
+
+import pytest
+
+from langbot.pkg.api.http.authz import WorkspaceRequiredError
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.platform.botmgr import PlatformManager, RuntimeBot
+
+
+WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
+WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
+BOT_A = '10000000-0000-0000-0000-00000000000a'
+BOT_B = '10000000-0000-0000-0000-00000000000b'
+
+
+def _context(workspace_uuid: str, bot_uuid: str, generation: int = 4) -> ExecutionContext:
+ return ExecutionContext(
+ instance_uuid='instance',
+ workspace_uuid=workspace_uuid,
+ placement_generation=generation,
+ bot_uuid=bot_uuid,
+ )
+
+
+def _runtime(application, workspace_uuid: str, bot_uuid: str) -> RuntimeBot:
+ entity = SimpleNamespace(
+ uuid=bot_uuid,
+ workspace_uuid=workspace_uuid,
+ name='Same Name',
+ enable=True,
+ pipeline_routing_rules=[],
+ use_pipeline_uuid=None,
+ )
+ return RuntimeBot(
+ ap=application,
+ bot_entity=entity,
+ adapter=SimpleNamespace(),
+ logger=SimpleNamespace(),
+ execution_context=_context(workspace_uuid, bot_uuid),
+ )
+
+
+class _WorkspaceService:
+ async def get_execution_binding(self, workspace_uuid, expected_generation=None):
+ if workspace_uuid not in {WORKSPACE_A, WORKSPACE_B} or expected_generation != 4:
+ raise ValueError('stale')
+ return SimpleNamespace(
+ instance_uuid='instance',
+ workspace_uuid=workspace_uuid,
+ placement_generation=4,
+ )
+
+
+@pytest.fixture
+def manager():
+ application = SimpleNamespace(workspace_service=_WorkspaceService())
+ platform_manager = PlatformManager(application)
+ platform_manager.bots = [
+ _runtime(application, WORKSPACE_A, BOT_A),
+ _runtime(application, WORKSPACE_B, BOT_B),
+ ]
+ return platform_manager
+
+
+@pytest.mark.asyncio
+async def test_runtime_lookup_cannot_guess_another_workspace_bot(manager):
+ assert await manager.get_bot_by_uuid(_context(WORKSPACE_A, BOT_A), BOT_A) is manager.bots[0]
+ assert await manager.get_bot_by_uuid(_context(WORKSPACE_B, BOT_A), BOT_A) is None
+ assert await manager.get_bot_by_uuid(_context(WORKSPACE_A, BOT_B), BOT_B) is None
+
+
+@pytest.mark.asyncio
+async def test_public_route_key_resolves_bound_runtime_and_rejects_non_opaque_input(manager):
+ assert await manager.resolve_public_bot(BOT_A) is manager.bots[0]
+ assert await manager.resolve_public_bot('Same Name') is None
+ assert await manager.resolve_public_bot('not-a-uuid') is None
+
+
+@pytest.mark.asyncio
+async def test_stale_runtime_generation_is_not_returned(manager):
+ assert await manager.get_bot_by_uuid(_context(WORKSPACE_A, BOT_A, generation=5), BOT_A) is None
+
+
+@pytest.mark.asyncio
+async def test_runtime_bot_revalidates_its_generation_before_handling_events(manager):
+ runtime_bot = manager.bots[0]
+
+ await runtime_bot.assert_execution_active()
+
+ runtime_bot.placement_generation = 5
+ with pytest.raises(ValueError, match='stale'):
+ await runtime_bot.assert_execution_active()
+
+
+def test_runtime_bot_rejects_workspace_mismatch():
+ application = SimpleNamespace()
+ entity = SimpleNamespace(
+ uuid=BOT_A,
+ workspace_uuid=WORKSPACE_A,
+ name='Bot',
+ enable=True,
+ pipeline_routing_rules=[],
+ use_pipeline_uuid=None,
+ )
+ with pytest.raises(WorkspaceRequiredError):
+ RuntimeBot(
+ ap=application,
+ bot_entity=entity,
+ adapter=SimpleNamespace(),
+ logger=SimpleNamespace(),
+ execution_context=_context(WORKSPACE_B, BOT_A),
+ )
diff --git a/tests/unit_tests/platform/test_http_bot_tenancy.py b/tests/unit_tests/platform/test_http_bot_tenancy.py
new file mode 100644
index 000000000..b00b2364d
--- /dev/null
+++ b/tests/unit_tests/platform/test_http_bot_tenancy.py
@@ -0,0 +1,66 @@
+from __future__ import annotations
+
+from types import SimpleNamespace
+
+import pytest
+
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.platform.sources.http_bot import HttpBotAdapter
+
+
+def _session(key):
+ session = SimpleNamespace()
+ session._langbot_session_key = key
+ return session
+
+
+def _adapter(app, execution_context) -> HttpBotAdapter:
+ adapter = HttpBotAdapter.model_construct(
+ config={'signature_required': False},
+ logger=SimpleNamespace(execution_context=execution_context),
+ bot_uuid='bot-a',
+ outbound_states={},
+ idempotency_cache={},
+ sync_waiters={},
+ )
+ object.__setattr__(adapter, 'ap', app)
+ return adapter
+
+
+@pytest.mark.asyncio
+async def test_http_bot_reset_removes_only_exact_execution_scope():
+ context = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=3,
+ bot_uuid='bot-a',
+ )
+ target_key = ('instance-a', 'workspace-a', 3, 'bot-a', 'person', 'shared-session')
+ retained_keys = [
+ ('instance-b', 'workspace-a', 3, 'bot-a', 'person', 'shared-session'),
+ ('instance-a', 'workspace-b', 3, 'bot-a', 'person', 'shared-session'),
+ ('instance-a', 'workspace-a', 4, 'bot-a', 'person', 'shared-session'),
+ ('instance-a', 'workspace-a', 3, 'bot-b', 'person', 'shared-session'),
+ ('instance-a', 'workspace-a', 3, 'bot-a', 'group', 'shared-session'),
+ ('instance-a', 'workspace-a', 3, 'bot-a', 'person', 'other-session'),
+ ]
+ sessions = [_session(target_key), *[_session(key) for key in retained_keys], SimpleNamespace()]
+ app = SimpleNamespace(sess_mgr=SimpleNamespace(session_list=sessions))
+ adapter = _adapter(app, context)
+
+ removed = await adapter._reset_session('person', 'shared-session')
+
+ assert removed is True
+ assert [getattr(session, '_langbot_session_key', None) for session in app.sess_mgr.session_list] == [
+ *retained_keys,
+ None,
+ ]
+
+
+@pytest.mark.asyncio
+async def test_http_bot_reset_fails_closed_without_trusted_scope():
+ app = SimpleNamespace(sess_mgr=SimpleNamespace(session_list=[]))
+ adapter = _adapter(app, None)
+
+ with pytest.raises(RuntimeError, match='trusted execution scope'):
+ await adapter._reset_session('person', 'shared-session')
diff --git a/tests/unit_tests/platform/test_openclaw_weixin_tenancy.py b/tests/unit_tests/platform/test_openclaw_weixin_tenancy.py
new file mode 100644
index 000000000..3bb1298a6
--- /dev/null
+++ b/tests/unit_tests/platform/test_openclaw_weixin_tenancy.py
@@ -0,0 +1,84 @@
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, Mock
+
+import pytest
+
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.platform.sources.openclaw_weixin import OpenClawWeixinAdapter
+
+
+def make_adapter(*, execution_context: ExecutionContext | None):
+ app = SimpleNamespace(
+ persistence_mgr=SimpleNamespace(execute_async=AsyncMock()),
+ workspace_service=SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ )
+ )
+ ),
+ )
+ logger = SimpleNamespace(
+ ap=app,
+ execution_context=execution_context,
+ warning=AsyncMock(),
+ )
+ adapter = OpenClawWeixinAdapter.model_construct(
+ config={'token': 'refreshed-token'},
+ logger=logger,
+ client=Mock(),
+ bot_account_id='',
+ listeners={},
+ name='openclaw-weixin',
+ )
+ adapter._bot_uuid = 'shared-bot-uuid'
+ return adapter, app, logger
+
+
+@pytest.mark.asyncio
+async def test_persist_config_scopes_duplicate_bot_uuid_to_workspace():
+ adapter, app, _ = make_adapter(
+ execution_context=ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ bot_uuid='shared-bot-uuid',
+ )
+ )
+
+ await adapter._persist_config()
+
+ app.workspace_service.get_execution_binding.assert_awaited_once_with(
+ 'workspace-a',
+ expected_generation=1,
+ )
+ statement = app.persistence_mgr.execute_async.await_args.args[0]
+ params = statement.compile().params
+ assert 'workspace-a' in params.values()
+ assert 'shared-bot-uuid' in params.values()
+ assert {'workspace_uuid', 'uuid'} <= {comparison.left.name for comparison in statement._where_criteria}
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ 'execution_context',
+ [
+ None,
+ ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ bot_uuid='another-bot-uuid',
+ ),
+ ],
+ ids=['missing-context', 'mismatched-bot'],
+)
+async def test_persist_config_fails_closed_without_matching_execution_context(execution_context):
+ adapter, app, logger = make_adapter(execution_context=execution_context)
+
+ await adapter._persist_config()
+
+ app.persistence_mgr.execute_async.assert_not_awaited()
+ logger.warning.assert_awaited_once()
diff --git a/tests/unit_tests/platform/test_routing_rules.py b/tests/unit_tests/platform/test_routing_rules.py
index 3928f6f11..359eab5b7 100644
--- a/tests/unit_tests/platform/test_routing_rules.py
+++ b/tests/unit_tests/platform/test_routing_rules.py
@@ -278,3 +278,19 @@ class TestResolvePipelineUuid:
uuid, routed = bot.resolve_pipeline_uuid('person', '123', 'normal message')
assert uuid == 'default-uuid'
assert routed is False
+
+ def test_websocket_task_override_does_not_mutate_bot_default(self):
+ bot = self._make_bot('default-uuid', [])
+ adapter = Mock()
+ adapter.get_pipeline_uuid_override.return_value = 'connection-pipeline'
+
+ pipeline_uuid, routed = bot.resolve_event_pipeline_uuid(
+ adapter,
+ 'person',
+ 'launcher',
+ 'hello',
+ )
+
+ assert pipeline_uuid == 'connection-pipeline'
+ assert routed is False
+ assert bot.bot_entity.use_pipeline_uuid == 'default-uuid'
diff --git a/tests/unit_tests/platform/test_websocket_adapter_attachments.py b/tests/unit_tests/platform/test_websocket_adapter_attachments.py
index 18138383d..438e4c5c6 100644
--- a/tests/unit_tests/platform/test_websocket_adapter_attachments.py
+++ b/tests/unit_tests/platform/test_websocket_adapter_attachments.py
@@ -4,89 +4,114 @@ The web debug client uploads Image / Voice / File components carrying a storage
key in ``path``. This helper resolves each to a base64 data URI (so multimodal
LLM input and the Box sandbox inbox have usable bytes), then deletes the
consumed storage object and clears ``path``. Covers mimetype selection per
-type and graceful error handling.
+type and fail-closed error handling.
"""
from __future__ import annotations
import base64
+from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
import pytest
+from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.platform.sources.websocket_adapter import WebSocketAdapter
+_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ pipeline_uuid='pipeline-a',
+)
+_UPLOAD_PREFIX = 'v1/instance-a/workspace-a/1/upload_image/'
+
+
+def _make_connection():
+ return SimpleNamespace(execution_context=_CONTEXT)
+
+
def _make_adapter(load_return=b'hello', load_side_effect=None):
provider = Mock()
provider.load = AsyncMock(return_value=load_return, side_effect=load_side_effect)
- provider.delete = AsyncMock()
+ storage_mgr = Mock()
+ storage_mgr.storage_provider = provider
+ storage_mgr.scoped_prefix.return_value = _UPLOAD_PREFIX
+ storage_mgr.is_scoped_object_key.return_value = True
+ storage_mgr.delete_scoped_object_key = AsyncMock()
ap = Mock()
- ap.storage_mgr.storage_provider = provider
+ ap.storage_mgr = storage_mgr
logger = Mock()
logger.error = AsyncMock()
+ logger.warning = AsyncMock()
# WebSocketAdapter is a pydantic model; bypass full __init__/validation.
adapter = WebSocketAdapter.model_construct(ap=ap, logger=logger)
- return adapter, provider
+ return adapter, storage_mgr, provider
@pytest.mark.asyncio
async def test_image_jpeg_mimetype_and_cleanup():
- adapter, provider = _make_adapter(load_return=b'\xff\xd8\xff')
- chain = [{'type': 'Image', 'path': 'storage://abc/photo.jpg'}]
+ adapter, storage_mgr, _ = _make_adapter(load_return=b'\xff\xd8\xff')
+ path = f'{_UPLOAD_PREFIX}photo.jpg'
+ chain = [{'type': 'Image', 'path': path}]
- await adapter._process_image_components(chain)
+ await adapter._process_image_components(_make_connection(), chain)
expected_b64 = base64.b64encode(b'\xff\xd8\xff').decode('utf-8')
assert chain[0]['base64'] == f'data:image/jpeg;base64,{expected_b64}'
assert chain[0]['path'] == '' # consumed
- provider.delete.assert_awaited_once_with('storage://abc/photo.jpg')
+ storage_mgr.delete_scoped_object_key.assert_awaited_once_with(
+ _CONTEXT,
+ path,
+ expected_owner_type='upload_image',
+ )
@pytest.mark.asyncio
async def test_image_defaults_to_png():
- adapter, _ = _make_adapter()
- chain = [{'type': 'Image', 'path': 'storage://abc/blob'}]
- await adapter._process_image_components(chain)
+ adapter, _, _ = _make_adapter()
+ chain = [{'type': 'Image', 'path': f'{_UPLOAD_PREFIX}blob'}]
+ await adapter._process_image_components(_make_connection(), chain)
assert chain[0]['base64'].startswith('data:image/png;base64,')
@pytest.mark.asyncio
async def test_voice_uses_guessed_or_wav_mimetype():
- adapter, _ = _make_adapter()
- chain = [{'type': 'Voice', 'path': 'storage://abc/clip.wav'}]
- await adapter._process_image_components(chain)
+ adapter, _, _ = _make_adapter()
+ chain = [{'type': 'Voice', 'path': f'{_UPLOAD_PREFIX}clip.wav'}]
+ await adapter._process_image_components(_make_connection(), chain)
assert chain[0]['base64'].startswith('data:audio/')
@pytest.mark.asyncio
async def test_file_uses_octet_stream_fallback():
- adapter, _ = _make_adapter()
- chain = [{'type': 'File', 'path': 'storage://abc/unknownblob'}]
- await adapter._process_image_components(chain)
+ adapter, _, _ = _make_adapter()
+ chain = [{'type': 'File', 'path': f'{_UPLOAD_PREFIX}unknownblob'}]
+ await adapter._process_image_components(_make_connection(), chain)
assert chain[0]['base64'].startswith('data:application/octet-stream;base64,')
@pytest.mark.asyncio
async def test_skips_components_without_path_or_unknown_type():
- adapter, provider = _make_adapter()
+ adapter, _, provider = _make_adapter()
chain = [
{'type': 'Image', 'path': ''}, # no path
{'type': 'Plain', 'path': 'storage://abc/x'}, # not a file component
{'type': 'At', 'target': '123'}, # no path key at all
]
- await adapter._process_image_components(chain)
+ await adapter._process_image_components(_make_connection(), chain)
provider.load.assert_not_awaited()
assert 'base64' not in chain[0]
assert 'base64' not in chain[1]
@pytest.mark.asyncio
-async def test_load_failure_is_logged_not_raised():
- adapter, _ = _make_adapter(load_side_effect=RuntimeError('storage down'))
- chain = [{'type': 'File', 'path': 'storage://abc/doc.pdf'}]
+async def test_load_failure_is_logged_and_aborts_processing():
+ adapter, _, _ = _make_adapter(load_side_effect=RuntimeError('storage down'))
+ chain = [{'type': 'File', 'path': f'{_UPLOAD_PREFIX}doc.pdf'}]
- # must not raise
- await adapter._process_image_components(chain)
+ with pytest.raises(RuntimeError, match='storage down'):
+ await adapter._process_image_components(_make_connection(), chain)
assert 'base64' not in chain[0]
adapter.logger.error.assert_awaited_once()
diff --git a/tests/unit_tests/platform/test_websocket_session_isolation.py b/tests/unit_tests/platform/test_websocket_session_isolation.py
index d86580000..110e6e0c3 100644
--- a/tests/unit_tests/platform/test_websocket_session_isolation.py
+++ b/tests/unit_tests/platform/test_websocket_session_isolation.py
@@ -9,7 +9,25 @@ import pytest
import langbot_plugin.api.entities.builtin.platform.events as platform_events
from langbot.pkg.platform.sources import websocket_adapter as websocket_adapter_module
from langbot.pkg.platform.sources.websocket_adapter import WebSocketAdapter, WebSocketMessage, WebSocketSession
-from langbot.pkg.platform.sources.websocket_manager import WebSocketConnectionManager, is_valid_session_id
+from langbot.pkg.platform.sources.websocket_manager import (
+ WebSocketConnectionManager,
+ WebSocketScope,
+ is_valid_session_id,
+)
+
+
+SCOPE_A = WebSocketScope('instance-a', 'workspace-a', 1)
+SCOPE_B = WebSocketScope('instance-a', 'workspace-b', 1)
+
+
+def _adapter_logger(scope: WebSocketScope = SCOPE_A):
+ logger = AsyncMock()
+ logger.execution_context = Mock(
+ instance_uuid=scope.instance_uuid,
+ workspace_uuid=scope.workspace_uuid,
+ placement_generation=scope.placement_generation,
+ )
+ return logger
@pytest.mark.asyncio
@@ -17,18 +35,21 @@ async def test_broadcast_only_reaches_connections_in_same_browser_session():
manager = WebSocketConnectionManager()
first = await manager.add_connection(
websocket=Mock(),
+ scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
session_id='session-a',
)
second = await manager.add_connection(
websocket=Mock(),
+ scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
session_id='session-b',
)
dashboard = await manager.add_connection(
websocket=Mock(),
+ scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
)
@@ -36,6 +57,7 @@ async def test_broadcast_only_reaches_connections_in_same_browser_session():
await manager.broadcast_to_pipeline(
'pipeline-1',
{'type': 'response'},
+ scope=SCOPE_A,
session_type='person',
session_id='session-a',
)
@@ -47,6 +69,7 @@ async def test_broadcast_only_reaches_connections_in_same_browser_session():
await manager.broadcast_to_pipeline(
'pipeline-1',
{'type': 'dashboard-response'},
+ scope=SCOPE_A,
session_type='person',
session_id=None,
)
@@ -56,19 +79,49 @@ async def test_broadcast_only_reaches_connections_in_same_browser_session():
assert second.send_queue.empty()
+@pytest.mark.asyncio
+async def test_pipeline_indexes_and_broadcasts_are_workspace_scoped():
+ manager = WebSocketConnectionManager()
+ workspace_a = await manager.add_connection(
+ websocket=Mock(),
+ scope=SCOPE_A,
+ pipeline_uuid='shared-pipeline',
+ session_type='person',
+ )
+ workspace_b = await manager.add_connection(
+ websocket=Mock(),
+ scope=SCOPE_B,
+ pipeline_uuid='shared-pipeline',
+ session_type='person',
+ )
+
+ await manager.broadcast_to_pipeline(
+ 'shared-pipeline',
+ {'type': 'workspace-a'},
+ scope=SCOPE_A,
+ )
+
+ assert await workspace_a.send_queue.get() == {'type': 'workspace-a'}
+ assert workspace_b.send_queue.empty()
+ assert await manager.get_connection(workspace_b.connection_id, scope=SCOPE_A) is None
+ assert await manager.get_connection(workspace_b.connection_id, scope=SCOPE_B) is workspace_b
+ assert manager.get_stats(scope=SCOPE_A)['total_connections'] == 1
+
+
@pytest.mark.asyncio
async def test_embed_event_uses_stable_session_launcher(monkeypatch):
manager = WebSocketConnectionManager()
session_id = '31c0f2e9-b115-4ee6-8f15-3e624d6456b1'
connection = await manager.add_connection(
websocket=Mock(),
+ scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
session_id=session_id,
)
monkeypatch.setattr(websocket_adapter_module, 'ws_connection_manager', manager)
- adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=AsyncMock())
+ adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=_adapter_logger())
adapter.websocket_person_session = WebSocketSession(id='person')
adapter.websocket_group_session = WebSocketSession(id='group')
received = []
@@ -92,13 +145,14 @@ async def test_embed_group_event_uses_stable_session_launcher(monkeypatch):
session_id = '31c0f2e9-b115-4ee6-8f15-3e624d6456b1'
connection = await manager.add_connection(
websocket=Mock(),
+ scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='group',
session_id=session_id,
)
monkeypatch.setattr(websocket_adapter_module, 'ws_connection_manager', manager)
- adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=AsyncMock())
+ adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=_adapter_logger())
adapter.websocket_person_session = WebSocketSession(id='person')
adapter.websocket_group_session = WebSocketSession(id='group')
received = []
@@ -118,6 +172,7 @@ async def test_embed_group_event_uses_stable_session_launcher(monkeypatch):
dashboard = await manager.add_connection(
websocket=Mock(),
+ scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='group',
)
@@ -138,30 +193,46 @@ async def test_stable_session_launcher_resolves_to_active_connection(monkeypatch
session_id = '31c0f2e9-b115-4ee6-8f15-3e624d6456b1'
await manager.add_connection(
websocket=Mock(),
+ scope=SCOPE_A,
pipeline_uuid='pipeline-2',
session_type='person',
session_id=session_id,
)
connection = await manager.add_connection(
websocket=Mock(),
+ scope=SCOPE_A,
pipeline_uuid='pipeline-1',
session_type='person',
session_id=session_id,
)
monkeypatch.setattr(websocket_adapter_module, 'ws_connection_manager', manager)
- adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=AsyncMock())
+ adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=_adapter_logger())
message_source = Mock()
message_source.sender.id = f'websocket_pipeline-1:{session_id}'
assert await adapter._get_message_context(message_source) == ('pipeline-1', session_id)
assert await adapter._get_connection_from_target(f'websocketgroup_pipeline-1:{session_id}') is connection
- assert await manager.get_connection_by_session_id(session_id, 'pipeline-1') is connection
+ assert (
+ await manager.get_connection_by_session_id(
+ session_id,
+ scope=SCOPE_A,
+ pipeline_uuid='pipeline-1',
+ )
+ is connection
+ )
await manager.remove_connection(connection.connection_id)
assert await adapter._get_message_context(message_source) == ('pipeline-1', session_id)
- assert await manager.get_connection_by_session_id(session_id, 'pipeline-1') is None
+ assert (
+ await manager.get_connection_by_session_id(
+ session_id,
+ scope=SCOPE_A,
+ pipeline_uuid='pipeline-1',
+ )
+ is None
+ )
def test_session_ids_must_be_canonical_random_uuids():
@@ -171,7 +242,7 @@ def test_session_ids_must_be_canonical_random_uuids():
def test_history_read_does_not_allocate_unknown_session():
- adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=AsyncMock())
+ adapter = WebSocketAdapter.model_construct(ap=Mock(), logger=_adapter_logger())
adapter.websocket_person_session = WebSocketSession(id='person')
adapter.websocket_group_session = WebSocketSession(id='group')
@@ -179,16 +250,70 @@ def test_history_read_does_not_allocate_unknown_session():
assert adapter.websocket_person_session.message_lists == {}
+@pytest.mark.asyncio
+async def test_attachment_key_must_belong_to_connection_upload_scope():
+ manager = WebSocketConnectionManager()
+ connection = await manager.add_connection(
+ websocket=Mock(),
+ scope=SCOPE_A,
+ pipeline_uuid='pipeline-1',
+ session_type='person',
+ )
+ storage_mgr = Mock()
+ storage_mgr.scoped_prefix.return_value = 'v1/current/upload_image/'
+ storage_mgr.is_scoped_object_key.return_value = True
+ storage_mgr.storage_provider.load = AsyncMock(return_value=b'image')
+ storage_mgr.delete_scoped_object_key = AsyncMock()
+ adapter = WebSocketAdapter.model_construct(
+ ap=Mock(storage_mgr=storage_mgr),
+ logger=_adapter_logger(),
+ )
+ message_chain = [{'type': 'Image', 'path': 'v1/current/upload_image/key.png'}]
+
+ await adapter._process_image_components(connection, message_chain)
+
+ assert message_chain[0]['base64'].startswith('data:image/png;base64,')
+ assert message_chain[0]['path'] == ''
+ storage_mgr.scoped_prefix.assert_called_once_with(
+ connection.execution_context,
+ owner_type='upload_image',
+ )
+ storage_mgr.is_scoped_object_key.assert_called_once_with(
+ 'v1/current/upload_image/key.png',
+ expected_owner_type='upload_image',
+ )
+ storage_mgr.delete_scoped_object_key.assert_awaited_once_with(
+ connection.execution_context,
+ 'v1/current/upload_image/key.png',
+ expected_owner_type='upload_image',
+ )
+
+ with pytest.raises(ValueError, match='does not belong'):
+ await adapter._process_image_components(
+ connection,
+ [{'type': 'File', 'path': 'v1/other/upload/key.txt'}],
+ )
+
+
def test_history_and_reset_are_scoped_to_browser_session():
matching_provider_session = Mock(
+ instance_uuid=SCOPE_A.instance_uuid,
+ workspace_uuid=SCOPE_A.workspace_uuid,
+ placement_generation=SCOPE_A.placement_generation,
launcher_type=Mock(value='person'),
launcher_id='websocket_pipeline-1:session-a',
)
matching_group_provider_session = Mock(
+ instance_uuid=SCOPE_A.instance_uuid,
+ workspace_uuid=SCOPE_A.workspace_uuid,
+ placement_generation=SCOPE_A.placement_generation,
launcher_type=Mock(value='group'),
launcher_id='websocketgroup_pipeline-1:session-a',
)
other_session = Mock(
+ instance_uuid=SCOPE_A.instance_uuid,
+ workspace_uuid=SCOPE_A.workspace_uuid,
+ placement_generation=SCOPE_A.placement_generation,
launcher_type=Mock(value='person'),
launcher_id='websocket_pipeline-1:session-b',
)
@@ -200,7 +325,7 @@ def test_history_and_reset_are_scoped_to_browser_session():
]
adapter = WebSocketAdapter.model_construct(
ap=ap,
- logger=AsyncMock(),
+ logger=_adapter_logger(),
)
adapter.websocket_person_session = Mock()
adapter.websocket_group_session = Mock()
diff --git a/tests/unit_tests/plugin/test_connector_methods.py b/tests/unit_tests/plugin/test_connector_methods.py
index 34cab5271..9839767f8 100644
--- a/tests/unit_tests/plugin/test_connector_methods.py
+++ b/tests/unit_tests/plugin/test_connector_methods.py
@@ -14,6 +14,14 @@ from unittest.mock import Mock, AsyncMock
from importlib import import_module
from tests.factories import text_query
+from langbot_plugin.entities.io.context import ActionContext
+
+
+TEST_ACTION_CONTEXT = ActionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+)
def get_connector_module():
@@ -69,6 +77,7 @@ class TestListPlugins:
connector = create_mock_connector()
connector.handler = AsyncMock()
+ connector.handler.require_bound_action_context = Mock(return_value=TEST_ACTION_CONTEXT)
connector.handler.list_plugins = AsyncMock(
return_value=[{'manifest': {'manifest': {'metadata': {'author': 'test', 'name': 'plugin'}}}}]
)
@@ -77,6 +86,8 @@ class TestListPlugins:
connector.handler.list_plugins.assert_called_once()
assert result == [{'manifest': {'manifest': {'metadata': {'author': 'test', 'name': 'plugin'}}}}]
+ statement = connector.ap.persistence_mgr.execute_async.await_args.args[0]
+ assert 'workspace-a' in statement.compile().params.values()
@pytest.mark.asyncio
async def test_filters_by_component_kinds(self):
@@ -85,6 +96,7 @@ class TestListPlugins:
connector = create_mock_connector()
connector.handler = AsyncMock()
+ connector.handler.require_bound_action_context = Mock(return_value=TEST_ACTION_CONTEXT)
connector.handler.list_plugins = AsyncMock(
return_value=[
{
@@ -112,6 +124,7 @@ class TestListPlugins:
connector = create_mock_connector()
connector.handler = AsyncMock()
+ connector.handler.require_bound_action_context = Mock(return_value=TEST_ACTION_CONTEXT)
connector.handler.list_plugins = AsyncMock(
return_value=[
{
diff --git a/tests/unit_tests/plugin/test_connector_ping.py b/tests/unit_tests/plugin/test_connector_ping.py
index 766e51f8d..f3251c8ab 100644
--- a/tests/unit_tests/plugin/test_connector_ping.py
+++ b/tests/unit_tests/plugin/test_connector_ping.py
@@ -5,7 +5,13 @@ from unittest.mock import AsyncMock
import pytest
+from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.plugin.connector import PluginRuntimeConnector, PluginRuntimeNotConnectedError
+from langbot_plugin.entities.io.context import ActionContext
+from langbot_plugin.runtime.security import (
+ PLUGIN_RUNTIME_CONTROL_TOKEN_ENV,
+ PLUGIN_RUNTIME_CONTROL_TOKEN_HEADER,
+)
def make_connector() -> PluginRuntimeConnector:
@@ -30,3 +36,193 @@ async def test_ping_plugin_runtime_delegates_to_connected_handler():
assert result == 'pong'
connector.handler.ping.assert_awaited_once()
+
+
+@pytest.mark.asyncio
+async def test_disabled_connector_validates_workspace_without_runtime_handler():
+ app = SimpleNamespace(
+ instance_config=SimpleNamespace(data={'plugin': {'enable': False}}),
+ workspace_service=SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=3,
+ )
+ )
+ ),
+ )
+ connector = PluginRuntimeConnector(app, AsyncMock())
+ request_context = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=3,
+ )
+
+ result = await connector.require_workspace_context(request_context)
+
+ assert result == request_context
+ app.workspace_service.get_execution_binding.assert_awaited_once_with(
+ 'workspace-a',
+ expected_generation=3,
+ )
+
+
+@pytest.mark.asyncio
+async def test_enabled_connector_reports_not_connected_after_workspace_validation():
+ connector = make_connector()
+ connector.ap.workspace_service = SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=3,
+ )
+ )
+ )
+ request_context = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=3,
+ )
+
+ with pytest.raises(PluginRuntimeNotConnectedError, match='Plugin runtime is not connected'):
+ await connector.require_workspace_context(request_context)
+
+
+@pytest.mark.asyncio
+async def test_connector_resolves_oss_singleton_binding_from_workspace_service():
+ connector = make_connector()
+ connector.ap.workspace_service = SimpleNamespace(
+ get_local_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=3,
+ )
+ )
+ )
+
+ assert await connector._resolve_action_context() == ActionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=3,
+ )
+
+
+@pytest.mark.asyncio
+async def test_edition_metadata_cannot_enable_multi_workspace_runtime_binding():
+ connector = make_connector()
+ connector.ap.instance_config.data['system'] = {'edition': 'cloud'}
+ connector.ap.workspace_service = SimpleNamespace(
+ policy=SimpleNamespace(multi_workspace_enabled=False),
+ get_local_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ )
+ ),
+ )
+
+ context = await connector._resolve_action_context()
+
+ assert context.workspace_uuid == 'workspace-a'
+
+
+def test_external_runtime_control_headers_require_strong_secret(monkeypatch):
+ monkeypatch.delenv(PLUGIN_RUNTIME_CONTROL_TOKEN_ENV, raising=False)
+ connector = make_connector()
+
+ with pytest.raises(PluginRuntimeNotConnectedError, match=PLUGIN_RUNTIME_CONTROL_TOKEN_ENV):
+ connector._control_headers(allow_generate=False)
+
+
+def test_local_runtime_control_headers_generate_ephemeral_secret(monkeypatch):
+ monkeypatch.delenv(PLUGIN_RUNTIME_CONTROL_TOKEN_ENV, raising=False)
+ connector = make_connector()
+
+ headers = connector._control_headers(allow_generate=True)
+
+ assert len(headers[PLUGIN_RUNTIME_CONTROL_TOKEN_HEADER]) >= 32
+
+
+@pytest.mark.asyncio
+async def test_connector_fails_closed_without_trusted_workspace_service():
+ connector = make_connector()
+
+ with pytest.raises(
+ RuntimeError,
+ match='Plugin Runtime Workspace binding is unavailable',
+ ):
+ await connector._resolve_action_context()
+
+
+@pytest.mark.asyncio
+async def test_cloud_connector_never_falls_back_to_ghost_local_workspace():
+ connector = make_connector()
+ connector.ap.instance_config.data['system'] = {'edition': 'cloud'}
+ get_local_binding = AsyncMock()
+ connector.ap.workspace_service = SimpleNamespace(
+ policy=SimpleNamespace(multi_workspace_enabled=True),
+ get_local_execution_binding=get_local_binding,
+ )
+
+ with pytest.raises(
+ RuntimeError,
+ match='require an explicit projected Workspace binding',
+ ):
+ await connector._resolve_action_context()
+
+ get_local_binding.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_cloud_connector_initialize_fails_before_opening_unbound_runtime():
+ connector = make_connector()
+ connector.ap.instance_config.data['system'] = {'edition': 'cloud'}
+ connector.ap.workspace_service = SimpleNamespace(
+ policy=SimpleNamespace(multi_workspace_enabled=True),
+ )
+
+ with pytest.raises(
+ RuntimeError,
+ match='require an explicit projected Workspace binding',
+ ):
+ await connector.initialize()
+
+
+@pytest.mark.asyncio
+async def test_explicit_cloud_binding_is_revalidated_against_projection():
+ app = SimpleNamespace(
+ instance_config=SimpleNamespace(
+ data={
+ 'plugin': {'enable': True},
+ 'system': {'edition': 'cloud'},
+ }
+ )
+ )
+ configured = ActionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-cloud-a',
+ placement_generation=7,
+ )
+ app.workspace_service = SimpleNamespace(
+ policy=SimpleNamespace(multi_workspace_enabled=True),
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-cloud-a',
+ placement_generation=7,
+ )
+ ),
+ get_local_execution_binding=AsyncMock(),
+ )
+ connector = PluginRuntimeConnector(app, AsyncMock(), action_context=configured)
+
+ assert await connector._resolve_action_context() == configured
+ app.workspace_service.get_execution_binding.assert_awaited_once_with(
+ 'workspace-cloud-a',
+ expected_generation=7,
+ )
+ app.workspace_service.get_local_execution_binding.assert_not_awaited()
diff --git a/tests/unit_tests/plugin/test_handler.py b/tests/unit_tests/plugin/test_handler.py
index a2fdddd33..cfcd40a8b 100644
--- a/tests/unit_tests/plugin/test_handler.py
+++ b/tests/unit_tests/plugin/test_handler.py
@@ -10,13 +10,56 @@ from unittest.mock import AsyncMock, MagicMock, Mock
import pytest
from langbot_plugin.entities.io.actions.enums import PluginToRuntimeAction
+from langbot_plugin.entities.io.context import ActionContext
def make_handler(app):
"""Create a RuntimeConnectionHandler with mocked external connection."""
from langbot.pkg.plugin.handler import RuntimeConnectionHandler
- return RuntimeConnectionHandler(Mock(), AsyncMock(return_value=True), app)
+ workspace_context = ActionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ )
+ app.workspace_service = SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid=workspace_context.instance_uuid,
+ workspace_uuid=workspace_context.workspace_uuid,
+ placement_generation=workspace_context.placement_generation,
+ )
+ )
+ )
+ runtime_handler = RuntimeConnectionHandler(
+ Mock(),
+ AsyncMock(return_value=True),
+ app,
+ workspace_context,
+ )
+ installation_uuid = runtime_handler._remember_installation(
+ workspace_context,
+ 'test-author',
+ 'test-plugin',
+ )
+ runtime_handler.bind_action_context(workspace_context.for_installation(installation_uuid))
+ query_pool = getattr(app, 'query_pool', None)
+ if query_pool is not None and hasattr(query_pool, 'cached_queries'):
+
+ def scoped_query(query):
+ if query is not None:
+ query.instance_uuid = workspace_context.instance_uuid
+ query.workspace_uuid = workspace_context.workspace_uuid
+ query.placement_generation = workspace_context.placement_generation
+ return query
+
+ query_pool.get_query = AsyncMock(
+ side_effect=lambda workspace_uuid, query_uuid: scoped_query(query_pool.cached_queries.get(query_uuid))
+ )
+ query_pool.get_query_by_legacy_id = AsyncMock(
+ side_effect=lambda workspace_uuid, query_id: scoped_query(query_pool.cached_queries.get(query_id))
+ )
+ return runtime_handler
class TestHandlerQueryVariables:
diff --git a/tests/unit_tests/plugin/test_handler_actions.py b/tests/unit_tests/plugin/test_handler_actions.py
index 2dbc1f4f9..e9a2b2863 100644
--- a/tests/unit_tests/plugin/test_handler_actions.py
+++ b/tests/unit_tests/plugin/test_handler_actions.py
@@ -8,13 +8,75 @@ from unittest.mock import AsyncMock, Mock
import pytest
from langbot_plugin.entities.io.actions.enums import PluginToRuntimeAction, RuntimeToLangBotAction
+from langbot_plugin.entities.io.context import ActionContext
+
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.storage.mgr import StorageMgr
-def make_handler(app):
+TEST_EXECUTION_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+)
+
+
+def canonical_binary_key(owner_type: str, owner: str, key: str) -> str:
+ return StorageMgr.canonical_binary_storage_key(
+ TEST_EXECUTION_CONTEXT,
+ owner_type=owner_type,
+ owner=owner,
+ key=key,
+ )
+
+
+def make_handler(app, workspace_context: ActionContext | None = None):
"""Create a RuntimeConnectionHandler with mocked external connection."""
from langbot.pkg.plugin.handler import RuntimeConnectionHandler
- return RuntimeConnectionHandler(Mock(), AsyncMock(return_value=True), app)
+ workspace_context = workspace_context or ActionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ )
+ app.workspace_service = SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid=workspace_context.instance_uuid,
+ workspace_uuid=workspace_context.workspace_uuid,
+ placement_generation=workspace_context.placement_generation,
+ )
+ )
+ )
+ runtime_handler = RuntimeConnectionHandler(
+ Mock(),
+ AsyncMock(return_value=True),
+ app,
+ workspace_context,
+ )
+ installation_uuid = runtime_handler._remember_installation(
+ workspace_context,
+ 'test-author',
+ 'test-plugin',
+ )
+ runtime_handler.bind_action_context(workspace_context.for_installation(installation_uuid))
+ query_pool = getattr(app, 'query_pool', None)
+ if query_pool is not None and hasattr(query_pool, 'cached_queries'):
+
+ def scoped_query(query):
+ if query is not None:
+ query.instance_uuid = workspace_context.instance_uuid
+ query.workspace_uuid = workspace_context.workspace_uuid
+ query.placement_generation = workspace_context.placement_generation
+ return query
+
+ query_pool.get_query = AsyncMock(
+ side_effect=lambda workspace_uuid, query_uuid: scoped_query(query_pool.cached_queries.get(query_uuid))
+ )
+ query_pool.get_query_by_legacy_id = AsyncMock(
+ side_effect=lambda workspace_uuid, query_id: scoped_query(query_pool.cached_queries.get(query_id))
+ )
+ return runtime_handler
def make_result(first_item=None):
@@ -34,6 +96,8 @@ class TestRagRerankAction:
def app(self):
mock_app = Mock()
mock_app.model_mgr = Mock()
+ mock_app.persistence_mgr = Mock()
+ mock_app.persistence_mgr.execute_async = AsyncMock(return_value=make_result(SimpleNamespace(uuid='rerank-1')))
mock_app.logger = Mock()
return mock_app
@@ -63,12 +127,16 @@ class TestRagRerankAction:
assert response.code == 0
assert response.data['results'] == [{'index': 1, 'relevance_score': 0.9}]
- app.model_mgr.get_rerank_model_by_uuid.assert_awaited_once_with('rerank-1')
+ app.model_mgr.get_rerank_model_by_uuid.assert_awaited_once_with(
+ TEST_EXECUTION_CONTEXT,
+ 'rerank-1',
+ )
provider.invoke_rerank.assert_awaited_once_with(
model=rerank_model,
query='hello',
documents=['a', 'b'],
extra_args={'return_documents': False},
+ execution_context=TEST_EXECUTION_CONTEXT,
)
@pytest.mark.asyncio
@@ -122,6 +190,7 @@ class TestInitializePluginSettings:
assert app.persistence_mgr.execute_async.await_count == 2
insert_params = compiled_params(app.persistence_mgr.execute_async.await_args_list[1].args[0])
assert insert_params == {
+ 'workspace_uuid': 'workspace-a',
'plugin_author': 'test-author',
'plugin_name': 'test-plugin',
'install_source': 'local',
@@ -218,7 +287,12 @@ class TestSetBinaryStorage:
assert response.code == 0
assert app.persistence_mgr.execute_async.await_count == 2
insert_params = compiled_params(app.persistence_mgr.execute_async.await_args_list[1].args[0])
- assert insert_params['unique_key'] == 'plugin:test-owner:test-key'
+ assert insert_params['workspace_uuid'] == 'workspace-a'
+ assert insert_params['unique_key'] == canonical_binary_key(
+ 'plugin',
+ 'test-author/test-plugin',
+ 'test-key',
+ )
assert insert_params['value'] == b'x' * 512
@pytest.mark.asyncio
@@ -231,7 +305,15 @@ class TestSetBinaryStorage:
assert response.code == 0
assert app.persistence_mgr.execute_async.await_count == 2
+ select_params = compiled_params(app.persistence_mgr.execute_async.await_args_list[0].args[0])
update_params = compiled_params(app.persistence_mgr.execute_async.await_args_list[1].args[0])
+ expected_key = canonical_binary_key(
+ 'plugin',
+ 'test-author/test-plugin',
+ 'test-key',
+ )
+ assert expected_key in select_params.values()
+ assert expected_key in update_params.values()
assert update_params['value'] == b'new'
@pytest.mark.asyncio
@@ -304,6 +386,11 @@ class TestGetPluginSettings:
'plugin_config': {},
'install_source': 'local',
'install_info': {},
+ 'installation_uuid': runtime_handler.derive_installation_uuid(
+ runtime_handler.require_bound_action_context(),
+ 'test-author',
+ 'test-plugin',
+ ),
}
@pytest.mark.asyncio
@@ -333,9 +420,90 @@ class TestGetPluginSettings:
'plugin_config': {'custom': 'config'},
'install_source': 'github',
'install_info': {'repo': 'test/repo'},
+ 'installation_uuid': runtime_handler.derive_installation_uuid(
+ runtime_handler.require_bound_action_context(),
+ 'test-author',
+ 'test-plugin',
+ ),
}
+class TestGetConfigFile:
+ """Plugin config files remain bound to the trusted runtime placement."""
+
+ WORKSPACE_A = '11111111-1111-4111-8111-111111111111'
+ WORKSPACE_B = '22222222-2222-4222-8222-222222222222'
+
+ @pytest.fixture
+ def app(self):
+ mock_app = Mock()
+ mock_app.persistence_mgr = Mock()
+ mock_app.persistence_mgr.execute_async = AsyncMock()
+ mock_app.storage_mgr = StorageMgr(mock_app)
+ mock_app.storage_mgr.storage_provider = SimpleNamespace(load=AsyncMock(return_value=b'plugin config bytes'))
+ mock_app.logger = Mock()
+ return mock_app
+
+ @staticmethod
+ def object_key(
+ *,
+ workspace_uuid: str = WORKSPACE_A,
+ placement_generation: int = 1,
+ owner_type: str = 'plugin_config',
+ ) -> str:
+ return StorageMgr.scoped_object_key(
+ ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid=workspace_uuid,
+ placement_generation=placement_generation,
+ ),
+ owner_type=owner_type,
+ owner=workspace_uuid,
+ key='config.json',
+ )
+
+ async def invoke(self, app, file_key: str):
+ runtime_handler = make_handler(
+ app,
+ ActionContext(
+ instance_uuid='instance-a',
+ workspace_uuid=self.WORKSPACE_A,
+ placement_generation=1,
+ ),
+ )
+ app.persistence_mgr.execute_async.return_value = make_result(
+ SimpleNamespace(config={'uploaded_file': file_key})
+ )
+ return await runtime_handler.actions[PluginToRuntimeAction.GET_CONFIG_FILE.value]({'file_key': file_key})
+
+ @pytest.mark.asyncio
+ async def test_loads_config_file_from_same_workspace_generation(self, app):
+ file_key = self.object_key()
+
+ response = await self.invoke(app, file_key)
+
+ assert response.code == 0
+ assert base64.b64decode(response.data['file_base64']) == b'plugin config bytes'
+ app.storage_mgr.storage_provider.load.assert_awaited_once_with(file_key)
+
+ @pytest.mark.asyncio
+ @pytest.mark.parametrize(
+ 'file_key',
+ [
+ object_key.__func__(workspace_uuid=WORKSPACE_B),
+ object_key.__func__(placement_generation=2),
+ object_key.__func__(owner_type='upload_document'),
+ ],
+ ids=['other-workspace', 'stale-generation', 'wrong-owner-type'],
+ )
+ async def test_rejects_config_file_outside_trusted_scope(self, app, file_key):
+ response = await self.invoke(app, file_key)
+
+ assert response.code != 0
+ assert 'Failed to load config file' in response.message
+ app.storage_mgr.storage_provider.load.assert_not_awaited()
+
+
class TestGetBinaryStorage:
"""Tests for get_binary_storage action handler."""
@@ -364,6 +532,16 @@ class TestGetBinaryStorage:
assert response.data == {
'value_base64': base64.b64encode(b'test binary content').decode('utf-8'),
}
+ statement_params = compiled_params(app.persistence_mgr.execute_async.await_args.args[0])
+ assert 'workspace-a' in statement_params.values()
+ assert (
+ canonical_binary_key(
+ 'plugin',
+ 'test-author/test-plugin',
+ 'test-key',
+ )
+ in statement_params.values()
+ )
@pytest.mark.asyncio
async def test_returns_error_when_not_found(self, app):
@@ -383,6 +561,63 @@ class TestGetBinaryStorage:
assert 'Storage with key test-key not found' in response.message
+class TestDeleteAndListBinaryStorage:
+ """Delete/list remain fenced to the trusted canonical owner scope."""
+
+ @pytest.fixture
+ def app(self):
+ mock_app = Mock()
+ mock_app.persistence_mgr = Mock()
+ mock_app.persistence_mgr.execute_async = AsyncMock()
+ return mock_app
+
+ @pytest.mark.asyncio
+ async def test_delete_uses_workspace_and_canonical_unique_key(self, app):
+ runtime_handler = make_handler(app)
+
+ response = await runtime_handler.actions[RuntimeToLangBotAction.DELETE_BINARY_STORAGE.value](
+ {
+ 'key': 'test-key',
+ 'owner_type': 'plugin',
+ 'owner': 'forged-owner',
+ }
+ )
+
+ assert response.code == 0
+ statement_params = compiled_params(app.persistence_mgr.execute_async.await_args.args[0])
+ assert 'workspace-a' in statement_params.values()
+ assert (
+ canonical_binary_key(
+ 'plugin',
+ 'test-author/test-plugin',
+ 'test-key',
+ )
+ in statement_params.values()
+ )
+ assert 'forged-owner' not in statement_params.values()
+
+ @pytest.mark.asyncio
+ async def test_list_keys_uses_trusted_plugin_owner(self, app):
+ result = Mock()
+ result.scalars.return_value.all.return_value = ['first', 'second']
+ app.persistence_mgr.execute_async.return_value = result
+ runtime_handler = make_handler(app)
+
+ response = await runtime_handler.actions[RuntimeToLangBotAction.GET_BINARY_STORAGE_KEYS.value](
+ {
+ 'owner_type': 'plugin',
+ 'owner': 'forged-owner',
+ }
+ )
+
+ assert response.code == 0
+ assert response.data == {'keys': ['first', 'second']}
+ statement_params = compiled_params(app.persistence_mgr.execute_async.await_args.args[0])
+ assert 'workspace-a' in statement_params.values()
+ assert 'test-author/test-plugin' in statement_params.values()
+ assert 'forged-owner' not in statement_params.values()
+
+
class TestHandlerQueryLookup:
"""Tests for query lookup in cached_queries."""
diff --git a/tests/unit_tests/plugin/test_handler_tenancy.py b/tests/unit_tests/plugin/test_handler_tenancy.py
new file mode 100644
index 000000000..ba5066fb6
--- /dev/null
+++ b/tests/unit_tests/plugin/test_handler_tenancy.py
@@ -0,0 +1,303 @@
+from __future__ import annotations
+
+import asyncio
+import json
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, Mock
+
+import pytest
+
+from langbot_plugin.entities.io.actions.enums import PluginToRuntimeAction
+from langbot_plugin.entities.io.context import ActionContext
+from langbot_plugin.entities.io.resp import ActionResponse
+from langbot_plugin.runtime.io.connection import Connection
+
+from langbot.pkg.plugin.handler import RuntimeConnectionHandler
+from langbot.pkg.api.http.context import ExecutionContext
+
+
+class EmptyResult:
+ def first(self):
+ return None
+
+ def all(self):
+ return []
+
+
+class RecordingConnection(Connection):
+ def __init__(self):
+ self.sent: list[str] = []
+
+ async def send(self, message: str) -> None:
+ self.sent.append(message)
+
+ async def receive(self) -> str:
+ raise NotImplementedError
+
+ async def close(self) -> None:
+ return None
+
+
+def workspace_context(workspace_uuid: str = 'workspace-a') -> ActionContext:
+ return ActionContext(
+ instance_uuid='instance-a',
+ workspace_uuid=workspace_uuid,
+ placement_generation=7,
+ )
+
+
+def make_handler(workspace_uuid: str = 'workspace-a'):
+ context = workspace_context(workspace_uuid)
+ app = SimpleNamespace(
+ persistence_mgr=SimpleNamespace(execute_async=AsyncMock(return_value=EmptyResult())),
+ logger=Mock(),
+ workspace_service=SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid=context.instance_uuid,
+ workspace_uuid=context.workspace_uuid,
+ placement_generation=context.placement_generation,
+ )
+ )
+ ),
+ )
+ runtime_handler = RuntimeConnectionHandler(
+ Mock(),
+ AsyncMock(return_value=True),
+ app,
+ context,
+ )
+ installation_uuid = runtime_handler._remember_installation(
+ context,
+ 'author-a',
+ 'plugin-a',
+ )
+ return runtime_handler, app, context.for_installation(installation_uuid)
+
+
+async def invoke_with_context(
+ runtime_handler: RuntimeConnectionHandler,
+ action_context: ActionContext,
+ action: PluginToRuntimeAction,
+ data: dict,
+):
+ token = runtime_handler._current_action_context.set(action_context)
+ try:
+ return await runtime_handler.actions[action.value](data)
+ finally:
+ runtime_handler._current_action_context.reset(token)
+
+
+@pytest.mark.asyncio
+async def test_plugin_action_requires_installation_capability():
+ runtime_handler, _app, _installation_context = make_handler()
+
+ with pytest.raises(ValueError, match='missing installation_uuid'):
+ await runtime_handler.actions[PluginToRuntimeAction.GET_LANGBOT_VERSION.value]({})
+
+
+@pytest.mark.asyncio
+async def test_plugin_action_rejects_forged_installation_capability():
+ runtime_handler, _app, installation_context = make_handler()
+ forged = installation_context.model_copy(update={'installation_uuid': 'forged-installation'})
+
+ with pytest.raises(
+ ValueError,
+ match='installation is not registered in this Workspace',
+ ):
+ await invoke_with_context(
+ runtime_handler,
+ forged,
+ PluginToRuntimeAction.GET_LANGBOT_VERSION,
+ {},
+ )
+
+
+@pytest.mark.asyncio
+async def test_plugin_action_rejects_stale_workspace_generation():
+ runtime_handler, app, installation_context = make_handler()
+ app.workspace_service.get_execution_binding.side_effect = ValueError('generation is fenced')
+
+ with pytest.raises(ValueError, match='generation is fenced'):
+ await invoke_with_context(
+ runtime_handler,
+ installation_context,
+ PluginToRuntimeAction.GET_LANGBOT_VERSION,
+ {},
+ )
+
+ app.workspace_service.get_execution_binding.assert_awaited_once_with(
+ 'workspace-a',
+ expected_generation=7,
+ )
+
+
+@pytest.mark.asyncio
+async def test_query_uuid_and_forged_payload_workspace_cannot_cross_tenants():
+ runtime_handler, app, installation_context = make_handler()
+ query_a = SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=7,
+ bot_uuid='bot-a',
+ )
+ query_b = SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-b',
+ placement_generation=7,
+ bot_uuid='bot-b',
+ )
+
+ async def get_query(workspace_uuid, query_uuid):
+ return {
+ ('workspace-a', 'query-a'): query_a,
+ ('workspace-b', 'query-b'): query_b,
+ }.get((workspace_uuid, query_uuid))
+
+ app.query_pool = SimpleNamespace(
+ get_query=AsyncMock(side_effect=get_query),
+ get_query_by_legacy_id=AsyncMock(),
+ )
+
+ response = await invoke_with_context(
+ runtime_handler,
+ installation_context,
+ PluginToRuntimeAction.GET_BOT_UUID,
+ {
+ 'query_id': 2,
+ 'query_uuid': 'query-b',
+ 'workspace_uuid': 'workspace-b',
+ },
+ )
+
+ assert response.code != 0
+ app.query_pool.get_query.assert_awaited_once_with('workspace-a', 'query-b')
+
+
+@pytest.mark.asyncio
+async def test_legacy_query_id_fallback_is_workspace_scoped():
+ runtime_handler, app, installation_context = make_handler()
+ query = SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=7,
+ bot_uuid='bot-a',
+ )
+ app.query_pool = SimpleNamespace(
+ get_query=AsyncMock(),
+ get_query_by_legacy_id=AsyncMock(return_value=query),
+ )
+
+ response = await invoke_with_context(
+ runtime_handler,
+ installation_context,
+ PluginToRuntimeAction.GET_BOT_UUID,
+ {'query_id': 19},
+ )
+
+ assert response.code == 0
+ assert response.data == {'bot_uuid': 'bot-a'}
+ app.query_pool.get_query_by_legacy_id.assert_awaited_once_with(
+ 'workspace-a',
+ 19,
+ )
+
+
+def test_runtime_connection_rejects_workspace_rebinding():
+ runtime_handler, _app, _installation_context = make_handler()
+
+ with pytest.raises(ValueError, match='another Workspace'):
+ runtime_handler.bind_action_context(workspace_context('workspace-b'))
+
+
+def test_inbound_installation_envelope_is_preserved_only_inside_bound_workspace():
+ runtime_handler, _app, installation_context = make_handler()
+
+ assert (
+ runtime_handler.validate_inbound_action_context(
+ PluginToRuntimeAction.GET_BOTS.value,
+ installation_context,
+ )
+ == installation_context
+ )
+ with pytest.raises(ValueError, match='does not match connection Workspace'):
+ runtime_handler.validate_inbound_action_context(
+ PluginToRuntimeAction.GET_BOTS.value,
+ workspace_context('workspace-b').for_installation('installation-b'),
+ )
+
+
+def test_installation_uuid_is_stable_and_workspace_specific():
+ context_a = workspace_context('workspace-a')
+ context_b = workspace_context('workspace-b')
+
+ first = RuntimeConnectionHandler.derive_installation_uuid(
+ context_a,
+ 'author-a',
+ 'plugin-a',
+ )
+ repeated = RuntimeConnectionHandler.derive_installation_uuid(
+ context_a,
+ 'author-a',
+ 'plugin-a',
+ )
+ other_workspace = RuntimeConnectionHandler.derive_installation_uuid(
+ context_b,
+ 'author-a',
+ 'plugin-a',
+ )
+
+ assert first == repeated
+ assert first != other_workspace
+
+
+@pytest.mark.asyncio
+async def test_plugin_vector_action_forwards_trusted_context_and_logical_collection():
+ runtime_handler, app, installation_context = make_handler()
+ app.rag_runtime_service = SimpleNamespace(vector_upsert=AsyncMock())
+
+ response = await invoke_with_context(
+ runtime_handler,
+ installation_context,
+ PluginToRuntimeAction.VECTOR_UPSERT,
+ {
+ 'workspace_uuid': 'workspace-forged',
+ 'collection_id': 'plugin-supplied-name',
+ 'vectors': [[0.1]],
+ 'ids': ['point-a'],
+ },
+ )
+
+ assert response.code == 0
+ execution_context = app.rag_runtime_service.vector_upsert.await_args.args[0]
+ assert execution_context == ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=7,
+ )
+ assert app.rag_runtime_service.vector_upsert.await_args.args[1] == 'plugin-supplied-name'
+
+
+@pytest.mark.asyncio
+async def test_host_to_runtime_action_carries_trusted_connector_context():
+ app = SimpleNamespace(logger=Mock())
+ context = workspace_context()
+ connection = RecordingConnection()
+ runtime_handler = RuntimeConnectionHandler(
+ connection,
+ AsyncMock(return_value=True),
+ app,
+ context,
+ )
+
+ task = asyncio.create_task(runtime_handler.set_runtime_config(None))
+ for _ in range(10):
+ if connection.sent:
+ break
+ await asyncio.sleep(0)
+ request = json.loads(connection.sent[0])
+ runtime_handler.resp_waiters[request['seq_id']].set_result(ActionResponse.success({}))
+ await task
+
+ assert request['data'] == {}
+ assert request['context'] == context.model_dump(exclude_none=True)
diff --git a/tests/unit_tests/provider/conftest.py b/tests/unit_tests/provider/conftest.py
index 13b44fd14..be56eb09f 100644
--- a/tests/unit_tests/provider/conftest.py
+++ b/tests/unit_tests/provider/conftest.py
@@ -16,6 +16,18 @@ from langbot.pkg.provider.modelmgr import token
from langbot.pkg.provider.modelmgr.modelmgr import ModelManager
from langbot.pkg.entity.persistence import model as persistence_model
from langbot.pkg.discover import engine as discover_engine
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.workspace.entities import WorkspaceExecutionBinding
+
+
+TEST_INSTANCE_UUID = 'test-instance'
+TEST_WORKSPACE_UUID = 'test-workspace'
+TEST_GENERATION = 1
+TEST_EXECUTION_CONTEXT = ExecutionContext(
+ instance_uuid=TEST_INSTANCE_UUID,
+ workspace_uuid=TEST_WORKSPACE_UUID,
+ placement_generation=TEST_GENERATION,
+)
class FakeProviderAPIRequester(requester.ProviderAPIRequester):
@@ -157,6 +169,26 @@ def mock_app_for_modelmgr():
app.llm_model_service = AsyncMock()
app.embedding_models_service = AsyncMock()
app.monitoring_service = AsyncMock()
+ app.workspace_service = SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ return_value=WorkspaceExecutionBinding(
+ instance_uuid=TEST_INSTANCE_UUID,
+ workspace_uuid=TEST_WORKSPACE_UUID,
+ placement_generation=TEST_GENERATION,
+ write_fenced=False,
+ state='active',
+ )
+ ),
+ get_local_execution_binding=AsyncMock(
+ return_value=WorkspaceExecutionBinding(
+ instance_uuid=TEST_INSTANCE_UUID,
+ workspace_uuid=TEST_WORKSPACE_UUID,
+ placement_generation=TEST_GENERATION,
+ write_fenced=False,
+ state='active',
+ )
+ ),
+ )
return app
@@ -184,6 +216,7 @@ def fake_persistence_data():
providers = [
persistence_model.ModelProvider(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid=provider_uuid,
name='Test Provider',
requester='fake-requester',
@@ -191,6 +224,7 @@ def fake_persistence_data():
api_keys=['test-api-key-1', 'test-api-key-2'],
),
persistence_model.ModelProvider(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid=provider_uuid2,
name='Test Provider 2',
requester='another-fake-requester',
@@ -201,6 +235,7 @@ def fake_persistence_data():
llm_models = [
persistence_model.LLMModel(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid='test-llm-uuid-1',
name='TestLLM-1',
provider_uuid=provider_uuid,
@@ -208,6 +243,7 @@ def fake_persistence_data():
extra_args={'temperature': 0.7},
),
persistence_model.LLMModel(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid='test-llm-uuid-2',
name='TestLLM-2',
provider_uuid=provider_uuid,
@@ -218,6 +254,7 @@ def fake_persistence_data():
embedding_models = [
persistence_model.EmbeddingModel(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid='test-embedding-uuid-1',
name='TestEmbedding-1',
provider_uuid=provider_uuid,
@@ -227,6 +264,7 @@ def fake_persistence_data():
rerank_models = [
persistence_model.RerankModel(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid='test-rerank-uuid-1',
name='TestRerank-1',
provider_uuid=provider_uuid2,
@@ -252,6 +290,7 @@ def runtime_provider(fake_persistence_data, mock_app_for_modelmgr):
requester_inst = FakeProviderAPIRequester(mock_app_for_modelmgr, {'base_url': provider_entity.base_url})
return requester.RuntimeProvider(
+ execution_context=TEST_EXECUTION_CONTEXT,
provider_entity=provider_entity,
token_mgr=token_mgr,
requester=requester_inst,
@@ -263,6 +302,7 @@ def runtime_llm_model(fake_persistence_data, runtime_provider):
"""Provides a RuntimeLLMModel instance for testing."""
model_entity = fake_persistence_data['llm_models'][0]
return requester.RuntimeLLMModel(
+ execution_context=TEST_EXECUTION_CONTEXT,
model_entity=model_entity,
provider=runtime_provider,
)
@@ -273,6 +313,7 @@ def runtime_embedding_model(fake_persistence_data, runtime_provider):
"""Provides a RuntimeEmbeddingModel instance for testing."""
model_entity = fake_persistence_data['embedding_models'][0]
return requester.RuntimeEmbeddingModel(
+ execution_context=TEST_EXECUTION_CONTEXT,
model_entity=model_entity,
provider=runtime_provider,
)
@@ -286,6 +327,7 @@ def runtime_rerank_model(fake_persistence_data, mock_app_for_modelmgr):
requester_inst = AnotherFakeRequester(mock_app_for_modelmgr, {'base_url': provider_entity.base_url})
provider = requester.RuntimeProvider(
+ execution_context=TEST_EXECUTION_CONTEXT,
provider_entity=provider_entity,
token_mgr=token_mgr,
requester=requester_inst,
@@ -293,6 +335,7 @@ def runtime_rerank_model(fake_persistence_data, mock_app_for_modelmgr):
model_entity = fake_persistence_data['rerank_models'][0]
return requester.RuntimeRerankModel(
+ execution_context=TEST_EXECUTION_CONTEXT,
model_entity=model_entity,
provider=provider,
)
diff --git a/tests/unit_tests/provider/runners/test_difysvapi_runner.py b/tests/unit_tests/provider/runners/test_difysvapi_runner.py
index f75938209..f6d414f34 100644
--- a/tests/unit_tests/provider/runners/test_difysvapi_runner.py
+++ b/tests/unit_tests/provider/runners/test_difysvapi_runner.py
@@ -252,42 +252,65 @@ class TestDifyHumanInputForms:
runner.dify_client.upload_file = AsyncMock(return_value={'id': 'upload-1'})
return runner
- def test_pending_forms_are_isolated_by_bot_and_pipeline(self):
+ def test_pending_forms_are_isolated_by_workspace_generation_bot_and_pipeline(self):
from langbot.pkg.provider.runners import difysvapi
query_a = MagicMock()
+ query_a.instance_uuid = 'instance-a'
+ query_a.workspace_uuid = 'workspace-a'
+ query_a.placement_generation = 1
query_a.bot_uuid = 'bot-a'
query_a.pipeline_uuid = 'pipeline-a'
query_a.session.launcher_type.value = 'person'
query_a.session.launcher_id = 'shared-user'
query_b = MagicMock()
+ query_b.instance_uuid = 'instance-a'
+ query_b.workspace_uuid = 'workspace-a'
+ query_b.placement_generation = 1
query_b.bot_uuid = 'bot-b'
query_b.pipeline_uuid = 'pipeline-a'
query_b.session.launcher_type.value = 'person'
query_b.session.launcher_id = 'shared-user'
query_c = MagicMock()
+ query_c.instance_uuid = 'instance-a'
+ query_c.workspace_uuid = 'workspace-a'
+ query_c.placement_generation = 1
query_c.bot_uuid = 'bot-a'
query_c.pipeline_uuid = 'pipeline-b'
query_c.session.launcher_type.value = 'person'
query_c.session.launcher_id = 'shared-user'
+ query_d = MagicMock()
+ query_d.instance_uuid = 'instance-a'
+ query_d.workspace_uuid = 'workspace-b'
+ query_d.placement_generation = 2
+ query_d.bot_uuid = 'bot-a'
+ query_d.pipeline_uuid = 'pipeline-a'
+ query_d.session.launcher_type.value = 'person'
+ query_d.session.launcher_id = 'shared-user'
+
key_a = difysvapi._session_key_from_query(query_a)
key_b = difysvapi._session_key_from_query(query_b)
key_c = difysvapi._session_key_from_query(query_c)
+ key_d = difysvapi._session_key_from_query(query_d)
difysvapi._PENDING_FORMS.clear()
difysvapi._set_pending_form(key_a, {'form_token': 'token-a', 'workflow_run_id': 'run-a'})
difysvapi._set_pending_form(key_b, {'form_token': 'token-b', 'workflow_run_id': 'run-b'})
difysvapi._set_pending_form(key_c, {'form_token': 'token-c', 'workflow_run_id': 'run-c'})
+ difysvapi._set_pending_form(key_d, {'form_token': 'token-d', 'workflow_run_id': 'run-d'})
assert key_a != key_b
assert key_a != key_c
+ assert key_a != key_d
assert difysvapi._get_pending_form_by_token(key_a, 'token-a') is not None
assert difysvapi._get_pending_form_by_token(key_a, 'token-b') is None
assert difysvapi._get_pending_form_by_token(key_a, 'token-c') is None
+ assert difysvapi._get_pending_form_by_token(key_a, 'token-d') is None
assert difysvapi._get_pending_form_by_token(key_b, 'token-b') is not None
assert difysvapi._get_pending_form_by_token(key_c, 'token-c') is not None
+ assert difysvapi._get_pending_form_by_token(key_d, 'token-d') is not None
assert difysvapi._get_latest_pending_form(key_a)['workflow_run_id'] == 'run-a'
assert difysvapi._get_latest_pending_form(key_b)['workflow_run_id'] == 'run-b'
assert difysvapi._get_latest_pending_form(key_c)['workflow_run_id'] == 'run-c'
diff --git a/tests/unit_tests/provider/test_localagent_sandbox_exec.py b/tests/unit_tests/provider/test_localagent_sandbox_exec.py
index 9bc155343..e912d4e27 100644
--- a/tests/unit_tests/provider/test_localagent_sandbox_exec.py
+++ b/tests/unit_tests/provider/test_localagent_sandbox_exec.py
@@ -10,6 +10,7 @@ import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
import langbot_plugin.api.entities.builtin.provider.message as provider_message
import langbot_plugin.api.entities.builtin.provider.session as provider_session
+from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.provider.runners.localagent import LocalAgentRunner, _StreamAccumulator
@@ -95,7 +96,7 @@ def make_query() -> pipeline_query.Query:
adapter = AsyncMock()
adapter.is_stream_output_supported = AsyncMock(return_value=False)
- return pipeline_query.Query.model_construct(
+ query = pipeline_query.Query.model_construct(
query_id='avg-query',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -122,6 +123,20 @@ def make_query() -> pipeline_query.Query:
use_llm_model_uuid='test-model-uuid',
variables={},
)
+ object.__setattr__(
+ query,
+ '_execution_context',
+ ExecutionContext(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-test',
+ placement_generation=1,
+ bot_uuid='bot-uuid',
+ pipeline_uuid='pipeline-uuid',
+ query_uuid='query-avg',
+ ),
+ )
+ object.__setattr__(query, 'query_uuid', 'query-avg')
+ return query
def test_stream_accumulator_merges_fragmented_tool_call_arguments():
diff --git a/tests/unit_tests/provider/test_mcp_box_integration.py b/tests/unit_tests/provider/test_mcp_box_integration.py
index 6b0c141ed..bf0d82a20 100644
--- a/tests/unit_tests/provider/test_mcp_box_integration.py
+++ b/tests/unit_tests/provider/test_mcp_box_integration.py
@@ -154,7 +154,16 @@ def mcp_module():
def _make_ap():
ap = Mock()
ap.logger = Mock()
+ ap.workspace_service = Mock()
+ ap.workspace_service.get_execution_binding = AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ )
+ )
ap.box_service = Mock()
+ ap.box_service.get_managed_process_websocket_connection = AsyncMock(return_value=('ws://box.example/process', {}))
return ap
@@ -166,6 +175,11 @@ def _make_session(mcp_module, server_config: dict, ap=None):
server_config=server_config,
enable=True,
ap=ap,
+ execution_context=mcp_module.ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ ),
)
@@ -417,7 +431,7 @@ class TestBuildBoxSessionPayload:
payload = s._build_box_session_payload('session-123')
assert payload['image'] == 'node:20'
assert payload['cpus'] == 2.0
- assert payload["memory_mb"] == 1024
+ assert payload['memory_mb'] == 1024
assert payload['pids_limit'] == 256
def test_none_fields_excluded(self, mcp_module):
@@ -591,6 +605,26 @@ class TestGetRuntimeInfoDict:
assert info['status'] == 'connecting'
assert 'box_session_id' not in info
+ def test_runtime_error_detail_never_echoes_secret_config(self, mcp_module):
+ s = _make_session(
+ mcp_module,
+ {
+ 'name': 'test',
+ 'uuid': 'test-uuid',
+ 'mode': 'invalid',
+ 'headers': {'Authorization': 'Bearer TOPSECRET'},
+ 'env': {'API_KEY': 'TOPSECRET'},
+ },
+ )
+ s.status = mcp_module.MCPSessionStatus.ERROR
+ s.error_message = f'Unknown MCP server mode: {s.server_config}'
+
+ info = s.get_runtime_info_dict()
+
+ assert info['error_message'] == 'MCP runtime failed'
+ assert info['error_code'] == 'runtime_error'
+ assert 'TOPSECRET' not in str(info)
+
def test_runtime_tools_include_parameters(self, mcp_module):
s = _make_session(
mcp_module,
@@ -774,12 +808,16 @@ async def test_init_box_stdio_server_stages_host_path_in_shared_workspace(mcp_mo
async def initialize(self):
return None
+ captured_transport = {}
+
@asynccontextmanager
- async def fake_websocket_client(_url: str):
+ async def fake_authenticated_websocket_client(url: str, headers: dict[str, str]):
+ captured_transport['url'] = url
+ captured_transport['headers'] = headers
yield ('read-stream', 'write-stream')
mcp_stdio_module.ClientSession = FakeClientSession
- mcp_stdio_module.websocket_client = fake_websocket_client
+ mcp_stdio_module.authenticated_websocket_client = fake_authenticated_websocket_client
ap = _make_ap()
ap.box_service.available = True
@@ -790,7 +828,17 @@ async def test_init_box_stdio_server_stages_host_path_in_shared_workspace(mcp_mo
execute=AsyncMock(return_value=SimpleNamespace(ok=True, stderr='', exit_code=0))
)
ap.box_service.start_managed_process = AsyncMock(return_value={})
- ap.box_service.get_managed_process_websocket_url = Mock(return_value='ws://box.example/process')
+ ap.box_service.get_managed_process_websocket_connection = AsyncMock(
+ return_value=(
+ 'ws://box.example/process',
+ {
+ 'X-LangBot-Box-Control-Token': 'secret-token',
+ 'X-LangBot-Instance-Id': 'instance-a',
+ 'X-LangBot-Workspace-Id': 'workspace-a',
+ 'X-LangBot-Placement-Generation': '1',
+ },
+ )
+ )
host_path = tmp_path / 'mcp-source'
host_path.mkdir()
@@ -814,7 +862,8 @@ async def test_init_box_stdio_server_stages_host_path_in_shared_workspace(mcp_mo
await session.exit_stack.aclose()
assert ap.box_service.create_session.await_count == 1
- session_payload = ap.box_service.create_session.await_args.args[0]
+ assert ap.box_service.create_session.await_args.args[0] == session.execution_context
+ session_payload = ap.box_service.create_session.await_args.args[1]
assert session_payload['session_id'] == 'mcp-shared'
assert 'host_path' not in session_payload
assert ap.box_service.build_spec.call_count == 1
@@ -824,11 +873,22 @@ async def test_init_box_stdio_server_stages_host_path_in_shared_workspace(mcp_mo
staged_file = tmp_path / 'shared-box-workspace' / '.mcp' / 'u1' / 'workspace' / 'server.py'
assert staged_file.read_text(encoding='utf-8') == 'print("hello")\n'
- process_payload = ap.box_service.start_managed_process.await_args.args[1]
+ assert ap.box_service.start_managed_process.await_args.args[0] == session.execution_context
+ process_payload = ap.box_service.start_managed_process.await_args.args[2]
assert process_payload['process_id'] == 'u1'
assert process_payload['command'] == 'python'
assert process_payload['args'] == ['/workspace/.mcp/u1/workspace/server.py']
assert process_payload['cwd'] == '/workspace/.mcp/u1/workspace'
+ assert captured_transport == {
+ 'url': 'ws://box.example/process',
+ 'headers': {
+ 'X-LangBot-Box-Control-Token': 'secret-token',
+ 'X-LangBot-Instance-Id': 'instance-a',
+ 'X-LangBot-Workspace-Id': 'workspace-a',
+ 'X-LangBot-Placement-Generation': '1',
+ },
+ }
+ assert 'secret-token' not in captured_transport['url']
@pytest.mark.asyncio
@@ -867,7 +927,7 @@ async def test_stdio_handshake_raises_coldstart_retry_while_process_alive(mcp_mo
ap.box_service.available = True
ap.box_service.create_session = AsyncMock(return_value={})
ap.box_service.start_managed_process = AsyncMock(return_value={})
- ap.box_service.get_managed_process_websocket_url = Mock(return_value='ws://box/p')
+ ap.box_service.get_managed_process_websocket_connection = AsyncMock(return_value=('ws://box/p', {}))
session = _make_session(
mcp_module,
@@ -933,7 +993,7 @@ async def test_stdio_handshake_raises_fatal_when_process_exited(mcp_module, tmp_
ap.box_service.available = True
ap.box_service.create_session = AsyncMock(return_value={})
ap.box_service.start_managed_process = AsyncMock(return_value={})
- ap.box_service.get_managed_process_websocket_url = Mock(return_value='ws://box/p')
+ ap.box_service.get_managed_process_websocket_connection = AsyncMock(return_value=('ws://box/p', {}))
session = _make_session(
mcp_module,
diff --git a/tests/unit_tests/provider/test_mcp_remote_transport.py b/tests/unit_tests/provider/test_mcp_remote_transport.py
index b2f1d2e14..ba53397c0 100644
--- a/tests/unit_tests/provider/test_mcp_remote_transport.py
+++ b/tests/unit_tests/provider/test_mcp_remote_transport.py
@@ -12,9 +12,17 @@ import pytest
from aiohttp import web
from mcp import types as mcp_types
+from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.provider.tools.loaders.mcp import RuntimeMCPSession
+TEST_EXECUTION_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+)
+
+
class _TransportProbe:
def __init__(self, streamable_status: int | None) -> None:
self.streamable_status = streamable_status
@@ -133,6 +141,7 @@ def _session(url: str, *, timeout: float = 2) -> RuntimeMCPSession:
{'uuid': 'srv-1', 'mode': 'remote', 'url': url, 'timeout': timeout},
True,
app,
+ TEST_EXECUTION_CONTEXT,
)
diff --git a/tests/unit_tests/provider/test_mcp_resources.py b/tests/unit_tests/provider/test_mcp_resources.py
index 773efe71c..ed073e1e9 100644
--- a/tests/unit_tests/provider/test_mcp_resources.py
+++ b/tests/unit_tests/provider/test_mcp_resources.py
@@ -8,6 +8,7 @@ import httpx
import pytest
from mcp import types as mcp_types
+from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.provider.tools.loaders.mcp import (
MCP_RESOURCE_CONTEXT_QUERY_KEY,
MCP_RESOURCE_TRACE_QUERY_KEY,
@@ -18,10 +19,30 @@ from langbot.pkg.provider.tools.loaders.mcp import (
RuntimeMCPSession,
)
from langbot.pkg.telemetry import features as telemetry_features
+from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError
+
+
+TEST_EXECUTION_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ query_uuid='query-a',
+)
def _app() -> SimpleNamespace:
- return SimpleNamespace(logger=Mock())
+ return SimpleNamespace(
+ logger=Mock(),
+ workspace_service=SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ )
+ )
+ ),
+ )
def _connected_session(
@@ -30,8 +51,15 @@ def _connected_session(
uuid: str = 'srv-1',
resources: list[dict] | None = None,
templates: list[dict] | None = None,
+ execution_context: ExecutionContext = TEST_EXECUTION_CONTEXT,
) -> RuntimeMCPSession:
- session = RuntimeMCPSession(name, {'uuid': uuid, 'mode': 'remote'}, True, _app())
+ session = RuntimeMCPSession(
+ name,
+ {'uuid': uuid, 'mode': 'remote'},
+ True,
+ _app(),
+ execution_context,
+ )
session.status = MCPSessionStatus.CONNECTED
session.session = SimpleNamespace(read_resource=AsyncMock())
session.resources = resources or [
@@ -51,8 +79,20 @@ def _connected_session(
return session
-def _query() -> SimpleNamespace:
- return SimpleNamespace(variables={})
+def _query(variables: dict | None = None, context: ExecutionContext = TEST_EXECUTION_CONTEXT) -> SimpleNamespace:
+ return SimpleNamespace(
+ instance_uuid=context.instance_uuid,
+ workspace_uuid=context.workspace_uuid,
+ placement_generation=context.placement_generation,
+ bot_uuid=context.bot_uuid,
+ pipeline_uuid=context.pipeline_uuid,
+ query_uuid=context.query_uuid,
+ variables=variables or {},
+ )
+
+
+def _register_session(loader: MCPLoader, session: RuntimeMCPSession) -> None:
+ loader.sessions[loader._session_key(session.execution_context, session.server_name)] = session
def _http_status_error(status_code: int) -> httpx.HTTPStatusError:
@@ -68,6 +108,7 @@ async def test_remote_transport_falls_back_to_sse_for_compatible_http_status_in_
{'uuid': 'srv-1', 'mode': 'remote', 'url': 'https://example.com/mcp'},
True,
_app(),
+ TEST_EXECUTION_CONTEXT,
)
session._init_streamable_http_server = AsyncMock(
side_effect=ExceptionGroup('transport failed', [_http_status_error(405)])
@@ -87,6 +128,7 @@ async def test_remote_transport_does_not_fallback_for_auth_http_status():
{'uuid': 'srv-1', 'mode': 'remote', 'url': 'https://example.com/mcp'},
True,
_app(),
+ TEST_EXECUTION_CONTEXT,
)
error = _http_status_error(403)
session._init_streamable_http_server = AsyncMock(side_effect=error)
@@ -244,10 +286,18 @@ def test_resource_uri_allowed_supports_listed_templates_conservatively():
async def test_mcp_loader_can_hide_synthetic_resource_tools():
loader = MCPLoader(_app())
session = _connected_session()
- loader.sessions = {'docs': session}
+ _register_session(loader, session)
- with_resource_tools = await loader.get_tools(['srv-1'], include_resource_tools=True)
- without_resource_tools = await loader.get_tools(['srv-1'], include_resource_tools=False)
+ with_resource_tools = await loader.get_tools(
+ TEST_EXECUTION_CONTEXT,
+ ['srv-1'],
+ include_resource_tools=True,
+ )
+ without_resource_tools = await loader.get_tools(
+ TEST_EXECUTION_CONTEXT,
+ ['srv-1'],
+ include_resource_tools=False,
+ )
assert {tool.name for tool in with_resource_tools} == {
MCP_TOOL_LIST_RESOURCES,
@@ -260,9 +310,9 @@ async def test_mcp_loader_can_hide_synthetic_resource_tools():
async def test_mcp_loader_refuses_resource_tool_calls_when_agent_read_disabled():
loader = MCPLoader(_app())
session = _connected_session()
- loader.sessions = {'docs': session}
- query = SimpleNamespace(
- variables={
+ _register_session(loader, session)
+ query = _query(
+ {
'_pipeline_bound_mcp_servers': ['srv-1'],
'_pipeline_mcp_resource_agent_read_enabled': False,
}
@@ -301,9 +351,10 @@ async def test_build_resource_context_for_query_uses_only_bound_attached_text_re
)
]
)
- loader.sessions = {'docs': docs, 'other': other}
- query = SimpleNamespace(
- variables={
+ _register_session(loader, docs)
+ _register_session(loader, other)
+ query = _query(
+ {
'_pipeline_bound_mcp_servers': ['srv-1'],
'_pipeline_mcp_resource_attachments': [
{'server_uuid': 'srv-1', 'server_name': 'docs', 'uri': 'file:///README.md', 'mode': 'pinned'},
@@ -321,3 +372,88 @@ async def test_build_resource_context_for_query_uses_only_bound_attached_text_re
assert query.variables[MCP_RESOURCE_CONTEXT_QUERY_KEY]['resource_count'] == 1
docs.session.read_resource.assert_awaited_once()
other.session.read_resource.assert_not_called()
+
+
+def test_mcp_loader_session_keys_do_not_collide_between_workspaces():
+ loader = MCPLoader(_app())
+ workspace_b = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-b',
+ placement_generation=1,
+ query_uuid='query-b',
+ )
+ session_a = _connected_session(name='docs', uuid='srv-a')
+ session_b = _connected_session(
+ name='docs',
+ uuid='srv-b',
+ execution_context=workspace_b,
+ )
+ _register_session(loader, session_a)
+ _register_session(loader, session_b)
+
+ assert len(loader.sessions) == 2
+ assert loader.get_session(TEST_EXECUTION_CONTEXT, 'docs') is session_a
+ assert loader.get_session(workspace_b, 'docs') is session_b
+ assert loader.get_session(TEST_EXECUTION_CONTEXT, 'docs') is not loader.get_session(workspace_b, 'docs')
+
+
+@pytest.mark.asyncio
+async def test_mcp_tool_result_is_discarded_when_generation_changes_during_call():
+ session = _connected_session()
+ binding = SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ )
+ session.ap.workspace_service.get_execution_binding.side_effect = [
+ binding,
+ binding,
+ WorkspaceGenerationMismatchError('generation changed during tool call'),
+ ]
+ session.session = SimpleNamespace(call_tool=AsyncMock(return_value=SimpleNamespace(isError=False, content=[])))
+
+ with pytest.raises(WorkspaceGenerationMismatchError):
+ await session.invoke_mcp_tool('side_effecting_tool', {})
+
+ session.session.call_tool.assert_awaited_once_with('side_effecting_tool', {})
+
+
+@pytest.mark.asyncio
+async def test_mcp_resource_cache_is_not_served_to_stale_generation():
+ session = _connected_session()
+ session._resource_cache[('file:///README.md', 10, None, False)] = {
+ 'cached_at': 0,
+ 'envelope': {'contents': [{'type': 'text', 'text': 'stale'}]},
+ }
+ session.ap.workspace_service.get_execution_binding.side_effect = WorkspaceGenerationMismatchError(
+ 'stale generation'
+ )
+
+ with pytest.raises(WorkspaceGenerationMismatchError):
+ await session.read_resource_envelope('file:///README.md', max_bytes=10)
+
+ session.session.read_resource.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_mcp_idle_lifecycle_stops_without_retry_after_generation_bump():
+ session = _connected_session()
+ session.server_config.update({'mode': 'remote', 'url': 'https://example.com/mcp'})
+ session._FENCE_POLL_INTERVAL = 0
+ session._init_remote_server = AsyncMock()
+ session.refresh = AsyncMock()
+ session._assert_execution_active = AsyncMock(
+ side_effect=[
+ None,
+ None,
+ None,
+ WorkspaceGenerationMismatchError('generation changed while idle'),
+ ]
+ )
+
+ await session._lifecycle_loop_with_retry()
+
+ session._init_remote_server.assert_awaited_once_with()
+ assert session.status == MCPSessionStatus.ERROR
+ assert session.error_message == 'Workspace execution binding is stale'
+ assert session._shutdown_event.is_set()
diff --git a/tests/unit_tests/provider/test_model_manager.py b/tests/unit_tests/provider/test_model_manager.py
index 015fd5450..a9aafff86 100644
--- a/tests/unit_tests/provider/test_model_manager.py
+++ b/tests/unit_tests/provider/test_model_manager.py
@@ -8,14 +8,22 @@ and error handling without calling real LLM APIs.
from __future__ import annotations
import pytest
-from unittest.mock import Mock
+from unittest.mock import AsyncMock, Mock
from langbot.pkg.provider.modelmgr.modelmgr import ModelManager
from langbot.pkg.provider.modelmgr import requester
from langbot.pkg.entity.persistence import model as persistence_model
from langbot.pkg.entity.errors import provider as provider_errors
from langbot.pkg.provider.modelmgr import token
-from tests.unit_tests.provider.conftest import _make_mock_result, _make_row_mock
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.workspace.entities import WorkspaceExecutionBinding
+from langbot.pkg.workspace.errors import WorkspaceGenerationMismatchError
+from tests.unit_tests.provider.conftest import (
+ TEST_EXECUTION_CONTEXT,
+ TEST_WORKSPACE_UUID,
+ _make_mock_result,
+ _make_row_mock,
+)
# ============================================================================
@@ -91,13 +99,15 @@ async def test_model_manager_load_models_from_db(fake_requester_registry, fake_p
# Check providers loaded
assert len(model_mgr.provider_dict) == 2
- assert fake_persistence_data['provider_uuid'] in model_mgr.provider_dict
- assert fake_persistence_data['provider_uuid2'] in model_mgr.provider_dict
+ assert {provider.provider_entity.uuid for provider in model_mgr.provider_dict.values()} == {
+ fake_persistence_data['provider_uuid'],
+ fake_persistence_data['provider_uuid2'],
+ }
# Check models loaded
- assert len(model_mgr.llm_models) == 2
- assert len(model_mgr.embedding_models) == 1
- assert len(model_mgr.rerank_models) == 1
+ assert len(model_mgr.llm_model_dict) == 2
+ assert len(model_mgr.embedding_model_dict) == 1
+ assert len(model_mgr.rerank_model_dict) == 1
@pytest.mark.asyncio
@@ -118,7 +128,7 @@ async def test_model_manager_load_provider_unknown_requester(mock_app_for_modelm
}
with pytest.raises(provider_errors.RequesterNotFoundError) as exc_info:
- await model_mgr.load_provider(provider_info)
+ await model_mgr.load_provider(TEST_EXECUTION_CONTEXT, provider_info)
assert exc_info.value.requester_name == 'non-existent-requester'
@@ -137,7 +147,7 @@ async def test_model_manager_load_provider_from_dict(fake_requester_registry):
'api_keys': ['dict-key'],
}
- runtime_provider = await model_mgr.load_provider(provider_info)
+ runtime_provider = await model_mgr.load_provider(TEST_EXECUTION_CONTEXT, provider_info)
assert runtime_provider.provider_entity.uuid == 'dict-provider-uuid'
assert runtime_provider.provider_entity.name == 'Dict Provider'
@@ -154,7 +164,7 @@ async def test_model_manager_load_provider_from_entity(fake_requester_registry,
provider_entity = fake_persistence_data['providers'][0]
- runtime_provider = await model_mgr.load_provider(provider_entity)
+ runtime_provider = await model_mgr.load_provider(TEST_EXECUTION_CONTEXT, provider_entity)
assert runtime_provider.provider_entity.uuid == provider_entity.uuid
assert runtime_provider.requester is not None
@@ -181,7 +191,7 @@ async def test_model_manager_get_model_by_uuid(fake_requester_registry, fake_per
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
- model = await model_mgr.get_model_by_uuid('test-llm-uuid-1')
+ model = await model_mgr.get_model_by_uuid(TEST_EXECUTION_CONTEXT, 'test-llm-uuid-1')
assert model.model_entity.uuid == 'test-llm-uuid-1'
assert model.model_entity.name == 'TestLLM-1'
@@ -194,7 +204,7 @@ async def test_model_manager_get_model_by_uuid_not_found(fake_requester_registry
await model_mgr.initialize()
with pytest.raises(ValueError) as exc_info:
- await model_mgr.get_model_by_uuid('unknown-model-uuid')
+ await model_mgr.get_model_by_uuid(TEST_EXECUTION_CONTEXT, 'unknown-model-uuid')
assert 'unknown-model-uuid' in str(exc_info.value)
@@ -215,7 +225,10 @@ async def test_model_manager_get_embedding_model_by_uuid(fake_requester_registry
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
- model = await model_mgr.get_embedding_model_by_uuid('test-embedding-uuid-1')
+ model = await model_mgr.get_embedding_model_by_uuid(
+ TEST_EXECUTION_CONTEXT,
+ 'test-embedding-uuid-1',
+ )
assert model.model_entity.uuid == 'test-embedding-uuid-1'
@@ -227,7 +240,10 @@ async def test_model_manager_get_embedding_model_by_uuid_not_found(fake_requeste
await model_mgr.initialize()
with pytest.raises(ValueError):
- await model_mgr.get_embedding_model_by_uuid('unknown-embedding-uuid')
+ await model_mgr.get_embedding_model_by_uuid(
+ TEST_EXECUTION_CONTEXT,
+ 'unknown-embedding-uuid',
+ )
@pytest.mark.asyncio
@@ -246,7 +262,7 @@ async def test_model_manager_get_rerank_model_by_uuid(fake_requester_registry, f
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
- model = await model_mgr.get_rerank_model_by_uuid('test-rerank-uuid-1')
+ model = await model_mgr.get_rerank_model_by_uuid(TEST_EXECUTION_CONTEXT, 'test-rerank-uuid-1')
assert model.model_entity.uuid == 'test-rerank-uuid-1'
@@ -258,7 +274,7 @@ async def test_model_manager_get_rerank_model_by_uuid_not_found(fake_requester_r
await model_mgr.initialize()
with pytest.raises(ValueError):
- await model_mgr.get_rerank_model_by_uuid('unknown-rerank-uuid')
+ await model_mgr.get_rerank_model_by_uuid(TEST_EXECUTION_CONTEXT, 'unknown-rerank-uuid')
# ============================================================================
@@ -282,12 +298,12 @@ async def test_model_manager_remove_llm_model(fake_requester_registry, fake_pers
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
- assert len(model_mgr.llm_models) == 2
+ assert len(model_mgr.llm_model_dict) == 2
- await model_mgr.remove_llm_model('test-llm-uuid-1')
+ await model_mgr.remove_llm_model(TEST_EXECUTION_CONTEXT, 'test-llm-uuid-1')
- assert len(model_mgr.llm_models) == 1
- assert model_mgr.llm_models[0].model_entity.uuid == 'test-llm-uuid-2'
+ assert len(model_mgr.llm_model_dict) == 1
+ assert next(iter(model_mgr.llm_model_dict.values())).model_entity.uuid == 'test-llm-uuid-2'
@pytest.mark.asyncio
@@ -306,12 +322,12 @@ async def test_model_manager_remove_llm_model_not_found(fake_requester_registry,
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
- original_count = len(model_mgr.llm_models)
+ original_count = len(model_mgr.llm_model_dict)
# Removing unknown model should do nothing (no error)
- await model_mgr.remove_llm_model('unknown-model-uuid')
+ await model_mgr.remove_llm_model(TEST_EXECUTION_CONTEXT, 'unknown-model-uuid')
- assert len(model_mgr.llm_models) == original_count
+ assert len(model_mgr.llm_model_dict) == original_count
@pytest.mark.asyncio
@@ -330,11 +346,11 @@ async def test_model_manager_remove_embedding_model(fake_requester_registry, fak
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
- assert len(model_mgr.embedding_models) == 1
+ assert len(model_mgr.embedding_model_dict) == 1
- await model_mgr.remove_embedding_model('test-embedding-uuid-1')
+ await model_mgr.remove_embedding_model(TEST_EXECUTION_CONTEXT, 'test-embedding-uuid-1')
- assert len(model_mgr.embedding_models) == 0
+ assert len(model_mgr.embedding_model_dict) == 0
@pytest.mark.asyncio
@@ -353,11 +369,11 @@ async def test_model_manager_remove_rerank_model(fake_requester_registry, fake_p
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
- assert len(model_mgr.rerank_models) == 1
+ assert len(model_mgr.rerank_model_dict) == 1
- await model_mgr.remove_rerank_model('test-rerank-uuid-1')
+ await model_mgr.remove_rerank_model(TEST_EXECUTION_CONTEXT, 'test-rerank-uuid-1')
- assert len(model_mgr.rerank_models) == 0
+ assert len(model_mgr.rerank_model_dict) == 0
@pytest.mark.asyncio
@@ -376,11 +392,17 @@ async def test_model_manager_remove_provider(fake_requester_registry, fake_persi
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
- assert fake_persistence_data['provider_uuid'] in model_mgr.provider_dict
+ assert any(
+ provider.provider_entity.uuid == fake_persistence_data['provider_uuid']
+ for provider in model_mgr.provider_dict.values()
+ )
- await model_mgr.remove_provider(fake_persistence_data['provider_uuid'])
+ await model_mgr.remove_provider(TEST_EXECUTION_CONTEXT, fake_persistence_data['provider_uuid'])
- assert fake_persistence_data['provider_uuid'] not in model_mgr.provider_dict
+ assert all(
+ provider.provider_entity.uuid != fake_persistence_data['provider_uuid']
+ for provider in model_mgr.provider_dict.values()
+ )
# ============================================================================
@@ -498,7 +520,7 @@ async def test_model_manager_init_temporary_runtime_llm_model(fake_requester_reg
'extra_args': {'temperature': 0.5},
}
- runtime_model = await model_mgr.init_temporary_runtime_llm_model(model_info)
+ runtime_model = await model_mgr.init_temporary_runtime_llm_model(TEST_EXECUTION_CONTEXT, model_info)
assert runtime_model.model_entity.uuid == 'temp-model-uuid'
assert runtime_model.model_entity.name == 'TempModel'
@@ -528,7 +550,10 @@ async def test_model_manager_init_temporary_runtime_embedding_model(fake_request
'extra_args': {'dimensions': 512},
}
- runtime_model = await model_mgr.init_temporary_runtime_embedding_model(model_info)
+ runtime_model = await model_mgr.init_temporary_runtime_embedding_model(
+ TEST_EXECUTION_CONTEXT,
+ model_info,
+ )
assert runtime_model.model_entity.uuid == 'temp-embedding-uuid'
assert runtime_model.model_entity.name == 'TempEmbedding'
@@ -553,7 +578,10 @@ async def test_model_manager_init_temporary_runtime_rerank_model(fake_requester_
'extra_args': {},
}
- runtime_model = await model_mgr.init_temporary_runtime_rerank_model(model_info)
+ runtime_model = await model_mgr.init_temporary_runtime_rerank_model(
+ TEST_EXECUTION_CONTEXT,
+ model_info,
+ )
assert runtime_model.model_entity.uuid == 'temp-rerank-uuid'
assert runtime_model.model_entity.name == 'TempRerank'
@@ -589,12 +617,16 @@ async def test_model_manager_reload_provider(fake_requester_registry, fake_persi
model_mgr.ap.persistence_mgr.execute_async = fake_execute
await model_mgr.initialize()
- original_provider = model_mgr.provider_dict[fake_persistence_data['provider_uuid']]
+ original_provider = await model_mgr.get_provider_by_uuid(
+ TEST_EXECUTION_CONTEXT,
+ fake_persistence_data['provider_uuid'],
+ )
original_base_url = original_provider.provider_entity.base_url
# Setup for reload - return updated provider
async def reload_execute(query):
updated_provider = persistence_model.ModelProvider(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid=fake_persistence_data['provider_uuid'],
name='Updated Provider',
requester='fake-requester',
@@ -605,9 +637,12 @@ async def test_model_manager_reload_provider(fake_requester_registry, fake_persi
model_mgr.ap.persistence_mgr.execute_async = reload_execute
- await model_mgr.reload_provider(fake_persistence_data['provider_uuid'])
+ await model_mgr.reload_provider(TEST_EXECUTION_CONTEXT, fake_persistence_data['provider_uuid'])
- updated_provider = model_mgr.provider_dict[fake_persistence_data['provider_uuid']]
+ updated_provider = await model_mgr.get_provider_by_uuid(
+ TEST_EXECUTION_CONTEXT,
+ fake_persistence_data['provider_uuid'],
+ )
assert updated_provider.provider_entity.base_url == 'https://updated.example.com'
assert updated_provider.provider_entity.base_url != original_base_url
@@ -624,7 +659,7 @@ async def test_model_manager_reload_provider_not_found(fake_requester_registry):
model_mgr.ap.persistence_mgr.execute_async = fake_execute
with pytest.raises(provider_errors.ProviderNotFoundError) as exc_info:
- await model_mgr.reload_provider('unknown-provider-uuid')
+ await model_mgr.reload_provider(TEST_EXECUTION_CONTEXT, 'unknown-provider-uuid')
assert exc_info.value.provider_name == 'unknown-provider-uuid'
@@ -643,7 +678,11 @@ async def test_model_manager_load_llm_model_with_provider(
model_entity = fake_persistence_data['llm_models'][0]
- runtime_model = await model_mgr.load_llm_model_with_provider(model_entity, runtime_provider)
+ runtime_model = await model_mgr.load_llm_model_with_provider(
+ TEST_EXECUTION_CONTEXT,
+ model_entity,
+ runtime_provider,
+ )
assert runtime_model.model_entity.uuid == model_entity.uuid
assert runtime_model.provider is runtime_provider
@@ -659,7 +698,11 @@ async def test_model_manager_load_llm_model_with_provider_from_row(
model_entity = fake_persistence_data['llm_models'][0]
row_mock = _make_row_mock(model_entity)
- runtime_model = await model_mgr.load_llm_model_with_provider(row_mock, runtime_provider)
+ runtime_model = await model_mgr.load_llm_model_with_provider(
+ TEST_EXECUTION_CONTEXT,
+ row_mock,
+ runtime_provider,
+ )
assert runtime_model.model_entity.uuid == model_entity.uuid
@@ -673,7 +716,11 @@ async def test_model_manager_load_embedding_model_with_provider(
model_entity = fake_persistence_data['embedding_models'][0]
- runtime_model = await model_mgr.load_embedding_model_with_provider(model_entity, runtime_provider)
+ runtime_model = await model_mgr.load_embedding_model_with_provider(
+ TEST_EXECUTION_CONTEXT,
+ model_entity,
+ runtime_provider,
+ )
assert runtime_model.model_entity.uuid == model_entity.uuid
assert runtime_model.provider is runtime_provider
@@ -692,6 +739,7 @@ async def test_model_manager_load_rerank_model_with_provider(fake_requester_regi
)
await requester_inst.initialize()
provider = requester.RuntimeProvider(
+ execution_context=TEST_EXECUTION_CONTEXT,
provider_entity=provider_entity,
token_mgr=token_mgr,
requester=requester_inst,
@@ -699,7 +747,11 @@ async def test_model_manager_load_rerank_model_with_provider(fake_requester_regi
model_entity = fake_persistence_data['rerank_models'][0]
- runtime_model = await model_mgr.load_rerank_model_with_provider(model_entity, provider)
+ runtime_model = await model_mgr.load_rerank_model_with_provider(
+ TEST_EXECUTION_CONTEXT,
+ model_entity,
+ provider,
+ )
assert runtime_model.model_entity.uuid == model_entity.uuid
assert runtime_model.provider is provider
@@ -723,6 +775,7 @@ async def test_model_manager_logs_warning_for_missing_provider(fake_requester_re
elif 'llm_models' in query_str:
# Return model with missing provider
fake_model = persistence_model.LLMModel(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid='model-with-missing-provider',
name='MissingProviderModel',
provider_uuid='missing-provider-uuid',
@@ -736,7 +789,7 @@ async def test_model_manager_logs_warning_for_missing_provider(fake_requester_re
await model_mgr.initialize()
# Should have logged warning and skipped the model
- assert len(model_mgr.llm_models) == 0
+ assert len(model_mgr.llm_model_dict) == 0
model_mgr.ap.logger.warning.assert_called()
@@ -750,6 +803,7 @@ async def test_model_manager_handles_requester_not_found_gracefully(fake_request
if 'model_providers' in query_str:
# Return provider with unknown requester
fake_provider = persistence_model.ModelProvider(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid='provider-with-unknown-requester',
name='Unknown Requester Provider',
requester='unknown-requester-name',
@@ -759,6 +813,7 @@ async def test_model_manager_handles_requester_not_found_gracefully(fake_request
return _make_mock_result([_make_row_mock(fake_provider)])
elif 'llm_models' in query_str:
fake_model = persistence_model.LLMModel(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid='model-uuid',
name='Model',
provider_uuid='provider-with-unknown-requester',
@@ -773,7 +828,7 @@ async def test_model_manager_handles_requester_not_found_gracefully(fake_request
# Provider should be skipped
assert len(model_mgr.provider_dict) == 0
- assert len(model_mgr.llm_models) == 0
+ assert len(model_mgr.llm_model_dict) == 0
model_mgr.ap.logger.warning.assert_called()
@@ -790,6 +845,85 @@ def test_requester_not_found_error_str():
assert error.requester_name == 'test-requester'
+@pytest.mark.asyncio
+async def test_runtime_cache_isolates_same_resource_uuid_between_workspaces(fake_requester_registry):
+ """A UUID collision cannot select another Workspace's runtime object."""
+
+ model_mgr = fake_requester_registry
+ await model_mgr.initialize()
+ contexts = {
+ workspace_uuid: ExecutionContext(
+ instance_uuid='test-instance',
+ workspace_uuid=workspace_uuid,
+ placement_generation=1,
+ )
+ for workspace_uuid in ('workspace-a', 'workspace-b')
+ }
+
+ async def resolve_binding(workspace_uuid, *, expected_generation=None):
+ assert expected_generation in (None, 1)
+ return WorkspaceExecutionBinding(
+ instance_uuid='test-instance',
+ workspace_uuid=workspace_uuid,
+ placement_generation=1,
+ write_fenced=False,
+ state='active',
+ )
+
+ model_mgr.ap.workspace_service.get_execution_binding = AsyncMock(side_effect=resolve_binding)
+ for workspace_uuid, context in contexts.items():
+ provider = await model_mgr.load_provider(
+ context,
+ {
+ 'uuid': 'shared-provider',
+ 'name': f'Provider {workspace_uuid}',
+ 'requester': 'fake-requester',
+ 'base_url': f'https://{workspace_uuid}.example.com',
+ 'api_keys': [],
+ },
+ )
+ await model_mgr.cache_provider(context, provider)
+ runtime_model = await model_mgr.load_llm_model_with_provider(
+ context,
+ persistence_model.LLMModel(
+ workspace_uuid=workspace_uuid,
+ uuid='shared-model',
+ name=f'Model {workspace_uuid}',
+ provider_uuid='shared-provider',
+ abilities=[],
+ extra_args={},
+ ),
+ provider,
+ )
+ await model_mgr.cache_llm_model(context, runtime_model)
+
+ workspace_a_model = await model_mgr.get_model_by_uuid(contexts['workspace-a'], 'shared-model')
+ workspace_b_model = await model_mgr.get_model_by_uuid(contexts['workspace-b'], 'shared-model')
+
+ assert workspace_a_model.model_entity.name == 'Model workspace-a'
+ assert workspace_b_model.model_entity.name == 'Model workspace-b'
+ assert workspace_a_model is not workspace_b_model
+
+
+@pytest.mark.asyncio
+async def test_runtime_cache_rejects_stale_placement_generation(fake_requester_registry):
+ """A stale generation is fenced before any cached model can be returned."""
+
+ model_mgr = fake_requester_registry
+ await model_mgr.initialize()
+ stale_context = TEST_EXECUTION_CONTEXT
+
+ async def reject_stale(_workspace_uuid, *, expected_generation=None):
+ if expected_generation == stale_context.placement_generation:
+ raise WorkspaceGenerationMismatchError('stale generation')
+ raise AssertionError('lookup must include the supplied generation')
+
+ model_mgr.ap.workspace_service.get_execution_binding = AsyncMock(side_effect=reject_stale)
+
+ with pytest.raises(WorkspaceGenerationMismatchError, match='stale generation'):
+ await model_mgr.get_model_by_uuid(stale_context, 'any-model')
+
+
def test_provider_not_found_error_str():
"""Test ProviderNotFoundError string representation."""
error = provider_errors.ProviderNotFoundError('test-provider')
diff --git a/tests/unit_tests/provider/test_model_service.py b/tests/unit_tests/provider/test_model_service.py
index b4e1b3ca8..ba184de2b 100644
--- a/tests/unit_tests/provider/test_model_service.py
+++ b/tests/unit_tests/provider/test_model_service.py
@@ -19,6 +19,8 @@ from langbot.pkg.provider.modelmgr import requester
from langbot.pkg.provider.modelmgr.modelmgr import ModelManager
from langbot.pkg.provider.modelmgr.token import TokenManager
from langbot.pkg.provider.runners.localagent import LocalAgentRunner
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.workspace.entities import WorkspaceExecutionBinding
def test_runtime_llm_model_data_preserves_uuid_after_update_payload_uuid_removed():
@@ -95,11 +97,22 @@ async def test_model_manager_initialize_skips_space_sync_after_timeout():
ap.discover = SimpleNamespace(get_components_by_kind=Mock(return_value=[]))
ap.instance_config = SimpleNamespace(data={'space': {'models_sync_timeout': 0.01}})
ap.logger = Mock()
+ binding = WorkspaceExecutionBinding(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-test',
+ placement_generation=1,
+ write_fenced=False,
+ state='active',
+ )
+ ap.workspace_service = SimpleNamespace(
+ get_local_execution_binding=AsyncMock(return_value=binding),
+ get_execution_binding=AsyncMock(return_value=binding),
+ )
mgr = ModelManager(ap)
mgr.load_models_from_db = AsyncMock()
- async def slow_sync():
+ async def slow_sync(_context):
await asyncio.sleep(1)
mgr.sync_new_models_from_space = AsyncMock(side_effect=slow_sync)
@@ -117,6 +130,14 @@ async def test_updated_llm_model_is_immediately_usable_by_local_agent_pipeline()
model_uuid = 'qwen-model-uuid'
provider_uuid = 'ollama-provider-uuid'
+ workspace_uuid = 'workspace-test'
+ execution_context = ExecutionContext(
+ instance_uuid='instance-test',
+ workspace_uuid=workspace_uuid,
+ placement_generation=1,
+ bot_uuid='bot-uuid',
+ pipeline_uuid='pipeline-uuid',
+ )
ap = SimpleNamespace()
ap.logger = Mock()
@@ -126,24 +147,63 @@ async def test_updated_llm_model_is_immediately_usable_by_local_agent_pipeline()
ap.plugin_connector = SimpleNamespace(
emit_event=AsyncMock(return_value=SimpleNamespace(event=SimpleNamespace(default_prompt=[], prompt=[])))
)
+ binding = WorkspaceExecutionBinding(
+ instance_uuid='instance-test',
+ workspace_uuid=workspace_uuid,
+ placement_generation=1,
+ write_fenced=False,
+ state='active',
+ )
+ ap.workspace_service = SimpleNamespace(get_execution_binding=AsyncMock(return_value=binding))
ap.model_mgr = ModelManager(ap)
- runtime_provider = Mock()
- ap.model_mgr.provider_dict = {provider_uuid: runtime_provider}
- ap.model_mgr.llm_models = [
- requester.RuntimeLLMModel(
- model_entity=persistence_model.LLMModel(
- uuid=model_uuid,
- name='old-qwen-name',
- provider_uuid=provider_uuid,
- abilities=[],
- extra_args={},
- ),
- provider=runtime_provider,
- )
- ]
+ runtime_provider = Mock(
+ execution_context=execution_context,
+ provider_entity=persistence_model.ModelProvider(
+ workspace_uuid=workspace_uuid,
+ uuid=provider_uuid,
+ name='Ollama',
+ requester='ollama',
+ base_url='http://localhost:11434',
+ api_keys=[],
+ ),
+ )
+ cache_key = ('instance-test', workspace_uuid, 1, provider_uuid)
+ ap.model_mgr.provider_dict = {cache_key: runtime_provider}
+ runtime_model = requester.RuntimeLLMModel(
+ execution_context=execution_context,
+ model_entity=persistence_model.LLMModel(
+ workspace_uuid=workspace_uuid,
+ uuid=model_uuid,
+ name='old-qwen-name',
+ provider_uuid=provider_uuid,
+ abilities=[],
+ extra_args={},
+ ),
+ provider=runtime_provider,
+ )
+ ap.model_mgr.llm_model_dict = {
+ ('instance-test', workspace_uuid, 1, model_uuid): runtime_model,
+ }
- await LLMModelsService(ap).update_llm_model(
+ ap.provider_service = SimpleNamespace(
+ get_provider=AsyncMock(return_value={'uuid': provider_uuid, 'workspace_uuid': workspace_uuid})
+ )
+ model_service = LLMModelsService(ap)
+ model_service.get_llm_model = AsyncMock(
+ return_value={
+ 'uuid': model_uuid,
+ 'workspace_uuid': workspace_uuid,
+ 'name': 'old-qwen-name',
+ 'provider_uuid': provider_uuid,
+ 'abilities': [],
+ 'context_length': None,
+ 'extra_args': {},
+ 'prefered_ranking': 0,
+ }
+ )
+ await model_service.update_llm_model(
+ workspace_uuid,
model_uuid,
{
'name': 'Qwen3.5-27B',
@@ -153,13 +213,17 @@ async def test_updated_llm_model_is_immediately_usable_by_local_agent_pipeline()
},
)
- runtime_model = await ap.model_mgr.get_model_by_uuid(model_uuid)
+ runtime_model = await ap.model_mgr.get_model_by_uuid(execution_context, model_uuid)
assert runtime_model.model_entity.uuid == model_uuid
assert runtime_model.model_entity.name == 'Qwen3.5-27B'
- session = SimpleNamespace(
+ session = provider_session.Session(
+ instance_uuid='instance-test',
+ workspace_uuid=workspace_uuid,
+ placement_generation=1,
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
+ bot_uuid='bot-uuid',
)
conversation = SimpleNamespace(
uuid='conversation-uuid',
@@ -194,6 +258,9 @@ async def test_updated_llm_model_is_immediately_usable_by_local_agent_pipeline()
'output': {'misc': {'remove-think': False}},
}
query = pipeline_query.Query.model_construct(
+ instance_uuid='instance-test',
+ workspace_uuid=workspace_uuid,
+ placement_generation=1,
query_id='query-id',
launcher_type=provider_session.LauncherTypes.PERSON,
launcher_id=12345,
@@ -215,6 +282,7 @@ async def test_updated_llm_model_is_immediately_usable_by_local_agent_pipeline()
resp_message_chain=None,
current_stage_name=None,
)
+ object.__setattr__(query, '_execution_context', execution_context)
result = await PreProcessor(ap).process(query, 'PreProcessor')
processed_query = result.new_query
diff --git a/tests/unit_tests/provider/test_requester_base.py b/tests/unit_tests/provider/test_requester_base.py
index 71c0da653..672c930be 100644
--- a/tests/unit_tests/provider/test_requester_base.py
+++ b/tests/unit_tests/provider/test_requester_base.py
@@ -15,6 +15,7 @@ from langbot.pkg.provider.modelmgr import requester
from langbot.pkg.provider.modelmgr import token
from langbot.pkg.entity.persistence import model as persistence_model
from langbot.pkg.provider.modelmgr.errors import RequesterError
+from tests.unit_tests.provider.conftest import TEST_EXECUTION_CONTEXT, TEST_WORKSPACE_UUID
# ============================================================================
@@ -134,6 +135,7 @@ async def test_requester_invoke_rerank_not_implemented():
# Create fake model
fake_provider_entity = persistence_model.ModelProvider(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid='provider-uuid',
name='Provider',
requester='test',
@@ -143,17 +145,20 @@ async def test_requester_invoke_rerank_not_implemented():
fake_token_mgr = token.TokenManager(name='test', tokens=[])
fake_requester = inst
fake_provider = requester.RuntimeProvider(
+ execution_context=TEST_EXECUTION_CONTEXT,
provider_entity=fake_provider_entity,
token_mgr=fake_token_mgr,
requester=fake_requester,
)
fake_model_entity = persistence_model.RerankModel(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid='model-uuid',
name='Model',
provider_uuid='provider-uuid',
extra_args={},
)
fake_model = requester.RuntimeRerankModel(
+ execution_context=TEST_EXECUTION_CONTEXT,
model_entity=fake_model_entity,
provider=fake_provider,
)
@@ -289,6 +294,7 @@ async def test_runtime_provider_invoke_llm_delegates(runtime_provider, runtime_l
resp_message_chain=None,
current_stage_name=None,
)
+ object.__setattr__(query, '_execution_context', TEST_EXECUTION_CONTEXT)
messages = [
provider_message.Message(role='user', content=[provider_message.ContentElement(type='text', text='Hello')])
@@ -332,6 +338,7 @@ async def test_runtime_provider_invoke_llm_stream_yields_chunks(runtime_provider
resp_message_chain=None,
current_stage_name=None,
)
+ object.__setattr__(query, '_execution_context', TEST_EXECUTION_CONTEXT)
messages = [
provider_message.Message(role='user', content=[provider_message.ContentElement(type='text', text='Hello')])
@@ -350,7 +357,11 @@ async def test_runtime_provider_invoke_embedding_returns_vectors(runtime_provide
"""Test RuntimeProvider.invoke_embedding returns embedding vectors."""
provider = runtime_provider
- result = await provider.invoke_embedding(runtime_embedding_model, ['text1', 'text2'])
+ result = await provider.invoke_embedding(
+ runtime_embedding_model,
+ ['text1', 'text2'],
+ execution_context=TEST_EXECUTION_CONTEXT,
+ )
assert len(result) == 2
assert result[0] == [0.1, 0.2, 0.3]
@@ -362,7 +373,12 @@ async def test_runtime_provider_invoke_rerank_returns_scores(runtime_provider, r
# Need to use the correct provider for rerank model
provider = runtime_rerank_model.provider
- result = await provider.invoke_rerank(runtime_rerank_model, 'query', ['doc1', 'doc2', 'doc3'])
+ result = await provider.invoke_rerank(
+ runtime_rerank_model,
+ 'query',
+ ['doc1', 'doc2', 'doc3'],
+ execution_context=TEST_EXECUTION_CONTEXT,
+ )
assert len(result) == 3
assert result[0]['index'] == 0
@@ -532,6 +548,7 @@ async def test_runtime_provider_invoke_llm_propagates_error(mock_app_for_modelmg
await requester_inst.initialize()
provider_entity = persistence_model.ModelProvider(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid='error-provider',
name='Error Provider',
requester='error-requester',
@@ -541,19 +558,25 @@ async def test_runtime_provider_invoke_llm_propagates_error(mock_app_for_modelmg
token_mgr = token.TokenManager(name='error-provider', tokens=['error-key'])
provider = requester.RuntimeProvider(
+ execution_context=TEST_EXECUTION_CONTEXT,
provider_entity=provider_entity,
token_mgr=token_mgr,
requester=requester_inst,
)
model_entity = persistence_model.LLMModel(
+ workspace_uuid=TEST_WORKSPACE_UUID,
uuid='error-model',
name='Error Model',
provider_uuid='error-provider',
abilities=[],
extra_args={},
)
- model = requester.RuntimeLLMModel(model_entity=model_entity, provider=provider)
+ model = requester.RuntimeLLMModel(
+ execution_context=TEST_EXECUTION_CONTEXT,
+ model_entity=model_entity,
+ provider=provider,
+ )
import langbot_plugin.api.entities.builtin.provider.message as provider_message
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
@@ -580,6 +603,7 @@ async def test_runtime_provider_invoke_llm_propagates_error(mock_app_for_modelmg
resp_message_chain=None,
current_stage_name=None,
)
+ object.__setattr__(query, '_execution_context', TEST_EXECUTION_CONTEXT)
messages = [
provider_message.Message(role='user', content=[provider_message.ContentElement(type='text', text='Hello')])
diff --git a/tests/unit_tests/provider/test_session_manager.py b/tests/unit_tests/provider/test_session_manager.py
index eca8cac8a..018917430 100644
--- a/tests/unit_tests/provider/test_session_manager.py
+++ b/tests/unit_tests/provider/test_session_manager.py
@@ -10,12 +10,73 @@ from __future__ import annotations
import pytest
import asyncio
+from types import SimpleNamespace
from unittest.mock import Mock
from importlib import import_module
import langbot_plugin.api.entities.builtin.provider.session as provider_session
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.pipeline.pool import (
+ ExecutionContextMismatchError,
+ ExecutionContextRequiredError,
+)
+
+
+TEST_CONTEXT = ExecutionContext(
+ instance_uuid='instance-test',
+ workspace_uuid='workspace-test',
+ placement_generation=1,
+)
+TEST_BOT_UUID = 'bot-123'
+
+
+def bind_query_context(query, *, context=TEST_CONTEXT, bot_uuid=TEST_BOT_UUID):
+ """Attach the trusted runtime scope expected by the session manager."""
+ query.bot_uuid = bot_uuid
+ query._execution_context = context
+ return query
+
+
+def bind_session_context(session, query):
+ """Make a mocked legacy Session belong to the Query execution scope."""
+ session.bot_uuid = query.bot_uuid
+ session._langbot_session_key = (
+ TEST_CONTEXT.instance_uuid,
+ TEST_CONTEXT.workspace_uuid,
+ TEST_CONTEXT.placement_generation,
+ query.bot_uuid,
+ query.launcher_type.value,
+ query.launcher_id,
+ )
+ return session
+
+
+def scoped_query(
+ *,
+ workspace_uuid='workspace-test',
+ bot_uuid=TEST_BOT_UUID,
+ placement_generation=1,
+ pipeline_uuid=None,
+):
+ """Create a small Query-like object with a complete trusted scope."""
+ context = ExecutionContext(
+ instance_uuid='instance-test',
+ workspace_uuid=workspace_uuid,
+ placement_generation=placement_generation,
+ bot_uuid=bot_uuid,
+ pipeline_uuid=pipeline_uuid,
+ query_uuid=f'query-{workspace_uuid}-{bot_uuid}',
+ )
+ return SimpleNamespace(
+ launcher_type=provider_session.LauncherTypes.PERSON,
+ launcher_id='same-launcher',
+ sender_id='same-sender',
+ bot_uuid=bot_uuid,
+ _execution_context=context,
+ )
+
def get_session_module():
"""Lazy import to avoid circular import issues."""
@@ -71,7 +132,7 @@ class TestSessionManagerGetSession:
query.launcher_type = provider_session.LauncherTypes.PERSON
query.launcher_id = '12345'
query.sender_id = '12345'
- return query
+ return bind_query_context(query)
@pytest.mark.asyncio
async def test_creates_new_session_when_not_found(self, mock_app_with_config, sample_query):
@@ -126,11 +187,13 @@ class TestSessionManagerGetSession:
query1.launcher_type = provider_session.LauncherTypes.PERSON
query1.launcher_id = 'user1'
query1.sender_id = 'user1'
+ bind_query_context(query1)
query2 = Mock(spec=pipeline_query.Query)
query2.launcher_type = provider_session.LauncherTypes.PERSON
query2.launcher_id = 'user2'
query2.sender_id = 'user2'
+ bind_query_context(query2)
session1 = await manager.get_session(query1)
session2 = await manager.get_session(query2)
@@ -149,11 +212,13 @@ class TestSessionManagerGetSession:
query1.launcher_type = provider_session.LauncherTypes.PERSON
query1.launcher_id = 'same_id'
query1.sender_id = 'same_id'
+ bind_query_context(query1)
query2 = Mock(spec=pipeline_query.Query)
query2.launcher_type = provider_session.LauncherTypes.GROUP
query2.launcher_id = 'same_id'
query2.sender_id = 'same_id'
+ bind_query_context(query2)
session1 = await manager.get_session(query1)
session2 = await manager.get_session(query2)
@@ -191,7 +256,7 @@ class TestSessionManagerGetConversation:
query.launcher_type = provider_session.LauncherTypes.PERSON
query.launcher_id = '12345'
query.sender_id = '12345'
- return query
+ return bind_query_context(query)
@pytest.mark.asyncio
async def test_creates_conversation_with_prompt(self, mock_app_with_config, sample_query, sample_session):
@@ -199,6 +264,7 @@ class TestSessionManagerGetConversation:
sessionmgr = get_session_module()
manager = sessionmgr.SessionManager(mock_app_with_config)
+ bind_session_context(sample_session, sample_query)
prompt_config = [{'role': 'system', 'content': 'You are a helpful assistant.'}]
pipeline_uuid = 'pipeline-123'
@@ -222,6 +288,7 @@ class TestSessionManagerGetConversation:
sessionmgr = get_session_module()
manager = sessionmgr.SessionManager(mock_app_with_config)
+ bind_session_context(sample_session, sample_query)
prompt_config = [{'role': 'system', 'content': 'You are a helpful assistant.'}]
pipeline_uuid = 'pipeline-123'
@@ -244,14 +311,15 @@ class TestSessionManagerGetConversation:
sessionmgr = get_session_module()
manager = sessionmgr.SessionManager(mock_app_with_config)
+ bind_session_context(sample_session, sample_query)
prompt_config = [{'role': 'system', 'content': 'You are a helpful assistant.'}]
# First call with pipeline1
- conv1 = await manager.get_conversation(sample_query, sample_session, prompt_config, 'pipeline-1', 'bot-1')
+ conv1 = await manager.get_conversation(sample_query, sample_session, prompt_config, 'pipeline-1', TEST_BOT_UUID)
# Second call with different pipeline should create new conversation
- conv2 = await manager.get_conversation(sample_query, sample_session, prompt_config, 'pipeline-2', 'bot-2')
+ conv2 = await manager.get_conversation(sample_query, sample_session, prompt_config, 'pipeline-2', TEST_BOT_UUID)
assert conv1 is not conv2
assert len(sample_session.conversations) == 2
@@ -263,6 +331,7 @@ class TestSessionManagerGetConversation:
sessionmgr = get_session_module()
manager = sessionmgr.SessionManager(mock_app_with_config)
+ bind_session_context(sample_session, sample_query)
prompt_config = [{'role': 'system', 'content': 'You are a helpful assistant.'}]
@@ -278,6 +347,7 @@ class TestSessionManagerGetConversation:
sessionmgr = get_session_module()
manager = sessionmgr.SessionManager(mock_app_with_config)
+ bind_session_context(sample_session, sample_query)
prompt_config = [{'role': 'system', 'content': 'System message'}, {'role': 'user', 'content': 'User message'}]
@@ -287,3 +357,72 @@ class TestSessionManagerGetConversation:
assert conversation.prompt.name == 'default'
assert len(conversation.prompt.messages) == 2
+
+
+class TestSessionManagerWorkspaceIsolation:
+ """Regression coverage for workspace, bot, and placement fencing."""
+
+ @staticmethod
+ def manager():
+ mock_app = Mock()
+ mock_app.instance_config.data = {'concurrency': {'session': 5}}
+ return get_session_module().SessionManager(mock_app)
+
+ @pytest.mark.asyncio
+ async def test_get_session_requires_trusted_query_scope(self):
+ query = SimpleNamespace(
+ launcher_type=provider_session.LauncherTypes.PERSON,
+ launcher_id='same-launcher',
+ sender_id='same-sender',
+ bot_uuid=TEST_BOT_UUID,
+ )
+
+ with pytest.raises(ExecutionContextRequiredError):
+ await self.manager().get_session(query)
+
+ @pytest.mark.asyncio
+ async def test_same_launcher_in_two_workspaces_does_not_share_session(self):
+ manager = self.manager()
+
+ first = await manager.get_session(scoped_query(workspace_uuid='workspace-a'))
+ second = await manager.get_session(scoped_query(workspace_uuid='workspace-b'))
+
+ assert first is not second
+ assert len(manager.session_list) == 2
+
+ @pytest.mark.asyncio
+ async def test_same_launcher_in_two_bots_does_not_share_session(self):
+ manager = self.manager()
+
+ first = await manager.get_session(scoped_query(bot_uuid='bot-a'))
+ second = await manager.get_session(scoped_query(bot_uuid='bot-b'))
+
+ assert first is not second
+
+ @pytest.mark.asyncio
+ async def test_new_placement_generation_does_not_reuse_old_session(self):
+ manager = self.manager()
+
+ first = await manager.get_session(scoped_query(placement_generation=1))
+ second = await manager.get_session(scoped_query(placement_generation=2))
+
+ assert first is not second
+
+ @pytest.mark.asyncio
+ async def test_conversation_rejects_session_from_another_workspace(self):
+ manager = self.manager()
+ query_a = scoped_query(workspace_uuid='workspace-a')
+ query_b = scoped_query(workspace_uuid='workspace-b')
+ session_a = await manager.get_session(query_a)
+
+ with pytest.raises(ExecutionContextMismatchError):
+ await manager.get_conversation(query_b, session_a, [], 'pipeline-1', TEST_BOT_UUID)
+
+ @pytest.mark.asyncio
+ async def test_conversation_rejects_substituted_bot_argument(self):
+ manager = self.manager()
+ query = scoped_query(bot_uuid='bot-a')
+ session = await manager.get_session(query)
+
+ with pytest.raises(ExecutionContextMismatchError):
+ await manager.get_conversation(query, session, [], 'pipeline-1', 'bot-b')
diff --git a/tests/unit_tests/provider/test_skill_tools.py b/tests/unit_tests/provider/test_skill_tools.py
index 9db7b945e..bdc2410d3 100644
--- a/tests/unit_tests/provider/test_skill_tools.py
+++ b/tests/unit_tests/provider/test_skill_tools.py
@@ -7,6 +7,37 @@ from unittest.mock import AsyncMock, Mock
import pytest
+from langbot.pkg.api.http.context import ExecutionContext
+
+
+_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ query_uuid='query-a',
+)
+
+
+def _make_query(*, variables=None, **kwargs):
+ return SimpleNamespace(
+ query_id=kwargs.pop('query_id', 'query-a'),
+ query_uuid=kwargs.pop('query_uuid', 'query-a'),
+ instance_uuid=kwargs.pop('instance_uuid', _CONTEXT.instance_uuid),
+ workspace_uuid=kwargs.pop('workspace_uuid', _CONTEXT.workspace_uuid),
+ placement_generation=kwargs.pop('placement_generation', _CONTEXT.placement_generation),
+ variables={} if variables is None else variables,
+ **kwargs,
+ )
+
+
+def _make_skill_manager(skills: dict[str, dict], **kwargs):
+ return SimpleNamespace(
+ skills=skills,
+ get_skills=Mock(return_value=skills),
+ get_skill_by_name=Mock(side_effect=lambda _context, name: skills.get(name)),
+ **kwargs,
+ )
+
def _make_ap(logger=None):
ap = SimpleNamespace()
@@ -51,14 +82,14 @@ class TestSkillManagerCache:
mgr = SkillManager(ap)
# Empty cache → returns False
- assert mgr.refresh_skill_from_disk('test-skill') is False
+ assert mgr.refresh_skill_from_disk(_CONTEXT, 'test-skill') is False
# Cache populated → returns True; method does NOT mutate the cache
cached = _make_skill_data(name='test-skill', instructions='Cached')
- mgr.skills['test-skill'] = cached
- assert mgr.refresh_skill_from_disk('test-skill') is True
- assert mgr.skills['test-skill'] is cached
- assert mgr.refresh_skill_from_disk('') is False
+ mgr._skills_by_scope[mgr._scope_key(_CONTEXT)] = {'test-skill': cached}
+ assert mgr.refresh_skill_from_disk(_CONTEXT, 'test-skill') is True
+ assert mgr.get_skills(_CONTEXT)['test-skill'] is cached
+ assert mgr.refresh_skill_from_disk(_CONTEXT, '') is False
@pytest.mark.asyncio
async def test_reload_skills_drops_box_skills_with_missing_package_root(self):
@@ -85,9 +116,9 @@ class TestSkillManagerCache:
ap.box_service = box_service
mgr = SkillManager(ap)
- await mgr.reload_skills()
+ await mgr.reload_skills(_CONTEXT)
- assert list(mgr.skills) == ['alive']
+ assert list(mgr.get_skills(_CONTEXT)) == ['alive']
# Warning fired with the dropped skill name so operators can see it.
warning_messages = [str(call.args[0]) for call in ap.logger.warning.call_args_list]
assert any('ghost' in msg and 'package_root missing' in msg for msg in warning_messages)
@@ -116,9 +147,9 @@ class TestSkillManagerCache:
ap.box_service = box_service
mgr = SkillManager(ap)
- await mgr.reload_skills()
+ await mgr.reload_skills(_CONTEXT)
- assert sorted(mgr.skills) == ['alpha', 'beta']
+ assert sorted(mgr.get_skills(_CONTEXT)) == ['alpha', 'beta']
# No skill dropped → no "package_root missing" warning.
warning_messages = [str(call.args[0]) for call in ap.logger.warning.call_args_list]
assert not any('package_root missing' in msg for msg in warning_messages)
@@ -141,12 +172,12 @@ class TestSkillActivationHelper:
ap = _make_ap()
mgr = SkillManager(ap)
- mgr.skills = {
+ mgr._skills_by_scope[mgr._scope_key(_CONTEXT)] = {
'primary': _make_skill_data(name='primary', instructions='Primary instructions'),
}
ap.skill_mgr = mgr
- query = SimpleNamespace(variables={})
+ query = _make_query()
assert register_activated_skill(ap, query, 'primary') is True
assert set(query.variables[ACTIVATED_SKILLS_KEY].keys()) == {'primary'}
@@ -159,10 +190,10 @@ class TestSkillActivationHelper:
ap = _make_ap()
mgr = SkillManager(ap)
- mgr.skills = {'primary': _make_skill_data(name='primary')}
+ mgr._skills_by_scope[mgr._scope_key(_CONTEXT)] = {'primary': _make_skill_data(name='primary')}
ap.skill_mgr = mgr
- query = SimpleNamespace(variables={})
+ query = _make_query()
assert register_activated_skill(ap, query, 'missing') is False
assert ACTIVATED_SKILLS_KEY not in query.variables
@@ -171,7 +202,7 @@ class TestSkillActivationHelper:
from langbot.pkg.skill.activation import register_activated_skill
ap = _make_ap() # no skill_mgr attribute
- query = SimpleNamespace(variables={})
+ query = _make_query()
assert register_activated_skill(ap, query, 'primary') is False
@@ -181,13 +212,13 @@ class TestSkillPathHelpers:
from langbot.pkg.provider.tools.loaders.skill import PIPELINE_BOUND_SKILLS_KEY, get_visible_skills
ap = _make_ap()
- ap.skill_mgr = SimpleNamespace(
- skills={
+ ap.skill_mgr = _make_skill_manager(
+ {
'visible': _make_skill_data(name='visible'),
'hidden': _make_skill_data(name='hidden'),
}
)
- query = SimpleNamespace(variables={PIPELINE_BOUND_SKILLS_KEY: ['visible']})
+ query = _make_query(variables={PIPELINE_BOUND_SKILLS_KEY: ['visible']})
result = get_visible_skills(ap, query)
@@ -202,13 +233,13 @@ class TestSkillPathHelpers:
)
ap = _make_ap()
- ap.skill_mgr = SimpleNamespace(
- skills={
+ ap.skill_mgr = _make_skill_manager(
+ {
'visible': _make_skill_data(name='visible'),
'hidden': _make_skill_data(name='hidden'),
}
)
- query = SimpleNamespace(variables={PIPELINE_BOUND_SKILLS_KEY: ['visible']})
+ query = _make_query(variables={PIPELINE_BOUND_SKILLS_KEY: ['visible']})
restored = restore_activated_skills(ap, query, ['visible', 'hidden', 'visible', ''])
@@ -223,8 +254,8 @@ class TestSkillPathHelpers:
)
ap = _make_ap()
- ap.skill_mgr = SimpleNamespace(skills={'demo': _make_skill_data(name='demo')})
- query = SimpleNamespace(variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']})
+ ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo')})
+ query = _make_query(variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']})
skill, rewritten = resolve_virtual_skill_path(
ap,
@@ -292,13 +323,10 @@ class TestSkillToolLoader:
skill = _make_skill_data(name='demo', package_root='/data/skills/demo', instructions='Step 1')
ap = _make_ap()
- ap.skill_mgr = SimpleNamespace(
- skills={'demo': skill},
- get_skill_by_name=lambda name: skill if name == 'demo' else None,
- )
+ ap.skill_mgr = _make_skill_manager({'demo': skill})
loader = SkillToolLoader(ap)
- query = SimpleNamespace(variables={})
+ query = _make_query()
result = await loader.invoke_tool(ACTIVATE_SKILL_TOOL_NAME, {'skill_name': 'demo'}, query)
@@ -317,10 +345,7 @@ class TestSkillToolLoader:
)
ap = _make_ap()
- ap.skill_mgr = SimpleNamespace(
- skills={'demo': _make_skill_data(name='demo')},
- get_skill_by_name=lambda name: None,
- )
+ ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo')})
loader = SkillToolLoader(ap)
@@ -328,7 +353,7 @@ class TestSkillToolLoader:
await loader.invoke_tool(
ACTIVATE_SKILL_TOOL_NAME,
{'skill_name': 'ghost'},
- SimpleNamespace(variables={}),
+ _make_query(),
)
@pytest.mark.asyncio
@@ -362,18 +387,19 @@ class TestSkillToolLoader:
result = await loader.invoke_tool(
REGISTER_SKILL_TOOL_NAME,
{'path': '/workspace/repo'},
- SimpleNamespace(),
+ _make_query(),
)
- ap.skill_service.scan_directory_async.assert_awaited_once_with(os.path.realpath(repo_dir))
+ ap.skill_service.scan_directory_async.assert_awaited_once_with(_CONTEXT, os.path.realpath(repo_dir))
ap.skill_service.create_skill.assert_awaited_once_with(
+ _CONTEXT,
{
'name': 'cloned-skill',
'display_name': 'Cloned Skill',
'description': 'Imported from clone',
'instructions': 'Do work',
'package_root': os.path.realpath(repo_dir),
- }
+ },
)
assert result['registered'] is True
assert result['skill_name'] == 'cloned-skill'
@@ -397,7 +423,7 @@ class TestSkillToolLoader:
await loader.invoke_tool(
REGISTER_SKILL_TOOL_NAME,
{'path': '/workspace/../../etc'},
- SimpleNamespace(),
+ _make_query(),
)
@pytest.mark.asyncio
@@ -417,7 +443,7 @@ class TestSkillToolLoader:
await loader.invoke_tool(
REGISTER_SKILL_TOOL_NAME,
{'path': '/workspace/foo'},
- SimpleNamespace(),
+ _make_query(),
)
@pytest.mark.asyncio
@@ -428,7 +454,7 @@ class TestSkillToolLoader:
ap.skill_mgr = SimpleNamespace(skills={})
ap.box_service = SimpleNamespace(
available=True,
- get_status=AsyncMock(return_value={'backend': {'available': False}}),
+ get_backend_status=AsyncMock(return_value={'backend': {'available': False}}),
)
loader = SkillToolLoader(ap)
@@ -443,10 +469,10 @@ class TestSkillToolLoader:
from langbot.pkg.provider.tools.loaders.skill_authoring import SkillToolLoader
ap = _make_ap()
- ap.skill_mgr = SimpleNamespace(skills={'demo': _make_skill_data(name='demo')})
+ ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo')})
ap.box_service = SimpleNamespace(
available=True,
- get_status=AsyncMock(return_value={'backend': {'available': True}}),
+ get_backend_status=AsyncMock(return_value={'backend': {'available': True}}),
)
loader = SkillToolLoader(ap)
@@ -472,13 +498,13 @@ class TestNativeToolLoaderSkillPaths:
ap = _make_ap()
ap.box_service = SimpleNamespace(available=True, default_workspace=tmpdir)
- ap.skill_mgr = SimpleNamespace(skills={'demo': _make_skill_data(name='demo', package_root=tmpdir)})
+ ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo', package_root=tmpdir)})
loader = NativeToolLoader(ap)
result = await loader.invoke_tool(
'read',
{'path': '/workspace/.skills/demo/SKILL.md'},
- SimpleNamespace(query_id='q1', variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']}),
+ _make_query(query_id='q1', variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']}),
)
assert result['ok'] is True
@@ -500,7 +526,7 @@ class TestNativeToolLoaderSkillPaths:
ap.skill_mgr = SimpleNamespace(refresh_skill_from_disk=Mock())
loader = NativeToolLoader(ap)
- query = SimpleNamespace(query_id='q1', launcher_type='person', launcher_id='123', variables={})
+ query = _make_query(query_id='q1', launcher_type='person', launcher_id='123')
register_activated_skill(query, _make_skill_data(name='demo', package_root=tmpdir))
result = await loader.invoke_tool(
@@ -516,7 +542,7 @@ class TestNativeToolLoaderSkillPaths:
tool_parameters = ap.box_service.execute_tool.await_args.args[0]
assert tool_parameters['command'] == 'python /workspace/.skills/demo/scripts/run.py'
assert tool_parameters['workdir'] == '/workspace/.skills/demo'
- ap.skill_mgr.refresh_skill_from_disk.assert_called_once_with('demo')
+ ap.skill_mgr.refresh_skill_from_disk.assert_called_once_with(_CONTEXT, 'demo')
@pytest.mark.asyncio
async def test_write_requires_skill_activation(self):
@@ -526,10 +552,10 @@ class TestNativeToolLoaderSkillPaths:
with tempfile.TemporaryDirectory() as tmpdir:
ap = _make_ap()
ap.box_service = SimpleNamespace(available=True, default_workspace=tmpdir)
- ap.skill_mgr = SimpleNamespace(skills={'demo': _make_skill_data(name='demo', package_root=tmpdir)})
+ ap.skill_mgr = _make_skill_manager({'demo': _make_skill_data(name='demo', package_root=tmpdir)})
loader = NativeToolLoader(ap)
- query = SimpleNamespace(query_id='q1', variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']})
+ query = _make_query(query_id='q1', variables={PIPELINE_BOUND_SKILLS_KEY: ['demo']})
with pytest.raises(ValueError, match='Skill "demo" is not available at this path'):
await loader.invoke_tool(
diff --git a/tests/unit_tests/provider/test_tool_manager.py b/tests/unit_tests/provider/test_tool_manager.py
index 0ae33115c..b61fc9d0f 100644
--- a/tests/unit_tests/provider/test_tool_manager.py
+++ b/tests/unit_tests/provider/test_tool_manager.py
@@ -14,6 +14,15 @@ from importlib import import_module
import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
import langbot_plugin.api.entities.builtin.pipeline.query as pipeline_query
+from langbot.pkg.api.http.context import ExecutionContext
+
+
+_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+)
+
def get_toolmgr_module():
"""Lazy import to avoid circular import issues."""
@@ -188,6 +197,10 @@ class TestToolManagerExecuteFuncCall:
def sample_query(self):
"""Create sample query for testing."""
query = Mock(spec=pipeline_query.Query)
+ query.bot_uuid = None
+ query.pipeline_uuid = None
+ query.query_uuid = None
+ query._execution_context = _CONTEXT
return query
@pytest.mark.asyncio
diff --git a/tests/unit_tests/provider/test_tool_manager_native.py b/tests/unit_tests/provider/test_tool_manager_native.py
index 2085a9e8c..2e45541d8 100644
--- a/tests/unit_tests/provider/test_tool_manager_native.py
+++ b/tests/unit_tests/provider/test_tool_manager_native.py
@@ -12,6 +12,14 @@ import langbot_plugin.api.entities.builtin.resource.tool as resource_tool
from langbot.pkg.provider.tools.loaders.native import NativeToolLoader
from langbot.pkg.provider.tools.toolmgr import ToolManager
+from langbot.pkg.api.http.context import ExecutionContext
+
+
+_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+)
class StubLoader:
@@ -72,7 +80,7 @@ async def test_tool_manager_omits_skill_authoring_tools_by_default():
manager.plugin_tool_loader = StubLoader([make_tool('plugin_tool')])
manager.mcp_tool_loader = StubLoader([make_tool('mcp_tool')])
- tools = await manager.get_all_tools()
+ tools = await manager.get_all_tools(_CONTEXT)
assert [tool.name for tool in tools] == ['exec', 'plugin_tool', 'mcp_tool']
@@ -85,7 +93,7 @@ async def test_tool_manager_includes_skill_authoring_tools_when_requested():
manager.plugin_tool_loader = StubLoader([make_tool('plugin_tool')])
manager.mcp_tool_loader = StubLoader([make_tool('mcp_tool')])
- tools = await manager.get_all_tools(include_skill_authoring=True)
+ tools = await manager.get_all_tools(_CONTEXT, include_skill_authoring=True)
assert [tool.name for tool in tools] == ['exec', 'activate', 'plugin_tool', 'mcp_tool']
@@ -102,7 +110,7 @@ async def test_tool_manager_catalog_labels_tool_sources():
)
manager.mcp_tool_loader = StubLoader([make_tool('mcp_tool')])
- catalog = await manager.get_tool_catalog(include_skill_authoring=True)
+ catalog = await manager.get_tool_catalog(_CONTEXT, include_skill_authoring=True)
assert [(item['name'], item['source'], item['source_name']) for item in catalog] == [
('exec', 'builtin', 'LangBot'),
@@ -139,7 +147,7 @@ async def test_native_tool_loader_hides_tools_when_box_unavailable():
async def test_native_tool_loader_exposes_all_tools_when_box_available():
box_service = SimpleNamespace(
available=True,
- get_status=AsyncMock(return_value={'backend': {'available': True}}),
+ get_backend_status=AsyncMock(return_value={'backend': {'available': True}}),
)
loader = NativeToolLoader(SimpleNamespace(box_service=box_service, logger=Mock()))
await loader.initialize()
@@ -156,15 +164,26 @@ async def test_native_tool_loader_exposes_all_tools_when_box_available():
def _make_loader_with_workspace(tmpdir: str) -> tuple[NativeToolLoader, Mock]:
logger = Mock()
- box_service = SimpleNamespace(available=True, default_workspace=tmpdir)
+ box_service = SimpleNamespace(
+ available=True,
+ default_workspace=tmpdir,
+ _tenant_workspace=Mock(return_value=tmpdir),
+ )
ap = SimpleNamespace(box_service=box_service, logger=logger)
return NativeToolLoader(ap), logger
-def _make_query() -> Mock:
- q = Mock()
- q.query_id = 'test-query-1'
- return q
+def _make_query() -> SimpleNamespace:
+ return SimpleNamespace(
+ query_id='test-query-1',
+ query_uuid='test-query-1',
+ instance_uuid=_CONTEXT.instance_uuid,
+ workspace_uuid=_CONTEXT.workspace_uuid,
+ placement_generation=_CONTEXT.placement_generation,
+ bot_uuid=None,
+ pipeline_uuid=None,
+ variables={},
+ )
@pytest.mark.asyncio
@@ -376,13 +395,13 @@ async def test_box_availability_helper_handles_unavailable_and_errors():
unavailable_backend = SimpleNamespace(
available=True,
- get_status=AsyncMock(return_value={'backend': {'available': False}}),
+ get_backend_status=AsyncMock(return_value={'backend': {'available': False}}),
)
assert await is_box_backend_available(SimpleNamespace(box_service=unavailable_backend)) is False
failing_backend = SimpleNamespace(
available=True,
- get_status=AsyncMock(side_effect=RuntimeError('box unavailable')),
+ get_backend_status=AsyncMock(side_effect=RuntimeError('box unavailable')),
)
assert await is_box_backend_available(SimpleNamespace(box_service=failing_backend)) is False
diff --git a/tests/unit_tests/rag/test_file_storage.py b/tests/unit_tests/rag/test_file_storage.py
index d4a6f2239..b2c1800c1 100644
--- a/tests/unit_tests/rag/test_file_storage.py
+++ b/tests/unit_tests/rag/test_file_storage.py
@@ -9,7 +9,27 @@ from unittest.mock import AsyncMock, Mock
import pytest
+from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.rag.knowledge.kbmgr import RuntimeKnowledgeBase
+from langbot.pkg.storage.mgr import StorageMgr
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
+
+
+WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
+CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid=WORKSPACE_A,
+ placement_generation=2,
+)
+
+
+def _upload_key(logical_key: str, *, context: ExecutionContext = CONTEXT) -> str:
+ return StorageMgr.scoped_object_key(
+ context,
+ owner_type='upload_document',
+ owner='account:test',
+ key=logical_key,
+ )
def _make_zip_bytes(entries: dict[str, bytes]) -> bytes:
@@ -25,26 +45,38 @@ def _make_app() -> Mock:
app = Mock()
app.logger = Mock()
app.task_mgr = Mock()
- app.storage_mgr = Mock()
- app.storage_mgr.storage_provider = Mock()
- app.storage_mgr.storage_provider.exists = AsyncMock(return_value=True)
- app.storage_mgr.storage_provider.load = AsyncMock()
- app.storage_mgr.storage_provider.save = AsyncMock()
- app.storage_mgr.storage_provider.size = AsyncMock(return_value=123)
- app.storage_mgr.storage_provider.delete = AsyncMock()
+ storage_mgr = StorageMgr(app)
+ storage_mgr.storage_provider = Mock()
+ storage_mgr.storage_provider.exists = AsyncMock(return_value=True)
+ storage_mgr.storage_provider.load = AsyncMock()
+ storage_mgr.storage_provider.save = AsyncMock()
+ storage_mgr.storage_provider.size = AsyncMock(return_value=123)
+ storage_mgr.storage_provider.delete = AsyncMock()
+ app.storage_mgr = storage_mgr
app.persistence_mgr = Mock()
app.persistence_mgr.execute_async = AsyncMock()
app.plugin_connector = Mock()
+ app.plugin_connector.require_workspace_context = AsyncMock(side_effect=lambda context: context)
+ app.workspace_service = SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ return_value=SimpleNamespace(
+ instance_uuid=CONTEXT.instance_uuid,
+ workspace_uuid=CONTEXT.workspace_uuid,
+ placement_generation=CONTEXT.placement_generation,
+ )
+ )
+ )
return app
def _make_kb(plugin_id: str | None = 'author/engine') -> RuntimeKnowledgeBase:
kb_entity = Mock()
kb_entity.uuid = 'test-kb-uuid'
+ kb_entity.workspace_uuid = WORKSPACE_A
kb_entity.collection_id = 'test-collection'
kb_entity.creation_settings = {}
kb_entity.knowledge_engine_plugin_id = plugin_id
- return RuntimeKnowledgeBase(_make_app(), kb_entity)
+ return RuntimeKnowledgeBase(_make_app(), kb_entity, CONTEXT)
class TestStoreFile:
@@ -58,27 +90,65 @@ class TestStoreFile:
kb.ap.task_mgr.create_user_task = Mock(side_effect=create_user_task)
- task_id = await kb.store_file('documents/test.pdf')
+ object_key = _upload_key('documents/test.pdf')
+ task_id = await kb.store_file(CONTEXT, object_key)
assert task_id == 'task-1'
- kb.ap.storage_mgr.storage_provider.exists.assert_awaited_once_with('documents/test.pdf')
+ kb.ap.storage_mgr.storage_provider.exists.assert_awaited_once_with(object_key)
kb.ap.persistence_mgr.execute_async.assert_awaited_once()
call_kwargs = kb.ap.task_mgr.create_user_task.call_args.kwargs
assert call_kwargs['kind'] == 'knowledge-operation'
- assert call_kwargs['name'] == 'knowledge-store-file-documents/test.pdf'
- assert call_kwargs['label'] == 'Store file documents/test.pdf'
+ assert call_kwargs['name'] == f'knowledge-store-file-{object_key}'
+ assert call_kwargs['label'] == f'Store file {object_key}'
@pytest.mark.asyncio
async def test_store_file_raises_when_source_file_missing(self):
kb = _make_kb()
kb.ap.storage_mgr.storage_provider.exists = AsyncMock(return_value=False)
- with pytest.raises(Exception, match='File missing.pdf not found'):
- await kb.store_file('missing.pdf')
+ object_key = _upload_key('missing.pdf')
+ with pytest.raises(WorkspaceNotFoundError, match='Upload not found'):
+ await kb.store_file(CONTEXT, object_key)
kb.ap.persistence_mgr.execute_async.assert_not_awaited()
kb.ap.task_mgr.create_user_task.assert_not_called()
+ @pytest.mark.asyncio
+ async def test_store_file_rejects_cross_workspace_upload_key(self):
+ kb = _make_kb()
+ other_context = ExecutionContext(
+ instance_uuid=CONTEXT.instance_uuid,
+ workspace_uuid='00000000-0000-0000-0000-00000000000b',
+ placement_generation=CONTEXT.placement_generation,
+ )
+
+ with pytest.raises(WorkspaceNotFoundError, match='Upload not found'):
+ await kb.store_file(CONTEXT, _upload_key('stolen.pdf', context=other_context))
+
+ kb.ap.storage_mgr.storage_provider.exists.assert_not_awaited()
+ kb.ap.persistence_mgr.execute_async.assert_not_awaited()
+
+ @pytest.mark.asyncio
+ async def test_store_file_rejects_stale_generation_and_wrong_owner_type(self):
+ kb = _make_kb()
+ stale_context = ExecutionContext(
+ instance_uuid=CONTEXT.instance_uuid,
+ workspace_uuid=CONTEXT.workspace_uuid,
+ placement_generation=CONTEXT.placement_generation + 1,
+ )
+ plugin_key = StorageMgr.scoped_object_key(
+ CONTEXT,
+ owner_type='plugin_config',
+ owner='plugin:test',
+ key='config.pdf',
+ )
+
+ for object_key in (_upload_key('stale.pdf', context=stale_context), plugin_key, 'raw.pdf'):
+ with pytest.raises(WorkspaceNotFoundError, match='Upload not found'):
+ await kb.store_file(CONTEXT, object_key)
+
+ kb.ap.storage_mgr.storage_provider.exists.assert_not_awaited()
+
class TestStoreZipFile:
@pytest.mark.asyncio
@@ -99,19 +169,22 @@ class TestStoreZipFile:
)
kb.store_file = AsyncMock(side_effect=['task-pdf', 'task-txt', 'task-md', 'task-html'])
- task_id = await kb._store_zip_file('archive.zip', parser_plugin_id='parser/plugin')
+ zip_key = _upload_key('archive.zip')
+ task_id = await kb._store_zip_file(CONTEXT, zip_key, parser_plugin_id='parser/plugin')
assert task_id == 'task-pdf'
assert kb.ap.storage_mgr.storage_provider.save.await_count == 4
saved_names = [call.args[0] for call in kb.ap.storage_mgr.storage_provider.save.await_args_list]
- assert any(name.startswith('doc1_') and name.endswith('.pdf') for name in saved_names)
- assert any(name.startswith('doc2_') and name.endswith('.txt') for name in saved_names)
- assert any(name.startswith('subdir_doc3_') and name.endswith('.md') for name in saved_names)
- assert any(name.startswith('page_') and name.endswith('.html') for name in saved_names)
- assert not any('image' in name for name in saved_names)
- assert not any('hidden' in name for name in saved_names)
- assert not any('__MACOSX' in name for name in saved_names)
- kb.ap.storage_mgr.storage_provider.delete.assert_awaited_once_with('archive.zip')
+ assert {name.rsplit('.', 1)[-1] for name in saved_names} == {'pdf', 'txt', 'md', 'html'}
+ for name in saved_names:
+ StorageMgr.require_scoped_object_key(
+ CONTEXT,
+ name,
+ expected_owner_type='upload_document',
+ )
+ forwarded_keys = [call.args[1] for call in kb.store_file.await_args_list]
+ assert forwarded_keys == saved_names
+ kb.ap.storage_mgr.storage_provider.delete.assert_awaited_once_with(zip_key)
@pytest.mark.asyncio
async def test_store_zip_file_raises_when_no_supported_files(self):
@@ -122,10 +195,10 @@ class TestStoreZipFile:
kb.store_file = AsyncMock()
with pytest.raises(Exception, match='No supported files found'):
- await kb._store_zip_file('archive.zip')
+ await kb._store_zip_file(CONTEXT, _upload_key('archive.zip'))
kb.store_file.assert_not_awaited()
- kb.ap.storage_mgr.storage_provider.delete.assert_awaited_once_with('archive.zip')
+ kb.ap.storage_mgr.storage_provider.delete.assert_awaited_once_with(_upload_key('archive.zip'))
class TestStoreFileTask:
@@ -133,29 +206,31 @@ class TestStoreFileTask:
async def test_store_file_task_marks_completed_and_cleans_storage(self):
kb = _make_kb()
kb._ingest_document = AsyncMock(return_value={'status': 'completed'})
- file_obj = SimpleNamespace(uuid='file-uuid', file_name='test.pdf', extension='pdf')
+ object_key = _upload_key('test.pdf')
+ file_obj = SimpleNamespace(uuid='file-uuid', file_name=object_key, extension='pdf')
task_context = Mock()
- await kb._store_file_task(file_obj, task_context)
+ await kb._store_file_task(CONTEXT, file_obj, task_context)
task_context.set_current_action.assert_called_once_with('Processing file')
- kb.ap.storage_mgr.storage_provider.size.assert_awaited_once_with('test.pdf')
+ kb.ap.storage_mgr.storage_provider.size.assert_awaited_once_with(object_key)
kb._ingest_document.assert_awaited_once()
assert kb.ap.persistence_mgr.execute_async.await_count == 2
- kb.ap.storage_mgr.storage_provider.delete.assert_awaited_once_with('test.pdf')
+ kb.ap.storage_mgr.storage_provider.delete.assert_awaited_once_with(object_key)
@pytest.mark.asyncio
async def test_store_file_task_marks_failed_and_cleans_storage(self):
kb = _make_kb()
kb._ingest_document = AsyncMock(return_value={'status': 'failed', 'error_message': 'parser failed'})
- file_obj = SimpleNamespace(uuid='file-uuid', file_name='bad.pdf', extension='pdf')
+ object_key = _upload_key('bad.pdf')
+ file_obj = SimpleNamespace(uuid='file-uuid', file_name=object_key, extension='pdf')
task_context = Mock()
with pytest.raises(Exception, match='parser failed'):
- await kb._store_file_task(file_obj, task_context)
+ await kb._store_file_task(CONTEXT, file_obj, task_context)
assert kb.ap.persistence_mgr.execute_async.await_count == 2
- kb.ap.storage_mgr.storage_provider.delete.assert_awaited_once_with('bad.pdf')
+ kb.ap.storage_mgr.storage_provider.delete.assert_awaited_once_with(object_key)
class TestDeleteDocument:
@@ -163,7 +238,7 @@ class TestDeleteDocument:
async def test_delete_document_returns_false_when_no_plugin_id(self):
kb = _make_kb(plugin_id=None)
- result = await kb._delete_document('doc-id')
+ result = await kb._delete_document(CONTEXT, 'doc-id')
assert result is False
@@ -172,7 +247,7 @@ class TestDeleteDocument:
kb = _make_kb()
kb.ap.plugin_connector.call_rag_delete_document = AsyncMock(return_value=True)
- result = await kb._delete_document('doc-id')
+ result = await kb._delete_document(CONTEXT, 'doc-id')
assert result is True
kb.ap.plugin_connector.call_rag_delete_document.assert_awaited_once_with(
@@ -184,7 +259,7 @@ class TestDeleteDocument:
kb = _make_kb()
kb.ap.plugin_connector.call_rag_delete_document = AsyncMock(side_effect=Exception('plugin error'))
- result = await kb._delete_document('doc-id')
+ result = await kb._delete_document(CONTEXT, 'doc-id')
assert result is False
kb.ap.logger.error.assert_called_once()
diff --git a/tests/unit_tests/rag/test_kbmgr.py b/tests/unit_tests/rag/test_kbmgr.py
index a1a16118d..ba8fa0faa 100644
--- a/tests/unit_tests/rag/test_kbmgr.py
+++ b/tests/unit_tests/rag/test_kbmgr.py
@@ -1,137 +1,395 @@
-"""Unit tests for RAG knowledge base manager.
-
-Tests cover:
-- RAGManager CRUD operations
-- RuntimeKnowledgeBase getters
-- Knowledge engine enrichment
-- KB loading and removal
-"""
+"""Tests for Workspace-scoped RAG manager and runtime knowledge bases."""
from __future__ import annotations
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, Mock
+
import pytest
-import uuid
-from unittest.mock import Mock, AsyncMock
-from importlib import import_module
+
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.entity.persistence.rag import KnowledgeBase
+from langbot.pkg.rag.knowledge.kbmgr import RAGManager, RuntimeKnowledgeBase
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
-def get_rag_module():
- """Lazy import to avoid circular import issues."""
- return import_module('langbot.pkg.rag.knowledge.kbmgr')
+CONTEXT_A = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=5,
+)
+CONTEXT_B = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-b',
+ placement_generation=5,
+)
-def create_mock_app():
- """Create mock Application for testing."""
- mock_app = Mock()
- mock_app.logger = Mock()
- 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
- mock_app.storage_mgr = Mock()
- mock_app.storage_mgr.storage_provider = AsyncMock()
- mock_app.task_mgr = AsyncMock()
- mock_app.task_mgr.create_user_task = Mock(return_value=Mock(id=1))
- return mock_app
+class _Result:
+ def __init__(self, rows=(), *, first=None):
+ self.rows = list(rows)
+ self._first = first
+
+ def all(self):
+ return self.rows
+
+ def first(self):
+ return self._first
-def create_mock_kb_entity():
- """Create mock KnowledgeBase entity."""
- mock_kb = Mock()
- mock_kb.uuid = str(uuid.uuid4())
- mock_kb.name = 'Test KB'
- mock_kb.description = 'Test description'
- mock_kb.knowledge_engine_plugin_id = 'author/engine'
- mock_kb.collection_id = mock_kb.uuid
- mock_kb.creation_settings = {}
- mock_kb.retrieval_settings = {}
- return mock_kb
+def _entity(*, kb_uuid='kb-a', workspace_uuid='workspace-a', plugin_id='author/engine'):
+ return KnowledgeBase(
+ uuid=kb_uuid,
+ workspace_uuid=workspace_uuid,
+ name='Test KB',
+ description='description',
+ knowledge_engine_plugin_id=plugin_id,
+ collection_id=kb_uuid,
+ creation_settings={},
+ retrieval_settings={},
+ )
+def _app():
+ return SimpleNamespace(
+ logger=Mock(),
+ persistence_mgr=SimpleNamespace(
+ execute_async=AsyncMock(return_value=_Result()),
+ serialize_model=Mock(
+ side_effect=lambda _model, row: {
+ 'uuid': row.uuid,
+ 'workspace_uuid': row.workspace_uuid,
+ 'name': row.name,
+ 'description': row.description,
+ 'knowledge_engine_plugin_id': row.knowledge_engine_plugin_id,
+ 'collection_id': row.collection_id,
+ 'creation_settings': row.creation_settings,
+ 'retrieval_settings': row.retrieval_settings,
+ }
+ ),
+ ),
+ plugin_connector=SimpleNamespace(
+ is_enable_plugin=True,
+ require_workspace_context=AsyncMock(side_effect=lambda context: context),
+ list_knowledge_engines=AsyncMock(
+ return_value=[
+ {
+ 'plugin_id': 'author/engine',
+ 'name': {'en_US': 'Engine'},
+ 'capabilities': ['doc_ingestion'],
+ }
+ ]
+ ),
+ rag_on_kb_create=AsyncMock(),
+ rag_on_kb_delete=AsyncMock(),
+ call_rag_ingest=AsyncMock(return_value={'status': 'success'}),
+ call_rag_retrieve=AsyncMock(return_value={'results': []}),
+ call_rag_delete_document=AsyncMock(return_value=True),
+ call_parser=AsyncMock(),
+ ),
+ workspace_service=SimpleNamespace(
+ get_execution_binding=AsyncMock(
+ side_effect=lambda workspace_uuid, **_kwargs: SimpleNamespace(
+ instance_uuid='instance-a',
+ workspace_uuid=workspace_uuid,
+ placement_generation=5,
+ )
+ )
+ ),
+ storage_mgr=SimpleNamespace(storage_provider=AsyncMock()),
+ task_mgr=SimpleNamespace(create_user_task=Mock(return_value=SimpleNamespace(id='task-a'))),
+ )
+
+
+@pytest.mark.asyncio
+async def test_create_binds_workspace_and_uses_tuple_runtime_key():
+ app = _app()
+ manager = RAGManager(app)
+
+ kb = await manager.create_knowledge_base(
+ CONTEXT_A,
+ name='Created',
+ knowledge_engine_plugin_id='author/engine',
+ creation_settings={'model': 'embedding-a'},
+ )
+
+ assert kb.workspace_uuid == 'workspace-a'
+ assert ('workspace-a', kb.uuid) in manager.knowledge_bases
+ app.plugin_connector.rag_on_kb_create.assert_awaited_once_with(
+ 'author/engine',
+ kb.uuid,
+ {'model': 'embedding-a'},
+ )
+
+
+@pytest.mark.asyncio
+async def test_create_rejects_unknown_engine_and_rolls_back_plugin_failure():
+ app = _app()
+ manager = RAGManager(app)
+ app.plugin_connector.list_knowledge_engines.return_value = []
+ with pytest.raises(ValueError, match='not found'):
+ await manager.create_knowledge_base(
+ CONTEXT_A,
+ name='Unknown',
+ knowledge_engine_plugin_id='missing/engine',
+ creation_settings={},
+ )
+
+ app.plugin_connector.list_knowledge_engines.return_value = [{'plugin_id': 'author/engine'}]
+ app.plugin_connector.rag_on_kb_create.side_effect = RuntimeError('plugin failed')
+ with pytest.raises(RuntimeError, match='plugin failed'):
+ await manager.create_knowledge_base(
+ CONTEXT_A,
+ name='Rollback',
+ knowledge_engine_plugin_id='author/engine',
+ creation_settings={},
+ )
+ assert manager.knowledge_bases == {}
+
+
+@pytest.mark.asyncio
+async def test_runtime_retrieve_carries_context_and_merges_settings_without_mutation():
+ app = _app()
+ entity = _entity()
+ entity.retrieval_settings = {'top_k': 10, 'model': 'default'}
+ app.plugin_connector.call_rag_retrieve.return_value = {
+ 'results': [
+ {
+ 'id': 'entry-a',
+ 'content': [{'type': 'text', 'text': 'hello'}],
+ 'metadata': {},
+ 'distance': 0.2,
+ }
+ ]
+ }
+ runtime = RuntimeKnowledgeBase(app, entity, CONTEXT_A)
+ overrides = {'top_k': 2, 'filters': {'file_id': 'file-a'}}
+
+ results = await runtime.retrieve(CONTEXT_A, 'query', settings=overrides)
+
+ assert results[0].id == 'entry-a'
+ assert overrides == {'top_k': 2, 'filters': {'file_id': 'file-a'}}
+ payload = app.plugin_connector.call_rag_retrieve.await_args.args[1]
+ assert payload['knowledge_base_id'] == 'kb-a'
+ assert payload['collection_id'] == 'kb-a'
+ assert payload['retrieval_settings']['top_k'] == 2
+ assert payload['filters'] == {'file_id': 'file-a'}
+
+
+@pytest.mark.asyncio
+async def test_runtime_rejects_cross_workspace_and_stale_contexts():
+ app = _app()
+ runtime = RuntimeKnowledgeBase(app, _entity(), CONTEXT_A)
+
+ with pytest.raises(WorkspaceNotFoundError):
+ await runtime.retrieve(CONTEXT_B, 'query')
+ stale = CONTEXT_A.__class__(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=4,
+ )
+ with pytest.raises(WorkspaceNotFoundError):
+ await runtime.retrieve(stale, 'query')
+ app.plugin_connector.call_rag_retrieve.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_runtime_retrieve_fails_before_bound_connector_on_generation_mismatch():
+ app = _app()
+ app.plugin_connector.require_workspace_context.side_effect = WorkspaceNotFoundError('Plugin resource not found')
+ runtime = RuntimeKnowledgeBase(app, _entity(), CONTEXT_A)
+
+ with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
+ await runtime.retrieve(CONTEXT_A, 'query')
+
+ app.plugin_connector.call_rag_retrieve.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ ('operation', 'connector_method'),
+ [
+ ('create', 'rag_on_kb_create'),
+ ('ingest', 'call_rag_ingest'),
+ ('delete', 'call_rag_delete_document'),
+ ],
+)
+async def test_runtime_plugin_mutations_fail_before_mismatched_connector(operation, connector_method):
+ app = _app()
+ app.plugin_connector.require_workspace_context.side_effect = WorkspaceNotFoundError('Plugin resource not found')
+ runtime = RuntimeKnowledgeBase(app, _entity(), CONTEXT_A)
+
+ with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
+ if operation == 'create':
+ await runtime._on_kb_create(CONTEXT_A)
+ elif operation == 'ingest':
+ await runtime._ingest_document(CONTEXT_A, {'filename': 'document.pdf'}, 'storage/path')
+ else:
+ await runtime._delete_document(CONTEXT_A, 'file-a')
+
+ getattr(app.plugin_connector, connector_method).assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_create_fails_before_engine_discovery_on_connector_workspace_mismatch():
+ app = _app()
+ app.plugin_connector.require_workspace_context.side_effect = WorkspaceNotFoundError('Plugin resource not found')
+ manager = RAGManager(app)
+
+ with pytest.raises(WorkspaceNotFoundError, match='Plugin resource not found'):
+ await manager.create_knowledge_base(
+ CONTEXT_A,
+ name='Other workspace',
+ knowledge_engine_plugin_id='author/engine',
+ creation_settings={},
+ )
+
+ app.plugin_connector.list_knowledge_engines.assert_not_awaited()
+ app.persistence_mgr.execute_async.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_ingestion_payload_uses_host_owned_kb_collection():
+ app = _app()
+ runtime = RuntimeKnowledgeBase(app, _entity(), CONTEXT_A)
+ metadata = {'filename': 'document.pdf'}
+
+ await runtime._ingest_document(CONTEXT_A, metadata, 'uploads/document.pdf')
+
+ payload = app.plugin_connector.call_rag_ingest.await_args.args[1]
+ assert payload['knowledge_base_id'] == 'kb-a'
+ assert payload['collection_id'] == 'kb-a'
+ assert payload['file_object']['metadata']['knowledge_base_id'] == 'kb-a'
+
+
+@pytest.mark.asyncio
+async def test_delete_file_checks_workspace_and_parent_before_plugin_call():
+ app = _app()
+ runtime = RuntimeKnowledgeBase(app, _entity(), CONTEXT_A)
+ app.persistence_mgr.execute_async.return_value = _Result(first=('file-a',))
+
+ await runtime.delete_file(CONTEXT_A, 'file-a')
+ app.plugin_connector.call_rag_delete_document.assert_awaited_once_with(
+ 'author/engine',
+ 'file-a',
+ 'kb-a',
+ )
+
+ missing_app = _app()
+ missing = RuntimeKnowledgeBase(missing_app, _entity(), CONTEXT_A)
+ with pytest.raises(WorkspaceNotFoundError):
+ await missing.delete_file(CONTEXT_A, 'file-other')
+ missing_app.plugin_connector.call_rag_delete_document.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_manager_get_remove_and_delete_are_workspace_scoped():
+ app = _app()
+ manager = RAGManager(app)
+ runtime = await manager.load_knowledge_base(CONTEXT_A, _entity())
+
+ assert await manager.get_knowledge_base_by_uuid(CONTEXT_A, 'kb-a') is runtime
+ assert await manager.get_knowledge_base_by_uuid(CONTEXT_B, 'kb-a') is None
+ await manager.remove_knowledge_base_from_runtime(CONTEXT_B, 'kb-a')
+ assert await manager.get_knowledge_base_by_uuid(CONTEXT_A, 'kb-a') is runtime
+ await manager.delete_knowledge_base(CONTEXT_A, 'kb-a')
+ app.plugin_connector.rag_on_kb_delete.assert_awaited_once()
+ assert await manager.get_knowledge_base_by_uuid(CONTEXT_A, 'kb-a') is None
+
+
+@pytest.mark.asyncio
+async def test_details_queries_require_workspace_and_enrich_engine():
+ app = _app()
+ row_a = _entity()
+ app.persistence_mgr.execute_async.return_value = _Result([row_a], first=row_a)
+ manager = RAGManager(app)
+
+ listed = await manager.get_all_knowledge_base_details(CONTEXT_A)
+ fetched = await manager.get_knowledge_base_details(CONTEXT_A, 'kb-a')
+ assert listed[0]['knowledge_engine']['plugin_id'] == 'author/engine'
+ assert fetched['knowledge_engine']['capabilities'] == ['doc_ingestion']
+
+
+@pytest.mark.asyncio
+async def test_load_dict_filters_computed_fields_and_requires_matching_workspace():
+ app = _app()
+ manager = RAGManager(app)
+ runtime = await manager.load_knowledge_base(
+ CONTEXT_A,
+ {
+ 'uuid': 'kb-a',
+ 'workspace_uuid': 'workspace-a',
+ 'name': 'KB',
+ 'description': '',
+ 'knowledge_engine_plugin_id': 'author/engine',
+ 'collection_id': 'kb-a',
+ 'creation_settings': {},
+ 'retrieval_settings': {},
+ 'knowledge_engine': {'computed': True},
+ },
+ )
+ assert runtime.get_uuid() == 'kb-a'
+
+ with pytest.raises(WorkspaceNotFoundError):
+ await manager.load_knowledge_base(CONTEXT_B, _entity())
+
+
+# Preserve the complete pre-tenancy RAG manager regression matrix. Every
+# runtime operation now carries the immutable ExecutionContext, and cache
+# assertions use the Workspace + resource tuple key introduced for isolation.
class TestRAGManagerCreateKnowledgeBase:
- """Tests for create_knowledge_base method."""
-
@pytest.mark.asyncio
async def test_creates_kb_with_valid_engine(self):
- """Test creates KB when engine plugin exists."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
-
- # Mock valid engine list
- mock_app.plugin_connector.list_knowledge_engines = AsyncMock(
- return_value=[{'plugin_id': 'author/engine', 'name': 'Engine'}]
- )
- mock_app.persistence_mgr.execute_async = AsyncMock()
- mock_app.plugin_connector.rag_on_kb_create = AsyncMock()
-
- manager = rag_module.RAGManager(mock_app)
+ app = _app()
+ manager = RAGManager(app)
kb = await manager.create_knowledge_base(
+ CONTEXT_A,
name='Test KB',
knowledge_engine_plugin_id='author/engine',
creation_settings={'model': 'test'},
)
assert kb.name == 'Test KB'
+ assert kb.workspace_uuid == CONTEXT_A.workspace_uuid
assert kb.knowledge_engine_plugin_id == 'author/engine'
@pytest.mark.asyncio
async def test_raises_when_engine_not_found(self):
- """Test raises ValueError when engine plugin not found."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
+ app.plugin_connector.list_knowledge_engines.return_value = []
- # Mock empty engine list
- mock_app.plugin_connector.list_knowledge_engines = AsyncMock(return_value=[])
-
- manager = rag_module.RAGManager(mock_app)
-
- with pytest.raises(ValueError) as exc_info:
- await manager.create_knowledge_base(
+ with pytest.raises(ValueError, match='not found'):
+ await RAGManager(app).create_knowledge_base(
+ CONTEXT_A,
name='Test KB',
knowledge_engine_plugin_id='unknown/engine',
creation_settings={},
)
- assert 'not found' in str(exc_info.value)
-
@pytest.mark.asyncio
async def test_rollback_on_plugin_create_failure(self):
- """Test that DB entry is rolled back when plugin create fails."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
+ app.plugin_connector.rag_on_kb_create.side_effect = RuntimeError('Plugin error')
+ manager = RAGManager(app)
- mock_app.plugin_connector.list_knowledge_engines = AsyncMock(return_value=[{'plugin_id': 'author/engine'}])
- mock_app.persistence_mgr.execute_async = AsyncMock()
- mock_app.plugin_connector.rag_on_kb_create = AsyncMock(side_effect=Exception('Plugin error'))
-
- manager = rag_module.RAGManager(mock_app)
-
- with pytest.raises(Exception):
+ with pytest.raises(RuntimeError, match='Plugin error'):
await manager.create_knowledge_base(
+ CONTEXT_A,
name='Test KB',
knowledge_engine_plugin_id='author/engine',
creation_settings={},
)
- # Should have called delete to rollback
- # Check that delete was called (for rollback)
- assert len(manager.knowledge_bases) == 0
+ assert manager.knowledge_bases == {}
+ assert app.persistence_mgr.execute_async.await_count == 2
@pytest.mark.asyncio
async def test_sets_default_retrieval_settings(self):
- """Test that empty retrieval_settings defaults to {}."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
- mock_app.plugin_connector.list_knowledge_engines = AsyncMock(return_value=[{'plugin_id': 'author/engine'}])
- mock_app.persistence_mgr.execute_async = AsyncMock()
- mock_app.plugin_connector.rag_on_kb_create = AsyncMock()
-
- manager = rag_module.RAGManager(mock_app)
-
- kb = await manager.create_knowledge_base(
+ kb = await RAGManager(app).create_knowledge_base(
+ CONTEXT_A,
name='Test KB',
knowledge_engine_plugin_id='author/engine',
creation_settings={},
@@ -142,419 +400,289 @@ class TestRAGManagerCreateKnowledgeBase:
@pytest.mark.asyncio
async def test_skips_validation_when_plugin_disabled(self):
- """Test that engine validation is skipped when plugin disabled."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_app.plugin_connector.is_enable_plugin = False
- mock_app.persistence_mgr.execute_async = AsyncMock()
- mock_app.plugin_connector.rag_on_kb_create = AsyncMock()
+ app = _app()
+ app.plugin_connector.is_enable_plugin = False
- manager = rag_module.RAGManager(mock_app)
-
- # Should not raise even though engine list would be empty
- kb = await manager.create_knowledge_base(
+ kb = await RAGManager(app).create_knowledge_base(
+ CONTEXT_A,
name='Test KB',
knowledge_engine_plugin_id='any/engine',
creation_settings={},
)
assert kb.knowledge_engine_plugin_id == 'any/engine'
+ app.plugin_connector.list_knowledge_engines.assert_not_awaited()
class TestRuntimeKnowledgeBaseOnKBCreate:
- """Tests for _on_kb_create method."""
-
@pytest.mark.asyncio
async def test_calls_plugin_on_create(self):
- """Test that plugin is notified on KB create."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
- mock_kb.creation_settings = {'model': 'test'}
+ app = _app()
+ entity = _entity()
+ entity.creation_settings = {'model': 'test'}
- mock_app.plugin_connector.rag_on_kb_create = AsyncMock()
+ await RuntimeKnowledgeBase(app, entity, CONTEXT_A)._on_kb_create(CONTEXT_A)
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
- await runtime_kb._on_kb_create()
-
- mock_app.plugin_connector.rag_on_kb_create.assert_called_once_with(
- 'author/engine', mock_kb.uuid, {'model': 'test'}
+ app.plugin_connector.rag_on_kb_create.assert_awaited_once_with(
+ 'author/engine',
+ entity.uuid,
+ {'model': 'test'},
)
@pytest.mark.asyncio
async def test_skips_when_no_plugin_id(self):
- """Test that create notification is skipped when no plugin."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
- mock_kb.knowledge_engine_plugin_id = None
+ app = _app()
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
- await runtime_kb._on_kb_create()
+ await RuntimeKnowledgeBase(
+ app,
+ _entity(plugin_id=None),
+ CONTEXT_A,
+ )._on_kb_create(CONTEXT_A)
- mock_app.plugin_connector.rag_on_kb_create.assert_not_called()
+ app.plugin_connector.rag_on_kb_create.assert_not_awaited()
@pytest.mark.asyncio
async def test_raises_on_plugin_error(self):
- """Test that exception is raised when plugin fails."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
+ app = _app()
+ app.plugin_connector.rag_on_kb_create.side_effect = RuntimeError('Plugin failed')
- mock_app.plugin_connector.rag_on_kb_create = AsyncMock(side_effect=Exception('Plugin failed'))
-
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
-
- with pytest.raises(Exception):
- await runtime_kb._on_kb_create()
+ with pytest.raises(RuntimeError, match='Plugin failed'):
+ await RuntimeKnowledgeBase(app, _entity(), CONTEXT_A)._on_kb_create(CONTEXT_A)
class TestRuntimeKnowledgeBaseDeleteFile:
- """Tests for delete_file method."""
-
@pytest.mark.asyncio
async def test_delete_file_calls_plugin_and_db(self):
- """Test that delete_file calls plugin and removes DB record."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
+ app = _app()
+ app.persistence_mgr.execute_async.return_value = _Result(first=('file-uuid',))
- mock_app.plugin_connector.call_rag_delete_document = AsyncMock(return_value=True)
+ await RuntimeKnowledgeBase(app, _entity(), CONTEXT_A).delete_file(
+ CONTEXT_A,
+ 'file-uuid',
+ )
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
- await runtime_kb.delete_file('file-uuid')
-
- mock_app.plugin_connector.call_rag_delete_document.assert_called_once()
- mock_app.persistence_mgr.execute_async.assert_called()
+ app.plugin_connector.call_rag_delete_document.assert_awaited_once_with(
+ 'author/engine',
+ 'file-uuid',
+ 'kb-a',
+ )
+ assert app.persistence_mgr.execute_async.await_count == 2
class TestRuntimeKnowledgeBaseIngestDocument:
- """Tests for _ingest_document method."""
-
@pytest.mark.asyncio
async def test_ingest_calls_plugin(self):
- """Test that ingest calls plugin connector."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
+ app = _app()
- mock_app.plugin_connector.call_rag_ingest = AsyncMock(return_value={'status': 'success'})
-
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
-
- result = await runtime_kb._ingest_document(
+ result = await RuntimeKnowledgeBase(app, _entity(), CONTEXT_A)._ingest_document(
+ CONTEXT_A,
{'filename': 'test.pdf'},
'storage/path',
)
- assert result['status'] == 'success'
- mock_app.plugin_connector.call_rag_ingest.assert_called_once()
+ assert result == {'status': 'success'}
+ app.plugin_connector.call_rag_ingest.assert_awaited_once()
@pytest.mark.asyncio
async def test_ingest_raises_when_no_plugin_id(self):
- """Test that ValueError is raised when no plugin ID."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
- mock_kb.knowledge_engine_plugin_id = None
+ app = _app()
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
-
- with pytest.raises(ValueError) as exc_info:
- await runtime_kb._ingest_document({'filename': 'test.pdf'}, 'path')
-
- assert 'Plugin ID required' in str(exc_info.value)
+ with pytest.raises(ValueError, match='Plugin ID required'):
+ await RuntimeKnowledgeBase(
+ app,
+ _entity(plugin_id=None),
+ CONTEXT_A,
+ )._ingest_document(CONTEXT_A, {'filename': 'test.pdf'}, 'path')
class TestRAGManagerLoadKnowledgeBasesFromDB:
- """Tests for load_knowledge_bases_from_db method."""
-
@pytest.mark.asyncio
async def test_loads_all_kbs_from_db(self):
- """Test that all KBs are loaded from database."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
+ kb1 = _entity(kb_uuid='kb-1')
+ kb2 = _entity(kb_uuid='kb-2')
+ app.persistence_mgr.execute_async.return_value = _Result([kb1, kb2])
+ manager = RAGManager(app)
- mock_kb1 = create_mock_kb_entity()
- mock_kb2 = create_mock_kb_entity()
- mock_app.persistence_mgr.execute_async = AsyncMock(
- return_value=Mock(all=Mock(return_value=[mock_kb1, mock_kb2]))
- )
-
- manager = rag_module.RAGManager(mock_app)
await manager.load_knowledge_bases_from_db()
- assert len(manager.knowledge_bases) == 2
+ assert set(manager.knowledge_bases) == {
+ ('workspace-a', 'kb-1'),
+ ('workspace-a', 'kb-2'),
+ }
@pytest.mark.asyncio
async def test_handles_load_error_gracefully(self):
- """Test that load errors are logged but not raised."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
+ app.persistence_mgr.execute_async.return_value = _Result([_entity()])
+ app.workspace_service.get_execution_binding.side_effect = RuntimeError('binding unavailable')
+ manager = RAGManager(app)
- # KB that will cause initialize to fail
- mock_kb = create_mock_kb_entity()
-
- mock_app.persistence_mgr.execute_async = AsyncMock(return_value=Mock(all=Mock(return_value=[mock_kb])))
-
- # Make initialize fail by having plugin_connector throw error
- mock_app.plugin_connector.rag_on_kb_create = AsyncMock(side_effect=Exception('Init failed'))
-
- manager = rag_module.RAGManager(mock_app)
- # Should not raise - errors are caught
await manager.load_knowledge_bases_from_db()
- # KB should still be loaded (initialize just passes)
- # The error would come from runtime_kb.initialize which we can't easily mock
- # So we just verify it doesn't crash
+ assert manager.knowledge_bases == {}
+ app.logger.error.assert_called_once()
class TestRuntimeKnowledgeBaseGetters:
- """Tests for RuntimeKnowledgeBase getter methods."""
-
def test_get_uuid_returns_entity_uuid(self):
- """Test get_uuid returns KB entity UUID."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
+ entity = _entity()
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
-
- assert runtime_kb.get_uuid() == mock_kb.uuid
+ assert RuntimeKnowledgeBase(_app(), entity, CONTEXT_A).get_uuid() == entity.uuid
def test_get_name_returns_entity_name(self):
- """Test get_name returns KB entity name."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
+ entity = _entity()
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
-
- assert runtime_kb.get_name() == mock_kb.name
+ assert RuntimeKnowledgeBase(_app(), entity, CONTEXT_A).get_name() == entity.name
def test_get_knowledge_engine_plugin_id_returns_plugin_id(self):
- """Test get_knowledge_engine_plugin_id returns plugin ID."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
+ runtime = RuntimeKnowledgeBase(_app(), _entity(), CONTEXT_A)
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
-
- assert runtime_kb.get_knowledge_engine_plugin_id() == 'author/engine'
+ assert runtime.get_knowledge_engine_plugin_id() == 'author/engine'
def test_get_knowledge_engine_plugin_id_returns_empty_when_none(self):
- """Test returns empty string when plugin_id is None."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
- mock_kb.knowledge_engine_plugin_id = None
+ runtime = RuntimeKnowledgeBase(_app(), _entity(plugin_id=None), CONTEXT_A)
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
-
- assert runtime_kb.get_knowledge_engine_plugin_id() == ''
+ assert runtime.get_knowledge_engine_plugin_id() == ''
class TestRuntimeKnowledgeBaseRetrieve:
- """Tests for RuntimeKnowledgeBase retrieve method."""
-
@pytest.mark.asyncio
async def test_retrieve_merges_settings(self):
- """Test that retrieve merges stored and request settings."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
- mock_kb.retrieval_settings = {'top_k': 10, 'model': 'default'}
+ app = _app()
+ entity = _entity()
+ entity.retrieval_settings = {'top_k': 10, 'model': 'default'}
+ app.plugin_connector.call_rag_retrieve.return_value = {
+ 'results': [
+ {
+ 'id': 'doc1',
+ 'content': [{'type': 'text', 'text': 'test content'}],
+ 'metadata': {},
+ 'distance': 0.1,
+ }
+ ]
+ }
- # Mock plugin connector response with valid RetrievalResultEntry fields
- # content must be list of ContentElement dicts
- mock_app.plugin_connector.call_rag_retrieve = AsyncMock(
- return_value={
- 'results': [
- {
- 'id': 'doc1',
- 'content': [{'type': 'text', 'text': 'test content'}],
- 'metadata': {},
- 'distance': 0.1,
- }
- ]
- }
+ results = await RuntimeKnowledgeBase(app, entity, CONTEXT_A).retrieve(
+ CONTEXT_A,
+ 'query text',
+ settings={'top_k': 20},
)
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
-
- # Override top_k in request
- results = await runtime_kb.retrieve('query text', settings={'top_k': 20})
-
assert len(results) == 1
- # Check that merged settings were passed (top_k overridden)
- call_args = mock_app.plugin_connector.call_rag_retrieve.call_args
- assert call_args[0][1]['retrieval_settings']['top_k'] == 20
+ payload = app.plugin_connector.call_rag_retrieve.await_args.args[1]
+ assert payload['retrieval_settings'] == {'top_k': 20, 'model': 'default'}
@pytest.mark.asyncio
async def test_retrieve_adds_default_top_k(self):
- """Test that default top_k=5 is added when not specified."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
- mock_kb.retrieval_settings = {}
+ app = _app()
- mock_app.plugin_connector.call_rag_retrieve = AsyncMock(return_value={'results': []})
+ await RuntimeKnowledgeBase(app, _entity(), CONTEXT_A).retrieve(
+ CONTEXT_A,
+ 'query text',
+ )
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
-
- await runtime_kb.retrieve('query text')
-
- call_args = mock_app.plugin_connector.call_rag_retrieve.call_args
- assert call_args[0][1]['retrieval_settings']['top_k'] == 5
+ payload = app.plugin_connector.call_rag_retrieve.await_args.args[1]
+ assert payload['retrieval_settings']['top_k'] == 5
@pytest.mark.asyncio
async def test_retrieve_converts_dict_to_entry(self):
- """Test that dict results are converted to RetrievalResultEntry."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
+ app = _app()
+ app.plugin_connector.call_rag_retrieve.return_value = {
+ 'results': [
+ {
+ 'id': 'doc1',
+ 'content': [{'type': 'text', 'text': 'test content'}],
+ 'metadata': {'source': 'file.pdf'},
+ 'distance': 0.15,
+ }
+ ]
+ }
- # Mock response with valid RetrievalResultEntry fields
- # content must be list of ContentElement dicts
- mock_app.plugin_connector.call_rag_retrieve = AsyncMock(
- return_value={
- 'results': [
- {
- 'id': 'doc1',
- 'content': [{'type': 'text', 'text': 'test content'}],
- 'metadata': {'source': 'file.pdf'},
- 'distance': 0.15,
- }
- ]
- }
+ results = await RuntimeKnowledgeBase(app, _entity(), CONTEXT_A).retrieve(
+ CONTEXT_A,
+ 'query',
)
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
-
- results = await runtime_kb.retrieve('query')
-
assert len(results) == 1
- # Result should be RetrievalResultEntry
- assert hasattr(results[0], 'content')
assert results[0].id == 'doc1'
+ assert hasattr(results[0], 'content')
class TestRuntimeKnowledgeBaseDispose:
- """Tests for RuntimeKnowledgeBase dispose method."""
-
@pytest.mark.asyncio
async def test_dispose_calls_on_kb_delete(self):
- """Test that dispose calls _on_kb_delete."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
+ app = _app()
- mock_app.plugin_connector.rag_on_kb_delete = AsyncMock()
+ await RuntimeKnowledgeBase(app, _entity(), CONTEXT_A).dispose(CONTEXT_A)
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
-
- await runtime_kb.dispose()
-
- mock_app.plugin_connector.rag_on_kb_delete.assert_called_once()
+ app.plugin_connector.rag_on_kb_delete.assert_awaited_once_with(
+ 'author/engine',
+ 'kb-a',
+ )
@pytest.mark.asyncio
async def test_dispose_skips_when_no_plugin_id(self):
- """Test that dispose skips when no plugin ID."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_kb = create_mock_kb_entity()
- mock_kb.knowledge_engine_plugin_id = None
+ app = _app()
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
+ await RuntimeKnowledgeBase(
+ app,
+ _entity(plugin_id=None),
+ CONTEXT_A,
+ ).dispose(CONTEXT_A)
- await runtime_kb.dispose()
-
- # Should not call plugin connector
- mock_app.plugin_connector.rag_on_kb_delete.assert_not_called()
+ app.plugin_connector.rag_on_kb_delete.assert_not_awaited()
class TestRAGManagerInit:
- """Tests for RAGManager initialization."""
-
def test_init_stores_app_reference(self):
- """Test that __init__ stores Application reference."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
- manager = rag_module.RAGManager(mock_app)
-
- assert manager.ap is mock_app
+ assert RAGManager(app).ap is app
def test_init_creates_empty_knowledge_bases_dict(self):
- """Test that knowledge_bases starts as empty dict."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
-
- manager = rag_module.RAGManager(mock_app)
-
- assert manager.knowledge_bases == {}
+ assert RAGManager(_app()).knowledge_bases == {}
class TestRAGManagerGetKnowledgeBase:
- """Tests for RAGManager get methods."""
-
@pytest.mark.asyncio
async def test_get_knowledge_base_by_uuid_returns_runtime_kb(self):
- """Test get_knowledge_base_by_uuid returns loaded KB."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
+ manager = RAGManager(app)
+ runtime = RuntimeKnowledgeBase(app, _entity(), CONTEXT_A)
+ manager.knowledge_bases[('workspace-a', 'kb-a')] = runtime
- manager = rag_module.RAGManager(mock_app)
- mock_kb = create_mock_kb_entity()
+ result = await manager.get_knowledge_base_by_uuid(CONTEXT_A, 'kb-a')
- # Manually add to knowledge_bases
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
- manager.knowledge_bases[mock_kb.uuid] = runtime_kb
-
- result = await manager.get_knowledge_base_by_uuid(mock_kb.uuid)
-
- assert result is runtime_kb
+ assert result is runtime
@pytest.mark.asyncio
async def test_get_knowledge_base_by_uuid_returns_none_when_not_found(self):
- """Test returns None when KB not in runtime."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
-
- manager = rag_module.RAGManager(mock_app)
-
- result = await manager.get_knowledge_base_by_uuid('nonexistent-uuid')
-
- assert result is None
+ assert (
+ await RAGManager(_app()).get_knowledge_base_by_uuid(
+ CONTEXT_A,
+ 'nonexistent-uuid',
+ )
+ is None
+ )
@pytest.mark.asyncio
async def test_remove_knowledge_base_from_runtime(self):
- """Test remove_knowledge_base_from_runtime removes KB."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
+ manager = RAGManager(app)
+ manager.knowledge_bases[('workspace-a', 'kb-a')] = RuntimeKnowledgeBase(
+ app,
+ _entity(),
+ CONTEXT_A,
+ )
- manager = rag_module.RAGManager(mock_app)
- mock_kb = create_mock_kb_entity()
+ await manager.remove_knowledge_base_from_runtime(CONTEXT_A, 'kb-a')
- # Add to knowledge_bases
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
- manager.knowledge_bases[mock_kb.uuid] = runtime_kb
-
- await manager.remove_knowledge_base_from_runtime(mock_kb.uuid)
-
- assert mock_kb.uuid not in manager.knowledge_bases
+ assert ('workspace-a', 'kb-a') not in manager.knowledge_bases
class TestRAGManagerEnrichKB:
- """Tests for _enrich_kb_dict method."""
-
def test_enrich_adds_engine_info_from_map(self):
- """Test that engine info is added from engine_map."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
-
- manager = rag_module.RAGManager(mock_app)
-
kb_dict = {'knowledge_engine_plugin_id': 'author/engine'}
engine_map = {
'author/engine': {
@@ -564,208 +692,133 @@ class TestRAGManagerEnrichKB:
}
}
- manager._enrich_kb_dict(kb_dict, engine_map)
+ RAGManager(_app())._enrich_kb_dict(kb_dict, engine_map)
- assert 'knowledge_engine' in kb_dict
assert kb_dict['knowledge_engine']['plugin_id'] == 'author/engine'
assert kb_dict['knowledge_engine']['capabilities'] == ['doc_ingestion', 'search']
def test_enrich_uses_fallback_when_engine_not_in_map(self):
- """Test that fallback info is used when engine not found."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
-
- manager = rag_module.RAGManager(mock_app)
-
kb_dict = {'knowledge_engine_plugin_id': 'unknown/engine'}
- engine_map = {}
- manager._enrich_kb_dict(kb_dict, engine_map)
+ RAGManager(_app())._enrich_kb_dict(kb_dict, {})
- assert 'knowledge_engine' in kb_dict
assert kb_dict['knowledge_engine']['plugin_id'] == 'unknown/engine'
assert kb_dict['knowledge_engine']['capabilities'] == []
def test_enrich_uses_fallback_when_no_plugin_id(self):
- """Test that fallback is used when no plugin ID."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
-
- manager = rag_module.RAGManager(mock_app)
-
kb_dict = {}
- engine_map = {}
- manager._enrich_kb_dict(kb_dict, engine_map)
+ RAGManager(_app())._enrich_kb_dict(kb_dict, {})
- assert 'knowledge_engine' in kb_dict
- # Should have Internal (Legacy) name
+ assert kb_dict['knowledge_engine']['plugin_id'] is None
assert 'en_US' in kb_dict['knowledge_engine']['name']
def test_enrich_converts_string_name_to_i18n(self):
- """Test that engine name is converted to i18n dict."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
-
- manager = rag_module.RAGManager(mock_app)
-
kb_dict = {'knowledge_engine_plugin_id': 'author/engine'}
engine_map = {
'author/engine': {
'plugin_id': 'author/engine',
- 'name': 'Simple Name', # String, not dict
+ 'name': 'Simple Name',
'capabilities': [],
}
}
- manager._enrich_kb_dict(kb_dict, engine_map)
+ RAGManager(_app())._enrich_kb_dict(kb_dict, engine_map)
- # Name should be converted to i18n dict
- engine_name = kb_dict['knowledge_engine']['name']
- assert isinstance(engine_name, dict)
- assert engine_name['en_US'] == 'Simple Name'
+ assert kb_dict['knowledge_engine']['name'] == {
+ 'en_US': 'Simple Name',
+ 'zh_Hans': 'Simple Name',
+ }
class TestRAGManagerDeleteKnowledgeBase:
- """Tests for delete_knowledge_base method."""
-
@pytest.mark.asyncio
async def test_delete_removes_from_runtime_and_disposes(self):
- """Test that delete removes KB and calls dispose."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
+ manager = RAGManager(app)
+ manager.knowledge_bases[('workspace-a', 'kb-a')] = RuntimeKnowledgeBase(
+ app,
+ _entity(),
+ CONTEXT_A,
+ )
- manager = rag_module.RAGManager(mock_app)
- mock_kb = create_mock_kb_entity()
+ await manager.delete_knowledge_base(CONTEXT_A, 'kb-a')
- # Add to knowledge_bases
- runtime_kb = rag_module.RuntimeKnowledgeBase(mock_app, mock_kb)
- manager.knowledge_bases[mock_kb.uuid] = runtime_kb
-
- await manager.delete_knowledge_base(mock_kb.uuid)
-
- assert mock_kb.uuid not in manager.knowledge_bases
+ assert ('workspace-a', 'kb-a') not in manager.knowledge_bases
+ app.plugin_connector.rag_on_kb_delete.assert_awaited_once()
@pytest.mark.asyncio
async def test_delete_logs_warning_when_not_in_runtime(self):
- """Test that warning is logged when KB not in runtime."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
- manager = rag_module.RAGManager(mock_app)
+ await RAGManager(app).delete_knowledge_base(CONTEXT_A, 'nonexistent-uuid')
- await manager.delete_knowledge_base('nonexistent-uuid')
-
- mock_app.logger.warning.assert_called_once()
+ app.logger.warning.assert_called_once()
class TestRAGManagerGetAllDetails:
- """Tests for get_all_knowledge_base_details method."""
-
@pytest.mark.asyncio
async def test_returns_empty_list_when_no_kbs(self):
- """Test returns empty list when no knowledge bases."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_app.persistence_mgr.execute_async = AsyncMock(return_value=Mock(all=Mock(return_value=[])))
-
- manager = rag_module.RAGManager(mock_app)
- result = await manager.get_all_knowledge_base_details()
-
- assert result == []
+ assert await RAGManager(_app()).get_all_knowledge_base_details(CONTEXT_A) == []
@pytest.mark.asyncio
async def test_enriches_each_kb_with_engine_info(self):
- """Test that each KB is enriched with engine info."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
+ app.persistence_mgr.execute_async.return_value = _Result([_entity()])
- # Mock DB result
- mock_kb_row = Mock()
- mock_app.persistence_mgr.execute_async = AsyncMock(return_value=Mock(all=Mock(return_value=[mock_kb_row])))
- mock_app.persistence_mgr.serialize_model = Mock(
- return_value={'uuid': 'kb1', 'knowledge_engine_plugin_id': 'author/engine'}
- )
- mock_app.plugin_connector.list_knowledge_engines = AsyncMock(
- return_value=[{'plugin_id': 'author/engine', 'name': 'Engine', 'capabilities': ['search']}]
- )
-
- manager = rag_module.RAGManager(mock_app)
- result = await manager.get_all_knowledge_base_details()
+ result = await RAGManager(app).get_all_knowledge_base_details(CONTEXT_A)
assert len(result) == 1
- assert 'knowledge_engine' in result[0]
+ assert result[0]['knowledge_engine']['plugin_id'] == 'author/engine'
class TestRAGManagerGetDetails:
- """Tests for get_knowledge_base_details method."""
-
@pytest.mark.asyncio
async def test_returns_none_when_kb_not_found(self):
- """Test returns None when KB doesn't exist."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
- mock_app.persistence_mgr.execute_async = AsyncMock(return_value=Mock(first=Mock(return_value=None)))
-
- manager = rag_module.RAGManager(mock_app)
- result = await manager.get_knowledge_base_details('nonexistent')
-
- assert result is None
+ assert (
+ await RAGManager(_app()).get_knowledge_base_details(
+ CONTEXT_A,
+ 'nonexistent',
+ )
+ is None
+ )
@pytest.mark.asyncio
async def test_returns_enriched_kb_dict(self):
- """Test returns enriched KB dict when found."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ app = _app()
+ app.persistence_mgr.execute_async.return_value = _Result(first=_entity())
- mock_kb_row = Mock()
- mock_app.persistence_mgr.execute_async = AsyncMock(return_value=Mock(first=Mock(return_value=mock_kb_row)))
- mock_app.persistence_mgr.serialize_model = Mock(
- return_value={'uuid': 'kb1', 'knowledge_engine_plugin_id': 'author/engine'}
- )
- mock_app.plugin_connector.list_knowledge_engines = AsyncMock(
- return_value=[{'plugin_id': 'author/engine', 'name': 'Engine', 'capabilities': []}]
- )
-
- manager = rag_module.RAGManager(mock_app)
- result = await manager.get_knowledge_base_details('kb1')
+ result = await RAGManager(app).get_knowledge_base_details(CONTEXT_A, 'kb-a')
assert result is not None
- assert 'knowledge_engine' in result
+ assert result['knowledge_engine']['plugin_id'] == 'author/engine'
class TestRAGManagerLoadKnowledgeBase:
- """Tests for load_knowledge_base method."""
-
@pytest.mark.asyncio
async def test_loads_kb_entity_into_runtime(self):
- """Test that KB entity is loaded into runtime."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
+ manager = RAGManager(_app())
- manager = rag_module.RAGManager(mock_app)
- mock_kb = create_mock_kb_entity()
+ result = await manager.load_knowledge_base(CONTEXT_A, _entity())
- result = await manager.load_knowledge_base(mock_kb)
-
- assert mock_kb.uuid in manager.knowledge_bases
- assert result.get_uuid() == mock_kb.uuid
+ assert ('workspace-a', 'kb-a') in manager.knowledge_bases
+ assert result.get_uuid() == 'kb-a'
@pytest.mark.asyncio
async def test_load_handles_dict_entity(self):
- """Test that dict entity is converted to KB object."""
- rag_module = get_rag_module()
- mock_app = create_mock_app()
-
- manager = rag_module.RAGManager(mock_app)
-
+ manager = RAGManager(_app())
kb_dict = {
'uuid': 'kb-uuid',
+ 'workspace_uuid': 'workspace-a',
'name': 'Test',
+ 'description': '',
'knowledge_engine_plugin_id': 'author/engine',
- 'knowledge_engine': {'name': 'should_be_filtered'}, # non-db field
+ 'collection_id': 'kb-uuid',
+ 'creation_settings': {},
+ 'retrieval_settings': {},
+ 'knowledge_engine': {'name': 'should_be_filtered'},
}
- await manager.load_knowledge_base(kb_dict)
+ await manager.load_knowledge_base(CONTEXT_A, kb_dict)
- assert 'kb-uuid' in manager.knowledge_bases
+ assert ('workspace-a', 'kb-uuid') in manager.knowledge_bases
diff --git a/tests/unit_tests/rag/test_runtime_service.py b/tests/unit_tests/rag/test_runtime_service.py
index 650b3bf2f..11bb3807a 100644
--- a/tests/unit_tests/rag/test_runtime_service.py
+++ b/tests/unit_tests/rag/test_runtime_service.py
@@ -1,474 +1,467 @@
-"""Tests for RAGRuntimeService.
-
-Tests the service that handles RAG-related requests from plugins,
-using mocked vector_db_mgr and storage_mgr.
-"""
+"""Tenant-aware tests for the plugin-facing RAG runtime service."""
from __future__ import annotations
-from unittest.mock import AsyncMock, MagicMock
+from types import SimpleNamespace
+from unittest.mock import AsyncMock
+
import pytest
-from tests.utils.import_isolation import isolated_sys_modules
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.rag.service.runtime import RAGRuntimeService
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
-class TestRAGRuntimeServiceVectorUpsert:
- """Tests for vector_upsert method."""
+WORKSPACE_UUID = '00000000-0000-0000-0000-00000000000a'
+CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid=WORKSPACE_UUID,
+ placement_generation=4,
+)
- def _create_mock_app(self):
- """Create mock app with vector_db_mgr and storage_mgr."""
- mock_app = MagicMock()
- mock_app.vector_db_mgr = MagicMock()
- mock_app.vector_db_mgr.upsert = AsyncMock()
- mock_app.storage_mgr = MagicMock()
- mock_app.storage_mgr.storage_provider = MagicMock()
- mock_app.storage_mgr.storage_provider.load = AsyncMock(return_value=b'content')
- return mock_app
- def _make_rag_import_mocks(self):
- """Create mocks needed for importing RAG service."""
- return {
- 'langbot.pkg.core.app': MagicMock(),
- 'langbot_plugin.api.entities.builtin.rag': MagicMock(),
- }
+class _ScalarResult:
+ def __init__(self, value):
+ self.value = value
+ def scalar_one_or_none(self):
+ return self.value
+
+ def first(self):
+ return None if self.value is None else (self.value,)
+
+
+def _app(*, kb_uuid='kb-a', file_exists=True):
+ persistence_results = [_ScalarResult(kb_uuid)]
+ if file_exists is not None:
+ persistence_results.append(_ScalarResult('file-a' if file_exists else None))
+ return SimpleNamespace(
+ workspace_service=SimpleNamespace(
+ get_execution_binding=AsyncMock(return_value=SimpleNamespace(instance_uuid='instance-a'))
+ ),
+ persistence_mgr=SimpleNamespace(execute_async=AsyncMock(side_effect=persistence_results)),
+ vector_db_mgr=SimpleNamespace(
+ upsert=AsyncMock(),
+ search=AsyncMock(return_value=[{'id': 'chunk-a'}]),
+ delete_by_file_id=AsyncMock(),
+ delete_by_filter=AsyncMock(return_value=3),
+ list_by_filter=AsyncMock(return_value=([{'id': 'chunk-a'}], 1)),
+ ),
+ storage_mgr=SimpleNamespace(
+ load_scoped_object_key=AsyncMock(return_value=b'content'),
+ storage_provider=SimpleNamespace(load=AsyncMock(return_value=b'content')),
+ ),
+ )
+
+
+@pytest.mark.asyncio
+async def test_vector_upsert_resolves_canonical_kb_and_forwards_trusted_context():
+ app = _app()
+ service = RAGRuntimeService(app)
+
+ await service.vector_upsert(
+ CONTEXT,
+ 'logical-collection',
+ [[0.1, 0.2]],
+ ['chunk-a'],
+ metadata=[{'file_id': 'file-a'}],
+ documents=['hello'],
+ )
+
+ app.vector_db_mgr.upsert.assert_awaited_once_with(
+ execution_context=CONTEXT,
+ knowledge_base_uuid='kb-a',
+ vectors=[[0.1, 0.2]],
+ ids=['chunk-a'],
+ metadata=[{'file_id': 'file-a'}],
+ documents=['hello'],
+ )
+
+
+@pytest.mark.asyncio
+async def test_vector_upsert_rejects_mismatched_lengths():
+ service = RAGRuntimeService(_app())
+ with pytest.raises(ValueError, match='vectors and ids'):
+ await service.vector_upsert(CONTEXT, 'kb-a', [[0.1]], ['a', 'b'])
+
+
+@pytest.mark.asyncio
+async def test_unknown_or_cross_workspace_collection_is_not_forwarded():
+ app = _app(kb_uuid=None)
+ service = RAGRuntimeService(app)
+
+ with pytest.raises(WorkspaceNotFoundError):
+ await service.vector_search(CONTEXT, 'kb-from-other-workspace', [0.1], 5)
+ app.vector_db_mgr.search.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_vector_search_forwards_all_search_options():
+ app = _app()
+ service = RAGRuntimeService(app)
+
+ result = await service.vector_search(
+ CONTEXT,
+ 'kb-a',
+ [0.1, 0.2],
+ 7,
+ filters={'file_id': 'file-a'},
+ search_type='hybrid',
+ query_text='hello',
+ vector_weight=0.7,
+ )
+
+ assert result == [{'id': 'chunk-a'}]
+ app.vector_db_mgr.search.assert_awaited_once_with(
+ execution_context=CONTEXT,
+ knowledge_base_uuid='kb-a',
+ query_vector=[0.1, 0.2],
+ limit=7,
+ filter={'file_id': 'file-a'},
+ search_type='hybrid',
+ query_text='hello',
+ vector_weight=0.7,
+ )
+
+
+@pytest.mark.asyncio
+async def test_vector_delete_and_list_stay_on_canonical_kb():
+ delete_app = _app()
+ delete_service = RAGRuntimeService(delete_app)
+ assert await delete_service.vector_delete(CONTEXT, 'logical', file_ids=['file-a']) == 1
+ delete_app.vector_db_mgr.delete_by_file_id.assert_awaited_once_with(
+ execution_context=CONTEXT,
+ knowledge_base_uuid='kb-a',
+ file_ids=['file-a'],
+ )
+
+ filter_app = _app()
+ filter_service = RAGRuntimeService(filter_app)
+ assert await filter_service.vector_delete(CONTEXT, 'logical', filters={'page': 1}) == 3
+ filter_app.vector_db_mgr.delete_by_filter.assert_awaited_once_with(
+ execution_context=CONTEXT,
+ knowledge_base_uuid='kb-a',
+ filter={'page': 1},
+ )
+
+ list_app = _app()
+ list_service = RAGRuntimeService(list_app)
+ assert await list_service.vector_list(CONTEXT, 'logical', {'page': 1}, 10, 2) == (
+ [{'id': 'chunk-a'}],
+ 1,
+ )
+ list_app.vector_db_mgr.list_by_filter.assert_awaited_once_with(
+ execution_context=CONTEXT,
+ knowledge_base_uuid='kb-a',
+ filter={'page': 1},
+ limit=10,
+ offset=2,
+ )
+
+
+@pytest.mark.asyncio
+async def test_file_stream_requires_workspace_owned_file():
+ app = _app(file_exists=True)
+ app.persistence_mgr.execute_async = AsyncMock(return_value=_ScalarResult('file-a'))
+ service = RAGRuntimeService(app)
+ assert await service.get_file_stream(CONTEXT, 'nested/file.pdf') == b'content'
+ app.storage_mgr.load_scoped_object_key.assert_awaited_once_with(
+ CONTEXT,
+ 'nested/file.pdf',
+ expected_owner_type='upload_document',
+ )
+
+ missing_app = _app(file_exists=False)
+ missing_app.persistence_mgr.execute_async = AsyncMock(return_value=_ScalarResult(None))
+ missing_service = RAGRuntimeService(missing_app)
+ with pytest.raises(WorkspaceNotFoundError):
+ await missing_service.get_file_stream(CONTEXT, 'other.pdf')
+ missing_app.storage_mgr.load_scoped_object_key.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+@pytest.mark.parametrize(
+ 'unsafe_path',
+ [
+ '',
+ '../secret.txt',
+ '/absolute/path.txt',
+ '..\\secret.txt',
+ 'nested\\..\\secret.txt',
+ '%2e%2e/secret.txt',
+ 'nested/%2e%2e/secret.txt',
+ 'C:\\secret.txt',
+ 'safe/\x00file.txt',
+ ],
+)
+async def test_file_stream_rejects_unsafe_paths_before_storage(unsafe_path):
+ app = _app(file_exists=True)
+ service = RAGRuntimeService(app)
+ with pytest.raises(ValueError, match='Invalid storage path'):
+ await service.get_file_stream(CONTEXT, unsafe_path)
+ app.storage_mgr.load_scoped_object_key.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_runtime_rejects_wrong_instance_binding():
+ app = _app()
+ app.workspace_service.get_execution_binding.return_value = SimpleNamespace(instance_uuid='instance-b')
+ service = RAGRuntimeService(app)
+
+ with pytest.raises(Exception, match='another LangBot instance'):
+ await service.vector_search(CONTEXT, 'kb-a', [0.1], 1)
+ app.persistence_mgr.execute_async.assert_not_awaited()
+
+
+# The following classes preserve the pre-tenancy regression scenarios. They
+# intentionally exercise the same inputs through the new trusted
+# ExecutionContext and assert the canonical Workspace-owned KB forwarded to the
+# vector layer.
+class TestRAGRuntimeServiceVectorUpsertRegression:
@pytest.mark.asyncio
async def test_vector_upsert_basic(self):
- """Basic vector upsert delegates to vector_db_mgr."""
- mock_app = self._create_mock_app()
+ app = _app()
+ service = RAGRuntimeService(app)
+ vectors = [[0.1, 0.2], [0.3, 0.4]]
+ ids = ['id1', 'id2']
- mocks = self._make_rag_import_mocks()
+ await service.vector_upsert(CONTEXT, 'test_collection', vectors, ids)
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- vectors = [[0.1, 0.2], [0.3, 0.4]]
- ids = ['id1', 'id2']
-
- await service.vector_upsert(
- collection_id='test_collection',
- vectors=vectors,
- ids=ids,
- )
-
- mock_app.vector_db_mgr.upsert.assert_called_once()
- call_args = mock_app.vector_db_mgr.upsert.call_args
- assert call_args.kwargs['collection_name'] == 'test_collection'
- assert call_args.kwargs['vectors'] == vectors
- assert call_args.kwargs['ids'] == ids
- # Default metadata is empty dicts
- assert call_args.kwargs['metadata'] == [{} for _ in vectors]
+ app.vector_db_mgr.upsert.assert_awaited_once_with(
+ execution_context=CONTEXT,
+ knowledge_base_uuid='kb-a',
+ vectors=vectors,
+ ids=ids,
+ metadata=[{}, {}],
+ documents=None,
+ )
@pytest.mark.asyncio
async def test_vector_upsert_with_metadata(self):
- """Vector upsert with provided metadata."""
- mock_app = self._create_mock_app()
+ app = _app()
+ metadata = [{'file_id': 'abc', 'page': 1}]
- mocks = self._make_rag_import_mocks()
+ await RAGRuntimeService(app).vector_upsert(
+ CONTEXT,
+ 'test',
+ [[0.1, 0.2]],
+ ['id1'],
+ metadata=metadata,
+ )
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- vectors = [[0.1, 0.2]]
- ids = ['id1']
- metadata = [{'file_id': 'abc', 'page': 1}]
-
- await service.vector_upsert(
- collection_id='test',
- vectors=vectors,
- ids=ids,
- metadata=metadata,
- )
-
- call_args = mock_app.vector_db_mgr.upsert.call_args
- assert call_args.kwargs['metadata'] == metadata
+ assert app.vector_db_mgr.upsert.await_args.kwargs['metadata'] == metadata
@pytest.mark.asyncio
async def test_vector_upsert_with_documents(self):
- """Vector upsert with documents for full-text search."""
- mock_app = self._create_mock_app()
+ app = _app()
+ documents = ['This is a test document']
- mocks = self._make_rag_import_mocks()
-
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- vectors = [[0.1, 0.2]]
- ids = ['id1']
- documents = ['This is a test document']
-
- await service.vector_upsert(
- collection_id='test',
- vectors=vectors,
- ids=ids,
- documents=documents,
- )
-
- call_args = mock_app.vector_db_mgr.upsert.call_args
- assert call_args.kwargs['documents'] == documents
-
-
-class TestRAGRuntimeServiceVectorSearch:
- """Tests for vector_search method."""
-
- def _create_mock_app(self):
- """Create mock app."""
- mock_app = MagicMock()
- mock_app.vector_db_mgr = MagicMock()
- mock_app.vector_db_mgr.search = AsyncMock(
- return_value=[
- {'id': 'id1', 'distance': 0.1, 'metadata': {'file_id': 'abc'}},
- {'id': 'id2', 'distance': 0.2, 'metadata': {'file_id': 'def'}},
- ]
+ await RAGRuntimeService(app).vector_upsert(
+ CONTEXT,
+ 'test',
+ [[0.1, 0.2]],
+ ['id1'],
+ documents=documents,
)
- return mock_app
- def _make_rag_import_mocks(self):
- return {
- 'langbot.pkg.core.app': MagicMock(),
- 'langbot_plugin.api.entities.builtin.rag': MagicMock(),
- }
+ assert app.vector_db_mgr.upsert.await_args.kwargs['documents'] == documents
+
+class TestRAGRuntimeServiceVectorSearchRegression:
@pytest.mark.asyncio
async def test_vector_search_basic(self):
- """Basic vector search delegates to vector_db_mgr."""
- mock_app = self._create_mock_app()
+ app = _app()
+ query_vector = [0.1, 0.2, 0.3]
- mocks = self._make_rag_import_mocks()
+ result = await RAGRuntimeService(app).vector_search(
+ CONTEXT,
+ 'test',
+ query_vector,
+ 5,
+ )
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- query_vector = [0.1, 0.2, 0.3]
-
- result = await service.vector_search(
- collection_id='test',
- query_vector=query_vector,
- top_k=5,
- )
-
- assert len(result) == 2
- mock_app.vector_db_mgr.search.assert_called_once()
- call_args = mock_app.vector_db_mgr.search.call_args
- assert call_args.kwargs['collection_name'] == 'test'
- assert call_args.kwargs['query_vector'] == query_vector
- assert call_args.kwargs['limit'] == 5
+ assert result == [{'id': 'chunk-a'}]
+ app.vector_db_mgr.search.assert_awaited_once_with(
+ execution_context=CONTEXT,
+ knowledge_base_uuid='kb-a',
+ query_vector=query_vector,
+ limit=5,
+ filter=None,
+ search_type='vector',
+ query_text='',
+ vector_weight=None,
+ )
@pytest.mark.asyncio
async def test_vector_search_with_filters(self):
- """Vector search with metadata filters."""
- mock_app = self._create_mock_app()
+ app = _app()
+ filters = {'file_id': 'abc'}
- mocks = self._make_rag_import_mocks()
+ await RAGRuntimeService(app).vector_search(
+ CONTEXT,
+ 'test',
+ [0.1, 0.2],
+ 10,
+ filters=filters,
+ )
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- filters = {'file_id': 'abc'}
-
- await service.vector_search(
- collection_id='test',
- query_vector=[0.1, 0.2],
- top_k=10,
- filters=filters,
- )
-
- call_args = mock_app.vector_db_mgr.search.call_args
- assert call_args.kwargs['filter'] == filters
+ assert app.vector_db_mgr.search.await_args.kwargs['filter'] == filters
@pytest.mark.asyncio
async def test_vector_search_hybrid_mode(self):
- """Vector search with hybrid search type."""
- mock_app = self._create_mock_app()
+ app = _app()
- mocks = self._make_rag_import_mocks()
+ await RAGRuntimeService(app).vector_search(
+ CONTEXT,
+ 'test',
+ [0.1, 0.2],
+ 10,
+ search_type='hybrid',
+ query_text='search query',
+ vector_weight=0.7,
+ )
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- await service.vector_search(
- collection_id='test',
- query_vector=[0.1, 0.2],
- top_k=10,
- search_type='hybrid',
- query_text='search query',
- vector_weight=0.7,
- )
-
- call_args = mock_app.vector_db_mgr.search.call_args
- assert call_args.kwargs['search_type'] == 'hybrid'
- assert call_args.kwargs['query_text'] == 'search query'
- assert call_args.kwargs['vector_weight'] == 0.7
+ kwargs = app.vector_db_mgr.search.await_args.kwargs
+ assert kwargs['search_type'] == 'hybrid'
+ assert kwargs['query_text'] == 'search query'
+ assert kwargs['vector_weight'] == 0.7
-class TestRAGRuntimeServiceVectorDelete:
- """Tests for vector_delete method."""
-
- def _create_mock_app(self):
- mock_app = MagicMock()
- mock_app.vector_db_mgr = MagicMock()
- mock_app.vector_db_mgr.delete_by_file_id = AsyncMock()
- mock_app.vector_db_mgr.delete_by_filter = AsyncMock(return_value=5)
- return mock_app
-
- def _make_rag_import_mocks(self):
- return {
- 'langbot.pkg.core.app': MagicMock(),
- 'langbot_plugin.api.entities.builtin.rag': MagicMock(),
- }
-
+class TestRAGRuntimeServiceVectorDeleteRegression:
@pytest.mark.asyncio
async def test_vector_delete_by_file_ids(self):
- """Delete by file_ids delegates to delete_by_file_id."""
- mock_app = self._create_mock_app()
+ app = _app()
- mocks = self._make_rag_import_mocks()
+ result = await RAGRuntimeService(app).vector_delete(
+ CONTEXT,
+ 'test',
+ file_ids=['file1', 'file2', 'file3'],
+ )
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- result = await service.vector_delete(
- collection_id='test',
- file_ids=['file1', 'file2', 'file3'],
- )
-
- assert result == 3 # Returns count of file_ids
- mock_app.vector_db_mgr.delete_by_file_id.assert_called_once()
- call_args = mock_app.vector_db_mgr.delete_by_file_id.call_args
- assert call_args.kwargs['collection_name'] == 'test'
- assert call_args.kwargs['file_ids'] == ['file1', 'file2', 'file3']
+ assert result == 3
+ app.vector_db_mgr.delete_by_file_id.assert_awaited_once_with(
+ execution_context=CONTEXT,
+ knowledge_base_uuid='kb-a',
+ file_ids=['file1', 'file2', 'file3'],
+ )
@pytest.mark.asyncio
async def test_vector_delete_by_filters(self):
- """Delete by filters delegates to delete_by_filter."""
- mock_app = self._create_mock_app()
+ app = _app()
+ filters = {'status': 'deleted'}
+ app.vector_db_mgr.delete_by_filter.return_value = 5
- mocks = self._make_rag_import_mocks()
+ result = await RAGRuntimeService(app).vector_delete(
+ CONTEXT,
+ 'test',
+ filters=filters,
+ )
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- filters = {'status': 'deleted'}
-
- result = await service.vector_delete(
- collection_id='test',
- filters=filters,
- )
-
- assert result == 5 # Returns count from delete_by_filter
- mock_app.vector_db_mgr.delete_by_filter.assert_called_once()
- call_args = mock_app.vector_db_mgr.delete_by_filter.call_args
- assert call_args.kwargs['collection_name'] == 'test'
- assert call_args.kwargs['filter'] == filters
+ assert result == 5
+ assert app.vector_db_mgr.delete_by_filter.await_args.kwargs['filter'] == filters
@pytest.mark.asyncio
async def test_vector_delete_no_params(self):
- """Delete with no params returns 0."""
- mock_app = self._create_mock_app()
+ app = _app()
- mocks = self._make_rag_import_mocks()
+ result = await RAGRuntimeService(app).vector_delete(CONTEXT, 'test')
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- result = await service.vector_delete(collection_id='test')
-
- assert result == 0
- mock_app.vector_db_mgr.delete_by_file_id.assert_not_called()
- mock_app.vector_db_mgr.delete_by_filter.assert_not_called()
+ assert result == 0
+ app.vector_db_mgr.delete_by_file_id.assert_not_awaited()
+ app.vector_db_mgr.delete_by_filter.assert_not_awaited()
-class TestRAGRuntimeServiceVectorList:
- """Tests for vector_list method."""
-
- def _create_mock_app(self):
- mock_app = MagicMock()
- mock_app.vector_db_mgr = MagicMock()
- mock_app.vector_db_mgr.list_by_filter = AsyncMock(
- return_value=([{'id': 'id1', 'metadata': {'file_id': 'abc'}}], 10)
- )
- return mock_app
-
- def _make_rag_import_mocks(self):
- return {
- 'langbot.pkg.core.app': MagicMock(),
- 'langbot_plugin.api.entities.builtin.rag': MagicMock(),
- }
-
+class TestRAGRuntimeServiceVectorListRegression:
@pytest.mark.asyncio
async def test_vector_list_basic(self):
- """Basic vector list delegates to vector_db_mgr."""
- mock_app = self._create_mock_app()
+ app = _app()
- mocks = self._make_rag_import_mocks()
+ items, total = await RAGRuntimeService(app).vector_list(CONTEXT, 'test')
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- items, total = await service.vector_list(
- collection_id='test',
- )
-
- assert len(items) == 1
- assert total == 10
- mock_app.vector_db_mgr.list_by_filter.assert_called_once()
- call_args = mock_app.vector_db_mgr.list_by_filter.call_args
- assert call_args.kwargs['collection_name'] == 'test'
- assert call_args.kwargs['limit'] == 20 # Default
- assert call_args.kwargs['offset'] == 0 # Default
+ assert items == [{'id': 'chunk-a'}]
+ assert total == 1
+ app.vector_db_mgr.list_by_filter.assert_awaited_once_with(
+ execution_context=CONTEXT,
+ knowledge_base_uuid='kb-a',
+ filter=None,
+ limit=20,
+ offset=0,
+ )
@pytest.mark.asyncio
async def test_vector_list_with_pagination(self):
- """Vector list with custom pagination."""
- mock_app = self._create_mock_app()
+ app = _app()
- mocks = self._make_rag_import_mocks()
+ await RAGRuntimeService(app).vector_list(CONTEXT, 'test', limit=50, offset=100)
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- await service.vector_list(
- collection_id='test',
- limit=50,
- offset=100,
- )
-
- call_args = mock_app.vector_db_mgr.list_by_filter.call_args
- assert call_args.kwargs['limit'] == 50
- assert call_args.kwargs['offset'] == 100
+ kwargs = app.vector_db_mgr.list_by_filter.await_args.kwargs
+ assert kwargs['limit'] == 50
+ assert kwargs['offset'] == 100
@pytest.mark.asyncio
async def test_vector_list_with_filters(self):
- """Vector list with metadata filters."""
- mock_app = self._create_mock_app()
+ app = _app()
+ filters = {'file_id': 'abc'}
- mocks = self._make_rag_import_mocks()
+ await RAGRuntimeService(app).vector_list(CONTEXT, 'test', filters=filters)
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- filters = {'file_id': 'abc'}
-
- await service.vector_list(
- collection_id='test',
- filters=filters,
- )
-
- call_args = mock_app.vector_db_mgr.list_by_filter.call_args
- assert call_args.kwargs['filter'] == filters
+ assert app.vector_db_mgr.list_by_filter.await_args.kwargs['filter'] == filters
-class TestRAGRuntimeServiceGetFileStream:
- """Tests for get_file_stream method."""
-
- def _create_mock_app(self):
- mock_app = MagicMock()
- mock_app.vector_db_mgr = MagicMock()
- mock_app.storage_mgr = MagicMock()
- mock_app.storage_mgr.storage_provider = MagicMock()
- mock_app.storage_mgr.storage_provider.load = AsyncMock(return_value=b'file content')
- return mock_app
-
- def _make_rag_import_mocks(self):
- return {
- 'langbot.pkg.core.app': MagicMock(),
- 'langbot_plugin.api.entities.builtin.rag': MagicMock(),
- }
-
+class TestRAGRuntimeServiceGetFileStreamRegression:
@pytest.mark.asyncio
async def test_get_file_stream_basic(self):
- """Get file stream loads from storage."""
- mock_app = self._create_mock_app()
+ app = _app(file_exists=True)
+ app.persistence_mgr.execute_async = AsyncMock(return_value=_ScalarResult('file-a'))
- mocks = self._make_rag_import_mocks()
+ result = await RAGRuntimeService(app).get_file_stream(
+ CONTEXT,
+ 'knowledge/files/doc.pdf',
+ )
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- result = await service.get_file_stream('knowledge/files/doc.pdf')
-
- assert result == b'file content'
- mock_app.storage_mgr.storage_provider.load.assert_called_once_with('knowledge/files/doc.pdf')
+ assert result == b'content'
+ app.storage_mgr.load_scoped_object_key.assert_awaited_once_with(
+ CONTEXT,
+ 'knowledge/files/doc.pdf',
+ expected_owner_type='upload_document',
+ )
@pytest.mark.asyncio
async def test_get_file_stream_empty_result(self):
- """Empty file returns empty bytes."""
- mock_app = self._create_mock_app()
- mock_app.storage_mgr.storage_provider.load = AsyncMock(return_value=None)
+ app = _app(file_exists=True)
+ app.persistence_mgr.execute_async = AsyncMock(return_value=_ScalarResult('file-a'))
+ app.storage_mgr.load_scoped_object_key.return_value = None
- mocks = self._make_rag_import_mocks()
+ result = await RAGRuntimeService(app).get_file_stream(CONTEXT, 'nonexistent.pdf')
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- result = await service.get_file_stream('nonexistent.pdf')
-
- assert result == b''
+ assert result == b''
@pytest.mark.asyncio
async def test_get_file_stream_normalizes_safe_path(self):
- """Safe relative paths are normalized before loading."""
- mock_app = self._create_mock_app()
+ app = _app(file_exists=True)
+ app.persistence_mgr.execute_async = AsyncMock(return_value=_ScalarResult('file-a'))
- mocks = self._make_rag_import_mocks()
+ result = await RAGRuntimeService(app).get_file_stream(
+ CONTEXT,
+ 'knowledge/./files/doc.pdf',
+ )
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- result = await service.get_file_stream('knowledge/./files/doc.pdf')
-
- assert result == b'file content'
- mock_app.storage_mgr.storage_provider.load.assert_called_once_with('knowledge/files/doc.pdf')
+ assert result == b'content'
+ app.storage_mgr.load_scoped_object_key.assert_awaited_once_with(
+ CONTEXT,
+ 'knowledge/files/doc.pdf',
+ expected_owner_type='upload_document',
+ )
@pytest.mark.asyncio
async def test_get_file_stream_path_traversal_blocked(self):
- """Path traversal attacks are blocked."""
- mock_app = self._create_mock_app()
+ app = _app(file_exists=True)
+ service = RAGRuntimeService(app)
- mocks = self._make_rag_import_mocks()
-
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- # Absolute path should raise ValueError
- with pytest.raises(ValueError, match='Invalid storage path'):
- await service.get_file_stream('/etc/passwd')
-
- # Path traversal should raise ValueError
- with pytest.raises(ValueError, match='Invalid storage path'):
- await service.get_file_stream('knowledge/../../../etc/passwd')
+ with pytest.raises(ValueError, match='Invalid storage path'):
+ await service.get_file_stream(CONTEXT, '/etc/passwd')
+ with pytest.raises(ValueError, match='Invalid storage path'):
+ await service.get_file_stream(CONTEXT, 'knowledge/../../../etc/passwd')
@pytest.mark.asyncio
@pytest.mark.parametrize(
@@ -486,36 +479,22 @@ class TestRAGRuntimeServiceGetFileStream:
],
)
async def test_get_file_stream_rejects_unsafe_paths(self, storage_path: str):
- """Unsafe runtime file paths are rejected before storage load."""
- mock_app = self._create_mock_app()
+ app = _app(file_exists=True)
- mocks = self._make_rag_import_mocks()
+ with pytest.raises(ValueError, match='Invalid storage path'):
+ await RAGRuntimeService(app).get_file_stream(CONTEXT, storage_path)
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- with pytest.raises(ValueError, match='Invalid storage path'):
- await service.get_file_stream(storage_path)
-
- mock_app.storage_mgr.storage_provider.load.assert_not_called()
+ app.storage_mgr.load_scoped_object_key.assert_not_awaited()
@pytest.mark.asyncio
async def test_get_file_stream_normalizes_path(self):
- """Valid paths with .. in filename (not traversal) should work."""
- mock_app = self._create_mock_app()
+ app = _app(file_exists=True)
+ app.persistence_mgr.execute_async = AsyncMock(return_value=_ScalarResult('file-a'))
- mocks = self._make_rag_import_mocks()
+ await RAGRuntimeService(app).get_file_stream(CONTEXT, 'knowledge/files/test.pdf')
- with isolated_sys_modules(mocks):
- from langbot.pkg.rag.service.runtime import RAGRuntimeService
-
- service = RAGRuntimeService(mock_app)
-
- # Path that contains '..' as part of filename (not traversal)
- # This should NOT raise - posixpath.normpath handles this
- # But the current implementation checks '..' in split('/')
- # Let's test a simple valid path
- await service.get_file_stream('knowledge/files/test.pdf')
- mock_app.storage_mgr.storage_provider.load.assert_called()
+ app.storage_mgr.load_scoped_object_key.assert_awaited_once_with(
+ CONTEXT,
+ 'knowledge/files/test.pdf',
+ expected_owner_type='upload_document',
+ )
diff --git a/tests/unit_tests/rag/test_tenant_isolation.py b/tests/unit_tests/rag/test_tenant_isolation.py
new file mode 100644
index 000000000..9824b8060
--- /dev/null
+++ b/tests/unit_tests/rag/test_tenant_isolation.py
@@ -0,0 +1,350 @@
+from __future__ import annotations
+
+import datetime
+from types import SimpleNamespace
+from unittest.mock import AsyncMock, Mock
+
+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.knowledge import KnowledgeService
+from langbot.pkg.entity.persistence.base import Base
+from langbot.pkg.entity.persistence.rag import File, KnowledgeBase
+from langbot.pkg.entity.persistence.workspace import Workspace, WorkspaceExecutionState
+from langbot.pkg.rag.knowledge.kbmgr import RAGManager
+from langbot.pkg.rag.service.runtime import RAGRuntimeService
+from langbot.pkg.vector.mgr import VectorDBManager
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
+from langbot.pkg.workspace.policy import CloudWorkspacePolicy, SingleWorkspacePolicy
+from langbot.pkg.workspace.service import WorkspaceService
+
+
+pytestmark = pytest.mark.asyncio
+
+INSTANCE_UUID = 'instance-rag-isolation'
+WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
+WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
+
+
+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
+ }
+
+
+class _RecordingVectorDatabase:
+ def __init__(self):
+ self.collections: list[str] = []
+ self.metadatas: list[list[dict]] = []
+ self.calls: list[tuple[str, str]] = []
+
+ async def add_embeddings(self, *, collection, ids, embeddings_list, metadatas, documents):
+ self.collections.append(collection)
+ self.metadatas.append(metadatas)
+ self.calls.append(('upsert', collection))
+
+ async def search(self, *, collection, **_kwargs):
+ self.calls.append(('search', collection))
+ return {'ids': [[]], 'distances': [[]], 'metadatas': [[]]}
+
+ async def delete_by_file_id(self, collection, _file_id):
+ self.calls.append(('delete_by_file_id', collection))
+
+ async def delete_collection(self, collection):
+ self.calls.append(('delete_collection', collection))
+
+ async def delete_by_filter(self, collection, _filter):
+ self.calls.append(('delete_by_filter', collection))
+ return 1
+
+ async def list_by_filter(self, collection, _filter, _limit, _offset):
+ self.calls.append(('list_by_filter', collection))
+ return [], 0
+
+
+@pytest.fixture
+async def tenant_rag(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "rag-tenant.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_UUID,
+ 'name': 'Workspace A',
+ 'slug': 'workspace-a',
+ 'source': 'local',
+ },
+ {
+ 'uuid': WORKSPACE_B,
+ 'instance_uuid': INSTANCE_UUID,
+ 'name': 'Workspace B',
+ 'slug': 'workspace-b',
+ 'source': 'cloud_projection',
+ },
+ ],
+ )
+ await connection.execute(
+ sqlalchemy.insert(WorkspaceExecutionState),
+ [
+ {
+ 'workspace_uuid': WORKSPACE_A,
+ 'instance_uuid': INSTANCE_UUID,
+ 'active_generation': 3,
+ 'state': 'active',
+ 'source': 'local',
+ 'write_fenced': False,
+ },
+ {
+ 'workspace_uuid': WORKSPACE_B,
+ 'instance_uuid': INSTANCE_UUID,
+ 'active_generation': 3,
+ 'state': 'active',
+ 'source': 'cloud',
+ 'write_fenced': False,
+ },
+ ],
+ )
+ await connection.execute(
+ sqlalchemy.insert(KnowledgeBase),
+ [
+ {
+ 'uuid': 'kb-a',
+ 'workspace_uuid': WORKSPACE_A,
+ 'name': 'Same Knowledge Base',
+ 'description': 'A',
+ 'knowledge_engine_plugin_id': 'author/engine',
+ 'collection_id': 'kb-a',
+ 'creation_settings': {},
+ 'retrieval_settings': {},
+ },
+ {
+ 'uuid': 'kb-b',
+ 'workspace_uuid': WORKSPACE_B,
+ 'name': 'Same Knowledge Base',
+ 'description': 'B',
+ 'knowledge_engine_plugin_id': 'author/engine',
+ 'collection_id': 'kb-b',
+ 'creation_settings': {},
+ 'retrieval_settings': {},
+ },
+ ],
+ )
+ await connection.execute(
+ sqlalchemy.insert(File),
+ [
+ {
+ 'uuid': 'file-a',
+ 'workspace_uuid': WORKSPACE_A,
+ 'kb_id': 'kb-a',
+ 'file_name': 'a.pdf',
+ 'extension': 'pdf',
+ },
+ {
+ 'uuid': 'file-b',
+ 'workspace_uuid': WORKSPACE_B,
+ 'kb_id': 'kb-b',
+ 'file_name': 'b.pdf',
+ 'extension': 'pdf',
+ },
+ ],
+ )
+
+ app = SimpleNamespace()
+ app.persistence_mgr = _PersistenceManager(engine)
+ app.logger = Mock()
+ app.workspace_policy = SingleWorkspacePolicy()
+ app.plugin_connector = SimpleNamespace(
+ is_enable_plugin=False,
+ rag_on_kb_create=AsyncMock(),
+ rag_on_kb_delete=AsyncMock(),
+ )
+ app.workspace_service = WorkspaceService(app, instance_uuid=INSTANCE_UUID)
+ app.rag_mgr = RAGManager(app)
+ await app.rag_mgr.initialize()
+ app.knowledge_service = KnowledgeService(app)
+
+ yield app, engine
+ await engine.dispose()
+
+
+def _context(workspace_uuid: str) -> ExecutionContext:
+ return ExecutionContext(
+ instance_uuid=INSTANCE_UUID,
+ workspace_uuid=workspace_uuid,
+ placement_generation=3,
+ )
+
+
+async def test_context_is_mandatory_and_same_names_are_isolated(tenant_rag):
+ app, _engine = tenant_rag
+
+ with pytest.raises(WorkspaceRequiredError):
+ await app.knowledge_service.get_knowledge_bases(None)
+
+ bases_a = await app.knowledge_service.get_knowledge_bases(_context(WORKSPACE_A))
+ bases_b = await app.knowledge_service.get_knowledge_bases(_context(WORKSPACE_B))
+ assert [(item['uuid'], item['name']) for item in bases_a] == [('kb-a', 'Same Knowledge Base')]
+ assert [(item['uuid'], item['name']) for item in bases_b] == [('kb-b', 'Same Knowledge Base')]
+
+
+async def test_cross_workspace_uuid_and_file_guessing_return_not_found(tenant_rag):
+ app, engine = tenant_rag
+ context_a = _context(WORKSPACE_A)
+
+ assert await app.knowledge_service.get_knowledge_base(context_a, 'kb-b') is None
+ with pytest.raises(WorkspaceNotFoundError):
+ await app.knowledge_service.update_knowledge_base(context_a, 'kb-b', {'name': 'stolen'})
+ with pytest.raises(WorkspaceNotFoundError):
+ await app.knowledge_service.delete_knowledge_base(context_a, 'kb-b')
+ with pytest.raises(WorkspaceNotFoundError):
+ await app.knowledge_service.get_files_by_knowledge_base(context_a, 'kb-b')
+ with pytest.raises(WorkspaceNotFoundError):
+ await app.knowledge_service.delete_file(context_a, 'kb-b', 'file-b')
+
+ async with engine.connect() as connection:
+ assert (
+ await connection.scalar(sqlalchemy.select(KnowledgeBase.name).where(KnowledgeBase.uuid == 'kb-b'))
+ == 'Same Knowledge Base'
+ )
+ assert await connection.scalar(sqlalchemy.select(File.uuid).where(File.uuid == 'file-b')) == 'file-b'
+
+
+async def test_runtime_rejects_cross_workspace_collection_reference(tenant_rag):
+ app, _engine = tenant_rag
+ app.vector_db_mgr = SimpleNamespace(upsert=AsyncMock())
+ service = RAGRuntimeService(app)
+
+ with pytest.raises(WorkspaceNotFoundError):
+ await service.vector_upsert(
+ _context(WORKSPACE_A),
+ 'kb-b',
+ vectors=[[0.1, 0.2]],
+ ids=['chunk-1'],
+ )
+ app.vector_db_mgr.upsert.assert_not_awaited()
+
+
+async def test_physical_vector_handles_do_not_collide_across_workspaces(tenant_rag):
+ app, _engine = tenant_rag
+ database = _RecordingVectorDatabase()
+ manager = VectorDBManager(app)
+ manager.vector_db = database
+
+ context_a = _context(WORKSPACE_A)
+ context_b = _context(WORKSPACE_B)
+ assert manager.physical_collection_name(context_a, 'same-kb-id') != manager.physical_collection_name(
+ context_b,
+ 'same-kb-id',
+ )
+
+ await manager.upsert(context_a, 'kb-a', [[0.1]], ['a'], metadata=[{'source': 'client'}])
+ await manager.upsert(context_b, 'kb-b', [[0.2]], ['b'], metadata=[{'source': 'client'}])
+ assert len(set(database.collections)) == 2
+ assert database.metadatas[0][0]['_langbot_workspace_uuid'] == WORKSPACE_A
+ assert database.metadatas[1][0]['_langbot_workspace_uuid'] == WORKSPACE_B
+
+
+async def test_migrated_local_kb_keeps_legacy_collection_for_every_vector_operation(tenant_rag):
+ app, engine = tenant_rag
+ legacy_collection = 'legacy-collection-kb-a'
+ async with engine.begin() as connection:
+ await connection.execute(
+ sqlalchemy.update(KnowledgeBase)
+ .where(KnowledgeBase.uuid == 'kb-a')
+ .values(
+ collection_id=legacy_collection,
+ legacy_vector_collection=True,
+ )
+ )
+
+ database = _RecordingVectorDatabase()
+ manager = VectorDBManager(app)
+ manager.vector_db = database
+ context = _context(WORKSPACE_A)
+
+ await manager.upsert(context, 'kb-a', [[0.1]], ['chunk-a'])
+ await manager.search(context, 'kb-a', [0.1], 3)
+ await manager.delete_by_file_id(context, 'kb-a', ['file-a'])
+ assert await manager.delete_by_filter(context, 'kb-a', {'file_id': 'file-a'}) == 1
+ assert await manager.list_by_filter(context, 'kb-a', {'file_id': 'file-a'}) == ([], 0)
+ await manager.delete_collection(context, 'kb-a')
+
+ assert database.calls == [
+ ('upsert', legacy_collection),
+ ('search', legacy_collection),
+ ('delete_by_file_id', legacy_collection),
+ ('delete_by_filter', legacy_collection),
+ ('list_by_filter', legacy_collection),
+ ('delete_collection', legacy_collection),
+ ]
+
+
+@pytest.mark.parametrize('deny_by', ['projected_workspace', 'multi_workspace_policy'])
+async def test_legacy_marker_is_ignored_outside_single_local_workspace(tenant_rag, deny_by):
+ app, engine = tenant_rag
+ workspace_uuid = WORKSPACE_B if deny_by == 'projected_workspace' else WORKSPACE_A
+ kb_uuid = 'kb-b' if deny_by == 'projected_workspace' else 'kb-a'
+ legacy_collection = f'legacy-{deny_by}'
+ async with engine.begin() as connection:
+ await connection.execute(
+ sqlalchemy.update(KnowledgeBase)
+ .where(KnowledgeBase.uuid == kb_uuid)
+ .values(
+ collection_id=legacy_collection,
+ legacy_vector_collection=True,
+ )
+ )
+ if deny_by == 'multi_workspace_policy':
+ app.workspace_policy = CloudWorkspacePolicy()
+
+ database = _RecordingVectorDatabase()
+ manager = VectorDBManager(app)
+ manager.vector_db = database
+ context = _context(workspace_uuid)
+ await manager.search(context, kb_uuid, [0.1], 3)
+
+ assert database.calls == [('search', manager.physical_collection_name(context, kb_uuid))]
+ assert database.calls[0][1] != legacy_collection
+ app.logger.warning.assert_called_once()
+
+
+async def test_stale_generation_is_rejected_before_vector_access(tenant_rag):
+ app, _engine = tenant_rag
+ database = _RecordingVectorDatabase()
+ manager = VectorDBManager(app)
+ manager.vector_db = database
+ stale = ExecutionContext(
+ instance_uuid=INSTANCE_UUID,
+ workspace_uuid=WORKSPACE_A,
+ placement_generation=2,
+ )
+
+ with pytest.raises(Exception, match='generation'):
+ await manager.upsert(stale, 'kb-a', [[0.1]], ['a'])
+ assert database.collections == []
diff --git a/tests/unit_tests/storage/test_workspace_scoping.py b/tests/unit_tests/storage/test_workspace_scoping.py
new file mode 100644
index 000000000..7d8407248
--- /dev/null
+++ b/tests/unit_tests/storage/test_workspace_scoping.py
@@ -0,0 +1,235 @@
+from __future__ import annotations
+
+from types import SimpleNamespace
+from unittest.mock import AsyncMock
+
+import pytest
+
+from langbot.pkg.api.http.authz import WorkspaceRequiredError
+from langbot.pkg.api.http.context import ExecutionContext
+from langbot.pkg.storage.mgr import StorageMgr
+
+
+WORKSPACE_A = '00000000-0000-0000-0000-00000000000a'
+WORKSPACE_B = '00000000-0000-0000-0000-00000000000b'
+
+
+def _context(workspace_uuid: str, generation: int = 7) -> ExecutionContext:
+ return ExecutionContext(
+ instance_uuid='instance',
+ workspace_uuid=workspace_uuid,
+ placement_generation=generation,
+ )
+
+
+def _context_for_instance(instance_uuid: str) -> ExecutionContext:
+ return ExecutionContext(
+ instance_uuid=instance_uuid,
+ workspace_uuid=WORKSPACE_A,
+ placement_generation=7,
+ )
+
+
+class _Provider:
+ def __init__(self):
+ self.values: dict[str, bytes] = {}
+
+ async def save(self, key: str, value: bytes):
+ self.values[key] = value
+
+ async def load(self, key: str) -> bytes:
+ return self.values[key]
+
+ async def exists(self, key: str) -> bool:
+ return key in self.values
+
+ async def size(self, key: str) -> int:
+ return len(self.values[key])
+
+ async def delete(self, key: str):
+ self.values.pop(key, None)
+
+
+class _WorkspaceService:
+ instance_uuid = 'instance'
+
+ async def get_execution_binding(self, workspace_uuid, expected_generation=None):
+ if workspace_uuid not in {WORKSPACE_A, WORKSPACE_B} or expected_generation != 7:
+ raise ValueError('inactive or stale binding')
+ return SimpleNamespace(
+ instance_uuid='instance',
+ workspace_uuid=workspace_uuid,
+ placement_generation=7,
+ )
+
+
+@pytest.fixture
+def manager():
+ application = SimpleNamespace(workspace_service=_WorkspaceService())
+ storage = StorageMgr(application)
+ storage.storage_provider = _Provider()
+ return storage
+
+
+def test_binary_storage_canonical_key_covers_every_scope_dimension(manager):
+ baseline = manager.canonical_binary_storage_key(
+ _context(WORKSPACE_A),
+ owner_type='plugin',
+ owner='author/name',
+ key='same-key',
+ )
+ assert baseline != manager.canonical_binary_storage_key(
+ _context(WORKSPACE_B),
+ owner_type='plugin',
+ owner='author/name',
+ key='same-key',
+ )
+ assert baseline != manager.canonical_binary_storage_key(
+ _context_for_instance('other-instance'),
+ owner_type='plugin',
+ owner='author/name',
+ key='same-key',
+ )
+ assert baseline != manager.canonical_binary_storage_key(
+ _context(WORKSPACE_A),
+ owner_type='workspace',
+ owner='author/name',
+ key='same-key',
+ )
+ assert baseline != manager.canonical_binary_storage_key(
+ _context(WORKSPACE_A),
+ owner_type='plugin',
+ owner='other/name',
+ key='same-key',
+ )
+ assert baseline != manager.canonical_binary_storage_key(
+ _context(WORKSPACE_A),
+ owner_type='plugin',
+ owner='author/name',
+ key='other-key',
+ )
+
+
+@pytest.mark.asyncio
+async def test_public_object_route_derives_trusted_workspace(manager):
+ object_key = await manager.save_scoped(
+ _context(WORKSPACE_A),
+ owner_type='upload',
+ owner='account:a',
+ key='photo.png',
+ value=b'image-a',
+ )
+ assert await manager.resolve_public_object(object_key, expected_owner_type='upload') == b'image-a'
+ assert await manager.resolve_public_object(object_key, expected_owner_type='plugin') is None
+
+ guessed_workspace_key = object_key.replace(WORKSPACE_A, WORKSPACE_B)
+ assert await manager.resolve_public_object(guessed_workspace_key, expected_owner_type='upload') is None
+
+ stale_generation_key = object_key.replace('/7/upload/', '/8/upload/')
+ assert await manager.resolve_public_object(stale_generation_key, expected_owner_type='upload') is None
+
+
+@pytest.mark.asyncio
+async def test_storage_operations_without_context_fail_closed(manager):
+ with pytest.raises(WorkspaceRequiredError):
+ await manager.save_scoped(
+ None,
+ owner_type='upload',
+ owner='account:a',
+ key='photo.png',
+ value=b'image',
+ )
+
+
+@pytest.mark.asyncio
+async def test_opaque_object_operations_reject_cross_scope_and_owner_type(manager):
+ object_key = await manager.save_scoped(
+ _context(WORKSPACE_A),
+ owner_type='upload',
+ owner='account:a',
+ key='document.pdf',
+ value=b'document-a',
+ )
+
+ assert await manager.exists_scoped_object_key(
+ _context(WORKSPACE_A),
+ object_key,
+ expected_owner_type='upload',
+ )
+ assert (
+ await manager.load_scoped_object_key(
+ _context(WORKSPACE_A),
+ object_key,
+ expected_owner_type='upload',
+ )
+ == b'document-a'
+ )
+ assert await manager.size_scoped_object_key(
+ _context(WORKSPACE_A),
+ object_key,
+ expected_owner_type='upload',
+ ) == len(b'document-a')
+
+ for wrong_context in (_context(WORKSPACE_B), _context(WORKSPACE_A, generation=8)):
+ with pytest.raises(WorkspaceRequiredError):
+ await manager.load_scoped_object_key(
+ wrong_context,
+ object_key,
+ expected_owner_type='upload',
+ )
+ with pytest.raises(WorkspaceRequiredError):
+ await manager.delete_scoped_object_key(
+ wrong_context,
+ object_key,
+ expected_owner_type='upload',
+ )
+
+ with pytest.raises(WorkspaceRequiredError):
+ await manager.load_scoped_object_key(
+ _context(WORKSPACE_A),
+ object_key,
+ expected_owner_type='plugin_config',
+ )
+
+ assert object_key in manager.storage_provider.values
+
+
+@pytest.mark.asyncio
+async def test_scoped_provider_is_not_touched_after_generation_is_fenced(manager):
+ object_key = await manager.save_scoped(
+ _context(WORKSPACE_A),
+ owner_type='upload',
+ owner='account:a',
+ key='document.pdf',
+ value=b'document-a',
+ )
+ manager.ap.workspace_service = SimpleNamespace(
+ get_execution_binding=AsyncMock(side_effect=ValueError('generation is stale'))
+ )
+ manager.storage_provider.load = AsyncMock(side_effect=AssertionError('provider must not be called'))
+
+ with pytest.raises(WorkspaceRequiredError, match='execution scope is unavailable'):
+ await manager.load_scoped_object_key(
+ _context(WORKSPACE_A),
+ object_key,
+ expected_owner_type='upload',
+ )
+ manager.storage_provider.load.assert_not_awaited()
+
+
+@pytest.mark.asyncio
+async def test_scoped_provider_rejects_incomplete_execution_binding(manager):
+ manager.ap.workspace_service = SimpleNamespace(
+ get_execution_binding=AsyncMock(return_value=SimpleNamespace(instance_uuid='instance'))
+ )
+ manager.storage_provider.save = AsyncMock(side_effect=AssertionError('provider must not be called'))
+
+ with pytest.raises(WorkspaceRequiredError, match='execution scope is unavailable'):
+ await manager.save_scoped(
+ _context(WORKSPACE_A),
+ owner_type='upload',
+ owner='account:a',
+ key='document.pdf',
+ value=b'document-a',
+ )
+ manager.storage_provider.save.assert_not_awaited()
diff --git a/tests/unit_tests/telemetry/test_heartbeat.py b/tests/unit_tests/telemetry/test_heartbeat.py
index 18d61f2d8..f13aa9a39 100644
--- a/tests/unit_tests/telemetry/test_heartbeat.py
+++ b/tests/unit_tests/telemetry/test_heartbeat.py
@@ -49,6 +49,7 @@ def make_app():
# skills
ap.skill_mgr = Mock()
ap.skill_mgr.skills = {'a': {}, 'b': {}, 'c': {}}
+ ap.skill_mgr.total_cached_skill_count.return_value = 3
return ap
diff --git a/tests/unit_tests/test_preproc.py b/tests/unit_tests/test_preproc.py
index 8f05277cb..d81f67af5 100644
--- a/tests/unit_tests/test_preproc.py
+++ b/tests/unit_tests/test_preproc.py
@@ -16,10 +16,21 @@ from langbot_plugin.api.entities.builtin.provider.message import Message
from langbot_plugin.api.entities.builtin.provider.prompt import Prompt
from langbot_plugin.api.entities.builtin.provider.session import Conversation, LauncherTypes, Session
+from langbot.pkg.api.http.context import ExecutionContext
+
+
+_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+ bot_uuid='bot-1',
+ pipeline_uuid='pipe-1',
+)
+
def _make_query() -> Query:
message_chain = MessageChain([Plain(text='create a skill')])
- return Query(
+ query = Query(
query_id=1,
launcher_type=LauncherTypes.PERSON,
launcher_id='launcher-1',
@@ -45,6 +56,8 @@ def _make_query() -> Query:
},
variables={},
)
+ object.__setattr__(query, '_execution_context', _CONTEXT)
+ return query
def _make_conversation() -> Conversation:
@@ -84,6 +97,8 @@ def _make_app(*, skill_service) -> SimpleNamespace:
get_pipeline=AsyncMock(return_value={'extensions_preferences': {'enable_all_skills': True}})
),
skill_mgr=SimpleNamespace(
+ ensure_loaded=AsyncMock(),
+ get_skills=Mock(return_value={}),
build_skill_aware_prompt_addition=Mock(return_value=''),
skills={},
),
@@ -119,6 +134,7 @@ async def test_preproc_enables_skill_authoring_tools_when_skill_service_availabl
assert result.result_type == entities_module.ResultType.CONTINUE
app.tool_mgr.get_all_tools.assert_awaited_once_with(
+ _CONTEXT,
None,
None,
include_skill_authoring=True,
@@ -137,6 +153,7 @@ async def test_preproc_disables_skill_authoring_tools_when_skill_service_missing
assert result.result_type == entities_module.ResultType.CONTINUE
app.tool_mgr.get_all_tools.assert_awaited_once_with(
+ _CONTEXT,
None,
None,
include_skill_authoring=False,
@@ -157,6 +174,7 @@ async def test_preproc_disables_mcp_resource_tools_when_agent_reading_is_disable
assert result.result_type == entities_module.ResultType.CONTINUE
app.tool_mgr.get_all_tools.assert_awaited_once_with(
+ _CONTEXT,
None,
None,
include_skill_authoring=True,
@@ -179,7 +197,10 @@ async def test_preproc_injects_skill_index_into_system_prompt():
result = await stage_process_capture(preproc_module, app, query)
assert result.result_type == entities_module.ResultType.CONTINUE
- app.skill_mgr.build_skill_aware_prompt_addition.assert_called_once_with(bound_skills=None)
+ app.skill_mgr.build_skill_aware_prompt_addition.assert_called_once_with(
+ _CONTEXT,
+ bound_skills=None,
+ )
head = query.prompt.messages[0]
assert head.role == 'system'
assert head.content.endswith(addendum)
@@ -206,7 +227,10 @@ async def test_preproc_respects_pipeline_bound_skills_subset():
result = await stage_process_capture(preproc_module, app, query)
assert result.result_type == entities_module.ResultType.CONTINUE
- app.skill_mgr.build_skill_aware_prompt_addition.assert_called_once_with(bound_skills=['only-this'])
+ app.skill_mgr.build_skill_aware_prompt_addition.assert_called_once_with(
+ _CONTEXT,
+ bound_skills=['only-this'],
+ )
assert query.variables.get('_pipeline_bound_skills') == ['only-this']
diff --git a/tests/unit_tests/test_skill_service.py b/tests/unit_tests/test_skill_service.py
index 6fd7d64f2..5e1d465a5 100644
--- a/tests/unit_tests/test_skill_service.py
+++ b/tests/unit_tests/test_skill_service.py
@@ -3,9 +3,23 @@ from unittest.mock import AsyncMock
import pytest
+from langbot.pkg.api.http.context import ExecutionContext
from langbot.pkg.api.http.service.skill import SkillService
+_CONTEXT = ExecutionContext(
+ instance_uuid='instance-a',
+ workspace_uuid='workspace-a',
+ placement_generation=1,
+)
+
+
+def _workspace_service():
+ return SimpleNamespace(
+ get_execution_binding=AsyncMock(return_value=SimpleNamespace(instance_uuid=_CONTEXT.instance_uuid))
+ )
+
+
class TestRequireBoxForWrite:
"""Box is the only source of truth for skills — there is no local
filesystem fallback. Every write and (most) read methods refuse cleanly
@@ -14,6 +28,7 @@ class TestRequireBoxForWrite:
def _ap_with_disabled_box(self):
return SimpleNamespace(
skill_mgr=SimpleNamespace(reload_skills=AsyncMock()),
+ workspace_service=_workspace_service(),
box_service=SimpleNamespace(
available=False,
enabled=False,
@@ -24,6 +39,7 @@ class TestRequireBoxForWrite:
def _ap_with_failed_box(self):
return SimpleNamespace(
skill_mgr=SimpleNamespace(reload_skills=AsyncMock()),
+ workspace_service=_workspace_service(),
box_service=SimpleNamespace(
available=False,
enabled=True,
@@ -35,55 +51,67 @@ class TestRequireBoxForWrite:
async def test_create_skill_refused_when_box_disabled(self):
service = SkillService(self._ap_with_disabled_box())
with pytest.raises(ValueError, match='disabled in config'):
- await service.create_skill({'name': 'x'})
+ await service.create_skill(_CONTEXT, {'name': 'x'})
@pytest.mark.asyncio
async def test_create_skill_refused_when_box_failed(self):
service = SkillService(self._ap_with_failed_box())
with pytest.raises(ValueError, match='docker daemon not running'):
- await service.create_skill({'name': 'x'})
+ await service.create_skill(_CONTEXT, {'name': 'x'})
@pytest.mark.asyncio
async def test_update_skill_refused_when_box_disabled(self):
service = SkillService(self._ap_with_disabled_box())
with pytest.raises(ValueError, match='Editing a skill requires the Box runtime'):
- await service.update_skill('x', {})
+ await service.update_skill(_CONTEXT, 'x', {})
@pytest.mark.asyncio
async def test_write_skill_file_refused_when_box_disabled(self):
service = SkillService(self._ap_with_disabled_box())
with pytest.raises(ValueError, match='Editing skill files requires the Box runtime'):
- await service.write_skill_file('x', 'a.txt', 'hi')
+ await service.write_skill_file(_CONTEXT, 'x', 'a.txt', 'hi')
@pytest.mark.asyncio
async def test_install_from_github_refused_when_box_disabled(self):
service = SkillService(self._ap_with_disabled_box())
with pytest.raises(ValueError, match='Installing a skill from GitHub'):
- await service.install_from_github({'owner': 'o', 'repo': 'r', 'asset_url': 'https://example/x.zip'})
+ await service.install_from_github(
+ _CONTEXT,
+ {'owner': 'o', 'repo': 'r', 'asset_url': 'https://example/x.zip'},
+ )
@pytest.mark.asyncio
async def test_install_from_zip_upload_refused_when_box_disabled(self):
service = SkillService(self._ap_with_disabled_box())
with pytest.raises(ValueError, match='Installing a skill from upload'):
- await service.install_from_zip_upload(file_bytes=b'', filename='x.zip')
+ await service.install_from_zip_upload(
+ _CONTEXT,
+ file_bytes=b'',
+ filename='x.zip',
+ )
@pytest.mark.asyncio
async def test_create_skill_refused_when_box_service_missing_entirely(self):
"""No ap.box_service attribute at all (truly minimal setup):
Box is the only source of truth, so creation must still refuse."""
- service = SkillService(SimpleNamespace(skill_mgr=SimpleNamespace(reload_skills=AsyncMock())))
+ service = SkillService(
+ SimpleNamespace(
+ skill_mgr=SimpleNamespace(reload_skills=AsyncMock()),
+ workspace_service=_workspace_service(),
+ )
+ )
with pytest.raises(ValueError, match='not initialised'):
- await service.create_skill({'name': 'x'})
+ await service.create_skill(_CONTEXT, {'name': 'x'})
@pytest.mark.asyncio
async def test_list_skills_returns_empty_when_box_unavailable(self):
"""list_skills should render an empty surface (not crash) so the
skills page can show a banner instead of a broken state."""
service = SkillService(self._ap_with_disabled_box())
- assert await service.list_skills() == []
+ assert await service.list_skills(_CONTEXT) == []
@pytest.mark.asyncio
async def test_read_skill_file_refused_when_box_unavailable(self):
service = SkillService(self._ap_with_disabled_box())
with pytest.raises(ValueError, match='Reading a skill file'):
- await service.read_skill_file('x', 'a.txt')
+ await service.read_skill_file(_CONTEXT, 'x', 'a.txt')
diff --git a/tests/unit_tests/utils/test_httpclient.py b/tests/unit_tests/utils/test_httpclient.py
index 0a102969a..da110afe2 100644
--- a/tests/unit_tests/utils/test_httpclient.py
+++ b/tests/unit_tests/utils/test_httpclient.py
@@ -25,6 +25,7 @@ class TestGetSession:
assert isinstance(session, aiohttp.ClientSession)
assert not session.closed
+ assert isinstance(session.cookie_jar, aiohttp.DummyCookieJar)
# Cleanup
await session.close()
@@ -144,3 +145,12 @@ class TestSessionPoolIntegration:
assert session is session2
await httpclient.close_all()
+
+ async def test_shared_session_does_not_persist_cookies(self):
+ """Shared transport pooling never creates cross-Workspace cookie state."""
+ session = httpclient.get_session()
+
+ session.cookie_jar.update_cookies({'workspace_session': 'secret'})
+
+ assert list(session.cookie_jar) == []
+ await httpclient.close_all()
diff --git a/tests/unit_tests/workspace/__init__.py b/tests/unit_tests/workspace/__init__.py
new file mode 100644
index 000000000..e69de29bb
diff --git a/tests/unit_tests/workspace/test_workspace_collaboration.py b/tests/unit_tests/workspace/test_workspace_collaboration.py
new file mode 100644
index 000000000..8020e7c41
--- /dev/null
+++ b/tests/unit_tests/workspace/test_workspace_collaboration.py
@@ -0,0 +1,310 @@
+from __future__ import annotations
+
+import asyncio
+import uuid
+from types import SimpleNamespace
+
+import pytest
+import sqlalchemy
+from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
+
+from langbot.pkg.entity.persistence.base import Base
+from langbot.pkg.entity.persistence.user import User
+from langbot.pkg.entity.persistence.workspace import (
+ Workspace,
+ WorkspaceExecutionState,
+ WorkspaceInvitation,
+ WorkspaceMembership,
+)
+from langbot.pkg.workspace.collaboration import (
+ InvitationEmailMismatchError,
+ InvitationRoleError,
+ InvitationUsedError,
+ LastOwnerError,
+ MembershipPermissionError,
+ WorkspaceCollaborationService,
+)
+from langbot.pkg.workspace.errors import WorkspaceNotFoundError
+from langbot.pkg.workspace.service import WorkspaceService
+from langbot.pkg.workspace.policy import CloudWorkspacePolicy
+
+
+pytestmark = pytest.mark.asyncio
+
+
+@pytest.fixture
+async def collaboration_context(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "workspace-collaboration.db"}')
+ async with engine.begin() as conn:
+ await conn.run_sync(Base.metadata.create_all)
+
+ persistence_mgr = SimpleNamespace(get_db_engine=lambda: engine)
+ application = SimpleNamespace(persistence_mgr=persistence_mgr)
+ workspace_service = WorkspaceService(application, instance_uuid='instance-collaboration-test')
+ service = WorkspaceCollaborationService(application, workspace_service)
+ session_factory = async_sessionmaker(engine, expire_on_commit=False)
+
+ async with session_factory() as session:
+ async with session.begin():
+ owner = User(
+ uuid=str(uuid.uuid4()),
+ user='owner@example.com',
+ normalized_email='owner@example.com',
+ password='owner-hash',
+ account_type='local',
+ )
+ session.add(owner)
+ await session.flush()
+ workspace, owner_membership = await workspace_service.bootstrap_local_account(
+ owner.uuid,
+ session=session,
+ )
+
+ yield service, workspace_service, session_factory, owner, workspace, owner_membership
+ await engine.dispose()
+
+
+async def _add_account(session_factory, email: str) -> User:
+ async with session_factory() as session:
+ async with session.begin():
+ account = User(
+ uuid=str(uuid.uuid4()),
+ user=email,
+ normalized_email=email.strip().casefold(),
+ password='member-hash',
+ account_type='local',
+ )
+ session.add(account)
+ await session.flush()
+ return account
+
+
+async def test_invitation_secret_is_hashed_and_acceptance_is_one_time(collaboration_context):
+ service, _, session_factory, _, workspace, owner_membership = collaboration_context
+ created = await service.create_invitation(
+ workspace.uuid,
+ owner_membership,
+ 'member@example.com',
+ 'developer',
+ )
+
+ assert created.token.startswith('lbi_')
+ assert created.invitation.token_hash != created.token
+ assert created.token not in created.invitation.token_hash
+
+ account = await _add_account(session_factory, 'MEMBER@example.com')
+ membership = await service.accept_invitation(created.token, account.uuid)
+ assert membership.workspace_uuid == workspace.uuid
+ assert membership.role == 'developer'
+
+ with pytest.raises(InvitationUsedError):
+ await service.accept_invitation(created.token, account.uuid)
+
+ async with session_factory() as session:
+ persisted = await session.get(WorkspaceInvitation, created.invitation.uuid)
+ assert persisted is not None
+ assert persisted.status == 'accepted'
+ assert not hasattr(persisted, 'token')
+
+
+async def test_concurrent_invitation_acceptance_creates_one_membership(collaboration_context):
+ service, _, session_factory, _, workspace, owner_membership = collaboration_context
+ created = await service.create_invitation(
+ workspace.uuid,
+ owner_membership,
+ 'race@example.com',
+ 'viewer',
+ )
+ account = await _add_account(session_factory, 'race@example.com')
+
+ results = await asyncio.gather(
+ service.accept_invitation(created.token, account.uuid),
+ service.accept_invitation(created.token, account.uuid),
+ return_exceptions=True,
+ )
+
+ assert sum(isinstance(result, WorkspaceMembership) for result in results) == 1
+ assert sum(isinstance(result, InvitationUsedError) for result in results) == 1
+ async with session_factory() as session:
+ count = await session.scalar(
+ sqlalchemy.select(sqlalchemy.func.count())
+ .select_from(WorkspaceMembership)
+ .where(
+ WorkspaceMembership.workspace_uuid == workspace.uuid,
+ WorkspaceMembership.account_uuid == account.uuid,
+ )
+ )
+ assert count == 1
+
+
+async def test_invitation_rejects_owner_role_and_email_mismatch(collaboration_context):
+ service, _, session_factory, _, workspace, owner_membership = collaboration_context
+ with pytest.raises(InvitationRoleError):
+ await service.create_invitation(
+ workspace.uuid,
+ owner_membership,
+ 'member@example.com',
+ 'owner',
+ )
+
+ created = await service.create_invitation(
+ workspace.uuid,
+ owner_membership,
+ 'expected@example.com',
+ 'viewer',
+ )
+ wrong_account = await _add_account(session_factory, 'wrong@example.com')
+ with pytest.raises(InvitationEmailMismatchError):
+ await service.accept_invitation(created.token, wrong_account.uuid)
+
+
+async def test_last_owner_cannot_be_demoted(collaboration_context):
+ service, _, session_factory, _, workspace, owner_membership = collaboration_context
+ created = await service.create_invitation(
+ workspace.uuid,
+ owner_membership,
+ 'second@example.com',
+ 'admin',
+ )
+ second = await _add_account(session_factory, 'second@example.com')
+ second_membership = await service.accept_invitation(created.token, second.uuid)
+
+ with pytest.raises(LastOwnerError):
+ await service.update_member_role(
+ workspace.uuid,
+ owner_membership.account_uuid,
+ 'admin',
+ owner_membership,
+ )
+
+ with pytest.raises(MembershipPermissionError):
+ await service.update_member_role(
+ workspace.uuid,
+ owner_membership.account_uuid,
+ 'viewer',
+ second_membership,
+ )
+
+ promoted = await service.update_member_role(
+ workspace.uuid,
+ second.uuid,
+ 'owner',
+ owner_membership,
+ )
+ assert promoted.role == 'owner'
+ demoted = await service.update_member_role(
+ workspace.uuid,
+ owner_membership.account_uuid,
+ 'admin',
+ owner_membership,
+ )
+ assert demoted.role == 'admin'
+
+
+async def test_workspace_selector_requires_membership(collaboration_context):
+ service, _, session_factory, _, workspace, _ = collaboration_context
+ outsider = await _add_account(session_factory, 'outsider@example.com')
+
+ with pytest.raises(WorkspaceNotFoundError):
+ await service.resolve_account_workspace(outsider.uuid, workspace.uuid)
+
+
+async def test_cloud_directory_requires_explicit_selector_and_rejects_local_mutation(
+ collaboration_context,
+):
+ _service, workspace_service, session_factory, owner, _workspace, _membership = collaboration_context
+ cloud_workspace_uuid = str(uuid.uuid4())
+ cloud_membership = WorkspaceMembership(
+ uuid=str(uuid.uuid4()),
+ workspace_uuid=cloud_workspace_uuid,
+ account_uuid=owner.uuid,
+ role='owner',
+ status='active',
+ projection_revision=4,
+ )
+ async with session_factory.begin() as session:
+ session.add(
+ Workspace(
+ uuid=cloud_workspace_uuid,
+ instance_uuid='instance-collaboration-test',
+ name='Projected Workspace',
+ slug='projected-workspace',
+ source='cloud_projection',
+ projection_revision=4,
+ )
+ )
+ session.add(
+ WorkspaceExecutionState(
+ workspace_uuid=cloud_workspace_uuid,
+ instance_uuid='instance-collaboration-test',
+ active_generation=8,
+ state='active',
+ write_fenced=False,
+ source='cloud',
+ desired_state_revision=8,
+ )
+ )
+ session.add(cloud_membership)
+
+ cloud_service = WorkspaceCollaborationService(
+ workspace_service.ap,
+ WorkspaceService(
+ workspace_service.ap,
+ policy=CloudWorkspacePolicy(),
+ instance_uuid='instance-collaboration-test',
+ ),
+ )
+ with pytest.raises(WorkspaceNotFoundError):
+ await cloud_service.resolve_account_workspace(owner.uuid, None)
+
+ access = await cloud_service.resolve_account_workspace(owner.uuid, cloud_workspace_uuid)
+ assert access.workspace.uuid == cloud_workspace_uuid
+ assert access.execution.placement_generation == 8
+
+ with pytest.raises(MembershipPermissionError, match='control plane'):
+ await cloud_service.create_invitation(
+ cloud_workspace_uuid,
+ cloud_membership,
+ 'member@example.com',
+ 'viewer',
+ )
+
+
+async def test_invitation_lock_registry_does_not_retain_sequential_tokens(collaboration_context):
+ service, *_ = collaboration_context
+
+ for index in range(100):
+ async with service._invitation_lock(f'token-{index}'):
+ pass
+
+ assert service._invitation_locks == {}
+
+
+async def test_invitation_lock_registry_keeps_waiters_serialized(collaboration_context):
+ service, *_ = collaboration_context
+ first_entered = asyncio.Event()
+ release_first = asyncio.Event()
+ order: list[str] = []
+
+ async def first_worker() -> None:
+ async with service._invitation_lock('same-token'):
+ order.append('first')
+ first_entered.set()
+ await release_first.wait()
+
+ async def second_worker() -> None:
+ await first_entered.wait()
+ async with service._invitation_lock('same-token'):
+ order.append('second')
+
+ first_task = asyncio.create_task(first_worker())
+ second_task = asyncio.create_task(second_worker())
+ await first_entered.wait()
+ await asyncio.sleep(0)
+
+ assert order == ['first']
+ release_first.set()
+ await asyncio.gather(first_task, second_task)
+
+ assert order == ['first', 'second']
+ assert service._invitation_locks == {}
diff --git a/tests/unit_tests/workspace/test_workspace_policy.py b/tests/unit_tests/workspace/test_workspace_policy.py
new file mode 100644
index 000000000..d2f747467
--- /dev/null
+++ b/tests/unit_tests/workspace/test_workspace_policy.py
@@ -0,0 +1,19 @@
+from langbot.pkg.workspace.policy import (
+ CloudWorkspacePolicy,
+ SingleWorkspacePolicy,
+ open_core_workspace_policy,
+)
+
+
+def test_open_source_policy_is_single_workspace() -> None:
+ policy = open_core_workspace_policy()
+
+ assert policy.workspace_limit == 1
+ assert policy.multi_workspace_enabled is False
+
+
+def test_cloud_policy_is_not_exported_as_the_default_policy() -> None:
+ # CloudWorkspacePolicy remains a contract for the future verified, closed
+ # bootstrap. Constructing it explicitly must not change the OSS default.
+ assert CloudWorkspacePolicy().multi_workspace_enabled is True
+ assert isinstance(open_core_workspace_policy(), SingleWorkspacePolicy)
diff --git a/tests/unit_tests/workspace/test_workspace_service.py b/tests/unit_tests/workspace/test_workspace_service.py
new file mode 100644
index 000000000..32f8875e8
--- /dev/null
+++ b/tests/unit_tests/workspace/test_workspace_service.py
@@ -0,0 +1,237 @@
+from __future__ import annotations
+
+import uuid
+from types import SimpleNamespace
+
+import pytest
+import sqlalchemy
+from sqlalchemy.ext.asyncio import async_sessionmaker, create_async_engine
+
+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,
+ WorkspaceMembership,
+ WorkspaceSource,
+)
+from langbot.pkg.workspace import (
+ WorkspaceExecutionUnavailableError,
+ WorkspaceGenerationMismatchError,
+ WorkspaceInvariantError,
+ WorkspaceLimitExceededError,
+ WorkspaceOwnerAlreadyExistsError,
+ WorkspaceService,
+)
+from langbot.pkg.workspace.policy import CloudWorkspacePolicy
+
+
+pytestmark = pytest.mark.asyncio
+
+
+@pytest.fixture
+async def workspace_test_context(tmp_path):
+ engine = create_async_engine(f'sqlite+aiosqlite:///{tmp_path / "workspace-service.db"}')
+ async with engine.begin() as conn:
+ await conn.run_sync(Base.metadata.create_all)
+
+ persistence_mgr = SimpleNamespace(get_db_engine=lambda: engine)
+ application = SimpleNamespace(persistence_mgr=persistence_mgr)
+ session_factory = async_sessionmaker(engine, expire_on_commit=False)
+ service = WorkspaceService(application, instance_uuid='instance_service_test')
+ yield service, session_factory
+ await engine.dispose()
+
+
+async def _insert_account(session, email: str) -> str:
+ account_uuid = str(uuid.uuid4())
+ session.add(
+ User(
+ uuid=account_uuid,
+ user=email,
+ normalized_email=email.strip().casefold(),
+ password='hashed-password',
+ account_type='local',
+ )
+ )
+ await session.flush()
+ return account_uuid
+
+
+async def test_account_uuid_default_supports_existing_core_insert_path(workspace_test_context):
+ _service, session_factory = workspace_test_context
+
+ async with session_factory.begin() as session:
+ await session.execute(
+ sqlalchemy.insert(User).values(
+ user='core-insert@example.com',
+ normalized_email='core-insert@example.com',
+ password='hashed-password',
+ account_type='local',
+ )
+ )
+
+ async with session_factory() as session:
+ account = await session.scalar(sqlalchemy.select(User))
+ assert account is not None
+ uuid.UUID(account.uuid)
+ assert account.status == 'active'
+ assert account.source == 'local'
+
+
+async def test_bootstrap_local_account_uses_callers_transaction(workspace_test_context):
+ service, session_factory = workspace_test_context
+
+ async with session_factory() as session:
+ async with session.begin():
+ account_uuid = await _insert_account(session, 'owner@example.com')
+ workspace, membership = await service.bootstrap_local_account(account_uuid, session=session)
+ assert session.in_transaction()
+ assert workspace.created_by_account_uuid == account_uuid
+ assert membership.account_uuid == account_uuid
+ assert membership.role == 'owner'
+
+ async with session_factory() as session:
+ persisted_workspace = await session.scalar(sqlalchemy.select(Workspace))
+ persisted_membership = await session.scalar(sqlalchemy.select(WorkspaceMembership))
+ execution_state = await session.scalar(sqlalchemy.select(WorkspaceExecutionState))
+ assert persisted_workspace is not None
+ assert persisted_membership is not None
+ assert execution_state is not None
+ assert execution_state.workspace_uuid == persisted_workspace.uuid
+ assert execution_state.instance_uuid == 'instance_service_test'
+ assert execution_state.active_generation == 1
+ assert execution_state.write_fenced is False
+
+
+async def test_ensure_singleton_workspace_is_idempotent(workspace_test_context):
+ service, session_factory = workspace_test_context
+
+ first = await service.ensure_singleton_workspace()
+ second = await service.ensure_singleton_workspace()
+
+ assert second.uuid == first.uuid
+ async with session_factory() as session:
+ assert await session.scalar(sqlalchemy.select(sqlalchemy.func.count()).select_from(Workspace)) == 1
+ assert (
+ await session.scalar(sqlalchemy.select(sqlalchemy.func.count()).select_from(WorkspaceExecutionState)) == 1
+ )
+
+
+async def test_create_local_workspace_rejects_second_workspace(workspace_test_context):
+ service, session_factory = workspace_test_context
+
+ await service.create_local_workspace(name='First', slug='first')
+ with pytest.raises(WorkspaceLimitExceededError) as exc_info:
+ await service.create_local_workspace(name='Second', slug='second')
+
+ assert exc_info.value.code == 'edition_limit'
+ async with session_factory() as session:
+ assert await session.scalar(sqlalchemy.select(sqlalchemy.func.count()).select_from(Workspace)) == 1
+
+
+async def test_initial_owner_cannot_be_claimed_by_another_account(workspace_test_context):
+ service, session_factory = workspace_test_context
+
+ async with session_factory() as session:
+ async with session.begin():
+ first_account_uuid = await _insert_account(session, 'first@example.com')
+ second_account_uuid = await _insert_account(session, 'second@example.com')
+ await service.bootstrap_local_account(first_account_uuid, session=session)
+
+ with pytest.raises(WorkspaceOwnerAlreadyExistsError):
+ await service.claim_initial_owner(second_account_uuid)
+
+ async with session_factory() as session:
+ owners = (
+ await session.scalars(sqlalchemy.select(WorkspaceMembership).where(WorkspaceMembership.role == 'owner'))
+ ).all()
+ assert len(owners) == 1
+ assert owners[0].account_uuid == first_account_uuid
+
+
+async def test_execution_binding_returns_persisted_generation(workspace_test_context):
+ service, session_factory = workspace_test_context
+ workspace = await service.ensure_singleton_workspace()
+
+ async with session_factory.begin() as session:
+ execution_state = await session.get(WorkspaceExecutionState, workspace.uuid)
+ execution_state.active_generation = 7
+
+ binding = await service.get_local_execution_binding(
+ workspace.uuid,
+ expected_generation=7,
+ )
+
+ assert binding.instance_uuid == 'instance_service_test'
+ assert binding.workspace_uuid == workspace.uuid
+ assert binding.placement_generation == 7
+ assert binding.state == 'active'
+ assert binding.write_fenced is False
+
+ with pytest.raises(WorkspaceGenerationMismatchError):
+ await service.get_local_execution_binding(workspace.uuid, expected_generation=6)
+
+
+async def test_execution_binding_fails_closed_when_fenced(workspace_test_context):
+ service, session_factory = workspace_test_context
+ workspace = await service.ensure_singleton_workspace()
+
+ async with session_factory.begin() as session:
+ execution_state = await session.get(WorkspaceExecutionState, workspace.uuid)
+ execution_state.write_fenced = True
+
+ with pytest.raises(WorkspaceExecutionUnavailableError):
+ await service.get_local_execution_context(workspace.uuid)
+
+
+async def test_cloud_projection_requires_explicit_general_binding(workspace_test_context):
+ service, session_factory = workspace_test_context
+ await service.ensure_singleton_workspace()
+ cloud_workspace_uuid = str(uuid.uuid4())
+
+ async with session_factory.begin() as session:
+ session.add(
+ Workspace(
+ uuid=cloud_workspace_uuid,
+ instance_uuid='instance_service_test',
+ name='Cloud Workspace',
+ slug='cloud-workspace',
+ source=WorkspaceSource.CLOUD_PROJECTION.value,
+ )
+ )
+ session.add(
+ WorkspaceExecutionState(
+ workspace_uuid=cloud_workspace_uuid,
+ instance_uuid='instance_service_test',
+ active_generation=9,
+ state='active',
+ write_fenced=False,
+ source=WorkspaceExecutionSource.CLOUD.value,
+ )
+ )
+
+ binding = await service.get_execution_binding(
+ cloud_workspace_uuid,
+ expected_generation=9,
+ )
+ assert binding.workspace_uuid == cloud_workspace_uuid
+ assert binding.placement_generation == 9
+
+ with pytest.raises(WorkspaceInvariantError):
+ await service.get_local_execution_binding(cloud_workspace_uuid)
+
+
+async def test_cloud_policy_never_creates_or_guesses_a_workspace(workspace_test_context):
+ service, _session_factory = workspace_test_context
+ cloud_service = WorkspaceService(
+ service.ap,
+ policy=CloudWorkspacePolicy(),
+ instance_uuid='instance_service_test',
+ )
+
+ with pytest.raises(WorkspaceLimitExceededError):
+ await cloud_service.create_local_workspace(name='Forbidden', slug='forbidden')
+
+ assert cloud_service.policy.multi_workspace_enabled is True
diff --git a/tests/utils/import_isolation.py b/tests/utils/import_isolation.py
index 9f2b3c583..d9d2b3dc5 100644
--- a/tests/utils/import_isolation.py
+++ b/tests/utils/import_isolation.py
@@ -24,6 +24,9 @@ from typing import Generator
from unittest.mock import MagicMock
+_MISSING = object()
+
+
class MockLifecycleControlScope(enum.Enum):
"""Mock enum for breaking circular import in core.entities."""
@@ -74,19 +77,29 @@ def isolated_sys_modules(
if name in sys.modules:
saved[name] = sys.modules[name]
- # Save original package attributes that will be updated
- saved_attrs: dict[str, tuple[str, object]] = {}
- for mock_name, (pkg_name, attr_name) in _PACKAGE_ATTRIBUTE_UPDATES.items():
- if mock_name in mocks and pkg_name in sys.modules:
- pkg = sys.modules[pkg_name]
- if hasattr(pkg, attr_name):
- saved_attrs[mock_name] = (pkg_name, getattr(pkg, attr_name))
+ # Importing a submodule also mutates its parent package attribute. Preserve
+ # that state for every mocked or cleared module, otherwise restoring only
+ # sys.modules leaves stale class/enum identities attached to the package.
+ saved_attrs: dict[str, tuple[str, str, object]] = {}
+ for name in touched:
+ pkg_name, separator, attr_name = name.rpartition('.')
+ if separator and pkg_name in sys.modules:
+ saved_attrs[name] = (
+ pkg_name,
+ attr_name,
+ getattr(sys.modules[pkg_name], attr_name, _MISSING),
+ )
try:
# Clear modules first (force re-import)
for name in clear:
if name not in mocks: # Don't clear if we're mocking it
sys.modules.pop(name, None)
+ saved_attr = saved_attrs.get(name)
+ if saved_attr is not None:
+ pkg_name, attr_name, _ = saved_attr
+ if pkg_name in sys.modules and hasattr(sys.modules[pkg_name], attr_name):
+ delattr(sys.modules[pkg_name], attr_name)
# Apply mocks
for name, module in mocks.items():
@@ -110,10 +123,16 @@ def isolated_sys_modules(
# Wasn't in sys.modules originally, remove it
sys.modules.pop(name, None)
- # Restore package attributes
- for mock_name, (pkg_name, original_value) in saved_attrs.items():
- if pkg_name in sys.modules:
- setattr(sys.modules[pkg_name], _PACKAGE_ATTRIBUTE_UPDATES[mock_name][1], original_value)
+ # Restore package attributes (or remove ones that did not previously
+ # exist), keeping package lookup consistent with restored sys.modules.
+ for pkg_name, attr_name, original_value in saved_attrs.values():
+ if pkg_name not in sys.modules:
+ continue
+ if original_value is _MISSING:
+ if hasattr(sys.modules[pkg_name], attr_name):
+ delattr(sys.modules[pkg_name], attr_name)
+ else:
+ setattr(sys.modules[pkg_name], attr_name, original_value)
def make_pipeline_handler_import_mocks() -> dict[str, MagicMock]:
diff --git a/uv.lock b/uv.lock
index 6ddc22e3e..d6f49aa86 100644
--- a/uv.lock
+++ b/uv.lock
@@ -2124,7 +2124,7 @@ requires-dist = [
{ name = "ebooklib", specifier = ">=0.18" },
{ name = "gewechat-client", specifier = ">=0.1.5" },
{ name = "html2text", specifier = ">=2024.2.26" },
- { name = "langbot-plugin", specifier = "==0.4.13" },
+ { name = "langbot-plugin", git = "https://github.com/langbot-app/langbot-plugin-sdk.git?rev=a1544b6b38a37ba72e3284f2836618144f0742c1" },
{ name = "langchain", specifier = ">=1.3.9" },
{ name = "langchain-core", specifier = ">=1.3.3" },
{ name = "langchain-text-splitters", specifier = ">=1.1.2" },
@@ -2189,8 +2189,8 @@ dev = [
[[package]]
name = "langbot-plugin"
-version = "0.4.13"
-source = { registry = "https://pypi.org/simple" }
+version = "0.4.15"
+source = { git = "https://github.com/langbot-app/langbot-plugin-sdk.git?rev=a1544b6b38a37ba72e3284f2836618144f0742c1#a1544b6b38a37ba72e3284f2836618144f0742c1" }
dependencies = [
{ name = "aiofiles" },
{ name = "aiohttp" },
@@ -2210,10 +2210,6 @@ dependencies = [
{ name = "watchdog" },
{ name = "websockets" },
]
-sdist = { url = "https://files.pythonhosted.org/packages/40/a6/1eaf77c3b81e9de3390c504c5f627dc41f43bff6df9aff0e1e31d796b6f0/langbot_plugin-0.4.13.tar.gz", hash = "sha256:f936340e67679c21f1e7e7f1447339f31a0a2c965db060ecfbd9d0c51bb0d6fe", size = 334887, upload-time = "2026-07-04T05:38:59.942Z" }
-wheels = [
- { url = "https://files.pythonhosted.org/packages/e9/bf/fc9671a7afbd933440c38403c84d918c1022fdeed16e22a6ab3b2aec83ff/langbot_plugin-0.4.13-py3-none-any.whl", hash = "sha256:9d45ebc7a7ee0413d6db9baa009fcbf0ad07e2e1753a6f0a27f37b8b665cd1ee", size = 221884, upload-time = "2026-07-04T05:38:58.525Z" },
-]
[[package]]
name = "langchain"
diff --git a/web/src/app/auth/space/callback/page.tsx b/web/src/app/auth/space/callback/page.tsx
index 8711cbd6f..19a9a710c 100644
--- a/web/src/app/auth/space/callback/page.tsx
+++ b/web/src/app/auth/space/callback/page.tsx
@@ -1,6 +1,10 @@
import { useEffect, useState, useCallback, Suspense, useRef } from 'react';
import { useNavigate, useSearchParams } from 'react-router-dom';
import { httpClient } from '@/app/infra/http/HttpClient';
+import {
+ beginAuthenticatedSession,
+ bootstrapWorkspaceSession,
+} from '@/app/infra/http';
import { toast } from 'sonner';
import { useTranslation } from 'react-i18next';
import {
@@ -32,19 +36,21 @@ const pendingSpaceOAuthLogins = new Map<
function getOrCreateSpaceOAuthLoginPromise(
authCode: string,
+ state: string,
): Promise {
- const pendingRequest = pendingSpaceOAuthLogins.get(authCode);
+ const requestKey = `${authCode}:${state}`;
+ const pendingRequest = pendingSpaceOAuthLogins.get(requestKey);
if (pendingRequest) {
return pendingRequest;
}
const requestPromise = httpClient
- .exchangeSpaceOAuthCode(authCode)
+ .exchangeSpaceOAuthCode(authCode, state)
.finally(() => {
- pendingSpaceOAuthLogins.delete(authCode);
+ pendingSpaceOAuthLogins.delete(requestKey);
});
- pendingSpaceOAuthLogins.set(authCode, requestPromise);
+ pendingSpaceOAuthLogins.set(requestKey, requestPromise);
return requestPromise;
}
@@ -64,23 +70,31 @@ function SpaceOAuthCallbackContent() {
const [localEmail, setLocalEmail] = useState('');
const handleOAuthCallback = useCallback(
- async (authCode: string) => {
+ async (authCode: string, state: string) => {
try {
- const response = await getOrCreateSpaceOAuthLoginPromise(authCode);
+ const response = await getOrCreateSpaceOAuthLoginPromise(
+ authCode,
+ state,
+ );
if (!isMountedRef.current) {
return;
}
- localStorage.setItem('token', response.token);
- if (response.user) {
- localStorage.setItem('userEmail', response.user);
+ beginAuthenticatedSession(response.token, response.user);
+ const workspaceResult = await bootstrapWorkspaceSession();
+ if (workspaceResult.status === 'unavailable') {
+ throw new Error('No Workspace is available for this Account');
}
setStatus('success');
toast.success(t('common.spaceLoginSuccess'));
// If wizard state exists, redirect back to wizard instead of home
const wizardState = localStorage.getItem('langbot_wizard_state');
- const redirectTo = wizardState ? '/wizard' : '/home';
+ const destination = wizardState ? '/wizard' : '/home';
+ const redirectTo =
+ workspaceResult.status === 'selection-required'
+ ? `/workspaces/select?returnTo=${encodeURIComponent(destination)}`
+ : destination;
setTimeout(() => {
navigate(redirectTo);
}, 1000);
@@ -113,14 +127,19 @@ function SpaceOAuthCallbackContent() {
return;
}
- localStorage.setItem('token', response.token);
- if (response.user) {
- localStorage.setItem('userEmail', response.user);
+ beginAuthenticatedSession(response.token, response.user);
+ const workspaceResult = await bootstrapWorkspaceSession();
+ if (workspaceResult.status === 'unavailable') {
+ throw new Error('No Workspace is available for this Account');
}
setStatus('success');
toast.success(t('account.bindSpaceSuccess'));
+ const redirectTo =
+ workspaceResult.status === 'selection-required'
+ ? '/workspaces/select?returnTo=%2Fhome'
+ : '/home';
setTimeout(() => {
- navigate('/home');
+ navigate(redirectTo);
}, 1000);
} catch (err) {
if (!isMountedRef.current) {
@@ -181,8 +200,12 @@ function SpaceOAuthCallbackContent() {
setLocalEmail(localStorage.getItem('userEmail') || '');
setStatus('confirm');
} else {
- // Normal login/register mode
- handleOAuthCallback(authCode);
+ if (!state) {
+ setStatus('error');
+ setErrorMessage(t('common.spaceLoginFailed'));
+ return;
+ }
+ handleOAuthCallback(authCode, state);
}
return () => {
isMountedRef.current = false;
diff --git a/web/src/app/home/bots/BotDetailContent.tsx b/web/src/app/home/bots/BotDetailContent.tsx
index 6c0a9ccc1..02ede4081 100644
--- a/web/src/app/home/bots/BotDetailContent.tsx
+++ b/web/src/app/home/bots/BotDetailContent.tsx
@@ -29,11 +29,17 @@ import { useTranslation } from 'react-i18next';
import { Settings, FileText, Users, RefreshCw, Trash2 } from 'lucide-react';
import { cn } from '@/lib/utils';
import { toast } from 'sonner';
+import { useCurrentWorkspace } from '@/app/infra/http';
export default function BotDetailContent({ id }: { id: string }) {
const isCreateMode = id === 'new';
const navigate = useNavigate();
const { t } = useTranslation();
+ const currentWorkspace = useCurrentWorkspace();
+ const canManage =
+ currentWorkspace?.permissions.includes('resource.manage') ?? false;
+ const canViewAudit =
+ currentWorkspace?.permissions.includes('audit.view') ?? false;
const { refreshBots, bots, setDetailEntityName } = useSidebarData();
// Set breadcrumb entity name
@@ -131,19 +137,23 @@ export default function BotDetailContent({ id }: { id: string }) {
{/* Header */}
{t('bots.createBot')}
-
+ {canManage && (
+
+ )}
{/* Content */}
@@ -164,6 +174,7 @@ export default function BotDetailContent({ id }: { id: string }) {
id="bot-enable-switch"
checked={botEnabled}
onCheckedChange={handleEnableToggle}
+ disabled={!canManage}
/>
- setShowDeleteConfirm(true)}
- className="shrink-0"
- >
-
- {t('common.delete')}
-
+ {canManage && (
+ setShowDeleteConfirm(true)}
+ className="shrink-0"
+ >
+
+ {t('common.delete')}
+
+ )}
@@ -185,14 +197,18 @@ export default function SkillDetailContent({ id }: { id: string }) {
)}
- handleImportedSkills([skillName])}
- onSkillUpdated={handleSkillUpdated}
- />
+
diff --git a/web/src/app/home/skills/page.tsx b/web/src/app/home/skills/page.tsx
index 9d50040db..b6529b995 100644
--- a/web/src/app/home/skills/page.tsx
+++ b/web/src/app/home/skills/page.tsx
@@ -7,9 +7,13 @@ import SkillForm from '@/app/home/skills/components/skill-form/SkillForm';
import { useSidebarData } from '@/app/home/components/home-sidebar/SidebarDataContext';
import { BoxUnavailableNotice } from '@/app/home/components/BoxUnavailableNotice';
import { useBoxStatus } from '@/app/infra/hooks/useBoxStatus';
+import { useCurrentWorkspace } from '@/app/infra/http';
export default function SkillsPage() {
const { t } = useTranslation();
+ const currentWorkspace = useCurrentWorkspace();
+ const canManage =
+ currentWorkspace?.permissions.includes('resource.manage') ?? false;
const navigate = useNavigate();
const [searchParams] = useSearchParams();
const detailId = searchParams.get('id');
@@ -61,7 +65,11 @@ export default function SkillsPage() {
{t('common.cancel')}
-
+
{t('common.save')}
@@ -72,13 +80,15 @@ export default function SkillsPage() {
)}
- {}}
- />
+
);
diff --git a/web/src/app/infra/entities/workspace.ts b/web/src/app/infra/entities/workspace.ts
new file mode 100644
index 000000000..9bd02f583
--- /dev/null
+++ b/web/src/app/infra/entities/workspace.ts
@@ -0,0 +1,51 @@
+export type WorkspaceRole =
+ | 'owner'
+ | 'admin'
+ | 'developer'
+ | 'operator'
+ | 'viewer';
+
+export interface Workspace {
+ uuid: string;
+ instance_uuid: string;
+ name: string;
+ slug: string;
+ type: 'personal' | 'team';
+ status: 'provisioning' | 'active' | 'suspended' | 'archived' | 'deleted';
+ source: 'local' | 'cloud_projection';
+}
+
+export interface WorkspaceMembership {
+ uuid: string;
+ workspace_uuid: string;
+ account_uuid: string;
+ email: string;
+ role: WorkspaceRole;
+ status: 'active' | 'disabled' | 'removed';
+ joined_at: string | null;
+ created_at: string;
+}
+
+export interface CurrentWorkspace {
+ workspace: Workspace;
+ membership: WorkspaceMembership;
+ permissions: string[];
+ placement_generation: number;
+}
+
+/** Account-scoped Workspace entry returned before a Workspace is selected. */
+export type WorkspaceBootstrapEntry = CurrentWorkspace;
+
+export interface WorkspaceBootstrapResponse {
+ workspaces: WorkspaceBootstrapEntry[];
+}
+
+export interface WorkspaceInvitation {
+ uuid: string;
+ workspace_uuid: string;
+ normalized_email: string;
+ role: Exclude;
+ status: 'pending' | 'accepted' | 'revoked' | 'expired';
+ expires_at: string;
+ created_at: string;
+}
diff --git a/web/src/app/infra/http/BackendClient.ts b/web/src/app/infra/http/BackendClient.ts
index 371b82942..efe70a6bd 100644
--- a/web/src/app/infra/http/BackendClient.ts
+++ b/web/src/app/infra/http/BackendClient.ts
@@ -1,4 +1,4 @@
-import { BaseHttpClient } from './BaseHttpClient';
+import { BaseHttpClient, type RequestConfig } from './BaseHttpClient';
import {
ApiRespProviderRequesters,
ApiRespProviderRequester,
@@ -61,6 +61,14 @@ import type { PluginLogEntry } from '@/app/infra/entities/plugin';
import type { I18nObject } from '@/app/infra/entities/common';
import { GetBotLogsRequest } from '@/app/infra/http/requestParam/bots/GetBotLogsRequest';
import { GetBotLogsResponse } from '@/app/infra/http/requestParam/bots/GetBotLogsResponse';
+import type {
+ CurrentWorkspace,
+ Workspace,
+ WorkspaceInvitation,
+ WorkspaceMembership,
+ WorkspaceBootstrapResponse,
+ WorkspaceRole,
+} from '@/app/infra/entities/workspace';
/**
* 后端服务客户端
@@ -1067,19 +1075,29 @@ export class BackendClient extends BaseHttpClient {
// ============ User API ============
public checkIfInited(): Promise<{ initialized: boolean }> {
- return this.get('/api/v1/user/init');
+ return this.get('/api/v1/user/init', undefined, { skipWorkspace: true });
}
public initUser(user: string, password: string): Promise