mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-08-08 20:30:59 +00:00
feat(cloud): provision workspace model catalog (#2376)
* feat(cloud): provision workspace model catalog * ci(cloud): pin model catalog adapter source * fix: make cloud model catalog sync recoverable * ci: pin cloud adapter source for release --------- Co-authored-by: dadachann <185672915+dadachann@users.noreply.github.com>
This commit is contained in:
@@ -5,6 +5,7 @@ import uuid
|
||||
import sqlalchemy
|
||||
from langbot_plugin.api.entities.builtin.provider import message as provider_message
|
||||
|
||||
from ....cloud.model_catalog import LANGBOT_MODELS_PROVIDER_REQUESTER
|
||||
from ....core import app
|
||||
from ....entity.persistence import model as persistence_model
|
||||
from ....entity.persistence import pipeline as persistence_pipeline
|
||||
@@ -113,6 +114,23 @@ async def _require_workspace_provider(
|
||||
return provider
|
||||
|
||||
|
||||
def _is_cloud_runtime(ap: app.Application) -> bool:
|
||||
mode = getattr(ap.persistence_mgr, 'mode', None)
|
||||
return getattr(mode, 'value', None) == 'cloud_runtime'
|
||||
|
||||
|
||||
async def _assert_cloud_managed_provider_mutable(
|
||||
ap: app.Application,
|
||||
context: TenantContext,
|
||||
provider_uuid: str,
|
||||
) -> None:
|
||||
if not _is_cloud_runtime(ap):
|
||||
return
|
||||
provider = await _require_workspace_provider(ap, context, provider_uuid)
|
||||
if provider.get('requester') == LANGBOT_MODELS_PROVIDER_REQUESTER:
|
||||
raise ValueError('LangBot Models is managed by Cloud and cannot be modified')
|
||||
|
||||
|
||||
async def _require_runtime_provider(
|
||||
ap: app.Application,
|
||||
context: TenantContext,
|
||||
@@ -213,6 +231,7 @@ class LLMModelsService:
|
||||
model_data['provider_uuid'] = provider_uuid
|
||||
|
||||
await _require_workspace_provider(self.ap, context, model_data['provider_uuid'])
|
||||
await _assert_cloud_managed_provider_mutable(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))
|
||||
@@ -291,11 +310,17 @@ class LLMModelsService:
|
||||
|
||||
return model_dict
|
||||
|
||||
async def update_llm_model(self, context: TenantContext, 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"""
|
||||
existing_model = await self.get_llm_model(context, model_uuid, include_secret=True)
|
||||
if existing_model is None:
|
||||
raise WorkspaceNotFoundError('Model not found')
|
||||
await _assert_cloud_managed_provider_mutable(self.ap, context, existing_model['provider_uuid'])
|
||||
model_data = model_data.copy()
|
||||
model_data.pop('uuid', None)
|
||||
model_data.pop('workspace_uuid', None)
|
||||
@@ -321,6 +346,7 @@ class LLMModelsService:
|
||||
|
||||
provider_uuid = model_data.get('provider_uuid', existing_model['provider_uuid'])
|
||||
await _require_workspace_provider(self.ap, context, provider_uuid)
|
||||
await _assert_cloud_managed_provider_mutable(self.ap, context, provider_uuid)
|
||||
await _validate_provider_supports(self.ap, context, provider_uuid, 'llm')
|
||||
|
||||
result = await self.ap.persistence_mgr.execute_async(
|
||||
@@ -355,6 +381,11 @@ class LLMModelsService:
|
||||
|
||||
async def delete_llm_model(self, context: TenantContext, model_uuid: str) -> None:
|
||||
"""Delete an LLM model"""
|
||||
if _is_cloud_runtime(self.ap):
|
||||
existing_model = await self.get_llm_model(context, model_uuid, include_secret=True)
|
||||
if existing_model is None:
|
||||
raise WorkspaceNotFoundError('Model not found')
|
||||
await _assert_cloud_managed_provider_mutable(self.ap, context, existing_model['provider_uuid'])
|
||||
result = await self.ap.persistence_mgr.execute_async(
|
||||
scope_statement(
|
||||
sqlalchemy.delete(persistence_model.LLMModel).where(persistence_model.LLMModel.uuid == model_uuid),
|
||||
@@ -448,7 +479,10 @@ class EmbeddingModelsService:
|
||||
return serialized if include_secret else [_redact_model_secrets(model) for model in serialized]
|
||||
|
||||
async def create_embedding_model(
|
||||
self, context: TenantContext, model_data: dict, preserve_uuid: bool = False
|
||||
self,
|
||||
context: TenantContext,
|
||||
model_data: dict,
|
||||
preserve_uuid: bool = False,
|
||||
) -> str:
|
||||
"""Create a new embedding model"""
|
||||
model_data = model_data.copy()
|
||||
@@ -472,6 +506,7 @@ class EmbeddingModelsService:
|
||||
model_data['provider_uuid'] = provider_uuid
|
||||
|
||||
await _require_workspace_provider(self.ap, context, model_data['provider_uuid'])
|
||||
await _assert_cloud_managed_provider_mutable(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(
|
||||
@@ -530,11 +565,17 @@ class EmbeddingModelsService:
|
||||
|
||||
return model_dict
|
||||
|
||||
async def update_embedding_model(self, context: TenantContext, 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"""
|
||||
existing_model = await self.get_embedding_model(context, model_uuid, include_secret=True)
|
||||
if existing_model is None:
|
||||
raise WorkspaceNotFoundError('Model not found')
|
||||
await _assert_cloud_managed_provider_mutable(self.ap, context, existing_model['provider_uuid'])
|
||||
model_data = model_data.copy()
|
||||
model_data.pop('uuid', None)
|
||||
model_data.pop('workspace_uuid', None)
|
||||
@@ -559,6 +600,7 @@ class EmbeddingModelsService:
|
||||
|
||||
provider_uuid = model_data.get('provider_uuid', existing_model['provider_uuid'])
|
||||
await _require_workspace_provider(self.ap, context, provider_uuid)
|
||||
await _assert_cloud_managed_provider_mutable(self.ap, context, provider_uuid)
|
||||
await _validate_provider_supports(self.ap, context, provider_uuid, 'text-embedding')
|
||||
|
||||
result = await self.ap.persistence_mgr.execute_async(
|
||||
@@ -593,6 +635,11 @@ class EmbeddingModelsService:
|
||||
|
||||
async def delete_embedding_model(self, context: TenantContext, model_uuid: str) -> None:
|
||||
"""Delete an embedding model"""
|
||||
if _is_cloud_runtime(self.ap):
|
||||
existing_model = await self.get_embedding_model(context, model_uuid, include_secret=True)
|
||||
if existing_model is None:
|
||||
raise WorkspaceNotFoundError('Model not found')
|
||||
await _assert_cloud_managed_provider_mutable(self.ap, context, existing_model['provider_uuid'])
|
||||
result = await self.ap.persistence_mgr.execute_async(
|
||||
scope_statement(
|
||||
sqlalchemy.delete(persistence_model.EmbeddingModel).where(
|
||||
@@ -685,7 +732,12 @@ class RerankModelsService:
|
||||
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, context: TenantContext, 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:
|
||||
@@ -708,6 +760,7 @@ class RerankModelsService:
|
||||
model_data['provider_uuid'] = provider_uuid
|
||||
|
||||
await _require_workspace_provider(self.ap, context, model_data['provider_uuid'])
|
||||
await _assert_cloud_managed_provider_mutable(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(
|
||||
@@ -766,11 +819,17 @@ class RerankModelsService:
|
||||
|
||||
return model_dict
|
||||
|
||||
async def update_rerank_model(self, context: TenantContext, 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"""
|
||||
existing_model = await self.get_rerank_model(context, model_uuid, include_secret=True)
|
||||
if existing_model is None:
|
||||
raise WorkspaceNotFoundError('Model not found')
|
||||
await _assert_cloud_managed_provider_mutable(self.ap, context, existing_model['provider_uuid'])
|
||||
model_data = model_data.copy()
|
||||
model_data.pop('uuid', None)
|
||||
model_data.pop('workspace_uuid', None)
|
||||
@@ -795,6 +854,7 @@ class RerankModelsService:
|
||||
|
||||
provider_uuid = model_data.get('provider_uuid', existing_model['provider_uuid'])
|
||||
await _require_workspace_provider(self.ap, context, provider_uuid)
|
||||
await _assert_cloud_managed_provider_mutable(self.ap, context, provider_uuid)
|
||||
await _validate_provider_supports(self.ap, context, provider_uuid, 'rerank')
|
||||
|
||||
result = await self.ap.persistence_mgr.execute_async(
|
||||
@@ -829,6 +889,11 @@ class RerankModelsService:
|
||||
|
||||
async def delete_rerank_model(self, context: TenantContext, model_uuid: str) -> None:
|
||||
"""Delete a rerank model"""
|
||||
if _is_cloud_runtime(self.ap):
|
||||
existing_model = await self.get_rerank_model(context, model_uuid, include_secret=True)
|
||||
if existing_model is None:
|
||||
raise WorkspaceNotFoundError('Model not found')
|
||||
await _assert_cloud_managed_provider_mutable(self.ap, context, existing_model['provider_uuid'])
|
||||
result = await self.ap.persistence_mgr.execute_async(
|
||||
scope_statement(
|
||||
sqlalchemy.delete(persistence_model.RerankModel).where(
|
||||
|
||||
@@ -5,6 +5,7 @@ import traceback
|
||||
|
||||
import sqlalchemy
|
||||
|
||||
from ....cloud.model_catalog import LANGBOT_MODELS_PROVIDER_REQUESTER
|
||||
from ....core import app
|
||||
from ....entity.persistence import model as persistence_model
|
||||
from ....workspace.errors import WorkspaceNotFoundError
|
||||
@@ -20,6 +21,20 @@ class ModelProviderService:
|
||||
def __init__(self, ap: app.Application) -> None:
|
||||
self.ap = ap
|
||||
|
||||
def _is_cloud_runtime(self) -> bool:
|
||||
mode = getattr(self.ap.persistence_mgr, 'mode', None)
|
||||
return getattr(mode, 'value', None) == 'cloud_runtime'
|
||||
|
||||
def _system_requester_is_reserved(self, requester: object) -> bool:
|
||||
return self._is_cloud_runtime() and requester == LANGBOT_MODELS_PROVIDER_REQUESTER
|
||||
|
||||
async def _assert_provider_mutable(self, context: TenantContext, provider_uuid: str) -> None:
|
||||
if not self._is_cloud_runtime():
|
||||
return
|
||||
provider = await self.get_provider(context, provider_uuid)
|
||||
if provider is not None and self._system_requester_is_reserved(provider.get('requester')):
|
||||
raise ValueError('LangBot Models is managed by Cloud and cannot be modified')
|
||||
|
||||
@staticmethod
|
||||
def _normalize_api_keys(api_keys: str | list[str] | tuple[str, ...] | None) -> list[str]:
|
||||
if api_keys is None:
|
||||
@@ -99,6 +114,8 @@ class ModelProviderService:
|
||||
async def create_provider(self, context: TenantContext, provider_data: dict) -> str:
|
||||
"""Create a new provider"""
|
||||
provider_data = provider_data.copy()
|
||||
if self._system_requester_is_reserved(provider_data.get('requester')):
|
||||
raise ValueError('space-chat-completions is reserved for the Cloud-managed LangBot Models provider')
|
||||
provider_data['uuid'] = str(uuid.uuid4())
|
||||
provider_data['workspace_uuid'] = require_workspace_uuid(context)
|
||||
provider_data['api_keys'] = self._normalize_api_keys(
|
||||
@@ -115,7 +132,10 @@ class ModelProviderService:
|
||||
|
||||
async def update_provider(self, context: TenantContext, provider_uuid: str, provider_data: dict) -> None:
|
||||
"""Update an existing provider"""
|
||||
await self._assert_provider_mutable(context, provider_uuid)
|
||||
provider_data = provider_data.copy()
|
||||
if self._system_requester_is_reserved(provider_data.get('requester')):
|
||||
raise ValueError('space-chat-completions is reserved for the Cloud-managed LangBot Models provider')
|
||||
provider_data.pop('uuid', None)
|
||||
provider_data.pop('workspace_uuid', None)
|
||||
if 'api_keys' in provider_data:
|
||||
@@ -145,6 +165,7 @@ class ModelProviderService:
|
||||
|
||||
async def delete_provider(self, context: TenantContext, provider_uuid: str) -> None:
|
||||
"""Delete a provider (only if no models reference it)"""
|
||||
await self._assert_provider_mutable(context, provider_uuid)
|
||||
workspace_uuid = require_workspace_uuid(context)
|
||||
# Check if any models use this provider
|
||||
llm_result = await self.ap.persistence_mgr.execute_async(
|
||||
@@ -245,6 +266,8 @@ class ModelProviderService:
|
||||
api_keys: list,
|
||||
) -> str:
|
||||
"""Find existing provider or create new one"""
|
||||
if self._system_requester_is_reserved(requester):
|
||||
raise ValueError('space-chat-completions is reserved for the Cloud-managed LangBot Models provider')
|
||||
workspace_uuid = require_workspace_uuid(context)
|
||||
api_keys = self._normalize_api_keys(restore_secret_placeholders(api_keys, sensitive=True))
|
||||
|
||||
|
||||
@@ -13,6 +13,7 @@ from typing import Any, Protocol, runtime_checkable
|
||||
from ..workspace.policy import CloudWorkspacePolicy, SingleWorkspacePolicy
|
||||
from .directory import DirectoryProjectionProvider, directory_projection_limits_from_config
|
||||
from .entitlements import EntitlementProvider, OpenSourceEntitlementProvider
|
||||
from .model_catalog import CloudModelCatalogProvider
|
||||
|
||||
|
||||
CLOUD_BOOTSTRAP_ENTRY_POINT = 'langbot.cloud_bootstrap'
|
||||
@@ -50,6 +51,7 @@ class OpenSourceDeployment:
|
||||
)
|
||||
directory_provider: None = None
|
||||
manifest_provider: None = None
|
||||
model_catalog_provider: None = None
|
||||
persistence_mode: str = 'oss_compat'
|
||||
required_vector_backend: str | None = None
|
||||
|
||||
@@ -80,6 +82,7 @@ class VerifiedCloudDeployment:
|
||||
entitlement_provider: EntitlementProvider
|
||||
directory_provider: DirectoryProjectionProvider
|
||||
manifest_provider: CloudManifestProvider
|
||||
model_catalog_provider: CloudModelCatalogProvider
|
||||
verification_key_id: str
|
||||
mode: str = dataclasses.field(default='cloud', init=False)
|
||||
workspace_policy: CloudWorkspacePolicy = dataclasses.field(default_factory=CloudWorkspacePolicy, init=False)
|
||||
@@ -110,6 +113,8 @@ class VerifiedCloudDeployment:
|
||||
raise CloudBootstrapError('Verified Cloud bootstrap did not provide a directory adapter')
|
||||
if not isinstance(self.manifest_provider, CloudManifestProvider):
|
||||
raise CloudBootstrapError('Verified Cloud bootstrap did not provide a Manifest renewal adapter')
|
||||
if not isinstance(self.model_catalog_provider, CloudModelCatalogProvider):
|
||||
raise CloudBootstrapError('Verified Cloud bootstrap did not provide a model catalog adapter')
|
||||
|
||||
def validate_instance_config(self, config: dict[str, Any]) -> None:
|
||||
try:
|
||||
|
||||
@@ -0,0 +1,315 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import uuid
|
||||
from datetime import datetime
|
||||
from typing import Any, Literal, Protocol, runtime_checkable
|
||||
|
||||
import sqlalchemy
|
||||
from pydantic import BaseModel, ConfigDict, Field, SecretStr, field_validator
|
||||
|
||||
from ..entity.persistence import model as persistence_model
|
||||
|
||||
|
||||
LANGBOT_MODELS_PROVIDER_REQUESTER = 'space-chat-completions'
|
||||
LANGBOT_MODELS_PROVIDER_NAME = 'LangBot Models'
|
||||
_MODEL_RESOURCE_NAMESPACE = uuid.UUID('94c703ca-1df5-4e91-bcd3-74ac65cb7921')
|
||||
_SUPPORTED_CATEGORIES = {'chat', 'embedding', 'rerank'}
|
||||
_MODEL_TABLES = (
|
||||
persistence_model.LLMModel,
|
||||
persistence_model.EmbeddingModel,
|
||||
persistence_model.RerankModel,
|
||||
)
|
||||
|
||||
|
||||
class CloudModelCatalogItem(BaseModel):
|
||||
model_config = ConfigDict(extra='forbid', frozen=True)
|
||||
|
||||
uuid: str = Field(min_length=1, max_length=255)
|
||||
model_id: str = Field(min_length=1, max_length=255)
|
||||
category: Literal['chat', 'embedding', 'rerank']
|
||||
llm_abilities: tuple[str, ...] = ()
|
||||
is_featured: bool = False
|
||||
featured_order: int = 0
|
||||
|
||||
@field_validator('llm_abilities')
|
||||
@classmethod
|
||||
def validate_abilities(cls, value: tuple[str, ...]) -> tuple[str, ...]:
|
||||
if any(not item.strip() or len(item) > 64 for item in value):
|
||||
raise ValueError('Model abilities must be non-empty strings of at most 64 characters')
|
||||
if len(set(value)) != len(value):
|
||||
raise ValueError('Model abilities must be unique')
|
||||
return value
|
||||
|
||||
|
||||
class CloudWorkspaceModelBilling(BaseModel):
|
||||
model_config = ConfigDict(extra='forbid', frozen=True)
|
||||
|
||||
workspace_uuid: str = Field(min_length=36, max_length=36)
|
||||
owner_account_uuid: str | None = Field(default=None, min_length=36, max_length=36)
|
||||
api_key: SecretStr | None = None
|
||||
|
||||
@field_validator('workspace_uuid')
|
||||
@classmethod
|
||||
def validate_uuid(cls, value: str) -> str:
|
||||
return str(uuid.UUID(value))
|
||||
|
||||
@field_validator('owner_account_uuid')
|
||||
@classmethod
|
||||
def validate_optional_uuid(cls, value: str | None) -> str | None:
|
||||
return None if value is None else str(uuid.UUID(value))
|
||||
|
||||
|
||||
class CloudModelCatalogSnapshot(BaseModel):
|
||||
model_config = ConfigDict(extra='forbid', frozen=True)
|
||||
|
||||
instance_uuid: str = Field(min_length=1, max_length=255)
|
||||
generated_at: datetime
|
||||
base_url: str = Field(min_length=1, max_length=512)
|
||||
models: tuple[CloudModelCatalogItem, ...]
|
||||
workspaces: tuple[CloudWorkspaceModelBilling, ...]
|
||||
|
||||
@field_validator('base_url')
|
||||
@classmethod
|
||||
def validate_base_url(cls, value: str) -> str:
|
||||
normalized = value.rstrip('/')
|
||||
if not normalized.startswith('https://'):
|
||||
raise ValueError('Cloud model gateway base URL must use HTTPS')
|
||||
return normalized
|
||||
|
||||
@field_validator('models')
|
||||
@classmethod
|
||||
def validate_models(cls, value: tuple[CloudModelCatalogItem, ...]) -> tuple[CloudModelCatalogItem, ...]:
|
||||
if len(value) > 500:
|
||||
raise ValueError('Cloud model catalog exceeds 500 models')
|
||||
identities = {(item.category, item.uuid) for item in value}
|
||||
if len(identities) != len(value):
|
||||
raise ValueError('Cloud model catalog contains duplicate model identities')
|
||||
return value
|
||||
|
||||
@field_validator('workspaces')
|
||||
@classmethod
|
||||
def validate_workspaces(
|
||||
cls, value: tuple[CloudWorkspaceModelBilling, ...]
|
||||
) -> tuple[CloudWorkspaceModelBilling, ...]:
|
||||
if len(value) > 10_000:
|
||||
raise ValueError('Cloud model catalog exceeds 10000 Workspaces')
|
||||
identities = {item.workspace_uuid for item in value}
|
||||
if len(identities) != len(value):
|
||||
raise ValueError('Cloud model catalog contains duplicate Workspaces')
|
||||
return value
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class CloudModelCatalogProvider(Protocol):
|
||||
async def fetch_model_catalog(self, instance_uuid: str) -> CloudModelCatalogSnapshot:
|
||||
"""Fetch and verify the complete model catalog and Workspace billing projection."""
|
||||
...
|
||||
|
||||
|
||||
def system_provider_uuid(workspace_uuid: str) -> str:
|
||||
workspace = str(uuid.UUID(workspace_uuid))
|
||||
return str(uuid.uuid5(_MODEL_RESOURCE_NAMESPACE, f'{workspace}:provider:{LANGBOT_MODELS_PROVIDER_REQUESTER}'))
|
||||
|
||||
|
||||
def system_model_uuid(workspace_uuid: str, category: str, upstream_uuid: str) -> str:
|
||||
workspace = str(uuid.UUID(workspace_uuid))
|
||||
if category not in _SUPPORTED_CATEGORIES:
|
||||
raise ValueError(f'Unsupported model category: {category}')
|
||||
if not upstream_uuid:
|
||||
raise ValueError('Upstream model UUID is required')
|
||||
return str(uuid.uuid5(_MODEL_RESOURCE_NAMESPACE, f'{workspace}:model:{category}:{upstream_uuid}'))
|
||||
|
||||
|
||||
class CloudModelCatalogSyncService:
|
||||
"""Reconcile Space-owned model catalog and Owner billing tokens into every Cloud Workspace."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
ap: Any,
|
||||
provider: CloudModelCatalogProvider,
|
||||
instance_uuid: str,
|
||||
*,
|
||||
sync_interval_seconds: float = 3600.0,
|
||||
) -> None:
|
||||
if not isinstance(provider, CloudModelCatalogProvider):
|
||||
raise TypeError('Cloud model catalog sync requires a CloudModelCatalogProvider')
|
||||
if sync_interval_seconds < 10:
|
||||
raise ValueError('Cloud model catalog sync interval must be at least 10 seconds')
|
||||
self.ap = ap
|
||||
self.provider = provider
|
||||
self.instance_uuid = instance_uuid
|
||||
self.sync_interval_seconds = float(sync_interval_seconds)
|
||||
# A tenant UoW commits one Workspace at a time. Keep a durable in-memory
|
||||
# convergence marker so a failed runtime reload is retried even when the
|
||||
# following database reconciliation is a no-op.
|
||||
self._runtime_reload_pending = False
|
||||
|
||||
async def initialize(self) -> None:
|
||||
await self.sync_once(reload_runtime=False)
|
||||
|
||||
async def run(self) -> None:
|
||||
while True:
|
||||
await asyncio.sleep(self.sync_interval_seconds)
|
||||
try:
|
||||
await self.sync_once(reload_runtime=True)
|
||||
except asyncio.CancelledError:
|
||||
raise
|
||||
except Exception as exc:
|
||||
# Exception messages can contain rendered SQL bound values,
|
||||
# including provider API keys. Log only the exception class.
|
||||
self.ap.logger.warning(f'Cloud model catalog synchronization failed ({type(exc).__name__})')
|
||||
|
||||
async def sync_once(self, *, reload_runtime: bool = True) -> dict[str, int]:
|
||||
summary = {'workspaces': 0, 'created': 0, 'updated': 0, 'deleted': 0}
|
||||
snapshot: CloudModelCatalogSnapshot | None = None
|
||||
sync_error: Exception | None = None
|
||||
reload_error: Exception | None = None
|
||||
try:
|
||||
snapshot = await self.provider.fetch_model_catalog(self.instance_uuid)
|
||||
if snapshot.instance_uuid != self.instance_uuid:
|
||||
raise ValueError('Cloud model catalog targets another LangBot instance')
|
||||
|
||||
bindings = await self.ap.workspace_service.list_active_execution_bindings()
|
||||
billing_by_workspace = {item.workspace_uuid: item for item in snapshot.workspaces}
|
||||
missing = sorted(
|
||||
binding.workspace_uuid for binding in bindings if binding.workspace_uuid not in billing_by_workspace
|
||||
)
|
||||
if missing:
|
||||
raise ValueError(
|
||||
f'Cloud model catalog is missing billing projections for {len(missing)} active Workspaces'
|
||||
)
|
||||
|
||||
for binding in bindings:
|
||||
counts = await self._sync_workspace(
|
||||
binding.workspace_uuid,
|
||||
snapshot,
|
||||
billing_by_workspace[binding.workspace_uuid],
|
||||
)
|
||||
summary['workspaces'] += 1
|
||||
workspace_changed = any(counts[key] > 0 for key in ('created', 'updated', 'deleted'))
|
||||
if workspace_changed:
|
||||
# _sync_workspace returns only after its tenant UoW commits.
|
||||
self._runtime_reload_pending = True
|
||||
for key in ('created', 'updated', 'deleted'):
|
||||
summary[key] += counts[key]
|
||||
except Exception as exc:
|
||||
sync_error = exc
|
||||
finally:
|
||||
model_mgr = getattr(self.ap, 'model_mgr', None)
|
||||
if reload_runtime and self._runtime_reload_pending and model_mgr is not None:
|
||||
try:
|
||||
await model_mgr.load_models_from_db()
|
||||
except Exception as exc:
|
||||
reload_error = exc
|
||||
else:
|
||||
self._runtime_reload_pending = False
|
||||
|
||||
if sync_error is not None:
|
||||
if reload_error is not None:
|
||||
raise sync_error from reload_error
|
||||
raise sync_error
|
||||
if reload_error is not None:
|
||||
raise reload_error
|
||||
|
||||
changed = any(summary[key] > 0 for key in ('created', 'updated', 'deleted'))
|
||||
if changed and snapshot is not None:
|
||||
self.ap.logger.info(
|
||||
'Cloud model catalog synchronized '
|
||||
f'({summary["workspaces"]} Workspaces, {len(snapshot.models)} models, '
|
||||
f'created={summary["created"]}, updated={summary["updated"]}, deleted={summary["deleted"]})'
|
||||
)
|
||||
return summary
|
||||
|
||||
async def _sync_workspace(
|
||||
self,
|
||||
workspace_uuid: str,
|
||||
snapshot: CloudModelCatalogSnapshot,
|
||||
billing: CloudWorkspaceModelBilling,
|
||||
) -> dict[str, int]:
|
||||
counts = {'created': 0, 'updated': 0, 'deleted': 0}
|
||||
provider_uuid = system_provider_uuid(workspace_uuid)
|
||||
desired_keys = [billing.api_key.get_secret_value()] if billing.api_key is not None else []
|
||||
|
||||
async with self.ap.persistence_mgr.tenant_uow(workspace_uuid) as uow:
|
||||
provider = await uow.session.scalar(
|
||||
sqlalchemy.select(persistence_model.ModelProvider).where(
|
||||
persistence_model.ModelProvider.uuid == provider_uuid
|
||||
)
|
||||
)
|
||||
provider_values = {
|
||||
'workspace_uuid': workspace_uuid,
|
||||
'name': LANGBOT_MODELS_PROVIDER_NAME,
|
||||
'requester': LANGBOT_MODELS_PROVIDER_REQUESTER,
|
||||
'base_url': snapshot.base_url,
|
||||
'api_keys': desired_keys,
|
||||
}
|
||||
if provider is None:
|
||||
provider = persistence_model.ModelProvider(uuid=provider_uuid, **provider_values)
|
||||
uow.session.add(provider)
|
||||
await uow.session.flush()
|
||||
counts['created'] += 1
|
||||
elif self._update_entity(provider, provider_values):
|
||||
counts['updated'] += 1
|
||||
|
||||
existing_by_table: dict[type, dict[str, Any]] = {}
|
||||
for table in _MODEL_TABLES:
|
||||
rows = (
|
||||
await uow.session.scalars(sqlalchemy.select(table).where(table.provider_uuid == provider_uuid))
|
||||
).all()
|
||||
existing_by_table[table] = {row.uuid: row for row in rows}
|
||||
|
||||
desired_ids: dict[type, set[str]] = {table: set() for table in _MODEL_TABLES}
|
||||
for item in snapshot.models:
|
||||
table, values = self._model_values(workspace_uuid, provider_uuid, item)
|
||||
model_uuid = system_model_uuid(workspace_uuid, item.category, item.uuid)
|
||||
desired_ids[table].add(model_uuid)
|
||||
existing = existing_by_table[table].get(model_uuid)
|
||||
if existing is None:
|
||||
uow.session.add(table(uuid=model_uuid, **values))
|
||||
counts['created'] += 1
|
||||
elif self._update_entity(existing, values):
|
||||
counts['updated'] += 1
|
||||
|
||||
for table, entities in existing_by_table.items():
|
||||
for model_uuid, entity in entities.items():
|
||||
if model_uuid not in desired_ids[table]:
|
||||
await uow.session.delete(entity)
|
||||
counts['deleted'] += 1
|
||||
|
||||
return counts
|
||||
|
||||
@staticmethod
|
||||
def _update_entity(entity: Any, values: dict[str, Any]) -> bool:
|
||||
changed = False
|
||||
for key, value in values.items():
|
||||
if getattr(entity, key) != value:
|
||||
setattr(entity, key, value)
|
||||
changed = True
|
||||
return changed
|
||||
|
||||
@staticmethod
|
||||
def _model_values(
|
||||
workspace_uuid: str,
|
||||
provider_uuid: str,
|
||||
item: CloudModelCatalogItem,
|
||||
) -> tuple[type, dict[str, Any]]:
|
||||
ranking = 100 - item.featured_order if item.is_featured else 0
|
||||
common = {
|
||||
'workspace_uuid': workspace_uuid,
|
||||
'name': item.model_id,
|
||||
'provider_uuid': provider_uuid,
|
||||
'extra_args': {},
|
||||
'prefered_ranking': ranking,
|
||||
}
|
||||
if item.category == 'chat':
|
||||
return persistence_model.LLMModel, {
|
||||
**common,
|
||||
'abilities': list(item.llm_abilities),
|
||||
'context_length': None,
|
||||
}
|
||||
if item.category == 'embedding':
|
||||
return persistence_model.EmbeddingModel, common
|
||||
if item.category == 'rerank':
|
||||
return persistence_model.RerankModel, common
|
||||
raise ValueError(f'Unsupported model category: {item.category}')
|
||||
@@ -54,6 +54,7 @@ from ..cloud import launch as cloud_launch_module
|
||||
from ..cloud import support_admin as cloud_support_admin_module
|
||||
from ..cloud import directory_projection as cloud_directory_projection_module
|
||||
from ..cloud import entitlements as cloud_entitlements_module
|
||||
from ..cloud import model_catalog as cloud_model_catalog_module
|
||||
from ..api.http.context import ExecutionContext, PrincipalContext, PrincipalType
|
||||
|
||||
|
||||
@@ -142,13 +143,12 @@ class Application:
|
||||
deployment: cloud_bootstrap_module.OpenSourceDeployment | cloud_bootstrap_module.VerifiedCloudDeployment = None
|
||||
|
||||
deployment_admission: cloud_bootstrap_module.DeploymentAdmissionGuard = None
|
||||
|
||||
directory_projection_service: cloud_directory_projection_module.DirectoryProjectionService | None = None
|
||||
cloud_model_catalog_service: cloud_model_catalog_module.CloudModelCatalogSyncService | None = None
|
||||
manifest_refresh_service: cloud_bootstrap_module.CloudManifestRefreshService | None = None
|
||||
|
||||
entitlement_resolver: cloud_entitlements_module.EntitlementResolver | None = None
|
||||
|
||||
directory_projection_service: cloud_directory_projection_module.DirectoryProjectionService | None = None
|
||||
|
||||
vector_db_mgr: vectordb_mgr.VectorDBManager = None
|
||||
|
||||
http_ctrl: http_controller.HTTPController = None
|
||||
@@ -306,6 +306,12 @@ class Application:
|
||||
name='cloud-directory-projection',
|
||||
scopes=[core_entities.LifecycleControlScope.APPLICATION],
|
||||
)
|
||||
if self.cloud_model_catalog_service is not None:
|
||||
self.task_mgr.create_task(
|
||||
self.cloud_model_catalog_service.run(),
|
||||
name='cloud-model-catalog-sync',
|
||||
scopes=[core_entities.LifecycleControlScope.APPLICATION],
|
||||
)
|
||||
if self.manifest_refresh_service is not None:
|
||||
self.task_mgr.create_task(
|
||||
self.manifest_refresh_service.run(),
|
||||
|
||||
@@ -46,6 +46,7 @@ from ...cloud import support_admin as cloud_support_admin_module
|
||||
from ...cloud.directory import directory_projection_limits_from_config
|
||||
from ...cloud.directory_projection import DirectoryProjectionService
|
||||
from ...cloud.entitlements import EntitlementResolver
|
||||
from ...cloud.model_catalog import CloudModelCatalogSyncService
|
||||
from ...api.http.context import ExecutionContext, PrincipalContext, PrincipalType
|
||||
from ...api.http.authz import WorkspaceRequiredError
|
||||
|
||||
@@ -176,6 +177,16 @@ class BuildAppStage(stage.BootingStage):
|
||||
# of repeating tenant validation for every manager.
|
||||
await workspace_service_inst.prime_startup_execution_bindings()
|
||||
|
||||
if not isinstance(deployment, cloud_bootstrap.VerifiedCloudDeployment):
|
||||
raise RuntimeError('Multi-Workspace runtime requires a verified Cloud deployment')
|
||||
cloud_model_catalog_service = CloudModelCatalogSyncService(
|
||||
ap,
|
||||
deployment.model_catalog_provider,
|
||||
constants.instance_id,
|
||||
)
|
||||
await cloud_model_catalog_service.initialize()
|
||||
ap.cloud_model_catalog_service = cloud_model_catalog_service
|
||||
|
||||
ap.workspace_collaboration_service = workspace_collaboration_module.WorkspaceCollaborationService(
|
||||
ap,
|
||||
workspace_service_inst,
|
||||
|
||||
Reference in New Issue
Block a user