mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-08-09 04:40:57 +00:00
feat(tenancy): add Workspace multi-tenant foundation (#2353)
* Document multi-tenant workspace architecture * Add OSS and commercial workspace boundaries * docs: redesign multi-tenant workspace architecture * feat(tenancy): implement workspace isolation * docs(tenancy): record verification evidence * docs(tenancy): revise single-instance SaaS topology * docs(tenancy): refine architecture options * docs: finalize cloud v2 multi-tenant decisions * feat(tenancy): establish cloud isolation foundations * feat(tenancy): harden shared cloud runtime boundaries * docs(tenancy): record final isolation verification * fix(tenancy): close isolation and permission gaps * docs(tenancy): record final isolation verification * feat(tenancy): connect cloud workspace control plane * fix(build): install git for pinned SDK * docs(cloud): update control plane verification * chore: update multi-tenant SDK pin * fix(cloud): skip legacy model sync during startup * test(cloud): preserve minimal model manager fixtures * fix(cloud): preserve authenticated account context * fix(cloud): reuse authenticated account for user info * feat(cloud): complete Workspace settings navigation * test(web): cover Workspace dropdown menu * feat(web): place workspace controls in sidebar * refactor(web): streamline workspace controls * style(web): format workspace layout test * fix(cloud): surface runtime and workspace plan status * fix(plugin): keep runtime identity stable across restarts * fix(ui): widen and center workspace switcher * fix(ui): hide roles from workspace switcher * fix(ui): align workspace switcher with sidebar entries * feat(workspace): add in-product collaboration and direct Cloud launch * style: format collaboration changes * fix(workspace): bind collaboration APIs to tenant UoW * fix(cloud): preserve Core-owned collaboration state * test(cloud): require Space identity for invite registration * feat(cloud): complete secure invitation experience * style(web): format invitation flows * fix(cloud): recover box runtime without unscoped skill reload * feat(oss): enforce invitation account and owner billing flows * style: format OSS account service * test(oss): cover invitation logout handoff * fix(oss): resolve workspace owner in scoped session * feat(cloud): harden multi-tenant runtime resources * fix(cloud): bound runtime restart storms * fix(cloud): eliminate periodic runtime CPU spikes * fix(cloud): enforce instance capacity ceilings * fix(cloud): scope public login capability discovery * fix(cloud): bound tenant maintenance and monitoring work * fix(runtime): bound tenant resource amplification * fix(deps): pin green multi-tenant plugin SDK * fix(cloud): handle unavailable skill capability * fix(security): require authentication for image file endpoint (H-2) - Changed /api/v1/files/image from AuthType.NONE to USER_TOKEN_OR_API_KEY - Added Permission.RESOURCE_VIEW requirement - Prevents unauthenticated cross-tenant file access via leaked keys - Fixes HIGH severity finding from multi-tenant security review docs: add comprehensive database migration guide - Complete migration steps for OSS → multi-tenant - Backup, execution, verification procedures - Rollback scenarios and recovery plans - Performance tuning recommendations * test: add comprehensive cross-tenant isolation tests Added 7 critical test scenarios for multi-tenant boundaries: - Cross-tenant bot access prevention - Viewer role read-only enforcement - Removed member immediate access revocation - Model provider credential isolation - WebSocket message isolation - Invitation token workspace scoping - Multi-workspace context validation These tests address P0-2 coverage gaps for: - workspaces.py (membership & invitation flows) - user.py (authentication & authorization) - websocket_chat.py (real-time isolation) - plugins.py (resource access control) docs: finalize database migration guide * fix(security): resolve M-1, M-2, M-3 security findings M-1: WebSocket authorization TOCTOU race (FIXED) - Changed _revalidate_websocket_authorization to return RequestContext - Ensures validated context is used immediately without race window - Prevents removed members from sending messages during revalidation gap M-2: Model Manager cache workspace isolation (VERIFIED) - Confirmed _CacheKey already uses 4-tuple: (instance, workspace, generation, resource) - Cache is properly scoped per workspace, no cross-tenant leakage possible - No code change needed, documented as working correctly M-3: Invitation lock workspace scoping (FIXED) - Changed lock key from token_digest to workspace_uuid:token_digest - Prevents DoS where attacker locks token in Workspace A to block Workspace B - Locks now isolated per workspace All MEDIUM severity findings from security review now resolved. * fix(cloud): unblock tenant CI and enforce knowledge quotas * fix(tenancy): scope rerank model sync --------- Co-authored-by: dadachann <185672915+dadachann@users.noreply.github.com>
This commit is contained in:
@@ -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')
|
||||
|
||||
Reference in New Issue
Block a user