mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-07-21 20:06:06 +00:00
feat(tenancy): implement workspace isolation
This commit is contained in:
@@ -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,
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user