feat(tenancy): add Workspace multi-tenant foundation (#2353)

* Document multi-tenant workspace architecture

* Add OSS and commercial workspace boundaries

* docs: redesign multi-tenant workspace architecture

* feat(tenancy): implement workspace isolation

* docs(tenancy): record verification evidence

* docs(tenancy): revise single-instance SaaS topology

* docs(tenancy): refine architecture options

* docs: finalize cloud v2 multi-tenant decisions

* feat(tenancy): establish cloud isolation foundations

* feat(tenancy): harden shared cloud runtime boundaries

* docs(tenancy): record final isolation verification

* fix(tenancy): close isolation and permission gaps

* docs(tenancy): record final isolation verification

* feat(tenancy): connect cloud workspace control plane

* fix(build): install git for pinned SDK

* docs(cloud): update control plane verification

* chore: update multi-tenant SDK pin

* fix(cloud): skip legacy model sync during startup

* test(cloud): preserve minimal model manager fixtures

* fix(cloud): preserve authenticated account context

* fix(cloud): reuse authenticated account for user info

* feat(cloud): complete Workspace settings navigation

* test(web): cover Workspace dropdown menu

* feat(web): place workspace controls in sidebar

* refactor(web): streamline workspace controls

* style(web): format workspace layout test

* fix(cloud): surface runtime and workspace plan status

* fix(plugin): keep runtime identity stable across restarts

* fix(ui): widen and center workspace switcher

* fix(ui): hide roles from workspace switcher

* fix(ui): align workspace switcher with sidebar entries

* feat(workspace): add in-product collaboration and direct Cloud launch

* style: format collaboration changes

* fix(workspace): bind collaboration APIs to tenant UoW

* fix(cloud): preserve Core-owned collaboration state

* test(cloud): require Space identity for invite registration

* feat(cloud): complete secure invitation experience

* style(web): format invitation flows

* fix(cloud): recover box runtime without unscoped skill reload

* feat(oss): enforce invitation account and owner billing flows

* style: format OSS account service

* test(oss): cover invitation logout handoff

* fix(oss): resolve workspace owner in scoped session

* feat(cloud): harden multi-tenant runtime resources

* fix(cloud): bound runtime restart storms

* fix(cloud): eliminate periodic runtime CPU spikes

* fix(cloud): enforce instance capacity ceilings

* fix(cloud): scope public login capability discovery

* fix(cloud): bound tenant maintenance and monitoring work

* fix(runtime): bound tenant resource amplification

* fix(deps): pin green multi-tenant plugin SDK

* fix(cloud): handle unavailable skill capability

* fix(security): require authentication for image file endpoint (H-2)

- Changed /api/v1/files/image from AuthType.NONE to USER_TOKEN_OR_API_KEY
- Added Permission.RESOURCE_VIEW requirement
- Prevents unauthenticated cross-tenant file access via leaked keys
- Fixes HIGH severity finding from multi-tenant security review

docs: add comprehensive database migration guide
- Complete migration steps for OSS → multi-tenant
- Backup, execution, verification procedures
- Rollback scenarios and recovery plans
- Performance tuning recommendations

* test: add comprehensive cross-tenant isolation tests

Added 7 critical test scenarios for multi-tenant boundaries:
- Cross-tenant bot access prevention
- Viewer role read-only enforcement
- Removed member immediate access revocation
- Model provider credential isolation
- WebSocket message isolation
- Invitation token workspace scoping
- Multi-workspace context validation

These tests address P0-2 coverage gaps for:
- workspaces.py (membership & invitation flows)
- user.py (authentication & authorization)
- websocket_chat.py (real-time isolation)
- plugins.py (resource access control)

docs: finalize database migration guide

* fix(security): resolve M-1, M-2, M-3 security findings

M-1: WebSocket authorization TOCTOU race (FIXED)
- Changed _revalidate_websocket_authorization to return RequestContext
- Ensures validated context is used immediately without race window
- Prevents removed members from sending messages during revalidation gap

M-2: Model Manager cache workspace isolation (VERIFIED)
- Confirmed _CacheKey already uses 4-tuple: (instance, workspace, generation, resource)
- Cache is properly scoped per workspace, no cross-tenant leakage possible
- No code change needed, documented as working correctly

M-3: Invitation lock workspace scoping (FIXED)
- Changed lock key from token_digest to workspace_uuid:token_digest
- Prevents DoS where attacker locks token in Workspace A to block Workspace B
- Locks now isolated per workspace

All MEDIUM severity findings from security review now resolved.

* fix(cloud): unblock tenant CI and enforce knowledge quotas

* fix(tenancy): scope rerank model sync

---------

Co-authored-by: dadachann <185672915+dadachann@users.noreply.github.com>
This commit is contained in:
RockChinQ
2026-07-30 21:43:35 +08:00
committed by GitHub
parent 463b120923
commit e1ac5e0fc8
468 changed files with 78320 additions and 13137 deletions
+317 -45
View File
@@ -7,6 +7,7 @@ and error handling without calling real LLM APIs.
from __future__ import annotations
import dataclasses
import pytest
from types import SimpleNamespace
from unittest.mock import AsyncMock, Mock
@@ -16,7 +17,15 @@ 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, WorkspaceInvariantError
from tests.unit_tests.provider.conftest import (
TEST_EXECUTION_CONTEXT,
TEST_WORKSPACE_UUID,
_make_mock_result,
_make_row_mock,
)
# ============================================================================
@@ -63,6 +72,20 @@ async def test_model_manager_skips_space_sync_when_disabled(mock_app_for_modelmg
app.space_service.get_models.assert_not_called()
@pytest.mark.asyncio
async def test_model_manager_skips_legacy_space_sync_in_cloud_runtime(mock_app_for_modelmgr):
"""Cloud startup must not resolve an OSS-local Workspace for legacy model sync."""
app = mock_app_for_modelmgr
app.instance_config.data = {'space': {'disable_models_service': False}}
app.persistence_mgr.mode = SimpleNamespace(value='cloud_runtime')
model_mgr = ModelManager(app)
model_mgr.load_models_from_db = AsyncMock()
await model_mgr.initialize()
app.workspace_service.get_local_execution_binding.assert_not_awaited()
@pytest.mark.asyncio
async def test_sync_new_models_from_space_creates_rerank_models(mock_app_for_modelmgr):
"""Space rerank entries are discovered and persisted under the shared provider."""
@@ -91,9 +114,10 @@ async def test_sync_new_models_from_space_creates_rerank_models(mock_app_for_mod
app.rerank_models_service.get_rerank_models = AsyncMock(return_value=[])
model_mgr = ModelManager(app)
await model_mgr.sync_new_models_from_space()
await model_mgr.sync_new_models_from_space(TEST_EXECUTION_CONTEXT)
app.rerank_models_service.create_rerank_model.assert_awaited_once_with(
TEST_EXECUTION_CONTEXT,
{
'uuid': 'rerank-model-uuid',
'name': 'Qwen3-Reranker-8B',
@@ -134,13 +158,33 @@ 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
async def test_empty_cloud_workspace_does_not_retain_generation(
mock_app_for_modelmgr,
):
model_mgr = ModelManager(mock_app_for_modelmgr)
await model_mgr._load_workspace_models(TEST_EXECUTION_CONTEXT)
assert model_mgr.provider_dict == {}
assert model_mgr.llm_model_dict == {}
assert model_mgr.embedding_model_dict == {}
assert model_mgr.rerank_model_dict == {}
assert model_mgr._scope_generations == {}
await model_mgr.resolve_execution_context(TEST_EXECUTION_CONTEXT)
assert model_mgr._scope_generations == {}
@pytest.mark.asyncio
@@ -161,7 +205,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'
@@ -180,7 +224,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'
@@ -197,7 +241,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
@@ -224,7 +268,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'
@@ -237,7 +281,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)
@@ -258,7 +302,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'
@@ -270,7 +317,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
@@ -289,7 +339,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'
@@ -301,7 +351,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')
# ============================================================================
@@ -325,12 +375,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
@@ -349,12 +399,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
@@ -373,11 +423,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
@@ -396,11 +446,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
@@ -419,11 +469,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()
)
# ============================================================================
@@ -541,7 +597,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'
@@ -571,7 +627,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'
@@ -596,7 +655,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'
@@ -632,12 +694,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',
@@ -648,9 +714,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
@@ -667,7 +736,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'
@@ -686,7 +755,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
@@ -702,7 +775,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
@@ -716,7 +793,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
@@ -735,6 +816,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,
@@ -742,7 +824,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
@@ -766,6 +852,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',
@@ -779,7 +866,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()
@@ -793,6 +880,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',
@@ -802,6 +890,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',
@@ -816,7 +905,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()
@@ -833,6 +922,189 @@ 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_generation_advance_prunes_superseded_model_runtime_objects():
class NoGlobalIterationDict(dict):
def __iter__(self):
raise AssertionError('generation advance scanned every model runtime')
def items(self):
raise AssertionError('generation advance scanned every model runtime')
def keys(self):
raise AssertionError('generation advance scanned every model runtime')
model_mgr = ModelManager(Mock())
old_context = ExecutionContext(
instance_uuid='instance-a',
workspace_uuid='workspace-a',
placement_generation=1,
)
new_context = dataclasses.replace(old_context, placement_generation=2)
model_mgr._observe_execution_context(old_context)
for cache in (
model_mgr.provider_dict,
model_mgr.llm_model_dict,
model_mgr.embedding_model_dict,
model_mgr.rerank_model_dict,
):
model_mgr._cache_set(
cache,
('instance-a', 'workspace-a', 1, 'resource-a'),
object(),
)
model_mgr._cache_set(
cache,
('instance-a', 'workspace-b', 1, 'resource-b'),
object(),
)
model_mgr.provider_dict = NoGlobalIterationDict(model_mgr.provider_dict)
model_mgr.llm_model_dict = NoGlobalIterationDict(model_mgr.llm_model_dict)
model_mgr.embedding_model_dict = NoGlobalIterationDict(model_mgr.embedding_model_dict)
model_mgr.rerank_model_dict = NoGlobalIterationDict(model_mgr.rerank_model_dict)
model_mgr._observe_execution_context(new_context)
for cache in (
model_mgr.provider_dict,
model_mgr.llm_model_dict,
model_mgr.embedding_model_dict,
model_mgr.rerank_model_dict,
):
assert ('instance-a', 'workspace-a', 1, 'resource-a') not in cache
assert ('instance-a', 'workspace-b', 1, 'resource-b') in cache
with pytest.raises(WorkspaceInvariantError, match='rolled back'):
model_mgr._observe_execution_context(old_context)
@pytest.mark.asyncio
async def test_generation_advance_closes_retired_provider_requester(
fake_requester_registry,
runtime_provider,
):
model_mgr = fake_requester_registry
runtime_provider.requester.aclose = AsyncMock()
await model_mgr.cache_provider(TEST_EXECUTION_CONTEXT, runtime_provider)
next_context = dataclasses.replace(
TEST_EXECUTION_CONTEXT,
placement_generation=2,
)
model_mgr.ap.workspace_service.get_execution_binding = AsyncMock(
return_value=WorkspaceExecutionBinding(
instance_uuid=next_context.instance_uuid,
workspace_uuid=next_context.workspace_uuid,
placement_generation=next_context.placement_generation,
write_fenced=False,
state='active',
)
)
await model_mgr.resolve_execution_context(next_context)
runtime_provider.requester.aclose.assert_awaited_once_with()
assert model_mgr.provider_dict == {}
assert model_mgr._scope_generations == {}
@pytest.mark.asyncio
async def test_model_manager_shutdown_closes_all_requesters_once(
fake_requester_registry,
runtime_provider,
):
model_mgr = fake_requester_registry
runtime_provider.requester.aclose = AsyncMock()
await model_mgr.cache_provider(TEST_EXECUTION_CONTEXT, runtime_provider)
await model_mgr.shutdown()
await model_mgr.shutdown()
runtime_provider.requester.aclose.assert_awaited_once_with()
assert model_mgr.provider_dict == {}
assert model_mgr.llm_model_dict == {}
assert model_mgr.embedding_model_dict == {}
assert model_mgr.rerank_model_dict == {}
def test_provider_not_found_error_str():
"""Test ProviderNotFoundError string representation."""
error = provider_errors.ProviderNotFoundError('test-provider')